Burn IR 中间表示:burn-ir 如何为张量计算提供跨后端抽象与可缓存的图执行基础
【免费下载链接】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-ir是 Burn 深度学习框架中负责定义"中间表示"(Intermediate Representation, IR)的 crate:它以纯数据形式描述张量(tensor)与张量操作(operation),使同一份计算描述可以跨不同执行目标(如远程后端)执行,并在真正执行前被优化和变换(例如算子融合)。读完本文,你将理解 IR 的四层数据结构(标量、张量、操作、图)、操作图边界的自动推导规则、BackendIr后端集成契约,以及HandleContainer在句柄生命周期、错误传播和自动调优(autotune)中的作用。
1. crate 定位与依赖结构
burn-ir在 Cargo.toml 中的声明为description = "Intermediate representation for the Burn framework",是 Burn 多 crate 架构中独立的一环。从依赖声明可以看出其设计取向:
| 依赖 | 用途 |
|---|---|
burn-backend | 提供Shape、DType、Slice、各类 ops trait 与Distribution,IR 的张量/操作元数据全部建立在这些类型之上 |
burn-std | 提供Arc等标准库替代实现,使错误句柄可以跨无指针宽度原子(如 thumbv6m)的目标共享 |
serde | 所有 IR 类型派生Serialize/Deserialize,这是 IR 可被传输(远程后端)与持久化的前提 |
hashbrown | 选用 no_std 兼容的哈希集合,用于图边界推导 |
特性(features)方面,default = ["std"],std特性向下传递到burn-backend/std与burn-std/std;此外还有tracing(透传burn-backend/tracing)用于日志追踪。lib.rs 顶部的#![cfg_attr(not(feature = "std"), no_std)]表明该 crate 本身支持 no_std 环境,模块列表为:
mod backend; // BackendIr 后端扩展 trait、TensorHandle、HandleKind mod builder; // 各类操作的 create 构造器(含形状/dtype 推导) mod graph; // GraphIr、图边界推导、GraphId / GraphBindings mod handle; // HandleContainer、句柄状态与错误传播 mod operation; // OperationIr 全量操作枚举 mod scalar; // ScalarIr 标量表示 mod tensor; // TensorId / TensorIr / TensorStatus从源码结构看,burn-ir本身不包含任何 kernel 或执行逻辑——它是一个"编译器前端"式的描述层:描述计算的是什么,而在哪执行由实现BackendIr的后端(如burn-fusion、burn-remote对应的 router 后端)负责。
2. 标量与张量 IR:最小描述单元
2.1 ScalarIr
scalar.rs 定义了可序列化的标量字面量:
pub enum ScalarIr { Float(f64), Int(i64), UInt(u64), Bool(bool), }源码注释明确区分了它与burn_backend::Scalar的分工:Scalar是运行时的字面量值,而ScalarIr是用于 IR 的可序列化字面量表示。两者通过From实现双向无损转换。ScalarIr::new会根据目标DType选择构造分支(float/int/uint/bool),对不支持的 dtype(如 QFloat)会直接unimplemented!。
值得注意的是 graph.rs 中GraphBindings的文档:相对化图(relative graph)会把标量替换为ScalarIr::UInt(placeholder)占位符,重放(replay)时再从 bindings 中还原具体值——这意味着同一张缓存图可以在每次调用时携带不同的标量参数(如 learning rate、epsilon)。
2.2 TensorId、TensorStatus 与 TensorIr
tensor.rs 中的TensorId是一个包装u64的轻量标识符(Copy + Hash + Serialize)。IR 张量不携带数据,只有元数据:
pub struct TensorIr { pub id: TensorId, // 张量唯一标识 pub shape: Shape, // 形状 pub status: TensorStatus, // 该张量被使用时所处的状态 pub dtype: DType, // 数据类型 }TensorStatus有三种取值,这是理解 IR 内存管理的关键:
ReadOnly:只读引用;ReadWrite:原地(inplace)可变——表示这是该张量的最后一次使用;NotInit:句柄尚不存在(尚未被任何计算写入)。
TensorIr的文档注释给出了一个典型的生命周期序列:同一个张量 id 在一次计算流中会经历NotInit → ReadOnly → ReadOnly → ReadWrite。执行器据此判断何时可以回收缓冲:当某次使用标记为ReadWrite时,读走的句柄不会被放回容器(见 handle.rs 的get_handle实现)。
3. 操作 IR:类型安全的操作全集
operation.rs 是整个 crate 最大的文件(5000+ 行),按张量类型分层组织操作枚举:
pub enum OperationIr { BaseFloat(BaseOperationIr), // 任意 float 张量上的基础操作 BaseInt(BaseOperationIr), // 任意 int 张量 BaseBool(BaseOperationIr), // 任意 bool 张量 NumericFloat(DType, NumericOperationIr), // 数值型 float 操作(带具体 dtype) NumericInt(DType, NumericOperationIr), Bool(BoolOperationIr), Int(IntOperationIr), Float(DType, FloatOperationIr), Module(ModuleOperationIr), // 模块级操作 Init(InitOperationIr), // 初始化 Custom(CustomOpIr), // 用户自定义操作 Drop(TensorIr), // 显式释放 Distributed(DistributedOperationIr), // 分布式集合通信 Activation(ActivationOperationIr), // 激活函数 }各分层覆盖了完整算子面:
BaseOperationIr:reshape/expand、swap_dims、permute、slice/slice_assign、select、gather/scatter/gather_nd/scatter_nd、cat、cast、zeros/ones/empty、mask_where/mask_fill、reduce all/any 等跨 float/int/bool 通用的操作;NumericOperationIr:add/sub/mul/div 及其 scalar 变体、sum/mean/max/min/prod 及其 dim 变体、cumsum/cumprod、argmax/topk/sort、pad、clamp 等;FloatOperationIr:exp、log、tanh、matmul、quantize/dequantize、grid_sample_2d 等 float 专属操作;IntOperationIr:位运算(and/or/xor/not/移位)、matmul、cast;ModuleOperationIr:模块级高层操作——linear 及三个 backward、conv1d/2d/3d 及对应 backward、deformable conv、conv transpose、avg/max/adaptive pool 及 backward、interpolate、embedding 及 backward、layer_norm、batch_norm、rfft/irfft、attention、ctc_loss 等;DistributedOperationIr:AllReduce与SyncCollective(后者特意建模成操作而非旁路调用,以便与 AllReduce 在操作流中保持有序)。
3.1 CustomOpIr:面向融合流与远程后端的扩展点
pub struct CustomOpIr { pub id: String, // 操作唯一标识 pub inputs: Vec<TensorIr>, pub outputs: Vec<TensorIr>, pub scalars: Vec<ScalarIr>, // 非张量标量参数,按声明顺序 }其scalars字段的文档说明了一个重要的运行时链路:自定义操作的标量参数会像其他标量一样参与融合的"相对化"过程,因此缓存的图可以在重放时携带新的标量值;远程后端(remote backend)依赖这一点把标量参数发往服务器,由注册的处理函数通过ScalarIr::elem读回。
3.2 builder.rs:带形状与 dtype 推导的构造器
builder.rs 通过impl_ir_create!宏为每个操作生成create(..., new_id)构造器,其核心不变式是:构造 IR 时就完成输出张量的 shape 与 dtype 推导和校验,从而在"执行前"发现错误。例如:
BinaryOpIr:shape = lhs.shape.broadcast(&rhs.shape),dtype 要求两输入一致,否则返回IrError::DTypeMismatch;MatmulOpIr:输出形状由calculate_matmul_output计算,并支持 float 与 QFloat 混合输入的create_mixed;MaxPool2dOpIr等:直接复用burn_backend::ops::conv::calculate_pool_output_shape推导输出形状,与后端 kernel 的计算保持一致;PermuteOpIr/SwapDimsOpIr:通过permute_quantized_dtype/swap_dims_quantized_dtype同步调整量化存储方案(QuantStore::PackedU32/PackedNative)的打包轴,保证量化张量重排后元数据仍然自洽;PadOpIr:构造时调用burn_backend::ops::validate_padding校验 padding 合法性。
这种"IR 构造即推导"的设计意味着执行器在解释执行时不需要再做形状检查,也意味着任何形状非法的程序会在 IR 构建阶段(而非 kernel 运行时)暴露出来。
4. 图 IR:操作序列、边界推导与缓存重放
graph.rs 将操作序列提升为图:
pub struct GraphIr { pub operations: Vec<OperationIr>, // 按执行顺序排列 pub inputs: Vec<TensorId>, // 图输入(首用顺序) pub outputs: Vec<TensorId>, // 图输出(产生顺序) }4.1 边界自动推导:GraphIr::classify
GraphIr::new(operations)会调用classify从操作序列推导输入/输出边界,规则(来自源码文档注释):
- 输入= 被图引用但不由计算操作产生的张量:外部数据、前一张图的结果。特别地,
OperationIr::Init不视为"产生"——initializer 句柄是带外(out of band)注册的,因此它们永远算作图输入; - 输出= 由计算操作产生且幸存的张量:被
ReadWrite原地消费、或被显式Drop的张量被排除; - 中间张量(产生后被原地消费或显式丢弃)不出现在任何一边;
- 输入按首用顺序、输出按产生顺序排列,保证构造是确定性的。
graph.rs 内的单元测试 验证了这些规则,例如graph_boundary_is_in_first_use_order断言:两个自定义操作串联后,中间结果tensor(5)被Drop排除出输出,输入按[9, 3]的首用顺序排列。classify是借用版本的分类(不克隆操作序列),文档说明它可用于验证显式声明的逻辑边界或估算图的绑定成本。
4.2 图缓存与重放:GraphId/GraphBindings
这是 IR 面向"跨目标执行 + 优化"的核心机制:
pub struct GraphId(pub u64); // 缓存图的注册 id pub struct GraphBindings { pub tensors: Vec<(TensorId, TensorId)>, // 边界张量:(相对id, 具体id) pub shapes: Vec<usize>, // 相对形状维度 -> 具体值 pub scalars: Vec<ScalarIr>, // 标量占位符 -> 具体值 pub ranges: Vec<Slice>, // 切片范围占位符 -> 具体范围 }从文档注释可以还原出完整的工作流:
- 首次执行时,router 后端(如 remote backend)把一段操作图相对化(relativize):张量 id 变成位置化 id,形状维度变稠密占位符,标量/切片范围变成占位符,然后以
GraphId注册缓存; - 之后的每次调用只发送
GraphBindings——仅包含边界张量的绑定、稠密的形状维度表、标量值与切片范围。中间张量 id 由重放端自行分配,每个张量(含中间张量)的具体形状都能从shapes表重建; - 由于绑定载荷只含边界信息,无论图内有多少个操作,每次调用的通信/传输成本保持很小。
形状维度表用Vec而非 map 的原因在注释中写明:相对维度 id 是稠密的0..N,且每个不同的维度值只存一份(多个张量可共享同一维度 id),天然就是普通表。
5. BackendIr:把 IR 接回具体后端
backend.rs 定义了后端扩展契约:
pub trait BackendIr: Backend { type Handle: Sync + Send + Clone; fn float_tensor(handle: TensorHandle<Self::Handle>) -> FloatTensor<Self>; fn int_tensor(handle: TensorHandle<Self::Handle>) -> IntTensor<Self>; fn bool_tensor(handle: TensorHandle<Self::Handle>) -> BoolTensor<Self>; fn quantized_tensor(handle: TensorHandle<Self::Handle>) -> QuantizedTensor<Self>; fn float_tensor_handle(tensor: FloatTensor<Self>) -> Self::Handle; fn int_tensor_handle(tensor: IntTensor<Self>) -> Self::Handle; fn bool_tensor_handle(tensor: BoolTensor<Self>) -> Self::Handle; fn quantized_tensor_handle(tensor: QuantizedTensor<Self>) -> Self::Handle; }配套的两个类型:
TensorHandle<H>:{ handle: H, shape: Shape },带形状的张量资源引用;HandleKind<B>:区分Float / Int / Bool / Quantized四种句柄种类,提供name()便于调试与传输。
这个 trait 的文档说明其用途是"允许既有的 Backend 使用 Burn 张量 IR 做编译(compilation)或其他目的"——即:实现了BackendIr的后端可以把GraphIr当作编译输入(做算子融合、图缓存),而普通后端(CPU、ndarray 等直接解释执行)则不需要它。从源码结构看,burn-fusion(融合后端)与burn-router/burn-remote(路由与远程后端)是消费 IR 的主要下游。
6. HandleContainer:句柄注册表、内存回收与错误传播
handle.rs 的HandleContainer<H>是 IR 执行侧的状态核心——它"把所有张量句柄集中管理,确保所有资源被最优使用"。每个 id 对应的条目有三种状态:
pub enum Handle<H> { NotInit, // 句柄尚未创建 Existing(H), // 句柄已创建 Errored(TensorError), // 写入该张量的工作从未执行 }6.1 状态驱动的生命周期
get_handle(id, status)是最热的读路径:ReadOnly读走后把句柄克隆放回容器(张量仍存活),ReadWrite读走即销毁(原地操作的最后一次使用),NotInit则 panic;free(tensor):仅对ReadWrite状态的张量释放条目,与状态机一一对应;register_*_tensor/register_handle:把新产生的张量登记进容器;fork():浅克隆整个容器,文档注明用途是 autotune——测试fork_is_isolated_from_original明确记录了"fork 中注册的输出句柄不会出现在原容器中"这一隔离语义。
6.2 零成本失败路径与错误传播
HandleContainer维护一个errored计数,所有错误检查都先经过has_errors()分支——在没有任何失败("几乎总是"的情况)时,错误检查只付出一次分支判断。TensorError的设计:
- 根因是
Arc<ExecutionError>:同一失败传播到下游的每个张量都共享同一错误对象(保留原始类型与 backtrace,而非压扁成消息),same_root用指针相等判断"是否同一失败"; depth记录该张量距离失败点被跳过了多少个操作;- 错误随张量自身的
Drop/最后一次读释放,因此错误集合天然被"仍然存活的张量"限定; - 关键恢复语义:
register_handle会顶掉Errored条目——源码注释说明这正是 autotune 依赖的性质:"失败的候选声明了它没写出的输出;成功的候选写入并清除声明,执行继续,下游无人被跳过"。claim_unwritten则只声明"无人写入且无人声明"的 id,避免用模糊声明覆盖精确声明。
handle.rs 内的 15 个单元测试 系统性覆盖了 fork 隔离、错误声明/恢复/传播、跨容器(跨设备)迁移后根因保持、计数漂移防护等场景,是理解这套错误模型的最好材料。
7. 端到端视角:IR 在 Burn 运行时中的位置
把上述模块串起来,burn-ir支撑的是这样一条链路(基于源码文档与结构推断):
- 描述:上层(融合 router 等)把一次前向/反向计算记录为
Vec<OperationIr>,构造GraphIr并自动得到输入/输出边界; - 抽象执行:实现
BackendIr的执行器按顺序解释操作,用HandleContainer管理句柄的注册、读取与释放,TensorStatus驱动缓冲回收; - 跨目标传输:因为
OperationIr、GraphIr、GraphBindings全部可 serde 序列化,操作流/图可以发往远程后端执行; - 优化与缓存:图被相对化后按
GraphId缓存,之后每次调用只传边界绑定,载荷大小与图内操作数解耦;标量与切片范围作为占位符参与,使同一缓存图支持逐次变化的参数; - 失败处理:kernel 编译/启动失败时,受影响张量被声明为
Errored,下游读取按"未写入"报错并共享根因,autotune 重试成功时声明被写入覆盖清除。
这也解释了 README 中两句话的落点:"抽象出计算允许跨不同目标执行(如 remote backend)"对应BackendIr+ serde 可序列化 +GraphBindings;"允许在执行前对张量计算做优化和变换(如算子融合)"对应GraphIr作为可分析、可变换的纯数据操作序列,以及CustomOpIr为融合流保留的扩展接口。
8. 适用边界与使用建议
burn-ir是框架内部 crate,普通 Burn 用户写模型/训练代码时不会直接 import 它;它面向的是后端/执行器开发者(例如实现融合策略、远程执行、自定义 router 后端的人);- 从 Cargo.toml 可见其版本、edition 均继承 workspace 配置(
crates/burn-ir/Cargo.toml中version.workspace = true),构建时以仓库根 workspace 为准,需要 Rust 工具链满足 workspace 的 MSRV; - 若要在自己的后端中使用 IR,最小集成路径是:让后端
Backend实现基础上补实现BackendIr(定义Handle类型及双向转换),执行器持有一个HandleContainer<Self::Handle>,按GraphIr.operations顺序解释OperationIr,并在读/写句柄时遵循TensorStatus语义; - 注意
TensorStatus的契约注释(handle.rs):传入的 status 必须与实际操作匹配,否则可能提前移除后续仍需要的句柄——这是使用该 crate 时最需要小心的不变式。
相关源码入口:模块总览、操作全集、图与缓存绑定、构造器与推导、后端契约、句柄与错误模型。
【免费下载链接】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),仅供参考