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
15 changes: 15 additions & 0 deletions docs/arch/tvmscript.rst
Original file line number Diff line number Diff line change
Expand Up @@ -136,6 +136,21 @@ underlying builders can also be used directly from Python. The transpiler theref
needs no separate mutable IR representation: it generates calls to this construction
protocol.

Ordinary IR constructors can be used directly in parsed source. Shared exports
such as ``I.Call`` and ``T.Range`` use the same constructor contracts as
``tvm.ir.Call`` and ``tvm.ir.Range``, including keyword arguments, source spans
and validation. ``Call(..., ty=...)`` supplies an explicit result type; omission
leaves ``Type.missing()`` for subsequent normalization. Use ``Call.unchecked``
explicitly for provisional calls that require later validation. Raw printed calls
use this form with their stored result type to preserve all fields.

Operations likewise retain their normal argument contracts. A dtype inferred from
operands is not an extra ``dtype`` keyword; operations with an explicit dtype
parameter accept it normally. Annotation shorthand, module metadata selectors,
frame construction and variadic dtype positioning remain explicit builder adapters.
They do not change the underlying IR constructor's validation or consult builder
state from ordinary construction.

Printing and round trips
------------------------

Expand Down
2 changes: 1 addition & 1 deletion docs/reference/api/python/script/ir_builder.rst
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ tvm.script.ir_builder
.. automodule:: tvm.script.ir_builder
:members:
:imported-members:
:exclude-members: DataTypeImm, FuncType, GenericConst, PrimType, Range, StringImm, StringType, Type
:exclude-members: Call, DataTypeImm, FuncType, GenericConst, PrimType, Range, StringImm, StringType, Type

tvm.relax.script.ir_builder.distributed
***************************************
Expand Down
2 changes: 1 addition & 1 deletion python/tvm/backend/cuda/script.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,7 @@ def __call__(self, *args, **kwds):
return tvm.ir.Call(
tvm.ir.Op.get("tirx.s_tir.cp_async_raw"),
[dst, dst_off, src, src_off, cp_size],
ret_ty=tvm.ir.PrimType(elem_dtype),
ty=tvm.ir.PrimType(elem_dtype),
)
raise TypeError(
"T.s_tir.cp_async_raw only accepts the printed 6-arg raw form; "
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -692,7 +692,7 @@ def _get_or_create_desc(sctx, s_buf, ldo, sdo, swizzle):
StringImm(desc_buf.scope()),
],
attrs=DictAttrs({}),
ret_ty=desc_buf.ty,
ty=desc_buf.ty,
),
),
Evaluate(encode_call),
Expand Down
6 changes: 3 additions & 3 deletions python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py
Original file line number Diff line number Diff line change
Expand Up @@ -1099,7 +1099,7 @@ def _make_lo_uniform(desc_buf):
StringImm(desc_lo.scope()),
],
attrs=DictAttrs({}),
ret_ty=desc_lo.ty,
ty=desc_lo.ty,
),
),
Bind(
Expand All @@ -1112,7 +1112,7 @@ def _make_lo_uniform(desc_buf):
StringImm(desc_hi.scope()),
],
attrs=DictAttrs({}),
ret_ty=desc_hi.ty,
ty=desc_hi.ty,
),
),
Evaluate(unpack),
Expand Down Expand Up @@ -1147,7 +1147,7 @@ def _make_desc(smem_buf, ldo, sdo, swizzle_val, name):
StringImm(desc_buf.scope()),
],
attrs=DictAttrs({}),
ret_ty=desc_buf.ty,
ty=desc_buf.ty,
),
),
Evaluate(encode_call),
Expand Down
2 changes: 1 addition & 1 deletion python/tvm/backend/trn/transform/naive_allocator.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,7 +98,7 @@ def allocate_buffer(op: Bind):
attrs=op.value.attrs,
ty_args=op.value.ty_args,
span=op.value.span,
ret_ty=new_buffer.ty,
ty=new_buffer.ty,
),
op.span,
)
Expand Down
2 changes: 1 addition & 1 deletion python/tvm/backend/trn/transform/private_buffer_alloc.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,7 @@ def visit_attr(op: AttrStmt):
StringImm(buffer.scope()),
],
attrs=DictAttrs({}),
ret_ty=buffer.ty,
ty=buffer.ty,
),
)
body = SeqStmt([allocation, body])
Expand Down
28 changes: 14 additions & 14 deletions python/tvm/ir/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -503,7 +503,7 @@ def __ne__(self, other) -> bool:
class Call(_CallableExprWithOp):
"""Core function call node.

When ``ret_ty`` is omitted, use a missing type for subsequent normalization.
When ``ty`` is omitted, use a missing type for subsequent normalization.
Builders may supply a known result type explicitly.
Operator validation runs during construction. Use :meth:`unchecked` for a
provisional call that will be checked after normalization.
Expand All @@ -522,14 +522,14 @@ def __init__(
attrs: "tvm.ir.Attrs | dict | None" = None,
ty_args: list["tvm.ir.Type"] | tuple["tvm.ir.Type", ...] | None = None,
span: Span | None = None,
ret_ty: "tvm.ir.Type | str | None" = None,
ty: "tvm.ir.Type | str | None" = None,
) -> None:
self.__init_handle_by_constructor__(
_ffi_api.Call, *self._normalize_constructor_args(op, args, attrs, ty_args, span, ret_ty)
_ffi_api.Call, *self._normalize_constructor_args(op, args, attrs, ty_args, span, ty)
)

@staticmethod
def _normalize_constructor_args(op, args, attrs, ty_args, span, ret_ty):
def _normalize_constructor_args(op, args, attrs, ty_args, span, ty):
# pylint: disable=import-outside-toplevel
from .attrs import DictAttrs
from .op import Op
Expand All @@ -539,15 +539,15 @@ def _normalize_constructor_args(op, args, attrs, ty_args, span, ret_ty):
op = Op.get(op)
if attrs is not None and isinstance(attrs, dict):
attrs = DictAttrs(attrs)
if ret_ty is None:
ret_ty = Type.missing()
if isinstance(ret_ty, str) and ret_ty == "handle":
ret_ty = PointerType(PrimType("void"))
elif ret_ty is not None and not isinstance(ret_ty, Type):
ret_ty = PrimType(ret_ty)
if ty is None:
ty = Type.missing()
if isinstance(ty, str) and ty == "handle":
ty = PointerType(PrimType("void"))
elif ty is not None and not isinstance(ty, Type):
ty = PrimType(ty)
if ty_args is None:
ty_args = []
return ret_ty, op, args, attrs, ty_args, span
return ty, op, args, attrs, ty_args, span

@staticmethod
def unchecked(
Expand All @@ -556,7 +556,7 @@ def unchecked(
attrs: "tvm.ir.Attrs | dict | None" = None,
ty_args: list["tvm.ir.Type"] | tuple["tvm.ir.Type", ...] | None = None,
span: Span | None = None,
ret_ty: "tvm.ir.Type | str | None" = None,
ty: "tvm.ir.Type | str | None" = None,
) -> "Call":
"""Construct a provisional Call without invoking its Op validator.

