Rust 实现深度学习推理引擎:从计算图到异构硬件加速

在 PyTorch 和 TensorFlow 垄断的 AI 框架领域,Rust 生态正在悄然崛起一条新路径。Candle(Hugging Face 出品)和 Burn(独立开源社区)代表了两种截然不同的设计哲学:一个是轻量级、模块化的推理引擎,另一个是重新创造轮子的全栈深度学习框架。本文将深入剖析如何用 Rust 构建一个生产级推理引擎的核心组件,涵盖张量抽象、计算图 JIT 编译、自动微分以及异构硬件后端调度。

一、为什么用 Rust 做深度学习

C++ 统治了深度学习框架的底层实现三十年,但这几年 Rust 开始在这个领域发力。原因很简单:

  1. 内存安全无代价:张量操作涉及大量临时内存分配和释放,Rust 的所有权系统可以在编译期消除 use-after-free 和数据竞争
  2. 零成本抽象:计算图中的算子融合(operator fusion)依赖泛型和内联优化,Rust 的表现不输 C++
  3. PyO3 胶水友好:可以通过 PyO3 暴露 Python 接口,让用户无感地替换 PyTorch 后端
  4. 优秀的并发模型:跨设备(CPU/GPU)数据传输和计算重叠需要精细的异步控制

Hugging Face 选择 Rust 重写 tokenizers 库后获得 2-5x 的性能提升,这证明了 Rust 在 AI 基础设施领域的潜力。Candle 作为其后续项目,目标是在推理场景下替代 PyTorch。

二、张量核心:Device + DType + Layout

一切从张量开始。一个生产级引擎需要精确区分数据类型和物理设备。

use std::fmt;

#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum DType {
    F16,   // IEEE 754 half-precision
    BF16,  // Brain floating point
    F32,   // IEEE 754 single
    F64,   // IEEE 754 double
    I32,   // 32-bit signed integer
    I64,   // 64-bit signed integer
    U8,    // 8-bit unsigned (quantized weights)
}

impl DType {
    pub fn size_in_bytes(&self) -> usize {
        match self {
            DType::F16 | DType::BF16 => 2,
            DType::F32 | DType::I32 => 4,
            DType::F64 | DType::I64 => 8,
            DType::U8 => 1,
        }
    }
}

#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Device {
    Cpu,
    Cuda { device_id: u32 },
    Metal { device_id: u32 },
}

impl fmt::Display for Device {
    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
        match self {
            Device::Cpu => write!(f, "cpu"),
            Device::Cuda { device_id } => write!(f, "cuda:{}", device_id),
            Device::Metal { device_id } => write!(f, "metal:{}", device_id),
        }
    }
}

接下来是存储层。我们需要一个能够跨设备共享内存的缓冲区抽象:

pub enum Storage {
    Cpu(Vec<u8>),
    Cuda(CudaBuffer),
    Metal(MetalBuffer),
}

pub struct Tensor {
    storage: Storage,
    shape: Vec<usize>,
    strides: Vec<usize>,
    dtype: DType,
    device: Device,
}

impl Tensor {
    pub fn new(data: Vec<f32>, shape: Vec<usize>, device: Device) -> Result<Self, TensorError> {
        let expected_len: usize = shape.iter().product();
        if data.len() != expected_len {
            return Err(TensorError::ShapeMismatch {
                expected: expected_len,
                got: data.len(),
            });
        }

        let strides = Self::compute_contiguous_strides(&shape);
        let storage = match device {
            Device::Cpu => {
                let bytes: Vec<u8> = data.into_iter()
                    .flat_map(|f| f.to_ne_bytes())
                    .collect();
                Storage::Cpu(bytes)
            }
            Device::Cuda { device_id } => {
                Storage::Cuda(CudaBuffer::from_slice(&data, device_id)?)
            }
            Device::Metal { device_id } => {
                Storage::Metal(MetalBuffer::from_slice(&data, device_id)?)
            }
        };

        Ok(Self { storage, shape, strides, dtype: DType::F32, device })
    }

    pub fn shape(&self) -> &[usize] { &self.shape }
    pub fn device(&self) -> &Device { &self.device }
    pub fn dtype(&self) -> DType { self.dtype }

