diff --git a/include/tvm/ir/base_expr.h b/include/tvm/ir/base_expr.h index 9ce81b197884..74e9c7a65fe9 100644 --- a/include/tvm/ir/base_expr.h +++ b/include/tvm/ir/base_expr.h @@ -371,6 +371,20 @@ class Expr : public ffi::ObjectRef { TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(Expr, ffi::ObjectRef, ExprNode); }; +/*! \brief Base node for traversable expressions eliminated during compilation. */ +class StagingExprNode : public ExprNode { + public: + static void RegisterReflection() { ffi::reflection::ObjectDef(); } + static constexpr const uint32_t _type_child_slots = 4; + TVM_FFI_DECLARE_OBJECT_INFO("ir.StagingExpr", StagingExprNode, ExprNode); +}; + +/*! \brief Managed reference to a staging expression. */ +class StagingExpr : public Expr { + public: + TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(StagingExpr, Expr, StagingExprNode); +}; + /*! * \brief Base node for opaque construction-time expressions. * diff --git a/include/tvm/ir/expr.h b/include/tvm/ir/expr.h index aa4b2dd44bfa..0936618f4d54 100644 --- a/include/tvm/ir/expr.h +++ b/include/tvm/ir/expr.h @@ -438,6 +438,48 @@ class PrimVar : public PrimExpr { static constexpr bool _type_container_is_exact = false; }; +/*! + * \brief A typed staging expression representing a lambda computation. + * + * LambdaExpr records computations such as reduction combiners and predication + * rules. Its body may describe computations on runtime values. + * + * Parameters are bound within the expression body, which may produce a scalar + * or tuple result. The lambda has a FuncType describing its parameter and + * return types. + * + * As a StagingExpr, LambdaExpr is eliminated during compilation and does not + * remain in executable IR. + */ +class LambdaExprNode : public StagingExprNode { + public: + /*! \brief Lambda-local parameter definitions. */ + ffi::Array vars; + /*! \brief Computation over the parameters and captured expressions. */ + Expr body; + + explicit LambdaExprNode(ffi::UnsafeInit tag) : body(tag) {} + explicit LambdaExprNode(Expr body) : body(std::move(body)) {} + /*! \brief Simultaneously substitute arguments for the bound parameters. */ + TVM_DLL Expr Apply(const ffi::Array& arguments) const; + + static void RegisterReflection() { + namespace refl = tvm::ffi::reflection; + refl::ObjectDef() + .def_ro("vars", &LambdaExprNode::vars, refl::AttachFieldFlag::SEqHashDefSimple()) + .def_ro("body", &LambdaExprNode::body); + } + static constexpr TVMFFISEqHashKind _type_s_eq_hash_kind = kTVMFFISEqHashKindTreeNode; + TVM_FFI_DECLARE_OBJECT_INFO_FINAL("ir.LambdaExpr", LambdaExprNode, StagingExprNode); +}; + +/*! \brief Managed reference to a typed staging lambda. */ +class LambdaExpr : public StagingExpr { + public: + TVM_DLL explicit LambdaExpr(ffi::Array vars, Expr body); + TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(LambdaExpr, StagingExpr, LambdaExprNode); +}; + class GlobalVar; /*! * \brief Global variable that lives in the top-level module. diff --git a/include/tvm/ir/expr_functor.h b/include/tvm/ir/expr_functor.h index 95bd226a7782..64df3d6e76df 100644 --- a/include/tvm/ir/expr_functor.h +++ b/include/tvm/ir/expr_functor.h @@ -73,6 +73,9 @@ class ExprFunctor { return (*vtable_)(node, this, std::forward(args)...); } + virtual R Dispatch_(const LambdaExprNode* node, Args... args) { + return DispatchDefault_(node, std::forward(args)...); + } virtual R Dispatch_(const OpaqueExprNode* node, Args... args) { return DispatchDefault_(node, std::forward(args)...); } @@ -230,6 +233,7 @@ class ExprFunctor { * \param vtable The table to initialize before adding derived registrations. */ static void InitVTable(VTable* vtable) { + SetDispatch(vtable); SetDispatch(vtable); SetDispatch(vtable); SetDispatch(vtable); diff --git a/include/tvm/s_tir/stmt.h b/include/tvm/s_tir/stmt.h index 3edce701ac73..5685694ec24c 100644 --- a/include/tvm/s_tir/stmt.h +++ b/include/tvm/s_tir/stmt.h @@ -254,9 +254,6 @@ constexpr const char* fragment_layout = "fragment_layout"; */ constexpr const char* loop_partition_hint = "loop_partition_hint"; -/*! \brief Mark of reduce scope */ -constexpr const char* reduce_scope = "reduce_scope"; - // ----------------------------------------------------------------------- // meta_schedule annotations // ----------------------------------------------------------------------- diff --git a/include/tvm/tirx/builtin.h b/include/tvm/tirx/builtin.h index 27033de7165c..784753f1c763 100644 --- a/include/tvm/tirx/builtin.h +++ b/include/tvm/tirx/builtin.h @@ -500,17 +500,33 @@ TVM_DLL const Op& tvm_warp_shuffle_xor(); TVM_DLL const Op& tvm_warp_activemask(); /*! - * \brief See pesudo code + * \brief Cross-thread reduction with an explicit typed combiner and identities. * - * void tvm_thread_allreduce(UIntImm size, Expr source0, ..., Expr cond, - * Var reduce_temp0, .., Var thread_idx1, ...) { - * // constraint by the other thread_idx remain the same. - * // reduce_temp is used to save intermediate result. - * reduce_temp0, ... = reduce(combiner, source0, ..., cond - * over [thread_idx1, thread_idx2] passed by any caller) - * } + * void tvm_thread_allreduce(LambdaExpr combine, Expr identity, Expr values, + * PrimExpr predicate, Expr destinations, Expr thread_axes); + * + * For N values, combine binds lhs[0:N] followed by rhs[0:N] and returns an + * N-element Tuple, or a scalar when N is one. Identity, values, destinations + * and thread_axes may each be a scalar or an explicit Tuple of fields. + * Each value, identity, pair of parameters and result have + * the same primitive type. Inactive inputs are replaced by their identities. + * Destinations are N tensor loads (optionally cast for boolean storage), and + * thread_axes are reduction thread variables or zero for simplified unit axes. + * Other thread indices remain fixed. The operation writes the reduced values + * to the destination tensors and returns void. */ TVM_DLL const Op& tvm_thread_allreduce(); + +/*! + * \brief View a scalar all-reduce operand/result as one field, or expose its Tuple fields. + * \param value The scalar expression or explicit Tuple. + * \return The fields without changing the expression's representation in IR. + */ +inline ffi::Array GetAllreduceFields(const Expr& value) { + if (const auto* tuple = value.as()) return tuple->fields; + return {value}; +} + // Metal cooperative_tensor intrinsics (MetalPerformancePrimitives / Metal 4) /*! diff --git a/include/tvm/tirx/tile_primitive.h b/include/tvm/tirx/tile_primitive.h index 70568d5c4e40..1716e29a03f3 100644 --- a/include/tvm/tirx/tile_primitive.h +++ b/include/tvm/tirx/tile_primitive.h @@ -33,51 +33,6 @@ namespace tvm { namespace tirx { -/*! - * \brief A reified Python lambda: a list of bound variables and a body over them. - * - * Used by tile primitive ops that take a per-element expression over the - * destination axes (e.g. ``tirx.tile.select``). ``vars`` are the abstract - * axis variables (lambda-bound); ``pred`` is the body referencing them. - * At lowering time the dispatch substitutes ``vars`` with the concrete - * instruction axes via ``Apply``. - */ -class LambdaExprNode : public ExprNode { - public: - explicit LambdaExprNode(ffi::UnsafeInit tag) : pred(tag) {} - - explicit LambdaExprNode(PrimExpr pred) : pred(std::move(pred)) {} - - /*! \brief The bound variables of the lambda. */ - Array vars; - /*! \brief The lambda body over ``vars``. */ - PrimExpr pred; - - /*! \brief Replace the bound variables with the given indices, returning the substituted body. */ - PrimExpr Apply(const Array& indices) const; - - static void RegisterReflection() { - namespace refl = tvm::ffi::reflection; - refl::ObjectDef() - .def_ro("vars", &LambdaExprNode::vars, refl::AttachFieldFlag::SEqHashDefPattern()) - .def_ro("pred", &LambdaExprNode::pred); - } - - static constexpr TVMFFISEqHashKind _type_s_eq_hash_kind = kTVMFFISEqHashKindTreeNode; - TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tirx.LambdaExpr", LambdaExprNode, ExprNode); -}; - -/*! - * \brief Managed reference to LambdaExprNode. - * \sa LambdaExprNode - */ -class LambdaExpr : public Expr { - public: - explicit LambdaExpr(Array vars, PrimExpr pred); - - TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(LambdaExpr, Expr, LambdaExprNode); -}; - /*! * \brief The type of the function that sanitizes the arguments of a TIRX operator. * \param op The operator. diff --git a/python/tvm/ir/__init__.py b/python/tvm/ir/__init__.py index 7703a6dc4043..7151ef6f1e21 100644 --- a/python/tvm/ir/__init__.py +++ b/python/tvm/ir/__init__.py @@ -54,6 +54,8 @@ ExprOperand, ExprWithOp, GlobalVar, + LambdaExpr, + StagingExpr, OpaqueExpr, Range, TensorLoad, diff --git a/python/tvm/ir/expr.py b/python/tvm/ir/expr.py index 1177f52c5631..9d67db2683e2 100644 --- a/python/tvm/ir/expr.py +++ b/python/tvm/ir/expr.py @@ -16,6 +16,9 @@ # under the License. """Common expressions data structures in the IR.""" +from collections.abc import Callable +from numbers import Number + import tvm_ffi import tvm @@ -59,6 +62,11 @@ def __getitem__(self, index): ) +@tvm_ffi.register_object("ir.StagingExpr") +class StagingExpr(Expr): + """A traversable expression eliminated before executable IR.""" + + @tvm_ffi.register_object("ir.OpaqueExpr") class OpaqueExpr(Expr): """Base class for opaque values that must be removed from finished IR.""" @@ -681,6 +689,83 @@ def __init__( self.__init_handle_by_constructor__(_ffi_api.Var, name, ty, span) +def _lambda_type(annotation): + """Normalize explicit types and the script scalar constructor forms.""" + if isinstance(annotation, tvm.ir.Type): + return annotation + if isinstance(annotation, str): + return Var("", annotation).ty + if isinstance(annotation, tvm.DataType): + return tvm.ir.PrimType(annotation) + dtype = getattr(annotation, "_dtype_str", None) + if dtype is not None: + return tvm.ir.PrimType(dtype) + if callable(annotation): + value = annotation() + if isinstance(value, Expr): + return value.ty + if isinstance(value, tvm.ir.Type): + return value + raise TypeError("Lambda parameter and return annotations must be explicit IR types") + + +def _lambda_result(value): + if isinstance(value, tuple | list): + return Tuple([_lambda_result(field) for field in value]) + if isinstance(value, Number): + value = const(value) + elif isinstance(value, str): + value = StringImm(value) + else: + value = tvm.runtime.convert(value) + if not isinstance(value, Expr): + raise TypeError("Lambda body must be an Expr or a tuple/list of expressions") + return value + + +@tvm_ffi.register_object("ir.LambdaExpr") +class LambdaExpr(StagingExpr, Scriptable): + """A typed staging expression representing a lambda computation. + + LambdaExpr records computations such as reduction combiners and predication + rules. Its body may describe computations on runtime values. + + Parameters are bound within the expression body, which may produce a scalar + or tuple result. The lambda has a FuncType describing its parameter and + return types. + + As a StagingExpr, LambdaExpr is eliminated during compilation and does not + remain in executable IR. + + Parameters + ---------- + parameter_types : list[Type] + Explicit parameter types, in callable argument order. Primitive dtype + strings and script scalar constructors are also accepted. + function : Callable + A callable evaluated once with one fresh typed Var per supplied type. + Tuple/list results become shared IR Tuple expressions. + ret_type : Type, optional + An exact return-type check. No implicit conversion or cast is inserted. + """ + + vars: list[Var] + body: Expr + + def __init__(self, parameter_types, function: Callable, *, ret_type=None): + variables = [Var(f"arg{i}", _lambda_type(ty)) for i, ty in enumerate(parameter_types)] + body = _lambda_result(function(*variables)) + if ret_type is not None: + expected = _lambda_type(ret_type) + if not tvm_ffi.structural_equal(expected, body.ty): + raise TypeError("LambdaExpr return annotation does not match the body type") + self.__init_handle_by_constructor__(_ffi_api.LambdaExpr, variables, body) + + def apply(self, arguments: list[Expr]) -> Expr: + """Substitute arguments simultaneously for the lambda's bound variables.""" + return _ffi_api.LambdaExprApply(self, [_lambda_result(arg) for arg in arguments]) + + @tvm_ffi.register_object("ir.Range") class Range(Node, Scriptable): """Represent a range in TVM. diff --git a/python/tvm/relax/frontend/nn/llm/_decode_kernels.py b/python/tvm/relax/frontend/nn/llm/_decode_kernels.py index 83bba4dfe270..bb38cf90f703 100644 --- a/python/tvm/relax/frontend/nn/llm/_decode_kernels.py +++ b/python/tvm/relax/frontend/nn/llm/_decode_kernels.py @@ -324,12 +324,7 @@ def batch_decode_paged_kv( with Ts.sblock("block_cross_thread"): Ts.reads(S_reduce_local[0]) Ts.writes(t0[0]) - T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) - T.tvm_thread_allreduce(T.uint32(1), S_reduce_local[0], True, t0[0], tx) + T.tvm_thread_allreduce(T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), (T.float32(0),), (S_reduce_local[0],), True, (t0[0],), (tx,)) S_local[j] = -5e4 if (iterator * bdz + tz) * bdy * tile_size_per_bdx + j < kv_chunk_len[0]: diff --git a/python/tvm/script/ir_builder/__init__.py b/python/tvm/script/ir_builder/__init__.py index 30f92ab098e1..327a55048b46 100644 --- a/python/tvm/script/ir_builder/__init__.py +++ b/python/tvm/script/ir_builder/__init__.py @@ -39,8 +39,8 @@ with_at_group_, ) from .frame import IRModuleFrame -from .ir import _get_dialect_builder as __getattr__ from .ir import ( + Lambda, constexpr, dtype, dynamic, @@ -50,6 +50,7 @@ module_global_infos, module_set_attr, ) +from .ir import _get_dialect_builder as __getattr__ from .parser_protocol import ( check_well_formed_, decl_function, @@ -68,6 +69,7 @@ "GenericConst", "IRBuilder", "IRModuleFrame", + "Lambda", "MissingType", "PrimType", "Range", diff --git a/python/tvm/script/ir_builder/ir.py b/python/tvm/script/ir_builder/ir.py index 6b31d7b174e7..d7f1b40dae74 100644 --- a/python/tvm/script/ir_builder/ir.py +++ b/python/tvm/script/ir_builder/ir.py @@ -19,7 +19,7 @@ from typing import NoReturn, TypeVar from tvm import DataType -from tvm.ir import DataTypeImm, GlobalInfo, Span, Var +from tvm.ir import DataTypeImm, GlobalInfo, LambdaExpr, Span, Var from tvm.runtime import Object as tvm_Object from . import _ffi_api @@ -189,3 +189,12 @@ def _get_dialect_builder(name: str): setattr(sys.modules["tvm.script.ir_builder"], name, module) return module raise AttributeError(f"module 'tvm.script.ir_builder' has no attribute {name!r}") + + +def Lambda(parameter_types, function, *, ret_type=None): # pylint: disable=invalid-name + """Build a shared staging lambda from explicit types and a Python callable. + + Scalar constructors such as ``T.float32`` may be used as parameter types. + An optional return annotation checks the body type without inserting casts. + """ + return LambdaExpr(parameter_types, function, ret_type=ret_type) diff --git a/python/tvm/tirx/__init__.py b/python/tvm/tirx/__init__.py index 705540f31e33..01989b0a3531 100644 --- a/python/tvm/tirx/__init__.py +++ b/python/tvm/tirx/__init__.py @@ -54,7 +54,7 @@ from .stmt import IfThenElse, Evaluate, stmt_seq, stmt_list from .stmt import BufferRegion, BufferRegionType from .stmt import ScopeIdDefStmt -from .tile_primitive import DispatchContext, LambdaExpr, TilePrimitiveCall +from .tile_primitive import DispatchContext, TilePrimitiveCall from .function import PrimFunc, IndexMap diff --git a/python/tvm/tirx/op.py b/python/tvm/tirx/op.py index 9e3dd06fd394..6220d308da93 100644 --- a/python/tvm/tirx/op.py +++ b/python/tvm/tirx/op.py @@ -811,20 +811,45 @@ def address_of(obj: Var | TensorLoad, span: Span | None = None) -> Expr: raise ValueError(f"Invalid object type: {type(obj)}") -def tvm_thread_allreduce(*freduce_args): - """Perform allreduce inside threadblock. +def tvm_thread_allreduce(combine, identity, values, predicate, destinations, thread_axes): + """Perform an all-reduce inside a thread block. Parameters ---------- - freduce_args : Expr - The args. + combine : tvm.ir.LambdaExpr + Typed combining lambda with parameters ordered as all left-hand values + followed by all right-hand values. Its body returns a scalar for one + result or a Tuple of results. + identity : Expr or Sequence[Expr] + Identity value for each reduction result. + values : Expr or Sequence[Expr] + Values contributed by the current thread. + predicate : PrimExpr + Boolean participation predicate. Inactive threads contribute identities. + destinations : Expr or Sequence[Expr] + Tensor loads identifying the destinations of the reduction results. + thread_axes : Expr or Sequence[Expr] + Thread variables participating in the reduction. Returns ------- call : Expr - The call expression. + The void call expression with six explicit operands. """ - return call_intrin("void", "tirx.tvm_thread_allreduce", *freduce_args) + + def as_operand(value): + return tvm.ir.Tuple(value) if isinstance(value, list | tuple | Array) else value + + return call_intrin( + "void", + "tirx.tvm_thread_allreduce", + combine, + as_operand(identity), + as_operand(values), + predicate, + as_operand(destinations), + as_operand(thread_axes), + ) def tvm_thread_invariant(cond): diff --git a/python/tvm/tirx/script/ir_builder/ir.py b/python/tvm/tirx/script/ir_builder/ir.py index fceb8ddbe3a5..372b8c9193b6 100644 --- a/python/tvm/tirx/script/ir_builder/ir.py +++ b/python/tvm/tirx/script/ir_builder/ir.py @@ -38,7 +38,7 @@ from tvm import tirx as tir from tvm.ir import Range, Type, is_prim_expr from tvm.script.ir_builder.base import MISSING, IRBuilder -from tvm.script.ir_builder.ir import meta_var +from tvm.script.ir_builder.ir import Lambda, meta_var from tvm.script.parser.protocol_registry import ( register_mutable_decl as _register_mutable_decl, ) @@ -1613,6 +1613,7 @@ def Ptr(dtype, storage_scope="global", *, span=None): "IntImm", "Iter", "IterVar", + "Lambda", "Layout", "LetAnnotation", "LocalVectorAnnotation", diff --git a/python/tvm/tirx/script/ir_builder/tirx.py b/python/tvm/tirx/script/ir_builder/tirx.py index 2348a88da79f..e9cec36ff636 100644 --- a/python/tvm/tirx/script/ir_builder/tirx.py +++ b/python/tvm/tirx/script/ir_builder/tirx.py @@ -21,8 +21,8 @@ import tvm import tvm.tirx.operator as tirx_op -from tvm.ir import Op, TensorRegion -from tvm.tirx import Expr, LambdaExpr, Var, buffer_data, is_tensor_var +from tvm.ir import LambdaExpr, Op, PrimType, TensorRegion +from tvm.tirx import Expr, Var, buffer_data, is_tensor_var from tvm.tirx.exec_scope import _SCOPE_KIND_TO_NAME, ExecScope from tvm.tirx.expr import FloatImm, IntImm from tvm.tirx.lang.alloc_pool import SMEMPool, TMEMPool @@ -1679,7 +1679,11 @@ def select( if is_tensor_var(false_value): false_value = _to_region(false_value) if not isinstance(pred, LambdaExpr): - pred = LambdaExpr(pred) + pred = LambdaExpr([PrimType("int32")] * len(dst.region), pred) + if len(pred.vars) != len(dst.region) or any(var.ty != PrimType("int32") for var in pred.vars): + raise TypeError("Tile select requires one int32 lambda parameter per destination axis") + if pred.ty.ret_type != PrimType("bool"): + raise TypeError("Tile select requires a scalar boolean lambda result") return f_insert(tirx_op.Select(dst, true_value, false_value, pred, scope=scope)) diff --git a/python/tvm/tirx/tile_primitive.py b/python/tvm/tirx/tile_primitive.py index 94c390f03c28..48049f4a0f6b 100644 --- a/python/tvm/tirx/tile_primitive.py +++ b/python/tvm/tirx/tile_primitive.py @@ -14,14 +14,12 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -"""TIRx tile primitive IR nodes: LambdaExpr, DispatchContext, TilePrimitiveCall. +"""TIRx tile primitive IR nodes: DispatchContext, TilePrimitiveCall. Mirrors the C++ header ``include/tvm/tirx/tile_primitive.h``. """ # pylint: disable=no-member -import inspect -from collections.abc import Callable from typing import Any, ClassVar import tvm_ffi @@ -37,27 +35,6 @@ from .stmt import Stmt -@register_object("tirx.LambdaExpr") -class LambdaExpr(Expr): - """A reified Python lambda: bound variables and a body over them. - - Used by tile primitive ops that take a per-element expression over the - destination axes (e.g. ``tirx.tile.select``). - """ - - vars: list[Var] - pred: Expr - - def __init__(self, f_pred: Callable[..., Expr]): - vars = [Var(name, "int32") for name in inspect.signature(f_pred).parameters] - pred = f_pred(*vars) - self.__init_handle_by_constructor__(_ffi_api.LambdaExpr, vars, pred) - - def apply(self, indices: list[Expr]) -> Expr: - """Substitute the bound variables with the given indices, returning the body.""" - return _ffi_api.LambdaExprApply(self, indices) - - @register_object("tirx.DispatchContext") class DispatchContext(Object, Scriptable): """DispatchContext node. diff --git a/src/ir/expr.cc b/src/ir/expr.cc index 7d07ab95d90a..6ab088995e9a 100644 --- a/src/ir/expr.cc +++ b/src/ir/expr.cc @@ -720,6 +720,44 @@ TVM_FFI_STATIC_INIT_BLOCK() { }); } +// LambdaExpr + +TVM_FFI_STATIC_INIT_BLOCK() { + StagingExprNode::RegisterReflection(); + LambdaExprNode::RegisterReflection(); +} + +Expr LambdaExprNode::Apply(const ffi::Array& arguments) const { + TVM_FFI_CHECK_EQ(arguments.size(), vars.size(), ValueError) << "LambdaExpr Apply arity mismatch"; + ffi::Map vmap; + for (size_t i = 0; i < vars.size(); ++i) vmap.Set(vars[i], arguments[i]); + return ffi::StructuralMap( + body, [&](const Var& var) -> Expr { return vmap.Get(var).value_or(var); }) + .cast(); +} + +LambdaExpr::LambdaExpr(ffi::Array vars, Expr body) : StagingExpr(ffi::UnsafeInit{}) { + auto n = ffi::make_object(std::move(body)); + ffi::Array types; + for (const Var& var : vars) types.push_back(var->ty); + n->ty = FuncType(types, n->body->ty); + n->vars = std::move(vars); + data_ = std::move(n); +} + +TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; + refl::GlobalDef().def("ir.LambdaExpr", + [](ffi::Array vars, Expr body) { return LambdaExpr(vars, body); }); +} + +TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; + refl::GlobalDef().def("ir.LambdaExprApply", [](LambdaExpr body, ffi::Array indices) { + return body->Apply(indices); + }); +} + // Tuple Tuple::Tuple(ffi::Array fields, Span span) : Expr(ffi::UnsafeInit{}) { ffi::Optional tuple_ty = [&]() -> ffi::Optional { diff --git a/src/s_tir/transform/lower_cross_thread_reduction.cc b/src/s_tir/transform/lower_cross_thread_reduction.cc index 098ebdf119a6..fde0ac24f7f7 100644 --- a/src/s_tir/transform/lower_cross_thread_reduction.cc +++ b/src/s_tir/transform/lower_cross_thread_reduction.cc @@ -24,6 +24,7 @@ #include #include #include +#include #include #include #include @@ -410,31 +411,35 @@ Stmt TransformReductionBlock(const SBlockRealizeNode* realize, } // Stmt 3: do cross-thread reduction { - // Step 3.1. Create the parameters to the intrinsic - ffi::Array parameters; - parameters.reserve(reduction_loops.size() + 4); - // 1-st argument: number of buffers - parameters.push_back(IntImm(PrimType::UInt(32), n_buffers)); - // Next `n_buffers` arguments: sources + // Step 3.1. Carry the reducer directly as typed staging operands. + ffi::Array combine_vars; + for (const Var& var : reducer->lhs) combine_vars.push_back(var); + for (const Var& var : reducer->rhs) combine_vars.push_back(var); + LambdaExpr combine(combine_vars, tvm::Tuple(reducer->result)); + ffi::Array values; if (it_buffers.has_value()) { for (int i = 0; i < n_buffers; ++i) { - parameters.push_back(MakeTensorLoad(it_buffers.value()[i], {IntImm::Int32(0)})); + values.push_back(MakeTensorLoad(it_buffers.value()[i], {IntImm::Int32(0)})); } } else { - parameters.insert(parameters.end(), combiner_rhs.begin(), combiner_rhs.end()); + values = combiner_rhs; } - // Next argument: predicate - parameters.push_back(IntImm::Bool(true)); - // Next `n_buffers` arguments: destinations + ffi::Array destinations; for (int i = 0; i < n_buffers; ++i) { - parameters.push_back(MakeTensorLoad(ct_buffers[i], {0})); + destinations.push_back(MakeTensorLoad(ct_buffers[i], {0})); } - // Next arguments: all the reduction threads + ffi::Array thread_axes; for (const ForNode* reduction_loop : reduction_loops) { if (reduction_loop->thread_binding.has_value()) { - parameters.push_back(reduction_loop->loop_var); + thread_axes.push_back(reduction_loop->loop_var); } } + ffi::Array parameters{combine, + tvm::Tuple(reducer->identity_element), + tvm::Tuple(values), + IntImm::Bool(true), + tvm::Tuple(destinations), + tvm::Tuple(thread_axes)}; // Step 3.2. Create the block and the block-realize. ffi::Array iter_vars{nullptr}; ffi::Array bindings{nullptr}; @@ -457,14 +462,10 @@ Stmt TransformReductionBlock(const SBlockRealizeNode* realize, /*writes=*/ct_buffer_regions, /*name_hint=*/block->name_hint + "_cross_thread", /*body=*/ - AttrStmt(/*node=*/reducer, - /*attr_key=*/s_tir::attr::reduce_scope, - /*value=*/IntImm::Int32(0), - /*body=*/ - Evaluate(Call(/*dtype=*/PrimType::Void(), - /*op=*/tirx::builtin::tvm_thread_allreduce(), - /*args=*/std::move(parameters)) - .as_or_throw()))))); + Evaluate(Call(/*dtype=*/PrimType::Void(), + /*op=*/tirx::builtin::tvm_thread_allreduce(), + /*args=*/std::move(parameters)) + .as_or_throw())))); } // Stmt 4: write cross-thread reduction result to the original buffer { diff --git a/src/script/printer/ir/prim_expr.cc b/src/script/printer/ir/prim_expr.cc index 9bf33fd3edfb..9e0ff08992ff 100644 --- a/src/script/printer/ir/prim_expr.cc +++ b/src/script/printer/ir/prim_expr.cc @@ -29,6 +29,26 @@ namespace script { namespace printer { namespace details { +ffi::Optional LambdaExprDocTranslate(DocTranslatorObj* d, ffi::AnyView input, + const ffi::Object*) { + const auto* lambda = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); + VarScope scope(d); + ffi::Array args; + ffi::Array types; + for (const Var& var : lambda->vars) { + args.push_back(VarDoc(d, var)); + types.push_back(TypeValue(d, var->ty, false)); + } + ExprDoc body = d->Translate(lambda->body).value(); + return NamespaceDoc("ir")->Attr("Lambda")->Call({ListDoc(types), LambdaDoc(args, body)}); +} + +TVM_FFI_STATIC_INIT_BLOCK() { + ffi::reflection::TypeAttrDef().attr( + kDocTranslate, FDocTranslate::FromNative<&LambdaExprDocTranslate>()); +} + namespace { template diff --git a/src/script/printer/script_printer.cc b/src/script/printer/script_printer.cc index 13be7ca2026d..04bbdcb85058 100644 --- a/src/script/printer/script_printer.cc +++ b/src/script/printer/script_printer.cc @@ -138,7 +138,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { RegisterScriptRepr(); RegisterScriptRepr(); RegisterScriptRepr(); - RegisterScriptRepr(); + RegisterScriptRepr(); RegisterScriptRepr(); RegisterScriptRepr(); RegisterScriptRepr(); diff --git a/src/tirx/ir/lambdaexpr.cc b/src/tirx/ir/lambdaexpr.cc deleted file mode 100644 index b73e6caab127..000000000000 --- a/src/tirx/ir/lambdaexpr.cc +++ /dev/null @@ -1,144 +0,0 @@ -/* - * Licensed to the Apache Software Foundation (ASF) under one - * or more contributor license agreements. See the NOTICE file - * distributed with this work for additional information - * regarding copyright ownership. The ASF licenses this file - * to you under the Apache License, Version 2.0 (the - * "License"); you may not use this file except in compliance - * with the License. You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, - * software distributed under the License is distributed on an - * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY - * KIND, either express or implied. See the License for the - * specific language governing permissions and limitations - * under the License. - */ - -/*! - * \file lambdaexpr.cc - * \brief Implementation of LambdaExpr, a reified lambda used by tile primitive ops. - */ - -#include -#include -#include - -#include -#include - -namespace tvm { -namespace tirx { - -namespace { - -ffi::Expected> LambdaExprVisit( - ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { - const auto* self = - ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); - TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ty)); - TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->WithDefRegionKind( - kTVMFFIDefRegionKindSimple, [&]() { return visitor->VisitExpected(self->vars); })); - return visitor->VisitExpected(self->pred); -} - -ffi::Expected> LambdaExprMutate(ffi::StructuralMutatorObj* mutator, - ffi::AnyView value) noexcept { - const auto* self = - ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); - // Parameters shadow substitutions from the surrounding expression. Save and restore - // their remaps so a binder rewrite stays local to this lambda's body. - std::vector saved; - for (const Var& var : self->vars) { - auto previous = mutator->VarRemapGetExpected(var); - TVM_FFI_S_MUTATE_MAYBE_EARLY_RETURN(previous); - saved.push_back(std::move(previous).value()); - } - auto result = [&]() -> ffi::Expected> { - for (const Var& var : self->vars) { - auto cleared = mutator->VarRemapSetExpected(var, nullptr); - TVM_FFI_S_MUTATE_MAYBE_EARLY_RETURN(cleared); - } - TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, - mutator->MutateExpected(self->ty)); - TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_vars, - mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() { - return mutator->MutateExpected(self->vars); - })); - auto vars = std::move(mapped_vars).ValueOrUnchanged(self->vars); - for (size_t i = 0; i < vars.size(); ++i) { - auto remap = mutator->VarRemapSetExpected(self->vars[i], vars[i]); - TVM_FFI_S_MUTATE_MAYBE_EARLY_RETURN(remap); - } - TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_pred, - mutator->MutateExpected(self->pred)); - if (vars.same_as(self->vars) && mapped_pred.UnchangedOrSameAs(self->pred) && - mapped_ty.UnchangedOrSameAs(self->ty)) { - return ffi::Unchanged(); - } - auto copy = ffi::make_object(*self); - copy->vars = std::move(vars); - copy->pred = std::move(mapped_pred).ValueOrUnchanged(self->pred); - copy->ty = std::move(mapped_ty).ValueOrUnchanged(self->ty); - return ffi::Any(std::move(copy)); - }(); - for (size_t i = 0; i < self->vars.size(); ++i) { - auto restored = mutator->VarRemapSetExpected(self->vars[i], saved[i]); - TVM_FFI_S_MUTATE_MAYBE_EARLY_RETURN(restored); - } - return result; -} - -} // namespace - -TVM_FFI_STATIC_INIT_BLOCK() { - namespace refl = tvm::ffi::reflection; - LambdaExprNode::RegisterReflection(); - refl::TypeAttrDef() - .attr(refl::type_attr::kStructuralVisit, - ffi::FStructuralVisit::FromNative<&LambdaExprVisit>()) - .attr(refl::type_attr::kStructuralMutate, - ffi::FStructuralMutate::FromNative<&LambdaExprMutate>()) - .attr(refl::type_attr::kStructuralMaybeInplaceMutate, - ffi::FStructuralMutate::FromNative<&LambdaExprMutate>()); -} - -PrimExpr LambdaExprNode::Apply(const ffi::Array& indices) const { - TVM_FFI_ICHECK_EQ(indices.size(), vars.size()); - - ffi::Map vmap; - - for (size_t i = 0; i < vars.size(); i++) { - vmap.Set(vars[i], indices[i]); - } - auto f_substitute = [&vmap](const Var& var) -> ffi::Expected> { - if (auto repl = vmap.Get(var)) return ffi::Any(*std::move(repl)); - return ffi::Unchanged(); - }; - return ffi::StructuralMap(std::move(pred), f_substitute) - .as_or_throw(); -} - -LambdaExpr::LambdaExpr(ffi::Array vars, PrimExpr pred) : Expr(ffi::UnsafeInit{}) { - auto n = ffi::make_object(std::move(pred)); - n->vars = std::move(vars); - data_ = std::move(n); -} - -TVM_FFI_STATIC_INIT_BLOCK() { - namespace refl = tvm::ffi::reflection; - refl::GlobalDef().def("tirx.LambdaExpr", - [](ffi::Array vars, PrimExpr pred) { return LambdaExpr(vars, pred); }); -} - -TVM_FFI_STATIC_INIT_BLOCK() { - namespace refl = tvm::ffi::reflection; - refl::GlobalDef().def("tirx.LambdaExprApply", [](LambdaExpr pred, ffi::Array indices) { - return pred->Apply(indices); - }); -} - -} // namespace tirx -} // namespace tvm diff --git a/src/tirx/ir/tir_visitor_with_path.cc b/src/tirx/ir/tir_visitor_with_path.cc index 77aa2e78abba..c0595a1168dc 100644 --- a/src/tirx/ir/tir_visitor_with_path.cc +++ b/src/tirx/ir/tir_visitor_with_path.cc @@ -314,7 +314,7 @@ void TIRVisitorWithPath::VisitLambda(const LambdaExprNode* op, AccessPath path) for (size_t i = 0; i < op->vars.size(); ++i) { context.push_back(WithDef(op->vars[i], path->Attr("vars")->ArrayItem(i))); } - Visit(op->pred, path->Attr("pred")); + Visit(op->body, path->Attr("body")); } void TIRVisitorWithPath::Dispatch_(const TupleNode* op, AccessPath path) { diff --git a/src/tirx/op/builtin.cc b/src/tirx/op/builtin.cc index 3309f9bded2c..1974e79cc033 100644 --- a/src/tirx/op/builtin.cc +++ b/src/tirx/op/builtin.cc @@ -23,6 +23,7 @@ * builtin intrinsic operators. */ #include +#include #include #include #include @@ -594,7 +595,13 @@ TVM_FFI_STATIC_INIT_BLOCK() { .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); OpDef("tirx.tvm_thread_allreduce") - .signature(sig::arg("size", "The size."), sig::var_args("args")) + .signature(sig::arg("combine", "The typed combining lambda."), + sig::arg("identity", "The identity values."), + sig::arg("values", "The reduction values."), + sig::arg("predicate", "Whether this thread contributes."), + sig::arg("destinations", "The destination tensor loads."), + sig::arg("thread_axes", "The reduction thread axes.")) + .set_attr("TFixedReturnType", PrimType::Void()) .set_attr("TScriptPrinterName", ffi::String("tirx.tvm_thread_allreduce")) .set_attr("TIRxOpCategory", ffi::String("builtin")) .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc index f92676e1eae4..47b713c7a628 100644 --- a/src/tirx/script/printer/expr.cc +++ b/src/tirx/script/printer/expr.cc @@ -169,20 +169,6 @@ TVM_FFI_STATIC_INIT_BLOCK() { kDocTranslate, FDocTranslate::FromNative<&IndexMapDocTranslate>()); } -ffi::Optional LambdaExprDocTranslate(DocTranslatorObj* d, ffi::AnyView input, - const ffi::Object*) { - const auto* lambda = - ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); - ffi::Array args; - for (const Var& var : lambda->vars) args.push_back(VarDoc(d, var)); - return LambdaDoc(args, d->Translate(lambda->pred).value()); -} - -TVM_FFI_STATIC_INIT_BLOCK() { - ffi::reflection::TypeAttrDef().attr( - kDocTranslate, FDocTranslate::FromNative<&LambdaExprDocTranslate>()); -} - ffi::Optional StorageSyncDocTranslate(DocTranslatorObj* d, ffi::AnyView input, const ffi::Object*) { const auto* call = diff --git a/src/tirx/transform/lower_thread_allreduce.h b/src/tirx/transform/lower_thread_allreduce.h index 81507ae1eb6b..4e03b7d92aee 100644 --- a/src/tirx/transform/lower_thread_allreduce.h +++ b/src/tirx/transform/lower_thread_allreduce.h @@ -22,11 +22,11 @@ #include #include +#include #include #include #include #include -#include #include #include #include @@ -83,18 +83,6 @@ class ThreadAllreduceBuilder final : public DialectMutator { return result; } - UnchangedOr Mutate_(const AttrStmtNode* op, InplaceMode inplace_mode) final { - if (op->attr_key == "reduce_scope") { - const te::CommReducerNode* combiner = op->node.as(); - TVM_FFI_ICHECK(combiner); - reduce_combiner_.push_back(combiner); - Stmt ret = DialectMutator::Mutate_(op, inplace_mode).ValueOrUnchanged(ffi::GetRef(op)); - reduce_combiner_.pop_back(); - return ret; - } else { - return DialectMutator::Mutate_(op, inplace_mode); - } - } UnchangedOr Mutate_(const EvaluateNode* op, InplaceMode inplace_mode) final { Stmt stmt = DialectMutator::Mutate_(op, inplace_mode).ValueOrUnchanged(ffi::GetRef(op)); op = stmt.as(); @@ -225,32 +213,41 @@ class ThreadAllreduceBuilder final : public DialectMutator { } }; + static ffi::Array ApplyCombiner(const LambdaExpr& combiner, + const ffi::Array& lhs, + const ffi::Array& rhs) { + ffi::Array arguments; + for (const PrimExpr& value : lhs) arguments.push_back(value); + for (const PrimExpr& value : rhs) arguments.push_back(value); + return builtin::GetAllreduceFields(combiner->Apply(arguments)).Map([](const Expr& value) { + return value.as_or_throw(); + }); + } + // make allreduce. Stmt MakeAllreduce(const CallNode* call) { - TVM_FFI_ICHECK(!reduce_combiner_.empty()); - const te::CommReducerNode* combiner = reduce_combiner_.back(); - size_t size = combiner->result.size(); - - const IntImmNode* size_of_args = call->args[0].as(); - TVM_FFI_ICHECK(size_of_args) << call->args[0]->GetTypeKey(); - TVM_FFI_ICHECK_EQ(size, size_of_args->value); - ffi::Array inits = combiner->identity_element; + LambdaExpr combiner = call->args[0].as_or_throw(); + ffi::Array inits = builtin::GetAllreduceFields(call->args[1]); + ffi::Array inputs = builtin::GetAllreduceFields(call->args[2]); + ffi::Array destinations = builtin::GetAllreduceFields(call->args[4]); + ffi::Array thread_axes = builtin::GetAllreduceFields(call->args[5]); + size_t size = inputs.size(); std::vector values; values.reserve(size); std::vector dtypes; dtypes.reserve(size); - PrimExpr cond = call->args[size + 1].as_or_throw(); + PrimExpr cond = call->args[3].as_or_throw(); for (size_t idx = 0; idx < size; ++idx) { - values.push_back(call->args[1 + idx].as_or_throw()); + values.push_back(inputs[idx].as_or_throw()); if (!is_one(cond)) { - values[idx] = Select(cond, values[idx], inits[idx]); + values[idx] = Select(cond, values[idx], inits[idx].as_or_throw()); } dtypes.push_back(values[idx].ty()); } std::vector buffers; buffers.reserve(size); for (size_t idx = 0; idx < size; ++idx) { - PrimExpr arg = call->args[2 + size + idx].as_or_throw(); + PrimExpr arg = destinations[idx].as_or_throw(); // Loads from boolean buffers may have cast nodes inserted by // earlier passes. if (auto cast = arg.as()) { @@ -260,8 +257,8 @@ class ThreadAllreduceBuilder final : public DialectMutator { } std::unordered_set reduce_set; - for (size_t i = 2 + 2 * size; i < call->args.size(); ++i) { - auto var = call->args[i].as(); + for (const Expr& axis : thread_axes) { + auto var = axis.as(); const VarNode* v = var.has_value() ? var.value().get() : nullptr; // The simply optimization replace a iteration variable with a constant // when extent of the iteration is 1. As threaded IterVar always started from 0, @@ -269,14 +266,14 @@ class ThreadAllreduceBuilder final : public DialectMutator { if (v) { reduce_set.insert(v); } else { - TVM_FFI_ICHECK(call->args[i].as() && call->args[i].as()->value == 0) - << "arg" << i << "should be a VarNode or IntImmNode"; + TVM_FFI_ICHECK(axis.as() && axis.as()->value == 0) + << "Reduction thread axis should be a VarNode or zero IntImmNode"; } } size_t nmatch = 0; std::vector vred, vpar; - std::map> thread_axes; + std::map> thread_axes_by_dim; for (const RegionStmtNode* launch : thread_extents_) { ThreadEntry e; IterVar iv(Range(), launch->body_params[0].as_or_throw(), IterVarType::kThreadIndex, @@ -291,7 +288,8 @@ class ThreadAllreduceBuilder final : public DialectMutator { e.extent = ptr->value.as().value(); bool is_reduce = reduce_set.count(iv->var.get()); nmatch += is_reduce; - auto [it, inserted] = thread_axes.emplace(e.scope.dim_index, std::make_pair(e, is_reduce)); + auto [it, inserted] = + thread_axes_by_dim.emplace(e.scope.dim_index, std::make_pair(e, is_reduce)); if (!inserted) { TVM_FFI_ICHECK_EQ(it->second.first.extent, e.extent) << "Incompatible extents for nested bindings of " << iv->thread_tag; @@ -303,7 +301,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { } } } - for (const auto& [dim, entry] : thread_axes) { + for (const auto& [dim, entry] : thread_axes_by_dim) { if (entry.first.extent != 1) (entry.second ? vred : vpar).push_back(entry.first); } TVM_FFI_ICHECK_EQ(nmatch, reduce_set.size()) @@ -537,7 +535,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { std::pair, std::vector> MakeWarpAllreduce( std::vector src_values, // std::vector dtypes, // - const te::CommReducerNode* combiner, // + const LambdaExpr& combiner, // PrimExpr reduce_index, int reduce_extent, // PrimExpr group_index, // PrimExpr mask, ffi::Optional predicate, // @@ -622,7 +620,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { } // Do reductions. - ffi::Array ret = (*combiner)(a, b); + ffi::Array ret = ApplyCombiner(combiner, a, b); // Store the reduction result to itself. std::vector stores; @@ -656,7 +654,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { } // make allreduce. - Stmt MakeBufAllreduce(const te::CommReducerNode* combiner, const std::vector& dtypes, + Stmt MakeBufAllreduce(const LambdaExpr& combiner, const std::vector& dtypes, const ffi::Array& shared_bufs, PrimExpr reduce_index, PrimExpr group_index, int reduce_extent, int group_extent, int contiguous_reduce_extent) { @@ -683,7 +681,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { TVM_FFI_ICHECK_EQ(a_load.ty(), dtypes[i]); a.push_back(a_load); } - ffi::Array ret = (*combiner)(a, b); + ffi::Array ret = ApplyCombiner(combiner, a, b); return ret; }; auto fstore = [&](const ffi::Array& ret) { @@ -904,7 +902,6 @@ class ThreadAllreduceBuilder final : public DialectMutator { // surrounding scope of thread extent. std::vector thread_extents_; - std::vector reduce_combiner_; // The load remap std::unordered_map load_remap_; // Internal analyzer diff --git a/src/tirx/transform/unsupported_dtype_legalize.cc b/src/tirx/transform/unsupported_dtype_legalize.cc index bc63f55163cf..a62965845846 100644 --- a/src/tirx/transform/unsupported_dtype_legalize.cc +++ b/src/tirx/transform/unsupported_dtype_legalize.cc @@ -24,8 +24,8 @@ #include #include #include +#include #include -#include #include #include #include @@ -144,7 +144,7 @@ class FP8ComputeLegalizePlanner : public ComputeLegalizePlanner { PrimExpr origin_b = \ PromoteToTarget(this->Mutate(op->b, inplace_mode).ValueOrUnchanged(op->b)); \ \ - if (origin_a.same_as(op->a) && origin_b.same_as(op->b)) { \ + if (origin_a.same_as(op->a) && origin_b.same_as(op->b) && !MatchType(op->ty)) { \ return ffi::Unchanged(); \ } else { \ return FUNC(origin_a, origin_b); \ @@ -201,7 +201,7 @@ class ComputeLegalizer : public StmtExprMutator { PrimExpr false_value = PromoteToTarget( this->Mutate(op->false_value, inplace_mode).ValueOrUnchanged(op->false_value)); if (condition_unchanged && true_value.same_as(op->true_value) && - false_value.same_as(op->false_value)) { + false_value.same_as(op->false_value) && !MatchType(op->ty)) { return ffi::Unchanged(); } else { return prim::Select(condition, true_value, false_value); @@ -211,7 +211,7 @@ class ComputeLegalizer : public StmtExprMutator { UnchangedOr Mutate_(const prim::BroadcastNode* op, InplaceMode inplace_mode) final { PrimExpr value = PromoteToTarget(this->Mutate(op->value, inplace_mode).ValueOrUnchanged(op->value)); - if (value.same_as(op->value)) { + if (value.same_as(op->value) && !MatchType(op->ty)) { return ffi::Unchanged(); } else { return prim::Broadcast(value, op->lanes); @@ -222,7 +222,7 @@ class ComputeLegalizer : public StmtExprMutator { auto vectors = op->vectors.Map([this](const PrimExpr& value) { return PromoteToTarget(Mutate(value).ValueOrUnchanged(value)); }); - if (vectors.same_as(op->vectors)) { + if (vectors.same_as(op->vectors) && !MatchType(op->ty)) { return ffi::Unchanged(); } else { return prim::Shuffle(vectors, op->indices); @@ -230,6 +230,9 @@ class ComputeLegalizer : public StmtExprMutator { } UnchangedOr Mutate_(const CallNode* op, InplaceMode inplace_mode) final { + if (op->op.same_as(builtin::tvm_thread_allreduce())) { + return LegalizeThreadAllreduce(op); + } if (op->op.same_as(builtin::alloc_tensor()) || op->op.same_as(builtin::decl_tensor())) { Call call = StmtExprMutator::Mutate_(op, inplace_mode) .ValueOrUnchanged(ffi::GetRef(op)) @@ -403,40 +406,6 @@ class ComputeLegalizer : public StmtExprMutator { if (mapped != nullptr) { return AttrStmt(mapped.as_or_throw(), op->attr_key, op->value, op->body); } - } else if (auto reducer = op->node.as()) { - auto reducer_mode = op->unique() && reducer->unique() ? inplace_mode : InplaceMode::kDisallow; - auto legalized_identity_elements = Mutate(reducer->identity_element, reducer_mode) - .as_or_throw>>() - .ValueOrUnchanged(reducer->identity_element); - - // Remap input variables - for (size_t i = 0; i < legalized_identity_elements.size(); i++) { - Var lhs_var = reducer->lhs[i]; - if (lhs_var->ty.as_or_throw() != legalized_identity_elements[i].ty()) { - VarRemapSet(lhs_var, lhs_var.CopyWithDType(legalized_identity_elements[i].ty())); - } - Var rhs_var = reducer->rhs[i]; - if (rhs_var->ty.as_or_throw() != legalized_identity_elements[i].ty()) { - VarRemapSet(rhs_var, rhs_var.CopyWithDType(legalized_identity_elements[i].ty())); - } - } - - auto legalized_results = Mutate(reducer->result, reducer_mode) - .as_or_throw>>() - .ValueOrUnchanged(reducer->result); - - auto legalized_lhs = reducer->lhs.Map([this](PrimVar var) { - auto mapped = VarRemapGet(var); - return mapped == nullptr ? var : mapped.as_or_throw(); - }); - - auto legalized_rhs = reducer->rhs.Map([this](PrimVar var) { - auto mapped = VarRemapGet(var); - return mapped == nullptr ? var : mapped.as_or_throw(); - }); - return AttrStmt(te::CommReducer(legalized_lhs, legalized_rhs, legalized_results, - legalized_identity_elements, reducer->span), - op->attr_key, op->value, op->body); } return ret; } @@ -453,6 +422,40 @@ class ComputeLegalizer : public StmtExprMutator { } private: + // Preserve scalar operands and explicit Tuple grouping while promoting computation. + Expr LegalizeThreadAllreduce(const CallNode* op) { + LambdaExpr combine = op->args[0].as_or_throw(); + auto map_operand = [](const Expr& operand, const auto& transform) -> Expr { + if (const auto* tuple = operand.as()) { + return tvm::Tuple(tuple->fields.Map(transform), operand->span); + } + return transform(operand); + }; + auto promote = [this](const Expr& value) -> Expr { + return PromoteToTarget(Mutate(value).ValueOrUnchanged(value).as_or_throw()); + }; + Expr identity = map_operand(op->args[1], promote); + Expr values = map_operand(op->args[2], promote); + ffi::Array value_fields = builtin::GetAllreduceFields(values); + ffi::Array vars; + ffi::Array arguments; + for (size_t i = 0; i < combine->vars.size(); ++i) { + Var promoted = combine->vars[i].CopyWithDType( + value_fields[i % value_fields.size()]->ty.as_or_throw()); + vars.push_back(promoted); + arguments.push_back(promoted); + } + LambdaExpr legalized_combine(vars, map_operand(combine->Apply(arguments), promote)); + auto mutate = [this](const Expr& value) { return Mutate(value).ValueOrUnchanged(value); }; + Expr predicate = mutate(op->args[3]); + // Destinations are lvalues: remap promoted allocations without adding compute casts. + Expr destinations = map_operand(op->args[4], mutate); + Expr axes = map_operand(op->args[5], mutate); + return Call(PrimType::Void(), op->op, + {legalized_combine, identity, values, predicate, destinations, axes}, op->attrs, + op->ty_args, op->span); + } + /*! * \brief promote value to target datatype F16/F32 and keep other values unchanged. * \param value The input value. diff --git a/tests/python/codegen/test_gpu_codegen_allreduce.py b/tests/python/codegen/test_gpu_codegen_allreduce.py index bbbbf62e5727..18f8f9d31417 100644 --- a/tests/python/codegen/test_gpu_codegen_allreduce.py +++ b/tests/python/codegen/test_gpu_codegen_allreduce.py @@ -27,9 +27,9 @@ def _reduce_module(d1, d2, d3, is_max=False): - reducer = T.comm_reducer( - (lambda x, y: T.max(x, y)) if is_max else (lambda x, y: x + y), - [T.float32(-3.4028234663852886e38 if is_max else 0)], + combine = T.Lambda( + [T.float32, T.float32], + (lambda x, y: (T.max(x, y),)) if is_max else (lambda x, y: (x + y,)), ) @I.ir_module @@ -41,10 +41,14 @@ def main(A: T.Tensor((1, d1, d2, d3), "float32"), B: T.Tensor((1, d1, d2), "floa for k in T.thread_binding(d2, thread="threadIdx.y"): for l in T.thread_binding(d3, thread="threadIdx.x"): reduced = T.alloc_tensor((1,), "float32", scope="local") - with T.attr(reducer, "reduce_scope", 0): - T.tvm_thread_allreduce( - T.uint32(1), A[i, j, k, l], True, reduced[0], l - ) + T.tvm_thread_allreduce( + combine, + (T.float32(-3.4028234663852886e38 if is_max else 0),), + (A[i, j, k, l],), + True, + (reduced[0],), + (l,), + ) if l == 0: B[i, j, k] = reduced[0] diff --git a/tests/python/codegen/test_target_codegen_cuda.py b/tests/python/codegen/test_target_codegen_cuda.py index 997798ae393f..8b33848e713f 100644 --- a/tests/python/codegen/test_target_codegen_cuda.py +++ b/tests/python/codegen/test_target_codegen_cuda.py @@ -533,10 +533,14 @@ def main(A: T.Tensor((n, m)), B: T.Tensor((n,))): for m_1 in range((m + nthd - 1) // nthd): if m_0 * ((m + nthd - 1) // nthd) + m_1 < m: partial[0] = partial[0] + A[i, m_0 * ((m + nthd - 1) // nthd) + m_1] - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), "reduce_scope", 0 - ): - T.tvm_thread_allreduce(T.uint32(1), partial[0], True, reduced[0], m_0) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (partial[0],), + True, + (reduced[0],), + (m_0,), + ) if m_0 == 0: B[i] = reduced[0] @@ -608,14 +612,14 @@ def main(A: T.Tensor((n, k0, k1)), B: T.Tensor((n,))): k1_0 * ((k1 + nthdy - 1) // nthdy) + k1_1, ] ) - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - 0, - ): - T.tvm_thread_allreduce( - T.uint32(1), partial[0], True, reduced[0], k0_0, k1_0 - ) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (partial[0],), + True, + (reduced[0],), + (k0_0, k1_0), + ) if k0_0 == 0 and k1_0 == 0: B[i] = reduced[0] diff --git a/tests/python/codegen/test_target_codegen_cuda_fp8.py b/tests/python/codegen/test_target_codegen_cuda_fp8.py index ec330668f8cd..ff53b524bc63 100644 --- a/tests/python/codegen/test_target_codegen_cuda_fp8.py +++ b/tests/python/codegen/test_target_codegen_cuda_fp8.py @@ -574,12 +574,14 @@ def moe_dequantize_gemv( ) * T.Cast("float16", scale[0]) ) - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float16(0)]), "reduce_scope", 0 - ): - T.tvm_thread_allreduce( - T.uint32(1), partial[0], True, reduced[0], reduction - ) + T.tvm_thread_allreduce( + T.Lambda([T.float16, T.float16], lambda x, y: (x + y,)), + (T.float16(0),), + (partial[0],), + True, + (reduced[0],), + (reduction,), + ) if reduction == 0: output[expert, block * 4 + spatial] = reduced[0] diff --git a/tests/python/s_tir/script/test_s_tir_script_printer.py b/tests/python/s_tir/script/test_s_tir_script_printer.py index 5ab098fdfe02..e3787402c568 100644 --- a/tests/python/s_tir/script/test_s_tir_script_printer.py +++ b/tests/python/s_tir/script/test_s_tir_script_printer.py @@ -1318,18 +1318,16 @@ def comm_reducer_single_reduce_group( for i in T.serial(0, 128): threadIdx_x = T.launch_thread("threadIdx.x", 128) reduce_temp0 = T.alloc_tensor((1,), scope="local") - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), "reduce_scope", T.int32(0) - ): - T.evaluate( - T.tvm_thread_allreduce( - T.uint32(1), - A[i * 128 + threadIdx_x], - True, - reduce_temp0.data, - threadIdx_x, - ) + T.evaluate( + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (A[i * 128 + threadIdx_x],), + True, + (reduce_temp0[0],), + (threadIdx_x,), ) + ) return comm_reducer_single_reduce_group @@ -1343,27 +1341,24 @@ def comm_reducer_multiple_reduce_groups( for i in T.serial(0, 128): threadIdx_x = T.launch_thread("threadIdx.x", 128) - reduce_temp0 = T.alloc_tensor((1,), scope="local") - with T.attr( - T.comm_reducer( - lambda x0, x1, y0, y1: ( - T.Select((x1 >= y1), x0, y0), - T.Select((x1 >= y1), x1, y1), + reduce_temp0 = T.alloc_tensor((1,), "int32", scope="local") + reduce_temp1 = T.alloc_tensor((1,), "float32", scope="local") + T.evaluate( + T.tvm_thread_allreduce( + T.Lambda( + [T.int32, T.float32, T.int32, T.float32], + lambda x0, x1, y0, y1: ( + T.Select(x1 >= y1, x0, y0), + T.Select(x1 >= y1, x1, y1), + ), ), - [T.int32(-1), T.min_value("float32")], - ), - "reduce_scope", - T.int32(0), - ): - T.evaluate( - T.tvm_thread_allreduce( - T.uint32(1), - A[i * 128 + threadIdx_x], - True, - reduce_temp0.data, - threadIdx_x, - ) + (T.int32(-1), T.min_value("float32")), + (i * 128 + threadIdx_x, A[i * 128 + threadIdx_x]), + True, + (reduce_temp0[0], reduce_temp1[0]), + (threadIdx_x,), ) + ) return comm_reducer_multiple_reduce_groups @@ -1386,32 +1381,26 @@ def multiple_commreducer() -> None: ) for ax0_1 in T.thread_binding(0, 32, thread="threadIdx.x"): with Ts.sblock("T_softmax_maxelem_cross_thread_reduction"): - T.attr( - T.comm_reducer(lambda x, y: T.max(x, y), [T.min_value("float32")]), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), - normal_reduce_temp0[0], + T.Lambda([T.float32, T.float32], lambda x, y: (T.max(x, y),)), + (T.min_value("float32"),), + (normal_reduce_temp0[0],), True, - reduce_temp0.data, - ax0_1, + (reduce_temp0[0],), + (ax0_1,), ) ) for ax0_1 in T.thread_binding(0, 32, thread="threadIdx.x"): with Ts.sblock("T_softmax_expsum_cross_thread_reduction"): - T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), "reduce_scope", T.int32(0) - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), - normal_reduce_temp1[0], + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (normal_reduce_temp1[0],), True, - reduce_temp1.data, - ax0_1, + (reduce_temp1[0],), + (ax0_1,), ) ) @@ -1884,30 +1873,25 @@ def func( T.tvm_warp_activemask(), A_warp_1[0], threadIdx_x % 4 * 8 + threadIdx_x // 4, 32, 32 ) + T.float32(1) red_buf0_1 = T.decl_tensor((1,), data=red_buf0.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - mask = T.alloc_tensor((1,), "uint32", scope="local") - t0 = T.alloc_tensor((1,), scope="local") - red_buf0_1[0] = A_warp_1[0] - mask_1 = T.decl_tensor((1,), "uint32", data=mask.data, scope="local") - mask_1[0] = T.tvm_warp_activemask() - t0_1 = T.decl_tensor((1,), data=t0.data, scope="local") - t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 16, 32, 32) - red_buf0_1[0] = red_buf0_1[0] + t0_1[0] - t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 8, 32, 32) - red_buf0_1[0] = red_buf0_1[0] + t0_1[0] - t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 4, 32, 32) - red_buf0_1[0] = red_buf0_1[0] + t0_1[0] - t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 2, 32, 32) - red_buf0_1[0] = red_buf0_1[0] + t0_1[0] - t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 1, 32, 32) - red_buf0_1[0] = red_buf0_1[0] + t0_1[0] - red_buf0_1[0] = T.tvm_warp_shuffle(mask_1[0], red_buf0_1[0], 0, 32, 32) - # NOTE(Zihao): test tvm_warp_shuffle_up - red_buf0_1[0] = T.tvm_warp_shuffle_up(mask_1[0], red_buf0_1[0], 0, 32, 32) + mask = T.alloc_tensor((1,), "uint32", scope="local") + t0 = T.alloc_tensor((1,), scope="local") + red_buf0_1[0] = A_warp_1[0] + mask_1 = T.decl_tensor((1,), "uint32", data=mask.data, scope="local") + mask_1[0] = T.tvm_warp_activemask() + t0_1 = T.decl_tensor((1,), data=t0.data, scope="local") + t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 16, 32, 32) + red_buf0_1[0] = red_buf0_1[0] + t0_1[0] + t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 8, 32, 32) + red_buf0_1[0] = red_buf0_1[0] + t0_1[0] + t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 4, 32, 32) + red_buf0_1[0] = red_buf0_1[0] + t0_1[0] + t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 2, 32, 32) + red_buf0_1[0] = red_buf0_1[0] + t0_1[0] + t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 1, 32, 32) + red_buf0_1[0] = red_buf0_1[0] + t0_1[0] + red_buf0_1[0] = T.tvm_warp_shuffle(mask_1[0], red_buf0_1[0], 0, 32, 32) + # NOTE(Zihao): test tvm_warp_shuffle_up + red_buf0_1[0] = T.tvm_warp_shuffle_up(mask_1[0], red_buf0_1[0], 0, 32, 32) if threadIdx_x == 0: C_1 = T.decl_tensor((1,), data=C) C_1[0] = red_buf0_1[0] @@ -2322,18 +2306,14 @@ def lowered_loop_split( with Ts.sblock("B_cross_thread_reduction"): Ts.reads([normal_reduce_temp0[0]]) Ts.writes([reduce_temp0[0]]) - T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), - normal_reduce_temp0[0], + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (normal_reduce_temp0[0],), True, - reduce_temp0.data, - ki, + (reduce_temp0[0],), + (ki,), ) ) with Ts.sblock("B_write_back"): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lift_thread_binding.py b/tests/python/s_tir/transform/test_s_tir_transform_lift_thread_binding.py index c15e9c1ce960..fcd1df9d78da 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lift_thread_binding.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lift_thread_binding.py @@ -64,8 +64,7 @@ def before(A: T.Tensor((32, 1, 128)), B: T.Tensor((32, n, 128)), C: T.Tensor((32 with Ts.sblock("NT_matmul_cross_thread"): Ts.reads(in_thread_D_local[0]) Ts.writes(cross_thread_D_local[0]) - T.attr(T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0)) - T.tvm_thread_allreduce(T.uint32(1), in_thread_D_local[0], T.bool(True), cross_thread_D_local[0], ax0_fused) + T.tvm_thread_allreduce(T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), (T.float32(0),), (in_thread_D_local[0],), T.bool(True), (cross_thread_D_local[0],), (ax0_fused,)) with Ts.sblock("NT_matmul_write_back"): Ts.where(ax0_fused == 0) Ts.reads(cross_thread_D_local[0]) @@ -118,8 +117,7 @@ def expected(A: T.Tensor((32, 1, 128), "float32"), B: T.Tensor((32, n, 128)), C: with Ts.sblock("NT_matmul_cross_thread"): Ts.reads(in_thread_D_local[0]) Ts.writes(cross_thread_D_local[0]) - T.attr(T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0)) - T.tvm_thread_allreduce(T.uint32(1), in_thread_D_local[0], T.bool(True), cross_thread_D_local[0], threadIdx_x) + T.tvm_thread_allreduce(T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), (T.float32(0),), (in_thread_D_local[0],), T.bool(True), (cross_thread_D_local[0],), (threadIdx_x,)) with Ts.sblock("NT_matmul_write_back"): Ts.where(threadIdx_x == 0) Ts.reads(cross_thread_D_local[0]) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py b/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py index 5543f3c7d9ac..43256bb119d5 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py @@ -81,18 +81,14 @@ def lowered_loop_split( with Ts.sblock("B_cross_thread_reduction"): Ts.reads([normal_reduce_temp0[0]]) Ts.writes([reduce_temp0[0]]) - T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), - normal_reduce_temp0[0], + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (normal_reduce_temp0[0],), True, - reduce_temp0[0], - ki, + (reduce_temp0[0],), + (ki,), ) ) with Ts.sblock("B_write_back"): @@ -130,12 +126,16 @@ def lowered_no_normal_reduction( vi, vk = Ts.axis.remap("SR", [i, k]) Ts.reads([A[vi, vk]]) Ts.writes([reduce_temp0[0]]) - T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), + T.evaluate( + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (A[vi, vk],), + True, + (reduce_temp0[0],), + (k,), + ) ) - T.evaluate(T.tvm_thread_allreduce(T.uint32(1), A[vi, vk], True, reduce_temp0[0], k)) with Ts.sblock("B_write_back"): vi = Ts.axis.spatial(128, i) Ts.where(k == 0) @@ -175,14 +175,14 @@ def lowered_two_bound_loops( vk = Ts.axis.reduce(128, ko * 32 + ki) Ts.reads([A[vi, vk]]) Ts.writes([reduce_temp0[0]]) - T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), A[vi, vk], True, reduce_temp0[0], ko, ki + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (A[vi, vk],), + True, + (reduce_temp0[0],), + (ko, ki), ) ) with Ts.sblock("B_write_back"): @@ -252,18 +252,14 @@ def lowered_multiple_blocks_under_reduction_loop( with Ts.sblock("B_cross_thread_reduction"): Ts.reads([normal_reduce_temp0[0]]) Ts.writes([reduce_temp0[0]]) - T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), - normal_reduce_temp0[0], + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (normal_reduce_temp0[0],), True, - reduce_temp0[0], - k0o, + (reduce_temp0[0],), + (k0o,), ) ) with Ts.sblock("B_write_back"): @@ -314,18 +310,14 @@ def lowered_with_block_predicate( with Ts.sblock("B_cross_thread_reduction"): Ts.reads([normal_reduce_temp0[0]]) Ts.writes([reduce_temp0[0]]) - T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), - normal_reduce_temp0[0], + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (normal_reduce_temp0[0],), True, - reduce_temp0[0], - ki, + (reduce_temp0[0],), + (ki,), ) ) with Ts.sblock("B_write_back"): @@ -414,20 +406,14 @@ def lowered_single_reduction_loop_with_block_predicate( with Ts.sblock("T_softmax_maxelem_cross_thread"): Ts.reads(in_thread_0[0]) Ts.writes(cross_thread_0[0]) - T.attr( - T.comm_reducer( - lambda x, y: T.max(x, y), [T.float32(-3.4028234663852886e38)] - ), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), - in_thread_0[0], + T.Lambda([T.float32, T.float32], lambda x, y: (T.max(x, y),)), + (T.float32(-3.4028234663852886e38),), + (in_thread_0[0],), True, - cross_thread_0[0], - ax1_1, + (cross_thread_0[0],), + (ax1_1,), ) ) with Ts.sblock("T_softmax_maxelem_write_back"): @@ -455,18 +441,14 @@ def lowered_single_reduction_loop_with_block_predicate( with Ts.sblock("T_softmax_expsum_cross_thread"): Ts.reads(in_thread_1[0]) Ts.writes(cross_thread_1[0]) - T.attr( - T.comm_reducer(lambda x_1, y_1: x_1 + y_1, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), - in_thread_1[0], + T.Lambda([T.float32, T.float32], lambda x_1, y_1: (x_1 + y_1,)), + (T.float32(0),), + (in_thread_1[0],), True, - cross_thread_1[0], - ax1_1, + (cross_thread_1[0],), + (ax1_1,), ) ) with Ts.sblock("T_softmax_expsum_write_back"): @@ -683,17 +665,13 @@ def lowered_spatial_reduction_with_shared_prefetch( with Ts.sblock("B_cross_thread"): Ts.reads(in_thread_C_local[0]) Ts.writes(cross_thread_C_local[0]) - T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.tvm_thread_allreduce( - T.uint32(1), - in_thread_C_local[0], + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (in_thread_C_local[0],), T.bool(True), - cross_thread_C_local[0], - ax2_1_1_fused, + (cross_thread_C_local[0],), + (ax2_1_1_fused,), ) with Ts.sblock("B_write_back"): v0 = Ts.axis.spatial(128, ax0_0_ax1_0_fused // 16 * 8 + ax0_1_ax1_1_fused // 8) @@ -755,13 +733,13 @@ def lowered_reduction_spatial_loop_predicate( with Ts.sblock("block_cross_thread"): Ts.reads(in_thread_B[0]) Ts.writes(cross_thread_B[0]) - T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.tvm_thread_allreduce( - T.uint32(1), in_thread_B[0], T.bool(True), cross_thread_B[0], k_1 + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (in_thread_B[0],), + T.bool(True), + (cross_thread_B[0],), + (k_1,), ) with Ts.sblock("block_write_back"): vi = Ts.axis.spatial(2, i_0 * 16 + i_1) @@ -913,12 +891,16 @@ def lowered_reducer_max( vi, vk = Ts.axis.remap("SR", [i, k]) Ts.reads([A[vi, vk]]) Ts.writes([reduce_temp0[0]]) - T.attr( - T.comm_reducer(lambda x, y: T.max(x, y), [T.min_value("float32")]), - "reduce_scope", - T.int32(0), + T.evaluate( + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (T.max(x, y),)), + (T.min_value("float32"),), + (A[vi, vk],), + True, + (reduce_temp0[0],), + (k,), + ) ) - T.evaluate(T.tvm_thread_allreduce(T.uint32(1), A[vi, vk], True, reduce_temp0[0], k)) with Ts.sblock("B_write_back"): vi = Ts.axis.spatial(128, i) Ts.where(k == 0) @@ -950,12 +932,16 @@ def lowered_zero_rank_buffer( vk = Ts.axis.reduce(128, k) Ts.reads([A[vk]]) Ts.writes([reduce_temp0[0]]) - T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), + T.evaluate( + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (A[vk],), + True, + (reduce_temp0[0],), + (k,), + ) ) - T.evaluate(T.tvm_thread_allreduce(T.uint32(1), A[vk], True, reduce_temp0[0], k)) with Ts.sblock("B_write_back"): Ts.reads([reduce_temp0[0]]) Ts.writes([B[()]]) @@ -1131,18 +1117,14 @@ def lowered_softmax( with Ts.sblock("T_softmax_maxelem_cross_thread_reduction"): Ts.reads([normal_reduce_temp0[0]]) Ts.writes([reduce_temp0[0]]) - T.attr( - T.comm_reducer(lambda x, y: T.max(x, y), [T.min_value("float32")]), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), - normal_reduce_temp0[0], + T.Lambda([T.float32, T.float32], lambda x, y: (T.max(x, y),)), + (T.min_value("float32"),), + (normal_reduce_temp0[0],), True, - reduce_temp0[0], - ax0_1, + (reduce_temp0[0],), + (ax0_1,), ) ) with Ts.sblock("T_softmax_maxelem_write_back"): @@ -1173,18 +1155,14 @@ def lowered_softmax( with Ts.sblock("T_softmax_expsum_cross_thread_reduction"): Ts.reads([normal_reduce_temp1[0]]) Ts.writes([reduce_temp1[0]]) - T.attr( - T.comm_reducer(lambda x_1, y_1: x_1 + y_1, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), - normal_reduce_temp1[0], + T.Lambda([T.float32, T.float32], lambda x_1, y_1: (x_1 + y_1,)), + (T.float32(0),), + (normal_reduce_temp1[0],), True, - reduce_temp1[0], - ax0_1, + (reduce_temp1[0],), + (ax0_1,), ) ) with Ts.sblock("T_softmax_expsum_write_back"): @@ -1279,26 +1257,20 @@ def lowered_argmax_split( with Ts.sblock("argmax_cross_thread"): Ts.reads(in_thread_argmax_v0[0], in_thread_argmax_v1[0]) Ts.writes(cross_thread_argmax_v0[0], cross_thread_argmax_v1[0]) - T.attr( - T.comm_reducer( - lambda x0, x1, y0, y1: ( - T.Select(x1 >= y1, x0, y0), - T.Select(x1 >= y1, x1, y1), - ), - [-1, T.float32(-3.4028234663852886e38)], - ), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(2), - in_thread_argmax_v0[0], - in_thread_argmax_v1[0], + T.Lambda( + [T.int32, T.float32, T.int32, T.float32], + lambda x0, x1, y0, y1: ( + T.Select(x1 >= y1, x0, y0), + T.Select(x1 >= y1, x1, y1), + ), + ), + (-1, T.float32(-3.4028234663852886e38)), + (in_thread_argmax_v0[0], in_thread_argmax_v1[0]), True, - cross_thread_argmax_v0[0], - cross_thread_argmax_v1[0], - i1_1, + (cross_thread_argmax_v0[0], cross_thread_argmax_v1[0]), + (i1_1,), ) ) with Ts.sblock("argmax_write_back"): @@ -1374,26 +1346,20 @@ def lowered_argmin_split_init_update_reordered( with Ts.sblock("argmin_cross_thread"): Ts.reads(in_thread_argmin_v0[0], in_thread_argmin_v1[0]) Ts.writes(cross_thread_argmin_v0[0], cross_thread_argmin_v1[0]) - T.attr( - T.comm_reducer( - lambda x0, x1, y0, y1: ( - T.Select(x1 <= y1, x0, y0), - T.Select(x1 <= y1, x1, y1), - ), - [-1, T.float32(3.4028234663852886e38)], - ), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(2), - in_thread_argmin_v0[0], - in_thread_argmin_v1[0], + T.Lambda( + [T.int32, T.float32, T.int32, T.float32], + lambda x0, x1, y0, y1: ( + T.Select(x1 <= y1, x0, y0), + T.Select(x1 <= y1, x1, y1), + ), + ), + (-1, T.float32(3.4028234663852886e38)), + (in_thread_argmin_v0[0], in_thread_argmin_v1[0]), True, - cross_thread_argmin_v0[0], - cross_thread_argmin_v1[0], - i1_1, + (cross_thread_argmin_v0[0], cross_thread_argmin_v1[0]), + (i1_1,), ) ) with Ts.sblock("argmin_write_back"): @@ -1501,22 +1467,17 @@ def lowered_layer_norm_tuple_sum( with Ts.sblock("data_red_temp_cross_thread"): Ts.reads(in_thread_data_red_temp_v0[0], in_thread_data_red_temp_v1[0]) Ts.writes(cross_thread_data_red_temp_v0[0], cross_thread_data_red_temp_v1[0]) - T.attr( - T.comm_reducer( - lambda x0, x1, y0, y1: (x0 + y0, x1 + y1), [T.float32(0), T.float32(0)] - ), - "reduce_scope", - T.int32(0), - ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(2), - in_thread_data_red_temp_v0[0], - in_thread_data_red_temp_v1[0], + T.Lambda( + [T.float32, T.float32, T.float32, T.float32], + lambda x0, x1, y0, y1: (x0 + y0, x1 + y1), + ), + (T.float32(0), T.float32(0)), + (in_thread_data_red_temp_v0[0], in_thread_data_red_temp_v1[0]), True, - cross_thread_data_red_temp_v0[0], - cross_thread_data_red_temp_v1[0], - i1_1, + (cross_thread_data_red_temp_v0[0], cross_thread_data_red_temp_v1[0]), + (i1_1,), ) ) with Ts.sblock("data_red_temp_write_back"): @@ -1580,13 +1541,13 @@ def lowered_thread_broadcast_1(A: T.Tensor((256, 256), "float32"), B: T.Tensor(( vi, vk = Ts.axis.remap("SR", [i, k]) Ts.reads(A[vi, vk]) Ts.writes(cross_thread_temp_local[0]) - T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.tvm_thread_allreduce( - T.uint32(1), A[vi, vk], T.bool(True), cross_thread_temp_local[0], k + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A[vi, vk],), + T.bool(True), + (cross_thread_temp_local[0],), + (k,), ) with Ts.sblock("sum_write_back"): vi = Ts.axis.spatial(256, i) @@ -1694,8 +1655,7 @@ def lowered_thread_broadcast_2(lv1605: T.Tensor((T.int64(1), T.int64(32), T.int6 with Ts.sblock("NT_matmul_cross_thread"): Ts.reads(in_thread_var_NT_matmul_intermediate_local[0]) Ts.writes(cross_thread_var_NT_matmul_intermediate_local[0]) - T.attr(T.comm_reducer(lambda x0, y0: x0 + y0, [T.float16(0)]), "reduce_scope", T.int32(0)) - T.tvm_thread_allreduce(T.uint32(1), in_thread_var_NT_matmul_intermediate_local[0], T.bool(True), cross_thread_var_NT_matmul_intermediate_local[0], ax0_fused) + T.tvm_thread_allreduce(T.Lambda([T.float16, T.float16], lambda x0, y0: (x0 + y0,)), (T.float16(0),), (in_thread_var_NT_matmul_intermediate_local[0],), T.bool(True), (cross_thread_var_NT_matmul_intermediate_local[0],), (ax0_fused,)) with Ts.sblock("NT_matmul_write_back"): v0 = Ts.axis.spatial(T.int64(32), ax0_ax1_fused // n) v1 = Ts.axis.spatial(n, ax0_ax1_fused % n) @@ -1754,13 +1714,13 @@ def lowered_no_thread_broadcast( vi, vk = Ts.axis.remap("SR", [i, k]) Ts.reads(A[vi, vk]) Ts.writes(cross_thread_temp_1_local[0]) - T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.tvm_thread_allreduce( - T.uint32(1), A[vi, vk], T.bool(True), cross_thread_temp_1_local[0], k + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A[vi, vk],), + T.bool(True), + (cross_thread_temp_1_local[0],), + (k,), ) with Ts.sblock("sum_write_back"): vi = Ts.axis.spatial(256, i) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py b/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py index 0ab0f63c2647..7be9050194e6 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py @@ -63,18 +63,14 @@ def main(A: T.Tensor((128, 32), "float32"), B: T.Tensor(128, "float32")): reduce = T.alloc_tensor((1,), scope="local") reduce_1 = T.decl_tensor(1, data=reduce.data, scope="local") - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), - A_flat[0], - T.bool(True), - reduce_1[0], - threadIdx_x, - ) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: x + y), + T.float32(0), + A_flat[0], + T.bool(True), + reduce_1[0], + threadIdx_x, + ) if threadIdx_x == 0: B[i] = reduce_1[0] @@ -102,18 +98,14 @@ def main(A: T.Tensor((128, 32), "float32"), B: T.Tensor(128, "float32")): reduce = T.decl_tensor(1, dtype="float32", scope="local") - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), - A_flat[0], - T.bool(True), - reduce[0], - threadIdx_x, - ) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (A_flat[0],), + T.bool(True), + (reduce[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B[i] = reduce[0] @@ -150,18 +142,14 @@ def main(A: T.Tensor((128, 128), "float32"), B: T.Tensor(128, "float32")): normal_reduce_1[0] + A_flat[i * 128 + ko * 32 + threadIdx_x] ) - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), - normal_reduce_1[0], - T.bool(True), - reduce_1[0], - threadIdx_x, - ) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (normal_reduce_1[0],), + T.bool(True), + (reduce_1[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B[i] = reduce_1[0] @@ -183,19 +171,15 @@ def main(A: T.Tensor((32, 32), "float32"), B: T.Tensor((32,), "float32")): cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 32) cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - A_1 = T.decl_tensor((1024,), data=A.data) - T.tvm_thread_allreduce( - T.uint32(1), - A_1[threadIdx_y * 32 + threadIdx_x], - T.bool(True), - cross_thread_B_1[0], - threadIdx_x, - ) + A_1 = T.decl_tensor((1024,), data=A.data) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A_1[threadIdx_y * 32 + threadIdx_x],), + T.bool(True), + (cross_thread_B_1[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B_1 = T.decl_tensor((32,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] @@ -218,19 +202,15 @@ def main(A: T.Tensor((4, 128), "float32"), B: T.Tensor((4,), "float32")): cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 128) cross_thread_B_alias = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - A_flat = T.decl_tensor((512,), data=A.data) - T.tvm_thread_allreduce( - T.uint32(1), - A_flat[threadIdx_y * 128 + threadIdx_x], - T.bool(True), - cross_thread_B[0], - threadIdx_x, - ) + A_flat = T.decl_tensor((512,), data=A.data) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A_flat[threadIdx_y * 128 + threadIdx_x],), + T.bool(True), + (cross_thread_B[0],), + (threadIdx_x,), + ) cross_thread_B_alias[0] = cross_thread_B[0] if threadIdx_x == 0: B_flat = T.decl_tensor((4,), data=B.data) @@ -257,19 +237,15 @@ def main(A: T.Tensor((4, 128), "float32"), B: T.Tensor((4,), "float32")): threadIdx_y = T.launch_thread("threadIdx.y", 4) cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 128) - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - A_flat = T.decl_tensor((512,), data=A.data) - T.tvm_thread_allreduce( - T.uint32(1), - A_flat[threadIdx_y * 128 + threadIdx_x], - T.bool(True), - cross_thread_B[0], - threadIdx_x, - ) + A_flat = T.decl_tensor((512,), data=A.data) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A_flat[threadIdx_y * 128 + threadIdx_x],), + T.bool(True), + (cross_thread_B[0],), + (threadIdx_x,), + ) cross_thread_B_alias = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") cross_thread_B[0] = cross_thread_B_alias[0] if threadIdx_x == 0: @@ -298,19 +274,15 @@ def main(A: T.Tensor((32, 8), "float32"), B: T.Tensor((32,), "float32")): cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 8) cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - A_1 = T.decl_tensor((256,), data=A.data) - T.tvm_thread_allreduce( - T.uint32(1), - A_1[threadIdx_y * 8 + threadIdx_x], - T.bool(True), - cross_thread_B_1[0], - threadIdx_x, - ) + A_1 = T.decl_tensor((256,), data=A.data) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A_1[threadIdx_y * 8 + threadIdx_x],), + T.bool(True), + (cross_thread_B_1[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B_1 = T.decl_tensor((32,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] @@ -333,19 +305,15 @@ def main(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128,), "float32")): threadIdx_x = T.launch_thread("threadIdx.x", 128) cross_thread_B = T.alloc_tensor((1,), scope="local") cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - A_1 = T.decl_tensor((16384,), data=A.data) - T.tvm_thread_allreduce( - T.uint32(1), - A_1[i * 128 + threadIdx_x], - T.bool(True), - cross_thread_B_1[0], - threadIdx_x, - ) + A_1 = T.decl_tensor((16384,), data=A.data) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A_1[i * 128 + threadIdx_x],), + T.bool(True), + (cross_thread_B_1[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B_1 = T.decl_tensor((128,), data=B.data) B_1[i] = cross_thread_B_1[0] @@ -368,15 +336,15 @@ def main(A: T.Tensor((1, 1024), "float32"), B: T.Tensor((1,), "float32")): threadIdx_x = T.launch_thread("threadIdx.x", 1024) cross_thread_B = T.alloc_tensor((1,), scope="local") cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - A_1 = T.decl_tensor((1024,), data=A.data) - T.tvm_thread_allreduce( - T.uint32(1), A_1[threadIdx_x], T.bool(True), cross_thread_B_1[0], threadIdx_x - ) + A_1 = T.decl_tensor((1024,), data=A.data) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A_1[threadIdx_x],), + T.bool(True), + (cross_thread_B_1[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B_1 = T.decl_tensor((1,), data=B.data) B_1[0] = cross_thread_B_1[0] @@ -400,19 +368,15 @@ def main(A: T.Tensor((4, 128), "float32"), B: T.Tensor((4,), "float32")): cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 128) cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - A_1 = T.decl_tensor((512,), data=A.data) - T.tvm_thread_allreduce( - T.uint32(1), - A_1[threadIdx_y * 128 + threadIdx_x], - T.bool(True), - cross_thread_B_1[0], - threadIdx_x, - ) + A_1 = T.decl_tensor((512,), data=A.data) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A_1[threadIdx_y * 128 + threadIdx_x],), + T.bool(True), + (cross_thread_B_1[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B_1 = T.decl_tensor((4,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] @@ -442,14 +406,14 @@ def main(A: T.Tensor((2, 70), "float32"), B: T.Tensor((2,), "float32")): A_1 = T.decl_tensor((140,), data=A.data) in_thread_B_1[0] = in_thread_B_1[0] + A_1[threadIdx_y * 70 + threadIdx_x] cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), in_thread_B_1[0], T.bool(True), cross_thread_B_1[0], threadIdx_x - ) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (in_thread_B_1[0],), + T.bool(True), + (cross_thread_B_1[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B_1 = T.decl_tensor((2,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] @@ -486,19 +450,15 @@ def main(A: T.Tensor((1, 1, 2, 128), "float32"), B: T.Tensor((1, 1, 2), "float32 threadIdx_y = T.launch_thread("threadIdx.y", 2) threadIdx_x = T.launch_thread("threadIdx.x", 128) cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - A_1 = T.decl_tensor((256,), data=A.data) - T.tvm_thread_allreduce( - T.uint32(1), - A_1[threadIdx_y * 128 + threadIdx_x], - T.bool(True), - cross_thread_B_1[0], - threadIdx_x, - ) + A_1 = T.decl_tensor((256,), data=A.data) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A_1[threadIdx_y * 128 + threadIdx_x],), + T.bool(True), + (cross_thread_B_1[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B_1 = T.decl_tensor((2,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] @@ -538,18 +498,14 @@ def main(A: T.Tensor((128, 32), "float32"), B: T.Tensor(128, "float32")): reduce = T.decl_tensor(1, data=reduce_data.data, scope="local") reduce_alias = T.decl_tensor(1, data=reduce.data, scope="local") - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), - A_flat[0], - T.bool(True), - reduce_alias[0], - threadIdx_x, - ) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + (A_flat[0],), + T.bool(True), + (reduce_alias[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B[i] = reduce_alias[0] @@ -588,19 +544,15 @@ def main(A: T.Tensor((1, 1, 2, 128), "float32"), B: T.Tensor((1, 1, 2), "float32 threadIdx_y = T.launch_thread("threadIdx.y", 2) threadIdx_x = T.launch_thread("threadIdx.x", 128) cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - A_1 = T.decl_tensor((256,), data=A.data) - T.tvm_thread_allreduce( - T.uint32(1), - A_1[threadIdx_y * 128 + threadIdx_x], - T.bool(True), - cross_thread_B_1[0], - threadIdx_x, - ) + A_1 = T.decl_tensor((256,), data=A.data) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (A_1[threadIdx_y * 128 + threadIdx_x],), + T.bool(True), + (cross_thread_B_1[0],), + (threadIdx_x,), + ) if threadIdx_x == 0: B_1 = T.decl_tensor((2,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] diff --git a/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py b/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py index e59aa85c082f..447b65cd318a 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py @@ -190,18 +190,14 @@ def func(A: T.Tensor((16 * 512), "float32")): A_temp_4 = T.bind(in_thread_A_temp_1[0] + A_shared_1[threadIdx_x + 384]) in_thread_A_temp_1[0] = A_temp_4 cross_thread_A_temp_1 = T.decl_tensor((1,), data=cross_thread_A_temp.data, scope="local") - with T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), - in_thread_A_temp_1[0], - T.bool(True), - cross_thread_A_temp_1[0], - threadIdx_x, - ) + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (in_thread_A_temp_1[0],), + T.bool(True), + (cross_thread_A_temp_1[0],), + (threadIdx_x,), + ) @Ts.prim_func(private=True) def expected(A: T.Tensor((8192,), "float32")): @@ -227,17 +223,13 @@ def expected(A: T.Tensor((8192,), "float32")): cross_thread_A_temp_1_1 = T.decl_tensor( (1,), data=cross_thread_A_temp_1.data, scope="local" ) - T.attr( - T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ) T.tvm_thread_allreduce( - T.uint32(1), - in_thread_A_temp_1_1[0], + T.Lambda([T.float32, T.float32], lambda x0, y0: (x0 + y0,)), + (T.float32(0),), + (in_thread_A_temp_1_1[0],), T.bool(True), - cross_thread_A_temp_1_1[0], - threadIdx_x, + (cross_thread_A_temp_1_1[0],), + (threadIdx_x,), ) mod = tvm.IRModule({"main": func}) diff --git a/tests/python/tirx-base/test_tir_op_types.py b/tests/python/tirx-base/test_tir_op_types.py index fad780f8cc27..83a3944b81c5 100644 --- a/tests/python/tirx-base/test_tir_op_types.py +++ b/tests/python/tirx-base/test_tir_op_types.py @@ -93,11 +93,12 @@ def test_tir_op_call_likely(): def test_tir_op_tvm_thread_allreduce(): - x = tirx.Var("x", "int32") - buffer = tirx.decl_tensor((128), "float32") - y = tirx.Var("y", "handle") - z = tirx.Var("z", "int32") - expr = tirx.tvm_thread_allreduce(x, buffer[0], True, y, z) + tensor = tirx.decl_tensor((128,), "float32") + axis = tirx.Var("thread", "int32") + combine = tvm.ir.LambdaExpr(["float32", "float32"], lambda lhs, rhs: (lhs + rhs,)) + expr = tirx.tvm_thread_allreduce( + combine, (tirx.const(0, "float32"),), (tensor[0],), True, (tensor[1],), (axis,) + ) assert expr.op.name == "tirx.tvm_thread_allreduce" diff --git a/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py b/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py index d11d11f2ac6f..4329ed2a210b 100644 --- a/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py +++ b/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py @@ -359,18 +359,14 @@ def main( reduce = T.decl_tensor(1, dtype="bfloat16", scope="local") - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.bfloat16(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), - A_flat[0], - T.bool(True), - reduce[0], - threadIdx_x, - ) + T.tvm_thread_allreduce( + T.Lambda([T.bfloat16, T.bfloat16], lambda x, y: (x + y,)), + (T.bfloat16(0),), + (A_flat[0],), + T.bool(True), + (reduce[0],), + (threadIdx_x,), + ) return Before @@ -388,13 +384,10 @@ def main( reduce = T.decl_tensor(1, dtype="float32", scope="local") - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + ( T.reinterpret( "float32", T.shift_left( @@ -402,10 +395,11 @@ def main( T.uint32(16), ), ), - T.bool(True), - reduce[0], - threadIdx_x, - ) + ), + T.bool(True), + (reduce[0],), + (threadIdx_x,), + ) return After @@ -423,13 +417,10 @@ def main( reduce = T.decl_tensor(1, dtype="float32", scope="local") - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), + T.tvm_thread_allreduce( + T.Lambda([T.float32, T.float32], lambda x, y: (x + y,)), + (T.float32(0),), + ( T.reinterpret( "float32", T.shift_left( @@ -437,10 +428,11 @@ def main( T.uint32(16), ), ), - T.bool(True), - reduce[0], - threadIdx_x, - ) + ), + T.bool(True), + (reduce[0],), + (threadIdx_x,), + ) return After @@ -467,18 +459,14 @@ def main( reduce = T.decl_tensor(1, dtype="bfloat16", scope="local") - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.bfloat16(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), - A_flat[0], - T.bool(True), - reduce[0], - threadIdx_x, - ) + T.tvm_thread_allreduce( + T.Lambda([T.bfloat16, T.bfloat16], lambda x, y: (x + y,)), + (T.bfloat16(0),), + (A_flat[0],), + T.bool(True), + (reduce[0],), + (threadIdx_x,), + ) return Before @@ -496,18 +484,14 @@ def main( reduce = T.decl_tensor(1, dtype="bfloat16", scope="local") - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.bfloat16(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), - A_flat[0], - T.bool(True), - reduce[0], - threadIdx_x, - ) + T.tvm_thread_allreduce( + T.Lambda([T.bfloat16, T.bfloat16], lambda x, y: (x + y,)), + (T.bfloat16(0),), + (A_flat[0],), + T.bool(True), + (reduce[0],), + (threadIdx_x,), + ) return After @@ -525,18 +509,14 @@ def main( reduce = T.decl_tensor(1, dtype="bfloat16", scope="local") - with T.attr( - T.comm_reducer(lambda x, y: x + y, [T.bfloat16(0)]), - "reduce_scope", - T.int32(0), - ): - T.tvm_thread_allreduce( - T.uint32(1), - A_flat[0], - T.bool(True), - reduce[0], - threadIdx_x, - ) + T.tvm_thread_allreduce( + T.Lambda([T.bfloat16, T.bfloat16], lambda x, y: (x + y,)), + (T.bfloat16(0),), + (A_flat[0],), + T.bool(True), + (reduce[0],), + (threadIdx_x,), + ) return After