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

本地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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 09:58:11