train_dinov2_upper_half_v0904.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321
  1. # -*- coding: utf-8 -*-
  2. """
  3. train_dinov2_upper_half.py - 上半区精灵识别检索模型训练脚本
  4. 架构目标(双区):
  5. 上半区模型(本脚本) = 识别精灵/年龄段/姿势 (替代 cards_master_v2 card_name_ch 文本依赖)
  6. 下半区模型(Layer3) = 版本/特效(球闪等)消歧
  7. 两者都用 dinov2-large, 输入均为 392x196 (letterbox 392x392 的上/下半)。
  8. 对齐 Layer3(train_dinov2_layer3_bottom.py) 框架, 仅改:
  9. - 数据路径 -> upper_data_v0904 (全库上半区, organize_upper_dataset.py 产出)
  10. - 输入虽同为 392x196, 但取的是上半区(rows 0:196), Layer3 取下半区(rows 196:392)
  11. - group = card_name_ch + card_no (同卡同号不同实例作困难负样本, 与 A4 一致)
  12. 其余(冻结18层/InfoNCE/HardNegativeBatchSampler/增强策略)完全复用 Layer3。
  13. 用法(pytorch环境, GPU):
  14. CUDA_VISIBLE_DEVICES=0 ~/miniconda3/envs/pytorch/bin/python train_dinov2_upper_half.py
  15. """
  16. import os
  17. os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
  18. import json
  19. import random
  20. from collections import defaultdict
  21. import cv2
  22. import numpy as np
  23. import torch
  24. import torch.nn as nn
  25. import torch.nn.functional as F
  26. import torch.optim as optim
  27. import albumentations as A
  28. from albumentations.pytorch import ToTensorV2
  29. from torch.utils.data import BatchSampler, DataLoader, Dataset
  30. from tqdm import tqdm
  31. from transformers import Dinov2Model
  32. # ================= 配置区域 =================
  33. DATA_ROOT = r"/home/user/顾工交接/wzj/upper_data_v0904"
  34. IMG_DIR = os.path.join(DATA_ROOT, "images")
  35. TRAIN_GROUPS_JSON = os.path.join(DATA_ROOT, "train_groups.json")
  36. VAL_GROUPS_JSON = os.path.join(DATA_ROOT, "val_groups.json")
  37. SAVE_DIR = r"/home/user/顾工交接/wzj/upper_model_output_v0904"
  38. BEST_MODEL_PATH = os.path.join(SAVE_DIR, "best_upper_half_model.pth")
  39. HF_MODEL_ID = "facebook/dinov2-large"
  40. IMG_HEIGHT = 196 # 上半区高度 (392 letterbox 的上50%, 与 Layer3 下半区对称)
  41. IMG_WIDTH = 392 # 全宽
  42. FREEZE_BLOCKS = 18 # dinov2-large 共24层, 冻结前18层(与 Layer3 一致)
  43. BATCH_SIZE = 224
  44. EPOCHS = 100
  45. LEARNING_RATE = 1e-5
  46. WEIGHT_DECAY = 1e-4
  47. TEMPERATURE = 0.07
  48. EVAL_INTERVAL = 1
  49. HARD_NEGATIVE_PROB = 0.5
  50. HARD_GROUP_SIZE = 6
  51. DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  52. # ================= 数据增强 (与 Layer3 一致) =================
  53. transform_clean = A.Compose([
  54. A.Resize(IMG_HEIGHT, IMG_WIDTH, interpolation=cv2.INTER_CUBIC),
  55. A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
  56. ToTensorV2(),
  57. ])
  58. transform_aug = A.Compose([
  59. A.Resize(IMG_HEIGHT, IMG_WIDTH, interpolation=cv2.INTER_CUBIC),
  60. A.Perspective(scale=(0.02, 0.05), p=0.4),
  61. A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.5),
  62. A.RandomGamma(gamma_limit=(85, 115), p=0.4),
  63. A.GaussNoise(std_range=(0.05, 0.15), p=0.3),
  64. A.ImageCompression(quality_range=(60, 95), p=0.3),
  65. A.SafeRotate(limit=5, border_mode=cv2.BORDER_CONSTANT, fill=0, p=0.4),
  66. A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
  67. ToTensorV2(),
  68. ])
  69. # ================= 1. 数据集 (原样复用 Layer3 MiningGroupDataset) =================
  70. class MiningGroupDataset(Dataset):
  71. def __init__(self, groups_json_path, img_dir, is_train):
  72. with open(groups_json_path, "r", encoding="utf-8") as f:
  73. groups = json.load(f)
  74. self.samples = []
  75. label_idx = 0
  76. for g in groups:
  77. group_key = g["group_id"]
  78. for cid in g["card_ids"]:
  79. path = os.path.join(img_dir, f"{cid}.jpg")
  80. if not os.path.exists(path):
  81. continue
  82. self.samples.append((path, label_idx, group_key))
  83. label_idx += 1
  84. self.group_to_indices = defaultdict(list)
  85. for idx, (_, _, group_key) in enumerate(self.samples):
  86. self.group_to_indices[group_key].append(idx)
  87. self.valid_hard_groups = [g for g, idxs in self.group_to_indices.items() if len(idxs) >= 2]
  88. tag = "train" if is_train else "val"
  89. print(f"[{tag}] samples={len(self.samples)} hard_negative_groups(>=2)={len(self.valid_hard_groups)}")
  90. def __len__(self):
  91. return len(self.samples)
  92. def __getitem__(self, idx):
  93. path, label, group_key = self.samples[idx]
  94. img = cv2.imread(path)
  95. if img is None:
  96. img = np.zeros((IMG_HEIGHT, IMG_WIDTH, 3), dtype=np.uint8)
  97. else:
  98. img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
  99. img_clean = transform_clean(image=img)["image"]
  100. img_aug = transform_aug(image=img)["image"]
  101. return img_clean, img_aug, label
  102. # ================= 2. HardNegativeBatchSampler (原样复用 A4/Layer3) =================
  103. class HardNegativeBatchSampler(BatchSampler):
  104. def __init__(self, dataset, batch_size, hard_negative_prob=0.35, hard_group_size=8, drop_last=True):
  105. self.dataset = dataset
  106. self.batch_size = batch_size
  107. self.hard_negative_prob = hard_negative_prob
  108. self.hard_group_size = min(hard_group_size, batch_size)
  109. self.drop_last = drop_last
  110. self.num_samples = len(dataset)
  111. self.all_indices = list(range(self.num_samples))
  112. def __iter__(self):
  113. unused = set(self.all_indices)
  114. while len(unused) >= self.batch_size:
  115. batch = []
  116. if len(self.dataset.valid_hard_groups) > 0 and random.random() < self.hard_negative_prob:
  117. candidate_groups = []
  118. for g in self.dataset.valid_hard_groups:
  119. available = [i for i in self.dataset.group_to_indices[g] if i in unused]
  120. if len(available) >= 2:
  121. candidate_groups.append((g, available))
  122. if candidate_groups:
  123. g, available = random.choice(candidate_groups)
  124. take_n = min(len(available), self.hard_group_size, self.batch_size)
  125. chosen_hard = random.sample(available, take_n)
  126. batch.extend(chosen_hard)
  127. for i in chosen_hard:
  128. unused.remove(i)
  129. remain = self.batch_size - len(batch)
  130. if remain > 0:
  131. if len(unused) < remain:
  132. if self.drop_last:
  133. break
  134. batch.extend(list(unused))
  135. unused.clear()
  136. else:
  137. chosen_rand = random.sample(list(unused), remain)
  138. batch.extend(chosen_rand)
  139. for i in chosen_rand:
  140. unused.remove(i)
  141. if len(batch) == self.batch_size:
  142. random.shuffle(batch)
  143. yield batch
  144. if not self.drop_last and unused:
  145. yield list(unused)
  146. def __len__(self):
  147. if self.drop_last:
  148. return self.num_samples // self.batch_size
  149. return (self.num_samples + self.batch_size - 1) // self.batch_size
  150. # ================= 3. 模型: DINOv2-large 冻结前18层 =================
  151. class Dinov2UpperHalfModel(nn.Module):
  152. def __init__(self, freeze_blocks=FREEZE_BLOCKS):
  153. super().__init__()
  154. print(f"Loading {HF_MODEL_ID} backbone ...")
  155. self.backbone = Dinov2Model.from_pretrained(HF_MODEL_ID)
  156. for param in self.backbone.parameters():
  157. param.requires_grad = False
  158. total_layers = len(self.backbone.encoder.layer)
  159. print(f"Backbone total layers: {total_layers}, freezing first {freeze_blocks}")
  160. for i in range(freeze_blocks, total_layers):
  161. for param in self.backbone.encoder.layer[i].parameters():
  162. param.requires_grad = True
  163. for param in self.backbone.layernorm.parameters():
  164. param.requires_grad = True
  165. def forward(self, x):
  166. outputs = self.backbone(x, interpolate_pos_encoding=True)
  167. cls_token = outputs.last_hidden_state[:, 0, :]
  168. return F.normalize(cls_token, p=2, dim=1)
  169. # ================= 4. InfoNCE =================
  170. def contrastive_loss(feat_clean, feat_aug, temperature=TEMPERATURE):
  171. batch_size = feat_clean.size(0)
  172. logits = torch.matmul(feat_clean, feat_aug.T) / temperature
  173. labels = torch.arange(batch_size).to(feat_clean.device)
  174. loss_1 = F.cross_entropy(logits, labels)
  175. loss_2 = F.cross_entropy(logits.T, labels)
  176. return (loss_1 + loss_2) / 2.0
  177. # ================= 5. 检索评估 =================
  178. def evaluate(model, val_dataset):
  179. model.eval()
  180. loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4)
  181. gallery_feats, gallery_labels = [], []
  182. query_feats, query_labels = [], []
  183. with torch.no_grad():
  184. for imgs_c, imgs_a, lbls in loader:
  185. imgs_c, imgs_a = imgs_c.to(DEVICE), imgs_a.to(DEVICE)
  186. with torch.amp.autocast("cuda"):
  187. f_c = model(imgs_c)
  188. f_a = model(imgs_a)
  189. gallery_feats.append(f_c.cpu())
  190. gallery_labels.append(lbls)
  191. query_feats.append(f_a.cpu())
  192. query_labels.append(lbls)
  193. gallery_feats = torch.cat(gallery_feats, dim=0)
  194. gallery_labels = torch.cat(gallery_labels, dim=0)
  195. query_feats = torch.cat(query_feats, dim=0)
  196. query_labels = torch.cat(query_labels, dim=0)
  197. sim_matrix = torch.matmul(query_feats, gallery_feats.T)
  198. k = min(5, gallery_feats.size(0))
  199. topk_idx = sim_matrix.topk(k, dim=1).indices
  200. correct_1, correct_5, total = 0, 0, query_labels.size(0)
  201. for idx in range(total):
  202. q_label = query_labels[idx]
  203. retrieved = gallery_labels[topk_idx[idx]]
  204. if q_label == retrieved[0]:
  205. correct_1 += 1
  206. if q_label in retrieved:
  207. correct_5 += 1
  208. return correct_1 / total, correct_5 / total
  209. # ================= 6. 主流程 =================
  210. def main():
  211. os.makedirs(SAVE_DIR, exist_ok=True)
  212. train_dataset = MiningGroupDataset(TRAIN_GROUPS_JSON, IMG_DIR, is_train=True)
  213. val_dataset = MiningGroupDataset(VAL_GROUPS_JSON, IMG_DIR, is_train=False)
  214. batch_sampler = HardNegativeBatchSampler(
  215. dataset=train_dataset,
  216. batch_size=BATCH_SIZE,
  217. hard_negative_prob=HARD_NEGATIVE_PROB,
  218. hard_group_size=HARD_GROUP_SIZE,
  219. drop_last=True,
  220. )
  221. train_loader = DataLoader(train_dataset, batch_sampler=batch_sampler, num_workers=16,
  222. pin_memory=True, persistent_workers=True, prefetch_factor=3)
  223. model = Dinov2UpperHalfModel().to(DEVICE)
  224. if torch.cuda.device_count() > 1:
  225. model = nn.DataParallel(model)
  226. optimizer = optim.AdamW(
  227. filter(lambda p: p.requires_grad, model.parameters()),
  228. lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY,
  229. )
  230. scaler = torch.amp.GradScaler("cuda")
  231. best_top1_acc = 0.0
  232. print("开始 上半区(精灵识别) 对比学习训练 ...")
  233. for epoch in range(EPOCHS):
  234. model.train()
  235. pbar = tqdm(train_loader, desc=f"Epoch {epoch + 1}/{EPOCHS} [Train]")
  236. running_loss = 0.0
  237. for i, (imgs_clean, imgs_aug, _) in enumerate(pbar):
  238. imgs_clean = imgs_clean.to(DEVICE)
  239. imgs_aug = imgs_aug.to(DEVICE)
  240. optimizer.zero_grad()
  241. with torch.amp.autocast("cuda"):
  242. feat_clean = model(imgs_clean)
  243. feat_aug = model(imgs_aug)
  244. loss = contrastive_loss(feat_clean, feat_aug, TEMPERATURE)
  245. scaler.scale(loss).backward()
  246. scaler.step(optimizer)
  247. scaler.update()
  248. running_loss += loss.item()
  249. pbar.set_postfix({"InfoNCE Loss": f"{running_loss / (i + 1):.4f}"})
  250. if (epoch + 1) % EVAL_INTERVAL == 0:
  251. top1_acc, top5_acc = evaluate(model, val_dataset)
  252. print(f"[Epoch {epoch + 1} Eval(held-out val)] Retrieval Top-1: {top1_acc:.2%} | Top-5: {top5_acc:.2%}")
  253. if top1_acc > best_top1_acc:
  254. best_top1_acc = top1_acc
  255. raw_model = model.module if hasattr(model, "module") else model
  256. torch.save(raw_model.state_dict(), BEST_MODEL_PATH)
  257. print(f"保存新的最佳模型! Top-1: {top1_acc:.2%} -> {BEST_MODEL_PATH}")
  258. if __name__ == "__main__":
  259. main()