brush-splat 深度解析:brush-sort 基于 WebGPU 间接调度的基数排序(Radix Sort)实现
2026/9/20 21:32:11 网站建设 项目流程

brush-splat 深度解析:brush-sort 基于 WebGPU 间接调度的基数排序(Radix Sort)实现

【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research

本文聚焦 brush_splat 仓库中的独立排序 crate ——brush-sort(对应文档 brush_splat/crates/brush-sort/README.md),完整讲解一个 WebGPU 兼容的 GPU 基数排序库的设计与实现。通过本文,读者将理解:它如何用五段 WGSL compute shader 流水线完成按位排序、如何用 indirect dispatch 对“元素数量由 GPU 运行时决定”的数组排序、如何通过sorting_bits控制排序位宽,以及它如何被上层渲染器用于高斯溅射的深度排序。

一、模块定位:为 GPU 渲染准备的排序原语

brush-sort是 brush_splat 仓库中一个职责单一的基础 crate(见 Cargo.toml),依赖brush-kernel(内核调度与工具函数)以及burn/burn-wgpu/burn-jit张量生态。它的官方 README 用两句话概括了全部核心定位:

WebGPU compatible radix sort. It's based on this implementation, which in turn is based on FidelityFX Radix sort. It allows sorting up to a given number of bits, and sorting an array with a GPU known number of elements using indirect dispatches.

翻译并拆解为三点设计目标:

  1. WebGPU 兼容:全部排序逻辑用 WGSL compute shader 实现,可在任意 WebGPU 后端(如浏览器与 wgpu 原生后端)上运行;
  2. 算法血统:实现参考了 googlefonts/compute-shader-101 的 PR 实现,而该实现又源于 AMD 的 FidelityFX Radix Sort(经典 GPU 基数排序范式:直方图统计 → 前缀和 → 散射);
  3. 两个杀手级能力:支持仅排序指定数量的比特位(sorting_bits);支持对“元素个数由 GPU 决定”的数组进行排序(通过 indirect dispatch 实现,避免 CPU-GPU 往返同步)。

在下游,它被 brush-render 以brush_sort::radix_argsort的形式调用,用于按深度重排 3D 高斯溅射,这正是整个渲染管线中“可见性排序”的关键一环。

二、算法总体设计:每趟 4 比特的 LSD 基数排序

从入口 brush_splat/crates/brush-sort/src/lib.rs 可以看到,排序以LSD(最低有效位优先)方式逐趟进行,每趟处理BITS_PER_PASS = 4位,因此趟数为sorting_bits.div_ceil(4)。当sorting_bits = 32时共 8 趟;sorting_bits = 16时共 4 趟。

核心常量定义在共享 shader 模块 sorting.wgsl 中,并在 mod.rs 中被自动生成为 Rust 常量:

常量含义
WG256每个 workgroup 的线程数(@workgroup_size(256, 1, 1)
BITS_PER_PASS4每趟排序的比特位数,即 radix 为 2⁴ = 16
BIN_COUNT16每趟的桶(bin)数量,1 << BITS_PER_PASS
ELEMENTS_PER_THREAD4每个线程单趟处理的元素数(循环展开)
BLOCK_SIZE1024每个 workgroup 处理的元素数,WG * ELEMENTS_PER_THREAD
HISTOGRAM_SIZE4096直方图缓冲区大小,WG * BIN_COUNT

单趟流程由5 段 compute shader组成,各自职责如下:

sort_count ──► sort_reduce ──► sort_scan ──► sort_scan_add ──► sort_scatter 统计直方图 块内规约 块上前缀和 全局偏移叠加 按偏移散射
  • sort_count.wgsl:每个 workgroup 对自己块内的元素统计 16 桶直方图(atomicAdd写入 workgroup 共享内存),随后将直方图写出为counts[bin * num_wgs + group_id]
  • sort_reduce.wgsl:把每个桶的多个 workgroup 直方图分段规约求和,得到reduced缓冲区,供下一步做全局前缀和;
  • sort_scan.wgsl:以单个 workgroupCubeCount::Static(1,1,1))对整个reduced数组做排他前缀和(exclusive prefix sum),利用lds(local data share)转置技巧配合希尔式 workgroup 内前缀和;
  • sort_scan_add.wgsl:将全局前缀偏移加回每个 workgroup 的直方图,得到每个 workgroup 每个桶的全局写入起点
  • sort_scatter.wgsl:每个线程对自己块内的元素按当前 4 位键值,结合bin_offset_cache与 workgroup 内直方图前缀和,计算目标位置并写出outout_values

