| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149 |
- """
- 构建图库(离线,一次性)
- 流程:读 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()
|