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

使用torchvision.transform处理DVS128Gesture数据集遇类型错误的问询

问题:DVS128Gesture数据集使用torchvision.transform报错

报错信息:

img should be PIL Image. Got <class 'numpy.lib.npyio.NpzFile'>

用户代码:

import torch
import torchvision
from spikingjelly.datasets.dvs128_gesture import DVS128Gesture

train_data = DVS128Gesture(root_dir, train=True, data_type='event',
                            transform=torchvision.transforms.Compose([
                                torchvision.transforms.Resize(32),
                                torchvision.transforms.Normalize((0.0,), (0.8,)),
                                torchvision.transforms.ToTensor()
                            ]))
test_data = DVS128Gesture(root_dir, train=False, data_type='event',
                           transform=torchvision.transforms.Compose([
                               torchvision.transforms.Resize(32),
                               torchvision.transforms.Normalize((0.0,), (0.8,)),
                               torchvision.transforms.ToTensor()
                           ]))

train_loader = torch.utils.data.DataLoader(train_data, batch_size=bs, shuffle=True)
test_loader = torch.utils.data.DataLoader(test_data, batch_size=bs, shuffle=True)

examples = enumerate(test_loader)
batch_idx, (example_data, example_targets) = next(examples)
example_data.shape

此前处理MNIST正常的代码:

train_data = torchvision.datasets.MNIST(root_dir, train=True, download=True,
                            transform=torchvision.transforms.Compose([
                                torchvision.transforms.Resize(28),
                                torchvision.transforms.ToTensor(),
                                torchvision.transforms.Normalize((0.0,), (0.8,))
                            ]))
test_data = torchvision.datasets.MNIST(root_dir, train=False, download=True,
                           transform=torchvision.transforms.Compose([
                               torchvision.transforms.Resize(28),
                               torchvision.transforms.ToTensor(),
                               torchvision.transforms.Normalize((0.0,), (0.8,))
                           ]))

报错栈信息:

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
Cell In [10], line 2
      1 examples = enumerate(test_loader)
----> 2 batch_idx, (example_data, example_targets) = next(examples)
      3 example_data.shape

File p:\Programs\Anaconda3\lib\site-packages\torch\utils\data\dataloader.py:681, in _BaseDataLoaderIter.__next__(self)
    678 if self._sampler_iter is None:
    679     # TODO(https://github.com/pytorch/pytorch/issues/76750)
    680     self._reset()  # type: ignore[call-arg]
---> 681 data = self._next_data()
    682 self._num_yielded += 1
    683 if self._dataset_kind == _DatasetKind.Iterable and \
    684         self._IterableDataset_len_called is not None and \
    685         self._num_yielded > self._IterableDataset_len_called:

File p:\Programs\Anaconda3\lib\site-packages\torch\utils\data\dataloader.py:721, in _SingleProcessDataLoaderIter._next_data(self)
    719 def _next_data(self):
    720     index = self._next_index()  # may raise StopIteration
---> 721     data = self._dataset_fetcher.fetch(index)  # may raise StopIteration
    722     if self._pin_memory:
    723         data = _utils.pin_memory.pin_memory(data, self._pin_memory_device)

File p:\Programs\Anaconda3\lib\site-packages\torch\utils\data\_utils\fetch.py:49, in _MapDatasetFetcher.fetch(self, possibly_batched_index)
     47 def fetch(self, possibly_batched_index):
     48     if self.auto_collation:
---> 49         data = [self.dataset[idx] for idx in possibly_batched_index]
     50     else:
     51         data = self.dataset[possibly_batched_index]

File p:\Programs\Anaconda3\lib\site-packages\torch\utils\data\_utils\fetch.py:49, in <listcomp>(.0)
     47 def fetch(self, possibly_batched_index):
     48     if self.auto_collation:
---> 49         data = [self.dataset[idx] for idx in possibly_batched_index]
     50     else:
     51         data = self.dataset[possibly_batched_index]

File p:\Programs\Anaconda3\lib\site-packages\torchvision\datasets\folder.py:232, in DatasetFolder.__getitem__(self, index)
    230 sample = self.loader(path)
    231 if self.transform is not None:
---> 232     sample = self.transform(sample)
    233 if self.target_transform is not None:
    234     target = self.target_transform(target)

File p:\Programs\Anaconda3\lib\site-packages\torchvision\transforms\transforms.py:94, in Compose.__call__(self, img)
     92 def __call__(self, img):
     93     for t in self.transforms:
---> 94         img = t(img)
     95     return img

File p:\Programs\Anaconda3\lib\site-packages\torch\nn\modules\module.py:1130, in Module._call_impl(self, *input, **kwargs)
   1126 # If we don't have any hooks, we want to skip the rest of the logic in
   1127 # this function, and just call forward.
   1128 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks
   1129         or _global_forward_hooks or _global_forward_pre_hooks):
-> 1130     return forward_call(*input, **kwargs)
   1131 # Do not call functions when jit is used
   1132 full_backward_hooks, non_full_backward_hooks = [], []