    fn compute_contiguous_strides(shape: &[usize]) -> Vec<usize> {
        let mut strides = vec![1usize; shape.len()];
        for i in (0..shape.len()-1).rev() {
            strides[i] = strides[i + 1] * shape[i + 1];
        }
        strides
    }

    /// 转置操作:只修改 strides,零拷贝
    pub fn transpose(&self, dim0: usize, dim1: usize) -> Result<Self, TensorError> {
        if dim0 >= self.shape.len() || dim1 >= self.shape.len() {
            return Err(TensorError::DimensionOutOfBounds);
        }
        let mut result = self.clone_metadata_only();
        result.shape.swap(dim0, dim1);
        result.strides.swap(dim0, dim1);
        Ok(result)
    }

    fn clone_metadata_only(&self) -> Self {
        Self {
            storage: match &self.storage {
                Storage::Cpu(buf) => Storage::Cpu(buf.clone()),
                Storage::Cuda(buf) => Storage::Cuda(buf.clone()),
                Storage::Metal(buf) => Storage::Metal(buf.clone()),
            },
            shape: self.shape.clone(),
            strides: self.strides.clone(),
            dtype: self.dtype,
            device: self.device,
        }
    }
}

这个设计有几个关键决策:

  • Storage 枚举将设备差异封装在内部,上层算子代码通过统一接口操作
  • 转置不拷贝数据,只修改 strides 语义
  • dtype + device 在编译期用于零成本分发(配合宏或泛型 specialization)

三、计算图:节点、边与拓扑排序

现代推理引擎的核心是计算图(computational graph)。我们需要将用户的高层表达式(如 y = softmax(Wx + b))转化为可在异构硬件上执行的有向无环图(DAG)。

use std::collections::{HashMap, HashSet};
use std::sync::Arc;

pub type NodeId = usize;

#[derive(Debug, Clone)]
pub enum Op {
    // 基础运算
    Add,
    MatMul,
    Mul,
    Exp,
    Log,
    Softmax { axis: usize },
    LayerNorm { eps: f32 },
    
    // 激活函数
    Relu,
    Gelu,
    Silu,     // SwiGLU / SiLU,LLaMA 标配
    
    // 内存操作
    Reshape { shape: Vec<usize> },
    Transpose { dim0: usize, dim1: usize },
    Concat { axis: usize },
    
    // 数据加载
    Literal(Tensor),      // 常量/权重
    Parameter(String),    // 可训练参数
    Input(String),        // 外部输入
}

#[derive(Debug)]
pub struct ComputeGraph {
    pub nodes: HashMap<NodeId, Op>,
    pub edges: Vec<(NodeId, NodeId)>,   // (from, to)
    pub inputs: HashMap<NodeId, Vec<NodeId>>,   // node -> its input dependencies
    pub outputs: HashMap<NodeId, Vec<NodeId>>,  // node -> nodes consuming its output
    next_id: NodeId,
}

impl ComputeGraph {
    pub fn new() -> Self {
        Self {
            nodes: HashMap::new(),
            edges: Vec::new(),
            inputs: HashMap::new(),
            outputs: HashMap::new(),
            next_id: 0,
        }
    }

    pub fn add_node(&mut self, op: Op) -> NodeId {
        let id = self.next_id;
        self.next_id += 1;
        self.nodes.insert(id, op);
        self.inputs.insert(id, Vec::new());
        self.outputs.insert(id, Vec::new());
        id
    }

    pub fn add_edge(&mut self, from: NodeId, to: NodeId) {
        self.edges.push((from, to));
        self.outputs.get_mut(&from).unwrap().push(to);
        self.inputs.get_mut(&to).unwrap().push(from);
    }

