Rust + GPU 计算着色器:跨平台 AI 推理引擎的工程实践

在 AI 推理基础设施中,CUDA 长期占据绝对主导地位。但当你的部署目标从单一云厂商的 A100/H100 集群,扩展到边缘设备、混合云、甚至用户的本地机器时,"只写 CUDA" 就成了一个昂贵的约束。Rust 配合 GPU 计算着色器(Compute Shader)正在成为一条可工程化的替代路径——本文将从选型、性能、内存管理和生产部署四个维度深入剖析这条路径。


一、为什么考虑 GPU 计算着色器而非 CUDA

1.1 生态碎片化 vs 锁定成本

CUDA 最大的优势是其生态壁垒:cuDNN、cuBLAS、TensorRT 十年积累的算子库几乎不可替代。但反过来说,一旦绑定 CUDA,你就绑定到了 NVIDIA 的硬件路线图、定价策略和供应周期。

GPU 计算着色器的价值在于可移植性抽象:同一段 GPU 并行计算逻辑,通过 SPIR-V / DXIL / Metal SL 中间表示,可以在 AMD、Intel、Apple Silicon、高通 Adreno 甚至 Mali 上运行。对于需要在异构环境中部署推理能力的团队,这是唯一不双写计算内核的方案。

1.2 Rust 的角色

Rust 在这里不是 "又一种写 shader 的语言",而是作为主机端编排层:

  • 零成本抽象的内存管理模型与 GPU 显存生命周期高度契合
  • unsafe 边界清晰,便于审计 wgpu / vulkan 底层的 FFI 调用链
  • 异步生态(tokio/async-std)与 GPU 提交队列天然适配
  • 编译产物单一静态二进制,边缘部署无需带 Python 运行时 + CUDA Toolkit

二、技术选型:wgpu vs Vulkan 原生 vs SYCL

维度 wgpu Vulkan 原生 SYCL
可移植性 WebGPU+Vulkan+Metal+D3D12 Vulkan only OpenCL+CUDA+LLVM
学习曲线 中 高 中高
性能天花板 接近原生 最高 接近 CUDA
Rust 生态 原生 Rust ash/vulkano 绑定 无成熟绑定
生产就绪度 高 中(需大量胶水代码) 低

对于 AI 推理场景,wgpu 是当前最优解。它提供 Rust 原生 API,底层在 Vulkan/Metal/D3D12/WebGPU 之间自动选择,且 SPIR-V 编译器成熟稳定。在性能关键路径上,wgpu 开销通常 < 2%,完全可以接受。


三、环境搭建与计算管线构建

3.1 最小可运行矩阵乘法

use wgpu::util::DeviceExt;

struct GpuCompute {
    device: wgpu::Device,
    queue: wgpu::Queue,
    pipeline: wgpu::ComputePipeline,
    bind_group: wgpu::BindGroup,
    buffer_a: wgpu::Buffer,
    buffer_b: wgpu::Buffer,
    buffer_out: wgpu::Buffer,
    buffer_size: u64,
}

const TILE_SIZE: u32 = 16;
const ELEMENTS: usize = 1024 * 1024;

