# -*- coding: utf-8 -*- """ train_dinov2_layer3_bottom.py - Layer3 背景/版本消歧模型训练脚本 对齐 A4(train_dinov2_pokemonCN2_.py)的训练框架,但做了以下改动(详见 D:/顾工交接/wzj/data_all/卡牌区域/球闪/Layer3专用背景分类模型训练方案.md): 1. backbone: facebook/dinov2-base -> facebook/dinov2-large 2. 冻结策略: 冻结前6/12层 -> 冻结前18/24层(大模型+小数据集,更保守) 3. 输入: 整卡392x392 -> 下半区392(宽)x196(高)(只训背景/版本消歧,不训物种识别) 4. 困难负样本来源: 文件夹名解析"卡名,卡号" -> 直接读 train_groups.json/val_groups.json (来自 find_confusable_pairs_v2.py 实测挖矿结果,group内的card_id互相是DINOv2现在会混淆的对象) 5. 数据增强: 去掉 RandomGlareA / RandomSunFlare(那是让模型对反光免疫,方向相反) 也去掉 MotionBlur/GaussianBlur/CoarseDropout(会抹掉反光纹理细节) 只保留:轻度亮度对比度、轻度旋转、轻度透视、轻度噪声/压缩伪影 6. 新增真正的held-out验证集(val_groups.json,46组/116张),而非训练集抽样自测 用法(pytorch环境,GPU): CUDA_VISIBLE_DEVICES=0 ~/miniconda3/envs/pytorch/bin/python train_dinov2_layer3_bottom.py """ import os os.environ["HF_ENDPOINT"] = "https://hf-mirror.com" import json import random from collections import defaultdict import cv2 import numpy as np import torch import torch.nn as nn import torch.nn.functional as F import torch.optim as optim import albumentations as A from albumentations.pytorch import ToTensorV2 from torch.utils.data import BatchSampler, DataLoader, Dataset from tqdm import tqdm from transformers import Dinov2Model # ================= 配置区域 ================= DATA_ROOT = r"/home/user/顾工交接/wzj/layer3_data" IMG_DIR = os.path.join(DATA_ROOT, "images") TRAIN_GROUPS_JSON = os.path.join(DATA_ROOT, "train_groups.json") VAL_GROUPS_JSON = os.path.join(DATA_ROOT, "val_groups.json") SAVE_DIR = r"/home/user/顾工交接/wzj/layer3_bg_model_output" BEST_MODEL_PATH = os.path.join(SAVE_DIR, "best_layer3_bottom_model.pth") HF_MODEL_ID = "facebook/dinov2-large" IMG_HEIGHT = 196 # 下半区裁剪高度 (392 letterbox 的下50%) IMG_WIDTH = 392 # 下半区裁剪宽度 (全宽) FREEZE_BLOCKS = 18 # dinov2-large 共24层,冻结前18层,只训后6层(比A4的12层冻一半更保守) BATCH_SIZE = 8 EPOCHS = 150 LEARNING_RATE = 1e-5 # 比A4的2e-5更低:大模型+小数据集,降低过拟合/震荡风险 WEIGHT_DECAY = 1e-4 TEMPERATURE = 0.07 EVAL_INTERVAL = 1 HARD_NEGATIVE_PROB = 0.5 # 数据集本身几乎全是困难负样本组,比例可以比A4(0.35)更高 HARD_GROUP_SIZE = 6 DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") # ================= 数据增强 ================= # 注意:不用 RandomGlareA / RandomSunFlare(那是A4用来让模型对反光"免疫"的, # Layer3 的目标恰好相反——要对反光纹理保持敏感),也不用 MotionBlur/GaussianBlur/ # CoarseDropout(会抹掉我们想学习的反光细节纹理)。 transform_clean = A.Compose([ A.Resize(IMG_HEIGHT, IMG_WIDTH, interpolation=cv2.INTER_CUBIC), A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ToTensorV2(), ]) transform_aug = A.Compose([ A.Resize(IMG_HEIGHT, IMG_WIDTH, interpolation=cv2.INTER_CUBIC), A.Perspective(scale=(0.02, 0.05), p=0.4), A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.5), A.RandomGamma(gamma_limit=(85, 115), p=0.4), A.GaussNoise(std_range=(0.05, 0.15), p=0.3), # 模拟相机ISO噪点 A.ImageCompression(quality_range=(60, 95), p=0.3), # 模拟JPEG压缩伪影 A.SafeRotate(limit=5, border_mode=cv2.BORDER_CONSTANT, fill=0, p=0.4), A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ToTensorV2(), ]) # ================= 1. 数据集:直接读挖矿产出的分组JSON ================= class MiningGroupDataset(Dataset): """ groups_json: [{"group_id": int, "card_ids": [...], ...}, ...] 每个 card_id 对应 IMG_DIR/{card_id}.jpg 一张背景裁剪图。 label = 该 card_id 的唯一序号(用于检索评估,判断"能不能认回自己") group_key = group_id(用于 HardNegativeBatchSampler,组内互为困难负样本) """ def __init__(self, groups_json_path, img_dir, is_train): with open(groups_json_path, "r", encoding="utf-8") as f: groups = json.load(f) self.samples = [] # (path, label_idx, group_key) label_idx = 0 for g in groups: group_key = g["group_id"] for cid in g["card_ids"]: path = os.path.join(img_dir, f"{cid}.jpg") if not os.path.exists(path): continue self.samples.append((path, label_idx, group_key)) label_idx += 1 self.group_to_indices = defaultdict(list) for idx, (_, _, group_key) in enumerate(self.samples): self.group_to_indices[group_key].append(idx) self.valid_hard_groups = [g for g, idxs in self.group_to_indices.items() if len(idxs) >= 2] tag = "train" if is_train else "val" print(f"[{tag}] samples={len(self.samples)} hard_negative_groups(>=2)={len(self.valid_hard_groups)}") def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label, group_key = self.samples[idx] img = cv2.imread(path) if img is None: img = np.zeros((IMG_HEIGHT, IMG_WIDTH, 3), dtype=np.uint8) else: img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img_clean = transform_clean(image=img)["image"] img_aug = transform_aug(image=img)["image"] return img_clean, img_aug, label # ================= 2. 困难负样本 BatchSampler(原样复用 A4 逻辑)================= class HardNegativeBatchSampler(BatchSampler): def __init__(self, dataset, batch_size, hard_negative_prob=0.35, hard_group_size=8, drop_last=True): self.dataset = dataset self.batch_size = batch_size self.hard_negative_prob = hard_negative_prob self.hard_group_size = min(hard_group_size, batch_size) self.drop_last = drop_last self.num_samples = len(dataset) self.all_indices = list(range(self.num_samples)) def __iter__(self): unused = set(self.all_indices) while len(unused) >= self.batch_size: batch = [] if len(self.dataset.valid_hard_groups) > 0 and random.random() < self.hard_negative_prob: candidate_groups = [] for g in self.dataset.valid_hard_groups: available = [i for i in self.dataset.group_to_indices[g] if i in unused] if len(available) >= 2: candidate_groups.append((g, available)) if candidate_groups: g, available = random.choice(candidate_groups) take_n = min(len(available), self.hard_group_size, self.batch_size) chosen_hard = random.sample(available, take_n) batch.extend(chosen_hard) for i in chosen_hard: unused.remove(i) remain = self.batch_size - len(batch) if remain > 0: if len(unused) < remain: if self.drop_last: break batch.extend(list(unused)) unused.clear() else: chosen_rand = random.sample(list(unused), remain) batch.extend(chosen_rand) for i in chosen_rand: unused.remove(i) if len(batch) == self.batch_size: random.shuffle(batch) yield batch if not self.drop_last and unused: yield list(unused) def __len__(self): if self.drop_last: return self.num_samples // self.batch_size return (self.num_samples + self.batch_size - 1) // self.batch_size # ================= 3. 模型:DINOv2-large,冻结前18层 ================= class Dinov2Layer3Model(nn.Module): def __init__(self, freeze_blocks=FREEZE_BLOCKS): super().__init__() print(f"Loading {HF_MODEL_ID} backbone ...") self.backbone = Dinov2Model.from_pretrained(HF_MODEL_ID) for param in self.backbone.parameters(): param.requires_grad = False total_layers = len(self.backbone.encoder.layer) print(f"Backbone total layers: {total_layers}, freezing first {freeze_blocks}") for i in range(freeze_blocks, total_layers): for param in self.backbone.encoder.layer[i].parameters(): param.requires_grad = True for param in self.backbone.layernorm.parameters(): param.requires_grad = True def forward(self, x): # 输入非正方形(392x196),需要开启位置编码插值 outputs = self.backbone(x, interpolate_pos_encoding=True) cls_token = outputs.last_hidden_state[:, 0, :] return F.normalize(cls_token, p=2, dim=1) # ================= 4. InfoNCE 对比损失(原样复用 A4)================= def contrastive_loss(feat_clean, feat_aug, temperature=TEMPERATURE): batch_size = feat_clean.size(0) logits = torch.matmul(feat_clean, feat_aug.T) / temperature labels = torch.arange(batch_size).to(feat_clean.device) loss_1 = F.cross_entropy(logits, labels) loss_2 = F.cross_entropy(logits.T, labels) return (loss_1 + loss_2) / 2.0 # ================= 5. 检索式评估(在真正held-out的val集上)================= def evaluate(model, val_dataset): model.eval() loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4) gallery_feats, gallery_labels = [], [] query_feats, query_labels = [], [] with torch.no_grad(): for imgs_c, imgs_a, lbls in loader: imgs_c, imgs_a = imgs_c.to(DEVICE), imgs_a.to(DEVICE) with torch.amp.autocast("cuda"): f_c = model(imgs_c) f_a = model(imgs_a) gallery_feats.append(f_c.cpu()) gallery_labels.append(lbls) query_feats.append(f_a.cpu()) query_labels.append(lbls) gallery_feats = torch.cat(gallery_feats, dim=0) gallery_labels = torch.cat(gallery_labels, dim=0) query_feats = torch.cat(query_feats, dim=0) query_labels = torch.cat(query_labels, dim=0) sim_matrix = torch.matmul(query_feats, gallery_feats.T) k = min(5, gallery_feats.size(0)) topk_idx = sim_matrix.topk(k, dim=1).indices correct_1, correct_5, total = 0, 0, query_labels.size(0) for idx in range(total): q_label = query_labels[idx] retrieved = gallery_labels[topk_idx[idx]] if q_label == retrieved[0]: correct_1 += 1 if q_label in retrieved: correct_5 += 1 return correct_1 / total, correct_5 / total # ================= 6. 主流程 ================= def main(): os.makedirs(SAVE_DIR, exist_ok=True) train_dataset = MiningGroupDataset(TRAIN_GROUPS_JSON, IMG_DIR, is_train=True) val_dataset = MiningGroupDataset(VAL_GROUPS_JSON, IMG_DIR, is_train=False) batch_sampler = HardNegativeBatchSampler( dataset=train_dataset, batch_size=BATCH_SIZE, hard_negative_prob=HARD_NEGATIVE_PROB, hard_group_size=HARD_GROUP_SIZE, drop_last=True, ) train_loader = DataLoader(train_dataset, batch_sampler=batch_sampler, num_workers=8, pin_memory=True) model = Dinov2Layer3Model().to(DEVICE) if torch.cuda.device_count() > 1: model = nn.DataParallel(model) optimizer = optim.AdamW( filter(lambda p: p.requires_grad, model.parameters()), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY, ) scaler = torch.amp.GradScaler("cuda") best_top1_acc = 0.0 print("开始 Layer3 (下半区/背景版本消歧) 对比学习训练 ...") for epoch in range(EPOCHS): model.train() pbar = tqdm(train_loader, desc=f"Epoch {epoch + 1}/{EPOCHS} [Train]") running_loss = 0.0 for i, (imgs_clean, imgs_aug, _) in enumerate(pbar): imgs_clean = imgs_clean.to(DEVICE) imgs_aug = imgs_aug.to(DEVICE) optimizer.zero_grad() with torch.amp.autocast("cuda"): feat_clean = model(imgs_clean) feat_aug = model(imgs_aug) loss = contrastive_loss(feat_clean, feat_aug, TEMPERATURE) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() running_loss += loss.item() pbar.set_postfix({"InfoNCE Loss": f"{running_loss / (i + 1):.4f}"}) if (epoch + 1) % EVAL_INTERVAL == 0: top1_acc, top5_acc = evaluate(model, val_dataset) print(f"[Epoch {epoch + 1} Eval(held-out val)] Retrieval Top-1: {top1_acc:.2%} | Top-5: {top5_acc:.2%}") if top1_acc > best_top1_acc: best_top1_acc = top1_acc raw_model = model.module if hasattr(model, "module") else model torch.save(raw_model.state_dict(), BEST_MODEL_PATH) print(f"保存新的最佳模型! Top-1: {top1_acc:.2%} -> {BEST_MODEL_PATH}") if __name__ == "__main__": main()