Expand All @@ -578,7 +578,7 @@ def unchecked(
Explicit type arguments; ``None`` gives an empty list.
span : Span or None
Source location of the Call.
ret_ty : tvm.ir.Type, str, or None
ty : tvm.ir.Type, str, or None
Result type. ``None`` uses ``Type.missing()`` for later inference;
``"handle"`` becomes a void pointer type and other strings become
primitive types.
Expand All @@ -589,7 +589,7 @@ def unchecked(
The provisional, unvalidated Call.
"""
return _ffi_api.CallUnchecked(
*Call._normalize_constructor_args(op, args, attrs, ty_args, span, ret_ty)
*Call._normalize_constructor_args(op, args, attrs, ty_args, span, ty)
)


Expand Down
2 changes: 1 addition & 1 deletion python/tvm/ir/prim/op.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,4 +80,4 @@ def clz(x):
y : Expr
The result.
"""
return Call("prim.clz", [x], ret_ty="int32")
return Call("prim.clz", [x], ty="int32")
2 changes: 1 addition & 1 deletion python/tvm/relax/frontend/nn/llm/_decode_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -329,7 +329,7 @@ def batch_decode_paged_kv(
"reduce_scope",
T.int32(0),
)
T.tvm_thread_allreduce(T.uint32(1), S_reduce_local[0], True, t0[0], tx, dtype="void")
T.tvm_thread_allreduce(T.uint32(1), 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
46 changes: 3 additions & 43 deletions python/tvm/relax/script/ir_builder/ir.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,6 @@
import numbers as _numbers
import re as _re

import tvm
from tvm import ir as _ir
from tvm import relax
from tvm import relax as _relax
Expand Down Expand Up @@ -248,48 +247,9 @@ def tuple(*fields: Expr) -> Expr:
return relax.Tuple(fields) # type: ignore[attr-defined] # pylint: disable=no-member


def shape(value: list[Expr]) -> Expr:
"""Create a ShapeExpr.
Parameters
----------
value : List[Expr]
The fields of the tuple.
Returns
-------
res : Expr
The result tuple.
"""
return relax.ShapeExpr(value) # pylint: disable=no-member # type: ignore


def prim_value(value: Expr | int | float) -> Expr:
"""Convert a value to a primitive expression.

Parameters
----------
value : Expr | int | float
The value to convert.

Returns
-------
res : Expr
The primitive expression.
"""
return relax.prim_value(value) # type: ignore[attr-defined] # pylint: disable=no-member


def str(value: py_str) -> Expr:
"""Create a string imm expression.
Parameters
----------
value : str
The value of the str.
Returns
-------
res : Expr
The result str.
"""
return tvm.ir.StringImm(value) # type: ignore[attr-defined] # pylint: disable=no-member
shape = ShapeExpr
prim_value = relax.prim_value
str = _ir.StringImm


@_resolve_global_info_args("vdevice", resolver=resolve_global_info_)
Expand Down
4 changes: 0 additions & 4 deletions python/tvm/s_tir/tensor_intrin/cuda.py
Original file line number Diff line number Diff line change
Expand Up @@ -851,7 +851,6 @@ def wmma_load_impl(
A.access_ptr("r", ptr_type=dtype),
s1,
layout,
dtype="void",
)
)

Expand Down Expand Up @@ -904,7 +903,6 @@ def wmma_fill_impl(
k_dim,
get_wmma_fragment_index(C, d1, m_dim, n_dim),
T.float32(0),
dtype="void",
)
)

Expand Down Expand Up @@ -969,7 +967,6 @@ def wmma_store_impl(
C.access_ptr("w", ptr_type=dtype),
s1,
"row_major",
dtype="void",
)
)

