[Backend] Keep the AllReduce XOR-butterfly rule in the algorithm, not in the backend

Follow-up to 9a928180, which answered the shared lowerer's questions with one
per-assertion backend predicate (Impl::AllReduceWidthRequiresPowerOfTwo). That
named a check rather than a capability, so every further check would have added
another such hook. Reshape it:

- CheckAllReduceWidth(reducing_threads, scale, op_name) keeps only the three
  checks that hold for every all-reduce (positive threads, positive scale, scale
  divides threads). Its signature and body are now identical to upstream's.
- The power-of-two requirement moves to CheckXorButterflyWidth(reducing_threads,
  scale), named after the algorithm that actually imposes it. It is called from
  the backends that emit an XOR-butterfly intrinsic -- cuda/rocm in
  MakeBatchAllReduce and MakeScalarAllReduce -- and never from shared code.
  Ascend does not call it: AscendAllReduce falls back to a shared-memory tree
  reduction (ub_reduce) for arbitrary widths.
- Impl::AllReduceWidthRequiresPowerOfTwo is gone from all five Impls, and the
  Make{Batch,Scalar}AllReduce signatures are untouched.

The remaining backend hook is Impl::AllReduceNeedsWorkspace, which asks a real
capability question ("does your lowering need a scratch buffer") rather than
restating an assertion.

The rule's diagnostic loses its "tl.reduce:" / "tl.finalize_reducer:" prefix
because the emit helpers do not know which tile op drove them. Threading that
label through would have been the only alternative, and it would have put a
caller-supplied string into the intrinsic-naming helpers; the message names the
algorithm instead, which is the actionable part.

src/backend/common/ is free of TargetIsAscend, and this is now true of the
AllReduce path without any target-shaped hook.

