
Burn 的 burn-tensor 张量核心库后端无关的 Tensor、Device 抽象与自动求值机制【免费下载链接】burnBurn is a next generation tensor library and Deep Learning Framework that doesnt compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn本文围绕 Burn 仓库中的 burn-tensor crate 展开这是整个框架执行张量运算的核心抽象层。读完本文你将理解TensorD, 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. Burns tensors also support auto-differentiation thanks to theAutodiffBackendtrait.这句话浓缩了 burn-tensor 的三大设计支柱核心抽象——Tensor、Shape、TensorData、Device等类型都在此 crate 中定义或再导出后端泛型——同一份张量代码可以跑在 CPU、CUDA、ROCm、WGPUVulkan/Metal/WebGPU、libtorch、嵌入式 no-std 环境等多种后端上自动求值——通过AutodiffBackendtrait 为张量叠加求导能力而无需更换张量类型本身。从 crates/burn-tensor/Cargo.toml 的元数据也可以看到该 crate 的自我定位categories [science, no-std, embedded, wasm]即它被设计为一个可运行到嵌入式与 Web 环境的张量库。其直接依赖仅四个同仓库 crateburn-std、burn-backend、burn-dispatch、burn-derive说明 burn-tensor 处在用户 API与后端实现之间的桥接位置。Tensor 数据结构Tensorconst D, K的维度与类型双层泛型张量的定义位于 crates/burn-tensor/src/tensor/api/base.rspub struct Tensorconst D: usize, K Float where K: Basic, { pub(crate) primitive: BridgeTensor, _kind: PhantomDataK, }两个泛型参数分别承担不同职责const D: usize秩rank张量的维度数被编码为常量泛型参数而非运行时数据。这意味着Tensor22 维矩阵与Tensor33 张量在类型层面就是不同类型reshape、slice等改变秩的操作必须在类型系统中显式体现许多维度不匹配的错误可以在编译期就被拦截。K张量种类默认值为Float。K约束了张量上允许的操作集合——整数张量不能做浮点数学运算布尔张量只能做逻辑运算。结构体的内部实现值得注意primitive字段的类型是BridgeTensor类型擦除后的张量句柄而非具体的后端张量类型_kind: PhantomDataK仅用于在类型层面标记种类。这种类型擦除 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 / FloatMathBool、Int、Float三个种类分别实现对应层级的 opsBasicOps、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 之间的转换。常用构造与切片 APIbase.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 IntoTensorCreationOptions作为第二个参数即设备 dtype的组合选项——这体现了 burn-tensor 的一个 API 惯例创建类操作把放到哪个设备、用什么精度统一收敛到 options 参数中而不是散落成多个参数。Device统一设备句柄与后端选择Tensor的每一个创建操作都需要指定设备。设备抽象定义在 crates/burn-tensor/src/device.rspub 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!宏把DispatchDeviceburn-dispatch层的真实设备类型封装成一个不透明 blobDevice本身对外承诺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工厂方法参数形态对应 featureDevice::cpu()无cpuDevice::cuda(index)/Device::rocm(index)整数索引或DeviceIndexcuda/rocmDevice::wgpu(kind)/Device::vulkan(kind)/Device::metal(kind)/Device::webgpu(kind)DeviceKind选择器wgpu/vulkan/metal/webgpuDevice::flex()/Device::ndarray()无flex/ndarrayDevice::libtorch()/Device::libtorch_cuda(index)/Device::libtorch_mps()/Device::libtorch_vulkan()可选索引tchDevice::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之上的张量 APIREADME 提到自动求值thanks to theAutodiffBackendtrait——该 trait 定义在burn-backendburn-tensor 在其上暴露了面向张量的求导 API全部位于 crates/burn-tensor/src/tensor/api/autodiff.rs由autodifffeature 门控。与旧式 Burn用AdBackendB包一层后端类型不同当前代码中自动求值配置在 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容器。要求张量被 trackedis_tracked()为真且会消费共享的求导图 tape——重复调用会 panicgrad(grads)查询该张量保留的梯度只读可重复调用grad_remove(mut grads)取出并移除梯度一次性场景可用它启用原地优化grad_replace(mut grads, new)用新张量替换该张量在grads中的梯度例如手动注入梯度is_tracked()判断该张量节点是否参与求导图。注意启用 autodiff与参与图是两回事不要求梯度的常数在 autodiff 上下文中也不 trackedis_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 的导出结构可以一览kindBool/Float/Int种类与密封 traitactivation激活函数直接操作张量如 relu、sigmoidsignalFFT、STFT、汉宁/汉明/布莱克曼窗等信号处理原语见 crates/burn-tensor/src/tensor/signal 下的fft.rs、stft.rs等文件gridmeshgrid、仿射网格等loss/stats张量级损失与统计工具quantization量化张量支持distributedstd only分布式张量reportstd 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 如何服务所有后端的问题类型层——Tensorconst D, K把秩与种类编码进类型系统维度错误编译期可见Bool/Float/Int的种类能力由密封 trait 阶梯约束句柄层——Device与内部BridgeTensor通过类型擦除隐藏burn-dispatch/cubecl 类型树配合*_impl非泛型辅助函数压低下游编译时间能力层——autodiff、capture、remote、fusion等以 feature 叠加同一套TensorAPI 在 CPU、GPUCUDA/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 doesnt compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考