mirror of
https://github.com/tile-ai/tilelang.git
synced 2026-10-03 06:48:18 +08:00
[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:
@@ -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.
|
||||
|
||||
@@ -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 << ", "
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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 &&
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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";
|
||||
|
||||
@@ -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";
|
||||
|
||||
Reference in New Issue
Block a user