Rust Candle 深度学习框架深度工程实践:从张量原语到 LLM 推理引擎

当 Python 生态成为 AI 研究的"舒适区",Rust 正在用零成本抽象和内存安全重新定义深度学习基础设施的边界。本文深入剖析 Hugging Face Candle 框架的核心架构,并通过完整代码示例展示如何从零构建一个生产级 LLM 推理引擎。

一、为什么深度学习需要 Rust

长期以来,Python 凭借 PyTorch 和 TensorFlow 占据 AI 生态主导地位。然而在生产部署环节,Python 暴露出三个根本性瓶颈:GIL 限制并发、解释器开销巨大、内存占用不可控。一个典型的 LLM 推理服务在 Python 运行时下,内存峰值往往达到模型权重的 3-5 倍,其中大量开销来自 Python 对象系统和 GC 不可预测性。

Candle 的出现并非要挑战 Python 在研究领域的地位,而是精准切入"研究与生产之间的死亡之谷"。Hugging Face 官方团队负责人在设计 Candle 时就明确了三个目标:取代 PyTorch 成为 Inference 后端;用 Rust 重写 HF 的 Transformers 推理路径;为嵌入式和边缘 AI 提供 C++ 之外的另一种选择。

// Candle 的哲学:一切皆为 Tensor
use candle_core::{Device, Tensor, DType;

fn main() -> anyhow::Result<()> {
    // 一行代码在 CUDA 上创建张量,无需 Python 上下文
    let dev = Device::new_cuda(0)?;
    let a = Tensor::arange(0f32, 12f32, &dev)?.reshape((3, 4))?;
    let b = Tensor::ones((4, 2), DType::F32, &dev)?;
    let c = a.matmul(&b)?;  // (3, 2) 矩阵乘法
    println!("{:?}", c.to_vec2::<f32>()?);
    Ok(())
}

二、Candle 核心架构解析

2.1 Tensor 系统:类型安全的张量代数

Candle 的 Tensor 设计遵循 Rust 的所有权和借用规则,确保张量操作的内存安全。每个 Tensor 实例携带形状、数据类型和设备信息作为类型参数,编译器就能在编译阶段捕获大量错误。

use candle_core::{Tensor, DType, Shape};

// Tensor 结构的核心抽象
// pub struct Tensor_ {
//     storage: Storage,      // 底层存储(CPU/CUDA/Metal)
//     layout: Layout,        // 步长和偏移(支持 view 而不拷贝)
//     dtype: DType,
//     shape: Shape,
// }

fn tensor_operations() -> anyhow::Result<()> {
    let dev = Device::Cpu;
    
    // 连续内存布局 vs 非连续视图
    let base = Tensor::from_slice(&[1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], (2, 3), &dev)?;
    let transposed = base.t()?;  // 不产生数据拷贝,仅修改 stride
    let contiguous = transposed.contiguous()?;  // 真正拷贝为连续内存
    
    // 广播语义与 NumPy 一致
    let a = Tensor::ones((3, 1), DType::F32, &dev)?;
    let b = Tensor::ones((1, 4), DType::F32, &dev)?;
    let c = (&a + &b)?;  // 广播为 (3, 4)
    Ok(())
}

2.2 设备抽象层:从 CPU 到异构统一内存

Candle 通过 Device 枚举实现跨设备统一 API,支持 CPU、CUDA、Metal 和 Sycl 四种后端。设备切换只需修改 Device 实例,代码无需任何改动。

use candle_core::Device;

fn multi_device_workflow(prefer_gpu: bool) -> anyhow::Result<()> {
    // 自动选择最佳可用设备
    let dev = if prefer_gpu {
        if let Ok(cuda) = Device::new_cuda(0) {
            cuda
        } else if let Ok(metal) = Device::new_metal(0) {
            metal
        } else {
            Device::Cpu
        }
    } else {
        Device::Cpu
    };
    
    // CUDA Stream 管理:非阻塞执行
    let s1 = Device::new_cuda(0)?;
    let s2 = Device::new_cuda(0)?;
    
    // 在独立 stream 上并行计算
    let a = Tensor::randn(0f32, 1f32, (1024, 1024), &s1)?;
    let b = Tensor::randn(0f32, 1f32, (1024, 1024), &s2)?;
    
    // 同步点确保计算完成
    s1.synchronize()?;
    s2.synchronize()?;
    Ok(())
}

2.3 Backend Trait:可扩展的计算后端

Candle 定义了 BackendStorage trait 作为所有计算操作的统一接口。新增硬件只需实现该 trait,无需修改上层逻辑。

// Backend trait 定义了所有设备必须实现的计算原语
// pub trait BackendStorage {
//     fn dtype(&self) -> DType;
//     fn device(&self) -> Device;
//     fn to_cpu_storage(&self) -> Result<CpuStorage>;
//     fn affine(&self, mul: f64, add: f64) -> Result<Self>;
//     fn powf(&self, exp: f64) -> Result<Self>;
//     fn elu(&self, alpha: f64) -> Result<Self>;
//     // ... 约 200 个计算原语
// }

三、从零构建 Transformer:工程级实现

3.1 注意力机制的数学与工程实现

Scaled Dot-Product Attention 是 Transformer 的心脏。Candle 实现中特别注意了内存效率和数值稳定性。

use candle_core::{Tensor, D};
use candle_nn as nn;

/// 多头注意力机制的工程级实现
struct MultiHeadAttention {
    q_proj: nn::Linear,
    k_proj: nn::Linear,
    v_proj: nn::Linear,
    o_proj: nn::Linear,
    num_heads: usize,
    head_dim: usize,
    span: tracing::Span,
}

impl MultiHeadAttention {
    fn new(vocab_size: usize, embed_dim: usize, num_heads: usize, vb: nn::VarBuilder) -> anyhow::Result<Self> {
        let head_dim = embed_dim / num_heads;
        Ok(Self {
            q_proj: nn::linear(embed_dim, embed_dim, vb.pp("q_proj"))?,
            k_proj: nn::linear(embed_dim, embed_dim, vb.pp("k_proj"))?,
            v_proj: nn::linear(embed_dim, embed_dim, vb.pp("v_proj"))?,
            o_proj: nn::linear(embed_dim, embed_dim, vb.pp("o_proj"))?,
            num_heads,
            head_dim,
            span: tracing::span!(tracing::Level::TRACE, "multi-head-attn"),
        })
    }
    
    /// 计算注意力,支持 KV-Cache 增量推理
    fn forward(&self, x: &Tensor, kv_cache: Option<&Tensor>, causal_mask: bool) -> anyhow::Result<(Tensor, Tensor)> {
        let _guard = self.span.enter();
        
        let (b_sz, seq_len, hidden) = x.dims3()?;
        
        // 投影 Q/K/V 并 reshape 为多头格式
        // (batch, seq, hidden) -> (batch, num_heads, seq, head_dim)
        let q = self.q_proj.forward(x)?
            .reshape((b_sz, seq_len, self.num_heads, self.head_dim))?
            .transpose(1, 2)?;
        let k = self.k_proj.forward(x)?
            .reshape((b_sz, seq_len, self.num_heads, self.head_dim))?
            .transpose(1, 2)?;
        let v = self.v_proj.forward(x)?
            .reshape((b_sz, seq_len, self.num_heads, self.head_dim))?
            .transpose(1, 2)?;
        
        // KV-Cache:增量推理时拼接历史 KV
        let (k, v) = match kv_cache {
            Some(cache) => {
                let past_k = cache.get(0)?;
                let past_v = cache.get(1)?;
                let k = Tensor::cat(&[&past_k, &k], D::Minus2)?;
                let v = Tensor::cat(&[&past_v, &v], D::Minus2)?;
                (k, v)
            }
            None => (k, v),
        };
        
        let updated_cache = Tensor::stack(&[&k, &v], 0)?;
        
        // 注意力分数:Q @ K^T / sqrt(d_k)
        let scale = (self.head_dim as f64).sqrt();
        let attn_weights = q.matmul(&k.t()?)?;
        let attn_weights = (attn_weights / scale)?;
        
        // Causal Mask:防止未来信息泄露
        let attn_weights = if causal_mask {
            let mask = create_causal_mask(seq_len, &attn_weights.device())?;
            attn_weights.broadcast_add(&mask)?
        } else {
            attn_weights
        };
        
        let attn_probs = candle_nn::ops::softmax(&attn_weights, D::Minus1)?;
        let xs = attn_probs.matmul(&v)?;
        
        // 合并多头输出
        let xs = xs.transpose(1, 2)?.contiguous()?
            .reshape((b_sz, seq_len, hidden))?;
        let xs = self.o_proj.forward(&xs)?;
        
        Ok((xs, updated_cache))
    }
}

fn create_causal_mask(seq_len: usize, dev: &candle_core::Device) -> anyhow::Result<Tensor> {
    let mut mask = vec![0f32; seq_len * seq_len];
    for i in 0..seq_len {
        for j in (i + 1)..seq_len {
            mask[i * seq_len + j] = f32::NEG_INFINITY;
        }
    }
    Tensor::from_vec(mask, (seq_len, seq_len), dev)
}

3.2 前馈网络与层归一化

/// SwiGLU 前馈网络:现代 LLM 的标准选择
struct FeedForward {
    gate_proj: nn::Linear,
    up_proj: nn::Linear,
    down_proj: nn::Linear,
}

impl FeedForward {
    fn new(dim: usize, hidden_dim: usize, vb: nn::VarBuilder) -> anyhow::Result<Self> {
        Ok(Self {
            gate_proj: nn::linear(dim, hidden_dim, vb.pp("gate_proj"))?,
            up_proj: nn::linear(dim, hidden_dim, vb.pp("up_proj"))?,
            down_proj: nn::linear(hidden_dim, dim, vb.pp("down_proj"))?,
        })
    }
    
    fn forward(&self, x: &Tensor) -> anyhow::Result<Tensor> {
        let gate = self.gate_proj.forward(x)?;
        let gate = candle_nn::ops::silu(&gate)?;
        let up = self.up_proj.forward(x)?;
        let hidden = (gate * up)?;
        self.down_proj.forward(&hidden)
    }
}

/// RMSNorm:比 LayerNorm 更高效,不计算均值
struct rms_norm {
    weight: Tensor,
    epsilon: f32,
}

impl RmsNorm {
    fn new(dim: usize, epsilon: f32, vb: nn::VarBuilder) -> anyhow::Result<Self> {
        Ok(Self {
            weight: vb.get(dim, "weight")?,
            epsilon,
        })
    }
    
    fn forward(&self, x: &Tensor) -> anyhow::Result<Tensor> {
        let norm = x.sqr()?.mean_keepdim(D::Minus1)?
            .add(self.epsilon)?.rsqrt()?;
        let x = (x * norm)?;
        (x * &self.weight)
    }
}

3.3 完整 Transformer Block 与模型组装

struct TransformerBlock {
    attn: MultiHeadAttention,
    ffn: FeedForward,
    attn_norm: RmsNorm,
    ffn_norm: RmsNorm,
}

impl TransformerBlock {
    fn forward(&self, x: &Tensor, kv_cache: Option<&Tensor>) -> anyhow::Result<(Tensor, Tensor)> {
        // Pre-Norm 架构(比 Post-Norm 更稳定)
        let residual = x.clone();
        let normalized = self.attn_norm.forward(x)?;
        let (attn_out, new_cache) = self.attn.forward(&normalized, kv_cache, true)?;
        let x = (residual + attn_out)?;
        
        let residual = x.clone();
        let normalized = self.ffn_norm.forward(&x)?;
        let ffn_out = self.ffn.forward(&normalized)?;
        let x = (residual + ffn_out)?;
        
        Ok((x, new_cache))
    }
}

/// 完整的 Llama-style 模型
struct LlamaModel {
    embedding: nn::Embedding,
    blocks: Vec<TransformerBlock>,
    norm: RmsNorm,
    lm_head: nn::Linear,
    num_layers: usize,
}

impl LlamaModel {
    fn forward(&self, input: &Tensor, caches: &[Option<Tensor>]) -> anyhow::Result<(Tensor, Vec<Option<Tensor>>)> {
        let mut x = self.embedding.forward(input)?;
        let mut new_caches = Vec::with_capacity(self.num_layers);
        
        for (i, block) in self.blocks.iter().enumerate() {
            let (out, cache) = block.forward(&x, caches[i].as_ref())?;
            x = out;
            new_caches.push(Some(cache));
        }
        
        let x = self.norm.forward(&x)?;
        // 共享 embedding 权重(LLM 中的常见技巧)
        let logits = self.lm_head.forward(&x)?;
        
        Ok((logits, new_caches))
    }
}

四、推理引擎的工程实现

4.1 高效 KV-Cache 管理

KV-Cache 是 LLM 推理性能的关键。增量生成的每一步,我们都希望复用历史 KV 值而非重新计算。

/// 分页 KV-Cache:灵感来自 vLLM 的 PagedAttention
struct PagedKvCache {
    block_size: usize,          // 默认 16 tokens
    num_blocks: usize,          // 总物理块数
    free_list: Vec<usize>,      // 空闲块索引
    block_table: Vec<Vec<usize>>, // sequence_id -> 物理块列表
    kv_tensor: Tensor,          // (num_blocks, 2, block_size, num_heads, head_dim)
}

impl PagedKvCache {
    fn new(num_blocks: usize, block_size: usize, num_heads: usize, head_dim: usize, dev: &Device) -> anyhow::Result<Self> {
        Ok(Self {
            block_size,
            num_blocks,
            free_list: (0..num_blocks).rev().collect(),
            block_table: Vec::new(),
            kv_tensor: Tensor::zeros((num_blocks, 2, block_size, num_heads, head_dim), DType::F32, dev)?,
        })
    }
    
    fn allocate(&mut self) -> anyhow::Result<usize> {
        self.free_list.pop().ok_or(anyhow::anyhow!("KV cache exhausted"))
    }
    
    fn deallocate(&mut self, seq_id: usize) {
        if let Some(blocks) = self.block_table.get(seq_id) {
            for &block in blocks {
                self.free_list.push(block);
            }
        }
        // 标记为已释放
    }
    
    /// 增量写入:将新 KV 值写入分配的块
    fn write(&mut self, seq_id: usize, token_pos: usize, k: &Tensor, v: &Tensor) -> anyhow::Result<()> {
        let block_idx = token_pos / self.block_size;
        let offset = token_pos % self.block_size;
        
        let blocks = &self.block_table[seq_id];
        let physical_block = blocks[block_idx];
        
        // 高效索引赋值:只写入当前 token 对应的槽位
        self.kv_tensor.get(physical_block)?.get(0)?.get(offset)?.copy_from(k)?;
        self.kv_tensor.get(physical_block)?.get(1)?.get(offset)?.copy_from(v)?;
        
        Ok(())
    }
}

4.2 采样策略:从贪婪到核采样

use rand::Rng;

enum SamplingStrategy {
    Greedy,
    TopK { k: usize, temperature: f32 },
    TopP { p: f32, temperature: f32 },
    Typical { mass: f32, temperature: f32 },
}

fn sample(logits: &Tensor, strategy: &SamplingStrategy) -> anyhow::Result<u32> {
    let logits = logits.to_vec1::<f32>()?;
    
    match strategy {
        SamplingStrategy::Greedy => {
            Ok(logits.iter().enumerate().max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
                .map(|(i, _)| i as u32).unwrap())
        }
        SamplingStrategy::TopK { k, temperature } => {
            // 1. 温度缩放
            let scaled: Vec<f32> = logits.iter().map(|x| x / temperature).collect();
            
            // 2. Top-K 过滤
            let mut indices: Vec<usize> = (0..scaled.len()).collect();
            indices.sort_by(|&a, &b| scaled[b].partial_cmp(&scaled[a]).unwrap());
            let top_k_indices = &indices[..(*k).min(indices.len())];
            
            // 3. Softmax 采样
            let max_logit = top_k_indices.iter().map(|&i| scaled[i]).fold(f32::NEG_INFINITY, f32::max);
            let exp_sum: f32 = top_k_indices.iter().map(|&i| (scaled[i] - max_logit).exp()).sum();
            let probs: Vec<f32> = top_k_indices.iter().map(|&i| (scaled[i] - max_logit).exp() / exp_sum).collect();
            
            let mut rng = rand::thread_rng();
            let mut cumsum = 0.0;
            let sample: f32 = rng.gen();
            for (i, &prob) in probs.iter().enumerate() {
                cumsum += prob;
                if cumsum >= sample {
                    return Ok(top_k_indices[i] as u32);
                }
            }
            Ok(top_k_indices[top_k_indices.len() - 1] as u32)
        }
        _ => Ok(0), // 简化处理
    }
}

4.3 分词器与流式输出

use tokenizers::Tokenizer;

struct LlmEngine {
    model: LlamaModel,
    tokenizer: Tokenizer,
    device: Device,
    kv_cache: PagedKvCache,
}

impl LlmEngine {
    fn generate(&mut self, prompt: &str, max_tokens: usize, strategy: SamplingStrategy) -> anyhow::Result<String> {
        let encoding = self.tokenizer.encode(prompt, true)?;
        let mut token_ids: Vec<u32> = encoding.get_ids().to_vec();
        
        // Prefill 阶段:处理完整 prompt
        let input = Tensor::new(&token_ids, &self.device)?;
        let (mut logits, mut caches) = self.model.forward(&input, &vec![None; self.model.num_layers])?;
        
        // Decode 阶段:逐 token 生成
        let mut generated = Vec::new();
        for _ in 0..max_tokens {
            let last_logits = logits.narrow(D::Minus2, logits.dim(D::Minus2)? - 1, 1)?.squeeze(D::Minus2)?;
            let next_id = sample(&last_logits, &strategy)?;
            
            if next_id == self.tokenizer.token_to_id("</s>").unwrap_or(0) {
                break;
            }
            
            generated.push(next_id);
            
            // 流式输出当前 token
            let piece = self.tokenizer.decode(&[next_id], true)?;
            print!("{}", piece);
            std::io::Write::flush(&mut std::io::stdout())?;
            
            // 增量前向传播
            let next_input = Tensor::new(&[next_id], &self.device)?;
            let (new_logits, new_caches) = self.model.forward(&next_input, &caches)?;
            logits = new_logits;
            caches = new_caches;
        }
        
        println!(); // 换行
        Ok(self.tokenizer.decode(&generated, true)?)
    }
}

五、生产部署的工程要点

5.1 模型量化:4bit GPTQ/AWQ 支持

Candle 内置了对多种量化格式的支持,大幅降低显存占用。

use candle_transformers::quantized_var_builder::VarBuilder;

fn load_quantized_model(path: &str) -> anyhow::Result<LlamaModel> {
    // 加载 GGUF 格式的量化模型
    let vb = VarBuilder::from_gguf(path, &Device::Cpu)?;
    
    // 模型自动识别量化类型(Q4_0, Q5_K_M, Q8_0 等)
    // 推理时自动反量化
    let config = candle_transformers::models::llama::Config::tiny_llama();
    let model = LlamaModel::new(&config, vb)?;
    
    Ok(model)
}

5.2 批调度与连续批处理

/// 连续批处理(Continuous Batching):提升吞吐量的关键
struct BatchScheduler {
    pending_requests: VecDeque<GenerationRequest>,
    active_requests: Vec<ActiveRequest>,
    max_batch_size: usize,
    kv_cache: PagedKvCache,
}

impl BatchScheduler {
    fn step(&mut self) -> Vec<GenerationResult> {
        // 1. 合并 prefill + decode 请求
        let batch = self.prepare_batch();
        
        // 2. 拼接所有输入张量
        let input_ids = self.pad_and_concat(&batch.inputs);
        
        // 3. 单次 forward 处理整个 batch
        let (logits, new_kv) = self.model.forward(&input_ids, &batch.kv_caches).unwrap();
        
        // 4. 分发结果
        let mut results = Vec::new();
        for (i, request) in batch.requests.iter().enumerate() {
            let logits_i = logits.get(i).unwrap();
            let next_token = sample(logits_i, &request.strategy).unwrap();
            
            if next_token == EOS_TOKEN || request.tokens.len() >= request.max_tokens {
                results.push(GenerationResult::Completed(request.id));
            } else {
                results.push(GenerationResult::Token(next_token));
            }
        }
        
        results
    }
}

5.3 性能监控与 Profiling

use std::time::Instant;

struct InferenceMetrics {
    prompt_tokens: u64,
    generated_tokens: u64,
    prefill_time_ms: f64,
    decode_time_per_token_ms: Vec<f64>,
    peak_memory_mb: f64,
}

impl InferenceMetrics {
    fn report(&self) {
        let total_time = self.prefill_time_ms + self.decode_time_per_token_ms.iter().sum::<f64>();
        let tokens_per_sec = self.generated_tokens as f64 / (total_time / 1000.0);
        
        println!("=== Performance Report ===");
        println!("Prompt tokens: {}", self.prompt_tokens);
        println!("Generated tokens: {}", self.generated_tokens);
        println!("Prefill: {:.2} ms", self.prefill_time_ms);
        println!("Decode: {:.2} ms/tok", self.decode_time_per_token_ms.iter().sum::<f64>() / self.decode_time_per_token_ms.len() as f64);
        println!("Throughput: {:.2} tok/s", tokens_per_sec);
        println!("Peak memory: {:.2} MB", self.peak_memory_mb);
    }
}

六、与 Python 生态的对比

| 维度 | PyTorch (Python) | Candle (Rust) |

| 启动时间 | 1-3s | 10-50ms |

| 内存开销 (7B 模型) | 14GB+ | 6-8GB |

| 推理延迟 (p50) | 2-5ms | 1-3ms |

| 并发处理 | GIL 受限 | 原生异步 |

| 部署依赖 | Python + pip 生态 | 单一二进制 |

| 编译时间 | 无 | 首次 5-15min |

| 热更新 | 易 | 需预编译 |

Candle 并非要完全替代 Python,而是在部署层面提供优势。研究阶段用 PyTorch 快速实验,生产部署切换到 Candle 降低 TCO——这是目前最实际的混合策略。

七、总结与展望

Candle 代表了一种新的深度学习基础设施范式:用 Rust 的类型系统和所有权模型,将"运行时错误"转化为"编译期保证"。随着模型推理成本在生产环境中占比越来越高,这种范式正在获得越来越多的工程认可。

对于希望构建生产级 AI 服务的团队,我推荐的落地路径是:首先在推理侧试点 Candle(成本降低 40%+),然后逐步将数据预处理和特征工程迁移到 Rust,最终实现"Python 研究 + Rust 部署"的双栈架构。

关键洞察:深度学习框架的下一个十年,属于能在内存安全、计算性能和工程可维护性之间找到平衡的语言。Rust 正在证明它可以是那个答案。
点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部