使用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'>
错误原因
- 数据类型不匹配:设置
data_type='event'时,DVS128Gesture返回的是NpzFile对象(存储事件数据的压缩文件),而torchvision的Resize等变换仅支持PIL Image或Tensor类型。MNIST默认返回PIL Image,因此之前的代码可正常运行,二者数据类型本质不同。 - 变换顺序错误:即使数据类型正确,
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
相关产品推荐
相关产品推荐

