build_gallery.py 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149
  1. """
  2. 构建图库(离线,一次性)
  3. 流程:读 cards_master_v2 → 下载卡牌标准图 → YOLO裁切 → DINOv2提取特征 → 保存索引
  4. 用法:
  5. python build_gallery.py # 全量构建(78272张)
  6. python build_gallery.py --limit 100 # 只构建100张(测试)
  7. python build_gallery.py --language 简中 # 只构建简中卡
  8. """
  9. import os
  10. import sys
  11. import time
  12. import json
  13. import argparse
  14. sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
  15. sys.stdout.reconfigure(encoding="utf-8")
  16. import config
  17. from modules import db_reader, image_downloader, csv_reader
  18. from modules.yolo_detector import CardDetector
  19. from modules.feature_extractor import CardFeatureExtractor
  20. from modules.vector_index import VectorIndex
  21. import numpy as np
  22. def build_yolo_cropped_cache(detector, img_path):
  23. """YOLO裁切后缓存到同目录 _crop.jpg,避免重复裁切"""
  24. crop_path = img_path.replace(".jpg", "_crop.jpg")
  25. if os.path.exists(crop_path):
  26. return crop_path
  27. crop = detector.detect_and_crop(img_path)
  28. if crop is None:
  29. return None
  30. import cv2
  31. cv2.imwrite(crop_path, cv2.cvtColor(crop, cv2.COLOR_RGB2BGR))
  32. return crop_path
  33. def main():
  34. parser = argparse.ArgumentParser()
  35. parser.add_argument("--limit", type=int, default=None, help="限制条数(测试用)")
  36. parser.add_argument("--language", type=str, default=None, help="按语言过滤,如 简中")
  37. parser.add_argument("--year", type=str, default=None, help="按年份过滤,如 1996")
  38. parser.add_argument("--no-yolo", action="store_true", help="跳过YOLO裁切(图库已是干净卡面时用)")
  39. parser.add_argument("--rebuild-features", action="store_true", help="强制重新提取特征")
  40. parser.add_argument("--csv", type=str, default=None, help="从CSV读取卡牌数据(服务器无DB时用)")
  41. args = parser.parse_args()
  42. config.ensure_dirs()
  43. # 1. 读卡牌元数据(CSV 优先,否则 PG)
  44. if args.csv or os.path.exists(config.CARD_MASTER_CSV):
  45. csv_path = args.csv or config.CARD_MASTER_CSV
  46. print(f"[1/4] 从 CSV 读取: {csv_path}")
  47. cards = csv_reader.read_card_master_csv(csv_path)
  48. # 应用 limit
  49. if args.limit:
  50. cards = cards[:args.limit]
  51. total = len(cards)
  52. print(f" 共 {total} 张卡牌")
  53. else:
  54. # 构建过滤条件
  55. filters = {}
  56. if args.language:
  57. filters["language"] = [args.language]
  58. if args.year:
  59. filters["year"] = [args.year]
  60. print(f"[1/4] 读取 cards_master_v2 ... 过滤: {filters or '无'}")
  61. cards = db_reader.read_card_master(config.PG_CONFIG, config.PG_TABLE, filters, args.limit)
  62. total = len(cards)
  63. print(f" 共 {total} 张卡牌")
  64. if total == 0:
  65. print("无数据,退出")
  66. return
  67. # 2. 下载图片
  68. print(f"[2/4] 下载卡牌标准图到 {config.GALLERY_IMG_DIR} ...")
  69. t0 = time.time()
  70. url_to_path = image_downloader.download_card_images_batch(
  71. [c["img_url"] for c in cards], config.GALLERY_IMG_DIR, config.DOWNLOAD_WORKERS
  72. )
  73. downloaded = sum(1 for v in url_to_path.values() if v)
  74. print(f" 下载完成 {downloaded}/{total},耗时 {time.time()-t0:.1f}s")
  75. # 3. YOLO裁切 + 特征提取
  76. print(f"[3/4] YOLO裁切 + DINOv2特征提取 ...")
  77. detector = None if args.no_yolo else CardDetector(config.YOLO_MODEL_PATH, conf=config.YOLO_CONF_THRESHOLD)
  78. extractor = CardFeatureExtractor(config.DINOV2_MODEL_PATH)
  79. print(f" 设备: {extractor.device}, 特征维度: {extractor.feature_dim}")
  80. all_feats = []
  81. valid_cards = [] # 成功提取特征的卡牌元数据
  82. batch_imgs, batch_meta = [], []
  83. def flush_batch():
  84. if not batch_imgs:
  85. return
  86. feats = extractor.extract(batch_imgs)
  87. for i, meta in enumerate(batch_meta):
  88. if not np.isnan(feats[i]).any():
  89. all_feats.append(feats[i])
  90. valid_cards.append(meta)
  91. for i, card in enumerate(cards):
  92. path = url_to_path.get(card["img_url"])
  93. if not path:
  94. continue
  95. img = image_downloader.load_image_rgb(path)
  96. if img is None:
  97. continue
  98. if detector is not None:
  99. img = detector.detect_and_crop(img)
  100. if img is None:
  101. continue
  102. batch_imgs.append(img)
  103. batch_meta.append(card)
  104. if len(batch_imgs) >= config.FEATURE_BATCH_SIZE:
  105. flush_batch()
  106. batch_imgs, batch_meta = [], []
  107. print(f" 进度 {i+1}/{total}, 已提取 {len(all_feats)}", end="\r")
  108. flush_batch()
  109. print(f"\n 特征提取完成,有效 {len(all_feats)}/{total}")
  110. if not all_feats:
  111. print("无有效特征,退出")
  112. return
  113. # 4. 构建并保存索引
  114. print(f"[4/4] 保存索引 ...")
  115. features = np.stack(all_feats).astype(np.float32)
  116. metas = [{"card_id": c["card_id"], "card_name_ch": c.get("card_name_ch"),
  117. "language": c.get("language"), "year": c.get("year"),
  118. "card_no": c.get("card_no"), "pg_label": c.get("pg_label")}
  119. for c in valid_cards]
  120. card_ids = [c["card_id"] for c in valid_cards]
  121. index = VectorIndex(features, card_ids, metas)
  122. index.save(config.GALLERY_FEATURES_PATH, config.GALLERY_META_PATH)
  123. print(f" 已保存: {config.GALLERY_FEATURES_PATH}")
  124. print(f" 特征矩阵: {features.shape}")
  125. print(f"\n完成!图库构建成功,共 {len(card_ids)} 张")
  126. if __name__ == "__main__":
  127. main()