使用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.
相关产品推荐
相关产品推荐