    /// Kahn 算法进行拓扑排序
    pub fn topological_sort(&self) -> Result<Vec<NodeId>, GraphError> {
        let mut in_degree: HashMap<NodeId, usize> = HashMap::new();
        for (node, deps) in &self.inputs {
            in_degree.insert(*node, deps.len());
        }

        let mut queue: Vec<NodeId> = in_degree.iter()
            .filter(|(_, &deg)| deg == 0)
            .map(|(&node, _)| node)
            .collect();

        let mut sorted = Vec::new();
        while let Some(node) = queue.pop() {
            sorted.push(node);
            if let Some(children) = self.outputs.get(&node) {
                for child in children {
                    let deg = in_degree.get_mut(child).unwrap();
                    *deg -= 1;
                    if *deg == 0 {
                        queue.push(*child);
                    }
                }
            }
        }

        if sorted.len() != self.nodes.len() {
            return Err(GraphError::CycleDetected);
        }
        Ok(sorted)
    }
}

四、算子融合:性能翻倍的关键

朴素执行计算图中每个节点的方式会产生大量中间张量,导致内存带宽成为瓶颈。算子融合(Operator Fusion)是推理优化的灵魂——将多个连续操作合并成一个内核调用。

考虑 Transformer 中典型的一行计算:output = softmax(Q @ K^T / sqrt(d_k)) @ V

如果逐算子执行,会产生 QK^T、scaled、softmax_out、attn_out 四个中间张量。融合后只需一个 Flash Attention 内核:

pub enum FusedOp {
    /// 将 LayerNorm 的线性变换直接融入下一层 MatMul
    /// 即: x_w = (x - mean) / std * gamma + bias; y = x_w @ W
    /// 融合为: y = ((x - mean) / std * gamma + bias) @ W
    {
        norm_params: LayerNormParams,
        matmul_params: MatMulParams,
    },
    
    /// GeLU 激活融合:y = 0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3)))
    /// 与 MatMul 融合可避免一次全局内存往返
    MatMulBiasAct {
        trans_b: bool,
        activation: Activation,
    },
    
    /// Flash Attention 的 QKV 融合
    FlashAttention {
        causal: bool,
        softmax_scale: f32,
    },
}

#[derive(Debug, Clone, Copy)]
pub enum Activation {
    None,
    Relu,
    Gelu,
    SiLU,
}

pub fn fuse_adjacent_ops(graph: &ComputeGraph, execution_order: &[NodeId]) -> Vec<FusedGroup> {
    let mut groups: Vec<FusedGroup> = Vec::new();
    let mut current_group: Vec<NodeId> = Vec::new();
    let mut i = 0;

    while i < execution_order.len() {
        let node_id = execution_order[i];
        let op = graph.nodes.get(&node_id).unwrap();

        if can_be_fused(op) {
            current_group.push(node_id);
        } else {
            if !current_group.is_empty() {
                if let Some(fused) = try_fuse_group(graph, &current_group) {
                    groups.push(fused);
                } else {
                    groups.push(FusedGroup::Unfused(current_group.clone()));
                }
                current_group.clear();
            }
            groups.push(FusedGroup::Single(node_id, op.clone()));
        }
        i += 1;
    }

    // 处理最后一组
    if !current_group.is_empty() {
        if let Some(fused) = try_fuse_group(graph, &current_group) {
            groups.push(fused);
        } else {
            groups.push(FusedGroup::Unfused(current_group));
        }
    }

    groups
}

fn can_be_fused(op: &Op) -> bool {
    matches!(op, 
        Op::LayerNorm { .. } | Op::Add | Op::Mul | Op::Gelu | Op::Relu | 
        Op::Silu | Op::Exp | Op::Softmat { .. }
    )
}

enum FusedGroup {
    Fused(Vec<NodeId>, FusedOp),
    Unfused(Vec<NodeId>),
    Single(NodeId, Op),
}

五、内存规划:arena 分配与 In-Place 优化

深度学习推理的最大内存开销不是模型权重,而是激活值(activations)。一个好的引擎需要在计算开始前规划好每块张量的内存偏移,复用不再需要的中间结果:

pub struct MemoryArena {
    buffer: Vec<u8>,
    capacity: usize,
    allocations: Vec<Allocation>,
    lru: LruCache<NodeId, usize>,  // node_id -> allocation slot
}

struct Allocation {
    offset: usize,
    size: usize,
    /// 最后一个使用这个 slot 的节点
    last_use: NodeId,
    /// 是否可复用
    reusable: bool,
}

pub struct MemoryPlanner<'a> {
    graph: &'a ComputeGraph,
    order: Vec<NodeId>,
    dtype_sizes: HashMap<NodeId, usize>,
}

