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 正在证明它可以是那个答案。

发表评论 取消回复