ab_test_v0904.py 5.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123
  1. # -*- coding: utf-8 -*-
  2. """ab_test_v0904.py - v0904 vs v819 同尺 A/B 对比(上半区+下半区)。
  3. query: v0904 held-out val 集(upper_data_v0904 的 val_groups,4,326 张现成半区裁剪图)
  4. gallery: 各自系统当时构建的 npy(v819=data/gallery_v819(88,384), v0904=data/gallery_v0904(91,116))
  5. query 特征各用对应版本权重提;检索指标 = card_id 级 Top-1/Top-5。
  6. head-to-head 只统计 card_id 在两个库都存在的 query;另报 v0904 全量(含新增卡)覆盖表现。
  7. """
  8. import os
  9. os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
  10. import json, time
  11. import numpy as np
  12. import cv2
  13. import torch
  14. import torch.nn as nn
  15. import torch.nn.functional as F
  16. from transformers import Dinov2Model
  17. WZJ = "/home/user/顾工交接/wzj"
  18. H, W = 196, 392
  19. MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
  20. STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
  21. DEV = "cuda:0"
  22. BS = 128
  23. class Back(nn.Module):
  24. def __init__(self):
  25. super().__init__()
  26. self.backbone = Dinov2Model.from_pretrained("facebook/dinov2-large")
  27. def forward(self, x):
  28. o = self.backbone(x, interpolate_pos_encoding=True)
  29. return F.normalize(o.last_hidden_state[:, 0, :], p=2, dim=1)
  30. def load_model(path):
  31. m = Back()
  32. sd = torch.load(path, map_location="cpu")
  33. m.load_state_dict(sd)
  34. return m.to(DEV).eval()
  35. def prep(paths):
  36. out = np.empty((len(paths), H, W, 3), dtype=np.float32)
  37. for i, p in enumerate(paths):
  38. img = cv2.imread(p)
  39. if img is None:
  40. raise RuntimeError("unreadable " + p)
  41. img = cv2.resize(cv2.cvtColor(img, cv2.COLOR_BGR2RGB), (W, H), interpolation=cv2.INTER_CUBIC)
  42. out[i] = (img.astype(np.float32) / 255.0 - MEAN) / STD
  43. return torch.from_numpy(out).permute(0, 3, 1, 2).contiguous()
  44. @torch.no_grad()
  45. def feats(model, paths):
  46. fs = []
  47. for i in range(0, len(paths), BS):
  48. fs.append(model(prep(paths[i:i+BS]).to(DEV)).cpu())
  49. return torch.cat(fs).numpy().astype(np.float32)
  50. def topk_metrics(qf, gf, q_cid, g_cid, ks=(1, 5)):
  51. qn, gn = qf / (np.linalg.norm(qf, axis=1, keepdims=True) + 1e-9), gf / (np.linalg.norm(gf, axis=1, keepdims=True) + 1e-9)
  52. t1 = t5 = 0
  53. for i in range(0, len(qn), 512):
  54. sims = qn[i:i+512] @ gn.T
  55. idx = np.argpartition(-sims, 4, axis=1)[:, :5] # top5 无序
  56. order = np.take_along_axis(sims, idx, axis=1).argsort(axis=1)[:, ::-1]
  57. top5 = idx[np.arange(len(idx))[:, None], order]
  58. hit = (g_cid[top5] == q_cid[i:i+512][:, None])
  59. t1 += int(hit[:, 0].sum()); t5 += int(hit.any(axis=1).sum())
  60. n = len(qn)
  61. return t1 / n, t5 / n
  62. def main():
  63. t0 = time.time()
  64. # ---------- 公共数据 ----------
  65. val_ids = []
  66. for g in json.load(open(f"{WZJ}/upper_data_v0904/val_groups.json", encoding="utf-8")):
  67. val_ids.extend(g["card_ids"]) # 字段名 card_ids 实为 sample_id
  68. meta = json.load(open(f"{WZJ}/data/train_meta_v0904.json", encoding="utf-8"))["records"]
  69. sid2cid = {r["sample_id"]: r["card_id"] for r in meta}
  70. q_cid_all = np.array([sid2cid[s] for s in val_ids])
  71. print(f"[q] val={len(val_ids)}", flush=True)
  72. gallery_dirs = {
  73. "v819": (f"{WZJ}/data/gallery_v819", f"{WZJ}/upper_model_output/best_upper_half_model.pth",
  74. f"{WZJ}/layer3_bg_model_output/best_layer3_bottom_model.pth"),
  75. "v0904": (f"{WZJ}/data/gallery_v0904", f"{WZJ}/upper_model_output_v0904/best_upper_half_model.pth",
  76. f"{WZJ}/layer3_bg_model_output_v0904/best_layer3_bottom_model.pth"),
  77. }
  78. g = {}
  79. for k, (d, up, lo) in gallery_dirs.items():
  80. meta_j = json.load(open(f"{d}/gallery_dual_meta.json", encoding="utf-8"))
  81. g_cids = meta_j["card_ids"] if "card_ids" in meta_j else [m["card_id"] for m in meta_j["metas"]]
  82. g[k] = {"u": np.load(f"{d}/gallery_upper_features.npy"),
  83. "l": np.load(f"{d}/gallery_lower_features.npy"),
  84. "cid": np.array(g_cids)}
  85. print(f"[g:{k}] n={len(g[k]['cid'])} upper={g[k]['u'].shape} lower={g[k]['l'].shape}", flush=True)
  86. # head-to-head: card_id 两库都存在的 query
  87. both = set(g["v819"]["cid"]) & set(g["v0904"]["cid"])
  88. mask_h2h = np.array([c in both for c in q_cid_all])
  89. ids_h2h = [s for s, m in zip(val_ids, mask_h2h) if m]
  90. print(f"[h2h] 双库共有card_id的query: {mask_h2h.sum()}/{len(val_ids)}", flush=True)
  91. for half, img_dir, gk in (("upper", f"{WZJ}/upper_data_v0904/images", "u"),
  92. ("lower", f"{WZJ}/lower_all_data_v0904/images", "l")):
  93. print(f"\n===== {half} 半区 A/B =====", flush=True)
  94. paths_all = [f"{img_dir}/{s}.jpg" for s in val_ids]
  95. for ver in ("v819", "v0904"):
  96. t1 = time.time()
  97. mp = gallery_dirs[ver][1 if half == "upper" else 2]
  98. model = load_model(mp)
  99. qf_all = feats(model, paths_all)
  100. # v0904 全量覆盖(含新增卡)
  101. m1, m5 = topk_metrics(qf_all, g[ver][gk], q_cid_all, g[ver]["cid"])
  102. # head-to-head
  103. idx = np.where(mask_h2h)[0]
  104. qf_h = qf_all[idx]
  105. h1, h5 = topk_metrics(qf_h, g[ver][gk], q_cid_all[idx], g[ver]["cid"])
  106. print(f"[{ver}] 全量4326: Top1={m1:.2%} Top5={m5:.2%} | h2h({len(idx)}): Top1={h1:.2%} Top5={h5:.2%} ({time.time()-t1:.0f}s)", flush=True)
  107. del model, qf_all
  108. torch.cuda.empty_cache()
  109. print(f"\nALL DONE {time.time()-t0:.0f}s", flush=True)
  110. if __name__ == "__main__":
  111. main()