This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch script/canonical-parser-df in repository https://gitbox.apache.org/repos/asf/tvm.git
commit 5e56b88a7217a33c1a7451defdce1d24acb238a8 Author: Tianqi Chen <[email protected]> AuthorDate: Wed Sep 23 12:55:09 2026 +0000 [DOCS][Script] Restore builder API details and executable JIT examples Document builder parameters, returned values, scope selection and source behavior at public call sites. Align JIT examples with explicit constexpr selection and postponed symbolic annotations. --- python/tvm/relax/script/builder/ir.py | 72 ++++++++++-- python/tvm/script/ir_builder/base.py | 24 +++- python/tvm/script/ir_builder/ir/ir.py | 37 ++++++- python/tvm/tirx/script/builder/ir.py | 197 ++++++++++++++++++++++++++------- python/tvm/tirx/script/builder/tirx.py | 32 +++++- python/tvm/tirx/script/jit.py | 39 +++++++ 6 files changed, 342 insertions(+), 59 deletions(-) diff --git a/python/tvm/relax/script/builder/ir.py b/python/tvm/relax/script/builder/ir.py index 15d0598904..ae531d28d9 100644 --- a/python/tvm/relax/script/builder/ir.py +++ b/python/tvm/relax/script/builder/ir.py @@ -291,17 +291,35 @@ def arg(name: py_str, ty: Type) -> Var: def func_name(name: py_str) -> None: - """Specify the name of the last function frame.""" + """Specify the name of the last function frame. + + Parameters + ---------- + name: str + The function name. + """ return _ffi_api.FuncName(name) # type: ignore[attr-defined] # pylint: disable=no-member def func_attr(attrs: dict[py_str, tvm_Object]) -> None: - """Specify the attrs of the last function frame.""" + """Specify the attrs of the last function frame. + + Parameters + ---------- + attrs: Dict[str, Object] + The function attrs. + """ return _ffi_api.FuncAttrs(attrs) # type: ignore[attr-defined] # pylint: disable=no-member def func_ret_type(ret_ty: Type) -> None: - """Specify the return type of the last function frame.""" + """Specify the return type of the last function frame. + + Parameters + ---------- + ret_ty: Type + The function return type. + """ return _ffi_api.FuncRetType(ret_ty) # type: ignore[attr-defined] # pylint: disable=no-member @@ -311,7 +329,13 @@ def func_ret_ty(ret_ty: Type) -> None: def func_ret_value(value: Expr) -> None: - """Specify the return value of the last function frame.""" + """Specify the return value of the last function frame. + + Parameters + ---------- + value: Expr + The function return value. + """ return _ffi_api.FuncRetValue(value) # type: ignore[attr-defined] # pylint: disable=no-member @@ -361,12 +385,24 @@ def rewriter(rewriter_mod: IRModule | type) -> PatternMatchingRewriter: def dataflow() -> frame.BindingBlockFrame: - """Start a dataflow binding block frame.""" + """Start a dataflow binding block frame. + + Returns + ------- + frame: frame.BindingBlockFrame + The created ir_builder Block frame. + """ return _ffi_api.Dataflow() # type: ignore[attr-defined] # pylint: disable=no-member def output(*vars: tuple[Var]) -> None: - """Expose the dataflow block output variables as global ones.""" + """Expose the dataflow block output variables as global ones. + + Parameters + ---------- + vars: Tuple[Var] + The output variables of a dataflow block. + """ return _ffi_api.DataflowBlockOutput(vars) # type: ignore[attr-defined] # pylint: disable=no-member @@ -620,7 +656,13 @@ def emit_with_ty( def SeqExpr() -> frame.SeqExprFrame: # pylint: disable=invalid-name - """Create a SeqExpr frame.""" + """Create a SeqExpr frame. + + Returns + ------- + res : frame.SeqExprFrame + The result SeqExprFrame + """ return _ffi_api.SeqExpr() # type: ignore[attr-defined] # pylint: disable=no-member @@ -650,12 +692,24 @@ def If(condition: Expr) -> frame.IfFrame: # pylint: disable=invalid-name def Then() -> frame.ThenFrame: # pylint: disable=invalid-name - """Create a then frame.""" + """Create a then frame. + + Returns + ------- + res : frame.ThenFrame + The result ThenFrame. + """ return _ffi_api.Then() # type: ignore[attr-defined] # pylint: disable=no-member def Else() -> frame.ElseFrame: # pylint: disable=invalid-name - """Create an else frame.""" + """Create an else frame. + + Returns + ------- + res : frame.ElseFrame + The result ElseFrame. + """ return _ffi_api.Else() # type: ignore[attr-defined] # pylint: disable=no-member diff --git a/python/tvm/script/ir_builder/base.py b/python/tvm/script/ir_builder/base.py index a781234cd5..671b5109f7 100644 --- a/python/tvm/script/ir_builder/base.py +++ b/python/tvm/script/ir_builder/base.py @@ -79,7 +79,13 @@ class IRBuilderFrame(_Object): _ffi_api.IRBuilderFrameExit(self) # type: ignore[attr-defined] # pylint: disable=no-member def add_callback(self, callback: Callable[[], None]) -> None: - """Add a callback method invoked when exiting the with-scope.""" + """Add a callback method invoked when exiting the with-scope. + + Parameters + ---------- + callback : Callable[[], None] + The callback method to be invoked. + """ _ffi_api.IRBuilderFrameAddCallback( # type: ignore[attr-defined] # pylint: disable=no-member self, callback ) @@ -135,12 +141,24 @@ class IRBuilder(_Object): @staticmethod def current() -> "IRBuilder": - """Get the current IRBuilder put in the with-scope.""" + """Get the current IRBuilder put in the with-scope. + + Returns + ------- + builder : IRBuilder + The current IRBuilder. + """ return _ffi_api.IRBuilderCurrent() # type: ignore[attr-defined] # pylint: disable=no-member @staticmethod def is_in_scope() -> bool: - """See if the current thread-local scope has an IRBuilder.""" + """See if the current thread-local scope has an IRBuilder. + + Returns + ------- + bool + Whether the current thread-local scope has an IRBuilder + """ return _ffi_api.IRBuilderIsInScope() # type: ignore[attr-defined] # pylint: disable=no-member def get(self) -> _Object: diff --git a/python/tvm/script/ir_builder/ir/ir.py b/python/tvm/script/ir_builder/ir/ir.py index 4befa2d6fd..61e4ae3461 100644 --- a/python/tvm/script/ir_builder/ir/ir.py +++ b/python/tvm/script/ir_builder/ir/ir.py @@ -38,7 +38,18 @@ if TYPE_CHECKING: else: class meta_var: # pylint: disable=invalid-name - """A value used only for TVMScript parser-time metaprogramming.""" + """A value used only for TVMScript parser-time metaprogramming. + + Assignments unwrap this object without emitting an IR binding. The + shared wrapper is exposed as ``I.meta_var``; dialect namespaces may + provide compatibility aliases to the same implementation. + For Relax, this is the explicit opt-out from default primitive binding emission. + + Parameters + ---------- + value : Any + The parser-time value. + """ def __init__(self, value: Any) -> None: self.value = value @@ -48,7 +59,13 @@ else: def ir_module() -> IRModuleFrame: - """Start a ir_module frame.""" + """Start a ir_module frame. + + Returns + ------- + frame: IRModuleFrame + The constructed frame. + """ return _ffi_api.IRModule() # type: ignore[attr-defined] # pylint: disable=no-member @@ -143,7 +160,13 @@ def module_set_attr( def module_global_infos(global_infos: dict[str, list[GlobalInfo]]) -> None: - """Specify the global infos of the ir_module frame.""" + """Specify the global infos of the ir_module frame. + + Parameters + ---------- + global_infos: Dict[str, List[GlobalInfo]] + The module global infos. + """ if IRBuilder.is_in_scope(): return _ffi_api.ModuleGlobalInfos(global_infos) # Keep native argument validation even before Python has applied the module decorator. @@ -201,7 +224,13 @@ def lookup_global_info(name: str, index: int) -> GlobalInfo: def dummy_global_info() -> "DummyGlobalInfo": - """Create a dummy global info expression.""" + """Create a dummy global info expression. + + Returns + ------- + res : DummyGlobalInfo + The result dummy global info. + """ from tvm.relax import DummyGlobalInfo # pylint: disable=import-outside-toplevel return DummyGlobalInfo() # type: ignore[attr-defined] # pylint: disable=no-member diff --git a/python/tvm/tirx/script/builder/ir.py b/python/tvm/tirx/script/builder/ir.py index e7091bb1c1..55ba6f4708 100644 --- a/python/tvm/tirx/script/builder/ir.py +++ b/python/tvm/tirx/script/builder/ir.py @@ -129,7 +129,12 @@ def cast(value, dtype, span=None): def _current_s_tir() -> bool: - """Return True if the innermost enclosing PrimFuncFrame has ``s_tir=True``.""" + """Return True if the innermost enclosing PrimFuncFrame has ``s_tir=True``. + + Gates the builder's default layout fill: ``s_tir=True`` PrimFuncs leave + ``layout=None`` so s_tir-style passes that do not touch layout round-trip + cleanly. Ordinary TIRx storage scopes instead use ``TileLayout(S[shape])``. + """ from tvm.script.ir_builder.base import IRBuilder # local import to avoid cycle if not IRBuilder.is_in_scope(): @@ -439,12 +444,24 @@ def arg(name: str, obj: Var | Buffer) -> Var | Buffer: def func_name(name: str) -> None: - """The PrimFunc naming statement.""" + """The PrimFunc naming statement. + + Parameters + ---------- + name : str + The name of the PrimFunc. + """ _ffi_api.FuncName(name) # type: ignore[attr-defined] # pylint: disable=no-member def func_attr(attrs: dict[str, Any]) -> None: - """The PrimFunc annotation statement.""" + """The PrimFunc annotation statement. + + Parameters + ---------- + attrs : Dict[str, Any] + The annotations of the PrimFunc. + """ _ffi_api.FuncAttrs(attrs) # type: ignore[attr-defined] # pylint: disable=no-member @@ -640,7 +657,13 @@ def device_entry() -> None: def elected(): - """Stub that rejects the removed ``T.elected()`` sugar.""" + """Stub that rejects the removed ``T.elected()`` sugar. + + Write the explicit form instead:: + + if T.cuda.elect_sync(): + ... # thread is the default scope + """ raise RuntimeError( "T.elected() is no longer available. Write explicitly: " "`if T.cuda.elect_sync(): ...` (thread is the default scope)" @@ -663,8 +686,10 @@ def scope_id( def cluster_id(extents: list[Expr | int] | None = None, dtype: str = "int32") -> BypassBind: - """Define a kernel→cluster scope id. Pass ``None`` (the default) to defer the extent; it - will be inferred at LowerTIRx from sibling ScopeIdDef closure. + """Define a kernel→cluster scope id. Pass ``None`` (the default) to defer the + extent; it will be inferred at LowerTIRx from sibling ScopeIdDef closure. + + ``dtype`` selects the dtype of the introduced vars (``"int32"`` or ``"uint32"``). """ ret = _ffi_api.ClusterId(extents, "kernel", dtype) # type: ignore[attr-defined] # pylint: disable=no-member if len(ret) == 1: @@ -675,8 +700,10 @@ def cluster_id(extents: list[Expr | int] | None = None, dtype: str = "int32") -> def cta_id( extents: list[Expr | int] | None = None, preferred=None, dtype: str = "int32" ) -> BypassBind: - """Define a kernel→cta scope id. Pass ``None`` (the default) to defer the extent; it will - be inferred at LowerTIRx from sibling ScopeIdDef closure. + """Define a kernel→cta scope id. Pass ``None`` (the default) to defer the + extent; it will be inferred at LowerTIRx from sibling ScopeIdDef closure. + + ``dtype`` selects the dtype of the introduced vars (``"int32"`` or ``"uint32"``). """ ret = _ffi_api.CtaId(extents, "kernel", preferred, dtype) # type: ignore[attr-defined] # pylint: disable=no-member if len(ret) == 1: @@ -687,8 +714,10 @@ def cta_id( def cta_id_in_cluster( extents: list[Expr | int] | None = None, preferred=None, dtype: str = "int32" ) -> BypassBind: - """Define a cluster→cta scope id. Pass ``None`` (the default) to defer the extent; it - will be inferred at LowerTIRx from sibling ScopeIdDef closure. + """Define a cluster→cta scope id. Pass ``None`` (the default) to defer the + extent; it will be inferred at LowerTIRx from sibling ScopeIdDef closure. + + ``dtype`` selects the dtype of the introduced vars (``"int32"`` or ``"uint32"``). """ ret = _ffi_api.CtaId(extents, "cluster", preferred, dtype) # type: ignore[attr-defined] # pylint: disable=no-member if len(ret) == 1: @@ -702,8 +731,10 @@ def cta_id_in_pair(dtype: str = "int32") -> BypassBind: def warpgroup_id(extents: list[Expr | int] | None = None, dtype: str = "int32") -> BypassBind: - """Define a cta→warpgroup scope id. Pass ``None`` (the default) to defer the extent; it - will be inferred at LowerTIRx from sibling closure. + """Define a cta→warpgroup scope id. Pass ``None`` (the default) to defer + the extent; it will be inferred at LowerTIRx from sibling closure. + + ``dtype`` selects the dtype of the introduced vars (``"int32"`` or ``"uint32"``). """ ret = _ffi_api.WarpgroupId(extents, "cta", dtype) # type: ignore[attr-defined] # pylint: disable=no-member if len(ret) == 1: @@ -712,8 +743,10 @@ def warpgroup_id(extents: list[Expr | int] | None = None, dtype: str = "int32") def warp_id(extents: list[Expr | int] | None = None, dtype: str = "int32") -> BypassBind: - """Define a cta→warp scope id. Pass ``None`` (the default) to defer the extent; it will - be inferred at LowerTIRx from sibling closure. + """Define a cta→warp scope id. Pass ``None`` (the default) to defer the + extent; it will be inferred at LowerTIRx from sibling closure. + + ``dtype`` selects the dtype of the introduced vars (``"int32"`` or ``"uint32"``). """ ret = _ffi_api.WarpId(extents, "cta", dtype) # type: ignore[attr-defined] # pylint: disable=no-member if len(ret) == 1: @@ -722,8 +755,10 @@ def warp_id(extents: list[Expr | int] | None = None, dtype: str = "int32") -> By def warp_id_in_wg(extents: list[Expr | int] | None = None, dtype: str = "int32") -> BypassBind: - """Define a warpgroup→warp scope id. Pass ``None`` (the default) to defer the extent; it - will be inferred at LowerTIRx from sibling closure. + """Define a warpgroup→warp scope id. Pass ``None`` (the default) to defer + the extent; it will be inferred at LowerTIRx from sibling closure. + + ``dtype`` selects the dtype of the introduced vars (``"int32"`` or ``"uint32"``). """ ret = _ffi_api.WarpId(extents, "warpgroup", dtype) # type: ignore[attr-defined] # pylint: disable=no-member if len(ret) == 1: @@ -732,8 +767,10 @@ def warp_id_in_wg(extents: list[Expr | int] | None = None, dtype: str = "int32") def lane_id(extents: list[Expr | int] | None = None, dtype: str = "int32") -> BypassBind: - """Define a warp→thread scope id. Pass ``None`` (the default) to defer the extent; it - will be inferred at LowerTIRx from sibling closure. + """Define a warp→thread scope id. Pass ``None`` (the default) to defer the + extent; it will be inferred at LowerTIRx from sibling closure. + + ``dtype`` selects the dtype of the introduced vars (``"int32"`` or ``"uint32"``). """ ret = _ffi_api.ThreadId(extents, "warp", dtype) # type: ignore[attr-defined] # pylint: disable=no-member if len(ret) == 1: @@ -742,8 +779,10 @@ def lane_id(extents: list[Expr | int] | None = None, dtype: str = "int32") -> By def thread_id(extents: list[Expr | int] | None = None, dtype: str = "int32") -> BypassBind: - """Define a cta→thread scope id. Pass ``None`` (the default) to defer the extent; it will - be inferred at LowerTIRx from sibling closure. + """Define a cta→thread scope id. Pass ``None`` (the default) to defer the + extent; it will be inferred at LowerTIRx from sibling closure. + + ``dtype`` selects the dtype of the introduced vars (``"int32"`` or ``"uint32"``). """ ret = _ffi_api.ThreadId(extents, "cta", dtype) # type: ignore[attr-defined] # pylint: disable=no-member if len(ret) == 1: @@ -754,8 +793,10 @@ def thread_id(extents: list[Expr | int] | None = None, dtype: str = "int32") -> def thread_id_in_wg( extents: list[Expr | int] | None = None, dtype: str = "int32" ) -> BypassBind: - """Define a warpgroup→thread scope id. Pass ``None`` (the default) to defer the extent; - it will be inferred at LowerTIRx from sibling closure. + """Define a warpgroup→thread scope id. Pass ``None`` (the default) to defer + the extent; it will be inferred at LowerTIRx from sibling closure. + + ``dtype`` selects the dtype of the introduced vars (``"int32"`` or ``"uint32"``). """ ret = _ffi_api.ThreadId(extents, "warpgroup", dtype) # type: ignore[attr-defined] # pylint: disable=no-member if len(ret) == 1: @@ -764,12 +805,24 @@ def thread_id_in_wg( def init() -> frame.BlockInitFrame: - """The block initialization statement.""" + """The block initialization statement. + + Returns + ------- + res : frame.BlockInitFrame + The BlockInitFrame. + """ return _ffi_api.Init() # type: ignore[attr-defined] # pylint: disable=no-member def where(predicate: Expr | int) -> None: - """The block predicate statement.""" + """The block predicate statement. + + Parameters + ---------- + predicate : Union[Expr, Literal[0, 1]] + The predicate condition. + """ if isinstance(predicate, bool): predicate = IntImm("bool", predicate) if isinstance(predicate, int): @@ -781,7 +834,13 @@ def where(predicate: Expr | int) -> None: def reads(*buffer_slices: list[TensorRegion | TensorLoad]) -> None: - """The block buffer region reading statement.""" + """The block buffer region reading statement. + + Parameters + ---------- + buffer_slices : List[Union[TensorRegion, TensorLoad]] + The array of buffer regions to read. + """ if len(buffer_slices) == 1: if isinstance(buffer_slices[0], tuple): buffer_slices = list(buffer_slices[0]) @@ -795,7 +854,13 @@ def reads(*buffer_slices: list[TensorRegion | TensorLoad]) -> None: def writes(*buffer_slices: list[TensorRegion | TensorLoad]) -> None: - """The block buffer region writing statement.""" + """The block buffer region writing statement. + + Parameters + ---------- + buffer_slices : List[Union[TensorRegion, TensorLoad]] + The array of buffer regions to write. + """ if len(buffer_slices) == 1: if isinstance(buffer_slices[0], tuple): buffer_slices = list(buffer_slices[0]) @@ -809,7 +874,13 @@ def writes(*buffer_slices: list[TensorRegion | TensorLoad]) -> None: def sblock_attr(attrs: dict[str, Any]) -> None: - """The block annotation statement (for non-tirx SBlock usage).""" + """The block annotation statement (for non-tirx SBlock usage). + + Parameters + ---------- + attrs : Dict[str, Any] + The annotation of the block. + """ return _ffi_api.BlockAttrs(attrs) # type: ignore[attr-defined] # pylint: disable=no-member @@ -1533,9 +1604,12 @@ def bind( class LetAnnotation: - """Marker for explicit LetStmt. Created by T.let or T.let[type]. Usage in TVMScript: x: - T.let[T.int32] = expr # LetStmt with explicit type x: T.let = expr # LetStmt with - auto-typed RHS + """Marker for explicit LetStmt. Created by ``T.let`` or ``T.let[type]``. + + Usage in TVMScript:: + + x: T.let[T.int32] = expr # LetStmt with explicit type + x: T.let = expr # LetStmt with auto-typed RHS """ def __init__(self, type_spec=None): @@ -1569,7 +1643,12 @@ let = LetAnnotation() # Singleton for T.let (no subscript) class LocalVectorAnnotation: - """Marker for local vector/tensor allocation via type annotation subscript.""" + """Marker for local vector/tensor allocation via type annotation subscript. + + Created when a DtypeConstructor is subscripted, e.g. ``T.float32[N]`` or + ``T.float32[M, N]``. The declaration protocol recognizes this annotation + and allocates local storage with ``T.alloc_local(shape=..., dtype=...)``. + """ __slots__ = ("dtype", "shape") @@ -1763,12 +1842,24 @@ def If(condition: Expr) -> frame.IfFrame: # pylint: disable=invalid-name def Then() -> frame.ThenFrame: # pylint: disable=invalid-name - """Create a then.""" + """Create a then. + + Returns + ------- + res : frame.ThenFrame + The result ThenFrame. + """ return _ffi_api.Then() # type: ignore[attr-defined] # pylint: disable=no-member def Else() -> frame.ElseFrame: # pylint: disable=invalid-name - """Create an else.""" + """Create an else. + + Returns + ------- + res : frame.ElseFrame + The result ElseFrame. + """ return _ffi_api.Else() # type: ignore[attr-defined] # pylint: disable=no-member @@ -2361,7 +2452,19 @@ def buffer_store( def evaluate(value: Expr) -> BypassEmit: - """Emit an evaluation and return a reference to its stored statement.""" + """Emit an evaluation and return a reference to its stored statement. + + Parameters + ---------- + value : Expr + The input expression to evaluate. + + Returns + ------- + result : BypassEmit + A receipt containing the emitted statement, so expression-statement + handling does not emit it again. + """ if isinstance(value, str): value = _StringImm(value) if isinstance(value, bool): @@ -2375,7 +2478,11 @@ def evaluate(value: Expr) -> BypassEmit: def _ffi_name_to_dtype(name: str) -> str: - """Convert an FFI type name to its TVM dtype string.""" + """Convert an FFI type name to its TVM dtype string. + + Examples: "Float32" -> "float32", "Int8x4" -> "int8x4", + "Float8E4M3" -> "float8_e4m3", "Float8E4M3B11FNUZ" -> "float8_e4m3b11fnuz". + """ import re # Insert underscore before E-notation in float8 names (E3M4, E4M3, etc.) @@ -2384,7 +2491,13 @@ def _ffi_name_to_dtype(name: str) -> str: def func_gen(name: str): - """Generate a DtypeConstructor for each Expr dtype.""" + """Generate a DtypeConstructor for each Expr dtype. + + Parameters + ---------- + name: str + The ffi function name to call, e.g. "Float32", "Int32". + """ return DtypeConstructor(name, _ffi_name_to_dtype(name)) @@ -2631,7 +2744,12 @@ def handle( def TensorMap() -> Var: # pylint: disable=invalid-name - """Create a TIRx var that represents a CUDA tensor-map descriptor.""" + """Create a TIRx var that represents a CUDA tensor-map descriptor. + + The host/runtime ABI passes a handle to descriptor storage. CUDA kernel + codegen lowers this type to ``const __grid_constant__ CUtensorMap`` when it + appears as a kernel parameter. + """ return _ffi_api.TensorMap() # type: ignore[attr-defined] # pylint: disable=no-member @@ -2910,7 +3028,10 @@ else: return cls def meta_class(cls): - """Decorator for utility classes used inside @T.prim_func.""" + """Decorator for utility classes used inside @T.prim_func. + + Instances of decorated classes are treated as parser meta values. + """ return _install_meta_class(cls) # pylint: disable=invalid-name diff --git a/python/tvm/tirx/script/builder/tirx.py b/python/tvm/tirx/script/builder/tirx.py index ec873276b0..851c56740b 100644 --- a/python/tvm/tirx/script/builder/tirx.py +++ b/python/tvm/tirx/script/builder/tirx.py @@ -32,7 +32,13 @@ from .ir import decl_buffer, meta_class def _normalize_scope(scope) -> ExecScope: - """Normalize a scope selector to an ``ExecScope``.""" + """Normalize a scope selector to an ``ExecScope``. + + Accepts an ``ExecScope`` (passed through), a scope-name ``str`` + (e.g. ``"warp"``, normalized via the FFI ctor / ``StringToScopeKind``), + or an ``int`` ``ScopeKind`` value. ``None`` resolves to the default + ``thread`` scope, keeping the default in one place. + """ if scope is None: return ExecScope("thread") if isinstance(scope, ExecScope): @@ -63,7 +69,10 @@ class ScopedOp: return self._fn(*args, scope=ExecScope("thread"), **kwargs) def _bind(self, scope: ExecScope): - """Return a callable that emits this op at ``scope``.""" + """Return a callable that emits this op at ``scope``. + + Used by :class:`ScopeNamespace`; not part of the user-facing surface. + """ return lambda *args, **kwargs: self._fn(*args, scope=scope, **kwargs) @@ -455,7 +464,12 @@ def cast( scope: ExecScope | None = None, **kwargs, ): - """Cast — overloaded.""" + """Cast — overloaded. + + 1. ``cast(value, dtype)`` — expression-level cast: returns ``T.cast(value, dtype)``. + Also accepts ``cast(value, dtype=...)`` as a kwarg form. + 2. ``cast(dst, src, workspace=..., dispatch=...)`` — buffer-level Cast operator. + """ # Expression-level cast: src is a dtype (str / DataType) — emit T.cast(value, dtype). from tvm import tirx as _tirx @@ -879,7 +893,11 @@ def max( scope: ExecScope | None = None, **kwargs, ): - """Max — overloaded.""" + """Max — overloaded. + + 1. ``max(a, b)`` — expression: returns ``tirx.max(a, b)``. + 2. ``max(dst, src, axes=, accum=)`` — reduction operator over buffers. + """ from tvm import tirx as _tirx if not _is_buffer_or_region(dst) or not _is_buffer_or_region(src): @@ -916,7 +934,11 @@ def min( scope: ExecScope | None = None, **kwargs, ): - """Min — overloaded.""" + """Min — overloaded. + + 1. ``min(a, b)`` — expression: returns ``tirx.min(a, b)``. + 2. ``min(dst, src, axes=, accum=)`` — reduction operator over buffers. + """ from tvm import tirx as _tirx if not _is_buffer_or_region(dst) or not _is_buffer_or_region(src): diff --git a/python/tvm/tirx/script/jit.py b/python/tvm/tirx/script/jit.py index 19175c9149..04f826c9f9 100644 --- a/python/tvm/tirx/script/jit.py +++ b/python/tvm/tirx/script/jit.py @@ -39,6 +39,45 @@ def make_jit(builder): """Create a JIT decorator for the canonical construction namespace.""" def jit(func=None, *, private=False, check_well_formed=True, is_stir=False, persistent=False): + """Decorator: capture the kernel and defer parsing until ``.specialize()``. + + Use ``@T.jit`` (instead of ``@T.prim_func``) when the kernel takes + compile-time parameters annotated with ``T.constexpr`` or runtime + parameters that may be removed with ``T.Optional``. The resulting object + exposes ``.specialize(**specialization_kwargs)``, which returns a + ``tvm.tirx.PrimFunc``. + + Example:: + + from __future__ import annotations + + from tvm.script import tirx as T + + @T.jit + def add( + A: T.Buffer((N,), "float32"), + B: T.Buffer((N,), "float32"), + *, + N: T.constexpr, + ): + for i in T.serial(N): + B[i] = A[i] + 1.0 + + kernel = add.specialize(N=1024) # returns a PrimFunc + + @T.jit + def guarded(optional: T.Optional(T.handle), out: T.handle): + output = T.match_buffer(out, (1,), "int32") + if T.constexpr(optional is not None): + value = T.match_buffer(optional, (1,), "int32") + output[0] = value[0] + else: + output[0] = 0 + + present = guarded.specialize() + absent = guarded.specialize(optional=None) + """ + def apply(function): from tvm.script.parser.entry import _definition_scope