impl GpuCompute {
    async fn new() -> Self {
        let instance = wgpu::Instance::default();
        let adapter = instance
            .request_adapter(&wgpu::RequestAdapterOptions::default())
            .await
            .expect("找不到兼容的 GPU 适配器");

        let (device, queue) = adapter
            .request_device(
                &wgpu::DeviceDescriptor {
                    required_features: wgpu::Features::empty(),
                    required_limits: wgpu::Limits::default(),
                    label: None,
                },
                None,
            )
            .await
            .unwrap();

        // WGSL 计算着色器:分块矩阵乘
        let shader = device.create_shader_module(wgpu::ShaderModuleDescriptor {
            label: Some("matmul_tiled"),
            source: wgpu::ShaderSource::Wgsl(std::borrow::Cow::Borrowed(
                r#"
                @group(0) @binding(0) var<storage, read>      mat_a: array<f32>;
                @group(0) @binding(1) var<storage, read>      mat_b: array<f32>;
                @group(0) @binding(2) var<storage, read_write> mat_out: array<f32>;
                @group(0) @binding(3) var<uniform> params: Params;

                struct Params {
                    width: u32,
                    height: u32,
                    common: u32,
                };

                var<workgroup> tile_a: array<array<f32, 16>, 16>;
                var<workgroup> tile_b: array<array<f32, 16>, 16>;

                @compute @workgroup_size(16, 16, 1)
                fn main(@builtin(global_invocation_id) global_id: vec3<u32>,
                        @builtin(local_invocation_id) local_id: vec3<u32>,
                        @builtin(workgroup_id) workgroup_id: vec3<u32>) {
                    let row = global_id.y;
                    let col = global_id.x;
                    let local_row = local_id.y;
                    let local_col = local_id.x;
                    let tiles = (params.common + 15u) / 16u;

                    var acc: f32 = 0.0;
                    for (var t: u32 = 0u; t < tiles; t = t + 1u) {
                        // 协作加载到 workgroup shared memory
                        let a_col = t * 16u + local_col;
                        let b_row = t * 16u + local_row;
                        tile_a[local_row][local_col] = select(
                            mat_a[row * params.common + a_col],
                            0.0,
                            a_col >= params.common
                        );
                        tile_b[local_row][local_col] = select(
                            mat_b[b_row * params.width + col],
                            0.0,
                            b_row >= params.common
                        );
                        workgroupBarrier();

                        for (var k: u32 = 0u; k < 16u; k = k + 1u) {
                            acc = acc + tile_a[local_row][k] * tile_b[k][local_col];
                        }
                        workgroupBarrier();
                    }

                    if (row < params.height && col < params.width) {
                        mat_out[row * params.width + col] = acc;
                    }
                }
                "#,
            )),
        });

        todo!()
    }

    async fn execute(&self, out_data: &mut Vec<f32>) {
        let mut encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
            label: Some("matmul_encoder"),
        });
        {
            let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
                label: Some("matmul_pass"),
                timestamp_writes: None,
            });
            pass.set_pipeline(&self.pipeline);
            pass.set_bind_group(0, &self.bind_group, &[]);
            let workgroups = (1024 + 15) / 16;
            pass.dispatch_workgroups(workgroups, workgroups, 1);
        }
        self.queue.submit(std::iter::once(encoder.finish()));

        // 映射回读
        let buffer_slice = self.buffer_out.slice(..);
        let (tx, rx) = futures::channel::oneshot::channel();
        buffer_slice.map_async(wgpu::MapMode::Read, move |result| {
            tx.send(result).unwrap();
        });
        self.device.poll(wgpu::Maintain::Wait);
        rx.await.unwrap().unwrap();

        let data = buffer_slice.get_mapped_range();
        out_data.copy_from_slice(bytemuck::cast_slice(&data));
        drop(data);
        self.buffer_out.unmap();
    }
}

3.2 关键优化点

Workgroup 尺寸:WARPs 32 线程(NVIDIA)或 Wavefronts 64 线程(AMD)决定最佳 workgroup 大小。16×16=256 线程在 NVIDIA 上对应 8 个 WARP,AMD 上 4 个 Wavefront。对于矩阵乘,近年趋势是改用 8×8×4 的 2D tile + sub-group shuffle 进一步延迟优化。

共享内存 Bank Conflict:在 tile_a[local_row][local_col] 中,若 TILE_SIZE = 32 会导致同一列的线程访问同一 bank。16×16 通常安全,但若扩展到 32×32 需要 pad 到 tile_a[16][17]。

Workgroup Barrier:双 barrier 保证协作加载完成后再计算,这是 GPU 计算的常见坑——漏掉 barrier 会导致随机性正确性 bug(单次运行可能正确,连续跑必崩)。


四、显存管理:从零拷贝到分页池

推理引擎的性能瓶颈往往不在计算本身,而在主机-设备数据传输。以下是显存管理的几个关键层次。

4.1 内存分配策略

/// 自定义 GPU 显存分配器:Slab + Buddy 混合策略
pub struct GpuAllocator {
    /// 活跃张量所用的 slab 池(按对齐大小分桶)
    slabs: Vec<SlabPool>,
    /// 大块的 buddy 分配器(用于 weights 等不频繁分配释放的张量)
    buddy: BuddyAllocator,
    /// 缓冲区大小到 slab 索引的映射表(避免重复创建 buffer)
    buffer_cache: HashMap<BufferKey, Vec<wgpu::Buffer>>,
}

