本地GPU训练模型遇TypeError:img应为PIL Image却得到dict类型
本地训练模型报错:
TypeError: img should be PIL Image. Got <class 'dict'> 问题场景
相同代码在Google Colab可正常运行,但本地GPU训练时,执行batch = next(iter(dataloader))触发错误。
运行代码
from datasets import load_dataset dataset = load_dataset("tglcourse/lsun_church_train", cache_dir='dataset') image_size = 256 channels = 3 batch_size = 1 from torchvision import transforms from torch.utils.data import DataLoader # 定义图像变换 transform = transforms.Compose([ # transforms.RandomHorizontalFlip(), transforms.Resize(image_size), transforms.CenterCrop(image_size), transforms.ToTensor(), transforms.Lambda(lambda t: (t * 2) - 1) ]) # 定义数据集变换函数 def transforms(examples): examples["pixel_values"] = [transform(image) for image in examples["image"]] del examples["image"] return examples transformed_dataset = dataset.with_transform(transforms).remove_columns("label") # 创建数据加载器 dataloader = DataLoader(transformed_dataset["train"], batch_size=batch_size, shuffle=True) # 获取批量数据 batch = next(iter(dataloader))
报错信息
简短报错:
TypeError: img should be PIL Image. Got <class 'dict'>
完整报错堆栈:
Traceback (most recent call last): File "diff_lsun_church.py", line 486, in <module> batch = next(iter(dataloader)) File "/home1/rishi_suman/.local/lib/python3.6/site-packages/torch/utils/data/dataloader.py", line 521, in __next__ data = self._next_data() File "/home1/rishi_suman/.local/lib/python3.6/site-packages/torch/utils/data/dataloader.py", line 561, in _next_data data = self._dataset_fetcher.fetch(index) # may raise StopIteration File "/home1/rishi_suman/.local/lib/python3.6/site-packages/torch/utils/data/_utils/fetch.py", line 49, in fetch data = [self.dataset[idx] for idx in possibly_batched_index] File "/home1/rishi_suman/.local/lib/python3.6/site-packages/torch/utils/data/_utils/fetch.py", line 49, in <listcomp> data = [self.dataset[idx] for idx in possibly_batched_index] File "/home1/rishi_suman/.local/lib/python3.6/site-packages/datasets/arrow_dataset.py", line 2166, in __getitem__ key, File "/home1/rishi_suman/.local/lib/python3.6/site-packages/datasets/arrow_dataset.py", line 2151, in _getitem pa_subtable, key, formatter=formatter, format_columns=format_columns, output_all_columns=output_all_columns File "/home1/rishi_suman/.local/lib/python3.6/site-packages/datasets/formatting/formatting.py", line 532, in format_table return formatter(pa_table, query_type=query_type) File "/home1/rishi_suman/.local/lib/python3.6/site-packages/datasets/formatting/formatting.py", line 281, in __call__ return self.format_row(pa_table) File "/home1/rishi_suman/.local/lib/python3.6/site-packages/datasets/formatting/formatting.py", line 387, in format_row formatted_batch = self.format_batch(pa_table) File "/home1/rishi_suman/.local/lib/python3.6/site-packages/datasets/formatting/formatting.py", line 418, in format_batch return self.transform(batch) File "diff_lsun_church.py", line 475, in transforms examples["pixel_values"] = [transform(image) for image in examples["image"]] File "diff_lsun_church.py", line 475, in <listcomp> examples["pixel_values"] = [transform(image) for image in examples["image"]] File "/home1/rishi_suman/.local/lib/python3.6/site-packages/torchvision/transforms/transforms.py", line 61, in __call__ img = t(img) File "/home1/rishi_suman/.local/lib/python3.6/site-packages/torch/nn/modules/module.py", line 1102, in _call_impl return forward_call(*input, **kwargs) File "/home1/rishi_suman/.local/lib/python3.6/site-packages/torchvision/transforms/transforms.py", line 304, in forward return F.resize(img, self.size, self.interpolation, self.max_size, self.antialias) File "/home1/rishi_suman/.local/lib/python3.6/site-packages/torchvision/transforms/functional.py", line 419, in resize return F_pil.resize(img, size=size, interpolation=pil_interpolation, max_size=max_size) File "/home1/rishi_suman/.local/lib/python3.6/site-packages/torchvision/transforms/functional_pil.py", line 233, in resize raise TypeError('img should be PIL Image. Got {}'.format(type(img))) TypeError: img should be PIL Image. Got <class 'dict'>
问题原因
本地安装的datasets库版本过旧,旧版本加载图片数据集时,examples["image"]中的元素是包含图片字节数据、路径等信息的字典,而非直接返回PIL Image对象;而Colab上的datasets版本较新,默认会自动将图片转换为PIL Image,因此代码能正常运行。
解决方案
方案1:升级datasets库到最新版本
执行以下命令升级:
pip install --upgrade datasets
升级后无需修改代码,load_dataset会自动返回PIL Image对象,代码可直接运行。
方案2:修改代码适配旧版本datasets
如果无法升级库,可在变换函数中手动将字典转换为PIL Image:
from datasets import load_dataset import io from PIL import Image dataset = load_dataset("tglcourse/lsun_church_train", cache_dir='dataset') image_size = 256 channels = 3 batch_size = 1 from torchvision import transforms from torch.utils.data import DataLoader # 定义图像变换 transform = transforms.Compose([ # transforms.RandomHorizontalFlip(), transforms.Resize(image_size), transforms.CenterCrop(image_size), transforms.ToTensor(), transforms.Lambda(lambda t: (t * 2) - 1) ]) # 定义数据集变换函数 def transforms(examples): # 从字典中提取字节数据并转换为PIL Image examples["pixel_values"] = [transform(Image.open(io.BytesIO(image['bytes']))) for image in examples["image"]] del examples["image"] return examples transformed_dataset = dataset.with_transform(transforms).remove_columns("label") # 创建数据加载器 dataloader = DataLoader(transformed_dataset["train"], batch_size=batch_size, shuffle=True) # 获取批量数据 batch = next(iter(dataloader))
内容的提问来源于stack exchange,提问作者Rishi Suman
相关产品推荐
相关产品推荐