impl<'a> MemoryPlanner<'a> {
    /// 基于线性扫描的最小内存规划
    /// 策略:按拓扑序分配,当某节点的所有消费者完成后,释放其 slot 供复用
    pub fn plan(&self) -> MemoryPlan {
        let mut plan = MemoryPlan::default();
        let mut live_slots: HashMap<NodeId, usize> = HashMap::new();
        let mut free_slots: Vec<usize> = Vec::new();
        let mut next_slot: usize = 0;

        for (step, &node_id) in self.order.iter().enumerate() {
            // 检查输入是否都可以释放
            if let Some(inputs) = self.graph.inputs.get(&node_id) {
                for input in inputs {
                    let consumers = self.graph.outputs.get(input).unwrap();
                    let all_consumed = consumers.iter().all(|c| {
                        self.order.iter().position(|x| x == c).unwrap() <= step
                    });
                    if all_consumed {
                        if let Some(slot) = live_slots.remove(input) {
                            free_slots.push(slot);
                        }
                    }
                }
            }

            // 为当前节点分配 output slot
            let needed_size = self.dtype_sizes[&node_id];
            let slot = if let Some(reusable) = free_slots.iter()
                .find(|&&s| plan.slot_size(s) >= needed_size)
                .copied()
            {
                // 复用已有 slot
                free_slots.retain(|&s| s != reusable);
                reusable
            } else {
                let s = next_slot;
                next_slot += 1;
                plan.ensure_capacity(s, needed_size);
                s
            };

            live_slots.insert(node_id, slot);
            plan.node_slots.insert(node_id, slot);
            plan.slot_sizes.insert(slot, plan.slot_size(slot).max(needed_size));
        }

        plan
    }
}

#[derive(Default)]
pub struct MemoryPlan {
    pub node_slots: HashMap<NodeId, usize>,
    pub slot_sizes: HashMap<NodeId, usize>,
}

impl MemoryPlan {
    fn slot_size(&self, slot: usize) -> usize {
        self.slot_sizes.get(&slot).copied().unwrap_or(0)
    }

    fn ensure_capacity(&mut self, slot: usize, size: usize) {
        let entry = self.slot_sizes.entry(slot).or_default();
        *entry = (*entry).max(size);
    }

    pub fn total_memory_required(&self) -> usize {
        self.slot_sizes.values().sum()
    }
}

六、多后端调度:统一 IR + 各平台特化

生产级推理引擎需要同时支持 CPU、CUDA、Metal、甚至 WebGPU。核心思路是定义中间表示(IR),然后每个后端提供特化实现:

/// 平台无关的算子中间表示
pub enum IrOp {
    Binary {
        kind: BinaryKind,
        lhs: IrTensor,
        rhs: IrTensor,
        out: IrTensor,
    },
    Unary {
        kind: UnaryKind,
        input: IrTensor,
        out: IrTensor,
    },
    MatMul {
        a: IrTensor,
        b: IrTensor,
        c: IrTensor,
        trans_a: bool,
        trans_b: bool,
    },
    Reduce {
        kind: ReduceKind,
        axis: usize,
        input: IrTensor,
        out: IrTensor,
    },
    Copy {
        src: IrTensor,
        dst: IrTensor,
    },
}

pub trait Backend: Send + Sync {
    fn name(&self) -> &str;
    fn supports_dtype(&self, dtype: DType) -> bool;
    fn execute(&self, ops: &[IrOp], memory: &mut MemoryArena) -> Result<(), BackendError>;
    fn synchronize(&self) -> Result<(), BackendError>;
}

// === CPU Backend (基于 SIMD: std::simd / packed_simd) ===
pub struct CpuBackend;

impl Backend for CpuBackend {
    fn name(&self) -> &str { "cpu" }
    
    fn supports_dtype(&self, dtype: DType) -> bool {
        matches!(dtype, DType::F32 | DType::F16 | DType::I32)
    }

