ESC
输入关键词搜索文章标题和内容

替代NeRF:3DGS如何用PyTorch实现照片级渲染

本文由 linuxROS 整理发布,首发于 linuxros.cn,转载请注明出处。

替代NeRF:3DGS如何用PyTorch实现照片级渲染

导读:NeRF用神经网络"记住"场景,渲染要几十秒。3D Gaussian Splatting用高斯椭球体表示场景,20000个"彩色椭圆球"拼出照片级真实感,渲染只需几十毫秒。gsplat(UC Berkeley开源,1.6k+ stars)是官方PyTorch实现。本文拆解核心算法:四元数→协方差→EWA投影→Alpha合成→梯度反传。


一、从NeRF到3DGS

1.1 NeRF的困境

NeRF(Neural Radiance Fields)用MLP神经网络表示场景:
- 输入:3D坐标 (x, y, z) + 视角方向 (θ, φ)
- 输出:该点的颜色 + 密度

问题:渲染一帧需要调用神经网络数百次,RTX 3090上也要几十秒。

1.2 3DGS的革命

3D Gaussian Splatting用高斯椭球体替代神经网络:

场景 = N个高斯椭球体,每个椭球体有位置、形状、颜色、不透明度。把椭球体"溅射"到2D图像平面上,叠加颜色,就是渲染结果。

flowchart LR subgraph 场景["N个高斯椭球体"] G1["位置+形状+颜色+不透明度"] G2["..."] GN["..."] end subgraph 渲染["渲染管线"] G1 --> P["3D→2D投影"] G2 --> P GN --> P P --> A["Alpha合成"] A --> I["渲染图像"] end subgraph 训练["端到端训练"] I --> L["L1+SSIM Loss"] L --> G["梯度反传"] G -.->|"更新参数"| G1 end style G1 fill:#E3F2FD,stroke:#1976D2 style P fill:#FFF8E1,stroke:#F57C00 style A fill:#F3E5F5,stroke:#7B1FA2 style I fill:#E8F5E9,stroke:#388E3C

1.3 性能对比

方法 渲染一帧 训练时间 质量
NeRF 几十秒 几小时 高
3DGS 几十毫秒 几十分钟 更高

二、高斯椭球体的数学表示

2.1 每个高斯的参数

参数 维度 含义
mean (3,) 椭球体中心位置 (x, y, z)
covariance (3, 3) 椭球体形状的协方差矩阵
color (3,) RGB颜色 (R, G, B)
opacity (1,) 不透明度 (0-1)

2.2 四元数 → 协方差矩阵

class GaussianCloud(nn.Module):
    def __init__(self, num_gaussians=1000, device='cuda'):
        super().__init__()
        self.means = nn.Parameter(
            torch.randn(num_gaussians, 3, device=device) * 2.0
        )
        self.raw_scales = nn.Parameter(
            torch.randn(num_gaussians, 3, device=device) * 0.1 - 2.0
        )
        self.raw_quats = nn.Parameter(
            torch.randn(num_gaussians, 4, device=device)
        )
        self.raw_opacity = nn.Parameter(
            torch.randn(num_gaussians, 1, device=device) * 0.5 + 1.0
        )

    def get_covariance(self):
        # 1. 缩放因子(指数保证正定)
        scales = torch.exp(self.raw_scales)

        # 2. 四元数归一化
        quats = self.raw_quats / self.raw_quats.norm(dim=-1, keepdim=True)
        r, x, y, z = quats[:, 0], quats[:, 1], quats[:, 2], quats[:, 3]

        # 3. 四元数 → 旋转矩阵
        R = torch.stack([
            1-2*(y*y+z*z), 2*(x*y-r*z), 2*(x*z+r*y),
            2*(x*y+r*z), 1-2*(x*x+z*z), 2*(y*z-r*x),
            2*(x*z-r*y), 2*(y*z+r*x), 1-2*(x*x+y*y)
        ], dim=-1).reshape(-1, 3, 3)

        # 4. 协方差 = RSS^T R^T
        S = torch.diag_embed(scales)
        L = R @ S
        return L @ L.transpose(-2, -1)

三、EWA投影:3D → 2D

3.1 投影流程

flowchart TB A["3D点 (x,y,z)"] --> B["世界→相机变换<br/>W2C矩阵"] B --> C["相机坐标系 (xc,yc,zc)"] C --> D["透视投影"] D --> E["2D像素 (u,v)"] subgraph 协方差["协方差传播"] F["3D协方差 Σ"] --> G["雅可比矩阵 J"] G --> H["2D协方差 Σ' = JΣJ^T"] end style B fill:#E3F2FD,stroke:#1976D2 style C fill:#FFF8E1,stroke:#F57C00 style H fill:#E8F5E9,stroke:#388E3C