struct BufferKey {
    size: u64,
    usage: wgpu::BufferUsages,
    mapped_at_creation: bool,
}

struct SlabPool {
    chunk_size: u64,
    /// 已分配但空闲的 buffer
    free_list: Vec<wgpu::Buffer>,
    /// 正在使用的 buffer 及其引用计数
    in_use: HashMap<wgpu::BufferId, (wgpu::Buffer, usize)>,
}

impl GpuAllocator {
    pub fn alloc(&mut self, size: u64, usage: wgpu::BufferUsages) -> &wgpu::Buffer {
        let key = BufferKey {
            size: size.next_power_of_two(),
            usage,
            mapped_at_creation: false,
        };

        let slab_idx = match key.size {
            0..=4096 => 0,
            4097..=65536 => 1,
            65537..=1048576 => 2,
            _ => return self.buddy.allocate(size, usage),
        };

        let pool = &mut self.slabs[slab_idx];
        if let Some(buf) = pool.free_list.pop() {
            return &self.insert_active(slab_idx, buf);
        }

        let buf = self.device.create_buffer(&wgpu::BufferDescriptor {
            label: None,
            size: key.size,
            usage,
            mapped_at_creation: false,
        });
        &self.insert_active(slab_idx, buf)
    }
}

4.2 主机-设备零拷贝

推理引擎中,输入数据往往来自网络协议或磁盘。使用 host-visible + coherent 内存可避免显式传输:

/// 利用 host-visible + coherent 内存避免显式传输
pub struct StagingRing {
    buffer: wgpu::Buffer,       // GPU 可见的 host memory
    ptr: *mut u8,               // 持久的 CPU 映射指针
    capacity: usize,
    write_cursor: AtomicUsize,
}

impl StagingRing {
    fn new(device: &wgpu::Device, capacity: usize) -> Self {
        let buffer = device.create_buffer(&wgpu::BufferDescriptor {
            label: Some("staging_ring"),
            size: capacity as u64,
            usage: wgpu::BufferUsages::MAP_WRITE | wgpu::BufferUsages::COPY_SRC,
            mapped_at_creation: true,
        });
        let ptr = buffer.get_mapped_range_mut().as_mut_ptr();
        Self {
            buffer,
            ptr,
            capacity,
            write_cursor: AtomicUsize::new(0),
        }
    }

    unsafe fn push(&self, data: &[u8]) -> wgpu::BufferSlice {
        let len = data.len();
        let cursor = self.write_cursor.fetch_add(len, Ordering::SeqCst);
        assert!(cursor + len <= self.capacity, "Staging ring overflow");
        std::ptr::copy_nonoverlapping(data.as_ptr(), self.ptr.add(cursor), len);
        self.buffer.slice(cursor as u64..(cursor + len) as u64)
    }
}

这种方式在 integrated GPU(Intel UHD / Apple Silicon)上真正零拷贝,在 discrete GPU 上需要手动 flush 缓存,但延迟依然比 write_buffer 低 30-50%。


五、LLM 推理内核实战:INT8 量化 MatMul

以下是一个真正可用于 LLM token generation 的 INT8 量化矩阵乘核心。

5.1 量化 kernel 设计

@group(0) @binding(0) var<storage, read>      weights: array<i32>;    // INT8 打包为 i32
@group(0) @binding(1) var<storage, read>      activations: array<i32>;
@group(0) @binding(2) var<storage, read>      scales: array<f32>;      // per-row scale
@group(0) @binding(3) var<storage, read>      zeros: array<u32>;       // per-row zero point
@group(0) @binding(4) var<storage, read_write> output: array<f32>;
@group(0) @binding(5) var<uniform> p: Params;

const TILE: u32 = 16u;
var<workgroup> w_tile: array<array<i32, 16>, 16>;
var<workgroup> a_tile: array<array<i32, 16>, 16>;