    fn execute(&self, ops: &[IrOp], memory: &mut MemoryArena) -> Result<(), BackendError> {
        for op in ops {
            match op {
                IrOp::Binary { kind, lhs, rhs, out } => {
                    cpu_execute_binary(*kind, lhs, rhs, out, memory)?;
                }
                IrOp::MatMul { a, b, c, trans_a, trans_b } => {
                    cpu_matmul(a, b, c, *trans_a, *trans_b, memory)?;
                }
                IrOp::Reduce { kind, axis, input, out } => {
                    cpu_reduce(*kind, *axis, input, out, memory)?;
                }
                _ => unimplemented!("CPU op: {:?}", op),
            }
        }
        Ok(())
    }

    fn synchronize(&self) -> Result<(), BackendError> {
        // CPU 执行是同步的
        Ok(())
    }
}

#[cfg(target_arch = "x86_64")]
fn cpu_matmul(
    a: &IrTensor, b: &IrTensor, c: &IrTensor,
    trans_a: bool, trans_b: bool,
    memory: &mut MemoryArena,
) -> Result<(), BackendError> {
    use std::simd::*;
    
    let m = if trans_a { a.shape[1] } else { a.shape[0] };
    let k = if trans_a { a.shape[0] } else { a.shape[1] };
    let n = if trans_b { b.shape[0] } else { b.shape[1] };
    
    let a_ptr: *const f32 = memory.tensor_ptr(a.id);
    let b_ptr: *const f32 = memory.tensor_ptr(b.id);
    let c_ptr: *mut f32 = memory.tensor_ptr_mut(c.id);
    
    const TILE_M: usize = 32;
    const TILE_N: usize = 32;
    const TILE_K: usize = 64;
    
    // 分块矩阵乘法以优化缓存局部性
    for i0 in (0..m).step_by(TILE_M) {
        for j0 in (0..n).step_by(TILE_N) {
            for k0 in (0..k).step_by(TILE_K) {
                let i_end = (i0 + TILE_M).min(m);
                let j_end = (j0 + TILE_N).min(n);
                let k_end = (k0 + TILE_K).min(k);
                
                for i in i0..i_end {
                    for j in (j0..j_end).step_by(8) {
                        unsafe {
                            let mut sum = f32x8::splat(0.0);
                            if j + 8 <= j_end {
                                for kk in k0..k_end {
                                    let a_val = if trans_a {
                                        *a_ptr.add(kk * a.shape[1] + i)
                                    } else {
                                        *a_ptr.add(i * k + kk)
                                    };
                                    let b_vec = f32x8::from_slice(
                                        std::slice::from_raw_parts(
                                            b_ptr.add(kk * n + j), 8
                                        )
                                    );
                                    sum += f32x8::splat(a_val) * b_vec;
                                }
                                sum.to_slice(std::slice::from_raw_parts_mut(
                                    c_ptr.add(i * n + j), 8
                                ));
                            } else {
                                // 边界处理
                                for jj in j..j_end {
                                    let mut sum = 0.0f32;
                                    for kk in k0..k_end {
                                        let a_val = if trans_a {
                                            *a_ptr.add(kk * a.shape[1] + i)
                                        } else {
                                            *a_ptr.add(i * k + kk)
                                        };
                                        let b_val = if trans_b {
                                            *b_ptr.add(jj * k + kk)
                                        } else {
                                            *b_ptr.add(kk * n + jj)
                                        };
                                        sum += a_val * b_val;
                                    }
                                    *c_ptr.add(i * n + jj) = sum;
                                }
                            }
                        }
                    }
                }
            }
        }
    }
    
    Ok(())
}

// === CUDA Backend (基于 cust / rustacuda) ===
#[cfg(feature = "cuda")]
pub struct CudaBackend {
    device_id: u32,
    stream: cuda::Stream,
    module_cache: HashMap<String, cuda::Module>,
}

#[cfg(feature = "cuda")]
impl Backend for CudaBackend {
    fn name(&self) -> &str { "cuda" }
    
    fn supports_dtype(&self, dtype: DType) -> bool {
        matches!(dtype, DType::F16 | DType::BF16 | DType::F32 | DType::I32)
    }

