楼层: 首页/ 软件技术/ Rust + AI 全栈/ 机器学习框架:candle / ort / tch-rs
02

机器学习框架:candle / ort / tch-rs

Inference Frameworks · candle, ort, tch-rs

candle 就是 Rust 世界的 PyTorch——Hugging Face 官方出品,纯 Rust 实现,没有 Python 依赖。ort 是 ONNX Runtime 的 Rust 绑定,跑跨框架导出的模型最方便。tch-rs 是 PyTorch 的 Rust 绑定,能直接加载 Python 训练好的 .pt 文件。

candle:纯 Rust 的大模型推理框架

论candle 是什么

Hugging Face 2023 年开源的 Rust 深度学习框架,目标很明确:把模型推理从 Python 解放出来。它不追求训练(训练还是 PyTorch 的活),专注推理——加载模型、前向计算、采样生成。原生支持 safetensors 和 GGUF 量化格式,能跑 LLaMA、Qwen、Mistral、Phi 这些主流架构。

为什么爽:纯 Rust,没有 libtorch 那 2GB 的依赖;交叉编译到 Android/iOS 毫无压力;内存安全,不会跑着跑着 segfault。

Cargo.toml(candle 三件套必须同版本)

[dependencies] candle-core = { version = "0.11", features = ["cuda"] } // CPU 部署去掉 cuda candle-nn = "0.11" candle-transformers = "0.11" tokenizers = "0.23" hf-hub = "1.0" anyhow = "1"

用 candle 跑一个 Qwen 量化模型(核心代码)

use anyhow::Result; use candle_core::{Device, Tensor}; use candle_transformers::models::quantized_qwen2::Model as Qwen2; use candle_transformers::generation::{LogitsProcessor, Sampling}; use hf_hub::{api::sync::Api, Repo, RepoType}; use tokenizers::Tokenizer; fn load_model() -> Result<(Qwen2, Tokenizer, Device)> { let device = Device::new_cuda(0).unwrap_or(Device::Cpu); let api = Api::new()?; let repo = Repo::with_revision( "Qwen/Qwen2-0.5B-Instruct-GGUF".into(), RepoType::Model, "main".into(), ); let model_file = api.model(repo.clone()).get("model-q4_k_m.gguf")?; let tok_file = api.model(repo).get("tokenizer.json")?; let tokenizer = Tokenizer::from_file(tok_file)?; let mut file = std::fs::File::open(model_file)?; let model = Qwen2::from_gguf(&mut file, &device)?; Ok((model, tokenizer, device)) } // 逐 token 生成(流式核心) fn generate(model: &mut Qwen2, tok: &Tokenizer, dev: &Device, prompt: &str) -> Result<String> { let tokens = tok.encode(prompt, true)?.get_ids().to_vec(); let mut input = Tensor::new(&tokens[..], dev)?.unsqueeze(0)?; let mut lp = LogitsProcessor::new(42, Some(0.7), Some(0.95)); let mut out = vec![]; for _ in 0..512 { let logits = model.forward(&input, 0)?.squeeze(0)?; let next = lp.sample(&logits)?; if next == tok.token_to_id("<|im_end|>").flatten().unwrap_or(151645) { break; } out.push(next); let n = Tensor::new(&[next], dev)?.unsqueeze(0)?; input = Tensor::cat(&[&input, &n], 1)?; } Ok(tok.decode(&out, true)?) }

GGUF 量化:用 1/4 内存跑 7B 模型

量化等级精度损失7B 模型内存适用场景
FP16无~14 GB服务器高精度推理
Q8_0极小~7 GB服务器 / 工作站
Q5_K_M较小~4.5 GB消费级显卡
Q4_K_M可接受~4 GB端侧 / 移动端(推荐)
Q3_K_M明显~3 GB极致压缩,不推荐生产

人话:量化就是把模型权重从 16 位浮点数压到 4 位整数,内存占 1/4,速度还更快(因为瓶颈是内存带宽)。Q4_K_M 是甜点位——精度损失普通人感觉不出来。

ort:跑 ONNX 模型(多模态首选)

ort 是 ONNX Runtime 的 Rust 绑定。ONNX 是跨框架的通用模型格式——PyTorch、TensorFlow、PaddlePaddle 训练的模型都能导出成 ONNX,一次导出处处运行。图像识别、语音、目标检测、文本 Embedding,ort 全包。

[dependencies] ort = { version = "2.0.0-rc.13", features = ["download-binaries"] } ndarray = "0.16"
use ort::{Session, inputs, GraphOptimizationLevel}; let session = Session::builder()? .with_optimization_level(GraphOptimizationLevel::Level3)? .commit_from_file("models/resnet50.onnx")?; // 推理:输入一个 ndarray 张量 let input = ndarray::Array4::<f32>::zeros((1, 3, 224, 224)); let outputs = session.run(ort::inputs!["input" => input.view()])?; let logits = outputs[0].try_extract_tensor::<f32>()?;

tch-rs:PyTorch 原生绑定

如果你手上有 Python 训练好的 .pt/.bin 模型,又没法导出 ONNX(比如自定义算子),用 tch-rs。它直接绑定 libtorch,所有 PyTorch 操作都能用。缺点:libtorch 体积约 2GB,部署包很大。除非必须,否则优先 candle 或 ort。

[dependencies] tch = "0.19" # 环境变量:指向 libtorch 解压目录 # export LIBTORCH=/path/to/libtorch

传统机器学习:linfa 与 smartcore

除了深度学习,Rust 还有传统 ML 库——linfa 对标 scikit-learn,smartcore 另一个选择。线性回归、逻辑回归、SVM、决策树、聚类、降维都有。不需要 GPU 的小模型场景,用它们比 candle 轻得多。

use linfa::prelude::*; use linfa_regression::LinearRegression; // 数据集:特征 X,目标 y let dataset = Dataset::new(x, y); // 训练线性回归 let model = LinearRegression::::default().fit(&dataset)?; // 预测 let pred = model.predict(&test_x);

框架选择速查表

场景选什么
LLM 推理(中文/英文)candle + GGUF 量化
图像识别 / 语音 / Embeddingort(ONNX 模型)
必须跑 Python 训的 .pttch-rs
线性回归 / 决策树 / 聚类linfa
端侧 / 移动端candle(纯 Rust 交叉编译)
记
本章小结

① candle 跑大语言模型(LLaMA/Qwen/Mistral),纯 Rust、GGUF 量化、可交叉编译。

② ort 跑 ONNX 模型,多模态(图像/语音/Embedding)首选。

③ tch-rs 只在必须跑 PyTorch 原生模型时用。

④ 量化选 Q4_K_M,内存降 4 倍,精度可接受。