You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.30 10:50:55