mine_layer3_step2_pairs_v0904.py 7.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204
  1. # -*- coding: utf-8 -*-
  2. """mine_layer3_step2_pairs.py - 混淆对挖矿:老下半区模型实测"同精灵不同版本"高相似对。
  3. 逻辑(对齐老方案):
  4. 版本组 = (card_name_ch, card_no, pg_label) 的样本集合
  5. 混淆判定 = 同 card_name_ch 下, 两个版本组的代表向量(组内均值归一) cos > 阈值
  6. 混淆group = 该name下互相连通(传递闭包)的版本组并集, 组内样本互为hard negative
  7. 排除: card_name_ch 为空 / 含"能量"
  8. 输出: layer3_data/{images, train_groups.json, val_groups.json, manifest.csv}
  9. 用法: python mine_layer3_step2_pairs.py --cos-threshold 0.90 [--topk-preview]
  10. """
  11. import argparse
  12. import json
  13. import os
  14. import shutil
  15. from collections import defaultdict
  16. import numpy as np
  17. WZJ = "/home/user/顾工交接/wzj"
  18. META = os.path.join(WZJ, "data/train_meta_v0904.json")
  19. FEATS = os.path.join(WZJ, "lower_all_feats_v0904.npy")
  20. IDS = os.path.join(WZJ, "lower_all_ids_v0904.json")
  21. SRC_IMG = os.path.join(WZJ, "lower_all_data_v0904/images")
  22. OUT_DIR = os.path.join(WZJ, "layer3_data_v0904")
  23. def main():
  24. ap = argparse.ArgumentParser()
  25. ap.add_argument("--cos-threshold", type=float, default=0.90)
  26. ap.add_argument("--exclude-same-no", action="store_true", default=True,
  27. help="排除同card_no不同系列的对(多为Deck Exclusives同图重复,标签噪音)")
  28. ap.add_argument("--max-per-version", type=int, default=4, help="每版本组最多采样张数(控规模)")
  29. ap.add_argument("--val-ratio", type=float, default=0.10)
  30. ap.add_argument("--seed", type=int, default=42)
  31. ap.add_argument("--topk-preview", action="store_true", help="只打印相似度分布不落盘")
  32. args = ap.parse_args()
  33. with open(META, encoding="utf-8") as f:
  34. meta = json.load(f)["records"]
  35. by_sid = {r["sample_id"]: r for r in meta}
  36. feats = np.load(FEATS)
  37. with open(IDS, encoding="utf-8") as f:
  38. ids = json.load(f)
  39. assert len(ids) == feats.shape[0]
  40. sid2idx = {s: i for i, s in enumerate(ids)}
  41. # 版本组: (name, card_no, pg_label) -> [sample_id]
  42. ver_groups = defaultdict(list)
  43. skipped_name = 0
  44. for s, i in sid2idx.items():
  45. r = by_sid.get(s)
  46. if r is None:
  47. continue
  48. name = (r.get("card_name_ch") or "").strip()
  49. if not name or "能量" in name:
  50. skipped_name += 1
  51. continue
  52. ver_groups[(name, r.get("card_no") or "", r.get("pg_label") or "")].append(s)
  53. print(f"[ver_groups] {len(ver_groups)} skipped_name={skipped_name}")
  54. # 每版本代表向量(先按max-per-version抽样再取均值)
  55. rng = np.random.RandomState(args.seed)
  56. keys = sorted(ver_groups.keys())
  57. reps = np.zeros((len(keys), feats.shape[1]), dtype=np.float32)
  58. kept_members = []
  59. for ki, k in enumerate(keys):
  60. members = ver_groups[k]
  61. if len(members) > args.max_per_version:
  62. members = list(rng.choice(members, args.max_per_version, replace=False))
  63. kept_members.append(members)
  64. v = feats[[sid2idx[m] for m in members]].mean(axis=0)
  65. reps[ki] = v / (np.linalg.norm(v) + 1e-8)
  66. # 按 name 聚类版本组下标
  67. name2vis = defaultdict(list)
  68. for ki, k in enumerate(keys):
  69. name2vis[k[0]].append(ki)
  70. multi_ver_names = {n: v for n, v in name2vis.items() if len(v) >= 2}
  71. print(f"[names] total={len(name2vis)} multi_version={len(multi_ver_names)}")
  72. # 同name下版本对相似度(排除同card_no重复图对)
  73. pair_sims = []
  74. n_excluded_dup = 0
  75. for n, vis in multi_ver_names.items():
  76. R = reps[vis]
  77. S = R @ R.T
  78. for a in range(len(vis)):
  79. for b in range(a + 1, len(vis)):
  80. if args.exclude_same_no and keys[vis[a]][1] == keys[vis[b]][1]:
  81. n_excluded_dup += 1
  82. continue
  83. pair_sims.append((float(S[a, b]), vis[a], vis[b], n))
  84. print(f"[pairs] excluded_same_no_dups={n_excluded_dup}")
  85. sims_arr = np.array([p[0] for p in pair_sims])
  86. print(f"[pairs] {len(pair_sims)} version-pairs; "
  87. f"p50={np.percentile(sims_arr,50):.3f} p75={np.percentile(sims_arr,75):.3f} "
  88. f"p90={np.percentile(sims_arr,90):.3f} p95={np.percentile(sims_arr,95):.3f} max={sims_arr.max():.3f}")
  89. if args.topk_preview:
  90. top = sorted(pair_sims, reverse=True)[:30]
  91. for s, a, b, n in top:
  92. print(f" {s:.4f} {n} | {keys[a][1:] if False else keys[a]} <-> {keys[b]}")
  93. return
  94. # 连通分量合并混淆版本组
  95. thr = args.cos_threshold
  96. confusable = [(s, a, b) for s, a, b, _ in pair_sims if s > thr]
  97. print(f"[confusable] pairs>thr({thr}): {len(confusable)}")
  98. parent = list(range(len(keys)))
  99. def find(x):
  100. while parent[x] != x:
  101. parent[x] = parent[parent[x]]
  102. x = parent[x]
  103. return x
  104. def union(a, b):
  105. ra, rb = find(a), find(b)
  106. if ra != rb:
  107. parent[rb] = ra
  108. for _, a, b in confusable:
  109. union(a, b)
  110. # 只保留含>=2版本且来自混淆合并的组
  111. merge_map = defaultdict(list)
  112. for ki in range(len(keys)):
  113. merge_map[find(ki)].append(ki)
  114. out_groups = []
  115. for root, kis in merge_map.items():
  116. if len(kis) < 2:
  117. continue
  118. members = []
  119. meta_info = None
  120. for ki in kis:
  121. members.extend(kept_members[ki])
  122. if meta_info is None:
  123. k = keys[ki]
  124. meta_info = {"name": k[0], "card_no": k[1], "pg_label": k[2]}
  125. if len(members) < 2:
  126. continue
  127. # language 取第一个样本
  128. m0 = by_sid.get(members[0], {})
  129. out_groups.append({
  130. "name": meta_info["name"],
  131. "pg_label": meta_info["pg_label"],
  132. "language": m0.get("language") or "",
  133. "card_ids": members,
  134. })
  135. n_imgs = sum(len(g["card_ids"]) for g in out_groups)
  136. print(f"[groups] confusable_groups={len(out_groups)} total_imgs={n_imgs}")
  137. # val split 按组
  138. gidx = list(range(len(out_groups)))
  139. rng.shuffle(gidx)
  140. n_val = int(len(out_groups) * args.val_ratio)
  141. val_set = set(gidx[:n_val])
  142. os.makedirs(os.path.join(OUT_DIR, "images"), exist_ok=True)
  143. def dump(sel_idx, path):
  144. out = []
  145. for gid, gi in enumerate(sorted(sel_idx)):
  146. g = out_groups[gi]
  147. out.append({"group_id": gid, **{k: g[k] for k in ("name", "pg_label", "language", "card_ids")}})
  148. with open(path, "w", encoding="utf-8") as f:
  149. json.dump(out, f, ensure_ascii=False)
  150. return out
  151. train_json = dump([i for i in gidx if i not in val_set], os.path.join(OUT_DIR, "train_groups.json"))
  152. val_json = dump(val_set, os.path.join(OUT_DIR, "val_groups.json"))
  153. # 拷贝图片 + manifest
  154. import csv
  155. man_rows = []
  156. n_copy = 0
  157. all_sel = set()
  158. for g in train_json + val_json:
  159. all_sel.update(g["card_ids"])
  160. for sid in sorted(all_sel):
  161. src = os.path.join(SRC_IMG, f"{sid}.jpg")
  162. dst = os.path.join(OUT_DIR, "images", f"{sid}.jpg")
  163. if os.path.exists(src) and not os.path.exists(dst):
  164. shutil.copy2(src, dst)
  165. n_copy += 1
  166. r = by_sid.get(sid, {})
  167. man_rows.append([sid, r.get("card_id") or "", r.get("card_name_ch") or "",
  168. r.get("card_no") or "", r.get("pg_label") or "", r.get("language") or ""])
  169. with open(os.path.join(OUT_DIR, "manifest.csv"), "w", encoding="utf-8-sig", newline="") as f:
  170. w = csv.writer(f)
  171. w.writerow(["sample_id", "card_id", "card_name_ch", "card_no", "pg_label", "language"])
  172. w.writerows(man_rows)
  173. print(f"[DONE] train_groups={len(train_json)}({sum(len(g['card_ids']) for g in train_json)} imgs) "
  174. f"val_groups={len(val_json)}({sum(len(g['card_ids']) for g in val_json)} imgs) "
  175. f"copied={n_copy} -> {OUT_DIR}")
  176. if __name__ == "__main__":
  177. main()