Burn 深度学习框架后端与设备(Backend Device)完全指南:从设备选择、数据迁移到自动微分上下文
2026/9/14 19:09:19 网站建设 项目流程

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 中,TensorModuleDevice并不针对某个后端做泛型化(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之类的泛型约束即可运行在任何后端之上。BackendAutodiffBackendtrait 仍然定义了底层的实现契约,但普通应用代码几乎不会直接接触它们——只有当你实现或扩展后端时才需要(参见 后端扩展指南)。

从源码结构看,这一设计在 crates/burn-tensor/src/device.rs 中得到印证:Device是一个薄封装结构,内部通过burn_std::obfuscate!宏将具体的DispatchDevice类型擦除为对齐的透明存储(device_opaque::Opaque),这样下游代码既不需要知道 dispatch 的类型树,又能通过as_dispatch()/into_dispatch()在需要时访问底层设备。

需要特别注意的是:每个设备构造器都对应一个 Cargo feature。例如启用wgpucudafeature 后,Device::wgpuDevice::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::cudaDevice::rocmDevice::metalDevice::vulkanDevice::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::wgpuDevice::vulkanDevice::metalDevice::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_deviceModule::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_autodiffgradient_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_allocationsmemory_install_poolsmemory_pool_reportmemory_pool_usage等,用于精细控制动态内存池。

这些设备方法最终都汇聚到 dispatch 层——crates/burn-dispatch/src/backend.rs 中Dispatchseedsyncad_enabledmemory_cleanupsupports_dtypeflush等接口定义了后端必须实现或可覆盖的行为契约。

设备设置:默认数据类型与一次性锁定

每个设备都有自己的运行时设置,包括默认的浮点、整数与布尔数据类型。通过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_dtypeint_dtypebool_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):

  1. 手动初始化:可在程序开头通过configure/set_default_dtypes设置一次;
  2. 默认初始化:若在手动初始化之前发生了任何操作(比如创建张量),设置将被永久锁定为默认值;
  3. 不可变性:一旦初始化,设置不可更改,以保证跨线程、跨操作的行为一致。

因此:务必在该设备上的第一次张量操作之前完成配置。初始化之后,默认 dtype 即被锁定,任何后续不兼容的配置都会返回DeviceError::AlreadyInitializedDeviceError定义于 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):CpuCudaRocmWgpuMetalVulkanWebGpuFlexNdArrayLibTorch以及携带服务器地址的Remote(String)。几个值得注意的细节:

  • DeviceType变体之间可以通过|运算符组合成DeviceFilter(也支持从单个DeviceTypeVec<DeviceType>转换),例如DeviceType::Cuda | DeviceType::remote_websocket("ws://host:3000")一次调用即可横跨本地 CUDA 与远程服务器(见 device.rs 的文档示例)。
  • 由于MetalVulkanWebGpu本质是 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::onestensor + tensor等;
  • Bridge:将张量句柄转换为 dispatch 值(bridge 操作持有DispatchDevice,仅需将其包装为统一的Device);
  • Dispatch:根据张量所属的设备选择对应的实现——Dispatch::syncDispatch::seedDispatch::enumerate等设备级操作同样由它分派(见 crates/burn-dispatch/src/backend.rs);
  • Backend:执行原语操作。

BackendAutodiffBackendtrait 仍然定义低层实现契约,但普通应用代码不需要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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询