PyTorch DataLoader迭代数轮后抛出RuntimeError报错求助
问题:DataLoader抛出RuntimeError: Trying to resize storage that is not resizable
自定义Dataset代码
class Dataset(Dataset): 'Characterizes a dataset for PyTorch' def __init__(self, input_feature_paths, target_feature_folder) -> None: self.input_feature_paths = input_feature_paths self.target_feature_folder = target_feature_folder def __len__(self): #return sum(1 for _ in self.input_feature_paths) return len(self.input_feature_paths) def __getitem__(self, index) -> None: input_feature_path = self.input_feature_paths[index] input_feature = load(input_feature_path, map_location='cpu') target_feature_path = self.target_feature_folder / input_feature_path.parts[-1] target_feature = load(target_feature_path, map_location='cpu') return input_feature.to(dtype=torch.float64), target_feature.to(dtype=torch.float64)
错误栈信息
Traceback (most recent call last): File "student_audio_feature_extractor.py", line 178, in <module> train(dt, input_frame) File "student_audio_feature_extractor.py", line 164, in train model, train_loss = train_step(model, train_loader, optimizer, criterion) File "student_audio_feature_extractor.py", line 80, in train_step for input_feature, target_feature in train_loader: File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/utils/data/dataloader.py", line 628, in __next__ data = self._next_data() File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/utils/data/dataloader.py", line 1313, in _next_data return self._process_data(data) File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/utils/data/dataloader.py", line 1359, in _process_data data.reraise() File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/_utils.py", line 543, in reraise raise exception RuntimeError: Caught RuntimeError in DataLoader worker process 4. Original Traceback (most recent call last): File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/utils/data/_utils/worker.py", line 302, in _worker_loop data = fetcher.fetch(index) File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/utils/data/_utils/fetch.py", line 61, in fetch return self.collate_fn(data) File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/utils/data/_utils/collate.py", line 265, in default_collate return collate(batch, collate_fn_map=default_collate_fn_map) File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/utils/data/_utils/collate.py", line 143, in collate return [collate(samples, collate_fn_map=collate_fn_map) for samples in transposed] # Backwards compatibility. File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/utils/data/_utils/collate.py", line 143, in <listcomp> return [collate(samples, collate_fn_map=collate_fn_map) for samples in transposed] # Backwards compatibility. File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/utils/data/_utils/collate.py", line 120, in collate return collate_fn_map[elem_type](batch, collate_fn_map=collate_fn_map) File "/home/visge/miniconda3/envs/zk_torch/lib/python3.8/site-packages/torch/utils/data/_utils/collate.py", line 162, in collate_tensor_fn out = elem.new(storage).resize_(len(batch), *list(elem.size())) RuntimeError: Trying to resize storage that is not resizable
问题原因
错误源于加载的Tensor使用了不可调整大小的存储,DataLoader默认的collate_fn在尝试将多个样本拼接成批量Tensor时,无法修改这类存储的大小。另外,若不同样本的Tensor形状不一致,也会触发类似问题。
解决办法
1. 转换为可调整的Tensor
在__getitem__中加载Tensor后,显式创建新的可调整Tensor,确保存储支持修改:
def __getitem__(self, index): input_feature_path = self.input_feature_paths[index] input_feature = load(input_feature_path, map_location='cpu') # 用torch.tensor重新创建可调整Tensor input_feature = torch.tensor(input_feature, dtype=torch.float64) target_feature_path = self.target_feature_folder / input_feature_path.parts[-1] target_feature = load(target_feature_path, map_location='cpu') target_feature = torch.tensor(target_feature, dtype=torch.float64) return input_feature, target_feature
或者使用.clone()方法复制Tensor,生成可调整的存储:
input_feature = load(input_feature_path, map_location='cpu').clone().to(dtype=torch.float64) target_feature = load(target_feature_path, map_location='cpu').clone().to(dtype=torch.float64)
2. 确保样本形状一致
检查所有加载的input_feature和target_feature形状完全相同,避免拼接时出错:
def __getitem__(self, index): # ... 加载代码 ... # 替换为你的预期形状 expected_input_shape = (128, 100) expected_target_shape = (64, 100) if input_feature.shape != expected_input_shape: raise ValueError(f"输入特征形状不匹配:{input_feature.shape} vs {expected_input_shape}") if target_feature.shape != expected_target_shape: raise ValueError(f"目标特征形状不匹配:{target_feature.shape} vs {expected_target_shape}") # ... 转换类型并返回 ...
3. 自定义collate_fn
如果默认拼接逻辑不适用,手动实现批量拼接函数:
def custom_collate(batch): # 从batch中取出所有输入和目标特征,堆叠成批量Tensor inputs = torch.stack([item[0] for item in batch]) targets = torch.stack([item[1] for item in batch]) return inputs, targets # 创建DataLoader时指定自定义collate_fn train_loader = torch.utils.data.DataLoader( dataset, batch_size=32, shuffle=True, collate_fn=custom_collate )
内容的提问来源于stack exchange,提问作者tealy
相关产品推荐
相关产品推荐

