| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123 |
- # -*- coding: utf-8 -*-
- """ab_test_v0904.py - v0904 vs v819 同尺 A/B 对比(上半区+下半区)。
- query: v0904 held-out val 集(upper_data_v0904 的 val_groups,4,326 张现成半区裁剪图)
- gallery: 各自系统当时构建的 npy(v819=data/gallery_v819(88,384), v0904=data/gallery_v0904(91,116))
- query 特征各用对应版本权重提;检索指标 = card_id 级 Top-1/Top-5。
- head-to-head 只统计 card_id 在两个库都存在的 query;另报 v0904 全量(含新增卡)覆盖表现。
- """
- import os
- os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
- import json, time
- import numpy as np
- import cv2
- import torch
- import torch.nn as nn
- import torch.nn.functional as F
- from transformers import Dinov2Model
- WZJ = "/home/user/顾工交接/wzj"
- H, W = 196, 392
- MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
- STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
- DEV = "cuda:0"
- BS = 128
- class Back(nn.Module):
- def __init__(self):
- super().__init__()
- self.backbone = Dinov2Model.from_pretrained("facebook/dinov2-large")
- def forward(self, x):
- o = self.backbone(x, interpolate_pos_encoding=True)
- return F.normalize(o.last_hidden_state[:, 0, :], p=2, dim=1)
- def load_model(path):
- m = Back()
- sd = torch.load(path, map_location="cpu")
- m.load_state_dict(sd)
- return m.to(DEV).eval()
- def prep(paths):
- out = np.empty((len(paths), H, W, 3), dtype=np.float32)
- for i, p in enumerate(paths):
- img = cv2.imread(p)
- if img is None:
- raise RuntimeError("unreadable " + p)
- img = cv2.resize(cv2.cvtColor(img, cv2.COLOR_BGR2RGB), (W, H), interpolation=cv2.INTER_CUBIC)
- out[i] = (img.astype(np.float32) / 255.0 - MEAN) / STD
- return torch.from_numpy(out).permute(0, 3, 1, 2).contiguous()
- @torch.no_grad()
- def feats(model, paths):
- fs = []
- for i in range(0, len(paths), BS):
- fs.append(model(prep(paths[i:i+BS]).to(DEV)).cpu())
- return torch.cat(fs).numpy().astype(np.float32)
- def topk_metrics(qf, gf, q_cid, g_cid, ks=(1, 5)):
- qn, gn = qf / (np.linalg.norm(qf, axis=1, keepdims=True) + 1e-9), gf / (np.linalg.norm(gf, axis=1, keepdims=True) + 1e-9)
- t1 = t5 = 0
- for i in range(0, len(qn), 512):
- sims = qn[i:i+512] @ gn.T
- idx = np.argpartition(-sims, 4, axis=1)[:, :5] # top5 无序
- order = np.take_along_axis(sims, idx, axis=1).argsort(axis=1)[:, ::-1]
- top5 = idx[np.arange(len(idx))[:, None], order]
- hit = (g_cid[top5] == q_cid[i:i+512][:, None])
- t1 += int(hit[:, 0].sum()); t5 += int(hit.any(axis=1).sum())
- n = len(qn)
- return t1 / n, t5 / n
- def main():
- t0 = time.time()
- # ---------- 公共数据 ----------
- val_ids = []
- for g in json.load(open(f"{WZJ}/upper_data_v0904/val_groups.json", encoding="utf-8")):
- val_ids.extend(g["card_ids"]) # 字段名 card_ids 实为 sample_id
- meta = json.load(open(f"{WZJ}/data/train_meta_v0904.json", encoding="utf-8"))["records"]
- sid2cid = {r["sample_id"]: r["card_id"] for r in meta}
- q_cid_all = np.array([sid2cid[s] for s in val_ids])
- print(f"[q] val={len(val_ids)}", flush=True)
- gallery_dirs = {
- "v819": (f"{WZJ}/data/gallery_v819", f"{WZJ}/upper_model_output/best_upper_half_model.pth",
- f"{WZJ}/layer3_bg_model_output/best_layer3_bottom_model.pth"),
- "v0904": (f"{WZJ}/data/gallery_v0904", f"{WZJ}/upper_model_output_v0904/best_upper_half_model.pth",
- f"{WZJ}/layer3_bg_model_output_v0904/best_layer3_bottom_model.pth"),
- }
- g = {}
- for k, (d, up, lo) in gallery_dirs.items():
- meta_j = json.load(open(f"{d}/gallery_dual_meta.json", encoding="utf-8"))
- g_cids = meta_j["card_ids"] if "card_ids" in meta_j else [m["card_id"] for m in meta_j["metas"]]
- g[k] = {"u": np.load(f"{d}/gallery_upper_features.npy"),
- "l": np.load(f"{d}/gallery_lower_features.npy"),
- "cid": np.array(g_cids)}
- print(f"[g:{k}] n={len(g[k]['cid'])} upper={g[k]['u'].shape} lower={g[k]['l'].shape}", flush=True)
- # head-to-head: card_id 两库都存在的 query
- both = set(g["v819"]["cid"]) & set(g["v0904"]["cid"])
- mask_h2h = np.array([c in both for c in q_cid_all])
- ids_h2h = [s for s, m in zip(val_ids, mask_h2h) if m]
- print(f"[h2h] 双库共有card_id的query: {mask_h2h.sum()}/{len(val_ids)}", flush=True)
- for half, img_dir, gk in (("upper", f"{WZJ}/upper_data_v0904/images", "u"),
- ("lower", f"{WZJ}/lower_all_data_v0904/images", "l")):
- print(f"\n===== {half} 半区 A/B =====", flush=True)
- paths_all = [f"{img_dir}/{s}.jpg" for s in val_ids]
- for ver in ("v819", "v0904"):
- t1 = time.time()
- mp = gallery_dirs[ver][1 if half == "upper" else 2]
- model = load_model(mp)
- qf_all = feats(model, paths_all)
- # v0904 全量覆盖(含新增卡)
- m1, m5 = topk_metrics(qf_all, g[ver][gk], q_cid_all, g[ver]["cid"])
- # head-to-head
- idx = np.where(mask_h2h)[0]
- qf_h = qf_all[idx]
- h1, h5 = topk_metrics(qf_h, g[ver][gk], q_cid_all[idx], g[ver]["cid"])
- 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)
- del model, qf_all
- torch.cuda.empty_cache()
- print(f"\nALL DONE {time.time()-t0:.0f}s", flush=True)
- if __name__ == "__main__":
- main()
|