3.2 EWA投影实现

def project_gaussians(means_3d, cov_3d, K, R, T, img_h, img_w):
    # 1. 世界坐标 → 相机坐标
    W2C = torch.eye(4, device=means_3d.device)
    W2C[:3, :3] = R
    W2C[:3, 3] = T

    pts_4d = torch.cat([
        means_3d,
        torch.ones_like(means_3d[:, :1])
    ], dim=-1)
    pts_cam = (W2C @ pts_4d.T).T[:, :3]

    # 2. 透视投影到像素坐标
    fx, fy, cx, cy = K[0,0], K[1,1], K[0,2], K[1,2]
    pts_2d_x = fx * pts_cam[:, 0] / pts_cam[:, 2] + cx
    pts_2d_y = fy * pts_cam[:, 1] / pts_cam[:, 2] + cy

    # 3. 透视投影的雅可比矩阵
    W = torch.zeros(means_3d.shape[0], 3, 3, device=means_3d.device)
    W[:, 0, 0] = fx / pts_cam[:, 2]
    W[:, 1, 1] = fy / pts_cam[:, 2]
    W[:, 0, 2] = -fx * pts_cam[:, 0] / (pts_cam[:, 2] ** 2)
    W[:, 1, 2] = -fy * pts_cam[:, 1] / (pts_cam[:, 2] ** 2)

    # 4. 2D协方差 = J × 3D协方差 × J^T
    cov_2d = W @ cov_3d @ W.transpose(-2, -1)
    cov_2d = cov_2d[:, :2, :2] + torch.eye(2) * 1e-4

    return pts_2d_x, pts_2d_y, cov_2d, pts_cam[:, 2]

四、Alpha合成渲染

4.1 体渲染公式

2D图像上每个像素的颜色,来自所有覆盖该像素的高斯椭球体的加权叠加:

C = Σ cᵢ × αᵢ × Πⱼ(1 - αⱼ)

4.2 Alpha合成实现

def render_gaussians(pts_x, pts_y, cov_2d, colors, opacity,
                      depths, img_h, img_w):
    # 1. 按深度排序(从远到近)
    sorted_idx = torch.argsort(depths, descending=False)

    rendered = torch.zeros(3, img_h, img_w, device=colors.device)
    accum_alpha = torch.zeros(img_h, img_w, device=colors.device)

    # 2. 预计算2D高斯的逆协方差
    det = (cov_2d[:, 0, 0] * cov_2d[:, 1, 1] -
           cov_2d[:, 0, 1] * cov_2d[:, 1, 0])
    inv_cov = torch.zeros_like(cov_2d)
    inv_cov[:, 0, 0] = cov_2d[:, 1, 1] / det
    inv_cov[:, 1, 1] = cov_2d[:, 0, 0] / det
    inv_cov[:, 0, 1] = -cov_2d[:, 0, 1] / det
    inv_cov[:, 1, 0] = -cov_2d[:, 1, 0] / det

    # 3. 逐像素合成
    for i in sorted_idx:
        # 计算该像素处的高斯权重
        dx = pts_x - pts_x[i]
        dy = pts_y - pts_y[i]
        g_val = torch.exp(-0.5 * (
            dx * (inv_cov[i, 0, 0] * dx + inv_cov[i, 0, 1] * dy) +
            dy * (inv_cov[i, 1, 0] * dx + inv_cov[i, 1, 1] * dy)
        ))

        # Alpha合成
        alpha = opacity[i] * g_val
        rendered += alpha * (1 - accum_alpha) * colors[i]
        accum_alpha += alpha * (1 - accum_alpha)

    return rendered

五、梯度反传:端到端可微

3DGS的核心优势:整个渲染管线可微,可以端到端优化。

flowchart TB A["初始化N个高斯"] --> B["四元数→旋转矩阵"] B --> C["计算协方差矩阵"] C --> D["EWA投影 3D→2D"] D --> E["深度排序"] E --> F["Alpha合成"] F --> G["渲染图像"] G --> H{"Loss计算<br/>L1 + SSIM"} H --> I["梯度反传"] I --> J["更新高斯参数"] J --> K{"收敛?"} K -->|"否"| D K -->|"是"| L["输出场景"] style A fill:#E3F2FD,stroke:#1976D2 style D fill:#FFF8E1,stroke:#F57C00 style G fill:#E8F5E9,stroke:#388E3C style H fill:#F3E5F5,stroke:#7B1FA2