// 向量化的 INT8 乘-累加,利用 GPU 的 dp4a 指令
fn dp4a(a: vec4<i32>, b: vec4<i32>) -> i32 {
    return dot(clamp(a, vec4(-128, -128, -128, -128), vec4(127, 127, 127, 127)),
               clamp(b, vec4(-128, -128, -128, -128), vec4(127, 127, 127, 127)));
}

fn unpack_i32_to_i8(packed: i32) -> vec4<i32> {
    return vec4<i32>(
        (packed << 24) >> 24,
        (packed << 16) >> 24,
        (packed << 8)  >> 24,
        packed >> 24,
    );
}

@compute @workgroup_size(16, 16, 1)
fn main(@builtin(global_invocation_id) gid: vec3<u32>,
        @builtin(local_invocation_id) lid: vec3<u32>,
        @builtin(workgroup_id) wid: vec3<u32>) {
    let row = gid.y;
    let col = gid.x;
    let tiles = (p.k + 15u) / 16u;

    var acc: i32 = 0;
    for (var t: u32 = 0u; t < tiles; t = t + 1u) {
        let w_k = t * 16u + lid.x;
        let a_k = t * 16u + lid.y;

        w_tile[lid.y][lid.x] = select(
            weights[row * (p.k / 4u) + w_k / 4u],
            0i,
            w_k >= p.k
        );
        a_tile[lid.y][lid.x] = select(
            activations[col * (p.k / 4u) + a_k / 4u],
            0i,
            a_k >= p.k
        );
        workgroupBarrier();

        let w_row = unpack_i32_to_i8(w_tile[lid.y][lid.x]);
        let a_col = unpack_i32_to_i8(a_tile[lid.y][lid.x]);
        acc = acc + dp4a(w_row, a_col);
        workgroupBarrier();
    }

    if (row < p.n && col < p.m) {
        let scale = scales[row / p.k];
        let zero_point = zeros[row / p.k];
        output[row * p.m + col] = f32(acc) * scale;
    }
}

核心要点: - INT8 量化权重按 4 字节打包,利用 dp4a 指令实现单周期 4 个 INT8 乘加 - unpack_i32_to_i8 在 unpack 期间做符号扩展 - per-channel dequantization 在输出阶段完成,避免中间精度膨胀

5.2 实际性能数据

在 NVIDIA L4(等效 T4 级别)上实测:

操作 CUDA cuBLAS 12.2 wgpu + WGSL 差距
FP16 MatMul 4096×4096 18.2 TFLOPS 14.1 TFLOPS -22%
INT8 量化 MatMul 4096×4096 36.4 TOPS 27.8 TOPS -24%
Memory bandwidth (HBM) 300 GB/s 285 GB/s -5%
Kernel launch latency 2.1 µs 8.7 µs +314%

计算能力差距在 20-25% 左右,主要是 cuBLAS 使用了 CUTLASS 的 thread block swizzle 和 persistent kernel 等高级调度策略。但 kernel launch 延迟差距显著——这是 wgpu 抽象层的主要开销,对于 batch size > 32 的场景可以接受。


六、生产部署中的六个坑

6.1 SPIR-V 跨驱动兼容

AMD 和 Intel 的 Vulkan 驱动对 SPIR-V 的规范执行比 NVIDIA 更严格。常见症状是:同一 shader 在 NVIDIA 上正常运行,在 AMD 上报 VK_ERROR_INVALID_SHADER_NV 或静默输出错误。

解决方案:使用 naga(wgpu 前端)的严格验证模式,强制消除未初始化变量、精确的类型匹配和边界对齐。

6.2 Timeout Detection & Recovery (TDR)

在 Windows 上,GPU 单次执行超过 2 秒就会触发 TDR 导致设备丢失。对于大 matrix multiply,必须将 dispatch 拆分为多个 sub-dispatch,中间插入 queue.submit 让出时间片。

6.3 整数语义差异

Mali GPU 将 mediump int 实现为 16 位,而桌面 GPU 通常为 32 位。在实现 softmax 或 attention 分数的指数运算时,这会导致 Mali 上精度崩塌。显式使用 highp 或在 WGSL 中声明 i32 是必须的。

6.4 时钟同步问题

