cascade_match.py 21 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443
  1. # -*- coding: utf-8 -*-
  2. """
  3. cascade_match.py - 双区级联卡牌检索(上半区召回 + 下半区重排)
  4. 查询图 → YOLO裁卡 → PIL letterbox 392 → 上/下半区特征
  5. → 上半区全库 Top-K 召回(物种级) → 下半区对候选算版本相似度
  6. → fusion = α·upper_sim + (1-α)·lower_sim 重排 → 最终 card_id
  7. 三种模式:
  8. --image <路径> 单图查询(打印 Top-K)
  9. --image-dir <目录> 目录批量查询
  10. --self-test N 全库自检索 sanity(抽 N 张 gallery 图,Top-1 应命中自己 card_id)
  11. --pairs <json> 版本消歧测试(test_pairs_valid.json:同精灵普通卡 vs 球闪卡),
  12. 内置 α∈{1.0,0.7,0.5,0.3,0.0} 扫描,输出各权重分开成功率
  13. 用法(pytorch 环境):
  14. python cascade_match.py --self-test 500
  15. python cascade_match.py --pairs test_pairs_valid.json --alpha 0.5
  16. python cascade_match.py --image /path/to/card.jpg
  17. """
  18. import os
  19. import sys
  20. import json
  21. import time
  22. import random
  23. import hashlib
  24. import argparse
  25. sys.stdout.reconfigure(encoding="utf-8")
  26. _THIS = os.path.dirname(os.path.abspath(__file__))
  27. _ROOT = os.path.dirname(_THIS)
  28. sys.path.insert(0, _ROOT)
  29. sys.path.insert(0, _THIS)
  30. import numpy as np
  31. from PIL import Image
  32. import config
  33. from modules.yolo_detector import CardDetector
  34. from modules.feature_extractor_dual import DualCardFeatureExtractor, letterbox392, split_half
  35. ALPHAS = [1.0, 0.7, 0.5, 0.3, 0.0] # pairs 扫描权重;1.0=纯上半区召回(不重排), 0.0=纯下半区
  36. def md5n(u):
  37. return hashlib.md5(u.encode()).hexdigest() + ".jpg"
  38. class DualCascadeMatcher:
  39. def __init__(self, alpha=None, top_k_recall=None, device=None, verbose=True, backend=None,
  40. detector_device=None, milvus_alias=None):
  41. """device: DINOv2 特征提取器设备(torch.device/str/None=自动)。
  42. detector_device: YOLO 分割模型设备(None=ultralytics 自动=CVD 掩码后第一张;
  43. 多 worker 分卡时传 torch.device(f"cuda:N")——ultralytics select_device 对
  44. torch.device 输入直接返回、不改写 CUDA_VISIBLE_DEVICES)。
  45. milvus_alias: pymilvus 连接 alias(None=默认连接)。多 worker 并发时各自
  46. 独立 alias,避免共享默认连接;单例服务不传即保持旧行为。"""
  47. self.alpha = config.CASCADE_ALPHA if alpha is None else alpha
  48. self.top_k_recall = config.CASCADE_TOP_K_RECALL if top_k_recall is None else top_k_recall
  49. self.verbose = verbose
  50. self.backend = backend or getattr(config, "CASCADE_BACKEND", "npy")
  51. assert self.backend in ("npy", "milvus", "gpu"), \
  52. f"未知 backend: {self.backend!r}(只支持 npy/milvus/gpu)"
  53. with open(config.GALLERY_DUAL_META_PATH, "r", encoding="utf-8") as f:
  54. m = json.load(f)
  55. self.card_ids = m["card_ids"]
  56. self.metas = m["metas"]
  57. if self.backend == "gpu" and device is None:
  58. # gpu 后端必须有确定设备:调用方没传时自动选(与 extractor 同规则)
  59. try:
  60. import torch as _t
  61. device = _t.device("cuda") if _t.cuda.is_available() else None
  62. except Exception:
  63. device = None
  64. if device is None:
  65. print("[Cascade] ⚠ backend=gpu 但无可用 CUDA,回退 npy", flush=True)
  66. self.backend = "npy"
  67. if self.backend == "npy":
  68. if self.verbose:
  69. print(f"[Cascade] 加载双区特征库(npy) ...", flush=True)
  70. self.gallery_upper = np.load(config.GALLERY_UPPER_FEATURES_PATH).astype(np.float32)
  71. self.gallery_lower = np.load(config.GALLERY_LOWER_FEATURES_PATH).astype(np.float32)
  72. assert self.gallery_upper.shape[0] == self.gallery_lower.shape[0] == len(self.card_ids), \
  73. "双区特征库行数与 meta 不一致!"
  74. elif self.backend == "gpu":
  75. # 双区特征库常驻 GPU 显存(fp32):单卡召回 torch 矩阵乘实测 0.4ms,
  76. # 替代 milvus FLAT 扫描+网络往返的 ~40ms/卡(2026-09-03 bench:Top-5 与
  77. # milvus 3/3 精确一致;fp16 图库无额外收益且有 tie 抖动,故用 fp32)。
  78. # 每 worker 显存 +724MB(2×88384×1024×4B),V100 16G 无压力。
  79. import torch
  80. if self.verbose:
  81. print(f"[Cascade] 加载双区特征库(GPU 常驻, device={device}) ...", flush=True)
  82. self.gallery_upper = torch.from_numpy(
  83. np.load(config.GALLERY_UPPER_FEATURES_PATH).astype(np.float32)).to(device)
  84. self.gallery_lower = torch.from_numpy(
  85. np.load(config.GALLERY_LOWER_FEATURES_PATH).astype(np.float32)).to(device)
  86. assert self.gallery_upper.shape[0] == self.gallery_lower.shape[0] == len(self.card_ids), \
  87. "双区特征库行数与 meta 不一致!"
  88. else:
  89. alias = milvus_alias or "default"
  90. if self.verbose:
  91. print(f"[Cascade] 连接 Milvus {config.MILVUS_HOST}:{config.MILVUS_PORT} "
  92. f"(alias={alias}) ...", flush=True)
  93. from pymilvus import connections, Collection
  94. if not connections.has_connection(alias):
  95. connections.connect(alias=alias, host=config.MILVUS_HOST, port=config.MILVUS_PORT)
  96. self.upper_coll = Collection(config.MILVUS_UPPER_COLLECTION, using=alias)
  97. self.lower_coll = Collection(config.MILVUS_LOWER_COLLECTION, using=alias)
  98. if self.verbose:
  99. print(f"[Cascade] 图库 {len(self.card_ids)} 行 K_recall={self.top_k_recall} "
  100. f"α={self.alpha} backend={self.backend}", flush=True)
  101. self.detector = CardDetector(config.YOLO_MODEL_PATH, device=detector_device)
  102. self.extractor = DualCardFeatureExtractor(
  103. config.UPPER_MODEL_PATH, config.LOWER_MODEL_PATH, device=device,
  104. )
  105. def _features_of(self, img_path, want_label=False, want_crop=False):
  106. """读图 → 裁卡 → letterbox → (上半区特征, 下半区特征[, 卡牌类型label][, 裁卡RGB])。
  107. want_label=True 时第三个返回值是 YOLO 检出的最大实例类别名(如 'pokemon'),供 card_type。
  108. want_crop=True 时额外返回裁卡 RGB numpy 数组(供 element 检测等下游)。"""
  109. try:
  110. arr = np.array(Image.open(img_path).convert("RGB"))
  111. except Exception:
  112. return (None, None, None, None) if (want_label and want_crop) else \
  113. (None, None, None) if want_label else (None, None)
  114. if want_label:
  115. crop, card_type = self.detector.detect_and_crop(arr, return_label=True)
  116. else:
  117. crop = self.detector.detect_and_crop(arr)
  118. card_type = None
  119. if crop is None:
  120. crop = arr
  121. lb = letterbox392(Image.fromarray(crop))
  122. up, lo = split_half(lb)
  123. q_u = self.extractor.extract_upper([up])[0]
  124. q_l = self.extractor.extract_lower([lo])[0]
  125. if np.isnan(q_u).any() or np.isnan(q_l).any():
  126. return (None, None, None, None) if (want_label and want_crop) else \
  127. (None, None, None) if want_label else (None, None)
  128. ret = (q_u, q_l)
  129. if want_label:
  130. ret = ret + (card_type,)
  131. if want_crop:
  132. ret = ret + (crop,)
  133. return ret
  134. def _query_raw_npy(self, q_u, q_l):
  135. sims_u = self.gallery_upper @ q_u # (N,)
  136. K = min(self.top_k_recall, sims_u.shape[0])
  137. cand = np.argpartition(-sims_u, K - 1)[:K] # 上半区召回 K 个候选 idx
  138. cand = cand[np.argsort(-sims_u[cand])] # 按 upper_sim 降序
  139. su = sims_u[cand]
  140. sl = self.gallery_lower[cand] @ q_l # 候选的下半区相似度
  141. cand_ids = [self.card_ids[int(c)] for c in cand]
  142. return cand_ids, su, sl
  143. def _query_raw_milvus(self, q_u, q_l):
  144. res = self.upper_coll.search(
  145. data=[q_u.tolist()], anns_field="vector",
  146. param={"metric_type": "COSINE", "params": {}},
  147. limit=self.top_k_recall, output_fields=["card_id"],
  148. )
  149. cand_ids = [hit.entity.get("card_id") for hit in res[0]]
  150. su = np.array([hit.distance for hit in res[0]], dtype=np.float32)
  151. id_list = ",".join(f'"{c}"' for c in cand_ids)
  152. lower_rows = self.lower_coll.query(expr=f"card_id in [{id_list}]", output_fields=["card_id", "vector"])
  153. lower_map = {r["card_id"]: np.array(r["vector"], dtype=np.float32) for r in lower_rows}
  154. sl = np.array([float(np.dot(lower_map[c], q_l)) if c in lower_map else -1.0 for c in cand_ids])
  155. return cand_ids, su, sl
  156. def _query_raw_gpu(self, q_u, q_l):
  157. """GPU 常驻图库召回(backend='gpu'):与 _query_raw_npy 同数学(归一化向量
  158. 点积=COSINE),只是库在显存、查询向量不过网。2026-09-03 bench:Top-5 与
  159. milvus 精确一致、单卡 0.4ms(milvus 40.9ms)。"""
  160. import torch
  161. dev = self.gallery_upper.device
  162. q_u_t = torch.from_numpy(q_u).to(dev)
  163. K = min(self.top_k_recall, self.gallery_upper.shape[0])
  164. tv = torch.topk(self.gallery_upper @ q_u_t, K)
  165. cand = tv.indices.cpu().numpy()
  166. cand_ids = [self.card_ids[int(c)] for c in cand]
  167. su = tv.values.detach().cpu().numpy().astype(np.float32)
  168. q_l_t = torch.from_numpy(q_l).to(dev)
  169. cand_t = torch.from_numpy(cand).to(dev)
  170. sl = (self.gallery_lower[cand_t] @ q_l_t).detach().cpu().numpy().astype(np.float32)
  171. return cand_ids, su, sl
  172. def _query_raw(self, q_u, q_l):
  173. """按 backend 分发检索。"""
  174. if self.backend == "npy":
  175. return self._query_raw_npy(q_u, q_l)
  176. if self.backend == "gpu":
  177. return self._query_raw_gpu(q_u, q_l)
  178. return self._query_raw_milvus(q_u, q_l)
  179. def query_raw(self, img_path, want_crop=False):
  180. """返回上半区召回的候选 card_id + 各候选的 upper_sim/lower_sim(供扫 α)。
  181. 额外带 card_type(YOLO 检出的卡牌类型 label);want_crop=True 时额外带 crop(裁卡RGB)。"""
  182. # _features_of 返回元组长度随请求变化(2/3/4),按 want_crop 分别解包
  183. if want_crop:
  184. q_u, q_l, card_type, crop = self._features_of(img_path, want_label=True, want_crop=True)
  185. else:
  186. q_u, q_l, card_type = self._features_of(img_path, want_label=True, want_crop=False)
  187. crop = None
  188. if q_u is None:
  189. return None
  190. cand_ids, su, sl = self._query_raw(q_u, q_l)
  191. ret = {"cand_ids": cand_ids, "upper_sim": su, "lower_sim": sl, "card_type": card_type}
  192. if want_crop:
  193. ret["crop"] = crop
  194. return ret
  195. def _features_from_crop(self, crop):
  196. """已裁卡 RGB → (q_u, q_l)。失败返回 (None, None)。"""
  197. if crop is None:
  198. return None, None
  199. lb = letterbox392(Image.fromarray(crop))
  200. up, lo = split_half(lb)
  201. q_u = self.extractor.extract_upper([up])[0]
  202. q_l = self.extractor.extract_lower([lo])[0]
  203. if np.isnan(q_u).any() or np.isnan(q_l).any():
  204. return None, None
  205. return q_u, q_l
  206. def query_all(self, img_path, alpha=None, top_n=5):
  207. """多卡查询:YOLO 一次检出全部实例,按阅读顺序各自检索。
  208. 返回 list[dict],每项结构同 query()(含 crop/box)。未检出返回 []。
  209. 后台 /match 仍走 query() 单卡,本方法仅给小程序 /match_fields。
  210. 特征提取按实例批量(上下半区各合一批前向):多卡图 GPU 利用率显著更高
  211. (FP32 bs=2 实测每 half 18.1→11.7ms),批量与逐张前向数值差 ~1e-6 量级,
  212. 不影响检索结果。"""
  213. a = self.alpha if alpha is None else alpha
  214. try:
  215. arr = np.array(Image.open(img_path).convert("RGB"))
  216. except Exception:
  217. return []
  218. instances = self.detector.detect_and_crop_all(arr)
  219. if not instances:
  220. return []
  221. ups, los = [], []
  222. for inst in instances:
  223. up, lo = split_half(letterbox392(Image.fromarray(inst["crop"])))
  224. ups.append(up)
  225. los.append(lo)
  226. q_us = self.extractor.extract_upper(ups) # (N,1024),失败行 NaN
  227. q_ls = self.extractor.extract_lower(los)
  228. out = []
  229. for i, inst in enumerate(instances):
  230. q_u, q_l = q_us[i], q_ls[i]
  231. if q_u is None or q_l is None or np.isnan(q_u).any() or np.isnan(q_l).any():
  232. continue
  233. cand_ids, su, sl = self._query_raw(q_u, q_l)
  234. raw = {"cand_ids": cand_ids, "upper_sim": su, "lower_sim": sl,
  235. "card_type": inst.get("label"), "crop": inst["crop"]}
  236. ranked = self.rank(raw, a)
  237. if not ranked:
  238. continue
  239. out.append({
  240. "predicted_card_id": ranked[0][0],
  241. "fusion_score": ranked[0][1],
  242. "upper_sim": ranked[0][2],
  243. "lower_sim": ranked[0][3],
  244. "card_type": inst.get("label"),
  245. "top_k": [{"card_id": r[0], "fusion": round(r[1], 4),
  246. "upper": round(r[2], 4), "lower": round(r[3], 4)}
  247. for r in ranked[:top_n]],
  248. "crop": inst["crop"],
  249. "box": inst.get("box"),
  250. })
  251. return out
  252. def rank(self, raw, alpha):
  253. """给定 raw 和 α,返回 [(card_id, fusion, upper_sim, lower_sim), ...] 按 fusion 降序。"""
  254. cand_ids = raw["cand_ids"]
  255. fusion = alpha * raw["upper_sim"] + (1 - alpha) * raw["lower_sim"]
  256. order = np.argsort(-fusion)
  257. out = []
  258. for j in order:
  259. out.append((cand_ids[j], float(fusion[j]),
  260. float(raw["upper_sim"][j]), float(raw["lower_sim"][j])))
  261. return out
  262. def query(self, img_path, alpha=None, top_n=5, want_crop=False):
  263. """单次查询(用 self.alpha),返回 dict(含 card_type=YOLO 卡牌类型 label)。
  264. want_crop=True 时额外返回 crop(裁卡 RGB numpy 数组,供 element 检测等)。"""
  265. a = self.alpha if alpha is None else alpha
  266. raw = self.query_raw(img_path, want_crop=want_crop)
  267. if raw is None:
  268. return None
  269. ranked = self.rank(raw, a)
  270. ret = {
  271. "predicted_card_id": ranked[0][0],
  272. "fusion_score": ranked[0][1],
  273. "upper_sim": ranked[0][2],
  274. "lower_sim": ranked[0][3],
  275. "card_type": raw.get("card_type"),
  276. "top_k": [{"card_id": r[0], "fusion": round(r[1], 4),
  277. "upper": round(r[2], 4), "lower": round(r[3], 4)} for r in ranked[:top_n]],
  278. }
  279. if want_crop:
  280. ret["crop"] = raw.get("crop")
  281. return ret
  282. # ===================== 模式:单图 / 目录 =====================
  283. def mode_images(matcher, paths, alpha):
  284. for p in paths:
  285. r = matcher.query(p, alpha=alpha)
  286. if r is None:
  287. print(f"[FAIL] {p}", flush=True)
  288. continue
  289. print(f"{os.path.basename(p):<40} -> {r['predicted_card_id']} "
  290. f"fusion={r['fusion_score']:.3f} (u={r['upper_sim']:.3f} l={r['lower_sim']:.3f})", flush=True)
  291. # ===================== 模式:自检索 sanity =====================
  292. def mode_self_test(matcher, n):
  293. import csv
  294. cards = list(csv.DictReader(open(config.CARD_MASTER_ALL_CSV, encoding="utf-8-sig")))
  295. # 只取有图、且 card_id 在当前(可能 limit)图库里的
  296. in_lib = set(matcher.card_ids)
  297. cands = [c for c in cards if c.get("card_id") in in_lib
  298. and os.path.exists(os.path.join(config.GALLERY_IMG_DIR, md5n(c.get("img_url", ""))))]
  299. random.seed(42)
  300. random.shuffle(cands)
  301. cands = cands[:n]
  302. print(f"[SelfTest] 抽样 {len(cands)} 张 gallery 图做自检索(Top-1 应命中同 card_id)", flush=True)
  303. hit = 0
  304. t0 = time.time()
  305. for i, c in enumerate(cands, 1):
  306. img_path = os.path.join(config.GALLERY_IMG_DIR, md5n(c["img_url"]))
  307. r = matcher.query(img_path)
  308. ok = r and r["predicted_card_id"] == c["card_id"]
  309. hit += int(bool(ok))
  310. if i % 100 == 0 or i == len(cands):
  311. print(f" [{i}/{len(cands)}] 命中率 {hit/i:.2%} ({(time.time()-t0)/i:.2f}s/张)", flush=True)
  312. report = {"mode": "self_test", "n": len(cands), "hit": hit, "top1_rate": hit / max(1, len(cands))}
  313. print(f"\n[SelfTest] Top-1 自命中: {hit}/{len(cands)} = {hit/max(1,len(cands)):.2%}", flush=True)
  314. _save_report(report, "self_test")
  315. return report
  316. # ===================== 模式:版本消歧 pairs(含 α 扫描)=====================
  317. def mode_pairs(matcher, pairs_json):
  318. from modules import image_downloader
  319. pairs = json.load(open(pairs_json, encoding="utf-8"))
  320. print(f"[Pairs] {len(pairs)} 对,每对查 url_a(普通)+url_b(球闪),扫 α={ALPHAS}", flush=True)
  321. # 每个查询缓存 raw(扫 α 复用),记录期望 card_id
  322. queries = [] # (expect_id, raw_or_None, side, name)
  323. for p in pairs:
  324. name = p.get("name", "")
  325. for side, prefix in [("a", "id_a"), ("b", "id_b")]:
  326. cid = p.get(prefix)
  327. url = p.get(f"url_{side}")
  328. if not cid or not url:
  329. continue
  330. local = image_downloader.download_card_image(url, config.QUERY_IMG_DIR)
  331. raw = matcher.query_raw(local) if local else None
  332. queries.append((cid, raw, side, name, url))
  333. tag = "OK" if raw else "FAIL_DL"
  334. print(f" [{name}/{side}] expect={cid} {tag}", flush=True)
  335. valid = [q for q in queries if q[1] is not None]
  336. print(f"[Pairs] 有效查询 {len(valid)}/{len(queries)}(下载失败 {len(queries)-len(valid)})", flush=True)
  337. # 扫 α:对每个 α,统计 Top-1 命中期望、混淆到对家、其他错
  338. table = []
  339. per_query_detail = []
  340. for alpha in ALPHAS:
  341. hit = confuse = other = 0
  342. for cid, raw, side, name, url in valid:
  343. ranked = matcher.rank(raw, alpha)
  344. pred = ranked[0][0]
  345. # 对家 card_id:a 的对家是同 pair 的 b
  346. pair = next(pp for pp in pairs if cid in (pp.get("id_a"), pp.get("id_b")))
  347. other_id = pair["id_b"] if cid == pair.get("id_a") else pair["id_a"]
  348. if pred == cid:
  349. hit += 1
  350. elif pred == other_id:
  351. confuse += 1
  352. else:
  353. other += 1
  354. if alpha == matcher.alpha:
  355. per_query_detail.append({"name": name, "side": side, "expect": cid,
  356. "pred": pred, "other": other_id,
  357. "upper": round(ranked[0][2], 4), "lower": round(ranked[0][3], 4)})
  358. n = len(valid)
  359. table.append({"alpha": alpha, "hit": hit, "confuse_a_b": confuse, "other_wrong": other,
  360. "hit_rate": round(hit / n, 4) if n else None,
  361. "confuse_rate": round(confuse / n, 4) if n else None})
  362. print(f" α={alpha:.1f}: 正确 {hit}/{n}={hit/max(1,n):.2%} 混淆a↔b {confuse} 其他错 {other}", flush=True)
  363. report = {"mode": "pairs", "n_pairs": len(pairs), "n_valid_queries": len(valid),
  364. "alpha_scan": table, "detail_at_alpha": per_query_detail}
  365. _save_report(report, "pairs")
  366. return report
  367. def _save_report(report, tag):
  368. out = os.path.join(config.DATA_DIR, f"cascade_test_report_{tag}.json")
  369. with open(out, "w", encoding="utf-8") as f:
  370. json.dump(report, f, ensure_ascii=False, indent=2)
  371. print(f"[Report] 已保存 {out}", flush=True)
  372. def main():
  373. ap = argparse.ArgumentParser()
  374. ap.add_argument("--image", type=str, default=None)
  375. ap.add_argument("--image-dir", type=str, default=None)
  376. ap.add_argument("--self-test", type=int, default=0)
  377. ap.add_argument("--pairs", type=str, default=None)
  378. ap.add_argument("--alpha", type=float, default=None)
  379. ap.add_argument("--top-k-recall", type=int, default=None)
  380. ap.add_argument("--backend", choices=["npy", "milvus"], default=None,
  381. help="检索后端,默认取 config.CASCADE_BACKEND(=npy)")
  382. args = ap.parse_args()
  383. config.ensure_dirs()
  384. matcher = DualCascadeMatcher(alpha=args.alpha, top_k_recall=args.top_k_recall, backend=args.backend)
  385. if args.self_test:
  386. mode_self_test(matcher, args.self_test)
  387. elif args.pairs:
  388. mode_pairs(matcher, args.pairs)
  389. elif args.image or args.image_dir:
  390. paths = [args.image] if args.image else sorted(
  391. os.path.join(args.image_dir, f) for f in os.listdir(args.image_dir)
  392. if f.lower().endswith((".jpg", ".jpeg", ".png")))
  393. mode_images(matcher, paths, args.alpha)
  394. else:
  395. ap.error("需指定 --image / --image-dir / --self-test / --pairs 之一")
  396. if __name__ == "__main__":
  397. main()