Gauss Splat 接口说明文档
概述
gauss_splat 是一个基于华为昇腾NPU的高斯溅射(Gaussian Splatting)渲染库,提供了完整的3D高斯渲染流水线。该库包含投影、排序、渲染等核心操作,支持可微分渲染。
对外接口列表
1. Rasterizer 类
功能: 高斯溅射光栅化器,提供完整的渲染流水线。
位置: gauss_splat.Rasterizer
主要方法
rasterization()
渲染3D高斯点云到2D图像。
参数:
cam(Tuple): 相机参数元组camtoworlds(Tensor): 相机到世界变换矩阵,形状 (C, 4, 4)Ks(Tensor): 相机内参矩阵,形状 (C, 3, 3)render_mode(str): 渲染模式,可选值: "RGB", "D", "ED", "RGB+D", "RGB+ED"
size(Tuple): 图像尺寸 (width, height)tile_size(int): 分块大小,默认为32active_sh_degree(int): 激活的球谐函数阶数splats(dict): 高斯点云字典,包含:means(Tensor): 高斯中心位置,形状 (N, 3)quats(Tensor): 四元数表示旋转,形状 (N, 4)scales(Tensor): 缩放参数(log空间),形状 (N, 3)opacities(Tensor): 不透明度(logits),形状 (N,)sh0(Tensor): 球谐系数第0阶,形状 (N, 1, 3)shN(Tensor): 球谐系数高阶,形状 (N, K, 3)
camera_model(str): 相机模型,可选值: "pinhole", "ortho", "fisheye",默认为"pinhole"
返回值:
render_colors(Tensor): 渲染的颜色图像,形状 (C, H, W, 3)render_depth(Tensor): 渲染的深度图像,形状 (C, H, W, 1)info(dict): 元数据字典,包含:gaussian_ids: 高斯点ID(当前为None)means2d: 2D投影坐标,形状 (B, C, 2, N)radii: 投影半径,形状 (B, C, 2, N)width: 图像宽度height: 图像高度n_cameras: 相机数量
示例:
rasterizer = gauss_splat.Rasterizer()
render_colors, render_depth, info = rasterizer.rasterization(
cam=(camtoworlds, Ks, "RGB"),
size=(width, height),
tile_size=32,
active_sh_degree=3,
splats=splats_dict,
camera_model="pinhole"
)
2. projection_three_dims_gaussian_fused
功能: 3D高斯投影融合操作,将3D高斯投影到2D平面并进行裁剪。
位置: gauss_splat.projection_three_dims_gaussian_fused
参数:
means(Tensor): 高斯中心位置,形状 (B, N, 3)colors(Tensor): 颜色值,形状 (B, 3, N)covars(Tensor, optional): 协方差矩阵,形状 (B, N, 3, 3)。与quat/scales互斥quat(Tensor, optional): 四元数表示旋转,形状 (B, N, 4)。与covars互斥scales(Tensor, optional): 缩放参数,形状 (B, N, 3)。与quat一起使用opacities(Tensor): 不透明度,形状 (B, N)viewmats(Tensor): 视图变换矩阵,形状 (B, C, 4, 4)ks(Tensor): 相机内参矩阵,形状 (B, C, 3, 3)width(int): 图像宽度height(int): 图像高度eps(float): 最小高斯半径阈值,默认0.3near_plane(float): 近平面距离,默认0.01far_plane(float): 远平面距离,默认1e10calc_compensations(bool): 是否计算补偿因子,默认Falsecamera_model(str): 相机模型,默认"pinhole"
返回值:
means2d(Tensor): 2D投影坐标,形状 (B, C, 2, N)depths(Tensor): 深度值,形状 (B, C, N)conics(Tensor): 2D协方差逆矩阵(锥形参数),形状 (B, C, 3, N)opacities(Tensor): 过滤后的不透明度,形状 (B, C, N)radius(Tensor): 投影半径,形状 (B, C, 2, N)covars2d(Tensor): 2D协方差矩阵,形状 (B, C, 3, N)colors(Tensor): 过滤后的颜色,形状 (B, C, 3, N)cnt(Tensor): 有效高斯点数量, 形状 (B, C)
示例:
means2d, depths, conics, opacities, radius, covars2d, colors, cnt = \
gauss_splat.projection_three_dims_gaussian_fused(
means=means,
colors=colors,
quat=quats,
scales=scales,
opacities=opacities,
viewmats=viewmats,
ks=Ks,
width=1920,
height=1080
)
3. calc_render
功能: 计算高斯溅射渲染,执行alpha混合和深度计算。
位置: gauss_splat.calc_render
参数:
means(Tensor): 2D高斯中心坐标,形状 (2, N)conic0s(Tensor): 协方差逆矩阵元素0,形状 (1, N)conic1s(Tensor): 协方差逆矩阵元素1,形状 (1, N)conic2s(Tensor): 协方差逆矩阵元素2,形状 (1, N)opacities(Tensor): 不透明度,形状 (1, N)colors(Tensor): 颜色,形状 (3, N)depths(Tensor, optional): 深度值,形状 (1, N)。如为None则不渲染深度tile_coords(Tensor): 分块坐标,形状 (tileNum, 2, nPixel)offsets(Tensor): 偏移量,形状 (vectorCnt + (TileNum * 2))sorted_gs_ids(Tensor): 排序后的高斯点ID,形状 (totalGauss)
返回值:
- 如果提供
depths:color(Tensor): 渲染的颜色图像,形状 (3, tileNum, nPixel)depth(Tensor): 渲染的深度图像,形状 (1, tileNum, nPixel)
- 如果不提供
depths:color(Tensor): 渲染的颜色图像,形状 (3, tileNum, nPixel)
示例:
render_colors, render_depths = gauss_splat.calc_render(
means=means2d,
conic0s=conics[:, 0],
conic1s=conics[:, 1],
conic2s=conics[:, 2],
opacities=opacities,
colors=colors,
depths=depths,
tile_coords=pix_coords,
offsets=lb_sched,
sorted_gs_ids=sorted_gs_id
)
4. spherical_harmonics
功能: 计算球谐函数,用于视角相关的颜色计算。
位置: gauss_splat.spherical_harmonics
参数:
degrees_to_use(int): 球谐函数阶数,范围0-4dirs(Tensor): 方向向量,形状 (B, N, 3)coeffs(Tensor): 球谐系数,形状 (B, N, K, 3),其中K=(degrees_to_use+1)²
返回值:
output(Tensor): 计算得到的颜色值,形状 (B, 3, N)
示例:
colors = gauss_splat.spherical_harmonics(
degrees_to_use=3,
dirs=view_directions, # (1, N, 3)
coeffs=sh_coefficients # (1, N, 16, 3)
)
5. gaussian_sort
功能: 对高斯点按深度进行排序,用于正确的alpha混合。
位置: gauss_splat.gaussian_sort
参数:
lb_sched(Tensor): 负载均衡调度张量,形状 (B, C, schedule_num)gaussian_cnt(Tensor): 每个tile的高斯点计数,形状 (B, C, tile_num, 1)depths(Tensor): 深度值,形状 (B, C, tile_num, N)gs_ids(Tensor): 高斯点ID,形状 (B, C, tile_num, N)sorted_offset(Tensor): 排序偏移量,形状 (B*C)max_tile_gauss(int): 单个tile最大高斯点数
返回值:
sorted_gs_ids(Tensor): 排序后的高斯点ID,一维展平张量,形状 (totalGauss)
示例:
sorted_gs_ids = gauss_splat.gaussian_sort(
lb_sched=lb_sched_tensor,
gaussian_cnt=tile_sums,
depths=tile_depths,
gs_ids=tile_gauss_ids,
sorted_offset=sorted_offset,
max_tile_gauss=max_tile_gauss
)
6. flash_gaussian_build_mask
功能: 构建高斯溅射的tile掩码,用于分块渲染优化。
位置: gauss_splat.flash_gaussian_build_mask
参数:
means2d(Tensor): 2D投影坐标,形状 (B, C, 2, N)opacity(Tensor): 不透明度,形状 (B, C, 1, N)conics(Tensor): 协方差逆矩阵,形状 (B, C, 3, N)covars2d(Tensor): 2D协方差矩阵,形状 (B, C, 3, N)depths(Tensor): 深度值,形状 (B, C, 1, N)cnt(Tensor): 有效高斯点计数,形状 (B, C)tile_grid(Tensor): tile网格坐标image_width(int): 图像宽度image_height(int): 图像高度tile_size(int): tile大小,默认64
返回值:
tile_sum(Tensor): 每个tile的高斯点和,形状(B, C, tile_num, 1)tile_offset(Tensor): tile偏移量,形状(B, C, tile_num, 1)tile_depths(Tensor): tile深度,形状(B, C, tile_num, N)gauss_index(Tensor): 高斯点索引,形状(B, C, tile_num, N)
示例:
tile_sums, tile_offsets, tile_depths, tile_gauss_ids = \
gauss_splat.flash_gaussian_build_mask(
means2d=means2d,
opacity=opacities,
conics=conics,
covars2d=covars2d,
depths=depths,
cnt=cnt,
tile_grid=tile_grid,
image_width=1920,
image_height=1080,
tile_size=32
)
7. gaussian_filter
功能: 根据深度范围和有效性过滤高斯点。
位置: gauss_splat.gaussian_filter
参数:
means(Tensor): 3D高斯中心位置,形状 (B, 3, N)colors(Tensor): 颜色值,形状 (B, 3, N)det(Tensor): 协方差行列式,形状 (B, C, N)opacities(Tensor): 不透明度,形状 (B, N)means2d(Tensor): 2D投影坐标,形状 (B, C, 2, N)depths(Tensor): 深度值,形状 (B, C, N)radius(Tensor): 投影半径,形状 (B, C, 2, N)conics(Tensor): 协方差逆矩阵,形状 (B, C, 3, N)covars2d(Tensor): 2D协方差矩阵,形状 (B, C, 3, N)compensations(Tensor, optional): 补偿因子,形状 (B, C, N)width(int): 图像宽度height(int): 图像高度near_plane(float): 近平面距离far_plane(float): 远平面距离
返回值:
means_culling: 过滤后的3D坐标,形状 (B, C, 3, N)colors_culling: 过滤后的颜色,形状 (B, C, 3, N)means2d_culling: 过滤后的2D坐标,形状 (B, C, 2, N)depths_culling: 过滤后的深度,形状 (B, C, N)radius_culling: 过滤后的半径,形状 (B, C, 2, N)covars2d_culling: 过滤后的2D协方差,形状 (B, C, 3, N)conics_culling: 过滤后的协方差逆矩阵,形状 (B, C, 3, N)opacities_culling: 过滤后的不透明度,形状 (B, C, N)proj_filter: 投影过滤器,形状 (B, C, ceil(N/8))cnt: 有效点数量,形状 (B, C)
8. get_render_schedule_cpp
功能: 获取渲染调度方案,用于负载均衡。
位置: gauss_splat.get_render_schedule_cpp
参数:
nums_tensor(Tensor): 每个tile的高斯点数量,形状 (B, C, T)num_bins(int): 分箱数量(向量处理器数量)
返回值:
lb_sched_tensor(Tensor): 负载均衡调度张量,形状 (B, C, M)
示例:
vector_num = acl.get_device_capability(0, 1)[0]
lb_sched_tensor = gauss_splat.get_render_schedule_cpp(
nums_tensor=tile_sums_cpu,
num_bins=vector_num
)
数据类型说明
张量设备要求
所有输入张量必须在NPU设备上,支持自动求导。
数据格式
- 位置坐标: 世界坐标系下的3D坐标 (x, y, z)
- 四元数: 用于表示旋转的四元数 (w, x, y, z)
- 缩放: log空间的缩放因子,渲染时通过
torch.exp(scales)转换 - 不透明度: logit形式,渲染时通过
torch.sigmoid(opacities)转换 - 球谐系数: 用于视角相关颜色,形状为 (N, K, 3),K为系数数量
相机模型
支持三种相机模型:
- pinhole: 针孔相机模型
- ortho: 正交投影
- fisheye: 鱼眼镜头
渲染模式
- RGB: 仅渲染颜色
- D: 仅渲染深度
- ED: 仅渲染边缘深度
- RGB+D: 渲染颜色和深度
- RGB+ED: 渲染颜色和边缘深度
完整渲染流水线示例
import torch
import gauss_splat
# 1. 初始化光栅化器
rasterizer = gauss_splat.Rasterizer()
# 2. 准备高斯点云数据
N = 10000 # 高斯点数量
splats = {
"means": torch.randn(N, 3, device='npu:0'), # 位置
"quats": torch.randn(N, 4, device='npu:0'), # 四元数
"scales": torch.randn(N, 3, device='npu:0'), # 缩放(log空间)
"opacities": torch.randn(N, device='npu:0'), # 不透明度(logits)
"sh0": torch.randn(N, 1, 3, device='npu:0'), # 球谐第0阶
"shN": torch.randn(N, 15, 3, device='npu:0'), # 球谐高阶
}
# 3. 准备相机参数
C = 1 # 相机数量
camtoworlds = torch.eye(4, device='npu:0').unsqueeze(0).expand(C, -1, -1)
Ks = torch.tensor([
[500.0, 0.0, 960.0],
[0.0, 500.0, 540.0],
[0.0, 0.0, 1.0]
], device='npu:0').unsqueeze(0).expand(C, -1, -1)
# 4. 渲染
render_colors, render_depths, info = rasterizer.rasterization(
cam=(camtoworlds, Ks, "RGB"),
size=(1920, 1080),
tile_size=32,
active_sh_degree=3,
splats=splats,
camera_model="pinhole"
)
print(f"渲染颜色形状: {render_colors.shape}") # (1, 1080, 1920, 3)
print(f"渲染深度形状: {render_depths.shape}") # (1, 1080, 1920, 1)
版本信息
- 版权所有 © 2026 华为技术有限公司
- 许可证: CANN Open Software License Agreement Version 2.0