- 人工智能
- 大模型
- 机器学习
- 深度学习
- 本地部署
- 模型推理服务
【免费下载链接】candle
Minimalist ML framework for Rust
本篇指南围绕 Candle 推理链路中的第一步——模型权重获取展开,讲解如何借助hf-hub从 Hugging Face Hub 下载预训练权重(以bert-base-uncased为例),将safetensors文件解析为candle-core的Tensor,并进一步打通内存映射加载(mmap)与多 GPU 张量并行(Tensor Parallel)场景下的按分片加载方案。读完本文,你将掌握一套可复用的"下载 → 加载 → 接入模型 → 分片"完整流程,并了解其底层 API 与测试验证。
为什么需要 hub:Candle 中的权重获取方式
Candle 自身不托管模型文件,绝大多数预训练模型的权重以safetensors(或老旧的pytorch_model.bin)格式存放在 Hugging Face Hub 上。因此,任何推理或微调流程的第一步都是从 Hub 拉取模型文件。Candle 官方在 candle-book/src/inference/hub.md 中给出了标准做法:使用官方维护的hf-hubRust crate,它可以处理模型库与数据集库的鉴权、缓存、断点续传与 revision 固定。
仓库中candle-examples的大量示例都依赖这一套流程,例如 bert 示例 会先下载config.json、tokenizer.json与model.safetensors,再构建VarBuilder加载模型;candle-examples甚至把 hub 的常用封装收敛到了 hub.rs,统一了缓存目录、revision 与下载进度条的处理。
安装依赖并下载第一个模型文件
在Cargo.toml中加入hf-hub:
cargo add hf-hub随后用下面的代码下载bert-base-uncased仓库中的model.safetensors:
use hf_hub::{split_id, HFClientSync}; use candle_core::Device; let api = HFClientSync::new().unwrap(); let (owner, name) = split_id("bert-base-uncased"); let repo = api.model(owner, name); let weights = repo.download_file().filename("model.safetensors").send().unwrap();这里有几个关键点:
split_id("bert-base-uncased")把owner/name形式的仓库 ID 拆成("bert-base-uncased", "")——对于没有组织前缀的仓库,owner 与 name 相同;对sentence-transformers/all-MiniLM-L6-v2这类 ID 则会正确拆为 owner 与 name 两部分。api.model(owner, name)返回一个针对"模型仓库"的句柄(HFRepositorySync<RepoTypeModel>),与之对应还有api.dataset(...)用于数据集仓库。download_file().filename("model.safetensors")是可链式配置的请求构造器,send()返回下载到本地缓存的文件路径(PathBuf),默认缓存目录遵循hf-hub约定,也支持通过环境变量(如HF_TOKEN、HF_HUB_CACHE)配置鉴权与缓存位置。
如果你使用的是异步运行时,hf-hub同样提供HFClient(async 版本),candle-book 的测试代码 candle-book/src/lib.rs#L13-L27(book_hub_1)中即采用:
use candle::Device; use hf_hub::{split_id, HFClient}; let api = HFClient::new().unwrap(); let (owner, name) = split_id("bert-base-uncased"); let repo = api.model(owner, name); let weights_filename = repo.download_file().filename("model.safetensors").send().await.unwrap(); let weights = candle::safetensors::load(weights_filename, &Device::Cpu).unwrap();两者的区别仅在同步/异步:HFClientSync直接阻塞返回PathBuf,HFClient返回Future需要.await,按你的运行时环境二选一即可。
将 safetensors 解析为 Tensor 集合
下载完成后,用candle_core::safetensors::load一次性把文件读入内存并解析:
let weights = candle_core::safetensors::load(weights, &Device::Cpu);load的返回值是HashMap<String, Tensor>,key 为张量名(例如bert.encoder.layer.0.attention.self.query.weight),value 为对应Tensor。其底层实现位于 candle-core/src/safetensors.rs#L408-L419:先用std::fs::read读入完整字节,再调用load_buffer通过SafeTensors::deserialize解析头部,并逐个张量转换为 Candle 的Tensor。转换过程(convert/convert_slice)会按safetensors的 dtype 与形状直接构造存储,因此F32、F16、BF16等常见权重格式都开箱即用。
值得注意:candle-book 的对应测试book_hub_1中有一条断言assert_eq!(weights.len(), 206),即bert-base-uncased的model.safetensors一共包含 206 个张量——你可以通过这个数字快速验证自己的加载流程是否完整。
把权重接入真实模型:以 BERT 的一个 Linear 层为例
拿到HashMap<String, Tensor>之后,就可以按张量名取用参数。文档中给出了最直接的用法——取 BERT 第一层 self-attention 中 query 投影的权重和偏置,构造一个Linear并前向计算:
use candle_core::{Device, Tensor, DType}; use candle_nn::{Linear, Module}; let weights = candle_core::safetensors::load(weights, &Device::Cpu).unwrap(); let weight = weights.get("bert.encoder.layer.0.attention.self.query.weight").unwrap(); let bias = weights.get("bert.encoder.layer.0.attention.self.query.bias").unwrap(); let linear = Linear::new(weight.clone(), Some(bias.clone())); let input_ids = Tensor::zeros((3, 768), DType::F32, &Device::Cpu).unwrap(); let output = linear.forward(&input_ids).unwrap();Linear::new来自 candle-nn,输入形状(3, 768)对应bert-base-uncased的 hidden size 768。前向得到(3, 768)的输出(该层输出维度与输入一致,因为 query 投影保持隐层维度)。
如果要在生产代码里完整加载 BERT 而非手写张量名,推荐直接复用仓库中的完整实现。可以参考 bert 示例,其build_model_and_tokenizer展示了真实工程的做法:
- 用
Api::new()构造客户端(默认从环境变量读取配置); - 用
.with_revision(revision)固定 commit/revision(如refs/pr/21); - 分别
get("config.json")、get("tokenizer.json")、get("model.safetensors"); - 通过
VarBuilder::from_mmaped_safetensors(&[weights_filename], DTYPE, &device)构建变量加载器,再BertModel::load(vb, &config)组装模型。
VarBuilder::from_mmaped_safetensors定义在 candle-nn/src/var_builder.rs#L642-L647,它内部调用candle::safetensors::MmapedSafetensors::multi支持同时映射多个权重文件,这也正好衔接到下一节的内存映射方案。
内存映射加载(mmap):更高效的大模型启动方式
对于动辄数 GB 的权重文件,"整文件读入内存"会产生不必要的分配与拷贝。Candle 支持借助memmap2将文件映射到虚拟内存,只在实际访问时按页调入:
cargo add memmap2use candle::Device; use hf_hub::{split_id, HFClientSync}; use memmap2::Mmap; use std::fs; let api = HFClientSync::new().unwrap(); let (owner, name) = split_id("bert-base-uncased"); let repo = api.model(owner, name); let weights_filename = repo.download_file().filename("model.safetensors").send().unwrap(); let file = fs::File::open(weights_filename).unwrap(); let mmap = unsafe { Mmap::map(&file).unwrap() }; let weights = candle::safetensors::load_buffer(&mmap[..], &Device::Cpu).unwrap();与load的区别在于:这里不再std::fs::read,而是直接把Mmap的字节切片交给load_buffer解析。这一用法在 candle-book 测试 candle-book/src/lib.rs#L31-L49(book_hub_2)中被验证,同样断言可解析出 206 个张量。
需要注意的是文档与源码中的三重提醒:
- unsafe:
Mmap::map是 unsafe 操作,语义上要求映射期间底层文件不被截断或改写。参见memmap2的 Safety 说明。实际上模型文件在推理期间不会被修改,且映射通常保持只读,因此该风险在常规场景下基本不触发,但仍应时刻留意。 - Windows / WSL 兼容性:内存映射在 Windows 与 WSL 环境下可能出现问题(社区已有相关 issue 反馈),跨平台项目需要评估。
- 网络挂载盘:如果权重文件位于网络挂载的磁盘(NFS 等),mmap 会触发更多小粒度读调用,性能反而明显变慢,此时整文件读入更合适。
进一步地,candle-core 在 safetensors.rs 中把 mmap 封装成了多种加载器,按需选用:
MmapedSafetensors:对单文件(new)或多文件(multi)做 mmap + 惰性解析,load(name, dev)按需取张量,这正是VarBuilder::from_mmaped_safetensors与ShardedVarBuilder的底层;SliceSafetensors/BufferedSafetensors:分别面向外部借用的字节切片与自持有的Vec<u8>缓冲;MmapedFile:仅映射文件、按需deserialize获取SafeTensors。
张量并行分片加载:每个 GPU 只读自己那份权重
在多 GPU 做张量并行(Tensor Parallel)以降低延迟时,每个 rank 其实只需要权重的一个切片。此时应直接使用safetensorscrate,按张量维度切出本卡需要的区间,而不是把整个张量都加载进来:
cargo add safetensorsuse candle::{DType, Device, Tensor}; use hf_hub::{split_id, HFClientSync}; use memmap2::Mmap; use safetensors::slice::IndexOp; use safetensors::SafeTensors; use std::fs; let api = HFClientSync::new().unwrap(); let (owner, name) = split_id("bert-base-uncased"); let repo = api.model(owner, name); let weights_filename = repo.download_file().filename("model.safetensors").send().unwrap(); let file = fs::File::open(weights_filename).unwrap(); let mmap = unsafe { Mmap::map(&file).unwrap() }; // 直接使用 safetensors 反序列化,拿到张量视图 let tensors = SafeTensors::deserialize(&mmap[..]).unwrap(); let view = tensors .tensor("bert.encoder.layer.0.attention.self.query.weight") .unwrap(); // 以 rank=1、world_size=4 为例,沿第 0 维切出本卡所需分片 VIEW[start..stop, :] let rank = 1; let world_size = 4; let dim = 0; let dtype = view.dtype(); let mut tp_shape = view.shape().to_vec(); let size = tp_shape[0]; if size % world_size != 0 { panic!("The dimension is not divisible by `world_size`"); } let block_size = size / world_size; let start = rank * block_size; let stop = (rank + 1) * block_size; // 一切按张量维度表达,字节偏移由 safetensors 自动处理 let iterator = view.slice(start..stop).unwrap(); tp_shape[dim] = block_size; // 将 safetensors 的 Dtype 转换为 candle 的 DType let dtype: DType = dtype.try_into().unwrap(); // 收集该分片的原始字节 let raw: Vec<u8> = iterator.into_iter().flatten().cloned().collect(); let tp_tensor = Tensor::from_raw_buffer(&raw, dtype, &tp_shape, &Device::Cpu).unwrap();这段代码的要点:
view.slice(start..stop)(来自safetensors::slice::IndexOp)在张量维度上切分,safetensors会自动把维度区间换算成字节偏移,无需手工计算data_offsets;- 每个 rank 只需设置自己的
rank与全局world_size,即可只把[start, stop)区间的数据收集到Vec<u8>; Tensor::from_raw_buffer(定义于 candle-core/src/safetensors.rs#L208-L288)按原始字节 + dtype + 形状直接构造张量,避免了一次中间分配(源码注释中的 TODO 也指出未来可进一步实现from_buffer_iterator以省去这段 CPU 拷贝);safetensors::Dtype到candle::DType的转换通过try_into()完成,映射关系可在 candle-core/src/safetensors.rs#L43-L64 中核对。
以bert-base-uncased的 query 权重(形状[768, 768])为例,world_size=4、rank=1时切出的分片形状为[192, 768]——candle-book 测试 candle-book/src/lib.rs#L107-L108(book_hub_3)正是用这两条断言验证了整个分片逻辑的正确性。
工程落地:缓存、revision 与进度条的参考实现
如果要在自己的项目中复刻 candle-examples 的完整下载体验,可以借鉴 candle-examples/src/hub.rs 的封装:
Api::with_cache_dir(cache_dir):通过hf_hub::HFClient::builder().cache_dir(...)自定义缓存目录;Repo::with_revision(revision):把后续所有下载固定到指定 revision,保证结果可复现;Repo::get(filename):先以local_files_only(true)尝试命中本地缓存,未命中才真正发起下载,并挂载StderrProgress进度处理器——在终端上以\r原地刷新百分比,重定向时按行输出文件名: 42% (42.0/100.0 MiB)形式的日志(该格式化逻辑与单位换算(B/KiB/MiB/GiB)均有单测覆盖于 hub.rs#L198-L234)。
小结
至此,一条完整的 Candle 权重接入链路已经打通:hf-hub下载(同步/异步、revision 固定、本地缓存)→candle::safetensors::load整文件解析为HashMap<String, Tensor>→ 按张量名接入Linear等candle-nn模块,或经VarBuilder::from_mmaped_safetensors直接加载进 BERT 等完整模型 → 面向大模型场景改用memmap2+load_buffer减少拷贝 → 面向多 GPU 张量并行场景用safetensors的slice只加载本 rank 分片。每一步都在仓库源码与 candle-book 测试(candle-book/src/lib.rs)中留下了可验证的实现与断言,你可以直接参考这些测试代码把流程复刻到自己的项目里。
- 人工智能
- 大模型
- 机器学习
- 深度学习
- 本地部署
- 模型推理服务
【免费下载链接】candle
Minimalist ML framework for Rust
相关推荐
Shortcircuit XT 免费采样器三平台支持全景指南:Windows、macOS、Linux
Shortcircuit XT 免费采样器三平台支持全景指南:Windows、macOS、Linux Shortcircuit XT 是一款由 Surge Sy
人工智能大模型机器学习深度学习本地部署模型推理服务Text Generation Inference 中的 Safetensors 权重格式:安全加载、张量并行分片与自动转换机制
Text Generation Inference 中的 Safetensors 权重格式:安全加载、张量并行分片与自动转换机制 Safetensors 是 T
模型推理服务大模型后端从 PyTorch、Transformers 与 Safetensors 三种途径加载 GPT-2 预训练权重(LLMs-from-scratch 实战指南)
从 PyTorch、Transformers 与 Safetensors 三种途径加载 GPT 2 预训练权重(LLMs from scratch 实战指南)
示例工程大模型人工智能
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考