diff --git a/tests/python/relax/test_transform_canonicalize_bindings.py b/tests/python/relax/test_transform_canonicalize_bindings.py index 1a4c1d1b539d..325d6d711b24 100644 --- a/tests/python/relax/test_transform_canonicalize_bindings.py +++ b/tests/python/relax/test_transform_canonicalize_bindings.py @@ -1417,9 +1417,9 @@ class Before: @Ts.prim_func(private=True) def transpose( - x: T.Buffer((1, n, 2, m, 2, 8), "float16"), - y: T.Buffer((1, n, m, 2, 2, 8), "float16"), - z: T.Buffer((k,), "float16"), + x: T.Tensor((1, n, 2, m, 2, 8), "float16"), + y: T.Tensor((1, n, m, 2, 2, 8), "float16"), + z: T.Tensor((k,), "float16"), ): for i0, i1, i2, i3, i4, i5 in T.grid(1, n, m, 2, 2, 8): with Ts.sblock("b"): @@ -1427,7 +1427,7 @@ def transpose( y[v0, v1, v2, v3, v4, v5] = x[v0, v1, v3, v2, v4, v5] @Ts.prim_func(private=True) - def add_scalar(x: T.Buffer((1, 8), "float16"), s: T.int64, y: T.Buffer((1, 8), "float16")): + def add_scalar(x: T.Tensor((1, 8), "float16"), s: T.int64, y: T.Tensor((1, 8), "float16")): for i in T.serial(8): with Ts.sblock("b"): vi = Ts.axis.spatial(8, i) @@ -1475,14 +1475,14 @@ def test_call_tir_out_ty_follows_arguments_through_chained_match_casts(): @I.ir_module class Before: @Ts.prim_func(private=True) - def copy(x: T.Buffer((n, 8), "float16"), y: T.Buffer((n, 8), "float16")): + def copy(x: T.Tensor((n, 8), "float16"), y: T.Tensor((n, 8), "float16")): for i, j in T.grid(n, 8): with Ts.sblock("b"): vi, vj = Ts.axis.remap("SS", [i, j]) y[vi, vj] = x[vi, vj] @Ts.prim_func(private=True) - def add_scalar(x: T.Buffer((1, 8), "float16"), s: T.int64, y: T.Buffer((1, 8), "float16")): + def add_scalar(x: T.Tensor((1, 8), "float16"), s: T.int64, y: T.Tensor((1, 8), "float16")): for i in T.serial(8): with Ts.sblock("b"): vi = Ts.axis.spatial(8, i)