据 MarkTechPost 报道,一篇发布于 2026 年 9 月 13 日的教程演示了如何用 JAX、Flax、Optax 以及 google-research/jax3d 提供的体积渲染原语,构建一套端到端的分层神经辐射场(NeRF)流程,涵盖合成多视图数据生成、前向渲染、模型训练、新视角合成与三维几何提取。
教程首先从一个解析场景构造合成多视图数据集,该场景同时包含体积几何与随视角变化的辐射度,并调用 jax3d 的 sample_along_rays 与 volume_rendering 建立前向渲染过程。随后实现带位置编码、跳跃连接、粗网络与细网络分离以及视角方向条件化的 NeRF,其中分层重要性采样通过 sample_piecewise_constant_pdf 完成。训练环节使用 JAX 的 JIT 编译、Adam 优化器、指数学习率衰减与梯度裁剪。
在环境准备上,教程通过 pip 安装 etils 的 array-types、epy、etree、enp 等扩展,以及 chex、flax、optax、scikit-image,并浅克隆 jax3d 仓库。为避免引入 gin、tfds 等依赖,代码用 importlib 按文件路径单独加载仓库中的 volume_rendering.py,而不是安装整个包;若模块加载失败,教程建议改用 etils 1.9.4 版本后重试。运行时脚本会打印 JAX 版本、设备类型与平台,并检查 sample_along_rays、volume_rendering、sample_piecewise_constant_pdf、sample_1d 等 API 是否可用。
默认配置面向 GPU:图像分辨率 64×64,训练视图 24 个、测试视图 3 个,相机半径 3.2,视场角 40 度,近平面 1.9、远平面 4.7;每条光线在真值场景中取 256 个采样点,粗网络与细网络各 64 个采样点;位置编码 10 阶、方向编码 4 阶,网络宽度 128、深度 6 层,第 3 层设置跳跃连接;训练批大小为每条 2048 条光线,共 2500 步,学习率从 5e-4 经指数衰减降至 5e-6,分块大小 4096,用于几何提取的 marching cubes 网格分辨率 96。脚本检测到运行在 CPU 上时会自动切换到小规模配置,包括 40×40 图像、14 个训练视图、400 步训练、128 个真值采样点、粗细网络各 32 个采样点、宽度 64、4 层网络、第 2 层跳跃连接、每批 1024 条光线、分块 1600、网格分辨率 64,并提示用户切换到 GPU 运行完整版本。
相机与光线生成部分,look_at 函数采用 OpenGL/NeRF 约定,即相机坐标系 +x 向右、+y 向上、光轴朝向 -z;orbit_poses 用黄金角方位角配合从 18 度到 58 度单调递增的仰角,在穹顶轨迹上生成分布较均匀的视角;rays_from_pose 返回形状为 H×W×3 的起点与方向,方向为单位向量,因此采样器给出的深度对应世界空间中的真实距离。
评估阶段使用峰值信噪比(PSNR)衡量新视角合成质量,并输出深度图与不透明度可视化、采样诊断信息、360 度环绕渲染结果,同时通过 marching cubes 从密度场中提取三维几何。教程以可复制运行的代码形式给出上述全部步骤,单篇内容即覆盖从数据构造到几何重建的完整链路。