Burn 的 burn-tensor 张量核心库:后端无关的 Tensor、Device 抽象与自动求值机制
【免费下载链接】burnBurn is a next generation tensor library and Deep Learning Framework that doesn't compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn
本文围绕 Burn 仓库中的 burn-tensor crate 展开,这是整个框架执行张量运算的核心抽象层。读完本文,你将理解Tensor<D, K>如何通过泛型与类型擦除实现"后端无关"、Device句柄如何统一选择 CPU/CUDA/WGPU 等计算设备,以及自动求值(autodiff)在此层之上的具体 API 与使用方式,并掌握通过 Cargo feature 组合出不同后端部署形态的方法。
核心定位:张量运算的统一抽象层
crates/burn-tensor/README.md 对 crate 的定位非常凝练:
This library provides the core abstractions required to run tensor operations with Burn.
Tensors are generic over the backend to allow users to perform operations using differentBackendimplementations. Burn's tensors also support auto-differentiation thanks to theAutodiffBackendtrait.
这句话浓缩了 burn-tensor 的三大设计支柱:
- 核心抽象——
Tensor、Shape、TensorData、Device等类型都在此 crate 中定义或再导出; - 后端泛型——同一份张量代码可以跑在 CPU、CUDA、ROCm、WGPU(Vulkan/Metal/WebGPU)、libtorch、嵌入式 no-std 环境等多种后端上;
- 自动求值——通过
AutodiffBackendtrait 为张量叠加求导能力,而无需更换张量类型本身。
从 crates/burn-tensor/Cargo.toml 的元数据也可以看到该 crate 的自我定位:categories = ["science", "no-std", "embedded", "wasm"],即它被设计为一个可运行到嵌入式与 Web 环境的张量库。其直接依赖仅四个同仓库 crate(burn-std、burn-backend、burn-dispatch、burn-derive),说明 burn-tensor 处在"用户 API"与"后端实现"之间的桥接位置。
Tensor 数据结构:Tensor<const D, K>的维度与类型双层泛型
张量的定义位于 crates/burn-tensor/src/tensor/api/base.rs:
pub struct Tensor<const D: usize, K = Float> where K: Basic, { pub(crate) primitive: BridgeTensor, _kind: PhantomData<K>, }两个泛型参数分别承担不同职责:
const D: usize(秩,rank):张量的维度数被编码为常量泛型参数而非运行时数据。这意味着Tensor<2>(2 维矩阵)与Tensor<3>(3 张量)在类型层面就是不同类型,reshape、slice等改变秩的操作必须在类型系统中显式体现,许多"维度不匹配"的错误可以在编译期就被拦截。K(张量种类):默认值为Float。K约束了张量上允许的操作集合——整数张量不能做浮点数学运算,布尔张量只能做逻辑运算。
结构体的内部实现值得注意:primitive字段的类型是BridgeTensor(类型擦除后的张量句柄),而非具体的后端张量类型;_kind: PhantomData<K>仅用于在类型层面标记种类。这种"类型擦除 + PhantomData"的组合是该 crate 编译时间优化的关键,见下文编译时间优化一节。
张量种类:Bool / Float / Int 的密封 trait 体系
K的可能取值定义在 crates/burn-tensor/src/tensor/kind.rs,通过一组分层 trait 精确描述每种张量"能做什么":
/// The base trait for any tensor kind. pub trait Basic: crate::ops::BasicOps {} /// Kinds that support numeric operations. pub trait Numeric: Basic + crate::ops::Numeric {} /// Kinds that support ordered operations. pub trait Ordered: Numeric + crate::ops::Ordered {} /// Kinds that support float math operations. pub trait FloatMath: Numeric + crate::ops::FloatMathOps {}从源码结构看,这套 trait 构成一个能力阶梯:Basic ⊂ Numeric ⊂ Ordered / FloatMath,Bool、Int、Float三个种类分别实现对应层级的 ops(BasicOps、Numeric、Ordered、FloatMathOps等)。文件头部注释明确标注这些 trait 是"sealed"(密封)的——外部代码无法自行实现Basic等 trait,只有 crate 内置的Bool、Float、Int可以作为K使用。这个封闭性保证了 burn-tensor 的 API 表面不会因为第三方自定义张量种类而失控。
Tensor的方法按K的能力分散在不同文件中:tensor/api/base.rs提供所有种类通用的方法(empty、zeros、from_data、shape、slice等),tensor/api/float.rs、tensor/api/int.rs、tensor/api/bool.rs分别挂接各种类特有的数学、索引、逻辑操作,cast.rs处理种类与 dtype 之间的转换。
常用构造与切片 API
base.rs的文档注释中给出了完整可运行的切片/索引示例(rustdoc测试代码):
use burn_tensor::Tensor; use burn_tensor::Int; let device = Default::default(); let tensor = Tensor::<2>::from_data( [ [3.0, 4.9, 2.0], [2.0, 1.9, 3.0], [6.0, 1.5, 7.0], [3.0, 4.9, 9.0], ], &device, ); // Slice: 取第 2、3 行 → shape [2, 3] let slice = tensor.clone().slice([1..3]); // Slice: 取前两行两列 → shape [2, 2] let slice = tensor.clone().slice([0..2, 0..2]); // select: 沿 dim 1 取第 0、2 列 → shape [4, 2] let indices = Tensor::<1, Int>::from_data([0, 2], &device); let indexed = tensor.select(1, indices);与之配套的创建方法族(empty/zeros/ones/from_data/arange/rand等)都接收impl Into<TensorCreationOptions>作为第二个参数,即"设备 + dtype"的组合选项——这体现了 burn-tensor 的一个 API 惯例:创建类操作把"放到哪个设备、用什么精度"统一收敛到 options 参数中,而不是散落成多个参数。
Device:统一设备句柄与后端选择
Tensor的每一个创建操作都需要指定设备。设备抽象定义在 crates/burn-tensor/src/device.rs:
pub struct Device { blob: device_opaque::Opaque, } // Aligned, type-erased storage for `DispatchDevice`. burn_std::obfuscate!( type: DispatchDevice, module: device_opaque, derives: [Send, Sync] );Device是一个类型擦除的高层设备句柄:内部用burn_std::obfuscate!宏把DispatchDevice(burn-dispatch层的真实设备类型)封装成一个不透明 blob,Device本身对外承诺Send + Sync。这样做有两个直接收益:
- API 稳定:下游 crate 看到的只有
Device,具体后端(CPU/CUDA/WGPU...)的类型树不会泄漏进公共接口; - 编译时间:与下文
*_impl辅助函数同理,避免下游泛型代码对 cubecl 类型树做单态化。
设备选择通过 Cargo feature + 工厂方法组合完成。device.rs中的文档注释给出了用法:
// 默认 CUDA 设备(需要 `cuda` feature) let device = Device::cuda(DeviceIndex::Default); // 硬件索引为 1 的 CUDA 设备 let device = Device::cuda(1); // 显式选择器的 WGPU 设备(`wgpu`/`vulkan`/`metal`/`webgpu`) let device = Device::wgpu(DeviceKind::DiscreteGpu(0)); // 任一已启用后端的默认设备 let device = Default::default();源码中实际提供按 feature 门控的工厂方法(均位于 crates/burn-tensor/src/device.rs):
| 工厂方法 | 参数形态 | 对应 feature |
|---|---|---|
Device::cpu() | 无 | cpu |
Device::cuda(index)/Device::rocm(index) | 整数索引或DeviceIndex | cuda/rocm |
Device::wgpu(kind)/Device::vulkan(kind)/Device::metal(kind)/Device::webgpu(kind) | DeviceKind选择器 | wgpu/vulkan/metal/webgpu |
Device::flex()/Device::ndarray() | 无 | flex/ndarray |
Device::libtorch()/Device::libtorch_cuda(index)/Device::libtorch_mps()/Device::libtorch_vulkan() | 可选索引 | tch |
Device::capture() | 无 | capture(图捕获后端) |
注意vulkan/metal/webgpu在 feature 层面都复用wgpu实现(见 Cargo.toml 中vulkan = ["wgpu", ...]的定义),区别仅在于设备选择策略与 kernel 编译目标。
Feature 组合即部署形态
crates/burn-tensor/Cargo.toml 的[features]表实际上是 burn-tensor 的"部署矩阵",可以归纳为四类:
- 后端选择:
cuda、rocm、wgpu、vulkan、metal、webgpu、cpu、flex、ndarray、tch。默认 feature 为["std", "burn-dispatch/default"]; - 能力开关:
autodiff(自动求值)、capture(图捕获)、fusion、autotune、simd、rayon等性能特性; - 远程计算:
remote(通过 Iroh 协议连接远端计算客户端)、remote-server(在本机托管远端计算服务,Wasm 兼容)、remote-websocket(旧版 WebSocket 传输); - 基础开关:
std(关闭即进入 no-std 模式)、tracing(操作追踪)。
一个典型的"纯推理 + 自动求值训练"配置就是同时启用cuda(或cpu)与autodiff;若目标是浏览器端推理,则启用webgpu即可——这正是categories = ["no-std", "embedded", "wasm"]所承诺的跨环境能力,由 feature 组合而非 fork 代码来区分。
自动求值:AutodiffBackend之上的张量 API
README 提到自动求值"thanks to theAutodiffBackendtrait"——该 trait 定义在burn-backend,burn-tensor 在其上暴露了面向张量的求导 API,全部位于 crates/burn-tensor/src/tensor/api/autodiff.rs,由autodifffeature 门控。
与旧式 Burn(用AdBackend<B>包一层后端类型)不同,当前代码中自动求值配置在 Device 上,而非张量类型参数上。device.rs的文档注释明确写道:"Autodiff support is configured on the device rather than through a separate type parameter." 完整的最小训练流程如下(来自源码 doc 示例):
let device = Device::default().autodiff(); // 在该设备上创建的张量即可参与求导图 let x = Tensor::<1>::from_floats([1.0, 2.0, 3.0], &device).require_grad(); // ... 用 x 计算 loss ... let grads = loss.backward(); let g = x.grad(&grads);autodiff.rs中暴露的核心方法:
| 方法 | 语义 |
|---|---|
backward() | 从该张量开始反向传播,返回Gradients容器。要求张量被 tracked(is_tracked()为真),且会消费共享的求导图 tape——重复调用会 panic |
grad(&grads) | 查询该张量保留的梯度(只读,可重复调用) |
grad_remove(&mut grads) | 取出并移除梯度,一次性场景可用它启用原地优化 |
grad_replace(&mut grads, new) | 用新张量替换该张量在grads中的梯度(例如手动注入梯度) |
is_tracked() | 判断该张量节点是否参与求导图。注意"启用 autodiff"与"参与图"是两回事:不要求梯度的常数在 autodiff 上下文中也不 tracked |
is_autodiff()/is_require_grad() | 分别报告"是否处于 autodiff 上下文"与"梯度是否被保留" |
Gradients本身同样是类型擦除容器(gradients_opaque::Opaque包装 dispatch 层的Gradients类型),保持了与Device一致的封装风格。反向传播的实际执行在backward_impl中一行转发:Dispatch::backward(...),即 burn-tensor 层只负责 API 校验(如assert!(self.is_tracked(), ...))与类型擦除,图构建与求导规则在burn-autodiff(经burn-dispatch转发)中实现——这符合仓库中 crates/burn-autodiff/src/backend.rs 所在的独立 crate 分工。
从kind.rs中未被 feature 门控删除的Autodifftrait 定义(pub trait Autodiff: Basic + crate::ops::BasicAutodiffOps {})可以推断:种类层级也为求导预留了能力位,Float种类在autodifffeature 下实现该 trait。
编译时间优化:*_impl辅助函数模式
burn-tensor 内部有一项对下游用户影响很大的工程约定,记录在 crates/burn-tensor/src/lib.rs 的 crate 级文档注释中:
公共泛型方法(如
tensor::api::float中的方法)在需要调用burn_dispatch时,会转发到文件底部一组名为*_impl的非泛型小函数。这些 helper 的签名只出现类型擦除的BridgeTensor——不出现任何burn_dispatch类型。由于 helper 不对D泛型,它们只被编译一次,公共泛型方法的 MIR 中不提及任何 dispatch 类型。下游 crate 单态化这些公共 API 时因此永远不需要解析 cubecl 类型树,大幅削减用户代码的编译时间。
这与 crates/burn-tensor/src/tensor/api/autodiff.rs 中的写法完全一致:backward、grad等公共方法体都很薄,最终落到backward_impl(p: &BridgeTensor)、grad_impl(...)等非泛型函数;device.rs中DispatchDevice的obfuscate!包装也是同一目的(注释原话:"it keeps the dispatch type tree out of downstream MIR")。对于要在大型项目中嵌入 Burn 的开发者,这个设计意味着TensorAPI 的下游编译时间不会因为 cubecl 这种"类型树很重"的依赖而爆炸。
模块全景与扩展生态
crate 的其余模块从 crates/burn-tensor/src/tensor/mod.rs 的导出结构可以一览:
kind:Bool/Float/Int种类与密封 trait;activation:激活函数直接操作张量(如 relu、sigmoid);signal:FFT、STFT、汉宁/汉明/布莱克曼窗等信号处理原语(见 crates/burn-tensor/src/tensor/signal 下的fft.rs、stft.rs等文件);grid:meshgrid、仿射网格等;loss/stats:张量级损失与统计工具;quantization:量化张量支持;distributed(std only):分布式张量;report(std only):内存池等运行时报告的再导出(SlicedPool、SlicedPoolReport等,源自 crates/burn-tensor/src/device.rs)。
mod.rs还大量再导出burn-std的类型(Shape、TensorData、DType、TensorReadError、Distribution等),因此用户只需use burn_tensor::*就能拿到形状、数据、容差、索引切片等全套配套类型,这是 burn-tensor 作为"唯一入口层"的典型体现。
此外,Tensor支持#[derive(Record)]生态所需的序列化:base.rs引入了serde::{Serialize, Deserialize, Serializer, Deserializer},张量可以被纳入 Burn 的Record状态体系(配合 burn-core 的模块序列化),这也是"核心抽象"定位的一部分——张量既是一等训练公民,也是可存取的一等状态公民。
小结
burn-tensor 用三层设计回答了"一份张量 API 如何服务所有后端"的问题:
- 类型层——
Tensor<const D, K>把秩与种类编码进类型系统,维度错误编译期可见;Bool/Float/Int的种类能力由密封 trait 阶梯约束; - 句柄层——
Device与内部BridgeTensor通过类型擦除隐藏burn-dispatch/cubecl 类型树,配合*_impl非泛型辅助函数压低下游编译时间; - 能力层——
autodiff、capture、remote、fusion等以 feature 叠加,同一套TensorAPI 在 CPU、GPU(CUDA/ROCm/WGPU 系)、libtorch、嵌入式 no-std 与 Web 之间切换,只改 Cargo feature 与Device工厂方法。
如需进一步阅读,可直接查看 Tensor 主 API 定义、Device 实现、autodiff API 以及 feature 配置;更完整的端到端用法(训练、推理、设备迁移)可参考 burn 主文档 与 basic-workflow 指南。
【免费下载链接】burnBurn is a next generation tensor library and Deep Learning Framework that doesn't compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考