六、gsplat官方库

根据gsplat官方文档(UC Berkeley开源,Apache 2.0协议):

来自 linuxros.cn · linuxROS
import torch
from gsplat import rasterization

# 初始化高斯
mean = torch.tensor([[0., 0., 0.01]], device="cuda")
quat = torch.tensor([[1., 0., 0., 0.]], device="cuda")
scale = torch.rand((1, 3), device="cuda")
color = torch.rand((1, 3), device="cuda")
opac = torch.ones((1,), device="cuda")

# 相机参数
view = torch.eye(4, device="cuda")[None]
K = torch.tensor([[[1., 0., 120.],
                   [0., 1., 120.],
                   [0., 0., 1.]]], device="cuda")

# 渲染
rgb_image, alpha, metadata = rasterization(
    mean, quat, scale, opac, color, view, K, 240, 240
)

七、性能验证

指标 数据
创建20000高斯 0.281s
计算协方差 0.194s
EWA投影 0.101s
渲染256×256 5.279s
梯度反传 5.889s
显存占用 20MB

八、实战踩坑

8.1 显存不足(OOM)

GTX 965M仅有4GB显存,大幅面渲染容易OOM:

# 降低分辨率
img_h, img_w = 128, 128  # 从256x256降低

# 减少高斯数量
num_gaussians = 5000    # 从20000减少

# 降低batch_size
batch_size = 1           # 最小batch

8.2 PyTorch CUDA版本不匹配

踩坑:使用清华镜像的CUDA 12.9 nvidia包,导致Bus Error。

原因:nvidia包版本必须与torch的CUDA编译版本精确匹配。

正确做法:

# 1. 先查torch的CUDA版本
python3 -c "import torch; print(torch.version.cuda)"
# 输出: 12.1

# 2. 安装匹配版本的nvidia包
pip install -i https://pypi.tuna.tsinghua.edu.cn/simple \
  nvidia-cublas-cu12==12.1.3.1 \
  nvidia-cuda-nvrtc-cu12==12.1.105 \
  nvidia-cuda-runtime-cu12==12.1.105 \
  ...

# 3. 再安装torch
pip install torch==2.5.1 torchvision==0.20.1 \
  --index-url https://download.pytorch.org/whl/cu121

8.3 性能优化

优化项 方法 效果
显存优化 降低分辨率至128×128 显存占用减少75%
计算优化 使用向量化操作 渲染速度提升2-3倍
梯度裁剪 torch.nn.utils.clip_grad_norm_ 训练稳定性提升

完整安装验证:

# WSL2环境验证
source /root/slam_venv/bin/activate

python3 -c "
import torch
print(f'CUDA available: {torch.cuda.is_available()}')
print(f'GPU: {torch.cuda.get_device_name(0)}')
print(f'CUDA version: {torch.version.cuda}')
print(f'PyTorch version: {torch.__version__}')
"
# 应输出: CUDA available: True, GTX 965M, 12.1, 2.5.1

8.4 gsplat官方库安装

pip install gsplat

# 验证
python3 -c "from gsplat import rasterization; print('gsplat OK')"

九、技术对比

维度 点云 NeRF 3DGS
表示方式 离散点 神经网络 高斯椭球体
渲染速度 快(光栅化) 慢(MLP推理) 极快(光栅化)
可微性 ❌ ✅ ✅
新视角合成 差 高 极高
内存占用 中 高 低

十、总结

3DGS的核心是"椭球体+投影+合成"三步走:

  • 椭球体表示:均值+协方差+颜色+不透明度
  • 四元数旋转:稳定计算旋转矩阵
  • EWA投影:雅可比矩阵近似3D→2D变换
  • Alpha合成:从远到近累积颜色和透明度
  • 端到端可微:L1+SSIM loss驱动参数优化

纯PyTorch实现证明了3DGS的核心算法不依赖专用光栅化器。gsplat官方库提供了CUDA加速版本,渲染速度更快。


看完不点赞?渲染明天就黑屏。看完不转发?Bug跟你到白头。推荐给朋友,一起掉坑到永久。评论区见,谁家的3DGS最丝滑?

版权声明

作者linuxROS
协议本作品采用 CC BY-NC-SA 4.0 许可协议:署名-非商业性使用-相同方式共享
关注欢迎关注微信公众号 linuxROS,获取更多机器人 / 嵌入式 / Linux 干货
返回首页