Burn 深度学习框架后端与设备(Backend Device)完全指南:从设备选择、数据迁移到自动微分上下文

发布时间:2026/9/14 19:09:29
Burn 深度学习框架后端与设备(Backend  Device)完全指南:从设备选择、数据迁移到自动微分上下文 Burn 深度学习框架后端与设备Backend 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/burnBurn 是一个将运行时设备Device作为一等公民的深度学习框架Tensor、Module 与 Device 三者构成用户侧的核心 API任何张量都携带一个标识其执行位置与方式的运行时Device从而让同一套模型与张量代码可以在 CPU、CUDA、ROCm、WGPUVulkan/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 使用第一块独立 GPUDevice::vulkan(Default::default())最佳 Vulkan 适配器Device::metal(Default::default())最佳 Metal 适配器Device::webgpu(Default::default())浏览器中的 WebGPU 设备Device::cuda(0)索引为 0 的 CUDA GPUDevice::cuda(DeviceIndex::Default)由后端挑选默认 CUDA GPUDevice::rocm(0)索引为 0 的 ROCm/HIP GPUDevice::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()需要flexfeatureDevice::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 IntoDeviceIndex而usize/u32/u64/i32/i64都实现了From转换其中负数会 panic见 device.rs所以最常见的写法就是直接传整数Device::cuda(0)等价于Device::cuda(DeviceIndex::Specified(0))而Device::cuda(DeviceIndex::Default)则把选择权交给后端。DeviceKindWGPU 家族的适配器选择DeviceKind服务于 WGPU 这类设备句柄是带标签枚举的灵活后端其变体与 CubeCL 的WgpuDevice一一对应但作为 Burn 自有枚举避免调用方直接依赖 cubecl见 device.rsDiscreteGpu(usize)系统中第 N 块独立 GPUIntegratedGpu(usize)系统中第 N 块集成 GPUVirtualGpu(usize)系统中第 N 块虚拟 GPUCpuCPU 适配器DefaultDevice默认值由 wgpu 的启发式规则挑选最佳可用设备优先高功耗 GPU并可通过环境变量CUBECL_WGPU_DEFAULT_DEVICE覆盖例如CUBECL_WGPU_DEFAULT_DEVICEIntegratedGpu(1)或CUBECL_WGPU_DEFAULT_DEVICECpuExisting(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()可启用带平衡Balancedcheckpointing 策略的 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外部传入上下文与GenericBurn 内部错误附带惰性解析的 backtrace两个变体。flush()提交排队的工作但不等待完成。其源码文档device.rs指出fusion 后端会缓存算子以构建优化、远程后端会批量攒积后经网络发送flush()强制立即排空这些队列而即时执行eager后端没有缓冲此操作是空操作。is_autodiff()报告设备是否关联了 autodiff。这只是新张量继承的设备上下文不代表任何张量参与计算图或保留梯度。gradient_checkpointing_strategy()返回当前生效的策略无 autodiff 时返回None。supports_dtype(dtype)报告设备是否支持该数据类型存储、转换与算术均支持。文档特别提醒某些类型可能可存储可转换但无算术支持——例如 Vulkan 设备上的 bf16SPIR-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未指定的项会在设备初始化时解析为设备特定默认值。它还实现了FromFloatDType、FromIntDType、FromBoolDType、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变体完全取决于应用启用的后端 featuredevice.rsCpu、Cuda、Rocm、Wgpu、Metal、Vulkan、WebGpu、Flex、NdArray、LibTorch以及携带服务器地址的Remote(String)。几个值得注意的细节DeviceType变体之间可以通过|运算符组合成DeviceFilter也支持从单个DeviceType或VecDeviceType转换例如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仅需将其包装为统一的DeviceDispatch根据张量所属的设备选择对应的实现——Dispatch::sync、Dispatch::seed、Dispatch::enumerate等设备级操作同样由它分派见 crates/burn-dispatch/src/backend.rsBackend执行原语操作。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 覆盖了训练、推理、多设备与远程执行的全部场景。掌握这张设备全景图你就能用同一套代码在笔记本的 CPUDevice::flex()、实验室的 CUDA 集群Device::cuda(0)与浏览器的 WebGPUDevice::webgpu(Default::default())之间自由切换并准确控制 autodiff、随机种子与内存行为。【免费下载链接】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),仅供参考