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

使用Huggingface Trainer训练时,DatasetDict预处理仅传入标签致KeyError

解决Hugging Face Trainer训练TorchVision模型时预处理函数丢失图像列的问题

问题原因

你遇到的KeyError是因为Hugging Face Trainer默认会根据模型的forward方法签名过滤数据集列:

  • TorchVision的MobileNetV3 forward方法仅接受位置参数(输入张量),没有对应的关键字参数(如pixel_values),因此Trainer无法识别img作为输入特征列,会自动丢弃该列,只保留label作为标签传递给预处理函数。
  • 同时,原代码未适配Transformers框架的标准列名(图像输入用pixel_values,标签用labels复数形式),进一步加剧了列识别问题。

解决方案

1. 统一数据集列名到Transformers规范

将原数据集的img重命名为pixel_values(Transformers图像模型标准输入列),label重命名为labels(Trainer默认标签列)。

2. 调整预处理函数适配批量处理

with_transform在批量加载数据时会传入包含列表的字典,预处理函数需要处理这种批量场景。

3. 包装TorchVision模型适配Trainer要求

Trainer期望模型接受关键字参数输入,并返回包含loss和logits的字典,因此需要对MobileNetV3进行简单包装,同时修改分类头适配CIFAR10的10分类任务。

修改后的完整代码

import torch
import numpy as np
from torchvision.models import mobilenet_v3_small, MobileNet_V3_Small_Weights
from transformers import TrainingArguments, Trainer
import evaluate
from datasets import load_dataset

# 加载预训练权重和预处理管道
weights = MobileNet_V3_Small_Weights.IMAGENET1K_V1
preprocess = weights.transforms()

# 加载数据集并统一列名
raw_data = load_dataset("cifar10")
raw_data = raw_data.rename_column("img", "pixel_values")
raw_data = raw_data.rename_column("label", "labels")

# 定义适配批量的预处理函数
def transform(examples):
    examples["pixel_values"] = [preprocess(img) for img in examples["pixel_values"]]
    return examples

data = raw_data.with_transform(transform)

# 包装MobileNetV3以适配Trainer
class MobileNetV3Wrapper(torch.nn.Module):
    def __init__(self, base_model):
        super().__init__()
        self.base_model = base_model
        # 修改分类头适配CIFAR10的10类(原模型为ImageNet 1000类)
        num_features = self.base_model.classifier[3].in_features
        self.base_model.classifier[3] = torch.nn.Linear(num_features, 10)
    
    def forward(self, pixel_values, labels=None):
        logits = self.base_model(pixel_values)
        if labels is not None:
            loss = torch.nn.functional.cross_entropy(logits, labels)
            return {"loss": loss, "logits": logits}
        return {"logits": logits}

# 初始化模型
base_model = mobilenet_v3_small(weights=weights)
model = MobileNetV3Wrapper(base_model)

# 训练参数与指标
training_args = TrainingArguments(
    output_dir="test_trainer",
    evaluation_strategy="epoch",
    per_device_train_batch_size=32,
    per_device_eval_batch_size=32,
    num_train_epochs=3,
    logging_dir="./logs",
)

accuracy = evaluate.load("accuracy")

def compute_metrics(eval_pred):
    logits, labels = eval_pred
    predictions = np.argmax(logits, axis=-1)
    return accuracy.compute(predictions=predictions, references=labels)

# 初始化Trainer并开始训练
trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=data["train"].select(range(5000)),
    eval_dataset=data["test"].select(range(1000)),
    compute_metrics=compute_metrics,
)

trainer.train()

关键说明

  • 列名适配:pixel_values和labels是Transformers框架约定的标准列名,Trainer会自动识别这些列作为输入和标签,不会被过滤。
  • 模型包装:包装后的模型接受pixel_values关键字参数,同时返回Trainer需要的loss和logits,确保训练流程正常运行。
  • 分类头修改:原MobileNetV3是为ImageNet 1000类设计的,必须修改最后一层线性层适配CIFAR10的10类,否则会出现维度不匹配错误。

内容的提问来源于stack exchange,提问作者Thibaut B.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 14:05:30