news 2026/10/1 2:08:09

Candle 入门实战:在 Rust 中构建并运行你的第一个 MNIST 分类模型(Hello World 指南)

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Candle 入门实战:在 Rust 中构建并运行你的第一个 MNIST 分类模型(Hello World 指南)
  • 人工智能
  • 大模型
  • 机器学习
  • 深度学习
  • 本地部署
  • 模型推理服务

【免费下载链接】candle

Minimalist ML framework for Rust

项目地址:https://gitcode.com/GitHub_Trending/ca/candle
点击查看免费下载

本文档是 Candle 官方电子书(candle-book)的开篇实战教程,目标只有一个:用最少的代码,在 Rust 中构建并运行一个能够面向 MNIST 手写数字数据集做分类的两层神经网络(MLP)。你将依次经历三种实现方式——纯Tensor手写、自定义Linear层、以及直接使用candle-nn提供的现成Linear——并在此过程中掌握 Candle 的Device、Tensor、Module等核心抽象的真实用法。读完本篇,你就能理解 Candle 模型"从张量运算到模块封装"的演进路径,为后续跑通真实模型(推理)与训练打下基础。

前置准备:创建工程并引入 candle-core

Candle 是一个极简主义的 Rust 机器学习框架,核心张量库candle-core提供了Device(设备抽象)、Tensor(张量)、Result等基础类型。首先创建并进入一个新工程:

cargo new myapp cd myapp

然后添加candle-core依赖。官方标准做法是通过 git 源引入(完整安装说明见 installation.md):

cargo add --git https://github.com/huggingface/candle.git candle-core

如果你打算在本地仓库工作区内调试,也可以改用path依赖直接指向本仓库的 candle-core 子目录。

安装指南中还提供了几种可选加速特性,按需添加:

特性适用场景安装命令
cudaNVIDIA GPU 加速cargo add --git https://github.com/huggingface/candle.git candle-core --features "cuda"
cutile实验性 cuTile CUDA 后端(要求 Rust 1.89+、CUDA 13.2+、NVIDIA driver r580+)cargo add --git https://github.com/huggingface/candle.git candle-core --features "cutile"
mklCPU 上更快的推理(Intel MKL)cargo add --git https://github.com/huggingface/candle.git candle-core --features "mkl"
metalmacOS 上的 Metal GPU 加速cargo add --git https://github.com/huggingface/candle.git candle-core --features "metal"

添加完成后先运行cargo build确认依赖可正常编译。对于 CUDA 特性,安装指南要求先用nvcc --version和nvidia-smi --query-gpu=compute_cap --format=csv确认驱动与显卡算力(例如输出8.9),也可用CUDA_COMPUTE_CAP=<compute cap>环境变量指定要编译的算力版本。

第一版:只用 Tensor 完成的两层网络

打开src/main.rs,填入下面的内容(源自 hello_world.md):