值得注意的是sort_scatter.wgsl内部还有一个两层技巧:先用 2 位一组的方式在 workgroup 共享内存中做局部计数排序(lds_sums[key_offset] = local_key两轮回写),再在最后一轮用 4 位桶散射,从而减少对全局内存的随机写压力。

三、入口函数 radix_argsort:类型签名与调度编排

对外唯一的主入口是 lib.rs 中的:

pub fn radix_argsort( input_keys: JitTensor<WgpuRuntime, u32>, input_values: JitTensor<WgpuRuntime, u32>, n_sort: JitTensor<WgpuRuntime, u32>, sorting_bits: u32, ) -> (JitTensor<WgpuRuntime, u32>, JitTensor<WgpuRuntime, u32>)

参数语义:

  • input_keys/input_values:待排序的键与伴随值(长度必须相等,函数内assert_eq!强制校验),两者会以(键, 值)对的形式一起重排;
  • n_sortGPU 端已知的有效元素数量(一个长度为 1 的 u32 张量),这正是“元素个数由 GPU 决定”的落点——排序过程中所有 workgroup 数量都由该值动态推导;
  • sorting_bits:需要排序的比特位数,入口断言assert!(sorting_bits <= 32)
  • 返回:排序后的键与值张量对。

调度编排的关键代码(lib.rs)如下:

let max_needed_wgs = max_n.div_ceil(BLOCK_SIZE); let num_wgs = create_dispatch_buffer(n_sort.clone(), [BLOCK_SIZE, 1, 1]); let num_reduce_wgs: Tensor<JitBackend<WgpuRuntime, f32, i32>, 1, Int> = Tensor::from_primitive(bitcast_tensor(create_dispatch_buffer( num_wgs.clone(), [BLOCK_SIZE, 1, 1], ))) * Tensor::from_ints([BIN_COUNT, 1, 1], device);
  • num_wgs是由n_sortcreate_dispatch_buffer生成的间接调度缓冲(见下文第四节);
  • num_reduce_wgs = num_wgs × BIN_COUNT为规约阶段需要的 workgroup 总数;
  • 缓冲区的分配按max_n(张量声明长度)计算,保证任何实际元素数都不会越界,而实际计算量由 GPU 侧动态截断——这是典型的“按最大分配、按实际调度”GPU 编程范式。

每趟循环(lib.rs)依次:

  1. create_uniform_buffer写入Uniforms { shift: pass * 4 },即本趟取出的比特窗口起点;
  2. 分配count_buf(大小max_needed_wgs * 16)与reduced_buf(大小BLOCK_SIZE);
  3. CubeCount::Dynamic依次派发SortCountSortReduceSortScanAddSortScatter,其中SortScanCubeCount::Static(1, 1, 1)单 workgroup 执行;
  4. 每趟结束后把输出张量交给下一趟作为输入(cur_keys = output_keys; cur_vals = output_values;),直到所有位趟完成。

这些 shader 均通过kernel_source_gen!宏(基于brush-wgsl的 naga_oil 编译链)在构建期生成,sort_countsort_scatter共享Uniforms { shift }结构,见 mod.rs。

四、间接调度(Indirect Dispatch)机制

“对 GPU 已知数量的元素排序”是 README 强调的核心能力,其底层是 brush-kernel 提供的create_dispatch_buffer

pub fn create_dispatch_buffer<R: JitRuntime>( thread_nums: JitTensor<R, u32>, wg_size: [u32; 3], ) -> JitTensor<R, u32>

