| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333 |
- # -*- 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()
|