nerf核心是体渲染流程而非mlp结构:需手动实现位置编码(60维坐标+24维方向)、gpu上计算δ并反向cumprod求权重、按射线组织数据加载器、两阶段采样及gamma校正等细节决定训练成败。

直接上手 PyTorch 实现完整 NeRF 模型不现实——它不是调一个 torch.nn.Module 就能跑起来的“模型”,而是一整套坐标映射、体渲染、采样策略和训练循环的组合。你真正需要的,是理解哪些模块必须自己写、哪些可以复用、以及哪几个地方一写错就根本训不出图像。
NeRF 的核心不是网络结构,而是体渲染流程
NeRF 的 MLP(通常叫 NeRFModel)本身非常简单:输入 (x, y, z, dx, dy, dz),输出 (σ, r, g, b)。但它的输入不是原始坐标,而是经过位置编码(positional encoding)后的 60 维向量(3D 坐标 + 3D 方向,各做 10 层 L=10 的傅里叶映射)。这个编码必须手动实现,PyTorch 没有现成函数。
-
torch.sin和torch.cos是唯一依赖,别想用torch.fft或其他库替代 - 编码维度必须严格对齐:坐标部分输出 60 维(3 × 2 × 10),方向部分 24 维(3 × 2 × 4),否则 MLP 输入尺寸错配会报
RuntimeError: size mismatch - 编码应在 GPU 上完成,不要在 dataloader 里做 CPU 编码再搬上 GPU,否则训练慢 3 倍以上
体积分(volume rendering)必须手写,不能靠 torch.nn.functional
NeRF 的前向过程本质是沿射线做分段常数积分:对每条射线上采样的 N 个点,先算出每个点的密度 σ 和颜色 c,再用 torch.cumprod(1 - σ * δ, dim=-1) 推出权重 w,最后加权求和得像素颜色。这个过程没有现成算子。
- δ(相邻采样点距离)不能用固定值,必须由射线参数和采样点坐标动态计算:
δ = (z_vals[:, 1:] - z_vals[:, :-1]),否则渲染结果发灰或出现条纹 - 累积乘积要从后往前(
torch.cumprod(1 - σ * δ, dim=-1, reverse=True)),否则权重归零顺序错乱,整个 batch 渲染全黑 - 为防梯度爆炸,
σ输出建议用F.relu而非torch.exp,后者在早期训练中极易让 loss 突增至nan
数据加载器必须按射线组织,不是按图像
NeRF 不喂整张图,而是把一张图的所有像素坐标 + 对应 pose 拆成上万条射线(rays_o, rays_d),每 batch 随机采样 1024–4096 条。这意味着你的 DataLoader 输出不是 (B, 3, H, W),而是 (N, 6)(前3维是原点,后3维是方向)+ 对应的真值 RGB (N, 3)。
- 不要用
torchvision.datasets.ImageFolder直接加载——它输出图像张量,你需要的是射线参数,必须自己解析 COLMAP 或 LLFF 输出的poses_bounds.npy和images/ - batch 内射线必须来自同一张图(否则视角混乱,loss 不降),但不同 batch 可跨图;用
sampler=RayBatchSampler类控制,别依赖默认随机打散 - 验证时别用 full-image render:哪怕只 render 128×128 区域,也要确保射线生成逻辑和训练完全一致,否则 PSNR 看着高,实际泛化差
最难调的永远不是 MLP 结构,而是采样点分布(coarse→fine 两阶段采样中的 pdf 归一化)、白平衡处理(很多开源实现漏掉对真值 RGB 做 gamma 校正)、以及 pose 矩阵的旋转部分是否用了右乘惯例(R @ point.T 还是 point @ R.T)。这些细节不打印中间变量、不画出 z_vals 分布图,根本看不出问题。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











