cascade_match_batch.py 5.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139
  1. # -*- coding: utf-8 -*-
  2. """
  3. cascade_match_batch.py — 双区级联(DINOv2-large 上半区召回 + 下半区重排)批量匹配。
  4. 交易平台图目录 → 每张出 card_id + 卡名/系列 + fusion/upper/lower 相似度 + Top5 候选,
  5. 写 CSV,同时写 <out_csv>.topk.json(每张的 Top-N 全部候选含 upper/lower + img_filename
  6. 供 HTML 可视化用)。这是把批量交易图匹配从 A4 单区**切换到双区级联**的生产入口,
  7. 取代 scripts/match_transactions.py(旧 A4 768 维)。
  8. 用法(pytorch 环境):
  9. python scripts/cascade_match_batch.py --image-dir <目录> --out-csv <path.csv> [--alpha 0.5] [--top-k-recall 30] [--top-n 5]
  10. """
  11. import os
  12. import sys
  13. import csv
  14. import json
  15. import time
  16. import hashlib
  17. import argparse
  18. sys.stdout.reconfigure(encoding="utf-8")
  19. _THIS = os.path.dirname(os.path.abspath(__file__))
  20. _ROOT = os.path.dirname(_THIS)
  21. sys.path.insert(0, _ROOT)
  22. sys.path.insert(0, _THIS)
  23. import config
  24. from cascade_match import DualCascadeMatcher
  25. FIELDS = ["id", "predicted_card_id", "card_name_ch", "series", "language", "year",
  26. "master_card_no", "fusion", "upper_sim", "lower_sim", "high_conf", "top5"]
  27. def build_card_id_to_img(card_master_csv):
  28. """card_master_all.csv → {card_id: (img_url, md5.jpg)},同 card_id 取首条。"""
  29. m = {}
  30. with open(card_master_csv, encoding="utf-8-sig") as f:
  31. for r in csv.DictReader(f):
  32. cid = r.get("card_id", "").strip()
  33. url = r.get("img_url", "").strip()
  34. if cid and url and cid not in m:
  35. m[cid] = (url, hashlib.md5(url.encode()).hexdigest() + ".jpg")
  36. return m
  37. def main():
  38. ap = argparse.ArgumentParser()
  39. ap.add_argument("--image-dir", required=True)
  40. ap.add_argument("--out-csv", required=True)
  41. ap.add_argument("--alpha", type=float, default=None)
  42. ap.add_argument("--top-k-recall", type=int, default=None)
  43. ap.add_argument("--top-n", type=int, default=5)
  44. ap.add_argument("--device", default=None)
  45. ap.add_argument("--backend", choices=["npy", "milvus"], default=None,
  46. help="检索后端,默认取 config.CASCADE_BACKEND(=npy),不传则跟现状完全一样")
  47. args = ap.parse_args()
  48. config.ensure_dirs()
  49. matcher = DualCascadeMatcher(alpha=args.alpha, top_k_recall=args.top_k_recall,
  50. device=args.device, backend=args.backend)
  51. id2meta = {m["card_id"]: m for m in matcher.metas}
  52. print(f"[batch] 读取 card_master_all.csv 建 card_id→img 映射 ...", flush=True)
  53. id2img = build_card_id_to_img(config.CARD_MASTER_ALL_CSV)
  54. print(f"[batch] card_master 覆盖 {len(id2img)} 个 card_id", flush=True)
  55. files = sorted(f for f in os.listdir(args.image_dir)
  56. if f.lower().endswith((".jpg", ".jpeg", ".png", ".webp")))
  57. print(f"[batch] {len(files)} 张待匹配 α={matcher.alpha} K_recall={matcher.top_k_recall} top_n={args.top_n}", flush=True)
  58. thr = config.SIMILARITY_THRESHOLD
  59. rows = []
  60. topk_dump = {}
  61. n_ok = n_fail = n_high = 0
  62. t0 = time.time()
  63. for i, fn in enumerate(files, 1):
  64. stem = os.path.splitext(fn)[0]
  65. path = os.path.join(args.image_dir, fn)
  66. raw = matcher.query_raw(path)
  67. if raw is None:
  68. rows.append({k: "" for k in FIELDS} | {"id": stem, "predicted_card_id": "FAIL"})
  69. topk_dump[stem] = []
  70. n_fail += 1
  71. continue
  72. ranked = matcher.rank(raw, matcher.alpha) # [(card_id, fusion, u, l), ...]
  73. top1 = ranked[0]
  74. meta = id2meta.get(top1[0], {})
  75. high = top1[1] >= thr
  76. n_high += int(high); n_ok += 1
  77. # 组装 topk.json 每条
  78. top_list = []
  79. for cid, fu, u, l in ranked[:args.top_n]:
  80. m = id2meta.get(cid, {})
  81. url, img_fn = id2img.get(cid, ("", ""))
  82. top_list.append({
  83. "card_id": cid, "fusion": round(float(fu), 4),
  84. "upper": round(float(u), 4), "lower": round(float(l), 4),
  85. "card_name_ch": m.get("card_name_ch", ""),
  86. "series": m.get("pg_label", ""),
  87. "language": m.get("language", ""),
  88. "year": m.get("year", ""),
  89. "master_card_no": m.get("card_no", ""),
  90. "img_filename": img_fn,
  91. })
  92. topk_dump[stem] = top_list
  93. rows.append({
  94. "id": stem,
  95. "predicted_card_id": top1[0],
  96. "card_name_ch": meta.get("card_name_ch", ""),
  97. "series": meta.get("pg_label", ""),
  98. "language": meta.get("language", ""),
  99. "year": meta.get("year", ""),
  100. "master_card_no": meta.get("card_no", ""),
  101. "fusion": f"{top1[1]:.4f}",
  102. "upper_sim": f"{top1[2]:.4f}",
  103. "lower_sim": f"{top1[3]:.4f}",
  104. "high_conf": "1" if high else "0",
  105. "top5": " | ".join(f"{t['card_id']}:{t['fusion']:.3f}" for t in top_list),
  106. })
  107. if i % 50 == 0 or i == len(files):
  108. print(f" [{i}/{len(files)}] ok={n_ok} fail={n_fail} 高置信={n_high} "
  109. f"({(time.time()-t0)/i:.2f}s/张)", flush=True)
  110. with open(args.out_csv, "w", encoding="utf-8-sig", newline="") as f:
  111. w = csv.DictWriter(f, fieldnames=FIELDS)
  112. w.writeheader()
  113. w.writerows(rows)
  114. topk_path = args.out_csv.rsplit(".", 1)[0] + ".topk.json"
  115. with open(topk_path, "w", encoding="utf-8") as f:
  116. json.dump(topk_dump, f, ensure_ascii=False, indent=2)
  117. print(f"\n[batch] 写出 {args.out_csv}", flush=True)
  118. print(f"[batch] 写出 {topk_path}", flush=True)
  119. print(f"[batch] 总 {len(files)} 成功 {n_ok} 失败 {n_fail} "
  120. f"高置信(fusion≥{thr}) {n_high} ({n_high/max(1,len(files))*100:.1f}%)", flush=True)
  121. if __name__ == "__main__":
  122. main()