| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139 |
- # -*- coding: utf-8 -*-
- """
- cascade_match_batch.py — 双区级联(DINOv2-large 上半区召回 + 下半区重排)批量匹配。
- 交易平台图目录 → 每张出 card_id + 卡名/系列 + fusion/upper/lower 相似度 + Top5 候选,
- 写 CSV,同时写 <out_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 <path.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()
|