PyTorch自动创建weight类型不匹配问题咨询:输入为torch.DoubleTensor
当你碰到RuntimeError: Expected object of type torch.DoubleTensor but found type torch.FloatTensor for argument #2 'weight'这类类型不匹配错误,且完全没有手动初始化weight时,以下几个核心因素会决定PyTorch自动生成的参数(比如weight)的数据类型:
全局默认浮点张量类型
PyTorch默认的浮点类型是torch.float32(对应FloatTensor),这是绝大多数层(比如nn.Linear、nn.Conv2d)自动创建参数时的默认选择,除非你主动修改全局配置。比如执行torch.set_default_dtype(torch.double)或者torch.set_default_tensor_type(torch.DoubleTensor)后,新创建的参数就会默认使用DoubleTensor类型。层初始化时显式指定的dtype参数
虽然你说没自行初始化weight,但如果在定义层的时候(比如nn.Linear(10, 20, dtype=torch.double))显式传递了dtype参数,PyTorch会直接用这个指定的类型来生成weight和bias。要是没指定,就会 fallback 到全局默认类型。模型/模块的
dtype属性
如果你给整个模型或模块设置了dtype(比如model = model.to(torch.double),或者创建模块时直接指定dtype),后续添加到该模块的子层,其自动创建的参数会继承这个模块的dtype设置。需要注意的是,to()方法是转换已有参数的类型,而之后新增的层才会直接继承模块的dtype。混合精度训练相关配置
如果你开启了自动混合精度(比如用torch.cuda.amp.autocast),这不会改变参数本身的类型,但可能会让输入在前向传播中被自动转换——不过这不是你当前问题的诱因,毕竟你输入是DoubleTensor,参数是FloatTensor,本质还是参数默认类型和输入不匹配导致的。
额外补充:输入张量的类型不会影响PyTorch自动创建的参数类型。PyTorch不会根据输入数据的类型自动调整参数类型,这也是为什么你输入DoubleTensor但参数是FloatTensor时会报错。解决这个问题的快速方法有三种:要么把输入转换为FloatTensor(input = input.float()),要么把模型参数转换为DoubleTensor(model = model.double()),或者修改全局默认浮点类型。
内容的提问来源于stack exchange,提问作者Eric Kani

