train_dinov2_layer3_bottom.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333
  1. # -*- coding: utf-8 -*-
  2. """
  3. train_dinov2_layer3_bottom.py - Layer3 背景/版本消歧模型训练脚本
  4. 对齐 A4(train_dinov2_pokemonCN2_.py)的训练框架,但做了以下改动(详见
  5. D:/顾工交接/wzj/data_all/卡牌区域/球闪/Layer3专用背景分类模型训练方案.md):
  6. 1. backbone: facebook/dinov2-base -> facebook/dinov2-large
  7. 2. 冻结策略: 冻结前6/12层 -> 冻结前18/24层(大模型+小数据集,更保守)
  8. 3. 输入: 整卡392x392 -> 下半区392(宽)x196(高)(只训背景/版本消歧,不训物种识别)
  9. 4. 困难负样本来源: 文件夹名解析"卡名,卡号" -> 直接读 train_groups.json/val_groups.json
  10. (来自 find_confusable_pairs_v2.py 实测挖矿结果,group内的card_id互相是DINOv2现在会混淆的对象)
  11. 5. 数据增强: 去掉 RandomGlareA / RandomSunFlare(那是让模型对反光免疫,方向相反)
  12. 也去掉 MotionBlur/GaussianBlur/CoarseDropout(会抹掉反光纹理细节)
  13. 只保留:轻度亮度对比度、轻度旋转、轻度透视、轻度噪声/压缩伪影
  14. 6. 新增真正的held-out验证集(val_groups.json,46组/116张),而非训练集抽样自测
  15. 用法(pytorch环境,GPU):
  16. CUDA_VISIBLE_DEVICES=0 ~/miniconda3/envs/pytorch/bin/python train_dinov2_layer3_bottom.py
  17. """
  18. import os
  19. os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
  20. import json
  21. import random
  22. from collections import defaultdict
  23. import cv2
  24. import numpy as np
  25. import torch
  26. import torch.nn as nn
  27. import torch.nn.functional as F
  28. import torch.optim as optim
  29. import albumentations as A
  30. from albumentations.pytorch import ToTensorV2
  31. from torch.utils.data import BatchSampler, DataLoader, Dataset
  32. from tqdm import tqdm
  33. from transformers import Dinov2Model
  34. # ================= 配置区域 =================
  35. DATA_ROOT = r"/home/user/顾工交接/wzj/layer3_data"
  36. IMG_DIR = os.path.join(DATA_ROOT, "images")
  37. TRAIN_GROUPS_JSON = os.path.join(DATA_ROOT, "train_groups.json")
  38. VAL_GROUPS_JSON = os.path.join(DATA_ROOT, "val_groups.json")
  39. SAVE_DIR = r"/home/user/顾工交接/wzj/layer3_bg_model_output"
  40. BEST_MODEL_PATH = os.path.join(SAVE_DIR, "best_layer3_bottom_model.pth")
  41. HF_MODEL_ID = "facebook/dinov2-large"
  42. IMG_HEIGHT = 196 # 下半区裁剪高度 (392 letterbox 的下50%)
  43. IMG_WIDTH = 392 # 下半区裁剪宽度 (全宽)
  44. FREEZE_BLOCKS = 18 # dinov2-large 共24层,冻结前18层,只训后6层(比A4的12层冻一半更保守)
  45. BATCH_SIZE = 8
  46. EPOCHS = 150
  47. LEARNING_RATE = 1e-5 # 比A4的2e-5更低:大模型+小数据集,降低过拟合/震荡风险
  48. WEIGHT_DECAY = 1e-4
  49. TEMPERATURE = 0.07
  50. EVAL_INTERVAL = 1
  51. HARD_NEGATIVE_PROB = 0.5 # 数据集本身几乎全是困难负样本组,比例可以比A4(0.35)更高
  52. HARD_GROUP_SIZE = 6
  53. DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  54. # ================= 数据增强 =================
  55. # 注意:不用 RandomGlareA / RandomSunFlare(那是A4用来让模型对反光"免疫"的,
  56. # Layer3 的目标恰好相反——要对反光纹理保持敏感),也不用 MotionBlur/GaussianBlur/
  57. # CoarseDropout(会抹掉我们想学习的反光细节纹理)。
  58. transform_clean = A.Compose([
  59. A.Resize(IMG_HEIGHT, IMG_WIDTH, interpolation=cv2.INTER_CUBIC),
  60. A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
  61. ToTensorV2(),
  62. ])
  63. transform_aug = A.Compose([
  64. A.Resize(IMG_HEIGHT, IMG_WIDTH, interpolation=cv2.INTER_CUBIC),
  65. A.Perspective(scale=(0.02, 0.05), p=0.4),
  66. A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.5),
  67. A.RandomGamma(gamma_limit=(85, 115), p=0.4),
  68. A.GaussNoise(std_range=(0.05, 0.15), p=0.3), # 模拟相机ISO噪点
  69. A.ImageCompression(quality_range=(60, 95), p=0.3), # 模拟JPEG压缩伪影
  70. A.SafeRotate(limit=5, border_mode=cv2.BORDER_CONSTANT, fill=0, p=0.4),
  71. A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
  72. ToTensorV2(),
  73. ])
  74. # ================= 1. 数据集:直接读挖矿产出的分组JSON =================
  75. class MiningGroupDataset(Dataset):
  76. """
  77. groups_json: [{"group_id": int, "card_ids": [...], ...}, ...]
  78. 每个 card_id 对应 IMG_DIR/{card_id}.jpg 一张背景裁剪图。
  79. label = 该 card_id 的唯一序号(用于检索评估,判断"能不能认回自己")
  80. group_key = group_id(用于 HardNegativeBatchSampler,组内互为困难负样本)
  81. """
  82. def __init__(self, groups_json_path, img_dir, is_train):
  83. with open(groups_json_path, "r", encoding="utf-8") as f:
  84. groups = json.load(f)
  85. self.samples = [] # (path, label_idx, group_key)
  86. label_idx = 0
  87. for g in groups:
  88. group_key = g["group_id"]
  89. for cid in g["card_ids"]:
  90. path = os.path.join(img_dir, f"{cid}.jpg")
  91. if not os.path.exists(path):
  92. continue
  93. self.samples.append((path, label_idx, group_key))
  94. label_idx += 1
  95. self.group_to_indices = defaultdict(list)
  96. for idx, (_, _, group_key) in enumerate(self.samples):
  97. self.group_to_indices[group_key].append(idx)
  98. self.valid_hard_groups = [g for g, idxs in self.group_to_indices.items() if len(idxs) >= 2]
  99. tag = "train" if is_train else "val"
  100. print(f"[{tag}] samples={len(self.samples)} hard_negative_groups(>=2)={len(self.valid_hard_groups)}")
  101. def __len__(self):
  102. return len(self.samples)
  103. def __getitem__(self, idx):
  104. path, label, group_key = self.samples[idx]
  105. img = cv2.imread(path)
  106. if img is None:
  107. img = np.zeros((IMG_HEIGHT, IMG_WIDTH, 3), dtype=np.uint8)
  108. else:
  109. img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
  110. img_clean = transform_clean(image=img)["image"]
  111. img_aug = transform_aug(image=img)["image"]
  112. return img_clean, img_aug, label
  113. # ================= 2. 困难负样本 BatchSampler(原样复用 A4 逻辑)=================
  114. class HardNegativeBatchSampler(BatchSampler):
  115. def __init__(self, dataset, batch_size, hard_negative_prob=0.35, hard_group_size=8, drop_last=True):
  116. self.dataset = dataset
  117. self.batch_size = batch_size
  118. self.hard_negative_prob = hard_negative_prob
  119. self.hard_group_size = min(hard_group_size, batch_size)
  120. self.drop_last = drop_last
  121. self.num_samples = len(dataset)
  122. self.all_indices = list(range(self.num_samples))
  123. def __iter__(self):
  124. unused = set(self.all_indices)
  125. while len(unused) >= self.batch_size:
  126. batch = []
  127. if len(self.dataset.valid_hard_groups) > 0 and random.random() < self.hard_negative_prob:
  128. candidate_groups = []
  129. for g in self.dataset.valid_hard_groups:
  130. available = [i for i in self.dataset.group_to_indices[g] if i in unused]
  131. if len(available) >= 2:
  132. candidate_groups.append((g, available))
  133. if candidate_groups:
  134. g, available = random.choice(candidate_groups)
  135. take_n = min(len(available), self.hard_group_size, self.batch_size)
  136. chosen_hard = random.sample(available, take_n)
  137. batch.extend(chosen_hard)
  138. for i in chosen_hard:
  139. unused.remove(i)
  140. remain = self.batch_size - len(batch)
  141. if remain > 0:
  142. if len(unused) < remain:
  143. if self.drop_last:
  144. break
  145. batch.extend(list(unused))
  146. unused.clear()
  147. else:
  148. chosen_rand = random.sample(list(unused), remain)
  149. batch.extend(chosen_rand)
  150. for i in chosen_rand:
  151. unused.remove(i)
  152. if len(batch) == self.batch_size:
  153. random.shuffle(batch)
  154. yield batch
  155. if not self.drop_last and unused:
  156. yield list(unused)
  157. def __len__(self):
  158. if self.drop_last:
  159. return self.num_samples // self.batch_size
  160. return (self.num_samples + self.batch_size - 1) // self.batch_size
  161. # ================= 3. 模型:DINOv2-large,冻结前18层 =================
  162. class Dinov2Layer3Model(nn.Module):
  163. def __init__(self, freeze_blocks=FREEZE_BLOCKS):
  164. super().__init__()
  165. print(f"Loading {HF_MODEL_ID} backbone ...")
  166. self.backbone = Dinov2Model.from_pretrained(HF_MODEL_ID)
  167. for param in self.backbone.parameters():
  168. param.requires_grad = False
  169. total_layers = len(self.backbone.encoder.layer)
  170. print(f"Backbone total layers: {total_layers}, freezing first {freeze_blocks}")
  171. for i in range(freeze_blocks, total_layers):
  172. for param in self.backbone.encoder.layer[i].parameters():
  173. param.requires_grad = True
  174. for param in self.backbone.layernorm.parameters():
  175. param.requires_grad = True
  176. def forward(self, x):
  177. # 输入非正方形(392x196),需要开启位置编码插值
  178. outputs = self.backbone(x, interpolate_pos_encoding=True)
  179. cls_token = outputs.last_hidden_state[:, 0, :]
  180. return F.normalize(cls_token, p=2, dim=1)
  181. # ================= 4. InfoNCE 对比损失(原样复用 A4)=================
  182. def contrastive_loss(feat_clean, feat_aug, temperature=TEMPERATURE):
  183. batch_size = feat_clean.size(0)
  184. logits = torch.matmul(feat_clean, feat_aug.T) / temperature
  185. labels = torch.arange(batch_size).to(feat_clean.device)
  186. loss_1 = F.cross_entropy(logits, labels)
  187. loss_2 = F.cross_entropy(logits.T, labels)
  188. return (loss_1 + loss_2) / 2.0
  189. # ================= 5. 检索式评估(在真正held-out的val集上)=================
  190. def evaluate(model, val_dataset):
  191. model.eval()
  192. loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4)
  193. gallery_feats, gallery_labels = [], []
  194. query_feats, query_labels = [], []
  195. with torch.no_grad():
  196. for imgs_c, imgs_a, lbls in loader:
  197. imgs_c, imgs_a = imgs_c.to(DEVICE), imgs_a.to(DEVICE)
  198. with torch.amp.autocast("cuda"):
  199. f_c = model(imgs_c)
  200. f_a = model(imgs_a)
  201. gallery_feats.append(f_c.cpu())
  202. gallery_labels.append(lbls)
  203. query_feats.append(f_a.cpu())
  204. query_labels.append(lbls)
  205. gallery_feats = torch.cat(gallery_feats, dim=0)
  206. gallery_labels = torch.cat(gallery_labels, dim=0)
  207. query_feats = torch.cat(query_feats, dim=0)
  208. query_labels = torch.cat(query_labels, dim=0)
  209. sim_matrix = torch.matmul(query_feats, gallery_feats.T)
  210. k = min(5, gallery_feats.size(0))
  211. topk_idx = sim_matrix.topk(k, dim=1).indices
  212. correct_1, correct_5, total = 0, 0, query_labels.size(0)
  213. for idx in range(total):
  214. q_label = query_labels[idx]
  215. retrieved = gallery_labels[topk_idx[idx]]
  216. if q_label == retrieved[0]:
  217. correct_1 += 1
  218. if q_label in retrieved:
  219. correct_5 += 1
  220. return correct_1 / total, correct_5 / total
  221. # ================= 6. 主流程 =================
  222. def main():
  223. os.makedirs(SAVE_DIR, exist_ok=True)
  224. train_dataset = MiningGroupDataset(TRAIN_GROUPS_JSON, IMG_DIR, is_train=True)
  225. val_dataset = MiningGroupDataset(VAL_GROUPS_JSON, IMG_DIR, is_train=False)
  226. batch_sampler = HardNegativeBatchSampler(
  227. dataset=train_dataset,
  228. batch_size=BATCH_SIZE,
  229. hard_negative_prob=HARD_NEGATIVE_PROB,
  230. hard_group_size=HARD_GROUP_SIZE,
  231. drop_last=True,
  232. )
  233. train_loader = DataLoader(train_dataset, batch_sampler=batch_sampler, num_workers=8, pin_memory=True)
  234. model = Dinov2Layer3Model().to(DEVICE)
  235. if torch.cuda.device_count() > 1:
  236. model = nn.DataParallel(model)
  237. optimizer = optim.AdamW(
  238. filter(lambda p: p.requires_grad, model.parameters()),
  239. lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY,
  240. )
  241. scaler = torch.amp.GradScaler("cuda")
  242. best_top1_acc = 0.0
  243. print("开始 Layer3 (下半区/背景版本消歧) 对比学习训练 ...")
  244. for epoch in range(EPOCHS):
  245. model.train()
  246. pbar = tqdm(train_loader, desc=f"Epoch {epoch + 1}/{EPOCHS} [Train]")
  247. running_loss = 0.0
  248. for i, (imgs_clean, imgs_aug, _) in enumerate(pbar):
  249. imgs_clean = imgs_clean.to(DEVICE)
  250. imgs_aug = imgs_aug.to(DEVICE)
  251. optimizer.zero_grad()
  252. with torch.amp.autocast("cuda"):
  253. feat_clean = model(imgs_clean)
  254. feat_aug = model(imgs_aug)
  255. loss = contrastive_loss(feat_clean, feat_aug, TEMPERATURE)
  256. scaler.scale(loss).backward()
  257. scaler.step(optimizer)
  258. scaler.update()
  259. running_loss += loss.item()
  260. pbar.set_postfix({"InfoNCE Loss": f"{running_loss / (i + 1):.4f}"})
  261. if (epoch + 1) % EVAL_INTERVAL == 0:
  262. top1_acc, top5_acc = evaluate(model, val_dataset)
  263. print(f"[Epoch {epoch + 1} Eval(held-out val)] Retrieval Top-1: {top1_acc:.2%} | Top-5: {top5_acc:.2%}")
  264. if top1_acc > best_top1_acc:
  265. best_top1_acc = top1_acc
  266. raw_model = model.module if hasattr(model, "module") else model
  267. torch.save(raw_model.state_dict(), BEST_MODEL_PATH)
  268. print(f"保存新的最佳模型! Top-1: {top1_acc:.2%} -> {BEST_MODEL_PATH}")
  269. if __name__ == "__main__":
  270. main()