diff --git a/src/target/sunmmio/sunmmio_codegen_tiles_loop.cc b/src/target/sunmmio/sunmmio_codegen_tiles_loop.cc index 410e15ddd2..63beb0d8e9 100644 --- a/src/target/sunmmio/sunmmio_codegen_tiles_loop.cc +++ b/src/target/sunmmio/sunmmio_codegen_tiles_loop.cc @@ -3670,18 +3670,13 @@ bool CodeGenTileLangSunMMIO::TryLowerTilesScope(const tir::ForNode *op) { select->false_value, select->dtype); } if (const auto *cast = expr.as()) { - SunMMIOValue value = lower_expr(cast->value, state, preferred_dtype); + // An explicit cast applies after its operand has been evaluated. + SunMMIOValue value = lower_expr(cast->value, state, std::nullopt); if (IsTileLike(value)) { DataType dst_dtype = CanonicalizeSuvmDType(cast->dtype).with_lanes(1); if (value.dtype == dst_dtype) { return value; } - if (preferred_dtype.has_value() && is_float_like_dtype(value.dtype) && - is_float_like_dtype(dst_dtype) && - value.dtype == - CanonicalizeSuvmDType(preferred_dtype.value()).with_lanes(1)) { - return value; - } SunMMIOType dst_type = MakeTileType(CanonicalizeSuvmDType(cast->dtype), ExtractStaticShape(value.type)); return builder_->Cast(NewValueName(), value, dst_type, diff --git a/testing/python/sunmmio/codegen/test_tile_ops_opt_validate.py b/testing/python/sunmmio/codegen/test_tile_ops_opt_validate.py index 1151c1c77d..ea28aa58f5 100644 --- a/testing/python/sunmmio/codegen/test_tile_ops_opt_validate.py +++ b/testing/python/sunmmio/codegen/test_tile_ops_opt_validate.py @@ -1,4 +1,5 @@ import os +import re import tilelang import tilelang.language as T @@ -17,6 +18,7 @@ # os.environ["SUNMMIO_TEST_LOG_IR"] = "1" LOOSE_OPT_ARGS = ("--verify-each",) +STRICT_OPT_ARGS = ("--verify-each", "--suvm-to-llvm-pipeline") def validate_sunmmio_codegen_loose(kernel, tmp_path, *, mlir_filename, expected_tokens=()): @@ -188,6 +190,38 @@ def main( return main +@target("Sunmmio") +def fp32_select_then_bf16_cast_test(m=32, n=32): + input_dtype = T.float32 + output_dtype = T.bfloat16 + shard_policy = T.MeshShardingPolicy() + tensor_shape = (m, n) + tensor_layout = make_zz_layout(tensor_shape, [0, 1], tensor_shape) + + @T.prim_func + def main( + A: T.MeshTensor(tensor_shape, shard_policy, input_dtype, layout=tensor_layout), # type: ignore + C: T.MeshTensor(tensor_shape, shard_policy, output_dtype, layout=tensor_layout), # type: ignore + ): + with T.Kernel(): + A_shared = T.alloc_shared(tensor_shape, input_dtype) + C_shared = T.alloc_shared(tensor_shape, output_dtype) + + T.copy(A, A_shared) + for i, j in T.Tiles(A_shared, parallel=True): + C_shared[i, j] = T.Cast( + output_dtype, + T.if_then_else( + A_shared[i, j] > T.float32(0), + A_shared[i, j], + T.float32(0), + ), + ) + T.copy(C_shared, C) + + return main + + def test_tile_elementwise_ops_2d_codegen_validates_with_npuir_opt(tmp_path): src = validate_sunmmio_codegen_with_npuir_opt( tile_elementwise_ops_2d_test(), @@ -244,5 +278,28 @@ def test_tile_elementwise_ops_codegen_validates_loose_with_npuir_opt(tmp_path): ) +def test_fp32_select_is_evaluated_before_bf16_cast(tmp_path): + src = validate_sunmmio_codegen_with_npuir_opt( + fp32_select_then_bf16_cast_test(), + tmp_path, + mlir_filename="fp32_select_then_bf16_cast_suvm.mlir", + expected_tokens=("suvm.tile.cmpf", "suvm.tile.select", "suvm.tile.cast"), + opt_args=STRICT_OPT_ARGS, + ) + + select = re.search( + r"(?P%[\w.]+) = suvm\.tile\.select .*" + r"!suvm\.tile<[^>]*xf32>, !suvm\.tile<[^>]*xf32>" + r" -> !suvm\.tile<[^>]*xf32>", + src, + ) + assert select, src + assert re.search( + rf"suvm\.tile\.cast {re.escape(select.group('result'))} : " + r"!suvm\.tile<[^>]*xf32> -> !suvm\.tile<[^>]*xbf16>", + src, + ), src + + if __name__ == "__main__": tilelang.testing.main() diff --git a/testing/python/sunmmio/jit/test_suvm_edit_session.py b/testing/python/sunmmio/jit/test_suvm_edit_session.py new file mode 100644 index 0000000000..7223198675 --- /dev/null +++ b/testing/python/sunmmio/jit/test_suvm_edit_session.py @@ -0,0 +1,305 @@ +from pathlib import Path +import sys +from types import ModuleType, SimpleNamespace + +import pytest + +import tilelang +from tilelang import tvm +from tilelang.jit.adapter.sunmmio import SunmmioKernelABI +from tilelang.jit.adapter.sunmmio import suvm_edit_session as session_module +from tilelang.jit.adapter.sunmmio import adapter as adapter_module +from tilelang.jit.adapter.sunmmio.adapter import ( + SunmmioKernelSuDeckAdapter, + SunmmioSunsimKernelAdapter, +) +from tilelang.jit.adapter.sunmmio.libgen import SunmmioKernelArtifact +from tilelang.jit.adapter.sunmmio.suvm_edit_session import SunmmioSuvmEditSession + + +def _empty_kernel(): + return tvm.tir.PrimFunc([], tvm.tir.Evaluate(0)).with_attr("global_symbol", "main") + + +def _empty_abi(): + return SunmmioKernelABI( + kernel_name="main_kernel", + public_arg_count=0, + public_param_names=(), + device_param_names=(), + device_param_dtypes=(), + runtime_scalars=(), + ) + + +def _fake_lowered(source: str): + device_func = tvm.tir.PrimFunc([], tvm.tir.Evaluate(0)).with_attr("tir.is_global_func", True) + return SimpleNamespace( + kernel_source=source, + host_mod=None, + device_mod=tvm.IRModule({"main_kernel": device_func}), + params=[], + ) + + +def _write_compile_inputs(session): + artifacts = session.artifacts + artifacts.path(session_module.ORIGINAL_MLIR).write_text("module {}\n", encoding="utf-8") + artifacts.edited_mlir.write_text("module { // edited\n}\n", encoding="utf-8") + artifacts.path(session_module.DEVICE_TIR).write_text("# device tir\n", encoding="utf-8") + session_module._write_json(artifacts.path(session_module.ABI_FILE), _empty_abi().to_json_dict()) + session_module._write_json( + artifacts.path(session_module.MANIFEST_FILE), + { + "schema_version": session_module.MANIFEST_SCHEMA_VERSION, + "target": "sunmmio", + "opt_level": 3, + "parameters": [], + }, + ) + + +def test_suvm_edit_session_repeated_emit_archives_previous_edit(tmp_path, monkeypatch): + lowered = _fake_lowered("module { // fresh\n}\n") + monkeypatch.setattr(tilelang, "lower", lambda *_args, **_kwargs: lowered) + monkeypatch.setattr( + session_module.SunmmioKernelABI, + "from_modules", + lambda **_kwargs: _empty_abi(), + ) + + session = SunmmioSuvmEditSession(tmp_path) + session.emit(_empty_kernel()) + session.artifacts.edited_mlir.write_text("module { // manual edit\n}\n", encoding="utf-8") + session.emit(_empty_kernel()) + + archives = list(tmp_path.glob("kernel.edited.*Z.mlir")) + assert len(archives) == 1 + assert "manual edit" in archives[0].read_text(encoding="utf-8") + assert session.artifacts.edited_mlir.read_text(encoding="utf-8") == lowered.kernel_source + + +def test_suvm_edit_session_lowering_failure_keeps_previous_edit(tmp_path, monkeypatch): + edited = tmp_path / "kernel.edited.mlir" + edited.write_text("module { // keep me\n}\n", encoding="utf-8") + monkeypatch.setattr( + tilelang, + "lower", + lambda *_args, **_kwargs: (_ for _ in ()).throw(RuntimeError("lowering failed")), + ) + + with pytest.raises(RuntimeError, match="lowering failed"): + SunmmioSuvmEditSession(tmp_path).emit(_empty_kernel()) + + assert "keep me" in edited.read_text(encoding="utf-8") + assert list(tmp_path.glob("kernel.edited.*Z.mlir")) == [] + + +def test_suvm_edit_session_compile_sunsim_uses_edited_mlir(tmp_path, monkeypatch): + session = SunmmioSuvmEditSession(tmp_path, timeout=17.0) + _write_compile_inputs(session) + + captured = {} + + class FakeGenerator: + def __init__(self, target, kernel_name, verbose=False): + captured.setdefault("kernel_names", []).append(kernel_name) + self.artifact = None + + def update_mlir_source(self, source): + captured["mlir"] = source + + def update_device_tir_source(self, source): + captured["tir"] = source + + def compile_lib(self, timeout, output_dir): + captured["timeout"] = timeout + elf = Path(output_dir) / "kernel.elf" + mlir = Path(output_dir) / "kernel.mlir" + llvm = Path(output_dir) / "kernel.ll" + elf.write_bytes(b"ELF") + mlir.write_text(captured["mlir"], encoding="utf-8") + llvm.write_text("define void @main_kernel() {}\n", encoding="utf-8") + self.artifact = SunmmioKernelArtifact( + elf_path=elf, + mlir_path=mlir, + llvm_ir_path=llvm, + build_dir=Path(output_dir), + runtime_kernel_name="main_kernel", + ) + + def load_lib(self, kernel_lib_path): + captured["loaded"] = Path(kernel_lib_path) + output_dir = Path(kernel_lib_path).parent + self.artifact = SunmmioKernelArtifact( + elf_path=Path(kernel_lib_path), + mlir_path=output_dir / "kernel.mlir", + llvm_ir_path=output_dir / "kernel.ll", + build_dir=output_dir, + runtime_kernel_name="main_kernel", + ) + + monkeypatch.setattr(SunmmioSuvmEditSession, "_validate_mlir", lambda *_args: None) + monkeypatch.setattr(session_module, "SunmmioSunsimLibraryGenerator", FakeGenerator) + monkeypatch.setattr(adapter_module, "SunmmioSunsimLibraryGenerator", FakeGenerator) + monkeypatch.setattr(session_module, "_validate_sunsim_elf_abi", lambda *_args: None) + + executable = session.compile_sunsim() + + assert isinstance(executable, SunmmioSunsimKernelAdapter) + assert captured == { + "kernel_names": ["main_kernel", "main_kernel"], + "mlir": "module { // edited\n}\n", + "tir": "# device tir\n", + "timeout": 17.0, + "loaded": tmp_path / "kernel.elf", + } + assert executable._artifact_parameter_kinds == () + assert not hasattr(executable, "_artifact_timeout") + assert executable.lib_generator.artifact.elf_path == tmp_path / "kernel.elf" + + runs = [] + fake_sunsim = SimpleNamespace( + Input=type("Input", (), {}), + Output=type("Output", (), {}), + Inout=type("Inout", (), {}), + Descriptor=type("Descriptor", (), {}), + run=lambda **kwargs: runs.append(kwargs), + ) + monkeypatch.setattr( + SunmmioSunsimKernelAdapter, + "_import_sunsim", + staticmethod(lambda: fake_sunsim), + ) + executable(timeout=31.0) + assert runs == [ + { + "elf": tmp_path / "kernel.elf", + "args": [], + "kernel_name": "main_kernel", + "timeout": 31.0, + } + ] + + +def test_suvm_edit_session_compile_uses_sudeck_runtime(tmp_path, monkeypatch): + session = SunmmioSuvmEditSession(tmp_path, timeout=19.0) + _write_compile_inputs(session) + captured = {} + + class FakeGenerator: + def __init__(self, target, kernel_name, verbose=False): + captured.setdefault("kernel_names", []).append(kernel_name) + self.artifact = None + self.pymodule = None + + def update_launcher_specs(self, specs): + captured.setdefault("launcher_specs", []).append(specs) + + def update_mlir_source(self, source): + captured["mlir"] = source + + def update_device_tir_source(self, source): + captured["tir"] = source + + def compile_lib(self, timeout, output_dir): + captured["timeout"] = timeout + output_dir = Path(output_dir) + (output_dir / "kernel.elf").write_bytes(b"ELF") + self.artifact = SunmmioKernelArtifact( + elf_path=output_dir / "kernel.elf", + mlir_path=output_dir / "kernel.mlir", + llvm_ir_path=output_dir / "kernel.ll", + build_dir=output_dir, + runtime_kernel_name="main_kernel", + ) + + def load_lib(self, kernel_lib_path): + captured["loaded"] = Path(kernel_lib_path) + output_dir = Path(kernel_lib_path).parent + self.artifact = SunmmioKernelArtifact( + elf_path=Path(kernel_lib_path), + mlir_path=output_dir / "kernel.mlir", + llvm_ir_path=output_dir / "kernel.ll", + build_dir=output_dir, + runtime_kernel_name="main_kernel", + ) + self.pymodule = SimpleNamespace(call=lambda *_args: None) + + monkeypatch.setattr(SunmmioSuvmEditSession, "_validate_mlir", lambda *_args: None) + monkeypatch.setattr(session_module, "SunmmioSuDeckLibraryGenerator", FakeGenerator) + monkeypatch.setattr(adapter_module, "SunmmioSuDeckLibraryGenerator", FakeGenerator) + + executable = session.compile() + + assert isinstance(executable, SunmmioKernelSuDeckAdapter) + assert captured == { + "kernel_names": ["main_kernel", "main_kernel"], + "launcher_specs": [[], []], + "mlir": "module { // edited\n}\n", + "tir": "# device tir\n", + "timeout": 19.0, + "loaded": tmp_path / "kernel.elf", + } + assert not hasattr(executable, "_artifact_parameter_kinds") + assert not hasattr(executable, "_artifact_timeout") + assert executable.lib_generator.artifact.elf_path == tmp_path / "kernel.elf" + + +def test_sudeck_adapter_loads_artifact_and_launches_torch_sunmmio_handles( + tmp_path, + monkeypatch, +): + stream = object() + torch_module = ModuleType("torch") + torch_module.sunmmio = SimpleNamespace( + current_stream=lambda: stream, + unsafe_get_sudeck_stream_handle=lambda value: 0x1234 if value is stream else 0, + ) + torch_sunmmio_module = ModuleType("torch_sunmmio") + torch_sunmmio_runtime = ModuleType("torch_sunmmio.sunmmio") + tensor = SimpleNamespace(shape=(8, 16), stride=lambda: (16, 1)) + torch_sunmmio_runtime.unsafe_get_sutensor_handle = lambda value: 0x5678 if value is tensor else 0 + monkeypatch.setitem(sys.modules, "torch", torch_module) + monkeypatch.setitem(sys.modules, "torch_sunmmio", torch_sunmmio_module) + monkeypatch.setitem(sys.modules, "torch_sunmmio.sunmmio", torch_sunmmio_runtime) + + calls = [] + abi = SunmmioKernelABI( + kernel_name="main_kernel", + public_arg_count=1, + public_param_names=("A",), + device_param_names=("A",), + device_param_dtypes=("handle",), + runtime_scalars=(), + ) + + class FakeGenerator: + def __init__(self, target, kernel_name, verbose=False): + self.artifact = None + self.pymodule = SimpleNamespace(call=lambda *args: calls.append(args)) + + def update_launcher_specs(self, specs): + assert specs == [("A", "tensor")] + + def load_lib(self, kernel_lib_path): + self.artifact = SunmmioKernelArtifact( + elf_path=Path(kernel_lib_path), + mlir_path=tmp_path / "kernel.mlir", + llvm_ir_path=tmp_path / "kernel.ll", + build_dir=tmp_path, + runtime_kernel_name="main_kernel", + ) + + elf_path = tmp_path / "kernel.elf" + elf_path.write_bytes(b"ELF") + monkeypatch.setattr(adapter_module, "SunmmioSuDeckLibraryGenerator", FakeGenerator) + executable = SunmmioKernelSuDeckAdapter.from_compiled_artifact( + target="sunmmio", + abi=abi, + kernel_lib_path=elf_path, + ) + + executable(tensor) + + assert calls == [(0x5678, 0x1234)] diff --git a/tilelang/jit/adapter/sunmmio/adapter.py b/tilelang/jit/adapter/sunmmio/adapter.py index 02a3b3f3e4..3baa28464c 100644 --- a/tilelang/jit/adapter/sunmmio/adapter.py +++ b/tilelang/jit/adapter/sunmmio/adapter.py @@ -68,6 +68,36 @@ def __init__( self.lib_generator.load_lib(kernel_lib_path) self._post_init() + @classmethod + def from_compiled_artifact( + cls, + *, + target: str | Target, + abi: SunmmioKernelABI, + kernel_lib_path: str | os.PathLike[str], + verbose: bool = False, + ): + """Load a compiled Sunmmio artifact without rebuilding its frontend kernel.""" + if not os.path.exists(kernel_lib_path): + raise FileNotFoundError(f"Compiled Sunmmio kernel artifact does not exist: {kernel_lib_path}") + + instance = cls.__new__(cls) + instance.params = [] + instance.result_idx = [] + instance.target = Target.canon_target(determine_target(target)) + if not target_is_sunmmio(instance.target): + raise ValueError(f"SunmmioKernelAdapter requires a Sunmmio target, got {instance.target}") + instance.ir_module = None + instance.abi = abi + instance.host_mod = None + instance.device_mod = None + instance.device_kernel_source = "" + instance.verbose = verbose + instance.lib_generator = instance._make_lib_generator(verbose) + instance.lib_generator.load_lib(kernel_lib_path) + instance._post_init() + return instance + @classmethod def from_database( cls, @@ -206,6 +236,31 @@ class SunmmioSunsimKernelAdapter(SunmmioKernelAdapter): def _make_lib_generator(self, verbose: bool) -> SunmmioSunsimLibraryGenerator: return SunmmioSunsimLibraryGenerator(self.target, self.kernel_name, verbose) + @classmethod + def from_compiled_artifact( + cls, + *, + target: str | Target, + abi: SunmmioKernelABI, + kernel_lib_path: str | os.PathLike[str], + parameter_kinds: Sequence[str], + verbose: bool = False, + ): + invalid_kinds = set(parameter_kinds) - {"tensor", "scalar"} + if invalid_kinds: + raise ValueError(f"Invalid Sunmmio parameter kinds: {sorted(invalid_kinds)}") + if len(parameter_kinds) != abi.public_arg_count: + raise ValueError(f"Sunmmio artifact has {len(parameter_kinds)} parameter kinds but ABI expects {abi.public_arg_count}.") + + instance = super().from_compiled_artifact( + target=target, + abi=abi, + kernel_lib_path=kernel_lib_path, + verbose=verbose, + ) + instance._artifact_parameter_kinds = tuple(parameter_kinds) + return instance + def _convert_torch_func(self) -> Callable[..., Any]: if self.result_idx: raise NotImplementedError( @@ -244,8 +299,11 @@ def _prepare_sunsim_args(self, args: Sequence[Any], sunsim) -> list[Any]: marker_types = (sunsim.Input, sunsim.Output, sunsim.Inout) descriptor_type = sunsim.Descriptor - for index, (arg, param) in enumerate(zip(args[: self.abi.public_arg_count], self.params)): - if param.is_scalar(): + parameter_kinds = getattr(self, "_artifact_parameter_kinds", None) + if parameter_kinds is None: + parameter_kinds = tuple("scalar" if param.is_scalar() else "tensor" for param in self.params) + for index, (arg, kind) in enumerate(zip(args[: self.abi.public_arg_count], parameter_kinds)): + if kind == "scalar": if isinstance(arg, marker_types): raise TypeError( f"Sunmmio sunsim argument {index} is a scalar slot, but got {type(arg).__name__}. " diff --git a/tilelang/jit/adapter/sunmmio/suvm_edit_session.py b/tilelang/jit/adapter/sunmmio/suvm_edit_session.py new file mode 100644 index 0000000000..4c4c04a3e3 --- /dev/null +++ b/tilelang/jit/adapter/sunmmio/suvm_edit_session.py @@ -0,0 +1,406 @@ +"""Persistent edit session for TileLang-generated Sunmmio SUVM MLIR.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass +from datetime import UTC, datetime +import difflib +import json +from pathlib import Path +import subprocess +from typing import Any + +import tilelang +from tilelang import tvm +from tilelang.engine.param import KernelParam +from tilelang.utils.target import determine_target + +from .adapter import SunmmioKernelSuDeckAdapter, SunmmioSunsimKernelAdapter +from .abi import SunmmioKernelABI +from .libgen import ( + SUNMMIO_KERNEL_ELF_FILE, + SUNMMIO_KERNEL_LLVM_IR_FILE, + SUNMMIO_KERNEL_MLIR_FILE, + SUNMMIO_KERNEL_OBJ_FILE, + SUNMMIO_SUDECK_LAUNCHER_CPP_FILE, + SUNMMIO_SUDECK_LAUNCHER_LIB_FILE, + SUNMMIO_SUDECK_LAUNCH_MODULE_FILE, + SunmmioKernelArtifact, + SunmmioSuDeckLibraryGenerator, + SunmmioSunsimLibraryGenerator, + find_npuir_tool, +) + + +ORIGINAL_MLIR = "kernel.original.mlir" +EDITED_MLIR = "kernel.edited.mlir" +DEVICE_TIR = "kernel.tir" +DEVICE_VALIDATED_MLIR = "kernel.device-validated.mlir" +LOWERED_MLIR = "kernel.lowered.mlir" +ABI_FILE = "abi.json" +DIFF_FILE = "kernel.diff" +MANIFEST_FILE = "manifest.json" +MANIFEST_SCHEMA_VERSION = 2 + +_COMPILED_ARTIFACTS = ( + DEVICE_VALIDATED_MLIR, + LOWERED_MLIR, + SUNMMIO_KERNEL_MLIR_FILE, + SUNMMIO_KERNEL_LLVM_IR_FILE, + SUNMMIO_KERNEL_ELF_FILE, + SUNMMIO_KERNEL_OBJ_FILE, + DIFF_FILE, + "main_thunk.cpp", + "main_thunk.o", + "CMakeLists.txt", + "device_sudeck.ld", + SUNMMIO_SUDECK_LAUNCHER_CPP_FILE, + SUNMMIO_SUDECK_LAUNCHER_LIB_FILE, + SUNMMIO_SUDECK_LAUNCH_MODULE_FILE, +) + + +@dataclass(frozen=True) +class SunmmioSuvmArtifacts: + """Stable artifact paths for a manual SUVM edit session.""" + + work_dir: Path + + def __post_init__(self) -> None: + object.__setattr__(self, "work_dir", Path(self.work_dir).resolve()) + + def path(self, name: str) -> Path: + return self.work_dir / name + + @property + def edited_mlir(self) -> Path: + return self.path(EDITED_MLIR) + + @property + def diff(self) -> Path: + return self.path(DIFF_FILE) + + @property + def llvm_ir(self) -> Path: + return self.path(SUNMMIO_KERNEL_LLVM_IR_FILE) + + @property + def elf(self) -> Path: + return self.path(SUNMMIO_KERNEL_ELF_FILE) + + def managed_paths(self) -> list[Path]: + sources = (ORIGINAL_MLIR, EDITED_MLIR, DEVICE_TIR, ABI_FILE, MANIFEST_FILE) + return [self.path(name) for name in (*sources, *_COMPILED_ARTIFACTS)] + + +@dataclass +class SunmmioSuvmEditSession: + """Emit editable SUVM MLIR and compile it into a callable Sunmmio kernel. + + A minimal two-stage workflow looks like this:: + + from pathlib import Path + + from tilelang.jit.adapter.sunmmio.suvm_edit_session import SunmmioSuvmEditSession + + session = SunmmioSuvmEditSession(Path("manual_suvm")) + + # Stage 1: lower a @tilelang.jit kernel only as far as SUVM MLIR. + artifacts = session.emit(jit_kernel.get_tir(...)) + print(f"Edit this file: {artifacts.edited_mlir}") + + # Stop here and manually edit manual_suvm/kernel.edited.mlir. + + # Stage 2: in a later run, reuse the same directory. This validates and + # compiles the edited MLIR, then accepts normal torch_sunmmio tensors. + kernel = session.compile() + kernel(a_dev, b_dev, c_dev) + + Do not call :meth:`emit` again after editing: a new emit starts a fresh + session input and archives the previous ``kernel.edited.mlir``. + """ + + work_dir: Path + target: str | tvm.target.Target = "sunmmio" + opt_level: int = 3 + timeout: float | None = 240.0 + + def __post_init__(self) -> None: + self.work_dir = Path(self.work_dir).resolve() + + @property + def artifacts(self) -> SunmmioSuvmArtifacts: + return SunmmioSuvmArtifacts(self.work_dir) + + def emit(self, kernel: tvm.tir.PrimFunc | tvm.IRModule) -> SunmmioSuvmArtifacts: + """Lower a kernel, archive the previous edit, and write a fresh editable MLIR.""" + if not isinstance(kernel, (tvm.tir.PrimFunc, tvm.IRModule)): + raise TypeError("SunmmioSuvmEditSession.emit expects a PrimFunc or IRModule.") + + resolved_target = determine_target(self.target, return_object=True) + with tvm.transform.PassContext(opt_level=self.opt_level), resolved_target: + lowered = tilelang.lower( + kernel, + target=resolved_target, + enable_host_codegen=False, + enable_device_compile=False, + ) + if not lowered.kernel_source or not lowered.kernel_source.strip(): + raise RuntimeError("Sunmmio lowering produced no SUVM MLIR.") + if lowered.device_mod is None: + raise RuntimeError("Sunmmio lowering produced no device TIR.") + + params = lowered.params or _kernel_params(kernel) + abi = SunmmioKernelABI.from_modules( + func_or_mod=kernel, + host_mod=lowered.host_mod, + device_mod=lowered.device_mod, + params=params, + ) + manifest = { + "schema_version": MANIFEST_SCHEMA_VERSION, + "target": str(resolved_target), + "opt_level": self.opt_level, + "parameters": [ + { + "name": name, + "kind": "scalar" if param.is_scalar() else "tensor", + "dtype": str(param.dtype), + "shape": [str(dim) for dim in param.shape], + } + for name, param in zip(abi.public_param_names, params) + ], + } + + artifacts = self.artifacts + _prepare_emit(artifacts) + artifacts.path(ORIGINAL_MLIR).write_text(lowered.kernel_source, encoding="utf-8") + artifacts.edited_mlir.write_text(lowered.kernel_source, encoding="utf-8") + artifacts.path(DEVICE_TIR).write_text(lowered.device_mod.script(), encoding="utf-8") + _write_json(artifacts.path(ABI_FILE), abi.to_json_dict()) + _write_json(artifacts.path(MANIFEST_FILE), manifest) + return artifacts + + def compile(self) -> SunmmioKernelSuDeckAdapter: + """Validate and compile the edited MLIR for the regular SuDeck runtime.""" + artifacts, manifest, abi, edited_source = self._prepare_compile() + target = determine_target(manifest["target"], return_object=True) + generator = SunmmioSuDeckLibraryGenerator(target, abi.kernel_name) + generator.update_launcher_specs( + [(name, "tensor" if dtype == "handle" else dtype) for name, dtype in zip(abi.device_param_names, abi.device_param_dtypes)] + ) + generator.update_mlir_source(edited_source) + tir_path = artifacts.path(DEVICE_TIR) + if tir_path.is_file(): + generator.update_device_tir_source(tir_path.read_text(encoding="utf-8")) + generator.compile_lib(timeout=self.timeout, output_dir=artifacts.work_dir) + if generator.artifact is None: + raise RuntimeError("Sunmmio SuDeck generator produced no artifact.") + return SunmmioKernelSuDeckAdapter.from_compiled_artifact( + target=target, + abi=abi, + kernel_lib_path=generator.artifact.elf_path, + ) + + def compile_sunsim(self) -> SunmmioSunsimKernelAdapter: + """Validate the edited MLIR and compile it into a callable sunsim ELF.""" + artifacts, manifest, abi, edited_source = self._prepare_compile() + target = determine_target(manifest["target"], return_object=True) + generator = SunmmioSunsimLibraryGenerator(target, abi.kernel_name) + generator.update_mlir_source(edited_source) + tir_path = artifacts.path(DEVICE_TIR) + if tir_path.is_file(): + generator.update_device_tir_source(tir_path.read_text(encoding="utf-8")) + generator.compile_lib(timeout=self.timeout, output_dir=artifacts.work_dir) + if generator.artifact is None: + raise RuntimeError("Sunmmio sunsim generator produced no artifact.") + _validate_sunsim_elf_abi(generator.artifact, abi) + return SunmmioSunsimKernelAdapter.from_compiled_artifact( + target=target, + abi=abi, + parameter_kinds=_parameter_kinds(manifest, abi), + kernel_lib_path=generator.artifact.elf_path, + ) + + def _prepare_compile( + self, + ) -> tuple[SunmmioSuvmArtifacts, dict[str, Any], SunmmioKernelABI, str]: + artifacts = self.artifacts + manifest = _load_manifest(artifacts) + abi = SunmmioKernelABI.from_json_dict(_read_json(artifacts.path(ABI_FILE), "Sunmmio ABI metadata")) + edited_source = _read_nonempty(artifacts.edited_mlir, "edited SUVM MLIR") + _write_diff(artifacts.path(ORIGINAL_MLIR), artifacts.edited_mlir, artifacts.diff) + self._validate_mlir(artifacts) + return artifacts, manifest, abi, edited_source + + def _validate_mlir(self, artifacts: SunmmioSuvmArtifacts) -> None: + npuir_opt = find_npuir_tool("npuir-opt") + _run_checked( + [ + npuir_opt, + artifacts.edited_mlir, + "--verify-each", + "--suvm-device-validate", + "-o", + artifacts.path(DEVICE_VALIDATED_MLIR), + ], + "NPU-IR device validation", + self.timeout, + ) + _run_checked( + [ + npuir_opt, + artifacts.edited_mlir, + "--verify-each", + "--suvm-to-llvm-pipeline", + "-o", + artifacts.path(LOWERED_MLIR), + ], + "NPU-IR full lowering validation", + self.timeout, + ) + + +def _kernel_params(kernel: tvm.tir.PrimFunc | tvm.IRModule) -> list[KernelParam]: + if isinstance(kernel, tvm.tir.PrimFunc): + function = kernel + else: + functions = [func for func in kernel.functions.values() if isinstance(func, tvm.tir.PrimFunc)] + if len(functions) != 1: + raise ValueError(f"Manual SUVM edit session requires one PrimFunc, got {len(functions)}.") + function = functions[0] + return [ + KernelParam.from_buffer(function.buffer_map[param]) if param in function.buffer_map else KernelParam.from_var(param) + for param in function.params + ] + + +def _prepare_emit(artifacts: SunmmioSuvmArtifacts) -> None: + artifacts.work_dir.mkdir(parents=True, exist_ok=True) + edited = artifacts.edited_mlir + if edited.is_file() or edited.is_symlink(): + archived = _timestamped_archive_path(edited) + edited.replace(archived) + print(f"Archived previous edit: {archived}") + elif edited.exists(): + raise IsADirectoryError(f"Edited SUVM MLIR path is not a file: {edited}") + + for path in artifacts.managed_paths(): + if path.is_file() or path.is_symlink(): + path.unlink() + elif path.exists(): + raise IsADirectoryError(f"Managed artifact path is not a file: {path}") + + +def _timestamped_archive_path(path: Path) -> Path: + timestamp = datetime.now(UTC).strftime("%Y%m%dT%H%M%S.%fZ") + archived = path.with_name(f"{path.stem}.{timestamp}{path.suffix}") + sequence = 1 + while archived.exists(): + archived = path.with_name(f"{path.stem}.{timestamp}.{sequence}{path.suffix}") + sequence += 1 + return archived + + +def _parameter_kinds(manifest: dict[str, Any], abi: SunmmioKernelABI) -> tuple[str, ...]: + parameters = manifest["parameters"] + if len(parameters) != abi.public_arg_count: + raise ValueError("Manual SUVM manifest parameter count does not match ABI metadata.") + kinds = [] + for index, (parameter, name) in enumerate(zip(parameters, abi.public_param_names, strict=True)): + if not isinstance(parameter, dict) or parameter.get("name") != name: + raise ValueError(f"Manual SUVM manifest parameter {index} does not match {name!r}.") + kind = parameter.get("kind") + if kind not in {"tensor", "scalar"}: + raise ValueError(f"Invalid parameter kind at index {index}: {kind!r}.") + kinds.append(kind) + return tuple(kinds) + + +def _validate_sunsim_elf_abi( + artifact: SunmmioKernelArtifact, + abi: SunmmioKernelABI, +) -> None: + """Check that an edited kernel kept the ABI captured during TileLang lowering.""" + from sunsim.notes import ArgumentKind, find_kernel + + try: + metadata = find_kernel(artifact.elf_path, abi.kernel_name) + except KeyError as exc: + raise ValueError(f"Edited kernel changed or removed ABI symbol {abi.kernel_name!r}: {exc}") from exc + if len(metadata.args) != abi.full_arg_count: + raise ValueError(f"Edited kernel ABI has {len(metadata.args)} arguments, but TileLang lowering recorded {abi.full_arg_count}.") + + mismatches = [] + for index, (argument, dtype) in enumerate(zip(metadata.args, abi.device_param_dtypes, strict=True)): + expected = ArgumentKind.GLOBAL_BUFFER if dtype == "handle" else ArgumentKind.BY_VALUE + if argument.kind != expected: + mismatches.append(f"arg {index} ({abi.device_param_names[index]}): expected {expected.name}, got {argument.kind.name}") + if mismatches: + raise ValueError("Edited kernel ABI is incompatible: " + "; ".join(mismatches)) + + +def _write_diff(original: Path, edited: Path, output: Path) -> None: + before = _read_nonempty(original, "original SUVM MLIR").splitlines(keepends=True) + after = _read_nonempty(edited, "edited SUVM MLIR").splitlines(keepends=True) + output.write_text( + "".join(difflib.unified_diff(before, after, fromfile=original.name, tofile=edited.name)), + encoding="utf-8", + ) + + +def _load_manifest(artifacts: SunmmioSuvmArtifacts) -> dict[str, Any]: + path = artifacts.path(MANIFEST_FILE) + manifest = _read_json(path, "manual SUVM manifest") + if manifest.get("schema_version") != MANIFEST_SCHEMA_VERSION: + raise ValueError(f"Unsupported manual SUVM manifest schema in {path}; rerun emit.") + if not isinstance(manifest.get("target"), str) or not isinstance(manifest.get("parameters"), list): + raise ValueError(f"Incomplete manual SUVM manifest: {path}") + return manifest + + +def _read_nonempty(path: Path, description: str) -> str: + try: + content = path.read_text(encoding="utf-8") + except FileNotFoundError: + raise FileNotFoundError(f"{description} does not exist: {path}") from None + if not content.strip(): + raise ValueError(f"{description} is empty: {path}") + return content + + +def _write_json(path: Path, value: dict[str, Any]) -> None: + path.write_text(json.dumps(value, indent=2, sort_keys=True) + "\n", encoding="utf-8") + + +def _read_json(path: Path, description: str) -> dict[str, Any]: + try: + value = json.loads(path.read_text(encoding="utf-8")) + except FileNotFoundError: + raise FileNotFoundError(f"{description} does not exist: {path}") from None + except json.JSONDecodeError as exc: + raise ValueError(f"Invalid {description} {path}: {exc}") from exc + if not isinstance(value, dict): + raise ValueError(f"{description} must contain a JSON object: {path}") + return value + + +def _run_checked( + command: Sequence[str | Path], + description: str, + timeout: float | None, +) -> None: + command_text = [str(part) for part in command] + try: + result = subprocess.run( + command_text, + capture_output=True, + text=True, + check=False, + timeout=timeout, + ) + except subprocess.TimeoutExpired as exc: + raise RuntimeError(f"{description} timed out after {timeout} seconds\ncommand: {' '.join(command_text)}") from exc + if result.returncode != 0: + raise RuntimeError(f"{description} failed\ncommand: {' '.join(command_text)}\nstdout:\n{result.stdout}\nstderr:\n{result.stderr}")