| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321 |
- # -*- coding: utf-8 -*-
- """
- train_dinov2_upper_half.py - 上半区精灵识别检索模型训练脚本
- 架构目标(双区):
- 上半区模型(本脚本) = 识别精灵/年龄段/姿势 (替代 cards_master_v2 card_name_ch 文本依赖)
- 下半区模型(Layer3) = 版本/特效(球闪等)消歧
- 两者都用 dinov2-large, 输入均为 392x196 (letterbox 392x392 的上/下半)。
- 对齐 Layer3(train_dinov2_layer3_bottom.py) 框架, 仅改:
- - 数据路径 -> upper_data_v0904 (全库上半区, organize_upper_dataset.py 产出)
- - 输入虽同为 392x196, 但取的是上半区(rows 0:196), Layer3 取下半区(rows 196:392)
- - group = card_name_ch + card_no (同卡同号不同实例作困难负样本, 与 A4 一致)
- 其余(冻结18层/InfoNCE/HardNegativeBatchSampler/增强策略)完全复用 Layer3。
- 用法(pytorch环境, GPU):
- CUDA_VISIBLE_DEVICES=0 ~/miniconda3/envs/pytorch/bin/python train_dinov2_upper_half.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/upper_data_v0904"
- 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/upper_model_output_v0904"
- BEST_MODEL_PATH = os.path.join(SAVE_DIR, "best_upper_half_model.pth")
- HF_MODEL_ID = "facebook/dinov2-large"
- IMG_HEIGHT = 196 # 上半区高度 (392 letterbox 的上50%, 与 Layer3 下半区对称)
- IMG_WIDTH = 392 # 全宽
- FREEZE_BLOCKS = 18 # dinov2-large 共24层, 冻结前18层(与 Layer3 一致)
- BATCH_SIZE = 224
- EPOCHS = 100
- LEARNING_RATE = 1e-5
- WEIGHT_DECAY = 1e-4
- TEMPERATURE = 0.07
- EVAL_INTERVAL = 1
- HARD_NEGATIVE_PROB = 0.5
- HARD_GROUP_SIZE = 6
- DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
- # ================= 数据增强 (与 Layer3 一致) =================
- 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),
- A.ImageCompression(quality_range=(60, 95), p=0.3),
- 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. 数据集 (原样复用 Layer3 MiningGroupDataset) =================
- class MiningGroupDataset(Dataset):
- 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 = []
- 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. HardNegativeBatchSampler (原样复用 A4/Layer3) =================
- 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 Dinov2UpperHalfModel(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):
- 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 =================
- 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. 检索评估 =================
- 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=16,
- pin_memory=True, persistent_workers=True, prefetch_factor=3)
- model = Dinov2UpperHalfModel().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("开始 上半区(精灵识别) 对比学习训练 ...")
- 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()
|