基于pytorch-metric-learning的LFW人脸对比模型训练失败排查请求
我正在尝试用pytorch-metric-learning库训练一个人脸对比模型,参考了官方MNIST示例代码,但试过不同的sampler、miner和loss函数后,模型始终无法有效训练。我是metric learning新手,恳请帮忙指出问题所在,谢谢!
训练代码
import pytorch_metric_learning as pml from torch.utils.data import DataLoader import torch.nn as nn import torch.nn.functional as F import torch.optim as optim from torchvision import datasets, transforms from torchvision.models import resnet18, ResNet18_Weights from torchvision.models import mobilenet_v3_large, MobileNet_V3_Large_Weights from torchvision.models import mobilenet_v3_small, MobileNet_V3_Small_Weights import torchvision import cv2 from pytorch_metric_learning import distances, losses, miners, reducers, testers,samplers from pytorch_metric_learning.utils.accuracy_calculator import AccuracyCalculator import numpy as np from torch.optim.lr_scheduler import StepLR class CFG: nepoch = 200 batch_size = 256 device="cuda" dataset_dir = "./tmp" sampler_m = 2 learning_rate = 0.001 weight_decay = 0.0005 step_size = 50 gamma = 0.5 num_workers = 16 class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.model = mobilenet_v3_small(weights = MobileNet_V3_Small_Weights.IMAGENET1K_V1) self.model.classifier = nn.Sequential( self.model.classifier[0], nn.Linear(1024,512, bias = True), nn.Dropout(p=0.25), #nn.ReLU(), #nn.BatchNorm1d(512), nn.Linear(512,256, bias = True), nn.Dropout(p=0.2), nn.Linear(256,128, bias = True) ) print(self.model) def forward(self,x): y = self.model(x) return y def get_all_embeddings(dataset, model): tester = testers.BaseTester() return tester.get_all_embeddings(dataset, model) def grid_image(batch): grid_imgs = torchvision.utils.make_grid(batch,normalize=True) grid_imgs = (grid_imgs * 255).cpu().numpy().astype(dtype = np.uint8) grid_imgs = np.transpose(grid_imgs,(1,2,0)) grid_imgs = cv2.cvtColor(grid_imgs,cv2.COLOR_RGB2BGR) return grid_imgs def step_iterio(model,optimizer,loss_func,mining_func,inputs,labes): data, labels = inputs.cuda(),labes.cuda() optimizer.zero_grad() embeddings = model(data) indices_tuple = mining_func(embeddings, labels) loss = loss_func(embeddings, labels, indices_tuple) loss.backward() optimizer.step() def train_lfw(): cfg = CFG() #cfg.batch_size = 16 # configure dataset t_train = transforms.Compose([ transforms.Resize((224,224)), transforms.ColorJitter(brightness=(0.5,1.5)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) t_test = transforms.Compose([ transforms.Resize((224,224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # dataset configure dataset_train = datasets.LFWPeople(cfg.dataset_dir,split = "train",transform = t_train) dataset_test = datasets.LFWPeople(cfg.dataset_dir,split = "test",transform = t_test) sampler = samplers.MPerClassSampler(dataset_train.targets,m=cfg.sampler_m,batch_size=cfg.batch_size,length_before_new_iter=10000) dataloader_train = DataLoader(dataset_train,batch_size=cfg.batch_size,sampler=sampler,num_workers=cfg.num_workers) #idx_to_class = {v:k for k,v in dataset_train.class_to_idx.items()} model = Net() model.cuda() optimizer = optim.Adam(model.parameters(), lr=cfg.learning_rate,weight_decay=cfg.weight_decay) lr_sheduler = StepLR(optimizer,step_size=cfg.step_size,gamma=cfg.gamma) #mining_func = miners.PairMarginMiner() #loss_func = losses.ContrastiveLoss(pos_margin=0.5,neg_margin=0.5) distance = distances.LpDistance() reducer = reducers.MeanReducer() loss_func = losses.TripletMarginLoss(margin=0.5, distance=distance, reducer=reducer) mining_func = miners.TripletMarginMiner(margin=0.5, distance=distance, type_of_triplets="all") accuracy_calculator = AccuracyCalculator(include=("precision_at_1",), k=1) for epoch in range(cfg.nepoch): model.train() for i_batch, (data, labels) in enumerate(dataloader_train): # print(labels) # for l in labels: # print(l," ",idx_to_class[l.item()]) data, labels = data.cuda(),labels.cuda() optimizer.zero_grad() embeddings = model(data) indices_tuple = mining_func(embeddings, labels) # print(indices_tuple[0]) # print(indices_tuple[1]) # print(indices_tuple[2]) # print(labels) loss = loss_func(embeddings, labels, indices_tuple) loss.backward() optimizer.step() if i_batch % 10 == 0: print( "Epoch {} Iteration {}: Loss = {}".format( epoch, i_batch, loss ) ) # grid_imgs = grid_image(data) # print(labels) # cv2.imshow("grid_imgs",grid_imgs) # cv2.waitKey() #break # testing lr_sheduler.step() model.eval() test_embeddings, test_labels = get_all_embeddings(dataset_test, model) #test_embeddings, test_labels = get_all_embeddings(dataset_train, model) test_labels = test_labels.squeeze(1) accuracies = accuracy_calculator.get_accuracy(test_embeddings, test_labels) print("Test set accuracy (Precision@1) = {}".format(accuracies["precision_at_1"])) if __name__ == "__main__": train_lfw()
特征提取头设计缺陷
你在MobileNetV3的classifier后堆叠了多层全连接,但注释掉了ReLU激活和BatchNorm1d,同时使用了高比例Dropout(0.25+0.2),这会严重限制特征表达能力,且Dropout在训练初期容易破坏特征学习。建议修改为:self.model.classifier = nn.Sequential( self.model.classifier[0], nn.Linear(1024,512, bias=True), nn.ReLU(), nn.BatchNorm1d(512), nn.Linear(512,256, bias=True), nn.ReLU(), nn.BatchNorm1d(256), nn.Linear(256,128, bias=True) )先去掉Dropout,待模型稳定后再考虑添加低比例(如0.1)的Dropout。
采样器与batch_size不兼容
MPerClassSampler要求batch_size必须是sampler_m * 类别数的整数倍。你的sampler_m=2、batch_size=256,意味着每个batch需要包含128个不同类别,但LFW训练集的人数(类别数)可能无法满足这个要求,导致采样异常。建议调小batch_size至64或128,确保batch_size % sampler_m == 0。Triplet Miner类型选择不当
使用type_of_triplets="all"会将所有可能的三元组纳入损失计算,其中大量易学习的三元组会稀释有效梯度,降低训练效率。建议改为挖掘难样本的模式:mining_func = miners.TripletMarginMiner(margin=0.5, distance=distance, type_of_triplets="semihard")或
"hard",聚焦难样本提升特征区分度。预训练模型未合理冻结
直接使用预训练MobileNetV3却未冻结底层特征层,初始学习率0.001会破坏预训练的通用特征。建议先冻结前几层,只训练新增的分类头:# 模型初始化后添加 for param in model.model.features[:-4].parameters(): param.requires_grad = False训练10-20轮后,解冻所有层并将学习率降至1e-5左右微调。
距离度量与损失参数不适配
人脸对比任务中,余弦相似度通常比Lp距离更有效,建议替换距离度量:distance = distances.CosineSimilarity()同时可将TripletMarginLoss的margin调整为0.3,更贴合人脸特征分布。
数据增强不足
当前仅使用ColorJitter,人脸任务需要针对性增强提升泛化能力:t_train = transforms.Compose([ transforms.Resize((224,224)), transforms.RandomHorizontalFlip(), transforms.RandomCrop(224, padding=10), transforms.ColorJitter(brightness=(0.5,1.5), contrast=0.2, saturation=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])评估指标不符合任务场景
LFW的标准任务是判断人脸对是否为同一人,而你使用的precision_at_1是检索任务的指标,无法准确反映模型性能。建议改用配对任务的评估方式,或实现LFW标准验证逻辑(计算准确率、ROC曲线等)。
内容的提问来源于stack exchange,提问作者Lucifer