    fn execute(&self, ops: &[IrOp], _memory: &mut MemoryArena) -> Result<(), BackendError> {
        for op in ops {
            match op {
                IrOp::MatMul { a, b, c, trans_a, trans_b } => {
                    let algo = cublas::GemmAlgo::default();
                    unsafe {
                        cublas::gemm_ex(
                            self.handle,
                            if *trans_a { cublas::Operation::T } else { cublas::Operation::N },
                            if *trans_b { cublas::Operation::T } else { cublas::Operation::N },
                            c.shape[1] as i32,
                            c.shape[0] as i32,
                            a.shape[1] as i32,
                            &1.0f32,
                            b.ptr as *const f32,
                            b.shape[1] as i32,
                            a.ptr as *const f32,
                            a.shape[1] as i32,
                            &0.0f32,
                            c.ptr as *mut f32,
                            c.shape[1] as i32,
                            algo,
                        )?;
                    }
                }
                IrOp::Binary { kind, lhs, rhs, out } => {
                    self.launch_elementwise_kernel(*kind, lhs, rhs, out)?;
                }
                _ => return Err(BackendError::UnsupportedOperation(self.name().into())),
            }
        }
        Ok(())
    }

    fn synchronize(&self) -> Result<(), BackendError> {
        self.stream.synchronize()?;
        Ok(())
    }
}

七、端到端示例:构建一个 Mini GPT 推理引擎

来看一个完整的示例:如何用我们上面构建的组件运行一个 GPT-2 124M 模型的推理。

pub struct TransformerBlock {
    attn_qkv: Tensor,      // [3 * hidden_size, hidden_size]
    attn_proj: Tensor,     // [hidden_size, hidden_size]
    ln1_weight: Tensor,    // [hidden_size]
    ln1_bias: Tensor,      // [hidden_size]
    ffn_inter: Tensor,     // [4 * hidden_size, hidden_size]
    ffn_output: Tensor,    // [hidden_size, 4 * hidden_size]
    ln2_weight: Tensor,    // [hidden_size]
    ln2_bias: Tensor,      // [hidden_size]
}

impl TransformerBlock {
    pub fn build_graph(&self, input: NodeId, graph: &mut ComputeGraph) -> NodeId {
        // self-attention
        let qkv = graph.add_node(Op::MatMul);
        let qkv_bias = graph.add_node(Op::Add);
        let qkv_split = graph.add_node(Op::Split { n: 3, axis: 2 });
        let attn_weights = graph.add_node(Op::MatMul);  // Q @ K^T
        let attn_scale = graph.add_node(Op::Mul);       // / sqrt(d_k)
        let attn_softmax = graph.add_node(Op::Softmax { axis: 3 });
        let attn_out = graph.add_node(Op::MatMul);      // attn @ V
        let attn_proj = graph.add_node(Op::MatMul);     // proj
        
        // residual
        let res1 = graph.add_node(Op::Add);
        
        // FFN with GeLU
        let ln2 = graph.add_node(Op::LayerNorm { eps: 1e-5 });
        let ffn1 = graph.add_node(Op::MatMul);
        let ffn_act = graph.add_node(Op::Silu);  // LLaMA 风格用 SiLU
        let ffn2 = graph.add_node(Op::MatMul);
        
        // residual 2
        let res2 = graph.add_node(Op::Add);

        graph.add_edge(input, qkv);
        // ... 连接所有边
        
        res2
    }

    pub fn run_inference(&self, input: &[f32], seq_len: usize) -> Vec<f32> {
        // 1. 构建计算图
        let mut graph = ComputeGraph::new();
        let input_node = graph.add_node(Op::Input("hidden_states".into()));
        let output_node = self.build_graph(input_node, &mut graph);
        
        // 2. 拓扑排序
        let order = graph.topological_sort().unwrap();
        
        // 3. 算子融合
        let fused_ops = fuse_adjacent_ops(&graph, &order);
        
        // 4. 内存规划
        let planner = MemoryPlanner { graph: &graph, order, dtype_sizes: self.compute_sizes() };
        let plan = planner.plan();
        
        // 5. 分配共享内存 arena
        let mut arena = MemoryArena::new(plan.total_memory_required());
        
        // 6. 转换为 IR
        let ir_ops = lower_to_ir(&fused_ops);
        
        // 7. 执行
        let backend = select_backend();
        backend.execute(&ir_ops, &mut arena).unwrap();
        
        // 8. 取回结果
        arena.read_output(output_node)
    }
    
