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

使用HuggingFace数据集应用PyTorch预训练权重变换时出错

问题:CIFAR100微调MobileNetV3时,HuggingFace Dataset with_transform预处理报错

在CIFAR100数据集上使用PyTorch+HuggingFace Datasets微调MobileNetV3模型时,调用.with_transform()将预训练权重的预处理变换应用到数据集出现错误:手动对单张图片执行预处理可正常运行,但通过数据集懒加载处理则触发类型错误。

可复现代码

import torch
from torchvision.models import MobileNet_V3_Small_Weights
from datasets import load_dataset
from matplotlib import pyplot as plt

weights = MobileNet_V3_Small_Weights.DEFAULT
preprocess = weights.transforms()

raw_data = load_dataset("cifar100")
data = raw_data.with_transform(preprocess)

raw_img = raw_data["train"][0]["img"]

fig, axes = plt.subplots(1, 3)

axes[0].imshow(raw_img)
axes[0].set_title("Raw image")

img = preprocess(raw_img).permute(1, 2, 0)    # 手动预处理图片可行
axes[1].imshow(img)
axes[1].set_title("Preprocessed image (manual)")

img = data["train"][0]["img"]                 # 从预处理数据集获取图片失败(懒加载)
axes[2].imshow(img.permute(1, 2, 0))
axes[2].set_title("Preprocessed image (dataset)")

plt.show()

报错信息

Traceback (most recent call last):
  File "C:\Users\thiba\OneDrive - McGill University\Internship\ECSE301\pytorch_test.py", line 23, in <module>
    img = data["train"][0]["img"]
          ~~~~~~~~~~~~~^^^
  File "C:\Python311\Lib\site-packages\datasets\arrow_dataset.py", line 2778, in __getitem__
    return self._getitem(key)
           ^^^^^^^^^^^^^^^^^^
  File "C:\Python311\Lib\site-packages\datasets\arrow_dataset.py", line 2763, in _getitem
    formatted_output = format_table(
                       ^^^^^^^^^^^^^
  File "C:\Python311\Lib\site-packages\datasets\formatting\formatting.py", line 624, in format_table
    return formatter(pa_table, query_type=query_type)
    return self.format_row(pa_table)
           ^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Python311\Lib\site-packages\datasets\formatting\formatting.py", line 480, in format_row
    formatted_batch = self.format_batch(pa_table)
                      ^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Python311\Lib\site-packages\datasets\formatting\formatting.py", line 510, in format_batch
    return self.transform(batch)
           ^^^^^^^^^^^^^^^^^^^^^
  File "C:\Python311\Lib\site-packages\torch\nn\modules\module.py", line 1501, in _call_impl
    return forward_call(*args, **kwargs)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Python311\Lib\site-packages\torchvision\transforms\_presets.py", line 58, in forward
    img = F.resize(img, self.resize_size, interpolation=self.interpolation, antialias=self.antialias)
          ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Python311\Lib\site-packages\torchvision\transforms\functional.py", line 476, in resize
    _, image_height, image_width = get_dimensions(img)
                                   ^^^^^^^^^^^^^^^^^^^
  File "C:\Python311\Lib\site-packages\torchvision\transforms\functional.py", line 78, in get_dimensions
    return F_pil.get_dimensions(img)
           ^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Python311\Lib\site-packages\torchvision\transforms\_functional_pil.py", line 31, in get_dimensions
    raise TypeError(f"Unexpected type {type(img)}")
TypeError: Unexpected type <class 'dict'>

原因分析

with_transform()会将**整个样本字典(包含img、label等字段)**传递给预处理函数,但torchvision的weights.transforms()返回的变换仅接受单张图片(PIL图像/Tensor)作为输入,因此会把字典当成输入触发类型错误。

解决方案

方法1:自定义批量处理函数(推荐,符合with_transform的批量处理逻辑)

编写一个处理函数,仅对样本中的img字段应用预处理,同时保留其他字段:

def apply_preprocess(batch):
    # 对批量中的每张图片应用预处理
    batch["img"] = [preprocess(img) for img in batch["img"]]
    return batch

# 替换原有的with_transform调用
data = raw_data.with_transform(apply_preprocess)

方法2:使用map()处理单样本

如果需要逐样本处理,可以用map()方法(默认逐样本处理):

data = raw_data.map(lambda sample: {"img": preprocess(sample["img"])})

修改后验证代码

import torch
from torchvision.models import MobileNet_V3_Small_Weights
from datasets import load_dataset
from matplotlib import pyplot as plt

weights = MobileNet_V3_Small_Weights.DEFAULT
preprocess = weights.transforms()

raw_data = load_dataset("cifar100")

# 使用自定义批量处理函数
def apply_preprocess(batch):
    batch["img"] = [preprocess(img) for img in batch["img"]]
    return batch

data = raw_data.with_transform(apply_preprocess)

raw_img = raw_data["train"][0]["img"]

fig, axes = plt.subplots(1, 3)

axes[0].imshow(raw_img)
axes[0].set_title("Raw image")

img = preprocess(raw_img).permute(1, 2, 0)
axes[1].imshow(img)
axes[1].set_title("Preprocessed image (manual)")

img = data["train"][0]["img"]
axes[2].imshow(img.permute(1, 2, 0))
axes[2].set_title("Preprocessed image (dataset)")

plt.show()

内容的提问来源于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.21 01:03:15