微调DeiT蒸馏图像分类模型时始终预测单一类别问题求助
我完全懂你现在的挫败感——折腾了这么久,不管是用平衡还是不平衡数据集,甚至试过用小数据集过拟合,模型却始终只预测0类,准确率刚好等于数据集里0类的占比,这确实让人摸不着头脑。咱们一步步拆解可能的问题,试试这些解决方案:
预处理不符合预训练模型的要求
你现在用自己定义的transform,但DeiT这类预训练模型有配套的DeiTImageProcessor,它包含了预训练时用到的归一化参数(特定的均值和方差),自定义的transform没做这个,会导致输入数据分布和预训练时完全不一致,模型根本没法正确提取特征。
替换成官方processor试试:processor = DeiTImageProcessor.from_pretrained(model_path) class myDataset(Dataset): def __init__(self, root_dir): self.root_dir = root_dir self.data = [] for label in os.listdir(root_dir): label_dir = os.path.join(root_dir, label) if os.path.isdir(label_dir): for file in os.listdir(label_dir): self.data.append((os.path.join(label_dir, file), int(label))) def __len__(self): return len(self.data) def __getitem__(self, idx): img_path, label = self.data[idx] image = Image.open(img_path).convert("RGB") # 确保图像是RGB格式 # 用processor处理图像 encoding = processor(images=image, return_tensors="pt") pixel_values = encoding['pixel_values'].squeeze() # 去除多余的batch维度 return pixel_values, label学习率设置过高,破坏预训练特征
你用的learning_rate = 0.01对于微调预训练模型来说实在太大了!预训练模型的参数已经经过大量数据学习到了通用特征,这么高的学习率会直接冲掉这些有用的特征,导致模型无法收敛,甚至输出完全随机(或者一直预测同一类)。
建议把学习率降到1e-5到1e-4之间,比如:learning_rate = 1e-5 optimizer = optim.Adam(model.parameters(), lr=learning_rate)还可以搭配学习率调度器,让学习率随训练逐步降低:
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.1) # 在训练循环里每个epoch后更新学习率 for epoch in range(num_epochs): # ... 训练代码 ... scheduler.step()分类头没有适配二分类任务
DeiT预训练模型的分类头是针对ImageNet的1000类设计的,而你的任务是二分类,直接用原分类头会导致输出维度(1000维)和任务需求(2维)不匹配,模型根本没法正确学习二分类的边界。
替换分类头的代码:model = DeiTForImageClassificationWithTeacher.from_pretrained(model_path) # 替换主分类头 num_features = model.classifier.in_features model.classifier = nn.Linear(num_features, 2).to(device) # 蒸馏版本还要替换distillation_classifier if hasattr(model, 'distillation_classifier'): model.distillation_classifier = nn.Linear(num_features, 2).to(device)训练过程中加入监控,定位问题
建议在训练每个epoch后,也检查一下训练集的准确率,如果训练集准确率也很低,那说明模型根本没学到东西,大概率是前面预处理、学习率或者分类头的问题;如果训练集准确率正常,测试集一直预测单一类,再考虑数据集分布或者验证集的问题。
先试试这几个方向,应该能解决模型一直预测单一类的问题。
备注:内容来源于stack exchange,提问作者shamilemir