Verified on 8x Ascend950DT (CANN 9.2.0): testing/ascend/ 1010 passed, 0 failed;
examples/ascend/ 96 passed, 0 failed; the non-power-of-two reduce widths (6,5)
and (7,3) still lower.
This commit is contained in:
LeiWang1999
2026-09-11 13:28:03 +08:00
parent 9a92818026
commit 4d40c1904e
9 changed files with 38 additions and 49 deletions
+10 -15
View File
@@ -2,17 +2,17 @@
* \file tl/ascend/op/ascend_allreduce_policy.h
* \brief Ascend answers to the two AllReduce questions the shared lowerers ask.
*
* The shared reduce/finalize lowerers do not branch on the target themselves;
* they ask the backend `Impl` for policy. Two of those answers depend on the
* all-reduce *algorithm* rather than on the IR, so they live here, next to the
* Ascend implementation they describe:
* The shared reduce lowerer does not branch on the target itself; it asks the
* backend `Impl` whether the lowering needs a shared-memory workspace. That
* answer depends on the all-reduce *algorithm* rather than on the IR, so it
* lives here, next to the Ascend implementation it describes. It has to mirror
* the constexpr dispatch in src/tl_templates/ascend/reduce.h, which is why it
* is stated once here instead of being restated in shared code.
*
* - whether the reduction width must be a power of two, and
* - whether the lowering needs a shared-memory workspace.
*
* Both are properties of `tl::AscendAllReduce`; the second has to mirror the
* constexpr dispatch in src/tl_templates/ascend/reduce.h, which is why it is
* stated once here instead of being restated in shared code.
* The XOR-butterfly power-of-two rule lives in backend/common/op/reduce.h as
* CheckXorButterflyWidth and is simply not called by Ascend: AscendAllReduce
* falls back to a shared-memory tree reduction (ub_reduce) for arbitrary
* widths, so there is nothing to enforce.
*/
#ifndef TVM_TL_ASCEND_OP_ASCEND_ALLREDUCE_POLICY_H_
@@ -22,11 +22,6 @@ namespace tvm {
namespace tl {
namespace ascend {
// AscendAllReduce falls back to a shared-memory tree reduction (ub_reduce) for
// widths that are not powers of two, so the XOR-butterfly power-of-two
// requirement does not apply.
inline constexpr bool kAllReduceWidthRequiresPowerOfTwo = false;
// Mirror AscendAllReduce<>::run()'s constexpr dispatch: the warp and cross-warp
// XOR-butterfly paths run entirely in registers, while the generic ub_reduce
// path needs the shared-memory workspace.
+2 -3
View File
@@ -89,9 +89,8 @@ Stmt LowerFinalizeReducer(const FinalizeReducerOpNode &op,
auto thread_offset = lower_args.thread_bounds->min;
Array<Stmt> step_stmts;
for (const auto &[reducing_threads, scale] : steps) {
backend::reduce::CheckAllReduceWidth(
reducing_threads, scale, "tl.finalize_reducer",
ascend::kAllReduceWidthRequiresPowerOfTwo);
backend::reduce::CheckAllReduceWidth(reducing_threads, scale,
"tl.finalize_reducer");
std::stringstream ss;
ss << "tl::AscendAllReduce<" << op_str << ", " << reducing_threads << ", "
-4
View File
@@ -18,10 +18,6 @@ using namespace tirx;
namespace ascend {
struct Reduce : backend::ReduceLowerer<Reduce> {
static bool AllReduceWidthRequiresPowerOfTwo(Target) {
return ascend::kAllReduceWidthRequiresPowerOfTwo;
}
static bool AllReduceNeedsWorkspace(int reducing_threads, int scale, Target) {
return ascend::AllReduceNeedsWorkspace(reducing_threads, scale);
}
+2 -3
View File
@@ -91,9 +91,8 @@ template <typename Impl> struct FinalizeReducerLowerer {
Array<Stmt> step_stmts;
for (const auto &[reducing_threads, scale] : steps) {
reduce::CheckAllReduceWidth(
reducing_threads, scale, "tl.finalize_reducer",
Impl::AllReduceWidthRequiresPowerOfTwo(lower_args.target));
reduce::CheckAllReduceWidth(reducing_threads, scale,
"tl.finalize_reducer");
bool use_batch = effective_batch > 1 &&
reducing_threads > Impl::WarpSize(lower_args.target);
+16 -16
View File
@@ -186,8 +186,7 @@ inline int GetPreferredVectorizedSize(DataType dt,
}
inline void CheckAllReduceWidth(int reducing_threads, int scale,
const char *op_name,
bool requires_power_of_two_width = true) {
const char *op_name) {
ICHECK_GT(reducing_threads, 0)
<< op_name << ": AllReduce threads must be positive, got "
<< reducing_threads;
@@ -196,17 +195,20 @@ inline void CheckAllReduceWidth(int reducing_threads, int scale,
ICHECK_EQ(reducing_threads % scale, 0)
<< op_name << ": AllReduce threads (" << reducing_threads
<< ") must be divisible by scale (" << scale << ")";
// The power-of-two requirement belongs to the XOR-butterfly shuffle
// reduction. A backend whose all-reduce handles arbitrary widths answers
// false through Impl::AllReduceWidthRequiresPowerOfTwo().
if (!requires_power_of_two_width) {
return;
}
}
// Additional requirement of the XOR-butterfly shuffle all-reduce, which has no
// fallback for widths that are not powers of two. It is deliberately not part
// of CheckAllReduceWidth: a backend that lowers to an all-reduce with an
// arbitrary-width fallback has nothing to satisfy here. It is called from the
// backends that emit an XOR-butterfly intrinsic and never from shared code, so
// this header stays free of target branching.
inline void CheckXorButterflyWidth(int reducing_threads, int scale) {
int logical_width = reducing_threads / scale;
int shift = 0;
ICHECK(tirx::is_const_power_of_two_integer(Integer(logical_width), &shift))
<< op_name << ": XOR-butterfly AllReduce requires logical_width "
<< "(threads / scale) to be a positive power of two, got "
<< "XOR-butterfly all-reduce requires logical_width (threads / scale) to "
"be a positive power of two, got "
<< logical_width << " (threads=" << reducing_threads
<< ", scale=" << scale << ")";
}
@@ -1012,9 +1014,8 @@ template <typename Impl> struct ReduceLowerer {
for (const auto &thread_step : reduce_plan.thread_steps) {
int reducing_threads = thread_step.ReducingThreads();
reduce::CheckAllReduceWidth(
reducing_threads, thread_step.scale, "tl.reduce",
Impl::AllReduceWidthRequiresPowerOfTwo(lower_args.target));
reduce::CheckAllReduceWidth(reducing_threads, thread_step.scale,
"tl.reduce");
int block_threads =
static_cast<int>(*as_const_int(lower_args.thread_bounds->extent));
auto thread_offset = lower_args.thread_bounds->min;
@@ -1207,9 +1208,8 @@ template <typename Impl> struct ReduceLowerer {
for (const auto &thread_step : reduce_plan.thread_steps) {
int reducing_threads = thread_step.ReducingThreads();
reduce::CheckAllReduceWidth(
reducing_threads, thread_step.scale, "tl.reduce",
Impl::AllReduceWidthRequiresPowerOfTwo(lower_args.target));
reduce::CheckAllReduceWidth(reducing_threads, thread_step.scale,
"tl.reduce");
auto thread_offset = lower_args.thread_bounds->min;
PrimExpr all_threads = lower_args.thread_bounds->extent;
if (reducing_threads > 32 &&
+2 -2
View File
@@ -17,8 +17,6 @@ using namespace tirx;
namespace cuda {
struct FinalizeReducer : backend::FinalizeReducerLowerer<FinalizeReducer> {
static bool AllReduceWidthRequiresPowerOfTwo(Target) { return true; }
static bool AllReduceNeedsWorkspace(int reducing_threads, int, Target) {
return reducing_threads > 32;
}
@@ -30,6 +28,7 @@ struct FinalizeReducer : backend::FinalizeReducerLowerer<FinalizeReducer> {
PrimExpr thread_offset,
PrimExpr all_threads, int batch,
int workspace_stride, Target target) {
backend::reduce::CheckXorButterflyWidth(reducing_threads, scale);
std::stringstream ss;
ss << "tl::AllReduce<" << reducer << ", " << reducing_threads << ", "
<< scale << ", " << thread_offset;
@@ -46,6 +45,7 @@ struct FinalizeReducer : backend::FinalizeReducerLowerer<FinalizeReducer> {
int reducing_threads, int scale,
PrimExpr thread_offset,
PrimExpr all_threads, Target target) {
backend::reduce::CheckXorButterflyWidth(reducing_threads, scale);
std::stringstream ss;
ss << "tl::AllReduce<" << reducer << ", " << reducing_threads << ", "
<< scale << ", " << thread_offset;
+2 -2
View File
@@ -19,8 +19,6 @@ using namespace tirx;
namespace cuda {
struct Reduce : backend::ReduceLowerer<Reduce> {
static bool AllReduceWidthRequiresPowerOfTwo(Target) { return true; }
static bool AllReduceNeedsWorkspace(int reducing_threads, int, Target) {
return reducing_threads > 32;
}
@@ -70,6 +68,7 @@ struct Reduce : backend::ReduceLowerer<Reduce> {
PrimExpr thread_offset,
PrimExpr all_threads, int batch,
int workspace_stride, Target target) {
backend::reduce::CheckXorButterflyWidth(reducing_threads, scale);
std::stringstream ss;
ss << "tl::AllReduce<" << reducer << ", " << reducing_threads << ", "
<< scale << ", " << thread_offset;
@@ -86,6 +85,7 @@ struct Reduce : backend::ReduceLowerer<Reduce> {
int reducing_threads, int scale,
PrimExpr thread_offset,
PrimExpr all_threads, Target target) {
backend::reduce::CheckXorButterflyWidth(reducing_threads, scale);
std::stringstream ss;
ss << "tl::AllReduce<" << reducer << ", " << reducing_threads << ", "
<< scale << ", " << thread_offset;
+2 -2
View File
@@ -17,8 +17,6 @@ using namespace tirx;
namespace rocm {
struct FinalizeReducer : backend::FinalizeReducerLowerer<FinalizeReducer> {
static bool AllReduceWidthRequiresPowerOfTwo(Target) { return true; }
static bool AllReduceNeedsWorkspace(int reducing_threads, int, Target) {
return reducing_threads > 32;
}
@@ -30,6 +28,7 @@ struct FinalizeReducer : backend::FinalizeReducerLowerer<FinalizeReducer> {
PrimExpr thread_offset, PrimExpr,
int batch, int workspace_stride,
Target) {
backend::reduce::CheckXorButterflyWidth(reducing_threads, scale);
std::stringstream ss;
ss << "tl::AllReduce<" << reducer << ", " << reducing_threads << ", "
<< scale << ", " << thread_offset << ", " << batch << ", "
@@ -41,6 +40,7 @@ struct FinalizeReducer : backend::FinalizeReducerLowerer<FinalizeReducer> {
int reducing_threads, int scale,
PrimExpr thread_offset, PrimExpr,
Target) {
backend::reduce::CheckXorButterflyWidth(reducing_threads, scale);
std::stringstream ss;
ss << "tl::AllReduce<" << reducer << ", " << reducing_threads << ", "
<< scale << ", " << thread_offset << ">::run";
+2 -2
View File
@@ -17,8 +17,6 @@ using namespace tirx;
namespace rocm {
struct Reduce : backend::ReduceLowerer<Reduce> {
static bool AllReduceWidthRequiresPowerOfTwo(Target) { return true; }
static bool AllReduceNeedsWorkspace(int reducing_threads, int, Target) {
return reducing_threads > 32;
}
@@ -34,6 +32,7 @@ struct Reduce : backend::ReduceLowerer<Reduce> {
PrimExpr thread_offset, PrimExpr,
int batch, int workspace_stride,
Target) {
backend::reduce::CheckXorButterflyWidth(reducing_threads, scale);
std::stringstream ss;
ss << "tl::AllReduce<" << reducer << ", " << reducing_threads << ", "
<< scale << ", " << thread_offset << ", " << batch << ", "
@@ -45,6 +44,7 @@ struct Reduce : backend::ReduceLowerer<Reduce> {
int reducing_threads, int scale,
PrimExpr thread_offset, PrimExpr,
Target) {
backend::reduce::CheckXorButterflyWidth(reducing_threads, scale);
std::stringstream ss;
ss << "tl::AllReduce<" << reducer << ", " << reducing_threads << ", "
<< scale << ", " << thread_offset << ">::run";