[Ascend] Hold the target scope at lower call sites, not inside lower

tilelang.lower expects its caller to hold the target scope: the
vectorize planner consults Target::Current(false) (fail-loud), the JIT
path and tools/compile_only already enter it, and upstream's own tests
wrap bare lower calls the same way. The fork's with-target inside
lower_to_host_device_ir papered over that contract for every caller;
drop it and restore tilelang/engine/lower.py to upstream byte-for-byte.

The thirteen Ascend tests whose lowering actually reaches the planner
now enter the scope explicitly at their call sites, as does the
annotate_unlimit_memory example. Also remove a stray module-scope test
invocation in test_simtvf_fragment_narrow_checker.py that ran the test
at import time, and conftest's unused pytest import.
This commit is contained in:
LeiWang1999
2026-09-15 01:04:27 +08:00
parent 88252f43e6
commit df4596878f
6 changed files with 34 additions and 25 deletions
@@ -72,5 +72,10 @@ if __name__ == "__main__":
print("Default L1 limit: 512 KB -- annotate_unlimit_memory required\n")
program = gemm_with_unlimit(M, N, K, block_M, block_N, block_K, dtype)
mod = tilelang.lower(program, target=args.target)
# tilelang.lower expects the caller to hold the target scope; passes that
# consult Target.current() read the Ascend vector capabilities through it.
from tvm.target import Target
with Target(args.target):
mod = tilelang.lower(program, target=args.target)
print(f'Compilation succeeded with target="{args.target}" and T.annotate_unlimit_memory("shared.l1").')
-1
View File
@@ -18,7 +18,6 @@ source-only lowering tests keep running on a host without an NPU.
from __future__ import annotations
import pytest
import tilelang.testing
@@ -9,7 +9,10 @@ from tilelang.engine.lower import lower_to_host_device_ir
def _run_to_optimize(func):
context = create_backend_context("ascend")
mod = tvm.IRModule({"main": func})
return lower_to_host_device_ir(mod, context)
# lower_to_host_device_ir expects its caller to hold the target scope;
# the vectorize planner consults Target.current().
with tvm.target.Target("ascend"):
return lower_to_host_device_ir(mod, context)
def test_fragment_narrow_checker_fragment():
@@ -42,6 +45,3 @@ def test_fragment_narrow_checker_rejects_fragment_leak():
with pytest.raises(ValueError, match="Parallel loops outside VF blocks are not supported on Ascend NPU"):
_run_to_optimize(leaked)
test_fragment_narrow_checker_rejects_fragment_leak()
@@ -258,10 +258,13 @@ def test_nd2nz_scatter(sd, dd):
def test_nd2nz_scatter_tcopy_codegen():
source = tilelang.lower(
nd2nz_scatter_test(torch.float32, torch.bfloat16, ROWS, COLS),
target="ascend",
).kernel_source
# tilelang.lower expects the caller to hold the target scope; the
# vectorize planner consults Target.current(). Same below.
with tvm.target.Target("ascend"):
source = tilelang.lower(
nd2nz_scatter_test(torch.float32, torch.bfloat16, ROWS, COLS),
target="ascend",
).kernel_source
assert "__global__ __vector__ void main_kernel" in source
assert "ascend_nd2nz_scatter<32, 128, float, bfloat16_t>" in source
@@ -269,37 +272,41 @@ def test_nd2nz_scatter_tcopy_codegen():
def test_nd2nz_scatter_tcopy_requires_padding_row():
with pytest.raises(ValueError, match="reserve exactly one padding row"):
with tvm.target.Target("ascend"), pytest.raises(ValueError, match="reserve exactly one padding row"):
tilelang.lower(_invalid_unpadded_nd2nz_copy(ROWS, COLS), target="ascend")
def test_staged_nd2nz_scatter_uses_loop_bounds():
source = tilelang.lower(_staged_nd2nz_copy(), target="ascend").kernel_source
with tvm.target.Target("ascend"):
source = tilelang.lower(_staged_nd2nz_copy(), target="ascend").kernel_source
assert "ascend_nd2nz_scatter<16, 16, half, half>" in source
def test_reshaped_source_tcopy_codegen():
source = tilelang.lower(_reshaped_source_nd2nz_copy(), target="ascend").kernel_source
with tvm.target.Target("ascend"):
source = tilelang.lower(_reshaped_source_nd2nz_copy(), target="ascend").kernel_source
assert "ascend_nd2nz_scatter<32, 128, bfloat16_t, float>" in source
def test_nd2nz_scatter_rejects_strided_source():
with pytest.raises(ValueError, match="compact trailing source matrix"):
with tvm.target.Target("ascend"), pytest.raises(ValueError, match="compact trailing source matrix"):
tilelang.lower(_strided_source_nd2nz_copy(), target="ascend")
def test_nd2nz_scatter_accepts_aligned_leading_stride():
source = tilelang.lower(_leading_strided_source_nd2nz_copy(ROWS * COLS), target="ascend").kernel_source
with tvm.target.Target("ascend"):
source = tilelang.lower(_leading_strided_source_nd2nz_copy(ROWS * COLS), target="ascend").kernel_source
assert "ascend_nd2nz_scatter<32, 128, float, float>" in source
def test_nd2nz_scatter_rejects_unaligned_source_address():
with pytest.raises(ValueError, match="32-byte-aligned source address"):
with tvm.target.Target("ascend"), pytest.raises(ValueError, match="32-byte-aligned source address"):
tilelang.lower(_leading_strided_source_nd2nz_copy(ROWS * COLS - 1), target="ascend")
def test_simtvf_copy_is_not_rewritten_to_simdvf_scatter():
source = tilelang.lower(_simtvf_nd2nz_copy(), target="ascend").kernel_source
with tvm.target.Target("ascend"):
source = tilelang.lower(_simtvf_nd2nz_copy(), target="ascend").kernel_source
assert "__simt_vf__" in source
assert "ascend_nd2nz_scatter" not in source
@@ -148,7 +148,9 @@ def test_reducer_v2_rejects_unsupported_bitwise_collectives(op, target):
T.finalize_reducer(partial, result)
T.copy(result, B)
with pytest.raises(Exception, match="bitand, bitor, and bitxor are not supported"):
# tilelang.lower expects the caller to hold the target scope; the
# vectorize planner consults Target.current().
with tilelang.tvm.target.Target(target), pytest.raises(Exception, match="bitand, bitor, and bitxor are not supported"):
tilelang.lower(kernel, target=target)
+3 -7
View File
@@ -92,7 +92,6 @@ def host_codegen(
def _prepare_device_codegen_mod(device_mod: tvm.IRModule) -> tvm.IRModule:
device_mod = tilelang.transform.LowerIntrin()(device_mod)
device_mod = tirx.transform.Simplify()(device_mod)
device_mod = tilelang.transform.HoistBroadcastValues()(device_mod)
return device_mod
@@ -131,13 +130,10 @@ def lower_to_host_device_ir(
_is_host_call = get_host_call(is_device_c=is_cpu_device_backend(target))
_is_device_call = get_device_call(is_device_c=is_cpu_device_backend(target))
# Keep target context active for passes that rely on Target::Current.
with target:
# Run backend-independent semantic checks before target-specific
# lowering.
PreLowerSemanticCheck(mod)
# Run backend-independent semantic checks before target-specific lowering.
PreLowerSemanticCheck(mod)
mod = context.lower(mod)
mod = context.lower(mod)
host_mod = tirx.transform.Filter(_is_host_call)(mod)
device_mod = tirx.transform.Filter(_is_device_call)(mod)