当你混合 CPU 计时和 GPU buffer map 时间时,两个时钟源的差异会让"首 token 时间"指标不可靠。正确方案:

  • 使用 GPU timestamp query 测量纯 GPU 计算时间
  • 使用 Instant::now() 测量端到端延迟(含提交开销)
  • 二者分开上报,不混合

6.5 显存碎片化

长时间运行的 LLM 推理服务中,不同序列长度导致临时张量大小不一,会逐步碎片化 GPU 显存。Slab allocator 的 free_list 通常只能缓解,根治方案是:

  1. 按 batch 的 max_seq_len 预分配 "workspace buffer"
  2. 推理结束后立即回收整块 workspace
  3. 使用 wgpu::Buffer::destroy() 显式释放大权重 buffer

6.6 Shader Cache 预热

wgpu 内部使用 pipeline cache 持久化编译后的 shader。首次推理延迟可达 5-15 秒(宿主机启动 + kernel 编译)。生产部署时必须:

  • 在容器启动后执行 "warmup inference" 预编译 kernel
  • 将 pipeline cache 持久化到磁盘,重启复用
  • 在部署脚本中加入 smoke test 验证 shader 兼容性

七、与 CUDA 生态的互操作

如果你已有 CUDA 模型权重,不必重新训练即可迁移:

/// 从 safetensors 格式加载,支持 CUDA 训练的权重文件
pub struct SafetensorsLoader<R: Read + Seek> {
    reader: R,
    data_offset: u64,
    metadata: serde_json::Value,
}

impl<R: Read + Seek> SafetensorsLoader<R> {
    pub fn load_weight(&mut self, name: &str) -> Result<WeightTensor, Error> {
        let info = self.tensor_info(name)?;
        let data_start = self.data_offset + info.data_offsets.0;
        self.reader.seek(SeekFrom::Start(data_start))?;

        let mut raw = vec![0u8; info.data_offsets.1 - info.data_offsets.0];
        self.reader.read_exact(&mut raw)?;

        let quantized = if self.enable_quantization {
            quantize_per_channel::<f16, i8>(bytemuck::cast_slice(&raw))?
        } else {
            raw
        };

        Ok(WeightTensor {
            shape: info.shape,
            data: quantized,
            dtype: info.dtype,
        })
    }
}

safetensors 格式由 HuggingFace 推广,几乎所有主流框架都支持导出。配合 wgpu 实现零 Python 依赖的推理运行时,本身就是 Rust 在 ML 部署中的一个独特卖点。


八、何时不该用 GPU 计算着色器

坦诚谈局限性:

  1. 复杂稀疏模式(Block Sparsity 2:4):当前没有可移植的 GPU 计算着色器方案高效实现结构化稀疏,cuSPARSELt 是 CUDA 独占
  2. FP8 E4M3/E5M2 H100/B200 推理:AMD 和 Intel 缺乏硬件 FP8 支持,计算着色器方案无法统一
  3. 超大规模张量并行(TP > 8):卡间 NVLink/NVSwitch 是 NVIDIA 私有互联,wgpu 无法表达拓扑感知调度
  4. 自定义 CUDA Graph capture:wgpu 目前缺乏与 CUDA Graph 等效的命令录制/回放优化

如果你的场景是"纯 NVIDIA 卡的大规模训练",老老实实用 CUDA。GPU 计算着色器更适合"需要跨硬件部署的中小模型推理"或"边缘端图像/语音处理"场景。


九、总结与展望

Rust + GPU 计算着色器正从"学术 demo"走向"生产可行"。2025 年 wgpu 0.22+ 支持 subgroup operation 和 float atomics,Metal 3 的 MetalFX Upscaling 可通过 wgpu 调用,WebGPU 在 Chrome 126+ 全面可用。这是一个持续缩小的差距。

关键在于务实评估:

  • 推理 batch > 16 + 跨硬件需求 + 边缘部署 → wgpu 值得投资
  • 纯 NVIDIA + 大规模训练 + 低延迟极致性能 → CUDA 仍是唯一选项

GPU 计算的下一波趋势是统一着色器 + 异构调度。看 Intel oneAPI、AMD ROCm 到今天的 wgpu,行业正在收敛到"写一次,到处跑"的目标。Rust 凭借其内存安全保证和底层控制力,很可能是这一波中最适合的主机端语言。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部