Expand Down Expand Up @@ -1075,7 +1072,6 @@ def wmma_sync_impl(
get_wmma_fragment_index(B, b1, b_shape_0, b_shape_1),
C.data,
get_wmma_fragment_index(C, c1, m_dim, n_dim),
dtype="void",
)
)

Expand Down
11 changes: 4 additions & 7 deletions python/tvm/s_tir/tensor_intrin/hexagon.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,26 +54,23 @@ def sync_dma_load_impl(
T.tvm_call_packed(
"device_api.hexagon.dma_copy_dltensor",
T.tvm_stack_make_array(
T.address_of(C[0], dtype="handle"),
T.tvm_stack_make_shape(size, dtype="handle"),
T.address_of(C[0]),
T.tvm_stack_make_shape(size),
0,
1,
C.dtype,
0,
dtype="handle",
),
T.tvm_stack_make_array(
T.address_of(A[0], dtype="handle"),
T.tvm_stack_make_shape(size, dtype="handle"),
T.address_of(A[0]),
T.tvm_stack_make_shape(size),
0,
1,
A.dtype,
0,
dtype="handle",
),
T.cast(size, dtype="int"),
False, # Do not use experimental bypass mode.
dtype="int32",
)
)

Expand Down
20 changes: 1 addition & 19 deletions python/tvm/script/ir_builder/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,8 @@
# under the License.
"""Shared TVMScript construction APIs and lazy language variant builders."""

from tvm.ir import Call as _IRCall
from tvm.ir import (
Call,
DataTypeImm,
FuncType,
GenericConst,
Expand All @@ -27,7 +27,6 @@
StringType,
Type,
make_node,
reinfer_type,
)

from .base import (
Expand Down Expand Up @@ -58,23 +57,6 @@
resolve_global_info_,
)


def Call(op, args, attrs=None, ty_args=None, ty=None): # pylint: disable=invalid-name
"""Construct a shared IR Call, inferring its result type when ``ty`` is omitted.

``ty_args`` are independent explicit inputs to the inference hook.
An explicit ``ty`` preserves the supplied Call fields without invoking its
operator validator, including provisional calls with missing input types.
Omitting ``ty`` requests inference followed by ordinary Call validation.
"""

if ty is None:
provisional = _IRCall.unchecked(op, args, attrs=attrs, ty_args=ty_args)
ty = reinfer_type(provisional)
return _IRCall(op, args, attrs=attrs, ty_args=ty_args, ret_ty=ty)
return _IRCall.unchecked(op, args, attrs=attrs, ty_args=ty_args, ret_ty=ty)


# Keep source namespaces independent of imported helper modules and lazy builders.
__all__ = [
"MISSING",
Expand Down
Loading
Loading