diff --git a/python/tvm/te/operation.py b/python/tvm/te/operation.py index 8da78a599c28..5279c46aebc2 100644 --- a/python/tvm/te/operation.py +++ b/python/tvm/te/operation.py @@ -326,7 +326,11 @@ def extern( if not isinstance(t, _tensor.Tensor): raise ValueError("expect inputs to be tensor") if in_buffers is None: - input_placeholders.append(tvm.tir.decl_buffer(t.shape, t.dtype, t.op.name)) + input_placeholders.append( + tvm.tir.decl_buffer( + t.shape, t.dtype, t.op.name, elem_offset=tvm.tir.Var("elem_offset", "int32") + ) + ) types.add(t.dtype) if dtype is None: @@ -339,7 +343,9 @@ def extern( if out_buffers is None: for shp, dt in zip(shape, dtype): - output_placeholders.append(tvm.tir.decl_buffer(shp, dt, name)) + output_placeholders.append( + tvm.tir.decl_buffer(shp, dt, name, elem_offset=tvm.tir.Var("elem_offset", "int32")) + ) body = fcompute(input_placeholders, output_placeholders) if isinstance(body, tvm.tir.PrimExpr): body = tvm.tir.Evaluate(body) diff --git a/tests/python/relay/test_op_level1.py b/tests/python/relay/test_op_level1.py index 3436bdd9f28d..4234c18c110f 100644 --- a/tests/python/relay/test_op_level1.py +++ b/tests/python/relay/test_op_level1.py @@ -820,5 +820,31 @@ def test_dense_rocm_sdot4(): np.testing.assert_equal(out, ref) +def test_extern_concat_injective_fuse(): + # This is a subgraph from MobileBERT, which crashes compilation if buffers created in te.extern(...) + # do not have their elem_offset explicitly set as a variable. + + # fmt: off + mod = tvm.parser.fromtext( + """ + #[version = "0.0.5"] + def @main(%p0844: Tensor[(1, 384), int64], %p1652: Tensor[(2016, 128), float16]) { + %1331 = cast(%p0844, dtype="int32"); + %1332 = take(%p1652, %1331, axis=0); + %1333 = strided_slice(%1332, begin=[0, 1, 0], end=[1, 384, 128], strides=[1, 1, 1], axes=None); + %1334 = strided_slice(%1332, begin=[0, 0, 0], end=[1, -1, 128], strides=[1, 1, 1], axes=None); + %1335 = nn.pad(%1333, 0, pad_width=[[0, 0], [0, 1], [0, 0]]); + %1336 = nn.pad(%1334, 0, pad_width=[[0, 0], [1, 0], [0, 0]]); + %1337 = (%1335, %1332, %1336); + %1338 = concatenate(%1337, axis=2); + reshape(%1338, newshape=[-1, 384]) + } + """ + ) + # fmt: on + + relay.build(mod, params={}, target="llvm") + + if __name__ == "__main__": pytest.main([__file__]) diff --git a/tests/python/unittest/test_te_create_primfunc.py b/tests/python/unittest/test_te_create_primfunc.py index d10fd2d23d47..4c216cdbc53a 100644 --- a/tests/python/unittest/test_te_create_primfunc.py +++ b/tests/python/unittest/test_te_create_primfunc.py @@ -216,9 +216,12 @@ def te_extern(): @T.prim_func def tir_extern(a: T.handle, b: T.handle, c: T.handle) -> None: T.func_attr({"global_symbol": "main", "tir.noalias": True}) - A = T.match_buffer(a, (128, 128)) - B = T.match_buffer(b, (128, 128)) - C = T.match_buffer(c, (128, 128)) + off1 = te.var("elem_offset") + off2 = te.var("elem_offset_1") + off3 = te.var("elem_offset_2") + A = T.match_buffer(a, (128, 128), elem_offset=off1) + B = T.match_buffer(b, (128, 128), elem_offset=off2) + C = T.match_buffer(c, (128, 128), elem_offset=off3) # body with T.block("C"): T.reads([A[0:128, 0:128], B[0:128, 0:128]]) @@ -232,7 +235,7 @@ def tir_extern(a: T.handle, b: T.handle, c: T.handle) -> None: 0, 2, 0.0, - 0, + off1, dtype="handle", ), T.tvm_stack_make_array( @@ -241,7 +244,7 @@ def tir_extern(a: T.handle, b: T.handle, c: T.handle) -> None: 0, 2, 0.0, - 0, + off2, dtype="handle", ), T.tvm_stack_make_array( @@ -250,7 +253,7 @@ def tir_extern(a: T.handle, b: T.handle, c: T.handle) -> None: 0, 2, 0.0, - 0, + off3, dtype="handle", ), 0,