PyTorch中RuntimeError: Could not infer dtype of generator错误排查
问题:PyTorch调用to_tensor时触发RuntimeError: Could not infer dtype of generator
在使用PyTorch构建训练数据时,自定义了shufflerow函数、Dataset类处理数据,同时实现to_tensor函数用于数据转Tensor并指定设备。但传入Dataset处理后的数据调用to_tensor时,出现类型解析错误。
相关代码
数据处理代码
def shufflerow(tensor1, tensor2, axis): row_perm = torch.rand(tensor1.shape[:axis+1]).argsort(axis) # 获取排列索引 for _ in range(tensor1.ndim-axis-1): row_perm.unsqueeze_(-1) row_perm = row_perm.repeat(*[1 for _ in range(axis+1)], *(tensor1.shape[axis+1:])) # 调整形状适配gather操作 return tensor1.gather(axis, row_perm),tensor2.gather(axis, row_perm) class Dataset: def __init__(self, observation, next_observation): self.data =(observation, next_observation) indices = torch.randperm(observation.shape[0]) self.train_samples = (observation[indices ,:], next_observation[indices ,:]) self.test_samples = shufflerow(observation, next_observation, 0)
to_tensor函数代码
def to_tensor(x, device): if torch.is_tensor(x): return x elif isinstance(x, np.ndarray): return torch.from_numpy(x).to(device=device, dtype=torch.float32) elif isinstance(x, list): if all(isinstance(item, np.ndarray) for item in x): return [torch.from_numpy(item).to(device=device, dtype=torch.float32) for item in x] elif isinstance(x, tuple): return (torch.from_numpy(item).to(device=device, dtype=torch.float32) for item in x) else: print(f"X:{x} and X's type{type(x)}") return torch.tensor(x).to(device=device, dtype=torch.float32)
错误栈
-> 1725 self._target_samples = to_tensor(true_samples) 1726 self._steps = [] /content/data_gen.py in to_tensor(x) 1368 else: 1369 print(f"X:{x} and X's type{type(x)}") -> 1370 return torch.tensor(x).to(device=device, dtype=torch.float32) X:<generator object to_tensor.<locals>.<genexpr> at 0x7f380235d6d0> and X's type<class 'generator'> RuntimeError: Could not infer dtype of generator
错误原因
- 生成器被误传入torch.tensor:
to_tensor处理tuple类型时,返回的是生成器表达式((...) for item in x),而非实际的tuple对象。如果后续代码将这个生成器再次传入to_tensor,生成器不属于函数中判断的任何类型(tensor、ndarray、list、tuple),会走到else分支,被传入torch.tensor(),而PyTorch无法将生成器直接转为Tensor,触发报错。 - tuple分支逻辑缺陷:当前tuple分支直接用
torch.from_numpy处理每个元素,若元素本身已经是torch.tensor,会导致额外错误;同时生成器的惰性特性会导致后续处理时类型不符合预期。
修复方案
方案1:将生成器转为tuple对象
修改to_tensor的tuple分支,把生成器表达式转为实际的tuple:
elif isinstance(x, tuple): return tuple(torch.from_numpy(item).to(device=device, dtype=torch.float32) for item in x)
方案2:递归调用to_tensor,兼容多种元素类型
更健壮的方式是对tuple中的每个元素递归调用to_tensor,自动兼容tensor、ndarray等多种类型:
def to_tensor(x, device): if torch.is_tensor(x): return x.to(device=device, dtype=torch.float32) # 补充设备和类型转换 elif isinstance(x, np.ndarray): return torch.from_numpy(x).to(device=device, dtype=torch.float32) elif isinstance(x, list): return [to_tensor(item, device) for item in x] elif isinstance(x, tuple): return tuple(to_tensor(item, device) for item in x) else: print(f"X:{x} and X's type{type(x)}") return torch.tensor(x).to(device=device, dtype=torch.float32)
内容的提问来源于stack exchange,提问作者Dalek
相关产品推荐
相关产品推荐

