From 75989509a35cabe63dc86d8313c9a6e712383167 Mon Sep 17 00:00:00 2001 From: georgebisbas Date: Tue, 15 Sep 2026 17:01:01 +0200 Subject: [PATCH] refactor(ir): extract LoweringBuilder from lower_composite_ops_pass MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit lower_composite_ops_pass.cpp has grown to 2845 lines and carries every composite collective lowering rule in one translation unit, so each new collective lands hundreds of lines in the same file and any edit risks every other rule. Plan 70 pays this down by splitting per-collective translation units without touching pass order (the reorder alternative was rejected in #1850). This is phase 1: move the shared LoweringBuilder scratchpad, its CommSetup result struct and the MakeNegation helper into a lower_composite/ module so the per-collective rules extracted next have a shared home to depend on. Pure code motion — the rules still resolve LoweringBuilder and CommSetup through using declarations, so no call site changes and no IR output changes. Mirrors public issue #2632, which proposes the same split. Also updates the pass's own architecture doc (en/zh), which still described a single-translation-unit layout and told contributors that "all edits stay in lower_composite_ops_pass.cpp" — now stale since CMakeLists.txt compiles a second TU and LoweringBuilder/CommSetup live there. --- CMakeLists.txt | 1 + docs/en/dev/passes/13-lower_composite_ops.md | 15 +- docs/zh/dev/passes/13-lower_composite_ops.md | 12 +- .../lower_composite_builder.cpp | 387 ++++++++++++++ .../lower_composite/lower_composite_builder.h | 264 ++++++++++ .../transforms/lower_composite_ops_pass.cpp | 496 +----------------- 6 files changed, 676 insertions(+), 499 deletions(-) create mode 100644 src/ir/transforms/lower_composite/lower_composite_builder.cpp create mode 100644 src/ir/transforms/lower_composite/lower_composite_builder.h diff --git a/CMakeLists.txt b/CMakeLists.txt index 23da24e010..5e05765aca 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -220,6 +220,7 @@ set(PYPTO_SOURCES src/ir/transforms/loop_invariant_mat_residency.cpp src/ir/transforms/inline_functions_pass.cpp src/ir/transforms/init_memref.cpp + src/ir/transforms/lower_composite/lower_composite_builder.cpp src/ir/transforms/lower_composite_ops_pass.cpp src/ir/transforms/materialize_tensor_strides_pass.cpp src/ir/transforms/ir_property.cpp diff --git a/docs/en/dev/passes/13-lower_composite_ops.md b/docs/en/dev/passes/13-lower_composite_ops.md index fe20939e00..bf8f2095aa 100644 --- a/docs/en/dev/passes/13-lower_composite_ops.md +++ b/docs/en/dev/passes/13-lower_composite_ops.md @@ -28,15 +28,22 @@ The empty `PassProperties` contract (`kLowerCompositeOpsProperties` in `include/ ## Architecture -The pass is a single translation unit, `src/ir/transforms/lower_composite_ops_pass.cpp`: +The pass spans two translation units: `LoweringBuilder` and its helpers were extracted into +`src/ir/transforms/lower_composite/lower_composite_builder.{h,cpp}` (plan 70) so the shared +scratchpad/control-flow machinery isn't tied to one giant file; the rule table and mutator stay in +`src/ir/transforms/lower_composite_ops_pass.cpp`, which `#include`s the new header: ```text -src/ir/transforms/lower_composite_ops_pass.cpp +src/ir/transforms/lower_composite/lower_composite_builder.h / .cpp + CommSetup — per-collective comm-domain/signal setup bundle LoweringBuilder — per-call scratchpad (Bind + primitive tile-op builders: tile.muls, tile.adds, tile.add, tile.sub, tile.mul, tile.maximum, tile.minimum, tile.cast + structured control-flow: EmitFor / EmitForReduce / EmitIf / EmitIfExpr + NotEq scalar guard) + MakeNegation — file-local scalar/tile negation helper used by builder rules + +src/ir/transforms/lower_composite_ops_pass.cpp CompositeLoweringFn — (call, visited_args, builder) -> result expr LowerRule — one rule function per composite op (LowerSinRule, LowerCosRule, LowerTensorAllReduceRule, ...) @@ -44,7 +51,9 @@ src/ir/transforms/lower_composite_ops_pass.cpp LowerCompositeOpsMutator — walks the function, looks up a rule per Call ``` -Adding a new single-result composite op (all edits stay in `lower_composite_ops_pass.cpp`): +Adding a new single-result composite op (rule + dispatch-table edits stay in +`lower_composite_ops_pass.cpp`; only touch `lower_composite_builder.{h,cpp}` if the rule needs a +new builder primitive): 1. Write a `LowerRule(call, args, builder)` function. It receives the original `CallPtr` (use `call->span_`, `call->kwargs_`, `call->op_->name_` as needed), the visited arg expressions (var-remap already applied), and a `LoweringBuilder` whose `Bind` helper appends an `AssignStmt` per intermediate temp. For rules that need control flow, use `builder.EmitFor` / `builder.EmitForReduce` / `builder.EmitIf` / `builder.EmitIfExpr` — each takes a body callback that receives a nested builder sharing the same temp counter, so emitted temps stay uniquely named regardless of nesting depth. `LowerTensorAllReduceRule` is the canonical example of a control-flow-bearing rule (ready barrier plus chunked remote_load+accumulate / barrier / store for mesh; `LowerTensorRingAllReduceRule` adds a chunked RS+AG ring schedule dispatched via a `mode` kwarg). 2. Add a `{"", &LowerRule}` row to `kRules` inside `LookupCompositeRule`. diff --git a/docs/zh/dev/passes/13-lower_composite_ops.md b/docs/zh/dev/passes/13-lower_composite_ops.md index a3c1d2d01d..f446976987 100644 --- a/docs/zh/dev/passes/13-lower_composite_ops.md +++ b/docs/zh/dev/passes/13-lower_composite_ops.md @@ -28,15 +28,21 @@ host-orchestrator 中的 `pld.tensor.allreduce` 调用会跳过本 Pass:`Synth ## 架构 (Architecture) -本 Pass 是单个翻译单元 (translation unit),即 `src/ir/transforms/lower_composite_ops_pass.cpp`: +本 Pass 现在跨两个翻译单元:`LoweringBuilder` 及其辅助设施已抽取到 +`src/ir/transforms/lower_composite/lower_composite_builder.{h,cpp}`(plan 70),使共享的暂存区/控制流机制不再绑定在一个巨型文件里;规则表和 mutator 仍留在 +`src/ir/transforms/lower_composite_ops_pass.cpp` 中,该文件 `#include` 新头文件: ```text -src/ir/transforms/lower_composite_ops_pass.cpp +src/ir/transforms/lower_composite/lower_composite_builder.h / .cpp + CommSetup — 每次集合通信调用的 comm-domain/signal 配置组合 LoweringBuilder — 单次调用的暂存区 (Bind + 基本 tile 算子构造器: tile.muls、tile.adds、tile.add、tile.sub、tile.mul、 tile.maximum、tile.minimum、tile.cast + 结构化控制流:EmitFor / EmitForReduce / EmitIf / EmitIfExpr + NotEq 标量比较) + MakeNegation — builder 规则使用的文件内标量/tile 取负辅助函数 + +src/ir/transforms/lower_composite_ops_pass.cpp CompositeLoweringFn — (call, visited_args, builder) -> 结果表达式 LowerRule — 每个组合算子一个规则函数(LowerSinRule、 LowerCosRule、LowerTensorAllReduceRule ...) @@ -44,7 +50,7 @@ src/ir/transforms/lower_composite_ops_pass.cpp LowerCompositeOpsMutator — 遍历函数,对每个 Call 查表 ``` -新增一个单结果组合算子的步骤(改动都留在 `lower_composite_ops_pass.cpp` 内): +新增一个单结果组合算子的步骤(规则与分发表改动留在 `lower_composite_ops_pass.cpp` 内;只有当规则需要新的 builder 基本能力时才需要改动 `lower_composite_builder.{h,cpp}`): 1. 写一个 `LowerRule(call, args, builder)` 函数。它接收原始 `CallPtr`(按需用 `call->span_`、`call->kwargs_`、`call->op_->name_`)、已 visit 过的参数表达式(已应用 var-remap)以及一个 `LoweringBuilder`,其 `Bind` 助手会为每个中间临时变量追加一条 `AssignStmt`。需要控制流的规则可以用 `builder.EmitFor` / `builder.EmitForReduce` / `builder.EmitIf` / `builder.EmitIfExpr`——每个都接收一个 body 回调,回调里收到的嵌套 builder 与外层共享同一个 temp 计数器,因此发射的临时变量名跨任意嵌套深度都唯一。`LowerTensorAllReduceRule` 是含控制流规则的范例(mesh 使用 ready 屏障,加分块 remote_load+accumulate / 屏障 / store;`LowerTensorRingAllReduceRule` 则通过 `mode` kwarg 分发,增加分块 RS+AG ring 调度)。 2. 在 `LookupCompositeRule` 的 `kRules` 里加一条 `{"", &LowerRule}`。 diff --git a/src/ir/transforms/lower_composite/lower_composite_builder.cpp b/src/ir/transforms/lower_composite/lower_composite_builder.cpp new file mode 100644 index 0000000000..5fa659fc26 --- /dev/null +++ b/src/ir/transforms/lower_composite/lower_composite_builder.cpp @@ -0,0 +1,387 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ + +#include "src/ir/transforms/lower_composite/lower_composite_builder.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "pypto/core/dtype.h" +#include "pypto/core/logging.h" +#include "pypto/ir/comm.h" +#include "pypto/ir/expr.h" +#include "pypto/ir/kind_traits.h" +#include "pypto/ir/op_registry.h" +#include "pypto/ir/scalar_expr.h" +#include "pypto/ir/span.h" +#include "pypto/ir/stmt.h" +#include "pypto/ir/transforms/utils/auto_name_utils.h" +#include "pypto/ir/transforms/utils/tile_conversion_utils.h" +#include "pypto/ir/type.h" + +namespace pypto { +namespace ir { +namespace lower_composite { + +namespace { + +/// Safe negation: folds ConstInt(-value) directly when possible so the +/// PyPTO printer->parser roundtrip (which folds ``Neg(ConstInt)`` into +/// ``ConstInt(-value)``) produces structurally equal IR. Runtime +/// expressions are still wrapped with ``Neg``. File-local: the only callers +/// are EmitEpilogueReset overloads below. +ExprPtr MakeNegation(const ExprPtr& value) { + if (auto c = As(value)) { + return std::make_shared(-c->value_, GetScalarDtype(value), value->span_); + } + return MakeNeg(value, value->span_); +} + +} // namespace + +LoweringBuilder::LoweringBuilder(std::string base_name, std::size_t& temp_counter) + : base_name_(std::move(base_name)), temp_counter_(temp_counter) {} + +LoweringBuilder::LoweringBuilder(std::string base_name, std::size_t& temp_counter, bool nested) + : base_name_(std::move(base_name)), temp_counter_(temp_counter), nested_(nested) {} + +ExprPtr LoweringBuilder::Bind(const std::string& qualifier, const ExprPtr& expr, const Span& span) { + auto var = std::make_shared(MakeTempName(qualifier), expr->GetType(), span); + stmts_.push_back(std::make_shared(var, expr, span)); + return var; +} + +void LoweringBuilder::EmitEval(const ExprPtr& expr, const Span& span) { + stmts_.push_back(std::make_shared(expr, span)); +} + +ExprPtr LoweringBuilder::Muls(const ExprPtr& x, float c, const Span& span) { + auto tile_type = As(x->GetType()); + INTERNAL_CHECK_SPAN(tile_type, span) << "tile.muls input must be TileType"; + auto scalar = std::make_shared(static_cast(c), tile_type->dtype_, span); + return OpRegistry::GetInstance().Create("tile.muls", {x, scalar}, {}, span); +} + +ExprPtr LoweringBuilder::Adds(const ExprPtr& x, float c, const Span& span) { + auto tile_type = As(x->GetType()); + INTERNAL_CHECK_SPAN(tile_type, span) << "tile.adds input must be TileType"; + auto scalar = std::make_shared(static_cast(c), tile_type->dtype_, span); + return OpRegistry::GetInstance().Create("tile.adds", {x, scalar}, {}, span); +} + +ExprPtr LoweringBuilder::Add(const ExprPtr& a, const ExprPtr& b, const Span& span) { + return OpRegistry::GetInstance().Create("tile.add", {a, b}, {}, span); +} + +ExprPtr LoweringBuilder::Sub(const ExprPtr& a, const ExprPtr& b, const Span& span) { + return OpRegistry::GetInstance().Create("tile.sub", {a, b}, {}, span); +} + +ExprPtr LoweringBuilder::Mul(const ExprPtr& a, const ExprPtr& b, const Span& span) { + return OpRegistry::GetInstance().Create("tile.mul", {a, b}, {}, span); +} + +ExprPtr LoweringBuilder::Reduce(ReduceOp op, const ExprPtr& a, const ExprPtr& b, const Span& span) { + const char* op_name; + switch (op) { + case ReduceOp::kSum: + op_name = "tile.add"; + break; + case ReduceOp::kMax: + op_name = "tile.maximum"; + break; + case ReduceOp::kMin: + op_name = "tile.minimum"; + break; + case ReduceOp::kProd: + op_name = "tile.mul"; + break; + default: + INTERNAL_CHECK_SPAN(false, span) + << "pld.tensor.allreduce lowering received unknown ReduceOp " << static_cast(op); + } + return OpRegistry::GetInstance().Create(op_name, {a, b}, {}, span); +} + +ExprPtr LoweringBuilder::Cast(const ExprPtr& x, DataType to, int mode, const Span& span) { + std::vector> kw = {{"target_type", to}, {"mode", mode}}; + return OpRegistry::GetInstance().Create("tile.cast", {x}, kw, span); +} + +ExprPtr LoweringBuilder::NotEq(const ExprPtr& left, const ExprPtr& right, const Span& span) { + return MakeNe(left, right, span); +} + +ExprPtr LoweringBuilder::Gt(const ExprPtr& left, const ExprPtr& right, const Span& span) { + return MakeGt(left, right, span); +} + +CommSetup LoweringBuilder::EmitCommSetup(const ExprPtr& comm_target, const Span& span) { + auto& reg = OpRegistry::GetInstance(); + CommSetup s; + s.ctx = Bind("ctx", reg.Create("pld.system.get_comm_ctx", {comm_target}, {}, span), span); + s.nranks_i32 = Bind("nranks", reg.Create("pld.system.nranks", {s.ctx}, {}, span), span); + s.nranks_idx = Bind("nranks_idx", std::make_shared(s.nranks_i32, DataType::INDEX, span), span); + s.my_rank = Bind("my_rank", reg.Create("pld.system.rank", {s.ctx}, {}, span), span); + return s; +} + +void LoweringBuilder::EmitNotifyAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, + NotifyOp notify_op, const ExprPtr& value, const std::string& suffix, + const Span& span) { + auto zero_idx = std::make_shared(0, DataType::INDEX, span); + auto one_idx = std::make_shared(1, DataType::INDEX, span); + auto my_offsets = tile_conversion_utils::MakeSignalOffsets(my_rank, span); + + EmitFor( + "peer" + suffix, zero_idx, nranks_idx, one_idx, + [&](LoweringBuilder& body, const VarPtr& peer) { + body.EmitIf( + body.NotEq(peer, my_rank, span), + [&](LoweringBuilder& then_body) { + auto call = + OpRegistry::GetInstance().Create("pld.system.notify", {signal, peer, my_offsets, value}, + {{"op", static_cast(notify_op)}}, span); + then_body.Bind("notify" + suffix + "_ret", call, span); + }, + /*else_fn=*/nullptr, span); + }, + span); +} + +void LoweringBuilder::EmitNotifyAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, + const ExprPtr& row_offset, NotifyOp notify_op, const ExprPtr& value, + const std::string& suffix, const Span& span) { + auto zero_idx = std::make_shared(0, DataType::INDEX, span); + auto one_idx = std::make_shared(1, DataType::INDEX, span); + auto my_offsets = tile_conversion_utils::MakeSignalOffsets(my_rank, row_offset, span); + + EmitFor( + "peer" + suffix, zero_idx, nranks_idx, one_idx, + [&](LoweringBuilder& body, const VarPtr& peer) { + body.EmitIf( + body.NotEq(peer, my_rank, span), + [&](LoweringBuilder& then_body) { + auto call = + OpRegistry::GetInstance().Create("pld.system.notify", {signal, peer, my_offsets, value}, + {{"op", static_cast(notify_op)}}, span); + then_body.Bind("notify" + suffix + "_ret", call, span); + }, + /*else_fn=*/nullptr, span); + }, + span); +} + +void LoweringBuilder::EmitWaitAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, + const ExprPtr& expected, const std::string& suffix, const Span& span) { + auto zero_idx = std::make_shared(0, DataType::INDEX, span); + auto one_idx = std::make_shared(1, DataType::INDEX, span); + + EmitFor( + "src" + suffix, zero_idx, nranks_idx, one_idx, + [&](LoweringBuilder& body, const VarPtr& src) { + auto src_offsets = tile_conversion_utils::MakeSignalOffsets(src, span); + body.EmitIf( + body.NotEq(src, my_rank, span), + [&](LoweringBuilder& then_body) { + auto call = OpRegistry::GetInstance().Create("pld.system.wait", {signal, src_offsets, expected}, + {{"cmp", static_cast(WaitCmp::kGe)}}, span); + then_body.Bind("wait" + suffix + "_ret", call, span); + }, + /*else_fn=*/nullptr, span); + }, + span); +} + +void LoweringBuilder::EmitWaitAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, + const ExprPtr& row_offset, const ExprPtr& expected, + const std::string& suffix, const Span& span) { + auto zero_idx = std::make_shared(0, DataType::INDEX, span); + auto one_idx = std::make_shared(1, DataType::INDEX, span); + + EmitFor( + "src" + suffix, zero_idx, nranks_idx, one_idx, + [&](LoweringBuilder& body, const VarPtr& src) { + auto src_offsets = tile_conversion_utils::MakeSignalOffsets(src, row_offset, span); + body.EmitIf( + body.NotEq(src, my_rank, span), + [&](LoweringBuilder& then_body) { + auto call = OpRegistry::GetInstance().Create("pld.system.wait", {signal, src_offsets, expected}, + {{"cmp", static_cast(WaitCmp::kGe)}}, span); + then_body.Bind("wait" + suffix + "_ret", call, span); + }, + /*else_fn=*/nullptr, span); + }, + span); +} + +int64_t LoweringBuilder::EmitBarrier(const ExprPtr& signal, const CommSetup& comm, const std::string& suffix, + const Span& span) { + INTERNAL_CHECK_SPAN(!nested_, span) + << "Internal error: EmitBarrier must only be called from a top-level lowering rule, not from inside " + << "EmitFor / EmitIf / EmitIfExpr bodies. Loop- or condition-resident barriers must " + << "emit notify/wait by hand with a call-local expected value."; + const int64_t generation = ++barrier_count_; + auto one_i32 = std::make_shared(1, DataType::INT32, span); + auto expected_i32 = std::make_shared(generation, DataType::INT32, span); + EmitNotifyAll(signal, comm.nranks_idx, comm.my_rank, NotifyOp::kAtomicAdd, one_i32, suffix, span); + EmitWaitAll(signal, comm.nranks_idx, comm.my_rank, expected_i32, suffix, span); + return generation; +} + +void LoweringBuilder::EmitEpilogueReset(const ExprPtr& signal, const CommSetup& comm, const ExprPtr& total, + const Span& span) { + INTERNAL_CHECK_SPAN(!nested_, span) + << "EmitEpilogueReset must only be called from a top-level lowering rule, exactly once " + "per call, after every EmitBarrier / hand-rolled notify-wait pair the rule issues."; + auto neg_total = MakeNegation(total); + auto zero_idx = std::make_shared(0, DataType::INDEX, span); + auto one_idx = std::make_shared(1, DataType::INDEX, span); + EmitFor( + "reset_src", zero_idx, comm.nranks_idx, one_idx, + [&](LoweringBuilder& body, const VarPtr& src) { + body.EmitIf( + body.NotEq(src, comm.my_rank, span), + [&](LoweringBuilder& then_body) { + auto src_offsets = tile_conversion_utils::MakeSignalOffsets(src, span); + auto call = OpRegistry::GetInstance().Create( + "pld.system.notify", {signal, comm.my_rank, src_offsets, neg_total}, + {{"op", static_cast(NotifyOp::kAtomicAdd)}}, span); + then_body.Bind("epilogue_reset_ret", call, span); + }, + /*else_fn=*/nullptr, span); + }, + span); +} + +void LoweringBuilder::EmitEpilogueReset(const ExprPtr& signal, const CommSetup& comm, const ExprPtr& num_rows, + const ExprPtr& total_per_row, const Span& span) { + INTERNAL_CHECK_SPAN(!nested_, span) + << "EmitEpilogueReset must only be called from a top-level lowering rule, exactly once per call."; + auto neg_total = MakeNegation(total_per_row); + auto zero_idx = std::make_shared(0, DataType::INDEX, span); + auto one_idx = std::make_shared(1, DataType::INDEX, span); + EmitFor( + "reset_row", zero_idx, num_rows, one_idx, + [&](LoweringBuilder& row_body, const VarPtr& row) { + row_body.EmitFor( + "reset_src", zero_idx, comm.nranks_idx, one_idx, + [&](LoweringBuilder& body, const VarPtr& src) { + body.EmitIf( + body.NotEq(src, comm.my_rank, span), + [&](LoweringBuilder& then_body) { + auto src_offsets = tile_conversion_utils::MakeSignalOffsets(src, row, span); + auto call = OpRegistry::GetInstance().Create( + "pld.system.notify", {signal, comm.my_rank, src_offsets, neg_total}, + {{"op", static_cast(NotifyOp::kAtomicAdd)}}, span); + then_body.Bind("epilogue_reset_ret", call, span); + }, + /*else_fn=*/nullptr, span); + }, + span); + }, + span); +} + +void LoweringBuilder::EmitFor(const std::string& loop_var_name, const ExprPtr& start, const ExprPtr& stop, + const ExprPtr& step, + const std::function& body_fn, + const Span& span) { + auto loop_var = std::make_shared(MakeTempName(loop_var_name), start->GetType(), span); + LoweringBuilder body_builder(base_name_, temp_counter_, /*nested=*/true); + body_fn(body_builder, loop_var); + auto body_stmt = WrapBodyStmts(body_builder.TakeStmts(), span); + stmts_.push_back(std::make_shared(loop_var, start, stop, step, std::vector{}, + body_stmt, std::vector{}, span)); +} + +ExprPtr LoweringBuilder::EmitForReduce( + const std::string& loop_var_name, const ExprPtr& start, const ExprPtr& stop, const ExprPtr& step, + const ExprPtr& init_value, + const std::function& body_fn, const Span& span) { + auto loop_var = std::make_shared(MakeTempName(loop_var_name), start->GetType(), span); + auto iter_arg = std::make_shared(MakeTempName(loop_var_name + "_acc"), init_value->GetType(), + init_value, span); + LoweringBuilder body_builder(base_name_, temp_counter_, /*nested=*/true); + ExprPtr yield_val = body_fn(body_builder, loop_var, iter_arg); + INTERNAL_CHECK_SPAN(yield_val, span) + << "EmitForReduce body_fn must return the next iteration's accumulator value"; + body_builder.stmts_.push_back(std::make_shared(std::vector{yield_val}, span)); + auto body_stmt = WrapBodyStmts(body_builder.TakeStmts(), span); + auto return_var = + std::make_shared(MakeTempName(loop_var_name + "_final"), init_value->GetType(), span); + stmts_.push_back(std::make_shared(loop_var, start, stop, step, std::vector{iter_arg}, + body_stmt, std::vector{return_var}, span)); + return return_var; +} + +void LoweringBuilder::EmitIf(const ExprPtr& cond, const std::function& then_fn, + const std::function& else_fn, const Span& span) { + LoweringBuilder then_builder(base_name_, temp_counter_, /*nested=*/true); + then_fn(then_builder); + auto then_body = WrapBodyStmts(then_builder.TakeStmts(), span); + + std::optional else_body = std::nullopt; + if (else_fn) { + LoweringBuilder else_builder(base_name_, temp_counter_, /*nested=*/true); + else_fn(else_builder); + else_body = WrapBodyStmts(else_builder.TakeStmts(), span); + } + stmts_.push_back(std::make_shared(cond, then_body, else_body, std::vector{}, span)); +} + +ExprPtr LoweringBuilder::EmitIfExpr(const ExprPtr& cond, + const std::function& then_fn, + const std::function& else_fn, + const Span& span) { + INTERNAL_CHECK_SPAN(then_fn && else_fn, span) + << "EmitIfExpr requires both then_fn and else_fn (the if must yield a value on every path)"; + LoweringBuilder then_builder(base_name_, temp_counter_, /*nested=*/true); + ExprPtr then_val = then_fn(then_builder); + INTERNAL_CHECK_SPAN(then_val, span) << "EmitIfExpr then_fn must return the yielded value"; + then_builder.stmts_.push_back(std::make_shared(std::vector{then_val}, span)); + auto then_body = WrapBodyStmts(then_builder.TakeStmts(), span); + + LoweringBuilder else_builder(base_name_, temp_counter_, /*nested=*/true); + ExprPtr else_val = else_fn(else_builder); + INTERNAL_CHECK_SPAN(else_val, span) << "EmitIfExpr else_fn must return the yielded value"; + else_builder.stmts_.push_back(std::make_shared(std::vector{else_val}, span)); + auto else_body = WrapBodyStmts(else_builder.TakeStmts(), span); + + auto return_var = std::make_shared(MakeTempName("if_res"), then_val->GetType(), span); + stmts_.push_back(std::make_shared(cond, then_body, std::optional(else_body), + std::vector{return_var}, span)); + return return_var; +} + +std::vector LoweringBuilder::TakeStmts() { return std::move(stmts_); } + +std::string LoweringBuilder::MakeTempName(const std::string& qualifier) { + return auto_name::BuildName(auto_name::GetBaseName(base_name_), qualifier, "tmp", + static_cast(temp_counter_++)); +} + +StmtPtr LoweringBuilder::WrapBodyStmts(std::vector body_stmts, const Span& span) { + if (body_stmts.empty()) return std::make_shared(std::vector{}, span); + if (body_stmts.size() == 1) return body_stmts.front(); + return std::make_shared(std::move(body_stmts), span); +} + +} // namespace lower_composite +} // namespace ir +} // namespace pypto diff --git a/src/ir/transforms/lower_composite/lower_composite_builder.h b/src/ir/transforms/lower_composite/lower_composite_builder.h new file mode 100644 index 0000000000..e707ed1f20 --- /dev/null +++ b/src/ir/transforms/lower_composite/lower_composite_builder.h @@ -0,0 +1,264 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ + +#ifndef SRC_IR_TRANSFORMS_LOWER_COMPOSITE_LOWER_COMPOSITE_BUILDER_H_ +#define SRC_IR_TRANSFORMS_LOWER_COMPOSITE_LOWER_COMPOSITE_BUILDER_H_ + +#include +#include +#include +#include +#include + +#include "pypto/core/dtype.h" +#include "pypto/ir/comm.h" +#include "pypto/ir/expr.h" +#include "pypto/ir/span.h" +#include "pypto/ir/stmt.h" + +namespace pypto { +namespace ir { +namespace lower_composite { + +// ============================================================================ +// CommSetup — result struct for LoweringBuilder::EmitCommSetup() +// ============================================================================ + +/// Holds bound expressions from the comm-setup preamble (ctx, nranks, my_rank). +/// Returned by LoweringBuilder::EmitCommSetup() for use in subsequent phases. +struct CommSetup { + ExprPtr ctx; ///< Result of pld.system.get_comm_ctx + ExprPtr nranks_i32; ///< Result of pld.system.nranks (INT32) + ExprPtr nranks_idx; ///< nranks cast to INDEX (for loop bounds) + ExprPtr my_rank; ///< Result of pld.system.rank (INT32) +}; + +// ============================================================================ +// LoweringBuilder +// +// Per-call scratchpad handed to a composite-lowering rule. A rule appends one +// ``AssignStmt`` per intermediate temp via ``Bind`` and returns the final +// result ``ExprPtr``; the mutator wraps that result in the original target +// ``Var`` (or a fresh result ``Var`` for ``ReturnStmt`` calls) before splicing +// the accumulated statements into the surrounding sequence. +// +// In addition to ``Bind`` and the primitive op builders, the builder exposes +// structured control-flow constructors — ``EmitFor`` / ``EmitForReduce`` / +// ``EmitIf`` / ``EmitIfExpr`` — that hand the body off to a nested builder +// callback. The nested builder shares this builder's temp counter so every +// emitted temp gets a unique name across the entire rule, regardless of +// nesting depth. +// +// The temp counter is borrowed from the mutator so unique temp names span +// distinct composite-op calls in the same function. Barrier generations are +// call-local (see the self-clearing credit-barrier protocol in +// lower_composite_ops_pass.cpp), so each LoweringBuilder instance — one per +// top-level composite-op call — owns its own ``barrier_count_`` that always +// starts at 0. +// ============================================================================ +class LoweringBuilder { + public: + /// @param base_name Name hint to derive temp names from (typically the + /// AssignStmt's LHS ``Var`` name). + /// @param temp_counter Reference to a mutator-owned counter; bumped per Bind. + LoweringBuilder(std::string base_name, std::size_t& temp_counter); + + LoweringBuilder(std::string base_name, std::size_t& temp_counter, bool nested); + + /// Append an ``AssignStmt`` binding a fresh ``Var`` to ``expr`` and return + /// the new ``Var`` so it can be used as input to subsequent ops. The + /// ``qualifier`` is woven into the temp name for debuggability. + ExprPtr Bind(const std::string& qualifier, const ExprPtr& expr, const Span& span); + + /// Append a side-effecting expression without manufacturing an unused SSA + /// result. This is used by destination-passing ops whose output buffers are + /// explicit operands. + void EmitEval(const ExprPtr& expr, const Span& span); + + // Primitive op builders -- type deduction is delegated to OpRegistry so the + // result preserves the input TileType's shape/layout/dtype. + ExprPtr Muls(const ExprPtr& x, float c, const Span& span); + ExprPtr Adds(const ExprPtr& x, float c, const Span& span); + ExprPtr Add(const ExprPtr& a, const ExprPtr& b, const Span& span); + ExprPtr Sub(const ExprPtr& a, const ExprPtr& b, const Span& span); + ExprPtr Mul(const ExprPtr& a, const ExprPtr& b, const Span& span); + ExprPtr Reduce(ReduceOp op, const ExprPtr& a, const ExprPtr& b, const Span& span); + ExprPtr Cast(const ExprPtr& x, DataType to, int mode, const Span& span); + + // ---- Scalar comparison helpers (yield BOOL-typed expressions, suitable as + // IfStmt conditions or loop guards). Delegated to the scalar_expr + // Make* helpers so operand promotion stays consistent with parser + // output. + ExprPtr NotEq(const ExprPtr& left, const ExprPtr& right, const Span& span); + + ExprPtr Gt(const ExprPtr& left, const ExprPtr& right, const Span& span); + + // ---- Collective-op helpers (DRY extraction for barrier/broadcast/allgather/ + // reduce_scatter/allreduce) ---- + + /// Emit comm-setup preamble: get_comm_ctx, nranks, rank. + /// Returns a CommSetup struct with the bound expressions for use in + /// subsequent phases. + CommSetup EmitCommSetup(const ExprPtr& comm_target, const Span& span); + + /// Emit notify-all loop: for peer in 0..nranks: if peer != my_rank: notify(...) + /// @param signal The signal DistributedTensor + /// @param nranks_idx Loop bound (INDEX-typed) + /// @param my_rank This rank's ID (INT32) + /// @param notify_op NotifyOp::kSet or NotifyOp::kAtomicAdd + /// @param value Value to notify (e.g., one_i32) + /// @param suffix Suffix for loop variable names (e.g., "" or "2" for re-notify) + /// @param span Source span for error reporting + void EmitNotifyAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, + NotifyOp notify_op, const ExprPtr& value, const std::string& suffix, const Span& span); + + /// Overload for 2D signal matrices (e.g. ring allreduce [2*(NR-1), NR]). + /// @param row_offset Row index expression for the 2D signal (e.g. ring step var) + void EmitNotifyAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, + const ExprPtr& row_offset, NotifyOp notify_op, const ExprPtr& value, + const std::string& suffix, const Span& span); + + /// Emit wait-all loop: for src in 0..nranks: if src != my_rank: wait(...) + /// @param signal The signal DistributedTensor + /// @param nranks_idx Loop bound (INDEX-typed) + /// @param my_rank This rank's ID (INT32) + /// @param expected Expected signal value — the barrier generation (INT32) + /// @param suffix Suffix for loop variable names (e.g., "" or "2" for re-wait) + /// @param span Source span for error reporting + void EmitWaitAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, + const ExprPtr& expected, const std::string& suffix, const Span& span); + + /// Overload for 2D signal matrices (e.g. ring allreduce [2*(NR-1), NR]). + /// @param row_offset Row index expression for the 2D signal (e.g. ring step var) + void EmitWaitAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, + const ExprPtr& row_offset, const ExprPtr& expected, const std::string& suffix, + const Span& span); + + // ---- Self-clearing credit barrier protocol (see lower_composite_ops_pass.cpp's file-header comment) ---- + + /// Emit one complete cross-rank barrier on ``signal``: ``AtomicAdd(1)`` into + /// every peer's cell, then wait for this call's generation on every peer's + /// cell. Returns the generation waited for (1-based, scoped to *this call* + /// only — every fresh ``LoweringBuilder`` starts counting at 0), so a rule + /// that fans out further barriers (the mesh allreduce's per-chunk barriers) + /// can continue the sequence from it, and so the rule can compute the total + /// credit count its ``EmitEpilogueReset`` call must subtract. + /// + /// Call this only from a rule's straight-line code — one call consumes exactly + /// one generation, so invoking it inside an ``EmitFor`` body would reserve a + /// single generation for a barrier that executes many times. Loop-resident + /// barriers must emit notify/wait by hand with a call-local expected value + /// (see the ring / mesh-chunked rules below). + int64_t EmitBarrier(const ExprPtr& signal, const CommSetup& comm, const std::string& suffix, + const Span& span); + + /// Self-clearing epilogue: subtract ``total`` from every non-self peer's + /// contribution to *my own* cells, restoring the signal to all-zero once + /// every rank has run its own epilogue. ``total`` is the number of + /// ``AtomicAdd(+1)`` notifies this call issued per peer (the sum of every + /// ``EmitBarrier`` / hand-rolled notify-wait pair the rule emitted) — it may + /// be a runtime-computed expression, not just a ``ConstInt``: + /// ``pld.system.notify``'s value only requires ``ScalarType``. + /// + /// This is a self-notify (``peer == my_rank``): the codegen path resolves + /// ``peer == my_rank`` via the same identity mapping ``pld.tile.put`` / + /// ``pld.tile.get`` already rely on for their self-rank case, so this lands + /// on the exact same hardware atomic as an incoming remote add. + /// + /// Call this exactly once per rule invocation, from top-level code only, + /// after every barrier the rule issues. + void EmitEpilogueReset(const ExprPtr& signal, const CommSetup& comm, const ExprPtr& total, + const Span& span); + + /// 2D-signal overload (ring allreduce, ``[2*(NR-1), NR]``): subtract + /// ``total_per_row`` from every non-self cell of every one of ``num_rows`` + /// rows. Ring credits every row independently (one per round / sub-chunk + /// sequence), and every row's sub-chunk loop shares the same bound, so one + /// symbolic ``total_per_row`` resets all rows uniformly. + void EmitEpilogueReset(const ExprPtr& signal, const CommSetup& comm, const ExprPtr& num_rows, + const ExprPtr& total_per_row, const Span& span); + + // ---- Structured control-flow constructors ---- + // + // Each method takes a body callback that receives a freshly-constructed + // nested ``LoweringBuilder`` scoped to the body region. The callback emits + // its body via the nested builder; this builder then drains the nested + // stmts, wraps them in a ``SeqStmts`` (when there is more than one), and + // emits the resulting ``ForStmt`` / ``IfStmt`` against its own ``stmts_``. + // + // The nested builder shares this builder's ``temp_counter_`` reference so + // emitted temp names stay unique across the entire rule regardless of + // nesting depth. + + /// Emit a side-effect-only ``for`` loop: + /// + /// for loop_var in range(start, stop, step): + /// + /// + /// ``body_fn`` receives a fresh body builder and the freshly-created loop + /// variable. The callback's return value is discarded — use this overload + /// for loops whose only purpose is side effects (e.g. issuing notify / + /// wait sequences). + void EmitFor(const std::string& loop_var_name, const ExprPtr& start, const ExprPtr& stop, + const ExprPtr& step, const std::function& body_fn, + const Span& span); + + /// Emit a reducing ``for`` loop with one loop-carried accumulator. The + /// body callback receives a nested builder, the loop variable, and the + /// accumulator (typed via ``init_value``); it returns the next iteration's + /// accumulator value. The method returns an expression holding the + /// post-loop accumulator, ready to feed into subsequent ops. + ExprPtr EmitForReduce(const std::string& loop_var_name, const ExprPtr& start, const ExprPtr& stop, + const ExprPtr& step, const ExprPtr& init_value, + const std::function& body_fn, + const Span& span); + + /// Emit a side-effect-only ``if`` statement: + /// + /// if cond: + /// + /// [else: + /// ] + /// + /// Pass ``nullptr`` for ``else_fn`` when there is no else branch. + void EmitIf(const ExprPtr& cond, const std::function& then_fn, + const std::function& else_fn, const Span& span); + + /// Emit a value-producing ``if`` statement. Both branches must yield a + /// value (via their body_fn's ExprPtr return); the method returns an + /// expression holding the chosen value, ready to feed into subsequent ops. + ExprPtr EmitIfExpr(const ExprPtr& cond, const std::function& then_fn, + const std::function& else_fn, const Span& span); + + /// Drain accumulated statements (called by the mutator after the rule + /// returns). + std::vector TakeStmts(); + + private: + std::string MakeTempName(const std::string& qualifier); + + // Wrap a sequence of body stmts into a single StmtPtr: pass through a sole + // stmt, wrap multiple into a SeqStmts, and synthesise an empty SeqStmts + // when the body is empty (a no-op body is still a valid loop / if branch). + static StmtPtr WrapBodyStmts(std::vector body_stmts, const Span& span); + + std::string base_name_; + std::size_t& temp_counter_; + bool nested_ = false; + int64_t barrier_count_ = 0; ///< Call-local generation counter; see EmitBarrier. + std::vector stmts_; +}; + +} // namespace lower_composite +} // namespace ir +} // namespace pypto + +#endif // SRC_IR_TRANSFORMS_LOWER_COMPOSITE_LOWER_COMPOSITE_BUILDER_H_ diff --git a/src/ir/transforms/lower_composite_ops_pass.cpp b/src/ir/transforms/lower_composite_ops_pass.cpp index 29f8937fc9..27bff884dc 100644 --- a/src/ir/transforms/lower_composite_ops_pass.cpp +++ b/src/ir/transforms/lower_composite_ops_pass.cpp @@ -38,42 +38,21 @@ #include "pypto/ir/transforms/base/mutator.h" #include "pypto/ir/transforms/pass_properties.h" #include "pypto/ir/transforms/passes.h" -#include "pypto/ir/transforms/utils/auto_name_utils.h" #include "pypto/ir/transforms/utils/mutable_copy.h" #include "pypto/ir/transforms/utils/op_predicates.h" #include "pypto/ir/transforms/utils/tensor_view_semantics.h" #include "pypto/ir/transforms/utils/tile_conversion_utils.h" #include "pypto/ir/type.h" #include "pypto/ir/type_inference.h" +#include "src/ir/transforms/lower_composite/lower_composite_builder.h" namespace pypto { namespace ir { namespace { -// ============================================================================ -// CommSetup — result struct for LoweringBuilder::EmitCommSetup() -// ============================================================================ - -/// Holds bound expressions from the comm-setup preamble (ctx, nranks, my_rank). -/// Returned by LoweringBuilder::EmitCommSetup() for use in subsequent phases. -struct CommSetup { - ExprPtr ctx; ///< Result of pld.system.get_comm_ctx - ExprPtr nranks_i32; ///< Result of pld.system.nranks (INT32) - ExprPtr nranks_idx; ///< nranks cast to INDEX (for loop bounds) - ExprPtr my_rank; ///< Result of pld.system.rank (INT32) -}; - -/// Safe negation: folds ConstInt(-value) directly when possible so the -/// PyPTO printer→parser roundtrip (which folds ``Neg(ConstInt)`` into -/// ``ConstInt(-value)``) produces structurally equal IR. Runtime -/// expressions are still wrapped with ``Neg``. -inline ExprPtr MakeNegation(const ExprPtr& value) { - if (auto c = As(value)) { - return std::make_shared(-c->value_, GetScalarDtype(value), value->span_); - } - return MakeNeg(value, value->span_); -} +using lower_composite::CommSetup; +using lower_composite::LoweringBuilder; // ============================================================================ // Self-clearing credit barrier — the shared, stateless barrier-signal protocol @@ -355,475 +334,6 @@ ExprPtr MakeCollectiveStageShape(const std::vector& transfer_shape, span); } -// ============================================================================ -// LoweringBuilder -// -// Per-call scratchpad handed to a composite-lowering rule. A rule appends one -// ``AssignStmt`` per intermediate temp via ``Bind`` and returns the final -// result ``ExprPtr``; the mutator wraps that result in the original target -// ``Var`` (or a fresh result ``Var`` for ``ReturnStmt`` calls) before splicing -// the accumulated statements into the surrounding sequence. -// -// In addition to ``Bind`` and the primitive op builders, the builder exposes -// structured control-flow constructors — ``EmitFor`` / ``EmitForReduce`` / -// ``EmitIf`` / ``EmitIfExpr`` — that hand the body off to a nested builder -// callback. The nested builder shares this builder's temp counter so every -// emitted temp gets a unique name across the entire rule, regardless of -// nesting depth. -// -// The temp counter is borrowed from the mutator so unique temp names span -// distinct composite-op calls in the same function. Barrier generations are -// call-local (see the self-clearing credit-barrier protocol above), so each -// LoweringBuilder instance — one per top-level composite-op call — owns its -// own ``barrier_count_`` that always starts at 0. -// ============================================================================ -class LoweringBuilder { - public: - /// @param base_name Name hint to derive temp names from (typically the - /// AssignStmt's LHS ``Var`` name). - /// @param temp_counter Reference to a mutator-owned counter; bumped per Bind. - LoweringBuilder(std::string base_name, std::size_t& temp_counter) - : base_name_(std::move(base_name)), temp_counter_(temp_counter) {} - - LoweringBuilder(std::string base_name, std::size_t& temp_counter, bool nested) - : base_name_(std::move(base_name)), temp_counter_(temp_counter), nested_(nested) {} - - /// Append an ``AssignStmt`` binding a fresh ``Var`` to ``expr`` and return - /// the new ``Var`` so it can be used as input to subsequent ops. The - /// ``qualifier`` is woven into the temp name for debuggability. - ExprPtr Bind(const std::string& qualifier, const ExprPtr& expr, const Span& span) { - auto var = std::make_shared(MakeTempName(qualifier), expr->GetType(), span); - stmts_.push_back(std::make_shared(var, expr, span)); - return var; - } - - /// Append a side-effecting expression without manufacturing an unused SSA - /// result. This is used by destination-passing ops whose output buffers are - /// explicit operands. - void EmitEval(const ExprPtr& expr, const Span& span) { - stmts_.push_back(std::make_shared(expr, span)); - } - - // Primitive op builders -- type deduction is delegated to OpRegistry so the - // result preserves the input TileType's shape/layout/dtype. - ExprPtr Muls(const ExprPtr& x, float c, const Span& span) { - auto tile_type = As(x->GetType()); - INTERNAL_CHECK_SPAN(tile_type, span) << "tile.muls input must be TileType"; - auto scalar = std::make_shared(static_cast(c), tile_type->dtype_, span); - return OpRegistry::GetInstance().Create("tile.muls", {x, scalar}, {}, span); - } - ExprPtr Adds(const ExprPtr& x, float c, const Span& span) { - auto tile_type = As(x->GetType()); - INTERNAL_CHECK_SPAN(tile_type, span) << "tile.adds input must be TileType"; - auto scalar = std::make_shared(static_cast(c), tile_type->dtype_, span); - return OpRegistry::GetInstance().Create("tile.adds", {x, scalar}, {}, span); - } - ExprPtr Add(const ExprPtr& a, const ExprPtr& b, const Span& span) { - return OpRegistry::GetInstance().Create("tile.add", {a, b}, {}, span); - } - ExprPtr Sub(const ExprPtr& a, const ExprPtr& b, const Span& span) { - return OpRegistry::GetInstance().Create("tile.sub", {a, b}, {}, span); - } - ExprPtr Mul(const ExprPtr& a, const ExprPtr& b, const Span& span) { - return OpRegistry::GetInstance().Create("tile.mul", {a, b}, {}, span); - } - ExprPtr Reduce(ReduceOp op, const ExprPtr& a, const ExprPtr& b, const Span& span) { - const char* op_name; - switch (op) { - case ReduceOp::kSum: - op_name = "tile.add"; - break; - case ReduceOp::kMax: - op_name = "tile.maximum"; - break; - case ReduceOp::kMin: - op_name = "tile.minimum"; - break; - case ReduceOp::kProd: - op_name = "tile.mul"; - break; - default: - INTERNAL_CHECK_SPAN(false, span) - << "pld.tensor.allreduce lowering received unknown ReduceOp " << static_cast(op); - } - return OpRegistry::GetInstance().Create(op_name, {a, b}, {}, span); - } - ExprPtr Cast(const ExprPtr& x, DataType to, int mode, const Span& span) { - std::vector> kw = {{"target_type", to}, {"mode", mode}}; - return OpRegistry::GetInstance().Create("tile.cast", {x}, kw, span); - } - - // ---- Scalar comparison helpers (yield BOOL-typed expressions, suitable as - // IfStmt conditions or loop guards). Delegated to the scalar_expr - // Make* helpers so operand promotion stays consistent with parser - // output. - ExprPtr NotEq(const ExprPtr& left, const ExprPtr& right, const Span& span) { - return MakeNe(left, right, span); - } - - ExprPtr Gt(const ExprPtr& left, const ExprPtr& right, const Span& span) { - return MakeGt(left, right, span); - } - - // ---- Collective-op helpers (DRY extraction for barrier/broadcast/allgather/ - // reduce_scatter/allreduce) ---- - - /// Emit comm-setup preamble: get_comm_ctx, nranks, rank. - /// Returns a CommSetup struct with the bound expressions for use in - /// subsequent phases. - CommSetup EmitCommSetup(const ExprPtr& comm_target, const Span& span) { - auto& reg = OpRegistry::GetInstance(); - CommSetup s; - s.ctx = Bind("ctx", reg.Create("pld.system.get_comm_ctx", {comm_target}, {}, span), span); - s.nranks_i32 = Bind("nranks", reg.Create("pld.system.nranks", {s.ctx}, {}, span), span); - s.nranks_idx = Bind("nranks_idx", std::make_shared(s.nranks_i32, DataType::INDEX, span), span); - s.my_rank = Bind("my_rank", reg.Create("pld.system.rank", {s.ctx}, {}, span), span); - return s; - } - - /// Emit notify-all loop: for peer in 0..nranks: if peer != my_rank: notify(...) - /// @param signal The signal DistributedTensor - /// @param nranks_idx Loop bound (INDEX-typed) - /// @param my_rank This rank's ID (INT32) - /// @param notify_op NotifyOp::kSet or NotifyOp::kAtomicAdd - /// @param value Value to notify (e.g., one_i32) - /// @param suffix Suffix for loop variable names (e.g., "" or "2" for re-notify) - /// @param span Source span for error reporting - void EmitNotifyAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, - NotifyOp notify_op, const ExprPtr& value, const std::string& suffix, const Span& span) { - auto zero_idx = std::make_shared(0, DataType::INDEX, span); - auto one_idx = std::make_shared(1, DataType::INDEX, span); - auto my_offsets = tile_conversion_utils::MakeSignalOffsets(my_rank, span); - - EmitFor( - "peer" + suffix, zero_idx, nranks_idx, one_idx, - [&](LoweringBuilder& body, const VarPtr& peer) { - body.EmitIf( - body.NotEq(peer, my_rank, span), - [&](LoweringBuilder& then_body) { - auto call = - OpRegistry::GetInstance().Create("pld.system.notify", {signal, peer, my_offsets, value}, - {{"op", static_cast(notify_op)}}, span); - then_body.Bind("notify" + suffix + "_ret", call, span); - }, - /*else_fn=*/nullptr, span); - }, - span); - } - - /// Overload for 2D signal matrices (e.g. ring allreduce [2*(NR-1), NR]). - /// @param row_offset Row index expression for the 2D signal (e.g. ring step var) - void EmitNotifyAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, - const ExprPtr& row_offset, NotifyOp notify_op, const ExprPtr& value, - const std::string& suffix, const Span& span) { - auto zero_idx = std::make_shared(0, DataType::INDEX, span); - auto one_idx = std::make_shared(1, DataType::INDEX, span); - auto my_offsets = tile_conversion_utils::MakeSignalOffsets(my_rank, row_offset, span); - - EmitFor( - "peer" + suffix, zero_idx, nranks_idx, one_idx, - [&](LoweringBuilder& body, const VarPtr& peer) { - body.EmitIf( - body.NotEq(peer, my_rank, span), - [&](LoweringBuilder& then_body) { - auto call = - OpRegistry::GetInstance().Create("pld.system.notify", {signal, peer, my_offsets, value}, - {{"op", static_cast(notify_op)}}, span); - then_body.Bind("notify" + suffix + "_ret", call, span); - }, - /*else_fn=*/nullptr, span); - }, - span); - } - - /// Emit wait-all loop: for src in 0..nranks: if src != my_rank: wait(...) - /// @param signal The signal DistributedTensor - /// @param nranks_idx Loop bound (INDEX-typed) - /// @param my_rank This rank's ID (INT32) - /// @param expected Expected signal value — the barrier generation (INT32) - /// @param suffix Suffix for loop variable names (e.g., "" or "2" for re-wait) - /// @param span Source span for error reporting - void EmitWaitAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, - const ExprPtr& expected, const std::string& suffix, const Span& span) { - auto zero_idx = std::make_shared(0, DataType::INDEX, span); - auto one_idx = std::make_shared(1, DataType::INDEX, span); - - EmitFor( - "src" + suffix, zero_idx, nranks_idx, one_idx, - [&](LoweringBuilder& body, const VarPtr& src) { - auto src_offsets = tile_conversion_utils::MakeSignalOffsets(src, span); - body.EmitIf( - body.NotEq(src, my_rank, span), - [&](LoweringBuilder& then_body) { - auto call = - OpRegistry::GetInstance().Create("pld.system.wait", {signal, src_offsets, expected}, - {{"cmp", static_cast(WaitCmp::kGe)}}, span); - then_body.Bind("wait" + suffix + "_ret", call, span); - }, - /*else_fn=*/nullptr, span); - }, - span); - } - - /// Overload for 2D signal matrices (e.g. ring allreduce [2*(NR-1), NR]). - /// @param row_offset Row index expression for the 2D signal (e.g. ring step var) - void EmitWaitAll(const ExprPtr& signal, const ExprPtr& nranks_idx, const ExprPtr& my_rank, - const ExprPtr& row_offset, const ExprPtr& expected, const std::string& suffix, - const Span& span) { - auto zero_idx = std::make_shared(0, DataType::INDEX, span); - auto one_idx = std::make_shared(1, DataType::INDEX, span); - - EmitFor( - "src" + suffix, zero_idx, nranks_idx, one_idx, - [&](LoweringBuilder& body, const VarPtr& src) { - auto src_offsets = tile_conversion_utils::MakeSignalOffsets(src, row_offset, span); - body.EmitIf( - body.NotEq(src, my_rank, span), - [&](LoweringBuilder& then_body) { - auto call = - OpRegistry::GetInstance().Create("pld.system.wait", {signal, src_offsets, expected}, - {{"cmp", static_cast(WaitCmp::kGe)}}, span); - then_body.Bind("wait" + suffix + "_ret", call, span); - }, - /*else_fn=*/nullptr, span); - }, - span); - } - - // ---- Self-clearing credit barrier protocol (see the file-header comment) ---- - - /// Emit one complete cross-rank barrier on ``signal``: ``AtomicAdd(1)`` into - /// every peer's cell, then wait for this call's generation on every peer's - /// cell. Returns the generation waited for (1-based, scoped to *this call* - /// only — every fresh ``LoweringBuilder`` starts counting at 0), so a rule - /// that fans out further barriers (the mesh allreduce's per-chunk barriers) - /// can continue the sequence from it, and so the rule can compute the total - /// credit count its ``EmitEpilogueReset`` call must subtract. - /// - /// Call this only from a rule's straight-line code — one call consumes exactly - /// one generation, so invoking it inside an ``EmitFor`` body would reserve a - /// single generation for a barrier that executes many times. Loop-resident - /// barriers must emit notify/wait by hand with a call-local expected value - /// (see the ring / mesh-chunked rules below). - int64_t EmitBarrier(const ExprPtr& signal, const CommSetup& comm, const std::string& suffix, - const Span& span) { - INTERNAL_CHECK_SPAN(!nested_, span) - << "Internal error: EmitBarrier must only be called from a top-level lowering rule, not from inside " - << "EmitFor / EmitIf / EmitIfExpr bodies. Loop- or condition-resident barriers must " - << "emit notify/wait by hand with a call-local expected value."; - const int64_t generation = ++barrier_count_; - auto one_i32 = std::make_shared(1, DataType::INT32, span); - auto expected_i32 = std::make_shared(generation, DataType::INT32, span); - EmitNotifyAll(signal, comm.nranks_idx, comm.my_rank, NotifyOp::kAtomicAdd, one_i32, suffix, span); - EmitWaitAll(signal, comm.nranks_idx, comm.my_rank, expected_i32, suffix, span); - return generation; - } - - /// Self-clearing epilogue: subtract ``total`` from every non-self peer's - /// contribution to *my own* cells, restoring the signal to all-zero once - /// every rank has run its own epilogue. ``total`` is the number of - /// ``AtomicAdd(+1)`` notifies this call issued per peer (the sum of every - /// ``EmitBarrier`` / hand-rolled notify-wait pair the rule emitted) — it may - /// be a runtime-computed expression, not just a ``ConstInt``: - /// ``pld.system.notify``'s value only requires ``ScalarType``. - /// - /// This is a self-notify (``peer == my_rank``): the codegen path resolves - /// ``peer == my_rank`` via the same identity mapping ``pld.tile.put`` / - /// ``pld.tile.get`` already rely on for their self-rank case, so this lands - /// on the exact same hardware atomic as an incoming remote add. - /// - /// Call this exactly once per rule invocation, from top-level code only, - /// after every barrier the rule issues. - void EmitEpilogueReset(const ExprPtr& signal, const CommSetup& comm, const ExprPtr& total, - const Span& span) { - INTERNAL_CHECK_SPAN(!nested_, span) - << "EmitEpilogueReset must only be called from a top-level lowering rule, exactly once " - "per call, after every EmitBarrier / hand-rolled notify-wait pair the rule issues."; - auto neg_total = MakeNegation(total); - auto zero_idx = std::make_shared(0, DataType::INDEX, span); - auto one_idx = std::make_shared(1, DataType::INDEX, span); - EmitFor( - "reset_src", zero_idx, comm.nranks_idx, one_idx, - [&](LoweringBuilder& body, const VarPtr& src) { - body.EmitIf( - body.NotEq(src, comm.my_rank, span), - [&](LoweringBuilder& then_body) { - auto src_offsets = tile_conversion_utils::MakeSignalOffsets(src, span); - auto call = OpRegistry::GetInstance().Create( - "pld.system.notify", {signal, comm.my_rank, src_offsets, neg_total}, - {{"op", static_cast(NotifyOp::kAtomicAdd)}}, span); - then_body.Bind("epilogue_reset_ret", call, span); - }, - /*else_fn=*/nullptr, span); - }, - span); - } - - /// 2D-signal overload (ring allreduce, ``[2*(NR-1), NR]``): subtract - /// ``total_per_row`` from every non-self cell of every one of ``num_rows`` - /// rows. Ring credits every row independently (one per round / sub-chunk - /// sequence), and every row's sub-chunk loop shares the same bound, so one - /// symbolic ``total_per_row`` resets all rows uniformly. - void EmitEpilogueReset(const ExprPtr& signal, const CommSetup& comm, const ExprPtr& num_rows, - const ExprPtr& total_per_row, const Span& span) { - INTERNAL_CHECK_SPAN(!nested_, span) - << "EmitEpilogueReset must only be called from a top-level lowering rule, exactly once per call."; - auto neg_total = MakeNegation(total_per_row); - auto zero_idx = std::make_shared(0, DataType::INDEX, span); - auto one_idx = std::make_shared(1, DataType::INDEX, span); - EmitFor( - "reset_row", zero_idx, num_rows, one_idx, - [&](LoweringBuilder& row_body, const VarPtr& row) { - row_body.EmitFor( - "reset_src", zero_idx, comm.nranks_idx, one_idx, - [&](LoweringBuilder& body, const VarPtr& src) { - body.EmitIf( - body.NotEq(src, comm.my_rank, span), - [&](LoweringBuilder& then_body) { - auto src_offsets = tile_conversion_utils::MakeSignalOffsets(src, row, span); - auto call = OpRegistry::GetInstance().Create( - "pld.system.notify", {signal, comm.my_rank, src_offsets, neg_total}, - {{"op", static_cast(NotifyOp::kAtomicAdd)}}, span); - then_body.Bind("epilogue_reset_ret", call, span); - }, - /*else_fn=*/nullptr, span); - }, - span); - }, - span); - } - - // ---- Structured control-flow constructors ---- - // - // Each method takes a body callback that receives a freshly-constructed - // nested ``LoweringBuilder`` scoped to the body region. The callback emits - // its body via the nested builder; this builder then drains the nested - // stmts, wraps them in a ``SeqStmts`` (when there is more than one), and - // emits the resulting ``ForStmt`` / ``IfStmt`` against its own ``stmts_``. - // - // The nested builder shares this builder's ``temp_counter_`` reference so - // emitted temp names stay unique across the entire rule regardless of - // nesting depth. - - /// Emit a side-effect-only ``for`` loop: - /// - /// for loop_var in range(start, stop, step): - /// - /// - /// ``body_fn`` receives a fresh body builder and the freshly-created loop - /// variable. The callback's return value is discarded — use this overload - /// for loops whose only purpose is side effects (e.g. issuing notify / - /// wait sequences). - void EmitFor(const std::string& loop_var_name, const ExprPtr& start, const ExprPtr& stop, - const ExprPtr& step, const std::function& body_fn, - const Span& span) { - auto loop_var = std::make_shared(MakeTempName(loop_var_name), start->GetType(), span); - LoweringBuilder body_builder(base_name_, temp_counter_, /*nested=*/true); - body_fn(body_builder, loop_var); - auto body_stmt = WrapBodyStmts(body_builder.TakeStmts(), span); - stmts_.push_back(std::make_shared(loop_var, start, stop, step, std::vector{}, - body_stmt, std::vector{}, span)); - } - - /// Emit a reducing ``for`` loop with one loop-carried accumulator. The - /// body callback receives a nested builder, the loop variable, and the - /// accumulator (typed via ``init_value``); it returns the next iteration's - /// accumulator value. The method returns an expression holding the - /// post-loop accumulator, ready to feed into subsequent ops. - ExprPtr EmitForReduce(const std::string& loop_var_name, const ExprPtr& start, const ExprPtr& stop, - const ExprPtr& step, const ExprPtr& init_value, - const std::function& body_fn, - const Span& span) { - auto loop_var = std::make_shared(MakeTempName(loop_var_name), start->GetType(), span); - auto iter_arg = std::make_shared(MakeTempName(loop_var_name + "_acc"), init_value->GetType(), - init_value, span); - LoweringBuilder body_builder(base_name_, temp_counter_, /*nested=*/true); - ExprPtr yield_val = body_fn(body_builder, loop_var, iter_arg); - INTERNAL_CHECK_SPAN(yield_val, span) - << "EmitForReduce body_fn must return the next iteration's accumulator value"; - body_builder.stmts_.push_back(std::make_shared(std::vector{yield_val}, span)); - auto body_stmt = WrapBodyStmts(body_builder.TakeStmts(), span); - auto return_var = - std::make_shared(MakeTempName(loop_var_name + "_final"), init_value->GetType(), span); - stmts_.push_back(std::make_shared(loop_var, start, stop, step, std::vector{iter_arg}, - body_stmt, std::vector{return_var}, span)); - return return_var; - } - - /// Emit a side-effect-only ``if`` statement: - /// - /// if cond: - /// - /// [else: - /// ] - /// - /// Pass ``nullptr`` for ``else_fn`` when there is no else branch. - void EmitIf(const ExprPtr& cond, const std::function& then_fn, - const std::function& else_fn, const Span& span) { - LoweringBuilder then_builder(base_name_, temp_counter_, /*nested=*/true); - then_fn(then_builder); - auto then_body = WrapBodyStmts(then_builder.TakeStmts(), span); - - std::optional else_body = std::nullopt; - if (else_fn) { - LoweringBuilder else_builder(base_name_, temp_counter_, /*nested=*/true); - else_fn(else_builder); - else_body = WrapBodyStmts(else_builder.TakeStmts(), span); - } - stmts_.push_back(std::make_shared(cond, then_body, else_body, std::vector{}, span)); - } - - /// Emit a value-producing ``if`` statement. Both branches must yield a - /// value (via their body_fn's ExprPtr return); the method returns an - /// expression holding the chosen value, ready to feed into subsequent ops. - ExprPtr EmitIfExpr(const ExprPtr& cond, const std::function& then_fn, - const std::function& else_fn, const Span& span) { - INTERNAL_CHECK_SPAN(then_fn && else_fn, span) - << "EmitIfExpr requires both then_fn and else_fn (the if must yield a value on every path)"; - LoweringBuilder then_builder(base_name_, temp_counter_, /*nested=*/true); - ExprPtr then_val = then_fn(then_builder); - INTERNAL_CHECK_SPAN(then_val, span) << "EmitIfExpr then_fn must return the yielded value"; - then_builder.stmts_.push_back(std::make_shared(std::vector{then_val}, span)); - auto then_body = WrapBodyStmts(then_builder.TakeStmts(), span); - - LoweringBuilder else_builder(base_name_, temp_counter_, /*nested=*/true); - ExprPtr else_val = else_fn(else_builder); - INTERNAL_CHECK_SPAN(else_val, span) << "EmitIfExpr else_fn must return the yielded value"; - else_builder.stmts_.push_back(std::make_shared(std::vector{else_val}, span)); - auto else_body = WrapBodyStmts(else_builder.TakeStmts(), span); - - auto return_var = std::make_shared(MakeTempName("if_res"), then_val->GetType(), span); - stmts_.push_back(std::make_shared(cond, then_body, std::optional(else_body), - std::vector{return_var}, span)); - return return_var; - } - - /// Drain accumulated statements (called by the mutator after the rule - /// returns). - std::vector TakeStmts() { return std::move(stmts_); } - - private: - std::string MakeTempName(const std::string& qualifier) { - return auto_name::BuildName(auto_name::GetBaseName(base_name_), qualifier, "tmp", - static_cast(temp_counter_++)); - } - - // Wrap a sequence of body stmts into a single StmtPtr: pass through a sole - // stmt, wrap multiple into a SeqStmts, and synthesise an empty SeqStmts - // when the body is empty (a no-op body is still a valid loop / if branch). - static StmtPtr WrapBodyStmts(std::vector body_stmts, const Span& span) { - if (body_stmts.empty()) return std::make_shared(std::vector{}, span); - if (body_stmts.size() == 1) return body_stmts.front(); - return std::make_shared(std::move(body_stmts), span); - } - - std::string base_name_; - std::size_t& temp_counter_; - bool nested_ = false; - int64_t barrier_count_ = 0; ///< Call-local generation counter; see EmitBarrier. - std::vector stmts_; -}; - // Signature for a composite-lowering rule. // // @param call Original composite-op Call. Rules read ``call->kwargs_``,