Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions include/tvm/ir/base_expr.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<StagingExprNode>(); }
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.
*
Expand Down
42 changes: 42 additions & 0 deletions include/tvm/ir/expr.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<Var> 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<Expr>& arguments) const;

static void RegisterReflection() {
namespace refl = tvm::ffi::reflection;
refl::ObjectDef<LambdaExprNode>()
.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<Var> 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.
Expand Down
4 changes: 4 additions & 0 deletions include/tvm/ir/expr_functor.h
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,9 @@ class ExprFunctor<R(const Expr&, Args...)> {
return (*vtable_)(node, this, std::forward<Args>(args)...);
}

virtual R Dispatch_(const LambdaExprNode* node, Args... args) {
return DispatchDefault_(node, std::forward<Args>(args)...);
}
virtual R Dispatch_(const OpaqueExprNode* node, Args... args) {
return DispatchDefault_(node, std::forward<Args>(args)...);
}
Expand Down Expand Up @@ -230,6 +233,7 @@ class ExprFunctor<R(const Expr&, Args...)> {
* \param vtable The table to initialize before adding derived registrations.
*/
static void InitVTable(VTable* vtable) {
SetDispatch<TSelf, LambdaExprNode>(vtable);
SetDispatch<TSelf, OpaqueExprNode>(vtable);
SetDispatch<TSelf, TupleNode>(vtable);
SetDispatch<TSelf, TupleGetItemNode>(vtable);
Expand Down
3 changes: 0 additions & 3 deletions include/tvm/s_tir/stmt.h
Original file line number Diff line number Diff line change
Expand Up @@ -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
// -----------------------------------------------------------------------
Expand Down
32 changes: 24 additions & 8 deletions include/tvm/tirx/builtin.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<Expr> GetAllreduceFields(const Expr& value) {
if (const auto* tuple = value.as<tvm::TupleNode>()) return tuple->fields;
return {value};
}

// Metal cooperative_tensor intrinsics (MetalPerformancePrimitives / Metal 4)

/*!
Expand Down
45 changes: 0 additions & 45 deletions include/tvm/tirx/tile_primitive.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<Var> 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<PrimExpr>& indices) const;

static void RegisterReflection() {
namespace refl = tvm::ffi::reflection;
refl::ObjectDef<LambdaExprNode>()
.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<Var> 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.
Expand Down
2 changes: 2 additions & 0 deletions python/tvm/ir/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,8 @@
ExprOperand,
ExprWithOp,
GlobalVar,
LambdaExpr,
StagingExpr,
OpaqueExpr,
Range,
TensorLoad,
Expand Down
85 changes: 85 additions & 0 deletions python/tvm/ir/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -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.
Expand Down
7 changes: 1 addition & 6 deletions python/tvm/relax/frontend/nn/llm/_decode_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down
4 changes: 3 additions & 1 deletion python/tvm/script/ir_builder/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -68,6 +69,7 @@
"GenericConst",
"IRBuilder",
"IRModuleFrame",
"Lambda",
"MissingType",
"PrimType",
"Range",
Expand Down
11 changes: 10 additions & 1 deletion python/tvm/script/ir_builder/ir.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
2 changes: 1 addition & 1 deletion python/tvm/tirx/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading
Loading