From c6c11303f520229eaf34f7fb481262e893a561f5 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Fri, 2 Oct 2026 18:08:37 +0000 Subject: [PATCH 1/2] [REFACTOR][SCRIPT] Share normal IR constructor contracts --- docs/arch/tvmscript.rst | 15 ++ python/tvm/backend/cuda/script.py | 2 +- .../tile_primitive/copy_async/tcgen05_cp.py | 2 +- .../cuda/tile_primitive/gemm_async/tcgen05.py | 6 +- .../backend/trn/transform/naive_allocator.py | 2 +- .../trn/transform/private_buffer_alloc.py | 2 +- python/tvm/ir/expr.py | 28 +-- python/tvm/ir/prim/op.py | 2 +- .../relax/frontend/nn/llm/_decode_kernels.py | 2 +- python/tvm/relax/script/ir_builder/ir.py | 46 +--- python/tvm/s_tir/tensor_intrin/cuda.py | 4 - python/tvm/s_tir/tensor_intrin/hexagon.py | 11 +- python/tvm/script/ir_builder/__init__.py | 20 +- python/tvm/tirx/op.py | 30 +-- python/tvm/tirx/script/ir_builder/ir.py | 19 +- python/tvm/tirx/script/ir_builder/op.py | 224 +++++++++--------- python/tvm/topi/gpu/sort.py | 6 +- src/relax/script/printer/call.cc | 20 +- src/s_tir/script/printer/expr.cc | 2 +- src/script/printer/ir/utils.cc | 13 +- src/script/printer/ir/utils.h | 2 +- src/tirx/script/printer/buffer.cc | 18 +- src/tirx/script/printer/expr.cc | 38 +-- .../codegen/test_target_codegen_llvm.py | 8 +- .../python/relax/test_analysis_well_formed.py | 4 +- tests/python/relax/test_expr.py | 8 +- .../test_transform_annotate_tir_op_pattern.py | 3 +- tests/python/relax/test_transform_fuse_ops.py | 2 +- .../test_transform_legalize_ops_image.py | 2 +- .../relax/test_transform_legalize_ops_nn.py | 18 +- .../test_transform_rewrite_cuda_graph.py | 8 +- .../analysis/test_sblock_access_region.py | 4 +- .../analysis/test_sblock_buffer_access_lca.py | 2 +- ...ule_feature_extractor_per_store_feature.py | 2 +- ...t_meta_schedule_postproc_rewrite_layout.py | 6 +- ...hedule_postproc_rewrite_reduction_block.py | 4 +- ...schedule_postproc_rewrite_unbound_block.py | 4 +- ..._meta_schedule_postproc_verify_gpu_code.py | 24 +- ...meta_schedule_schedule_rule_auto_inline.py | 18 +- ...le_schedule_rule_cross_thread_reduction.py | 39 ++- ...schedule_rule_parallel_vectorize_unroll.py | 14 +- .../test_meta_schedule_trace_apply.py | 36 +-- .../schedule/test_tir_schedule_cache_index.py | 17 -- .../test_tir_schedule_cache_read_write.py | 12 - .../schedule/test_tir_schedule_compute_at.py | 30 ++- .../test_tir_schedule_compute_inline.py | 12 +- .../test_tir_schedule_decompose_padding.py | 5 +- .../schedule/test_tir_schedule_pad_einsum.py | 6 +- .../schedule/test_tir_schedule_partition.py | 2 +- .../schedule/test_tir_schedule_reindex.py | 3 - .../schedule/test_tir_schedule_reorder.py | 4 +- .../schedule/test_tir_schedule_rfactor.py | 12 +- .../schedule/test_tir_schedule_sampling.py | 1 - .../schedule/test_tir_schedule_split_fuse.py | 6 +- .../test_tir_schedule_state_cached_flags.py | 16 +- .../schedule/test_tir_schedule_tensorize.py | 10 +- .../test_tir_schedule_transform_layout.py | 23 +- .../s_tir/script/test_s_tir_script_blocks.py | 2 +- .../s_tir/script/test_s_tir_script_printer.py | 86 ++----- ...t_s_tir_transform_compact_buffer_region.py | 34 ++- .../test_s_tir_transform_hoist_expression.py | 4 +- ...t_s_tir_transform_inject_ptx_async_copy.py | 4 +- ..._tir_transform_inject_software_pipeline.py | 2 +- ..._transform_lower_cross_thread_reduction.py | 44 +--- ...tir_transform_memhammer_lower_auto_copy.py | 4 - ..._plan_update_buffer_allocation_location.py | 4 +- .../test_s_tir_transform_remove_undef.py | 12 +- ...tir_transform_renormalize_split_pattern.py | 12 +- .../script/test_constructor_contracts.py | 114 +++++++++ tests/python/script/test_source_locations.py | 4 +- tests/python/te/test_te_create_primfunc.py | 13 +- .../test_tir_analysis_undefined_vars.py | 6 +- .../test_binding_type_propagation.py | 6 +- tests/python/tirx-base/test_tir_buffer.py | 2 +- .../python/tirx-base/test_tir_constructor.py | 14 +- .../test_tir_inline_private_functions.py | 3 +- .../test_tir_transform_convert_ssa.py | 4 +- .../test_tir_transform_narrow_datatype.py | 3 +- .../test_tir_transform_simplify.py | 16 +- .../test_tir_transform_storage_rewrite.py | 4 +- .../python/tirx/codegen/test_codegen_cuda.py | 2 +- .../tirx/test_buffer_data_reinfer_type.py | 2 +- 82 files changed, 590 insertions(+), 690 deletions(-) create mode 100644 tests/python/script/test_constructor_contracts.py diff --git a/docs/arch/tvmscript.rst b/docs/arch/tvmscript.rst index d7d64fb3d2d4..fc29376a8d2d 100644 --- a/docs/arch/tvmscript.rst +++ b/docs/arch/tvmscript.rst @@ -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 ------------------------ diff --git a/python/tvm/backend/cuda/script.py b/python/tvm/backend/cuda/script.py index 28f6109b9929..048805fb2324 100644 --- a/python/tvm/backend/cuda/script.py +++ b/python/tvm/backend/cuda/script.py @@ -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; " diff --git a/python/tvm/backend/cuda/tile_primitive/copy_async/tcgen05_cp.py b/python/tvm/backend/cuda/tile_primitive/copy_async/tcgen05_cp.py index 86117c69a2ea..32b14cd2cf2c 100644 --- a/python/tvm/backend/cuda/tile_primitive/copy_async/tcgen05_cp.py +++ b/python/tvm/backend/cuda/tile_primitive/copy_async/tcgen05_cp.py @@ -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), diff --git a/python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py b/python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py index f108b673c73d..d76eda7c08fe 100644 --- a/python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py +++ b/python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py @@ -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( @@ -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), @@ -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), diff --git a/python/tvm/backend/trn/transform/naive_allocator.py b/python/tvm/backend/trn/transform/naive_allocator.py index 4bf6d1403e4c..4a768558d9f9 100644 --- a/python/tvm/backend/trn/transform/naive_allocator.py +++ b/python/tvm/backend/trn/transform/naive_allocator.py @@ -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, ) diff --git a/python/tvm/backend/trn/transform/private_buffer_alloc.py b/python/tvm/backend/trn/transform/private_buffer_alloc.py index a74d283d10a1..ee3a01254299 100644 --- a/python/tvm/backend/trn/transform/private_buffer_alloc.py +++ b/python/tvm/backend/trn/transform/private_buffer_alloc.py @@ -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]) diff --git a/python/tvm/ir/expr.py b/python/tvm/ir/expr.py index 9c261c38c366..177029702ece 100644 --- a/python/tvm/ir/expr.py +++ b/python/tvm/ir/expr.py @@ -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. @@ -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 @@ -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( @@ -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. @@ -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. @@ -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) ) diff --git a/python/tvm/ir/prim/op.py b/python/tvm/ir/prim/op.py index 0562e888a446..da4d4b04e314 100644 --- a/python/tvm/ir/prim/op.py +++ b/python/tvm/ir/prim/op.py @@ -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") diff --git a/python/tvm/relax/frontend/nn/llm/_decode_kernels.py b/python/tvm/relax/frontend/nn/llm/_decode_kernels.py index d5098a35c585..e7671175239a 100644 --- a/python/tvm/relax/frontend/nn/llm/_decode_kernels.py +++ b/python/tvm/relax/frontend/nn/llm/_decode_kernels.py @@ -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]: diff --git a/python/tvm/relax/script/ir_builder/ir.py b/python/tvm/relax/script/ir_builder/ir.py index 195d0b79c344..42d491a92502 100644 --- a/python/tvm/relax/script/ir_builder/ir.py +++ b/python/tvm/relax/script/ir_builder/ir.py @@ -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 @@ -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_) diff --git a/python/tvm/s_tir/tensor_intrin/cuda.py b/python/tvm/s_tir/tensor_intrin/cuda.py index 23011d5330de..f94edb47c399 100644 --- a/python/tvm/s_tir/tensor_intrin/cuda.py +++ b/python/tvm/s_tir/tensor_intrin/cuda.py @@ -851,7 +851,6 @@ def wmma_load_impl( A.access_ptr("r", ptr_type=dtype), s1, layout, - dtype="void", ) ) @@ -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", ) ) @@ -969,7 +967,6 @@ def wmma_store_impl( C.access_ptr("w", ptr_type=dtype), s1, "row_major", - dtype="void", ) ) @@ -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", ) ) diff --git a/python/tvm/s_tir/tensor_intrin/hexagon.py b/python/tvm/s_tir/tensor_intrin/hexagon.py index b300c0232864..a1c477dfe27e 100644 --- a/python/tvm/s_tir/tensor_intrin/hexagon.py +++ b/python/tvm/s_tir/tensor_intrin/hexagon.py @@ -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", ) ) diff --git a/python/tvm/script/ir_builder/__init__.py b/python/tvm/script/ir_builder/__init__.py index 90743bde5f62..9e1f1f7f0fd9 100644 --- a/python/tvm/script/ir_builder/__init__.py +++ b/python/tvm/script/ir_builder/__init__.py @@ -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, @@ -27,7 +27,6 @@ StringType, Type, make_node, - reinfer_type, ) from .base import ( @@ -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", diff --git a/python/tvm/tirx/op.py b/python/tvm/tirx/op.py index 4eff21210ae8..917834b75036 100644 --- a/python/tvm/tirx/op.py +++ b/python/tvm/tirx/op.py @@ -138,14 +138,14 @@ def _pack_buffer(buf, span=None): "tirx.tvm_stack_make_shape", buf.ty.shape, span=span, - ret_ty=PointerType(tvm.ir.PrimType("int64")), + ty=PointerType(tvm.ir.PrimType("int64")), ) strides = ( Call( "tirx.tvm_stack_make_shape", buf.ty.strides, span=span, - ret_ty=PointerType(tvm.ir.PrimType("int64")), + ty=PointerType(tvm.ir.PrimType("int64")), ) if buf.ty.strides else 0 @@ -158,7 +158,7 @@ def _pack_buffer(buf, span=None): const(0, dtype=buf.ty.dtype), buf.ty.elem_offset, ] - return Call(Op.get("tirx.tvm_stack_make_array"), pack_args, span=span, ret_ty="handle") + return Call(Op.get("tirx.tvm_stack_make_array"), pack_args, span=span, ty="handle") def call_packed_lowered(*args, span=None): @@ -190,7 +190,7 @@ def call_packed_lowered(*args, span=None): _pack_buffer(x) if is_buffer_var(x) else _reject_buffer_region(x, "call_packed_lowered") for x in args ] - return Call(Op.get("tirx.tvm_call_packed_lowered"), call_args, span=span, ret_ty="int32") + return Call(Op.get("tirx.tvm_call_packed_lowered"), call_args, span=span, ty="int32") def call_cpacked_lowered(*args, span=None): @@ -219,7 +219,7 @@ def call_cpacked_lowered(*args, span=None): _pack_buffer(x) if is_buffer_var(x) else _reject_buffer_region(x, "call_cpacked_lowered") for x in args ] - return Call(Op.get("tirx.tvm_call_cpacked_lowered"), call_args, span=span, ret_ty="int32") + return Call(Op.get("tirx.tvm_call_cpacked_lowered"), call_args, span=span, ty="int32") def call_packed(*args, span=None): @@ -253,7 +253,7 @@ def call_packed(*args, span=None): _pack_buffer(x) if is_buffer_var(x) else _reject_buffer_region(x, "call_packed") for x in args ] - return Call(Op.get("tirx.tvm_call_packed"), call_args, span=span, ret_ty="int32") + return Call(Op.get("tirx.tvm_call_packed"), call_args, span=span, ty="int32") @tvm_ffi.register_object("tirx.CallFFIKernelAttr") @@ -278,7 +278,7 @@ def call_ffi_kernel(*args, launch_params, ret_ty="int32", span=None): "tirx.call_ffi_kernel", args, attrs=CallFFIKernelAttr(launch_params), - ret_ty=ret_ty, + ty=ret_ty, span=span, ) @@ -335,7 +335,7 @@ def tensormap_encode_tiled( attrs=TensorMapEncodeTiledAttr( descriptor_dtype, rank, interleave, swizzle, l2_promotion, oob_fill, force_cu_dtype ), - ret_ty="int32", + ty="int32", span=span, ) @@ -367,7 +367,7 @@ def call_cpacked(*args, span=None): _pack_buffer(x) if is_buffer_var(x) else _reject_buffer_region(x, "call_cpacked") for x in args ] - return Call(Op.get("tirx.tvm_call_cpacked"), call_args, span=span, ret_ty="int32") + return Call(Op.get("tirx.tvm_call_cpacked"), call_args, span=span, ty="int32") def call_intrin(dtype: str | tvm.ir.Type, func_name, *args, attrs=None, span=None): @@ -401,7 +401,7 @@ def call_intrin(dtype: str | tvm.ir.Type, func_name, *args, attrs=None, span=Non if isinstance(func_name, str): func_name = _canonical_device_intrin_name(func_name) args = tuple(_reject_buffer_region(arg, "call_intrin") for arg in args) - return Call(func_name, args, attrs=attrs, span=span, ret_ty=dtype) + return Call(func_name, args, attrs=attrs, span=span, ty=dtype) def call_pure_extern(dtype, func_name, *args, span=None): @@ -430,7 +430,7 @@ def call_pure_extern(dtype, func_name, *args, span=None): Op.get("tirx.call_pure_extern"), [func_name, *(_reject_buffer_region(arg, "call_pure_extern") for arg in args)], span=span, - ret_ty=dtype, + ty=dtype, ) @@ -460,7 +460,7 @@ def call_extern(dtype, func_name, *args, span=None): Op.get("tirx.call_extern"), [func_name, *(_reject_buffer_region(arg, "call_extern") for arg in args)], span=span, - ret_ty=dtype, + ty=dtype, ) @@ -781,20 +781,20 @@ def address_of(obj: Buffer | TensorLoad | Var, span: Span | None = None) -> Expr "tirx.address_of", [buffer_load], span=span, - ret_ty=_buffer_element_pointer_type(obj), + ty=_buffer_element_pointer_type(obj), ) elif isinstance(obj, Var): if _is_tensormap_var(obj): return call_intrin("uint64", "tirx.address_of", obj, span=span) if not isinstance(obj.ty, tvm.ir.PrimType): raise TypeError(f"address_of expects a scalar or TensorMap Var, but got {obj.ty}") - return Call("tirx.address_of", [obj], span=span, ret_ty=PointerType(obj.ty)) + return Call("tirx.address_of", [obj], span=span, ty=PointerType(obj.ty)) elif isinstance(obj, TensorLoad): return Call( "tirx.address_of", [obj], span=span, - ret_ty=_buffer_element_pointer_type(obj.source), + ty=_buffer_element_pointer_type(obj.source), ) else: raise ValueError(f"Invalid object type: {type(obj)}") diff --git a/python/tvm/tirx/script/ir_builder/ir.py b/python/tvm/tirx/script/ir_builder/ir.py index 3b7725991b9f..4e8fda52f827 100644 --- a/python/tvm/tirx/script/ir_builder/ir.py +++ b/python/tvm/tirx/script/ir_builder/ir.py @@ -36,7 +36,7 @@ from tvm import DataType, ir from tvm import tirx as tir -from tvm.ir import TensorLoad, Type, is_prim_expr +from tvm.ir import Range, TensorLoad, 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.parser.protocol_registry import ( @@ -641,7 +641,7 @@ def _normalize_ann_value(v): ir.StringImm(buf.scope()), ], attrs=ir.DictAttrs(norm_annotations), - ret_ty=buf.ty, + ty=buf.ty, ) _ffi_api.AddToParent(tir.Bind(buf, allocation)) return buf @@ -1619,21 +1619,6 @@ def target( return Target(target_config, host) -def Range(begin: Expr, end: Expr) -> ir.Range: # pylint: disable=invalid-name - """ - Create a Range object. - - Parameters - ---------- - begin : Expr - The begin value of the range. - - end : Optional[Expr] - The end value of the range. - """ - return ir.Range(begin, end) - - if TYPE_CHECKING: C = TypeVar("C") diff --git a/python/tvm/tirx/script/ir_builder/op.py b/python/tvm/tirx/script/ir_builder/op.py index f2173f361f40..7b145e0590d7 100644 --- a/python/tvm/tirx/script/ir_builder/op.py +++ b/python/tvm/tirx/script/ir_builder/op.py @@ -21,7 +21,7 @@ import functools import inspect from collections.abc import Callable -from typing import Any, ParamSpec, TypeVar +from typing import Any import tvm_ffi as _ffi @@ -95,10 +95,10 @@ def _call_global(func: ir.GlobalVar, *args: Expr) -> Call: declaration = module_frame.functions[func] if isinstance(declaration, tir.PrimFunc): # The Relax-facing signature may erase pointer results to Any. - return Call(func, args, ret_ty=declaration.ret_type) + return Call(func, args, ty=declaration.ret_type) break if isinstance(func.ty, ir.FuncType): - return Call(func, args, ret_ty=func.ty.ret_type) + return Call(func, args, ty=func.ty.ret_type) return Call(func, args) @@ -186,24 +186,14 @@ def comm_reducer(combiner: Callable, identity: list[Expr]) -> CommReducer: return CommReducer(args[: num_args // 2], args[num_args // 2 :], res, identity) -T = TypeVar("T") +def _op_wrapper(func): + """Retain the normal call contract while attaching namespace printer metadata.""" - -P = ParamSpec("P") - - -def _op_wrapper(func: Callable[P, T]) -> Callable[P, T]: @functools.wraps(func) - def wrapped(*args, **kwargs) -> T: - if "dtype" in kwargs: - kwargs.pop("dtype") + def wrapped(*args, **kwargs): return func(*args, **kwargs) - # Expose underlying tir op name for printer registration - try: - wrapped.__tir_op_name__ = getattr(func, "__name__", None) - except Exception: # pragma: no cover - pass + wrapped.__tir_op_name__ = getattr(func, "__name__", None) return wrapped @@ -345,274 +335,274 @@ def _register_tir_namespace_printer_names(): _register_tir_namespace_printer_names() -abs = _op_wrapper(_tir_op.abs) # pylint: disable=redefined-builtin +abs = _tir_op.abs # pylint: disable=redefined-builtin -acos = _op_wrapper(_tir_op.acos) +acos = _tir_op.acos -acosh = _op_wrapper(_tir_op.acosh) +acosh = _tir_op.acosh -address_of = _op_wrapper(_tir_op.address_of) +address_of = _tir_op.address_of -asin = _op_wrapper(_tir_op.asin) +asin = _tir_op.asin -asinh = _op_wrapper(_tir_op.asinh) +asinh = _tir_op.asinh -atan = _op_wrapper(_tir_op.atan) +atan = _tir_op.atan -atan2 = _op_wrapper(_tir_op.atan2) +atan2 = _tir_op.atan2 -atanh = _op_wrapper(_tir_op.atanh) +atanh = _tir_op.atanh -bitwise_and = _op_wrapper(_tir_op.bitwise_and) +bitwise_and = _tir_op.bitwise_and -bitwise_not = _op_wrapper(_tir_op.bitwise_not) +bitwise_not = _tir_op.bitwise_not -bitwise_or = _op_wrapper(_tir_op.bitwise_or) +bitwise_or = _tir_op.bitwise_or -bitwise_xor = _op_wrapper(_tir_op.bitwise_xor) +bitwise_xor = _tir_op.bitwise_xor -ceil = _op_wrapper(_tir_op.ceil) +ceil = _tir_op.ceil -clz = _op_wrapper(_tir_op.clz) +clz = _tir_op.clz -copysign = _op_wrapper(_tir_op.copysign) +copysign = _tir_op.copysign -cos = _op_wrapper(_tir_op.cos) +cos = _tir_op.cos -cosh = _op_wrapper(_tir_op.cosh) +cosh = _tir_op.cosh -erf = _op_wrapper(_tir_op.erf) +erf = _tir_op.erf -exp = _op_wrapper(_tir_op.exp) +exp = _tir_op.exp -exp2 = _op_wrapper(_tir_op.exp2) +exp2 = _tir_op.exp2 -exp10 = _op_wrapper(_tir_op.exp10) +exp10 = _tir_op.exp10 -filter = _op_wrapper(_tir_op.filter) # pylint: disable=redefined-builtin +filter = _tir_op.filter # pylint: disable=redefined-builtin -selector = _op_wrapper(_tir_op.selector) +selector = _tir_op.selector -floor = _op_wrapper(_tir_op.floor) +floor = _tir_op.floor -ceildiv = _op_wrapper(_tir_op.ceildiv) +ceildiv = _tir_op.ceildiv -floordiv = _op_wrapper(_tir_op.floordiv) +floordiv = _tir_op.floordiv -floormod = _op_wrapper(_tir_op.floormod) +floormod = _tir_op.floormod -fmod = _op_wrapper(_tir_op.fmod) +fmod = _tir_op.fmod -fma = _op_wrapper(_tir_op.fma) +fma = _tir_op.fma -hypot = _op_wrapper(_tir_op.hypot) +hypot = _tir_op.hypot -if_then_else = _op_wrapper(_tir_op.if_then_else) +if_then_else = _tir_op.if_then_else -infinity = _op_wrapper(_tir_op.infinity) +infinity = _tir_op.infinity -isfinite = _op_wrapper(_tir_op.isfinite) +isfinite = _tir_op.isfinite -isinf = _op_wrapper(_tir_op.isinf) +isinf = _tir_op.isinf -isnan = _op_wrapper(_tir_op.isnan) +isnan = _tir_op.isnan -isnullptr = _op_wrapper(_tir_op.isnullptr) +isnullptr = _tir_op.isnullptr -ldexp = _op_wrapper(_tir_op.ldexp) +ldexp = _tir_op.ldexp -likely = _op_wrapper(_tir_op.likely) +likely = _tir_op.likely -log = _op_wrapper(_tir_op.log) +log = _tir_op.log -log1p = _op_wrapper(_tir_op.log1p) +log1p = _tir_op.log1p -log2 = _op_wrapper(_tir_op.log2) +log2 = _tir_op.log2 -log10 = _op_wrapper(_tir_op.log10) +log10 = _tir_op.log10 -max_value = _op_wrapper(_tir_op.max_value) +max_value = _tir_op.max_value -min_value = _op_wrapper(_tir_op.min_value) +min_value = _tir_op.min_value -nearbyint = _op_wrapper(_tir_op.nearbyint) +nearbyint = _tir_op.nearbyint -nextafter = _op_wrapper(_tir_op.nextafter) +nextafter = _tir_op.nextafter -popcount = _op_wrapper(_tir_op.popcount) +popcount = _tir_op.popcount -pow = _op_wrapper(_tir_op.pow) # pylint: disable=redefined-builtin +pow = _tir_op.pow # pylint: disable=redefined-builtin -q_multiply_shift = _op_wrapper(_tir_op.q_multiply_shift) +q_multiply_shift = _tir_op.q_multiply_shift -q_multiply_shift_per_axis = _op_wrapper(_tir_op.q_multiply_shift_per_axis) +q_multiply_shift_per_axis = _tir_op.q_multiply_shift_per_axis -round = _op_wrapper(_tir_op.round) # pylint: disable=redefined-builtin +round = _tir_op.round # pylint: disable=redefined-builtin -rsqrt = _op_wrapper(_tir_op.rsqrt) +rsqrt = _tir_op.rsqrt -shift_left = _op_wrapper(_tir_op.shift_left) +shift_left = _tir_op.shift_left -shift_right = _op_wrapper(_tir_op.shift_right) +shift_right = _tir_op.shift_right -sigmoid = _op_wrapper(_tir_op.sigmoid) +sigmoid = _tir_op.sigmoid -sin = _op_wrapper(_tir_op.sin) +sin = _tir_op.sin -sinh = _op_wrapper(_tir_op.sinh) +sinh = _tir_op.sinh -sqrt = _op_wrapper(_tir_op.sqrt) +sqrt = _tir_op.sqrt -tan = _op_wrapper(_tir_op.tan) +tan = _tir_op.tan -tanh = _op_wrapper(_tir_op.tanh) +tanh = _tir_op.tanh -thread_return = _op_wrapper(_tir_op.thread_return) +thread_return = _tir_op.thread_return -trunc = _op_wrapper(_tir_op.trunc) +trunc = _tir_op.trunc -truncdiv = _op_wrapper(_tir_op.truncdiv) +truncdiv = _tir_op.truncdiv -truncmod = _op_wrapper(_tir_op.truncmod) +truncmod = _tir_op.truncmod -tvm_access_ptr = _op_wrapper(_tir_op.tvm_access_ptr) +tvm_access_ptr = _tir_op.tvm_access_ptr -ptr_byte_offset = _op_wrapper(_tir_op.ptr_byte_offset) +ptr_byte_offset = _tir_op.ptr_byte_offset -tvm_throw_last_error = _op_wrapper(_tir_op.tvm_throw_last_error) +tvm_throw_last_error = _tir_op.tvm_throw_last_error -print_buffer = _op_wrapper(_tir_op.print_buffer) +print_buffer = _tir_op.print_buffer -tvm_stack_alloca = _op_wrapper(_tir_op.tvm_stack_alloca) +tvm_stack_alloca = _tir_op.tvm_stack_alloca -tvm_stack_make_shape = _op_wrapper(_tir_op.tvm_stack_make_shape) +tvm_stack_make_shape = _tir_op.tvm_stack_make_shape -tvm_stack_make_array = _op_wrapper(_tir_op.tvm_stack_make_array) +tvm_stack_make_array = _tir_op.tvm_stack_make_array -call_packed = _op_wrapper(_tir_op.call_packed) +call_packed = _tir_op.call_packed -call_ffi_kernel = _op_wrapper(_tir_op.call_ffi_kernel) +call_ffi_kernel = _tir_op.call_ffi_kernel -tensormap_encode_tiled = _op_wrapper(_tir_op.tensormap_encode_tiled) +tensormap_encode_tiled = _tir_op.tensormap_encode_tiled -call_cpacked = _op_wrapper(_tir_op.call_cpacked) +call_cpacked = _tir_op.call_cpacked -call_packed_lowered = _op_wrapper(_tir_op.call_packed_lowered) +call_packed_lowered = _tir_op.call_packed_lowered -call_cpacked_lowered = _op_wrapper(_tir_op.call_cpacked_lowered) +call_cpacked_lowered = _tir_op.call_cpacked_lowered -handle_add_byte_offset = _op_wrapper(_tir_op.handle_add_byte_offset) +handle_add_byte_offset = _tir_op.handle_add_byte_offset -tvm_struct_set = _op_wrapper(_tir_op.tvm_struct_set) +tvm_struct_set = _tir_op.tvm_struct_set tvm_struct_get = _tir_op.tvm_struct_get -tvm_thread_invariant = _op_wrapper(_tir_op.tvm_thread_invariant) +tvm_thread_invariant = _tir_op.tvm_thread_invariant -tvm_thread_allreduce = _op_wrapper(_tir_op.tvm_thread_allreduce) +tvm_thread_allreduce = _tir_op.tvm_thread_allreduce -tvm_load_matrix_sync = _op_wrapper(_tir_op.tvm_load_matrix_sync) +tvm_load_matrix_sync = _tir_op.tvm_load_matrix_sync -tvm_mma_sync = _op_wrapper(_tir_op.tvm_mma_sync) +tvm_mma_sync = _tir_op.tvm_mma_sync -tvm_bmma_sync = _op_wrapper(_tir_op.tvm_bmma_sync) +tvm_bmma_sync = _tir_op.tvm_bmma_sync -tvm_fill_fragment = _op_wrapper(_tir_op.tvm_fill_fragment) +tvm_fill_fragment = _tir_op.tvm_fill_fragment -tvm_store_matrix_sync = _op_wrapper(_tir_op.tvm_store_matrix_sync) +tvm_store_matrix_sync = _tir_op.tvm_store_matrix_sync tvm_storage_sync = _tir_op.tvm_storage_sync -tvm_kernel_replace_point = _op_wrapper(_tir_op.tvm_kernel_replace_point) +tvm_kernel_replace_point = _tir_op.tvm_kernel_replace_point tvm_warp_shuffle = _tir_op.tvm_warp_shuffle @@ -630,34 +620,34 @@ def _register_tir_namespace_printer_names(): tvm_warp_activemask = _tir_op.tvm_warp_activemask -cooperative_tensor_fill = _op_wrapper(_tir_op.cooperative_tensor_fill) +cooperative_tensor_fill = _tir_op.cooperative_tensor_fill -cooperative_tensor_load = _op_wrapper(_tir_op.cooperative_tensor_load) +cooperative_tensor_load = _tir_op.cooperative_tensor_load -cooperative_tensor_store = _op_wrapper(_tir_op.cooperative_tensor_store) +cooperative_tensor_store = _tir_op.cooperative_tensor_store -cooperative_tensor_multiply_accumulate = _op_wrapper(_tir_op.cooperative_tensor_multiply_accumulate) +cooperative_tensor_multiply_accumulate = _tir_op.cooperative_tensor_multiply_accumulate -assume = _op_wrapper(_tir_op.assume) +assume = _tir_op.assume -undef = _op_wrapper(_tir_op.undef) +undef = _tir_op.undef -TVMBackendAllocWorkspace = _op_wrapper(_tir_op.TVMBackendAllocWorkspace) +TVMBackendAllocWorkspace = _tir_op.TVMBackendAllocWorkspace -TVMBackendFreeWorkspace = _op_wrapper(_tir_op.TVMBackendFreeWorkspace) +TVMBackendFreeWorkspace = _tir_op.TVMBackendFreeWorkspace -vscale = _op_wrapper(_tir_op.vscale) +vscale = _tir_op.vscale -ignore_loop_partition = _op_wrapper(_tir_op.ignore_loop_partition) +ignore_loop_partition = _tir_op.ignore_loop_partition reinterpret = _dtype_forward(_tir_op.reinterpret) @@ -693,10 +683,10 @@ def _register_tir_namespace_printer_names(): masked_load = _dtype_forward(_tir_op.masked_load) -masked_store = _op_wrapper(_tir_op.masked_store) +masked_store = _tir_op.masked_store -dp4a = _dtype_forward(_tir_op.dp4a) +dp4a = _tir_op.dp4a broadcast = Broadcast diff --git a/python/tvm/topi/gpu/sort.py b/python/tvm/topi/gpu/sort.py index 65dbf4da729f..5b44906f2acd 100644 --- a/python/tvm/topi/gpu/sort.py +++ b/python/tvm/topi/gpu/sort.py @@ -154,9 +154,7 @@ def _odd_even_sort( [tid + n], ) - T.evaluate( - tvm.ir.Call("tirx.tvm_storage_sync", [tvm.ir.StringImm("shared")], ret_ty="void") - ) + T.evaluate(tvm.ir.Call("tirx.tvm_storage_sync", [tvm.ir.StringImm("shared")], ty="void")) idxm = tvm.tirx.indexmod # OddEvenTransposeSort @@ -185,7 +183,7 @@ def _odd_even_sort( ) T.buffer_store(tmp_values_swap, temp_values[0], [tid + n + 1]) T.evaluate( - tvm.ir.Call("tirx.tvm_storage_sync", [tvm.ir.StringImm("shared")], ret_ty="void") + tvm.ir.Call("tirx.tvm_storage_sync", [tvm.ir.StringImm("shared")], ty="void") ) ## Copy sorted data to output diff --git a/src/relax/script/printer/call.cc b/src/relax/script/printer/call.cc index 4a7b514dccb8..d6ffdd3dde98 100644 --- a/src/relax/script/printer/call.cc +++ b/src/relax/script/printer/call.cc @@ -172,7 +172,7 @@ ffi::Optional CallTIRDocTranslate(DocTranslatorObj* d, ffi::AnyView inp ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); if (!HasRelaxCallResult(call, destination) || call->args.size() != 2 || !call->args[1].as() || call->ty_args.size() != 1) { - return RawCall(d, call, false); + return RawCall(d, call); } const Op& op = call->op.as_or_throw(); const auto* inplace = call->attrs.as(); @@ -180,27 +180,27 @@ ffi::Optional CallTIRDocTranslate(DocTranslatorObj* d, ffi::AnyView inp if (op->name == "relax.call_tir_inplace" ? inplace == nullptr : op->name == "relax.call_tir_with_grad" ? grad == nullptr : call->attrs.defined()) { - return RawCall(d, call, false); + return RawCall(d, call); } const Type& output = call->ty_args[0]; const auto* tuple = output.as(); // The constructor unwraps a single-element output list into a tensor type. - if (tuple && tuple->fields.size() == 1) return RawCall(d, call, false); + if (tuple && tuple->fields.size() == 1) return RawCall(d, call); ffi::Array output_types = tuple ? tuple->fields : ffi::Array{output}; bool distributed = !output_types.empty() && output_types[0].as(); - if (distributed && op->name != "relax.call_tir") return RawCall(d, call, false); + if (distributed && op->name != "relax.call_tir") return RawCall(d, call); ffi::Array output_docs; for (const Type& type : output_types) { if (distributed) { const auto* tensor = type.as(); if (!tensor || !tensor->tensor_ty->shape.as()) { - return RawCall(d, call, false); + return RawCall(d, call); } } else { const auto* tensor = type.as(); - if (!tensor || !tensor->shape.as()) return RawCall(d, call, false); + if (!tensor || !tensor->shape.as()) return RawCall(d, call); } output_docs.push_back(d->Translate(type).value()); } @@ -251,9 +251,9 @@ ffi::Optional CallDefaultDocTranslate(DocTranslatorObj* d, ffi::AnyView ffi::Array args; for (const Expr& arg : call->args) args.push_back(d->Translate(arg).value()); if (auto doc = RelaxCallDocTranslate(d, call, args)) return doc; - return RawCall(d, call, false, args); + return RawCall(d, call, args); } - return RawCall(d, call, false); + return RawCall(d, call); } if ((call->op.as() || call->op.as() || call->op.as()) && @@ -275,7 +275,7 @@ ffi::Optional CallDefaultDocTranslate(DocTranslatorObj* d, ffi::AnyView } } if (!inferred || !ffi::StructuralEqual()(inferred.value(), call->ty)) { - return RawCall(d, call, false); + return RawCall(d, call); } const Type& result_type = inferred.value(); if (auto doc = TIRCallPrefixDocTranslate(d, call)) return doc; @@ -283,7 +283,7 @@ ffi::Optional CallDefaultDocTranslate(DocTranslatorObj* d, ffi::AnyView for (const Expr& arg : call->args) args.push_back(d->Translate(arg).value()); if (auto doc = FFIKernelDocTranslate(d, call, result_type, args)) return doc; if (auto doc = TIRCallDocTranslate(d, call, result_type, args)) return doc; - return RawCall(d, call, true, args); + return RawCall(d, call, args); } ffi::Optional CallDocTranslate(DocTranslatorObj* d, ffi::AnyView input, diff --git a/src/s_tir/script/printer/expr.cc b/src/s_tir/script/printer/expr.cc index 0d78e365f084..e624638e8a5d 100644 --- a/src/s_tir/script/printer/expr.cc +++ b/src/s_tir/script/printer/expr.cc @@ -33,7 +33,7 @@ ffi::Optional CpAsyncRawDocTranslate(DocTranslatorObj* d, ffi::AnyView // This constructor takes an element dtype followed by the five stored operands. // The dtype is carried by the Call itself, not derived from a pointer argument. if (!CanTranslateExplicitResultCall(call) || call->args.size() != 5) { - return RawCall(d, call, false); + return RawCall(d, call); } ffi::Array args = {TypeValue(d, call->ty)}; for (const Expr& arg : call->args) { diff --git a/src/script/printer/ir/utils.cc b/src/script/printer/ir/utils.cc index 50b8ed2fb9ec..2840ddcb3b52 100644 --- a/src/script/printer/ir/utils.cc +++ b/src/script/printer/ir/utils.cc @@ -161,7 +161,7 @@ ExprDoc CallAttrsValue(DocTranslatorObj* d, const Attrs& attrs) { } // namespace // The explicit fallback retains every field, including typed attribute objects. -ExprDoc RawCall(DocTranslatorObj* d, const CallNode* call, bool infer_result, +ExprDoc RawCall(DocTranslatorObj* d, const CallNode* call, ffi::Optional> translated_args) { ffi::Optional op = call->op.as(); ffi::Array args; @@ -188,11 +188,12 @@ ExprDoc RawCall(DocTranslatorObj* d, const CallNode* call, bool infer_result, keys.push_back("ty_args"); values.push_back(ListDoc(types)); } - if (!infer_result) { - keys.push_back("ty"); - values.push_back(TypeValue(d, call->ty)); - } - return NamespaceDoc("ir")->Attr("Call")->Call({callee, ListDoc(args)}, keys, values); + keys.push_back("ty"); + values.push_back(TypeValue(d, call->ty)); + return NamespaceDoc("ir") + ->Attr("Call") + ->Attr("unchecked") + ->Call({callee, ListDoc(args)}, keys, values); } // The query aliases the active frame; candidate classification must own a copy. diff --git a/src/script/printer/ir/utils.h b/src/script/printer/ir/utils.h index 2dd96e46ec74..d7ed37392866 100644 --- a/src/script/printer/ir/utils.h +++ b/src/script/printer/ir/utils.h @@ -40,7 +40,7 @@ ExprDoc NamedCallCallee(const ffi::String& canonical_name); ExprDoc TypeValue(DocTranslatorObj* d, const Type& type, bool dtype_literal = true); ExprDoc MaterializeCallArgument(DocTranslatorObj* d, const Expr& arg, ExprDoc doc); -ExprDoc RawCall(DocTranslatorObj* d, const CallNode* call, bool infer_result, +ExprDoc RawCall(DocTranslatorObj* d, const CallNode* call, ffi::Optional> translated_args = std::nullopt); ExprDoc AnyValue(DocTranslatorObj* d, ffi::AnyView value); ffi::Dict CopyImplicitDefs(DocTranslatorObj* d); diff --git a/src/tirx/script/printer/buffer.cc b/src/tirx/script/printer/buffer.cc index c6d69053fcd7..d5f965c8bc42 100644 --- a/src/tirx/script/printer/buffer.cc +++ b/src/tirx/script/printer/buffer.cc @@ -53,18 +53,18 @@ ffi::Optional BufferOperationDocTranslate(DocTranslatorObj* d, ffi::Any ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); // Surface buffer constructors emit a binding, so they cannot replace an // allocation expression nested inside another Call or statement. - if (!destination) return RawCall(d, call, false); + if (!destination) return RawCall(d, call); TVM_FFI_CHECK(destination->IsInstance(), TypeError) << "Buffer operation destination must be a variable"; auto var = ffi::GetRef(static_cast(destination)); - if (!ffi::StructuralEqual()(var->ty, call->ty)) return RawCall(d, call, false); + if (!ffi::StructuralEqual()(var->ty, call->ty)) return RawCall(d, call); bool is_alloc = call->op.same_as(tirx::builtin::alloc_buffer()); size_t shape_index = is_alloc ? 0 : 1; auto buffer = call->ty.as(); if (!buffer || call->args.size() != shape_index + 3 || !call->ty_args.empty() || call->attrs.defined() != is_alloc || (call->attrs.defined() && !call->attrs.as())) { - return RawCall(d, call, false); + return RawCall(d, call); } auto shape = call->args[shape_index].as(); auto dtype = call->args[shape_index + 1].as(); @@ -73,7 +73,7 @@ ffi::Optional BufferOperationDocTranslate(DocTranslatorObj* d, ffi::Any dtype->value != buffer.value()->dtype->dtype || scope->value != buffer.value()->storage_scope || scope->value.empty() || buffer.value()->data_alignment <= 0 || buffer.value()->offset_factor == 0) { - return RawCall(d, call, false); + return RawCall(d, call); } ffi::Map annotations; if (is_alloc) annotations = call->attrs.as_or_throw()->dict; @@ -81,7 +81,7 @@ ffi::Optional BufferOperationDocTranslate(DocTranslatorObj* d, ffi::Any if (value.type_index() == ffi::TypeIndex::kTVMFFIInt || value.type_index() == ffi::TypeIndex::kTVMFFIBool || value.type_index() == ffi::TypeIndex::kTVMFFIFloat) { - return RawCall(d, call, false); + return RawCall(d, call); } } ffi::Optional data = is_alloc ? std::nullopt : ffi::Optional(call->args[0]); @@ -89,7 +89,7 @@ ffi::Optional BufferOperationDocTranslate(DocTranslatorObj* d, ffi::Any (!ffi::StructuralEqual()(buffer.value()->elem_offset, IntImm(PrimType(buffer.value()->DefaultIndexType()), 0)) || buffer.value()->offset_factor != 1 || (!is_alloc && scope->value != "tmem"))) { - return RawCall(d, call, false); + return RawCall(d, call); } if (!is_alloc && scope->value == "tmem") { const auto* pointer = data.value().as(); @@ -98,15 +98,15 @@ ffi::Optional BufferOperationDocTranslate(DocTranslatorObj* d, ffi::Any pointer->attrs.defined() || !pointer->ty_args.empty() || !ffi::StructuralEqual()(pointer->ty, buffer.value()->DataPointerType()) || !ffi::StructuralEqual()(pointer->args[0], buffer.value()->allocated_addr[0])) { - return RawCall(d, call, false); + return RawCall(d, call); } data = std::nullopt; } else if (!is_alloc) { const auto* pointer = data.value()->ty.as(); - if (!pointer || pointer->storage_scope != scope->value) return RawCall(d, call, false); + if (!pointer || pointer->storage_scope != scope->value) return RawCall(d, call); } CallDoc rhs = d->Translate(buffer.value()).value().as_or_throw(); - if (rhs->callee.as_or_throw()->name != "Buffer") return RawCall(d, call, false); + if (rhs->callee.as_or_throw()->name != "Buffer") return RawCall(d, call); ffi::String method = is_alloc ? "alloc_buffer" : "decl_buffer"; if (is_alloc && (scope->value == "local" || scope->value == "shared")) { method = scope->value == "local" ? "alloc_local" : "alloc_shared"; diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc index 4b106b672b2c..959fe8d1a21b 100644 --- a/src/tirx/script/printer/expr.cc +++ b/src/tirx/script/printer/expr.cc @@ -188,7 +188,7 @@ ffi::Optional StorageSyncDocTranslate(DocTranslatorObj* d, ffi::AnyView const auto* call = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); if (!CanTranslateExplicitResultCall(call) || call->args.empty() || call->args.size() > 3) { - return RawCall(d, call, false); + return RawCall(d, call); } ffi::Array args; for (const Expr& arg : call->args) { @@ -218,10 +218,10 @@ ffi::Optional CallExternDocTranslate(DocTranslatorObj* d, ffi::AnyView ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); if (!call->op.same_as(tirx::builtin::call_extern()) || !CanTranslateExplicitResultCall(call) || call->args.empty()) { - return RawCall(d, call, false); + return RawCall(d, call); } const auto* name = call->args[0].as(); - if (!name) return RawCall(d, call, false); + if (!name) return RawCall(d, call); ExprDoc name_doc = LiteralDoc::Str(name->value, std::nullopt); d->RecordOrigin(name_doc, call->args[0]); ffi::Array args = {TypeValue(d, call->ty), name_doc}; @@ -244,11 +244,11 @@ ffi::Optional CUDAFuncCallDocTranslate(DocTranslatorObj* d, ffi::AnyVie static const Op cuda_func_call = Op::Get("tirx.cuda.func_call"); if (!call->op.same_as(cuda_func_call) || !CanTranslateExplicitResultCall(call) || call->args.size() < 2) { - return RawCall(d, call, false); + return RawCall(d, call); } const auto* name = call->args[0].as(); const auto* source = call->args.back().as(); - if (!name || !source) return RawCall(d, call, false); + if (!name || !source) return RawCall(d, call); ExprDoc name_doc = LiteralDoc::Str(name->value, std::nullopt); ExprDoc source_doc = LiteralDoc::Str(source->value, std::nullopt); d->RecordOrigin(name_doc, call->args[0]); @@ -280,10 +280,10 @@ ffi::Optional CUDAInstructionDescriptorDocTranslate(DocTranslatorObj* d ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); if (!CanTranslateExplicitResultCall(call) || call->args.size() != (block_scaled ? 17 : 14) || !ffi::StructuralEqual()(call->ty, PrimType::Void())) { - return RawCall(d, call, false); + return RawCall(d, call); } Op op = call->op.as_or_throw(); - if (op->args_info.size() != call->args.size()) return RawCall(d, call, false); + if (op->args_info.size() != call->args.size()) return RawCall(d, call); constexpr size_t optional_begin = block_scaled ? 13 : 9; ffi::Array keys; ffi::Array values; @@ -330,23 +330,23 @@ ffi::Optional LLVMIntrinsicDocTranslate(DocTranslatorObj* d, ffi::AnyVi ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); if (call->attrs.defined() || !call->ty_args.empty() || call->args.empty() || !call->ty.as()) { - return RawCall(d, call, false); + return RawCall(d, call); } const auto* id = call->args[0].as(); // The named constructor uses an int32 intrinsic identifier. Other stored // representations retain their exact operand type through the full Call. - if (!id || !ffi::StructuralEqual()(id->ty, PrimType::Int(32))) return RawCall(d, call, false); + if (!id || !ffi::StructuralEqual()(id->ty, PrimType::Int(32))) return RawCall(d, call); auto lookup = ffi::Function::GetGlobal("target.llvm_get_intrinsic_name"); auto reverse = ffi::Function::GetGlobal("target.llvm_lookup_intrinsic_id"); - if (!lookup || !reverse) return RawCall(d, call, false); + if (!lookup || !reverse) return RawCall(d, call); ffi::String name; try { name = (*lookup)(static_cast(id->value)).cast(); if (name.empty() || (*reverse)(name).cast() != static_cast(id->value)) { - return RawCall(d, call, false); + return RawCall(d, call); } } catch (const ffi::Error&) { - return RawCall(d, call, false); + return RawCall(d, call); } ExprDoc name_doc = LiteralDoc::Str(name, std::nullopt); d->RecordOrigin(name_doc, call->args[0]); @@ -373,21 +373,21 @@ ffi::Optional GetActiveLaneMaskDocTranslate(DocTranslatorObj* d, ffi::A ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); if (!call->op.same_as(tirx::builtin::get_active_lane_mask()) || call->attrs.defined() || !call->ty_args.empty() || call->args.size() != 2) { - return RawCall(d, call, false); + return RawCall(d, call); } auto result = call->ty.as(); if (!result || !(result.value().MatchesCode(DLDataTypeCode::kDLBool) || result.value().MatchesElementType(DLDataTypeCode::kDLUInt, 1)) || !(result.value().IsScalableVector() || result.value().IsFixedLengthVector())) { - return RawCall(d, call, false); + return RawCall(d, call); } ffi::Array args = {TypeValue(d, call->ty)}; for (const Expr& arg : call->args) { auto type = arg->ty.as(); if (!type || !type.value().IsScalar() || !type.value().MatchesCode(DLDataTypeCode::kDLInt, DLDataTypeCode::kDLUInt)) { - return RawCall(d, call, false); + return RawCall(d, call); } args.push_back(MaterializeCallArgument(d, arg, d->Translate(arg).value())); } @@ -452,7 +452,7 @@ ffi::Optional TIRCallPrefixDocTranslate(DocTranslatorObj* d, const Call // A standalone address Call must remain an expression after reparsing. if (IsPTXAddressCall(call) || ((is_ptx || descriptor) && !ffi::StructuralEqual()(call->ty, PrimType::Void()))) { - return RawCall(d, call, false); + return RawCall(d, call); } } if (auto op = call->op.as(); op && op.value()->name == "tirx.isnan") { @@ -463,7 +463,7 @@ ffi::Optional TIRCallPrefixDocTranslate(DocTranslatorObj* d, const Call !input_type.value().MatchesCode(DLDataTypeCode::kDLFloat) || (input_type.value().bits() != 32 && input_type.value().bits() != 64) || call->args[0].as()) { - return RawCall(d, call, false); + return RawCall(d, call); } } if (call->op.same_as(tirx::builtin::buffer_data()) && !call->attrs.defined() && @@ -479,7 +479,7 @@ ffi::Optional FFIKernelDocTranslate(DocTranslatorObj* d, const CallNode const ffi::Array& args) { if (call->op.same_as(tirx::builtin::call_ffi_kernel())) { const auto* attrs = call->attrs.as(); - if (!attrs || !call->ty_args.empty()) return RawCall(d, call, true, args); + if (!attrs || !call->ty_args.empty()) return RawCall(d, call, args); ffi::Array launch_params; for (const ffi::String& param : attrs->launch_params) { launch_params.push_back(LiteralDoc::Str(param, std::nullopt)); @@ -536,7 +536,7 @@ ffi::Optional TIRCallDocTranslate(DocTranslatorObj* d, const CallNode* is_ptx && address && IsPTXAddressCall(address)) { // Use the address constructor only when it preserves its operands. auto translated = ConsumedPTXAddressDocTranslate(d, address); - if (!translated) return RawCall(d, call, false); + if (!translated) return RawCall(d, call); argument = translated.value(); } bool reads_operand_type = diff --git a/tests/python/codegen/test_target_codegen_llvm.py b/tests/python/codegen/test_target_codegen_llvm.py index ec0b9aa14603..6a5c0dad4a47 100644 --- a/tests/python/codegen/test_target_codegen_llvm.py +++ b/tests/python/codegen/test_target_codegen_llvm.py @@ -78,7 +78,7 @@ class Module: @T.prim_func def main(A: T.handle("float32")): A_buf = T.decl_buffer((4,), "float32", data=A) - T.evaluate(T.Call("tirx.prefetch", [T.address_of(A_buf[0]), 0, 3, 1], ret_ty="void")) + T.evaluate(T.Call("tirx.prefetch", [T.address_of(A_buf[0]), 0, 3, 1], ty="void")) fcode = tvm.compile(Module) @@ -1115,7 +1115,7 @@ def main(): T.Call( tvm.ir.Op.get("tirx.tvm_call_packed"), ["dummy_function_name"], - ret_ty="void", + ty="void", ) # Error occurred during build, as part of @@ -1137,7 +1137,7 @@ def test_call_packed_without_string_arg(): class Module: @T.prim_func def main(A: T.Buffer(1, "float32")): - T.Call(tvm.ir.Op.get("tirx.tvm_call_packed"), [A.data], ret_ty="int32") + T.Call(tvm.ir.Op.get("tirx.tvm_call_packed"), [A.data], ty="int32") with pytest.raises(RuntimeError): built = tvm.compile(Module, target="llvm") @@ -1151,7 +1151,7 @@ def test_call_extern_returning_void(): class Module: @T.prim_func def main(): - T.Call(tvm.ir.Op.get("tirx.call_extern"), ["dummy_function_name"], ret_ty="void") + T.Call(tvm.ir.Op.get("tirx.call_extern"), ["dummy_function_name"], ty="void") built = tvm.compile(Module, target="llvm") diff --git a/tests/python/relax/test_analysis_well_formed.py b/tests/python/relax/test_analysis_well_formed.py index a6fad5df7dd6..4e9e8d829ef7 100644 --- a/tests/python/relax/test_analysis_well_formed.py +++ b/tests/python/relax/test_analysis_well_formed.py @@ -176,7 +176,7 @@ def test_unchecked_call_constructor(): assert call.span.same_as(span) assert isinstance(call.attrs, tvm.ir.DictAttrs) assert len(call.ty_args) == 0 - assert isinstance(tvm.ir.Call.unchecked(op, [x], ret_ty="handle").ty, tvm.ir.PointerType) + assert isinstance(tvm.ir.Call.unchecked(op, [x], ty="handle").ty, tvm.ir.PointerType) with pytest.raises(TypeError, match="skip_validate"): tvm.ir.Call(op, [x], skip_validate=True) @@ -669,7 +669,7 @@ def test_impure_in_dataflow_block(): # The throwing form surfaces the offending impure call in its message. with pytest.raises(Exception) as excinfo: rx.analysis.well_formed(mod) - assert 'I.Call("relax.print", ["{}", x], ty=R.Tuple())' in str(excinfo.value) + assert 'I.Call.unchecked("relax.print", ["{}", x], ty=R.Tuple())' in str(excinfo.value) def test_well_formed_function(): diff --git a/tests/python/relax/test_expr.py b/tests/python/relax/test_expr.py index 1830abcbf426..d3bda196fce7 100644 --- a/tests/python/relax/test_expr.py +++ b/tests/python/relax/test_expr.py @@ -390,14 +390,14 @@ def test_call_reinfer_type_from_current_inputs(): assert tvm.ir.reinfer_type(initial).same_as(original_ty) _check_type_missing(initial.ty) - unchanged = rx.Call("relax.abs", [x], ret_ty=original_ty) + unchanged = rx.Call("relax.abs", [x], ty=original_ty) assert tvm.ir.reinfer_type(unchanged).same_as(unchanged.ty) equivalent_ty = rx.TensorType(original_ty.shape, original_ty.dtype) - equivalent = rx.Call("relax.abs", [x], ret_ty=equivalent_ty) + equivalent = rx.Call("relax.abs", [x], ty=equivalent_ty) assert tvm.ir.reinfer_type(equivalent).same_as(equivalent_ty) - stale = rx.Call("relax.abs", [y], ret_ty=original_ty) + stale = rx.Call("relax.abs", [y], ty=original_ty) _check_equal(tvm.ir.reinfer_type(stale), changed_ty) assert tvm.ir.reinfer_type(stale).same_as(changed_ty) _check_equal(stale.ty, original_ty) @@ -405,7 +405,7 @@ def test_call_reinfer_type_from_current_inputs(): with pytest.raises(tvm.error.InternalError, match="type is not populated"): tvm.ir.reinfer_type(rx.Call("relax.abs", [rx.Var("untyped")])) with pytest.raises(tvm.error.InternalError, match="type is not populated"): - tvm.ir.reinfer_type(rx.Call("relax.abs", [rx.Var("untyped")], ret_ty=original_ty)) + tvm.ir.reinfer_type(rx.Call("relax.abs", [rx.Var("untyped")], ty=original_ty)) with pytest.raises(ValueError, match="No context-free"): tvm.ir.reinfer_type(rx.Call.unchecked("relax.matmul", [x, x])) func = rx.Var("func", rx.FuncType([original_ty], original_ty)) diff --git a/tests/python/relax/test_transform_annotate_tir_op_pattern.py b/tests/python/relax/test_transform_annotate_tir_op_pattern.py index c279eb1f8261..82b1aa6c6a08 100644 --- a/tests/python/relax/test_transform_annotate_tir_op_pattern.py +++ b/tests/python/relax/test_transform_annotate_tir_op_pattern.py @@ -283,7 +283,6 @@ def max_pool2d( 1 <= ax2 and ax2 < 113 and 1 <= ax3 and ax3 < 113, rxplaceholder_1[ax0, ax1, ax2 - 1, ax3 - 1], T.float32(-3.4028234663852886e38), - dtype="float32", ) for i0, i1, i2, i3, i4, i5 in T.grid(1, 64, 56, 56, 3, 3): with Ts.sblock("tensor"): @@ -336,7 +335,7 @@ def softmax( Ts.reads(rxplaceholder_1[i0_10, i1_5], T_softmax_maxelem_1[i0_10]) Ts.writes(T_softmax_exp_1[i0_10, i1_5]) T_softmax_exp_1[i0_10, i1_5] = T.exp( - rxplaceholder_1[i0_10, i1_5] - T_softmax_maxelem_1[i0_10], dtype="float32" + rxplaceholder_1[i0_10, i1_5] - T_softmax_maxelem_1[i0_10] ) for i0_11, i1_6 in T.grid(16, 16): with Ts.sblock("T_softmax_expsum"): diff --git a/tests/python/relax/test_transform_fuse_ops.py b/tests/python/relax/test_transform_fuse_ops.py index 600a51a3c9bc..d4afb8b0ad52 100644 --- a/tests/python/relax/test_transform_fuse_ops.py +++ b/tests/python/relax/test_transform_fuse_ops.py @@ -911,7 +911,7 @@ def layer_norm(A: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(A[ax0, ax1, ax2, ax3], rxplaceholder_red_temp_v0[ax0, ax1], rxplaceholder_red_temp_v1[ax0, ax1], gamma[ax2, ax3], beta[ax2, ax3]) Ts.writes(T_layer_norm[ax0, ax1, ax2, ax3]) - T_layer_norm[ax0, ax1, ax2, ax3] = (A[ax0, ax1, ax2, ax3] - rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.05)) * T.rsqrt(rxplaceholder_red_temp_v1[ax0, ax1] * T.float32(0.05) - rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.05) * (rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.05)) + T.float32(1e-05), dtype="float32") * gamma[ax2, ax3] + beta[ax2, ax3] + T_layer_norm[ax0, ax1, ax2, ax3] = (A[ax0, ax1, ax2, ax3] - rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.05)) * T.rsqrt(rxplaceholder_red_temp_v1[ax0, ax1] * T.float32(0.05) - rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.05) * (rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.05)) + T.float32(1e-05)) * gamma[ax2, ax3] + beta[ax2, ax3] @Ts.prim_func(private=True) def relu(A: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32"), B: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32")): diff --git a/tests/python/relax/test_transform_legalize_ops_image.py b/tests/python/relax/test_transform_legalize_ops_image.py index 25da182ac3b4..142f2dfbdead 100644 --- a/tests/python/relax/test_transform_legalize_ops_image.py +++ b/tests/python/relax/test_transform_legalize_ops_image.py @@ -101,7 +101,7 @@ def resize2d(rxplaceholder: T.Buffer([n_resize2d, c_resize2d, h_resize2d, w_resi i0_1, i1_1, i2_1, i3_1, i4_1 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(rxplaceholder[i0_1, i1_1, T.int64(0) : T.max(h_resize2d, T.int64(1)), T.int64(0) : T.max(w_resize2d, T.int64(1)), i4_1]) Ts.writes(resize[i0_1, i1_1, i2_1, i3_1, i4_1]) - resize[i0_1, i1_1, i2_1, i3_1, i4_1] = rxplaceholder[i0_1, i1_1, T.max(T.min(T.Cast("int64", T.round(T.Cast("float32", h_resize2d) / T.Cast("float32", oh_resize2d) * T.Cast("float32", i2_1), dtype="float32")), h_resize2d - T.int64(1)), T.int64(0)), T.max(T.min(T.Cast("int64", T.round(T.Cast("float32", w_resize2d) / T.Cast("float32", ow_resize2d) * T.Cast("float32", i3_1), dtype="float32")), w_resize2d - T.int64(1)), T.int64(0)), i4_1] + resize[i0_1, i1_1, i2_1, i3_1, i4_1] = rxplaceholder[i0_1, i1_1, T.max(T.min(T.Cast("int64", T.round(T.Cast("float32", h_resize2d) / T.Cast("float32", oh_resize2d) * T.Cast("float32", i2_1))), h_resize2d - T.int64(1)), T.int64(0)), T.max(T.min(T.Cast("int64", T.round(T.Cast("float32", w_resize2d) / T.Cast("float32", ow_resize2d) * T.Cast("float32", i3_1))), w_resize2d - T.int64(1)), T.int64(0)), i4_1] # fmt: on mod = LegalizeOps()(Resize2D) diff --git a/tests/python/relax/test_transform_legalize_ops_nn.py b/tests/python/relax/test_transform_legalize_ops_nn.py index a957b9cfe127..d00ae6912adc 100644 --- a/tests/python/relax/test_transform_legalize_ops_nn.py +++ b/tests/python/relax/test_transform_legalize_ops_nn.py @@ -282,7 +282,7 @@ def conv2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(128), T.int64(28), T.int i0_1, i1_1, i2_1, i3_1 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[i0_1, i1_1, i2_1 - T.int64(1), i3_1 - T.int64(1)]) Ts.writes(pad_temp[i0_1, i1_1, i2_1, i3_1]) - pad_temp[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(T.int64(1) <= i2_1 and i2_1 < T.int64(29) and T.int64(1) <= i3_1 and i3_1 < T.int64(29), rxplaceholder[i0_1, i1_1, i2_1 - T.int64(1), i3_1 - T.int64(1)], T.float32(0), dtype="float32") + pad_temp[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(T.int64(1) <= i2_1 and i2_1 < T.int64(29) and T.int64(1) <= i3_1 and i3_1 < T.int64(29), rxplaceholder[i0_1, i1_1, i2_1 - T.int64(1), i3_1 - T.int64(1)], T.float32(0)) for i0, i1, i2, i3, i4, i5, i6 in T.grid(T.int64(2), T.int64(64), T.int64(13), T.int64(13), T.int64(16), T.int64(3), T.int64(3)): with Ts.sblock("group_conv2d_nchw"): nn, ff, yy, xx, rc, ry, rx = Ts.axis.remap("SSSSRRR", [i0, i1, i2, i3, i4, i5, i6]) @@ -870,7 +870,7 @@ def max_pool2d(rxplaceholder: T.Buffer((T.int64(4), T.int64(112), T.int64(112), ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[ax0, ax1 - T.int64(1), ax2 - T.int64(1), ax3]) Ts.writes(pad_temp[ax0, ax1, ax2, ax3]) - pad_temp[ax0, ax1, ax2, ax3] = T.if_then_else(T.int64(1) <= ax1 and ax1 < T.int64(113) and T.int64(1) <= ax2 and ax2 < T.int64(113), rxplaceholder[ax0, ax1 - T.int64(1), ax2 - T.int64(1), ax3], T.float32(-3.4028234663852886e+38), dtype="float32") + pad_temp[ax0, ax1, ax2, ax3] = T.if_then_else(T.int64(1) <= ax1 and ax1 < T.int64(113) and T.int64(1) <= ax2 and ax2 < T.int64(113), rxplaceholder[ax0, ax1 - T.int64(1), ax2 - T.int64(1), ax3], T.float32(-3.4028234663852886e+38)) for i0, i1, i2, i3, i4, i5 in T.grid(T.int64(4), T.int64(56), T.int64(56), T.int64(6), T.int64(3), T.int64(3)): with Ts.sblock("pool_max"): ax0, ax1, ax2, ax3, rv0, rv1 = Ts.axis.remap("SSSSRR", [i0, i1, i2, i3, i4, i5]) @@ -945,7 +945,7 @@ def max_pool2d(rxplaceholder: T.Buffer((T.int64(4), T.int64(6), T.int64(112), T. ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[ax0, ax1, ax2 - T.int64(1), ax3 - T.int64(1)]) Ts.writes(pad_temp[ax0, ax1, ax2, ax3]) - pad_temp[ax0, ax1, ax2, ax3] = T.if_then_else(T.int64(1) <= ax2 and ax2 < T.int64(113) and T.int64(1) <= ax3 and ax3 < T.int64(113), rxplaceholder[ax0, ax1, ax2 - T.int64(1), ax3 - T.int64(1)], T.float32(-3.4028234663852886e+38), dtype="float32") + pad_temp[ax0, ax1, ax2, ax3] = T.if_then_else(T.int64(1) <= ax2 and ax2 < T.int64(113) and T.int64(1) <= ax3 and ax3 < T.int64(113), rxplaceholder[ax0, ax1, ax2 - T.int64(1), ax3 - T.int64(1)], T.float32(-3.4028234663852886e+38)) for i0, i1, i2, i3, i4, i5 in T.grid(T.int64(4), T.int64(6), T.int64(38), T.int64(38), T.int64(3), T.int64(3)): with Ts.sblock("pool_max"): ax0, ax1, ax2, ax3, rv0, rv1 = Ts.axis.remap("SSSSRR", [i0, i1, i2, i3, i4, i5]) @@ -1906,7 +1906,7 @@ def softmax(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(16), T.int6 i0_2, i1_2, i2_2, i3_1 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[i0_2, i1_2, i2_2, i3_1], T_softmax_maxelem[i0_2, i1_2, i3_1]) Ts.writes(T_softmax_exp[i0_2, i1_2, i2_2, i3_1]) - T_softmax_exp[i0_2, i1_2, i2_2, i3_1] = T.exp(rxplaceholder[i0_2, i1_2, i2_2, i3_1] - T_softmax_maxelem[i0_2, i1_2, i3_1], dtype="float32") + T_softmax_exp[i0_2, i1_2, i2_2, i3_1] = T.exp(rxplaceholder[i0_2, i1_2, i2_2, i3_1] - T_softmax_maxelem[i0_2, i1_2, i3_1]) for i0_3, i1_3, i2_3, i3 in T.grid(T.int64(2), T.int64(3), T.int64(32), T.int64(16)): with Ts.sblock("T_softmax_expsum"): i0_4, i1_4, i2_4, k = Ts.axis.remap("SSSR", [i0_3, i1_3, i2_3, i3]) @@ -1975,7 +1975,7 @@ def softmax(rxplaceholder: T.Buffer([a_softmax, b_softmax, c_softmax], dtype='fl i0_2, i1_2, i2_1 = Ts.axis.remap("SSS", [i0, i1, i2]) Ts.reads(rxplaceholder[i0_2, i1_2, i2_1], T_softmax_maxelem[i0_2, i1_2]) Ts.writes(T_softmax_exp[i0_2, i1_2, i2_1]) - T_softmax_exp[i0_2, i1_2, i2_1] = T.exp(rxplaceholder[i0_2, i1_2, i2_1] - T_softmax_maxelem[i0_2, i1_2], dtype="float32") + T_softmax_exp[i0_2, i1_2, i2_1] = T.exp(rxplaceholder[i0_2, i1_2, i2_1] - T_softmax_maxelem[i0_2, i1_2]) for i0_3, i1_3, i2 in T.grid(a_softmax, b_softmax, c_softmax): with Ts.sblock("T_softmax_expsum"): i0_4, i1_4, k = Ts.axis.remap("SSR", [i0_3, i1_3, i2]) @@ -2033,14 +2033,14 @@ def log_softmax(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(16), T. Ts.writes(compute_1[i0_2, i1_2, i2_2]) with Ts.init(): compute_1[i0_2, i1_2, i2_2] = T.float32(0) - compute_1[i0_2, i1_2, i2_2] = compute_1[i0_2, i1_2, i2_2] + T.exp(rxplaceholder[i0_2, i1_2, k, i2_2] - T_softmax_maxelem[i0_2, i1_2, i2_2], dtype="float32") + compute_1[i0_2, i1_2, i2_2] = compute_1[i0_2, i1_2, i2_2] + T.exp(rxplaceholder[i0_2, i1_2, k, i2_2] - T_softmax_maxelem[i0_2, i1_2, i2_2]) for i0_3, i1_3, i2_3, i3 in T.grid(T.int64(2), T.int64(3), T.int64(16), T.int64(32)): with Ts.sblock("compute_1"): i0_4, i1_4, i2_4, i3_1 = Ts.axis.remap("SSSS", [i0_3, i1_3, i2_3, i3]) Ts.reads(rxplaceholder[i0_4, i1_4, i2_4, i3_1], T_softmax_maxelem[i0_4, i1_4, i3_1], compute_1[i0_4, i1_4, i3_1]) Ts.writes(compute[i0_4, i1_4, i2_4, i3_1]) Ts.sblock_attr({"axis": 2}) - compute[i0_4, i1_4, i2_4, i3_1] = (rxplaceholder[i0_4, i1_4, i2_4, i3_1] - T_softmax_maxelem[i0_4, i1_4, i3_1] - T.log(compute_1[i0_4, i1_4, i3_1], dtype="float32")) + compute[i0_4, i1_4, i2_4, i3_1] = (rxplaceholder[i0_4, i1_4, i2_4, i3_1] - T_softmax_maxelem[i0_4, i1_4, i3_1] - T.log(compute_1[i0_4, i1_4, i3_1])) # fmt: on mod = LegalizeOps()(LogSoftmax) @@ -2096,14 +2096,14 @@ def log_softmax(rxplaceholder: T.Buffer([a_log_softmax, b_log_softmax, c_log_sof Ts.writes(compute_1[v_i0, v_i1]) with Ts.init(): compute_1[v_i0, v_i1] = T.float32(0) - compute_1[v_i0, v_i1] = compute_1[v_i0, v_i1] + T.exp(rxplaceholder[v_i0, v_i1, v_k] - T_softmax_maxelem[v_i0, v_i1], dtype="float32") + compute_1[v_i0, v_i1] = compute_1[v_i0, v_i1] + T.exp(rxplaceholder[v_i0, v_i1, v_k] - T_softmax_maxelem[v_i0, v_i1]) for i0, i1, i2 in T.grid(a_log_softmax, b_log_softmax, c_log_softmax): with Ts.sblock("compute_1"): v_i0, v_i1, v_i2 = Ts.axis.remap("SSS", [i0, i1, i2]) Ts.reads(rxplaceholder[v_i0, v_i1, v_i2], T_softmax_maxelem[v_i0, v_i1], compute_1[v_i0, v_i1],) Ts.writes(compute[v_i0, v_i1, v_i2]) Ts.sblock_attr({"axis": 2}) - compute[v_i0, v_i1, v_i2] = (rxplaceholder[v_i0, v_i1, v_i2] - T_softmax_maxelem[v_i0, v_i1] - T.log(compute_1[v_i0, v_i1], dtype="float32")) + compute[v_i0, v_i1, v_i2] = (rxplaceholder[v_i0, v_i1, v_i2] - T_softmax_maxelem[v_i0, v_i1] - T.log(compute_1[v_i0, v_i1])) # fmt: on mod = LegalizeOps()(LogSoftmax) diff --git a/tests/python/relax/test_transform_rewrite_cuda_graph.py b/tests/python/relax/test_transform_rewrite_cuda_graph.py index cf953f6f28e9..1e160e4040a4 100644 --- a/tests/python/relax/test_transform_rewrite_cuda_graph.py +++ b/tests/python/relax/test_transform_rewrite_cuda_graph.py @@ -48,7 +48,7 @@ def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T with Ts.sblock("compute"): i0 = Ts.axis.spatial(T.int64(2), (i0_i1_fused_0 * T.int64(8) + i0_i1_fused_1) // T.int64(4)) i1 = Ts.axis.spatial(T.int64(4), (i0_i1_fused_0 * T.int64(8) + i0_i1_fused_1) % T.int64(4)) - compute[i0, i1] = T.exp(rxplaceholder[i0, i1], dtype="float32") + compute[i0, i1] = T.exp(rxplaceholder[i0, i1]) @R.function def main(x: R.Tensor((2, 4), dtype="float32")) -> R.Tensor((2,4), dtype="float32"): @@ -92,7 +92,7 @@ def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T i1 = Ts.axis.spatial(T.int64(4), (i0_i1_fused_0 * T.int64(8) + i0_i1_fused_1) % T.int64(4)) Ts.reads(rxplaceholder[i0, i1]) Ts.writes(compute[i0, i1]) - compute[i0, i1] = T.exp(rxplaceholder[i0, i1], dtype="float32") + compute[i0, i1] = T.exp(rxplaceholder[i0, i1]) @R.function(private=True) def cuda_graph_alloc() -> R.Tuple(R.Any, R.Any, R.Any): @@ -162,7 +162,7 @@ def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T i1 = Ts.axis.spatial(T.int64(4), (i0_i1_fused_0 * T.int64(8) + i0_i1_fused_1) % T.int64(4)) Ts.reads(rxplaceholder[i0, i1]) Ts.writes(compute[i0, i1]) - compute[i0, i1] = T.exp(rxplaceholder[i0, i1], dtype="float32") + compute[i0, i1] = T.exp(rxplaceholder[i0, i1]) @R.function def main(x: R.Tensor((2, 4), dtype="float32")) -> R.Tensor((2, 4), dtype="float32"): @@ -265,7 +265,7 @@ def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T with Ts.sblock("compute"): i0 = Ts.axis.spatial(T.int64(2), (i0_i1_fused_0 * T.int64(8) + i0_i1_fused_1) // T.int64(4)) i1 = Ts.axis.spatial(T.int64(4), (i0_i1_fused_0 * T.int64(8) + i0_i1_fused_1) % T.int64(4)) - compute[i0, i1] = T.exp(rxplaceholder[i0, i1], dtype="float32") + compute[i0, i1] = T.exp(rxplaceholder[i0, i1]) @R.function def main(x: R.Tensor((2, 4), dtype="float32")) -> R.Tensor((2,4), dtype="float32"): diff --git a/tests/python/s_tir/analysis/test_sblock_access_region.py b/tests/python/s_tir/analysis/test_sblock_access_region.py index 4e590a0f7949..ea27472ebdf2 100644 --- a/tests/python/s_tir/analysis/test_sblock_access_region.py +++ b/tests/python/s_tir/analysis/test_sblock_access_region.py @@ -153,7 +153,7 @@ def access_in_if_then_else_func() -> None: Ts.reads([A[0:5]]) Ts.writes([B[0:8]]) for i in T.serial(0, 8): - B[i] = T.if_then_else(i < 5, A[i], 0.0, dtype="float32") + B[i] = T.if_then_else(i < 5, A[i], 0.0) @Ts.prim_func @@ -221,7 +221,7 @@ def access_of_padding_pattern() -> None: Ts.reads([X[vi - 2, vj - 2]]) Ts.writes([X_pad[vi, vj]]) X_pad[vi, vj] = T.if_then_else( - 2 <= vi and vi < 30 and 2 <= vj and vj < 30, X[vi - 2, vj - 2], 0.0, dtype="float32" + 2 <= vi and vi < 30 and 2 <= vj and vj < 30, X[vi - 2, vj - 2], 0.0 ) with Ts.sblock("padding_reverse"): vi, vj = Ts.axis.remap("SS", [i, j]) diff --git a/tests/python/s_tir/analysis/test_sblock_buffer_access_lca.py b/tests/python/s_tir/analysis/test_sblock_buffer_access_lca.py index 642d744c8123..32b13f998166 100644 --- a/tests/python/s_tir/analysis/test_sblock_buffer_access_lca.py +++ b/tests/python/s_tir/analysis/test_sblock_buffer_access_lca.py @@ -59,7 +59,7 @@ def buffer_opaque_access( for j in range(0, 16): T.evaluate(A[i * 16 + j]) for j in range(0, 16): - T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, T.float32(0), dtype="handle")) + T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, T.float32(0))) for i, j in T.grid(16, 16): with Ts.sblock(): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_feature_extractor_per_store_feature.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_feature_extractor_per_store_feature.py index b9f21a84fd8e..919e11e1e7c6 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_feature_extractor_per_store_feature.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_feature_extractor_per_store_feature.py @@ -73,7 +73,7 @@ def main(placeholder: T.Buffer((1, 16, 7, 7, 32), "float32"), placeholder_1: T.B ax4 = Ts.axis.spatial(512, i0_i1_i2_i3_i4_fused % 512) Ts.reads(placeholder[0, (ax4 * 49 + ax2 * 7 + ax3) % 25088 // 1568, (ax2 * 7 + ax3) % 49 // 7, ax3 % 7, (ax4 * 49 + ax2 * 7 + ax3) % 1568 // 49], placeholder_1[(ax4 * 49 + ax2 * 7 + ax3) % 25088]) Ts.writes(T_layout_trans[ax0, ax1, ax2, ax3, ax4]) - T_layout_trans[ax0, ax1, ax2, ax3, ax4] = T.if_then_else(ax0 < 1 and ax1 * 512 + ax4 < 512 and ax2 < 7 and ax3 < 7, T.Select(T.float32(0) < T.if_then_else(T.LT(0, 1) and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 < 512 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7 < 7 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7 < 7, placeholder[0, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 // 32, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 % 32], T.float32(0), dtype="float32"), T.if_then_else(T.LT(0, 1) and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 < 512 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7 < 7 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7 < 7, placeholder[0, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 // 32, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 % 32], T.float32(0), dtype="float32"), T.if_then_else(T.LT(0, 1) and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 < 512 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7 < 7 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7 < 7, placeholder[0, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 // 32, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 % 32], T.float32(0), dtype="float32") * placeholder_1[((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088]), T.float32(0), dtype="float32") + T_layout_trans[ax0, ax1, ax2, ax3, ax4] = T.if_then_else(ax0 < 1 and ax1 * 512 + ax4 < 512 and ax2 < 7 and ax3 < 7, T.Select(T.float32(0) < T.if_then_else(T.LT(0, 1) and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 < 512 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7 < 7 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7 < 7, placeholder[0, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 // 32, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 % 32], T.float32(0)), T.if_then_else(T.LT(0, 1) and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 < 512 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7 < 7 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7 < 7, placeholder[0, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 // 32, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 % 32], T.float32(0)), T.if_then_else(T.LT(0, 1) and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 < 512 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7 < 7 and ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7 < 7, placeholder[0, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 // 32, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 49 // 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 7, ((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088 % 25088 // 49 % 32], T.float32(0)) * placeholder_1[((ax1 * 512 + ax4) * 49 + ax2 * 7 + ax3) % 25088]), T.float32(0)) # fmt: on # pylint: enable=invalid-name,no-member,line-too-long,too-many-nested-blocks,no-self-argument diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_layout.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_layout.py index 0b491ac7a40b..9d03d103e92a 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_layout.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_layout.py @@ -242,7 +242,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((3, 3, 64, 64), i3 = Ts.axis.spatial(64, ax3_fused) Ts.reads(p0[i0, i1 - 1, i2 - 1, i3]) Ts.writes(pad_temp[i0, i1, i2, i3]) - pad_temp[i0, i1, i2, i3] = T.if_then_else(1 <= i1 and i1 < 57 and 1 <= i2 and i2 < 57, p0[i0, i1 - 1, i2 - 1, i3], T.float32(0), dtype="float32") + pad_temp[i0, i1, i2, i3] = T.if_then_else(1 <= i1 and i1 < 57 and 1 <= i2 and i2 < 57, p0[i0, i1 - 1, i2 - 1, i3], T.float32(0)) for i3_0 in T.serial(16): for ax0_ax1_ax2_ax3_fused in T.serial(57600): with Ts.sblock("pad_temp_global"): @@ -327,7 +327,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((3, 3, 64, 64), i3 = Ts.axis.spatial(64, ax3_fused) Ts.reads(p0[i0, i1 - 1, i2 - 1, i3]) Ts.writes(pad_temp[i0, i1, i2, i3]) - pad_temp[i0, i1, i2, i3] = T.if_then_else(1 <= i1 and i1 < 57 and 1 <= i2 and i2 < 57, p0[i0, i1 - 1, i2 - 1, i3], T.float32(0), dtype="float32") + pad_temp[i0, i1, i2, i3] = T.if_then_else(1 <= i1 and i1 < 57 and 1 <= i2 and i2 < 57, p0[i0, i1 - 1, i2 - 1, i3], T.float32(0)) for i3_0 in T.serial(16): for ax0_ax1_ax2_ax3_fused in T.serial(57600): with Ts.sblock("pad_temp_global"): @@ -419,7 +419,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((3, 3, 64, 64), i3 = Ts.axis.spatial(64, ax3_fused) Ts.reads(p0[i0, i1 - 1, i2 - 1, i3]) Ts.writes(pad_temp[i0, i1, i2, i3]) - pad_temp[i0, i1, i2, i3] = T.if_then_else(1 <= i1 and i1 < 57 and 1 <= i2 and i2 < 57, p0[i0, i1 - 1, i2 - 1, i3], T.float32(0), dtype="float32") + pad_temp[i0, i1, i2, i3] = T.if_then_else(1 <= i1 and i1 < 57 and 1 <= i2 and i2 < 57, p0[i0, i1 - 1, i2 - 1, i3], T.float32(0)) for i3_0 in T.serial(16): for ax0_ax1_ax2_ax3_fused in T.serial(57600): with Ts.sblock("pad_temp_global"): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_reduction_block.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_reduction_block.py index d889ba9cd0be..b515cb06c6bd 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_reduction_block.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_reduction_block.py @@ -177,7 +177,7 @@ def main(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256) Ts.writes(T_softmax_expsum_shared[i0_2]) with Ts.init(): T_softmax_expsum_shared[i0_2] = T.float32(0) - T_softmax_expsum_shared[i0_2] = T_softmax_expsum_shared[i0_2] + T.exp(A[i0_2, k] - T_softmax_maxelem_shared[i0_2], dtype="float32") + T_softmax_expsum_shared[i0_2] = T_softmax_expsum_shared[i0_2] + T.exp(A[i0_2, k] - T_softmax_maxelem_shared[i0_2]) for i1_0 in T.serial(8): for i1_1 in T.thread_binding(32, thread="threadIdx.x"): with Ts.sblock("T_softmax_norm"): @@ -186,7 +186,7 @@ def main(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256) Ts.reads(A[i0_3, i1], T_softmax_maxelem_shared[i0_3], T_softmax_expsum_shared[i0_3]) Ts.writes(T_softmax_norm[i0_3, i1]) Ts.sblock_attr({"axis":1}) - T_softmax_norm[i0_3, i1] = T.exp(A[i0_3, i1] - T_softmax_maxelem_shared[i0_3], dtype="float32") / T_softmax_expsum_shared[i0_3] + T_softmax_norm[i0_3, i1] = T.exp(A[i0_3, i1] - T_softmax_maxelem_shared[i0_3]) / T_softmax_expsum_shared[i0_3] # pylint: enable=no-member,invalid-name,unused-variable,no-self-argument,line-too-long,chained-comparison,not-callable,too-many-nested-blocks # fmt: on diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_unbound_block.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_unbound_block.py index f35e2ee83ccd..100087058788 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_unbound_block.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_unbound_block.py @@ -86,7 +86,7 @@ def main(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer((1,), "float32")) -> for i0 in T.serial(1): with Ts.sblock("D"): b = Ts.axis.S(1, i0) - D[b] = T.sqrt(C[b], dtype="float32") + D[b] = T.sqrt(C[b]) @tvm.script.ir_module @@ -107,7 +107,7 @@ def main(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer((1,), "float32")) -> for i0_fused_1 in T.thread_binding(1, thread="threadIdx.x"): with Ts.sblock("D"): b = Ts.axis.S(1, 0) - D[b] = T.sqrt(C[b], dtype="float32") + D[b] = T.sqrt(C[b]) @tvm.script.ir_module diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_gpu_code.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_gpu_code.py index b0b5709e78fb..b90100292193 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_gpu_code.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_gpu_code.py @@ -77,7 +77,7 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * 1 <= blockIdx_z // 14 + ry and blockIdx_z // 14 + ry < 15 and 1 <= rx + blockIdx_z % 14 and rx + blockIdx_z % 14 < 15, A[T.ramp(ry * 917504 + blockIdx_z * 65536 + rx * 65536 + rc_outer * 2048 + threadIdx_y * 256 + blockIdx_x * 64 + threadIdx_x * 8 + ax3_inner_outer * 4 - 983040, 1, 4)], T.broadcast(T.float32(0), 4), - dtype="float32x4", + ) for rc_inner in T.serial(0, 8): for ax3 in T.serial(0, 8): @@ -121,7 +121,7 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * 1 <= blockIdx_z // 14 + ry and blockIdx_z // 14 + ry < 15 and 1 <= rx + blockIdx_z % 14 and rx + blockIdx_z % 14 < 15, A[T.ramp(ry * 917504 + blockIdx_z * 65536 + rx * 65536 + rc_outer * 2048 + threadIdx_y * 256 + blockIdx_x * 64 + threadIdx_x * 8 + ax3_inner_outer * 4 - 983040, 1, 4)], T.broadcast(T.float32(0), 4), - dtype="float32x4", + ) for rc_inner in T.serial(0, 8): for ax3 in T.serial(0, 8): @@ -161,7 +161,7 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * 1 <= blockIdx_z // 14 + ry and blockIdx_z // 14 + ry < 15 and 1 <= rx + blockIdx_z % 14 and rx + blockIdx_z % 14 < 15, A[T.ramp(ry * 917504 + blockIdx_z * 65536 + rx * 65536 + rc_outer * 2048 + threadIdx_y * 256 + blockIdx_x * 64 + threadIdx_x * 8 + ax3_inner_outer * 4 - 983040, 1, 4)], T.broadcast(T.float32(0), 4), - dtype="float32x4", + ) # Access of the last element of Apad_shared prevents # buffer compacting from reducing the amount of shared @@ -205,7 +205,7 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * 1 <= blockIdx_z // 14 + ry and blockIdx_z // 14 + ry < 15 and 1 <= rx + blockIdx_z % 14 and rx + blockIdx_z % 14 < 15, A[T.ramp(ry * 917504 + blockIdx_z * 65536 + rx * 65536 + rc_outer * 2048 + threadIdx_y * 256 + blockIdx_x * 64 + threadIdx_x * 8 + ax3_inner_outer * 4 - 983040, 1, 4)], T.broadcast(T.float32(0), 4), - dtype="float32x4", + ) for rc_inner in T.serial(0, 8): for ax3 in T.serial(0, 8): @@ -456,7 +456,7 @@ def GMMCUDATensorCore( 16, C.elem_offset // 256 + C.elem_offset % 256 // 16, T.float32(0), - dtype="handle", + ) ) for ax3_0_0 in T.serial(32): @@ -574,11 +574,11 @@ def GMMCUDATensorCore( A.elem_offset, s1 * 16, 1, - dtype="handle", + ), s1, "row_major", - dtype="handle", + ) ) for ax0_0, ax1_0 in T.grid(1, 4): @@ -627,11 +627,11 @@ def GMMCUDATensorCore( A_1.elem_offset, s1_1 * 16, 1, - dtype="handle", + ), s1_1, "row_major", - dtype="handle", + ) ) for ax0_3, ax1_0_3, ax2_0_3, ax3_0_2, ax0_4, ax1_0_4, ax2_0_4 in T.grid( @@ -713,7 +713,7 @@ def GMMCUDATensorCore( B.elem_offset // 256, C_3.data, C_3.elem_offset // 256 + C_3.elem_offset % 256 // 16, - dtype="handle", + ) ) for ax0_0, ax1_0 in T.grid(4, 4): @@ -761,11 +761,11 @@ def GMMCUDATensorCore( C_4.elem_offset, s1_2 * 16, 2, - dtype="handle", + ), s1_2, "row_major", - dtype="handle", + ) ) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_inline.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_inline.py index 8d1b994be971..cc4189b71bef 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_inline.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_inline.py @@ -44,7 +44,7 @@ def main(X: T.Buffer([1, 512, 56, 56], dtype='float32'), W: T.Buffer([512, 512, for i0, i1, i2, i3 in T.grid(1, 512, 58, 58): with Ts.sblock("pad_temp"): i0_1, i1_1, i2_1, i3_1 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) - pad_temp[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(i2_1 >= 1 and i2_1 < 57 and i3_1 >= 1 and i3_1 < 57, X[i0_1, i1_1, i2_1 - 1, i3_1 - 1], T.float32(0), dtype="float32") + pad_temp[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(i2_1 >= 1 and i2_1 < 57 and i3_1 >= 1 and i3_1 < 57, X[i0_1, i1_1, i2_1 - 1, i3_1 - 1], T.float32(0)) for i0, i1, i2, i3, i4, i5, i6 in T.grid(1, 512, 56, 56, 512, 3, 3): with Ts.sblock("compute"): nn, ff, yy, xx, rc, ry, rx = Ts.axis.remap("SSSSRRR", [i0, i1, i2, i3, i4, i5, i6]) @@ -78,7 +78,7 @@ def main(X: T.Buffer([1, 512, 56, 56], dtype='float32'), W: T.Buffer([512, 512, for i0, i1, i2, i3 in T.grid(1, 512, 58, 58): with Ts.sblock("pad_temp"): i0_1, i1_1, i2_1, i3_1 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) - pad_temp[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(i2_1 >= 1 and i2_1 < 57 and i3_1 >= 1 and i3_1 < 57, X[i0_1, i1_1, i2_1 - 1, i3_1 - 1], T.float32(0), dtype="float32") + pad_temp[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(i2_1 >= 1 and i2_1 < 57 and i3_1 >= 1 and i3_1 < 57, X[i0_1, i1_1, i2_1 - 1, i3_1 - 1], T.float32(0)) for i0, i1, i2, i3, i4, i5, i6 in T.grid(1, 512, 56, 56, 512, 3, 3): with Ts.sblock("compute"): nn, ff, yy, xx, rc, ry, rx = Ts.axis.remap("SSSSRRR", [i0, i1, i2, i3, i4, i5, i6]) @@ -103,7 +103,7 @@ def main(X: T.Buffer([1, 512, 56, 56], dtype='float32'), W: T.Buffer([512, 512, for i0, i1, i2, i3 in T.grid(1, 512, 58, 58): with Ts.sblock("pad_temp"): i0_1, i1_1, i2_1, i3_1 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) - pad_temp[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(i2_1 >= 1 and i2_1 < 57 and i3_1 >= 1 and i3_1 < 57, X[i0_1, i1_1, i2_1 - 1, i3_1 - 1], T.float32(0), dtype="float32") + pad_temp[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(i2_1 >= 1 and i2_1 < 57 and i3_1 >= 1 and i3_1 < 57, X[i0_1, i1_1, i2_1 - 1, i3_1 - 1], T.float32(0)) for i0_0_i1_0_i2_0_i3_0_fused in T.thread_binding(0, 224, thread="blockIdx.x"): for i0_1_i1_1_i2_1_i3_1_fused in T.thread_binding(0, 2, thread="vthread.x"): for i0_2_i1_2_i2_2_i3_2_fused in T.thread_binding(0, 8, thread="threadIdx.x"): @@ -165,7 +165,7 @@ def main(X: T.Buffer((1, 512, 56, 56), "float32"), W: T.Buffer((512, 512, 3, 3), ry, rx = Ts.axis.remap("RR", [i5_0, i6_2]) with Ts.init(): compute_local[nn, ff, yy, xx] = T.float32(0) - compute_local[nn, ff, yy, xx] = compute_local[nn, ff, yy, xx] + T.if_then_else(yy + ry >= 1 and yy + ry < 57 and xx + rx >= 1 and xx + rx < 57, X[nn, rc, yy + ry - 1, xx + rx - 1], T.float32(0), dtype="float32") * W[ff, rc, ry, rx] + compute_local[nn, ff, yy, xx] = compute_local[nn, ff, yy, xx] + T.if_then_else(yy + ry >= 1 and yy + ry < 57 and xx + rx >= 1 and xx + rx < 57, X[nn, rc, yy + ry - 1, xx + rx - 1], T.float32(0)) * W[ff, rc, ry, rx] for ax0, ax1, ax2, ax3 in T.grid(1, 8, 2, 28): with Ts.sblock("compute_local"): v0 = Ts.axis.spatial(1, ax0) @@ -190,7 +190,7 @@ def main(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256) for i0, i1 in T.grid(256, 256): with Ts.sblock("T_softmax_exp"): i0_2, i1_1 = Ts.axis.remap("SS", [i0, i1]) - T_softmax_exp[i0_2, i1_1] = T.exp(A[i0_2, i1_1] - T_softmax_maxelem[i0_2], dtype="float32") + T_softmax_exp[i0_2, i1_1] = T.exp(A[i0_2, i1_1] - T_softmax_maxelem[i0_2]) for i0_3, i1 in T.grid(256, 256): with Ts.sblock("T_softmax_expsum"): i0_4, k = Ts.axis.remap("SR", [i0_3, i1]) @@ -219,11 +219,11 @@ def main(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256) i0_2, k = Ts.axis.remap("SR", [i0, i1]) with Ts.init(): T_softmax_expsum[i0_2] = T.float32(0) - T_softmax_expsum[i0_2] = T_softmax_expsum[i0_2] + T.exp(A[i0_2, k] - T_softmax_maxelem[i0_2], dtype="float32") + T_softmax_expsum[i0_2] = T_softmax_expsum[i0_2] + T.exp(A[i0_2, k] - T_softmax_maxelem[i0_2]) for i0_3, i1 in T.grid(256, 256): with Ts.sblock("T_softmax_norm"): i0_4, i1_1 = Ts.axis.remap("SS", [i0_3, i1]) - T_softmax_norm[i0_4, i1_1] = T.exp(A[i0_4, i1_1] - T_softmax_maxelem[i0_4], dtype="float32") / T_softmax_expsum[i0_4] + T_softmax_norm[i0_4, i1_1] = T.exp(A[i0_4, i1_1] - T_softmax_maxelem[i0_4]) / T_softmax_expsum[i0_4] @tvm.script.ir_module class BeforePureSpatial: @@ -374,7 +374,7 @@ def main(p0: T.Buffer((16, 14, 14, 256), "int8"), p1: T.Buffer((1024, 1, 1, 256) i0_2, i1_2, i2_2, i3_2 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(T_add[i0_2, i1_2, i2_2, i3_2], p4[i3_2], p5[i3_2], p6[i3_2]) Ts.writes(compute_1[i0_2, i1_2, i2_2, i3_2]) - compute_1[i0_2, i1_2, i2_2, i3_2] = T.q_multiply_shift_per_axis(T_add[i0_2, i1_2, i2_2, i3_2], p4[i3_2], p5[i3_2], p6[i3_2], 31, False, True, dtype="int32") + compute_1[i0_2, i1_2, i2_2, i3_2] = T.q_multiply_shift_per_axis(T_add[i0_2, i1_2, i2_2, i3_2], p4[i3_2], p5[i3_2], p6[i3_2], 31, False, True) for i0_3, i1_3, i2_3, i3_3 in T.grid(16, 14, 14, 1024): with Ts.sblock("T_add_1"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0_3, i1_3, i2_3, i3_3]) @@ -398,7 +398,7 @@ def main(p0: T.Buffer((16, 14, 14, 256), "int8"), p1: T.Buffer((1024, 1, 1, 256) i0_8, i1_8, i2_8, i3_8 = Ts.axis.remap("SSSS", [i0_7, i1_7, i2_7, i3_7]) Ts.reads(T_subtract_1[i0_8, i1_8, i2_8, i3_8]) Ts.writes(compute_3[i0_8, i1_8, i2_8, i3_8]) - compute_3[i0_8, i1_8, i2_8, i3_8] = T.q_multiply_shift(T_subtract_1[i0_8, i1_8, i2_8, i3_8], 1408572815, 31, 1, dtype="int32") + compute_3[i0_8, i1_8, i2_8, i3_8] = T.q_multiply_shift(T_subtract_1[i0_8, i1_8, i2_8, i3_8], 1408572815, 31, 1) for i0_9, i1_9, i2_9, i3_9 in T.grid(16, 14, 14, 1024): with Ts.sblock("T_add_2"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0_9, i1_9, i2_9, i3_9]) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_cross_thread_reduction.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_cross_thread_reduction.py index a07044c5b3e7..dc0c298ac58d 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_cross_thread_reduction.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_cross_thread_reduction.py @@ -49,15 +49,14 @@ def main( with Ts.init(): T_softmax_expsum[i0_2] = T.float32(0) T_softmax_expsum[i0_2] = T_softmax_expsum[i0_2] + T.exp( - A[i0_2, k] - T_softmax_maxelem[i0_2], dtype="float32" + A[i0_2, k] - T_softmax_maxelem[i0_2] ) for i0_3, i1 in T.grid(256, 256): with Ts.sblock("T_softmax_norm"): i0_4, i1_1 = Ts.axis.remap("SS", [i0_3, i1]) Ts.sblock_attr({"axis": 1}) T_softmax_norm[i0_4, i1_1] = ( - T.exp(A[i0_4, i1_1] - T_softmax_maxelem[i0_4], dtype="float32") - / T_softmax_expsum[i0_4] + T.exp(A[i0_4, i1_1] - T_softmax_maxelem[i0_4]) / T_softmax_expsum[i0_4] ) @@ -87,9 +86,7 @@ def softmax_mn_0( i0_2, i1_1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(A[i0_2, i1_1], T_softmax_maxelem[i0_2]) Ts.writes(T_softmax_exp[i0_2, i1_1]) - T_softmax_exp[i0_2, i1_1] = T.exp( - A[i0_2, i1_1] - T_softmax_maxelem[i0_2], dtype="float32" - ) + T_softmax_exp[i0_2, i1_1] = T.exp(A[i0_2, i1_1] - T_softmax_maxelem[i0_2]) for i0_3, i1 in T.grid(256, 256): with Ts.sblock("T_softmax_expsum"): i0_4, k = Ts.axis.remap("SR", [i0_3, i1]) @@ -140,7 +137,7 @@ def softmax_mn_1( Ts.reads(A[i0_2, i1], T_softmax_maxelem_shared[i0_2]) Ts.writes(T_softmax_exp[i0_2, i1]) T_softmax_exp[i0_2, i1] = T.exp( - A[i0_2, i1] - T_softmax_maxelem_shared[i0_2], dtype="float32" + A[i0_2, i1] - T_softmax_maxelem_shared[i0_2] ) for i0_3, i1 in T.grid(256, 256): with Ts.sblock("T_softmax_expsum"): @@ -182,9 +179,7 @@ def softmax_mn_2( i0_2, i1_1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(A[i0_2, i1_1], T_softmax_maxelem[i0_2]) Ts.writes(T_softmax_exp[i0_2, i1_1]) - T_softmax_exp[i0_2, i1_1] = T.exp( - A[i0_2, i1_1] - T_softmax_maxelem[i0_2], dtype="float32" - ) + T_softmax_exp[i0_2, i1_1] = T.exp(A[i0_2, i1_1] - T_softmax_maxelem[i0_2]) for i0_3 in T.serial(256): for ax0, ax1_0 in T.grid(1, 32): for ax1_1 in T.thread_binding(8, thread="threadIdx.x"): @@ -244,7 +239,7 @@ def softmax_mn_3( Ts.reads(A[i0_2, i1], T_softmax_maxelem_shared[i0_2]) Ts.writes(T_softmax_exp[i0_2, i1]) T_softmax_exp[i0_2, i1] = T.exp( - A[i0_2, i1] - T_softmax_maxelem_shared[i0_2], dtype="float32" + A[i0_2, i1] - T_softmax_maxelem_shared[i0_2] ) for i0_3 in T.serial(256): for ax0, ax1_0 in T.grid(1, 32): @@ -320,7 +315,7 @@ def softmax_mn_after_inline_0( with Ts.init(): T_softmax_expsum[i0_2] = T.float32(0) T_softmax_expsum[i0_2] = T_softmax_expsum[i0_2] + T.exp( - A[i0_2, k] - T_softmax_maxelem[i0_2], dtype="float32" + A[i0_2, k] - T_softmax_maxelem[i0_2] ) for i0_3, i1 in T.grid(256, 256): with Ts.sblock("T_softmax_norm"): @@ -329,8 +324,7 @@ def softmax_mn_after_inline_0( Ts.writes(T_softmax_norm[i0_4, i1_1]) Ts.sblock_attr({"axis": 1}) T_softmax_norm[i0_4, i1_1] = ( - T.exp(A[i0_4, i1_1] - T_softmax_maxelem[i0_4], dtype="float32") - / T_softmax_expsum[i0_4] + T.exp(A[i0_4, i1_1] - T_softmax_maxelem[i0_4]) / T_softmax_expsum[i0_4] ) @Ts.prim_func @@ -357,7 +351,7 @@ def softmax_mn_after_inline_1( with Ts.init(): T_softmax_expsum[i0_2] = T.float32(0) T_softmax_expsum[i0_2] = T_softmax_expsum[i0_2] + T.exp( - A[i0_2, k] - T_softmax_maxelem[i0_2], dtype="float32" + A[i0_2, k] - T_softmax_maxelem[i0_2] ) for i0_3, i1 in T.grid(256, 256): with Ts.sblock("T_softmax_norm"): @@ -366,8 +360,7 @@ def softmax_mn_after_inline_1( Ts.writes(T_softmax_norm[i0_4, i1_1]) Ts.sblock_attr({"axis": 1}) T_softmax_norm[i0_4, i1_1] = ( - T.exp(A[i0_4, i1_1] - T_softmax_maxelem[i0_4], dtype="float32") - / T_softmax_expsum[i0_4] + T.exp(A[i0_4, i1_1] - T_softmax_maxelem[i0_4]) / T_softmax_expsum[i0_4] ) @Ts.prim_func @@ -396,7 +389,7 @@ def softmax_mn_after_inline_2( with Ts.init(): T_softmax_expsum_shared[i0_2] = T.float32(0) T_softmax_expsum_shared[i0_2] = T_softmax_expsum_shared[i0_2] + T.exp( - A[i0_2, k] - T_softmax_maxelem[i0_2], dtype="float32" + A[i0_2, k] - T_softmax_maxelem[i0_2] ) for i1_0 in T.serial(1): for i1_1 in T.thread_binding(512, thread="threadIdx.x"): @@ -410,7 +403,7 @@ def softmax_mn_after_inline_2( Ts.writes(T_softmax_norm[i0_4, i1_1_1]) Ts.sblock_attr({"axis": 1}) T_softmax_norm[i0_4, i1_1_1] = ( - T.exp(A[i0_4, i1_1_1] - T_softmax_maxelem[i0_4], dtype="float32") + T.exp(A[i0_4, i1_1_1] - T_softmax_maxelem[i0_4]) / T_softmax_expsum_shared[i0_4] ) @@ -445,7 +438,7 @@ def softmax_mn_after_inline_3( with Ts.init(): T_softmax_expsum_shared[i0_2] = T.float32(0) T_softmax_expsum_shared[i0_2] = T_softmax_expsum_shared[i0_2] + T.exp( - A[i0_2, k] - T_softmax_maxelem_shared[i0_2], dtype="float32" + A[i0_2, k] - T_softmax_maxelem_shared[i0_2] ) for i1_0 in T.serial(1): for i1_1 in T.thread_binding(512, thread="threadIdx.x"): @@ -461,7 +454,7 @@ def softmax_mn_after_inline_3( Ts.writes(T_softmax_norm[i0_4, i1_1_1]) Ts.sblock_attr({"axis": 1}) T_softmax_norm[i0_4, i1_1_1] = ( - T.exp(A[i0_4, i1_1_1] - T_softmax_maxelem_shared[i0_4], dtype="float32") + T.exp(A[i0_4, i1_1_1] - T_softmax_maxelem_shared[i0_4]) / T_softmax_expsum_shared[i0_4] ) @@ -518,7 +511,7 @@ def batch_norm_bmn_0(A: T.Buffer((1, 512, 512), "float32"), D: T.Buffer(1, "floa b = Ts.axis.spatial(1, i0) Ts.reads(C[b]) Ts.writes(D[b]) - D[b] = T.sqrt(C[b], dtype="float32") + D[b] = T.sqrt(C[b]) @Ts.prim_func def batch_norm_bmn_1(A: T.Buffer((1, 512, 512), "float32"), D: T.Buffer(1, "float32")) -> None: @@ -545,7 +538,7 @@ def batch_norm_bmn_1(A: T.Buffer((1, 512, 512), "float32"), D: T.Buffer(1, "floa b = Ts.axis.spatial(1, i0_0 * 256 + i0_1) Ts.reads(C_shared[b]) Ts.writes(D[b]) - D[b] = T.sqrt(C_shared[b], dtype="float32") + D[b] = T.sqrt(C_shared[b]) decision_0 = [] # type: ignore decision_1 = [ diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_parallel_vectorize_unroll.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_parallel_vectorize_unroll.py index 524cd5c25481..3d40ecf663fe 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_parallel_vectorize_unroll.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_parallel_vectorize_unroll.py @@ -96,7 +96,7 @@ def main(placeholder: T.Buffer((1, 13, 13, 3, 85), "float32"), placeholder_1: T. ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_strided_slice_with_axes[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_sigmoid[ax0, ax1, ax2, ax3, ax4]) - T_sigmoid[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_sigmoid[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 52, 52, 3, 80): with Ts.sblock("T_strided_slice_with_axes_1"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -108,7 +108,7 @@ def main(placeholder: T.Buffer((1, 13, 13, 3, 85), "float32"), placeholder_1: T. ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_strided_slice_with_axes_1[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_sigmoid_1[ax0, ax1, ax2, ax3, ax4]) - T_sigmoid_1[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes_1[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_sigmoid_1[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes_1[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 52, 52, 3, 80): with Ts.sblock("T_multiply"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -132,7 +132,7 @@ def main(placeholder: T.Buffer((1, 13, 13, 3, 85), "float32"), placeholder_1: T. ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_strided_slice_with_axes_2[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_sigmoid_2[ax0, ax1, ax2, ax3, ax4]) - T_sigmoid_2[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes_2[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_sigmoid_2[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes_2[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 26, 26, 3, 80): with Ts.sblock("T_strided_slice_with_axes_3"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -144,7 +144,7 @@ def main(placeholder: T.Buffer((1, 13, 13, 3, 85), "float32"), placeholder_1: T. ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_strided_slice_with_axes_3[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_sigmoid_3[ax0, ax1, ax2, ax3, ax4]) - T_sigmoid_3[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes_3[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_sigmoid_3[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes_3[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 26, 26, 3, 80): with Ts.sblock("T_multiply_1"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -168,7 +168,7 @@ def main(placeholder: T.Buffer((1, 13, 13, 3, 85), "float32"), placeholder_1: T. ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_strided_slice_with_axes_4[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_sigmoid_4[ax0, ax1, ax2, ax3, ax4]) - T_sigmoid_4[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes_4[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_sigmoid_4[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes_4[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 13, 13, 3, 80): with Ts.sblock("T_strided_slice_with_axes_5"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -180,7 +180,7 @@ def main(placeholder: T.Buffer((1, 13, 13, 3, 85), "float32"), placeholder_1: T. ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_strided_slice_with_axes_5[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_sigmoid_5[ax0, ax1, ax2, ax3, ax4]) - T_sigmoid_5[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes_5[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_sigmoid_5[ax0, ax1, ax2, ax3, ax4] = T.sigmoid(T_strided_slice_with_axes_5[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 13, 13, 3, 80): with Ts.sblock("T_multiply_2"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -198,7 +198,7 @@ def main(placeholder: T.Buffer((1, 13, 13, 3, 85), "float32"), placeholder_1: T. ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(T_reshape[ax0 - 2535, ax1], T_reshape_1[ax0 - 507, ax1], T_reshape_2[ax0, ax1]) Ts.writes(T_concat[ax0, ax1]) - T_concat[ax0, ax1] = T.if_then_else(2535 <= ax0, T_reshape[ax0 - 2535, ax1], T.if_then_else(507 <= ax0, T_reshape_1[ax0 - 507, ax1], T_reshape_2[ax0, ax1], dtype="float32"), dtype="float32") + T_concat[ax0, ax1] = T.if_then_else(2535 <= ax0, T_reshape[ax0 - 2535, ax1], T.if_then_else(507 <= ax0, T_reshape_1[ax0 - 507, ax1], T_reshape_2[ax0, ax1])) for i0, i1 in T.grid(80, 10647): with Ts.sblock("T_transpose"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py index 516421704690..bcc2775ea9e1 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py @@ -439,7 +439,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(T_add_1[ax0, ax1, ax2, ax3], p6[0, 0, 0, ax3]) Ts.writes(T_right_shift[ax0, ax1, ax2, ax3]) - T_right_shift[ax0, ax1, ax2, ax3] = T.shift_right(T_add_1[ax0, ax1, ax2, ax3], p6[0, 0, 0, ax3], dtype="int64") + T_right_shift[ax0, ax1, ax2, ax3] = T.shift_right(T_add_1[ax0, ax1, ax2, ax3], p6[0, 0, 0, ax3]) for i0, i1, i2, i3 in T.grid(16, 56, 56, 256): with Ts.sblock("T_cast_1"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) @@ -481,7 +481,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " i0_7, i1_7, i2_7, i3_7 = Ts.axis.remap("SSSS", [i0_6, i1_6, i2_6, i3_6]) Ts.reads(T_subtract_1[i0_7, i1_7, i2_7, i3_7]) Ts.writes(compute[i0_7, i1_7, i2_7, i3_7]) - compute[i0_7, i1_7, i2_7, i3_7] = T.q_multiply_shift(T_subtract_1[i0_7, i1_7, i2_7, i3_7], 1963325822, 31, 1, dtype="int32") + compute[i0_7, i1_7, i2_7, i3_7] = T.q_multiply_shift(T_subtract_1[i0_7, i1_7, i2_7, i3_7], 1963325822, 31, 1) @tvm.script.ir_module class Conv2dInt8_target: @@ -558,7 +558,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(T_add_1[ax0, ax1, ax2, ax3], p6[0, 0, 0, ax3]) Ts.writes(T_right_shift[ax0, ax1, ax2, ax3]) - T_right_shift[ax0, ax1, ax2, ax3] = T.shift_right(T_add_1[ax0, ax1, ax2, ax3], p6[0, 0, 0, ax3], dtype="int64") + T_right_shift[ax0, ax1, ax2, ax3] = T.shift_right(T_add_1[ax0, ax1, ax2, ax3], p6[0, 0, 0, ax3]) for i0, i1, i2, i3 in T.grid(16, 56, 56, 256): with Ts.sblock("T_cast_1"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) @@ -600,7 +600,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " i0_7, i1_7, i2_7, i3_7 = Ts.axis.remap("SSSS", [i0_6, i1_6, i2_6, i3_6]) Ts.reads(T_subtract_1[i0_7, i1_7, i2_7, i3_7]) Ts.writes(compute_2[i0_7, i1_7, i2_7, i3_7]) - compute_2[i0_7, i1_7, i2_7, i3_7] = T.q_multiply_shift(T_subtract_1[i0_7, i1_7, i2_7, i3_7], 1098990753, 31, 1, dtype="int32") + compute_2[i0_7, i1_7, i2_7, i3_7] = T.q_multiply_shift(T_subtract_1[i0_7, i1_7, i2_7, i3_7], 1098990753, 31, 1) for i0_8, i1_8, i2_8, i3_8 in T.grid(16, 56, 56, 256): with Ts.sblock("T_add_3"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0_8, i1_8, i2_8, i3_8]) @@ -825,7 +825,7 @@ def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_add_1[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_floor[ax0, ax1, ax2, ax3, ax4]) - T_floor[ax0, ax1, ax2, ax3, ax4] = T.floor(T_add_1[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_floor[ax0, ax1, ax2, ax3, ax4] = T.floor(T_add_1[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 128, 7, 7, 16): with Ts.sblock("T_cast_1"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -878,7 +878,7 @@ def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_add_2[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_floor_1[ax0, ax1, ax2, ax3, ax4]) - T_floor_1[ax0, ax1, ax2, ax3, ax4] = T.floor(T_add_2[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_floor_1[ax0, ax1, ax2, ax3, ax4] = T.floor(T_add_2[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 128, 7, 7, 16): with Ts.sblock("T_cast_4"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -995,7 +995,7 @@ def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_add_1[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_floor[ax0, ax1, ax2, ax3, ax4]) - T_floor[ax0, ax1, ax2, ax3, ax4] = T.floor(T_add_1[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_floor[ax0, ax1, ax2, ax3, ax4] = T.floor(T_add_1[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 128, 7, 7, 16): with Ts.sblock("T_cast_1"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -1048,7 +1048,7 @@ def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_add_2[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_floor_1[ax0, ax1, ax2, ax3, ax4]) - T_floor_1[ax0, ax1, ax2, ax3, ax4] = T.floor(T_add_2[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_floor_1[ax0, ax1, ax2, ax3, ax4] = T.floor(T_add_2[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 128, 7, 7, 16): with Ts.sblock("T_cast_4"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -1088,7 +1088,7 @@ def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) Ts.reads(T_add_3[ax0, ax1, ax2, ax3, ax4]) Ts.writes(T_floor_2[ax0, ax1, ax2, ax3, ax4]) - T_floor_2[ax0, ax1, ax2, ax3, ax4] = T.floor(T_add_3[ax0, ax1, ax2, ax3, ax4], dtype="float32") + T_floor_2[ax0, ax1, ax2, ax3, ax4] = T.floor(T_add_3[ax0, ax1, ax2, ax3, ax4]) for i0, i1, i2, i3, i4 in T.grid(1, 128, 7, 7, 16): with Ts.sblock("T_cast_6"): ax0, ax1, ax2, ax3, ax4 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) @@ -1185,7 +1185,7 @@ def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, ax3_1, ax4 = Ts.axis.remap("SS", [ax3, ax4_fused]) Ts.reads(conv2d_NCHWc_int8[ax0_1, ax1_1, ax2_1, ax3_1, ax4], p2[ax0_1, ax1_1, 0, 0, ax4], p3[ax0_1, ax1_1, 0, 0, ax4], p4[0], p5[ax0_1, ax1_1, ax2_1, ax3_1, ax4]) Ts.writes(T_cast[ax0_1, ax1_1, ax2_1, ax3_1, ax4]) - T_cast[ax0_1, ax1_1, ax2_1, ax3_1, ax4] = T.cast(T.max(T.min(T.cast(T.max(T.min(T.cast(T.floor(T.float32(0.95489668846130371) * (T.cast(T.cast(T.max(T.min(T.cast(T.floor(T.cast(conv2d_NCHWc_int8[ax0_1, ax1_1, ax2_1, ax3_1, ax4] + p2[ax0_1, ax1_1, 0, 0, ax4], "float32") * p3[ax0_1, ax1_1, 0, 0, ax4] + T.float32(65.5), dtype="float32"), "int32"), 255), 0), "uint8"), "float32") - p4[0]) + T.float32(0.5), dtype="float32"), "int32") + T.cast(T.floor(T.float32(0.71245479583740234) * T.cast(p5[ax0_1, ax1_1, ax2_1, ax3_1, ax4], "float32") + T.float32(0.5), dtype="float32"), "int32"), 255), 0), "uint8"), T.uint8(255)), T.uint8(0)), "int32") + T_cast[ax0_1, ax1_1, ax2_1, ax3_1, ax4] = T.cast(T.max(T.min(T.cast(T.max(T.min(T.cast(T.floor(T.float32(0.95489668846130371) * (T.cast(T.cast(T.max(T.min(T.cast(T.floor(T.cast(conv2d_NCHWc_int8[ax0_1, ax1_1, ax2_1, ax3_1, ax4] + p2[ax0_1, ax1_1, 0, 0, ax4], "float32") * p3[ax0_1, ax1_1, 0, 0, ax4] + T.float32(65.5)), "int32"), 255), 0), "uint8"), "float32") - p4[0]) + T.float32(0.5)), "int32") + T.cast(T.floor(T.float32(0.71245479583740234) * T.cast(p5[ax0_1, ax1_1, ax2_1, ax3_1, ax4], "float32") + T.float32(0.5)), "int32"), 255), 0), "uint8"), T.uint8(255)), T.uint8(0)), "int32") return Conv2dInt8_NCHWc_scheduled @@ -1212,7 +1212,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((6, 6, 64, 64), Ts.reads(p0[i0_1, i1_1 - 1, i2_1 - 1, i3_1]) Ts.writes(data_pad[i0_1, i1_1, i2_1, i3_1]) Ts.sblock_attr({"schedule_rule":"None"}) - data_pad[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(1 <= i1_1 and i1_1 < 57 and 1 <= i2_1 and i2_1 < 57, p0[i0_1, i1_1 - 1, i2_1 - 1, i3_1], T.float32(0), dtype="float32") + data_pad[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(1 <= i1_1 and i1_1 < 57 and 1 <= i2_1 and i2_1 < 57, p0[i0_1, i1_1 - 1, i2_1 - 1, i3_1], T.float32(0)) for i0, i1, i2, i3 in T.grid(6, 6, 196, 64): with Ts.sblock("input_tile"): eps, nu, p, ci = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) @@ -1304,7 +1304,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((6, 6, 64, 64), Ts.reads(p0[i0_1, i1_1 - 1, i2_1 - 1, i3_1]) Ts.writes(data_pad[i0_1, i1_1, i2_1, i3_1]) Ts.sblock_attr({"schedule_rule":"None"}) - data_pad[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(1 <= i1_1 and i1_1 < 57 and 1 <= i2_1 and i2_1 < 57, p0[i0_1, i1_1 - 1, i2_1 - 1, i3_1], T.float32(0), dtype="float32") + data_pad[i0_1, i1_1, i2_1, i3_1] = T.if_then_else(1 <= i1_1 and i1_1 < 57 and 1 <= i2_1 and i2_1 < 57, p0[i0_1, i1_1 - 1, i2_1 - 1, i3_1], T.float32(0)) for i0, i1, i2, i3 in T.grid(6, 6, 196, 64): with Ts.sblock("input_tile"): eps, nu, p, ci = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) @@ -1403,7 +1403,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((6, 6, 64, 64), Ts.reads(p0[p // 196, p % 196 // 14 * 4 + eps - 1, p % 14 * 4 + nu - 1, ci]) Ts.writes(input_tile_local[eps, nu, p, ci]) Ts.sblock_attr({"schedule_rule":"None"}) - input_tile_local[eps, nu, p, ci] = T.if_then_else(1 <= p % 196 // 14 * 4 + eps and p % 196 // 14 * 4 + eps < 57 and 1 <= p % 14 * 4 + nu and p % 14 * 4 + nu < 57, p0[p // 196, p % 196 // 14 * 4 + eps - 1, p % 14 * 4 + nu - 1, ci], T.float32(0), dtype="float32") + input_tile_local[eps, nu, p, ci] = T.if_then_else(1 <= p % 196 // 14 * 4 + eps and p % 196 // 14 * 4 + eps < 57 and 1 <= p % 14 * 4 + nu and p % 14 * 4 + nu < 57, p0[p // 196, p % 196 // 14 * 4 + eps - 1, p % 14 * 4 + nu - 1, ci], T.float32(0)) for i0 in T.unroll(6): for i1 in T.unroll(6): with Ts.sblock("data_pack_init"): @@ -1564,7 +1564,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " i0_2, i1_2, i2_2, i3_2 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(T_add[i0_2, i1_2, i2_2, i3_2], p4[i3_2], p5[i3_2], p6[i3_2]) Ts.writes(compute_1[i0_2, i1_2, i2_2, i3_2]) - compute_1[i0_2, i1_2, i2_2, i3_2] = T.q_multiply_shift_per_axis(T_add[i0_2, i1_2, i2_2, i3_2], p4[i3_2], p5[i3_2], p6[i3_2], 31, False, True, dtype="int32") + compute_1[i0_2, i1_2, i2_2, i3_2] = T.q_multiply_shift_per_axis(T_add[i0_2, i1_2, i2_2, i3_2], p4[i3_2], p5[i3_2], p6[i3_2], 31, False, True) for i0_3, i1_3, i2_3, i3_3 in T.grid(16, 56, 56, 256): with Ts.sblock("T_add_1"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0_3, i1_3, i2_3, i3_3]) @@ -1588,7 +1588,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " i0_8, i1_8, i2_8, i3_8 = Ts.axis.remap("SSSS", [i0_7, i1_7, i2_7, i3_7]) Ts.reads(T_subtract_1[i0_8, i1_8, i2_8, i3_8]) Ts.writes(compute[i0_8, i1_8, i2_8, i3_8]) - compute[i0_8, i1_8, i2_8, i3_8] = T.q_multiply_shift(T_subtract_1[i0_8, i1_8, i2_8, i3_8], 1963325822, 31, 1, dtype="int32") + compute[i0_8, i1_8, i2_8, i3_8] = T.q_multiply_shift(T_subtract_1[i0_8, i1_8, i2_8, i3_8], 1963325822, 31, 1) @tvm.script.ir_module class Conv2dInt8_with_predicate_target: @@ -1640,7 +1640,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " i0_2, i1_2, i2_2, i3_2 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(T_add[i0_2, i1_2, i2_2, i3_2], p4[i3_2], p5[i3_2], p6[i3_2]) Ts.writes(compute_1[i0_2, i1_2, i2_2, i3_2]) - compute_1[i0_2, i1_2, i2_2, i3_2] = T.q_multiply_shift_per_axis(T_add[i0_2, i1_2, i2_2, i3_2], p4[i3_2], p5[i3_2], p6[i3_2], 31, False, True, dtype="int32") + compute_1[i0_2, i1_2, i2_2, i3_2] = T.q_multiply_shift_per_axis(T_add[i0_2, i1_2, i2_2, i3_2], p4[i3_2], p5[i3_2], p6[i3_2], 31, False, True) for i0_3, i1_3, i2_3, i3_3 in T.grid(16, 56, 56, 256): with Ts.sblock("T_add_1"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0_3, i1_3, i2_3, i3_3]) @@ -1664,13 +1664,13 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " i0_8, i1_8, i2_8, i3_8 = Ts.axis.remap("SSSS", [i0_7, i1_7, i2_7, i3_7]) Ts.reads(T_subtract_1[i0_8, i1_8, i2_8, i3_8]) Ts.writes(compute_3[i0_8, i1_8, i2_8, i3_8]) - compute_3[i0_8, i1_8, i2_8, i3_8] = T.q_multiply_shift(T_subtract_1[i0_8, i1_8, i2_8, i3_8], 1457846997, 31, 0, dtype="int32") + compute_3[i0_8, i1_8, i2_8, i3_8] = T.q_multiply_shift(T_subtract_1[i0_8, i1_8, i2_8, i3_8], 1457846997, 31, 0) for i0_9, i1_9, i2_9, i3_9 in T.grid(16, 56, 56, 256): with Ts.sblock("compute_3"): i0_10, i1_10, i2_10, i3_10 = Ts.axis.remap("SSSS", [i0_9, i1_9, i2_9, i3_9]) Ts.reads(p9[i0_10, i1_10, i2_10, i3_10]) Ts.writes(compute_4[i0_10, i1_10, i2_10, i3_10]) - compute_4[i0_10, i1_10, i2_10, i3_10] = T.q_multiply_shift(p9[i0_10, i1_10, i2_10, i3_10], 2101000910, 31, 0, dtype="int32") + compute_4[i0_10, i1_10, i2_10, i3_10] = T.q_multiply_shift(p9[i0_10, i1_10, i2_10, i3_10], 2101000910, 31, 0) for i0_11, i1_11, i2_11, i3_11 in T.grid(16, 56, 56, 256): with Ts.sblock("T_add_2"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0_11, i1_11, i2_11, i3_11]) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_cache_index.py b/tests/python/s_tir/schedule/test_tir_schedule_cache_index.py index 9b5313a75fce..271501bc42e3 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_cache_index.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_cache_index.py @@ -91,7 +91,6 @@ def bilinear_resize( (T.Cast("float32", i2_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), 39, @@ -106,7 +105,6 @@ def bilinear_resize( (T.Cast("float32", i3_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), 39, @@ -127,7 +125,6 @@ def bilinear_resize( T.floor( (T.Cast("float32", i3_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), ) @@ -146,7 +143,6 @@ def bilinear_resize( (T.Cast("float32", i2_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), 39, @@ -161,7 +157,6 @@ def bilinear_resize( (T.Cast("float32", i3_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ) + 1, @@ -181,7 +176,6 @@ def bilinear_resize( T.floor( (T.Cast("float32", i3_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), ) @@ -199,7 +193,6 @@ def bilinear_resize( T.floor( (T.Cast("float32", i2_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), ) @@ -219,7 +212,6 @@ def bilinear_resize( (T.Cast("float32", i2_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ) + 1, @@ -235,7 +227,6 @@ def bilinear_resize( (T.Cast("float32", i3_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), 39, @@ -256,7 +247,6 @@ def bilinear_resize( T.floor( (T.Cast("float32", i3_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), ) @@ -275,7 +265,6 @@ def bilinear_resize( (T.Cast("float32", i2_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ) + 1, @@ -291,7 +280,6 @@ def bilinear_resize( (T.Cast("float32", i3_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ) + 1, @@ -311,7 +299,6 @@ def bilinear_resize( T.floor( (T.Cast("float32", i3_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), ) @@ -327,7 +314,6 @@ def bilinear_resize( T.floor( (T.Cast("float32", i2_1) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), ) @@ -357,7 +343,6 @@ def cached_bilinear_resize( T.floor( (T.Cast("float32", v0) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ), ) @@ -371,7 +356,6 @@ def cached_bilinear_resize( "int32", T.floor( (T.Cast("float32", v0) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ) for ax0 in T.serial(80): @@ -383,7 +367,6 @@ def cached_bilinear_resize( "int32", T.floor( (T.Cast("float32", v0) + T.float32(0.5)) * T.float32(0.5) - T.float32(0.5), - dtype="float32", ), ) for i0, i1, i2, i3 in T.grid(1, 3, 80, 80): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_cache_read_write.py b/tests/python/s_tir/schedule/test_tir_schedule_cache_read_write.py index f42cd552376a..33ae465b9ffc 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_cache_read_write.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_cache_read_write.py @@ -280,11 +280,9 @@ def opaque_access( vi * 2048 + vj * 16, 128, 1, - dtype="handle", ), 128, "row_major", - dtype="handle", ) ) for i, j in T.grid(8, 8): @@ -325,11 +323,9 @@ def opaque_access( A0.elem_offset, A0.strides[0], 1, - dtype="handle", ), 128, "row_major", - dtype="handle", ) ) @@ -601,11 +597,9 @@ def cache_read_opaque_access( vi * 2048 + vj * 16, 128, 1, - dtype="handle", ), 128, "row_major", - dtype="handle", ) ) for i, j in T.grid(8, 8): @@ -646,11 +640,9 @@ def cache_read_opaque_access( A0.elem_offset, A0.strides[0], 1, - dtype="handle", ), 128, "row_major", - dtype="handle", ) ) @@ -966,11 +958,9 @@ def cache_write_opaque_access( vi * 2048 + vj * 16, 128, 1, - dtype="handle", ), 128, "row_major", - dtype="handle", ) ) for i, j in T.grid(8, 8): @@ -1011,11 +1001,9 @@ def cache_write_opaque_access( A0.elem_offset, A0.strides[0], 1, - dtype="handle", ), 128, "row_major", - dtype="handle", ) ) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_compute_at.py b/tests/python/s_tir/schedule/test_tir_schedule_compute_at.py index 99983a533d83..f2659a81c287 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_compute_at.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_compute_at.py @@ -739,7 +739,7 @@ def read_out_of_bound(A: T.Buffer([16], 'float32'), C: T.Buffer([16], 'float32') with Ts.sblock("C"): v = Ts.axis.S(16, j) Ts.reads(B[v : v + 2]) - C[v] = T.if_then_else(v < 15, T.max(B[v], B[v + 1]), B[v], dtype="float32") + C[v] = T.if_then_else(v < 15, T.max(B[v], B[v + 1]), B[v]) @Ts.prim_func def read_out_of_bound_after_compute_at(A: T.Buffer([16], 'float32'), C: T.Buffer([16], 'float32')) -> None: @@ -755,7 +755,7 @@ def read_out_of_bound_after_compute_at(A: T.Buffer([16], 'float32'), C: T.Buffer with Ts.sblock("C"): v = Ts.axis.S(16, j) Ts.reads([B[v : v + 2]]) - C[v] = T.if_then_else(v < 15, T.max(B[v], B[v + 1]), B[v], dtype="float32") + C[v] = T.if_then_else(v < 15, T.max(B[v], B[v + 1]), B[v]) @Ts.prim_func def multi_reduction(A: T.Buffer((16, 16), "float32"), C: T.Buffer((), "float32")): @@ -808,11 +808,11 @@ def tiled_pooling_read_cache(X: T.Buffer([224, 224], dtype='float32'), Y: T.Buff with Ts.init(): Y[h, w] = 0.0 Y[h, w] = T.max(Y[h, w], T.if_then_else( - T.likely(1 <= h + kh, dtype="bool") and \ - T.likely(h + kh < 225, dtype="bool") and \ - T.likely(1 <= w + kw, dtype="bool") and \ - T.likely(w + kw < 225, dtype="bool"), - cache[h + kh - 1, w + kw - 1], 0.0, dtype="float32")) + T.likely(1 <= h + kh) and \ + T.likely(h + kh < 225) and \ + T.likely(1 <= w + kw) and \ + T.likely(w + kw < 225), + cache[h + kh - 1, w + kw - 1], 0.0)) @Ts.prim_func def tiled_pooling_read_cache_after_compute_at(X: T.Buffer([224, 224], dtype='float32'), Y: T.Buffer([224, 224], dtype='float32')) -> None: @@ -833,11 +833,11 @@ def tiled_pooling_read_cache_after_compute_at(X: T.Buffer([224, 224], dtype='flo with Ts.init(): Y[h, w] = 0.0 Y[h, w] = T.max(Y[h, w], T.if_then_else( - T.likely(1 <= h + kh, dtype="bool") and \ - T.likely(h + kh < 225, dtype="bool") and \ - T.likely(1 <= w + kw, dtype="bool") and \ - T.likely(w + kw < 225, dtype="bool"), - cache[h + kh - 1, w + kw - 1], 0.0, dtype="float32")) + T.likely(1 <= h + kh) and \ + T.likely(h + kh < 225) and \ + T.likely(1 <= w + kw) and \ + T.likely(w + kw < 225), + cache[h + kh - 1, w + kw - 1], 0.0)) @Ts.prim_func def non_uniform_tiled_conv(x: T.Buffer((1, 3, 100, 100), "float32"), @@ -905,7 +905,7 @@ def concat_two_elemwise(x: T.Buffer((16,), "float32"), for i in T.serial(24): with Ts.sblock("T_concat"): ax = Ts.axis.spatial(24, i) - T_concat[ax] = T.if_then_else(16 <= ax, T_add_2[ax - 16], T_add_1[ax], dtype="float32") + T_concat[ax] = T.if_then_else(16 <= ax, T_add_2[ax - 16], T_add_1[ax]) @Ts.prim_func def concat_two_elemwise_after_compute_at(x: T.Buffer((16,), "float32"), @@ -924,7 +924,7 @@ def concat_two_elemwise_after_compute_at(x: T.Buffer((16,), "float32"), T_add_2[ax] = y[ax] + T.float32(2) with Ts.sblock("T_concat"): ax = Ts.axis.spatial(24, i) - T_concat[ax] = T.if_then_else(16 <= ax, T_add_2[ax - 16], T_add_1[ax], dtype="float32") + T_concat[ax] = T.if_then_else(16 <= ax, T_add_2[ax - 16], T_add_1[ax]) @Ts.prim_func def floordiv_and_floormod_indices(X: T.Buffer([16, 16]), Y: T.Buffer([256])) -> None: @@ -1476,7 +1476,6 @@ def multi_producers_conv( 3 <= i2_1 and i2_1 < 227 and 3 <= i3_1 and i3_1 < 227, data[i0_1, i1_1, i2_1 - 3, i3_1 - 3], T.int8(0), - dtype="int8", ) for i0 in T.serial(1): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 7, 7): @@ -1519,7 +1518,6 @@ def multi_producers_after_compute_at( 3 <= i2_1 and i2_1 < 227 and 3 <= i3_1 and i3_1 < 227, data[i0_1, i1_1, i2_1 - 3, i3_1 - 3], T.int8(0), - dtype="int8", ) for ax0, ax1, ax2, ax3 in T.grid(16, 3, 7, 7): with Ts.sblock("wbuf"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_compute_inline.py b/tests/python/s_tir/schedule/test_tir_schedule_compute_inline.py index 5fc490ff0ae8..786758d3ed5f 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_compute_inline.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_compute_inline.py @@ -548,7 +548,7 @@ def exp_exp_opaque_access_with_tvm_access_ptr( i0_1 = Ts.axis.spatial(16, i0) Ts.reads(x[i0_1]) Ts.writes(compute_1[i0_1]) - compute_1[i0_1] = T.exp(x[i0_1], dtype="float16") + compute_1[i0_1] = T.exp(x[i0_1]) for i0 in T.serial(16): with Ts.sblock("compute_1"): i0_2 = Ts.axis.spatial(16, i0) @@ -557,7 +557,6 @@ def exp_exp_opaque_access_with_tvm_access_ptr( T.evaluate(lookup_table.access_ptr("r")) compute[i0_2] = T.exp( compute_1[i0_2], - dtype="float16", ) @@ -576,8 +575,7 @@ def exp_exp_opaque_access_with_tvm_access_ptr_inlined( Ts.writes(compute[i0_1]) T.evaluate(lookup_table.access_ptr("r")) compute[i0_1] = T.exp( - T.exp(x[i0_1], dtype="float16"), - dtype="float16", + T.exp(x[i0_1]), ) @@ -673,7 +671,7 @@ def elementwise_producer_not_cover_consumer( for i, j in T.grid(256, 128): with Ts.sblock("C"): vi, vj = Ts.axis.remap("SS", [i, j]) - D[vi, vj] = T.if_then_else(vi >= 128, B[vi - 128, vj], T.float32(0), dtype="float32") + D[vi, vj] = T.if_then_else(vi >= 128, B[vi - 128, vj], T.float32(0)) @Ts.prim_func @@ -835,13 +833,13 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " v1 = Ts.axis.spatial(256, ax2_0_0_ax3_0_0_fused % 4 * 64 + (ax1_0 * 256 + ax1_1 * 64 + ax1_2 * 2 + ax1_3)) Ts.reads(p7[()], conv2d_nhwc_reindex_shared[v0, v1], p2[0, 0, 0, v1], p3[0, 0, 0, v1], p4[v1], p5[v1], p6[v1], p8[0]) Ts.writes(compute_3[v0 // 3136, v0 % 3136 // 56, v0 % 56, v1]) - compute_3[v0 // 3136, v0 % 3136 // 56, v0 % 56, v1] = T.q_multiply_shift(T.max(T.min(p7[()] + T.q_multiply_shift_per_axis(conv2d_nhwc_reindex_shared[v0, v1] - p2[0, 0, 0, v1] + p3[0, 0, 0, v1], p4[v1], p5[v1], p6[v1], 31, False, True, dtype="int32"), 255), 0) - p8[0], 1457846997, 31, 0, dtype="int32") + compute_3[v0 // 3136, v0 % 3136 // 56, v0 % 56, v1] = T.q_multiply_shift(T.max(T.min(p7[()] + T.q_multiply_shift_per_axis(conv2d_nhwc_reindex_shared[v0, v1] - p2[0, 0, 0, v1] + p3[0, 0, 0, v1], p4[v1], p5[v1], p6[v1], 31, False, True), 255), 0) - p8[0], 1457846997, 31, 0) for i0_12, i1_12, i2_12, i3_12 in T.grid(16, 56, 56, 256): with Ts.sblock("compute_4"): i0_13, i1_13, i2_13, i3_13 = Ts.axis.remap("SSSS", [i0_12, i1_12, i2_12, i3_12]) Ts.reads(compute_3[i0_13, i1_13, i2_13, i3_13], p9[i0_13, i1_13, i2_13, i3_13]) Ts.writes(compute[i0_13, i1_13, i2_13, i3_13]) - compute[i0_13, i1_13, i2_13, i3_13] = T.max(T.min(compute_3[i0_13, i1_13, i2_13, i3_13] + T.q_multiply_shift(p9[i0_13, i1_13, i2_13, i3_13], 2101000910, 31, 0, dtype="int32"), 255), 0) + compute[i0_13, i1_13, i2_13, i3_13] = T.max(T.min(compute_3[i0_13, i1_13, i2_13, i3_13] + T.q_multiply_shift(p9[i0_13, i1_13, i2_13, i3_13], 2101000910, 31, 0), 255), 0) @tvm.script.ir_module class Conv2dInt8_TensorCore_with_predicate_after: diff --git a/tests/python/s_tir/schedule/test_tir_schedule_decompose_padding.py b/tests/python/s_tir/schedule/test_tir_schedule_decompose_padding.py index 52535e682403..d51cb3d417d2 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_decompose_padding.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_decompose_padding.py @@ -95,7 +95,7 @@ def before_decompose(x: T.Buffer(128, "int32"), y: T.Buffer(140, "int32")): for i in range(140): with Ts.sblock("block"): vi = Ts.axis.remap("S", [i]) - y[vi] = T.if_then_else(vi >= 6 and vi < 134, x[vi - 6], 0, dtype="int32") + y[vi] = T.if_then_else(vi >= 6 and vi < 134, x[vi - 6], 0) @Ts.prim_func def after_decompose(x: T.Buffer(128, "int32"), y: T.Buffer(140, "int32")): @@ -130,7 +130,6 @@ def sum_pool_2d( 3 <= ax2 and ax2 < 228 and 3 <= ax3 and ax3 < 228, x[ax0, ax1, ax2 - 3, ax3 - 3], T.int8(0), - dtype="int8", ) for i0, i1, i2, i3, i4, i5 in T.grid(1, 16, 225, 225, 7, 7): with Ts.sblock("tensor"): @@ -372,7 +371,6 @@ def pad_op( 3 <= ax2 and ax2 < 228 and 3 <= ax3 and ax3 < 228, x[ax0, ax1, ax2 - 3, ax3 - 3], T.int8(0), - dtype="int8", ) @Ts.prim_func @@ -412,7 +410,6 @@ def trivial_pad( 0 <= ax2 and ax2 < 225 and 0 <= ax3 and ax3 < 225, x[ax0, ax1, ax2, ax3], T.int8(0), - dtype="int8", ) sch = tvm.s_tir.Schedule(trivial_pad, debug_mask="all") diff --git a/tests/python/s_tir/schedule/test_tir_schedule_pad_einsum.py b/tests/python/s_tir/schedule/test_tir_schedule_pad_einsum.py index 41e63b437351..e99efe06cb90 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_pad_einsum.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_pad_einsum.py @@ -74,15 +74,13 @@ def matmul_expected( i, j = Ts.axis.remap("SS", [i0, i1]) Ts.reads(A[i, j]) Ts.writes(A_shared_padded[i, j]) - A_shared_padded[i, j] = T.if_then_else(j < 127, A[i, j], T.float32(0), dtype="float32") + A_shared_padded[i, j] = T.if_then_else(j < 127, A[i, j], T.float32(0)) for i0, i1 in T.grid(128, 128): with Ts.sblock("B"): i, j = Ts.axis.remap("SS", [i0, i1]) Ts.reads(B[i, j]) Ts.writes(B_shared_padded[i, j]) - B_shared_padded[i, j] = T.if_then_else( - i < 127 and j < 127, B[i, j], T.float32(0), dtype="float32" - ) + B_shared_padded[i, j] = T.if_then_else(i < 127 and j < 127, B[i, j], T.float32(0)) for i0, i1, i2 in T.grid(128, 128, 128): with Ts.sblock("C_shared"): i, j, k = Ts.axis.remap("SSR", [i0, i1, i2]) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_partition.py b/tests/python/s_tir/schedule/test_tir_schedule_partition.py index 30b4f2608190..a126eaf33fe8 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_partition.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_partition.py @@ -287,7 +287,7 @@ def opaque_access(A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float vi, vj = Ts.axis.remap("SS", [i, j]) Ts.reads([]) Ts.writes([B[0:16, 0:16]]) - T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, vi * 16 + vj, dtype="handle")) + T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, vi * 16 + vj)) @Ts.prim_func diff --git a/tests/python/s_tir/schedule/test_tir_schedule_reindex.py b/tests/python/s_tir/schedule/test_tir_schedule_reindex.py index 9cedc0dd6c97..05042b40bd8e 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_reindex.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_reindex.py @@ -69,7 +69,6 @@ def conv2d_nhwc( ((((i1_1 >= 3) and (i1_1 < 227)) and (i2_1 >= 3)) and (i2_1 < 227)), Input[i0_1, (i1_1 - 3), (i2_1 - 3), i3_1], T.float32(0), - dtype="float32", ) for i0, i1, i2, i3, i4, i5, i6 in T.grid(1, 112, 112, 64, 7, 7, 3): with Ts.sblock("conv2d_nhwc"): @@ -97,7 +96,6 @@ def conv2d_nhwc_reindex_data( ((((i1_1 >= 3) and (i1_1 < 227)) and (i2_1 >= 3)) and (i2_1 < 227)), Input[i0_1, (i1_1 - 3), (i2_1 - 3), i3_1], T.float32(0), - dtype="float32", ) for i0, i1, i2, i3, i4, i5 in T.grid(1, 112, 112, 7, 7, 3): with Ts.sblock("ReindexInput"): @@ -130,7 +128,6 @@ def conv2d_nhwc_reindex_weight( i1_1 >= 3 and i1_1 < 227 and i2_1 >= 3 and i2_1 < 227, inputs[i0_1, i1_1 - 3, i2_1 - 3, i3_1], T.float32(0), - dtype="float32", ) for ax3, ax4, ax5, ax6 in T.grid(64, 7, 7, 3): with Ts.sblock("weight_reindex"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_reorder.py b/tests/python/s_tir/schedule/test_tir_schedule_reorder.py index 96aad565d5af..6938cf3ef1f0 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_reorder.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_reorder.py @@ -163,7 +163,7 @@ def opaque_access(A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float vi, vj = Ts.axis.remap("SS", [i, j]) Ts.reads([]) Ts.writes([B[0:16, 0:16]]) - T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, vi * 16 + vj, dtype="handle")) + T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, vi * 16 + vj)) @Ts.prim_func @@ -181,7 +181,7 @@ def opaque_access_reorder( vi, vj = Ts.axis.remap("SS", [i, j]) Ts.reads([]) Ts.writes([B[0:16, 0:16]]) - T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, vi * 16 + vj, dtype="handle")) + T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, vi * 16 + vj)) # pylint: enable=no-member,invalid-name,unused-variable diff --git a/tests/python/s_tir/schedule/test_tir_schedule_rfactor.py b/tests/python/s_tir/schedule/test_tir_schedule_rfactor.py index aed92ae9ebfe..0519fd5fe2a5 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_rfactor.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_rfactor.py @@ -195,7 +195,7 @@ def transformed_square_sum_square_root(A: T.Buffer([16, 256, 256]), D: T.Buffer( b_1 = Ts.axis.S(16, i0_1) Ts.reads([C[b_1]]) Ts.writes([D[b_1]]) - D[b_1] = T.sqrt(C[b_1], dtype="float32") + D[b_1] = T.sqrt(C[b_1]) @Ts.prim_func @@ -222,7 +222,7 @@ def square_sum_square_root_rfactor(A: T.Buffer([16, 256, 256]), D: T.Buffer([16] for i0_2 in T.serial(0, 16): with Ts.sblock("D"): b_2 = Ts.axis.S(16, i0_2) - D[b_2] = T.sqrt(C[b_2], dtype="float32") + D[b_2] = T.sqrt(C[b_2]) @Ts.prim_func @@ -242,7 +242,7 @@ def transformed_square_sum_square_root_factor_one_1( for i0_1 in T.serial(0, 16): with Ts.sblock("D"): b_1 = Ts.axis.S(16, i0_1) - D[b_1] = T.sqrt(C[b_1], dtype="float32") + D[b_1] = T.sqrt(C[b_1]) @Ts.prim_func @@ -269,7 +269,7 @@ def square_sum_square_root_factor_one_1_rfactor( for i0_1 in T.serial(16): with Ts.sblock("D"): b_1 = Ts.axis.spatial(16, i0_1) - D[b_1] = T.sqrt(C[b_1], dtype="float32") + D[b_1] = T.sqrt(C[b_1]) @Ts.prim_func @@ -289,7 +289,7 @@ def transformed_square_sum_square_root_factor_one_2( for i0_1 in T.serial(0, 16): with Ts.sblock("D"): b_1 = Ts.axis.S(16, i0_1) - D[b_1] = T.sqrt(C[b_1], dtype="float32") + D[b_1] = T.sqrt(C[b_1]) @Ts.prim_func @@ -316,7 +316,7 @@ def square_sum_square_root_factor_one_2_rfactor( for i0_1 in T.serial(16): with Ts.sblock("D"): b_1 = Ts.axis.spatial(16, i0_1) - D[b_1] = T.sqrt(C[b_1], dtype="float32") + D[b_1] = T.sqrt(C[b_1]) @Ts.prim_func diff --git a/tests/python/s_tir/schedule/test_tir_schedule_sampling.py b/tests/python/s_tir/schedule/test_tir_schedule_sampling.py index 42f6acfd0b9b..3820cd0aa210 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_sampling.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_sampling.py @@ -54,7 +54,6 @@ def tiled_conv2d_with_padding( 3 <= i1_1 and i1_1 < 227 and 3 <= i2_1 and i2_1 < 227, inputs[i0_1, i1_1 - 3, i2_1 - 3, i3_1], T.float32(0), - dtype="float32", ) for ( i0_0, diff --git a/tests/python/s_tir/schedule/test_tir_schedule_split_fuse.py b/tests/python/s_tir/schedule/test_tir_schedule_split_fuse.py index ad3663571e14..06aa8c363a84 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_split_fuse.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_split_fuse.py @@ -276,7 +276,7 @@ def opaque_access(A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float vi, vj = Ts.axis.remap("SS", [i, j]) Ts.reads([]) Ts.writes([B[0:16, 0:16]]) - T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, vi * 16 + vj, dtype="handle")) + T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, vi * 16 + vj)) @Ts.prim_func @@ -294,7 +294,7 @@ def opaque_access_fused(A: T.Buffer([16, 16]), B: T.Buffer([16, 16])) -> None: vj = Ts.axis.S(16, T.floormod(i_j_fused, 16)) Ts.reads([]) Ts.writes([B[0:16, 0:16]]) - T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, ((vi * 16) + vj), dtype="handle")) + T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, ((vi * 16) + vj))) @Ts.prim_func @@ -312,7 +312,7 @@ def opaque_access_split(A: T.Buffer((16, 16)), B: T.Buffer((16, 16))) -> None: vj = Ts.axis.S(16, j0 * 4 + j1) Ts.reads([]) Ts.writes([B[0:16, 0:16]]) - T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, ((vi * 16) + vj), dtype="handle")) + T.evaluate(T.tvm_fill_fragment(B.data, 16, 16, 16, 0, ((vi * 16) + vj))) @Ts.prim_func diff --git a/tests/python/s_tir/schedule/test_tir_schedule_state_cached_flags.py b/tests/python/s_tir/schedule/test_tir_schedule_state_cached_flags.py index 8e390883d1f5..04f6ce389ec5 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_state_cached_flags.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_state_cached_flags.py @@ -103,7 +103,7 @@ def loop_carried_dependency(A: T.Buffer((128,)), B: T.Buffer((128,)), C: T.Buffe B[vi] = A[vi] * 2.0 with Ts.sblock("C"): vi = Ts.axis.S(128, i) - C[vi] = T.if_then_else(vi >= 1, B[vi - 1] + 1.0, 0.0, dtype="float32") + C[vi] = T.if_then_else(vi >= 1, B[vi - 1] + 1.0, 0.0) @Ts.prim_func def concatenate_multi_producer(A: T.Buffer((128,)), B: T.Buffer((128,))) -> None: @@ -309,13 +309,13 @@ def non_perfect_tiling_cache(X: T.Buffer([224, 224], dtype='float32'), Y: T.Buff Y[h, w] = T.max( Y[h, w], T.if_then_else( - T.likely(1 <= h + kh, dtype="bool") - and T.likely(h + kh < 225, dtype="bool") - and T.likely(1 <= w + kw, dtype="bool") - and T.likely(w + kw < 225, dtype="bool"), + T.likely(1 <= h + kh) + and T.likely(h + kh < 225) + and T.likely(1 <= w + kw) + and T.likely(w + kw < 225), cache[h + kh - 1, w + kw - 1], 0.0, - dtype="float32", + ), ) @@ -346,13 +346,13 @@ def matmul_relu_padding(A: T.Buffer((127, 127), "float16"), B: T.Buffer((127, 12 v0, v1, v2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) Ts.reads(A[v0, v2]) Ts.writes(A_reindex[v0, v2]) - A_reindex[v0, v2] = T.if_then_else(v0 < 127 and v2 < 127, A[v0, v2], T.float16(0), dtype="float16") + A_reindex[v0, v2] = T.if_then_else(v0 < 127 and v2 < 127, A[v0, v2], T.float16(0)) for ax0, ax1, ax2 in T.grid(1, 128, 128): with Ts.sblock("B_reindex"): v0, v1, v2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) Ts.reads(B[v2, v1]) Ts.writes(B_reindex[v2, v1]) - B_reindex[v2, v1] = T.if_then_else(v2 < 127 and v1 < 127, B[v2, v1], T.float16(0), dtype="float16") + B_reindex[v2, v1] = T.if_then_else(v2 < 127 and v1 < 127, B[v2, v1], T.float16(0)) for ax0_0_0_ax1_0_0_fused in T.thread_binding(2, thread="blockIdx.y"): for ax0_0_1_ax1_0_1_fused in T.thread_binding(1, thread="blockIdx.x"): for ax0_0_2_ax1_0_2_fused in T.thread_binding(16, thread="threadIdx.y"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_tensorize.py b/tests/python/s_tir/schedule/test_tir_schedule_tensorize.py index b519742d5d37..37b703740cd1 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_tensorize.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_tensorize.py @@ -70,7 +70,7 @@ def mma_intrin(A: T.Buffer((16, 16), align=64, offset_factor=1), B: T.Buffer((16 B.elem_offset // 256, C.data, C.elem_offset // 256, - dtype="handle", + ) ) @@ -224,7 +224,7 @@ def tensorized_matmul(A: T.Buffer([128, 128], elem_offset=0, align=64, offset_fa T.floordiv(B_sub.elem_offset, 256), C_sub.data, T.floordiv(C_sub.elem_offset, 256), - dtype="handle", + ) ) @@ -295,7 +295,7 @@ def tensorized_batch_matmul_mma( T.floordiv(B_sub.elem_offset, 256), C_sub.data, T.floordiv(C_sub.elem_offset, 256), - dtype="handle", + ) ) @@ -447,7 +447,7 @@ def annotated_tensorized_matmul(A: T.Buffer([128, 128], elem_offset=0, align=64, T.floordiv(B_sub.elem_offset, 256), C_sub.data, T.floordiv(C_sub.elem_offset, 256), - dtype="handle", + ) ) @@ -769,7 +769,7 @@ def tensorized_matmul_int64_shape( T.floordiv(B_sub.elem_offset, T.int64(256)), C_sub.data, T.floordiv(C_sub.elem_offset, T.int64(256)), - dtype="handle", + ) ) # fmt: on diff --git a/tests/python/s_tir/schedule/test_tir_schedule_transform_layout.py b/tests/python/s_tir/schedule/test_tir_schedule_transform_layout.py index 5edb28ae189d..5060d47dd62e 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_transform_layout.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_transform_layout.py @@ -120,7 +120,7 @@ def conv2d_nhwc( ((((i1_1 >= 3) and (i1_1 < 227)) and (i2_1 >= 3)) and (i2_1 < 227)), Input[i0_1, (i1_1 - 3), (i2_1 - 3), i3_1], T.float32(0), - dtype="float32", + ) for i0, i1, i2, i3, i4, i5, i6 in T.grid(1, 112, 112, 64, 7, 7, 3): with Ts.sblock("conv2d_nhwc"): @@ -148,7 +148,7 @@ def conv2d_nhwc_transformed( i1_1 >= 3 and i1_1 < 227 and i2_1 >= 3 and i2_1 < 227, Input[i0_1, i1_1 - 3, i2_1 - 3, i3_1], T.float32(0), - dtype="float32", + ) for ax0, ax1, ax2 in T.grid(12544, 64, 147): with Ts.sblock("conv2d_nhwc"): @@ -711,9 +711,7 @@ def expected_func(A: T.Buffer(14, dtype)): for i, j in T.grid(4, 4): with Ts.sblock("block"): vi, vj = Ts.axis.remap("SS", [i, j]) - B[vi, vj] = T.if_then_else( - vi == 3 and 2 <= vj, pad_value_imm, A[vi * 4 + vj], dtype=dtype - ) + B[vi, vj] = T.if_then_else(vi == 3 and 2 <= vj, pad_value_imm, A[vi * 4 + vj]) Before = tvm.IRModule({"main": before_func}) Expected = tvm.IRModule({"main": expected_func}) @@ -794,9 +792,9 @@ def main(A: T.Buffer((14, 32), "int32")): with Ts.sblock("block"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) with Ts.init(): - B[vi, vj] = T.if_then_else(vi == 3 and 2 <= vj, 0, 0, dtype="int32") + B[vi, vj] = T.if_then_else(vi == 3 and 2 <= vj, 0, 0) B[vi, vj] = T.if_then_else( - vi == 3 and 2 <= vj, 0, B[vi, vj] + A[vi * 4 + vj, vk], dtype="int32" + vi == 3 and 2 <= vj, 0, B[vi, vj] + A[vi * 4 + vj, vk] ) sch = tvm.s_tir.Schedule(Before) @@ -830,12 +828,10 @@ class Expected: def main(A: T.Buffer((14, 32), "int32")): B = Ts.sblock_alloc_buffer([4, 4], "int32") for i, j in T.grid(4, 4): - B[i, j] = T.if_then_else(i == 3 and 2 <= j, 0, 0, dtype="int32") + B[i, j] = T.if_then_else(i == 3 and 2 <= j, 0, 0) for k in T.serial(32): with Ts.sblock("block"): - B[i, j] = T.if_then_else( - i == 3 and 2 <= j, 0, B[i, j] + A[i * 4 + j, k], dtype="int32" - ) + B[i, j] = T.if_then_else(i == 3 and 2 <= j, 0, B[i, j] + A[i * 4 + j, k]) sch = tvm.s_tir.Schedule(Before) sch.transform_layout( @@ -959,9 +955,7 @@ def main(A: T.Buffer(14, "int32")): for i, j in T.grid(4, 4): with Ts.sblock("block"): vi, vj = Ts.axis.remap("SS", [i, j]) - B[vi, vj] = T.if_then_else( - vi == 3 and 2 <= vj, vi + vj, A[vi * 4 + vj], dtype="int32" - ) + B[vi, vj] = T.if_then_else(vi == 3 and 2 <= vj, vi + vj, A[vi * 4 + vj]) sch = tvm.s_tir.Schedule(Before) sch.transform_layout( @@ -1086,7 +1080,6 @@ def main(A: T.Buffer(16, "int32"), n: T.int32): and 16 % n <= (vj + vi * n) % n, 0, A[vj + vi * n], - dtype="int32", ) sch = tvm.s_tir.Schedule(Before) diff --git a/tests/python/s_tir/script/test_s_tir_script_blocks.py b/tests/python/s_tir/script/test_s_tir_script_blocks.py index 59bd72ea7e74..5e5570c6c854 100644 --- a/tests/python/s_tir/script/test_s_tir_script_blocks.py +++ b/tests/python/s_tir/script/test_s_tir_script_blocks.py @@ -1029,7 +1029,7 @@ def different_access_indices( ] ) with Ts.init(): - B[vj, vi] = T.exp(B[vj, vi], dtype="float32") + B[vj, vi] = T.exp(B[vj, vi]) B[vi, vj] = B[vi, vj] + A[vi, vj, vk] 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 f0aa416d6cc3..3554a2ea61d1 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 @@ -98,7 +98,7 @@ def test_prim_func_symbolic_alloc_buffer_roundtrip(): tvm.ir.StringImm(buf.scope()), ], attrs=tvm.ir.DictAttrs({}), - ret_ty=buf.ty, + ty=buf.ty, ), ), tirx.Evaluate(tirx.BufferLoad(buf, [0])), @@ -780,16 +780,8 @@ def func( T.launch_thread(by, 4) T.launch_thread(ty, 4) T.launch_thread(tz, 2) - T.evaluate( - T.tvm_fill_fragment( - Conv_wmma_accumulator.data, 16, 16, 16, 0, T.float32(0), dtype="handle" - ) - ) - T.evaluate( - T.tvm_fill_fragment( - Conv_wmma_accumulator.data, 16, 16, 16, 7, T.float32(0), dtype="handle" - ) - ) + T.evaluate(T.tvm_fill_fragment(Conv_wmma_accumulator.data, 16, 16, 16, 0, T.float32(0))) + T.evaluate(T.tvm_fill_fragment(Conv_wmma_accumulator.data, 16, 16, 16, 7, T.float32(0))) for ic_outer in T.serial(0, 8): for kh in T.serial(0, 3): for ax2 in T.serial(0, 3): @@ -831,7 +823,6 @@ def func( ), ], T.float16(0), - dtype="float16", ) ) T.launch_thread(tx, 32) @@ -872,7 +863,6 @@ def func( ), ], T.float16(0), - dtype="float16", ) ) with T.launch_thread(tx, 32): @@ -927,11 +917,9 @@ def func( (((ty * 3072) + (kw * 512)) + (ic_inner * 256)), 256, 1, - dtype="handle", ), 16, "row_major", - dtype="handle", ) ) T.evaluate( @@ -947,11 +935,9 @@ def func( ((((ty * 3072) + (kw * 512)) + (ic_inner * 256)) + 1536), 256, 1, - dtype="handle", ), 16, "row_major", - dtype="handle", ) ) T.evaluate( @@ -967,11 +953,9 @@ def func( (((kw * 4096) + (ic_inner * 2048)) + (tz * 1024)), 256, 1, - dtype="handle", ), 16, "row_major", - dtype="handle", ) ) T.evaluate( @@ -987,11 +971,9 @@ def func( ((((kw * 4096) + (ic_inner * 2048)) + (tz * 1024)) + 768), 256, 1, - dtype="handle", ), 16, "row_major", - dtype="handle", ) ) T.evaluate( @@ -1004,7 +986,6 @@ def func( 0, Conv_wmma_accumulator.data, 0, - dtype="handle", ) ) T.evaluate( @@ -1017,7 +998,6 @@ def func( 3, Conv_wmma_accumulator.data, 7, - dtype="handle", ) ) T.evaluate( @@ -1036,11 +1016,9 @@ def func( ), 256, 2, - dtype="handle", ), 16, "row_major", - dtype="handle", ) ) T.evaluate( @@ -1062,11 +1040,9 @@ def func( ), 256, 2, - dtype="handle", ), 16, "row_major", - dtype="handle", ) ) @@ -1095,11 +1071,9 @@ def opt_conv_tensorcore_mod_host( } ) # body - stack_tcode_data: T.let[T.handle("int32")] = T.tvm_stack_alloca( - "arg_tcode", 10, dtype="handle" - ) + stack_tcode_data: T.let[T.handle("int32")] = T.tvm_stack_alloca("arg_tcode", 10) stack_tcode = T.decl_buffer([9], "int32", data=stack_tcode_data) - stack_value: T.let[T.handle] = T.tvm_stack_alloca("arg_value", 10, dtype="handle") + stack_value: T.let[T.handle] = T.tvm_stack_alloca("arg_value", 10) assert num_args == 3, "default_function: num_args should be 3" arg0: T.let[T.handle] = T.tvm_struct_get(args, 0, 12, dtype="handle") arg0_code: T.let[T.int32] = arg_type_ids[0] @@ -1141,7 +1115,7 @@ def opt_conv_tensorcore_mod_host( assert 14 == T.cast(arg0_shape[1], "int32"), ( "Argument arg0.shape[1] has an unsatisfied constraint" ) - if not (T.isnullptr(arg0_strides.data, dtype="bool")): + if not (T.isnullptr(arg0_strides.data)): assert ( ( ( @@ -1173,33 +1147,31 @@ def opt_conv_tensorcore_mod_host( assert dev_id == T.tvm_struct_get(arg2, 0, 9, dtype="int32"), ( "Argument arg2.device_id has an unsatisfied constraint" ) - T.evaluate(T.tvm_struct_set(stack_value, 0, 12, T.cast(2, "int64"), dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 0, 12, T.cast(2, "int64"))) stack_tcode[0] = 0 - T.evaluate(T.tvm_struct_set(stack_value, 1, 12, T.cast(dev_id, "int64"), dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 1, 12, T.cast(dev_id, "int64"))) stack_tcode[1] = 0 - T.evaluate(T.tvm_call_packed_lowered("__tvm_set_device", stack_value, 0, 2, dtype="int32")) + T.evaluate(T.tvm_call_packed_lowered("__tvm_set_device", stack_value, 0, 2)) T.attr(0, "compute_scope", "default_function_compute_") - T.evaluate(T.tvm_struct_set(stack_value, 0, 12, A, dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 0, 12, A)) stack_tcode[0] = 3 - T.evaluate(T.tvm_struct_set(stack_value, 1, 12, W, dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 1, 12, W)) stack_tcode[1] = 3 - T.evaluate(T.tvm_struct_set(stack_value, 2, 12, Conv, dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 2, 12, Conv)) stack_tcode[2] = 3 - T.evaluate(T.tvm_struct_set(stack_value, 3, 12, T.cast(196, "int64"), dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 3, 12, T.cast(196, "int64"))) stack_tcode[3] = 0 - T.evaluate(T.tvm_struct_set(stack_value, 4, 12, T.cast(2, "int64"), dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 4, 12, T.cast(2, "int64"))) stack_tcode[4] = 0 - T.evaluate(T.tvm_struct_set(stack_value, 5, 12, T.cast(4, "int64"), dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 5, 12, T.cast(4, "int64"))) stack_tcode[5] = 0 - T.evaluate(T.tvm_struct_set(stack_value, 6, 12, T.cast(4, "int64"), dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 6, 12, T.cast(4, "int64"))) stack_tcode[6] = 0 - T.evaluate(T.tvm_struct_set(stack_value, 7, 12, T.cast(2, "int64"), dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 7, 12, T.cast(2, "int64"))) stack_tcode[7] = 0 - T.evaluate(T.tvm_struct_set(stack_value, 8, 12, T.cast(32, "int64"), dtype="int32")) + T.evaluate(T.tvm_struct_set(stack_value, 8, 12, T.cast(32, "int64"))) stack_tcode[8] = 0 - T.evaluate( - T.tvm_call_packed_lowered("default_function_kernel0", stack_value, 0, 9, dtype="int32") - ) + T.evaluate(T.tvm_call_packed_lowered("default_function_kernel0", stack_value, 0, 9)) return opt_conv_tensorcore_mod_host @@ -1331,7 +1303,6 @@ def primfunc_with_allocate_annotations( ) ], T.uint8(0), - dtype="uint8", ), ) for ax0_ax1_fused_5 in T.serial(0, 56): @@ -1364,7 +1335,6 @@ def comm_reducer_single_reduce_group( True, reduce_temp0.data, threadIdx_x, - dtype="handle", ) ) @@ -1400,7 +1370,6 @@ def comm_reducer_multiple_reduce_groups( True, reduce_temp0.data, threadIdx_x, - dtype="handle", ) ) @@ -1429,7 +1398,6 @@ def multiple_commreducer() -> None: True, reduce_temp0.data, ax0_1, - dtype="handle", ) ) for ax0_1 in T.thread_binding(0, 32, thread="threadIdx.x"): @@ -1444,7 +1412,6 @@ def multiple_commreducer() -> None: True, reduce_temp1.data, ax0_1, - dtype="handle", ) ) @@ -1612,8 +1579,8 @@ def pointer_type(): def func_with_ptr_type_annotations(x: T.handle("int32"), y: T.handle("int32", "shared")): xx = T.alloc_buffer((16,), "int32") yy = T.alloc_buffer((16,), "int32", scope="shared") - a: T.let[T.handle("int32")] = T.address_of(xx[0], dtype="handle") - b: T.let[T.handle("int32", "shared")] = T.address_of(yy[0], dtype="handle") + a: T.let[T.handle("int32")] = T.address_of(xx[0]) + b: T.let[T.handle("int32", "shared")] = T.address_of(yy[0]) T.evaluate(T.call_extern("copy", a, b, dtype="")) return func_with_ptr_type_annotations @@ -1802,9 +1769,9 @@ def func( ax0, ax1, ax2 = Ts.axis.remap("SSS", [i0, i1, i2]) Ts.reads(placeholder[ax0, ax1, ax2]) Ts.writes(T_isinf[ax0, ax1, ax2]) - T_isinf[ax0, ax1, ax2] = T.fabs( - placeholder[ax0, ax1, ax2], dtype="float32" - ) == T.float32("inf") and not (T.isnan(placeholder[ax0, ax1, ax2], dtype="bool")) + T_isinf[ax0, ax1, ax2] = T.fabs(placeholder[ax0, ax1, ax2]) == T.float32( + "inf" + ) and not (T.isnan(placeholder[ax0, ax1, ax2])) return func @@ -1980,14 +1947,12 @@ def tir_packed_call(A: T.Buffer(16)): "tvm_test_cpacked", T.tvm_stack_make_array( A.data, - T.tvm_stack_make_shape(16, dtype="handle"), + T.tvm_stack_make_shape(16), T.reinterpret(T.uint64(0), dtype="handle"), T.uint32(1), T.Cast("float32", 0), 0, - dtype="handle", ), - dtype="int32", ) ) @@ -2369,7 +2334,6 @@ def lowered_loop_split( True, reduce_temp0.data, ki, - dtype="handle", ) ) with Ts.sblock("B_write_back"): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_compact_buffer_region.py b/tests/python/s_tir/transform/test_s_tir_transform_compact_buffer_region.py index 44fb4a5ddc8f..2fd48723f660 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_compact_buffer_region.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_compact_buffer_region.py @@ -435,7 +435,6 @@ def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((20, 20), "float32")) - 2 <= i and i < 18 and 2 <= j and j < 18, B[i - 2, j - 2], 0.0, - dtype="float32", ) @Ts.prim_func @@ -453,7 +452,6 @@ def expected( 2 <= i and i < 18 and 2 <= j and j < 18, B[i - 2, j - 2], 0.0, - dtype="float32", ) @@ -471,13 +469,12 @@ def before( Y[h, w] = T.max( Y[h, w], T.if_then_else( - T.likely(1 <= h + kh, dtype="bool") - and T.likely(h + kh < 225, dtype="bool") - and T.likely(1 <= w + kw, dtype="bool") - and T.likely(w + kw < 225, dtype="bool"), + T.likely(1 <= h + kh) + and T.likely(h + kh < 225) + and T.likely(1 <= w + kw) + and T.likely(w + kw < 225), cache[h + kh - 1, w + kw - 1], 0.0, - dtype="float32", ), ) @@ -492,13 +489,12 @@ def expected(X: T.Buffer((224, 224), "float32"), Y: T.Buffer((224, 224), "float3 Y[h, w] = T.max( Y[h, w], T.if_then_else( - T.likely(1 <= h + kh, dtype="bool") - and T.likely(h + kh < 225, dtype="bool") - and T.likely(1 <= w + kw, dtype="bool") - and T.likely(w + kw < 225, dtype="bool"), + T.likely(1 <= h + kh) + and T.likely(h + kh < 225) + and T.likely(1 <= w + kw) + and T.likely(w + kw < 225), cache[h + kh - 1, w + kw - 1], 0.0, - dtype="float32", ), ) @@ -779,17 +775,16 @@ def before(X: T.Buffer((64, 112, 112), "int32"), Y: T.Buffer((64, 56, 56), "int3 Y[h_o * 4 + h_i, w_o * 4 + w_i, c] = T.max( Y[h_o * 4 + h_i, w_o * 4 + w_i, c], T.if_then_else( - T.likely(1 <= (h_o * 4 + h_i) * 2 + kh, dtype="bool") - and T.likely((h_o * 4 + h_i) * 2 + kh < 113, dtype="bool") - and T.likely(1 <= (w_o * 4 + w_i) * 2 + kw, dtype="bool") - and T.likely((w_o * 4 + w_i) * 2 + kw < 113, dtype="bool"), + T.likely(1 <= (h_o * 4 + h_i) * 2 + kh) + and T.likely((h_o * 4 + h_i) * 2 + kh < 113) + and T.likely(1 <= (w_o * 4 + w_i) * 2 + kw) + and T.likely((w_o * 4 + w_i) * 2 + kw < 113), X_cache[ (h_o * 4 + h_i) * 2 + kh - 1, (w_o * 4 + w_i) * 2 + kw - 1, c, ], 0, - dtype="int32", ), ) @@ -831,15 +826,14 @@ def expected(X: T.Buffer((64, 112, 112), "int32"), Y: T.Buffer((64, 56, 56), "in Y[h_o * 4 + h_i, w_o * 4 + w_i, c] = T.max( Y[h_o * 4 + h_i, w_o * 4 + w_i, c], T.if_then_else( - T.likely(1 <= h_o * 8 + h_i * 2 + kh, dtype="bool") - and T.likely(1 <= w_o * 8 + w_i * 2 + kw, dtype="bool"), + T.likely(1 <= h_o * 8 + h_i * 2 + kh) + and T.likely(1 <= w_o * 8 + w_i * 2 + kw), X_cache[ h_o * 8 + h_i * 2 + kh - T.max(0, h_o * 8 - 1) - 1, w_o * 8 + w_i * 2 + kw - T.max(0, w_o * 8 - 1) - 1, c, ], 0, - dtype="int32", ), ) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_hoist_expression.py b/tests/python/s_tir/transform/test_s_tir_transform_hoist_expression.py index 400c24c0ad16..fe141edf1a4f 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_hoist_expression.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_hoist_expression.py @@ -478,7 +478,7 @@ def test_hoist_if_else_expr(): @Ts.prim_func(private=True) def before(A: T.Buffer((4, 4), "float32")): for i, j in T.grid(4, 4): - A[i, j] = T.if_then_else(i < 2, 1.0, 2.0, dtype="float32") + A[i, j] = T.if_then_else(i < 2, 1.0, 2.0) @Ts.prim_func(private=True) def expected(A: T.Buffer((4, 4), "float32")): @@ -498,7 +498,7 @@ def test_suppress_hoist_if_else_expr(): @Ts.prim_func(private=True) def before(A: T.Buffer((4, 4), "float32")): for i, j in T.grid(4, 4): - A[i, j] = T.if_then_else(i < 2, 1.0, 2.0, dtype="float32") + A[i, j] = T.if_then_else(i < 2, 1.0, 2.0) expected = before diff --git a/tests/python/s_tir/transform/test_s_tir_transform_inject_ptx_async_copy.py b/tests/python/s_tir/transform/test_s_tir_transform_inject_ptx_async_copy.py index 485154788e2d..af1ebe08ac10 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_inject_ptx_async_copy.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_inject_ptx_async_copy.py @@ -401,13 +401,13 @@ def simple_compute( Ts.reads(A[tx, i]) Ts.writes(A_shared[tx, 0]) A_shared[tx, 0] = T.if_then_else( - 1 <= i and i < 15, A[tx, i - 1], T.float32(0), dtype="float32" + 1 <= i and i < 15, A[tx, i - 1], T.float32(0) ) with Ts.sblock(): Ts.reads(B[tx, i]) Ts.writes(B_shared[tx, 0]) B_shared[tx, 0] = T.if_then_else( - 1 <= i and i < 15, B[tx, i - 1], T.float32(0), dtype="float32" + 1 <= i and i < 15, B[tx, i - 1], T.float32(0) ) with Ts.sblock(): Ts.reads(A_shared[tx, 0], B_shared[tx, 0]) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_inject_software_pipeline.py b/tests/python/s_tir/transform/test_s_tir_transform_inject_software_pipeline.py index 33e3d268ea62..6e093cc430f4 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_inject_software_pipeline.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_inject_software_pipeline.py @@ -1489,7 +1489,7 @@ def ref(A: T.Buffer((16, 16), "float32"), D: T.Buffer((16, 16), "float32")) -> N with T.attr( 0, "async_wait_inflight_count", - T.if_then_else(i + 16 - 1 < 16, 1, 0, dtype="int32"), + T.if_then_else(i + 16 - 1 < 16, 1, 0), ): D[tx, i - 2 + 16] = C[(i - 2 + 16) % 2, tx, 0] + T.float32(1) 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 a2aaab25383e..c375eb7c1ddc 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 @@ -93,7 +93,6 @@ def lowered_loop_split( True, reduce_temp0[0], ki, - dtype="handle", ) ) with Ts.sblock("B_write_back"): @@ -136,11 +135,7 @@ def lowered_no_normal_reduction( "reduce_scope", T.int32(0), ) - T.evaluate( - T.tvm_thread_allreduce( - T.uint32(1), A[vi, vk], True, reduce_temp0[0], k, dtype="handle" - ) - ) + 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) @@ -187,7 +182,7 @@ def lowered_two_bound_loops( ) T.evaluate( T.tvm_thread_allreduce( - T.uint32(1), A[vi, vk], True, reduce_temp0[0], ko, ki, dtype="handle" + T.uint32(1), A[vi, vk], True, reduce_temp0[0], ko, ki ) ) with Ts.sblock("B_write_back"): @@ -269,7 +264,6 @@ def lowered_multiple_blocks_under_reduction_loop( True, reduce_temp0[0], k0o, - dtype="handle", ) ) with Ts.sblock("B_write_back"): @@ -332,7 +326,6 @@ def lowered_with_block_predicate( True, reduce_temp0[0], ki, - dtype="handle", ) ) with Ts.sblock("B_write_back"): @@ -374,7 +367,7 @@ def single_reduction_loop_with_block_predicate( with Ts.init(): T_softmax_expsum_shared[i0_2] = T.float32(0) T_softmax_expsum_shared[i0_2] = T_softmax_expsum_shared[i0_2] + T.exp( - A[i0_2, k] - T_softmax_maxelem_shared[i0_2], dtype="float32" + A[i0_2, k] - T_softmax_maxelem_shared[i0_2] ) for i1_0 in T.serial(1): for i1_1 in T.thread_binding(512, thread="threadIdx.x"): @@ -388,7 +381,7 @@ def single_reduction_loop_with_block_predicate( Ts.writes(T_softmax_norm[i0_3, i1]) Ts.sblock_attr({"axis": 1}) T_softmax_norm[i0_3, i1] = ( - T.exp(A[i0_3, i1] - T_softmax_maxelem_shared[i0_3], dtype="float32") + T.exp(A[i0_3, i1] - T_softmax_maxelem_shared[i0_3]) / T_softmax_expsum_shared[i0_3] ) @@ -435,7 +428,6 @@ def lowered_single_reduction_loop_with_block_predicate( True, cross_thread_0[0], ax1_1, - dtype="handle", ) ) with Ts.sblock("T_softmax_maxelem_write_back"): @@ -458,7 +450,7 @@ def lowered_single_reduction_loop_with_block_predicate( Ts.reads(A[i0_3, k], T_softmax_maxelem_shared[i0_3]) Ts.writes(in_thread_1[0]) in_thread_1[0] = in_thread_1[0] + T.exp( - A[i0_3, k] - T_softmax_maxelem_shared[i0_3], dtype="float32" + A[i0_3, k] - T_softmax_maxelem_shared[i0_3] ) with Ts.sblock("T_softmax_expsum_cross_thread"): Ts.reads(in_thread_1[0]) @@ -475,7 +467,6 @@ def lowered_single_reduction_loop_with_block_predicate( True, cross_thread_1[0], ax1_1, - dtype="handle", ) ) with Ts.sblock("T_softmax_expsum_write_back"): @@ -496,7 +487,7 @@ def lowered_single_reduction_loop_with_block_predicate( Ts.writes(T_softmax_norm[i0_5, i1]) Ts.sblock_attr({"axis": 1}) T_softmax_norm[i0_5, i1] = ( - T.exp(A[i0_5, i1] - T_softmax_maxelem_shared[i0_5], dtype="float32") + T.exp(A[i0_5, i1] - T_softmax_maxelem_shared[i0_5]) / T_softmax_expsum_shared[i0_5] ) @@ -927,11 +918,7 @@ def lowered_reducer_max( "reduce_scope", T.int32(0), ) - T.evaluate( - T.tvm_thread_allreduce( - T.uint32(1), A[vi, vk], True, reduce_temp0[0], k, dtype="handle" - ) - ) + 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) @@ -968,9 +955,7 @@ def lowered_zero_rank_buffer( "reduce_scope", T.int32(0), ) - T.evaluate( - T.tvm_thread_allreduce(T.uint32(1), A[vk], True, reduce_temp0[0], k, dtype="handle") - ) + 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[()]]) @@ -1096,7 +1081,7 @@ def softmax( with Ts.init(): T_softmax_expsum_shared[i0_2] = T.float32(0) T_softmax_expsum_shared[i0_2] = T_softmax_expsum_shared[i0_2] + T.exp( - A[i0_2, k] - T_softmax_maxelem_shared[i0_2], dtype="float32" + A[i0_2, k] - T_softmax_maxelem_shared[i0_2] ) for i1_0 in T.serial(0, 8): for i1_1 in T.thread_binding(0, 32, thread="threadIdx.x"): @@ -1115,7 +1100,6 @@ def softmax( T_softmax_norm[i0_3, i1] = ( T.exp( A[i0_3, i1] - T_softmax_maxelem_shared[i0_3], - dtype="float32", ) / T_softmax_expsum_shared[i0_3] ) @@ -1159,7 +1143,6 @@ def lowered_softmax( True, reduce_temp0[0], ax0_1, - dtype="handle", ) ) with Ts.sblock("T_softmax_maxelem_write_back"): @@ -1185,7 +1168,7 @@ def lowered_softmax( ) Ts.writes([normal_reduce_temp1[0]]) normal_reduce_temp1[0] = normal_reduce_temp1[0] + T.exp( - A[i0_3, k] - T_softmax_maxelem_shared[i0_3], dtype="float32" + A[i0_3, k] - T_softmax_maxelem_shared[i0_3] ) with Ts.sblock("T_softmax_expsum_cross_thread_reduction"): Ts.reads([normal_reduce_temp1[0]]) @@ -1202,7 +1185,6 @@ def lowered_softmax( True, reduce_temp1[0], ax0_1, - dtype="handle", ) ) with Ts.sblock("T_softmax_expsum_write_back"): @@ -1228,7 +1210,6 @@ def lowered_softmax( T_softmax_norm[i0_5, i1] = ( T.exp( A[i0_5, i1] - T_softmax_maxelem_shared[i0_5], - dtype="float32", ) / T_softmax_expsum_shared[i0_5] ) @@ -1318,7 +1299,6 @@ def lowered_argmax_split( cross_thread_argmax_v0[0], cross_thread_argmax_v1[0], i1_1, - dtype="handle", ) ) with Ts.sblock("argmax_write_back"): @@ -1414,7 +1394,6 @@ def lowered_argmin_split_init_update_reordered( cross_thread_argmin_v0[0], cross_thread_argmin_v1[0], i1_1, - dtype="handle", ) ) with Ts.sblock("argmin_write_back"): @@ -1473,7 +1452,6 @@ def layer_norm_tuple_sum( * T.float32(0.0013020833333333333) * (data_red_temp_v0[ax0] * T.float32(0.0013020833333333333)) + T.float32(1.0000000000000001e-05), - dtype="float32", ) * gamma[ax1] + bias[ax1] @@ -1539,7 +1517,6 @@ def lowered_layer_norm_tuple_sum( cross_thread_data_red_temp_v0[0], cross_thread_data_red_temp_v1[0], i1_1, - dtype="handle", ) ) with Ts.sblock("data_red_temp_write_back"): @@ -1570,7 +1547,6 @@ def lowered_layer_norm_tuple_sum( * T.float32(0.0013020833333333333) * (data_red_temp_v0[ax0] * T.float32(0.0013020833333333333)) + T.float32(1.0000000000000001e-05), - dtype="float32", ) * gamma[ax1] + bias[ax1] diff --git a/tests/python/s_tir/transform/test_s_tir_transform_memhammer_lower_auto_copy.py b/tests/python/s_tir/transform/test_s_tir_transform_memhammer_lower_auto_copy.py index eb3443bf80e3..52f66e62b3be 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_memhammer_lower_auto_copy.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_memhammer_lower_auto_copy.py @@ -486,11 +486,9 @@ def main() -> None: src.elem_offset, s1 * 16, 1, - dtype="handle", ), s1, "row_major", - dtype="handle", ) ) @@ -569,11 +567,9 @@ def main() -> None: tgt.elem_offset, s1 * 16, 2, - dtype="handle", ), s1, "row_major", - dtype="handle", ) ) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_plan_update_buffer_allocation_location.py b/tests/python/s_tir/transform/test_s_tir_transform_plan_update_buffer_allocation_location.py index 5f0c5edcba48..66650165ba81 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_plan_update_buffer_allocation_location.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_plan_update_buffer_allocation_location.py @@ -253,7 +253,7 @@ def before(A: T.Buffer((8, 8, 8), "int32"), B: T.Buffer((8, 8, 8), "int32")): vj = Ts.axis.opaque(8, j) B[vi, vj, vk] = ( C[vi, vj, vk] - + T.if_then_else(0 < vj, C[vi, vj - 1, vk], 0, dtype="int32") + + T.if_then_else(0 < vj, C[vi, vj - 1, vk], 0) + D[vi, vj, vk] ) @@ -280,7 +280,7 @@ def after(A: T.Buffer((8, 8, 8), "int32"), B: T.Buffer((8, 8, 8), "int32")) -> N vj = Ts.axis.opaque(8, j) B[vi, vj, vk] = ( C[vi, vj, vk] - + T.if_then_else(0 < vj, C[vi, vj - 1, vk], 0, dtype="int32") + + T.if_then_else(0 < vj, C[vi, vj - 1, vk], 0) + D[vi, vj, vk] ) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_remove_undef.py b/tests/python/s_tir/transform/test_s_tir_transform_remove_undef.py index 0d4e5793a37d..20682648a6f2 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_remove_undef.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_remove_undef.py @@ -31,7 +31,7 @@ def test_remove_store_undef(): class Before: @Ts.prim_func def main(A: T.Buffer(1, "int32")): - A[0] = T.undef(dtype="int32") + A[0] = T.undef() @I.ir_module class Expected: @@ -50,7 +50,7 @@ def test_remove_store_undef_expression(): class Before: @Ts.prim_func def main(A: T.Buffer(1, "int32")): - A[0] = 1 + T.undef(dtype="int32") + A[0] = 1 + T.undef() @I.ir_module class Expected: @@ -69,7 +69,7 @@ def test_keep_other_call_nodes(): class Before: @Ts.prim_func def main(A: T.Buffer(1, "int32"), n: T.int32): - A[0] = T.shift_left(n, 1, dtype="int32") + A[0] = T.shift_left(n, 1) Expected = Before @@ -84,7 +84,7 @@ def test_remove_let_undef(): class Before: @Ts.prim_func def main(A: T.Buffer(1, "int32")): - val: T.let[T.int32] = T.undef(dtype="int32") + val: T.let[T.int32] = T.undef() A[0] = val @I.ir_module @@ -104,7 +104,7 @@ def test_raise_error_for_undef_as_store_indices(): class Before: @Ts.prim_func def main(A: T.Buffer(1, "int32")): - val: T.let[T.int32] = T.undef(dtype="int32") + val: T.let[T.int32] = T.undef() A[val] = 5 with pytest.raises(RuntimeError): @@ -122,7 +122,7 @@ def test_raise_error_for_undef_as_load_indices(): class Before: @Ts.prim_func def main(A: T.Buffer(1, "int32"), B: T.Buffer(1, "int32")): - B[0] = A[T.undef(dtype="int32")] + B[0] = A[T.undef()] with pytest.raises(RuntimeError): tvm.s_tir.transform.RemoveStoreUndef()(Before) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_renormalize_split_pattern.py b/tests/python/s_tir/transform/test_s_tir_transform_renormalize_split_pattern.py index 3d0ff2d7f740..cd98b4076bf8 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_renormalize_split_pattern.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_renormalize_split_pattern.py @@ -47,11 +47,11 @@ def main(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 51 conv2d_transpose_nhwc_local[i1_4_init * 4 + i2_3_init * 2 + i2_4_init] = T.float32(0) for i6_0 in T.serial(16): for ax0_ax1_ax2_ax3_fused_0 in T.serial(24): - PadInput_shared[ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x] = T.if_then_else(128 <= ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x and ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x < 640 and 1 <= blockIdx_x // 32 * 2 + (ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x) % 128 // 32 and blockIdx_x // 32 * 2 + (ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x) % 128 // 32 < 5, inputs_flat[blockIdx_x // 32 * 1024 + ax0_ax1_ax2_ax3_fused_0 * 512 + i6_0 * 32 + threadIdx_x - 2560], T.float32(0), dtype="float32") + PadInput_shared[ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x] = T.if_then_else(128 <= ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x and ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x < 640 and 1 <= blockIdx_x // 32 * 2 + (ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x) % 128 // 32 and blockIdx_x // 32 * 2 + (ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x) % 128 // 32 < 5, inputs_flat[blockIdx_x // 32 * 1024 + ax0_ax1_ax2_ax3_fused_0 * 512 + i6_0 * 32 + threadIdx_x - 2560], T.float32(0)) for ax0_ax1_ax2_ax3_fused_0 in T.serial(32): weight_shared[T.ramp(ax0_ax1_ax2_ax3_fused_0 * 128 + threadIdx_x * 4, 1, 4)] = weight_flat[T.ramp((ax0_ax1_ax2_ax3_fused_0 * 128 + threadIdx_x * 4) // 256 * 131072 + i6_0 * 8192 + (ax0_ax1_ax2_ax3_fused_0 * 128 + threadIdx_x * 4) % 256 // 8 * 256 + blockIdx_x % 32 * 8 + threadIdx_x % 2 * 4, 1, 4)] for i6_1, i2_3, i4_2, i5_2, i6_2, i1_4, i2_4 in T.grid(4, 2, 4, 4, 8, 2, 2): - conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] = conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] + T.if_then_else((i1_4 + i4_2) % 2 == 0 and (i2_4 + i5_2) % 2 == 0, PadInput_shared[threadIdx_x // 8 * 128 + (i1_4 + i4_2) // 2 * 128 + (i2_4 + i5_2) // 2 * 32 + i2_3 * 32 + i6_1 * 8 + i6_2], T.float32(0), dtype="float32") * weight_shared[i6_1 * 64 + i6_2 * 8 + threadIdx_x % 8 + 3840 - i5_2 * 256 - i4_2 * 1024] + conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] = conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] + T.if_then_else((i1_4 + i4_2) % 2 == 0 and (i2_4 + i5_2) % 2 == 0, PadInput_shared[threadIdx_x // 8 * 128 + (i1_4 + i4_2) // 2 * 128 + (i2_4 + i5_2) // 2 * 32 + i2_3 * 32 + i6_1 * 8 + i6_2], T.float32(0)) * weight_shared[i6_1 * 64 + i6_2 * 8 + threadIdx_x % 8 + 3840 - i5_2 * 256 - i4_2 * 1024] for ax1, ax2 in T.grid(2, 4): conv2d_transpose_nhwc_flat[threadIdx_x // 8 * 4096 + ax1 * 2048 + blockIdx_x // 32 * 1024 + ax2 * 256 + blockIdx_x % 32 * 8 + threadIdx_x % 8] = conv2d_transpose_nhwc_local[ax1 * 4 + ax2] @@ -78,11 +78,11 @@ def main(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 51 conv2d_transpose_nhwc_local[i1_4_init * 4 + i2_3_init * 2 + i2_4_init] = T.float32(0) for i6_0 in T.serial(16): for ax0_ax1_ax2_ax3_fused_0 in T.serial(24): - PadInput_shared[ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x] = T.if_then_else(1 <= (ax0_ax1_ax2_ax3_fused_0 + threadIdx_x // 32) // 4 and (ax0_ax1_ax2_ax3_fused_0 + threadIdx_x // 32) // 20 < 1 and 1 <= blockIdx_x // 32 * 2 + (ax0_ax1_ax2_ax3_fused_0 + threadIdx_x // 32) % 4 and (blockIdx_x // 32 * 2 + (ax0_ax1_ax2_ax3_fused_0 + threadIdx_x // 32) % 4) // 5 < 1, inputs_flat[blockIdx_x // 32 * 1024 + ax0_ax1_ax2_ax3_fused_0 * 512 + i6_0 * 32 + threadIdx_x - 2560], T.float32(0), dtype="float32") + PadInput_shared[ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x] = T.if_then_else(1 <= (ax0_ax1_ax2_ax3_fused_0 + threadIdx_x // 32) // 4 and (ax0_ax1_ax2_ax3_fused_0 + threadIdx_x // 32) // 20 < 1 and 1 <= blockIdx_x // 32 * 2 + (ax0_ax1_ax2_ax3_fused_0 + threadIdx_x // 32) % 4 and (blockIdx_x // 32 * 2 + (ax0_ax1_ax2_ax3_fused_0 + threadIdx_x // 32) % 4) // 5 < 1, inputs_flat[blockIdx_x // 32 * 1024 + ax0_ax1_ax2_ax3_fused_0 * 512 + i6_0 * 32 + threadIdx_x - 2560], T.float32(0)) for ax0_ax1_ax2_ax3_fused_0 in T.serial(32): weight_shared[T.ramp(ax0_ax1_ax2_ax3_fused_0 * 128 + threadIdx_x * 4, 1, 4)] = weight_flat[T.ramp((ax0_ax1_ax2_ax3_fused_0 + threadIdx_x * 4 // 128) // 2 * 131072 + i6_0 * 8192 + (ax0_ax1_ax2_ax3_fused_0 * 16 + threadIdx_x * 4 // 8) % 32 * 256 + blockIdx_x % 32 * 8 + threadIdx_x % 2 * 4, 1, 4)] for i6_1, i2_3, i4_2, i5_2, i6_2, i1_4, i2_4 in T.grid(4, 2, 4, 4, 8, 2, 2): - conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] = conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] + T.if_then_else((i1_4 + i4_2) % 2 == 0 and (i2_4 + i5_2) % 2 == 0, PadInput_shared[threadIdx_x // 8 * 128 + (i1_4 + i4_2) // 2 * 128 + (i2_4 + i5_2) // 2 * 32 + i2_3 * 32 + i6_1 * 8 + i6_2], T.float32(0), dtype="float32") * weight_shared[i6_1 * 64 + i6_2 * 8 + threadIdx_x % 8 + 3840 - i5_2 * 256 - i4_2 * 1024] + conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] = conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] + T.if_then_else((i1_4 + i4_2) % 2 == 0 and (i2_4 + i5_2) % 2 == 0, PadInput_shared[threadIdx_x // 8 * 128 + (i1_4 + i4_2) // 2 * 128 + (i2_4 + i5_2) // 2 * 32 + i2_3 * 32 + i6_1 * 8 + i6_2], T.float32(0)) * weight_shared[i6_1 * 64 + i6_2 * 8 + threadIdx_x % 8 + 3840 - i5_2 * 256 - i4_2 * 1024] for ax1, ax2 in T.grid(2, 4): conv2d_transpose_nhwc_flat[threadIdx_x // 8 * 4096 + ax1 * 2048 + blockIdx_x // 32 * 1024 + ax2 * 256 + blockIdx_x % 32 * 8 + threadIdx_x % 8] = conv2d_transpose_nhwc_local[ax1 * 4 + ax2] @@ -109,11 +109,11 @@ def main(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 51 conv2d_transpose_nhwc_local[i1_4_init * 4 + i2_3_init * 2 + i2_4_init] = T.float32(0) for i6_0 in T.serial(16): for ax0_ax1_ax2_ax3_fused_0 in T.serial(24): - PadInput_shared[ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x] = T.if_then_else(4 <= ax0_ax1_ax2_ax3_fused_0 and ax0_ax1_ax2_ax3_fused_0 < 20 and 1 <= blockIdx_x // 32 * 2 + ax0_ax1_ax2_ax3_fused_0 % 4 and blockIdx_x // 32 * 2 + ax0_ax1_ax2_ax3_fused_0 % 4 < 5, inputs_flat[blockIdx_x // 32 * 1024 + ax0_ax1_ax2_ax3_fused_0 * 512 + i6_0 * 32 + threadIdx_x - 2560], T.float32(0), dtype="float32") + PadInput_shared[ax0_ax1_ax2_ax3_fused_0 * 32 + threadIdx_x] = T.if_then_else(4 <= ax0_ax1_ax2_ax3_fused_0 and ax0_ax1_ax2_ax3_fused_0 < 20 and 1 <= blockIdx_x // 32 * 2 + ax0_ax1_ax2_ax3_fused_0 % 4 and blockIdx_x // 32 * 2 + ax0_ax1_ax2_ax3_fused_0 % 4 < 5, inputs_flat[blockIdx_x // 32 * 1024 + ax0_ax1_ax2_ax3_fused_0 * 512 + i6_0 * 32 + threadIdx_x - 2560], T.float32(0)) for ax0_ax1_ax2_ax3_fused_0 in T.serial(32): weight_shared[T.ramp(ax0_ax1_ax2_ax3_fused_0 * 128 + threadIdx_x * 4, 1, 4)] = weight_flat[T.ramp(ax0_ax1_ax2_ax3_fused_0 // 2 * 131072 + i6_0 * 8192 + ax0_ax1_ax2_ax3_fused_0 % 2 * 4096 + threadIdx_x // 2 * 256 + blockIdx_x % 32 * 8 + threadIdx_x % 2 * 4, 1, 4)] for i6_1, i2_3, i4_2, i5_2, i6_2, i1_4, i2_4 in T.grid(4, 2, 4, 4, 8, 2, 2): - conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] = conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] + T.if_then_else((i1_4 + i4_2) % 2 == 0 and (i2_4 + i5_2) % 2 == 0, PadInput_shared[threadIdx_x // 8 * 128 + (i1_4 + i4_2) // 2 * 128 + (i2_4 + i5_2) // 2 * 32 + i2_3 * 32 + i6_1 * 8 + i6_2], T.float32(0), dtype="float32") * weight_shared[i6_1 * 64 + i6_2 * 8 + threadIdx_x % 8 + 3840 - i5_2 * 256 - i4_2 * 1024] + conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] = conv2d_transpose_nhwc_local[i1_4 * 4 + i2_3 * 2 + i2_4] + T.if_then_else((i1_4 + i4_2) % 2 == 0 and (i2_4 + i5_2) % 2 == 0, PadInput_shared[threadIdx_x // 8 * 128 + (i1_4 + i4_2) // 2 * 128 + (i2_4 + i5_2) // 2 * 32 + i2_3 * 32 + i6_1 * 8 + i6_2], T.float32(0)) * weight_shared[i6_1 * 64 + i6_2 * 8 + threadIdx_x % 8 + 3840 - i5_2 * 256 - i4_2 * 1024] for ax1, ax2 in T.grid(2, 4): conv2d_transpose_nhwc_flat[threadIdx_x // 8 * 4096 + ax1 * 2048 + blockIdx_x // 32 * 1024 + ax2 * 256 + blockIdx_x % 32 * 8 + threadIdx_x % 8] = conv2d_transpose_nhwc_local[ax1 * 4 + ax2] diff --git a/tests/python/script/test_constructor_contracts.py b/tests/python/script/test_constructor_contracts.py new file mode 100644 index 000000000000..33dbf12fffc4 --- /dev/null +++ b/tests/python/script/test_constructor_contracts.py @@ -0,0 +1,114 @@ +# 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. +"""Normal IR constructors retain their contracts in TVMScript.""" + +import pytest + +import tvm +import tvm.testing +from tvm import ir, relax, tirx +from tvm.script import ir as I +from tvm.script import relax as R +from tvm.script import tirx as T + + +def test_call_type_and_validation_contract(): + x = ir.Var("x", "float32") + span = ir.Span(ir.SourceName("constructor"), 1, 1, 1, 5) + for constructor in (ir.Call, I.Call): + call = constructor("tirx.exp", [x], span=span) + assert isinstance(call, I.Call) + assert call.ty.is_missing() + assert call.span.same_as(span) + ir.assert_structural_equal( + constructor("tirx.exp", [x], ty="float32"), + ir.Call("tirx.exp", [x], ty=ir.PrimType("float32")), + ) + with pytest.raises(TypeError): + constructor("tirx.exp", [], ty="float32") + provisional = constructor.unchecked("tirx.exp", [], ty="float32", span=span) + assert not provisional.args + assert provisional.span.same_as(span) + assert provisional.ty.dtype == "float32" + + +@pytest.mark.parametrize("ty", [None, "float32", "handle"]) +def test_raw_call_preserves_provisional_fields(ty): + # Invalid arguments and type arguments must survive the raw representation. + call = ir.Call.unchecked("tirx.exp", [], attrs={"tag": 7}, ty_args=[ir.StringType()], ty=ty) + source = call.script() + assert "I.Call.unchecked(" in source + restored = eval(source, {"I": I, "R": R, "T": T}) + ir.assert_structural_equal(call, restored) + + +def test_raw_call_preserves_checked_fields(): + call = ir.Call("tirx.exp", [tirx.FloatImm("float32", 0)], attrs={"tag": 7}, ty="float32") + source = call.script() + assert "I.Call.unchecked(" in source + ir.assert_structural_equal(call, eval(source, {"I": I, "R": R, "T": T})) + + +def test_range_and_value_constructor_parameters(): + span = ir.Span(ir.SourceName("constructor"), 1, 1, 1, 5) + for args in [(5,), (2, 5)]: + actual = T.Range(*args, span=span) + ir.assert_structural_equal(actual, ir.Range(*args, span=span)) + assert actual.span.same_as(span) + ir.assert_structural_equal(T.Range.from_min_extent(2, 3), ir.Range(2, 5)) + for actual, expected in [ + (R.shape([2, 3], span=span), relax.ShapeExpr([2, 3], span=span)), + (R.str("value", span=span), ir.StringImm("value", span=span)), + ]: + ir.assert_structural_equal(actual, expected) + assert actual.span.same_as(span) + ir.assert_structural_equal(R.prim_value(3, dtype="int32"), relax.prim_value(3, dtype="int32")) + + +def test_operation_dtype_keywords_match_normal_constructors(): + for name in ("min_value", "max_value", "infinity"): + ir.assert_structural_equal( + getattr(T, name)(dtype="float32"), getattr(tirx, name)(dtype="float32") + ) + x = ir.Var("x", "float32") + for operation in (T.exp, tirx.exp): + with pytest.raises(TypeError): + operation(x, dtype="float32") + vector = ir.Var("v", "int8x4") + ir.assert_structural_equal(T.dp4a(vector, vector), tirx.dp4a(vector, vector)) + + +def test_normal_constructors_in_parsed_function(): + @T.prim_func + def function(x: ir.PrimType("float32")): + T.evaluate(ir.Call("tirx.exp", [x], ty="float32")) + + call = function.body.value + ir.assert_structural_equal( + call, ir.Call("tirx.exp", [function.params[0]], ty=ir.PrimType("float32")) + ) + assert call.span is not None + + @R.function + def identity(x: relax.TensorType([2], "float32")) -> relax.TensorType([2], "float32"): + return x + + ir.assert_structural_equal(identity.params[0].ty, relax.TensorType([2], "float32")) + + +if __name__ == "__main__": + tvm.testing.main() diff --git a/tests/python/script/test_source_locations.py b/tests/python/script/test_source_locations.py index 788ce5606a41..b93f62a7d1cb 100644 --- a/tests/python/script/test_source_locations.py +++ b/tests/python/script/test_source_locations.py @@ -111,7 +111,7 @@ def main(): def calls(spanned_language): M = spanned_language.M M.evaluate = lambda value: value - M.call_extern = lambda dtype, name: ir.Call(ir.GlobalVar(name), [], ret_ty=dtype) + M.call_extern = lambda dtype, name: ir.Call(ir.GlobalVar(name), [], ty=dtype) return M @@ -193,7 +193,7 @@ def gallery(monkeypatch, spanned_language): monkeypatch.setitem(globals(), "__name__", "__main__") M = spanned_language.M M.store = lambda value: ir.Call( - ir.GlobalVar("store"), [ir.prim.IntImm("int32", value)], ret_ty="int32" + ir.GlobalVar("store"), [ir.prim.IntImm("int32", value)], ty="int32" ) return M diff --git a/tests/python/te/test_te_create_primfunc.py b/tests/python/te/test_te_create_primfunc.py index dd84a630f0c9..a703736b0029 100644 --- a/tests/python/te/test_te_create_primfunc.py +++ b/tests/python/te/test_te_create_primfunc.py @@ -178,7 +178,6 @@ def tir_conv2d( 1 <= yy and yy < 15 and 1 <= xx and xx < 15, A[nn, cc, yy - 1, xx - 1], 0.0, - dtype="float32", ) for n, f, y, x, kc, ky, kx in T.grid(16, 32, 14, 14, 16, 3, 3): with Ts.sblock("B"): @@ -260,34 +259,30 @@ def tir_extern( "tvm.contrib.cblas.matmul", T.tvm_stack_make_array( A.data, - T.tvm_stack_make_shape(128, 128, dtype="handle"), + T.tvm_stack_make_shape(128, 128), 0, 2, 0.0, off1, - dtype="handle", ), T.tvm_stack_make_array( B.data, - T.tvm_stack_make_shape(128, 128, dtype="handle"), + T.tvm_stack_make_shape(128, 128), 0, 2, 0.0, off2, - dtype="handle", ), T.tvm_stack_make_array( C.data, - T.tvm_stack_make_shape(128, 128, dtype="handle"), + T.tvm_stack_make_shape(128, 128), 0, 2, 0.0, off3, - dtype="handle", ), 0, 0, - dtype="int32", ) ) @@ -762,7 +757,6 @@ def tir_resize2d_symbolic( "int64", T.round( T.float32(128) / T.Cast("float32", oh) * T.Cast("float32", v_i2), - dtype="float32", ), ), T.int64(127), @@ -775,7 +769,6 @@ def tir_resize2d_symbolic( "int64", T.round( T.float32(128) / T.Cast("float32", ow) * T.Cast("float32", v_i3), - dtype="float32", ), ), T.int64(127), diff --git a/tests/python/tirx-analysis/test_tir_analysis_undefined_vars.py b/tests/python/tirx-analysis/test_tir_analysis_undefined_vars.py index da204b70ef75..82343213a886 100644 --- a/tests/python/tirx-analysis/test_tir_analysis_undefined_vars.py +++ b/tests/python/tirx-analysis/test_tir_analysis_undefined_vars.py @@ -45,7 +45,7 @@ def test_decl_buffer_data_is_use(): tvm.ir.DataTypeImm(tvm.DataType(buf.dtype)), tvm.ir.StringImm(buf.scope()), ], - ret_ty=buf.ty, + ty=buf.ty, ), ) stmt = tirx.SeqStmt([decl, body]) @@ -80,7 +80,7 @@ def test_decl_buffer_elem_offset_is_use(): tvm.ir.DataTypeImm(tvm.DataType(buf.dtype)), tvm.ir.StringImm(buf.scope()), ], - ret_ty=buf.ty, + ty=buf.ty, ), ) stmt = tirx.SeqStmt([decl, body]) @@ -113,7 +113,7 @@ def test_alloc_buffer_data_is_def(): tvm.ir.StringImm(buf.scope()), ], attrs=tvm.ir.DictAttrs({}), - ret_ty=buf.ty, + ty=buf.ty, ), ) stmt = tirx.SeqStmt([alloc, body]) diff --git a/tests/python/tirx-base/test_binding_type_propagation.py b/tests/python/tirx-base/test_binding_type_propagation.py index 3f76ae95c5f0..ac31f8877d3f 100644 --- a/tests/python/tirx-base/test_binding_type_propagation.py +++ b/tests/python/tirx-base/test_binding_type_propagation.py @@ -63,7 +63,7 @@ def test_bind_uses_type_identity_after_real_value_mutation(inplace): var = ir.Var("y", binder_ty) input_var = ir.Var("x", result_ty) replacement = ir.Var("replacement", result_ty) - call = ir.Call("relax.exp", [input_var], ret_ty=result_ty) + call = ir.Call("relax.exp", [input_var], ty=result_ty) stmt = tirx.SeqStmt([tirx.Bind(var, call), tirx.Evaluate(var)]) if inplace: stmt = stmt._move() @@ -100,7 +100,7 @@ def test_unchanged_equal_but_distinct_types_do_not_rebind(inplace): binder_ty = relax.TensorType(result_ty.shape, result_ty.dtype) var = ir.Var("y", binder_ty) x = ir.Var("x", result_ty) - stmt = tirx.Bind(var, ir.Call("relax.exp", [x], ret_ty=result_ty)) + stmt = tirx.Bind(var, ir.Call("relax.exp", [x], ty=result_ty)) if inplace: stmt = stmt._move() @@ -146,7 +146,7 @@ def test_symbolic_call_type_rewrite_remaps_later_use(inplace): tensor_ty = relax.TensorType([n], "float32") x = ir.Var("x", tensor_ty) y = ir.Var("y", tensor_ty) - call = ir.Call("relax.exp", [x], ret_ty=tensor_ty) + call = ir.Call("relax.exp", [x], ty=tensor_ty) stmt = tirx.SeqStmt([tirx.Bind(y, call), tirx.Evaluate(y)]) if inplace: stmt = stmt._move() diff --git a/tests/python/tirx-base/test_tir_buffer.py b/tests/python/tirx-base/test_tir_buffer.py index 04732707bf8c..87066ef0ec29 100644 --- a/tests/python/tirx-base/test_tir_buffer.py +++ b/tests/python/tirx-base/test_tir_buffer.py @@ -94,7 +94,7 @@ def test_decl_buffer_physical_data_binding(): tvm.ir.DataTypeImm(tvm.DataType(buffer.dtype)), tvm.ir.StringImm(buffer.scope()), ], - ret_ty=buffer.ty, + ty=buffer.ty, ), ) assert decl.var.same_as(buffer) diff --git a/tests/python/tirx-base/test_tir_constructor.py b/tests/python/tirx-base/test_tir_constructor.py index 97fb392132bf..3c24f75f89cd 100644 --- a/tests/python/tirx-base/test_tir_constructor.py +++ b/tests/python/tirx-base/test_tir_constructor.py @@ -129,7 +129,7 @@ def test_expr_constructor(): assert x.vectors[0] == a assert x.indices[0].value == 0 - x = tvm.ir.Call("tirx.call_extern", [tvm.ir.StringImm("xyz"), a], ret_ty="float32") + x = tvm.ir.Call("tirx.call_extern", [tvm.ir.StringImm("xyz"), a], ty="float32") assert isinstance(x, tvm.ir.Call) assert tvm.ir.is_prim_expr(x) assert x.ty == tvm.ir.PrimType("float32") @@ -142,7 +142,7 @@ def test_expr_constructor(): "tirx.call_extern", [tvm.ir.StringImm("xyz"), attr_arg], attrs={"disable_tma": True}, - ret_ty="float32", + ty="float32", ) assert x_with_attrs.attrs["disable_tma"] is True assert not tvm_ffi.structural_equal(x, x_with_attrs) @@ -174,7 +174,7 @@ def test_expr_constructor(): "tirx.call_extern", [tvm.ir.StringImm("xyz"), attr_arg], attrs={"disable_tma": False}, - ret_ty="float32", + ty="float32", ) assert not expr_deep_equal(x_with_attrs, x_with_other_attrs) @@ -186,7 +186,7 @@ def call_with(arg): return tvm.ir.Call( "tirx.call_extern", [tvm.ir.StringImm("tuple_arg"), arg], - ret_ty="int32", + ty="int32", ) assert expr_deep_equal(call_with(tuple_arg), call_with(same_tuple_arg)) @@ -201,13 +201,13 @@ def call_with(arg): inner_if = tvm.ir.Call( "prim.if_then_else", [cond1, tvm.tirx.IntImm("int32", 1), tvm.tirx.IntImm("int32", 0)], - ret_ty="int32", + ty="int32", ) outer_if = tvm.ir.Call( "prim.if_then_else", [cond0, inner_if, tvm.tirx.IntImm("int32", 0)], attrs={"keep": True}, - ret_ty="int32", + ty="int32", ) simplified = tvm.tirx.transform.StmtSimplify()( tvm.IRModule({"main": tvm.tirx.PrimFunc([], tvm.tirx.Evaluate(outer_if))}) @@ -333,7 +333,7 @@ def test_stmt_constructor(): tvm.ir.StringImm(buf.scope()), ], attrs=tvm.ir.DictAttrs({}), - ret_ty=buf.ty, + ty=buf.ty, ), ) assert _is_buffer_binding(x, "tirx.alloc_buffer") diff --git a/tests/python/tirx-transform/test_tir_inline_private_functions.py b/tests/python/tirx-transform/test_tir_inline_private_functions.py index 9914bee6837e..6f0c24ea7bd5 100644 --- a/tests/python/tirx-transform/test_tir_inline_private_functions.py +++ b/tests/python/tirx-transform/test_tir_inline_private_functions.py @@ -224,12 +224,11 @@ def main(A: T.Buffer(16, "float32")): Before.subroutine( T.tvm_stack_make_array( A.data, - T.tvm_stack_make_shape(*A.ty.shape, dtype="handle"), + T.tvm_stack_make_shape(*A.ty.shape), 0, len(A.ty.shape), 0.0, A.ty.elem_offset, - dtype="handle", ) ) diff --git a/tests/python/tirx-transform/test_tir_transform_convert_ssa.py b/tests/python/tirx-transform/test_tir_transform_convert_ssa.py index 4d978f911902..e27ab33afe24 100644 --- a/tests/python/tirx-transform/test_tir_transform_convert_ssa.py +++ b/tests/python/tirx-transform/test_tir_transform_convert_ssa.py @@ -548,7 +548,7 @@ def test_shared_shape_var_in_buffer_params_and_alloc_buffer(): tvm.ir.StringImm(C.scope()), ], attrs=tvm.ir.DictAttrs({}), - ret_ty=C.ty, + ty=C.ty, ), ), tirx.Evaluate(1), @@ -591,7 +591,7 @@ def test_reused_loop_var_in_decl_buffer_elem_offset(): tvm.ir.DataTypeImm(tvm.DataType(buffer.dtype)), tvm.ir.StringImm(buffer.scope()), ], - ret_ty=buffer.ty, + ty=buffer.ty, ), ), tirx.Evaluate(tirx.BufferLoad(buffer, [0])), diff --git a/tests/python/tirx-transform/test_tir_transform_narrow_datatype.py b/tests/python/tirx-transform/test_tir_transform_narrow_datatype.py index 250b4ab6bec1..d26cd17fe8be 100644 --- a/tests/python/tirx-transform/test_tir_transform_narrow_datatype.py +++ b/tests/python/tirx-transform/test_tir_transform_narrow_datatype.py @@ -196,7 +196,6 @@ def before(A: T.Buffer((128,), "float32"), B: T.Buffer((130,), "float32")): i * T.int64(65) + j >= T.int64(0) and i * T.int64(65) + j < T.int64(128), A[i * T.int64(65) + j], 0.0, - dtype="float32", ) @T.prim_func @@ -206,7 +205,7 @@ def expected_after(A: T.Buffer(128, "float32"), B: T.Buffer(130, "float32")): A[i * 65 + j] = T.float32(0) for i, j in T.grid(2, 65): B[i * 65 + j] = T.if_then_else( - i * 65 + j >= 0 and i * 65 + j < 128, A[i * 65 + j], T.float32(0), dtype="float32" + i * 65 + j >= 0 and i * 65 + j < 128, A[i * 65 + j], T.float32(0) ) after = tvm.tirx.transform.NarrowDataType(32)( diff --git a/tests/python/tirx-transform/test_tir_transform_simplify.py b/tests/python/tirx-transform/test_tir_transform_simplify.py index 73d7b98730df..3c2b7dbcfb9f 100644 --- a/tests/python/tirx-transform/test_tir_transform_simplify.py +++ b/tests/python/tirx-transform/test_tir_transform_simplify.py @@ -486,7 +486,7 @@ def test_if_then_else_expr(): def before(A: T.Buffer(16, "float32")): for i in T.serial(16): if i < 12: - A[i] = T.if_then_else(i < 12, 1.0, 2.0, dtype="float32") + A[i] = T.if_then_else(i < 12, 1.0, 2.0) @T.prim_func(private=True) def expected(A: T.Buffer(16, "float32")): @@ -503,9 +503,7 @@ def test_ceil_log2_int(): @T.prim_func(private=True) def before(A: T.Buffer(1, "int32")): - A[0] = T.cast( - T.ceil(T.log2(T.cast(14, "float64"), dtype="float64"), dtype="float64"), dtype="int32" - ) + A[0] = T.cast(T.ceil(T.log2(T.cast(14, "float64"))), dtype="int32") @T.prim_func(private=True) def expected(A: T.Buffer(1, "int32")): @@ -526,7 +524,7 @@ def test_left_ceil_log2_lower_bound(): def before(A: T.Buffer(16, "float32")): for i in T.serial(16): x: T.let[T.int32] = T.cast( - T.ceil(T.log2(T.cast(i + 1024 + 1, "float64"), dtype="float64"), dtype="float64"), + T.ceil(T.log2(T.cast(i + 1024 + 1, "float64"))), dtype="int32", ) if x == 11: @@ -556,7 +554,7 @@ def test_left_shift_lower_bound(): @T.prim_func(private=True) def before(A: T.Buffer(16, "float32")): for i in T.serial(16): - if T.shift_left(1, i, dtype="int32") >= 1: + if T.shift_left(1, i) >= 1: A[i] = 0.0 @T.prim_func(private=True) @@ -579,7 +577,7 @@ def test_left_shift_upper_bound(): @T.prim_func(private=True) def before(A: T.Buffer(16, "float32")): for i in T.serial(16): - if T.shift_left(31, i, dtype="int32") <= 1015808: + if T.shift_left(31, i) <= 1015808: A[i] = 0.0 @T.prim_func(private=True) @@ -602,7 +600,7 @@ def test_left_shift_of_negative_value(): @T.prim_func(private=True) def before(A: T.Buffer(16, "float32")): for i in T.serial(16): - if -64 <= T.shift_left(-i, 4, dtype="int32"): + if -64 <= T.shift_left(-i, 4): A[i] = 0.0 expected = before @@ -622,7 +620,7 @@ def test_left_shift_by_negative_value(): @T.prim_func(private=True) def before(A: T.Buffer(16, "float32")): for i in T.serial(16): - if T.shift_left(16, -i, dtype="int32") <= 16: + if T.shift_left(16, -i) <= 16: A[i] = 0.0 expected = before diff --git a/tests/python/tirx-transform/test_tir_transform_storage_rewrite.py b/tests/python/tirx-transform/test_tir_transform_storage_rewrite.py index 5fd92e535127..d3ad057317c2 100644 --- a/tests/python/tirx-transform/test_tir_transform_storage_rewrite.py +++ b/tests/python/tirx-transform/test_tir_transform_storage_rewrite.py @@ -364,7 +364,7 @@ def func(A: T.Buffer((8,), "float32")): for i in range(8): B = T.alloc_buffer((1,)) B[0] = 3.14 - x: T.let[T.float32] = T.exp(B[0], dtype="float32") + x: T.let[T.float32] = T.exp(B[0]) A[i] = (x + 1.0) / (x - 1.0) @T.prim_func @@ -372,7 +372,7 @@ def func_rewritten(A: T.Buffer((8,), "float32")) -> None: B = T.alloc_buffer((1,)) for i in range(8): B[0] = 3.14 - x: T.let[T.float32] = T.exp(B[0], dtype="float32") + x: T.let[T.float32] = T.exp(B[0]) A[i] = (x + 1.0) / (x - 1.0) mod = tvm.tirx.transform.StorageRewrite()( diff --git a/tests/python/tirx/codegen/test_codegen_cuda.py b/tests/python/tirx/codegen/test_codegen_cuda.py index a9ea171582ec..2e0123c7bd84 100644 --- a/tests/python/tirx/codegen/test_codegen_cuda.py +++ b/tests/python/tirx/codegen/test_codegen_cuda.py @@ -160,7 +160,7 @@ def test_vector_access_ptr_preserves_packed_offset(monkeypatch): tvm.ir.DataTypeImm(tvm.DataType(buffer.dtype)), tvm.ir.StringImm(buffer.scope()), ], - ret_ty=buffer.ty, + ty=buffer.ty, ), ), tvm.tirx.Evaluate(tvm.tirx.call_extern("void", "consume", access_ptr)), diff --git a/tests/python/tirx/test_buffer_data_reinfer_type.py b/tests/python/tirx/test_buffer_data_reinfer_type.py index f4bbfdb9dea6..45439ec59db8 100644 --- a/tests/python/tirx/test_buffer_data_reinfer_type.py +++ b/tests/python/tirx/test_buffer_data_reinfer_type.py @@ -27,7 +27,7 @@ def test_buffer_data_reinfer_type_from_rewritten_argument(): stale_type = tirx.buffer_data_pointer_type(global_buffer) expected_type = tirx.buffer_data_pointer_type(local_buffer) - call = ir.Call("tirx.buffer_data", [local_buffer], ret_ty=stale_type) + call = ir.Call("tirx.buffer_data", [local_buffer], ty=stale_type) ir.assert_structural_equal(ir.reinfer_type(call), expected_type) ir.assert_structural_equal(call.ty, stale_type) From 862f6ef6942e0c70436fc6da7ed4d19c04af1a3f Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Fri, 2 Oct 2026 19:49:18 +0000 Subject: [PATCH 2/2] [DOCS] Keep the shared Call alias out of duplicate API indexing --- docs/reference/api/python/script/ir_builder.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/reference/api/python/script/ir_builder.rst b/docs/reference/api/python/script/ir_builder.rst index 797511815593..35ff6bfa957f 100644 --- a/docs/reference/api/python/script/ir_builder.rst +++ b/docs/reference/api/python/script/ir_builder.rst @@ -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 ***************************************