ViTForImageClassification多标签分类数据集与训练报错修复求助
多标签ViT分类系统改造报错解决
问题概述
将单分类ViTForImageClassification模型改造为多标签分类时,出现RuntimeError: result type Float can't be cast to the desired output type Long错误,核心原因是多标签分类的标签张量类型不匹配,同时数据集处理、评估逻辑存在多处不符合多标签要求的问题。
报错核心原因
多标签分类使用binary_cross_entropy_with_logits损失函数,要求标签张量为Float类型,但当前代码中生成的标签是Long类型;此外数据集标签收集、评估逻辑也存在错误。
代码修改方案
1. 修正数据集标签收集逻辑
原代码中标签收集存在错误,且无需手动维护类别列表,直接使用MultiLabelBinarizer生成的类别即可:
file_names = [] labels = [] for file in sorted((Path('/content/dataset').glob('*/*.*'))): folder = str(file).split('/')[-2].split('.')[0] label = folder.split('-') labels.append([x + '.class' for x in label]) file_names.append(str(file)) df = pd.DataFrame.from_dict({"image": file_names, "label": labels}) mlb = MultiLabelBinarizer() mlb_result = mlb.fit_transform(df['label']) df_final = pd.concat([df['image'], pd.DataFrame(mlb_result, columns=mlb.classes_)], axis=1) dataset = Dataset.from_pandas(df_final).cast_column("image", Image()) # 直接使用MultiLabelBinarizer生成的类别,避免手动维护错误 labels_list = list(mlb.classes_) label2id, id2label = {label:i for i,label in enumerate(labels_list)}, {i:label for i,label in enumerate(labels_list)} dataset = dataset.train_test_split(test_size=0.8, shuffle=True) train_data = dataset['train'] test_data = dataset['test']
2. 修正数据整理器(collate_fn)的标签类型
将标签张量转换为Float类型,匹配多标签损失函数的要求:
def collate_fn(examples): pixel_values = torch.stack([example["pixel_values"] for example in examples]) # 提取所有标签列 label_cols = [col for col in examples[0].keys() if col not in ['image', 'pixel_values']] # 转换为Float类型张量 labels = torch.tensor([[example[col] for col in label_cols] for example in examples], dtype=torch.float) return {"pixel_values": pixel_values, "labels": labels}
3. 修正评估函数逻辑
多标签分类不能使用argmax(单分类逻辑),需对模型输出做sigmoid后取阈值判断正负类,同时使用多标签适用的评估指标:
# 加载多标签适用的F1指标 metric = evaluate.load("f1", average="micro") def compute_metrics(eval_pred): predictions, label_ids = eval_pred # 对预测结果做sigmoid,以0.5为阈值生成二值标签 predicted_labels = torch.sigmoid(torch.tensor(predictions)) > 0.5 # 计算多标签F1分数 return metric.compute(predictions=predicted_labels.numpy(), references=label_ids, average="micro")
4. 模型初始化保持正确
原模型初始化逻辑无误,确保指定problem_type="multi_label_classification":
model = ViTForImageClassification.from_pretrained( model_str, num_labels=len(labels_list), problem_type="multi_label_classification" ) model.config.id2label = id2label model.config.label2id = label2id
验证修改
运行训练前可先验证标签类型:
# 取一个样本batch检查标签类型 sample_batch = collate_fn(train_data[:2]) print("Labels dtype:", sample_batch['labels'].dtype) # 输出应为torch.float32
确认类型正确后,执行trainer.evaluate()和trainer.train()即可解决报错。
内容的提问来源于stack exchange,提问作者Kayer
相关产品推荐
相关产品推荐

