feature_extractor_dual.py 6.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148
  1. # -*- coding: utf-8 -*-
  2. """
  3. 双区 DINOv2-large 特征提取模块(上半区·精灵识别 + 下半区·Layer3 版本/特效)。
  4. 两个半区模型完全同构:facebook/dinov2-large,取 CLS token,L2 归一化,
  5. **无 BN Neck**,forward 必须开 interpolate_pos_encoding=True(输入 196×392 非正方形)。
  6. 预处理严格对齐训练侧(stage1_upper_half.py 的 letterbox + train_dinov2_*_half.py 的 transform_clean):
  7. 原图/裁切图 → CardDetector.detect_and_crop(RGB 数组) → PIL letterbox 392×392
  8. → 上半区 rows[0:196] / 下半区 rows[196:392](196×392)
  9. → ImageNet 归一化 → backbone(interpolate_pos_encoding=True) → CLS → L2
  10. 建库与查询必须用同一套预处理(本模块),否则相似度不可比。
  11. """
  12. import os
  13. import numpy as np
  14. from PIL import Image
  15. os.environ.setdefault("HF_ENDPOINT", "https://hf-mirror.com")
  16. HF_MODEL_ID = "facebook/dinov2-large"
  17. IMG_SIZE = 392
  18. HALF = 196 # 392 // 2,与 stage1 / Layer3 stage3 完全一致
  19. IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
  20. IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
  21. FEATURE_DIM = 1024 # dinov2-large CLS 维度
  22. # ===================== 预处理(与训练严格对齐,纯函数,不依赖 torch)=====================
  23. def letterbox392(pil_img):
  24. """与 stage1_upper_half.py / Layer3 stage3_finalize_bottom_crop_norot.letterbox 完全一致。
  25. PIL BICUBIC、min-ratio、黑色(0,0,0)填充、居中、输出 392×392。"""
  26. img = pil_img.convert("RGB")
  27. w, h = img.size
  28. r = min(IMG_SIZE / h, IMG_SIZE / w)
  29. nw, nh = int(w * r), int(h * r)
  30. img = img.resize((nw, nh), Image.BICUBIC)
  31. canvas = Image.new("RGB", (IMG_SIZE, IMG_SIZE), (0, 0, 0))
  32. canvas.paste(img, ((IMG_SIZE - nw) // 2, (IMG_SIZE - nh) // 2))
  33. return canvas
  34. def split_half(pil_letterboxed):
  35. """从 392×392 letterbox 图取上下半区(各 196×392)。
  36. 上半区 crop((0,0,392,196)) = rows[0:196];下半区 crop((0,196,392,392)) = rows[196:392]。
  37. 两者来自同一 letterbox,card_id 一一对应。"""
  38. upper = pil_letterboxed.crop((0, 0, IMG_SIZE, HALF))
  39. lower = pil_letterboxed.crop((0, HALF, IMG_SIZE, IMG_SIZE))
  40. return upper, lower
  41. def half_to_chw(pil_half):
  42. """PIL 196×392 半区 → ImageNet 归一化 CHW float32(对齐 A.Resize(196,392)+Normalize+ToTensorV2)。
  43. 若尺寸非 196×392,用 INTER_CUBIC 补 resize(与训练 transform 一致)。"""
  44. import cv2
  45. arr = np.array(pil_half.convert("RGB"))
  46. if arr.shape[:2] != (HALF, IMG_SIZE):
  47. arr = cv2.resize(arr, (IMG_SIZE, HALF), interpolation=cv2.INTER_CUBIC)
  48. arr = arr.astype(np.float32) / 255.0
  49. arr = (arr - IMAGENET_MEAN) / IMAGENET_STD
  50. return arr.transpose(2, 0, 1) # CHW
  51. class DualCardFeatureExtractor:
  52. """双区 DINOv2-large 卡牌特征提取器:上半区认精灵、下半区认版本/特效。"""
  53. def __init__(self, upper_pth, lower_pth, device=None, batch_size=64, hf_model_id=HF_MODEL_ID):
  54. import torch
  55. self.torch = torch
  56. self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
  57. self.batch_size = batch_size
  58. self.feature_dim = FEATURE_DIM
  59. print(f"[DualExtractor] 加载上半区模型: {upper_pth}", flush=True)
  60. self.upper = self._load_half_net(upper_pth, hf_model_id)
  61. print(f"[DualExtractor] 加载下半区模型: {lower_pth}", flush=True)
  62. self.lower = self._load_half_net(lower_pth, hf_model_id)
  63. print(f"[DualExtractor] device={self.device} batch_size={batch_size}", flush=True)
  64. def _load_half_net(self, pth_path, hf_model_id):
  65. """加载单个半区 dinov2-large backbone 并装入训练好的权重。
  66. 训练模型结构为 self.backbone=Dinov2Model,故 .pth 的 key 带 'backbone.' 前缀,这里剥离后装入裸 Dinov2Model。"""
  67. import torch
  68. from transformers import Dinov2Model
  69. backbone = Dinov2Model.from_pretrained(hf_model_id)
  70. state = torch.load(pth_path, map_location="cpu")
  71. state = {k[len("backbone."):]: v for k, v in state.items() if k.startswith("backbone.")}
  72. missing, unexpected = backbone.load_state_dict(state, strict=False)
  73. if missing:
  74. print(f" [warn] missing keys(前5): {list(missing)[:5]}", flush=True)
  75. if unexpected:
  76. print(f" [warn] unexpected keys(前5): {list(unexpected)[:5]}", flush=True)
  77. backbone.to(self.device).eval()
  78. return backbone
  79. # ---------- 静态预处理入口(建库/查询共用)----------
  80. @staticmethod
  81. def letterbox(pil_img):
  82. return letterbox392(pil_img)
  83. @staticmethod
  84. def split_half(pil_letterboxed):
  85. return split_half(pil_letterboxed)
  86. @staticmethod
  87. def crop_to_halves(crop_rgb):
  88. """detect_and_crop 返回的 RGB numpy 数组 → (上半区PIL, 下半区PIL)。"""
  89. lb = letterbox392(Image.fromarray(crop_rgb))
  90. return split_half(lb)
  91. # ---------- 批量提特征 ----------
  92. def _forward(self, backbone, pil_halves):
  93. """对一批 PIL 半区图提 CLS 并 L2 归一化,返回 (N, 1024) float32,失败行 NaN。"""
  94. import torch.nn.functional as F
  95. torch = self.torch
  96. out = np.full((len(pil_halves), self.feature_dim), np.nan, dtype=np.float32)
  97. tensors, valid = [], []
  98. for i, p in enumerate(pil_halves):
  99. if p is None:
  100. continue
  101. try:
  102. tensors.append(torch.from_numpy(half_to_chw(p)).float())
  103. valid.append(i)
  104. except Exception:
  105. continue
  106. if not tensors:
  107. return out
  108. for s in range(0, len(tensors), self.batch_size):
  109. chunk = torch.stack(tensors[s:s + self.batch_size]).to(self.device)
  110. with torch.no_grad():
  111. # transformers>=4.5x 的 Dinov2Model.forward 已去掉 interpolate_pos_encoding 参数
  112. # (非标准分辨率时内部自动插值);硬传会 TypeError。
  113. try:
  114. cls = backbone(chunk, interpolate_pos_encoding=True).last_hidden_state[:, 0, :]
  115. except TypeError:
  116. cls = backbone(chunk).last_hidden_state[:, 0, :]
  117. feat = F.normalize(cls, p=2, dim=1).cpu().numpy()
  118. for k, vi in enumerate(valid[s:s + self.batch_size]):
  119. out[vi] = feat[k]
  120. return out
  121. def extract_upper(self, pil_uppers):
  122. """pil_uppers: list[PIL 196×392](或 None)→ (N, 1024) L2 归一化。"""
  123. return self._forward(self.upper, pil_uppers)
  124. def extract_lower(self, pil_lowers):
  125. """pil_lowers: list[PIL 196×392](或 None)→ (N, 1024) L2 归一化。"""
  126. return self._forward(self.lower, pil_lowers)