Burn 深度学习框架后端与设备(Backend & 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 是一个将运行时设备(Device)作为一等公民的深度学习框架:Tensor、Module 与 Device 三者构成用户侧的核心 API,任何张量都携带一个标识其执行位置与方式的运行时Device,从而让同一套模型与张量代码可以在 CPU、CUDA、ROCm、WGPU(Vulkan/Metal/WebGPU)乃至远程设备之间无缝切换。本文将基于 burn-book 的 Backend and Device 章节,结合 burn-tensor 的设备实现 与 burn-std 的设备设置模块 源码,系统讲解设备构造器的选择、张量与模块的设备迁移、autodiff 上下文的开启与关闭、设备默认数据类型的配置,以及 Tensor → Bridge → Dispatch → Backend 的底层执行栈,帮助你写出可移植、可复现且具备多硬件部署能力的 Burn 应用。
核心设计:后端不再是泛型参数
在 Burn 中,Tensor、Module和Device并不针对某个后端做泛型化(generic over a backend)。相反,每一个张量都携带一个运行时的Device,由它来决定张量上的操作在哪里、以何种方式执行:
use burn::tensor::{Device, Tensor}; let device = Device::wgpu(Default::default()); let tensor = Tensor::<2>::ones([2, 3], &device); // 同一套模型和张量 API 可以无缝切换到另一个后端。 let other_device = Device::cuda(0); let tensor = tensor.to_device(&other_device);这意味着你编写的训练、推理代码不需要引入B: Backend之类的泛型约束即可运行在任何后端之上。Backend与AutodiffBackendtrait 仍然定义了底层的实现契约,但普通应用代码几乎不会直接接触它们——只有当你实现或扩展后端时才需要(参见 后端扩展指南)。
从源码结构看,这一设计在 crates/burn-tensor/src/device.rs 中得到印证:Device是一个薄封装结构,内部通过burn_std::obfuscate!宏将具体的DispatchDevice类型擦除为对齐的透明存储(device_opaque::Opaque),这样下游代码既不需要知道 dispatch 的类型树,又能通过as_dispatch()/into_dispatch()在需要时访问底层设备。
需要特别注意的是:每个设备构造器都对应一个 Cargo feature。例如启用wgpu与cudafeature 后,Device::wgpu与Device::cuda才可用。这是 Burn 按需裁剪后端支持的基础机制。
选择设备:构造器全景与 feature 对应关系
Device为构建中启用的后端提供对应的构造器。常见选择如下表所示:
| 构造器 | 目标 |
|---|---|
Device::wgpu(Default::default()) | 可用的最佳 WGPU 适配器 |
Device::wgpu(DeviceKind::DiscreteGpu(0)) | 通过 WGPU 使用第一块独立 GPU |
Device::vulkan(Default::default()) | 最佳 Vulkan 适配器 |
Device::metal(Default::default()) | 最佳 Metal 适配器 |
Device::webgpu(Default::default()) | 浏览器中的 WebGPU 设备 |
Device::cuda(0) | 索引为 0 的 CUDA GPU |
Device::cuda(DeviceIndex::Default) | 由后端挑选默认 CUDA GPU |
Device::rocm(0) | 索引为 0 的 ROCm/HIP GPU |
Device::cpu() | CubeCL CPU 后端 |
Device::flex() | Flex CPU 后端 |
Device::ndarray() | NdArray CPU 后端(已弃用) |
Device::libtorch() | LibTorch CPU 后端(已弃用) |
Device::libtorch_cuda(0) | 索引为 0 的 LibTorch CUDA GPU(已弃用) |
Device::libtorch_mps() | LibTorch Metal Performance Shaders(已弃用) |
Device::libtorch_vulkan() | LibTorch Vulkan 设备(已弃用) |
结合 crates/burn-tensor/src/device.rs 的源码可以确认每个构造器的 feature 门槛与具体行为:
Device::cpu()需要cpufeature,返回 CubeCL CPU 运行时设备;Device::cuda(index)需要cudafeature,内部通过CudaDevice::new(index.into().resolve())解析硬件索引;Device::rocm(index)需要rocmfeature,语义与cuda相同;Device::flex()需要flexfeature;Device::ndarray()、Device::libtorch()及其变体自 0.22.0 起标记为#[deprecated],源码中的弃用说明明确建议改用 CubeCL 系后端(Device::cuda、Device::rocm、Device::metal、Device::vulkan、Device::cpu)或Device::flex()——在编写新代码时应优先选择这些非弃用构造器。
DeviceIndex:索引型后端的硬件选择
DeviceIndex是索引型后端(CUDA、ROCm)的硬件选择器,定义于 crates/burn-tensor/src/device.rs,仅有两个变体:
Specified(usize):指定具体硬件索引;Default(默认值):由后端自行选择默认设备(通常是索引 0)。
由于后端工厂方法接受impl Into<DeviceIndex>,而usize/u32/u64/i32/i64都实现了From转换(其中负数会 panic,见 device.rs),所以最常见的写法就是直接传整数:Device::cuda(0)等价于Device::cuda(DeviceIndex::Specified(0)),而Device::cuda(DeviceIndex::Default)则把选择权交给后端。
DeviceKind:WGPU 家族的适配器选择
DeviceKind服务于 WGPU 这类设备句柄是带标签枚举的灵活后端,其变体与 CubeCL 的WgpuDevice一一对应,但作为 Burn 自有枚举避免调用方直接依赖 cubecl(见 device.rs):
DiscreteGpu(usize):系统中第 N 块独立 GPU;IntegratedGpu(usize):系统中第 N 块集成 GPU;VirtualGpu(usize):系统中第 N 块虚拟 GPU;Cpu:CPU 适配器;DefaultDevice(默认值):由 wgpu 的启发式规则挑选最佳可用设备(优先"高功耗" GPU),并可通过环境变量CUBECL_WGPU_DEFAULT_DEVICE覆盖,例如CUBECL_WGPU_DEFAULT_DEVICE=IntegratedGpu(1)或CUBECL_WGPU_DEFAULT_DEVICE=Cpu;Existing(u32):复用外部已创建的 wgpu 环境(例如与 egui、bevy 共存时通过 wgpu runtime 的init_device初始化),便于资源进出 CubeCL。
组合使用示例:
use burn::tensor::{Device, DeviceIndex, DeviceKind}; let default_wgpu = Device::wgpu(Default::default()); let discrete_wgpu = Device::wgpu(DeviceKind::DiscreteGpu(0)); let integrated_wgpu = Device::wgpu(DeviceKind::IntegratedGpu(0)); let default_cuda = Device::cuda(DeviceIndex::Default); let second_cuda = Device::cuda(1);源码中的wgpu_device()辅助函数(device.rs)展示了Device::wgpu、Device::vulkan、Device::metal、Device::webgpu这四个构造器如何共享同一映射逻辑——它们只是 feature 门槛不同,运行时返回的都是 WGPU 设备:着色语言编译器(WGSL/SPIR-V/MSL)根据启用的 feature 在运行时自动选择,而不是由构造器固定。
远程设备:跨机器执行
当启用对应的远程 feature 时,Burn 还支持远程设备:
Device::remote_websocket(address, index)(remote-websocketfeature):连接burn-remote的 WebSocket 服务器,index选择服务器上的第几个设备——相同地址不同索引指向同一主机上的不同设备;Device::remote_iroh(endpoint, peer, index)(remotefeature,非 wasm 环境):通过 Iroh 点对点通道连接计算服务器,peer是服务器的身份标识;wasm 环境需改用异步版本remote_iroh_async。
从源码看(device.rs),远程构造器在返回前会调用device.connect()建立连接(这是获取设备默认设置所必需的),此外还提供了携带授权凭证的remote_iroh_authorized系列变体,供要求凭据的服务器使用。
使用设备:创建张量、初始化模块与迁移
Tensor 的创建方法接收&Device;模块在初始化参数时也会收到同一个设备:
let device = Device::cuda(0); let input = Tensor::<2>::zeros([32, 128], &device); let model = ModelConfig::new().init(&device); let output = model.forward(input);迁移已有值时,使用Tensor::to_device或Module::to_device。涉及多个张量的运算要求设备兼容,因此在组合之前必须显式移动它们:
let cpu = Device::flex(); let gpu = Device::cuda(0); let tensor = Tensor::<2>::ones([2, 3], &cpu); let tensor = tensor.to_device(&gpu); let model = model.to_device(&gpu);Tensor::to_device的实现位于 crates/burn-tensor/src/tensor/api/base.rs,它作为张量 API 的一部分为所有张量类型提供统一的跨设备移动能力。
Module::fork 与 Module::to_device 的区别
若要移动一个可训练模块并在目标设备上优化其参数,应使用Module::fork;而Module::to_device会保留与源参数的连接(例如与 autodiff 图的关联)。关于模块转移语义的完整说明见 Module 方法章节,其中以表格形式列出了module.fork(device)、module.to_device(device)、module.no_grad()、module.freeze()、module.into_record()等全部模块方法及其与 PyTorch 的对应关系。
Autodiff 与执行控制
设备为新建张量提供 autodiff 与 checkpointing 的默认上下文。每个张量携带自己的上下文,之后可以独立改变。Tensor::to_device会保留源张量的上下文,无论目标设备的默认值如何。调用autodiff会返回一个"新建张量默认开启 autodiff"的设备:
let device = Device::wgpu(Default::default()); let training_device = device.autodiff(); assert!(training_device.is_autodiff()); let inference_device = training_device.without_autodiff(); assert!(!inference_device.is_autodiff());关键语义(均可从 device.rs 源码确认):
autodiff()与without_autodiff()都是幂等的:对已开启 autodiff 的设备再次调用autodiff()会原样返回并保留其梯度 checkpointing 策略;重复调用并不会启用高阶求导——Burn 目前仅支持一阶自动微分。- 历史上的
inner()方法等价于without_autodiff(),它反映的是旧的"后端装饰器"模型,而without_autodiff直接描述了设备的运行时属性。 - 链式调用
autodiff().gradient_checkpointing()可启用带"平衡"(Balanced)checkpointing 策略的 autodiff;该策略对标记为内存受限(memory-bound)的算子在前向传播时重新计算激活值,而对计算受限(compute-bound)的算子仍缓存输出,从而降低峰值内存。若设备未开启 autodiff 就调用gradient_checkpointing(),会触发 panic(源码测试gradient_checkpointing_requires_autodiff验证了这一行为,见 device.rs)。 - 设备相等性只比较计算资源,忽略 autodiff 与 checkpointing 设置。因此判断张量是否参与自动微分应使用
tensor.is_autodiff()——与一个 autodiff 设备相等并不代表该张量已开启 autodiff。PartialEq的实现文档明确写道:当执行上下文也重要时,请检查Device::is_autodiff与gradient_checkpointing_strategy()(见 device.rs)。
以下是协调执行时常用的方法清单:
seed(seed):为设备上的随机操作设置种子。在涉及随机性的张量操作(如Tensor::random)之前调用可让单线程程序的结果可复现;注意某些后端可能将种子全局应用,而非严格限定在该设备上。sync():等待排队的工作完成,若发生执行错误则返回ExecutionError。返回的ExecutionError类型定义于 crates/burn-std/src/device_settings.rs,包含WithContext(外部传入上下文)与Generic(Burn 内部错误,附带惰性解析的 backtrace)两个变体。flush():提交排队的工作但不等待完成。其源码文档(device.rs)指出:fusion 后端会缓存算子以构建优化、远程后端会批量攒积后经网络发送,flush()强制立即排空这些队列;而即时执行(eager)后端没有缓冲,此操作是空操作。is_autodiff():报告设备是否关联了 autodiff。这只是新张量继承的设备上下文,不代表任何张量参与计算图或保留梯度。gradient_checkpointing_strategy():返回当前生效的策略;无 autodiff 时返回None。supports_dtype(dtype):报告设备是否支持该数据类型(存储、转换与算术均支持)。文档特别提醒:某些类型可能"可存储可转换但无算术支持"——例如 Vulkan 设备上的 bf16,SPIR-V 的SPV_KHR_bfloat16仅允许转换、点积和 cooperative-matrix 使用,直接计算会得到后端相关的垃圾结果,因此选择低精度前务必检查:
let dtype = if device.supports_dtype(FloatDType::BF16) { FloatDType::BF16 } else { FloatDType::F32 };memory_cleanup():请求后端释放未使用的缓存分配。回收多少取决于分配器实现,不保证一定释放内存(源码文档明确说明这一点,见 device.rs)。相关能力还有memory_persistent_allocations、memory_install_pools、memory_pool_report、memory_pool_usage等,用于精细控制动态内存池。
这些设备方法最终都汇聚到 dispatch 层——crates/burn-dispatch/src/backend.rs 中Dispatch的seed、sync、ad_enabled、memory_cleanup、supports_dtype、flush等接口定义了后端必须实现或可覆盖的行为契约。
设备设置:默认数据类型与一次性锁定
每个设备都有自己的运行时设置,包括默认的浮点、整数与布尔数据类型。通过settings()查看、configure()设置:
use burn::tensor::{Device, DeviceConfig, FloatDType, IntDType}; let mut device = Device::cuda(0); device.configure( DeviceConfig::default() .float_dtype(FloatDType::F16) .int_dtype(IntDType::I32), )?; let settings = device.settings();DeviceConfig定义于 crates/burn-tensor/src/device.rs,是用户提供的部分配置:float_dtype、int_dtype、bool_dtype均为Option,未指定的项会在设备初始化时解析为设备特定默认值。它还实现了From<FloatDType>、From<IntDType>、From<BoolDType>、From<(FloatDType, IntDType)>等转换,因此device.configure(FloatDType::F16)?这样省略包装器的写法同样合法。
configure的实现(device.rs)会读取设备的默认值,用配置项覆盖后调用set_default_dtypes,最终写入一个全局注册表。该注册表强制执行严格的初始化语义(见 crates/burn-std/src/device_settings.rs):
- 手动初始化:可在程序开头通过
configure/set_default_dtypes设置一次; - 默认初始化:若在手动初始化之前发生了任何操作(比如创建张量),设置将被永久锁定为默认值;
- 不可变性:一旦初始化,设置不可更改,以保证跨线程、跨操作的行为一致。
因此:务必在该设备上的第一次张量操作之前完成配置。初始化之后,默认 dtype 即被锁定,任何后续不兼容的配置都会返回DeviceError::AlreadyInitialized。DeviceError定义于 crates/burn-std/src/device_settings.rs,共两个变体:UnsupportedDType(设备不支持请求的数据类型)与AlreadyInitialized(设置已初始化且不可更改)。
枚举设备:动态选择硬件与多设备训练
Device::enumerate可以发现所有匹配过滤器(DeviceFilter)的已启用设备,适合动态选择硬件或搭建多设备训练:
use burn::tensor::{Device, DeviceFilter, DeviceType}; let devices = Device::enumerate( DeviceFilter::new() .with(DeviceType::Cuda) .with(DeviceType::Wgpu), ); let devices = devices.into_vec();从源码看,可用的DeviceType变体完全取决于应用启用的后端 feature(device.rs):Cpu、Cuda、Rocm、Wgpu、Metal、Vulkan、WebGpu、Flex、NdArray、LibTorch以及携带服务器地址的Remote(String)。几个值得注意的细节:
DeviceType变体之间可以通过|运算符组合成DeviceFilter(也支持从单个DeviceType或Vec<DeviceType>转换),例如DeviceType::Cuda | DeviceType::remote_websocket("ws://host:3000")一次调用即可横跨本地 CUDA 与远程服务器(见 device.rs 的文档示例)。- 由于
Metal、Vulkan、WebGpu本质是 wgpu 配上特定着色器编译器,三者都会枚举 wgpu 设备;push_cube辅助函数会去重,避免同一硬件被重复列出。 Remote变体比较特殊:它不是按后端类型 id 枚举,而是在运行时连接服务器查询它暴露了多少设备。- 枚举结果
Devices支持批量操作:.autodiff()为所有设备开启 autodiff、.configure(config)批量配置默认 dtype(见 device.rs),并且解引用为&[Device],可直接迭代。
执行栈:Tensor → Bridge → Dispatch → Backend
在底层,一次操作会流经Tensor → Bridge → Dispatch → Backend四层栈:
- Tensor:为应用提供稳定、与后端无关的 API,即你日常调用的
Tensor::ones、tensor + tensor等; - Bridge:将张量句柄转换为 dispatch 值(bridge 操作持有
DispatchDevice,仅需将其包装为统一的Device); - Dispatch:根据张量所属的设备选择对应的实现——
Dispatch::sync、Dispatch::seed、Dispatch::enumerate等设备级操作同样由它分派(见 crates/burn-dispatch/src/backend.rs); - Backend:执行原语操作。
Backend与AutodiffBackendtrait 仍然定义低层实现契约,但普通应用代码不需要B: Backend之类的约束。只有实现或扩展后端时才需要接触这些 trait,相关指引见 后端扩展(Backend Extension)。这种"设备即上下文"的设计,配合Device的透明类型擦除(device.rs),正是 Burn 能够在不牺牲灵活性的前提下同时支持多种硬件、且让用户代码保持简洁的关键。
总结
Burn 的Device是一个贯穿全生命周期的运行时概念:从Device::cuda(0)、Device::wgpu(DeviceKind::DiscreteGpu(0))这类按 feature 启用的构造器,到autodiff()/without_autodiff()的幂等上下文切换,再到configure()的一次性 dtype 锁定与enumerate()的硬件发现,设备 API 覆盖了训练、推理、多设备与远程执行的全部场景。掌握这张设备全景图,你就能用同一套代码在笔记本的 CPU(Device::flex())、实验室的 CUDA 集群(Device::cuda(0))与浏览器的 WebGPU(Device::webgpu(Default::default()))之间自由切换,并准确控制 autodiff、随机种子与内存行为。
【免费下载链接】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),仅供参考