From d7ed3a764743d3cec72f300a6debcd6aaa09e9f2 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Sat, 3 Oct 2026 20:15:10 -0400 Subject: [PATCH] [FEAT] Add the host-compiled IR lifecycle for clients A client can now declare NEEDS_INTERPRETER = False and receive, for every traced launch, the kernels Triton compiles for each config instead of the interpreted run. The compile runs on the host for a configured GPU target, so nothing needs a GPU. Core lifecycle (tilelens/core/client.py, trace.py): - LaunchCall / LaunchEvent; the begin_launch, abort_launch, before_launch, after_launch and compile_failed hooks; each begun launch ends in exactly one of finalize or abort_launch. begin_launch now creates the per-launch Launch; a concurrent launch of one trace from another thread is refused, and so is a second traced launch capturing the same JITFunction. - Client declarations NEEDS_INTERPRETER, IR_STAGES, LAUNCH and ir_target. Interpreter hooks reach interpreting clients only. - ClientManager.ir_capture routes every JITFunction.run call of the traced kernel through a host compile and delivers one event per (specialization, binding fingerprint) per launch; a failing compile is compile_failed data and never fails the launch; a call that does not bind raises as untraced. - TritonTrace compiles every (pruned) config through a separate IR runner chain. An IR-only launch never launches the real kernel; a mixed trace compiles for the IR clients, then interprets for the eager ones. - Each runner chain's Autotuner copies drop the call's arguments (nargs) and restore_value clones (restore_copies) when the call returns or raises, so no caller tensor outlives a launch that raised. Host compile (tilelens/core/host_compile.py): binds and packs a call with the JIT's own binder and compiles it through the backend's TTIR passes for the target, with the front end's target queries answered by that target and every device query refused. IR targets are "cuda:89" by default (TILELENS_IR_TARGET, Client.ir_target); IR mode runs on Triton 3.8 only, checked in one place (config.ir_triton_unsupported). tilelens.ir: IRClient (a finalize template over a per-launch ArtifactLog), LaunchBinding/TensorFacts, a content-addressed ParseCache and the IRVerdict records, which saved traces round-trip. ParseCache's default reader is the TTIR reader, which this change does not add: until it exists, a lookup with the default reader is an error outcome. The tests use stub readers. Not supported yet, refused instead of half-done: - running the real kernel after the host compile: an IR client must declare LAUNCH = "skip"; - stages other than "ttir": IR_STAGES naming another stage is refused when the client is added; - several IR targets in one trace: refused when the launch compiles; - whole-pipeline compiles: TRITON_KERNEL_OVERRIDE, USE_IR_LOC, an ir_override option and a custom pipeline make the host compile raise HostCompileUnavailable. CI installs triton==3.8.0 and triton_kernels at v3.8.0, with UV_NO_SYNC so `uv run` keeps that pin. --- .github/workflows/benchmark.yml | 12 +- .github/workflows/python-app.yml | 16 +- tests/conftest.py | 33 + tests/end_to_end/test_ir_client.py | 274 +++++ tests/unit/ir/__init__.py | 0 tests/unit/ir/test_host_compile.py | 842 ++++++++++++++ tests/unit/ir/test_ir_capture.py | 853 ++++++++++++++ tests/unit/test_ir_lifecycle.py | 1682 ++++++++++++++++++++++++++++ tilelens/core/client.py | 800 ++++++++++++- tilelens/core/config.py | 37 + tilelens/core/host_compile.py | 885 +++++++++++++++ tilelens/core/trace.py | 298 ++++- tilelens/core/trace_io.py | 13 +- tilelens/ir/__init__.py | 41 + tilelens/ir/capture.py | 277 +++++ tilelens/ir/client.py | 124 ++ tilelens/ir/launch.py | 227 ++++ tilelens/ir/verdict.py | 130 +++ 18 files changed, 6460 insertions(+), 84 deletions(-) create mode 100644 tests/end_to_end/test_ir_client.py create mode 100644 tests/unit/ir/__init__.py create mode 100644 tests/unit/ir/test_host_compile.py create mode 100644 tests/unit/ir/test_ir_capture.py create mode 100644 tests/unit/test_ir_lifecycle.py create mode 100644 tilelens/core/host_compile.py create mode 100644 tilelens/ir/__init__.py create mode 100644 tilelens/ir/capture.py create mode 100644 tilelens/ir/client.py create mode 100644 tilelens/ir/launch.py create mode 100644 tilelens/ir/verdict.py diff --git a/.github/workflows/benchmark.yml b/.github/workflows/benchmark.yml index 45d6b1312..4efd0d5c2 100644 --- a/.github/workflows/benchmark.yml +++ b/.github/workflows/benchmark.yml @@ -16,6 +16,10 @@ concurrency: jobs: benchmark: runs-on: ubuntu-latest + env: + # Keep `uv run` from re-syncing the environment to the lockfile, which + # would undo the Triton 3.8 pin installed with `uv pip install`. + UV_NO_SYNC: "1" steps: - name: Checkout PR branch @@ -44,15 +48,15 @@ jobs: run: | cd main-branch uv sync --extra test - uv pip install --pre -U triton - uv pip install "git+https://github.com/triton-lang/triton.git#subdirectory=python/triton_kernels" + uv pip install triton==3.8.0 + uv pip install "git+https://github.com/triton-lang/triton.git@v3.8.0#subdirectory=python/triton_kernels" - name: Install PR dependencies run: | cd pr-branch uv sync --extra test - uv pip install --pre -U triton - uv pip install "git+https://github.com/triton-lang/triton.git#subdirectory=python/triton_kernels" + uv pip install triton==3.8.0 + uv pip install "git+https://github.com/triton-lang/triton.git@v3.8.0#subdirectory=python/triton_kernels" - name: Run interleaved A/B benchmarks run: | diff --git a/.github/workflows/python-app.yml b/.github/workflows/python-app.yml index 1919ccd17..96eba7688 100644 --- a/.github/workflows/python-app.yml +++ b/.github/workflows/python-app.yml @@ -21,6 +21,10 @@ concurrency: jobs: build: runs-on: ubuntu-latest + env: + # Keep `uv run` from re-syncing the environment to the lockfile, which + # would undo the Triton 3.8 pin installed with `uv pip install`. + UV_NO_SYNC: "1" steps: - uses: actions/checkout@v3 @@ -48,11 +52,11 @@ jobs: cd tilelens uv sync --extra test - - name: Upgrade Triton to latest from main + - name: Install Triton 3.8 run: | cd tilelens - uv pip install --pre -U triton - uv pip install "git+https://github.com/triton-lang/triton.git#subdirectory=python/triton_kernels" + uv pip install triton==3.8.0 + uv pip install "git+https://github.com/triton-lang/triton.git@v3.8.0#subdirectory=python/triton_kernels" - name: Run frontend tests run: | @@ -70,11 +74,11 @@ jobs: cd tilelens uv sync --extra test --extra nki - - name: Upgrade Triton to latest from main (NKI) + - name: Install Triton 3.8 (NKI) run: | cd tilelens - uv pip install --pre -U triton - uv pip install "git+https://github.com/triton-lang/triton.git#subdirectory=python/triton_kernels" + uv pip install triton==3.8.0 + uv pip install "git+https://github.com/triton-lang/triton.git@v3.8.0#subdirectory=python/triton_kernels" - name: Run full (Triton + NKI) pytest suite run: | diff --git a/tests/conftest.py b/tests/conftest.py index 53387a160..7ec6b7ddd 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -14,6 +14,39 @@ def pytest_addoption(parser): ) +@pytest.fixture +def unreachable_driver(monkeypatch): + """``unreachable_driver(message)`` makes Triton's active driver + unreachable, as on a machine without a GPU: any question to it raises + ``AssertionError(message)``. One call reaches the real driver: unloading + a module an earlier test loaded on a real GPU, which Triton 3.8's + CompiledKernel.__del__ does through the driver whenever that kernel is + collected (e.g. when tilelens.clear() drops the launch holding it).""" + from triton.runtime.driver import driver + + owner = type(driver) + real = owner.__dict__["active"] + + def refuse(message: str) -> None: + class Utils: + def unload_module(self, module): + return real.__get__(driver, owner).utils.unload_module(module) + + def __getattr__(self, name): + raise AssertionError(message) + + class Active: + utils = Utils() + + def __getattr__(self, name): + raise AssertionError(message) + + stand_in = Active() + monkeypatch.setattr(owner, "active", property(lambda self: stand_in)) + + return refuse + + @pytest.fixture(scope="session", params=["cpu"]) def device(request): return request.param diff --git a/tests/end_to_end/test_ir_client.py b/tests/end_to_end/test_ir_client.py new file mode 100644 index 000000000..ac3e83a66 --- /dev/null +++ b/tests/end_to_end/test_ir_client.py @@ -0,0 +1,274 @@ +"""End-to-end tests of the IR client layer: a toy IRClient under tilelens.trace +on real kernels, compiled on the host for the default target (CPU tensors, +no GPU), gets an ArtifactLog per launch and puts its IRVerdict into +Launch.records. Counterparts on fake events live in +tests/unit/ir/test_ir_capture.py. +""" + +from __future__ import annotations + +import importlib + +import pytest +import torch +import triton +import triton.language as tl +from triton.backends.compiler import GPUTarget +from triton.compiler.errors import CompileTimeAssertionFailure + +import tilelens +from tilelens.core.config import DEFAULT_IR_TARGET +from tilelens.ir import ConfigVerdict, IRClient, IRVerdict, ParseCache, Refusal + +trace_module = importlib.import_module("tilelens.core.trace") +config_module = importlib.import_module("tilelens.core.config") + + +def _real_compiles_available() -> bool: + # Triton imported under TRITON_INTERPRET=1 builds its own standard library + # as InterpretedFunctions, so nothing can compile for real in-process. No + # GPU is needed: IR mode compiles on the host. + import triton.language.standard as tl_standard + from triton.runtime.jit import JITFunction + + return isinstance(tl_standard.cdiv, JITFunction) + + +pytestmark = pytest.mark.skipif( + not _real_compiles_available(), + reason="Triton was imported under TRITON_INTERPRET=1: nothing compiles in-process", +) + + +@pytest.fixture(autouse=True) +def _no_driver(unreachable_driver): + """IR mode needs no GPU: Triton's driver is unreachable here, as on + a machine without one (where it raises "0 active drivers").""" + unreachable_driver("IR mode queried Triton's driver") + + +@pytest.fixture(autouse=True) +def _default_ir_target(monkeypatch): + """The default IR target, whatever TILELENS_IR_TARGET the caller + set: in the process config, and in any Config read from the environment.""" + for name in ("TILELENS_IR_TARGET", "TRITON_VIZ_IR_TARGET"): + monkeypatch.delenv(name, raising=False) + monkeypatch.setattr(config_module.config, "ir_target", DEFAULT_IR_TARGET) + + +@pytest.fixture(autouse=True) +def _real_jit(monkeypatch): + # tests/unit/test_multithreading.py sets TRITON_INTERPRET=1 at import time, + # and a traced launch's patch scope restores knobs.runtime.interpret as an + # explicit override. These tests need @triton.jit to build real + # JITFunctions, so pin the knob off and put back exactly what was there. + from triton import knobs + + monkeypatch.delenv("TRITON_INTERPRET", raising=False) + missing = object() + previous = knobs.runtime.__dict__.get("interpret", missing) + knobs.runtime.__dict__["interpret"] = False + yield + if previous is missing: + knobs.runtime.__dict__.pop("interpret", None) + else: + knobs.runtime.__dict__["interpret"] = previous + + +class _StubRefusal(Exception): + def __init__(self, message, kind): + super().__init__(message) + self.kind = kind + + +class _StubReader: + """Counts the texts it is asked to parse; its graph is the text.""" + + def __init__(self): + self.texts: list[str] = [] + + def __call__(self, text): + self.texts.append(text) + return text + + +class _ToyIR(IRClient): + """Parses each specialization's TTIR through a ParseCache; one + ConfigVerdict per compiled or failed config.""" + + NAME = "toy_ir" + LAUNCH = "skip" + IR_STAGES = frozenset({"ttir"}) + + def __init__(self, reader=None): + super().__init__() + self.parses = ( + ParseCache() if reader is None else ParseCache(reader, refusal=_StubRefusal) + ) + self.logs: list[tuple] = [] + self.outcomes: list = [] + + def analyze_launch(self, log): + self.logs.append((log.specializations, log.failures)) + per_config = [] + for spec in log.specializations: + outcome = self.parses.get(spec.artifacts.stages["ttir"]) + self.outcomes.append(outcome) + if outcome.refusal is not None: + refusal = Refusal.from_exception(outcome.refusal) + per_config.append( + ConfigVerdict(spec.specialization, spec.config, "refused", refusal) + ) + else: + status = "parsed" if outcome.error is None else "error" + per_config.append( + ConfigVerdict(spec.specialization, spec.config, status) + ) + for failure in log.failures: + per_config.append(ConfigVerdict(None, failure.config, "compile-failed")) + return [], IRVerdict(self.NAME, "ok", per_config=per_config) + + def on_analysis_error(self, exc): + return IRVerdict(self.NAME, "error", notes=[f"{type(exc).__name__}: {exc}"]) + + +def _make_add_one(): + @triton.jit + def add_one(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + return add_one + + +def _make_autotuned(*blocks): + @triton.autotune( + configs=[triton.Config({"BLOCK": b}, num_warps=1) for b in blocks], + key=["n"], + ) + @triton.heuristics({"EVEN": lambda args: args["n"] % args["BLOCK"] == 0}) + @triton.jit + def add_one_tuned(x_ptr, out_ptr, n, BLOCK: tl.constexpr, EVEN: tl.constexpr): + tl.static_assert(BLOCK <= 32) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + return add_one_tuned + + +def _grid(meta): + return (triton.cdiv(meta["n"], meta["BLOCK"]),) + + +def _inputs(n=64): + x = torch.arange(n, dtype=torch.float32) + return x, torch.zeros_like(x) + + +def test_launch_records_carry_the_ir_verdict(): + reader = _StubReader() + ir = _ToyIR(reader) + traced = tilelens.trace(ir)(_make_add_one()) + x, out = _inputs() + + kernel = traced[(4,)](x, out, 64, BLOCK=16) + + # LAUNCH="skip": compiled and analyzed, never launched. + assert torch.equal(out, torch.zeros_like(x)) + launch = trace_module.launches[-1] + assert launch.records == [ir.last_verdict] + verdict = ir.last_verdict + assert verdict == IRVerdict( + "toy_ir", + "ok", + per_config=(ConfigVerdict(kernel.hash, {}, "parsed"),), + ) + (text,) = reader.texts + assert text == kernel.asm["ttir"] and "tt.func" in text + + ((spec,), failures) = ir.logs[0] + assert failures == () + assert spec.artifacts.stages.keys() == {"ttir"} + meta = spec.artifacts.meta + assert (meta["backend"], meta["name"], meta["num_warps"]) == ("cuda", "add_one", 4) + # Compiled for the default target, through TTIR only: no shared-memory + # size yet. + assert (meta["arch"], meta["shared"]) == (89, None) # the default target + (binding,) = spec.bindings + assert binding.error is None + assert binding.tensors.keys() == {"x_ptr", "out_ptr"} + facts = binding.tensors["x_ptr"] + assert (facts.data_ptr, facts.numel, facts.elem_size) == (x.data_ptr(), 64, 4) + assert (facts.shape, facts.strides, facts.dtype) == ((64,), (1,), "torch.float32") + assert facts.allocation_interval() == (x.data_ptr(), x.data_ptr() + 256) + assert dict(binding.params) == {"n": 64} + assert dict(binding.constexprs) == {"BLOCK": 16} + assert binding.grid == (4, 1, 1) + + +def test_autotune_gives_one_config_verdict_per_config(): + reader = _StubReader() + ir = _ToyIR(reader) + traced = tilelens.trace(ir)(_make_autotuned(16, 32)) + x, out = _inputs() + + traced[_grid](x, out, 64) + first = ir.last_verdict + traced[_grid](x, out, 64) + + for verdict in (first, ir.last_verdict): + assert [c.status for c in verdict.per_config] == ["parsed", "parsed"] + configs = [c.config for c in verdict.per_config] + assert [(c["BLOCK"], c["num_warps"], c["EVEN"]) for c in configs] == [ + (16, 1, True), + (32, 1, True), + ] + assert len({c.specialization for c in verdict.per_config}) == 2 + # The second launch finds both TTIR texts in the parse cache. + assert len(reader.texts) == 2 + assert [launch.records for launch in trace_module.launches[-2:]] == [ + [first], + [ir.last_verdict], + ] + + +def test_a_config_that_fails_to_compile_is_recorded(): + ir = _ToyIR(_StubReader()) + # BLOCK=64 trips the kernel's static_assert. + traced = tilelens.trace(ir)(_make_autotuned(16, 64)) + x, out = _inputs() + + traced[_grid](x, out, 64) + + parsed, failed = ir.last_verdict.per_config + assert (parsed.status, parsed.config["BLOCK"]) == ("parsed", 16) + assert (failed.status, failed.specialization) == ("compile-failed", None) + assert (failed.config["BLOCK"], failed.config["num_warps"]) == (64, 1) + ((spec,), (failure,)) = ir.logs[0] + assert spec.config["BLOCK"] == 16 + assert isinstance(failure.error, CompileTimeAssertionFailure) + assert failure.target == GPUTarget("cuda", 89, 32) # the default + + +def test_another_triton_release_compiles_nothing_and_goes_on(monkeypatch): + """IR mode runs on Triton 3.8 only: on any other release the host + compile is unavailable, so the launch's config fails to compile, saying + why, and the program goes on.""" + from tilelens.core.host_compile import host_compile_unavailable, triton_api + + reader = _StubReader() + ir = _ToyIR(reader) + traced = tilelens.trace(ir)(_make_add_one()) + x, out = _inputs() + triton_api.cache_clear() # a refusal is not cached + monkeypatch.setattr(triton, "__version__", "3.6.0") + + assert traced[(4,)](x, out, 64, BLOCK=16) is None + + (config,) = ir.last_verdict.per_config + assert config.status == "compile-failed" and reader.texts == [] + ((), (failure,)) = ir.logs[0] + unavailable = host_compile_unavailable(failure.error) + assert "supports Triton 3.8.x only" in str(unavailable) diff --git a/tests/unit/ir/__init__.py b/tests/unit/ir/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/ir/test_host_compile.py b/tests/unit/ir/test_host_compile.py new file mode 100644 index 000000000..e5ad6f091 --- /dev/null +++ b/tests/unit/ir/test_host_compile.py @@ -0,0 +1,842 @@ +"""tilelens.core.host_compile: IR targets and the host compile. + +CPU only, and no driver: every compile here runs with Triton's driver made +unreachable (``no_driver``), as on a machine without a GPU. Kernels are +built inside the tests so their JITFunctions are real even when +TRITON_INTERPRET was set during collection. +""" + +from __future__ import annotations + +import importlib +import re + +import pytest +import torch +import triton +import triton.language as tl +from triton.backends.compiler import GPUTarget +from triton.compiler.errors import CompileTimeAssertionFailure + +from tilelens.core.config import DEFAULT_IR_TARGET, Config +from tilelens.core.host_compile import ( + HostCompileUnavailable, + HostCompiler, + HostKernel, + default_ir_target, + format_ir_target, + parse_ir_target, + resolve_ir_target, + target_queried, + triton_api, +) + +config_module = importlib.import_module("tilelens.core.config") + + +def _real_compiles_available() -> bool: + # Triton imported under TRITON_INTERPRET=1 builds its own standard library + # as InterpretedFunctions, so nothing can compile for real in-process. + import triton.language.standard as tl_standard + from triton.runtime.jit import JITFunction + + return isinstance(tl_standard.cdiv, JITFunction) + + +needs_compiles = pytest.mark.skipif( + not _real_compiles_available(), + reason="Triton was imported under TRITON_INTERPRET=1: nothing compiles in-process", +) + + +@pytest.fixture(autouse=True) +def _real_jit(monkeypatch): + # tests/unit/test_multithreading.py sets TRITON_INTERPRET=1 at import + # time; pin the knob off so @triton.jit builds real JITFunctions. + from triton import knobs + + monkeypatch.delenv("TRITON_INTERPRET", raising=False) + missing = object() + previous = knobs.runtime.__dict__.get("interpret", missing) + knobs.runtime.__dict__["interpret"] = False + yield + if previous is missing: + knobs.runtime.__dict__.pop("interpret", None) + else: + knobs.runtime.__dict__["interpret"] = previous + + +@pytest.fixture +def no_driver(unreachable_driver): + """Make Triton's active driver unreachable, as without a GPU (where it + raises "0 active drivers"): any driver query fails the test.""" + unreachable_driver("the host compile queried Triton's driver") + + +CUDA80 = GPUTarget("cuda", 80, 32) +# The default IR target (sm89: the first that compiles fp8e4nv). +CUDA89 = GPUTarget("cuda", 89, 32) + + +# ======== targets ========= + + +@pytest.mark.parametrize( + "spec, target", + [ + ("cuda:80", CUDA80), + ("cuda:90", GPUTarget("cuda", 90, 32)), + (" CUDA:120 ", GPUTarget("cuda", 120, 32)), + ("cuda:80:64", GPUTarget("cuda", 80, 64)), + ("hip:gfx942", GPUTarget("hip", "gfx942", 64)), + ("hip:gfx90a", GPUTarget("hip", "gfx90a", 64)), + ("hip:gfx1100", GPUTarget("hip", "gfx1100", 32)), + ("hip:gfx1100:64", GPUTarget("hip", "gfx1100", 64)), + (GPUTarget("hip", "gfx950", 64), GPUTarget("hip", "gfx950", 64)), + ], +) +def test_parse_ir_target_reads_the_documented_forms(spec, target): + assert parse_ir_target(spec) == target + assert parse_ir_target(format_ir_target(target)) == target + + +@pytest.mark.parametrize( + "spec", + [ + "", + "cuda", + "cuda:", + "cuda:sm80", + "sm80", + "80", + "cuda:80:", + "rocm:gfx942", + "hip:942", + "hip:gfx942:x", + 80, + None, + ("cuda", 80, 32), + GPUTarget("cuda", "80", 32), + GPUTarget("hip", 942, 64), + GPUTarget("cpu", "x86", 1), + GPUTarget("cuda", 80, 0), + # A warp size is positive in either form, a capability an int >= 70 + # (a bool is no int here), a gfx arch gfx. + "cuda:80:0", + "hip:gfx942:0", + "cuda:0", + "cuda:60", + "hip:gfx9", + GPUTarget("cuda", True, 32), + GPUTarget("cuda", 80, True), + GPUTarget("cuda", 60, 32), + GPUTarget("hip", "gfx9", 64), + ], +) +def test_parse_ir_target_rejects_what_names_no_target(spec): + with pytest.raises(ValueError, match="invalid IR target .*expected 'cuda:"): + parse_ir_target(spec) + + +def test_format_ir_target_names_the_warp_size_only_when_it_is_not_the_default(): + assert format_ir_target(CUDA80) == "cuda:80" + assert format_ir_target(GPUTarget("cuda", 80, 64)) == "cuda:80:64" + assert format_ir_target(GPUTarget("hip", "gfx942", 64)) == "hip:gfx942" + assert format_ir_target(GPUTarget("hip", "gfx1100", 64)) == "hip:gfx1100:64" + + +def test_the_default_target_is_cuda89(): + assert DEFAULT_IR_TARGET == "cuda:89" + assert default_ir_target() == CUDA89 + + +def test_resolve_takes_the_clients_target_else_the_configured_one(monkeypatch): + monkeypatch.setattr(config_module.config, "ir_target", DEFAULT_IR_TARGET) + assert resolve_ir_target("cuda:90") == GPUTarget("cuda", 90, 32) + assert resolve_ir_target(None) == CUDA89 + monkeypatch.setattr(config_module.config, "ir_target", "hip:gfx942") + assert resolve_ir_target(None) == GPUTarget("hip", "gfx942", 64) + # A client's own target wins over the configured one. + assert resolve_ir_target(CUDA80) == CUDA80 + monkeypatch.setattr(config_module.config, "ir_target", "gfx942") + with pytest.raises(ValueError, match=r"TILELENS_IR_TARGET\) is 'gfx942'"): + resolve_ir_target(None) + with pytest.raises(ValueError, match="invalid IR target 'gfx942'"): + resolve_ir_target("gfx942") + + +@pytest.mark.parametrize( + "env, expected", + [ + ({}, "cuda:89"), + ({"TILELENS_IR_TARGET": "cuda:90"}, "cuda:90"), + ], +) +def test_the_configured_target_comes_from_the_environment(monkeypatch, env, expected): + monkeypatch.delenv("TILELENS_IR_TARGET", raising=False) + monkeypatch.delenv("TRITON_VIZ_IR_TARGET", raising=False) + for name, value in env.items(): + monkeypatch.setenv(name, value) + assert Config().ir_target == expected + + +# ======== the host compile ========= + + +def _make_scalars(): + @triton.jit + def scalars(x_ptr, n, flag, scale, none_arg, pair, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + pair[0] + n + vals = tl.load(x_ptr + offs, mask=offs < pair[1]) * scale + if flag: + tl.store(x_ptr + offs, vals, mask=offs < pair[1]) + + return scalars + + +def _make_copy(): + @triton.jit + def copy(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask), mask=mask) + + return copy + + +def _signature(ttir: str) -> dict[str, str]: + """The TTIR entry function's arguments: name -> type.""" + header = re.search(r"tt\.func public @\w+\((.*?)\) attributes", ttir, re.S) + assert header is not None, ttir + return dict(re.findall(r"%([\w.]+): ([^\s{,)]+)", header.group(1))) + + +@needs_compiles +@pytest.mark.parametrize( + "value, ttir_type", + [ + (2**31 - 1, "i32"), + (2**31, "i64"), + (-(2**31), "i32"), + (-(2**31) - 1, "i64"), + (2**32, "i64"), + (2**63, "i64"), # u64 in the JIT's signature; TTIR integers are signless + (16, "i32"), + (1, None), # the equal-to-1 specialization: a constexpr, no argument + ], +) +def test_integers_are_typed_by_value_as_the_jit_types_them(no_driver, value, ttir_type): + kernel = HostCompiler().compile( + _make_copy(), + (torch.zeros(64), torch.zeros(64), value), + {"BLOCK": 16}, + target=CUDA80, + ) + assert _signature(kernel.asm["ttir"]).get("n") == ttir_type + + +@needs_compiles +def test_a_host_compile_stops_after_ttir(no_driver): + x = torch.zeros(64) + kernel = HostCompiler().compile( + _make_scalars(), + (x, 5, True, 1.5, None, (3, 4)), + {"BLOCK": 16, "num_warps": 2}, + target=CUDA80, + ) + assert isinstance(kernel, HostKernel) + assert list(kernel.asm) == ["ttir"] + assert kernel.name == kernel.metadata.name == "scalars" + assert kernel.target == kernel.metadata.target == CUDA80 + assert (kernel.metadata.num_warps, kernel.metadata.hash) == (2, kernel.hash) + # Nothing after TTIR ran: no shared-memory size, no binary. + assert not hasattr(kernel.metadata, "shared") + # bool i1, float f32, a tuple one argument per item; None is a constexpr. + # the tuple's items are named by their path in it + first, second = "pair.0", "pair.1" + assert _signature(kernel.asm["ttir"]) == { + "x_ptr": "!tt.ptr", + "n": "i32", + "flag": "i1", + "scale": "f32", + first: "i32", + second: "i32", + } + # Divisibility by 16 is specialized as the JIT does. + assert "tt.divisibility = 16" in kernel.asm["ttir"].split("attributes")[0] + + +@needs_compiles +def test_compiles_are_cached_per_call_and_target(no_driver): + compiler = HostCompiler() + copy = _make_copy() + x, out = torch.zeros(64), torch.zeros(64) + + def compile(n=64, target=CUDA80, **kwargs): + return compiler.compile( + copy, (x, out, n), {"BLOCK": 16, **kwargs}, target=target + ) + + first = compile() + assert compile() is first + # Another tensor with the same specialization is the same kernel. + assert compiler.compile(copy, (torch.zeros(8), out, 64), {"BLOCK": 16}, target=CUDA80) is first # fmt: skip + assert compile(n=80) is first # 80 % 16 == 0: specialized alike + assert compile(n=65) is not first + assert compile(num_warps=8) is not first + hip = compile(target=GPUTarget("hip", "gfx942", 64)) + assert hip is not first and hip.hash != first.hash + assert hip.target == GPUTarget("hip", "gfx942", 64) + + +@needs_compiles +def test_targets_compile_their_own_ttir(no_driver): + """A tensor descriptor is rewritten to pointers below sm90 only, as the + target's backend does (the target decides, not the machine).""" + tensor_descriptor = pytest.importorskip("triton.tools.tensor_descriptor") + + @triton.jit + def bump(desc, BLOCK: tl.constexpr): + desc.store([0, 0], desc.load([0, 0]) + 1) + + desc = tensor_descriptor.TensorDescriptor.from_tensor(torch.zeros(64, 64), [16, 16]) + compiler = HostCompiler() + sm80, sm90 = ( + compiler.compile(bump, (desc,), {"BLOCK": 16}, target=parse_ir_target(t)) + for t in ("cuda:80", "cuda:90") + ) + assert "tt.descriptor_load" not in sm80.asm["ttir"] + assert "tt.descriptor_load" in sm90.asm["ttir"] + + +@needs_compiles +def test_compile_errors_are_the_jits(no_driver): + @triton.jit + def bounded(x_ptr, BLOCK: tl.constexpr): + tl.static_assert(BLOCK <= 32) + tl.store(x_ptr + tl.arange(0, BLOCK), 1.0) + + compiler = HostCompiler() + x = torch.zeros(64) + with pytest.raises(CompileTimeAssertionFailure): + compiler.compile(bounded, (x,), {"BLOCK": 64}, target=CUDA80) + with pytest.raises(KeyError, match="unrecognised"): + compiler.compile(bounded, (x,), {"BLOCK": 16, "bogus": 1}, target=CUDA80) + with pytest.raises(TypeError): + compiler.compile(bounded, (), {"BLOCK": 16}, target=CUDA80) + # A target-specific option check (num_ctas > 1 needs sm90). + with pytest.raises(ValueError, match="num_ctas"): + compiler.compile(bounded, (x,), {"BLOCK": 16, "num_ctas": 2}, target=CUDA80) + assert ( + compiler.compile( + bounded, + (x,), + {"BLOCK": 16, "num_ctas": 2}, + target=parse_ir_target("cuda:90"), + ).metadata.num_ctas + == 2 + ) + + +_SCALE = tl.constexpr(2) + + +@needs_compiles +def test_a_changed_global_is_refused_like_the_jit_refuses_it(no_driver, monkeypatch): + @triton.jit + def scaled(x_ptr, BLOCK: tl.constexpr): + tl.store(x_ptr + tl.arange(0, BLOCK) * _SCALE, 1.0) + + compiler = HostCompiler() + call = ((torch.zeros(64),), {"BLOCK": 16}) + compiler.compile(scaled, *call, target=CUDA80) + monkeypatch.setitem(globals(), "_SCALE", tl.constexpr(3)) + # The cached kernel read the old value: stale, not handed out. + with pytest.raises(RuntimeError, match="_SCALE has changed since we compiled"): + compiler.compile(scaled, *call, target=CUDA80) + + +def test_what_is_no_jit_function_cannot_be_host_compiled(): + with pytest.raises(HostCompileUnavailable, match="has no 'signature'"): + HostCompiler().compile(object(), (), {}, target=CUDA80) + + +def _clear_api_caches(): + from tilelens.core import host_compile + + triton_api.cache_clear() + host_compile._self_test_target.cache_clear() + + +def test_another_triton_release_is_refused(monkeypatch): + _clear_api_caches() + monkeypatch.setattr(triton, "__version__", "3.6.0") + try: + with pytest.raises( + HostCompileUnavailable, match=r"Triton 3\.8\.x only; .* is 3\.6\.0" + ): + triton_api() + finally: + monkeypatch.undo() + _clear_api_caches() + + +@needs_compiles +@pytest.mark.parametrize( + "knob, named", + [ + ("override", "TRITON_KERNEL_OVERRIDE"), + ("use_ir_loc", "USE_IR_LOC"), + ("ir_override", "'ir_override' compile option"), + ("custom_pipeline", "a custom pipeline"), + ], +) +def test_what_only_the_whole_pipeline_applies_is_refused( + no_driver, monkeypatch, knob, named +): + """A knob under which triton.compile changes the TTIR, or keys the + kernel apart, in stages a host compile does not run: refused as + HostCompileUnavailable, never a TTIR the device would not run.""" + from triton import knobs + + call = (_make_copy(), (torch.zeros(64), torch.zeros(64), 64)) + options = {"BLOCK": 16} + with monkeypatch.context() as patch: + if knob == "override": + patch.setattr(knobs.compilation, "override", True) + elif knob == "use_ir_loc": + patch.setattr(knobs.compilation, "use_ir_loc", "ttir") + elif knob == "ir_override": + options["ir_override"] = "kernel.ttir" + else: + patch.setattr( + knobs.runtime, "add_stages_inspection_hook", lambda *args: None + ) + with pytest.raises(HostCompileUnavailable, match=re.escape(named)): + HostCompiler().compile(*call, options, target=CUDA80) + # Without it, the same call compiles. + assert HostCompiler().compile(*call, {"BLOCK": 16}, target=CUDA80) + + +# ======== the target the front end sees ========= + + +class _Machine: + """A stand-in for Triton's active driver on a machine with a GPU of + ``target``, counting the target queries it answers.""" + + def __init__(self, target): + self.target = target + self.queries = 0 + + def get_current_target(self): + self.queries += 1 + return self.target + + def get_current_device(self): + return 0 + + def get_current_stream(self, device=None): + return 0 + + +def _on_machine(monkeypatch, machine): + """Make ``machine`` Triton's active driver; None: no GPU (Triton then + raises "0 active drivers", which tl.target_info reads as no target).""" + from triton.runtime.driver import driver + + def active(self): + if machine is None: + raise RuntimeError("0 active drivers ([]). There should only be one.") + return machine + + monkeypatch.setattr(type(driver), "active", property(active)) + + +def _make_target_branches(): + @triton.jit + def branches(x_ptr, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + if tl.target_info.is_cuda(): + tl.store(x_ptr + offs, 1.0) + if tl.target_info.cuda_capability_geq(8, 9): + tl.store(x_ptr + offs, 2.0) + if tl.target_info.is_hip(): + tl.store(x_ptr + offs, 3.0) + + return branches + + +def _stored(ttir: str) -> set[float]: + return { + float(v) + for v in re.findall( + r"arith\.constant dense<([-+.e0-9]+)> : tensor<16xf32>", ttir + ) + } + + +@needs_compiles +@pytest.mark.parametrize( + "spec, stored", + [ + ("cuda:80", {1.0}), + ("cuda:89", {1.0, 2.0}), + ("cuda:90", {1.0, 2.0}), + ("hip:gfx942", {3.0}), + ], +) +def test_the_front_end_asks_the_compiles_target(no_driver, spec, stored): + """tl.target_info reads Triton's driver; a host compile answers it with + the compile's target (never the machine's, and here there is none).""" + kernel = HostCompiler().compile( + _make_target_branches(), + (torch.zeros(16),), + {"BLOCK": 16}, + target=parse_ir_target(spec), + ) + assert _stored(kernel.asm["ttir"]) == stored + + +@needs_compiles +def test_the_ttir_does_not_depend_on_the_machine(monkeypatch): + """Without a GPU, or on any GPU, a target's TTIR (and hash) is the same: + the machine's driver is never asked for its target.""" + machines = [ + None, + _Machine(GPUTarget("cuda", 89, 32)), + _Machine(GPUTarget("cuda", 120, 32)), + _Machine(GPUTarget("hip", "gfx942", 64)), + ] + targets = [parse_ir_target(t) for t in ("cuda:80", "cuda:90", "hip:gfx942")] + seen: dict = {} + for machine in machines: + _on_machine(monkeypatch, machine) + kernel = _make_target_branches() + for target in targets: + compiled = HostCompiler().compile( + kernel, + (torch.zeros(16),), + {"BLOCK": 16}, + target=target, + ) + seen.setdefault(target, set()).add((compiled.hash, compiled.asm["ttir"])) + assert machine is None or machine.queries == 0 + assert all(len(compiles) == 1 for compiles in seen.values()), seen + + +@needs_compiles +def test_native_tma_is_the_targets(no_driver): + """The semantic's native-TMA check (a 16-bit descriptor atomic_min) + reads the compile's target too: fine for cuda:90, refused for cuda:80 + as on an sm80 device.""" + from triton.compiler.errors import CompilationError + + tensor_descriptor = pytest.importorskip("triton.tools.tensor_descriptor") + + @triton.jit + def shrink(desc, BLOCK: tl.constexpr): + desc.atomic_min([0, 0], desc.load([0, 0])) + + desc = tensor_descriptor.TensorDescriptor.from_tensor( + torch.zeros(64, 64, dtype=torch.float16), [16, 16] + ) + compiler = HostCompiler() + sm90 = compiler.compile( + shrink, (desc,), {"BLOCK": 16}, target=parse_ir_target("cuda:90") + ) + assert "tt.descriptor_reduce" in sm90.asm["ttir"] + with pytest.raises(CompilationError, match="native tma") as raised: + compiler.compile(shrink, (desc,), {"BLOCK": 16}, target=CUDA80) + # The front end asked for the target, and the answer is what failed. + assert target_queried(raised.value) + + +@needs_compiles +def test_a_compile_error_says_whether_the_front_end_asked_for_the_target(no_driver): + """target_queried marks a kernel's compile error when Triton's front end + had asked for the compile's target before it (here tl.target_info in a + static_assert). A failure no target query decides is not marked, even + one the target's compile options decide (num_ctas > 1 below sm90).""" + + @triton.jit + def bounded(x_ptr, BLOCK: tl.constexpr): + tl.static_assert(BLOCK <= 32) + tl.store(x_ptr + tl.arange(0, BLOCK), 1.0) + + @triton.jit + def hopper_only(x_ptr, BLOCK: tl.constexpr): + tl.static_assert(tl.target_info.cuda_capability_geq(9, 0)) + tl.store(x_ptr + tl.arange(0, BLOCK), 1.0) + + def error(kernel, block, target, **options): + with pytest.raises(Exception) as raised: + HostCompiler().compile( + kernel, + (torch.zeros(64),), + {"BLOCK": block, **options}, + target=target, + ) + return raised.value + + too_big = error(bounded, 64, CUDA89) + assert isinstance(too_big, CompileTimeAssertionFailure) + assert not target_queried(too_big) + for_hopper = error(hopper_only, 16, CUDA89) + assert isinstance(for_hopper, CompileTimeAssertionFailure) + assert target_queried(for_hopper) + two_ctas = error(bounded, 16, CUDA89, num_ctas=2) + assert isinstance(two_ctas, ValueError) and "num_ctas > 1" in str(two_ctas) + assert not target_queried(two_ctas) + # Each compile answers for itself: the same kernel compiles for sm90. + HostCompiler().compile( + hopper_only, + (torch.zeros(64),), + {"BLOCK": 16}, + target=parse_ir_target("cuda:90"), + ) + # No host compile raised these. + assert not target_queried(ValueError("x")) and not target_queried(None) + + +@needs_compiles +def test_a_device_query_the_kernel_catches_still_counts(no_driver): + """The host compile refuses a device query; the kernel's code may catch + that and fall back to an answer of its own ("no big shared memory"), + which a GPU might not give: a compile error after it is marked as one + that asked, like one after the target query.""" + from triton.runtime.jit import constexpr_function + + @constexpr_function + def has_big_smem(): + from triton.runtime import driver + + try: + properties = driver.active.utils.get_device_properties(0) + except Exception: + return False + return properties["max_shared_mem"] >= 200_000 + + @triton.jit + def big_smem_only(x_ptr, BLOCK: tl.constexpr): + tl.static_assert(has_big_smem()) + tl.store(x_ptr + tl.arange(0, BLOCK), 1.0) + + with pytest.raises(CompileTimeAssertionFailure) as raised: + HostCompiler().compile( + big_smem_only, + (torch.zeros(16),), + {"BLOCK": 16}, + target=CUDA89, + ) + assert target_queried(raised.value) + + +@needs_compiles +def test_a_kernel_that_keeps_its_target_answer_is_marked_after_it_asked(no_driver): + """A kernel's code may keep the target answer (a memo) and ask only on + its first compile: a later compile of the same kernel for the same + target, failing on the kept answer, is marked too. Another kernel that + never asked is not.""" + from triton.runtime.jit import constexpr_function + + @constexpr_function + def is_hopper(_memo={}): # noqa: B006 the kernel's own memo + if "arch" not in _memo: + from triton.runtime import driver + + _memo["arch"] = driver.active.get_current_target().arch + return _memo["arch"] >= 90 + + @triton.jit + def hopper_only(x_ptr, HOPPER: tl.constexpr, BLOCK: tl.constexpr): + if HOPPER: + tl.static_assert(is_hopper()) + else: + tl.static_assert(is_hopper() or True) + tl.store(x_ptr + tl.arange(0, BLOCK), 1.0) + + @triton.jit + def bounded(x_ptr, BLOCK: tl.constexpr): + tl.static_assert(BLOCK <= 32) + tl.store(x_ptr + tl.arange(0, BLOCK), 1.0) + + compiler = HostCompiler() + x = torch.zeros(64) + # Asks for the target, and keeps the answer. + compiler.compile( + hopper_only, + (x,), + {"HOPPER": False, "BLOCK": 16}, + target=CUDA89, + ) + assert is_hopper.fn.__defaults__[0] == {"arch": 89} + with pytest.raises(CompileTimeAssertionFailure) as raised: + compiler.compile( + hopper_only, + (x,), + {"HOPPER": True, "BLOCK": 16}, + target=CUDA89, + ) + assert target_queried(raised.value) + with pytest.raises(CompileTimeAssertionFailure) as raised: + compiler.compile(bounded, (x,), {"BLOCK": 64}, target=CUDA89) + assert not target_queried(raised.value) + + +@needs_compiles +def test_a_call_that_does_not_bind_is_marked_as_such(no_driver): + """A call the JIT's binder rejects (a missing argument) fails on any + device, whatever the target answers: bind_failed marks it, and it is + never marked as having asked for the target, even for a kernel whose + earlier compile asked. An option the target's backend does not know + fails later, in the compile, and is no bind failure (another backend + may know it).""" + from tilelens.core.host_compile import bind_failed + + @triton.jit + def on_cuda(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + if tl.target_info.is_cuda(): + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + compiler = HostCompiler() + x = torch.zeros(16) + compiler.compile(on_cuda, (x, 16), {"BLOCK": 16}, target=CUDA89) + with pytest.raises(TypeError, match="required positional argument: 'n'") as raised: + compiler.compile(on_cuda, (x,), {"BLOCK": 16}, target=CUDA89) + assert bind_failed(raised.value) and not target_queried(raised.value) + with pytest.raises(KeyError, match="waves_per_eu") as raised: + compiler.compile( + on_cuda, (x, 16), {"BLOCK": 16, "waves_per_eu": 2}, target=CUDA89 + ) + assert not bind_failed(raised.value) + assert not bind_failed(ValueError("x")) and not bind_failed(None) + + +@needs_compiles +def test_a_call_the_jit_cannot_key_is_marked_as_a_bind_failure(no_driver): + """JITFunction.run keys the call right after binding it + (compute_cache_key: the bound specialization and the call's options), + on any device: an unhashable constexpr value fails there, whatever the + target, so it is the call's own error too, raised as untraced.""" + from tilelens.core.host_compile import bind_failed, unknown_options + + kernel = _make_copy() + x = torch.zeros(16) + for target in (CUDA89, GPUTarget("hip", "gfx942", 64)): + with pytest.raises(TypeError, match="unhashable type: 'list'") as raised: + HostCompiler().compile(kernel, (x, x, 16), {"BLOCK": [16]}, target=target) + assert bind_failed(raised.value) and not unknown_options(raised.value) + + +@needs_compiles +def test_an_option_no_backend_of_the_target_knows_is_named(no_driver): + """The JIT's KeyError for a keyword that is neither a parameter nor an + option of the target's backend says which keywords (unknown_options); + it stays no bind failure: another backend may know them.""" + from tilelens.core.host_compile import bind_failed, unknown_options + + kernel = _make_copy() + x = torch.zeros(16) + hip = GPUTarget("hip", "gfx942", 64) + cases = [ + (CUDA89, {"bogus": 1}, ("bogus",)), + (CUDA89, {"waves_per_eu": 2}, ("waves_per_eu",)), + (hip, {"maxnreg": 64}, ("maxnreg",)), + (hip, {"bogus": 1, "maxnreg": 64}, ("bogus", "maxnreg")), + ] + for target, options, names in cases: + with pytest.raises(KeyError, match="unrecognised") as raised: + HostCompiler().compile( + kernel, (x, x, 16), {"BLOCK": 16, **options}, target=target + ) + assert unknown_options(raised.value) == names + assert not bind_failed(raised.value) + HostCompiler().compile( + kernel, (x, x, 16), {"BLOCK": 16, "waves_per_eu": 2}, target=hip + ) + assert unknown_options(KeyError("x")) == () and unknown_options(None) == () + + +def _device_is_zero(): + # A host function a constexpr function may call from a kernel (marked + # like tl.target_info.current_target) that asks Triton's driver for the + # device, as a compile never should on the host. + return triton.runtime.driver.active.get_current_device() == 0 + + +_device_is_zero.__triton_builtin__ = True # type: ignore[attr-defined] + + +@needs_compiles +def test_a_device_query_while_compiling_is_refused(no_driver): + from triton.compiler.errors import CompilationError + from triton.runtime.jit import constexpr_function + + from tilelens.core.host_compile import host_compile_unavailable + + @constexpr_function + def on_device_zero(): + return _device_is_zero() + + @triton.jit + def device_dependent(x_ptr, BLOCK: tl.constexpr): + if on_device_zero(): + tl.store(x_ptr + tl.arange(0, BLOCK), 1.0) + + # Triton's code generator re-raises it as the kernel's CompilationError. + with pytest.raises(CompilationError) as raised: + HostCompiler().compile( + device_dependent, (torch.zeros(16),), {"BLOCK": 16}, target=CUDA80 + ) + unavailable = host_compile_unavailable(raised.value) + assert isinstance(unavailable, HostCompileUnavailable) + assert "'get_current_device'" in str(unavailable) + # A kernel's own compile error has none behind it, nor has one raised + # "from None" while handling it. + assert host_compile_unavailable(CompilationError("src", None, "bad")) is None + try: + try: + raise HostCompileUnavailable("no device") + except HostCompileUnavailable: + raise KeyError("the kernel's") from None + except KeyError as exc: + assert host_compile_unavailable(exc) is None + + +@needs_compiles +def test_override_arch_does_not_reach_a_host_compile(no_driver, monkeypatch): + """TRITON_OVERRIDE_ARCH retargets the JIT; a host compile stays for the + target it was asked for, a hip one included.""" + from triton.compiler.errors import CompilationError + + monkeypatch.setenv("TRITON_OVERRIDE_ARCH", "sm90") + + @triton.jit + def to_fp8(x_ptr, out_ptr, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs).to(tl.float8e4nv).to(tl.float32)) + + compiler = HostCompiler() + copy = _make_copy() + call = ((torch.zeros(64), torch.zeros(64), 64), {"BLOCK": 16}) + assert compiler.compile(copy, *call, target=CUDA80).metadata.arch == "sm80" + # sm80's rules: no num_ctas > 1, no fp8e4nv. + with pytest.raises(ValueError, match="num_ctas > 1 requires NVIDIA SM90"): + compiler.compile(copy, call[0], {**call[1], "num_ctas": 2}, target=CUDA80) + with pytest.raises(CompilationError, match="fp8e4nv not supported"): + compiler.compile(to_fp8, call[0][:2], {"BLOCK": 16}, target=CUDA80) + hip = compiler.compile(copy, *call, target=parse_ir_target("hip:gfx942")) + assert hip.metadata.arch == "gfx942" + + # Where "arch" is the kernel's own parameter the arch cannot be pinned, + # and a compile for another arch is refused, not mislabeled. + @triton.jit + def with_arch(x_ptr, arch, BLOCK: tl.constexpr): + tl.store(x_ptr + tl.arange(0, BLOCK), arch) + + with pytest.raises(HostCompileUnavailable, match="arch 'sm90', not 'sm80'"): + compiler.compile( + with_arch, (torch.zeros(16), 1.0), {"BLOCK": 16}, target=CUDA80 + ) diff --git a/tests/unit/ir/test_ir_capture.py b/tests/unit/ir/test_ir_capture.py new file mode 100644 index 000000000..c3ac0992e --- /dev/null +++ b/tests/unit/ir/test_ir_capture.py @@ -0,0 +1,853 @@ +"""The IR client layer (tilelens.ir.{launch,capture,verdict,client}) on fake +LaunchEvents: launch binding, the per-launch artifact log, the parse cache, +verdict records and the IRClient finalize template. No GPU; the real-kernel +counterparts live in tests/end_to_end/test_ir_client.py. +""" + +from __future__ import annotations + +import dataclasses +import enum +import gc +import importlib +import pickle +import subprocess +import sys +import types +import weakref +from pathlib import Path +from types import SimpleNamespace + +import numpy as np +import pytest +import torch +import triton +import triton.language as tl + +import tilelens +from tilelens.core.client import ClientManager, LaunchCall +from tilelens.core.data import Launch +from tilelens.ir import ( + ArtifactLog, + CompileFailure, + ConfigVerdict, + IRClient, + IRVerdict, + ParseCache, + ParseOutcome, + Refusal, + SourceLocation, + TensorFacts, + bind_launch, +) + +trace_module = importlib.import_module("tilelens.core.trace") +REPO = Path(__file__).resolve().parents[3] + + +# ======== fakes ========= + + +class _Refusal(Exception): + """Stands in for the TTIR reader's UnsupportedTTIR.""" + + def __init__(self, message, kind, line_no=None, loc=None): + super().__init__(message) + self.message = message + self.kind = kind + self.line_no = line_no + self.loc = loc + + +class _FakeKernel: + def __init__(self, key, *, asm=None, metadata=True): + self.hash = f"hash-{key}" + self.asm = ( + {"ttir": f"// ttir {key}", "ttgir": f"// ttgir {key}", "cubin": b"\x7fELF"} + if asm is None + else asm + ) + if metadata: + self.metadata = SimpleNamespace( + target=SimpleNamespace(backend="cuda", arch=89, warp_size=32), + num_warps=4, + num_stages=3, + shared=512, + name=f"kernel_{key}", + hash=self.hash, + ) + + +@triton.jit +def _kernel(x_ptr, out_ptr, n, flag, scale, BLOCK: tl.constexpr, EVEN: tl.constexpr): + pass + + +@triton.jit +def _tuple_kernel(ptrs, n): + pass + + +class _Descriptor: + """A descriptor-style argument: the kernel addresses its .base tensor.""" + + def __init__(self, base): + self.base = base + + +def _grid(meta): + return (triton.cdiv(meta["n"], meta["BLOCK"]),) + + +def _event(args, kwargs, *, grid=(1,), kernel=None, error=None, target=None): + # The core's own event builder: bound_args and resolved_grid as + # ir_capture computes them. + return ClientManager._launch_event( + _kernel, args, kwargs, grid, kernel, error=error, target=target + ) + + +def _call(**kwargs): + return LaunchCall(jit_fn=_kernel, args=(), kwargs=kwargs, grid=None, capture=True) + + +def _tensors(): + x = torch.arange(64, dtype=torch.float32) + out = torch.zeros(64, dtype=torch.float32) + return x, out + + +class _ToyIR(IRClient): + NAME = "toy_ir" + LAUNCH = "skip" + IR_STAGES = frozenset({"ttir"}) + + def __init__(self, analyze=None): + super().__init__() + self.calls: list = [] + self._analyze = analyze + + def analyze_launch(self, log): + self.calls.append(("analyze", log.call, log.specializations, log.failures)) + if self._analyze is not None: + return self._analyze(log) + per_config = [ + ConfigVerdict(spec.specialization, spec.config, "seen") + for spec in log.specializations + ] + return ["report"], IRVerdict(self.NAME, "ok", per_config=per_config) + + def on_analysis_error(self, exc): + self.calls.append(("error", exc)) + return IRVerdict(self.NAME, "error", notes=[repr(exc)]) + + +# ======== launch binding ========= + + +def test_tensor_facts_read_the_view_and_its_storage(): + base = torch.arange(12, dtype=torch.float32).reshape(3, 4) + view = base[1:, 1:] + binding = bind_launch( + _event((view, base, 1, False, 0.0), {"BLOCK": 1, "EVEN": True}) + ) + facts = binding.tensors["x_ptr"] + + assert facts == TensorFacts( + data_ptr=base.data_ptr() + 5 * 4, + elem_size=4, + numel=6, + shape=(2, 3), + strides=(4, 1), + dtype="torch.float32", + contiguous=False, + storage_data_ptr=base.data_ptr(), + storage_nbytes=48, + ) + # A strided view's allocation is its storage, not numel * elem_size. + assert facts.allocation_interval() == (base.data_ptr(), base.data_ptr() + 48) + + +_FACTS = dict( + data_ptr=1024, + elem_size=4, + numel=8, + shape=(8,), + strides=(1,), + dtype="torch.float32", + contiguous=True, +) + + +@pytest.mark.parametrize( + "overrides, interval", + [ + # Without storage metadata only a contiguous view's extent is known. + ({}, (1024, 1056)), + ({"contiguous": False}, None), + # Partial or inconsistent storage metadata never falls back to numel. + ({"storage_data_ptr": 1024}, None), + ({"storage_data_ptr": 2048, "storage_nbytes": 64}, None), + ({"storage_data_ptr": 1024, "storage_nbytes": 16}, None), + ({"storage_data_ptr": 1000, "storage_nbytes": 100}, (1000, 1100)), + ({"elem_size": 0}, None), + ], +) +def test_allocation_interval_refuses_unknown_extents(overrides, interval): + assert TensorFacts(**{**_FACTS, **overrides}).allocation_interval() == interval + + +def test_bind_launch_splits_arguments_by_kind(): + x, out = _tensors() + # The caller passed BLOCK; a Heuristics layer added EVEN and an + # Autotuner config num_warps. + event = _event( + (x, _Descriptor(out), 64, True, 0.5), + {"BLOCK": 16, "EVEN": True, "num_warps": 4}, + grid=_grid, + ) + binding = bind_launch(event, _call(BLOCK=16)) + + assert binding.error is None + assert dict(binding.params) == {"n": 64, "flag": 1} + assert not isinstance(binding.params["flag"], bool) + assert binding.tensors.keys() == {"x_ptr", "out_ptr"} + assert binding.tensors["out_ptr"].data_ptr == out.data_ptr() + assert binding.tensors["x_ptr"].numel == 64 + # Floats are no binding fact; constexprs keep their values. + assert "scale" not in binding.params + assert dict(binding.constexprs) == {"BLOCK": 16, "EVEN": True} + assert dict(binding.config) == {"EVEN": True, "num_warps": 4} + assert binding.raw_grid is _grid + assert binding.grid == (4, 1, 1) + with pytest.raises(TypeError): + binding.params["n"] = 1 # type: ignore[index] + + +def test_bind_launch_without_the_call_counts_every_kwarg_as_config(): + x, out = _tensors() + binding = bind_launch( + _event((x, out, 64, False, 0.5), {"BLOCK": 16, "EVEN": False}) + ) + assert dict(binding.config) == {"BLOCK": 16, "EVEN": False} + # A heuristic overriding a caller kwarg with another value is config too. + heuristic = bind_launch( + _event((x, out, 64, False, 0.5), {"BLOCK": 32, "EVEN": False}), + _call(BLOCK=16, EVEN=False), + ) + assert dict(heuristic.config) == {"BLOCK": 32} + + +def test_config_kwargs_tell_the_callers_scalars_by_value(): + x, out = _tensors() + big = 10**6 + recomputed = int(str(big)) # equal, but another object + assert recomputed is not big + step = torch.tensor(1) + binding = bind_launch( + _event( + (x, out, 64, False, 0.5), + {"BLOCK": recomputed, "EVEN": True, "step": torch.tensor(1), "mode": 1}, + ), + _call(BLOCK=big, EVEN=True, step=step, mode=True), + ) + # An equal plain scalar of the same type is the caller's; another object + # of any other kind, or a value of another type, is config. + assert binding.config.keys() == {"step", "mode"} + + +def test_bind_launch_leaves_out_tuple_arguments(): + x, out = _tensors() + event = ClientManager._launch_event(_tuple_kernel, ((x, out), 64), {}, (1,), None) + binding = bind_launch(event) + # Two TTIR pointer arguments, but no binding fact and no error: a + # consumer must treat them as unknown (see LaunchBinding). + assert event.bound_args["ptrs"] == (x, out) + assert dict(binding.tensors) == {} + assert dict(binding.params) == {"n": 64} + assert binding.error is None + + +def test_bind_launch_never_raises(): + x, out = _tensors() + + class _Unreadable: + def data_ptr(self): + return 0 + + def element_size(self): + return 4 + + def numel(self): + raise RuntimeError("no numel") + + binding = bind_launch( + _event((_Unreadable(), out, 64, False, 0.5), {"BLOCK": 4, "EVEN": True}) + ) + assert binding.error == "argument 'x_ptr': RuntimeError: no numel" + assert binding.tensors.keys() == {"out_ptr"} + assert dict(binding.params) == {"n": 64, "flag": 0} + + # An unresolvable grid is no error; a non-integer resolved grid is. + unresolved = bind_launch( + _event((x, out, 64, False, 0.5), {"BLOCK": 4, "EVEN": True}, grid=None) + ) + assert unresolved.grid is None and unresolved.error is None + broken = dataclasses.replace( + _event((x, out, 64, False, 0.5), {"BLOCK": 4, "EVEN": True}), + resolved_grid=("wide", 1, 1), + ) + assert bind_launch(broken).grid is None + assert "grid ('wide', 1, 1): TypeError" in bind_launch(broken).error + + # Not an event at all: everything unreadable, still a binding. + nothing = bind_launch(SimpleNamespace(jit_fn=None)) # type: ignore[arg-type] + assert nothing.error is not None and not nothing.tensors + + +@pytest.mark.parametrize( + "resolved, grid", + [ + ((np.int64(3), torch.tensor(2), 1), (3, 2, 1)), + # The untraced launch rejects a float grid; it is not truncated. + ((2.7, 1, 1), None), + ], +) +def test_bind_launch_converts_grid_dims_as_the_launcher_does(resolved, grid): + x, out = _tensors() + event = dataclasses.replace( + _event((x, out, 64, False, 0.5), {"BLOCK": 4, "EVEN": True}), + resolved_grid=resolved, + ) + binding = bind_launch(event) + assert binding.grid == grid + assert (binding.error is None) == (grid is not None) + if grid is None: + assert "TypeError: 'float' object cannot be interpreted" in binding.error + + +# ======== artifact log ========= + + +def test_artifact_log_keeps_declared_stages_meta_and_bindings(): + x, out = _tensors() + log = ArtifactLog({"ttir", "ptx"}) + call = _call() + log.reset(call) + kernel_a, kernel_b = _FakeKernel("a"), _FakeKernel("b") + # Config A, config B, then A again with another grid (another binding). + log.record( + _event((x, out, 64, False, 0.5), {"BLOCK": 4, "EVEN": True}, kernel=kernel_a) + ) + log.record( + _event((x, out, 64, False, 0.5), {"BLOCK": 8, "EVEN": True}, kernel=kernel_b) + ) + log.record( + _event( + (x, out, 64, False, 0.5), + {"BLOCK": 4, "EVEN": True}, + grid=(2,), + kernel=kernel_a, + ) + ) + + assert log.call is call + spec_a, spec_b = log.specializations + assert (spec_a.specialization, spec_b.specialization) == ("hash-a", "hash-b") + assert len(spec_a.bindings) == 2 and len(spec_b.bindings) == 1 + # Only the declared stages the kernel has: no ttgir, no ptx. + assert dict(spec_a.artifacts.stages) == {"ttir": "// ttir a"} + assert dict(spec_a.artifacts.meta) == { + "backend": "cuda", + "arch": 89, + "num_warps": 4, + "num_stages": 3, + "shared": 512, + "name": "kernel_a", + "config": {"BLOCK": 4, "EVEN": True}, + } + assert spec_a.config == {"BLOCK": 4, "EVEN": True} + assert spec_b.config == {"BLOCK": 8, "EVEN": True} + assert spec_a.artifacts.error is None + assert log.failures == () + + +def test_artifact_log_records_compile_failures_with_their_config(): + from triton.backends.compiler import GPUTarget + + x, out = _tensors() + log = ArtifactLog({"ttir"}) + log.reset(_call()) + compile_error = RuntimeError("static_assert failed") + option_error = ValueError("num_ctas > 1 requires NVIDIA SM90+") + target = GPUTarget("cuda", 80, 32) + args = (x, out, 64, False, 0.5) + log.record_failure( + _event(args, {"BLOCK": 64, "EVEN": True}, error=compile_error, target=target) + ) + # An event built outside the core names no target. + log.record_failure(_event(args, {"BLOCK": 8, "EVEN": True}, error=option_error)) + + assert log.failures == ( + CompileFailure(compile_error, {"BLOCK": 64, "EVEN": True}, target, _kernel), + CompileFailure(option_error, {"BLOCK": 8, "EVEN": True}, None, _kernel), + ) + assert log.specializations == () + + +def test_artifact_log_contains_unreadable_kernels(): + x, out = _tensors() + + class _BrokenAsm(_FakeKernel): + @property + def asm(self): + raise RuntimeError("asm gone") + + @asm.setter + def asm(self, value): + pass + + log = ArtifactLog({"ttir", "sass"}) + log.reset(_call()) + args = (x, out, 64, False, 0.5) + log.record(_event(args, {"BLOCK": 4, "EVEN": True}, kernel=_BrokenAsm("a"))) + log.record( + _event( + args, + {"BLOCK": 8, "EVEN": True}, + kernel=_FakeKernel("b", asm={"cubin": b""}, metadata=False), + ) + ) + + broken, bare = log.specializations + assert broken.artifacts.error == "asm: RuntimeError: asm gone" + assert broken.artifacts.meta["name"] == "kernel_a" + # A declared stage the kernel lacks is simply absent; no metadata at all + # is an error. + assert dict(bare.artifacts.stages) == {} + assert bare.artifacts.error.startswith("metadata: AttributeError") + assert bare.artifacts.meta["num_warps"] is None + + +def test_artifact_log_reset_forgets_the_launch(): + x, out = _tensors() + log = ArtifactLog({"ttir"}) + log.reset(_call()) + args = (x, out, 64, False, 0.5) + log.record(_event(args, {"BLOCK": 4, "EVEN": True}, kernel=_FakeKernel("a"))) + log.record_failure(_event(args, {"BLOCK": 64, "EVEN": True}, error=RuntimeError())) + log.reset() + assert (log.call, log.specializations, log.failures) == (None, (), ()) + + +# ======== parse cache ========= + + +class _StubReader: + def __init__(self, result=None): + self.calls: list = [] + self.result = result + + def __call__(self, text, **options): + self.calls.append((text, options)) + if isinstance(self.result, BaseException): + raise self.result + return ("graph", text, tuple(sorted(options.items()))) + + +def test_parse_cache_parses_each_text_once_per_options(): + reader = _StubReader() + cache = ParseCache(reader, refusal=_Refusal) + + first = cache.get("module a") + assert first == ParseOutcome(graph=("graph", "module a", ())) + assert cache.get("module a") is first + cache.get("module b") + cache.get("module a", keep_going=True) + cache.get("module a", keep_going=True) + assert reader.calls == [ + ("module a", {}), + ("module b", {}), + ("module a", {"keep_going": True}), + ] + + +def test_parse_cache_keeps_refusals_with_their_kind(): + def refuse(): + raise _Refusal( + "scf.while", "control-flow", line_no=7, loc=SourceLocation("k.py", 3, 1) + ) + + try: + refuse() + except _Refusal as exc: + refusal = exc + assert refusal.__traceback__ is not None + reader = _StubReader(refusal) + cache = ParseCache(reader, refusal=_Refusal) + + outcome = cache.get("module") + assert outcome.graph is None and outcome.error is None + assert outcome.refusal is refusal and refusal.kind == "control-flow" + # A cached refusal keeps no frames alive. + assert refusal.__traceback__ is None + assert cache.get("module") is outcome + assert len(reader.calls) == 1 + assert Refusal.from_exception(outcome.refusal) == Refusal( + "control-flow", "scf.while", 7, SourceLocation("k.py", 3, 1) + ) + + +class _Held: + """A reader-frame local whose lifetime a test watches.""" + + +def test_a_cached_refusal_keeps_no_frame_of_its_chain_alive(): + held = [] + + def reader(text): + local = _Held() + held.append(weakref.ref(local)) + try: + raise KeyError("walk") + except KeyError: + # A cause next to the implicit context: both chain links hold + # a traceback into this frame. + raise _Refusal("scf.while", "control-flow") from ValueError("cause") + + cache = ParseCache(reader, refusal=_Refusal) + try: + raise LookupError("the caller's") + except LookupError as caller: + refusal = cache.get("module").refusal + # The exception the caller was handling is the caller's: untouched, + # and no longer chained to the cached refusal. + assert caller.__traceback__ is not None + assert isinstance(refusal.__cause__, ValueError) + walk = refusal.__context__ + assert isinstance(walk, KeyError) and walk.__context__ is None + assert all(e.__traceback__ is None for e in (refusal, refusal.__cause__, walk)) + gc.collect() + assert held[0]() is None + + +def test_parse_cache_reports_other_errors_without_caching_them(): + reader = _StubReader(RecursionError("too deep")) + cache = ParseCache(reader, refusal=_Refusal) + + assert cache.get("module") == ParseOutcome(error="RecursionError: too deep") + assert cache.get("module").error == "RecursionError: too deep" + assert len(reader.calls) == 2 + # Without a refusal type, a refusal-shaped exception is just an error. + plain = ParseCache(_StubReader(_Refusal("x", "control-flow")), refusal=KeyError) + assert plain.get("module").error == "_Refusal: x" + + +def test_parse_cache_never_raises(): + reader = _StubReader() + cache = ParseCache(reader) + # A lone surrogate hashes (as "?") and parses. + assert cache.get("module \ud800").graph == ("graph", "module \ud800", ()) + assert cache.get(None).error.startswith("AttributeError") # type: ignore[arg-type] + assert cache.get("module", layout=[1]).error.startswith("TypeError") + # Every option name reaches the reader, "text" and "self" included. + assert cache.get("module", text=1).error.startswith("TypeError: _StubReader") + options = ParseCache(lambda text, /, **options: options, refusal=_Refusal) + assert options.get("module", text=1, self=2).graph == {"text": 1, "self": 2} + + +@pytest.fixture +def default_reader(monkeypatch): + """The module ParseCache resolves its default reader from: the real + tilelens.ir.ttir_reader when it exists (its parse_ttir replaced), else a + stand-in with the same two names.""" + try: + module = importlib.import_module("tilelens.ir.ttir_reader") + except ModuleNotFoundError as exc: + if exc.name != "tilelens.ir.ttir_reader": + raise + module = types.ModuleType("tilelens.ir.ttir_reader") + module.UnsupportedTTIR = _Refusal # type: ignore[attr-defined] + monkeypatch.setitem(sys.modules, "tilelens.ir.ttir_reader", module) + readers = [] + + def install(result=None): + reader = _StubReader(result) + readers.append(reader) + monkeypatch.setattr(module, "parse_ttir", reader, raising=False) + return reader + + install.module = module # type: ignore[attr-defined] + return install + + +def test_parse_cache_without_a_reader_module_is_an_error_outcome(monkeypatch): + """While tilelens.ir.ttir_reader cannot be imported, a lookup with the + default reader reports that as an error outcome, never raises, and + caches nothing.""" + monkeypatch.setitem(sys.modules, "tilelens.ir.ttir_reader", None) + cache = ParseCache() + + outcome = cache.get("module") + + assert outcome.graph is None and outcome.refusal is None + assert outcome.error.startswith("ModuleNotFoundError: ") + assert "tilelens.ir.ttir_reader" in outcome.error + assert cache.get("module") is not outcome # not cached: retried + + +def test_parse_cache_resolves_its_default_reader_at_each_lookup(default_reader): + cache = ParseCache() + first = default_reader() + cache.get("module") + cache.get("module") + # A replaced reader is another reader: parsed again, keyed apart. + second = default_reader() + cache.get("module") + assert (len(first.calls), len(second.calls)) == (1, 1) + + # The default refusal type is the reader module's UnsupportedTTIR. + unsupported = default_reader.module.UnsupportedTTIR + refusal = unsupported(kind="control-flow", message="no") + default_reader(refusal) + assert cache.get("refused").refusal is refusal + + +# ======== verdict records ========= + + +def test_verdicts_are_plain_frozen_picklable_records(): + config = {"BLOCK": 16} + per_config = [ConfigVerdict("hash-a", config, "proved", n_reports=0)] + verdict = IRVerdict( + "toy_ir", + "ok", + scope="launch", + refusal=Refusal("control-flow", "scf.while", 3, SourceLocation("k.py", 1, 1)), + per_config=per_config, + notes=["note"], + ) + + assert verdict.per_config == (ConfigVerdict("hash-a", {"BLOCK": 16}, "proved"),) + assert verdict.notes == ("note",) + config["BLOCK"] = 32 # the verdict holds its own copy + assert verdict.per_config[0].config == {"BLOCK": 16} + with pytest.raises(dataclasses.FrozenInstanceError): + verdict.status = "races" # type: ignore[misc] + assert pickle.loads(pickle.dumps(verdict)) == verdict + + +def test_verdicts_are_not_hashable_and_take_no_bare_strings(): + # Frozen, but a config dict has no hash: no verdict claims one. + with pytest.raises(TypeError, match="unhashable"): + hash(ConfigVerdict("hash-a", {"BLOCK": 1}, "ok")) + with pytest.raises(TypeError, match="unhashable"): + hash(IRVerdict("toy_ir", "ok")) + # A str is a sequence, but never the notes or configs meant. + with pytest.raises(TypeError, match="notes takes a sequence"): + IRVerdict("toy_ir", "ok", notes="solver timed out") + with pytest.raises(TypeError, match="per_config takes a sequence"): + IRVerdict("toy_ir", "ok", per_config="hash-a") # type: ignore[arg-type] + + +class _Kind(str, enum.Enum): + CONTROL_FLOW = "control-flow" + + +def test_verdicts_round_trip_through_a_saved_trace(tmp_path, monkeypatch): + # Every verdict field is a value a trace can hold (trace_io + # registers the tilelens.ir.verdict records). + refusal = Refusal(_Kind.CONTROL_FLOW, "scf.while", 3, SourceLocation("k.py", 1, 1)) + verdict = IRVerdict( + "toy_ir", + "unsupported", + scope="launch", + refusal=refusal, + per_config=[ + ConfigVerdict("hash-a", {"BLOCK": 16, "num_warps": 4}, "proved"), + ConfigVerdict(None, {"BLOCK": 64}, "refused", refusal, n_reports=2), + ], + notes=["note"], + ) + saved = [Launch(grid=(4, 1, 1), records=["report", verdict])] + monkeypatch.setattr(trace_module, "launches", saved) + + tilelens.save(tmp_path / "trace.tvz") + (launch,) = tilelens.load(tmp_path / "trace.tvz") + + assert launch.records == ["report", verdict] + # A str-valued kind enum is saved as its string. + kind = launch.records[1].refusal.kind + assert isinstance(kind, str) and not isinstance(kind, enum.Enum) + + +def test_refusal_from_exception_reads_the_structured_fields(): + loc = SourceLocation("k.py", 2) + assert Refusal.from_exception(_Refusal("m", "call", 4, loc)) == Refusal( + "call", "m", 4, loc + ) + + class _Bare(Exception): + kind = "inline-asm" + + assert Refusal.from_exception(_Bare("impure")) == Refusal("inline-asm", "impure") + + +# ======== IRClient ========= + + +def test_ir_client_is_abstract_and_inert(): + class _Partial(IRClient): + NAME = "partial" + + def analyze_launch(self, log): + return [], IRVerdict(self.NAME, "ok") + + with pytest.raises(TypeError, match="abstract"): + _Partial() # type: ignore[abstract] + + ir = _ToyIR() + assert ir.NEEDS_INTERPRETER is False + assert ir.artifacts.stages == {"ttir"} + assert ir.last_verdict is None + # The interpreter path does nothing, and the warmup vote declines. + assert ir.pre_warmup_callback(_kernel) is False + assert ir.pre_run_callback(_kernel) is False + assert ir.post_run_callback(_kernel) is False + ops = ir.register_op_callback(object) # type: ignore[arg-type] + assert (ops.before_callback, ops.after_callback, ops.op_overrider) == (None,) * 3 + loops = ir.register_for_loop_callback() + assert all(getattr(loops, f.name) is None for f in dataclasses.fields(loops)) + manager = ClientManager([ir]) + assert manager.ir_clients() == [ir] and manager.interpreting_clients() == [] + + +def _run_launch(manager, ir, *, events=(), failures=()): + call = _call() + manager.begin_launch(call) + for event in events: + ir.before_launch(event) + for event in failures: + ir.compile_failed(event) + manager.finalize() + return call + + +def test_finalize_returns_the_reports_then_the_verdict(): + x, out = _tensors() + ir = _ToyIR() + manager = ClientManager([ir]) + args = (x, out, 64, False, 0.5) + events = [ + _event(args, {"BLOCK": 4, "EVEN": True}, kernel=_FakeKernel("a")), + _event(args, {"BLOCK": 8, "EVEN": True}, kernel=_FakeKernel("b")), + ] + failure = _event(args, {"BLOCK": 64, "EVEN": True}, error=RuntimeError("bad")) + call = _run_launch(manager, ir, events=events, failures=[failure]) + + ((_, seen_call, specs, failures),) = ir.calls + assert seen_call is call + assert [s.specialization for s in specs] == ["hash-a", "hash-b"] + assert [f.config for f in failures] == [{"BLOCK": 64, "EVEN": True}] + verdict = manager.launch.records[-1] + assert manager.launch.records == ["report", verdict] + assert verdict is ir.last_verdict + assert [c.config for c in verdict.per_config] == [ + {"BLOCK": 4, "EVEN": True}, + {"BLOCK": 8, "EVEN": True}, + ] + # The log is released once the launch is finalized. + assert (ir.artifacts.call, ir.artifacts.specializations) == (None, ()) + + +def test_a_launch_with_nothing_captured_is_the_subclass_call(): + # TRITON_INTERPRET / an InterpretedFunction runner / Gluon / NKI: no + # JITFunction, so the log stays empty, and only log.call says why. + def analyze(log): + if not log.call.capture: + refusal = Refusal("no-capture", "no compiled kernel to read") + return [], IRVerdict(_ToyIR.NAME, "unsupported", refusal=refusal) + return [], IRVerdict(_ToyIR.NAME, "ok") + + ir = _ToyIR(analyze) + manager = ClientManager([ir]) + call = LaunchCall(jit_fn=None, args=(), kwargs={}, grid=(4,), capture=False) + manager.begin_launch(call) + manager.finalize() + + assert ir.calls == [("analyze", call, (), ())] + assert manager.launch.records == [ir.last_verdict] + assert ir.last_verdict.refusal.kind == "no-capture" + + +def test_analysis_exceptions_go_to_the_client_handler(): + boom = RuntimeError("solver crashed") + + def analyze(log): + raise boom + + ir = _ToyIR(analyze) + manager = ClientManager([ir]) + _run_launch(manager, ir) + + assert ir.calls[-1] == ("error", boom) + assert manager.launch.records == [ir.last_verdict] + assert ir.last_verdict.status == "error" + + +def test_an_exiting_analysis_propagates_and_releases_the_log(): + x, out = _tensors() + + def analyze(log): + raise SystemExit(1) # e.g. abort_on_error + + ir = _ToyIR(analyze) + manager = ClientManager([ir]) + manager.begin_launch(_call()) + ir.before_launch( + _event( + (x, out, 64, False, 0.5), + {"BLOCK": 4, "EVEN": True}, + kernel=_FakeKernel("a"), + ) + ) + with pytest.raises(SystemExit): + manager.finalize() + assert ir.last_verdict is None + assert ir.artifacts.specializations == () + + +def test_each_launch_starts_from_a_clean_log_and_no_verdict(): + x, out = _tensors() + ir = _ToyIR() + manager = ClientManager([ir]) + args = (x, out, 64, False, 0.5) + _run_launch( + manager, + ir, + events=[_event(args, {"BLOCK": 4, "EVEN": True}, kernel=_FakeKernel("a"))], + ) + assert ir.last_verdict is not None + + # An aborted launch leaves no verdict and nothing recorded behind. + call = _call() + manager.begin_launch(call) + assert ir.last_verdict is None and ir.artifacts.call is call + ir.before_launch(_event(args, {"BLOCK": 8, "EVEN": True}, kernel=_FakeKernel("b"))) + manager.abort_launch(RuntimeError("launch failed")) + assert ir.artifacts.specializations == () + + _run_launch(manager, ir) + assert ir.calls[-1][2] == () + + +def test_importing_the_ir_layer_imports_no_triton(): + code = ( + "import sys\n" + "import tilelens.ir as ir\n" + "import tilelens.ir.capture, tilelens.ir.launch, tilelens.ir.verdict\n" + "for name in ir.__all__:\n" + " getattr(ir, name)\n" + "assert 'triton' not in sys.modules, sorted(m for m in sys.modules if 'triton' in m)\n" + ) + subprocess.run([sys.executable, "-c", code], check=True, cwd=REPO) diff --git a/tests/unit/test_ir_lifecycle.py b/tests/unit/test_ir_lifecycle.py new file mode 100644 index 000000000..682678fa4 --- /dev/null +++ b/tests/unit/test_ir_lifecycle.py @@ -0,0 +1,1682 @@ +"""CPU-only tests of the core IR lifecycle: client declarations, ClientManager +dispatch rules, the ``ir_capture`` run wrapper and TritonTrace's runner handling. + +The core compiles IR kernels on the host (tilelens.core.host_compile). +These tests pin call sequences, so a fake stands in for that compile: a +``_FakeJit`` compiles through its ``fake_compile``, and a real JITFunction +gets one from ``fake_compile(jit_fn)`` (``_install_fake_run``); a host +compile nothing faked fails the test. The real host compile is tested in +tests/unit/ir/test_host_compile.py, and end to end under a toy IR client in +tests/end_to_end/test_ir_client.py. +""" +import ast +import gc +import importlib +import inspect +import re +import threading +import weakref +from types import SimpleNamespace + +import pytest +import torch +import triton +import triton.language as tl +from triton.compiler.errors import CompileTimeAssertionFailure +from triton.runtime.interpreter import InterpretedFunction + +import tilelens +from tilelens.clients import Sanitizer, Tracer +from tilelens.core.callbacks import ForLoopCallbacks, OpCallbacks +from tilelens.core.client import ( + Client, + ClientManager, + LanguagePatchedError, + LaunchCall, + LaunchEvent, + _resolve_grid, +) +from tilelens.core.config import DEFAULT_IR_TARGET, config as tilelens_config +from tilelens.core.data import Store +from tilelens.core.frontend.base import LANG_PATCH_SCOPES, get_frontend +from tilelens.core.host_compile import HostCompiler, default_ir_target +from tilelens.core.trace import ( + GluonTrace, + NKITrace, + TraceInterface, + TritonTrace, +) + +# `tilelens.core.trace` the attribute is the trace() decorator; the module +# holds the `launches` list. +trace_module = importlib.import_module("tilelens.core.trace") + + +@pytest.fixture(autouse=True) +def _real_jit(monkeypatch): + # tests/unit/test_multithreading.py sets TRITON_INTERPRET=1 at import time, + # and a traced launch's patch scope restores knobs.runtime.interpret as an + # explicit override. These tests need @triton.jit to build real + # JITFunctions, so pin the knob off and put back exactly what was there. + from triton import knobs + + monkeypatch.delenv("TRITON_INTERPRET", raising=False) + missing = object() + previous = knobs.runtime.__dict__.get("interpret", missing) + knobs.runtime.__dict__["interpret"] = False + yield + if previous is missing: + knobs.runtime.__dict__.pop("interpret", None) + else: + knobs.runtime.__dict__["interpret"] = previous + + +@pytest.fixture(autouse=True) +def _default_ir_target(monkeypatch): + # Whatever TILELENS_IR_TARGET the caller has set. + monkeypatch.setattr(tilelens_config, "ir_target", DEFAULT_IR_TARGET) + + +@pytest.fixture(autouse=True) +def _fake_host_compile(monkeypatch): + """Route the core's host compile to the jit_fn's ``fake_compile`` (see + the module docstring); every compile is recorded in ``compiles`` as + (jit_fn, target).""" + compiles: list[tuple] = [] + + def compile(self, jit_fn, args, kwargs, *, target): + fake = getattr(jit_fn, "fake_compile", None) + assert fake is not None, f"unexpected host compile of {jit_fn!r}" + compiles.append((jit_fn, target)) + return fake(*args, **kwargs) + + monkeypatch.setattr(HostCompiler, "compile", compile) + return compiles + + +# ======== Fake clients ========= + + +class _EagerClient(Client): + """Interpreting client that records every callback it receives.""" + + NAME = "eager" + + def __init__(self, *, warmup_vote=False, loop_overrider=None, records=()): + super().__init__() + self.calls: list = [] + self.stores = 0 + self.warmup_vote = warmup_vote + self.loop_overrider = loop_overrider + self.records = list(records) + self.on_store = self._on_store + + def _on_store(self, *args, **kwargs): + self.stores += 1 + + def pre_run_callback(self, fn): + self.calls.append("pre_run") + return True + + def post_run_callback(self, fn): + self.calls.append("post_run") + return True + + def arg_callback(self, name, arg, arg_cvt): + self.calls.append(("arg", name)) + + def grid_callback(self, grid): + self.calls.append(("grid", grid)) + + def grid_idx_callback(self, grid_idx): + self.calls.append("grid_idx") + + def register_op_callback(self, op_type, *args, **kwargs): + if op_type is Store: + return OpCallbacks(before_callback=self.on_store) + return OpCallbacks() + + def register_for_loop_callback(self): + return ForLoopCallbacks(loop_iter_overrider=self.loop_overrider) + + def finalize(self): + self.calls.append("finalize") + return list(self.records) + + def pre_warmup_callback(self, jit_fn, *args, **kwargs): + self.calls.append("pre_warmup") + return self.warmup_vote + + def post_warmup_callback(self, jit_fn, ret): + self.calls.append(("post_warmup", ret)) + + def begin_launch(self, call): + self.calls.append("begin") + + def abort_launch(self, exc): + self.calls.append(("abort", type(exc))) + + def before_launch(self, event): + self.calls.append("before_launch") + + +class _SiblingEagerClient(Client): + """A second interpreting client class, unrelated to _EagerClient.""" + + NAME = "sibling_eager" + + def __init__(self, loop_overrider=None): + super().__init__() + self.loop_overrider = loop_overrider + + def pre_run_callback(self, fn): + return True + + def post_run_callback(self, fn): + return True + + def arg_callback(self, name, arg, arg_cvt): + pass + + def grid_callback(self, grid): + pass + + def grid_idx_callback(self, grid_idx): + pass + + def register_op_callback(self, op_type, *args, **kwargs): + return OpCallbacks() + + def register_for_loop_callback(self): + return ForLoopCallbacks(loop_iter_overrider=self.loop_overrider) + + def finalize(self): + return [] + + def pre_warmup_callback(self, jit_fn, *args, **kwargs): + return False + + def post_warmup_callback(self, jit_fn, ret): + pass + + +class _IRClient(Client): + """IR client: records lifecycle hooks; interpreter hooks must never fire.""" + + NEEDS_INTERPRETER = False + IR_STAGES = frozenset({"ttir"}) + LAUNCH = "skip" + + def __init__(self, log=None, *, raise_in_before=None, records=()): + super().__init__() + self.log = [] if log is None else log + self.events: list[LaunchEvent] = [] + self.failures: list[LaunchEvent] = [] + self.finalized: list[list[LaunchEvent]] = [] + self.launch_calls: list[LaunchCall] = [] + self.raise_in_before = raise_in_before + self.records = list(records) + + def begin_launch(self, call): + self.log.append("begin") + self.launch_calls.append(call) + self.events = [] + self.failures = [] + + def abort_launch(self, exc): + self.log.append(("abort", type(exc))) + + def before_launch(self, event): + self.log.append("before") + if self.raise_in_before is not None: + raise self.raise_in_before + self.events.append(event) + + def after_launch(self, event): + self.log.append("after") + + def compile_failed(self, event): + self.log.append(("compile_failed", type(event.error))) + self.failures.append(event) + + def finalize(self): + self.log.append("finalize") + self.finalized.append(list(self.events)) + return list(self.records) + + def pre_warmup_callback(self, jit_fn, *args, **kwargs): + self.log.append("pre_warmup") + return False + + def post_warmup_callback(self, jit_fn, ret): + self.log.append("post_warmup") + + def _unreachable(self, *args, **kwargs): + raise AssertionError(f"interpreter hook reached IR client {self.NAME}") + + pre_run_callback = _unreachable + post_run_callback = _unreachable + arg_callback = _unreachable + grid_callback = _unreachable + grid_idx_callback = _unreachable + register_op_callback = _unreachable + register_for_loop_callback = _unreachable + + +class _SkipIRClient(_IRClient): + NAME = "ir_skip" + + +class _PeerIRClient(_IRClient): + """A second IR client class, unrelated to _SkipIRClient.""" + + NAME = "ir_peer" + + +# ======== Fake compile ========= + + +class _FakeKernel: + def __init__(self, key): + self.hash = f"hash-{key}" + self.asm = {"ttir": f"// ttir {key}"} + + def _init_handles(self): + # A CompiledKernel loads its binary here; IR mode never does. + raise AssertionError("IR mode loaded a kernel") + + +def _fake_kernel_signature(x_ptr, n, BLOCK=4): + pass + + +class _FakeJit: + """Stands in for a JITFunction: the host compile calls fake_compile; its + ``run``, the real launch, must never be reached.""" + + signature = inspect.signature(_fake_kernel_signature) + + def __init__(self, log=None, *, compile_error=None): + self.log = [] if log is None else log + self.compile_error = compile_error + + def fake_compile(self, *args, **kwargs): + self.log.append("compile") + if self.compile_error is not None: + raise self.compile_error + return _FakeKernel(kwargs.get("BLOCK", 4)) + + def run(self, *args, grid, warmup, **kwargs): + raise AssertionError("a traced launch with IR clients ran the kernel") + + +def _install_fake_run(monkeypatch, jit_fn, run): + """Install ``run(*args, grid, warmup, **kwargs)`` on a real JITFunction + as both its ``run`` (what a voted warmup compiles through) and its fake + host compile (warmup=True, grid=None: a host compile needs no grid).""" + monkeypatch.setattr(jit_fn, "run", run, raising=False) + monkeypatch.setattr( + jit_fn, + "fake_compile", + lambda *args, **kwargs: run(*args, grid=None, warmup=True, **kwargs), + raising=False, + ) + + +@pytest.fixture +def fake_compile(monkeypatch): + """Record a real JITFunction's host compiles and runs, in order, instead + of performing them. + + ``compile_error(kwargs)`` may return an exception for a compile to + raise, e.g. per config. + """ + + def install(jit_fn, *, fail_first=False, compile_error=None): + calls: list[SimpleNamespace] = [] + + def run(*args, grid, warmup, **kwargs): + calls.append( + SimpleNamespace(args=args, grid=grid, warmup=warmup, kwargs=kwargs) + ) + if fail_first and len(calls) == 1: + raise RuntimeError("compile failed") + if warmup and compile_error is not None: + error = compile_error(kwargs) + if error is not None: + raise error + return _FakeKernel(tuple(sorted(kwargs.items()))) + + _install_fake_run(monkeypatch, jit_fn, run) + return calls + + return install + + +def _fake_bench(kernel_call, quantiles): + # An Autotuner do_bench that needs no GPU: every config ties. + kernel_call() + return [1.0, 1.0, 1.0] + + +def _static_assert_failure(): + return CompileTimeAssertionFailure(None, ast.Pass(), "static_assert failed") + + +def _call(**overrides): + fields: dict = dict(jit_fn=None, args=(), kwargs={}, grid=(1,), capture=False) + fields.update(overrides) + return LaunchCall(**fields) + + +def _make_plain_kernel(): + @triton.jit + def add_one(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + return add_one + + +def _make_autotuned_kernel(**autotune_kwargs): + @triton.autotune( + configs=[triton.Config({"BLOCK": 4}), triton.Config({"BLOCK": 8})], + key=["n"], + **autotune_kwargs, + ) + @triton.heuristics({"EVEN": lambda args: args["n"] % args["BLOCK"] == 0}) + @triton.jit + def add_one_tuned(x_ptr, out_ptr, n, BLOCK: tl.constexpr, EVEN: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + return add_one_tuned + + +def _grid(meta): + return (triton.cdiv(meta["n"], meta["BLOCK"]),) + + +def _dummy_lang_fn(): + """Provides tl globals for patch_lang in patch_run tests.""" + return tl.arange(0, 1) + + +def _in_thread(fn, *args): + """Run ``fn(*args)`` on another host thread; its result or exception.""" + outcome: dict = {} + + def target(): + try: + outcome["result"] = fn(*args) + except BaseException as exc: + outcome["error"] = exc + + worker = threading.Thread(target=target) + worker.start() + worker.join(30) + assert not worker.is_alive() + return outcome + + +# ======== declarations and composition ========= + + +def test_client_declaration_defaults(): + client = _EagerClient() + assert client.NEEDS_INTERPRETER is True + assert client.IR_STAGES == frozenset() + assert client.LAUNCH == "indifferent" + assert not hasattr(client, "collect_asm") + assert not hasattr(client, "asm_info") + for existing in (Sanitizer(), Tracer()): + assert existing.NEEDS_INTERPRETER is True + assert existing.LAUNCH == "indifferent" + + +class _NotSkippingIRClient(_IRClient): + NAME = "ir_not_skipping" + LAUNCH = "indifferent" + + +def test_add_clients_refuses_an_ir_client_that_does_not_skip_the_launch(): + """A traced launch with IR clients never runs the real kernel, so an IR + client must declare LAUNCH="skip"; the refused batch inserts nothing.""" + manager = ClientManager([_SkipIRClient(), _EagerClient()]) + + with pytest.raises(ValueError, match="must declare LAUNCH = 'skip'"): + manager.add_clients([_PeerIRClient(), _NotSkippingIRClient()]) + + assert list(manager.clients) == ["ir_skip", "eager"] + + # An interpreting client's LAUNCH is not read. + class _EagerSkip(_EagerClient): + NAME = "eager_skip" + LAUNCH = "skip" + + manager.add_clients([_EagerSkip()]) + assert list(manager.clients) == ["ir_skip", "eager", "eager_skip"] + + +def test_add_clients_keeps_the_duplicate_rule(): + first = _SkipIRClient() + manager = ClientManager([first, _PeerIRClient(), _EagerClient()]) + manager.add_clients([_SkipIRClient()]) + + assert list(manager.clients) == ["ir_skip", "ir_peer", "eager"] + assert manager.clients["ir_skip"] is first + + +def test_add_clients_rejects_unknown_launch_value(): + class _BadIRClient(_IRClient): + NAME = "ir_bad" + LAUNCH = "maybe" + + with pytest.raises(ValueError, match="LAUNCH must be one of"): + ClientManager([_BadIRClient()]) + + +def test_trace_decorator_refuses_an_ir_client_that_does_not_skip(): + traced = tilelens.trace(_SkipIRClient())(_make_plain_kernel()) + + with pytest.raises(ValueError, match="must declare LAUNCH = 'skip'"): + tilelens.trace(_NotSkippingIRClient())(traced) + + assert list(traced.client_manager.clients) == ["ir_skip"] + + +@pytest.mark.parametrize( + "stages, named", [({"TTIR"}, "['TTIR']"), ({"ttir", "ttgir"}, "['ttgir']")] +) +def test_add_clients_refuses_stages_a_host_compile_does_not_produce(stages, named): + """An IR_STAGES name other than "ttir" (a later stage, or a misspelled + one) is the client's bug: a ValueError when the client is added, never + a compile failure.""" + + class _Stages(_IRClient): + NAME = "ir_stages" + IR_STAGES = frozenset(stages) + + with pytest.raises(ValueError, match=re.escape(f"_Stages.IR_STAGES: {named}")): + ClientManager([_Stages()]) + + # Declaring none is fine (the client reads no stage). + class _NoStages(_IRClient): + NAME = "ir_no_stages" + IR_STAGES = frozenset() + + ClientManager([_NoStages()]) + + +def test_client_partition(): + eager, ir = _EagerClient(), _SkipIRClient() + manager = ClientManager([eager, ir]) + + assert manager.interpreting_clients() == [eager] + assert manager.ir_clients() == [ir] + + +# ======== patch_run ========= + + +def _first_op(): + frontend = get_frontend("triton") + namespace, attrs = next(iter(frontend.namespaces.items())) + attr = next(iter(attrs)) + return frontend, namespace, attr + + +def test_patch_run_registers_ops_only_for_interpreting_clients(): + eager = _EagerClient() + # The IR client would raise if asked for op or loop callbacks. + manager = ClientManager([eager, _PeerIRClient()]) + frontend, namespace, attr = _first_op() + original = frontend.original_ops[namespace][attr] + store_patches = [ + (ns, name) + for ns, attrs in frontend.namespaces.items() + for name, op_type in attrs.items() + if op_type is Store + ] + assert store_patches + + with manager.patch_run(_dummy_lang_fn, frontend_name="triton"): + for ns, name in store_patches: + assert getattr(ns, name).before_callback is eager.on_store + + assert getattr(namespace, attr) is original + + +# ======== interpreter callbacks ========= + + +def test_interpreter_callbacks_reach_only_interpreting_clients(): + eager = _EagerClient() + manager = ClientManager([eager, _PeerIRClient()]) + tensor = torch.zeros(1) + + assert manager.pre_run_callback(_dummy_lang_fn) is True + assert manager.post_run_callback(_dummy_lang_fn) is True + manager.arg_callback("x_ptr", tensor, tensor) + manager.grid_callback((2, 1, 1)) + manager.grid_idx_callback((0, 0, 0)) + + assert eager.calls == [ + "pre_run", + "post_run", + ("arg", "x_ptr"), + ("grid", (2, 1, 1)), + "grid_idx", + ] + assert tensor in manager.launch.tensors + assert manager.launch.grid == (2, 1, 1) + + +def test_run_votes_without_interpreting_clients_keep_the_grid_running(): + manager = ClientManager([_PeerIRClient()]) + + assert manager.pre_run_callback(_dummy_lang_fn) is True + assert manager.post_run_callback(_dummy_lang_fn) is True + + +# ======== finalize, begin/abort ========= + + +def test_begin_and_abort_fan_out_to_every_client(): + log: list = [] + eager, ir = _EagerClient(), _PeerIRClient(log) + manager = ClientManager([eager, ir]) + call = _call() + + manager.begin_launch(call) + manager.abort_launch(KeyError("x")) + + assert eager.calls == ["begin", ("abort", KeyError)] + assert log == ["begin", ("abort", KeyError)] + assert ir.launch_calls == [call] + + +def test_each_launch_gets_its_own_launch_record(): + manager = ClientManager([_EagerClient(records=["record"])]) + manager.begin_launch(_call()) + first = manager.launch + manager.arg_callback("x_ptr", torch.zeros(1), None) + manager.finalize() + + manager.begin_launch(_call()) + + assert manager.launch is not first + assert manager.launch.records == [] and not manager.launch.tensors + assert first.records == ["record"] and len(first.tensors) == 1 + + +def test_abort_hook_failure_never_masks_the_launch_exception(): + class _BrokenAbort(_EagerClient): + NAME = "broken_abort" + + def abort_launch(self, exc): + raise RuntimeError("abort hook failed") + + log: list = [] + manager = ClientManager([_BrokenAbort(), _PeerIRClient(log)]) + launch_exc = ValueError("launch failed") + + if hasattr(launch_exc, "add_note"): + manager.abort_launch(launch_exc) + assert any("abort hook failed" in n for n in launch_exc.__notes__) + else: + with pytest.warns(RuntimeWarning, match="abort hook failed"): + manager.abort_launch(launch_exc) + # Every client still got the abort. + assert log == [("abort", ValueError)] + + +def test_abort_hook_interrupt_propagates_after_every_client(): + class _Interrupting(_EagerClient): + NAME = "interrupting" + + def abort_launch(self, exc): + raise KeyboardInterrupt + + log: list = [] + manager = ClientManager([_Interrupting(), _PeerIRClient(log)]) + launch_exc = ValueError("launch failed") + + with pytest.raises(KeyboardInterrupt) as info: + manager.abort_launch(launch_exc) + + assert info.value.__cause__ is launch_exc + assert log == [("abort", ValueError)] + + +def test_no_abort_after_finalize_started(): + log: list = [] + manager = ClientManager([_PeerIRClient(log)]) + manager.begin_launch(_call()) + manager.finalize() + + manager.abort_launch(SystemExit(3)) + + assert log == ["begin", "finalize"] + + +def test_begin_failure_aborts_exactly_the_clients_that_began(): + class _BrokenBegin(_PeerIRClient): + fail = True + + def begin_launch(self, call): + super().begin_launch(call) + if self.fail: + raise KeyError("begin failed") + + log: list = [] + first, broken, last = _SkipIRClient(log), _BrokenBegin(log), _EagerClient() + manager = ClientManager([first, broken, last]) + + with pytest.raises(KeyError): + manager.begin_launch(_call()) + + # The failing client began (and may hold partial state); `last` never did. + assert log == ["begin", "begin", ("abort", KeyError), ("abort", KeyError)] + assert last.calls == [] + # No launch was left open, so another host thread may begin one. + broken.fail = False + assert "error" not in _in_thread(manager.begin_launch, _call()) + + +def test_begin_launch_refuses_another_threads_launch_without_touching_it(): + log: list = [] + manager = ClientManager([_SkipIRClient(log)]) + manager.begin_launch(_call()) + launch = manager.launch + + refused = _in_thread(manager.begin_launch, _call())["error"] + # A stray abort from that thread does not reach this launch either. + _in_thread(manager.abort_launch, refused) + + assert isinstance(refused, RuntimeError) + assert "another host thread" in str(refused) + assert manager.launch is launch + assert log == ["begin"] + + # Once this launch ends, the other thread may begin the next one. + manager.finalize() + assert "error" not in _in_thread(manager.begin_launch, _call()) + assert log == ["begin", "finalize", "begin"] + + +# ======== ir_capture ========= + + +def test_ir_capture_compiles_without_launching(_fake_host_compile): + log: list = [] + ir, eager = _SkipIRClient(log), _EagerClient() + manager = ClientManager([ir, eager]) + jit_fn = _FakeJit(log) + x = torch.zeros(10) + + with manager.ir_capture(jit_fn): + assert "run" in vars(jit_fn) + ret = jit_fn.run(x, 10, grid=_grid, warmup=False, BLOCK=4, num_warps=2) + + assert "run" not in vars(jit_fn) + # A launching call compiles too; nothing launches (_FakeJit.run raises). + assert log == ["compile", "before", "after"] + # One host compile, for the default target. + target = default_ir_target() + assert _fake_host_compile == [(jit_fn, target)] + (event,) = ir.events + assert event.target == target + assert ret is event.kernel + assert event.jit_fn is jit_fn + assert event.args == (x, 10) + assert dict(event.kwargs) == {"BLOCK": 4, "num_warps": 2} + assert dict(event.bound_args) == {"x_ptr": x, "n": 10, "BLOCK": 4} + assert event.grid is _grid + assert event.resolved_grid == (3, 1, 1) + assert event.specialization == "hash-4" + assert "ttir" in event.kernel.asm + # Launch.grid from the binding, and no IR event for the interpreting + # peer. The binding never adds tensors (none may outlive the launch): + # an interpreted run's arg_callback records its own. + assert not manager.launch.tensors + assert manager.launch.grid == (3, 1, 1) + assert "before_launch" not in eager.calls + + +def test_ir_capture_restores_on_error_and_does_not_double_wrap(): + log: list = [] + manager = ClientManager([_SkipIRClient(log, raise_in_before=ValueError("stop"))]) + jit_fn = _FakeJit(log) + + with pytest.raises(ValueError, match="stop"): + with manager.ir_capture(jit_fn): + wrapper = jit_fn.run + with manager.ir_capture(jit_fn): + assert jit_fn.run is wrapper + assert jit_fn.run is wrapper + jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=True) + + # before_launch raised: no after_launch, wrapper removed. + assert log == ["compile", "before"] + assert "run" not in vars(jit_fn) + + +def test_ir_capture_refuses_another_owner_and_ignores_other_threads(): + class _LaunchingJit(_FakeJit): + def run(self, *args, grid, warmup, **kwargs): + self.log.append("launch") + return "launched" + + log: list = [] + first = ClientManager([_SkipIRClient(log)]) + second = ClientManager([_PeerIRClient()]) + jit_fn = _LaunchingJit(log) + + with first.ir_capture(jit_fn): + with pytest.raises(RuntimeError, match="already being captured"): + with second.ir_capture(jit_fn): + pass + outcome = _in_thread( + lambda: jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=False) + ) + + # The other thread's call went straight to the original run. + assert outcome == {"result": "launched"} + assert log == ["launch"] + + +def test_ir_capture_delivers_each_specialization_once_per_launch(): + log: list = [] + ir = _SkipIRClient(log) + manager = ClientManager([ir]) + jit_fn = _FakeJit(log) + x = torch.zeros(4) + + with manager.ir_capture(jit_fn): + jit_fn.run(x, 4, grid=(1,), warmup=True, BLOCK=4) + with manager.ir_capture(jit_fn): + # e.g. the configs of an autotuned kernel, compiled again. + for block in (4, 4, 8, 4): + jit_fn.run(x, 4, grid=(1,), warmup=True, BLOCK=block) + + assert [e.specialization for e in ir.events] == ["hash-4", "hash-8"] + assert log.count("compile") == 5 # every call still compiled + + # The next traced launch reports its specializations again. + manager.begin_launch(_call(capture=True)) + with manager.ir_capture(jit_fn): + jit_fn.run(x, 4, grid=(1,), warmup=True, BLOCK=4) + assert [e.specialization for e in ir.events] == ["hash-4"] + + +class _Opaque: + """A value the binding fingerprint knows nothing about (weakref-able).""" + + +def test_ir_capture_delivers_each_binding_of_a_specialization(): + """Calls compiling to one kernel ("hash-4") are told apart by their + binding fingerprint, never by tensor data.""" + ir = _SkipIRClient() + manager = ClientManager([ir]) + jit_fn = _FakeJit() + x, y = torch.zeros(8), torch.zeros(8) + opaque, items = _Opaque(), [1] + + def delivered(*args, grid=(1,), **kwargs) -> bool: + before = len(ir.events) + jit_fn.run(*args, grid=grid, warmup=True, **kwargs) + return len(ir.events) > before + + with manager.ir_capture(jit_fn): + assert delivered(x, 4) + assert not delivered(x, 4) # the same call again + assert not delivered(x, 4, grid=_grid) # a callable grid: (1, 1, 1) + assert delivered(x, 5) # a scalar's value + assert delivered(x, 4.0) # ... and type + assert delivered(x, 4, grid=(2,)) # the grid + assert delivered(x, 4, num_warps=8) # a kwarg (compile option) + assert delivered(y, 4) # a tensor's data_ptr + assert delivered(x[:4], 4) # ... shape + assert delivered(x[::2], 4) # ... strides + assert delivered(x.view(torch.int32), 4) # ... dtype + x.add_(1) + assert not delivered(x, 4) # never its data + assert delivered(x, (4, 5)) # a tuple, item by item + assert not delivered(x, (4, 5)) + assert delivered(x, opaque) # anything else by identity + assert not delivered(x, opaque) + assert delivered(x, items) + assert delivered(x, [1]) # equal, but another object + + assert {e.specialization for e in ir.events} == {"hash-4"} + assert [e.resolved_grid for e in ir.events][:4] == [(1, 1, 1)] * 3 + [(2, 1, 1)] + + +class _ConstexprJit(_FakeJit): + """A _FakeJit whose BLOCK is a tl.constexpr parameter.""" + + params = [ + SimpleNamespace(name=name, is_constexpr=name == "BLOCK") + for name in _FakeJit.signature.parameters + ] + + +class _FreshDType: + """Equal to every other instance in all but identity, as a tl.dtype a + heuristic builds per call is (the fake kernel's hash holds the repr).""" + + def __repr__(self): + return "fp32" + + +def test_a_constexpr_argument_counts_only_through_the_specialization(): + """Triton hashes a constexpr argument into the kernel, so the binding + fingerprint leaves it out: an equal but fresh constexpr object per call + adds no binding, passed by keyword or positionally; a non-constexpr + argument still counts by identity.""" + ir = _SkipIRClient() + manager = ClientManager([ir]) + jit_fn = _ConstexprJit() + x = torch.zeros(8) + + def delivered(*args, **kwargs) -> bool: + before = len(ir.events) + jit_fn.run(*args, grid=(1,), warmup=True, **kwargs) + return len(ir.events) > before + + with manager.ir_capture(jit_fn): + assert delivered(x, 4, BLOCK=_FreshDType()) # compiles "hash-fp32" + assert not delivered(x, 4, BLOCK=_FreshDType()) + assert delivered(x, 5, BLOCK=_FreshDType()) # a runtime argument + assert delivered(x, 4, _FreshDType()) # "hash-4": BLOCK is no kwarg + assert not delivered(x, 4, _FreshDType()) + opaque = _Opaque() + assert delivered(x, opaque, BLOCK=_FreshDType()) + assert delivered(x, _Opaque(), BLOCK=_FreshDType()) + + specializations = [e.specialization for e in ir.events] + assert specializations == ["hash-fp32"] * 2 + ["hash-4"] + ["hash-fp32"] * 2 + # Each event still carries the call's own constexpr object. + assert all(isinstance(e.bound_args["BLOCK"], _FreshDType) for e in ir.events) + + +def test_a_pinned_value_outlives_its_call_only_until_the_launch_ends(): + """An unknown value's identity token keeps the object alive for the + launch, so a fresh object per call never reuses a delivered id; once + the launch ends (finalize or abort), nothing holds it.""" + + class _Forgetful(_SkipIRClient): + def before_launch(self, event): + self.log.append(event.bound_args["n"].__class__.__name__) + + for end in ("finalize", "abort"): + ir = _Forgetful() + manager = ClientManager([ir]) + jit_fn = _FakeJit() + manager.begin_launch(_call(capture=True)) + refs = [] + with manager.ir_capture(jit_fn): + for _ in range(3): + value = _Opaque() + refs.append(weakref.ref(value)) + jit_fn.run(torch.zeros(4), value, grid=(1,), warmup=True) + del value + # Three distinct objects, three events: none was freed mid-launch. + assert ir.log.count("_Opaque") == 3 + assert all(ref() is not None for ref in refs) + if end == "finalize": + manager.finalize() + else: + manager.abort_launch(RuntimeError("launch failed")) + gc.collect() + assert all(ref() is None for ref in refs) + + +def test_a_failing_compile_is_reported_as_data(): + log: list = [] + ir = _SkipIRClient(log) + manager = ClientManager([ir]) + x = torch.zeros(4) + error = RuntimeError("the front end failed") + broken = _FakeJit(log, compile_error=error) + + with manager.ir_capture(broken) as window: + assert broken.run(x, 4, grid=(1,), warmup=True) is None + assert (window.compiled, window.failures) == (0, [error]) + + assert ir.events == [] + (failure,) = ir.failures + assert failure.error is error and failure.kernel is None + assert failure.specialization is None + assert failure.target == default_ir_target() + assert dict(failure.bound_args) == {"x_ptr": x, "n": 4, "BLOCK": 4} + + +def test_a_failing_config_is_reported_once_per_launch(): + """The same failing call compiled again in a launch is no news; another + call, or the next launch, is.""" + log: list = [] + ir = _SkipIRClient(log) + manager = ClientManager([ir]) + broken = _FakeJit(log, compile_error=_static_assert_failure()) + x = torch.zeros(4) + + with manager.ir_capture(broken) as window: + for block in (8, 8, 16): + assert broken.run(x, 4, grid=(1,), warmup=True, BLOCK=block) is None + + assert [dict(f.kwargs)["BLOCK"] for f in ir.failures] == [8, 16] + assert len(window.failures) == 3 + manager.begin_launch(_call(capture=True)) + with manager.ir_capture(broken): + broken.run(x, 4, grid=(1,), warmup=True, BLOCK=8) + assert [dict(f.kwargs)["BLOCK"] for f in ir.failures] == [8] + + +class _HipIRClient(_PeerIRClient): + def __init__(self, log=None): + super().__init__(log) + self.ir_target = "hip:gfx942" + + +def test_one_trace_compiles_for_one_target(_fake_host_compile): + """IR clients that name different targets are refused before anything + compiles; a target set explicitly to the default's value is the + default.""" + from triton.backends.compiler import GPUTarget + + cuda89 = GPUTarget("cuda", 89, 32) + ir, hip = _SkipIRClient(), _HipIRClient() + manager = ClientManager([ir, hip]) + jit_fn = _FakeJit() + + with pytest.raises( + ValueError, + match=re.escape("these name several (cuda:89: ir_skip; hip:gfx942: ir_peer)"), + ): + with manager.ir_capture(jit_fn): + pass + assert _fake_host_compile == [] + + hip.ir_target = cuda89 + with manager.ir_capture(jit_fn): + jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=True) + assert _fake_host_compile == [(jit_fn, cuda89)] + # Both IR clients get the one event. + assert ir.events[0] is hip.events[0] + assert ir.events[0].target == cuda89 + + +def test_the_configured_target_is_the_default(monkeypatch, _fake_host_compile): + from triton.backends.compiler import GPUTarget + + monkeypatch.setattr(tilelens_config, "ir_target", "cuda:90") + ir = _SkipIRClient() + manager = ClientManager([ir]) + jit_fn = _FakeJit() + with manager.ir_capture(jit_fn): + jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=True) + assert ir.events[0].target == GPUTarget("cuda", 90, 32) + + # A spec naming no target is an error before anything compiles, not + # a compile failure. + monkeypatch.setattr(tilelens_config, "ir_target", "cuda:sm90") + with pytest.raises(ValueError, match="TILELENS_IR_TARGET.*'cuda:sm90'"): + with manager.ir_capture(jit_fn): + pass + assert len(_fake_host_compile) == 1 and ir.failures == [] + + +def test_an_invalid_configured_target_fails_the_launch(fake_compile, monkeypatch): + monkeypatch.setattr(tilelens_config, "ir_target", "sm80") + log: list = [] + traced = tilelens.trace(_SkipIRClient(log))(_make_plain_kernel()) + calls = fake_compile(traced.jit_fn) + + with pytest.raises(ValueError, match="no IR target"): + traced[(1,)](torch.zeros(4), torch.zeros(4), 4, BLOCK=4) + assert calls == [] and log == ["begin", ("abort", ValueError)] + + +def test_ir_capture_compiles_on_the_real_arguments(): + ir = _SkipIRClient() + manager = ClientManager([ir]) + received: list = [] + + class _RecordingJit(_FakeJit): + def fake_compile(self, *args, **kwargs): + received.append((args, dict(kwargs))) + return super().fake_compile(*args, **kwargs) + + jit_fn = _RecordingJit() + x = torch.zeros(4) + + def real_args(fn, args, kwargs): + assert fn is jit_fn + return (args[0], 8), {**kwargs, "BLOCK": 8} + + with manager.ir_capture(jit_fn, real_args=real_args): + jit_fn.run(x, "traced", grid=(1,), warmup=True) + + assert [(a[1], k) for a, k in received] == [(8, {"BLOCK": 8})] + # The event describes the call as made; the kernel is the compiled one. + (event,) = ir.events + assert event.args == (x, "traced") and dict(event.kwargs) == {} + assert event.specialization == "hash-8" + + +@pytest.fixture +def patched_language(): + """An interpreted traced launch's language patch, as if active on + another host thread.""" + scopes = LANG_PATCH_SCOPES.setdefault("triton", []) + scope = object() + scopes.append(scope) + yield + scopes.remove(scope) + + +def test_a_compile_refuses_while_the_language_is_patched(patched_language): + log: list = [] + ir = _SkipIRClient(log) + manager = ClientManager([ir]) + jit_fn = _FakeJit(log) + + # Reported as data: no compile ran, so the refusal has its own type (a + # RuntimeError). + with manager.ir_capture(jit_fn) as window: + assert jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=True) is None + (failure,) = ir.failures + assert failure.error is window.failures[0] + assert isinstance(failure.error, LanguagePatchedError) + assert isinstance(failure.error, RuntimeError) + assert "language patched" in str(failure.error) + # Nothing was compiled. + assert log == [("compile_failed", LanguagePatchedError)] + + +def test_an_ir_only_launch_raises_the_patched_language_refusal( + fake_compile, patched_language +): + """No compile's outcome, so not a compile failure that lets a launch + survive: the IR-only launch fails (concurrent traced launches that + mix interpretation and real compiles are unsupported). The IR client + saw it as data first.""" + log: list = [] + ir = _SkipIRClient(log) + traced = tilelens.trace(ir)(_make_plain_kernel()) + calls = fake_compile(traced.jit_fn) + + with pytest.raises(LanguagePatchedError, match="language patched"): + traced[(2,)](torch.zeros(8), torch.zeros(8), 8, BLOCK=4) + assert log == [ + "begin", + ("compile_failed", LanguagePatchedError), + ("abort", LanguagePatchedError), + ] + assert calls == [] + + +def test_resolve_grid(): + assert _resolve_grid((2,), {}) == (2, 1, 1) + assert _resolve_grid((2, 3, 4), {}) == (2, 3, 4) + assert _resolve_grid(lambda meta: (meta["n"], 2), {"n": 5}) == (5, 2, 1) + assert _resolve_grid(None, {}) is None + assert _resolve_grid(lambda meta: (meta["missing"],), {}) is None + assert _resolve_grid((1, 1, 1, 1), {}) is None + + +# ======== TritonTrace.run lifecycle ========= + + +def test_ir_only_launch_compiles_every_config_without_running(fake_compile): + log: list = [] + ir = _SkipIRClient(log) + traced = tilelens.trace(ir)(_make_autotuned_kernel()) + calls = fake_compile(traced.jit_fn) + x = torch.arange(8, dtype=torch.float32) + out = torch.zeros(8) + + traced[_grid](x, out, 8) + + assert torch.equal(out, torch.zeros(8)) + assert [c.warmup for c in calls] == [True, True] + assert log == ["begin", "before", "after", "before", "after", "finalize"] + (events,) = ir.finalized + assert [e.kwargs["BLOCK"] for e in events] == [4, 8] + assert [e.kwargs["EVEN"] for e in events] == [True, True] + assert [e.resolved_grid for e in events] == [(2, 1, 1), (1, 1, 1)] + assert len({e.specialization for e in events}) == 2 + (call,) = ir.launch_calls + assert call.jit_fn is traced.jit_fn and call.capture is True + assert call.args == (x, out, 8) and dict(call.kwargs) == {} and call.grid is _grid + # The fake compile is back in place; the capture wrapper is gone. + assert "run" in vars(traced.jit_fn) + assert not getattr(traced.jit_fn.run, "_tilelens_ir_capture", False) + + # The next launch compiles every config again, whatever the autotune + # cache holds. + traced[_grid](x, out, 8) + assert [e.kwargs["BLOCK"] for e in ir.finalized[1]] == [4, 8] + + +@pytest.mark.parametrize( + "failing", + [ + # A config Autotuner._bench drops when it fails like this, + {8: _static_assert_failure}, + # every config, + {4: _static_assert_failure, 8: _static_assert_failure}, + # an error the autotuner does not tolerate. + {8: lambda: ValueError("bad config")}, + ], + ids=["one-config", "every-config", "untolerated-error"], +) +def test_ir_only_compile_failures_never_fail_the_launch(fake_compile, failing): + """A config that fails to compile for the IR target is data for + the IR clients, whatever the error and even when no config compiled; the + skipped launch returns as it does when every config compiles (None for + an autotuned kernel).""" + log: list = [] + ir = _SkipIRClient(log) + traced = tilelens.trace(ir)(_make_autotuned_kernel()) + fake_compile( + traced.jit_fn, + compile_error=lambda kw: failing[kw["BLOCK"]]() + if kw["BLOCK"] in failing + else None, + ) + x, out = torch.zeros(8), torch.zeros(8) + + assert traced[_grid](x, out, 8) is None + + assert log[0] == "begin" and log[-1] == "finalize" + assert not [e for e in log if isinstance(e, tuple) and e[0] == "abort"] + (events,) = ir.finalized + assert [e.kwargs["BLOCK"] for e in events] == [ + b for b in (4, 8) if b not in failing + ] + # Every failing config reached the IR client as data, once. + assert sorted(f.kwargs["BLOCK"] for f in ir.failures) == sorted(failing) + + +def test_a_failed_host_compile_ends_the_launch_normally(fake_compile): + """A plain kernel's only config fails to compile: the skipped launch + returns None (the host-compiled kernel it returns otherwise), and the + next launch compiles as if nothing had happened.""" + log: list = [] + ir = _SkipIRClient(log) + traced = tilelens.trace(ir)(_make_plain_kernel()) + calls = fake_compile(traced.jit_fn, fail_first=True) + args = (torch.zeros(8), torch.zeros(8), 8) + + assert traced[(2,)](*args, BLOCK=4) is None + assert log == ["begin", ("compile_failed", RuntimeError), "finalize"] + assert not getattr(traced.jit_fn.run, "_tilelens_ir_capture", False) + + log.clear() + assert isinstance(traced[(2,)](*args, BLOCK=4), _FakeKernel) + + assert log == ["begin", "before", "after", "finalize"] + assert [len(events) for events in ir.finalized] == [0, 1] + assert len(calls) == 2 + + +def test_mixed_trace_compiles_for_ir_and_interprets_for_eager(fake_compile): + log: list = [] + ir, eager = _SkipIRClient(log), _EagerClient() + traced = tilelens.trace(ir)(_make_plain_kernel()) + traced = tilelens.trace(eager)(traced) + calls = fake_compile(traced.jit_fn) + x = torch.arange(8, dtype=torch.float32) + out = torch.zeros(8) + + traced[(2,)](x, out, 8, BLOCK=4) + + # IR client: one host compile, no real launch. + assert [c.warmup for c in calls] == [True] + (events,) = ir.finalized + assert len(events) == 1 + # Interpreting client: the full interpreted run, which wrote the output. + assert eager.stores == 2 + assert eager.calls.count("pre_run") == 2 + assert "finalize" in eager.calls + torch.testing.assert_close(out, x + 1) + # Both clients vote on the legacy warmup. + assert "pre_warmup" in eager.calls and "pre_warmup" in log + + +def test_mixed_trace_survives_ir_compile_failures(fake_compile): + # E.g. a kernel the host compile rejects, which the interpreter runs. + log: list = [] + ir, eager = _SkipIRClient(log), _EagerClient() + traced = tilelens.trace(eager)(tilelens.trace(ir)(_make_plain_kernel())) + fake_compile( + traced.jit_fn, + compile_error=lambda kw: RuntimeError("the host compile failed"), + ) + x = torch.arange(8, dtype=torch.float32) + out = torch.zeros(8) + + traced[(2,)](x, out, 8, BLOCK=4) + + torch.testing.assert_close(out, x + 1) + assert eager.stores == 2 and "finalize" in eager.calls + assert ir.finalized == [[]] + assert [type(f.error) for f in ir.failures] == [RuntimeError] + assert not any(isinstance(e, tuple) and e[0] == "abort" for e in log) + + +# ======== a call that does not bind raises, as untraced ========= + + +def _bind_failure(): + """The TypeError the host compile raises for a call missing ``n``, + marked as a bind failure (tilelens.core.host_compile.bind_failed).""" + from tilelens.core import host_compile + + exc = TypeError("dynamic_func() missing 1 required positional argument: 'n'") + host_compile._mark_bind_failed(exc) + return exc + + +@pytest.mark.parametrize("warmup", [True, False]) +def test_ir_capture_raises_a_call_that_does_not_bind(warmup): + """A bind failure is the call's own error, which JITFunction.run raises + on any device: ir_capture raises that very exception, whatever the + call's warmup flag. No compile_failed event, and the capture is + removed.""" + log: list = [] + manager = ClientManager([_SkipIRClient(log)]) + unbound = _bind_failure() + jit_fn = _FakeJit(log, compile_error=unbound) + + with manager.ir_capture(jit_fn) as window: + with pytest.raises(TypeError) as raised: + jit_fn.run(torch.zeros(4), grid=(1,), warmup=warmup) + + assert raised.value is unbound + assert log == ["compile"] + assert (window.compiled, window.failures) == (0, []) + assert "run" not in vars(jit_fn) + + +@pytest.mark.parametrize( + "make, launch", + [ + (_make_plain_kernel, lambda k, x, out: k[(2,)](x, out, 8, BLOCK=4)), + ( + lambda: _make_autotuned_kernel(do_bench=_fake_bench), + lambda k, x, out: k[_grid](x, out, 8), + ), + ], + ids=["plain", "autotuned"], +) +def test_a_traced_call_that_does_not_bind_raises_as_untraced( + fake_compile, make, launch +): + """An IR-only launch raises the bind failure ("the program goes on" is + for kernel compile failures only), from the first compile + of the compile-only pass: the IR client's launch is aborted, never + finalized, nothing launches or is recorded, and the next launch runs as + if nothing had happened.""" + log: list = [] + ir = _SkipIRClient(log) + traced = tilelens.trace(ir)(make()) + unbound = _bind_failure() + failing = [unbound] + calls = fake_compile( + traced.jit_fn, compile_error=lambda kw: failing.pop() if failing else None + ) + x, out = torch.zeros(8), torch.zeros(8) + launches = len(trace_module.launches) + + with pytest.raises(TypeError) as raised: + launch(traced, x, out) + + assert raised.value is unbound + assert log == ["begin", ("abort", TypeError)] + assert [c.warmup for c in calls] == [True] + assert len(trace_module.launches) == launches + assert not getattr(traced.jit_fn.run, "_tilelens_ir_capture", False) + + log.clear() + launch(traced, x, out) + assert log[0] == "begin" and log[-1] == "finalize" + assert len(trace_module.launches) == launches + 1 + + +def test_a_mixed_trace_raises_a_call_that_does_not_bind_before_interpreting( + fake_compile, +): + """A mixed trace (IR and eager clients): the IR clients' compile pass + raises the bind failure before the interpreter runs (which would fail on + the same call too), so the untraced JIT's error is the one raised; every + client's launch is aborted.""" + log: list = [] + ir, eager = _SkipIRClient(log), _EagerClient() + traced = tilelens.trace(eager)(tilelens.trace(ir)(_make_plain_kernel())) + unbound = _bind_failure() + fake_compile(traced.jit_fn, compile_error=lambda kw: unbound) + x, out = torch.arange(8, dtype=torch.float32), torch.zeros(8) + + with pytest.raises(TypeError) as raised: + traced[(2,)](x, out, 8, BLOCK=4) + + assert raised.value is unbound + assert log == ["begin", ("abort", TypeError)] + assert eager.calls == ["begin", ("abort", TypeError)] + assert eager.stores == 0 and torch.equal(out, torch.zeros(8)) + + +def test_a_launch_failing_in_finalize_is_not_aborted(fake_compile): + class _ExitingIR(_SkipIRClient): + def finalize(self): + super().finalize() + raise SystemExit(3) + + log: list = [] + traced = TritonTrace(_make_plain_kernel(), _ExitingIR(log)) + traced.add_client(_PeerIRClient(log)) + fake_compile(traced.jit_fn) + + with pytest.raises(SystemExit): + traced[(1,)](torch.zeros(4), torch.zeros(4), 4, BLOCK=4) + + # Each launch ends in finalize or abort, never both. + assert log.count("finalize") == 2 + assert not any(isinstance(e, tuple) and e[0] == "abort" for e in log) + + +def test_each_traced_launch_is_recorded_separately(fake_compile): + class _VerdictIR(_SkipIRClient): + def finalize(self): + super().finalize() + return [f"verdict-{len(self.finalized)}"] + + traced = tilelens.trace(_VerdictIR())(_make_plain_kernel()) + fake_compile(traced.jit_fn) + before = len(trace_module.launches) + a, b = torch.zeros(8), torch.zeros(16) + + traced[(2,)](a, a, 8, BLOCK=4) + traced[(4,)](b, b, 16, BLOCK=4) + + first, second = trace_module.launches[before:] + assert first is not second + assert first.records == ["verdict-1"] and second.records == ["verdict-2"] + # Launch.grid from the IR binding, per launch; an IR-only launch + # records no tensors. + assert not first.tensors and not second.tensors + assert (first.grid, second.grid) == ((2, 1, 1), (4, 1, 1)) + + +def test_cli_shape_traces_the_autotuner_over_an_inner_trace(fake_compile): + # The CLI wrappers turn every @triton.jit into a TritonTrace and wrap the + # Autotuner built on it again: TritonTrace(Autotuner(TritonTrace(JIT))). + inner_ir, outer_ir = _SkipIRClient(), _SkipIRClient() + jit_fn = _make_plain_kernel() + inner = TritonTrace(jit_fn, inner_ir) + user = triton.autotune( + configs=[triton.Config({"BLOCK": 4}), triton.Config({"BLOCK": 8})], + key=["n"], + )(inner) + before = dict(vars(user)) + outer = TritonTrace(user, outer_ir) + fake_compile(outer.jit_fn) + x, out = torch.zeros(8), torch.zeros(8) + + outer[_grid](x, out, 8) + + assert outer.jit_fn is inner.jit_fn is jit_fn + (events,) = outer_ir.finalized + assert [e.kwargs["BLOCK"] for e in events] == [4, 8] + assert inner_ir.log == [] + assert torch.equal(out, torch.zeros(8)) + assert vars(user).keys() == before.keys() + assert all(vars(user)[k] is v for k, v in before.items()) + + +def test_a_skipped_launch_returns_the_kernel_only_without_an_autotuner(fake_compile): + @triton.heuristics({"BLOCK": lambda args: 4}) + @triton.jit + def heur_kernel(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + + args = (torch.zeros(8), torch.zeros(8), 8) + plain = tilelens.trace(_SkipIRClient())(_make_plain_kernel()) + fake_compile(plain.jit_fn) + heur = tilelens.trace(_SkipIRClient())(heur_kernel) + fake_compile(heur.jit_fn) + tuned = tilelens.trace(_SkipIRClient())(_make_autotuned_kernel()) + fake_compile(tuned.jit_fn) + + # One config: the kernel the untraced launch would return. + assert isinstance(plain[(2,)](*args, BLOCK=4), _FakeKernel) + assert isinstance(heur[(2,)](*args), _FakeKernel) + # Autotuned: no config was picked. + assert tuned[_grid](*args) is None + + +def test_an_autotuned_call_passing_a_tuned_parameter_raises_as_untraced( + fake_compile, +): + """The untraced launch's autotuning refuses a call that passes one of a + config's meta-parameters itself (Autotuner._bench's ValueError). The IR + compiles go through Autotuner.warmup instead, and refuse it the same + way, not with a TypeError naming a tilelens internal.""" + log: list = [] + traced = tilelens.trace(_SkipIRClient(log))(_make_autotuned_kernel()) + calls = fake_compile(traced.jit_fn) + args = (torch.zeros(8), torch.zeros(8), 8) + + with pytest.raises(ValueError, match="Conflicting meta-parameters: BLOCK"): + traced[_grid](*args, BLOCK=4) + + assert calls == [] and log == ["begin", ("abort", ValueError)] + + +class _ForgetfulIRClient(_IRClient): + """An IR client that keeps nothing of a launch's arguments.""" + + NAME = "ir_forgetful" + + def begin_launch(self, call): + self.log.append("begin") + + def before_launch(self, event): + self.log.append("before") + + +def _raise_in_pre_hook(nargs): + raise RuntimeError("pre_hook failed") + + +@pytest.mark.parametrize("path", ["ir", "interpreted"]) +def test_a_launch_that_raises_keeps_no_caller_tensor(path): + """Autotuner.run and .warmup keep the call's arguments in ``nargs`` + until they return. Once a launch has raised, none of the trace's + autotuner copies keeps the caller's tensors.""" + if path == "ir": + # The IR compiles refuse a call passing a tuned meta-parameter. + traced = tilelens.trace(_ForgetfulIRClient())(_make_autotuned_kernel()) + kwargs, error = {"BLOCK": 4}, "Conflicting meta-parameters" + else: + # The interpreted autotuning raises in a config's pre_hook. + kernel = triton.autotune( + configs=[triton.Config({"BLOCK": 4}, pre_hook=_raise_in_pre_hook)], + key=["n"], + )(_make_plain_kernel()) + traced = tilelens.trace(_EagerClient())(kernel) + kwargs, error = {}, "pre_hook failed" + + def launch() -> list[weakref.ref]: + # The caller's references end with this frame. + x, out = torch.zeros(8), torch.zeros(8) + with pytest.raises((ValueError, RuntimeError), match=error): + traced[_grid](x, out, 8, **kwargs) + return [weakref.ref(x), weakref.ref(out)] + + refs = launch() + gc.collect() + + for runner in (traced.runner, traced.warmup_runner, traced.ir_runner): + assert getattr(runner, "nargs", None) is None + assert [ref() for ref in refs] == [None, None] + + +def test_launch_grid_is_the_grid_the_compiled_configs_share(fake_compile): + # Nothing launches: configs that disagree on the grid leave it open, a + # grid they share is the launch's. + args = (torch.zeros(8), torch.zeros(8), 8) + traced = tilelens.trace(_SkipIRClient())(_make_autotuned_kernel()) + fake_compile(traced.jit_fn) + traced[_grid](*args) + assert trace_module.launches[-1].grid is None + traced[(3,)](*args) + assert trace_module.launches[-1].grid == (3, 1, 1) + + +def test_launch_tensors_have_one_representation_per_launch(fake_compile): + x = torch.arange(8, dtype=torch.float32) + out = torch.zeros(8) + + ir_only = tilelens.trace(_SkipIRClient())(_make_plain_kernel()) + fake_compile(ir_only.jit_fn) + ir_only[(2,)](x, out, 8, BLOCK=4) + # None: holding the caller's device tensors would keep them alive + # after the launch; the IR clients' records hold the facts they need. + assert not trace_module.launches[-1].tensors + assert trace_module.launches[-1].grid == (2, 1, 1) + + mixed = tilelens.trace(_EagerClient())( + tilelens.trace(_SkipIRClient())(_make_plain_kernel()) + ) + fake_compile(mixed.jit_fn) + mixed[(2,)](x, out, 8, BLOCK=4) + # Only the interpreter's host copies, whose addresses the eager + # clients' records use; not the caller's tensors on top. + tensors = trace_module.launches[-1].tensors + assert len(tensors) == 2 + assert not {id(t) for t in tensors} & {id(x), id(out)} + + +def test_ir_compiles_are_not_gated_by_an_instance_warmup_patch( + fake_compile, monkeypatch +): + # E.g. a vote gate someone left on the JITFunction, declining every + # compile: IR compiles go through JITFunction's own warmup. + def declining_warmup(*args, **kwargs): + return None + + args = (torch.zeros(8), torch.zeros(8), 8) + for traced, grid, kwargs in ( + (tilelens.trace(_SkipIRClient())(_make_plain_kernel()), (2,), {"BLOCK": 4}), + (tilelens.trace(_SkipIRClient())(_make_autotuned_kernel()), _grid, {}), + ): + calls = fake_compile(traced.jit_fn) + monkeypatch.setattr(traced.jit_fn, "warmup", declining_warmup, raising=False) + traced[grid](*args, **kwargs) + (ir,) = traced.client_manager.ir_clients() + assert ir.finalized[-1] and calls + + +def test_ir_launch_compiles_on_untraced_arguments(fake_compile): + helper = tilelens.trace(_SiblingEagerClient())(triton.jit(_unwrap_leaf)) + + @triton.jit + def apply(x_ptr, out_ptr, n, FN: tl.constexpr, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + + ir = _SkipIRClient() + traced = tilelens.trace(ir)(apply) + calls = fake_compile(traced.jit_fn) + traced[(1,)](torch.zeros(4), torch.zeros(4), 4, FN=helper, BLOCK=4) + + # The host compile sees the JITFunction; events describe the call as + # made. + assert [c.kwargs["FN"] for c in calls] == [helper.jit_fn] + assert all(e.kwargs["FN"] is helper for e in ir.finalized[0]) + + +def test_ir_clients_without_a_jit_function(): + ir = _SkipIRClient() + traced = TritonTrace(InterpretedFunction(_make_plain_kernel().fn), ir) + assert traced.jit_fn is None + x = torch.arange(8, dtype=torch.float32) + out = torch.zeros(8) + + assert traced[(2,)](x, out, 8, BLOCK=4) is None + + # No compiled kernel, so no events, and nothing runs. + assert ir.finalized == [[]] + (call,) = ir.launch_calls + assert call.jit_fn is None and call.capture is False + assert torch.equal(out, torch.zeros(8)) + + +@pytest.mark.parametrize("trace_cls", [GluonTrace, NKITrace]) +def test_gluon_and_nki_traces_run_the_launch_lifecycle(trace_cls): + # Built without __init__: the Gluon simulation and NKI are not importable + # everywhere, and an IR-only trace never reaches them. + class _BrokenBegin(_SkipIRClient): + def begin_launch(self, call): + super().begin_launch(call) + raise KeyError("begin failed") + + log: list = [] + traced = trace_cls.__new__(trace_cls) + TraceInterface.__init__(traced, _SkipIRClient(log)) + + # Only IR clients: nothing is interpreted. + assert traced[(2,)](torch.zeros(4)) is None + assert log == ["begin", "finalize"] + (call,) = traced.client_manager.clients["ir_skip"].launch_calls + assert call.jit_fn is None and call.capture is False and call.grid == (2,) + + log = [] + traced = trace_cls.__new__(trace_cls) + TraceInterface.__init__(traced, _BrokenBegin(log)) + with pytest.raises(KeyError): + traced[(2,)](torch.zeros(4)) + assert log == ["begin", ("abort", KeyError)] + + +# ======== nested traced calls ========= + + +def _unwrap_leaf(x): + return x + 1 + + +def _kernel_calling_nested_leaf(x_ptr, out_ptr, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(out_ptr + offs, _nested_traced_leaf(tl.load(x_ptr + offs))) # noqa: F821 + + +def test_nested_traced_calls_compare_only_interpreting_clients(fake_compile): + # The CLI shape: the helper is traced with the eager client only, the + # kernel with it and an IR client, which takes no part in the + # interpreted run. + module_globals = globals() + module_globals["_nested_traced_leaf"] = tilelens.trace(_EagerClient())( + triton.jit(_unwrap_leaf) + ) + try: + traced = tilelens.trace(_EagerClient())( + tilelens.trace(_SkipIRClient())(triton.jit(_kernel_calling_nested_leaf)) + ) + fake_compile(traced.jit_fn) + x = torch.arange(4, dtype=torch.float32) + out = torch.zeros(4) + + traced[(1,)](x, out, BLOCK=4) + + torch.testing.assert_close(out, x + 1) + finally: + module_globals.pop("_nested_traced_leaf", None) diff --git a/tilelens/core/client.py b/tilelens/core/client.py index f58fcf1b9..0f0bd478e 100644 --- a/tilelens/core/client.py +++ b/tilelens/core/client.py @@ -1,9 +1,14 @@ from contextlib import AbstractContextManager, contextmanager, nullcontext from abc import ABC, abstractmethod -from typing import ClassVar, Any -from collections.abc import Callable +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import ClassVar, Any, Literal +from collections.abc import Callable, Hashable, Mapping +import inspect +import operator import threading +import warnings from .data import Op, Launch from .patch import ( @@ -18,23 +23,133 @@ from functools import wraps from .callbacks import OpCallbacks, ForLoopCallbacks from .patch import patch_lang, unpatch_lang -from .frontend.base import get_frontend +from .frontend.base import LANG_PATCH_SCOPES, get_frontend from .config import config as cfg +from .host_compile import ( + HostCompiler, + bind_failed, + format_ir_target, + resolve_ir_target, +) + +LaunchPreference = Literal["skip", "indifferent"] +LAUNCH_PREFERENCES: tuple[LaunchPreference, ...] = ("skip", "indifferent") +# The compiler stages an IR client can read: the host compile runs Triton's +# front end and the backend's "ttir" passes, nothing after them. +IR_STAGE_NAMES = frozenset({"ttir"}) # (jit_fn, args, kwargs) -> (args, kwargs): the arguments a real compile of # one JITFunction call must see, supplied by the trace. RealArgs = Callable[[Any, tuple, dict], tuple[tuple, dict]] +@dataclass(frozen=True, eq=False) +class LaunchCall: + """One traced launch as the caller made it, delivered to ``begin_launch``. + + Mechanism-only data. ``eq=False`` keeps identity comparison, since + field-wise equality would compare tensors. + """ + + # The traced JITFunction; None when the trace has none (TRITON_INTERPRET / + # InterpretedFunction runner, Gluon, NKI). + jit_fn: Any + args: tuple + # The caller's keyword arguments, excluding ``grid`` and ``warmup``. + # Config kwargs added by Autotuner/Heuristics layers appear only in + # LaunchEvent.kwargs. + kwargs: Mapping[str, Any] + grid: Any + # Whether IR clients receive compile events for this launch. False when + # there is no JITFunction (``jit_fn`` is None). + capture: bool + + +@dataclass(frozen=True, eq=False) +class LaunchEvent: + """One call into a traced JITFunction's ``run``, as delivered to IR clients. + + Mechanism-only data: what was compiled, for which target, and how the + call was bound. What the compiled kernel means is left to each client. + ``eq=False`` keeps identity comparison, since field-wise equality would + compare tensors. + """ + + jit_fn: Any + # Positional arguments as passed to JITFunction.run. + args: tuple + # Keyword arguments, including Autotuner/Heuristics config kwargs, + # excluding ``grid`` and ``warmup``. + kwargs: Mapping[str, Any] + # The grid as passed: a tuple, a callable, or None (e.g. for a warmup). + grid: Any + # ``grid`` canonicalized to three dims, a callable resolved against + # ``bound_args`` as JITFunction.run does; None if it cannot be resolved. + resolved_grid: tuple[Any, Any, Any] | None + # Kernel parameter name -> value, defaults applied. + bound_args: Mapping[str, Any] + # The kernel compiled on the host for ``target`` through its TTIR + # (tilelens.core.host_compile): ``.asm["ttir"]`` the text, ``.metadata`` + # the compile metadata, ``.hash`` the specialization. Never loaded or + # launched. None in a compile_failed event. + kernel: Any + # Identity of the compiled specialization (``kernel.hash``: what + # triton.compile names the kernel for ``target``); None when nothing + # was compiled. + specialization: Hashable + # compile_failed only: the exception the host compile raised: the + # kernel's compile error for ``target``, or, when the host compile could + # not run at all, a HostCompileUnavailable or an error raised from one + # (tilelens.core.host_compile's host_compile_unavailable tells them + # apart; its target_queried marks an error after the front end had + # asked the driver, in this compile or an earlier one of the kernel for + # the target), or a LanguagePatchedError (no compile ran). Never the + # call's own bind error (host_compile.bind_failed): the core raises + # that as the untraced call does (see ClientManager.ir_capture). + error: BaseException | None = None + # The GPUTarget ``kernel`` was compiled for: the receiving clients' + # (Client.ir_target). None only for an event built outside the core. + target: Any = None + + +@dataclass +class CaptureWindow: + """What one ``ClientManager.ir_capture`` window observed.""" + + # Host compiles that produced a kernel (one per call). + compiled: int = 0 + # Host compile exceptions, in call order (one per call). + failures: list[BaseException] = field(default_factory=list) + + class Client(ABC): NAME: ClassVar[str] + # Whether the client consumes the interpreted run (op/loop callbacks, + # pre/post_run votes, arg/grid callbacks). IR clients set this to False + # and receive compiled kernels through before_launch/after_launch. + NEEDS_INTERPRETER: ClassVar[bool] = True + # Compiler stages (keys of the compiled kernel's asm) an IR client reads. + # The core compiles each call on the host through the backend's "ttir" + # passes, so "ttir" is the one stage there is; ClientManager.add_clients + # refuses an IR client that declares another. + IR_STAGES: ClassVar[frozenset[str]] = frozenset() + # An IR client's real-launch preference. The core never runs the real + # kernel of a traced launch that has IR clients: they get the kernels + # compiled on the host, and the interpreted run stands in for the launch + # when an interpreting client shares the trace. So an IR client must + # declare "skip" (ClientManager.add_clients refuses one that does not); + # an interpreting client's LAUNCH changes nothing. + LAUNCH: ClassVar[LaunchPreference] = "indifferent" + # The target an IR client's kernels are compiled for: a triton + # GPUTarget or a spec such as "cuda:90" or "hip:gfx942" (see + # tilelens.core.host_compile.parse_ir_target); None for the configured + # default (tilelens.config.ir_target: TILELENS_IR_TARGET, else + # "cuda:89"). A client may set it per instance. The IR clients of one + # trace must compile for the same target (see ClientManager.ir_target). + ir_target: Any = None def __init__(self) -> None: - # Whether this client needs ASM information from kernel warmup - self.collect_asm: bool = False - # Storage for ASM information if collected - self.asm_info: dict | None = None # Thread-local scratch space for per-thread callback state self._thread_local = threading.local() # Lock for serializing shared state where needed @@ -106,6 +221,66 @@ def pre_warmup_callback(self, jit_fn: Callable, *args, **kwargs) -> bool: def post_warmup_callback(self, jit_fn: Callable, ret: Any) -> None: ... + # Each begun launch ends in exactly one of finalize() or abort_launch(). + # A client whose begin_launch raised still gets abort_launch; clients + # after it in the trace never begin that launch. + + def begin_launch(self, call: LaunchCall) -> None: + """Called before every traced launch; reset per-launch state here.""" + + def abort_launch(self, exc: BaseException) -> None: + """Called when a traced launch raises before finalize; ``exc`` is + re-raised afterwards and finalize() is not called for this launch.""" + + def before_launch(self, event: LaunchEvent) -> None: + """IR clients: ``event.kernel`` was compiled on the host for the + client's target (``event.target``). + + Fires once per traced launch for each distinct (specialization, + binding fingerprint), whichever call produced it: a compile-only + warmup or a config of the launch's autotuning. + The fingerprint summarizes how the call was bound, never tensor + data: the resolved grid, and every argument and kwarg (config kwargs + and compile options such as num_warps included) except those to + tl.constexpr parameters, which the specialization already tells + apart (so a heuristic handing out an equal but fresh constexpr + object per call adds no binding). An int/bool/float/str/None value + counts by type and value, a tuple item by item, a tensor by its + data_ptr, shape, strides and dtype; any other value counts by + identity, so an equal but distinct object is another binding. Two + configs that compile to one kernel but differ in a runtime argument + or the grid thus get an event each, while repeated identical calls + share one. + + A TritonTrace compiles every config before the launch goes on, so + each config is seen on every launch, whatever the autotune cache + holds. + """ + + def after_launch(self, event: LaunchEvent) -> None: + """IR clients: the call described by ``event`` has finished. Only a + call that fired before_launch gets one; the real kernel never ran + (see LAUNCH), so it follows before_launch right away.""" + + def compile_failed(self, event: LaunchEvent) -> None: + """IR clients: a call failed to compile on the host for the + client's target (``event.target``); ``event.error`` is the + exception, ``event.kernel`` is None. + + Fires once per traced launch for each distinct failing call (its + arguments and kwargs, constexprs included). Nothing is loaded: a + kernel the device could not run (e.g. too much shared memory) still + compiles, and is delivered like any other. A failing host compile + never fails the launch: the target is the IR client's choice, not + the machine's, and the launch goes on without the config. Deciding + what a failure means for the client's result is the client's call. + + A call that does not bind the kernel's parameters is no compile + failure and never gets here: the launch raises the binder's error, + as the untraced call does on any device (see + ClientManager.ir_capture), and the clients get abort_launch. + """ + def _set_thread_local(self, key: str, value: Any) -> None: setattr(self._thread_local, key, value) @@ -158,6 +333,137 @@ def remove(self, jit_fn: Any) -> None: top = top.below +_MISSING = object() + + +@contextmanager +def _instance_attr(obj: Any, name: str, value: Any): + """Install ``value`` as an instance attribute of ``obj`` for the scope, then + restore exactly what was there (deleting it if the class provided it).""" + previous = getattr(obj, "__dict__", {}).get(name, _MISSING) + setattr(obj, name, value) + try: + yield + finally: + if previous is _MISSING: + obj.__dict__.pop(name, None) + else: + setattr(obj, name, previous) + + +def _bind_launch_args( + jit_fn: Any, args: tuple, kwargs: Mapping[str, Any] +) -> dict[str, Any]: + """Parameter name -> value for one run() call, the way JITFunction's + binder builds ``bound_args`` (bind, then apply defaults; non-parameter + kwargs such as num_warps are compile options, not arguments).""" + signature = getattr(jit_fn, "signature", None) + if not isinstance(signature, inspect.Signature): + return {} + params = {k: v for k, v in kwargs.items() if k in signature.parameters} + try: + bound = signature.bind(*args, **params) + except TypeError: + return {} + bound.apply_defaults() + return dict(bound.arguments) + + +def _resolve_grid(grid: Any, bound_args: Mapping[str, Any]) -> tuple | None: + """Canonicalize a launch grid to three dims; None if it cannot be resolved.""" + if grid is None: + return None + try: + resolved = tuple(grid(dict(bound_args)) if callable(grid) else grid) + except Exception: + return None + if not 1 <= len(resolved) <= 3: + return None + return resolved + (1,) * (3 - len(resolved)) + + +def _specialization(kernel: Any) -> Hashable: + specialization = getattr(kernel, "hash", None) + return id(kernel) if specialization is None else specialization + + +def _fingerprint_value(value: Any, pinned: dict[int, Any]) -> Hashable: + """A hashable summary of one call value that holds no user object (see + Client.before_launch): a plain scalar by type and value, a tensor by + data_ptr, shape, strides and dtype (never its data), a tuple item by + item. Anything else is a type+id token; the object is put in ``pinned`` + so its id cannot be reused by another object while the tokens are + compared, and a distinct object never shares a token.""" + if value is None: + return None + if isinstance(value, (bool, int)): + return (type(value), int(value)) + if isinstance(value, float): + # hex() tells -0.0 from 0.0 and makes NaN equal to itself. + return (type(value), float.hex(value)) + if isinstance(value, str): + return (type(value), str(value)) + if isinstance(value, tuple): + return (type(value), tuple(_fingerprint_value(v, pinned) for v in value)) + if hasattr(value, "data_ptr"): + try: + return ( + "tensor", + int(value.data_ptr()), + tuple(int(size) for size in value.shape), + tuple(int(stride) for stride in value.stride()), + str(value.dtype), + ) + except Exception: + pass + pinned[id(value)] = value + return (type(value), id(value)) + + +def _constexpr_params(jit_fn: Any) -> tuple[frozenset[int], frozenset[str]]: + """The positions and names of ``jit_fn``'s tl.constexpr parameters.""" + constexprs = [ + (index, param.name) + for index, param in enumerate(getattr(jit_fn, "params", None) or ()) + if getattr(param, "is_constexpr", False) + ] + return ( + frozenset(index for index, _ in constexprs), + frozenset(name for _, name in constexprs), + ) + + +def _grid_fingerprint(resolved: tuple | None, pinned: dict[int, Any]) -> Hashable: + if resolved is None: + return None + try: + # The launcher reads each dim as an index, so e.g. a numpy int dim a + # grid callable returns afresh per call is the same grid each time. + return tuple(operator.index(dim) for dim in resolved) + except Exception: + return _fingerprint_value(resolved, pinned) + + +class LanguagePatchedError(RuntimeError): + """A host compile refused to start while an interpreted traced launch + has the language patched (see _refuse_patched_language): no compile + ran, so it is no kernel's compile error.""" + + +def _refuse_patched_language() -> None: + """Raise LanguagePatchedError before a compile while an interpreted + traced launch has triton.language patched (patch_lang is process-wide, + so the code generator would run on the interpreter's builtins).""" + patched = [name for name in ("triton", "gluon") if LANG_PATCH_SCOPES.get(name)] + if patched: + raise LanguagePatchedError( + "a Triton compile cannot run while an interpreted traced " + f"launch has the {'/'.join(patched)} language patched (e.g. on " + "another host thread); concurrent traced launches that mix " + "interpretation and host compiles are not supported." + ) + + class ClientManager: def __init__(self, clients: list[Client] | None = None): self.clients: dict[str, Client] = {} @@ -165,7 +471,35 @@ def __init__(self, clients: list[Client] | None = None): self.add_clients(clients) self.launch = Launch() self._lock = threading.Lock() + # Compiles every IR client's kernels on the host; its cache + # lives as long as the trace. + self.compiler = HostCompiler() + # The host thread whose launch is in flight (begin_launch until + # finalize or abort_launch), and the lock guarding it. + self._launch_owner: int | None = None + self._owner_lock = threading.Lock() self._clear_loop_hooks() + self._reset_launch_state() + + def _reset_launch_state(self) -> None: + # Per traced launch: which (target, specialization, binding + # fingerprint) keys IR clients were already given, which (target, + # call) compile failures, the objects their identity tokens stand + # for, which parameters the fingerprint leaves out, and the grids of + # the compiled kernels Launch.grid is settled from. Nothing here + # refers to a caller's tensor once the launch has ended. + self._delivered: set[tuple[Any, Hashable, Hashable]] = set() + self._failed: set[tuple[Any, Hashable]] = set() + self._pinned: dict[int, Any] = {} + # id(jit_fn) -> (jit_fn, its constexpr positions, their names). + self._constexprs: dict[int, tuple[Any, frozenset[int], frozenset[str]]] = {} + self._compiled_grids: set[tuple] = set() + self._finalize_started = False + + def _release_pinned(self) -> None: + # The launch has ended: its fingerprints are compared no more, so + # the objects kept alive for their identity tokens can go. + self._pinned = {} def _lock_context(self): if cfg.num_sms > 1: @@ -176,13 +510,161 @@ def get_client(self, name: str) -> Client | None: return self.clients.get(name) def add_clients(self, new_clients_list: list[Client]) -> None: + # Check every addition before inserting any, so a rejected batch + # leaves the manager unchanged. + additions: dict[str, Client] = {} for new_client in new_clients_list: duplicate = any( isinstance(existing_client, new_client.__class__) - for existing_client in self.clients.values() + for existing_client in (*self.clients.values(), *additions.values()) ) if not duplicate: - self.clients[new_client.NAME] = new_client + additions[new_client.NAME] = new_client + for client in additions.values(): + self._check_declarations(client) + self.clients.update(additions) + + @staticmethod + def _check_declarations(client: Client) -> None: + # Core only checks what it can deliver; what a launch means to a + # client stays with the client. + name = type(client).__name__ + if client.LAUNCH not in LAUNCH_PREFERENCES: + raise ValueError( + f"{name}.LAUNCH must be one of {LAUNCH_PREFERENCES}, got " + f"{client.LAUNCH!r}" + ) + if client.NEEDS_INTERPRETER: + return + if client.LAUNCH != "skip": + raise ValueError( + f"{name}.LAUNCH is {client.LAUNCH!r}, but a traced launch with " + "IR clients never runs the real kernel: an IR client must " + "declare LAUNCH = 'skip'" + ) + unknown = set(client.IR_STAGES) - IR_STAGE_NAMES + if unknown: + raise ValueError( + f"{name}.IR_STAGES: {sorted(unknown)} are no stage of a host " + f"compile, which produces {sorted(IR_STAGE_NAMES)} only" + ) + + def interpreting_clients(self) -> list[Client]: + return [c for c in self.clients.values() if c.NEEDS_INTERPRETER] + + def ir_clients(self) -> list[Client]: + return [c for c in self.clients.values() if not c.NEEDS_INTERPRETER] + + def ir_target(self) -> Any: + """The GPUTarget the IR clients' kernels are compiled for: their + Client.ir_target, else the configured default; None without IR + clients. Raises ValueError for a target spec that names no target, + or for IR clients that name different targets: a trace compiles + each call for one target.""" + targets: dict[Any, list[str]] = {} + for client in self.ir_clients(): + target = resolve_ir_target(client.ir_target) + targets.setdefault(target, []).append(client.NAME) + if len(targets) > 1: + named = "; ".join( + f"{format_ir_target(target)}: {', '.join(names)}" + for target, names in targets.items() + ) + raise ValueError( + "the IR clients of one trace must compile for one target, " + f"these name several ({named}); trace the kernel once per " + "target instead" + ) + return next(iter(targets), None) + + def begin_launch(self, call: LaunchCall) -> None: + """Start one traced launch: a fresh Launch and per-launch state, then + every client's begin_launch. + + While a launch begun on another host thread is still in flight, this + raises RuntimeError before changing anything or telling any client: + concurrent launches of one trace are not supported. If a client's + begin_launch raises, the clients whose begin_launch was called get + abort_launch and the exception propagates; no launch is left open. + """ + self._claim_launch() + # Every launch gets its own Launch, so the entries TraceInterface + # appends to `launches` stay distinct and tilelens.clear() releases + # the tensors an interpreted run recorded in them (an IR-only launch + # records none). + self.launch = Launch() + self._reset_launch_state() + begun: list[Client] = [] + try: + for client in self.clients.values(): + begun.append(client) + client.begin_launch(call) + except BaseException as exc: + self._abort_clients(begun, exc) + raise + + def _claim_launch(self) -> None: + thread = threading.get_ident() + with self._owner_lock: + if self._launch_owner not in (None, thread): + raise RuntimeError( + "this trace is already running a launch on another host " + "thread; concurrent launches of one traced kernel are not " + "supported." + ) + self._launch_owner = thread + + def _release_launch(self) -> None: + with self._owner_lock: + if self._launch_owner == threading.get_ident(): + self._launch_owner = None + + def abort_launch(self, exc: BaseException) -> None: + """Deliver ``exc`` to every client's abort_launch. + + Nothing is sent once finalize has started: every client is finalized + by then, and each launch ends in either finalize or abort. Nor is + anything sent while another host thread's launch is in flight (ours + has ended already). A failing hook never replaces ``exc``, which the + caller re-raises; the failure is attached to it as a note. Only a + hook's KeyboardInterrupt or SystemExit is raised, after every client + got the abort. + """ + thread = threading.get_ident() + with self._owner_lock: + if self._launch_owner not in (None, thread): + return + if self._finalize_started: + self._launch_owner = None + return + # Held until every client got the abort, so no other thread's + # begin_launch resets the state in between. + self._launch_owner = thread + self._abort_clients(list(self.clients.values()), exc) + + def _abort_clients(self, clients: list[Client], exc: BaseException) -> None: + interrupt: BaseException | None = None + try: + for client in clients: + try: + client.abort_launch(exc) + except Exception as hook_exc: + message = ( + f"{type(client).__name__}.abort_launch raised " + f"{type(hook_exc).__name__}: {hook_exc}" + ) + if hasattr(exc, "add_note"): + exc.add_note(message) + else: # Python 3.10 + warnings.warn(message, RuntimeWarning, stacklevel=3) + except BaseException as hook_exc: + if interrupt is None: + interrupt = hook_exc + finally: + self._release_pinned() + self._release_launch() + if interrupt is not None: + raise interrupt from exc @contextmanager def patch_warmup( @@ -225,19 +707,259 @@ def vote(warmup, args, kwargs): finally: gate.remove(jit_fn) + @contextmanager + def ir_capture(self, jit_fn, *, real_args: RealArgs | None = None): + """Route every ``jit_fn.run`` call through a host compile, then + IR-client dispatch. Nothing launches: a call returns the kernel + compiled on the host, or None when it failed to compile. + + Only the traced JITFunction instance is touched (an instance attribute, + restored on exit). Autotuner/Heuristics layers reach it through their + ``fn.run`` calls, so every config they warm up goes through the + capture; IR clients get one event per distinct (specialization, + binding fingerprint) per traced launch (see Client.before_launch). + The compile never enters the original ``run``: ``self.compiler`` + compiles the call on the host for the IR clients' target + (``ir_target``), no driver or device involved, so the user's + pre_run_hooks never fire. A callable grid is resolved once more per + captured call, for the events and the fingerprint; one that raises + (or gives no 1-3 dim grid) is recorded as an unknown grid + (``resolved_grid`` None), never raised here. Compiles see the + arguments ``real_args`` maps the call to; events and fingerprints + describe the call as made. + + A failing host compile is delivered through compile_failed, never + raised. A host compile that could not run at all raises (or is + caused by) HostCompileUnavailable, which tells it from a kernel's + compile error; a compile refused while the language is patched is + delivered as a LanguagePatchedError. Yields the CaptureWindow. A + target spec that names no target, or IR clients that name different + targets, raise ValueError here, before anything is compiled. + + The one host compile error raised instead is a call that does not + bind the kernel's parameters (host_compile.bind_failed: a missing, + extra or misnamed argument, or a call the JIT cannot key, e.g. an + unhashable constexpr value): the JIT's binder or its cache key + raised it, as ``JITFunction.run`` does for the call on any device, + before any target or compile had a say, so the call raises that very + exception before any client hears of the call. An option the + target's backend does not know (a keyword that names no parameter) + is no bind failure: another backend may know it, so it is a compile + failure like any other (host_compile.unknown_options names it). + + On exit the window settles Launch.grid: the grid every kernel + compiled during the launch shares (what the launch would have + used), or None when configs disagree on it (a skipped autotuned + launch picks no config). Each event keeps its own ``resolved_grid``. + + Calls from other host threads pass through untouched, and a second + capture of the same jit_fn by another trace or thread is refused: + concurrent traced launches sharing a JITFunction are unsupported, and + a host compile refuses to start while an interpreted traced launch + has triton.language patched. + """ + window = CaptureWindow() + current = getattr(jit_fn, "run", None) + if current is None or not self.ir_clients(): + yield window + return + owner = getattr(current, "_tilelens_ir_capture", None) + thread = threading.get_ident() + if owner is not None: + manager, owner_thread, outer = owner + if manager is not self or owner_thread != thread: + raise RuntimeError( + f"{jit_fn!r} is already being captured by another traced " + "launch; concurrent traced launches sharing one " + "JITFunction are not supported." + ) + # Nested in our own capture: the outer window stays in charge. + yield outer + return + target = self.ir_target() + orig_run = current + + def run(*args, grid, warmup, **kwargs): + if threading.get_ident() != thread: + # Another host thread's launch (e.g. a peer trace's warmup + # compile) is not ours to capture. + return orig_run(*args, grid=grid, warmup=warmup, **kwargs) + return self._captured_run( + window, target, jit_fn, real_args, args, kwargs, grid + ) + + run._tilelens_ir_capture = (self, thread, window) # type: ignore[attr-defined] + with _instance_attr(jit_fn, "run", run): + yield window + self._settle_launch_grid() + + def _captured_run(self, window, target, jit_fn, real_args, args, kwargs, grid): + if real_args is None: + compile_args, compile_kwargs = args, kwargs + else: + compile_args, compile_kwargs = real_args(jit_fn, args, kwargs) + bound_args = _bind_launch_args(jit_fn, args, kwargs) + resolved_grid = _resolve_grid(grid, bound_args) + try: + _refuse_patched_language() + kernel = self.compiler.compile( + jit_fn, compile_args, compile_kwargs, target=target + ) + except Exception as exc: + if bind_failed(exc): + # The call's own error, whatever the target: raised as the + # untraced JITFunction.run raises it. + raise + window.failures.append(exc) + self._compile_failed( + target, jit_fn, args, kwargs, grid, exc, bound_args, resolved_grid + ) + return None + window.compiled += 1 + fingerprint = self._binding_fingerprint(jit_fn, args, kwargs, resolved_grid) + key = (target, _specialization(kernel), fingerprint) + if key not in self._delivered: + self._delivered.add(key) + event = self._launch_event( + jit_fn, + args, + kwargs, + grid, + kernel, + target=target, + bound_args=bound_args, + resolved_grid=resolved_grid, + ) + self._record_compiled_grid(event) + clients = self.ir_clients() + self._dispatch_ir("before_launch", event, clients) + self._dispatch_ir("after_launch", event, clients) + return kernel + + def _binding_fingerprint(self, jit_fn, args, kwargs, resolved_grid) -> Hashable: + """The binding part of the dedup key (see Client.before_launch). + Arguments to tl.constexpr parameters are left out: Triton hashes + each into the kernel, so the specialization already tells them + apart.""" + pinned = self._pinned + cached = self._constexprs.get(id(jit_fn)) + if cached is None or cached[0] is not jit_fn: + cached = self._constexprs[id(jit_fn)] = (jit_fn, *_constexpr_params(jit_fn)) + _, positions, names = cached + return ( + tuple( + None if index in positions else _fingerprint_value(arg, pinned) + for index, arg in enumerate(args) + ), + tuple( + sorted( + (name, _fingerprint_value(value, pinned)) + for name, value in kwargs.items() + if name not in names + ) + ), + _grid_fingerprint(resolved_grid, pinned), + ) + + def _compile_failed( + self, target, jit_fn, args, kwargs, grid, error, bound_args, resolved_grid + ): + # Once per launch for each failing call (constexprs included, as no + # specialization tells configs apart here): the same call compiled + # again in the launch is not news. + pinned = self._pinned + call = ( + tuple(_fingerprint_value(arg, pinned) for arg in args), + tuple( + sorted( + (name, _fingerprint_value(value, pinned)) + for name, value in kwargs.items() + ) + ), + ) + if (target, call) in self._failed: + return + self._failed.add((target, call)) + event = self._launch_event( + jit_fn, + args, + kwargs, + grid, + None, + error=error, + target=target, + bound_args=bound_args, + resolved_grid=resolved_grid, + ) + self._dispatch_ir("compile_failed", event, self.ir_clients()) + + @staticmethod + def _launch_event( + jit_fn, + args, + kwargs, + grid, + kernel, + error=None, + *, + target: Any = None, + bound_args: dict[str, Any] | None = None, + resolved_grid: Any = _MISSING, + ) -> LaunchEvent: + # ``bound_args`` / ``resolved_grid``: already computed for the call. + if bound_args is None: + bound_args = _bind_launch_args(jit_fn, args, kwargs) + if resolved_grid is _MISSING: + resolved_grid = _resolve_grid(grid, bound_args) + return LaunchEvent( + jit_fn=jit_fn, + args=tuple(args), + kwargs=MappingProxyType(dict(kwargs)), + grid=grid, + resolved_grid=resolved_grid, + bound_args=MappingProxyType(bound_args), + kernel=kernel, + specialization=None if kernel is None else _specialization(kernel), + error=error, + target=target, + ) + + def _record_compiled_grid(self, event: LaunchEvent) -> None: + # Launch.tensors is not filled from the binding: + # an interpreted run records (arg_callback) the host copies its eager + # clients' records point into, and an IR-only launch records none, so + # no device tensor outlives its launch in tilelens.launches. IR + # clients keep the tensor facts they need in their own records. + if event.resolved_grid is not None: + with self._lock_context(): + self._compiled_grids.add(event.resolved_grid) + + def _settle_launch_grid(self) -> None: + # See ir_capture: the grid every compiled kernel shares, else None. + grids = self._compiled_grids + resolved = next(iter(grids)) if len(grids) == 1 else None + with self._lock_context(): + self.launch.grid = resolved + + @staticmethod + def _dispatch_ir(hook: str, event: LaunchEvent, clients) -> None: + # Runs on the launching host thread, never on interpreter workers. + for client in clients: + getattr(client, hook)(event) + @contextmanager def patch_run(self, fn, frontend_name: str): frontend = get_frontend(frontend_name) namespaces = frontend.namespaces - # Every launch gets its own Launch, so the entries TraceInterface - # appends to `launches` stay distinct. - self.launch = Launch() + # IR clients take no part in op/loop registration: their empty + # callbacks would otherwise replace an interpreting peer's patches. + interpreting = self.interpreting_clients() with patch_calls(frontend_name): lang_patched = False try: # Collect all for-loop callbacks from clients all_loop_callbacks = [] - for client in self.clients.values(): + for client in interpreting: for namespace, attrs in namespaces.items(): # patch ops for attr, op in attrs.items(): callbacks = client.register_op_callback(op) @@ -265,48 +987,56 @@ def patch_run(self, fn, frontend_name: str): def pre_run_callback(self, fn: Callable) -> bool: with self._lock_context(): - rets = [client.pre_run_callback(fn) for client in self.clients.values()] + rets = [c.pre_run_callback(fn) for c in self.interpreting_clients()] return all(rets) if rets else True def post_run_callback(self, fn: Callable) -> bool: with self._lock_context(): - rets = [client.post_run_callback(fn) for client in self.clients.values()] - return any(rets) + rets = [c.post_run_callback(fn) for c in self.interpreting_clients()] + # With no interpreting voter, keep running the whole grid. + return any(rets) if rets else True def finalize(self) -> None: - with self._lock_context(): - self.launch.records = [] - # Finalize every client even if a peer raises (SystemExit - # included), so none carries this launch's state into the next; - # then re-raise the first failure. - first_exc: BaseException | None = None - for client in self.clients.values(): - try: - # client may introduce tensors not declared in kernel args (e.g. tracer recording a tensor allocation) - self.launch.tensors.update(getattr(client, "tensors", []) or []) - self.launch.records += client.finalize() - except BaseException as exc: - if first_exc is None: - first_exc = exc - if first_exc is not None: - raise first_exc + """Finalize every client into self.launch. This ends the launch: + another host thread may begin the next one right after.""" + try: + with self._lock_context(): + self._finalize_started = True + self.launch.records = [] + # Finalize every client even if a peer raises (SystemExit + # included), so none carries this launch's state into the + # next; then re-raise the first failure. + first_exc: BaseException | None = None + for client in self.clients.values(): + try: + # client may introduce tensors not declared in kernel args (e.g. tracer recording a tensor allocation) + self.launch.tensors.update(getattr(client, "tensors", []) or []) + self.launch.records += client.finalize() + except BaseException as exc: + if first_exc is None: + first_exc = exc + if first_exc is not None: + raise first_exc + finally: + self._release_pinned() + self._release_launch() def arg_callback(self, name, arg, arg_cvt): with self._lock_context(): if hasattr(arg, "data_ptr"): self.launch.tensors.add(arg) - for client in self.clients.values(): + for client in self.interpreting_clients(): client.arg_callback(name, arg, arg_cvt) def grid_callback(self, grid: tuple[int]): with self._lock_context(): self.launch.grid = grid - for client in self.clients.values(): + for client in self.interpreting_clients(): client.grid_callback(grid) def grid_idx_callback(self, grid_idx: tuple[int, ...]): with self._lock_context(): - for client in self.clients.values(): + for client in self.interpreting_clients(): client.grid_idx_callback(grid_idx) # --- For-loop callback management --- diff --git a/tilelens/core/config.py b/tilelens/core/config.py index 9163d60bd..de7861f8b 100644 --- a/tilelens/core/config.py +++ b/tilelens/core/config.py @@ -1,6 +1,15 @@ import os +# The target IR mode compiles kernels for unless a client or +# TILELENS_IR_TARGET says otherwise: GPUTarget("cuda", 89, 32), so a +# result never depends on the machine it was computed on. sm89 (Ada) is the +# first capability Triton compiles fp8e4nv for, and still has no native TMA +# (sm90+), so tensor descriptors are lowered to pointer math the reader +# analyzes. +DEFAULT_IR_TARGET = "cuda:89" + + def _get_env(env: str, default: str) -> str: """Prefer TileLens settings, falling back to the former variable names.""" if env.startswith("TILELENS_"): @@ -54,6 +63,14 @@ class Config: - sanitizer_report_max_segments: SANITIZER_REPORT_MAX_SEGMENTS, max number of address segments to list verbatim in the OOB report before truncating to a head/tail summary. Affects display only (min 2). + - ir_target: TILELENS_IR_TARGET, the target IR mode compiles kernels for + when the IR client names none (DEFAULT_IR_TARGET, "cuda:89", if unset): + e.g. "cuda:90" or "hip:gfx942", see + tilelens.core.host_compile.parse_ir_target. An IR client's own target + (Client.ir_target) wins over it; a value that names no target is + reported when a traced launch compiles. The IR target also wins over + TRITON_OVERRIDE_ARCH, which retargets only the JIT's own (device) + compiles. """ def __init__(self) -> None: @@ -87,6 +104,26 @@ def reset(self) -> None: self.sanitizer_report_max_segments: int = _get_int_env( "SANITIZER_REPORT_MAX_SEGMENTS", 8, minimum=2 ) + self.ir_target: str = _get_env("TILELENS_IR_TARGET", DEFAULT_IR_TARGET) config = Config() + + +# The one Triton minor release IR mode runs on: the host compile +# (tilelens.core.host_compile) uses its private API, so on any other release +# it refuses. +IR_TRITON_RELEASE = "3.8" + + +def ir_triton_unsupported() -> str | None: + """Why IR mode cannot run on the installed Triton; None on Triton 3.8.x.""" + import triton + + version = triton.__version__ + if version.split(".")[:2] == IR_TRITON_RELEASE.split("."): + return None + return ( + f"IR mode supports Triton {IR_TRITON_RELEASE}.x only; the installed " + f"Triton is {version}" + ) diff --git a/tilelens/core/host_compile.py b/tilelens/core/host_compile.py new file mode 100644 index 000000000..4826e65d4 --- /dev/null +++ b/tilelens/core/host_compile.py @@ -0,0 +1,885 @@ +"""Host compile: one JITFunction call compiled for a GPUTarget without a GPU. + +IR mode reads the kernels Triton compiles, but the analysis is CPU work, so +the compile is too: :class:`HostCompiler` binds a call with the JIT's own +binder (``create_function_from_signature`` for the target's backend: the +signature types, i32 / i64 / u64 integers by value, the equal-to-1 and +divisibility specializations, tuples, tensor descriptors, constexprs, +``do_not_specialize``), packs it with ``JITFunction._pack_args`` and +compiles an ``ASTSource`` for the target, exactly as ``JITFunction.run`` +would on a device of that target. No driver is queried and nothing is +loaded or launched: no ``get_current_device``, no stream, no +``_init_handles``. The JIT runtime's own hooks are not called either (the +function's ``pre_run_hooks``, ``knobs.runtime.jit_cache_hook`` / +``jit_post_compile_hook``, async compile mode): the host compile is no JIT +run. + +The target is the caller's, whatever the machine has. While a thread host +compiles, Triton's ``driver.active`` answers that thread's target query +(``get_current_target``: what ``tl.target_info.is_cuda()`` / +``cuda_capability_geq()`` / ``is_hip()`` and the front end's own target +checks read) with the compile's target, and refuses any device query with +:class:`HostCompileUnavailable`; other threads, and this one outside its +compile, see Triton's own driver. The compile options name the target's +arch, so ``TRITON_OVERRIDE_ARCH`` does not reach a host compile either. A +compile that raises after its front end asked the driver anything (the +target, or a device query it refused and the kernel's code may have caught), +or after an earlier compile of the same kernel for the target did, says so +(:func:`target_queried`): what failed may be the target's answer. A call +that does not bind the kernel's parameters says that instead +(:func:`bind_failed`): the JIT raises the same for it on any device. + +The pipeline stops after the backend's ``ttir`` passes (the same passes +``triton.compile`` runs): the front end, then TTIR, never ``ttgir`` / +``llir`` / the binary. Such a truncated compile is kept in memory only (a +``HostKernel``), never in Triton's on-disk cache, whose entries must hold +the whole pipeline. What only ``triton.compile``'s whole pipeline applies +to the TTIR is refused as :class:`HostCompileUnavailable` rather than +silently left out: ``TRITON_KERNEL_OVERRIDE``, ``USE_IR_LOC``, an +``ir_override`` compile option, and a custom pipeline +(``knobs.runtime.add_stages_inspection_hook``, which also keys the +kernel). ``TRITON_KERNEL_DUMP`` changes no stage; a host compile is just +not dumped. + +A ``HostKernel`` has ``.asm`` (``"ttir"`` -> text), ``.metadata`` (a +namedtuple: ``target``, ``name``, the compile options, ...) and ``.hash``, +the specialization: what ``triton.compile`` names the kernel for that +target. + +The APIs used are private to Triton. :func:`triton_api` checks that they +exist and that the front end's target queries can be scoped, and the first +compile for each target host-compiles a small built-in kernel first, so a +changed API fails as :class:`HostCompileUnavailable` naming it rather than +as an error blamed on the user's kernel. They are Triton 3.8's: on any +other release triton_api refuses (tilelens.core.config.IR_TRITON_RELEASE). +A ``CompiledKernel.__del__`` that unloads through the driver, which may run +in the middle of a host compile, reaches the machine's driver (see +_UnloadOnlyUtils). + +Importing this module does not import Triton. +""" + +from __future__ import annotations + +import functools +import hashlib +import inspect +import linecache +import re +import threading +from collections import namedtuple +from collections.abc import Hashable, Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass, field +from types import MappingProxyType, SimpleNamespace +from typing import Any + +from . import config as config_module +from .config import DEFAULT_IR_TARGET + + +class HostCompileUnavailable(RuntimeError): + """The host compile cannot run: the installed Triton lacks (or changed) + an API it uses, or the compile asked for something only a device has. + Never a kernel's own compile error.""" + + +# The attributes a host compile's exception carries when the front end had +# asked the driver before it was raised (see target_queried), and when the +# call did not bind the kernel's parameters (see bind_failed). +_TARGET_QUERIED = "_tilelens_target_queried" +_BIND_FAILED = "_tilelens_bind_failed" +# The keyword arguments no option of the target's backend names (see +# unknown_options). +_UNKNOWN_OPTIONS = "_tilelens_unknown_options" + + +def target_queried(exc: BaseException | None) -> bool: + """Whether ``exc`` was raised by a host compile whose front end had + asked Triton's driver anything before it failed, so the failure may + follow from the target's answer: the target + (``driver.active.get_current_target()``: ``tl.target_info``, the + tensor-descriptor lowering's native-TMA check, a constexpr function + asking the driver), or anything else, which the host compile refuses + (a device query the kernel's code caught, falling back to an answer of + its own, is still a question about the device). Also true when an + earlier compile of the same kernel for the same target by the same + HostCompiler had asked: the kernel's code may keep the answer (a memo) + and not ask again. False for any other exception, a bind failure + (bind_failed) included. What the compile options derive from the target + (e.g. its fp8 types, or whether ``num_ctas > 1`` is allowed) is no + query. An answer the kernel's code keeps from a compile this + HostCompiler did not run (another trace's, another target's, the + untraced program's), and never asks for again, cannot be seen.""" + return exc is not None and getattr(exc, _TARGET_QUERIED, False) is True + + +def bind_failed(exc: BaseException | None) -> bool: + """Whether ``exc`` was raised by a host compile while binding the call + to the kernel's parameters (the JIT's binder: a missing or unexpected + argument, an argument of a type Triton cannot pass) or keying it + (``compute_cache_key``: e.g. an unhashable constexpr value): no target, + and no compile, decides it, so ``JITFunction.run`` raises the same for + the call on any device, and a traced launch raises it as is (see + tilelens.core.client.ClientManager.ir_capture). False for any other + exception, e.g. the KeyError for a keyword that names neither a + parameter nor an option of the target's backend (another backend may + know it, see unknown_options), and for a HostCompileUnavailable.""" + return exc is not None and getattr(exc, _BIND_FAILED, False) is True + + +def _mark(exc: BaseException, attr: str) -> None: + try: + setattr(exc, attr, True) + except Exception: # an exception type that takes no attribute + pass + + +def _mark_target_queried(exc: BaseException) -> None: + _mark(exc, _TARGET_QUERIED) + + +def _mark_bind_failed(exc: BaseException) -> None: + _mark(exc, _BIND_FAILED) + + +def unknown_options(exc: BaseException | None) -> tuple[str, ...]: + """The call's keyword arguments that name neither a parameter of the + kernel nor a compile option of the target's backend, when ``exc`` is + the KeyError the JIT raises for them (``JITFunction._pack_args``); () + for any other exception. Such a call fails on every device whose + backend does not know them (a misspelled option: on every GPU), yet + another backend may know them (e.g. HIP's ``waves_per_eu``), so the + compile failure is the target's, not the call's (not bind_failed).""" + names = getattr(exc, _UNKNOWN_OPTIONS, ()) if exc is not None else () + return names if isinstance(names, tuple) else () + + +def _mark_unknown_options( + exc: BaseException, jit_fn: Any, backend: Any, kwargs: Mapping[str, Any] +) -> None: + # JITFunction._pack_args's own check: a keyword in neither the parsed + # options nor the signature. Parsing again is how the JIT reads the + # options; if parsing is what failed, nothing is marked. + try: + known = vars(backend.parse_options(dict(kwargs))) + except Exception: + return + params = {param.name for param in jit_fn.params} + names = tuple(k for k in kwargs if k not in known and k not in params) + if names: + try: + setattr(exc, _UNKNOWN_OPTIONS, names) + except Exception: + pass + + +def _mark_call_error(exc: BaseException) -> None: + """Mark ``exc``, raised while binding or keying the call, as the call's + own error (bind_failed), unless the host compile could not run.""" + if host_compile_unavailable(exc) is None: + _mark_bind_failed(exc) + + +def host_compile_unavailable(exc: BaseException) -> HostCompileUnavailable | None: + """The HostCompileUnavailable behind ``exc``: ``exc`` itself, or one it + was raised from or while handling (Triton's code generator re-raises + what a kernel's code raised as a CompilationError from it); None if + there is none, i.e. ``exc`` is the kernel's own compile error.""" + seen: set[int] = set() + link: BaseException | None = exc + while link is not None and id(link) not in seen: + if isinstance(link, HostCompileUnavailable): + return link + seen.add(id(link)) + # The chain a traceback shows: the cause, else the unsuppressed context. + if link.__cause__ is not None: + link = link.__cause__ + else: + link = None if link.__suppress_context__ else link.__context__ + return None + + +# ─────────────────────────── targets ─────────────────────────── + +_TARGET_FORMS = ( + "'cuda:' (e.g. 'cuda:80', 'cuda:90'), " + "'hip:' (e.g. 'hip:gfx942'), either optionally followed by " + "':', or a triton.backends.compiler.GPUTarget" +) +_RE_CUDA = re.compile(r"cuda:(\d+)(?::(\d+))?") +# gfx: gfx90a, gfx942, gfx1100, ... +_RE_GFX = r"gfx\d{1,2}[0-9a-z]{2}" +_RE_HIP = re.compile(rf"hip:({_RE_GFX})(?::(\d+))?") +# Volta: no Triton release targets an older NVIDIA GPU. +_MIN_CUDA_CAPABILITY = 70 + + +def _is_int(value: Any) -> bool: + # A bool is an int, but no capability or warp size. + return isinstance(value, int) and not isinstance(value, bool) + + +def _checked_target(target: Any, spec: Any) -> Any: + backend, arch, warp_size = target.backend, target.arch, target.warp_size + valid = ( + backend == "cuda" + and _is_int(arch) + and arch >= _MIN_CUDA_CAPABILITY + or backend == "hip" + and isinstance(arch, str) + and re.fullmatch(_RE_GFX, arch) is not None + ) + if not valid or not _is_int(warp_size) or warp_size <= 0: + raise ValueError( + f"invalid IR target {spec!r}: expected {_TARGET_FORMS}; a CUDA " + f"compute capability is at least {_MIN_CUDA_CAPABILITY}, a warp " + "size positive" + ) + return target + + +@functools.lru_cache(maxsize=64) +def _parse_target_spec(spec: str) -> Any: + from triton.backends.compiler import GPUTarget + + text = spec.strip().lower() + if match := _RE_CUDA.fullmatch(text): + capability, warp_size = match.group(1), match.group(2) + target = GPUTarget("cuda", int(capability), int(warp_size) if warp_size else 32) + elif match := _RE_HIP.fullmatch(text): + gfx, warp_size = match.group(1), match.group(2) + # CDNA (gfx9*) runs 64-wide wavefronts, RDNA 32-wide. + default = 64 if gfx.startswith("gfx9") else 32 + target = GPUTarget("hip", gfx, int(warp_size) if warp_size else default) + else: + raise ValueError(f"invalid IR target {spec!r}: expected {_TARGET_FORMS}") + return _checked_target(target, spec) + + +def parse_ir_target(spec: Any) -> Any: + """The ``GPUTarget`` an IR target spec names: a ``GPUTarget`` itself, or + a string such as ``"cuda:89"``, ``"cuda:90"``, ``"hip:gfx942"`` or + ``"hip:gfx1100:32"``. Raises ValueError for anything else, a CUDA + compute capability below 70 or a warp size that is not positive + included.""" + from triton.backends.compiler import GPUTarget + + if isinstance(spec, GPUTarget): + return _checked_target(spec, spec) + if isinstance(spec, str): + return _parse_target_spec(spec) + raise ValueError(f"invalid IR target {spec!r}: expected {_TARGET_FORMS}") + + +def format_ir_target(target: Any) -> str: + """A GPUTarget as the spec parse_ir_target reads back, e.g. + ``"cuda:89"``; the warp size only where it is not the default.""" + backend = getattr(target, "backend", None) + arch = getattr(target, "arch", None) + warp_size = getattr(target, "warp_size", None) + if backend not in ("cuda", "hip"): + return repr(target) + spec = f"{backend}:{arch}" + try: + default = _parse_target_spec(spec).warp_size + except ValueError: + return repr(target) + return spec if warp_size == default else f"{spec}:{warp_size}" + + +def resolve_ir_target(requested: Any = None) -> Any: + """The ``GPUTarget`` for a client's ``ir_target``: ``requested`` when it + is set, else the configured default (``tilelens.config.ir_target``, from + ``TILELENS_IR_TARGET``, else ``"cuda:89"``).""" + if requested is not None: + return parse_ir_target(requested) + spec = config_module.config.ir_target + try: + return parse_ir_target(spec) + except ValueError as exc: + raise ValueError( + f"tilelens.config.ir_target (TILELENS_IR_TARGET) is {spec!r}, which " + f"is no IR target: {exc}" + ) from None + + +def default_ir_target() -> Any: + """``GPUTarget("cuda", 89, 32)`` (sm89 is the first capability Triton + compiles fp8e4nv for).""" + return parse_ir_target(DEFAULT_IR_TARGET) + + +def _target_arch(target: Any) -> str | None: + """``target``'s ``arch`` compile option, as its backend's + parse_options derives it unless TRITON_OVERRIDE_ARCH says otherwise; + None for a backend this module does not know.""" + if target.backend == "cuda": + return f"sm{target.arch}" + if target.backend == "hip": + return str(target.arch) + return None + + +# ─────────────────── the target Triton's front end sees ─────────────────── + +_MISSING = object() + + +class _TargetDriver: + """``triton.runtime.driver.active`` on a thread while it host-compiles: + it answers the target query with the compile's target, so Triton's + front end (``tl.target_info``, its own target checks, a user's + constexpr function) sees the target the kernel is compiled for, never + the machine's device. The host has no device, stream or device + property to give, so anything else is refused, with one exception: + ``CompiledKernel.__del__`` unloads a module through the driver, see + _UnloadOnlyUtils.""" + + def __init__(self, target: Any, set_aside: Any) -> None: + self._target = target + # A context manager factory that sets this thread's scope aside + # (_ScopedActiveDriver.set_aside), for unloading (_UnloadOnlyUtils). + self._set_aside = set_aside + # Whether the driver was asked anything (see target_queried): the + # target, or a question it refuses, which the kernel's code may + # catch and answer itself (e.g. "no big shared memory"), so a + # failure after it may be the device's all the same. + self.queried = False + + def get_current_target(self) -> Any: + self.queried = True + return self._target + + @property + def utils(self) -> Any: + return _UnloadOnlyUtils(self, self._set_aside) + + def _refuse(self, name: str) -> Any: + self.queried = True + raise HostCompileUnavailable( + f"compiling for {format_ir_target(self._target)} on the host, " + f"Triton asked its driver for {name!r}: a host compile has no " + "device to ask and answers only the target query" + ) + + def __getattr__(self, name: str) -> Any: + if name.startswith("__"): + raise AttributeError(name) + return self._refuse(name) + + +class _UnloadOnlyUtils: + """``driver.active.utils`` on a thread while it host-compiles. + ``CompiledKernel.__del__`` unloads a loaded module through it: a kernel + a real launch loaded can be collected on any thread, in the middle of a + host compile too. ``unload_module`` releases the module through the + driver the thread has outside its compile, which loaded it; it asks + nothing about the device, so the compile is not marked as having asked + (see target_queried). Anything else is refused as any device query.""" + + def __init__(self, scoped: _TargetDriver, set_aside: Any) -> None: + self._scoped = scoped + self._set_aside = set_aside + + def unload_module(self, module: Any) -> Any: + from triton.runtime.driver import driver + + with self._set_aside(): + return driver.active.utils.unload_module(module) + + def __getattr__(self, name: str) -> Any: + if name.startswith("__"): + raise AttributeError(name) + return self._scoped._refuse(f"utils.{name}") + + +class _ScopedActiveDriver: + """Thread-scoped ``driver.active`` (see _TargetDriver). + + While any thread host-compiles, the DriverConfig class's ``active`` + property is wrapped: a compiling thread gets its _TargetDriver, every + other thread (and the compiling one outside its compile) whatever + ``active`` was before, i.e. Triton's own driver or a test's stand-in. + The last compile to end puts the class attribute back, unless someone + replaced the wrapper in the meantime. Replacing the process-wide active + driver instead would hand the target driver to another thread's real + launch. + """ + + def __init__(self) -> None: + self._lock = threading.Lock() + self._local = threading.local() + self._depth = 0 + self._owner: Any = None + self._previous: Any = _MISSING + self._wrapper: Any = None + + @contextmanager + def targeting(self, config_cls: type, target: Any) -> Iterator[_TargetDriver]: + """Scope ``driver.active`` on this thread to a _TargetDriver for + ``target``, which is yielded.""" + with self._lock: + if self._depth == 0: + self._install(config_cls) + self._depth += 1 + saved = getattr(self._local, "driver", None) + scoped = self._local.driver = _TargetDriver(target, self.set_aside) + try: + yield scoped + finally: + self._local.driver = saved + with self._lock: + self._depth -= 1 + if self._depth == 0: + self._uninstall() + + @contextmanager + def set_aside(self) -> Iterator[None]: + """This thread's scope set aside: ``driver.active`` is what it is + outside every host compile of the thread.""" + saved = getattr(self._local, "driver", None) + self._local.driver = None + try: + yield + finally: + self._local.driver = saved + + def _install(self, config_cls: type) -> None: + fallback = inspect.getattr_static(config_cls, "active") + local = self._local + + def active(config: Any) -> Any: + scoped = getattr(local, "driver", None) + if scoped is not None: + return scoped + return fallback.__get__(config, type(config)) + + self._owner = config_cls + self._previous = config_cls.__dict__.get("active", _MISSING) + self._wrapper = property(active) + setattr(config_cls, "active", self._wrapper) + + def _uninstall(self) -> None: + owner, wrapper = self._owner, self._wrapper + if owner is not None and owner.__dict__.get("active") is wrapper: + if self._previous is _MISSING: + delattr(owner, "active") + else: + setattr(owner, "active", self._previous) + self._owner, self._previous, self._wrapper = None, _MISSING, None + + +_SCOPED_DRIVER = _ScopedActiveDriver() + + +# ─────────────────────────── Triton's API ─────────────────────────── + + +def _unavailable(version: str, what: str) -> HostCompileUnavailable: + return HostCompileUnavailable( + f"IR mode compiles kernels on the host with Triton's private compile " + f"API, and on Triton {version} {what}" + ) + + +@functools.lru_cache(maxsize=1) +def triton_api() -> SimpleNamespace: + """The Triton internals the host compile uses, checked for presence, + and the front end's target queries checked to answer a scoped target + (see _TargetDriver). Raises HostCompileUnavailable naming what is + missing or does not behave so; a failure is not cached.""" + import triton + + unsupported = config_module.ir_triton_unsupported() + if unsupported is not None: + raise HostCompileUnavailable(unsupported) + + def missing(what: str) -> HostCompileUnavailable: + return _unavailable(triton.__version__, f"it lacks {what}") + + try: + from triton import knobs + from triton._C.libtriton import get_cache_invalidating_env_vars, ir + from triton.backends.compiler import GPUTarget + from triton.compiler import ASTSource, get_cache_key, make_backend + from triton.compiler.compiler import filter_traceback + from triton.runtime.driver import driver + from triton.runtime.jit import ( + JITFunction, + compute_cache_key, + create_function_from_signature, + ) + except ImportError as exc: + raise missing(str(exc)) from exc + for owner, name, attr in ( + (ir, "triton._C.libtriton.ir", "context"), + (ir, "triton._C.libtriton.ir", "load_dialects"), + (ASTSource, "ASTSource", "make_ir"), + (knobs.runtime, "knobs.runtime", "debug"), + (knobs.compilation, "knobs.compilation", "instrumentation_mode"), + ): + if not hasattr(owner, attr): + raise missing(f"{name}.{attr}") + if not isinstance(inspect.getattr_static(type(driver), "active", None), property): + raise missing( + "triton.runtime.driver.driver.active as a property of its class, " + "which the host compile scopes to answer the target query" + ) + for target in (GPUTarget("cuda", 80, 32), GPUTarget("cuda", 90, 32)): + try: + with _SCOPED_DRIVER.targeting(type(driver), target): + wrong = _unscoped_target_queries(target) + except Exception as exc: + raise _unavailable( + triton.__version__, + "its front end's target queries could not be asked " + f"({type(exc).__name__}: {exc})", + ) from exc + if wrong: + raise _unavailable( + triton.__version__, + f"its front end's target queries answer {wrong} while compiling " + f"for {format_ir_target(target)} on the host", + ) + return SimpleNamespace( + version=triton.__version__, + knobs=knobs, + ir=ir, + get_cache_invalidating_env_vars=get_cache_invalidating_env_vars, + GPUTarget=GPUTarget, + ASTSource=ASTSource, + get_cache_key=get_cache_key, + make_backend=make_backend, + compute_cache_key=compute_cache_key, + create_function_from_signature=create_function_from_signature, + filter_traceback=filter_traceback, + driver_config=type(driver), + JITFunction=JITFunction, + ) + + +def _unscoped_target_queries(target: Any) -> dict[str, Any]: + """The front end's target queries that do not answer ``target`` under + its scope, with what they answer.""" + from triton.language import target_info + from triton.language.semantic import TritonSemantic + + wrong: dict[str, Any] = {} + if (got := target_info.current_target()) != target: + wrong["tl.target_info.current_target()"] = got + # It reads nothing of the semantic, only the driver's target. + native = TritonSemantic._has_native_tma(None) + if native != (target.backend == "cuda" and target.arch >= 90): + wrong["TritonSemantic._has_native_tma()"] = native + return wrong + + +# Host-compiled before the first compile for each target: scalars only (it +# needs no tensor, and no name from triton.language); ``one`` takes the +# equal-to-1 constexpr specialization. Its source is registered with +# linecache under a name of its own, so the JIT reads it from there and not +# from this file, which may have changed on disk since it was imported. +_SELF_TEST_SOURCE = """\ +def _self_test_kernel(n, flag, one): + if flag: + n = n * one +""" +_SELF_TEST_FILE = "" + + +def _self_test_jit_function(api: SimpleNamespace) -> Any: + lines = _SELF_TEST_SOURCE.splitlines(keepends=True) + linecache.cache[_SELF_TEST_FILE] = ( + len(_SELF_TEST_SOURCE), + None, + lines, + _SELF_TEST_FILE, + ) + namespace: dict[str, Any] = {"__name__": __name__} + exec(compile(_SELF_TEST_SOURCE, _SELF_TEST_FILE, "exec"), namespace) + return api.JITFunction(namespace["_self_test_kernel"]) + + +@functools.lru_cache(maxsize=None) +def _self_test_target(target: Any) -> None: + """Host-compile the built-in _self_test_kernel for ``target`` through its TTIR; + raise HostCompileUnavailable if that fails. Only a success is cached.""" + api = triton_api() + try: + kernel = HostCompiler().compile( + _self_test_jit_function(api), + (5, True, 1), + {}, + target=target, + _self_test=True, + ) + text = kernel.asm["ttir"] + except Exception as exc: + raise _unavailable( + api.version, + f"a built-in test kernel failed to host-compile for " + f"{format_ir_target(target)} ({type(exc).__name__}: {exc})", + ) from exc + if "tt.func" not in text or "_self_test_kernel" not in text: + raise _unavailable( + api.version, + f"a built-in test kernel host-compiled for {format_ir_target(target)} " + "to no TTIR function", + ) + + +# ─────────────────────────── the artifact ─────────────────────────── + + +@dataclass(frozen=True, eq=False) +class HostKernel: + """A kernel compiled on the host through its TTIR (see the module + docstring): what a ``CompiledKernel`` holds of it, never loaded.""" + + # The specialization: what triton.compile names this kernel. + hash: str + name: str + # Stage -> text, every stage compiled, in pipeline order: "ttir", and + # any stage a backend runs before it. + asm: Mapping[str, str | bytes] = field(repr=False) + # A namedtuple, as CompiledKernel.metadata: "target", "name", "hash", + # the compile options and whatever the compiled stages added. + metadata: Any = field(repr=False) + + @property + def target(self) -> Any: + return self.metadata.target + + +# The stage a host compile stops after. +_TTIR = "ttir" + + +def _unsupported_pipeline( + api: SimpleNamespace, kwargs: Mapping[str, Any] +) -> str | None: + """What makes ``triton.compile`` change the TTIR (or the kernel's hash) + in a way only its whole pipeline applies, which the host compile does + not run; None when nothing does.""" + compilation = api.knobs.compilation + if getattr(compilation, "override", False): + return "TRITON_KERNEL_OVERRIDE (knobs.compilation.override)" + if getattr(compilation, "use_ir_loc", None): + return "USE_IR_LOC (knobs.compilation.use_ir_loc)" + if kwargs.get("ir_override"): + return "the 'ir_override' compile option" + if getattr(api.knobs.runtime, "add_stages_inspection_hook", None) is not None: + return "a custom pipeline (knobs.runtime.add_stages_inspection_hook)" + return None + + +def _check_used_globals(jit_fn: Any) -> None: + # JITFunction.run's check, for every kernel handed out: a kernel + # compiled before a global it reads changed is stale. + not_present = object() + for (name, _), (value, globals_dict) in jit_fn.used_global_vals.items(): + if (new := globals_dict.get(name, not_present)) != value: + raise RuntimeError( + f"Global variable {name} has changed since we compiled this " + f"kernel, from {value} to {new}" + ) + + +class HostCompiler: + """Host compiles with an in-process cache (one per trace: the + ClientManager's ``compiler``), keyed by the JIT's own specialization + key (``compute_cache_key``: the bound specialization and the call's + compile options) and the target.""" + + def __init__(self) -> None: + # (id(jit_fn), target) -> (jit_fn, backend, binder, key cache). + self._binders: dict[Hashable, tuple[Any, Any, Any, dict]] = {} + # (id(jit_fn), target) -> jit_fn, for each kernel a compile of which + # for the target asked the driver (see target_queried). + self._asked: dict[Hashable, Any] = {} + # (id(jit_fn), specialization key, target) -> (jit_fn, kernel). + self._kernels: dict[Hashable, tuple[Any, Any]] = {} + + def compile( + self, + jit_fn: Any, + args: tuple, + kwargs: Mapping[str, Any], + *, + target: Any, + _self_test: bool = False, + ) -> HostKernel: + """Compile the call ``jit_fn.run(*args, **kwargs)`` would compile on + a device of ``target`` (a GPUTarget), through its TTIR. Raises what + the JIT's bind, pack or compile raises (a bind failure marked as + such, see bind_failed; any other error marked when this compile, or + an earlier one of ``jit_fn`` for ``target``, had asked the driver, + see target_queried), or HostCompileUnavailable. (``_self_test``: the + built-in test compile, which skips the checks it is part of.)""" + api = triton_api() + for attr in ("signature", "params", "_pack_args", "used_global_vals"): + if not hasattr(jit_fn, attr): + raise HostCompileUnavailable( + f"cannot host-compile {jit_fn!r}: it has no {attr!r} " + f"(a JITFunction of Triton {api.version} has)" + ) + if not _self_test: + _self_test_target(target) + asked_key = (id(jit_fn), target) + with _SCOPED_DRIVER.targeting(api.driver_config, target) as scoped: + try: + return self._compile(api, jit_fn, args, kwargs, target, _self_test) + except Exception as exc: + asked_before = self._asked.get(asked_key) is jit_fn + if not bind_failed(exc) and (scoped.queried or asked_before): + _mark_target_queried(exc) + raise + finally: + if scoped.queried: + self._asked[asked_key] = jit_fn + + def _compile( + self, + api: SimpleNamespace, + jit_fn: Any, + args: tuple, + kwargs: Mapping[str, Any], + target: Any, + self_test: bool, + ) -> HostKernel: + backend, binder, key_cache = self._binder(api, jit_fn, target) + # What JITFunction.run adds to every call's options. + kwargs = dict(kwargs) + kwargs["debug"] = ( + kwargs.get("debug", getattr(jit_fn, "debug", None)) + or api.knobs.runtime.debug + ) + kwargs["instrumentation_mode"] = api.knobs.compilation.instrumentation_mode + # The target's arch as a compile option, which the backend's + # parse_options takes over TRITON_OVERRIDE_ARCH: a host compile is + # for the target it was asked for. Not where "arch" is the + # call's own (a launch option, or a kernel parameter). + arch = _target_arch(target) + if ( + arch is not None + and "arch" not in kwargs + and all(param.name != "arch" for param in jit_fn.params) + ): + kwargs["arch"] = arch + try: + bound_args, specialization, options = binder(*args, **kwargs) + except Exception as exc: + # The call's own error (see bind_failed); the backend only adds + # its tensor-alignment flags to the specialization. + _mark_call_error(exc) + raise + try: + cache_key = api.compute_cache_key(key_cache, specialization, options) + except Exception as exc: + # The call's own error too (e.g. an unhashable constexpr value): + # the key is the bound specialization and the call's options, + # which JITFunction.run keys the call by right after its binder, + # on any device. + _mark_call_error(exc) + raise + unsupported = None if self_test else _unsupported_pipeline(api, kwargs) + if unsupported is not None: + raise HostCompileUnavailable( + f"under {unsupported}, triton.compile changes the kernel in its " + "whole pipeline, and a host compile runs the front end and the " + "TTIR passes only" + ) + key = (id(jit_fn), cache_key, target) + cached = self._kernels.get(key) + if cached is not None and cached[0] is jit_fn: + kernel = cached[1] + else: + try: + options, signature, constexprs, attrs = jit_fn._pack_args( + backend, kwargs, bound_args, specialization, options + ) + except KeyError as exc: + _mark_unknown_options(exc, jit_fn, backend, kwargs) + raise + compiled_arch = getattr(options, "arch", arch) + if compiled_arch != arch: + raise HostCompileUnavailable( + f"the call's compile options name arch {compiled_arch!r}, " + f"not {arch!r} of the IR target {format_ir_target(target)} " + "(an 'arch' launch option, or TRITON_OVERRIDE_ARCH with a " + "kernel parameter named 'arch'), so its host compile would " + "not be for the target" + ) + source = api.ASTSource(jit_fn, signature, constexprs, attrs) + kernel = self._compile_source(api, source, backend, target, options) + self._kernels[key] = (jit_fn, kernel) + _check_used_globals(jit_fn) + return kernel + + def _binder(self, api: SimpleNamespace, jit_fn: Any, target: Any) -> tuple: + entry = self._binders.get((id(jit_fn), target)) + if entry is None or entry[0] is not jit_fn: + backend = api.make_backend(target) + binder = api.create_function_from_signature( + jit_fn.signature, jit_fn.params, backend + ) + entry = self._binders[(id(jit_fn), target)] = (jit_fn, backend, binder, {}) + return entry[1:] + + @staticmethod + def _compile_source( + api: SimpleNamespace, + source: Any, + backend: Any, + target: Any, + options: Any, + ) -> HostKernel: + """triton.compile's front half: the front end, then the backend's + stages through "ttir".""" + pipeline: dict[str, Any] = {} + backend.add_stages(pipeline, options, source.language) + names = list(pipeline) + if _TTIR not in names: + raise HostCompileUnavailable( + f"the {target.backend} backend of Triton {api.version} has no " + f"{_TTIR!r} stage to stop after (its stages: {', '.join(names)})" + ) + env_vars = api.get_cache_invalidating_env_vars() + key = api.get_cache_key(source, backend, options, env_vars) + digest = hashlib.sha256(key.encode("utf-8")).hexdigest() + metadata = { + "hash": digest, + "target": target, + **options.__dict__, + **env_vars, + "triton_version": api.version, + } + # Keep the context referenced until every module of it is gone. + context = api.ir.context() + api.ir.load_dialects(context) + backend.load_dialects(context) + codegen_fns = backend.get_codegen_implementation(options) + module_map = backend.get_module_map() + try: + module = source.make_ir(target, options, codegen_fns, module_map, context) + except Exception as exc: + api.filter_traceback(exc) + raise + asm: dict[str, str | bytes] = {} + for name in names[: names.index(_TTIR) + 1]: + module = pipeline[name](module, metadata) + asm[name] = module if isinstance(module, (str, bytes)) else str(module) + del module + # A later stage names the entry point; up to here it is the kernel's. + metadata.setdefault("name", source.name) + kernel_metadata = namedtuple( # type: ignore[misc] + "KernelMetadata", sorted(metadata) + )(**metadata) + del context + return HostKernel( + hash=digest, + name=metadata["name"], + asm=MappingProxyType(asm), + metadata=kernel_metadata, + ) diff --git a/tilelens/core/trace.py b/tilelens/core/trace.py index ca27e967a..aa4cde172 100644 --- a/tilelens/core/trace.py +++ b/tilelens/core/trace.py @@ -2,13 +2,14 @@ import inspect from contextlib import contextmanager from collections.abc import Callable +from types import MappingProxyType from typing import Any from ..utils.traceback_utils import CODE_KEYS, get_code_key from .config import config as cfg from ..clients import Sanitizer, Profiler, RaceDetector, Tracer from ..clients.race_detector.race_detector import NullRaceDetector -from .client import ClientManager, Client +from .client import ClientManager, Client, LaunchCall, LanguagePatchedError from .data import Launch import types @@ -21,6 +22,20 @@ def _without_warmup(kwargs: dict[str, Any]) -> dict[str, Any]: return {k: v for k, v in kwargs.items() if k != "warmup"} +def _launch_call( + jit_fn: Any, args: tuple, kwargs: dict[str, Any], *, capture: bool +) -> LaunchCall: + return LaunchCall( + jit_fn=jit_fn, + args=tuple(args), + kwargs=MappingProxyType( + {k: v for k, v in kwargs.items() if k not in ("grid", "warmup")} + ), + grid=kwargs.get("grid"), + capture=capture, + ) + + def _rebind_closure(fn: Any, old: Any, new: Any) -> Any: """Return ``fn`` with the closure cells that hold ``old`` pointing at ``new``.""" closure = getattr(fn, "__closure__", None) @@ -69,8 +84,24 @@ def add_client(self, new_client: str | Client) -> None: self.client_manager.add_clients([self._normalize_client(new_client)]) def finalize(self): + # Take the Launch first: once finalize ends the launch, another host + # thread may begin the next one on this manager. + launch = self.client_manager.launch self.client_manager.finalize() - launches.append(self.client_manager.launch) + launches.append(launch) + + @contextmanager + def _launch_scope(self, call: LaunchCall): + """begin_launch, then abort_launch if the launch raises. A refused or + failing begin_launch cleans up after itself and is not aborted: the + refusal must not reach the clients of another thread's launch.""" + mgr = self.client_manager + mgr.begin_launch(call) + try: + yield + except BaseException as exc: + mgr.abort_launch(exc) + raise class LaunchInterface: @@ -119,26 +150,85 @@ def _warmup_runner(self, runner: Any, jit_fn: Any | None) -> Any | None: return None return self._rebuild_runner(runner, jit_fn, interpreted=False) - def _rebuild_runner(self, runner: Any, leaf: Any, *, interpreted: bool) -> Any: + def _ir_runner(self, runner: Any, jit_fn: Any | None) -> Any | None: + if jit_fn is None: + return None + return self._rebuild_runner(runner, _IRLeaf(jit_fn), interpreted=False, ir=True) + + def _autotuned(self, runner: Any) -> bool: + """Whether an Autotuner layer sits anywhere in ``runner``'s chain.""" + while self._is_autotuner(runner) or self._is_heuristics(runner): + if self._is_autotuner(runner): + return True + runner = runner.fn + return False + + def _drop_autotuner_args(self, runner: Any) -> None: + """Clear the per-call tensors Autotuner layers in ``runner``'s chain + of copies may still hold: ``nargs`` and ``restore_copies``. + + Autotuner.run and .warmup keep the call's arguments in ``nargs`` + until they return, and a benchmark call keeps its restore_value + clones in ``restore_copies`` until its post_hook, which _bench skips + for a KeyboardInterrupt. So a launch that raises would leave the + caller's tensors, or device clones of them, on a copy that outlives + the launch. + """ + while self._is_autotuner(runner) or self._is_heuristics(runner): + if self._is_autotuner(runner): + runner.nargs = None + if hasattr(runner, "restore_copies"): + runner.restore_copies = {} + runner = runner.fn + + def _rebuild_runner( + self, runner: Any, leaf: Any, *, interpreted: bool, ir: bool = False + ) -> Any: """Rebuild the Autotuner/Heuristics chain of ``runner`` on top of ``leaf``. Every layer is shallow-copied down to the kernel, which ``leaf`` replaces, so the user's runner is never mutated. No deepcopy: a real JITFunction holds an RLock. A nested trace is looked through so the - layers it wraps are kept. + layers it wraps are kept. ``ir``: the chain whose warmup stands in + for the launch's autotuning (the IR clients' compiles). """ if isinstance(runner, (TritonTrace, GluonTrace)): runner = runner.fn if not (self._is_autotuner(runner) or self._is_heuristics(runner)): return leaf layer = copy.copy(runner) - layer.fn = self._rebuild_runner(runner.fn, leaf, interpreted=interpreted) + layer.fn = self._rebuild_runner(runner.fn, leaf, interpreted=interpreted, ir=ir) if self._is_autotuner(layer): self._isolate_autotuner(runner, layer, interpreted=interpreted) + if ir: + layer.prune_configs = self._refusing_conflicts(layer) elif not interpreted: layer.warmup = self._heuristics_warmup(layer) return layer + @staticmethod + def _refusing_conflicts(layer: Any) -> Callable: + # The launch's autotuning benchmarks each pruned config first, and + # Autotuner._bench refuses a call that passes one of the config's + # meta-parameters itself, with this ValueError (as Triton words it). + # Autotuner.warmup, through which the IR clients' compiles + # go instead, would pass the keyword twice (a TypeError naming the + # IR leaf), so the IR chain's copy checks the pruned configs first. + prune = layer.prune_configs + + def prune_configs(kwargs): + pruned = prune(kwargs) + for config in pruned: + conflicts = kwargs.keys() & config.kwargs.keys() + if conflicts: + raise ValueError( + f"Conflicting meta-parameters: {', '.join(conflicts)}." + " Make sure that you don't re-define auto-tuned symbols." + ) + return pruned + + return prune_configs + def _isolate_autotuner( self, original: Any, layer: Any, *, interpreted: bool ) -> None: @@ -232,7 +322,10 @@ def unpack_kernel( else: self.jit_fn, self.base_fn, self.interpreted_fn = unpack_kernel(runner) self.runner = self._interpreter_runner(runner, self.interpreted_fn) + # The real chain for the interpreted launches' warmup votes, and one + # for the IR clients' host compiles, which no warmup patch gates. self.warmup_runner = self._warmup_runner(runner, self.jit_fn) + self.ir_runner = self._ir_runner(runner, self.jit_fn) self.arg_names = runner.arg_names @@ -243,12 +336,98 @@ def unpack_kernel( self._copy_callable_attrs(runner, self.base_fn, src_fallback=self.jit_fn) def run(self, *args, **kwargs): + mgr = self.client_manager + has_ir = bool(mgr.ir_clients()) + capture = has_ir and self.jit_fn is not None + call = _launch_call(self.jit_fn, args, kwargs, capture=capture) + with self._launch_scope(call): + if not has_ir: + return self._run_interpreted(*args, **kwargs) + if mgr.interpreting_clients(): + if capture: + # Mixed trace: host-compile every config for the IR + # clients, then the interpreter produces the outputs (no + # real launch, no device). An IR-side compile failure is + # data for the IR clients and never stops the eager + # peers. A call that does not bind the kernel's + # parameters raises here, as in an IR-only trace, before + # the interpreter runs: the interpreted run would fail on + # the same call (with Python's own TypeError for the + # kernel function), and the error raised is the one the + # untraced JIT raises. + self._compile_for_ir(args, kwargs) + return self._run_interpreted(*args, **kwargs) + # Only IR clients, none of which runs the real kernel (see + # Client.LAUNCH). Without a JITFunction (TRITON_INTERPRET / an + # InterpretedFunction runner; call.capture is False) there is no + # compiled kernel for them either, so nothing runs. + ret = self._run_compiled(*args, **kwargs) if capture else None + self.finalize() + return ret + + def _compile_for_ir(self, args, kwargs): + """Compile every (pruned) config through the IR runner's warmup; the + capture host-compiles each call for the IR clients' target and turns + it into an IR event, or a compile_failed event, without launching or + touching a device. Returns the warmup result and the capture window: + for a plain or @heuristics kernel the host-compiled kernel (None if + it failed to compile). + + A config that fails to compile never fails the launch, not + even when no config compiled: the failure is the IR target's, which + the IR client chose, and the IR clients get it through + compile_failed. A call that does not bind the kernel's parameters + is no compile failure: it raises the JIT binder's error, as the + untraced call does (see ClientManager.ir_capture). + """ + runner = self.ir_runner + assert runner is not None # built whenever jit_fn is set + try: + with ( + self._real_compile_window(), + self.client_manager.ir_capture( + self.jit_fn, real_args=_untraced_call_args + ) as window, + ): + ret = runner.warmup(*args, **_without_warmup(kwargs)) + finally: + self._drop_autotuner_args(runner) + return ret, window + + def _run_compiled(self, *args, **kwargs): + """IR-only trace: no interpreter, no real launch. Every config is + compiled for the IR clients, so what they see never depends on the + autotune cache or on benchmark timing. + + Returns what it can of the untraced return value: the host-compiled + kernel when no Autotuner is involved (its only config; never loaded, + it cannot launch; None if it failed to compile), None for an + autotuned kernel (no config was picked). Nothing needs a GPU, and + the user's pre_run_hooks never fire. A config that failed to compile + for the IR target does not fail the launch: the IR clients get it + through compile_failed. Only a compile refused while an interpreted + traced launch has the language patched (LanguagePatchedError, no + compile's outcome: concurrent traced launches that mix + interpretation and host compiles are unsupported) fails the + launch, and a call that does not bind the kernel's parameters, + which raises the JIT binder's error as the untraced call does. + """ + ret, window = self._compile_for_ir(args, kwargs) + refused = [e for e in window.failures if isinstance(e, LanguagePatchedError)] + if refused: + raise refused[0] + return None if self._autotuned(self.ir_runner) else ret + + def _run_interpreted(self, *args, **kwargs): self._voted_warmup(*args, **kwargs) with self.client_manager.patch_run(self.base_fn, frontend_name="triton"): kwargs.update({"client_manager": self.client_manager}) kwargs.update({"jit_fn": self.jit_fn}) - ret = self.runner.run(*args, **kwargs) + try: + ret = self.runner.run(*args, **kwargs) + finally: + self._drop_autotuner_args(self.runner) self.finalize() return ret @@ -281,9 +460,9 @@ def __call__(self, *args, **kwargs): "tilelens does not unwrap; use its JITFunction " f"({self.__name__}.jit_fn) there instead." ) - # check that client sets match for calling and called functions - outer_clients = set(outer_client_manager.clients) - inner_clients = set(self.client_manager.clients) + # Only interpreting clients take part in the interpreted run. + outer_clients = {c.NAME for c in outer_client_manager.interpreting_clients()} + inner_clients = {c.NAME for c in self.client_manager.interpreting_clients()} if outer_clients != inner_clients: raise RuntimeError( "nested traced calls require matching clients; " @@ -304,7 +483,38 @@ def _voted_warmup(self, *args, **kwargs): compile_context=self._real_compile_window, real_args=_untraced_call_args, ): - return self.warmup_runner.warmup(*args, **_without_warmup(kwargs)) + try: + return self.warmup_runner.warmup(*args, **_without_warmup(kwargs)) + finally: + self._drop_autotuner_args(self.warmup_runner) + + +class _IRLeaf: + """The traced JITFunction at the bottom of the IR runner chain. + + Its warmup is the JITFunction class's, so no instance-level warmup patch + (patch_warmup's vote gate, or anyone else's) decides whether an IR + config compiles; everything else, ``run`` (where ir_capture sits, and + host-compiles instead of entering JITFunction.run) included, is the + JITFunction instance's. Its ``fn`` is the JITFunction too, as Triton's + Autotuner expects when it follows ``.fn`` from its own down to the + JITFunction it tunes. + """ + + def __init__(self, jit_fn: Any) -> None: + self.jit_fn = jit_fn + + @property + def fn(self) -> Any: + return self.jit_fn + + def warmup(self, *args, **kwargs): + return type(self.jit_fn).warmup(self.jit_fn, *args, **kwargs) + + def __getattr__(self, name: str) -> Any: + if name == "jit_fn": # not set yet (e.g. mid-copy): no recursion + raise AttributeError(name) + return getattr(self.jit_fn, name) def _untraced_call_args( @@ -521,23 +731,29 @@ def run(self, *args, pre_trace=True, platform_target="trn1", **kwargs): if you want full python flexibility inside kernels (e.g. importing modules inside a kernel). Does nothing if self.frontend_name == 'nki'. """ - if self.frontend_name == "nki_beta2" and pre_trace: - import nki - - kwargs.pop("warmup", None) - grid = kwargs.pop("grid", None) - nki.trace(self.func, grid=grid, platform_target=platform_target).specialize( - *args, **kwargs - ) - kwargs["grid"] = grid - with self.client_manager.patch_run( - self.func, - frontend_name=self.frontend_name, - ): - kwargs.update({"client_manager": self.client_manager}) - ret = self.interpreter_fn.run(*args, **kwargs) - self.finalize() - return ret + with self._launch_scope(_launch_call(None, args, kwargs, capture=False)): + if not self.client_manager.interpreting_clients(): + # Only IR clients: there is no compiled kernel here, and the + # real kernel never runs for them, so nothing runs. + self.finalize() + return None + if self.frontend_name == "nki_beta2" and pre_trace: + import nki + + kwargs.pop("warmup", None) + grid = kwargs.pop("grid", None) + nki.trace( + self.func, grid=grid, platform_target=platform_target + ).specialize(*args, **kwargs) + kwargs["grid"] = grid + with self.client_manager.patch_run( + self.func, + frontend_name=self.frontend_name, + ): + kwargs.update({"client_manager": self.client_manager}) + ret = self.interpreter_fn.run(*args, **kwargs) + self.finalize() + return ret class GluonTrace(LaunchInterface, TraceInterface, KernelTraceSupport): @@ -579,17 +795,23 @@ def run(self, *args, **kwargs): "GluonTrace.run() missing required keyword argument: 'grid'" ) - with self.client_manager.patch_run(self.base_fn, frontend_name="gluon"): - try: - ret = self.runner.run( - *args, - **kwargs, - client_manager=self.client_manager, - ) - finally: - self.client_manager.post_run_callback(self.base_fn) - self.finalize() - return ret + with self._launch_scope(_launch_call(None, args, kwargs, capture=False)): + if not self.client_manager.interpreting_clients(): + # Only IR clients: there is no compiled kernel here, and the + # real kernel never runs for them, so nothing runs. + self.finalize() + return None + with self.client_manager.patch_run(self.base_fn, frontend_name="gluon"): + try: + ret = self.runner.run( + *args, + **kwargs, + client_manager=self.client_manager, + ) + finally: + self.client_manager.post_run_callback(self.base_fn) + self.finalize() + return ret def __call__(self, *args, **kwargs): return self.fn(*args, **kwargs) diff --git a/tilelens/core/trace_io.py b/tilelens/core/trace_io.py index c95723c9c..aa233be65 100644 --- a/tilelens/core/trace_io.py +++ b/tilelens/core/trace_io.py @@ -10,6 +10,8 @@ from ..clients.profiler import data as profiler_data from ..clients.sanitizer import data as sanitizer_data +from ..ir import launch as ir_launch +from ..ir import verdict as ir_verdict from ..utils import traceback_utils from . import data as trace_data from .data import Launch, TensorSnapshot @@ -22,7 +24,16 @@ ArrayMap = dict[str, np.ndarray] _TRACE_CLASSES = { f"{cls.__module__}:{cls.__qualname__}": cls - for module in (trace_data, profiler_data, sanitizer_data, traceback_utils) + for module in ( + trace_data, + profiler_data, + sanitizer_data, + traceback_utils, + # Pure data, like verdict (neither imports Triton): TensorFacts is + # registered in its own right, not only as sanitizer_data's import. + ir_launch, + ir_verdict, + ) for cls in vars(module).values() if isinstance(cls, type) and is_dataclass(cls) } diff --git a/tilelens/ir/__init__.py b/tilelens/ir/__init__.py new file mode 100644 index 000000000..0bc6c7ead --- /dev/null +++ b/tilelens/ir/__init__.py @@ -0,0 +1,41 @@ +"""Compiled-IR layer: what IR-mode clients get of the kernels Triton compiles (TTIR). + +Exports resolve on first access, so importing ``tilelens.ir`` imports neither +Triton nor its MLIR bindings. +""" + +from __future__ import annotations + +from importlib import import_module +from typing import Any + + +_EXPORTS: dict[str, tuple[str, str]] = { + "IRClient": ("tilelens.ir.client", "IRClient"), + "ArtifactLog": ("tilelens.ir.capture", "ArtifactLog"), + "CompiledArtifacts": ("tilelens.ir.capture", "CompiledArtifacts"), + "CompiledSpecialization": ("tilelens.ir.capture", "CompiledSpecialization"), + "CompileFailure": ("tilelens.ir.capture", "CompileFailure"), + "ParseCache": ("tilelens.ir.capture", "ParseCache"), + "ParseOutcome": ("tilelens.ir.capture", "ParseOutcome"), + "LaunchBinding": ("tilelens.ir.launch", "LaunchBinding"), + "TensorFacts": ("tilelens.ir.launch", "TensorFacts"), + "bind_launch": ("tilelens.ir.launch", "bind_launch"), + "IRVerdict": ("tilelens.ir.verdict", "IRVerdict"), + "ConfigVerdict": ("tilelens.ir.verdict", "ConfigVerdict"), + "Refusal": ("tilelens.ir.verdict", "Refusal"), + "SourceLocation": ("tilelens.ir.verdict", "SourceLocation"), +} + +__all__ = list(_EXPORTS) + + +def __getattr__(name: str) -> Any: + try: + module_name, attr_name = _EXPORTS[name] + except KeyError as exc: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") from exc + + value = getattr(import_module(module_name), attr_name) + globals()[name] = value + return value diff --git a/tilelens/ir/capture.py b/tilelens/ir/capture.py new file mode 100644 index 000000000..38165eec3 --- /dev/null +++ b/tilelens/ir/capture.py @@ -0,0 +1,277 @@ +"""Per-launch compiled artifacts, and a content-addressed parse cache. + +``ArtifactLog`` records what the core's IR hooks delivered during one traced +launch: per compiled specialization its declared IR stages and compile +metadata plus the LaunchBindings seen for it, and every compile failure. +``ParseCache`` runs a reader once per distinct text and keeps what it gave, +a graph or a typed refusal, so a refusal's kind survives cache hits. + +Mechanism only: which specialization counts, what an absent stage means and +what a refusal or an error becomes are the client's calls. Neither class +raises for a bad kernel, text or reader. + +Importing this module does not import Triton or the TTIR reader. +""" + +from __future__ import annotations + +import builtins +import hashlib +import sys +from collections.abc import Callable, Hashable, Iterable, Mapping +from dataclasses import dataclass +from importlib import import_module +from types import MappingProxyType +from typing import TYPE_CHECKING, Any + +from .launch import LaunchBinding, bind_launch, config_kwargs + +if TYPE_CHECKING: + from ..core.client import LaunchCall, LaunchEvent + + +@dataclass(frozen=True) +class CompiledArtifacts: + """What one compiled specialization left in ``kernel.asm`` and + ``kernel.metadata``.""" + + # The declared stages the kernel holds: text, or bytes for a binary + # stage. A declared stage the kernel lacks (e.g. under + # TRITON_STORE_BINARY_ONLY) is absent. + stages: Mapping[str, Any] + # "backend", "arch", "num_warps", "num_stages", "shared", "name" from + # the compile metadata (None where it has none, e.g. "shared" for a + # kernel compiled only through TTIR), and "config": the config kwargs of + # the call that first produced the specialization. + meta: Mapping[str, Any] + # What could not be read, "; "-joined; stages and meta hold the rest. + error: str | None = None + + +@dataclass(frozen=True) +class CompiledSpecialization: + """One specialization a traced launch compiled, and every binding it was + delivered with (one per before_launch event).""" + + specialization: Hashable + artifacts: CompiledArtifacts + bindings: tuple[LaunchBinding, ...] + + @property + def config(self) -> Mapping[str, Any]: + return self.artifacts.meta["config"] + + +@dataclass(frozen=True) +class CompileFailure: + """A call of the launch that failed to compile (compile_failed).""" + + # The exception the host compile raised, whole (a CompilationError + # with its source excerpt and the errors it was raised from). + error: BaseException | None + config: Mapping[str, Any] + # The GPUTarget the compile was for (LaunchEvent.target). + target: Any = None + # The JITFunction that failed to compile (LaunchEvent.jit_fn). + jit_fn: Any = None + + +_METADATA_FIELDS = ("num_warps", "num_stages", "shared", "name") + + +def _describe(exc: BaseException) -> str: + return f"{type(exc).__name__}: {exc}" + + +def _read_artifacts( + kernel: Any, stages: frozenset[str], config: Mapping[str, Any] +) -> CompiledArtifacts: + texts: dict[str, Any] = {} + meta: dict[str, Any] = dict.fromkeys(("backend", "arch", *_METADATA_FIELDS)) + meta["config"] = config + errors: list[str] = [] + try: + asm = kernel.asm + for stage in sorted(stages): + try: + texts[stage] = asm[stage] + except KeyError: + pass + except Exception as exc: # e.g. "sass" needs cuobjdump + errors.append(f"asm[{stage!r}]: {_describe(exc)}") + except Exception as exc: + errors.append(f"asm: {_describe(exc)}") + try: + metadata = kernel.metadata + target = getattr(metadata, "target", None) + meta["backend"] = getattr(target, "backend", None) + meta["arch"] = getattr(target, "arch", None) + for name in _METADATA_FIELDS: + meta[name] = getattr(metadata, name, None) + except Exception as exc: + errors.append(f"metadata: {_describe(exc)}") + return CompiledArtifacts( + stages=MappingProxyType(texts), + meta=MappingProxyType(meta), + error="; ".join(errors) if errors else None, + ) + + +class ArtifactLog: + """What one traced launch compiled, for an IR client that reads + ``stages`` of each kernel. + + ``reset(call)`` starts a launch; ``record`` takes each before_launch + event and ``record_failure`` each compile_failed event. Specializations + and failures keep the order they were first seen in (the autotuner's + config order). + """ + + def __init__(self, stages: Iterable[str]) -> None: + self.stages = frozenset(stages) + self.reset() + + def reset(self, call: LaunchCall | None = None) -> None: + """Forget everything recorded; ``call`` is the launch about to start + (it tells config kwargs from the caller's own, see config_kwargs).""" + self.call = call + self._compiled: dict[ + Hashable, tuple[CompiledArtifacts, list[LaunchBinding]] + ] = {} + self._failures: list[CompileFailure] = [] + + def record(self, event: LaunchEvent) -> None: + binding = bind_launch(event, self.call) + entry = self._compiled.get(event.specialization) + if entry is None: + artifacts = _read_artifacts(event.kernel, self.stages, binding.config) + self._compiled[event.specialization] = (artifacts, [binding]) + else: + entry[1].append(binding) + + def record_failure(self, event: LaunchEvent) -> None: + self._failures.append( + CompileFailure( + error=event.error, + config=MappingProxyType(config_kwargs(event, self.call)), + target=getattr(event, "target", None), + jit_fn=getattr(event, "jit_fn", None), + ) + ) + + @property + def specializations(self) -> tuple[CompiledSpecialization, ...]: + return tuple( + CompiledSpecialization(specialization, artifacts, tuple(bindings)) + for specialization, (artifacts, bindings) in self._compiled.items() + ) + + @property + def failures(self) -> tuple[CompileFailure, ...]: + return tuple(self._failures) + + +@dataclass(frozen=True) +class ParseOutcome: + """A reader's result for one text: exactly one of the three is set, + unless the reader itself returned None.""" + + graph: Any = None + # The reader's refusal exception (the TTIR reader's UnsupportedTTIR), + # tracebacks dropped. + refusal: BaseException | None = None + # Any other exception the reader raised, as "Type: message". + error: str | None = None + + +def content_key(text: str) -> str: + """Stable SHA-256 of an IR text (a lone surrogate hashes as "?").""" + return hashlib.sha256(text.encode("utf-8", errors="replace")).hexdigest() + + +def _default_reader() -> Callable[..., Any]: + # Resolved on every lookup, so a monkeypatched reader is picked up (and + # keyed apart by its identity). + return import_module(".ttir_reader", __package__).parse_ttir + + +def _default_refusal() -> type[BaseException] | None: + try: + return import_module(".ttir_reader", __package__).UnsupportedTTIR + except Exception: # no reader module: nothing can be its refusal + return None + + +_EXCEPTION_GROUP = getattr(builtins, "BaseExceptionGroup", None) # Python >= 3.11 + + +def _without_frames(exc: BaseException, outer: BaseException | None) -> BaseException: + # A cached refusal outlives its parse; its traceback, and those of every + # exception chained to it, would keep the reader's frames alive. The + # exception the caller was handling when it asked (``outer``, the chain's + # implicit context) is the caller's: unlinked, never cleared. + seen: set[int] = set() + stack = [exc] + while stack: + link = stack.pop() + if id(link) in seen: + continue + seen.add(id(link)) + link.__traceback__ = None + if outer is not None and link.__context__ is outer: + link.__context__ = None + if outer is not None and link.__cause__ is outer: + link.__cause__ = None + stack.extend(x for x in (link.__cause__, link.__context__) if x is not None) + if _EXCEPTION_GROUP is not None and isinstance(link, _EXCEPTION_GROUP): + stack.extend(getattr(link, "exceptions", ())) + return exc + + +class ParseCache: + """Parse each distinct IR text once per reader and options. + + ``reader(text, **options)`` returns a graph or raises ``refusal`` (by + default the TTIR reader's UnsupportedTTIR) to decline the text; any other + exception is reported as an error and not cached, so a later lookup + retries it. The default reader, ``tilelens.ir.ttir_reader.parse_ttir``, + is imported at each lookup, not with this module; while that module + cannot be imported, a lookup is an error outcome like any other reader + error. Never raises (``Exception``s only; an interrupt still + propagates); ``text`` is positional-only, so any option name reaches + the reader. + """ + + def __init__( + self, + reader: Callable[..., Any] | None = None, + *, + refusal: type[BaseException] | None = None, + ) -> None: + self._reader = reader + self._refusal = refusal + self._outcomes: dict[Hashable, ParseOutcome] = {} + + def get(self, text: str, /, **options: Hashable) -> ParseOutcome: + outer = sys.exc_info()[1] + try: + reader = self._reader if self._reader is not None else _default_reader() + key = ( + content_key(text), + reader, + tuple(sorted(options.items())), + ) + cached = self._outcomes.get(key) + except Exception as exc: + return ParseOutcome(error=_describe(exc)) + if cached is not None: + return cached + try: + outcome = ParseOutcome(graph=reader(text, **options)) + except Exception as exc: + refusal = self._refusal if self._refusal is not None else _default_refusal() + if refusal is None or not isinstance(exc, refusal): + return ParseOutcome(error=_describe(exc)) + outcome = ParseOutcome(refusal=_without_frames(exc, outer)) + self._outcomes[key] = outcome + return outcome diff --git a/tilelens/ir/client.py b/tilelens/ir/client.py new file mode 100644 index 000000000..461bda818 --- /dev/null +++ b/tilelens/ir/client.py @@ -0,0 +1,124 @@ +"""The IR client base: lifecycle only, no analysis defaults. + +An ``IRClient`` takes no part in the interpreted run: its interpreter-path +methods are inert and it declares ``NEEDS_INTERPRETER = False``, so the core +hands it compiled kernels through ``before_launch`` / ``compile_failed``, +which fill its per-launch ``ArtifactLog``. ``finalize`` is a template: +``analyze_launch(log)`` returns the reports and the verdict, also for a +launch nothing could be captured for (see analyze_launch); an ``Exception`` +from it goes to ``on_analysis_error``, which returns the verdict instead (an +interrupt or ``SystemExit`` propagates). + +It returns the reports followed by the verdict, which ``ClientManager`` +puts into ``Launch.records``, and keeps the verdict as ``last_verdict`` +(None until a launch finalizes). A subclass declares ``NAME``, +``IR_STAGES`` and ``LAUNCH``, may set ``ir_target`` (the target its kernels +are compiled for on the host; the configured default otherwise) and +implements the two hooks; statuses, refusal meanings, caches and report +printing are all its own. +""" + +from __future__ import annotations + +from abc import abstractmethod +from collections.abc import Callable +from typing import Any, ClassVar + +from ..core.callbacks import ForLoopCallbacks, OpCallbacks +from ..core.client import Client, LaunchCall, LaunchEvent +from ..core.data import Op +from .capture import ArtifactLog +from .verdict import IRVerdict + + +class IRClient(Client): + NEEDS_INTERPRETER: ClassVar[bool] = False + + def __init__(self) -> None: + super().__init__() + self.artifacts = ArtifactLog(self.IR_STAGES) + # The last finalized launch's verdict (it is also in Launch.records). + self.last_verdict: IRVerdict | None = None + + # ── the client's analysis ──────────────────────────────────────── + + @abstractmethod + def analyze_launch(self, log: ArtifactLog) -> tuple[list, IRVerdict]: + """Analyze one traced launch: the reports and the verdict. + + ``log.call`` is the launch's LaunchCall. When ``log.call.capture`` is + False, nothing was compiled or recorded: the trace has no JITFunction + (``log.call.jit_fn`` is None: TRITON_INTERPRET, an InterpretedFunction + runner, Gluon, NKI). An empty log then says nothing about the + kernel; what such a launch gets is the subclass's call. + """ + + @abstractmethod + def on_analysis_error(self, exc: Exception) -> IRVerdict: + """The verdict for a launch whose analysis raised ``exc``.""" + + # ── launch lifecycle ───────────────────────────────────────────── + # A subclass overriding one of these calls super(). + + def begin_launch(self, call: LaunchCall) -> None: + self.artifacts.reset(call) + self.last_verdict = None + + def abort_launch(self, exc: BaseException) -> None: + self.artifacts.reset() + + def before_launch(self, event: LaunchEvent) -> None: + self.artifacts.record(event) + + def compile_failed(self, event: LaunchEvent) -> None: + self.artifacts.record_failure(event) + + def finalize(self) -> list: + try: + reports, verdict = self._verdict() + finally: + # The log holds compile exceptions (and their frames); the next + # launch starts from a fresh one anyway. + self.artifacts.reset() + self.last_verdict = verdict + return [*reports, verdict] + + def _verdict(self) -> tuple[list, IRVerdict]: + try: + reports, verdict = self.analyze_launch(self.artifacts) + return list(reports), verdict + except Exception as exc: + return [], self.on_analysis_error(exc) + + # ── inert interpreter path: the core calls none of these for an IR + # client except the warmup vote, which declines (IR compiles go through + # ir_capture) ── + + def pre_run_callback(self, fn: Callable) -> bool: + return False + + def post_run_callback(self, fn: Callable) -> bool: + return False + + def arg_callback(self, name: str, arg: Any, arg_cvt: Any) -> None: + pass + + def grid_callback(self, grid: tuple[int, ...]) -> None: + pass + + def grid_idx_callback(self, grid_idx: tuple[int, ...]) -> None: + pass + + def register_op_callback( + self, op_type: type[Op], *args: Any, **kwargs: Any + ) -> OpCallbacks: + return OpCallbacks() + + def register_for_loop_callback(self) -> ForLoopCallbacks: + return ForLoopCallbacks() + + def pre_warmup_callback(self, jit_fn: Callable, *args: Any, **kwargs: Any) -> bool: + return False + + def post_warmup_callback(self, jit_fn: Callable, ret: Any) -> None: + pass diff --git a/tilelens/ir/launch.py b/tilelens/ir/launch.py new file mode 100644 index 000000000..ea21100e1 --- /dev/null +++ b/tilelens/ir/launch.py @@ -0,0 +1,227 @@ +"""What one traced launch bound its kernel parameters to. + +A ``LaunchBinding`` is built from a core ``LaunchEvent``: its ``bound_args`` +split into integer scalars, tensor facts and constexprs, plus the grid and the +config kwargs an Autotuner/Heuristics layer added. Mechanism only: which of +these facts an analysis trusts (the view footprint or the allocation, whether +non-contiguous tensors are refused, what a missing fact means) is the client's +call. Building a binding never raises; a fact that cannot be read leaves the +argument out and names it in ``LaunchBinding.error``. + +A binding is not a complete account of the kernel's arguments: see +``LaunchBinding`` for what it leaves out. + +Importing this module does not import Triton. +""" + +from __future__ import annotations + +import operator +from collections.abc import Mapping +from dataclasses import dataclass +from types import MappingProxyType +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from ..core.client import LaunchCall, LaunchEvent + + +@dataclass(frozen=True) +class TensorFacts: + """Launch-time facts about one tensor argument, read without touching + its values.""" + + # The view's first element; already includes the storage offset, so a + # lowering must never add that offset again. + data_ptr: int + elem_size: int # bytes + numel: int + shape: tuple[int, ...] + strides: tuple[int, ...] # in elements + dtype: str # str(tensor.dtype), e.g. "torch.float32" + contiguous: bool + # The underlying allocation, independent of the view's data_ptr, shape + # and strides; None when the tensor exposes no storage. + storage_data_ptr: int | None = None + storage_nbytes: int | None = None + + def allocation_interval(self) -> tuple[int, int] | None: + """Verified byte bounds [start, end) of the allocation, or None when + the address extent is unknown. + + Without storage metadata only a contiguous view's own extent is + known. Partial or inconsistent storage metadata never falls back to + numel, which could silently deactivate valid accesses. + """ + if self.elem_size <= 0 or self.numel < 0 or self.data_ptr < 0: + return None + if self.storage_data_ptr is None and self.storage_nbytes is None: + if not self.contiguous: + return None + return self.data_ptr, self.data_ptr + self.numel * self.elem_size + if self.storage_data_ptr is None or self.storage_nbytes is None: + return None + start, size = self.storage_data_ptr, self.storage_nbytes + end = start + size + if start < 0 or size < 0 or not start <= self.data_ptr <= end: + return None + if self.numel and self.data_ptr + self.elem_size > end: + return None + if self.contiguous and self.data_ptr + self.numel * self.elem_size > end: + return None + return start, end + + +@dataclass(frozen=True) +class LaunchBinding: + """One call's kernel parameters, by name, as a launch bound them. + + Only int/bool scalars, tensors and constexprs are bound. Arguments of + other kinds (floats, None, tuples, ...) are left out without an error, + although a tuple argument is several TTIR function arguments (e.g. two + pointers). A descriptor-style argument is bound as its ``.base`` tensor + alone: the shape, stride and flag fields it adds to the TTIR function + are not bound. So a consumer must treat a TTIR function argument with no + entry in ``params`` or ``tensors`` as unknown (e.g. refuse an access that + depends on it), never as unconstrained. + """ + + # Non-constexpr int and bool arguments (bools as 0/1). + params: Mapping[str, int] + # Tensor arguments; a descriptor-style argument is recorded as its + # ``.base`` tensor (see above). + tensors: Mapping[str, TensorFacts] + # Arguments to tl.constexpr parameters, as passed. + constexprs: Mapping[str, Any] + # The grid as passed: a tuple, a callable, or None. + raw_grid: Any + # The grid canonicalized to three int dims; None if it cannot be resolved, + # or if a dim is no integer (named in ``error``). + grid: tuple[int, int, int] | None + # The keyword arguments Autotuner/Heuristics layers added to the + # caller's call (see config_kwargs). + config: Mapping[str, Any] + # The facts that could not be read, "; "-joined, e.g. "argument 'x': + # AttributeError: ...". Arguments of kinds a binding does not record are + # no error (see above). + error: str | None = None + + +def tensor_facts(value: Any) -> TensorFacts: + """Read the TensorFacts of a torch-like tensor. Raises if a fact is + unreadable; bind_launch contains that.""" + storage_data_ptr = storage_nbytes = None + untyped_storage = getattr(value, "untyped_storage", None) + if callable(untyped_storage): + try: + storage = untyped_storage() + storage_data_ptr = int(storage.data_ptr()) + storage_nbytes = int(storage.nbytes()) + except Exception: # duck-typed tensors without a storage + storage_data_ptr = storage_nbytes = None + return TensorFacts( + data_ptr=int(value.data_ptr()), + elem_size=int(value.element_size()), + numel=int(value.numel()), + shape=tuple(int(size) for size in value.shape), + strides=tuple(int(stride) for stride in value.stride()), + dtype=str(value.dtype), + contiguous=bool(value.is_contiguous()), + storage_data_ptr=storage_data_ptr, + storage_nbytes=storage_nbytes, + ) + + +_SCALARS = (bool, int, float, str) + + +def _is_passed(passed: Any, value: Any) -> bool: + # A layer that recomputes a caller's scalar to an equal value may hand on + # another object (e.g. an int above the small-int cache); anything else + # (tensors, callables) is the caller's only as the same object. + if passed is value: + return True + return type(passed) is type(value) and type(value) in _SCALARS and passed == value + + +def config_kwargs(event: LaunchEvent, call: LaunchCall | None) -> dict[str, Any]: + """The kwargs of ``event`` its launch's caller did not pass: what the + Autotuner/Heuristics layers added (config kwargs, num_warps, ... and + heuristic values that differ from the caller's). Without ``call`` every + kwarg counts.""" + if call is None: + return dict(event.kwargs) + passed = call.kwargs + return { + name: value + for name, value in event.kwargs.items() + if name not in passed or not _is_passed(passed[name], value) + } + + +def _constexpr_names(jit_fn: Any) -> frozenset[str]: + return frozenset( + param.name + for param in getattr(jit_fn, "params", None) or () + if getattr(param, "is_constexpr", False) + ) + + +def _described_tensor(value: Any) -> Any: + # A descriptor-style argument (e.g. triton.tools.tensor_descriptor. + # TensorDescriptor) addresses its .base tensor. + base = getattr(value, "base", None) + if base is not None and hasattr(base, "data_ptr"): + return base + return value + + +def _int_grid(resolved: Any) -> tuple[int, int, int] | None: + if resolved is None: + return None + # operator.index, as the launcher converts: a float dim is an error, not + # truncated into a grid the untraced launch would reject. + x, y, z = (operator.index(dim) for dim in resolved) + return x, y, z + + +def bind_launch(event: LaunchEvent, call: LaunchCall | None = None) -> LaunchBinding: + """Bind ``event`` (its ``bound_args`` and ``resolved_grid``); ``call`` is + the launch's LaunchCall, which tells config kwargs from the caller's. + Never raises.""" + params: dict[str, int] = {} + tensors: dict[str, TensorFacts] = {} + constexprs: dict[str, Any] = {} + errors: list[str] = [] + grid = None + config: dict[str, Any] = {} + try: + constexpr_names = _constexpr_names(event.jit_fn) + for name, value in event.bound_args.items(): + try: + if name in constexpr_names: + constexprs[name] = value + continue + value = _described_tensor(value) + if hasattr(value, "data_ptr"): + tensors[name] = tensor_facts(value) + elif isinstance(value, (bool, int)): + params[name] = int(value) + except Exception as exc: + errors.append(f"argument {name!r}: {type(exc).__name__}: {exc}") + try: + grid = _int_grid(event.resolved_grid) + except Exception as exc: + errors.append(f"grid {event.resolved_grid!r}: {type(exc).__name__}: {exc}") + config = config_kwargs(event, call) + except Exception as exc: + errors.append(f"{type(exc).__name__}: {exc}") + return LaunchBinding( + params=MappingProxyType(params), + tensors=MappingProxyType(tensors), + constexprs=MappingProxyType(constexprs), + raw_grid=getattr(event, "grid", None), + grid=grid, + config=MappingProxyType(config), + error="; ".join(errors) if errors else None, + ) diff --git a/tilelens/ir/verdict.py b/tilelens/ir/verdict.py new file mode 100644 index 000000000..3bb88c5e2 --- /dev/null +++ b/tilelens/ir/verdict.py @@ -0,0 +1,130 @@ +"""The structured record an IR client adds to ``Launch.records``. + +Plain frozen records: statuses and scopes are strings in the reporting +client's own vocabulary, and nothing here aggregates, ranks or interprets +them. Fields hold only strings, ints, tuples, dicts and these records, so +verdicts are picklable and ``tilelens.save()`` holds them, provided +the config values are ones a trace can hold too. +``SourceLocation`` and ``Refusal`` are hashable; ``ConfigVerdict`` and +``IRVerdict`` compare by value but are not hashable (a config dict has no +hash). + +Importing this module does not import Triton. +""" + +from __future__ import annotations + +from collections.abc import Hashable, Mapping +from dataclasses import dataclass +from typing import Any + + +def _plain_str(value: Any, what: str) -> str: + # A str subclass (e.g. the reader's TTIRKind enum) is saved as, and + # loads back as, its plain string; holding that string keeps a record's + # type, equality, hash and repr the same on both sides of a save. + if not isinstance(value, str): + raise TypeError(f"{what} must be a str, not {type(value).__name__}") + return str.__str__(value) + + +@dataclass(frozen=True) +class SourceLocation: + """A user-source location: file, 1-based line, and column if known.""" + + file: str + line: int + col: int | None = None + + +@dataclass(frozen=True) +class Refusal: + """Why (part of) a launch was not analyzed: a reader's UnsupportedTTIR, + or a refusal a client defines itself.""" + + kind: str + message: str + # The refused op's line in the IR text, and its user-source location. + # ``loc`` also takes any object with ``file`` and ``line`` (and + # optionally ``col``) attributes, such as the TTIR reader's source loc, + # and holds it as a SourceLocation. + line_no: int | None = None + loc: SourceLocation | None = None + + def __post_init__(self) -> None: + object.__setattr__(self, "kind", _plain_str(self.kind, "Refusal.kind")) + object.__setattr__(self, "message", _plain_str(self.message, "Refusal.message")) + loc = self.loc + if loc is None or isinstance(loc, SourceLocation): + return + if not (hasattr(loc, "file") and hasattr(loc, "line")): + raise TypeError( + "Refusal.loc must be a SourceLocation, None or an object with " + f"file and line attributes, not {type(loc).__name__}" + ) + object.__setattr__( + self, "loc", SourceLocation(loc.file, loc.line, getattr(loc, "col", None)) + ) + + @classmethod + def from_exception(cls, exc: BaseException) -> Refusal: + """Copy a refusal exception's structured fields (``kind``, and + ``message`` / ``line_no`` / ``loc`` where it has them); the reader's + source loc becomes a SourceLocation.""" + message = getattr(exc, "message", None) + return cls( + kind=getattr(exc, "kind"), + message=str(exc) if message is None else message, + line_no=getattr(exc, "line_no", None), + loc=getattr(exc, "loc", None), + ) + + +@dataclass(frozen=True) +class ConfigVerdict: + """One analyzed (or failed) config of a launch.""" + + # The compiled specialization (kernel hash); None for a config that + # produced no kernel. + specialization: Hashable + # The config kwargs the Autotuner/Heuristics layers added. + config: Mapping[str, Any] + status: str + refusal: Refusal | None = None + n_reports: int = 0 + + __hash__ = None # type: ignore[assignment] + + def __post_init__(self) -> None: + object.__setattr__(self, "config", dict(self.config)) + + +@dataclass(frozen=True) +class IRVerdict: + """One IR client's result for one traced launch.""" + + client: str # the client's NAME + status: str + # What the status holds for (e.g. a proof's quantifier scope). + scope: str | None = None + refusal: Refusal | None = None + per_config: tuple[ConfigVerdict, ...] = () + notes: tuple[str, ...] = () + + __hash__ = None # type: ignore[assignment] + + def __post_init__(self) -> None: + for name in ("per_config", "notes"): + if isinstance(getattr(self, name), str): + # tuple() would split it into characters + raise TypeError(f"IRVerdict.{name} takes a sequence, not a str") + per_config = tuple(self.per_config) + for item in per_config: + if not isinstance(item, ConfigVerdict): + raise TypeError( + "IRVerdict.per_config items must be ConfigVerdicts, " + f"not {type(item).__name__}" + ) + notes = tuple(_plain_str(note, "an IRVerdict note") for note in self.notes) + object.__setattr__(self, "per_config", per_config) + object.__setattr__(self, "notes", notes)