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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 16:26:26