如何在NeRF代码中转换double为float?及原代码无报错原因咨询
NeRF模型数据类型不匹配问题解决
代码片段
class NeRF(nn.Module): def __init__(self, D=8, W=256, input_ch=60, input_ch_views=24, output_ch=4, skip=4, use_view_dirs=True): ... def forward(self, input_pts, input_views=None): h = input_pts.clone() for i, _ in enumerate(self.pts_linears): h = F.relu(self.pts_linears[i](h)) if i == self.skip: h = torch.cat([input_pts, h], -1) ...
报错信息
---> 40 h = F.relu(self.pts_linears[i](h)) ... RuntimeError: expected scalar type Float but found Double
变量转float类型的方法
有两种核心处理方向,按需选择:
转换输入数据类型:
在传入模型forward之前,将输入张量转为float32:input_pts = input_pts.float() if input_views is not None: input_views = input_views.float()或者在
forward方法开头直接处理:def forward(self, input_pts, input_views=None): input_pts = input_pts.float() h = input_pts.clone() # 后续代码不变注意如果
input_views参与运算,也要同步转换类型。转换模型参数类型:
初始化模型后,将整个模型的参数转为float32(默认即为float32,仅当模型参数被意外修改为其他类型时需要):model = NeRF().float()
原代码未报错的原因
原代码运行时,输入数据的张量类型(float32)和模型参数的默认类型(float32)完全一致,PyTorch允许同类型张量进行运算,因此不会触发类型不匹配的报错。
本次报错是因为后续输入数据被改为double64类型(比如从numpy加载数据时,numpy默认类型为double;或者手动调用了.double()方法),但模型参数仍保持默认的float32,两者类型不兼容,从而抛出RuntimeError。
内容的提问来源于stack exchange,提问作者lavinal
相关产品推荐
相关产品推荐

