# -*- coding: utf-8 -*- """ cascade_match_batch.py — 双区级联(DINOv2-large 上半区召回 + 下半区重排)批量匹配。 交易平台图目录 → 每张出 card_id + 卡名/系列 + fusion/upper/lower 相似度 + Top5 候选, 写 CSV,同时写 .topk.json(每张的 Top-N 全部候选含 upper/lower + img_filename 供 HTML 可视化用)。这是把批量交易图匹配从 A4 单区**切换到双区级联**的生产入口, 取代 scripts/match_transactions.py(旧 A4 768 维)。 用法(pytorch 环境): python scripts/cascade_match_batch.py --image-dir <目录> --out-csv [--alpha 0.5] [--top-k-recall 30] [--top-n 5] """ import os import sys import csv import json import time import hashlib import argparse sys.stdout.reconfigure(encoding="utf-8") _THIS = os.path.dirname(os.path.abspath(__file__)) _ROOT = os.path.dirname(_THIS) sys.path.insert(0, _ROOT) sys.path.insert(0, _THIS) import config from cascade_match import DualCascadeMatcher FIELDS = ["id", "predicted_card_id", "card_name_ch", "series", "language", "year", "master_card_no", "fusion", "upper_sim", "lower_sim", "high_conf", "top5"] def build_card_id_to_img(card_master_csv): """card_master_all.csv → {card_id: (img_url, md5.jpg)},同 card_id 取首条。""" m = {} with open(card_master_csv, encoding="utf-8-sig") as f: for r in csv.DictReader(f): cid = r.get("card_id", "").strip() url = r.get("img_url", "").strip() if cid and url and cid not in m: m[cid] = (url, hashlib.md5(url.encode()).hexdigest() + ".jpg") return m def main(): ap = argparse.ArgumentParser() ap.add_argument("--image-dir", required=True) ap.add_argument("--out-csv", required=True) ap.add_argument("--alpha", type=float, default=None) ap.add_argument("--top-k-recall", type=int, default=None) ap.add_argument("--top-n", type=int, default=5) ap.add_argument("--device", default=None) ap.add_argument("--backend", choices=["npy", "milvus"], default=None, help="检索后端,默认取 config.CASCADE_BACKEND(=npy),不传则跟现状完全一样") args = ap.parse_args() config.ensure_dirs() matcher = DualCascadeMatcher(alpha=args.alpha, top_k_recall=args.top_k_recall, device=args.device, backend=args.backend) id2meta = {m["card_id"]: m for m in matcher.metas} print(f"[batch] 读取 card_master_all.csv 建 card_id→img 映射 ...", flush=True) id2img = build_card_id_to_img(config.CARD_MASTER_ALL_CSV) print(f"[batch] card_master 覆盖 {len(id2img)} 个 card_id", flush=True) files = sorted(f for f in os.listdir(args.image_dir) if f.lower().endswith((".jpg", ".jpeg", ".png", ".webp"))) print(f"[batch] {len(files)} 张待匹配 α={matcher.alpha} K_recall={matcher.top_k_recall} top_n={args.top_n}", flush=True) thr = config.SIMILARITY_THRESHOLD rows = [] topk_dump = {} n_ok = n_fail = n_high = 0 t0 = time.time() for i, fn in enumerate(files, 1): stem = os.path.splitext(fn)[0] path = os.path.join(args.image_dir, fn) raw = matcher.query_raw(path) if raw is None: rows.append({k: "" for k in FIELDS} | {"id": stem, "predicted_card_id": "FAIL"}) topk_dump[stem] = [] n_fail += 1 continue ranked = matcher.rank(raw, matcher.alpha) # [(card_id, fusion, u, l), ...] top1 = ranked[0] meta = id2meta.get(top1[0], {}) high = top1[1] >= thr n_high += int(high); n_ok += 1 # 组装 topk.json 每条 top_list = [] for cid, fu, u, l in ranked[:args.top_n]: m = id2meta.get(cid, {}) url, img_fn = id2img.get(cid, ("", "")) top_list.append({ "card_id": cid, "fusion": round(float(fu), 4), "upper": round(float(u), 4), "lower": round(float(l), 4), "card_name_ch": m.get("card_name_ch", ""), "series": m.get("pg_label", ""), "language": m.get("language", ""), "year": m.get("year", ""), "master_card_no": m.get("card_no", ""), "img_filename": img_fn, }) topk_dump[stem] = top_list rows.append({ "id": stem, "predicted_card_id": top1[0], "card_name_ch": meta.get("card_name_ch", ""), "series": meta.get("pg_label", ""), "language": meta.get("language", ""), "year": meta.get("year", ""), "master_card_no": meta.get("card_no", ""), "fusion": f"{top1[1]:.4f}", "upper_sim": f"{top1[2]:.4f}", "lower_sim": f"{top1[3]:.4f}", "high_conf": "1" if high else "0", "top5": " | ".join(f"{t['card_id']}:{t['fusion']:.3f}" for t in top_list), }) if i % 50 == 0 or i == len(files): print(f" [{i}/{len(files)}] ok={n_ok} fail={n_fail} 高置信={n_high} " f"({(time.time()-t0)/i:.2f}s/张)", flush=True) with open(args.out_csv, "w", encoding="utf-8-sig", newline="") as f: w = csv.DictWriter(f, fieldnames=FIELDS) w.writeheader() w.writerows(rows) topk_path = args.out_csv.rsplit(".", 1)[0] + ".topk.json" with open(topk_path, "w", encoding="utf-8") as f: json.dump(topk_dump, f, ensure_ascii=False, indent=2) print(f"\n[batch] 写出 {args.out_csv}", flush=True) print(f"[batch] 写出 {topk_path}", flush=True) print(f"[batch] 总 {len(files)} 成功 {n_ok} 失败 {n_fail} " f"高置信(fusion≥{thr}) {n_high} ({n_high/max(1,len(files))*100:.1f}%)", flush=True) if __name__ == "__main__": main()