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

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

错误原因

  1. 生成器被误传入torch.tensor:to_tensor处理tuple类型时,返回的是生成器表达式((...) for item in x),而非实际的tuple对象。如果后续代码将这个生成器再次传入to_tensor,生成器不属于函数中判断的任何类型(tensor、ndarray、list、tuple),会走到else分支,被传入torch.tensor(),而PyTorch无法将生成器直接转为Tensor,触发报错。
  2. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 14:45:33