    fn compute_sizes(&self) -> HashMap<NodeId, usize> {
        // 预计算每个节点的输出大小
        todo!()
    }
}

八、性能调优实战总结

在实际部署 Rust 推理引擎时,以下调优手段可以带来显著性能提升:

  1. 算子融合策略:Transformer 中优先融合 MatMul + BiasAdd + Activation 三元组,减少 50%+ 的内存带宽开销
  2. 量化推理:使用 4-bit GPTQ/AWQ 量化将模型内存降低 4x,推理吞吐提升 2-3x
  3. KV Cache 预分配:LLM 推理中 Key-Value Cache 按最大序列长度预分配,避免推理过程中内存分配
  4. Continuous Batching:动态批次组合提高 GPU 利用率(类似 vLLM 的策略)
  5. 分页注意力:将 KV Cache 拆分为固定大小的块,支持超长序列推理
/// KV Cache 的分页管理(类似 vLLM 的 PagedAttention)
pub struct PagedKvCache {
    block_size: usize,           // 每个 block 存放的 token 数
    num_blocks: usize,           // 总 block 数
    free_blocks: Vec<usize>,     // 空闲 block 索引
    /// seq_id -> 已分配的 block 列表
    allocations: HashMap<usize, Vec<usize>>,
    /// 物理显存中 block 数据
    key_blocks: HashMap<usize, Tensor>,
    value_blocks: HashMap<usize, Tensor>,
}

impl PagedKvCache {
    pub fn allocate(&mut self, seq_id: usize, num_tokens: usize) -> Result<(), CacheError> {
        let blocks_needed = (num_tokens + self.block_size - 1) / self.block_size;
        if blocks_needed > self.free_blocks.len() {
            return Err(CacheError::OutOfMemory);
        }
        
        let allocated: Vec<usize> = self.free_blocks
            .drain(..blocks_needed)
            .collect();
        
        for block_id in &allocated {
            // 初始化 block
            self.key_blocks.insert(*block_id, Tensor::zeros(/* shape */));
            self.value_blocks.insert(*block_id, Tensor::zeros(/* shape */));
        }
        
        self.allocations.insert(seq_id, allocated);
        Ok(())
    }
}

九、生态现状与前景展望

截至 2026 年中,Rust AI 推理生态已经形成了几个明确的阵营:

  • Candle(Hugging Face):优势在于 Hugging Face 生态集成,transformers 模型直接加载;劣势是算子支持不如 PyTorch 丰富
  • Burn(独立社区):设计最野心勃勃,支持训练+推理,自带 JIT 编译后端(基于 CubeCL),在 Apple Silicon 上通过 Metal 后端表现出色
  • dfdx:早期 Rust DL 框架,类型状态机设计精巧但 API 过度面向训练
  • ort(ONNX Runtime 绑定):生产环境最稳妥的选择,通过 C FFI 调用微软 ONNX Runtime

从工程实践角度看,Rust 推理引擎在以下场景有明显优势:

  • 边缘设备部署:Rust 的交叉编译能力 + 静态链接,生成单个无依赖可执行文件,远优于 Python 运行时
  • 高并发服务端:配合 tokio 异步运行时,可以在一个进程中并行处理数百个推理请求
  • 安全敏感的推理场景:可信执行环境(TEE)中,Rust 的内存安全保证避免 side-channel 攻击面

总结

构建 Rust 深度学习推理引擎不是要取代 PyTorch,而是在内存安全、部署便利性、并发性能等方面创造独特的工程价值。核心挑战在于如何平衡:

  • 抽象 vs 性能:计算图的 JIT 编译需要足够的静态信息来生成优化代码
  • 通用性 vs 特化性:Flash Attention 等高度优化的算子不能只通过通用抽象实现
  • 生态系统成熟度: weights 格式转换(safetensors)、模型定义文件(config.json)等标准化的支撑

Rust 已经在这个赛道上证明了它的价值。对于希望在生产环境中部署高性能、低开销推理服务的工程师来说,这个方向值得持续投入。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部