它通过一个单 workgroup 的CreateDispatchBuffer内核,把thread_nums(如n_sort张量)与 workgroup 尺寸换算成(group_count_x, 1, 1)三元组写入缓冲,随后client.execute_unchecked(...)CubeCount::Dynamic(...)派发,GPU 驱动会在设备端读取该缓冲再决定实际派发多少个 workgroup。整个链路 CPU 侧零同步,元素数量完全由 GPU 流水线内部决定——这正是 brush-splat 渲染管线中“可见高斯数量由上一阶段内核算出、下一阶段排序内核直接消费”的设计基础。

配套工具函数还包括:

  • bitcast_tensor(brush-kernel/src/lib.rs):仅重解释张量数据类型而不做真实转换,排序前后用它在u32i32/f32之间切换视图;
  • create_uniform_buffer(brush-kernel/src/lib.rs):把Uniforms { shift }等 POD 结构按字节写入缓冲区,供 WGSL@group(0) @binding(0)读取;
  • create_tensor(brush-kernel/src/lib.rs):从 client 按形状预留存储缓冲。

五、测试验证:从 15 元素小数组到万级高斯数据

lib.rs 内置了两个#[cfg(test)]单元测试,验证正确性:

  • test_sorting:循环 128 次,每次用 15 个元素的键数组(包含大数2^24 + 123、重复键6123等边界情况),以radix_argsort(keys, values, num_points, 32)全 32 位排序,随后与 Rust 标准库argsort参考实现逐元素比对键与值;
  • test_sorting_big:模拟“若干高斯分布区间”的 10000 条随机数据(rng.gen_range(i..i+150)产生起始点、rng.gen_range(start..start+250)产生区间长度,概率性产出键),验证大规模数据下的排序正确性。

测试中的values采用key * 2 + 5的确定性映射,可同时严格验证“键值对一起重排”的伴随值语义。运行方式为标准的cargo test -p brush-sort

六、实战案例:brush-render 中的深度排序

brush-sort在仓库中的实际消费方是 brush_splat/crates/brush-render/src/render.rs:

let (_, global_from_compact_gid) = tracing::info_span!("DepthSort", sync_burn = true) .in_scope(|| { // Interpret the depth as a u32. This is fine for a radix sort, as long as the depth > 0.0, // which we know to be the case given how we cull splats. radix_argsort( bitcast_tensor(depths), global_from_presort_gid, num_visible.clone(), 32, ) });

这里把每个高斯溅射的深度值depths重解释为u32键、把溅射的全局 ID 张量global_from_presort_gid作为伴随值,num_visible(从投影阶段 uniforms 中读出的可见溅射数,GPU 端计算)作为n_sort,一次 32 位radix_argsort即可得到按深度重排后的紧凑 ID 列表,供后续ProjectVisible等阶段按从前到后的顺序处理。注释中的关键前提是:由于已做视锥剔除,深度恒为正值,位重解释不会破坏排序单调性。在 render.rs 处还存在另一处调用,说明该原语在渲染管线中被复用于不同的排序需求。

七、小结与扩展阅读

brush-sort用约五段 WGSL shader 与一个薄薄的 Rust 入口,完整实现了可复用的 GPU 基数排序:以 4 位为基数的 LSD 多趟排序、sorting_bits位宽裁剪、以及基于CubeCount::Dynamic的间接调度,使其能在不打断 GPU 流水线的前提下处理运行时才知道长度的数组。它既是理解 FidelityFX Radix Sort 思路的极佳精简参考实现,也是 brush-splat 渲染管线中深度排序的直接基石。

若想继续深入,可按以下路径阅读:

  • 排序常量与数学基础:shaders/sorting.wgsl
  • 单趟五段 shader 的完整实现:sort_count.wgsl、sort_reduce.wgsl、sort_scan.wgsl、sort_scan_add.wgsl、sort_scatter.wgsl
  • 入口与测试:src/lib.rs
  • 间接调度与缓冲工具:brush-kernel/src/lib.rs
  • 上层渲染调用场景:brush-render/src/render.rs

【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询