| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204 |
- # -*- 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_v0904.json")
- FEATS = os.path.join(WZJ, "lower_all_feats_v0904.npy")
- IDS = os.path.join(WZJ, "lower_all_ids_v0904.json")
- SRC_IMG = os.path.join(WZJ, "lower_all_data_v0904/images")
- OUT_DIR = os.path.join(WZJ, "layer3_data_v0904")
- 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()
|