# -*- coding: utf-8 -*- """mine_layer3_step2_pairs.py - 混淆对挖矿:老下半区模型实测"同精灵不同版本"高相似对。 逻辑(对齐老方案): 版本组 = (card_name_ch, card_no, pg_label) 的样本集合 混淆判定 = 同 card_name_ch 下, 两个版本组的代表向量(组内均值归一) cos > 阈值 混淆group = 该name下互相连通(传递闭包)的版本组并集, 组内样本互为hard negative 排除: card_name_ch 为空 / 含"能量" 输出: layer3_data/{images, train_groups.json, val_groups.json, manifest.csv} 用法: python mine_layer3_step2_pairs.py --cos-threshold 0.90 [--topk-preview] """ import argparse import json import os import shutil from collections import defaultdict import numpy as np WZJ = "/home/user/顾工交接/wzj" META = os.path.join(WZJ, "data/train_meta_88407.json") FEATS = os.path.join(WZJ, "lower_all_feats.npy") IDS = os.path.join(WZJ, "lower_all_ids.json") SRC_IMG = os.path.join(WZJ, "lower_all_data/images") OUT_DIR = os.path.join(WZJ, "layer3_data") def main(): ap = argparse.ArgumentParser() ap.add_argument("--cos-threshold", type=float, default=0.90) ap.add_argument("--exclude-same-no", action="store_true", default=True, help="排除同card_no不同系列的对(多为Deck Exclusives同图重复,标签噪音)") ap.add_argument("--max-per-version", type=int, default=4, help="每版本组最多采样张数(控规模)") ap.add_argument("--val-ratio", type=float, default=0.10) ap.add_argument("--seed", type=int, default=42) ap.add_argument("--topk-preview", action="store_true", help="只打印相似度分布不落盘") args = ap.parse_args() with open(META, encoding="utf-8") as f: meta = json.load(f)["records"] by_sid = {r["sample_id"]: r for r in meta} feats = np.load(FEATS) with open(IDS, encoding="utf-8") as f: ids = json.load(f) assert len(ids) == feats.shape[0] sid2idx = {s: i for i, s in enumerate(ids)} # 版本组: (name, card_no, pg_label) -> [sample_id] ver_groups = defaultdict(list) skipped_name = 0 for s, i in sid2idx.items(): r = by_sid.get(s) if r is None: continue name = (r.get("card_name_ch") or "").strip() if not name or "能量" in name: skipped_name += 1 continue ver_groups[(name, r.get("card_no") or "", r.get("pg_label") or "")].append(s) print(f"[ver_groups] {len(ver_groups)} skipped_name={skipped_name}") # 每版本代表向量(先按max-per-version抽样再取均值) rng = np.random.RandomState(args.seed) keys = sorted(ver_groups.keys()) reps = np.zeros((len(keys), feats.shape[1]), dtype=np.float32) kept_members = [] for ki, k in enumerate(keys): members = ver_groups[k] if len(members) > args.max_per_version: members = list(rng.choice(members, args.max_per_version, replace=False)) kept_members.append(members) v = feats[[sid2idx[m] for m in members]].mean(axis=0) reps[ki] = v / (np.linalg.norm(v) + 1e-8) # 按 name 聚类版本组下标 name2vis = defaultdict(list) for ki, k in enumerate(keys): name2vis[k[0]].append(ki) multi_ver_names = {n: v for n, v in name2vis.items() if len(v) >= 2} print(f"[names] total={len(name2vis)} multi_version={len(multi_ver_names)}") # 同name下版本对相似度(排除同card_no重复图对) pair_sims = [] n_excluded_dup = 0 for n, vis in multi_ver_names.items(): R = reps[vis] S = R @ R.T for a in range(len(vis)): for b in range(a + 1, len(vis)): if args.exclude_same_no and keys[vis[a]][1] == keys[vis[b]][1]: n_excluded_dup += 1 continue pair_sims.append((float(S[a, b]), vis[a], vis[b], n)) print(f"[pairs] excluded_same_no_dups={n_excluded_dup}") sims_arr = np.array([p[0] for p in pair_sims]) print(f"[pairs] {len(pair_sims)} version-pairs; " f"p50={np.percentile(sims_arr,50):.3f} p75={np.percentile(sims_arr,75):.3f} " f"p90={np.percentile(sims_arr,90):.3f} p95={np.percentile(sims_arr,95):.3f} max={sims_arr.max():.3f}") if args.topk_preview: top = sorted(pair_sims, reverse=True)[:30] for s, a, b, n in top: print(f" {s:.4f} {n} | {keys[a][1:] if False else keys[a]} <-> {keys[b]}") return # 连通分量合并混淆版本组 thr = args.cos_threshold confusable = [(s, a, b) for s, a, b, _ in pair_sims if s > thr] print(f"[confusable] pairs>thr({thr}): {len(confusable)}") parent = list(range(len(keys))) def find(x): while parent[x] != x: parent[x] = parent[parent[x]] x = parent[x] return x def union(a, b): ra, rb = find(a), find(b) if ra != rb: parent[rb] = ra for _, a, b in confusable: union(a, b) # 只保留含>=2版本且来自混淆合并的组 merge_map = defaultdict(list) for ki in range(len(keys)): merge_map[find(ki)].append(ki) out_groups = [] for root, kis in merge_map.items(): if len(kis) < 2: continue members = [] meta_info = None for ki in kis: members.extend(kept_members[ki]) if meta_info is None: k = keys[ki] meta_info = {"name": k[0], "card_no": k[1], "pg_label": k[2]} if len(members) < 2: continue # language 取第一个样本 m0 = by_sid.get(members[0], {}) out_groups.append({ "name": meta_info["name"], "pg_label": meta_info["pg_label"], "language": m0.get("language") or "", "card_ids": members, }) n_imgs = sum(len(g["card_ids"]) for g in out_groups) print(f"[groups] confusable_groups={len(out_groups)} total_imgs={n_imgs}") # val split 按组 gidx = list(range(len(out_groups))) rng.shuffle(gidx) n_val = int(len(out_groups) * args.val_ratio) val_set = set(gidx[:n_val]) os.makedirs(os.path.join(OUT_DIR, "images"), exist_ok=True) def dump(sel_idx, path): out = [] for gid, gi in enumerate(sorted(sel_idx)): g = out_groups[gi] out.append({"group_id": gid, **{k: g[k] for k in ("name", "pg_label", "language", "card_ids")}}) with open(path, "w", encoding="utf-8") as f: json.dump(out, f, ensure_ascii=False) return out train_json = dump([i for i in gidx if i not in val_set], os.path.join(OUT_DIR, "train_groups.json")) val_json = dump(val_set, os.path.join(OUT_DIR, "val_groups.json")) # 拷贝图片 + manifest import csv man_rows = [] n_copy = 0 all_sel = set() for g in train_json + val_json: all_sel.update(g["card_ids"]) for sid in sorted(all_sel): src = os.path.join(SRC_IMG, f"{sid}.jpg") dst = os.path.join(OUT_DIR, "images", f"{sid}.jpg") if os.path.exists(src) and not os.path.exists(dst): shutil.copy2(src, dst) n_copy += 1 r = by_sid.get(sid, {}) man_rows.append([sid, r.get("card_id") or "", r.get("card_name_ch") or "", r.get("card_no") or "", r.get("pg_label") or "", r.get("language") or ""]) with open(os.path.join(OUT_DIR, "manifest.csv"), "w", encoding="utf-8-sig", newline="") as f: w = csv.writer(f) w.writerow(["sample_id", "card_id", "card_name_ch", "card_no", "pg_label", "language"]) w.writerows(man_rows) print(f"[DONE] train_groups={len(train_json)}({sum(len(g['card_ids']) for g in train_json)} imgs) " f"val_groups={len(val_json)}({sum(len(g['card_ids']) for g in val_json)} imgs) " f"copied={n_copy} -> {OUT_DIR}") if __name__ == "__main__": main()