diff --git a/.github/workflows/benchmark.yml b/.github/workflows/benchmark.yml index 45d6b131..4efd0d5c 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 1919ccd1..96eba768 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 53387a16..7ec6b7dd 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 00000000..ac3e83a6 --- /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 00000000..e69de29b diff --git a/tests/unit/ir/test_host_compile.py b/tests/unit/ir/test_host_compile.py new file mode 100644 index 00000000..e5ad6f09 --- /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 00000000..c3ac0992 --- /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 00000000..682678fa --- /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 f58fcf1b..0f0bd478 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 9163d60b..de7861f8 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 00000000..4826e65d --- /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 ca27e967..aa4cde17 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 c95723c9..aa233be6 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 00000000..0bc6c7ea --- /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 00000000..38165eec --- /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 00000000..461bda81 --- /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 00000000..ea21100e --- /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 00000000..3bb88c5e --- /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)