在 Burn 中新增一个张量算子(Operation)的完整实践指南
在 Burn 中新增一个张量算子Operation的完整实践指南【免费下载链接】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 开源仓库的贡献者指南与真实源码以powf算子为例完整讲解如何在 Burn 中从零新增一个张量算子从burn-tensor的用户侧 API 与后端 trait 定义到burn-autodiff的自动微分实现含偏导数推导与梯度检查点分类再到burn-fusion、burn-ir、burn-cubecl等 JIT/GPU 后端的分发与代码生成最后落到burn-backend-tests的测试套件。读完本文你将掌握 Burn 算子在 API 层 → trait 层 → 各后端实现 → 测试 全链路中的每个落点与约定并能在自己的分支上为 Burn 提交一个完整、可测试、可自动微分的新算子。说明本文以 contributor-book/src/guides/adding-a-new-operation-to-burn.md 为骨架所有路径与行号均依据当前仓库实际布局核对与文档写作时引用的 commit如9f31281、0ee2021相比部分文件位置与行号已随版本演进发生变化但分层结构与命名约定保持稳定。一、总体路线图一个新算子要经过哪些层Burn 的算子不是在一个地方一次性实现完的而是沿着一条从用户可见 API到具体硬件计算的链路逐层落子。以powf为例完整路径如下burn-tensorAPI 层在Tensor结构体上暴露powf方法并在FloatTensorOps/IntTensorOps/QTensorOpstrait 中声明底层操作burn-autodiff自动微分层为算子实现反向传播backward pass计算左右偏导数并声明其计算/内存绑定属性目标后端如burn-flex、burn-ndarray真正执行浮点/整数计算的实现burn-fusionburn-irburn-cubecl融合/JIT/GPU 层将算子注册进融合中间表示IR并映射为 GPU 可执行指令WGSL/CPP/SPIR-Vburn-backend-tests测试层为浮点、整数、量化张量分别补充测试用例。其中向burn-tensor添加算子与向burn-autodiff添加反向传播是核心两步后者难度最高。二、第一步在 burn-tensor 中添加算子2.1 找到正确的 traitnumeric 与类型专属 traitburn-tensor是所有后端必须实现的张量操作定义处核心位于 crates/burn-backend/src/backend/ops/tensor.rsFloatTensorOps与 crates/burn-backend/src/backend/ops/int_tensor.rsIntTensorOps。这两个 trait 由不同的burn-*后端实现如果 trait 提供了默认实现后端可以只覆盖需要优化的部分。关键约定如下跨类型共享的算子Int与Float张量都需要的数值类算子如pow、add同时放入两个 trait并以类型前缀命名。例如powf在FloatTensorOps中命名为float_powf在IntTensorOps中命名为int_powi。类型专属算子只对某一种张量有意义的算子如sin、cos对Int/Bool无意义直接在对应的 API 文件中实现。当前仓库中sin位于 crates/burn-tensor/src/tensor/api/float.rs。注意文档写作时曾提到burn-tensor/src/tensor/api/numeric.rs中存在Numerictrait 与其对应的 int/float 实现在当前仓库中数值算子的公共 API 已分别收敛到 crates/burn-tensor/src/tensor/api/float.rs 与 crates/burn-tensor/src/tensor/api/int.rspowi相关代码见 numeric.rs 第 720-750 行。分层职责不变Tensor是只有一个primitive字段的结构体定义于 crates/burn-tensor/src/tensor/api/base.rs由KindBool/Float/Int之一定义于 crates/burn-tensor/src/tensor/kind.rs约束API 方法最终会分发给对应类型的后端 op。2.2 用户侧 API 与底层实现分离以powf为例用户侧 API 与底层实现是分离的两层公开 APIcrates/burn-tensor/src/tensor/api/float.rspub fn powf(self, other: Self) - Self { check!(TensorCheck::binary_ops_ew(Powf, self, other)); Tensor::new(powf_impl(self.primitive, other.primitive)) } pub fn powf_scalarE: ElementConversion(self, other: E) - Self { let rhs Scalar::new(other, self.dtype()); Tensor::new(powf_scalar_impl(self.primitive, rhs)) }方法在调用底层实现前先做形状/广播检查TensorCheck::binary_ops_ew失败时 panicpowf_scalar接受任意实现了ElementConversion的标量类型统一转换为Scalar后分发。底层分发实现crates/burn-tensor/src/tensor/api/float.rs#L1248-L1287是体现量化支持的关键fn powf_impl(lhs: BridgeTensor, rhs: BridgeTensor) - BridgeTensor { let (lkind, lhs) lhs.into_parts(); let (rkind, rhs) rhs.into_parts(); match (lkind, rkind) { (BridgeKind::Float, BridgeKind::Float) { BridgeTensor::float(Dispatch::float_powf(lhs, rhs)) } (BridgeKind::QFloat, BridgeKind::QFloat) match Dispatch::q_powf(lhs, rhs) { TensorPrimitive::Float(out) BridgeTensor::float(out), TensorPrimitive::QFloat(out) BridgeTensor::qfloat(out), }, (BridgeKind::QFloat, BridgeKind::Float) { let dtype rhs.dtype(); BridgeTensor::float(Dispatch::float_powf( Dispatch::dequantize(lhs, dtype.into()), rhs, )) } (BridgeKind::Float, BridgeKind::QFloat) { let dtype lhs.dtype(); BridgeTensor::float(Dispatch::float_powf( lhs, Dispatch::dequantize(rhs, dtype.into()), )) } _ panic!(Should be Float primitive kind), } }这段代码印证了文档中的核心设计Float 张量的 primitive 由TensorPrimitive枚举表示见 crates/burn-tensor/src/tensor/api/kind.rs可以携带Float或QFloat量化浮点两种变体powf_impl根据左右操作数的类型组合正确分发到浮点 opfloat_powf或量化 opq_powf混合场景则先dequantize再走浮点路径。2.3 Int 版本用浮点实现加两次 cast对于Int张量int_powi的实现方式是复用浮点实现并在两侧各做一次类型转换见 crates/burn-backend/src/backend/ops/int_tensor.rs#L510-L521fn int_powi(lhs: IntTensorB, rhs: IntTensorB) - IntTensorB { let dtype lhs.dtype(); let float_dtype get_device_settings::B(lhs.device()).float_dtype; B::float_into_int( B::float_powi(B::int_into_float(lhs, float_dtype), rhs), dtype.into(), ) }即Int→Float使用设备设置指定的float_dtype→ 执行浮点幂运算 → 结果转回Int。因此后续其他后端的实现只需要聚焦浮点版本。2.4 量化QTensor约定q_* 前缀与默认实现引入量化浮点张量后量化算子遵循同样的约定量化 op 以q_为前缀例如q_powf、q_reshape对应浮点的float_*定义于 crates/burn-backend/src/backend/ops/qtensor.rs。大多数量化算子带有默认实现逻辑为反量化输入 → 在浮点张量上执行运算 → 量化输出。以q_powf为例qtensor.rs 第 745 行附近fn q_powf(lhs: QuantizedTensorB, rhs: QuantizedTensorB) - TensorPrimitiveB { // 默认实现dequantize 两个输入 - float_powf - quantize 输出 // 后端可以在需要时覆盖此实现 }后端只有在需要针对性优化如直接在量化域计算、减少精度损失时才覆盖这些默认实现。三、第二步在 burn-autodiff 中添加反向传播burn-autodiff是让其他后端获得自动微分能力的包装层实现于 crates/burn-autodiff/src/ops/tensor.rs这是整个流程中最不直观的一步。以float_powf的 autodiff 实现第 4015-4088 行为例需要依次完成以下工作3.1 五个必要步骤定义 backward 单元结构体实现一个反向传播函数。对于powf这类逐元素二元运算使用binary辅助函数来自同目录backward.rs最后两个参数是两个闭包分别定义左、右偏导数fn float_powf(lhs: FloatTensorSelf, rhs: FloatTensorSelf) - FloatTensorSelf { #[derive(Debug)] struct PowF; retro_binary!(RetroPowf, B::float_powf); implB: Backend BackwardB, 2 for PowF { type State (NodeId, NodeId, BinaryOpsBroadcast); fn backward( self, ops: OpsSelf::State, 2, grads: mut Gradients, checkpointer: mut Checkpointer, ) { let (lhs_id, rhs_id, broadcast) ops.state; let lhs: B::FloatTensorPrimitive checkpointer.retrieve_node_output(lhs_id); let rhs: B::FloatTensorPrimitive checkpointer.retrieve_node_output(rhs_id); // lhs 与 rhs 分别被左右两侧偏导使用按 parents 规格复制所需份数 let [rhs_4lhs, rhs_4rhs] duplicate(ops.parents, Some(rhs)); let [lhs_4lhs, lhs_4rhs] duplicate(ops.parents, Some(lhs)); binary::B, _, _( ops.parents, ops.node, grads, |grad| { // rhs*(lhs.val**(rhs-1))*grad let rhs1 rhs_4lhs.unwrap(); let rhs2 rhs1.clone(); let lhs lhs_4lhs.unwrap(); let tmp B::float_powf(lhs, B::float_sub_scalar(rhs1, 1.0.into())); let value B::float_mul(tmp, rhs2); let grad B::float_mul(grad, value); broadcast.backward_lhs::B(grad) }, |grad| { // lhs**rhs * ln(lhs) * grad let rhs rhs_4rhs.unwrap(); let lhs1 lhs_4rhs.unwrap(); let lhs2 lhs1.clone(); let tmp B::float_powf(lhs1, rhs); let value B::float_mul(tmp, B::float_log(lhs2)); let grad B::float_mul(grad, value); broadcast.backward_rhs::B(grad) }, ); } } // ...见下方 prepare 流程 }定义 tracked / untracked 行为tracked时算子进入 autodiff 图并注册 backward 执行untracked时直接以普通方式调用底层函数let broadcast BinaryOpsBroadcast::new::B(lhs.primitive, rhs.primitive); match PowF .prepare::C([lhs.node.clone(), rhs.node.clone()]) .memory_bound() .retro_forward(RetroPowf::B::new(lhs.node.id, rhs.node.id)) .parents([lhs, rhs]) .stateful() { OpsKind::Tracked(mut prep) { let lhs_state prep.checkpoint(lhs); let rhs_state prep.checkpoint(rhs); prep.finish( (lhs_state, rhs_state, broadcast), B::float_powf(lhs.primitive, rhs.primitive), ) } OpsKind::UnTracked(prep) prep.finish(B::float_powf(lhs.primitive, rhs.primitive)), }状态保存策略被跟踪时算子必须保存足够信息以在反向时高效计算。信息轻量如 shape直接存入 state反向需要输入值的应使用 checkpoint 而不是直接保存——prep.checkpoint(lhs)会按检查点策略在反向时惰性提供输入。powf的反向同时需要lhs与rhs所以两个输入都被 checkpoint。区分 compute-bound 与 memory-bound梯度检查点分类API 定义于 crates/burn-autodiff/src/ops/base.rs.compute_bound()计算密集型算子如 matmul、卷积即使开启 checkpoint 也会保存输出反向时不重算.memory_bound()轻量逐元素算子如powf每个元素只做一次小运算反向时用父节点输出重算前向更省内存因此不保存整个前向输出。被注册为 memory-bound 的算子必须实现两件事.parents()声明其父节点以及提供一个实现RetroForward的结构体在反向时用父节点输出重算前向。powf正是通过retro_binary!(RetroPowf, B::float_powf)宏生成RetroPowf并调用B::float_powf完成重算tensor.rs 第 4019 行。处理广播BinaryOpsBroadcast记录了前向时左右操作数的广播关系反向时分别通过broadcast.backward_lhs/backward_rhs将梯度还原为原始形状。实操建议文档原话的印证上述步骤大多是样板代码最省力的做法是复制一个结构类似的已有算子改掉结构体名确保两侧都能拿到所需数据需要对方张量的副本时就clone其内容。3.2 偏导数速查以 pow 为例如果你不熟悉偏导数计算这里是pow算子必要的微积分基础。pow是二元算子左、右闭包分别是关于左、右张量的偏导数。定义算子为函数 \(f(x,y)x^{y}\)其中 \(x\) 是左张量、\(y\) 是右张量计算偏导时把另一个变量视为常数对 \(x\) 的偏导左闭包\(\frac{\partial}{\partial x}(x^{y}) y \cdot x^{y-1}\)对 \(y\) 的偏导右闭包\(\frac{\partial}{\partial y}(x^{y}) x^{y} \cdot \ln(x)\)对照上面 autodiff 代码中的注释左闭包实现rhs*(lhs**(rhs-1))*grad右闭包实现lhs**rhs * ln(lhs) * grad与公式一一对应。3.3 autodiff 测试autodiff 算子的测试覆盖在burn-backend-tests的 autodiff 测试目录中见 crates/burn-backend-tests/tests/autodiff/具体测试方法参考 contributor-book/src/getting-started/testing.md 的说明。四、第三步在其他后端实现算子4.1 普通后端burn-flex、burn-ndarray对目标后端而言算子实现通常直截了当——这里就是真正发生计算的地方。以burn-flex的powf浮点实现为例crates/burn-flex/src/ops/float.rs#L924-L931fn float_powf(lhs: FloatTensorFlex, rhs: FloatTensorFlex) - FloatTensorFlex { binary_op(lhs, rhs, |a: f32, b| a.powf(b), |a: f64, b| a.powf(b), None) } fn float_powf_scalar_impl(tensor: FloatTensorFlex, value: Scalar) - FloatTensorFlex { let exp value.elem::f64() as f32; // 转换为标量 scalar_op(tensor, exp, |a: f32, b| a.powf(b), |a: f64, b| a.powf(b)) }通过binary_op/scalar_op辅助函数分别给出f32与f64两个闭包即可。同样地burn-ndarray的实现位于 crates/burn-ndarray/src/ops/tensor.rs。版本说明文档写作时burn-tch已被标记为弃用并将于未来版本移除因此新算子不需要实现 LibTorch 后端当前仓库中 crates/burn-tch 的powf实现仅存在于历史代码中。4.2 融合与 JIT 后端burn-fusion / burn-irburn-fusion与 JIT 后端并非目标后端而是为其他后端提供内核融合与即时编译能力的包装层。添加算子不涉及任何计算只需描述生成代码长什么样大部分可以从现有函数复制粘贴调整。文档以powf加入burn-fusion为例给出 4 个落点当前仓库对应如下融合浮点 ops在 crates/burn-fusion/src/ops/tensor.rs当前float_powf位于第 2586 行附近注册powfIR 枚举在 crates/burn-ir/src/operation.rs 的FloatOperationIr枚举中新增Powf(BinaryOpIr)当前第 214 行与PowfScalar(ScalarOpIr)当前第 156 行IR 枚举的 trait 实现在同一文件中为FloatOperationIr的各 trait如输入/输出收集、融合判定补齐Powf/PowfScalar分支当前第 3275-3596 行附近有大量FloatOperationIr::Powf...的 match 分支融合 stream context在 crates/burn-fusion/src/stream/context.rs 的匹配中处理新算子。从源码看burn-ir中每个融合算子都需要为算子收集输入节点inputs、输出节点outputs以及描述如何由输入构造输出construct等方法提供分支这正是描述生成代码长什么样的落点。4.3 GPU 后端burn-cubecl / cubeclcubecl负责算子在 GPU 上的编译与执行burn-cubecl则是 Burn 对接 cubecl 的实现层。cubecl 处理 tensor-scalar 运算的方式是把两者都变换为一系列向量化标量运算因此powf的 tensor 版本可以复用已有实现。当前仓库的落点FloatTensorOps 实现float_powf与float_powf_scalar_impl位于 crates/burn-cubecl/src/ops/tensor.rs第 590、815 行附近其中标量版本通过Vector::powf生成向量化代码numeric 辅助实际的计算辅助函数位于 crates/burn-cubecl/src/ops/numeric.rs融合代码生成在 crates/burn-cubecl-fusion/src/engine/codegen/ir.rs 中定义 GPU 融合内核的打印形式当前第 217 行FuseOp::Powf(args) write!(f, {} powf({}, {}), ...)。对于需要复杂内核、无法直接映射为底层指令的函数直接使用cube宏编写自定义内核即可。4.4 WGSL 特殊处理案例文档特别提到Burn 团队为powf生成了自定义 WGSL 代码原因在于 WGSL 原生pow函数的边界行为问题——例如0^0应为 1负数取偶次幂应为正数。实现策略是尽可能复用现有逻辑在最后一步根据 rhs 的操作数类型var type分支处理。这个细节提醒我们跨后端实现时不仅要写对数学公式还要注意各目标语言WGSL、CPP、SPIR-V对边界语义的差异必要时需要定制指令映射。五、第四步补充测试burn-backend-tests新算子必须配套测试测试统一放在burn-backend-testscrate 中因为该 crate 会被所有后端共用。5.1 浮点与整数算子测试测试文件放在crates/burn-backend-tests/tests/tensor/{float|int}/ops/{op_name}.rs模块名登记进对应的crates/burn-backend-tests/tests/tensor/{float|int}/ops/mod.rs。以powf为例crates/burn-backend-tests/tests/tensor/float/ops/powf.rs并在 mod.rs 第 109-110 行 登记了mod powf; mod powf_scalar;#[test] fn should_support_powf_ops() { let data TensorData::from([[1.0, 1.0, 2.0], [3.0, 4.0, 5.0]]); let tensor TestTensor::2::from_data(data, Default::default()); let pow TensorData::from([[1.0, 1.0, 2.0], [3.0, 4.0, 2.0]]); let tensor_pow TestTensor::2::from_data(pow, Default::default()); let output tensor.powf(tensor_pow); let expected TensorData::from([[1.0, 1.0, 4.0], [27.0, 256.0, 25.0]]); output .into_data() .assert_approx_eq::FloatElem(expected, Tolerance::default()); }该文件中还包含负指数should_support_neg_power、负数底数偶次幂、形状不兼容 panicshould_panic_powf_incompatible_shapes、广播should_support_powf_broadcasted等边界测试powf_scalar.rs则覆盖标量版本。5.2 量化算子测试q_*如果浮点算子支持量化通常会同步添加QTensorOps对应实现带默认实现测试流程类似测试放在crates/burn-backend-tests/tests/tensor/float/quantization/ops/extended/{op_name}.rs当前仓库powf.rs、powf_scalar.rs均存在模块名登记进同目录的mod.rs注意extended测试套件仅对支持原生非打包量化后端的后端启用见 crates/burn-backend-tests/tests/tensor/float/quantization/ops/mod.rs 顶部注释。量化测试的关键约定输入与期望输出一律用浮点值定义测试对输出调用.dequantize()后再与期望比对。虽然这隐含假设了量化/反量化本身是正确的但让测试可读性大增——本质上这些测试是验证张量运算对量化保持不变性在量化误差范围内。例如 powf.rs 中let tensor QTensor::2::int8([[1.0, 1.0, 2.0], [3.0, 4.0, 5.0]]); let tensor_pow QTensor::2::int8([[1.0, 1.0, 2.0], [3.0, 4.0, 2.0]]); let output tensor.powf(tensor_pow); let expected TensorData::from([[1.0, 1.0, 4.0], [27.0, 256.0, 25.0]]); output .dequantize() .into_data() .assert_approx_eq::FloatElem(expected, Tolerance::rel_abs(4e-2, 1e-2));注意文档原话的补充测试尽量选用可以无失真或失真很小地完成量化/反量化的浮点值但结果始终取决于算子本身——例如张量乘积会使数值范围显著放大导致结果的反/量化误差增大。因此量化测试的容差需要按算子特性谨慎设定上例使用相对容差4e-2、绝对容差1e-2。六、总结完整的提交清单按照本文流程为一个新算子提交完整实现时逐项检查burn-tensor API 层在 crates/burn-tensor/src/tensor/api/ 暴露公开方法类型共享算子按{type}_op_name前缀在 crates/burn-backend/src/backend/ops/tensor.rsFloat、int_tensor.rsInt声明支持量化则在 qtensor.rs 添加q_*默认实现autodiff 层在 crates/burn-autodiff/src/ops/tensor.rs 实现 backward左右偏导闭包、checkpoint 状态、compute_bound/memory_bound分类、RetroForward重算目标后端在 burn-flex、burn-ndarray 等实现计算融合/JIT/GPU 层在 burn-ir 注册 IR 分支、burn-fusion 融合 ops、burn-cubecl 与 burn-cubecl-fusion 生成 GPU 代码测试浮点/整数测试进crates/burn-backend-tests/tests/tensor/{float|int}/ops/量化测试进crates/burn-backend-tests/tests/tensor/float/quantization/ops/extended/并同步登记进各自的mod.rs。对照这份清单你就能完整地给 Burn 新增一个算子。更多架构背景可参考 contributor-book/src/project-architecture/tensor.mdTensor 分层架构与 contributor-book/src/guides/README.md贡献指南总览。【免费下载链接】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),仅供参考