""" 构建图库(离线,一次性) 流程:读 cards_master_v2 → 下载卡牌标准图 → YOLO裁切 → DINOv2提取特征 → 保存索引 用法: python build_gallery.py # 全量构建(78272张) python build_gallery.py --limit 100 # 只构建100张(测试) python build_gallery.py --language 简中 # 只构建简中卡 """ import os import sys import time import json import argparse sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) sys.stdout.reconfigure(encoding="utf-8") import config from modules import db_reader, image_downloader, csv_reader from modules.yolo_detector import CardDetector from modules.feature_extractor import CardFeatureExtractor from modules.vector_index import VectorIndex import numpy as np def build_yolo_cropped_cache(detector, img_path): """YOLO裁切后缓存到同目录 _crop.jpg,避免重复裁切""" crop_path = img_path.replace(".jpg", "_crop.jpg") if os.path.exists(crop_path): return crop_path crop = detector.detect_and_crop(img_path) if crop is None: return None import cv2 cv2.imwrite(crop_path, cv2.cvtColor(crop, cv2.COLOR_RGB2BGR)) return crop_path def main(): parser = argparse.ArgumentParser() parser.add_argument("--limit", type=int, default=None, help="限制条数(测试用)") parser.add_argument("--language", type=str, default=None, help="按语言过滤,如 简中") parser.add_argument("--year", type=str, default=None, help="按年份过滤,如 1996") parser.add_argument("--no-yolo", action="store_true", help="跳过YOLO裁切(图库已是干净卡面时用)") parser.add_argument("--rebuild-features", action="store_true", help="强制重新提取特征") parser.add_argument("--csv", type=str, default=None, help="从CSV读取卡牌数据(服务器无DB时用)") args = parser.parse_args() config.ensure_dirs() # 1. 读卡牌元数据(CSV 优先,否则 PG) if args.csv or os.path.exists(config.CARD_MASTER_CSV): csv_path = args.csv or config.CARD_MASTER_CSV print(f"[1/4] 从 CSV 读取: {csv_path}") cards = csv_reader.read_card_master_csv(csv_path) # 应用 limit if args.limit: cards = cards[:args.limit] total = len(cards) print(f" 共 {total} 张卡牌") else: # 构建过滤条件 filters = {} if args.language: filters["language"] = [args.language] if args.year: filters["year"] = [args.year] print(f"[1/4] 读取 cards_master_v2 ... 过滤: {filters or '无'}") cards = db_reader.read_card_master(config.PG_CONFIG, config.PG_TABLE, filters, args.limit) total = len(cards) print(f" 共 {total} 张卡牌") if total == 0: print("无数据,退出") return # 2. 下载图片 print(f"[2/4] 下载卡牌标准图到 {config.GALLERY_IMG_DIR} ...") t0 = time.time() url_to_path = image_downloader.download_card_images_batch( [c["img_url"] for c in cards], config.GALLERY_IMG_DIR, config.DOWNLOAD_WORKERS ) downloaded = sum(1 for v in url_to_path.values() if v) print(f" 下载完成 {downloaded}/{total},耗时 {time.time()-t0:.1f}s") # 3. YOLO裁切 + 特征提取 print(f"[3/4] YOLO裁切 + DINOv2特征提取 ...") detector = None if args.no_yolo else CardDetector(config.YOLO_MODEL_PATH, conf=config.YOLO_CONF_THRESHOLD) extractor = CardFeatureExtractor(config.DINOV2_MODEL_PATH) print(f" 设备: {extractor.device}, 特征维度: {extractor.feature_dim}") all_feats = [] valid_cards = [] # 成功提取特征的卡牌元数据 batch_imgs, batch_meta = [], [] def flush_batch(): if not batch_imgs: return feats = extractor.extract(batch_imgs) for i, meta in enumerate(batch_meta): if not np.isnan(feats[i]).any(): all_feats.append(feats[i]) valid_cards.append(meta) for i, card in enumerate(cards): path = url_to_path.get(card["img_url"]) if not path: continue img = image_downloader.load_image_rgb(path) if img is None: continue if detector is not None: img = detector.detect_and_crop(img) if img is None: continue batch_imgs.append(img) batch_meta.append(card) if len(batch_imgs) >= config.FEATURE_BATCH_SIZE: flush_batch() batch_imgs, batch_meta = [], [] print(f" 进度 {i+1}/{total}, 已提取 {len(all_feats)}", end="\r") flush_batch() print(f"\n 特征提取完成,有效 {len(all_feats)}/{total}") if not all_feats: print("无有效特征,退出") return # 4. 构建并保存索引 print(f"[4/4] 保存索引 ...") features = np.stack(all_feats).astype(np.float32) metas = [{"card_id": c["card_id"], "card_name_ch": c.get("card_name_ch"), "language": c.get("language"), "year": c.get("year"), "card_no": c.get("card_no"), "pg_label": c.get("pg_label")} for c in valid_cards] card_ids = [c["card_id"] for c in valid_cards] index = VectorIndex(features, card_ids, metas) index.save(config.GALLERY_FEATURES_PATH, config.GALLERY_META_PATH) print(f" 已保存: {config.GALLERY_FEATURES_PATH}") print(f" 特征矩阵: {features.shape}") print(f"\n完成!图库构建成功,共 {len(card_ids)} 张") if __name__ == "__main__": main()