mirror of
https://github.com/tile-ai/tilelang.git
synced 2026-10-03 06:48:18 +08:00
[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:
@@ -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").')
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user