File p:\Programs\Anaconda3\lib\site-packages\torchvision\transforms\transforms.py:349, in Resize.forward(self, img)
    341 def forward(self, img):
    342     """
    343     Args:
    344         img (PIL Image or Tensor): Image to be scaled.
   (...)
    347         PIL Image or Tensor: Rescaled image.
    348     """
---> 349     return F.resize(img, self.size, self.interpolation, self.max_size, self.antialias)

File p:\Programs\Anaconda3\lib\site-packages\torchvision\transforms\functional.py:430, in resize(img, size, interpolation, max_size, antialias)
    428         warnings.warn("Anti-alias option is always applied for PIL Image input. Argument antialias is ignored.")
    429     pil_interpolation = pil_modes_mapping[interpolation]
---> 430     return F_pil.resize(img, size=size, interpolation=pil_interpolation, max_size=max_size)
    432 return F_t.resize(img, size=size, interpolation=interpolation.value, max_size=max_size, antialias=antialias)

File p:\Programs\Anaconda3\lib\site-packages\torchvision\transforms\functional_pil.py:249, in resize(img, size, interpolation, max_size)
    240 @torch.jit.unused
    241 def resize(
    242     img: Image.Image,
   (...)
    245     max_size: Optional[int] = None,
    246 ) -> Image.Image:
    248     if not _is_pil_image(img):
---> 249         raise TypeError(f"img should be PIL Image. Got {type(img)}")
    250     if not (isinstance(size, int) or (isinstance(size, Sequence) and len(size) in (1, 2))):
    251         raise TypeError(f"Got inappropriate size arg: {size}")

TypeError: img should be PIL Image. Got <class 'numpy.lib.npyio.NpzFile'>

错误原因

  1. 数据类型不匹配:设置data_type='event'时,DVS128Gesture返回的是NpzFile对象(存储事件数据的压缩文件),而torchvision的Resize等变换仅支持PIL Image或Tensor类型。MNIST默认返回PIL Image,因此之前的代码可正常运行,二者数据类型本质不同。
  2. 变换顺序错误:即使数据类型正确,Normalize需在ToTensor之后执行——Normalize是对Tensor进行的操作;Resize虽可处理PIL Image或Tensor,但前提是输入类型符合要求。

解决方案

方案一:改用帧数据(直接适配torchvision变换)

若目标是处理帧数据,将data_type改为'frame',数据集会返回已处理好的帧Tensor,再配合正确的变换顺序即可:

import torch
import torchvision
from spikingjelly.datasets.dvs128_gesture import DVS128Gesture

train_data = DVS128Gesture(root_dir, train=True, data_type='frame',
                            transform=torchvision.transforms.Compose([
                                torchvision.transforms.Resize((32,32)),
                                torchvision.transforms.ToTensor(),
                                torchvision.transforms.Normalize((0.0,), (0.8,))
                            ]))
test_data = DVS128Gesture(root_dir, train=False, data_type='frame',
                           transform=torchvision.transforms.Compose([
                               torchvision.transforms.Resize((32,32)),
                               torchvision.transforms.ToTensor(),
                               torchvision.transforms.Normalize((0.0,), (0.8,))
                           ]))

方案二:自定义变换处理事件数据

若必须使用data_type='event',需先自定义变换提取NpzFile中的数据,转换为Tensor或PIL Image后再应用torchvision变换:

import torch
import torchvision
from spikingjelly.datasets.dvs128_gesture import DVS128Gesture
import numpy as np
from PIL import Image

# 自定义变换:从NpzFile中提取事件数据并转为二元帧
class ProcessDVSData:
    def __call__(self, npz_file):
        # 加载事件数据
        events = np.load(npz_file)
        # 构建128x128的二元帧示例(根据需求调整事件转帧逻辑)
        frame = np.zeros((128,128), dtype=np.float32)
        # 累加事件极性(示例逻辑,需根据任务调整)
        frame[events['y'], events['x']] += events['p']
        # 转为二元帧(大于0为1,否则为0)
        frame = (frame > 0).astype(np.float32)
        # 转为PIL Image
        return Image.fromarray(frame * 255).convert('L')

train_data = DVS128Gesture(root_dir, train=True, data_type='event',
                            transform=torchvision.transforms.Compose([
                                ProcessDVSData(),
                                torchvision.transforms.Resize((32,32)),
                                torchvision.transforms.ToTensor(),
                                torchvision.transforms.Normalize((0.0,), (0.8,))
                            ]))
test_data = DVS128Gesture(root_dir, train=False, data_type='event',
                           transform=torchvision.transforms.Compose([
                               ProcessDVSData(),
                               torchvision.transforms.Resize((32,32)),
                               torchvision.transforms.ToTensor(),
                               torchvision.transforms.Normalize((0.0,), (0.8,))
                           ]))

注意:事件转帧的逻辑需根据任务需求调整,上述代码为基础示例。

内容的提问来源于stack exchange,提问作者Roy Meijer

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 22:55:48