use candle_core::{Device, Result, Tensor}; struct Model { first: Tensor, second: Tensor, } impl Model { fn forward(&self, image: &Tensor) -> Result<Tensor> { let x = image.matmul(&self.first)?; let x = x.relu()?; x.matmul(&self.second) } } fn main() -> Result<()> { // Use Device::new_cuda(0)?; to use the GPU. let device = Device::Cpu; let first = Tensor::randn(0f32, 1.0, (784, 100), &device)?; let second = Tensor::randn(0f32, 1.0, (100, 10), &device)?; let model = Model { first, second }; let dummy_image = Tensor::randn(0f32, 1.0, (1, 784), &device)?; let digit = model.forward(&dummy_image)?; println!("Digit {digit:?} digit"); Ok(()) }

这段代码只有四层逻辑,却是理解 Candle 的绝佳入口:

1. 设备(Device)抽象。Device::Cpu代表在 CPU 上计算;想切换 GPU,只需把注释里的Device::new_cuda(0)?换上来即可。从 device.rs 的源码可以看到,Device是一个枚举,支持Cpu、Cuda、Metal三种后端,new_cuda(ordinal)通过CudaDevice::new(ordinal)创建第ordinal块 GPU 设备。此外还有更省心的Device::cuda_if_available(0)?(device.rs),它会先探测 CUDA 是否可用,可用则创建 CUDA 设备,否则自动回退到 CPU——这一行我们会在第二版里用到。

2. 随机初始化。Tensor::randn(0f32, 1.0, (784, 100), &device)?创建一个形状为(784, 100)的张量,元素服从均值为0、标准差为1.0的标准正态分布。从 tensor.rs 的签名pub fn randn<S: Into<Shape>, T: crate::FloatDType>(mean: T, std: T, s: S, device: &Device)可以看到,均值与标准差类型与张量数据类型一致(这里用0f32而不是0,就是为了让类型推断落在f32上)。

3. 形状的语义。MNIST 的每张灰度图是 28×28 像素,展平后就是 784 个元素;数据集共有 10 类(数字 0–9)。所以第一层权重(784, 100)把输入从 784 维投影到 100 维隐藏空间,第二层权重(100, 10)再投影到 10 维的类别得分空间。dummy_image的形状(1, 784)代表"一张 784 维的样本"。

4. 运算与错误处理。image.matmul(&self.first)?是矩阵乘法,relu()是 ReLU 激活函数,两者都是张量的内建方法(matmul定义在 tensor.rs)。Candle 沿用了 Rust 生态的?运算符做错误传播,所有可能失败的算子都返回Result<Tensor>,因此在main里声明返回类型Result<()>,并把每个运算都用?接住——这也是 Candle 的一个核心设计:任何设备后端(CPU/CUDA/Metal)都可能产生错误,必须在类型系统中显式处理。

编译运行:

cargo run --release

程序会打印一个随机的 10 维向量(当前权重完全随机,还没有学习任何东西),例如Digit [0.31, -0.12, ...] digit。至此,你的第一个 Candle 模型已经跑起来了。

第二版:自定义一个带偏置的 Linear 层

真实网络通常还需要偏置(bias)。原文档借此演示"如何用张量运算自己拼装层"——先定义经典的Linear结构:

use candle_core::{Device, Result, Tensor}; struct Linear { weight: Tensor, bias: Tensor, } impl Linear { fn forward(&self, x: &Tensor) -> Result<Tensor> { let x = x.matmul(&self.weight)?; x.broadcast_add(&self.bias) } } struct Model { first: Linear, second: Linear, } impl Model { fn forward(&self, image: &Tensor) -> Result<Tensor> { let x = self.first.forward(image)?; let x = x.relu()?; self.second.forward(&x) } }

注意这里forward的结果:x.matmul(&self.weight)?之后得到一个形状为(batch, out_dim)的结果,而bias的形状是(out_dim,),两者无法直接相加。broadcast_add正是 Candle 的广播加法——它会自动把偏置向量沿着 batch 维度"展开",实现y = x @ w + b的经典线性变换语义。Model的forward则由两个Linear与中间的relu组成。

对应的main函数改为用Device::cuda_if_available(0)?选择设备,并分别创建两层的权重与偏置:

fn main() -> Result<()> { // Use Device::new_cuda(0)?; to use the GPU. // Use Device::Cpu; to use the CPU. let device = Device::cuda_if_available(0)?; // Creating a dummy model let weight = Tensor::randn(0f32, 1.0, (784, 100), &device)?; let bias = Tensor::randn(0f32, 1.0, (100, ), &device)?; let first = Linear{weight, bias}; let weight = Tensor::randn(0f32, 1.0, (100, 10), &device)?; let bias = Tensor::randn(0f32, 1.0, (10, ), &device)?; let second = Linear{weight, bias}; let model = Model { first, second }; let dummy_image = Tensor::randn(0f32, 1.0, (1, 784), &device)?; // Inference on the model let digit = model.forward(&dummy_image)?; println!("Digit {digit:?} digit"); Ok(()) }

这段代码展示了 Candle 组装自定义层的完整套路:层就是一个持有Tensor权重的结构体,forward就是张量运算的组合,模型则是层的嵌套。这也是官方文档强调的"这是创建你自己的层的好方法"。

第三版:直接使用 candle-nn 的 Linear

自己写层虽然直观,但经典的层在candle-nn中大多已有实现。添加依赖:

cargo add --git https://github.com/huggingface/candle.git candle-nn

然后改写示例(原文档的第三个版本):

use candle_core::{Device, Result, Tensor}; use candle_nn::{Linear, Module}; struct Model { first: Linear, second: Linear, } impl Model { fn forward(&self, image: &Tensor) -> Result<Tensor> { let x = self.first.forward(image)?; let x = x.relu()?; self.second.forward(&x) } } fn main() -> Result<()> { // Use Device::new_cuda(0)?; to use the GPU. let device = Device::Cpu; // This has changed (784, 100) -> (100, 784) ! let weight = Tensor::randn(0f32, 1.0, (100, 784), &device)?; let bias = Tensor::randn(0f32, 1.0, (100, ), &device)?; let first = Linear::new(weight, Some(bias)); let weight = Tensor::randn(0f32, 1.0, (10, 100), &device)?; let bias = Tensor::randn(0f32, 1.0, (10, ), &device)?; let second = Linear::new(weight, Some(bias)); let model = Model { first, second }; let dummy_image = Tensor::randn(0f32, 1.0, (1, 784), &device)?; let digit = model.forward(&dummy_image)?; println!("Digit {digit:?} digit"); Ok(()) }

与原版相比有两个关键差异,值得重点理解:

差异一:权重形状反了——(100, 784)而不是(784, 100)。原文档明确说明:candle-nn的Linear是按照 PyTorch 的布局习惯设计的,为了最大化复用现有模型(PyTorch 权重文件),它"使用权重的转置而不是权重本身"。看 linear.rs 的forward实现即可印证:

impl super::Module for Linear { fn forward(&self, x: &Tensor) -> candle::Result<Tensor> { let x = match *x.dims() { [b1, b2, m, k] => { /* ... */ } [bsize, m, k] => { /* ... */ } _ => { let w = self.weight.t()?; // 先转置 x.matmul(&w)? } }; match &self.bias { None => Ok(x), Some(bias) => x.broadcast_add(bias), } } }

所以Linear::new(weight, bias)要求你传入的权重形状是(out_dim, in_dim)——也就是(100, 784)——在forward内部通过weight.t()转置成(784, 100)再与输入做matmul,最终输出形状为(batch, out_dim)。这一设计让 PyTorch 导出的权重(safetensors 格式)可以直接原样喂给 Candle。

差异二:偏置是可选的。Linear::new的签名是pub fn new(weight: Tensor, bias: Option<Tensor>) -> Self(linear.rs),用Some(bias)/None表达是否启用偏置。当偏置为None时,forward直接返回x.matmul(&w)的结果,不加任何广播加法。

此外,Linear的forward对多 batch 维度输入做了优化:从源码可以看到,对于[b1, b2, m, k]或[bsize, m, k]形状的输入,如果输入是连续的(is_contiguous()),它会先reshape成二维再matmul,并在注释中说明"广播式 matmul 在 cuda 和 cpu 后端上比标准 matmul 慢得多"——这是 Candle 为推理性能做的典型优化。

Moduletrait 与ModuleT。上面的代码中use candle_nn::{Linear, Module},其中Moduletrait 定义在 lib.rs:

pub trait Module { fn forward(&self, xs: &Tensor) -> Result<Tensor>; }

任何实现了Module的类型都可以用统一的forward(&Tensor) -> Result<Tensor>接口被串联、嵌套。Candle 还为闭包、Option<&M>自动实现了Module,并提供了带训练标志的ModuleTtrait(forward_t(&self, xs, train)),用于区分训练与推理行为(如 Dropout)。这也是为什么官方示例可以写成self.first.forward(image)?这种完全一致的调用风格。

linear()工厂函数与默认初始化。除了手搓权重再Linear::new,candle-nn还提供了更省事的工厂函数(linear.rs):

pub fn linear(in_dim: usize, out_dim: usize, vb: crate::VarBuilder) -> Result<Linear> { let init_ws = crate::init::DEFAULT_KAIMING_NORMAL; let ws = vb.get_with_hints((out_dim, in_dim), "weight", init_ws)?; let bound = 1. / (in_dim as f64).sqrt(); let init_bs = crate::Init::Uniform { lo: -bound, up: bound }; let bs = vb.get_with_hints(out_dim, "bias", init_bs)?; Ok(Linear::new(ws, Some(bs))) }

它会从VarBuilder(变量构建器,通常用于从 safetensors 权重文件或随机初始化中取参)按默认名字"weight"、"bias"取张量,权重用 Kaiming Normal 初始化、偏置用±1/√in_dim均匀分布初始化;linear_no_bias则跳过偏置。这也是真实项目(例如仓库里大量的模型示例)中创建Linear的标准姿势。

下一步:从 Hello World 走向真实模型

原文档在结尾给出了三条进阶路径,也值得我们在此一一展开:

1. 换成卷积网络。原文档建议读者动手把示例中的Linear换成Conv2d来构建经典卷积网络——candle_nn中已实现Conv2d(见 conv.rs),配合candle-examples/examples/mnist-training中的训练示例(mnist-training)可以对照 MLP 与 CNN 的差异。

2. 准备 MNIST 数据集。真实的 MNIST 包含 60,000 张训练图与 10,000 张测试图,每张图展平为 784 维向量、共 10 类。仓库中 candle-book/src/lib.rs 展示了如何用hf_hub下载 MNIST 的 parquet 文件、解码图片并把整份数据集装入内存张量的完整代码(train_images形状[60_000, 784]、test_images形状[10_000, 784]),而 candle-datasets 则提供了candle_datasets::vision::Dataset这样的现成数据结构,方便直接喂给模型训练。相关训练代码参见 training/mnist.md 与 training/training.md。

3. 对标 PyTorch。如果你是 PyTorch 用户,cheatsheet.md 提供了一张 Candle 与 PyTorch 的逐项对照表(内容摘自仓库根目录 README.md 的 "Cheatsheet" 锚点),覆盖张量创建(Tensor::newvstorch.Tensor)、索引(tensor.i((.., ..4))vstensor[:, :4])、视图(reshapevsview)、设备迁移(to_devicevsto(device="cuda"))、dtype 转换、以及 safetensors 的保存与加载等高频操作,是快速上手 Candle 的最短路径。

4. 运行真实模型。当你理解了这个最小骨架后,就可以进入"运行现有模型"阶段:inference/inference.md 讲解如何加载 safetensors 权重并用VarBuilder组装出 BERT、LLaMA 等真实模型,hub.md 则讲解从 Hugging Face Hub 下载权重的流程;仓库中 candle-examples 下的上百个模型示例(bert、llama、qwen、whisper、stable-diffusion 等)都可以作为实战参考。

从一段 20 行的张量代码,到理解设备抽象、广播语义、PyTorch 兼容的权重布局与Module接口,你已经走完了 Candle 学习曲线的第一段。接下来无论是啃 cheatsheet、跑推理还是写训练循环,都只是在这套骨架上添砖加瓦。

  • 人工智能
  • 大模型
  • 机器学习
  • 深度学习
  • 本地部署
  • 模型推理服务

【免费下载链接】candle

Minimalist ML framework for Rust

项目地址:https://gitcode.com/GitHub_Trending/ca/candle
点击查看免费下载
上一篇:todo[bot]性能优化:大规模项目的自动化Issue管理策略
下一篇:单模型双模式革命:Qwen3-14B-FP8如何重新定义企业AI部署成本

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/1 2:06:46

SQL面试常见问题:查询及删除重复记录的方法与TaoToken实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/10/1 2:06:03

个人开发者单卡RTX 3090实战:GPT-2预训练与领域适配全流程

1. 为什么个人开发者也要走一遍LLM全流程很多人一提到大模型&#xff0c;第一反应就是“这玩意儿得几百张卡才能玩”。我一开始也这么想&#xff0c;直到自己用一张RTX 3090把GPT-2从预训练一路做到领域适配&#xff0c;才发现个人开发者和工业级团队之间的差距&#xff0c;其实…

作者头像 李华
网站建设 2026/10/1 2:05:46

OpenDaylight安装避坑指南:Java环境、版本兼容与Karaf启动全解析

1. 这不是普通软件安装&#xff1a;OpenDaylight 是网络操作系统&#xff0c;装错一步就卡在 Karaf 控制台里出不来OpenDaylight&#xff08;ODL&#xff09;不是你点几下“下一步”就能装好的桌面应用。它本质是一个基于 OSGi 架构的、面向 SDN&#xff08;软件定义网络&#…

作者头像 李华