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

