diff --git a/.codespellrc b/.codespellrc new file mode 100644 index 000000000..ed130d4ed --- /dev/null +++ b/.codespellrc @@ -0,0 +1,5 @@ +[codespell] +# Fixed vocabulary, not typos: `olt` is an MLIR arith.cmpf predicate (ordered +# less-than), `aranges` is the plural of tl.arange, `lits` names the walk +# layer's string-literal spans (tilelens/ir/_mlir_walk.py). +ignore-words-list = aranges,lits,olt diff --git a/README.md b/README.md index 0efc0527e..bd52af756 100644 --- a/README.md +++ b/README.md @@ -124,6 +124,12 @@ uv sync --extra test # tests but no NKI support * To run core TileLens tests, run `pytest tests/`. * (if NKI installed) To run NKI-specific tests, run `pytest tests/ -m nki`. * To run all tests (Triton + NKI), run `pytest tests/ -m ""`. +* On a Triton release outside the tested window of the compiled sanitizer + (`tilelens.core.config.TESTED_TRITON_VERSIONS`), its IR-mode tests (marked + `ir_mode`, see `tests/conftest.py`) are skipped, saying why; the version-gate + tests (`tests/unit/test_ir_version_gate.py`) still run. Set + `TILELENS_IR_ALLOW_UNTESTED_TRITON=1` to run them all, e.g. to validate a new + release before adding it to the window. * To run visualizer web UI tests, run `npm run test:frontend`. ## Working with Examples @@ -191,6 +197,101 @@ Analyze kernels across visualization, profiling, and sanitization with a single - Profiler: flags non-unrolled loops, inefficient mask usage, and missing buffer_load optimizations while tracking load/store byte counts with low-overhead sampling. - Sanitizer: symbolically checks tensor memory accesses for out-of-bounds errors and emits reports with tensor metadata, call stack, and expression trees; optional fake-memory storage avoids real reads. +### Compiled sanitizer + +`Sanitizer(compile=True)` checks each launch against the kernel Triton compiles +for it (its TTIR) instead of interpreting it: out-of-bounds accesses (against the +tensor's view, strides included), integer-width overflows in address, mask and +branch arithmetic, and divisions by zero, each with a witness (program ids, lanes, +loop iteration). The kernel is compiled on the host, exactly as the JIT would +compile the launch but only through the TTIR stage (through the whole pipeline +under `TRITON_KERNEL_DUMP`, `TRITON_KERNEL_OVERRIDE` or `USE_IR_LOC`), so **no GPU +is needed**: CPU tensors work, and so does a machine without a driver. From the +CLI, give the flag before the script name (the legacy `triton-sanitizer` alias +takes it too): + +```sh +tile-sanitizer --compile my_script.py --my-script-flag +``` + +```py +from tilelens.clients import Sanitizer + + +@tilelens.trace(Sanitizer(compile=True)) +@triton.jit +def kernel(x_ptr, n, BLOCK: tl.constexpr): + ... +``` + +- Kernels are compiled and checked, **not run**: their outputs are never written, + so a script that checks its own results fails under `--compile`, and a launch + whose arguments the script computes from an earlier kernel's output is checked + with those unwritten values. +- Kernels are compiled for a fixed target, `cuda:89` (sm89, e.g. RTX 4090) by + default, so a verdict does not depend on the machine it was computed on. + Choose another with `Sanitizer(compile=True, target="cuda:90")` (or + `"hip:gfx942"`, or a Triton `GPUTarget`), or for every client that names none + with the environment variable `TILELENS_IR_TARGET=cuda:90` (e.g. for + `tile-sanitizer --compile`). The kernel's + own target queries (`tl.target_info.is_cuda()`, `cuda_capability_geq()`, + `is_hip()`) answer for that target too, and `TRITON_OVERRIDE_ARCH` does not + apply: the target is the one named. The TTIR can differ between targets (e.g. + tensor descriptors, target-dependent branches), and a verdict holds for the + target it was checked for. To check a launch for several targets, stack one + sanitizer per target (`@tilelens.trace(Sanitizer(compile=True, + target="cuda:90"))` over `@tilelens.trace(Sanitizer(compile=True, + target="cuda:80"))`): each keeps its own verdict (`last_verdict`), and the + launch's records hold them in trace order. +- Targets also differ in what compiles at all: `fp8e4nv` (`torch.float8_e4m3fn`, + `tl.float8e4nv`) needs `cuda:89` or later, `num_ctas > 1` and 16-bit tensor + descriptor atomic min/max need `cuda:90`. A kernel or autotune config that fails + to compile for the target never stops the script (it was not going to run + anyway): it is reported `unsupported`, kind `compile-failed`, naming the target + and how to choose another, since it may run on a GPU of another kind unchecked. + Name a target it compiles for to check it. Only a failure no target compiles + past (a failing `tl.static_assert`, or a Python construct Triton never compiles, + unless the kernel asked Triton's driver anything first, e.g. through + `tl.target_info` or a device query it catches, or an earlier compile of the + kernel for the target did) is just a note in the verdict: that config never + launches (the autotuner skips it too). A target answer the kernel's own code + keeps from outside the check (another trace or target, the untraced program) + and never asks for again cannot be seen. A launch none of whose configs + compiled is `unsupported` (`compile-failed`), and its notes are printed with + it. Each report names where the kernel failed (`file:line`) and the innermost + error in one line. +- A call that does not match the kernel's signature (a missing, extra or + misnamed argument), or that Triton cannot key (e.g. an unhashable constexpr + value), is a bug in the call, not a compile failure: it raises the very + `TypeError` the untraced call raises, on any GPU, and the script stops there + (under `tile-sanitizer --compile` with a traceback and exit status 1). So does + a trace that also interprets the kernel (e.g. with the `Tracer`), before the + interpreter runs, which would fail on the call too. An autotuned call that + passes an autotuned meta-parameter itself raises the autotuner's own + `Conflicting meta-parameters` `ValueError`. A keyword that names no parameter + is a compile option, which the target may not know (e.g. `waves_per_eu`, a + HIP option, under a CUDA target): that is the target's `compile-failed`, and + the report names the keyword (misspelled, it fails on every GPU untraced too). + All of this holds on a tested Triton release (below), or with the override: + on another release nothing is compiled, so nothing binds the call, and the + launch is `unsupported` (`untested-triton-version`), a call that does not + bind included (a trace that also interprets the kernel raises the + interpreter's own error for it). +- Each launch gets an `IRVerdict` in its records: `ok` is a proof for that + launch's scalar arguments, grid and tensors (`scope="launch"`); `violations` + comes with the findings; `unsupported` names what was not checked (e.g. a + data-dependent address, a construct the TTIR reader does not model, a Z3 + query that timed out, which can depend on the machine's load, a config that + failed to compile for the target, `compile-failed`, or a compile the host + cannot run, `host-compile-unavailable`). An autotuned + launch checks every config, each with its own arguments and grid and with a + `ConfigVerdict` of its own, also when several configs compile to one kernel. +- With the default `abort_on_error=True` the findings are printed and the + process exits with status 1; `ENABLE_SANITIZER=0` leaves kernels untraced. +- Tested on Triton 3.6 and 3.8; on another release every launch is + `unsupported` (`untested-triton-version`), nothing compiled, unless + `TILELENS_IR_ALLOW_UNTESTED_TRITON=1` is set. + ### Save and load traces ```py diff --git a/tests/conformance/__init__.py b/tests/conformance/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/conformance/_corpus.py b/tests/conformance/_corpus.py new file mode 100644 index 000000000..6bdfb6156 --- /dev/null +++ b/tests/conformance/_corpus.py @@ -0,0 +1,1753 @@ +"""Kernel + launch corpus of the TTIR reader conformance suite (D10b). + +Each :class:`Case` is one kernel launch: ``build(device)`` returns +``(grid, args, kwargs)`` and ``kernel`` is the ``@triton.jit`` function to +launch (an autotuned kernel contributes one case per config, the config's +kwargs in ``kwargs``). The same module is imported by the TTIR capture +child (real compile), by the interpreter child (``TRITON_INTERPRET=1``, so +the decorators build InterpretedFunctions there) and by the test module +(case names only), so it must import neither tilelens nor anything that +depends on how ``@triton.jit`` resolved. + +Every tensor argument owns its storage, zero-padded by ``PAD`` elements on +each side: the interpreter child attributes an address to the argument +whose storage holds it, and a moderately out-of-bounds lane (an unmasked +ragged tile, ...) still lands inside its own argument's padding. + +Rules for kernels here: every memory-op call on ONE source line, at most +one per line (sites are keyed by line; for a call wrapped over several +lines Triton 3.6 locates the op at the line of the argument its code +generator visited last while the interpreter's frame reports the call's, +so the interpreter child refuses such a call; 3.8 locates it at the +call's first line), and launches small enough for the interpreter. The audit +regression kernels are the goldens' own +(``tests/golden/ir/reader_kernels.py``), imported by path. + +Not a test module: pytest imports it (python_files = *.py) and finds nothing. +""" + +from __future__ import annotations + +import importlib.util +import math +import os +import sys +from dataclasses import dataclass +from typing import Any, Callable + +import torch +import triton +import triton.language as tl + +HERE = os.path.dirname(os.path.abspath(__file__)) +READER_KERNELS = os.path.join(HERE, "..", "golden", "ir", "reader_kernels.py") +PAD = 1 << 12 # elements of zero padding on each side of every tensor + + +def _load_reader_kernels(): + name = "tilelens_conformance_reader_kernels" + module = sys.modules.get(name) + if module is None: + spec = importlib.util.spec_from_file_location( + name, os.path.abspath(READER_KERNELS) + ) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module # Triton reads the jit fn's module + spec.loader.exec_module(module) + return module + + +RK = _load_reader_kernels() + + +def buf( + numel: int, dtype=torch.float32, dev: str = "cpu", init: Any = "arange" +) -> torch.Tensor: + """A 1-D tensor of ``numel`` elements inside its own zero-padded storage.""" + base = torch.zeros(numel + 2 * PAD, dtype=dtype, device=dev) + t = base[PAD : PAD + numel] + if init == "arange" and numel: + t.copy_(torch.arange(numel, device=dev).to(dtype)) + elif init is not None and init != "arange": + t.copy_(torch.as_tensor(init, dtype=dtype, device=dev)) + return t + + +def T( + *shape: int, dtype=torch.float32, dev: str = "cpu", init: Any = "arange" +) -> torch.Tensor: + return buf(math.prod(shape), dtype, dev, init).view(*shape) + + +@dataclass(frozen=True) +class Case: + name: str + kernel: Any # the JITFunction (InterpretedFunction under TRITON_INTERPRET=1) + build: Callable[[str], tuple] # device -> (grid, args, kwargs) + group: str + note: str = "" + + +CASES: list[Case] = [] + + +def case(name: str, kernel: Any, group: str, note: str = ""): + def deco(build): + if any(c.name == name for c in CASES): + raise ValueError(f"duplicate case {name!r}") + CASES.append(Case(name, kernel, build, group, note)) + return build + + return deco + + +def by_name() -> dict[str, Case]: + return {c.name: c for c in CASES} + + +# ════════════════════════════ a: 1-D ════════════════════════════ + + +@triton.jit +def k_add(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr): + pid = tl.program_id(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + x = tl.load(x_ptr + offs, mask=mask) + y = tl.load(y_ptr + offs, mask=mask) + tl.store(out_ptr + offs, x + y, mask=mask) + + +@case("a_add_masked", k_add, "a") +def _(dev): + n = 300 + return (3,), (T(n, dev=dev), T(n, dev=dev), T(n, dev=dev), n), {"BLOCK": 128} + + +@triton.jit +def k_copy(x_ptr, out_ptr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + v = tl.load(x_ptr + offs) + tl.store(out_ptr + offs, v) + + +@case("a_copy_exact", k_copy, "a") +def _(dev): + return (4,), (T(256, dev=dev), T(256, dev=dev)), {"BLOCK": 64} + + +@case( + "a_copy_ragged", + k_copy, + "a", + "unmasked ragged tail: out-of-bounds lanes land in the padding", +) +def _(dev): + return (3,), (T(150, dev=dev), T(150, dev=dev)), {"BLOCK": 64} + + +@triton.jit +def k_mask_le(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + m = offs <= n + v = tl.load(x_ptr + offs, mask=m) + tl.store(out_ptr + offs, v, mask=m) + + +@case("a_mask_le", k_mask_le, "a") +def _(dev): + return (4,), (T(100, dev=dev), T(100, dev=dev), 99), {"BLOCK": 32} + + +@triton.jit +def k_mask_or_pid0(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + pid = tl.program_id(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + m = (offs < n) | (pid == 0) + v = tl.load(x_ptr + offs, mask=m) + tl.store(out_ptr + offs, v, mask=m) + + +@case("a_mask_or_pid0", k_mask_or_pid0, "a") +def _(dev): + return (2,), (T(10, dev=dev), T(10, dev=dev), 10), {"BLOCK": 16} + + +@triton.jit +def k_mask_window(x_ptr, lo, hi, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + m = (offs >= lo) & (offs < hi) + tl.store(x_ptr + offs, 1.0, mask=m) + + +@case("a_mask_window", k_mask_window, "a") +def _(dev): + return (3,), (T(96, dev=dev), 5, 77), {"BLOCK": 32} + + +@triton.jit +def k_load_other(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + v = tl.load(x_ptr + offs, mask=offs < n, other=-1.0) + tl.store(out_ptr + offs, v) + + +@case("a_load_other", k_load_other, "a") +def _(dev): + return (2,), (T(40, dev=dev), T(64, dev=dev), 40), {"BLOCK": 32} + + +@triton.jit +def k_arange_start(x_ptr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * 64 + tl.arange(8, 8 + BLOCK) + tl.store(x_ptr + offs, 2.0) + + +@case("a_arange_start", k_arange_start, "a", "make_range with a non-zero start") +def _(dev): + return (3,), (T(192, dev=dev),), {"BLOCK": 32} + + +@triton.jit +def k_two_ranges_one_dim(x_ptr): + offs = tl.arange(0, 16) + tl.arange(16, 32) + tl.store(x_ptr + offs + tl.program_id(0) * 64, 1.0) + + +@case( + "a_two_ranges_one_dim", + k_two_ranges_one_dim, + "a", + "two make_ranges share one lane: 2i + 16, not i + j + 16", +) +def _(dev): + return (2,), (T(128, dev=dev),), {} + + +@triton.jit +def k_negative_shift(x_ptr, out_ptr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + v = tl.load(x_ptr + offs - 3, mask=offs >= 3) + tl.store(out_ptr + offs, v) + + +@case("a_negative_shift", k_negative_shift, "a") +def _(dev): + return (2,), (T(64, dev=dev), T(64, dev=dev)), {"BLOCK": 32} + + +@triton.jit +def k_scalar_per_program(x_ptr, out_ptr): + pid = tl.program_id(0) + v = tl.load(x_ptr + pid * 2) + tl.store(out_ptr + pid, v) + + +@case("a_scalar_per_program", k_scalar_per_program, "a") +def _(dev): + return (5,), (T(10, dev=dev), T(5, dev=dev)), {} + + +@case("a_copy_int8", k_copy, "a", "1-byte elements") +def _(dev): + return ( + (2,), + (T(64, dtype=torch.int8, dev=dev), T(64, dtype=torch.int8, dev=dev)), + {"BLOCK": 32}, + ) + + +@case("a_copy_fp16", k_copy, "a", "2-byte elements") +def _(dev): + return ( + (2,), + (T(64, dtype=torch.float16, dev=dev), T(64, dtype=torch.float16, dev=dev)), + {"BLOCK": 32}, + ) + + +@case("a_copy_bool", k_copy, "a", "i1 pointees, 1 byte each") +def _(dev): + return ( + (2,), + ( + T(64, dtype=torch.bool, dev=dev, init=None), + T(64, dtype=torch.bool, dev=dev, init=None), + ), + {"BLOCK": 32}, + ) + + +@case("a_copy_bf16", k_copy, "a", "2-byte bf16 elements") +def _(dev): + return ( + (2,), + (T(64, dtype=torch.bfloat16, dev=dev), T(64, dtype=torch.bfloat16, dev=dev)), + {"BLOCK": 32}, + ) + + +@case("a_copy_f64", k_copy, "a", "8-byte float elements") +def _(dev): + return ( + (2,), + (T(64, dtype=torch.float64, dev=dev), T(64, dtype=torch.float64, dev=dev)), + {"BLOCK": 32}, + ) + + +@case("a_copy_i64", k_copy, "a", "8-byte integer elements") +def _(dev): + return ( + (2,), + (T(64, dtype=torch.int64, dev=dev), T(64, dtype=torch.int64, dev=dev)), + {"BLOCK": 32}, + ) + + +@triton.jit +def k_same_width_ptr_bitcast(x_ptr, out_ptr, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + v = tl.load(x_ptr.to(tl.pointer_type(tl.int32)) + offs) + tl.store(out_ptr + offs, v) + + +@case( + "a_same_width_ptr_bitcast", + k_same_width_ptr_bitcast, + "a", + "an f32 argument read through an i32 pointer: the element width stays", +) +def _(dev): + return (1,), (T(16, dev=dev), T(16, dtype=torch.int32, dev=dev)), {"BLOCK": 16} + + +@triton.jit +def k_i64_limit(x_ptr, big, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(x_ptr + tl.minimum(offs, big - 4294967296), 1.0, mask=offs < big) + + +@case("a_i64_param", k_i64_limit, "a", "an i64 scalar argument (2**32 + 40)") +def _(dev): + return (2,), (T(64, dev=dev), 2**32 + 40), {"BLOCK": 32} + + +@triton.jit +def k_int64_offsets(x_ptr, stride, BLOCK: tl.constexpr): + pid = tl.program_id(0).to(tl.int64) + offs = pid * stride + tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 3.0) + + +@case("a_int64_offsets", k_int64_offsets, "a", "extsi to i64 before the row stride") +def _(dev): + return (4,), (T(4 * 40, dev=dev), 40), {"BLOCK": 32} + + +@triton.jit +def k_stride_param(x_ptr, s, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + v = tl.load(x_ptr + offs * s, mask=offs < n) + tl.store(x_ptr + offs * s + 1, v, mask=offs < n) + + +@case("a_stride_param", k_stride_param, "a") +def _(dev): + return (2,), (T(3 * 50, dev=dev), 3, 50), {"BLOCK": 32} + + +@triton.jit +def k_hints(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + offs = tl.max_contiguous(tl.multiple_of(offs, BLOCK), BLOCK) + tl.assume(n > 0) + v = tl.load(x_ptr + offs, mask=offs < n) + tl.store(out_ptr + offs, v, mask=offs < n) + + +@case("a_hints", k_hints, "a", "multiple_of / max_contiguous / assume") +def _(dev): + return (2,), (T(50, dev=dev), T(50, dev=dev), 50), {"BLOCK": 32} + + +@triton.jit +def _store_helper(p, offs, n): + tl.store(p + offs, 1.0, mask=offs < n) + + +@triton.jit +def _load_then_store_helper(p, offs, n): + v = tl.load(p + offs, mask=offs < n) + _store_helper(p + 64, offs, n) + return v + + +@triton.jit +def k_inlined_helpers(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + v = _load_then_store_helper(x_ptr, offs, n) + tl.store(out_ptr + offs, v, mask=offs < n) + + +@case( + "a_inlined_helpers", + k_inlined_helpers, + "a", + "memory ops in two levels of inlined @triton.jit helpers (callsite locs)", +) +def _(dev): + return (2,), (T(128, dev=dev), T(64, dev=dev), 50), {"BLOCK": 32} + + +@triton.jit +def k_debug_ops(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.static_assert(BLOCK % 16 == 0) + tl.device_assert(n > 0, "n must be positive") + tl.debug_barrier() + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + +@case("a_debug_ops", k_debug_ops, "a", "static_assert / device_assert / debug_barrier") +def _(dev): + return (2,), (T(64, dev=dev), 40), {"BLOCK": 32} + + +# ═════════════════ b: min / max / where / integer division ═════════════════ + + +@triton.jit +def k_clamp_min(x_ptr, out_ptr, lim, BLOCK: tl.constexpr): + offs = tl.minimum(tl.program_id(0) * BLOCK + tl.arange(0, BLOCK), lim) + v = tl.load(x_ptr + offs) + tl.store(out_ptr + offs, v) + + +@case("b_clamp_min", k_clamp_min, "b") +def _(dev): + return (4,), (T(100, dev=dev), T(100, dev=dev), 99), {"BLOCK": 32} + + +@triton.jit +def k_clamp_both(x_ptr, lo, hi, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + c = tl.maximum(tl.minimum(offs - 4, hi), lo) + tl.store(x_ptr + c, 1.0) + + +@case("b_clamp_both", k_clamp_both, "b") +def _(dev): + return (3,), (T(80, dev=dev), 2, 70), {"BLOCK": 32} + + +@triton.jit +def k_where_offsets(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + offs = tl.where(offs < n, offs, n - 1) + tl.store(x_ptr + offs, 1.0) + + +@case("b_where_offsets", k_where_offsets, "b") +def _(dev): + return (4,), (T(100, dev=dev), 100), {"BLOCK": 32} + + +@case( + "b_where_pointer", + RK.where_pointer, + "b", + "arith.select over two pointers of one base", +) +def _(dev): + return (1,), (T(128, dev=dev), 9), {} + + +@triton.jit +def k_modulo(x_ptr, out_ptr, m, BLOCK: tl.constexpr): + offs = (tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)) % m + v = tl.load(x_ptr + offs) + tl.store(out_ptr + offs, v) + + +@case("b_modulo", k_modulo, "b") +def _(dev): + return (4,), (T(100, dev=dev), T(100, dev=dev), 37), {"BLOCK": 32} + + +@triton.jit +def k_div_mod_2d(x_ptr, W, stride, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + row = offs // W + col = offs % W + tl.store(x_ptr + row * stride + col, 1.0, mask=offs < n) + + +@case("b_div_mod_2d", k_div_mod_2d, "b", "flat index -> (row, col) by runtime divisor") +def _(dev): + return (3,), (T(8 * 20, dev=dev), 12, 20, 90), {"BLOCK": 32} + + +@triton.jit +def k_signed_div(x_ptr, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) - 20 + tl.store(x_ptr + offs // 4 + 8, 1.0) + + +@case( + "b_signed_div_trunc", + k_signed_div, + "b", + "divsi truncates toward zero on negative dividends", +) +def _(dev): + return (1,), (T(64, dev=dev),), {"BLOCK": 64} + + +@triton.jit +def k_signed_rem(x_ptr, d, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) - 20 + tl.store(x_ptr + offs % d + 10, 1.0) + + +@case("b_signed_rem", k_signed_rem, "b", "remsi takes the dividend's sign") +def _(dev): + return (1,), (T(64, dev=dev), 7), {"BLOCK": 64} + + +@case("b_unsigned_index", RK.unsigned_index, "b", "divui, cmpi ult, extui") +def _(dev): + return (7,), (T(8, dev=dev), 5), {} + + +@triton.jit +def k_bool_extui(x_ptr, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + offs * 2 + (offs > 5).to(tl.int32), 1.0) + + +@case("b_bool_extui", k_bool_extui, "b", "extui of an i1 compare") +def _(dev): + return (1,), (T(64, dev=dev),), {"BLOCK": 32} + + +@triton.jit +def k_unsigned_min(x_ptr, lim, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + tl.minimum(offs.to(tl.uint32), lim.to(tl.uint32)), 1.0) + + +@case("b_unsigned_min", k_unsigned_min, "b", "minui, then extui to i64 for the address") +def _(dev): + return (1,), (T(64, dev=dev), 20), {"BLOCK": 32} + + +@triton.jit +def k_unsigned_cmp(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs.to(tl.uint32) < n.to(tl.uint32)) + + +@case( + "b_unsigned_cmp_highbit", + k_unsigned_cmp, + "b", + "cmpi ult with n = -1 (2**32 - 1 unsigned): the unsigned-operand obligation fails", +) +def _(dev): + return (1,), (T(16, dev=dev), -1), {"BLOCK": 16} + + +@triton.jit +def k_nested_select(x_ptr, a, b, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + o = tl.where(offs < a, offs, tl.where(offs < b, offs * 2, 0)) + tl.store(x_ptr + o, 1.0) + + +@case("b_nested_select", k_nested_select, "b") +def _(dev): + return (1,), (T(128, dev=dev), 5, 20), {"BLOCK": 32} + + +# ══════════════════════ c: 2-D tiles and tensor views ══════════════════════ + + +@triton.jit +def k_tile2d( + x_ptr, out_ptr, M, N, sxm, sxn, sym, syn, BM: tl.constexpr, BN: tl.constexpr +): + # x: the input view, y: the output + rm = tl.program_id(0) * BM + tl.arange(0, BM) + rn = tl.program_id(1) * BN + tl.arange(0, BN) + m = (rm[:, None] < M) & (rn[None, :] < N) + v = tl.load(x_ptr + rm[:, None] * sxm + rn[None, :] * sxn, mask=m) + tl.store(out_ptr + rm[:, None] * sym + rn[None, :] * syn, v, mask=m) + + +@case("c_tile2d", k_tile2d, "c") +def _(dev): + M, N = 20, 24 + x, out = T(M, N, dev=dev), T(M, N, dev=dev) + return (2, 2), (x, out, M, N, *x.stride(), *out.stride()), {"BM": 16, "BN": 16} + + +@case( + "c_strided_view", k_tile2d, "c", "x[:, ::2]: a non-contiguous view, strides (2N, 2)" +) +def _(dev): + M, N = 12, 10 + x = T(M, 2 * N, dev=dev)[:, ::2] + out = T(M, N, dev=dev) + return (1, 1), (x, out, M, N, *x.stride(), *out.stride()), {"BM": 16, "BN": 16} + + +@case("c_transposed_view", k_tile2d, "c", "x.t(): strides (1, M)") +def _(dev): + M, N = 12, 20 + x = T(N, M, dev=dev).t() + out = T(M, N, dev=dev) + return (1, 2), (x, out, M, N, *x.stride(), *out.stride()), {"BM": 16, "BN": 16} + + +@case("c_expand_view_stride0", k_tile2d, "c", "x.expand(M, N): a stride-0 dim") +def _(dev): + M, N = 6, 10 + x = T(N, dev=dev).expand(M, N) + out = T(M, N, dev=dev) + return (1, 1), (x, out, M, N, *x.stride(), *out.stride()), {"BM": 8, "BN": 16} + + +@triton.jit +def k_block_ptr(x_ptr, out_ptr, M, N, sm, sn, BM: tl.constexpr, BN: tl.constexpr): + pid = tl.program_id(0) + src = tl.make_block_ptr(x_ptr, (M, N), (sm, sn), (pid * BM, 0), (BM, BN), (1, 0)) + v = tl.load(src, boundary_check=(0, 1), padding_option="zero") + dst = tl.make_block_ptr(out_ptr, (M, N), (sm, sn), (pid * BM, 0), (BM, BN), (1, 0)) + tl.store(dst, v, boundary_check=(0, 1)) + + +@case( + "c_block_ptr", + k_block_ptr, + "c", + "block pointers, rewritten to pointer tiles + masks in TTIR", +) +def _(dev): + M, N = 20, 12 + x, out = T(M, N, dev=dev), T(M, N, dev=dev) + return (3,), (x, out, M, N, *x.stride()), {"BM": 8, "BN": 16} + + +@triton.jit +def k_block_ptr_loop(x_ptr, out_ptr, K, BK: tl.constexpr): + src = tl.make_block_ptr(x_ptr, (K,), (1,), (0,), (BK,), (0,)) + acc = tl.zeros((BK,), dtype=tl.float32) + for k in range(0, tl.cdiv(K, BK)): + acc += tl.load(src, boundary_check=(0,), padding_option="zero") + src = tl.advance(src, (BK,)) + tl.store(out_ptr + tl.arange(0, BK), acc) + + +@case( + "c_block_ptr_loop", + k_block_ptr_loop, + "c", + "an advanced block pointer: the loop carries integer offsets", +) +def _(dev): + return (1,), (T(40, dev=dev), T(16, dev=dev), 40), {"BK": 16} + + +@triton.jit +def k_broadcast_row(x_ptr, b_ptr, out_ptr, M, N, BM: tl.constexpr, BN: tl.constexpr): + rm = tl.arange(0, BM) + rn = tl.arange(0, BN) + bias = tl.load(b_ptr + rn, mask=rn < N) + m = (rm[:, None] < M) & (rn[None, :] < N) + v = tl.load(x_ptr + rm[:, None] * N + rn[None, :], mask=m) + tl.store(out_ptr + rm[:, None] * N + rn[None, :], v + bias[None, :], mask=m) + + +@case("c_broadcast_row", k_broadcast_row, "c") +def _(dev): + M, N = 6, 12 + return ( + (1,), + (T(M, N, dev=dev), T(N, dev=dev), T(M, N, dev=dev), M, N), + {"BM": 8, "BN": 16}, + ) + + +@triton.jit +def k_broadcast_col(x_ptr, out_ptr, M, BM: tl.constexpr, BN: tl.constexpr): + rm = tl.arange(0, BM) + rn = tl.arange(0, BN) + col = tl.load(x_ptr + rm[:, None] + rn[None, :] * 0, mask=rm[:, None] < M) + tl.store(out_ptr + rm[:, None] * BN + rn[None, :], col, mask=rm[:, None] < M) + + +@case( + "c_broadcast_col", + k_broadcast_col, + "c", + "one column broadcast along dim 1 (stride 0)", +) +def _(dev): + return (1,), (T(8, dev=dev), T(8 * 8, dev=dev), 7), {"BM": 8, "BN": 8} + + +@case( + "c_tile3d_shared_arange", + RK.tile3d_shared_arange, + "c", + "one make_range on all three dims", +) +def _(dev): + return (1,), (T(64, dev=dev),), {"N": 4} + + +@triton.jit +def k_one_range_two_dims(x_ptr, N: tl.constexpr): + r = tl.arange(0, N) + tl.store(x_ptr + r[:, None] * (2 * N) + r[None, :] + tl.program_id(0) * N, 1.0) + + +@case("c_one_range_two_dims", k_one_range_two_dims, "c") +def _(dev): + return (2,), (T(8 * 16, dev=dev),), {"N": 8} + + +@triton.jit +def k_softmax_rows(x_ptr, out_ptr, n_cols, stride, BLOCK: tl.constexpr): + row = tl.program_id(0) + cols = tl.arange(0, BLOCK) + x = tl.load(x_ptr + row * stride + cols, mask=cols < n_cols, other=-float("inf")) + e = tl.exp(x - tl.max(x, axis=0)) + tl.store(out_ptr + row * stride + cols, e / tl.sum(e, axis=0), mask=cols < n_cols) + + +@case("c_softmax_rows", k_softmax_rows, "c") +def _(dev): + R, C = 5, 13 + return (R,), (T(R, 16, dev=dev), T(R, 16, dev=dev), C, 16), {"BLOCK": 16} + + +@triton.jit +def k_transpose_store(x_ptr, out_ptr, M, N, BM: tl.constexpr, BN: tl.constexpr): + rm = tl.arange(0, BM) + rn = tl.arange(0, BN) + src = x_ptr + rm[:, None] * N + rn[None, :] + dst = out_ptr + rn[:, None] * M + rm[None, :] + v = tl.load(src, mask=(rm[:, None] < M) & (rn[None, :] < N)) + tl.store(dst, tl.trans(v), mask=(rn[:, None] < N) & (rm[None, :] < M)) + + +@case("c_transpose_store", k_transpose_store, "c") +def _(dev): + M, N = 6, 10 + return (1,), (T(M, N, dev=dev), T(N, M, dev=dev), M, N), {"BM": 8, "BN": 16} + + +@triton.jit +def k_row_sum(x_ptr, out_ptr, M, N, BM: tl.constexpr, BN: tl.constexpr): + rm = tl.program_id(0) * BM + tl.arange(0, BM) + rn = tl.arange(0, BN) + m = (rm[:, None] < M) & (rn[None, :] < N) + v = tl.load(x_ptr + rm[:, None] * N + rn[None, :], mask=m, other=0.0) + tl.store(out_ptr + rm, tl.sum(v, axis=1), mask=rm < M) + + +@case("c_row_sum", k_row_sum, "c", "a reduction's 1-D result stored by the 1-D range") +def _(dev): + M, N = 11, 12 + return (2,), (T(M, N, dev=dev), T(M, dev=dev), M, N), {"BM": 8, "BN": 16} + + +# ═══════════════════════ d: program ids and the grid ═══════════════════════ + + +@triton.jit +def k_pid3d(x_ptr, BLOCK: tl.constexpr): + base = tl.program_id(0) * 100 + tl.program_id(1) * 20 + tl.program_id(2) * 7 + tl.store(x_ptr + base + tl.arange(0, BLOCK), 1.0) + + +@case("d_pid3d", k_pid3d, "d", "program ids on axes 0, 1 and 2") +def _(dev): + return (2, 3, 2), (T(200, dev=dev),), {"BLOCK": 4} + + +@triton.jit +def k_num_programs(x_ptr, out_ptr): + p0 = tl.program_id(0) + p1 = tl.program_id(1) + v = tl.load(x_ptr + p1 * tl.num_programs(0) + p0) + tl.store(out_ptr + p0 * tl.num_programs(1) + p1, v) + + +@case("d_num_programs", k_num_programs, "d") +def _(dev): + return (3, 2), (T(6, dev=dev), T(6, dev=dev)), {} + + +@triton.jit +def k_num_programs_3d(x_ptr, BLOCK: tl.constexpr): + np0 = tl.num_programs(0) + flat = ( + tl.program_id(2) * tl.num_programs(1) + tl.program_id(1) + ) * np0 + tl.program_id(0) + tl.store(x_ptr + flat * BLOCK + tl.arange(0, BLOCK), 1.0) + + +@case("d_num_programs_3d", k_num_programs_3d, "d", "num_programs on axes 0, 1 and 2") +def _(dev): + return (2, 2, 3), (T(12 * 4, dev=dev),), {"BLOCK": 4} + + +@triton.jit +def k_num_programs_unread_axis(x_ptr, BLOCK: tl.constexpr): + # graph.pid_axes must list axis 1 though no program_id reads it + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(x_ptr + offs * tl.num_programs(1), 1.0) + + +@case( + "d_num_programs_unread_axis", + k_num_programs_unread_axis, + "d", + "num_programs(1) without program_id(1): every program on axis 1 stores alike", +) +def _(dev): + return (2, 3), (T(24, dev=dev),), {"BLOCK": 4} + + +@triton.jit +def k_grid_stride(x_ptr, out_ptr, n, n_tiles, BLOCK: tl.constexpr): + for tile in range(tl.program_id(0), n_tiles, tl.num_programs(0)): + offs = tile * BLOCK + tl.arange(0, BLOCK) + v = tl.load(x_ptr + offs, mask=offs < n) + tl.store(out_ptr + offs, v * 2, mask=offs < n) + + +@case( + "d_grid_stride_loop", + k_grid_stride, + "d", + "for tile in range(pid, n_tiles, num_programs)", +) +def _(dev): + n, B = 150, 16 + return (3,), (T(n, dev=dev), T(n, dev=dev), n, triton.cdiv(n, B)), {"BLOCK": B} + + +@triton.jit +def k_grouped_order( + c_ptr, M, N, BM: tl.constexpr, BN: tl.constexpr, GROUP_M: tl.constexpr +): + pid = tl.program_id(0) + num_pid_m = tl.cdiv(M, BM) + num_pid_n = tl.cdiv(N, BN) + num_pid_in_group = GROUP_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_M + group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_M) + pid_m = first_pid_m + (pid % num_pid_in_group) % group_size_m + pid_n = (pid % num_pid_in_group) // group_size_m + rm = pid_m * BM + tl.arange(0, BM) + rn = pid_n * BN + tl.arange(0, BN) + m = (rm[:, None] < M) & (rn[None, :] < N) + tl.store(c_ptr + rm[:, None] * N + rn[None, :], 1.0, mask=m) + + +@case( + "d_grouped_order", + k_grouped_order, + "d", + "grouped (swizzled) tile order: div / mod / min on pid", +) +def _(dev): + M, N = 40, 24 + grid = (triton.cdiv(M, 8) * triton.cdiv(N, 8),) + return grid, (T(M, N, dev=dev), M, N), {"BM": 8, "BN": 8, "GROUP_M": 2} + + +@triton.jit +def k_pid_axis1_rows(x_ptr, stride, n_cols, BLOCK: tl.constexpr): + row = tl.program_id(1) + cols = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(x_ptr + row * stride + cols, 1.0, mask=cols < n_cols) + + +@case("d_pid_axis1_rows", k_pid_axis1_rows, "d") +def _(dev): + return (2, 4), (T(4, 40, dev=dev), 40, 27), {"BLOCK": 16} + + +# ═══════════════════ e: loops and loop-carried pointers ═══════════════════ + + +@triton.jit +def k_loop_ptr_advance(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + p = x_ptr + offs + for k in range(0, n): + v = tl.load(p) + tl.store(out_ptr + k * BLOCK + offs, v) + p += BLOCK + + +@case("e_loop_ptr_advance", k_loop_ptr_advance, "e") +def _(dev): + return (1,), (T(5 * 16, dev=dev), T(5 * 16, dev=dev), 5), {"BLOCK": 16} + + +@triton.jit +def k_matmul( + a_ptr, + b_ptr, + c_ptr, + M, + N, + K, + sam, + sak, + sbk, + sbn, + scm, + scn, + BM: tl.constexpr, + BN: tl.constexpr, + BK: tl.constexpr, +): + rm = tl.program_id(0) * BM + tl.arange(0, BM) + rn = tl.program_id(1) * BN + tl.arange(0, BN) + rk = tl.arange(0, BK) + a_ptrs = a_ptr + rm[:, None] * sam + rk[None, :] * sak + b_ptrs = b_ptr + rk[:, None] * sbk + rn[None, :] * sbn + acc = tl.zeros((BM, BN), dtype=tl.float32) + for k in range(0, tl.cdiv(K, BK)): + k_left = K - k * BK + a = tl.load(a_ptrs, mask=(rm[:, None] < M) & (rk[None, :] < k_left), other=0.0) + b = tl.load(b_ptrs, mask=(rk[:, None] < k_left) & (rn[None, :] < N), other=0.0) + acc += tl.dot(a, b) + a_ptrs += BK * sak + b_ptrs += BK * sbk + c_mask = (rm[:, None] < M) & (rn[None, :] < N) + tl.store(c_ptr + rm[:, None] * scm + rn[None, :] * scn, acc, mask=c_mask) + + +@case( + "e_matmul", + k_matmul, + "e", + "two 2-D pointer tiles advanced by BK * stride, K-tail masks", +) +def _(dev): + M, N, K = 20, 18, 40 + a, b, c = T(M, K, dev=dev), T(K, N, dev=dev), T(M, N, dev=dev) + args = (a, b, c, M, N, K, *a.stride(), *b.stride(), *c.stride()) + return (2, 2), args, {"BM": 16, "BN": 16, "BK": 16} + + +@triton.jit +def k_loop_iv_offset(x_ptr, lo, hi): + for k in range(lo, hi): + tl.store(x_ptr + k * 2 + tl.program_id(0), 1.0) + + +@case( + "e_loop_iv_offset", + k_loop_iv_offset, + "e", + "runtime lower bound, the induction variable in the address", +) +def _(dev): + return (2,), (T(40, dev=dev), 3, 11), {} + + +@triton.jit +def k_loop_step(x_ptr, lo, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + for k in range(lo, n, 3): + tl.store(x_ptr + k * BLOCK + offs, 1.0) + + +@case("e_loop_step3", k_loop_step, "e") +def _(dev): + return (1,), (T(20 * 4, dev=dev), 2, 17), {"BLOCK": 4} + + +@triton.jit +def k_zero_trip(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(out_ptr + tl.program_id(0), 0.0) + p = x_ptr + offs + for k in range(0, n): + tl.store(p, 1.0) + p += BLOCK + + +@case("e_zero_trip", k_zero_trip, "e", "n = 0: the loop's store has no footprint") +def _(dev): + return (2,), (T(64, dev=dev), T(2, dev=dev), 0), {"BLOCK": 16} + + +@case("e_some_trips", k_zero_trip, "e") +def _(dev): + return (2,), (T(64, dev=dev), T(2, dev=dev), 3), {"BLOCK": 16} + + +@case( + "e_expand_iterarg_3d", + RK.expand_iterarg_3d, + "e", + "a loop-carried [N, N] tile expanded to 3-D", +) +def _(dev): + x = buf(64, dev=dev) + return (1,), (x, T(64, dev=dev), 2), {"N": 4} + + +@case( + "e_expand_iterarg_mask", + RK.expand_iterarg_mask, + "e", + "a loop-carried 1-D tile expanded to 2-D, masked", +) +def _(dev): + return (1,), (T(64, dev=dev), T(16, dev=dev), 3, 2), {"N": 4} + + +@case("e_two_step_advance", RK.loop_two_step_advance, "e", "two addptrs per iteration") +def _(dev): + return (1,), (T(64, dev=dev), 5, 4), {} + + +@triton.jit +def k_if_in_loop(x_ptr, n): + for k in range(0, n): + if k % 2 == 0: + tl.store(x_ptr + k, 1.0) + + +@case("e_if_in_loop", k_if_in_loop, "e", "an scf.if on the induction variable") +def _(dev): + return (1,), (T(16, dev=dev), 9), {} + + +@triton.jit +def k_iv_mask(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + p = x_ptr + offs + for k in range(0, tl.cdiv(n, BLOCK)): + tl.store(p, 1.0, mask=offs < n - k * BLOCK) + p += BLOCK + + +@case("e_iv_mask", k_iv_mask, "e", "a mask on the induction variable and the lane") +def _(dev): + return (1,), (T(64, dev=dev), 45), {"BLOCK": 16} + + +@triton.jit +def k_static_range_in_loop(x_ptr, n): + p = x_ptr + for k in range(0, n): + for j in tl.static_range(3): + tl.store(p + j, 1.0) + p += 4 + + +@case( + "e_static_range_in_loop", + k_static_range_in_loop, + "e", + "an unrolled inner loop: three stores on one line", +) +def _(dev): + return (1,), (T(32, dev=dev), 5), {} + + +@triton.jit +def k_pid_trips(x_ptr): + pid = tl.program_id(0) + for k in range(0, pid + 1): + tl.store(x_ptr + pid * 8 + k, 1.0) + + +@case( + "e_pid_dependent_trips", k_pid_trips, "e", "a trip count that differs per program" +) +def _(dev): + return (5,), (T(40, dev=dev),), {} + + +@triton.jit +def k_unsigned_loop(x_ptr, n): + for k in range(0, n.to(tl.uint32)): + tl.store(x_ptr + k, 1.0) + + +@case("e_unsigned_loop", k_unsigned_loop, "e", "an unsigned induction variable") +def _(dev): + return (1,), (T(16, dev=dev), 7), {} + + +@triton.jit +def k_tl_range(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + for k in tl.range(0, n, num_stages=2): + tl.store(x_ptr + k * BLOCK + offs, 1.0) + + +@case("e_tl_range", k_tl_range, "e", "tl.range with num_stages") +def _(dev): + return (1,), (T(6 * 8, dev=dev), 6), {"BLOCK": 8} + + +@triton.jit +def k_static_unroll(x_ptr, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + for j in tl.static_range(4): + tl.store(x_ptr + j * 2 * BLOCK + offs, 1.0) + + +@case( + "e_static_unroll", + k_static_unroll, + "e", + "four unrolled stores on one line, no scf.for", +) +def _(dev): + return (1,), (T(8 * 8, dev=dev),), {"BLOCK": 8} + + +@triton.jit +def k_negative_delta(x_ptr, out_ptr, n): + p = x_ptr + 60 + for k in range(0, n): + v = tl.load(p) + tl.store(out_ptr + k, v) + p -= 3 + + +@case("e_negative_delta", k_negative_delta, "e") +def _(dev): + return (1,), (T(64, dev=dev), T(16, dev=dev), 12), {} + + +@triton.jit +def k_per_lane_delta(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + p = x_ptr + offs + for k in range(0, n): + tl.store(p, 1.0) + p += offs + 1 + + +@case("e_per_lane_delta", k_per_lane_delta, "e", "a loop-invariant per-lane advance") +def _(dev): + return (1,), (T(64, dev=dev), 4), {"BLOCK": 8} + + +@triton.jit +def k_expand_axis1(x_ptr, out_ptr, n, N: tl.constexpr): + r = tl.arange(0, N) + p = x_ptr + r * N + for k in range(0, n): + v = tl.load(p[:, None] + r[None, :]) + tl.store(out_ptr + r[:, None] * N + r[None, :] + k * N * N, v) + p += 1 + + +@case( + "e_expand_axis1", k_expand_axis1, "e", "a loop-carried 1-D tile expanded at axis 1" +) +def _(dev): + return (1,), (T(40, dev=dev), T(3 * 16, dev=dev), 3), {"N": 4} + + +@triton.jit +def k_expand_per_lane_delta(x_ptr, out_ptr, n, N: tl.constexpr): + r = tl.arange(0, N) + p = x_ptr + r + for k in range(0, n): + v = tl.load(p[None, :] + r[:, None] * N) + tl.store(out_ptr + r[:, None] * N + r[None, :] + k * N * N, v) + p += r + 1 + + +@case( + "e_expand_per_lane_delta", + k_expand_per_lane_delta, + "e", + "a loop-carried 1-D tile expanded to 2-D whose delta is per lane", +) +def _(dev): + return (1,), (T(64, dev=dev), T(3 * 16, dev=dev), 3), {"N": 4} + + +@triton.jit +def k_negative_step(x_ptr, out_ptr, n): + for k in range(n - 1, -1, -1): + v = tl.load(x_ptr + k) + tl.store(out_ptr + (n - 1 - k), v) + + +@case("e_negative_step", k_negative_step, "e", "range(n - 1, -1, -1)") +def _(dev): + return (1,), (T(16, dev=dev), T(16, dev=dev), 9), {} + + +@triton.jit +def k_negative_lower(x_ptr, n): + for k in range(-3, n): + tl.store(x_ptr + k + 3, 1.0) + + +@case("e_negative_lower", k_negative_lower, "e", "a negative lower bound") +def _(dev): + return (1,), (T(16, dev=dev), 6), {} + + +@triton.jit +def k_zero_trip_invariant(x_ptr, out_ptr, n): + tl.store(out_ptr + tl.program_id(0), 0.0) + for k in range(0, n): + tl.store(x_ptr + tl.program_id(0), 1.0) + + +@case( + "e_zero_trip_invariant", + k_zero_trip_invariant, + "e", + "n = 0: an in-loop store that reads no loop value still has no footprint", +) +def _(dev): + return (2,), (T(4, dev=dev), T(2, dev=dev), 0), {} + + +# ═══════════════════════════ f: structured if ═══════════════════════════ + + +@triton.jit +def k_if_else_pid(a_ptr, b_ptr): + pid = tl.program_id(0) + if pid % 2 == 0: + tl.store(a_ptr + pid, 1.0) + else: + tl.store(b_ptr + pid // 2, 2.0) + + +@case("f_if_else_pid", k_if_else_pid, "f") +def _(dev): + return (5,), (T(5, dev=dev), T(5, dev=dev)), {} + + +@triton.jit +def k_if_param(x_ptr, n, BLOCK: tl.constexpr): + pid = tl.program_id(0) + if pid < n: + tl.store(x_ptr + pid * BLOCK + tl.arange(0, BLOCK), 1.0) + + +@case("f_if_param", k_if_param, "f") +def _(dev): + return (4,), (T(4 * 8, dev=dev), 3), {"BLOCK": 8} + + +@triton.jit +def k_nested_if(x_ptr, lo, hi): + pid = tl.program_id(0) + if pid > lo: + if pid < hi: + tl.store(x_ptr + pid, 1.0) + else: + tl.store(x_ptr + pid + 10, 2.0) + + +@case("f_nested_if", k_nested_if, "f") +def _(dev): + return (6,), (T(16, dev=dev), 1, 4), {} + + +@triton.jit +def k_if_pointer_result(x_ptr, n): + pid = tl.program_id(0) + if pid == 0: + p = x_ptr + 1 + else: + p = x_ptr + pid * 3 + n + tl.store(p, 1.0) + + +@case( + "f_if_pointer_result", + k_if_pointer_result, + "f", + "an scf.if yielding a pointer of one base", +) +def _(dev): + return (3,), (T(16, dev=dev), 2), {} + + +# ═════════════════════════ g: atomics (results unused) ═════════════════════════ + + +@triton.jit +def k_atomic_add_masked(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.atomic_add(x_ptr + offs, 1.0, mask=offs < n) + + +@case("g_atomic_add_masked", k_atomic_add_masked, "g") +def _(dev): + return (3,), (T(70, dev=dev), 70), {"BLOCK": 32} + + +@triton.jit +def k_atomic_counter(cnt_ptr, x_ptr): + tl.atomic_add(cnt_ptr, 1) + tl.store(x_ptr + tl.program_id(0), 1.0) + + +@case("g_atomic_counter", k_atomic_counter, "g", "a scalar counter every program bumps") +def _(dev): + return (4,), (T(1, dtype=torch.int32, dev=dev, init=None), T(4, dev=dev)), {} + + +@triton.jit +def k_atomic_max_int(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.atomic_max(x_ptr + offs % 8, offs - 10, mask=offs < n) + + +@case("g_atomic_max_int", k_atomic_max_int, "g") +def _(dev): + return (1,), (T(8, dtype=torch.int32, dev=dev, init=None), 20), {"BLOCK": 32} + + +@triton.jit +def k_atomic_cas(lock_ptr): + tl.atomic_cas(lock_ptr + tl.program_id(0) * 2, 0, 1) + + +@case("g_atomic_cas", k_atomic_cas, "g") +def _(dev): + return (4,), (T(8, dtype=torch.int32, dev=dev, init=None),), {} + + +@triton.jit +def k_atomic_xchg_2d(x_ptr, M, N: tl.constexpr): + rm = tl.arange(0, 8) + rn = tl.arange(0, N) + tl.atomic_xchg(x_ptr + rm[:, None] * N + rn[None, :], 5, mask=rm[:, None] < M) + + +@case("g_atomic_xchg_2d", k_atomic_xchg_2d, "g") +def _(dev): + return (1,), (T(8 * 4, dtype=torch.int32, dev=dev, init=None), 6), {"N": 4} + + +@triton.jit +def k_atomic_histogram(h_ptr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.atomic_add(h_ptr + offs % 8, 1) + + +@case("g_atomic_histogram", k_atomic_histogram, "g") +def _(dev): + return (2,), (T(8, dtype=torch.int32, dev=dev, init=None),), {"BLOCK": 16} + + +@triton.jit +def k_atomic_in_loop(x_ptr, n): + for k in range(0, n): + tl.atomic_add(x_ptr + k * 2 + tl.program_id(0), 1.0) + + +@case("g_atomic_in_loop", k_atomic_in_loop, "g") +def _(dev): + return (2,), (T(16, dev=dev, init=None), 6), {} + + +@triton.jit +def k_atomic_result_value(cnt_ptr, out_ptr): + old = tl.atomic_add(cnt_ptr, 1) + tl.store(out_ptr + tl.program_id(0), old) + + +@case( + "g_atomic_result_as_value", + k_atomic_result_value, + "g", + "an observation stored as data, not an address", +) +def _(dev): + return ( + (4,), + (T(1, dtype=torch.int32, dev=dev, init=None), T(4, dtype=torch.int32, dev=dev)), + {}, + ) + + +# ═════════════════════════ h: autotune-style configs ═════════════════════════ + + +@triton.autotune( + configs=[triton.Config({"BLOCK": b}) for b in (16, 32, 64)], + key=["n"], +) +@triton.jit +def k_autotuned(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + v = tl.load(x_ptr + offs, mask=offs < n) + tl.store(out_ptr + offs, v + 1.0, mask=offs < n) + + +def _autotune_case(cfg_kwargs: dict): + def build(dev): + n = 100 + return ( + (triton.cdiv(n, cfg_kwargs["BLOCK"]),), + (T(n, dev=dev), T(n, dev=dev), n), + dict(cfg_kwargs), + ) + + return build + + +for _i, _cfg in enumerate(k_autotuned.configs): + case(f"h_autotune_cfg{_i}", k_autotuned.fn, "h", f"config {_cfg.kwargs}")( + _autotune_case(dict(_cfg.kwargs)) + ) + + +# ═══════════ i: accesses the model over-approximates (skipped, not compared) ═══════════ + + +@triton.jit +def k_datadep_mask(flags_ptr, out_ptr, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + f = tl.load(flags_ptr + offs) + tl.store(out_ptr + offs, 1.0, mask=f != 0) + + +@case( + "i_datadep_mask", + k_datadep_mask, + "i", + "a mask from loaded data: dropped, the store skipped", +) +def _(dev): + return ( + (1,), + ( + T(16, dtype=torch.int32, dev=dev, init=[i % 3 for i in range(16)]), + T(16, dev=dev), + ), + {"BLOCK": 16}, + ) + + +@triton.jit +def k_atomic_max_float(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.atomic_max(x_ptr + offs % 8, offs.to(tl.float32) - 10.0, mask=offs < n) + tl.store(out_ptr + offs, 1.0, mask=offs < n) + + +@case( + "i_atomic_max_float", + k_atomic_max_float, + "i", + "float max: two integer atomics masked by the value's sign", +) +def _(dev): + return (1,), (T(8, dev=dev, init=None), T(32, dev=dev), 20), {"BLOCK": 32} + + +@triton.jit +def k_guarded_branch(x_ptr, out_ptr): + pid = tl.program_id(0) + v = tl.load(x_ptr + pid) + if v > 2.0: + tl.store(out_ptr + pid, v) + + +@case( + "i_guarded_branch", + k_guarded_branch, + "i", + "a branch on loaded data: the store is guarded, skipped", +) +def _(dev): + return (5,), (T(5, dev=dev), T(5, dev=dev)), {} + + +@triton.jit +def k_observed_mask_path(cnt_ptr, out_ptr): + old = tl.atomic_add(cnt_ptr, 1) + tl.store(out_ptr + tl.program_id(0), 1, mask=old < 2) + if old == 0: + tl.store(out_ptr + 8, 2) + + +@case( + "i_observed_mask_path", + k_observed_mask_path, + "i", + "an observation in a mask and in a branch condition", +) +def _(dev): + return ( + (4,), + ( + T(1, dtype=torch.int32, dev=dev, init=None), + T(16, dtype=torch.int32, dev=dev), + ), + {}, + ) + + +# ═══════════ j: shapes outside the reader's model (refuse OR conform) ═══════════ + + +@triton.jit +def k_two_loops(x_ptr, n): + for k in range(0, n): + tl.store(x_ptr + k, 1.0) + for j in range(0, 2 * n): + tl.store(x_ptr + 32 + j, 2.0) + + +@case("j_two_loops", k_two_loops, "j", "two sequential scf.for") +def _(dev): + return (1,), (T(64, dev=dev), 5), {} + + +@triton.jit +def k_nested_loops(x_ptr, n): + for k in range(0, n): + for j in range(0, 3): + tl.store(x_ptr + k * 4 + j, 1.0) + + +@case("j_nested_loops", k_nested_loops, "j", "an scf.for in an scf.for") +def _(dev): + return (1,), (T(32, dev=dev), 5), {} + + +@triton.jit +def k_loop_under_if(x_ptr, n): + if tl.program_id(0) == 0: + for k in range(0, n): + tl.store(x_ptr + k, 1.0) + + +@case("j_loop_under_if", k_loop_under_if, "j", "an scf.for under an scf.if") +def _(dev): + return (2,), (T(16, dev=dev), 5), {} + + +@triton.jit +def k_while(x_ptr, n): + k = 0 + while k < n: + tl.store(x_ptr + k, 1.0) + k += 1 + + +@case("j_while", k_while, "j", "an scf.while") +def _(dev): + return (1,), (T(16, dev=dev), 5), {} + + +@triton.jit +def k_csr_bound(rowptr_ptr, x_ptr, out_ptr): + pid = tl.program_id(0) + lo = tl.load(rowptr_ptr + pid) + hi = tl.load(rowptr_ptr + pid + 1) + acc = 0.0 + for k in range(lo, hi): + acc += tl.load(x_ptr + k) + tl.store(out_ptr + pid, acc) + + +@case("j_csr_bound", k_csr_bound, "j", "CSR rows: loop bounds loaded from memory") +def _(dev): + rowptr = T(4, dtype=torch.int32, dev=dev, init=[0, 2, 5, 9]) + return (3,), (rowptr, T(9, dev=dev), T(3, dev=dev)), {} + + +@triton.jit +def k_gather(x_ptr, idx_ptr, out_ptr, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + i = tl.load(idx_ptr + offs) + v = tl.load(x_ptr + i) + tl.store(out_ptr + offs, v) + + +@case("j_gather", k_gather, "j", "a gather: the index is loaded from memory") +def _(dev): + idx = T(16, dtype=torch.int32, dev=dev, init=[(i * 5) % 16 for i in range(16)]) + return (1,), (T(16, dev=dev), idx, T(16, dev=dev)), {"BLOCK": 16} + + +# ═══════════════ r: audit regression kernels (refuse OR conform) ═══════════════ +# The shapes #361's reader misread (ir_mode_audit/probes/), from the goldens. + + +@case("r_p1_variant_delta", RK.p1_variant_delta, "r", "p1: p += k") +def _(dev): + return (1,), (T(100, dev=dev), T(1, dev=dev), 4), {} + + +@case("r_p2_swap", RK.p2_swap, "r", "p2: p, q = q, p") +def _(dev): + return (1,), (T(128, dev=dev), T(128, dev=dev), 4), {} + + +@case("r_p3_call_guarded", RK.p3_call_guarded, "r", "p3: noinline call under if") +def _(dev): + return (8,), (T(8, dev=dev), 4), {} + + +@case("r_p3_call_offset", RK.p3_call_offset, "r", "p3: noinline call, actual != formal") +def _(dev): + return (4,), (T(128, dev=dev), 4), {} + + +@case( + "r_p3_call_formals", + RK.p3_call_formals, + "r", + "p3: noinline call, formals match no caller name", +) +def _(dev): + return (4,), (T(128, dev=dev), 4), {} + + +@case( + "r_p4_observed_direct", + RK.p4_observed_direct, + "r", + "p4: an address from an atomic observation", +) +def _(dev): + return (3,), (T(1, dtype=torch.int32, dev=dev, init=None), T(16, dev=dev), 8), {} + + +@case( + "r_p4_observed_loop", + RK.p4_observed_loop, + "r", + "p4: the observation in a loop-carried offset0", +) +def _(dev): + return (3,), (T(1, dtype=torch.int32, dev=dev, init=None), T(16, dev=dev), 8), {} + + +@case( + "r_p4_observed_delta", + RK.p4_observed_delta, + "r", + "p4: the observation in the loop delta", +) +def _(dev): + return (3,), (T(1, dtype=torch.int32, dev=dev, init=None), T(64, dev=dev), 8), {} + + +@case( + "r_trunci_alias_pid0", + RK.rv_trunci_alias, + "r", + "trunci of pid * 2**32 with pid 0 only: fits", +) +def _(dev): + return (1,), (T(4, dtype=torch.int32, dev=dev),), {} + + +@case( + "r_trunci_alias_wrap", + RK.rv_trunci_alias, + "r", + "trunci of pid * 2**32: every program stores x[0]", +) +def _(dev): + return (3,), (T(4, dtype=torch.int32, dev=dev),), {} + + +@case("r_i32_wrap_small", RK.rv_i32_wrap, "r", "(pid * S) * S without a wrap") +def _(dev): + return (4,), (T(64, dtype=torch.int32, dev=dev), 3), {} + + +@case("r_i32_wrap_wrap", RK.rv_i32_wrap, "r", "(pid * 65536) * 65536 wraps to 0 in i32") +def _(dev): + return (2,), (T(4, dtype=torch.int32, dev=dev), 65536), {} + + +@triton.jit +def k_iv_wrap(x_ptr, lo, n, STEP: tl.constexpr): + # the induction variable's increment wraps in i32 when n is near INT32_MAX + for k in range(lo, n, STEP): + tl.store(x_ptr + (k - lo) // STEP, 1.0) + + +@case("r_iv_wrap_small", k_iv_wrap, "r", "the loop increment stays in i32") +def _(dev): + return (1,), (T(16, dev=dev), 0, 10 << 20), {"STEP": 1 << 20} + + +@case("r_iv_wrap_wrap", k_iv_wrap, "r", "upper - 1 + step overflows i32") +def _(dev): + step = 1 << 20 + return (1,), (T(16, dev=dev), 2**31 - 3 * step - 7, 2**31 - 1), {"STEP": step} + + +@case("r_inline_asm_store", RK.rv_inline_asm_store, "r", "an impure asm st.global") +def _(dev): + return (1,), (T(8, dtype=torch.int32, dev=dev),), {"OFF": 4096} + + +@case( + "r_pure_asm_int_addr", + RK.pure_asm_int_addr, + "r", + "a pure asm handed the address as an integer", +) +def _(dev): + return (1,), (T(8, dtype=torch.int32, dev=dev),), {} + + +@case( + "r_loop_observed_advance", + RK.loop_observed_advance, + "r", + "the advance is an atomic observed in the loop", +) +def _(dev): + return (1,), (T(1, dtype=torch.int32, dev=dev, init=None), T(64, dev=dev), 4), {} + + +@case( + "r_int_iterarg_offset", + RK.int_iterarg_offset, + "r", + "an integer offset carried by the loop", +) +def _(dev): + return (1,), (T(64, dev=dev), 3), {"B": 8} + + +@case( + "r_observed_lanes", + RK.observed_lanes, + "r", + "two lanes of one tensor atomic's old values", +) +def _(dev): + return (1,), (T(4, dtype=torch.int32, dev=dev, init=None), T(64, dev=dev)), {"N": 4} diff --git a/tests/conformance/_interp_footprint.py b/tests/conformance/_interp_footprint.py new file mode 100644 index 000000000..4ec4b6407 --- /dev/null +++ b/tests/conformance/_interp_footprint.py @@ -0,0 +1,334 @@ +"""The footprint Triton's interpreter actually touches, per access site. + + python _interp_footprint.py OUT.json CASE [CASE ...] + +Runs each corpus case in a subprocess under ``TRITON_INTERPRET=1`` on CPU +tensors (``CUDA_VISIBLE_DEVICES=""``): every program and loop iteration +executes with the interpreter's fixed-width numpy integers. The +interpreter builder's loads, stores and atomics are instrumented (the +approach of ir_mode_audit/probes_phase3/soundness/oracle.py): each active +lane's address is attributed to the tensor argument whose storage holds it +and recorded as a ``(pid_0, pid_1, pid_2, element offset)`` point, the +offset relative to that argument's ``data_ptr()`` and the program the +builder's ``grid_idx``, keyed ``(argument, kind, source line)`` like the +static side, the line being the innermost frame in a corpus source file. +Each site also records the element width (bits) of the pointers it +accessed. + +Lanes outside every argument's storage (wild) are masked off, or for a +CAS redirected to scratch, so the interpreter never touches unmapped +memory; they are counted and make the case an error. Nothing here imports +tilelens: the interpreter is the independent oracle. + +Not a test module: pytest imports it (python_files = *.py) and finds nothing. +""" + +from __future__ import annotations + +import json +import os +import signal +import subprocess +import sys +import tempfile +import time +import traceback +from types import FrameType +from typing import Any, Sequence + +HERE = os.path.dirname(os.path.abspath(__file__)) +CASE_TIMEOUT_S = int(os.environ.get("TILELENS_CONFORMANCE_CASE_TIMEOUT", "300")) + + +def run_interpreter( + names: Sequence[str], *, timeout: float = 3600.0 +) -> dict[str, dict[str, Any]]: + """case name -> {"sites": [{"arg", "kind", "line", "file", "bits", "points"}], + "errors": [...], "wild": n}; a point is [pid_0, pid_1, pid_2, element offset].""" + env = dict(os.environ, TRITON_INTERPRET="1", CUDA_VISIBLE_DEVICES="") + with tempfile.TemporaryDirectory() as tmp: + out = os.path.join(tmp, "interp.json") + proc = subprocess.run( + [sys.executable, os.path.abspath(__file__), out, *names], + env=env, + capture_output=True, + text=True, + timeout=timeout, + ) + if proc.returncode != 0 or not os.path.exists(out): + raise RuntimeError( + f"interpreter child failed ({proc.returncode}):\n{proc.stderr[-4000:]}" + ) + with open(out, encoding="utf-8") as f: + return json.load(f) + + +# ─────────────────────────── the child ─────────────────────────── + + +class _State: + def __init__(self) -> None: + self.files: frozenset[str] = frozenset() + self.reset() + + def reset(self) -> None: + # (arg, kind, line, file) -> (pid_0, pid_1, pid_2, element offset) points + self.sites: dict[tuple[str, str, int, str], set[tuple[int, int, int, int]]] = {} + # (arg, kind, line, file) -> element widths (bits) of the accessed pointers + self.bits: dict[tuple[str, str, int, str], set[int]] = {} + self.errors: list[str] = [] + self.wild = 0 + # (name, data_ptr, element size, storage lo, storage hi) per tensor argument + self.regions: list[tuple[str, int, int, int, int]] = [] + + def error(self, msg: str) -> None: + if msg not in self.errors: + self.errors.append(msg) + + +STATE = _State() +_REALPATH: dict[str, str] = {} +_POSITIONS: dict[Any, list] = {} # code object -> its co_positions() + + +def _site_line() -> tuple[int, str] | None: + """(line, file) of the innermost frame in a corpus source file. The + call it is in must sit on one line: Triton locates a wrapped call's op + at another of its lines than the frame reports.""" + f: FrameType | None = sys._getframe(2) + while f is not None: + code = f.f_code + path = _REALPATH.get(code.co_filename) + if path is None: + path = _REALPATH[code.co_filename] = os.path.realpath(code.co_filename) + if path in STATE.files: + positions = _POSITIONS.get(code) + if positions is None: + positions = _POSITIONS[code] = list(code.co_positions()) + start, end = positions[f.f_lasti // 2][:2] + if start != end: + STATE.error( + f"a memory-op call spans lines {start}-{end} of {path}: " + "keep it on one line" + ) + return f.f_lineno, path + f = f.f_back + return None + + +def _bind(bound: dict[str, Any]) -> None: + import torch + + STATE.regions = [] + seen: dict[int, str] = {} + for name, v in bound.items(): + if not isinstance(v, torch.Tensor): + continue + storage = v.untyped_storage() + lo = storage.data_ptr() + if lo in seen: + STATE.error( + f"arguments {seen[lo]!r} and {name!r} share a storage: attribution is ambiguous" + ) + seen[lo] = name + STATE.regions.append( + (name, v.data_ptr(), v.element_size(), lo, lo + storage.nbytes()) + ) + + +def _record(kind: str, ptrs, mask, grid_idx): + """Record the active lanes of one memory op run by program ``grid_idx``; + return the lane mask with wild lanes off.""" + import numpy as np + + p = np.asarray(ptrs.data).astype(np.uint64) + if mask is None: + m = np.ones(p.shape, dtype=bool) + else: + # a TensorHandle, or a bare array (materialized block pointers) + raw = mask if isinstance(mask, np.ndarray) else mask.data + m = np.broadcast_to(np.asarray(raw).astype(bool), p.shape).copy() + bits = int(ptrs.get_element_ty().primitive_bitwidth) + width = max(1, bits // 8) + flat_p, flat_m = p.ravel(), m.ravel() + owner = np.full(flat_p.shape, -1, dtype=np.int64) + for i, (_, _, _, lo, hi) in enumerate(STATE.regions): + owner[ + (flat_p >= np.uint64(lo)) & (flat_p + np.uint64(width) <= np.uint64(hi)) + ] = i + site = _site_line() + if site is None and flat_m.any(): + STATE.error(f"{kind}: no frame in a corpus file") + wild = flat_m & (owner < 0) + if wild.any(): + STATE.wild += int(wild.sum()) + STATE.error( + f"{kind} at line {site[0] if site else '?'}: {int(wild.sum())} lanes outside every argument" + ) + if site is not None: + line, path = site + pid = tuple(int(g) for g in grid_idx) + for i, (name, base, elem, _, _) in enumerate(STATE.regions): + sel = flat_m & (owner == i) + if not sel.any(): + continue + rel = flat_p[sel].astype(np.int64) - np.int64(base) + if (rel % elem).any(): + STATE.error( + f"{kind} at line {line}: an address misaligned to {name!r}'s elements" + ) + key = (name, kind, line, path) + STATE.sites.setdefault(key, set()).update( + (*pid, off) for off in (rel // elem).tolist() + ) + STATE.bits.setdefault(key, set()).add(bits) + return (flat_m & (owner >= 0)).reshape(p.shape) + + +def _install() -> None: + import numpy as np + from triton.runtime import interpreter as I + + B, TH = I.InterpreterBuilder, I.TensorHandle + scratch = np.zeros(64, dtype=np.uint64) + + orig_init = I.GridExecutor._init_args_hst + + def _init_args_hst(self, args_dev, kwargs): + import inspect + + args_hst, kwargs_hst = orig_init(self, args_dev, kwargs) + _bind(inspect.getcallargs(self.fn, *args_hst, **kwargs_hst)) + return args_hst, kwargs_hst + + I.GridExecutor._init_args_hst = _init_args_hst + + # NumPy >= 2.4 refuses int() of a 1-element 1-D array, which the + # interpreter's tensor.__index__ does for every loop bound. + orig_patch_tensor = I._patch_lang_tensor + + def _patch_lang_tensor(tensor, scope): + orig_patch_tensor(tensor, scope) + scope.set_attr( + tensor, + "__index__", + lambda self: int(np.asarray(self.handle.data).reshape(-1)[0]), + ) + + I._patch_lang_tensor = _patch_lang_tensor + + def _mask(m, mask): + if isinstance(mask, np.ndarray): + return m + return TH(m, mask.dtype) if mask is not None else TH(m, I.tl.int1) + + orig_load = B.create_masked_load + + def create_masked_load(self, ptrs, mask, *rest, **kw): + return orig_load( + self, + ptrs, + _mask(_record("load", ptrs, mask, self.grid_idx), mask), + *rest, + **kw, + ) + + orig_store = B.create_masked_store + + def create_masked_store(self, ptrs, value, mask, *rest, **kw): + return orig_store( + self, + ptrs, + value, + _mask(_record("store", ptrs, mask, self.grid_idx), mask), + *rest, + **kw, + ) + + orig_rmw = B.create_atomic_rmw + + def create_atomic_rmw(self, rmw_op, ptr, val, mask, *rest, **kw): + return orig_rmw( + self, + rmw_op, + ptr, + val, + _mask(_record("atomic_rmw", ptr, mask, self.grid_idx), mask), + *rest, + **kw, + ) + + orig_cas = B.create_atomic_cas + + def create_atomic_cas(self, ptr, cmp, val, *rest, **kw): + m = _record("atomic_cas", ptr, None, self.grid_idx) + if not m.all(): + data = np.asarray(ptr.data).astype(np.uint64).copy() + data[~m] = np.uint64(scratch.ctypes.data) + ptr = TH(data, ptr.dtype) + return orig_cas(self, ptr, cmp, val, *rest, **kw) + + B.create_masked_load = create_masked_load + B.create_masked_store = create_masked_store + B.create_atomic_rmw = create_atomic_rmw + B.create_atomic_cas = create_atomic_cas + # Block-pointer and descriptor loads / stores materialize their pointers + # and go through create_masked_load / create_masked_store (Triton 3.6). + + +def _alarm(signum, frame): + raise TimeoutError(f"case exceeded {CASE_TIMEOUT_S}s") + + +def main(argv: list[str]) -> int: + out_path, names = argv[0], argv[1:] + assert os.environ.get("TRITON_INTERPRET") == "1", "run through run_interpreter()" + sys.path.insert(0, HERE) + from _ttir_capture import load_corpus # noqa: E402 - the child's own import path + + _install() + corpus = load_corpus() + STATE.files = frozenset( + {os.path.realpath(corpus.__file__), os.path.realpath(corpus.READER_KERNELS)} + ) + cases = corpus.by_name() + signal.signal(signal.SIGALRM, _alarm) + results: dict[str, dict[str, Any]] = {} + for name in names: + STATE.reset() + t0 = time.monotonic() + signal.alarm(CASE_TIMEOUT_S) + try: + c = cases[name] + grid, args, kwargs = c.build("cpu") + c.kernel[grid](*args, **kwargs) + except BaseException as e: # noqa: BLE001 - reported per case + if isinstance(e, KeyboardInterrupt): + raise + STATE.error(f"{type(e).__name__}: {str(e)[:1000]}") + STATE.errors.append(traceback.format_exc()[-2000:]) + finally: + signal.alarm(0) + results[name] = { + "sites": [ + { + "arg": a, + "kind": k, + "line": line, + "file": path, + "bits": sorted(STATE.bits[(a, k, line, path)]), + "points": sorted(points), + } + for (a, k, line, path), points in STATE.sites.items() + ], + "errors": list(STATE.errors), + "wild": STATE.wild, + "seconds": round(time.monotonic() - t0, 3), + } + with open(out_path, "w", encoding="utf-8") as f: + json.dump(results, f) + return 0 + + +if __name__ == "__main__": + sys.exit(main(sys.argv[1:])) diff --git a/tests/conformance/_static_footprint.py b/tests/conformance/_static_footprint.py new file mode 100644 index 000000000..3996438d7 --- /dev/null +++ b/tests/conformance/_static_footprint.py @@ -0,0 +1,603 @@ +"""A concrete evaluator of the TTIR reader's ``AccessGraph`` (D10b; D17 note). + +The static side of the conformance suite: for one launch (grid + scalar +arguments) it enumerates every program id x arange lane x loop iteration +and returns, per access site ``(base_param, kind, source line)``, the set +of ``(pid_0, pid_1, pid_2, element offset)`` points (the offset relative +to that base argument) of the lanes that execute: a footprint per program +instance, as #361's ``differential.static_footprints`` compared it and the +race detector needs it (D17), never merged across programs. Each site also +records the ``AccessEvent.elem_bits`` of its accesses. Ported from #361's +evaluator and extended to this reader's graph: + +* every ``graph.iter_args`` entry, the per-axis expanded tiles included + (an ``IterArgOffset`` is ``offset0 + k * delta`` at iteration ``k``); +* ``IntCast`` and the signed / unsigned ``Bin`` and ``Cmp`` spellings; +* lanes as TTIR broadcasting defines them: every tensor one access + combines has the access's shape, so all aranges along one dim with one + extent index the SAME position there (``tl.arange(0, 16) + + tl.arange(16, 32)`` is ``2i + 16``, not ``i + j + 16`` as #361's + per-``(ssa, dim)`` meshgrid had it), and an extent-1 arange broadcast + along a longer dim stays at position 0; +* terms are evaluated with UNBOUNDED integers (int64 while every operand + bound stays below 2**62, Python ints past that), and the reader's + ``width_obligations`` are checked concretely under their role's + discharge discipline (loop bounds unconditionally, the loop increment + where the loop runs, path where the loop iteration runs, mask under the + path, offset under path and mask). The unbounded reading is the IR's + fixed-width arithmetic only while the obligations hold, so an access + with a failing obligation is reported and never compared. + +Accesses without an exact concrete footprint are excluded, never compared: +``mask_dropped`` / ``guarded`` accesses (skipped: the model deliberately +over-approximates them), accesses whose terms reach an atomic observation +(``Observed``: interleaving-dependent), a ``DataDep`` (no value), a +failing width obligation, and a division by zero or non-positive loop +step (undefined in the IR). An excluded access excludes its whole site +key, so a site is compared only when every access on it is. + +The reader's graph-aware ``mentions_observed`` must agree with this +module's own walk (a disagreement is a :class:`ContractError`). This module +deliberately shares no code with the compiled sanitizer +(``tilelens/clients/sanitizer/compiled/oob.py``): the suite checks the +reader, and the sanitizer is one of its consumers. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Iterable, Iterator, Mapping + +import numpy as np + +from tilelens.ir.ttir_reader import ( + AccessEvent, + AccessGraph, + Arange, + Bin, + BoolBin, + Cmp, + Const, + DataDep, + IntCast, + IterArgOffset, + LoopVar, + Not, + NumPrograms, + Observed, + Param, + Pid, + Select, + mentions_observed, + width_obligations, +) + +SiteKey = tuple[str, str, int] # (base_param, kind, source line) +Point = tuple[int, int, int, int] # (pid_0, pid_1, pid_2, element offset) + +# Why an access is not compared. +MASK_DROPPED = "mask-dropped" +GUARDED = "guarded" +OBSERVED = "observed" +DATA_DEPENDENT = "data-dependent" +OBLIGATION = "obligation" +UNDEFINED = "undefined" +SKIPPED = frozenset({MASK_DROPPED, GUARDED}) # over-approximated by design + +_LIMIT = 1 << 62 # int64 evaluation stays exact while bounds stay below this +_MAX_POINTS = 1 << 24 # (pid, iteration, lane) points one access may enumerate + + +class ContractError(AssertionError): + """The graph breaks a contract the reader documents.""" + + +class EvaluationError(Exception): + """The launch cannot be evaluated (a missing argument, a term outside + the vocabulary, a space too large to enumerate).""" + + +@dataclass(frozen=True) +class Exclusion: + access: int # index into graph.accesses + key: SiteKey + reason: str + detail: str + + +@dataclass(frozen=True) +class ObligationFailure: + access: int + key: SiteKey + role: str + bits: int + signed: bool + ttir_line: int | None + source_line: int | None + points: int # failing (pid, iteration, lane) points + example: int # one failing value + + +@dataclass +class StaticFootprint: + # compared sites: key -> (program, element offset) of the lanes that execute + sites: dict[SiteKey, set[Point]] = field(default_factory=dict) + # excluded sites: key -> why (every excluded access on it) + excluded: dict[SiteKey, list[Exclusion]] = field(default_factory=dict) + obligation_failures: list[ObligationFailure] = field(default_factory=list) + # every site's source file (the key holds its line) + files: dict[SiteKey, set[str]] = field(default_factory=dict) + # every site's element widths (AccessEvent.elem_bits of its accesses) + bits: dict[SiteKey, set[int]] = field(default_factory=dict) + + +def site_key(access: AccessEvent) -> SiteKey: + if access.loc is None: + raise EvaluationError( + f"access at TTIR line {access.line_no} has no source location" + ) + return (access.base_param, access.kind, access.loc.line) + + +# ─────────────────────────── graph walk ─────────────────────────── + + +def _kids(t: object, graph: AccessGraph) -> tuple: + """The terms ``t``'s value is computed from: operands, a loop-carried + pointer's ``offset0`` / ``delta``, the loop's lower bound and step for + the induction variable, a DataDep's modelable ``keep``.""" + if isinstance(t, (Bin, Cmp, BoolBin)): + return (t.a, t.b) + if isinstance(t, Select): + return (t.cond, t.t, t.f) + if isinstance(t, Not): + return (t.a,) + if isinstance(t, IntCast): + return (t.x,) + if isinstance(t, IterArgOffset): + info = graph.iter_args[t.arg_id] + if info.arg_id != t.arg_id: + raise ContractError(f"iter_args[{t.arg_id}].arg_id is {info.arg_id}") + return (info.offset0, info.delta) + if isinstance(t, LoopVar): + if graph.loop is None: + raise ContractError("an induction variable without a loop") + return (graph.loop.lower, graph.loop.step) + if isinstance(t, DataDep): + return () if t.keep is None else (t.keep,) + if isinstance(t, (Const, Pid, NumPrograms, Arange, Param, Observed)): + return () + raise EvaluationError(f"term outside the vocabulary: {type(t).__name__}") + + +def _walk(roots: Iterable[object], graph: AccessGraph) -> Iterator[object]: + """Every node reachable from ``roots``, each once by identity (terms can + be deeper than the recursion limit).""" + seen: set[int] = set() + stack = [r for r in roots if r is not None] + while stack: + t = stack.pop() + if id(t) in seen: + continue + seen.add(id(t)) + yield t + stack.extend(_kids(t, graph)) + + +# ─────────────────────────── unbounded integer arrays ─────────────────────────── + + +def _leaf(v: int) -> np.ndarray: + return np.asarray(v, dtype=np.int64 if -_LIMIT < v < _LIMIT else object) + + +def _int(a: np.ndarray) -> np.ndarray: + return a.astype(np.int64) if a.dtype == np.bool_ else a + + +def _truth(a: np.ndarray) -> np.ndarray: + return a if a.dtype == np.bool_ else np.asarray(a != 0, dtype=np.bool_) + + +def _bound(a: np.ndarray) -> int: + """max |a| as a Python int.""" + if a.size == 0: + return 0 + if a.dtype == np.bool_: + return 1 + return max(abs(int(a.max())), abs(int(a.min()))) + + +def _obj(a: np.ndarray) -> np.ndarray: + return a if a.dtype == object else a.astype(object) + + +def _compact(a: object) -> np.ndarray: + """``a`` as an array (arithmetic on 0-d arrays yields scalars), int64 + again once its values fit.""" + arr = np.asarray(a) + if arr.dtype == object and _bound(arr) < _LIMIT: + return arr.astype(np.int64) + return arr + + +def _wide(a: np.ndarray, b: np.ndarray, bound: int) -> tuple[np.ndarray, np.ndarray]: + return (_obj(a), _obj(b)) if bound >= _LIMIT else (a, b) + + +def _tdiv(a: np.ndarray, b: np.ndarray) -> np.ndarray: + """Quotient truncated toward zero (arith.divsi; numpy's // floors).""" + q = np.abs(a) // np.abs(b) + return np.where((a < 0) != (b < 0), -q, q) + + +def _unsigned(a: np.ndarray, bits: int) -> np.ndarray: + return _compact(_obj(a) % (1 << bits)) + + +def _signed(a: np.ndarray, bits: int) -> np.ndarray: + a = _obj(a) + half = 1 << (bits - 1) + return _compact((a + half) % (1 << bits) - half) + + +# ─────────────────────────── one access ─────────────────────────── + + +class _Access: + """The evaluation space of one access: axes (pid_0, pid_1, pid_2, + iteration, lane positions ...), every array broadcasting over them.""" + + def __init__( + self, + graph: AccessGraph, + access: AccessEvent, + params: Mapping[str, int], + grid: tuple[int, int, int], + lanes: list[tuple[int, int]], + ) -> None: + self.graph = graph + self.access = access + self.params = params + self.grid = grid + self.lanes = {key: 4 + i for i, key in enumerate(lanes)} + self.ndim = 4 + len(lanes) + self.shape = [*grid, 1, *(extent for _, extent in lanes)] + self.iteration: np.ndarray | None = None + # id(term) -> (term, value); holding the term keeps its id unique + self.memo: dict[int, tuple[object, np.ndarray]] = {} + # divisors: (node, zero mask) for every division evaluated + self.divisions: list[tuple[object, np.ndarray]] = [] + + def axis(self, axis: int, n: int) -> np.ndarray: + shape = [1] * self.ndim + shape[axis] = n + return np.arange(n, dtype=np.int64).reshape(shape) + + def set_iterations(self, trips: int) -> None: + self.shape[3] = trips + self.iteration = self.axis(3, trips) + + def param(self, name: str) -> np.ndarray: + if name not in self.params: + raise EvaluationError(f"scalar argument {name!r} has no launch value") + arg = self.graph.arg(name) + bits = arg.int_bits if arg is not None else 0 + v = int(self.params[name]) + if bits == 1: + return np.asarray(v != 0) # i1 terms are booleans + if bits > 1: + half = 1 << (bits - 1) + v = (v + half) % (1 << bits) - half # the IR's signed reading + return _leaf(v) + + def value(self, root: object) -> np.ndarray: + stack: list[tuple[object, bool]] = [(root, False)] + memo = self.memo + while stack: + t, ready = stack.pop() + if id(t) in memo: + continue + kids = self.operands(t) + missing = [k for k in kids if id(k) not in memo] + if missing and not ready: + stack.append((t, True)) + stack.extend((k, False) for k in missing) + continue + memo[id(t)] = (t, np.asarray(self.apply(t, [memo[id(k)][1] for k in kids]))) + return memo[id(root)][1] + + def operands(self, t: object) -> tuple: + # an Observed / DataDep leaf has no value: apply() refuses it + return () if isinstance(t, DataDep) else _kids(t, self.graph) + + def apply(self, t: object, v: list[np.ndarray]) -> np.ndarray: + if isinstance(t, Const): + return _leaf(int(t.value)) + if isinstance(t, Pid): + return self.axis(t.axis, self.grid[t.axis]) + if isinstance(t, NumPrograms): + return _leaf(self.grid[t.axis]) + if isinstance(t, Param): + return self.param(t.name) + if isinstance(t, Arange): + key = (t.dim, t.end - t.start) + return self.axis(self.lanes[key], t.end - t.start) + t.start + if isinstance(t, (LoopVar, IterArgOffset)): + if self.iteration is None: + raise ContractError(f"{type(t).__name__} in an access outside the loop") + # offset0 + k * delta, lower + k * step + base, step = _int(v[0]), _int(v[1]) + k, step = _wide(self.iteration, step, _bound(self.iteration) * _bound(step)) + scaled = k * step + base, scaled = _wide(base, scaled, _bound(base) + _bound(scaled)) + return _compact(base + scaled) + if isinstance(t, Bin): + return self.bin(t, _int(v[0]), _int(v[1])) + if isinstance(t, Cmp): + return self.cmp(t, v[0], v[1]) + if isinstance(t, BoolBin): + a, b = _truth(v[0]), _truth(v[1]) + return a & b if t.op == "and" else a | b + if isinstance(t, Select): + return np.where(_truth(v[0]), v[1], v[2]) + if isinstance(t, Not): + return ~_truth(v[0]) + if isinstance(t, IntCast): + # the model's value (exact while the cast's obligation holds) + return _int(v[0]) + if isinstance(t, Observed): + raise ContractError( + f"evaluation reached Observed({t.access_index}), which the walk did not report" + ) + if isinstance(t, DataDep): + raise ContractError( + f"evaluation reached DataDep({t.why!r}), which the walk did not report" + ) + raise EvaluationError(f"term outside the vocabulary: {type(t).__name__}") + + def divisor(self, t: Bin, b: np.ndarray) -> np.ndarray: + zero = np.asarray(b == 0) + if zero.any(): + self.divisions.append((t, zero)) + b = np.where(zero, 1, b) + return b + + def bin(self, t: Bin, a: np.ndarray, b: np.ndarray) -> np.ndarray: + op = t.op + if op in ("+", "-"): + a, b = _wide(a, b, _bound(a) + _bound(b)) + return _compact(a + b if op == "+" else a - b) + if op == "*": + a, b = _wide(a, b, _bound(a) * _bound(b)) + return _compact(a * b) + if op in ("min", "max"): + return np.minimum(a, b) if op == "min" else np.maximum(a, b) + if op in ("//", "%"): + b = self.divisor(t, b) + q = _tdiv(a, b) + return q if op == "//" else _compact(_obj(a) - _obj(b) * _obj(q)) + if op in ("u//", "u%", "umin", "umax"): + if t.bits is None: + raise ContractError(f"unsigned op {op} without a width") + ua, ub = _unsigned(a, t.bits), _unsigned(b, t.bits) + if op == "umin": + r = np.minimum(ua, ub) + elif op == "umax": + r = np.maximum(ua, ub) + else: + ub = self.divisor(t, ub) + r = ua // ub if op == "u//" else ua % ub + return _signed(r, t.bits) + raise EvaluationError(f"unknown integer op {op!r}") + + def cmp(self, t: Cmp, a: np.ndarray, b: np.ndarray) -> np.ndarray: + a, b = _int(a), _int(b) + pred = t.pred + if pred[0] == "u": + if not t.bits: + raise ContractError(f"unsigned predicate {pred} without a width") + a, b = _unsigned(a, t.bits), _unsigned(b, t.bits) + pred = "s" + pred[1:] + if pred == "eq": + return np.asarray(a == b) + if pred == "ne": + return np.asarray(a != b) + ops = { + "slt": np.less, + "sle": np.less_equal, + "sgt": np.greater, + "sge": np.greater_equal, + } + if pred not in ops: + raise EvaluationError(f"unknown predicate {t.pred!r}") + return np.asarray(ops[pred](a, b), dtype=np.bool_) + + def full(self, a: np.ndarray, *, pids_only: bool = False) -> np.ndarray: + """``a`` over the whole space, or over the program ids alone (the + loop's bounds, computed once per program whatever the trip count).""" + shape = self.shape[:3] + [1] * (self.ndim - 3) if pids_only else self.shape + return np.broadcast_to(a, tuple(shape)) + + +def _lane_keys(nodes: Iterable[object]) -> list[tuple[int, int]]: + return sorted({(n.dim, n.end - n.start) for n in nodes if isinstance(n, Arange)}) + + +def _check_obligations( + ev: _Access, index: int, key: SiteKey, bound_ids: set[int], valid, path, mask, trips +) -> list[ObligationFailure]: + """The access's width obligations, each where its role says it matters.""" + out = [] + for ob in width_obligations(ev.graph, ev.access): + pids_only = ob.role == "loop" + if pids_only: + # the bounds: unconditional; the increment (a term of its own): + # where the loop runs at least once + cond = ( + np.asarray(True) if id(ob.term) in bound_ids else np.asarray(trips > 0) + ) + elif ob.role == "path": + cond = valid + elif ob.role == "mask": + cond = valid & path + elif ob.role == "offset": + cond = valid & path & mask + else: + raise ContractError(f"unknown obligation role {ob.role!r}") + v = _int(ev.value(ob.term)) + if ob.signed: + lo, hi = -(1 << (ob.bits - 1)), 1 << (ob.bits - 1) + else: + lo, hi = 0, 1 << ob.bits + fits = np.asarray((v >= lo) & (v < hi), dtype=np.bool_) + bad = ev.full(cond, pids_only=pids_only) & ~ev.full(fits, pids_only=pids_only) + if bad.any(): + example = int(ev.full(v, pids_only=pids_only)[bad].flat[0]) + out.append( + ObligationFailure( + index, + key, + ob.role, + ob.bits, + ob.signed, + ob.line_no, + ob.loc.line if ob.loc is not None else None, + int(bad.sum()), + example, + ) + ) + return out + + +def _evaluate( + graph: AccessGraph, + index: int, + params: Mapping[str, int], + grid: tuple[int, int, int], + out: StaticFootprint, +) -> tuple[set[Point] | None, list[Exclusion]]: + """One access's footprint, or None with the reasons it is excluded.""" + access = graph.accesses[index] + key = site_key(access) + if access.mask_dropped or access.guarded: + why = [MASK_DROPPED] * access.mask_dropped + [GUARDED] * access.guarded + return None, [ + Exclusion(index, key, r, "over-approximated by the model") for r in why + ] + loop = graph.loop + if access.in_loop and loop is None: + raise ContractError(f"access {index} is in_loop but the graph has no loop") + roots = [access.offset, access.mask, access.path] + if access.in_loop: + roots += [loop.lower, loop.upper, loop.step] # type: ignore[union-attr] + nodes = list(_walk(roots, graph)) + observed = any(isinstance(n, Observed) for n in nodes) + reader_says = any(t is not None and mentions_observed(t, graph) for t in roots) + if observed != reader_says: + raise ContractError( + f"access {index} (line {key[2]}): mentions_observed says {reader_says}, the graph walk {observed}" + ) + if observed: + return None, [ + Exclusion( + index, + key, + OBSERVED, + "reads an atomic observation (interleaving-dependent)", + ) + ] + deps = [n.why for n in nodes if isinstance(n, DataDep)] + if deps: + return None, [ + Exclusion(index, key, DATA_DEPENDENT, "; ".join(sorted(set(deps)))) + ] + + ev = _Access(graph, access, params, grid, _lane_keys(nodes)) + trips = np.asarray(1) + bound_ids: set[int] = set() + if access.in_loop: + assert loop is not None + bound_ids = {id(n) for n in _walk((loop.lower, loop.upper, loop.step), graph)} + lower, upper, step = ( + _int(ev.value(b)) for b in (loop.lower, loop.upper, loop.step) + ) + if (np.asarray(step) <= 0).any(): + return None, [Exclusion(index, key, UNDEFINED, "a non-positive loop step")] + trips = np.maximum(_obj(upper) - _obj(lower) + _obj(step) - 1, 0) // _obj(step) + trips = _compact(np.asarray(trips)) + ev.set_iterations(int(np.max(trips)) if trips.size else 0) + valid = np.asarray(ev.iteration < trips) + else: + valid = np.asarray(True) + points = int(np.prod(ev.shape)) + if points > _MAX_POINTS: + raise EvaluationError( + f"access {index} spans {points} points (limit {_MAX_POINTS})" + ) + + path = ( + _truth(ev.value(access.path)) if access.path is not None else np.asarray(True) + ) + mask = ( + _truth(ev.value(access.mask)) if access.mask is not None else np.asarray(True) + ) + offset = _int(ev.value(access.offset)) + failures = _check_obligations(ev, index, key, bound_ids, valid, path, mask, trips) + if failures: + out.obligation_failures += failures + detail = ", ".join( + f"{f.role} i{f.bits} at source line {f.source_line}: e.g. {f.example}" + for f in failures + ) + return None, [Exclusion(index, key, OBLIGATION, detail)] + for node, zero in ev.divisions: + # a divisor of the loop's bounds counts in every program, others + # at the iterations that run + hit = ( + ev.full(zero, pids_only=True).any() + if id(node) in bound_ids + else ev.full(zero & valid).any() + ) + if hit: + return None, [ + Exclusion( + index, + key, + UNDEFINED, + f"a division by zero (TTIR line {getattr(node, 'line_no', None)})", + ) + ] + active = ev.full(valid & path & mask) + # np.nonzero and boolean indexing both walk the space in C order + p0, p1, p2 = (i.tolist() for i in np.nonzero(active)[:3]) + offsets = ev.full(offset)[active].tolist() + return { + (a, b, c, int(o)) for a, b, c, o in zip(p0, p1, p2, offsets, strict=True) + }, [] + + +def static_footprint( + graph: AccessGraph, + params: Mapping[str, int], + grid: tuple[int, ...], +) -> StaticFootprint: + """The model's footprint of one launch: ``params`` maps every scalar + kernel argument to its launch value, ``grid`` is the launch grid.""" + grid3 = tuple(int(g) for g in grid) + (1,) * (3 - len(grid)) + assert len(grid3) == 3 and all(g >= 1 for g in grid3), grid + out = StaticFootprint() + for index in range(len(graph.accesses)): + offsets, why = _evaluate(graph, index, params, grid3, out) # type: ignore[arg-type] + access = graph.accesses[index] + key = site_key(access) + out.files.setdefault(key, set()).add(access.loc.file) # type: ignore[union-attr] + out.bits.setdefault(key, set()).add(access.elem_bits) + if why: + out.excluded.setdefault(key, []).extend(why) + else: + assert offsets is not None + out.sites.setdefault(key, set()).update(offsets) + for key in out.excluded: + out.sites.pop(key, None) + return out diff --git a/tests/conformance/_ttir_capture.py b/tests/conformance/_ttir_capture.py new file mode 100644 index 000000000..a05ea7a76 --- /dev/null +++ b/tests/conformance/_ttir_capture.py @@ -0,0 +1,215 @@ +"""Compile corpus cases to TTIR in a clean subprocess. + + python _ttir_capture.py OUT.json CASE [CASE ...] + +The TTIR is what IR mode reads (D25): host-compiled by +``tilelens.core.host_compile.HostCompiler`` for the default target +``GPUTarget("cuda", 89, 32)`` (D26), through the ``ttir`` stage, with no +device involved; ``source`` is ``"host"``. The suite thereby validates the +host compile path together with the reader on every Triton it runs on, CPU +only included. The JIT's own compile of the same launch, for the same +target, is recorded apart (``jit_ttir``, ``jit_hash``, ``jit_target``; the +host's hash is ``hash``) for the explicit host-vs-JIT check +(test_host_ttir_is_the_jit_ttir): ``jit_fn.warmup(...)`` with a stand-in +for Triton's active driver that reports a cuda:89 device, so it needs no GPU +either. Launches are built on CPU tensors, so a record does not depend on +the machine. Each record also holds the launch the static evaluator needs: +the resolved 3-D grid and every integer argument by name. + +A subprocess, because the test process may have imported Triton under +``TRITON_INTERPRET=1`` (tests/unit/test_multithreading.py sets it during +collection), where nothing compiles for real. The child imports tilelens +from this checkout (``PYTHONPATH``), and compiles into a Triton cache of +this checkout's own (:func:`cache_dir`): Triton keys a compiled kernel by +its source text and first line, not its file, so a cache shared by two +checkouts of the repository hands the second one TTIR whose ``loc()`` +entries name the first checkout's files. + +Not a test module: pytest imports it (python_files = *.py) and finds nothing. +""" + +from __future__ import annotations + +import hashlib +import json +import os +import subprocess +import sys +import tempfile +import time +import traceback +from typing import Any, Sequence + +HERE = os.path.dirname(os.path.abspath(__file__)) +CORPUS = os.path.join(HERE, "_corpus.py") +REPO = os.path.dirname(os.path.dirname(HERE)) + + +def cache_dir() -> str: + """The capture child's Triton cache: a subdirectory, named after this + checkout's path, of the cache Triton would use (``TRITON_CACHE_DIR``, + else ``$TRITON_HOME/.triton/cache``), so warm runs stay warm.""" + root = os.environ.get("TRITON_CACHE_DIR") or os.path.join( + os.environ.get("TRITON_HOME") or os.path.expanduser("~"), ".triton", "cache" + ) + tag = hashlib.sha1(os.path.realpath(HERE).encode()).hexdigest()[:12] + return os.path.join(root, f"tilelens-conformance-{tag}") + + +def capture( + names: Sequence[str], *, timeout: float = 1800.0 +) -> dict[str, dict[str, Any]]: + """case name -> {"ttir", "hash", "source", "jit_ttir", "jit_hash", + "jit_target", "grid", "params", "file"} or {"error"}.""" + env = {k: v for k, v in os.environ.items() if k != "TRITON_INTERPRET"} + env["TRITON_CACHE_DIR"] = cache_dir() + env["PYTHONPATH"] = os.pathsep.join( + [REPO, *filter(None, [os.environ.get("PYTHONPATH")])] + ) + with tempfile.TemporaryDirectory() as tmp: + out = os.path.join(tmp, "ttir.json") + proc = subprocess.run( + [sys.executable, os.path.abspath(__file__), out, *names], + env=env, + capture_output=True, + text=True, + timeout=timeout, + ) + if proc.returncode != 0 or not os.path.exists(out): + raise RuntimeError( + f"TTIR capture child failed ({proc.returncode}):\n{proc.stderr[-4000:]}" + ) + with open(out, encoding="utf-8") as f: + return json.load(f) + + +# ─────────────────────────── the child ─────────────────────────── + + +def load_corpus(): + import importlib.util + + name = "tilelens_conformance_corpus" + spec = importlib.util.spec_from_file_location(name, CORPUS) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +def launch_record(kernel, grid, args, kwargs) -> dict[str, Any]: + """The launch as the static evaluator reads it: integer arguments by + Python parameter name (constexprs included; the reader has them folded).""" + import inspect + + import torch + + fn = kernel.fn + bound = inspect.signature(fn).bind(*args, **kwargs) + bound.apply_defaults() + params: dict[str, int] = {} + tensors: dict[str, Any] = {} + for name, v in bound.arguments.items(): + if isinstance(v, torch.Tensor): + tensors[name] = { + "shape": list(v.shape), + "stride": list(v.stride()), + "dtype": str(v.dtype), + } + elif isinstance(v, (bool, int)): + params[name] = int(v) + grid3 = [int(g) for g in grid] + [1] * (3 - len(grid)) + return { + "grid": grid3, + "params": params, + "tensors": tensors, + "file": os.path.realpath(fn.__code__.co_filename), + } + + +class StandInDriver: + """What JITFunction.run asks Triton's active driver for, as on a machine + whose device 0 is a GPU of ``target``: nothing is launched or loaded on + a warmup, so the JIT compiles without a GPU.""" + + def __init__(self, target) -> None: + self.target = target + + def get_current_device(self) -> int: + return 0 + + def get_current_stream(self, device=None) -> int: + return 0 + + def get_current_target(self): + return self.target + + +def jit_compile(kernel, grid, args, kwargs) -> tuple[str, str, list]: + """The JIT's own compile of the launch for the default target (the + stand-in driver's): its TTIR, hash and target.""" + from triton.runtime.driver import driver + + from tilelens.core.host_compile import default_ir_target + + previous = driver._active + driver.set_active(StandInDriver(default_ir_target())) + try: + compiled = kernel.warmup(*args, grid=grid, **kwargs) + finally: + driver._active = previous + target = compiled.metadata.target + return ( + compiled.asm["ttir"], + compiled.hash, + [target.backend, target.arch, target.warp_size], + ) + + +def host_compile(kernel, args, kwargs) -> tuple[str, str]: + """The TTIR IR mode reads, and its hash: the launch host-compiled for + the default target, as the core compiles it (D25, D26).""" + from tilelens.core.host_compile import HostCompiler, default_ir_target + + compiled = HostCompiler().compile( + kernel, tuple(args), kwargs, target=default_ir_target(), stages={"ttir"} + ) + return compiled.asm["ttir"], compiled.hash + + +def main(argv: list[str]) -> int: + out_path, names = argv[0], argv[1:] + from triton.runtime.jit import JITFunction + + corpus = load_corpus() + cases = corpus.by_name() + results: dict[str, dict[str, Any]] = {} + for name in names: + t0 = time.monotonic() + try: + c = cases[name] + if not isinstance(c.kernel, JITFunction): + raise TypeError( + f"{name}: kernel is {type(c.kernel).__name__}, not a JITFunction" + ) + grid, args, kwargs = c.build("cpu") + rec = launch_record(c.kernel, grid, args, kwargs) + ttir, digest = host_compile(c.kernel, args, kwargs) + rec.update(ttir=ttir, hash=digest, source="host") + ttir, digest, target = jit_compile(c.kernel, grid, args, kwargs) + rec.update(jit_ttir=ttir, jit_hash=digest, jit_target=target) + except Exception as e: # noqa: BLE001 - reported per case + rec = { + "error": f"{type(e).__name__}: {e}", + "traceback": traceback.format_exc()[-3000:], + } + rec["seconds"] = round(time.monotonic() - t0, 3) + results[name] = rec + with open(out_path, "w", encoding="utf-8") as f: + json.dump(results, f) + return 0 + + +if __name__ == "__main__": + sys.exit(main(sys.argv[1:])) diff --git a/tests/conformance/conftest.py b/tests/conformance/conftest.py new file mode 100644 index 000000000..41d1a7f5f --- /dev/null +++ b/tests/conformance/conftest.py @@ -0,0 +1,79 @@ +"""Terminal summary of the reader conformance suite (D10b): the counts a +TESTED_TRITON_VERSIONS decision reads, aggregated from the per-case +``conformance`` user property of test_reader_conformance.py.""" + +from __future__ import annotations + +from collections import Counter + + +def pytest_terminal_summary(terminalreporter) -> None: + cases: list[tuple[str, dict]] = [] + for reports in terminalreporter.stats.values(): + for report in reports: + if getattr(report, "when", None) != "call": + continue + for key, props in getattr(report, "user_properties", ()): + if key == "conformance": + cases.append((report.outcome, props)) + if not cases: + return + import triton + + from tilelens.core.config import untested_triton_version + + outcomes = Counter(p.get("outcome", "error") for _, p in cases) + refused = Counter( + o.split(":", 1)[1] for o in outcomes.elements() if o.startswith("refused:") + ) + accepted = [p for _, p in cases if not p.get("outcome", "").startswith("refused:")] + excluded: Counter = Counter() + for p in accepted: + excluded.update(p.get("excluded", {})) + seconds = max((p["seconds"] for _, p in cases), key=lambda s: sum(s.values())) + w = terminalreporter.write_line + terminalreporter.section("TTIR reader conformance (D10b)") + untested = ( + " (outside TESTED_TRITON_VERSIONS: failures expected, xfail)" + if untested_triton_version() + else "" + ) + w( + f"triton {triton.__version__}{untested}; " + f"TTIR source: {', '.join(sorted({str(p.get('source')) for _, p in cases}))}" + ) + compared = sum(1 for p in accepted if p.get("exercised_sites")) + w( + f"cases {len(cases)}: compared {compared} (conforming {outcomes['conform']}, mismatching " + f"{outcomes['mismatch']}), accepted but not compared {outcomes['not-compared']}, " + f"errors {outcomes['error']}; failed tests {sum(1 for o, _ in cases if o != 'passed')}" + ) + w(f"refused {sum(refused.values())}: {dict(sorted(refused.items()))}") + w( + f"sites compared {sum(p.get('compared_sites', 0) for p in accepted)} " + f"({sum(p.get('points', 0) for p in accepted)} (program, offset) points); " + f"skipped accesses {sum(p.get('skipped_accesses', 0) for p in accepted)}; " + f"excluded sites {dict(sorted(excluded.items()))}" + ) + w( + f"obligation-violating launches {sum(1 for p in accepted if p.get('obligation_failures'))}" + ) + w( + "runtime: " + + ", ".join(f"{k} {v:.1f}s" for k, v in seconds.items()) + + f" (total {sum(seconds.values()):.1f}s)" + ) + compared_jit = [ + (report.outcome, props) + for reports in terminalreporter.stats.values() + for report in reports + if getattr(report, "when", None) == "call" + for key, props in getattr(report, "user_properties", ()) + if key == "host_vs_jit" + ] + if compared_jit: + w( + f"host vs JIT compile (cuda:89, stand-in driver): {len(compared_jit)} " + f"compared, {sum(1 for _, p in compared_jit if p['same_text'])} the same " + "kernel (hash and TTIR text)" + ) diff --git a/tests/conformance/test_reader_conformance.py b/tests/conformance/test_reader_conformance.py new file mode 100644 index 000000000..69fd0fd3c --- /dev/null +++ b/tests/conformance/test_reader_conformance.py @@ -0,0 +1,486 @@ +"""D10b: the TTIR reader's static footprint equals what Triton's interpreter touches. + +For every (kernel, launch) of the corpus (``_corpus.py``): + +1. the launch is compiled to TTIR in a clean subprocess (``_ttir_capture``) + the way IR mode compiles it: on the host, by + ``tilelens.core.host_compile``, for the default target + ``GPUTarget("cuda", 89, 32)`` (D25, D26), CPU only or not; +2. :func:`tilelens.ir.ttir_reader.parse_ttir` reads it, and the concrete + evaluator (``_static_footprint``) enumerates every program id x arange + lane x loop iteration of the AccessGraph with unbounded integers, the + reader's width obligations checked concretely; +3. the same launch runs under ``TRITON_INTERPRET=1`` in another subprocess + (``_interp_footprint``), its loads / stores / atomics instrumented; +4. per access site ``(base argument, kind, source line)`` the + ``(pid_0, pid_1, pid_2, element offset)`` points must be EQUAL: each + program's own footprint, no subset slack, either way. The site's + ``AccessEvent.elem_bits`` must equal the element width of the pointers + the interpreter accessed there (the byte model is offset x width), and + ``graph.pid_axes`` the axes of the TTIR's ``tt.get_program_id`` / + ``tt.get_num_programs`` ops. + +Accesses the model over-approximates by design (``mask_dropped``, +``guarded``) are skipped, sites that read an atomic observation or whose +launch breaks a width obligation are excluded; each is reported and pinned +by the ``EXPECTED`` table below, as are the reader's refusals: a refusal, an +exclusion or a case that exercises no access, anywhere the table does not +list it, fails, and so does a reader crash (on its own case). For the audit +regression kernels (``r_*``) and the representational limits (``j_*``) the +requirement is "the reader refuses with the listed kind OR the footprints +conform", so reverting a reader fix makes this suite fail; a case the +interpreter cannot run (inline asm) must refuse. + +The capture child also compiles each launch with the JIT itself, for the +same target (a stand-in driver reports a cuda:89 device, so no GPU is +needed), and test_host_ttir_is_the_jit_ttir checks explicitly that the +host compile is the JIT's: the same hash and the same TTIR text, so the +host compile's binder and stage cut cannot change what the reader sees. + +Version policy: this suite runs on the INSTALLED Triton. A Triton minor +release may be added to ``tilelens.core.config.TESTED_TRITON_VERSIONS`` +(the IR-mode gate, D10b) only when this suite passes on it (the +host-vs-JIT check included; neither needs a GPU), run with +``TILELENS_IR_ALLOW_UNTESTED_TRITON=1``. Without that override, on +a release outside the table the cases still run but are expected to fail +(``xfail``, not strict): a CI job that installs the newest Triton stays +green and still shows how far it conforms. The capture uses tilelens's own +host compile (the private Triton API it leans on is feature-checked there), +so this suite validates that path on each release too; the oracle leans on +the interpreter builder's memory methods and ``grid_idx``; adapting it to a +new release is part of adding it. +""" + +from __future__ import annotations + +import hashlib +import json +import os +import re +import time +import traceback +from collections import Counter +from dataclasses import dataclass, field +from typing import Any, Mapping + +import pytest + +from tilelens.core.config import TESTED_TRITON_VERSIONS, untested_triton_version +from tilelens.ir import _mlir_walk +from tilelens.ir import ttir_reader +from tilelens.ir.ttir_reader import TTIRKind, UnsupportedTTIR, parse_ttir + +from . import _corpus +from . import _interp_footprint as interp_side +from . import _static_footprint as S +from . import _ttir_capture as capture_side + +NAMES = [c.name for c in _corpus.CASES] + + +@dataclass(frozen=True) +class Expect: + # The reader may refuse with this kind; otherwise the case must conform. + refusal: TTIRKind | None = None + # The reader MUST refuse (with ``refusal``): the interpreter cannot run + # the kernel, so accepting it can never conform. + must_refuse: bool = False + # Accepted: exclusion reason (joined by "+" when one site has several) -> + # number of excluded sites. + excluded: Mapping[str, int] = field(default_factory=dict) + # Accepted: at least one compared site with a non-empty footprint. False + # only where every site is excluded by design. + compared: bool = True + why: str = "" + + +_OBSERVED_2 = {S.OBSERVED: 2} +EXPECTED: dict[str, Expect] = { + # ── audit regressions (ir_mode_audit/probes/): refuse with this kind OR conform ── + "r_p1_variant_delta": Expect(TTIRKind.LOOP_VARIANT_ADVANCE, why="p1: p += k"), + "r_p2_swap": Expect(TTIRKind.LOOP_VARIANT_ADVANCE, why="p2: p, q = q, p"), + "r_p3_call_guarded": Expect(TTIRKind.CALL, why="p3: noinline call"), + "r_p3_call_offset": Expect(TTIRKind.CALL, why="p3: noinline call"), + "r_p3_call_formals": Expect(TTIRKind.CALL, why="p3: noinline call"), + # Triton's interpreter cannot run inline asm: accepting these is a + # reader regression, never a conformance question + "r_inline_asm_store": Expect( + TTIRKind.INLINE_ASM, must_refuse=True, why="an impure asm st.global" + ), + "r_pure_asm_int_addr": Expect( + TTIRKind.INLINE_ASM, must_refuse=True, why="a pure asm handed an address" + ), + "r_loop_observed_advance": Expect( + TTIRKind.LOOP_VARIANT_ADVANCE, why="advance by an in-loop observation" + ), + "r_int_iterarg_offset": Expect( + TTIRKind.LOOP_VARIANT_ADVANCE, why="an integer offset carried by the loop" + ), + "r_observed_lanes": Expect( + TTIRKind.INDIRECT_ADDRESS, why="two lanes of one tensor observation" + ), + # p4: the reader represents observations by design; its graph-aware + # mentions_observed must flag the loads / stores (the evaluator checks it + # against its own walk), the atomic itself is compared. + "r_p4_observed_direct": Expect(excluded=_OBSERVED_2), + "r_p4_observed_loop": Expect(excluded=_OBSERVED_2), + "r_p4_observed_delta": Expect(excluded=_OBSERVED_2), + # D9: a launch where a width obligation fails is reported, not compared + # (the *_small / *_pid0 launches of the same kernels are compared). + "r_trunci_alias_wrap": Expect( + excluded={S.OBLIGATION: 1}, compared=False, why="trunci(pid * 2**32)" + ), + "r_i32_wrap_wrap": Expect( + excluded={S.OBLIGATION: 1}, compared=False, why="(pid * S) * S wraps" + ), + "r_iv_wrap_wrap": Expect( + excluded={S.OBLIGATION: 1}, compared=False, why="the loop increment wraps" + ), + "b_unsigned_cmp_highbit": Expect( + excluded={S.OBLIGATION: 1}, + compared=False, + why="cmpi ult reads n = -1 as 2**32 - 1: the unsigned-operand obligation fails", + ), + # ── representational limits of the reader: refuse with this kind OR conform ── + "a_copy_bool": Expect( + TTIRKind.OTHER, + why="a *i1 argument is accessed through a tt.bitcast to *i8: the reader reads 1 vs 8 bits as a " + "width change, though both are one byte in memory", + ), + "c_block_ptr_loop": Expect( + TTIRKind.LOOP_VARIANT_ADVANCE, + why="an advanced block pointer's offsets are integer iter_args (3.6: rewrite_tensor_pointer; " + "3.8: the frontend lowers tl.make_block_ptr to pointer arithmetic)", + ), + "j_two_loops": Expect(TTIRKind.NESTED_LOOP, why="two sequential scf.for"), + "j_nested_loops": Expect(TTIRKind.NESTED_LOOP, why="an scf.for in an scf.for"), + "j_loop_under_if": Expect(TTIRKind.CONTROL_FLOW, why="an scf.for under an scf.if"), + "j_while": Expect(TTIRKind.CONTROL_FLOW, why="an scf.while"), + "j_csr_bound": Expect( + TTIRKind.DATA_DEPENDENT_BOUND, why="loop bounds loaded from memory" + ), + "j_gather": Expect(TTIRKind.INDIRECT_ADDRESS, why="an index loaded from memory"), + # ── accesses the model over-approximates or cannot evaluate concretely ── + "i_datadep_mask": Expect(excluded={S.MASK_DROPPED: 1}), + "i_guarded_branch": Expect(excluded={S.GUARDED: 1}), + "i_atomic_max_float": Expect( + excluded={S.MASK_DROPPED: 1}, why="two atomics masked by the value's sign" + ), + "i_observed_mask_path": Expect(excluded=_OBSERVED_2), +} + + +def test_expectation_table_names_corpus_cases(): + assert not set(EXPECTED) - set(NAMES) + assert len(NAMES) == len(set(NAMES)) + assert all(e.refusal is not None for e in EXPECTED.values() if e.must_refuse) + + +# ─────────────────────────── running the corpus ─────────────────────────── + + +@dataclass +class Outcome: + record: dict[str, Any] # the capture child's record + refusal: UnsupportedTTIR | None = None + reader_error: str | None = None # parse_ttir raised something else + pid_axes: frozenset[int] = frozenset() + static: S.StaticFootprint | None = None + static_error: str | None = None + interp: dict[str, Any] | None = None + + +@dataclass +class Run: + outcomes: dict[str, Outcome] + seconds: dict[str, float] + # the children's results, as they came (shared between xdist workers) + records: dict[str, Any] + interp: dict[str, Any] + + +def _selected(request) -> list[str]: + names = [] + for item in request.session.items: + callspec = getattr(item, "callspec", None) + if ( + getattr(item, "module", None) is request.module + and callspec is not None + and "name" in callspec.params + ): + names.append(callspec.params["name"]) + return list(dict.fromkeys(names)) or NAMES + + +def _run( + names: list[str], + records: dict[str, Any] | None = None, + interp: dict[str, Any] | None = None, +) -> Run: + """Capture (unless ``records`` are given), read, evaluate and interpret + (the accepted cases ``interp`` lacks) every case in ``names``.""" + t0 = time.monotonic() + if records is None: + records = capture_side.capture(names) + t1 = time.monotonic() + outcomes: dict[str, Outcome] = {} + for name in names: + out = outcomes[name] = Outcome(records[name]) + if out.record.get("error"): + continue + try: + graph = parse_ttir(out.record["ttir"]) + except UnsupportedTTIR as e: + out.refusal = e + continue + except Exception as e: # noqa: BLE001 - reported by the case's test + out.reader_error = ( + f"{type(e).__name__}: {e}\n{traceback.format_exc()[-3000:]}" + ) + continue + out.pid_axes = graph.pid_axes + try: + out.static = S.static_footprint( + graph, out.record["params"], out.record["grid"] + ) + except Exception as e: # noqa: BLE001 - reported by the case's test + out.static_error = f"{type(e).__name__}: {e}" + t2 = time.monotonic() + accepted = [ + n + for n, o in outcomes.items() + if o.static is not None and not EXPECTED.get(n, Expect()).must_refuse + ] + interp = dict(interp or {}) + missing = [n for n in accepted if n not in interp] + if missing: + interp.update(interp_side.run_interpreter(missing)) + for name in accepted: + outcomes[name].interp = interp[name] + t3 = time.monotonic() + seconds = {"capture": t1 - t0, "static": t2 - t1, "interpreter": t3 - t2} + return Run(outcomes, seconds, records, interp) + + +def _shared_result(tmp_path_factory, names: list[str]) -> str | None: + """Where pytest-xdist workers share the run: each worker collects the + whole module, so without it every worker that runs any case would + capture, read and interpret the whole corpus. None outside xdist.""" + if os.environ.get("PYTEST_XDIST_WORKER") is None: + return None + digest = hashlib.sha1("\0".join(names).encode()).hexdigest()[:12] + # the base temp directory's parent is common to one run's workers + root = tmp_path_factory.getbasetemp().parent + return str(root / f"tilelens-conformance-{digest}.json") + + +@pytest.fixture(scope="module") +def conformance(request, tmp_path_factory) -> Run: + names = _selected(request) + shared = _shared_result(tmp_path_factory, names) + if shared is None: + return _run(names) + try: + from filelock import FileLock # a torch dependency + except ImportError: + return _run(names) + with FileLock(shared + ".lock"): + if os.path.exists(shared): + with open(shared, encoding="utf-8") as f: + done = json.load(f) + return _run(names, done["records"], done["interp"]) + run = _run(names) + with open(shared, "w", encoding="utf-8") as f: + json.dump({"records": run.records, "interp": run.interp}, f) + return run + + +def _reasons(static: S.StaticFootprint) -> Counter: + return Counter( + "+".join(sorted({e.reason for e in why})) for why in static.excluded.values() + ) + + +def _describe(key, s: set[S.Point], d: set[S.Point]) -> str: + only_s, only_d = sorted(s - d), sorted(d - s) + return ( + f"{key}: (pid_0, pid_1, pid_2, offset) static-only {only_s[:8]} ({len(only_s)}), " + f"interpreter-only {only_d[:8]} ({len(only_d)}); |static|={len(s)} |interpreter|={len(d)}" + ) + + +# tt.get_program_id / tt.get_num_programs, custom (``x``) or generic +# (``"() <{axis = 0 : i32}>``) form +_RE_PID_OP = re.compile( + r"\btt\.get_(?:program_id|num_programs)(?:\s+([xyz])\b|\"\(\)\s*<\{axis\s*=\s*(\d+))" +) + + +def _ttir_pid_axes(text: str) -> frozenset[int]: + return frozenset( + "xyz".index(word) if word else int(num) + for word, num in _RE_PID_OP.findall(text) + ) + + +# D10b version policy (see the module docstring) +_UNTESTED = untested_triton_version() + + +@pytest.mark.xfail( + _UNTESTED is not None, + reason=f"Triton {_UNTESTED} is outside TESTED_TRITON_VERSIONS {TESTED_TRITON_VERSIONS}; " + "set TILELENS_IR_ALLOW_UNTESTED_TRITON=1 to require conformance", + strict=False, +) +@pytest.mark.parametrize("name", NAMES) +def test_reader_conformance(name, conformance, record_property): + out = conformance.outcomes[name] + exp = EXPECTED.get(name, Expect()) + rec = out.record + props: dict[str, Any] = { + "source": rec.get("source"), + "seconds": conformance.seconds, + } + record_property("conformance", props) + assert not rec.get( + "error" + ), f"TTIR capture failed: {rec.get('error')}\n{rec.get('traceback', '')}" + assert out.reader_error is None, f"the reader crashed: {out.reader_error}" + + if out.refusal is not None: + e = out.refusal + props["outcome"] = f"refused:{e.kind}" + assert ( + exp.refusal is not None + ), f"unexpected refusal ({e.kind}): {e.message} (TTIR line {e.line_no})" + assert ( + e.kind is exp.refusal + ), f"refused as {e.kind}, expected {exp.refusal}: {e.message}" + return + + props["outcome"] = "error" + assert ( + not exp.must_refuse + ), f"the reader accepted a kernel it must refuse as {exp.refusal}: {exp.why}" + axes = _ttir_pid_axes(rec["ttir"]) + assert ( + out.pid_axes == axes + ), f"graph.pid_axes is {sorted(out.pid_axes)}; the TTIR's program-id ops read axes {sorted(axes)}" + assert out.static_error is None, f"static evaluation failed: {out.static_error}" + static, run = out.static, out.interp + assert static is not None and run is not None + assert not run["errors"] and not run["wild"], ( + "interpreter run failed:\n" + "\n".join(run["errors"]) + ) + + # sites are keyed by line: every one must sit in the kernel's own file + kernel_file = rec["file"] + for key, files in static.files.items(): + assert {os.path.realpath(f) for f in files} == { + kernel_file + }, f"{key} sits in {files}, not {kernel_file}" + dynamic: dict[S.SiteKey, set[S.Point]] = {} + widths: dict[S.SiteKey, set[int]] = {} + for site in run["sites"]: + arg, kind, line = site["arg"], site["kind"], site["line"] + assert ( + site["file"] == kernel_file + ), f"the interpreter's {kind} of {arg!r} sits in {site['file']}:{line}, not {kernel_file}" + dynamic.setdefault((arg, kind, line), set()).update( + tuple(p) for p in site["points"] + ) + widths.setdefault((arg, kind, line), set()).update(site["bits"]) + + compared = (set(static.sites) | set(dynamic)) - set(static.excluded) + mismatches = [ + _describe(k, static.sites.get(k, set()), dynamic.get(k, set())) + for k in sorted(compared) + if static.sites.get(k, set()) != dynamic.get(k, set()) + ] + exercised = [k for k in compared if dynamic.get(k)] + reasons = _reasons(static) + if mismatches: + outcome = "mismatch" + elif exercised: + outcome = "conform" + else: + outcome = ( + "not-compared" # every site excluded (the table says so, or the test fails) + ) + props.update( + outcome=outcome, + compared_sites=len(compared), + exercised_sites=len(exercised), + points=sum(len(dynamic.get(k, ())) for k in compared), + excluded=dict(reasons), + skipped_accesses=sum( + e.reason in S.SKIPPED for why in static.excluded.values() for e in why + ), + obligation_failures=len(static.obligation_failures), + ) + assert not mismatches, "footprints differ:\n" + "\n".join(mismatches) + bad_widths = [ + f"{k}: the reader's elem_bits {sorted(static.bits.get(k, ()))}, " + f"the interpreter's pointers {sorted(widths[k])}" + for k in sorted(compared) + if k in widths and static.bits.get(k) != widths[k] + ] + assert not bad_widths, "element widths differ:\n" + "\n".join(bad_widths) + assert reasons == Counter(exp.excluded), ( + f"excluded sites {dict(reasons)}, expected {dict(exp.excluded)}: " + + "; ".join( + f"{k}: {[(e.reason, e.detail) for e in v]}" + for k, v in static.excluded.items() + ) + ) + if exp.compared: + assert ( + exercised + ), f"no compared site touched memory (compared: {sorted(compared)})" + else: + assert ( + not compared + ), f"expected every site excluded, compared {sorted(compared)}" + + +# The op tl.debug_barrier() prints, per Triton release: a_debug_ops must +# read it, and the release's reader vocabulary must hold it as inert. +_BARRIER_OP = {"3.6": "gpu.barrier", "3.8": "ttg.barrier"} + + +@pytest.mark.xfail( + _UNTESTED is not None, + reason=f"Triton {_UNTESTED} is outside TESTED_TRITON_VERSIONS {TESTED_TRITON_VERSIONS}; " + "set TILELENS_IR_ALLOW_UNTESTED_TRITON=1 to require conformance", + strict=False, +) +def test_the_debug_barrier_case_reads_the_release_barrier(conformance): + rec = conformance.outcomes["a_debug_ops"].record + assert not rec.get("error"), rec.get("error") + release = _mlir_walk.triton_release()[0] + op = _BARRIER_OP[release] + assert re.search(rf"^\s*{re.escape(op)}\b", rec["ttir"], re.M), op + assert op in ttir_reader._VOCABULARIES[release].inert + assert conformance.outcomes["a_debug_ops"].refusal is None + + +@pytest.mark.xfail( + _UNTESTED is not None, + reason=f"Triton {_UNTESTED} is outside TESTED_TRITON_VERSIONS {TESTED_TRITON_VERSIONS}; " + "set TILELENS_IR_ALLOW_UNTESTED_TRITON=1 to require conformance", + strict=False, +) +@pytest.mark.parametrize("name", NAMES) +def test_host_ttir_is_the_jit_ttir(name, conformance, record_property): + """D25: IR mode reads host-compiled TTIR. The JIT's own compile of the + same launch for the same target (cuda:89, a stand-in driver's) must be + the same kernel: the same hash and the same TTIR text.""" + rec = conformance.outcomes[name].record + if rec.get("error"): + pytest.skip("the capture failed (test_reader_conformance reports it)") + same = rec["ttir"] == rec["jit_ttir"] and rec["hash"] == rec["jit_hash"] + record_property("host_vs_jit", {"same_text": same}) + assert rec["jit_target"] == ["cuda", 89, 32] + assert rec["hash"] == rec["jit_hash"] + assert rec["ttir"] == rec["jit_ttir"] diff --git a/tests/conftest.py b/tests/conftest.py index 53387a160..722e5958d 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,5 +1,36 @@ +from __future__ import annotations + +from pathlib import Path + import pytest +TESTS = Path(__file__).resolve().parent + +# ─────────── IR mode on a Triton outside the tested window (D29) ─────────── +# +# IR mode (tilelens.ir, Sanitizer(compile=True), the host compile) runs only on +# the Triton releases in tilelens.core.config.TESTED_TRITON_VERSIONS unless +# TILELENS_IR_ALLOW_UNTESTED_TRITON=1 says otherwise (D10b), and so do its +# tests: on any other release every test marked IR_MODE is skipped, with a +# reason naming the installed Triton and the window. Running them with the +# override is how a release joins the window. The modules below are marked +# here; a test module elsewhere opts in with ``pytestmark = pytest.mark.ir_mode``. +# The tests that the gate itself refuses correctly live in +# tests/unit/test_ir_version_gate.py, which is not marked and runs on every +# release; the reader conformance suite (tests/conformance/) is not marked +# either: it runs everywhere, as a non-strict xfail outside the window. +IR_MODE = "ir_mode" +# Paths relative to tests/: an entry ending in "/" covers every module below +# that directory, one ending in "*" every module whose path it prefixes. +IR_MODE_MODULES: tuple[str, ...] = ( + "unit/ir/", + "unit/sanitizer_compiled/", + "unit/test_ir_lifecycle.py", + "end_to_end/test_ir_*", + "end_to_end/test_compiled_sanitizer.py", + "end_to_end/test_host_compile.py", +) + def pytest_addoption(parser): group = parser.getgroup("tilelens") @@ -14,6 +45,96 @@ def pytest_addoption(parser): ) +def pytest_configure(config): + config.addinivalue_line( + "markers", + f"{IR_MODE}: a test of IR mode, skipped on a Triton release outside " + "tilelens.core.config.TESTED_TRITON_VERSIONS unless " + "TILELENS_IR_ALLOW_UNTESTED_TRITON=1 (D29)", + ) + + +def is_ir_mode_module(path: Path) -> bool: + """Whether the test module at ``path`` is one of IR_MODE_MODULES.""" + try: + relative = Path(path).resolve().relative_to(TESTS).as_posix() + except ValueError: # not under tests/ + return False + for entry in IR_MODE_MODULES: + if entry.endswith(("/", "*")): + if relative.startswith(entry.rstrip("*")): + return True + elif relative == entry: + return True + return False + + +def ir_mode_skip_reason() -> str | None: + """Why the IR-mode tests skip on the installed Triton; None when they + run (a release in the window, or the override set).""" + from tilelens.core.config import TESTED_TRITON_VERSIONS, untested_triton_version + + version = untested_triton_version() + if version is None: + return None + window = ", ".join(f"{release}.x" for release in TESTED_TRITON_VERSIONS) + return ( + f"IR mode is not tested on the installed Triton {version}: the tested " + "window (tilelens.core.config.TESTED_TRITON_VERSIONS) is Triton " + f"{window}; set TILELENS_IR_ALLOW_UNTESTED_TRITON=1 to run the IR-mode " + "tests anyway (D29)" + ) + + +def pytest_collection_modifyitems(config, items): + marked = [] + for item in items: + if is_ir_mode_module(item.path): + item.add_marker(IR_MODE) + if item.get_closest_marker(IR_MODE) is not None: + marked.append(item) + reason = ir_mode_skip_reason() if marked else None + if reason is not None: + # First of the item's skipif marks, so its reason is the one given + # even where another (e.g. "needs a CUDA GPU") holds as well. + skip = pytest.mark.skipif(True, reason=reason) + for item in marked: + item.add_marker(skip, append=False) + + +@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_compiled_sanitizer.py b/tests/end_to_end/test_compiled_sanitizer.py new file mode 100644 index 000000000..69991b5c6 --- /dev/null +++ b/tests/end_to_end/test_compiled_sanitizer.py @@ -0,0 +1,1882 @@ +"""Sanitizer(compile=True) on real kernels: each launch is compiled on the host +for the client's target (D25, D26), checked against its TTIR and not run, on +CPU tensors and with Triton's driver unreachable: no GPU is needed. Ports +#361's tests/end_to_end/test_compiled_sanitizer.py onto the IR lifecycle. +Counterparts on fake launches live in +tests/unit/sanitizer_compiled/test_client.py. +""" + +from __future__ import annotations + +import importlib +import inspect +import os +import subprocess +import sys +import threading +from pathlib import Path + +import pytest +import torch +import triton +import triton.language as tl +from triton.backends.compiler import GPUTarget + +import tilelens +from tilelens.clients import Sanitizer, Tracer +from tilelens.clients.sanitizer.compiled import CompiledSanitizer +from tilelens.clients.sanitizer.data import CompiledSanitizerRecord +from tilelens.core.config import DEFAULT_IR_TARGET, Config +from tilelens.core.data import Load, Store +from tilelens.ir import IRVerdict +from tilelens.ir.verdict import SourceLocation + +trace_module = importlib.import_module("tilelens.core.trace") +config_module = importlib.import_module("tilelens.core.config") +REPO = Path(__file__).resolve().parents[2] + + +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 (D25). + 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 (D25): 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 (D26), whatever TILELENS_IR_TARGET the caller + set: in the process config, and in any Config read from the environment + (configured_target below sets its own).""" + 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 + + +def _sanitizer() -> CompiledSanitizer: + det = Sanitizer(compile=True, abort_on_error=False) + assert isinstance(det, CompiledSanitizer) + return det + + +def _line(kernel, needle: str) -> int: + """The source line of ``kernel`` (a JITFunction) holding ``needle``.""" + lines, start = inspect.getsourcelines(kernel.fn) + (index,) = [i for i, line in enumerate(lines) if needle in line] + return start + index + + +def _lane(record: CompiledSanitizerRecord) -> int: + (lane,) = [v for k, v in record.witness.items() if k.startswith("arange_")] + return lane + + +def _make_add(): + @triton.jit + def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr): + pid = tl.program_id(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + x = tl.load(x_ptr + offs, mask=mask) + y = tl.load(y_ptr + offs, mask=mask) + tl.store(out_ptr + offs, x + y, mask=mask) + + return add_kernel + + +def _make_add_nomask(): + @triton.jit + def add_nomask(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + pid = tl.program_id(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + x = tl.load(x_ptr + offs) # no mask: OOB on a ragged tail + tl.store(out_ptr + offs, x) + + return add_nomask + + +# ======== proofs and findings ========= + + +def test_an_in_bounds_launch_is_proved_and_nothing_runs(): + det = _sanitizer() + traced = tilelens.trace(det)(_make_add()) + n = 3000 # a ragged tail, masked + x, y = torch.randn(n), torch.randn(n) + out = torch.zeros(n) + + kernel = traced[(triton.cdiv(n, 1024),)](x, y, out, n, BLOCK=1024) + + assert det.last_status == "ok" and det.records == [] + verdict = det.last_verdict + assert trace_module.launches[-1].records == [verdict] + (config,) = verdict.per_config + assert (config.specialization, config.status) == (kernel.hash, "ok") + assert verdict.refusal is None and verdict.notes == () + # LAUNCH="skip": compiled and checked, never run. + assert torch.count_nonzero(out) == 0 + + +def test_an_unmasked_tail_is_reported_at_its_line_and_device_address(): + det = _sanitizer() + add_nomask = _make_add_nomask() + traced = tilelens.trace(det)(add_nomask) + n = 3000 + x, out = torch.randn(n), torch.zeros(n) + + traced[(triton.cdiv(n, 1024),)](x, out, n, BLOCK=1024) + + assert det.last_status == "violations" + load, store = det.records + assert trace_module.launches[-1].records == [load, store, det.last_verdict] + assert [(r.kind, r.op_type, r.tensor_name) for r in (load, store)] == [ + ("out-of-bounds", Load, "x_ptr"), + ("out-of-bounds", Store, "out_ptr"), + ] + for record, tensor, needle in ((load, x, "tl.load"), (store, out, "tl.store")): + offset = record.violation_offset + assert n <= offset < 3 * 1024 + assert record.witness["pid_0"] == 2 + assert 2 * 1024 + _lane(record) == offset + # the address of the offending element + assert record.violation_address == tensor.data_ptr() + offset * 4 + assert record.tensor_facts.data_ptr == tensor.data_ptr() + (tb,) = record.user_code_tracebacks + assert (tb.filename, tb.lineno, tb.func_name) == ( + __file__, + _line(add_nomask, needle), + "add_nomask", + ) + assert needle in tb.line_of_code + assert torch.count_nonzero(out) == 0 + + +def test_each_launch_is_checked_against_its_own_arguments(): + det = _sanitizer() + traced = tilelens.trace(det)(_make_add_nomask()) + + # An exact multiple of BLOCK: in bounds. + x, out = torch.randn(4096), torch.zeros(4096) + traced[(4,)](x, out, 4096, BLOCK=1024) + assert (det.last_status, det.records) == ("ok", []) + + # A ragged tail with the same specialization: out of bounds. + x, out = torch.randn(3000), torch.zeros(3000) + traced[(3,)](x, out, 3000, BLOCK=1024) + assert det.last_status == "violations" and len(det.records) == 2 + + +def test_a_loop_advancing_a_pointer_is_checked_per_iteration(): + @triton.jit + def loop_sum(x_ptr, out_ptr, n_iters, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + ptrs = x_ptr + offs + acc = tl.zeros((BLOCK,), tl.float32) + for _ in range(n_iters): + acc += tl.load(ptrs) + ptrs += BLOCK + tl.store(out_ptr + offs, acc) + + det = _sanitizer() + traced = tilelens.trace(det)(loop_sum) + x, out = torch.randn(64), torch.zeros(16) + + traced[(1,)](x, out, 4, BLOCK=16) # 4 * 16 == 64 + assert (det.last_status, det.records) == ("ok", []) + traced[(1,)](x, out, 0, BLOCK=16) # a zero-trip loop reads nothing + assert (det.last_status, det.records) == ("ok", []) + + traced[(1,)](x, out, 5, BLOCK=16) + (record,) = det.records + assert (record.kind, record.op_type, record.tensor_name) == ( + "out-of-bounds", + Load, + "x_ptr", + ) + assert record.witness["iter_loop"] == 4 + assert record.violation_offset == 4 * 16 + _lane(record) + assert record.user_code_tracebacks[0].lineno == _line(loop_sum, "tl.load") + + +def test_a_store_loop_without_an_accumulator_is_checked(): + @triton.jit + def store_loop(out_ptr, iters, BLOCK: tl.constexpr): + for i in range(0, iters): + offs = i * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.full((BLOCK,), 1.0, tl.float32)) + + det = _sanitizer() + traced = tilelens.trace(det)(store_loop) + out = torch.zeros(16) + traced[(1,)](out, 4, BLOCK=4) # 4 * 4 == 16, exactly fits + assert (det.last_status, det.records) == ("ok", []) + traced[(1,)](out, 6, BLOCK=4) # 6 * 4 == 24 > 16 + assert det.last_status == "violations" + assert {r.op_type for r in det.records} == {Store} + assert torch.count_nonzero(out) == 0 + + +def test_a_strided_view_is_checked_against_its_elements(): + @triton.jit + def strided(x_ptr, out_ptr, n, stride, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs * stride, mask=mask), mask=mask) + + @triton.jit + def ignores_stride(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask), mask=mask) + + x = torch.randn(64, 2)[:, 0] # stride 2: every other float + out = torch.zeros(64) + assert not x.is_contiguous() + + det = _sanitizer() + tilelens.trace(det)(strided)[(1,)](x, out, 64, x.stride(0), BLOCK=64) + assert (det.last_status, det.records) == ("ok", []) + + # Offsets 0..63 land in the view's gaps (D12: gaps are out of bounds). + tilelens.trace(det)(ignores_stride)[(1,)](x, out, 64, BLOCK=64) + (record,) = det.records + assert (record.kind, record.op_type) == ("out-of-bounds", Load) + assert record.violation_offset % 2 == 1 + assert (record.tensor_facts.strides, record.tensor_facts.contiguous) == ( + (2,), + False, + ) + + +def test_an_i32_wrap_is_an_integer_overflow(): + @triton.jit + def i32_wrap(x_ptr, S): + pid = tl.program_id(0) + off = (pid * S) * S # wraps in i32 for S = 65536 + tl.store(x_ptr + off, 1) + + det = _sanitizer() + traced = tilelens.trace(det)(i32_wrap) + x = torch.zeros(16, dtype=torch.int32) + + traced[(4,)](x, 2) + assert (det.last_status, det.records) == ("ok", []) + + traced[(4,)](x, 65536) + (record,) = det.records + assert (record.kind, record.op_type, record.tensor_name) == ( + "integer-overflow", + Store, + "x_ptr", + ) + assert (record.violation_offset, record.violation_address) == (None, None) + assert not -(1 << 31) <= record.witness["value"] < 1 << 31 + assert record.user_code_tracebacks[0].lineno == _line(i32_wrap, "off = ") + assert torch.count_nonzero(x) == 0 + + +def test_a_division_by_a_zero_argument_is_reported(): + @triton.jit + def divide(x_ptr, d): + pid = tl.program_id(0) + q = pid // d + tl.store(x_ptr + q, 1) + + det = _sanitizer() + traced = tilelens.trace(det)(divide) + x = torch.zeros(2, dtype=torch.int32) + + traced[(4,)](x, 2) + assert (det.last_status, det.records) == ("ok", []) + + traced[(4,)](x, 0) + (record,) = det.records + assert (record.kind, record.op_type) == ("division-by-zero", Store) + (tb,) = record.user_code_tracebacks + assert (tb.lineno, tb.line_of_code.strip()) == ( + _line(divide, "q = "), + "q = pid // d", + ) + + +BIG32 = (1 << 31) - 1 + + +def _make_circular(): + """The audit corpus's N01 / N06: at pid 1, t = pid + BIG wraps in i32, so + the kernel divides by 3 (-3) where the unbounded reading divides by 0.""" + + @triton.jit + def circ_offset(x_ptr, BIG): + pid = tl.program_id(0) + t = pid + BIG + d = tl.where(t > BIG, 0, 3) + off = (t // d).to(tl.int64) - BIG // 3 + tl.store(x_ptr + off, 1.0) # pid 1: a wild store + + @triton.jit + def circ_loop(x_ptr, BIG): + pid = tl.program_id(0) + t = pid + BIG + d = tl.where(t > BIG, 0, -3) + hi = tl.minimum((t // d) - BIG // -3 + 1, 4) + for i in range(0, hi): + tl.store(x_ptr + i, 1.0) # pid 1: x[1..3] + + return circ_offset, circ_loop + + +@pytest.mark.parametrize("which", [0, 1], ids=["offset", "loop"]) +def test_a_wrap_that_decides_its_own_divisor_is_never_a_proof(which): + kernel = _make_circular()[which] + det = _sanitizer() + traced = tilelens.trace(det)(kernel) + x = torch.zeros(1) + + traced[(1,)](x, BIG32) # pid 0 alone: offset 0, one iteration + assert (det.last_status, det.records) == ("ok", []) + + traced[(2,)](x, BIG32) + assert det.last_status == "violations" + (record,) = det.records + assert (record.kind, record.witness["pid_0"]) == ("integer-overflow", 1) + assert record.user_code_tracebacks[0].lineno == _line(kernel, "t = pid + BIG") + + +def test_a_wrap_a_where_discards_is_no_finding(): + """The audit's p11: i * S wraps on the lanes the where discards.""" + + @triton.jit + def guarded_scale(x_ptr, n, S, BLOCK: tl.constexpr): + i = tl.arange(0, BLOCK) + off = tl.where(i < n, (i * S) // S, 0) + tl.store(x_ptr + off, 1.0) + + det = _sanitizer() + traced = tilelens.trace(det)(guarded_scale) + x = torch.zeros(256) + traced[(1,)](x, 100, 1 << 24, BLOCK=256) + assert (det.last_status, det.records) == ("ok", []) + traced[(1,)](x, 200, 1 << 24, BLOCK=256) # lanes 128..199 read i * S: wraps + assert [r.kind for r in det.records] == ["integer-overflow"] + + +def test_launches_on_two_host_threads_check_at_the_same_time(): + """Two traced kernels, each with its own sanitizer, launched from two + host threads at once: their Z3 checks overlap (one shared Z3 context + segfaulted here).""" + kernels = [_make_add_nomask(), _make_add_nomask()] + dets = [_sanitizer(), _sanitizer()] + traced = [tilelens.trace(d)(k) for d, k in zip(dets, kernels)] + tensors = [(torch.randn(4096), torch.zeros(4096)) for _ in kernels] + for t, (x, out) in zip(traced, tensors): # compile first, one at a time + t[(3,)](x, out, 3000, BLOCK=1024) + barrier = threading.Barrier(len(kernels)) + errors: list[BaseException] = [] + statuses: list[list] = [[] for _ in kernels] + + def launch(i): + try: + barrier.wait() + for k in range(6): + # another n each time: a check of its own, not a remembered one + traced[i][(5,)](*tensors[i], 4097 + k, BLOCK=1024) + statuses[i].append(dets[i].last_status) + except BaseException as exc: # noqa: BLE001 - reported below + errors.append(exc) + + threads = [threading.Thread(target=launch, args=(i,)) for i in range(len(kernels))] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + assert errors == [] + assert statuses == [["violations"] * 6] * len(kernels) + + +def test_a_grouped_swizzle_matmul_is_modeled(): + """The tutorial-03 grouped swizzle (``//``, ``%`` and ``min`` over launch + quantities) with the ``% M`` / ``% N`` row clamps removed (TritonBench's + matmul_triton2): a proof when M, N, K cover the blocks, the real OOB + otherwise.""" + + @triton.jit + def swizzle_matmul( + a_ptr, b_ptr, c_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, + BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, + GROUP_M: tl.constexpr, + ): # fmt: skip + pid = tl.program_id(0) + num_pid_m = tl.cdiv(M, BLOCK_M) + num_pid_n = tl.cdiv(N, BLOCK_N) + num_pid_in_group = GROUP_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + offs_am = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) # no `% M` clamp + offs_bn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) # no `% N` clamp + offs_k = tl.arange(0, BLOCK_K) + a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak + b_ptrs = b_ptr + offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + for k in range(0, tl.cdiv(K, BLOCK_K)): + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * BLOCK_K, other=0.0) + acc += tl.dot(a, b) + a_ptrs += BLOCK_K * stride_ak + b_ptrs += BLOCK_K * stride_bk + offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) + tl.store(c_ptrs, acc, mask=c_mask) + + def launch(m, n, k): + det = _sanitizer() + a = torch.randn(m, k) + b = torch.randn(k, n) + c = torch.empty(m, n) + grid = (triton.cdiv(m, 32) * triton.cdiv(n, 32),) + tilelens.trace(det)(swizzle_matmul)[grid]( + a, b, c, m, n, k, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), + c.stride(0), c.stride(1), + BLOCK_M=32, BLOCK_N=32, BLOCK_K=32, GROUP_M=8, + ) # fmt: skip + return det + + clean = launch(64, 64, 64) + assert (clean.last_status, clean.records) == ("ok", []), clean.last_verdict + # M = N = K = 16 < BLOCK: the K-only masks leave rows 16..31 of A (and + # cols 16..31 of B) unguarded. + buggy = launch(16, 16, 16) + assert buggy.last_status == "violations", buggy.last_verdict + assert {r.tensor_name for r in buggy.records} >= {"a_ptr", "b_ptr"} + + +# ======== branches: path conditions and abstentions ========= + + +def test_a_modeled_branch_condition_leaves_no_false_witness(): + """``if t > 0: load(p + t * n_cols + offs - n_cols)`` never reads offset + -n_cols: the t == 0 iteration takes the other branch.""" + + @triton.jit + def guarded_scan(x_ptr, out_ptr, n_steps, n_cols, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n_cols + acc = tl.zeros((BLOCK,), tl.float32) + for i in range(n_steps): + t = n_steps - 1 - i + if t > 0: + prev = tl.load(x_ptr + t * n_cols + offs - n_cols, mask=mask, other=0) + else: + prev = tl.zeros((BLOCK,), tl.float32) + acc += prev + tl.store(out_ptr + offs, acc, mask=mask) + + det = _sanitizer() + x, out = torch.randn(5 * 8), torch.zeros(8) + tilelens.trace(det)(guarded_scan)[(1,)](x, out, 5, 8, BLOCK=8) + assert (det.last_status, det.records) == ("ok", []), det.last_verdict + + +def test_a_data_dependent_branch_abstains(): + """A possible OOB under a branch on loaded data is unsupported: never a + witness from a branch that may not run, never "ok".""" + + @triton.jit + def flag_gated(flag_ptr, x_ptr, out_ptr, n_cols, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n_cols + flag = tl.load(flag_ptr) + acc = tl.zeros((BLOCK,), tl.float32) + if flag > 0: + acc = tl.load(x_ptr + offs - n_cols, mask=mask, other=0) + tl.store(out_ptr + offs, acc, mask=mask) + + det = _sanitizer() + flag = torch.zeros(1, dtype=torch.int32) + x, out = torch.randn(8), torch.zeros(8) + tilelens.trace(det)(flag_gated)[(1,)](flag, x, out, 8, BLOCK=8) + assert (det.last_status, det.records) == ("unsupported", []) + refusal = det.last_verdict.refusal + assert refusal.kind == "unmodelable-condition" + assert refusal.loc.line == _line(flag_gated, "offs - n_cols") + + +def test_an_unguarded_oob_is_reported_beside_a_branch(): + @triton.jit + def mixed(x_ptr, out_ptr, n, flag, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + x = tl.load(x_ptr + offs) # unguarded: OOB when numel < BLOCK + if flag > 0: + x += tl.load(x_ptr + offs, mask=offs < n, other=0) # guarded, safe + tl.store(out_ptr + offs, x, mask=offs < n) + + det = _sanitizer() + x, out = torch.randn(8), torch.zeros(8) + tilelens.trace(det)(mixed)[(1,)](x, out, 8, 1, BLOCK=16) + assert det.last_status == "violations", det.last_verdict + (record,) = det.records + assert record.op_type is Load + assert record.user_code_tracebacks[0].lineno == _line(mixed, "unguarded") + + +# ======== what cannot be checked ========= + + +def test_a_gather_is_unsupported_at_its_line(): + @triton.jit + def gather(idx_ptr, src_ptr, out_ptr, n, BLOCK: tl.constexpr): + pid = tl.program_id(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + idx = tl.load(idx_ptr + offs, mask=mask) + vals = tl.load(src_ptr + idx, mask=mask) # data-dependent address + tl.store(out_ptr + offs, vals, mask=mask) + + det = Sanitizer(compile=True) # abort_on_error: unsupported never exits + n = 1024 + idx = torch.zeros(n, dtype=torch.int32) + src, out = torch.randn(n), torch.zeros(n) + tilelens.trace(det)(gather)[(4,)](idx, src, out, n, BLOCK=256) + + assert (det.last_status, det.records) == ("unsupported", []) + refusal = det.last_verdict.refusal + assert refusal.kind == "indirect-address" + assert (refusal.loc.file, refusal.loc.line) == ( + __file__, + _line(gather, "src_ptr + idx"), + ) + assert refusal.message.startswith(f"{__file__}:{refusal.loc.line}: ") + assert det.last_verdict.per_config[0].refusal == refusal + + +def test_nested_loops_are_unsupported(): + @triton.jit + def nested(in_ptr, out_ptr, M, N, BLOCK: tl.constexpr): + for i in range(0, M): + for j in range(0, N): + offs = (i * N + j) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(in_ptr + offs)) + + det = _sanitizer() + inp, out = torch.randn(20), torch.zeros(20) + tilelens.trace(det)(nested)[(1,)](inp, out, 8, 2, BLOCK=4) + assert (det.last_status, det.records) == ("unsupported", []) + assert det.last_verdict.refusal.kind == "nested-loop" + + +def _release() -> str: + """The installed Triton's minor release, e.g. "3.6".""" + return ".".join(triton.__version__.split(".")[:2]) + + +# How a kernel taking a tuple, or a host TensorDescriptor, is refused, per +# Triton release: 3.6 names every TTIR argument the parameter flattens to by +# the parameter's own name, which the reader refuses (two parameters of one +# name); 3.8 names each by its path in the tuple (``ptrs.0``), which reads, +# and binds to no launch argument (its host descriptor's leaves still repeat +# a name, ``d.shape.0``). Never checked, never a finding either way. +_AGGREGATE_REFUSALS = { + "3.6": { + "tuple of pointers": "other", + "tuple of ints": "other", + "descriptor": "other", + }, + "3.8": { + "tuple of pointers": "missing-binding", + "tuple of ints": "missing-binding", + "descriptor": "other", + }, +} + + +def _tuple_of_pointers(): + @triton.jit + def tuple_of_pointers(ptrs, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + vals = tl.load(ptrs[0] + offs) # unmasked: out of bounds of 32 elements + tl.store(ptrs[1] + offs, vals, mask=offs < n) + + return tuple_of_pointers, ((torch.zeros(32), torch.zeros(64)), 64), {"BLOCK": 64} + + +def _tuple_of_ints(): + @triton.jit + def tuple_of_ints(x_ptr, bounds, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + bounds[0] + tl.store(x_ptr + offs, 1.0, mask=offs < bounds[1]) # out of bounds from 8 + + return tuple_of_ints, (torch.zeros(64), (8, 72)), {"BLOCK": 64} + + +def _host_descriptor(): + from triton.tools.tensor_descriptor import TensorDescriptor + + @triton.jit + def descriptor_copy(desc, out_ptr, BM: tl.constexpr, BN: tl.constexpr): + offs = tl.arange(0, BM)[:, None] * BN + tl.arange(0, BN)[None, :] + tl.store(out_ptr + offs, desc.load([0, 0])) + + desc = TensorDescriptor.from_tensor(torch.zeros(64, 64), [32, 32]) + return descriptor_copy, (desc, torch.zeros(16 * 16)), {"BM": 32, "BN": 32} + + +@pytest.mark.parametrize( + "case, make", + [ + ("tuple of pointers", _tuple_of_pointers), + ("tuple of ints", _tuple_of_ints), + ("descriptor", _host_descriptor), + ], +) +def test_a_tuple_or_host_descriptor_parameter_is_unsupported(case, make): + refusals = _AGGREGATE_REFUSALS.get(_release()) + if refusals is None: + pytest.fail(f"no aggregate-parameter refusals for Triton {_release()}") + kernel, args, kwargs = make() + det = _sanitizer() + tilelens.trace(det)(kernel)[(1,)](*args, **kwargs) + assert (det.last_status, det.records) == ("unsupported", []), det.last_verdict + assert det.last_verdict.refusal.kind == refusals[case], det.last_verdict.refusal + + +# ======== configs (D3) ========= + + +def test_autotune_reports_the_one_config_that_goes_out_of_bounds(): + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_warps=1), + triton.Config({"BLOCK": 128}, num_warps=1), + ], + key=["n"], + ) + @triton.jit + def copy(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) # no mask + + det = _sanitizer() + x, out = torch.randn(64), torch.zeros(64) + ret = tilelens.trace(det)(copy)[lambda meta: (triton.cdiv(64, meta["BLOCK"]),)]( + x, out, 64 + ) + + assert ret is None # a skipped autotuned launch picks no config + verdict = det.last_verdict + assert verdict.status == "violations" + assert [(c.config["BLOCK"], c.status) for c in verdict.per_config] == [ + (16, "ok"), + (128, "violations"), + ] + assert {r.config["BLOCK"] for r in det.records} == {128} + assert verdict.per_config[1].n_reports == len(det.records) == 2 + assert torch.count_nonzero(out) == 0 + + +def test_configs_compiling_to_one_kernel_are_each_checked(): + """D22 (the audit's D3 probe): S is a runtime int that 2 and 3 + specialize alike, so both configs compile to one kernel; only S=3, on + its larger grid, goes out of bounds. Deduplicating events by kernel + alone would check S=2's binding only and prove the launch.""" + + @triton.autotune( + configs=[ + triton.Config({"S": 2}, num_warps=1), + triton.Config({"S": 3}, num_warps=1), + ], + key=["n"], + ) + @triton.jit + def strided_copy(x_ptr, out_ptr, n, S, BLOCK: tl.constexpr): + # Program p covers [p * BLOCK * S, p * BLOCK * S + BLOCK). + offs = tl.program_id(0) * BLOCK * S + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + + det = _sanitizer() + x, out = torch.randn(64), torch.zeros(64) + # S=2: 2 programs, up to element 47; S=3: 4 programs, up to 159. + tilelens.trace(det)(strided_copy)[lambda meta: (4 if meta["S"] == 3 else 2,)]( + x, out, 64, BLOCK=16 + ) + + verdict = det.last_verdict + assert verdict.status == "violations" + assert [(c.config["S"], c.status) for c in verdict.per_config] == [ + (2, "ok"), + (3, "violations"), + ] + assert len({c.specialization for c in verdict.per_config}) == 1 + assert {(r.tensor_name, r.config["S"]) for r in det.records} == { + ("x_ptr", 3), + ("out_ptr", 3), + } + assert torch.count_nonzero(out) == 0 + + +def test_a_config_that_fails_to_compile_is_a_note(): + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_warps=1), + triton.Config({"BLOCK": 64}, num_warps=1), + ], + key=["n"], + ) + @triton.jit + def add_one(x_ptr, out_ptr, n, BLOCK: 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) + + det = _sanitizer() + x, out = torch.randn(64), torch.zeros(64) + tilelens.trace(det)(add_one)[lambda meta: (triton.cdiv(64, meta["BLOCK"]),)]( + x, out, 64 + ) + + verdict = det.last_verdict + assert verdict.status == "ok" + (config,) = verdict.per_config + assert config.config["BLOCK"] == 16 + (note,) = verdict.notes + assert "'BLOCK': 64" in note and "CompileTimeAssertionFailure" in note + + +# ======== composition, persistence and the process exit ========= + + +def test_a_mixed_trace_with_the_eager_tracer(): + """D4b: the compiled sanitizer checks the compiled kernel; the tracer + gets the interpreted run, which writes the outputs.""" + det, tracer = _sanitizer(), Tracer() + traced = tilelens.trace(tracer)(tilelens.trace(det)(_make_add())) + n = 100 + x, y = torch.randn(n), torch.randn(n) + out = torch.zeros(n) + + traced[(1,)](x, y, out, n, BLOCK=128) + + assert det.last_status == "ok" + records = trace_module.launches[-1].records + assert det.last_verdict in records + assert {type(r) for r in records if not isinstance(r, IRVerdict)} >= {Load, Store} + torch.testing.assert_close(out, x + y) + + +def test_a_real_launch_round_trips_through_a_saved_trace(tmp_path): + det = _sanitizer() + traced = tilelens.trace(det)(_make_add_nomask()) + x, out = torch.randn(3000), torch.zeros(3000) + traced[(3,)](x, out, 3000, BLOCK=1024) + launch = trace_module.launches[-1] + assert launch.records[:-1] == det.records and len(det.records) == 2 + + saved = list(trace_module.launches) + trace_module.launches[:] = [launch] + try: + path = tilelens.save(tmp_path / "trace.zip") + (loaded,) = tilelens.load(path) + finally: + trace_module.launches[:] = saved + assert loaded.records == launch.records + assert loaded.records[-1] == det.last_verdict + + +# Triton's on-disk cache keys a kernel by its source and first line, not its +# file, and a cached TTIR keeps the locs of the file first compiled: the +# comment makes each script's kernel its own (a TTIR-only host compile never +# reaches that cache, a whole-pipeline one does). +_OOB_SCRIPT = """\ +import torch, triton, triton.language as tl +{prelude} + +{decorator} +@triton.jit +def add_nomask(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + # {tag} + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + +n = {n} +x = torch.randn(n) +out = torch.zeros(n) +add_nomask[(triton.cdiv(n, 1024),)](x, out, n, BLOCK=1024) +print("launch returned") +""" + + +def _run(script: Path, *argv: str, **env_vars: str) -> subprocess.CompletedProcess: + # No GPU in the child either (D25), and no IR target but env_vars'. This + # checkout first on the path, then whatever the caller put there (e.g. + # another Triton release). + unset = ("TRITON_INTERPRET", "TILELENS_IR_TARGET", "TRITON_VIZ_IR_TARGET") + env = {k: v for k, v in os.environ.items() if k not in unset} + path = os.pathsep.join([str(REPO), *filter(None, [os.environ.get("PYTHONPATH")])]) + env.update(PYTHONPATH=path, CUDA_VISIBLE_DEVICES="", **env_vars) + return subprocess.run( + [sys.executable, *argv], capture_output=True, text=True, env=env, cwd=REPO + ) + + +@pytest.mark.parametrize("n", [3000, 4096]) +def test_abort_on_error_exits_after_reporting(tmp_path, n): + script = tmp_path / "oob.py" + script.write_text( + _OOB_SCRIPT.format( + prelude="import tilelens\nfrom tilelens.clients import Sanitizer", + decorator="@tilelens.trace(Sanitizer(compile=True))", + n=n, + tag=script, + ) + ) + proc = _run(script, str(script)) + if n == 4096: # in bounds + assert proc.returncode == 0, proc.stderr + assert "launch returned" in proc.stdout + return + assert proc.returncode == 1, proc.stderr + assert proc.stdout.count("Out-Of-Bounds Access Detected") == 2 + lines = script.read_text().splitlines() + (line,) = [i for i, text in enumerate(lines, 1) if "tl.store(" in text] + assert f"File: {script}, Line: {line}, in add_nomask" in proc.stdout + assert "launch returned" not in proc.stdout + + +@pytest.mark.parametrize("command", ["tile-sanitizer", "triton-sanitizer"]) +def test_the_cli_compile_flag_runs_the_compiled_sanitizer(tmp_path, command): + script = tmp_path / "oob.py" + script.write_text(_OOB_SCRIPT.format(prelude="", decorator="", n=3000, tag=script)) + cli = ( + f"import sys; sys.argv = [{command!r}, '--compile', {str(script)!r}]; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + proc = _run(script, "-c", cli) + assert proc.returncode == 1, proc.stderr + assert proc.stdout.count("(compiled sanitizer)") == 2 + assert f"File: {script}, " in proc.stdout + assert "launch returned" not in proc.stdout + + +# ======== the target (D26) ========= + + +def _make_two_cta_configs(): + # num_ctas=2 is an sm90+ option: a compile for the default cuda:89 (or + # any target below sm90) rejects the config. + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_ctas=1), + triton.Config({"BLOCK": 16}, num_ctas=2), + ], + key=["n"], + ) + @triton.jit + def copy_ctas(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + + return copy_ctas + + +def test_the_default_target_is_cuda89_whatever_the_machine(): + det = _sanitizer() + assert det.ir_target is None # the configured default + x, out = torch.zeros(4096), torch.zeros(4096) + kernel = tilelens.trace(det)(_make_add_nomask())[(4,)](x, out, 4096, BLOCK=1024) + assert kernel.target == GPUTarget("cuda", 89, 32) + assert list(kernel.asm) == ["ttir"] # only what the sanitizer reads + assert det.last_status == "ok" + + +def _assert_compile_failed(refusal, target, error): + """A D27 refusal: the config failed to compile for ``target``, with + ``error``, and the refusal says how to check it for another target. It + is one line, starting where in the kernel's source file it failed.""" + assert refusal.kind == "compile-failed" + assert f"failed to compile for {target} (" in refusal.message + assert error in refusal.message + assert "Sanitizer(compile=True, target=...)" in refusal.message + assert "TILELENS_IR_TARGET" in refusal.message + assert "\n" not in refusal.message + loc = refusal.loc + assert loc is not None and refusal.message.startswith(f"{loc.file}:{loc.line}: ") + + +def test_a_target_passed_to_the_sanitizer_is_the_one_checked(): + x, out = torch.zeros(64), torch.zeros(64) + # For the default cuda:89 the sm90 config fails to compile: unchecked, + # and it may run on an sm90 GPU, so the launch is unsupported (D27). + det = _sanitizer() + tilelens.trace(det)(_make_two_cta_configs())[(4,)](x, out, 64) + assert det.last_status == "unsupported" + one, two = det.last_verdict.per_config + assert (one.config["num_ctas"], one.status) == (1, "ok") + assert (two.config["num_ctas"], two.status) == (2, "unsupported") + + for target in ("cuda:90", GPUTarget("cuda", 90, 32)): + det = Sanitizer(compile=True, abort_on_error=False, target=target) + assert det.ir_target == GPUTarget("cuda", 90, 32) + tilelens.trace(det)(_make_two_cta_configs())[(4,)](x, out, 64) + assert det.last_status == "ok", det.last_verdict + per_config = det.last_verdict.per_config + assert [c.config["num_ctas"] for c in per_config] == [1, 2] + + +# ======== a kernel that fails to compile for the target (D27) ========= + + +def test_a_config_that_needs_two_ctas_is_unsupported_and_the_program_goes_on(capsys): + """An autotuned kernel one config of which needs num_ctas=2 (sm90+): that + config is unsupported, kind compile-failed, naming the target and how to + name another; the other is checked; the launch returns (None, as for any + skipped autotuned launch) even with abort_on_error, and the program + goes on to its next launch.""" + det = Sanitizer(compile=True) # abort_on_error=True + kernel = _make_two_cta_configs() + traced = tilelens.trace(det)(kernel) + x, out = torch.zeros(64), torch.zeros(64) + + assert traced[(4,)](x, out, 64) is None + + verdict = det.last_verdict + assert (verdict.status, verdict.scope, verdict.notes) == ("unsupported", None, ()) + ok, failed = verdict.per_config + assert (ok.config["num_ctas"], ok.status) == (1, "ok") + assert (failed.config["num_ctas"], failed.status) == (2, "unsupported") + assert failed.specialization is None and verdict.refusal == failed.refusal + _assert_compile_failed( + failed.refusal, "cuda:89", "num_ctas > 1 requires NVIDIA SM90" + ) + assert det.records == [] + # No source position for an option error: the kernel's def line. + assert failed.refusal.loc == SourceLocation( + kernel.fn.fn.__code__.co_filename, _line(kernel.fn, "def copy_ctas(") + ) + (line,) = capsys.readouterr().out.splitlines() + assert line.startswith("[CompiledSanitizer] not checked (config {") + assert ( + f": compile-failed: {failed.refusal.loc.file}:{failed.refusal.loc.line}: " + "it failed to compile for cuda:89 (ValueError: num_ctas > 1 requires" + ) in line + + # The program goes on: its next launch is checked like any other. + traced[(4,)](x, out, 64) + assert det.last_status == "unsupported" + assert [c.status for c in det.last_verdict.per_config] == ["ok", "unsupported"] + + +def test_a_plain_kernel_launched_with_two_ctas_is_unsupported(): + """A plain kernel whose only config fails to compile for the target: the + launch is unsupported compile-failed and returns None (no kernel to + return), and the next launch compiles as usual.""" + det = Sanitizer(compile=True) # abort_on_error=True + traced = tilelens.trace(det)(_make_add_nomask()) + x, out = torch.zeros(4096), torch.zeros(4096) + + assert traced[(4,)](x, out, 4096, BLOCK=1024, num_ctas=2) is None + + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + (failed,) = verdict.per_config + assert (failed.specialization, failed.config, failed.status) == ( + None, + {}, + "unsupported", + ) + assert verdict.refusal == failed.refusal + _assert_compile_failed( + failed.refusal, "cuda:89", "num_ctas > 1 requires NVIDIA SM90" + ) + + kernel = traced[(4,)](x, out, 4096, BLOCK=1024) + assert kernel.target == GPUTarget("cuda", 89, 32) + assert det.last_status == "ok" + + +def _make_block_asserting_configs(*blocks): + @triton.autotune( + configs=[triton.Config({"BLOCK": block}, num_warps=1) for block in blocks], + key=["n"], + ) + @triton.jit + def add_one(x_ptr, out_ptr, n, BLOCK: 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 + + +def test_a_launch_none_of_whose_configs_compiled_is_unsupported(): + """Every config failed: unsupported compile-failed, whether each could + compile for another target (num_ctas=2 here: its refusal) or for none + (a failing tl.static_assert: noted), and the launch returns.""" + x, out = torch.zeros(64), torch.zeros(64) + + def grid(meta): + return (triton.cdiv(64, meta["BLOCK"]),) + + det = Sanitizer(compile=True) + kernel = _make_block_asserting_configs(64, 128) + assert tilelens.trace(det)(kernel)[grid](x, out, 64) is None + verdict = det.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "compile-failed") + assert verdict.per_config == () + first, second = verdict.notes + assert first.startswith("config {'BLOCK': 64, ") + assert second.startswith("config {'BLOCK': 128, ") + # Where the assertion is, and the error in one line. + path = kernel.fn.fn.__code__.co_filename + at = f"{path}:{_line(kernel.fn, 'tl.static_assert(')}: " + assert all( + f"was not checked: {at}it failed to compile for cuda:89 " + "(CompileTimeAssertionFailure), " in note + for note in verdict.notes + ) + assert verdict.refusal.message == ( + f"{path}:{_line(kernel.fn, 'def add_one(')}: no config of the launch " + "compiled for cuda:89, so nothing was checked: each failed with an " + "error of its own code whatever the target (see the notes)" + ) + + @triton.autotune(configs=[triton.Config({"BLOCK": 16}, num_ctas=2)], key=["n"]) + @triton.jit + def copy_two_ctas(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + + det = Sanitizer(compile=True) + assert tilelens.trace(det)(copy_two_ctas)[(4,)](x, out, 64) is None + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + (failed,) = verdict.per_config + assert verdict.refusal == failed.refusal + _assert_compile_failed(failed.refusal, "cuda:89", "num_ctas > 1") + + +def test_a_static_assert_on_the_target_is_the_targets(): + """A tl.static_assert that asks for the target (tl.target_info) fails + for this target only: the config may run on the user's GPU, so it is + unsupported, not a note; a target it compiles for checks it.""" + + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16, "HOPPER": False}), + triton.Config({"BLOCK": 16, "HOPPER": True}), + ], + key=["n"], + ) + @triton.jit + def maybe_hopper(x_ptr, n, BLOCK: tl.constexpr, HOPPER: tl.constexpr): + tl.static_assert(not HOPPER or tl.target_info.cuda_capability_geq(9, 0)) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + x = torch.zeros(64) + det = _sanitizer() + tilelens.trace(det)(maybe_hopper)[(4,)](x, 64) + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + ok, failed = verdict.per_config + assert (ok.config["HOPPER"], ok.status) == (False, "ok") + assert (failed.config["HOPPER"], failed.status) == (True, "unsupported") + _assert_compile_failed(failed.refusal, "cuda:89", "CompileTimeAssertionFailure") + + det = Sanitizer(compile=True, abort_on_error=False, target="cuda:90") + tilelens.trace(det)(maybe_hopper)[(4,)](x, 64) + assert det.last_status == "ok" + assert [c.status for c in det.last_verdict.per_config] == ["ok", "ok"] + + +def _make_device_asserting_configs(asks): + # Two configs, both calling ``asks()`` (the first compiles first); the + # second, which goes out of bounds (64 lanes, no mask, over 16 + # elements), holds only where ``asks()`` answers yes. + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16, "BIG": False}), + triton.Config({"BLOCK": 64, "BIG": True}), + ], + key=["n"], + ) + @triton.jit + def big_only(x_ptr, n, BLOCK: tl.constexpr, BIG: tl.constexpr): + yes: tl.constexpr = asks() + tl.static_assert(not BIG or yes) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0) + + return big_only + + +def test_a_static_assert_on_a_caught_device_query_is_the_devices(): + """The host compile refuses a device query; a kernel that catches that + and falls back to an answer of its own ("no big shared memory") fails + a static_assert a GPU with more might pass (an H100 would run the out + of bounds config): that config was not checked, so the launch is never + "ok" (D27), and it is no note.""" + 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 + + kernel = _make_device_asserting_configs(has_big_smem) + det = _sanitizer() + tilelens.trace(det)(kernel)[(1,)](torch.zeros(16), 16) + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + ok, failed = verdict.per_config + assert (ok.config["BIG"], ok.status) == (False, "ok") + assert (failed.config["BIG"], failed.status) == (True, "unsupported") + _assert_compile_failed(failed.refusal, "cuda:89", "CompileTimeAssertionFailure") + assert failed.refusal.loc.line == _line(kernel.fn, "tl.static_assert(") + + +def test_a_static_assert_on_a_kept_target_answer_is_the_targets(): + """A kernel's code may keep the target answer after its first compile + asked (a memo): a later config failing a static_assert on the kept + answer, without asking again, is the target's all the same. At a + target it holds for, the config is checked (and goes out of bounds).""" + 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 + + kernel = _make_device_asserting_configs(is_hopper) + det = _sanitizer() + tilelens.trace(det)(kernel)[(1,)](torch.zeros(16), 16) + assert is_hopper.fn.__defaults__[0] == {"arch": 89} + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + ok, failed = verdict.per_config + assert (ok.config["BIG"], ok.status) == (False, "ok") + assert (failed.config["BIG"], failed.status) == (True, "unsupported") + _assert_compile_failed(failed.refusal, "cuda:89", "CompileTimeAssertionFailure") + + is_hopper.fn.__defaults__[0].clear() + det = Sanitizer(compile=True, abort_on_error=False, target="cuda:90") + tilelens.trace(det)(kernel)[(1,)](torch.zeros(16), 16) + assert det.last_status == "violations" + assert {r.config["BIG"] for r in det.records} == {True} + + +# ======== calls that do not bind the kernel's parameters (D28) ========= + + +def _make_two_args(): + @triton.jit + def two_args(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + return two_args + + +# Calls of two_args(x_ptr, n, BLOCK) that do not bind: the arguments after +# x_ptr, and the keyword arguments. +_UNBOUND_CALLS = { + "missing-argument": ((), {"BLOCK": 64}), + "extra-argument": ((64, 64, 7), {}), + "misnamed-keyword": ((), {"N": 64, "BLOCK": 64}), + "repeated-keyword": ((64,), {"n": 64, "BLOCK": 64}), + # binds, but the JIT cannot key the call (compute_cache_key), on any GPU + "unhashable-constexpr": ((64,), {"BLOCK": [64]}), +} + + +class _StandInDriver: + """What JITFunction.run asks Triton's driver for before it binds a call, + on a machine with a GPU of the default IR target.""" + + def get_current_device(self): + return 0 + + def get_current_stream(self, device=None): + return 0 + + def get_current_target(self): + return GPUTarget("cuda", 89, 32) + + +def _untraced_error(kernel, args, kwargs) -> BaseException: + """What the untraced JIT raises for ``kernel[grid](*args, **kwargs)``, + for a call that does not bind: JITFunction.run's own binder, on a + stand-in driver (no GPU: the call fails before anything is compiled or + launched).""" + from triton.runtime.driver import driver + + owner = type(driver) + saved = owner.__dict__["active"] + stand_in = _StandInDriver() + owner.active = property(lambda self: stand_in) + try: + kernel.warmup(*args, grid=(1,), **kwargs) + except Exception as exc: + return exc + finally: + owner.active = saved + raise AssertionError("the untraced call bound") + + +@pytest.mark.parametrize("case", list(_UNBOUND_CALLS)) +def test_a_call_that_does_not_bind_raises_as_untraced(case, capsys): + """D28: a call that does not match the kernel's signature is a bug in + the call, whatever the target: the launch raises the untraced JIT's own + error (the same type and message) instead of reporting the launch as + not checked. Nothing is printed or recorded, and the trace goes on.""" + kernel = _make_two_args() + x = torch.zeros(64) + rest, kwargs = _UNBOUND_CALLS[case] + expected = _untraced_error(kernel, (x, *rest), kwargs) + assert type(expected) is TypeError + det = Sanitizer(compile=True) # abort_on_error=True + traced = tilelens.trace(det)(kernel) + launches = len(trace_module.launches) + + with pytest.raises(TypeError) as raised: + traced[(1,)](x, *rest, **kwargs) + + assert (type(raised.value), str(raised.value)) == (type(expected), str(expected)) + assert det.last_verdict is None and det.records == [] + assert len(trace_module.launches) == launches + assert capsys.readouterr().out == "" + # The trace is not left mid-launch: a call that binds is checked. + traced[(1,)](x, 64, BLOCK=64) + assert det.last_status == "ok" + + +def test_an_autotuned_call_that_does_not_bind_raises_as_untraced(): + """Through the autotuner too: the first config's compile raises.""" + + @triton.autotune( + configs=[triton.Config({"BLOCK": 16}), triton.Config({"BLOCK": 32})], + key=["n"], + ) + @triton.jit + def tuned(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + x = torch.zeros(64) + # The JIT's binder, as the autotuner calls it for its first config. + expected = _untraced_error(tuned.fn, (x,), {"BLOCK": 16}) + det = _sanitizer() + + with pytest.raises(TypeError) as raised: + tilelens.trace(det)(tuned)[(1,)](x) + + assert str(raised.value) == str(expected) + assert det.last_verdict is None + + +def _make_tuned(): + @triton.autotune( + configs=[triton.Config({"BLOCK": 16}), triton.Config({"BLOCK": 32})], + key=["n"], + ) + @triton.jit + def tuned(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + return tuned + + +def test_an_autotuned_call_passing_a_tuned_parameter_raises_as_untraced(): + """A call that passes an autotuned meta-parameter itself: the untraced + launch's autotuning refuses it before benchmarking (Autotuner._bench's + ValueError; no device involved), and so does a traced launch, IR-only + or mixed, not a TypeError naming a tilelens internal; the trace goes + on.""" + x = torch.zeros(64) + with pytest.raises(ValueError) as untraced: + _make_tuned()[(1,)](x, 64, BLOCK=64) + assert "Conflicting meta-parameters: BLOCK" in str(untraced.value) + det = _sanitizer() + traced = tilelens.trace(det)(_make_tuned()) + + with pytest.raises(ValueError) as raised: + traced[(1,)](x, 64, BLOCK=64) + + assert str(raised.value) == str(untraced.value) + assert det.last_verdict is None + traced[(1,)](x, 64) + assert det.last_status == "ok" + mixed = tilelens.trace(Tracer())(tilelens.trace(_sanitizer())(_make_tuned())) + with pytest.raises(ValueError) as raised: + mixed[(1,)](x, 64, BLOCK=64) + assert str(raised.value) == str(untraced.value) + assert torch.equal(x, torch.zeros(64)) + + +def test_a_mixed_trace_raises_the_untraced_error_before_interpreting(): + """D28 in a mixed trace (D4b): the call raises the untraced JIT's error + from the compile pass, before the interpreter runs, which would fail on + the same call too (with Python's own TypeError for the kernel's + function), so no output is written either way.""" + kernel = _make_two_args() + x = torch.zeros(64) + expected = _untraced_error(kernel, (x,), {"BLOCK": 64}) + traced = tilelens.trace(Tracer())(tilelens.trace(_sanitizer())(kernel)) + + with pytest.raises(TypeError) as raised: + traced[(1,)](x, BLOCK=64) + + assert str(raised.value) == str(expected) + assert torch.equal(x, torch.zeros(64)) + with pytest.raises(TypeError, match="'n'"): + tilelens.trace(Tracer())(_make_two_args())[(1,)](x, BLOCK=64) + assert torch.equal(x, torch.zeros(64)) + + +def test_an_option_the_target_does_not_know_is_a_compile_failure(): + """A keyword naming no parameter is a compile option, which the + target's backend may not know (waves_per_eu is a HIP option): no bind + failure, but the target's compile failure (D27), so the launch is + unsupported, the program goes on, and another target checks it.""" + x = torch.zeros(64) + det = _sanitizer() + call = dict(BLOCK=64, waves_per_eu=2) + + assert tilelens.trace(det)(_make_two_args())[(1,)](x, 64, **call) is None + + refusal = det.last_verdict.refusal + assert (det.last_status, refusal.kind) == ("unsupported", "compile-failed") + assert "failed to compile for cuda:89 (KeyError:" in refusal.message + assert "waves_per_eu" in refusal.message and "TILELENS_IR_TARGET" in refusal.message + hip = Sanitizer(compile=True, abort_on_error=False, target="hip:gfx942") + tilelens.trace(hip)(_make_two_args())[(1,)](x, 64, **call) + assert hip.last_status == "ok" + + +@pytest.mark.parametrize("target", ["cuda:89", "hip:gfx942"]) +def test_an_unknown_option_is_named_not_blamed_on_the_kernel(target): + """A keyword no backend knows (a misspelled option) fails for every + target: the refusal says the call passes an option the target lacks, + that a misspelled one fails on every GPU, and that only a target of a + backend with that option can check it; it does not claim the kernel may + compile for another target.""" + det = Sanitizer(compile=True, abort_on_error=False, target=target) + + assert ( + tilelens.trace(det)(_make_two_args())[(1,)]( + torch.zeros(64), 64, BLOCK=64, bogus=1 + ) + is None + ) + + refusal = det.last_verdict.refusal + assert (det.last_status, refusal.kind) == ("unsupported", "compile-failed") + assert ( + f"the call passes 'bogus', neither a parameter of the kernel nor a " + f"compile option for {target}, so it was not checked" in refusal.message + ) + assert "a misspelled option: on every GPU" in refusal.message + assert "name a target of that backend" in refusal.message + assert "a kernel can compile for one target and fail for another" not in ( + refusal.message + ) + + +_UNBOUND_SCRIPT = """\ +import torch, triton, triton.language as tl + + +@triton.jit +def two_args(x_ptr, n, BLOCK: tl.constexpr): + # {tag} + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + +x = torch.zeros(64) +print("before the launch") +two_args[(1,)]({call}) # the launch +print("launch returned") +""" + + +@pytest.mark.parametrize( + "case", ["missing-argument", "extra-argument", "misnamed-keyword"] +) +def test_the_cli_raises_a_call_that_does_not_bind(tmp_path, case): + """tile-sanitizer --compile: the script fails at the call as it does + untraced, with a non-zero exit status and the traceback, whose last + line is the untraced JIT's error.""" + rest, kwargs = _UNBOUND_CALLS[case] + expected = _untraced_error(_make_two_args(), (torch.zeros(64), *rest), kwargs) + call = ", ".join( + ["x", *map(repr, rest), *(f"{k}={v!r}" for k, v in kwargs.items())] + ) + script = tmp_path / "unbound.py" + script.write_text(_UNBOUND_SCRIPT.format(tag=script, call=call)) + cli = ( + f"import sys; sys.argv = ['tile-sanitizer', '--compile', {str(script)!r}]; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + + proc = _run(script, "-c", cli) + + assert proc.returncode == 1, proc.stderr + assert proc.stdout.splitlines() == ["before the launch"] + lines = script.read_text().splitlines() + (line,) = [i for i, text in enumerate(lines, 1) if "# the launch" in text] + assert "Traceback (most recent call last):" in proc.stderr + assert f'File "{script}", line {line}, in ' in proc.stderr + assert proc.stderr.rstrip().splitlines()[-1] == f"TypeError: {expected}" + + +def test_a_mixed_trace_still_interprets_after_a_compile_failure(): + """D4b + D27: the interpreted peer runs (and writes the outputs) while + the compiled sanitizer reports the config it could not compile.""" + det, tracer = _sanitizer(), Tracer() + traced = tilelens.trace(tracer)(tilelens.trace(det)(_make_add_nomask())) + x, out = torch.randn(64), torch.zeros(64) + + traced[(1,)](x, out, 64, BLOCK=64, num_ctas=2) + + torch.testing.assert_close(out, x) + verdict = det.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "compile-failed") + assert any(isinstance(r, Store) for r in trace_module.launches[-1].records) + + +def test_a_hip_target_compiles_with_its_own_backend(): + det = Sanitizer(compile=True, abort_on_error=False, target="hip:gfx942") + x, out = torch.zeros(3000), torch.zeros(3000) + tilelens.trace(det)(_make_add_nomask())[(3,)](x, out, 3000, BLOCK=1024) + assert det.last_status == "violations" + assert {r.tensor_name for r in det.records} == {"x_ptr", "out_ptr"} + + +# The target-dependent branches of a kernel are the IR target's (D26). + + +def _make_unmasked_on_cuda(): + @triton.jit + def unmasked_on_cuda(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if tl.target_info.is_cuda(): + tl.store(x_ptr + offs, 5.0) # what every CUDA device runs + else: + tl.store(x_ptr + offs, 6.0, mask=offs < n) + + return unmasked_on_cuda + + +def _make_unmasked_from_sm89(): + @triton.jit + def unmasked_from_sm89(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if tl.target_info.cuda_capability_geq(8, 9): + tl.store(x_ptr + offs, 1.0) + else: + tl.store(x_ptr + offs, 2.0, mask=offs < n) + + return unmasked_from_sm89 + + +def _make_unmasked_on_hip(): + @triton.jit + def unmasked_on_hip(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if tl.target_info.is_hip(): + tl.store(x_ptr + offs, 3.0) + else: + tl.store(x_ptr + offs, 4.0, mask=offs < n) + + return unmasked_on_hip + + +# (kernel, target, status): 128 lanes over 64 elements, so the unmasked +# branch goes out of bounds. Without a GPU, tl.target_info used to read +# "no target": the default target's CUDA branch was never checked (a false +# "ok"). +TARGET_BRANCHES = [ + (_make_unmasked_on_cuda, None, "violations"), + (_make_unmasked_on_cuda, "hip:gfx942", "ok"), + (_make_unmasked_from_sm89, None, "violations"), + (_make_unmasked_from_sm89, "cuda:80", "ok"), + (_make_unmasked_from_sm89, "cuda:89", "violations"), + (_make_unmasked_from_sm89, "cuda:90", "violations"), + (_make_unmasked_on_hip, None, "ok"), + (_make_unmasked_on_hip, "hip:gfx942", "violations"), +] + + +def _check_branches(make, target): + det = Sanitizer(compile=True, abort_on_error=False, target=target) + tilelens.trace(det)(make())[(8,)](torch.zeros(64), 64, BLOCK=16) + return det.last_status, sorted( + (r.kind, r.tensor_name, r.violation_offset, tuple(sorted(r.witness.items()))) + for r in det.records + ) + + +@pytest.mark.parametrize( + "make, target, status", + TARGET_BRANCHES, + ids=[f"{m.__name__[6:]}-{t or 'default'}" for m, t, _ in TARGET_BRANCHES], +) +def test_target_dependent_branches_are_the_ir_targets(make, target, status): + assert _check_branches(make, target)[0] == status + + +def test_sanitizers_for_two_targets_share_one_trace(): + """A second compiled sanitizer, for another target, is not dropped from + the trace: one launch is checked for each target, each with its own + verdict.""" + sm80 = Sanitizer(compile=True, abort_on_error=False, target="cuda:80") + sm90 = Sanitizer(compile=True, abort_on_error=False, target="cuda:90") + traced = tilelens.trace(sm90)(tilelens.trace(sm80)(_make_unmasked_from_sm89())) + assert traced.client_manager.ir_clients() == [sm80, sm90] + + traced[(8,)](torch.zeros(64), 64, BLOCK=16) + + assert (sm80.last_status, sm80.records) == ("ok", []) + assert sm90.last_status == "violations" and sm90.records + assert trace_module.launches[-1].records == [ + sm80.last_verdict, + *sm90.records, + sm90.last_verdict, + ] + + +class _Machine: + """A stand-in for Triton's active driver on a machine with a GPU of + ``target``; it must never be asked during IR mode's compile.""" + + def __init__(self, target): + self.target = target + + def get_current_target(self): + raise AssertionError("IR mode asked the machine's driver for its target") + + +@pytest.mark.parametrize( + "machine", + [ + GPUTarget("cuda", 89, 32), + GPUTarget("cuda", 90, 32), + GPUTarget("hip", "gfx942", 64), + ], + ids=["sm89", "sm90", "gfx942"], +) +def test_a_verdict_does_not_depend_on_the_machine(monkeypatch, machine): + """Every verdict and finding above is the same on a GPU of any kind as + without one (this module's default: the driver is unreachable).""" + without_gpu = [_check_branches(make, target) for make, target, _ in TARGET_BRANCHES] + from triton.runtime.driver import driver + + stand_in = _Machine(machine) + monkeypatch.setattr(type(driver), "active", property(lambda self: stand_in)) + on_gpu = [_check_branches(make, target) for make, target, _ in TARGET_BRANCHES] + assert on_gpu == without_gpu + assert [status for status, _ in on_gpu] == [s for _, _, s in TARGET_BRANCHES] + + +def _make_to_fp8e4nv(): + @triton.jit + def to_fp8e4nv(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + x = tl.load(x_ptr + offs, mask=mask) + tl.store(out_ptr + offs, x.to(tl.float8e4nv).to(tl.float32), mask=mask) + + return to_fp8e4nv + + +def _make_load_fp8e4nv(): + @triton.jit + def load_fp8e4nv(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).to(tl.float32), mask=mask + ) + + return load_fp8e4nv + + +FP8E4NV_KERNELS = pytest.mark.parametrize( + "make, dtype", + [(_make_to_fp8e4nv, torch.float32), (_make_load_fp8e4nv, torch.float8_e4m3fn)], + ids=["cast", "tensor"], +) + + +@FP8E4NV_KERNELS +def test_fp8e4nv_compiles_under_the_default_target(make, dtype): + """D26 amended: the default target, cuda:89, is the first with fp8e4nv, + so such kernels are checked without naming a target.""" + x, out = torch.zeros(64).to(dtype), torch.zeros(64) + det = _sanitizer() + kernel = tilelens.trace(det)(make())[(4,)](x, out, 64, BLOCK=16) + assert kernel.target == GPUTarget("cuda", 89, 32) + assert det.last_status == "ok", det.last_verdict + + +@FP8E4NV_KERNELS +def test_an_explicit_cuda80_refuses_fp8e4nv_as_unsupported(make, dtype): + """cuda:80 has no fp8e4nv: the launch fails to compile there (whatever + this machine's GPU), which is unsupported compile-failed naming cuda:80 + and how to name another target, never an exception (D27).""" + x, out = torch.zeros(64).to(dtype), torch.zeros(64) + det = Sanitizer(compile=True, target="cuda:80") # abort_on_error=True + assert tilelens.trace(det)(make())[(4,)](x, out, 64, BLOCK=16) is None + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + (failed,) = verdict.per_config + assert verdict.refusal == failed.refusal + _assert_compile_failed(failed.refusal, "cuda:80", "fp8e4nv not supported") + + +def _device_is_zero(): + # A host function a kernel's constexpr function calls (marked like + # tl.target_info.current_target), asking Triton's driver for a device. + return triton.runtime.driver.active.get_current_device() == 0 + + +_device_is_zero.__triton_builtin__ = True # type: ignore[attr-defined] + + +def test_a_host_compile_that_cannot_run_is_unsupported_not_raised(): + """A compile that asks for a device cannot run on the host: no error of + the kernel's (a GPU would compile it), so the launch is unsupported and + goes on, never failed.""" + from triton.runtime.jit import constexpr_function + + @constexpr_function + def on_device_zero(): + return _device_is_zero() + + @triton.jit + def device_dependent(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if on_device_zero(): + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + det = _sanitizer() + tilelens.trace(det)(device_dependent)[(4,)](torch.zeros(64), 64, BLOCK=16) + verdict = det.last_verdict + assert (verdict.status, verdict.refusal.kind) == ( + "unsupported", + "host-compile-unavailable", + ) + assert "'get_current_device'" in verdict.refusal.message + assert verdict.notes == () + + +@pytest.fixture +def configured_target(monkeypatch): + """Set TILELENS_IR_TARGET and reload the process config from it.""" + + def set_target(value): + monkeypatch.setenv("TILELENS_IR_TARGET", value) + monkeypatch.setattr(config_module, "config", Config()) + + return set_target + + +def test_the_environment_sets_the_default_target(configured_target): + configured_target("cuda:90") + x, out = torch.zeros(64), torch.zeros(64) + det = _sanitizer() + tilelens.trace(det)(_make_two_cta_configs())[(4,)](x, out, 64) + assert [c.config["num_ctas"] for c in det.last_verdict.per_config] == [1, 2] + + # A target the client names wins over the environment's. + det = Sanitizer(compile=True, abort_on_error=False, target="cuda:80") + tilelens.trace(det)(_make_two_cta_configs())[(4,)](x, out, 64) + assert det.last_status == "unsupported" + _assert_compile_failed(det.last_verdict.refusal, "cuda:80", "num_ctas > 1") + + +@pytest.mark.parametrize("spec", ["sm90", "cuda:", "hip:942", "cuda:90:x", 90]) +def test_a_spec_that_names_no_target_is_refused_up_front(spec): + with pytest.raises(ValueError, match=f"invalid IR target {spec!r}"): + Sanitizer(compile=True, target=spec) + + +def test_a_configured_spec_that_names_no_target_fails_the_launch(configured_target): + configured_target("gfx942") + det = _sanitizer() + x, out = torch.zeros(64), torch.zeros(64) + with pytest.raises(ValueError, match=r"TILELENS_IR_TARGET\) is 'gfx942'"): + tilelens.trace(det)(_make_add_nomask())[(1,)](x, out, 64, BLOCK=64) + assert det.last_verdict is None + + +_TARGET_SCRIPT = """\ +import torch, triton, triton.language as tl + +@triton.autotune( + configs=[ + triton.Config({{"BLOCK": 16}}, num_ctas=1), + triton.Config({{"BLOCK": 16}}, num_ctas=2), + ], + key=["n"], +) +@triton.jit +def copy_ctas(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + # {tag} + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + +x, out = torch.zeros(64), torch.zeros(64) +copy_ctas[(4,)](x, out, 64) +print("launch returned") +""" + + +@pytest.mark.parametrize( + "target, returncode, output", + [ + ("cuda:90", 0, "launch returned"), + # D27: the sm90 config is reported as not checked, and the script + # goes on. + ( + None, + 0, + ": it failed to compile for cuda:89 (ValueError: " + "num_ctas > 1 requires NVIDIA SM90", + ), + ("sm90", 1, "TILELENS_IR_TARGET) is 'sm90'"), + ], +) +def test_the_cli_reads_the_target_from_the_environment( + tmp_path, target, returncode, output +): + script = tmp_path / "ctas.py" + script.write_text(_TARGET_SCRIPT.format(tag=script)) + cli = ( + f"import sys; sys.argv = ['tile-sanitizer', '--compile', {str(script)!r}]; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + env = {} if target is None else {"TILELENS_IR_TARGET": target} + proc = _run(script, "-c", cli, **env) + assert proc.returncode == returncode, proc.stderr + assert output in proc.stdout + proc.stderr + assert ("launch returned" in proc.stdout) == (returncode == 0) + + +_UNCOMPILABLE_SCRIPT = """\ +import torch, triton, triton.language as tl + +@triton.jit +def copy(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + # {tag} + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=offs < n), mask=offs < n) + +@triton.autotune(configs=[triton.Config({{"BLOCK": 64}})], key=["n"]) +@triton.jit +def bounded(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + tl.static_assert(BLOCK <= 32) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=offs < n), mask=offs < n) + +x, out = torch.zeros(64), torch.zeros(64) +copy[(1,)](x, out, 64, BLOCK=64, num_ctas=2) +print("first launch returned") +bounded[(1,)](x, out, 64) +print("second launch returned") +""" + + +def test_the_cli_goes_on_past_kernels_that_fail_to_compile(tmp_path): + """tile-sanitizer --compile on kernels no config of which compiles for + the target (num_ctas=2 below sm90; a failing tl.static_assert): no + finding, so exit status 0, each launch reported as not checked, in one + line naming where the kernel failed. A launch whose configs all failed + as notes prints the notes too: they say why.""" + script = tmp_path / "uncompilable.py" + script.write_text(_UNCOMPILABLE_SCRIPT.format(tag=script)) + cli = ( + f"import sys; sys.argv = ['tile-sanitizer', '--compile', {str(script)!r}]; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + proc = _run(script, "-c", cli) + assert proc.returncode == 0, proc.stderr + lines = proc.stdout.splitlines() + assert "first launch returned" in lines and "second launch returned" in lines + source = script.read_text().splitlines() + + def at(needle): + (line,) = [i for i, text in enumerate(source, 1) if needle in text] + return f"{script}:{line}: " + + printed = [line for line in lines if line.startswith("[CompiledSanitizer]")] + first, *unchecked = printed + assert first.startswith( + "[CompiledSanitizer] not checked: compile-failed: " + f"{at('def copy(')}it failed to compile for cuda:89 (ValueError: num_ctas " + "> 1 requires NVIDIA SM90" + ) + assert unchecked == [ + "[CompiledSanitizer] not checked: compile-failed: " + f"{at('def bounded(')}no config of the launch compiled for cuda:89, so " + "nothing was checked: each failed with an error of its own code whatever " + "the target (see the notes)", + "[CompiledSanitizer] note: config {'BLOCK': 64, 'num_warps': 4, " + "'num_ctas': 1, 'num_stages': 3} was not checked: " + f"{at('tl.static_assert(')}it failed to compile for cuda:89 " + "(CompileTimeAssertionFailure), an error of its own code whatever the " + "target, so it never launches", + ] + assert "Traceback" not in proc.stderr diff --git a/tests/end_to_end/test_core.py b/tests/end_to_end/test_core.py index b98b77714..94eb7f0bd 100644 --- a/tests/end_to_end/test_core.py +++ b/tests/end_to_end/test_core.py @@ -16,19 +16,18 @@ def test_trace_decorator_add_clients(): Test goal: 1. Apply @trace("sanitizer") and @trace("profiler") to add the Sanitizer and Profiler clients. 2. Apply @trace("tracer") to append a Tracer client. - 3. Apply @trace(("sanitizer",)) with a duplicate Sanitizer, which should be - ignored by the de-duplication logic. + 3. Apply @trace("sanitizer") over a Sanitizer instance: the name asks for a + default Sanitizer, which the one already in the trace serves. The final Trace object should contain exactly one instance each of - Sanitizer, Profiler, and Tracer (total = 3 clients). + Sanitizer, Profiler, and Tracer (total = 3 clients). A second Sanitizer + instance, whose settings would be lost, is refused. """ @tilelens.trace("sanitizer") @tilelens.trace("profiler") @tilelens.trace("tracer") - @tilelens.trace( - Sanitizer(abort_on_error=True) - ) # Duplicate Sanitizer (should be ignored) + @tilelens.trace(Sanitizer(abort_on_error=False)) @triton.jit def my_kernel(x_ptr, y_ptr, out_ptr, BLOCK_SIZE: tl.constexpr): pid = tl.program_id(0) @@ -41,11 +40,14 @@ def my_kernel(x_ptr, y_ptr, out_ptr, BLOCK_SIZE: tl.constexpr): assert isinstance(my_kernel, TritonTrace) # Verify client de-duplication and addition logic - clients = my_kernel.client_manager.clients - assert len(clients) == 3 - assert sum(c == "sanitizer" for c in clients) == 1 - assert sum(c == "profiler" for c in clients) == 1 - assert sum(c == "tracer" for c in clients) == 1 + names = [c.NAME for c in my_kernel.client_manager.clients] + assert sorted(names) == ["profiler", "sanitizer", "tracer"] + # The instance's own settings were kept, not a default's. + assert my_kernel.client_manager.get_client("sanitizer").abort_on_error is False + + with pytest.raises(ValueError, match="interpreting client named 'sanitizer'"): + tilelens.trace(Sanitizer(abort_on_error=True))(my_kernel) + assert len(my_kernel.client_manager.clients) == 3 def test_trace_decorator_supports_gluon_frontend(): diff --git a/tests/end_to_end/test_host_compile.py b/tests/end_to_end/test_host_compile.py new file mode 100644 index 000000000..d24caa170 --- /dev/null +++ b/tests/end_to_end/test_host_compile.py @@ -0,0 +1,377 @@ +"""Host compile vs the JIT's own compile (D25), for one target, on the CPU. + +IR mode compiles on the host for its target instead of through +``JITFunction.run``. The JIT compiles a launch for whatever Triton's active +driver says the device is, and a stand-in driver (a device id, a stream and +a target, nothing more) lets it compile without a GPU, for any target: the +oracle for what HostCompiler lifts from ``JITFunction.run`` and +``triton.compile`` (the binder call, the options the JIT adds, the +truncated compile's hash key), wherever these tests run. + +For every launch below and each of cuda:80, cuda:90 and hip:gfx942, the +JIT's kernel and the host's must be one kernel: the same hash (Triton's +name for the specialization), the same TTIR text, the same text for every +deeper stage, and the same compiled-sanitizer verdict, down to each +finding's witness. The launches cover the JIT's integer typing at the +i32/i64 boundary (2**31 - 1, 2**31, -2**31, 2**32), the equal-to-1 +specialization, bool and float arguments, a tensor descriptor and the +front end's target queries (tl.target_info). +""" + +from __future__ import annotations + +import pytest +import torch +import triton +import triton.language as tl +from triton.backends.compiler import GPUTarget + +import tilelens +from tilelens.clients import Sanitizer +from tilelens.core.client import ClientManager, LaunchCall +from tilelens.core.host_compile import HostCompiler, parse_ir_target +from tilelens.ir.ttir_reader import UnsupportedTTIR, parse_ttir + + +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) + + +pytestmark = pytest.mark.skipif( + not _real_compiles_available(), + reason="Triton was imported under TRITON_INTERPRET=1: nothing compiles in-process", +) + +CUDA80 = GPUTarget("cuda", 80, 32) +# IR target -> the stand-in's device id: JITFunction.device_caches keeps a +# target per device, so every target gets a device of its own. +TARGETS = {"cuda:80": 0, "cuda:90": 1, "hip:gfx942": 2} + + +@pytest.fixture(autouse=True) +def _real_jit(monkeypatch, tmp_path_factory): + # tests/unit/test_multithreading.py sets TRITON_INTERPRET=1 at import + # time; pin the knob off so @triton.jit builds real JITFunctions. The + # JIT's compiles go to a cache of this module's own: its kernels are + # compiled here, never read back from another checkout's cache. + from triton import knobs + + monkeypatch.delenv("TRITON_INTERPRET", raising=False) + monkeypatch.setenv( + "TRITON_CACHE_DIR", str(tmp_path_factory.getbasetemp() / "jit-cache") + ) + 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 _StandInDriver: + """What JITFunction.run asks Triton's active driver for, on a machine + whose device ``device`` is a GPU of ``target``.""" + + def __init__(self, target, device): + self.target, self.device = target, device + + def get_current_device(self): + return self.device + + def get_current_stream(self, device=None): + return 0 + + def get_current_target(self): + return self.target + + +@pytest.fixture(params=list(TARGETS)) +def target(request, monkeypatch): + """The IR target, and a stand-in driver for a device of it.""" + from triton.runtime.driver import driver + + target = parse_ir_target(request.param) + stand_in = _StandInDriver(target, TARGETS[request.param]) + monkeypatch.setattr(type(driver), "active", property(lambda self: stand_in)) + return target + + +def _masked_copy(): + @triton.jit + def masked_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 masked_copy + + +def _strided_store(): + @triton.jit + def strided_store(x_ptr, S, BLOCK: tl.constexpr): + off = tl.program_id(0) * S # an i32 product for an i32 S + tl.store(x_ptr + off + tl.arange(0, BLOCK), 1.0) + + return strided_store + + +def _flagged_scale(): + @triton.jit + def flagged_scale(x_ptr, out_ptr, flag, scale, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if flag: + tl.store(out_ptr + offs, tl.load(x_ptr + offs) * scale) # unmasked + + return flagged_scale + + +def _target_branches(): + @triton.jit + def target_branches(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if tl.target_info.cuda_capability_geq(9, 0): + tl.store(x_ptr + offs, 1.0) # unmasked from sm90 + elif tl.target_info.is_hip(): + tl.store(x_ptr + offs, 2.0) # unmasked on HIP + else: + tl.store(x_ptr + offs, 3.0, mask=offs < n) + + return target_branches + + +def _descriptor_bump(): + @triton.jit + def descriptor_bump(desc, BLOCK: tl.constexpr): + desc.store([0, 0], desc.load([0, 0]) + 1) + + return descriptor_bump + + +def _descriptor(): + from triton.tools.tensor_descriptor import TensorDescriptor + + return (TensorDescriptor.from_tensor(torch.zeros(64, 64), [16, 16]),) + + +BOUNDARY = [2**31 - 1, 2**31, -(2**31), 2**32, 1, 0, 64] + + +def _cases(): + """(id, kernel factory, args builder, kwargs, grid).""" + cases = [] + for n in BOUNDARY: + cases.append( + ( + f"masked_copy-n={n}", + _masked_copy, + lambda n=n: (torch.zeros(64), torch.zeros(64), n), + {"BLOCK": 16}, + (8,), # 128 lanes over 64 elements: in bounds iff n <= 64 + ) + ) + cases.append( + ( + f"strided_store-S={n}", + _strided_store, + lambda n=n: (torch.zeros(64), n), + {"BLOCK": 16}, + (3,), + ) + ) + for flag in (True, False): + for scale in (1.5, -0.0): + cases.append( + ( + f"flagged_scale-flag={flag}-scale={scale}", + _flagged_scale, + lambda flag=flag, scale=scale: ( + torch.zeros(64), + torch.zeros(64), + flag, + scale, + ), + {"BLOCK": 16}, + (8,), # 128 lanes over 64 elements: out of bounds if flag + ) + ) + cases.append( + ( + "target_branches", + _target_branches, + lambda: (torch.zeros(64), 64), + {"BLOCK": 16}, + (8,), # out of bounds on sm90+ and HIP only + ) + ) + cases.append( + ("descriptor_bump", _descriptor_bump, _descriptor, {"BLOCK": 16}, (1,)) + ) + return cases + + +CASES = _cases() + + +def _reading(text: str): + try: + return parse_ttir(text) + except UnsupportedTTIR as e: + return ("refused", e.kind, str(e.loc)) + + +def _verdict_from(kernel, jit_fn, args, kwargs, grid): + """The compiled sanitizer's verdict and findings for one launch whose + compiled kernel is ``kernel``, delivered as the core delivers it.""" + san = Sanitizer(compile=True, abort_on_error=False) + manager = ClientManager([san]) + manager.begin_launch( + LaunchCall(jit_fn=jit_fn, args=args, kwargs=kwargs, grid=grid, capture=True) + ) + event = ClientManager._launch_event( + jit_fn, args, kwargs, grid, kernel, False, target=kernel.metadata.target + ) + manager._dispatch_ir("before_launch", event, manager.ir_clients()) + manager.finalize() + return san.last_verdict, san.records + + +def _summary(verdict, records): + return ( + verdict.status, + verdict.scope, + None if verdict.refusal is None else verdict.refusal.kind, + tuple((c.specialization, c.status, c.n_reports) for c in verdict.per_config), + sorted( + ( + r.kind, + r.op_type.__name__, + r.tensor_name, + r.violation_offset, + tuple(sorted(r.witness.items())), + ) + for r in records + ), + ) + + +@pytest.mark.parametrize( + "make, build, kwargs, grid", [c[1:] for c in CASES], ids=[c[0] for c in CASES] +) +def test_the_jit_and_the_host_compile_the_same_kernel( + target, make, build, kwargs, grid +): + kernel = make() + args = build() + jit = kernel.warmup(*args, grid=grid, **kwargs) + host = HostCompiler().compile(kernel, args, kwargs, target=target, stages={"ttir"}) + + assert jit.metadata.target == host.target == target + assert host.hash == jit.hash + assert host.asm["ttir"] == jit.asm["ttir"] + from_host = _verdict_from(host, kernel, args, kwargs, grid) + assert _summary(*from_host) == _summary( + *_verdict_from(jit, kernel, args, kwargs, grid) + ) + + # A traced launch (the host path end to end, which never asks the + # driver) reaches that verdict too. + san = Sanitizer(compile=True, abort_on_error=False, target=target) + tilelens.trace(san)(kernel)[grid](*args, **kwargs) + assert _summary(san.last_verdict, san.records) == _summary(*from_host) + + +def test_every_host_stage_is_the_jits(target): + """Past TTIR too: the host compile through the stage before the binary + holds the JIT's text for every stage, and the binary itself is + triton.compile's (the same hash).""" + kernel = _masked_copy() + args, kwargs = (torch.zeros(64), torch.zeros(64), 64), {"BLOCK": 16} + jit = kernel.warmup(*args, grid=(4,), **kwargs) + compiler = HostCompiler() + *stages, binary = [s for s in jit.asm if s != "source"] + host = compiler.compile(kernel, args, kwargs, target=target, stages={stages[-1]}) + assert list(host.asm) == stages + for stage in stages: + assert host.asm[stage] == jit.asm[stage], stage + full = compiler.compile(kernel, args, kwargs, target=target, stages={binary}) + assert host.hash == full.hash == jit.hash + assert full.asm[binary] == jit.asm[binary] + + +def test_the_cases_are_not_vacuous(): + """The launches exercise both verdicts, every finding kind they can, + both integer widths, and branches the target decides.""" + statuses, kinds, widths = set(), set(), set() + for _, make, build, kwargs, grid in CASES: + kernel = make() + args = build() + host = HostCompiler().compile( + kernel, args, kwargs, target=CUDA80, stages={"ttir"} + ) + graph = _reading(host.asm["ttir"]) + if not isinstance(graph, tuple): + widths |= {a.int_bits for a in graph.func_args if a.int_bits} + verdict, records = _verdict_from(host, kernel, args, kwargs, grid) + statuses.add(verdict.status) + kinds |= {r.kind for r in records} + assert {"ok", "violations"} <= statuses + assert kinds == {"out-of-bounds", "integer-overflow"} + assert {32, 64} <= widths + + by_target = { + spec: HostCompiler() + .compile( + _target_branches(), + (torch.zeros(64), 64), + {"BLOCK": 16}, + target=parse_ir_target(spec), + ) + .asm["ttir"] + for spec in TARGETS + } + assert len(set(by_target.values())) == len(TARGETS) + + +class _PipelineHook: + """A custom pipeline (``knobs.runtime.add_stages_inspection_hook``) in + both of its calling conventions: called with no arguments (Triton 3.8's + JITFunction.run and triton.compile) it names the pipeline, a (key, hash) + pair; called by a backend's add_stages it leaves the stages as they + are.""" + + def __init__(self, name: str) -> None: + self.name = name + + def __call__(self, *args): + if not args: + return (f"-pipeline-{self.name}", f"{self.name}0") + return None + + +def test_under_a_custom_pipeline_the_host_compiles_the_jits_kernel(target, monkeypatch): + """Triton 3.8's JIT keys a kernel by a custom pipeline too (the + specialization and triton.compile's cache key): the host compile names + it alike, through TTIR and through the whole pipeline, for each + pipeline.""" + from triton import knobs + + for name in ("one", "two"): + monkeypatch.setattr( + knobs.runtime, "add_stages_inspection_hook", _PipelineHook(name) + ) + kernel = _masked_copy() + args, kwargs = (torch.zeros(64), torch.zeros(64), 64), {"BLOCK": 16} + jit = kernel.warmup(*args, grid=(4,), **kwargs) + compiler = HostCompiler() + host = compiler.compile(kernel, args, kwargs, target=target, stages={"ttir"}) + binary = [s for s in jit.asm if s != "source"][-1] + full = compiler.compile(kernel, args, kwargs, target=target, stages={binary}) + assert host.hash == full.hash == jit.hash + assert host.asm["ttir"] == jit.asm["ttir"] diff --git a/tests/end_to_end/test_ir_client.py b/tests/end_to_end/test_ir_client.py new file mode 100644 index 000000000..9253879ae --- /dev/null +++ b/tests/end_to_end/test_ir_client.py @@ -0,0 +1,301 @@ +"""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 (D25, D26: 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, Config +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 (D25). + 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 (D25): 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 (D26), 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 on_refusal(self, refusal): + return IRVerdict(self.NAME, "unsupported", refusal=refusal) + + +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) + + +@pytest.fixture +def untested_triton(monkeypatch): + monkeypatch.setattr(triton, "__version__", "3.5.0") + + +@pytest.fixture +def allow_untested_triton(monkeypatch): + # The D10b gate reads the process config, which reads the environment. + monkeypatch.setenv("TILELENS_IR_ALLOW_UNTESTED_TRITON", "1") + monkeypatch.setattr(config_module, "config", Config()) + + +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_the_env_override_runs_an_untested_triton( + untested_triton, allow_untested_triton +): + """The override runs IR mode on the (pretended) untested release, host + compile included. The gate's refusals are checked on every release in + tests/unit/test_ir_version_gate.py.""" + reader = _StubReader() + ir = _ToyIR(reader) + traced = tilelens.trace(ir)(_make_add_one()) + x, out = _inputs() + + traced[(4,)](x, out, 64, BLOCK=16) + + assert [c.status for c in ir.last_verdict.per_config] == ["parsed"] + assert len(reader.texts) == 1 + + +def test_the_real_reader_through_the_parse_cache(): + ir = _ToyIR() # ParseCache's default reader: tilelens.ir.ttir_reader.parse_ttir + traced = tilelens.trace(ir)(_make_add_one()) + x, out = _inputs() + + traced[(4,)](x, out, 64, BLOCK=16) + traced[(4,)](x, out, 64, BLOCK=16) + + first, second = ir.outcomes + assert (first.error, first.refusal) == (None, None) + assert first.graph.kernel_name == "add_one" + # The second launch is a cache hit. + assert second is first + (config,) = ir.last_verdict.per_config + assert config.status == "parsed" diff --git a/tests/end_to_end/test_ir_lifecycle_compiled.py b/tests/end_to_end/test_ir_lifecycle_compiled.py new file mode 100644 index 000000000..69f210d6f --- /dev/null +++ b/tests/end_to_end/test_ir_lifecycle_compiled.py @@ -0,0 +1,1090 @@ +"""End-to-end tests of the core IR lifecycle on real kernels: IR clients receive +kernels compiled on the host (D25) through ClientManager.ir_capture, with or +without the real launch, across plain, @heuristics and @autotune∘@heuristics +kernels. Counterparts on a fake compile live in tests/unit/test_ir_lifecycle.py. + +Only what needs a device runs on one: a real launch (an IR client declaring +LAUNCH="run"), device-memory accounting, and an interpreting client's voted +warmup (the JIT's own compile). Everything else runs on CPU tensors with +Triton's driver unreachable (``_no_driver``), as on a machine without a GPU. +""" + +import importlib +import types + +import pytest +import torch +import triton +import triton.language as tl +from triton.backends.compiler import GPUTarget +from triton.compiler import CompiledKernel +from triton.compiler.errors import CompilationError, CompileTimeAssertionFailure + +import tilelens +from tilelens.core.callbacks import ForLoopCallbacks, OpCallbacks +from tilelens.core.client import Client +from tilelens.core.config import DEFAULT_IR_TARGET +from tilelens.core.data import Store +from tilelens.core.host_compile import HostCompiler, HostKernel + +# `tilelens.core.trace` the attribute is the trace() decorator; the module +# holds the `launches` list. +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. + 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", +) + +GPU_REASON = "a real launch (or the JIT's own compile) needs a CUDA GPU" +# Marks what needs a device; everything else runs without Triton's driver. +needs_gpu = pytest.mark.skipif(not torch.cuda.is_available(), reason=GPU_REASON) +# The default IR target (D26, amended), and one passed explicitly. +CUDA89 = GPUTarget("cuda", 89, 32) +CUDA80 = GPUTarget("cuda", 80, 32) + + +@pytest.fixture(autouse=True) +def _no_driver(request, unreachable_driver): + """Unless the test is marked needs_gpu: Triton's driver is unreachable, + as on a machine without a GPU, so any driver query on the IR path fails + the test (D25).""" + if any( + mark.kwargs.get("reason") == GPU_REASON + for mark in request.node.iter_markers("skipif") + ): + return + unreachable_driver("the IR path queried Triton's driver") + + +@pytest.fixture(autouse=True) +def _default_ir_target(monkeypatch): + """The default IR target (D26), 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 _IRClient(Client): + NEEDS_INTERPRETER = False + IR_STAGES = frozenset({"ttir"}) + + def __init__(self): + super().__init__() + self.log: list = [] + self.events: list = [] + self.failures: list = [] + self.finalized: list = [] + + def begin_launch(self, call): + self.log.append("begin") + self.events = [] + self.failures = [] + + def abort_launch(self, exc): + self.log.append("abort") + + def before_launch(self, event): + self.log.append("before") + self.events.append(event) + + def after_launch(self, event): + self.log.append("after") + + def compile_failed(self, event): + self.log.append("compile_failed") + self.failures.append(event) + + def finalize(self): + self.log.append("finalize") + self.finalized.append(list(self.events)) + return [] + + def pre_warmup_callback(self, jit_fn, *args, **kwargs): + return False + + def post_warmup_callback(self, jit_fn, ret): + pass + + 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" + LAUNCH = "skip" + + +class _RunIRClient(_IRClient): + NAME = "ir_run" + LAUNCH = "run" + + +class _EagerCounter(Client): + """Interpreting client counting stores; optionally votes for a warmup.""" + + NAME = "eager_counter" + + def __init__(self, warmup_vote=False): + super().__init__() + self.stores = 0 + self.warmup_vote = warmup_vote + self.warmups: list = [] + + def _on_store(self, *args, **kwargs): + self.stores += 1 + + 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): + if op_type is Store: + return OpCallbacks(before_callback=self._on_store) + return OpCallbacks() + + def register_for_loop_callback(self): + return ForLoopCallbacks() + + def finalize(self): + return [] + + def pre_warmup_callback(self, jit_fn, *args, **kwargs): + return self.warmup_vote + + def post_warmup_callback(self, jit_fn, ret): + self.warmups.append(ret) + + +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(**autotune_kwargs): + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_warps=1), + triton.Config({"BLOCK": 32}, num_warps=2), + ], + 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) + if EVEN: + tl.store(out_ptr + offs, tl.load(x_ptr + offs) + 1) + else: + 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 _grid64(meta): + # The interpreter hands grid callables tensor-converted runtime args, so + # interpreted launches may only read constexprs here. + return (triton.cdiv(64, meta["BLOCK"]),) + + +def _device(ir_cls) -> str: + # A real launch needs device tensors; a compile does not. + return "cuda" if ir_cls.LAUNCH == "run" else "cpu" + + +def _inputs(n=64, device="cpu"): + x = torch.arange(n, dtype=torch.float32, device=device) + return x, torch.zeros_like(x) + + +def _synchronize(): + if torch.cuda.is_available(): + torch.cuda.synchronize() + + +# Each IR client class, the one that launches for real only with a GPU. +IR_CLASSES = [_SkipIRClient, pytest.param(_RunIRClient, marks=needs_gpu)] + + +def test_ir_only_skip_compiles_on_the_host_without_launching(): + ir = _SkipIRClient() + kernel = _make_add_one() + hooks = [] + kernel.add_pre_run_hook(lambda *args, **kwargs: hooks.append(1)) + traced = tilelens.trace(ir)(kernel) + x, out = _inputs() + + ret = traced[(4,)](x, out, 64, BLOCK=16) + + assert torch.equal(out, torch.zeros_like(x)) + assert ir.log == ["begin", "before", "after", "finalize"] + (events,) = ir.finalized + (event,) = events + assert event.launched is False + # Compiled on the host for the default target, through TTIR only. + assert isinstance(event.kernel, HostKernel) + assert event.target == event.kernel.target == CUDA89 + assert list(event.kernel.asm) == ["ttir"] + assert "tt.func" in event.kernel.asm["ttir"] + assert event.resolved_grid == (4, 1, 1) + assert event.specialization == event.kernel.hash + assert "run" not in vars(traced.jit_fn) + # JITFunction.run was never entered: no pre_run_hook fired. + assert hooks == [] + # One config: the launch returns its kernel, as the untraced launch + # would, and the Launch carries the grid, but no tensor (D23). + assert ret is event.kernel + launch = trace_module.launches[-1] + assert launch.grid == (4, 1, 1) + assert not launch.tensors + + +def test_a_relaunch_reuses_the_traces_host_compile(monkeypatch): + compiles = [] + compile = HostCompiler._compile_source + + def counting(*args, **kwargs): + compiles.append(1) + return compile(*args, **kwargs) + + monkeypatch.setattr(HostCompiler, "_compile_source", staticmethod(counting)) + ir = _SkipIRClient() + traced = tilelens.trace(ir)(_make_add_one()) + x, out = _inputs() + + first = traced[(4,)](x, out, 64, BLOCK=16) + again = traced[(4,)](torch.zeros(64), torch.zeros(64), 64, BLOCK=16) + other = traced[(4,)](x, out, 63, BLOCK=16) # 63: not divisible by 16 + + assert again is first and other is not first + assert len(compiles) == 2 + assert [len(events) for events in ir.finalized] == [1, 1, 1] + assert ir.finalized[0][0].specialization == ir.finalized[1][0].specialization + + +@needs_gpu +def test_ir_only_run_launches_the_real_kernel(): + ir = _RunIRClient() + kernel = _make_add_one() + hooks = [] + kernel.add_pre_run_hook(lambda *args, **kwargs: hooks.append(1)) + traced = tilelens.trace(ir)(kernel) + x, out = _inputs(device="cuda") + + ret = traced[(4,)](x, out, 64, BLOCK=16) + torch.cuda.synchronize() + + torch.testing.assert_close(out, x + 1) + (events,) = ir.finalized + # Compiled first (launched=False), then seen again before the launch; + # both carry the host compile, the launch compiled its own device kernel. + assert [e.launched for e in events] == [False, True] + assert events[0].kernel is events[1].kernel + assert events[0].kernel.target == CUDA89 + assert isinstance(ret, CompiledKernel) and ret.module is not None + # Only the real launch entered JITFunction.run. + assert hooks == [1] + + +@pytest.mark.parametrize("ir_cls", IR_CLASSES) +def test_autotune_over_heuristics_reports_every_config(ir_cls): + user = _make_autotuned() + ir = ir_cls() + traced = tilelens.trace(ir)(user) + x, out = _inputs(device=_device(ir_cls)) + + rets = [traced[_grid](x, out, 64), traced[_grid](x, out, 64)] + grids = [launch.grid for launch in trace_module.launches[-2:]] + _synchronize() + + # Every launch reports every config compile-only, whatever the autotune + # cache holds or the benchmark picked. + for events in ir.finalized: + compiled = [e for e in events if not e.launched] + assert [e.kwargs["BLOCK"] for e in compiled] == [16, 32] + assert {e.kwargs["EVEN"] for e in compiled} == {True} + assert len({e.specialization for e in compiled}) == 2 + first, second = ir.finalized + if ir_cls is _SkipIRClient: + assert len(first) == len(second) == 2 + assert torch.equal(out, torch.zeros_like(x)) + # No config was picked: nothing to return, and the configs' grids + # differ, so the Launch has none. + assert rets == [None, None] and grids == [None, None] + else: + # Benchmarking launches each config (once reported, however many + # benchmark calls); the cached second launch only the winner. + assert sorted(e.kwargs["BLOCK"] for e in first if e.launched) == [16, 32] + assert len([e for e in second if e.launched]) == 1 + torch.testing.assert_close(out, x + 1) + # The winner's kernel and grid, on the benchmarking launch too. + winner = traced.ir_runner.best_config.kwargs["BLOCK"] + assert all(ret.hash == rets[1].hash for ret in rets) + assert grids == [(64 // winner, 1, 1)] * 2 + assert user.cache == {} + + +def _make_runtime_stride_configs(): + # S is a runtime int that 2 and 3 specialize alike: both configs compile + # to one kernel, yet cover other elements (and here grids). + @triton.autotune( + configs=[ + triton.Config({"S": 2}, num_warps=1), + triton.Config({"S": 3}, num_warps=1), + ], + key=["n"], + ) + @triton.jit + def strided_copy(x_ptr, out_ptr, n, S, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK * S + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask), mask=mask) + + return strided_copy + + +def _grid_per_stride(meta): + return (64 // (16 * meta["S"]),) # S=2: 2 programs, S=3: 1 + + +@pytest.mark.parametrize("ir_cls", IR_CLASSES) +def test_configs_compiling_to_one_kernel_each_get_an_event(ir_cls): + """D22: the dedup key holds the binding, so the second config is not + lost behind the first one's (specialization, launched).""" + ir = ir_cls() + traced = tilelens.trace(ir)(_make_runtime_stride_configs()) + x, out = _inputs(device=_device(ir_cls)) + + traced[_grid_per_stride](x, out, 64, BLOCK=16) + _synchronize() + + (events,) = ir.finalized + compiled = [e for e in events if not e.launched] + assert [(e.kwargs["S"], e.resolved_grid) for e in compiled] == [ + (2, (2, 1, 1)), + (3, (1, 1, 1)), + ] + assert len({e.specialization for e in events}) == 1 + if ir_cls is _RunIRClient: + # Each config benchmarked for real, seen once more as launched. + assert sorted(e.kwargs["S"] for e in events if e.launched) == [2, 3] + + +@needs_gpu +def test_benchmark_repetitions_share_one_event_per_config(): + """An autotuned "run" launch benchmarks every config with many real + calls; each config's calls share one binding, so the event count stays + two per config (compile-only, then launched).""" + user = _make_autotuned() + calls = [] + user.fn.fn.add_pre_run_hook(lambda *args, **kwargs: calls.append(1)) + ir = _RunIRClient() + traced = tilelens.trace(ir)(user) + x, out = _inputs(device="cuda") + + traced[_grid](x, out, 64) + torch.cuda.synchronize() + + (events,) = ir.finalized + assert sorted((e.launched, e.kwargs["BLOCK"]) for e in events) == [ + (False, 16), + (False, 32), + (True, 16), + (True, 32), + ] + # One run() entry per benchmark repetition and for the final launch (the + # host compiles enter none): far more calls than events. + assert len(calls) > 10 * len(events) + torch.testing.assert_close(out, x + 1) + + +@needs_gpu +def test_a_fresh_constexpr_object_per_call_adds_no_event(): + """A heuristic building its tl.dtype per call hands every benchmark + repetition an equal but distinct constexpr object. Triton hashes it into + the kernel, so it adds no binding: two events per config, as above.""" + + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_warps=1), + triton.Config({"BLOCK": 32}, num_warps=2), + ], + key=["n"], + ) + @triton.heuristics({"DT": lambda args: tl.dtype("fp32")}) + @triton.jit + def cast_add_one(x_ptr, out_ptr, n, DT: tl.constexpr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + value = tl.load(x_ptr + offs, mask=mask).to(DT) + 1 + tl.store(out_ptr + offs, value, mask=mask) + + ir = _RunIRClient() + traced = tilelens.trace(ir)(cast_add_one) + x, out = _inputs(device="cuda") + + traced[_grid](x, out, 64) + torch.cuda.synchronize() + + (events,) = ir.finalized + assert sorted((e.launched, e.kwargs["BLOCK"]) for e in events) == [ + (False, 16), + (False, 32), + (True, 16), + (True, 32), + ] + torch.testing.assert_close(out, x + 1) + + +class _ForgetfulRunIRClient(_RunIRClient): + """Keeps nothing of a launch, as a harness-side client would.""" + + NAME = "ir_run_forgetful" + + def before_launch(self, event): + self.log.append("before") + + def finalize(self): + self.log.append("finalize") + return [] + + +def _compiled_sanitizer(): + from tilelens.clients import Sanitizer + + return Sanitizer(compile=True, abort_on_error=False) + + +def test_ir_only_launches_retain_no_tensor(): + """D23, without a GPU: a harness launching an IR-only trace in a loop + with fresh tensors, calling tilelens.clear() after each launch, holds on + to none of them (the grid callable closes over them, too); neither does + the trace's host-compile cache.""" + import gc + import weakref + + traced = tilelens.trace(_compiled_sanitizer())(_make_add_one()) + refs = [] + + def launch(): + x = torch.empty(4096) + out = torch.empty_like(x) + + def grid(meta): + return (triton.cdiv(x.numel(), meta["BLOCK"]),) + + traced[grid](x, out, 4096, BLOCK=1024) + assert not trace_module.launches[-1].tensors + tilelens.clear() + refs.extend((weakref.ref(x), weakref.ref(out))) + + for _ in range(3): + launch() + gc.collect() + + assert [ref() for ref in refs] == [None] * 6 + + +@needs_gpu +@pytest.mark.parametrize( + "make_client, make_kernel", + [ + (_compiled_sanitizer, _make_add_one), + (_ForgetfulRunIRClient, _make_add_one), + (_ForgetfulRunIRClient, _make_autotuned), + ], + ids=["compiled-sanitizer", "run", "run-autotuned"], +) +def test_ir_only_launches_retain_no_device_memory(make_client, make_kernel): + """D23: a harness launching an IR-only trace in a loop with fresh + device tensors, calling tilelens.clear() after each launch, holds on to + none of them (the grid callable closes over them, too).""" + traced = tilelens.trace(make_client())(make_kernel()) + n = 16 * 2**20 # 64 MiB of float32 + # The autotuned kernel's configs set BLOCK themselves. + kwargs = {"BLOCK": 1024} if make_kernel is _make_add_one else {} + + def launch(): + x = torch.empty(n, device="cuda") + out = torch.empty_like(x) + + def grid(meta): + return (triton.cdiv(x.numel(), meta["BLOCK"]),) + + traced[grid](x, out, n, **kwargs) + assert not trace_module.launches[-1].tensors + tilelens.clear() + + # Measured from before the first launch: the manager's Launch is + # replaced per launch, so tensors it held would be the last launch's. + torch.cuda.synchronize() + base = torch.cuda.memory_allocated() + for _ in range(10): + launch() + torch.cuda.synchronize() + + assert torch.cuda.memory_allocated() - base < 2**20 + + +class _InterruptingRunIRClient(_ForgetfulRunIRClient): + """The user hits Ctrl+C in a launch's first autotune benchmark call.""" + + NAME = "ir_run_interrupting" + + def before_launch(self, event): + super().before_launch(event) + if event.launched: + raise KeyboardInterrupt + + +@needs_gpu +def test_an_interrupted_benchmark_retains_no_device_memory(): + """D23: Triton's _bench skips the post_hook that drops a benchmark + call's restore_value clones for a KeyboardInterrupt; the aborted launch + keeps neither those device clones nor the caller's tensors.""" + traced = tilelens.trace(_InterruptingRunIRClient())( + _make_autotuned(restore_value=["x_ptr"]) + ) + n = 16 * 2**20 # 64 MiB of float32 + + def launch(): + x = torch.empty(n, device="cuda") + out = torch.empty_like(x) + with pytest.raises(KeyboardInterrupt): + traced[_grid](x, out, n) + tilelens.clear() + + torch.cuda.synchronize() + base = torch.cuda.memory_allocated() + for _ in range(3): + launch() + torch.cuda.synchronize() + + assert torch.cuda.memory_allocated() - base < 2**20 + + +@pytest.mark.parametrize("ir_cls", IR_CLASSES) +def test_plain_heuristics_kernel_fires_events(ir_cls): + @triton.heuristics({"BLOCK": lambda args: 16}) + @triton.jit + def heur_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) + + ir = ir_cls() + traced = tilelens.trace(ir)(heur_add_one) + x, out = _inputs(device=_device(ir_cls)) + + traced[_grid](x, out, 64) + _synchronize() + + (events,) = ir.finalized + assert [e.kwargs["BLOCK"] for e in events if not e.launched] == [16] + assert events[0].resolved_grid == (4, 1, 1) + if ir_cls is _RunIRClient: + torch.testing.assert_close(out, x + 1) + else: + assert torch.equal(out, torch.zeros_like(x)) + + +@pytest.mark.parametrize("ir_cls", IR_CLASSES) +def test_a_compile_error_is_data_and_the_next_launch_is_clean(ir_cls): + @triton.jit + def bounded(x_ptr, out_ptr, n, BLOCK: 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) + + ir = ir_cls() + traced = tilelens.trace(ir)(bounded) + x, out = _inputs(device=_device(ir_cls)) + + # The only config failed to host-compile: reported as data (D27). A + # skipped launch then ends normally; a real launch compiles for the + # device, which fails as the untraced launch does. + if ir_cls is _RunIRClient: + with pytest.raises(CompileTimeAssertionFailure): + traced[(1,)](x, out, 64, BLOCK=64) + assert ir.log == ["begin", "compile_failed", "abort"] + else: + assert traced[(1,)](x, out, 64, BLOCK=64) is None + assert ir.log == ["begin", "compile_failed", "finalize"] + assert ir.finalized == [[]] + (failure,) = ir.failures + assert isinstance(failure.error, CompileTimeAssertionFailure) + assert failure.target == CUDA89 + assert "run" not in vars(traced.jit_fn) + + ir.log.clear() + traced[(4,)](x, out, 64, BLOCK=16) + _synchronize() + + if ir_cls is _RunIRClient: + assert ir.log == ["begin"] + ["before", "after"] * 2 + ["finalize"] + assert [len(events) for events in ir.finalized] == [2] + torch.testing.assert_close(out, x + 1) + else: + assert ir.log == ["begin", "before", "after", "finalize"] + # The failed launch finalized with nothing compiled. + assert [len(events) for events in ir.finalized] == [0, 1] + + +def _make_bounded_autotuned(): + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_warps=1), + # Fails to compile (static_assert) ... + triton.Config({"BLOCK": 64}, num_warps=1), + # ... compiles, but needs more threads than a block can have. + triton.Config({"BLOCK": 16}, num_warps=64), + ], + key=["n"], + ) + @triton.jit + def bounded(x_ptr, out_ptr, n, BLOCK: 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 bounded + + +def test_a_config_that_fails_to_compile_is_reported_and_one_too_big_is_analyzed(): + """D25: nothing is loaded, so a config some device could not run + (num_warps=64: more threads than a block can have) is analyzed like any + other; it can only add findings, never hide one.""" + ir = _SkipIRClient() + traced = tilelens.trace(ir)(_make_bounded_autotuned()) + x, out = _inputs() + + traced[_grid](x, out, 64) + + (events,) = ir.finalized + assert [(e.kwargs["BLOCK"], e.kwargs["num_warps"]) for e in events] == [ + (16, 1), + (16, 64), + ] + ((config, error),) = [ + ((f.kwargs["BLOCK"], f.kwargs["num_warps"]), f.error) for f in ir.failures + ] + assert config == (64, 1) and isinstance(error, CompileTimeAssertionFailure) + assert torch.equal(out, torch.zeros_like(x)) + + +@needs_gpu +def test_the_real_launch_drops_what_the_device_cannot_run(): + """Under "run" the device decides: the JIT's own load of the num_warps=64 + config raises OutOfResources, which the autotuner's benchmark absorbs + after the host-compiled config was delivered.""" + ir = _RunIRClient() + traced = tilelens.trace(ir)(_make_bounded_autotuned()) + x, out = _inputs(device="cuda") + + traced[_grid](x, out, 64) + torch.cuda.synchronize() + + (events,) = ir.finalized + compiled = {(e.kwargs["BLOCK"], e.kwargs["num_warps"]) for e in events} + assert compiled == {(16, 1), (16, 64)} + # The failing config is reported once, although benchmarked again. + assert [(f.kwargs["BLOCK"], f.kwargs["num_warps"]) for f in ir.failures] == [ + (64, 1) + ] + # Only the winner reached after_launch as a real launch; (16, 64) raised. + assert ir.log.count("after") == len(events) - 1 + torch.testing.assert_close(out, x + 1) + + +def _times_two(): + # `other=` exercises the semantic._load_legacy path (`other.handle if + # other else None`) that a leaked tensor.__bool__ would break. + @triton.jit + def times_two(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n + loaded = tl.load(x_ptr + offs, mask=mask, other=-1.0) + tl.store(out_ptr + offs, loaded * 2, mask=mask) + + return times_two + + +def test_mixed_trace_compiles_for_ir_and_interprets_for_eager(): + kernel = _make_add_one() + ir, eager = _SkipIRClient(), _EagerCounter() + traced = tilelens.trace(eager)(tilelens.trace(ir)(kernel)) + x, out = _inputs() + + traced[(4,)](x, out, 64, BLOCK=16) + + (events,) = ir.finalized + assert [e.launched for e in events] == [False] + assert list(events[0].kernel.asm) == ["ttir"] + assert eager.stores == 4 + torch.testing.assert_close(out, x + 1) + # Launch.tensors: the interpreter's copies only, not the caller's + # tensors on top. + tensors = trace_module.launches[-1].tensors + assert len(tensors) == 2 and not {id(t) for t in tensors} & {id(x), id(out)} + + # D4b regression: interpreter patches must not leak into later compiles, + # of another kernel or of this one (a new BLOCK forces a recompile). + compiler = HostCompiler() + for jit_fn, block in ((_times_two(), 64), (kernel, 32)): + compiled = compiler.compile( + jit_fn, (x, out, 64), {"BLOCK": block}, target=CUDA80, stages={"ttir"} + ) + assert "tt.store" in compiled.asm["ttir"] + + +@needs_gpu +def test_interpreter_patches_do_not_leak_into_later_real_launches(): + kernel = _make_add_one() + traced = tilelens.trace(_EagerCounter())(tilelens.trace(_SkipIRClient())(kernel)) + x, out = _inputs(device="cuda") + traced[(4,)](x, out, 64, BLOCK=16) + + doubled = torch.zeros_like(x) + _times_two()[(1,)](x, doubled, 64, BLOCK=64) + again = torch.zeros_like(x) + kernel[(2,)](x, again, 64, BLOCK=32) + torch.cuda.synchronize() + torch.testing.assert_close(doubled, x * 2) + torch.testing.assert_close(again, x + 1) + + +@pytest.mark.parametrize( + "client_cls", + [_SkipIRClient, _EagerCounter, pytest.param(_RunIRClient, marks=needs_gpu)], +) +def test_trace_leaves_the_users_autotuner_unchanged(client_cls): + user = _make_autotuned() + heuristics, jit_fn = user.fn, user.fn.fn + before = dict(vars(user)) + before_heuristics = dict(vars(heuristics)) + traced = tilelens.trace(client_cls())(user) + x, out = _inputs(device=_device(client_cls)) + + traced[_grid64](x, out, 64) + _synchronize() + + assert vars(user).keys() == before.keys() + assert all(vars(user)[k] is v for k, v in before.items()) + assert vars(heuristics).keys() == before_heuristics.keys() + assert all(vars(heuristics)[k] is v for k, v in before_heuristics.items()) + assert user.fn is heuristics and heuristics.fn is jit_fn + assert user.cache == {} + + +@needs_gpu +def test_the_users_autotuner_still_autotunes_for_real_after_a_trace(): + user = _make_autotuned() + x, out = _inputs(device="cuda") + tilelens.trace(_RunIRClient())(user)[_grid64](x, out, 64) + + untraced = torch.zeros_like(x) + user[_grid64](x, untraced, 64) + torch.cuda.synchronize() + torch.testing.assert_close(untraced, x + 1) + assert len(user.cache) == 1 + assert user.best_config in user.configs + + +def _add_one_helper(x): + return x + 1 + + +def _kernel_with_helper(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, + _traced_add_one_helper(tl.load(x_ptr + offs, mask=mask)), # noqa: F821 + mask=mask, + ) + + +@pytest.fixture +def helper_kernel(): + """A kernel whose device function is a module global wrapped by + tilelens.trace, as the CLI wrappers do for every @triton.jit function. + The real code generator only accepts JITFunctions, so real compiles must + see the unwrapped binding. Built here, not at import, so the jits are real + even when TRITON_INTERPRET was set during collection.""" + module_globals = globals() + helper = tilelens.trace(_EagerCounter())(triton.jit(_add_one_helper)) + module_globals["_traced_add_one_helper"] = helper + try: + yield triton.jit(_kernel_with_helper), helper + finally: + module_globals.pop("_traced_add_one_helper", None) + + +@pytest.mark.parametrize("ir_cls", IR_CLASSES) +def test_traced_device_function_compiles_under_ir_capture(helper_kernel, ir_cls): + kernel, helper = helper_kernel + ir = ir_cls() + traced = tilelens.trace(ir)(kernel) + x, out = _inputs(device=_device(ir_cls)) + + traced[(4,)](x, out, 64, BLOCK=16) + _synchronize() + + (events,) = ir.finalized + if ir_cls is _RunIRClient: + torch.testing.assert_close(out, x + 1) + assert [e.launched for e in events] == [False, True] + else: + assert [e.launched for e in events] == [False] + assert "tt.func" in events[0].kernel.asm["ttir"] + assert globals()["_traced_add_one_helper"] is helper + + +def _kernel_with_package_helper(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, + _helper_pkg.api.add_one(tl.load(x_ptr + offs, mask=mask)), # noqa: F821 + mask=mask, + ) + + +@pytest.mark.parametrize("ir_cls", IR_CLASSES) +def test_traced_device_function_behind_a_package_path_compiles(ir_cls): + # `pkg.api.add_one`, where `api` re-exports the traced helper: Triton + # resolves it through two module attributes. + module_globals = globals() + helper = tilelens.trace(_EagerCounter())(triton.jit(_add_one_helper)) + pkg = types.ModuleType("tilelens_test_helper_pkg") + pkg.api = types.ModuleType("tilelens_test_helper_pkg.api") + pkg.api.add_one = helper + module_globals["_helper_pkg"] = pkg + try: + ir = ir_cls() + traced = tilelens.trace(ir)(triton.jit(_kernel_with_package_helper)) + x, out = _inputs(device=_device(ir_cls)) + + traced[(4,)](x, out, 64, BLOCK=16) + _synchronize() + + (events,) = ir.finalized + assert "tt.func" in events[0].kernel.asm["ttir"] + assert pkg.api.add_one is helper + if ir_cls is _RunIRClient: + torch.testing.assert_close(out, x + 1) + else: + assert torch.equal(out, torch.zeros_like(x)) + finally: + module_globals.pop("_helper_pkg", None) + + +@needs_gpu +def test_traced_device_function_compiles_in_the_interpreted_warmup(helper_kernel): + # Interpreting clients that vote for a warmup compile for real, through + # the JIT (e.g. the profiler); the traced helper must resolve there too. + kernel, helper = helper_kernel + eager = _EagerCounter(warmup_vote=True) + traced = tilelens.trace(eager)(kernel) + x, out = _inputs(device="cuda") + + traced[(4,)](x, out, 64, BLOCK=16) + torch.cuda.synchronize() + + assert len(eager.warmups) == 1 and "ttir" in eager.warmups[0].asm + assert eager.stores == 4 + torch.testing.assert_close(out, x + 1) + assert globals()["_traced_add_one_helper"] is helper + + +def _make_apply_fn(): + @triton.jit + def apply_fn(x_ptr, out_ptr, n, FN: tl.constexpr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, FN(tl.load(x_ptr + offs, mask=mask)), mask=mask) + + return apply_fn + + +def _make_apply_default(): + # Triton's code generator evaluates parameter defaults in the kernel's + # globals, so the default is a module global (bound by the caller). + @triton.jit + def apply_default( + x_ptr, + out_ptr, + n, + BLOCK: tl.constexpr, + FN: tl.constexpr = _traced_default_helper, # noqa: F821 + ): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, FN(tl.load(x_ptr + offs, mask=mask)), mask=mask) + + return apply_default + + +def _make_apply_first(): + @triton.jit + def apply_first(x_ptr, out_ptr, n, FNS: tl.constexpr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, FNS[0](tl.load(x_ptr + offs, mask=mask)), mask=mask) + + return apply_first + + +def _assert_real_compiles_still_work(monkeypatch, *, launch: bool): + # D4b: the interpreter's triton.language patches must not have leaked. A + # never-seen kernel, host-compiled (never disk-cached); with ``launch`` + # (a needs_gpu test) also compiled and launched by the JIT (no disk-cache + # shortcut). + monkeypatch.setenv("TRITON_ALWAYS_COMPILE", "1") + + @triton.jit + def times_two(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) * 2, mask=mask) + + x, out = _inputs() + compiled = HostCompiler().compile( + times_two, (x, out, 64), {"BLOCK": 16}, target=CUDA80, stages={"ttir"} + ) + assert "tt.store" in compiled.asm["ttir"] + if not launch: + return + x, out = _inputs(device="cuda") + times_two[(4,)](x, out, 64, BLOCK=16) + torch.cuda.synchronize() + torch.testing.assert_close(out, x * 2) + + +@pytest.mark.parametrize("ir_cls", IR_CLASSES) +@pytest.mark.parametrize("passing", ["keyword", "default", "tuple"]) +def test_traced_helper_passed_as_an_argument_compiles(ir_cls, passing, monkeypatch): + # The CLI shape (every @triton.jit is a TritonTrace), with the helper + # reaching the real compile as a constexpr argument, which the globals + # unwrap cannot see. + helper = tilelens.trace(_EagerCounter())(triton.jit(_add_one_helper)) + if passing == "keyword": + kernel, extra = _make_apply_fn(), {"FN": helper} + elif passing == "default": + monkeypatch.setitem(globals(), "_traced_default_helper", helper) + kernel, extra = _make_apply_default(), {} + else: + kernel, extra = _make_apply_first(), {"FNS": (helper,)} + ir = ir_cls() + x, out = _inputs(device=_device(ir_cls)) + + tilelens.trace(ir)(kernel)[(4,)](x, out, 64, BLOCK=16, **extra) + _synchronize() + + assert ir.failures == [] + (events,) = ir.finalized + assert "tt.func" in events[0].kernel.asm["ttir"] + if ir_cls is _RunIRClient: + torch.testing.assert_close(out, x + 1) + _assert_real_compiles_still_work(monkeypatch, launch=ir_cls is _RunIRClient) + + +@needs_gpu +def test_voted_warmup_compiles_with_a_traced_helper_argument(monkeypatch): + # Interpreting clients that vote for a warmup compile for real, through + # the JIT. + helper = tilelens.trace(_EagerCounter())(triton.jit(_add_one_helper)) + eager = _EagerCounter(warmup_vote=True) + x, out = _inputs(device="cuda") + + tilelens.trace(eager)(_make_apply_fn())[(4,)](x, out, 64, BLOCK=16, FN=helper) + torch.cuda.synchronize() + + assert len(eager.warmups) == 1 and "ttir" in eager.warmups[0].asm + torch.testing.assert_close(out, x + 1) + _assert_real_compiles_still_work(monkeypatch, launch=True) + + +def test_a_traced_callee_the_compile_still_reaches_fails_it_cleanly(monkeypatch): + # Should a compile reach a TritonTrace anyway (here: with the argument + # mapping switched off), the trace refuses to interpret there. The + # compile fails; triton.language is left alone for later compiles. + monkeypatch.setattr( + trace_module, "_untraced_call_args", lambda jit_fn, args, kwargs: (args, kwargs) + ) + helper = tilelens.trace(_EagerCounter())(triton.jit(_add_one_helper)) + ir = _SkipIRClient() + x, out = _inputs() + + # A compile failure like any other (D27): data, and the skipped launch + # ends normally. + tilelens.trace(ir)(_make_apply_fn())[(4,)](x, out, 64, BLOCK=16, FN=helper) + + assert ir.log == ["begin", "compile_failed", "finalize"] + (failure,) = ir.failures + assert isinstance(failure.error, CompilationError) + assert "outside a traced launch" in str(failure.error) + _assert_real_compiles_still_work(monkeypatch, launch=False) diff --git a/tests/end_to_end/test_ir_smoke.py b/tests/end_to_end/test_ir_smoke.py new file mode 100644 index 000000000..a0b507d3b --- /dev/null +++ b/tests/end_to_end/test_ir_smoke.py @@ -0,0 +1,417 @@ +"""Smoke of the IR layers end to end: a toy IRClient under tilelens.trace +parses every specialization's TTIR, compiled on the host (D25: CPU tensors, +no GPU), through the real ParseCache and tilelens.ir.ttir_reader.parse_ttir, +and puts its IRVerdict into Launch.records. Each kernel is launched twice; +the second launch must find its texts in the parse cache. Per-layer tests +live in tests/unit/ir/ and tests/end_to_end/test_ir_client.py. +""" + +from __future__ import annotations + +import importlib +import inspect + +import pytest +import torch +import triton +import triton.language as tl + +import tilelens +from tilelens.ir import ConfigVerdict, IRClient, IRVerdict, ParseCache, Refusal +from tilelens.ir import _mlir_walk, ttir_reader +from tilelens.ir.ttir_reader import Const, IterArgOffset, Param + +trace_module = importlib.import_module("tilelens.core.trace") + + +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 (D25). + 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 (D25): 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 _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 +def parsed_texts(monkeypatch): + """Every text the real reader parses. ParseCache resolves its default + reader at each lookup, so the spy stands in for it for the whole test.""" + texts: list[str] = [] + parse_ttir = ttir_reader.parse_ttir + + def spy(text): + texts.append(text) + return parse_ttir(text) + + monkeypatch.setattr(ttir_reader, "parse_ttir", spy) + return texts + + +class _ParsingIR(IRClient): + """Parses each specialization's TTIR; a refused config makes the launch + "unsupported" with the first refusal.""" + + NAME = "parsing_ir" + LAUNCH = "skip" + IR_STAGES = frozenset({"ttir"}) + + def __init__(self): + super().__init__() + self.parses = ParseCache() + # Per finalized launch: (specializations, parse outcomes). + self.launches: list[tuple] = [] + + def analyze_launch(self, log): + assert log.failures == (), log.failures + specs = log.specializations + outcomes = [self.parses.get(spec.artifacts.stages["ttir"]) for spec in specs] + self.launches.append((specs, outcomes)) + per_config = [] + for spec, outcome in zip(specs, outcomes): + assert outcome.error is None, outcome.error + if outcome.refusal is None: + per_config.append( + ConfigVerdict(spec.specialization, spec.config, "parsed") + ) + else: + refusal = Refusal.from_exception(outcome.refusal) + per_config.append( + ConfigVerdict(spec.specialization, spec.config, "refused", refusal) + ) + refusals = [c.refusal for c in per_config if c.refusal is not None] + if refusals: + return [], IRVerdict( + self.NAME, "unsupported", refusal=refusals[0], per_config=per_config + ) + return [], IRVerdict(self.NAME, "parsed", per_config=per_config) + + def on_analysis_error(self, exc): + return IRVerdict(self.NAME, "error", notes=[f"{type(exc).__name__}: {exc}"]) + + def on_refusal(self, refusal): + return IRVerdict(self.NAME, "unsupported", refusal=refusal) + + +def _launch_twice(kernel, grid, *args, **kwargs): + """Trace ``kernel`` with a fresh _ParsingIR and launch it twice; checks + what every launch must hold and returns the client and the first + launch's graphs (None for a refused text), in specialization order.""" + ir = _ParsingIR() + traced = tilelens.trace(ir)(kernel) + verdicts = [] + for _ in range(2): + traced[grid](*args, **kwargs) + verdicts.append(ir.last_verdict) + + for verdict in verdicts: + assert verdict.status != "error", verdict.notes + first, second = verdicts + assert [launch.records for launch in trace_module.launches[-2:]] == [ + [first], + [second], + ] + assert second == first + (specs, outcomes), (specs_again, outcomes_again) = ir.launches + assert [s.specialization for s in specs_again] == [s.specialization for s in specs] + # The second launch is a parse-cache hit: the same outcome objects, a + # refusal's kind included. + assert all(a is b for a, b in zip(outcomes_again, outcomes, strict=True)) + return ir, [outcome.graph for outcome in outcomes] + + +def _line_of(jit_fn, needle: str) -> int: + """The source line of ``jit_fn`` that contains ``needle``.""" + lines, start = inspect.getsourcelines(jit_fn.fn) + (offset,) = [i for i, line in enumerate(lines) if needle in line] + return start + offset + + +def _vector(n=64, dtype=torch.float32): + return torch.arange(n, dtype=dtype) + + +def _make_plain(): + @triton.jit + def plain(x_ptr, out_ptr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs) + 1) + + return plain + + +def _make_masked(): + @triton.jit + def masked(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 masked + + +def _make_row_sum(): + @triton.jit + def row_sum(x_ptr, out_ptr, n_cols, BLOCK: tl.constexpr): + row = tl.program_id(0) + ptrs = x_ptr + row * n_cols + tl.arange(0, BLOCK) + acc = tl.zeros((BLOCK,), dtype=tl.float32) + for _ in range(0, n_cols, BLOCK): + acc += tl.load(ptrs) + ptrs += BLOCK + tl.store(out_ptr + row, tl.sum(acc)) + + return row_sum + + +def _make_autotuned(): + @triton.autotune( + configs=[triton.Config({"BLOCK": b}, num_warps=1) for b in (16, 32)], + key=["n"], + ) + @triton.jit + def masked_tuned(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 masked_tuned + + +def _make_gather(): + @triton.jit + def gather(x_ptr, idx_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + idx = tl.load(idx_ptr + offs, mask=mask, other=0) + tl.store(out_ptr + offs, tl.load(x_ptr + idx, mask=mask), mask=mask) + + return gather + + +def _make_calls(): + # Triton 3.6 passes only scalars to a noinline function. + @triton.jit(noinline=True) + def store_one(out_ptr, i): + tl.store(out_ptr + i, 1.0) + + @triton.jit + def calls(out_ptr): + store_one(out_ptr, tl.program_id(0)) + + return calls + + +def _assert_skipped(out): + # LAUNCH="skip": compiled and analyzed, never launched. + assert torch.equal(out, torch.zeros_like(out)) + + +def test_plain_kernel(): + kernel = _make_plain() + x, out = _vector(), torch.zeros(64) + + ir, (graph,) = _launch_twice(kernel, (4,), x, out, BLOCK=16) + + _assert_skipped(out) + (config,) = ir.last_verdict.per_config + assert (config.status, config.config) == ("parsed", {}) + assert trace_module.launches[-1].grid == (4, 1, 1) + assert graph.kernel_name == "plain" + assert [(a.kind, a.base_param, a.mask) for a in graph.accesses] == [ + ("load", "x_ptr", None), + ("store", "out_ptr", None), + ] + line = _line_of(kernel, "tl.store(") + source = (kernel.fn.__code__.co_filename, line) + assert [(a.loc.file, a.loc.line) for a in graph.accesses] == [source, source] + assert graph.loop is None and graph.pid_axes == {0} + + +def test_masked_kernel(): + x, out = _vector(), torch.zeros(64) + + ir, (graph,) = _launch_twice(_make_masked(), (4,), x, out, 64, BLOCK=16) + + _assert_skipped(out) + assert ir.last_verdict.status == "parsed" + assert [(a.kind, a.base_param) for a in graph.accesses] == [ + ("load", "x_ptr"), + ("store", "out_ptr"), + ] + assert all(a.mask is not None for a in graph.accesses) + assert graph.arg("n").int_bits == 32 and graph.loop is None + + +def test_loop_with_a_pointer_iter_arg(): + x, out = _vector(4 * 64), torch.zeros(4) + + ir, (graph,) = _launch_twice(_make_row_sum(), (4,), x, out, 64, BLOCK=16) + + _assert_skipped(out) + assert ir.last_verdict.status == "parsed" + loop = graph.loop + assert (loop.lower, loop.upper, loop.step) == (Const(0), Param("n_cols"), Const(16)) + (iter_arg,) = graph.iter_args + assert (iter_arg.base_param, iter_arg.delta) == ("x_ptr", Const(16)) + load, store = graph.accesses + assert (load.kind, load.in_loop, load.offset) == ("load", True, IterArgOffset(0)) + assert (store.kind, store.in_loop) == ("store", False) + ((spec,), _), _ = ir.launches + (binding,) = spec.bindings + assert dict(binding.params) == {"n_cols": 64} + + +def test_autotune_parses_both_configs(parsed_texts): + x, out = _vector(), torch.zeros(64) + + ir, graphs = _launch_twice( + _make_autotuned(), lambda meta: (triton.cdiv(64, meta["BLOCK"]),), x, out, 64 + ) + + _assert_skipped(out) + per_config = ir.last_verdict.per_config + assert [(c.status, c.config["BLOCK"]) for c in per_config] == [ + ("parsed", 16), + ("parsed", 32), + ] + assert len({c.specialization for c in per_config}) == 2 + # One parse per config, none on the second launch. + assert len(parsed_texts) == 2 + assert [g.kernel_name for g in graphs] == ["masked_tuned"] * 2 + # A skipped autotuned launch picks no config, and the configs' grids differ. + assert trace_module.launches[-1].grid is None + + +def test_gather_refuses_as_indirect_address(parsed_texts): + kernel = _make_gather() + x, out = _vector(), torch.zeros(64) + idx = torch.arange(64, dtype=torch.int32) + + ir, graphs = _launch_twice(kernel, (4,), x, idx, out, 64, BLOCK=16) + + _assert_skipped(out) + verdict = ir.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "indirect-address") + assert verdict.per_config[0].refusal == verdict.refusal + assert verdict.refusal.loc.line == _line_of(kernel, "x_ptr + idx") + assert graphs == [None] and len(parsed_texts) == 1 + + +def test_noinline_call_refuses_as_call(parsed_texts): + kernel = _make_calls() + out = torch.zeros(4) + + ir, graphs = _launch_twice(kernel, (4,), out) + + _assert_skipped(out) + verdict = ir.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "call") + assert "store_one" in verdict.refusal.message + assert verdict.refusal.loc.line == _line_of(kernel, "store_one(") + assert graphs == [None] and len(parsed_texts) == 1 + + +def _release() -> str: + return _mlir_walk.triton_release()[0] + + +def _make_barrier(): + @triton.jit + def barrier(x_ptr, out_ptr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + v = tl.load(x_ptr + offs) + tl.debug_barrier() + tl.store(out_ptr + offs, v) + + return barrier + + +# The op tl.debug_barrier() prints, per Triton release; each release's +# reader holds its own barrier inert. +_BARRIER_OP = {"3.6": "gpu.barrier", "3.8": "ttg.barrier all"} + + +def test_debug_barrier_is_inert(parsed_texts): + x, out = _vector(), torch.zeros(64) + + ir, (graph,) = _launch_twice(_make_barrier(), (4,), x, out, BLOCK=16) + + _assert_skipped(out) + assert ir.last_verdict.status == "parsed" + assert [(a.kind, a.base_param) for a in graph.accesses] == [ + ("load", "x_ptr"), + ("store", "out_ptr"), + ] + (text,) = parsed_texts + assert _BARRIER_OP[_release()] in text + + +def _make_tuple_args(): + @triton.jit + def pair_copy(ptrs, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(ptrs[1] + offs, tl.load(ptrs[0] + offs, mask=mask), mask=mask) + + return pair_copy + + +# The names of a tuple parameter's flattened TTIR arguments, per Triton +# release: 3.6 names every leaf by the parameter (the reader refuses two +# parameters of one name), 3.8 by its path. +_TUPLE_LEAVES = {"3.6": None, "3.8": ("ptrs.0", "ptrs.1")} + + +def test_tuple_parameter_leaves(parsed_texts): + x, out = _vector(), torch.zeros(64) + + ir, (graph,) = _launch_twice(_make_tuple_args(), (4,), (x, out), 64, BLOCK=16) + + _assert_skipped(out) + leaves = _TUPLE_LEAVES[_release()] + verdict = ir.last_verdict + if leaves is None: + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "other") + assert "two parameters named 'ptrs'" in verdict.refusal.message + assert graph is None + return + assert verdict.status == "parsed" + assert [a.name for a in graph.func_args] == [*leaves, "n"] + assert [(a.kind, a.base_param) for a in graph.accesses] == [ + ("load", leaves[0]), + ("store", leaves[1]), + ] diff --git a/tests/end_to_end/test_race_detector.py b/tests/end_to_end/test_race_detector.py index 7aed6c316..6e876b706 100644 --- a/tests/end_to_end/test_race_detector.py +++ b/tests/end_to_end/test_race_detector.py @@ -125,7 +125,7 @@ def test_string_dispatch_and_manager_lookup(_isolate_race_detector_cfg): x = torch.zeros(8, dtype=torch.float32) traced[(1,)](x, BLOCK=8) - rd = traced.client_manager.clients["race_detector"] + rd = traced.client_manager.get_client("race_detector") assert isinstance(rd, SymbolicRaceDetector) assert len(rd.records) == 2 diff --git a/tests/golden/ir/expected.json b/tests/golden/ir/expected.json new file mode 100644 index 000000000..4c3f73ca7 --- /dev/null +++ b/tests/golden/ir/expected.json @@ -0,0 +1,5038 @@ +{ + "adv_cf_blockargs.ttir": { + "stats": { + "ops": 40, + "implicit_ops": 0, + "blocks": 9, + "values": 35, + "funcs": 1, + "ssa_edges": 48, + "result_locs": 28, + "arg_locs": 5, + "bind_attrs": 4, + "type_checks": 6, + "needed_attrs": 21, + "cf_edges": 8, + "pred_checks": 6 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='sgt'", + 3 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=100", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "arith.constant.value=7", + 1 + ], + [ + "cf.cond_br -> 2 successors", + 4 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 12 + ] + ], + "funcs": [ + "cf_blockargs" + ] + }, + "adv_consts.ttir": { + "stats": { + "ops": 51, + "implicit_ops": 0, + "blocks": 2, + "values": 42, + "funcs": 1, + "ssa_edges": 60, + "result_locs": 38, + "arg_locs": 4, + "bind_attrs": 4, + "type_checks": 16, + "needed_attrs": 23, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpf.predicate='une'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 14 + ], + [ + "arith.constant.value=('float', '-2.14748365E+9')", + 1 + ], + [ + "arith.constant.value=('float', '-3.000000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '-7.000000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '0x7FC00000')", + 1 + ], + [ + "arith.constant.value=('float', '0xFF800000')", + 1 + ], + [ + "arith.constant.value=('float', '1.000000e-30')", + 1 + ], + [ + "arith.constant.value=-1", + 2 + ], + [ + "arith.constant.value=-9223372036854775807", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=128", + 1 + ], + [ + "arith.constant.value=192", + 1 + ], + [ + "arith.constant.value=4294967295", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "arith.constant.value=9223372036854775807", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 13 + ] + ], + "funcs": [ + "consts" + ] + }, + "adv_descs.ttir": { + "stats": { + "ops": 15, + "implicit_ops": 0, + "blocks": 2, + "values": 12, + "funcs": 1, + "ssa_edges": 29, + "result_locs": 9, + "arg_locs": 3, + "bind_attrs": 3, + "type_checks": 4, + "needed_attrs": 8, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "tt.descriptor_reduce.kind='add'", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "descs" + ] + }, + "adv_hinted.ttir": { + "stats": { + "ops": 15, + "implicit_ops": 0, + "blocks": 3, + "values": 15, + "funcs": 1, + "ssa_edges": 15, + "result_locs": 10, + "arg_locs": 5, + "bind_attrs": 4, + "type_checks": 3, + "needed_attrs": 8, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=32", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "tt.reduce.axis=0", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "hinted" + ] + }, + "adv_multi_func.ttir": { + "stats": { + "ops": 41, + "implicit_ops": 2, + "blocks": 8, + "values": 38, + "funcs": 4, + "ssa_edges": 57, + "result_locs": 27, + "arg_locs": 9, + "bind_attrs": 17, + "type_checks": 5, + "needed_attrs": 27, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 2 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 2 + ], + [ + "arith.constant.value=3", + 1 + ], + [ + "arith.constant.value=True", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='add'", + 1 + ], + [ + "tt.atomic_rmw.scope='gpu'", + 1 + ], + [ + "tt.atomic_rmw.sem='acq_rel'", + 1 + ], + [ + "zero-result op with a text loc", + 15 + ] + ], + "funcs": [ + "multi_func", + "adv_kernels._nl_pair__Pi32_i32__(2,)cconstexpr_3_", + "adv_kernels._nl_loop__Pi32_i32__", + "adv_kernels._nl_pair__Pi32_i32_i32__" + ] + }, + "adv_multi_result.ttir": { + "stats": { + "ops": 46, + "implicit_ops": 0, + "blocks": 5, + "values": 51, + "funcs": 1, + "ssa_edges": 66, + "result_locs": 43, + "arg_locs": 3, + "bind_attrs": 6, + "type_checks": 7, + "needed_attrs": 16, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='sgt'", + 1 + ], + [ + "arith.constant.value=('float', '1.000000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '2.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=4", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 7 + ] + ], + "funcs": [ + "multi_result" + ] + }, + "adv_names.ttir": { + "stats": { + "ops": 16, + "implicit_ops": 0, + "blocks": 2, + "values": 15, + "funcs": 1, + "ssa_edges": 17, + "result_locs": 12, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 3, + "needed_attrs": 9, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=2", + 1 + ], + [ + "arith.constant.value=3", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "names" + ] + }, + "adv_nest3.ttir": { + "stats": { + "ops": 125, + "implicit_ops": 3, + "blocks": 16, + "values": 119, + "funcs": 1, + "ssa_edges": 190, + "result_locs": 104, + "arg_locs": 5, + "bind_attrs": 9, + "type_checks": 12, + "needed_attrs": 36, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 3 + ], + [ + "arith.cmpi.predicate='sgt'", + 3 + ], + [ + "arith.cmpi.predicate='slt'", + 3 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '2.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 3 + ], + [ + "arith.constant.value=2", + 2 + ], + [ + "arith.constant.value=3", + 1 + ], + [ + "arith.constant.value=5", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 5 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 21 + ] + ], + "funcs": [ + "nest3" + ] + }, + "adv_reduce3.ttir": { + "stats": { + "ops": 83, + "implicit_ops": 0, + "blocks": 7, + "values": 97, + "funcs": 1, + "ssa_edges": 126, + "result_locs": 75, + "arg_locs": 20, + "bind_attrs": 6, + "type_checks": 16, + "needed_attrs": 31, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpf.predicate='oeq'", + 2 + ], + [ + "arith.cmpf.predicate='ogt'", + 1 + ], + [ + "arith.cmpf.predicate='olt'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 2 + ], + [ + "arith.constant.splat=True", + 3 + ], + [ + "arith.constant.value=('float', '2.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 2 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=16", + 1 + ], + [ + "arith.constant.value=32", + 2 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.expand_dims.axis=0", + 1 + ], + [ + "tt.expand_dims.axis=1", + 2 + ], + [ + "tt.make_range.end=16", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.start=0", + 2 + ], + [ + "tt.reduce.axis=1", + 3 + ], + [ + "tt.scan.axis=1", + 1 + ], + [ + "zero-result op with a text loc", + 12 + ] + ], + "funcs": [ + "reduce3" + ] + }, + "adv_views.ttir": { + "stats": { + "ops": 33, + "implicit_ops": 0, + "blocks": 2, + "values": 30, + "funcs": 1, + "ssa_edges": 35, + "result_locs": 28, + "arg_locs": 2, + "bind_attrs": 5, + "type_checks": 10, + "needed_attrs": 21, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 3 + ], + [ + "arith.constant.value=16", + 1 + ], + [ + "arith.constant.value=50", + 1 + ], + [ + "arith.constant.value=63", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "tt.expand_dims.axis=0", + 1 + ], + [ + "tt.expand_dims.axis=1", + 1 + ], + [ + "tt.make_range.end=16", + 1 + ], + [ + "tt.make_range.end=4", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 3 + ], + [ + "tt.reshape.allow_reorder=True", + 2 + ], + [ + "tt.trans.order=(1, 0)", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "views" + ] + }, + "adv_while_nested.ttir": { + "stats": { + "ops": 23, + "implicit_ops": 0, + "blocks": 7, + "values": 24, + "funcs": 1, + "ssa_edges": 32, + "result_locs": 15, + "arg_locs": 5, + "bind_attrs": 4, + "type_checks": 3, + "needed_attrs": 10, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=2", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "zero-result op with a text loc", + 9 + ] + ], + "funcs": [ + "while_nested" + ] + }, + "adv_zero_result.ttir": { + "stats": { + "ops": 28, + "implicit_ops": 1, + "blocks": 3, + "values": 23, + "funcs": 1, + "ssa_edges": 31, + "result_locs": 20, + "arg_locs": 3, + "bind_attrs": 7, + "type_checks": 5, + "needed_attrs": 20, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "arith.constant.value=True", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='add'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='max'", + 1 + ], + [ + "tt.atomic_rmw.scope='gpu'", + 2 + ], + [ + "tt.atomic_rmw.sem='acq_rel'", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 7 + ] + ], + "funcs": [ + "zero_result" + ] + }, + "crafted_attr_dicts.ttir": { + "stats": { + "ops": 14, + "implicit_ops": 0, + "blocks": 3, + "values": 12, + "funcs": 1, + "ssa_edges": 13, + "result_locs": 0, + "arg_locs": 0, + "bind_attrs": 3, + "type_checks": 5, + "needed_attrs": 9, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=-3", + 1 + ], + [ + "arith.constant.value=16", + 1 + ], + [ + "tt.expand_dims.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "tt.reduce.axis=0", + 1 + ] + ], + "funcs": [ + "k" + ] + }, + "crafted_deep_nest.ttir": { + "stats": { + "ops": 30, + "implicit_ops": 0, + "blocks": 11, + "values": 31, + "funcs": 1, + "ssa_edges": 43, + "result_locs": 0, + "arg_locs": 0, + "bind_attrs": 3, + "type_checks": 4, + "needed_attrs": 13, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpf.predicate='ogt'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 2 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 2 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "tt.reduce.axis=0", + 1 + ] + ], + "funcs": [ + "k" + ] + }, + "crafted_empty_bodies.ttir": { + "stats": { + "ops": 17, + "implicit_ops": 6, + "blocks": 8, + "values": 6, + "funcs": 1, + "ssa_edges": 10, + "result_locs": 2, + "arg_locs": 3, + "bind_attrs": 3, + "type_checks": 2, + "needed_attrs": 6, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "zero-result op with a text loc", + 9 + ] + ], + "funcs": [ + "k" + ] + }, + "crafted_empty_else.ttir": { + "stats": { + "ops": 8, + "implicit_ops": 2, + "blocks": 4, + "values": 3, + "funcs": 1, + "ssa_edges": 3, + "result_locs": 0, + "arg_locs": 0, + "bind_attrs": 3, + "type_checks": 1, + "needed_attrs": 4, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=0", + 1 + ] + ], + "funcs": [ + "k" + ] + }, + "crafted_empty_for.ttir": { + "stats": { + "ops": 8, + "implicit_ops": 1, + "blocks": 3, + "values": 5, + "funcs": 1, + "ssa_edges": 5, + "result_locs": 0, + "arg_locs": 0, + "bind_attrs": 3, + "type_checks": 2, + "needed_attrs": 6, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ] + ], + "funcs": [ + "k" + ] + }, + "crafted_fwd_ref_cf.ttir": { + "stats": { + "ops": 8, + "implicit_ops": 0, + "blocks": 4, + "values": 4, + "funcs": 1, + "ssa_edges": 6, + "result_locs": 0, + "arg_locs": 0, + "bind_attrs": 3, + "type_checks": 0, + "needed_attrs": 5, + "cf_edges": 2, + "pred_checks": 0 + }, + "census": [ + [ + "cf.br -> 1 successors", + 2 + ] + ], + "funcs": [ + "k" + ] + }, + "crafted_locs.ttir": { + "stats": { + "ops": 8, + "implicit_ops": 0, + "blocks": 2, + "values": 6, + "funcs": 1, + "ssa_edges": 8, + "result_locs": 4, + "arg_locs": 2, + "bind_attrs": 3, + "type_checks": 1, + "needed_attrs": 4, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=1", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "k" + ] + }, + "crafted_odd_names.ttir": { + "stats": { + "ops": 10, + "implicit_ops": 0, + "blocks": 2, + "values": 8, + "funcs": 1, + "ssa_edges": 12, + "result_locs": 0, + "arg_locs": 0, + "bind_attrs": 3, + "type_checks": 1, + "needed_attrs": 4, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=-1", + 1 + ] + ], + "funcs": [ + "k" + ] + }, + "crafted_same_dest.ttir": { + "stats": { + "ops": 10, + "implicit_ops": 0, + "blocks": 5, + "values": 9, + "funcs": 1, + "ssa_edges": 14, + "result_locs": 0, + "arg_locs": 0, + "bind_attrs": 3, + "type_checks": 0, + "needed_attrs": 7, + "cf_edges": 5, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "cf.br -> 1 successors", + 1 + ], + [ + "cf.cond_br -> 2 successors", + 2 + ] + ], + "funcs": [ + "k" + ] + }, + "crafted_symbols_strings.ttir": { + "stats": { + "ops": 13, + "implicit_ops": 0, + "blocks": 3, + "values": 10, + "funcs": 2, + "ssa_edges": 15, + "result_locs": 0, + "arg_locs": 0, + "bind_attrs": 13, + "type_checks": 0, + "needed_attrs": 13, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='sgt'", + 1 + ], + [ + "tt.elementwise_inline_asm.packed_element=1", + 1 + ] + ], + "funcs": [ + "f{%x} \"q\" (a)", + "k" + ] + }, + "crafted_unicode_strings.ttir": { + "stats": { + "ops": 6, + "implicit_ops": 0, + "blocks": 2, + "values": 3, + "funcs": 1, + "ssa_edges": 3, + "result_locs": 1, + "arg_locs": 2, + "bind_attrs": 7, + "type_checks": 0, + "needed_attrs": 5, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "核" + ] + }, + "golden_add_sm80.ttir": { + "stats": { + "ops": 21, + "implicit_ops": 0, + "blocks": 2, + "values": 21, + "funcs": 1, + "ssa_edges": 26, + "result_locs": 17, + "arg_locs": 4, + "bind_attrs": 5, + "type_checks": 2, + "needed_attrs": 10, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=1024", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=1024", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "add_kernel" + ] + }, + "golden_add_sm90.ttir": { + "stats": { + "ops": 21, + "implicit_ops": 0, + "blocks": 2, + "values": 21, + "funcs": 1, + "ssa_edges": 26, + "result_locs": 17, + "arg_locs": 4, + "bind_attrs": 5, + "type_checks": 2, + "needed_attrs": 10, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=1024", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=1024", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "add_kernel" + ] + }, + "golden_atomic_fmax_sm80.ttir": { + "stats": { + "ops": 29, + "implicit_ops": 0, + "blocks": 2, + "values": 29, + "funcs": 1, + "ssa_edges": 35, + "result_locs": 26, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 6, + "needed_attrs": 20, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='ne'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 4 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=256", + 1 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=True", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='max'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='umin'", + 1 + ], + [ + "tt.atomic_rmw.scope='gpu'", + 2 + ], + [ + "tt.atomic_rmw.sem='acq_rel'", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 3 + ] + ], + "funcs": [ + "atomic_fmax_kernel" + ] + }, + "golden_atomic_fmax_sm90.ttir": { + "stats": { + "ops": 29, + "implicit_ops": 0, + "blocks": 2, + "values": 29, + "funcs": 1, + "ssa_edges": 35, + "result_locs": 26, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 6, + "needed_attrs": 20, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='ne'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 4 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=256", + 1 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=True", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='max'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='umin'", + 1 + ], + [ + "tt.atomic_rmw.scope='gpu'", + 2 + ], + [ + "tt.atomic_rmw.sem='acq_rel'", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 3 + ] + ], + "funcs": [ + "atomic_fmax_kernel" + ] + }, + "golden_atomic_sm80.ttir": { + "stats": { + "ops": 21, + "implicit_ops": 0, + "blocks": 2, + "values": 20, + "funcs": 1, + "ssa_edges": 26, + "result_locs": 17, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 4, + "needed_attrs": 17, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=256", + 1 + ], + [ + "arith.constant.value=True", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='exch'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='fadd'", + 1 + ], + [ + "tt.atomic_rmw.scope='gpu'", + 2 + ], + [ + "tt.atomic_rmw.sem='acq_rel'", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "atomic_kernel" + ] + }, + "golden_atomic_sm90.ttir": { + "stats": { + "ops": 21, + "implicit_ops": 0, + "blocks": 2, + "values": 20, + "funcs": 1, + "ssa_edges": 26, + "result_locs": 17, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 4, + "needed_attrs": 17, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=256", + 1 + ], + [ + "arith.constant.value=True", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='exch'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='fadd'", + 1 + ], + [ + "tt.atomic_rmw.scope='gpu'", + 2 + ], + [ + "tt.atomic_rmw.sem='acq_rel'", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "atomic_kernel" + ] + }, + "golden_cas_sm80.ttir": { + "stats": { + "ops": 7, + "implicit_ops": 0, + "blocks": 2, + "values": 5, + "funcs": 1, + "ssa_edges": 5, + "result_locs": 3, + "arg_locs": 2, + "bind_attrs": 3, + "type_checks": 2, + "needed_attrs": 7, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "tt.atomic_cas.scope='gpu'", + 1 + ], + [ + "tt.atomic_cas.sem='acq_rel'", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "cas_kernel" + ] + }, + "golden_cas_sm90.ttir": { + "stats": { + "ops": 7, + "implicit_ops": 0, + "blocks": 2, + "values": 5, + "funcs": 1, + "ssa_edges": 5, + "result_locs": 3, + "arg_locs": 2, + "bind_attrs": 3, + "type_checks": 2, + "needed_attrs": 7, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "tt.atomic_cas.scope='gpu'", + 1 + ], + [ + "tt.atomic_cas.sem='acq_rel'", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "cas_kernel" + ] + }, + "golden_early_return_loaded_sm80.ttir": { + "stats": { + "ops": 23, + "implicit_ops": 0, + "blocks": 4, + "values": 21, + "funcs": 1, + "ssa_edges": 25, + "result_locs": 17, + "arg_locs": 4, + "bind_attrs": 5, + "type_checks": 3, + "needed_attrs": 13, + "cf_edges": 2, + "pred_checks": 2 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=-1", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "cf.cond_br -> 2 successors", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "early_return_loaded_kernel" + ] + }, + "golden_early_return_pid_sm80.ttir": { + "stats": { + "ops": 20, + "implicit_ops": 0, + "blocks": 4, + "values": 18, + "funcs": 1, + "ssa_edges": 22, + "result_locs": 14, + "arg_locs": 4, + "bind_attrs": 4, + "type_checks": 2, + "needed_attrs": 11, + "cf_edges": 2, + "pred_checks": 2 + }, + "census": [ + [ + "arith.cmpi.predicate='sge'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "cf.cond_br -> 2 successors", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "early_return_pid_kernel" + ] + }, + "golden_gather_sm80.ttir": { + "stats": { + "ops": 22, + "implicit_ops": 0, + "blocks": 2, + "values": 22, + "funcs": 1, + "ssa_edges": 26, + "result_locs": 18, + "arg_locs": 4, + "bind_attrs": 5, + "type_checks": 4, + "needed_attrs": 12, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=256", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "gather_kernel" + ] + }, + "golden_gather_sm90.ttir": { + "stats": { + "ops": 22, + "implicit_ops": 0, + "blocks": 2, + "values": 22, + "funcs": 1, + "ssa_edges": 26, + "result_locs": 18, + "arg_locs": 4, + "bind_attrs": 5, + "type_checks": 4, + "needed_attrs": 12, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=256", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "gather_kernel" + ] + }, + "golden_grid_stride_sm80.ttir": { + "stats": { + "ops": 17, + "implicit_ops": 1, + "blocks": 3, + "values": 16, + "funcs": 1, + "ssa_edges": 18, + "result_locs": 11, + "arg_locs": 4, + "bind_attrs": 4, + "type_checks": 2, + "needed_attrs": 9, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=4", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "grid_stride_kernel" + ] + }, + "golden_guard_then_loop_sm80.ttir": { + "stats": { + "ops": 24, + "implicit_ops": 1, + "blocks": 5, + "values": 21, + "funcs": 1, + "ssa_edges": 24, + "result_locs": 16, + "arg_locs": 4, + "bind_attrs": 4, + "type_checks": 4, + "needed_attrs": 13, + "cf_edges": 2, + "pred_checks": 2 + }, + "census": [ + [ + "arith.cmpi.predicate='sge'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "cf.cond_br -> 2 successors", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 7 + ] + ], + "funcs": [ + "guard_then_loop_kernel" + ] + }, + "golden_if_else_load_sm80.ttir": { + "stats": { + "ops": 25, + "implicit_ops": 0, + "blocks": 4, + "values": 23, + "funcs": 1, + "ssa_edges": 29, + "result_locs": 19, + "arg_locs": 4, + "bind_attrs": 5, + "type_checks": 3, + "needed_attrs": 12, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=256", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "if_else_load_kernel" + ] + }, + "golden_if_else_load_sm90.ttir": { + "stats": { + "ops": 25, + "implicit_ops": 0, + "blocks": 4, + "values": 23, + "funcs": 1, + "ssa_edges": 29, + "result_locs": 19, + "arg_locs": 4, + "bind_attrs": 5, + "type_checks": 3, + "needed_attrs": 12, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=256", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "if_else_load_kernel" + ] + }, + "golden_if_else_offset_sm80.ttir": { + "stats": { + "ops": 16, + "implicit_ops": 0, + "blocks": 2, + "values": 15, + "funcs": 1, + "ssa_edges": 18, + "result_locs": 12, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 2, + "needed_attrs": 9, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "if_else_offset_kernel" + ] + }, + "golden_if_else_offset_sm90.ttir": { + "stats": { + "ops": 16, + "implicit_ops": 0, + "blocks": 2, + "values": 15, + "funcs": 1, + "ssa_edges": 18, + "result_locs": 12, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 2, + "needed_attrs": 9, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "if_else_offset_kernel" + ] + }, + "golden_loop_under_if_sm80.ttir": { + "stats": { + "ops": 26, + "implicit_ops": 2, + "blocks": 4, + "values": 22, + "funcs": 1, + "ssa_edges": 27, + "result_locs": 18, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 4, + "needed_attrs": 12, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "loop_under_if_kernel" + ] + }, + "golden_matmul_bp_s3_sm80.ttir": { + "stats": { + "ops": 116, + "implicit_ops": 0, + "blocks": 3, + "values": 126, + "funcs": 1, + "ssa_edges": 151, + "result_locs": 113, + "arg_locs": 9, + "bind_attrs": 5, + "type_checks": 21, + "needed_attrs": 46, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='sge'", + 6 + ], + [ + "arith.cmpi.predicate='slt'", + 6 + ], + [ + "arith.constant.splat=True", + 5 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 6 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 2 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.expand_dims.axis=0", + 3 + ], + [ + "tt.expand_dims.axis=1", + 3 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.end=64", + 2 + ], + [ + "tt.make_range.start=0", + 3 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_blockptr_kernel" + ] + }, + "golden_matmul_bp_s3_sm90.ttir": { + "stats": { + "ops": 116, + "implicit_ops": 0, + "blocks": 3, + "values": 126, + "funcs": 1, + "ssa_edges": 151, + "result_locs": 113, + "arg_locs": 9, + "bind_attrs": 5, + "type_checks": 21, + "needed_attrs": 46, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='sge'", + 6 + ], + [ + "arith.cmpi.predicate='slt'", + 6 + ], + [ + "arith.constant.splat=True", + 5 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 6 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 2 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.expand_dims.axis=0", + 3 + ], + [ + "tt.expand_dims.axis=1", + 3 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.end=64", + 2 + ], + [ + "tt.make_range.start=0", + 3 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_blockptr_kernel" + ] + }, + "golden_matmul_s1_sm80.ttir": { + "stats": { + "ops": 75, + "implicit_ops": 0, + "blocks": 3, + "values": 85, + "funcs": 1, + "ssa_edges": 99, + "result_locs": 72, + "arg_locs": 9, + "bind_attrs": 5, + "type_checks": 15, + "needed_attrs": 31, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 4 + ], + [ + "arith.constant.splat=True", + 4 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 3 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 2 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.expand_dims.axis=0", + 2 + ], + [ + "tt.expand_dims.axis=1", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 2 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_kernel" + ] + }, + "golden_matmul_s1_sm90.ttir": { + "stats": { + "ops": 75, + "implicit_ops": 0, + "blocks": 3, + "values": 85, + "funcs": 1, + "ssa_edges": 99, + "result_locs": 72, + "arg_locs": 9, + "bind_attrs": 5, + "type_checks": 15, + "needed_attrs": 31, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 4 + ], + [ + "arith.constant.splat=True", + 4 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 3 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 2 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.expand_dims.axis=0", + 2 + ], + [ + "tt.expand_dims.axis=1", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 2 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_kernel" + ] + }, + "golden_matmul_s3_sm80.ttir": { + "stats": { + "ops": 75, + "implicit_ops": 0, + "blocks": 3, + "values": 85, + "funcs": 1, + "ssa_edges": 99, + "result_locs": 72, + "arg_locs": 9, + "bind_attrs": 5, + "type_checks": 15, + "needed_attrs": 31, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 4 + ], + [ + "arith.constant.splat=True", + 4 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 3 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 2 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.expand_dims.axis=0", + 2 + ], + [ + "tt.expand_dims.axis=1", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 2 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_kernel" + ] + }, + "golden_matmul_s3_sm90.ttir": { + "stats": { + "ops": 75, + "implicit_ops": 0, + "blocks": 3, + "values": 85, + "funcs": 1, + "ssa_edges": 99, + "result_locs": 72, + "arg_locs": 9, + "bind_attrs": 5, + "type_checks": 15, + "needed_attrs": 31, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 4 + ], + [ + "arith.constant.splat=True", + 4 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 3 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 2 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.expand_dims.axis=0", + 2 + ], + [ + "tt.expand_dims.axis=1", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 2 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_kernel" + ] + }, + "golden_matmul_tma_s1_sm90.ttir": { + "stats": { + "ops": 31, + "implicit_ops": 0, + "blocks": 3, + "values": 34, + "funcs": 1, + "ssa_edges": 50, + "result_locs": 26, + "arg_locs": 6, + "bind_attrs": 3, + "type_checks": 7, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 2 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_tma_kernel" + ] + }, + "golden_matmul_tma_s3_sm90.ttir": { + "stats": { + "ops": 31, + "implicit_ops": 0, + "blocks": 3, + "values": 34, + "funcs": 1, + "ssa_edges": 50, + "result_locs": 26, + "arg_locs": 6, + "bind_attrs": 3, + "type_checks": 7, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 2 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_tma_kernel" + ] + }, + "golden_matmul_tma_ws_s3_sm90.ttir": { + "stats": { + "ops": 31, + "implicit_ops": 0, + "blocks": 3, + "values": 34, + "funcs": 1, + "ssa_edges": 50, + "result_locs": 26, + "arg_locs": 6, + "bind_attrs": 3, + "type_checks": 7, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 2 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_tma_ws_kernel" + ] + }, + "golden_nested_guard_merge_sm80.ttir": { + "stats": { + "ops": 25, + "implicit_ops": 0, + "blocks": 7, + "values": 21, + "funcs": 1, + "ssa_edges": 27, + "result_locs": 16, + "arg_locs": 5, + "bind_attrs": 4, + "type_checks": 3, + "needed_attrs": 16, + "cf_edges": 7, + "pred_checks": 5 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='sge'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "cf.br -> 1 successors", + 1 + ], + [ + "cf.cond_br -> 2 successors", + 3 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 9 + ] + ], + "funcs": [ + "nested_guard_merge_kernel" + ] + }, + "golden_nested_loops_sm80.ttir": { + "stats": { + "ops": 18, + "implicit_ops": 2, + "blocks": 4, + "values": 16, + "funcs": 1, + "ssa_edges": 21, + "result_locs": 10, + "arg_locs": 4, + "bind_attrs": 4, + "type_checks": 2, + "needed_attrs": 9, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "nested_loops_kernel" + ] + }, + "golden_pid_branch_sm80.ttir": { + "stats": { + "ops": 21, + "implicit_ops": 1, + "blocks": 3, + "values": 18, + "funcs": 1, + "ssa_edges": 22, + "result_locs": 15, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 3, + "needed_attrs": 11, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=256", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "pid_branch_kernel" + ] + }, + "golden_pid_branch_sm90.ttir": { + "stats": { + "ops": 21, + "implicit_ops": 1, + "blocks": 3, + "values": 18, + "funcs": 1, + "ssa_edges": 22, + "result_locs": 15, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 3, + "needed_attrs": 11, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=256", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=256", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "pid_branch_kernel" + ] + }, + "golden_sequential_loops_sm80.ttir": { + "stats": { + "ops": 29, + "implicit_ops": 1, + "blocks": 4, + "values": 28, + "funcs": 1, + "ssa_edges": 34, + "result_locs": 22, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 5, + "needed_attrs": 13, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "sequential_loops_kernel" + ] + }, + "golden_tile2d_sm80.ttir": { + "stats": { + "ops": 40, + "implicit_ops": 0, + "blocks": 2, + "values": 42, + "funcs": 1, + "ssa_edges": 49, + "result_locs": 36, + "arg_locs": 6, + "bind_attrs": 4, + "type_checks": 6, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 2 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '2.000000e+00')", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "tt.expand_dims.axis=0", + 1 + ], + [ + "tt.expand_dims.axis=1", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "tile2d_kernel" + ] + }, + "golden_tile2d_sm90.ttir": { + "stats": { + "ops": 40, + "implicit_ops": 0, + "blocks": 2, + "values": 42, + "funcs": 1, + "ssa_edges": 49, + "result_locs": 36, + "arg_locs": 6, + "bind_attrs": 4, + "type_checks": 6, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 2 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '2.000000e+00')", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "tt.expand_dims.axis=0", + 1 + ], + [ + "tt.expand_dims.axis=1", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "tile2d_kernel" + ] + }, + "kernel_deep_chain.ttir": { + "stats": { + "ops": 1207, + "implicit_ops": 0, + "blocks": 2, + "values": 1205, + "funcs": 1, + "ssa_edges": 2404, + "result_locs": 1203, + "arg_locs": 2, + "bind_attrs": 3, + "type_checks": 1, + "needed_attrs": 5, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=('float', '1.000000e+00')", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "deep_chain" + ] + }, + "kernel_dot_precisions.ttir": { + "stats": { + "ops": 23, + "implicit_ops": 0, + "blocks": 2, + "values": 22, + "funcs": 1, + "ssa_edges": 27, + "result_locs": 19, + "arg_locs": 3, + "bind_attrs": 5, + "type_checks": 5, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=16", + 1 + ], + [ + "tt.dot.inputPrecision='ieee'", + 1 + ], + [ + "tt.dot.inputPrecision='tf32x3'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 2 + ], + [ + "tt.expand_dims.axis=0", + 1 + ], + [ + "tt.expand_dims.axis=1", + 1 + ], + [ + "tt.make_range.end=16", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "dot_precisions" + ] + }, + "kernel_dot_scaled.ttir": { + "stats": { + "ops": 50, + "implicit_ops": 0, + "blocks": 2, + "values": 51, + "funcs": 1, + "ssa_edges": 58, + "result_locs": 46, + "arg_locs": 5, + "bind_attrs": 7, + "type_checks": 13, + "needed_attrs": 23, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 5 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=128", + 2 + ], + [ + "arith.constant.value=2", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "tt.expand_dims.axis=0", + 3 + ], + [ + "tt.expand_dims.axis=1", + 2 + ], + [ + "tt.make_range.end=128", + 1 + ], + [ + "tt.make_range.end=2", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 3 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "dot_scaled_k" + ] + }, + "kernel_eps_consts.ttir": { + "stats": { + "ops": 17, + "implicit_ops": 0, + "blocks": 2, + "values": 16, + "funcs": 1, + "ssa_edges": 17, + "result_locs": 13, + "arg_locs": 3, + "bind_attrs": 5, + "type_checks": 3, + "needed_attrs": 9, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '9.99999996E-13')", + 1 + ], + [ + "arith.constant.value=('float', '9.99999997E-7')", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "eps_consts" + ] + }, + "kernel_unicode_msgs.ttir": { + "stats": { + "ops": 14, + "implicit_ops": 0, + "blocks": 2, + "values": 9, + "funcs": 1, + "ssa_edges": 12, + "result_locs": 8, + "arg_locs": 1, + "bind_attrs": 7, + "type_checks": 3, + "needed_attrs": 10, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpf.predicate='ogt'", + 1 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '1.000000e+00')", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "unicode_msgs" + ] + }, + "nat_dead_if.ttir": { + "stats": { + "ops": 16, + "implicit_ops": 2, + "blocks": 4, + "values": 11, + "funcs": 1, + "ssa_edges": 16, + "result_locs": 8, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 1, + "needed_attrs": 7, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "dead_if" + ] + }, + "nat_empty_loop.ttir": { + "stats": { + "ops": 8, + "implicit_ops": 0, + "blocks": 2, + "values": 7, + "funcs": 1, + "ssa_edges": 7, + "result_locs": 4, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 0, + "needed_attrs": 5, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "empty_loop" + ] + }, + "nat_empty_then.ttir": { + "stats": { + "ops": 12, + "implicit_ops": 2, + "blocks": 4, + "values": 8, + "funcs": 1, + "ssa_edges": 10, + "result_locs": 5, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 0, + "needed_attrs": 6, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "empty_then" + ] + }, + "nat_hint_arange_const.ttir": { + "stats": { + "ops": 10, + "implicit_ops": 0, + "blocks": 2, + "values": 8, + "funcs": 1, + "ssa_edges": 9, + "result_locs": 6, + "arg_locs": 2, + "bind_attrs": 4, + "type_checks": 1, + "needed_attrs": 5, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=16", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "hint_arange_const" + ] + }, + "nat_hint_scalar_const.ttir": { + "stats": { + "ops": 12, + "implicit_ops": 0, + "blocks": 2, + "values": 10, + "funcs": 1, + "ssa_edges": 11, + "result_locs": 8, + "arg_locs": 2, + "bind_attrs": 4, + "type_checks": 2, + "needed_attrs": 7, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=64", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "hint_scalar_const" + ] + }, + "nat_k_uni.ttir": { + "stats": { + "ops": 10, + "implicit_ops": 0, + "blocks": 2, + "values": 7, + "funcs": 1, + "ssa_edges": 8, + "result_locs": 6, + "arg_locs": 1, + "bind_attrs": 4, + "type_checks": 2, + "needed_attrs": 7, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '1.000000e+00')", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "k_uni" + ] + }, + "nat_uni_params.ttir": { + "stats": { + "ops": 10, + "implicit_ops": 0, + "blocks": 2, + "values": 8, + "funcs": 1, + "ssa_edges": 10, + "result_locs": 6, + "arg_locs": 2, + "bind_attrs": 4, + "type_checks": 1, + "needed_attrs": 7, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "uni_params" + ] + }, + "spike_atomics.ttir": { + "stats": { + "ops": 33, + "implicit_ops": 0, + "blocks": 2, + "values": 33, + "funcs": 1, + "ssa_edges": 40, + "result_locs": 30, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 13, + "needed_attrs": 44, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 6 + ], + [ + "arith.constant.value=('float', '1.000000e+00')", + 1 + ], + [ + "arith.constant.value=-3", + 1 + ], + [ + "arith.constant.value=1", + 2 + ], + [ + "arith.constant.value=2", + 2 + ], + [ + "arith.constant.value=240", + 1 + ], + [ + "arith.constant.value=3", + 1 + ], + [ + "arith.constant.value=5", + 1 + ], + [ + "arith.constant.value=7", + 1 + ], + [ + "arith.constant.value=9", + 1 + ], + [ + "arith.constant.value=True", + 1 + ], + [ + "tt.atomic_cas.scope='cta'", + 1 + ], + [ + "tt.atomic_cas.sem='acq_rel'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='add'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='and'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='exch'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='fadd'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='max'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='min'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='or'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='xor'", + 1 + ], + [ + "tt.atomic_rmw.scope='cta'", + 1 + ], + [ + "tt.atomic_rmw.scope='gpu'", + 5 + ], + [ + "tt.atomic_rmw.scope='sys'", + 2 + ], + [ + "tt.atomic_rmw.sem='acq_rel'", + 4 + ], + [ + "tt.atomic_rmw.sem='acquire'", + 1 + ], + [ + "tt.atomic_rmw.sem='relaxed'", + 2 + ], + [ + "tt.atomic_rmw.sem='release'", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 3 + ] + ], + "funcs": [ + "atomics" + ] + }, + "spike_casts.ttir": { + "stats": { + "ops": 26, + "implicit_ops": 0, + "blocks": 2, + "values": 25, + "funcs": 1, + "ssa_edges": 30, + "result_locs": 22, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 2, + "needed_attrs": 10, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 2 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "casts" + ] + }, + "spike_dot.ttir": { + "stats": { + "ops": 26, + "implicit_ops": 0, + "blocks": 2, + "values": 25, + "funcs": 1, + "ssa_edges": 30, + "result_locs": 22, + "arg_locs": 3, + "bind_attrs": 5, + "type_checks": 5, + "needed_attrs": 13, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.expand_dims.axis=0", + 1 + ], + [ + "tt.expand_dims.axis=1", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "dot" + ] + }, + "spike_early_return.ttir": { + "stats": { + "ops": 20, + "implicit_ops": 0, + "blocks": 4, + "values": 18, + "funcs": 1, + "ssa_edges": 22, + "result_locs": 14, + "arg_locs": 4, + "bind_attrs": 4, + "type_checks": 2, + "needed_attrs": 11, + "cf_edges": 2, + "pred_checks": 2 + }, + "census": [ + [ + "arith.cmpi.predicate='sge'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "cf.cond_br -> 2 successors", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "early_return" + ] + }, + "spike_early_return_loop.ttir": { + "stats": { + "ops": 27, + "implicit_ops": 0, + "blocks": 5, + "values": 25, + "funcs": 1, + "ssa_edges": 28, + "result_locs": 20, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 5, + "needed_attrs": 14, + "cf_edges": 2, + "pred_checks": 2 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=3", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "cf.cond_br -> 2 successors", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 7 + ] + ], + "funcs": [ + "early_return_loop" + ] + }, + "spike_for_ptr_iterargs.ttir": { + "stats": { + "ops": 29, + "implicit_ops": 0, + "blocks": 3, + "values": 35, + "funcs": 1, + "ssa_edges": 42, + "result_locs": 26, + "arg_locs": 5, + "bind_attrs": 5, + "type_checks": 6, + "needed_attrs": 14, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 3 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=2", + 1 + ], + [ + "arith.constant.value=64", + 2 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "for_ptr_iterargs" + ] + }, + "spike_i64_index.ttir": { + "stats": { + "ops": 21, + "implicit_ops": 0, + "blocks": 2, + "values": 21, + "funcs": 1, + "ssa_edges": 22, + "result_locs": 17, + "arg_locs": 4, + "bind_attrs": 4, + "type_checks": 3, + "needed_attrs": 9, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=-4294967296", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "i64_index" + ] + }, + "spike_if_yield.ttir": { + "stats": { + "ops": 32, + "implicit_ops": 0, + "blocks": 4, + "values": 31, + "funcs": 1, + "ssa_edges": 40, + "result_locs": 27, + "arg_locs": 4, + "bind_attrs": 5, + "type_checks": 5, + "needed_attrs": 14, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '2.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=2", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "if_yield" + ] + }, + "spike_inline_asm.ttir": { + "stats": { + "ops": 11, + "implicit_ops": 0, + "blocks": 2, + "values": 10, + "funcs": 1, + "ssa_edges": 10, + "result_locs": 8, + "arg_locs": 2, + "bind_attrs": 10, + "type_checks": 1, + "needed_attrs": 14, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "tt.elementwise_inline_asm.packed_element=1", + 2 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 3 + ] + ], + "funcs": [ + "inline_asm" + ] + }, + "spike_misc.ttir": { + "stats": { + "ops": 24, + "implicit_ops": 0, + "blocks": 2, + "values": 21, + "funcs": 1, + "ssa_edges": 26, + "result_locs": 18, + "arg_locs": 3, + "bind_attrs": 6, + "type_checks": 4, + "needed_attrs": 14, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpf.predicate='ogt'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '-1.500000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "tt.get_num_programs.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "misc" + ] + }, + "spike_nested_for.ttir": { + "stats": { + "ops": 23, + "implicit_ops": 0, + "blocks": 4, + "values": 24, + "funcs": 1, + "ssa_edges": 29, + "result_locs": 17, + "arg_locs": 3, + "bind_attrs": 4, + "type_checks": 5, + "needed_attrs": 12, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 2 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "nested_for" + ] + }, + "spike_noinline_call.ttir": { + "stats": { + "ops": 41, + "implicit_ops": 2, + "blocks": 6, + "values": 36, + "funcs": 3, + "ssa_edges": 47, + "result_locs": 27, + "arg_locs": 9, + "bind_attrs": 12, + "type_checks": 7, + "needed_attrs": 24, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 3 + ], + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '2.000000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '3.000000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '4.000000e+00')", + 1 + ], + [ + "arith.constant.value=1", + 2 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 12 + ] + ], + "funcs": [ + "noinline_call", + "corpus_kernels._noinline_store__Pfp32_i32_i32__(2,)cconstexpr_2_d_0_", + "corpus_kernels._noinline_store__Pfp32_i32_i32__(2,)cconstexpr_4_d_0_" + ] + }, + "spike_reduce_scan.ttir": { + "stats": { + "ops": 40, + "implicit_ops": 0, + "blocks": 6, + "values": 42, + "funcs": 1, + "ssa_edges": 56, + "result_locs": 30, + "arg_locs": 12, + "bind_attrs": 5, + "type_checks": 8, + "needed_attrs": 17, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpf.predicate='oeq'", + 1 + ], + [ + "arith.cmpf.predicate='ogt'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=2", + 1 + ], + [ + "arith.constant.value=3", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "tt.reduce.axis=0", + 3 + ], + [ + "tt.scan.axis=0", + 1 + ], + [ + "zero-result op with a text loc", + 11 + ] + ], + "funcs": [ + "reduce_scan" + ] + }, + "spike_spin_while.ttir": { + "stats": { + "ops": 19, + "implicit_ops": 0, + "blocks": 6, + "values": 15, + "funcs": 1, + "ssa_edges": 19, + "result_locs": 10, + "arg_locs": 4, + "bind_attrs": 6, + "type_checks": 3, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 2 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=True", + 1 + ], + [ + "tt.atomic_cas.scope='gpu'", + 1 + ], + [ + "tt.atomic_cas.sem='acquire'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='exch'", + 1 + ], + [ + "tt.atomic_rmw.scope='gpu'", + 1 + ], + [ + "tt.atomic_rmw.sem='release'", + 1 + ], + [ + "zero-result op with a text loc", + 9 + ] + ], + "funcs": [ + "spin_while" + ] + }, + "spike_tile2d_i64.ttir": { + "stats": { + "ops": 39, + "implicit_ops": 0, + "blocks": 2, + "values": 40, + "funcs": 1, + "ssa_edges": 44, + "result_locs": 35, + "arg_locs": 5, + "bind_attrs": 4, + "type_checks": 6, + "needed_attrs": 16, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='slt'", + 2 + ], + [ + "arith.constant.value=16", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "tt.expand_dims.axis=0", + 1 + ], + [ + "tt.expand_dims.axis=1", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "tt.make_range.end=16", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.start=0", + 2 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "tile2d_i64" + ] + } +} diff --git a/tests/golden/ir/expected_3.8.json b/tests/golden/ir/expected_3.8.json new file mode 100644 index 000000000..aed9123ab --- /dev/null +++ b/tests/golden/ir/expected_3.8.json @@ -0,0 +1,684 @@ +{ + "adv_descs.ttir": { + "stats": { + "ops": 15, + "implicit_ops": 0, + "blocks": 2, + "values": 12, + "funcs": 1, + "ssa_edges": 29, + "result_locs": 9, + "arg_locs": 3, + "bind_attrs": 3, + "type_checks": 4, + "needed_attrs": 8, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "tt.descriptor_reduce.kind='add'", + 1 + ], + [ + "tt.make_range.end=32", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "descs" + ] + }, + "adv_zero_result.ttir": { + "stats": { + "ops": 28, + "implicit_ops": 1, + "blocks": 3, + "values": 23, + "funcs": 1, + "ssa_edges": 31, + "result_locs": 20, + "arg_locs": 3, + "bind_attrs": 8, + "type_checks": 5, + "needed_attrs": 21, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpi.predicate='eq'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "arith.constant.value=True", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='add'", + 1 + ], + [ + "tt.atomic_rmw.rmw_op='max'", + 1 + ], + [ + "tt.atomic_rmw.scope='gpu'", + 2 + ], + [ + "tt.atomic_rmw.sem='acq_rel'", + 2 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 7 + ] + ], + "funcs": [ + "zero_result" + ] + }, + "golden_matmul_tma_s1_sm90.ttir": { + "stats": { + "ops": 31, + "implicit_ops": 0, + "blocks": 3, + "values": 34, + "funcs": 1, + "ssa_edges": 50, + "result_locs": 26, + "arg_locs": 6, + "bind_attrs": 3, + "type_checks": 7, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 2 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_tma_kernel" + ] + }, + "golden_matmul_tma_s3_sm90.ttir": { + "stats": { + "ops": 31, + "implicit_ops": 0, + "blocks": 3, + "values": 34, + "funcs": 1, + "ssa_edges": 50, + "result_locs": 26, + "arg_locs": 6, + "bind_attrs": 3, + "type_checks": 7, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 2 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_tma_kernel" + ] + }, + "golden_matmul_tma_ws_s3_sm90.ttir": { + "stats": { + "ops": 31, + "implicit_ops": 0, + "blocks": 3, + "values": 34, + "funcs": 1, + "ssa_edges": 50, + "result_locs": 26, + "arg_locs": 6, + "bind_attrs": 3, + "type_checks": 7, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=0", + 1 + ], + [ + "arith.constant.value=1", + 2 + ], + [ + "arith.constant.value=31", + 1 + ], + [ + "arith.constant.value=32", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "scf.for.unsignedCmp=False", + 1 + ], + [ + "tt.dot.inputPrecision='tf32'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=1", + 1 + ], + [ + "zero-result op with a text loc", + 5 + ] + ], + "funcs": [ + "matmul_tma_ws_kernel" + ] + }, + "kernel_deep_chain.ttir": { + "stats": { + "ops": 1207, + "implicit_ops": 0, + "blocks": 2, + "values": 1205, + "funcs": 1, + "ssa_edges": 2404, + "result_locs": 1203, + "arg_locs": 2, + "bind_attrs": 3, + "type_checks": 1, + "needed_attrs": 5, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.value=('float', '1.000000e+00')", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "deep_chain" + ] + }, + "kernel_dot_precisions.ttir": { + "stats": { + "ops": 23, + "implicit_ops": 0, + "blocks": 2, + "values": 22, + "funcs": 1, + "ssa_edges": 27, + "result_locs": 19, + "arg_locs": 3, + "bind_attrs": 5, + "type_checks": 5, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=16", + 1 + ], + [ + "tt.dot.inputPrecision='ieee'", + 1 + ], + [ + "tt.dot.inputPrecision='tf32x3'", + 1 + ], + [ + "tt.dot.maxNumImpreciseAcc=0", + 2 + ], + [ + "tt.expand_dims.axis=0", + 1 + ], + [ + "tt.expand_dims.axis=1", + 1 + ], + [ + "tt.make_range.end=16", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "dot_precisions" + ] + }, + "kernel_dot_scaled.ttir": { + "stats": { + "ops": 50, + "implicit_ops": 0, + "blocks": 2, + "values": 51, + "funcs": 1, + "ssa_edges": 58, + "result_locs": 46, + "arg_locs": 5, + "bind_attrs": 7, + "type_checks": 13, + "needed_attrs": 23, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 5 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=128", + 2 + ], + [ + "arith.constant.value=2", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "tt.expand_dims.axis=0", + 3 + ], + [ + "tt.expand_dims.axis=1", + 2 + ], + [ + "tt.make_range.end=128", + 1 + ], + [ + "tt.make_range.end=2", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 3 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "dot_scaled_k" + ] + }, + "kernel_eps_consts.ttir": { + "stats": { + "ops": 17, + "implicit_ops": 0, + "blocks": 2, + "values": 16, + "funcs": 1, + "ssa_edges": 17, + "result_locs": 13, + "arg_locs": 3, + "bind_attrs": 5, + "type_checks": 3, + "needed_attrs": 9, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.constant.splat=True", + 1 + ], + [ + "arith.constant.value=('float', '9.99999996E-13')", + 1 + ], + [ + "arith.constant.value=('float', '9.99999997E-7')", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 4 + ] + ], + "funcs": [ + "eps_consts" + ] + }, + "kernel_unicode_msgs.ttir": { + "stats": { + "ops": 14, + "implicit_ops": 0, + "blocks": 2, + "values": 9, + "funcs": 1, + "ssa_edges": 12, + "result_locs": 8, + "arg_locs": 1, + "bind_attrs": 7, + "type_checks": 3, + "needed_attrs": 10, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpf.predicate='ogt'", + 1 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '1.000000e+00')", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "unicode_msgs" + ] + }, + "spike_misc.ttir": { + "stats": { + "ops": 24, + "implicit_ops": 0, + "blocks": 2, + "values": 21, + "funcs": 1, + "ssa_edges": 26, + "result_locs": 18, + "arg_locs": 3, + "bind_attrs": 7, + "type_checks": 4, + "needed_attrs": 15, + "cf_edges": 0, + "pred_checks": 0 + }, + "census": [ + [ + "arith.cmpf.predicate='ogt'", + 1 + ], + [ + "arith.cmpi.predicate='slt'", + 1 + ], + [ + "arith.constant.splat=True", + 2 + ], + [ + "arith.constant.value=('float', '-1.500000e+00')", + 1 + ], + [ + "arith.constant.value=('float', '0.000000e+00')", + 1 + ], + [ + "arith.constant.value=64", + 1 + ], + [ + "tt.get_num_programs.axis=0", + 1 + ], + [ + "tt.get_program_id.axis=0", + 1 + ], + [ + "tt.make_range.end=64", + 1 + ], + [ + "tt.make_range.start=0", + 1 + ], + [ + "zero-result op with a text loc", + 6 + ] + ], + "funcs": [ + "misc" + ] + } +} diff --git a/tests/golden/ir/generate_reader_ttir.py b/tests/golden/ir/generate_reader_ttir.py new file mode 100644 index 000000000..cbacb40bc --- /dev/null +++ b/tests/golden/ir/generate_reader_ttir.py @@ -0,0 +1,109 @@ +"""Regenerate the TTIR reader's regression goldens of the installed Triton release. + + python tests/golden/ir/generate_reader_ttir.py [NAME ...] + +Host-compiles each kernel of ``reader_kernels.py`` (ASTSource + +``triton.compile`` for ``GPUTarget("cuda", 80, 32)``, no GPU needed) into +``reader_ttir/.ttir`` under Triton 3.6 (the base release, whose +goldens every release's tests read) and into ``reader_ttir_/`` +under any other (that release's own printing, which shadows the base +golden of the same name under that release), using a throwaway Triton +cache. These goldens sit apart from ``ttir/``, whose every file the +walk-layer tests pin. + +Not a test module: pytest imports it (python_files = *.py) and finds nothing. +""" + +from __future__ import annotations + +import importlib.util +import os +import sys +import tempfile +from typing import Any + +HERE = os.path.dirname(os.path.abspath(__file__)) +BASE_RELEASE = "3.6" # the release that printed reader_ttir/ + +_F32 = "*fp32" +# name -> (signature, constexprs) +SPECS: dict[str, tuple[dict[str, str], dict[str, Any]]] = { + "p1_variant_delta": ({"x_ptr": _F32, "out_ptr": _F32, "n": "i32"}, {}), + "p2_swap": ({"a_ptr": _F32, "b_ptr": _F32, "n": "i32"}, {}), + "p3_call_guarded": ({"x_ptr": _F32, "n": "i32"}, {}), + "p3_call_offset": ({"x_ptr": _F32, "n": "i32"}, {}), + "p3_call_formals": ({"x_ptr": _F32, "n": "i32"}, {}), + "p4_observed_direct": ({"cnt_ptr": "*i32", "x_ptr": _F32, "n": "i32"}, {}), + "p4_observed_loop": ({"cnt_ptr": "*i32", "x_ptr": _F32, "n": "i32"}, {}), + "p4_observed_delta": ({"cnt_ptr": "*i32", "x_ptr": _F32, "n": "i32"}, {}), + "rv_trunci_alias": ({"x_ptr": "*i32"}, {}), + "rv_i32_wrap": ({"x_ptr": "*i32", "S": "i32"}, {}), + "unsigned_index": ({"x_ptr": _F32, "n": "i32"}, {}), + "rv_inline_asm_store": ({"x_ptr": "*i32", "OFF": "constexpr"}, {"OFF": 4096}), + "loop_two_step_advance": ({"x_ptr": _F32, "n": "i32", "s": "i32"}, {}), + "loop_observed_advance": ({"cnt_ptr": "*i32", "x_ptr": _F32, "n": "i32"}, {}), + "where_pointer": ({"x_ptr": _F32, "n": "i32"}, {}), + "tile3d_shared_arange": ({"x_ptr": _F32, "N": "constexpr"}, {"N": 4}), + "expand_iterarg_3d": ( + {"x_ptr": _F32, "out_ptr": _F32, "n": "i32", "N": "constexpr"}, + {"N": 4}, + ), + "expand_iterarg_mask": ( + {"x_ptr": _F32, "out_ptr": _F32, "n": "i32", "M": "i32", "N": "constexpr"}, + {"N": 4}, + ), + "int_iterarg_offset": ({"x_ptr": _F32, "n": "i32", "B": "constexpr"}, {"B": 8}), + "iv_wrap": ({"x_ptr": _F32, "lo": "i32", "n": "i32"}, {}), + "pure_asm_int_addr": ({"x_ptr": "*i32"}, {}), + "observed_lanes": ({"cnt_ptr": "*i32", "x_ptr": _F32, "N": "constexpr"}, {"N": 4}), +} + + +def _kernels(): + spec = importlib.util.spec_from_file_location( + "reader_kernels", os.path.join(HERE, "reader_kernels.py") + ) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + sys.modules["reader_kernels"] = module # Triton reads the jit fn's module + spec.loader.exec_module(module) + return module + + +def main() -> int: + import triton + from triton.backends.compiler import GPUTarget + from triton.compiler import ASTSource + + only = set(sys.argv[1:]) + kernels = _kernels() + failed = 0 + release = ".".join(triton.__version__.split(".")[:2]) + out = os.path.join( + HERE, "reader_ttir" if release == BASE_RELEASE else f"reader_ttir_{release}" + ) + os.makedirs(out, exist_ok=True) + with tempfile.TemporaryDirectory() as cache: + os.environ["TRITON_CACHE_DIR"] = cache + for name, (sig, consts) in SPECS.items(): + if only and name not in only: + continue + src = ASTSource(fn=getattr(kernels, name), signature=sig, constexprs=consts) + try: + k = triton.compile(src, target=GPUTarget("cuda", 80, 32)) + except Exception as e: # noqa: BLE001 + failed += 1 + print( + f"[{name}] FAILED: {type(e).__name__}: {str(e)[:300]}", + file=sys.stderr, + ) + continue + path = os.path.join(out, f"{name}.ttir") + with open(path, "w", encoding="utf-8") as f: + f.write(k.asm["ttir"]) + print(f"[{name}] wrote {path}") + return 1 if failed else 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/golden/ir/generate_ttir.py b/tests/golden/ir/generate_ttir.py new file mode 100644 index 000000000..2f9027480 --- /dev/null +++ b/tests/golden/ir/generate_ttir.py @@ -0,0 +1,406 @@ +"""Regenerate the kernel-derived TTIR goldens of the installed Triton release. + + python tests/golden/ir/generate_ttir.py [NAME ...] + +Host-compiles each kernel below (ASTSource + ``triton.compile``, no GPU +needed, a throwaway Triton cache) into the installed release's directory: +``ttir/`` for 3.6, the base release, ``ttir_/`` for any other. The +goldens' locs name the kernels' lines: a line added above the first kernel +moves every loc, and the base goldens no longer regenerate byte for byte, +so notes (BASE_RELEASE, the copies, RESPELLED) go below the kernels. + +Not a test module: pytest imports it (python_files = *.py) and finds nothing. +""" + +from __future__ import annotations + +import os +import sys +import tempfile +from typing import Any + +import triton +import triton.language as tl + +HERE = os.path.dirname(os.path.abspath(__file__)) + + +@triton.jit +def dot_precisions(a_ptr, b_ptr, c_ptr, BLOCK: tl.constexpr): + # fp32 inputs: `ieee` is the printer-elided default, tf32x3 prints + offs = tl.arange(0, BLOCK) + idx = offs[:, None] * BLOCK + offs[None, :] + a = tl.load(a_ptr + idx) + b = tl.load(b_ptr + idx) + c = tl.dot(a, b, input_precision="ieee") + d = tl.dot(a, b, input_precision="tf32x3") + tl.store(c_ptr + idx, c + d) + + +@triton.jit +def eps_consts(x_ptr, s_ptr, out_ptr, BLOCK: tl.constexpr): + # uppercase-E float literals: a scalar constant and a dense splat + offs = tl.arange(0, BLOCK) + x = tl.load(x_ptr + offs) + s = tl.load(s_ptr) + 1e-6 + tl.store(out_ptr + offs, x * s + 1e-12) + + +@triton.jit +def unicode_msgs(x_ptr, BLOCK: tl.constexpr): + # a non-ASCII tt.assert message (debug=True); device_print prefixes must be ASCII + offs = tl.arange(0, BLOCK) + x = tl.load(x_ptr + offs) + tl.device_assert(x > 0, "错误: π must be > 0") + tl.device_print("x=", x) + tl.store(x_ptr + offs, x + 1) + + +@triton.jit +def deep_chain(out_ptr, s, N: tl.constexpr): + # a def chain deeper than Python's default recursion limit + pid = tl.program_id(0) + off = pid + for _ in tl.static_range(N): + off = off * s + pid + tl.store(out_ptr + off, 1.0) + + +@triton.jit +def dot_scaled_k(a_ptr, as_ptr, b_ptr, bs_ptr, c_ptr, M: tl.constexpr, N: tl.constexpr, K: tl.constexpr): # fmt: skip + # tt.dot_scaled prints `%a scale %as, %b scale %bs, %c`: ODS (a, b, c, as, bs) + rm = tl.arange(0, M) + rn = tl.arange(0, N) + rk = tl.arange(0, K) + rs = tl.arange(0, K // 32) + a = tl.load(a_ptr + rm[:, None] * K + rk[None, :]) + b = tl.load(b_ptr + rk[:, None] * N + rn[None, :]) + a_scale = tl.load(as_ptr + rm[:, None] * (K // 32) + rs[None, :]) + b_scale = tl.load(bs_ptr + rn[:, None] * (K // 32) + rs[None, :]) + c = tl.dot_scaled(a, a_scale, "e4m3", b, b_scale, "e4m3") + tl.store(c_ptr + rm[:, None] * N + rn[None, :], c) + + +# ── the release directories, the copied goldens, RESPELLED ── +# +# BASE_RELEASE printed ttir/, whose goldens are read under every release; a +# later release's ttir_/ (e.g. ttir_3.8/) holds that release's own +# printing of the same goldens, which shadows the base copy of the same name +# under that release. +# +# ``SPECS`` are the ``kernel_`` goldens. The other base goldens are +# copies: ``golden_*`` from #361's ``tests/golden/ttgir/*.ttir``, ``spike_*`` +# from the D10a spike corpus, ``adv_*`` / ``nat_*`` from its independent +# review, and ``crafted_*`` are hand-written TTIR for shapes no kernel prints +# reliably (empty region bodies, generic form, quoted symbols, loc forms, cf +# edge cases). ``RESPELLED`` rebuilds from their kernels, for every release +# but the base one, the copies a later release prints differently in a way +# its tests must see: a syntax the base copy does not parse under (3.8: the +# ``!tt.tensordesc`` type), or an op that release's reader reads differently +# (3.8: ``tl.debug_barrier()`` is ``ttg.barrier all``, not ``gpu.barrier``). + +BASE_RELEASE = "3.6" # the release that printed ttir/ + +# The SPECS kernels keep the lines ttir/'s goldens were printed from (their +# locs), dot_scaled_k's one-line signature included (hence ``fmt: skip``). + + +# ── RESPELLED: the kernels behind base copies a later printer respells ── + + +@triton.jit +def descs(a_ptr, M, N, BM: tl.constexpr, BN: tl.constexpr): + # adv_descs (the D10a review's adv_kernels.py): device-side tensor + # descriptors, descriptor_load / store / reduce / gather / scatter + d = tl.make_tensor_descriptor(a_ptr, [M, N], [N, 1], [BM, BN]) + x = d.load([0, BN]) + d.store([BM, 0], x) + d.atomic_add([BM, BN], x) + d1 = tl.make_tensor_descriptor(a_ptr, [M, N], [N, 1], [1, BN]) + rows = tl.arange(0, BM) + g = d1.gather(rows, 0) + d1.scatter(g, rows, BN) + + +@triton.jit +def matmul_tma_kernel( + a_ptr, + b_ptr, + c_ptr, + M, + N, + K, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, +): + # golden_matmul_tma_s{1,3}_sm90 (#361's tests/golden/ttgir/generate_golden.py) + pid_m = tl.program_id(0) + pid_n = tl.program_id(1) + a_desc = tl.make_tensor_descriptor( + a_ptr, shape=[M, K], strides=[K, 1], block_shape=[BLOCK_M, BLOCK_K] + ) + b_desc = tl.make_tensor_descriptor( + b_ptr, shape=[K, N], strides=[N, 1], block_shape=[BLOCK_K, BLOCK_N] + ) + c_desc = tl.make_tensor_descriptor( + c_ptr, shape=[M, N], strides=[N, 1], block_shape=[BLOCK_M, BLOCK_N] + ) + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + for k in range(0, tl.cdiv(K, BLOCK_K)): + a = a_desc.load([pid_m * BLOCK_M, k * BLOCK_K]) + b = b_desc.load([k * BLOCK_K, pid_n * BLOCK_N]) + acc += tl.dot(a, b) + c_desc.store([pid_m * BLOCK_M, pid_n * BLOCK_N], acc.to(tl.float16)) + + +@triton.jit +def matmul_tma_ws_kernel( + a_ptr, + b_ptr, + c_ptr, + M, + N, + K, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, +): + # golden_matmul_tma_ws_s3_sm90 (#361): the warp-specialized loop + pid_m = tl.program_id(0) + pid_n = tl.program_id(1) + a_desc = tl.make_tensor_descriptor( + a_ptr, shape=[M, K], strides=[K, 1], block_shape=[BLOCK_M, BLOCK_K] + ) + b_desc = tl.make_tensor_descriptor( + b_ptr, shape=[K, N], strides=[N, 1], block_shape=[BLOCK_K, BLOCK_N] + ) + c_desc = tl.make_tensor_descriptor( + c_ptr, shape=[M, N], strides=[N, 1], block_shape=[BLOCK_M, BLOCK_N] + ) + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + for k in tl.range(0, tl.cdiv(K, BLOCK_K), warp_specialize=True): + a = a_desc.load([pid_m * BLOCK_M, k * BLOCK_K]) + b = b_desc.load([k * BLOCK_K, pid_n * BLOCK_N]) + acc += tl.dot(a, b) + c_desc.store([pid_m * BLOCK_M, pid_n * BLOCK_N], acc.to(tl.float16)) + + +@triton.jit +def zero_result(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + # adv_zero_result (the D10a review's adv_kernels.py): zero-result ops + pid = tl.program_id(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + v = tl.load(x_ptr + offs, mask=offs < n) + tl.device_assert(offs < n + BLOCK, "offs {bad} } { loc(") + tl.device_print("pid=", pid, offs, hex=True) + tl.debug_barrier() + if pid == 0: + tl.atomic_add(out_ptr, 1) + tl.atomic_max(out_ptr + offs, v.to(tl.int32), mask=offs < n) + tl.store(out_ptr + offs, v.to(tl.int32), mask=offs < n) + + +@triton.jit +def misc(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + # spike_misc (the D10a spike's corpus_kernels.py) + pid = tl.program_id(0) + npg = tl.num_programs(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + offs = tl.max_contiguous(tl.multiple_of(offs, BLOCK), BLOCK) + v = tl.load(x_ptr + offs, mask=offs < n, other=-1.5) + tl.debug_barrier() + tl.device_print("v{x} loc(", npg) + tl.store(out_ptr + offs, tl.where(v > 0, v, 0.0), mask=offs < n) + + +# (kernel, signature, constexprs, compute capability, compile options, +# ASTSource attrs) +_Spec = tuple[Any, dict[str, str], dict[str, int], int, dict[str, Any], dict] + +SPECS: dict[str, _Spec] = { + "dot_precisions": ( + dot_precisions, + {"a_ptr": "*fp32", "b_ptr": "*fp32", "c_ptr": "*fp32", "BLOCK": "constexpr"}, + {"BLOCK": 16}, + 80, + {}, + {}, + ), + "eps_consts": ( + eps_consts, + {"x_ptr": "*fp32", "s_ptr": "*fp32", "out_ptr": "*fp32", "BLOCK": "constexpr"}, + {"BLOCK": 64}, + 80, + {}, + {}, + ), + "unicode_msgs": ( + unicode_msgs, + {"x_ptr": "*fp32", "BLOCK": "constexpr"}, + {"BLOCK": 64}, + 80, + {"debug": True}, + {}, + ), + "deep_chain": ( + deep_chain, + {"out_ptr": "*fp32", "s": "i32", "N": "constexpr"}, + {"N": 600}, + 80, + {}, + {}, + ), + "dot_scaled": ( + dot_scaled_k, + { + "a_ptr": "*fp8e4nv", + "as_ptr": "*u8", + "b_ptr": "*fp8e4nv", + "bs_ptr": "*u8", + "c_ptr": "*fp32", + "M": "constexpr", + "N": "constexpr", + "K": "constexpr", + }, + {"M": 128, "N": 128, "K": 64}, + 100, + {}, + {}, + ), +} + + +_TMA_SIG = { + "a_ptr": "*fp16", + "b_ptr": "*fp16", + "c_ptr": "*fp16", + "M": "i32", + "N": "i32", + "K": "i32", + "BLOCK_M": "constexpr", + "BLOCK_N": "constexpr", + "BLOCK_K": "constexpr", +} +_TMA_CONST = {"BLOCK_M": 64, "BLOCK_N": 64, "BLOCK_K": 32} +# divisibility 16 on the pointers and M, N, K, as #361's generator set it +_TMA_ATTRS = {(i,): [["tt.divisibility", 16]] for i in range(6)} + +RESPELLED: dict[str, _Spec] = { + "adv_zero_result": ( + zero_result, + {"x_ptr": "*fp32", "out_ptr": "*i32", "n": "i32", "BLOCK": "constexpr"}, + {"BLOCK": 64}, + 80, + {}, + {}, + ), + "spike_misc": ( + misc, + {"x_ptr": "*fp32", "out_ptr": "*fp32", "n": "i32", "BLOCK": "constexpr"}, + {"BLOCK": 64}, + 80, + {"num_stages": 1}, + {}, + ), + "adv_descs": ( + descs, + { + "a_ptr": "*fp16", + "M": "i32", + "N": "i32", + "BM": "constexpr", + "BN": "constexpr", + }, + {"BM": 32, "BN": 32}, + 100, + {}, + {}, + ), + "golden_matmul_tma_s1_sm90": ( + matmul_tma_kernel, + _TMA_SIG, + _TMA_CONST, + 90, + {"num_stages": 1}, + _TMA_ATTRS, + ), + "golden_matmul_tma_s3_sm90": ( + matmul_tma_kernel, + _TMA_SIG, + _TMA_CONST, + 90, + {"num_stages": 3}, + _TMA_ATTRS, + ), + "golden_matmul_tma_ws_s3_sm90": ( + matmul_tma_ws_kernel, + _TMA_SIG, + _TMA_CONST, + 90, + {"num_stages": 3}, + _TMA_ATTRS, + ), +} + + +def release() -> str: + """The installed Triton's minor release ("3.6").""" + return ".".join(triton.__version__.split(".")[:2]) + + +def out_dir(rel: str) -> str: + """The golden directory of release ``rel``.""" + return os.path.join(HERE, "ttir" if rel == BASE_RELEASE else f"ttir_{rel}") + + +def jobs(rel: str) -> dict[str, _Spec]: + """Golden name -> spec of what release ``rel`` prints into out_dir(rel).""" + todo = {f"kernel_{name}": spec for name, spec in SPECS.items()} + if rel != BASE_RELEASE: # the base release holds these as copies + todo.update(RESPELLED) + return todo + + +def ttir(spec: _Spec) -> str: + """The TTIR the installed Triton prints for ``spec`` (a host compile).""" + from triton.backends.compiler import GPUTarget + from triton.compiler import ASTSource + + fn, sig, consts, cc, options, attrs = spec + src = ASTSource(fn=fn, signature=sig, constexprs=consts, attrs=attrs) + k = triton.compile( + src, target=GPUTarget("cuda", cc, 32), options={"num_warps": 4, **options} + ) + return k.asm["ttir"] + + +def main() -> int: + only = set(sys.argv[1:]) + rel = release() + out = out_dir(rel) + failed = 0 + os.makedirs(out, exist_ok=True) + with tempfile.TemporaryDirectory() as cache: + os.environ["TRITON_CACHE_DIR"] = cache + for name, spec in jobs(rel).items(): + if only and name not in only and name.removeprefix("kernel_") not in only: + continue + try: + text = ttir(spec) + except Exception as e: # noqa: BLE001 + failed += 1 + print( + f"[{name}] FAILED: {type(e).__name__}: {str(e)[:300]}", + file=sys.stderr, + ) + continue + path = os.path.join(out, f"{name}.ttir") + with open(path, "w", encoding="utf-8") as f: + f.write(text) + print(f"[{name}] wrote {path}") + return 1 if failed else 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/golden/ir/reader_kernels.py b/tests/golden/ir/reader_kernels.py new file mode 100644 index 000000000..07a55963f --- /dev/null +++ b/tests/golden/ir/reader_kernels.py @@ -0,0 +1,248 @@ +"""Kernels behind the TTIR reader's regression goldens in ``reader_ttir/``. + +The audit probes (``ir_mode_audit/probes/``: p1_variant_delta, p2_swap, +p3_noinline_call, p4_observed_iterarg, rv_trunci_alias, rv_i32_wrap, +rv_inline_asm_store), each a shape #361's reader misread, plus a few +shapes the new reader models on purpose. ``generate_reader_ttir.py`` +host-compiles them; ``tests/unit/ir/test_ttir_reader.py`` reads the +result. Editing a kernel moves its source lines: regenerate the goldens. + +Not a test module: pytest imports it (python_files = *.py) and finds nothing. +""" + +from __future__ import annotations + +import triton +import triton.language as tl + + +# ── p1: a loop-carried pointer advanced by a loop-VARIANT amount ── +@triton.jit +def p1_variant_delta(x_ptr, out_ptr, n): + p = x_ptr + 20 + for k in range(-8, n): # n is a runtime scalar + v = tl.load(p) + tl.store(out_ptr, v) + p += k # the advance depends on the induction variable + + +# ── p2: two pointer iter_args swapped in the scf.yield ── +@triton.jit +def p2_swap(a_ptr, b_ptr, n): + p = a_ptr + q = b_ptr + 100 + for i in range(0, n): + tl.store(p, 1.0) + p, q = q, p # ping-pong between two buffers + + +# ── p3: a noinline tt.call (the callee body is another tt.func) ── +@triton.jit(noinline=True) +def _p3_helper(x_ptr, pid): + tl.store(x_ptr + pid, 1.0) + + +@triton.jit +def p3_call_guarded(x_ptr, n): + # callee formal names match caller names; the call sits under `if pid < n` + pid = tl.program_id(0) + if pid < n: + _p3_helper(x_ptr, pid) + + +@triton.jit +def p3_call_offset(x_ptr, n): + # the ACTUAL argument differs from the caller value sharing the formal's name + pid = tl.program_id(0) + _p3_helper(x_ptr, pid + 100) + + +@triton.jit(noinline=True) +def _p3_helper_c(dst, i): + tl.store(dst + i, 1.0) + + +@triton.jit +def p3_call_formals(x_ptr, n): + # callee formals match no caller name + pid = tl.program_id(0) + _p3_helper_c(x_ptr, pid + 100) + + +# ── p4: an atomic observation reaching an address ── +@triton.jit +def p4_observed_direct(cnt_ptr, x_ptr, n): + old = tl.atomic_add(cnt_ptr, 1) + p = x_ptr + old % n + v = tl.load(p) + tl.store(p, v + 1.0) + + +@triton.jit +def p4_observed_loop(cnt_ptr, x_ptr, n): + # the same address, carried through an scf.for iter_arg (offset0) + old = tl.atomic_add(cnt_ptr, 1) + p = x_ptr + old % n + for i in range(0, 4): + v = tl.load(p) + tl.store(p, v + 1.0) + p += 1 + + +@triton.jit +def p4_observed_delta(cnt_ptr, x_ptr, n): + # the observation in the per-iteration DELTA instead of offset0 + old = tl.atomic_add(cnt_ptr, 1) + step = old % n + p = x_ptr + for i in range(0, 4): + v = tl.load(p) + tl.store(p, v + 1.0) + p += step + + +# ── D9: integer widths ── +@triton.jit +def rv_trunci_alias(x_ptr): + # trunc_i32(pid_i64 * 2**32) is 0 for every pid: every program stores x[0] + pid = tl.program_id(0).to(tl.int64) + off = (pid * 4294967296).to(tl.int32) + tl.store(x_ptr + off, 1) + + +@triton.jit +def rv_i32_wrap(x_ptr, S): + # (pid * S) * S wraps to 0 in i32 for S = 65536 + pid = tl.program_id(0) + off = (pid * S) * S + tl.store(x_ptr + off, 1) + + +@triton.jit +def unsigned_index(x_ptr, n): + # divui, cmpi ult and extui read their operands unsigned + pid = tl.program_id(0).to(tl.uint32) + q = pid // 3 + m = pid < n.to(tl.uint32) + tl.store(x_ptr + q, 1.0, mask=m) + + +# ── inline asm ── +@triton.jit +def rv_inline_asm_store(x_ptr, OFF: tl.constexpr): + # an impure asm st.global through a pointer cast to i64: a memory access + # the graph cannot see + p = (x_ptr + OFF).to(tl.int64, bitcast=True) + v = tl.full((), 7, tl.int32) + tl.inline_asm_elementwise( + "st.global.b32 [$1], $2; mov.b32 $0, 0;", + "=r,l,r", + [p, v], + dtype=tl.int32, + is_pure=False, + pack=1, + ) + + +# ── shapes the reader models ── +@triton.jit +def loop_two_step_advance(x_ptr, n, s): + # two addptrs per iteration: the advance is their (loop-invariant) sum + p = x_ptr + for i in range(0, n): + tl.store(p, 1.0) + p += s + p += 2 + + +@triton.jit +def loop_observed_advance(cnt_ptr, x_ptr, n): + # the advance is an atomic observed inside the loop: loop-variant + p = x_ptr + for i in range(0, n): + tl.store(p, 1.0) + p += tl.atomic_add(cnt_ptr, 1) + + +@triton.jit +def where_pointer(x_ptr, n): + # arith.select over two pointers of one base + offs = tl.arange(0, 16) + p = tl.where(offs < n, x_ptr + offs, x_ptr + 100) + tl.store(p, 1.0) + + +@triton.jit +def tile3d_shared_arange(x_ptr, N: tl.constexpr): + # one make_range on all three dims of a tile: three lane variables + r = tl.arange(0, N) + off = r[:, None, None] * (N * N) + r[None, :, None] * N + r[None, None, :] + tl.store(x_ptr + off, 1.0) + + +# ── review probes (ir_mode_audit/probes_phase2/ttir-reader/k_kernels.py) ── +@triton.jit +def expand_iterarg_3d(x_ptr, out_ptr, n, N: tl.constexpr): + # a loop-carried [N, N] pointer tile expanded to 3D inside the loop, next + # to the same make_range at dim 0: the load reads j*N + l - i*N + r = tl.arange(0, N) + p = x_ptr + r[:, None] * N + r[None, :] + for k in range(0, n): + q = p[None, :, :] - r[:, None, None] * N + v = tl.load(q) + o = out_ptr + r[:, None, None] * N * N + r[None, :, None] * N + r[None, None, :] + tl.store(o, v) + p += 1 + + +@triton.jit +def expand_iterarg_mask(x_ptr, out_ptr, n, M, N: tl.constexpr): + # a loop-carried 1D pointer expanded to 2D inside the loop, masked by the + # same make_range at the pointer's lane (dim 1) + r = tl.arange(0, N) + p = x_ptr + r + for k in range(0, n): + q = p[None, :] + r[:, None] * 0 + v = tl.load(q, mask=r[None, :] < M) + tl.store(out_ptr + r[:, None] * N + r[None, :], v) + p += N + + +@triton.jit +def int_iterarg_offset(x_ptr, n, B: tl.constexpr): + # an integer offset carried by the loop: a loop-variant address + offs = tl.arange(0, B) + for k in range(0, n): + tl.store(x_ptr + offs, 1.0) + offs += B + + +@triton.jit +def iv_wrap(x_ptr, lo, n): + # the induction variable's increment wraps in i32 when n is near INT32_MAX + for k in range(lo, n, 1 << 20): + tl.store(x_ptr + k, 1.0) + + +@triton.jit +def pure_asm_int_addr(x_ptr): + # a "pure" asm handed the address as an integer stores through it + a = x_ptr.to(tl.int64, bitcast=False) + r = tl.inline_asm_elementwise( + "st.global.b32 [$1], $2; mov.b32 $0, 0;", + "=r,l,r", + [a, tl.full([], 7, tl.int32)], + dtype=tl.int32, + is_pure=True, + pack=1, + ) + tl.store(x_ptr + 4096, r) + + +@triton.jit +def observed_lanes(cnt_ptr, x_ptr, N: tl.constexpr): + # two lanes of one tensor atomic's old values in one address + r = tl.arange(0, N) + old = tl.atomic_add(cnt_ptr + r, 1) + c = tl.minimum(tl.maximum(old, 0), 10) + tl.store(x_ptr + c[:, None] - c[None, :], 1.0) diff --git a/tests/golden/ir/reader_ttir/expand_iterarg_3d.ttir b/tests/golden/ir/reader_ttir/expand_iterarg_3d.ttir new file mode 100644 index 000000000..3eefcca98 --- /dev/null +++ b/tests/golden/ir/reader_ttir/expand_iterarg_3d.ttir @@ -0,0 +1,94 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":185:0) +#loc25 = loc("x_ptr"(#loc)) +#loc26 = loc("out_ptr"(#loc)) +#loc27 = loc("n"(#loc)) +module { + tt.func public @expand_iterarg_3d(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<0> : tensor<4x1x1xi32> loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + %cst_0 = arith.constant dense<1> : tensor<4x4xi32> loc(#loc1) + %cst_1 = arith.constant dense<4> : tensor<1x4x1xi32> loc(#loc1) + %cst_2 = arith.constant dense<4> : tensor<4x1x1xi32> loc(#loc1) + %p = arith.constant dense<4> : tensor<4x1xi32> loc(#loc28) + %r = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc29) + %p_3 = tt.expand_dims %r {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc30) + %p_4 = arith.muli %p_3, %p : tensor<4x1xi32> loc(#loc28) + %p_5 = tt.splat %x_ptr : !tt.ptr -> tensor<4x1x!tt.ptr> loc(#loc31) + %p_6 = tt.addptr %p_5, %p_4 : tensor<4x1x!tt.ptr>, tensor<4x1xi32> loc(#loc31) + %p_7 = tt.expand_dims %r {axis = 0 : i32} : tensor<4xi32> -> tensor<1x4xi32> loc(#loc32) + %p_8 = tt.broadcast %p_6 : tensor<4x1x!tt.ptr> -> tensor<4x4x!tt.ptr> loc(#loc33) + %p_9 = tt.broadcast %p_7 : tensor<1x4xi32> -> tensor<4x4xi32> loc(#loc33) + %p_10 = tt.addptr %p_8, %p_9 : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc33) + %p_11 = scf.for %k = %c0_i32 to %n step %c1_i32 iter_args(%p_12 = %p_10) -> (tensor<4x4x!tt.ptr>) : i32 { + %q = tt.expand_dims %p_12 {axis = 0 : i32} : tensor<4x4x!tt.ptr> -> tensor<1x4x4x!tt.ptr> loc(#loc35) + %q_13 = tt.expand_dims %p_3 {axis = 2 : i32} : tensor<4x1xi32> -> tensor<4x1x1xi32> loc(#loc36) + %q_14 = arith.muli %q_13, %cst_2 : tensor<4x1x1xi32> loc(#loc37) + %q_15 = tt.broadcast %q : tensor<1x4x4x!tt.ptr> -> tensor<4x4x4x!tt.ptr> loc(#loc38) + %q_16 = arith.subi %cst, %q_14 : tensor<4x1x1xi32> loc(#loc38) + %q_17 = tt.broadcast %q_16 : tensor<4x1x1xi32> -> tensor<4x4x4xi32> loc(#loc38) + %q_18 = tt.addptr %q_15, %q_17 : tensor<4x4x4x!tt.ptr>, tensor<4x4x4xi32> loc(#loc38) + %v = tt.load %q_18 : tensor<4x4x4x!tt.ptr> loc(#loc39) + %o = arith.muli %q_14, %cst_2 : tensor<4x1x1xi32> loc(#loc40) + %o_19 = tt.splat %out_ptr : !tt.ptr -> tensor<4x1x1x!tt.ptr> loc(#loc41) + %o_20 = tt.addptr %o_19, %o : tensor<4x1x1x!tt.ptr>, tensor<4x1x1xi32> loc(#loc41) + %o_21 = tt.expand_dims %p_7 {axis = 2 : i32} : tensor<1x4xi32> -> tensor<1x4x1xi32> loc(#loc42) + %o_22 = arith.muli %o_21, %cst_1 : tensor<1x4x1xi32> loc(#loc43) + %o_23 = tt.broadcast %o_20 : tensor<4x1x1x!tt.ptr> -> tensor<4x4x1x!tt.ptr> loc(#loc44) + %o_24 = tt.broadcast %o_22 : tensor<1x4x1xi32> -> tensor<4x4x1xi32> loc(#loc44) + %o_25 = tt.addptr %o_23, %o_24 : tensor<4x4x1x!tt.ptr>, tensor<4x4x1xi32> loc(#loc44) + %o_26 = tt.expand_dims %p_7 {axis = 1 : i32} : tensor<1x4xi32> -> tensor<1x1x4xi32> loc(#loc45) + %o_27 = tt.broadcast %o_25 : tensor<4x4x1x!tt.ptr> -> tensor<4x4x4x!tt.ptr> loc(#loc46) + %o_28 = tt.broadcast %o_26 : tensor<1x1x4xi32> -> tensor<4x4x4xi32> loc(#loc46) + %o_29 = tt.addptr %o_27, %o_28 : tensor<4x4x4x!tt.ptr>, tensor<4x4x4xi32> loc(#loc46) + tt.store %o_29, %v : tensor<4x4x4x!tt.ptr> loc(#loc21) + %p_30 = tt.addptr %p_12, %cst_0 : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc47) + scf.yield %p_30 : tensor<4x4x!tt.ptr> loc(#loc23) + } loc(#loc34) + tt.return loc(#loc24) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":190:22) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":189:29) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":188:21) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":189:18) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":189:16) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":189:35) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":189:33) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":191:14) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":191:30) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":191:47) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":191:28) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":192:20) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:45) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:22) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:51) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:68) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:49) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:74) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:72) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":194:20) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":195:13) +#loc23 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":195:8) +#loc24 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":190:4) +#loc28 = loc("p"(#loc3)) +#loc29 = loc("r"(#loc4)) +#loc30 = loc("p"(#loc5)) +#loc31 = loc("p"(#loc6)) +#loc32 = loc("p"(#loc7)) +#loc33 = loc("p"(#loc8)) +#loc34 = loc("p"(#loc2)) +#loc35 = loc("q"(#loc9)) +#loc36 = loc("q"(#loc10)) +#loc37 = loc("q"(#loc11)) +#loc38 = loc("q"(#loc12)) +#loc39 = loc("v"(#loc13)) +#loc40 = loc("o"(#loc14)) +#loc41 = loc("o"(#loc15)) +#loc42 = loc("o"(#loc16)) +#loc43 = loc("o"(#loc17)) +#loc44 = loc("o"(#loc18)) +#loc45 = loc("o"(#loc19)) +#loc46 = loc("o"(#loc20)) +#loc47 = loc("p"(#loc22)) diff --git a/tests/golden/ir/reader_ttir/expand_iterarg_mask.ttir b/tests/golden/ir/reader_ttir/expand_iterarg_mask.ttir new file mode 100644 index 000000000..8e0bb92f7 --- /dev/null +++ b/tests/golden/ir/reader_ttir/expand_iterarg_mask.ttir @@ -0,0 +1,62 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":199:0) +#loc18 = loc("x_ptr"(#loc)) +#loc19 = loc("out_ptr"(#loc)) +#loc20 = loc("n"(#loc)) +#loc21 = loc("M"(#loc)) +module { + tt.func public @expand_iterarg_mask(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc)), %M: i32 loc("M"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant dense<4> : tensor<4xi32> loc(#loc2) + %cst_0 = arith.constant dense<4> : tensor<4x1xi32> loc(#loc2) + %r = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc22) + %p = tt.splat %x_ptr : !tt.ptr -> tensor<4x!tt.ptr> loc(#loc23) + %p_1 = tt.addptr %p, %r : tensor<4x!tt.ptr>, tensor<4xi32> loc(#loc23) + %p_2 = scf.for %k = %c0_i32 to %n step %c1_i32 iter_args(%p_3 = %p_1) -> (tensor<4x!tt.ptr>) : i32 { + %q = tt.expand_dims %p_3 {axis = 0 : i32} : tensor<4x!tt.ptr> -> tensor<1x4x!tt.ptr> loc(#loc25) + %q_4 = tt.broadcast %q : tensor<1x4x!tt.ptr> -> tensor<4x4x!tt.ptr> loc(#loc26) + %v = tt.expand_dims %r {axis = 0 : i32} : tensor<4xi32> -> tensor<1x4xi32> loc(#loc27) + %v_5 = tt.splat %M : i32 -> tensor<1x4xi32> loc(#loc28) + %v_6 = arith.cmpi slt, %v, %v_5 : tensor<1x4xi32> loc(#loc28) + %v_7 = tt.broadcast %v_6 : tensor<1x4xi1> -> tensor<4x4xi1> loc(#loc29) + %v_8 = tt.load %q_4, %v_7 : tensor<4x4x!tt.ptr> loc(#loc29) + %0 = tt.expand_dims %r {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc10) + %1 = arith.muli %0, %cst_0 : tensor<4x1xi32> loc(#loc11) + %2 = tt.splat %out_ptr : !tt.ptr -> tensor<4x1x!tt.ptr> loc(#loc12) + %3 = tt.addptr %2, %1 : tensor<4x1x!tt.ptr>, tensor<4x1xi32> loc(#loc12) + %4 = tt.broadcast %3 : tensor<4x1x!tt.ptr> -> tensor<4x4x!tt.ptr> loc(#loc13) + %5 = tt.broadcast %v : tensor<1x4xi32> -> tensor<4x4xi32> loc(#loc13) + %6 = tt.addptr %4, %5 : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc13) + tt.store %6, %v_8 : tensor<4x4x!tt.ptr> loc(#loc14) + %p_9 = tt.addptr %p_3, %cst : tensor<4x!tt.ptr>, tensor<4xi32> loc(#loc30) + scf.yield %p_9 : tensor<4x!tt.ptr> loc(#loc16) + } loc(#loc24) + tt.return loc(#loc17) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":204:22) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":202:21) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":203:16) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":205:14) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":205:25) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":206:30) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":206:41) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":206:20) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":207:29) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":207:40) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":207:27) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":207:44) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":207:56) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":208:13) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":208:8) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":204:4) +#loc22 = loc("r"(#loc3)) +#loc23 = loc("p"(#loc4)) +#loc24 = loc("p"(#loc1)) +#loc25 = loc("q"(#loc5)) +#loc26 = loc("q"(#loc6)) +#loc27 = loc("v"(#loc7)) +#loc28 = loc("v"(#loc8)) +#loc29 = loc("v"(#loc9)) +#loc30 = loc("p"(#loc15)) diff --git a/tests/golden/ir/reader_ttir/int_iterarg_offset.ttir b/tests/golden/ir/reader_ttir/int_iterarg_offset.ttir new file mode 100644 index 000000000..ae28901e5 --- /dev/null +++ b/tests/golden/ir/reader_ttir/int_iterarg_offset.ttir @@ -0,0 +1,31 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":212:0) +#loc9 = loc("x_ptr"(#loc)) +#loc10 = loc("n"(#loc)) +module { + tt.func public @int_iterarg_offset(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant dense<8> : tensor<8xi32> loc(#loc2) + %cst_0 = arith.constant dense<1.000000e+00> : tensor<8xf32> loc(#loc2) + %offs = tt.make_range {end = 8 : i32, start = 0 : i32} : tensor<8xi32> loc(#loc11) + %offs_1 = scf.for %k = %c0_i32 to %n step %c1_i32 iter_args(%offs_2 = %offs) -> (tensor<8xi32>) : i32 { + %0 = tt.splat %x_ptr : !tt.ptr -> tensor<8x!tt.ptr> loc(#loc4) + %1 = tt.addptr %0, %offs_2 : tensor<8x!tt.ptr>, tensor<8xi32> loc(#loc4) + tt.store %1, %cst_0 : tensor<8x!tt.ptr> loc(#loc5) + %offs_3 = arith.addi %offs_2, %cst : tensor<8xi32> loc(#loc13) + scf.yield %offs_3 : tensor<8xi32> loc(#loc7) + } loc(#loc12) + tt.return loc(#loc8) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":215:22) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":214:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":216:25) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":216:31) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":217:16) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":217:8) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":215:4) +#loc11 = loc("offs"(#loc3)) +#loc12 = loc("offs"(#loc1)) +#loc13 = loc("offs"(#loc6)) diff --git a/tests/golden/ir/reader_ttir/iv_wrap.ttir b/tests/golden/ir/reader_ttir/iv_wrap.ttir new file mode 100644 index 000000000..60aa74087 --- /dev/null +++ b/tests/golden/ir/reader_ttir/iv_wrap.ttir @@ -0,0 +1,20 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":221:0) +#loc6 = loc("x_ptr"(#loc)) +#loc7 = loc("lo"(#loc)) +#loc8 = loc("n"(#loc)) +module { + tt.func public @iv_wrap(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %lo: i32 loc("lo"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %c1048576_i32 = arith.constant 1048576 : i32 loc(#loc2) + scf.for %k = %lo to %n step %c1048576_i32 : i32 { + %0 = tt.addptr %x_ptr, %k : !tt.ptr, i32 loc(#loc3) + tt.store %0, %cst : !tt.ptr loc(#loc4) + } loc(#loc2) + tt.return loc(#loc5) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":223:26) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":224:25) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":224:28) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":223:4) diff --git a/tests/golden/ir/reader_ttir/loop_observed_advance.ttir b/tests/golden/ir/reader_ttir/loop_observed_advance.ttir new file mode 100644 index 000000000..f94d4e926 --- /dev/null +++ b/tests/golden/ir/reader_ttir/loop_observed_advance.ttir @@ -0,0 +1,29 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":159:0) +#loc8 = loc("cnt_ptr"(#loc)) +#loc9 = loc("x_ptr"(#loc)) +#loc10 = loc("n"(#loc)) +module { + tt.func public @loop_observed_advance(%cnt_ptr: !tt.ptr loc("cnt_ptr"(#loc)), %x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %true = arith.constant true loc(#loc1) + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + %p = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%p_0 = %x_ptr) -> (!tt.ptr) : i32 { + tt.store %p_0, %cst : !tt.ptr loc(#loc3) + %p_1 = tt.atomic_rmw add, acq_rel, gpu, %cnt_ptr, %c1_i32, %true : (!tt.ptr, i32, i1) -> i32 loc(#loc12) + %p_2 = tt.addptr %p_0, %p_1 : !tt.ptr, i32 loc(#loc13) + scf.yield %p_2 : !tt.ptr loc(#loc6) + } loc(#loc11) + tt.return loc(#loc7) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":162:22) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":163:20) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":164:36) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":164:13) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":164:8) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":162:4) +#loc11 = loc("p"(#loc2)) +#loc12 = loc("p"(#loc4)) +#loc13 = loc("p"(#loc5)) diff --git a/tests/golden/ir/reader_ttir/loop_two_step_advance.ttir b/tests/golden/ir/reader_ttir/loop_two_step_advance.ttir new file mode 100644 index 000000000..0c594e77c --- /dev/null +++ b/tests/golden/ir/reader_ttir/loop_two_step_advance.ttir @@ -0,0 +1,29 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":149:0) +#loc8 = loc("x_ptr"(#loc)) +#loc9 = loc("n"(#loc)) +#loc10 = loc("s"(#loc)) +module { + tt.func public @loop_two_step_advance(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc)), %s: i32 loc("s"(#loc))) attributes {noinline = false} { + %c2_i32 = arith.constant 2 : i32 loc(#loc1) + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %p = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%p_0 = %x_ptr) -> (!tt.ptr) : i32 { + tt.store %p_0, %cst : !tt.ptr loc(#loc3) + %p_1 = tt.addptr %p_0, %s : !tt.ptr, i32 loc(#loc12) + %p_2 = tt.addptr %p_1, %c2_i32 : !tt.ptr, i32 loc(#loc13) + scf.yield %p_2 : !tt.ptr loc(#loc6) + } loc(#loc11) + tt.return loc(#loc7) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":152:22) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":153:20) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":154:13) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":155:13) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":155:8) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":152:4) +#loc11 = loc("p"(#loc2)) +#loc12 = loc("p"(#loc4)) +#loc13 = loc("p"(#loc5)) diff --git a/tests/golden/ir/reader_ttir/observed_lanes.ttir b/tests/golden/ir/reader_ttir/observed_lanes.ttir new file mode 100644 index 000000000..33dc6d91e --- /dev/null +++ b/tests/golden/ir/reader_ttir/observed_lanes.ttir @@ -0,0 +1,45 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":243:0) +#loc12 = loc("cnt_ptr"(#loc)) +#loc13 = loc("x_ptr"(#loc)) +module { + tt.func public @observed_lanes(%cnt_ptr: !tt.ptr loc("cnt_ptr"(#loc)), %x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<0> : tensor<1x4xi32> loc(#loc1) + %cst_0 = arith.constant dense<1.000000e+00> : tensor<4x4xf32> loc(#loc2) + %c = arith.constant dense<10> : tensor<4xi32> loc(#loc14) + %c_1 = arith.constant dense<0> : tensor<4xi32> loc(#loc15) + %old = arith.constant dense : tensor<4xi1> loc(#loc16) + %old_2 = arith.constant dense<1> : tensor<4xi32> loc(#loc16) + %r = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc17) + %old_3 = tt.splat %cnt_ptr : !tt.ptr -> tensor<4x!tt.ptr> loc(#loc18) + %old_4 = tt.addptr %old_3, %r : tensor<4x!tt.ptr>, tensor<4xi32> loc(#loc18) + %old_5 = tt.atomic_rmw add, acq_rel, gpu, %old_4, %old_2, %old : (tensor<4x!tt.ptr>, tensor<4xi32>, tensor<4xi1>) -> tensor<4xi32> loc(#loc16) + %c_6 = arith.maxsi %old_5, %c_1 : tensor<4xi32> loc(#loc15) + %c_7 = arith.minsi %c_6, %c : tensor<4xi32> loc(#loc14) + %0 = tt.expand_dims %c_7 {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc8) + %1 = tt.splat %x_ptr : !tt.ptr -> tensor<4x1x!tt.ptr> loc(#loc9) + %2 = tt.addptr %1, %0 : tensor<4x1x!tt.ptr>, tensor<4x1xi32> loc(#loc9) + %3 = tt.expand_dims %c_7 {axis = 0 : i32} : tensor<4xi32> -> tensor<1x4xi32> loc(#loc10) + %4 = tt.broadcast %2 : tensor<4x1x!tt.ptr> -> tensor<4x4x!tt.ptr> loc(#loc1) + %5 = arith.subi %cst, %3 : tensor<1x4xi32> loc(#loc1) + %6 = tt.broadcast %5 : tensor<1x4xi32> -> tensor<4x4xi32> loc(#loc1) + %7 = tt.addptr %4, %6 : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc1) + tt.store %7, %cst_0 : tensor<4x4x!tt.ptr> loc(#loc2) + tt.return loc(#loc11) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":248:34) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":248:46) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":247:39) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":247:35) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":246:37) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":245:21) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":246:34) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":248:23) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":248:21) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":248:36) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":248:4) +#loc14 = loc("c"(#loc3)) +#loc15 = loc("c"(#loc4)) +#loc16 = loc("old"(#loc5)) +#loc17 = loc("r"(#loc6)) +#loc18 = loc("old"(#loc7)) diff --git a/tests/golden/ir/reader_ttir/p1_variant_delta.ttir b/tests/golden/ir/reader_ttir/p1_variant_delta.ttir new file mode 100644 index 000000000..15e9b81e8 --- /dev/null +++ b/tests/golden/ir/reader_ttir/p1_variant_delta.ttir @@ -0,0 +1,30 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":21:0) +#loc8 = loc("x_ptr"(#loc)) +#loc9 = loc("out_ptr"(#loc)) +#loc10 = loc("n"(#loc)) +module { + tt.func public @p1_variant_delta(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c-8_i32 = arith.constant -8 : i32 loc(#loc1) + %p = arith.constant 20 : i32 loc(#loc11) + %p_0 = tt.addptr %x_ptr, %p : !tt.ptr, i32 loc(#loc11) + %p_1 = scf.for %k = %c-8_i32 to %n step %c1_i32 iter_args(%p_2 = %p_0) -> (!tt.ptr) : i32 { + %v = tt.load %p_2 : !tt.ptr loc(#loc13) + tt.store %out_ptr, %v : !tt.ptr loc(#loc4) + %p_3 = tt.addptr %p_2, %k : !tt.ptr, i32 loc(#loc14) + scf.yield %p_3 : !tt.ptr loc(#loc6) + } loc(#loc12) + tt.return loc(#loc7) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":23:23) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":22:16) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":24:20) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":25:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":26:13) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":26:8) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":23:4) +#loc11 = loc("p"(#loc2)) +#loc12 = loc("p"(#loc1)) +#loc13 = loc("v"(#loc3)) +#loc14 = loc("p"(#loc5)) diff --git a/tests/golden/ir/reader_ttir/p2_swap.ttir b/tests/golden/ir/reader_ttir/p2_swap.ttir new file mode 100644 index 000000000..d6eb4ec0c --- /dev/null +++ b/tests/golden/ir/reader_ttir/p2_swap.ttir @@ -0,0 +1,27 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":31:0) +#loc7 = loc("a_ptr"(#loc)) +#loc8 = loc("b_ptr"(#loc)) +#loc9 = loc("n"(#loc)) +module { + tt.func public @p2_swap(%a_ptr: !tt.ptr loc("a_ptr"(#loc)), %b_ptr: !tt.ptr loc("b_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant 1.000000e+00 : f32 loc(#loc2) + %q = arith.constant 100 : i32 loc(#loc10) + %q_0 = tt.addptr %b_ptr, %q : !tt.ptr, i32 loc(#loc10) + %q_1:2 = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%p = %a_ptr, %q_2 = %q_0) -> (!tt.ptr, !tt.ptr) : i32 { + tt.store %p, %cst : !tt.ptr loc(#loc4) + scf.yield %q_2, %p : !tt.ptr, !tt.ptr loc(#loc5) + } loc(#loc12) + tt.return loc(#loc6) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":34:22) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":33:16) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":35:20) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":36:8) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":34:4) +#loc10 = loc("q"(#loc3)) +#loc11 = loc("p"(#loc1)) +#loc12 = loc("q"(#loc11)) diff --git a/tests/golden/ir/reader_ttir/p3_call_formals.ttir b/tests/golden/ir/reader_ttir/p3_call_formals.ttir new file mode 100644 index 000000000..e60aec8be --- /dev/null +++ b/tests/golden/ir/reader_ttir/p3_call_formals.ttir @@ -0,0 +1,30 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":66:0) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":61:0) +#loc10 = loc("x_ptr"(#loc)) +#loc11 = loc("n"(#loc)) +#loc13 = loc("dst"(#loc6)) +#loc14 = loc("i"(#loc6)) +module { + tt.func public @p3_call_formals(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c100_i32 = arith.constant 100 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc12) + %0 = arith.addi %pid, %c100_i32 : i32 loc(#loc3) + tt.call @reader_kernels._p3_helper_c__Pfp32_i32__(%x_ptr, %0) : (!tt.ptr, i32) -> () loc(#loc4) + tt.return loc(#loc5) + } loc(#loc) + tt.func private @reader_kernels._p3_helper_c__Pfp32_i32__(%dst: !tt.ptr loc("dst"(#loc6)), %i: i32 loc("i"(#loc6))) attributes {noinline = true} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc7) + %0 = tt.addptr %dst, %i : !tt.ptr, i32 loc(#loc8) + tt.store %0, %cst : !tt.ptr loc(#loc7) + tt.return loc(#loc9) + } loc(#loc6) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":68:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":69:30) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":69:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":69:4) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":62:22) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":62:19) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":62:4) +#loc12 = loc("pid"(#loc2)) diff --git a/tests/golden/ir/reader_ttir/p3_call_guarded.ttir b/tests/golden/ir/reader_ttir/p3_call_guarded.ttir new file mode 100644 index 000000000..41b20e704 --- /dev/null +++ b/tests/golden/ir/reader_ttir/p3_call_guarded.ttir @@ -0,0 +1,31 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":46:0) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":41:0) +#loc10 = loc("x_ptr"(#loc)) +#loc11 = loc("n"(#loc)) +#loc13 = loc("x_ptr"(#loc6)) +#loc14 = loc("pid"(#loc6)) +module { + tt.func public @p3_call_guarded(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %pid = tt.get_program_id x : i32 loc(#loc12) + %0 = arith.cmpi slt, %pid, %n : i32 loc(#loc2) + scf.if %0 { + tt.call @reader_kernels._p3_helper__Pfp32_i32__(%x_ptr, %pid) : (!tt.ptr, i32) -> () loc(#loc4) + } loc(#loc3) + tt.return loc(#loc5) + } loc(#loc) + tt.func private @reader_kernels._p3_helper__Pfp32_i32__(%x_ptr: !tt.ptr loc("x_ptr"(#loc6)), %pid: i32 loc("pid"(#loc6))) attributes {noinline = true} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc7) + %0 = tt.addptr %x_ptr, %pid : !tt.ptr, i32 loc(#loc8) + tt.store %0, %cst : !tt.ptr loc(#loc7) + tt.return loc(#loc9) + } loc(#loc6) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":48:24) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":49:13) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":49:7) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":50:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":49:4) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":42:26) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":42:21) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":42:4) +#loc12 = loc("pid"(#loc1)) diff --git a/tests/golden/ir/reader_ttir/p3_call_offset.ttir b/tests/golden/ir/reader_ttir/p3_call_offset.ttir new file mode 100644 index 000000000..c9733a82d --- /dev/null +++ b/tests/golden/ir/reader_ttir/p3_call_offset.ttir @@ -0,0 +1,30 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":54:0) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":41:0) +#loc10 = loc("x_ptr"(#loc)) +#loc11 = loc("n"(#loc)) +#loc13 = loc("x_ptr"(#loc6)) +#loc14 = loc("pid"(#loc6)) +module { + tt.func public @p3_call_offset(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c100_i32 = arith.constant 100 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc12) + %0 = arith.addi %pid, %c100_i32 : i32 loc(#loc3) + tt.call @reader_kernels._p3_helper__Pfp32_i32__(%x_ptr, %0) : (!tt.ptr, i32) -> () loc(#loc4) + tt.return loc(#loc5) + } loc(#loc) + tt.func private @reader_kernels._p3_helper__Pfp32_i32__(%x_ptr: !tt.ptr loc("x_ptr"(#loc6)), %pid: i32 loc("pid"(#loc6))) attributes {noinline = true} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc7) + %0 = tt.addptr %x_ptr, %pid : !tt.ptr, i32 loc(#loc8) + tt.store %0, %cst : !tt.ptr loc(#loc7) + tt.return loc(#loc9) + } loc(#loc6) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":56:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":57:28) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":57:22) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":57:4) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":42:26) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":42:21) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":42:4) +#loc12 = loc("pid"(#loc2)) diff --git a/tests/golden/ir/reader_ttir/p4_observed_delta.ttir b/tests/golden/ir/reader_ttir/p4_observed_delta.ttir new file mode 100644 index 000000000..4d7ebfcc8 --- /dev/null +++ b/tests/golden/ir/reader_ttir/p4_observed_delta.ttir @@ -0,0 +1,38 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":93:0) +#loc11 = loc("cnt_ptr"(#loc)) +#loc12 = loc("x_ptr"(#loc)) +#loc13 = loc("n"(#loc)) +module { + tt.func public @p4_observed_delta(%cnt_ptr: !tt.ptr loc("cnt_ptr"(#loc)), %x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c4_i32 = arith.constant 4 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant 1.000000e+00 : f32 loc(#loc2) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %old = arith.constant true loc(#loc14) + %old_0 = tt.atomic_rmw add, acq_rel, gpu, %cnt_ptr, %c1_i32, %old : (!tt.ptr, i32, i1) -> i32 loc(#loc14) + %step = arith.remsi %old_0, %n : i32 loc(#loc15) + %p = scf.for %i = %c0_i32 to %c4_i32 step %c1_i32 iter_args(%p_1 = %x_ptr) -> (!tt.ptr) : i32 { + %v = tt.load %p_1 : !tt.ptr loc(#loc17) + %0 = arith.addf %v, %cst : f32 loc(#loc6) + tt.store %p_1, %0 : !tt.ptr loc(#loc7) + %p_2 = tt.addptr %p_1, %step : !tt.ptr, i32 loc(#loc18) + scf.yield %p_2 : !tt.ptr loc(#loc9) + } loc(#loc16) + tt.return loc(#loc10) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":98:22) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":95:33) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":96:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":99:20) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":100:24) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":100:20) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":101:13) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":101:8) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":98:4) +#loc14 = loc("old"(#loc3)) +#loc15 = loc("step"(#loc4)) +#loc16 = loc("p"(#loc1)) +#loc17 = loc("v"(#loc5)) +#loc18 = loc("p"(#loc8)) diff --git a/tests/golden/ir/reader_ttir/p4_observed_direct.ttir b/tests/golden/ir/reader_ttir/p4_observed_direct.ttir new file mode 100644 index 000000000..81319f990 --- /dev/null +++ b/tests/golden/ir/reader_ttir/p4_observed_direct.ttir @@ -0,0 +1,30 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":74:0) +#loc9 = loc("cnt_ptr"(#loc)) +#loc10 = loc("x_ptr"(#loc)) +#loc11 = loc("n"(#loc)) +module { + tt.func public @p4_observed_direct(%cnt_ptr: !tt.ptr loc("cnt_ptr"(#loc)), %x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %old = arith.constant 1 : i32 loc(#loc12) + %old_0 = arith.constant true loc(#loc12) + %old_1 = tt.atomic_rmw add, acq_rel, gpu, %cnt_ptr, %old, %old_0 : (!tt.ptr, i32, i1) -> i32 loc(#loc12) + %p = arith.remsi %old_1, %n : i32 loc(#loc13) + %p_2 = tt.addptr %x_ptr, %p : !tt.ptr, i32 loc(#loc14) + %v = tt.load %p_2 : !tt.ptr loc(#loc15) + %0 = arith.addf %v, %cst : f32 loc(#loc6) + tt.store %p_2, %0 : !tt.ptr loc(#loc7) + tt.return loc(#loc8) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":75:33) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":76:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":76:16) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":77:16) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":78:20) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":78:16) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":78:4) +#loc12 = loc("old"(#loc2)) +#loc13 = loc("p"(#loc3)) +#loc14 = loc("p"(#loc4)) +#loc15 = loc("v"(#loc5)) diff --git a/tests/golden/ir/reader_ttir/p4_observed_loop.ttir b/tests/golden/ir/reader_ttir/p4_observed_loop.ttir new file mode 100644 index 000000000..d6b64da98 --- /dev/null +++ b/tests/golden/ir/reader_ttir/p4_observed_loop.ttir @@ -0,0 +1,41 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":82:0) +#loc12 = loc("cnt_ptr"(#loc)) +#loc13 = loc("x_ptr"(#loc)) +#loc14 = loc("n"(#loc)) +module { + tt.func public @p4_observed_loop(%cnt_ptr: !tt.ptr loc("cnt_ptr"(#loc)), %x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c4_i32 = arith.constant 4 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant 1.000000e+00 : f32 loc(#loc2) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %old = arith.constant true loc(#loc15) + %old_0 = tt.atomic_rmw add, acq_rel, gpu, %cnt_ptr, %c1_i32, %old : (!tt.ptr, i32, i1) -> i32 loc(#loc15) + %p = arith.remsi %old_0, %n : i32 loc(#loc16) + %p_1 = tt.addptr %x_ptr, %p : !tt.ptr, i32 loc(#loc17) + %p_2 = scf.for %i = %c0_i32 to %c4_i32 step %c1_i32 iter_args(%p_3 = %p_1) -> (!tt.ptr) : i32 { + %v = tt.load %p_3 : !tt.ptr loc(#loc19) + %0 = arith.addf %v, %cst : f32 loc(#loc7) + tt.store %p_3, %0 : !tt.ptr loc(#loc8) + %p_4 = tt.addptr %p_3, %c1_i32 : !tt.ptr, i32 loc(#loc20) + scf.yield %p_4 : !tt.ptr loc(#loc10) + } loc(#loc18) + tt.return loc(#loc11) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":86:22) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":84:33) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":85:22) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":85:16) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":87:20) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":88:24) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":88:20) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":89:13) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":89:8) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":86:4) +#loc15 = loc("old"(#loc3)) +#loc16 = loc("p"(#loc4)) +#loc17 = loc("p"(#loc5)) +#loc18 = loc("p"(#loc1)) +#loc19 = loc("v"(#loc6)) +#loc20 = loc("p"(#loc9)) diff --git a/tests/golden/ir/reader_ttir/pure_asm_int_addr.ttir b/tests/golden/ir/reader_ttir/pure_asm_int_addr.ttir new file mode 100644 index 000000000..17e54d9e2 --- /dev/null +++ b/tests/golden/ir/reader_ttir/pure_asm_int_addr.ttir @@ -0,0 +1,22 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":228:0) +#loc7 = loc("x_ptr"(#loc)) +module { + tt.func public @pure_asm_int_addr(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %c4096_i32 = arith.constant 4096 : i32 loc(#loc1) + %r = arith.constant 7 : i32 loc(#loc8) + %a = tt.ptr_to_int %x_ptr : !tt.ptr -> i64 loc(#loc9) + %r_0 = tt.elementwise_inline_asm "st.global.b32 [$1], $2; mov.b32 $0, 0;" {constraints = "=r,l,r", packed_element = 1 : i32, pure = true} %a, %r : i64, i32 -> i32 loc(#loc10) + %0 = tt.addptr %x_ptr, %c4096_i32 : !tt.ptr, i32 loc(#loc1) + tt.store %0, %r_0 : !tt.ptr loc(#loc5) + tt.return loc(#loc6) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":239:21) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":234:27) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":230:17) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":234:8) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":239:27) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":239:4) +#loc8 = loc("r"(#loc2)) +#loc9 = loc("a"(#loc3)) +#loc10 = loc("r"(#loc4)) diff --git a/tests/golden/ir/reader_ttir/rv_i32_wrap.ttir b/tests/golden/ir/reader_ttir/rv_i32_wrap.ttir new file mode 100644 index 000000000..1ff9adbbc --- /dev/null +++ b/tests/golden/ir/reader_ttir/rv_i32_wrap.ttir @@ -0,0 +1,23 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":114:0) +#loc7 = loc("x_ptr"(#loc)) +#loc8 = loc("S"(#loc)) +module { + tt.func public @rv_i32_wrap(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %S: i32 loc("S"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc9) + %off = arith.muli %pid, %S : i32 loc(#loc10) + %off_0 = arith.muli %off, %S : i32 loc(#loc11) + %0 = tt.addptr %x_ptr, %off_0 : !tt.ptr, i32 loc(#loc5) + tt.store %0, %c1_i32 : !tt.ptr loc(#loc1) + tt.return loc(#loc6) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":118:26) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":116:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":117:17) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":117:22) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":118:21) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":118:4) +#loc9 = loc("pid"(#loc2)) +#loc10 = loc("off"(#loc3)) +#loc11 = loc("off"(#loc4)) diff --git a/tests/golden/ir/reader_ttir/rv_inline_asm_store.ttir b/tests/golden/ir/reader_ttir/rv_inline_asm_store.ttir new file mode 100644 index 000000000..694a219fd --- /dev/null +++ b/tests/golden/ir/reader_ttir/rv_inline_asm_store.ttir @@ -0,0 +1,20 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":132:0) +#loc6 = loc("x_ptr"(#loc)) +module { + tt.func public @rv_inline_asm_store(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %v = arith.constant 7 : i32 loc(#loc7) + %p = arith.constant 4096 : i32 loc(#loc8) + %p_0 = tt.addptr %x_ptr, %p : !tt.ptr, i32 loc(#loc8) + %p_1 = tt.ptr_to_int %p_0 : !tt.ptr -> i64 loc(#loc9) + %0 = tt.elementwise_inline_asm "st.global.b32 [$1], $2; mov.b32 $0, 0;" {constraints = "=r,l,r", packed_element = 1 : i32, pure = false} %p_1, %v : i64, i32 -> i32 loc(#loc4) + tt.return loc(#loc5) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":136:23) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":135:17) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":135:25) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":140:8) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":137:4) +#loc7 = loc("v"(#loc1)) +#loc8 = loc("p"(#loc2)) +#loc9 = loc("p"(#loc3)) diff --git a/tests/golden/ir/reader_ttir/rv_trunci_alias.ttir b/tests/golden/ir/reader_ttir/rv_trunci_alias.ttir new file mode 100644 index 000000000..464fd554e --- /dev/null +++ b/tests/golden/ir/reader_ttir/rv_trunci_alias.ttir @@ -0,0 +1,27 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":106:0) +#loc9 = loc("x_ptr"(#loc)) +module { + tt.func public @rv_trunci_alias(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c4294967296_i64 = arith.constant 4294967296 : i64 loc(#loc2) + %pid = tt.get_program_id x : i32 loc(#loc10) + %pid_0 = arith.extsi %pid : i32 to i64 loc(#loc11) + %off = arith.muli %pid_0, %c4294967296_i64 : i64 loc(#loc12) + %off_1 = arith.trunci %off : i64 to i32 loc(#loc13) + %0 = tt.addptr %x_ptr, %off_1 : !tt.ptr, i32 loc(#loc7) + tt.store %0, %c1_i32 : !tt.ptr loc(#loc1) + tt.return loc(#loc8) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":110:26) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":108:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":108:30) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":109:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":109:32) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":110:21) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":110:4) +#loc10 = loc("pid"(#loc3)) +#loc11 = loc("pid"(#loc4)) +#loc12 = loc("off"(#loc5)) +#loc13 = loc("off"(#loc6)) diff --git a/tests/golden/ir/reader_ttir/tile3d_shared_arange.ttir b/tests/golden/ir/reader_ttir/tile3d_shared_arange.ttir new file mode 100644 index 000000000..de80ea7ea --- /dev/null +++ b/tests/golden/ir/reader_ttir/tile3d_shared_arange.ttir @@ -0,0 +1,46 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":176:0) +#loc12 = loc("x_ptr"(#loc)) +module { + tt.func public @tile3d_shared_arange(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<1.000000e+00> : tensor<4x4x4xf32> loc(#loc1) + %off = arith.constant dense<4> : tensor<1x4x1xi32> loc(#loc13) + %off_0 = arith.constant dense<16> : tensor<4x1x1xi32> loc(#loc14) + %r = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc15) + %off_1 = tt.expand_dims %r {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc16) + %off_2 = tt.expand_dims %off_1 {axis = 2 : i32} : tensor<4x1xi32> -> tensor<4x1x1xi32> loc(#loc16) + %off_3 = arith.muli %off_2, %off_0 : tensor<4x1x1xi32> loc(#loc14) + %off_4 = tt.expand_dims %r {axis = 0 : i32} : tensor<4xi32> -> tensor<1x4xi32> loc(#loc17) + %off_5 = tt.expand_dims %off_4 {axis = 2 : i32} : tensor<1x4xi32> -> tensor<1x4x1xi32> loc(#loc17) + %off_6 = arith.muli %off_5, %off : tensor<1x4x1xi32> loc(#loc13) + %off_7 = tt.broadcast %off_3 : tensor<4x1x1xi32> -> tensor<4x4x1xi32> loc(#loc18) + %off_8 = tt.broadcast %off_6 : tensor<1x4x1xi32> -> tensor<4x4x1xi32> loc(#loc18) + %off_9 = arith.addi %off_7, %off_8 : tensor<4x4x1xi32> loc(#loc18) + %off_10 = tt.expand_dims %off_4 {axis = 1 : i32} : tensor<1x4xi32> -> tensor<1x1x4xi32> loc(#loc19) + %off_11 = tt.broadcast %off_9 : tensor<4x4x1xi32> -> tensor<4x4x4xi32> loc(#loc20) + %off_12 = tt.broadcast %off_10 : tensor<1x1x4xi32> -> tensor<4x4x4xi32> loc(#loc20) + %off_13 = arith.addi %off_11, %off_12 : tensor<4x4x4xi32> loc(#loc20) + %0 = tt.splat %x_ptr : !tt.ptr -> tensor<4x4x4x!tt.ptr> loc(#loc10) + %1 = tt.addptr %0, %off_13 : tensor<4x4x4x!tt.ptr>, tensor<4x4x4xi32> loc(#loc10) + tt.store %1, %cst : tensor<4x4x4x!tt.ptr> loc(#loc1) + tt.return loc(#loc11) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":180:26) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":179:58) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":179:30) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":178:21) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":179:12) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":179:41) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":179:39) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":179:64) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":179:62) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":180:21) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":180:4) +#loc13 = loc("off"(#loc2)) +#loc14 = loc("off"(#loc3)) +#loc15 = loc("r"(#loc4)) +#loc16 = loc("off"(#loc5)) +#loc17 = loc("off"(#loc6)) +#loc18 = loc("off"(#loc7)) +#loc19 = loc("off"(#loc8)) +#loc20 = loc("off"(#loc9)) diff --git a/tests/golden/ir/reader_ttir/unsigned_index.ttir b/tests/golden/ir/reader_ttir/unsigned_index.ttir new file mode 100644 index 000000000..0812805d4 --- /dev/null +++ b/tests/golden/ir/reader_ttir/unsigned_index.ttir @@ -0,0 +1,26 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":122:0) +#loc8 = loc("x_ptr"(#loc)) +#loc9 = loc("n"(#loc)) +module { + tt.func public @unsigned_index(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %c3_i32 = arith.constant 3 : i32 loc(#loc2) + %pid = tt.get_program_id x : i32 loc(#loc10) + %q = arith.divui %pid, %c3_i32 : i32 loc(#loc11) + %m = arith.cmpi ult, %pid, %n : i32 loc(#loc12) + %0 = arith.extui %q : i32 to i64 loc(#loc6) + %1 = tt.addptr %x_ptr, %0 : !tt.ptr, i64 loc(#loc6) + tt.store %1, %cst, %m : !tt.ptr loc(#loc1) + tt.return loc(#loc7) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":127:24) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":124:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":125:15) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":126:14) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":127:21) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":127:4) +#loc10 = loc("pid"(#loc3)) +#loc11 = loc("q"(#loc4)) +#loc12 = loc("m"(#loc5)) diff --git a/tests/golden/ir/reader_ttir/where_pointer.ttir b/tests/golden/ir/reader_ttir/where_pointer.ttir new file mode 100644 index 000000000..5464fd53c --- /dev/null +++ b/tests/golden/ir/reader_ttir/where_pointer.ttir @@ -0,0 +1,31 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":168:0) +#loc8 = loc("x_ptr"(#loc)) +#loc9 = loc("n"(#loc)) +module { + tt.func public @where_pointer(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<1.000000e+00> : tensor<16xf32> loc(#loc1) + %p = arith.constant 100 : i32 loc(#loc10) + %offs = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32> loc(#loc11) + %p_0 = tt.splat %n : i32 -> tensor<16xi32> loc(#loc12) + %p_1 = arith.cmpi slt, %offs, %p_0 : tensor<16xi32> loc(#loc12) + %p_2 = tt.splat %x_ptr : !tt.ptr -> tensor<16x!tt.ptr> loc(#loc13) + %p_3 = tt.addptr %p_2, %offs : tensor<16x!tt.ptr>, tensor<16xi32> loc(#loc13) + %p_4 = tt.addptr %x_ptr, %p : !tt.ptr, i32 loc(#loc10) + %p_5 = tt.splat %p_4 : !tt.ptr -> tensor<16x!tt.ptr> loc(#loc14) + %p_6 = arith.select %p_1, %p_3, %p_5 : tensor<16xi1>, tensor<16x!tt.ptr> loc(#loc14) + tt.store %p_6, %cst : tensor<16x!tt.ptr> loc(#loc1) + tt.return loc(#loc7) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":172:16) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":171:49) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":170:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":171:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":171:35) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":171:41) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":172:4) +#loc10 = loc("p"(#loc2)) +#loc11 = loc("offs"(#loc3)) +#loc12 = loc("p"(#loc4)) +#loc13 = loc("p"(#loc5)) +#loc14 = loc("p"(#loc6)) diff --git a/tests/golden/ir/reader_ttir_3.8/expand_iterarg_3d.ttir b/tests/golden/ir/reader_ttir_3.8/expand_iterarg_3d.ttir new file mode 100644 index 000000000..2e82deeb3 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/expand_iterarg_3d.ttir @@ -0,0 +1,78 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":185:1) +#loc16 = loc("x_ptr"(#loc)) +#loc17 = loc("out_ptr"(#loc)) +#loc18 = loc("n"(#loc)) +module { + tt.func public @expand_iterarg_3d(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<0> : tensor<4x1x1xi32> loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + %cst_0 = arith.constant dense<1> : tensor<4x4xi32> loc(#loc1) + %cst_1 = arith.constant dense<4> : tensor<1x4x1xi32> loc(#loc1) + %cst_2 = arith.constant dense<4> : tensor<4x1x1xi32> loc(#loc1) + %p = arith.constant dense<4> : tensor<4x1xi32> loc(#loc19) + %r = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc20) + %p_3 = tt.expand_dims %r {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc19) + %p_4 = arith.muli %p_3, %p : tensor<4x1xi32> loc(#loc19) + %p_5 = tt.splat %x_ptr : !tt.ptr -> tensor<4x1x!tt.ptr> loc(#loc21) + %p_6 = tt.addptr %p_5, %p_4 : tensor<4x1x!tt.ptr>, tensor<4x1xi32> loc(#loc21) + %p_7 = tt.expand_dims %r {axis = 0 : i32} : tensor<4xi32> -> tensor<1x4xi32> loc(#loc22) + %p_8 = tt.broadcast %p_6 : tensor<4x1x!tt.ptr> -> tensor<4x4x!tt.ptr> loc(#loc21) + %p_9 = tt.broadcast %p_7 : tensor<1x4xi32> -> tensor<4x4xi32> loc(#loc21) + %p_10 = tt.addptr %p_8, %p_9 : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc21) + %p_11 = scf.for %k = %c0_i32 to %n step %c1_i32 iter_args(%p_12 = %p_10) -> (tensor<4x4x!tt.ptr>) : i32 { + %q = tt.expand_dims %p_12 {axis = 0 : i32} : tensor<4x4x!tt.ptr> -> tensor<1x4x4x!tt.ptr> loc(#loc24) + %q_13 = tt.expand_dims %p_3 {axis = 2 : i32} : tensor<4x1xi32> -> tensor<4x1x1xi32> loc(#loc25) + %q_14 = arith.muli %q_13, %cst_2 : tensor<4x1x1xi32> loc(#loc25) + %q_15 = tt.broadcast %q : tensor<1x4x4x!tt.ptr> -> tensor<4x4x4x!tt.ptr> loc(#loc24) + %q_16 = arith.subi %cst, %q_14 : tensor<4x1x1xi32> loc(#loc24) + %q_17 = tt.broadcast %q_16 : tensor<4x1x1xi32> -> tensor<4x4x4xi32> loc(#loc24) + %q_18 = tt.addptr %q_15, %q_17 : tensor<4x4x4x!tt.ptr>, tensor<4x4x4xi32> loc(#loc24) + %v = tt.load %q_18 : tensor<4x4x4x!tt.ptr> loc(#loc26) + %o = arith.muli %q_14, %cst_2 : tensor<4x1x1xi32> loc(#loc27) + %o_19 = tt.splat %out_ptr : !tt.ptr -> tensor<4x1x1x!tt.ptr> loc(#loc28) + %o_20 = tt.addptr %o_19, %o : tensor<4x1x1x!tt.ptr>, tensor<4x1x1xi32> loc(#loc28) + %o_21 = tt.expand_dims %p_7 {axis = 2 : i32} : tensor<1x4xi32> -> tensor<1x4x1xi32> loc(#loc29) + %o_22 = arith.muli %o_21, %cst_1 : tensor<1x4x1xi32> loc(#loc29) + %o_23 = tt.broadcast %o_20 : tensor<4x1x1x!tt.ptr> -> tensor<4x4x1x!tt.ptr> loc(#loc28) + %o_24 = tt.broadcast %o_22 : tensor<1x4x1xi32> -> tensor<4x4x1xi32> loc(#loc28) + %o_25 = tt.addptr %o_23, %o_24 : tensor<4x4x1x!tt.ptr>, tensor<4x4x1xi32> loc(#loc28) + %o_26 = tt.expand_dims %p_7 {axis = 1 : i32} : tensor<1x4xi32> -> tensor<1x1x4xi32> loc(#loc30) + %o_27 = tt.broadcast %o_25 : tensor<4x4x1x!tt.ptr> -> tensor<4x4x4x!tt.ptr> loc(#loc28) + %o_28 = tt.broadcast %o_26 : tensor<1x1x4xi32> -> tensor<4x4x4xi32> loc(#loc28) + %o_29 = tt.addptr %o_27, %o_28 : tensor<4x4x4x!tt.ptr>, tensor<4x4x4xi32> loc(#loc28) + tt.store %o_29, %v : tensor<4x4x4x!tt.ptr> loc(#loc14) + %p_30 = tt.addptr %p_12, %cst_0 : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc31) + scf.yield %p_30 : tensor<4x4x!tt.ptr> loc(#loc2) + } loc(#loc23) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":190:5) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":189:17) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":188:9) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":189:9) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":189:34) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":191:13) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":191:29) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":192:13) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:23) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:13) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:50) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":193:73) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":194:9) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":195:9) +#loc19 = loc("p"(#loc3)) +#loc20 = loc("r"(#loc4)) +#loc21 = loc("p"(#loc5)) +#loc22 = loc("p"(#loc6)) +#loc23 = loc("p"(#loc2)) +#loc24 = loc("q"(#loc7)) +#loc25 = loc("q"(#loc8)) +#loc26 = loc("v"(#loc9)) +#loc27 = loc("o"(#loc10)) +#loc28 = loc("o"(#loc11)) +#loc29 = loc("o"(#loc12)) +#loc30 = loc("o"(#loc13)) +#loc31 = loc("p"(#loc15)) diff --git a/tests/golden/ir/reader_ttir_3.8/expand_iterarg_mask.ttir b/tests/golden/ir/reader_ttir_3.8/expand_iterarg_mask.ttir new file mode 100644 index 000000000..fb7a0358e --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/expand_iterarg_mask.ttir @@ -0,0 +1,54 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":199:1) +#loc12 = loc("x_ptr"(#loc)) +#loc13 = loc("out_ptr"(#loc)) +#loc14 = loc("n"(#loc)) +#loc15 = loc("M"(#loc)) +module { + tt.func public @expand_iterarg_mask(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc)), %M: i32 loc("M"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant dense<4> : tensor<4xi32> loc(#loc2) + %cst_0 = arith.constant dense<4> : tensor<4x1xi32> loc(#loc2) + %r = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc16) + %p = tt.splat %x_ptr : !tt.ptr -> tensor<4x!tt.ptr> loc(#loc17) + %p_1 = tt.addptr %p, %r : tensor<4x!tt.ptr>, tensor<4xi32> loc(#loc17) + %p_2 = scf.for %k = %c0_i32 to %n step %c1_i32 iter_args(%p_3 = %p_1) -> (tensor<4x!tt.ptr>) : i32 { + %q = tt.expand_dims %p_3 {axis = 0 : i32} : tensor<4x!tt.ptr> -> tensor<1x4x!tt.ptr> loc(#loc19) + %q_4 = tt.broadcast %q : tensor<1x4x!tt.ptr> -> tensor<4x4x!tt.ptr> loc(#loc19) + %v = tt.expand_dims %r {axis = 0 : i32} : tensor<4xi32> -> tensor<1x4xi32> loc(#loc20) + %v_5 = tt.splat %M : i32 -> tensor<1x4xi32> loc(#loc20) + %v_6 = arith.cmpi slt, %v, %v_5 : tensor<1x4xi32> loc(#loc20) + %v_7 = tt.broadcast %v_6 : tensor<1x4xi1> -> tensor<4x4xi1> loc(#loc21) + %v_8 = tt.load %q_4, %v_7 : tensor<4x4x!tt.ptr> loc(#loc21) + %0 = tt.expand_dims %r {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc8) + %1 = arith.muli %0, %cst_0 : tensor<4x1xi32> loc(#loc8) + %2 = tt.splat %out_ptr : !tt.ptr -> tensor<4x1x!tt.ptr> loc(#loc9) + %3 = tt.addptr %2, %1 : tensor<4x1x!tt.ptr>, tensor<4x1xi32> loc(#loc9) + %4 = tt.broadcast %3 : tensor<4x1x!tt.ptr> -> tensor<4x4x!tt.ptr> loc(#loc9) + %5 = tt.broadcast %v : tensor<1x4xi32> -> tensor<4x4xi32> loc(#loc9) + %6 = tt.addptr %4, %5 : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc9) + tt.store %6, %v_8 : tensor<4x4x!tt.ptr> loc(#loc10) + %p_9 = tt.addptr %p_3, %cst : tensor<4x!tt.ptr>, tensor<4xi32> loc(#loc22) + scf.yield %p_9 : tensor<4x!tt.ptr> loc(#loc1) + } loc(#loc18) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":204:5) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":202:9) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":203:9) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":205:13) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":206:29) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":206:13) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":207:28) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":207:18) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":207:9) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":208:9) +#loc16 = loc("r"(#loc3)) +#loc17 = loc("p"(#loc4)) +#loc18 = loc("p"(#loc1)) +#loc19 = loc("q"(#loc5)) +#loc20 = loc("v"(#loc6)) +#loc21 = loc("v"(#loc7)) +#loc22 = loc("p"(#loc11)) diff --git a/tests/golden/ir/reader_ttir_3.8/int_iterarg_offset.ttir b/tests/golden/ir/reader_ttir_3.8/int_iterarg_offset.ttir new file mode 100644 index 000000000..cb91e78eb --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/int_iterarg_offset.ttir @@ -0,0 +1,29 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":212:1) +#loc7 = loc("x_ptr"(#loc)) +#loc8 = loc("n"(#loc)) +module { + tt.func public @int_iterarg_offset(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant dense<8> : tensor<8xi32> loc(#loc2) + %cst_0 = arith.constant dense<1.000000e+00> : tensor<8xf32> loc(#loc2) + %offs = tt.make_range {end = 8 : i32, start = 0 : i32} : tensor<8xi32> loc(#loc9) + %offs_1 = scf.for %k = %c0_i32 to %n step %c1_i32 iter_args(%offs_2 = %offs) -> (tensor<8xi32>) : i32 { + %0 = tt.splat %x_ptr : !tt.ptr -> tensor<8x!tt.ptr> loc(#loc4) + %1 = tt.addptr %0, %offs_2 : tensor<8x!tt.ptr>, tensor<8xi32> loc(#loc4) + tt.store %1, %cst_0 : tensor<8x!tt.ptr> loc(#loc5) + %offs_3 = arith.addi %offs_2, %cst : tensor<8xi32> loc(#loc11) + scf.yield %offs_3 : tensor<8xi32> loc(#loc1) + } loc(#loc10) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":215:5) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":214:12) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":216:18) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":216:9) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":217:9) +#loc9 = loc("offs"(#loc3)) +#loc10 = loc("offs"(#loc1)) +#loc11 = loc("offs"(#loc6)) diff --git a/tests/golden/ir/reader_ttir_3.8/iv_wrap.ttir b/tests/golden/ir/reader_ttir_3.8/iv_wrap.ttir new file mode 100644 index 000000000..e32022258 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/iv_wrap.ttir @@ -0,0 +1,19 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":221:1) +#loc5 = loc("x_ptr"(#loc)) +#loc6 = loc("lo"(#loc)) +#loc7 = loc("n"(#loc)) +module { + tt.func public @iv_wrap(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %lo: i32 loc("lo"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %c1048576_i32 = arith.constant 1048576 : i32 loc(#loc2) + scf.for %k = %lo to %n step %c1048576_i32 : i32 { + %0 = tt.addptr %x_ptr, %k : !tt.ptr, i32 loc(#loc3) + tt.store %0, %cst : !tt.ptr loc(#loc4) + } loc(#loc2) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":223:5) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":224:18) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":224:9) diff --git a/tests/golden/ir/reader_ttir_3.8/loop_observed_advance.ttir b/tests/golden/ir/reader_ttir_3.8/loop_observed_advance.ttir new file mode 100644 index 000000000..9ba86e667 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/loop_observed_advance.ttir @@ -0,0 +1,27 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":159:1) +#loc6 = loc("cnt_ptr"(#loc)) +#loc7 = loc("x_ptr"(#loc)) +#loc8 = loc("n"(#loc)) +module { + tt.func public @loop_observed_advance(%cnt_ptr: !tt.ptr loc("cnt_ptr"(#loc)), %x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %true = arith.constant true loc(#loc1) + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + %p = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%p_0 = %x_ptr) -> (!tt.ptr) : i32 { + tt.store %p_0, %cst : !tt.ptr loc(#loc3) + %p_1 = tt.atomic_rmw add, acq_rel, gpu, %cnt_ptr, %c1_i32, %true : (!tt.ptr, i32, i1) -> i32 loc(#loc10) + %p_2 = tt.addptr %p_0, %p_1 : !tt.ptr, i32 loc(#loc11) + scf.yield %p_2 : !tt.ptr loc(#loc2) + } loc(#loc9) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":162:5) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":163:9) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":164:14) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":164:9) +#loc9 = loc("p"(#loc2)) +#loc10 = loc("p"(#loc4)) +#loc11 = loc("p"(#loc5)) diff --git a/tests/golden/ir/reader_ttir_3.8/loop_two_step_advance.ttir b/tests/golden/ir/reader_ttir_3.8/loop_two_step_advance.ttir new file mode 100644 index 000000000..383543cd6 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/loop_two_step_advance.ttir @@ -0,0 +1,27 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":149:1) +#loc6 = loc("x_ptr"(#loc)) +#loc7 = loc("n"(#loc)) +#loc8 = loc("s"(#loc)) +module { + tt.func public @loop_two_step_advance(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc)), %s: i32 loc("s"(#loc))) attributes {noinline = false} { + %c2_i32 = arith.constant 2 : i32 loc(#loc1) + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %p = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%p_0 = %x_ptr) -> (!tt.ptr) : i32 { + tt.store %p_0, %cst : !tt.ptr loc(#loc3) + %p_1 = tt.addptr %p_0, %s : !tt.ptr, i32 loc(#loc10) + %p_2 = tt.addptr %p_1, %c2_i32 : !tt.ptr, i32 loc(#loc11) + scf.yield %p_2 : !tt.ptr loc(#loc2) + } loc(#loc9) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":152:5) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":153:9) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":154:9) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":155:9) +#loc9 = loc("p"(#loc2)) +#loc10 = loc("p"(#loc4)) +#loc11 = loc("p"(#loc5)) diff --git a/tests/golden/ir/reader_ttir_3.8/observed_lanes.ttir b/tests/golden/ir/reader_ttir_3.8/observed_lanes.ttir new file mode 100644 index 000000000..07442967f --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/observed_lanes.ttir @@ -0,0 +1,43 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":243:1) +#loc10 = loc("cnt_ptr"(#loc)) +#loc11 = loc("x_ptr"(#loc)) +module { + tt.func public @observed_lanes(%cnt_ptr: !tt.ptr loc("cnt_ptr"(#loc)), %x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<0> : tensor<1x4xi32> loc(#loc1) + %cst_0 = arith.constant dense<1.000000e+00> : tensor<4x4xf32> loc(#loc2) + %c = arith.constant dense<10> : tensor<4xi32> loc(#loc12) + %c_1 = arith.constant dense<0> : tensor<4xi32> loc(#loc13) + %old = arith.constant dense : tensor<4xi1> loc(#loc14) + %old_2 = arith.constant dense<1> : tensor<4xi32> loc(#loc14) + %r = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc15) + %old_3 = tt.splat %cnt_ptr : !tt.ptr -> tensor<4x!tt.ptr> loc(#loc16) + %old_4 = tt.addptr %old_3, %r : tensor<4x!tt.ptr>, tensor<4xi32> loc(#loc16) + %old_5 = tt.atomic_rmw add, acq_rel, gpu, %old_4, %old_2, %old : (tensor<4x!tt.ptr>, tensor<4xi32>, tensor<4xi1>) -> tensor<4xi32> loc(#loc14) + %c_6 = arith.maxsi %old_5, %c_1 : tensor<4xi32> loc(#loc13) + %c_7 = arith.minsi %c_6, %c : tensor<4xi32> loc(#loc12) + %0 = tt.expand_dims %c_7 {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc8) + %1 = tt.splat %x_ptr : !tt.ptr -> tensor<4x1x!tt.ptr> loc(#loc1) + %2 = tt.addptr %1, %0 : tensor<4x1x!tt.ptr>, tensor<4x1xi32> loc(#loc1) + %3 = tt.expand_dims %c_7 {axis = 0 : i32} : tensor<4xi32> -> tensor<1x4xi32> loc(#loc9) + %4 = tt.broadcast %2 : tensor<4x1x!tt.ptr> -> tensor<4x4x!tt.ptr> loc(#loc1) + %5 = arith.subi %cst, %3 : tensor<1x4xi32> loc(#loc1) + %6 = tt.broadcast %5 : tensor<1x4xi32> -> tensor<4x4xi32> loc(#loc1) + %7 = tt.addptr %4, %6 : tensor<4x4x!tt.ptr>, tensor<4x4xi32> loc(#loc1) + tt.store %7, %cst_0 : tensor<4x4x!tt.ptr> loc(#loc2) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":248:14) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":248:5) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":247:9) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":247:20) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":246:11) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":245:9) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":246:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":248:22) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":248:35) +#loc12 = loc("c"(#loc3)) +#loc13 = loc("c"(#loc4)) +#loc14 = loc("old"(#loc5)) +#loc15 = loc("r"(#loc6)) +#loc16 = loc("old"(#loc7)) diff --git a/tests/golden/ir/reader_ttir_3.8/p1_variant_delta.ttir b/tests/golden/ir/reader_ttir_3.8/p1_variant_delta.ttir new file mode 100644 index 000000000..8d07f9bce --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/p1_variant_delta.ttir @@ -0,0 +1,28 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":21:1) +#loc6 = loc("x_ptr"(#loc)) +#loc7 = loc("out_ptr"(#loc)) +#loc8 = loc("n"(#loc)) +module { + tt.func public @p1_variant_delta(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c-8_i32 = arith.constant -8 : i32 loc(#loc1) + %p = arith.constant 20 : i32 loc(#loc9) + %p_0 = tt.addptr %x_ptr, %p : !tt.ptr, i32 loc(#loc9) + %p_1 = scf.for %k = %c-8_i32 to %n step %c1_i32 iter_args(%p_2 = %p_0) -> (!tt.ptr) : i32 { + %v = tt.load %p_2 : !tt.ptr loc(#loc11) + tt.store %out_ptr, %v : !tt.ptr loc(#loc4) + %p_3 = tt.addptr %p_2, %k : !tt.ptr, i32 loc(#loc12) + scf.yield %p_3 : !tt.ptr loc(#loc1) + } loc(#loc10) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":23:5) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":22:9) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":24:13) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":25:9) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":26:9) +#loc9 = loc("p"(#loc2)) +#loc10 = loc("p"(#loc1)) +#loc11 = loc("v"(#loc3)) +#loc12 = loc("p"(#loc5)) diff --git a/tests/golden/ir/reader_ttir_3.8/p2_swap.ttir b/tests/golden/ir/reader_ttir_3.8/p2_swap.ttir new file mode 100644 index 000000000..487f48c1e --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/p2_swap.ttir @@ -0,0 +1,25 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":31:1) +#loc5 = loc("a_ptr"(#loc)) +#loc6 = loc("b_ptr"(#loc)) +#loc7 = loc("n"(#loc)) +module { + tt.func public @p2_swap(%a_ptr: !tt.ptr loc("a_ptr"(#loc)), %b_ptr: !tt.ptr loc("b_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant 1.000000e+00 : f32 loc(#loc2) + %q = arith.constant 100 : i32 loc(#loc8) + %q_0 = tt.addptr %b_ptr, %q : !tt.ptr, i32 loc(#loc8) + %q_1:2 = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%p = %a_ptr, %q_2 = %q_0) -> (!tt.ptr, !tt.ptr) : i32 { + tt.store %p, %cst : !tt.ptr loc(#loc4) + scf.yield %q_2, %p : !tt.ptr, !tt.ptr loc(#loc1) + } loc(#loc10) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":34:5) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":33:9) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":35:9) +#loc8 = loc("q"(#loc3)) +#loc9 = loc("p"(#loc1)) +#loc10 = loc("q"(#loc9)) diff --git a/tests/golden/ir/reader_ttir_3.8/p3_call_formals.ttir b/tests/golden/ir/reader_ttir_3.8/p3_call_formals.ttir new file mode 100644 index 000000000..d6b704b4f --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/p3_call_formals.ttir @@ -0,0 +1,28 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":66:1) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":61:1) +#loc8 = loc("x_ptr"(#loc)) +#loc9 = loc("n"(#loc)) +#loc11 = loc("dst"(#loc5)) +#loc12 = loc("i"(#loc5)) +module { + tt.func public @p3_call_formals(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c100_i32 = arith.constant 100 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc10) + %0 = arith.addi %pid, %c100_i32 : i32 loc(#loc3) + tt.call @reader_kernels._p3_helper_c__Pfp32_i32(%x_ptr, %0) : (!tt.ptr, i32) -> () loc(#loc4) + tt.return loc(#loc) + } loc(#loc) + tt.func private @reader_kernels._p3_helper_c__Pfp32_i32(%dst: !tt.ptr loc("dst"(#loc5)), %i: i32 loc("i"(#loc5))) attributes {noinline = true} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc6) + %0 = tt.addptr %dst, %i : !tt.ptr, i32 loc(#loc7) + tt.store %0, %cst : !tt.ptr loc(#loc6) + tt.return loc(#loc5) + } loc(#loc5) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":68:11) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":69:25) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":69:5) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":62:5) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":62:14) +#loc10 = loc("pid"(#loc2)) diff --git a/tests/golden/ir/reader_ttir_3.8/p3_call_guarded.ttir b/tests/golden/ir/reader_ttir_3.8/p3_call_guarded.ttir new file mode 100644 index 000000000..0b8a24bd9 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/p3_call_guarded.ttir @@ -0,0 +1,29 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":46:1) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":41:1) +#loc8 = loc("x_ptr"(#loc)) +#loc9 = loc("n"(#loc)) +#loc11 = loc("x_ptr"(#loc5)) +#loc12 = loc("pid"(#loc5)) +module { + tt.func public @p3_call_guarded(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %pid = tt.get_program_id x : i32 loc(#loc10) + %0 = arith.cmpi slt, %pid, %n : i32 loc(#loc2) + scf.if %0 { + tt.call @reader_kernels._p3_helper__Pfp32_i32(%x_ptr, %pid) : (!tt.ptr, i32) -> () loc(#loc4) + } loc(#loc3) + tt.return loc(#loc) + } loc(#loc) + tt.func private @reader_kernels._p3_helper__Pfp32_i32(%x_ptr: !tt.ptr loc("x_ptr"(#loc5)), %pid: i32 loc("pid"(#loc5))) attributes {noinline = true} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc6) + %0 = tt.addptr %x_ptr, %pid : !tt.ptr, i32 loc(#loc7) + tt.store %0, %cst : !tt.ptr loc(#loc6) + tt.return loc(#loc5) + } loc(#loc5) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":48:11) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":49:8) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":49:5) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":50:9) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":42:5) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":42:14) +#loc10 = loc("pid"(#loc1)) diff --git a/tests/golden/ir/reader_ttir_3.8/p3_call_offset.ttir b/tests/golden/ir/reader_ttir_3.8/p3_call_offset.ttir new file mode 100644 index 000000000..83f9e2fbd --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/p3_call_offset.ttir @@ -0,0 +1,28 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":54:1) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":41:1) +#loc8 = loc("x_ptr"(#loc)) +#loc9 = loc("n"(#loc)) +#loc11 = loc("x_ptr"(#loc5)) +#loc12 = loc("pid"(#loc5)) +module { + tt.func public @p3_call_offset(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c100_i32 = arith.constant 100 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc10) + %0 = arith.addi %pid, %c100_i32 : i32 loc(#loc3) + tt.call @reader_kernels._p3_helper__Pfp32_i32(%x_ptr, %0) : (!tt.ptr, i32) -> () loc(#loc4) + tt.return loc(#loc) + } loc(#loc) + tt.func private @reader_kernels._p3_helper__Pfp32_i32(%x_ptr: !tt.ptr loc("x_ptr"(#loc5)), %pid: i32 loc("pid"(#loc5))) attributes {noinline = true} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc6) + %0 = tt.addptr %x_ptr, %pid : !tt.ptr, i32 loc(#loc7) + tt.store %0, %cst : !tt.ptr loc(#loc6) + tt.return loc(#loc5) + } loc(#loc5) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":56:11) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":57:23) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":57:5) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":42:5) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":42:14) +#loc10 = loc("pid"(#loc2)) diff --git a/tests/golden/ir/reader_ttir_3.8/p4_observed_delta.ttir b/tests/golden/ir/reader_ttir_3.8/p4_observed_delta.ttir new file mode 100644 index 000000000..d31a4e680 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/p4_observed_delta.ttir @@ -0,0 +1,36 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":93:1) +#loc9 = loc("cnt_ptr"(#loc)) +#loc10 = loc("x_ptr"(#loc)) +#loc11 = loc("n"(#loc)) +module { + tt.func public @p4_observed_delta(%cnt_ptr: !tt.ptr loc("cnt_ptr"(#loc)), %x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c4_i32 = arith.constant 4 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant 1.000000e+00 : f32 loc(#loc2) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %old = arith.constant true loc(#loc12) + %old_0 = tt.atomic_rmw add, acq_rel, gpu, %cnt_ptr, %c1_i32, %old : (!tt.ptr, i32, i1) -> i32 loc(#loc12) + %step = arith.remsi %old_0, %n : i32 loc(#loc13) + %p = scf.for %i = %c0_i32 to %c4_i32 step %c1_i32 iter_args(%p_1 = %x_ptr) -> (!tt.ptr) : i32 { + %v = tt.load %p_1 : !tt.ptr loc(#loc15) + %0 = arith.addf %v, %cst : f32 loc(#loc6) + tt.store %p_1, %0 : !tt.ptr loc(#loc7) + %p_2 = tt.addptr %p_1, %step : !tt.ptr, i32 loc(#loc16) + scf.yield %p_2 : !tt.ptr loc(#loc1) + } loc(#loc14) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":98:5) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":95:11) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":96:12) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":99:13) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":100:21) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":100:9) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":101:9) +#loc12 = loc("old"(#loc3)) +#loc13 = loc("step"(#loc4)) +#loc14 = loc("p"(#loc1)) +#loc15 = loc("v"(#loc5)) +#loc16 = loc("p"(#loc8)) diff --git a/tests/golden/ir/reader_ttir_3.8/p4_observed_direct.ttir b/tests/golden/ir/reader_ttir_3.8/p4_observed_direct.ttir new file mode 100644 index 000000000..e05ce4094 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/p4_observed_direct.ttir @@ -0,0 +1,29 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":74:1) +#loc8 = loc("cnt_ptr"(#loc)) +#loc9 = loc("x_ptr"(#loc)) +#loc10 = loc("n"(#loc)) +module { + tt.func public @p4_observed_direct(%cnt_ptr: !tt.ptr loc("cnt_ptr"(#loc)), %x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %old = arith.constant 1 : i32 loc(#loc11) + %old_0 = arith.constant true loc(#loc11) + %old_1 = tt.atomic_rmw add, acq_rel, gpu, %cnt_ptr, %old, %old_0 : (!tt.ptr, i32, i1) -> i32 loc(#loc11) + %p = arith.remsi %old_1, %n : i32 loc(#loc12) + %p_2 = tt.addptr %x_ptr, %p : !tt.ptr, i32 loc(#loc13) + %v = tt.load %p_2 : !tt.ptr loc(#loc14) + %0 = arith.addf %v, %cst : f32 loc(#loc6) + tt.store %p_2, %0 : !tt.ptr loc(#loc7) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":75:11) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":76:17) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":76:9) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":77:9) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":78:17) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":78:5) +#loc11 = loc("old"(#loc2)) +#loc12 = loc("p"(#loc3)) +#loc13 = loc("p"(#loc4)) +#loc14 = loc("v"(#loc5)) diff --git a/tests/golden/ir/reader_ttir_3.8/p4_observed_loop.ttir b/tests/golden/ir/reader_ttir_3.8/p4_observed_loop.ttir new file mode 100644 index 000000000..85291defc --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/p4_observed_loop.ttir @@ -0,0 +1,39 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":82:1) +#loc10 = loc("cnt_ptr"(#loc)) +#loc11 = loc("x_ptr"(#loc)) +#loc12 = loc("n"(#loc)) +module { + tt.func public @p4_observed_loop(%cnt_ptr: !tt.ptr loc("cnt_ptr"(#loc)), %x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c4_i32 = arith.constant 4 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant 1.000000e+00 : f32 loc(#loc2) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %old = arith.constant true loc(#loc13) + %old_0 = tt.atomic_rmw add, acq_rel, gpu, %cnt_ptr, %c1_i32, %old : (!tt.ptr, i32, i1) -> i32 loc(#loc13) + %p = arith.remsi %old_0, %n : i32 loc(#loc14) + %p_1 = tt.addptr %x_ptr, %p : !tt.ptr, i32 loc(#loc15) + %p_2 = scf.for %i = %c0_i32 to %c4_i32 step %c1_i32 iter_args(%p_3 = %p_1) -> (!tt.ptr) : i32 { + %v = tt.load %p_3 : !tt.ptr loc(#loc17) + %0 = arith.addf %v, %cst : f32 loc(#loc7) + tt.store %p_3, %0 : !tt.ptr loc(#loc8) + %p_4 = tt.addptr %p_3, %c1_i32 : !tt.ptr, i32 loc(#loc18) + scf.yield %p_4 : !tt.ptr loc(#loc1) + } loc(#loc16) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":86:5) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":84:11) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":85:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":85:9) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":87:13) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":88:21) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":88:9) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":89:9) +#loc13 = loc("old"(#loc3)) +#loc14 = loc("p"(#loc4)) +#loc15 = loc("p"(#loc5)) +#loc16 = loc("p"(#loc1)) +#loc17 = loc("v"(#loc6)) +#loc18 = loc("p"(#loc9)) diff --git a/tests/golden/ir/reader_ttir_3.8/pure_asm_int_addr.ttir b/tests/golden/ir/reader_ttir_3.8/pure_asm_int_addr.ttir new file mode 100644 index 000000000..98d5c7be0 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/pure_asm_int_addr.ttir @@ -0,0 +1,21 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":228:1) +#loc6 = loc("x_ptr"(#loc)) +module { + tt.func public @pure_asm_int_addr(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %c4096_i32 = arith.constant 4096 : i32 loc(#loc1) + %r = arith.constant 7 : i32 loc(#loc7) + %a = tt.ptr_to_int %x_ptr : !tt.ptr -> i64 loc(#loc8) + %r_0 = tt.elementwise_inline_asm "st.global.b32 [$1], $2; mov.b32 $0, 0;" {constraints = "=r,l,r", packed_element = 1 : i32, pure = true} %a, %r : i64, i32 -> i32 loc(#loc9) + %0 = tt.addptr %x_ptr, %c4096_i32 : !tt.ptr, i32 loc(#loc1) + tt.store %0, %r_0 : !tt.ptr loc(#loc5) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":239:14) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":234:13) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":230:9) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":231:9) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":239:5) +#loc7 = loc("r"(#loc2)) +#loc8 = loc("a"(#loc3)) +#loc9 = loc("r"(#loc4)) diff --git a/tests/golden/ir/reader_ttir_3.8/rv_i32_wrap.ttir b/tests/golden/ir/reader_ttir_3.8/rv_i32_wrap.ttir new file mode 100644 index 000000000..b80d48f5c --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/rv_i32_wrap.ttir @@ -0,0 +1,22 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":114:1) +#loc6 = loc("x_ptr"(#loc)) +#loc7 = loc("S"(#loc)) +module { + tt.func public @rv_i32_wrap(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %S: i32 loc("S"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc8) + %off = arith.muli %pid, %S : i32 loc(#loc9) + %off_0 = arith.muli %off, %S : i32 loc(#loc10) + %0 = tt.addptr %x_ptr, %off_0 : !tt.ptr, i32 loc(#loc5) + tt.store %0, %c1_i32 : !tt.ptr loc(#loc1) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":118:5) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":116:11) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":117:12) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":117:11) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":118:14) +#loc8 = loc("pid"(#loc2)) +#loc9 = loc("off"(#loc3)) +#loc10 = loc("off"(#loc4)) diff --git a/tests/golden/ir/reader_ttir_3.8/rv_inline_asm_store.ttir b/tests/golden/ir/reader_ttir_3.8/rv_inline_asm_store.ttir new file mode 100644 index 000000000..9cf437497 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/rv_inline_asm_store.ttir @@ -0,0 +1,19 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":132:1) +#loc5 = loc("x_ptr"(#loc)) +module { + tt.func public @rv_inline_asm_store(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %v = arith.constant 7 : i32 loc(#loc6) + %p = arith.constant 4096 : i32 loc(#loc7) + %p_0 = tt.addptr %x_ptr, %p : !tt.ptr, i32 loc(#loc7) + %p_1 = tt.ptr_to_int %p_0 : !tt.ptr -> i64 loc(#loc8) + %0 = tt.elementwise_inline_asm "st.global.b32 [$1], $2; mov.b32 $0, 0;" {constraints = "=r,l,r", packed_element = 1 : i32, pure = false} %p_1, %v : i64, i32 -> i32 loc(#loc4) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":136:9) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":135:10) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":135:9) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":137:5) +#loc6 = loc("v"(#loc1)) +#loc7 = loc("p"(#loc2)) +#loc8 = loc("p"(#loc3)) diff --git a/tests/golden/ir/reader_ttir_3.8/rv_trunci_alias.ttir b/tests/golden/ir/reader_ttir_3.8/rv_trunci_alias.ttir new file mode 100644 index 000000000..486ac869a --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/rv_trunci_alias.ttir @@ -0,0 +1,24 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":106:1) +#loc7 = loc("x_ptr"(#loc)) +module { + tt.func public @rv_trunci_alias(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c4294967296_i64 = arith.constant 4294967296 : i64 loc(#loc2) + %pid = tt.get_program_id x : i32 loc(#loc8) + %pid_0 = arith.extsi %pid : i32 to i64 loc(#loc8) + %off = arith.muli %pid_0, %c4294967296_i64 : i64 loc(#loc9) + %off_1 = arith.trunci %off : i64 to i32 loc(#loc10) + %0 = tt.addptr %x_ptr, %off_1 : !tt.ptr, i32 loc(#loc6) + tt.store %0, %c1_i32 : !tt.ptr loc(#loc1) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":110:5) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":108:11) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":109:12) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":109:11) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":110:14) +#loc8 = loc("pid"(#loc3)) +#loc9 = loc("off"(#loc4)) +#loc10 = loc("off"(#loc5)) diff --git a/tests/golden/ir/reader_ttir_3.8/tile3d_shared_arange.ttir b/tests/golden/ir/reader_ttir_3.8/tile3d_shared_arange.ttir new file mode 100644 index 000000000..1ab272f6e --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/tile3d_shared_arange.ttir @@ -0,0 +1,37 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":176:1) +#loc7 = loc("x_ptr"(#loc)) +module { + tt.func public @tile3d_shared_arange(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<1.000000e+00> : tensor<4x4x4xf32> loc(#loc1) + %off = arith.constant dense<4> : tensor<1x4x1xi32> loc(#loc8) + %off_0 = arith.constant dense<16> : tensor<4x1x1xi32> loc(#loc9) + %r = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc10) + %off_1 = tt.expand_dims %r {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc9) + %off_2 = tt.expand_dims %off_1 {axis = 2 : i32} : tensor<4x1xi32> -> tensor<4x1x1xi32> loc(#loc9) + %off_3 = arith.muli %off_2, %off_0 : tensor<4x1x1xi32> loc(#loc9) + %off_4 = tt.expand_dims %r {axis = 0 : i32} : tensor<4xi32> -> tensor<1x4xi32> loc(#loc8) + %off_5 = tt.expand_dims %off_4 {axis = 2 : i32} : tensor<1x4xi32> -> tensor<1x4x1xi32> loc(#loc8) + %off_6 = arith.muli %off_5, %off : tensor<1x4x1xi32> loc(#loc8) + %off_7 = tt.broadcast %off_3 : tensor<4x1x1xi32> -> tensor<4x4x1xi32> loc(#loc9) + %off_8 = tt.broadcast %off_6 : tensor<1x4x1xi32> -> tensor<4x4x1xi32> loc(#loc9) + %off_9 = arith.addi %off_7, %off_8 : tensor<4x4x1xi32> loc(#loc9) + %off_10 = tt.expand_dims %off_4 {axis = 1 : i32} : tensor<1x4xi32> -> tensor<1x1x4xi32> loc(#loc11) + %off_11 = tt.broadcast %off_9 : tensor<4x4x1xi32> -> tensor<4x4x4xi32> loc(#loc9) + %off_12 = tt.broadcast %off_10 : tensor<1x1x4xi32> -> tensor<4x4x4xi32> loc(#loc9) + %off_13 = arith.addi %off_11, %off_12 : tensor<4x4x4xi32> loc(#loc9) + %0 = tt.splat %x_ptr : !tt.ptr -> tensor<4x4x4x!tt.ptr> loc(#loc6) + %1 = tt.addptr %0, %off_13 : tensor<4x4x4x!tt.ptr>, tensor<4x4x4xi32> loc(#loc6) + tt.store %1, %cst : tensor<4x4x4x!tt.ptr> loc(#loc1) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":180:5) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":179:40) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":179:11) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":178:9) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":179:63) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":180:14) +#loc8 = loc("off"(#loc2)) +#loc9 = loc("off"(#loc3)) +#loc10 = loc("r"(#loc4)) +#loc11 = loc("off"(#loc5)) diff --git a/tests/golden/ir/reader_ttir_3.8/unsigned_index.ttir b/tests/golden/ir/reader_ttir_3.8/unsigned_index.ttir new file mode 100644 index 000000000..d44bbadb0 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/unsigned_index.ttir @@ -0,0 +1,25 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":122:1) +#loc7 = loc("x_ptr"(#loc)) +#loc8 = loc("n"(#loc)) +module { + tt.func public @unsigned_index(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %c3_i32 = arith.constant 3 : i32 loc(#loc2) + %pid = tt.get_program_id x : i32 loc(#loc9) + %q = arith.divui %pid, %c3_i32 : i32 loc(#loc10) + %m = arith.cmpi ult, %pid, %n : i32 loc(#loc11) + %0 = arith.extui %q : i32 to i64 loc(#loc6) + %1 = tt.addptr %x_ptr, %0 : !tt.ptr, i64 loc(#loc6) + tt.store %1, %cst, %m : !tt.ptr loc(#loc1) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":127:5) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":124:11) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":125:9) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":126:9) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":127:14) +#loc9 = loc("pid"(#loc3)) +#loc10 = loc("q"(#loc4)) +#loc11 = loc("m"(#loc5)) diff --git a/tests/golden/ir/reader_ttir_3.8/where_pointer.ttir b/tests/golden/ir/reader_ttir_3.8/where_pointer.ttir new file mode 100644 index 000000000..3747bb235 --- /dev/null +++ b/tests/golden/ir/reader_ttir_3.8/where_pointer.ttir @@ -0,0 +1,30 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":168:1) +#loc7 = loc("x_ptr"(#loc)) +#loc8 = loc("n"(#loc)) +module { + tt.func public @where_pointer(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<1.000000e+00> : tensor<16xf32> loc(#loc1) + %p = arith.constant 100 : i32 loc(#loc9) + %offs = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32> loc(#loc10) + %p_0 = tt.splat %n : i32 -> tensor<16xi32> loc(#loc11) + %p_1 = arith.cmpi slt, %offs, %p_0 : tensor<16xi32> loc(#loc11) + %p_2 = tt.splat %x_ptr : !tt.ptr -> tensor<16x!tt.ptr> loc(#loc12) + %p_3 = tt.addptr %p_2, %offs : tensor<16x!tt.ptr>, tensor<16xi32> loc(#loc12) + %p_4 = tt.addptr %x_ptr, %p : !tt.ptr, i32 loc(#loc9) + %p_5 = tt.splat %p_4 : !tt.ptr -> tensor<16x!tt.ptr> loc(#loc13) + %p_6 = arith.select %p_1, %p_3, %p_5 : tensor<16xi1>, tensor<16x!tt.ptr> loc(#loc13) + tt.store %p_6, %cst : tensor<16x!tt.ptr> loc(#loc1) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":172:5) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":171:42) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":170:12) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":171:18) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":171:28) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/reader_kernels.py":171:9) +#loc9 = loc("p"(#loc2)) +#loc10 = loc("offs"(#loc3)) +#loc11 = loc("p"(#loc4)) +#loc12 = loc("p"(#loc5)) +#loc13 = loc("p"(#loc6)) diff --git a/tests/golden/ir/ttir/adv_cf_blockargs.ttir b/tests/golden/ir/ttir/adv_cf_blockargs.ttir new file mode 100644 index 000000000..b9a029314 --- /dev/null +++ b/tests/golden/ir/ttir/adv_cf_blockargs.ttir @@ -0,0 +1,86 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":70:0) +#loc1 = loc(unknown) +#loc25 = loc("x_ptr"(#loc)) +#loc26 = loc("out_ptr"(#loc)) +#loc27 = loc("n"(#loc)) +#loc28 = loc("t"(#loc)) +module { + tt.func public @cf_blockargs(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc)), %t: i32 loc("t"(#loc))) attributes {noinline = false} { + %c100_i32 = arith.constant 100 : i32 loc(#loc1) + %c7_i32 = arith.constant 7 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc29) + %s = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%s_4 = %c0_i32) -> (i32) : i32 { + %s_5 = arith.muli %i, %pid : i32 loc(#loc31) + %s_6 = arith.addi %s_4, %s_5 : i32 loc(#loc32) + scf.yield %s_6 : i32 loc(#loc6) + } loc(#loc30) + %0 = arith.cmpi sgt, %s, %t : i32 loc(#loc7) + cf.cond_br %0, ^bb1, ^bb4 loc(#loc7) + ^bb1: // pred: ^bb0 + %1 = arith.cmpi eq, %pid, %c1_i32 : i32 loc(#loc8) + cf.cond_br %1, ^bb2, ^bb3 loc(#loc8) + ^bb2: // 2 preds: ^bb1, ^bb5 + tt.return loc(#loc9) + ^bb3: // pred: ^bb1 + %2 = tt.addptr %out_ptr, %pid : !tt.ptr, i32 loc(#loc10) + %3 = arith.sitofp %s : i32 to f32 loc(#loc11) + tt.store %2, %3 : !tt.ptr loc(#loc11) + tt.return loc(#loc12) + ^bb4: // pred: ^bb0 + %offs = arith.muli %pid, %c64_i32 : i32 loc(#loc33) + %offs_0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc34) + %offs_1 = tt.splat %offs : i32 -> tensor<64xi32> loc(#loc35) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<64xi32> loc(#loc35) + %4 = arith.cmpi sgt, %pid, %c7_i32 : i32 loc(#loc16) + cf.cond_br %4, ^bb5, ^bb6(%s : i32) loc(#loc16) + ^bb5: // pred: ^bb4 + %s_3 = arith.addi %s, %c1_i32 : i32 loc(#loc36) + %5 = arith.cmpi sgt, %s_3, %c100_i32 : i32 loc(#loc18) + cf.cond_br %5, ^bb2, ^bb6(%s_3 : i32) loc(#loc18) + ^bb6(%6: i32 loc(unknown)): // 2 preds: ^bb4, ^bb5 + %7 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc19) + %8 = tt.addptr %7, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc19) + %9 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc20) + %10 = tt.addptr %9, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc20) + %11 = tt.load %10 : tensor<64x!tt.ptr> loc(#loc21) + %12 = arith.sitofp %6 : i32 to f32 loc(#loc22) + %13 = tt.splat %12 : f32 -> tensor<64xf32> loc(#loc22) + %14 = arith.addf %11, %13 : tensor<64xf32> loc(#loc22) + tt.store %8, %14 : tensor<64x!tt.ptr> loc(#loc23) + tt.return loc(#loc24) + } loc(#loc) +} loc(#loc) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":71:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":73:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":74:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":74:13) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":74:8) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":75:11) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":76:18) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":77:12) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":78:27) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":78:32) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":79:8) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":80:17) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":80:38) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":80:25) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":81:13) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":82:16) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":83:15) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":85:23) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":85:45) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":85:37) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":85:53) +#loc23 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":85:29) +#loc24 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":85:4) +#loc29 = loc("pid"(#loc2)) +#loc30 = loc("s"(#loc3)) +#loc31 = loc("s"(#loc4)) +#loc32 = loc("s"(#loc5)) +#loc33 = loc("offs"(#loc13)) +#loc34 = loc("offs"(#loc14)) +#loc35 = loc("offs"(#loc15)) +#loc36 = loc("s"(#loc17)) diff --git a/tests/golden/ir/ttir/adv_consts.ttir b/tests/golden/ir/ttir/adv_consts.ttir new file mode 100644 index 000000000..2618addb5 --- /dev/null +++ b/tests/golden/ir/ttir/adv_consts.ttir @@ -0,0 +1,98 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":157:0) +#loc37 = loc("x_ptr"(#loc)) +#loc38 = loc("i8_ptr"(#loc)) +#loc39 = loc("i64_ptr"(#loc)) +#loc40 = loc("u32_ptr"(#loc)) +module { + tt.func public @consts(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %i8_ptr: !tt.ptr loc("i8_ptr"(#loc)), %i64_ptr: !tt.ptr loc("i64_ptr"(#loc)), %u32_ptr: !tt.ptr loc("u32_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<-3.000000e+00> : tensor<64xf32> loc(#loc1) + %c-9223372036854775807_i64 = arith.constant -9223372036854775807 : i64 loc(#loc2) + %cst_0 = arith.constant dense<4294967295> : tensor<64xi64> loc(#loc3) + %cst_1 = arith.constant dense<-1> : tensor<64xi8> loc(#loc4) + %cst_2 = arith.constant dense<1.000000e-30> : tensor<64xf32> loc(#loc5) + %cst_3 = arith.constant dense<-2.14748365E+9> : tensor<64xf32> loc(#loc6) + %cst_4 = arith.constant dense<-7.000000e+00> : tensor<64xf32> loc(#loc7) + %cst_5 = arith.constant dense<192> : tensor<64xi32> loc(#loc8) + %cst_6 = arith.constant dense<1> : tensor<64xi32> loc(#loc9) + %cst_7 = arith.constant dense<-1> : tensor<64xi32> loc(#loc9) + %big = arith.constant dense<9223372036854775807> : tensor<64xi64> loc(#loc41) + %cst_8 = arith.constant dense<128> : tensor<64xi32> loc(#loc9) + %cst_9 = arith.constant dense<0xFF800000> : tensor<64xf32> loc(#loc11) + %cst_10 = arith.constant dense<0x7FC00000> : tensor<64xf32> loc(#loc11) + %cst_11 = arith.constant dense<64> : tensor<64xi32> loc(#loc9) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc42) + %v = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc43) + %v_12 = tt.addptr %v, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc43) + %v_13 = tt.load %v_12 : tensor<64x!tt.ptr> loc(#loc44) + %0 = arith.addf %v_13, %cst_4 : tensor<64xf32> loc(#loc7) + %1 = arith.addf %0, %cst_3 : tensor<64xf32> loc(#loc15) + tt.store %v_12, %1 : tensor<64x!tt.ptr> loc(#loc16) + %2 = tt.addptr %v_12, %cst_11 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc17) + %3 = arith.cmpf une, %v_13, %v_13 : tensor<64xf32> loc(#loc18) + %4 = arith.select %3, %cst_10, %cst_9 : tensor<64xi1>, tensor<64xf32> loc(#loc11) + tt.store %2, %4 : tensor<64x!tt.ptr> loc(#loc19) + %5 = tt.addptr %v_12, %cst_8 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc20) + tt.store %5, %cst_2 : tensor<64x!tt.ptr> loc(#loc21) + %6 = tt.splat %i8_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc22) + %7 = tt.addptr %6, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc22) + tt.store %7, %cst_1 : tensor<64x!tt.ptr> loc(#loc23) + %8 = tt.splat %i64_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc24) + %9 = tt.addptr %8, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc24) + tt.store %9, %cst_0 : tensor<64x!tt.ptr> loc(#loc25) + %10 = tt.addptr %i64_ptr, %c-9223372036854775807_i64 : !tt.ptr, i64 loc(#loc2) + %11 = tt.splat %10 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc26) + %12 = tt.addptr %11, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc26) + tt.store %12, %big : tensor<64x!tt.ptr> loc(#loc27) + %13 = tt.splat %u32_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc28) + %14 = tt.addptr %13, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc28) + tt.store %14, %cst_7 : tensor<64x!tt.ptr> loc(#loc29) + %15 = tt.addptr %14, %cst_11 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc30) + tt.store %15, %cst_6 : tensor<64x!tt.ptr> loc(#loc31) + %16 = arith.cmpi slt, %offs, %cst_7 : tensor<64xi32> loc(#loc32) + %17 = tt.addptr %14, %cst_8 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc33) + tt.store %17, %cst_6, %16 : tensor<64x!tt.ptr> loc(#loc34) + %18 = tt.addptr %v_12, %cst_5 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc8) + tt.store %18, %cst : tensor<64x!tt.ptr> loc(#loc35) + tt.return loc(#loc36) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":175:44) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":169:23) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":168:43) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":165:62) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":164:76) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":162:43) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":162:31) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":175:28) +#loc9 = loc(unknown) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":166:48) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":163:66) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":158:24) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":159:24) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":159:16) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":162:37) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":162:27) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":163:28) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":163:49) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":163:35) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":164:28) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":164:39) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":165:22) +#loc23 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":165:28) +#loc24 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":168:23) +#loc25 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":168:29) +#loc26 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":169:45) +#loc27 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":169:51) +#loc28 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":170:23) +#loc29 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":170:29) +#loc30 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":172:30) +#loc31 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":172:37) +#loc32 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":173:85) +#loc33 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":173:30) +#loc34 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":173:41) +#loc35 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":175:39) +#loc36 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":175:4) +#loc41 = loc("big"(#loc10)) +#loc42 = loc("offs"(#loc12)) +#loc43 = loc("v"(#loc13)) +#loc44 = loc("v"(#loc14)) diff --git a/tests/golden/ir/ttir/adv_descs.ttir b/tests/golden/ir/ttir/adv_descs.ttir new file mode 100644 index 000000000..fe82dd485 --- /dev/null +++ b/tests/golden/ir/ttir/adv_descs.ttir @@ -0,0 +1,36 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":234:0) +#loc11 = loc("a_ptr"(#loc)) +#loc12 = loc("M"(#loc)) +#loc13 = loc("N"(#loc)) +module { + tt.func public @descs(%a_ptr: !tt.ptr loc("a_ptr"(#loc)), %M: i32 loc("M"(#loc)), %N: i32 loc("N"(#loc))) attributes {noinline = false} { + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c1_i64 = arith.constant 1 : i64 loc(#loc1) + %d = arith.extsi %N : i32 to i64 loc(#loc14) + %d_0 = tt.make_tensor_descriptor %a_ptr, [%M, %N], [%d, %c1_i64] : , > loc(#loc14) + %x = tt.descriptor_load %d_0[%c0_i32, %c32_i32] : !tt.tensordesc> -> tensor<32x32xf16> loc(#loc15) + tt.descriptor_store %d_0[%c32_i32, %c0_i32], %x : !tt.tensordesc>, tensor<32x32xf16> loc(#loc4) + tt.descriptor_reduce add, %d_0[%c32_i32, %c32_i32], %x : !tt.tensordesc>, tensor<32x32xf16> loc(#loc5) + %d1 = tt.make_tensor_descriptor %a_ptr, [%M, %N], [%d, %c1_i64] : , > loc(#loc16) + %rows = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc17) + %g = tt.descriptor_gather %d1[%rows, %c0_i32] : (!tt.tensordesc>, tensor<32xi32>, i32) -> tensor<32x32xf16> loc(#loc18) + tt.descriptor_scatter %d1[%rows, %c32_i32], %g : !tt.tensordesc>, tensor<32xi32>, i32, tensor<32x32xf16> loc(#loc9) + tt.return loc(#loc10) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":235:57) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":236:15) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":237:21) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":238:27) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":239:58) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":240:24) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":241:24) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":242:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":242:4) +#loc14 = loc("d"(#loc2)) +#loc15 = loc("x"(#loc3)) +#loc16 = loc("d1"(#loc6)) +#loc17 = loc("rows"(#loc7)) +#loc18 = loc("g"(#loc8)) diff --git a/tests/golden/ir/ttir/adv_hinted.ttir b/tests/golden/ir/ttir/adv_hinted.ttir new file mode 100644 index 000000000..749e3886e --- /dev/null +++ b/tests/golden/ir/ttir/adv_hinted.ttir @@ -0,0 +1,41 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":196:0) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":199:15) +#loc7 = loc(unknown) +#loc12 = loc("x_ptr"(#loc)) +#loc13 = loc("out_ptr"(#loc)) +#loc14 = loc("n"(#loc)) +#loc18 = loc("s"(#loc6)) +#loc20 = loc(callsite(#loc7 at #loc18)) +module { + tt.func public @hinted(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc15) + %x = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc16) + %x_0 = tt.addptr %x, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc16) + %x_1 = tt.load %x_0 : tensor<64x!tt.ptr> loc(#loc17) + %s = "tt.reduce"(%offs) <{axis = 0 : i32}> ({ + ^bb0(%s_2: i32 loc(callsite(#loc7 at #loc18)), %s_3: i32 loc(callsite(#loc7 at #loc18))): + %s_4 = arith.addi %s_2, %s_3 : i32 loc(#loc21) + tt.reduce.return %s_4 : i32 loc(#loc19) + }) : (tensor<64xi32>) -> i32 loc(#loc19) + %0 = tt.addptr %out_ptr, %s : !tt.ptr, i32 loc(#loc9) + %1 = tt.addptr %0, %c32_i32 : !tt.ptr, i32 loc(#loc1) + %2 = tt.splat %1 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc1) + tt.store %2, %x_1 : tensor<64x!tt.ptr> loc(#loc10) + tt.return loc(#loc11) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":202:27) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":197:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":198:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":198:16) +#loc5 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":293:36) +#loc8 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":263:15) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":202:23) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":202:30) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":202:4) +#loc15 = loc("offs"(#loc2)) +#loc16 = loc("x"(#loc3)) +#loc17 = loc("x"(#loc4)) +#loc19 = loc(callsite(#loc5 at #loc18)) +#loc21 = loc(callsite(#loc8 at #loc19)) diff --git a/tests/golden/ir/ttir/adv_multi_func.ttir b/tests/golden/ir/ttir/adv_multi_func.ttir new file mode 100644 index 000000000..fde9d9ef4 --- /dev/null +++ b/tests/golden/ir/ttir/adv_multi_func.ttir @@ -0,0 +1,83 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":107:0) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":89:0) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":96:0) +#loc26 = loc("ptr"(#loc)) +#loc27 = loc("n"(#loc)) +#loc29 = loc("ptr"(#loc8)) +#loc30 = loc("a"(#loc8)) +#loc31 = loc("ptr"(#loc17)) +#loc32 = loc("n"(#loc17)) +#loc35 = loc("b"(#loc8)) +module { + tt.func public @multi_func(%ptr: !tt.ptr loc("ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %0:2 = tt.call @"adv_kernels._nl_pair__Pi32_i32__(2,)cconstexpr_3_"(%ptr, %n) : (!tt.ptr, i32) -> (i32, i32) loc(#loc2) + %r = tt.call @adv_kernels._nl_loop__Pi32_i32__(%ptr, %0#0) : (!tt.ptr, i32) -> i32 loc(#loc28) + %1 = tt.addptr %ptr, %0#1 : !tt.ptr, i32 loc(#loc4) + tt.store %1, %r : !tt.ptr loc(#loc5) + %2 = tt.addptr %ptr, %c1_i32 : !tt.ptr, i32 loc(#loc1) + %3:2 = tt.call @adv_kernels._nl_pair__Pi32_i32_i32__(%2, %r, %0#1) : (!tt.ptr, i32, i32) -> (i32, i32) loc(#loc6) + tt.return loc(#loc7) + } loc(#loc) + tt.func private @"adv_kernels._nl_pair__Pi32_i32__(2,)cconstexpr_3_"(%ptr: !tt.ptr loc("ptr"(#loc8)), %a: i32 loc("a"(#loc8))) -> (i32, i32) attributes {noinline = true} { + %c3_i32 = arith.constant 3 : i32 loc(#loc9) + %0 = arith.cmpi slt, %a, %c3_i32 : i32 loc(#loc10) + scf.if %0 { + %3 = tt.addptr %ptr, %a : !tt.ptr, i32 loc(#loc12) + tt.store %3, %c3_i32 : !tt.ptr loc(#loc13) + } loc(#loc11) + %1 = arith.addi %a, %c3_i32 : i32 loc(#loc14) + %2 = arith.muli %a, %c3_i32 : i32 loc(#loc15) + tt.return %1, %2 : i32, i32 loc(#loc16) + } loc(#loc8) + tt.func private @adv_kernels._nl_loop__Pi32_i32__(%ptr: !tt.ptr loc("ptr"(#loc17)), %n: i32 loc("n"(#loc17))) -> i32 attributes {noinline = true} { + %true = arith.constant true loc(#loc9) + %c0_i32 = arith.constant 0 : i32 loc(#loc9) + %c1_i32 = arith.constant 1 : i32 loc(#loc9) + %s = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%s_0 = %c0_i32) -> (i32) : i32 { + %s_1 = arith.addi %s_0, %i : i32 loc(#loc34) + %2 = tt.addptr %ptr, %i : !tt.ptr, i32 loc(#loc20) + %3 = tt.atomic_rmw add, acq_rel, gpu, %2, %c1_i32, %true : (!tt.ptr, i32, i1) -> i32 loc(#loc21) + scf.yield %s_1 : i32 loc(#loc22) + } loc(#loc33) + %0:2 = tt.call @adv_kernels._nl_pair__Pi32_i32_i32__(%ptr, %s, %n) : (!tt.ptr, i32, i32) -> (i32, i32) loc(#loc23) + %1 = arith.subi %0#0, %0#1 : i32 loc(#loc24) + tt.return %1 : i32 loc(#loc25) + } loc(#loc17) + tt.func private @adv_kernels._nl_pair__Pi32_i32_i32__(%ptr: !tt.ptr loc("ptr"(#loc8)), %a: i32 loc("a"(#loc8)), %b: i32 loc("b"(#loc8))) -> (i32, i32) attributes {noinline = true} { + %0 = arith.cmpi slt, %a, %b : i32 loc(#loc10) + scf.if %0 { + %3 = tt.addptr %ptr, %a : !tt.ptr, i32 loc(#loc12) + tt.store %3, %b : !tt.ptr loc(#loc13) + } loc(#loc11) + %1 = arith.addi %a, %b : i32 loc(#loc14) + %2 = arith.muli %a, %b : i32 loc(#loc15) + tt.return %1, %2 : i32, i32 loc(#loc16) + } loc(#loc8) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":111:19) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":108:28) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":109:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":110:19) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":110:22) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":111:25) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":111:4) +#loc9 = loc(unknown) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":90:11) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":90:7) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":91:23) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":91:26) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":92:15) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":92:22) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":92:11) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":98:22) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":99:13) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":100:28) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":100:31) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":100:8) +#loc23 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":101:28) +#loc24 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":102:15) +#loc25 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":102:11) +#loc28 = loc("r"(#loc3)) +#loc33 = loc("s"(#loc18)) +#loc34 = loc("s"(#loc19)) diff --git a/tests/golden/ir/ttir/adv_multi_result.ttir b/tests/golden/ir/ttir/adv_multi_result.ttir new file mode 100644 index 000000000..d15ff7515 --- /dev/null +++ b/tests/golden/ir/ttir/adv_multi_result.ttir @@ -0,0 +1,105 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":132:0) +#loc29 = loc("x_ptr"(#loc)) +#loc30 = loc("out_ptr"(#loc)) +#loc31 = loc("n"(#loc)) +module { + tt.func public @multi_result(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c4_i32 = arith.constant 4 : i32 loc(#loc1) + %cst = arith.constant 2.000000e+00 : f32 loc(#loc2) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %c1 = arith.constant 1.000000e+00 : f32 loc(#loc32) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + %y = arith.constant 64 : i32 loc(#loc33) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc34) + %x = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc35) + %x_0 = tt.addptr %x, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc35) + %x_1 = tt.load %x_0 : tensor<64x!tt.ptr> loc(#loc36) + %y_2 = tt.addptr %x_ptr, %y : !tt.ptr, i32 loc(#loc33) + %y_3 = tt.splat %y_2 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc37) + %y_4 = tt.addptr %y_3, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc37) + %y_5 = tt.load %y_4 : tensor<64x!tt.ptr> loc(#loc38) + %j = tt.join %x_1, %y_5 : tensor<64xf32> -> tensor<64x2xf32> loc(#loc39) + %r1, %r1_6 = tt.split %j : tensor<64x2xf32> -> tensor<64xf32> loc(#loc51) + %r2:4 = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%c0 = %c0_i32, %c1_7 = %c1, %c2 = %offs, %c3 = %x_ptr) -> (i32, f32, tensor<64xi32>, !tt.ptr) : i32 { + %c0_8 = arith.addi %c0, %i : i32 loc(#loc42) + %c1_9 = arith.mulf %c1_7, %cst : f32 loc(#loc43) + %c2_10 = tt.splat %i : i32 -> tensor<64xi32> loc(#loc44) + %c2_11 = arith.addi %c2, %c2_10 : tensor<64xi32> loc(#loc44) + %c3_12 = tt.addptr %c3, %c1_i32 : !tt.ptr, i32 loc(#loc45) + scf.yield %c0_8, %c1_9, %c2_11, %c3_12 : i32, f32, tensor<64xi32>, !tt.ptr loc(#loc17) + } loc(#loc53) + %0 = arith.cmpi sgt, %n, %c4_i32 : i32 loc(#loc1) + %1 = arith.select %0, %r1, %r1_6 : tensor<64xf32> loc(#loc18) + %2 = arith.select %0, %r1_6, %r1 : tensor<64xf32> loc(#loc18) + %3 = scf.if %0 -> (i32) { + scf.yield %r2#0 : i32 loc(#loc18) + } else { + %r2_7 = arith.addi %r2#0, %c1_i32 : i32 loc(#loc46) + scf.yield %r2_7 : i32 loc(#loc19) + } loc(#loc18) + %4 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc20) + %5 = tt.addptr %4, %r2#2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc20) + %6 = arith.addf %1, %2 : tensor<64xf32> loc(#loc21) + %7 = arith.sitofp %3 : i32 to f32 loc(#loc22) + %8 = tt.splat %7 : f32 -> tensor<64xf32> loc(#loc22) + %9 = arith.addf %6, %8 : tensor<64xf32> loc(#loc22) + %10 = tt.splat %r2#1 : f32 -> tensor<64xf32> loc(#loc23) + %11 = arith.addf %9, %10 : tensor<64xf32> loc(#loc23) + %12 = tt.splat %r2#3 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc24) + %13 = tt.addptr %12, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc24) + %14 = tt.load %13 : tensor<64x!tt.ptr> loc(#loc25) + %15 = arith.addf %11, %14 : tensor<64xf32> loc(#loc26) + tt.store %5, %15 : tensor<64x!tt.ptr> loc(#loc27) + tt.return loc(#loc28) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":147:11) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":139:9) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":135:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":133:24) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":134:24) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":134:16) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":135:32) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":135:16) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":136:19) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":137:20) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":142:22) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":143:14) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":144:14) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":145:18) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":146:18) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":146:8) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":147:7) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":150:32) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":151:23) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":151:32) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":151:37) +#loc23 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":151:42) +#loc24 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":151:60) +#loc25 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":151:55) +#loc26 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":151:47) +#loc27 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":151:27) +#loc28 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":151:4) +#loc32 = loc("c1"(#loc3)) +#loc33 = loc("y"(#loc4)) +#loc34 = loc("offs"(#loc5)) +#loc35 = loc("x"(#loc6)) +#loc36 = loc("x"(#loc7)) +#loc37 = loc("y"(#loc8)) +#loc38 = loc("y"(#loc9)) +#loc39 = loc("j"(#loc10)) +#loc40 = loc("r0"(#loc11)) +#loc41 = loc("c0"(#loc12)) +#loc42 = loc("c0"(#loc13)) +#loc43 = loc("c1"(#loc14)) +#loc44 = loc("c2"(#loc15)) +#loc45 = loc("c3"(#loc16)) +#loc46 = loc("r2"(#loc19)) +#loc47 = loc("r1"(#loc40)) +#loc48 = loc("c1"(#loc41)) +#loc49 = loc("r0"(#loc47)) +#loc50 = loc("c2"(#loc48)) +#loc51 = loc("r1"(#loc49)) +#loc52 = loc("c3"(#loc50)) +#loc53 = loc("r2"(#loc52)) diff --git a/tests/golden/ir/ttir/adv_names.ttir b/tests/golden/ir/ttir/adv_names.ttir new file mode 100644 index 000000000..c5842cdc7 --- /dev/null +++ b/tests/golden/ir/ttir/adv_names.ttir @@ -0,0 +1,37 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":207:0) +#loc10 = loc("x_ptr"(#loc)) +#loc11 = loc("out_ptr"(#loc)) +#loc12 = loc("n"(#loc)) +module { + tt.func public @names(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %to = arith.constant dense<3> : tensor<64xi32> loc(#loc13) + %loc = arith.constant dense<2> : tensor<64xi32> loc(#loc14) + %_CF80 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc15) + %loc_0 = arith.muli %_CF80, %loc : tensor<64xi32> loc(#loc14) + %to_1 = arith.addi %_CF80, %to : tensor<64xi32> loc(#loc13) + %iter_args = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc16) + %iter_args_2 = tt.addptr %iter_args, %loc_0 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc16) + %iter_args_3 = tt.load %iter_args_2 : tensor<64x!tt.ptr> loc(#loc17) + %true = tt.splat %n : i32 -> tensor<64xi32> loc(#loc18) + %true_4 = arith.cmpi slt, %to_1, %true : tensor<64xi32> loc(#loc18) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc7) + %1 = tt.addptr %0, %to_1 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc7) + tt.store %1, %iter_args_3, %true_4 : tensor<64x!tt.ptr> loc(#loc8) + tt.return loc(#loc9) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":211:14) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":209:15) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":208:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":212:32) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":212:24) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":213:16) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":214:23) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":214:27) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":214:4) +#loc13 = loc("to"(#loc1)) +#loc14 = loc("loc"(#loc2)) +#loc15 = loc("\CF\80"(#loc3)) +#loc16 = loc("iter_args"(#loc4)) +#loc17 = loc("iter_args"(#loc5)) +#loc18 = loc("true"(#loc6)) diff --git a/tests/golden/ir/ttir/adv_nest3.ttir b/tests/golden/ir/ttir/adv_nest3.ttir new file mode 100644 index 000000000..e679a7874 --- /dev/null +++ b/tests/golden/ir/ttir/adv_nest3.ttir @@ -0,0 +1,196 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":15:0) +#loc32 = loc("x_ptr"(#loc)) +#loc33 = loc("y_ptr"(#loc)) +#loc34 = loc("out_ptr"(#loc)) +#loc35 = loc("n"(#loc)) +#loc36 = loc("m"(#loc)) +module { + tt.func public @nest3(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %y_ptr: !tt.ptr loc("y_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc)), %m: i32 loc("m"(#loc))) attributes {noinline = false} { + %acc = arith.constant dense<0.000000e+00> : tensor<64xf32> loc(#loc53) + %c2_i32 = arith.constant 2 : i32 loc(#loc3) + %c1_i32 = arith.constant 1 : i32 loc(#loc4) + %cst = arith.constant dense<2.000000e+00> : tensor<64xf32> loc(#loc3) + %c5_i32 = arith.constant 5 : i32 loc(#loc3) + %c64_i32 = arith.constant 64 : i32 loc(#loc3) + %c3_i32 = arith.constant 3 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc38) + %p = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc39) + %p_0 = tt.addptr %p, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc39) + %acc_1 = arith.subi %n, %c0_i32 : i32 loc(#loc40) + %acc_2 = arith.constant 1 : i32 loc(#loc40) + %acc_3 = arith.subi %c1_i32, %acc_2 : i32 loc(#loc40) + %acc_4 = arith.addi %acc_1, %acc_3 : i32 loc(#loc40) + %acc_5 = arith.divui %acc_4, %c1_i32 : i32 loc(#loc40) + %acc_6 = arith.constant 2 : i32 loc(#loc40) + %acc_7 = arith.remsi %acc_5, %acc_6 : i32 loc(#loc40) + %acc_8 = arith.subi %acc_5, %acc_7 : i32 loc(#loc40) + %acc_9 = arith.muli %acc_8, %c1_i32 : i32 loc(#loc40) + %acc_10 = arith.addi %c0_i32, %acc_9 : i32 loc(#loc40) + %acc_11 = arith.muli %c1_i32, %acc_6 : i32 loc(#loc40) + %acc_12 = scf.for %i = %c0_i32 to %acc_10 step %acc_11 iter_args(%acc_14 = %acc) -> (tensor<64xf32>) : i32 { + %2 = arith.remsi %i, %c3_i32 : i32 loc(#loc7) + %3 = arith.cmpi eq, %2, %c0_i32 : i32 loc(#loc8) + %4:2 = scf.if %3 -> (tensor<64xf32>, tensor<64x!tt.ptr>) { + %acc_22 = scf.for %j = %i to %m step %c2_i32 iter_args(%s = %acc_14) -> (tensor<64xf32>) : i32 { + %v = arith.subi %m, %j : i32 loc(#loc42) + %v_24 = tt.splat %v : i32 -> tensor<64xi32> loc(#loc43) + %v_25 = arith.cmpi slt, %offs, %v_24 : tensor<64xi32> loc(#loc43) + %v_26 = arith.muli %j, %c64_i32 : i32 loc(#loc44) + %v_27 = tt.splat %v_26 : i32 -> tensor<64xi32> loc(#loc45) + %v_28 = tt.addptr %p_0, %v_27 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc45) + %v_29 = tt.load %v_28, %v_25 : tensor<64x!tt.ptr> loc(#loc46) + %8 = arith.cmpi sgt, %j, %c5_i32 : i32 loc(#loc16) + scf.if %8 { + %9 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc18) + %10 = tt.addptr %9, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc18) + %11 = tt.splat %j : i32 -> tensor<64xi32> loc(#loc19) + %12 = tt.addptr %10, %11 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc19) + tt.store %12, %v_29 : tensor<64x!tt.ptr> loc(#loc20) + } loc(#loc17) + %s_30 = arith.addf %s, %v_29 : tensor<64xf32> loc(#loc47) + scf.yield %s_30 : tensor<64xf32> loc(#loc22) + } {tt.num_stages = 2 : i32} loc(#loc54) + %q = tt.splat %y_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc48) + %q_23 = tt.addptr %q, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc55) + scf.yield %acc_22, %q_23 : tensor<64xf32>, tensor<64x!tt.ptr> loc(#loc55) + } else { + %acc_22 = arith.mulf %acc_14, %cst : tensor<64xf32> loc(#loc56) + %q = tt.splat %i : i32 -> tensor<64xi32> loc(#loc50) + %q_23 = tt.addptr %p_0, %q : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc57) + scf.yield %acc_22, %q_23 : tensor<64xf32>, tensor<64x!tt.ptr> loc(#loc50) + } loc(#loc9) + %acc_15 = tt.load %4#1 : tensor<64x!tt.ptr> loc(#loc51) + %acc_16 = arith.addf %4#0, %acc_15 : tensor<64xf32> loc(#loc52) + %acc_17 = arith.constant 1 : i32 loc(#loc40) + %acc_18 = arith.muli %c1_i32, %acc_17 : i32 loc(#loc40) + %acc_19 = arith.addi %i, %acc_18 : i32 loc(#loc40) + %5 = arith.remsi %acc_19, %c3_i32 : i32 loc(#loc7) + %6 = arith.cmpi eq, %5, %c0_i32 : i32 loc(#loc8) + %7:2 = scf.if %6 -> (tensor<64xf32>, tensor<64x!tt.ptr>) { + %acc_22 = scf.for %j = %acc_19 to %m step %c2_i32 iter_args(%s = %acc_16) -> (tensor<64xf32>) : i32 { + %v = arith.subi %m, %j : i32 loc(#loc42) + %v_24 = tt.splat %v : i32 -> tensor<64xi32> loc(#loc43) + %v_25 = arith.cmpi slt, %offs, %v_24 : tensor<64xi32> loc(#loc43) + %v_26 = arith.muli %j, %c64_i32 : i32 loc(#loc44) + %v_27 = tt.splat %v_26 : i32 -> tensor<64xi32> loc(#loc45) + %v_28 = tt.addptr %p_0, %v_27 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc45) + %v_29 = tt.load %v_28, %v_25 : tensor<64x!tt.ptr> loc(#loc46) + %8 = arith.cmpi sgt, %j, %c5_i32 : i32 loc(#loc16) + scf.if %8 { + %9 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc18) + %10 = tt.addptr %9, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc18) + %11 = tt.splat %j : i32 -> tensor<64xi32> loc(#loc19) + %12 = tt.addptr %10, %11 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc19) + tt.store %12, %v_29 : tensor<64x!tt.ptr> loc(#loc20) + } loc(#loc17) + %s_30 = arith.addf %s, %v_29 : tensor<64xf32> loc(#loc47) + scf.yield %s_30 : tensor<64xf32> loc(#loc22) + } {tt.num_stages = 2 : i32} loc(#loc54) + %q = tt.splat %y_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc48) + %q_23 = tt.addptr %q, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc55) + scf.yield %acc_22, %q_23 : tensor<64xf32>, tensor<64x!tt.ptr> loc(#loc55) + } else { + %acc_22 = arith.mulf %acc_16, %cst : tensor<64xf32> loc(#loc56) + %q = tt.splat %acc_19 : i32 -> tensor<64xi32> loc(#loc50) + %q_23 = tt.addptr %p_0, %q : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc57) + scf.yield %acc_22, %q_23 : tensor<64xf32>, tensor<64x!tt.ptr> loc(#loc50) + } loc(#loc9) + %acc_20 = tt.load %7#1 : tensor<64x!tt.ptr> loc(#loc51) + %acc_21 = arith.addf %7#0, %acc_20 : tensor<64xf32> loc(#loc52) + scf.yield %acc_21 : tensor<64xf32> loc(#loc28) + } {tt.disallow_acc_multi_buffer, tt.flatten, tt.num_stages = 3 : i32} loc(#loc40) + %acc_13 = scf.for %i = %acc_10 to %n step %c1_i32 iter_args(%acc_14 = %acc_12) -> (tensor<64xf32>) : i32 { + %2 = arith.remsi %i, %c3_i32 : i32 loc(#loc7) + %3 = arith.cmpi eq, %2, %c0_i32 : i32 loc(#loc8) + %4:2 = scf.if %3 -> (tensor<64xf32>, tensor<64x!tt.ptr>) { + %acc_17 = scf.for %j = %i to %m step %c2_i32 iter_args(%s = %acc_14) -> (tensor<64xf32>) : i32 { + %v = arith.subi %m, %j : i32 loc(#loc42) + %v_19 = tt.splat %v : i32 -> tensor<64xi32> loc(#loc43) + %v_20 = arith.cmpi slt, %offs, %v_19 : tensor<64xi32> loc(#loc43) + %v_21 = arith.muli %j, %c64_i32 : i32 loc(#loc44) + %v_22 = tt.splat %v_21 : i32 -> tensor<64xi32> loc(#loc45) + %v_23 = tt.addptr %p_0, %v_22 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc45) + %v_24 = tt.load %v_23, %v_20 : tensor<64x!tt.ptr> loc(#loc46) + %5 = arith.cmpi sgt, %j, %c5_i32 : i32 loc(#loc16) + scf.if %5 { + %6 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc18) + %7 = tt.addptr %6, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc18) + %8 = tt.splat %j : i32 -> tensor<64xi32> loc(#loc19) + %9 = tt.addptr %7, %8 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc19) + tt.store %9, %v_24 : tensor<64x!tt.ptr> loc(#loc20) + } loc(#loc17) + %s_25 = arith.addf %s, %v_24 : tensor<64xf32> loc(#loc47) + scf.yield %s_25 : tensor<64xf32> loc(#loc22) + } {tt.num_stages = 2 : i32} loc(#loc54) + %q = tt.splat %y_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc48) + %q_18 = tt.addptr %q, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc55) + scf.yield %acc_17, %q_18 : tensor<64xf32>, tensor<64x!tt.ptr> loc(#loc55) + } else { + %acc_17 = arith.mulf %acc_14, %cst : tensor<64xf32> loc(#loc56) + %q = tt.splat %i : i32 -> tensor<64xi32> loc(#loc50) + %q_18 = tt.addptr %p_0, %q : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc57) + scf.yield %acc_17, %q_18 : tensor<64xf32>, tensor<64x!tt.ptr> loc(#loc50) + } loc(#loc9) + %acc_15 = tt.load %4#1 : tensor<64x!tt.ptr> loc(#loc51) + %acc_16 = arith.addf %4#0, %acc_15 : tensor<64xf32> loc(#loc52) + scf.yield %acc_16 : tensor<64xf32> loc(#loc28) + } {tt.disallow_acc_multi_buffer, tt.flatten, tt.num_stages = 1 : i32} loc(#loc40) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc29) + %1 = tt.addptr %0, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc29) + tt.store %1, %acc_13 : tensor<64x!tt.ptr> loc(#loc30) + tt.return loc(#loc31) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":129:31) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":17:19) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":19:78) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":16:24) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":18:16) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":20:15) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":20:20) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":20:11) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":22:39) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":23:59) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":23:55) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":23:36) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":23:32) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":23:28) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":24:23) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":24:19) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":25:39) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":25:46) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":25:49) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":26:21) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":26:16) +#loc23 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":28:24) +#loc24 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":30:24) +#loc25 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":31:31) +#loc26 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":32:23) +#loc27 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":32:15) +#loc28 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":32:8) +#loc29 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":33:23) +#loc30 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":33:29) +#loc31 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":33:4) +#loc37 = loc("acc"(#loc2)) +#loc38 = loc("offs"(#loc5)) +#loc39 = loc("p"(#loc6)) +#loc40 = loc("acc"(#loc4)) +#loc41 = loc("s"(#loc10)) +#loc42 = loc("v"(#loc11)) +#loc43 = loc("v"(#loc12)) +#loc44 = loc("v"(#loc13)) +#loc45 = loc("v"(#loc14)) +#loc46 = loc("v"(#loc15)) +#loc47 = loc("s"(#loc21)) +#loc48 = loc("q"(#loc23)) +#loc49 = loc("acc"(#loc24)) +#loc50 = loc("q"(#loc25)) +#loc51 = loc("acc"(#loc26)) +#loc52 = loc("acc"(#loc27)) +#loc53 = loc(callsite(#loc1 at #loc37)) +#loc54 = loc("acc"(#loc41)) +#loc55 = loc("q"(#loc48)) +#loc56 = loc("acc"(#loc49)) +#loc57 = loc("q"(#loc50)) diff --git a/tests/golden/ir/ttir/adv_reduce3.ttir b/tests/golden/ir/ttir/adv_reduce3.ttir new file mode 100644 index 000000000..d8784a677 --- /dev/null +++ b/tests/golden/ir/ttir/adv_reduce3.ttir @@ -0,0 +1,218 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":50:0) +#loc6 = loc(unknown) +#loc33 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":57:15) +#loc50 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":257:24) +#loc51 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":63:25) +#loc65 = loc("x_ptr"(#loc)) +#loc66 = loc("i_ptr"(#loc)) +#loc67 = loc("out_ptr"(#loc)) +#loc68 = loc("n"(#loc)) +#loc89 = loc("s"(#loc33)) +#loc94 = loc("tot"(#loc51)) +#loc110 = loc(callsite(#loc6 at #loc89)) +#loc112 = loc(callsite(#loc50 at #loc94)) +#loc115 = loc(callsite(#loc6 at #loc112)) +module { + tt.func public @reduce3(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %i_ptr: !tt.ptr loc("i_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %tot = arith.constant dense<0> : tensor<16xi32> loc(#loc103) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %c16_i32 = arith.constant 16 : i32 loc(#loc4) + %cst = arith.constant dense<2.000000e+00> : tensor<16x32xf32> loc(#loc5) + %cst_0 = arith.constant dense<32> : tensor<16x1xi32> loc(#loc6) + %c32_i32 = arith.constant 32 : i32 loc(#loc6) + %rm = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32> loc(#loc70) + %rn = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc71) + %x = tt.expand_dims %rm {axis = 1 : i32} : tensor<16xi32> -> tensor<16x1xi32> loc(#loc72) + %x_1 = arith.muli %x, %cst_0 : tensor<16x1xi32> loc(#loc73) + %x_2 = tt.splat %x_ptr : !tt.ptr -> tensor<16x1x!tt.ptr> loc(#loc74) + %x_3 = tt.addptr %x_2, %x_1 : tensor<16x1x!tt.ptr>, tensor<16x1xi32> loc(#loc74) + %x_4 = tt.expand_dims %rn {axis = 0 : i32} : tensor<32xi32> -> tensor<1x32xi32> loc(#loc75) + %x_5 = tt.broadcast %x_3 : tensor<16x1x!tt.ptr> -> tensor<16x32x!tt.ptr> loc(#loc76) + %x_6 = tt.broadcast %x_4 : tensor<1x32xi32> -> tensor<16x32xi32> loc(#loc76) + %x_7 = tt.addptr %x_5, %x_6 : tensor<16x32x!tt.ptr>, tensor<16x32xi32> loc(#loc76) + %x_8 = tt.load %x_7 : tensor<16x32x!tt.ptr> loc(#loc77) + %idx = tt.splat %i_ptr : !tt.ptr -> tensor<16x1x!tt.ptr> loc(#loc78) + %idx_9 = tt.addptr %idx, %x_1 : tensor<16x1x!tt.ptr>, tensor<16x1xi32> loc(#loc78) + %idx_10 = tt.broadcast %idx_9 : tensor<16x1x!tt.ptr> -> tensor<16x32x!tt.ptr> loc(#loc79) + %idx_11 = tt.addptr %idx_10, %x_6 : tensor<16x32x!tt.ptr>, tensor<16x32xi32> loc(#loc79) + %idx_12 = tt.load %idx_11 : tensor<16x32x!tt.ptr> loc(#loc80) + %0 = arith.mulf %x_8, %cst : tensor<16x32xf32> loc(#loc5) + %1:3 = "tt.reduce"(%x_8, %idx_12, %0) <{axis = 1 : i32}> ({ + ^bb0(%arg4: f32 loc(unknown), %arg5: i32 loc(unknown), %arg6: f32 loc(unknown), %arg7: f32 loc(unknown), %arg8: i32 loc(unknown), %arg9: f32 loc(unknown)): + %take = arith.cmpf ogt, %arg4, %arg7 : f32 loc(#loc104) + %take_15 = arith.cmpf oeq, %arg4, %arg7 : f32 loc(#loc105) + %take_16 = arith.cmpi slt, %arg5, %arg8 : i32 loc(#loc106) + %take_17 = arith.andi %take_15, %take_16 : i1 loc(#loc107) + %take_18 = arith.ori %take, %take_17 : i1 loc(#loc108) + %22 = arith.select %take_18, %arg4, %arg7 : f32 loc(#loc86) + %23 = arith.select %take_18, %arg5, %arg8 : i32 loc(#loc87) + %24 = arith.addf %arg6, %arg9 : f32 loc(#loc88) + tt.reduce.return %22, %23, %24 : f32, i32, f32 loc(#loc18) + }) : (tensor<16x32xf32>, tensor<16x32xi32>, tensor<16x32xf32>) -> (tensor<16xf32>, tensor<16xi32>, tensor<16xf32>) loc(#loc18) + %2 = tt.splat %out_ptr : !tt.ptr -> tensor<16x!tt.ptr> loc(#loc27) + %3 = tt.addptr %2, %rm : tensor<16x!tt.ptr>, tensor<16xi32> loc(#loc27) + %4 = arith.sitofp %1#1 : tensor<16xi32> to tensor<16xf32> loc(#loc28) + %5 = arith.addf %1#0, %4 : tensor<16xf32> loc(#loc29) + %6 = arith.addf %5, %1#2 : tensor<16xf32> loc(#loc30) + tt.store %3, %6 : tensor<16x!tt.ptr> loc(#loc31) + %s = "tt.reduce"(%x_8) <{axis = 1 : i32}> ({ + ^bb0(%s_15: f32 loc(callsite(#loc6 at #loc89)), %s_16: f32 loc(callsite(#loc6 at #loc89))): + %s_17 = arith.addf %s_15, %s_16 : f32 loc(#loc113) + tt.reduce.return %s_17 : f32 loc(#loc109) + }) : (tensor<16x32xf32>) -> tensor<16xf32> loc(#loc109) + %s_13 = tt.expand_dims %s {axis = 1 : i32} : tensor<16xf32> -> tensor<16x1xf32> loc(#loc111) + %7 = tt.addptr %out_ptr, %c16_i32 : !tt.ptr, i32 loc(#loc4) + %8 = tt.splat %7 : !tt.ptr -> tensor<16x1x!tt.ptr> loc(#loc36) + %9 = tt.addptr %8, %x : tensor<16x1x!tt.ptr>, tensor<16x1xi32> loc(#loc36) + %10 = tt.broadcast %9 : tensor<16x1x!tt.ptr> -> tensor<16x32x!tt.ptr> loc(#loc37) + %11 = tt.broadcast %s_13 : tensor<16x1xf32> -> tensor<16x32xf32> loc(#loc38) + tt.store %10, %11 : tensor<16x32x!tt.ptr> loc(#loc38) + %12:2 = "tt.scan"(%x_8, %idx_12) <{axis = 1 : i32, reverse = true}> ({ + ^bb0(%arg4: f32 loc(unknown), %arg5: i32 loc(unknown), %arg6: f32 loc(unknown), %arg7: i32 loc(unknown)): + %22 = arith.addf %arg4, %arg6 : f32 loc(#loc90) + %23 = arith.maxsi %arg5, %arg7 : i32 loc(#loc91) + tt.scan.return %22, %23 : f32, i32 loc(#loc39) + }) : (tensor<16x32xf32>, tensor<16x32xi32>) -> (tensor<16x32xf32>, tensor<16x32xi32>) loc(#loc39) + %13 = tt.addptr %out_ptr, %c32_i32 : !tt.ptr, i32 loc(#loc42) + %14 = tt.splat %13 : !tt.ptr -> tensor<16x1x!tt.ptr> loc(#loc43) + %15 = tt.addptr %14, %x_1 : tensor<16x1x!tt.ptr>, tensor<16x1xi32> loc(#loc43) + %16 = tt.broadcast %15 : tensor<16x1x!tt.ptr> -> tensor<16x32x!tt.ptr> loc(#loc44) + %17 = tt.addptr %16, %x_6 : tensor<16x32x!tt.ptr>, tensor<16x32xi32> loc(#loc44) + %18 = arith.sitofp %12#1 : tensor<16x32xi32> to tensor<16x32xf32> loc(#loc45) + %19 = arith.addf %12#0, %18 : tensor<16x32xf32> loc(#loc46) + tt.store %17, %19 : tensor<16x32x!tt.ptr> loc(#loc47) + %tot_14 = scf.for %k = %c0_i32 to %n step %c1_i32 iter_args(%tot_15 = %tot) -> (tensor<16xi32>) : i32 { + %tot_16 = arith.sitofp %k : i32 to f32 loc(#loc93) + %tot_17 = tt.splat %tot_16 : f32 -> tensor<16x32xf32> loc(#loc93) + %tot_18 = arith.addf %x_8, %tot_17 : tensor<16x32xf32> loc(#loc93) + %tot_19:2 = "tt.reduce"(%tot_18, %x_6) <{axis = 1 : i32}> ({ + ^bb0(%tot_21: f32 loc(callsite(#loc6 at #loc112)), %tot_22: i32 loc(callsite(#loc6 at #loc112)), %tot_23: f32 loc(callsite(#loc6 at #loc112)), %tot_24: i32 loc(callsite(#loc6 at #loc112))): + %tie = arith.cmpf oeq, %tot_21, %tot_23 : f32 loc(#loc117) + %tie_25 = arith.cmpi slt, %tot_22, %tot_24 : i32 loc(#loc118) + %tie_26 = arith.andi %tie, %tie_25 : i1 loc(#loc119) + %lt = arith.cmpf olt, %tot_21, %tot_23 : f32 loc(#loc120) + %lt_27 = arith.ori %lt, %tie_26 : i1 loc(#loc121) + %value_ret = arith.select %lt_27, %tot_21, %tot_23 : f32 loc(#loc122) + %index_ret = arith.select %lt_27, %tot_22, %tot_24 : i32 loc(#loc123) + tt.reduce.return %value_ret, %index_ret : f32, i32 loc(#loc114) + }) : (tensor<16x32xf32>, tensor<16x32xi32>) -> (tensor<16xf32>, tensor<16xi32>) loc(#loc114) + %tot_20 = arith.addi %tot_15, %tot_19#1 : tensor<16xi32> loc(#loc102) + scf.yield %tot_20 : tensor<16xi32> loc(#loc61) + } loc(#loc92) + %20 = tt.splat %i_ptr : !tt.ptr -> tensor<16x!tt.ptr> loc(#loc62) + %21 = tt.addptr %20, %rm : tensor<16x!tt.ptr>, tensor<16xi32> loc(#loc62) + tt.store %21, %tot_14 : tensor<16x!tt.ptr> loc(#loc63) + tt.return loc(#loc64) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":129:31) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":61:19) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":62:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":58:23) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":55:37) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":51:22) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":52:22) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":53:27) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":53:38) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":53:24) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":53:46) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":53:43) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":53:16) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":54:26) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":54:45) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":54:18) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":55:46) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":38:17) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":38:31) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":38:43) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":38:38) +#loc23 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":38:24) +#loc24 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":39:30) +#loc25 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":39:54) +#loc26 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":39:64) +#loc27 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":56:23) +#loc28 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":56:36) +#loc29 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":56:31) +#loc30 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":56:50) +#loc31 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":56:27) +#loc32 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":293:36) +#loc34 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":263:15) +#loc35 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":287:0) +#loc36 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":58:28) +#loc37 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":58:42) +#loc38 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":58:59) +#loc39 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":59:44) +#loc40 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":44:16) +#loc41 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":44:35) +#loc42 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":60:23) +#loc43 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":60:32) +#loc44 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":60:51) +#loc45 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":60:73) +#loc46 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":60:68) +#loc47 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":60:64) +#loc48 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":63:29) +#loc49 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":240:58) +#loc52 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":208:24) +#loc53 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":219:59) +#loc54 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":208:44) +#loc55 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":208:35) +#loc56 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":211:18) +#loc57 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":211:28) +#loc58 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":212:39) +#loc59 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":213:39) +#loc60 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":63:15) +#loc61 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":63:8) +#loc62 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":64:21) +#loc63 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":64:25) +#loc64 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":64:4) +#loc69 = loc("tot"(#loc2)) +#loc70 = loc("rm"(#loc7)) +#loc71 = loc("rn"(#loc8)) +#loc72 = loc("x"(#loc9)) +#loc73 = loc("x"(#loc10)) +#loc74 = loc("x"(#loc11)) +#loc75 = loc("x"(#loc12)) +#loc76 = loc("x"(#loc13)) +#loc77 = loc("x"(#loc14)) +#loc78 = loc("idx"(#loc15)) +#loc79 = loc("idx"(#loc16)) +#loc80 = loc("idx"(#loc17)) +#loc81 = loc("take"(#loc19)) +#loc82 = loc("take"(#loc20)) +#loc83 = loc("take"(#loc21)) +#loc84 = loc("take"(#loc22)) +#loc85 = loc("take"(#loc23)) +#loc86 = loc(callsite(#loc24 at #loc18)) +#loc87 = loc(callsite(#loc25 at #loc18)) +#loc88 = loc(callsite(#loc26 at #loc18)) +#loc90 = loc(callsite(#loc40 at #loc39)) +#loc91 = loc(callsite(#loc41 at #loc39)) +#loc92 = loc("tot"(#loc3)) +#loc93 = loc("tot"(#loc48)) +#loc95 = loc("tie"(#loc52)) +#loc96 = loc("tie"(#loc54)) +#loc97 = loc("tie"(#loc55)) +#loc98 = loc("lt"(#loc56)) +#loc99 = loc("lt"(#loc57)) +#loc100 = loc("value_ret"(#loc58)) +#loc101 = loc("index_ret"(#loc59)) +#loc102 = loc("tot"(#loc60)) +#loc103 = loc(callsite(#loc1 at #loc69)) +#loc104 = loc(callsite(#loc81 at #loc18)) +#loc105 = loc(callsite(#loc82 at #loc18)) +#loc106 = loc(callsite(#loc83 at #loc18)) +#loc107 = loc(callsite(#loc84 at #loc18)) +#loc108 = loc(callsite(#loc85 at #loc18)) +#loc109 = loc(callsite(#loc32 at #loc89)) +#loc111 = loc(callsite(#loc35 at #loc89)) +#loc113 = loc(callsite(#loc34 at #loc109)) +#loc114 = loc(callsite(#loc49 at #loc112)) +#loc116 = loc(callsite(#loc53 at #loc114)) +#loc117 = loc(callsite(#loc95 at #loc116)) +#loc118 = loc(callsite(#loc96 at #loc116)) +#loc119 = loc(callsite(#loc97 at #loc116)) +#loc120 = loc(callsite(#loc98 at #loc116)) +#loc121 = loc(callsite(#loc99 at #loc116)) +#loc122 = loc(callsite(#loc100 at #loc116)) +#loc123 = loc(callsite(#loc101 at #loc116)) diff --git a/tests/golden/ir/ttir/adv_views.ttir b/tests/golden/ir/ttir/adv_views.ttir new file mode 100644 index 000000000..eba42ca2b --- /dev/null +++ b/tests/golden/ir/ttir/adv_views.ttir @@ -0,0 +1,76 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":219:0) +#loc23 = loc("x_ptr"(#loc)) +#loc24 = loc("out_ptr"(#loc)) +module { + tt.func public @views(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc))) attributes {noinline = false} { + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %g = arith.constant dense<63> : tensor<64xi32> loc(#loc25) + %q = arith.constant dense<16> : tensor<4x1xi32> loc(#loc26) + %m2 = arith.constant dense<50> : tensor<64xi32> loc(#loc27) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc28) + %p = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc29) + %p_0 = tt.addptr %p, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc29) + %p2 = tt.reshape %p_0 allow_reorder : tensor<64x!tt.ptr> -> tensor<16x4x!tt.ptr> loc(#loc30) + %m2_1 = arith.cmpi slt, %offs, %m2 : tensor<64xi32> loc(#loc27) + %m2_2 = tt.reshape %m2_1 allow_reorder : tensor<64xi1> -> tensor<16x4xi1> loc(#loc31) + %v = tt.load %p2, %m2_2 : tensor<16x4x!tt.ptr> loc(#loc32) + %t = tt.trans %v {order = array} : tensor<16x4xf32> -> tensor<4x16xf32> loc(#loc33) + %q_3 = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> loc(#loc34) + %q_4 = tt.expand_dims %q_3 {axis = 1 : i32} : tensor<4xi32> -> tensor<4x1xi32> loc(#loc35) + %q_5 = arith.muli %q_4, %q : tensor<4x1xi32> loc(#loc26) + %q_6 = tt.splat %out_ptr : !tt.ptr -> tensor<4x1x!tt.ptr> loc(#loc36) + %q_7 = tt.addptr %q_6, %q_5 : tensor<4x1x!tt.ptr>, tensor<4x1xi32> loc(#loc36) + %q_8 = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32> loc(#loc37) + %q_9 = tt.expand_dims %q_8 {axis = 0 : i32} : tensor<16xi32> -> tensor<1x16xi32> loc(#loc38) + %q_10 = tt.broadcast %q_7 : tensor<4x1x!tt.ptr> -> tensor<4x16x!tt.ptr> loc(#loc39) + %q_11 = tt.broadcast %q_9 : tensor<1x16xi32> -> tensor<4x16xi32> loc(#loc39) + %q_12 = tt.addptr %q_10, %q_11 : tensor<4x16x!tt.ptr>, tensor<4x16xi32> loc(#loc39) + tt.store %q_12, %t : tensor<4x16x!tt.ptr> loc(#loc17) + %g_13 = arith.subi %g, %offs : tensor<64xi32> loc(#loc25) + %g_14 = tt.gather %offs[%g_13] {axis = 0 : i32} : (tensor<64xi32>, tensor<64xi32>) -> tensor<64xi32> loc(#loc40) + %0 = tt.addptr %out_ptr, %c64_i32 : !tt.ptr, i32 loc(#loc1) + %1 = tt.splat %0 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc19) + %2 = tt.addptr %1, %g_14 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc19) + %3 = tt.load %p_0 : tensor<64x!tt.ptr> loc(#loc20) + tt.store %2, %3 : tensor<64x!tt.ptr> loc(#loc21) + tt.return loc(#loc22) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":229:23) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":228:38) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":226:46) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":223:27) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":220:24) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":221:16) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":222:23) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":223:31) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":224:16) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":225:17) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":226:31) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":226:34) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":226:18) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":226:73) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":226:85) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":226:60) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":227:16) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":228:44) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":229:31) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":229:42) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":229:34) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":229:4) +#loc25 = loc("g"(#loc2)) +#loc26 = loc("q"(#loc3)) +#loc27 = loc("m2"(#loc4)) +#loc28 = loc("offs"(#loc5)) +#loc29 = loc("p"(#loc6)) +#loc30 = loc("p2"(#loc7)) +#loc31 = loc("m2"(#loc8)) +#loc32 = loc("v"(#loc9)) +#loc33 = loc("t"(#loc10)) +#loc34 = loc("q"(#loc11)) +#loc35 = loc("q"(#loc12)) +#loc36 = loc("q"(#loc13)) +#loc37 = loc("q"(#loc14)) +#loc38 = loc("q"(#loc15)) +#loc39 = loc("q"(#loc16)) +#loc40 = loc("g"(#loc18)) diff --git a/tests/golden/ir/ttir/adv_while_nested.ttir b/tests/golden/ir/ttir/adv_while_nested.ttir new file mode 100644 index 000000000..cc8ccbecb --- /dev/null +++ b/tests/golden/ir/ttir/adv_while_nested.ttir @@ -0,0 +1,63 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":180:0) +#loc4 = loc("i") +#loc5 = loc("acc") +#loc19 = loc("p_ptr"(#loc)) +#loc20 = loc("out_ptr"(#loc)) +#loc21 = loc("n"(#loc)) +module { + tt.func public @while_nested(%p_ptr: !tt.ptr loc("p_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c2_i32 = arith.constant 2 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %acc:2 = scf.while (%i = %c0_i32, %acc_0 = %c0_i32) : (i32, i32) -> (i32, i32) { + %0 = arith.cmpi slt, %i, %n : i32 loc(#loc3) + scf.condition(%0) %i, %acc_0 : i32, i32 loc(#loc3) + } do { + ^bb0(%i: i32 loc("i"), %acc_0: i32 loc("acc")): + %0 = arith.remsi %i, %c2_i32 : i32 loc(#loc6) + %1 = arith.cmpi eq, %0, %c0_i32 : i32 loc(#loc7) + %2 = scf.if %1 -> (i32) { + %acc_2 = scf.for %k = %c0_i32 to %i step %c1_i32 iter_args(%acc_3 = %acc_0) -> (i32) : i32 { + %acc_4 = tt.addptr %p_ptr, %k : !tt.ptr, i32 loc(#loc24) + %acc_5 = tt.load %acc_4 : !tt.ptr loc(#loc25) + %acc_6 = arith.addi %acc_3, %acc_5 : i32 loc(#loc26) + scf.yield %acc_6 : i32 loc(#loc13) + } loc(#loc30) + scf.yield %acc_2 : i32 loc(#loc30) + } else { + %acc_2 = arith.subi %acc_0, %c1_i32 : i32 loc(#loc31) + scf.yield %acc_2 : i32 loc(#loc27) + } loc(#loc8) + %i_1 = arith.addi %i, %c1_i32 : i32 loc(#loc28) + scf.yield %i_1, %2 : i32, i32 loc(#loc16) + } loc(#loc29) + tt.store %out_ptr, %acc#1 : !tt.ptr loc(#loc17) + tt.return loc(#loc18) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":183:4) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":183:14) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":184:15) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":184:20) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":184:11) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":185:30) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":186:39) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":186:31) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":186:23) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":186:16) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":188:19) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":189:13) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":189:8) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":190:22) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":190:4) +#loc22 = loc("i"(#loc2)) +#loc23 = loc("acc"(#loc9)) +#loc24 = loc("acc"(#loc10)) +#loc25 = loc("acc"(#loc11)) +#loc26 = loc("acc"(#loc12)) +#loc27 = loc("acc"(#loc14)) +#loc28 = loc("i"(#loc15)) +#loc29 = loc("acc"(#loc22)) +#loc30 = loc("acc"(#loc23)) +#loc31 = loc("acc"(#loc27)) diff --git a/tests/golden/ir/ttir/adv_zero_result.ttir b/tests/golden/ir/ttir/adv_zero_result.ttir new file mode 100644 index 000000000..400026111 --- /dev/null +++ b/tests/golden/ir/ttir/adv_zero_result.ttir @@ -0,0 +1,59 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":117:0) +#loc19 = loc("x_ptr"(#loc)) +#loc20 = loc("out_ptr"(#loc)) +#loc21 = loc("n"(#loc)) +module { + tt.func public @zero_result(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %true = arith.constant true loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc22) + %offs = arith.muli %pid, %c64_i32 : i32 loc(#loc23) + %offs_0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc24) + %offs_1 = tt.splat %offs : i32 -> tensor<64xi32> loc(#loc25) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<64xi32> loc(#loc25) + %v = tt.splat %n : i32 -> tensor<64xi32> loc(#loc26) + %v_3 = arith.cmpi slt, %offs_2, %v : tensor<64xi32> loc(#loc26) + %v_4 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc27) + %v_5 = tt.addptr %v_4, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc27) + %v_6 = tt.load %v_5, %v_3 : tensor<64x!tt.ptr> loc(#loc28) + tt.print " pid=: " {hex = true, isSigned = array} : %pid, %offs_2 : i32, tensor<64xi32> loc(#loc10) + gpu.barrier loc(#loc11) + %0 = arith.cmpi eq, %pid, %c0_i32 : i32 loc(#loc2) + scf.if %0 { + %5 = tt.atomic_rmw add, acq_rel, gpu, %out_ptr, %c1_i32, %true : (!tt.ptr, i32, i1) -> i32 loc(#loc13) + } loc(#loc12) + %1 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc14) + %2 = tt.addptr %1, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc14) + %3 = arith.fptosi %v_6 : tensor<64xf32> to tensor<64xi32> loc(#loc15) + %4 = tt.atomic_rmw max, acq_rel, gpu, %2, %3, %v_3 : (tensor<64x!tt.ptr>, tensor<64xi32>, tensor<64xi1>) -> tensor<64xi32> loc(#loc16) + tt.store %2, %3, %v_3 : tensor<64x!tt.ptr> loc(#loc17) + tt.return loc(#loc18) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":124:14) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":118:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":119:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":119:38) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":119:25) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":120:42) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":120:24) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":120:16) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":122:33) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":123:4) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":124:7) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":125:31) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":126:28) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":126:39) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":126:34) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":127:29) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/adv_kernels.py":127:4) +#loc22 = loc("pid"(#loc3)) +#loc23 = loc("offs"(#loc4)) +#loc24 = loc("offs"(#loc5)) +#loc25 = loc("offs"(#loc6)) +#loc26 = loc("v"(#loc7)) +#loc27 = loc("v"(#loc8)) +#loc28 = loc("v"(#loc9)) diff --git a/tests/golden/ir/ttir/crafted_attr_dicts.ttir b/tests/golden/ir/ttir/crafted_attr_dicts.ttir new file mode 100644 index 000000000..bfd6db3ba --- /dev/null +++ b/tests/golden/ir/ttir/crafted_attr_dicts.ttir @@ -0,0 +1,18 @@ +module { + tt.func public @k(%p: !tt.ptr) attributes {noinline = false} { + %c = arith.constant {tt.divisibility = dense<16> : tensor<1xi32>} 16 : i32 + %d = arith.constant {axis = 5 : i32} -3 : i32 + %r = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> + %s = "tt.reduce"(%r) <{axis = 0 : i32}> ({ + ^bb0(%a: i32, %b: i32): + %t = arith.addi %a, %b : i32 + tt.reduce.return %t : i32 + }) {tt.divisibility = dense<16> : tensor<1xi32>, axis_note = 3 : i32} : (tensor<64xi32>) -> i32 + %e = tt.expand_dims %r {axis = 0 : i32, tt.note = "axis = 1"} : tensor<64xi32> -> tensor<1x64xi32> + %q = tt.addptr %p, %s : !tt.ptr, i32 + %q2 = tt.addptr %q, %c : !tt.ptr, i32 + %q3 = tt.addptr %q2, %d : !tt.ptr, i32 + tt.store %q3, %s : !tt.ptr + tt.return + } +} diff --git a/tests/golden/ir/ttir/crafted_deep_nest.ttir b/tests/golden/ir/ttir/crafted_deep_nest.ttir new file mode 100644 index 000000000..174de5379 --- /dev/null +++ b/tests/golden/ir/ttir/crafted_deep_nest.ttir @@ -0,0 +1,43 @@ +module { + tt.func public @k(%p: !tt.ptr, %n: i32) attributes {noinline = false} { + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %r = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> + %v = arith.sitofp %r : tensor<64xi32> to tensor<64xf32> + %s:2 = "tt.reduce"(%v, %r) <{axis = 0 : i32}> ({ + ^bb0(%a: f32, %ai: i32, %b: f32, %bi: i32): + %gt = arith.cmpf ogt, %a, %b : f32 + %o:2 = scf.if %gt -> (f32, i32) { + %acc = scf.for %i = %c0 to %n step %c1 iter_args(%t = %a) -> (f32) : i32 { + %lt = arith.cmpi slt, %i, %ai : i32 + %u = scf.if %lt -> (f32) { + %w = arith.addf %t, %b : f32 + scf.yield %w : f32 + } else { + scf.yield %t : f32 + } + scf.yield %u : f32 + } + scf.yield %acc, %ai : f32, i32 + } else { + scf.yield %b, %bi : f32, i32 + } + tt.reduce.return %o#0, %o#1 : f32, i32 + }) : (tensor<64xf32>, tensor<64xi32>) -> (f32, i32) + %z = scf.for %i = %c0 to %n step %c1 iter_args(%t = %c0) -> (i32) : i32 { + %wr = scf.while (%x = %t) : (i32) -> i32 { + %c = arith.cmpi slt, %x, %n : i32 + scf.condition(%c) %x : i32 + } do { + ^bb0(%y: i32): + %y1 = arith.addi %y, %c1 : i32 + scf.yield %y1 : i32 + } + scf.yield %wr : i32 + } + %q = tt.addptr %p, %s#1 : !tt.ptr, i32 + %q2 = tt.addptr %q, %z : !tt.ptr, i32 + tt.store %q2, %s#0 : !tt.ptr + tt.return + } +} diff --git a/tests/golden/ir/ttir/crafted_empty_bodies.ttir b/tests/golden/ir/ttir/crafted_empty_bodies.ttir new file mode 100644 index 000000000..077bdf4d2 --- /dev/null +++ b/tests/golden/ir/ttir/crafted_empty_bodies.ttir @@ -0,0 +1,24 @@ +#loc = loc("k.py":1:0) +module { + tt.func public @k(%p: !tt.ptr loc("p"(#loc)), %c: i1 loc("c"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c0 = arith.constant 0 : i32 loc(#loc1) + %c1 = arith.constant 1 : i32 loc(#loc1) + scf.if %c { + } else { + tt.store %p, %c0 : !tt.ptr loc(#loc2) + } loc(#loc1) + scf.if %c { + tt.store %p, %c1 : !tt.ptr loc(#loc2) + } else { + } loc(#loc1) + scf.for %i = %c0 to %n step %c1 : i32 { + } loc(#loc3) + scf.if %c { + } loc(#loc1) + tt.return loc(#loc4) + } loc(#loc) +} loc(#loc) +#loc1 = loc("k.py":2:4) +#loc2 = loc("k.py":3:8) +#loc3 = loc("k.py":4:4) +#loc4 = loc("k.py":5:4) diff --git a/tests/golden/ir/ttir/crafted_empty_else.ttir b/tests/golden/ir/ttir/crafted_empty_else.ttir new file mode 100644 index 000000000..ff28d8fb9 --- /dev/null +++ b/tests/golden/ir/ttir/crafted_empty_else.ttir @@ -0,0 +1,10 @@ +module { + tt.func public @k(%p: !tt.ptr, %c: i1) attributes {noinline = false} { + %c0 = arith.constant 0 : i32 + scf.if %c { + tt.store %p, %c0 : !tt.ptr + } else { + } + tt.return + } +} diff --git a/tests/golden/ir/ttir/crafted_empty_for.ttir b/tests/golden/ir/ttir/crafted_empty_for.ttir new file mode 100644 index 000000000..00138145a --- /dev/null +++ b/tests/golden/ir/ttir/crafted_empty_for.ttir @@ -0,0 +1,10 @@ +module { + tt.func public @k(%p: !tt.ptr, %n: i32) attributes {noinline = false} { + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + scf.for %i = %c0 to %n step %c1 : i32 { + } + tt.store %p, %c0 : !tt.ptr + tt.return + } +} diff --git a/tests/golden/ir/ttir/crafted_fwd_ref_cf.ttir b/tests/golden/ir/ttir/crafted_fwd_ref_cf.ttir new file mode 100644 index 000000000..3dc2d1196 --- /dev/null +++ b/tests/golden/ir/ttir/crafted_fwd_ref_cf.ttir @@ -0,0 +1,12 @@ +module { + tt.func public @k(%p: !tt.ptr, %n: i32) attributes {noinline = false} { + cf.br ^bb2 + ^bb1: + %q = tt.addptr %p, %x : !tt.ptr, i32 + tt.store %q, %x : !tt.ptr + tt.return + ^bb2: + %x = arith.addi %n, %n : i32 + cf.br ^bb1 + } +} diff --git a/tests/golden/ir/ttir/crafted_generic_form.ttir b/tests/golden/ir/ttir/crafted_generic_form.ttir new file mode 100644 index 000000000..5e6847581 --- /dev/null +++ b/tests/golden/ir/ttir/crafted_generic_form.ttir @@ -0,0 +1,11 @@ +module { + tt.func public @k(%p: !tt.ptr, %a: i32, %b: i32) attributes {noinline = false} { + %c = "arith.cmpi"(%a, %b) <{predicate = 2 : i64}> : (i32, i32) -> i1 + %x = "arith.select"(%c, %a, %b) : (i1, i32, i32) -> i32 + %o = "tt.atomic_rmw"(%p, %x) <{atomic_rmw_op = 5 : i32, scope = 1 : i32, sem = 4 : i32}> : (!tt.ptr, i32) -> i32 + %pid = "tt.get_program_id"() <{axis = 1 : i32}> : () -> i32 + %q = tt.addptr %p, %pid : !tt.ptr, i32 + tt.store %q, %o : !tt.ptr + tt.return + } +} diff --git a/tests/golden/ir/ttir/crafted_locs.ttir b/tests/golden/ir/ttir/crafted_locs.ttir new file mode 100644 index 000000000..471a94981 --- /dev/null +++ b/tests/golden/ir/ttir/crafted_locs.ttir @@ -0,0 +1,16 @@ +#loc = loc("k.py":1:0) +#loc1 = loc("k.py":2:5) +#loc2 = loc("k.py":3:6) +#loc9 = loc(fused<"meta">[#loc1, #loc2]) +#loc10 = loc(callsite(#loc1 at #loc9)) +module { + tt.func public @k(%p: !tt.ptr loc("p"(#loc)), %n: i32 loc(unknown)) attributes {noinline = false} { + %c1 = arith.constant 1 : i32 loc(fused[#loc1, "x.py":7:3]) + %a = arith.addi %n, %c1 : i32 loc(#loc9) + %b = arith.addi %a, %c1 : i32 loc(callsite("inner"("y.py":3:4) at callsite(#loc1 at #loc2))) + %q = tt.addptr %p, %b : !tt.ptr, i32 loc("q.py":9:9) + tt.store %q, %a : !tt.ptr loc(#loc10) + tt.return loc(#loc11) + } loc(#loc) +} loc(#loc) +#loc11 = loc("k.py":12:1) diff --git a/tests/golden/ir/ttir/crafted_odd_names.ttir b/tests/golden/ir/ttir/crafted_odd_names.ttir new file mode 100644 index 000000000..6c0207606 --- /dev/null +++ b/tests/golden/ir/ttir/crafted_odd_names.ttir @@ -0,0 +1,12 @@ +module { + tt.func public @k(%p: !tt.ptr, %n: i32) attributes {noinline = false} { + %c-1_i32 = arith.constant -1 : i32 + %1 = arith.addi %n, %c-1_i32 : i32 + %10 = arith.addi %1, %1 : i32 + %a.b = arith.addi %10, %1 : i32 + %a$c = arith.muli %a.b, %10 : i32 + %q = tt.addptr %p, %a$c : !tt.ptr, i32 + tt.store %q, %1 : !tt.ptr + tt.return + } +} diff --git a/tests/golden/ir/ttir/crafted_same_dest.ttir b/tests/golden/ir/ttir/crafted_same_dest.ttir new file mode 100644 index 000000000..fd0bf142b --- /dev/null +++ b/tests/golden/ir/ttir/crafted_same_dest.ttir @@ -0,0 +1,15 @@ +module { + tt.func public @k(%p: !tt.ptr, %a: i32, %b: i32) attributes {noinline = false} { + %c = arith.cmpi slt, %a, %b : i32 + cf.cond_br %c, ^bb1(%a : i32), ^bb1(%b : i32) + ^bb1(%x: i32): + %y = arith.addi %x, %a : i32 + cf.cond_br %c, ^bb2(%y, %x : i32, i32), ^bb3 + ^bb2(%u: i32, %v: i32): + %q = tt.addptr %p, %u : !tt.ptr, i32 + tt.store %q, %v : !tt.ptr + cf.br ^bb3 + ^bb3: + tt.return + } +} diff --git a/tests/golden/ir/ttir/crafted_symbols_strings.ttir b/tests/golden/ir/ttir/crafted_symbols_strings.ttir new file mode 100644 index 000000000..fa408b5eb --- /dev/null +++ b/tests/golden/ir/ttir/crafted_symbols_strings.ttir @@ -0,0 +1,16 @@ +module { + tt.func private @"f{%x} \22q\22 (a)"(%a: i32, %b: i32) -> (i32, i32) attributes {noinline = true} { + %s = arith.addi %a, %b : i32 + tt.return %s, %a : i32, i32 + } + tt.func public @k(%p: !tt.ptr, %n: i32) attributes {noinline = false} { + %r:2 = tt.call @"f{%x} \22q\22 (a)"(%n, %n) : (i32, i32) -> (i32, i32) + %y = tt.elementwise_inline_asm "{ mov.u32 $0, %tid.x; } // loc(\22x\22) }" {constraints = "=r,r", packed_element = 1 : i32, pure = true} %r#0 : i32 -> i32 + %c = arith.cmpi sgt, %y, %r#1 : i32 + tt.assert %c, "bad { } loc( \22 %x" : i1 + tt.print " p={%d} " {hex = false, isSigned = array} : %y : i32 + %q = tt.addptr %p, %y : !tt.ptr, i32 + tt.store %q, %r#1 : !tt.ptr + tt.return + } +} diff --git a/tests/golden/ir/ttir/crafted_unicode_strings.ttir b/tests/golden/ir/ttir/crafted_unicode_strings.ttir new file mode 100644 index 000000000..3e03ba265 --- /dev/null +++ b/tests/golden/ir/ttir/crafted_unicode_strings.ttir @@ -0,0 +1,10 @@ +#loc = loc("/tmp/\E5\86\85\E6\A0\B8/k.py":1:0) +#loc1 = loc("/tmp/\E5\86\85\E6\A0\B8/k.py":2:4) +module { + tt.func public @"\E6\A0\B8"(%p: !tt.ptr loc("\CF\80_ptr"(#loc)), %c: i1 loc("\E6\95\B0"(#loc))) attributes {noinline = false} { + %v = tt.load %p : !tt.ptr loc("\E5\80\BC"(#loc1)) + tt.assert %c, "\E9\94\99\E8\AF\AF \22quoted\22 \\ back" : i1 loc(#loc1) + tt.print "\CF\80=" {hex = false, isSigned = array} : %v : f32 loc(#loc1) + tt.return loc(#loc1) + } loc(#loc) +} loc(#loc) diff --git a/tests/golden/ir/ttir/golden_add_sm80.ttir b/tests/golden/ir/ttir/golden_add_sm80.ttir new file mode 100644 index 000000000..91b80235d --- /dev/null +++ b/tests/golden/ir/ttir/golden_add_sm80.ttir @@ -0,0 +1,51 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":108:0) +#loc15 = loc("x_ptr"(#loc)) +#loc16 = loc("y_ptr"(#loc)) +#loc17 = loc("out_ptr"(#loc)) +#loc18 = loc("n_elements"(#loc)) +module { + tt.func public @add_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %y_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("y_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %c1024_i32 = arith.constant 1024 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc19) + %offs = arith.muli %pid, %c1024_i32 : i32 loc(#loc20) + %offs_0 = tt.make_range {end = 1024 : i32, start = 0 : i32} : tensor<1024xi32> loc(#loc21) + %offs_1 = tt.splat %offs : i32 -> tensor<1024xi32> loc(#loc22) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<1024xi32> loc(#loc22) + %mask = tt.splat %n_elements : i32 -> tensor<1024xi32> loc(#loc23) + %mask_3 = arith.cmpi slt, %offs_2, %mask : tensor<1024xi32> loc(#loc23) + %x = tt.splat %x_ptr : !tt.ptr -> tensor<1024x!tt.ptr> loc(#loc24) + %x_4 = tt.addptr %x, %offs_2 : tensor<1024x!tt.ptr>, tensor<1024xi32> loc(#loc24) + %x_5 = tt.load %x_4, %mask_3 : tensor<1024x!tt.ptr> loc(#loc25) + %y = tt.splat %y_ptr : !tt.ptr -> tensor<1024x!tt.ptr> loc(#loc26) + %y_6 = tt.addptr %y, %offs_2 : tensor<1024x!tt.ptr>, tensor<1024xi32> loc(#loc26) + %y_7 = tt.load %y_6, %mask_3 : tensor<1024x!tt.ptr> loc(#loc27) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<1024x!tt.ptr> loc(#loc11) + %1 = tt.addptr %0, %offs_2 : tensor<1024x!tt.ptr>, tensor<1024xi32> loc(#loc11) + %2 = arith.addf %x_5, %y_7 : tensor<1024xf32> loc(#loc12) + tt.store %1, %2, %mask_3 : tensor<1024x!tt.ptr> loc(#loc13) + tt.return loc(#loc14) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":109:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":110:17) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":110:43) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":110:30) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":111:18) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":112:24) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":112:16) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":113:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":113:16) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":114:23) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":114:33) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":114:29) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":114:4) +#loc19 = loc("pid"(#loc2)) +#loc20 = loc("offs"(#loc3)) +#loc21 = loc("offs"(#loc4)) +#loc22 = loc("offs"(#loc5)) +#loc23 = loc("mask"(#loc6)) +#loc24 = loc("x"(#loc7)) +#loc25 = loc("x"(#loc8)) +#loc26 = loc("y"(#loc9)) +#loc27 = loc("y"(#loc10)) diff --git a/tests/golden/ir/ttir/golden_add_sm90.ttir b/tests/golden/ir/ttir/golden_add_sm90.ttir new file mode 100644 index 000000000..91b80235d --- /dev/null +++ b/tests/golden/ir/ttir/golden_add_sm90.ttir @@ -0,0 +1,51 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":108:0) +#loc15 = loc("x_ptr"(#loc)) +#loc16 = loc("y_ptr"(#loc)) +#loc17 = loc("out_ptr"(#loc)) +#loc18 = loc("n_elements"(#loc)) +module { + tt.func public @add_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %y_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("y_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %c1024_i32 = arith.constant 1024 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc19) + %offs = arith.muli %pid, %c1024_i32 : i32 loc(#loc20) + %offs_0 = tt.make_range {end = 1024 : i32, start = 0 : i32} : tensor<1024xi32> loc(#loc21) + %offs_1 = tt.splat %offs : i32 -> tensor<1024xi32> loc(#loc22) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<1024xi32> loc(#loc22) + %mask = tt.splat %n_elements : i32 -> tensor<1024xi32> loc(#loc23) + %mask_3 = arith.cmpi slt, %offs_2, %mask : tensor<1024xi32> loc(#loc23) + %x = tt.splat %x_ptr : !tt.ptr -> tensor<1024x!tt.ptr> loc(#loc24) + %x_4 = tt.addptr %x, %offs_2 : tensor<1024x!tt.ptr>, tensor<1024xi32> loc(#loc24) + %x_5 = tt.load %x_4, %mask_3 : tensor<1024x!tt.ptr> loc(#loc25) + %y = tt.splat %y_ptr : !tt.ptr -> tensor<1024x!tt.ptr> loc(#loc26) + %y_6 = tt.addptr %y, %offs_2 : tensor<1024x!tt.ptr>, tensor<1024xi32> loc(#loc26) + %y_7 = tt.load %y_6, %mask_3 : tensor<1024x!tt.ptr> loc(#loc27) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<1024x!tt.ptr> loc(#loc11) + %1 = tt.addptr %0, %offs_2 : tensor<1024x!tt.ptr>, tensor<1024xi32> loc(#loc11) + %2 = arith.addf %x_5, %y_7 : tensor<1024xf32> loc(#loc12) + tt.store %1, %2, %mask_3 : tensor<1024x!tt.ptr> loc(#loc13) + tt.return loc(#loc14) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":109:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":110:17) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":110:43) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":110:30) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":111:18) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":112:24) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":112:16) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":113:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":113:16) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":114:23) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":114:33) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":114:29) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":114:4) +#loc19 = loc("pid"(#loc2)) +#loc20 = loc("offs"(#loc3)) +#loc21 = loc("offs"(#loc4)) +#loc22 = loc("offs"(#loc5)) +#loc23 = loc("mask"(#loc6)) +#loc24 = loc("x"(#loc7)) +#loc25 = loc("x"(#loc8)) +#loc26 = loc("y"(#loc9)) +#loc27 = loc("y"(#loc10)) diff --git a/tests/golden/ir/ttir/golden_atomic_fmax_sm80.ttir b/tests/golden/ir/ttir/golden_atomic_fmax_sm80.ttir new file mode 100644 index 000000000..820632269 --- /dev/null +++ b/tests/golden/ir/ttir/golden_atomic_fmax_sm80.ttir @@ -0,0 +1,53 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":147:0) +#loc12 = loc("x_ptr"(#loc)) +#loc13 = loc("out_ptr"(#loc)) +#loc14 = loc("n_elements"(#loc)) +module { + tt.func public @atomic_fmax_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense : tensor<256xi1> loc(#loc1) + %cst_0 = arith.constant dense<0> : tensor<256xi32> loc(#loc1) + %cst_1 = arith.constant dense<31> : tensor<256xi32> loc(#loc1) + %v = arith.constant dense<0.000000e+00> : tensor<256xf32> loc(#loc15) + %c256_i32 = arith.constant 256 : i32 loc(#loc3) + %pid = tt.get_program_id x : i32 loc(#loc16) + %offs = arith.muli %pid, %c256_i32 : i32 loc(#loc17) + %offs_2 = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc18) + %offs_3 = tt.splat %offs : i32 -> tensor<256xi32> loc(#loc19) + %offs_4 = arith.addi %offs_3, %offs_2 : tensor<256xi32> loc(#loc19) + %mask = tt.splat %n_elements : i32 -> tensor<256xi32> loc(#loc20) + %mask_5 = arith.cmpi slt, %offs_4, %mask : tensor<256xi32> loc(#loc20) + %v_6 = tt.splat %x_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc21) + %v_7 = tt.addptr %v_6, %offs_4 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc21) + %v_8 = tt.load %v_7, %mask_5, %v : tensor<256x!tt.ptr> loc(#loc15) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc10) + %1 = tt.addptr %0, %offs_4 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc10) + %2 = tt.bitcast %v_8 : tensor<256xf32> -> tensor<256xi32> loc(#loc1) + %3 = tt.bitcast %1 : tensor<256x!tt.ptr> -> tensor<256x!tt.ptr> loc(#loc1) + %4 = arith.shrui %2, %cst_1 : tensor<256xi32> loc(#loc1) + %5 = arith.cmpi ne, %4, %cst_0 : tensor<256xi32> loc(#loc1) + %6 = arith.xori %5, %cst : tensor<256xi1> loc(#loc1) + %7 = arith.andi %mask_5, %6 : tensor<256xi1> loc(#loc1) + %8 = tt.atomic_rmw max, acq_rel, gpu, %3, %2, %7 : (tensor<256x!tt.ptr>, tensor<256xi32>, tensor<256xi1>) -> tensor<256xi32> loc(#loc1) + %9 = arith.andi %mask_5, %5 : tensor<256xi1> loc(#loc1) + %10 = tt.atomic_rmw umin, acq_rel, gpu, %3, %2, %9 : (tensor<256x!tt.ptr>, tensor<256xi32>, tensor<256xi1>) -> tensor<256xi32> loc(#loc1) + tt.return loc(#loc11) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":155:34) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":154:16) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":151:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":152:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":152:38) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":152:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":153:18) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":154:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":155:28) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":155:4) +#loc15 = loc("v"(#loc2)) +#loc16 = loc("pid"(#loc4)) +#loc17 = loc("offs"(#loc5)) +#loc18 = loc("offs"(#loc6)) +#loc19 = loc("offs"(#loc7)) +#loc20 = loc("mask"(#loc8)) +#loc21 = loc("v"(#loc9)) diff --git a/tests/golden/ir/ttir/golden_atomic_fmax_sm90.ttir b/tests/golden/ir/ttir/golden_atomic_fmax_sm90.ttir new file mode 100644 index 000000000..820632269 --- /dev/null +++ b/tests/golden/ir/ttir/golden_atomic_fmax_sm90.ttir @@ -0,0 +1,53 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":147:0) +#loc12 = loc("x_ptr"(#loc)) +#loc13 = loc("out_ptr"(#loc)) +#loc14 = loc("n_elements"(#loc)) +module { + tt.func public @atomic_fmax_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense : tensor<256xi1> loc(#loc1) + %cst_0 = arith.constant dense<0> : tensor<256xi32> loc(#loc1) + %cst_1 = arith.constant dense<31> : tensor<256xi32> loc(#loc1) + %v = arith.constant dense<0.000000e+00> : tensor<256xf32> loc(#loc15) + %c256_i32 = arith.constant 256 : i32 loc(#loc3) + %pid = tt.get_program_id x : i32 loc(#loc16) + %offs = arith.muli %pid, %c256_i32 : i32 loc(#loc17) + %offs_2 = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc18) + %offs_3 = tt.splat %offs : i32 -> tensor<256xi32> loc(#loc19) + %offs_4 = arith.addi %offs_3, %offs_2 : tensor<256xi32> loc(#loc19) + %mask = tt.splat %n_elements : i32 -> tensor<256xi32> loc(#loc20) + %mask_5 = arith.cmpi slt, %offs_4, %mask : tensor<256xi32> loc(#loc20) + %v_6 = tt.splat %x_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc21) + %v_7 = tt.addptr %v_6, %offs_4 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc21) + %v_8 = tt.load %v_7, %mask_5, %v : tensor<256x!tt.ptr> loc(#loc15) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc10) + %1 = tt.addptr %0, %offs_4 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc10) + %2 = tt.bitcast %v_8 : tensor<256xf32> -> tensor<256xi32> loc(#loc1) + %3 = tt.bitcast %1 : tensor<256x!tt.ptr> -> tensor<256x!tt.ptr> loc(#loc1) + %4 = arith.shrui %2, %cst_1 : tensor<256xi32> loc(#loc1) + %5 = arith.cmpi ne, %4, %cst_0 : tensor<256xi32> loc(#loc1) + %6 = arith.xori %5, %cst : tensor<256xi1> loc(#loc1) + %7 = arith.andi %mask_5, %6 : tensor<256xi1> loc(#loc1) + %8 = tt.atomic_rmw max, acq_rel, gpu, %3, %2, %7 : (tensor<256x!tt.ptr>, tensor<256xi32>, tensor<256xi1>) -> tensor<256xi32> loc(#loc1) + %9 = arith.andi %mask_5, %5 : tensor<256xi1> loc(#loc1) + %10 = tt.atomic_rmw umin, acq_rel, gpu, %3, %2, %9 : (tensor<256x!tt.ptr>, tensor<256xi32>, tensor<256xi1>) -> tensor<256xi32> loc(#loc1) + tt.return loc(#loc11) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":155:34) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":154:16) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":151:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":152:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":152:38) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":152:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":153:18) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":154:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":155:28) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":155:4) +#loc15 = loc("v"(#loc2)) +#loc16 = loc("pid"(#loc4)) +#loc17 = loc("offs"(#loc5)) +#loc18 = loc("offs"(#loc6)) +#loc19 = loc("offs"(#loc7)) +#loc20 = loc("mask"(#loc8)) +#loc21 = loc("v"(#loc9)) diff --git a/tests/golden/ir/ttir/golden_atomic_sm80.ttir b/tests/golden/ir/ttir/golden_atomic_sm80.ttir new file mode 100644 index 000000000..220c65434 --- /dev/null +++ b/tests/golden/ir/ttir/golden_atomic_sm80.ttir @@ -0,0 +1,48 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":133:0) +#loc14 = loc("x_ptr"(#loc)) +#loc15 = loc("out_ptr"(#loc)) +#loc16 = loc("n_elements"(#loc)) +module { + tt.func public @atomic_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %old = arith.constant dense : tensor<256xi1> loc(#loc17) + %v = arith.constant dense<0.000000e+00> : tensor<256xf32> loc(#loc18) + %c256_i32 = arith.constant 256 : i32 loc(#loc3) + %pid = tt.get_program_id x : i32 loc(#loc19) + %offs = arith.muli %pid, %c256_i32 : i32 loc(#loc20) + %offs_0 = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc21) + %offs_1 = tt.splat %offs : i32 -> tensor<256xi32> loc(#loc22) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<256xi32> loc(#loc22) + %mask = tt.splat %n_elements : i32 -> tensor<256xi32> loc(#loc23) + %mask_3 = arith.cmpi slt, %offs_2, %mask : tensor<256xi32> loc(#loc23) + %v_4 = tt.splat %x_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc24) + %v_5 = tt.addptr %v_4, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc24) + %v_6 = tt.load %v_5, %mask_3, %v : tensor<256x!tt.ptr> loc(#loc18) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc10) + %1 = tt.addptr %0, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc10) + %2 = tt.atomic_rmw fadd, acq_rel, gpu, %1, %v_6, %mask_3 : (tensor<256x!tt.ptr>, tensor<256xf32>, tensor<256xi1>) -> tensor<256xf32> loc(#loc11) + %old_7 = tt.atomic_rmw exch, acq_rel, gpu, %1, %v_6, %old : (tensor<256x!tt.ptr>, tensor<256xf32>, tensor<256xi1>) -> tensor<256xf32> loc(#loc17) + tt.store %v_5, %old_7, %mask_3 : tensor<256x!tt.ptr> loc(#loc12) + tt.return loc(#loc13) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":142:41) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:16) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:38) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":139:18) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":141:28) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":141:34) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":143:27) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":143:4) +#loc17 = loc("old"(#loc1)) +#loc18 = loc("v"(#loc2)) +#loc19 = loc("pid"(#loc4)) +#loc20 = loc("offs"(#loc5)) +#loc21 = loc("offs"(#loc6)) +#loc22 = loc("offs"(#loc7)) +#loc23 = loc("mask"(#loc8)) +#loc24 = loc("v"(#loc9)) diff --git a/tests/golden/ir/ttir/golden_atomic_sm90.ttir b/tests/golden/ir/ttir/golden_atomic_sm90.ttir new file mode 100644 index 000000000..220c65434 --- /dev/null +++ b/tests/golden/ir/ttir/golden_atomic_sm90.ttir @@ -0,0 +1,48 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":133:0) +#loc14 = loc("x_ptr"(#loc)) +#loc15 = loc("out_ptr"(#loc)) +#loc16 = loc("n_elements"(#loc)) +module { + tt.func public @atomic_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %old = arith.constant dense : tensor<256xi1> loc(#loc17) + %v = arith.constant dense<0.000000e+00> : tensor<256xf32> loc(#loc18) + %c256_i32 = arith.constant 256 : i32 loc(#loc3) + %pid = tt.get_program_id x : i32 loc(#loc19) + %offs = arith.muli %pid, %c256_i32 : i32 loc(#loc20) + %offs_0 = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc21) + %offs_1 = tt.splat %offs : i32 -> tensor<256xi32> loc(#loc22) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<256xi32> loc(#loc22) + %mask = tt.splat %n_elements : i32 -> tensor<256xi32> loc(#loc23) + %mask_3 = arith.cmpi slt, %offs_2, %mask : tensor<256xi32> loc(#loc23) + %v_4 = tt.splat %x_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc24) + %v_5 = tt.addptr %v_4, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc24) + %v_6 = tt.load %v_5, %mask_3, %v : tensor<256x!tt.ptr> loc(#loc18) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc10) + %1 = tt.addptr %0, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc10) + %2 = tt.atomic_rmw fadd, acq_rel, gpu, %1, %v_6, %mask_3 : (tensor<256x!tt.ptr>, tensor<256xf32>, tensor<256xi1>) -> tensor<256xf32> loc(#loc11) + %old_7 = tt.atomic_rmw exch, acq_rel, gpu, %1, %v_6, %old : (tensor<256x!tt.ptr>, tensor<256xf32>, tensor<256xi1>) -> tensor<256xf32> loc(#loc17) + tt.store %v_5, %old_7, %mask_3 : tensor<256x!tt.ptr> loc(#loc12) + tt.return loc(#loc13) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":142:41) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:16) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:38) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":139:18) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":141:28) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":141:34) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":143:27) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":143:4) +#loc17 = loc("old"(#loc1)) +#loc18 = loc("v"(#loc2)) +#loc19 = loc("pid"(#loc4)) +#loc20 = loc("offs"(#loc5)) +#loc21 = loc("offs"(#loc6)) +#loc22 = loc("offs"(#loc7)) +#loc23 = loc("mask"(#loc8)) +#loc24 = loc("v"(#loc9)) diff --git a/tests/golden/ir/ttir/golden_cas_sm80.ttir b/tests/golden/ir/ttir/golden_cas_sm80.ttir new file mode 100644 index 000000000..592ac9b57 --- /dev/null +++ b/tests/golden/ir/ttir/golden_cas_sm80.ttir @@ -0,0 +1,16 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":146:0) +#loc4 = loc("lock_ptr"(#loc)) +#loc5 = loc("out_ptr"(#loc)) +module { + tt.func public @cas_kernel(%lock_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("lock_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc))) attributes {noinline = false} { + %old = arith.constant 0 : i32 loc(#loc6) + %old_0 = arith.constant 1 : i32 loc(#loc6) + %old_1 = tt.atomic_cas acq_rel, gpu, %lock_ptr, %old, %old_0 : (!tt.ptr, i32, i32) -> i32 loc(#loc6) + tt.store %out_ptr, %old_1 : !tt.ptr loc(#loc2) + tt.return loc(#loc3) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":148:37) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":149:22) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":149:4) +#loc6 = loc("old"(#loc1)) diff --git a/tests/golden/ir/ttir/golden_cas_sm90.ttir b/tests/golden/ir/ttir/golden_cas_sm90.ttir new file mode 100644 index 000000000..592ac9b57 --- /dev/null +++ b/tests/golden/ir/ttir/golden_cas_sm90.ttir @@ -0,0 +1,16 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":146:0) +#loc4 = loc("lock_ptr"(#loc)) +#loc5 = loc("out_ptr"(#loc)) +module { + tt.func public @cas_kernel(%lock_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("lock_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc))) attributes {noinline = false} { + %old = arith.constant 0 : i32 loc(#loc6) + %old_0 = arith.constant 1 : i32 loc(#loc6) + %old_1 = tt.atomic_cas acq_rel, gpu, %lock_ptr, %old, %old_0 : (!tt.ptr, i32, i32) -> i32 loc(#loc6) + tt.store %out_ptr, %old_1 : !tt.ptr loc(#loc2) + tt.return loc(#loc3) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":148:37) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":149:22) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":149:4) +#loc6 = loc("old"(#loc1)) diff --git a/tests/golden/ir/ttir/golden_early_return_loaded_sm80.ttir b/tests/golden/ir/ttir/golden_early_return_loaded_sm80.ttir new file mode 100644 index 000000000..9b38c5066 --- /dev/null +++ b/tests/golden/ir/ttir/golden_early_return_loaded_sm80.ttir @@ -0,0 +1,56 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":342:0) +#loc16 = loc("x_ptr"(#loc)) +#loc17 = loc("idx_ptr"(#loc)) +#loc18 = loc("out_ptr"(#loc)) +#loc19 = loc("n"(#loc)) +module { + tt.func public @early_return_loaded_kernel(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %idx_ptr: !tt.ptr loc("idx_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %c-1_i32 = arith.constant -1 : i32 loc(#loc2) + %pid = tt.get_program_id x : i32 loc(#loc20) + %y = tt.addptr %idx_ptr, %pid : !tt.ptr, i32 loc(#loc21) + %y_0 = tt.load %y : !tt.ptr loc(#loc22) + %0 = arith.cmpi eq, %y_0, %c-1_i32 : i32 loc(#loc2) + cf.cond_br %0, ^bb1, ^bb2 loc(#loc2) + ^bb1: // pred: ^bb0 + tt.return loc(#loc6) + ^bb2: // pred: ^bb0 + %offs = arith.muli %pid, %c64_i32 : i32 loc(#loc23) + %offs_1 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc24) + %offs_2 = tt.splat %offs : i32 -> tensor<64xi32> loc(#loc25) + %offs_3 = arith.addi %offs_2, %offs_1 : tensor<64xi32> loc(#loc25) + %m = tt.splat %n : i32 -> tensor<64xi32> loc(#loc26) + %m_4 = arith.cmpi slt, %offs_3, %m : tensor<64xi32> loc(#loc26) + %v = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc27) + %v_5 = tt.addptr %v, %offs_3 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc27) + %v_6 = tt.load %v_5, %m_4 : tensor<64x!tt.ptr> loc(#loc28) + %1 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc13) + %2 = tt.addptr %1, %offs_3 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc13) + tt.store %2, %v_6, %m_4 : tensor<64x!tt.ptr> loc(#loc14) + tt.return loc(#loc15) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":345:12) +#loc3 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":343:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":344:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":344:16) +#loc6 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":346:8) +#loc7 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":347:17) +#loc8 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":347:38) +#loc9 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":347:25) +#loc10 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":348:15) +#loc11 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":349:24) +#loc12 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":349:16) +#loc13 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":350:23) +#loc14 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":350:29) +#loc15 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":350:4) +#loc20 = loc("pid"(#loc3)) +#loc21 = loc("y"(#loc4)) +#loc22 = loc("y"(#loc5)) +#loc23 = loc("offs"(#loc7)) +#loc24 = loc("offs"(#loc8)) +#loc25 = loc("offs"(#loc9)) +#loc26 = loc("m"(#loc10)) +#loc27 = loc("v"(#loc11)) +#loc28 = loc("v"(#loc12)) diff --git a/tests/golden/ir/ttir/golden_early_return_pid_sm80.ttir b/tests/golden/ir/ttir/golden_early_return_pid_sm80.ttir new file mode 100644 index 000000000..37859e371 --- /dev/null +++ b/tests/golden/ir/ttir/golden_early_return_pid_sm80.ttir @@ -0,0 +1,48 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":331:0) +#loc14 = loc("x_ptr"(#loc)) +#loc15 = loc("out_ptr"(#loc)) +#loc16 = loc("n"(#loc)) +#loc17 = loc("T"(#loc)) +module { + tt.func public @early_return_pid_kernel(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc)), %T: i32 loc("T"(#loc))) attributes {noinline = false} { + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc18) + %0 = arith.muli %pid, %c64_i32 : i32 loc(#loc3) + %1 = arith.cmpi sge, %0, %T : i32 loc(#loc4) + cf.cond_br %1, ^bb1, ^bb2 loc(#loc4) + ^bb1: // pred: ^bb0 + tt.return loc(#loc5) + ^bb2: // pred: ^bb0 + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc19) + %offs_0 = tt.splat %0 : i32 -> tensor<64xi32> loc(#loc20) + %offs_1 = arith.addi %offs_0, %offs : tensor<64xi32> loc(#loc20) + %m = tt.splat %n : i32 -> tensor<64xi32> loc(#loc21) + %m_2 = arith.cmpi slt, %offs_1, %m : tensor<64xi32> loc(#loc21) + %v = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc22) + %v_3 = tt.addptr %v, %offs_1 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc22) + %v_4 = tt.load %v_3, %m_2 : tensor<64x!tt.ptr> loc(#loc23) + %2 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc11) + %3 = tt.addptr %2, %offs_1 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc11) + tt.store %3, %v_4, %m_2 : tensor<64x!tt.ptr> loc(#loc12) + tt.return loc(#loc13) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":332:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":333:13) +#loc4 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":333:22) +#loc5 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":334:8) +#loc6 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":335:38) +#loc7 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":335:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":336:15) +#loc9 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":337:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":337:16) +#loc11 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":338:23) +#loc12 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":338:29) +#loc13 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":338:4) +#loc18 = loc("pid"(#loc2)) +#loc19 = loc("offs"(#loc6)) +#loc20 = loc("offs"(#loc7)) +#loc21 = loc("m"(#loc8)) +#loc22 = loc("v"(#loc9)) +#loc23 = loc("v"(#loc10)) diff --git a/tests/golden/ir/ttir/golden_gather_sm80.ttir b/tests/golden/ir/ttir/golden_gather_sm80.ttir new file mode 100644 index 000000000..ecaa02ff6 --- /dev/null +++ b/tests/golden/ir/ttir/golden_gather_sm80.ttir @@ -0,0 +1,51 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":133:0) +#loc14 = loc("idx_ptr"(#loc)) +#loc15 = loc("src_ptr"(#loc)) +#loc16 = loc("out_ptr"(#loc)) +#loc17 = loc("n_elements"(#loc)) +module { + tt.func public @gather_kernel(%idx_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("idx_ptr"(#loc)), %src_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("src_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %vals = arith.constant dense<0.000000e+00> : tensor<256xf32> loc(#loc18) + %idx = arith.constant dense<0> : tensor<256xi32> loc(#loc19) + %c256_i32 = arith.constant 256 : i32 loc(#loc3) + %pid = tt.get_program_id x : i32 loc(#loc20) + %offs = arith.muli %pid, %c256_i32 : i32 loc(#loc21) + %offs_0 = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc22) + %offs_1 = tt.splat %offs : i32 -> tensor<256xi32> loc(#loc23) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<256xi32> loc(#loc23) + %mask = tt.splat %n_elements : i32 -> tensor<256xi32> loc(#loc24) + %mask_3 = arith.cmpi slt, %offs_2, %mask : tensor<256xi32> loc(#loc24) + %idx_4 = tt.splat %idx_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc25) + %idx_5 = tt.addptr %idx_4, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc25) + %idx_6 = tt.load %idx_5, %mask_3, %idx : tensor<256x!tt.ptr> loc(#loc19) + %vals_7 = tt.splat %src_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc26) + %vals_8 = tt.addptr %vals_7, %idx_6 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc26) + %vals_9 = tt.load %vals_8, %mask_3, %vals : tensor<256x!tt.ptr> loc(#loc18) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc11) + %1 = tt.addptr %0, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc11) + tt.store %1, %vals_9, %mask_3 : tensor<256x!tt.ptr> loc(#loc12) + tt.return loc(#loc13) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:19) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":139:18) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":136:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:38) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:18) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":139:28) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:29) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":141:23) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":141:29) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":141:4) +#loc18 = loc("vals"(#loc1)) +#loc19 = loc("idx"(#loc2)) +#loc20 = loc("pid"(#loc4)) +#loc21 = loc("offs"(#loc5)) +#loc22 = loc("offs"(#loc6)) +#loc23 = loc("offs"(#loc7)) +#loc24 = loc("mask"(#loc8)) +#loc25 = loc("idx"(#loc9)) +#loc26 = loc("vals"(#loc10)) diff --git a/tests/golden/ir/ttir/golden_gather_sm90.ttir b/tests/golden/ir/ttir/golden_gather_sm90.ttir new file mode 100644 index 000000000..ecaa02ff6 --- /dev/null +++ b/tests/golden/ir/ttir/golden_gather_sm90.ttir @@ -0,0 +1,51 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":133:0) +#loc14 = loc("idx_ptr"(#loc)) +#loc15 = loc("src_ptr"(#loc)) +#loc16 = loc("out_ptr"(#loc)) +#loc17 = loc("n_elements"(#loc)) +module { + tt.func public @gather_kernel(%idx_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("idx_ptr"(#loc)), %src_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("src_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %vals = arith.constant dense<0.000000e+00> : tensor<256xf32> loc(#loc18) + %idx = arith.constant dense<0> : tensor<256xi32> loc(#loc19) + %c256_i32 = arith.constant 256 : i32 loc(#loc3) + %pid = tt.get_program_id x : i32 loc(#loc20) + %offs = arith.muli %pid, %c256_i32 : i32 loc(#loc21) + %offs_0 = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc22) + %offs_1 = tt.splat %offs : i32 -> tensor<256xi32> loc(#loc23) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<256xi32> loc(#loc23) + %mask = tt.splat %n_elements : i32 -> tensor<256xi32> loc(#loc24) + %mask_3 = arith.cmpi slt, %offs_2, %mask : tensor<256xi32> loc(#loc24) + %idx_4 = tt.splat %idx_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc25) + %idx_5 = tt.addptr %idx_4, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc25) + %idx_6 = tt.load %idx_5, %mask_3, %idx : tensor<256x!tt.ptr> loc(#loc19) + %vals_7 = tt.splat %src_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc26) + %vals_8 = tt.addptr %vals_7, %idx_6 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc26) + %vals_9 = tt.load %vals_8, %mask_3, %vals : tensor<256x!tt.ptr> loc(#loc18) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc11) + %1 = tt.addptr %0, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc11) + tt.store %1, %vals_9, %mask_3 : tensor<256x!tt.ptr> loc(#loc12) + tt.return loc(#loc13) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:19) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":139:18) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":136:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:38) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:18) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":139:28) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:29) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":141:23) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":141:29) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":141:4) +#loc18 = loc("vals"(#loc1)) +#loc19 = loc("idx"(#loc2)) +#loc20 = loc("pid"(#loc4)) +#loc21 = loc("offs"(#loc5)) +#loc22 = loc("offs"(#loc6)) +#loc23 = loc("offs"(#loc7)) +#loc24 = loc("mask"(#loc8)) +#loc25 = loc("idx"(#loc9)) +#loc26 = loc("vals"(#loc10)) diff --git a/tests/golden/ir/ttir/golden_grid_stride_sm80.ttir b/tests/golden/ir/ttir/golden_grid_stride_sm80.ttir new file mode 100644 index 000000000..d894065f6 --- /dev/null +++ b/tests/golden/ir/ttir/golden_grid_stride_sm80.ttir @@ -0,0 +1,41 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":403:0) +#loc12 = loc("x_ptr"(#loc)) +#loc13 = loc("out_ptr"(#loc)) +#loc14 = loc("n_rows"(#loc)) +#loc15 = loc("stride"(#loc)) +module { + tt.func public @grid_stride_kernel(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n_rows: i32 loc("n_rows"(#loc)), %stride: i32 loc("stride"(#loc))) attributes {noinline = false} { + %c4_i32 = arith.constant 4 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc16) + %cols = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc17) + scf.for %row = %pid to %n_rows step %c4_i32 : i32 { + %v = arith.muli %row, %stride : i32 loc(#loc18) + %v_0 = tt.addptr %x_ptr, %v : !tt.ptr, i32 loc(#loc19) + %v_1 = tt.splat %v_0 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc20) + %v_2 = tt.addptr %v_1, %cols : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc20) + %v_3 = tt.load %v_2 : tensor<64x!tt.ptr> loc(#loc21) + %0 = tt.addptr %out_ptr, %v : !tt.ptr, i32 loc(#loc8) + %1 = tt.splat %0 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc9) + %2 = tt.addptr %1, %cols : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc9) + tt.store %2, %v_3 : tensor<64x!tt.ptr> loc(#loc10) + } loc(#loc1) + tt.return loc(#loc11) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":408:34) +#loc2 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":406:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":407:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":409:34) +#loc5 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":409:28) +#loc6 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":409:43) +#loc7 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":409:20) +#loc8 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":410:27) +#loc9 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":410:42) +#loc10 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":410:48) +#loc11 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":408:4) +#loc16 = loc("pid"(#loc2)) +#loc17 = loc("cols"(#loc3)) +#loc18 = loc("v"(#loc4)) +#loc19 = loc("v"(#loc5)) +#loc20 = loc("v"(#loc6)) +#loc21 = loc("v"(#loc7)) diff --git a/tests/golden/ir/ttir/golden_guard_then_loop_sm80.ttir b/tests/golden/ir/ttir/golden_guard_then_loop_sm80.ttir new file mode 100644 index 000000000..d8fad4ad2 --- /dev/null +++ b/tests/golden/ir/ttir/golden_guard_then_loop_sm80.ttir @@ -0,0 +1,55 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":371:0) +#loc16 = loc("x_ptr"(#loc)) +#loc17 = loc("out_ptr"(#loc)) +#loc18 = loc("n"(#loc)) +#loc19 = loc("T"(#loc)) +module { + tt.func public @guard_then_loop_kernel(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc)), %T: i32 loc("T"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc20) + %0 = arith.muli %pid, %c64_i32 : i32 loc(#loc3) + %1 = arith.cmpi sge, %0, %T : i32 loc(#loc4) + cf.cond_br %1, ^bb1, ^bb2 loc(#loc4) + ^bb1: // pred: ^bb0 + tt.return loc(#loc5) + ^bb2: // pred: ^bb0 + scf.for %k = %c0_i32 to %n step %c1_i32 : i32 { + %offs = arith.muli %k, %T : i32 loc(#loc21) + %offs_0 = arith.addi %0, %offs : i32 loc(#loc22) + %offs_1 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc23) + %offs_2 = tt.splat %offs_0 : i32 -> tensor<64xi32> loc(#loc24) + %offs_3 = arith.addi %offs_2, %offs_1 : tensor<64xi32> loc(#loc24) + %v = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc25) + %v_4 = tt.addptr %v, %offs_3 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc25) + %v_5 = tt.load %v_4 : tensor<64x!tt.ptr> loc(#loc26) + %2 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc13) + %3 = tt.addptr %2, %offs_3 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc13) + tt.store %3, %v_5 : tensor<64x!tt.ptr> loc(#loc14) + } loc(#loc6) + tt.return loc(#loc15) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":372:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":373:13) +#loc4 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":373:22) +#loc5 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":374:8) +#loc6 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":375:22) +#loc7 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":376:33) +#loc8 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":376:29) +#loc9 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":376:50) +#loc10 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":376:37) +#loc11 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":377:28) +#loc12 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":377:20) +#loc13 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":378:27) +#loc14 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":378:33) +#loc15 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":375:4) +#loc20 = loc("pid"(#loc2)) +#loc21 = loc("offs"(#loc7)) +#loc22 = loc("offs"(#loc8)) +#loc23 = loc("offs"(#loc9)) +#loc24 = loc("offs"(#loc10)) +#loc25 = loc("v"(#loc11)) +#loc26 = loc("v"(#loc12)) diff --git a/tests/golden/ir/ttir/golden_if_else_load_sm80.ttir b/tests/golden/ir/ttir/golden_if_else_load_sm80.ttir new file mode 100644 index 000000000..81028b9c4 --- /dev/null +++ b/tests/golden/ir/ttir/golden_if_else_load_sm80.ttir @@ -0,0 +1,60 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":192:0) +#loc16 = loc("x_ptr"(#loc)) +#loc17 = loc("y_ptr"(#loc)) +#loc18 = loc("out_ptr"(#loc)) +#loc19 = loc("n_elements"(#loc)) +module { + tt.func public @if_else_load_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %y_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("y_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c256_i32 = arith.constant 256 : i32 loc(#loc2) + %pid = tt.get_program_id x : i32 loc(#loc20) + %offs = arith.muli %pid, %c256_i32 : i32 loc(#loc21) + %offs_0 = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc22) + %offs_1 = tt.splat %offs : i32 -> tensor<256xi32> loc(#loc23) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<256xi32> loc(#loc23) + %mask = tt.splat %n_elements : i32 -> tensor<256xi32> loc(#loc24) + %mask_3 = arith.cmpi slt, %offs_2, %mask : tensor<256xi32> loc(#loc24) + %0 = arith.cmpi eq, %pid, %c0_i32 : i32 loc(#loc1) + %1 = scf.if %0 -> (tensor<256xf32>) { + %v = tt.splat %x_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc25) + %v_4 = tt.addptr %v, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc25) + %v_5 = tt.load %v_4, %mask_3 : tensor<256x!tt.ptr> loc(#loc29) + scf.yield %v_5 : tensor<256xf32> loc(#loc29) + } else { + %v = tt.splat %y_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc27) + %v_4 = tt.addptr %v, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc27) + %v_5 = tt.load %v_4, %mask_3 : tensor<256x!tt.ptr> loc(#loc30) + scf.yield %v_5 : tensor<256xf32> loc(#loc28) + } loc(#loc8) + %2 = tt.splat %out_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc13) + %3 = tt.addptr %2, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc13) + tt.store %3, %1, %mask_3 : tensor<256x!tt.ptr> loc(#loc14) + tt.return loc(#loc15) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":199:14) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":196:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":197:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":197:38) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":197:25) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":198:18) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":199:7) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":200:28) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":200:20) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":202:28) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":202:20) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":203:23) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":203:29) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":203:4) +#loc20 = loc("pid"(#loc3)) +#loc21 = loc("offs"(#loc4)) +#loc22 = loc("offs"(#loc5)) +#loc23 = loc("offs"(#loc6)) +#loc24 = loc("mask"(#loc7)) +#loc25 = loc("v"(#loc9)) +#loc26 = loc("v"(#loc10)) +#loc27 = loc("v"(#loc11)) +#loc28 = loc("v"(#loc12)) +#loc29 = loc("v"(#loc26)) +#loc30 = loc("v"(#loc28)) diff --git a/tests/golden/ir/ttir/golden_if_else_load_sm90.ttir b/tests/golden/ir/ttir/golden_if_else_load_sm90.ttir new file mode 100644 index 000000000..81028b9c4 --- /dev/null +++ b/tests/golden/ir/ttir/golden_if_else_load_sm90.ttir @@ -0,0 +1,60 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":192:0) +#loc16 = loc("x_ptr"(#loc)) +#loc17 = loc("y_ptr"(#loc)) +#loc18 = loc("out_ptr"(#loc)) +#loc19 = loc("n_elements"(#loc)) +module { + tt.func public @if_else_load_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %y_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("y_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c256_i32 = arith.constant 256 : i32 loc(#loc2) + %pid = tt.get_program_id x : i32 loc(#loc20) + %offs = arith.muli %pid, %c256_i32 : i32 loc(#loc21) + %offs_0 = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc22) + %offs_1 = tt.splat %offs : i32 -> tensor<256xi32> loc(#loc23) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<256xi32> loc(#loc23) + %mask = tt.splat %n_elements : i32 -> tensor<256xi32> loc(#loc24) + %mask_3 = arith.cmpi slt, %offs_2, %mask : tensor<256xi32> loc(#loc24) + %0 = arith.cmpi eq, %pid, %c0_i32 : i32 loc(#loc1) + %1 = scf.if %0 -> (tensor<256xf32>) { + %v = tt.splat %x_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc25) + %v_4 = tt.addptr %v, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc25) + %v_5 = tt.load %v_4, %mask_3 : tensor<256x!tt.ptr> loc(#loc29) + scf.yield %v_5 : tensor<256xf32> loc(#loc29) + } else { + %v = tt.splat %y_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc27) + %v_4 = tt.addptr %v, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc27) + %v_5 = tt.load %v_4, %mask_3 : tensor<256x!tt.ptr> loc(#loc30) + scf.yield %v_5 : tensor<256xf32> loc(#loc28) + } loc(#loc8) + %2 = tt.splat %out_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc13) + %3 = tt.addptr %2, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc13) + tt.store %3, %1, %mask_3 : tensor<256x!tt.ptr> loc(#loc14) + tt.return loc(#loc15) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":199:14) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":196:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":197:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":197:38) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":197:25) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":198:18) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":199:7) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":200:28) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":200:20) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":202:28) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":202:20) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":203:23) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":203:29) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":203:4) +#loc20 = loc("pid"(#loc3)) +#loc21 = loc("offs"(#loc4)) +#loc22 = loc("offs"(#loc5)) +#loc23 = loc("offs"(#loc6)) +#loc24 = loc("mask"(#loc7)) +#loc25 = loc("v"(#loc9)) +#loc26 = loc("v"(#loc10)) +#loc27 = loc("v"(#loc11)) +#loc28 = loc("v"(#loc12)) +#loc29 = loc("v"(#loc26)) +#loc30 = loc("v"(#loc28)) diff --git a/tests/golden/ir/ttir/golden_if_else_offset_sm80.ttir b/tests/golden/ir/ttir/golden_if_else_offset_sm80.ttir new file mode 100644 index 000000000..e0b1adc13 --- /dev/null +++ b/tests/golden/ir/ttir/golden_if_else_offset_sm80.ttir @@ -0,0 +1,40 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":178:0) +#loc13 = loc("x_ptr"(#loc)) +#loc14 = loc("out_ptr"(#loc)) +#loc15 = loc("n_elements"(#loc)) +#loc21 = loc("base"(#loc15)) +module { + tt.func public @if_else_offset_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %base: i32 {tt.divisibility = 16 : i32} loc("base"(#loc15))) attributes {noinline = false} { + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc16) + %offs = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc17) + %0 = arith.cmpi eq, %pid, %c0_i32 : i32 loc(#loc4) + %1 = arith.select %0, %c0_i32, %base : i32 loc(#loc5) + %v = tt.addptr %x_ptr, %1 : !tt.ptr, i32 loc(#loc18) + %v_0 = tt.splat %v : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc19) + %v_1 = tt.addptr %v_0, %offs : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc19) + %v_2 = tt.load %v_1 : tensor<256x!tt.ptr> loc(#loc20) + %2 = tt.addptr %out_ptr, %1 : !tt.ptr, i32 loc(#loc9) + %3 = tt.splat %2 : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc10) + %4 = tt.addptr %3, %offs : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc10) + tt.store %4, %v_2 : tensor<256x!tt.ptr> loc(#loc11) + tt.return loc(#loc12) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":181:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":182:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":183:14) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":183:7) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":187:24) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":187:31) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":187:16) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":188:23) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":188:30) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":188:36) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":188:4) +#loc16 = loc("pid"(#loc2)) +#loc17 = loc("offs"(#loc3)) +#loc18 = loc("v"(#loc6)) +#loc19 = loc("v"(#loc7)) +#loc20 = loc("v"(#loc8)) diff --git a/tests/golden/ir/ttir/golden_if_else_offset_sm90.ttir b/tests/golden/ir/ttir/golden_if_else_offset_sm90.ttir new file mode 100644 index 000000000..e0b1adc13 --- /dev/null +++ b/tests/golden/ir/ttir/golden_if_else_offset_sm90.ttir @@ -0,0 +1,40 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":178:0) +#loc13 = loc("x_ptr"(#loc)) +#loc14 = loc("out_ptr"(#loc)) +#loc15 = loc("n_elements"(#loc)) +#loc21 = loc("base"(#loc15)) +module { + tt.func public @if_else_offset_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %base: i32 {tt.divisibility = 16 : i32} loc("base"(#loc15))) attributes {noinline = false} { + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc16) + %offs = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc17) + %0 = arith.cmpi eq, %pid, %c0_i32 : i32 loc(#loc4) + %1 = arith.select %0, %c0_i32, %base : i32 loc(#loc5) + %v = tt.addptr %x_ptr, %1 : !tt.ptr, i32 loc(#loc18) + %v_0 = tt.splat %v : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc19) + %v_1 = tt.addptr %v_0, %offs : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc19) + %v_2 = tt.load %v_1 : tensor<256x!tt.ptr> loc(#loc20) + %2 = tt.addptr %out_ptr, %1 : !tt.ptr, i32 loc(#loc9) + %3 = tt.splat %2 : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc10) + %4 = tt.addptr %3, %offs : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc10) + tt.store %4, %v_2 : tensor<256x!tt.ptr> loc(#loc11) + tt.return loc(#loc12) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":181:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":182:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":183:14) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":183:7) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":187:24) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":187:31) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":187:16) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":188:23) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":188:30) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":188:36) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":188:4) +#loc16 = loc("pid"(#loc2)) +#loc17 = loc("offs"(#loc3)) +#loc18 = loc("v"(#loc6)) +#loc19 = loc("v"(#loc7)) +#loc20 = loc("v"(#loc8)) diff --git a/tests/golden/ir/ttir/golden_loop_under_if_sm80.ttir b/tests/golden/ir/ttir/golden_loop_under_if_sm80.ttir new file mode 100644 index 000000000..ea036277f --- /dev/null +++ b/tests/golden/ir/ttir/golden_loop_under_if_sm80.ttir @@ -0,0 +1,52 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":414:0) +#loc17 = loc("x_ptr"(#loc)) +#loc18 = loc("out_ptr"(#loc)) +#loc19 = loc("n"(#loc)) +module { + tt.func public @loop_under_if_kernel(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc20) + %offs = arith.muli %pid, %c64_i32 : i32 loc(#loc21) + %offs_0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc22) + %offs_1 = tt.splat %offs : i32 -> tensor<64xi32> loc(#loc23) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<64xi32> loc(#loc23) + %0 = arith.cmpi eq, %pid, %c0_i32 : i32 loc(#loc6) + scf.if %0 { + scf.for %i = %c0_i32 to %n step %c1_i32 : i32 { + %1 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc9) + %2 = tt.addptr %1, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc9) + %3 = arith.muli %i, %c64_i32 : i32 loc(#loc10) + %4 = tt.splat %3 : i32 -> tensor<64xi32> loc(#loc11) + %5 = tt.addptr %2, %4 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc11) + %6 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc12) + %7 = tt.addptr %6, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc12) + %8 = tt.addptr %7, %4 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc13) + %9 = tt.load %8 : tensor<64x!tt.ptr> loc(#loc14) + tt.store %5, %9 : tensor<64x!tt.ptr> loc(#loc15) + } loc(#loc8) + } loc(#loc7) + tt.return loc(#loc16) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":415:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":416:17) +#loc4 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":416:38) +#loc5 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":416:25) +#loc6 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":417:14) +#loc7 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":417:7) +#loc8 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":418:26) +#loc9 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":419:31) +#loc10 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":419:42) +#loc11 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":419:38) +#loc12 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":419:65) +#loc13 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":419:72) +#loc14 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":419:57) +#loc15 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":419:49) +#loc16 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":417:4) +#loc20 = loc("pid"(#loc2)) +#loc21 = loc("offs"(#loc3)) +#loc22 = loc("offs"(#loc4)) +#loc23 = loc("offs"(#loc5)) diff --git a/tests/golden/ir/ttir/golden_matmul_bp_s3_sm80.ttir b/tests/golden/ir/ttir/golden_matmul_bp_s3_sm80.ttir new file mode 100644 index 000000000..098c32d2c --- /dev/null +++ b/tests/golden/ir/ttir/golden_matmul_bp_s3_sm80.ttir @@ -0,0 +1,168 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":67:0) +#loc22 = loc("a_ptr"(#loc)) +#loc23 = loc("b_ptr"(#loc)) +#loc24 = loc("c_ptr"(#loc)) +#loc25 = loc("M"(#loc)) +#loc26 = loc("N"(#loc)) +#loc27 = loc("K"(#loc)) +#loc28 = loc("stride_am"(#loc)) +#loc29 = loc("stride_bk"(#loc)) +#loc30 = loc("stride_cm"(#loc)) +module { + tt.func public @matmul_blockptr_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc)), %stride_am: i32 {tt.divisibility = 16 : i32} loc("stride_am"(#loc)), %stride_bk: i32 {tt.divisibility = 16 : i32} loc("stride_bk"(#loc)), %stride_cm: i32 {tt.divisibility = 16 : i32} loc("stride_cm"(#loc))) attributes {noinline = false} { + %c32_i64 = arith.constant 32 : i64 loc(#loc1) + %cst = arith.constant dense<0> : tensor<1x64xi64> loc(#loc1) + %cst_0 = arith.constant dense<0> : tensor<32x1xi64> loc(#loc1) + %cst_1 = arith.constant dense<0> : tensor<1x32xi64> loc(#loc1) + %cst_2 = arith.constant dense<0> : tensor<64x1xi64> loc(#loc1) + %c0_i64 = arith.constant 0 : i64 loc(#loc1) + %c31_i32 = arith.constant 31 : i32 loc(#loc31) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %cst_3 = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc32) + %pid_n = tt.get_program_id y : i32 loc(#loc33) + %a_bp = arith.muli %pid_m, %c64_i32 : i32 loc(#loc34) + %a_bp_4 = arith.extsi %M : i32 to i64 loc(#loc35) + %a_bp_5 = arith.extsi %K : i32 to i64 loc(#loc35) + %a_bp_6 = arith.extsi %stride_am : i32 to i64 loc(#loc35) + %a_bp_7 = arith.extsi %a_bp : i32 to i64 loc(#loc35) + %b_bp = arith.muli %pid_n, %c64_i32 : i32 loc(#loc36) + %b_bp_8 = arith.extsi %N : i32 to i64 loc(#loc37) + %b_bp_9 = arith.extsi %stride_bk : i32 to i64 loc(#loc37) + %b_bp_10 = arith.extsi %b_bp : i32 to i64 loc(#loc37) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc38) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc39) + %acc:3 = scf.for %acc_11 = %c0_i32 to %1 step %c1_i32 iter_args(%a_bp_12 = %c0_i64, %b_bp_13 = %c0_i64, %arg12 = %cst_3) -> (i64, i64, tensor<64x64xf32>) : i32 { + %a = tt.splat %a_ptr : !tt.ptr -> tensor<64x32x!tt.ptr> loc(#loc41) + %a_14 = tt.splat %a_bp_7 : i64 -> tensor<64xi64> loc(#loc41) + %a_15 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc41) + %a_16 = arith.extsi %a_15 : tensor<64xi32> to tensor<64xi64> loc(#loc41) + %a_17 = arith.addi %a_14, %a_16 : tensor<64xi64> loc(#loc41) + %a_18 = tt.expand_dims %a_17 {axis = 1 : i32} : tensor<64xi64> -> tensor<64x1xi64> loc(#loc41) + %a_19 = tt.splat %a_bp_6 : i64 -> tensor<64x1xi64> loc(#loc41) + %a_20 = arith.muli %a_18, %a_19 : tensor<64x1xi64> loc(#loc41) + %a_21 = tt.broadcast %a_20 : tensor<64x1xi64> -> tensor<64x32xi64> loc(#loc41) + %a_22 = tt.splat %a_bp_12 : i64 -> tensor<32xi64> loc(#loc41) + %a_23 = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc41) + %a_24 = arith.extsi %a_23 : tensor<32xi32> to tensor<32xi64> loc(#loc41) + %a_25 = arith.addi %a_22, %a_24 : tensor<32xi64> loc(#loc41) + %a_26 = tt.expand_dims %a_25 {axis = 0 : i32} : tensor<32xi64> -> tensor<1x32xi64> loc(#loc41) + %a_27 = tt.broadcast %a_26 : tensor<1x32xi64> -> tensor<64x32xi64> loc(#loc41) + %a_28 = arith.addi %a_21, %a_27 : tensor<64x32xi64> loc(#loc41) + %a_29 = tt.addptr %a, %a_28 : tensor<64x32x!tt.ptr>, tensor<64x32xi64> loc(#loc41) + %a_30 = arith.cmpi sge, %a_18, %cst_2 : tensor<64x1xi64> loc(#loc41) + %a_31 = tt.splat %a_bp_4 : i64 -> tensor<64x1xi64> loc(#loc41) + %a_32 = arith.cmpi slt, %a_18, %a_31 : tensor<64x1xi64> loc(#loc41) + %a_33 = arith.andi %a_30, %a_32 : tensor<64x1xi1> loc(#loc41) + %a_34 = tt.broadcast %a_33 : tensor<64x1xi1> -> tensor<64x32xi1> loc(#loc41) + %a_35 = arith.cmpi sge, %a_26, %cst_1 : tensor<1x32xi64> loc(#loc41) + %a_36 = tt.splat %a_bp_5 : i64 -> tensor<1x32xi64> loc(#loc41) + %a_37 = arith.cmpi slt, %a_26, %a_36 : tensor<1x32xi64> loc(#loc41) + %a_38 = arith.andi %a_35, %a_37 : tensor<1x32xi1> loc(#loc41) + %a_39 = tt.broadcast %a_38 : tensor<1x32xi1> -> tensor<64x32xi1> loc(#loc41) + %a_40 = arith.andi %a_34, %a_39 : tensor<64x32xi1> loc(#loc41) + %a_41 = tt.load %a_29, %a_40 : tensor<64x32x!tt.ptr> loc(#loc41) + %b = tt.splat %b_ptr : !tt.ptr -> tensor<32x64x!tt.ptr> loc(#loc42) + %b_42 = tt.splat %b_bp_13 : i64 -> tensor<32xi64> loc(#loc42) + %b_43 = arith.addi %b_42, %a_24 : tensor<32xi64> loc(#loc42) + %b_44 = tt.expand_dims %b_43 {axis = 1 : i32} : tensor<32xi64> -> tensor<32x1xi64> loc(#loc42) + %b_45 = tt.splat %b_bp_9 : i64 -> tensor<32x1xi64> loc(#loc42) + %b_46 = arith.muli %b_44, %b_45 : tensor<32x1xi64> loc(#loc42) + %b_47 = tt.broadcast %b_46 : tensor<32x1xi64> -> tensor<32x64xi64> loc(#loc42) + %b_48 = tt.splat %b_bp_10 : i64 -> tensor<64xi64> loc(#loc42) + %b_49 = arith.addi %b_48, %a_16 : tensor<64xi64> loc(#loc42) + %b_50 = tt.expand_dims %b_49 {axis = 0 : i32} : tensor<64xi64> -> tensor<1x64xi64> loc(#loc42) + %b_51 = tt.broadcast %b_50 : tensor<1x64xi64> -> tensor<32x64xi64> loc(#loc42) + %b_52 = arith.addi %b_47, %b_51 : tensor<32x64xi64> loc(#loc42) + %b_53 = tt.addptr %b, %b_52 : tensor<32x64x!tt.ptr>, tensor<32x64xi64> loc(#loc42) + %b_54 = arith.cmpi sge, %b_44, %cst_0 : tensor<32x1xi64> loc(#loc42) + %b_55 = tt.splat %a_bp_5 : i64 -> tensor<32x1xi64> loc(#loc42) + %b_56 = arith.cmpi slt, %b_44, %b_55 : tensor<32x1xi64> loc(#loc42) + %b_57 = arith.andi %b_54, %b_56 : tensor<32x1xi1> loc(#loc42) + %b_58 = tt.broadcast %b_57 : tensor<32x1xi1> -> tensor<32x64xi1> loc(#loc42) + %b_59 = arith.cmpi sge, %b_50, %cst : tensor<1x64xi64> loc(#loc42) + %b_60 = tt.splat %b_bp_8 : i64 -> tensor<1x64xi64> loc(#loc42) + %b_61 = arith.cmpi slt, %b_50, %b_60 : tensor<1x64xi64> loc(#loc42) + %b_62 = arith.andi %b_59, %b_61 : tensor<1x64xi1> loc(#loc42) + %b_63 = tt.broadcast %b_62 : tensor<1x64xi1> -> tensor<32x64xi1> loc(#loc42) + %b_64 = arith.andi %b_58, %b_63 : tensor<32x64xi1> loc(#loc42) + %b_65 = tt.load %b_53, %b_64 : tensor<32x64x!tt.ptr> loc(#loc42) + %acc_66 = tt.dot %a_41, %b_65, %arg12, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc43) + %a_bp_67 = arith.addi %a_bp_12, %c32_i64 : i64 loc(#loc44) + %b_bp_68 = arith.addi %b_bp_13, %c32_i64 : i64 loc(#loc45) + scf.yield %a_bp_67, %b_bp_68, %acc_66 : i64, i64, tensor<64x64xf32> loc(#loc17) + } loc(#loc48) + %c_bp = arith.extsi %stride_cm : i32 to i64 loc(#loc46) + %2 = arith.truncf %acc#2 : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc19) + %3 = tt.splat %c_ptr : !tt.ptr -> tensor<64x64x!tt.ptr> loc(#loc20) + %4 = tt.splat %a_bp_7 : i64 -> tensor<64xi64> loc(#loc20) + %5 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc20) + %6 = arith.extsi %5 : tensor<64xi32> to tensor<64xi64> loc(#loc20) + %7 = arith.addi %4, %6 : tensor<64xi64> loc(#loc20) + %8 = tt.expand_dims %7 {axis = 1 : i32} : tensor<64xi64> -> tensor<64x1xi64> loc(#loc20) + %9 = tt.splat %c_bp : i64 -> tensor<64x1xi64> loc(#loc20) + %10 = arith.muli %8, %9 : tensor<64x1xi64> loc(#loc20) + %11 = tt.broadcast %10 : tensor<64x1xi64> -> tensor<64x64xi64> loc(#loc20) + %12 = tt.splat %b_bp_10 : i64 -> tensor<64xi64> loc(#loc20) + %13 = arith.addi %12, %6 : tensor<64xi64> loc(#loc20) + %14 = tt.expand_dims %13 {axis = 0 : i32} : tensor<64xi64> -> tensor<1x64xi64> loc(#loc20) + %15 = tt.broadcast %14 : tensor<1x64xi64> -> tensor<64x64xi64> loc(#loc20) + %16 = arith.addi %11, %15 : tensor<64x64xi64> loc(#loc20) + %17 = tt.addptr %3, %16 : tensor<64x64x!tt.ptr>, tensor<64x64xi64> loc(#loc20) + %18 = arith.cmpi sge, %8, %cst_2 : tensor<64x1xi64> loc(#loc20) + %19 = tt.splat %a_bp_4 : i64 -> tensor<64x1xi64> loc(#loc20) + %20 = arith.cmpi slt, %8, %19 : tensor<64x1xi64> loc(#loc20) + %21 = arith.andi %18, %20 : tensor<64x1xi1> loc(#loc20) + %22 = tt.broadcast %21 : tensor<64x1xi1> -> tensor<64x64xi1> loc(#loc20) + %23 = arith.cmpi sge, %14, %cst : tensor<1x64xi64> loc(#loc20) + %24 = tt.splat %b_bp_8 : i64 -> tensor<1x64xi64> loc(#loc20) + %25 = arith.cmpi slt, %14, %24 : tensor<1x64xi64> loc(#loc20) + %26 = arith.andi %23, %25 : tensor<1x64xi1> loc(#loc20) + %27 = tt.broadcast %26 : tensor<1x64xi1> -> tensor<64x64xi1> loc(#loc20) + %28 = arith.andi %22, %27 : tensor<64x64xi1> loc(#loc20) + tt.store %17, %2, %28 : tensor<64x64x!tt.ptr> loc(#loc20) + tt.return loc(#loc21) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":90:33) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":90:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":81:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":82:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":84:48) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":84:81) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":87:51) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":87:81) +#loc10 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:17) +#loc11 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:30) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":91:20) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":92:20) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":93:25) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":94:32) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":95:32) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":95:8) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":102:8) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":104:26) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":104:19) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":104:4) +#loc31 = loc(callsite(#loc1 at #loc2)) +#loc32 = loc("pid_m"(#loc4)) +#loc33 = loc("pid_n"(#loc5)) +#loc34 = loc("a_bp"(#loc6)) +#loc35 = loc("a_bp"(#loc7)) +#loc36 = loc("b_bp"(#loc8)) +#loc37 = loc("b_bp"(#loc9)) +#loc38 = loc(callsite(#loc10 at #loc2)) +#loc39 = loc(callsite(#loc11 at #loc2)) +#loc40 = loc("a_bp"(#loc3)) +#loc41 = loc("a"(#loc12)) +#loc42 = loc("b"(#loc13)) +#loc43 = loc("acc"(#loc14)) +#loc44 = loc("a_bp"(#loc15)) +#loc45 = loc("b_bp"(#loc16)) +#loc46 = loc("c_bp"(#loc18)) +#loc47 = loc("b_bp"(#loc40)) +#loc48 = loc("acc"(#loc47)) diff --git a/tests/golden/ir/ttir/golden_matmul_bp_s3_sm90.ttir b/tests/golden/ir/ttir/golden_matmul_bp_s3_sm90.ttir new file mode 100644 index 000000000..098c32d2c --- /dev/null +++ b/tests/golden/ir/ttir/golden_matmul_bp_s3_sm90.ttir @@ -0,0 +1,168 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":67:0) +#loc22 = loc("a_ptr"(#loc)) +#loc23 = loc("b_ptr"(#loc)) +#loc24 = loc("c_ptr"(#loc)) +#loc25 = loc("M"(#loc)) +#loc26 = loc("N"(#loc)) +#loc27 = loc("K"(#loc)) +#loc28 = loc("stride_am"(#loc)) +#loc29 = loc("stride_bk"(#loc)) +#loc30 = loc("stride_cm"(#loc)) +module { + tt.func public @matmul_blockptr_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc)), %stride_am: i32 {tt.divisibility = 16 : i32} loc("stride_am"(#loc)), %stride_bk: i32 {tt.divisibility = 16 : i32} loc("stride_bk"(#loc)), %stride_cm: i32 {tt.divisibility = 16 : i32} loc("stride_cm"(#loc))) attributes {noinline = false} { + %c32_i64 = arith.constant 32 : i64 loc(#loc1) + %cst = arith.constant dense<0> : tensor<1x64xi64> loc(#loc1) + %cst_0 = arith.constant dense<0> : tensor<32x1xi64> loc(#loc1) + %cst_1 = arith.constant dense<0> : tensor<1x32xi64> loc(#loc1) + %cst_2 = arith.constant dense<0> : tensor<64x1xi64> loc(#loc1) + %c0_i64 = arith.constant 0 : i64 loc(#loc1) + %c31_i32 = arith.constant 31 : i32 loc(#loc31) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %cst_3 = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc32) + %pid_n = tt.get_program_id y : i32 loc(#loc33) + %a_bp = arith.muli %pid_m, %c64_i32 : i32 loc(#loc34) + %a_bp_4 = arith.extsi %M : i32 to i64 loc(#loc35) + %a_bp_5 = arith.extsi %K : i32 to i64 loc(#loc35) + %a_bp_6 = arith.extsi %stride_am : i32 to i64 loc(#loc35) + %a_bp_7 = arith.extsi %a_bp : i32 to i64 loc(#loc35) + %b_bp = arith.muli %pid_n, %c64_i32 : i32 loc(#loc36) + %b_bp_8 = arith.extsi %N : i32 to i64 loc(#loc37) + %b_bp_9 = arith.extsi %stride_bk : i32 to i64 loc(#loc37) + %b_bp_10 = arith.extsi %b_bp : i32 to i64 loc(#loc37) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc38) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc39) + %acc:3 = scf.for %acc_11 = %c0_i32 to %1 step %c1_i32 iter_args(%a_bp_12 = %c0_i64, %b_bp_13 = %c0_i64, %arg12 = %cst_3) -> (i64, i64, tensor<64x64xf32>) : i32 { + %a = tt.splat %a_ptr : !tt.ptr -> tensor<64x32x!tt.ptr> loc(#loc41) + %a_14 = tt.splat %a_bp_7 : i64 -> tensor<64xi64> loc(#loc41) + %a_15 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc41) + %a_16 = arith.extsi %a_15 : tensor<64xi32> to tensor<64xi64> loc(#loc41) + %a_17 = arith.addi %a_14, %a_16 : tensor<64xi64> loc(#loc41) + %a_18 = tt.expand_dims %a_17 {axis = 1 : i32} : tensor<64xi64> -> tensor<64x1xi64> loc(#loc41) + %a_19 = tt.splat %a_bp_6 : i64 -> tensor<64x1xi64> loc(#loc41) + %a_20 = arith.muli %a_18, %a_19 : tensor<64x1xi64> loc(#loc41) + %a_21 = tt.broadcast %a_20 : tensor<64x1xi64> -> tensor<64x32xi64> loc(#loc41) + %a_22 = tt.splat %a_bp_12 : i64 -> tensor<32xi64> loc(#loc41) + %a_23 = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc41) + %a_24 = arith.extsi %a_23 : tensor<32xi32> to tensor<32xi64> loc(#loc41) + %a_25 = arith.addi %a_22, %a_24 : tensor<32xi64> loc(#loc41) + %a_26 = tt.expand_dims %a_25 {axis = 0 : i32} : tensor<32xi64> -> tensor<1x32xi64> loc(#loc41) + %a_27 = tt.broadcast %a_26 : tensor<1x32xi64> -> tensor<64x32xi64> loc(#loc41) + %a_28 = arith.addi %a_21, %a_27 : tensor<64x32xi64> loc(#loc41) + %a_29 = tt.addptr %a, %a_28 : tensor<64x32x!tt.ptr>, tensor<64x32xi64> loc(#loc41) + %a_30 = arith.cmpi sge, %a_18, %cst_2 : tensor<64x1xi64> loc(#loc41) + %a_31 = tt.splat %a_bp_4 : i64 -> tensor<64x1xi64> loc(#loc41) + %a_32 = arith.cmpi slt, %a_18, %a_31 : tensor<64x1xi64> loc(#loc41) + %a_33 = arith.andi %a_30, %a_32 : tensor<64x1xi1> loc(#loc41) + %a_34 = tt.broadcast %a_33 : tensor<64x1xi1> -> tensor<64x32xi1> loc(#loc41) + %a_35 = arith.cmpi sge, %a_26, %cst_1 : tensor<1x32xi64> loc(#loc41) + %a_36 = tt.splat %a_bp_5 : i64 -> tensor<1x32xi64> loc(#loc41) + %a_37 = arith.cmpi slt, %a_26, %a_36 : tensor<1x32xi64> loc(#loc41) + %a_38 = arith.andi %a_35, %a_37 : tensor<1x32xi1> loc(#loc41) + %a_39 = tt.broadcast %a_38 : tensor<1x32xi1> -> tensor<64x32xi1> loc(#loc41) + %a_40 = arith.andi %a_34, %a_39 : tensor<64x32xi1> loc(#loc41) + %a_41 = tt.load %a_29, %a_40 : tensor<64x32x!tt.ptr> loc(#loc41) + %b = tt.splat %b_ptr : !tt.ptr -> tensor<32x64x!tt.ptr> loc(#loc42) + %b_42 = tt.splat %b_bp_13 : i64 -> tensor<32xi64> loc(#loc42) + %b_43 = arith.addi %b_42, %a_24 : tensor<32xi64> loc(#loc42) + %b_44 = tt.expand_dims %b_43 {axis = 1 : i32} : tensor<32xi64> -> tensor<32x1xi64> loc(#loc42) + %b_45 = tt.splat %b_bp_9 : i64 -> tensor<32x1xi64> loc(#loc42) + %b_46 = arith.muli %b_44, %b_45 : tensor<32x1xi64> loc(#loc42) + %b_47 = tt.broadcast %b_46 : tensor<32x1xi64> -> tensor<32x64xi64> loc(#loc42) + %b_48 = tt.splat %b_bp_10 : i64 -> tensor<64xi64> loc(#loc42) + %b_49 = arith.addi %b_48, %a_16 : tensor<64xi64> loc(#loc42) + %b_50 = tt.expand_dims %b_49 {axis = 0 : i32} : tensor<64xi64> -> tensor<1x64xi64> loc(#loc42) + %b_51 = tt.broadcast %b_50 : tensor<1x64xi64> -> tensor<32x64xi64> loc(#loc42) + %b_52 = arith.addi %b_47, %b_51 : tensor<32x64xi64> loc(#loc42) + %b_53 = tt.addptr %b, %b_52 : tensor<32x64x!tt.ptr>, tensor<32x64xi64> loc(#loc42) + %b_54 = arith.cmpi sge, %b_44, %cst_0 : tensor<32x1xi64> loc(#loc42) + %b_55 = tt.splat %a_bp_5 : i64 -> tensor<32x1xi64> loc(#loc42) + %b_56 = arith.cmpi slt, %b_44, %b_55 : tensor<32x1xi64> loc(#loc42) + %b_57 = arith.andi %b_54, %b_56 : tensor<32x1xi1> loc(#loc42) + %b_58 = tt.broadcast %b_57 : tensor<32x1xi1> -> tensor<32x64xi1> loc(#loc42) + %b_59 = arith.cmpi sge, %b_50, %cst : tensor<1x64xi64> loc(#loc42) + %b_60 = tt.splat %b_bp_8 : i64 -> tensor<1x64xi64> loc(#loc42) + %b_61 = arith.cmpi slt, %b_50, %b_60 : tensor<1x64xi64> loc(#loc42) + %b_62 = arith.andi %b_59, %b_61 : tensor<1x64xi1> loc(#loc42) + %b_63 = tt.broadcast %b_62 : tensor<1x64xi1> -> tensor<32x64xi1> loc(#loc42) + %b_64 = arith.andi %b_58, %b_63 : tensor<32x64xi1> loc(#loc42) + %b_65 = tt.load %b_53, %b_64 : tensor<32x64x!tt.ptr> loc(#loc42) + %acc_66 = tt.dot %a_41, %b_65, %arg12, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc43) + %a_bp_67 = arith.addi %a_bp_12, %c32_i64 : i64 loc(#loc44) + %b_bp_68 = arith.addi %b_bp_13, %c32_i64 : i64 loc(#loc45) + scf.yield %a_bp_67, %b_bp_68, %acc_66 : i64, i64, tensor<64x64xf32> loc(#loc17) + } loc(#loc48) + %c_bp = arith.extsi %stride_cm : i32 to i64 loc(#loc46) + %2 = arith.truncf %acc#2 : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc19) + %3 = tt.splat %c_ptr : !tt.ptr -> tensor<64x64x!tt.ptr> loc(#loc20) + %4 = tt.splat %a_bp_7 : i64 -> tensor<64xi64> loc(#loc20) + %5 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc20) + %6 = arith.extsi %5 : tensor<64xi32> to tensor<64xi64> loc(#loc20) + %7 = arith.addi %4, %6 : tensor<64xi64> loc(#loc20) + %8 = tt.expand_dims %7 {axis = 1 : i32} : tensor<64xi64> -> tensor<64x1xi64> loc(#loc20) + %9 = tt.splat %c_bp : i64 -> tensor<64x1xi64> loc(#loc20) + %10 = arith.muli %8, %9 : tensor<64x1xi64> loc(#loc20) + %11 = tt.broadcast %10 : tensor<64x1xi64> -> tensor<64x64xi64> loc(#loc20) + %12 = tt.splat %b_bp_10 : i64 -> tensor<64xi64> loc(#loc20) + %13 = arith.addi %12, %6 : tensor<64xi64> loc(#loc20) + %14 = tt.expand_dims %13 {axis = 0 : i32} : tensor<64xi64> -> tensor<1x64xi64> loc(#loc20) + %15 = tt.broadcast %14 : tensor<1x64xi64> -> tensor<64x64xi64> loc(#loc20) + %16 = arith.addi %11, %15 : tensor<64x64xi64> loc(#loc20) + %17 = tt.addptr %3, %16 : tensor<64x64x!tt.ptr>, tensor<64x64xi64> loc(#loc20) + %18 = arith.cmpi sge, %8, %cst_2 : tensor<64x1xi64> loc(#loc20) + %19 = tt.splat %a_bp_4 : i64 -> tensor<64x1xi64> loc(#loc20) + %20 = arith.cmpi slt, %8, %19 : tensor<64x1xi64> loc(#loc20) + %21 = arith.andi %18, %20 : tensor<64x1xi1> loc(#loc20) + %22 = tt.broadcast %21 : tensor<64x1xi1> -> tensor<64x64xi1> loc(#loc20) + %23 = arith.cmpi sge, %14, %cst : tensor<1x64xi64> loc(#loc20) + %24 = tt.splat %b_bp_8 : i64 -> tensor<1x64xi64> loc(#loc20) + %25 = arith.cmpi slt, %14, %24 : tensor<1x64xi64> loc(#loc20) + %26 = arith.andi %23, %25 : tensor<1x64xi1> loc(#loc20) + %27 = tt.broadcast %26 : tensor<1x64xi1> -> tensor<64x64xi1> loc(#loc20) + %28 = arith.andi %22, %27 : tensor<64x64xi1> loc(#loc20) + tt.store %17, %2, %28 : tensor<64x64x!tt.ptr> loc(#loc20) + tt.return loc(#loc21) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":90:33) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":90:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":81:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":82:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":84:48) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":84:81) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":87:51) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":87:81) +#loc10 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:17) +#loc11 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:30) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":91:20) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":92:20) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":93:25) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":94:32) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":95:32) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":95:8) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":102:8) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":104:26) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":104:19) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":104:4) +#loc31 = loc(callsite(#loc1 at #loc2)) +#loc32 = loc("pid_m"(#loc4)) +#loc33 = loc("pid_n"(#loc5)) +#loc34 = loc("a_bp"(#loc6)) +#loc35 = loc("a_bp"(#loc7)) +#loc36 = loc("b_bp"(#loc8)) +#loc37 = loc("b_bp"(#loc9)) +#loc38 = loc(callsite(#loc10 at #loc2)) +#loc39 = loc(callsite(#loc11 at #loc2)) +#loc40 = loc("a_bp"(#loc3)) +#loc41 = loc("a"(#loc12)) +#loc42 = loc("b"(#loc13)) +#loc43 = loc("acc"(#loc14)) +#loc44 = loc("a_bp"(#loc15)) +#loc45 = loc("b_bp"(#loc16)) +#loc46 = loc("c_bp"(#loc18)) +#loc47 = loc("b_bp"(#loc40)) +#loc48 = loc("acc"(#loc47)) diff --git a/tests/golden/ir/ttir/golden_matmul_s1_sm80.ttir b/tests/golden/ir/ttir/golden_matmul_s1_sm80.ttir new file mode 100644 index 000000000..72c5a1f80 --- /dev/null +++ b/tests/golden/ir/ttir/golden_matmul_s1_sm80.ttir @@ -0,0 +1,172 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":28:0) +#loc44 = loc("a_ptr"(#loc)) +#loc45 = loc("b_ptr"(#loc)) +#loc46 = loc("c_ptr"(#loc)) +#loc47 = loc("M"(#loc)) +#loc48 = loc("N"(#loc)) +#loc49 = loc("K"(#loc)) +#loc50 = loc("stride_am"(#loc)) +#loc51 = loc("stride_bk"(#loc)) +#loc52 = loc("stride_cm"(#loc)) +module { + tt.func public @matmul_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc)), %stride_am: i32 {tt.divisibility = 16 : i32} loc("stride_am"(#loc)), %stride_bk: i32 {tt.divisibility = 16 : i32} loc("stride_bk"(#loc)), %stride_cm: i32 {tt.divisibility = 16 : i32} loc("stride_cm"(#loc))) attributes {noinline = false} { + %c31_i32 = arith.constant 31 : i32 loc(#loc53) + %cst = arith.constant dense<0.000000e+00> : tensor<32x64xf16> loc(#loc1) + %cst_0 = arith.constant dense<0.000000e+00> : tensor<64x32xf16> loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %cst_1 = arith.constant dense<32> : tensor<64x32xi32> loc(#loc1) + %cst_2 = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc54) + %pid_n = tt.get_program_id y : i32 loc(#loc55) + %offs_m = arith.muli %pid_m, %c64_i32 : i32 loc(#loc56) + %offs_m_3 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc57) + %offs_m_4 = tt.splat %offs_m : i32 -> tensor<64xi32> loc(#loc58) + %offs_m_5 = arith.addi %offs_m_4, %offs_m_3 : tensor<64xi32> loc(#loc58) + %offs_n = arith.muli %pid_n, %c64_i32 : i32 loc(#loc59) + %offs_n_6 = tt.splat %offs_n : i32 -> tensor<64xi32> loc(#loc60) + %offs_n_7 = arith.addi %offs_n_6, %offs_m_3 : tensor<64xi32> loc(#loc60) + %offs_k = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc61) + %a_ptrs = tt.expand_dims %offs_m_5 {axis = 1 : i32} : tensor<64xi32> -> tensor<64x1xi32> loc(#loc62) + %a_ptrs_8 = tt.splat %stride_am : i32 -> tensor<64x1xi32> loc(#loc63) + %a_ptrs_9 = arith.muli %a_ptrs, %a_ptrs_8 : tensor<64x1xi32> loc(#loc63) + %a_ptrs_10 = tt.splat %a_ptr : !tt.ptr -> tensor<64x1x!tt.ptr> loc(#loc64) + %a_ptrs_11 = tt.addptr %a_ptrs_10, %a_ptrs_9 : tensor<64x1x!tt.ptr>, tensor<64x1xi32> loc(#loc64) + %a_ptrs_12 = tt.expand_dims %offs_k {axis = 0 : i32} : tensor<32xi32> -> tensor<1x32xi32> loc(#loc65) + %a_ptrs_13 = tt.broadcast %a_ptrs_11 : tensor<64x1x!tt.ptr> -> tensor<64x32x!tt.ptr> loc(#loc66) + %a_ptrs_14 = tt.broadcast %a_ptrs_12 : tensor<1x32xi32> -> tensor<64x32xi32> loc(#loc66) + %a_ptrs_15 = tt.addptr %a_ptrs_13, %a_ptrs_14 : tensor<64x32x!tt.ptr>, tensor<64x32xi32> loc(#loc66) + %b_ptrs = tt.expand_dims %offs_k {axis = 1 : i32} : tensor<32xi32> -> tensor<32x1xi32> loc(#loc67) + %b_ptrs_16 = tt.splat %stride_bk : i32 -> tensor<32x1xi32> loc(#loc68) + %b_ptrs_17 = arith.muli %b_ptrs, %b_ptrs_16 : tensor<32x1xi32> loc(#loc68) + %b_ptrs_18 = tt.splat %b_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc69) + %b_ptrs_19 = tt.addptr %b_ptrs_18, %b_ptrs_17 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc69) + %b_ptrs_20 = tt.expand_dims %offs_n_7 {axis = 0 : i32} : tensor<64xi32> -> tensor<1x64xi32> loc(#loc70) + %b_ptrs_21 = tt.broadcast %b_ptrs_19 : tensor<32x1x!tt.ptr> -> tensor<32x64x!tt.ptr> loc(#loc71) + %b_ptrs_22 = tt.broadcast %b_ptrs_20 : tensor<1x64xi32> -> tensor<32x64xi32> loc(#loc71) + %b_ptrs_23 = tt.addptr %b_ptrs_21, %b_ptrs_22 : tensor<32x64x!tt.ptr>, tensor<32x64xi32> loc(#loc71) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc72) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc73) + %acc:3 = scf.for %k = %c0_i32 to %1 step %c1_i32 iter_args(%a_ptrs_36 = %a_ptrs_15, %b_ptrs_37 = %b_ptrs_23, %acc_38 = %cst_2) -> (tensor<64x32x!tt.ptr>, tensor<32x64x!tt.ptr>, tensor<64x64xf32>) : i32 { + %a = arith.muli %k, %c32_i32 : i32 loc(#loc75) + %a_39 = arith.subi %K, %a : i32 loc(#loc76) + %a_40 = tt.splat %a_39 : i32 -> tensor<1x32xi32> loc(#loc77) + %a_41 = arith.cmpi slt, %a_ptrs_12, %a_40 : tensor<1x32xi32> loc(#loc77) + %a_42 = tt.broadcast %a_41 : tensor<1x32xi1> -> tensor<64x32xi1> loc(#loc78) + %a_43 = tt.load %a_ptrs_36, %a_42, %cst_0 : tensor<64x32x!tt.ptr> loc(#loc78) + %b = tt.splat %a_39 : i32 -> tensor<32x1xi32> loc(#loc79) + %b_44 = arith.cmpi slt, %b_ptrs, %b : tensor<32x1xi32> loc(#loc79) + %b_45 = tt.broadcast %b_44 : tensor<32x1xi1> -> tensor<32x64xi1> loc(#loc80) + %b_46 = tt.load %b_ptrs_37, %b_45, %cst : tensor<32x64x!tt.ptr> loc(#loc80) + %acc_47 = tt.dot %a_43, %b_46, %acc_38, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc81) + %a_ptrs_48 = tt.addptr %a_ptrs_36, %cst_1 : tensor<64x32x!tt.ptr>, tensor<64x32xi32> loc(#loc82) + %b_ptrs_49 = arith.muli %stride_bk, %c32_i32 : i32 loc(#loc83) + %b_ptrs_50 = tt.splat %b_ptrs_49 : i32 -> tensor<32x64xi32> loc(#loc84) + %b_ptrs_51 = tt.addptr %b_ptrs_37, %b_ptrs_50 : tensor<32x64x!tt.ptr>, tensor<32x64xi32> loc(#loc84) + scf.yield %a_ptrs_48, %b_ptrs_51, %acc_47 : tensor<64x32x!tt.ptr>, tensor<32x64x!tt.ptr>, tensor<64x64xf32> loc(#loc34) + } loc(#loc93) + %c = arith.truncf %acc#2 : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc85) + %c_ptrs = tt.splat %stride_cm : i32 -> tensor<64x1xi32> loc(#loc86) + %c_ptrs_24 = arith.muli %a_ptrs, %c_ptrs : tensor<64x1xi32> loc(#loc86) + %c_ptrs_25 = tt.splat %c_ptr : !tt.ptr -> tensor<64x1x!tt.ptr> loc(#loc87) + %c_ptrs_26 = tt.addptr %c_ptrs_25, %c_ptrs_24 : tensor<64x1x!tt.ptr>, tensor<64x1xi32> loc(#loc87) + %c_ptrs_27 = tt.broadcast %c_ptrs_26 : tensor<64x1x!tt.ptr> -> tensor<64x64x!tt.ptr> loc(#loc88) + %c_ptrs_28 = tt.broadcast %b_ptrs_20 : tensor<1x64xi32> -> tensor<64x64xi32> loc(#loc88) + %c_ptrs_29 = tt.addptr %c_ptrs_27, %c_ptrs_28 : tensor<64x64x!tt.ptr>, tensor<64x64xi32> loc(#loc88) + %c_mask = tt.splat %M : i32 -> tensor<64x1xi32> loc(#loc89) + %c_mask_30 = arith.cmpi slt, %a_ptrs, %c_mask : tensor<64x1xi32> loc(#loc89) + %c_mask_31 = tt.splat %N : i32 -> tensor<1x64xi32> loc(#loc90) + %c_mask_32 = arith.cmpi slt, %b_ptrs_20, %c_mask_31 : tensor<1x64xi32> loc(#loc90) + %c_mask_33 = tt.broadcast %c_mask_30 : tensor<64x1xi1> -> tensor<64x64xi1> loc(#loc91) + %c_mask_34 = tt.broadcast %c_mask_32 : tensor<1x64xi1> -> tensor<64x64xi1> loc(#loc91) + %c_mask_35 = arith.andi %c_mask_33, %c_mask_34 : tensor<64x64xi1> loc(#loc91) + tt.store %c_ptrs_29, %c, %c_mask_35 : tensor<64x64x!tt.ptr> loc(#loc42) + tt.return loc(#loc43) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":53:33) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":53:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":42:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":43:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:21) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:44) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:31) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":46:21) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":46:31) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":47:26) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:28) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:39) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:21) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:58) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:51) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:28) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:39) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:21) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:58) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:51) +#loc22 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:17) +#loc23 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:30) +#loc24 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:59) +#loc25 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:55) +#loc26 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:51) +#loc27 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:20) +#loc28 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":55:51) +#loc29 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":55:20) +#loc30 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":56:25) +#loc31 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":57:18) +#loc32 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:28) +#loc33 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:18) +#loc34 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:8) +#loc35 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":60:15) +#loc36 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:39) +#loc37 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:21) +#loc38 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:51) +#loc39 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:32) +#loc40 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:56) +#loc41 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:38) +#loc42 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":63:21) +#loc43 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":63:4) +#loc53 = loc(callsite(#loc1 at #loc2)) +#loc54 = loc("pid_m"(#loc4)) +#loc55 = loc("pid_n"(#loc5)) +#loc56 = loc("offs_m"(#loc6)) +#loc57 = loc("offs_m"(#loc7)) +#loc58 = loc("offs_m"(#loc8)) +#loc59 = loc("offs_n"(#loc9)) +#loc60 = loc("offs_n"(#loc10)) +#loc61 = loc("offs_k"(#loc11)) +#loc62 = loc("a_ptrs"(#loc12)) +#loc63 = loc("a_ptrs"(#loc13)) +#loc64 = loc("a_ptrs"(#loc14)) +#loc65 = loc("a_ptrs"(#loc15)) +#loc66 = loc("a_ptrs"(#loc16)) +#loc67 = loc("b_ptrs"(#loc17)) +#loc68 = loc("b_ptrs"(#loc18)) +#loc69 = loc("b_ptrs"(#loc19)) +#loc70 = loc("b_ptrs"(#loc20)) +#loc71 = loc("b_ptrs"(#loc21)) +#loc72 = loc(callsite(#loc22 at #loc2)) +#loc73 = loc(callsite(#loc23 at #loc2)) +#loc74 = loc("a_ptrs"(#loc3)) +#loc75 = loc("a"(#loc24)) +#loc76 = loc("a"(#loc25)) +#loc77 = loc("a"(#loc26)) +#loc78 = loc("a"(#loc27)) +#loc79 = loc("b"(#loc28)) +#loc80 = loc("b"(#loc29)) +#loc81 = loc("acc"(#loc30)) +#loc82 = loc("a_ptrs"(#loc31)) +#loc83 = loc("b_ptrs"(#loc32)) +#loc84 = loc("b_ptrs"(#loc33)) +#loc85 = loc("c"(#loc35)) +#loc86 = loc("c_ptrs"(#loc36)) +#loc87 = loc("c_ptrs"(#loc37)) +#loc88 = loc("c_ptrs"(#loc38)) +#loc89 = loc("c_mask"(#loc39)) +#loc90 = loc("c_mask"(#loc40)) +#loc91 = loc("c_mask"(#loc41)) +#loc92 = loc("b_ptrs"(#loc74)) +#loc93 = loc("acc"(#loc92)) diff --git a/tests/golden/ir/ttir/golden_matmul_s1_sm90.ttir b/tests/golden/ir/ttir/golden_matmul_s1_sm90.ttir new file mode 100644 index 000000000..72c5a1f80 --- /dev/null +++ b/tests/golden/ir/ttir/golden_matmul_s1_sm90.ttir @@ -0,0 +1,172 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":28:0) +#loc44 = loc("a_ptr"(#loc)) +#loc45 = loc("b_ptr"(#loc)) +#loc46 = loc("c_ptr"(#loc)) +#loc47 = loc("M"(#loc)) +#loc48 = loc("N"(#loc)) +#loc49 = loc("K"(#loc)) +#loc50 = loc("stride_am"(#loc)) +#loc51 = loc("stride_bk"(#loc)) +#loc52 = loc("stride_cm"(#loc)) +module { + tt.func public @matmul_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc)), %stride_am: i32 {tt.divisibility = 16 : i32} loc("stride_am"(#loc)), %stride_bk: i32 {tt.divisibility = 16 : i32} loc("stride_bk"(#loc)), %stride_cm: i32 {tt.divisibility = 16 : i32} loc("stride_cm"(#loc))) attributes {noinline = false} { + %c31_i32 = arith.constant 31 : i32 loc(#loc53) + %cst = arith.constant dense<0.000000e+00> : tensor<32x64xf16> loc(#loc1) + %cst_0 = arith.constant dense<0.000000e+00> : tensor<64x32xf16> loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %cst_1 = arith.constant dense<32> : tensor<64x32xi32> loc(#loc1) + %cst_2 = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc54) + %pid_n = tt.get_program_id y : i32 loc(#loc55) + %offs_m = arith.muli %pid_m, %c64_i32 : i32 loc(#loc56) + %offs_m_3 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc57) + %offs_m_4 = tt.splat %offs_m : i32 -> tensor<64xi32> loc(#loc58) + %offs_m_5 = arith.addi %offs_m_4, %offs_m_3 : tensor<64xi32> loc(#loc58) + %offs_n = arith.muli %pid_n, %c64_i32 : i32 loc(#loc59) + %offs_n_6 = tt.splat %offs_n : i32 -> tensor<64xi32> loc(#loc60) + %offs_n_7 = arith.addi %offs_n_6, %offs_m_3 : tensor<64xi32> loc(#loc60) + %offs_k = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc61) + %a_ptrs = tt.expand_dims %offs_m_5 {axis = 1 : i32} : tensor<64xi32> -> tensor<64x1xi32> loc(#loc62) + %a_ptrs_8 = tt.splat %stride_am : i32 -> tensor<64x1xi32> loc(#loc63) + %a_ptrs_9 = arith.muli %a_ptrs, %a_ptrs_8 : tensor<64x1xi32> loc(#loc63) + %a_ptrs_10 = tt.splat %a_ptr : !tt.ptr -> tensor<64x1x!tt.ptr> loc(#loc64) + %a_ptrs_11 = tt.addptr %a_ptrs_10, %a_ptrs_9 : tensor<64x1x!tt.ptr>, tensor<64x1xi32> loc(#loc64) + %a_ptrs_12 = tt.expand_dims %offs_k {axis = 0 : i32} : tensor<32xi32> -> tensor<1x32xi32> loc(#loc65) + %a_ptrs_13 = tt.broadcast %a_ptrs_11 : tensor<64x1x!tt.ptr> -> tensor<64x32x!tt.ptr> loc(#loc66) + %a_ptrs_14 = tt.broadcast %a_ptrs_12 : tensor<1x32xi32> -> tensor<64x32xi32> loc(#loc66) + %a_ptrs_15 = tt.addptr %a_ptrs_13, %a_ptrs_14 : tensor<64x32x!tt.ptr>, tensor<64x32xi32> loc(#loc66) + %b_ptrs = tt.expand_dims %offs_k {axis = 1 : i32} : tensor<32xi32> -> tensor<32x1xi32> loc(#loc67) + %b_ptrs_16 = tt.splat %stride_bk : i32 -> tensor<32x1xi32> loc(#loc68) + %b_ptrs_17 = arith.muli %b_ptrs, %b_ptrs_16 : tensor<32x1xi32> loc(#loc68) + %b_ptrs_18 = tt.splat %b_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc69) + %b_ptrs_19 = tt.addptr %b_ptrs_18, %b_ptrs_17 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc69) + %b_ptrs_20 = tt.expand_dims %offs_n_7 {axis = 0 : i32} : tensor<64xi32> -> tensor<1x64xi32> loc(#loc70) + %b_ptrs_21 = tt.broadcast %b_ptrs_19 : tensor<32x1x!tt.ptr> -> tensor<32x64x!tt.ptr> loc(#loc71) + %b_ptrs_22 = tt.broadcast %b_ptrs_20 : tensor<1x64xi32> -> tensor<32x64xi32> loc(#loc71) + %b_ptrs_23 = tt.addptr %b_ptrs_21, %b_ptrs_22 : tensor<32x64x!tt.ptr>, tensor<32x64xi32> loc(#loc71) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc72) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc73) + %acc:3 = scf.for %k = %c0_i32 to %1 step %c1_i32 iter_args(%a_ptrs_36 = %a_ptrs_15, %b_ptrs_37 = %b_ptrs_23, %acc_38 = %cst_2) -> (tensor<64x32x!tt.ptr>, tensor<32x64x!tt.ptr>, tensor<64x64xf32>) : i32 { + %a = arith.muli %k, %c32_i32 : i32 loc(#loc75) + %a_39 = arith.subi %K, %a : i32 loc(#loc76) + %a_40 = tt.splat %a_39 : i32 -> tensor<1x32xi32> loc(#loc77) + %a_41 = arith.cmpi slt, %a_ptrs_12, %a_40 : tensor<1x32xi32> loc(#loc77) + %a_42 = tt.broadcast %a_41 : tensor<1x32xi1> -> tensor<64x32xi1> loc(#loc78) + %a_43 = tt.load %a_ptrs_36, %a_42, %cst_0 : tensor<64x32x!tt.ptr> loc(#loc78) + %b = tt.splat %a_39 : i32 -> tensor<32x1xi32> loc(#loc79) + %b_44 = arith.cmpi slt, %b_ptrs, %b : tensor<32x1xi32> loc(#loc79) + %b_45 = tt.broadcast %b_44 : tensor<32x1xi1> -> tensor<32x64xi1> loc(#loc80) + %b_46 = tt.load %b_ptrs_37, %b_45, %cst : tensor<32x64x!tt.ptr> loc(#loc80) + %acc_47 = tt.dot %a_43, %b_46, %acc_38, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc81) + %a_ptrs_48 = tt.addptr %a_ptrs_36, %cst_1 : tensor<64x32x!tt.ptr>, tensor<64x32xi32> loc(#loc82) + %b_ptrs_49 = arith.muli %stride_bk, %c32_i32 : i32 loc(#loc83) + %b_ptrs_50 = tt.splat %b_ptrs_49 : i32 -> tensor<32x64xi32> loc(#loc84) + %b_ptrs_51 = tt.addptr %b_ptrs_37, %b_ptrs_50 : tensor<32x64x!tt.ptr>, tensor<32x64xi32> loc(#loc84) + scf.yield %a_ptrs_48, %b_ptrs_51, %acc_47 : tensor<64x32x!tt.ptr>, tensor<32x64x!tt.ptr>, tensor<64x64xf32> loc(#loc34) + } loc(#loc93) + %c = arith.truncf %acc#2 : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc85) + %c_ptrs = tt.splat %stride_cm : i32 -> tensor<64x1xi32> loc(#loc86) + %c_ptrs_24 = arith.muli %a_ptrs, %c_ptrs : tensor<64x1xi32> loc(#loc86) + %c_ptrs_25 = tt.splat %c_ptr : !tt.ptr -> tensor<64x1x!tt.ptr> loc(#loc87) + %c_ptrs_26 = tt.addptr %c_ptrs_25, %c_ptrs_24 : tensor<64x1x!tt.ptr>, tensor<64x1xi32> loc(#loc87) + %c_ptrs_27 = tt.broadcast %c_ptrs_26 : tensor<64x1x!tt.ptr> -> tensor<64x64x!tt.ptr> loc(#loc88) + %c_ptrs_28 = tt.broadcast %b_ptrs_20 : tensor<1x64xi32> -> tensor<64x64xi32> loc(#loc88) + %c_ptrs_29 = tt.addptr %c_ptrs_27, %c_ptrs_28 : tensor<64x64x!tt.ptr>, tensor<64x64xi32> loc(#loc88) + %c_mask = tt.splat %M : i32 -> tensor<64x1xi32> loc(#loc89) + %c_mask_30 = arith.cmpi slt, %a_ptrs, %c_mask : tensor<64x1xi32> loc(#loc89) + %c_mask_31 = tt.splat %N : i32 -> tensor<1x64xi32> loc(#loc90) + %c_mask_32 = arith.cmpi slt, %b_ptrs_20, %c_mask_31 : tensor<1x64xi32> loc(#loc90) + %c_mask_33 = tt.broadcast %c_mask_30 : tensor<64x1xi1> -> tensor<64x64xi1> loc(#loc91) + %c_mask_34 = tt.broadcast %c_mask_32 : tensor<1x64xi1> -> tensor<64x64xi1> loc(#loc91) + %c_mask_35 = arith.andi %c_mask_33, %c_mask_34 : tensor<64x64xi1> loc(#loc91) + tt.store %c_ptrs_29, %c, %c_mask_35 : tensor<64x64x!tt.ptr> loc(#loc42) + tt.return loc(#loc43) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":53:33) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":53:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":42:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":43:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:21) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:44) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:31) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":46:21) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":46:31) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":47:26) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:28) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:39) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:21) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:58) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:51) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:28) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:39) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:21) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:58) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:51) +#loc22 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:17) +#loc23 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:30) +#loc24 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:59) +#loc25 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:55) +#loc26 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:51) +#loc27 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:20) +#loc28 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":55:51) +#loc29 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":55:20) +#loc30 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":56:25) +#loc31 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":57:18) +#loc32 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:28) +#loc33 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:18) +#loc34 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:8) +#loc35 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":60:15) +#loc36 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:39) +#loc37 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:21) +#loc38 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:51) +#loc39 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:32) +#loc40 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:56) +#loc41 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:38) +#loc42 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":63:21) +#loc43 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":63:4) +#loc53 = loc(callsite(#loc1 at #loc2)) +#loc54 = loc("pid_m"(#loc4)) +#loc55 = loc("pid_n"(#loc5)) +#loc56 = loc("offs_m"(#loc6)) +#loc57 = loc("offs_m"(#loc7)) +#loc58 = loc("offs_m"(#loc8)) +#loc59 = loc("offs_n"(#loc9)) +#loc60 = loc("offs_n"(#loc10)) +#loc61 = loc("offs_k"(#loc11)) +#loc62 = loc("a_ptrs"(#loc12)) +#loc63 = loc("a_ptrs"(#loc13)) +#loc64 = loc("a_ptrs"(#loc14)) +#loc65 = loc("a_ptrs"(#loc15)) +#loc66 = loc("a_ptrs"(#loc16)) +#loc67 = loc("b_ptrs"(#loc17)) +#loc68 = loc("b_ptrs"(#loc18)) +#loc69 = loc("b_ptrs"(#loc19)) +#loc70 = loc("b_ptrs"(#loc20)) +#loc71 = loc("b_ptrs"(#loc21)) +#loc72 = loc(callsite(#loc22 at #loc2)) +#loc73 = loc(callsite(#loc23 at #loc2)) +#loc74 = loc("a_ptrs"(#loc3)) +#loc75 = loc("a"(#loc24)) +#loc76 = loc("a"(#loc25)) +#loc77 = loc("a"(#loc26)) +#loc78 = loc("a"(#loc27)) +#loc79 = loc("b"(#loc28)) +#loc80 = loc("b"(#loc29)) +#loc81 = loc("acc"(#loc30)) +#loc82 = loc("a_ptrs"(#loc31)) +#loc83 = loc("b_ptrs"(#loc32)) +#loc84 = loc("b_ptrs"(#loc33)) +#loc85 = loc("c"(#loc35)) +#loc86 = loc("c_ptrs"(#loc36)) +#loc87 = loc("c_ptrs"(#loc37)) +#loc88 = loc("c_ptrs"(#loc38)) +#loc89 = loc("c_mask"(#loc39)) +#loc90 = loc("c_mask"(#loc40)) +#loc91 = loc("c_mask"(#loc41)) +#loc92 = loc("b_ptrs"(#loc74)) +#loc93 = loc("acc"(#loc92)) diff --git a/tests/golden/ir/ttir/golden_matmul_s3_sm80.ttir b/tests/golden/ir/ttir/golden_matmul_s3_sm80.ttir new file mode 100644 index 000000000..72c5a1f80 --- /dev/null +++ b/tests/golden/ir/ttir/golden_matmul_s3_sm80.ttir @@ -0,0 +1,172 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":28:0) +#loc44 = loc("a_ptr"(#loc)) +#loc45 = loc("b_ptr"(#loc)) +#loc46 = loc("c_ptr"(#loc)) +#loc47 = loc("M"(#loc)) +#loc48 = loc("N"(#loc)) +#loc49 = loc("K"(#loc)) +#loc50 = loc("stride_am"(#loc)) +#loc51 = loc("stride_bk"(#loc)) +#loc52 = loc("stride_cm"(#loc)) +module { + tt.func public @matmul_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc)), %stride_am: i32 {tt.divisibility = 16 : i32} loc("stride_am"(#loc)), %stride_bk: i32 {tt.divisibility = 16 : i32} loc("stride_bk"(#loc)), %stride_cm: i32 {tt.divisibility = 16 : i32} loc("stride_cm"(#loc))) attributes {noinline = false} { + %c31_i32 = arith.constant 31 : i32 loc(#loc53) + %cst = arith.constant dense<0.000000e+00> : tensor<32x64xf16> loc(#loc1) + %cst_0 = arith.constant dense<0.000000e+00> : tensor<64x32xf16> loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %cst_1 = arith.constant dense<32> : tensor<64x32xi32> loc(#loc1) + %cst_2 = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc54) + %pid_n = tt.get_program_id y : i32 loc(#loc55) + %offs_m = arith.muli %pid_m, %c64_i32 : i32 loc(#loc56) + %offs_m_3 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc57) + %offs_m_4 = tt.splat %offs_m : i32 -> tensor<64xi32> loc(#loc58) + %offs_m_5 = arith.addi %offs_m_4, %offs_m_3 : tensor<64xi32> loc(#loc58) + %offs_n = arith.muli %pid_n, %c64_i32 : i32 loc(#loc59) + %offs_n_6 = tt.splat %offs_n : i32 -> tensor<64xi32> loc(#loc60) + %offs_n_7 = arith.addi %offs_n_6, %offs_m_3 : tensor<64xi32> loc(#loc60) + %offs_k = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc61) + %a_ptrs = tt.expand_dims %offs_m_5 {axis = 1 : i32} : tensor<64xi32> -> tensor<64x1xi32> loc(#loc62) + %a_ptrs_8 = tt.splat %stride_am : i32 -> tensor<64x1xi32> loc(#loc63) + %a_ptrs_9 = arith.muli %a_ptrs, %a_ptrs_8 : tensor<64x1xi32> loc(#loc63) + %a_ptrs_10 = tt.splat %a_ptr : !tt.ptr -> tensor<64x1x!tt.ptr> loc(#loc64) + %a_ptrs_11 = tt.addptr %a_ptrs_10, %a_ptrs_9 : tensor<64x1x!tt.ptr>, tensor<64x1xi32> loc(#loc64) + %a_ptrs_12 = tt.expand_dims %offs_k {axis = 0 : i32} : tensor<32xi32> -> tensor<1x32xi32> loc(#loc65) + %a_ptrs_13 = tt.broadcast %a_ptrs_11 : tensor<64x1x!tt.ptr> -> tensor<64x32x!tt.ptr> loc(#loc66) + %a_ptrs_14 = tt.broadcast %a_ptrs_12 : tensor<1x32xi32> -> tensor<64x32xi32> loc(#loc66) + %a_ptrs_15 = tt.addptr %a_ptrs_13, %a_ptrs_14 : tensor<64x32x!tt.ptr>, tensor<64x32xi32> loc(#loc66) + %b_ptrs = tt.expand_dims %offs_k {axis = 1 : i32} : tensor<32xi32> -> tensor<32x1xi32> loc(#loc67) + %b_ptrs_16 = tt.splat %stride_bk : i32 -> tensor<32x1xi32> loc(#loc68) + %b_ptrs_17 = arith.muli %b_ptrs, %b_ptrs_16 : tensor<32x1xi32> loc(#loc68) + %b_ptrs_18 = tt.splat %b_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc69) + %b_ptrs_19 = tt.addptr %b_ptrs_18, %b_ptrs_17 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc69) + %b_ptrs_20 = tt.expand_dims %offs_n_7 {axis = 0 : i32} : tensor<64xi32> -> tensor<1x64xi32> loc(#loc70) + %b_ptrs_21 = tt.broadcast %b_ptrs_19 : tensor<32x1x!tt.ptr> -> tensor<32x64x!tt.ptr> loc(#loc71) + %b_ptrs_22 = tt.broadcast %b_ptrs_20 : tensor<1x64xi32> -> tensor<32x64xi32> loc(#loc71) + %b_ptrs_23 = tt.addptr %b_ptrs_21, %b_ptrs_22 : tensor<32x64x!tt.ptr>, tensor<32x64xi32> loc(#loc71) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc72) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc73) + %acc:3 = scf.for %k = %c0_i32 to %1 step %c1_i32 iter_args(%a_ptrs_36 = %a_ptrs_15, %b_ptrs_37 = %b_ptrs_23, %acc_38 = %cst_2) -> (tensor<64x32x!tt.ptr>, tensor<32x64x!tt.ptr>, tensor<64x64xf32>) : i32 { + %a = arith.muli %k, %c32_i32 : i32 loc(#loc75) + %a_39 = arith.subi %K, %a : i32 loc(#loc76) + %a_40 = tt.splat %a_39 : i32 -> tensor<1x32xi32> loc(#loc77) + %a_41 = arith.cmpi slt, %a_ptrs_12, %a_40 : tensor<1x32xi32> loc(#loc77) + %a_42 = tt.broadcast %a_41 : tensor<1x32xi1> -> tensor<64x32xi1> loc(#loc78) + %a_43 = tt.load %a_ptrs_36, %a_42, %cst_0 : tensor<64x32x!tt.ptr> loc(#loc78) + %b = tt.splat %a_39 : i32 -> tensor<32x1xi32> loc(#loc79) + %b_44 = arith.cmpi slt, %b_ptrs, %b : tensor<32x1xi32> loc(#loc79) + %b_45 = tt.broadcast %b_44 : tensor<32x1xi1> -> tensor<32x64xi1> loc(#loc80) + %b_46 = tt.load %b_ptrs_37, %b_45, %cst : tensor<32x64x!tt.ptr> loc(#loc80) + %acc_47 = tt.dot %a_43, %b_46, %acc_38, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc81) + %a_ptrs_48 = tt.addptr %a_ptrs_36, %cst_1 : tensor<64x32x!tt.ptr>, tensor<64x32xi32> loc(#loc82) + %b_ptrs_49 = arith.muli %stride_bk, %c32_i32 : i32 loc(#loc83) + %b_ptrs_50 = tt.splat %b_ptrs_49 : i32 -> tensor<32x64xi32> loc(#loc84) + %b_ptrs_51 = tt.addptr %b_ptrs_37, %b_ptrs_50 : tensor<32x64x!tt.ptr>, tensor<32x64xi32> loc(#loc84) + scf.yield %a_ptrs_48, %b_ptrs_51, %acc_47 : tensor<64x32x!tt.ptr>, tensor<32x64x!tt.ptr>, tensor<64x64xf32> loc(#loc34) + } loc(#loc93) + %c = arith.truncf %acc#2 : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc85) + %c_ptrs = tt.splat %stride_cm : i32 -> tensor<64x1xi32> loc(#loc86) + %c_ptrs_24 = arith.muli %a_ptrs, %c_ptrs : tensor<64x1xi32> loc(#loc86) + %c_ptrs_25 = tt.splat %c_ptr : !tt.ptr -> tensor<64x1x!tt.ptr> loc(#loc87) + %c_ptrs_26 = tt.addptr %c_ptrs_25, %c_ptrs_24 : tensor<64x1x!tt.ptr>, tensor<64x1xi32> loc(#loc87) + %c_ptrs_27 = tt.broadcast %c_ptrs_26 : tensor<64x1x!tt.ptr> -> tensor<64x64x!tt.ptr> loc(#loc88) + %c_ptrs_28 = tt.broadcast %b_ptrs_20 : tensor<1x64xi32> -> tensor<64x64xi32> loc(#loc88) + %c_ptrs_29 = tt.addptr %c_ptrs_27, %c_ptrs_28 : tensor<64x64x!tt.ptr>, tensor<64x64xi32> loc(#loc88) + %c_mask = tt.splat %M : i32 -> tensor<64x1xi32> loc(#loc89) + %c_mask_30 = arith.cmpi slt, %a_ptrs, %c_mask : tensor<64x1xi32> loc(#loc89) + %c_mask_31 = tt.splat %N : i32 -> tensor<1x64xi32> loc(#loc90) + %c_mask_32 = arith.cmpi slt, %b_ptrs_20, %c_mask_31 : tensor<1x64xi32> loc(#loc90) + %c_mask_33 = tt.broadcast %c_mask_30 : tensor<64x1xi1> -> tensor<64x64xi1> loc(#loc91) + %c_mask_34 = tt.broadcast %c_mask_32 : tensor<1x64xi1> -> tensor<64x64xi1> loc(#loc91) + %c_mask_35 = arith.andi %c_mask_33, %c_mask_34 : tensor<64x64xi1> loc(#loc91) + tt.store %c_ptrs_29, %c, %c_mask_35 : tensor<64x64x!tt.ptr> loc(#loc42) + tt.return loc(#loc43) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":53:33) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":53:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":42:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":43:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:21) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:44) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:31) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":46:21) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":46:31) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":47:26) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:28) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:39) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:21) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:58) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:51) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:28) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:39) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:21) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:58) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:51) +#loc22 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:17) +#loc23 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:30) +#loc24 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:59) +#loc25 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:55) +#loc26 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:51) +#loc27 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:20) +#loc28 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":55:51) +#loc29 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":55:20) +#loc30 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":56:25) +#loc31 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":57:18) +#loc32 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:28) +#loc33 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:18) +#loc34 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:8) +#loc35 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":60:15) +#loc36 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:39) +#loc37 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:21) +#loc38 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:51) +#loc39 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:32) +#loc40 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:56) +#loc41 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:38) +#loc42 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":63:21) +#loc43 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":63:4) +#loc53 = loc(callsite(#loc1 at #loc2)) +#loc54 = loc("pid_m"(#loc4)) +#loc55 = loc("pid_n"(#loc5)) +#loc56 = loc("offs_m"(#loc6)) +#loc57 = loc("offs_m"(#loc7)) +#loc58 = loc("offs_m"(#loc8)) +#loc59 = loc("offs_n"(#loc9)) +#loc60 = loc("offs_n"(#loc10)) +#loc61 = loc("offs_k"(#loc11)) +#loc62 = loc("a_ptrs"(#loc12)) +#loc63 = loc("a_ptrs"(#loc13)) +#loc64 = loc("a_ptrs"(#loc14)) +#loc65 = loc("a_ptrs"(#loc15)) +#loc66 = loc("a_ptrs"(#loc16)) +#loc67 = loc("b_ptrs"(#loc17)) +#loc68 = loc("b_ptrs"(#loc18)) +#loc69 = loc("b_ptrs"(#loc19)) +#loc70 = loc("b_ptrs"(#loc20)) +#loc71 = loc("b_ptrs"(#loc21)) +#loc72 = loc(callsite(#loc22 at #loc2)) +#loc73 = loc(callsite(#loc23 at #loc2)) +#loc74 = loc("a_ptrs"(#loc3)) +#loc75 = loc("a"(#loc24)) +#loc76 = loc("a"(#loc25)) +#loc77 = loc("a"(#loc26)) +#loc78 = loc("a"(#loc27)) +#loc79 = loc("b"(#loc28)) +#loc80 = loc("b"(#loc29)) +#loc81 = loc("acc"(#loc30)) +#loc82 = loc("a_ptrs"(#loc31)) +#loc83 = loc("b_ptrs"(#loc32)) +#loc84 = loc("b_ptrs"(#loc33)) +#loc85 = loc("c"(#loc35)) +#loc86 = loc("c_ptrs"(#loc36)) +#loc87 = loc("c_ptrs"(#loc37)) +#loc88 = loc("c_ptrs"(#loc38)) +#loc89 = loc("c_mask"(#loc39)) +#loc90 = loc("c_mask"(#loc40)) +#loc91 = loc("c_mask"(#loc41)) +#loc92 = loc("b_ptrs"(#loc74)) +#loc93 = loc("acc"(#loc92)) diff --git a/tests/golden/ir/ttir/golden_matmul_s3_sm90.ttir b/tests/golden/ir/ttir/golden_matmul_s3_sm90.ttir new file mode 100644 index 000000000..72c5a1f80 --- /dev/null +++ b/tests/golden/ir/ttir/golden_matmul_s3_sm90.ttir @@ -0,0 +1,172 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":28:0) +#loc44 = loc("a_ptr"(#loc)) +#loc45 = loc("b_ptr"(#loc)) +#loc46 = loc("c_ptr"(#loc)) +#loc47 = loc("M"(#loc)) +#loc48 = loc("N"(#loc)) +#loc49 = loc("K"(#loc)) +#loc50 = loc("stride_am"(#loc)) +#loc51 = loc("stride_bk"(#loc)) +#loc52 = loc("stride_cm"(#loc)) +module { + tt.func public @matmul_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc)), %stride_am: i32 {tt.divisibility = 16 : i32} loc("stride_am"(#loc)), %stride_bk: i32 {tt.divisibility = 16 : i32} loc("stride_bk"(#loc)), %stride_cm: i32 {tt.divisibility = 16 : i32} loc("stride_cm"(#loc))) attributes {noinline = false} { + %c31_i32 = arith.constant 31 : i32 loc(#loc53) + %cst = arith.constant dense<0.000000e+00> : tensor<32x64xf16> loc(#loc1) + %cst_0 = arith.constant dense<0.000000e+00> : tensor<64x32xf16> loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %cst_1 = arith.constant dense<32> : tensor<64x32xi32> loc(#loc1) + %cst_2 = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc54) + %pid_n = tt.get_program_id y : i32 loc(#loc55) + %offs_m = arith.muli %pid_m, %c64_i32 : i32 loc(#loc56) + %offs_m_3 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc57) + %offs_m_4 = tt.splat %offs_m : i32 -> tensor<64xi32> loc(#loc58) + %offs_m_5 = arith.addi %offs_m_4, %offs_m_3 : tensor<64xi32> loc(#loc58) + %offs_n = arith.muli %pid_n, %c64_i32 : i32 loc(#loc59) + %offs_n_6 = tt.splat %offs_n : i32 -> tensor<64xi32> loc(#loc60) + %offs_n_7 = arith.addi %offs_n_6, %offs_m_3 : tensor<64xi32> loc(#loc60) + %offs_k = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc61) + %a_ptrs = tt.expand_dims %offs_m_5 {axis = 1 : i32} : tensor<64xi32> -> tensor<64x1xi32> loc(#loc62) + %a_ptrs_8 = tt.splat %stride_am : i32 -> tensor<64x1xi32> loc(#loc63) + %a_ptrs_9 = arith.muli %a_ptrs, %a_ptrs_8 : tensor<64x1xi32> loc(#loc63) + %a_ptrs_10 = tt.splat %a_ptr : !tt.ptr -> tensor<64x1x!tt.ptr> loc(#loc64) + %a_ptrs_11 = tt.addptr %a_ptrs_10, %a_ptrs_9 : tensor<64x1x!tt.ptr>, tensor<64x1xi32> loc(#loc64) + %a_ptrs_12 = tt.expand_dims %offs_k {axis = 0 : i32} : tensor<32xi32> -> tensor<1x32xi32> loc(#loc65) + %a_ptrs_13 = tt.broadcast %a_ptrs_11 : tensor<64x1x!tt.ptr> -> tensor<64x32x!tt.ptr> loc(#loc66) + %a_ptrs_14 = tt.broadcast %a_ptrs_12 : tensor<1x32xi32> -> tensor<64x32xi32> loc(#loc66) + %a_ptrs_15 = tt.addptr %a_ptrs_13, %a_ptrs_14 : tensor<64x32x!tt.ptr>, tensor<64x32xi32> loc(#loc66) + %b_ptrs = tt.expand_dims %offs_k {axis = 1 : i32} : tensor<32xi32> -> tensor<32x1xi32> loc(#loc67) + %b_ptrs_16 = tt.splat %stride_bk : i32 -> tensor<32x1xi32> loc(#loc68) + %b_ptrs_17 = arith.muli %b_ptrs, %b_ptrs_16 : tensor<32x1xi32> loc(#loc68) + %b_ptrs_18 = tt.splat %b_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc69) + %b_ptrs_19 = tt.addptr %b_ptrs_18, %b_ptrs_17 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc69) + %b_ptrs_20 = tt.expand_dims %offs_n_7 {axis = 0 : i32} : tensor<64xi32> -> tensor<1x64xi32> loc(#loc70) + %b_ptrs_21 = tt.broadcast %b_ptrs_19 : tensor<32x1x!tt.ptr> -> tensor<32x64x!tt.ptr> loc(#loc71) + %b_ptrs_22 = tt.broadcast %b_ptrs_20 : tensor<1x64xi32> -> tensor<32x64xi32> loc(#loc71) + %b_ptrs_23 = tt.addptr %b_ptrs_21, %b_ptrs_22 : tensor<32x64x!tt.ptr>, tensor<32x64xi32> loc(#loc71) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc72) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc73) + %acc:3 = scf.for %k = %c0_i32 to %1 step %c1_i32 iter_args(%a_ptrs_36 = %a_ptrs_15, %b_ptrs_37 = %b_ptrs_23, %acc_38 = %cst_2) -> (tensor<64x32x!tt.ptr>, tensor<32x64x!tt.ptr>, tensor<64x64xf32>) : i32 { + %a = arith.muli %k, %c32_i32 : i32 loc(#loc75) + %a_39 = arith.subi %K, %a : i32 loc(#loc76) + %a_40 = tt.splat %a_39 : i32 -> tensor<1x32xi32> loc(#loc77) + %a_41 = arith.cmpi slt, %a_ptrs_12, %a_40 : tensor<1x32xi32> loc(#loc77) + %a_42 = tt.broadcast %a_41 : tensor<1x32xi1> -> tensor<64x32xi1> loc(#loc78) + %a_43 = tt.load %a_ptrs_36, %a_42, %cst_0 : tensor<64x32x!tt.ptr> loc(#loc78) + %b = tt.splat %a_39 : i32 -> tensor<32x1xi32> loc(#loc79) + %b_44 = arith.cmpi slt, %b_ptrs, %b : tensor<32x1xi32> loc(#loc79) + %b_45 = tt.broadcast %b_44 : tensor<32x1xi1> -> tensor<32x64xi1> loc(#loc80) + %b_46 = tt.load %b_ptrs_37, %b_45, %cst : tensor<32x64x!tt.ptr> loc(#loc80) + %acc_47 = tt.dot %a_43, %b_46, %acc_38, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc81) + %a_ptrs_48 = tt.addptr %a_ptrs_36, %cst_1 : tensor<64x32x!tt.ptr>, tensor<64x32xi32> loc(#loc82) + %b_ptrs_49 = arith.muli %stride_bk, %c32_i32 : i32 loc(#loc83) + %b_ptrs_50 = tt.splat %b_ptrs_49 : i32 -> tensor<32x64xi32> loc(#loc84) + %b_ptrs_51 = tt.addptr %b_ptrs_37, %b_ptrs_50 : tensor<32x64x!tt.ptr>, tensor<32x64xi32> loc(#loc84) + scf.yield %a_ptrs_48, %b_ptrs_51, %acc_47 : tensor<64x32x!tt.ptr>, tensor<32x64x!tt.ptr>, tensor<64x64xf32> loc(#loc34) + } loc(#loc93) + %c = arith.truncf %acc#2 : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc85) + %c_ptrs = tt.splat %stride_cm : i32 -> tensor<64x1xi32> loc(#loc86) + %c_ptrs_24 = arith.muli %a_ptrs, %c_ptrs : tensor<64x1xi32> loc(#loc86) + %c_ptrs_25 = tt.splat %c_ptr : !tt.ptr -> tensor<64x1x!tt.ptr> loc(#loc87) + %c_ptrs_26 = tt.addptr %c_ptrs_25, %c_ptrs_24 : tensor<64x1x!tt.ptr>, tensor<64x1xi32> loc(#loc87) + %c_ptrs_27 = tt.broadcast %c_ptrs_26 : tensor<64x1x!tt.ptr> -> tensor<64x64x!tt.ptr> loc(#loc88) + %c_ptrs_28 = tt.broadcast %b_ptrs_20 : tensor<1x64xi32> -> tensor<64x64xi32> loc(#loc88) + %c_ptrs_29 = tt.addptr %c_ptrs_27, %c_ptrs_28 : tensor<64x64x!tt.ptr>, tensor<64x64xi32> loc(#loc88) + %c_mask = tt.splat %M : i32 -> tensor<64x1xi32> loc(#loc89) + %c_mask_30 = arith.cmpi slt, %a_ptrs, %c_mask : tensor<64x1xi32> loc(#loc89) + %c_mask_31 = tt.splat %N : i32 -> tensor<1x64xi32> loc(#loc90) + %c_mask_32 = arith.cmpi slt, %b_ptrs_20, %c_mask_31 : tensor<1x64xi32> loc(#loc90) + %c_mask_33 = tt.broadcast %c_mask_30 : tensor<64x1xi1> -> tensor<64x64xi1> loc(#loc91) + %c_mask_34 = tt.broadcast %c_mask_32 : tensor<1x64xi1> -> tensor<64x64xi1> loc(#loc91) + %c_mask_35 = arith.andi %c_mask_33, %c_mask_34 : tensor<64x64xi1> loc(#loc91) + tt.store %c_ptrs_29, %c, %c_mask_35 : tensor<64x64x!tt.ptr> loc(#loc42) + tt.return loc(#loc43) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":53:33) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":53:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":42:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":43:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:21) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:44) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":45:31) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":46:21) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":46:31) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":47:26) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:28) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:39) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:21) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:58) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":49:51) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:28) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:39) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:21) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:58) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":50:51) +#loc22 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:17) +#loc23 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:30) +#loc24 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:59) +#loc25 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:55) +#loc26 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:51) +#loc27 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":54:20) +#loc28 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":55:51) +#loc29 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":55:20) +#loc30 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":56:25) +#loc31 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":57:18) +#loc32 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:28) +#loc33 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:18) +#loc34 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":58:8) +#loc35 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":60:15) +#loc36 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:39) +#loc37 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:21) +#loc38 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":61:51) +#loc39 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:32) +#loc40 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:56) +#loc41 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":62:38) +#loc42 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":63:21) +#loc43 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":63:4) +#loc53 = loc(callsite(#loc1 at #loc2)) +#loc54 = loc("pid_m"(#loc4)) +#loc55 = loc("pid_n"(#loc5)) +#loc56 = loc("offs_m"(#loc6)) +#loc57 = loc("offs_m"(#loc7)) +#loc58 = loc("offs_m"(#loc8)) +#loc59 = loc("offs_n"(#loc9)) +#loc60 = loc("offs_n"(#loc10)) +#loc61 = loc("offs_k"(#loc11)) +#loc62 = loc("a_ptrs"(#loc12)) +#loc63 = loc("a_ptrs"(#loc13)) +#loc64 = loc("a_ptrs"(#loc14)) +#loc65 = loc("a_ptrs"(#loc15)) +#loc66 = loc("a_ptrs"(#loc16)) +#loc67 = loc("b_ptrs"(#loc17)) +#loc68 = loc("b_ptrs"(#loc18)) +#loc69 = loc("b_ptrs"(#loc19)) +#loc70 = loc("b_ptrs"(#loc20)) +#loc71 = loc("b_ptrs"(#loc21)) +#loc72 = loc(callsite(#loc22 at #loc2)) +#loc73 = loc(callsite(#loc23 at #loc2)) +#loc74 = loc("a_ptrs"(#loc3)) +#loc75 = loc("a"(#loc24)) +#loc76 = loc("a"(#loc25)) +#loc77 = loc("a"(#loc26)) +#loc78 = loc("a"(#loc27)) +#loc79 = loc("b"(#loc28)) +#loc80 = loc("b"(#loc29)) +#loc81 = loc("acc"(#loc30)) +#loc82 = loc("a_ptrs"(#loc31)) +#loc83 = loc("b_ptrs"(#loc32)) +#loc84 = loc("b_ptrs"(#loc33)) +#loc85 = loc("c"(#loc35)) +#loc86 = loc("c_ptrs"(#loc36)) +#loc87 = loc("c_ptrs"(#loc37)) +#loc88 = loc("c_ptrs"(#loc38)) +#loc89 = loc("c_mask"(#loc39)) +#loc90 = loc("c_mask"(#loc40)) +#loc91 = loc("c_mask"(#loc41)) +#loc92 = loc("b_ptrs"(#loc74)) +#loc93 = loc("acc"(#loc92)) diff --git a/tests/golden/ir/ttir/golden_matmul_tma_s1_sm90.ttir b/tests/golden/ir/ttir/golden_matmul_tma_s1_sm90.ttir new file mode 100644 index 000000000..33837338b --- /dev/null +++ b/tests/golden/ir/ttir/golden_matmul_tma_s1_sm90.ttir @@ -0,0 +1,78 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":108:0) +#loc23 = loc("a_ptr"(#loc)) +#loc24 = loc("b_ptr"(#loc)) +#loc25 = loc("c_ptr"(#loc)) +#loc26 = loc("M"(#loc)) +#loc27 = loc("N"(#loc)) +#loc28 = loc("K"(#loc)) +module { + tt.func public @matmul_tma_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc))) attributes {noinline = false} { + %c31_i32 = arith.constant 31 : i32 loc(#loc29) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %cst = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %c1_i64 = arith.constant 1 : i64 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc30) + %pid_n = tt.get_program_id y : i32 loc(#loc31) + %a_desc = arith.extsi %K : i32 to i64 loc(#loc32) + %a_desc_0 = tt.make_tensor_descriptor %a_ptr, [%M, %K], [%a_desc, %c1_i64] : , > loc(#loc32) + %b_desc = arith.extsi %N : i32 to i64 loc(#loc33) + %b_desc_1 = tt.make_tensor_descriptor %b_ptr, [%K, %N], [%b_desc, %c1_i64] : , > loc(#loc33) + %c_desc = tt.make_tensor_descriptor %c_ptr, [%M, %N], [%b_desc, %c1_i64] : , > loc(#loc34) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc35) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc36) + %acc = scf.for %k = %c0_i32 to %1 step %c1_i32 iter_args(%acc_2 = %cst) -> (tensor<64x64xf32>) : i32 { + %a = arith.muli %pid_m, %c64_i32 : i32 loc(#loc38) + %a_3 = arith.muli %k, %c32_i32 : i32 loc(#loc39) + %a_4 = tt.descriptor_load %a_desc_0[%a, %a_3] : !tt.tensordesc> -> tensor<64x32xf16> loc(#loc40) + %b = arith.muli %pid_n, %c64_i32 : i32 loc(#loc41) + %b_5 = tt.descriptor_load %b_desc_1[%a_3, %b] : !tt.tensordesc> -> tensor<32x64xf16> loc(#loc42) + %acc_6 = tt.dot %a_4, %b_5, %acc_2, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc43) + scf.yield %acc_6 : tensor<64x64xf32> loc(#loc17) + } loc(#loc37) + %2 = arith.muli %pid_m, %c64_i32 : i32 loc(#loc18) + %3 = arith.muli %pid_n, %c64_i32 : i32 loc(#loc19) + %4 = arith.truncf %acc : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc20) + tt.descriptor_store %c_desc[%2, %3], %4 : !tt.tensordesc>, tensor<64x64xf16> loc(#loc21) + tt.return loc(#loc22) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":136:33) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":136:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":124:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":127:8) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":130:8) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":133:8) +#loc9 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:17) +#loc10 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:30) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:33) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:46) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:24) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:46) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:24) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":139:25) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":139:8) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:26) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:43) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:60) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:53) +#loc22 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:4) +#loc29 = loc(callsite(#loc1 at #loc2)) +#loc30 = loc("pid_m"(#loc4)) +#loc31 = loc("pid_n"(#loc5)) +#loc32 = loc("a_desc"(#loc6)) +#loc33 = loc("b_desc"(#loc7)) +#loc34 = loc("c_desc"(#loc8)) +#loc35 = loc(callsite(#loc9 at #loc2)) +#loc36 = loc(callsite(#loc10 at #loc2)) +#loc37 = loc("acc"(#loc3)) +#loc38 = loc("a"(#loc11)) +#loc39 = loc("a"(#loc12)) +#loc40 = loc("a"(#loc13)) +#loc41 = loc("b"(#loc14)) +#loc42 = loc("b"(#loc15)) +#loc43 = loc("acc"(#loc16)) diff --git a/tests/golden/ir/ttir/golden_matmul_tma_s3_sm90.ttir b/tests/golden/ir/ttir/golden_matmul_tma_s3_sm90.ttir new file mode 100644 index 000000000..33837338b --- /dev/null +++ b/tests/golden/ir/ttir/golden_matmul_tma_s3_sm90.ttir @@ -0,0 +1,78 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":108:0) +#loc23 = loc("a_ptr"(#loc)) +#loc24 = loc("b_ptr"(#loc)) +#loc25 = loc("c_ptr"(#loc)) +#loc26 = loc("M"(#loc)) +#loc27 = loc("N"(#loc)) +#loc28 = loc("K"(#loc)) +module { + tt.func public @matmul_tma_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc))) attributes {noinline = false} { + %c31_i32 = arith.constant 31 : i32 loc(#loc29) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %cst = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %c1_i64 = arith.constant 1 : i64 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc30) + %pid_n = tt.get_program_id y : i32 loc(#loc31) + %a_desc = arith.extsi %K : i32 to i64 loc(#loc32) + %a_desc_0 = tt.make_tensor_descriptor %a_ptr, [%M, %K], [%a_desc, %c1_i64] : , > loc(#loc32) + %b_desc = arith.extsi %N : i32 to i64 loc(#loc33) + %b_desc_1 = tt.make_tensor_descriptor %b_ptr, [%K, %N], [%b_desc, %c1_i64] : , > loc(#loc33) + %c_desc = tt.make_tensor_descriptor %c_ptr, [%M, %N], [%b_desc, %c1_i64] : , > loc(#loc34) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc35) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc36) + %acc = scf.for %k = %c0_i32 to %1 step %c1_i32 iter_args(%acc_2 = %cst) -> (tensor<64x64xf32>) : i32 { + %a = arith.muli %pid_m, %c64_i32 : i32 loc(#loc38) + %a_3 = arith.muli %k, %c32_i32 : i32 loc(#loc39) + %a_4 = tt.descriptor_load %a_desc_0[%a, %a_3] : !tt.tensordesc> -> tensor<64x32xf16> loc(#loc40) + %b = arith.muli %pid_n, %c64_i32 : i32 loc(#loc41) + %b_5 = tt.descriptor_load %b_desc_1[%a_3, %b] : !tt.tensordesc> -> tensor<32x64xf16> loc(#loc42) + %acc_6 = tt.dot %a_4, %b_5, %acc_2, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc43) + scf.yield %acc_6 : tensor<64x64xf32> loc(#loc17) + } loc(#loc37) + %2 = arith.muli %pid_m, %c64_i32 : i32 loc(#loc18) + %3 = arith.muli %pid_n, %c64_i32 : i32 loc(#loc19) + %4 = arith.truncf %acc : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc20) + tt.descriptor_store %c_desc[%2, %3], %4 : !tt.tensordesc>, tensor<64x64xf16> loc(#loc21) + tt.return loc(#loc22) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":136:33) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":136:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":124:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":127:8) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":130:8) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":133:8) +#loc9 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:17) +#loc10 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:30) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:33) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:46) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":137:24) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:46) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":138:24) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":139:25) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":139:8) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:26) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:43) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:60) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:53) +#loc22 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":140:4) +#loc29 = loc(callsite(#loc1 at #loc2)) +#loc30 = loc("pid_m"(#loc4)) +#loc31 = loc("pid_n"(#loc5)) +#loc32 = loc("a_desc"(#loc6)) +#loc33 = loc("b_desc"(#loc7)) +#loc34 = loc("c_desc"(#loc8)) +#loc35 = loc(callsite(#loc9 at #loc2)) +#loc36 = loc(callsite(#loc10 at #loc2)) +#loc37 = loc("acc"(#loc3)) +#loc38 = loc("a"(#loc11)) +#loc39 = loc("a"(#loc12)) +#loc40 = loc("a"(#loc13)) +#loc41 = loc("b"(#loc14)) +#loc42 = loc("b"(#loc15)) +#loc43 = loc("acc"(#loc16)) diff --git a/tests/golden/ir/ttir/golden_matmul_tma_ws_s3_sm90.ttir b/tests/golden/ir/ttir/golden_matmul_tma_ws_s3_sm90.ttir new file mode 100644 index 000000000..396b89f31 --- /dev/null +++ b/tests/golden/ir/ttir/golden_matmul_tma_ws_s3_sm90.ttir @@ -0,0 +1,78 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":144:0) +#loc23 = loc("a_ptr"(#loc)) +#loc24 = loc("b_ptr"(#loc)) +#loc25 = loc("c_ptr"(#loc)) +#loc26 = loc("M"(#loc)) +#loc27 = loc("N"(#loc)) +#loc28 = loc("K"(#loc)) +module { + tt.func public @matmul_tma_ws_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc))) attributes {noinline = false} { + %c31_i32 = arith.constant 31 : i32 loc(#loc29) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %cst = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %c1_i64 = arith.constant 1 : i64 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc30) + %pid_n = tt.get_program_id y : i32 loc(#loc31) + %a_desc = arith.extsi %K : i32 to i64 loc(#loc32) + %a_desc_0 = tt.make_tensor_descriptor %a_ptr, [%M, %K], [%a_desc, %c1_i64] : , > loc(#loc32) + %b_desc = arith.extsi %N : i32 to i64 loc(#loc33) + %b_desc_1 = tt.make_tensor_descriptor %b_ptr, [%K, %N], [%b_desc, %c1_i64] : , > loc(#loc33) + %c_desc = tt.make_tensor_descriptor %c_ptr, [%M, %N], [%b_desc, %c1_i64] : , > loc(#loc34) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc35) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc36) + %acc = scf.for %k = %c0_i32 to %1 step %c1_i32 iter_args(%acc_2 = %cst) -> (tensor<64x64xf32>) : i32 { + %a = arith.muli %pid_m, %c64_i32 : i32 loc(#loc38) + %a_3 = arith.muli %k, %c32_i32 : i32 loc(#loc39) + %a_4 = tt.descriptor_load %a_desc_0[%a, %a_3] : !tt.tensordesc> -> tensor<64x32xf16> loc(#loc40) + %b = arith.muli %pid_n, %c64_i32 : i32 loc(#loc41) + %b_5 = tt.descriptor_load %b_desc_1[%a_3, %b] : !tt.tensordesc> -> tensor<32x64xf16> loc(#loc42) + %acc_6 = tt.dot %a_4, %b_5, %acc_2, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc43) + scf.yield %acc_6 : tensor<64x64xf32> loc(#loc17) + } {tt.warp_specialize} loc(#loc37) + %2 = arith.muli %pid_m, %c64_i32 : i32 loc(#loc18) + %3 = arith.muli %pid_n, %c64_i32 : i32 loc(#loc19) + %4 = arith.truncf %acc : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc20) + tt.descriptor_store %c_desc[%2, %3], %4 : !tt.tensordesc>, tensor<64x64xf16> loc(#loc21) + tt.return loc(#loc22) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":172:36) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":172:46) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":160:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":161:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":163:8) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":166:8) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":169:8) +#loc9 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:17) +#loc10 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":43:30) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":173:33) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":173:46) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":173:24) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":174:46) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":174:24) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":175:25) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":175:8) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":176:26) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":176:43) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":176:60) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":176:53) +#loc22 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":176:4) +#loc29 = loc(callsite(#loc1 at #loc2)) +#loc30 = loc("pid_m"(#loc4)) +#loc31 = loc("pid_n"(#loc5)) +#loc32 = loc("a_desc"(#loc6)) +#loc33 = loc("b_desc"(#loc7)) +#loc34 = loc("c_desc"(#loc8)) +#loc35 = loc(callsite(#loc9 at #loc2)) +#loc36 = loc(callsite(#loc10 at #loc2)) +#loc37 = loc("acc"(#loc3)) +#loc38 = loc("a"(#loc11)) +#loc39 = loc("a"(#loc12)) +#loc40 = loc("a"(#loc13)) +#loc41 = loc("b"(#loc14)) +#loc42 = loc("b"(#loc15)) +#loc43 = loc("acc"(#loc16)) diff --git a/tests/golden/ir/ttir/golden_nested_guard_merge_sm80.ttir b/tests/golden/ir/ttir/golden_nested_guard_merge_sm80.ttir new file mode 100644 index 000000000..fd4f5b094 --- /dev/null +++ b/tests/golden/ir/ttir/golden_nested_guard_merge_sm80.ttir @@ -0,0 +1,59 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":354:0) +#loc1 = loc(unknown) +#loc16 = loc("x_ptr"(#loc)) +#loc17 = loc("out_ptr"(#loc)) +#loc18 = loc("n"(#loc)) +#loc19 = loc("T"(#loc)) +module { + tt.func public @nested_guard_merge_kernel(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc)), %T: i32 loc("T"(#loc))) attributes {noinline = false} { + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc20) + %base = arith.muli %pid, %c64_i32 : i32 loc(#loc21) + %0 = arith.cmpi sge, %pid, %T : i32 loc(#loc4) + cf.cond_br %0, ^bb1, ^bb2 loc(#loc4) + ^bb1: // 2 preds: ^bb0, ^bb3 + tt.return loc(#loc5) + ^bb2: // pred: ^bb0 + %1 = arith.cmpi eq, %pid, %c0_i32 : i32 loc(#loc6) + cf.cond_br %1, ^bb3, ^bb4 loc(#loc6) + ^bb3: // pred: ^bb2 + %2 = arith.cmpi slt, %n, %c0_i32 : i32 loc(#loc7) + cf.cond_br %2, ^bb1, ^bb5(%c0_i32 : i32) loc(#loc7) + ^bb4: // pred: ^bb2 + %base_0 = arith.addi %base, %n : i32 loc(#loc22) + cf.br ^bb5(%base_0 : i32) loc(#loc22) + ^bb5(%3: i32 loc(unknown)): // 2 preds: ^bb3, ^bb4 + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc23) + %offs_1 = tt.splat %3 : i32 -> tensor<64xi32> loc(#loc24) + %offs_2 = arith.addi %offs_1, %offs : tensor<64xi32> loc(#loc24) + %v = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc25) + %v_3 = tt.addptr %v, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc25) + %v_4 = tt.load %v_3 : tensor<64x!tt.ptr> loc(#loc26) + %4 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc13) + %5 = tt.addptr %4, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc13) + tt.store %5, %v_4 : tensor<64x!tt.ptr> loc(#loc14) + tt.return loc(#loc15) + } loc(#loc) +} loc(#loc) +#loc2 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":355:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":356:17) +#loc4 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":357:14) +#loc5 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":358:8) +#loc6 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":359:14) +#loc7 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":361:15) +#loc8 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":364:22) +#loc9 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":365:31) +#loc10 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":365:18) +#loc11 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":366:24) +#loc12 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":366:16) +#loc13 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":367:23) +#loc14 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":367:29) +#loc15 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":367:4) +#loc20 = loc("pid"(#loc2)) +#loc21 = loc("base"(#loc3)) +#loc22 = loc("base"(#loc8)) +#loc23 = loc("offs"(#loc9)) +#loc24 = loc("offs"(#loc10)) +#loc25 = loc("v"(#loc11)) +#loc26 = loc("v"(#loc12)) diff --git a/tests/golden/ir/ttir/golden_nested_loops_sm80.ttir b/tests/golden/ir/ttir/golden_nested_loops_sm80.ttir new file mode 100644 index 000000000..5cabeb410 --- /dev/null +++ b/tests/golden/ir/ttir/golden_nested_loops_sm80.ttir @@ -0,0 +1,45 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":382:0) +#loc14 = loc("x_ptr"(#loc)) +#loc15 = loc("out_ptr"(#loc)) +#loc16 = loc("n"(#loc)) +#loc17 = loc("m"(#loc)) +module { + tt.func public @nested_loops_kernel(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc)), %m: i32 loc("m"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc18) + scf.for %i = %c0_i32 to %n step %c1_i32 : i32 { + scf.for %j = %c0_i32 to %m step %c1_i32 : i32 { + %offs = arith.muli %pid, %n : i32 loc(#loc19) + %offs_0 = arith.addi %offs, %i : i32 loc(#loc20) + %offs_1 = arith.muli %offs_0, %m : i32 loc(#loc21) + %offs_2 = arith.addi %offs_1, %j : i32 loc(#loc22) + %v = tt.addptr %x_ptr, %offs_2 : !tt.ptr, i32 loc(#loc23) + %v_3 = tt.load %v : !tt.ptr loc(#loc24) + %0 = tt.addptr %out_ptr, %offs_2 : !tt.ptr, i32 loc(#loc11) + tt.store %0, %v_3 : !tt.ptr loc(#loc12) + } loc(#loc4) + } loc(#loc3) + tt.return loc(#loc13) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":383:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":384:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":385:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":386:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":386:30) +#loc7 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":386:35) +#loc8 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":386:39) +#loc9 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":387:32) +#loc10 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":387:24) +#loc11 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":388:31) +#loc12 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":388:37) +#loc13 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":384:4) +#loc18 = loc("pid"(#loc2)) +#loc19 = loc("offs"(#loc5)) +#loc20 = loc("offs"(#loc6)) +#loc21 = loc("offs"(#loc7)) +#loc22 = loc("offs"(#loc8)) +#loc23 = loc("v"(#loc9)) +#loc24 = loc("v"(#loc10)) diff --git a/tests/golden/ir/ttir/golden_pid_branch_sm80.ttir b/tests/golden/ir/ttir/golden_pid_branch_sm80.ttir new file mode 100644 index 000000000..8f80ade64 --- /dev/null +++ b/tests/golden/ir/ttir/golden_pid_branch_sm80.ttir @@ -0,0 +1,47 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":166:0) +#loc14 = loc("x_ptr"(#loc)) +#loc15 = loc("out_ptr"(#loc)) +#loc16 = loc("n_elements"(#loc)) +module { + tt.func public @pid_branch_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c256_i32 = arith.constant 256 : i32 loc(#loc2) + %pid = tt.get_program_id x : i32 loc(#loc17) + %offs = arith.muli %pid, %c256_i32 : i32 loc(#loc18) + %offs_0 = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc19) + %offs_1 = tt.splat %offs : i32 -> tensor<256xi32> loc(#loc20) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<256xi32> loc(#loc20) + %mask = tt.splat %n_elements : i32 -> tensor<256xi32> loc(#loc21) + %mask_3 = arith.cmpi slt, %offs_2, %mask : tensor<256xi32> loc(#loc21) + %v = tt.splat %x_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc22) + %v_4 = tt.addptr %v, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc22) + %v_5 = tt.load %v_4, %mask_3 : tensor<256x!tt.ptr> loc(#loc23) + %0 = arith.cmpi eq, %pid, %c0_i32 : i32 loc(#loc1) + scf.if %0 { + %1 = tt.splat %out_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc11) + %2 = tt.addptr %1, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc11) + tt.store %2, %v_5, %mask_3 : tensor<256x!tt.ptr> loc(#loc12) + } loc(#loc10) + tt.return loc(#loc13) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":173:14) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":169:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":170:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":170:38) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":170:25) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":171:18) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":172:24) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":172:16) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":173:7) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":174:27) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":174:33) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":173:4) +#loc17 = loc("pid"(#loc3)) +#loc18 = loc("offs"(#loc4)) +#loc19 = loc("offs"(#loc5)) +#loc20 = loc("offs"(#loc6)) +#loc21 = loc("mask"(#loc7)) +#loc22 = loc("v"(#loc8)) +#loc23 = loc("v"(#loc9)) diff --git a/tests/golden/ir/ttir/golden_pid_branch_sm90.ttir b/tests/golden/ir/ttir/golden_pid_branch_sm90.ttir new file mode 100644 index 000000000..8f80ade64 --- /dev/null +++ b/tests/golden/ir/ttir/golden_pid_branch_sm90.ttir @@ -0,0 +1,47 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":166:0) +#loc14 = loc("x_ptr"(#loc)) +#loc15 = loc("out_ptr"(#loc)) +#loc16 = loc("n_elements"(#loc)) +module { + tt.func public @pid_branch_kernel(%x_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("x_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %n_elements: i32 {tt.divisibility = 16 : i32} loc("n_elements"(#loc))) attributes {noinline = false} { + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c256_i32 = arith.constant 256 : i32 loc(#loc2) + %pid = tt.get_program_id x : i32 loc(#loc17) + %offs = arith.muli %pid, %c256_i32 : i32 loc(#loc18) + %offs_0 = tt.make_range {end = 256 : i32, start = 0 : i32} : tensor<256xi32> loc(#loc19) + %offs_1 = tt.splat %offs : i32 -> tensor<256xi32> loc(#loc20) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<256xi32> loc(#loc20) + %mask = tt.splat %n_elements : i32 -> tensor<256xi32> loc(#loc21) + %mask_3 = arith.cmpi slt, %offs_2, %mask : tensor<256xi32> loc(#loc21) + %v = tt.splat %x_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc22) + %v_4 = tt.addptr %v, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc22) + %v_5 = tt.load %v_4, %mask_3 : tensor<256x!tt.ptr> loc(#loc23) + %0 = arith.cmpi eq, %pid, %c0_i32 : i32 loc(#loc1) + scf.if %0 { + %1 = tt.splat %out_ptr : !tt.ptr -> tensor<256x!tt.ptr> loc(#loc11) + %2 = tt.addptr %1, %offs_2 : tensor<256x!tt.ptr>, tensor<256xi32> loc(#loc11) + tt.store %2, %v_5, %mask_3 : tensor<256x!tt.ptr> loc(#loc12) + } loc(#loc10) + tt.return loc(#loc13) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":173:14) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":169:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":170:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":170:38) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":170:25) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":171:18) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":172:24) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":172:16) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":173:7) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":174:27) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":174:33) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":173:4) +#loc17 = loc("pid"(#loc3)) +#loc18 = loc("offs"(#loc4)) +#loc19 = loc("offs"(#loc5)) +#loc20 = loc("offs"(#loc6)) +#loc21 = loc("mask"(#loc7)) +#loc22 = loc("v"(#loc8)) +#loc23 = loc("v"(#loc9)) diff --git a/tests/golden/ir/ttir/golden_sequential_loops_sm80.ttir b/tests/golden/ir/ttir/golden_sequential_loops_sm80.ttir new file mode 100644 index 000000000..006c3bba1 --- /dev/null +++ b/tests/golden/ir/ttir/golden_sequential_loops_sm80.ttir @@ -0,0 +1,68 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":392:0) +#loc21 = loc("x_ptr"(#loc)) +#loc22 = loc("out_ptr"(#loc)) +#loc23 = loc("n"(#loc)) +module { + tt.func public @sequential_loops_kernel(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %acc = arith.constant dense<0.000000e+00> : tensor<64xf32> loc(#loc35) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %c64_i32 = arith.constant 64 : i32 loc(#loc3) + %pid = tt.get_program_id x : i32 loc(#loc25) + %offs = arith.muli %pid, %c64_i32 : i32 loc(#loc26) + %offs_0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc27) + %offs_1 = tt.splat %offs : i32 -> tensor<64xi32> loc(#loc28) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<64xi32> loc(#loc28) + %acc_3 = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%acc_4 = %acc) -> (tensor<64xf32>) : i32 { + %acc_5 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc30) + %acc_6 = tt.addptr %acc_5, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc30) + %acc_7 = arith.muli %i, %c64_i32 : i32 loc(#loc31) + %acc_8 = tt.splat %acc_7 : i32 -> tensor<64xi32> loc(#loc32) + %acc_9 = tt.addptr %acc_6, %acc_8 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc32) + %acc_10 = tt.load %acc_9 : tensor<64x!tt.ptr> loc(#loc33) + %acc_11 = arith.addf %acc_4, %acc_10 : tensor<64xf32> loc(#loc34) + scf.yield %acc_11 : tensor<64xf32> loc(#loc14) + } loc(#loc29) + scf.for %j = %c0_i32 to %n step %c1_i32 : i32 { + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc16) + %1 = tt.addptr %0, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc16) + %2 = arith.muli %j, %c64_i32 : i32 loc(#loc17) + %3 = tt.splat %2 : i32 -> tensor<64xi32> loc(#loc18) + %4 = tt.addptr %1, %3 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc18) + tt.store %4, %acc_3 : tensor<64x!tt.ptr> loc(#loc19) + } loc(#loc15) + tt.return loc(#loc20) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":129:31) +#loc2 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":395:19) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":393:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":394:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":394:38) +#loc7 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":394:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":396:22) +#loc9 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":397:31) +#loc10 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":397:42) +#loc11 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":397:38) +#loc12 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":397:23) +#loc13 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":397:15) +#loc14 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":397:8) +#loc15 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":398:22) +#loc16 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":399:27) +#loc17 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":399:38) +#loc18 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":399:34) +#loc19 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":399:45) +#loc20 = loc("/home/hwu27/workspace/triton-viz-route3/tests/golden/ttgir/generate_golden.py":398:4) +#loc24 = loc("acc"(#loc2)) +#loc25 = loc("pid"(#loc4)) +#loc26 = loc("offs"(#loc5)) +#loc27 = loc("offs"(#loc6)) +#loc28 = loc("offs"(#loc7)) +#loc29 = loc("acc"(#loc8)) +#loc30 = loc("acc"(#loc9)) +#loc31 = loc("acc"(#loc10)) +#loc32 = loc("acc"(#loc11)) +#loc33 = loc("acc"(#loc12)) +#loc34 = loc("acc"(#loc13)) +#loc35 = loc(callsite(#loc1 at #loc24)) diff --git a/tests/golden/ir/ttir/golden_tile2d_sm80.ttir b/tests/golden/ir/ttir/golden_tile2d_sm80.ttir new file mode 100644 index 000000000..87a63ffc1 --- /dev/null +++ b/tests/golden/ir/ttir/golden_tile2d_sm80.ttir @@ -0,0 +1,91 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":118:0) +#loc24 = loc("in_ptr"(#loc)) +#loc25 = loc("out_ptr"(#loc)) +#loc26 = loc("M"(#loc)) +#loc27 = loc("N"(#loc)) +#loc28 = loc("stride_m"(#loc)) +#loc29 = loc("stride_n"(#loc)) +module { + tt.func public @tile2d_kernel(%in_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("in_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %stride_m: i32 {tt.divisibility = 16 : i32} loc("stride_m"(#loc)), %stride_n: i32 {tt.divisibility = 16 : i32} loc("stride_n"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<2.000000e+00> : tensor<32x32xf32> loc(#loc1) + %vals = arith.constant dense<0.000000e+00> : tensor<32x32xf32> loc(#loc30) + %c32_i32 = arith.constant 32 : i32 loc(#loc3) + %pid_m = tt.get_program_id x : i32 loc(#loc31) + %pid_n = tt.get_program_id y : i32 loc(#loc32) + %offs_m = arith.muli %pid_m, %c32_i32 : i32 loc(#loc33) + %offs_m_0 = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc34) + %offs_m_1 = tt.splat %offs_m : i32 -> tensor<32xi32> loc(#loc35) + %offs_m_2 = arith.addi %offs_m_1, %offs_m_0 : tensor<32xi32> loc(#loc35) + %offs_n = arith.muli %pid_n, %c32_i32 : i32 loc(#loc36) + %offs_n_3 = tt.splat %offs_n : i32 -> tensor<32xi32> loc(#loc37) + %offs_n_4 = arith.addi %offs_n_3, %offs_m_0 : tensor<32xi32> loc(#loc37) + %ptrs = tt.expand_dims %offs_m_2 {axis = 1 : i32} : tensor<32xi32> -> tensor<32x1xi32> loc(#loc38) + %ptrs_5 = tt.splat %stride_m : i32 -> tensor<32x1xi32> loc(#loc39) + %ptrs_6 = arith.muli %ptrs, %ptrs_5 : tensor<32x1xi32> loc(#loc39) + %ptrs_7 = tt.splat %in_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc40) + %ptrs_8 = tt.addptr %ptrs_7, %ptrs_6 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc40) + %ptrs_9 = tt.expand_dims %offs_n_4 {axis = 0 : i32} : tensor<32xi32> -> tensor<1x32xi32> loc(#loc41) + %ptrs_10 = tt.splat %stride_n : i32 -> tensor<1x32xi32> loc(#loc42) + %ptrs_11 = arith.muli %ptrs_9, %ptrs_10 : tensor<1x32xi32> loc(#loc42) + %ptrs_12 = tt.broadcast %ptrs_8 : tensor<32x1x!tt.ptr> -> tensor<32x32x!tt.ptr> loc(#loc43) + %ptrs_13 = tt.broadcast %ptrs_11 : tensor<1x32xi32> -> tensor<32x32xi32> loc(#loc43) + %ptrs_14 = tt.addptr %ptrs_12, %ptrs_13 : tensor<32x32x!tt.ptr>, tensor<32x32xi32> loc(#loc43) + %mask = tt.splat %M : i32 -> tensor<32x1xi32> loc(#loc44) + %mask_15 = arith.cmpi slt, %ptrs, %mask : tensor<32x1xi32> loc(#loc44) + %mask_16 = tt.splat %N : i32 -> tensor<1x32xi32> loc(#loc45) + %mask_17 = arith.cmpi slt, %ptrs_9, %mask_16 : tensor<1x32xi32> loc(#loc45) + %mask_18 = tt.broadcast %mask_15 : tensor<32x1xi1> -> tensor<32x32xi1> loc(#loc46) + %mask_19 = tt.broadcast %mask_17 : tensor<1x32xi1> -> tensor<32x32xi1> loc(#loc46) + %mask_20 = arith.andi %mask_18, %mask_19 : tensor<32x32xi1> loc(#loc46) + %vals_21 = tt.load %ptrs_14, %mask_20, %vals : tensor<32x32x!tt.ptr> loc(#loc30) + %optrs = tt.splat %out_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc47) + %optrs_22 = tt.addptr %optrs, %ptrs_6 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc47) + %optrs_23 = tt.broadcast %optrs_22 : tensor<32x1x!tt.ptr> -> tensor<32x32x!tt.ptr> loc(#loc48) + %optrs_24 = tt.addptr %optrs_23, %ptrs_13 : tensor<32x32x!tt.ptr>, tensor<32x32xi32> loc(#loc48) + %0 = arith.mulf %vals_21, %cst : tensor<32x32xf32> loc(#loc1) + tt.store %optrs_24, %0, %mask_20 : tensor<32x32x!tt.ptr> loc(#loc22) + tt.return loc(#loc23) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":129:27) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":127:19) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":121:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":122:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":123:21) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":123:44) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":123:31) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":124:21) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":124:31) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:27) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:38) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:20) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:56) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:67) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:49) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":126:30) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":126:54) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":126:36) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":128:22) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":128:51) +#loc22 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":129:20) +#loc23 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":129:4) +#loc30 = loc("vals"(#loc2)) +#loc31 = loc("pid_m"(#loc4)) +#loc32 = loc("pid_n"(#loc5)) +#loc33 = loc("offs_m"(#loc6)) +#loc34 = loc("offs_m"(#loc7)) +#loc35 = loc("offs_m"(#loc8)) +#loc36 = loc("offs_n"(#loc9)) +#loc37 = loc("offs_n"(#loc10)) +#loc38 = loc("ptrs"(#loc11)) +#loc39 = loc("ptrs"(#loc12)) +#loc40 = loc("ptrs"(#loc13)) +#loc41 = loc("ptrs"(#loc14)) +#loc42 = loc("ptrs"(#loc15)) +#loc43 = loc("ptrs"(#loc16)) +#loc44 = loc("mask"(#loc17)) +#loc45 = loc("mask"(#loc18)) +#loc46 = loc("mask"(#loc19)) +#loc47 = loc("optrs"(#loc20)) +#loc48 = loc("optrs"(#loc21)) diff --git a/tests/golden/ir/ttir/golden_tile2d_sm90.ttir b/tests/golden/ir/ttir/golden_tile2d_sm90.ttir new file mode 100644 index 000000000..87a63ffc1 --- /dev/null +++ b/tests/golden/ir/ttir/golden_tile2d_sm90.ttir @@ -0,0 +1,91 @@ +#loc = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":118:0) +#loc24 = loc("in_ptr"(#loc)) +#loc25 = loc("out_ptr"(#loc)) +#loc26 = loc("M"(#loc)) +#loc27 = loc("N"(#loc)) +#loc28 = loc("stride_m"(#loc)) +#loc29 = loc("stride_n"(#loc)) +module { + tt.func public @tile2d_kernel(%in_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("in_ptr"(#loc)), %out_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("out_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %stride_m: i32 {tt.divisibility = 16 : i32} loc("stride_m"(#loc)), %stride_n: i32 {tt.divisibility = 16 : i32} loc("stride_n"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<2.000000e+00> : tensor<32x32xf32> loc(#loc1) + %vals = arith.constant dense<0.000000e+00> : tensor<32x32xf32> loc(#loc30) + %c32_i32 = arith.constant 32 : i32 loc(#loc3) + %pid_m = tt.get_program_id x : i32 loc(#loc31) + %pid_n = tt.get_program_id y : i32 loc(#loc32) + %offs_m = arith.muli %pid_m, %c32_i32 : i32 loc(#loc33) + %offs_m_0 = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc34) + %offs_m_1 = tt.splat %offs_m : i32 -> tensor<32xi32> loc(#loc35) + %offs_m_2 = arith.addi %offs_m_1, %offs_m_0 : tensor<32xi32> loc(#loc35) + %offs_n = arith.muli %pid_n, %c32_i32 : i32 loc(#loc36) + %offs_n_3 = tt.splat %offs_n : i32 -> tensor<32xi32> loc(#loc37) + %offs_n_4 = arith.addi %offs_n_3, %offs_m_0 : tensor<32xi32> loc(#loc37) + %ptrs = tt.expand_dims %offs_m_2 {axis = 1 : i32} : tensor<32xi32> -> tensor<32x1xi32> loc(#loc38) + %ptrs_5 = tt.splat %stride_m : i32 -> tensor<32x1xi32> loc(#loc39) + %ptrs_6 = arith.muli %ptrs, %ptrs_5 : tensor<32x1xi32> loc(#loc39) + %ptrs_7 = tt.splat %in_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc40) + %ptrs_8 = tt.addptr %ptrs_7, %ptrs_6 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc40) + %ptrs_9 = tt.expand_dims %offs_n_4 {axis = 0 : i32} : tensor<32xi32> -> tensor<1x32xi32> loc(#loc41) + %ptrs_10 = tt.splat %stride_n : i32 -> tensor<1x32xi32> loc(#loc42) + %ptrs_11 = arith.muli %ptrs_9, %ptrs_10 : tensor<1x32xi32> loc(#loc42) + %ptrs_12 = tt.broadcast %ptrs_8 : tensor<32x1x!tt.ptr> -> tensor<32x32x!tt.ptr> loc(#loc43) + %ptrs_13 = tt.broadcast %ptrs_11 : tensor<1x32xi32> -> tensor<32x32xi32> loc(#loc43) + %ptrs_14 = tt.addptr %ptrs_12, %ptrs_13 : tensor<32x32x!tt.ptr>, tensor<32x32xi32> loc(#loc43) + %mask = tt.splat %M : i32 -> tensor<32x1xi32> loc(#loc44) + %mask_15 = arith.cmpi slt, %ptrs, %mask : tensor<32x1xi32> loc(#loc44) + %mask_16 = tt.splat %N : i32 -> tensor<1x32xi32> loc(#loc45) + %mask_17 = arith.cmpi slt, %ptrs_9, %mask_16 : tensor<1x32xi32> loc(#loc45) + %mask_18 = tt.broadcast %mask_15 : tensor<32x1xi1> -> tensor<32x32xi1> loc(#loc46) + %mask_19 = tt.broadcast %mask_17 : tensor<1x32xi1> -> tensor<32x32xi1> loc(#loc46) + %mask_20 = arith.andi %mask_18, %mask_19 : tensor<32x32xi1> loc(#loc46) + %vals_21 = tt.load %ptrs_14, %mask_20, %vals : tensor<32x32x!tt.ptr> loc(#loc30) + %optrs = tt.splat %out_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc47) + %optrs_22 = tt.addptr %optrs, %ptrs_6 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc47) + %optrs_23 = tt.broadcast %optrs_22 : tensor<32x1x!tt.ptr> -> tensor<32x32x!tt.ptr> loc(#loc48) + %optrs_24 = tt.addptr %optrs_23, %ptrs_13 : tensor<32x32x!tt.ptr>, tensor<32x32xi32> loc(#loc48) + %0 = arith.mulf %vals_21, %cst : tensor<32x32xf32> loc(#loc1) + tt.store %optrs_24, %0, %mask_20 : tensor<32x32x!tt.ptr> loc(#loc22) + tt.return loc(#loc23) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":129:27) +#loc2 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":127:19) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":121:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":122:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":123:21) +#loc7 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":123:44) +#loc8 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":123:31) +#loc9 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":124:21) +#loc10 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":124:31) +#loc11 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:27) +#loc12 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:38) +#loc13 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:20) +#loc14 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:56) +#loc15 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:67) +#loc16 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":125:49) +#loc17 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":126:30) +#loc18 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":126:54) +#loc19 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":126:36) +#loc20 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":128:22) +#loc21 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":128:51) +#loc22 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":129:20) +#loc23 = loc("/home/hwu27/workspace/triton-viz/tests/golden/ttgir/generate_golden.py":129:4) +#loc30 = loc("vals"(#loc2)) +#loc31 = loc("pid_m"(#loc4)) +#loc32 = loc("pid_n"(#loc5)) +#loc33 = loc("offs_m"(#loc6)) +#loc34 = loc("offs_m"(#loc7)) +#loc35 = loc("offs_m"(#loc8)) +#loc36 = loc("offs_n"(#loc9)) +#loc37 = loc("offs_n"(#loc10)) +#loc38 = loc("ptrs"(#loc11)) +#loc39 = loc("ptrs"(#loc12)) +#loc40 = loc("ptrs"(#loc13)) +#loc41 = loc("ptrs"(#loc14)) +#loc42 = loc("ptrs"(#loc15)) +#loc43 = loc("ptrs"(#loc16)) +#loc44 = loc("mask"(#loc17)) +#loc45 = loc("mask"(#loc18)) +#loc46 = loc("mask"(#loc19)) +#loc47 = loc("optrs"(#loc20)) +#loc48 = loc("optrs"(#loc21)) diff --git a/tests/golden/ir/ttir/kernel_deep_chain.ttir b/tests/golden/ir/ttir/kernel_deep_chain.ttir new file mode 100644 index 000000000..4c544b0b7 --- /dev/null +++ b/tests/golden/ir/ttir/kernel_deep_chain.ttir @@ -0,0 +1,1221 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":60:0) +#loc7 = loc("out_ptr"(#loc)) +#loc8 = loc("s"(#loc)) +module { + tt.func public @deep_chain(%out_ptr: !tt.ptr loc("out_ptr"(#loc)), %s: i32 loc("s"(#loc))) attributes {noinline = false} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc9) + %off = arith.muli %pid, %s : i32 loc(#loc10) + %off_0 = arith.addi %off, %pid : i32 loc(#loc11) + %off_1 = arith.muli %off_0, %s : i32 loc(#loc10) + %off_2 = arith.addi %off_1, %pid : i32 loc(#loc11) + %off_3 = arith.muli %off_2, %s : i32 loc(#loc10) + %off_4 = arith.addi %off_3, %pid : i32 loc(#loc11) + %off_5 = arith.muli %off_4, %s : i32 loc(#loc10) + %off_6 = arith.addi %off_5, %pid : i32 loc(#loc11) + %off_7 = arith.muli %off_6, %s : i32 loc(#loc10) + %off_8 = arith.addi %off_7, %pid : i32 loc(#loc11) + %off_9 = arith.muli %off_8, %s : i32 loc(#loc10) + %off_10 = arith.addi %off_9, %pid : i32 loc(#loc11) + %off_11 = arith.muli %off_10, %s : i32 loc(#loc10) + %off_12 = arith.addi %off_11, %pid : i32 loc(#loc11) + %off_13 = arith.muli %off_12, %s : i32 loc(#loc10) + %off_14 = arith.addi %off_13, %pid : i32 loc(#loc11) + %off_15 = arith.muli %off_14, %s : i32 loc(#loc10) + %off_16 = arith.addi %off_15, %pid : i32 loc(#loc11) + %off_17 = arith.muli %off_16, %s : i32 loc(#loc10) + %off_18 = arith.addi %off_17, %pid : i32 loc(#loc11) + %off_19 = arith.muli %off_18, %s : i32 loc(#loc10) + %off_20 = arith.addi %off_19, %pid : i32 loc(#loc11) + %off_21 = arith.muli %off_20, %s : i32 loc(#loc10) + %off_22 = arith.addi %off_21, %pid : i32 loc(#loc11) + %off_23 = arith.muli %off_22, %s : i32 loc(#loc10) + %off_24 = arith.addi %off_23, %pid : i32 loc(#loc11) + %off_25 = arith.muli %off_24, %s : i32 loc(#loc10) + %off_26 = arith.addi %off_25, %pid : i32 loc(#loc11) + %off_27 = arith.muli %off_26, %s : i32 loc(#loc10) + %off_28 = arith.addi %off_27, %pid : i32 loc(#loc11) + %off_29 = arith.muli %off_28, %s : i32 loc(#loc10) + %off_30 = arith.addi %off_29, %pid : i32 loc(#loc11) + %off_31 = arith.muli %off_30, %s : i32 loc(#loc10) + %off_32 = arith.addi %off_31, %pid : i32 loc(#loc11) + %off_33 = arith.muli %off_32, %s : i32 loc(#loc10) + %off_34 = arith.addi %off_33, %pid : i32 loc(#loc11) + %off_35 = arith.muli %off_34, %s : i32 loc(#loc10) + %off_36 = arith.addi %off_35, %pid : i32 loc(#loc11) + %off_37 = arith.muli %off_36, %s : i32 loc(#loc10) + %off_38 = arith.addi %off_37, %pid : i32 loc(#loc11) + %off_39 = arith.muli %off_38, %s : i32 loc(#loc10) + %off_40 = arith.addi %off_39, %pid : i32 loc(#loc11) + %off_41 = arith.muli %off_40, %s : i32 loc(#loc10) + %off_42 = arith.addi %off_41, %pid : i32 loc(#loc11) + %off_43 = arith.muli %off_42, %s : i32 loc(#loc10) + %off_44 = arith.addi %off_43, %pid : i32 loc(#loc11) + %off_45 = arith.muli %off_44, %s : i32 loc(#loc10) + %off_46 = arith.addi %off_45, %pid : i32 loc(#loc11) + %off_47 = arith.muli %off_46, %s : i32 loc(#loc10) + %off_48 = arith.addi %off_47, %pid : i32 loc(#loc11) + %off_49 = arith.muli %off_48, %s : i32 loc(#loc10) + %off_50 = arith.addi %off_49, %pid : i32 loc(#loc11) + %off_51 = arith.muli %off_50, %s : i32 loc(#loc10) + %off_52 = arith.addi %off_51, %pid : i32 loc(#loc11) + %off_53 = arith.muli %off_52, %s : i32 loc(#loc10) + %off_54 = arith.addi %off_53, %pid : i32 loc(#loc11) + %off_55 = arith.muli %off_54, %s : i32 loc(#loc10) + %off_56 = arith.addi %off_55, %pid : i32 loc(#loc11) + %off_57 = arith.muli %off_56, %s : i32 loc(#loc10) + %off_58 = arith.addi %off_57, %pid : i32 loc(#loc11) + %off_59 = arith.muli %off_58, %s : i32 loc(#loc10) + %off_60 = arith.addi %off_59, %pid : i32 loc(#loc11) + %off_61 = arith.muli %off_60, %s : i32 loc(#loc10) + %off_62 = arith.addi %off_61, %pid : i32 loc(#loc11) + %off_63 = arith.muli %off_62, %s : i32 loc(#loc10) + %off_64 = arith.addi %off_63, %pid : i32 loc(#loc11) + %off_65 = arith.muli %off_64, %s : i32 loc(#loc10) + %off_66 = arith.addi %off_65, %pid : i32 loc(#loc11) + %off_67 = arith.muli %off_66, %s : i32 loc(#loc10) + %off_68 = arith.addi %off_67, %pid : i32 loc(#loc11) + %off_69 = arith.muli %off_68, %s : i32 loc(#loc10) + %off_70 = arith.addi %off_69, %pid : i32 loc(#loc11) + %off_71 = arith.muli %off_70, %s : i32 loc(#loc10) + %off_72 = arith.addi %off_71, %pid : i32 loc(#loc11) + %off_73 = arith.muli %off_72, %s : i32 loc(#loc10) + %off_74 = arith.addi %off_73, %pid : i32 loc(#loc11) + %off_75 = arith.muli %off_74, %s : i32 loc(#loc10) + %off_76 = arith.addi %off_75, %pid : i32 loc(#loc11) + %off_77 = arith.muli %off_76, %s : i32 loc(#loc10) + %off_78 = arith.addi %off_77, %pid : i32 loc(#loc11) + %off_79 = arith.muli %off_78, %s : i32 loc(#loc10) + %off_80 = arith.addi %off_79, %pid : i32 loc(#loc11) + %off_81 = arith.muli %off_80, %s : i32 loc(#loc10) + %off_82 = arith.addi %off_81, %pid : i32 loc(#loc11) + %off_83 = arith.muli %off_82, %s : i32 loc(#loc10) + %off_84 = arith.addi %off_83, %pid : i32 loc(#loc11) + %off_85 = arith.muli %off_84, %s : i32 loc(#loc10) + %off_86 = arith.addi %off_85, %pid : i32 loc(#loc11) + %off_87 = arith.muli %off_86, %s : i32 loc(#loc10) + %off_88 = arith.addi %off_87, %pid : i32 loc(#loc11) + %off_89 = arith.muli %off_88, %s : i32 loc(#loc10) + %off_90 = arith.addi %off_89, %pid : i32 loc(#loc11) + %off_91 = arith.muli %off_90, %s : i32 loc(#loc10) + %off_92 = arith.addi %off_91, %pid : i32 loc(#loc11) + %off_93 = arith.muli %off_92, %s : i32 loc(#loc10) + %off_94 = arith.addi %off_93, %pid : i32 loc(#loc11) + %off_95 = arith.muli %off_94, %s : i32 loc(#loc10) + %off_96 = arith.addi %off_95, %pid : i32 loc(#loc11) + %off_97 = arith.muli %off_96, %s : i32 loc(#loc10) + %off_98 = arith.addi %off_97, %pid : i32 loc(#loc11) + %off_99 = arith.muli %off_98, %s : i32 loc(#loc10) + %off_100 = arith.addi %off_99, %pid : i32 loc(#loc11) + %off_101 = arith.muli %off_100, %s : i32 loc(#loc10) + %off_102 = arith.addi %off_101, %pid : i32 loc(#loc11) + %off_103 = arith.muli %off_102, %s : i32 loc(#loc10) + %off_104 = arith.addi %off_103, %pid : i32 loc(#loc11) + %off_105 = arith.muli %off_104, %s : i32 loc(#loc10) + %off_106 = arith.addi %off_105, %pid : i32 loc(#loc11) + %off_107 = arith.muli %off_106, %s : i32 loc(#loc10) + %off_108 = arith.addi %off_107, %pid : i32 loc(#loc11) + %off_109 = arith.muli %off_108, %s : i32 loc(#loc10) + %off_110 = arith.addi %off_109, %pid : i32 loc(#loc11) + %off_111 = arith.muli %off_110, %s : i32 loc(#loc10) + %off_112 = arith.addi %off_111, %pid : i32 loc(#loc11) + %off_113 = arith.muli %off_112, %s : i32 loc(#loc10) + %off_114 = arith.addi %off_113, %pid : i32 loc(#loc11) + %off_115 = arith.muli %off_114, %s : i32 loc(#loc10) + %off_116 = arith.addi %off_115, %pid : i32 loc(#loc11) + %off_117 = arith.muli %off_116, %s : i32 loc(#loc10) + %off_118 = arith.addi %off_117, %pid : i32 loc(#loc11) + %off_119 = arith.muli %off_118, %s : i32 loc(#loc10) + %off_120 = arith.addi %off_119, %pid : i32 loc(#loc11) + %off_121 = arith.muli %off_120, %s : i32 loc(#loc10) + %off_122 = arith.addi %off_121, %pid : i32 loc(#loc11) + %off_123 = arith.muli %off_122, %s : i32 loc(#loc10) + %off_124 = arith.addi %off_123, %pid : i32 loc(#loc11) + %off_125 = arith.muli %off_124, %s : i32 loc(#loc10) + %off_126 = arith.addi %off_125, %pid : i32 loc(#loc11) + %off_127 = arith.muli %off_126, %s : i32 loc(#loc10) + %off_128 = arith.addi %off_127, %pid : i32 loc(#loc11) + %off_129 = arith.muli %off_128, %s : i32 loc(#loc10) + %off_130 = arith.addi %off_129, %pid : i32 loc(#loc11) + %off_131 = arith.muli %off_130, %s : i32 loc(#loc10) + %off_132 = arith.addi %off_131, %pid : i32 loc(#loc11) + %off_133 = arith.muli %off_132, %s : i32 loc(#loc10) + %off_134 = arith.addi %off_133, %pid : i32 loc(#loc11) + %off_135 = arith.muli %off_134, %s : i32 loc(#loc10) + %off_136 = arith.addi %off_135, %pid : i32 loc(#loc11) + %off_137 = arith.muli %off_136, %s : i32 loc(#loc10) + %off_138 = arith.addi %off_137, %pid : i32 loc(#loc11) + %off_139 = arith.muli %off_138, %s : i32 loc(#loc10) + %off_140 = arith.addi %off_139, %pid : i32 loc(#loc11) + %off_141 = arith.muli %off_140, %s : i32 loc(#loc10) + %off_142 = arith.addi %off_141, %pid : i32 loc(#loc11) + %off_143 = arith.muli %off_142, %s : i32 loc(#loc10) + %off_144 = arith.addi %off_143, %pid : i32 loc(#loc11) + %off_145 = arith.muli %off_144, %s : i32 loc(#loc10) + %off_146 = arith.addi %off_145, %pid : i32 loc(#loc11) + %off_147 = arith.muli %off_146, %s : i32 loc(#loc10) + %off_148 = arith.addi %off_147, %pid : i32 loc(#loc11) + %off_149 = arith.muli %off_148, %s : i32 loc(#loc10) + %off_150 = arith.addi %off_149, %pid : i32 loc(#loc11) + %off_151 = arith.muli %off_150, %s : i32 loc(#loc10) + %off_152 = arith.addi %off_151, %pid : i32 loc(#loc11) + %off_153 = arith.muli %off_152, %s : i32 loc(#loc10) + %off_154 = arith.addi %off_153, %pid : i32 loc(#loc11) + %off_155 = arith.muli %off_154, %s : i32 loc(#loc10) + %off_156 = arith.addi %off_155, %pid : i32 loc(#loc11) + %off_157 = arith.muli %off_156, %s : i32 loc(#loc10) + %off_158 = arith.addi %off_157, %pid : i32 loc(#loc11) + %off_159 = arith.muli %off_158, %s : i32 loc(#loc10) + %off_160 = arith.addi %off_159, %pid : i32 loc(#loc11) + %off_161 = arith.muli %off_160, %s : i32 loc(#loc10) + %off_162 = arith.addi %off_161, %pid : i32 loc(#loc11) + %off_163 = arith.muli %off_162, %s : i32 loc(#loc10) + %off_164 = arith.addi %off_163, %pid : i32 loc(#loc11) + %off_165 = arith.muli %off_164, %s : i32 loc(#loc10) + %off_166 = arith.addi %off_165, %pid : i32 loc(#loc11) + %off_167 = arith.muli %off_166, %s : i32 loc(#loc10) + %off_168 = arith.addi %off_167, %pid : i32 loc(#loc11) + %off_169 = arith.muli %off_168, %s : i32 loc(#loc10) + %off_170 = arith.addi %off_169, %pid : i32 loc(#loc11) + %off_171 = arith.muli %off_170, %s : i32 loc(#loc10) + %off_172 = arith.addi %off_171, %pid : i32 loc(#loc11) + %off_173 = arith.muli %off_172, %s : i32 loc(#loc10) + %off_174 = arith.addi %off_173, %pid : i32 loc(#loc11) + %off_175 = arith.muli %off_174, %s : i32 loc(#loc10) + %off_176 = arith.addi %off_175, %pid : i32 loc(#loc11) + %off_177 = arith.muli %off_176, %s : i32 loc(#loc10) + %off_178 = arith.addi %off_177, %pid : i32 loc(#loc11) + %off_179 = arith.muli %off_178, %s : i32 loc(#loc10) + %off_180 = arith.addi %off_179, %pid : i32 loc(#loc11) + %off_181 = arith.muli %off_180, %s : i32 loc(#loc10) + %off_182 = arith.addi %off_181, %pid : i32 loc(#loc11) + %off_183 = arith.muli %off_182, %s : i32 loc(#loc10) + %off_184 = arith.addi %off_183, %pid : i32 loc(#loc11) + %off_185 = arith.muli %off_184, %s : i32 loc(#loc10) + %off_186 = arith.addi %off_185, %pid : i32 loc(#loc11) + %off_187 = arith.muli %off_186, %s : i32 loc(#loc10) + %off_188 = arith.addi %off_187, %pid : i32 loc(#loc11) + %off_189 = arith.muli %off_188, %s : i32 loc(#loc10) + %off_190 = arith.addi %off_189, %pid : i32 loc(#loc11) + %off_191 = arith.muli %off_190, %s : i32 loc(#loc10) + %off_192 = arith.addi %off_191, %pid : i32 loc(#loc11) + %off_193 = arith.muli %off_192, %s : i32 loc(#loc10) + %off_194 = arith.addi %off_193, %pid : i32 loc(#loc11) + %off_195 = arith.muli %off_194, %s : i32 loc(#loc10) + %off_196 = arith.addi %off_195, %pid : i32 loc(#loc11) + %off_197 = arith.muli %off_196, %s : i32 loc(#loc10) + %off_198 = arith.addi %off_197, %pid : i32 loc(#loc11) + %off_199 = arith.muli %off_198, %s : i32 loc(#loc10) + %off_200 = arith.addi %off_199, %pid : i32 loc(#loc11) + %off_201 = arith.muli %off_200, %s : i32 loc(#loc10) + %off_202 = arith.addi %off_201, %pid : i32 loc(#loc11) + %off_203 = arith.muli %off_202, %s : i32 loc(#loc10) + %off_204 = arith.addi %off_203, %pid : i32 loc(#loc11) + %off_205 = arith.muli %off_204, %s : i32 loc(#loc10) + %off_206 = arith.addi %off_205, %pid : i32 loc(#loc11) + %off_207 = arith.muli %off_206, %s : i32 loc(#loc10) + %off_208 = arith.addi %off_207, %pid : i32 loc(#loc11) + %off_209 = arith.muli %off_208, %s : i32 loc(#loc10) + %off_210 = arith.addi %off_209, %pid : i32 loc(#loc11) + %off_211 = arith.muli %off_210, %s : i32 loc(#loc10) + %off_212 = arith.addi %off_211, %pid : i32 loc(#loc11) + %off_213 = arith.muli %off_212, %s : i32 loc(#loc10) + %off_214 = arith.addi %off_213, %pid : i32 loc(#loc11) + %off_215 = arith.muli %off_214, %s : i32 loc(#loc10) + %off_216 = arith.addi %off_215, %pid : i32 loc(#loc11) + %off_217 = arith.muli %off_216, %s : i32 loc(#loc10) + %off_218 = arith.addi %off_217, %pid : i32 loc(#loc11) + %off_219 = arith.muli %off_218, %s : i32 loc(#loc10) + %off_220 = arith.addi %off_219, %pid : i32 loc(#loc11) + %off_221 = arith.muli %off_220, %s : i32 loc(#loc10) + %off_222 = arith.addi %off_221, %pid : i32 loc(#loc11) + %off_223 = arith.muli %off_222, %s : i32 loc(#loc10) + %off_224 = arith.addi %off_223, %pid : i32 loc(#loc11) + %off_225 = arith.muli %off_224, %s : i32 loc(#loc10) + %off_226 = arith.addi %off_225, %pid : i32 loc(#loc11) + %off_227 = arith.muli %off_226, %s : i32 loc(#loc10) + %off_228 = arith.addi %off_227, %pid : i32 loc(#loc11) + %off_229 = arith.muli %off_228, %s : i32 loc(#loc10) + %off_230 = arith.addi %off_229, %pid : i32 loc(#loc11) + %off_231 = arith.muli %off_230, %s : i32 loc(#loc10) + %off_232 = arith.addi %off_231, %pid : i32 loc(#loc11) + %off_233 = arith.muli %off_232, %s : i32 loc(#loc10) + %off_234 = arith.addi %off_233, %pid : i32 loc(#loc11) + %off_235 = arith.muli %off_234, %s : i32 loc(#loc10) + %off_236 = arith.addi %off_235, %pid : i32 loc(#loc11) + %off_237 = arith.muli %off_236, %s : i32 loc(#loc10) + %off_238 = arith.addi %off_237, %pid : i32 loc(#loc11) + %off_239 = arith.muli %off_238, %s : i32 loc(#loc10) + %off_240 = arith.addi %off_239, %pid : i32 loc(#loc11) + %off_241 = arith.muli %off_240, %s : i32 loc(#loc10) + %off_242 = arith.addi %off_241, %pid : i32 loc(#loc11) + %off_243 = arith.muli %off_242, %s : i32 loc(#loc10) + %off_244 = arith.addi %off_243, %pid : i32 loc(#loc11) + %off_245 = arith.muli %off_244, %s : i32 loc(#loc10) + %off_246 = arith.addi %off_245, %pid : i32 loc(#loc11) + %off_247 = arith.muli %off_246, %s : i32 loc(#loc10) + %off_248 = arith.addi %off_247, %pid : i32 loc(#loc11) + %off_249 = arith.muli %off_248, %s : i32 loc(#loc10) + %off_250 = arith.addi %off_249, %pid : i32 loc(#loc11) + %off_251 = arith.muli %off_250, %s : i32 loc(#loc10) + %off_252 = arith.addi %off_251, %pid : i32 loc(#loc11) + %off_253 = arith.muli %off_252, %s : i32 loc(#loc10) + %off_254 = arith.addi %off_253, %pid : i32 loc(#loc11) + %off_255 = arith.muli %off_254, %s : i32 loc(#loc10) + %off_256 = arith.addi %off_255, %pid : i32 loc(#loc11) + %off_257 = arith.muli %off_256, %s : i32 loc(#loc10) + %off_258 = arith.addi %off_257, %pid : i32 loc(#loc11) + %off_259 = arith.muli %off_258, %s : i32 loc(#loc10) + %off_260 = arith.addi %off_259, %pid : i32 loc(#loc11) + %off_261 = arith.muli %off_260, %s : i32 loc(#loc10) + %off_262 = arith.addi %off_261, %pid : i32 loc(#loc11) + %off_263 = arith.muli %off_262, %s : i32 loc(#loc10) + %off_264 = arith.addi %off_263, %pid : i32 loc(#loc11) + %off_265 = arith.muli %off_264, %s : i32 loc(#loc10) + %off_266 = arith.addi %off_265, %pid : i32 loc(#loc11) + %off_267 = arith.muli %off_266, %s : i32 loc(#loc10) + %off_268 = arith.addi %off_267, %pid : i32 loc(#loc11) + %off_269 = arith.muli %off_268, %s : i32 loc(#loc10) + %off_270 = arith.addi %off_269, %pid : i32 loc(#loc11) + %off_271 = arith.muli %off_270, %s : i32 loc(#loc10) + %off_272 = arith.addi %off_271, %pid : i32 loc(#loc11) + %off_273 = arith.muli %off_272, %s : i32 loc(#loc10) + %off_274 = arith.addi %off_273, %pid : i32 loc(#loc11) + %off_275 = arith.muli %off_274, %s : i32 loc(#loc10) + %off_276 = arith.addi %off_275, %pid : i32 loc(#loc11) + %off_277 = arith.muli %off_276, %s : i32 loc(#loc10) + %off_278 = arith.addi %off_277, %pid : i32 loc(#loc11) + %off_279 = arith.muli %off_278, %s : i32 loc(#loc10) + %off_280 = arith.addi %off_279, %pid : i32 loc(#loc11) + %off_281 = arith.muli %off_280, %s : i32 loc(#loc10) + %off_282 = arith.addi %off_281, %pid : i32 loc(#loc11) + %off_283 = arith.muli %off_282, %s : i32 loc(#loc10) + %off_284 = arith.addi %off_283, %pid : i32 loc(#loc11) + %off_285 = arith.muli %off_284, %s : i32 loc(#loc10) + %off_286 = arith.addi %off_285, %pid : i32 loc(#loc11) + %off_287 = arith.muli %off_286, %s : i32 loc(#loc10) + %off_288 = arith.addi %off_287, %pid : i32 loc(#loc11) + %off_289 = arith.muli %off_288, %s : i32 loc(#loc10) + %off_290 = arith.addi %off_289, %pid : i32 loc(#loc11) + %off_291 = arith.muli %off_290, %s : i32 loc(#loc10) + %off_292 = arith.addi %off_291, %pid : i32 loc(#loc11) + %off_293 = arith.muli %off_292, %s : i32 loc(#loc10) + %off_294 = arith.addi %off_293, %pid : i32 loc(#loc11) + %off_295 = arith.muli %off_294, %s : i32 loc(#loc10) + %off_296 = arith.addi %off_295, %pid : i32 loc(#loc11) + %off_297 = arith.muli %off_296, %s : i32 loc(#loc10) + %off_298 = arith.addi %off_297, %pid : i32 loc(#loc11) + %off_299 = arith.muli %off_298, %s : i32 loc(#loc10) + %off_300 = arith.addi %off_299, %pid : i32 loc(#loc11) + %off_301 = arith.muli %off_300, %s : i32 loc(#loc10) + %off_302 = arith.addi %off_301, %pid : i32 loc(#loc11) + %off_303 = arith.muli %off_302, %s : i32 loc(#loc10) + %off_304 = arith.addi %off_303, %pid : i32 loc(#loc11) + %off_305 = arith.muli %off_304, %s : i32 loc(#loc10) + %off_306 = arith.addi %off_305, %pid : i32 loc(#loc11) + %off_307 = arith.muli %off_306, %s : i32 loc(#loc10) + %off_308 = arith.addi %off_307, %pid : i32 loc(#loc11) + %off_309 = arith.muli %off_308, %s : i32 loc(#loc10) + %off_310 = arith.addi %off_309, %pid : i32 loc(#loc11) + %off_311 = arith.muli %off_310, %s : i32 loc(#loc10) + %off_312 = arith.addi %off_311, %pid : i32 loc(#loc11) + %off_313 = arith.muli %off_312, %s : i32 loc(#loc10) + %off_314 = arith.addi %off_313, %pid : i32 loc(#loc11) + %off_315 = arith.muli %off_314, %s : i32 loc(#loc10) + %off_316 = arith.addi %off_315, %pid : i32 loc(#loc11) + %off_317 = arith.muli %off_316, %s : i32 loc(#loc10) + %off_318 = arith.addi %off_317, %pid : i32 loc(#loc11) + %off_319 = arith.muli %off_318, %s : i32 loc(#loc10) + %off_320 = arith.addi %off_319, %pid : i32 loc(#loc11) + %off_321 = arith.muli %off_320, %s : i32 loc(#loc10) + %off_322 = arith.addi %off_321, %pid : i32 loc(#loc11) + %off_323 = arith.muli %off_322, %s : i32 loc(#loc10) + %off_324 = arith.addi %off_323, %pid : i32 loc(#loc11) + %off_325 = arith.muli %off_324, %s : i32 loc(#loc10) + %off_326 = arith.addi %off_325, %pid : i32 loc(#loc11) + %off_327 = arith.muli %off_326, %s : i32 loc(#loc10) + %off_328 = arith.addi %off_327, %pid : i32 loc(#loc11) + %off_329 = arith.muli %off_328, %s : i32 loc(#loc10) + %off_330 = arith.addi %off_329, %pid : i32 loc(#loc11) + %off_331 = arith.muli %off_330, %s : i32 loc(#loc10) + %off_332 = arith.addi %off_331, %pid : i32 loc(#loc11) + %off_333 = arith.muli %off_332, %s : i32 loc(#loc10) + %off_334 = arith.addi %off_333, %pid : i32 loc(#loc11) + %off_335 = arith.muli %off_334, %s : i32 loc(#loc10) + %off_336 = arith.addi %off_335, %pid : i32 loc(#loc11) + %off_337 = arith.muli %off_336, %s : i32 loc(#loc10) + %off_338 = arith.addi %off_337, %pid : i32 loc(#loc11) + %off_339 = arith.muli %off_338, %s : i32 loc(#loc10) + %off_340 = arith.addi %off_339, %pid : i32 loc(#loc11) + %off_341 = arith.muli %off_340, %s : i32 loc(#loc10) + %off_342 = arith.addi %off_341, %pid : i32 loc(#loc11) + %off_343 = arith.muli %off_342, %s : i32 loc(#loc10) + %off_344 = arith.addi %off_343, %pid : i32 loc(#loc11) + %off_345 = arith.muli %off_344, %s : i32 loc(#loc10) + %off_346 = arith.addi %off_345, %pid : i32 loc(#loc11) + %off_347 = arith.muli %off_346, %s : i32 loc(#loc10) + %off_348 = arith.addi %off_347, %pid : i32 loc(#loc11) + %off_349 = arith.muli %off_348, %s : i32 loc(#loc10) + %off_350 = arith.addi %off_349, %pid : i32 loc(#loc11) + %off_351 = arith.muli %off_350, %s : i32 loc(#loc10) + %off_352 = arith.addi %off_351, %pid : i32 loc(#loc11) + %off_353 = arith.muli %off_352, %s : i32 loc(#loc10) + %off_354 = arith.addi %off_353, %pid : i32 loc(#loc11) + %off_355 = arith.muli %off_354, %s : i32 loc(#loc10) + %off_356 = arith.addi %off_355, %pid : i32 loc(#loc11) + %off_357 = arith.muli %off_356, %s : i32 loc(#loc10) + %off_358 = arith.addi %off_357, %pid : i32 loc(#loc11) + %off_359 = arith.muli %off_358, %s : i32 loc(#loc10) + %off_360 = arith.addi %off_359, %pid : i32 loc(#loc11) + %off_361 = arith.muli %off_360, %s : i32 loc(#loc10) + %off_362 = arith.addi %off_361, %pid : i32 loc(#loc11) + %off_363 = arith.muli %off_362, %s : i32 loc(#loc10) + %off_364 = arith.addi %off_363, %pid : i32 loc(#loc11) + %off_365 = arith.muli %off_364, %s : i32 loc(#loc10) + %off_366 = arith.addi %off_365, %pid : i32 loc(#loc11) + %off_367 = arith.muli %off_366, %s : i32 loc(#loc10) + %off_368 = arith.addi %off_367, %pid : i32 loc(#loc11) + %off_369 = arith.muli %off_368, %s : i32 loc(#loc10) + %off_370 = arith.addi %off_369, %pid : i32 loc(#loc11) + %off_371 = arith.muli %off_370, %s : i32 loc(#loc10) + %off_372 = arith.addi %off_371, %pid : i32 loc(#loc11) + %off_373 = arith.muli %off_372, %s : i32 loc(#loc10) + %off_374 = arith.addi %off_373, %pid : i32 loc(#loc11) + %off_375 = arith.muli %off_374, %s : i32 loc(#loc10) + %off_376 = arith.addi %off_375, %pid : i32 loc(#loc11) + %off_377 = arith.muli %off_376, %s : i32 loc(#loc10) + %off_378 = arith.addi %off_377, %pid : i32 loc(#loc11) + %off_379 = arith.muli %off_378, %s : i32 loc(#loc10) + %off_380 = arith.addi %off_379, %pid : i32 loc(#loc11) + %off_381 = arith.muli %off_380, %s : i32 loc(#loc10) + %off_382 = arith.addi %off_381, %pid : i32 loc(#loc11) + %off_383 = arith.muli %off_382, %s : i32 loc(#loc10) + %off_384 = arith.addi %off_383, %pid : i32 loc(#loc11) + %off_385 = arith.muli %off_384, %s : i32 loc(#loc10) + %off_386 = arith.addi %off_385, %pid : i32 loc(#loc11) + %off_387 = arith.muli %off_386, %s : i32 loc(#loc10) + %off_388 = arith.addi %off_387, %pid : i32 loc(#loc11) + %off_389 = arith.muli %off_388, %s : i32 loc(#loc10) + %off_390 = arith.addi %off_389, %pid : i32 loc(#loc11) + %off_391 = arith.muli %off_390, %s : i32 loc(#loc10) + %off_392 = arith.addi %off_391, %pid : i32 loc(#loc11) + %off_393 = arith.muli %off_392, %s : i32 loc(#loc10) + %off_394 = arith.addi %off_393, %pid : i32 loc(#loc11) + %off_395 = arith.muli %off_394, %s : i32 loc(#loc10) + %off_396 = arith.addi %off_395, %pid : i32 loc(#loc11) + %off_397 = arith.muli %off_396, %s : i32 loc(#loc10) + %off_398 = arith.addi %off_397, %pid : i32 loc(#loc11) + %off_399 = arith.muli %off_398, %s : i32 loc(#loc10) + %off_400 = arith.addi %off_399, %pid : i32 loc(#loc11) + %off_401 = arith.muli %off_400, %s : i32 loc(#loc10) + %off_402 = arith.addi %off_401, %pid : i32 loc(#loc11) + %off_403 = arith.muli %off_402, %s : i32 loc(#loc10) + %off_404 = arith.addi %off_403, %pid : i32 loc(#loc11) + %off_405 = arith.muli %off_404, %s : i32 loc(#loc10) + %off_406 = arith.addi %off_405, %pid : i32 loc(#loc11) + %off_407 = arith.muli %off_406, %s : i32 loc(#loc10) + %off_408 = arith.addi %off_407, %pid : i32 loc(#loc11) + %off_409 = arith.muli %off_408, %s : i32 loc(#loc10) + %off_410 = arith.addi %off_409, %pid : i32 loc(#loc11) + %off_411 = arith.muli %off_410, %s : i32 loc(#loc10) + %off_412 = arith.addi %off_411, %pid : i32 loc(#loc11) + %off_413 = arith.muli %off_412, %s : i32 loc(#loc10) + %off_414 = arith.addi %off_413, %pid : i32 loc(#loc11) + %off_415 = arith.muli %off_414, %s : i32 loc(#loc10) + %off_416 = arith.addi %off_415, %pid : i32 loc(#loc11) + %off_417 = arith.muli %off_416, %s : i32 loc(#loc10) + %off_418 = arith.addi %off_417, %pid : i32 loc(#loc11) + %off_419 = arith.muli %off_418, %s : i32 loc(#loc10) + %off_420 = arith.addi %off_419, %pid : i32 loc(#loc11) + %off_421 = arith.muli %off_420, %s : i32 loc(#loc10) + %off_422 = arith.addi %off_421, %pid : i32 loc(#loc11) + %off_423 = arith.muli %off_422, %s : i32 loc(#loc10) + %off_424 = arith.addi %off_423, %pid : i32 loc(#loc11) + %off_425 = arith.muli %off_424, %s : i32 loc(#loc10) + %off_426 = arith.addi %off_425, %pid : i32 loc(#loc11) + %off_427 = arith.muli %off_426, %s : i32 loc(#loc10) + %off_428 = arith.addi %off_427, %pid : i32 loc(#loc11) + %off_429 = arith.muli %off_428, %s : i32 loc(#loc10) + %off_430 = arith.addi %off_429, %pid : i32 loc(#loc11) + %off_431 = arith.muli %off_430, %s : i32 loc(#loc10) + %off_432 = arith.addi %off_431, %pid : i32 loc(#loc11) + %off_433 = arith.muli %off_432, %s : i32 loc(#loc10) + %off_434 = arith.addi %off_433, %pid : i32 loc(#loc11) + %off_435 = arith.muli %off_434, %s : i32 loc(#loc10) + %off_436 = arith.addi %off_435, %pid : i32 loc(#loc11) + %off_437 = arith.muli %off_436, %s : i32 loc(#loc10) + %off_438 = arith.addi %off_437, %pid : i32 loc(#loc11) + %off_439 = arith.muli %off_438, %s : i32 loc(#loc10) + %off_440 = arith.addi %off_439, %pid : i32 loc(#loc11) + %off_441 = arith.muli %off_440, %s : i32 loc(#loc10) + %off_442 = arith.addi %off_441, %pid : i32 loc(#loc11) + %off_443 = arith.muli %off_442, %s : i32 loc(#loc10) + %off_444 = arith.addi %off_443, %pid : i32 loc(#loc11) + %off_445 = arith.muli %off_444, %s : i32 loc(#loc10) + %off_446 = arith.addi %off_445, %pid : i32 loc(#loc11) + %off_447 = arith.muli %off_446, %s : i32 loc(#loc10) + %off_448 = arith.addi %off_447, %pid : i32 loc(#loc11) + %off_449 = arith.muli %off_448, %s : i32 loc(#loc10) + %off_450 = arith.addi %off_449, %pid : i32 loc(#loc11) + %off_451 = arith.muli %off_450, %s : i32 loc(#loc10) + %off_452 = arith.addi %off_451, %pid : i32 loc(#loc11) + %off_453 = arith.muli %off_452, %s : i32 loc(#loc10) + %off_454 = arith.addi %off_453, %pid : i32 loc(#loc11) + %off_455 = arith.muli %off_454, %s : i32 loc(#loc10) + %off_456 = arith.addi %off_455, %pid : i32 loc(#loc11) + %off_457 = arith.muli %off_456, %s : i32 loc(#loc10) + %off_458 = arith.addi %off_457, %pid : i32 loc(#loc11) + %off_459 = arith.muli %off_458, %s : i32 loc(#loc10) + %off_460 = arith.addi %off_459, %pid : i32 loc(#loc11) + %off_461 = arith.muli %off_460, %s : i32 loc(#loc10) + %off_462 = arith.addi %off_461, %pid : i32 loc(#loc11) + %off_463 = arith.muli %off_462, %s : i32 loc(#loc10) + %off_464 = arith.addi %off_463, %pid : i32 loc(#loc11) + %off_465 = arith.muli %off_464, %s : i32 loc(#loc10) + %off_466 = arith.addi %off_465, %pid : i32 loc(#loc11) + %off_467 = arith.muli %off_466, %s : i32 loc(#loc10) + %off_468 = arith.addi %off_467, %pid : i32 loc(#loc11) + %off_469 = arith.muli %off_468, %s : i32 loc(#loc10) + %off_470 = arith.addi %off_469, %pid : i32 loc(#loc11) + %off_471 = arith.muli %off_470, %s : i32 loc(#loc10) + %off_472 = arith.addi %off_471, %pid : i32 loc(#loc11) + %off_473 = arith.muli %off_472, %s : i32 loc(#loc10) + %off_474 = arith.addi %off_473, %pid : i32 loc(#loc11) + %off_475 = arith.muli %off_474, %s : i32 loc(#loc10) + %off_476 = arith.addi %off_475, %pid : i32 loc(#loc11) + %off_477 = arith.muli %off_476, %s : i32 loc(#loc10) + %off_478 = arith.addi %off_477, %pid : i32 loc(#loc11) + %off_479 = arith.muli %off_478, %s : i32 loc(#loc10) + %off_480 = arith.addi %off_479, %pid : i32 loc(#loc11) + %off_481 = arith.muli %off_480, %s : i32 loc(#loc10) + %off_482 = arith.addi %off_481, %pid : i32 loc(#loc11) + %off_483 = arith.muli %off_482, %s : i32 loc(#loc10) + %off_484 = arith.addi %off_483, %pid : i32 loc(#loc11) + %off_485 = arith.muli %off_484, %s : i32 loc(#loc10) + %off_486 = arith.addi %off_485, %pid : i32 loc(#loc11) + %off_487 = arith.muli %off_486, %s : i32 loc(#loc10) + %off_488 = arith.addi %off_487, %pid : i32 loc(#loc11) + %off_489 = arith.muli %off_488, %s : i32 loc(#loc10) + %off_490 = arith.addi %off_489, %pid : i32 loc(#loc11) + %off_491 = arith.muli %off_490, %s : i32 loc(#loc10) + %off_492 = arith.addi %off_491, %pid : i32 loc(#loc11) + %off_493 = arith.muli %off_492, %s : i32 loc(#loc10) + %off_494 = arith.addi %off_493, %pid : i32 loc(#loc11) + %off_495 = arith.muli %off_494, %s : i32 loc(#loc10) + %off_496 = arith.addi %off_495, %pid : i32 loc(#loc11) + %off_497 = arith.muli %off_496, %s : i32 loc(#loc10) + %off_498 = arith.addi %off_497, %pid : i32 loc(#loc11) + %off_499 = arith.muli %off_498, %s : i32 loc(#loc10) + %off_500 = arith.addi %off_499, %pid : i32 loc(#loc11) + %off_501 = arith.muli %off_500, %s : i32 loc(#loc10) + %off_502 = arith.addi %off_501, %pid : i32 loc(#loc11) + %off_503 = arith.muli %off_502, %s : i32 loc(#loc10) + %off_504 = arith.addi %off_503, %pid : i32 loc(#loc11) + %off_505 = arith.muli %off_504, %s : i32 loc(#loc10) + %off_506 = arith.addi %off_505, %pid : i32 loc(#loc11) + %off_507 = arith.muli %off_506, %s : i32 loc(#loc10) + %off_508 = arith.addi %off_507, %pid : i32 loc(#loc11) + %off_509 = arith.muli %off_508, %s : i32 loc(#loc10) + %off_510 = arith.addi %off_509, %pid : i32 loc(#loc11) + %off_511 = arith.muli %off_510, %s : i32 loc(#loc10) + %off_512 = arith.addi %off_511, %pid : i32 loc(#loc11) + %off_513 = arith.muli %off_512, %s : i32 loc(#loc10) + %off_514 = arith.addi %off_513, %pid : i32 loc(#loc11) + %off_515 = arith.muli %off_514, %s : i32 loc(#loc10) + %off_516 = arith.addi %off_515, %pid : i32 loc(#loc11) + %off_517 = arith.muli %off_516, %s : i32 loc(#loc10) + %off_518 = arith.addi %off_517, %pid : i32 loc(#loc11) + %off_519 = arith.muli %off_518, %s : i32 loc(#loc10) + %off_520 = arith.addi %off_519, %pid : i32 loc(#loc11) + %off_521 = arith.muli %off_520, %s : i32 loc(#loc10) + %off_522 = arith.addi %off_521, %pid : i32 loc(#loc11) + %off_523 = arith.muli %off_522, %s : i32 loc(#loc10) + %off_524 = arith.addi %off_523, %pid : i32 loc(#loc11) + %off_525 = arith.muli %off_524, %s : i32 loc(#loc10) + %off_526 = arith.addi %off_525, %pid : i32 loc(#loc11) + %off_527 = arith.muli %off_526, %s : i32 loc(#loc10) + %off_528 = arith.addi %off_527, %pid : i32 loc(#loc11) + %off_529 = arith.muli %off_528, %s : i32 loc(#loc10) + %off_530 = arith.addi %off_529, %pid : i32 loc(#loc11) + %off_531 = arith.muli %off_530, %s : i32 loc(#loc10) + %off_532 = arith.addi %off_531, %pid : i32 loc(#loc11) + %off_533 = arith.muli %off_532, %s : i32 loc(#loc10) + %off_534 = arith.addi %off_533, %pid : i32 loc(#loc11) + %off_535 = arith.muli %off_534, %s : i32 loc(#loc10) + %off_536 = arith.addi %off_535, %pid : i32 loc(#loc11) + %off_537 = arith.muli %off_536, %s : i32 loc(#loc10) + %off_538 = arith.addi %off_537, %pid : i32 loc(#loc11) + %off_539 = arith.muli %off_538, %s : i32 loc(#loc10) + %off_540 = arith.addi %off_539, %pid : i32 loc(#loc11) + %off_541 = arith.muli %off_540, %s : i32 loc(#loc10) + %off_542 = arith.addi %off_541, %pid : i32 loc(#loc11) + %off_543 = arith.muli %off_542, %s : i32 loc(#loc10) + %off_544 = arith.addi %off_543, %pid : i32 loc(#loc11) + %off_545 = arith.muli %off_544, %s : i32 loc(#loc10) + %off_546 = arith.addi %off_545, %pid : i32 loc(#loc11) + %off_547 = arith.muli %off_546, %s : i32 loc(#loc10) + %off_548 = arith.addi %off_547, %pid : i32 loc(#loc11) + %off_549 = arith.muli %off_548, %s : i32 loc(#loc10) + %off_550 = arith.addi %off_549, %pid : i32 loc(#loc11) + %off_551 = arith.muli %off_550, %s : i32 loc(#loc10) + %off_552 = arith.addi %off_551, %pid : i32 loc(#loc11) + %off_553 = arith.muli %off_552, %s : i32 loc(#loc10) + %off_554 = arith.addi %off_553, %pid : i32 loc(#loc11) + %off_555 = arith.muli %off_554, %s : i32 loc(#loc10) + %off_556 = arith.addi %off_555, %pid : i32 loc(#loc11) + %off_557 = arith.muli %off_556, %s : i32 loc(#loc10) + %off_558 = arith.addi %off_557, %pid : i32 loc(#loc11) + %off_559 = arith.muli %off_558, %s : i32 loc(#loc10) + %off_560 = arith.addi %off_559, %pid : i32 loc(#loc11) + %off_561 = arith.muli %off_560, %s : i32 loc(#loc10) + %off_562 = arith.addi %off_561, %pid : i32 loc(#loc11) + %off_563 = arith.muli %off_562, %s : i32 loc(#loc10) + %off_564 = arith.addi %off_563, %pid : i32 loc(#loc11) + %off_565 = arith.muli %off_564, %s : i32 loc(#loc10) + %off_566 = arith.addi %off_565, %pid : i32 loc(#loc11) + %off_567 = arith.muli %off_566, %s : i32 loc(#loc10) + %off_568 = arith.addi %off_567, %pid : i32 loc(#loc11) + %off_569 = arith.muli %off_568, %s : i32 loc(#loc10) + %off_570 = arith.addi %off_569, %pid : i32 loc(#loc11) + %off_571 = arith.muli %off_570, %s : i32 loc(#loc10) + %off_572 = arith.addi %off_571, %pid : i32 loc(#loc11) + %off_573 = arith.muli %off_572, %s : i32 loc(#loc10) + %off_574 = arith.addi %off_573, %pid : i32 loc(#loc11) + %off_575 = arith.muli %off_574, %s : i32 loc(#loc10) + %off_576 = arith.addi %off_575, %pid : i32 loc(#loc11) + %off_577 = arith.muli %off_576, %s : i32 loc(#loc10) + %off_578 = arith.addi %off_577, %pid : i32 loc(#loc11) + %off_579 = arith.muli %off_578, %s : i32 loc(#loc10) + %off_580 = arith.addi %off_579, %pid : i32 loc(#loc11) + %off_581 = arith.muli %off_580, %s : i32 loc(#loc10) + %off_582 = arith.addi %off_581, %pid : i32 loc(#loc11) + %off_583 = arith.muli %off_582, %s : i32 loc(#loc10) + %off_584 = arith.addi %off_583, %pid : i32 loc(#loc11) + %off_585 = arith.muli %off_584, %s : i32 loc(#loc10) + %off_586 = arith.addi %off_585, %pid : i32 loc(#loc11) + %off_587 = arith.muli %off_586, %s : i32 loc(#loc10) + %off_588 = arith.addi %off_587, %pid : i32 loc(#loc11) + %off_589 = arith.muli %off_588, %s : i32 loc(#loc10) + %off_590 = arith.addi %off_589, %pid : i32 loc(#loc11) + %off_591 = arith.muli %off_590, %s : i32 loc(#loc10) + %off_592 = arith.addi %off_591, %pid : i32 loc(#loc11) + %off_593 = arith.muli %off_592, %s : i32 loc(#loc10) + %off_594 = arith.addi %off_593, %pid : i32 loc(#loc11) + %off_595 = arith.muli %off_594, %s : i32 loc(#loc10) + %off_596 = arith.addi %off_595, %pid : i32 loc(#loc11) + %off_597 = arith.muli %off_596, %s : i32 loc(#loc10) + %off_598 = arith.addi %off_597, %pid : i32 loc(#loc11) + %off_599 = arith.muli %off_598, %s : i32 loc(#loc10) + %off_600 = arith.addi %off_599, %pid : i32 loc(#loc11) + %off_601 = arith.muli %off_600, %s : i32 loc(#loc10) + %off_602 = arith.addi %off_601, %pid : i32 loc(#loc11) + %off_603 = arith.muli %off_602, %s : i32 loc(#loc10) + %off_604 = arith.addi %off_603, %pid : i32 loc(#loc11) + %off_605 = arith.muli %off_604, %s : i32 loc(#loc10) + %off_606 = arith.addi %off_605, %pid : i32 loc(#loc11) + %off_607 = arith.muli %off_606, %s : i32 loc(#loc10) + %off_608 = arith.addi %off_607, %pid : i32 loc(#loc11) + %off_609 = arith.muli %off_608, %s : i32 loc(#loc10) + %off_610 = arith.addi %off_609, %pid : i32 loc(#loc11) + %off_611 = arith.muli %off_610, %s : i32 loc(#loc10) + %off_612 = arith.addi %off_611, %pid : i32 loc(#loc11) + %off_613 = arith.muli %off_612, %s : i32 loc(#loc10) + %off_614 = arith.addi %off_613, %pid : i32 loc(#loc11) + %off_615 = arith.muli %off_614, %s : i32 loc(#loc10) + %off_616 = arith.addi %off_615, %pid : i32 loc(#loc11) + %off_617 = arith.muli %off_616, %s : i32 loc(#loc10) + %off_618 = arith.addi %off_617, %pid : i32 loc(#loc11) + %off_619 = arith.muli %off_618, %s : i32 loc(#loc10) + %off_620 = arith.addi %off_619, %pid : i32 loc(#loc11) + %off_621 = arith.muli %off_620, %s : i32 loc(#loc10) + %off_622 = arith.addi %off_621, %pid : i32 loc(#loc11) + %off_623 = arith.muli %off_622, %s : i32 loc(#loc10) + %off_624 = arith.addi %off_623, %pid : i32 loc(#loc11) + %off_625 = arith.muli %off_624, %s : i32 loc(#loc10) + %off_626 = arith.addi %off_625, %pid : i32 loc(#loc11) + %off_627 = arith.muli %off_626, %s : i32 loc(#loc10) + %off_628 = arith.addi %off_627, %pid : i32 loc(#loc11) + %off_629 = arith.muli %off_628, %s : i32 loc(#loc10) + %off_630 = arith.addi %off_629, %pid : i32 loc(#loc11) + %off_631 = arith.muli %off_630, %s : i32 loc(#loc10) + %off_632 = arith.addi %off_631, %pid : i32 loc(#loc11) + %off_633 = arith.muli %off_632, %s : i32 loc(#loc10) + %off_634 = arith.addi %off_633, %pid : i32 loc(#loc11) + %off_635 = arith.muli %off_634, %s : i32 loc(#loc10) + %off_636 = arith.addi %off_635, %pid : i32 loc(#loc11) + %off_637 = arith.muli %off_636, %s : i32 loc(#loc10) + %off_638 = arith.addi %off_637, %pid : i32 loc(#loc11) + %off_639 = arith.muli %off_638, %s : i32 loc(#loc10) + %off_640 = arith.addi %off_639, %pid : i32 loc(#loc11) + %off_641 = arith.muli %off_640, %s : i32 loc(#loc10) + %off_642 = arith.addi %off_641, %pid : i32 loc(#loc11) + %off_643 = arith.muli %off_642, %s : i32 loc(#loc10) + %off_644 = arith.addi %off_643, %pid : i32 loc(#loc11) + %off_645 = arith.muli %off_644, %s : i32 loc(#loc10) + %off_646 = arith.addi %off_645, %pid : i32 loc(#loc11) + %off_647 = arith.muli %off_646, %s : i32 loc(#loc10) + %off_648 = arith.addi %off_647, %pid : i32 loc(#loc11) + %off_649 = arith.muli %off_648, %s : i32 loc(#loc10) + %off_650 = arith.addi %off_649, %pid : i32 loc(#loc11) + %off_651 = arith.muli %off_650, %s : i32 loc(#loc10) + %off_652 = arith.addi %off_651, %pid : i32 loc(#loc11) + %off_653 = arith.muli %off_652, %s : i32 loc(#loc10) + %off_654 = arith.addi %off_653, %pid : i32 loc(#loc11) + %off_655 = arith.muli %off_654, %s : i32 loc(#loc10) + %off_656 = arith.addi %off_655, %pid : i32 loc(#loc11) + %off_657 = arith.muli %off_656, %s : i32 loc(#loc10) + %off_658 = arith.addi %off_657, %pid : i32 loc(#loc11) + %off_659 = arith.muli %off_658, %s : i32 loc(#loc10) + %off_660 = arith.addi %off_659, %pid : i32 loc(#loc11) + %off_661 = arith.muli %off_660, %s : i32 loc(#loc10) + %off_662 = arith.addi %off_661, %pid : i32 loc(#loc11) + %off_663 = arith.muli %off_662, %s : i32 loc(#loc10) + %off_664 = arith.addi %off_663, %pid : i32 loc(#loc11) + %off_665 = arith.muli %off_664, %s : i32 loc(#loc10) + %off_666 = arith.addi %off_665, %pid : i32 loc(#loc11) + %off_667 = arith.muli %off_666, %s : i32 loc(#loc10) + %off_668 = arith.addi %off_667, %pid : i32 loc(#loc11) + %off_669 = arith.muli %off_668, %s : i32 loc(#loc10) + %off_670 = arith.addi %off_669, %pid : i32 loc(#loc11) + %off_671 = arith.muli %off_670, %s : i32 loc(#loc10) + %off_672 = arith.addi %off_671, %pid : i32 loc(#loc11) + %off_673 = arith.muli %off_672, %s : i32 loc(#loc10) + %off_674 = arith.addi %off_673, %pid : i32 loc(#loc11) + %off_675 = arith.muli %off_674, %s : i32 loc(#loc10) + %off_676 = arith.addi %off_675, %pid : i32 loc(#loc11) + %off_677 = arith.muli %off_676, %s : i32 loc(#loc10) + %off_678 = arith.addi %off_677, %pid : i32 loc(#loc11) + %off_679 = arith.muli %off_678, %s : i32 loc(#loc10) + %off_680 = arith.addi %off_679, %pid : i32 loc(#loc11) + %off_681 = arith.muli %off_680, %s : i32 loc(#loc10) + %off_682 = arith.addi %off_681, %pid : i32 loc(#loc11) + %off_683 = arith.muli %off_682, %s : i32 loc(#loc10) + %off_684 = arith.addi %off_683, %pid : i32 loc(#loc11) + %off_685 = arith.muli %off_684, %s : i32 loc(#loc10) + %off_686 = arith.addi %off_685, %pid : i32 loc(#loc11) + %off_687 = arith.muli %off_686, %s : i32 loc(#loc10) + %off_688 = arith.addi %off_687, %pid : i32 loc(#loc11) + %off_689 = arith.muli %off_688, %s : i32 loc(#loc10) + %off_690 = arith.addi %off_689, %pid : i32 loc(#loc11) + %off_691 = arith.muli %off_690, %s : i32 loc(#loc10) + %off_692 = arith.addi %off_691, %pid : i32 loc(#loc11) + %off_693 = arith.muli %off_692, %s : i32 loc(#loc10) + %off_694 = arith.addi %off_693, %pid : i32 loc(#loc11) + %off_695 = arith.muli %off_694, %s : i32 loc(#loc10) + %off_696 = arith.addi %off_695, %pid : i32 loc(#loc11) + %off_697 = arith.muli %off_696, %s : i32 loc(#loc10) + %off_698 = arith.addi %off_697, %pid : i32 loc(#loc11) + %off_699 = arith.muli %off_698, %s : i32 loc(#loc10) + %off_700 = arith.addi %off_699, %pid : i32 loc(#loc11) + %off_701 = arith.muli %off_700, %s : i32 loc(#loc10) + %off_702 = arith.addi %off_701, %pid : i32 loc(#loc11) + %off_703 = arith.muli %off_702, %s : i32 loc(#loc10) + %off_704 = arith.addi %off_703, %pid : i32 loc(#loc11) + %off_705 = arith.muli %off_704, %s : i32 loc(#loc10) + %off_706 = arith.addi %off_705, %pid : i32 loc(#loc11) + %off_707 = arith.muli %off_706, %s : i32 loc(#loc10) + %off_708 = arith.addi %off_707, %pid : i32 loc(#loc11) + %off_709 = arith.muli %off_708, %s : i32 loc(#loc10) + %off_710 = arith.addi %off_709, %pid : i32 loc(#loc11) + %off_711 = arith.muli %off_710, %s : i32 loc(#loc10) + %off_712 = arith.addi %off_711, %pid : i32 loc(#loc11) + %off_713 = arith.muli %off_712, %s : i32 loc(#loc10) + %off_714 = arith.addi %off_713, %pid : i32 loc(#loc11) + %off_715 = arith.muli %off_714, %s : i32 loc(#loc10) + %off_716 = arith.addi %off_715, %pid : i32 loc(#loc11) + %off_717 = arith.muli %off_716, %s : i32 loc(#loc10) + %off_718 = arith.addi %off_717, %pid : i32 loc(#loc11) + %off_719 = arith.muli %off_718, %s : i32 loc(#loc10) + %off_720 = arith.addi %off_719, %pid : i32 loc(#loc11) + %off_721 = arith.muli %off_720, %s : i32 loc(#loc10) + %off_722 = arith.addi %off_721, %pid : i32 loc(#loc11) + %off_723 = arith.muli %off_722, %s : i32 loc(#loc10) + %off_724 = arith.addi %off_723, %pid : i32 loc(#loc11) + %off_725 = arith.muli %off_724, %s : i32 loc(#loc10) + %off_726 = arith.addi %off_725, %pid : i32 loc(#loc11) + %off_727 = arith.muli %off_726, %s : i32 loc(#loc10) + %off_728 = arith.addi %off_727, %pid : i32 loc(#loc11) + %off_729 = arith.muli %off_728, %s : i32 loc(#loc10) + %off_730 = arith.addi %off_729, %pid : i32 loc(#loc11) + %off_731 = arith.muli %off_730, %s : i32 loc(#loc10) + %off_732 = arith.addi %off_731, %pid : i32 loc(#loc11) + %off_733 = arith.muli %off_732, %s : i32 loc(#loc10) + %off_734 = arith.addi %off_733, %pid : i32 loc(#loc11) + %off_735 = arith.muli %off_734, %s : i32 loc(#loc10) + %off_736 = arith.addi %off_735, %pid : i32 loc(#loc11) + %off_737 = arith.muli %off_736, %s : i32 loc(#loc10) + %off_738 = arith.addi %off_737, %pid : i32 loc(#loc11) + %off_739 = arith.muli %off_738, %s : i32 loc(#loc10) + %off_740 = arith.addi %off_739, %pid : i32 loc(#loc11) + %off_741 = arith.muli %off_740, %s : i32 loc(#loc10) + %off_742 = arith.addi %off_741, %pid : i32 loc(#loc11) + %off_743 = arith.muli %off_742, %s : i32 loc(#loc10) + %off_744 = arith.addi %off_743, %pid : i32 loc(#loc11) + %off_745 = arith.muli %off_744, %s : i32 loc(#loc10) + %off_746 = arith.addi %off_745, %pid : i32 loc(#loc11) + %off_747 = arith.muli %off_746, %s : i32 loc(#loc10) + %off_748 = arith.addi %off_747, %pid : i32 loc(#loc11) + %off_749 = arith.muli %off_748, %s : i32 loc(#loc10) + %off_750 = arith.addi %off_749, %pid : i32 loc(#loc11) + %off_751 = arith.muli %off_750, %s : i32 loc(#loc10) + %off_752 = arith.addi %off_751, %pid : i32 loc(#loc11) + %off_753 = arith.muli %off_752, %s : i32 loc(#loc10) + %off_754 = arith.addi %off_753, %pid : i32 loc(#loc11) + %off_755 = arith.muli %off_754, %s : i32 loc(#loc10) + %off_756 = arith.addi %off_755, %pid : i32 loc(#loc11) + %off_757 = arith.muli %off_756, %s : i32 loc(#loc10) + %off_758 = arith.addi %off_757, %pid : i32 loc(#loc11) + %off_759 = arith.muli %off_758, %s : i32 loc(#loc10) + %off_760 = arith.addi %off_759, %pid : i32 loc(#loc11) + %off_761 = arith.muli %off_760, %s : i32 loc(#loc10) + %off_762 = arith.addi %off_761, %pid : i32 loc(#loc11) + %off_763 = arith.muli %off_762, %s : i32 loc(#loc10) + %off_764 = arith.addi %off_763, %pid : i32 loc(#loc11) + %off_765 = arith.muli %off_764, %s : i32 loc(#loc10) + %off_766 = arith.addi %off_765, %pid : i32 loc(#loc11) + %off_767 = arith.muli %off_766, %s : i32 loc(#loc10) + %off_768 = arith.addi %off_767, %pid : i32 loc(#loc11) + %off_769 = arith.muli %off_768, %s : i32 loc(#loc10) + %off_770 = arith.addi %off_769, %pid : i32 loc(#loc11) + %off_771 = arith.muli %off_770, %s : i32 loc(#loc10) + %off_772 = arith.addi %off_771, %pid : i32 loc(#loc11) + %off_773 = arith.muli %off_772, %s : i32 loc(#loc10) + %off_774 = arith.addi %off_773, %pid : i32 loc(#loc11) + %off_775 = arith.muli %off_774, %s : i32 loc(#loc10) + %off_776 = arith.addi %off_775, %pid : i32 loc(#loc11) + %off_777 = arith.muli %off_776, %s : i32 loc(#loc10) + %off_778 = arith.addi %off_777, %pid : i32 loc(#loc11) + %off_779 = arith.muli %off_778, %s : i32 loc(#loc10) + %off_780 = arith.addi %off_779, %pid : i32 loc(#loc11) + %off_781 = arith.muli %off_780, %s : i32 loc(#loc10) + %off_782 = arith.addi %off_781, %pid : i32 loc(#loc11) + %off_783 = arith.muli %off_782, %s : i32 loc(#loc10) + %off_784 = arith.addi %off_783, %pid : i32 loc(#loc11) + %off_785 = arith.muli %off_784, %s : i32 loc(#loc10) + %off_786 = arith.addi %off_785, %pid : i32 loc(#loc11) + %off_787 = arith.muli %off_786, %s : i32 loc(#loc10) + %off_788 = arith.addi %off_787, %pid : i32 loc(#loc11) + %off_789 = arith.muli %off_788, %s : i32 loc(#loc10) + %off_790 = arith.addi %off_789, %pid : i32 loc(#loc11) + %off_791 = arith.muli %off_790, %s : i32 loc(#loc10) + %off_792 = arith.addi %off_791, %pid : i32 loc(#loc11) + %off_793 = arith.muli %off_792, %s : i32 loc(#loc10) + %off_794 = arith.addi %off_793, %pid : i32 loc(#loc11) + %off_795 = arith.muli %off_794, %s : i32 loc(#loc10) + %off_796 = arith.addi %off_795, %pid : i32 loc(#loc11) + %off_797 = arith.muli %off_796, %s : i32 loc(#loc10) + %off_798 = arith.addi %off_797, %pid : i32 loc(#loc11) + %off_799 = arith.muli %off_798, %s : i32 loc(#loc10) + %off_800 = arith.addi %off_799, %pid : i32 loc(#loc11) + %off_801 = arith.muli %off_800, %s : i32 loc(#loc10) + %off_802 = arith.addi %off_801, %pid : i32 loc(#loc11) + %off_803 = arith.muli %off_802, %s : i32 loc(#loc10) + %off_804 = arith.addi %off_803, %pid : i32 loc(#loc11) + %off_805 = arith.muli %off_804, %s : i32 loc(#loc10) + %off_806 = arith.addi %off_805, %pid : i32 loc(#loc11) + %off_807 = arith.muli %off_806, %s : i32 loc(#loc10) + %off_808 = arith.addi %off_807, %pid : i32 loc(#loc11) + %off_809 = arith.muli %off_808, %s : i32 loc(#loc10) + %off_810 = arith.addi %off_809, %pid : i32 loc(#loc11) + %off_811 = arith.muli %off_810, %s : i32 loc(#loc10) + %off_812 = arith.addi %off_811, %pid : i32 loc(#loc11) + %off_813 = arith.muli %off_812, %s : i32 loc(#loc10) + %off_814 = arith.addi %off_813, %pid : i32 loc(#loc11) + %off_815 = arith.muli %off_814, %s : i32 loc(#loc10) + %off_816 = arith.addi %off_815, %pid : i32 loc(#loc11) + %off_817 = arith.muli %off_816, %s : i32 loc(#loc10) + %off_818 = arith.addi %off_817, %pid : i32 loc(#loc11) + %off_819 = arith.muli %off_818, %s : i32 loc(#loc10) + %off_820 = arith.addi %off_819, %pid : i32 loc(#loc11) + %off_821 = arith.muli %off_820, %s : i32 loc(#loc10) + %off_822 = arith.addi %off_821, %pid : i32 loc(#loc11) + %off_823 = arith.muli %off_822, %s : i32 loc(#loc10) + %off_824 = arith.addi %off_823, %pid : i32 loc(#loc11) + %off_825 = arith.muli %off_824, %s : i32 loc(#loc10) + %off_826 = arith.addi %off_825, %pid : i32 loc(#loc11) + %off_827 = arith.muli %off_826, %s : i32 loc(#loc10) + %off_828 = arith.addi %off_827, %pid : i32 loc(#loc11) + %off_829 = arith.muli %off_828, %s : i32 loc(#loc10) + %off_830 = arith.addi %off_829, %pid : i32 loc(#loc11) + %off_831 = arith.muli %off_830, %s : i32 loc(#loc10) + %off_832 = arith.addi %off_831, %pid : i32 loc(#loc11) + %off_833 = arith.muli %off_832, %s : i32 loc(#loc10) + %off_834 = arith.addi %off_833, %pid : i32 loc(#loc11) + %off_835 = arith.muli %off_834, %s : i32 loc(#loc10) + %off_836 = arith.addi %off_835, %pid : i32 loc(#loc11) + %off_837 = arith.muli %off_836, %s : i32 loc(#loc10) + %off_838 = arith.addi %off_837, %pid : i32 loc(#loc11) + %off_839 = arith.muli %off_838, %s : i32 loc(#loc10) + %off_840 = arith.addi %off_839, %pid : i32 loc(#loc11) + %off_841 = arith.muli %off_840, %s : i32 loc(#loc10) + %off_842 = arith.addi %off_841, %pid : i32 loc(#loc11) + %off_843 = arith.muli %off_842, %s : i32 loc(#loc10) + %off_844 = arith.addi %off_843, %pid : i32 loc(#loc11) + %off_845 = arith.muli %off_844, %s : i32 loc(#loc10) + %off_846 = arith.addi %off_845, %pid : i32 loc(#loc11) + %off_847 = arith.muli %off_846, %s : i32 loc(#loc10) + %off_848 = arith.addi %off_847, %pid : i32 loc(#loc11) + %off_849 = arith.muli %off_848, %s : i32 loc(#loc10) + %off_850 = arith.addi %off_849, %pid : i32 loc(#loc11) + %off_851 = arith.muli %off_850, %s : i32 loc(#loc10) + %off_852 = arith.addi %off_851, %pid : i32 loc(#loc11) + %off_853 = arith.muli %off_852, %s : i32 loc(#loc10) + %off_854 = arith.addi %off_853, %pid : i32 loc(#loc11) + %off_855 = arith.muli %off_854, %s : i32 loc(#loc10) + %off_856 = arith.addi %off_855, %pid : i32 loc(#loc11) + %off_857 = arith.muli %off_856, %s : i32 loc(#loc10) + %off_858 = arith.addi %off_857, %pid : i32 loc(#loc11) + %off_859 = arith.muli %off_858, %s : i32 loc(#loc10) + %off_860 = arith.addi %off_859, %pid : i32 loc(#loc11) + %off_861 = arith.muli %off_860, %s : i32 loc(#loc10) + %off_862 = arith.addi %off_861, %pid : i32 loc(#loc11) + %off_863 = arith.muli %off_862, %s : i32 loc(#loc10) + %off_864 = arith.addi %off_863, %pid : i32 loc(#loc11) + %off_865 = arith.muli %off_864, %s : i32 loc(#loc10) + %off_866 = arith.addi %off_865, %pid : i32 loc(#loc11) + %off_867 = arith.muli %off_866, %s : i32 loc(#loc10) + %off_868 = arith.addi %off_867, %pid : i32 loc(#loc11) + %off_869 = arith.muli %off_868, %s : i32 loc(#loc10) + %off_870 = arith.addi %off_869, %pid : i32 loc(#loc11) + %off_871 = arith.muli %off_870, %s : i32 loc(#loc10) + %off_872 = arith.addi %off_871, %pid : i32 loc(#loc11) + %off_873 = arith.muli %off_872, %s : i32 loc(#loc10) + %off_874 = arith.addi %off_873, %pid : i32 loc(#loc11) + %off_875 = arith.muli %off_874, %s : i32 loc(#loc10) + %off_876 = arith.addi %off_875, %pid : i32 loc(#loc11) + %off_877 = arith.muli %off_876, %s : i32 loc(#loc10) + %off_878 = arith.addi %off_877, %pid : i32 loc(#loc11) + %off_879 = arith.muli %off_878, %s : i32 loc(#loc10) + %off_880 = arith.addi %off_879, %pid : i32 loc(#loc11) + %off_881 = arith.muli %off_880, %s : i32 loc(#loc10) + %off_882 = arith.addi %off_881, %pid : i32 loc(#loc11) + %off_883 = arith.muli %off_882, %s : i32 loc(#loc10) + %off_884 = arith.addi %off_883, %pid : i32 loc(#loc11) + %off_885 = arith.muli %off_884, %s : i32 loc(#loc10) + %off_886 = arith.addi %off_885, %pid : i32 loc(#loc11) + %off_887 = arith.muli %off_886, %s : i32 loc(#loc10) + %off_888 = arith.addi %off_887, %pid : i32 loc(#loc11) + %off_889 = arith.muli %off_888, %s : i32 loc(#loc10) + %off_890 = arith.addi %off_889, %pid : i32 loc(#loc11) + %off_891 = arith.muli %off_890, %s : i32 loc(#loc10) + %off_892 = arith.addi %off_891, %pid : i32 loc(#loc11) + %off_893 = arith.muli %off_892, %s : i32 loc(#loc10) + %off_894 = arith.addi %off_893, %pid : i32 loc(#loc11) + %off_895 = arith.muli %off_894, %s : i32 loc(#loc10) + %off_896 = arith.addi %off_895, %pid : i32 loc(#loc11) + %off_897 = arith.muli %off_896, %s : i32 loc(#loc10) + %off_898 = arith.addi %off_897, %pid : i32 loc(#loc11) + %off_899 = arith.muli %off_898, %s : i32 loc(#loc10) + %off_900 = arith.addi %off_899, %pid : i32 loc(#loc11) + %off_901 = arith.muli %off_900, %s : i32 loc(#loc10) + %off_902 = arith.addi %off_901, %pid : i32 loc(#loc11) + %off_903 = arith.muli %off_902, %s : i32 loc(#loc10) + %off_904 = arith.addi %off_903, %pid : i32 loc(#loc11) + %off_905 = arith.muli %off_904, %s : i32 loc(#loc10) + %off_906 = arith.addi %off_905, %pid : i32 loc(#loc11) + %off_907 = arith.muli %off_906, %s : i32 loc(#loc10) + %off_908 = arith.addi %off_907, %pid : i32 loc(#loc11) + %off_909 = arith.muli %off_908, %s : i32 loc(#loc10) + %off_910 = arith.addi %off_909, %pid : i32 loc(#loc11) + %off_911 = arith.muli %off_910, %s : i32 loc(#loc10) + %off_912 = arith.addi %off_911, %pid : i32 loc(#loc11) + %off_913 = arith.muli %off_912, %s : i32 loc(#loc10) + %off_914 = arith.addi %off_913, %pid : i32 loc(#loc11) + %off_915 = arith.muli %off_914, %s : i32 loc(#loc10) + %off_916 = arith.addi %off_915, %pid : i32 loc(#loc11) + %off_917 = arith.muli %off_916, %s : i32 loc(#loc10) + %off_918 = arith.addi %off_917, %pid : i32 loc(#loc11) + %off_919 = arith.muli %off_918, %s : i32 loc(#loc10) + %off_920 = arith.addi %off_919, %pid : i32 loc(#loc11) + %off_921 = arith.muli %off_920, %s : i32 loc(#loc10) + %off_922 = arith.addi %off_921, %pid : i32 loc(#loc11) + %off_923 = arith.muli %off_922, %s : i32 loc(#loc10) + %off_924 = arith.addi %off_923, %pid : i32 loc(#loc11) + %off_925 = arith.muli %off_924, %s : i32 loc(#loc10) + %off_926 = arith.addi %off_925, %pid : i32 loc(#loc11) + %off_927 = arith.muli %off_926, %s : i32 loc(#loc10) + %off_928 = arith.addi %off_927, %pid : i32 loc(#loc11) + %off_929 = arith.muli %off_928, %s : i32 loc(#loc10) + %off_930 = arith.addi %off_929, %pid : i32 loc(#loc11) + %off_931 = arith.muli %off_930, %s : i32 loc(#loc10) + %off_932 = arith.addi %off_931, %pid : i32 loc(#loc11) + %off_933 = arith.muli %off_932, %s : i32 loc(#loc10) + %off_934 = arith.addi %off_933, %pid : i32 loc(#loc11) + %off_935 = arith.muli %off_934, %s : i32 loc(#loc10) + %off_936 = arith.addi %off_935, %pid : i32 loc(#loc11) + %off_937 = arith.muli %off_936, %s : i32 loc(#loc10) + %off_938 = arith.addi %off_937, %pid : i32 loc(#loc11) + %off_939 = arith.muli %off_938, %s : i32 loc(#loc10) + %off_940 = arith.addi %off_939, %pid : i32 loc(#loc11) + %off_941 = arith.muli %off_940, %s : i32 loc(#loc10) + %off_942 = arith.addi %off_941, %pid : i32 loc(#loc11) + %off_943 = arith.muli %off_942, %s : i32 loc(#loc10) + %off_944 = arith.addi %off_943, %pid : i32 loc(#loc11) + %off_945 = arith.muli %off_944, %s : i32 loc(#loc10) + %off_946 = arith.addi %off_945, %pid : i32 loc(#loc11) + %off_947 = arith.muli %off_946, %s : i32 loc(#loc10) + %off_948 = arith.addi %off_947, %pid : i32 loc(#loc11) + %off_949 = arith.muli %off_948, %s : i32 loc(#loc10) + %off_950 = arith.addi %off_949, %pid : i32 loc(#loc11) + %off_951 = arith.muli %off_950, %s : i32 loc(#loc10) + %off_952 = arith.addi %off_951, %pid : i32 loc(#loc11) + %off_953 = arith.muli %off_952, %s : i32 loc(#loc10) + %off_954 = arith.addi %off_953, %pid : i32 loc(#loc11) + %off_955 = arith.muli %off_954, %s : i32 loc(#loc10) + %off_956 = arith.addi %off_955, %pid : i32 loc(#loc11) + %off_957 = arith.muli %off_956, %s : i32 loc(#loc10) + %off_958 = arith.addi %off_957, %pid : i32 loc(#loc11) + %off_959 = arith.muli %off_958, %s : i32 loc(#loc10) + %off_960 = arith.addi %off_959, %pid : i32 loc(#loc11) + %off_961 = arith.muli %off_960, %s : i32 loc(#loc10) + %off_962 = arith.addi %off_961, %pid : i32 loc(#loc11) + %off_963 = arith.muli %off_962, %s : i32 loc(#loc10) + %off_964 = arith.addi %off_963, %pid : i32 loc(#loc11) + %off_965 = arith.muli %off_964, %s : i32 loc(#loc10) + %off_966 = arith.addi %off_965, %pid : i32 loc(#loc11) + %off_967 = arith.muli %off_966, %s : i32 loc(#loc10) + %off_968 = arith.addi %off_967, %pid : i32 loc(#loc11) + %off_969 = arith.muli %off_968, %s : i32 loc(#loc10) + %off_970 = arith.addi %off_969, %pid : i32 loc(#loc11) + %off_971 = arith.muli %off_970, %s : i32 loc(#loc10) + %off_972 = arith.addi %off_971, %pid : i32 loc(#loc11) + %off_973 = arith.muli %off_972, %s : i32 loc(#loc10) + %off_974 = arith.addi %off_973, %pid : i32 loc(#loc11) + %off_975 = arith.muli %off_974, %s : i32 loc(#loc10) + %off_976 = arith.addi %off_975, %pid : i32 loc(#loc11) + %off_977 = arith.muli %off_976, %s : i32 loc(#loc10) + %off_978 = arith.addi %off_977, %pid : i32 loc(#loc11) + %off_979 = arith.muli %off_978, %s : i32 loc(#loc10) + %off_980 = arith.addi %off_979, %pid : i32 loc(#loc11) + %off_981 = arith.muli %off_980, %s : i32 loc(#loc10) + %off_982 = arith.addi %off_981, %pid : i32 loc(#loc11) + %off_983 = arith.muli %off_982, %s : i32 loc(#loc10) + %off_984 = arith.addi %off_983, %pid : i32 loc(#loc11) + %off_985 = arith.muli %off_984, %s : i32 loc(#loc10) + %off_986 = arith.addi %off_985, %pid : i32 loc(#loc11) + %off_987 = arith.muli %off_986, %s : i32 loc(#loc10) + %off_988 = arith.addi %off_987, %pid : i32 loc(#loc11) + %off_989 = arith.muli %off_988, %s : i32 loc(#loc10) + %off_990 = arith.addi %off_989, %pid : i32 loc(#loc11) + %off_991 = arith.muli %off_990, %s : i32 loc(#loc10) + %off_992 = arith.addi %off_991, %pid : i32 loc(#loc11) + %off_993 = arith.muli %off_992, %s : i32 loc(#loc10) + %off_994 = arith.addi %off_993, %pid : i32 loc(#loc11) + %off_995 = arith.muli %off_994, %s : i32 loc(#loc10) + %off_996 = arith.addi %off_995, %pid : i32 loc(#loc11) + %off_997 = arith.muli %off_996, %s : i32 loc(#loc10) + %off_998 = arith.addi %off_997, %pid : i32 loc(#loc11) + %off_999 = arith.muli %off_998, %s : i32 loc(#loc10) + %off_1000 = arith.addi %off_999, %pid : i32 loc(#loc11) + %off_1001 = arith.muli %off_1000, %s : i32 loc(#loc10) + %off_1002 = arith.addi %off_1001, %pid : i32 loc(#loc11) + %off_1003 = arith.muli %off_1002, %s : i32 loc(#loc10) + %off_1004 = arith.addi %off_1003, %pid : i32 loc(#loc11) + %off_1005 = arith.muli %off_1004, %s : i32 loc(#loc10) + %off_1006 = arith.addi %off_1005, %pid : i32 loc(#loc11) + %off_1007 = arith.muli %off_1006, %s : i32 loc(#loc10) + %off_1008 = arith.addi %off_1007, %pid : i32 loc(#loc11) + %off_1009 = arith.muli %off_1008, %s : i32 loc(#loc10) + %off_1010 = arith.addi %off_1009, %pid : i32 loc(#loc11) + %off_1011 = arith.muli %off_1010, %s : i32 loc(#loc10) + %off_1012 = arith.addi %off_1011, %pid : i32 loc(#loc11) + %off_1013 = arith.muli %off_1012, %s : i32 loc(#loc10) + %off_1014 = arith.addi %off_1013, %pid : i32 loc(#loc11) + %off_1015 = arith.muli %off_1014, %s : i32 loc(#loc10) + %off_1016 = arith.addi %off_1015, %pid : i32 loc(#loc11) + %off_1017 = arith.muli %off_1016, %s : i32 loc(#loc10) + %off_1018 = arith.addi %off_1017, %pid : i32 loc(#loc11) + %off_1019 = arith.muli %off_1018, %s : i32 loc(#loc10) + %off_1020 = arith.addi %off_1019, %pid : i32 loc(#loc11) + %off_1021 = arith.muli %off_1020, %s : i32 loc(#loc10) + %off_1022 = arith.addi %off_1021, %pid : i32 loc(#loc11) + %off_1023 = arith.muli %off_1022, %s : i32 loc(#loc10) + %off_1024 = arith.addi %off_1023, %pid : i32 loc(#loc11) + %off_1025 = arith.muli %off_1024, %s : i32 loc(#loc10) + %off_1026 = arith.addi %off_1025, %pid : i32 loc(#loc11) + %off_1027 = arith.muli %off_1026, %s : i32 loc(#loc10) + %off_1028 = arith.addi %off_1027, %pid : i32 loc(#loc11) + %off_1029 = arith.muli %off_1028, %s : i32 loc(#loc10) + %off_1030 = arith.addi %off_1029, %pid : i32 loc(#loc11) + %off_1031 = arith.muli %off_1030, %s : i32 loc(#loc10) + %off_1032 = arith.addi %off_1031, %pid : i32 loc(#loc11) + %off_1033 = arith.muli %off_1032, %s : i32 loc(#loc10) + %off_1034 = arith.addi %off_1033, %pid : i32 loc(#loc11) + %off_1035 = arith.muli %off_1034, %s : i32 loc(#loc10) + %off_1036 = arith.addi %off_1035, %pid : i32 loc(#loc11) + %off_1037 = arith.muli %off_1036, %s : i32 loc(#loc10) + %off_1038 = arith.addi %off_1037, %pid : i32 loc(#loc11) + %off_1039 = arith.muli %off_1038, %s : i32 loc(#loc10) + %off_1040 = arith.addi %off_1039, %pid : i32 loc(#loc11) + %off_1041 = arith.muli %off_1040, %s : i32 loc(#loc10) + %off_1042 = arith.addi %off_1041, %pid : i32 loc(#loc11) + %off_1043 = arith.muli %off_1042, %s : i32 loc(#loc10) + %off_1044 = arith.addi %off_1043, %pid : i32 loc(#loc11) + %off_1045 = arith.muli %off_1044, %s : i32 loc(#loc10) + %off_1046 = arith.addi %off_1045, %pid : i32 loc(#loc11) + %off_1047 = arith.muli %off_1046, %s : i32 loc(#loc10) + %off_1048 = arith.addi %off_1047, %pid : i32 loc(#loc11) + %off_1049 = arith.muli %off_1048, %s : i32 loc(#loc10) + %off_1050 = arith.addi %off_1049, %pid : i32 loc(#loc11) + %off_1051 = arith.muli %off_1050, %s : i32 loc(#loc10) + %off_1052 = arith.addi %off_1051, %pid : i32 loc(#loc11) + %off_1053 = arith.muli %off_1052, %s : i32 loc(#loc10) + %off_1054 = arith.addi %off_1053, %pid : i32 loc(#loc11) + %off_1055 = arith.muli %off_1054, %s : i32 loc(#loc10) + %off_1056 = arith.addi %off_1055, %pid : i32 loc(#loc11) + %off_1057 = arith.muli %off_1056, %s : i32 loc(#loc10) + %off_1058 = arith.addi %off_1057, %pid : i32 loc(#loc11) + %off_1059 = arith.muli %off_1058, %s : i32 loc(#loc10) + %off_1060 = arith.addi %off_1059, %pid : i32 loc(#loc11) + %off_1061 = arith.muli %off_1060, %s : i32 loc(#loc10) + %off_1062 = arith.addi %off_1061, %pid : i32 loc(#loc11) + %off_1063 = arith.muli %off_1062, %s : i32 loc(#loc10) + %off_1064 = arith.addi %off_1063, %pid : i32 loc(#loc11) + %off_1065 = arith.muli %off_1064, %s : i32 loc(#loc10) + %off_1066 = arith.addi %off_1065, %pid : i32 loc(#loc11) + %off_1067 = arith.muli %off_1066, %s : i32 loc(#loc10) + %off_1068 = arith.addi %off_1067, %pid : i32 loc(#loc11) + %off_1069 = arith.muli %off_1068, %s : i32 loc(#loc10) + %off_1070 = arith.addi %off_1069, %pid : i32 loc(#loc11) + %off_1071 = arith.muli %off_1070, %s : i32 loc(#loc10) + %off_1072 = arith.addi %off_1071, %pid : i32 loc(#loc11) + %off_1073 = arith.muli %off_1072, %s : i32 loc(#loc10) + %off_1074 = arith.addi %off_1073, %pid : i32 loc(#loc11) + %off_1075 = arith.muli %off_1074, %s : i32 loc(#loc10) + %off_1076 = arith.addi %off_1075, %pid : i32 loc(#loc11) + %off_1077 = arith.muli %off_1076, %s : i32 loc(#loc10) + %off_1078 = arith.addi %off_1077, %pid : i32 loc(#loc11) + %off_1079 = arith.muli %off_1078, %s : i32 loc(#loc10) + %off_1080 = arith.addi %off_1079, %pid : i32 loc(#loc11) + %off_1081 = arith.muli %off_1080, %s : i32 loc(#loc10) + %off_1082 = arith.addi %off_1081, %pid : i32 loc(#loc11) + %off_1083 = arith.muli %off_1082, %s : i32 loc(#loc10) + %off_1084 = arith.addi %off_1083, %pid : i32 loc(#loc11) + %off_1085 = arith.muli %off_1084, %s : i32 loc(#loc10) + %off_1086 = arith.addi %off_1085, %pid : i32 loc(#loc11) + %off_1087 = arith.muli %off_1086, %s : i32 loc(#loc10) + %off_1088 = arith.addi %off_1087, %pid : i32 loc(#loc11) + %off_1089 = arith.muli %off_1088, %s : i32 loc(#loc10) + %off_1090 = arith.addi %off_1089, %pid : i32 loc(#loc11) + %off_1091 = arith.muli %off_1090, %s : i32 loc(#loc10) + %off_1092 = arith.addi %off_1091, %pid : i32 loc(#loc11) + %off_1093 = arith.muli %off_1092, %s : i32 loc(#loc10) + %off_1094 = arith.addi %off_1093, %pid : i32 loc(#loc11) + %off_1095 = arith.muli %off_1094, %s : i32 loc(#loc10) + %off_1096 = arith.addi %off_1095, %pid : i32 loc(#loc11) + %off_1097 = arith.muli %off_1096, %s : i32 loc(#loc10) + %off_1098 = arith.addi %off_1097, %pid : i32 loc(#loc11) + %off_1099 = arith.muli %off_1098, %s : i32 loc(#loc10) + %off_1100 = arith.addi %off_1099, %pid : i32 loc(#loc11) + %off_1101 = arith.muli %off_1100, %s : i32 loc(#loc10) + %off_1102 = arith.addi %off_1101, %pid : i32 loc(#loc11) + %off_1103 = arith.muli %off_1102, %s : i32 loc(#loc10) + %off_1104 = arith.addi %off_1103, %pid : i32 loc(#loc11) + %off_1105 = arith.muli %off_1104, %s : i32 loc(#loc10) + %off_1106 = arith.addi %off_1105, %pid : i32 loc(#loc11) + %off_1107 = arith.muli %off_1106, %s : i32 loc(#loc10) + %off_1108 = arith.addi %off_1107, %pid : i32 loc(#loc11) + %off_1109 = arith.muli %off_1108, %s : i32 loc(#loc10) + %off_1110 = arith.addi %off_1109, %pid : i32 loc(#loc11) + %off_1111 = arith.muli %off_1110, %s : i32 loc(#loc10) + %off_1112 = arith.addi %off_1111, %pid : i32 loc(#loc11) + %off_1113 = arith.muli %off_1112, %s : i32 loc(#loc10) + %off_1114 = arith.addi %off_1113, %pid : i32 loc(#loc11) + %off_1115 = arith.muli %off_1114, %s : i32 loc(#loc10) + %off_1116 = arith.addi %off_1115, %pid : i32 loc(#loc11) + %off_1117 = arith.muli %off_1116, %s : i32 loc(#loc10) + %off_1118 = arith.addi %off_1117, %pid : i32 loc(#loc11) + %off_1119 = arith.muli %off_1118, %s : i32 loc(#loc10) + %off_1120 = arith.addi %off_1119, %pid : i32 loc(#loc11) + %off_1121 = arith.muli %off_1120, %s : i32 loc(#loc10) + %off_1122 = arith.addi %off_1121, %pid : i32 loc(#loc11) + %off_1123 = arith.muli %off_1122, %s : i32 loc(#loc10) + %off_1124 = arith.addi %off_1123, %pid : i32 loc(#loc11) + %off_1125 = arith.muli %off_1124, %s : i32 loc(#loc10) + %off_1126 = arith.addi %off_1125, %pid : i32 loc(#loc11) + %off_1127 = arith.muli %off_1126, %s : i32 loc(#loc10) + %off_1128 = arith.addi %off_1127, %pid : i32 loc(#loc11) + %off_1129 = arith.muli %off_1128, %s : i32 loc(#loc10) + %off_1130 = arith.addi %off_1129, %pid : i32 loc(#loc11) + %off_1131 = arith.muli %off_1130, %s : i32 loc(#loc10) + %off_1132 = arith.addi %off_1131, %pid : i32 loc(#loc11) + %off_1133 = arith.muli %off_1132, %s : i32 loc(#loc10) + %off_1134 = arith.addi %off_1133, %pid : i32 loc(#loc11) + %off_1135 = arith.muli %off_1134, %s : i32 loc(#loc10) + %off_1136 = arith.addi %off_1135, %pid : i32 loc(#loc11) + %off_1137 = arith.muli %off_1136, %s : i32 loc(#loc10) + %off_1138 = arith.addi %off_1137, %pid : i32 loc(#loc11) + %off_1139 = arith.muli %off_1138, %s : i32 loc(#loc10) + %off_1140 = arith.addi %off_1139, %pid : i32 loc(#loc11) + %off_1141 = arith.muli %off_1140, %s : i32 loc(#loc10) + %off_1142 = arith.addi %off_1141, %pid : i32 loc(#loc11) + %off_1143 = arith.muli %off_1142, %s : i32 loc(#loc10) + %off_1144 = arith.addi %off_1143, %pid : i32 loc(#loc11) + %off_1145 = arith.muli %off_1144, %s : i32 loc(#loc10) + %off_1146 = arith.addi %off_1145, %pid : i32 loc(#loc11) + %off_1147 = arith.muli %off_1146, %s : i32 loc(#loc10) + %off_1148 = arith.addi %off_1147, %pid : i32 loc(#loc11) + %off_1149 = arith.muli %off_1148, %s : i32 loc(#loc10) + %off_1150 = arith.addi %off_1149, %pid : i32 loc(#loc11) + %off_1151 = arith.muli %off_1150, %s : i32 loc(#loc10) + %off_1152 = arith.addi %off_1151, %pid : i32 loc(#loc11) + %off_1153 = arith.muli %off_1152, %s : i32 loc(#loc10) + %off_1154 = arith.addi %off_1153, %pid : i32 loc(#loc11) + %off_1155 = arith.muli %off_1154, %s : i32 loc(#loc10) + %off_1156 = arith.addi %off_1155, %pid : i32 loc(#loc11) + %off_1157 = arith.muli %off_1156, %s : i32 loc(#loc10) + %off_1158 = arith.addi %off_1157, %pid : i32 loc(#loc11) + %off_1159 = arith.muli %off_1158, %s : i32 loc(#loc10) + %off_1160 = arith.addi %off_1159, %pid : i32 loc(#loc11) + %off_1161 = arith.muli %off_1160, %s : i32 loc(#loc10) + %off_1162 = arith.addi %off_1161, %pid : i32 loc(#loc11) + %off_1163 = arith.muli %off_1162, %s : i32 loc(#loc10) + %off_1164 = arith.addi %off_1163, %pid : i32 loc(#loc11) + %off_1165 = arith.muli %off_1164, %s : i32 loc(#loc10) + %off_1166 = arith.addi %off_1165, %pid : i32 loc(#loc11) + %off_1167 = arith.muli %off_1166, %s : i32 loc(#loc10) + %off_1168 = arith.addi %off_1167, %pid : i32 loc(#loc11) + %off_1169 = arith.muli %off_1168, %s : i32 loc(#loc10) + %off_1170 = arith.addi %off_1169, %pid : i32 loc(#loc11) + %off_1171 = arith.muli %off_1170, %s : i32 loc(#loc10) + %off_1172 = arith.addi %off_1171, %pid : i32 loc(#loc11) + %off_1173 = arith.muli %off_1172, %s : i32 loc(#loc10) + %off_1174 = arith.addi %off_1173, %pid : i32 loc(#loc11) + %off_1175 = arith.muli %off_1174, %s : i32 loc(#loc10) + %off_1176 = arith.addi %off_1175, %pid : i32 loc(#loc11) + %off_1177 = arith.muli %off_1176, %s : i32 loc(#loc10) + %off_1178 = arith.addi %off_1177, %pid : i32 loc(#loc11) + %off_1179 = arith.muli %off_1178, %s : i32 loc(#loc10) + %off_1180 = arith.addi %off_1179, %pid : i32 loc(#loc11) + %off_1181 = arith.muli %off_1180, %s : i32 loc(#loc10) + %off_1182 = arith.addi %off_1181, %pid : i32 loc(#loc11) + %off_1183 = arith.muli %off_1182, %s : i32 loc(#loc10) + %off_1184 = arith.addi %off_1183, %pid : i32 loc(#loc11) + %off_1185 = arith.muli %off_1184, %s : i32 loc(#loc10) + %off_1186 = arith.addi %off_1185, %pid : i32 loc(#loc11) + %off_1187 = arith.muli %off_1186, %s : i32 loc(#loc10) + %off_1188 = arith.addi %off_1187, %pid : i32 loc(#loc11) + %off_1189 = arith.muli %off_1188, %s : i32 loc(#loc10) + %off_1190 = arith.addi %off_1189, %pid : i32 loc(#loc11) + %off_1191 = arith.muli %off_1190, %s : i32 loc(#loc10) + %off_1192 = arith.addi %off_1191, %pid : i32 loc(#loc11) + %off_1193 = arith.muli %off_1192, %s : i32 loc(#loc10) + %off_1194 = arith.addi %off_1193, %pid : i32 loc(#loc11) + %off_1195 = arith.muli %off_1194, %s : i32 loc(#loc10) + %off_1196 = arith.addi %off_1195, %pid : i32 loc(#loc11) + %off_1197 = arith.muli %off_1196, %s : i32 loc(#loc10) + %off_1198 = arith.addi %off_1197, %pid : i32 loc(#loc11) + %0 = tt.addptr %out_ptr, %off_1198 : !tt.ptr, i32 loc(#loc5) + tt.store %0, %cst : !tt.ptr loc(#loc1) + tt.return loc(#loc6) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":66:28) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":62:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":65:20) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":65:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":66:23) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":66:4) +#loc9 = loc("pid"(#loc2)) +#loc10 = loc("off"(#loc3)) +#loc11 = loc("off"(#loc4)) diff --git a/tests/golden/ir/ttir/kernel_dot_precisions.ttir b/tests/golden/ir/ttir/kernel_dot_precisions.ttir new file mode 100644 index 000000000..cc128d22d --- /dev/null +++ b/tests/golden/ir/ttir/kernel_dot_precisions.ttir @@ -0,0 +1,55 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":29:0) +#loc16 = loc("a_ptr"(#loc)) +#loc17 = loc("b_ptr"(#loc)) +#loc18 = loc("c_ptr"(#loc)) +module { + tt.func public @dot_precisions(%a_ptr: !tt.ptr loc("a_ptr"(#loc)), %b_ptr: !tt.ptr loc("b_ptr"(#loc)), %c_ptr: !tt.ptr loc("c_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<0.000000e+00> : tensor<16x16xf32> loc(#loc1) + %idx = arith.constant dense<16> : tensor<16x1xi32> loc(#loc19) + %offs = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32> loc(#loc20) + %idx_0 = tt.expand_dims %offs {axis = 1 : i32} : tensor<16xi32> -> tensor<16x1xi32> loc(#loc21) + %idx_1 = arith.muli %idx_0, %idx : tensor<16x1xi32> loc(#loc19) + %idx_2 = tt.expand_dims %offs {axis = 0 : i32} : tensor<16xi32> -> tensor<1x16xi32> loc(#loc22) + %idx_3 = tt.broadcast %idx_1 : tensor<16x1xi32> -> tensor<16x16xi32> loc(#loc23) + %idx_4 = tt.broadcast %idx_2 : tensor<1x16xi32> -> tensor<16x16xi32> loc(#loc23) + %idx_5 = arith.addi %idx_3, %idx_4 : tensor<16x16xi32> loc(#loc23) + %a = tt.splat %a_ptr : !tt.ptr -> tensor<16x16x!tt.ptr> loc(#loc24) + %a_6 = tt.addptr %a, %idx_5 : tensor<16x16x!tt.ptr>, tensor<16x16xi32> loc(#loc24) + %a_7 = tt.load %a_6 : tensor<16x16x!tt.ptr> loc(#loc25) + %b = tt.splat %b_ptr : !tt.ptr -> tensor<16x16x!tt.ptr> loc(#loc26) + %b_8 = tt.addptr %b, %idx_5 : tensor<16x16x!tt.ptr>, tensor<16x16xi32> loc(#loc26) + %b_9 = tt.load %b_8 : tensor<16x16x!tt.ptr> loc(#loc27) + %c = tt.dot %a_7, %b_9, %cst : tensor<16x16xf32> * tensor<16x16xf32> -> tensor<16x16xf32> loc(#loc28) + %0 = tt.splat %c_ptr : !tt.ptr -> tensor<16x16x!tt.ptr> loc(#loc12) + %1 = tt.addptr %0, %idx_5 : tensor<16x16x!tt.ptr>, tensor<16x16xi32> loc(#loc12) + %d = tt.dot %a_7, %b_9, %c, inputPrecision = tf32x3 : tensor<16x16xf32> * tensor<16x16xf32> -> tensor<16x16xf32> loc(#loc29) + tt.store %1, %d : tensor<16x16x!tt.ptr> loc(#loc14) + tt.return loc(#loc15) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":32:26) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":31:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":32:15) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":32:39) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":32:34) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":33:24) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":33:16) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":34:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":34:16) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":35:18) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":37:21) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":36:18) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":37:26) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":37:4) +#loc19 = loc("idx"(#loc2)) +#loc20 = loc("offs"(#loc3)) +#loc21 = loc("idx"(#loc4)) +#loc22 = loc("idx"(#loc5)) +#loc23 = loc("idx"(#loc6)) +#loc24 = loc("a"(#loc7)) +#loc25 = loc("a"(#loc8)) +#loc26 = loc("b"(#loc9)) +#loc27 = loc("b"(#loc10)) +#loc28 = loc("c"(#loc11)) +#loc29 = loc("d"(#loc13)) diff --git a/tests/golden/ir/ttir/kernel_dot_scaled.ttir b/tests/golden/ir/ttir/kernel_dot_scaled.ttir new file mode 100644 index 000000000..d86d79dbd --- /dev/null +++ b/tests/golden/ir/ttir/kernel_dot_scaled.ttir @@ -0,0 +1,112 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":70:0) +#loc31 = loc("a_ptr"(#loc)) +#loc32 = loc("as_ptr"(#loc)) +#loc33 = loc("b_ptr"(#loc)) +#loc34 = loc("bs_ptr"(#loc)) +#loc35 = loc("c_ptr"(#loc)) +module { + tt.func public @dot_scaled_k(%a_ptr: !tt.ptr loc("a_ptr"(#loc)), %as_ptr: !tt.ptr loc("as_ptr"(#loc)), %b_ptr: !tt.ptr loc("b_ptr"(#loc)), %bs_ptr: !tt.ptr loc("bs_ptr"(#loc)), %c_ptr: !tt.ptr loc("c_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<128> : tensor<128x1xi32> loc(#loc1) + %c = arith.constant dense<0.000000e+00> : tensor<128x128xf32> loc(#loc36) + %cst_0 = arith.constant dense<2> : tensor<128x1xi32> loc(#loc3) + %b = arith.constant dense<128> : tensor<64x1xi32> loc(#loc37) + %a = arith.constant dense<64> : tensor<128x1xi32> loc(#loc38) + %rm = tt.make_range {end = 128 : i32, start = 0 : i32} : tensor<128xi32> loc(#loc39) + %rk = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc40) + %rs = tt.make_range {end = 2 : i32, start = 0 : i32} : tensor<2xi32> loc(#loc41) + %a_1 = tt.expand_dims %rm {axis = 1 : i32} : tensor<128xi32> -> tensor<128x1xi32> loc(#loc42) + %a_2 = arith.muli %a_1, %a : tensor<128x1xi32> loc(#loc38) + %a_3 = tt.splat %a_ptr : !tt.ptr -> tensor<128x1x!tt.ptr> loc(#loc43) + %a_4 = tt.addptr %a_3, %a_2 : tensor<128x1x!tt.ptr>, tensor<128x1xi32> loc(#loc43) + %a_5 = tt.expand_dims %rk {axis = 0 : i32} : tensor<64xi32> -> tensor<1x64xi32> loc(#loc44) + %a_6 = tt.broadcast %a_4 : tensor<128x1x!tt.ptr> -> tensor<128x64x!tt.ptr> loc(#loc45) + %a_7 = tt.broadcast %a_5 : tensor<1x64xi32> -> tensor<128x64xi32> loc(#loc45) + %a_8 = tt.addptr %a_6, %a_7 : tensor<128x64x!tt.ptr>, tensor<128x64xi32> loc(#loc45) + %a_9 = tt.load %a_8 : tensor<128x64x!tt.ptr> loc(#loc46) + %b_10 = tt.expand_dims %rk {axis = 1 : i32} : tensor<64xi32> -> tensor<64x1xi32> loc(#loc47) + %b_11 = arith.muli %b_10, %b : tensor<64x1xi32> loc(#loc37) + %b_12 = tt.splat %b_ptr : !tt.ptr -> tensor<64x1x!tt.ptr> loc(#loc48) + %b_13 = tt.addptr %b_12, %b_11 : tensor<64x1x!tt.ptr>, tensor<64x1xi32> loc(#loc48) + %b_14 = tt.expand_dims %rm {axis = 0 : i32} : tensor<128xi32> -> tensor<1x128xi32> loc(#loc49) + %b_15 = tt.broadcast %b_13 : tensor<64x1x!tt.ptr> -> tensor<64x128x!tt.ptr> loc(#loc50) + %b_16 = tt.broadcast %b_14 : tensor<1x128xi32> -> tensor<64x128xi32> loc(#loc50) + %b_17 = tt.addptr %b_15, %b_16 : tensor<64x128x!tt.ptr>, tensor<64x128xi32> loc(#loc50) + %b_18 = tt.load %b_17 : tensor<64x128x!tt.ptr> loc(#loc51) + %a_scale = arith.muli %a_1, %cst_0 : tensor<128x1xi32> loc(#loc52) + %a_scale_19 = tt.splat %as_ptr : !tt.ptr -> tensor<128x1x!tt.ptr> loc(#loc53) + %a_scale_20 = tt.addptr %a_scale_19, %a_scale : tensor<128x1x!tt.ptr>, tensor<128x1xi32> loc(#loc53) + %a_scale_21 = tt.expand_dims %rs {axis = 0 : i32} : tensor<2xi32> -> tensor<1x2xi32> loc(#loc54) + %a_scale_22 = tt.broadcast %a_scale_20 : tensor<128x1x!tt.ptr> -> tensor<128x2x!tt.ptr> loc(#loc55) + %a_scale_23 = tt.broadcast %a_scale_21 : tensor<1x2xi32> -> tensor<128x2xi32> loc(#loc55) + %a_scale_24 = tt.addptr %a_scale_22, %a_scale_23 : tensor<128x2x!tt.ptr>, tensor<128x2xi32> loc(#loc55) + %a_scale_25 = tt.load %a_scale_24 : tensor<128x2x!tt.ptr> loc(#loc56) + %b_scale = tt.splat %bs_ptr : !tt.ptr -> tensor<128x1x!tt.ptr> loc(#loc57) + %b_scale_26 = tt.addptr %b_scale, %a_scale : tensor<128x1x!tt.ptr>, tensor<128x1xi32> loc(#loc57) + %b_scale_27 = tt.broadcast %b_scale_26 : tensor<128x1x!tt.ptr> -> tensor<128x2x!tt.ptr> loc(#loc58) + %b_scale_28 = tt.addptr %b_scale_27, %a_scale_23 : tensor<128x2x!tt.ptr>, tensor<128x2xi32> loc(#loc58) + %b_scale_29 = tt.load %b_scale_28 : tensor<128x2x!tt.ptr> loc(#loc59) + %c_30 = tt.dot_scaled %a_9 scale %a_scale_25, %b_18 scale %b_scale_29, %c lhs = e4m3 rhs = e4m3 {fastMath = false} : tensor<128x64xf8E4M3FN>, tensor<128x2xi8> * tensor<64x128xf8E4M3FN>, tensor<128x2xi8> -> tensor<128x128xf32> loc(#loc36) + %0 = arith.muli %a_1, %cst : tensor<128x1xi32> loc(#loc1) + %1 = tt.splat %c_ptr : !tt.ptr -> tensor<128x1x!tt.ptr> loc(#loc27) + %2 = tt.addptr %1, %0 : tensor<128x1x!tt.ptr>, tensor<128x1xi32> loc(#loc27) + %3 = tt.broadcast %2 : tensor<128x1x!tt.ptr> -> tensor<128x128x!tt.ptr> loc(#loc28) + %4 = tt.broadcast %b_14 : tensor<1x128xi32> -> tensor<128x128xi32> loc(#loc28) + %5 = tt.addptr %3, %4 : tensor<128x128x!tt.ptr>, tensor<128x128xi32> loc(#loc28) + tt.store %5, %c_30 : tensor<128x128x!tt.ptr> loc(#loc29) + tt.return loc(#loc30) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":81:35) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":80:54) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":77:38) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":76:38) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":72:22) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":74:22) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":75:22) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":76:27) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":76:24) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":76:45) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":76:42) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":76:16) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":77:27) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":77:24) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":77:45) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":77:42) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":77:16) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":78:46) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":78:31) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":78:60) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":78:57) +#loc23 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":78:22) +#loc24 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":79:31) +#loc25 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":79:57) +#loc26 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":79:22) +#loc27 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":81:21) +#loc28 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":81:39) +#loc29 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":81:52) +#loc30 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":81:4) +#loc36 = loc("c"(#loc2)) +#loc37 = loc("b"(#loc4)) +#loc38 = loc("a"(#loc5)) +#loc39 = loc("rm"(#loc6)) +#loc40 = loc("rk"(#loc7)) +#loc41 = loc("rs"(#loc8)) +#loc42 = loc("a"(#loc9)) +#loc43 = loc("a"(#loc10)) +#loc44 = loc("a"(#loc11)) +#loc45 = loc("a"(#loc12)) +#loc46 = loc("a"(#loc13)) +#loc47 = loc("b"(#loc14)) +#loc48 = loc("b"(#loc15)) +#loc49 = loc("b"(#loc16)) +#loc50 = loc("b"(#loc17)) +#loc51 = loc("b"(#loc18)) +#loc52 = loc("a_scale"(#loc19)) +#loc53 = loc("a_scale"(#loc20)) +#loc54 = loc("a_scale"(#loc21)) +#loc55 = loc("a_scale"(#loc22)) +#loc56 = loc("a_scale"(#loc23)) +#loc57 = loc("b_scale"(#loc24)) +#loc58 = loc("b_scale"(#loc25)) +#loc59 = loc("b_scale"(#loc26)) diff --git a/tests/golden/ir/ttir/kernel_eps_consts.ttir b/tests/golden/ir/ttir/kernel_eps_consts.ttir new file mode 100644 index 000000000..d4df5ad58 --- /dev/null +++ b/tests/golden/ir/ttir/kernel_eps_consts.ttir @@ -0,0 +1,39 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":41:0) +#loc12 = loc("x_ptr"(#loc)) +#loc13 = loc("s_ptr"(#loc)) +#loc14 = loc("out_ptr"(#loc)) +module { + tt.func public @eps_consts(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %s_ptr: !tt.ptr loc("s_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<9.99999996E-13> : tensor<64xf32> loc(#loc1) + %cst_0 = arith.constant 9.99999997E-7 : f32 loc(#loc2) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc15) + %x = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc16) + %x_1 = tt.addptr %x, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc16) + %x_2 = tt.load %x_1 : tensor<64x!tt.ptr> loc(#loc17) + %s = tt.load %s_ptr : !tt.ptr loc(#loc18) + %s_3 = arith.addf %s, %cst_0 : f32 loc(#loc19) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc8) + %1 = tt.addptr %0, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc8) + %2 = tt.splat %s_3 : f32 -> tensor<64xf32> loc(#loc9) + %3 = arith.mulf %x_2, %2 : tensor<64xf32> loc(#loc9) + %4 = arith.addf %3, %cst : tensor<64xf32> loc(#loc1) + tt.store %1, %4 : tensor<64x!tt.ptr> loc(#loc10) + tt.return loc(#loc11) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":46:37) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":43:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":44:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":44:16) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":45:16) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":45:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":46:23) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":46:33) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":46:29) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":46:4) +#loc15 = loc("offs"(#loc3)) +#loc16 = loc("x"(#loc4)) +#loc17 = loc("x"(#loc5)) +#loc18 = loc("s"(#loc6)) +#loc19 = loc("s"(#loc7)) diff --git a/tests/golden/ir/ttir/kernel_unicode_msgs.ttir b/tests/golden/ir/ttir/kernel_unicode_msgs.ttir new file mode 100644 index 000000000..6895e2e36 --- /dev/null +++ b/tests/golden/ir/ttir/kernel_unicode_msgs.ttir @@ -0,0 +1,30 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":50:0) +#loc10 = loc("x_ptr"(#loc)) +module { + tt.func public @unicode_msgs(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<0.000000e+00> : tensor<64xf32> loc(#loc1) + %cst_0 = arith.constant dense<1.000000e+00> : tensor<64xf32> loc(#loc2) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc11) + %x = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc12) + %x_1 = tt.addptr %x, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc12) + %x_2 = tt.load %x_1 : tensor<64x!tt.ptr> loc(#loc13) + %0 = arith.cmpf ogt, %x_2, %cst : tensor<64xf32> loc(#loc1) + tt.assert %0, "\E9\94\99\E8\AF\AF: \CF\80 must be > 0" : tensor<64xi1> loc(#loc6) + tt.print " x=: " {hex = false, isSigned = array} : %x_2 : tensor<64xf32> loc(#loc7) + %1 = arith.addf %x_2, %cst_0 : tensor<64xf32> loc(#loc2) + tt.store %x_1, %1 : tensor<64x!tt.ptr> loc(#loc8) + tt.return loc(#loc9) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":54:25) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":56:31) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":52:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":53:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":53:16) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":54:28) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":55:26) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":56:27) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":56:4) +#loc11 = loc("offs"(#loc3)) +#loc12 = loc("x"(#loc4)) +#loc13 = loc("x"(#loc5)) diff --git a/tests/golden/ir/ttir/nat_dead_if.ttir b/tests/golden/ir/ttir/nat_dead_if.ttir new file mode 100644 index 000000000..469929ed7 --- /dev/null +++ b/tests/golden/ir/ttir/nat_dead_if.ttir @@ -0,0 +1,36 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":47:0) +#loc12 = loc("x_ptr"(#loc)) +#loc13 = loc("out_ptr"(#loc)) +#loc14 = loc("n"(#loc)) +module { + tt.func public @dead_if(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc15) + %v = tt.addptr %x_ptr, %pid : !tt.ptr, i32 loc(#loc16) + %v_0 = tt.load %v : !tt.ptr loc(#loc17) + %0 = arith.cmpi slt, %pid, %n : i32 loc(#loc5) + scf.if %0 { + } else { + %3 = tt.addptr %out_ptr, %pid : !tt.ptr, i32 loc(#loc7) + tt.store %3, %v_0 : !tt.ptr loc(#loc8) + } loc(#loc6) + %1 = tt.addptr %out_ptr, %pid : !tt.ptr, i32 loc(#loc9) + %2 = tt.addptr %1, %c1_i32 : !tt.ptr, i32 loc(#loc1) + tt.store %2, %v_0 : !tt.ptr loc(#loc10) + tt.return loc(#loc11) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":54:29) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":48:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":49:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":49:16) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":50:13) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":50:7) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":53:27) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":53:32) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":54:23) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":54:32) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":54:4) +#loc15 = loc("pid"(#loc2)) +#loc16 = loc("v"(#loc3)) +#loc17 = loc("v"(#loc4)) diff --git a/tests/golden/ir/ttir/nat_empty_loop.ttir b/tests/golden/ir/ttir/nat_empty_loop.ttir new file mode 100644 index 000000000..d76c033a0 --- /dev/null +++ b/tests/golden/ir/ttir/nat_empty_loop.ttir @@ -0,0 +1,21 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":37:0) +#loc7 = loc("x_ptr"(#loc)) +#loc8 = loc("out_ptr"(#loc)) +#loc9 = loc("n"(#loc)) +module { + tt.func public @empty_loop(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %pid = tt.get_program_id x : i32 loc(#loc10) + %0 = tt.addptr %out_ptr, %pid : !tt.ptr, i32 loc(#loc2) + %1 = tt.addptr %x_ptr, %pid : !tt.ptr, i32 loc(#loc3) + %2 = tt.load %1 : !tt.ptr loc(#loc4) + tt.store %0, %2 : !tt.ptr loc(#loc5) + tt.return loc(#loc6) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":38:24) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":43:23) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":43:44) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":43:36) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":43:28) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":43:4) +#loc10 = loc("pid"(#loc1)) diff --git a/tests/golden/ir/ttir/nat_empty_then.ttir b/tests/golden/ir/ttir/nat_empty_then.ttir new file mode 100644 index 000000000..891ddc66b --- /dev/null +++ b/tests/golden/ir/ttir/nat_empty_then.ttir @@ -0,0 +1,27 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":28:0) +#loc9 = loc("x_ptr"(#loc)) +#loc10 = loc("out_ptr"(#loc)) +#loc11 = loc("n"(#loc)) +module { + tt.func public @empty_then(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %pid = tt.get_program_id x : i32 loc(#loc12) + %0 = arith.cmpi slt, %pid, %n : i32 loc(#loc2) + scf.if %0 { + } else { + %1 = tt.addptr %out_ptr, %pid : !tt.ptr, i32 loc(#loc4) + %2 = tt.addptr %x_ptr, %pid : !tt.ptr, i32 loc(#loc5) + %3 = tt.load %2 : !tt.ptr loc(#loc6) + tt.store %1, %3 : !tt.ptr loc(#loc7) + } loc(#loc3) + tt.return loc(#loc8) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":29:24) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":30:13) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":30:7) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":33:27) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":33:48) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":33:40) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":33:32) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":30:4) +#loc12 = loc("pid"(#loc1)) diff --git a/tests/golden/ir/ttir/nat_hint_arange_const.ttir b/tests/golden/ir/ttir/nat_hint_arange_const.ttir new file mode 100644 index 000000000..ee7803f99 --- /dev/null +++ b/tests/golden/ir/ttir/nat_hint_arange_const.ttir @@ -0,0 +1,21 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":15:0) +#loc7 = loc("x_ptr"(#loc)) +#loc8 = loc("out_ptr"(#loc)) +module { + tt.func public @hint_arange_const(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc))) attributes {noinline = false} { + %c16_i32 = arith.constant 16 : i32 loc(#loc1) + %0 = tt.addptr %out_ptr, %c16_i32 : !tt.ptr, i32 loc(#loc2) + %1 = tt.splat %0 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc2) + %2 = tt.addptr %x_ptr, %c16_i32 : !tt.ptr, i32 loc(#loc3) + %3 = tt.splat %2 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc3) + %4 = tt.load %3 : tensor<64x!tt.ptr> loc(#loc4) + tt.store %1, %4 : tensor<64x!tt.ptr> loc(#loc5) + tt.return loc(#loc6) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":17:23) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":17:45) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":17:37) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":17:29) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":17:4) diff --git a/tests/golden/ir/ttir/nat_hint_scalar_const.ttir b/tests/golden/ir/ttir/nat_hint_scalar_const.ttir new file mode 100644 index 000000000..17e26216c --- /dev/null +++ b/tests/golden/ir/ttir/nat_hint_scalar_const.ttir @@ -0,0 +1,27 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":8:0) +#loc9 = loc("x_ptr"(#loc)) +#loc10 = loc("out_ptr"(#loc)) +module { + tt.func public @hint_scalar_const(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc))) attributes {noinline = false} { + %c = arith.constant {tt.divisibility = dense<64> : tensor<1xi32>} 64 : i32 loc(#loc11) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc12) + %0 = tt.addptr %out_ptr, %c : !tt.ptr, i32 loc(#loc3) + %1 = tt.splat %0 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc4) + %2 = tt.addptr %1, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc4) + %3 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc5) + %4 = tt.addptr %3, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc5) + %5 = tt.load %4 : tensor<64x!tt.ptr> loc(#loc6) + tt.store %2, %5 : tensor<64x!tt.ptr> loc(#loc7) + tt.return loc(#loc8) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":9:39) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":10:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":11:23) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":11:27) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":11:49) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":11:41) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":11:33) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels.py":11:4) +#loc11 = loc("c"(#loc1)) +#loc12 = loc("offs"(#loc2)) diff --git a/tests/golden/ir/ttir/nat_k_uni.ttir b/tests/golden/ir/ttir/nat_k_uni.ttir new file mode 100644 index 000000000..596a4c9e0 --- /dev/null +++ b/tests/golden/ir/ttir/nat_k_uni.ttir @@ -0,0 +1,21 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/\E5\86\85\E6\A0\B8/k_uni.py":7:0) +#loc7 = loc("x_ptr"(#loc)) +module { + tt.func public @k_uni(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<1.000000e+00> : tensor<64xf32> loc(#loc1) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc8) + %0 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc3) + %1 = tt.addptr %0, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc3) + %2 = tt.load %1 : tensor<64x!tt.ptr> loc(#loc4) + %3 = arith.addf %2, %cst : tensor<64xf32> loc(#loc1) + tt.store %1, %3 : tensor<64x!tt.ptr> loc(#loc5) + tt.return loc(#loc6) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/\E5\86\85\E6\A0\B8/k_uni.py":9:51) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/\E5\86\85\E6\A0\B8/k_uni.py":8:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/\E5\86\85\E6\A0\B8/k_uni.py":9:21) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/\E5\86\85\E6\A0\B8/k_uni.py":9:35) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/\E5\86\85\E6\A0\B8/k_uni.py":9:27) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/\E5\86\85\E6\A0\B8/k_uni.py":9:4) +#loc8 = loc("offs"(#loc2)) diff --git a/tests/golden/ir/ttir/nat_uni_params.ttir b/tests/golden/ir/ttir/nat_uni_params.ttir new file mode 100644 index 000000000..fda4160d4 --- /dev/null +++ b/tests/golden/ir/ttir/nat_uni_params.ttir @@ -0,0 +1,22 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels2.py":7:0) +#loc7 = loc("\CF\80_ptr"(#loc)) +#loc8 = loc("\E6\95\B0_n"(#loc)) +module { + tt.func public @uni_params(%_CF80_ptr: !tt.ptr loc("\CF\80_ptr"(#loc)), %_E695B0_n: i32 loc("\E6\95\B0_n"(#loc))) attributes {noinline = false} { + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc9) + %0 = tt.splat %_E695B0_n : i32 -> tensor<64xi32> loc(#loc2) + %1 = arith.cmpi slt, %offs, %0 : tensor<64xi32> loc(#loc2) + %2 = tt.splat %_CF80_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc3) + %3 = tt.addptr %2, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc3) + %4 = tt.load %3 : tensor<64x!tt.ptr> loc(#loc4) + tt.store %3, %4, %1 : tensor<64x!tt.ptr> loc(#loc5) + tt.return loc(#loc6) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels2.py":8:24) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels2.py":9:64) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels2.py":9:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels2.py":9:36) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels2.py":9:28) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/nat_kernels2.py":9:4) +#loc9 = loc("offs"(#loc1)) diff --git a/tests/golden/ir/ttir/spike_atomics.ttir b/tests/golden/ir/ttir/spike_atomics.ttir new file mode 100644 index 000000000..5af5b97f2 --- /dev/null +++ b/tests/golden/ir/ttir/spike_atomics.ttir @@ -0,0 +1,59 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":84:0) +#loc18 = loc("p_ptr"(#loc)) +#loc19 = loc("q_ptr"(#loc)) +#loc20 = loc("n"(#loc)) +module { + tt.func public @atomics(%p_ptr: !tt.ptr loc("p_ptr"(#loc)), %q_ptr: !tt.ptr loc("q_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c3_i32 = arith.constant 3 : i32 loc(#loc1) + %c9_i32 = arith.constant 9 : i32 loc(#loc2) + %true = arith.constant true loc(#loc3) + %old = arith.constant 5 : i32 loc(#loc21) + %cst = arith.constant dense<2> : tensor<64xi32> loc(#loc5) + %c2_i32 = arith.constant 2 : i32 loc(#loc3) + %cst_0 = arith.constant dense<1> : tensor<64xi32> loc(#loc6) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %cst_1 = arith.constant dense<240> : tensor<64xi32> loc(#loc7) + %cst_2 = arith.constant dense<-3> : tensor<64xi32> loc(#loc8) + %cst_3 = arith.constant dense<7> : tensor<64xi32> loc(#loc9) + %cst_4 = arith.constant dense<1.000000e+00> : tensor<64xf32> loc(#loc10) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc22) + %m = tt.splat %n : i32 -> tensor<64xi32> loc(#loc23) + %m_5 = arith.cmpi slt, %offs, %m : tensor<64xi32> loc(#loc23) + %0 = tt.splat %p_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc13) + %1 = tt.addptr %0, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc13) + %2 = tt.atomic_rmw fadd, relaxed, cta, %1, %cst_4, %m_5 : (tensor<64x!tt.ptr>, tensor<64xf32>, tensor<64xi1>) -> tensor<64xf32> loc(#loc10) + %3 = tt.splat %q_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc14) + %4 = tt.addptr %3, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc14) + %5 = tt.atomic_rmw max, release, sys, %4, %cst_3, %m_5 : (tensor<64x!tt.ptr>, tensor<64xi32>, tensor<64xi1>) -> tensor<64xi32> loc(#loc9) + %6 = tt.atomic_rmw min, acquire, gpu, %4, %cst_2, %m_5 : (tensor<64x!tt.ptr>, tensor<64xi32>, tensor<64xi1>) -> tensor<64xi32> loc(#loc8) + %7 = tt.atomic_rmw and, acq_rel, gpu, %4, %cst_1, %m_5 : (tensor<64x!tt.ptr>, tensor<64xi32>, tensor<64xi1>) -> tensor<64xi32> loc(#loc7) + %8 = tt.atomic_rmw or, acq_rel, gpu, %4, %cst_0, %m_5 : (tensor<64x!tt.ptr>, tensor<64xi32>, tensor<64xi1>) -> tensor<64xi32> loc(#loc6) + %9 = tt.atomic_rmw xor, acq_rel, gpu, %4, %cst, %m_5 : (tensor<64x!tt.ptr>, tensor<64xi32>, tensor<64xi1>) -> tensor<64xi32> loc(#loc5) + %old_6 = tt.atomic_rmw exch, relaxed, sys, %q_ptr, %old, %true : (!tt.ptr, i32, i1) -> i32 loc(#loc21) + %10 = tt.addptr %q_ptr, %c1_i32 : !tt.ptr, i32 loc(#loc15) + %11 = tt.atomic_cas acq_rel, cta, %10, %old_6, %c9_i32 : (!tt.ptr, i32, i32) -> i32 loc(#loc2) + %12 = tt.addptr %q_ptr, %c2_i32 : !tt.ptr, i32 loc(#loc16) + %13 = tt.atomic_rmw add, acq_rel, gpu, %12, %c3_i32, %true : (!tt.ptr, i32, i1) -> i32 loc(#loc1) + tt.return loc(#loc17) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":95:29) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":94:34) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":93:32) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":92:32) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":91:31) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":90:32) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":89:32) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":88:32) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":87:32) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":85:24) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":86:15) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":87:26) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":88:26) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":94:26) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":95:26) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":95:4) +#loc21 = loc("old"(#loc4)) +#loc22 = loc("offs"(#loc11)) +#loc23 = loc("m"(#loc12)) diff --git a/tests/golden/ir/ttir/spike_casts.ttir b/tests/golden/ir/ttir/spike_casts.ttir new file mode 100644 index 000000000..91d086d2d --- /dev/null +++ b/tests/golden/ir/ttir/spike_casts.ttir @@ -0,0 +1,63 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":182:0) +#loc19 = loc("x_ptr"(#loc)) +#loc20 = loc("out_ptr"(#loc)) +#loc21 = loc("n"(#loc)) +module { + tt.func public @casts(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %offs = tt.get_program_id x : i32 loc(#loc22) + %offs_0 = arith.muli %offs, %c64_i32 : i32 loc(#loc23) + %offs_1 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc24) + %offs_2 = tt.splat %offs_0 : i32 -> tensor<64xi32> loc(#loc25) + %offs_3 = arith.addi %offs_2, %offs_1 : tensor<64xi32> loc(#loc25) + %o16 = arith.trunci %offs_3 : tensor<64xi32> to tensor<64xi16> loc(#loc26) + %o64 = arith.extsi %o16 : tensor<64xi16> to tensor<64xi64> loc(#loc27) + %u8 = arith.trunci %offs_3 : tensor<64xi32> to tensor<64xi8> loc(#loc28) + %idx = arith.extui %u8 : tensor<64xi8> to tensor<64xi64> loc(#loc34) + %idx_4 = arith.addi %o64, %idx : tensor<64xi64> loc(#loc29) + %v = arith.extsi %n : i32 to i64 loc(#loc31) + %v_5 = tt.splat %v : i64 -> tensor<64xi64> loc(#loc31) + %v_6 = arith.cmpi slt, %idx_4, %v_5 : tensor<64xi64> loc(#loc31) + %v_7 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc32) + %v_8 = tt.addptr %v_7, %idx_4 : tensor<64x!tt.ptr>, tensor<64xi64> loc(#loc32) + %v_9 = tt.load %v_8, %v_6 : tensor<64x!tt.ptr> loc(#loc33) + %0 = tt.splat %n : i32 -> tensor<64xi32> loc(#loc14) + %1 = arith.cmpi slt, %offs_3, %0 : tensor<64xi32> loc(#loc14) + %2 = arith.extsi %offs_3 : tensor<64xi32> to tensor<64xi64> loc(#loc15) + %3 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc16) + %4 = tt.addptr %3, %2 : tensor<64x!tt.ptr>, tensor<64xi64> loc(#loc16) + tt.store %4, %v_9, %1 : tensor<64x!tt.ptr> loc(#loc17) + tt.return loc(#loc18) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":183:25) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":183:30) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":183:51) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":183:38) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":184:18) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":185:17) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":186:17) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":187:16) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":187:22) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":188:40) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":188:24) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":188:16) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":189:57) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":189:31) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":189:23) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":189:42) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":189:4) +#loc22 = loc("offs"(#loc2)) +#loc23 = loc("offs"(#loc3)) +#loc24 = loc("offs"(#loc4)) +#loc25 = loc("offs"(#loc5)) +#loc26 = loc("o16"(#loc6)) +#loc27 = loc("o64"(#loc7)) +#loc28 = loc("u8"(#loc8)) +#loc29 = loc("idx"(#loc9)) +#loc30 = loc("idx"(#loc10)) +#loc31 = loc("v"(#loc11)) +#loc32 = loc("v"(#loc12)) +#loc33 = loc("v"(#loc13)) +#loc34 = loc(fused[#loc29, #loc30]) diff --git a/tests/golden/ir/ttir/spike_dot.ttir b/tests/golden/ir/ttir/spike_dot.ttir new file mode 100644 index 000000000..ff30398d9 --- /dev/null +++ b/tests/golden/ir/ttir/spike_dot.ttir @@ -0,0 +1,59 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":130:0) +#loc17 = loc("a_ptr"(#loc)) +#loc18 = loc("b_ptr"(#loc)) +#loc19 = loc("c_ptr"(#loc)) +module { + tt.func public @dot(%a_ptr: !tt.ptr loc("a_ptr"(#loc)), %b_ptr: !tt.ptr loc("b_ptr"(#loc)), %c_ptr: !tt.ptr loc("c_ptr"(#loc))) attributes {noinline = false} { + %c = arith.constant dense<0.000000e+00> : tensor<32x32xf32> loc(#loc20) + %cst = arith.constant dense<32> : tensor<32x1xi32> loc(#loc2) + %rm = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc21) + %a = tt.expand_dims %rm {axis = 1 : i32} : tensor<32xi32> -> tensor<32x1xi32> loc(#loc22) + %a_0 = arith.muli %a, %cst : tensor<32x1xi32> loc(#loc23) + %a_1 = tt.splat %a_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc24) + %a_2 = tt.addptr %a_1, %a_0 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc24) + %a_3 = tt.expand_dims %rm {axis = 0 : i32} : tensor<32xi32> -> tensor<1x32xi32> loc(#loc25) + %a_4 = tt.broadcast %a_2 : tensor<32x1x!tt.ptr> -> tensor<32x32x!tt.ptr> loc(#loc26) + %a_5 = tt.broadcast %a_3 : tensor<1x32xi32> -> tensor<32x32xi32> loc(#loc26) + %a_6 = tt.addptr %a_4, %a_5 : tensor<32x32x!tt.ptr>, tensor<32x32xi32> loc(#loc26) + %a_7 = tt.load %a_6 : tensor<32x32x!tt.ptr> loc(#loc27) + %b = tt.splat %b_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc28) + %b_8 = tt.addptr %b, %a_0 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc28) + %b_9 = tt.broadcast %b_8 : tensor<32x1x!tt.ptr> -> tensor<32x32x!tt.ptr> loc(#loc29) + %b_10 = tt.addptr %b_9, %a_5 : tensor<32x32x!tt.ptr>, tensor<32x32xi32> loc(#loc29) + %b_11 = tt.load %b_10 : tensor<32x32x!tt.ptr> loc(#loc30) + %c_12 = tt.dot %a_7, %b_11, %c, inputPrecision = tf32 : tensor<32x32xf16> * tensor<32x32xf16> -> tensor<32x32xf32> loc(#loc20) + %0 = tt.splat %c_ptr : !tt.ptr -> tensor<32x1x!tt.ptr> loc(#loc13) + %1 = tt.addptr %0, %a_0 : tensor<32x1x!tt.ptr>, tensor<32x1xi32> loc(#loc13) + %2 = tt.broadcast %1 : tensor<32x1x!tt.ptr> -> tensor<32x32x!tt.ptr> loc(#loc14) + %3 = tt.addptr %2, %a_5 : tensor<32x32x!tt.ptr>, tensor<32x32xi32> loc(#loc14) + tt.store %3, %c_12 : tensor<32x32x!tt.ptr> loc(#loc15) + tt.return loc(#loc16) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":136:18) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":131:22) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":134:27) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":134:38) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":134:24) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":134:46) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":134:43) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":134:16) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":135:24) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":135:43) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":135:16) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":137:21) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":137:40) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":137:53) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":137:4) +#loc20 = loc("c"(#loc1)) +#loc21 = loc("rm"(#loc3)) +#loc22 = loc("a"(#loc4)) +#loc23 = loc("a"(#loc5)) +#loc24 = loc("a"(#loc6)) +#loc25 = loc("a"(#loc7)) +#loc26 = loc("a"(#loc8)) +#loc27 = loc("a"(#loc9)) +#loc28 = loc("b"(#loc10)) +#loc29 = loc("b"(#loc11)) +#loc30 = loc("b"(#loc12)) diff --git a/tests/golden/ir/ttir/spike_early_return.ttir b/tests/golden/ir/ttir/spike_early_return.ttir new file mode 100644 index 000000000..2e793ed99 --- /dev/null +++ b/tests/golden/ir/ttir/spike_early_return.ttir @@ -0,0 +1,46 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":47:0) +#loc14 = loc("x_ptr"(#loc)) +#loc15 = loc("out_ptr"(#loc)) +#loc16 = loc("n"(#loc)) +#loc17 = loc("T"(#loc)) +module { + tt.func public @early_return(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc)), %T: i32 loc("T"(#loc))) attributes {noinline = false} { + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc18) + %0 = arith.muli %pid, %c64_i32 : i32 loc(#loc3) + %1 = arith.cmpi sge, %0, %T : i32 loc(#loc4) + cf.cond_br %1, ^bb1, ^bb2 loc(#loc4) + ^bb1: // pred: ^bb0 + tt.return loc(#loc5) + ^bb2: // pred: ^bb0 + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc19) + %offs_0 = tt.splat %0 : i32 -> tensor<64xi32> loc(#loc20) + %offs_1 = arith.addi %offs_0, %offs : tensor<64xi32> loc(#loc20) + %m = tt.splat %n : i32 -> tensor<64xi32> loc(#loc21) + %m_2 = arith.cmpi slt, %offs_1, %m : tensor<64xi32> loc(#loc21) + %2 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc9) + %3 = tt.addptr %2, %offs_1 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc9) + %4 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc10) + %5 = tt.addptr %4, %offs_1 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc10) + %6 = tt.load %5, %m_2 : tensor<64x!tt.ptr> loc(#loc11) + tt.store %3, %6, %m_2 : tensor<64x!tt.ptr> loc(#loc12) + tt.return loc(#loc13) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":48:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":49:13) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":49:22) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":50:8) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":51:38) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":51:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":52:15) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":53:23) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":53:45) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":53:37) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":53:29) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":53:4) +#loc18 = loc("pid"(#loc2)) +#loc19 = loc("offs"(#loc6)) +#loc20 = loc("offs"(#loc7)) +#loc21 = loc("m"(#loc8)) diff --git a/tests/golden/ir/ttir/spike_early_return_loop.ttir b/tests/golden/ir/ttir/spike_early_return_loop.ttir new file mode 100644 index 000000000..cf2a01870 --- /dev/null +++ b/tests/golden/ir/ttir/spike_early_return_loop.ttir @@ -0,0 +1,58 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":59:0) +#loc17 = loc("x_ptr"(#loc)) +#loc18 = loc("out_ptr"(#loc)) +#loc19 = loc("n"(#loc)) +module { + tt.func public @early_return_loop(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c3_i32 = arith.constant 3 : i32 loc(#loc2) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc20) + %offs = arith.muli %pid, %c64_i32 : i32 loc(#loc21) + %offs_0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc22) + %offs_1 = tt.splat %offs : i32 -> tensor<64xi32> loc(#loc23) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<64xi32> loc(#loc23) + %0 = arith.cmpi eq, %pid, %c3_i32 : i32 loc(#loc2) + cf.cond_br %0, ^bb1, ^bb2 loc(#loc2) + ^bb1: // pred: ^bb0 + tt.return loc(#loc7) + ^bb2: // pred: ^bb0 + %s = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%s_3 = %c0_i32) -> (i32) : i32 { + %s_4 = arith.addi %s_3, %i : i32 loc(#loc25) + scf.yield %s_4 : i32 loc(#loc10) + } loc(#loc24) + %1 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc11) + %2 = tt.addptr %1, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc11) + %3 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc12) + %4 = tt.addptr %3, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc12) + %5 = tt.load %4 : tensor<64x!tt.ptr> loc(#loc13) + %6 = arith.sitofp %s : i32 to f32 loc(#loc14) + %7 = tt.splat %6 : f32 -> tensor<64xf32> loc(#loc14) + %8 = arith.addf %5, %7 : tensor<64xf32> loc(#loc14) + tt.store %2, %8 : tensor<64x!tt.ptr> loc(#loc15) + tt.return loc(#loc16) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":62:14) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":60:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":61:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":61:38) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":61:25) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":63:8) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":65:22) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":66:13) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":66:8) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":67:23) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":67:45) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":67:37) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":67:53) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":67:29) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":67:4) +#loc20 = loc("pid"(#loc3)) +#loc21 = loc("offs"(#loc4)) +#loc22 = loc("offs"(#loc5)) +#loc23 = loc("offs"(#loc6)) +#loc24 = loc("s"(#loc8)) +#loc25 = loc("s"(#loc9)) diff --git a/tests/golden/ir/ttir/spike_for_ptr_iterargs.ttir b/tests/golden/ir/ttir/spike_for_ptr_iterargs.ttir new file mode 100644 index 000000000..8b256cfd6 --- /dev/null +++ b/tests/golden/ir/ttir/spike_for_ptr_iterargs.ttir @@ -0,0 +1,71 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":17:0) +#loc19 = loc("a_ptr"(#loc)) +#loc20 = loc("b_ptr"(#loc)) +#loc21 = loc("out_ptr"(#loc)) +#loc22 = loc("K"(#loc)) +#loc23 = loc("stride_k"(#loc)) +module { + tt.func public @for_ptr_iterargs(%a_ptr: !tt.ptr loc("a_ptr"(#loc)), %b_ptr: !tt.ptr loc("b_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %K: i32 loc("K"(#loc)), %stride_k: i32 loc("stride_k"(#loc))) attributes {noinline = false} { + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %cst = arith.constant dense<64> : tensor<64xi32> loc(#loc2) + %cst_0 = arith.constant dense<0.000000e+00> : tensor<64xf32> loc(#loc2) + %b_ptrs = arith.constant dense<2> : tensor<64xi32> loc(#loc24) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc25) + %a_ptrs = tt.splat %a_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc26) + %a_ptrs_1 = tt.addptr %a_ptrs, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc26) + %b_ptrs_2 = arith.muli %offs, %b_ptrs : tensor<64xi32> loc(#loc24) + %b_ptrs_3 = tt.splat %b_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc27) + %b_ptrs_4 = tt.addptr %b_ptrs_3, %b_ptrs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc27) + %acc:3 = scf.for %k = %c0_i32 to %K step %c64_i32 iter_args(%a_ptrs_5 = %a_ptrs_1, %b_ptrs_6 = %b_ptrs_4, %acc_7 = %cst_0) -> (tensor<64x!tt.ptr>, tensor<64x!tt.ptr>, tensor<64xf32>) : i32 { + %m = arith.subi %K, %k : i32 loc(#loc29) + %m_8 = tt.splat %m : i32 -> tensor<64xi32> loc(#loc30) + %m_9 = arith.cmpi slt, %offs, %m_8 : tensor<64xi32> loc(#loc30) + %acc_10 = tt.load %a_ptrs_5, %m_9, %cst_0 : tensor<64x!tt.ptr> loc(#loc31) + %acc_11 = tt.load %b_ptrs_6, %m_9, %cst_0 : tensor<64x!tt.ptr> loc(#loc32) + %acc_12 = arith.mulf %acc_10, %acc_11 : tensor<64xf32> loc(#loc33) + %acc_13 = arith.addf %acc_7, %acc_12 : tensor<64xf32> loc(#loc34) + %a_ptrs_14 = tt.addptr %a_ptrs_5, %cst : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc35) + %b_ptrs_15 = tt.splat %stride_k : i32 -> tensor<64xi32> loc(#loc36) + %b_ptrs_16 = tt.addptr %b_ptrs_6, %b_ptrs_15 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc36) + scf.yield %a_ptrs_14, %b_ptrs_16, %acc_13 : tensor<64x!tt.ptr>, tensor<64x!tt.ptr>, tensor<64xf32> loc(#loc15) + } loc(#loc38) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc16) + %1 = tt.addptr %0, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc16) + tt.store %1, %acc#2 : tensor<64x!tt.ptr> loc(#loc17) + tt.return loc(#loc18) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":22:25) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":20:28) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":18:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":19:21) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":20:21) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":23:23) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":23:19) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":24:23) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":24:60) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":24:52) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":24:15) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":25:18) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":26:18) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":26:8) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":27:23) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":27:29) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":27:4) +#loc24 = loc("b_ptrs"(#loc3)) +#loc25 = loc("offs"(#loc4)) +#loc26 = loc("a_ptrs"(#loc5)) +#loc27 = loc("b_ptrs"(#loc6)) +#loc28 = loc("a_ptrs"(#loc1)) +#loc29 = loc("m"(#loc7)) +#loc30 = loc("m"(#loc8)) +#loc31 = loc("acc"(#loc9)) +#loc32 = loc("acc"(#loc10)) +#loc33 = loc("acc"(#loc11)) +#loc34 = loc("acc"(#loc12)) +#loc35 = loc("a_ptrs"(#loc13)) +#loc36 = loc("b_ptrs"(#loc14)) +#loc37 = loc("b_ptrs"(#loc28)) +#loc38 = loc("acc"(#loc37)) diff --git a/tests/golden/ir/ttir/spike_i64_index.ttir b/tests/golden/ir/ttir/spike_i64_index.ttir new file mode 100644 index 000000000..af6ebb589 --- /dev/null +++ b/tests/golden/ir/ttir/spike_i64_index.ttir @@ -0,0 +1,51 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":194:0) +#loc14 = loc("x_ptr"(#loc)) +#loc15 = loc("out_ptr"(#loc)) +#loc16 = loc("stride"(#loc)) +#loc17 = loc("big"(#loc)) +module { + tt.func public @i64_index(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %stride: i64 loc("stride"(#loc)), %big: i64 loc("big"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<-4294967296> : tensor<64xi64> loc(#loc1) + %base = arith.constant 1 : i64 loc(#loc18) + %pid = tt.get_program_id x : i32 loc(#loc19) + %pid_0 = arith.extsi %pid : i32 to i64 loc(#loc20) + %base_1 = arith.muli %pid_0, %stride : i64 loc(#loc21) + %base_2 = arith.addi %base_1, %big : i64 loc(#loc22) + %base_3 = arith.subi %base_2, %base : i64 loc(#loc18) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc23) + %offs_4 = arith.extsi %offs : tensor<64xi32> to tensor<64xi64> loc(#loc24) + %offs_5 = tt.splat %base_3 : i64 -> tensor<64xi64> loc(#loc24) + %offs_6 = arith.addi %offs_5, %offs_4 : tensor<64xi64> loc(#loc24) + %v = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc25) + %v_7 = tt.addptr %v, %offs_6 : tensor<64x!tt.ptr>, tensor<64xi64> loc(#loc25) + %v_8 = tt.load %v_7 : tensor<64x!tt.ptr> loc(#loc26) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc11) + %1 = arith.addi %offs_6, %cst : tensor<64xi64> loc(#loc27) + %2 = tt.addptr %0, %1 : tensor<64x!tt.ptr>, tensor<64xi64> loc(#loc27) + tt.store %2, %v_8 : tensor<64x!tt.ptr> loc(#loc12) + tt.return loc(#loc13) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":199:30) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":196:32) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":195:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":195:30) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":196:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":196:26) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":197:31) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":197:18) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":198:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":198:16) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":199:23) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":199:42) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":199:4) +#loc18 = loc("base"(#loc2)) +#loc19 = loc("pid"(#loc3)) +#loc20 = loc("pid"(#loc4)) +#loc21 = loc("base"(#loc5)) +#loc22 = loc("base"(#loc6)) +#loc23 = loc("offs"(#loc7)) +#loc24 = loc("offs"(#loc8)) +#loc25 = loc("v"(#loc9)) +#loc26 = loc("v"(#loc10)) +#loc27 = loc(fused[#loc1, #loc11]) diff --git a/tests/golden/ir/ttir/spike_if_yield.ttir b/tests/golden/ir/ttir/spike_if_yield.ttir new file mode 100644 index 000000000..f772113d9 --- /dev/null +++ b/tests/golden/ir/ttir/spike_if_yield.ttir @@ -0,0 +1,77 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":32:0) +#loc20 = loc("x_ptr"(#loc)) +#loc21 = loc("y_ptr"(#loc)) +#loc22 = loc("out_ptr"(#loc)) +#loc23 = loc("n"(#loc)) +module { + tt.func public @if_yield(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %y_ptr: !tt.ptr loc("y_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<2.000000e+00> : tensor<64xf32> loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + %c2_i32 = arith.constant 2 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc24) + %offs = arith.muli %pid, %c64_i32 : i32 loc(#loc25) + %offs_0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc26) + %offs_1 = tt.splat %offs : i32 -> tensor<64xi32> loc(#loc27) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<64xi32> loc(#loc27) + %m = tt.splat %n : i32 -> tensor<64xi32> loc(#loc28) + %m_3 = arith.cmpi slt, %offs_2, %m : tensor<64xi32> loc(#loc28) + %0 = arith.remsi %pid, %c2_i32 : i32 loc(#loc8) + %1 = arith.cmpi eq, %0, %c0_i32 : i32 loc(#loc2) + %2:2 = scf.if %1 -> (tensor<64x!tt.ptr>, tensor<64xf32>) { + %v = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc29) + %v_4 = tt.addptr %v, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc29) + %v_5 = tt.load %v_4, %m_3 : tensor<64x!tt.ptr> loc(#loc37) + %dst = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc31) + %dst_6 = tt.addptr %dst, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc38) + scf.yield %dst_6, %v_5 : tensor<64x!tt.ptr>, tensor<64xf32> loc(#loc38) + } else { + %v = tt.splat %y_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc32) + %v_4 = tt.addptr %v, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc32) + %v_5 = tt.load %v_4, %m_3 : tensor<64x!tt.ptr> loc(#loc33) + %v_6 = arith.mulf %v_5, %cst : tensor<64xf32> loc(#loc39) + %dst = tt.addptr %out_ptr, %n : !tt.ptr, i32 loc(#loc35) + %dst_7 = tt.splat %dst : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc36) + %dst_8 = tt.addptr %dst_7, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc40) + scf.yield %dst_8, %v_6 : tensor<64x!tt.ptr>, tensor<64xf32> loc(#loc36) + } loc(#loc9) + tt.store %2#0, %2#1, %m_3 : tensor<64x!tt.ptr> loc(#loc18) + tt.return loc(#loc19) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":36:18) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":33:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":34:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":34:38) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":34:25) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":35:15) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":36:13) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":36:7) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":37:28) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":37:20) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":38:24) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":40:28) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":40:20) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":40:44) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":41:24) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":41:28) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":42:18) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":42:4) +#loc24 = loc("pid"(#loc3)) +#loc25 = loc("offs"(#loc4)) +#loc26 = loc("offs"(#loc5)) +#loc27 = loc("offs"(#loc6)) +#loc28 = loc("m"(#loc7)) +#loc29 = loc("v"(#loc10)) +#loc30 = loc("v"(#loc11)) +#loc31 = loc("dst"(#loc12)) +#loc32 = loc("v"(#loc13)) +#loc33 = loc("v"(#loc14)) +#loc34 = loc("v"(#loc15)) +#loc35 = loc("dst"(#loc16)) +#loc36 = loc("dst"(#loc17)) +#loc37 = loc("v"(#loc30)) +#loc38 = loc("dst"(#loc31)) +#loc39 = loc("v"(#loc34)) +#loc40 = loc("dst"(#loc36)) diff --git a/tests/golden/ir/ttir/spike_inline_asm.ttir b/tests/golden/ir/ttir/spike_inline_asm.ttir new file mode 100644 index 000000000..41672383e --- /dev/null +++ b/tests/golden/ir/ttir/spike_inline_asm.ttir @@ -0,0 +1,27 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":159:0) +#loc8 = loc("x_ptr"(#loc)) +#loc9 = loc("out_ptr"(#loc)) +module { + tt.func public @inline_asm(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc))) attributes {noinline = false} { + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc10) + %x = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc11) + %x_0 = tt.addptr %x, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc11) + %x_1 = tt.load %x_0 : tensor<64x!tt.ptr> loc(#loc12) + %y = tt.elementwise_inline_asm "{ .reg .u32 t; mov.u32 t, %tid.x; shl.b32 $0, $1, 3; }" {constraints = "=r,r", packed_element = 1 : i32, pure = true} %x_1 : tensor<64xi32> -> tensor<64xi32> loc(#loc13) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc5) + %1 = tt.addptr %0, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc5) + %2 = tt.elementwise_inline_asm "st.global.b32 [$1], $2; mov.u32 $0, 0;" {constraints = "=r,l,r", packed_element = 1 : i32, pure = false} %1, %y : tensor<64x!tt.ptr>, tensor<64xi32> -> tensor<64xi32> loc(#loc6) + tt.return loc(#loc7) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":160:24) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":161:24) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":161:16) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":165:8) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":173:19) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":173:8) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":170:4) +#loc10 = loc("offs"(#loc1)) +#loc11 = loc("x"(#loc2)) +#loc12 = loc("x"(#loc3)) +#loc13 = loc("y"(#loc4)) diff --git a/tests/golden/ir/ttir/spike_misc.ttir b/tests/golden/ir/ttir/spike_misc.ttir new file mode 100644 index 000000000..11e0c5117 --- /dev/null +++ b/tests/golden/ir/ttir/spike_misc.ttir @@ -0,0 +1,54 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":205:0) +#loc17 = loc("x_ptr"(#loc)) +#loc18 = loc("out_ptr"(#loc)) +#loc19 = loc("n"(#loc)) +module { + tt.func public @misc(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<0.000000e+00> : tensor<64xf32> loc(#loc1) + %v = arith.constant dense<-1.500000e+00> : tensor<64xf32> loc(#loc20) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc21) + %npg = tt.get_num_programs x : i32 loc(#loc22) + %offs = arith.muli %pid, %c64_i32 : i32 loc(#loc23) + %offs_0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc24) + %offs_1 = tt.splat %offs : i32 -> tensor<64xi32> loc(#loc25) + %offs_2 = arith.addi %offs_1, %offs_0 {tt.contiguity = dense<64> : tensor<1xi32>, tt.divisibility = dense<64> : tensor<1xi32>} : tensor<64xi32> loc(#loc25) + %v_3 = tt.splat %n : i32 -> tensor<64xi32> loc(#loc26) + %v_4 = arith.cmpi slt, %offs_2, %v_3 : tensor<64xi32> loc(#loc26) + %v_5 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc27) + %v_6 = tt.addptr %v_5, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc27) + %v_7 = tt.load %v_6, %v_4, %v : tensor<64x!tt.ptr> loc(#loc20) + gpu.barrier loc(#loc10) + tt.print " v{x} loc(: " {hex = false, isSigned = array} : %npg : i32 loc(#loc11) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc12) + %1 = tt.addptr %0, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc12) + %2 = arith.cmpf ogt, %v_7, %cst : tensor<64xf32> loc(#loc13) + %3 = arith.select %2, %v_7, %cst : tensor<64xi1>, tensor<64xf32> loc(#loc14) + tt.store %1, %3, %v_4 : tensor<64x!tt.ptr> loc(#loc15) + tt.return loc(#loc16) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":210:16) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":206:24) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":207:26) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":208:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":208:38) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":208:25) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":210:42) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":210:24) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":211:4) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":212:33) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":213:23) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":213:42) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":213:48) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":213:29) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":213:4) +#loc20 = loc("v"(#loc2)) +#loc21 = loc("pid"(#loc3)) +#loc22 = loc("npg"(#loc4)) +#loc23 = loc("offs"(#loc5)) +#loc24 = loc("offs"(#loc6)) +#loc25 = loc("offs"(#loc7)) +#loc26 = loc("v"(#loc8)) +#loc27 = loc("v"(#loc9)) diff --git a/tests/golden/ir/ttir/spike_nested_for.ttir b/tests/golden/ir/ttir/spike_nested_for.ttir new file mode 100644 index 000000000..c90121328 --- /dev/null +++ b/tests/golden/ir/ttir/spike_nested_for.ttir @@ -0,0 +1,61 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":218:0) +#loc19 = loc("x_ptr"(#loc)) +#loc20 = loc("out_ptr"(#loc)) +#loc21 = loc("n"(#loc)) +module { + tt.func public @nested_for(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %acc = arith.constant dense<0.000000e+00> : tensor<64xf32> loc(#loc33) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc4) + %c64_i32 = arith.constant 64 : i32 loc(#loc3) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc23) + %acc_0 = scf.for %i = %c0_i32 to %n step %c1_i32 iter_args(%acc_1 = %acc) -> (tensor<64xf32>) : i32 { + %acc_2 = scf.for %j = %i to %n step %c1_i32 iter_args(%acc_3 = %acc_1) -> (tensor<64xf32>) : i32 { + %acc_4 = arith.muli %i, %n : i32 loc(#loc26) + %acc_5 = arith.addi %acc_4, %j : i32 loc(#loc27) + %acc_6 = arith.muli %acc_5, %c64_i32 : i32 loc(#loc28) + %acc_7 = tt.addptr %x_ptr, %acc_6 : !tt.ptr, i32 loc(#loc29) + %acc_8 = tt.splat %acc_7 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc30) + %acc_9 = tt.addptr %acc_8, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc30) + %acc_10 = tt.load %acc_9 : tensor<64x!tt.ptr> loc(#loc31) + %acc_11 = arith.addf %acc_3, %acc_10 : tensor<64xf32> loc(#loc32) + scf.yield %acc_11 : tensor<64xf32> loc(#loc14) + } loc(#loc25) + scf.yield %acc_2 : tensor<64xf32> loc(#loc15) + } loc(#loc24) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc16) + %1 = tt.addptr %0, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc16) + tt.store %1, %acc_0 : tensor<64x!tt.ptr> loc(#loc17) + tt.return loc(#loc18) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":129:31) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":220:19) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":221:22) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":219:24) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":222:26) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":223:40) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":223:44) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":223:49) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":223:35) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":223:57) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":223:27) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":223:19) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":223:12) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":222:8) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":224:23) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":224:29) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":224:4) +#loc22 = loc("acc"(#loc2)) +#loc23 = loc("offs"(#loc5)) +#loc24 = loc("acc"(#loc4)) +#loc25 = loc("acc"(#loc6)) +#loc26 = loc("acc"(#loc7)) +#loc27 = loc("acc"(#loc8)) +#loc28 = loc("acc"(#loc9)) +#loc29 = loc("acc"(#loc10)) +#loc30 = loc("acc"(#loc11)) +#loc31 = loc("acc"(#loc12)) +#loc32 = loc("acc"(#loc13)) +#loc33 = loc(callsite(#loc1 at #loc22)) diff --git a/tests/golden/ir/ttir/spike_noinline_call.ttir b/tests/golden/ir/ttir/spike_noinline_call.ttir new file mode 100644 index 000000000..2b5849d75 --- /dev/null +++ b/tests/golden/ir/ttir/spike_noinline_call.ttir @@ -0,0 +1,84 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":113:0) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":99:0) +#loc23 = loc("x_ptr"(#loc)) +#loc24 = loc("out_ptr"(#loc)) +#loc25 = loc("n"(#loc)) +#loc35 = loc("ptr"(#loc16)) +#loc36 = loc("pid"(#loc16)) +#loc37 = loc("n"(#loc16)) +module { + tt.func public @noinline_call(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %v = arith.constant dense<3.000000e+00> : tensor<64xf32> loc(#loc38) + %c64_i32 = arith.constant 64 : i32 loc(#loc3) + %offs = tt.get_program_id x : i32 loc(#loc27) + %offs_0 = arith.muli %offs, %c64_i32 : i32 loc(#loc28) + %offs_1 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc29) + %offs_2 = tt.splat %offs_0 : i32 -> tensor<64xi32> loc(#loc30) + %offs_3 = arith.addi %offs_2, %offs_1 : tensor<64xi32> loc(#loc30) + %v_4 = tt.splat %n : i32 -> tensor<64xi32> loc(#loc31) + %v_5 = arith.cmpi slt, %offs_3, %v_4 : tensor<64xi32> loc(#loc31) + %v_6 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc32) + %v_7 = tt.addptr %v_6, %offs_3 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc32) + %v_8 = tt.load %v_7, %v_5 : tensor<64x!tt.ptr> loc(#loc33) + %v_9 = arith.mulf %v_8, %v : tensor<64xf32> loc(#loc38) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc11) + %1 = tt.addptr %0, %offs_3 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc11) + tt.store %1, %v_9, %v_5 : tensor<64x!tt.ptr> loc(#loc12) + %nxt = tt.call @"corpus_kernels._noinline_store__Pfp32_i32_i32__(2,)cconstexpr_2_d_0_"(%out_ptr, %offs, %n) : (!tt.ptr, i32, i32) -> i32 loc(#loc34) + %2 = tt.call @"corpus_kernels._noinline_store__Pfp32_i32_i32__(2,)cconstexpr_4_d_0_"(%x_ptr, %nxt, %n) : (!tt.ptr, i32, i32) -> i32 loc(#loc14) + tt.return loc(#loc15) + } loc(#loc) + tt.func private @"corpus_kernels._noinline_store__Pfp32_i32_i32__(2,)cconstexpr_2_d_0_"(%ptr: !tt.ptr loc("ptr"(#loc16)), %pid: i32 loc("pid"(#loc16)), %n: i32 loc("n"(#loc16))) -> i32 attributes {noinline = true} { + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %cst = arith.constant 2.000000e+00 : f32 loc(#loc3) + %0 = arith.cmpi slt, %pid, %n : i32 loc(#loc17) + scf.if %0 { + %2 = tt.addptr %ptr, %pid : !tt.ptr, i32 loc(#loc19) + tt.store %2, %cst : !tt.ptr loc(#loc20) + } loc(#loc18) + %1 = arith.addi %pid, %c1_i32 : i32 loc(#loc21) + tt.return %1 : i32 loc(#loc22) + } loc(#loc16) + tt.func private @"corpus_kernels._noinline_store__Pfp32_i32_i32__(2,)cconstexpr_4_d_0_"(%ptr: !tt.ptr loc("ptr"(#loc16)), %pid: i32 loc("pid"(#loc16)), %n: i32 loc("n"(#loc16))) -> i32 attributes {noinline = true} { + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %cst = arith.constant 4.000000e+00 : f32 loc(#loc3) + %0 = arith.cmpi slt, %pid, %n : i32 loc(#loc17) + scf.if %0 { + %2 = tt.addptr %ptr, %pid : !tt.ptr, i32 loc(#loc19) + tt.store %2, %cst : !tt.ptr loc(#loc20) + } loc(#loc18) + %1 = arith.addi %pid, %c1_i32 : i32 loc(#loc21) + tt.return %1 : i32 loc(#loc22) + } loc(#loc16) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":108:15) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":115:23) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":114:25) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":114:30) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":114:51) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":114:38) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":115:57) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":115:39) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":115:31) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":116:23) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":116:29) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":117:58) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":118:37) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":118:4) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":101:13) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":101:7) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":102:23) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":102:28) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":103:17) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":103:11) +#loc26 = loc("v"(#loc2)) +#loc27 = loc("offs"(#loc4)) +#loc28 = loc("offs"(#loc5)) +#loc29 = loc("offs"(#loc6)) +#loc30 = loc("offs"(#loc7)) +#loc31 = loc("v"(#loc8)) +#loc32 = loc("v"(#loc9)) +#loc33 = loc("v"(#loc10)) +#loc34 = loc("nxt"(#loc13)) +#loc38 = loc(callsite(#loc1 at #loc26)) diff --git a/tests/golden/ir/ttir/spike_reduce_scan.ttir b/tests/golden/ir/ttir/spike_reduce_scan.ttir new file mode 100644 index 000000000..c683c89c8 --- /dev/null +++ b/tests/golden/ir/ttir/spike_reduce_scan.ttir @@ -0,0 +1,117 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":148:0) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":151:29) +#loc9 = loc(unknown) +#loc18 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":198:26) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":153:36) +#loc32 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":154:43) +#loc35 = loc("x_ptr"(#loc)) +#loc36 = loc("out_ptr"(#loc)) +#loc41 = loc(callsite(#loc9 at #loc8)) +#loc45 = loc(callsite(#loc18 at #loc19)) +#loc54 = loc(callsite(#loc9 at #loc32)) +#loc57 = loc(callsite(#loc9 at #loc45)) +module { + tt.func public @reduce_scan(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc))) attributes {noinline = false} { + %c3_i32 = arith.constant 3 : i32 loc(#loc1) + %c2_i32 = arith.constant 2 : i32 loc(#loc2) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc37) + %x = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc38) + %x_0 = tt.addptr %x, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc38) + %x_1 = tt.load %x_0 : tensor<64x!tt.ptr> loc(#loc39) + %0 = "tt.reduce"(%x_1) <{axis = 0 : i32}> ({ + ^bb0(%arg2: f32 loc(callsite(#loc9 at #loc8)), %arg3: f32 loc(callsite(#loc9 at #loc8))): + %10 = arith.addf %arg2, %arg3 : f32 loc(#loc55) + tt.reduce.return %10 : f32 loc(#loc40) + }) : (tensor<64xf32>) -> f32 loc(#loc40) + tt.store %out_ptr, %0 : !tt.ptr loc(#loc11) + %1 = tt.addptr %out_ptr, %c1_i32 : !tt.ptr, i32 loc(#loc3) + %2 = "tt.reduce"(%x_1) <{axis = 0 : i32}> ({ + ^bb0(%arg2: f32 loc(unknown), %arg3: f32 loc(unknown)): + %10 = math.absf %arg2 : f32 loc(#loc42) + %11 = math.absf %arg3 : f32 loc(#loc43) + %12 = arith.maxnumf %10, %11 : f32 loc(#loc44) + tt.reduce.return %12 : f32 loc(#loc12) + }) : (tensor<64xf32>) -> f32 loc(#loc12) + tt.store %1, %2 : !tt.ptr loc(#loc16) + %3 = tt.addptr %out_ptr, %c2_i32 : !tt.ptr, i32 loc(#loc2) + %4:2 = "tt.reduce"(%x_1, %offs) <{axis = 0 : i32}> ({ + ^bb0(%arg2: f32 loc(callsite(#loc9 at #loc45)), %arg3: i32 loc(callsite(#loc9 at #loc45)), %arg4: f32 loc(callsite(#loc9 at #loc45)), %arg5: i32 loc(callsite(#loc9 at #loc45))): + %tie = arith.cmpf oeq, %arg2, %arg4 : f32 loc(#loc60) + %tie_2 = arith.cmpi slt, %arg3, %arg5 : i32 loc(#loc61) + %tie_3 = arith.andi %tie, %tie_2 : i1 loc(#loc62) + %gt = arith.cmpf ogt, %arg2, %arg4 : f32 loc(#loc63) + %gt_4 = arith.ori %gt, %tie_3 : i1 loc(#loc64) + %v_ret = arith.select %gt_4, %arg2, %arg4 : f32 loc(#loc65) + %i_ret = arith.select %gt_4, %arg3, %arg5 : i32 loc(#loc66) + tt.reduce.return %v_ret, %i_ret : f32, i32 loc(#loc56) + }) : (tensor<64xf32>, tensor<64xi32>) -> (f32, i32) loc(#loc56) + %5 = arith.sitofp %4#1 : i32 to f32 loc(#loc28) + tt.store %3, %5 : !tt.ptr loc(#loc29) + %6 = tt.addptr %out_ptr, %c3_i32 : !tt.ptr, i32 loc(#loc1) + %7 = tt.splat %6 : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc30) + %8 = tt.addptr %7, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc30) + %9 = "tt.scan"(%x_1) <{axis = 0 : i32, reverse = false}> ({ + ^bb0(%arg2: f32 loc(callsite(#loc9 at #loc32)), %arg3: f32 loc(callsite(#loc9 at #loc32))): + %10 = arith.addf %arg2, %arg3 : f32 loc(#loc58) + tt.scan.return %10 : f32 loc(#loc53) + }) : (tensor<64xf32>) -> tensor<64xf32> loc(#loc53) + tt.store %8, %9 : tensor<64x!tt.ptr> loc(#loc33) + tt.return loc(#loc34) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":154:23) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":153:23) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":152:23) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":149:24) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":150:24) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":150:16) +#loc7 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":293:36) +#loc10 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":263:15) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":151:22) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":152:42) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":142:29) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":142:40) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":142:33) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":152:26) +#loc17 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":181:58) +#loc20 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":149:24) +#loc21 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":160:59) +#loc22 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":149:44) +#loc23 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":149:35) +#loc24 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":152:18) +#loc25 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":152:28) +#loc26 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":153:35) +#loc27 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":154:35) +#loc28 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":153:50) +#loc29 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":153:26) +#loc30 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":154:27) +#loc31 = loc("/home/hwu27/workspace/triton-viz/.venv/lib/python3.12/site-packages/triton/language/standard.py":343:60) +#loc33 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":154:33) +#loc34 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":154:4) +#loc37 = loc("offs"(#loc4)) +#loc38 = loc("x"(#loc5)) +#loc39 = loc("x"(#loc6)) +#loc40 = loc(callsite(#loc7 at #loc8)) +#loc42 = loc(callsite(#loc13 at #loc12)) +#loc43 = loc(callsite(#loc14 at #loc12)) +#loc44 = loc(callsite(#loc15 at #loc12)) +#loc46 = loc("tie"(#loc20)) +#loc47 = loc("tie"(#loc22)) +#loc48 = loc("tie"(#loc23)) +#loc49 = loc("gt"(#loc24)) +#loc50 = loc("gt"(#loc25)) +#loc51 = loc("v_ret"(#loc26)) +#loc52 = loc("i_ret"(#loc27)) +#loc53 = loc(callsite(#loc31 at #loc32)) +#loc55 = loc(callsite(#loc10 at #loc40)) +#loc56 = loc(callsite(#loc17 at #loc45)) +#loc58 = loc(callsite(#loc10 at #loc53)) +#loc59 = loc(callsite(#loc21 at #loc56)) +#loc60 = loc(callsite(#loc46 at #loc59)) +#loc61 = loc(callsite(#loc47 at #loc59)) +#loc62 = loc(callsite(#loc48 at #loc59)) +#loc63 = loc(callsite(#loc49 at #loc59)) +#loc64 = loc(callsite(#loc50 at #loc59)) +#loc65 = loc(callsite(#loc51 at #loc59)) +#loc66 = loc(callsite(#loc52 at #loc59)) diff --git a/tests/golden/ir/ttir/spike_spin_while.ttir b/tests/golden/ir/ttir/spike_spin_while.ttir new file mode 100644 index 000000000..96ebb6103 --- /dev/null +++ b/tests/golden/ir/ttir/spike_spin_while.ttir @@ -0,0 +1,47 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":72:0) +#loc10 = loc("v") +#loc15 = loc("lock_ptr"(#loc)) +#loc16 = loc("flag_ptr"(#loc)) +#loc17 = loc("out_ptr"(#loc)) +module { + tt.func public @spin_while(%lock_ptr: !tt.ptr loc("lock_ptr"(#loc)), %flag_ptr: !tt.ptr loc("flag_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc))) attributes {noinline = false} { + %true = arith.constant true loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc2) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + scf.while : () -> () { + %1 = tt.atomic_cas acquire, gpu, %lock_ptr, %c0_i32, %c1_i32 : (!tt.ptr, i32, i32) -> i32 loc(#loc4) + %2 = arith.cmpi eq, %1, %c1_i32 : i32 loc(#loc5) + scf.condition(%2) loc(#loc5) + } do { + scf.yield loc(#loc6) + } loc(#loc3) + %v = tt.load %flag_ptr {isVolatile = true} : !tt.ptr loc(#loc18) + %v_0 = scf.while (%v_1 = %v) : (i32) -> i32 { + %1 = arith.cmpi eq, %v_1, %c0_i32 : i32 loc(#loc9) + scf.condition(%1) %v_1 : i32 loc(#loc9) + } do { + ^bb0(%v_1: i32 loc("v")): + %v_2 = tt.load %flag_ptr {isVolatile = true} : !tt.ptr loc(#loc20) + scf.yield %v_2 : i32 loc(#loc12) + } loc(#loc19) + tt.store %out_ptr, %v_0 : !tt.ptr loc(#loc13) + %0 = tt.atomic_rmw exch, release, gpu, %lock_ptr, %c0_i32, %true : (!tt.ptr, i32, i1) -> i32 loc(#loc1) + tt.return loc(#loc14) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":79:29) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":73:4) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":73:37) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":73:71) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":74:8) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":75:16) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":76:4) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":76:15) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":77:20) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":77:8) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":78:22) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":79:4) +#loc18 = loc("v"(#loc7)) +#loc19 = loc("v"(#loc8)) +#loc20 = loc("v"(#loc11)) diff --git a/tests/golden/ir/ttir/spike_tile2d_i64.ttir b/tests/golden/ir/ttir/spike_tile2d_i64.ttir new file mode 100644 index 000000000..7cf75b946 --- /dev/null +++ b/tests/golden/ir/ttir/spike_tile2d_i64.ttir @@ -0,0 +1,85 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":118:0) +#loc23 = loc("x_ptr"(#loc)) +#loc24 = loc("out_ptr"(#loc)) +#loc25 = loc("M"(#loc)) +#loc26 = loc("N"(#loc)) +#loc27 = loc("stride_m"(#loc)) +module { + tt.func public @tile2d_i64(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %M: i32 loc("M"(#loc)), %N: i32 loc("N"(#loc)), %stride_m: i64 loc("stride_m"(#loc))) attributes {noinline = false} { + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %rm = arith.constant 16 : i64 loc(#loc28) + %pid_m = tt.get_program_id x : i32 loc(#loc29) + %pid_m_0 = arith.extsi %pid_m : i32 to i64 loc(#loc30) + %pid_n = tt.get_program_id y : i32 loc(#loc31) + %rm_1 = arith.muli %pid_m_0, %rm : i64 loc(#loc28) + %rm_2 = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32> loc(#loc32) + %rm_3 = arith.extsi %rm_2 : tensor<16xi32> to tensor<16xi64> loc(#loc33) + %rm_4 = tt.splat %rm_1 : i64 -> tensor<16xi64> loc(#loc33) + %rm_5 = arith.addi %rm_4, %rm_3 : tensor<16xi64> loc(#loc33) + %rn = arith.muli %pid_n, %c32_i32 : i32 loc(#loc34) + %rn_6 = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc35) + %rn_7 = tt.splat %rn : i32 -> tensor<32xi32> loc(#loc36) + %rn_8 = arith.addi %rn_7, %rn_6 : tensor<32xi32> loc(#loc36) + %offs = tt.expand_dims %rm_5 {axis = 1 : i32} : tensor<16xi64> -> tensor<16x1xi64> loc(#loc37) + %offs_9 = tt.splat %stride_m : i64 -> tensor<16x1xi64> loc(#loc38) + %offs_10 = arith.muli %offs, %offs_9 : tensor<16x1xi64> loc(#loc38) + %offs_11 = tt.expand_dims %rn_8 {axis = 0 : i32} : tensor<32xi32> -> tensor<1x32xi32> loc(#loc39) + %offs_12 = arith.extsi %offs_11 : tensor<1x32xi32> to tensor<1x32xi64> loc(#loc40) + %offs_13 = tt.broadcast %offs_10 : tensor<16x1xi64> -> tensor<16x32xi64> loc(#loc40) + %offs_14 = tt.broadcast %offs_12 : tensor<1x32xi64> -> tensor<16x32xi64> loc(#loc40) + %offs_15 = arith.addi %offs_13, %offs_14 : tensor<16x32xi64> loc(#loc40) + %m = arith.extsi %M : i32 to i64 loc(#loc41) + %m_16 = tt.splat %m : i64 -> tensor<16x1xi64> loc(#loc41) + %m_17 = arith.cmpi slt, %offs, %m_16 : tensor<16x1xi64> loc(#loc41) + %m_18 = tt.splat %N : i32 -> tensor<1x32xi32> loc(#loc42) + %m_19 = arith.cmpi slt, %offs_11, %m_18 : tensor<1x32xi32> loc(#loc42) + %m_20 = tt.broadcast %m_17 : tensor<16x1xi1> -> tensor<16x32xi1> loc(#loc43) + %m_21 = tt.broadcast %m_19 : tensor<1x32xi1> -> tensor<16x32xi1> loc(#loc43) + %m_22 = arith.andi %m_20, %m_21 : tensor<16x32xi1> loc(#loc43) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<16x32x!tt.ptr> loc(#loc18) + %1 = tt.addptr %0, %offs_15 : tensor<16x32x!tt.ptr>, tensor<16x32xi64> loc(#loc18) + %2 = tt.splat %x_ptr : !tt.ptr -> tensor<16x32x!tt.ptr> loc(#loc19) + %3 = tt.addptr %2, %offs_15 : tensor<16x32x!tt.ptr>, tensor<16x32xi64> loc(#loc19) + %4 = tt.load %3, %m_22 : tensor<16x32x!tt.ptr> loc(#loc20) + tt.store %1, %4, %m_22 : tensor<16x32x!tt.ptr> loc(#loc21) + tt.return loc(#loc22) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":121:17) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":119:26) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":119:32) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":120:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":121:35) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":121:22) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":122:17) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":122:35) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":122:22) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":123:14) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":123:25) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":123:39) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":123:36) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":124:23) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":124:43) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":124:29) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":125:23) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":125:45) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":125:37) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":125:29) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/corpus_kernels.py":125:4) +#loc28 = loc("rm"(#loc2)) +#loc29 = loc("pid_m"(#loc3)) +#loc30 = loc("pid_m"(#loc4)) +#loc31 = loc("pid_n"(#loc5)) +#loc32 = loc("rm"(#loc6)) +#loc33 = loc("rm"(#loc7)) +#loc34 = loc("rn"(#loc8)) +#loc35 = loc("rn"(#loc9)) +#loc36 = loc("rn"(#loc10)) +#loc37 = loc("offs"(#loc11)) +#loc38 = loc("offs"(#loc12)) +#loc39 = loc("offs"(#loc13)) +#loc40 = loc("offs"(#loc14)) +#loc41 = loc("m"(#loc15)) +#loc42 = loc("m"(#loc16)) +#loc43 = loc("m"(#loc17)) diff --git a/tests/golden/ir/ttir_3.8/adv_descs.ttir b/tests/golden/ir/ttir_3.8/adv_descs.ttir new file mode 100644 index 000000000..869d16a57 --- /dev/null +++ b/tests/golden/ir/ttir_3.8/adv_descs.ttir @@ -0,0 +1,35 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":112:1) +#loc10 = loc("a_ptr"(#loc)) +#loc11 = loc("M"(#loc)) +#loc12 = loc("N"(#loc)) +module { + tt.func public @descs(%a_ptr: !tt.ptr loc("a_ptr"(#loc)), %M: i32 loc("M"(#loc)), %N: i32 loc("N"(#loc))) attributes {noinline = false} { + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc1) + %c1_i64 = arith.constant 1 : i64 loc(#loc1) + %d = arith.extsi %N : i32 to i64 loc(#loc13) + %d_0 = tt.make_tensor_descriptor %a_ptr, [%M, %N], [%d, %c1_i64] : , <32x32xf16> loc(#loc13) + %x = tt.descriptor_load %d_0[%c0_i32, %c32_i32] : !tt.tensordesc<32x32xf16> -> tensor<32x32xf16> loc(#loc14) + tt.descriptor_store %d_0[%c32_i32, %c0_i32], %x : !tt.tensordesc<32x32xf16>, tensor<32x32xf16> loc(#loc4) + tt.descriptor_reduce add, %d_0[%c32_i32, %c32_i32], %x : !tt.tensordesc<32x32xf16>, tensor<32x32xf16> loc(#loc5) + %d1 = tt.make_tensor_descriptor %a_ptr, [%M, %N], [%d, %c1_i64] : , <1x32xf16> loc(#loc15) + %rows = tt.make_range {end = 32 : i32, start = 0 : i32} : tensor<32xi32> loc(#loc16) + %g = tt.descriptor_gather %d1[%rows, %c0_i32] : (!tt.tensordesc<1x32xf16>, tensor<32xi32>, i32) -> tensor<32x32xf16> loc(#loc17) + tt.descriptor_scatter %d1[%rows, %c32_i32], %g : !tt.tensordesc<1x32xf16>, tensor<32xi32>, i32, tensor<32x32xf16> loc(#loc9) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":115:9) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":116:9) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":117:5) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":118:5) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":119:10) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":120:12) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":121:9) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":122:5) +#loc13 = loc("d"(#loc2)) +#loc14 = loc("x"(#loc3)) +#loc15 = loc("d1"(#loc6)) +#loc16 = loc("rows"(#loc7)) +#loc17 = loc("g"(#loc8)) diff --git a/tests/golden/ir/ttir_3.8/adv_zero_result.ttir b/tests/golden/ir/ttir_3.8/adv_zero_result.ttir new file mode 100644 index 000000000..25c614caa --- /dev/null +++ b/tests/golden/ir/ttir_3.8/adv_zero_result.ttir @@ -0,0 +1,56 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":190:1) +#loc17 = loc("x_ptr"(#loc)) +#loc18 = loc("out_ptr"(#loc)) +#loc19 = loc("n"(#loc)) +module { + tt.func public @zero_result(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %true = arith.constant true loc(#loc1) + %c1_i32 = arith.constant 1 : i32 loc(#loc1) + %c0_i32 = arith.constant 0 : i32 loc(#loc2) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc20) + %offs = arith.muli %pid, %c64_i32 : i32 loc(#loc21) + %offs_0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc22) + %offs_1 = tt.splat %offs : i32 -> tensor<64xi32> loc(#loc21) + %offs_2 = arith.addi %offs_1, %offs_0 : tensor<64xi32> loc(#loc21) + %v = tt.splat %n : i32 -> tensor<64xi32> loc(#loc23) + %v_3 = arith.cmpi slt, %offs_2, %v : tensor<64xi32> loc(#loc23) + %v_4 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc24) + %v_5 = tt.addptr %v_4, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc24) + %v_6 = tt.load %v_5, %v_3 : tensor<64x!tt.ptr> loc(#loc25) + tt.print " pid=: " {hex = true, isSigned = array} : %pid, %offs_2 : i32, tensor<64xi32> loc(#loc9) + ttg.barrier all loc(#loc10) + %0 = arith.cmpi eq, %pid, %c0_i32 : i32 loc(#loc2) + scf.if %0 { + %5 = tt.atomic_rmw add, acq_rel, gpu, %out_ptr, %c1_i32, %true : (!tt.ptr, i32, i1) -> i32 loc(#loc12) + } loc(#loc11) + %1 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc13) + %2 = tt.addptr %1, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc13) + %3 = arith.fptosi %v_6 : tensor<64xf32> to tensor<64xi32> loc(#loc14) + %4 = tt.atomic_rmw max, acq_rel, gpu, %2, %3, %v_3 : (tensor<64x!tt.ptr>, tensor<64xi32>, tensor<64xi1>) -> tensor<64xi32> loc(#loc15) + tt.store %2, %3, %v_3 : tensor<64x!tt.ptr> loc(#loc16) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":198:8) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":192:11) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":193:12) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":193:26) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":194:36) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":194:17) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":194:9) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":196:5) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":197:5) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":198:5) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":199:9) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":200:19) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":200:35) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":200:5) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":201:5) +#loc20 = loc("pid"(#loc3)) +#loc21 = loc("offs"(#loc4)) +#loc22 = loc("offs"(#loc5)) +#loc23 = loc("v"(#loc6)) +#loc24 = loc("v"(#loc7)) +#loc25 = loc("v"(#loc8)) diff --git a/tests/golden/ir/ttir_3.8/golden_matmul_tma_s1_sm90.ttir b/tests/golden/ir/ttir_3.8/golden_matmul_tma_s1_sm90.ttir new file mode 100644 index 000000000..f0bc31313 --- /dev/null +++ b/tests/golden/ir/ttir_3.8/golden_matmul_tma_s1_sm90.ttir @@ -0,0 +1,76 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":126:1) +#loc21 = loc("a_ptr"(#loc)) +#loc22 = loc("b_ptr"(#loc)) +#loc23 = loc("c_ptr"(#loc)) +#loc24 = loc("M"(#loc)) +#loc25 = loc("N"(#loc)) +#loc26 = loc("K"(#loc)) +module { + tt.func public @matmul_tma_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc))) attributes {noinline = false} { + %c31_i32 = arith.constant 31 : i32 loc(#loc27) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %cst = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %c1_i64 = arith.constant 1 : i64 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc28) + %pid_n = tt.get_program_id y : i32 loc(#loc29) + %a_desc = arith.extsi %K : i32 to i64 loc(#loc30) + %a_desc_0 = tt.make_tensor_descriptor %a_ptr, [%M, %K], [%a_desc, %c1_i64] : , <64x32xf16> loc(#loc30) + %b_desc = arith.extsi %N : i32 to i64 loc(#loc31) + %b_desc_1 = tt.make_tensor_descriptor %b_ptr, [%K, %N], [%b_desc, %c1_i64] : , <32x64xf16> loc(#loc31) + %c_desc = tt.make_tensor_descriptor %c_ptr, [%M, %N], [%b_desc, %c1_i64] : , <64x64xf16> loc(#loc32) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc33) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc34) + %acc = scf.for %k = %c0_i32 to %1 step %c1_i32 iter_args(%acc_2 = %cst) -> (tensor<64x64xf32>) : i32 { + %a = arith.muli %pid_m, %c64_i32 : i32 loc(#loc36) + %a_3 = arith.muli %k, %c32_i32 : i32 loc(#loc37) + %a_4 = tt.descriptor_load %a_desc_0[%a, %a_3] : !tt.tensordesc<64x32xf16> -> tensor<64x32xf16> loc(#loc38) + %b = arith.muli %pid_n, %c64_i32 : i32 loc(#loc39) + %b_5 = tt.descriptor_load %b_desc_1[%a_3, %b] : !tt.tensordesc<32x64xf16> -> tensor<32x64xf16> loc(#loc40) + %acc_6 = tt.dot %a_4, %b_5, %acc_2, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc41) + scf.yield %acc_6 : tensor<64x64xf32> loc(#loc3) + } loc(#loc35) + %2 = arith.muli %pid_m, %c64_i32 : i32 loc(#loc17) + %3 = arith.muli %pid_n, %c64_i32 : i32 loc(#loc18) + %4 = arith.truncf %acc : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc19) + tt.descriptor_store %c_desc[%2, %3], %4 : !tt.tensordesc<64x64xf16>, tensor<64x64xf16> loc(#loc20) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":150:23) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":150:5) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":138:13) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":139:13) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":140:14) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":143:14) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":146:14) +#loc9 = loc("/tmp/claude-1003/-home-hwu27-workspace-triton-viz/7d3c8012-b668-4397-a82f-ef186562dd00/scratchpad/triton38_overlay/triton/language/standard.py":43:13) +#loc10 = loc("/tmp/claude-1003/-home-hwu27-workspace-triton-viz/7d3c8012-b668-4397-a82f-ef186562dd00/scratchpad/triton38_overlay/triton/language/standard.py":43:12) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":151:26) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":151:43) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":151:13) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":152:39) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":152:13) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":153:16) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":154:19) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":154:36) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":154:54) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":154:5) +#loc27 = loc(callsite(#loc1 at #loc2)) +#loc28 = loc("pid_m"(#loc4)) +#loc29 = loc("pid_n"(#loc5)) +#loc30 = loc("a_desc"(#loc6)) +#loc31 = loc("b_desc"(#loc7)) +#loc32 = loc("c_desc"(#loc8)) +#loc33 = loc(callsite(#loc9 at #loc2)) +#loc34 = loc(callsite(#loc10 at #loc2)) +#loc35 = loc("acc"(#loc3)) +#loc36 = loc("a"(#loc11)) +#loc37 = loc("a"(#loc12)) +#loc38 = loc("a"(#loc13)) +#loc39 = loc("b"(#loc14)) +#loc40 = loc("b"(#loc15)) +#loc41 = loc("acc"(#loc16)) diff --git a/tests/golden/ir/ttir_3.8/golden_matmul_tma_s3_sm90.ttir b/tests/golden/ir/ttir_3.8/golden_matmul_tma_s3_sm90.ttir new file mode 100644 index 000000000..f0bc31313 --- /dev/null +++ b/tests/golden/ir/ttir_3.8/golden_matmul_tma_s3_sm90.ttir @@ -0,0 +1,76 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":126:1) +#loc21 = loc("a_ptr"(#loc)) +#loc22 = loc("b_ptr"(#loc)) +#loc23 = loc("c_ptr"(#loc)) +#loc24 = loc("M"(#loc)) +#loc25 = loc("N"(#loc)) +#loc26 = loc("K"(#loc)) +module { + tt.func public @matmul_tma_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc))) attributes {noinline = false} { + %c31_i32 = arith.constant 31 : i32 loc(#loc27) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %cst = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %c1_i64 = arith.constant 1 : i64 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc28) + %pid_n = tt.get_program_id y : i32 loc(#loc29) + %a_desc = arith.extsi %K : i32 to i64 loc(#loc30) + %a_desc_0 = tt.make_tensor_descriptor %a_ptr, [%M, %K], [%a_desc, %c1_i64] : , <64x32xf16> loc(#loc30) + %b_desc = arith.extsi %N : i32 to i64 loc(#loc31) + %b_desc_1 = tt.make_tensor_descriptor %b_ptr, [%K, %N], [%b_desc, %c1_i64] : , <32x64xf16> loc(#loc31) + %c_desc = tt.make_tensor_descriptor %c_ptr, [%M, %N], [%b_desc, %c1_i64] : , <64x64xf16> loc(#loc32) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc33) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc34) + %acc = scf.for %k = %c0_i32 to %1 step %c1_i32 iter_args(%acc_2 = %cst) -> (tensor<64x64xf32>) : i32 { + %a = arith.muli %pid_m, %c64_i32 : i32 loc(#loc36) + %a_3 = arith.muli %k, %c32_i32 : i32 loc(#loc37) + %a_4 = tt.descriptor_load %a_desc_0[%a, %a_3] : !tt.tensordesc<64x32xf16> -> tensor<64x32xf16> loc(#loc38) + %b = arith.muli %pid_n, %c64_i32 : i32 loc(#loc39) + %b_5 = tt.descriptor_load %b_desc_1[%a_3, %b] : !tt.tensordesc<32x64xf16> -> tensor<32x64xf16> loc(#loc40) + %acc_6 = tt.dot %a_4, %b_5, %acc_2, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc41) + scf.yield %acc_6 : tensor<64x64xf32> loc(#loc3) + } loc(#loc35) + %2 = arith.muli %pid_m, %c64_i32 : i32 loc(#loc17) + %3 = arith.muli %pid_n, %c64_i32 : i32 loc(#loc18) + %4 = arith.truncf %acc : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc19) + tt.descriptor_store %c_desc[%2, %3], %4 : !tt.tensordesc<64x64xf16>, tensor<64x64xf16> loc(#loc20) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":150:23) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":150:5) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":138:13) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":139:13) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":140:14) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":143:14) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":146:14) +#loc9 = loc("/tmp/claude-1003/-home-hwu27-workspace-triton-viz/7d3c8012-b668-4397-a82f-ef186562dd00/scratchpad/triton38_overlay/triton/language/standard.py":43:13) +#loc10 = loc("/tmp/claude-1003/-home-hwu27-workspace-triton-viz/7d3c8012-b668-4397-a82f-ef186562dd00/scratchpad/triton38_overlay/triton/language/standard.py":43:12) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":151:26) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":151:43) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":151:13) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":152:39) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":152:13) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":153:16) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":154:19) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":154:36) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":154:54) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":154:5) +#loc27 = loc(callsite(#loc1 at #loc2)) +#loc28 = loc("pid_m"(#loc4)) +#loc29 = loc("pid_n"(#loc5)) +#loc30 = loc("a_desc"(#loc6)) +#loc31 = loc("b_desc"(#loc7)) +#loc32 = loc("c_desc"(#loc8)) +#loc33 = loc(callsite(#loc9 at #loc2)) +#loc34 = loc(callsite(#loc10 at #loc2)) +#loc35 = loc("acc"(#loc3)) +#loc36 = loc("a"(#loc11)) +#loc37 = loc("a"(#loc12)) +#loc38 = loc("a"(#loc13)) +#loc39 = loc("b"(#loc14)) +#loc40 = loc("b"(#loc15)) +#loc41 = loc("acc"(#loc16)) diff --git a/tests/golden/ir/ttir_3.8/golden_matmul_tma_ws_s3_sm90.ttir b/tests/golden/ir/ttir_3.8/golden_matmul_tma_ws_s3_sm90.ttir new file mode 100644 index 000000000..b7dbdb331 --- /dev/null +++ b/tests/golden/ir/ttir_3.8/golden_matmul_tma_ws_s3_sm90.ttir @@ -0,0 +1,76 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":158:1) +#loc21 = loc("a_ptr"(#loc)) +#loc22 = loc("b_ptr"(#loc)) +#loc23 = loc("c_ptr"(#loc)) +#loc24 = loc("M"(#loc)) +#loc25 = loc("N"(#loc)) +#loc26 = loc("K"(#loc)) +module { + tt.func public @matmul_tma_ws_kernel(%a_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("a_ptr"(#loc)), %b_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("b_ptr"(#loc)), %c_ptr: !tt.ptr {tt.divisibility = 16 : i32} loc("c_ptr"(#loc)), %M: i32 {tt.divisibility = 16 : i32} loc("M"(#loc)), %N: i32 {tt.divisibility = 16 : i32} loc("N"(#loc)), %K: i32 {tt.divisibility = 16 : i32} loc("K"(#loc))) attributes {noinline = false} { + %c31_i32 = arith.constant 31 : i32 loc(#loc27) + %c1_i32 = arith.constant 1 : i32 loc(#loc3) + %c0_i32 = arith.constant 0 : i32 loc(#loc3) + %cst = arith.constant dense<0.000000e+00> : tensor<64x64xf32> loc(#loc1) + %c32_i32 = arith.constant 32 : i32 loc(#loc1) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %c1_i64 = arith.constant 1 : i64 loc(#loc1) + %pid_m = tt.get_program_id x : i32 loc(#loc28) + %pid_n = tt.get_program_id y : i32 loc(#loc29) + %a_desc = arith.extsi %K : i32 to i64 loc(#loc30) + %a_desc_0 = tt.make_tensor_descriptor %a_ptr, [%M, %K], [%a_desc, %c1_i64] : , <64x32xf16> loc(#loc30) + %b_desc = arith.extsi %N : i32 to i64 loc(#loc31) + %b_desc_1 = tt.make_tensor_descriptor %b_ptr, [%K, %N], [%b_desc, %c1_i64] : , <32x64xf16> loc(#loc31) + %c_desc = tt.make_tensor_descriptor %c_ptr, [%M, %N], [%b_desc, %c1_i64] : , <64x64xf16> loc(#loc32) + %0 = arith.addi %K, %c31_i32 : i32 loc(#loc33) + %1 = arith.divsi %0, %c32_i32 : i32 loc(#loc34) + %acc = scf.for %k = %c0_i32 to %1 step %c1_i32 iter_args(%acc_2 = %cst) -> (tensor<64x64xf32>) : i32 { + %a = arith.muli %pid_m, %c64_i32 : i32 loc(#loc36) + %a_3 = arith.muli %k, %c32_i32 : i32 loc(#loc37) + %a_4 = tt.descriptor_load %a_desc_0[%a, %a_3] : !tt.tensordesc<64x32xf16> -> tensor<64x32xf16> loc(#loc38) + %b = arith.muli %pid_n, %c64_i32 : i32 loc(#loc39) + %b_5 = tt.descriptor_load %b_desc_1[%a_3, %b] : !tt.tensordesc<32x64xf16> -> tensor<32x64xf16> loc(#loc40) + %acc_6 = tt.dot %a_4, %b_5, %acc_2, inputPrecision = tf32 : tensor<64x32xf16> * tensor<32x64xf16> -> tensor<64x64xf32> loc(#loc41) + scf.yield %acc_6 : tensor<64x64xf32> loc(#loc3) + } {tt.warp_specialize} loc(#loc35) + %2 = arith.muli %pid_m, %c64_i32 : i32 loc(#loc17) + %3 = arith.muli %pid_n, %c64_i32 : i32 loc(#loc18) + %4 = arith.truncf %acc : tensor<64x64xf32> to tensor<64x64xf16> loc(#loc19) + tt.descriptor_store %c_desc[%2, %3], %4 : !tt.tensordesc<64x64xf16>, tensor<64x64xf16> loc(#loc20) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":182:26) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":182:5) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":170:13) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":171:13) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":172:14) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":175:14) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":178:14) +#loc9 = loc("/tmp/claude-1003/-home-hwu27-workspace-triton-viz/7d3c8012-b668-4397-a82f-ef186562dd00/scratchpad/triton38_overlay/triton/language/standard.py":43:13) +#loc10 = loc("/tmp/claude-1003/-home-hwu27-workspace-triton-viz/7d3c8012-b668-4397-a82f-ef186562dd00/scratchpad/triton38_overlay/triton/language/standard.py":43:12) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":183:26) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":183:43) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":183:13) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":184:39) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":184:13) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":185:16) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":186:19) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":186:36) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":186:54) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":186:5) +#loc27 = loc(callsite(#loc1 at #loc2)) +#loc28 = loc("pid_m"(#loc4)) +#loc29 = loc("pid_n"(#loc5)) +#loc30 = loc("a_desc"(#loc6)) +#loc31 = loc("b_desc"(#loc7)) +#loc32 = loc("c_desc"(#loc8)) +#loc33 = loc(callsite(#loc9 at #loc2)) +#loc34 = loc(callsite(#loc10 at #loc2)) +#loc35 = loc("acc"(#loc3)) +#loc36 = loc("a"(#loc11)) +#loc37 = loc("a"(#loc12)) +#loc38 = loc("a"(#loc13)) +#loc39 = loc("b"(#loc14)) +#loc40 = loc("b"(#loc15)) +#loc41 = loc("acc"(#loc16)) diff --git a/tests/golden/ir/ttir_3.8/kernel_deep_chain.ttir b/tests/golden/ir/ttir_3.8/kernel_deep_chain.ttir new file mode 100644 index 000000000..df5bf1184 --- /dev/null +++ b/tests/golden/ir/ttir_3.8/kernel_deep_chain.ttir @@ -0,0 +1,1218 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":60:1) +#loc5 = loc("out_ptr"(#loc)) +#loc6 = loc("s"(#loc)) +module { + tt.func public @deep_chain(%out_ptr: !tt.ptr loc("out_ptr"(#loc)), %s: i32 loc("s"(#loc))) attributes {noinline = false} { + %cst = arith.constant 1.000000e+00 : f32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc7) + %off = arith.muli %pid, %s : i32 loc(#loc8) + %off_0 = arith.addi %off, %pid : i32 loc(#loc8) + %off_1 = arith.muli %off_0, %s : i32 loc(#loc8) + %off_2 = arith.addi %off_1, %pid : i32 loc(#loc8) + %off_3 = arith.muli %off_2, %s : i32 loc(#loc8) + %off_4 = arith.addi %off_3, %pid : i32 loc(#loc8) + %off_5 = arith.muli %off_4, %s : i32 loc(#loc8) + %off_6 = arith.addi %off_5, %pid : i32 loc(#loc8) + %off_7 = arith.muli %off_6, %s : i32 loc(#loc8) + %off_8 = arith.addi %off_7, %pid : i32 loc(#loc8) + %off_9 = arith.muli %off_8, %s : i32 loc(#loc8) + %off_10 = arith.addi %off_9, %pid : i32 loc(#loc8) + %off_11 = arith.muli %off_10, %s : i32 loc(#loc8) + %off_12 = arith.addi %off_11, %pid : i32 loc(#loc8) + %off_13 = arith.muli %off_12, %s : i32 loc(#loc8) + %off_14 = arith.addi %off_13, %pid : i32 loc(#loc8) + %off_15 = arith.muli %off_14, %s : i32 loc(#loc8) + %off_16 = arith.addi %off_15, %pid : i32 loc(#loc8) + %off_17 = arith.muli %off_16, %s : i32 loc(#loc8) + %off_18 = arith.addi %off_17, %pid : i32 loc(#loc8) + %off_19 = arith.muli %off_18, %s : i32 loc(#loc8) + %off_20 = arith.addi %off_19, %pid : i32 loc(#loc8) + %off_21 = arith.muli %off_20, %s : i32 loc(#loc8) + %off_22 = arith.addi %off_21, %pid : i32 loc(#loc8) + %off_23 = arith.muli %off_22, %s : i32 loc(#loc8) + %off_24 = arith.addi %off_23, %pid : i32 loc(#loc8) + %off_25 = arith.muli %off_24, %s : i32 loc(#loc8) + %off_26 = arith.addi %off_25, %pid : i32 loc(#loc8) + %off_27 = arith.muli %off_26, %s : i32 loc(#loc8) + %off_28 = arith.addi %off_27, %pid : i32 loc(#loc8) + %off_29 = arith.muli %off_28, %s : i32 loc(#loc8) + %off_30 = arith.addi %off_29, %pid : i32 loc(#loc8) + %off_31 = arith.muli %off_30, %s : i32 loc(#loc8) + %off_32 = arith.addi %off_31, %pid : i32 loc(#loc8) + %off_33 = arith.muli %off_32, %s : i32 loc(#loc8) + %off_34 = arith.addi %off_33, %pid : i32 loc(#loc8) + %off_35 = arith.muli %off_34, %s : i32 loc(#loc8) + %off_36 = arith.addi %off_35, %pid : i32 loc(#loc8) + %off_37 = arith.muli %off_36, %s : i32 loc(#loc8) + %off_38 = arith.addi %off_37, %pid : i32 loc(#loc8) + %off_39 = arith.muli %off_38, %s : i32 loc(#loc8) + %off_40 = arith.addi %off_39, %pid : i32 loc(#loc8) + %off_41 = arith.muli %off_40, %s : i32 loc(#loc8) + %off_42 = arith.addi %off_41, %pid : i32 loc(#loc8) + %off_43 = arith.muli %off_42, %s : i32 loc(#loc8) + %off_44 = arith.addi %off_43, %pid : i32 loc(#loc8) + %off_45 = arith.muli %off_44, %s : i32 loc(#loc8) + %off_46 = arith.addi %off_45, %pid : i32 loc(#loc8) + %off_47 = arith.muli %off_46, %s : i32 loc(#loc8) + %off_48 = arith.addi %off_47, %pid : i32 loc(#loc8) + %off_49 = arith.muli %off_48, %s : i32 loc(#loc8) + %off_50 = arith.addi %off_49, %pid : i32 loc(#loc8) + %off_51 = arith.muli %off_50, %s : i32 loc(#loc8) + %off_52 = arith.addi %off_51, %pid : i32 loc(#loc8) + %off_53 = arith.muli %off_52, %s : i32 loc(#loc8) + %off_54 = arith.addi %off_53, %pid : i32 loc(#loc8) + %off_55 = arith.muli %off_54, %s : i32 loc(#loc8) + %off_56 = arith.addi %off_55, %pid : i32 loc(#loc8) + %off_57 = arith.muli %off_56, %s : i32 loc(#loc8) + %off_58 = arith.addi %off_57, %pid : i32 loc(#loc8) + %off_59 = arith.muli %off_58, %s : i32 loc(#loc8) + %off_60 = arith.addi %off_59, %pid : i32 loc(#loc8) + %off_61 = arith.muli %off_60, %s : i32 loc(#loc8) + %off_62 = arith.addi %off_61, %pid : i32 loc(#loc8) + %off_63 = arith.muli %off_62, %s : i32 loc(#loc8) + %off_64 = arith.addi %off_63, %pid : i32 loc(#loc8) + %off_65 = arith.muli %off_64, %s : i32 loc(#loc8) + %off_66 = arith.addi %off_65, %pid : i32 loc(#loc8) + %off_67 = arith.muli %off_66, %s : i32 loc(#loc8) + %off_68 = arith.addi %off_67, %pid : i32 loc(#loc8) + %off_69 = arith.muli %off_68, %s : i32 loc(#loc8) + %off_70 = arith.addi %off_69, %pid : i32 loc(#loc8) + %off_71 = arith.muli %off_70, %s : i32 loc(#loc8) + %off_72 = arith.addi %off_71, %pid : i32 loc(#loc8) + %off_73 = arith.muli %off_72, %s : i32 loc(#loc8) + %off_74 = arith.addi %off_73, %pid : i32 loc(#loc8) + %off_75 = arith.muli %off_74, %s : i32 loc(#loc8) + %off_76 = arith.addi %off_75, %pid : i32 loc(#loc8) + %off_77 = arith.muli %off_76, %s : i32 loc(#loc8) + %off_78 = arith.addi %off_77, %pid : i32 loc(#loc8) + %off_79 = arith.muli %off_78, %s : i32 loc(#loc8) + %off_80 = arith.addi %off_79, %pid : i32 loc(#loc8) + %off_81 = arith.muli %off_80, %s : i32 loc(#loc8) + %off_82 = arith.addi %off_81, %pid : i32 loc(#loc8) + %off_83 = arith.muli %off_82, %s : i32 loc(#loc8) + %off_84 = arith.addi %off_83, %pid : i32 loc(#loc8) + %off_85 = arith.muli %off_84, %s : i32 loc(#loc8) + %off_86 = arith.addi %off_85, %pid : i32 loc(#loc8) + %off_87 = arith.muli %off_86, %s : i32 loc(#loc8) + %off_88 = arith.addi %off_87, %pid : i32 loc(#loc8) + %off_89 = arith.muli %off_88, %s : i32 loc(#loc8) + %off_90 = arith.addi %off_89, %pid : i32 loc(#loc8) + %off_91 = arith.muli %off_90, %s : i32 loc(#loc8) + %off_92 = arith.addi %off_91, %pid : i32 loc(#loc8) + %off_93 = arith.muli %off_92, %s : i32 loc(#loc8) + %off_94 = arith.addi %off_93, %pid : i32 loc(#loc8) + %off_95 = arith.muli %off_94, %s : i32 loc(#loc8) + %off_96 = arith.addi %off_95, %pid : i32 loc(#loc8) + %off_97 = arith.muli %off_96, %s : i32 loc(#loc8) + %off_98 = arith.addi %off_97, %pid : i32 loc(#loc8) + %off_99 = arith.muli %off_98, %s : i32 loc(#loc8) + %off_100 = arith.addi %off_99, %pid : i32 loc(#loc8) + %off_101 = arith.muli %off_100, %s : i32 loc(#loc8) + %off_102 = arith.addi %off_101, %pid : i32 loc(#loc8) + %off_103 = arith.muli %off_102, %s : i32 loc(#loc8) + %off_104 = arith.addi %off_103, %pid : i32 loc(#loc8) + %off_105 = arith.muli %off_104, %s : i32 loc(#loc8) + %off_106 = arith.addi %off_105, %pid : i32 loc(#loc8) + %off_107 = arith.muli %off_106, %s : i32 loc(#loc8) + %off_108 = arith.addi %off_107, %pid : i32 loc(#loc8) + %off_109 = arith.muli %off_108, %s : i32 loc(#loc8) + %off_110 = arith.addi %off_109, %pid : i32 loc(#loc8) + %off_111 = arith.muli %off_110, %s : i32 loc(#loc8) + %off_112 = arith.addi %off_111, %pid : i32 loc(#loc8) + %off_113 = arith.muli %off_112, %s : i32 loc(#loc8) + %off_114 = arith.addi %off_113, %pid : i32 loc(#loc8) + %off_115 = arith.muli %off_114, %s : i32 loc(#loc8) + %off_116 = arith.addi %off_115, %pid : i32 loc(#loc8) + %off_117 = arith.muli %off_116, %s : i32 loc(#loc8) + %off_118 = arith.addi %off_117, %pid : i32 loc(#loc8) + %off_119 = arith.muli %off_118, %s : i32 loc(#loc8) + %off_120 = arith.addi %off_119, %pid : i32 loc(#loc8) + %off_121 = arith.muli %off_120, %s : i32 loc(#loc8) + %off_122 = arith.addi %off_121, %pid : i32 loc(#loc8) + %off_123 = arith.muli %off_122, %s : i32 loc(#loc8) + %off_124 = arith.addi %off_123, %pid : i32 loc(#loc8) + %off_125 = arith.muli %off_124, %s : i32 loc(#loc8) + %off_126 = arith.addi %off_125, %pid : i32 loc(#loc8) + %off_127 = arith.muli %off_126, %s : i32 loc(#loc8) + %off_128 = arith.addi %off_127, %pid : i32 loc(#loc8) + %off_129 = arith.muli %off_128, %s : i32 loc(#loc8) + %off_130 = arith.addi %off_129, %pid : i32 loc(#loc8) + %off_131 = arith.muli %off_130, %s : i32 loc(#loc8) + %off_132 = arith.addi %off_131, %pid : i32 loc(#loc8) + %off_133 = arith.muli %off_132, %s : i32 loc(#loc8) + %off_134 = arith.addi %off_133, %pid : i32 loc(#loc8) + %off_135 = arith.muli %off_134, %s : i32 loc(#loc8) + %off_136 = arith.addi %off_135, %pid : i32 loc(#loc8) + %off_137 = arith.muli %off_136, %s : i32 loc(#loc8) + %off_138 = arith.addi %off_137, %pid : i32 loc(#loc8) + %off_139 = arith.muli %off_138, %s : i32 loc(#loc8) + %off_140 = arith.addi %off_139, %pid : i32 loc(#loc8) + %off_141 = arith.muli %off_140, %s : i32 loc(#loc8) + %off_142 = arith.addi %off_141, %pid : i32 loc(#loc8) + %off_143 = arith.muli %off_142, %s : i32 loc(#loc8) + %off_144 = arith.addi %off_143, %pid : i32 loc(#loc8) + %off_145 = arith.muli %off_144, %s : i32 loc(#loc8) + %off_146 = arith.addi %off_145, %pid : i32 loc(#loc8) + %off_147 = arith.muli %off_146, %s : i32 loc(#loc8) + %off_148 = arith.addi %off_147, %pid : i32 loc(#loc8) + %off_149 = arith.muli %off_148, %s : i32 loc(#loc8) + %off_150 = arith.addi %off_149, %pid : i32 loc(#loc8) + %off_151 = arith.muli %off_150, %s : i32 loc(#loc8) + %off_152 = arith.addi %off_151, %pid : i32 loc(#loc8) + %off_153 = arith.muli %off_152, %s : i32 loc(#loc8) + %off_154 = arith.addi %off_153, %pid : i32 loc(#loc8) + %off_155 = arith.muli %off_154, %s : i32 loc(#loc8) + %off_156 = arith.addi %off_155, %pid : i32 loc(#loc8) + %off_157 = arith.muli %off_156, %s : i32 loc(#loc8) + %off_158 = arith.addi %off_157, %pid : i32 loc(#loc8) + %off_159 = arith.muli %off_158, %s : i32 loc(#loc8) + %off_160 = arith.addi %off_159, %pid : i32 loc(#loc8) + %off_161 = arith.muli %off_160, %s : i32 loc(#loc8) + %off_162 = arith.addi %off_161, %pid : i32 loc(#loc8) + %off_163 = arith.muli %off_162, %s : i32 loc(#loc8) + %off_164 = arith.addi %off_163, %pid : i32 loc(#loc8) + %off_165 = arith.muli %off_164, %s : i32 loc(#loc8) + %off_166 = arith.addi %off_165, %pid : i32 loc(#loc8) + %off_167 = arith.muli %off_166, %s : i32 loc(#loc8) + %off_168 = arith.addi %off_167, %pid : i32 loc(#loc8) + %off_169 = arith.muli %off_168, %s : i32 loc(#loc8) + %off_170 = arith.addi %off_169, %pid : i32 loc(#loc8) + %off_171 = arith.muli %off_170, %s : i32 loc(#loc8) + %off_172 = arith.addi %off_171, %pid : i32 loc(#loc8) + %off_173 = arith.muli %off_172, %s : i32 loc(#loc8) + %off_174 = arith.addi %off_173, %pid : i32 loc(#loc8) + %off_175 = arith.muli %off_174, %s : i32 loc(#loc8) + %off_176 = arith.addi %off_175, %pid : i32 loc(#loc8) + %off_177 = arith.muli %off_176, %s : i32 loc(#loc8) + %off_178 = arith.addi %off_177, %pid : i32 loc(#loc8) + %off_179 = arith.muli %off_178, %s : i32 loc(#loc8) + %off_180 = arith.addi %off_179, %pid : i32 loc(#loc8) + %off_181 = arith.muli %off_180, %s : i32 loc(#loc8) + %off_182 = arith.addi %off_181, %pid : i32 loc(#loc8) + %off_183 = arith.muli %off_182, %s : i32 loc(#loc8) + %off_184 = arith.addi %off_183, %pid : i32 loc(#loc8) + %off_185 = arith.muli %off_184, %s : i32 loc(#loc8) + %off_186 = arith.addi %off_185, %pid : i32 loc(#loc8) + %off_187 = arith.muli %off_186, %s : i32 loc(#loc8) + %off_188 = arith.addi %off_187, %pid : i32 loc(#loc8) + %off_189 = arith.muli %off_188, %s : i32 loc(#loc8) + %off_190 = arith.addi %off_189, %pid : i32 loc(#loc8) + %off_191 = arith.muli %off_190, %s : i32 loc(#loc8) + %off_192 = arith.addi %off_191, %pid : i32 loc(#loc8) + %off_193 = arith.muli %off_192, %s : i32 loc(#loc8) + %off_194 = arith.addi %off_193, %pid : i32 loc(#loc8) + %off_195 = arith.muli %off_194, %s : i32 loc(#loc8) + %off_196 = arith.addi %off_195, %pid : i32 loc(#loc8) + %off_197 = arith.muli %off_196, %s : i32 loc(#loc8) + %off_198 = arith.addi %off_197, %pid : i32 loc(#loc8) + %off_199 = arith.muli %off_198, %s : i32 loc(#loc8) + %off_200 = arith.addi %off_199, %pid : i32 loc(#loc8) + %off_201 = arith.muli %off_200, %s : i32 loc(#loc8) + %off_202 = arith.addi %off_201, %pid : i32 loc(#loc8) + %off_203 = arith.muli %off_202, %s : i32 loc(#loc8) + %off_204 = arith.addi %off_203, %pid : i32 loc(#loc8) + %off_205 = arith.muli %off_204, %s : i32 loc(#loc8) + %off_206 = arith.addi %off_205, %pid : i32 loc(#loc8) + %off_207 = arith.muli %off_206, %s : i32 loc(#loc8) + %off_208 = arith.addi %off_207, %pid : i32 loc(#loc8) + %off_209 = arith.muli %off_208, %s : i32 loc(#loc8) + %off_210 = arith.addi %off_209, %pid : i32 loc(#loc8) + %off_211 = arith.muli %off_210, %s : i32 loc(#loc8) + %off_212 = arith.addi %off_211, %pid : i32 loc(#loc8) + %off_213 = arith.muli %off_212, %s : i32 loc(#loc8) + %off_214 = arith.addi %off_213, %pid : i32 loc(#loc8) + %off_215 = arith.muli %off_214, %s : i32 loc(#loc8) + %off_216 = arith.addi %off_215, %pid : i32 loc(#loc8) + %off_217 = arith.muli %off_216, %s : i32 loc(#loc8) + %off_218 = arith.addi %off_217, %pid : i32 loc(#loc8) + %off_219 = arith.muli %off_218, %s : i32 loc(#loc8) + %off_220 = arith.addi %off_219, %pid : i32 loc(#loc8) + %off_221 = arith.muli %off_220, %s : i32 loc(#loc8) + %off_222 = arith.addi %off_221, %pid : i32 loc(#loc8) + %off_223 = arith.muli %off_222, %s : i32 loc(#loc8) + %off_224 = arith.addi %off_223, %pid : i32 loc(#loc8) + %off_225 = arith.muli %off_224, %s : i32 loc(#loc8) + %off_226 = arith.addi %off_225, %pid : i32 loc(#loc8) + %off_227 = arith.muli %off_226, %s : i32 loc(#loc8) + %off_228 = arith.addi %off_227, %pid : i32 loc(#loc8) + %off_229 = arith.muli %off_228, %s : i32 loc(#loc8) + %off_230 = arith.addi %off_229, %pid : i32 loc(#loc8) + %off_231 = arith.muli %off_230, %s : i32 loc(#loc8) + %off_232 = arith.addi %off_231, %pid : i32 loc(#loc8) + %off_233 = arith.muli %off_232, %s : i32 loc(#loc8) + %off_234 = arith.addi %off_233, %pid : i32 loc(#loc8) + %off_235 = arith.muli %off_234, %s : i32 loc(#loc8) + %off_236 = arith.addi %off_235, %pid : i32 loc(#loc8) + %off_237 = arith.muli %off_236, %s : i32 loc(#loc8) + %off_238 = arith.addi %off_237, %pid : i32 loc(#loc8) + %off_239 = arith.muli %off_238, %s : i32 loc(#loc8) + %off_240 = arith.addi %off_239, %pid : i32 loc(#loc8) + %off_241 = arith.muli %off_240, %s : i32 loc(#loc8) + %off_242 = arith.addi %off_241, %pid : i32 loc(#loc8) + %off_243 = arith.muli %off_242, %s : i32 loc(#loc8) + %off_244 = arith.addi %off_243, %pid : i32 loc(#loc8) + %off_245 = arith.muli %off_244, %s : i32 loc(#loc8) + %off_246 = arith.addi %off_245, %pid : i32 loc(#loc8) + %off_247 = arith.muli %off_246, %s : i32 loc(#loc8) + %off_248 = arith.addi %off_247, %pid : i32 loc(#loc8) + %off_249 = arith.muli %off_248, %s : i32 loc(#loc8) + %off_250 = arith.addi %off_249, %pid : i32 loc(#loc8) + %off_251 = arith.muli %off_250, %s : i32 loc(#loc8) + %off_252 = arith.addi %off_251, %pid : i32 loc(#loc8) + %off_253 = arith.muli %off_252, %s : i32 loc(#loc8) + %off_254 = arith.addi %off_253, %pid : i32 loc(#loc8) + %off_255 = arith.muli %off_254, %s : i32 loc(#loc8) + %off_256 = arith.addi %off_255, %pid : i32 loc(#loc8) + %off_257 = arith.muli %off_256, %s : i32 loc(#loc8) + %off_258 = arith.addi %off_257, %pid : i32 loc(#loc8) + %off_259 = arith.muli %off_258, %s : i32 loc(#loc8) + %off_260 = arith.addi %off_259, %pid : i32 loc(#loc8) + %off_261 = arith.muli %off_260, %s : i32 loc(#loc8) + %off_262 = arith.addi %off_261, %pid : i32 loc(#loc8) + %off_263 = arith.muli %off_262, %s : i32 loc(#loc8) + %off_264 = arith.addi %off_263, %pid : i32 loc(#loc8) + %off_265 = arith.muli %off_264, %s : i32 loc(#loc8) + %off_266 = arith.addi %off_265, %pid : i32 loc(#loc8) + %off_267 = arith.muli %off_266, %s : i32 loc(#loc8) + %off_268 = arith.addi %off_267, %pid : i32 loc(#loc8) + %off_269 = arith.muli %off_268, %s : i32 loc(#loc8) + %off_270 = arith.addi %off_269, %pid : i32 loc(#loc8) + %off_271 = arith.muli %off_270, %s : i32 loc(#loc8) + %off_272 = arith.addi %off_271, %pid : i32 loc(#loc8) + %off_273 = arith.muli %off_272, %s : i32 loc(#loc8) + %off_274 = arith.addi %off_273, %pid : i32 loc(#loc8) + %off_275 = arith.muli %off_274, %s : i32 loc(#loc8) + %off_276 = arith.addi %off_275, %pid : i32 loc(#loc8) + %off_277 = arith.muli %off_276, %s : i32 loc(#loc8) + %off_278 = arith.addi %off_277, %pid : i32 loc(#loc8) + %off_279 = arith.muli %off_278, %s : i32 loc(#loc8) + %off_280 = arith.addi %off_279, %pid : i32 loc(#loc8) + %off_281 = arith.muli %off_280, %s : i32 loc(#loc8) + %off_282 = arith.addi %off_281, %pid : i32 loc(#loc8) + %off_283 = arith.muli %off_282, %s : i32 loc(#loc8) + %off_284 = arith.addi %off_283, %pid : i32 loc(#loc8) + %off_285 = arith.muli %off_284, %s : i32 loc(#loc8) + %off_286 = arith.addi %off_285, %pid : i32 loc(#loc8) + %off_287 = arith.muli %off_286, %s : i32 loc(#loc8) + %off_288 = arith.addi %off_287, %pid : i32 loc(#loc8) + %off_289 = arith.muli %off_288, %s : i32 loc(#loc8) + %off_290 = arith.addi %off_289, %pid : i32 loc(#loc8) + %off_291 = arith.muli %off_290, %s : i32 loc(#loc8) + %off_292 = arith.addi %off_291, %pid : i32 loc(#loc8) + %off_293 = arith.muli %off_292, %s : i32 loc(#loc8) + %off_294 = arith.addi %off_293, %pid : i32 loc(#loc8) + %off_295 = arith.muli %off_294, %s : i32 loc(#loc8) + %off_296 = arith.addi %off_295, %pid : i32 loc(#loc8) + %off_297 = arith.muli %off_296, %s : i32 loc(#loc8) + %off_298 = arith.addi %off_297, %pid : i32 loc(#loc8) + %off_299 = arith.muli %off_298, %s : i32 loc(#loc8) + %off_300 = arith.addi %off_299, %pid : i32 loc(#loc8) + %off_301 = arith.muli %off_300, %s : i32 loc(#loc8) + %off_302 = arith.addi %off_301, %pid : i32 loc(#loc8) + %off_303 = arith.muli %off_302, %s : i32 loc(#loc8) + %off_304 = arith.addi %off_303, %pid : i32 loc(#loc8) + %off_305 = arith.muli %off_304, %s : i32 loc(#loc8) + %off_306 = arith.addi %off_305, %pid : i32 loc(#loc8) + %off_307 = arith.muli %off_306, %s : i32 loc(#loc8) + %off_308 = arith.addi %off_307, %pid : i32 loc(#loc8) + %off_309 = arith.muli %off_308, %s : i32 loc(#loc8) + %off_310 = arith.addi %off_309, %pid : i32 loc(#loc8) + %off_311 = arith.muli %off_310, %s : i32 loc(#loc8) + %off_312 = arith.addi %off_311, %pid : i32 loc(#loc8) + %off_313 = arith.muli %off_312, %s : i32 loc(#loc8) + %off_314 = arith.addi %off_313, %pid : i32 loc(#loc8) + %off_315 = arith.muli %off_314, %s : i32 loc(#loc8) + %off_316 = arith.addi %off_315, %pid : i32 loc(#loc8) + %off_317 = arith.muli %off_316, %s : i32 loc(#loc8) + %off_318 = arith.addi %off_317, %pid : i32 loc(#loc8) + %off_319 = arith.muli %off_318, %s : i32 loc(#loc8) + %off_320 = arith.addi %off_319, %pid : i32 loc(#loc8) + %off_321 = arith.muli %off_320, %s : i32 loc(#loc8) + %off_322 = arith.addi %off_321, %pid : i32 loc(#loc8) + %off_323 = arith.muli %off_322, %s : i32 loc(#loc8) + %off_324 = arith.addi %off_323, %pid : i32 loc(#loc8) + %off_325 = arith.muli %off_324, %s : i32 loc(#loc8) + %off_326 = arith.addi %off_325, %pid : i32 loc(#loc8) + %off_327 = arith.muli %off_326, %s : i32 loc(#loc8) + %off_328 = arith.addi %off_327, %pid : i32 loc(#loc8) + %off_329 = arith.muli %off_328, %s : i32 loc(#loc8) + %off_330 = arith.addi %off_329, %pid : i32 loc(#loc8) + %off_331 = arith.muli %off_330, %s : i32 loc(#loc8) + %off_332 = arith.addi %off_331, %pid : i32 loc(#loc8) + %off_333 = arith.muli %off_332, %s : i32 loc(#loc8) + %off_334 = arith.addi %off_333, %pid : i32 loc(#loc8) + %off_335 = arith.muli %off_334, %s : i32 loc(#loc8) + %off_336 = arith.addi %off_335, %pid : i32 loc(#loc8) + %off_337 = arith.muli %off_336, %s : i32 loc(#loc8) + %off_338 = arith.addi %off_337, %pid : i32 loc(#loc8) + %off_339 = arith.muli %off_338, %s : i32 loc(#loc8) + %off_340 = arith.addi %off_339, %pid : i32 loc(#loc8) + %off_341 = arith.muli %off_340, %s : i32 loc(#loc8) + %off_342 = arith.addi %off_341, %pid : i32 loc(#loc8) + %off_343 = arith.muli %off_342, %s : i32 loc(#loc8) + %off_344 = arith.addi %off_343, %pid : i32 loc(#loc8) + %off_345 = arith.muli %off_344, %s : i32 loc(#loc8) + %off_346 = arith.addi %off_345, %pid : i32 loc(#loc8) + %off_347 = arith.muli %off_346, %s : i32 loc(#loc8) + %off_348 = arith.addi %off_347, %pid : i32 loc(#loc8) + %off_349 = arith.muli %off_348, %s : i32 loc(#loc8) + %off_350 = arith.addi %off_349, %pid : i32 loc(#loc8) + %off_351 = arith.muli %off_350, %s : i32 loc(#loc8) + %off_352 = arith.addi %off_351, %pid : i32 loc(#loc8) + %off_353 = arith.muli %off_352, %s : i32 loc(#loc8) + %off_354 = arith.addi %off_353, %pid : i32 loc(#loc8) + %off_355 = arith.muli %off_354, %s : i32 loc(#loc8) + %off_356 = arith.addi %off_355, %pid : i32 loc(#loc8) + %off_357 = arith.muli %off_356, %s : i32 loc(#loc8) + %off_358 = arith.addi %off_357, %pid : i32 loc(#loc8) + %off_359 = arith.muli %off_358, %s : i32 loc(#loc8) + %off_360 = arith.addi %off_359, %pid : i32 loc(#loc8) + %off_361 = arith.muli %off_360, %s : i32 loc(#loc8) + %off_362 = arith.addi %off_361, %pid : i32 loc(#loc8) + %off_363 = arith.muli %off_362, %s : i32 loc(#loc8) + %off_364 = arith.addi %off_363, %pid : i32 loc(#loc8) + %off_365 = arith.muli %off_364, %s : i32 loc(#loc8) + %off_366 = arith.addi %off_365, %pid : i32 loc(#loc8) + %off_367 = arith.muli %off_366, %s : i32 loc(#loc8) + %off_368 = arith.addi %off_367, %pid : i32 loc(#loc8) + %off_369 = arith.muli %off_368, %s : i32 loc(#loc8) + %off_370 = arith.addi %off_369, %pid : i32 loc(#loc8) + %off_371 = arith.muli %off_370, %s : i32 loc(#loc8) + %off_372 = arith.addi %off_371, %pid : i32 loc(#loc8) + %off_373 = arith.muli %off_372, %s : i32 loc(#loc8) + %off_374 = arith.addi %off_373, %pid : i32 loc(#loc8) + %off_375 = arith.muli %off_374, %s : i32 loc(#loc8) + %off_376 = arith.addi %off_375, %pid : i32 loc(#loc8) + %off_377 = arith.muli %off_376, %s : i32 loc(#loc8) + %off_378 = arith.addi %off_377, %pid : i32 loc(#loc8) + %off_379 = arith.muli %off_378, %s : i32 loc(#loc8) + %off_380 = arith.addi %off_379, %pid : i32 loc(#loc8) + %off_381 = arith.muli %off_380, %s : i32 loc(#loc8) + %off_382 = arith.addi %off_381, %pid : i32 loc(#loc8) + %off_383 = arith.muli %off_382, %s : i32 loc(#loc8) + %off_384 = arith.addi %off_383, %pid : i32 loc(#loc8) + %off_385 = arith.muli %off_384, %s : i32 loc(#loc8) + %off_386 = arith.addi %off_385, %pid : i32 loc(#loc8) + %off_387 = arith.muli %off_386, %s : i32 loc(#loc8) + %off_388 = arith.addi %off_387, %pid : i32 loc(#loc8) + %off_389 = arith.muli %off_388, %s : i32 loc(#loc8) + %off_390 = arith.addi %off_389, %pid : i32 loc(#loc8) + %off_391 = arith.muli %off_390, %s : i32 loc(#loc8) + %off_392 = arith.addi %off_391, %pid : i32 loc(#loc8) + %off_393 = arith.muli %off_392, %s : i32 loc(#loc8) + %off_394 = arith.addi %off_393, %pid : i32 loc(#loc8) + %off_395 = arith.muli %off_394, %s : i32 loc(#loc8) + %off_396 = arith.addi %off_395, %pid : i32 loc(#loc8) + %off_397 = arith.muli %off_396, %s : i32 loc(#loc8) + %off_398 = arith.addi %off_397, %pid : i32 loc(#loc8) + %off_399 = arith.muli %off_398, %s : i32 loc(#loc8) + %off_400 = arith.addi %off_399, %pid : i32 loc(#loc8) + %off_401 = arith.muli %off_400, %s : i32 loc(#loc8) + %off_402 = arith.addi %off_401, %pid : i32 loc(#loc8) + %off_403 = arith.muli %off_402, %s : i32 loc(#loc8) + %off_404 = arith.addi %off_403, %pid : i32 loc(#loc8) + %off_405 = arith.muli %off_404, %s : i32 loc(#loc8) + %off_406 = arith.addi %off_405, %pid : i32 loc(#loc8) + %off_407 = arith.muli %off_406, %s : i32 loc(#loc8) + %off_408 = arith.addi %off_407, %pid : i32 loc(#loc8) + %off_409 = arith.muli %off_408, %s : i32 loc(#loc8) + %off_410 = arith.addi %off_409, %pid : i32 loc(#loc8) + %off_411 = arith.muli %off_410, %s : i32 loc(#loc8) + %off_412 = arith.addi %off_411, %pid : i32 loc(#loc8) + %off_413 = arith.muli %off_412, %s : i32 loc(#loc8) + %off_414 = arith.addi %off_413, %pid : i32 loc(#loc8) + %off_415 = arith.muli %off_414, %s : i32 loc(#loc8) + %off_416 = arith.addi %off_415, %pid : i32 loc(#loc8) + %off_417 = arith.muli %off_416, %s : i32 loc(#loc8) + %off_418 = arith.addi %off_417, %pid : i32 loc(#loc8) + %off_419 = arith.muli %off_418, %s : i32 loc(#loc8) + %off_420 = arith.addi %off_419, %pid : i32 loc(#loc8) + %off_421 = arith.muli %off_420, %s : i32 loc(#loc8) + %off_422 = arith.addi %off_421, %pid : i32 loc(#loc8) + %off_423 = arith.muli %off_422, %s : i32 loc(#loc8) + %off_424 = arith.addi %off_423, %pid : i32 loc(#loc8) + %off_425 = arith.muli %off_424, %s : i32 loc(#loc8) + %off_426 = arith.addi %off_425, %pid : i32 loc(#loc8) + %off_427 = arith.muli %off_426, %s : i32 loc(#loc8) + %off_428 = arith.addi %off_427, %pid : i32 loc(#loc8) + %off_429 = arith.muli %off_428, %s : i32 loc(#loc8) + %off_430 = arith.addi %off_429, %pid : i32 loc(#loc8) + %off_431 = arith.muli %off_430, %s : i32 loc(#loc8) + %off_432 = arith.addi %off_431, %pid : i32 loc(#loc8) + %off_433 = arith.muli %off_432, %s : i32 loc(#loc8) + %off_434 = arith.addi %off_433, %pid : i32 loc(#loc8) + %off_435 = arith.muli %off_434, %s : i32 loc(#loc8) + %off_436 = arith.addi %off_435, %pid : i32 loc(#loc8) + %off_437 = arith.muli %off_436, %s : i32 loc(#loc8) + %off_438 = arith.addi %off_437, %pid : i32 loc(#loc8) + %off_439 = arith.muli %off_438, %s : i32 loc(#loc8) + %off_440 = arith.addi %off_439, %pid : i32 loc(#loc8) + %off_441 = arith.muli %off_440, %s : i32 loc(#loc8) + %off_442 = arith.addi %off_441, %pid : i32 loc(#loc8) + %off_443 = arith.muli %off_442, %s : i32 loc(#loc8) + %off_444 = arith.addi %off_443, %pid : i32 loc(#loc8) + %off_445 = arith.muli %off_444, %s : i32 loc(#loc8) + %off_446 = arith.addi %off_445, %pid : i32 loc(#loc8) + %off_447 = arith.muli %off_446, %s : i32 loc(#loc8) + %off_448 = arith.addi %off_447, %pid : i32 loc(#loc8) + %off_449 = arith.muli %off_448, %s : i32 loc(#loc8) + %off_450 = arith.addi %off_449, %pid : i32 loc(#loc8) + %off_451 = arith.muli %off_450, %s : i32 loc(#loc8) + %off_452 = arith.addi %off_451, %pid : i32 loc(#loc8) + %off_453 = arith.muli %off_452, %s : i32 loc(#loc8) + %off_454 = arith.addi %off_453, %pid : i32 loc(#loc8) + %off_455 = arith.muli %off_454, %s : i32 loc(#loc8) + %off_456 = arith.addi %off_455, %pid : i32 loc(#loc8) + %off_457 = arith.muli %off_456, %s : i32 loc(#loc8) + %off_458 = arith.addi %off_457, %pid : i32 loc(#loc8) + %off_459 = arith.muli %off_458, %s : i32 loc(#loc8) + %off_460 = arith.addi %off_459, %pid : i32 loc(#loc8) + %off_461 = arith.muli %off_460, %s : i32 loc(#loc8) + %off_462 = arith.addi %off_461, %pid : i32 loc(#loc8) + %off_463 = arith.muli %off_462, %s : i32 loc(#loc8) + %off_464 = arith.addi %off_463, %pid : i32 loc(#loc8) + %off_465 = arith.muli %off_464, %s : i32 loc(#loc8) + %off_466 = arith.addi %off_465, %pid : i32 loc(#loc8) + %off_467 = arith.muli %off_466, %s : i32 loc(#loc8) + %off_468 = arith.addi %off_467, %pid : i32 loc(#loc8) + %off_469 = arith.muli %off_468, %s : i32 loc(#loc8) + %off_470 = arith.addi %off_469, %pid : i32 loc(#loc8) + %off_471 = arith.muli %off_470, %s : i32 loc(#loc8) + %off_472 = arith.addi %off_471, %pid : i32 loc(#loc8) + %off_473 = arith.muli %off_472, %s : i32 loc(#loc8) + %off_474 = arith.addi %off_473, %pid : i32 loc(#loc8) + %off_475 = arith.muli %off_474, %s : i32 loc(#loc8) + %off_476 = arith.addi %off_475, %pid : i32 loc(#loc8) + %off_477 = arith.muli %off_476, %s : i32 loc(#loc8) + %off_478 = arith.addi %off_477, %pid : i32 loc(#loc8) + %off_479 = arith.muli %off_478, %s : i32 loc(#loc8) + %off_480 = arith.addi %off_479, %pid : i32 loc(#loc8) + %off_481 = arith.muli %off_480, %s : i32 loc(#loc8) + %off_482 = arith.addi %off_481, %pid : i32 loc(#loc8) + %off_483 = arith.muli %off_482, %s : i32 loc(#loc8) + %off_484 = arith.addi %off_483, %pid : i32 loc(#loc8) + %off_485 = arith.muli %off_484, %s : i32 loc(#loc8) + %off_486 = arith.addi %off_485, %pid : i32 loc(#loc8) + %off_487 = arith.muli %off_486, %s : i32 loc(#loc8) + %off_488 = arith.addi %off_487, %pid : i32 loc(#loc8) + %off_489 = arith.muli %off_488, %s : i32 loc(#loc8) + %off_490 = arith.addi %off_489, %pid : i32 loc(#loc8) + %off_491 = arith.muli %off_490, %s : i32 loc(#loc8) + %off_492 = arith.addi %off_491, %pid : i32 loc(#loc8) + %off_493 = arith.muli %off_492, %s : i32 loc(#loc8) + %off_494 = arith.addi %off_493, %pid : i32 loc(#loc8) + %off_495 = arith.muli %off_494, %s : i32 loc(#loc8) + %off_496 = arith.addi %off_495, %pid : i32 loc(#loc8) + %off_497 = arith.muli %off_496, %s : i32 loc(#loc8) + %off_498 = arith.addi %off_497, %pid : i32 loc(#loc8) + %off_499 = arith.muli %off_498, %s : i32 loc(#loc8) + %off_500 = arith.addi %off_499, %pid : i32 loc(#loc8) + %off_501 = arith.muli %off_500, %s : i32 loc(#loc8) + %off_502 = arith.addi %off_501, %pid : i32 loc(#loc8) + %off_503 = arith.muli %off_502, %s : i32 loc(#loc8) + %off_504 = arith.addi %off_503, %pid : i32 loc(#loc8) + %off_505 = arith.muli %off_504, %s : i32 loc(#loc8) + %off_506 = arith.addi %off_505, %pid : i32 loc(#loc8) + %off_507 = arith.muli %off_506, %s : i32 loc(#loc8) + %off_508 = arith.addi %off_507, %pid : i32 loc(#loc8) + %off_509 = arith.muli %off_508, %s : i32 loc(#loc8) + %off_510 = arith.addi %off_509, %pid : i32 loc(#loc8) + %off_511 = arith.muli %off_510, %s : i32 loc(#loc8) + %off_512 = arith.addi %off_511, %pid : i32 loc(#loc8) + %off_513 = arith.muli %off_512, %s : i32 loc(#loc8) + %off_514 = arith.addi %off_513, %pid : i32 loc(#loc8) + %off_515 = arith.muli %off_514, %s : i32 loc(#loc8) + %off_516 = arith.addi %off_515, %pid : i32 loc(#loc8) + %off_517 = arith.muli %off_516, %s : i32 loc(#loc8) + %off_518 = arith.addi %off_517, %pid : i32 loc(#loc8) + %off_519 = arith.muli %off_518, %s : i32 loc(#loc8) + %off_520 = arith.addi %off_519, %pid : i32 loc(#loc8) + %off_521 = arith.muli %off_520, %s : i32 loc(#loc8) + %off_522 = arith.addi %off_521, %pid : i32 loc(#loc8) + %off_523 = arith.muli %off_522, %s : i32 loc(#loc8) + %off_524 = arith.addi %off_523, %pid : i32 loc(#loc8) + %off_525 = arith.muli %off_524, %s : i32 loc(#loc8) + %off_526 = arith.addi %off_525, %pid : i32 loc(#loc8) + %off_527 = arith.muli %off_526, %s : i32 loc(#loc8) + %off_528 = arith.addi %off_527, %pid : i32 loc(#loc8) + %off_529 = arith.muli %off_528, %s : i32 loc(#loc8) + %off_530 = arith.addi %off_529, %pid : i32 loc(#loc8) + %off_531 = arith.muli %off_530, %s : i32 loc(#loc8) + %off_532 = arith.addi %off_531, %pid : i32 loc(#loc8) + %off_533 = arith.muli %off_532, %s : i32 loc(#loc8) + %off_534 = arith.addi %off_533, %pid : i32 loc(#loc8) + %off_535 = arith.muli %off_534, %s : i32 loc(#loc8) + %off_536 = arith.addi %off_535, %pid : i32 loc(#loc8) + %off_537 = arith.muli %off_536, %s : i32 loc(#loc8) + %off_538 = arith.addi %off_537, %pid : i32 loc(#loc8) + %off_539 = arith.muli %off_538, %s : i32 loc(#loc8) + %off_540 = arith.addi %off_539, %pid : i32 loc(#loc8) + %off_541 = arith.muli %off_540, %s : i32 loc(#loc8) + %off_542 = arith.addi %off_541, %pid : i32 loc(#loc8) + %off_543 = arith.muli %off_542, %s : i32 loc(#loc8) + %off_544 = arith.addi %off_543, %pid : i32 loc(#loc8) + %off_545 = arith.muli %off_544, %s : i32 loc(#loc8) + %off_546 = arith.addi %off_545, %pid : i32 loc(#loc8) + %off_547 = arith.muli %off_546, %s : i32 loc(#loc8) + %off_548 = arith.addi %off_547, %pid : i32 loc(#loc8) + %off_549 = arith.muli %off_548, %s : i32 loc(#loc8) + %off_550 = arith.addi %off_549, %pid : i32 loc(#loc8) + %off_551 = arith.muli %off_550, %s : i32 loc(#loc8) + %off_552 = arith.addi %off_551, %pid : i32 loc(#loc8) + %off_553 = arith.muli %off_552, %s : i32 loc(#loc8) + %off_554 = arith.addi %off_553, %pid : i32 loc(#loc8) + %off_555 = arith.muli %off_554, %s : i32 loc(#loc8) + %off_556 = arith.addi %off_555, %pid : i32 loc(#loc8) + %off_557 = arith.muli %off_556, %s : i32 loc(#loc8) + %off_558 = arith.addi %off_557, %pid : i32 loc(#loc8) + %off_559 = arith.muli %off_558, %s : i32 loc(#loc8) + %off_560 = arith.addi %off_559, %pid : i32 loc(#loc8) + %off_561 = arith.muli %off_560, %s : i32 loc(#loc8) + %off_562 = arith.addi %off_561, %pid : i32 loc(#loc8) + %off_563 = arith.muli %off_562, %s : i32 loc(#loc8) + %off_564 = arith.addi %off_563, %pid : i32 loc(#loc8) + %off_565 = arith.muli %off_564, %s : i32 loc(#loc8) + %off_566 = arith.addi %off_565, %pid : i32 loc(#loc8) + %off_567 = arith.muli %off_566, %s : i32 loc(#loc8) + %off_568 = arith.addi %off_567, %pid : i32 loc(#loc8) + %off_569 = arith.muli %off_568, %s : i32 loc(#loc8) + %off_570 = arith.addi %off_569, %pid : i32 loc(#loc8) + %off_571 = arith.muli %off_570, %s : i32 loc(#loc8) + %off_572 = arith.addi %off_571, %pid : i32 loc(#loc8) + %off_573 = arith.muli %off_572, %s : i32 loc(#loc8) + %off_574 = arith.addi %off_573, %pid : i32 loc(#loc8) + %off_575 = arith.muli %off_574, %s : i32 loc(#loc8) + %off_576 = arith.addi %off_575, %pid : i32 loc(#loc8) + %off_577 = arith.muli %off_576, %s : i32 loc(#loc8) + %off_578 = arith.addi %off_577, %pid : i32 loc(#loc8) + %off_579 = arith.muli %off_578, %s : i32 loc(#loc8) + %off_580 = arith.addi %off_579, %pid : i32 loc(#loc8) + %off_581 = arith.muli %off_580, %s : i32 loc(#loc8) + %off_582 = arith.addi %off_581, %pid : i32 loc(#loc8) + %off_583 = arith.muli %off_582, %s : i32 loc(#loc8) + %off_584 = arith.addi %off_583, %pid : i32 loc(#loc8) + %off_585 = arith.muli %off_584, %s : i32 loc(#loc8) + %off_586 = arith.addi %off_585, %pid : i32 loc(#loc8) + %off_587 = arith.muli %off_586, %s : i32 loc(#loc8) + %off_588 = arith.addi %off_587, %pid : i32 loc(#loc8) + %off_589 = arith.muli %off_588, %s : i32 loc(#loc8) + %off_590 = arith.addi %off_589, %pid : i32 loc(#loc8) + %off_591 = arith.muli %off_590, %s : i32 loc(#loc8) + %off_592 = arith.addi %off_591, %pid : i32 loc(#loc8) + %off_593 = arith.muli %off_592, %s : i32 loc(#loc8) + %off_594 = arith.addi %off_593, %pid : i32 loc(#loc8) + %off_595 = arith.muli %off_594, %s : i32 loc(#loc8) + %off_596 = arith.addi %off_595, %pid : i32 loc(#loc8) + %off_597 = arith.muli %off_596, %s : i32 loc(#loc8) + %off_598 = arith.addi %off_597, %pid : i32 loc(#loc8) + %off_599 = arith.muli %off_598, %s : i32 loc(#loc8) + %off_600 = arith.addi %off_599, %pid : i32 loc(#loc8) + %off_601 = arith.muli %off_600, %s : i32 loc(#loc8) + %off_602 = arith.addi %off_601, %pid : i32 loc(#loc8) + %off_603 = arith.muli %off_602, %s : i32 loc(#loc8) + %off_604 = arith.addi %off_603, %pid : i32 loc(#loc8) + %off_605 = arith.muli %off_604, %s : i32 loc(#loc8) + %off_606 = arith.addi %off_605, %pid : i32 loc(#loc8) + %off_607 = arith.muli %off_606, %s : i32 loc(#loc8) + %off_608 = arith.addi %off_607, %pid : i32 loc(#loc8) + %off_609 = arith.muli %off_608, %s : i32 loc(#loc8) + %off_610 = arith.addi %off_609, %pid : i32 loc(#loc8) + %off_611 = arith.muli %off_610, %s : i32 loc(#loc8) + %off_612 = arith.addi %off_611, %pid : i32 loc(#loc8) + %off_613 = arith.muli %off_612, %s : i32 loc(#loc8) + %off_614 = arith.addi %off_613, %pid : i32 loc(#loc8) + %off_615 = arith.muli %off_614, %s : i32 loc(#loc8) + %off_616 = arith.addi %off_615, %pid : i32 loc(#loc8) + %off_617 = arith.muli %off_616, %s : i32 loc(#loc8) + %off_618 = arith.addi %off_617, %pid : i32 loc(#loc8) + %off_619 = arith.muli %off_618, %s : i32 loc(#loc8) + %off_620 = arith.addi %off_619, %pid : i32 loc(#loc8) + %off_621 = arith.muli %off_620, %s : i32 loc(#loc8) + %off_622 = arith.addi %off_621, %pid : i32 loc(#loc8) + %off_623 = arith.muli %off_622, %s : i32 loc(#loc8) + %off_624 = arith.addi %off_623, %pid : i32 loc(#loc8) + %off_625 = arith.muli %off_624, %s : i32 loc(#loc8) + %off_626 = arith.addi %off_625, %pid : i32 loc(#loc8) + %off_627 = arith.muli %off_626, %s : i32 loc(#loc8) + %off_628 = arith.addi %off_627, %pid : i32 loc(#loc8) + %off_629 = arith.muli %off_628, %s : i32 loc(#loc8) + %off_630 = arith.addi %off_629, %pid : i32 loc(#loc8) + %off_631 = arith.muli %off_630, %s : i32 loc(#loc8) + %off_632 = arith.addi %off_631, %pid : i32 loc(#loc8) + %off_633 = arith.muli %off_632, %s : i32 loc(#loc8) + %off_634 = arith.addi %off_633, %pid : i32 loc(#loc8) + %off_635 = arith.muli %off_634, %s : i32 loc(#loc8) + %off_636 = arith.addi %off_635, %pid : i32 loc(#loc8) + %off_637 = arith.muli %off_636, %s : i32 loc(#loc8) + %off_638 = arith.addi %off_637, %pid : i32 loc(#loc8) + %off_639 = arith.muli %off_638, %s : i32 loc(#loc8) + %off_640 = arith.addi %off_639, %pid : i32 loc(#loc8) + %off_641 = arith.muli %off_640, %s : i32 loc(#loc8) + %off_642 = arith.addi %off_641, %pid : i32 loc(#loc8) + %off_643 = arith.muli %off_642, %s : i32 loc(#loc8) + %off_644 = arith.addi %off_643, %pid : i32 loc(#loc8) + %off_645 = arith.muli %off_644, %s : i32 loc(#loc8) + %off_646 = arith.addi %off_645, %pid : i32 loc(#loc8) + %off_647 = arith.muli %off_646, %s : i32 loc(#loc8) + %off_648 = arith.addi %off_647, %pid : i32 loc(#loc8) + %off_649 = arith.muli %off_648, %s : i32 loc(#loc8) + %off_650 = arith.addi %off_649, %pid : i32 loc(#loc8) + %off_651 = arith.muli %off_650, %s : i32 loc(#loc8) + %off_652 = arith.addi %off_651, %pid : i32 loc(#loc8) + %off_653 = arith.muli %off_652, %s : i32 loc(#loc8) + %off_654 = arith.addi %off_653, %pid : i32 loc(#loc8) + %off_655 = arith.muli %off_654, %s : i32 loc(#loc8) + %off_656 = arith.addi %off_655, %pid : i32 loc(#loc8) + %off_657 = arith.muli %off_656, %s : i32 loc(#loc8) + %off_658 = arith.addi %off_657, %pid : i32 loc(#loc8) + %off_659 = arith.muli %off_658, %s : i32 loc(#loc8) + %off_660 = arith.addi %off_659, %pid : i32 loc(#loc8) + %off_661 = arith.muli %off_660, %s : i32 loc(#loc8) + %off_662 = arith.addi %off_661, %pid : i32 loc(#loc8) + %off_663 = arith.muli %off_662, %s : i32 loc(#loc8) + %off_664 = arith.addi %off_663, %pid : i32 loc(#loc8) + %off_665 = arith.muli %off_664, %s : i32 loc(#loc8) + %off_666 = arith.addi %off_665, %pid : i32 loc(#loc8) + %off_667 = arith.muli %off_666, %s : i32 loc(#loc8) + %off_668 = arith.addi %off_667, %pid : i32 loc(#loc8) + %off_669 = arith.muli %off_668, %s : i32 loc(#loc8) + %off_670 = arith.addi %off_669, %pid : i32 loc(#loc8) + %off_671 = arith.muli %off_670, %s : i32 loc(#loc8) + %off_672 = arith.addi %off_671, %pid : i32 loc(#loc8) + %off_673 = arith.muli %off_672, %s : i32 loc(#loc8) + %off_674 = arith.addi %off_673, %pid : i32 loc(#loc8) + %off_675 = arith.muli %off_674, %s : i32 loc(#loc8) + %off_676 = arith.addi %off_675, %pid : i32 loc(#loc8) + %off_677 = arith.muli %off_676, %s : i32 loc(#loc8) + %off_678 = arith.addi %off_677, %pid : i32 loc(#loc8) + %off_679 = arith.muli %off_678, %s : i32 loc(#loc8) + %off_680 = arith.addi %off_679, %pid : i32 loc(#loc8) + %off_681 = arith.muli %off_680, %s : i32 loc(#loc8) + %off_682 = arith.addi %off_681, %pid : i32 loc(#loc8) + %off_683 = arith.muli %off_682, %s : i32 loc(#loc8) + %off_684 = arith.addi %off_683, %pid : i32 loc(#loc8) + %off_685 = arith.muli %off_684, %s : i32 loc(#loc8) + %off_686 = arith.addi %off_685, %pid : i32 loc(#loc8) + %off_687 = arith.muli %off_686, %s : i32 loc(#loc8) + %off_688 = arith.addi %off_687, %pid : i32 loc(#loc8) + %off_689 = arith.muli %off_688, %s : i32 loc(#loc8) + %off_690 = arith.addi %off_689, %pid : i32 loc(#loc8) + %off_691 = arith.muli %off_690, %s : i32 loc(#loc8) + %off_692 = arith.addi %off_691, %pid : i32 loc(#loc8) + %off_693 = arith.muli %off_692, %s : i32 loc(#loc8) + %off_694 = arith.addi %off_693, %pid : i32 loc(#loc8) + %off_695 = arith.muli %off_694, %s : i32 loc(#loc8) + %off_696 = arith.addi %off_695, %pid : i32 loc(#loc8) + %off_697 = arith.muli %off_696, %s : i32 loc(#loc8) + %off_698 = arith.addi %off_697, %pid : i32 loc(#loc8) + %off_699 = arith.muli %off_698, %s : i32 loc(#loc8) + %off_700 = arith.addi %off_699, %pid : i32 loc(#loc8) + %off_701 = arith.muli %off_700, %s : i32 loc(#loc8) + %off_702 = arith.addi %off_701, %pid : i32 loc(#loc8) + %off_703 = arith.muli %off_702, %s : i32 loc(#loc8) + %off_704 = arith.addi %off_703, %pid : i32 loc(#loc8) + %off_705 = arith.muli %off_704, %s : i32 loc(#loc8) + %off_706 = arith.addi %off_705, %pid : i32 loc(#loc8) + %off_707 = arith.muli %off_706, %s : i32 loc(#loc8) + %off_708 = arith.addi %off_707, %pid : i32 loc(#loc8) + %off_709 = arith.muli %off_708, %s : i32 loc(#loc8) + %off_710 = arith.addi %off_709, %pid : i32 loc(#loc8) + %off_711 = arith.muli %off_710, %s : i32 loc(#loc8) + %off_712 = arith.addi %off_711, %pid : i32 loc(#loc8) + %off_713 = arith.muli %off_712, %s : i32 loc(#loc8) + %off_714 = arith.addi %off_713, %pid : i32 loc(#loc8) + %off_715 = arith.muli %off_714, %s : i32 loc(#loc8) + %off_716 = arith.addi %off_715, %pid : i32 loc(#loc8) + %off_717 = arith.muli %off_716, %s : i32 loc(#loc8) + %off_718 = arith.addi %off_717, %pid : i32 loc(#loc8) + %off_719 = arith.muli %off_718, %s : i32 loc(#loc8) + %off_720 = arith.addi %off_719, %pid : i32 loc(#loc8) + %off_721 = arith.muli %off_720, %s : i32 loc(#loc8) + %off_722 = arith.addi %off_721, %pid : i32 loc(#loc8) + %off_723 = arith.muli %off_722, %s : i32 loc(#loc8) + %off_724 = arith.addi %off_723, %pid : i32 loc(#loc8) + %off_725 = arith.muli %off_724, %s : i32 loc(#loc8) + %off_726 = arith.addi %off_725, %pid : i32 loc(#loc8) + %off_727 = arith.muli %off_726, %s : i32 loc(#loc8) + %off_728 = arith.addi %off_727, %pid : i32 loc(#loc8) + %off_729 = arith.muli %off_728, %s : i32 loc(#loc8) + %off_730 = arith.addi %off_729, %pid : i32 loc(#loc8) + %off_731 = arith.muli %off_730, %s : i32 loc(#loc8) + %off_732 = arith.addi %off_731, %pid : i32 loc(#loc8) + %off_733 = arith.muli %off_732, %s : i32 loc(#loc8) + %off_734 = arith.addi %off_733, %pid : i32 loc(#loc8) + %off_735 = arith.muli %off_734, %s : i32 loc(#loc8) + %off_736 = arith.addi %off_735, %pid : i32 loc(#loc8) + %off_737 = arith.muli %off_736, %s : i32 loc(#loc8) + %off_738 = arith.addi %off_737, %pid : i32 loc(#loc8) + %off_739 = arith.muli %off_738, %s : i32 loc(#loc8) + %off_740 = arith.addi %off_739, %pid : i32 loc(#loc8) + %off_741 = arith.muli %off_740, %s : i32 loc(#loc8) + %off_742 = arith.addi %off_741, %pid : i32 loc(#loc8) + %off_743 = arith.muli %off_742, %s : i32 loc(#loc8) + %off_744 = arith.addi %off_743, %pid : i32 loc(#loc8) + %off_745 = arith.muli %off_744, %s : i32 loc(#loc8) + %off_746 = arith.addi %off_745, %pid : i32 loc(#loc8) + %off_747 = arith.muli %off_746, %s : i32 loc(#loc8) + %off_748 = arith.addi %off_747, %pid : i32 loc(#loc8) + %off_749 = arith.muli %off_748, %s : i32 loc(#loc8) + %off_750 = arith.addi %off_749, %pid : i32 loc(#loc8) + %off_751 = arith.muli %off_750, %s : i32 loc(#loc8) + %off_752 = arith.addi %off_751, %pid : i32 loc(#loc8) + %off_753 = arith.muli %off_752, %s : i32 loc(#loc8) + %off_754 = arith.addi %off_753, %pid : i32 loc(#loc8) + %off_755 = arith.muli %off_754, %s : i32 loc(#loc8) + %off_756 = arith.addi %off_755, %pid : i32 loc(#loc8) + %off_757 = arith.muli %off_756, %s : i32 loc(#loc8) + %off_758 = arith.addi %off_757, %pid : i32 loc(#loc8) + %off_759 = arith.muli %off_758, %s : i32 loc(#loc8) + %off_760 = arith.addi %off_759, %pid : i32 loc(#loc8) + %off_761 = arith.muli %off_760, %s : i32 loc(#loc8) + %off_762 = arith.addi %off_761, %pid : i32 loc(#loc8) + %off_763 = arith.muli %off_762, %s : i32 loc(#loc8) + %off_764 = arith.addi %off_763, %pid : i32 loc(#loc8) + %off_765 = arith.muli %off_764, %s : i32 loc(#loc8) + %off_766 = arith.addi %off_765, %pid : i32 loc(#loc8) + %off_767 = arith.muli %off_766, %s : i32 loc(#loc8) + %off_768 = arith.addi %off_767, %pid : i32 loc(#loc8) + %off_769 = arith.muli %off_768, %s : i32 loc(#loc8) + %off_770 = arith.addi %off_769, %pid : i32 loc(#loc8) + %off_771 = arith.muli %off_770, %s : i32 loc(#loc8) + %off_772 = arith.addi %off_771, %pid : i32 loc(#loc8) + %off_773 = arith.muli %off_772, %s : i32 loc(#loc8) + %off_774 = arith.addi %off_773, %pid : i32 loc(#loc8) + %off_775 = arith.muli %off_774, %s : i32 loc(#loc8) + %off_776 = arith.addi %off_775, %pid : i32 loc(#loc8) + %off_777 = arith.muli %off_776, %s : i32 loc(#loc8) + %off_778 = arith.addi %off_777, %pid : i32 loc(#loc8) + %off_779 = arith.muli %off_778, %s : i32 loc(#loc8) + %off_780 = arith.addi %off_779, %pid : i32 loc(#loc8) + %off_781 = arith.muli %off_780, %s : i32 loc(#loc8) + %off_782 = arith.addi %off_781, %pid : i32 loc(#loc8) + %off_783 = arith.muli %off_782, %s : i32 loc(#loc8) + %off_784 = arith.addi %off_783, %pid : i32 loc(#loc8) + %off_785 = arith.muli %off_784, %s : i32 loc(#loc8) + %off_786 = arith.addi %off_785, %pid : i32 loc(#loc8) + %off_787 = arith.muli %off_786, %s : i32 loc(#loc8) + %off_788 = arith.addi %off_787, %pid : i32 loc(#loc8) + %off_789 = arith.muli %off_788, %s : i32 loc(#loc8) + %off_790 = arith.addi %off_789, %pid : i32 loc(#loc8) + %off_791 = arith.muli %off_790, %s : i32 loc(#loc8) + %off_792 = arith.addi %off_791, %pid : i32 loc(#loc8) + %off_793 = arith.muli %off_792, %s : i32 loc(#loc8) + %off_794 = arith.addi %off_793, %pid : i32 loc(#loc8) + %off_795 = arith.muli %off_794, %s : i32 loc(#loc8) + %off_796 = arith.addi %off_795, %pid : i32 loc(#loc8) + %off_797 = arith.muli %off_796, %s : i32 loc(#loc8) + %off_798 = arith.addi %off_797, %pid : i32 loc(#loc8) + %off_799 = arith.muli %off_798, %s : i32 loc(#loc8) + %off_800 = arith.addi %off_799, %pid : i32 loc(#loc8) + %off_801 = arith.muli %off_800, %s : i32 loc(#loc8) + %off_802 = arith.addi %off_801, %pid : i32 loc(#loc8) + %off_803 = arith.muli %off_802, %s : i32 loc(#loc8) + %off_804 = arith.addi %off_803, %pid : i32 loc(#loc8) + %off_805 = arith.muli %off_804, %s : i32 loc(#loc8) + %off_806 = arith.addi %off_805, %pid : i32 loc(#loc8) + %off_807 = arith.muli %off_806, %s : i32 loc(#loc8) + %off_808 = arith.addi %off_807, %pid : i32 loc(#loc8) + %off_809 = arith.muli %off_808, %s : i32 loc(#loc8) + %off_810 = arith.addi %off_809, %pid : i32 loc(#loc8) + %off_811 = arith.muli %off_810, %s : i32 loc(#loc8) + %off_812 = arith.addi %off_811, %pid : i32 loc(#loc8) + %off_813 = arith.muli %off_812, %s : i32 loc(#loc8) + %off_814 = arith.addi %off_813, %pid : i32 loc(#loc8) + %off_815 = arith.muli %off_814, %s : i32 loc(#loc8) + %off_816 = arith.addi %off_815, %pid : i32 loc(#loc8) + %off_817 = arith.muli %off_816, %s : i32 loc(#loc8) + %off_818 = arith.addi %off_817, %pid : i32 loc(#loc8) + %off_819 = arith.muli %off_818, %s : i32 loc(#loc8) + %off_820 = arith.addi %off_819, %pid : i32 loc(#loc8) + %off_821 = arith.muli %off_820, %s : i32 loc(#loc8) + %off_822 = arith.addi %off_821, %pid : i32 loc(#loc8) + %off_823 = arith.muli %off_822, %s : i32 loc(#loc8) + %off_824 = arith.addi %off_823, %pid : i32 loc(#loc8) + %off_825 = arith.muli %off_824, %s : i32 loc(#loc8) + %off_826 = arith.addi %off_825, %pid : i32 loc(#loc8) + %off_827 = arith.muli %off_826, %s : i32 loc(#loc8) + %off_828 = arith.addi %off_827, %pid : i32 loc(#loc8) + %off_829 = arith.muli %off_828, %s : i32 loc(#loc8) + %off_830 = arith.addi %off_829, %pid : i32 loc(#loc8) + %off_831 = arith.muli %off_830, %s : i32 loc(#loc8) + %off_832 = arith.addi %off_831, %pid : i32 loc(#loc8) + %off_833 = arith.muli %off_832, %s : i32 loc(#loc8) + %off_834 = arith.addi %off_833, %pid : i32 loc(#loc8) + %off_835 = arith.muli %off_834, %s : i32 loc(#loc8) + %off_836 = arith.addi %off_835, %pid : i32 loc(#loc8) + %off_837 = arith.muli %off_836, %s : i32 loc(#loc8) + %off_838 = arith.addi %off_837, %pid : i32 loc(#loc8) + %off_839 = arith.muli %off_838, %s : i32 loc(#loc8) + %off_840 = arith.addi %off_839, %pid : i32 loc(#loc8) + %off_841 = arith.muli %off_840, %s : i32 loc(#loc8) + %off_842 = arith.addi %off_841, %pid : i32 loc(#loc8) + %off_843 = arith.muli %off_842, %s : i32 loc(#loc8) + %off_844 = arith.addi %off_843, %pid : i32 loc(#loc8) + %off_845 = arith.muli %off_844, %s : i32 loc(#loc8) + %off_846 = arith.addi %off_845, %pid : i32 loc(#loc8) + %off_847 = arith.muli %off_846, %s : i32 loc(#loc8) + %off_848 = arith.addi %off_847, %pid : i32 loc(#loc8) + %off_849 = arith.muli %off_848, %s : i32 loc(#loc8) + %off_850 = arith.addi %off_849, %pid : i32 loc(#loc8) + %off_851 = arith.muli %off_850, %s : i32 loc(#loc8) + %off_852 = arith.addi %off_851, %pid : i32 loc(#loc8) + %off_853 = arith.muli %off_852, %s : i32 loc(#loc8) + %off_854 = arith.addi %off_853, %pid : i32 loc(#loc8) + %off_855 = arith.muli %off_854, %s : i32 loc(#loc8) + %off_856 = arith.addi %off_855, %pid : i32 loc(#loc8) + %off_857 = arith.muli %off_856, %s : i32 loc(#loc8) + %off_858 = arith.addi %off_857, %pid : i32 loc(#loc8) + %off_859 = arith.muli %off_858, %s : i32 loc(#loc8) + %off_860 = arith.addi %off_859, %pid : i32 loc(#loc8) + %off_861 = arith.muli %off_860, %s : i32 loc(#loc8) + %off_862 = arith.addi %off_861, %pid : i32 loc(#loc8) + %off_863 = arith.muli %off_862, %s : i32 loc(#loc8) + %off_864 = arith.addi %off_863, %pid : i32 loc(#loc8) + %off_865 = arith.muli %off_864, %s : i32 loc(#loc8) + %off_866 = arith.addi %off_865, %pid : i32 loc(#loc8) + %off_867 = arith.muli %off_866, %s : i32 loc(#loc8) + %off_868 = arith.addi %off_867, %pid : i32 loc(#loc8) + %off_869 = arith.muli %off_868, %s : i32 loc(#loc8) + %off_870 = arith.addi %off_869, %pid : i32 loc(#loc8) + %off_871 = arith.muli %off_870, %s : i32 loc(#loc8) + %off_872 = arith.addi %off_871, %pid : i32 loc(#loc8) + %off_873 = arith.muli %off_872, %s : i32 loc(#loc8) + %off_874 = arith.addi %off_873, %pid : i32 loc(#loc8) + %off_875 = arith.muli %off_874, %s : i32 loc(#loc8) + %off_876 = arith.addi %off_875, %pid : i32 loc(#loc8) + %off_877 = arith.muli %off_876, %s : i32 loc(#loc8) + %off_878 = arith.addi %off_877, %pid : i32 loc(#loc8) + %off_879 = arith.muli %off_878, %s : i32 loc(#loc8) + %off_880 = arith.addi %off_879, %pid : i32 loc(#loc8) + %off_881 = arith.muli %off_880, %s : i32 loc(#loc8) + %off_882 = arith.addi %off_881, %pid : i32 loc(#loc8) + %off_883 = arith.muli %off_882, %s : i32 loc(#loc8) + %off_884 = arith.addi %off_883, %pid : i32 loc(#loc8) + %off_885 = arith.muli %off_884, %s : i32 loc(#loc8) + %off_886 = arith.addi %off_885, %pid : i32 loc(#loc8) + %off_887 = arith.muli %off_886, %s : i32 loc(#loc8) + %off_888 = arith.addi %off_887, %pid : i32 loc(#loc8) + %off_889 = arith.muli %off_888, %s : i32 loc(#loc8) + %off_890 = arith.addi %off_889, %pid : i32 loc(#loc8) + %off_891 = arith.muli %off_890, %s : i32 loc(#loc8) + %off_892 = arith.addi %off_891, %pid : i32 loc(#loc8) + %off_893 = arith.muli %off_892, %s : i32 loc(#loc8) + %off_894 = arith.addi %off_893, %pid : i32 loc(#loc8) + %off_895 = arith.muli %off_894, %s : i32 loc(#loc8) + %off_896 = arith.addi %off_895, %pid : i32 loc(#loc8) + %off_897 = arith.muli %off_896, %s : i32 loc(#loc8) + %off_898 = arith.addi %off_897, %pid : i32 loc(#loc8) + %off_899 = arith.muli %off_898, %s : i32 loc(#loc8) + %off_900 = arith.addi %off_899, %pid : i32 loc(#loc8) + %off_901 = arith.muli %off_900, %s : i32 loc(#loc8) + %off_902 = arith.addi %off_901, %pid : i32 loc(#loc8) + %off_903 = arith.muli %off_902, %s : i32 loc(#loc8) + %off_904 = arith.addi %off_903, %pid : i32 loc(#loc8) + %off_905 = arith.muli %off_904, %s : i32 loc(#loc8) + %off_906 = arith.addi %off_905, %pid : i32 loc(#loc8) + %off_907 = arith.muli %off_906, %s : i32 loc(#loc8) + %off_908 = arith.addi %off_907, %pid : i32 loc(#loc8) + %off_909 = arith.muli %off_908, %s : i32 loc(#loc8) + %off_910 = arith.addi %off_909, %pid : i32 loc(#loc8) + %off_911 = arith.muli %off_910, %s : i32 loc(#loc8) + %off_912 = arith.addi %off_911, %pid : i32 loc(#loc8) + %off_913 = arith.muli %off_912, %s : i32 loc(#loc8) + %off_914 = arith.addi %off_913, %pid : i32 loc(#loc8) + %off_915 = arith.muli %off_914, %s : i32 loc(#loc8) + %off_916 = arith.addi %off_915, %pid : i32 loc(#loc8) + %off_917 = arith.muli %off_916, %s : i32 loc(#loc8) + %off_918 = arith.addi %off_917, %pid : i32 loc(#loc8) + %off_919 = arith.muli %off_918, %s : i32 loc(#loc8) + %off_920 = arith.addi %off_919, %pid : i32 loc(#loc8) + %off_921 = arith.muli %off_920, %s : i32 loc(#loc8) + %off_922 = arith.addi %off_921, %pid : i32 loc(#loc8) + %off_923 = arith.muli %off_922, %s : i32 loc(#loc8) + %off_924 = arith.addi %off_923, %pid : i32 loc(#loc8) + %off_925 = arith.muli %off_924, %s : i32 loc(#loc8) + %off_926 = arith.addi %off_925, %pid : i32 loc(#loc8) + %off_927 = arith.muli %off_926, %s : i32 loc(#loc8) + %off_928 = arith.addi %off_927, %pid : i32 loc(#loc8) + %off_929 = arith.muli %off_928, %s : i32 loc(#loc8) + %off_930 = arith.addi %off_929, %pid : i32 loc(#loc8) + %off_931 = arith.muli %off_930, %s : i32 loc(#loc8) + %off_932 = arith.addi %off_931, %pid : i32 loc(#loc8) + %off_933 = arith.muli %off_932, %s : i32 loc(#loc8) + %off_934 = arith.addi %off_933, %pid : i32 loc(#loc8) + %off_935 = arith.muli %off_934, %s : i32 loc(#loc8) + %off_936 = arith.addi %off_935, %pid : i32 loc(#loc8) + %off_937 = arith.muli %off_936, %s : i32 loc(#loc8) + %off_938 = arith.addi %off_937, %pid : i32 loc(#loc8) + %off_939 = arith.muli %off_938, %s : i32 loc(#loc8) + %off_940 = arith.addi %off_939, %pid : i32 loc(#loc8) + %off_941 = arith.muli %off_940, %s : i32 loc(#loc8) + %off_942 = arith.addi %off_941, %pid : i32 loc(#loc8) + %off_943 = arith.muli %off_942, %s : i32 loc(#loc8) + %off_944 = arith.addi %off_943, %pid : i32 loc(#loc8) + %off_945 = arith.muli %off_944, %s : i32 loc(#loc8) + %off_946 = arith.addi %off_945, %pid : i32 loc(#loc8) + %off_947 = arith.muli %off_946, %s : i32 loc(#loc8) + %off_948 = arith.addi %off_947, %pid : i32 loc(#loc8) + %off_949 = arith.muli %off_948, %s : i32 loc(#loc8) + %off_950 = arith.addi %off_949, %pid : i32 loc(#loc8) + %off_951 = arith.muli %off_950, %s : i32 loc(#loc8) + %off_952 = arith.addi %off_951, %pid : i32 loc(#loc8) + %off_953 = arith.muli %off_952, %s : i32 loc(#loc8) + %off_954 = arith.addi %off_953, %pid : i32 loc(#loc8) + %off_955 = arith.muli %off_954, %s : i32 loc(#loc8) + %off_956 = arith.addi %off_955, %pid : i32 loc(#loc8) + %off_957 = arith.muli %off_956, %s : i32 loc(#loc8) + %off_958 = arith.addi %off_957, %pid : i32 loc(#loc8) + %off_959 = arith.muli %off_958, %s : i32 loc(#loc8) + %off_960 = arith.addi %off_959, %pid : i32 loc(#loc8) + %off_961 = arith.muli %off_960, %s : i32 loc(#loc8) + %off_962 = arith.addi %off_961, %pid : i32 loc(#loc8) + %off_963 = arith.muli %off_962, %s : i32 loc(#loc8) + %off_964 = arith.addi %off_963, %pid : i32 loc(#loc8) + %off_965 = arith.muli %off_964, %s : i32 loc(#loc8) + %off_966 = arith.addi %off_965, %pid : i32 loc(#loc8) + %off_967 = arith.muli %off_966, %s : i32 loc(#loc8) + %off_968 = arith.addi %off_967, %pid : i32 loc(#loc8) + %off_969 = arith.muli %off_968, %s : i32 loc(#loc8) + %off_970 = arith.addi %off_969, %pid : i32 loc(#loc8) + %off_971 = arith.muli %off_970, %s : i32 loc(#loc8) + %off_972 = arith.addi %off_971, %pid : i32 loc(#loc8) + %off_973 = arith.muli %off_972, %s : i32 loc(#loc8) + %off_974 = arith.addi %off_973, %pid : i32 loc(#loc8) + %off_975 = arith.muli %off_974, %s : i32 loc(#loc8) + %off_976 = arith.addi %off_975, %pid : i32 loc(#loc8) + %off_977 = arith.muli %off_976, %s : i32 loc(#loc8) + %off_978 = arith.addi %off_977, %pid : i32 loc(#loc8) + %off_979 = arith.muli %off_978, %s : i32 loc(#loc8) + %off_980 = arith.addi %off_979, %pid : i32 loc(#loc8) + %off_981 = arith.muli %off_980, %s : i32 loc(#loc8) + %off_982 = arith.addi %off_981, %pid : i32 loc(#loc8) + %off_983 = arith.muli %off_982, %s : i32 loc(#loc8) + %off_984 = arith.addi %off_983, %pid : i32 loc(#loc8) + %off_985 = arith.muli %off_984, %s : i32 loc(#loc8) + %off_986 = arith.addi %off_985, %pid : i32 loc(#loc8) + %off_987 = arith.muli %off_986, %s : i32 loc(#loc8) + %off_988 = arith.addi %off_987, %pid : i32 loc(#loc8) + %off_989 = arith.muli %off_988, %s : i32 loc(#loc8) + %off_990 = arith.addi %off_989, %pid : i32 loc(#loc8) + %off_991 = arith.muli %off_990, %s : i32 loc(#loc8) + %off_992 = arith.addi %off_991, %pid : i32 loc(#loc8) + %off_993 = arith.muli %off_992, %s : i32 loc(#loc8) + %off_994 = arith.addi %off_993, %pid : i32 loc(#loc8) + %off_995 = arith.muli %off_994, %s : i32 loc(#loc8) + %off_996 = arith.addi %off_995, %pid : i32 loc(#loc8) + %off_997 = arith.muli %off_996, %s : i32 loc(#loc8) + %off_998 = arith.addi %off_997, %pid : i32 loc(#loc8) + %off_999 = arith.muli %off_998, %s : i32 loc(#loc8) + %off_1000 = arith.addi %off_999, %pid : i32 loc(#loc8) + %off_1001 = arith.muli %off_1000, %s : i32 loc(#loc8) + %off_1002 = arith.addi %off_1001, %pid : i32 loc(#loc8) + %off_1003 = arith.muli %off_1002, %s : i32 loc(#loc8) + %off_1004 = arith.addi %off_1003, %pid : i32 loc(#loc8) + %off_1005 = arith.muli %off_1004, %s : i32 loc(#loc8) + %off_1006 = arith.addi %off_1005, %pid : i32 loc(#loc8) + %off_1007 = arith.muli %off_1006, %s : i32 loc(#loc8) + %off_1008 = arith.addi %off_1007, %pid : i32 loc(#loc8) + %off_1009 = arith.muli %off_1008, %s : i32 loc(#loc8) + %off_1010 = arith.addi %off_1009, %pid : i32 loc(#loc8) + %off_1011 = arith.muli %off_1010, %s : i32 loc(#loc8) + %off_1012 = arith.addi %off_1011, %pid : i32 loc(#loc8) + %off_1013 = arith.muli %off_1012, %s : i32 loc(#loc8) + %off_1014 = arith.addi %off_1013, %pid : i32 loc(#loc8) + %off_1015 = arith.muli %off_1014, %s : i32 loc(#loc8) + %off_1016 = arith.addi %off_1015, %pid : i32 loc(#loc8) + %off_1017 = arith.muli %off_1016, %s : i32 loc(#loc8) + %off_1018 = arith.addi %off_1017, %pid : i32 loc(#loc8) + %off_1019 = arith.muli %off_1018, %s : i32 loc(#loc8) + %off_1020 = arith.addi %off_1019, %pid : i32 loc(#loc8) + %off_1021 = arith.muli %off_1020, %s : i32 loc(#loc8) + %off_1022 = arith.addi %off_1021, %pid : i32 loc(#loc8) + %off_1023 = arith.muli %off_1022, %s : i32 loc(#loc8) + %off_1024 = arith.addi %off_1023, %pid : i32 loc(#loc8) + %off_1025 = arith.muli %off_1024, %s : i32 loc(#loc8) + %off_1026 = arith.addi %off_1025, %pid : i32 loc(#loc8) + %off_1027 = arith.muli %off_1026, %s : i32 loc(#loc8) + %off_1028 = arith.addi %off_1027, %pid : i32 loc(#loc8) + %off_1029 = arith.muli %off_1028, %s : i32 loc(#loc8) + %off_1030 = arith.addi %off_1029, %pid : i32 loc(#loc8) + %off_1031 = arith.muli %off_1030, %s : i32 loc(#loc8) + %off_1032 = arith.addi %off_1031, %pid : i32 loc(#loc8) + %off_1033 = arith.muli %off_1032, %s : i32 loc(#loc8) + %off_1034 = arith.addi %off_1033, %pid : i32 loc(#loc8) + %off_1035 = arith.muli %off_1034, %s : i32 loc(#loc8) + %off_1036 = arith.addi %off_1035, %pid : i32 loc(#loc8) + %off_1037 = arith.muli %off_1036, %s : i32 loc(#loc8) + %off_1038 = arith.addi %off_1037, %pid : i32 loc(#loc8) + %off_1039 = arith.muli %off_1038, %s : i32 loc(#loc8) + %off_1040 = arith.addi %off_1039, %pid : i32 loc(#loc8) + %off_1041 = arith.muli %off_1040, %s : i32 loc(#loc8) + %off_1042 = arith.addi %off_1041, %pid : i32 loc(#loc8) + %off_1043 = arith.muli %off_1042, %s : i32 loc(#loc8) + %off_1044 = arith.addi %off_1043, %pid : i32 loc(#loc8) + %off_1045 = arith.muli %off_1044, %s : i32 loc(#loc8) + %off_1046 = arith.addi %off_1045, %pid : i32 loc(#loc8) + %off_1047 = arith.muli %off_1046, %s : i32 loc(#loc8) + %off_1048 = arith.addi %off_1047, %pid : i32 loc(#loc8) + %off_1049 = arith.muli %off_1048, %s : i32 loc(#loc8) + %off_1050 = arith.addi %off_1049, %pid : i32 loc(#loc8) + %off_1051 = arith.muli %off_1050, %s : i32 loc(#loc8) + %off_1052 = arith.addi %off_1051, %pid : i32 loc(#loc8) + %off_1053 = arith.muli %off_1052, %s : i32 loc(#loc8) + %off_1054 = arith.addi %off_1053, %pid : i32 loc(#loc8) + %off_1055 = arith.muli %off_1054, %s : i32 loc(#loc8) + %off_1056 = arith.addi %off_1055, %pid : i32 loc(#loc8) + %off_1057 = arith.muli %off_1056, %s : i32 loc(#loc8) + %off_1058 = arith.addi %off_1057, %pid : i32 loc(#loc8) + %off_1059 = arith.muli %off_1058, %s : i32 loc(#loc8) + %off_1060 = arith.addi %off_1059, %pid : i32 loc(#loc8) + %off_1061 = arith.muli %off_1060, %s : i32 loc(#loc8) + %off_1062 = arith.addi %off_1061, %pid : i32 loc(#loc8) + %off_1063 = arith.muli %off_1062, %s : i32 loc(#loc8) + %off_1064 = arith.addi %off_1063, %pid : i32 loc(#loc8) + %off_1065 = arith.muli %off_1064, %s : i32 loc(#loc8) + %off_1066 = arith.addi %off_1065, %pid : i32 loc(#loc8) + %off_1067 = arith.muli %off_1066, %s : i32 loc(#loc8) + %off_1068 = arith.addi %off_1067, %pid : i32 loc(#loc8) + %off_1069 = arith.muli %off_1068, %s : i32 loc(#loc8) + %off_1070 = arith.addi %off_1069, %pid : i32 loc(#loc8) + %off_1071 = arith.muli %off_1070, %s : i32 loc(#loc8) + %off_1072 = arith.addi %off_1071, %pid : i32 loc(#loc8) + %off_1073 = arith.muli %off_1072, %s : i32 loc(#loc8) + %off_1074 = arith.addi %off_1073, %pid : i32 loc(#loc8) + %off_1075 = arith.muli %off_1074, %s : i32 loc(#loc8) + %off_1076 = arith.addi %off_1075, %pid : i32 loc(#loc8) + %off_1077 = arith.muli %off_1076, %s : i32 loc(#loc8) + %off_1078 = arith.addi %off_1077, %pid : i32 loc(#loc8) + %off_1079 = arith.muli %off_1078, %s : i32 loc(#loc8) + %off_1080 = arith.addi %off_1079, %pid : i32 loc(#loc8) + %off_1081 = arith.muli %off_1080, %s : i32 loc(#loc8) + %off_1082 = arith.addi %off_1081, %pid : i32 loc(#loc8) + %off_1083 = arith.muli %off_1082, %s : i32 loc(#loc8) + %off_1084 = arith.addi %off_1083, %pid : i32 loc(#loc8) + %off_1085 = arith.muli %off_1084, %s : i32 loc(#loc8) + %off_1086 = arith.addi %off_1085, %pid : i32 loc(#loc8) + %off_1087 = arith.muli %off_1086, %s : i32 loc(#loc8) + %off_1088 = arith.addi %off_1087, %pid : i32 loc(#loc8) + %off_1089 = arith.muli %off_1088, %s : i32 loc(#loc8) + %off_1090 = arith.addi %off_1089, %pid : i32 loc(#loc8) + %off_1091 = arith.muli %off_1090, %s : i32 loc(#loc8) + %off_1092 = arith.addi %off_1091, %pid : i32 loc(#loc8) + %off_1093 = arith.muli %off_1092, %s : i32 loc(#loc8) + %off_1094 = arith.addi %off_1093, %pid : i32 loc(#loc8) + %off_1095 = arith.muli %off_1094, %s : i32 loc(#loc8) + %off_1096 = arith.addi %off_1095, %pid : i32 loc(#loc8) + %off_1097 = arith.muli %off_1096, %s : i32 loc(#loc8) + %off_1098 = arith.addi %off_1097, %pid : i32 loc(#loc8) + %off_1099 = arith.muli %off_1098, %s : i32 loc(#loc8) + %off_1100 = arith.addi %off_1099, %pid : i32 loc(#loc8) + %off_1101 = arith.muli %off_1100, %s : i32 loc(#loc8) + %off_1102 = arith.addi %off_1101, %pid : i32 loc(#loc8) + %off_1103 = arith.muli %off_1102, %s : i32 loc(#loc8) + %off_1104 = arith.addi %off_1103, %pid : i32 loc(#loc8) + %off_1105 = arith.muli %off_1104, %s : i32 loc(#loc8) + %off_1106 = arith.addi %off_1105, %pid : i32 loc(#loc8) + %off_1107 = arith.muli %off_1106, %s : i32 loc(#loc8) + %off_1108 = arith.addi %off_1107, %pid : i32 loc(#loc8) + %off_1109 = arith.muli %off_1108, %s : i32 loc(#loc8) + %off_1110 = arith.addi %off_1109, %pid : i32 loc(#loc8) + %off_1111 = arith.muli %off_1110, %s : i32 loc(#loc8) + %off_1112 = arith.addi %off_1111, %pid : i32 loc(#loc8) + %off_1113 = arith.muli %off_1112, %s : i32 loc(#loc8) + %off_1114 = arith.addi %off_1113, %pid : i32 loc(#loc8) + %off_1115 = arith.muli %off_1114, %s : i32 loc(#loc8) + %off_1116 = arith.addi %off_1115, %pid : i32 loc(#loc8) + %off_1117 = arith.muli %off_1116, %s : i32 loc(#loc8) + %off_1118 = arith.addi %off_1117, %pid : i32 loc(#loc8) + %off_1119 = arith.muli %off_1118, %s : i32 loc(#loc8) + %off_1120 = arith.addi %off_1119, %pid : i32 loc(#loc8) + %off_1121 = arith.muli %off_1120, %s : i32 loc(#loc8) + %off_1122 = arith.addi %off_1121, %pid : i32 loc(#loc8) + %off_1123 = arith.muli %off_1122, %s : i32 loc(#loc8) + %off_1124 = arith.addi %off_1123, %pid : i32 loc(#loc8) + %off_1125 = arith.muli %off_1124, %s : i32 loc(#loc8) + %off_1126 = arith.addi %off_1125, %pid : i32 loc(#loc8) + %off_1127 = arith.muli %off_1126, %s : i32 loc(#loc8) + %off_1128 = arith.addi %off_1127, %pid : i32 loc(#loc8) + %off_1129 = arith.muli %off_1128, %s : i32 loc(#loc8) + %off_1130 = arith.addi %off_1129, %pid : i32 loc(#loc8) + %off_1131 = arith.muli %off_1130, %s : i32 loc(#loc8) + %off_1132 = arith.addi %off_1131, %pid : i32 loc(#loc8) + %off_1133 = arith.muli %off_1132, %s : i32 loc(#loc8) + %off_1134 = arith.addi %off_1133, %pid : i32 loc(#loc8) + %off_1135 = arith.muli %off_1134, %s : i32 loc(#loc8) + %off_1136 = arith.addi %off_1135, %pid : i32 loc(#loc8) + %off_1137 = arith.muli %off_1136, %s : i32 loc(#loc8) + %off_1138 = arith.addi %off_1137, %pid : i32 loc(#loc8) + %off_1139 = arith.muli %off_1138, %s : i32 loc(#loc8) + %off_1140 = arith.addi %off_1139, %pid : i32 loc(#loc8) + %off_1141 = arith.muli %off_1140, %s : i32 loc(#loc8) + %off_1142 = arith.addi %off_1141, %pid : i32 loc(#loc8) + %off_1143 = arith.muli %off_1142, %s : i32 loc(#loc8) + %off_1144 = arith.addi %off_1143, %pid : i32 loc(#loc8) + %off_1145 = arith.muli %off_1144, %s : i32 loc(#loc8) + %off_1146 = arith.addi %off_1145, %pid : i32 loc(#loc8) + %off_1147 = arith.muli %off_1146, %s : i32 loc(#loc8) + %off_1148 = arith.addi %off_1147, %pid : i32 loc(#loc8) + %off_1149 = arith.muli %off_1148, %s : i32 loc(#loc8) + %off_1150 = arith.addi %off_1149, %pid : i32 loc(#loc8) + %off_1151 = arith.muli %off_1150, %s : i32 loc(#loc8) + %off_1152 = arith.addi %off_1151, %pid : i32 loc(#loc8) + %off_1153 = arith.muli %off_1152, %s : i32 loc(#loc8) + %off_1154 = arith.addi %off_1153, %pid : i32 loc(#loc8) + %off_1155 = arith.muli %off_1154, %s : i32 loc(#loc8) + %off_1156 = arith.addi %off_1155, %pid : i32 loc(#loc8) + %off_1157 = arith.muli %off_1156, %s : i32 loc(#loc8) + %off_1158 = arith.addi %off_1157, %pid : i32 loc(#loc8) + %off_1159 = arith.muli %off_1158, %s : i32 loc(#loc8) + %off_1160 = arith.addi %off_1159, %pid : i32 loc(#loc8) + %off_1161 = arith.muli %off_1160, %s : i32 loc(#loc8) + %off_1162 = arith.addi %off_1161, %pid : i32 loc(#loc8) + %off_1163 = arith.muli %off_1162, %s : i32 loc(#loc8) + %off_1164 = arith.addi %off_1163, %pid : i32 loc(#loc8) + %off_1165 = arith.muli %off_1164, %s : i32 loc(#loc8) + %off_1166 = arith.addi %off_1165, %pid : i32 loc(#loc8) + %off_1167 = arith.muli %off_1166, %s : i32 loc(#loc8) + %off_1168 = arith.addi %off_1167, %pid : i32 loc(#loc8) + %off_1169 = arith.muli %off_1168, %s : i32 loc(#loc8) + %off_1170 = arith.addi %off_1169, %pid : i32 loc(#loc8) + %off_1171 = arith.muli %off_1170, %s : i32 loc(#loc8) + %off_1172 = arith.addi %off_1171, %pid : i32 loc(#loc8) + %off_1173 = arith.muli %off_1172, %s : i32 loc(#loc8) + %off_1174 = arith.addi %off_1173, %pid : i32 loc(#loc8) + %off_1175 = arith.muli %off_1174, %s : i32 loc(#loc8) + %off_1176 = arith.addi %off_1175, %pid : i32 loc(#loc8) + %off_1177 = arith.muli %off_1176, %s : i32 loc(#loc8) + %off_1178 = arith.addi %off_1177, %pid : i32 loc(#loc8) + %off_1179 = arith.muli %off_1178, %s : i32 loc(#loc8) + %off_1180 = arith.addi %off_1179, %pid : i32 loc(#loc8) + %off_1181 = arith.muli %off_1180, %s : i32 loc(#loc8) + %off_1182 = arith.addi %off_1181, %pid : i32 loc(#loc8) + %off_1183 = arith.muli %off_1182, %s : i32 loc(#loc8) + %off_1184 = arith.addi %off_1183, %pid : i32 loc(#loc8) + %off_1185 = arith.muli %off_1184, %s : i32 loc(#loc8) + %off_1186 = arith.addi %off_1185, %pid : i32 loc(#loc8) + %off_1187 = arith.muli %off_1186, %s : i32 loc(#loc8) + %off_1188 = arith.addi %off_1187, %pid : i32 loc(#loc8) + %off_1189 = arith.muli %off_1188, %s : i32 loc(#loc8) + %off_1190 = arith.addi %off_1189, %pid : i32 loc(#loc8) + %off_1191 = arith.muli %off_1190, %s : i32 loc(#loc8) + %off_1192 = arith.addi %off_1191, %pid : i32 loc(#loc8) + %off_1193 = arith.muli %off_1192, %s : i32 loc(#loc8) + %off_1194 = arith.addi %off_1193, %pid : i32 loc(#loc8) + %off_1195 = arith.muli %off_1194, %s : i32 loc(#loc8) + %off_1196 = arith.addi %off_1195, %pid : i32 loc(#loc8) + %off_1197 = arith.muli %off_1196, %s : i32 loc(#loc8) + %off_1198 = arith.addi %off_1197, %pid : i32 loc(#loc8) + %0 = tt.addptr %out_ptr, %off_1198 : !tt.ptr, i32 loc(#loc4) + tt.store %0, %cst : !tt.ptr loc(#loc1) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":66:5) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":62:11) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":65:15) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":66:14) +#loc7 = loc("pid"(#loc2)) +#loc8 = loc("off"(#loc3)) diff --git a/tests/golden/ir/ttir_3.8/kernel_dot_precisions.ttir b/tests/golden/ir/ttir_3.8/kernel_dot_precisions.ttir new file mode 100644 index 000000000..c70d4986b --- /dev/null +++ b/tests/golden/ir/ttir_3.8/kernel_dot_precisions.ttir @@ -0,0 +1,50 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":29:1) +#loc13 = loc("a_ptr"(#loc)) +#loc14 = loc("b_ptr"(#loc)) +#loc15 = loc("c_ptr"(#loc)) +module { + tt.func public @dot_precisions(%a_ptr: !tt.ptr loc("a_ptr"(#loc)), %b_ptr: !tt.ptr loc("b_ptr"(#loc)), %c_ptr: !tt.ptr loc("c_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<0.000000e+00> : tensor<16x16xf32> loc(#loc1) + %idx = arith.constant dense<16> : tensor<16x1xi32> loc(#loc16) + %offs = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32> loc(#loc17) + %idx_0 = tt.expand_dims %offs {axis = 1 : i32} : tensor<16xi32> -> tensor<16x1xi32> loc(#loc16) + %idx_1 = arith.muli %idx_0, %idx : tensor<16x1xi32> loc(#loc16) + %idx_2 = tt.expand_dims %offs {axis = 0 : i32} : tensor<16xi32> -> tensor<1x16xi32> loc(#loc18) + %idx_3 = tt.broadcast %idx_1 : tensor<16x1xi32> -> tensor<16x16xi32> loc(#loc16) + %idx_4 = tt.broadcast %idx_2 : tensor<1x16xi32> -> tensor<16x16xi32> loc(#loc16) + %idx_5 = arith.addi %idx_3, %idx_4 : tensor<16x16xi32> loc(#loc16) + %a = tt.splat %a_ptr : !tt.ptr -> tensor<16x16x!tt.ptr> loc(#loc19) + %a_6 = tt.addptr %a, %idx_5 : tensor<16x16x!tt.ptr>, tensor<16x16xi32> loc(#loc19) + %a_7 = tt.load %a_6 : tensor<16x16x!tt.ptr> loc(#loc20) + %b = tt.splat %b_ptr : !tt.ptr -> tensor<16x16x!tt.ptr> loc(#loc21) + %b_8 = tt.addptr %b, %idx_5 : tensor<16x16x!tt.ptr>, tensor<16x16xi32> loc(#loc21) + %b_9 = tt.load %b_8 : tensor<16x16x!tt.ptr> loc(#loc22) + %c = tt.dot %a_7, %b_9, %cst : tensor<16x16xf32> * tensor<16x16xf32> -> tensor<16x16xf32> loc(#loc23) + %0 = tt.splat %c_ptr : !tt.ptr -> tensor<16x16x!tt.ptr> loc(#loc10) + %1 = tt.addptr %0, %idx_5 : tensor<16x16x!tt.ptr>, tensor<16x16xi32> loc(#loc10) + %d = tt.dot %a_7, %b_9, %c, inputPrecision = tf32x3 : tensor<16x16xf32> * tensor<16x16xf32> -> tensor<16x16xf32> loc(#loc24) + tt.store %1, %d : tensor<16x16x!tt.ptr> loc(#loc12) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":32:11) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":31:12) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":32:35) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":33:17) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":33:9) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":34:17) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":34:9) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":35:9) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":37:14) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":36:9) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":37:5) +#loc16 = loc("idx"(#loc2)) +#loc17 = loc("offs"(#loc3)) +#loc18 = loc("idx"(#loc4)) +#loc19 = loc("a"(#loc5)) +#loc20 = loc("a"(#loc6)) +#loc21 = loc("b"(#loc7)) +#loc22 = loc("b"(#loc8)) +#loc23 = loc("c"(#loc9)) +#loc24 = loc("d"(#loc11)) diff --git a/tests/golden/ir/ttir_3.8/kernel_dot_scaled.ttir b/tests/golden/ir/ttir_3.8/kernel_dot_scaled.ttir new file mode 100644 index 000000000..bd12b8cc9 --- /dev/null +++ b/tests/golden/ir/ttir_3.8/kernel_dot_scaled.ttir @@ -0,0 +1,98 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":70:1) +#loc23 = loc("a_ptr"(#loc)) +#loc24 = loc("as_ptr"(#loc)) +#loc25 = loc("b_ptr"(#loc)) +#loc26 = loc("bs_ptr"(#loc)) +#loc27 = loc("c_ptr"(#loc)) +module { + tt.func public @dot_scaled_k(%a_ptr: !tt.ptr loc("a_ptr"(#loc)), %as_ptr: !tt.ptr loc("as_ptr"(#loc)), %b_ptr: !tt.ptr loc("b_ptr"(#loc)), %bs_ptr: !tt.ptr loc("bs_ptr"(#loc)), %c_ptr: !tt.ptr loc("c_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<128> : tensor<128x1xi32> loc(#loc1) + %c = arith.constant dense<0.000000e+00> : tensor<128x128xf32> loc(#loc28) + %cst_0 = arith.constant dense<2> : tensor<128x1xi32> loc(#loc3) + %b = arith.constant dense<128> : tensor<64x1xi32> loc(#loc29) + %a = arith.constant dense<64> : tensor<128x1xi32> loc(#loc30) + %rm = tt.make_range {end = 128 : i32, start = 0 : i32} : tensor<128xi32> loc(#loc31) + %rk = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc32) + %rs = tt.make_range {end = 2 : i32, start = 0 : i32} : tensor<2xi32> loc(#loc33) + %a_1 = tt.expand_dims %rm {axis = 1 : i32} : tensor<128xi32> -> tensor<128x1xi32> loc(#loc30) + %a_2 = arith.muli %a_1, %a : tensor<128x1xi32> loc(#loc30) + %a_3 = tt.splat %a_ptr : !tt.ptr -> tensor<128x1x!tt.ptr> loc(#loc34) + %a_4 = tt.addptr %a_3, %a_2 : tensor<128x1x!tt.ptr>, tensor<128x1xi32> loc(#loc34) + %a_5 = tt.expand_dims %rk {axis = 0 : i32} : tensor<64xi32> -> tensor<1x64xi32> loc(#loc35) + %a_6 = tt.broadcast %a_4 : tensor<128x1x!tt.ptr> -> tensor<128x64x!tt.ptr> loc(#loc34) + %a_7 = tt.broadcast %a_5 : tensor<1x64xi32> -> tensor<128x64xi32> loc(#loc34) + %a_8 = tt.addptr %a_6, %a_7 : tensor<128x64x!tt.ptr>, tensor<128x64xi32> loc(#loc34) + %a_9 = tt.load %a_8 : tensor<128x64x!tt.ptr> loc(#loc36) + %b_10 = tt.expand_dims %rk {axis = 1 : i32} : tensor<64xi32> -> tensor<64x1xi32> loc(#loc29) + %b_11 = arith.muli %b_10, %b : tensor<64x1xi32> loc(#loc29) + %b_12 = tt.splat %b_ptr : !tt.ptr -> tensor<64x1x!tt.ptr> loc(#loc37) + %b_13 = tt.addptr %b_12, %b_11 : tensor<64x1x!tt.ptr>, tensor<64x1xi32> loc(#loc37) + %b_14 = tt.expand_dims %rm {axis = 0 : i32} : tensor<128xi32> -> tensor<1x128xi32> loc(#loc38) + %b_15 = tt.broadcast %b_13 : tensor<64x1x!tt.ptr> -> tensor<64x128x!tt.ptr> loc(#loc37) + %b_16 = tt.broadcast %b_14 : tensor<1x128xi32> -> tensor<64x128xi32> loc(#loc37) + %b_17 = tt.addptr %b_15, %b_16 : tensor<64x128x!tt.ptr>, tensor<64x128xi32> loc(#loc37) + %b_18 = tt.load %b_17 : tensor<64x128x!tt.ptr> loc(#loc39) + %a_scale = arith.muli %a_1, %cst_0 : tensor<128x1xi32> loc(#loc40) + %a_scale_19 = tt.splat %as_ptr : !tt.ptr -> tensor<128x1x!tt.ptr> loc(#loc41) + %a_scale_20 = tt.addptr %a_scale_19, %a_scale : tensor<128x1x!tt.ptr>, tensor<128x1xi32> loc(#loc41) + %a_scale_21 = tt.expand_dims %rs {axis = 0 : i32} : tensor<2xi32> -> tensor<1x2xi32> loc(#loc42) + %a_scale_22 = tt.broadcast %a_scale_20 : tensor<128x1x!tt.ptr> -> tensor<128x2x!tt.ptr> loc(#loc41) + %a_scale_23 = tt.broadcast %a_scale_21 : tensor<1x2xi32> -> tensor<128x2xi32> loc(#loc41) + %a_scale_24 = tt.addptr %a_scale_22, %a_scale_23 : tensor<128x2x!tt.ptr>, tensor<128x2xi32> loc(#loc41) + %a_scale_25 = tt.load %a_scale_24 : tensor<128x2x!tt.ptr> loc(#loc43) + %b_scale = tt.splat %bs_ptr : !tt.ptr -> tensor<128x1x!tt.ptr> loc(#loc44) + %b_scale_26 = tt.addptr %b_scale, %a_scale : tensor<128x1x!tt.ptr>, tensor<128x1xi32> loc(#loc44) + %b_scale_27 = tt.broadcast %b_scale_26 : tensor<128x1x!tt.ptr> -> tensor<128x2x!tt.ptr> loc(#loc44) + %b_scale_28 = tt.addptr %b_scale_27, %a_scale_23 : tensor<128x2x!tt.ptr>, tensor<128x2xi32> loc(#loc44) + %b_scale_29 = tt.load %b_scale_28 : tensor<128x2x!tt.ptr> loc(#loc45) + %c_30 = tt.dot_scaled %a_9 scale %a_scale_25, %b_18 scale %b_scale_29, %c lhs = e4m3 rhs = e4m3 {fastMath = false} : tensor<128x64xf8E4M3FN>, tensor<128x2xi8> * tensor<64x128xf8E4M3FN>, tensor<128x2xi8> -> tensor<128x128xf32> loc(#loc28) + %0 = arith.muli %a_1, %cst : tensor<128x1xi32> loc(#loc1) + %1 = tt.splat %c_ptr : !tt.ptr -> tensor<128x1x!tt.ptr> loc(#loc21) + %2 = tt.addptr %1, %0 : tensor<128x1x!tt.ptr>, tensor<128x1xi32> loc(#loc21) + %3 = tt.broadcast %2 : tensor<128x1x!tt.ptr> -> tensor<128x128x!tt.ptr> loc(#loc21) + %4 = tt.broadcast %b_14 : tensor<1x128xi32> -> tensor<128x128xi32> loc(#loc21) + %5 = tt.addptr %3, %4 : tensor<128x128x!tt.ptr>, tensor<128x128xi32> loc(#loc21) + tt.store %5, %c_30 : tensor<128x128x!tt.ptr> loc(#loc22) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":81:22) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":80:9) +#loc3 = loc(unknown) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":77:25) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":76:25) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":72:10) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":74:10) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":75:10) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":76:17) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":76:43) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":76:9) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":77:17) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":77:43) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":77:9) +#loc15 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":78:32) +#loc16 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":78:23) +#loc17 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":78:58) +#loc18 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":78:15) +#loc19 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":79:23) +#loc20 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":79:15) +#loc21 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":81:14) +#loc22 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":81:5) +#loc28 = loc("c"(#loc2)) +#loc29 = loc("b"(#loc4)) +#loc30 = loc("a"(#loc5)) +#loc31 = loc("rm"(#loc6)) +#loc32 = loc("rk"(#loc7)) +#loc33 = loc("rs"(#loc8)) +#loc34 = loc("a"(#loc9)) +#loc35 = loc("a"(#loc10)) +#loc36 = loc("a"(#loc11)) +#loc37 = loc("b"(#loc12)) +#loc38 = loc("b"(#loc13)) +#loc39 = loc("b"(#loc14)) +#loc40 = loc("a_scale"(#loc15)) +#loc41 = loc("a_scale"(#loc16)) +#loc42 = loc("a_scale"(#loc17)) +#loc43 = loc("a_scale"(#loc18)) +#loc44 = loc("b_scale"(#loc19)) +#loc45 = loc("b_scale"(#loc20)) diff --git a/tests/golden/ir/ttir_3.8/kernel_eps_consts.ttir b/tests/golden/ir/ttir_3.8/kernel_eps_consts.ttir new file mode 100644 index 000000000..d450c17c0 --- /dev/null +++ b/tests/golden/ir/ttir_3.8/kernel_eps_consts.ttir @@ -0,0 +1,35 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":41:1) +#loc9 = loc("x_ptr"(#loc)) +#loc10 = loc("s_ptr"(#loc)) +#loc11 = loc("out_ptr"(#loc)) +module { + tt.func public @eps_consts(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %s_ptr: !tt.ptr loc("s_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<9.99999996E-13> : tensor<64xf32> loc(#loc1) + %cst_0 = arith.constant 9.99999997E-7 : f32 loc(#loc2) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc12) + %x = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc13) + %x_1 = tt.addptr %x, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc13) + %x_2 = tt.load %x_1 : tensor<64x!tt.ptr> loc(#loc14) + %s = tt.load %s_ptr : !tt.ptr loc(#loc15) + %s_3 = arith.addf %s, %cst_0 : f32 loc(#loc15) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc7) + %1 = tt.addptr %0, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc7) + %2 = tt.splat %s_3 : f32 -> tensor<64xf32> loc(#loc1) + %3 = arith.mulf %x_2, %2 : tensor<64xf32> loc(#loc1) + %4 = arith.addf %3, %cst : tensor<64xf32> loc(#loc1) + tt.store %1, %4 : tensor<64x!tt.ptr> loc(#loc8) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":46:30) +#loc2 = loc(unknown) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":43:12) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":44:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":44:9) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":45:9) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":46:14) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":46:5) +#loc12 = loc("offs"(#loc3)) +#loc13 = loc("x"(#loc4)) +#loc14 = loc("x"(#loc5)) +#loc15 = loc("s"(#loc6)) diff --git a/tests/golden/ir/ttir_3.8/kernel_unicode_msgs.ttir b/tests/golden/ir/ttir_3.8/kernel_unicode_msgs.ttir new file mode 100644 index 000000000..ed4c67da2 --- /dev/null +++ b/tests/golden/ir/ttir_3.8/kernel_unicode_msgs.ttir @@ -0,0 +1,29 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":50:1) +#loc9 = loc("x_ptr"(#loc)) +module { + tt.func public @unicode_msgs(%x_ptr: !tt.ptr loc("x_ptr"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<0.000000e+00> : tensor<64xf32> loc(#loc1) + %cst_0 = arith.constant dense<1.000000e+00> : tensor<64xf32> loc(#loc2) + %offs = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc10) + %x = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc11) + %x_1 = tt.addptr %x, %offs : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc11) + %x_2 = tt.load %x_1 : tensor<64x!tt.ptr> loc(#loc12) + %0 = arith.cmpf ogt, %x_2, %cst : tensor<64xf32> loc(#loc1) + tt.assert %0, "\E9\94\99\E8\AF\AF: \CF\80 must be > 0" : tensor<64xi1> loc(#loc6) + tt.print " x=: " {hex = false, isSigned = array} : %x_2 : tensor<64xf32> loc(#loc7) + %1 = arith.addf %x_2, %cst_0 : tensor<64xf32> loc(#loc2) + tt.store %x_1, %1 : tensor<64x!tt.ptr> loc(#loc8) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":54:22) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":56:28) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":52:12) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":53:17) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":53:9) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":54:5) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":55:5) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":56:5) +#loc10 = loc("offs"(#loc3)) +#loc11 = loc("x"(#loc4)) +#loc12 = loc("x"(#loc5)) diff --git a/tests/golden/ir/ttir_3.8/spike_misc.ttir b/tests/golden/ir/ttir_3.8/spike_misc.ttir new file mode 100644 index 000000000..5a9c0ee94 --- /dev/null +++ b/tests/golden/ir/ttir_3.8/spike_misc.ttir @@ -0,0 +1,51 @@ +#loc = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":205:1) +#loc15 = loc("x_ptr"(#loc)) +#loc16 = loc("out_ptr"(#loc)) +#loc17 = loc("n"(#loc)) +module { + tt.func public @misc(%x_ptr: !tt.ptr loc("x_ptr"(#loc)), %out_ptr: !tt.ptr loc("out_ptr"(#loc)), %n: i32 loc("n"(#loc))) attributes {noinline = false} { + %cst = arith.constant dense<0.000000e+00> : tensor<64xf32> loc(#loc1) + %v = arith.constant dense<-1.500000e+00> : tensor<64xf32> loc(#loc18) + %c64_i32 = arith.constant 64 : i32 loc(#loc1) + %pid = tt.get_program_id x : i32 loc(#loc19) + %npg = tt.get_num_programs x : i32 loc(#loc20) + %offs = arith.muli %pid, %c64_i32 : i32 loc(#loc21) + %offs_0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> loc(#loc22) + %offs_1 = tt.splat %offs : i32 -> tensor<64xi32> loc(#loc21) + %offs_2 = arith.addi %offs_1, %offs_0 {tt.contiguity = dense<64> : tensor<1xi32>, tt.divisibility = dense<64> : tensor<1xi32>} : tensor<64xi32> loc(#loc21) + %v_3 = tt.splat %n : i32 -> tensor<64xi32> loc(#loc23) + %v_4 = arith.cmpi slt, %offs_2, %v_3 : tensor<64xi32> loc(#loc23) + %v_5 = tt.splat %x_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc24) + %v_6 = tt.addptr %v_5, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc24) + %v_7 = tt.load %v_6, %v_4, %v : tensor<64x!tt.ptr> loc(#loc18) + ttg.barrier all loc(#loc9) + tt.print " v{x} loc(: " {hex = false, isSigned = array} : %npg : i32 loc(#loc10) + %0 = tt.splat %out_ptr : !tt.ptr -> tensor<64x!tt.ptr> loc(#loc11) + %1 = tt.addptr %0, %offs_2 : tensor<64x!tt.ptr>, tensor<64xi32> loc(#loc11) + %2 = arith.cmpf ogt, %v_7, %cst : tensor<64xf32> loc(#loc12) + %3 = arith.select %2, %v_7, %cst : tensor<64xi1>, tensor<64xf32> loc(#loc13) + tt.store %1, %3, %v_4 : tensor<64x!tt.ptr> loc(#loc14) + tt.return loc(#loc) + } loc(#loc) +} loc(#loc) +#loc1 = loc(unknown) +#loc2 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":211:9) +#loc3 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":207:11) +#loc4 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":208:11) +#loc5 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":209:12) +#loc6 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":209:26) +#loc7 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":211:36) +#loc8 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":211:17) +#loc9 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":212:5) +#loc10 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":213:5) +#loc11 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":214:14) +#loc12 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":214:39) +#loc13 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":214:30) +#loc14 = loc("/home/hwu27/workspace/triton-viz-ir-mode/tests/golden/ir/generate_ttir.py":214:5) +#loc18 = loc("v"(#loc2)) +#loc19 = loc("pid"(#loc3)) +#loc20 = loc("npg"(#loc4)) +#loc21 = loc("offs"(#loc5)) +#loc22 = loc("offs"(#loc6)) +#loc23 = loc("v"(#loc7)) +#loc24 = loc("v"(#loc8)) diff --git a/tests/unit/ir/__init__.py b/tests/unit/ir/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/ir/_goldens.py b/tests/unit/ir/_goldens.py new file mode 100644 index 000000000..e638b3fb8 --- /dev/null +++ b/tests/unit/ir/_goldens.py @@ -0,0 +1,89 @@ +"""The TTIR goldens each Triton release reads, and the pins of each text. + +The base goldens (``tests/golden/ir/ttir/`` and ``reader_ttir/``) were +printed by Triton 3.6, the base release; every release reads them. A later +release's own printing of a golden lives in ``ttir_/`` / +``reader_ttir_/`` (generated by ``generate_ttir.py`` / +``generate_reader_ttir.py`` under that release) and shadows the base copy +of the same name under that release. A text's pins are those of the +release that printed it: ``expected.json`` for the base goldens, +``expected_.json`` for a release's own. + +Only reads ``triton.__version__`` at import: importing this module never +fails on a release without a walk-layer table (the IR tests then skip, +D29, or fail closed under the override). + +Not a test module: pytest imports it (python_files = *.py) and finds nothing. +""" + +from __future__ import annotations + +import json +from pathlib import Path + +from tilelens.ir import _mlir_walk as W + +REPO = Path(__file__).resolve().parents[3] +GOLDEN = REPO / "tests" / "golden" / "ir" +BASE_RELEASE = "3.6" # the release that printed ttir/ and reader_ttir/ +RELEASE = W.triton_release()[0] # the installed Triton's + +# Base goldens a release's parser rejects, because its printer spells a +# construct differently: name -> a fragment of the parser's diagnostic. +# Each is shadowed by that release's own printing (generate_ttir.py's +# RESPELLED), so the release keeps the coverage. +BASE_UNPARSABLE: dict[str, dict[str, str]] = { + # `!tt.tensordesc>` is `!tt.tensordesc<32x32xf16>` on 3.8 + "3.8": { + name: "tensor descriptors must not wrap tensor types" + for name in ( + "adv_descs.ttir", + "golden_matmul_tma_s1_sm90.ttir", + "golden_matmul_tma_s3_sm90.ttir", + "golden_matmul_tma_ws_s3_sm90.ttir", + ) + }, +} + + +def own_dir(kind: str, release: str) -> Path: + """The directory of ``release``'s own goldens of ``kind`` ("ttir" or + "reader_ttir"); the base release's are the base directory.""" + return GOLDEN / (kind if release == BASE_RELEASE else f"{kind}_{release}") + + +def releases_with_goldens(kind: str = "ttir") -> list[str]: + """Every release with goldens of ``kind``, the base release first.""" + own = sorted( + p.name[len(kind) + 1 :] for p in GOLDEN.glob(f"{kind}_*") if p.is_dir() + ) + return [BASE_RELEASE, *own] + + +def texts(kind: str, release: str = RELEASE) -> dict[str, Path]: + """name -> the golden ``release`` reads: its own printing where it has + one, else the base golden.""" + got = {p.name: p for p in sorted(own_dir(kind, BASE_RELEASE).glob("*.ttir"))} + if release != BASE_RELEASE: + got.update({p.name: p for p in sorted(own_dir(kind, release).glob("*.ttir"))}) + return dict(sorted(got.items())) + + +def printed_by(path: Path) -> str: + """The release that printed a golden.""" + for kind in ("reader_ttir", "ttir"): + if path.parent.name.startswith(f"{kind}_"): + return path.parent.name[len(kind) + 1 :] + return BASE_RELEASE + + +def pins_path(release: str) -> Path: + """The pins of the texts ``release`` printed (``ttir/`` goldens).""" + return GOLDEN / ( + "expected.json" if release == BASE_RELEASE else f"expected_{release}.json" + ) + + +def pins(release: str) -> dict: + path = pins_path(release) + return json.loads(path.read_text(encoding="utf-8")) if path.exists() else {} diff --git a/tests/unit/ir/_oracle_ttir_reader_361.py b/tests/unit/ir/_oracle_ttir_reader_361.py new file mode 100644 index 000000000..156d2315d --- /dev/null +++ b/tests/unit/ir/_oracle_ttir_reader_361.py @@ -0,0 +1,2134 @@ +# Differential oracle for tests/unit/ir/test_ttir_reader.py -- NOT shipped code. +# Verbatim copy of #361's regex TTIR reader, triton_viz/clients/common/ttir_reader.py +# at origin/race-detector-z3-demo 62c6d7c (sha256 3a82e07d3dc2aae16f2d3098894e411c +# 04779ae61c2f787eb2cb15278ff7a4df). Nothing below this header is edited: keep it +# byte-identical so the oracle stays #361's behaviour. It defines no test_* names. +"""Textual TTIR reader shared by the compiled-mode clients. + +Parses the pre-optimization Triton IR (TTIR) of one kernel specialization +into an ``AccessGraph``: the kernel's function arguments, every global +memory access (``tt.load`` / ``tt.store`` / ``tt.atomic_rmw`` / +``tt.atomic_cas``) as an *element offset* expression +relative to a base pointer argument, the mask guarding it, and the loop +structure. Scalar arguments (``n_elements``, ``M``, strides, ...) stay +symbolic (``Param`` nodes) and are substituted with concrete launch values +later; ``tl.constexpr`` values are already folded into TTIR constants. + +Why TTIR (not TTGIR): element addressing is cleanest here, before +layouts/pipelining add noise, and TTIR has no indirect loads unless the +kernel itself gathers — the data-dependent case, marked with ``DataDep``. + +This module is mechanism-only: it parses and flags (``DataDep`` markers, +``guarded`` accesses, ``UnsupportedTTIR``); what to do about a flagged or +unsupported kernel — report it, fall back to the interpreter, ... — is the +policy of each client that consumes the graph (sanitizer OOB checking, +race-detector global-memory front-end). + +Address model: ``tt.addptr(base, off)`` accumulates an ELEMENT offset; the +byte address is ``base.data_ptr() + offset * elem_size``. An access is OOB +iff, for some program id / arange lane / loop iteration with its mask true, +the element offset escapes ``[0, numel)`` of its base tensor. +""" + +from __future__ import annotations + +import re +from dataclasses import dataclass, field, replace + + +class UnsupportedTTIR(Exception): + """Raised for constructs outside the compiled-mode v1 model + (indirect/data-dependent addressing, block pointers, nested loops, ...). + The client converts this into an ``unsupported`` status (empty records) — + never a silent wrong verdict. v1 does not auto-fall back to interpreted + checking; run the eager ``Sanitizer()`` to check an unsupported kernel. + + ``kind`` is a stable, machine-readable class of the limitation — the + hybrid tier selector routes on it (an "indirect-address" kernel goes to + the interpreter front-end) and the evaluation reports its distribution: + "indirect-address" | "data-dependent-bound" | "nested-loop" | + "out-of-vocabulary" | "control-flow" | "block-pointer" | + "unmodelable-condition" | "data-dependent-mask" | + "cas-value" | "spin-shape" | "other". + + "spin-shape" (spec C1.1): an ``scf.while`` that is not the recognized + await form — the reason string names exactly which clause broke + (carried values, extra memory ops, non-comparison condition, ...). + """ + + def __init__(self, msg: str, kind: str = "other") -> None: + super().__init__(msg) + self.kind = kind + + +# ─────────────────────────── address-expression terms ─────────────────────────── +# A small lazily-evaluated tree. Leaves that are only known at launch time +# (scalar kernel args) are Param nodes; pid / arange / loop variables become +# free Z3 variables with range constraints in the OOB query. + + +@dataclass(frozen=True) +class Const: + value: int + + +@dataclass(frozen=True) +class Pid: + axis: int # 0=x, 1=y, 2=z + + +@dataclass(frozen=True) +class NumPrograms: + """``tt.get_num_programs axis`` — the launch grid size along ``axis``. + Uniform across program instances, but it PARAMETERIZES the kernel's + behavior by the grid (last-block gates compare an atomic observation + against it), so parsing one records the axis in ``pid_axes``: the + verdict must stay symbolic along that dim. The race encoder lowers it + to the SAME ``grid_`` variable the solver's symbolic grid uses.""" + + axis: int + + +@dataclass(frozen=True) +class Arange: + ssa: str # unique per make_range site + start: int + end: int + # Which tensor dimension this lane index varies along. -1 = 1D / not yet + # placed; 0/1 set by expand_dims. A single make_range reused for both the + # row and column of a 2D tile (triton does this) must become TWO + # independent variables — keyed by (ssa, dim) — or the modeled footprint + # collapses to the diagonal (the same collapse bug fixed in dynamic mode). + dim: int = -1 + + +@dataclass(frozen=True) +class Param: + name: str # scalar kernel argument, substituted per launch + + +@dataclass(frozen=True) +class IterArgOffset: + """The element-offset contribution of a loop-carried pointer at the + current iteration: ``offset0 + k * delta`` (resolved from the graph's + loop info at eval time).""" + + arg_id: int + + +@dataclass(frozen=True) +class LoopVar: + """The scf.for induction variable; a free variable in [lower, upper) + in the OOB query (e.g. it appears in masks like ``K - k*BLOCK_K``).""" + + loop_ssa: str + + +@dataclass(frozen=True) +class Bin: + op: str # + - * // % min max (// and % truncate toward zero: divsi/remsi) + a: "Term" + b: "Term" + + +@dataclass(frozen=True) +class Cmp: + pred: str # slt/sle/sgt/sge/eq/ne + a: "Term" + b: "Term" + + +@dataclass(frozen=True) +class BoolBin: + op: str # and / or + a: "Term" + b: "Term" + + +@dataclass(frozen=True) +class Select: + cond: "Term" + t: "Term" + f: "Term" + + +@dataclass(frozen=True) +class Not: + """Boolean negation — the path condition of an scf.if else-region.""" + + a: "Term" + + +# Sentinel for a value loaded from memory (tt.load result) or computed from +# loaded data (arith.*f, tt.dot, ...). If one ever reaches an address or mask +# it means data-dependent addressing → unsupported. +@dataclass(frozen=True) +class DataDep: + why: str = "value derived from loaded data" + # For a boolean ``and`` with one unmodelable operand: the modelable + # conjunct(s). The true value implies ``keep``, so a mask position may + # use ``keep`` as a sound over-approximation instead of dropping the + # whole mask (multipath only; the access still counts as widened). + keep: "Term | None" = None + + +@dataclass(frozen=True) +class Loaded: + """The VALUE of an integer ``tt.load`` (Route 2, the L2 reader mode): + lane-wise ``snapshot[base][offset]`` over the launch's pre-launch + contents of the source tensor, ``other`` (or a free value) on masked + lanes. Bound only under ``parse_ttir(multipath=True)``; single-path + keeps :class:`DataDep` for every loaded value. The encoder turns it + into an SMT-array Select over the tensor's snapshot and marks the + verdict content-qualified; it refuses by name when the launch carries + no snapshot for the source (float, too large, non-contiguous) or when + the kernel writes the source tensor (the read-only-source premise the + interpreter frontend enforces by fail-stop). Consumers that walk terms + descend into ``offset``, ``mask`` and ``other``.""" + + access_index: int + base_param: str + offset: "Term" + mask: "Term | None" + other: "Term | None" + + +@dataclass(frozen=True) +class Observed: + """The OLD value observed by the atomic at ``graph.accesses[access_index]`` + (spec part B): a fresh per-program-instance symbol, NOT a function of + other leaves. The reader binds an INTEGER-typed ``tt.atomic_rmw`` / + ``tt.atomic_cas`` result to this instead of ``DataDep`` so downstream + masks and branch conditions stay modelable; float-typed atomic results + keep the DataDep fallback (the value model is Int-sort only). + + Consumer policy (mechanism lives here, policy with each client): + * race-detector global encoder: interns one Z3 var per index, ties it + to the record's ``old_value`` (rf-justified) when the observation is + modelable, and fails closed on address uses of unmodeled ones; + * sanitizer OOB: a free variable — sound widening for proofs, with + the mask_dropped-style witness abstention; + * differential (C3): no concrete value exists — the access is + excluded SYMMETRICALLY from both sides of the diff. + """ + + access_index: int + + +_TERM_CHILDREN = ("a", "b", "cond", "t", "f", "offset", "mask", "other") + + +def mentions_observed(term: object) -> bool: + """True when ``term`` contains an :class:`Observed` leaf.""" + if isinstance(term, Observed): + return True + for attr in _TERM_CHILDREN: + sub = getattr(term, attr, None) + if sub is not None and mentions_observed(sub): + return True + return False + + +def observed_indices(term: object) -> set[int]: + """Access indices of every :class:`Observed` leaf in ``term``.""" + out: set[int] = set() + if isinstance(term, Observed): + out.add(term.access_index) + return out + for attr in _TERM_CHILDREN: + sub = getattr(term, attr, None) + if sub is not None: + out |= observed_indices(sub) + return out + + +# DataDep is also the generic unknown-value top (unresolved SSA, loop +# accumulators, unmodeled ops, ...). Only these ``why`` prefixes mean the +# value truly derives from MEMORY CONTENTS — the per-term policy classifies +# just those as indirection (the interpreter-front-end route); the rest are +# modeling gaps and keep the default kind. +_MEMORY_WHYS = ( + "loaded value", + "atomic result", + "arith over loaded data", + "cmpi over loaded data", + "select over loaded data", + "bool op over loaded data", + "float/reduction value", +) + + +def _from_memory(v: object) -> bool: + return isinstance(v, DataDep) and v.why.startswith(_MEMORY_WHYS) + + +Term = ( + Const + | Pid + | NumPrograms + | Arange + | Param + | IterArgOffset + | LoopVar + | Bin + | Cmp + | BoolBin + | Select + | Not + | DataDep + | Observed + | Loaded +) + + +def mentions_loaded(term: object) -> bool: + """True when ``term`` contains a :class:`Loaded` leaf.""" + if isinstance(term, Loaded): + return True + for attr in _TERM_CHILDREN: + sub = getattr(term, attr, None) + if sub is not None and mentions_loaded(sub): + return True + return False + + +def loaded_leaves(term: object, out: "list[Loaded] | None" = None) -> "list[Loaded]": + """Every :class:`Loaded` leaf of ``term`` (outer before inner).""" + if out is None: + out = [] + if isinstance(term, Loaded): + out.append(term) + for attr in _TERM_CHILDREN: + sub = getattr(term, attr, None) + if sub is not None: + loaded_leaves(sub, out) + return out + + +@dataclass(frozen=True) +class PtrValue: + """A pointer-typed SSA value: base argument + accumulated element + offset (a single lane's offset; arange/loop free vars cover all lanes + and iterations in the query).""" + + base_param: str + offset: Term + + +# ─────────────────────────── graph structures ─────────────────────────── + + +@dataclass(frozen=True) +class FuncArg: + name: str + is_ptr: bool + elem_bits: int # for ptr args: pointee width; 0 for scalars + # Float-typed pointee (f*/bf*): atomic results on it stay DataDep — the + # Int-sort observation model must not carry float values (spec B.5). + elem_float: bool = False + # Dependency order (D2/D3): indices (into AccessGraph.accesses) of the + # load / atomic accesses whose result this access's value, mask, or + # compare operand consumes through element-wise ops only. + deps: tuple[int, ...] = () + + +@dataclass(frozen=True) +class SourceLoc: + file: str + line: int + col: int + + +@dataclass(frozen=True) +class AtomicInfo: + """Atomicity metadata for ``tt.atomic_rmw`` / ``tt.atomic_cas`` accesses.""" + + rmw_op: str | None # "fadd", "max", "exch", ... ; None for CAS + sem: str # memory semantic: "acq_rel", "relaxed", ... + scope: str # sync scope: "gpu", "cta", "sys" + + +@dataclass(frozen=True) +class AccessEvent: + kind: str # "load" | "store" | "atomic_rmw" | "atomic_cas" + base_param: str + offset: Term + mask: Term | None # None = unconditional access + elem_bits: int + loc: SourceLoc | None + line_no: int + # True when some enclosing scf.if condition could NOT be modeled (it + # derives from loaded data). The access is then checked as if + # unconditional: UNSAT stays a sound proof, but a SAT model may sit in a + # branch the launch never takes, so it must not be reported as a witness + # (check_graph turns it into ``unsupported``). Modeled conditions ride + # in ``path`` instead and do not set this flag. + guarded: bool = False + # Conjunction of the MODELED enclosing branch conditions, with + # else-regions negated (Not). The access executes iff path ∧ mask, so a + # SAT model under both constraints is a real, reachable witness. + path: Term | None = None + # True when the access sits inside the scf.for body: it executes once + # per iteration — and NOT AT ALL when the launch's trip count is zero, + # which consumers must model (a zero-trip loop has no footprint). + in_loop: bool = False + # Present iff kind is atomic_*: an atomic is a read AND a write of its + # footprint (RMW), which is what is_read/is_write encode for consumers + # that build read/write event pairs (the race detector front-end). + atomic: AtomicInfo | None = None + # True when the printed mask operand derived from loaded data and was + # over-approximated as FREE (mask=None): dropping a constraint only + # widens the modeled footprint, so UNSAT stays a sound proof — but a SAT + # model may pick a lane the real mask disables, so it follows the same + # uncertainty discipline as ``guarded`` (never reported as a witness). + mask_dropped: bool = False + # For atomics: the printed VALUE operand (tt.atomic_rmw val / + # tt.atomic_cas val) as a Term, or None when it is not modelable + # (loaded data). The race encoder models the RMW write part from it. + atomic_val: "Term | None" = None + # For tt.atomic_cas only: the compare operand. + atomic_cmp: "Term | None" = None + # Float-typed pointee: the observation is never modeled (spec B.5). + elem_float: bool = False + # The await abstraction (spec C1): True when this access is the single + # kept read of a recognized scf.while spin loop. ``exit_pred`` is the + # loop's EXIT predicate over Observed(this access) — asserted on the + # event, justified by termination (in any terminating execution the + # final iteration's read observed the exit value). Dropped iterations + # lose no conflict pairs because the race encoder emits a PRE-EXIT + # REPRESENTATIVE alongside the poll: a value-model-free twin carrying + # each failed iteration's footprint and modes with a subset of its + # happens-before edges (global_records._pre_exit_representative). + # Verdicts over await-bearing kernels are therefore conditional on + # termination (surfaced as ``assumes_termination``). + awaited: bool = False + exit_pred: "Term | None" = None + # The enclosing scf.for loops (LoopInfo.loop_ssa), outermost first; + # ``in_loop == bool(loops)``. A multipath graph (see parse_ttir) may + # nest several; a single-path graph has at most one. + loops: tuple[str, ...] = () + + @property + def is_read(self) -> bool: + return self.kind != "store" + + @property + def is_write(self) -> bool: + return self.kind != "load" + + +@dataclass(frozen=True) +class IterArgInfo: + arg_id: int + base_param: str + offset0: Term + delta: Term # per-iteration element advance + # The scf.for this iter_arg belongs to (its LoopInfo.loop_ssa). Empty + # for graphs built before multi-loop capture existed: consumers then + # resolve it against the graph's single loop. + loop_ssa: str = "" + + +@dataclass(frozen=True) +class LoopInfo: + loop_ssa: str + induction_var: str + lower: Term + upper: Term + step: Term + + +@dataclass(frozen=True) +class LoopTokenConflict: + """A cross-iteration pair whose non-aliasing remains to be established. + + Access indices refer to the original graph, including a write paired + with itself. The reader cannot discharge these pairs from formal names: + encoding must prove the applicable allocation non-aliasing premise. + """ + + loop_ssa: str + first: int + second: int + + +@dataclass +class AccessGraph: + kernel_name: str + func_args: list[FuncArg] + accesses: list[AccessEvent] + loop: LoopInfo | None + iter_args: dict[int, IterArgInfo] = field(default_factory=dict) + # Every pid axis with a parsed tt.get_program_id — recorded at PARSE + # time, before any DataDep swallowing. Consumers deciding grid coverage + # must use THIS set, not the axes that happen to survive into modeled + # address/mask terms: a pid read into a stored value, a dropped mask, or + # an unmodeled branch condition still distinguishes the blocks' behavior. + pid_axes: set[int] = field(default_factory=set) + # Tile-level fences (``gpu.barrier``, the TTIR lowering of + # ``tl.debug_barrier``), as program_seq positions: a fence recorded + # after k accesses sits at k - 0.5, strictly between access k-1 and + # access k in the encoder's dense integer seq (paper + # design-fence-order.md, option A). + fences: list[float] = field(default_factory=list) + # Which reader produced the graph. cuTile exposes explicit token + # reachability below; a token edge is not a full fence cut. + frontend: str = "triton" + # Every scf.for of the kernel in textual (opening) order, outer before + # inner. ``loop`` above stays the single loop when there is exactly one + # (the pre-multipath consumers read it) and is None otherwise. + loops: list[LoopInfo] = field(default_factory=list) + # True when parsed with ``multipath=True`` (Route 3): block path + # predicates for the cf.* graph and multiple loops are modeled; the + # single-loop consumers (sanitizer OOB, differential) must not be fed + # such a graph. + multipath: bool = False + # Number of cf.* blocks modeled (0 for a structured kernel). + cf_blocks: int = 0 + # Transitive, operation-level token order, keyed by access indices. + # None selects the Triton fence/dependency discipline; an empty dict + # means explicit token semantics with no ordered pairs. A None guard + # is unconditional; otherwise it is a scalar per-instance condition. + # Memory masks do not gate token propagation through an operation. + token_order: dict[tuple[int, int], Term | None] | None = None + # cuTile pairs not serialized across this loop's iteration boundary. + # Consumers must discharge every pair before using shared iterators in + # same-instance queries; distinct formal names alone are insufficient. + loop_token_conflicts: list[LoopTokenConflict] = field(default_factory=list) + # The address reader treats integer width casts as transparent. CAS + # success needs exact values: ordinary CAS encoding may consume this + # conservative graph-wide fact to reject potentially changing casts. + has_value_changing_integer_casts: bool = False + # True when an address term was rewritten using a scalar param's + # CAPTURED value (the cuTile reader's exact bitwise lowering: a shift + # count or mask that is only known at the launch). The rewrite is + # exact for THIS launch's parameters and says nothing about others, + # so the tier selector must not attempt T0 on such a graph. + param_pinned: bool = False + + def arg(self, name: str) -> FuncArg | None: + for a in self.func_args: + if a.name == name: + return a + return None + + +# ─────────────────────────── regexes ─────────────────────────── + +# `#N` is a result index into a multi-result op (`%acc#2` = third result of +# `%acc:3 = scf.for ...`). It must be part of the operand token or lines like +# `tt.store %ptrs, %acc#2, %mask` fail to match the store regex and fail +# closed even though the stored VALUE plays no part in address math. The env +# never defines `%x#N` names, so val() resolves them to DataDep("unresolved +# SSA") — sound in every consuming position (mask → dropped and flagged +# ``mask_dropped``, i.e. proof-only; addptr/ptr → unsupported). +# `-` is part of the token class: negative constants print as `%c-1_i32`, +# and truncating at the hyphen made every kernel with one fail closed. +_SSA = r"%[-\w.]+(?:#\d+)?" +_DTYPE_BITS = { + "f64": 64, "f32": 32, "f16": 16, "bf16": 16, "f8": 8, + "i64": 64, "i32": 32, "i16": 16, "i8": 8, "i1": 1, + "u64": 64, "u32": 32, + # MLIR spells the fp8 families out (torchao's quant kernels take + # fp8 pointers); all are one byte wide + "f8E4M3FN": 8, "f8E5M2": 8, "f8E4M3FNUZ": 8, "f8E5M2FNUZ": 8, + "f8E4M3B11FNUZ": 8, "f8E8M0FNU": 8, +} # fmt: skip + +_RE_LOC_FILE = re.compile(r'^(#loc\d*) = loc\("([^"]+)":(\d+):(\d+)\)') +_RE_LOC_NAME = re.compile(r'^(#loc\d*) = loc\("[^"]+"\((#loc\d*)\)\)') +_RE_LOC_CALLSITE = re.compile(r"^(#loc\d*) = loc\(callsite\((#loc\d*) at (#loc\d*)\)\)") +_RE_LOC_TRAILER = re.compile(r"loc\((#loc\d*|#loc)\)\s*$") +_RE_FUNC = re.compile(r"tt\.func\s+\w+\s+@(\w+)\((.*)\)\s*attributes") +_RE_RESULT = re.compile(rf"^({_SSA})(?::\d+)?\s*=\s*(.*)$") +_RE_GET_PID = re.compile(r"^tt\.get_program_id (\w+)") +_RE_GET_NPROG = re.compile(r"^tt\.get_num_programs (\w+)") +_RE_MAKE_RANGE = re.compile( + r"^tt\.make_range \{end = (-?\d+) : i32, start = (-?\d+) : i32\}" +) +_RE_CONST_INT = re.compile(r"^arith\.constant (-?\d+) : i\d+") +_RE_CONST_DENSE = re.compile(r"^arith\.constant dense<(-?\d+)> : tensor") +_RE_CONST_DENSE_BOOL = re.compile(r"^arith\.constant dense<(true|false)> : tensor") +_RE_CONST_BOOL = re.compile(r"^arith\.constant (true|false)\b") +_RE_SPLAT = re.compile(rf"^tt\.splat ({_SSA}) : ([^-]+)->") +_RE_EXPAND = re.compile(rf"^tt\.expand_dims ({_SSA}) \{{axis = (\d+)") +_RE_BROADCAST = re.compile(rf"^tt\.broadcast ({_SSA})") +_RE_ADDPTR = re.compile(rf"^tt\.addptr ({_SSA}), ({_SSA})") +_RE_BIN = re.compile( + rf"^arith\.(muli|addi|subi|divsi|remsi|minsi|maxsi) ({_SSA}), ({_SSA})" +) +_RE_CMPI = re.compile(rf"^arith\.cmpi (\w+), ({_SSA}), ({_SSA})") +# andi/ori operate on any integer width; only the i1 form is boolean logic. +# The printed result type distinguishes them (": tensor<..xi1>" / ": i1"). +_RE_BOOLBIN = re.compile(rf"^arith\.(andi|ori) ({_SSA}), ({_SSA})\s*:\s*(\S+)") +_RE_SELECT = re.compile(rf"^arith\.select ({_SSA}), ({_SSA}), ({_SSA})") +_RE_EXT = re.compile(rf"^arith\.(extsi|trunci|extui) ({_SSA})") +# Only the known custom, three-operand dot spelling has a positional C +# contribution. Do not infer an accumulator slot for generic syntax, +# dot_scaled, extra operands, or unknown attributes. The optional integer +# accuracy attribute is an attr-dict entry, not a fourth SSA operand. +_RE_DOT = re.compile( + rf"^tt\.dot\s+({_SSA}),\s*({_SSA}),\s*({_SSA})" + r"(?:,\s*inputPrecision\s*=\s*(?:ieee|tf32|tf32x3|bf16x3|bf16x6))?" + r"(?:\s*\{\s*maxNumImpreciseAcc\s*=\s*\d+\s*:\s*i32\s*\})?\s*:" +) +# Trailing attributes print in TWO spellings: a dict (`{isVolatile = +# true}` for volatile spin reads) or bare assignments (`cacheModifier = +# ca` — liger's cache-hinted loads); both are irrelevant to the footprint. +_RE_LOAD = re.compile( + rf"^tt\.load ({_SSA})((?:, {_SSA})*)\s*" + rf"(?:\{{[^}}]*\}})?(?:\s+\w+\s*=\s*\w+)*\s*(?::|loc|$)" +) +_RE_STORE = re.compile( + rf"^tt\.store ({_SSA}), ({_SSA})((?:, {_SSA})*)\s*" + rf"(?:\{{[^}}]*\}})?(?:\s+\w+\s*=\s*\w+)*\s*(?::|loc|$)" +) +# Atomic RMW prints (op, sem, scope, ptr, val, mask); an unmasked tl.atomic_* +# still carries a mask operand (a dense constant), so the group is +# always present. CAS prints (sem, scope, ptr, cmp, val) — no mask exists. +_RE_ATOMIC_RMW = re.compile( + rf"^tt\.atomic_rmw (\w+), (\w+), (\w+), ({_SSA}), ({_SSA}), ({_SSA})\s*(?::|loc|$)" +) +_RE_ATOMIC_CAS = re.compile( + rf"^tt\.atomic_cas (\w+), (\w+), ({_SSA}), ({_SSA}), ({_SSA})\s*(?::|loc|$)" +) +_RE_PTR_ELEM = re.compile(r"!tt\.ptr<(\w+)>") +_RE_SCF_FOR = re.compile( + rf"^scf\.for ({_SSA}) = ({_SSA}) to ({_SSA}) step ({_SSA})" + # iter_args + "-> (types)" appear only when the loop yields values; a + # pure-side-effect loop (e.g. a store loop, no accumulator) ends at the + # ": i32 {" type annotation with no arrow. Match both, or the loop is + # missed and its induction var leaks as an unbound (data-dependent) SSA. + rf"(?: iter_args\((.*?)\))?\s*(?:->|:)" +) +_RE_SCF_YIELD = re.compile(r"^scf\.yield (.*?)\s*:") +_RE_SCF_IF = re.compile(rf"^scf\.if ({_SSA})") +# The await shape (C1.1): only the argument-free, result-free spin form is +# accepted; anything carrying values is refused as "spin-shape". +_RE_SCF_WHILE_SPIN = re.compile(r"^scf\.while\s*:\s*\(\)\s*->\s*\(\)\s*\{") +_RE_SCF_CONDITION = re.compile(rf"^scf\.condition\(({_SSA})\)") +# Unstructured control flow (Route 3, multipath): Triton lowers an ``if`` +# that contains a ``return`` through basic blocks instead of scf.if +# (code_generator.visit_if_top_level). Block labels carry optional +# parameters (``^bb5(%3: i32 loc(unknown)):``), branch targets optional +# operands (``^bb5(%c0_i32 : i32)``). +_RE_BLOCK_LABEL = re.compile(r"^\^(bb\d+)(?:\((.*)\))?:") +_RE_COND_BR = re.compile( + rf"^cf\.cond_br ({_SSA}), \^(bb\d+)(?:\((.*?)\))?, \^(bb\d+)(?:\((.*?)\))?" +) +_RE_BR = re.compile(r"^cf\.br \^(bb\d+)(?:\((.*?)\))?") + + +@dataclass +class _IfFrame: + """Walker state for one open scf.if region.""" + + cond: "Term | None" # modeled condition; None → accesses stay `guarded` + res: str | None # single-result SSA name ("%r"), if the if yields + branch: str = "then" + # Yield VALUES resolved at the yield line — then/else regions legally + # reuse the same SSA names, so resolving at close time would read the + # else-region's overwrites. + then_vals: "list[object] | None" = None + else_vals: "list[object] | None" = None + + +@dataclass +class _ForFrame: + """Walker state for one open scf.for region: its bounds, the iter_args + in declaration order, the arg ids of the pointer-typed ones, and the + body's yield operands (resolved at the loop's own scf.yield).""" + + ssa: str + ind: str + lower: "Term" + upper: "Term" + step: "Term" + order: int = 0 # opening order (outer loops open first) + iter_arg_ssa: list = field(default_factory=list) # (arg_ssa, init_ssa) + ptr_arg_ids: list = field(default_factory=list) + body_yields: list = field(default_factory=list) + + +@dataclass +class _Block: + """Walker state for one basic block of the unstructured cf.* graph + (multipath only). ``edges`` accumulates the incoming edges recorded at + the predecessors' branch lines as (path, exact, operand values): the + path is the predecessor's block predicate conjoined with the branch + condition (negated for the false target), ``exact`` is False when that + condition could not be modeled (loaded data) and the path is then only + the predecessor's predicate, an over-approximation. At the label the + block predicate is the disjunction of the incoming paths; a block with + an inexact edge is ``guarded`` (its accesses are widened) and binds its + parameters to DataDep (a Select over inexact paths would pick a wrong + VALUE, not a wider footprint).""" + + name: str + n_preds: int + edges: list = field(default_factory=list) + pred: "Term | None" = None + guarded: bool = False + # False when the label was reached before every predecessor's branch + # (a block placed before one of its predecessors): its predicate is + # unknown, so any access or branch inside refuses by name. Triton's + # lowering only does this for the shared return-only block. + resolved: bool = True + # True once the block's terminator (cf.br / cf.cond_br / tt.return) + # was seen: nothing after it is reachable. + terminated: bool = False + + +@dataclass +class _WhileFrame: + """Walker state for one open scf.while spin candidate (C1.1). + + The CONDITION region ("before") holds the awaited re-read plus its + address bookkeeping and ends at scf.condition; the BODY region ("do") + must be pure bookkeeping (scf.yield only). Any clause violation refuses + the kernel with kind="spin-shape" naming the clause.""" + + open_line: int + stage: str = "cond" # "cond" → "body" + n_accesses_before: int = 0 + cond_val: object | None = None # resolved AT the scf.condition line + + +def _branch_state(frames: list) -> "tuple[bool, Term | None, bool, tuple[str, ...]]": + """(guarded, path, in_loop, loops) for an access under the open frames: + ``guarded`` if any enclosing condition is unmodeled; ``path`` is the + conjunction of the modeled ones (else-regions negated); ``in_loop`` when + an scf.for body encloses the access; ``loops`` the enclosing scf.for + loops' ssa names, outermost first.""" + guarded = False + path: Term | None = None + loops: list[str] = [] + for f in frames: + if isinstance(f, _ForFrame): + loops.append(f.ssa) + continue + if not isinstance(f, _IfFrame): + continue + if f.cond is None: + guarded = True + continue + c: Term = f.cond if f.branch == "then" else Not(f.cond) + path = c if path is None else BoolBin("and", path, c) + return guarded, path, bool(loops), tuple(loops) + + +def _conj(a: "Term | None", b: "Term | None") -> "Term | None": + """None-aware conjunction (None = true).""" + if a is None: + return b + if b is None: + return a + return BoolBin("and", a, b) + + +def _disj(a: "Term | None", b: "Term | None") -> "Term | None": + """None-aware disjunction (None = true).""" + if a is None or b is None: + return None + return BoolBin("or", a, b) + + +def _arg_ssas(inner: "str | None") -> list[str]: + """SSA names of a block-operand list (``%a : i32, %b : i32``) or a + label parameter list (``%3: i32 loc(unknown), ...``).""" + if not inner: + return [] + out: list[str] = [] + for part in inner.split(","): + part = part.strip() + if part.startswith("%"): + out.append(part.split(":")[0].strip()) + return out + + +def _merge_block_param(edges: list, index: int) -> object: + """The value of block parameter ``index`` as a Select over the incoming + edges (every edge exact): the last edge is the fallback, each earlier + edge selects its value under its own path. Pointers merge when they + share a base (a Select over offsets); anything else stays DataDep, so + an address use fails closed.""" + vals = [e[2][index] if index < len(e[2]) else None for e in edges] + if not vals or any(v is None for v in vals): + return DataDep("block argument") + paths = [e[0] for e in edges] + if all(isinstance(v, PtrValue) for v in vals): + bases = {v.base_param for v in vals} # type: ignore[union-attr] + if len(bases) != 1: + return DataDep("block argument merging different bases") + sel: Term = vals[-1].offset # type: ignore[union-attr] + for epath, v in reversed(list(zip(paths, vals))[:-1]): + off: Term = v.offset # type: ignore[union-attr] + sel = off if epath is None else Select(epath, off, sel) + return PtrValue(vals[0].base_param, sel) # type: ignore[union-attr] + if any(isinstance(v, (DataDep, PtrValue)) for v in vals): + return DataDep("block argument") + term: Term = vals[-1] # type: ignore[assignment] + for epath, v in reversed(list(zip(paths, vals))[:-1]): + term = v if epath is None else Select(epath, v, term) # type: ignore[assignment,arg-type] + return term + + +def _prescan_blocks(lines: list[str]) -> dict[str, int]: + """Predecessor counts of every block of the function's cf.* graph, and + the acyclicity check. Block labels and cf.* terminators live only in + the function's own region (Triton never places them inside scf + regions: a ``return`` inside a loop is a compile error), so a flat + scan tracking the current label is exact. Raises (kind control-flow) + on a cycle: Triton never emits one, and a cyclic graph has no block + predicates.""" + cur = "bb0" + edges: dict[str, list[str]] = {} + seen_func = False + depth = 0 # anonymous op regions (tt.reduce combine blocks) are skipped + for raw in lines: + line = raw.strip() + if not seen_func: + seen_func = _RE_FUNC.search(line) is not None + continue + if line.endswith("({"): + depth += 1 + continue + if line.startswith("})") and depth: + depth -= 1 + continue + if depth: + continue + lm = _RE_BLOCK_LABEL.match(line) + if lm: + cur = lm.group(1) + edges.setdefault(cur, []) + continue + cm = _RE_COND_BR.match(line) + if cm: + edges.setdefault(cur, []).extend([cm.group(2), cm.group(4)]) + continue + bm = _RE_BR.match(line) + if bm: + edges.setdefault(cur, []).append(bm.group(1)) + n_preds: dict[str, int] = {} + for src, dsts in edges.items(): + for d in dsts: + n_preds[d] = n_preds.get(d, 0) + 1 + # DFS cycle check from the entry block + state: dict[str, int] = {} + + def visit(b: str, depth: int) -> None: + if depth > 10_000: + raise UnsupportedTTIR("cf graph too deep", kind="control-flow") + state[b] = 1 + for d in edges.get(b, []): + st = state.get(d, 0) + if st == 1: + raise UnsupportedTTIR( + f"cyclic cf.* control flow through ^{d} is unsupported", + kind="control-flow", + ) + if st == 0: + visit(d, depth + 1) + state[b] = 2 + + visit("bb0", 0) + return n_preds + + +def _elem_bits(type_str: str) -> int: + m = _RE_PTR_ELEM.search(type_str) + if m: + return _DTYPE_BITS.get(m.group(1), 0) + return 0 + + +def _elem_is_float(type_str: str) -> bool: + m = _RE_PTR_ELEM.search(type_str) + return m is not None and m.group(1).startswith(("f", "bf")) + + +def _split_ssa(text: str) -> list[str]: + return [t.strip() for t in text.split(",") if t.strip().startswith("%")] + + +class _LocTable: + def __init__(self) -> None: + self._file: dict[str, tuple[str, int, int]] = {} + self._alias: dict[str, str] = {} + + def add(self, line: str) -> bool: + m = _RE_LOC_FILE.match(line) + if m: + self._file[m.group(1)] = (m.group(2), int(m.group(3)), int(m.group(4))) + return True + m = _RE_LOC_NAME.match(line) + if m: + self._alias[m.group(1)] = m.group(2) + return True + m = _RE_LOC_CALLSITE.match(line) + if m: + # The memory operation belongs to the callee. The caller is + # useful stack context, not a substitute access source site. + # resolve() already bounds alias recursion and refuses unknowns. + self._alias[m.group(1)] = m.group(2) + return True + if line.startswith("#loc") and "= loc(" in line: + return True + return False + + def resolve(self, loc_id: str | None, _d: int = 0) -> SourceLoc | None: + if loc_id is None or _d > 8: + return None + if loc_id in self._file: + f, ln, col = self._file[loc_id] + return SourceLoc(f, ln, col) + if loc_id in self._alias: + return self.resolve(self._alias[loc_id], _d + 1) + return None + + +def parse_ttir(text: str, *, multipath: bool = False) -> AccessGraph: + """Parse one TTIR module into an AccessGraph. + + Raises :class:`UnsupportedTTIR` for indirect addressing, block pointers, + nested/while loops, or any op outside the v1 address vocabulary that + feeds a pointer. + + ``multipath=True`` (Route 3, the ladder's L2) lifts two structural + boundaries of the single-path model and is otherwise byte-identical: + * the unstructured ``cf.*`` graph Triton emits for an ``if`` that + contains a ``return`` (early-exit guards) gets block path + predicates: every access conjoins its block's predicate into + ``path`` exactly as it conjoins an enclosing scf.if condition, and + block parameters bind to a Select over the incoming edges' values; + * several ``scf.for`` loops (nested, sequential, under an scf.if or + a block predicate) each get their own induction variable + (``AccessGraph.loops``, ``AccessEvent.loops``). + Every new code path starts at a refusal site of the single-path model + (the ``cf.*`` raise, the second-loop raise), so a kernel without those + constructs is parsed identically in both modes. + """ + locs = _LocTable() + kernel_name = "" + func_args: list[FuncArg] = [] + # SSA name -> value: Term (int/bool), PtrValue, or DataDep + env: dict[str, object] = {} + accesses: list[AccessEvent] = [] + fences: list[float] = [] + # Dependency provenance: SSA name -> {access index it derives from: + # still position-preserving?}. Loads and atomics seed it; element-wise + # ops propagate it; position-changing ops clear the flag. The flag is + # kept PER SOURCE so a scalar operand that arrived through tt.splat + # does not strip the positional flag off the tile operand next to it. + prov: dict[str, dict[int, bool]] = {} + + def _prov_of( + ssa_names, position_preserving=True, *, simultaneous=False + ) -> dict[int, bool]: + merged: dict[int, bool] = {} + for name in ssa_names: + got = prov.get(name) + if not got: + continue + for idx, flag in got.items(): + positional = flag and position_preserving + if simultaneous: + # All operands contribute to this element. A positional + # path remains a dependency even if another path from + # the SAME load permutes elements (e.g. x + rotate(x)). + merged[idx] = merged.get(idx, False) or positional + else: + # Alternatives must not borrow an inactive arm's + # positional path. Keep their conservative intersection. + merged[idx] = merged.get(idx, True) and positional + return merged + + def _deps_of(*ssa_names) -> tuple[int, ...]: + return tuple(sorted(idx for idx, flag in _prov_of(ssa_names).items() if flag)) + + loop: LoopInfo | None = None + loops: list[tuple[int, LoopInfo]] = [] # (opening order, loop) + loops_opened = 0 + iter_args: dict[int, IterArgInfo] = {} + next_arg_id = 0 + # Block walk (multipath): the implicit entry block, the block table + # filled from the pre-scan on the first cf.* line, the current block. + entry = _Block("bb0", n_preds=0) + blocks: dict[str, _Block] = {} + n_preds: dict[str, int] | None = None + cur = entry + + lines = text.splitlines() + # This does not change generic address parsing. Ordinary CAS's exact + # value gate uses the fact, including a narrowing-then-widening chain + # that would otherwise look like a same-width atomic observation. + changing_integer_casts = False + has_cas = "tt.atomic_cas " in text + # Loc aliases live at the bottom; collect them in this same scan. + for raw in lines: + line = raw.strip() + locs.add(line) + if not has_cas: + continue + result = _RE_RESULT.match(line) + if result is None: + continue + body = result.group(2) + cast = _RE_EXT.match(body) + if cast is None: + continue + source = re.search(r":\s*(?:tensor<[^>]*x)?i(\d+)>?\s+to\s+", body) + source_bits = int(source.group(1)) if source else 0 + # Only bool-to-integer zero extension is needed by the admitted + # CAS shapes. Other signedness/width changes remain conservative. + value_preserving = cast.group(1) == "extui" and source_bits == 1 + changing_integer_casts |= not value_preserving + + def val(name: str) -> object: + v = env.get(name) + if v is None: + # Unknown SSA reaching an address/mask: be conservative. + return DataDep(f"unresolved SSA {name}") + return v + + def as_term(v: object, ctx: str) -> Term: + if isinstance(v, DataDep): + raise UnsupportedTTIR(f"{ctx}: data-dependent ({v.why})") + if isinstance(v, PtrValue): + raise UnsupportedTTIR(f"{ctx}: pointer used as integer") + return v # type: ignore[return-value] + + def parse_func_args(arg_text: str) -> None: + for m in re.finditer(r"(%[\w.]+): (!tt\.ptr<\w+>|i\d+|f\d+)", arg_text): + name, ty = m.group(1)[1:], m.group(2) + is_ptr = ty.startswith("!tt.ptr") + bits = _elem_bits(ty) if is_ptr else 0 + fa = FuncArg( + name=name, + is_ptr=is_ptr, + elem_bits=bits, + elem_float=_elem_is_float(ty) if is_ptr else False, + ) + func_args.append(fa) + # Pointer args seed addptr chains; scalar args are Param leaves. + env[f"%{name}"] = PtrValue(name, Const(0)) if is_ptr else Param(name) + + def base_elem_bits(param: str) -> int: + fa = next((a for a in func_args if a.name == param), None) + return fa.elem_bits if fa else 0 + + def base_elem_float(param: str) -> bool: + fa = next((a for a in func_args if a.name == param), None) + return fa.elem_float if fa else True # unknown pointee: fail closed + + def operand_term(v: object) -> "Term | None": + """An atomic cmp/val operand as a Term, or None when unmodelable.""" + return None if isinstance(v, (DataDep, PtrValue)) else v # type: ignore[return-value] + + def loaded_binding(acc: AccessEvent, idx: int, extra: str) -> object: + """Route 2: the value of an integer load whose mask is modeled; + float pointees and dropped masks stay DataDep (a masked-off lane + holds ``other`` or an undefined value, which only a modeled mask + can keep apart from the snapshot value).""" + if acc.elem_float or acc.mask_dropped: + return DataDep("loaded value") + trailing = _split_ssa(extra) if extra else [] + other_t: Term | None = None + if len(trailing) > 1: + ov = val(trailing[1]) + if not isinstance(ov, (DataDep, PtrValue)): + other_t = ov # type: ignore[assignment] + return Loaded(idx, acc.base_param, acc.offset, acc.mask, other_t) + + def observed_result_binding() -> object: + """The env value for the just-recorded access's result: Observed + for an integer-typed access (spec part B / the await re-read), + DataDep otherwise (float pointees stay outside the Int model).""" + if accesses and not accesses[-1].elem_float: + return Observed(len(accesses) - 1) + return DataDep("atomic result") + + # ── body parse (single function; loop handled inline) ── + # Region stack: "for" | _IfFrame. Tracking scf.if frames keeps the + # walker's brace accounting honest (an if's closing brace inside a loop + # must not be mistaken for the loop's close, nor its scf.yield for the + # loop's yield), carries the modeled branch condition for the accesses + # inside (``path``), and marks accesses under an UNMODELED condition as + # ``guarded``. + frames: list = [] + pid_axes: set[int] = set() + + def access_state() -> "tuple[bool, Term | None, bool, tuple[str, ...]]": + """_branch_state plus the current block's predicate (multipath): + the block predicate is the outermost conjunct of ``path`` and an + inexact block widens the access. Identical to _branch_state while + the walk is in the entry block.""" + guarded, path, in_loop, loops_ = _branch_state(frames) + if cur is entry: + return guarded, path, in_loop, loops_ + if not cur.resolved: + raise UnsupportedTTIR( + f"block ^{cur.name} is entered from a later block " + "(non-Triton block order)", + kind="control-flow", + ) + if cur.terminated: + raise UnsupportedTTIR( + f"access after the terminator of ^{cur.name}", + kind="control-flow", + ) + return guarded or cur.guarded, _conj(cur.pred, path), in_loop, loops_ + + def block_for(name: str) -> _Block: + assert n_preds is not None + blk = blocks.get(name) + if blk is None: + blk = _Block(name, n_preds=n_preds.get(name, 0)) + blocks[name] = blk + return blk + + def record_edge( + target: str, path: "Term | None", exact: bool, inner: "str | None" + ) -> None: + # Operand values resolve NOW: they are SSA names of the branching + # block, which the target's parameter binding must not re-read. + block_for(target).edges.append( + (path, exact, [val(s) for s in _arg_ssas(inner)]) + ) + + # Depth of anonymous OP regions (``"tt.reduce"(...) ({`` ... ``})``): + # their ``^bb0(...)`` combine-block labels belong to the op, not to the + # function's cf.* graph, and stay ignored exactly as in single-path. + op_region_depth = 0 + + for line_no, raw in enumerate(lines, start=1): + line = raw.strip() + if not line or line.startswith("#"): + continue + if line.endswith("({"): + op_region_depth += 1 + elif line.startswith("})") and op_region_depth: + op_region_depth -= 1 + m = _RE_FUNC.search(line) + if m and not kernel_name: + kernel_name = m.group(1) + parse_func_args(m.group(2)) + continue + if not kernel_name: + continue + + loc_m = _RE_LOC_TRAILER.search(line) + loc = locs.resolve(loc_m.group(1)) if loc_m else None + + rm = _RE_RESULT.match(line) + res = rm.group(1) if rm else None + body = rm.group(2) if rm else line + + # ---- scf.while body region: pure bookkeeping only (C1.1) ---- + # Placed FIRST so stray ops in the "do" region are refused before + # any other handler could record them; brace lines fall through to + # the region-close logic below. + if ( + frames + and isinstance(frames[-1], _WhileFrame) + and frames[-1].stage == "body" + and not line.startswith("}") + ): + if body.startswith("scf.yield"): + continue + raise UnsupportedTTIR( + f"line {line_no}: spin-loop body must be pure bookkeeping " + f"(scf.yield), found: {body.split(' ', 1)[0]}", + kind="spin-shape", + ) + + # ---- scf.while (the await abstraction, C1) ---- + if body.startswith("scf.while"): + if res is not None or not _RE_SCF_WHILE_SPIN.match(body): + raise UnsupportedTTIR( + f"line {line_no}: scf.while carries values (iter args or " + "results) — only the argument-free spin form is the " + "await shape", + kind="spin-shape", + ) + if any(isinstance(f, _WhileFrame) for f in frames): + raise UnsupportedTTIR( + f"line {line_no}: nested spin loops are not the await " "shape", + kind="spin-shape", + ) + frames.append( + _WhileFrame(open_line=line_no, n_accesses_before=len(accesses)) + ) + continue + + cm = _RE_SCF_CONDITION.match(body) + if cm: + top = frames[-1] if frames else None + if not (isinstance(top, _WhileFrame) and top.stage == "cond"): + raise UnsupportedTTIR( + f"line {line_no}: scf.condition outside a spin loop", + kind="control-flow", + ) + # Resolve NOW: region SSA names must not be re-read at close. + top.cond_val = val(cm.group(1)) + continue + + if line.startswith("} do") and frames and isinstance(frames[-1], _WhileFrame): + top = frames[-1] + if top.cond_val is None: + raise UnsupportedTTIR( + f"line {line_no}: spin loop without scf.condition", + kind="spin-shape", + ) + top.stage = "body" + continue + + # ---- scf.for ---- + fm = _RE_SCF_FOR.match(body) + if fm: + # ``loop`` is only set at the closing brace, so a second + # SEQUENTIAL loop is caught by it — but a NESTED loop opens while + # the outer one is still in flight (loop is still None), so guard + # on open frames too. Nested loops carry independent induction + # variables the single-loop model cannot represent, and a loop + # under an scf.if runs a condition-dependent iteration count; + # reject rather than silently mis-bound the induction var. + # Multipath (Route 3) lifts exactly this refusal: every loop gets + # its own induction variable and a loop under a condition + # carries that condition in its records' path. A loop inside a + # spin loop stays refused (the await shape has no body ops). + second_loop = loop is not None or bool(frames) + in_spin = any(isinstance(f, _WhileFrame) for f in frames) + if second_loop and (not multipath or in_spin): + raise UnsupportedTTIR( + f"line {line_no}: multiple/nested loops", + # A loop under an scf.if runs a branch-dependent + # iteration count — a control-flow limitation, not one + # more induction variable. + kind=( + "control-flow" + if any(isinstance(f, _IfFrame) for f in frames) + else "nested-loop" + ), + ) + ind, lo, up, st, iters = fm.groups() + pairs: list[tuple[str, str]] = [] + if iters: + pairs = list(re.findall(rf"({_SSA}) = ({_SSA})", iters)) + bound_terms: dict[str, Term] = {} + for label, ssa in (("lower", lo), ("upper", up), ("step", st)): + bv = val(ssa) + if isinstance(bv, DataDep): + # The CSR shape: for k in range(loaded_start, loaded_end). + raise UnsupportedTTIR( + f"loop {label} bound: data-dependent ({bv.why})", + kind="data-dependent-bound" if _from_memory(bv) else "other", + ) + if mentions_observed(bv): + # A trip count driven by an atomic observation is a + # dynamic work-fetch loop — outside the single-loop + # model (looped RMW fetch is a B+C1 stretch item). + raise UnsupportedTTIR( + f"loop {label} bound depends on an atomic observation", + kind="data-dependent-bound", + ) + bound_terms[label] = as_term(bv, f"loop {label}") + # The first loop keeps the historical "%loop" name; further + # loops (multipath only) need distinct names for their + # LoopVar / LoopInfo identity. + if loops_opened == 0: + loop_ssa = res or "%loop" # the single-path identity, unchanged + else: + # MLIR restarts value numbering per region, so two loops WITH + # results in sibling regions (then/else arms) print the same + # name; the line number keeps every later loop distinct + # (multipath only: single-path refuses a second loop). + loop_ssa = f"{res or '%loop'}@{line_no}" + frame = _ForFrame( + ssa=loop_ssa, + ind=ind, + lower=bound_terms["lower"], + upper=bound_terms["upper"], + step=bound_terms["step"], + order=loops_opened, + ) + loops_opened += 1 + # Bind induction var as a loop free variable. + env[ind] = LoopVar(loop_ssa) + # Bind ptr iter_args to IterArgOffset; ignore non-ptr (accumulators). + for arg_ssa, init_ssa in pairs: + iv = val(init_ssa) + if isinstance(iv, PtrValue): + arg_id = next_arg_id + next_arg_id += 1 + iter_args[arg_id] = IterArgInfo( + arg_id=arg_id, + base_param=iv.base_param, + offset0=iv.offset, + delta=Const(0), # filled at yield + loop_ssa=loop_ssa, + ) + env[arg_ssa] = PtrValue(iv.base_param, IterArgOffset(arg_id)) + frame.iter_arg_ssa.append((arg_ssa, init_ssa)) + frame.ptr_arg_ids.append(arg_id) + else: + env[arg_ssa] = DataDep("loop accumulator") + frame.iter_arg_ssa.append((arg_ssa, init_ssa)) + frames.append(frame) + continue + + # ---- scf.if: track the region and model its condition ---- + if body.startswith("scf.if"): + im = _RE_SCF_IF.match(body) + cond_t: Term | None = None + if im: + cv = val(im.group(1)) + # A pointer can't be a condition; loaded data (DataDep) + # can't be modeled → the region stays pessimistically + # ``guarded`` exactly as before this feature. + if not isinstance(cv, (DataDep, PtrValue)): + cond_t = cv # type: ignore[assignment] + frames.append(_IfFrame(cond=cond_t, res=res)) + # Fallback binding; upgraded to Select at the closing brace when + # the condition and both branches' single yield are modelable. + if res is not None: + env[res] = DataDep("scf.if result") + continue + + # A region close prints as ``}``, ``} loc(...)``, ``} else {`` or, + # with op attributes (``tl.range(num_stages=...)``), ``} {tt.num_stages + # = 2 : i32} loc(...)``; the last form used to leave its loop frame + # open until the function's own close (an access placed between the + # two would have been mis-attributed to the loop). + if frames and ( + line == "}" + or line.startswith("} loc") + or line.startswith("} else") + or line.startswith("} {") + ): + if line.startswith("} else"): + # The then-region closes and the else-region opens: the same + # if frame stays on the stack with its condition negated for + # the accesses that follow. + top = frames[-1] + if not isinstance(top, _IfFrame): + raise UnsupportedTTIR(f"line {line_no}: unexpected `else`") + top.branch = "else" + continue + popped = frames.pop() + if isinstance(popped, _WhileFrame): + _finalize_await(popped, accesses, line_no) + continue + if isinstance(popped, _IfFrame): + if ( + popped.res is not None + and popped.cond is not None + and popped.then_vals is not None + and popped.else_vals is not None + and len(popped.then_vals) == 1 + and len(popped.else_vals) == 1 + ): + tv, ev = popped.then_vals[0], popped.else_vals[0] + # Yielded pointers or loaded data keep the DataDep + # fallback (a stored VALUE never enters address math; + # an address use of the result then fails closed). + if not any(isinstance(x, (DataDep, PtrValue)) for x in (tv, ev)): + env[popped.res] = Select( + popped.cond, + as_term(tv, "scf.if yield"), + as_term(ev, "scf.if yield"), + ) + continue + # A "for" frame closed: resolve deltas from the yields, positionally. + assert isinstance(popped, _ForFrame) + ptr_idx = 0 + for pos, (arg_ssa, _init) in enumerate(popped.iter_arg_ssa): + if not isinstance(env.get(arg_ssa), PtrValue): + continue + if pos >= len(popped.body_yields): + raise UnsupportedTTIR("loop yield/iter_arg count mismatch") + yssa = popped.body_yields[pos] + yv = env.get(yssa) + if not isinstance(yv, PtrValue): + raise UnsupportedTTIR("loop yields a non-pointer for a ptr arg") + aid = popped.ptr_arg_ids[ptr_idx] + delta = _extract_loop_delta(yv.offset, aid) + if delta is None: + raise UnsupportedTTIR( + f"loop pointer advance for arg {aid} is not a " + "simple monotonic addptr" + ) + info = iter_args[aid] + iter_args[aid] = IterArgInfo( + info.arg_id, info.base_param, info.offset0, delta, popped.ssa + ) + ptr_idx += 1 + closed = LoopInfo( + loop_ssa=popped.ssa, + induction_var=popped.ind, + lower=popped.lower, + upper=popped.upper, + step=popped.step, + ) + loops.append((popped.order, closed)) + if loop is None: + loop = closed + continue + + ym = _RE_SCF_YIELD.match(body) + if ym and frames and isinstance(frames[-1], _ForFrame): + # Only the loop's own yield resolves iter-arg deltas; an scf.if's + # yield inside the loop body must not clobber it. + frames[-1].body_yields = _split_ssa(ym.group(1)) + continue + if ym and frames and isinstance(frames[-1], _IfFrame): + # Resolve yield VALUES here, not at the closing brace: then/else + # regions legally reuse the same SSA names, so a close-time + # lookup would read the else-region's overwrites. + fr = frames[-1] + vals = [val(s) for s in _split_ssa(ym.group(1))] + if fr.branch == "then": + fr.then_vals = vals + else: + fr.else_vals = vals + continue + + # ---- the unstructured cf.* graph (multipath, Route 3) ---- + # Block predicates are computed in the walk order: Triton creates + # the blocks of an if-with-return in topological order (then, else, + # nested blocks, merge), so a label normally sees every incoming + # edge; the exception (a shared return-only block placed before a + # later predecessor) is tolerated only while nothing inside needs + # the predicate (access_state refuses otherwise). + if ( + multipath + and op_region_depth == 0 + and (body.startswith("cf.") or line.startswith("^bb")) + ): + if n_preds is None: + n_preds = _prescan_blocks(lines) + if frames: + raise UnsupportedTTIR( + f"line {line_no}: cf.* control flow inside an scf region", + kind="control-flow", + ) + lbm = _RE_BLOCK_LABEL.match(line) + if lbm: + blk = block_for(lbm.group(1)) + params = _arg_ssas(lbm.group(2)) + if len(blk.edges) < blk.n_preds: + blk.resolved = False + for prm in params: + env[prm] = DataDep("block argument of an unresolved block") + cur = blk + continue + pred: Term | None = None + exact_all = True + for i, (epath, exact, _vals) in enumerate(blk.edges): + pred = epath if i == 0 else _disj(pred, epath) + exact_all = exact_all and exact + if not blk.edges: + # Unreachable block (no predecessor): no execution + # enters it. Keep it inert rather than fabricating + # accesses; Triton does not emit such blocks. + pred, exact_all = Cmp("ne", Const(0), Const(0)), True + blk.pred = pred + blk.guarded = not exact_all + for pi, prm in enumerate(params): + env[prm] = ( + _merge_block_param(blk.edges, pi) + if exact_all + else DataDep("block argument") + ) + cur = blk + continue + if cur.terminated or not cur.resolved: + raise UnsupportedTTIR( + f"line {line_no}: branch in an unresolved or terminated " + f"block ^{cur.name}", + kind="control-flow", + ) + cbm = _RE_COND_BR.match(body) + if cbm: + cv = val(cbm.group(1)) + base_exact = not cur.guarded + if not isinstance(cv, (DataDep, PtrValue)): + cond: Term = cv # type: ignore[assignment] + record_edge( + cbm.group(2), _conj(cur.pred, cond), base_exact, cbm.group(3) + ) + record_edge( + cbm.group(4), + _conj(cur.pred, Not(cond)), + base_exact, + cbm.group(5), + ) + else: + # Loaded-data condition: both targets stay reachable + # under the predecessor's predicate alone (widening). + record_edge(cbm.group(2), cur.pred, False, cbm.group(3)) + record_edge(cbm.group(4), cur.pred, False, cbm.group(5)) + cur.terminated = True + continue + brm = _RE_BR.match(body) + if brm: + record_edge(brm.group(1), cur.pred, not cur.guarded, brm.group(2)) + cur.terminated = True + continue + raise UnsupportedTTIR( + f"line {line_no}: control flow {body.split(' ', 1)[0]} is unsupported", + kind="control-flow", + ) + + # ---- other control flow: fail closed ---- + # scf.for, scf.if and the scf.while await shape are region-tracked + # above. Anything else that steers control flow (unstructured cf.*) + # would be flat-scanned as if it executed unconditionally — reject + # the kernel instead. + if body.startswith(("scf.", "cf.")) and not body.startswith( + ("scf.for", "scf.if", "scf.yield") + ): + raise UnsupportedTTIR( + f"line {line_no}: control flow {body.split(' ', 1)[0]} is unsupported", + kind="control-flow", + ) + + # ---- dependency provenance (D2/D3) ---- + # Every value-producing statement that is not a memory access + # inherits the provenance of its SSA operands; ops that move + # elements between positions (or aggregate them) drop the + # position-preserving flag. Accesses seed/consume it below. + if res is not None and not body.startswith( + ("tt.load", "tt.store", "tt.atomic_") + ): + operands = [t for t in re.findall(_SSA, body)] + position_preserving = body.startswith( + ( + "arith.", + "math.", + "tt.addptr", + "tt.bitcast", + "tt.fp_to_fp", + "tt.int_to_ptr", + "tt.ptr_to_int", + "tt.clampf", + "tt.precise_", + "tt.mulhiui", + # libdevice / inline-asm calls are element-wise by + # construction (liger's tanh-based GeGLU backward) + "tt.extern_elementwise", + "tt.elementwise_inline_asm", + ) + ) + dot = _RE_DOT.match(body) + if dot: + # d[i,j] = product(A, B)[i,j] + C[i,j]: only the + # accumulator's existing positional path represents D3. + # Keep all A/B paths non-positional, then restore C's + # flags so an A/C-shared source retains its valid C path. + # This does not model dot values, move positions, or add + # ordering for the matrix product's inputs. + merged = _prov_of(dot.group(1, 2), False) + merged.update(_prov_of((dot.group(3),))) + else: + merged = _prov_of( + operands, + position_preserving, + simultaneous=position_preserving + and not body.startswith("arith.select"), + ) + selection = _RE_SELECT.match(body) + if selection: + # Only one value arm contributes. Preserve an arm-derived + # dependency unconditionally only if BOTH arms have it. + # The condition itself is evaluated at every position. + condition, true_arm, false_arm = selection.groups() + condition_prov = prov.get(condition, {}) + true_prov = prov.get(true_arm, {}) + false_prov = prov.get(false_arm, {}) + for idx in merged: + merged[idx] = condition_prov.get(idx, False) or ( + true_prov.get(idx, False) and false_prov.get(idx, False) + ) + elif body.startswith("arith.select"): + merged = {idx: False for idx in merged} + if merged: + prov[res] = merged + + # ---- tile-level fence ---- + if body.startswith("gpu.barrier"): + fences.append(len(accesses) - 0.5) + continue + + # ---- value-producing ops ---- + handled = _parse_value_op( + body, res, env, val, as_term, base_elem_bits, pid_axes + ) + if handled: + continue + + # ---- accesses ---- + lm = _RE_LOAD.match(body) + if lm: + guarded, path, in_loop, loops_ = access_state() + _record_access( + "load", + lm.group(1), + lm.group(2), + guarded, + env, + val, + accesses, + base_elem_bits, + loc, + line_no, + path=path, + in_loop=in_loop, + loops=loops_, + keep_partial_mask=multipath, + base_elem_float=base_elem_float, + ) + if res is not None: + in_while_cond = any( + isinstance(f, _WhileFrame) and f.stage == "cond" for f in frames + ) + # A spin re-read's value IS an observation (the await's + # exit predicate is asserted over it, C1.2); everywhere + # else a loaded value stays DataDep, except under the L2 + # reader mode, where an integer load with a modeled mask + # becomes a Loaded term (Route 2: its value is a Select + # over the launch's snapshot of the source tensor). + if in_while_cond: + env[res] = observed_result_binding() + elif multipath: + env[res] = loaded_binding( + accesses[-1], len(accesses) - 1, lm.group(2) + ) + else: + env[res] = DataDep("loaded value") + if res is not None: + prov[res] = {len(accesses) - 1: True} + continue + sm = _RE_STORE.match(body) + if sm: + if any(isinstance(f, _WhileFrame) for f in frames): + raise UnsupportedTTIR( + f"line {line_no}: store inside a spin loop is not the " + "await shape", + kind="spin-shape", + ) + guarded, path, in_loop, loops_ = access_state() + _record_access( + "store", + sm.group(1), + sm.group(3), + guarded, + env, + val, + accesses, + base_elem_bits, + loc, + line_no, + path=path, + in_loop=in_loop, + loops=loops_, + keep_partial_mask=multipath, + ) + object.__setattr__( + accesses[-1], + "deps", + _deps_of(sm.group(2), *re.findall(_SSA, sm.group(3) or "")), + ) + continue + am = _RE_ATOMIC_RMW.match(body) + if am: + guarded, path, in_loop, loops_ = access_state() + _record_access( + "atomic_rmw", + am.group(4), + am.group(6), # the mask operand + guarded, + env, + val, + accesses, + base_elem_bits, + loc, + line_no, + atomic=AtomicInfo(am.group(1), am.group(2), am.group(3)), + path=path, + in_loop=in_loop, + loops=loops_, + keep_partial_mask=multipath, + atomic_val=operand_term(val(am.group(5))), + base_elem_float=base_elem_float, + ) + if res is not None: + env[res] = observed_result_binding() + object.__setattr__(accesses[-1], "deps", _deps_of(am.group(5), am.group(6))) + if res is not None: + prov[res] = {len(accesses) - 1: True} + continue + am = _RE_ATOMIC_CAS.match(body) + if am: + guarded, path, in_loop, loops_ = access_state() + _record_access( + "atomic_cas", + am.group(3), + "", # CAS has no mask operand: unconditional footprint + guarded, + env, + val, + accesses, + base_elem_bits, + loc, + line_no, + atomic=AtomicInfo(None, am.group(1), am.group(2)), + path=path, + in_loop=in_loop, + loops=loops_, + keep_partial_mask=multipath, + atomic_val=operand_term(val(am.group(5))), + atomic_cmp=operand_term(val(am.group(4))), + base_elem_float=base_elem_float, + ) + if res is not None: + env[res] = observed_result_binding() + object.__setattr__(accesses[-1], "deps", _deps_of(am.group(4), am.group(5))) + if res is not None: + prov[res] = {len(accesses) - 1: True} + continue + + # ---- fail closed on unrecognized memory ops ---- + # A tt.load/tt.store/tt.atomic_* syntax variant the regexes above did + # not match must NOT fall through to the value/DataDep handling below: + # a store has no result so it would be silently dropped, and an + # atomic's access would go unchecked while its result becomes a + # harmless-looking DataDep. Either way check_graph would then prove + # "ok" without having checked a real access. Bail to unsupported + # instead so the proof stays sound. + if body.startswith( + ( + "tt.load", + "tt.store", + "tt.atomic_", + "tt.descriptor_", + "tt.experimental_descriptor_", + ) + ): + raise UnsupportedTTIR( + f"line {line_no}: unsupported memory op syntax: {body[:60]}", + kind="out-of-vocabulary", + ) + + # ---- ops whose result is just data (ignored) ---- + if res is not None and ( + body.startswith( + ( + "arith.addf", + "arith.mulf", + "arith.subf", + "arith.divf", + "arith.cmpf", + "tt.dot", + "arith.truncf", + "arith.extf", + "arith.sitofp", + "tt.reduce", + "math.", + ) + ) + ): + env[res] = DataDep("float/reduction value") + continue + if body.startswith(("tt.return", "tt.reduce.return")): + if multipath and body.startswith("tt.return") and not frames: + cur.terminated = True + continue + if body.startswith("tt.make_block_ptr") or body.startswith("tt.advance"): + raise UnsupportedTTIR( + f"line {line_no}: block pointers are unsupported", + kind="block-pointer", + ) + # Unknown op producing a value used downstream → conservative DataDep. + if res is not None: + env[res] = DataDep(f"unmodeled op at line {line_no}") + + if not kernel_name: + raise UnsupportedTTIR("no tt.func found (not TTIR?)") + + ordered = [lp for _o, lp in sorted(loops, key=lambda t: t[0])] + return AccessGraph( + kernel_name=kernel_name, + has_value_changing_integer_casts=changing_integer_casts, + func_args=func_args, + accesses=accesses, + loop=loop if len(ordered) == 1 else None, + iter_args=iter_args, + pid_axes=pid_axes, + loops=ordered, + multipath=multipath, + cf_blocks=len(blocks), + fences=fences, + ) + + +def _set_arange_dim(v: object, dim: int) -> object: + """Tag every Arange in an integer expression with the tensor dimension + it varies along (set by expand_dims). Non-Arange leaves pass through.""" + if isinstance(v, Arange): + return Arange(v.ssa, v.start, v.end, dim if v.dim < 0 else v.dim) + if isinstance(v, Bin): + return Bin(v.op, _set_arange_dim(v.a, dim), _set_arange_dim(v.b, dim)) # type: ignore[arg-type] + if isinstance(v, Cmp): + return Cmp(v.pred, _set_arange_dim(v.a, dim), _set_arange_dim(v.b, dim)) # type: ignore[arg-type] + if isinstance(v, BoolBin): + return BoolBin(v.op, _set_arange_dim(v.a, dim), _set_arange_dim(v.b, dim)) # type: ignore[arg-type] + if isinstance(v, Select): + return Select( + _set_arange_dim(v.cond, dim), # type: ignore[arg-type] + _set_arange_dim(v.t, dim), # type: ignore[arg-type] + _set_arange_dim(v.f, dim), # type: ignore[arg-type] + ) + if isinstance(v, Not): + return Not(_set_arange_dim(v.a, dim)) # type: ignore[arg-type] + if isinstance(v, Loaded): + # the loaded tile's lanes follow the consumer's dimension exactly + # like an arange's (an expand_dims of the loaded value) + return Loaded( + v.access_index, + v.base_param, + _set_arange_dim(v.offset, dim), # type: ignore[arg-type] + None if v.mask is None else _set_arange_dim(v.mask, dim), # type: ignore[arg-type] + None if v.other is None else _set_arange_dim(v.other, dim), # type: ignore[arg-type] + ) + if isinstance(v, DataDep) and v.keep is not None: + # The kept conjunct of a mixed ``and`` must follow the tile's + # dimension like any other lane term, or its Arange would name a + # lane variable the address never uses (a vacuous mask). + return DataDep(v.why, keep=_set_arange_dim(v.keep, dim)) # type: ignore[arg-type] + if isinstance(v, PtrValue): + # A POINTER tile expanded (``base[None, :] + off[:, None]``): its + # offset's lanes follow the new dimension exactly like an integer + # tile's. Left untagged, the address kept the 1-D variable while + # the mask (an i1 tile expanded the same way) got the 2-D one, so + # the two copies could differ in a lane the address never reads: + # a phantom intra-instance WAW (aiter's causal_conv1d update + # kernels, found by Route 2's change surface). + return PtrValue(v.base_param, _set_arange_dim(v.offset, dim)) # type: ignore[arg-type] + return v + + +def _finalize_await(frame: _WhileFrame, accesses: list, line_no: int) -> None: + """Validate the C1.1 shape contract at the spin loop's closing brace and + stamp the kept read with ``awaited`` + the EXIT predicate. + + ``scf.condition(c)`` continues WHILE c holds, so the exit predicate is + ``Not(c)`` — for ``while load(flag) != 1`` that is ``flag == 1``; for + the CAS form ``while cas(lock,0,1) != 0`` it is ``old == 0`` (success). + Memory order/scope stay exactly as the op was written: a relaxed spin + must yield no synchronizes-with edge — that IS the missing-acquire bug + the detector exists to find.""" + where = f"line {frame.open_line} (scf.while)" + n_new = len(accesses) - frame.n_accesses_before + if n_new != 1: + raise UnsupportedTTIR( + f"{where}: the spin condition must re-read exactly one location " + f"(found {n_new} memory accesses)", + kind="spin-shape", + ) + idx = len(accesses) - 1 + acc = accesses[idx] + if acc.elem_float: + raise UnsupportedTTIR( + f"{where}: the awaited location is float-typed (the observation " + "model is Int-sort only)", + kind="spin-shape", + ) + # The await encoding keeps ONE read and drops every earlier iteration — + # sound only when the re-read is side-effect-free on the awaited + # location. A plain load never writes; a CAS writes exactly once (on + # success — the single modeled write). A mutating RMW re-read + # (atomic_add(flag, 1) spins) writes on EVERY dropped iteration: the + # loop can terminate by observing its OWN increments, and modeling the + # exit value as read-from a release writer fabricates a + # synchronizes-with edge (adversarial finding: self-satisfying spin + # proved a real data race away). Accept an RMW only when its written + # value provably equals the observation: add/or/xor with a constant 0. + if acc.kind == "atomic_rmw": + op = ((acc.atomic.rmw_op if acc.atomic else None) or "").lower() + identity = op in ("add", "or", "xor") and acc.atomic_val == Const(0) + if not identity: + raise UnsupportedTTIR( + f"{where}: the spin re-read MUTATES the awaited location " + f"(atomic {op or '?'} with a non-identity operand); dropped " + "iterations would lose real writes", + kind="spin-shape", + ) + cv = frame.cond_val + if not isinstance(cv, Cmp): + raise UnsupportedTTIR( + f"{where}: the spin condition is not a comparison over the " "awaited read", + kind="spin-shape", + ) + a_is_obs = isinstance(cv.a, Observed) and cv.a.access_index == idx + b_is_obs = isinstance(cv.b, Observed) and cv.b.access_index == idx + expected = cv.b if a_is_obs else cv.a + if a_is_obs == b_is_obs or idx in observed_indices(expected): + raise UnsupportedTTIR( + f"{where}: the spin condition must compare the awaited read " + "against a loop-invariant expected value", + kind="spin-shape", + ) + accesses[idx] = replace(acc, awaited=True, exit_pred=Not(cv)) + + +def _extract_loop_delta(offset: Term, arg_id: int) -> Term | None: + """From a yielded pointer offset of the shape + ``IterArgOffset(arg_id) + delta`` (any association), pull out ``delta``.""" + if isinstance(offset, IterArgOffset): + return Const(0) + if isinstance(offset, Bin) and offset.op == "+": + if isinstance(offset.a, IterArgOffset) and offset.a.arg_id == arg_id: + return offset.b + if isinstance(offset.b, IterArgOffset) and offset.b.arg_id == arg_id: + return offset.a + return None + + +def _parse_value_op(body, res, env, val, as_term, base_elem_bits, pid_axes) -> bool: + """Parse one address-structure value op into env. Returns True if handled.""" + if res is None: + return False + + m = _RE_GET_PID.match(body) + if m: + axis = {"x": 0, "y": 1, "z": 2}.get(m.group(1)) + if axis is None: + # Printer drift must surface as the designed error, not a bare + # KeyError escaping into the client's launch teardown. + raise UnsupportedTTIR( + f"unknown program-id axis {m.group(1)!r}", + kind="out-of-vocabulary", + ) + # Parse-time record (see AccessGraph.pid_axes): the read counts even + # if this value never survives into a modeled term. + pid_axes.add(axis) + env[res] = Pid(axis) + return True + m = _RE_GET_NPROG.match(body) + if m: + axis = {"x": 0, "y": 1, "z": 2}.get(m.group(1)) + if axis is None: + raise UnsupportedTTIR( + f"unknown num-programs axis {m.group(1)!r}", + kind="out-of-vocabulary", + ) + # The verdict depends on this grid dim (see NumPrograms): keep the + # axis symbolic even when no pid read distinguishes blocks along it. + pid_axes.add(axis) + env[res] = NumPrograms(axis) + return True + m = _RE_MAKE_RANGE.match(body) + if m: + env[res] = Arange(res, int(m.group(2)), int(m.group(1))) + return True + m = _RE_CONST_INT.match(body) + if m: + env[res] = Const(int(m.group(1))) + return True + m = _RE_CONST_DENSE.match(body) + if m: + env[res] = Const(int(m.group(1))) + return True + m = _RE_CONST_DENSE_BOOL.match(body) or _RE_CONST_BOOL.match(body) + if m: + # i1 constants (e.g. the dense mask of an unmasked atomic). + # Const(0/1) in a boolean position is coerced by the evaluator. + env[res] = Const(1 if m.group(1) == "true" else 0) + return True + if body.startswith("arith.constant"): + env[res] = DataDep("float/array constant") + return True + m = _RE_SPLAT.match(body) + if m: + env[res] = val(m.group(1)) # replicate scalar / seed ptr + return True + m = _RE_EXPAND.match(body) + if m and body.startswith("tt.expand_dims"): + # axis is the inserted size-1 dim; the lane index varies along the + # OTHER dim (1 - axis for a 1D->2D expand). Tag every Arange inside. + axis = int(m.group(2)) + env[res] = _set_arange_dim(val(m.group(1)), 1 - axis) + return True + m = _RE_BROADCAST.match(body) + if m and body.startswith("tt.broadcast"): + env[res] = val(m.group(1)) # shape change, value passthrough + return True + m = _RE_EXT.match(body) + if m: + env[res] = val(m.group(2)) # width change, value passthrough + return True + m = _RE_ADDPTR.match(body) + if m: + base, off = val(m.group(1)), val(m.group(2)) + if not isinstance(base, PtrValue): + raise UnsupportedTTIR( + "addptr base is not a pointer", + kind="indirect-address" if _from_memory(base) else "other", + ) + if isinstance(off, DataDep): + # A value in an address chain that cannot be modeled: a free + # address makes the query meaningless, so this stays + # whole-kernel unsupported. Only offsets truly derived from + # MEMORY CONTENTS classify as indirection (the interpreter + # front-end route); modeling gaps (loop accumulators, unmodeled + # ops, ...) keep the default kind so the buckets stay honest. + raise UnsupportedTTIR( + f"addptr offset: data-dependent ({off.why})", + kind="indirect-address" if _from_memory(off) else "other", + ) + off_t = as_term(off, "addptr offset") + env[res] = PtrValue(base.base_param, Bin("+", base.offset, off_t)) + return True + m = _RE_BIN.match(body) + if m: + op = { + "muli": "*", + "addi": "+", + "subi": "-", + "divsi": "//", + "remsi": "%", + "minsi": "min", + "maxsi": "max", + }[m.group(1)] + a, b = val(m.group(2)), val(m.group(3)) + if isinstance(a, DataDep) or isinstance(b, DataDep): + env[res] = DataDep("arith over loaded data") + else: + env[res] = Bin(op, as_term(a, "arith"), as_term(b, "arith")) + return True + m = _RE_CMPI.match(body) + if m: + a, b = val(m.group(2)), val(m.group(3)) + if isinstance(a, DataDep) or isinstance(b, DataDep): + env[res] = DataDep("cmpi over loaded data") + else: + env[res] = Cmp(m.group(1), as_term(a, "cmpi"), as_term(b, "cmpi")) + return True + m = _RE_BOOLBIN.match(body) + if m: + ty = m.group(4) + if not (ty == "i1" or ty.endswith("i1>")): + # Wide-int andi/ori is BITWISE arithmetic, not boolean logic; + # modeling it as And/Or would silently corrupt address math + # (e.g. ``offs & 8`` collapsing to a {0,1} truth value). Degrade + # to DataDep so an address use fails closed as unsupported. + env[res] = DataDep(f"bitwise arith.{m.group(1)} on non-i1 type {ty}") + return True + a, b = val(m.group(2)), val(m.group(3)) + if isinstance(a, DataDep) or isinstance(b, DataDep): + keep: Term | None = None + if m.group(1) == "andi": + # ``modelable ∧ unmodelable`` implies ``modelable``: remember + # the modelable conjunct(s) so a mask can keep them. + parts = [] + for x in (a, b): + if isinstance(x, DataDep): + if x.keep is not None: + parts.append(x.keep) + elif not isinstance(x, PtrValue): + parts.append(x) + for part in parts: + keep = part if keep is None else BoolBin("and", keep, part) + env[res] = DataDep("bool op over loaded data", keep=keep) + else: + env[res] = BoolBin( + "and" if m.group(1) == "andi" else "or", + as_term(a, "bool"), + as_term(b, "bool"), + ) + return True + m = _RE_SELECT.match(body) + if m: + c, t, f = val(m.group(1)), val(m.group(2)), val(m.group(3)) + if any(isinstance(x, DataDep) for x in (c, t, f)): + env[res] = DataDep("select over loaded data") + else: + env[res] = Select( + as_term(c, "select"), as_term(t, "select"), as_term(f, "select") + ) + return True + return False + + +def _record_access( + kind, + ptr_ssa, + extra_ops, + guarded, + env, + val, + accesses, + base_elem_bits, + loc, + line_no, + atomic=None, + path=None, + in_loop=False, + atomic_val=None, + atomic_cmp=None, + base_elem_float=None, + loops=(), + keep_partial_mask=False, +) -> None: + ptr = val(ptr_ssa) + if not isinstance(ptr, PtrValue): + raise UnsupportedTTIR( + f"line {line_no}: {kind} of a non-pointer value", + kind="indirect-address" if _from_memory(ptr) else "other", + ) + # Mask: for load it's the first trailing operand; for store the operand + # after value. _RE_LOAD captures trailing ", %x" groups; for store the + # caller passed the post-value trailing operands. + mask: Term | None = None + mask_dropped = False + trailing = _split_ssa(extra_ops) if extra_ops else [] + if trailing: + mv = val(trailing[0]) + if isinstance(mv, DataDep): + # Mask derived from loaded data: over-approximate it as free + # (any lane may be active) instead of failing the whole kernel. + # See AccessEvent.mask_dropped for the soundness discipline. + # Multipath keeps the modelable conjuncts of a mixed ``and`` + # (``bounds_mask and loaded_guard``): still an over-approximation + # (the access stays widened), but one that no longer activates + # lanes the bounds mask excludes, which is what turned such + # rows into phantom overlaps. + mask_dropped = True + if keep_partial_mask and mv.keep is not None: + mask = mv.keep + elif isinstance(mv, PtrValue): + raise UnsupportedTTIR(f"line {line_no}: pointer as mask") + else: + mask = mv # type: ignore[assignment] + accesses.append( + AccessEvent( + kind=kind, + base_param=ptr.base_param, + offset=ptr.offset, + mask=mask, + elem_bits=base_elem_bits(ptr.base_param), + loc=loc, + line_no=line_no, + guarded=guarded, + atomic=atomic, + path=path, + mask_dropped=mask_dropped, + in_loop=in_loop, + atomic_val=atomic_val, + atomic_cmp=atomic_cmp, + elem_float=(base_elem_float(ptr.base_param) if base_elem_float else False), + loops=tuple(loops), + ) + ) diff --git a/tests/unit/ir/test_host_compile.py b/tests/unit/ir/test_host_compile.py new file mode 100644 index 000000000..8e35ab45c --- /dev/null +++ b/tests/unit/ir/test_host_compile.py @@ -0,0 +1,1302 @@ +"""tilelens.core.host_compile: IR targets (D26) and the host compile (D25). + +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 dataclasses +import importlib +import re +import sys + +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") + + +@pytest.fixture +def private_triton_cache(tmp_path, monkeypatch): + """Give whole-pipeline compiles (triton.compile) an empty disk cache. + + Triton keys its disk cache by kernel source and start line, not by file + path, while the cached TTIR's #loc names the file it was compiled from. A + second checkout of the repo would otherwise read the first checkout's + kernels, and TTIR text comparisons would see the other checkout's path.""" + monkeypatch.setenv("TRITON_CACHE_DIR", str(tmp_path / "triton-cache")) + + +CUDA80 = GPUTarget("cuda", 80, 32) +# The default IR target (D26, amended). +CUDA89 = GPUTarget("cuda", 89, 32) + + +# ======== targets (D26) ========= + + +@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"), + # The former Triton-Viz name still works; the TileLens one wins. + ({"TRITON_VIZ_IR_TARGET": "hip:gfx942"}, "hip:gfx942"), + ( + {"TILELENS_IR_TARGET": "cuda:90", "TRITON_VIZ_IR_TARGET": "hip:gfx942"}, + "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 (D25) ========= + + +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))) + + +def _release() -> str: + """The installed Triton's minor release, e.g. "3.6".""" + return ".".join(triton.__version__.split(".")[:2]) + + +def _for_release(table: dict, what: str): + """``table``'s row for the installed Triton; a release it has no row for + fails the test, naming what to add.""" + release = _release() + if release not in table: + pytest.fail(f"no {what} for Triton {release}: add its row") + return table[release] + + +# The TTIR arguments a tuple parameter ``pair`` of two ints flattens to, per +# Triton release: 3.6 names both by the parameter (uniqued by the printer), +# 3.8 each by its path in the tuple. +_TUPLE_ARGUMENT_NAMES = { + "3.6": ("pair", "pair_0"), + "3.8": ("pair.0", "pair.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, + stages={"ttir"}, + ) + assert _signature(kernel.asm["ttir"]).get("n") == ttir_type + + +@needs_compiles +def test_a_ttir_request_compiles_only_through_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, + stages={"ttir"}, + ) + 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. + first, second = _for_release(_TUPLE_ARGUMENT_NAMES, "tuple argument names") + 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_deeper_stages_and_the_whole_pipeline_share_the_specialization( + no_driver, private_triton_cache +): + compiler = HostCompiler() + call = ((torch.zeros(64), torch.zeros(64), 64), {"BLOCK": 16}) + copy = _make_copy() + ttir = compiler.compile(copy, *call, target=CUDA80, stages={"ttir"}) + llir = compiler.compile(copy, *call, target=CUDA80, stages={"ttir", "llir"}) + assert list(llir.asm) == ["ttir", "ttgir", "llir"] + assert isinstance(llir.metadata.shared, int) + # The binary stage, or one derived from it, is triton.compile's: an + # unloaded CompiledKernel. + full = compiler.compile(copy, *call, target=CUDA80, stages={"cubin"}) + sass = compiler.compile(copy, *call, target=CUDA80, stages={"sass"}) + for compiled in (full, sass): + assert type(compiled).__name__ == "CompiledKernel" + assert {"source", "ttir", "cubin"} <= set(compiled.asm) + assert compiled.module is None # never loaded + # "source", the front end's module, needs no pass at all. + source = compiler.compile(copy, *call, target=CUDA80, stages={"source"}) + with_ttir = compiler.compile(copy, *call, target=CUDA80, stages={"source", "ttir"}) + assert isinstance(source, HostKernel) and list(source.asm) == ["source"] + assert list(with_ttir.asm) == ["source", "ttir"] + assert source.asm["source"] == with_ttir.asm["source"] == full.asm["source"] + assert ttir.hash == llir.hash == full.hash == sass.hash == source.hash + assert ttir.asm["ttir"] == llir.asm["ttir"] == full.asm["ttir"] + + +@needs_compiles +def test_compiles_are_cached_per_call_target_and_stages(no_driver): + compiler = HostCompiler() + copy = _make_copy() + x, out = torch.zeros(64), torch.zeros(64) + + def compile(n=64, target=CUDA80, stages=("ttir",), **kwargs): + return compiler.compile( + copy, (x, out, n), {"BLOCK": 16, **kwargs}, target=target, stages=stages + ) + + 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, stages=("ttir",)) 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 + assert compile(stages=("ttgir",)) 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 (D26: 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 test_a_triton_without_the_private_api_is_named(monkeypatch): + import triton.runtime.jit as jit_module + + triton_api.cache_clear() + monkeypatch.delattr(jit_module, "create_function_from_signature") + try: + with pytest.raises( + HostCompileUnavailable, + match=r"lacks .*create_function_from_signature", + ): + triton_api() + finally: + monkeypatch.undo() + triton_api.cache_clear() + assert triton_api().create_function_from_signature is not None + + +@needs_compiles +def test_a_stage_rewriting_knob_compiles_the_whole_pipeline(no_driver, monkeypatch): + # TRITON_KERNEL_OVERRIDE: only triton.compile reads the override files. + from triton import knobs + + monkeypatch.setattr(knobs.compilation, "override", True) + kernel = HostCompiler().compile( + _make_copy(), + (torch.zeros(64), torch.zeros(64), 64), + {"BLOCK": 16}, + target=CUDA80, + stages={"ttir"}, + ) + assert type(kernel).__name__ == "CompiledKernel" and "cubin" in kernel.asm + + +# ======== the target the front end sees (D26) ========= + + +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), + stages={"ttir"}, + ) + 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, + stages={"ttir"}, + ) + 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, + stages={"ttir"}, + ) + 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"), + stages={"ttir"}, + ) + # 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, + stages={"ttir"}, + ) + 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, + stages={"ttir"}, + ) + 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, + stages={"ttir"}, + ) + assert target_queried(raised.value) + with pytest.raises(CompileTimeAssertionFailure) as raised: + compiler.compile(bounded, (x,), {"BLOCK": 64}, target=CUDA89, stages={"ttir"}) + 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, stages={"ttir"}) + with pytest.raises(TypeError, match="required positional argument: 'n'") as raised: + compiler.compile(on_cuda, (x,), {"BLOCK": 16}, target=CUDA89, stages={"ttir"}) + 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 (D28).""" + 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, stages={"ttir"} + ) + 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 + + +_PAUSE: dict = {} + + +def _pause_point(): + # Called by the kernel below during its compile (see _device_is_zero). + _PAUSE["during"] = triton.runtime.driver.active.get_current_target() + _PAUSE["entered"].set() + assert _PAUSE["release"].wait(30) + return True + + +_pause_point.__triton_builtin__ = True # type: ignore[attr-defined] + + +@needs_compiles +def test_other_threads_see_tritons_driver_while_a_thread_compiles(monkeypatch): + """The target answer is scoped to the compiling thread: a real launch + on another thread meanwhile still gets the machine's driver.""" + import inspect + import threading + + from triton.runtime.driver import driver + from triton.runtime.jit import constexpr_function + + machine = _Machine(GPUTarget("cuda", 89, 32)) + _on_machine(monkeypatch, machine) + machine_active = inspect.getattr_static(type(driver), "active") + monkeypatch.setattr( + sys.modules[__name__], + "_PAUSE", + {"entered": threading.Event(), "release": threading.Event()}, + ) + + @constexpr_function + def pause(): + return _pause_point() + + @triton.jit + def paused(x_ptr, BLOCK: tl.constexpr): + if pause(): + tl.store(x_ptr + tl.arange(0, BLOCK), 1.0) + + errors: list = [] + + def compile_(): + try: + HostCompiler().compile( + paused, (torch.zeros(16),), {"BLOCK": 16}, target=CUDA80 + ) + except BaseException as exc: # reported below + errors.append(exc) + + worker = threading.Thread(target=compile_) + worker.start() + try: + assert _PAUSE["entered"].wait(30), errors + # Mid-compile: the compiling thread sees its target, this one the + # machine's driver, which nobody asked. + assert _PAUSE["during"] == CUDA80 + assert driver.active is machine and machine.queries == 0 + finally: + _PAUSE["release"].set() + worker.join(30) + assert errors == [] + # The compile is over: the class attribute is what it was. + assert inspect.getattr_static(type(driver), "active") is machine_active + + +@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 (D26), 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 + ) + + +def test_check_stages_names_the_stages_a_target_holds(): + compiler = HostCompiler() + compiler.check_stages( + CUDA80, {"source", "ttir", "ttgir", "llir", "ptx", "cubin", "sass"} + ) + hip = parse_ir_target("hip:gfx942") + compiler.check_stages(hip, {"source", "ttir", "ttgir", "llir", "amdgcn", "hsaco"}) + for target, stages, unknown in [ + (CUDA80, {"TTIR"}, "['TTIR']"), + (CUDA80, {"ttir", "bogus"}, "['bogus']"), + (hip, {"sass", "ptx"}, "['ptx', 'sass']"), + ]: + with pytest.raises( + ValueError, match=re.escape(f"IR stages {unknown} are no stage") + ): + compiler.check_stages(target, stages) + + +def _clear_api_caches(): + from tilelens.core import host_compile + + triton_api.cache_clear() + host_compile._self_test_target.cache_clear() + + +@needs_compiles +def test_a_changed_compile_api_fails_the_self_test(no_driver, monkeypatch): + """A private API that changed shape fails every host compile as + HostCompileUnavailable naming it, not as the kernel's error.""" + import triton.runtime.jit as jit_module + + real = jit_module.create_function_from_signature + + def two_results(sig, params, backend): + binder = real(sig, params, backend) + return lambda *args, **kwargs: binder(*args, **kwargs)[:2] + + _clear_api_caches() + monkeypatch.setattr(jit_module, "create_function_from_signature", two_results) + try: + with pytest.raises( + HostCompileUnavailable, + match=r"built-in test kernel failed to host-compile for cuda:80 \(ValueError", + ): + HostCompiler().compile( + _make_copy(), + (torch.zeros(64), torch.zeros(64), 64), + {"BLOCK": 16}, + target=CUDA80, + ) + finally: + monkeypatch.undo() + _clear_api_caches() + assert HostCompiler().compile( + _make_copy(), + (torch.zeros(64), torch.zeros(64), 64), + {"BLOCK": 16}, + target=CUDA80, + ) + + +def test_a_target_query_the_scope_does_not_reach_is_named(monkeypatch): + from triton.language import target_info + + _clear_api_caches() + monkeypatch.setattr(target_info, "current_target", lambda: None) + try: + with pytest.raises( + HostCompileUnavailable, + match=r"target queries answer \{'tl.target_info.current_target\(\)': None\}", + ): + triton_api() + finally: + monkeypatch.undo() + _clear_api_caches() + assert triton_api().driver_config is not None + + +@needs_compiles +def test_the_self_test_does_not_read_this_packages_files(no_driver, monkeypatch): + """The built-in kernel's source is its own: a host_compile.py changed on + disk since import (an editable install being edited) does not break it.""" + import linecache + + from tilelens.core import host_compile + + path = host_compile.__file__ + monkeypatch.setitem(linecache.cache, path, (8, None, ["x = 1\n"], path)) + _clear_api_caches() + try: + host_compile._self_test_target(CUDA80) + finally: + monkeypatch.undo() + _clear_api_caches() + + +# ======== what Triton releases differ in ========= + +# Per Triton release, what its JIT runtime does that the host compile mirrors +# or answers (tilelens.core.host_compile._RELEASE_RUNTIMES): whether +# JITFunction.run and triton.compile key a kernel by a custom pipeline +# (knobs.runtime.add_stages_inspection_hook), and whether +# CompiledKernel.__del__ unloads a loaded module through the driver. +_RELEASE_RUNTIME = { + "3.6": {"stages_hook_keys": False, "unloads_on_del": False}, + "3.8": {"stages_hook_keys": True, "unloads_on_del": True}, +} + + +def _detected_runtime(): + from triton.compiler import compile as triton_compile + from triton.compiler.compiler import CompiledKernel + from triton.runtime.jit import JITFunction + + from tilelens.core import host_compile + + return host_compile._detected_runtime(triton_compile, JITFunction, CompiledKernel) + + +def test_the_installed_releases_runtime_is_known_and_its_code_agrees(): + expected = _for_release(_RELEASE_RUNTIME, "JIT runtime") + api = triton_api() + assert api.runtime is not None, api.runtime_unknown + assert dataclasses.asdict(api.runtime) == expected + assert api.unloads is expected["unloads_on_del"] + # What the installed code shows, independently of the table. + assert _detected_runtime() == expected + + +class _PipelineHook: + """A custom pipeline (``knobs.runtime.add_stages_inspection_hook``) in + both of its calling conventions: called with no arguments (Triton 3.8's + JITFunction.run and triton.compile) it names the pipeline, a (key, hash) + pair; called by a backend's add_stages it leaves the stages as they + are.""" + + def __init__(self, name: str) -> None: + self.name = name + self.arities: list[int] = [] + + def __call__(self, *args): + self.arities.append(len(args)) + if not args: + return (f"-pipeline-{self.name}", f"{self.name}0") + return None + + +@needs_compiles +def test_a_custom_pipeline_names_the_kernel_as_triton_compile_does( + no_driver, monkeypatch, private_triton_cache +): + """Under a custom pipeline a TTIR-only compile's hash is still the one + triton.compile gives the kernel (its whole pipeline), and a release that + keys kernels by the pipeline (3.8) compiles one kernel per pipeline.""" + from triton import knobs + + keyed = _for_release(_RELEASE_RUNTIME, "JIT runtime")["stages_hook_keys"] + call = (_make_copy(), (torch.zeros(64), torch.zeros(64), 64), {"BLOCK": 16}) + plain = HostCompiler().compile(*call, target=CUDA80, stages={"ttir"}) + hashes = {plain.hash} + for name in ("one", "two"): + hook = _PipelineHook(name) + monkeypatch.setattr(knobs.runtime, "add_stages_inspection_hook", hook) + compiler = HostCompiler() + ttir = compiler.compile(*call, target=CUDA80, stages={"ttir"}) + full = compiler.compile(*call, target=CUDA80, stages={"cubin"}) + assert type(full).__name__ == "CompiledKernel" + assert ttir.hash == full.hash + assert ttir.asm["ttir"] == full.asm["ttir"] == plain.asm["ttir"] + # Asked with no arguments only where the release keys by it. + assert (0 in hook.arities) is keyed and 5 in hook.arities + hashes.add(ttir.hash) + assert len(hashes) == (3 if keyed else 1) + + +class _MachineUtils: + """A stand-in for the machine driver's ``utils``: records the modules + it unloads, and refuses to be asked about the device.""" + + def __init__(self) -> None: + self.unloaded: list = [] + + def unload_module(self, module): + self.unloaded.append(module) + + def get_device_properties(self, device): + raise AssertionError("the host compile asked the machine about its device") + + +@pytest.mark.parametrize("unloads", [False, True]) +def test_the_scoped_driver_unloads_through_the_machine_only_where_asked( + monkeypatch, unloads +): + """Inside a host compile the driver's ``utils`` are refused as a device + query, unless the compile's release unloads a collected kernel's module + through them (``unloads``): then ``unload_module`` reaches the machine's + driver and asks nothing, while the rest of ``utils`` is still refused.""" + from triton.runtime.driver import driver + + from tilelens.core.host_compile import _SCOPED_DRIVER + + machine = _Machine(CUDA89) + machine.utils = _MachineUtils() + _on_machine(monkeypatch, machine) + with _SCOPED_DRIVER.targeting(type(driver), CUDA80, unloads=unloads) as scoped: + if unloads: + driver.active.utils.unload_module("module") + assert machine.utils.unloaded == ["module"] and not scoped.queried + # The scope is back once the module is unloaded. + assert driver.active is scoped + with pytest.raises( + HostCompileUnavailable, match="asked its driver for 'utils.get_device" + ): + driver.active.utils.get_device_properties(0) + else: + with pytest.raises( + HostCompileUnavailable, match="asked its driver for 'utils':" + ): + driver.active.utils.unload_module("module") + assert machine.utils.unloaded == [] + assert scoped.queried + assert driver.active is machine and machine.queries == 0 + + +_COLLECTED: dict = {} + + +def _collect_a_loaded_kernel(): + # Called by the kernel below during its compile (see _device_is_zero): + # drops the last reference to a kernel a real launch had loaded, as the + # cyclic GC may at any point of a compile, which runs the kernel's + # CompiledKernel.__del__ right here, on the compiling thread. + _COLLECTED["kernels"].clear() + return True + + +_collect_a_loaded_kernel.__triton_builtin__ = True # type: ignore[attr-defined] + + +@needs_compiles +def test_a_kernel_collected_mid_compile_is_unloaded_by_the_machines_driver( + monkeypatch, +): + """Triton 3.8's CompiledKernel.__del__ unloads a loaded module through + ``driver.active``, and a real launch's kernel may be collected in the + middle of a host compile on the thread: the module goes back to the + machine's driver (none leaks), and the compile, failing after it, is not + taken for one whose front end asked the device (D27). A release whose + CompiledKernel has no such __del__ (3.6) unloads nothing.""" + from triton.compiler.compiler import CompiledKernel + from triton.runtime.jit import constexpr_function + + unloads = _for_release(_RELEASE_RUNTIME, "JIT runtime")["unloads_on_del"] + finalizer = CompiledKernel.__dict__.get("__del__") + assert (finalizer is not None) is unloads + machine = _Machine(CUDA89) + machine.utils = _MachineUtils() + _on_machine(monkeypatch, machine) + + class LoadedKernel: + # What CompiledKernel.__del__ reads of a kernel whose module is loaded. + function, name, metadata_group, hash = None, "loaded", {}, "0" * 64 + + def __init__(self) -> None: + self.module = "loaded module" + + if finalizer is not None: + LoadedKernel.__del__ = finalizer # type: ignore[attr-defined] + monkeypatch.setitem(_COLLECTED, "kernels", [LoadedKernel()]) + unraisable: list = [] + monkeypatch.setattr(sys, "unraisablehook", unraisable.append) + + @constexpr_function + def collected(): + return _collect_a_loaded_kernel() + + @triton.jit + def bounded(x_ptr, BLOCK: tl.constexpr): + tl.static_assert(collected() and BLOCK <= 32) + tl.store(x_ptr + tl.arange(0, BLOCK), 1.0) + + with pytest.raises(CompileTimeAssertionFailure) as raised: + HostCompiler().compile( + bounded, (torch.zeros(64),), {"BLOCK": 64}, target=CUDA89, stages={"ttir"} + ) + assert _COLLECTED["kernels"] == [] and unraisable == [] + assert machine.utils.unloaded == (["loaded module"] if unloads else []) + assert not target_queried(raised.value) + assert machine.queries == 0 + + +@needs_compiles +@pytest.mark.parametrize("rows", ["no row", "a row its code contradicts"]) +def test_a_custom_pipeline_on_an_unknown_release_runtime_is_refused( + no_driver, monkeypatch, rows +): + """A release with no _RELEASE_RUNTIMES row, or whose code does not do + what its row says, fails closed: its host compiles go on, unless a custom + pipeline is set, whose part in the kernel's name the host compile could + not mirror; and the driver's ``utils`` are refused.""" + from triton import knobs + + from tilelens.core import host_compile + + table = {} + if rows != "no row": + detected = _detected_runtime() + table[_release()] = host_compile._ReleaseRuntime( + stages_hook_keys=not detected["stages_hook_keys"], + unloads_on_del=not detected["unloads_on_del"], + ) + _clear_api_caches() + monkeypatch.setattr(host_compile, "_RELEASE_RUNTIMES", table) + try: + api = triton_api() + assert api.runtime is None and api.unloads is False + why = "has no row" if rows == "no row" else "does not do what" + assert why in api.runtime_unknown + call = (_make_copy(), (torch.zeros(64), torch.zeros(64), 64), {"BLOCK": 16}) + compiler = HostCompiler() + compiler.compile(*call, target=CUDA80, stages={"ttir"}) + monkeypatch.setattr( + knobs.runtime, "add_stages_inspection_hook", _PipelineHook("unknown") + ) + with pytest.raises( + HostCompileUnavailable, + match=rf"add_stages_inspection_hook is set, .*not known \(.*{why}", + ): + compiler.compile(*call, target=CUDA80, stages={"ttir"}) + finally: + monkeypatch.undo() + _clear_api_caches() diff --git a/tests/unit/ir/test_ir_capture.py b/tests/unit/ir/test_ir_capture.py new file mode 100644 index 000000000..108beeb64 --- /dev/null +++ b/tests/unit/ir/test_ir_capture.py @@ -0,0 +1,855 @@ +"""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, launched=False, 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, launched, 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)]) + + def on_refusal(self, refusal): + self.calls.append(("refusal", refusal)) + return IRVerdict(self.NAME, "unsupported", refusal=refusal) + + +# ======== 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, False + ) + 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 compile-only, config B compile-only, then A's real launch. + 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}, + kernel=kernel_a, + launched=True, + ) + ) + + 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_keys_on_the_triton_version(monkeypatch): + reader = _StubReader() + cache = ParseCache(reader) + cache.get("module") + monkeypatch.setattr(triton, "__version__", "3.7.0") + cache.get("module") + assert len(reader.calls) == 2 + + +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_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 (D20: 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/ir/test_lowering.py b/tests/unit/ir/test_lowering.py new file mode 100644 index 000000000..be1b6f832 --- /dev/null +++ b/tests/unit/ir/test_lowering.py @@ -0,0 +1,417 @@ +"""tilelens.ir.lowering: the Term -> Z3 lowering shared by the compiled-mode +clients (D13). + +CPU only: terms are built by hand over small AccessGraphs and lowered with +a test TermLeaves whose leaves are free Z3 variables (or constants), then +checked by Z3 equivalence or by evaluation; kernel_deep_chain comes from +the golden the installed release reads. The compiled sanitizer's own +results on top of the lowering are pinned in tests/unit/sanitizer_compiled/. +""" + +from __future__ import annotations + +import sys + +import pytest +import z3 + +from tilelens.ir.lowering import Lowerer, children, fold +from tilelens.ir.ttir_reader import ( + AccessGraph, + Arange, + Bin, + BoolBin, + Cmp, + Const, + DataDep, + IntCast, + IterArgInfo, + IterArgOffset, + LoopInfo, + LoopVar, + Not, + NumPrograms, + Observed, + Param, + Pid, + PtrValue, + Select, + parse_ttir, +) + +from . import _goldens as G + +LOOP = LoopInfo("%loop", "%i", lower=Param("lo"), upper=Param("hi"), step=Param("st")) + + +class _Refusal(Exception): + """A client's own refusal, raised from a leaf.""" + + +class _Leaves: + """Every leaf a free Int of ``ctx`` named after it (a Param in + ``params`` the constant instead); records the names it made.""" + + def __init__(self, ctx: z3.Context | None = None, params: dict | None = None): + self.ctx = ctx + self.params = params or {} + self.made: list[str] = [] + + def _var(self, name: str) -> z3.ArithRef: + self.made.append(name) + return z3.Int(name, self.ctx) + + def param(self, t): + if t.name in self.params: + return z3.IntVal(self.params[t.name], self.ctx) + return self._var(t.name) + + def pid(self, t): + return self._var(f"pid_{t.axis}") + + def num_programs(self, t): + return self._var(f"grid_{t.axis}") + + def arange(self, t): + return self._var(f"arange_{t.start}_{t.end}_d{t.dim}") + + def iteration(self, loop_ssa): + return self._var(f"k{loop_ssa}") + + def observed(self, t): + return self._var(f"observed_{t.access_index}") + + def data_dep(self, t): + raise _Refusal(t.why) + + +def _graph(loop: LoopInfo | None = None, iter_args=()) -> AccessGraph: + return AccessGraph("k", (), (), loop, iter_args) + + +def _lower(term, graph: AccessGraph | None = None, leaves: _Leaves | None = None): + return Lowerer(graph or _graph(LOOP), leaves or _Leaves()).lower(term) + + +def _v(name: str, ctx: z3.Context | None = None) -> z3.ArithRef: + return z3.Int(name, ctx) + + +def _proved(claim) -> bool: + solver = z3.Solver(ctx=claim.ctx) + solver.add(z3.Not(claim)) + return solver.check() == z3.unsat + + +def _eval(term) -> int | bool: + e = z3.simplify(_lower(term)) + if z3.is_bool(e): + assert z3.is_true(e) or z3.is_false(e), e + return z3.is_true(e) + return e.as_long() + + +X, Y = Param("x"), Param("y") + + +# ─────────────────────────── the Z3 context ─────────────────────────── + + +def _every_kind() -> tuple[object, AccessGraph]: + """One term reaching every kind of the algebra but DataDep.""" + graph = _graph(LOOP, (IterArgInfo(0, "p", Param("o0"), Const(4), "%loop"),)) + lane = Bin("+", Arange("%r", 0, 16, dim=0), IntCast("extsi", 32, 64, Pid(1))) + moved = Bin("*", Bin("//", IterArgOffset(0), NumPrograms(0)), LoopVar("%loop")) + cond = BoolBin("or", Cmp("ult", X, Const(8)), Not(Cmp("eq", Observed(2), Y))) + return Select(cond, Bin("umin", lane, moved), Bin("%", X, Const(3))), graph + + +@pytest.mark.parametrize("own_context", [False, True]) +def test_terms_are_made_in_the_leaves_context(own_context, monkeypatch): + """ctx=None lowers into Z3's main context, a given Context into that + Context (a constant included), and the lowering creates none.""" + ctx = z3.Context() if own_context else None + expected = ctx if own_context else z3.main_ctx() + term, graph = _every_kind() + + def no_context(*args, **kwargs): + raise AssertionError("the lowering created a z3.Context") + + monkeypatch.setattr(z3.Context, "__init__", no_context) + lowerer = Lowerer(graph, _Leaves(ctx)) + lowered = [lowerer.lower(term), lowerer.value(term), lowerer.cond(term)] + lowered += [lowerer.lower(Const(7)), lowerer.cond(Const(1))] + lowered += [lowerer.value(Cmp("slt", X, Y)), lowerer.cond(Not(Const(0)))] + monkeypatch.undo() + assert all(e.ctx is expected for e in lowered) + assert z3.is_int(lowered[0]) and z3.is_bool(lowered[2]) + + +def test_main_context_terms_take_a_main_context_substitution(): + """A client that renames variables with z3.substitute over pairs of the + main context gets its rename on terms lowered with ctx=None.""" + lowered = _lower(Bin("*", Pid(1), Const(3))) + renamed = z3.substitute(lowered, (_v("pid_1"), _v("pid_1_copy"))) + assert _proved(renamed == _v("pid_1_copy") * 3) + + +# ─────────────────────────── iteration and the memo ─────────────────────────── + + +def test_a_chain_deeper_than_the_recursion_limit_lowers(): + graph = _graph() + chain = Pid(0) + for _ in range(5000): + chain = Bin("+", chain, Const(1), 32) + limit = sys.getrecursionlimit() + sys.setrecursionlimit(1000) + try: + lowered = Lowerer(graph, _Leaves()).value(chain) + finally: + sys.setrecursionlimit(limit) + assert _proved(lowered == _v("pid_0") + 5000) + + +def test_kernel_deep_chain_lowers_at_the_default_recursion_limit(): + """kernel_deep_chain's offset nests more than 1000 levels (N = 600 in + generate_ttir.py: off = off * s + pid, 600 times from off = pid), where + the generated == / hash raise at Python's default limit.""" + text = G.texts("ttir")["kernel_deep_chain.ttir"].read_text(encoding="utf-8") + graph = parse_ttir(text) + (store,) = graph.accesses + applied: list[object] = [] + limit = sys.getrecursionlimit() + sys.setrecursionlimit(1000) + try: + offset = Lowerer(graph, _Leaves(params={"s": 1})).value(store.offset) + fold(store.offset, graph, lambda t, _: applied.append(t), {}) + finally: + sys.setrecursionlimit(limit) + assert sum(isinstance(t, Bin) for t in applied) > 1000 + assert _proved(offset == 601 * _v("pid_0")) + + +def test_each_term_is_lowered_once_by_identity(): + """The memo is keyed by identity: one Param object read twice is one + leaf call, two equal Param objects are two; a Lowerer keeps its memo + across calls.""" + leaves = _Leaves() + lowerer = Lowerer(_graph(), leaves) + n = Param("n") + shared = Bin("+", n, n) + lowerer.lower(shared) + lowerer.lower(Bin("*", shared, n)) + assert leaves.made == ["n"] + lowerer.lower(Bin("+", Param("n"), Param("n"))) + assert leaves.made == ["n", "n", "n"] + + +def test_fold_applies_children_first_and_each_term_once(): + graph = _graph(LOOP, (IterArgInfo(0, "p", Const(2), Const(3)),)) + two = Const(2) + root = Bin("+", Bin("*", two, two), IterArgOffset(0)) + order: list[object] = [] + + def apply(t, values): + order.append(t) + if isinstance(t, Const): + return t.value + if isinstance(t, IterArgOffset): + return values[0] + 10 * values[1] # iteration 10 + return values[0] * values[1] if t.op == "*" else values[0] + values[1] + + memo: dict = {} + assert fold(root, graph, apply, memo) == 2 * 2 + (2 + 10 * 3) + ids = [id(t) for t in order] + assert len(ids) == 6 and ids.count(id(two)) == 1 + assert ids.index(id(two)) < ids.index(id(root.a)) < ids.index(id(root)) + # the memo holds every term: a second fold applies nothing + assert fold(root, graph, apply, memo) == 36 and len(order) == 6 + + +def test_children_is_the_value_dependency_relation(): + info = IterArgInfo(0, "p", Param("o0"), Param("d")) + graph = _graph(LOOP, (info,)) + less = Cmp("slt", X, Y) + + def kids(t) -> list[int]: + return [id(k) for k in children(t, graph)] + + assert ( + kids(Bin("+", X, Y)) + == kids(less) + == kids(BoolBin("or", X, Y)) + == [id(X), id(Y)] + ) + assert kids(Select(less, X, Y)) == [id(less), id(X), id(Y)] + assert kids(Not(less)) == kids(IntCast("trunci", 64, 32, less)) == [id(less)] + assert kids(IterArgOffset(0)) == [id(info.offset0), id(info.delta)] + assert kids(LoopVar("%loop")) == [id(LOOP.lower), id(LOOP.step)] + # a DataDep's keep is not its value + assert kids(DataDep(keep=less)) == [] + for leaf in (Const(1), Pid(0), NumPrograms(1), Arange("%r", 0, 4), X, Observed(0)): + assert kids(leaf) == [] + + +# ─────────────────────────── operator semantics ─────────────────────────── + + +@pytest.mark.parametrize( + "unsigned, signed", [("u//", "//"), ("u%", "%"), ("umin", "min"), ("umax", "max")] +) +def test_unsigned_ops_read_as_their_signed_twins(unsigned, signed): + lowerer = Lowerer(_graph(), _Leaves()) + assert _proved( + lowerer.lower(Bin(unsigned, X, Y)) == lowerer.lower(Bin(signed, X, Y)) + ) + + +@pytest.mark.parametrize( + "unsigned, signed", [("ult", "slt"), ("ule", "sle"), ("ugt", "sgt"), ("uge", "sge")] +) +def test_unsigned_predicates_read_as_their_signed_twins(unsigned, signed): + lowerer = Lowerer(_graph(), _Leaves()) + got, twin = lowerer.lower(Cmp(unsigned, X, Y)), lowerer.lower(Cmp(signed, X, Y)) + assert z3.is_bool(got) and _proved(got == twin) + + +@pytest.mark.parametrize( + "pred, expected", + [ + ("slt", lambda x, y: x < y), + ("sle", lambda x, y: x <= y), + ("sgt", lambda x, y: x > y), + ("sge", lambda x, y: x >= y), + ("eq", lambda x, y: x == y), + ("ne", lambda x, y: x != y), + ], +) +def test_predicates(pred, expected): + assert _proved(_lower(Cmp(pred, X, Y)) == expected(_v("x"), _v("y"))) + + +@pytest.mark.parametrize("kind", ["trunci", "extsi", "extui"]) +def test_an_int_cast_reads_as_its_operand(kind): + assert z3.eq(_lower(IntCast(kind, 64, 32, X)), _v("x")) + # of an i1: the compare's 0/1 + cast = _lower(IntCast(kind, 1, 32, Cmp("slt", X, Y))) + assert z3.is_int(cast) and _proved(cast == z3.If(_v("x") < _v("y"), 1, 0)) + + +def test_i1_values_are_bool_or_0_1_by_position(): + x, y = _v("x"), _v("y") + lowerer = Lowerer(_graph(), _Leaves()) + less = Cmp("slt", X, Y) + # an integer position: 0/1 + assert _proved(lowerer.lower(Bin("+", less, Const(1))) == z3.If(x < y, 1, 0) + 1) + assert _proved(lowerer.value(less) == z3.If(x < y, 1, 0)) + # a boolean position: an i1 constant (dense) and an Int are != 0 + assert _proved(lowerer.lower(BoolBin("and", Const(1), less)) == (x < y)) + assert _proved(lowerer.lower(BoolBin("or", X, Const(0))) == (x != 0)) + assert _proved(lowerer.lower(Not(Const(0)))) + assert _proved(lowerer.cond(X) == (x != 0)) + # a compare of an i1 with an integer reads the i1 as 0/1 + assert _proved(lowerer.lower(Cmp("eq", less, Const(1))) == (x < y)) + # a Select: its condition is boolean; Bool arms stay Bool + both = lowerer.lower(Select(X, less, Cmp("eq", X, Y))) + assert z3.is_bool(both) and _proved(both == z3.If(x != 0, x < y, x == y)) + # arms of different sorts are Int + mixed = lowerer.lower(Select(less, Cmp("eq", X, Const(0)), Const(5))) + assert z3.is_int(mixed) + assert _proved(mixed == z3.If(x < y, z3.If(x == 0, 1, 0), 5)) + + +@pytest.mark.parametrize( + "op, a, b, expected", + [ + ("min", 3, -2, -2), + ("min", -2, 3, -2), + ("max", 3, -2, 3), + ("max", -2, 3, 3), + ("umin", 4, 9, 4), + ("umax", 4, 9, 9), + ("+", 7, -9, -2), + ("-", 7, -9, 16), + ("*", 7, -9, -63), + ], +) +def test_arithmetic(op, a, b, expected): + assert _eval(Bin(op, Const(a), Const(b), 32)) == expected + + +@pytest.mark.parametrize( + "a, b, quotient, remainder", + [ + (7, 2, 3, 1), + (-7, 2, -3, -1), + (7, -2, -3, 1), + (-7, -2, 3, -1), + (6, 3, 2, 0), + (-6, 3, -2, 0), + (0, -5, 0, 0), + (-1, 5, 0, -1), + ], +) +def test_division_truncates_toward_zero(a, b, quotient, remainder): + """arith.divsi / remsi (and divui / remui on the non-negative operands + their obligations leave): the remainder has the dividend's sign.""" + assert _eval(Bin("//", Const(a), Const(b), 32)) == quotient + assert _eval(Bin("%", Const(a), Const(b), 32)) == remainder + if a >= 0 and b >= 0: + assert _eval(Bin("u//", Const(a), Const(b), 32)) == quotient + assert _eval(Bin("u%", Const(a), Const(b), 32)) == remainder + + +def test_loop_terms_read_the_clients_iteration(): + """LoopVar is lower + k * step, IterArgOffset offset0 + k * delta, with + k the leaves' iteration of that loop: an IterArgInfo without a loop_ssa + is the graph's loop's.""" + graph = _graph( + LOOP, + ( + IterArgInfo(0, "p", Param("o0"), Param("d0")), + IterArgInfo(1, "p", Param("o1"), Param("d1"), "%other"), + ), + ) + leaves = _Leaves() + lowerer = Lowerer(graph, leaves) + k, other = _v("k%loop"), _v("k%other") + assert _proved(lowerer.lower(LoopVar("%loop")) == _v("lo") + k * _v("st")) + assert _proved(lowerer.lower(IterArgOffset(0)) == _v("o0") + k * _v("d0")) + assert _proved(lowerer.lower(IterArgOffset(1)) == _v("o1") + other * _v("d1")) + assert "hi" not in leaves.made # the upper bound is not the variable's value + + +# ─────────────────────────── leaves and bugs ─────────────────────────── + + +def test_a_leaf_refusal_propagates_unchanged(): + leaves = _Leaves() + with pytest.raises(_Refusal, match="loaded value"): + Lowerer(_graph(), leaves).lower(BoolBin("and", X, DataDep("loaded value"))) + # a DataDep's keep is not lowered + with pytest.raises(_Refusal): + Lowerer(_graph(), leaves).cond(DataDep(keep=Cmp("slt", Param("kept"), Y))) + assert "kept" not in leaves.made + + +@pytest.mark.parametrize( + "term, error, match", + [ + (object(), TypeError, "unknown term object"), + (PtrValue("p", Const(0)), TypeError, "unknown term PtrValue"), + (Bin("**", X, Y), ValueError, "unknown integer op"), + (Cmp("olt", X, Y), ValueError, "unknown cmpi predicate"), + (BoolBin("xor", X, Y), ValueError, "unknown boolean op"), + ], +) +def test_a_term_outside_the_algebra_is_a_bug(term, error, match): + with pytest.raises(error, match=match): + _lower(term) + + +@pytest.mark.parametrize("term", [LoopVar("%loop"), IterArgOffset(0)]) +def test_a_loop_term_without_a_loop_is_a_bug(term): + graph = _graph(None, (IterArgInfo(0, "p", Const(0), Const(1)),)) + with pytest.raises(ValueError, match="without a loop"): + Lowerer(graph, _Leaves()).lower(term) diff --git a/tests/unit/ir/test_mlir_walk.py b/tests/unit/ir/test_mlir_walk.py new file mode 100644 index 000000000..b53080318 --- /dev/null +++ b/tests/unit/ir/test_mlir_walk.py @@ -0,0 +1,1854 @@ +"""tilelens.ir._mlir_walk: the bindings + text alignment layer under the TTIR reader. + +Goldens live in tests/golden/ir/ttir/ (printed by Triton 3.6, read under every +release) and tests/golden/ir/ttir_/ (a later release's own printing, +which shadows the base golden of the same name under that release; see +_goldens.py; provenance: tests/golden/ir/generate_ttir.py). Their pinned counts +and text-only attribute census are pinned per printing release, in +tests/golden/ir/expected.json (3.6) and expected_.json. Regenerate the +pins of the installed release's own goldens after an intended change with +``TILELENS_IR_REGEN=1 pytest tests/unit/ir/test_mlir_walk.py -k regen``. +""" + +from __future__ import annotations + +import collections +import copy +import dataclasses +import json +import os +import pickle +import random +import re +import subprocess +import sys +import tempfile +import threading +from pathlib import Path + +import pytest + +from tilelens.ir import _mlir_walk as W + +from . import _goldens as G + +REPO = G.REPO +GOLDEN = G.GOLDEN +TTIR = GOLDEN / "ttir" +# fail closed by design: the generic op form the printer emits only for +# modules that fail to verify +MISALIGNED = {"crafted_generic_form.ttir"} +# name -> the golden the installed release reads (its own printing first) +GOLDENS = G.texts("ttir") +FILES = sorted(GOLDENS) +ALIGNED = [f for f in FILES if f not in MISALIGNED] +BASE_FILES = sorted(p.name for p in TTIR.glob("*.ttir")) +# Every pinned text the installed release parses, checked against the pins +# of the release that printed it: the base goldens (but those its parser +# rejects, BASE_UNPARSABLE) and its own; the base ones keep their names. +_REFUSED_HERE = G.BASE_UNPARSABLE.get(G.RELEASE, {}) +PINNED: dict[str, Path] = { + **{ + n: TTIR / n + for n in BASE_FILES + if n not in MISALIGNED and n not in _REFUSED_HERE + }, + **{ + f"{p.parent.name}/{n}": p + for n, p in GOLDENS.items() + if G.printed_by(p) != G.BASE_RELEASE and n not in MISALIGNED + }, +} + + +def _text(name: str) -> str: + return GOLDENS[name].read_text(encoding="utf-8") + + +def _printed_by(name: str) -> str: + return G.printed_by(GOLDENS[name]) + + +@pytest.fixture(autouse=True) +def _fresh_cache(): + W._CACHE.clear() + yield + W._CACHE.clear() + + +# ─────────────────────────── goldens ─────────────────────────── + + +def _text_only(table: W.Printer) -> dict[str, tuple[str, ...]]: + """The text-only attributes of ``table``'s release: the text is their + only source, so golden pins guard their extraction (spike condition 5).""" + bound = {(o, a) for o, g in table.bind_attrs.items() for _, a in g} + out = { + op: tuple(k for k in keys if (op, k) not in bound) + for op, keys in table.needed.items() + if op not in ("cf.br", "cf.cond_br", "tt.func", "tt.call") + } + out["arith.constant"] = ("value", "splat") + out["tt.descriptor_reduce"] = ("kind",) + return out + + +def _census(m: W.Module) -> list[list]: + text_only = _text_only(W.PRINTERS[m.release]) + c: collections.Counter[str] = collections.Counter() + for op in m.ops: + for key in text_only.get(op.name, ()): + if key in op.attrs: + c[f"{op.name}.{key}={op.attrs[key]!r}"] += 1 + if not op.results and not op.implicit and op.loc is not None: + c["zero-result op with a text loc"] += 1 + if "successors" in op.attrs: + c[f"{op.name} -> {len(op.attrs['successors'])} successors"] += 1 + return sorted([k, n] for k, n in c.items()) + + +def _pins(m: W.Module) -> dict: + return { + "stats": dict(m.stats), + "census": _census(m), + "funcs": [f.sym_name for f in m.funcs], + } + + +@pytest.mark.skipif( + not os.environ.get("TILELENS_IR_REGEN"), + reason="set TILELENS_IR_REGEN=1 to rewrite pins", +) +def test_regen_pins(): + """The pins of the goldens the installed release printed (its own + directory; the base directory under the base release).""" + own = [f for f in ALIGNED if _printed_by(f) == G.RELEASE] + assert own, f"Triton {G.RELEASE} printed no golden: run generate_ttir.py first" + G.pins_path(G.RELEASE).write_text( + json.dumps( + {f: _pins(W._walk(_text(f))) for f in own}, indent=1, ensure_ascii=False + ) + + "\n" + ) + + +def _generator(): + import importlib.util + + spec = importlib.util.spec_from_file_location( + "_generate_ttir", GOLDEN / "generate_ttir.py" + ) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +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) + + +# Where a golden's locs name the generator: the checkout's own path. +_GENERATOR_LOC = re.compile(r'loc\("[^"]*generate_ttir\.py"') +# Where they name Triton's own sources (tl.cdiv, tl.zeros, ...): the path +# Triton is installed at, which differs between machines. +_TRITON_LOC = re.compile(r'loc\("[^"]*/triton/') + + +def _portable(ttir: str) -> str: + """``ttir`` with the paths that depend on the machine taken out.""" + return _TRITON_LOC.sub('loc("/', _GENERATOR_LOC.sub('loc("G"', ttir)) + + +@pytest.mark.skipif( + not _real_compiles_available(), + reason="Triton was imported under TRITON_INTERPRET=1: nothing compiles in-process", +) +def test_the_kernel_goldens_regenerate_byte_for_byte(monkeypatch, tmp_path): + """generate_ttir.py prints, under the installed release, exactly the + goldens it wrote into that release's directory (the generator's and + Triton's install paths aside): its locs name the kernels' lines, so a + line added above them shows here, not as goldens that no longer + regenerate.""" + # tests/unit/test_multithreading.py sets TRITON_INTERPRET=1 at import + # time, under which @triton.jit builds InterpretedFunctions: pin the knob + # off while the generator's kernels are built, as the compile tests do. + from triton import knobs + + monkeypatch.delenv("TRITON_INTERPRET", raising=False) + missing = object() + previous = knobs.runtime.__dict__.get("interpret", missing) + knobs.runtime.__dict__["interpret"] = False + try: + gen = _generator() + finally: + if previous is missing: + knobs.runtime.__dict__.pop("interpret", None) + else: + knobs.runtime.__dict__["interpret"] = previous + monkeypatch.setenv("TRITON_CACHE_DIR", str(tmp_path)) + out = Path(gen.out_dir(G.RELEASE)) + todo = gen.jobs(G.RELEASE) + assert todo and all((out / f"{name}.ttir").is_file() for name in todo) + # Its compile takes ~25 s under 3.8's pipeline; a line added above it + # moves the locs of kernel_dot_scaled, below it, too. + del todo["kernel_deep_chain"] + for name, spec in todo.items(): + want = (out / f"{name}.ttir").read_text(encoding="utf-8") + got = gen.ttir(spec) + assert _portable(got) == _portable(want), name + + +@pytest.mark.parametrize("label", PINNED) +def test_golden_aligns_with_pinned_counts(label): + path = PINNED[label] + name = path.name + pinned = G.pins(G.printed_by(path)) + m = W._walk(path.read_text(encoding="utf-8")) + assert name in pinned, "new golden: regenerate the pins" + assert _pins(m) == pinned[name] + # the counters agree with the records + assert m.stats["ops"] == len(m.ops) and m.stats["values"] == len(m.values) + assert m.stats["ssa_edges"] == sum(len(op.operands) for op in m.ops) + assert m.stats["implicit_ops"] == sum(op.implicit for op in m.ops) + + +def test_every_golden_is_pinned(): + # every release's own goldens are pinned by that release (checked here + # under any release: no walk), and each shadows a base golden + for release in G.releases_with_goldens("ttir"): + own = sorted(p.name for p in G.own_dir("ttir", release).glob("*.ttir")) + assert sorted(G.pins(release)) == [f for f in own if f not in MISALIGNED] + assert set(own) <= set(BASE_FILES), release + # a base golden a release's parser rejects is one it prints itself + for release, refused in G.BASE_UNPARSABLE.items(): + own = {p.name for p in G.own_dir("ttir", release).glob("*.ttir")} + assert set(refused) <= own & set(BASE_FILES), release + # the pinned corpus: #361 goldens, spike corpus, review corpora, kernels, crafted + prefixes = collections.Counter(f.split("_", 1)[0] for f in BASE_FILES) + assert prefixes == { + "golden": 35, + "spike": 15, + "adv": 12, + "nat": 7, + "kernel": 5, + "crafted": 12, + } + + +def test_base_goldens_a_release_respells_do_not_parse_under_it(): + """A base golden in BASE_UNPARSABLE is refused by the installed + release's parser (never misread), and the release reads its own + printing instead.""" + for name, fragment in G.BASE_UNPARSABLE.get(G.RELEASE, {}).items(): + with pytest.raises(W.ModuleParseError) as e: + W._walk((TTIR / name).read_text(encoding="utf-8")) + assert fragment in e.value.diagnostic and e.value.line_no is not None + assert _printed_by(name) == G.RELEASE, name + + +def test_corpus_totals(): + totals: collections.Counter[str] = collections.Counter() + for name in ALIGNED: + totals.update(W._walk(_text(name)).stats) + assert totals["ops"] > 3500 and totals["ssa_edges"] > 5000 + assert ( + totals["implicit_ops"] >= 20 + and totals["cf_edges"] >= 30 + and totals["pred_checks"] >= 20 + ) + + +@pytest.mark.parametrize("name", ["crafted_generic_form.ttir"]) +def test_generic_form_fails_closed(name): + with pytest.raises(W.MisalignedModule) as e: + W._walk(_text(name)) + # the generic form's integer enums never satisfy the text-only extractors + assert any("arith.cmpi" in p for p in e.value.problems) + assert any("tt.atomic_rmw" in p for p in e.value.problems) + assert e.value.line_no == 3 + + +def test_structure_records_are_consistent(): + for name in ALIGNED: + m = W._walk(_text(name)) + assert m.ops[0].name == "builtin.module" and m.ops[0].path == () + assert [op.index for op in m.ops] == list(range(len(m.ops))) + assert [b.index for b in m.blocks] == list(range(len(m.blocks))) + assert [v.index for v in m.values] == list(range(len(m.values))) + for op in m.ops: + for v in op.results: + assert m.values[v].op == op.index + for k, blocks in enumerate(op.regions): + for pos, b in enumerate(blocks): + blk = m.blocks[b] + assert (blk.op, blk.region, blk.position) == (op.index, k, pos) + for i in blk.ops: + assert m.ops[i].path == op.path + (b,) + if op.path: + assert m.blocks[op.block].ops[op.position] == op.index + for v in m.values: + assert (v.op is None) != (v.block is None) + # pre-order: an op's regions come after it, its uses are in-module values + for op in m.ops: + assert all(0 <= v < len(m.values) for v in op.operands) + assert op.operand_types == tuple(m.values[v].type for v in op.operands) + + +# ─────────────────────────── an independent attribute extractor ─────────────────────────── + + +def _must(fn, pattern: str, s: str) -> re.Match: + m = fn(pattern, s) + assert m is not None, (pattern, s) + return m + + +def _indep(name: str, line: str) -> dict: + """Deliberately naive per-op regexes over the raw header line (strings kept): + a second extractor for the text-only attributes (the review's differential).""" + body = line.split(" = ", 1)[1] if re.match(r"\s*%[^=]*= ", line) else line.strip() + out: dict = {} + if name == "arith.cmpi" or name == "arith.cmpf": + out["predicate"] = _must(re.match, r"arith\.cmp[if] (\w+),", body).group(1) + elif name in ("tt.get_program_id", "tt.get_num_programs"): + out["axis"] = "xyz".index( + _must(re.match, r"tt\.get_\w+ ([xyz]) ", body).group(1) + ) + elif name == "tt.make_range": + out["start"] = int(_must(re.search, r"start = (-?\d+) : i32", body).group(1)) + out["end"] = int(_must(re.search, r"end = (-?\d+) : i32", body).group(1)) + elif name == "tt.atomic_rmw": + out.update( + zip( + ("rmw_op", "sem", "scope"), + _must( + re.match, r"tt\.atomic_rmw (\w+), (\w+), (\w+), %", body + ).groups(), + ) + ) + elif name == "tt.atomic_cas": + out.update( + zip( + ("sem", "scope"), + _must(re.match, r"tt\.atomic_cas (\w+), (\w+), %", body).groups(), + ) + ) + elif name in ("tt.expand_dims", "tt.reduce", "tt.scan"): + out["axis"] = int(_must(re.search, r"axis = (-?\d+) : i32", body).group(1)) + elif name == "tt.trans": + out["order"] = tuple( + int(x) + for x in _must(re.search, r"order = array", body) + .group(1) + .split(",") + ) + elif name == "tt.dot": + m = re.search(r", inputPrecision = (\w+) :", body) + out["inputPrecision"] = m.group(1) if m else "ieee" + elif name == "tt.reshape": + out["allow_reorder"] = " allow_reorder " in body + elif name == "scf.for": + out["unsignedCmp"] = body.startswith("scf.for unsigned ") + elif name == "arith.constant": + c = re.match( + r"arith\.constant (?:\{[^}]*\} )?(dense<)?(-?\d+|true|false)>? : ", + body + " : ", + ) + if c and c.group(2) in ("true", "false"): + out["value"] = c.group(2) == "true" + elif c: + out["value"] = int(c.group(2)) + return out + + +def test_independent_extractor_agrees_on_every_golden(): + checked = 0 + for name in ALIGNED: + lines = _text(name).splitlines() + for op in W._walk(_text(name)).ops: + if op.implicit or op.line_no is None: + continue + for k, v in _indep(op.name, lines[op.line_no - 1]).items(): + assert op.attrs.get(k) == v, (name, op.line_no, op.name, k) + checked += 1 + assert checked >= 700 + + +# ─────────────────────────── mutation sensitivity ─────────────────────────── +# The bindings parse the original text while the text layer reads a mutated +# copy: a text layer that mis-reads the module must never align silently. + +_OPLINE = re.compile( + r"^\s+(?:%[-\w.$]+(?::\d+)?(?:, %[-\w.$]+)* = )?[a-z_]+\.[\w.]+ .*loc\(.*\)\s*$" +) +_TERMINATOR = re.compile( + r"tt\.return|scf\.yield|cf\.br|cf\.cond_br|scf\.condition|reduce\.return|scan\.return" +) + + +def _mutants(text: str, rng: random.Random) -> list[tuple[str, str, str]]: + lines = text.splitlines() + idx = [ + i + for i, ln in enumerate(lines) + if _OPLINE.match(ln) and not ln.rstrip().endswith("{") + ] + out = [] + pairs = [ + (i, j) + for i, j in zip(idx, idx[1:]) + if j == i + 1 + and not _TERMINATOR.search(lines[j]) + and lines[i].split(" loc(")[0] != lines[j].split(" loc(")[0] + ] + same = [ + (i, j) + for i, j in pairs + if lines[i].split("=")[-1].split()[0] == lines[j].split("=")[-1].split()[0] + ] + for label, pool in (("swap-same-name", same), ("swap-any", pairs)): + for i, j in rng.sample(pool, min(2, len(pool))): + m = lines[:] + m[i], m[j] = m[j], m[i] + out.append((label, f"lines {i + 1}<->{j + 1}", "\n".join(m))) + for i in rng.sample(idx, min(2, len(idx))): + out.append( + ("drop-line", f"line {i + 1}", "\n".join(lines[:i] + lines[i + 1 :])) + ) + defs = re.findall(r"^\s+(%[-\w.$]+) = ", text, re.M) + use_lines = [i for i in idx if re.search(r"= [a-z_.]+ .*%", lines[i])] + for i in rng.sample(use_lines, min(2, len(use_lines))): + lhs, rhs = lines[i].split(" = ", 1) + uses = re.findall(r"%[-\w.$]+", rhs) + cands = [d for d in defs if uses and d != uses[0]] + if not cands: + continue + m = lines[:] + m[i] = ( + lhs + + " = " + + re.sub( + re.escape(uses[0]) + r"(?![-\w.$])", rng.choice(cands), rhs, count=1 + ) + ) + out.append(("repoint-use", f"line {i + 1}", "\n".join(m))) + edits = ( + ( + "typed-attr", + r"(tt\.make_range \{end = )(\d+)", + lambda g: f"{g.group(1)}{int(g.group(2)) * 2}", + ), + ( + "vocab-predicate", + r"(arith\.cmpi )(\w+)(,)", + lambda g: f"{g.group(1)}{g.group(2)}x{g.group(3)}", + ), + ( + "vocab-sem", + r"(tt\.atomic_\w+ (?:\w+, )?)(relaxed|acquire|release|acq_rel)(,)", + lambda g: f"{g.group(1)}seq_cst,", + ), + ( + "vocab-scope", + r"(tt\.atomic_\w+ (?:\w+, )?\w+, )(gpu|cta|sys)(,)", + lambda g: f"{g.group(1)}device,", + ), + ("vocab-rmw", r"(tt\.atomic_rmw )(\w+)(,)", lambda g: f"{g.group(1)}nand,"), + ( + "vocab-axis", + r"(tt\.get_program_id )([xyz])( :)", + lambda g: f"{g.group(1)}w :", + ), + ("vocab-precision", r"(inputPrecision = )(\w+)", lambda g: f"{g.group(1)}fp8"), + ( + "pred-comment", + r"(// pred: \^bb)(\d+)", + lambda g: f"{g.group(1)}{int(g.group(2)) + 7}", + ), + # a NEEDED attribute the text no longer prints + ( + "drop-needed", + r"(tt\.make_range \{end = \d+ : i32), start = \d+ : i32\}", + lambda g: g.group(1) + "}", + ), + ( + "drop-needed", + r"(tt\.trans %[-\w.$]+) \{order = array\}", + lambda g: g.group(1), + ), + ) + for label, pat, rep in edits: + for i, ln in enumerate(lines): + if re.search(pat, ln): + m = lines[:] + m[i] = re.sub(pat, rep, ln, count=1) + out.append((label, f"line {i + 1}", "\n".join(m))) + break + # a result-bearing op's loc: the bindings' result locs are its second + # source (Op.loc / callers / loc_name, and Value.name of the results) + aliases = dict(re.findall(r"^(#loc\d*) = loc\((.*)\)$", text, re.M)) + resolve = W._LocParser(aliases, "") + trees = {a: resolve.parse(f"loc({a})") for a in aliases} + sited: list[tuple[int, re.Match[str]]] = [] + for i in idx: + lm = re.search(r" loc\((#loc\d*)\)$", lines[i].rstrip()) + if lm and re.match(r"\s+%", lines[i]) and lm.group(1) in trees: + sited.append((i, lm)) + for i, lm in rng.sample(sited, min(2, len(sited))): + others = sorted(a for a in trees if trees[a] != trees[lm.group(1)]) + ln = lines[i].rstrip() + mt = lines[:] + mt[i] = ln[: lm.start(1)] + rng.choice(others) + ln[lm.end(1) :] + out.append(("repoint-loc", f"line {i + 1}", "\n".join(mt))) + def_line = { + a: k for k, ln in enumerate(lines) for a in re.findall(r"^(#loc\d*) =", ln) + } + for i, lm in rng.sample(sited, min(2, len(sited))): + k = def_line[lm.group(1)] + garbled = re.sub( + r'(":)(\d+)(:\d+\)\s*)$', + lambda g: f"{g.group(1)}{int(g.group(2)) + 1000}{g.group(3)}", + lines[k], + ) + if garbled == lines[k]: # a name loc: rename it + garbled = re.sub(r'^(#loc\d* = loc\(")', r"\1garbled_", lines[k]) + if garbled == lines[k]: # loc(unknown), callsite(...), fused[...] + continue + mt = lines[:] + mt[k] = garbled + out.append( + ("garble-loc-alias", f"line {k + 1} (used on line {i + 1})", "\n".join(mt)) + ) + return out + + +# the problem each targeted mutation class must be caught by (the structural +# classes may trip any of several checks) +_REASON = { + "typed-attr": "make_range", + "vocab-predicate": "closed vocabulary", + "vocab-sem": "closed vocabulary", + "vocab-scope": "closed vocabulary", + "vocab-rmw": "closed vocabulary", + "vocab-axis": "closed vocabulary", + "vocab-precision": "closed vocabulary", + "pred-comment": "printer preds", + "drop-needed": "not recovered", + "repoint-loc": "loc differs", + "garble-loc-alias": "loc differs", +} + + +def test_mutations_are_detected(): + rng = random.Random(0) + by_class: collections.Counter[str] = collections.Counter() + missed = [] + for name in ALIGNED: + text = _text(name) + if name == "kernel_deep_chain.ttir": + continue # 1200-op chain: slow and adds nothing here + for cls, desc, mt in _mutants(text, rng): + by_class[cls] += 1 + try: + W._walk(text, scan_text=mt) + except W.MisalignedModule as e: + assert e.problems and isinstance(e.line_no, (int, type(None))) + reason = _REASON.get(cls) + if reason is None or any(reason in p for p in e.problems): + continue + missed.append(f"{name} {cls} {desc}: caught only by {e.problems[:2]}") + continue + missed.append(f"{name} {cls} {desc}") + assert not missed + for cls in ( + "swap-same-name", + "swap-any", + "drop-line", + "repoint-use", + "typed-attr", + "vocab-predicate", + "drop-needed", + "repoint-loc", + "garble-loc-alias", + ): + assert by_class[cls] >= 20, by_class + for cls in ( + "vocab-sem", + "vocab-scope", + "vocab-rmw", + "vocab-axis", + "vocab-precision", + "pred-comment", + ): + assert by_class[cls] >= 2, by_class + + +def test_mutation_reports_the_mutated_line(): + text = _text("golden_add_sm80.ttir") + lines = text.splitlines() + i = next(k for k, ln in enumerate(lines) if "arith.cmpi slt" in ln) + lines[i] = lines[i].replace("arith.cmpi slt", "arith.cmpi lt") + with pytest.raises(W.MisalignedModule) as e: + W._walk(text, scan_text="\n".join(lines)) + assert e.value.line_no == i + 1 + assert "closed vocabulary" in e.value.problems[0] + + +def test_in_vocabulary_predicate_swap_is_a_single_source_attribute(): + # documented blind spot (spike: 0/32): an in-vocabulary spelling has no + # second source; only the golden census above guards its extraction + text = _text("golden_add_sm80.ttir") + m = W._walk(text, scan_text=text.replace("arith.cmpi slt", "arith.cmpi sge", 1)) + assert "sge" in [op.attrs.get("predicate") for op in m.ops] + + +@pytest.mark.parametrize( + "name, old, new, match", + [ + # a zero-result op's loc has no second source: a garbled trailer must + # not read as "no loc" while the rest of the module prints locs + ("golden_add_sm80.ttir", "tt.return loc(", "tt.return oc(", "prints no loc"), + ("golden_add_sm80.ttir", "} loc(#loc)", "} loc#(#loc)", "prints no loc"), + # a garbled value must not pass as a recovered attribute + ("golden_add_sm80.ttir", "start = 0 : i32", "start = 0 : im32", "is not a int"), + ( + "adv_views.ttir", + "order = array", + "order = arrayx", + "is not a tuple", + ), + # a garbled key must not fall back to the elided default (ieee) + ( + "golden_matmul_s1_sm80.ttir", + "inputPrecision = tf32", + "inputPrecision4 = tf32", + "tt.dot assignments", + ), + ], +) +def test_garbled_text_is_not_misread(name, old, new, match): + text = _text(name) + assert old in text + with pytest.raises(W.MisalignedModule, match=match): + W._walk(text, scan_text=text.replace(old, new, 1)) + + +# Single-source fields: the text is their only source, so a text layer that +# misreads them still aligns (the golden census and the independent extractor +# guard their extraction instead). Every other field of an aligned reading +# must equal the baseline reading. +_SINGLE_SOURCE_ATTRS = frozenset( + { + # constant values (beyond the signed-range check) + ("arith.constant", "value"), + ("arith.constant", "literal"), + ("arith.constant", "splat"), + # in-vocabulary enum swaps + ("arith.cmpi", "predicate"), + ("arith.cmpf", "predicate"), + ("tt.atomic_rmw", "rmw_op"), + ("tt.atomic_rmw", "sem"), + ("tt.atomic_rmw", "scope"), + ("tt.atomic_cas", "sem"), + ("tt.atomic_cas", "scope"), + ("tt.get_program_id", "axis"), + ("tt.get_num_programs", "axis"), + ("tt.descriptor_reduce", "kind"), + # a deleted clause reads as the printer-elided default + ("tt.dot", "inputPrecision"), + ("tt.dot", "maxNumImpreciseAcc"), + ("tt.reshape", "allow_reorder"), + ("scf.for", "unsignedCmp"), + # no getter and no type relation + ("tt.elementwise_inline_asm", "packed_element"), + } +) +_NO_ATTRS = W._FrozenMap({}) + + +def _op_shape(op: W.Op) -> W.Op: + """The op without the fields _silent_diffs compares on its own.""" + return dataclasses.replace( + op, + attrs=_NO_ATTRS, + line_no=None, + end_line=None, + loc=None, + callers=(), + loc_name=None, + ) + + +def _has_pred_comments(text: str) -> bool: + return re.search(r"^\s*\^.*//", text, re.M) is not None + + +def _silent_diffs(base: W.Module, got: W.Module, pred_comments: bool) -> list[str]: + """What an aligned reading ``got`` reads differently from ``base``, minus + the single-source fields: line numbers, zero-result op locs, entry-block + labels, func argument attrs, discardable (non-NEEDED) attrs, the + attributes in _SINGLE_SOURCE_ATTRS, the order of a cf.cond_br's + successors, and any cf successor when the text prints no pred comments.""" + shape = lambda m: (len(m.ops), len(m.blocks), len(m.values), len(m.funcs)) # noqa: E731 + if shape(base) != shape(got): + return [f"record counts {shape(base)} -> {shape(got)}"] + out = [f"value {a.index}" for a, b in zip(base.values, got.values) if a != b] + for oa, ob in zip(base.ops, got.ops): + where = f"op {oa.index} {oa.name} (line {oa.line_no})" + # a zero-result op's loc has no second source + site = (oa.loc, oa.callers, oa.loc_name) + if (oa.results or oa.implicit) and site != (ob.loc, ob.callers, ob.loc_name): + out.append(f"{where}: loc {oa.loc} -> {ob.loc}") + if _op_shape(oa) != _op_shape(ob): + out.append(f"{where}: structure") + for k in sorted(set(oa.attrs) | set(ob.attrs)): + va, vb = oa.attrs.get(k), ob.attrs.get(k) + if va == vb and type(va) is type(vb): + continue + if ( + k not in W.PRINTERS[base.release].needed.get(oa.name, ()) + or (oa.name, k) in _SINGLE_SOURCE_ATTRS + ): + continue + if k == "successors" and ( + not pred_comments + or (oa.name == "cf.cond_br" and sorted(va or ()) == sorted(vb or ())) + ): + continue + out.append(f"{where}: {k} {va!r} -> {vb!r}") + for ba, bb in zip(base.blocks, got.blocks): + if ba.position == 0: # an entry block's label is never referenced + ba, bb = ( + dataclasses.replace(ba, label=None), + dataclasses.replace(bb, label=None), + ) + if ba != bb: + out.append(f"block {ba.index} ({ba.label})") + for fa, fb in zip(base.funcs, got.funcs): + strip = lambda f: dataclasses.replace( # noqa: E731 + f, args=tuple(dataclasses.replace(x, attrs=_NO_ATTRS) for x in f.args) + ) + if strip(fa) != strip(fb): + out.append(f"func {fa.sym_name}") + return out + + +def test_fuzzed_text_layer_fails_closed(): + """Random character / line edits to the text-layer input: every outcome is + MisalignedModule, or an aligned Module that reads like the baseline in + every field with a second source; never another exception.""" + rng = random.Random(1) + pool = list('%^{}()<>[],:="#@ \\x0123456789abcdefgilnorstxyzE.-_/') + [ + "loc(", + "dense<", + "->", + "\n", + "//", + ] + names = [n for n in ALIGNED if n != "kernel_deep_chain.ttir"] + outcomes: collections.Counter[str] = collections.Counter() + baseline: dict[str, W.Module] = {} + silent = [] + for _ in range(400): + name = rng.choice(names) + text = s = _text(name) + for _ in range(rng.randint(1, 3)): + k = rng.randrange(len(s)) + r = rng.random() + if r < 0.4: + s = s[:k] + s[k + 1 :] + elif r < 0.8: + s = s[:k] + rng.choice(pool) + s[k:] + else: + lines = s.split("\n") + i = rng.randrange(len(lines)) + s = "\n".join(lines[:i] + [lines[i]] + lines[i:]) + try: + got = W._walk(text, scan_text=s) + except W.MisalignedModule: + outcomes["misaligned"] += 1 + continue + outcomes["aligned"] += 1 + if name not in baseline: + baseline[name] = W._walk(text) + diffs = _silent_diffs(baseline[name], got, _has_pred_comments(text)) + if diffs: + silent.append((name, diffs[:3])) + assert not silent + assert outcomes["misaligned"] > 250 and outcomes["aligned"] > 50 + + +def test_silent_diff_sees_what_the_cross_checks_guard(): + # the fuzz test's comparison flags a result op's loc, a NEEDED attribute + # with a second source, and a cf.br retarget under pred comments + text = _text("golden_nested_guard_merge_sm80.ttir") + base = W._walk(text) + load = next(op for op in base.ops if op.name == "tt.load") + moved = dataclasses.replace(load, loc=W.SourceLoc("elsewhere.py", 1, 1)) + ops = list(base.ops) + ops[load.index] = moved + assert _silent_diffs(base, dataclasses.replace(base, ops=tuple(ops)), True) + rng_op = next(op for op in base.ops if op.name == "tt.make_range") + ops = list(base.ops) + ops[rng_op.index] = dataclasses.replace( + rng_op, attrs=W._FrozenMap({**rng_op.attrs, "start": 1}) + ) + assert _silent_diffs(base, dataclasses.replace(base, ops=tuple(ops)), True) + (br,) = _ops(base, "cf.br") + ops = list(base.ops) + ops[br.index] = dataclasses.replace( + br, attrs=W._FrozenMap({"successors": (br.attrs["successors"][0] - 1,)}) + ) + changed = dataclasses.replace(base, ops=tuple(ops)) + assert _silent_diffs(base, changed, True) + assert not _silent_diffs(base, changed, False) + + +def test_pred_comments_are_checked_when_present(): + text = _text("golden_nested_guard_merge_sm80.ttir") + assert W._walk(text).stats["pred_checks"] == 5 + stripped = re.sub(r"[ \t]*//[^\n]*", "", text) + assert ( + W._walk(stripped).stats["pred_checks"] == 0 + ) # no comments: the check is skipped + # one comment missing while the printer's others are present: misaligned + one_gone = re.sub(r"(\^bb3:)\s*//[^\n]*", r"\1", text, count=1) + with pytest.raises(W.MisalignedModule, match="no predecessor comment"): + W._walk(text, scan_text=one_gone) + with pytest.raises(W.MisalignedModule, match="printer preds"): + W._walk( + text, + scan_text=text.replace("// 2 preds: ^bb3, ^bb4", "// 2 preds: ^bb3, ^bb3"), + ) + + +def test_successor_arity_is_checked(): + text = _text("golden_nested_guard_merge_sm80.ttir") + # drop a successor's operand group in the text layer only + bad = re.sub(r"(cf\.br \^bb5)\(%[-\w.$]+ : i32\)", r"\1", text, count=1) + assert bad != text + with pytest.raises(W.MisalignedModule): + W._walk(text, scan_text=bad) + + +def test_successor_operand_groups_are_checked(): + # an operand moved from one successor group to the other keeps the total + # count and the SSA order: only the per-successor arity sees it + text = _text("crafted_same_dest.ttir") + old = "cf.cond_br %c, ^bb1(%a : i32), ^bb1(%b : i32)" + assert old in text + moved = text.replace(old, "cf.cond_br %c, ^bb1(%a, %b : i32, i32), ^bb1") + with pytest.raises(W.MisalignedModule, match="successor operands"): + W._walk(text, scan_text=moved) + + +# ─────────────────────────── review fixes, one by one ─────────────────────────── + + +def _ops(m: W.Module, name: str) -> list[W.Op]: + return [op for op in m.ops if op.name == name] + + +def test_unicode_source_path(): + m = W._walk(_text("nat_k_uni.ttir")) + files = {op.loc.file for op in m.ops if op.loc is not None} + assert files == { + "/home/hwu27/workspace/triton-viz-ir-mode/ir_mode_audit/spike/review/内核/k_uni.py" + } + + +def test_unicode_parameter_names_come_from_namelocs(): + m = W._walk(_text("nat_uni_params.ttir")) + (f,) = m.funcs + # the printed SSA names are sanitised (%_CF80_ptr, %_E695B0_n) + assert [(a.index, a.name, a.type) for a in f.args] == [ + (0, "π_ptr", "!tt.ptr"), + (1, "数_n", "i32"), + ] + assert [m.values[a.value].name for a in f.args] == ["π_ptr", "数_n"] + assert m.blocks[m.ops[f.op].regions[0][0]].arg_names == ("π_ptr", "数_n") + + +def test_unicode_strings_cross_checked_against_bindings(): + m = W._walk(_text("kernel_unicode_msgs.ttir")) + (a,) = _ops(m, "tt.assert") + assert a.attrs["message"] == "错误: π must be > 0" + m = W._walk(_text("crafted_unicode_strings.ttir")) + assert [f.sym_name for f in m.funcs] == ["核"] + assert [a.name for a in m.funcs[0].args] == ["π_ptr", "数"] + assert _ops(m, "tt.assert")[0].attrs["message"] == '错误 "quoted" \\ back' + assert _ops(m, "tt.print")[0].attrs["prefix"] == "π=" + assert _ops(m, "tt.load")[0].loc_name == "值" + assert {op.loc.file for op in m.ops if op.loc} == {"/tmp/内核/k.py"} + + +@pytest.mark.parametrize( + "body, value", + [ + ("\\E9\\94\\99", "错"), + ("a\\22b\\\\c", 'a"b\\c'), + ("tab\\09nl\\0A", "tab\tnl\n"), + ("\\n\\t", "\n\t"), + ("plain", "plain"), + ], +) +def test_unescape_decodes_utf8_bytes(body, value): + assert W._unescape(body) == value + + +@pytest.mark.parametrize("body", ["\\E9\\94", "\\q", "\\"]) +def test_unescape_rejects_malformed(body): + with pytest.raises(ValueError): + W._unescape(body) + + +def test_uppercase_exponent_floats(): + m = W._walk(_text("kernel_eps_consts.ttir")) + consts = {op.attrs["literal"]: dict(op.attrs) for op in _ops(m, "arith.constant")} + assert consts["9.99999997E-7"] == { + "literal": "9.99999997E-7", + "value": ("float", "9.99999997E-7"), + } + assert consts["dense<9.99999996E-13>"]["value"] == ("float", "9.99999996E-13") + assert consts["dense<9.99999996E-13>"]["splat"] is True + m = W._walk(_text("adv_consts.ttir")) + by_lit = {op.attrs["literal"]: op.attrs for op in _ops(m, "arith.constant")} + assert by_lit["dense<-2.14748365E+9>"]["splat"] is True + assert by_lit["dense<0x7FC00000>"]["value"] == ( + "float", + "0x7FC00000", + ) # NaN bit pattern + assert by_lit["-9223372036854775807"]["value"] == -9223372036854775807 + + +def test_non_splat_dense_constants(): + text = _text("crafted_attr_dicts.ttir").replace( + "%r = tt.make_range", + "%ds = arith.constant dense<[1, 2]> : tensor<2xi32>\n %r = tt.make_range", + 1, + ) + m = W._walk(text) + (c,) = [ + op + for op in _ops(m, "arith.constant") + if op.attrs["literal"].startswith("dense") + ] + assert c.attrs["value"] == ("dense", "[1, 2]") and c.attrs["splat"] is False + + +def test_leading_attr_dict_on_arith_constant(): + m = W._walk(_text("crafted_attr_dicts.ttir")) + values = [op.attrs["value"] for op in _ops(m, "arith.constant")] + assert values == [16, -3] + m = W._walk(_text("nat_hint_scalar_const.ttir")) + assert 64 in [op.attrs["value"] for op in _ops(m, "arith.constant")] + + +def test_generic_closer_attr_dict_is_read(): + (red,) = _ops(W._walk(_text("crafted_attr_dicts.ttir")), "tt.reduce") + assert red.attrs["axis"] == 0 and red.attrs["axis_note"] == 3 + + +def test_tt_dot_default_precision(): + m = W._walk(_text("kernel_dot_precisions.ttir")) + got = [ + (op.attrs["inputPrecision"], op.attrs["maxNumImpreciseAcc"]) + for op in _ops(m, "tt.dot") + ] + assert got == [("ieee", 0), ("tf32x3", 0)] + (d,) = _ops(W._walk(_text("spike_dot.ttir")), "tt.dot") + assert d.attrs["inputPrecision"] == "tf32" + + +def test_reshape_allow_reorder(): + m = W._walk(_text("adv_views.ttir")) + assert [op.attrs["allow_reorder"] for op in _ops(m, "tt.reshape")] == [True, True] + text = _text("adv_views.ttir").replace(" allow_reorder", "", 1) + assert [op.attrs["allow_reorder"] for op in _ops(W._walk(text), "tt.reshape")] == [ + False, + True, + ] + with pytest.raises(W.MisalignedModule, match="tt.reshape keyword"): + W._walk( + _text("adv_views.ttir"), + scan_text=_text("adv_views.ttir").replace("allow_reorder", "reorder", 1), + ) + + +@pytest.mark.parametrize( + "name, shapes", + [ + ("crafted_empty_else.ttir", {"scf.if": [(1, 1)]}), + ("crafted_empty_for.ttir", {"scf.for": [(1,)]}), + ("nat_empty_then.ttir", {"scf.if": [(1, 1)]}), + ( + "crafted_empty_bodies.ttir", + {"scf.if": [(1, 1), (1, 1), (1, 0)], "scf.for": [(1,)]}, + ), + ], +) +def test_empty_region_bodies(name, shapes): + m = W._walk(_text(name)) + for opname, want in shapes.items(): + assert [tuple(len(r) for r in op.regions) for op in _ops(m, opname)] == want + for blk in m.blocks: + if m.ops[blk.op].name in ("scf.if", "scf.for"): + last = m.ops[blk.ops[-1]] + assert last.name == "scf.yield" + if last.implicit: + assert last.line_no is None and last.operands == () and last.loc is None + + +def test_empty_for_body_keeps_its_induction_variable(): + m = W._walk(_text("crafted_empty_for.ttir")) + (loop,) = _ops(m, "scf.for") + (body,) = loop.regions[0] + assert len(m.blocks[body].args) == 1 and m.blocks[body].arg_types == ("i32",) + assert [m.ops[i].implicit for i in m.blocks[body].ops] == [True] + + +# The tensordesc types of adv_descs as each release prints them. +_DESC_TYPES = { + "3.6": ("!tt.tensordesc>", "!tt.tensordesc>"), + "3.8": ("!tt.tensordesc<32x32xf16>", "!tt.tensordesc<1x32xf16>"), +} + + +def test_descriptor_operand_order(): + m = W._walk(_text("adv_descs.ttir")) + tile, row = _DESC_TYPES[_printed_by("adv_descs.ttir")] + (store,) = _ops(m, "tt.descriptor_store") + (red,) = _ops(m, "tt.descriptor_reduce") + (scatter,) = _ops(m, "tt.descriptor_scatter") + types = lambda op: [m.values[v].type for v in op.operands] # noqa: E731 + # ODS order (desc, src, indices...), printed `%desc[%i, %j], %src` + assert types(store) == [tile, "tensor<32x32xf16>", "i32", "i32"] + assert types(red) == types(store) and red.attrs["kind"] == "add" + # descriptor_scatter prints in ODS order (desc, x_offsets, y_offset, src) + assert types(scatter) == [row, "tensor<32xi32>", "i32", "tensor<32x32xf16>"] + + +def test_dot_scaled_operand_order(): + m = W._walk(_text("kernel_dot_scaled.ttir")) + (d,) = _ops(m, "tt.dot_scaled") + # ODS (a, b, c, a_scale, b_scale), printed `%a scale %as, %b scale %bs, %c` + assert [m.values[v].type for v in d.operands] == [ + "tensor<128x64xf8E4M3FN>", + "tensor<64x128xf8E4M3FN>", + "tensor<128x128xf32>", + "tensor<128x2xi8>", + "tensor<128x2xi8>", + ] + assert [m.values[v].name for v in d.operands[3:]] == ["a_scale", "b_scale"] + + +def test_cf_successors_are_block_indices(): + m = W._walk(_text("golden_nested_guard_merge_sm80.ttir")) + for op in m.ops: + if op.name in ("cf.br", "cf.cond_br"): + parent = m.blocks[op.block] + for s in op.attrs["successors"]: + dest = m.blocks[s] + assert (dest.op, dest.region) == ( + parent.op, + parent.region, + ) and dest.position > 0 + (br,) = _ops(m, "cf.br") + assert m.blocks[br.attrs["successors"][0]].label == "^bb5" + same = W._walk(_text("crafted_same_dest.ttir")) + first = _ops(same, "cf.cond_br")[0] + assert len(set(first.attrs["successors"])) == 1 # both edges into ^bb1 + + +def test_parser_invented_locs_are_absent(): + m = W._walk(_text("spike_spin_while.ttir")) + whiles = _ops(m, "scf.while") + before = m.blocks[whiles[1].regions[0][0]] + assert before.arg_names == (None,) # `scf.while (%v_1 = %v)` prints no arg loc + for op in m.ops: + for loc in (op.loc, *op.callers): + assert loc is None or not loc.file.startswith( + ("/proc/self/fd", tempfile.gettempdir()) + ) + + +def test_callsite_and_fused_locs(): + m = W._walk(_text("crafted_locs.ttir")) + b = [op for op in m.ops if op.name == "arith.addi"][1] + assert b.loc == W.SourceLoc("y.py", 3, 4) and b.loc_name == "inner" + assert b.callers == (W.SourceLoc("k.py", 2, 5), W.SourceLoc("k.py", 3, 6)) + (store,) = _ops(m, "tt.store") + assert store.loc == W.SourceLoc("k.py", 2, 5) and store.callers == ( + W.SourceLoc("k.py", 2, 5), + ) + (ret,) = _ops(m, "tt.return") + assert ret.loc == W.SourceLoc("k.py", 12, 1) # alias defined after its use + + +def test_quoted_symbols_and_strings_with_syntax_chars(): + m = W._walk(_text("crafted_symbols_strings.ttir")) + assert [(f.sym_name, f.visibility) for f in m.funcs] == [ + ('f{%x} "q" (a)', "private"), + ("k", "public"), + ] + (call,) = _ops(m, "tt.call") + assert call.attrs["callee"] == 'f{%x} "q" (a)' and len(call.results) == 2 + (asm,) = _ops(m, "tt.elementwise_inline_asm") + assert asm.attrs["asm_string"] == '{ mov.u32 $0, %tid.x; } // loc("x") }' + assert asm.attrs["pure"] is True and asm.attrs["packed_element"] == 1 + + +def test_atomic_enums(): + m = W._walk(_text("spike_atomics.ttir")) + assert [ + (op.attrs["rmw_op"], op.attrs["sem"], op.attrs["scope"]) + for op in _ops(m, "tt.atomic_rmw") + ] == [ + ("fadd", "relaxed", "cta"), + ("max", "release", "sys"), + ("min", "acquire", "gpu"), + ("and", "acq_rel", "gpu"), + ("or", "acq_rel", "gpu"), + ("xor", "acq_rel", "gpu"), + ("exch", "relaxed", "sys"), + ("add", "acq_rel", "gpu"), + ] + assert [ + (op.attrs["sem"], op.attrs["scope"]) for op in _ops(m, "tt.atomic_cas") + ] == [("acq_rel", "cta")] + + +def test_program_id_axes(): + m = W._walk(_text("golden_tile2d_sm80.ttir")) + assert sorted(op.attrs["axis"] for op in _ops(m, "tt.get_program_id")) == [0, 1] + + +def test_deep_def_chain(): + m = W._walk(_text("kernel_deep_chain.ttir")) + assert m.stats["ops"] > 1200 and m.stats["ssa_edges"] > 2400 + + +def test_deep_region_nesting_is_iterative(): + depth = 400 # well past Python's default recursion limit per nested frame + lines = [ + "module {", + " tt.func public @k(%c: i1, %p: !tt.ptr) attributes {noinline = false} {", + ] + lines.append(" %v = arith.constant 7 : i32") + lines += [" scf.if %c {"] * depth + lines.append(" tt.store %p, %v : !tt.ptr") + lines += [" }"] * depth + lines += [" tt.return", " }", "}"] + m = W._walk("\n".join(lines)) + assert len(_ops(m, "scf.if")) == depth and m.stats["implicit_ops"] == depth + assert max(len(op.path) for op in m.ops) == depth + 2 + + +def test_deep_loc_alias_chain_fails_closed(): + n = 300 + aliases = ['#loc0 = loc("k.py":1:1)'] + [ + f'#loc{i} = loc("n{i}"(#loc{i - 1}))' for i in range(1, n) + ] + text = "\n".join( + aliases + + [ + "module {", + " tt.func public @k(%p: !tt.ptr) attributes {noinline = false} {", + f" %v = arith.constant 7 : i32 loc(#loc{n - 1})", + " tt.store %p, %v : !tt.ptr", + " tt.return", + " }", + "}", + ] + ) + with pytest.raises(W.MisalignedModule, match="too deep"): + W._walk(text) + + +def _const_module(lit: str, ty: str) -> str: + return ( + "module {\n" + " tt.func public @k() attributes {noinline = false} {\n" + f" %c = arith.constant {lit} : {ty}\n" + " tt.return\n" + " }\n" + "}\n" + ) + + +@pytest.mark.parametrize( + "lit, ty, match", + [ + # the parser accepts these, but MLIR holds the bits as a negative + # value (the printer re-prints -1, -2147483648, -1, dense<-1>, -1) + ("4294967295", "i32", "signed"), + ("2147483648", "i32", "signed"), + ("255", "i8", "signed"), + ("dense<4294967295>", "tensor<4xi32>", "signed"), + ("18446744073709551615", "i64", "signed"), + ("dense<1>", "tensor<4xi1>", "true / false"), # printed dense + ], +) +def test_constant_outside_the_printed_signed_range_is_refused(lit, ty, match): + with pytest.raises(W.MisalignedModule, match=match): + W._walk(_const_module(lit, ty)) + + +@pytest.mark.parametrize( + "lit, ty, value", + [ + ("-1", "i32", -1), + ("2147483647", "i32", 2**31 - 1), + ("-2147483648", "i32", -(2**31)), + ("-128", "i8", -128), + ("dense<-1>", "tensor<4xi32>", -1), + ("-9223372036854775808", "i64", -(2**63)), + ("9223372036854775807", "index", 2**63 - 1), + ], +) +def test_constant_in_the_printed_signed_range(lit, ty, value): + (c,) = _ops(W._walk(_const_module(lit, ty)), "arith.constant") + assert c.attrs["value"] == value + + +_FOR = """module { + tt.func public @k(%p: !tt.ptr, %n: i32) attributes {noinline = false} { + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + scf.for KW%i = %c0 to %n step %c1 : i32 { + %q = tt.addptr %p, %i : !tt.ptr, i32 + tt.store %q, %i : !tt.ptr + } + tt.return + } +} +""" + + +def test_scf_for_unsigned_compare_is_recorded(): + signed, unsigned = _FOR.replace("KW", ""), _FOR.replace("KW", "unsigned ") + (loop,) = _ops(W._walk(signed), "scf.for") + assert loop.attrs["unsignedCmp"] is False + (loop,) = _ops(W._walk(unsigned), "scf.for") + assert loop.attrs["unsignedCmp"] is True + # an unknown header keyword is a misread, never an ignored word + with pytest.raises(W.MisalignedModule, match="scf.for header"): + W._walk(unsigned, scan_text=_FOR.replace("KW", "signless ")) + with pytest.raises(W.MisalignedModule, match="scf.for header"): + W._walk(signed, scan_text=_FOR.replace("KW", "").replace(" step ", " stride ")) + + +def test_ttgir_is_refused(): + text = "#blocked = #ttg.blocked<{sizePerThread = [1], threadsPerWarp = [32], warpsPerCTA = [4], order = [0]}>\n" + text += _text("golden_add_sm80.ttir") + with pytest.raises((W.MisalignedModule, W.ModuleParseError)): + W._walk(text) + + +# ─────────────────────────── parse errors and the parse input ─────────────────────────── + + +def test_parse_error_carries_the_diagnostic(capfd): + text = _text("golden_add_sm80.ttir") + lines = text.splitlines() + i = next(k for k, ln in enumerate(lines) if "tt.make_range {" in ln) + lines[i] = lines[i].replace("tt.make_range {", "tt.make_range_v2 {") + with pytest.raises(W.ModuleParseError) as e: + W.walk_module("\n".join(lines)) + assert ( + "tt.make_range_v2" in e.value.diagnostic and "is unknown" in e.value.diagnostic + ) + assert "" in e.value.diagnostic and "/proc/self/fd" not in e.value.diagnostic + assert e.value.line_no == i + 1 + out, err = capfd.readouterr() + assert "make_range_v2" not in err # captured, not leaked onto the process stderr + + +@pytest.mark.parametrize( + "old, new", + [ + ("{noinline = false}", "{noinline = }"), # malformed attribute dict + ( + "arith.addf %2, %cst : tensor<64xf32>", + "arith.addf %2, %cst : tensor<32xf32>", + ), # operand type mismatch + ], +) +def test_malformed_text_is_a_parse_error(old, new): + text = _text("nat_k_uni.ttir") + assert old in text + with pytest.raises(W.ModuleParseError): + W._walk(text.replace(old, new, 1)) + + +def test_unwritable_tmpdir_is_a_parse_error(monkeypatch, tmp_path): + ro = tmp_path / "ro" + ro.mkdir() + ro.chmod(0o500) + if os.access(ro, os.W_OK): + pytest.skip("directory permissions are not enforced (root?)") + monkeypatch.delattr(os, "memfd_create", raising=False) + monkeypatch.setattr(tempfile, "tempdir", str(ro)) + with pytest.raises(W.ModuleParseError, match="temporary file"): + W._walk(_text("nat_k_uni.ttir")) + ro.chmod(0o700) + + +def test_tempfile_fallback(monkeypatch, tmp_path): + monkeypatch.delattr(os, "memfd_create", raising=False) + monkeypatch.setattr(tempfile, "tempdir", str(tmp_path)) + m = W._walk(_text("spike_spin_while.ttir")) + assert m.stats["ops"] == 19 + assert list(tmp_path.iterdir()) == [] # the parse input is removed + # parser-invented locs name the temp path: filtered on both sides + assert m.blocks[_ops(m, "scf.while")[1].regions[0][0]].arg_names == (None,) + with pytest.raises(W.ModuleParseError) as e: + W._walk( + _text("nat_k_uni.ttir").replace("tt.make_range {", "tt.make_range_v2 {", 1) + ) + assert str(tmp_path) not in e.value.diagnostic and "" in e.value.diagnostic + + +def test_no_fd_leak(): + before = len(os.listdir("/proc/self/fd")) + for _ in range(20): + W._walk(_text("nat_k_uni.ttir")) + with pytest.raises(W.ModuleParseError): + W._walk( + _text("nat_k_uni.ttir").replace("tt.make_range {", "tt.make_range_v2 {", 1) + ) + assert len(os.listdir("/proc/self/fd")) == before + + +# ─────────────────────────── per-release printer tables ─────────────────────────── + + +def test_printer_tables_are_keyed_by_release(): + assert set(W.PRINTERS) == {"3.6", "3.8"} + for release, table in W.PRINTERS.items(): + assert table.release == release + # every leading keyword has a closed vocabulary, and every table + # names only what its fields hold + for op, keys in table.keywords.items(): + assert all((op, k) in table.vocab for k in keys), (release, op) + for (op, key), ints in table.bind_ints.items(): + assert ("int", key) in table.bind_attrs.get(op, ()), (release, op) + assert set(ints) == table.vocab[(op, key)], (release, op) + assert W.printer("3.6") is W.PRINTERS["3.6"] + assert W.printer().release == G.RELEASE + + +def test_the_3_8_table_is_3_6_plus_the_audited_changes(): + """3.8's printer is 3.6's (audited) but for ttg.barrier and the missing + block-pointer types: any other difference must be a deliberate edit.""" + t36, t38 = W.PRINTERS["3.6"], W.PRINTERS["3.8"] + barrier = {"ttg.barrier"} + for field in ("needed", "keywords", "bind_attrs", "defaults", "printed_order"): + a, b = dict(getattr(t36, field)), dict(getattr(t38, field)) + assert {k: v for k, v in b.items() if k not in barrier} == a, field + for field in ("vocab", "bind_ints", "attr_types"): + a, b = dict(getattr(t36, field)), dict(getattr(t38, field)) + assert {k: v for k, v in b.items() if k[0] not in barrier} == a, field + assert t38.keywords["ttg.barrier"] == ("addrSpace",) + assert t38.bind_attrs["ttg.barrier"] == (("int", "addrSpace"),) + assert dict(t38.bind_ints[("ttg.barrier", "addrSpace")]) == { + "none": 0, + "local": 1, + "global_read": 2, + "global_write": 4, + "tensor_read": 8, + "tensor_write": 16, + "all": 31, + } + assert t36.elides_yield == t38.elides_yield + assert (t36.block_pointer_types, t38.block_pointer_types) == (True, False) + + +@pytest.mark.parametrize("version", ["3.7.0", "3.9.0rc1", "4.0.0", "3.60.1"]) +def test_an_unknown_release_fails_closed_before_any_parse(monkeypatch, version): + import triton + + def no_parse(_data, _table): + raise AssertionError("parsed a text for an unknown release") + + monkeypatch.setattr(W, "_bind_walk", no_parse) + monkeypatch.setattr(triton, "__version__", version) + release = ".".join(version.split(".")[:2]) + for walk in (W.walk_module, W._walk): + with pytest.raises(W.UnknownTritonRelease) as e: + walk(_text("golden_add_sm80.ttir")) + assert (e.value.release, e.value.version) == (release, version) + assert f"no printer table for Triton {version}" in e.value.message + assert "3.6.x, 3.8.x" in e.value.message + assert not W._CACHE # nothing cached for it + with pytest.raises(W.UnknownTritonRelease): + W.printer(release) + + +def test_the_cache_is_keyed_by_release(monkeypatch): + import triton + + text = _text("golden_add_sm80.ttir") + here = W.walk_module(text) + assert here.release == G.RELEASE + # a (pretended) other release with the same table: its own entry + other = dataclasses.replace(W.printer(), release="9.9") + monkeypatch.setattr(W, "PRINTERS", {**W.PRINTERS, "9.9": other}) + monkeypatch.setattr(triton, "__version__", "9.9.0") + there = W.walk_module(text) + assert there.release == "9.9" and there is not here and len(W._CACHE) == 2 + assert dataclasses.replace(there, release=here.release) == here + + +# The barrier tl.debug_barrier() prints, per release. +_BARRIER = {"3.6": "gpu.barrier", "3.8": "ttg.barrier all"} + + +def _barrier_module(barrier: str) -> str: + return ( + "module {\n" + " tt.func public @k(%p: !tt.ptr) attributes {noinline = false} {\n" + f" {barrier}\n" + " %v = tt.load %p : !tt.ptr\n" + " tt.return\n" + " }\n" + "}\n" + ) + + +def test_the_release_barrier_aligns(): + m = W._walk(_barrier_module(_BARRIER[G.RELEASE])) + (barrier,) = [op for op in m.ops if op.name.endswith(".barrier")] + assert barrier.name == _BARRIER[G.RELEASE].split()[0] and not barrier.results + if G.RELEASE == "3.8": + assert dict(barrier.attrs) == {"addrSpace": "all"} + # addrSpace is needed, and read by the bindings too + plain = W._walk(_barrier_module("")).stats + assert m.stats["bind_attrs"] - plain["bind_attrs"] == 1 + assert m.stats["needed_attrs"] - plain["needed_attrs"] == 1 + # the other release's barrier: 3.6's parser has no ttg dialect, 3.8's + # still has gpu.barrier (its reader refuses it, see test_ttir_reader) + for release, spelling in _BARRIER.items(): + if release == G.RELEASE: + continue + if G.RELEASE == "3.6": + with pytest.raises(W.ModuleParseError, match="is unknown"): + W._walk(_barrier_module(spelling)) + else: + assert _ops(W._walk(_barrier_module(spelling)), "gpu.barrier") + + +_ONLY_3_8 = pytest.mark.skipif( + G.RELEASE != "3.8", reason="needs Triton 3.8's bindings (ttg.barrier)" +) + + +@_ONLY_3_8 +@pytest.mark.parametrize( + "word, bits", + sorted(W.PRINTERS["3.8"].bind_ints[("ttg.barrier", "addrSpace")].items()), +) +def test_ttg_barrier_addr_space_is_cross_checked(word, bits): + text = _barrier_module(f"ttg.barrier {word}") + (barrier,) = _ops(W._walk(text), "ttg.barrier") + assert barrier.attrs["addrSpace"] == word + # an in-vocabulary swap in the text layer: the bindings' bits see it + other = "none" if word != "none" else "all" + with pytest.raises( + W.MisalignedModule, match=f"attr addrSpace: text '{other}' vs bindings {bits}" + ): + W._walk(text, scan_text=_barrier_module(f"ttg.barrier {other}")) + + +@_ONLY_3_8 +def test_ttg_barrier_spellings_outside_the_one_keyword_syntax_misalign(): + # a bit combination prints as `a|b`: not one keyword, so not read + with pytest.raises(W.MisalignedModule, match="leading keyword"): + W._walk(_barrier_module("ttg.barrier local|global_read")) + text = _barrier_module("ttg.barrier all") + with pytest.raises(W.MisalignedModule, match="closed vocabulary"): + W._walk(text, scan_text=_barrier_module("ttg.barrier every")) + + +# two spellings that once reached (and aborted) 3.8's parser: the type split +# over two lines, and a comment between `ptr<` and `tensor` +_SPLIT_NEWLINE = "!tt.ptr<\n tensor<32xf32>>" +_SPLIT_COMMENT = "!tt.ptr< // c\n tensor<32xf32>>" + +_BLOCK_PTR = """module { + tt.func public @k(%p: !tt.ptr, %b: TYPE) attributes {noinline = false} { + tt.return + } +} +""" + + +@pytest.mark.parametrize( + "text, line", + [ + (_BLOCK_PTR.replace("TYPE", "!tt.ptr>"), 2), + (_BLOCK_PTR.replace("TYPE", "!tt.ptr, 1>"), 2), + (_BLOCK_PTR.replace("TYPE", "!tt.ptr< tensor <4xf32>>"), 2), + (_BLOCK_PTR.replace("TYPE", "!tt>>"), 2), + (_BLOCK_PTR.replace("TYPE", "!tt.ptr>>"), 2), + # split over lines, or around a comment: no single line holds it + (_BLOCK_PTR.replace("TYPE", _SPLIT_NEWLINE), 2), + (_BLOCK_PTR.replace("TYPE", _SPLIT_COMMENT), 2), + (_BLOCK_PTR.replace("TYPE", "!tt.ptr\n>"), 2), + (_BLOCK_PTR.replace("TYPE", '!tt.ptr< // "q>\n // b\n tensor<4xf32>>'), 2), + (_BLOCK_PTR.replace("TYPE", "!tt>>"), 2), + # a type alias may hide the pointee + ("!t = tensor<4xf32>\n" + _BLOCK_PTR.replace("TYPE", "!tt.ptr"), 1), + ("!t // c\n = tensor<4xf32>\n" + _BLOCK_PTR.replace("TYPE", "!tt.ptr"), 1), + ( + "// c\n\n !t\n = tensor<4xf32>\n" + + _BLOCK_PTR.replace("TYPE", "!tt.ptr"), + 3, + ), + ], +) +def test_block_pointer_types_are_refused_before_the_3_8_parser(monkeypatch, text, line): + """3.8's parser aborts the process on a block-pointer type: the walk + refuses such a text before the bindings see it (also under 3.6's + bindings: the screen runs first), and 3.6's table does not screen.""" + + def no_parse(_data, _table): + raise AssertionError("the bindings got a block-pointer type") + + monkeypatch.setattr(W, "_bind_walk", no_parse) + with pytest.raises(W.ModuleParseError, match="no block pointers") as e: + W._walk(text, table=W.PRINTERS["3.8"]) + assert e.value.line_no == line and f"line {line}:" in e.value.diagnostic + W._screen(text, W.PRINTERS["3.6"]) # 3.6 has block pointers: no screen + + +def test_the_block_pointer_screen_reads_types_not_strings(): + text = _module_with_print('"ptr>"), + W.PRINTERS["3.8"], + ) + W._screen(_module_with_print('"ptr< // "'), W.PRINTERS["3.8"]) + in_string = _BLOCK_PTR.replace("TYPE", "i32").replace( + "{noinline = false}", + '{noinline = false, s = "a // b", t = !tt.ptr<\n tensor<4xf32>>}', + ) + with pytest.raises(W.ModuleParseError, match="line 2: a block-pointer type"): + W._screen(in_string, W.PRINTERS["3.8"]) + unterminated = text.replace('"ptr str: + return ( + "module {\n" + " tt.func public @k(%x: i32) attributes {noinline = false} {\n" + f" tt.print {prefix} {{hex = false, isSigned = array}} : %x : i32\n" + " tt.return\n" + " }\n" + "}\n" + ) + + +_ABORT = r""" +import sys, warnings +warnings.filterwarnings("ignore") +sys.path.insert(0, sys.argv[1]) +from tilelens.ir import _mlir_walk as W +text = sys.argv[2] +try: + m = W.walk_module(text) + print("aligned", m.release) +except W.ModuleParseError as e: + print("refused", e.line_no) +except W.MisalignedModule: + print("misaligned") +""" + + +@pytest.mark.parametrize( + "pointee, aligned", + [ + ("!tt.ptr>", True), + # the text layer reads a type within one line and no comments + # outside block labels: these misalign where the parser takes them + (_SPLIT_NEWLINE, False), + (_SPLIT_COMMENT, False), + ], +) +def test_a_block_pointer_text_never_kills_the_process(pointee, aligned): + text = _BLOCK_PTR.replace("TYPE", pointee) + r = subprocess.run( + [sys.executable, "-c", _ABORT, str(REPO), text], + capture_output=True, + text=True, + timeout=120, + ) + assert r.returncode == 0, (r.returncode, r.stderr[-2000:]) + if not W.printer().block_pointer_types: + want = "refused 2" + else: + want = f"aligned {G.RELEASE}" if aligned else "misaligned" + assert r.stdout.strip() == want + + +# ─────────────────────────── cache, immutability, lifetime ─────────────────────────── + + +def test_cache_returns_the_same_module_and_caches_misalignment(monkeypatch): + text = _text("golden_add_sm80.ttir") + assert W.walk_module(text) is W.walk_module(text) + bad = _text("crafted_generic_form.ttir") + with pytest.raises(W.MisalignedModule) as first: + W.walk_module(bad) + + def no_parse(_data, _table): + raise AssertionError("parsed twice") + + monkeypatch.setattr(W, "_bind_walk", no_parse) + with pytest.raises(W.MisalignedModule) as again: + W.walk_module(bad) + assert ( + again.value.problems == first.value.problems and again.value is not first.value + ) + W.walk_module(text) + + +def test_cache_is_bounded(): + for k in range(W._CACHE_SIZE + 5): + W.walk_module(_text("nat_k_uni.ttir") + f"\n// {k}\n") + assert len(W._CACHE) == W._CACHE_SIZE + + +def test_module_is_immutable(): + m = W._walk(_text("spike_atomics.ttir")) + op = next(op for op in m.ops if op.attrs) + with pytest.raises(TypeError): + op.attrs["x"] = 1 # type: ignore[index] + with pytest.raises(dataclasses.FrozenInstanceError): + op.name = "x" # type: ignore[misc] + with pytest.raises(TypeError): + m.stats["ops"] = 0 # type: ignore[index] + + +def test_records_hash_pickle_and_deep_copy(): + m = W._walk(_text("spike_atomics.ttir")) + assert pickle.loads(pickle.dumps(m)) == m and copy.deepcopy(m) == m + again = W._walk(_text("spike_atomics.ttir")) + assert again == m and hash(again) == hash(m) + memo = {op: op.index for op in m.ops} # ops are usable as keys + assert [memo[op] for op in again.ops] == list(range(len(m.ops))) + assert {f.args[0]: 1 for f in m.funcs if f.args} + + +_PLAIN = (int, str, bool, float, type(None)) + + +def _assert_plain(root) -> int: + """Every object reachable from ``root`` is plain Python data.""" + n = 0 + stack = [root] + while stack: + x = stack.pop() + n += 1 + if isinstance(x, _PLAIN): + continue + if dataclasses.is_dataclass(x) and not isinstance(x, type): + assert type(x).__module__ == W.__name__, type(x) + stack += [getattr(x, f.name) for f in dataclasses.fields(x)] + elif isinstance(x, tuple): + stack += list(x) + elif isinstance(x, W._FrozenMap): + stack += list(x.keys()) + list(x.values()) + else: + raise AssertionError( + f"non-plain object {type(x).__module__}.{type(x).__qualname__}" + ) + return n + + +def test_no_binding_object_escapes(): + for name in ( + "golden_matmul_s3_sm80.ttir", + "spike_reduce_scan.ttir", + "adv_cf_blockargs.ttir", + ): + assert _assert_plain(W._walk(_text(name))) > 100 + + +def test_threads_walk_consistently(): + names = ALIGNED[:12] + ref = {n: _pins(W._walk(_text(n))) for n in names} + errors: list[str] = [] + + def work(k: int) -> None: + for n in names[k % 3 :] + names[: k % 3]: + if _pins(W._walk(_text(n))) != ref[n]: + errors.append(n) + + threads = [threading.Thread(target=work, args=(k,)) for k in range(4)] + for t in threads: + t.start() + for t in threads: + t.join() + assert not errors + + +_HAZARD = r""" +import gc, sys +sys.path.insert(0, sys.argv[1]) +from tilelens.ir import _mlir_walk as W +texts = [open(p, encoding="utf-8").read() for p in sys.argv[2:]] +first = [W.walk_module(t) for t in texts] +sig = [[(op.name, op.operands, op.results, op.result_types) for op in m.ops] for m in first] +W._CACHE.clear() +del first +gc.collect() +# everything the walks created is gone; read types and locs again, walk again +again = [W._walk(t) for t in texts] +gc.collect() +assert sig == [[(op.name, op.operands, op.results, op.result_types) for op in m.ops] for m in again] +for m in again: + for v in m.values: + v.type.encode() +del again +gc.collect() +try: + W._walk(texts[0].replace("tt.make_range {", "tt.make_range_v2 {", 1)) +except W.ModuleParseError: + pass +gc.collect() +print("ok", len(texts)) +""" + + +def test_subprocess_walk_then_drop_everything_is_safe(): + paths = [ + str(TTIR / n) + for n in ( + "golden_matmul_s3_sm80.ttir", + "spike_spin_while.ttir", + "nat_k_uni.ttir", + ) + ] + for _ in range( + 3 + ): # the unsafe pattern crashed nondeterministically (SIGSEGV / SIGBUS / hang) + r = subprocess.run( + [sys.executable, "-c", _HAZARD, str(REPO), *paths], + capture_output=True, + text=True, + timeout=120, + ) + assert r.returncode == 0, r.stderr[-2000:] + assert r.stdout.strip().endswith("ok 3") + + +_FORK = r""" +import os, signal, sys, threading, warnings +warnings.filterwarnings("ignore") +sys.path.insert(0, sys.argv[1]) +from tilelens.ir import _mlir_walk as W +text = open(sys.argv[2], encoding="utf-8").read() +inside, release = threading.Event(), threading.Event() + +def mid_walk(): + # another thread inside a walk: both locks held, fd 2 redirected + with W._CACHE_LOCK, W._PARSE_LOCK, W._StderrCapture(): + inside.set() + release.wait() + +holder = threading.Thread(target=mid_walk) +holder.start() +inside.wait() +pid = os.fork() +if pid == 0: + signal.alarm(60) # a deadlocked walk dies (SIGALRM) instead of hanging + os.write(2, b"child stderr\n") + W.walk_module(text) + os._exit(0) +_, status = os.waitpid(pid, 0) +release.set() +holder.join() +print("child exit", os.waitstatus_to_exitcode(status)) +""" + + +@pytest.mark.skipif(not hasattr(os, "fork"), reason="needs fork") +def test_fork_inside_a_parse_window_is_safe(): + r = subprocess.run( + [sys.executable, "-c", _FORK, str(REPO), str(TTIR / "nat_k_uni.ttir")], + capture_output=True, + text=True, + timeout=120, + ) + assert r.returncode == 0, r.stderr[-2000:] + assert r.stdout.strip() == "child exit 0" # the child's walk took fresh locks + assert "child stderr" in r.stderr # and its fd 2 is the real stderr again + + +def test_other_fd2_output_during_a_parse_is_passed_on(capfd, monkeypatch): + enter = W._StderrCapture.__enter__ + + def enter_then_write(self): + got = enter(self) + os.write(2, b"another thread's line\n") # lands in the capture buffer + return got + + monkeypatch.setattr(W._StderrCapture, "__enter__", enter_then_write) + W._walk(_text("nat_k_uni.ttir")) + assert "another thread's line" in capfd.readouterr().err + + +_RSS = r""" +import gc, sys, warnings +warnings.filterwarnings("ignore") +sys.path.insert(0, sys.argv[1]) +from tilelens.ir import _mlir_walk as W +text = open(sys.argv[2], encoding="utf-8").read() + +def rss_kib(): + with open("/proc/self/status") as f: + for line in f: + if line.startswith("VmRSS:"): + return int(line.split()[1]) + +for _ in range(100): + W._walk(text) +gc.collect() +before = rss_kib() +for _ in range(300): + W._walk(text) +gc.collect() +print((rss_kib() - before) / 300) +""" + + +@pytest.mark.skipif(not os.path.exists("/proc/self/status"), reason="needs /proc") +def test_parse_memory_is_reclaimed(): + # without the body-block erase each parse of this 12 KiB text leaked + # 12-20 KiB; with it about 1-2 KiB (the empty module op, allocator noise) + r = subprocess.run( + [ + sys.executable, + "-c", + _RSS, + str(REPO), + str(TTIR / "golden_matmul_s3_sm80.ttir"), + ], + capture_output=True, + text=True, + timeout=300, + ) + assert r.returncode == 0, r.stderr[-2000:] + per_parse_kib = float(r.stdout.split()[-1]) + assert per_parse_kib < 8.0 diff --git a/tests/unit/ir/test_ttir_reader.py b/tests/unit/ir/test_ttir_reader.py new file mode 100644 index 000000000..d4ac0aad3 --- /dev/null +++ b/tests/unit/ir/test_ttir_reader.py @@ -0,0 +1,1442 @@ +"""tilelens.ir.ttir_reader: the TTIR -> AccessGraph reader on top of _mlir_walk. + +Goldens: tests/golden/ir/ttir/ (the walk layer's corpus) and +tests/golden/ir/reader_ttir/ (the audit probes, the phase-2 review probes +and reader-specific shapes, host-compiled from +tests/golden/ir/reader_kernels.py by tests/golden/ir/generate_reader_ttir.py), +each read as the installed release prints it where it has its own printing +(ttir_/, reader_ttir_/: see _goldens.py). The differential +oracle is #361's regex reader, vendored verbatim as _oracle_ttir_reader_361.py. +""" + +from __future__ import annotations + +import copy +import dataclasses +import hashlib +import itertools +import os +import pickle +import subprocess +import sys +from pathlib import Path + +import pytest + +from tilelens.ir import ParseCache, Refusal, SourceLocation +from tilelens.ir import _mlir_walk as W +from tilelens.ir import ttir_reader as R +from tilelens.ir.ttir_reader import ( + AccessGraph, + Arange, + Bin, + Cmp, + Const, + DataDep, + IntCast, + IterArgOffset, + Observed, + Param, + Pid, + Select, + TTIRKind, + UnsupportedTTIR, + mentions_observed, + observed_indices, + parse_ttir, + width_obligations, +) + +from . import _goldens as G +from . import _oracle_ttir_reader_361 as O + +REPO = Path(__file__).resolve().parents[3] +GOLDEN = REPO / "tests" / "golden" / "ir" +KERNELS = GOLDEN / "reader_kernels.py" +DIRS = ("ttir", "reader_ttir") +# "/" -> the golden the installed release reads +PATHS = {f"{d}/{n}": p for d in DIRS for n, p in G.texts(d).items()} +FILES = list(PATHS) + + +def _path(name: str) -> Path: + return PATHS[name] + + +def _text(name: str) -> str: + return _path(name).read_text(encoding="utf-8") + + +def _graph(name: str) -> AccessGraph: + return parse_ttir(_text(name)) + + +def _refusal(name_or_text: str) -> UnsupportedTTIR: + text = _text(name_or_text) if name_or_text.endswith(".ttir") else name_or_text + with pytest.raises(UnsupportedTTIR) as info: + parse_ttir(text) + return info.value + + +def _module(body: str, args: str = "%p: !tt.ptr, %n: i32", extra: str = "") -> str: + """A minimal TTIR module (no locs) around ``body``'s op lines.""" + lines = "\n ".join(line.strip() for line in body.strip().splitlines()) + return ( + f"module {{\n tt.func public @k({args}) attributes {{noinline = false}} {{\n" + f" {lines}\n tt.return\n }}\n{extra}}}\n" + ) + + +# Where a refusal's loc points, per release: 3.8's frontend locates an op +# at its own AST node (scf.yield at the loop header, a call spanning lines +# at its first line), 3.6's at the last child it visited. +_REFUSAL_SITE = { + "p1_variant_delta": {"3.6": "p += k", "3.8": "for k in range(-8, n):"}, + "loop_observed_advance": {"3.6": "atomic_add", "3.8": "for i in range(0, n):"}, + # the asm call's operand list / its first line + "rv_inline_asm_store": {"3.6": "[p, v]", "3.8": "tl.inline_asm_elementwise("}, +} + + +def _source_line(loc) -> str: + """The reader_kernels.py source line a probe golden's loc points at.""" + assert loc is not None and loc.file.endswith("reader_kernels.py"), loc + return KERNELS.read_text(encoding="utf-8").splitlines()[loc.line - 1] + + +@pytest.fixture(autouse=True) +def _fresh_cache(): + W._CACHE.clear() + yield + W._CACHE.clear() + + +# ─────────────────────────── refusal form (D8) ─────────────────────────── + + +def test_kinds_are_the_representational_ones(): + assert {k.value for k in TTIRKind} == { + "untested-triton-version", + "indirect-address", + "data-dependent-bound", + "nested-loop", + "control-flow", + "block-pointer", + "out-of-vocabulary", + "call", + "loop-variant-advance", + "inline-asm", + "reader-misalignment", + "unparsable", + "other", + } + # a kind is its string, also when formatted + assert ( + TTIRKind.CALL == "call" and f"{TTIRKind.CALL}" == str(TTIRKind.CALL) == "call" + ) + + +def test_unsupported_ttir_carries_structured_fields(): + loc = W.SourceLoc("k.py", 3, 4) + e = UnsupportedTTIR("control-flow", "scf.while", line_no=7, loc=loc) + assert e.kind is TTIRKind.CONTROL_FLOW + assert (e.message, e.line_no, e.loc, str(e)) == ("scf.while", 7, loc, "scf.while") + back = pickle.loads(pickle.dumps(e)) + assert (back.kind, back.message, back.line_no, back.loc) == ( + e.kind, + e.message, + 7, + loc, + ) + # client-owned kinds are not reader kinds + for kind in ("spin-shape", "cas-value", "data-dependent-mask"): + with pytest.raises(ValueError): + UnsupportedTTIR(kind, "no") + # the verdict record copies the fields, no string round-trip; the + # private walk loc becomes the public SourceLocation (D20) + r = Refusal.from_exception(e) + assert (r.kind, r.message, r.line_no, r.loc) == ( + "control-flow", + "scf.while", + 7, + SourceLocation("k.py", 3, 4), + ) + assert type(r.loc) is SourceLocation + + +def test_parse_cache_reads_with_this_reader(): + cache = ParseCache() + parsed = cache.get(_text("ttir/golden_add_sm80.ttir")) + assert isinstance(parsed.graph, AccessGraph) and parsed.error is None + refused = cache.get(_text("reader_ttir/p2_swap.ttir")) + assert isinstance(refused.refusal, UnsupportedTTIR) and refused.error is None + assert refused.refusal.kind is TTIRKind.LOOP_VARIANT_ADVANCE + + +def test_walk_failures_map_to_kinds(): + e = _refusal("ttir/crafted_generic_form.ttir") + assert e.kind is TTIRKind.READER_MISALIGNMENT + assert e.line_no == 3 and "arith.cmpi" in e.message + e = _refusal("module {\n tt.func public @k( {\n}\n") + assert e.kind is TTIRKind.UNPARSABLE + assert e.line_no == 2 and "" in e.message + + +def test_every_walk_table_release_has_a_reader_vocabulary(): + assert set(R._VOCABULARIES) == set(W.PRINTERS) + + +@pytest.mark.parametrize("version", ["3.7.0", "3.9.0", "4.0.0"]) +def test_an_unknown_release_is_refused_as_untested(monkeypatch, version): + import triton + + monkeypatch.setattr(triton, "__version__", version) + e = _refusal("ttir/golden_add_sm80.ttir") + assert e.kind is TTIRKind.UNTESTED_TRITON_VERSION + assert f"no printer table for Triton {version}" in e.message + assert (e.line_no, e.loc) == (None, None) + + +def test_a_release_without_a_reader_vocabulary_is_refused(monkeypatch): + """A walk-layer table alone does not make the reader read a release.""" + vocabularies = {k: v for k, v in R._VOCABULARIES.items() if k != G.RELEASE} + monkeypatch.setattr(R, "_VOCABULARIES", vocabularies) + e = _refusal("ttir/golden_add_sm80.ttir") + assert e.kind is TTIRKind.UNTESTED_TRITON_VERSION + assert f"no op vocabulary for Triton {G.RELEASE}" in e.message + + +def test_refusals_carry_the_op_line_and_source_loc(): + e = _refusal("reader_ttir/p3_call_offset.ttir") + assert e.kind is TTIRKind.CALL and "_p3_helper" in e.message + assert ( + "tt.call" + in _text("reader_ttir/p3_call_offset.ttir").splitlines()[e.line_no - 1] + ) + assert "_p3_helper(" in _source_line(e.loc) + + +# ─────────────────────────── the differential oracle ─────────────────────────── +# #361's regex reader (single-path) against the new reader on every golden. +# Where both accept, the access inventory must match after _normalize; the +# table lists every file where the two legitimately differ, with the reason +# and exactly what differs: the normalized fields (see _diff) when both +# accept, else which reader refuses. + +_NEW = "new refuses" +_OLD = "#361 refuses" +_CALL = "fix 3: tt.call -> call (#361 reads the callee in the caller's env)" +_BITCAST = "atomics through a same-width tt.bitcast pointer (#361: refused)" +_NAMELOC = "a parameter without NameLoc is arg (#361: its printed name)" +_EXPANDED = ( + "a loop-carried pointer tile expanded in the loop gets its own iter_args " + "entry with re-placed lanes (#361 keeps the stale dims: a false in-bounds)" +) +EXPECTED_DIFF: dict[str, tuple[str, str | frozenset[str]]] = { + # the audit's soundness fixes + "reader_ttir/p1_variant_delta.ttir": ( + "fix 1: advance by the induction var -> loop-variant-advance " + "(#361: delta=LoopVar)", + _NEW, + ), + "reader_ttir/loop_observed_advance.ttir": ( + "fix 1: advance by an atomic observed in the loop -> loop-variant-advance", + _NEW, + ), + "reader_ttir/p2_swap.ttir": ( + "fix 2: swapped pointer iter_args -> loop-variant-advance (#361: delta 0)", + _NEW, + ), + "reader_ttir/p3_call_guarded.ttir": (_CALL, _NEW), + "reader_ttir/p3_call_offset.ttir": (_CALL, _NEW), + "reader_ttir/rv_inline_asm_store.ttir": ( + "fix 7: impure inline asm -> inline-asm (#361 misses its st.global)", + _NEW, + ), + "ttir/spike_inline_asm.ttir": ("fix 7: impure inline asm -> inline-asm", _NEW), + "reader_ttir/pure_asm_int_addr.ttir": ( + "fix 7: a pure asm handed the address as an integer -> inline-asm", + _NEW, + ), + "reader_ttir/tile3d_shared_arange.ttir": ( + "expand_dims tracks each lane's position (#361 keeps the first placement, " + "collapsing dims 1 and 2 of one make_range into one lane variable)", + frozenset({"events[0].offset"}), + ), + "reader_ttir/expand_iterarg_3d.ttir": ( + _EXPANDED, + frozenset({"events[0].offset", "events[1].offset", "n_iter_args"}), + ), + "reader_ttir/expand_iterarg_mask.ttir": ( + _EXPANDED, + frozenset({"events[0].offset", "n_iter_args"}), + ), + "reader_ttir/observed_lanes.ttir": ( + "two lanes of one tensor atomic's observations in an address " + "(#361: one symbol for both, a false in-bounds)", + _NEW, + ), + # #361 weaknesses the walk-based reader does not share + "reader_ttir/unsigned_index.ttir": ("divui is Bin('u//') (#361: unmodeled)", _OLD), + "reader_ttir/loop_two_step_advance.ttir": ( + "two addptrs per iteration: delta is their sum (#361: refused)", + _OLD, + ), + "reader_ttir/where_pointer.ttir": ( + "arith.select of same-base pointers selects offsets (#361: refused)", + _OLD, + ), + "ttir/spike_if_yield.ttir": ( + "a same-base pointer yielded by scf.if selects offsets (#361: refused)", + _OLD, + ), + "ttir/golden_atomic_fmax_sm80.ttir": (_BITCAST, _OLD), + "ttir/golden_atomic_fmax_sm90.ttir": (_BITCAST, _OLD), + "ttir/nat_hint_scalar_const.ttir": ( + "arith.constant with a leading attr dict (#361's regex: DataDep)", + _OLD, + ), + "ttir/crafted_odd_names.ttir": ( + "values by identity (#361's regexes miss the names)", + _OLD, + ), + "ttir/crafted_unicode_strings.ttir": ( + "quoted func symbol (#361: no tt.func found)", + _OLD, + ), + "ttir/nat_k_uni.ttir": ( + "a non-ASCII path decodes as UTF-8 (#361: raw \\XX escapes)", + frozenset({"events[0].loc", "events[1].loc"}), + ), + "ttir/nat_uni_params.ttir": ( + "parameters by NameLoc (π_ptr, 数_n), not printed names", + frozenset({"args", "events[0].base", "events[1].base", "events[1].mask"}), + ), + "ttir/crafted_empty_else.ttir": ( + _NAMELOC, + frozenset({"args", "events[0].base", "events[0].path"}), + ), + "ttir/crafted_empty_for.ttir": ( + _NAMELOC, + frozenset({"args", "events[0].base", "loop"}), + ), + "ttir/crafted_locs.ttir": (_NAMELOC, frozenset({"args", "events[0].offset"})), +} + +# The new reader's refusal kind for every refused golden (the rest parse). +REFUSED = { + "ttir/adv_cf_blockargs.ttir": "control-flow", + "ttir/adv_descs.ttir": "out-of-vocabulary", + "ttir/adv_hinted.ttir": "other", # an integer tt.reduce feeds an address + "ttir/adv_multi_func.ttir": "call", + "ttir/adv_multi_result.ttir": "other", # a loop result feeds an address + "ttir/adv_nest3.ttir": "control-flow", # a loop under an scf.if + "ttir/adv_views.ttir": "other", # tt.reshape of a pointer tile + "ttir/adv_while_nested.ttir": "control-flow", + "ttir/crafted_attr_dicts.ttir": "other", + "ttir/crafted_deep_nest.ttir": "control-flow", + "ttir/crafted_fwd_ref_cf.ttir": "control-flow", + "ttir/crafted_generic_form.ttir": "reader-misalignment", + "ttir/crafted_same_dest.ttir": "control-flow", + "ttir/crafted_symbols_strings.ttir": "call", + "ttir/golden_early_return_loaded_sm80.ttir": "control-flow", + "ttir/golden_early_return_pid_sm80.ttir": "control-flow", + "ttir/golden_gather_sm80.ttir": "indirect-address", + "ttir/golden_gather_sm90.ttir": "indirect-address", + "ttir/golden_guard_then_loop_sm80.ttir": "control-flow", + "ttir/golden_loop_under_if_sm80.ttir": "control-flow", + # integer offsets carried by the loop (rewritten block pointers) + "ttir/golden_matmul_bp_s3_sm80.ttir": "loop-variant-advance", + "ttir/golden_matmul_bp_s3_sm90.ttir": "loop-variant-advance", + "ttir/golden_matmul_tma_s1_sm90.ttir": "out-of-vocabulary", + "ttir/golden_matmul_tma_s3_sm90.ttir": "out-of-vocabulary", + "ttir/golden_matmul_tma_ws_s3_sm90.ttir": "out-of-vocabulary", + "ttir/golden_nested_guard_merge_sm80.ttir": "control-flow", + "ttir/golden_nested_loops_sm80.ttir": "nested-loop", + "ttir/golden_sequential_loops_sm80.ttir": "nested-loop", + "ttir/spike_early_return.ttir": "control-flow", + "ttir/spike_early_return_loop.ttir": "control-flow", + "ttir/spike_inline_asm.ttir": "inline-asm", + "ttir/spike_nested_for.ttir": "nested-loop", + "ttir/spike_noinline_call.ttir": "call", + "ttir/spike_spin_while.ttir": "control-flow", + "reader_ttir/p1_variant_delta.ttir": "loop-variant-advance", + "reader_ttir/loop_observed_advance.ttir": "loop-variant-advance", + "reader_ttir/p2_swap.ttir": "loop-variant-advance", + "reader_ttir/p3_call_guarded.ttir": "call", + "reader_ttir/p3_call_offset.ttir": "call", + "reader_ttir/p3_call_formals.ttir": "call", + "reader_ttir/rv_inline_asm_store.ttir": "inline-asm", + "reader_ttir/int_iterarg_offset.ttir": "loop-variant-advance", + "reader_ttir/pure_asm_int_addr.ttir": "inline-asm", + "reader_ttir/observed_lanes.ttir": "indirect-address", +} + + +def _tokens(term, graph, canon: dict) -> tuple: + """Flat pre-order tokens of a term from either reader: IntCast and the + D9 widths dropped, make_range sites renamed by first appearance, the + (single) loop's identity dropped, and an IterArgOffset replaced by its + base, offset0 and delta. Iterative (kernel_deep_chain is deep).""" + out: list[tuple] = [] + stack = [term] + while stack: + t = stack.pop() + if t is None: + out.append(("None",)) + continue + name = type(t).__name__ + if name == "IntCast": + stack.append(t.x) + elif name in ("Bin", "BoolBin"): + out.append((name, t.op)) + stack += [t.b, t.a] + elif name == "Cmp": + out.append((name, t.pred)) + stack += [t.b, t.a] + elif name == "Select": + out.append((name,)) + stack += [t.f, t.t, t.cond] + elif name == "Not": + out.append((name,)) + stack.append(t.a) + elif name == "Const": + out.append((name, t.value)) + elif name in ("Pid", "NumPrograms"): + out.append((name, t.axis)) + elif name == "Arange": + out.append( + (name, canon.setdefault(t.ssa, len(canon)), t.start, t.end, t.dim) + ) + elif name == "Param": + out.append((name, t.name)) + elif name == "LoopVar": + out.append((name,)) + elif name == "IterArgOffset": + info = graph.iter_args[t.arg_id] + out.append((name, info.base_param)) + stack += [info.delta, info.offset0] + elif name == "Observed": + out.append((name, t.access_index)) + elif name == "DataDep": + out.append((name,)) + else: + raise AssertionError(f"unexpected term {name}") + return tuple(out) + + +def _normalize(graph) -> dict: + """What both readers must agree on, as plain comparable data. Left out: + FuncArg.int_bits (new), and elem_float of non-atomic accesses, which + the new reader sets from the pointee on every access (#361: atomics + only).""" + canon: dict = {} + + def tok(t): + return _tokens(t, graph, canon) + + events = [ + { + "kind": a.kind, + "base": a.base_param, + "offset": tok(a.offset), + "mask": tok(a.mask), + "path": tok(a.path), + "in_loop": a.in_loop, + "atomic": None + if a.atomic is None + else (a.atomic.rmw_op, a.atomic.sem, a.atomic.scope), + "elem_bits": a.elem_bits, + "loc": None if a.loc is None else (a.loc.file, a.loc.line, a.loc.col), + "line_no": a.line_no, + "guarded": a.guarded, + "mask_dropped": a.mask_dropped, + "atomic_val": tok(a.atomic_val), + "atomic_cmp": tok(a.atomic_cmp), + "elem_float": a.elem_float if a.atomic is not None else None, + } + for a in graph.accesses + ] + loop = graph.loop + return { + "kernel": graph.kernel_name, + "args": [ + (a.name, a.is_ptr, a.elem_bits, a.elem_float) for a in graph.func_args + ], + "pid_axes": sorted(graph.pid_axes), + "loop": None + if loop is None + else [tok(loop.lower), tok(loop.upper), tok(loop.step)], + "n_iter_args": len(graph.iter_args), + "events": events, + } + + +def _diff(mine: dict, theirs: dict) -> set[str]: + """The normalized fields that differ: ``events[i].`` per event + when both have the same number of events, else top-level keys.""" + out = {k for k in mine if k != "events" and mine[k] != theirs[k]} + if len(mine["events"]) != len(theirs["events"]): + return out | {"events"} + for i, (a, b) in enumerate(zip(mine["events"], theirs["events"])): + out |= {f"events[{i}].{k}" for k in a if a[k] != b[k]} + return out + + +def _run(reader, text: str): + try: + return reader.parse_ttir(text), None + except reader.UnsupportedTTIR as e: + return None, e + + +@pytest.mark.parametrize("name", FILES) +def test_differential_oracle(name): + text = _text(name) + mine, mine_refusal = _run(R, text) + theirs, theirs_refusal = _run(O, text) + # the new reader's outcome is pinned + assert (None if mine_refusal is None else mine_refusal.kind) == REFUSED.get(name) + reason, differs = EXPECTED_DIFF.get(name, ("", frozenset())) + if mine is not None and theirs is not None: + # exactly the listed fields differ; everything else still matches + assert _diff(_normalize(mine), _normalize(theirs)) == differs, reason + elif mine is None and theirs is None: + assert not differs, reason + else: + outcome = _NEW if mine is None else _OLD + assert ( + outcome == differs + ), f"{name}: new {mine_refusal!r} vs #361 {theirs_refusal!r} ({reason})" + + +# How the reader of a later release reads a base golden it prints itself +# (the base text, printed by 3.6, shadowed by its own): name -> refusal kind, +# where it differs from the release's own printing. +BASE_UNDER: dict[str, dict[str, str]] = { + "3.8": { + # 3.6's gpu.barrier: 3.8's frontend emits ttg.barrier + "ttir/adv_zero_result.ttir": "out-of-vocabulary", + "ttir/spike_misc.ttir": "out-of-vocabulary", + # the 3.6 !tt.tensordesc spelling + "ttir/adv_descs.ttir": "unparsable", + "ttir/golden_matmul_tma_s1_sm90.ttir": "unparsable", + "ttir/golden_matmul_tma_s3_sm90.ttir": "unparsable", + "ttir/golden_matmul_tma_ws_s3_sm90.ttir": "unparsable", + }, +} + + +def test_shadowed_base_goldens_read_as_pinned(): + """The installed release reads each base golden it prints itself as its + own printing reads (REFUSED), but where BASE_UNDER pins the difference.""" + shadowed = [n for n in FILES if G.printed_by(_path(n)) != G.BASE_RELEASE] + for name in shadowed: + got = _run(R, (GOLDEN / name).read_text(encoding="utf-8"))[1] + want = BASE_UNDER.get(G.RELEASE, {}).get(name, REFUSED.get(name)) + assert (None if got is None else got.kind) == want, (name, got) + if want == "out-of-vocabulary" and name in BASE_UNDER.get("3.8", {}): + assert "op gpu.barrier is not TTIR" in got.message + # a release other than the base one reads its own printings + assert G.RELEASE == G.BASE_RELEASE or shadowed + + +def test_base_under_names_shadowed_goldens(): + for release, table in BASE_UNDER.items(): + for name in table: + d, n = name.split("/") + assert (G.own_dir(d, release) / n).is_file(), (release, name) + + +def test_oracle_corpus_coverage(): + # the tables name real goldens, and most files are compared event by event + assert set(EXPECTED_DIFF) <= set(FILES) and set(REFUSED) <= set(FILES) + both = [ + f + for f in FILES + if f not in REFUSED and f not in EXPECTED_DIFF and _run(O, _text(f))[0] + ] + assert len(both) >= 45 + assert sum(len(_graph(f).accesses) for f in both) >= 120 + + +def test_oracle_is_the_verbatim_361_reader(): + source = (Path(__file__).parent / "_oracle_ttir_reader_361.py").read_bytes() + body = source.split(b"\n", 5)[5] + assert hashlib.sha256(body).hexdigest() == ( + "3a82e07d3dc2aae16f2d3098894e411c04779ae61c2f787eb2cb15278ff7a4df" + ) + + +# ─────────────────────────── the audit probes ─────────────────────────── + + +def test_p1_loop_variant_advance_refuses(): + e = _refusal("reader_ttir/p1_variant_delta.ttir") + assert e.kind is TTIRKind.LOOP_VARIANT_ADVANCE and "loop-variant" in e.message + assert _REFUSAL_SITE["p1_variant_delta"][G.RELEASE] in _source_line(e.loc) + # #361 accepted it with a LoopVar advance (sanitizer: false 'ok') + g = O.parse_ttir(_text("reader_ttir/p1_variant_delta.ttir")) + assert isinstance(g.iter_args[0].delta, O.LoopVar) + + +def test_p2_swapped_iter_args_refuse(): + e = _refusal("reader_ttir/p2_swap.ttir") + assert ( + e.kind is TTIRKind.LOOP_VARIANT_ADVANCE + and "not advanced from itself" in e.message + ) + g = O.parse_ttir(_text("reader_ttir/p2_swap.ttir")) + assert [i.delta for i in g.iter_args.values()] == [O.Const(0), O.Const(0)] + + +def test_observation_inside_the_loop_is_a_variant_advance(): + e = _refusal("reader_ttir/loop_observed_advance.ttir") + assert e.kind is TTIRKind.LOOP_VARIANT_ADVANCE + assert _REFUSAL_SITE["loop_observed_advance"][G.RELEASE] in _source_line(e.loc) + + +@pytest.mark.parametrize( + "name", ["p3_call_guarded", "p3_call_offset", "p3_call_formals"] +) +def test_p3_calls_refuse(name): + e = _refusal(f"reader_ttir/{name}.ttir") + assert e.kind is TTIRKind.CALL and "_p3_helper" in e.message + + +def test_quoted_callee_and_a_second_function_refuse_as_call(): + e = _refusal("ttir/crafted_symbols_strings.ttir") + assert e.kind is TTIRKind.CALL and "'f{%x} \"q\" (a)'" in e.message + # a second tt.func refuses even when nothing calls it + extra = ( + " tt.func private @h(%q: !tt.ptr) attributes {noinline = true} {\n" + " tt.return\n }\n" + ) + e = _refusal(_module("", extra=extra)) + assert e.kind is TTIRKind.CALL and "'h'" in e.message + + +def test_p4_walkers_resolve_loop_carried_pointers(): + direct = _graph("reader_ttir/p4_observed_direct.ttir") + loop = _graph("reader_ttir/p4_observed_loop.ttir") + delta = _graph("reader_ttir/p4_observed_delta.ttir") + for g in (direct, loop, delta): + load = next(a for a in g.accesses if a.kind == "load") + assert mentions_observed(load.offset, g) + assert observed_indices(load.offset, g) == {0} + # the address reaches the observation only through the iter_arg (its + # offset0, resp. its delta): #361's term-local walker misses it + for g, part, name in ((loop, "offset0", "loop"), (delta, "delta", "delta")): + load = next(a for a in g.accesses if a.kind == "load") + assert isinstance(load.offset, IterArgOffset) + assert mentions_observed(getattr(g.iter_args[0], part), g) + theirs = O.parse_ttir(_text(f"reader_ttir/p4_observed_{name}.ttir")) + theirs_load = next(a for a in theirs.accesses if a.kind == "load") + assert not O.mentions_observed(theirs_load.offset) + + +def test_walkers_descend_datadep_keep(): + g = AccessGraph("k", (), (), None) + kept = DataDep("bool op over loaded data", keep=Cmp("eq", Observed(3), Const(0))) + assert mentions_observed(kept, g) and observed_indices(kept, g) == {3} + assert not mentions_observed(DataDep("loaded value"), g) + + +def test_rv_inline_asm_refuses(): + e = _refusal("reader_ttir/rv_inline_asm_store.ttir") + assert e.kind is TTIRKind.INLINE_ASM and "side effects" in e.message + assert _REFUSAL_SITE["rv_inline_asm_store"][G.RELEASE] in _source_line(e.loc) + # a pure asm with a pointer operand is still an address handed to asm + text = _text("ttir/spike_inline_asm.ttir").replace("pure = false", "pure = true") + e = _refusal(text) + assert e.kind is TTIRKind.INLINE_ASM and "pointer operand" in e.message + # a pure asm over data is plain data + text = "\n".join( + line + for line in _text("ttir/spike_inline_asm.ttir").splitlines() + if "pure = false" not in line + ) + assert [a.kind for a in parse_ttir(text).accesses] == ["load"] + + +# ─────────────────────────── the phase-2 review probes ─────────────────────────── +# ir_mode_audit/probes_phase2/ttir-reader/: compiled kernels (now goldens in +# reader_ttir/) and hand-written TTIR. + + +def _footprint(g, access, params: dict, iters: int) -> set[int]: + """The modeled element offsets of ``access`` over iterations + ``0 .. iters-1`` with every arange lane free, the lanes keyed by + (make_range, dim) as a consumer keys them; mask applied.""" + lanes = sorted( + { + (n.ssa, n.dim, n.start, n.end) + for n in R._nodes((access.offset, access.mask), g.iter_args) + if isinstance(n, Arange) + } + ) + + def ev(t, env): + if isinstance(t, Const): + return t.value + if isinstance(t, Param): + return env[t.name] + if isinstance(t, Arange): + return env[(t.ssa, t.dim)] + if isinstance(t, IterArgOffset): + info = g.iter_args[t.arg_id] + return ev(info.offset0, env) + env["k"] * ev(info.delta, env) + if isinstance(t, IntCast): + return ev(t.x, env) + a, b = ev(t.a, env), ev(t.b, env) + if isinstance(t, Cmp): + return int({"slt": a < b}[t.pred]) + return {"+": a + b, "-": a - b, "*": a * b}[t.op] + + out = set() + for k in range(iters): + for values in itertools.product(*(range(s, e) for _, _, s, e in lanes)): + env = dict(params, k=k) + env.update({(ssa, dim): v for (ssa, dim, _, _), v in zip(lanes, values)}) + if access.mask is None or ev(access.mask, env): + out.add(ev(access.offset, env)) + return out + + +def test_expanded_loop_carried_tiles_keep_their_lanes(): + # k1: [N, N] pointer tile expanded to 3D in the loop, next to the same + # make_range at dim 0 (#361 kept the tile's dims 0/1: a 4-offset footprint) + g = _graph("reader_ttir/expand_iterarg_3d.ttir") + load = g.accesses[0] + assert load.kind == "load" and load.in_loop + (tile,) = [ + n for n in R._nodes((load.offset,), None) if isinstance(n, IterArgOffset) + ] + expanded = g.iter_args[tile.arg_id] + assert expanded.base_param == "x_ptr" and expanded.delta == Const(1) + n = 4 + want = { + j * n + m - i * n + k + for i, j, m in itertools.product(range(n), repeat=3) + for k in range(2) + } + assert _footprint(g, load, {}, 2) == want # [-12, 16], every OOB offset kept + # k2: 1D pointer expanded to 2D, masked at its own lane (dim 1) + g = _graph("reader_ttir/expand_iterarg_mask.ttir") + load = g.accesses[0] + assert load.offset == IterArgOffset(1) and g.iter_args[1].delta == Const(4) + assert _footprint(g, load, {"M": 2}, 2) == {0, 1, 4, 5} + + +def test_tensor_observations_do_not_meet_across_lanes(): + # k8: c[:, None] - c[None, :] over one tensor atomic's old values + e = _refusal("reader_ttir/observed_lanes.ttir") + assert e.kind is TTIRKind.INDIRECT_ADDRESS and "across lanes" in e.message + assert "c[:, None] - c[None, :]" in _source_line(e.loc) + # #361 wrote both lanes as one symbol, offset (0 + X) + (0 - X) with the + # same X: every lane stores to offset 0, a false in-bounds + off = O.parse_ttir(_text("reader_ttir/observed_lanes.ttir")).accesses[-1].offset + assert off.b.op == "-" and off.a.b == off.b.b + assert O.mentions_observed(off.a.b) + # an expanded tile of a pointer that advances by per-lane observations + body = """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %r = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> + %ps = tt.splat %p : !tt.ptr -> tensor<4x!tt.ptr> + %q = tt.addptr %ps, %r : tensor<4x!tt.ptr>, tensor<4xi32> + %one = arith.constant dense<1> : tensor<4xi32> + %old = tt.atomic_rmw add, acq_rel, gpu, %q, %one : (tensor<4x!tt.ptr>, tensor<4xi32>) -> tensor<4xi32> + %res = scf.for %k = %c0 to %n step %c1 iter_args(%a = %q) -> (tensor<4x!tt.ptr>) : i32 { + %e = tt.expand_dims %a {axis = 0 : i32} : tensor<4x!tt.ptr> -> tensor<1x4x!tt.ptr> + %v = tt.load %e : tensor<1x4x!tt.ptr> + %a2 = tt.addptr %a, %old : tensor<4x!tt.ptr>, tensor<4xi32> + scf.yield %a2 : tensor<4x!tt.ptr> + }""" + e = _refusal(_module(body, args="%p: !tt.ptr, %n: i32")) + assert e.kind is TTIRKind.OTHER and "per-lane atomic result" in e.message + + +def test_loop_carried_integers_make_addresses_loop_variant(): + # k3: offs += B carried by the loop + e = _refusal("reader_ttir/int_iterarg_offset.ttir") + assert ( + e.kind is TTIRKind.LOOP_VARIANT_ADVANCE and "carried by the loop" in e.message + ) + # H4: a pointer advanced by a loop-carried integer + body = """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %r:2 = scf.for %i = %c0 to %n step %c1 iter_args(%a = %p, %s = %c0) -> (!tt.ptr, i32) : i32 { + %v = tt.load %a : !tt.ptr + %a2 = tt.addptr %a, %s : !tt.ptr, i32 + %s2 = arith.addi %s, %c1 : i32 + scf.yield %a2, %s2 : !tt.ptr, i32 + }""" + assert _refusal(_module(body)).kind is TTIRKind.LOOP_VARIANT_ADVANCE + + +def test_an_address_handed_to_opaque_ops_as_an_integer_refuses(): + # k6: tl.inline_asm_elementwise(..., is_pure=True) given x_ptr.to(tl.int64) + e = _refusal("reader_ttir/pure_asm_int_addr.ttir") + assert e.kind is TTIRKind.INLINE_ASM and "tt.ptr_to_int" in e.message + # through arithmetic with loaded data, and through a loop-carried value + args = "%p: !tt.ptr, %n: i32" + head = """ + %x = tt.load %p : !tt.ptr + %i = tt.ptr_to_int %p : !tt.ptr -> i64 + %a = arith.addi %i, %x : i64""" + asm = ( + '%y = tt.elementwise_inline_asm "mov.b64 $0, $1;" {constraints = "=l,l", ' + "packed_element = 1 : i32, pure = true} %a : i64 -> i64" + ) + assert _refusal(_module(head + "\n" + asm, args=args)).kind is TTIRKind.INLINE_ASM + loop = """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %z = arith.constant 0 : i64 + %i = tt.ptr_to_int %p : !tt.ptr -> i64 + %a = scf.for %k = %c0 to %n step %c1 iter_args(%s = %z) -> (i64) : i32 { + %s2 = arith.addi %s, %i : i64 + scf.yield %s2 : i64 + }""" + assert _refusal(_module(loop + "\n" + asm, args=args)).kind is TTIRKind.INLINE_ASM + extern = ( + '%y = tt.extern_elementwise %a {libname = "", libpath = "", pure = true, ' + 'symbol = "f"} : (i64) -> i64' + ) + e = _refusal(_module(head + "\n" + extern, args=args)) + assert e.kind is TTIRKind.OUT_OF_VOCABULARY and "tt.ptr_to_int" in e.message + # loaded data carries no address: an asm over it stays plain data + data = head.replace("%i, %x", "%x, %x") + "\n" + asm + assert [a.kind for a in parse_ttir(_module(data, args=args)).accesses] == ["load"] + + +def test_llvm_and_gpu_ops_outside_the_inert_ones_refuse(): + # H1a-c: memory effects through ops of the llvm dialect + cases = { + "llvm.inttoptr": """ + %i = tt.ptr_to_int %p : !tt.ptr -> i64 + %lp = llvm.inttoptr %i : i64 to !llvm.ptr<1> + %c = arith.constant 7 : i32 + %old = llvm.atomicrmw add %lp, %c monotonic : !llvm.ptr<1>, i32""", + "llvm.inline_asm": """ + %i = tt.ptr_to_int %p : !tt.ptr -> i64 + %c = arith.constant 7 : i32 + %r = llvm.inline_asm has_side_effects "st.global.b32 [$1], $2; mov.b32 $0, 0;", "=r,l,r" %i, %c : (i64, i32) -> i32""", + } + for name, body in cases.items(): + e = _refusal(_module(body, args="%p: !tt.ptr, %n: i32")) + assert e.kind is TTIRKind.OUT_OF_VOCABULARY and name in e.message + # inside a combine region too + body = """ + %r = tt.make_range {end = 4 : i32, start = 0 : i32} : tensor<4xi32> + %s = "tt.reduce"(%r) <{axis = 0 : i32}> ({ + ^bb0(%a: i32, %b: i32): + %i = tt.ptr_to_int %p : !tt.ptr -> i64 + %lp = llvm.inttoptr %i : i64 to !llvm.ptr<1> + %y = arith.addi %a, %b : i32 + tt.reduce.return %y : i32 + }) : (tensor<4xi32>) -> i32""" + e = _refusal(_module(body, args="%p: !tt.ptr, %n: i32")) + assert e.kind is TTIRKind.OUT_OF_VOCABULARY and "llvm.inttoptr" in e.message + # the inert ones stay inert: the release's barrier (3.8's ttg.barrier is + # the one ttg op accepted); 3.8 still parses gpu.barrier, which its + # frontend never emits: refused + barrier, other = _BARRIERS[G.RELEASE] + g = parse_ttir(_module(f"{barrier}\n%v = tt.load %p : !tt.ptr")) + assert [a.kind for a in g.accesses] == ["load"] + e = _refusal(_module(f"{other}\n%v = tt.load %p : !tt.ptr")) + assert e.kind is _OTHER_BARRIER[G.RELEASE], e.message + if G.RELEASE == "3.8": + assert "op gpu.barrier is not TTIR" in e.message + e = _refusal(_module("ttg.local_barrier\n%v = tt.load %p : !tt.ptr")) + assert e.kind in (TTIRKind.UNPARSABLE, TTIRKind.OUT_OF_VOCABULARY) + + +# release -> (the barrier tl.debug_barrier() prints, the other release's) +_BARRIERS = { + "3.6": ("gpu.barrier", "ttg.barrier all"), + "3.8": ("ttg.barrier all", "gpu.barrier"), +} +# what the other release's barrier gets: 3.6 has no ttg dialect +_OTHER_BARRIER = {"3.6": TTIRKind.UNPARSABLE, "3.8": TTIRKind.OUT_OF_VOCABULARY} + + +def test_pointers_outside_global_memory_refuse(): + # H8: a shared-memory (address space 3) pointer argument + e = _refusal( + _module( + "%c1 = arith.constant 1 : i32\ntt.store %p, %c1 : !tt.ptr", + args="%p: !tt.ptr, %n: i32", + ) + ) + assert e.kind is TTIRKind.OUT_OF_VOCABULARY and "address space" in e.message + # a pointer made into another address space + body = """ + %c = arith.constant 0 : i64 + %q = tt.int_to_ptr %c : i64 -> !tt.ptr + %v = tt.load %q : !tt.ptr""" + e = _refusal(_module(body)) + assert e.kind is TTIRKind.OUT_OF_VOCABULARY and "address space" in e.message + assert e.line_no == 4 # the tt.int_to_ptr + + +def test_generic_load_reads_its_mask_by_operand_segment(): + # G1: generic tt.load with `other` and no mask; an i1 `other` is no mask + load = ( + '%v = "tt.load"(%p, %f) <{boundaryCheck = array, cache = 1 : i32, ' + "evict = 1 : i32, isVolatile = false, operandSegmentSizes = array}> " + ": (!tt.ptr, i1) -> i1" + ) + for segments, mask in (("1, 0, 1", None), ("1, 1, 0", Const(0))): + g = parse_ttir( + _module( + "%f = arith.constant false\n" + load.replace("SEG", segments), + args="%p: !tt.ptr, %n: i32", + ) + ) + (access,) = g.accesses + assert (access.mask, access.mask_dropped) == (mask, False), segments + + +# ─────────────────────────── modeled shapes ─────────────────────────── + + +def test_two_step_advance_is_the_sum(): + g = _graph("reader_ttir/loop_two_step_advance.ttir") + (info,) = g.iter_args + assert (info.base_param, info.offset0, info.delta) == ( + "x_ptr", + Const(0), + Bin("+", Param("s"), Const(2)), + ) + assert g.accesses[0].offset == IterArgOffset(0) and g.accesses[0].in_loop + + +def test_same_base_pointer_selects(): + g = _graph("reader_ttir/where_pointer.ttir") + (store,) = g.accesses + assert store.base_param == "x_ptr" and isinstance(store.offset, Select) + g = _graph("ttir/spike_if_yield.ttir") + store = g.accesses[-1] + assert store.kind == "store" and store.base_param == "out_ptr" + assert isinstance(store.offset, Select) and isinstance(store.offset.cond, Cmp) + + +def test_three_dims_of_one_make_range_are_three_lanes(): + (store,) = _graph("reader_ttir/tile3d_shared_arange.ttir").accesses + ranges = [n for n in R._nodes((store.offset,), None) if isinstance(n, Arange)] + assert len({r.ssa for r in ranges}) == 1 + assert sorted(r.dim for r in ranges) == [0, 1, 2] + + +def test_same_width_pointer_bitcast_keeps_the_base(): + g = _graph("ttir/golden_atomic_fmax_sm80.ttir") + atomics = [a for a in g.accesses if a.kind == "atomic_rmw"] + assert [a.atomic.rmw_op for a in atomics] == ["max", "umin"] + assert all( + a.base_param == "out_ptr" and a.elem_bits == 32 and not a.elem_float + for a in atomics + ) + # the atomics' mask reads loaded data + assert all(a.mask_dropped and a.mask is None for a in atomics) + # a bitcast to another element width changes what an element offset means + body = """ + %q = tt.bitcast %p : !tt.ptr -> !tt.ptr + %v = tt.load %q : !tt.ptr""" + e = _refusal(_module(body)) + assert e.kind is TTIRKind.OTHER and "element width" in e.message + + +# The block-pointer case per release: 3.6 has the ops and types; 3.8 has +# neither (its parser aborts on the type), so the walk refuses the text +# before parsing. +_BLOCK_POINTER_KIND = {"3.6": "block-pointer", "3.8": "unparsable"} + + +def test_other_refusals(): + cases = { + _BLOCK_POINTER_KIND[G.RELEASE]: """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i64 + %c32 = arith.constant 32 : i64 + %bp = tt.make_tensor_ptr %p, [%c32, %c32], [%c32, %c1], [%c0, %c0] {order = array} : > + %v = tt.load %bp : !tt.ptr>""", + "out-of-vocabulary": """ + %r = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> + %v = arith.sitofp %r : tensor<64xi32> to tensor<64xf32> + %s = "tt.reduce"(%v) <{axis = 0 : i32}> ({ + ^bb0(%a: f32, %b: f32): + %x = tt.load %p : !tt.ptr + %y = arith.addf %a, %x : f32 + tt.reduce.return %y : f32 + }) : (tensor<64xf32>) -> f32""", + "control-flow": "%c0 = arith.constant 0 : i32\n" + "%b = arith.cmpi slt, %n, %c0 : i32\n" + 'cf.assert %b, "neg"', + } + for kind, body in cases.items(): + assert _refusal(_module(body)).kind == kind, kind + if G.RELEASE == "3.8": + e = _refusal(_module(cases["unparsable"])) + assert "no block pointers" in e.message and e.line_no == 7 + e = _refusal( + _module( + "%x = tt.load %p : !tt.ptr\n%y = tt.extern_elementwise %x " + '{libname = "", libpath = "", pure = false, symbol = "foo"} : (f32) -> f32' + ) + ) + assert e.kind is TTIRKind.OUT_OF_VOCABULARY and "'foo'" in e.message + # deep structured nesting refuses instead of exhausting the recursion limit + depth = 250 + nest = ( + "".join("scf.if %c {\n" for _ in range(depth)) + + "tt.store %p, %f : !tt.ptr\n" + + "}\n" * depth + ) + e = _refusal( + _module( + "%f = arith.constant 1.0 : f32\n" + nest, args="%p: !tt.ptr, %c: i1" + ) + ) + assert e.kind is TTIRKind.OTHER and "nested deeper" in e.message + + +def test_pure_extern_and_signed_i1_compare_are_data(): + g = parse_ttir( + _module( + "%x = tt.load %p : !tt.ptr\n%y = tt.extern_elementwise %x " + '{libname = "", libpath = "", pure = true, symbol = "__nv_expf"} : (f32) -> f32\n' + "tt.store %p, %y : !tt.ptr" + ) + ) + assert [a.kind for a in g.accesses] == ["load", "store"] + # a signed compare of i1 reads true as -1: not the boolean model, so the + # mask is dropped (widened), never misread + g = parse_ttir( + _module( + """ + %c0 = arith.constant 0 : i32 + %b = arith.cmpi slt, %n, %c0 : i32 + %t = arith.constant true + %c = arith.cmpi slt, %b, %t : i1 + %v = tt.load %p, %c : !tt.ptr""" + ) + ) + (load,) = g.accesses + assert load.mask is None and load.mask_dropped + + +def test_matmul_loop_iter_args_and_params(): + g = _graph("ttir/golden_matmul_s3_sm80.ttir") + assert g.kernel_name == "matmul_kernel" and g.loop is not None + assert [i.base_param for i in g.iter_args] == ["a_ptr", "b_ptr"] + assert all( + i.arg_id == k and i.loop_ssa == g.loop.loop_ssa + for k, i in enumerate(g.iter_args) + ) + assert Param("K") in list(R._nodes((g.loop.upper,), None)) + loads = [a for a in g.accesses if a.kind == "load"] + assert [a.offset for a in loads] == [IterArgOffset(0), IterArgOffset(1)] + assert all(a.in_loop for a in loads) and not g.accesses[-1].in_loop + assert g.loop.bits == 32 and not g.loop.unsigned and g.loop.line_no is not None + assert g.arg("K").int_bits == 32 and g.arg("a_ptr").elem_bits == 16 + + +# ─────────────────────────── integer widths (D9) ─────────────────────────── + + +def _value(t, env: dict) -> int: + """The unbounded-integer reading of a (shallow) term: casts as identity.""" + if isinstance(t, Const): + return t.value + if isinstance(t, Pid): + return env["pid"] + if isinstance(t, Param): + return env[t.name] + if isinstance(t, IntCast): + return _value(t.x, env) + if isinstance(t, Cmp): + a, b = _value(t.a, env), _value(t.b, env) + return int({"slt": a < b, "ult": a < b, "eq": a == b}[t.pred]) + if isinstance(t, Bin): + a, b = _value(t.a, env), _value(t.b, env) + ops = {"+": a + b, "*": a * b, "-": a - b} + if t.op in ("u//", "//"): + return int(a / b) + if t.op == "%": + return a - b * int(a / b) + return ops[t.op] + raise AssertionError(type(t).__name__) + + +def _holds(ob, env: dict) -> bool: + v = _value(ob.term, env) + if ob.signed: + return -(1 << (ob.bits - 1)) <= v < (1 << (ob.bits - 1)) + return 0 <= v < (1 << ob.bits) + + +def _integer_nodes(g): + for a in g.accesses: + yield from R._nodes((a.offset, a.mask, a.path), g.iter_args) + if g.loop is not None: + yield from R._nodes((g.loop.lower, g.loop.upper, g.loop.step), None) + + +@pytest.mark.parametrize("name", [f for f in FILES if f not in REFUSED]) +def test_every_integer_op_carries_its_width(name): + for n in _integer_nodes(_graph(name)): + if isinstance(n, Bin): + # bits=None only for the element-offset sums of tt.addptr + assert isinstance(n.bits, int) or (n.bits is None and n.op == "+") + elif isinstance(n, Cmp): + assert isinstance(n.bits, int) and n.bits >= 1 + elif isinstance(n, IntCast): + assert n.kind in ("trunci", "extsi", "extui") and n.src_bits != n.dst_bits + + +def test_i32_wrap_obligations(): + g = _graph("reader_ttir/rv_i32_wrap.ttir") + (store,) = g.accesses + inner = Bin("*", Pid(0), Param("S"), 32) + outer = Bin("*", inner, Param("S"), 32) + assert store.offset == Bin("+", Const(0), outer) + obs = width_obligations(g, store) + assert [(o.term, o.bits, o.signed) for o in obs] == [ + (outer, 32, True), + (inner, 32, True), + ] + assert all("pid * S" in _source_line(o.loc) for o in obs) + text = _text("reader_ttir/rv_i32_wrap.ttir").splitlines() + assert all("arith.muli" in text[o.line_no - 1] for o in obs) + # (pid * 65536) * 65536 is 0 in i32: the unbounded model is exact only for pid 0 + assert all(_holds(o, {"pid": 0, "S": 65536}) for o in obs) + assert [_holds(o, {"pid": 1, "S": 65536}) for o in obs] == [False, True] + + +def test_trunci_obligation(): + g = _graph("reader_ttir/rv_trunci_alias.ttir") + (store,) = g.accesses + wide = Bin("*", IntCast("extsi", 32, 64, Pid(0)), Const(1 << 32), 64) + assert store.offset == Bin("+", Const(0), IntCast("trunci", 64, 32, wide)) + obs = width_obligations(g, store) + assert [(o.term, o.bits, o.signed) for o in obs] == [ + (wide, 32, True), + (wide, 64, True), + ] + trunc = obs[0] + assert ( + "arith.trunci" + in _text("reader_ttir/rv_trunci_alias.ttir").splitlines()[trunc.line_no - 1] + ) + assert ".to(tl.int32)" in _source_line(trunc.loc) + # trunc_i32(pid * 2**32) is 0 for every pid: exact only for pid 0 + assert _holds(trunc, {"pid": 0}) and not _holds(trunc, {"pid": 1}) + assert _holds(obs[1], {"pid": 1}) + + +def test_unsigned_reads_need_non_negative_operands(): + g = _graph("reader_ttir/unsigned_index.ttir") + (store,) = g.accesses + quotient = Bin("u//", Pid(0), Const(3), 32) + assert store.offset == Bin("+", Const(0), IntCast("extui", 32, 64, quotient)) + assert store.mask == Cmp("ult", Pid(0), Param("n"), 32) + obs = {(o.term, o.bits, o.signed) for o in width_obligations(g, store)} + assert obs == { + (quotient, 32, False), # the extui operand + (quotient, 32, True), # the divui result + (Pid(0), 32, False), + (Const(3), 32, False), + (Param("n"), 32, False), # cmpi ult reads n unsigned + } + obs = width_obligations(g, store) + assert all(_holds(o, {"pid": 5, "n": 8}) for o in obs) + # n = -1 reads as 2**32 - 1 under ult: the unbounded model is not exact + assert not all(_holds(o, {"pid": 5, "n": -1}) for o in obs) + + +def test_narrowing_casts_and_extui_in_spike_casts(): + g = _graph("ttir/spike_casts.ttir") + load = g.accesses[0] + casts = [n for n in R._nodes((load.offset,), None) if isinstance(n, IntCast)] + assert {(c.kind, c.src_bits, c.dst_bits) for c in casts} == { + ("trunci", 32, 16), + ("extsi", 16, 64), + ("trunci", 32, 8), + ("extui", 8, 64), + } + obs = [ + (o.bits, o.signed) for o in width_obligations(g, load) if o.line_no is not None + ] + assert (16, True) in obs and (8, True) in obs and (8, False) in obs + + +def test_i1_casts_and_unsigned_loops(): + # extsi of an i1 maps true to -1: read as 0 - extui(b) + g = parse_ttir( + _module( + """ + %c0 = arith.constant 0 : i32 + %b = arith.cmpi slt, %n, %c0 : i32 + %e = arith.extsi %b : i1 to i32 + %q = tt.addptr %p, %e : !tt.ptr, i32 + %v = tt.load %q : !tt.ptr""" + ) + ) + b = Cmp("slt", Param("arg1"), Const(0), 32) + assert g.accesses[0].offset == Bin( + "+", Const(0), Bin("-", Const(0), IntCast("extui", 1, 32, b), 32) + ) + # trunci to i1 keeps only bit 0: exact for 0 / 1 + g = parse_ttir( + _module("%b = arith.trunci %n : i32 to i1\n%v = tt.load %p, %b : !tt.ptr") + ) + ((ob,),) = [width_obligations(g, a) for a in g.accesses] + assert (ob.term, ob.bits, ob.signed) == (Param("arg1"), 1, False) + # an unsigned loop compare needs non-negative bounds + g = parse_ttir( + _module( + """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + scf.for unsigned %i = %c0 to %n step %c1 : i32 { + %q = tt.addptr %p, %i : !tt.ptr, i32 + %v = tt.load %q : !tt.ptr + }""" + ) + ) + assert g.loop.unsigned and g.loop.bits == 32 + obs = [ + (o.term, o.bits, o.signed, o.line_no) + for o in width_obligations(g, g.accesses[0]) + ] + latch = Bin("+", Bin("-", Param("arg1"), Const(1), 32), Const(1), 32) + assert obs == [ + (Const(0), 32, False, 5), + (Param("arg1"), 32, False, 5), + (Const(1), 32, False, 5), + (latch, 32, False, 5), # the increment does not wrap, read unsigned + ] + + +def test_loop_increment_obligation(): + # k4: for k in range(lo, n, 1 << 20) wraps its induction variable when + # n is near INT32_MAX (the GPU then runs iterations at negative k) + g = _graph("reader_ttir/iv_wrap.ttir") + (store,) = g.accesses + latch = Bin("+", Bin("-", Param("n"), Const(1), 32), Const(1 << 20), 32) + obs = width_obligations(g, store) + assert [(o.term, o.bits, o.signed, o.role) for o in obs] == [ + (latch, 32, True, "loop") + ] + assert "range(lo, n, 1 << 20)" in _source_line(obs[0].loc) + assert ( + "scf.for" in _text("reader_ttir/iv_wrap.ttir").splitlines()[obs[0].line_no - 1] + ) + assert _holds(obs[0], {"n": 1000}) + assert not _holds(obs[0], {"n": (1 << 31) - (1 << 19)}) + # an access outside the loop gets no loop obligations + g = _graph("ttir/golden_matmul_s3_sm80.ttir") + assert not any(o.role == "loop" for o in width_obligations(g, g.accesses[-1])) + assert any(o.role == "loop" for o in width_obligations(g, g.accesses[0])) + + +def test_signed_remainder_needs_its_quotient_to_fit(): + # H9: remsi(INT_MIN, -1) is undefined, though its value 0 fits + body = """ + %cm = arith.constant -2147483648 : i32 + %cn = arith.constant -1 : i32 + %r = arith.remsi %cm, %cn : i32 + %q = tt.addptr %p, %r : !tt.ptr, i32 + %v = tt.load %q : !tt.ptr""" + g = parse_ttir(_module(body)) + (load,) = g.accesses + lo = Const(-(1 << 31)) + obs = width_obligations(g, load) + assert [(o.term, o.bits, o.signed, o.line_no) for o in obs] == [ + (Bin("%", lo, Const(-1), 32), 32, True, 5), + (Bin("//", lo, Const(-1), 32), 32, True, 5), + ] + assert _holds(obs[0], {}) and not _holds(obs[1], {}) + + +def test_obligation_roles(): + # a node the path, mask and offset share is listed once, under the most + # restrictive role (path, then mask, then offset) + body = """ + %c4 = arith.constant 4 : i32 + %pid = tt.get_program_id x : i32 + %o = arith.muli %pid, %c4 : i32 + %s = arith.subi %n, %c4 : i32 + %b = arith.cmpi slt, %o, %s : i32 + scf.if %b { + %t = arith.addi %o, %c4 : i32 + %m = arith.cmpi slt, %t, %n : i32 + %o2 = arith.addi %t, %c4 : i32 + %q = tt.addptr %p, %o2 : !tt.ptr, i32 + %v = tt.load %q, %m : !tt.ptr + }""" + g = parse_ttir(_module(body)) + (load,) = g.accesses + o = Bin("*", Pid(0), Const(4), 32) + s = Bin("-", Param("arg1"), Const(4), 32) + t = Bin("+", o, Const(4), 32) + o2 = Bin("+", t, Const(4), 32) + assert [(ob.term, ob.role) for ob in width_obligations(g, load)] == [ + (o, "path"), + (s, "path"), + (t, "mask"), + (o2, "offset"), + ] + + +def test_obligation_sites_do_not_affect_term_equality(): + a = Bin("+", Pid(0), Const(1), 32, line_no=3, loc=W.SourceLoc("a.py", 1, 1)) + b = Bin("+", Pid(0), Const(1), 32, line_no=9, loc=None) + assert a == b and hash(a) == hash(b) and repr(a) == repr(b) + assert Bin("+", Pid(0), Const(1), 64) != a # the width is part of the value + + +def test_deep_terms_stay_iterative(): + g = _graph("ttir/kernel_deep_chain.ttir") + (store,) = g.accesses + assert sum(1 for n in R._nodes((store.offset,), None) if isinstance(n, Bin)) > 1000 + obs = width_obligations(g, store) + assert len(obs) > 1000 and all(o.signed and o.bits == 32 for o in obs) + assert not mentions_observed(store.offset, g) + # pickle round-trips such a graph; the generated ==, hash (and repr, + # deepcopy) recurse, as the module docstring says + assert _fingerprint(pickle.loads(pickle.dumps(g))) == _fingerprint(g) + again = _graph("ttir/kernel_deep_chain.ttir").accesses[0].offset + # at Python's default limit: importing tilelens.visualizer.draw (e.g. from + # tests/unit/test_trace_io.py at collection) raises it process-wide + limit = sys.getrecursionlimit() + sys.setrecursionlimit(1000) + try: + with pytest.raises(RecursionError): + hash(store.offset) + with pytest.raises(RecursionError): + _ = store.offset == again + finally: + sys.setrecursionlimit(limit) + + +# ─────────────────────────── frozen, deterministic graphs ─────────────────────────── + + +_FROZEN_LEAVES = (int, str, float, bool, type(None), TTIRKind) + + +def _assert_deeply_frozen(obj) -> None: + stack = [obj] + while stack: + x = stack.pop() + if isinstance(x, _FROZEN_LEAVES): + continue + if isinstance(x, (tuple, frozenset)): + stack.extend(x) + elif dataclasses.is_dataclass(x): + assert type(x).__dataclass_params__.frozen, type(x).__name__ + stack.extend(getattr(x, f.name) for f in dataclasses.fields(x)) + else: + raise AssertionError(f"mutable {type(x).__name__} in the graph") + + +def _fingerprint(obj) -> str: + """Every field (the compare=False sites included), iteratively.""" + out: list[str] = [] + stack = [obj] + while stack: + x = stack.pop() + if dataclasses.is_dataclass(x): + out.append(type(x).__name__) + stack.extend(reversed([getattr(x, f.name) for f in dataclasses.fields(x)])) + elif isinstance(x, (tuple, frozenset)): + items = sorted(x) if isinstance(x, frozenset) else list(x) + out.append(f"{type(x).__name__}{len(items)}") + stack.extend(reversed(items)) + else: + out.append(repr(x)) + return hashlib.sha256("\x00".join(out).encode()).hexdigest() + + +ACCEPTED = [f for f in FILES if f not in REFUSED] + + +def test_graphs_are_frozen(): + g = _graph("ttir/golden_matmul_s3_sm80.ttir") + for obj, attr in ( + (g, "accesses"), + (g.accesses[0], "offset"), + (g.accesses[0].offset, "arg_id"), + (g.iter_args[0], "delta"), + (g.loop, "upper"), + (g.func_args[0], "name"), + ): + with pytest.raises(dataclasses.FrozenInstanceError): + setattr(obj, attr, None) + for name in ACCEPTED: + _assert_deeply_frozen(_graph(name)) + # hand-built graphs are coerced to immutable containers too + built = AccessGraph("k", [], [], None, iter_args=[], pid_axes={0}) + assert (built.func_args, built.accesses, built.iter_args, built.pid_axes) == ( + (), + (), + (), + frozenset({0}), + ) + + +def test_parse_is_deterministic(): + for name in ACCEPTED: + first = _fingerprint(_graph(name)) + W._CACHE.clear() # a fresh bindings parse + assert _fingerprint(_graph(name)) == first, name + + +SMALL = [ + "ttir/golden_matmul_s3_sm80.ttir", + "ttir/spike_atomics.ttir", + "reader_ttir/p4_observed_loop.ttir", +] + + +def test_graphs_hash_pickle_and_copy(): + for name in SMALL: + g = _graph(name) + W._CACHE.clear() + again = _graph(name) + assert g == again and hash(g) == hash(again) + assert pickle.loads(pickle.dumps(g)) == g + assert copy.deepcopy(g) == g + assert _fingerprint(pickle.loads(pickle.dumps(g))) == _fingerprint(g) + + +def test_parse_is_deterministic_across_processes(): + script = ( + "import hashlib, pickle, sys\n" + "from tilelens.ir.ttir_reader import parse_ttir\n" + "for p in sys.argv[1:]:\n" + " g = parse_ttir(open(p, encoding='utf-8').read())\n" + " print(hashlib.sha256(pickle.dumps(g, protocol=4)).hexdigest())\n" + ) + paths = [str(_path(n)) for n in SMALL] + runs = [] + for seed in ("0", "12345"): + env = dict(os.environ, PYTHONHASHSEED=seed) + out = subprocess.run( + [sys.executable, "-c", script, *paths], + cwd=REPO, + env=env, + capture_output=True, + text=True, + timeout=300, + ) + assert out.returncode == 0, out.stderr[-2000:] + runs.append(out.stdout.split()) + here = [ + hashlib.sha256(pickle.dumps(_graph(n), protocol=4)).hexdigest() for n in SMALL + ] + assert runs[0] == runs[1] == here diff --git a/tests/unit/ir/test_verdict_io.py b/tests/unit/ir/test_verdict_io.py new file mode 100644 index 000000000..2537a50b9 --- /dev/null +++ b/tests/unit/ir/test_verdict_io.py @@ -0,0 +1,305 @@ +"""Persistence of the IR-mode records (D20): the tilelens.ir.verdict records +and the compiled sanitizer's findings are plain data that tilelens.save() / +tilelens.load() round-trip, with a public source-location type in place of +the reader's private one. No GPU. +""" + +from __future__ import annotations + +import dataclasses +import enum +import gc +import importlib +import subprocess +import sys +import weakref +import zipfile +from pathlib import Path +from types import MappingProxyType, SimpleNamespace + +import numpy as np +import pytest +import torch + +import tilelens +from tilelens.clients.sanitizer.data import CompiledSanitizerRecord +from tilelens.core.data import Launch, Load, Store +from tilelens.ir import _mlir_walk as W +from tilelens.ir.launch import TensorFacts, tensor_facts +from tilelens.ir.ttir_reader import TTIRKind, UnsupportedTTIR +from tilelens.ir.verdict import ConfigVerdict, IRVerdict, Refusal, SourceLocation +from tilelens.utils.traceback_utils import TracebackInfo + +trace_module = importlib.import_module("tilelens.core.trace") +REPO = Path(__file__).resolve().parents[3] + + +@pytest.mark.parametrize( + "module", ["tilelens.ir.verdict", "tilelens.clients.sanitizer.data"] +) +def test_record_modules_import_without_triton(module): + code = ( + "import sys\n" + f"import {module}\n" + "loaded = sorted(m for m in sys.modules if m.split('.')[0] == 'triton')\n" + "assert not loaded, loaded\n" + ) + subprocess.run([sys.executable, "-c", code], cwd=REPO, check=True) + + +# ======== source locations and refusals ========= + + +def test_refusal_holds_the_reader_loc_as_a_public_source_location(): + exc = UnsupportedTTIR( + TTIRKind.CONTROL_FLOW, "scf.while", line_no=7, loc=W.SourceLoc("k.py", 3, 5) + ) + refusal = Refusal.from_exception(exc) + + assert refusal == Refusal( + "control-flow", "scf.while", 7, SourceLocation("k.py", 3, 5) + ) + assert type(refusal.loc) is SourceLocation + # The reader's enum kind is held as its plain string, as a load gives it. + assert refusal.kind == TTIRKind.CONTROL_FLOW + assert not isinstance(refusal.kind, enum.Enum) + # Building one directly converts the reader's loc the same way; an + # object without a column gets None. + assert Refusal("call", "m", loc=W.SourceLoc("k.py", 3, 5)).loc == refusal.loc + no_col = Refusal("call", "m", loc=SimpleNamespace(file="k.py", line=2)) + assert no_col.loc == SourceLocation("k.py", 2, None) + assert Refusal("call", "m").loc is None + + +@pytest.mark.parametrize( + "fields, match", + [ + ({"loc": ("k.py", 3, 5)}, "Refusal.loc must be a SourceLocation"), + ({"loc": "k.py:3"}, "Refusal.loc must be a SourceLocation"), + ({"kind": None}, "Refusal.kind must be a str"), + ({"message": ValueError("m")}, "Refusal.message must be a str"), + ], +) +def test_refusal_rejects_fields_a_trace_cannot_hold(fields, match): + with pytest.raises(TypeError, match=match): + Refusal(**{"kind": "call", "message": "m", **fields}) + + +def test_record_hashes_agree_with_equality(): + loc = SourceLocation("k.py", 3, 5) + from_enum = Refusal(TTIRKind.CALL, "m", 4, W.SourceLoc("k.py", 3, 5)) + plain = Refusal("call", "m", 4, loc) + assert from_enum == plain and hash(from_enum) == hash(plain) + assert len({from_enum, plain}) == 1 + assert hash(loc) == hash(SourceLocation("k.py", 3, 5)) + # The verdicts hold a config dict: explicitly unhashable, not a hash + # that fails only for some field values. + for verdict in (ConfigVerdict("h", {}, "ok"), IRVerdict("toy_ir", "ok")): + assert type(verdict).__hash__ is None + with pytest.raises(TypeError, match="unhashable"): + hash(verdict) + + +class _Note(str, enum.Enum): + TIMEOUT = "solver timed out" + + +def test_verdict_notes_are_a_tuple_of_plain_strings(): + verdict = IRVerdict("toy_ir", "ok", notes=(n for n in ["a", _Note.TIMEOUT])) + assert verdict.notes == ("a", "solver timed out") + assert not any(isinstance(note, enum.Enum) for note in verdict.notes) + assert IRVerdict("toy_ir", "ok").notes == () + with pytest.raises(TypeError, match="notes takes a sequence, not a str"): + IRVerdict("toy_ir", "ok", notes="solver timed out") + with pytest.raises(TypeError, match="IRVerdict note must be a str, not bytes"): + IRVerdict("toy_ir", "ok", notes=[b"raw"]) # type: ignore[list-item] + with pytest.raises(TypeError, match="IRVerdict note must be a str, not int"): + IRVerdict("toy_ir", "ok", notes=b"ab") # type: ignore[arg-type] + with pytest.raises(TypeError, match="per_config items must be ConfigVerdicts"): + IRVerdict("toy_ir", "ok", per_config=[{"BLOCK": 16}]) # type: ignore[list-item] + + +# ======== save / load ========= + + +def _save_and_load(tmp_path, monkeypatch, records): + monkeypatch.setattr( + trace_module, "launches", [Launch(grid=(4, 1, 1), records=records)] + ) + path = tilelens.save(tmp_path / "trace.tvz") + with zipfile.ZipFile(path) as archive: + manifest = archive.read("manifest.json").decode() + (launch,) = tilelens.load(path) + return launch, manifest + + +def _refused_verdict() -> IRVerdict: + refusal = Refusal.from_exception( + UnsupportedTTIR( + "control-flow", "scf.while", line_no=7, loc=W.SourceLoc("k.py", 3, 5) + ) + ) + return IRVerdict( + "sanitizer_ir", + "unsupported", + scope="launch", + refusal=refusal, + per_config=[ + ConfigVerdict( + "hash-a", MappingProxyType({"BLOCK": 16, "num_warps": 4}), "proved" + ), + ConfigVerdict(None, {"BLOCK": 64}, "refused", refusal, n_reports=2), + ConfigVerdict( + "hash-c", + {"BLOCK": 32}, + "refused", + Refusal("solver-unknown", "timeout after 10 s"), + ), + ], + notes=["1 of 3 configs proved"], + ) + + +def test_a_saved_trace_holds_ir_verdicts(tmp_path, monkeypatch): + verdict = _refused_verdict() + + launch, manifest = _save_and_load(tmp_path, monkeypatch, ["report", verdict]) + + assert launch.records == ["report", verdict] + loaded = launch.records[1] + assert type(loaded) is IRVerdict and loaded is not verdict + assert type(loaded.per_config) is tuple and type(loaded.notes) is tuple + assert [type(c.config) for c in loaded.per_config] == [dict] * 3 + assert type(loaded.refusal.loc) is SourceLocation + assert not isinstance(loaded.refusal.kind, enum.Enum) + assert hash(loaded.refusal) == hash(verdict.refusal) + assert loaded.per_config[1].refusal == loaded.refusal + # The manifest names only public record types. + assert "tilelens.ir.verdict:SourceLocation" in manifest + assert "_mlir_walk" not in manifest + + +def _findings(facts: TensorFacts) -> list[CompiledSanitizerRecord]: + traceback = TracebackInfo("k.py", 3, "kernel", "x = tl.load(x_ptr + offs)") + return [ + CompiledSanitizerRecord( + kind="out-of-bounds", + op_type=Load, + tensor_name="x_ptr", + tensor_facts=facts, + witness={"pid_0": np.int64(3), "arange_0_d0": 5}, + config=MappingProxyType({"BLOCK": 64}), + user_code_tracebacks=(traceback,), + violation_offset=np.int64(-2), + violation_address=facts.data_ptr - 2 * facts.elem_size, + detail="mask is live at the witness", + ), + CompiledSanitizerRecord( + kind="integer-overflow", + op_type=Store, + tensor_name="out_ptr", + tensor_facts=facts, + witness={"pid_0": 2**20}, + config={}, + user_code_tracebacks=[traceback], + detail="pid * 4096 overflows i32", + ), + CompiledSanitizerRecord( + kind="division-by-zero", + op_type=Load, + tensor_name="x_ptr", + tensor_facts=facts, + witness={"pid_0": 0}, + config={"BLOCK": 16}, + user_code_tracebacks=[], + ), + ] + + +def test_compiled_sanitizer_records_hold_no_tensor(): + tensor = torch.arange(24, dtype=torch.float32).reshape(4, 6)[:, ::2] + tensor_ref = weakref.ref(tensor) + records = _findings(tensor_facts(tensor)) + + first = records[0] + assert first.tensor_facts.shape == (4, 3) + assert first.tensor_facts.strides == (6, 2) + assert first.tensor_facts.dtype == "torch.float32" + assert first.witness == {"pid_0": 3, "arange_0_d0": 5} + assert all(isinstance(value, int) for value in first.witness.values()) + assert isinstance(first.violation_offset, int) + assert isinstance(first.config, dict) + assert isinstance(first.user_code_tracebacks, list) + # A record outlives its launch without keeping the tensor alive. + del tensor + gc.collect() + assert tensor_ref() is None + with pytest.raises(TypeError, match="unhashable"): + hash(first) + + +@pytest.mark.parametrize( + "overrides, error, match", + [ + ({"kind": "oob"}, ValueError, "unknown compiled sanitizer finding"), + ({"op_type": object}, TypeError, "op_type must be Load or Store"), + ({"witness": {"pid_0": 1.5}}, TypeError, "float"), + ({"violation_offset": 2.0}, TypeError, "float"), + ], +) +def test_compiled_sanitizer_records_reject_values_a_trace_cannot_hold( + overrides, error, match +): + facts = tensor_facts(torch.zeros(4)) + fields = { + "kind": "out-of-bounds", + "op_type": Load, + "tensor_name": "x_ptr", + "tensor_facts": facts, + "witness": {}, + "config": {}, + "user_code_tracebacks": [], + **overrides, + } + with pytest.raises(error, match=match): + CompiledSanitizerRecord(**fields) + + +def test_a_saved_trace_holds_compiled_sanitizer_records(tmp_path, monkeypatch): + records = _findings(tensor_facts(torch.arange(24.0).reshape(4, 6)[:, ::2])) + verdict = IRVerdict( + "sanitizer_ir", + "findings", + per_config=[ + ConfigVerdict("hash-a", {"BLOCK": 64}, "findings", n_reports=len(records)) + ], + ) + + launch, manifest = _save_and_load(tmp_path, monkeypatch, [*records, verdict]) + + assert launch.records == [*records, verdict] + loaded = launch.records[:-1] + assert [type(r) for r in loaded] == [CompiledSanitizerRecord] * 3 + assert [r.op_type for r in loaded] == [Load, Store, Load] + for record in loaded: + assert type(record.tensor_facts) is TensorFacts + assert all(type(tb) is TracebackInfo for tb in record.user_code_tracebacks) + assert not any( + isinstance(getattr(record, f.name), torch.Tensor) + for f in dataclasses.fields(record) + ) + # Nothing was stored as a tensor payload. + assert '"kind": "tensor"' not in manifest + + +def test_trace_io_registers_tensor_facts_in_its_own_right(): + """TensorFacts is registered from tilelens.ir.launch, not only as a name + tilelens.clients.sanitizer.data happens to import.""" + code = ( + "import tilelens.clients.sanitizer.data as data\n" + "del data.TensorFacts\n" + "from tilelens.core import trace_io\n" + "from tilelens.ir.launch import TensorFacts\n" + "assert trace_io._TRACE_CLASSES['tilelens.ir.launch:TensorFacts'] is TensorFacts\n" + ) + subprocess.run([sys.executable, "-c", code], cwd=REPO, check=True) diff --git a/tests/unit/sanitizer_compiled/__init__.py b/tests/unit/sanitizer_compiled/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/sanitizer_compiled/test_client.py b/tests/unit/sanitizer_compiled/test_client.py new file mode 100644 index 000000000..ada321411 --- /dev/null +++ b/tests/unit/sanitizer_compiled/test_client.py @@ -0,0 +1,1061 @@ +"""tilelens.clients.sanitizer.compiled.client: the CompiledSanitizer, its +factory and its reports, on fake launches. + +CPU only: fake compiled kernels hold golden TTIR (tests/golden/ir/) or small +TTIR texts whose locs point into a kernel source the test writes, and the +launch events are the core's own, built from CPU tensors. The real-kernel +counterparts (host-compiled, CPU too) live in +tests/end_to_end/test_compiled_sanitizer.py. +""" + +from __future__ import annotations + +import ast +import importlib +import inspect +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +import torch +import triton.language as tl + +import tilelens +import tilelens.ir +from tilelens.clients import CompiledSanitizer as ExportedCompiledSanitizer +from tilelens.clients import CompiledSanitizerRecord as ExportedRecord +from tilelens.clients import Sanitizer +from tilelens.clients.sanitizer.compiled import CompiledSanitizer, SanitizerKind +from tilelens.clients.sanitizer.data import CompiledSanitizerRecord +from tilelens.clients.sanitizer.sanitizer import NullSanitizer, SymbolicSanitizer +from tilelens.core.client import ClientManager, LaunchCall +from tilelens.core.config import config as cfg +from tilelens.core.data import Load, Store +from tilelens.ir import ConfigVerdict, IRClient, IRVerdict, ParseCache +from tilelens.ir.capture import CompiledArtifacts, CompiledSpecialization +from tilelens.ir.launch import tensor_facts +from tilelens.ir.ttir_reader import TTIRKind, UnsupportedTTIR, parse_ttir +from tilelens.ir.verdict import SourceLocation + +client_module = importlib.import_module("tilelens.clients.sanitizer.compiled.client") +trace_module = importlib.import_module("tilelens.core.trace") + +GOLDEN = Path(__file__).resolve().parents[2] / "golden" / "ir" +ADD_TTIR = (GOLDEN / "ttir" / "golden_add_sm80.ttir").read_text(encoding="utf-8") +GATHER_TTIR = (GOLDEN / "ttir" / "golden_gather_sm80.ttir").read_text(encoding="utf-8") + + +# ======== fakes ========= + + +class _FakeJit: + """What the core reads of a JITFunction to bind a launch.""" + + def __init__(self, fn, constexprs=()): + self.signature = inspect.signature(fn) + self.params = [ + SimpleNamespace(name=name, is_constexpr=name in constexprs) + for name in self.signature.parameters + ] + + +def _add_kernel(x_ptr, y_ptr, out_ptr, n_elements, BLOCK_SIZE): + pass + + +def _gather_kernel(idx_ptr, src_ptr, out_ptr, n_elements, BLOCK_SIZE): + pass + + +def _div_kernel(p_ptr, d): + pass + + +ADD = _FakeJit(_add_kernel, {"BLOCK_SIZE"}) +GATHER = _FakeJit(_gather_kernel, {"BLOCK_SIZE"}) +DIV = _FakeJit(_div_kernel) + + +class _FakeKernel: + """A CompiledKernel as the artifact log reads it.""" + + def __init__(self, key, asm): + self.hash = f"hash-{key}" + self.asm = asm + self.metadata = SimpleNamespace( + target=SimpleNamespace(backend="cuda", arch=89), + num_warps=4, + num_stages=3, + shared=0, + name=f"kernel_{key}", + ) + + +def _compiled(jit, args, kwargs=None, *, ttir, key="a", grid=(1,)): + """A before_launch event: ``jit`` compiled (to ``ttir``) for this call.""" + asm = {} if ttir is None else {"ttir": ttir} + return ClientManager._launch_event( + jit, args, dict(kwargs or {}), grid, _FakeKernel(key, asm), False + ) + + +def _failed(jit, args, kwargs, error, *, target=None): + """A compile_failed event.""" + return ClientManager._launch_event( + jit, args, dict(kwargs), (1,), None, False, error=error, target=target + ) + + +def _call(jit, *, capture=True): + return LaunchCall(jit_fn=jit, args=(), kwargs={}, grid=None, capture=capture) + + +def _launch(san, jit=ADD, *, compiled=(), failures=(), capture=True, peers=()): + """One traced launch through a ClientManager, as the core drives it.""" + manager = ClientManager([san, *peers]) + manager.begin_launch(_call(jit, capture=capture)) + for event in compiled: + manager._dispatch_ir("before_launch", event, manager.ir_clients()) + for event in failures: + manager._dispatch_ir("compile_failed", event, manager.ir_clients()) + manager.finalize() + return manager.launch + + +def _add_args(numel=4096, n=4096): + x, y, out = (torch.zeros(numel) for _ in range(3)) + return (x, y, out, n) + + +def _add_event(numel=4096, n=4096, *, blocks=4, key="a", kwargs=None): + """golden add (BLOCK_SIZE 1024 folded, masked by n) over ``blocks`` + programs.""" + kwargs = {"BLOCK_SIZE": 1024} if kwargs is None else kwargs + args = _add_args(numel, n) + return _compiled(ADD, args, kwargs, ttir=ADD_TTIR, key=key, grid=(blocks,)) + + +def _oob_event(**kwargs): + """golden add with every lane of 5 * 1024 active over 4096 elements.""" + return _add_event(n=10**6, blocks=5, **kwargs) + + +def _gather_event(key="g"): + args = (torch.zeros(64, dtype=torch.int32), torch.zeros(64), torch.zeros(64), 64) + return _compiled(GATHER, args, {"BLOCK_SIZE": 1024}, ttir=GATHER_TTIR, key=key) + + +# The divide kernel: a source file whose lines the TTIR locs point at. +DIV_SOURCE = """\ +@triton.jit +def div_kernel(p_ptr, d): + pid = tl.program_id(0) + q = pid // d + tl.load(p_ptr + q) +""" + + +def _div_ttir(path: Path) -> str: + path.write_text(DIV_SOURCE, encoding="utf-8") + f = str(path) + return f"""module {{ + tt.func public @div_kernel(%p: !tt.ptr loc("p_ptr"("{f}":2:0)), %d: i32 loc("d"("{f}":2:0))) attributes {{noinline = false}} {{ + %pid = tt.get_program_id x : i32 loc("{f}":3:10) + %q = arith.divsi %pid, %d : i32 loc("{f}":4:8) + %a = tt.addptr %p, %q : !tt.ptr, i32 loc("{f}":5:12) + %v = tt.load %a : !tt.ptr loc("{f}":5:4) + tt.return loc("{f}":5:4) + }} loc("{f}":2:0) +}} loc("{f}":2:0) +""" + + +def _div_event(ttir, d, *, numel=4, blocks=4): + return _compiled( + DIV, (torch.zeros(numel, dtype=torch.int32), d), ttir=ttir, grid=(blocks,) + ) + + +class _Peer(IRClient): + """Another IR client in the same trace.""" + + NAME = "peer" + IR_STAGES = frozenset() + + def __init__(self): + super().__init__() + self.finalized = 0 + + def analyze_launch(self, log): + self.finalized += 1 + return [], IRVerdict(self.NAME, "ok") + + def on_analysis_error(self, exc): + raise AssertionError(exc) + + def on_refusal(self, refusal): + return IRVerdict(self.NAME, "unsupported", refusal=refusal) + + +class _RunIR(_Peer): + NAME = "run_ir" + LAUNCH = "run" + + +@pytest.fixture +def _isolate_sanitizer_cfg(): + saved = cfg.enable_sanitizer + yield + cfg.enable_sanitizer = saved + + +@pytest.fixture +def quiet(): + return CompiledSanitizer(abort_on_error=False) + + +# ======== the factory and the declarations ========= + + +def test_factory_dispatches_on_compile(_isolate_sanitizer_cfg): + cfg.enable_sanitizer = True + compiled = Sanitizer(compile=True) + assert type(compiled) is CompiledSanitizer + # A virtual subclass: a sanitizer mode, not an eager one. + assert isinstance(compiled, Sanitizer) + assert not issubclass(CompiledSanitizer, SymbolicSanitizer) + assert (compiled.abort_on_error, compiled.timeout_ms) == (True, 10_000) + tuned = Sanitizer(compile=True, abort_on_error=False, timeout_ms=50) + assert (tuned.abort_on_error, tuned.timeout_ms) == (False, 50) + + for eager in (Sanitizer(), Sanitizer(compile=False, abort_on_error=False)): + assert type(eager) is SymbolicSanitizer + assert Sanitizer(compile=False, abort_on_error=False).abort_on_error is False + # ``compile`` is a keyword, and the eager class is never the compiled one. + with pytest.raises(TypeError): + Sanitizer(True, True) + with pytest.raises(TypeError, match="Sanitizer\\(compile=True\\)"): + SymbolicSanitizer(compile=True) + for timeout_ms in (0, -1, 1.5, True): + with pytest.raises(ValueError, match="positive int"): + CompiledSanitizer(timeout_ms=timeout_ms) + + +def test_factory_initializes_the_eager_sanitizer_once( + _isolate_sanitizer_cfg, monkeypatch +): + cfg.enable_sanitizer = True + calls = [] + original = SymbolicSanitizer.__init__ + + def counting(self, *args, **kwargs): + calls.append(kwargs) + original(self, *args, **kwargs) + + monkeypatch.setattr(SymbolicSanitizer, "__init__", counting) + Sanitizer(compile=False, abort_on_error=False) + assert calls == [{"compile": False, "abort_on_error": False}] + + +def test_the_disable_flag_wins_over_compile(_isolate_sanitizer_cfg): + cfg.enable_sanitizer = False + off = Sanitizer(compile=True, abort_on_error=False) + assert type(off) is NullSanitizer + # trace() leaves a kernel traced with it untraced. + kernel = MagicMock() + assert tilelens.trace(off)(kernel) is kernel + # So it does with a CompiledSanitizer built directly, like an explicit + # SymbolicSanitizer(): the kernel then runs, and is not silently skipped. + assert tilelens.trace(CompiledSanitizer(abort_on_error=False))(kernel) is kernel + assert tilelens.trace(SymbolicSanitizer(abort_on_error=False))(kernel) is kernel + + +def test_declarations_and_composition(): + san = CompiledSanitizer() + assert san.NAME == "compiled_sanitizer" != SymbolicSanitizer.NAME + assert (san.IR_STAGES, san.LAUNCH, san.NEEDS_INTERPRETER) == ( + frozenset({"ttir"}), + "skip", + False, + ) + assert ExportedCompiledSanitizer is CompiledSanitizer + assert ExportedRecord is CompiledSanitizerRecord + assert tilelens.ir.SourceLocation is SourceLocation + # The eager sanitizer can share its trace (D4b); a client that needs the + # real launch cannot (D4a). + manager = ClientManager([san, SymbolicSanitizer(abort_on_error=False)]) + assert [c.NAME for c in manager.clients] == ["compiled_sanitizer", "sanitizer"] + with pytest.raises(RuntimeError, match="disagree on whether the real kernel"): + ClientManager([san, _RunIR()]) + + +def test_the_status_is_none_until_a_launch_is_finalized(quiet): + assert (quiet.last_status, quiet.last_verdict, quiet.records) == (None, None, []) + _launch(quiet, compiled=[_oob_event()]) + assert quiet.last_status == "violations" + manager = ClientManager([quiet]) + manager.begin_launch(_call(ADD)) + assert (quiet.last_status, quiet.last_verdict, quiet.records) == (None, None, []) + + +# ======== proofs and findings ========= + + +def test_an_in_bounds_launch_is_ok(capsys): + san = CompiledSanitizer() # abort_on_error: nothing to report + launch = _launch(san, compiled=[_add_event()]) + + verdict = san.last_verdict + assert san.last_status == "ok" and san.records == [] + assert launch.records == [verdict] + assert verdict == IRVerdict( + "compiled_sanitizer", + "ok", + # the proof holds for this launch's arguments, grid and tensors + scope="launch", + per_config=(ConfigVerdict("hash-a", {"BLOCK_SIZE": 1024}, "ok"),), + ) + assert capsys.readouterr().out == "" + + +def test_findings_become_records_after_which_the_verdict_follows(quiet, capsys): + event = _oob_event() + launch = _launch(quiet, compiled=[event]) + + records = quiet.records + assert launch.records == [*records, quiet.last_verdict] + assert [(r.kind, r.op_type, r.tensor_name) for r in records] == [ + ("out-of-bounds", Load, "x_ptr"), + ("out-of-bounds", Load, "y_ptr"), + ("out-of-bounds", Store, "out_ptr"), + ] + accesses = parse_ttir(ADD_TTIR).accesses + for record, access, tensor in zip(records, accesses, event.bound_args.values()): + assert record.tensor_facts == tensor_facts(tensor) + assert record.config == {"BLOCK_SIZE": 1024} + offset = record.violation_offset + assert 4096 <= offset < 5120 + assert record.violation_address == tensor.data_ptr() + offset * 4 + assert record.witness["pid_0"] == 4 + (lane,) = [v for k, v in record.witness.items() if k.startswith("arange_")] + assert 4 * 1024 + lane == offset + (tb,) = record.user_code_tracebacks + assert (tb.filename, tb.lineno, tb.func_name) == ( + access.loc.file, + access.loc.line, + "add_kernel", + ) + assert f"{offset}" in record.detail + (config,) = quiet.last_verdict.per_config + assert (config.status, config.n_reports, config.refusal) == ("violations", 3, None) + # Printed only with abort_on_error or TILELENS_VERBOSE. + assert capsys.readouterr().out == "" + + +def test_a_division_by_zero_record_points_at_the_division(quiet, tmp_path): + ttir = _div_ttir(tmp_path / "k.py") + _launch(quiet, DIV, compiled=[_div_event(ttir, 0)]) + + (record,) = quiet.records + assert (record.kind, record.op_type, record.tensor_name) == ( + "division-by-zero", + Load, + "p_ptr", + ) + assert (record.violation_offset, record.violation_address) == (None, None) + assert set(record.witness) == {"pid_0", "pid_1", "pid_2"} + (tb,) = record.user_code_tracebacks + assert (tb.lineno, tb.func_name, tb.line_of_code.strip()) == ( + 4, + "div_kernel", + "q = pid // d", + ) + assert "divisor" in record.detail + + # d = 1: no division by zero, but pid 1..3 read past the one element. + _launch(quiet, DIV, compiled=[_div_event(ttir, 1, numel=1)]) + (record,) = quiet.records + assert (record.kind, record.violation_offset) == ("out-of-bounds", 1) + assert record.user_code_tracebacks[0].line_of_code.strip() == "tl.load(p_ptr + q)" + + +def test_a_finding_without_a_source_location_names_its_ttir_line(quiet): + text = ( + "module {\n tt.func public @k(%p: !tt.ptr, %d: i32) " + "attributes {noinline = false} {\n" + " %pid = tt.get_program_id x : i32\n" + " %q = arith.divsi %pid, %d : i32\n" + " %a = tt.addptr %p, %q : !tt.ptr, i32\n" + " %v = tt.load %a : !tt.ptr\n" + " tt.return\n }\n}\n" + ) + nameless = _FakeJit(lambda arg0, arg1: None) + event = _compiled(nameless, (torch.zeros(4, dtype=torch.int32), 0), ttir=text) + _launch(quiet, nameless, compiled=[event]) + + (record,) = quiet.records + assert record.user_code_tracebacks == [] + assert record.detail.endswith("(TTIR line 4)") + + +def test_a_loop_finding_on_an_unbound_pointer_has_no_tensor_facts(quiet): + # The loop's bound divides by zero before its store reads anything; the + # store's pointer is bound to no tensor (e.g. a tuple argument's part). + text = ( + "module {\n tt.func public @k(%p: !tt.ptr, %d: i32) " + "attributes {noinline = false} {\n" + " %c0 = arith.constant 0 : i32\n" + " %c1 = arith.constant 1 : i32\n" + " %c8 = arith.constant 8 : i32\n" + " %u = arith.divsi %c8, %d : i32\n" + " scf.for %i = %c0 to %u step %c1 : i32 {\n" + " %a = tt.addptr %p, %i : !tt.ptr, i32\n" + " tt.store %a, %c0 : !tt.ptr\n" + " }\n" + " tt.return\n }\n}\n" + ) + nameless = _FakeJit(lambda arg0, arg1: None) + _launch(quiet, nameless, compiled=[_compiled(nameless, (None, 0), ttir=text)]) + + (record,) = quiet.records + assert (record.kind, record.op_type, record.tensor_name) == ( + "division-by-zero", + Store, + "arg0", + ) + assert record.tensor_facts is None + # The store itself could not be checked. + assert quiet.last_verdict.refusal.kind == "missing-binding" + assert quiet.last_status == "violations" + + +# ======== configs, refusals and the union (D3) ========= + + +def test_the_status_is_the_union_over_configs(quiet): + _launch( + quiet, + compiled=[ + _add_event(key="a", kwargs={"BLOCK_SIZE": 1024, "num_warps": 4}), + _oob_event(key="b", kwargs={"BLOCK_SIZE": 1024, "num_warps": 8}), + ], + ) + verdict = quiet.last_verdict + assert verdict.status == "violations" and verdict.refusal is None + assert [(c.specialization, c.status, c.n_reports) for c in verdict.per_config] == [ + ("hash-a", "ok", 0), + ("hash-b", "violations", 3), + ] + assert {r.config["num_warps"] for r in quiet.records} == {8} + + # An unsupported config makes an otherwise clean launch unsupported, and + # its refusal is the verdict's; a finding elsewhere still wins. + _launch(quiet, compiled=[_add_event(), _gather_event()]) + assert quiet.last_status == "unsupported" + assert quiet.last_verdict.refusal == quiet.last_verdict.per_config[1].refusal + _launch(quiet, compiled=[_gather_event(), _oob_event()]) + verdict = quiet.last_verdict + assert [c.status for c in verdict.per_config] == ["unsupported", "violations"] + assert verdict.status == "violations" + # Not all there is: part of the launch was not checked. + assert verdict.refusal.kind == "indirect-address" + + +def test_configs_sharing_one_kernel_are_checked_and_reported_apart(quiet): + """D22: two configs that compile to one kernel ("hash-a") but set a + runtime int kwarg differently each get their own verdict, checked + against their own binding; findings carry their config.""" + x, y, out, _ = _add_args() + + def config_event(n, blocks): + kwargs = {"BLOCK_SIZE": 1024, "n_elements": n} + return _compiled(ADD, (x, y, out), kwargs, ttir=ADD_TTIR, grid=(blocks,)) + + _launch(quiet, compiled=[config_event(4096, 4), config_event(10**6, 5)]) + + verdict = quiet.last_verdict + assert verdict.status == "violations" + assert [ + (c.specialization, c.config["n_elements"], c.status, c.n_reports) + for c in verdict.per_config + ] == [("hash-a", 4096, "ok", 0), ("hash-a", 10**6, "violations", 3)] + assert {r.config["n_elements"] for r in quiet.records} == {10**6} + + # A kernel the reader refuses leaves every config sharing it unchecked. + args = (torch.zeros(64, dtype=torch.int32), torch.zeros(64), torch.zeros(64)) + events = [ + _compiled(GATHER, args, {"BLOCK_SIZE": 1024, "n_elements": n}, ttir=GATHER_TTIR) + for n in (32, 64) + ] + _launch(quiet, GATHER, compiled=events) + first, second = quiet.last_verdict.per_config + assert [c.config["n_elements"] for c in (first, second)] == [32, 64] + assert first.status == second.status == "unsupported" + assert first.refusal == second.refusal and first.refusal.kind == "indirect-address" + + +def test_configs_are_told_apart_by_their_values_not_their_reprs(quiet): + """Bindings of one kernel group into configs by value: a tensor config + kwarg (e.g. a heuristic's view) by its data_ptr, shape, strides and + dtype, and reported by those facts, never its data; other values that + print alike stay apart unless they are equal.""" + + def event(value, *, oob): + args = _add_args(n=10**6 if oob else 4096) + kwargs = {"BLOCK_SIZE": 1024, "V": value} + return _compiled(ADD, args, kwargs, ttir=ADD_TTIR, grid=(5 if oob else 4,)) + + def per_config(first, second): + _launch(quiet, compiled=[event(first, oob=False), event(second, oob=True)]) + return [(c.status, c.n_reports) for c in quiet.last_verdict.per_config] + + apart = [("ok", 0), ("violations", 3)] + together = [("violations", 3)] + view = torch.zeros(4) + assert per_config(view, torch.zeros(4)) == apart # one repr, two tensors + assert per_config(_Opaque(), _Opaque()) == apart # one repr, two objects + assert per_config(view, view) == together + assert per_config(tl.dtype("fp32"), tl.dtype("fp32")) == together # equal + + per_config(view, view) + (config,) = quiet.last_verdict.per_config + assert config.config["V"] == ( + f"" + ) + assert {r.config["V"] for r in quiet.records} == {config.config["V"]} + + +def test_a_kernel_delivered_without_a_binding_is_never_ok(quiet): + """Nothing was checked, so nothing is "ok": ArtifactLog delivers every + kernel with a binding, but another producer might not.""" + artifacts = CompiledArtifacts(stages={"ttir": ADD_TTIR}, meta={"config": {}}) + spec = CompiledSpecialization("hash-x", artifacts, bindings=()) + ((records, verdict),) = quiet._check_specialization(spec) + assert records == [] and verdict.n_reports == 0 + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "internal-error") + + +def test_a_reader_refusal_keeps_its_kind_and_loc_on_cache_hits(quiet): + texts = [] + + def reader(text): + texts.append(text) + return parse_ttir(text) + + quiet.parses = ParseCache(reader, refusal=UnsupportedTTIR) + _launch(quiet, GATHER, compiled=[_gather_event()]) + first = quiet.last_verdict + _launch(quiet, GATHER, compiled=[_gather_event()]) + + assert texts == [GATHER_TTIR] # read once, across launches + assert quiet.last_verdict == first + assert (first.status, quiet.records) == ("unsupported", []) + refusal = first.refusal + assert refusal == first.per_config[0].refusal + # The reader's enum kind, held as its plain string. + assert refusal.kind == "indirect-address" and not isinstance(refusal.kind, TTIRKind) + assert isinstance(refusal.loc, SourceLocation) and refusal.line_no is not None + assert refusal.message.startswith(f"{refusal.loc.file}:{refusal.loc.line}: ") + + +def test_an_abstention_is_unsupported_with_the_sanitizers_kind(quiet): + # A float n_elements has no binding: the mask reading it is unknown. + x, y, out, _ = _add_args() + event = _compiled(ADD, (x, y, out, 4096.0), {"BLOCK_SIZE": 1024}, ttir=ADD_TTIR) + _launch(quiet, compiled=[event]) + + verdict = quiet.last_verdict + assert verdict.status == "unsupported" and quiet.records == [] + assert verdict.refusal.kind == SanitizerKind.MISSING_BINDING == "missing-binding" + assert "n_elements" in verdict.refusal.message + + +def _static_assert_failure(): + from triton.compiler.errors import CompileTimeAssertionFailure + + return CompileTimeAssertionFailure(None, ast.Pass(), "BLOCK_SIZE <= 1024") + + +def _wrapped(cause): + """``cause`` as Triton's code generator re-raises an error of a called + @jit helper or builtin: a CompilationError raised from it.""" + from triton.compiler.errors import CompilationError + + wrapped = CompilationError(None, ast.Pass(), None) + wrapped.__cause__ = cause + return wrapped + + +def test_a_compile_failure_the_target_may_explain_is_unsupported(quiet): + """D27: a config that failed to compile for the IR target may compile for + (and launch on) the user's GPU, unchecked: unsupported, kind + compile-failed, naming the target and how to name another; never "ok", + never a note.""" + from triton.backends.compiler import GPUTarget + + args = _add_args() + cuda89 = GPUTarget("cuda", 89, 32) + failures = [ + _failed( + ADD, + args, + {"BLOCK_SIZE": 2048}, + ValueError("num_ctas > 1 requires NVIDIA SM90+ (Hopper)"), + target=cuda89, + ), + # An fp8 type the target lacks: the code generator wraps it. + _failed( + ADD, + args, + {"BLOCK_SIZE": 4096}, + _wrapped(ValueError("type fp8e4nv not supported in this architecture")), + target=cuda89, + ), + # An event built outside the core names no target. + _failed(ADD, args, {"BLOCK_SIZE": 512}, RuntimeError("bad option")), + ] + _launch(quiet, compiled=[_add_event()], failures=failures) + + verdict = quiet.last_verdict + assert (verdict.status, verdict.scope, verdict.notes) == ("unsupported", None, ()) + ok, *failed = verdict.per_config + assert ok.status == "ok" + assert [(c.specialization, c.config, c.status, c.refusal.kind) for c in failed] == [ + (None, {"BLOCK_SIZE": 2048}, "unsupported", "compile-failed"), + (None, {"BLOCK_SIZE": 4096}, "unsupported", "compile-failed"), + (None, {"BLOCK_SIZE": 512}, "unsupported", "compile-failed"), + ] + assert verdict.refusal == failed[0].refusal + assert failed[0].refusal.message == ( + "it failed to compile for cuda:89 (ValueError: num_ctas > 1 requires " + "NVIDIA SM90+ (Hopper)), so it was not checked; a kernel can compile for " + "one target and fail for another, so it may launch on a GPU of another " + "kind: to check it, name a target it compiles for " + "(Sanitizer(compile=True, target=...), or TILELENS_IR_TARGET)" + ) + # The innermost error, not the wrapper's source excerpt. + assert failed[1].refusal.message.startswith( + "it failed to compile for cuda:89 (ValueError: type fp8e4nv not supported " + "in this architecture), so it was not checked;" + ) + assert failed[2].refusal.message.startswith( + "it failed to compile (RuntimeError: bad option), so it was not checked;" + ) + + +def test_a_failure_no_target_compiles_past_is_a_note(quiet): + """A failing tl.static_assert (also in a called helper) or a construct + Triton never compiles fails for every target: the config never launches + anywhere, so it is only noted, and the launch can be "ok". Not so once + the compile had asked for its target (e.g. tl.target_info), or for any + other error.""" + from triton.backends.compiler import GPUTarget + from triton.compiler.errors import UnsupportedLanguageConstruct + + from tilelens.core import host_compile + + cuda89 = GPUTarget("cuda", 89, 32) + args = _add_args() + nowhere = [ + _static_assert_failure(), + _wrapped(_static_assert_failure()), + UnsupportedLanguageConstruct(None, ast.Pass(), "nested function"), + ] + failures = [ + _failed(ADD, args, {"BLOCK_SIZE": 2048 * (i + 1)}, error, target=cuda89) + for i, error in enumerate(nowhere) + ] + _launch(quiet, compiled=[_add_event()], failures=failures) + + verdict = quiet.last_verdict + assert verdict.status == "ok" and len(verdict.per_config) == 1 + assert len(verdict.notes) == 3 + assert verdict.notes[0] == ( + "config {'BLOCK_SIZE': 2048} was not checked: it failed to compile for " + "cuda:89 (CompileTimeAssertionFailure: BLOCK_SIZE <= 1024), an error of " + "its own code whatever the target, so it never launches" + ) + # Through the helper's wrapper: the assertion's own message. + assert "(CompileTimeAssertionFailure: BLOCK_SIZE <= 1024)" in verdict.notes[1] + assert "(UnsupportedLanguageConstruct: nested function)" in verdict.notes[2] + + # The same errors after a target query, or wrapped around another error, + # may be the target's. + queried = _static_assert_failure() + host_compile._mark_target_queried(queried) + maybe = [queried, _wrapped(ValueError("x")), None] + failures = [ + _failed(ADD, args, {"BLOCK_SIZE": 2048 * (i + 1)}, error, target=cuda89) + for i, error in enumerate(maybe) + ] + _launch(quiet, compiled=[_add_event()], failures=failures) + verdict = quiet.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + assert [c.refusal.kind for c in verdict.per_config[1:]] == ["compile-failed"] * 3 + + +def test_a_host_compile_that_could_not_run_is_unsupported_not_a_note(quiet): + """HostCompileUnavailable (Triton's compile API, or a compile that asked + for a device) says nothing about the kernel: the config may launch, so + it is unsupported, never a note that it cannot launch there.""" + from triton.backends.compiler import GPUTarget + from triton.compiler.errors import CompilationError + + from tilelens.core.host_compile import HostCompileUnavailable + + cuda80 = GPUTarget("cuda", 80, 32) + unavailable = HostCompileUnavailable("Triton asked its driver for 'utils'") + # Triton's code generator re-raises what the kernel's code raised. + wrapped = CompilationError("def add(...)", None, repr(unavailable)) + wrapped.__cause__ = unavailable + failures = [ + _failed(ADD, _add_args(), {"BLOCK_SIZE": 2048}, unavailable, target=cuda80), + _failed(ADD, _add_args(), {"BLOCK_SIZE": 4096}, wrapped, target=cuda80), + ] + _launch(quiet, compiled=[_add_event()], failures=failures) + verdict = quiet.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + ok, *refused = verdict.per_config + assert ok.status == "ok" + assert [(c.config, c.status, c.refusal.kind) for c in refused] == [ + ({"BLOCK_SIZE": 2048}, "unsupported", "host-compile-unavailable"), + ({"BLOCK_SIZE": 4096}, "unsupported", "host-compile-unavailable"), + ] + assert all("asked its driver for 'utils'" in c.refusal.message for c in refused) + + # Nothing compiled: the launch's refusal says why. + _launch(quiet, failures=failures[:1]) + verdict = quiet.last_verdict + assert (verdict.status, verdict.refusal.kind) == ( + "unsupported", + SanitizerKind.HOST_COMPILE_UNAVAILABLE, + ) + assert len(verdict.per_config) == 1 and verdict.notes == () + + +def test_a_launch_that_compiled_nothing_is_unsupported(quiet): + # No JITFunction (TRITON_INTERPRET, Gluon, NKI): nothing was captured. + _launch(quiet, capture=False) + verdict = quiet.last_verdict + assert (verdict.status, verdict.refusal.kind) == ( + "unsupported", + "no-compiled-kernel", + ) + assert "TRITON_INTERPRET" in verdict.refusal.message + + # Every config failed to compile for the target (D27; a mixed trace goes + # on to interpret the launch): the first one's refusal. + from triton.backends.compiler import GPUTarget + + cuda89 = GPUTarget("cuda", 89, 32) + failure = _failed( + ADD, _add_args(), {"BLOCK_SIZE": 2048}, RuntimeError("bad"), target=cuda89 + ) + _launch(quiet, failures=[failure]) + verdict = quiet.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "compile-failed") + (config,) = verdict.per_config + assert config.refusal == verdict.refusal and verdict.notes == () + assert "for cuda:89 (RuntimeError: bad)" in verdict.refusal.message + assert "TILELENS_IR_TARGET" in verdict.refusal.message + + # Every config failed with an error no target compiles past: compile-failed + # too, the notes saying why. + failures = [ + _failed(ADD, _add_args(), {"BLOCK_SIZE": b}, _static_assert_failure(), target=t) + for b, t in ((2048, cuda89), (4096, None)) + ] + _launch(quiet, failures=failures) + verdict = quiet.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "compile-failed") + assert verdict.per_config == () and len(verdict.notes) == 2 + assert verdict.refusal.message == ( + "no config of the launch compiled for cuda:89, so nothing was checked: " + "each failed with an error of its own code whatever the target (see the " + "notes)" + ) + + +def test_errors_no_target_decides_name_no_target(quiet): + """A call that does not bind the kernel's parameters (the host compile + marks it, see bind_failed) and a compile refused while the language is + patched (no compile ran) are unsupported, never notes, but their + refusals send nobody to another target: the untraced call raises the + bind error on any GPU, and the refused compile says nothing about the + kernel. (A traced launch raises a bind failure instead, D28; only an + event built outside the core, as here, delivers one.)""" + from triton.backends.compiler import GPUTarget + + from tilelens.core import host_compile + from tilelens.core.client import LanguagePatchedError + + cuda89 = GPUTarget("cuda", 89, 32) + unbound = TypeError("dynamic_func() missing 1 required positional argument: 'n'") + host_compile._mark_bind_failed(unbound) + patched = LanguagePatchedError("a Triton compile cannot run while ...") + failures = [ + _failed(ADD, _add_args(), {"BLOCK_SIZE": b}, error, target=cuda89) + for b, error in ((2048, unbound), (4096, patched)) + ] + _launch(quiet, compiled=[_add_event()], failures=failures) + + verdict = quiet.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + ok, bind, refused = verdict.per_config + assert ok.status == "ok" + assert (bind.status, bind.refusal.kind) == ("unsupported", "compile-failed") + assert bind.refusal.message == ( + "the call does not bind to the kernel's parameters (TypeError: " + "dynamic_func() missing 1 required positional argument: 'n'), so it was " + "not checked; Triton raises this error for the call whatever the GPU" + ) + assert (refused.status, refused.refusal.kind) == ( + "unsupported", + "host-compile-unavailable", + ) + assert refused.refusal.message == "a Triton compile cannot run while ..." + for config in (bind, refused): + assert "TILELENS_IR_TARGET" not in config.refusal.message + + +def test_a_launch_no_config_of_which_compiled_prints_its_notes(capsys): + """The refusal of a launch every config of which failed only as a note + says "see the notes": they are printed with it, each once.""" + from triton.backends.compiler import GPUTarget + + cuda89 = GPUTarget("cuda", 89, 32) + det = CompiledSanitizer() # abort_on_error: prints what was not checked + failures = [ + _failed( + ADD, _add_args(), {"BLOCK_SIZE": b}, _static_assert_failure(), target=cuda89 + ) + for b in (2048, 4096) + ] + for _ in range(2): + _launch(det, failures=failures) + assert det.last_verdict.per_config == () and len(det.last_verdict.notes) == 2 + + lines = capsys.readouterr().out.splitlines() + assert lines == [ + "[CompiledSanitizer] not checked: compile-failed: no config of the launch " + "compiled for cuda:89, so nothing was checked: each failed with an error " + "of its own code whatever the target (see the notes)", + *(f"[CompiledSanitizer] note: {note}" for note in det.last_verdict.notes), + ] + assert "(CompileTimeAssertionFailure: BLOCK_SIZE <= 1024)" in lines[1] + + +def test_a_kernel_without_ttir_is_unsupported(quiet): + _launch(quiet, compiled=[_compiled(ADD, _add_args(), ttir=None)]) + (config,) = quiet.last_verdict.per_config + assert (config.status, config.refusal.kind) == ("unsupported", "no-ttir") + + +def test_analysis_bugs_are_contained_per_config(quiet, monkeypatch): + def reader(text): + if "gather_kernel" in text: + raise KeyError("reader bug") + return parse_ttir(text) + + quiet.parses = ParseCache(reader, refusal=UnsupportedTTIR) + real_check = client_module.check_graph + + def check_graph(graph, binding, **kw): + if binding.grid == (2, 1, 1): + raise ZeroDivisionError("evaluator bug") + return real_check(graph, binding, **kw) + + monkeypatch.setattr(client_module, "check_graph", check_graph) + _launch( + quiet, + compiled=[_gather_event(), _add_event(blocks=2, key="b"), _oob_event()], + ) + reader_bug, evaluator_bug, checked = quiet.last_verdict.per_config + assert (reader_bug.refusal.kind, evaluator_bug.refusal.kind) == ( + "internal-error", + "internal-error", + ) + assert "KeyError" in reader_bug.refusal.message + assert evaluator_bug.refusal.message == "ZeroDivisionError: evaluator bug" + # The other config's findings still count. + assert checked.status == "violations" and len(quiet.records) == 3 + assert quiet.last_status == "violations" + + +def test_a_bug_outside_the_configs_is_an_internal_error(quiet, monkeypatch): + def broken(failure): + raise RuntimeError("note bug") + + monkeypatch.setattr(client_module, "_failure_refusal", broken) + failure = _failed(ADD, _add_args(), {"BLOCK_SIZE": 2048}, RuntimeError("bad")) + launch = _launch(quiet, compiled=[_oob_event()], failures=[failure]) + + verdict = quiet.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "internal-error") + assert verdict.refusal.message == "RuntimeError: note bug" + assert quiet.records == [] and launch.records == [verdict] + + +def test_a_repeated_launch_reuses_its_check(quiet, monkeypatch): + """A training loop: fresh tensors of the same shapes each launch. The + check runs once; each launch's records hold its own tensors' addresses.""" + calls = [] + real = client_module.check_graph + + def counting(graph, binding, **kwargs): + calls.append(binding) + return real(graph, binding, **kwargs) + + monkeypatch.setattr(client_module, "check_graph", counting) + addresses, events = [], [] # the events hold the tensors: fresh addresses + for _ in range(3): + event = _oob_event() + events.append(event) + _launch(quiet, compiled=[event]) + x = event.bound_args["x_ptr"] + (record,) = [r for r in quiet.records if r.tensor_name == "x_ptr"] + assert record.violation_address == x.data_ptr() + record.violation_offset * 4 + assert record.tensor_facts == tensor_facts(x) + addresses.append(record.violation_address) + assert len(calls) == 1 and len(set(addresses)) == 3 + assert quiet.last_status == "violations" + # Another argument (or shape) is another check. + _launch(quiet, compiled=[_add_event()]) + _launch(quiet, compiled=[_add_event(numel=8192, n=8192, blocks=8)]) + assert len(calls) == 3 and quiet.last_status == "ok" + + +def test_an_interrupt_during_the_check_is_not_contained(quiet, monkeypatch): + """Ctrl+C is the user's: never an internal-error verdict.""" + + def interrupted(graph, binding, **kwargs): + raise KeyboardInterrupt + + monkeypatch.setattr(client_module, "check_graph", interrupted) + with pytest.raises(KeyboardInterrupt): + _launch(quiet, compiled=[_oob_event()]) + + +def test_records_belong_to_the_last_launch(quiet): + _launch(quiet, compiled=[_oob_event()]) + assert len(quiet.records) == 3 + _launch(quiet, compiled=[_add_event()]) + assert (quiet.last_status, quiet.records) == ("ok", []) + + +# ======== reporting and abort_on_error ========= + + +def test_abort_on_error_reports_then_exits_once_every_client_finalized(capsys): + san, peer = CompiledSanitizer(), _Peer() + with pytest.raises(SystemExit) as exc_info: + _launch(san, compiled=[_oob_event()], peers=[peer]) + + assert exc_info.value.code == 1 + assert peer.finalized == 1 + assert san.last_status == "violations" and len(san.records) == 3 + out = capsys.readouterr().out + assert out.count("Out-Of-Bounds Access Detected") == 3 + assert "Tensor Arg: x_ptr" in out and "Operation: Store" in out + assert "Witness: pid_0=4" in out and "Config: {'BLOCK_SIZE': 1024}" in out + assert f"{san.records[0].violation_address:#x}" in out + + +def test_abort_on_error_reports_but_never_exits_for_unchecked_parts(capsys): + san = CompiledSanitizer() + _launch(san, GATHER, compiled=[_gather_event()]) + assert san.last_status == "unsupported" + (line,) = capsys.readouterr().out.splitlines() + assert line.startswith( + "[CompiledSanitizer] not checked (config {'BLOCK_SIZE': 1024}): " + "indirect-address: " + ) + # Once per client: a kernel launched in a loop does not repeat it. + _launch(san, GATHER, compiled=[_gather_event()]) + assert san.last_status == "unsupported" + assert capsys.readouterr().out == "" + # A launch-level refusal is printed too. + _launch(san, capture=False) + assert "not checked: no-compiled-kernel: " in capsys.readouterr().out + + +def test_every_config_that_failed_to_compile_is_printed_once(capsys): + """Failed configs have no kernel: each is printed once per config and + message, and the launch never exits for them (D27).""" + san = CompiledSanitizer() + failures = [ + _failed(ADD, _add_args(), {"BLOCK_SIZE": b}, ValueError("not here")) + for b in (2048, 4096) + ] + _launch(san, compiled=[_add_event()], failures=failures) + assert san.last_status == "unsupported" + lines = capsys.readouterr().out.splitlines() + assert [line.split(": compile-failed: ")[0] for line in lines] == [ + "[CompiledSanitizer] not checked (config {'BLOCK_SIZE': 2048})", + "[CompiledSanitizer] not checked (config {'BLOCK_SIZE': 4096})", + ] + _launch(san, compiled=[_add_event()], failures=failures) + assert capsys.readouterr().out == "" + + +def test_an_unchecked_op_prints_once_whatever_its_message(capsys): + """Once per compiled kernel, kind and op: a withheld finding's message + names an element offset that changes with the tensors' sizes.""" + printed: set = set() + loc = SourceLocation("k.py", 7) + + def verdict(message, kind="unmodelable-condition", line_no=3): + refusal = client_module.Refusal(kind, message, line_no, loc) + return IRVerdict( + "compiled_sanitizer", + "unsupported", + refusal=refusal, + per_config=(ConfigVerdict("hash-a", {}, "unsupported", refusal),), + ) + + client_module.print_unchecked(verdict("at element offset 4096"), printed) + client_module.print_unchecked(verdict("at element offset 8192"), printed) + assert len(capsys.readouterr().out.splitlines()) == 1 + client_module.print_unchecked(verdict("offset 1", line_no=4), printed) + client_module.print_unchecked(verdict("offset 1", "data-dependent-mask"), printed) + assert len(capsys.readouterr().out.splitlines()) == 2 + + +def test_verbose_prints_without_aborting(quiet, monkeypatch, capsys, tmp_path): + monkeypatch.setattr(cfg, "verbose", True) + _launch(quiet, DIV, compiled=[_div_event(_div_ttir(tmp_path / "k.py"), 0)]) + out = capsys.readouterr().out + assert "Division By Zero Detected" in out + assert "Code: q = pid // d" in out + assert "Invalid access detected" not in out + + +# ======== persistence ========= + + +class _Opaque: + def __repr__(self): + return "" + + +def test_a_launch_round_trips_through_a_saved_trace(quiet, tmp_path): + kwargs = {"BLOCK_SIZE": 1024, "DTYPE": _Opaque()} + launch = _launch(quiet, compiled=[_oob_event(kwargs=kwargs), _gather_event()]) + # A config value a trace cannot hold is kept as its repr. + assert quiet.records[0].config == {"BLOCK_SIZE": 1024, "DTYPE": ""} + assert quiet.last_verdict.per_config[0].config["DTYPE"] == "" + + saved = list(trace_module.launches) + trace_module.launches[:] = [launch] + try: + path = tilelens.save(tmp_path / "trace.zip") + (loaded,) = tilelens.load(path) + finally: + trace_module.launches[:] = saved + assert loaded.records == launch.records + assert all(isinstance(r, CompiledSanitizerRecord) for r in loaded.records[:-1]) + assert loaded.records[-1].per_config[1].refusal.loc == ( + launch.records[-1].per_config[1].refusal.loc + ) diff --git a/tests/unit/sanitizer_compiled/test_oob.py b/tests/unit/sanitizer_compiled/test_oob.py new file mode 100644 index 000000000..b27bc83c2 --- /dev/null +++ b/tests/unit/sanitizer_compiled/test_oob.py @@ -0,0 +1,1337 @@ +"""tilelens.clients.sanitizer.compiled.oob: the compiled sanitizer's checks. + +CPU only: graphs come from the TTIR reader on the goldens in tests/golden/ir/ +(ttir/ and reader_ttir/) or on small TTIR texts, or are built by hand; +launches are synthetic LaunchBindings, with TensorFacts read from CPU +tensors where a view's layout matters. Ports the #361 cases of +tests/unit/test_compiled_sanitizer_oob.py onto the new reader and binding. +""" + +from __future__ import annotations + +import gc +import pickle +import subprocess +import sys +from pathlib import Path +from types import MappingProxyType + +import pytest +import torch +import z3 + +from tilelens.clients.sanitizer.compiled import oob +from tilelens.clients.sanitizer.compiled.oob import ( + CheckResult, + Finding, + SanitizerKind, + _in_view, + check_graph, + launch_key, + readdressed, +) +from tilelens.clients.symbolic_engine import SymbolicClient +from tilelens.ir.launch import LaunchBinding, TensorFacts, tensor_facts +from tilelens.ir.ttir_reader import ( + AccessEvent, + AccessGraph, + Arange, + AtomicInfo, + Bin, + BoolBin, + Cmp, + Const, + DataDep, + FuncArg, + LoopInfo, + LoopVar, + Observed, + Param, + Pid, + TTIRKind, + parse_ttir, +) +from tilelens.ir.verdict import Refusal, SourceLocation + +K = SanitizerKind +REPO = Path(__file__).resolve().parents[3] +GOLDEN = REPO / "tests" / "golden" / "ir" +ADD = "ttir/golden_add_sm80.ttir" +MATMUL = "ttir/golden_matmul_s3_sm80.ttir" +TILE2D = "ttir/golden_tile2d_sm80.ttir" +PTR = 0x1000 + + +def _text(name: str) -> str: + return (GOLDEN / name).read_text(encoding="utf-8") + + +def _graph(name: str) -> AccessGraph: + return parse_ttir(_text(name)) + + +def _module( + body: str, args: str = "%p: !tt.ptr, %n: i32" +) -> tuple[AccessGraph, str]: + """A minimal TTIR module (no locs, so its parameters read as arg0, + arg1, ...) around ``body``'s op lines, and its text.""" + lines = "\n ".join(line.strip() for line in body.strip().splitlines()) + text = ( + f"module {{\n tt.func public @k({args}) attributes {{noinline = false}} {{\n" + f" {lines}\n tt.return\n }}\n}}\n" + ) + return parse_ttir(text), text + + +def _line(text: str, needle: str) -> int: + (line,) = [i for i, t in enumerate(text.splitlines(), 1) if needle in t] + return line + + +def _facts(numel: int, elem_size: int = 4, data_ptr: int = PTR) -> TensorFacts: + """A contiguous 1D tensor.""" + return TensorFacts( + data_ptr, elem_size, numel, (numel,), (1,), "torch.float32", True + ) + + +def _bind(grid=(1, 1, 1), params=None, **tensors) -> LaunchBinding: + return LaunchBinding( + params=MappingProxyType(dict(params or {})), + tensors=MappingProxyType(tensors), + constexprs=MappingProxyType({}), + raw_grid=grid, + grid=grid, + config=MappingProxyType({}), + ) + + +def _add(grid, n, **tensors) -> CheckResult: + return check_graph(_graph(ADD), _bind(grid, {"n_elements": n}, **tensors)) + + +def _all(numel: int, *names: str, elem_size: int = 4) -> dict[str, TensorFacts]: + return {name: _facts(numel, elem_size) for name in names} + + +ADD_TENSORS = ("x_ptr", "y_ptr", "out_ptr") + + +def _clean(result: CheckResult) -> None: + assert result.findings == () + assert result.refusal is None and result.abstained == () + + +def _found(result: CheckResult) -> list[tuple[str, int]]: + return [(f.kind, f.access_index) for f in result.findings] + + +def _synthetic(*accesses: AccessEvent, loop: LoopInfo | None = None, args=("p",)): + return AccessGraph( + kernel_name="synthetic", + func_args=[FuncArg(a, True, 32) for a in args], + accesses=accesses, + loop=loop, + ) + + +def _access(offset, *, kind="load", base="p", mask=None, line=1, **kw) -> AccessEvent: + return AccessEvent(kind, base, offset, mask, 32, None, line, **kw) + + +# ─────────────────────────── add (1D, masked) ─────────────────────────── + + +def test_add_in_bounds_is_clean(): + _clean(_add((4, 1, 1), 4096, **_all(4096, *ADD_TENSORS))) + + +def test_add_unmasked_tail_is_out_of_bounds(): + """A mask bound (n_elements) past the tensor leaves the last block's + tail unguarded.""" + r = _add((5, 1, 1), 10**9, **_all(4096, *ADD_TENSORS)) + assert r.refusal is None and r.abstained == () + assert _found(r) == [ + ("out-of-bounds", 0), + ("out-of-bounds", 1), + ("out-of-bounds", 2), + ] + text = _text(ADD).splitlines() + for f, kind in zip(r.findings, ("load", "load", "store")): + assert f.access_kind == kind and f.base_param == ADD_TENSORS[f.access_index] + assert f.violation_offset >= 4096 + assert f.violation_address == PTR + f.violation_offset * 4 + # the witness is the state that reaches the offset + lane = next(v for k, v in f.witness.items() if k.startswith("arange_")) + assert f.witness["pid_0"] * 1024 + lane == f.violation_offset + assert f"tt.{kind}" in text[f.line_no - 1] + assert isinstance(f.loc, SourceLocation) and f.loc.line > 0 + assert "outside the tensor's 4096 elements" in f.detail + + +# ─────────────────────────── matmul (loop, 2D) ─────────────────────────── + +_MATMUL_PARAMS = { + "M": 128, + "N": 128, + "K": 128, + "stride_am": 128, + "stride_bk": 128, + "stride_cm": 128, +} + + +def test_matmul_in_bounds_and_oversized_grid(): + tensors = _all(128 * 128, "a_ptr", "b_ptr", "c_ptr", elem_size=2) + _clean(check_graph(_graph(MATMUL), _bind((2, 2, 1), _MATMUL_PARAMS, **tensors))) + # The A load has no row mask (only a K mask): too many row blocks + # (pid_m = 2 with M = 128, BLOCK_M = 64) read rows past M. + r = check_graph(_graph(MATMUL), _bind((3, 2, 1), _MATMUL_PARAMS, **tensors)) + assert ("out-of-bounds", 0) in _found(r) and r.refusal is None + a = r.findings[0] + assert a.base_param == "a_ptr" and a.witness["pid_0"] == 2 + assert 0 <= a.witness["iter_loop"] < 4 # K / BLOCK_K iterations run + assert a.violation_address == PTR + a.violation_offset * 2 + + +def test_modeled_branch_path_constrains_the_witness(): + """Only program 0 stores (``if pid == 0``): the path is modeled, so a + later program's out-of-range offsets are no witness.""" + g = _graph("ttir/golden_pid_branch_sm80.ttir") + b = _bind((4, 1, 1), {"n_elements": 1024}, x_ptr=_facts(1024), out_ptr=_facts(256)) + _clean(check_graph(g, b)) + b = _bind((4, 1, 1), {"n_elements": 1024}, x_ptr=_facts(1024), out_ptr=_facts(255)) + ((f,),) = [check_graph(g, b).findings] + assert f.access_index == 1 and f.witness["pid_0"] == 0 and f.violation_offset == 255 + + +def test_grid_stride_loop_with_a_pid_lower_bound(): + """``for row in range(pid, n_rows, 4)``: a loop bound that depends on + the program id stays symbolic.""" + g = _graph("ttir/golden_grid_stride_sm80.ttir") + params = {"n_rows": 8, "stride": 64} + _clean(check_graph(g, _bind((4, 1, 1), params, **_all(8 * 64, "x_ptr", "out_ptr")))) + r = check_graph( + g, _bind((4, 1, 1), params, x_ptr=_facts(7 * 64), out_ptr=_facts(8 * 64)) + ) + ((f,),) = [r.findings] + assert f.base_param == "x_ptr" and f.violation_offset >= 7 * 64 + # row = pid + 4 * iter reaches row 7 only + assert f.witness["pid_0"] + 4 * f.witness["iter_loop"] == 7 + + +def test_expanded_loop_carried_tiles_keep_their_lanes(): + """q[i, j, l] = x + j*N + l - i*N + k (N = 4, k the iteration): the + reader's expanded iter_arg tile, one make_range on three dims.""" + g = _graph("reader_ttir/expand_iterarg_3d.ttir") + r = check_graph( + g, _bind((1, 1, 1), {"n": 2}, x_ptr=_facts(1000), out_ptr=_facts(64)) + ) + ((f,),) = [r.findings] + assert f.access_index == 0 and f.violation_offset < 0 + lane = {int(k[-1]): v for k, v in f.witness.items() if k.startswith("arange_")} + assert sorted(lane) == [0, 1, 2] + k = f.witness["iter_loop"] + assert 0 <= k < 2 + assert lane[1] * 4 + lane[2] - lane[0] * 4 + k == f.violation_offset + + +def test_select_of_pointers(): + """``tl.where(offs < n, x + offs, x + 100)`` over 16 lanes.""" + g = _graph("reader_ttir/where_pointer.ttir") + _clean(check_graph(g, _bind(params={"n": 8}, x_ptr=_facts(101)))) + _clean(check_graph(g, _bind(params={"n": 16}, x_ptr=_facts(16)))) + ((f,),) = [check_graph(g, _bind(params={"n": 8}, x_ptr=_facts(100))).findings] + assert f.violation_offset == 100 + + +def test_three_lanes_of_one_make_range(): + g = _graph("reader_ttir/tile3d_shared_arange.ttir") + _clean(check_graph(g, _bind(x_ptr=_facts(64)))) + ((f,),) = [check_graph(g, _bind(x_ptr=_facts(63))).findings] + assert f.violation_offset == 63 + + +def test_aranges_along_one_dim_share_the_lane(): + """tl.arange(1, 65) - tl.arange(0, 64) is 1 on every lane: two + make_range sites along one dim index one lane position, so positions 1 + and 0 of the two are no witness (the real kernel reads x[1] only).""" + g, _ = _module( + """ + %r1 = tt.make_range {end = 65 : i32, start = 1 : i32} : tensor<64xi32> + %r0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> + %d = arith.subi %r1, %r0 : tensor<64xi32> + %ps = tt.splat %p : !tt.ptr -> tensor<64x!tt.ptr> + %a = tt.addptr %ps, %d : tensor<64x!tt.ptr>, tensor<64xi32> + %v = tt.load %a : tensor<64x!tt.ptr>""" + ) + _clean(check_graph(g, _bind(arg0=_facts(2)))) + ((f,),) = [check_graph(g, _bind(arg0=_facts(1))).findings] + assert f.violation_offset == 1 + # the witness names each arange by its range, at the one lane + assert f.witness["arange_1_65"] == f.witness["arange_0_64"] + 1 + + +def test_a_broadcast_extent_one_arange_keeps_its_position(): + """tl.arange(5, 6) broadcast to 64 lanes plus tl.arange(0, 64): the + extent-1 range stays at its one position, every lane of the other + range stays free.""" + g, _ = _module( + """ + %r5 = tt.make_range {end = 6 : i32, start = 5 : i32} : tensor<1xi32> + %b = tt.broadcast %r5 : tensor<1xi32> -> tensor<64xi32> + %r0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> + %o = arith.addi %b, %r0 : tensor<64xi32> + %ps = tt.splat %p : !tt.ptr -> tensor<64x!tt.ptr> + %a = tt.addptr %ps, %o : tensor<64x!tt.ptr>, tensor<64xi32> + %v = tt.load %a : tensor<64x!tt.ptr>""" + ) + _clean(check_graph(g, _bind(arg0=_facts(69)))) + ((f,),) = [check_graph(g, _bind(arg0=_facts(68))).findings] + assert f.violation_offset == 68 + assert (f.witness["arange_5_6"], f.witness["arange_0_64"]) == (5, 63) + + +# ─────────────────── synthetic graphs (ported from #361) ─────────────────── + + +def test_reused_arange_rows_and_cols_are_independent(): + """One make_range on the row and the column dim is two variables: with + offset = row - col, collapsing them makes the offset 0 (never OOB).""" + r = "%shared_range" + offset = Bin("-", Arange(r, 0, 4, dim=0), Arange(r, 0, 4, dim=1)) + ((f,),) = [check_graph(_synthetic(_access(offset)), _bind(p=_facts(64))).findings] + assert f.violation_offset < 0 + + +def _store_loop(lower, step, upper, offset): + return _synthetic( + _access(offset, kind="store", base="out", in_loop=True), + loop=LoopInfo("%loop", "%k", lower=lower, upper=upper, step=step), + args=("out",), + ) + + +def test_loop_nonzero_lower_no_false_positive(): + """for k in range(1, n): store(out + (k - 1)) writes 0..n-2: the + iteration model never runs k = 0.""" + g = _store_loop( + Const(1), Const(1), Param("n"), Bin("-", LoopVar("%loop"), Const(1)) + ) + _clean(check_graph(g, _bind(params={"n": 8}, out=_facts(7)))) + + +def test_loop_step_skips_unrun_iterations(): + """for k in range(0, n, 2): store(out + k) writes only even offsets.""" + g = _store_loop(Const(0), Const(2), Param("n"), LoopVar("%loop")) + _clean(check_graph(g, _bind(params={"n": 8}, out=_facts(7)))) + ((f,),) = [check_graph(g, _bind(params={"n": 8}, out=_facts(6))).findings] + assert f.violation_offset == 6 and f.witness["iter_loop"] == 3 + + +def test_descending_loop_refuses_non_positive_step(): + g = _store_loop(Const(10), Const(-1), Const(0), LoopVar("%loop")) + r = check_graph(g, _bind(out=_facts(16))) + assert r.findings == () and r.abstained == ((0, K.NON_POSITIVE_STEP),) + assert r.refusal.kind == "non-positive-step" and "step is -1" in r.refusal.message + + +def test_step_that_depends_on_the_program_id(): + lower, upper = Const(0), Const(8) + ok = _store_loop(lower, Bin("+", Pid(0), Const(1)), upper, LoopVar("%loop")) + _clean(check_graph(ok, _bind((4, 1, 1), out=_facts(8)))) + bad = _store_loop(lower, Bin("-", Pid(0), Const(1)), upper, LoopVar("%loop")) + r = check_graph(bad, _bind((4, 1, 1), out=_facts(8))) + assert r.abstained == ((0, K.NON_POSITIVE_STEP),) + assert "step can be" in r.refusal.message + + +def _flat_load(offset, numel, grid_x): + return check_graph( + _synthetic(_access(offset)), _bind((grid_x, 1, 1), p=_facts(numel)) + ) + + +@pytest.mark.parametrize( + "offset, numel, grid_x, oob", + [ + # pid % 8 stays in [0, 8): the grouped-swizzle `pid % group_size_m` + (Bin("%", Pid(0), Const(8)), 8, 64, None), + (Bin("%", Pid(0), Const(8)), 7, 64, 7), + # min(pid, 9) never leaves [0, 10) + (Bin("min", Pid(0), Const(9)), 10, 1000, None), + (Bin("min", Pid(0), Const(9)), 9, 1000, 9), + # max(pid, 5) pins the floor at 5 + (Bin("max", Pid(0), Const(5)), 6, 4, None), + (Bin("max", Pid(0), Const(5)), 5, 4, 5), + # divsi truncates: (0 - 1) // 2 == 0, not Euclidean -1 + (Bin("//", Bin("-", Const(0), Pid(0)), Const(2)), 4, 2, None), + # remsi keeps the dividend's sign: (0 - 1) % 2 == -1 + (Bin("%", Bin("-", Const(0), Pid(0)), Const(2)), 2, 2, -1), + ], +) +def test_integer_semantics(offset, numel, grid_x, oob): + r = _flat_load(offset, numel, grid_x) + assert r.refusal is None + assert [f.violation_offset for f in r.findings] == ([] if oob is None else [oob]) + + +# ─────────────────── uncertainty: guarded, dropped, observed ─────────────────── + + +def test_guarded_access_unsat_is_still_a_proof(): + g = _synthetic(_access(Pid(0), guarded=True)) + _clean(check_graph(g, _bind((4, 1, 1), p=_facts(4)))) + + +def test_guarded_access_sat_abstains(): + """SAT under a branch condition the reader could not model may be a + branch the launch never takes: abstain, never a witness.""" + g = _synthetic(_access(Bin("-", Pid(0), Const(1)), guarded=True)) + r = check_graph(g, _bind((4, 1, 1), p=_facts(4))) + assert r.findings == () and r.abstained == ((0, K.UNMODELABLE_CONDITION),) + assert r.refusal.kind == "unmodelable-condition" + assert "possible out-of-bounds" in r.refusal.message + assert r.refusal.message.startswith("TTIR line 1: ") + + +def test_exact_finding_kept_alongside_guarded_abstention(): + g = _synthetic( + _access(Bin("-", Pid(0), Const(1)), guarded=True), + _access(Bin("+", Pid(0), Const(100)), line=2), + ) + r = check_graph(g, _bind((4, 1, 1), p=_facts(4))) + assert _found(r) == [("out-of-bounds", 1)] and r.findings[0].line_no == 2 + assert r.findings[0].violation_offset >= 100 + assert r.abstained == ((0, K.UNMODELABLE_CONDITION),) + + +def test_mask_dropped_abstains_and_exact_findings_stay(): + """atomic_fmax's atomics are masked by loaded data (dropped as free).""" + g = _graph("ttir/golden_atomic_fmax_sm80.ttir") + assert [a.mask_dropped for a in g.accesses] == [False, True, True] + r = check_graph( + g, + _bind((4, 1, 1), {"n_elements": 1024}, x_ptr=_facts(2048), out_ptr=_facts(10)), + ) + assert r.findings == () + assert r.abstained == ((1, K.DATA_DEPENDENT_MASK), (2, K.DATA_DEPENDENT_MASK)) + assert ( + r.refusal.kind == "data-dependent-mask" + and r.refusal.line_no == g.accesses[1].line_no + ) + # the exact load's finding is kept (#361 dropped it behind the abstention) + r = check_graph( + g, _bind((4, 1, 1), {"n_elements": 1024}, x_ptr=_facts(10), out_ptr=_facts(10)) + ) + assert _found(r) == [("out-of-bounds", 0)] + assert r.abstained == ((1, K.DATA_DEPENDENT_MASK), (2, K.DATA_DEPENDENT_MASK)) + # UNSAT behind a dropped mask is still a proof + _clean( + check_graph( + g, + _bind( + (4, 1, 1), + {"n_elements": 1024}, + x_ptr=_facts(1024), + out_ptr=_facts(1024), + ), + ) + ) + + +def test_observation_gated_mask_and_path_abstain(): + atomic = AccessEvent( + "atomic_rmw", + "c", + Const(0), + None, + 32, + None, + 1, + atomic=AtomicInfo("add", "acq_rel", "gpu"), + ) + gate = Cmp("slt", Observed(0), Const(4), 32) + g = AccessGraph( + "k", + [FuncArg("c", True, 32), FuncArg("p", True, 32)], + [ + atomic, + _access(Const(10), mask=gate, line=2), + _access(Const(10), path=gate, line=3), + _access(Const(0), mask=gate, line=4), # in bounds: a proof + ], + None, + ) + r = check_graph(g, _bind(c=_facts(1), p=_facts(4))) + assert r.findings == () + assert r.abstained == ((1, K.DATA_DEPENDENT_MASK), (2, K.UNMODELABLE_CONDITION)) + + +@pytest.mark.parametrize( + "name", ["p4_observed_direct", "p4_observed_loop", "p4_observed_delta"] +) +def test_observation_in_address_refuses(name): + """An atomic's old value in an address (directly, in a loop-carried + pointer's offset0 or in its delta) would make any address reachable.""" + g = _graph(f"reader_ttir/{name}.ttir") + r = check_graph(g, _bind((4, 1, 1), {"n": 8}, cnt_ptr=_facts(1), x_ptr=_facts(8))) + assert r.findings == () + assert r.abstained == ((1, K.OBSERVATION_IN_ADDRESS), (2, K.OBSERVATION_IN_ADDRESS)) + assert r.refusal.kind == "observation-in-address" + assert r.refusal.line_no == g.accesses[1].line_no + assert r.refusal.message.endswith( + "the address depends on the value an atomic observed" + ) + + +# ─────────────────────────── bindings and loops ─────────────────────────── + + +def test_missing_bindings_refuse(): + tensors = _all(4096, *ADD_TENSORS) + r = check_graph(_graph(ADD), _bind((4, 1, 1), **tensors)) # no n_elements + assert r.findings == () and r.abstained == tuple( + (i, K.MISSING_BINDING) for i in range(3) + ) + assert "scalar argument 'n_elements' has no launch binding" in r.refusal.message + b = _bind((4, 1, 1), {"n_elements": 4096}, **tensors) + b = LaunchBinding( + b.params, b.tensors, b.constexprs, None, None, b.config, "grid: boom" + ) + r = check_graph(_graph(ADD), b) + assert r.abstained == tuple((i, K.MISSING_BINDING) for i in range(3)) + assert "grid is unknown (unreadable: grid: boom)" in r.refusal.message + + +def test_an_unbound_loop_bound_refuses_the_loop(): + g = _graph("reader_ttir/iv_wrap.ttir") + r = check_graph(g, _bind(params={"lo": 0}, x_ptr=_facts(10))) + assert r.findings == () and r.abstained == ((0, K.MISSING_BINDING),) + assert r.refusal.line_no == g.loop.line_no + assert "scalar argument 'n'" in r.refusal.message + + +def test_unmodeled_values_and_unusable_facts_refuse(): + """Defensive: the reader never lets a DataDep into an address, and + bind_launch never builds such facts.""" + g = _synthetic(_access(Bin("+", Const(0), DataDep("loaded value")))) + r = check_graph(g, _bind(p=_facts(4))) + assert r.abstained == ((0, K.UNMODELED_VALUE),) + assert "an unmodeled value (loaded value)" in r.refusal.message + bad = TensorFacts(PTR, 4, 4, (4,), (), "torch.float32", False) + r = check_graph(_synthetic(_access(Const(0))), _bind(p=bad)) + assert r.abstained == ((0, K.MISSING_BINDING),) + assert "unusable" in r.refusal.message + + +def test_exact_finding_kept_alongside_a_refusal(): + """A tensor without a binding refuses its access only: the other + accesses' exact findings are returned with it.""" + r = _add((5, 1, 1), 10**9, **_all(4096, "x_ptr", "out_ptr")) + assert _found(r) == [("out-of-bounds", 0), ("out-of-bounds", 2)] + assert r.abstained == ((1, K.MISSING_BINDING),) + assert r.refusal == Refusal( + kind="missing-binding", + message=r.refusal.message, + line_no=r.findings[0].line_no + 3, + loc=r.refusal.loc, + ) + assert "pointer argument 'y_ptr' has no tensor binding" in r.refusal.message + + +_LOOP_THEN_TAIL = """ + %c0 = arith.constant 0 : i32 + %c10 = arith.constant 10 : i32 + %c100 = arith.constant 100 : i32 + scf.for %i = %c0 to %c10 step %n : i32 { + %a = tt.addptr %p, %i : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + } + %b = tt.addptr %p, %c100 : !tt.ptr, i32 + tt.store %b, %c0 : !tt.ptr +""" + + +@pytest.mark.parametrize("step", [0, -1]) +def test_non_positive_step_refuses_only_the_loop(step): + g, text = _module(_LOOP_THEN_TAIL) + r = check_graph(g, _bind(params={"arg1": step}, arg0=_facts(10))) + assert r.abstained == ((0, K.NON_POSITIVE_STEP),) + assert r.refusal.line_no == _line(text, "scf.for") + assert _found(r) == [("out-of-bounds", 1)] and r.findings[0].violation_offset == 100 + r = check_graph(g, _bind(params={"arg1": 2}, arg0=_facts(10))) + assert _found(r) == [("out-of-bounds", 1)] and r.abstained == () + + +def test_zero_trip_loop_has_no_footprint(): + """An access in the loop whose offset does not mention the induction + variable still runs only on iterations that run (#361 gave a zero-trip + loop's body a witness).""" + g, _ = _module( + """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %c100 = arith.constant 100 : i32 + scf.for %i = %c0 to %n step %c1 : i32 { + %a = tt.addptr %p, %c100 : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""" + ) + assert g.accesses[0].in_loop + _clean(check_graph(g, _bind(params={"arg1": 0}, arg0=_facts(10)))) + _clean(check_graph(g, _bind(params={"arg1": -5}, arg0=_facts(10)))) + ((f,),) = [check_graph(g, _bind(params={"arg1": 1}, arg0=_facts(10))).findings] + assert f.violation_offset == 100 and f.witness["iter_loop"] == 0 + + +def test_solver_unknown_abstains(): + """x^3 + y^3 == z^3 with positive x, y, z (no solution, which Z3 cannot + prove in time) gates an always-OOB offset: unknown, never a proof.""" + x, y, z = Pid(0), Pid(1), Pid(2) + + def cube(t): + return Bin("*", Bin("*", t, t), t) + + def pos(t): + return Cmp("sgt", t, Const(0)) + + fermat = BoolBin( + "and", + BoolBin("and", pos(x), pos(y)), + BoolBin("and", pos(z), Cmp("eq", Bin("+", cube(x), cube(y)), cube(z))), + ) + g = _synthetic(_access(Const(-1), mask=fermat)) + before = z3.get_param("timeout") + r = check_graph(g, _bind((1 << 20,) * 3, p=_facts(4)), timeout_ms=50) + assert z3.get_param("timeout") == before # a per-Solver timeout only (D11) + assert r.findings == () and r.abstained == ((0, K.SOLVER_UNKNOWN),) + assert r.refusal.kind == "solver-unknown" + assert "could not decide the out-of-bounds query" in r.refusal.message + + +# ─────────────────────────── view footprint (D12) ─────────────────────────── + + +def test_strided_view_gap_is_out_of_bounds(): + """x[::2] has 4096 elements two apart: the odd offsets between them are + outside the view, although every offset is below numel.""" + x = torch.empty(8192)[::2] + tensors = {"x_ptr": tensor_facts(x), **_all(4096, "y_ptr", "out_ptr")} + r = _add((4, 1, 1), 4096, **tensors) + ((f,),) = [r.findings] + assert (f.kind, f.access_index) == ("out-of-bounds", 0) and r.refusal is None + assert f.violation_offset % 2 == 1 and 0 < f.violation_offset < 4096 + assert f.violation_address == x.data_ptr() + f.violation_offset * 4 + assert "shape (4096,) and strides (2,)" in f.detail + + +def _tile2d(stride_m: int, stride_n: int, inp, out) -> CheckResult: + params = {"M": 64, "N": 64, "stride_m": stride_m, "stride_n": stride_n} + return check_graph( + _graph(TILE2D), + _bind((2, 2, 1), params, in_ptr=tensor_facts(inp), out_ptr=tensor_facts(out)), + ) + + +def test_strided_views_indexed_by_their_strides_are_clean(): + # a column slice (strides (128, 1)) and a transpose (a dense permutation) + sliced = torch.empty(64, 128)[:, :64] + _clean(_tile2d(128, 1, sliced, torch.empty(64, 128)[:, :64])) + t = torch.empty(64, 64).t() + _clean(_tile2d(1, 64, t, torch.empty(64, 64).t())) + # indexing the slice as if it were dense reads the gaps + r = _tile2d(64, 1, sliced, torch.empty(64, 64)) + ((f,),) = [r.findings] + assert f.access_index == 0 and 64 <= f.violation_offset % 128 < 128 + + +def test_broadcast_stride0_tensor(): + """A row expanded to 64x64 (strides (0, 1)) has 64 distinct elements: + offsets past them are out of bounds although below numel (4096).""" + row = torch.empty(64).expand(64, 64) + assert tensor_facts(row).numel == 4096 + _clean(_tile2d(0, 1, row, torch.empty(64, 64))) + r = _tile2d(64, 1, row, torch.empty(64, 64)) + ((f,),) = [r.findings] + assert f.access_index == 0 and 64 <= f.violation_offset < 4096 + + +def _views(): + base = torch.empty(64) + return { + "contiguous": base, + "step2": base[::2], + "offset_step3": base[3:40:3], + "column_slice": base.view(8, 8)[:, :5], + "transpose": base.view(8, 8).t(), + "grid_slice": base.view(8, 8)[::2, 1::3], + "broadcast_rows": torch.empty(5).expand(3, 5), + "broadcast_cols": torch.empty(4, 1).expand(4, 6), + "gappy": base.as_strided((3, 3), (5, 1)), + "overlapping": base.as_strided((3, 4), (2, 1)), + "scalar": torch.empty(()), + "empty": torch.empty(0), + "empty_2d": torch.empty(3, 0), + "channels_last": torch.empty(2, 3, 2, 2).to(memory_format=torch.channels_last), + "expand_to_0": torch.empty(1).expand(0), + "size1_odd_stride": base.as_strided((1, 4), (100, 1)), + "mixed_broadcast": base.as_strided((3, 4, 5), (0, 10, 2)), + "duplicate_strides": base.as_strided((3, 3), (3, 3)), + } + + +def _eager_offsets(t: torch.Tensor) -> set[int]: + """The element offsets the eager sanitizer admits for ``t``: the byte + segments SymbolicClient._tensor_physical_addresses dispatches to + (contiguous, storage-contiguous, inner-stride-1 slices, per element).""" + if not t.numel(): + return set() + base, item = t.data_ptr(), t.element_size() + return { + (a - base) // item + for start, end, _ in SymbolicClient._tensor_physical_addresses(None, "t", t) + for a in range(start, end + 1) + if (a - base) % item == 0 + } + + +@pytest.mark.parametrize("name", list(_views())) +def test_view_footprint_matches_the_eager_legal_set(name): + """The element offsets _in_view admits are exactly the ones the eager + sanitizer's own dispatch admits (D12).""" + t = _views()[name] + facts = tensor_facts(t) + eager = _eager_offsets(t) + for e in range(-3, max([70, *eager]) + 8): + solver = z3.Solver() + solver.add(_in_view(z3.IntVal(e), facts)) + assert (solver.check() == z3.sat) == (e in eager), (name, e) + + +def test_empty_tensor(): + """No element is legal: every executed access is out of bounds, and an + access that never executes is none.""" + g = _graph("ttir/golden_cas_sm80.ttir") + empty = TensorFacts(PTR, 4, 0, (0,), (1,), "torch.int32", True) + r = check_graph(g, _bind(lock_ptr=empty, out_ptr=_facts(1))) + ((f,),) = [r.findings] + assert (f.access_index, f.access_kind, f.violation_offset) == (0, "atomic_cas", 0) + assert "the empty tensor" in f.detail + _clean(_add((4, 1, 1), 0, **{n: empty for n in ADD_TENSORS})) + + +def test_reinterpreted_width_checks_every_byte(): + """An f32 pointer over a byte tensor (a width-changing reinterpret): + offsets count f32 elements, the view counts bytes.""" + g = _graph(ADD) + bytes_ = {n: _facts(4096, elem_size=1) for n in ADD_TENSORS} + _clean(check_graph(g, _bind((1, 1, 1), {"n_elements": 1024}, **bytes_))) + short = {**bytes_, "x_ptr": _facts(4095, elem_size=1)} + r = check_graph(g, _bind((1, 1, 1), {"n_elements": 1024}, **short)) + ((f,),) = [r.findings] + assert f.access_index == 0 and f.violation_offset == 1023 + assert f.violation_address == PTR + 1023 * 4 + + +# ─────────────────────────── integer widths (D9) ─────────────────────────── + + +def test_i32_wrap_is_integer_overflow_not_out_of_bounds(): + """(pid * S) * S wraps to 0 in i32 for S = 65536: every program stores + x[0]. The unbounded reading's offsets are an overflow, not OOB.""" + g = _graph("reader_ttir/rv_i32_wrap.ttir") + r = check_graph(g, _bind((4, 1, 1), {"S": 65536}, x_ptr=_facts(1))) + ((f,),) = [r.findings] + assert (f.kind, f.access_index) == ("integer-overflow", 0) and r.refusal is None + assert f.violation_offset is None and f.violation_address is None + assert not -(1 << 31) <= f.witness["value"] < 1 << 31 + assert ( + "arith.muli" + in _text("reader_ttir/rv_i32_wrap.ttir").splitlines()[f.line_no - 1] + ) + assert "i32 range" in f.detail + # no wrap for a small S: a clean launch + _clean(check_graph(g, _bind((4, 1, 1), {"S": 2}, x_ptr=_facts(16)))) + + +def test_trunci_is_integer_overflow_not_out_of_bounds(): + """trunc_i32(pid_i64 * 2**32) is 0 for every pid.""" + g = _graph("reader_ttir/rv_trunci_alias.ttir") + r = check_graph(g, _bind((2, 1, 1), x_ptr=_facts(1))) + ((f,),) = [r.findings] + assert (f.kind, f.witness["pid_0"], f.witness["value"]) == ( + "integer-overflow", + 1, + 1 << 32, + ) + assert ( + "arith.trunci" + in _text("reader_ttir/rv_trunci_alias.ttir").splitlines()[f.line_no - 1] + ) + assert "truncation to i32" in f.detail + _clean(check_graph(g, _bind((1, 1, 1), x_ptr=_facts(1)))) + + +def test_loop_increment_overflow(): + """range(lo, n, 1 << 20) wraps its induction variable near INT32_MAX.""" + g = _graph("reader_ttir/iv_wrap.ttir") + n = (1 << 31) - (1 << 19) + r = check_graph(g, _bind(params={"lo": 0, "n": n}, x_ptr=_facts(1 << 31))) + ((f,),) = [r.findings] + assert (f.kind, f.access_index) == ("integer-overflow", 0) + assert "scf.for" in _text("reader_ttir/iv_wrap.ttir").splitlines()[f.line_no - 1] + assert "induction-variable increment" in f.detail + _clean(check_graph(g, _bind(params={"lo": 0, "n": 1000}, x_ptr=_facts(1000)))) + # a zero-trip loop never increments + _clean(check_graph(g, _bind(params={"lo": n, "n": n}, x_ptr=_facts(1)))) + + +def test_unsigned_read_of_a_negative_param(): + """``pid < n.to(tl.uint32)`` reads n unsigned; the model reads the i32 + argument signed (2**32 - 1 and -1 are one i32).""" + g = _graph("reader_ttir/unsigned_index.ttir") + _clean(check_graph(g, _bind((8, 1, 1), {"n": 8}, x_ptr=_facts(3)))) + for n in (-1, (1 << 32) - 1): + r = check_graph(g, _bind((8, 1, 1), {"n": n}, x_ptr=_facts(3))) + ((f,),) = [r.findings] + assert f.kind == "integer-overflow" and f.witness["value"] == -1 + assert ( + "cmpi ult" + in _text("reader_ttir/unsigned_index.ttir").splitlines()[f.line_no - 1] + ) + + +def test_one_finding_per_op_site(): + """The add kernel's three accesses share one mask (pid * 1024 + lane): + its wrap is reported once.""" + r = _add((1 << 22, 1, 1), (1 << 31) - 1, **_all(1 << 32, *ADD_TENSORS)) + assert _found(r) == [("integer-overflow", 0)] and r.refusal is None + + +# A wrap that decides its own role's divisor (the audit corpus's N kernels, +# with BIG = 2**31 - 1): t = pid + BIG, d = where(t > BIG, 0, -3). At pid 1, +# t wraps in i32, so the kernel divides by -3 and r = t // d - BIG // -3 is +# 1431655764 (0 at pid 0); in the unbounded reading d is 0 there. Checking +# the wrap assuming the divisor non-zero, and the divisor assuming no wrap, +# neither query has a model: that was a false proof in every role. +_CIRCULAR = """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %c4 = arith.constant 4 : i32 + %cm3 = arith.constant -3 : i32 + %pid = tt.get_program_id x : i32 + %t = arith.addi %pid, %n : i32 + %gt = arith.cmpi sgt, %t, %n : i32 + %d = arith.select %gt, %c0, %cm3 : i32 + %q = arith.divsi %t, %d : i32 + %b3 = arith.divsi %n, %cm3 : i32 + %r = arith.subi %q, %b3 : i32 +""" +_CIRCULAR_ACCESS = { + # x + r + "offset": """ + %a = tt.addptr %p, %r : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr""", + # if r > 0: x[1] + "path": """ + %pos = arith.cmpi sgt, %r, %c0 : i32 + scf.if %pos { + %a = tt.addptr %p, %c1 : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""", + # x + arange(16), masked to lanes < r + 1 + "mask": """ + %lim = arith.addi %r, %c1 : i32 + %offs = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32> + %sl = tt.splat %lim : i32 -> tensor<16xi32> + %m = arith.cmpi slt, %offs, %sl : tensor<16xi32> + %ps = tt.splat %p : !tt.ptr -> tensor<16x!tt.ptr> + %a = tt.addptr %ps, %offs : tensor<16x!tt.ptr>, tensor<16xi32> + %z = arith.constant dense<0> : tensor<16xi32> + tt.store %a, %z, %m : tensor<16x!tt.ptr>""", + # for i in range(min(r + 1, 4)): x[i] + "loop": """ + %lim = arith.addi %r, %c1 : i32 + %hi = arith.minsi %lim, %c4 : i32 + scf.for %i = %c0 to %hi step %c1 : i32 { + %a = tt.addptr %p, %i : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""", +} + + +@pytest.mark.parametrize("role", list(_CIRCULAR_ACCESS)) +def test_a_role_never_assumes_its_own_conditions(role): + """Each role is one joint query: the wrap at pid 1 is found, as the + innermost failure (the divisor's zero is only the unbounded reading's).""" + g, text = _module(_CIRCULAR + _CIRCULAR_ACCESS[role]) + big = (1 << 31) - 1 + r = check_graph(g, _bind((2, 1, 1), {"arg1": big}, arg0=_facts(1))) + assert _found(r) == [("integer-overflow", 0)] and r.refusal is None + (f,) = r.findings + assert f.line_no == _line(text, "arith.addi %pid") + assert (f.witness["pid_0"], f.witness["value"]) == (1, 1 << 31) + # pid 0 alone is in bounds, and proved so + _clean(check_graph(g, _bind((1, 1, 1), {"arg1": big}, arg0=_facts(1)))) + + +def test_a_divisor_zero_where_a_sibling_wraps_is_found(): + """The reverse dependence: pid 1 both wraps (pid + BIG) and divides by + zero (d = 1 - pid); either is a finding, never a proof.""" + g, _ = _module( + """ + %c1 = arith.constant 1 : i32 + %pid = tt.get_program_id x : i32 + %t = arith.addi %pid, %n : i32 + %d = arith.subi %c1, %pid : i32 + %q = arith.divsi %pid, %d : i32 + %o = arith.subi %t, %n : i32 + %s = arith.addi %o, %q : i32 + %a = tt.addptr %p, %s : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + r = check_graph(g, _bind((2, 1, 1), {"arg1": (1 << 31) - 1}, arg0=_facts(1))) + (f,) = r.findings + assert f.kind in ("integer-overflow", "division-by-zero") + assert f.witness["pid_0"] == 1 and r.refusal is None + + +_SELECT_ARM = """ + %c0 = arith.constant 0 : i32 + %c10 = arith.constant 10 : i32 + %big = arith.constant 1073741824 : i32 + %pid = tt.get_program_id x : i32 + %is0 = arith.cmpi eq, %pid, %c0 : i32 + %m = arith.muli %pid, %big : i32 +""" + + +def test_a_wrap_in_the_arm_a_select_discards_is_no_finding(): + """off = where(pid == 0, pid * 2**30, pid): pid * 2**30 wraps from pid 2 + on, but only pid 0 reads it (the audit's H14 kernel).""" + g, _ = _module( + _SELECT_ARM + + """ + %o = arith.select %is0, %m, %pid : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + _clean(check_graph(g, _bind((4, 1, 1), arg0=_facts(4)))) + + +def test_a_wrap_in_the_arm_a_select_takes_is_a_finding(): + g, text = _module( + _SELECT_ARM + + """ + %o = arith.select %is0, %pid, %m : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + (f,) = check_graph(g, _bind((4, 1, 1), arg0=_facts(1 << 32))).findings + assert (f.kind, f.line_no) == ("integer-overflow", _line(text, "arith.muli")) + assert f.witness["pid_0"] in (2, 3) + + +def test_a_value_also_read_directly_is_checked_unguarded(): + """pid * 2**30 sits in the select's discarded arm of the offset, but the + mask reads it directly: its wrap matters where the mask is computed.""" + g, text = _module( + _SELECT_ARM + + """ + %lt = arith.cmpi slt, %m, %c10 : i32 + %o = arith.select %is0, %m, %pid : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + %v = tt.load %a, %lt : !tt.ptr""" + ) + (f,) = check_graph(g, _bind((4, 1, 1), arg0=_facts(4))).findings + assert (f.kind, f.line_no) == ("integer-overflow", _line(text, "arith.muli")) + + +# A loop bound (the step; the lower bound) wraps in the arm of a select on +# %m, taken where m < 100; the loop reads the bound itself. +_BOUND_IN_ARM = { + "step truncated": ( + """ + %c0 = arith.constant 0 : i32 + %c100 = arith.constant 100 : i32 + %c1000 = arith.constant 1000 : i32 + %small = arith.cmpi slt, %m, %c100 : i32 + %t = arith.trunci %n : i32 to i8 + %e = arith.extsi %t : i8 to i32 + %sel = arith.select %small, %e, %c0 : i32 + scf.for %i = %c0 to %c1000 step %n : i32 { + %o = arith.addi %i, %sel : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""", + 300, + "arith.trunci", + ), + "lower bound read unsigned": ( + """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %c8 = arith.constant 8 : i32 + %c64 = arith.constant 64 : i32 + %c100 = arith.constant 100 : i32 + %hi = arith.addi %n, %c8 : i32 + %small = arith.cmpi slt, %m, %c100 : i32 + %u = arith.minui %n, %c8 : i32 + %sel = arith.select %small, %u, %c0 : i32 + scf.for %i = %n to %hi step %c1 : i32 { + %j = arith.addi %i, %c64 : i32 + %o = arith.addi %j, %sel : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""", + -10, + "arith.minui", + ), +} + + +@pytest.mark.parametrize("case", list(_BOUND_IN_ARM)) +def test_a_loop_bound_wrapping_in_a_discarded_arm_is_no_finding(case): + """H14 for a loop bound: the loop reads the bound itself, but that read + is the loop's (checked with the loop's obligations), so the access + reads the wrapping op only through the arm: its wrap matters only + where the select takes the arm.""" + body, n, op = _BOUND_IN_ARM[case] + g, text = _module(body, "%p: !tt.ptr, %n: i32, %m: i32") + + def at(m): + return check_graph(g, _bind(params={"arg1": n, "arg2": m}, arg0=_facts(1300))) + + _clean(at(500)) + (f,) = at(5).findings + assert (f.kind, f.line_no) == ("integer-overflow", _line(text, op)) + assert f.witness["value"] == n + + +def test_undefined_divisions_count_in_either_arm(): + """where(n != 0, pid // n, 0) divides by zero in the discarded arm too + (the audit's p12: the GPU run faults), and so does INT_MIN // -1.""" + g, text = _module( + """ + %c0 = arith.constant 0 : i32 + %cm1 = arith.constant -1 : i32 + %pid = tt.get_program_id x : i32 + %nz = arith.cmpi ne, %n, %c0 : i32 + %q = arith.divsi %pid, %n : i32 + %o = arith.select %nz, %q, %c0 : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + r = check_graph(g, _bind((4, 1, 1), {"arg1": 0}, arg0=_facts(4))) + assert _found(r) == [("division-by-zero", 0)] + _clean(check_graph(g, _bind((4, 1, 1), {"arg1": 2}, arg0=_facts(4)))) + g, text = _module( + """ + %c0 = arith.constant 0 : i32 + %c5 = arith.constant 5 : i32 + %cm1 = arith.constant -1 : i32 + %pid = tt.get_program_id x : i32 + %never = arith.cmpi eq, %pid, %c5 : i32 + %q = arith.divsi %n, %cm1 : i32 + %o = arith.select %never, %q, %c0 : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + (f,) = check_graph(g, _bind(params={"arg1": -(1 << 31)}, arg0=_facts(1))).findings + assert (f.kind, f.line_no) == ("integer-overflow", _line(text, "arith.divsi")) + assert f.witness["value"] == 1 << 31 and "the result of '//'" in f.detail + + +# ─────────────────────────── division by zero (D21) ─────────────────────────── + + +def test_division_by_zero_in_an_offset(): + g, text = _module( + """ + %pid = tt.get_program_id x : i32 + %q = arith.divsi %pid, %n : i32 + %a = tt.addptr %p, %q : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + r = check_graph(g, _bind((4, 1, 1), {"arg1": 0}, arg0=_facts(2))) + assert _found(r) == [("division-by-zero", 0)] and r.refusal is None + assert r.findings[0].line_no == _line(text, "arith.divsi") + assert "divisor of this division" in r.findings[0].detail + _clean(check_graph(g, _bind((4, 1, 1), {"arg1": 2}, arg0=_facts(2)))) + + +def test_division_by_zero_in_a_mask(): + g, text = _module( + """ + %pid = tt.get_program_id x : i32 + %c1 = arith.constant 1 : i32 + %r = arith.remsi %pid, %n : i32 + %m = arith.cmpi slt, %r, %c1 : i32 + %v = tt.load %p, %m : !tt.ptr""" + ) + r = check_graph(g, _bind((4, 1, 1), {"arg1": 0}, arg0=_facts(1))) + assert _found(r) == [("division-by-zero", 0)] + assert r.findings[0].line_no == _line(text, "arith.remsi") + assert "remainder" in r.findings[0].detail + _clean(check_graph(g, _bind((4, 1, 1), {"arg1": 1}, arg0=_facts(1)))) + + +def test_division_by_zero_in_a_loop_bound(): + g, text = _module( + """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %c8 = arith.constant 8 : i32 + %u = arith.divsi %c8, %n : i32 + scf.for %i = %c0 to %u step %c1 : i32 { + %a = tt.addptr %p, %i : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""" + ) + r = check_graph(g, _bind(params={"arg1": 0}, arg0=_facts(4))) + assert _found(r) == [("division-by-zero", 0)] + assert r.findings[0].line_no == _line(text, "arith.divsi") + _clean(check_graph(g, _bind(params={"arg1": 2}, arg0=_facts(4)))) + ((f,),) = [check_graph(g, _bind(params={"arg1": 2}, arg0=_facts(3))).findings] + assert f.kind == "out-of-bounds" and f.violation_offset == 3 + + +# ─────────────────────────── robustness ─────────────────────────── + + +def test_deep_terms_are_lowered_iteratively(): + """kernel_deep_chain's offset nests more than 1000 levels: at Python's + default recursion limit the generated == / hash would raise.""" + g = _graph("ttir/kernel_deep_chain.ttir") + limit = sys.getrecursionlimit() + sys.setrecursionlimit(1000) + try: + clean = check_graph(g, _bind((4, 1, 1), {"s": 1}, out_ptr=_facts(1 << 20))) + wrap = check_graph(g, _bind((4, 1, 1), {"s": 3}, out_ptr=_facts(1 << 20))) + finally: + sys.setrecursionlimit(limit) + _clean(clean) + assert _found(wrap) == [("integer-overflow", 0)] + + +def test_a_launch_key_ignores_only_the_addresses(): + """Launches whose bindings share a launch_key get one result, the + findings' addresses moved to each launch's tensors; nothing else in a + result (a withheld finding's message included) holds an address.""" + g = _graph(ADD) + + def at(ptr): + return _bind( + (5, 1, 1), + {"n_elements": 10**9}, + **{n: _facts(4096, data_ptr=ptr) for n in ADD_TENSORS}, + ) + + first, moved = at(PTR), at(PTR + 0x10000) + assert launch_key(first) == launch_key(moved) + result = check_graph(g, first) + again = readdressed(result, g, moved) + assert again == check_graph(g, moved) != result + assert [f.violation_address for f in again.findings] == [ + PTR + 0x10000 + f.violation_offset * 4 for f in result.findings + ] + assert launch_key(at(PTR)) != launch_key(_bind((5, 1, 1), {"n_elements": 1})) + guarded = _synthetic(_access(Bin("-", Pid(0), Const(1)), guarded=True)) + messages = { + check_graph(guarded, _bind((4, 1, 1), p=_facts(4, data_ptr=ptr))).refusal + for ptr in (PTR, PTR + 0x10000) + } + assert len(messages) == 1 and "0x" not in messages.pop().message + + +@pytest.mark.parametrize("timeout_ms", [0, -1, 2.5, True]) +def test_a_timeout_must_be_positive(timeout_ms): + # Z3 reads 0 and below as no timeout at all. + with pytest.raises(ValueError, match="positive int"): + check_graph(_graph(ADD), _bind(), timeout_ms=timeout_ms) + + +class _InterruptedSolver: + """A Solver whose query Z3's Ctrl+C handler cancelled.""" + + def __init__(self, ctx=None): + pass + + def set(self, **kwargs): + pass + + def add(self, *formulas): + pass + + def check(self): + return z3.unknown + + def reason_unknown(self): + return "interrupted from keyboard" + + +def test_a_ctrl_c_during_a_query_is_raised(monkeypatch): + """Z3 turns a Ctrl+C into an unknown; it is the user's interrupt, never + a solver-unknown abstention.""" + monkeypatch.setattr(oob, "Solver", _InterruptedSolver) + with pytest.raises(KeyboardInterrupt): + _add((4, 1, 1), 4096, **_all(4096, *ADD_TENSORS)) + + +_SIGINT_SCRIPT = """ +import os, signal, threading, time +from types import MappingProxyType +from tilelens.clients.sanitizer.compiled.oob import check_graph +from tilelens.ir.launch import LaunchBinding, TensorFacts +from tilelens.ir.ttir_reader import ( + AccessEvent, AccessGraph, Bin, BoolBin, Cmp, Const, FuncArg, Pid, +) + +def cube(t): + return Bin("*", Bin("*", t, t), t) + +x, y, z = Pid(0), Pid(1), Pid(2) +pos = [Cmp("sgt", t, Const(0)) for t in (x, y, z)] +fermat = BoolBin("and", BoolBin("and", pos[0], pos[1]), BoolBin( + "and", pos[2], Cmp("eq", Bin("+", cube(x), cube(y)), cube(z)))) +graph = AccessGraph("k", [FuncArg("p", True, 32)], + [AccessEvent("load", "p", Const(-1), fermat, 32, None, 1)], None) +facts = TensorFacts(4096, 4, 4, (4,), (1,), "torch.float32", True) +grid = (1 << 20,) * 3 +binding = LaunchBinding(MappingProxyType({}), MappingProxyType({"p": facts}), + MappingProxyType({}), grid, grid, MappingProxyType({})) +threading.Timer(1.5, os.kill, (os.getpid(), signal.SIGINT)).start() +start = time.monotonic() +try: + result = check_graph(graph, binding, timeout_ms=60_000) +except KeyboardInterrupt as exc: + print("interrupted", repr(exc), round(time.monotonic() - start)) +else: + print("returned", result.abstained) +""" + + +def test_a_real_ctrl_c_stops_a_hard_query(): + """A SIGINT 1.5 s into a query Z3 cannot decide in a minute (the Fermat + mask of test_solver_unknown_abstains): Z3 cancels it, and the check + raises KeyboardInterrupt at once instead of abstaining and going on.""" + proc = subprocess.run( + [sys.executable, "-c", _SIGINT_SCRIPT], + capture_output=True, + text=True, + cwd=REPO, + timeout=120, + ) + assert proc.returncode == 0, proc.stderr + assert proc.stdout.startswith( + "interrupted KeyboardInterrupt('Z3 query" + ), proc.stdout + assert int(proc.stdout.split()[-1]) < 30 + + +@pytest.mark.parametrize( + "params", + [ + # a loop (its iteration lowers the bounds) and a finding's witness + _MATMUL_PARAMS, + # the loop refused: its bound reads K, which has no binding + {name: v for name, v in _MATMUL_PARAMS.items() if name != "K"}, + ], + ids=["finding", "refused-loop"], +) +def test_a_check_leaves_no_z3_object_to_the_cyclic_gc(params): + """A check's Z3 context and terms are freed by reference counting when + check_graph returns: the cyclic GC would free them later, on whichever + host thread it runs, inside that thread's own Z3 call (concurrent + checks hung or crashed).""" + tensors = _all(128 * 128, "a_ptr", "b_ptr", "c_ptr", elem_size=2) + graph = _graph(MATMUL) + binding = _bind((3, 2, 1), params, **tensors) + + def z3_objects() -> int: + # type(): isinstance would read __class__, which some objects warn on + z3_types = (z3.Context, z3.AstRef) + return sum(issubclass(type(o), z3_types) for o in gc.get_objects()) + + check_graph(graph, binding) + gc.collect() + gc.disable() + try: + before = z3_objects() + result = check_graph(graph, binding) + after = z3_objects() + finally: + gc.enable() + if "K" in params: + assert result.findings + else: + assert K.MISSING_BINDING in {kind for _, kind in result.abstained} + assert after == before + + +_THREADS_SCRIPT = """ +import sys, threading +from pathlib import Path +from types import MappingProxyType +from tilelens.clients.sanitizer.compiled.oob import check_graph +from tilelens.ir.launch import LaunchBinding, TensorFacts +from tilelens.ir.ttir_reader import parse_ttir + +golden = Path("tests/golden/ir/ttir") +graph = parse_ttir((golden / "golden_add_sm80.ttir").read_text()) +facts = TensorFacts(4096, 4, 4096, (4096,), (1,), "torch.float32", True) +grid = (5, 1, 1) +binding = LaunchBinding( + MappingProxyType({"n_elements": 10**6}), + MappingProxyType(dict.fromkeys(("x_ptr", "y_ptr", "out_ptr"), facts)), + MappingProxyType({}), grid, grid, MappingProxyType({})) +errors, kinds = [], set() + +def worker(): + try: + for _ in range(25): + kinds.add(tuple(f.kind for f in check_graph(graph, binding).findings)) + except Exception as exc: + errors.append(repr(exc)) + +threads = [threading.Thread(target=worker) for _ in range(4)] +for t in threads: + t.start() +for t in threads: + t.join() +print(errors[:3], sorted(kinds)) +""" + + +def test_checks_on_several_host_threads_at_once(): + """Two traced kernels finalized on different host threads check at the + same time: each check has a Z3 context of its own (one context shared + by threads failed with Z3 errors or segfaulted). In a subprocess, so a + crash is a failure here and not the test run's end.""" + proc = subprocess.run( + [sys.executable, "-c", _THREADS_SCRIPT], + capture_output=True, + text=True, + cwd=REPO, + timeout=600, + ) + assert proc.returncode == 0, proc.stderr[-2000:] + oob_kinds = ("out-of-bounds",) * 3 + assert proc.stdout.strip() == f"[] [{oob_kinds!r}]" + + +def test_results_are_plain_data(): + r = _add((5, 1, 1), 10**9, **_all(4096, "x_ptr", "out_ptr")) + again = pickle.loads(pickle.dumps(r)) + assert again == r and isinstance(again.findings[0], Finding) + # a Refusal holds its kind as a plain str; the abstentions keep the enum + assert not isinstance(r.refusal.kind, SanitizerKind) + assert ( + r.refusal.kind == "missing-binding" and r.abstained[0][1] is K.MISSING_BINDING + ) + with pytest.raises(TypeError): + hash(r.findings[0]) + # the sanitizer's kinds are its own, apart from the reader's + assert not {k.value for k in SanitizerKind} & {k.value for k in TTIRKind} + assert f"{K.SOLVER_UNKNOWN}" == str(K.SOLVER_UNKNOWN) == "solver-unknown" diff --git a/tests/unit/test_ir_lifecycle.py b/tests/unit/test_ir_lifecycle.py new file mode 100644 index 000000000..1d4900694 --- /dev/null +++ b/tests/unit/test_ir_lifecycle.py @@ -0,0 +1,2474 @@ +"""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, D25). +These tests pin call sequences, so a fake stands in for that compile: a +``_FakeJit`` compiles through its ``fake_compile`` and launches through its +``run``, and a real JITFunction gets both 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 (against +the JIT's own compile in tests/end_to_end/test_host_compile.py) and end to +end in tests/end_to_end/test_ir_lifecycle_compiled.py. +""" + +import ast +import gc +import importlib +import inspect +import threading +import types +import weakref +from contextlib import contextmanager +from types import SimpleNamespace + +import pytest +import torch +import triton +import triton.language as tl +from triton.compiler.errors import CompileTimeAssertionFailure +from triton.runtime import Autotuner +from triton.runtime.autotuner import Heuristics +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, + KernelTraceSupport, + NKITrace, + TraceInterface, + TritonTrace, + _untraced_call_args, + _unwrapped_trace_globals, +) + +# `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, stages).""" + compiles: list[tuple] = [] + + def compile(self, jit_fn, args, kwargs, *, target, stages=()): + 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, frozenset(stages))) + 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 _OtherEagerClient(_EagerClient): + NAME = "other_eager" + + +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"}) + + 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" + LAUNCH = "skip" + + +class _RunIRClient(_IRClient): + NAME = "ir_run" + LAUNCH = "run" + + +class _IndifferentIRClient(_IRClient): + NAME = "ir_indifferent" + + +# ======== 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 (D25). + 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, + a real launch run().""" + + 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): + self.log.append("compile" if warmup else "launch") + return None if warmup else "launched" + + +def _install_fake_run(monkeypatch, jit_fn, run): + """Install ``run(*args, grid, warmup, **kwargs)`` on a real JITFunction + as both its launch (``run``, warmup=False) 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 (``warmup=True``) and + launches (``warmup=False``), in order, instead of running 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 _grid8(meta): + # The interpreter hands grid callables tensor-converted runtime args, so + # interpreted launches may only read constexprs here. + return (triton.cdiv(8, 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 + + +# ======== 1 / 2a-2b: 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" + + +def test_add_clients_rejects_skip_run_conflict_before_inserting(): + manager = ClientManager([_SkipIRClient(), _EagerClient()]) + + with pytest.raises(RuntimeError, match="Trace the kernel twice"): + manager.add_clients([_IndifferentIRClient(), _RunIRClient()]) + + # Nothing from the rejected batch was inserted. + assert [c.NAME for c in manager.clients] == ["ir_skip", "eager"] + + with pytest.raises(RuntimeError, match="LAUNCH='skip'"): + ClientManager([_RunIRClient(), _SkipIRClient()]) + + +def test_launch_conflict_check_sees_every_ir_client(): + # A same-NAME IR client joins the first rather than replacing it, so + # their conflict is seen. + class _RunInSkipSlot(_IRClient): + NAME = "ir_skip" + LAUNCH = "run" + + first = _SkipIRClient() + manager = ClientManager([first]) + with pytest.raises(RuntimeError, match="cannot share one trace"): + manager.add_clients([_RunInSkipSlot()]) + assert manager.clients == [first] + + # An interpreting client's LAUNCH takes no part in the vote. + class _EagerSkip(_EagerClient): + NAME = "eager_skip" + LAUNCH = "skip" + + manager = ClientManager([_EagerSkip()]) + manager.add_clients([_RunIRClient()]) + assert manager.launch_policy() == "run" + + +def test_add_clients_keeps_every_ir_client_instance(): + # Adding a client already in the trace changes nothing; another + # instance of its class is kept, not dropped (e.g. one per target). + first, second = _SkipIRClient(), _SkipIRClient() + indifferent, eager = _IndifferentIRClient(), _EagerClient() + manager = ClientManager([first, indifferent, eager]) + manager.add_clients([first, second, second]) + + assert manager.clients == [first, indifferent, eager, second] + assert manager.ir_clients() == [first, indifferent, second] + assert manager.get_client("ir_skip") is first + + +def test_add_clients_refuses_a_second_interpreting_client_of_one_name(): + # One interpreted run serves one client per NAME: another one, of the + # same class or not, raises rather than being dropped, and nothing from + # its batch is inserted. + class _SameName(_SiblingEagerClient): + NAME = "eager" + + first = _EagerClient() + manager = ClientManager([first]) + for duplicate in (_EagerClient(), _SameName()): + with pytest.raises(ValueError, match="interpreting client named 'eager'"): + manager.add_clients([_IndifferentIRClient(), duplicate]) + assert manager.clients == [first] + manager.add_clients([first]) + assert manager.clients == [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_rejects_conflicting_launch_preferences(): + traced = tilelens.trace(_SkipIRClient())(_make_plain_kernel()) + + with pytest.raises(RuntimeError, match="cannot share one trace"): + tilelens.trace(_RunIRClient())(traced) + + assert [c.NAME for c in traced.client_manager.clients] == ["ir_skip"] + + +def test_client_partition_and_launch_policy(): + eager, ir = _EagerClient(), _IndifferentIRClient() + manager = ClientManager([eager, ir]) + + assert manager.interpreting_clients() == [eager] + assert manager.ir_clients() == [ir] + assert manager.launch_policy() == "run" + + manager.add_clients([_SkipIRClient()]) + assert manager.launch_policy() == "skip" + + # Only IR clients' declarations drive the policy. + assert ClientManager([_EagerClient()]).launch_policy() == "run" + + +# ======== 2c: patch_warmup ========= + + +class _FakeWarmupJit: + def __init__(self): + self.warmups: list[dict] = [] + + def warmup(self, *args, **kwargs): + self.warmups.append(kwargs) + return "compiled" + + +def test_patch_warmup_polls_every_client_and_compiles_on_any_vote(): + voter, abstainer = _EagerClient(warmup_vote=True), _OtherEagerClient() + ir = _IndifferentIRClient() + manager = ClientManager([voter, abstainer, ir]) + jit_fn = _FakeWarmupJit() + + with manager.patch_warmup(jit_fn): + ret = jit_fn.warmup(1, grid=(1,), warmup=False) + + assert ret == "compiled" + assert jit_fn.warmups == [{"grid": (1,)}] + # No short-circuit after the first True vote; every client votes and + # every client sees the result. + assert voter.calls == ["pre_warmup", ("post_warmup", "compiled")] + assert abstainer.calls == ["pre_warmup", ("post_warmup", "compiled")] + assert ir.log == ["pre_warmup", "post_warmup"] + assert "warmup" not in vars(jit_fn) + + +def test_patch_warmup_skips_compile_without_votes(): + client = _EagerClient() + manager = ClientManager([client, _OtherEagerClient()]) + jit_fn = _FakeWarmupJit() + + with manager.patch_warmup(jit_fn): + assert jit_fn.warmup(1, grid=(1,)) is None + + assert jit_fn.warmups == [] + assert client.calls == ["pre_warmup"] + + +def test_patch_warmup_enters_the_compile_context_only_for_a_real_compile(): + entered: list = [] + + @contextmanager + def compile_context(): + entered.append("enter") + yield + entered.append("exit") + + jit_fn = _FakeWarmupJit() + abstaining = ClientManager([_EagerClient()]) + with abstaining.patch_warmup(jit_fn, compile_context=compile_context): + assert jit_fn.warmup(1, grid=(1,)) is None + assert entered == [] + + voting = ClientManager([_EagerClient(warmup_vote=True)]) + with voting.patch_warmup(jit_fn, compile_context=compile_context): + assert jit_fn.warmup(1, grid=(1,)) == "compiled" + assert entered == ["enter", "exit"] + + +def test_patch_warmup_scopes_on_two_threads_vote_apart_and_leave_no_gate(): + # Two traces sharing a JITFunction warm up on two host threads, their + # scopes interleaved: A opens, B opens, A closes, B closes. + a_client = _EagerClient() + b_client = _OtherEagerClient(warmup_vote=True) + a_manager, b_manager = ClientManager([a_client]), ClientManager([b_client]) + jit_fn = _FakeWarmupJit() + a_in, b_in, a_out = threading.Event(), threading.Event(), threading.Event() + results: dict = {} + + def thread_a(): + with a_manager.patch_warmup(jit_fn): + a_in.set() + b_in.wait(10) + results["a"] = jit_fn.warmup("a", grid=(1,)) + a_out.set() + + def thread_b(): + a_in.wait(10) + with b_manager.patch_warmup(jit_fn): + b_in.set() + a_out.wait(10) + results["b"] = jit_fn.warmup("b", grid=(1,)) + # A thread with no scope of its own is not gated. + results["other"] = _in_thread(jit_fn.warmup, "other")["result"] + + threads = [threading.Thread(target=thread_a), threading.Thread(target=thread_b)] + for thread in threads: + thread.start() + for thread in threads: + thread.join(30) + + # Each call was voted on by its own thread's trace only. + assert results == {"a": None, "b": "compiled", "other": "compiled"} + assert a_client.calls == ["pre_warmup"] + assert b_client.calls == ["pre_warmup", ("post_warmup", "compiled")] + # The last scope to close removed the gate: an untraced warmup compiles + # and polls nobody. + assert "warmup" not in vars(jit_fn) + assert jit_fn.warmup("untraced", grid=(1,)) == "compiled" + assert a_client.calls == ["pre_warmup"] and len(b_client.calls) == 2 + + +def test_patch_warmup_shares_one_gate_and_puts_back_what_was_there(): + jit_fn = _FakeWarmupJit() + + def users_warmup(*args, **kwargs): + return "users" + + jit_fn.warmup = users_warmup + manager = ClientManager([_EagerClient(warmup_vote=True)]) + + with manager.patch_warmup(jit_fn): + gate = jit_fn.warmup + with ClientManager([_OtherEagerClient()]).patch_warmup(jit_fn): + # Nested on one thread: the same gate, the inner scope votes. + assert jit_fn.warmup is gate + assert jit_fn.warmup(1, grid=(1,)) is None + assert jit_fn.warmup is gate + assert jit_fn.warmup(1, grid=(1,)) == "users" + + assert vars(jit_fn)["warmup"] is users_warmup + + +def test_patch_warmup_compiles_on_the_real_arguments(): + client = _EagerClient(warmup_vote=True) + manager = ClientManager([client]) + jit_fn = _FakeWarmupJit() + mapped: list = [] + + def real_args(fn, args, kwargs): + mapped.append((fn, args, dict(kwargs))) + return args, {**kwargs, "FN": "untraced"} + + with manager.patch_warmup(jit_fn, real_args=real_args): + jit_fn.warmup("x", grid=(1,), FN="traced", warmup=False) + + # The votes saw the call as made; only the compile got the mapping. + assert client.calls[0] == "pre_warmup" + assert mapped == [(jit_fn, ("x",), {"grid": (1,), "FN": "traced"})] + assert jit_fn.warmups == [{"grid": (1,), "FN": "untraced"}] + + +# ======== 2d: 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, _IndifferentIRClient()]) + 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 + + +def test_patch_run_loop_hook_conflict_leaves_nothing_patched(): + manager = ClientManager( + [ + _EagerClient(loop_overrider=lambda site, idx: idx), + _SiblingEagerClient(loop_overrider=lambda site, idx: idx), + ] + ) + frontend, namespace, attr = _first_op() + original = frontend.original_ops[namespace][attr] + scopes_before = len(LANG_PATCH_SCOPES.get("triton", [])) + + with pytest.raises(RuntimeError, match="Only one loop_iter overrider"): + with manager.patch_run(_dummy_lang_fn, frontend_name="triton"): + pass + + assert getattr(namespace, attr) is original + assert frontend._patch_calls_scope == 0 + assert not frontend._loop_ast_patched + assert len(LANG_PATCH_SCOPES.get("triton", [])) == scopes_before + assert manager._iter_overrider is None + + +# ======== 2e: interpreter callbacks ========= + + +def test_interpreter_callbacks_reach_only_interpreting_clients(): + eager = _EagerClient() + manager = ClientManager([eager, _IndifferentIRClient()]) + 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([_IndifferentIRClient()]) + + assert manager.pre_run_callback(_dummy_lang_fn) is True + assert manager.post_run_callback(_dummy_lang_fn) is True + + +# ======== 2f-2g: finalize, begin/abort ========= + + +def test_finalize_runs_every_client_and_reraises_first_exception(): + class _Exiting(_EagerClient): + NAME = "exiting" + + def finalize(self): + super().finalize() + raise SystemExit(3) + + class _Failing(_SiblingEagerClient): + NAME = "failing" + + def finalize(self): + self.finalized = True + raise ValueError("second failure") + + record = object() + exiting, healthy, failing = ( + _Exiting(), + _OtherEagerClient(records=[record]), + _Failing(), + ) + manager = ClientManager([exiting, healthy, failing]) + + with pytest.raises(SystemExit) as info: + manager.finalize() + + assert info.value.code == 3 + assert "finalize" in healthy.calls + assert failing.finalized + assert manager.launch.records == [record] + + +def test_begin_and_abort_fan_out_to_every_client(): + log: list = [] + eager, ir = _EagerClient(), _IndifferentIRClient(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(), _IndifferentIRClient(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(), _IndifferentIRClient(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([_IndifferentIRClient(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(_IndifferentIRClient): + 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"] + + +# ======== 2h: ir_capture ========= + + +def test_ir_capture_skip_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) + assert log == ["compile", "before", "after"] + # One host compile, for the default target (D26), through the stages + # the IR client declares. + target = default_ir_target() + assert _fake_host_compile == [(jit_fn, target, frozenset({"ttir"}))] + (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.launched is False + assert event.specialization == "hash-4" + assert "ttir" in event.kernel.asm + # D5 surface, and no IR event for the interpreting peer. The binding + # never adds tensors (D23): 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_run_policy_launches_after_before_launch(): + log: list = [] + ir = _RunIRClient(log) + manager = ClientManager([ir]) + jit_fn = _FakeJit(log) + + with manager.ir_capture(jit_fn): + ret = jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=False) + + assert ret == "launched" + assert log == ["compile", "before", "launch", "after"] + assert ir.events[0].launched is True + assert ir.events[0].resolved_grid == (1, 1, 1) + assert dict(ir.events[0].bound_args)["BLOCK"] == 4 # default applied + + +def test_ir_capture_warmup_call_never_launches(): + log: list = [] + manager = ClientManager([_RunIRClient(log)]) + jit_fn = _FakeJit(log) + + with manager.ir_capture(jit_fn): + jit_fn.run(torch.zeros(4), 4, grid=None, warmup=True) + + assert log == ["compile", "before", "after"] + assert manager.get_client("ir_run").events[0].launched is False + assert manager.get_client("ir_run").events[0].resolved_grid is None + + +def test_ir_capture_restores_on_error_and_does_not_double_wrap(): + log: list = [] + manager = ClientManager([_RunIRClient(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=False) + + # before_launch raised: no launch, no after_launch, wrapper removed. + assert log == ["compile", "before"] + assert "run" not in vars(jit_fn) + + +def test_a_failing_host_compile_never_stops_a_real_launch(): + """The host compile is not the device's: a launching call still + launches, and the JIT's own compile decides (D25).""" + log: list = [] + skip = ClientManager([_SkipIRClient(log)]) + jit_fn = _FakeJit(log, compile_error=_static_assert_failure()) + with skip.ir_capture(jit_fn): + assert jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=False) is None + assert log == ["compile", ("compile_failed", CompileTimeAssertionFailure)] + + log.clear() + run = ClientManager([_RunIRClient(log)]) + jit_fn = _FakeJit(log, compile_error=_static_assert_failure()) + with run.ir_capture(jit_fn): + assert jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=False) == "launched" + assert log == [ + "compile", + ("compile_failed", CompileTimeAssertionFailure), + "launch", + ] + + +def test_ir_capture_delivers_each_specialization_once_per_launch_mode(): + log: list = [] + ir = _RunIRClient(log) + manager = ClientManager([ir]) + jit_fn = _FakeJit(log) + x = torch.zeros(4) + + with manager.ir_capture(jit_fn, compile_only=True): + jit_fn.run(x, 4, grid=(1,), warmup=True, BLOCK=4) + with manager.ir_capture(jit_fn): + # e.g. an autotuner benchmarking two configs, then launching one. + for block in (4, 4, 8, 4): + jit_fn.run(x, 4, grid=(1,), warmup=False, BLOCK=block) + + assert [(e.specialization, e.launched) for e in ir.events] == [ + ("hash-4", False), + ("hash-4", True), + ("hash-8", True), + ] + assert log.count("launch") == 4 # every call still launched + + # 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=False, BLOCK=4) + assert [(e.specialization, e.launched) for e in ir.events] == [("hash-4", True)] + + +class _Opaque: + """A value the binding fingerprint knows nothing about (weakref-able).""" + + +def test_ir_capture_delivers_each_binding_of_a_specialization(): + """D22: 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, compile_only=True): + 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, compile_only=True): + 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, compile_only=True): + 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_compile_only_window_reports_failures_as_data(): + log: list = [] + ir = _RunIRClient(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, compile_only=True) 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 and failure.launched is False + assert failure.target == default_ir_target() + assert dict(failure.bound_args) == {"x_ptr": x, "n": 4, "BLOCK": 4} + + +def test_a_kernel_the_device_could_not_load_is_still_delivered(): + """D25: nothing is loaded, so a config too big for some device (the + fake's _init_handles would raise) is analyzed like any other.""" + ir = _SkipIRClient() + manager = ClientManager([ir]) + jit_fn = _FakeJit() + with manager.ir_capture(jit_fn, compile_only=True) as window: + kernel = jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=True) + assert window.compiled == 1 and ir.failures == [] + (event,) = ir.events + assert event.kernel is kernel and event.specialization == "hash-4" + + +def test_a_failing_config_is_reported_once_per_launch(): + """A launch window's benchmark call of a config the compile-only pass + already reported is no news; another call, or the next launch, is.""" + log: list = [] + ir = _RunIRClient(log) + manager = ClientManager([ir]) + broken = _FakeJit(log, compile_error=_static_assert_failure()) + x = torch.zeros(4) + + with manager.ir_capture(broken, compile_only=True) as window: + broken.run(x, 4, grid=(1,), warmup=True, BLOCK=8) + with manager.ir_capture(broken): + for block in (8, 8, 16): + assert broken.run(x, 4, grid=(1,), warmup=False, BLOCK=block) == "launched" + + assert [dict(f.kwargs)["BLOCK"] for f in ir.failures] == [8, 16] + assert len(window.failures) == 1 + assert log.count("launch") == 3 + manager.begin_launch(_call(capture=True)) + with manager.ir_capture(broken, compile_only=True): + broken.run(x, 4, grid=(1,), warmup=True, BLOCK=8) + assert [dict(f.kwargs)["BLOCK"] for f in ir.failures] == [8] + + +class _TTGIRClient(_SkipIRClient): + NAME = "ir_ttgir" + IR_STAGES = frozenset({"ttgir"}) + + +class _HipIRClient(_SkipIRClient): + NAME = "ir_hip" + + def __init__(self, log=None): + super().__init__(log) + self.ir_target = "hip:gfx942" + + +def test_each_target_compiles_once_and_reaches_only_its_clients(_fake_host_compile): + """D26: one host compile per distinct target, through the latest stage + its clients declare; each client sees only its own target's events.""" + from triton.backends.compiler import GPUTarget + + cuda89, gfx942 = GPUTarget("cuda", 89, 32), GPUTarget("hip", "gfx942", 64) + ttir, hip, ttgir = _SkipIRClient(), _HipIRClient(), _TTGIRClient() + manager = ClientManager([ttir, hip, ttgir]) + jit_fn = _FakeJit(compile_error=None) + + with manager.ir_capture(jit_fn, compile_only=True): + jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=True) + + assert _fake_host_compile == [ + (jit_fn, cuda89, frozenset({"ttir", "ttgir"})), + (jit_fn, gfx942, frozenset({"ttir"})), + ] + assert [e.target for e in ttir.events] == [cuda89] + assert [e.target for e in ttgir.events] == [cuda89] + assert ttir.events[0] is ttgir.events[0] + assert [e.target for e in hip.events] == [gfx942] + + # A target set explicitly to the default's value shares its compile. + _fake_host_compile.clear() + hip.ir_target = cuda89 + manager.begin_launch(_call(capture=True)) + with manager.ir_capture(jit_fn, compile_only=True): + jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=True) + assert [t for _, t, _ in _fake_host_compile] == [cuda89] + assert [e.target for e in hip.events] == [cuda89] + + +def test_instances_of_one_ir_client_class_keep_their_own_targets( + _fake_host_compile, +): + """Stacked traces of one IR client class, one instance per target, keep + both instances: each target compiles once and reaches only its own.""" + from triton.backends.compiler import GPUTarget + + cuda80, cuda90 = GPUTarget("cuda", 80, 32), GPUTarget("cuda", 90, 32) + sm80, sm90 = _SkipIRClient(), _SkipIRClient() + sm80.ir_target, sm90.ir_target = "cuda:80", "cuda:90" + traced = tilelens.trace(sm90)(tilelens.trace(sm80)(_make_plain_kernel())) + manager = traced.client_manager + assert manager.clients == [sm80, sm90] + jit_fn = _FakeJit(compile_error=None) + + with manager.ir_capture(jit_fn, compile_only=True): + jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=True) + + assert [t for _, t, _ in _fake_host_compile] == [cuda80, cuda90] + assert [e.target for e in sm80.events] == [cuda80] + assert [e.target for e in sm90.events] == [cuda90] + + +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, compile_only=True): + 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, compile_only=True): + pass + assert len(_fake_host_compile) == 1 and ir.failures == [] + + +class _MisspelledStageClient(_SkipIRClient): + NAME = "ir_misspelled" + IR_STAGES = frozenset({"TTIR"}) + + +def test_an_ir_stage_no_kernel_holds_is_refused_before_compiling( + _fake_host_compile, +): + """An IR_STAGES name the target's kernels never hold is the client's + bug: a ValueError before anything compiles, never a silent compile of + the whole pipeline or a compile failure.""" + from triton.backends.compiler import GPUTarget + + manager = ClientManager([_MisspelledStageClient()]) + jit_fn = _FakeJit() + with pytest.raises( + ValueError, match=r"_MisspelledStageClient.IR_STAGES: .*\['TTIR'\]" + ): + with manager.ir_capture(jit_fn, compile_only=True): + pass + assert _fake_host_compile == [] + # Names are checked against the client's own target: "sass" is CUDA's. + manager.compiler.check_stages(GPUTarget("cuda", 80, 32), {"sass"}) + with pytest.raises(ValueError, match=r"\['sass'\] are no stage .*hip:gfx942"): + manager.compiler.check_stages(GPUTarget("hip", "gfx942", 64), {"sass"}) + + +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_refuses_another_owner_and_ignores_other_threads(): + log: list = [] + first = ClientManager([_SkipIRClient(log)]) + second = ClientManager([_RunIRClient()]) + jit_fn = _FakeJit(log) + results: list = [] + + with first.ir_capture(jit_fn): + with pytest.raises(RuntimeError, match="already being captured"): + with second.ir_capture(jit_fn): + pass + worker = threading.Thread( + target=lambda: results.append( + jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=False) + ) + ) + worker.start() + worker.join() + + # The other thread's call went straight to the original run. + assert results == ["launched"] + assert log == ["launch"] + + +def test_ir_capture_compiles_and_launches_on_the_real_arguments(): + ir = _RunIRClient() + manager = ClientManager([ir]) + received: list = [] + + class _RecordingJit(_FakeJit): + def run(self, *args, grid, warmup, **kwargs): + received.append((warmup, args, dict(kwargs))) + return super().run(*args, grid=grid, warmup=warmup, **kwargs) + + def fake_compile(self, *args, **kwargs): + received.append((True, 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=False) + + assert [(w, a[1], k) for w, a, k in received] == [ + (True, 8, {"BLOCK": 8}), + (False, 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_real_compiles_refuse_while_the_language_is_patched(patched_language): + log: list = [] + ir = _RunIRClient(log) + manager = ClientManager([ir]) + jit_fn = _FakeJit(log) + x = torch.zeros(4) + + # A compile reports the refusal as data (no compile ran: its own type, + # a RuntimeError) ... + with manager.ir_capture(jit_fn, compile_only=True) as window: + assert jit_fn.run(x, 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) + # ... a real launch raises it (the same call's compile failure is not + # reported again), and so does a voted warmup. + with manager.ir_capture(jit_fn): + with pytest.raises(RuntimeError, match="language patched"): + jit_fn.run(x, 4, grid=(1,), warmup=False) + assert len(ir.failures) == 1 + warmup_jit = _FakeWarmupJit() + with ClientManager([_EagerClient(warmup_vote=True)]).patch_warmup(warmup_jit): + with pytest.raises(RuntimeError, match="language patched"): + warmup_jit.warmup(x, grid=(1,)) + + # Nothing was compiled or launched. + assert log == [("compile_failed", LanguagePatchedError)] + assert warmup_jit.warmups == [] + + +def test_an_ir_only_launch_raises_the_patched_language_refusal( + fake_compile, patched_language +): + """No compile's outcome, so not a compile failure D27 lets a launch + survive: the IR-only launch fails, as it did before D27 (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 + + +# ======== 3a: runner chain rebuild ========= + + +def test_trace_does_not_mutate_the_users_autotuner_chain(): + user = _make_autotuned_kernel(restore_value=["out_ptr"]) + heuristics, jit_fn = user.fn, user.fn.fn + before = dict(vars(user)) + before_heuristics = dict(vars(heuristics)) + + traced = tilelens.trace(_EagerClient())(user) + + assert vars(user).keys() == before.keys() + assert all(vars(user)[k] is v for k, v in before.items()) + assert all(vars(heuristics)[k] is v for k, v in before_heuristics.items()) + assert vars(heuristics).keys() == before_heuristics.keys() + + # Interpreter chain: copies of both layers over the InterpretedFunction. + runner = traced.runner + assert isinstance(runner, Autotuner) and runner is not user + assert isinstance(runner.fn, Heuristics) and runner.fn is not heuristics + assert isinstance(runner.fn.fn, InterpretedFunction) + assert runner.fn.fn is traced.interpreted_fn + assert runner._do_bench is KernelTraceSupport.dummy_benchmarker + + # Real chain: copies of both layers over the user's JITFunction. + real = traced.warmup_runner + assert isinstance(real, Autotuner) and real is not user and real is not runner + assert isinstance(real.fn, Heuristics) and real.fn is not heuristics + assert real.fn.fn is jit_fn is traced.jit_fn + assert real._do_bench is user._do_bench + + # Per-run state is private to each copy. + assert len({id(user.cache), id(runner.cache), id(real.cache)}) == 3 + assert runner.cache_results is False and real.cache_results is False + + +def test_rebuilt_autotuner_restore_hooks_bind_to_the_copy(): + user = _make_autotuned_kernel(restore_value=["out_ptr"]) + traced = tilelens.trace(_EagerClient())(user) + out = torch.ones(2) + nargs = {"out_ptr": out} + + traced.runner.pre_hook(nargs) + out.zero_() + traced.runner.post_hook(nargs, exception=None) + + assert torch.equal(out, torch.ones(2)) + assert "restore_copies" in vars(traced.runner) + assert "restore_copies" not in vars(user) + + +def test_trace_refuses_an_autotuner_whose_default_hooks_it_cannot_isolate(): + user = _make_autotuned_kernel(restore_value=["out_ptr"]) + # As if Triton's default hook were a bound method of the Autotuner: the + # closure rebinding cannot point it at the traced copy. + user.pre_hook = types.MethodType(lambda self, kwargs, reset_only=False: 0, user) + + with pytest.raises(RuntimeError, match="could not be isolated"): + tilelens.trace(_EagerClient())(user) + + +def test_interpreter_copy_drops_a_benchmarker_cached_on_the_user_autotuner(): + user = _make_autotuned_kernel() + sentinel = object() + user.__dict__["do_bench"] = sentinel + + traced = tilelens.trace(_EagerClient())(user) + + assert traced.runner.do_bench is KernelTraceSupport.dummy_benchmarker + assert user.do_bench is sentinel + + +def test_autotune_over_heuristics_interprets_every_layer(): + # The interpreter chain used to drop the Heuristics layer under an + # Autotuner, so the heuristic constexpr never reached the kernel. + user = _make_autotuned_kernel() + traced = tilelens.trace(_EagerClient())(user) + x = torch.arange(8, dtype=torch.float32) + out = torch.zeros(8) + + traced[_grid8](x, out, 8) + + torch.testing.assert_close(out, x + 1) + assert user.cache == {} + + +def test_heuristics_warmup_reaches_the_warmup_votes(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)) + + client = _EagerClient(warmup_vote=True) + traced = tilelens.trace(client)(heur_kernel) + calls = fake_compile(traced.jit_fn) + + ret = traced.warmup(torch.zeros(4), torch.zeros(4), 4, grid=(1,)) + + assert isinstance(ret, _FakeKernel) + assert [c.warmup for c in calls] == [True] + assert calls[0].kwargs["BLOCK"] == 4 + assert client.calls[0] == "pre_warmup" + assert client.calls[1][0] == "post_warmup" + + +# ======== 3b: TritonTrace.run lifecycle ========= + + +def test_ir_only_skip_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 + assert not any(e.launched for e in events) + (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) + + +def test_ir_only_run_launches_through_the_real_runner(fake_compile): + ir = _RunIRClient() + traced = tilelens.trace(ir)(_make_plain_kernel()) + calls = fake_compile(traced.jit_fn) + + ret = traced[(2,)](torch.zeros(8), torch.zeros(8), 8, BLOCK=4) + + assert isinstance(ret, _FakeKernel) + # Compile-only pass, then the launch window's host compile and the real + # launch. + assert [c.warmup for c in calls] == [True, True, False] + (events,) = ir.finalized + assert [e.launched for e in events] == [False, True] + assert {e.specialization for e in events} == {ret.hash} + assert events[1].resolved_grid == (2, 1, 1) + assert events[1].grid == (2,) + + +def test_run_policy_reports_every_config_on_every_launch(fake_compile): + ir = _RunIRClient() + user = _make_autotuned_kernel(do_bench=_fake_bench) + traced = tilelens.trace(ir)(user) + fake_compile(traced.jit_fn) + x, out = torch.zeros(8), torch.zeros(8) + + traced[_grid](x, out, 8) + traced[_grid](x, out, 8) + + first, second = ir.finalized + # Independent of the autotune cache and of benchmark timing (D3). + for events in (first, second): + assert [e.kwargs["BLOCK"] for e in events if not e.launched] == [4, 8] + # Real launches: benchmarking launches each config once; the second, + # cached launch only the winner. + assert sorted(e.kwargs["BLOCK"] for e in first if e.launched) == [4, 8] + assert [e.kwargs["BLOCK"] for e in second if e.launched] == [4] + assert user.cache == {} + + +@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): + """D27: 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): + """D27 for a plain kernel: its only config failed, so 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_under_run_the_device_decides_after_a_failed_host_compile(fake_compile): + """Under "run" a failed host compile does not stop the real launch, + which compiles its own kernel for the device (here the fake device's + compile succeeds); the failure is reported once.""" + ir = _RunIRClient() + traced = tilelens.trace(ir)(_make_plain_kernel()) + calls = fake_compile( + traced.jit_fn, compile_error=lambda kw: ValueError("not for this target") + ) + + ret = traced[(2,)](torch.zeros(8), torch.zeros(8), 8, BLOCK=4) + + assert isinstance(ret, _FakeKernel) + # Compile-only pass, the launch window's host compile, the real launch. + assert [c.warmup for c in calls] == [True, True, False] + assert ir.finalized == [[]] + (failure,) = ir.failures + assert isinstance(failure.error, ValueError) and not failure.launched + + +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 compile-only event, no real launch. + assert [c.warmup for c in calls] == [True] + (events,) = ir.finalized + assert [e.launched for e in events] == [False] + # 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) + + +# ======== D28: 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("compile_only", [True, False]) +@pytest.mark.parametrize("ir_cls", [_SkipIRClient, _RunIRClient]) +def test_ir_capture_raises_a_call_that_does_not_bind(ir_cls, compile_only): + """A bind failure is the call's own error, which JITFunction.run raises + on any device: ir_capture raises that very exception, compile-only or + not, under either launch policy. No compile_failed event, no real + launch, and the capture is removed.""" + log: list = [] + manager = ClientManager([ir_cls(log)]) + unbound = _bind_failure() + jit_fn = _FakeJit(log, compile_error=unbound) + + with manager.ir_capture(jit_fn, compile_only=compile_only) as window: + with pytest.raises(TypeError) as raised: + jit_fn.run(torch.zeros(4), grid=(1,), warmup=compile_only) + + assert raised.value is unbound + assert log == ["compile"] + assert (window.compiled, window.failures) == (0, []) + assert "run" not in vars(jit_fn) + + +@pytest.mark.parametrize("ir_cls", [_SkipIRClient, _RunIRClient]) +@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, ir_cls, make, launch +): + """D28: an IR-only launch raises the bind failure (D27's "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 = ir_cls(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, +): + """D28 in a mixed trace (D4b): 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(_IndifferentIRClient(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 (D23). + 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_launch_grid_is_the_grid_the_launch_ran_with(fake_compile): + args = (torch.zeros(8), torch.zeros(8), 8) + # Run policy: benchmarking launches every config, then the winner (ties + # pick the first); Launch.grid is the winner's grid. + run = tilelens.trace(_RunIRClient())(_make_autotuned_kernel(do_bench=_fake_bench)) + calls = fake_compile(run.jit_fn) + run[_grid](*args) + assert [c.kwargs["BLOCK"] for c in calls if not c.warmup] == [4, 8, 4] + assert trace_module.launches[-1].grid == (2, 1, 1) + + # Skip policy: nothing launched; configs that disagree on the grid + # leave it open, a grid they share is the launch's. + skip = tilelens.trace(_SkipIRClient())(_make_autotuned_kernel()) + fake_compile(skip.jit_fn) + skip[_grid](*args) + assert trace_module.launches[-1].grid is None + skip[(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 (D23): 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)} + + +class _ForgetfulSkipIR(_SkipIRClient): + """Keeps nothing of a launch past its hooks; ``fail`` makes after_launch + raise (under "run" only after a real launch).""" + + NAME = "forgetful_skip" + + def __init__(self, fail=False): + super().__init__() + self.fail = fail + + def begin_launch(self, call): + self.log.append("begin") + + def before_launch(self, event): + self.log.append("before") + + def after_launch(self, event): + if self.fail and (event.launched or self.LAUNCH == "skip"): + raise RuntimeError("after_launch failed") + + def finalize(self): + self.log.append("finalize") + return [] + + +class _ForgetfulRunIR(_ForgetfulSkipIR): + NAME = "forgetful_run" + LAUNCH = "run" + + +@pytest.mark.parametrize("outcome", ["finalized", "aborted"]) +@pytest.mark.parametrize("ir_cls", [_ForgetfulSkipIR, _ForgetfulRunIR]) +@pytest.mark.parametrize( + "make", + [_make_plain_kernel, lambda: _make_autotuned_kernel(do_bench=_fake_bench)], + ids=["plain", "autotuned"], +) +def test_an_ir_only_launch_keeps_no_caller_tensor(make, ir_cls, outcome, monkeypatch): + """D23: once an IR-only launch has ended and tilelens.clear() ran, + nothing of the trace refers to the caller's tensors or to a grid + callable closing over them: not the Launch the manager keeps, its dedup + keys or last-launch grid, nor the trace's copy of the autotuner.""" + ir = ir_cls(fail=outcome == "aborted") + traced = tilelens.trace(ir)(make()) + + def run(*args, grid, warmup, **kwargs): + return _FakeKernel(kwargs.get("BLOCK", 4)) # keeps no argument + + _install_fake_run(monkeypatch, traced.jit_fn, run) + kwargs = {"BLOCK": 4} if make is _make_plain_kernel else {} + + def launch() -> list[weakref.ref]: + # The caller's references end with this frame. + x, out = torch.zeros(8), torch.zeros(8) + + def grid(meta): + return (triton.cdiv(x.numel(), meta["BLOCK"]),) + + if outcome == "aborted": + with pytest.raises(RuntimeError, match="after_launch failed"): + traced[grid](x, out, 8, **kwargs) + else: + traced[grid](x, out, 8, **kwargs) + assert not trace_module.launches[-1].tensors + assert ir.log[-1] == "finalize" + return [weakref.ref(x), weakref.ref(out)] + + refs = launch() + tilelens.clear() + gc.collect() + + assert [ref() for ref in refs] == [None, None] + + +def test_an_interrupted_benchmark_keeps_no_restore_value_clone(monkeypatch): + """D23: Autotuner._bench runs the post_hook that drops a benchmark + call's restore_value clones only for an Exception, so a Ctrl+C in the + call skips it; the trace's autotuner copy keeps neither the clones nor + the caller's tensors once the launch has ended.""" + clones: list[weakref.ref] = [] + + class _Interrupting(_ForgetfulRunIR): + def before_launch(self, event): + super().before_launch(event) + if event.launched: # a benchmark call: the user hits Ctrl+C + copies = traced.ir_runner.restore_copies + clones.extend(weakref.ref(clone) for clone in copies.values()) + raise KeyboardInterrupt + + traced = tilelens.trace(_Interrupting())( + _make_autotuned_kernel(do_bench=_fake_bench, restore_value=["out_ptr"]) + ) + + def run(*args, grid, warmup, **kwargs): + return _FakeKernel(kwargs.get("BLOCK", 4)) # keeps no argument + + _install_fake_run(monkeypatch, traced.jit_fn, run) + + def launch() -> list[weakref.ref]: + x, out = torch.zeros(8), torch.zeros(8) + with pytest.raises(KeyboardInterrupt): + traced[_grid](x, out, 8) + return [weakref.ref(x), weakref.ref(out)] + + refs = launch() + tilelens.clear() + gc.collect() + + assert len(clones) == 1 # the pre_hook's clone of out_ptr + assert [ref() for ref in refs + clones] == [None, None, None] + + +@pytest.mark.parametrize("mixed", [False, True], ids=["eager", "mixed"]) +def test_an_interpreted_launch_that_raises_keeps_no_caller_tensor(mixed, monkeypatch): + """Once an interpreted autotuned launch has raised (here in a config + pre_hook), the trace's autotuner copies keep none of the caller's + tensors: a mixed launch (D4b) is held to the IR-only guarantee (D23), + and an eager one to the same.""" + + def boom(nargs): + raise RuntimeError("pre_hook boom") + + @triton.autotune(configs=[triton.Config({"BLOCK": 4}, pre_hook=boom)], key=["n"]) + @triton.jit + def add_one_hooked(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) + + kernel: object = add_one_hooked + if mixed: + kernel = tilelens.trace(_ForgetfulSkipIR())(kernel) + traced = tilelens.trace(_EagerClient())(kernel) + + def run(*args, grid, warmup, **kwargs): + return _FakeKernel(kwargs.get("BLOCK", 4)) # keeps no argument + + _install_fake_run(monkeypatch, traced.jit_fn, run) + + def launch() -> list[weakref.ref]: + x, out = torch.zeros(8), torch.zeros(8) + with pytest.raises(RuntimeError, match="pre_hook boom"): + traced[_grid8](x, out, 8) + return [weakref.ref(x), weakref.ref(out)] + + refs = launch() + tilelens.clear() + gc.collect() + + assert traced.runner.nargs is None + assert [ref() for ref in refs] == [None, None] + + +def test_a_concurrent_launch_of_one_trace_is_refused_before_it_begins(monkeypatch): + log: list = [] + ir = _SkipIRClient(log) + traced = tilelens.trace(ir)(_make_autotuned_kernel()) + second: list = [] + + def run(*args, grid, warmup, **kwargs): + if kwargs["BLOCK"] == 8 and not second: + # Between this launch's two configs, launch again from another + # host thread. + second.append(_in_thread(traced[_grid], *args)) + return _FakeKernel(kwargs["BLOCK"]) + + _install_fake_run(monkeypatch, traced.jit_fn, run) + traced[_grid](torch.zeros(8), torch.zeros(8), 8) + + (outcome,) = second + assert isinstance(outcome["error"], RuntimeError) + assert "another host thread" in str(outcome["error"]) + # The first launch saw neither a second begin nor an abort, and kept + # every config. + assert log == ["begin", "before", "after", "before", "after", "finalize"] + (events,) = ir.finalized + assert [e.kwargs["BLOCK"] for e in events] == [4, 8] + + +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 = _RunIRClient() + traced = tilelens.trace(ir)(apply) + calls = fake_compile(traced.jit_fn) + traced[(1,)](torch.zeros(4), torch.zeros(4), 4, FN=helper, BLOCK=4) + + # Compile-only pass, the launch window's compile, the launch: all on the + # JITFunction. + assert [c.warmup for c in calls] == [True, True, False] + assert all(c.kwargs["FN"] is helper.jit_fn for c in calls) + # Events describe the call as made. + assert all(e.kwargs["FN"] is helper for e in ir.finalized[0]) + + # The interpreted launches' voted warmup compiles on them too. + voter = _EagerClient(warmup_vote=True) + traced = tilelens.trace(voter)(apply) + calls = fake_compile(traced.jit_fn) + traced[(1,)](torch.zeros(4), torch.zeros(4), 4, FN=helper, BLOCK=4) + assert [c.kwargs["FN"] for c in calls] == [helper.jit_fn] + assert voter.calls[:2] == ["begin", "pre_warmup"] + assert isinstance(voter.calls[2][1], _FakeKernel) # post_warmup + + +@pytest.mark.parametrize( + "ir_cls, expect_written", [(_SkipIRClient, False), (_RunIRClient, True)] +) +def test_ir_clients_without_a_jit_function(ir_cls, expect_written): + ir = ir_cls() + 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) + + traced[(2,)](x, out, 8, BLOCK=4) + + # No compiled kernel, so no events; "run" interprets the whole grid. + assert ir.finalized == [[]] + (call,) = ir.launch_calls + assert call.jit_fn is None and call.capture is False + if expect_written: + torch.testing.assert_close(out, x + 1) + else: + 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 a skip-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, one of them skipping: nothing is interpreted. + assert traced[(2,)](torch.zeros(4)) is None + assert log == ["begin", "finalize"] + (call,) = traced.client_manager.get_client("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)] + + +# ======== 3c: unwrapped trace globals ========= + + +def _unwrap_leaf(x): + return x + 1 + + +def _unwrap_helper(x): + return _unwrap_traced_leaf(x) # noqa: F821 + + +def test_unwrapped_trace_globals_swaps_only_what_the_kernel_reaches(): + module_globals = globals() + leaf = tilelens.trace(_SiblingEagerClient())(triton.jit(_unwrap_leaf)) + helper = tilelens.trace(_SiblingEagerClient())(triton.jit(_unwrap_helper)) + unrelated = tilelens.trace(_SiblingEagerClient())(_make_plain_kernel()) + # A package whose `api` re-exports `impl`'s binding. + pkg = types.ModuleType("tilelens_test_pkg") + pkg.api = types.ModuleType("tilelens_test_pkg.api") + pkg.impl = types.ModuleType("tilelens_test_pkg.impl") + pkg.api.helper = pkg.impl.helper = helper + kernel_globals = {"helper": helper, "pkg": pkg, "unrelated": unrelated, "keep": 1} + exec("def kernel_fn():\n return helper, pkg.api.helper\n", kernel_globals) + module_globals["_unwrap_traced_leaf"] = leaf + module_globals["_unwrap_traced_unrelated"] = unrelated + + try: + with pytest.raises(KeyError): + with _unwrapped_trace_globals(kernel_globals["kernel_fn"]): + # Direct, through a two-level module path, and transitively + # through the traced helper's own globals. + assert kernel_globals["helper"] is helper.jit_fn + assert pkg.api.helper is helper.jit_fn + assert module_globals["_unwrap_traced_leaf"] is leaf.jit_fn + # Not reachable from the kernel's code: left alone. + assert pkg.impl.helper is helper + assert kernel_globals["unrelated"] is unrelated + assert module_globals["_unwrap_traced_unrelated"] is unrelated + assert kernel_globals["keep"] == 1 + raise KeyError("restore on error") + + assert kernel_globals["helper"] is helper + assert pkg.api.helper is helper + assert module_globals["_unwrap_traced_leaf"] is leaf + finally: + module_globals.pop("_unwrap_traced_leaf", None) + module_globals.pop("_unwrap_traced_unrelated", None) + + +def test_unwrapped_trace_globals_covers_names_bound_to_traced_defaults(): + # Triton's dependency walker resolves a parameter default expression + # (``FN=helper``) in the kernel's globals; the name is not in co_names. + helper = tilelens.trace(_SiblingEagerClient())(triton.jit(_unwrap_leaf)) + kernel_globals = {"helper": helper, "alias": helper, "other": helper.jit_fn} + exec("def kernel_fn(FN=helper):\n return FN\n", kernel_globals) + + with _unwrapped_trace_globals(kernel_globals["kernel_fn"]): + assert kernel_globals["helper"] is helper.jit_fn + assert kernel_globals["alias"] is helper.jit_fn + assert kernel_globals["helper"] is kernel_globals["alias"] is helper + + +def test_untraced_call_args_unwraps_arguments_tuples_and_defaults(): + helper = tilelens.trace(_SiblingEagerClient())(triton.jit(_unwrap_leaf)) + raw = triton.jit(_unwrap_leaf) + # A trace without a JITFunction has nothing to unwrap to. + no_jit = TritonTrace(InterpretedFunction(_unwrap_leaf), _SiblingEagerClient()) + + @triton.jit + def kernel( + x_ptr, + FN: tl.constexpr, + FNS: tl.constexpr, + ACT: tl.constexpr = helper, + N: tl.constexpr = 1, + ): + pass + + x = torch.zeros(1) + args, kwargs = _untraced_call_args( + kernel, (x, helper), {"FNS": (raw, helper), "num_warps": 4} + ) + assert args[0] is x and args[1] is helper.jit_fn + assert kwargs["FNS"][0] is raw and kwargs["FNS"][1] is helper.jit_fn + # The traced default is passed explicitly; plain defaults are left alone. + assert kwargs["ACT"] is helper.jit_fn + assert "N" not in kwargs and kwargs["num_warps"] == 4 + + fns = (raw, no_jit) + args, kwargs = _untraced_call_args(kernel, (x, no_jit, fns, raw), {}) + assert args[1] is no_jit and args[2] is fns and args[3] is raw + assert kwargs == {} + + +def test_a_traced_function_refuses_to_run_outside_an_interpreted_launch(): + helper = tilelens.trace(_SiblingEagerClient())(triton.jit(_unwrap_leaf)) + program_id = tl.program_id + + # E.g. a real compile reaching the trace as a plain Python callee. + with pytest.raises(TypeError, match="outside a traced launch's interpreter"): + helper(1) + + # The interpreter never ran, so triton.language was never patched. + assert tl.program_id is program_id + + +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 (D4b). + 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/tests/unit/test_ir_version_gate.py b/tests/unit/test_ir_version_gate.py new file mode 100644 index 000000000..b7dfebb55 --- /dev/null +++ b/tests/unit/test_ir_version_gate.py @@ -0,0 +1,453 @@ +"""The D10b version gate, on every Triton release (D29). + +IR mode runs only on the Triton releases in +``tilelens.core.config.TESTED_TRITON_VERSIONS`` unless +``TILELENS_IR_ALLOW_UNTESTED_TRITON=1``. Outside that window the IR-mode test +modules skip (tests/conftest.py marks them ``ir_mode``); this module is +deliberately not one of them. It checks that the gate refuses correctly +wherever it runs, on a release in the window or not, with or without the +override in the environment: each test that expects a refusal pins the +override off (``gate``), and nothing here compiles on a release the gate +refuses, since the refusal comes first. It also checks the conftest's own +gating of the IR-mode tests. +""" + +from __future__ import annotations + +import importlib +import os +import subprocess +import sys +from pathlib import Path + +import pytest +import torch +import triton +import triton.language as tl + +import tilelens +from tilelens.clients import Sanitizer +from tilelens.core.client import Client, ClientManager, LaunchCall +from tilelens.core.config import ( + DEFAULT_IR_TARGET, + TESTED_TRITON_VERSIONS, + Config, + config as tilelens_config, + untested_triton_version, +) +from tilelens.core.host_compile import HostCompiler +from tilelens.ir import IRClient, IRVerdict + +trace_module = importlib.import_module("tilelens.core.trace") +TESTS = Path(__file__).resolve().parents[1] +REPO = TESTS.parent +OVERRIDE_VARS = ( + "TILELENS_IR_ALLOW_UNTESTED_TRITON", + "TRITON_VIZ_IR_ALLOW_UNTESTED_TRITON", +) +# A release outside the window: older than any release the window will hold. +UNTESTED = "3.5.1" + + +@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(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 +def gate(monkeypatch): + """The gate as it stands without the override, whatever the caller's + environment says; ``gate(version)`` pretends ``version`` is installed.""" + monkeypatch.setattr(tilelens_config, "ir_allow_untested_triton", False) + return lambda version: monkeypatch.setattr(triton, "__version__", version) + + +@pytest.fixture +def no_compile(monkeypatch): + """Fail any host compile: a refused launch compiles nothing.""" + + def compile(self, jit_fn, *args, **kwargs): + raise AssertionError(f"host compile of {jit_fn!r} on a refused Triton") + + monkeypatch.setattr(HostCompiler, "compile", compile) + + +@pytest.fixture +def no_driver(unreachable_driver): + """IR mode never asks Triton's driver (D25), refused or not.""" + unreachable_driver("IR mode queried Triton's driver") + + +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 _window() -> str: + return ", ".join(f"{release}.x" for release in TESTED_TRITON_VERSIONS) + + +# ======== the window ========= + + +def test_the_window_is_by_minor_release(gate, monkeypatch): + for version, expected in ( + ("3.6.0", None), + ("3.6.1+git1234abc", None), + ("3.8.0", None), + ("3.8.1", None), + ("3.5.1", "3.5.1"), + ("3.7.0", "3.7.0"), + ("3.9.0", "3.9.0"), + ("3.60.0", "3.60.0"), + ("4.0.0", "4.0.0"), + ): + gate(version) + assert untested_triton_version() == expected + monkeypatch.setattr(tilelens_config, "ir_allow_untested_triton", True) + assert untested_triton_version() is None + + +@pytest.mark.parametrize("prefix", ["TILELENS_", "TRITON_VIZ_"]) +def test_the_override_is_read_like_every_tilelens_env_flag(monkeypatch, prefix): + for name in OVERRIDE_VARS: + monkeypatch.delenv(name, raising=False) + assert Config().ir_allow_untested_triton is False + monkeypatch.setenv(f"{prefix}IR_ALLOW_UNTESTED_TRITON", "1") + assert Config().ir_allow_untested_triton is True + + +# ======== the core: nothing is captured ========= + + +class _RecordingIR(Client): + """An IR client that records the lifecycle; nothing interprets.""" + + NAME = "recording_ir" + NEEDS_INTERPRETER = False + IR_STAGES = frozenset({"ttir"}) + LAUNCH = "skip" + + def __init__(self): + super().__init__() + self.log: list = [] + self.calls: list[LaunchCall] = [] + + def begin_launch(self, call): + self.log.append("begin") + self.calls.append(call) + + def before_launch(self, event): + self.log.append("before") + + def compile_failed(self, event): + self.log.append("compile_failed") + + def finalize(self): + self.log.append("finalize") + return [] + + def pre_warmup_callback(self, jit_fn, *args, **kwargs): + return False + + def post_warmup_callback(self, jit_fn, ret): + pass + + def _unreachable(self, *args, **kwargs): + raise AssertionError("an interpreter hook reached an IR client") + + pre_run_callback = post_run_callback = arg_callback = _unreachable + grid_callback = grid_idx_callback = _unreachable + register_op_callback = register_for_loop_callback = _unreachable + + +def test_the_core_captures_nothing_on_an_untested_triton(gate, monkeypatch): + """Outside the window LaunchCall.capture is False: no host compile, no + IR event, and a skipped launch runs nothing; the override captures.""" + compiles: list = [] + + class _Kernel: + hash = "hash" + asm = {"ttir": "// ttir"} + + def compile(self, jit_fn, args, kwargs, *, target, stages=()): + compiles.append(jit_fn) + return _Kernel() + + monkeypatch.setattr(HostCompiler, "compile", compile) + gate(UNTESTED) + ir = _RecordingIR() + traced = tilelens.trace(ir)(_make_copy()) + x, out = torch.ones(8), torch.zeros(8) + + assert traced[(2,)](x, out, 8, BLOCK=4) is None + assert compiles == [] and ir.log == ["begin", "finalize"] + (call,) = ir.calls + assert call.jit_fn is traced.jit_fn and call.capture is False + assert torch.equal(out, torch.zeros(8)) + + monkeypatch.setattr(tilelens_config, "ir_allow_untested_triton", True) + traced[(2,)](x, out, 8, BLOCK=4) + assert compiles == [traced.jit_fn] and ir.calls[-1].capture is True + assert ir.log[2:] == ["begin", "before", "finalize"] + + +# ======== IR clients: refused before any analysis ========= + + +class _ToyIR(IRClient): + NAME = "toy_ir" + LAUNCH = "skip" + IR_STAGES = frozenset({"ttir"}) + + def __init__(self): + super().__init__() + self.calls: list = [] + + def analyze_launch(self, log): + self.calls.append("analyze") + return [], IRVerdict(self.NAME, "ok") + + def on_analysis_error(self, exc): + self.calls.append(("error", exc)) + return IRVerdict(self.NAME, "error", notes=[repr(exc)]) + + def on_refusal(self, refusal): + self.calls.append(("refusal", refusal)) + return IRVerdict(self.NAME, "unsupported", refusal=refusal) + + +def _toy_launch(ir): + manager = ClientManager([ir]) + call = LaunchCall(jit_fn=None, args=(), kwargs={}, grid=(1,), capture=True) + manager.begin_launch(call) + manager.finalize() + return manager + + +def test_an_ir_client_refuses_before_any_analysis(gate): + gate(UNTESTED) + ir = _ToyIR() + manager = _toy_launch(ir) + + ((hook, refusal),) = ir.calls + assert hook == "refusal" and refusal.kind == "untested-triton-version" + assert refusal.message == ( + f"IR mode is tested on Triton {_window()}, not {UNTESTED}; " + "set TILELENS_IR_ALLOW_UNTESTED_TRITON=1 to run it anyway" + ) + assert manager.launch.records == [ir.last_verdict] + assert ir.last_verdict == IRVerdict("toy_ir", "unsupported", refusal=refusal) + + +def test_an_ir_client_analyzes_under_the_override(gate, monkeypatch): + gate(UNTESTED) + monkeypatch.setattr(tilelens_config, "ir_allow_untested_triton", True) + ir = _ToyIR() + _toy_launch(ir) + assert ir.calls == ["analyze"] and ir.last_verdict.status == "ok" + + +# ======== the compiled sanitizer ========= + + +def test_the_compiled_sanitizer_refuses_an_untested_triton(gate, no_compile, no_driver): + gate(UNTESTED) + det = Sanitizer(compile=True, abort_on_error=False) + x, out = torch.ones(8), torch.zeros(8) + + assert tilelens.trace(det)(_make_copy())[(2,)](x, out, 8, BLOCK=4) is None + + verdict = det.last_verdict + assert (verdict.status, verdict.refusal.kind) == ( + "unsupported", + "untested-triton-version", + ) + assert f"not {UNTESTED}" in verdict.refusal.message + assert det.records == [] and trace_module.launches[-1].records == [verdict] + assert torch.equal(out, torch.zeros(8)) # nothing ran + + +def test_a_call_that_does_not_bind_is_not_bound_on_an_untested_triton( + gate, no_compile, no_driver +): + """D28 holds where IR mode runs: on an untested release the JIT's + binder (private API) is never reached, so a call that does not bind + the kernel's parameters is one more refused launch, not a TypeError + (the untraced call would raise it).""" + gate(UNTESTED) + det = Sanitizer(compile=True, abort_on_error=False) + x, out = torch.ones(8), torch.zeros(8) + + assert tilelens.trace(det)(_make_copy())[(2,)](x, out, BLOCK=4) is None + + assert (det.last_status, det.last_verdict.refusal.kind) == ( + "unsupported", + "untested-triton-version", + ) + + +def test_the_installed_triton_is_gated_as_the_window_says(gate, no_driver): + """The installed release, not a pretended one: refused exactly when it + is outside the window (then nothing compiles), analyzed otherwise.""" + det = Sanitizer(compile=True, abort_on_error=False) + x, out = torch.ones(8), torch.zeros(8) + + tilelens.trace(det)(_make_copy())[(2,)](x, out, 8, BLOCK=4) + + refusal = det.last_verdict.refusal + kind = None if refusal is None else refusal.kind + if untested_triton_version() is None: + assert kind != "untested-triton-version" + else: + assert kind == "untested-triton-version" + assert f"not {triton.__version__};" in refusal.message + + +_CLI_SCRIPT = """\ +import torch, triton, triton.language as tl + +triton.__version__ = {version!r} # an untested release, as far as IR mode knows + + +@triton.jit +def copy(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) # unmasked: OOB if checked + + +x, out = torch.ones(6), torch.zeros(6) +copy[(1,)](x, out, 6, BLOCK=8) +print("launch returned", out.sum().item()) +""" + + +def test_the_cli_reports_an_untested_triton_and_goes_on(tmp_path): + """tile-sanitizer --compile on an untested release: each launch is + reported as not checked, the kernel does not run, the script goes on.""" + script = tmp_path / "untested.py" + script.write_text(_CLI_SCRIPT.format(version=UNTESTED)) + cli = ( + f"import sys; sys.argv = ['tile-sanitizer', '--compile', {str(script)!r}]; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + # The override unset: the gate as it stands; this checkout first on the + # path, then whatever the caller put there (e.g. another Triton release). + unset = ("TRITON_INTERPRET", *OVERRIDE_VARS) + env = {k: v for k, v in os.environ.items() if k not in unset} + path = os.pathsep.join([str(REPO), *filter(None, [os.environ.get("PYTHONPATH")])]) + env.update(PYTHONPATH=path, CUDA_VISIBLE_DEVICES="") + proc = subprocess.run( + [sys.executable, "-c", cli], capture_output=True, text=True, env=env, cwd=REPO + ) + assert proc.returncode == 0, proc.stderr + assert proc.stdout.splitlines() == [ + "[CompiledSanitizer] not checked: untested-triton-version: IR mode is " + f"tested on Triton {_window()}, not {UNTESTED}; set " + "TILELENS_IR_ALLOW_UNTESTED_TRITON=1 to run it anyway", + "launch returned 0.0", + ] + + +# ======== the IR-mode tests' own gate (tests/conftest.py) ========= + + +@pytest.fixture +def tests_conftest(request): + """tests/conftest.py, as pytest loaded it.""" + path = TESTS / "conftest.py" + (module,) = [ + plugin + for plugin in request.config.pluginmanager.get_plugins() + if getattr(plugin, "__file__", None) and Path(plugin.__file__).resolve() == path + ] + return module + + +def test_the_ir_mode_modules_are_the_d29_list(tests_conftest): + ir_mode = [ + "unit/ir/test_host_compile.py", + "unit/ir/test_ir_capture.py", + "unit/ir/test_lowering.py", + "unit/ir/test_mlir_walk.py", + "unit/ir/test_ttir_reader.py", + "unit/ir/test_verdict_io.py", + "unit/sanitizer_compiled/test_client.py", + "unit/sanitizer_compiled/test_oob.py", + "unit/test_ir_lifecycle.py", + "end_to_end/test_ir_client.py", + "end_to_end/test_ir_lifecycle_compiled.py", + "end_to_end/test_ir_smoke.py", + "end_to_end/test_compiled_sanitizer.py", + "end_to_end/test_host_compile.py", + ] + others = [ + "unit/test_ir_version_gate.py", # this module: runs everywhere + "conformance/test_reader_conformance.py", # a non-strict xfail instead + "unit/test_client_manager.py", + "unit/test_wrapper.py", + "unit/test_sanitizer.py", + "end_to_end/test_sanitizer.py", + ] + for relative in ir_mode: + assert (TESTS / relative).is_file(), relative + assert tests_conftest.is_ir_mode_module(TESTS / relative), relative + for relative in others: + assert not tests_conftest.is_ir_mode_module(TESTS / relative), relative + + +def test_the_skip_reason_names_the_release_and_the_window( + tests_conftest, gate, monkeypatch +): + gate(UNTESTED) + assert tests_conftest.ir_mode_skip_reason() == ( + f"IR mode is not tested on the installed Triton {UNTESTED}: the tested " + "window (tilelens.core.config.TESTED_TRITON_VERSIONS) is Triton " + f"{_window()}; set TILELENS_IR_ALLOW_UNTESTED_TRITON=1 to run the " + "IR-mode tests anyway (D29)" + ) + monkeypatch.setattr(tilelens_config, "ir_allow_untested_triton", True) + assert tests_conftest.ir_mode_skip_reason() is None + monkeypatch.setattr(tilelens_config, "ir_allow_untested_triton", False) + gate(f"{TESTED_TRITON_VERSIONS[0]}.0") + assert tests_conftest.ir_mode_skip_reason() is None + + +def test_this_session_gates_the_ir_mode_tests(tests_conftest, request): + """Every IR-mode test collected with this one skips with the conftest's + reason exactly when the installed Triton (and the override as the + environment sets it) says so; this module's tests never do.""" + reason = tests_conftest.ir_mode_skip_reason() + for item in request.session.items: + skipifs = list(item.iter_markers("skipif")) + gated = [m for m in skipifs if "(D29)" in str(m.kwargs.get("reason"))] + if item.get_closest_marker(tests_conftest.IR_MODE) is None or reason is None: + assert gated == [], item.nodeid + else: + # The first skipif pytest evaluates: its reason is the one shown. + assert gated == skipifs[:1], item.nodeid + assert gated[0].args == (True,) and gated[0].kwargs == {"reason": reason} + assert request.node.get_closest_marker(tests_conftest.IR_MODE) is None diff --git a/tests/unit/test_wrapper.py b/tests/unit/test_wrapper.py index 12c478964..ab5d3e3d8 100644 --- a/tests/unit/test_wrapper.py +++ b/tests/unit/test_wrapper.py @@ -1,5 +1,7 @@ +import os import subprocess import sys +from pathlib import Path import pytest from unittest.mock import MagicMock, patch @@ -7,15 +9,20 @@ import tilelens from tilelens.core.config import config as cfg from tilelens.core.trace import TraceInterface +from tilelens.clients.sanitizer.compiled import CompiledSanitizer from tilelens.wrapper import ( + COMPILE_NOTE, create_patched_jit, create_patched_autotune, sanitizer_wrapper, + compiled_sanitizer_wrapper, profiler_wrapper, apply_sanitizer, apply_profiler, ) +REPO = Path(__file__).resolve().parents[2] + @pytest.fixture def _isolate_cli_active(): @@ -60,6 +67,24 @@ def test_sanitizer_wrapper_accepts_frontend(): assert result == "wrapped_kernel" +def test_compiled_sanitizer_wrapper_traces_with_the_compiled_sanitizer(monkeypatch): + monkeypatch.setattr(cfg, "enable_sanitizer", True) + mock_kernel = MagicMock() + mock_kernel.__name__ = "test_kernel" + + with patch("tilelens.wrapper.tilelens.trace") as mock_trace: + mock_decorator = MagicMock(return_value="wrapped_kernel") + mock_trace.return_value = mock_decorator + + result = compiled_sanitizer_wrapper(mock_kernel, frontend="gluon") + + client = mock_trace.call_args.kwargs["client"] + assert isinstance(client, CompiledSanitizer) and client.abort_on_error + assert mock_trace.call_args.kwargs["frontend"] == "gluon" + mock_decorator.assert_called_once_with(mock_kernel) + assert result == "wrapped_kernel" + + def test_profiler_wrapper_applies_trace(): mock_kernel = MagicMock() mock_kernel.__name__ = "test_kernel" @@ -256,3 +281,60 @@ def test_wrapper_imports_without_pytest(): code = "import sys; sys.modules['pytest'] = None; import tilelens.wrapper" proc = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True) assert proc.returncode == 0, proc.stderr + + +# ======== tile-sanitizer --compile (D14) =========== + +_CLIENTS_SCRIPT = """\ +import sys +import triton + + +@triton.jit +def kernel(x_ptr): + pass + + +print(sorted(c.NAME for c in kernel.client_manager.clients), sys.argv[1:]) +""" + + +def _run_cli(argv): + """Run apply_sanitizer in a subprocess (it patches triton.jit for good) + as the command ``argv[0]`` with ``argv[1:]``.""" + code = ( + f"import sys; sys.argv = {argv!r}; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + env = {k: v for k, v in os.environ.items() if k != "TRITON_INTERPRET"} + return subprocess.run( + [sys.executable, "-c", code], capture_output=True, text=True, cwd=REPO, env=env + ) + + +@pytest.mark.parametrize( + "command, before, after, clients", + [ + ("tile-sanitizer", ["--compile"], [], ["compiled_sanitizer"]), + ("triton-sanitizer", ["--compile"], ["--compile"], ["compiled_sanitizer"]), + # After the script name the flag is the script's own argument. + ("tile-sanitizer", [], ["--compile"], ["sanitizer"]), + ], +) +def test_the_compile_flag_before_the_script_selects_the_compiled_sanitizer( + tmp_path, command, before, after, clients +): + script = tmp_path / "script.py" + script.write_text(_CLIENTS_SCRIPT) + proc = _run_cli([command, *before, str(script), *after]) + assert proc.returncode == 0, proc.stderr + assert proc.stdout.strip() == f"{clients} {after}" + # Under the flag the user is told, once, that kernels do not run (D2). + assert proc.stderr.count(COMPILE_NOTE) == (1 if before else 0) + + +def test_the_compile_flag_without_a_script_prints_the_usage(): + proc = _run_cli(["tile-sanitizer", "--compile"]) + assert proc.returncode == 1 + assert "Usage: tile-sanitizer [--compile] [args...]" in proc.stdout + assert "kernels are not run" in proc.stdout diff --git a/tilelens/clients/__init__.py b/tilelens/clients/__init__.py index d4810590a..2ae7ba322 100644 --- a/tilelens/clients/__init__.py +++ b/tilelens/clients/__init__.py @@ -10,7 +10,15 @@ "OpTypeCounts": ("tilelens.clients.profiler.data", "OpTypeCounts"), "RaceDetector": ("tilelens.clients.race_detector.race_detector", "RaceDetector"), "Sanitizer": ("tilelens.clients.sanitizer.sanitizer", "Sanitizer"), + "CompiledSanitizer": ( + "tilelens.clients.sanitizer.compiled.client", + "CompiledSanitizer", + ), "OutOfBoundsRecord": ("tilelens.clients.sanitizer.data", "OutOfBoundsRecord"), + "CompiledSanitizerRecord": ( + "tilelens.clients.sanitizer.data", + "CompiledSanitizerRecord", + ), "SymbolicExpr": ("tilelens.clients.symbolic_engine", "SymbolicExpr"), "SymbolicClient": ("tilelens.clients.symbolic_engine", "SymbolicClient"), "RangeWrapper": ("tilelens.clients.symbolic_engine", "RangeWrapper"), diff --git a/tilelens/clients/sanitizer/compiled/__init__.py b/tilelens/clients/sanitizer/compiled/__init__.py new file mode 100644 index 000000000..e1ca894d7 --- /dev/null +++ b/tilelens/clients/sanitizer/compiled/__init__.py @@ -0,0 +1,42 @@ +"""The compiled-mode sanitizer (``Sanitizer(compile=True)``): out-of-bounds, +integer-overflow and division-by-zero checks over a kernel's compiled TTIR, +instantiated per launch with its LaunchBinding. + +``oob`` is the evaluator: an ``AccessGraph`` from ``tilelens.ir.ttir_reader`` +and a launch's binding in, a ``CheckResult`` out. ``client`` is the trace +client, ``CompiledSanitizer``, which runs it on every config a launch +compiled. + +Exports resolve on first access, so importing this package imports neither +Z3 nor Triton. +""" + +from __future__ import annotations + +from importlib import import_module +from typing import Any + + +_EXPORTS: dict[str, tuple[str, str]] = { + "CompiledSanitizer": ( + "tilelens.clients.sanitizer.compiled.client", + "CompiledSanitizer", + ), + "CheckResult": ("tilelens.clients.sanitizer.compiled.oob", "CheckResult"), + "Finding": ("tilelens.clients.sanitizer.compiled.oob", "Finding"), + "SanitizerKind": ("tilelens.clients.sanitizer.compiled.oob", "SanitizerKind"), + "check_graph": ("tilelens.clients.sanitizer.compiled.oob", "check_graph"), +} + +__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/clients/sanitizer/compiled/client.py b/tilelens/clients/sanitizer/compiled/client.py new file mode 100644 index 000000000..af8c404de --- /dev/null +++ b/tilelens/clients/sanitizer/compiled/client.py @@ -0,0 +1,776 @@ +"""The compiled sanitizer client (``Sanitizer(compile=True)``). + +An ``IRClient`` that reads each compiled config's TTIR and checks the +launch against it with ``oob.check_graph``, without running the kernel +(``LAUNCH = "skip"``, D2): output tensors are left untouched, and addresses +are the launch's tensor addresses. The TTIR is compiled on the host for the +client's target (D25, D26: ``target=``, else ``TILELENS_IR_TARGET``, else +``cuda:89``), so no GPU is needed and a result never depends on the machine: +CPU tensors are checked like device ones. + +Every config the launch compiled is checked, autotune benchmark configs +included (D3), each against the bindings the core delivered for it, and +gets its own ``ConfigVerdict``: configs that compile to one kernel but bind +other runtime arguments or grids are checked and reported apart (D22). + +A config that failed to compile never fails the launch (D27): the target is +this client's choice, not the machine's. What the failure means (see +_compiles_nowhere): + +- an error no target compiles past (a failing ``tl.static_assert``, a Python + construct Triton's code generator never accepts), from a compile that + never asked Triton's driver anything (nor did an earlier compile of the + kernel for the target, whose answer the kernel's code may keep): the + config never launches anywhere (Triton's autotuner drops it too), so it + is only noted; +- any other compile error may be the target's (``num_ctas > 1`` below + cuda:90, an fp8 type the target lacks, 16-bit descriptor atomic min/max + without native TMA, ...): the config may launch on the user's GPU and was + not checked, so it is unsupported, kind ``"compile-failed"``, with a + refusal naming the target and how to name another; +- a call that does not bind the kernel's parameters (a missing, extra or + misnamed argument) is no compile failure: the core raises the JIT + binder's error from the launch, as the untraced call raises it on any GPU + (D28), so it never reaches this client from a traced launch; one + delivered anyway (an event built outside the core) is unsupported, kind + ``"compile-failed"``, naming no other target, which would not help; +- a host compile that could not run at all (Triton's compile API, or a + compile that asked for a device: ``HostCompileUnavailable``; or one + refused while an interpreted traced launch has the language patched: + ``LanguagePatchedError``) says nothing about the kernel: unsupported, + kind ``"host-compile-unavailable"``. + +Each refusal and note names where the kernel failed (``file:line`` of the +kernel's code the error points at, else of its ``def``) and the innermost +error in one line; the whole error stays on the ArtifactLog's +CompileFailure. A launch no config of which compiled is unsupported, with +the first failed config's refusal (``"compile-failed"`` when every config +was only noted; its notes are printed with it). +The launch's status is the union: ``"violations"`` if any config has a +finding, else ``"unsupported"`` if any config was not (fully) checked, else +``"ok"``. Nothing is ``"ok"`` before a check says so. + +What ``"ok"`` covers (``IRVerdict.scope == "launch"``): this launch, with +the scalar arguments, grid and tensors it was called with, as compiled for +the target (a kernel's TTIR can differ between targets, e.g. for tensor +descriptors below sm90). Since no kernel +runs, its outputs are never written: a later launch whose arguments the +host computes from them (a count, an offset, a size read back) is checked +with the values the untraced program would not have used, so a program's +launches can each be "ok" while the untraced program goes out of bounds. +In a trace shared with an eager client that runs the interpreter (D4b), the +interpreted run happens before this analysis, unchecked by it. + +A check is remembered: a later launch of the same compiled kernel with the +same scalar arguments, grid, and tensor shapes, strides and element sizes +(only the addresses differ, as in a training loop) reuses it, with the +findings' addresses moved to the new tensors. + +Each finding becomes a ``CompiledSanitizerRecord``; the launch's +``IRVerdict`` follows them in ``Launch.records`` (D5). With +``abort_on_error`` every finding is printed as this client finalizes the +launch, which then raises ``SystemExit(1)``: the core still finalizes the +trace's other clients and exits after them, but the launch keeps none of +this client's records and is not added to ``tilelens.launches``. On a host +thread other than the main one the exit ends only that thread (the +process's exit status is unchanged). Without ``abort_on_error`` findings +are printed only under ``TILELENS_VERBOSE=1``. The parts of a launch that +were not checked are printed alongside, each op once per client and +compiled kernel (a kernel launched in a loop would repeat them). +""" + +from __future__ import annotations + +import sys +from collections import OrderedDict +from collections.abc import Hashable, Mapping +from dataclasses import replace +from typing import Any, ClassVar + +from ....core.client import LanguagePatchedError, LaunchCall +from ....core.config import config as cfg +from ....core.data import Load, Store +from ....core.host_compile import ( + bind_failed, + format_ir_target, + host_compile_unavailable, + parse_ir_target, + target_queried, + unknown_options, +) +from ....ir.capture import ( + ArtifactLog, + CompiledSpecialization, + CompileFailure, + ParseCache, +) +from ....ir.client import IRClient +from ....ir.launch import LaunchBinding, TensorFacts, tensor_facts +from ....ir.verdict import ConfigVerdict, IRVerdict, Refusal, SourceLocation +from ....utils.traceback_utils import location_to_traceback_info +from ..data import CompiledSanitizerRecord +from ..sanitizer import Sanitizer +from .oob import ( + CheckResult, + Finding, + SanitizerKind, + check_graph, + describe_site, + launch_key, + readdressed, +) + +# Checks remembered per client (see the module docstring), least recently +# used dropped first. +_REMEMBERED_CHECKS = 256 + + +class CompiledSanitizer(IRClient): + """Out-of-bounds, integer-width and division-by-zero checks over a + kernel's compiled TTIR. + + ``records`` (the last launch's findings), ``last_verdict`` and + ``last_status`` (its status, None before a launch is finalized) are the + compatibility view of the last launch. + """ + + NAME = "compiled_sanitizer" + IR_STAGES: ClassVar[frozenset[str]] = frozenset({"ttir"}) + LAUNCH = "skip" + + def __init__( + self, + abort_on_error: bool = True, + timeout_ms: int = 10_000, + target: Any = None, + ) -> None: + """``target``: what the kernels are compiled for, e.g. ``"cuda:89"``, + ``"cuda:90"``, ``"hip:gfx942"`` or a triton ``GPUTarget`` (see + tilelens.core.host_compile.parse_ir_target); None for the configured + default (``TILELENS_IR_TARGET``, else ``"cuda:89"``). A spec that + names no target raises ValueError here.""" + super().__init__() + # Parsed now, so a bad spec fails at construction, not at a launch. + self.ir_target = None if target is None else parse_ir_target(target) + if ( + isinstance(timeout_ms, bool) + or not isinstance(timeout_ms, int) + or timeout_ms <= 0 + ): + # Z3 reads a timeout of 0 or below as none at all. + raise ValueError(f"timeout_ms must be a positive int, not {timeout_ms!r}") + self.abort_on_error = abort_on_error + # Bounds each Z3 query (D11). + self.timeout_ms = timeout_ms + # Kept across launches: a TTIR text is read once, and a refusal keeps + # its kind on every later hit. + self.parses = ParseCache() + self.records: list[CompiledSanitizerRecord] = [] + # Remembered checks: (id(graph), timeout, launch_key) -> (graph, + # result); holding the graph keeps its id unique. + self._checks: OrderedDict[Hashable, tuple[Any, CheckResult]] = OrderedDict() + # The not-checked parts already printed (see print_unchecked). + self._printed: set[Hashable] = set() + + @property + def last_status(self) -> str | None: + verdict = self.last_verdict + return None if verdict is None else verdict.status + + # ── launch lifecycle ───────────────────────────────────────────── + + def begin_launch(self, call: LaunchCall) -> None: + super().begin_launch(call) + self.records = [] + + def finalize(self) -> list: + out = super().finalize() + verdict = self.last_verdict + assert verdict is not None + if self.abort_on_error or cfg.verbose: + for record in self.records: + print_compiled_record(record) + print_unchecked(verdict, self._printed) + if self.abort_on_error and self.records: + sys.exit(1) + return out + + # ── the analysis ───────────────────────────────────────────────── + + def analyze_launch(self, log: ArtifactLog) -> tuple[list, IRVerdict]: + # Each failed config (see the module docstring): unsupported, unless + # it compiles for no target, which is only noted. + failed: list[ConfigVerdict] = [] + notes: list[str] = [] + for failure in log.failures: + if _compiles_nowhere(failure.error): + notes.append(_nowhere_note(failure)) + else: + failed.append( + ConfigVerdict( + None, + _saveable(failure.config), + "unsupported", + _failure_refusal(failure), + ) + ) + specializations = log.specializations + if not specializations: + if failed: + refusal = failed[0].refusal + elif log.failures: + refusal = _none_compiled(log) + else: + refusal = Refusal( + SanitizerKind.NO_COMPILED_KERNEL, _nothing_compiled(log) + ) + return [], IRVerdict( + self.NAME, + "unsupported", + refusal=refusal, + per_config=tuple(failed), + notes=tuple(notes), + ) + records: list[CompiledSanitizerRecord] = [] + per_config: list[ConfigVerdict] = [] + for spec in specializations: + for found, verdict in self._check_specialization(spec): + records += found + per_config.append(verdict) + per_config += failed + statuses = {verdict.status for verdict in per_config} + status = next(s for s in ("violations", "unsupported", "ok") if s in statuses) + # The first refusal: under "violations" too, where it says the + # findings may not be all there is. + refusal = next((c.refusal for c in per_config if c.refusal is not None), None) + self.records = records + return records, IRVerdict( + self.NAME, + status, + # A proof or a finding holds for this launch's arguments only. + scope=None if status == "unsupported" else "launch", + refusal=refusal, + per_config=tuple(per_config), + notes=tuple(notes), + ) + + def _check_specialization( + self, spec: CompiledSpecialization + ) -> list[tuple[list[CompiledSanitizerRecord], ConfigVerdict]]: + """One (records, ConfigVerdict) per config among ``spec``'s + bindings, in the order first seen (see _configs_of): each config is + checked against its own bindings only.""" + configs = _configs_of(spec) + try: + graph = self._graph_of(spec) + except Exception as exc: + # A bug in this kernel's analysis; the other kernels still count. + graph = Refusal(SanitizerKind.INTERNAL_ERROR, _describe(exc)) + if isinstance(graph, Refusal): + return [ + ([], ConfigVerdict(spec.specialization, config, "unsupported", graph)) + for config, _ in configs + ] + return [ + self._check_config(spec.specialization, graph, config, bindings) + for config, bindings in configs + ] + + def _graph_of(self, spec: CompiledSpecialization) -> Any: + """The access graph read from ``spec``'s TTIR, or the Refusal + saying why there is none.""" + text = spec.artifacts.stages.get("ttir") + if not isinstance(text, str): + why = spec.artifacts.error or "the compiled kernel holds no TTIR" + return Refusal(SanitizerKind.NO_TTIR, why) + outcome = self.parses.get(text) + if outcome.refusal is not None: + # The reader's kind and fields, the message prefixed with the + # location like the sanitizer's own refusals. + refused = Refusal.from_exception(outcome.refusal) + where = describe_site(refused.line_no, refused.loc) + return replace(refused, message=f"{where}: {refused.message}") + if outcome.error is not None: + return Refusal( + SanitizerKind.INTERNAL_ERROR, + f"the TTIR reader failed: {outcome.error}", + ) + return outcome.graph + + def _check_config( + self, + specialization: Hashable, + graph: Any, + config: dict[str, Any], + bindings: list[LaunchBinding], + ) -> tuple[list[CompiledSanitizerRecord], ConfigVerdict]: + if not bindings: + # Nothing to check against: never "ok" (the core delivers every + # compiled kernel with a binding; another producer might not). + return [], ConfigVerdict( + specialization, + config, + "unsupported", + Refusal( + SanitizerKind.INTERNAL_ERROR, + "no binding was delivered for this kernel, so nothing was checked", + ), + ) + try: + records: list[CompiledSanitizerRecord] = [] + refusal: Refusal | None = None + for binding in bindings: + result = self._check(graph, binding) + records += [ + _record(finding, graph.kernel_name, binding, config) + for finding in result.findings + ] + if refusal is None: + refusal = result.refusal + except Exception as exc: + # A bug in this config's analysis; the other configs still count. + refusal = Refusal(SanitizerKind.INTERNAL_ERROR, _describe(exc)) + return [], ConfigVerdict(specialization, config, "unsupported", refusal) + if records: + status = "violations" + else: + status = "ok" if refusal is None else "unsupported" + return records, ConfigVerdict( + specialization, config, status, refusal, n_reports=len(records) + ) + + def _check(self, graph: Any, binding: LaunchBinding) -> CheckResult: + """check_graph, remembered (see the module docstring).""" + key = (id(graph), self.timeout_ms, launch_key(binding)) + hit = self._checks.get(key) + if hit is not None and hit[0] is graph: + self._checks.move_to_end(key) + return readdressed(hit[1], graph, binding) + result = check_graph(graph, binding, timeout_ms=self.timeout_ms) + self._checks[key] = (graph, result) + if len(self._checks) > _REMEMBERED_CHECKS: + self._checks.popitem(last=False) + return result + + def on_refusal(self, refusal: Refusal) -> IRVerdict: + return IRVerdict(self.NAME, "unsupported", refusal=refusal) + + def on_analysis_error(self, exc: Exception) -> IRVerdict: + self.records = [] + return IRVerdict( + self.NAME, + "unsupported", + refusal=Refusal(SanitizerKind.INTERNAL_ERROR, _describe(exc)), + ) + + +def _describe(exc: BaseException) -> str: + return f"{type(exc).__name__}: {exc}" + + +def _nothing_compiled(log: ArtifactLog) -> str: + if log.call is not None and not log.call.capture: + return ( + "no compiled kernel to check: the launch has no JITFunction " + "(TRITON_INTERPRET=1, an InterpretedFunction runner, Gluon or NKI); " + "the eager Sanitizer() checks such launches" + ) + return "no config of the launch was compiled, so nothing was checked" + + +# How to check a kernel for another target, as a refusal or note says it. +_NAME_A_TARGET = "Sanitizer(compile=True, target=...), or TILELENS_IR_TARGET" + + +def _compiles_nowhere(error: BaseException | None) -> bool: + """Whether a config's host compile error says the config compiles for + no target, so it never launches anywhere (a note), rather than perhaps + only not for this client's target (unsupported, compile-failed; D27). + + The rule, by exception type: the error the kernel's code raised (below + the CompilationErrors Triton's code generator wraps a called @jit + helper's or builtin's error in, following ``__cause__``) is a failing + ``tl.static_assert`` (CompileTimeAssertionFailure, which Triton's + autotuner drops a config for too) or a Python construct the code + generator never accepts (UnsupportedLanguageConstruct), and the compile + never asked Triton's driver anything (tilelens.core.host_compile. + target_queried: nor did an earlier compile of the kernel for the target): + a static_assert on ``tl.target_info``, on a device query the kernel's + code caught, or in a branch taken for the target only, is the target's. + Every other error may be the target's: Triton raises a target's + refusals (``num_ctas > 1`` below sm90, an fp8 type the target lacks, + 16-bit descriptor atomic min/max without native TMA, a dot shape the + target's MMA lacks, a PTXAS failure) as plain ValueError / + AssertionError / TypeError / RuntimeError, often wrapped in a + CompilationError like an error of the kernel's own code, so none of + them is taken to hold for every target. So is an unknown error (None). + A call that does not bind (bind_failed) never launches either, but the + untraced program raises for it (so does a traced launch, D28: only an + event built outside the core delivers one), so it is no note; a host + compile that could not run (HostCompileUnavailable, LanguagePatchedError) + is no compile error at all. + """ + if error is None or host_compile_unavailable(error) is not None: + return False + from triton.compiler.errors import ( + CompilationError, + CompileTimeAssertionFailure, + UnsupportedLanguageConstruct, + ) + + target_free = (CompileTimeAssertionFailure, UnsupportedLanguageConstruct) + seen: set[int] = set() + link: BaseException = error + while ( + isinstance(link, CompilationError) + and not isinstance(link, target_free) + and link.__cause__ is not None + and id(link) not in seen + ): + seen.add(id(link)) + link = link.__cause__ + return isinstance(link, target_free) and not target_queried(error) + + +def _for_target(failure: CompileFailure) -> str: + if failure.target is None: + return "" + return f" for {format_ir_target(failure.target)}" + + +def _error_chain(error: BaseException) -> list[BaseException]: + """``error`` and the errors behind it, outermost first, as Triton's code + generator nests them: a CompilationError raised from (``__cause__``), + or ``from None`` while handling (the suppressed ``__context__``), the + error of a called @jit helper, a builtin or the kernel's own code. Only + CompilationErrors are followed.""" + from triton.compiler.errors import CompilationError + + chain = [error] + link: BaseException | None = error + while isinstance(link, CompilationError): + link = link.__cause__ or ( + link.__context__ if link.__suppress_context__ else None + ) + if link is None or any(link is seen for seen in chain): + break + chain.append(link) + return chain + + +def _summary(error: BaseException) -> str: + """One line for a compile error: the innermost error's type and the + first line of its message (a CompilationError's own message, without + the source excerpt it formats around it).""" + from triton.compiler.errors import CompilationError + + inner = _error_chain(error)[-1] + if isinstance(inner, CompilationError): + text = str(getattr(inner, "error_message", None) or "") + else: + text = str(inner) + first = next((line.strip() for line in text.splitlines() if line.strip()), "") + return f"{type(inner).__name__}: {first}" if first else type(inner).__name__ + + +def _kernel_site(jit_fn: Any, error: BaseException | None = None) -> Any: + """Where in ``jit_fn``'s source file ``error`` was raised: the line of + the kernel's own code the code generator names for it (for an error in + a called @jit helper, the call), else the kernel's ``def`` line; None + for a JITFunction without source (e.g. a stand-in).""" + fn = getattr(jit_fn, "fn", None) + code = getattr(fn, "__code__", None) + start = getattr(jit_fn, "starting_line_number", None) + raw_src = getattr(jit_fn, "raw_src", None) + if code is None or not isinstance(start, int) or not raw_src: + return None + # The line of ``def`` after the decorators, as Triton counts it. + offset = next( + (i for i, line in enumerate(raw_src) if line.strip().startswith("def ")), 0 + ) + def_line = start + offset + src = getattr(jit_fn, "src", None) + for link in [] if error is None else _error_chain(error): + lineno = getattr(getattr(link, "node", None), "lineno", None) + if src is not None and getattr(link, "src", None) == src and lineno: + # The code generator's lines count from the ``def`` line. + return SourceLocation(code.co_filename, def_line + lineno - 1) + return SourceLocation(code.co_filename, def_line) + + +def _at(site: Any) -> str: + return "" if site is None else f"{site.file}:{site.line}: " + + +def _failure_refusal(failure: CompileFailure) -> Refusal: + """A failed config's refusal: why its host compile failed (see + _compiles_nowhere), where, and how to check it for another target where + another target may compile it.""" + error = failure.error + site = _kernel_site(failure.jit_fn, error) + unavailable = None if error is None else host_compile_unavailable(error) + if unavailable is not None: + return Refusal( + SanitizerKind.HOST_COMPILE_UNAVAILABLE, + f"{_at(site)}{unavailable}", + loc=site, + ) + if isinstance(error, LanguagePatchedError): + # No compile ran, so it says nothing about the kernel or the target. + return Refusal( + SanitizerKind.HOST_COMPILE_UNAVAILABLE, f"{_at(site)}{error}", loc=site + ) + if error is not None and bind_failed(error): + # The call's own error: the untraced call raises it on any GPU, and + # a traced launch does too (D28); only an event built outside the + # core gets here. + return Refusal( + SanitizerKind.COMPILE_FAILED, + f"{_at(site)}the call does not bind to the kernel's parameters " + f"({_summary(error)}), so it was not checked; Triton raises this " + "error for the call whatever the GPU", + loc=site, + ) + why = "" if error is None else f" ({_summary(error)})" + names = unknown_options(error) + if names: + # No kernel code failed: the call names options the target's backend + # lacks. Only a target of a backend that has them can compile it. + listed = ", ".join(repr(name) for name in names) + return Refusal( + SanitizerKind.COMPILE_FAILED, + f"{_at(site)}it failed to compile{_for_target(failure)}{why}: the " + f"call passes {listed}, neither a parameter of the kernel nor a " + f"compile option{_for_target(failure)}, so it was not checked; " + "Triton raises this for the call on every GPU whose backend lacks " + "that option (a misspelled option: on every GPU); if it is another " + "backend's option, name a target of that backend to check it " + f"({_NAME_A_TARGET})", + loc=site, + ) + return Refusal( + SanitizerKind.COMPILE_FAILED, + f"{_at(site)}it failed to compile{_for_target(failure)}{why}, so it was " + "not checked; a kernel can compile for one target and fail for another, " + "so it may launch on a GPU of another kind: to check it, name a target " + f"it compiles for ({_NAME_A_TARGET})", + loc=site, + ) + + +def _nowhere_note(failure: CompileFailure) -> str: + config = f"config {_saveable(failure.config)}" if failure.config else "the kernel" + error = failure.error + why = "" if error is None else f" ({_summary(error)})" + return ( + f"{config} was not checked: {_at(_kernel_site(failure.jit_fn, error))}it " + f"failed to compile{_for_target(failure)}{why}, an error of its own code " + "whatever the target, so it never launches" + ) + + +def _none_compiled(log: ArtifactLog) -> Refusal: + targets = sorted( + {format_ir_target(f.target) for f in log.failures if f.target is not None} + ) + where = f" for {', '.join(targets)}" if targets else "" + site = _kernel_site(log.failures[0].jit_fn) if log.failures else None + return Refusal( + SanitizerKind.COMPILE_FAILED, + f"{_at(site)}no config of the launch compiled{where}, so nothing was " + "checked: each failed with an error of its own code whatever the target " + "(see the notes)", + loc=site, + ) + + +_PLAIN = (str, int, float, bool, type(None)) + + +def _tensor_of(value: Any) -> TensorFacts | None: + """``value``'s tensor facts, None if it is no (readable) tensor.""" + if not hasattr(value, "data_ptr"): + return None + try: + return tensor_facts(value) + except Exception: + return None + + +def _saveable(config: Mapping[str, Any]) -> dict[str, Any]: + """The config kwargs, each value a saved trace cannot hold put as text: + a tensor (e.g. a heuristic's view) described by its facts, never its + data, anything else (e.g. a heuristic's tl.dtype) as its repr.""" + + def describe(value: Any) -> Any: + if isinstance(value, _PLAIN): + return value + facts = _tensor_of(value) + if facts is None: + return repr(value) + return ( + f"" + ) + + return {name: describe(value) for name, value in config.items()} + + +def _config_key(config: Mapping[str, Any]) -> Hashable: + """What tells two bindings' config kwargs apart: a plain value by type + and value, a tuple item by item, a tensor by its data_ptr, shape, + strides and dtype (never its data), any other hashable value by type + and equality, anything else by identity (the bindings hold the values + while the keys are compared).""" + + def key(value: Any) -> Hashable: + if isinstance(value, float): + return (float, value.hex()) # -0.0 apart from 0.0, NaN equal + if isinstance(value, _PLAIN): + return (type(value), value) + if isinstance(value, tuple): + return (type(value), tuple(key(item) for item in value)) + facts = _tensor_of(value) + if facts is not None: + return ("tensor", facts.data_ptr, facts.shape, facts.strides, facts.dtype) + try: + hash(value) + except Exception: + return ("id", id(value)) + return (type(value), value) + + return tuple(sorted((name, key(value)) for name, value in config.items())) + + +def _configs_of( + spec: CompiledSpecialization, +) -> list[tuple[dict[str, Any], list[LaunchBinding]]]: + """``spec``'s bindings grouped by their config kwargs (see _config_key), + each group with its saveable config, in the order first seen, so the + first group's config is ``spec.config``. Configs that compile to one + kernel (e.g. differing only in a runtime int kwarg) share a + specialization but not their bindings (D22).""" + groups: dict[Hashable, tuple[dict[str, Any], list[LaunchBinding]]] = {} + for binding in spec.bindings: + key = _config_key(binding.config) + group = groups.get(key) + if group is None: + groups[key] = (_saveable(binding.config), [binding]) + else: + group[1].append(binding) + return list(groups.values()) or [(_saveable(spec.config), [])] + + +def _record( + finding: Finding, kernel_name: str, binding: LaunchBinding, config: dict[str, Any] +) -> CompiledSanitizerRecord: + loc = finding.loc + tracebacks = ( + [] + if loc is None + else [location_to_traceback_info((loc.file, loc.line, kernel_name))] + ) + detail = finding.detail + if loc is None and finding.line_no is not None: + detail = f"{detail} (TTIR line {finding.line_no})" + return CompiledSanitizerRecord( + kind=finding.kind, + # An atomic reads and writes; it is reported as the write. + op_type=Load if finding.access_kind == "load" else Store, + tensor_name=finding.base_param, + tensor_facts=binding.tensors.get(finding.base_param), + witness=dict(finding.witness), + config=config, + user_code_tracebacks=tracebacks, + violation_offset=finding.violation_offset, + violation_address=finding.violation_address, + detail=detail, + ) + + +# ─────────────────────────── reporting ─────────────────────────── + +_TITLES = { + "out-of-bounds": "Out-Of-Bounds Access Detected", + "integer-overflow": "Integer Width Overflow Detected", + "division-by-zero": "Division By Zero Detected", +} + + +def print_compiled_record(record: CompiledSanitizerRecord) -> None: + """Print one compiled sanitizer finding, in the eager report's layout.""" + rule = "=" * 60 + print(rule) + print(f"{_TITLES[record.kind]:^60}".rstrip()) + print(f"{'(compiled sanitizer)':^60}".rstrip()) + print(rule) + print(f"Operation: {record.op_type.__name__}") + print(f"Tensor Arg: {record.tensor_name}") + facts = record.tensor_facts + if facts is not None: + print( + f"Tensor Info: dtype={facts.dtype}, shape={facts.shape}, " + f"strides={facts.strides}, contiguous={facts.contiguous}" + ) + print(f"Tensor base memory address: {facts.data_ptr:#x}") + if record.config: + print(f"Config: {record.config}") + for tb in record.user_code_tracebacks: + print(f"File: {tb.filename}, Line: {tb.lineno}, in {tb.func_name}") + print(f" Code: {tb.line_of_code.strip()}") + print("-" * 60) + if record.violation_address is not None: + print( + f"Invalid access detected at address: {record.violation_address:#x} " + f"(element offset {record.violation_offset})" + ) + witness = ", ".join(f"{name}={value}" for name, value in record.witness.items()) + print(f"Witness: {witness}") + if record.detail: + print(f"Detail: {record.detail}") + print(rule) + + +def print_unchecked(verdict: IRVerdict, printed: set[Hashable]) -> None: + """Print one line per part of a launch the compiled sanitizer did not + check (each refused config, else the launch's own refusal) that is not + in ``printed`` yet, and add it there: once per specialization, kind and + op, whatever launch-specific numbers the message holds; a config that + compiled to nothing (a failed compile) once per config and message. + The launch's own refusal is followed by its notes, each once: the + refusal of a launch no config of which compiled says why in them.""" + refused = [ + (c.specialization, c.config, c.refusal) + for c in verdict.per_config + if c.refusal is not None + ] + notes: tuple[str, ...] = () + if not refused and verdict.refusal is not None: + refused = [(None, {}, verdict.refusal)] + notes = verdict.notes + for specialization, config, refusal in refused: + ident: Hashable = specialization + if specialization is None: + ident = ( + tuple(sorted((name, repr(value)) for name, value in config.items())), + refusal.message, + ) + key = (ident, refusal.kind, refusal.line_no, refusal.loc) + if key in printed: + continue + printed.add(key) + where = f" (config {config})" if config else "" + print( + f"[CompiledSanitizer] not checked{where}: " + f"{refusal.kind}: {refusal.message}" + ) + for note in notes: + if ("note", note) in printed: + continue + printed.add(("note", note)) + print(f"[CompiledSanitizer] note: {note}") + + +# A sanitizer mode like the eager one: isinstance(..., Sanitizer) holds, so +# trace()'s ENABLE_SANITIZER=0 escape hatch covers a CompiledSanitizer() too. +Sanitizer.register(CompiledSanitizer) diff --git a/tilelens/clients/sanitizer/compiled/oob.py b/tilelens/clients/sanitizer/compiled/oob.py new file mode 100644 index 000000000..a49986f24 --- /dev/null +++ b/tilelens/clients/sanitizer/compiled/oob.py @@ -0,0 +1,1155 @@ +"""The compiled sanitizer's per-launch checks over an AccessGraph. + +Given a kernel's ``AccessGraph`` (``tilelens.ir.ttir_reader``) and one +launch's ``LaunchBinding``, each access becomes a few Z3 queries over the +launch's free variables (program ids, the access's lane positions, the +loop's iteration index), with its scalar arguments and grid substituted as +constants: + +* ``out-of-bounds`` (D12): the access executes (its loop runs, its path and + mask hold) at an element offset outside the view footprint of its + tensor, the eager sanitizer's legal set: the offsets + ``sum(i_d * stride_d)`` with ``0 <= i_d < size_d``, a stride-0 dim + counting once, relative to the view's ``data_ptr``; +* ``integer-overflow`` (D9): a width obligation the access depends on can + fail, so the IR's fixed-width arithmetic is not the unbounded reading; +* ``division-by-zero`` (D21): a ``//`` or ``%`` the access reaches (from its + offset, mask or path, or from its loop's bounds) can divide by zero. + +The unbounded reading of a term is exact only where the obligations below +it hold and its divisors are non-zero. The reader's roles order that +dependence: the loop's bounds, then its increment (only when the loop +runs), then the access's path, its mask (under the path) and its offset +(under path and mask). Each role is discharged by ONE query, "some +obligation or division of this role fails", under the safety of the roles +before it and never under its own (a wrap can decide a sibling divisor, and +the reverse); the out-of-bounds query assumes every role's safety. A +wrapped offset is therefore an integer overflow, never an out-of-bounds +witness that only the unbounded reading has. + +A value in an arm of a ``Select`` (``tl.where``) matters only where the +select picks that arm, so the width obligations of the access's path, mask +and offset are discharged under the conditions of the arms their terms are +read through: a wrap the select discards is harmless (``arith`` ops wrap, +with defined results). Divisions are not: the IR computes both arms, and a +zero divisor (or ``INT_MIN // -1``) in the discarded one is still +undefined, so they are checked wherever the access evaluates them. (The +loop's own obligations are always checked unguarded.) + +Lanes: every tensor one access's terms combine has the access's shape +(TTIR broadcasts explicitly, and only size-1 dims), so all aranges along +one dim index the same position there: an arange's value is its start plus +the lane position of its dim (one position per dim and extent; an extent-1 +arange broadcast along a longer dim stays at position 0). + +Findings are exact: a SAT model is a concrete launch state, and the op a +role reports is the innermost one failing in the model (so one wrap is not +reported at the ops computed from it). A model that may be unreachable is +withheld and the access *abstains* instead (D21's list kept from #361): a +branch condition the reader could not model (``guarded``) or one reading an +atomic observation, a mask it dropped (``mask_dropped``) or one reading an +observation. Z3's ``unknown`` abstains too (D11), each query on a fresh +Solver with its own timeout, so whether a hard query ends in +``solver-unknown`` can depend on the machine's load. An access the model +cannot check at all is refused (an unbound argument, a non-positive loop +step, an address that depends on an atomic observation). Abstentions never +suppress another access's findings: a ``CheckResult`` carries both, and +only one with neither is a proof for the launch. + +Each check builds its Z3 terms in a Z3 context of its own, so checks may +run on several host threads at once (one Z3 context is not thread-safe). A +Ctrl+C that Z3 caught during a query is raised as ``KeyboardInterrupt``, +never taken for an unknown. + +Terms are lowered by the shared ``tilelens.ir.lowering`` (the reader's +operator semantics) with ``_Env`` as its leaves: the launch's constants, and +the free variables above with their range premises. Terms can be deeper +than Python's recursion limit and their generated ``==`` / ``hash`` +recurse, so the walks over them are iterative and keyed by term identity. +""" + +from __future__ import annotations + +import weakref +from collections.abc import Hashable, Iterable, Iterator, Mapping, Sequence +from dataclasses import dataclass, replace +from enum import Enum +from typing import Any, Literal, NoReturn + +from z3 import ( + And, + ArithRef, + BoolRef, + BoolVal, + Context, + Exists, + Implies, + Int, + IntVal, + ModelRef, + Or, + Solver, + Sum, + is_false, + is_int_value, + sat, + simplify, + unknown, + unsat, +) +from z3 import Not as Z3Not + +from ....ir.launch import LaunchBinding, TensorFacts +from ....ir.lowering import Lowerer, children +from ....ir.ttir_reader import ( + AccessEvent, + AccessGraph, + Arange, + Bin, + DataDep, + LoopInfo, + LoopVar, + NumPrograms, + Observed, + Param, + Pid, + Select, + WidthObligation, + mentions_observed, + width_obligations, +) + +# The one import site of the source-location type findings carry. +from ....ir.verdict import Refusal, SourceLocation +from ..data import CompiledFindingKind + + +class SanitizerKind(str, Enum): + """Why the compiled sanitizer left (part of) a launch undecided. The + sanitizer's own kinds, apart from the reader's ``TTIRKind``.""" + + # A kernel argument the access (or its loop) reads has no launch + # binding, or the launch grid is unknown (D21: never unconstrained). + MISSING_BINDING = "missing-binding" + # The loop step is not positive (or may not be) for this launch. + NON_POSITIVE_STEP = "non-positive-step" + # Z3 answered unknown (a timeout included) (D11). + SOLVER_UNKNOWN = "solver-unknown" + # A witness sits under a branch condition that is not modeled (loaded + # data, or an atomic observation read as a free value). + UNMODELABLE_CONDITION = "unmodelable-condition" + # A witness sits behind a mask that is not modeled (dropped as loaded + # data, or reading an atomic observation). + DATA_DEPENDENT_MASK = "data-dependent-mask" + # The address depends on the value an atomic observed. + OBSERVATION_IN_ADDRESS = "observation-in-address" + # A term the reader marked unmodelable (DataDep) reached a lowered + # address, mask, path or loop bound; the reader never lets it. + UNMODELED_VALUE = "unmodeled-value" + # The client's own (compiled/client.py): the launch compiled nothing to + # read (no JITFunction, or none of its configs was delivered); a config + # failed to compile for the IR target with an error that need not hold + # for another target (D27: it may launch on the user's GPU), or its call + # does not bind the kernel's parameters (only in an event built outside + # the core, which raises such a call, D28), or no config compiled at all; a + # compiled kernel holds no TTIR; the host compile could not run for a + # config (the installed Triton's compile API, a compile that asked for a + # device, or one refused while an interpreted launch has the language + # patched: no error of the kernel's, so the config may well launch); the + # analysis itself failed (a bug, contained). + NO_COMPILED_KERNEL = "no-compiled-kernel" + COMPILE_FAILED = "compile-failed" + NO_TTIR = "no-ttir" + HOST_COMPILE_UNAVAILABLE = "host-compile-unavailable" + INTERNAL_ERROR = "internal-error" + + def __str__(self) -> str: + return self.value + + def __format__(self, spec: str) -> str: + return format(self.value, spec) + + +@dataclass(frozen=True) +class Finding: + """One exact finding: its witness is a concrete state of the launch.""" + + kind: CompiledFindingKind + access_index: int # into graph.accesses + access_kind: str # "load" | "store" | "atomic_rmw" | "atomic_cas" + base_param: str + # The op the finding is about: the access (out-of-bounds), the op or + # cast whose width does not hold (integer-overflow), the division + # (division-by-zero); the access's when that op's site is unknown. + line_no: int | None + loc: SourceLocation | None + # The free variables by name: pid_; arange__, the + # value of that tl.arange at the witness lane (with _d for a dim + # of a 2-D or wider tile); iter_loop, the loop's 0-based iteration. An + # integer-overflow adds "value", the out-of-range value. + witness: Mapping[str, int] + # out-of-bounds only: the element offset (in the accessed pointer's + # elements) from the view's data_ptr, and its byte address. + violation_offset: int | None = None + violation_address: int | None = None + detail: str = "" + + __hash__ = None # type: ignore[assignment] + + def __post_init__(self) -> None: + object.__setattr__(self, "witness", dict(self.witness)) + + +@dataclass(frozen=True) +class CheckResult: + """What ``check_graph`` decided for one launch: every exact finding, + and every undecided (access index, kind), whose first one is + ``refusal``. Only a result with neither proves the launch in bounds.""" + + findings: tuple[Finding, ...] + refusal: Refusal | None + abstained: tuple[tuple[int, SanitizerKind], ...] + + __hash__ = None # type: ignore[assignment] + + +def check_graph( + graph: AccessGraph, binding: LaunchBinding, *, timeout_ms: int = 10_000 +) -> CheckResult: + """Check every access of ``graph`` under ``binding`` (see the module + docstring); ``timeout_ms`` (a positive int) bounds each Z3 query. + + A limit of the model is returned as an abstention, never raised; an + exception means a bug (e.g. a graph that breaks the reader's + invariants), or is the KeyboardInterrupt of a Ctrl+C.""" + _check_timeout(timeout_ms) + return _Check(graph, binding, timeout_ms).run() + + +def launch_key(binding: LaunchBinding) -> Hashable: + """Everything ``check_graph`` reads of ``binding`` but the tensors' + addresses: for one graph (and timeout), bindings with one key get the + same CheckResult up to the findings' ``violation_address`` (see + ``readdressed``); no message holds an address.""" + return ( + tuple(sorted(binding.params.items())), + binding.grid, + tuple( + sorted( + (name, (f.elem_size, f.numel, f.shape, f.strides, f.contiguous)) + for name, f in binding.tensors.items() + ) + ), + binding.error, + ) + + +def readdressed( + result: CheckResult, graph: AccessGraph, binding: LaunchBinding +) -> CheckResult: + """``result``, the check of ``graph`` under a binding with the + ``launch_key`` of ``binding``, with its findings' byte addresses in + ``binding``'s tensors.""" + findings = tuple( + finding + if finding.violation_offset is None + else replace( + finding, + violation_address=binding.tensors[finding.base_param].data_ptr + + finding.violation_offset + * _pointee_bytes(graph.accesses[finding.access_index].elem_bits), + ) + for finding in result.findings + ) + return CheckResult(findings, result.refusal, result.abstained) + + +def describe_site(line_no: int | None, loc: Any) -> str: + """Where a refused or reported op is: ``file:line`` of its user-source + location, else its TTIR line.""" + if loc is not None: + return f"{loc.file}:{loc.line}" + return f"TTIR line {line_no}" if line_no is not None else "the kernel" + + +def _check_timeout(timeout_ms: Any) -> None: + # Z3 reads a timeout of 0 or below as none at all. + if ( + isinstance(timeout_ms, bool) + or not isinstance(timeout_ms, int) + or timeout_ms <= 0 + ): + raise ValueError(f"timeout_ms must be a positive int, not {timeout_ms!r}") + + +# ─────────────────────────── lowering ─────────────────────────── + + +class _Refused(Exception): + """An access (or the loop) the model cannot check: becomes an + abstention; never escapes check_graph.""" + + def __init__( + self, + kind: SanitizerKind, + message: str, + line_no: int | None = None, + loc: Any = None, + ) -> None: + super().__init__(message) + self.kind = kind + self.message = message + self.line_no = line_no + self.loc = loc + + def at(self, line_no: int | None, loc: Any) -> _Refused: + if self.line_no is None and self.loc is None: + self.line_no, self.loc = line_no, loc + return self + + +def _operands(t: object, graph: AccessGraph) -> tuple: + """The nodes the checks' walks descend to from ``t``: the lowering's + ``children``, but none of a LoopVar. The loop's bounds belong to the + loop's family (checked there unguarded, assumed by every access in the + loop, kept out of its accesses' divisions by ``loop_ids``), so a bound + node an access reads through a Select arm keeps that arm's guard.""" + return () if isinstance(t, LoopVar) else children(t, graph) + + +def _walk( + roots: Iterable[object], graph: AccessGraph, seen: set[int] +) -> Iterator[object]: + """The nodes reachable from ``roots`` in pre-order, skipping (and + adding to) ``seen`` by identity. Iterative.""" + stack = [r for r in reversed(list(roots)) if r is not None] + while stack: + t = stack.pop() + if id(t) in seen: + continue + seen.add(id(t)) + yield t + stack.extend(reversed(_operands(t, graph))) + + +_DIVISIONS = frozenset({"//", "%", "u//", "u%"}) + + +def _divisions( + roots: Iterable[object], graph: AccessGraph, seen: set[int] +) -> list[Bin]: + return [ + n + for n in _walk(roots, graph, seen) + if isinstance(n, Bin) and n.op in _DIVISIONS + ] + + +def _signed(value: int, bits: int) -> int: + """The IR's signed reading of an argument of ``bits`` (an i1 stays 0/1, + the boolean model).""" + if bits <= 1: + return value + half = 1 << (bits - 1) + return (value + half) % (1 << bits) - half + + +def _lane_name(t: Arange) -> str: + # Named like the eager engine's arange variables; a tile's dim added + # (a 1-D range has dim -1). + name = f"arange_{t.start}_{t.end}" + return name if t.dim < 0 else f"{name}_d{t.dim}" + + +class _Env: + """The Z3 variables of one family of queries (one access, or the + loop's own checks), their range premises, and the family's lowering: + the env is its leaves (``tilelens.ir.lowering.TermLeaves``).""" + + def __init__( + self, + graph: AccessGraph, + binding: LaunchBinding, + grid: tuple[int, int, int], + ctx: Context, + ) -> None: + self.graph = graph + self.binding = binding + self.grid = grid + self.ctx = ctx + # Range premises of the variables created so far. + self.premises: list[BoolRef] = [] + self.pids = tuple(Int(f"pid_{axis}", ctx) for axis in range(3)) + for pid, size in zip(self.pids, grid): + self.premises += [pid >= 0, pid < size] + # (dim, extent) -> the lane's position along that dim (see arange). + self.positions: dict[tuple[int, int], ArithRef] = {} + # witness name -> the arange's value at the lane + self.lanes: dict[str, ArithRef] = {} + self.k: ArithRef | None = None # the loop's iteration index, once used + self._observed: dict[int, ArithRef] = {} + # The lowering reaches its leaves through a proxy: a strong + # reference back would be a cycle, keeping the check's Z3 context + # and terms alive until the cyclic GC frees them, on whichever host + # thread it runs, inside that thread's own Z3 call. + self.lowering = Lowerer(graph, weakref.proxy(self)) + + # ── leaves ── + + def param(self, t: Param) -> ArithRef: + if t.name not in self.binding.params: + # Raised outside an except block: a context exception's traceback + # would reach the frames holding this env (see check_loop). + raise _Refused( + SanitizerKind.MISSING_BINDING, + f"scalar argument {t.name!r} has no launch binding" + + _binding_error(self.binding), + ) + value = self.binding.params[t.name] + arg = self.graph.arg(t.name) + bits = arg.int_bits if arg is not None else 0 + return IntVal(_signed(value, bits), self.ctx) + + def pid(self, t: Pid) -> ArithRef: + return self.pids[t.axis] + + def num_programs(self, t: NumPrograms) -> ArithRef: + return IntVal(self.grid[t.axis], self.ctx) + + def arange(self, t: Arange) -> ArithRef: + """``t``'s value at the access's lane: its start plus the lane's + position along its dim, one position per dim and extent (see the + module docstring).""" + extent = t.end - t.start + key = (t.dim, extent) + pos = self.positions.get(key) + if pos is None: + pos = Int(f"lane_d{t.dim}_n{extent}", self.ctx) + self.positions[key] = pos + self.premises += [pos >= 0, pos < extent] + value = pos + t.start + self.lanes.setdefault(_lane_name(t), value) + return value + + def iteration(self, loop_ssa: str) -> ArithRef: + """``k`` of the graph's loop (it has at most one): see + ``loop_iteration``.""" + loop = self._loop() + if loop_ssa != loop.loop_ssa: + raise ValueError( + f"kernel {self.graph.kernel_name!r}: loop {loop_ssa!r} is not " + f"the graph's loop {loop.loop_ssa!r}" + ) + return self.loop_iteration() + + def observed(self, t: Observed) -> ArithRef: + """An atomic observation: a free value of the atomic's width.""" + index = t.access_index + v = self._observed.get(index) + if v is None: + v = Int(f"observed_{index}", self.ctx) + self._observed[index] = v + bits = self.graph.accesses[index].elem_bits + if bits > 1: + half = 1 << (bits - 1) + self.premises += [v >= -half, v < half] + return v + + def data_dep(self, t: DataDep) -> NoReturn: + raise _Refused(SanitizerKind.UNMODELED_VALUE, f"an unmodeled value ({t.why})") + + # ── the loop ── + + def bounds(self) -> tuple[ArithRef, ArithRef, ArithRef]: + loop = self._loop() + return self.value(loop.lower), self.value(loop.upper), self.value(loop.step) + + def loop_iteration(self) -> ArithRef: + """The loop's 0-based iteration index ``k``, with its premise ``k >= + 0 and lower + k*step < upper``: only iterations that run, and none + when the launch's trip count is zero.""" + if self.k is None: + loop = self._loop() + lower, upper, step = self.bounds() + k = Int(f"iter_{loop.loop_ssa.strip('%')}", self.ctx) + self.premises += [k >= 0, lower + k * step < upper] + self.k = k + return self.k + + def _loop(self) -> LoopInfo: + loop = self.graph.loop + if loop is None: + raise ValueError( + f"kernel {self.graph.kernel_name!r}: a loop term without a loop" + ) + return loop + + # ── terms ── + + def value(self, term: object) -> ArithRef: + return self.lowering.value(term) + + def cond(self, term: object) -> BoolRef: + return self.lowering.cond(term) + + +class _Dag: + """The nodes one family of queries reads from ``roots``: each one's + rank (children before parents) and, when built with the family's env + (whose memo holds the lowered roots), its guard: the condition under + which its value reaches a root through the arms of Selects (a node + without one reaches a root directly). Iterative, keyed by identity.""" + + def __init__( + self, roots: Sequence[object], graph: AccessGraph, env: _Env | None = None + ) -> None: + self.graph = graph + self.order: dict[int, int] = {} + nodes: list[object] = [] + stack = [(r, False) for r in reversed(roots) if r is not None] + while stack: + t, done = stack.pop() + if id(t) in self.order: + continue + if done: + self.order[id(t)] = len(nodes) + nodes.append(t) + continue + stack.append((t, True)) + stack.extend( + (k, False) + for k in reversed(_operands(t, graph)) + if id(k) not in self.order + ) + self.guards: dict[int, BoolRef] = {} + if env is None: + return + direct = {id(r) for r in roots if r is not None} + shares: dict[int, list[BoolRef]] = {} + for t in reversed(nodes): # every parent before its children + parts = shares.pop(id(t), []) + guard = None + if id(t) not in direct: + guard = parts[0] if len(parts) == 1 else Or(parts) + self.guards[id(t)] = guard + edges: list[tuple[object, BoolRef | None]] + if isinstance(t, Select): + c = env.cond(t.cond) + edges = [(t.cond, guard)] + for arm, taken in ((t.t, c), (t.f, Z3Not(c))): + edges.append((arm, taken if guard is None else And(guard, taken))) + else: + edges = [(k, guard) for k in _operands(t, graph)] + for kid, kid_guard in edges: + if kid_guard is None: + direct.add(id(kid)) + elif id(kid) not in direct: + shares.setdefault(id(kid), []).append(kid_guard) + + def rank(self, term: object) -> float: + """Children first, and a Select's condition before its arms; a term + the walk does not hold (an obligation's own, e.g. a remainder's + quotient) right after its children.""" + index = self.order.get(id(term)) + if index is not None: + return index + kids = [ + self.order[id(k)] + for k in _operands(term, self.graph) + if id(k) in self.order + ] + return max(kids) + 0.5 if kids else len(self.order) + + +# ─────────────────────────── view footprint (D12) ─────────────────────────── + + +def _in_view(e: ArithRef, facts: TensorFacts) -> BoolRef: + """``e`` (an element index from the view's data_ptr) is the offset of one + of the view's elements: ``sum(i_d * stride_d)``, ``0 <= i_d < size_d``.""" + if facts.numel == 0: + return BoolVal(False, e.ctx) + if facts.contiguous: + return And(e >= 0, e < facts.numel) + # Size-1 dims add nothing and stride-0 dims alias one element: drop both. + # (stride, size), largest stride first + dims = sorted( + ((st, sz) for sz, st in zip(facts.shape, facts.strides) if sz != 1 and st != 0), + reverse=True, + ) + if any(st < 0 for st, _ in dims): + return _any_index(e, dims) + extent = 1 # dense (a permutation of a contiguous layout): an interval + for st, sz in reversed(dims): + if st != extent: + break + extent *= sz + else: + return And(e >= 0, e < extent) + # Without overlap (each stride exceeds the reach of the smaller ones), + # the indices are the greedy quotients: quantifier-free and exact. + reach = 0 + for st, sz in reversed(dims): + if st <= reach: + return _any_index(e, dims) + reach += (sz - 1) * st + conds = [e >= 0] + rest = e + for st, sz in dims: + conds.append(rest / st < sz) + rest = rest % st + conds.append(rest == 0) + return And(conds) + + +def _any_index(e: ArithRef, dims: list[tuple[int, int]]) -> BoolRef: + """The stride equation itself, for overlapping or negative strides (as + a quantifier, Z3 may answer unknown).""" + idx = [Int(f"view_index_{d}", e.ctx) for d in range(len(dims))] + body = [i >= 0 for i in idx] + [i < sz for i, (_, sz) in zip(idx, dims)] + body.append(e == Sum([i * st for i, (st, _) in zip(idx, dims)])) + return Exists(idx, And(body)) + + +def _legal(offset: ArithRef, facts: TensorFacts, width: int) -> BoolRef: + """An access of ``width`` bytes at element ``offset`` touches only bytes + of the view's elements.""" + if width == facts.elem_size: + return _in_view(offset, facts) + # A pointer whose element width differs from the tensor's (a + # reinterpreting view): every byte it touches must be in an element. + first = offset * width + return And( + [ + And(first + j >= 0, _in_view((first + j) / facts.elem_size, facts)) + for j in range(width) + ] + ) + + +def _footprint(facts: TensorFacts) -> str: + if facts.numel == 0: + return "the empty tensor" + if facts.contiguous: + return f"the tensor's {facts.numel} elements" + return f"the view of shape {facts.shape} and strides {facts.strides}" + + +def _unusable(facts: TensorFacts) -> str | None: + if facts.elem_size <= 0 or facts.numel < 0: + return f"element size {facts.elem_size}, numel {facts.numel}" + if len(facts.shape) != len(facts.strides): + return f"shape {facts.shape} with strides {facts.strides}" + return None + + +# ─────────────────────────── the checks ─────────────────────────── + + +def _fits(value: ArithRef, bits: int, signed: bool) -> BoolRef: + if signed: + half = 1 << (bits - 1) + return And(value >= -half, value < half) + return And(value >= 0, value < (1 << bits)) + + +def _undefined_when_wide(ob: WidthObligation) -> bool: + """A signed quotient that does not fit (``INT_MIN // -1``, a + remainder's quotient included) is undefined in the IR, not a wrap: like + a zero divisor, it counts in either arm of a Select.""" + return isinstance(ob.term, Bin) and ob.term.op == "//" + + +def _is_increment(ob: WidthObligation, loop: LoopInfo, bound_ids: set[int]) -> bool: + """The loop's increment obligation (``upper - 1 + step``, which + width_obligations builds from the bound nodes themselves): the one loop + obligation that holds only when the loop runs.""" + t = ob.term + return ( + id(t) not in bound_ids + and isinstance(t, Bin) + and t.op == "+" + and t.b is loop.step + and isinstance(t.a, Bin) + and t.a.op == "-" + and t.a.a is loop.upper + ) + + +def _location(loc: Any) -> SourceLocation | None: + if loc is None or isinstance(loc, SourceLocation): + return loc + return SourceLocation(loc.file, loc.line, getattr(loc, "col", None)) + + +def _binding_error(binding: LaunchBinding) -> str: + return f" (unreadable: {binding.error})" if binding.error else "" + + +def _pointee_bytes(elem_bits: int) -> int: + return max(1, (elem_bits + 7) // 8) + + +_WITHHELD = { + SanitizerKind.UNMODELABLE_CONDITION: "under a branch condition that is not " + "modeled (loaded data, or an atomic observation read as a free value)", + SanitizerKind.DATA_DEPENDENT_MASK: "behind a mask that is not modeled " + "(loaded data, or an atomic observation read as a free value)", +} + +# What Z3 answers for a query its Ctrl+C handler cancelled (Python's own +# handler does not run during the query). +_INTERRUPTED = "interrupted from keyboard" + + +@dataclass(frozen=True, eq=False) +class _Condition: + """One width obligation or division of a role: ``safe`` holds where it + does not fail; the lowest ``rank`` is the innermost failure.""" + + kind: Literal["integer-overflow", "division-by-zero"] + site: WidthObligation | Bin + safe: BoolRef + rank: tuple + + +class _Check: + def __init__( + self, graph: AccessGraph, binding: LaunchBinding, timeout_ms: int + ) -> None: + self.graph = graph + self.binding = binding + self.timeout_ms = timeout_ms + # This check's own Z3 context (see the module docstring). + self.ctx = Context() + self.findings: list[Finding] = [] + self.abstained: list[tuple[int, SanitizerKind]] = [] + self.refusal: Refusal | None = None + # (finding kind, site) already reported: one finding per op site. + self.reported: set[tuple[str, object]] = set() + self._alive: list[object] = [] # reported sites keyed by id + # The loop's obligations and divisions, shared by its accesses. + self.loop_obs: tuple[WidthObligation, ...] = () + self.loop_divs: tuple[Bin, ...] = () + self.loop_ids: set[int] = set() + + def run(self) -> CheckResult: + graph, grid = self.graph, self.binding.grid + if grid is None: + for i, access in enumerate(graph.accesses): + self.abstain( + i, + _Refused( + SanitizerKind.MISSING_BINDING, + "the launch grid is unknown" + _binding_error(self.binding), + access.line_no, + access.loc, + ), + ) + return self.result() + in_loop = [i for i, a in enumerate(graph.accesses) if a.in_loop] + loop_refusal = None + if graph.loop is not None and in_loop: + loop_refusal = self.check_loop(in_loop[0], grid) + for i, access in enumerate(graph.accesses): + if access.in_loop and loop_refusal is not None: + self.abstain(i, loop_refusal) + continue + try: + self.check_access(i, access, grid) + except _Refused as r: + self.abstain(i, r.at(access.line_no, access.loc)) + return self.result() + + def result(self) -> CheckResult: + return CheckResult(tuple(self.findings), self.refusal, tuple(self.abstained)) + + # ── the loop, once for all of its accesses ── + + def check_loop(self, first: int, grid: tuple[int, int, int]) -> _Refused | None: + """Check the loop's step, bounds and increment; findings go to the + loop's ``first`` access. A refusal refuses every access in the loop. + No Select guards here: the loop's accesses assume its obligations + unguarded.""" + graph = self.graph + loop = graph.loop + assert loop is not None + env = _Env(graph, self.binding, grid, self.ctx) + try: + _lower, _upper, step = env.bounds() + refusal = self.step_refusal(env, step) + except _Refused as r: + # A copy that was never raised: keeping ``r`` would keep its + # traceback, which reaches this frame and so ``env``, a cycle + # leaving the check's Z3 context to the cyclic GC (see _Env). + refusal = _Refused(r.kind, r.message, r.line_no, r.loc) + if refusal is not None: + return refusal.at(loop.line_no, loop.loc) + roots = (loop.lower, loop.upper, loop.step) + dag = _Dag(roots, graph) + self.loop_ids = {id(n) for n in _walk(roots, graph, set())} + self.loop_divs = tuple(_divisions(roots, graph, set())) + self.loop_obs = tuple( + ob + for ob in width_obligations(graph, graph.accesses[first]) + if ob.role == "loop" + ) + increment = [ + ob for ob in self.loop_obs if _is_increment(ob, loop, self.loop_ids) + ] + bounds = [ + ob for ob in self.loop_obs if not _is_increment(ob, loop, self.loop_ids) + ] + # The bounds are computed whether or not the loop runs; the + # increment only matters when it does. + assumed = self.role(env, first, bounds, self.loop_divs, [], None, dag) + env.loop_iteration() + self.role(env, first, increment, (), assumed, None, dag) + return None + + def step_refusal(self, env: _Env, step: ArithRef) -> _Refused | None: + s = simplify(step) + if is_int_value(s): + if s.as_long() > 0: + return None + return _Refused( + SanitizerKind.NON_POSITIVE_STEP, + f"the loop step is {s.as_long()}; only positive steps are modeled", + ) + status, model, reason = self.solve([*env.premises, step <= 0]) + if status == unsat: + return None + if status == sat: + assert model is not None + return _Refused( + SanitizerKind.NON_POSITIVE_STEP, + f"the loop step can be {model.eval(step, model_completion=True)} " + f"(at {self.witness(env, model)}); only positive steps are modeled", + ) + return _Refused( + SanitizerKind.SOLVER_UNKNOWN, + f"Z3 could not decide whether the loop step is positive ({reason})", + ) + + # ── one access ── + + def check_access( + self, index: int, access: AccessEvent, grid: tuple[int, int, int] + ) -> None: + graph = self.graph + if mentions_observed(access.offset, graph): + # A free observation would make any address reachable. + raise _Refused( + SanitizerKind.OBSERVATION_IN_ADDRESS, + "the address depends on the value an atomic observed", + ) + facts = self.binding.tensors.get(access.base_param) + if facts is None: + raise _Refused( + SanitizerKind.MISSING_BINDING, + f"pointer argument {access.base_param!r} has no tensor binding" + + _binding_error(self.binding), + ) + problem = _unusable(facts) + if problem is not None: + raise _Refused( + SanitizerKind.MISSING_BINDING, + f"the tensor facts of {access.base_param!r} are unusable ({problem})", + ) + env = _Env(graph, self.binding, grid, self.ctx) + assumed: list[BoolRef] = [] + if access.in_loop: + # Even an offset without the induction variable executes only + # on iterations that run: none for a zero-trip loop. + env.loop_iteration() + assumed += [ + _fits(env.value(ob.term), ob.bits, ob.signed) for ob in self.loop_obs + ] + assumed += [env.value(d.b) != 0 for d in self.loop_divs] + # Lower every root first: a refusal leaves no partial result. + offset = env.value(access.offset) + path = env.cond(access.path) if access.path is not None else None + mask = env.cond(access.mask) if access.mask is not None else None + # One DAG over the three roles: a node's guard covers every root it + # reaches, so a node one role reads through a Select arm and another + # directly is checked wherever it is read. + dag = _Dag((access.path, access.mask, access.offset), graph, env) + obs: dict[str, list[WidthObligation]] = {"path": [], "mask": [], "offset": []} + for ob in width_obligations(graph, access): + if ob.role != "loop": + obs[ob.role].append(ob) + seen = set(self.loop_ids) if access.in_loop else set() + divs = { + role: _divisions((root,), graph, seen) + for role, root in ( + ("path", access.path), + ("mask", access.mask), + ("offset", access.offset), + ) + } + observed_path = access.path is not None and mentions_observed( + access.path, graph + ) + observed_mask = access.mask is not None and mentions_observed( + access.mask, graph + ) + + def uncertain(role: str) -> SanitizerKind | None: + if access.guarded or observed_path: + return SanitizerKind.UNMODELABLE_CONDITION + if role != "path" and observed_mask: + return SanitizerKind.DATA_DEPENDENT_MASK + if role == "offset" and access.mask_dropped: + return SanitizerKind.DATA_DEPENDENT_MASK + return None + + for role, condition in (("path", path), ("mask", mask), ("offset", None)): + assumed = self.role( + env, index, obs[role], divs[role], assumed, uncertain(role), dag + ) + if condition is not None: + assumed.append(condition) + self.out_of_bounds( + env, index, access, facts, offset, assumed, uncertain("offset") + ) + + def role( + self, + env: _Env, + index: int, + obs: Sequence[WidthObligation], + divs: Sequence[Bin], + assumed: list[BoolRef], + uncertain: SanitizerKind | None, + dag: _Dag, + ) -> list[BoolRef]: + """Discharge one role's width obligations and divisions under + ``assumed`` (the earlier roles' safety) in one joint query: some + condition of the role fails. None of the role's own conditions is + assumed while they are checked, only the sites reported already (so + a wrap another access reported does not resurface at the ops + computed from it). At most one finding of each kind: the model's + innermost failure, then a query for the other kind with that site + assumed. Returns ``assumed`` plus the role's safety.""" + conds: list[_Condition] = [] + for i, ob in enumerate(obs): + fits = _fits(env.value(ob.term), ob.bits, ob.signed) + guard = None if _undefined_when_wide(ob) else dag.guards.get(id(ob.term)) + conds.append( + _Condition( + "integer-overflow", + ob, + fits if guard is None else Implies(guard, fits), + # Children first (a select's condition before its arms), + # and on one term the op's own result before the reads + # of the ops that read it. + (dag.rank(ob.term), -i), + ) + ) + for d in divs: + nonzero = env.value(d.b) != 0 + conds.append(_Condition("division-by-zero", d, nonzero, (dag.rank(d), 0))) + held = [c.safe for c in conds if self.is_reported(c.kind, c.site)] + pending = [c for c in conds if not self.is_reported(c.kind, c.site)] + while pending: + status, model, reason = self.solve( + [ + *env.premises, + *assumed, + *held, + Or([Z3Not(c.safe) for c in pending]), + ] + ) + if status == unknown: + self.unknown(index, "an integer-overflow or division-by-zero", reason) + if status != sat: + break + assert model is not None + failed = [c for c in pending if is_false(model.eval(c.safe, True))] + chosen = min(failed, key=lambda c: c.rank) if failed else pending[0] + self.found(chosen, env, model, index, uncertain) + held.append(chosen.safe) + pending = [c for c in pending if c.kind != chosen.kind] + return [*assumed, *(c.safe for c in conds)] + + def out_of_bounds( + self, + env: _Env, + index: int, + access: AccessEvent, + facts: TensorFacts, + offset: ArithRef, + assumed: list[BoolRef], + uncertain: SanitizerKind | None, + ) -> None: + width = _pointee_bytes(access.elem_bits) + status, model, reason = self.solve( + [*env.premises, *assumed, Z3Not(_legal(offset, facts, width))] + ) + if status == unknown: + self.unknown(index, "the out-of-bounds", reason) + if status != sat: + return + assert model is not None + off = model.eval(offset, True).as_long() + # No address in the text: it is the same for every launch with the + # same launch_key (the address is violation_address). + detail = ( + f"{access.kind} of {access.base_param!r} at element offset {off} " + f"is outside {_footprint(facts)}" + ) + if uncertain is not None: + self.withhold(index, uncertain, "out-of-bounds", detail) + return + self.findings.append( + Finding( + kind="out-of-bounds", + access_index=index, + access_kind=access.kind, + base_param=access.base_param, + line_no=access.line_no, + loc=_location(access.loc), + witness=self.witness(env, model), + violation_offset=off, + violation_address=facts.data_ptr + off * width, + detail=detail, + ) + ) + + # ── results ── + + def found( + self, + cond: _Condition, + env: _Env, + model: ModelRef, + index: int, + uncertain: SanitizerKind | None, + ) -> None: + site = cond.site + extra: dict[str, int] = {} + if isinstance(site, WidthObligation): + value = model.eval(env.value(site.term), True).as_long() + detail = self.overflow_detail(site, value) + extra["value"] = value + else: + what = "remainder" if site.op in ("%", "u%") else "division" + detail = f"the divisor of this {what} ({site.op!r}) can be 0" + if uncertain is not None: + self.withhold(index, uncertain, cond.kind, detail) + return + self.reported.add(self.site_key(cond.kind, site)) + self._alive.append(site) + access = self.graph.accesses[index] + line_no, loc = site.line_no, site.loc + if line_no is None: + line_no, loc = access.line_no, access.loc + self.findings.append( + Finding( + kind=cond.kind, + access_index=index, + access_kind=access.kind, + base_param=access.base_param, + line_no=line_no, + loc=_location(loc), + witness={**self.witness(env, model), **extra}, + detail=detail, + ) + ) + + @staticmethod + def site_key(kind: str, site: WidthObligation | Bin) -> tuple[str, object]: + """One finding per op: by its TTIR line, else by the term's identity + (``found`` keeps a reported term alive, so its id stays unique).""" + if site.line_no is not None: + return kind, site.line_no + return kind, id(site.term if isinstance(site, WidthObligation) else site) + + def is_reported(self, kind: str, site: WidthObligation | Bin) -> bool: + return self.site_key(kind, site) in self.reported + + def withhold( + self, index: int, kind: SanitizerKind, finding: str, detail: str + ) -> None: + self.abstain( + index, + _Refused( + kind, + f"possible {finding} {_WITHHELD[kind]}, so the witness may be " + f"unreachable: {detail}", + ), + ) + + def unknown(self, index: int, query: str, reason: str | None) -> None: + self.abstain( + index, + _Refused( + SanitizerKind.SOLVER_UNKNOWN, + f"Z3 could not decide {query} query ({reason})", + ), + ) + + def abstain(self, index: int, refused: _Refused) -> None: + access = self.graph.accesses[index] + refused.at(access.line_no, access.loc) + entry = (index, refused.kind) + if entry not in self.abstained: + self.abstained.append(entry) + if self.refusal is None: + self.refusal = Refusal( + kind=refused.kind.value, + message=f"{describe_site(refused.line_no, refused.loc)}: " + f"{refused.message}", + line_no=refused.line_no, + loc=_location(refused.loc), + ) + + def overflow_detail(self, ob: WidthObligation, value: int) -> str: + if ob.signed: + half = 1 << (ob.bits - 1) + bounds = f"i{ob.bits} range [{-half}, {half})" + else: + bounds = f"unsigned i{ob.bits} range [0, {1 << ob.bits})" + # The obligation's origin (see width_obligations): an op's result, a + # trunci operand (the other signed ones), or an unsigned read. + loop = self.graph.loop + if loop is not None and _is_increment(ob, loop, self.loop_ids): + what = "the loop's induction-variable increment" + elif isinstance(ob.term, Bin) and ob.signed and ob.term.bits == ob.bits: + what = f"the result of {ob.term.op!r}" + elif ob.signed: + what = f"the operand of a truncation to i{ob.bits}" + else: + what = "a value read as unsigned" + return ( + f"{what} can be {value}, outside the {bounds}: the IR's " + "fixed-width arithmetic differs from the unbounded reading" + ) + + def witness(self, env: _Env, model: ModelRef) -> dict[str, int]: + def val(v: ArithRef) -> int: + return model.eval(v, model_completion=True).as_long() + + out = {f"pid_{axis}": val(pid) for axis, pid in enumerate(env.pids)} + for name, lane in env.lanes.items(): + out[name] = val(lane) + if env.k is not None: + out[str(env.k)] = val(env.k) + return out + + def solve(self, formulas: list[Any]) -> tuple[Any, ModelRef | None, str | None]: + """One query on a fresh Solver with its own timeout (D11; never the + process-global z3.set_param). A query Z3 cancelled for a Ctrl+C + raises KeyboardInterrupt.""" + solver = Solver(ctx=self.ctx) + solver.set(timeout=self.timeout_ms) + solver.add(*formulas) + status = solver.check() + if status == sat: + return status, solver.model(), None + if status == unknown: + reason = solver.reason_unknown() + if reason == _INTERRUPTED: + raise KeyboardInterrupt(f"Z3 query cancelled ({reason})") + return status, None, reason + return status, None, None diff --git a/tilelens/clients/sanitizer/data.py b/tilelens/clients/sanitizer/data.py index 61b68e57c..b27f0ad54 100644 --- a/tilelens/clients/sanitizer/data.py +++ b/tilelens/clients/sanitizer/data.py @@ -1,11 +1,13 @@ from ...core.data import Store, Load +import operator import numpy as np from numpy.typing import NDArray from dataclasses import dataclass -from typing import Any +from typing import Any, Literal, get_args import torch import z3 +from ...ir.launch import TensorFacts from ...utils.traceback_utils import TracebackInfo @@ -54,3 +56,54 @@ class OutOfBoundsRecordZ3(OutOfBoundsRecord): violation_address: int symbolic_expr: Any = None # Optional symbolic expression tree tensor_name: str | None = None + + +CompiledFindingKind = Literal["out-of-bounds", "integer-overflow", "division-by-zero"] + + +@dataclass +class CompiledSanitizerRecord: + """One finding of the compiled sanitizer (``Sanitizer(compile=True)``) + on one access, with a witness: an out-of-bounds address (D12's view + footprint), an address/mask/path term that can overflow its declared + integer width (D9), or one that can divide by zero (D21). + + Plain data: the tensor is described by its launch-time facts, never held, + since records outlive the launch and must not pin device memory. + """ + + kind: CompiledFindingKind + # The accessing op; an atomic (a read and a write) is reported as Store. + op_type: type[Store | Load] + # The kernel parameter of the accessed tensor, and its facts at launch; + # None when the launch bound no tensor to it (a finding in the loop's + # bounds is attributed to the loop's first access, whatever it reads). + tensor_name: str + tensor_facts: TensorFacts | None + # The free variables' values in the witness (program ids, arange lanes, + # loop iterations), by name. + witness: dict[str, int] + # The config kwargs of the config (binding) the finding is in (D3, D22). + config: dict[str, Any] + user_code_tracebacks: list[TracebackInfo] + # For "out-of-bounds": the element offset from the view's data_ptr and + # its byte address; None for the other kinds. + violation_offset: int | None = None + violation_address: int | None = None + detail: str | None = None + + def __post_init__(self) -> None: + if self.kind not in get_args(CompiledFindingKind): + raise ValueError(f"unknown compiled sanitizer finding: {self.kind!r}") + if self.op_type not in (Load, Store): + raise TypeError(f"op_type must be Load or Store, not {self.op_type!r}") + # Ints and plain containers only, so a saved trace holds the record. + self.witness = { + name: operator.index(value) for name, value in self.witness.items() + } + self.config = dict(self.config) + self.user_code_tracebacks = list(self.user_code_tracebacks) + for name in ("violation_offset", "violation_address"): + value = getattr(self, name) + if value is not None: + setattr(self, name, operator.index(value)) diff --git a/tilelens/clients/sanitizer/sanitizer.py b/tilelens/clients/sanitizer/sanitizer.py index 730d57040..d6b55232a 100644 --- a/tilelens/clients/sanitizer/sanitizer.py +++ b/tilelens/clients/sanitizer/sanitizer.py @@ -57,8 +57,10 @@ class Sanitizer(Client): """ - Factory class that returns the concrete sanitizer implementation - based on the value of ``cfg.enable_sanitizer``. + Factory class that returns the concrete sanitizer implementation: + ``Sanitizer()`` the eager one, ``Sanitizer(compile=True)`` the compiled + one (``CompiledSanitizer``, a virtual subclass); both are the + ``NullSanitizer`` while ``cfg.enable_sanitizer`` is off. """ NAME = "sanitizer" @@ -67,13 +69,23 @@ class Sanitizer(Client): def __new__(cls: type[SanitizerT], *args: Any, **kwargs: Any) -> SanitizerT: if cls is Sanitizer: + # The disable flag wins over the mode: trace() leaves a kernel + # traced with a NullSanitizer untraced, compile=True or not. + if kwargs.pop("compile", False) and cfg.enable_sanitizer: + from .compiled.client import CompiledSanitizer + + # Only a virtual Sanitizer subclass, so Python does not call + # its __init__ after __new__: call it here, without ``compile``. + compiled = object.__new__(CompiledSanitizer) + CompiledSanitizer.__init__(compiled, *args, **kwargs) + return cast(SanitizerT, compiled) target_cls = cast( type["Sanitizer"], SymbolicSanitizer if cfg.enable_sanitizer else NullSanitizer, ) - obj = object.__new__(target_cls) - cast(Any, target_cls).__init__(obj, *args, **kwargs) - return cast(SanitizerT, obj) + # A Sanitizer subclass: Python calls its __init__ once this + # returns, with the call's own arguments, ``compile`` included. + return cast(SanitizerT, object.__new__(target_cls)) return cast(SanitizerT, object.__new__(cls)) def __init__(self, abort_on_error: bool = True, *args, **kwargs): @@ -154,7 +166,14 @@ def __hash__(self) -> int: class SymbolicSanitizer(Sanitizer, SymbolicClient): - def __init__(self, abort_on_error: bool = True): + # ``compile`` is the factory's mode switch: Sanitizer(compile=False) + # reaches this __init__ with it. The eager sanitizer is never compiled. + def __init__(self, abort_on_error: bool = True, *, compile: bool = False): + if compile: + raise TypeError( + "SymbolicSanitizer is the eager sanitizer; the compiled one is " + "Sanitizer(compile=True)" + ) super().__init__(abort_on_error=abort_on_error) self.records: list[OutOfBoundsRecordZ3] = [] self.cache_args: list[Any] = [] diff --git a/tilelens/core/client.py b/tilelens/core/client.py index 940dbf506..bee44f790 100644 --- a/tilelens/core/client.py +++ b/tilelens/core/client.py @@ -1,9 +1,14 @@ -from contextlib import contextmanager, nullcontext +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,18 +23,155 @@ 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, + HostCompileUnavailable, + bind_failed, + resolve_ir_target, +) + + +LaunchPreference = Literal["skip", "run", "indifferent"] +LAUNCH_PREFERENCES: tuple[LaunchPreference, ...] = ("skip", "run", "indifferent") + +# (jit_fn, args, kwargs) -> (args, kwargs): the arguments a compile or real +# launch 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), or when the installed + # Triton is outside IR mode's tested window (``jit_fn`` is set; see + # tilelens.core.config.untested_triton_version, D10b). + 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`` (tilelens.core. + # host_compile, D25): ``.asm`` holds every stage through the latest one + # the target's IR clients declare, ``.metadata`` the compile metadata, + # ``.hash`` the specialization. Never loaded or launched; a real launch + # (``launched``) compiles its own device kernel through the JIT. None in + # a compile_failed event. + kernel: Any + # Whether a real device launch follows this event. + launched: bool + # 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 (D28, 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.""" + + # A compile-only window never launches. + compile_only: bool + # Whether calls in this window perform the real launch. + launch: bool + # Host compiles that produced a kernel (one per call and target). + compiled: int = 0 + # Host compile exceptions, in call order (one per call and target). + failures: list[BaseException] = field(default_factory=list) + + +@dataclass(frozen=True) +class CompileGroup: + """The IR clients whose kernels are compiled for one target.""" + + # A triton GPUTarget. + target: Any + # The union of the clients' IR_STAGES: the compile stops after the + # latest of them. + stages: frozenset[str] + clients: tuple["Client", ...] class Client(ABC): + # Names the client's records and its ClientManager.get_client lookup. A + # trace holds one interpreting client per NAME, since the interpreted + # run serves one; IR clients may repeat a NAME (see ir_target). 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. + # Core compiles each call, per target, through the latest stage the + # target's IR clients declare: just the front end and the "ttir" passes + # when that is all they read, the first stage when none is declared, + # and the whole pipeline (triton.compile) for its last stage or a name + # it does not produce itself ("source", "sass"); see + # tilelens.core.host_compile. + IR_STAGES: ClassVar[frozenset[str]] = frozenset() + # An IR client's real-launch preference. ClientManager.add_clients + # rejects traces whose IR clients mix "skip" and "run". + LAUNCH: ClassVar[LaunchPreference] = "indifferent" + # The target an IR client's kernels are compiled for (D26): 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, so one trace can hold an + # instance of a client class per target. Core compiles once per distinct + # target and gives each client only its own target's events. + 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 @@ -101,6 +243,75 @@ 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``); ``event.launched`` says whether + the real launch follows. + + Fires once per traced launch for each distinct (target, + specialization, launched, binding fingerprint), whichever call + produced it: a compile-only warmup, an autotune benchmark call or + the final launch. + 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 (D22), while repeated identical + calls (autotune benchmark repetitions, the benchmarked winner's + final launch) share one. + + A TritonTrace compiles every config compile-only (launched=False) + before any real launch, so each config is seen that way on every + launch; under the "run" policy, a config's first real launch with a + given binding is seen again with launched=True. + """ + + 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; a real launch that raises + gets none, and the exception reaches abort_launch unless Triton's + autotuner absorbs it.""" + + 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) and target, whichever + window made it. Nothing is loaded (D25): 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 (D27): the target is the IR client's choice, not the + machine's. A launch that skips the real launch goes on without the + config; under "run" the real launch compiles its own kernel for the + device, whose outcome follows the untraced program (e.g. Triton's + autotuner drops a config whose real compile fails). 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 (D28, 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) @@ -116,14 +327,205 @@ def grid_idx(self, value: tuple[int, ...] | None) -> None: self._set_thread_local("grid_idx", value) +_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 compile or real launch 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 real compiles are not supported." + ) + + +# Guards installing and removing the shared warmup gates (patch_warmup). +_WARMUP_GATES_LOCK = threading.Lock() + + +def _install_warmup_gate(jit_fn: Any) -> Callable: + """Install the warmup gate patch_warmup shares on ``jit_fn`` (called + with _WARMUP_GATES_LOCK held) and return it.""" + original = jit_fn.warmup + # Open scopes by host thread, innermost last: (manager, compile_context, + # real_args). + scopes: dict[int, list[tuple]] = {} + + @wraps(original) + def gate(*args, **kwargs): + stack = scopes.get(threading.get_ident()) + if not stack: + return original(*args, **kwargs) + manager, compile_context, real_args = stack[-1] + return manager._warmup_by_vote( + jit_fn, original, compile_context, real_args, args, kwargs + ) + + gate._tilelens_warmup_scopes = scopes # type: ignore[attr-defined] + gate._tilelens_warmup_previous = getattr( # type: ignore[attr-defined] + jit_fn, "__dict__", {} + ).get("warmup", _MISSING) + jit_fn.warmup = gate + return gate + + class ClientManager: def __init__(self, clients: list[Client] | None = None): - self.clients: dict[str, Client] = {} + # In trace order: the order clients are called and finalized in. + self.clients: list[Client] = [] if clients: self.add_clients(clients) self.launch = Launch() self._lock = threading.Lock() + # Compiles every IR client's kernels on the host (D25); 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, launched, 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 what + # Launch.grid is settled from. Nothing here refers to a caller's + # tensor once the launch has ended (D23). + self._delivered: set[tuple[Any, Hashable, bool, 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._last_launch_grid: Any = _MISSING + 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: @@ -131,71 +533,582 @@ def _lock_context(self): return nullcontext() def get_client(self, name: str) -> Client | None: - return self.clients.get(name) + """The first client in trace order whose NAME is ``name``.""" + return next((c for c in self.clients if c.NAME == name), None) def add_clients(self, new_clients_list: list[Client]) -> None: + """Append each client, in order; one already in the trace (the same + object) is skipped. Every other client is kept or the call raises: + nothing is dropped silently.""" + # Validate the whole resulting set before inserting anything, so a + # rejected composition leaves the manager unchanged. + resulting = list(self.clients) for new_client in new_clients_list: - duplicate = any( - isinstance(existing_client, new_client.__class__) - for existing_client in self.clients.values() + if any(new_client is client for client in resulting): + continue + if new_client.NEEDS_INTERPRETER: + taken = next( + ( + c + for c in resulting + if c.NEEDS_INTERPRETER and c.NAME == new_client.NAME + ), + None, + ) + if taken is not None: + raise ValueError( + "this trace already has an interpreting client named " + f"{new_client.NAME!r} ({type(taken).__name__}); one " + "interpreted run serves one client per name, so " + f"another ({type(new_client).__name__}) cannot share " + "the trace. Trace the kernel twice instead, e.g. " + "tilelens.trace(a)(kernel) and tilelens.trace(b)(kernel), " + "and launch each; stacked trace decorators merge into " + "one trace." + ) + resulting.append(new_client) + self._check_launch_preferences(resulting) + self.clients = resulting + + @staticmethod + def _check_launch_preferences(clients: list[Client]) -> None: + # Core only compares the IR clients' declarations (D4a); what a launch + # means to a client stays with the client. + for client in clients: + if client.LAUNCH not in LAUNCH_PREFERENCES: + raise ValueError( + f"{type(client).__name__}.LAUNCH must be one of " + f"{LAUNCH_PREFERENCES}, got {client.LAUNCH!r}" + ) + ir = [c for c in clients if not c.NEEDS_INTERPRETER] + skip = [c.NAME for c in ir if c.LAUNCH == "skip"] + run = [c.NAME for c in ir if c.LAUNCH == "run"] + if skip and run: + raise RuntimeError( + f"IR clients {skip} (LAUNCH='skip') and {run} (LAUNCH='run') " + "disagree on whether the real kernel launches, so they cannot " + "share one trace. Trace the kernel twice instead, e.g. " + "tilelens.trace(a)(kernel) and tilelens.trace(b)(kernel), and " + "launch each; stacked trace decorators merge into one trace." ) - if not duplicate: - self.clients[new_client.NAME] = new_client + + def interpreting_clients(self) -> list[Client]: + return [c for c in self.clients if c.NEEDS_INTERPRETER] + + def ir_clients(self) -> list[Client]: + return [c for c in self.clients if not c.NEEDS_INTERPRETER] + + def compile_groups(self) -> list[CompileGroup]: + """The IR clients grouped by the target their kernels are compiled + for (Client.ir_target, else the configured default), in trace order. + Raises ValueError for a target spec that names no target, or for an + IR_STAGES name no kernel compiled for the client's target holds + (see HostCompiler.check_stages).""" + groups: dict[Any, tuple[set[str], list[Client]]] = {} + for client in self.ir_clients(): + target = resolve_ir_target(client.ir_target) + try: + self.compiler.check_stages(target, client.IR_STAGES) + except HostCompileUnavailable: + # Nothing compiles for the target: each compile says so, as + # compile_failed data for the client, never as the launch's + # error. + pass + except ValueError as exc: + raise ValueError(f"{type(client).__name__}.IR_STAGES: {exc}") from None + stages, clients = groups.setdefault(target, (set(), [])) + stages.update(client.IR_STAGES) + clients.append(client) + return [ + CompileGroup(target, frozenset(stages), tuple(clients)) + for target, (stages, clients) in groups.items() + ] + + def launch_policy(self) -> Literal["skip", "run"]: + """Return "skip" if any IR client declares skip, else "run" + (add_clients has already rejected skip-vs-run conflicts).""" + if any(c.LAUNCH == "skip" for c in self.ir_clients()): + return "skip" + return "run" + + 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, D23). + self.launch = Launch() + self._reset_launch_state() + begun: list[Client] = [] + try: + for client in self.clients: + 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), 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(self, jit_fn): + def patch_warmup( + self, + jit_fn, + compile_context: Callable[[], AbstractContextManager] = nullcontext, + real_args: RealArgs | None = None, + ): + """Gate ``jit_fn.warmup`` on this manager's warmup votes, for the + calls this host thread makes during the scope. The real compile, and + only it, runs inside ``compile_context()``, on the arguments + ``real_args`` maps the call to. + + One gate per jit_fn serves every open scope, whichever trace or + thread opened it: a call is voted on by the innermost scope of its + own host thread, and goes straight to the original warmup on a + thread with none. The last scope to close removes the gate and puts + back what was there before the first one opened. + """ if not hasattr(jit_fn, "warmup"): yield return - - def patcher(fn): - @wraps(fn) - def wrapped(*args, **kwargs): - if all( - not client.pre_warmup_callback(jit_fn, *args, **kwargs) - for client in self.clients.values() - ): - return None - kwargs.pop("warmup", None) - ret = fn(*args, **kwargs) - for client in self.clients.values(): - client.post_warmup_callback(jit_fn, ret) - return ret - - return wrapped - - jit_fn.warmup = patcher(jit_fn.warmup) + thread = threading.get_ident() + with _WARMUP_GATES_LOCK: + gate = jit_fn.warmup + if getattr(gate, "_tilelens_warmup_scopes", None) is None: + gate = _install_warmup_gate(jit_fn) + scopes = gate._tilelens_warmup_scopes + scopes.setdefault(thread, []).append((self, compile_context, real_args)) try: yield finally: - jit_fn.warmup = jit_fn.warmup.__wrapped__ + with _WARMUP_GATES_LOCK: + stack = scopes[thread] + stack.pop() + if not stack: + del scopes[thread] + instance = getattr(jit_fn, "__dict__", {}) + if not scopes and instance.get("warmup") is gate: + previous = gate._tilelens_warmup_previous + if previous is _MISSING: + del instance["warmup"] + else: + jit_fn.warmup = previous + + def _warmup_by_vote( + self, jit_fn, warmup, compile_context, real_args, args, kwargs + ) -> Any: + # Every client votes; a vote may carry per-launch side effects, so do + # not short-circuit on the first True. + votes = [ + client.pre_warmup_callback(jit_fn, *args, **kwargs) + for client in self.clients + ] + if not any(votes): + return None + kwargs.pop("warmup", None) + _refuse_patched_language() + if real_args is not None: + args, kwargs = real_args(jit_fn, args, kwargs) + with compile_context(): + ret = warmup(*args, **kwargs) + for client in self.clients: + client.post_warmup_callback(jit_fn, ret) + return ret + + @contextmanager + def ir_capture( + self, + jit_fn, + *, + compile_only: bool = False, + real_args: RealArgs | None = None, + ): + """Route every ``jit_fn.run`` call through a host compile per target, + then IR-client dispatch, then the real launch if the launch policy + allows it. + + 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, benchmark or launch + goes through the capture; IR clients get one event per distinct + (target, specialization, launched, 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 + each group's target (compile_groups), no driver or device involved + (D25). Only a call that launches enters it, once, to launch (user + pre_run_hooks fire then, as untraced); a warmup call (``warmup=True``) + never launches. 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: a real launch calls it again and raises as + untraced, while a skipped one goes on, its clients seeing a launch + with no grid. Compiles and launches 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 call that does not launch then returns None, a launching + call still launches, so the device compile's own outcome decides (it + raises Triton's error, which e.g. the autotuner handles). 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. A call that does not launch returns the + first group's kernel. Yields the CaptureWindow. A target spec that + names no target, or an IR_STAGES name the target's kernels never + hold, raises 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 (D28), compile-only or not, before any client hears of the + call and before any real launch. 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 of its last real + launch; without one, the grid every kernel it compiled 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 compile or real launch refuses to start while an interpreted + traced launch has triton.language patched. + """ + launch = not compile_only and self.launch_policy() == "run" + window = CaptureWindow(compile_only=compile_only, launch=launch) + current = getattr(jit_fn, "run", None) + if current is None: + 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 + groups = self.compile_groups() + 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, groups, jit_fn, orig_run, real_args, args, kwargs, grid, warmup + ) + + 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, groups, jit_fn, orig_run, real_args, args, kwargs, grid, warmup + ): + launched = window.launch and not warmup + if real_args is None: + run_args, run_kwargs = args, kwargs + else: + run_args, run_kwargs = real_args(jit_fn, args, kwargs) + bound_args = _bind_launch_args(jit_fn, args, kwargs) + resolved_grid = _resolve_grid(grid, bound_args) + fingerprint: Any = _MISSING + first = None + delivered: list[tuple[CompileGroup, LaunchEvent]] = [] + for group in groups: + try: + _refuse_patched_language() + kernel = self.compiler.compile( + jit_fn, + run_args, + run_kwargs, + target=group.target, + stages=group.stages, + ) + except Exception as exc: + if bind_failed(exc): + # The call's own error, whatever the target (D28): raised + # as the untraced JITFunction.run raises it. + raise + window.failures.append(exc) + self._compile_failed( + group, jit_fn, args, kwargs, grid, exc, bound_args, resolved_grid + ) + continue + window.compiled += 1 + if first is None: + first = kernel + if fingerprint is _MISSING: + fingerprint = self._binding_fingerprint( + jit_fn, args, kwargs, resolved_grid + ) + key = (group.target, _specialization(kernel), launched, fingerprint) + if key in self._delivered: + continue + self._delivered.add(key) + event = self._launch_event( + jit_fn, + args, + kwargs, + grid, + kernel, + launched, + target=group.target, + bound_args=bound_args, + resolved_grid=resolved_grid, + ) + self._record_compiled_grid(event) + self._dispatch_ir("before_launch", event, group.clients) + delivered.append((group, event)) + if launched: + _refuse_patched_language() + ret = orig_run(*run_args, grid=grid, warmup=False, **run_kwargs) + # Every real launch counts for Launch.grid, delivered or not. + self._last_launch_grid = resolved_grid + else: + ret = first + for group, event in delivered: + self._dispatch_ir("after_launch", event, group.clients) + return ret + + 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, group, 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) and target: a launch + # window's benchmark call of a config the compile-only pass already + # reported 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 (group.target, call) in self._failed: + return + self._failed.add((group.target, call)) + event = self._launch_event( + jit_fn, + args, + kwargs, + grid, + None, + launched=False, + error=error, + target=group.target, + bound_args=bound_args, + resolved_grid=resolved_grid, + ) + self._dispatch_ir("compile_failed", event, group.clients) + + @staticmethod + def _launch_event( + jit_fn, + args, + kwargs, + grid, + kernel, + launched, + 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, + launched=launched, + 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 (D23, amending D5): + # 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 not event.launched and 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 last real launch's grid, else the grid every + # compiled kernel shares, else None. + if self._last_launch_grid is not _MISSING: + resolved = self._last_launch_grid + self._last_launch_grid = _MISSING + elif len(self._compiled_grids) == 1: + (resolved,) = self._compiled_grids + else: + resolved = 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 + # 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): - # Collect all for-loop callbacks from clients - all_loop_callbacks = [] - for client in self.clients.values(): - for namespace, attrs in namespaces.items(): # patch ops - for attr, op in attrs.items(): - callbacks = client.register_op_callback(op) - patch_op( - namespace, - attr, - callbacks, - frontend_name=frontend_name, - ) - all_loop_callbacks.append(client.register_for_loop_callback()) - - self._populate_loop_hooks(all_loop_callbacks) - patch_for_loop(frontend_name) - patch_lang(fn, frontend_name, client_manager=self) + lang_patched = False try: + # Collect all for-loop callbacks from clients + all_loop_callbacks = [] + for client in interpreting: + for namespace, attrs in namespaces.items(): # patch ops + for attr, op in attrs.items(): + callbacks = client.register_op_callback(op) + patch_op( + namespace, + attr, + callbacks, + frontend_name=frontend_name, + ) + all_loop_callbacks.append(client.register_for_loop_callback()) + + self._populate_loop_hooks(all_loop_callbacks) + patch_for_loop(frontend_name) + patch_lang(fn, frontend_name, client_manager=self) + lang_patched = True yield finally: - unpatch_lang(frontend_name) + if lang_patched: + unpatch_lang(frontend_name) for namespace, attrs in namespaces.items(): for attr, op in attrs.items(): unpatch_op(namespace, attr, frontend_name) @@ -204,38 +1117,55 @@ 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 = [] - for client in self.clients.values(): - # 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() + """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 (e.g. SystemExit + # from an abort), then re-raise the first failure. + first_exc: BaseException | None = None + for client in self.clients: + try: + # client may introduce tensors not declared in kernel args (e.g. tracer recording a tensor allocation) + self.launch.tensors.update(getattr(client, "tensors", []) or []) + self.launch.records += client.finalize() + except BaseException as exc: + if first_exc is None: + first_exc = exc + if first_exc is not None: + raise first_exc + finally: + self._release_pinned() + self._release_launch() def arg_callback(self, name, arg, arg_cvt): with self._lock_context(): if hasattr(arg, "data_ptr"): self.launch.tensors.add(arg) - for client in self.clients.values(): + for client in self.interpreting_clients(): client.arg_callback(name, arg, arg_cvt) def grid_callback(self, grid: tuple[int]): with self._lock_context(): self.launch.grid = grid - for client in self.clients.values(): + for client in self.interpreting_clients(): client.grid_callback(grid) def grid_idx_callback(self, grid_idx: tuple[int, ...]): with self._lock_context(): - for client in self.clients.values(): + for client in self.interpreting_clients(): client.grid_idx_callback(grid_idx) # --- For-loop callback management --- diff --git a/tilelens/core/config.py b/tilelens/core/config.py index 9163d60bd..62b2d9c28 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 (D26): 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,17 @@ 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_allow_untested_triton: TILELENS_IR_ALLOW_UNTESTED_TRITON, runs IR + mode on a Triton release outside TESTED_TRITON_VERSIONS (see + untested_triton_version). + - 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. A client's own target + (e.g. Sanitizer(compile=True, 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 +107,36 @@ def reset(self) -> None: self.sanitizer_report_max_segments: int = _get_int_env( "SANITIZER_REPORT_MAX_SEGMENTS", 8, minimum=2 ) + self.ir_allow_untested_triton: bool = _is_one( + "TILELENS_IR_ALLOW_UNTESTED_TRITON" + ) + self.ir_target: str = _get_env("TILELENS_IR_TARGET", DEFAULT_IR_TARGET) config = Config() + + +# Triton minor releases IR mode is tested on (D10b). IR mode relies on +# private Triton API: the host compile in tilelens.core.host_compile (the +# JIT's binder and argument packing, the compiler's stages) and the MLIR +# bindings behind the TTIR reader. A release joins after its IR-mode tests, +# the reader conformance suite, a bulk walk of its TTIR and the differential +# soundness corpus pass under TILELENS_IR_ALLOW_UNTESTED_TRITON=1 (D29); each +# release also needs its rows in the per-release tables (the walk layer's +# PRINTERS, the reader's _VOCABULARIES, the host compile's _RELEASE_RUNTIMES). +TESTED_TRITON_VERSIONS: tuple[str, ...] = ("3.6", "3.8") + + +def untested_triton_version() -> str | None: + """The installed Triton's version when IR mode must not run on it: its + minor release is outside TESTED_TRITON_VERSIONS and + TILELENS_IR_ALLOW_UNTESTED_TRITON is not set. None when IR mode may run. + """ + if config.ir_allow_untested_triton: + return None + import triton + + version = triton.__version__ + if ".".join(version.split(".")[:2]) in TESTED_TRITON_VERSIONS: + return None + return version diff --git a/tilelens/core/host_compile.py b/tilelens/core/host_compile.py new file mode 100644 index 000000000..7215c4fda --- /dev/null +++ b/tilelens/core/host_compile.py @@ -0,0 +1,1113 @@ +"""Host compile: one JITFunction call compiled for a GPUTarget without a GPU +(D25, D26). + +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 at the latest stage the caller asks for: a TTIR-only +request runs the front end and the backend's ``ttir`` passes (the same +passes ``triton.compile`` runs), never ``ttgir`` / ``llir`` / the binary; +``"source"`` (the front end's module before any pass) needs no pass. 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. A +request for the pipeline's last stage or for ``"sass"`` (disassembled from +the binary), and any request under a knob that rewrites or dumps stages +(``TRITON_KERNEL_OVERRIDE``, ``TRITON_KERNEL_DUMP``, ``USE_IR_LOC``, an +``ir_override`` option), is compiled by ``triton.compile`` itself, which +returns an (unloaded) ``CompiledKernel``; no device is involved either. +:meth:`HostCompiler.check_stages` rejects a stage name no kernel compiled +for the target holds. + +Either artifact has ``.asm`` (stage -> text, or bytes for a binary), +``.metadata`` (a namedtuple: ``target``, ``name``, the compile options, ...) +and ``.hash``, the specialization: what ``triton.compile`` names the kernel +for that target, whatever stage the compile stopped at. + +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 (IR mode's version gate, D10b, +still bounds the Triton releases this runs on). Where releases' JIT +runtimes differ in what the host compile mirrors, ``_RELEASE_RUNTIMES`` +says per Triton minor release what each does, and triton_api checks the +installed release's row against its code: a custom pipeline +(``knobs.runtime.add_stages_inspection_hook``), which Triton 3.8's JIT and +``triton.compile`` also key a kernel by, and a ``CompiledKernel.__del__`` +that unloads through the driver (3.8), which may run in the middle of a host +compile. A release with no row, or one whose code does not match its row, +fails closed: a host compile under a custom pipeline is refused, and the +driver's ``utils`` are refused like any device query. + +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, Iterable, Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass, field +from types import CodeType, 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 (D28, 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 (D26) ─────────────────────────── + +_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)`` (D26, amended: 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 on a + release whose ``CompiledKernel.__del__`` unloads its module through the + driver (``_ReleaseRuntime.unloads_on_del``): see _UnloadOnlyUtils.""" + + def __init__(self, target: Any, set_aside: Any = None) -> None: + self._target = target + # A context manager factory that sets this thread's scope aside + # (_ScopedActiveDriver.set_aside), when the driver's ``utils`` are + # to unload modules (see _UnloadOnlyUtils); None refuses them. + 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: + if self._set_aside is None: + self._refuse("utils") + 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, on a + Triton release whose ``CompiledKernel.__del__`` unloads a loaded module + through it (``_ReleaseRuntime.unloads_on_del``, Triton 3.8): 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, *, unloads: bool = False + ) -> Iterator[_TargetDriver]: + """Scope ``driver.active`` on this thread to a _TargetDriver for + ``target``, which is yielded. ``unloads``: its ``utils`` unload + modules through the driver outside the scope (see + _UnloadOnlyUtils) instead of being refused.""" + 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 if unloads else None + ) + 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() + + +# ─────────────────── what Triton releases differ in ─────────────────── + + +@dataclass(frozen=True) +class _ReleaseRuntime: + """What a Triton minor release's JIT runtime does, where releases + differ, that the host compile mirrors (``stages_hook_keys``) or answers + (``unloads_on_del``). Each field is checked against the installed + Triton's code (_detected_runtime) before it is relied on.""" + + # JITFunction.run and triton.compile call + # knobs.runtime.add_stages_inspection_hook with no arguments for a + # (key, hash) pair: the JIT appends the string + # '("custom_pipeline", )' to the call's specialization (so to its + # cache key and to the ASTSource's attributes), triton.compile appends + # to the kernel's cache key (so to its hash). False: only a + # backend's add_stages calls the hook, as it does in the host compile's + # own pipeline too. + stages_hook_keys: bool + # CompiledKernel.__del__ unloads a loaded module through + # driver.active.utils.unload_module: a kernel a real launch loaded may be + # collected while a thread host-compiles (see _UnloadOnlyUtils). + unloads_on_del: bool + + +# Keyed by Triton minor release. On a release with no row here, or one whose +# installed code does not do what its row says (a changed private API), the +# runtime is unknown and fails closed: a host compile while +# knobs.runtime.add_stages_inspection_hook is set raises +# HostCompileUnavailable (its kernel's hash might not be the JIT's), and +# ``driver.active.utils`` is refused inside a host compile like any device +# query. +_RELEASE_RUNTIMES: Mapping[str, _ReleaseRuntime] = MappingProxyType( + { + "3.6": _ReleaseRuntime(stages_hook_keys=False, unloads_on_del=False), + "3.8": _ReleaseRuntime(stages_hook_keys=True, unloads_on_del=True), + } +) +_STAGES_HOOK = "add_stages_inspection_hook" + + +def _code_names(fn: Any) -> frozenset[str] | None: + """The names ``fn``'s code (and the code nested in it) reads; None when + ``fn`` is no Python function.""" + code = getattr(fn, "__code__", None) + if code is None: + return None + names: set[str] = set() + pending = [code] + while pending: + current = pending.pop() + names.update(current.co_names) + pending.extend(c for c in current.co_consts if isinstance(c, CodeType)) + return frozenset(names) + + +def _detected_runtime( + compile_fn: Any, jit_function: type, compiled_kernel: type +) -> dict[str, bool | None]: + """What the installed Triton's code does of each _ReleaseRuntime field: + True or False, or None where it cannot tell (e.g. only one of + triton.compile and JITFunction.run reads the stages-inspection hook).""" + readers = [ + _code_names(compile_fn), + _code_names(inspect.getattr_static(jit_function, "run", None)), + ] + reads_hook = {names is not None and _STAGES_HOOK in names for names in readers} + finalizer = _code_names(inspect.getattr_static(compiled_kernel, "__del__", None)) + return { + "stages_hook_keys": reads_hook.pop() if len(reads_hook) == 1 else None, + "unloads_on_del": finalizer is not None and "unload_module" in finalizer, + } + + +def _release_runtime( + version: str, detected: Mapping[str, bool | None] +) -> tuple[_ReleaseRuntime | None, str]: + """The installed release's row of _RELEASE_RUNTIMES when its code does + what the row says, else None; with why not, for the error that names + it.""" + release = ".".join(version.split(".")[:2]) + row = _RELEASE_RUNTIMES.get(release) + if row is None: + return None, ( + f"Triton {release} has no row in " + "tilelens.core.host_compile._RELEASE_RUNTIMES" + ) + differ = sorted( + name for name, value in detected.items() if getattr(row, name) != value + ) + if differ: + found = ", ".join(f"{name}={detected[name]}" for name in differ) + return None, ( + f"the installed Triton {version} does not do what the Triton " + f"{release} row of tilelens.core.host_compile._RELEASE_RUNTIMES " + f"says (its code shows {found})" + ) + return row, "" + + +# ─────────────────────────── 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} (IR mode is tested on the " + "releases in tilelens.core.config.TESTED_TRITON_VERSIONS)" + ) + + +@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 + + 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, Language + from triton.compiler import ASTSource, compile, get_cache_key, make_backend + from triton.compiler.compiler import CompiledKernel + 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"), + (Language, "triton.backends.compiler.Language", "TRITON"), + ): + 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" + ) + try: + from triton.compiler.compiler import filter_traceback + except ImportError: # only trims a front-end error's traceback + + def filter_traceback(e: BaseException) -> None: # type: ignore[misc] + pass + + runtime, runtime_unknown = _release_runtime( + triton.__version__, _detected_runtime(compile, JITFunction, CompiledKernel) + ) + unloads = runtime is not None and runtime.unloads_on_del + for target in (GPUTarget("cuda", 80, 32), GPUTarget("cuda", 90, 32)): + try: + with _SCOPED_DRIVER.targeting(type(driver), target, unloads=unloads): + 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, + Language=Language, + ASTSource=ASTSource, + compile=compile, + 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, + # The installed release's _RELEASE_RUNTIMES row, None when unknown + # (then runtime_unknown says why). + runtime=runtime, + runtime_unknown=runtime_unknown, + unloads=unloads, + ) + + +def _unscoped_target_queries(target: Any) -> dict[str, Any]: + """The front end's target queries (where the installed Triton has them) + that do not answer ``target`` under its scope, with what they answer.""" + wrong: dict[str, Any] = {} + try: + from triton.language import target_info + except ImportError: + target_info = None + current_target = getattr(target_info, "current_target", None) + if current_target is not None and (got := current_target()) != target: + wrong["tl.target_info.current_target()"] = got + try: + from triton.language.semantic import TritonSemantic + except ImportError: + TritonSemantic = None + has_native_tma = getattr(TritonSemantic, "_has_native_tma", None) + if has_native_tma is not None: + # It reads nothing of the semantic, only the driver's target. + native = 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, + stages={"ttir"}, + _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 up to a stage (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 (bytes for a binary stage), every stage compiled, in + # pipeline order, "source" first when asked for. + 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 + + +# Stages a compiled kernel holds besides its backend's pipeline: the front +# end's own module (what triton.compile keeps as "source"), and the CUDA +# binary's disassembly (CompiledKernel.asm derives "sass" from "cubin"). +_SOURCE_STAGE = "source" +_DERIVED_STAGES = {"sass": "cubin"} + + +def _full_pipeline_forced(api: SimpleNamespace, options: Any) -> bool: + # Knobs that make triton.compile rewrite or dump stages, which only it + # implements. + compilation = api.knobs.compilation + return bool( + getattr(compilation, "override", False) + or getattr(compilation, "dump_ir", False) + or getattr(compilation, "use_ir_loc", None) + or getattr(options, "ir_override", None) + ) + + +def _keying_stages_hook(api: SimpleNamespace) -> Any: + """``knobs.runtime.add_stages_inspection_hook`` where the installed + release keys kernels by it (``_ReleaseRuntime.stages_hook_keys``); None + where no hook is set or the release does not (its backends' + ``add_stages`` still call it, in a host compile's pipeline too). Raises + HostCompileUnavailable for a hook set on a release whose runtime is + unknown: the host compile's kernel might not be the one the JIT names.""" + hook = getattr(api.knobs.runtime, _STAGES_HOOK, None) + if hook is None: + return None + if api.runtime is None: + raise HostCompileUnavailable( + f"knobs.runtime.{_STAGES_HOOK} is set, and how Triton " + f"{api.version}'s JIT keys a kernel by it is not known " + f"({api.runtime_unknown}), so a host compile would not name its " + "kernel as the JIT does" + ) + return hook if api.runtime.stages_hook_keys else 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), the target and the requested stages.""" + + 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, stages) -> (jit_fn, kernel). + self._kernels: dict[Hashable, tuple[Any, Any]] = {} + # target -> the stages a kernel compiled for it can hold. + self._stage_names: dict[Hashable, frozenset[str]] = {} + + def check_stages(self, target: Any, stages: Iterable[str]) -> None: + """Raise ValueError naming each of ``stages`` no kernel compiled for + ``target`` holds: a stage of the backend's pipeline (as its + ``add_stages`` builds it for a Triton kernel), ``"source"``, or one + derived from a pipeline stage (``"sass"``) is fine. Raises + HostCompileUnavailable if the target's stages cannot be listed (as + its compiles would).""" + known = self._stage_names_of(target) + unknown = set(stages) - known + if unknown: + raise ValueError( + f"IR stages {sorted(unknown)} are no stage of a kernel compiled " + f"for {format_ir_target(target)}, which holds " + f"{', '.join(sorted(known))}" + ) + + def _stage_names_of(self, target: Any) -> frozenset[str]: + names = self._stage_names.get(target) + if names is None: + api = triton_api() + arch = _target_arch(target) + pipeline: dict[str, Any] = {} + try: + backend = api.make_backend(target) + with _SCOPED_DRIVER.targeting( + api.driver_config, target, unloads=api.unloads + ): + options = backend.parse_options( + {} if arch is None else {"arch": arch} + ) + backend.add_stages(pipeline, options, api.Language.TRITON) + except Exception as exc: + raise _unavailable( + api.version, + f"the stages of a kernel compiled for {format_ir_target(target)} " + f"cannot be listed ({type(exc).__name__}: {exc})", + ) from exc + derived = { + name for name, base in _DERIVED_STAGES.items() if base in pipeline + } + names = self._stage_names[target] = frozenset( + {*pipeline, _SOURCE_STAGE, *derived} + ) + return names + + def compile( + self, + jit_fn: Any, + args: tuple, + kwargs: Mapping[str, Any], + *, + target: Any, + stages: Iterable[str] = (), + _self_test: bool = False, + ) -> Any: + """Compile the call ``jit_fn.run(*args, **kwargs)`` would compile on + a device of ``target`` (a GPUTarget), through the latest of + ``stages`` (nothing requested: the first stage). 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 always stops at the requested stage.)""" + 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, unloads=api.unloads + ) as scoped: + try: + return self._compile( + api, jit_fn, args, kwargs, target, frozenset(stages), _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, + stages: frozenset[str], + truncate: bool, + ) -> Any: + 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 (D26). 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 + stages_hook = _keying_stages_hook(api) + if stages_hook is not None: + # JITFunction.run's field for a custom pipeline, as it spells it. + _, pipeline_hash = stages_hook() + specialization.append(f'("custom_pipeline", {pipeline_hash})') + 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 + key = (id(jit_fn), cache_key, target, stages) + 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, stages, truncate, stages_hook + ) + 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, + stages: frozenset[str], + truncate: bool, + stages_hook: Any = None, + ) -> Any: + pipeline: dict[str, Any] = {} + backend.add_stages(pipeline, options, source.language) + names = list(pipeline) + # Nothing requested: the first stage; "source" alone: no pass. + wanted = stages - {_SOURCE_STAGE} if stages else frozenset(names[:1]) + full = not truncate and ( + not wanted <= set(names) # "sass" + or max(map(names.index, wanted), default=-1) == len(names) - 1 + or _full_pipeline_forced(api, options) + ) + if full: + return api.compile(source, target=target, options=options.__dict__) + last = max(map(names.index, wanted), default=-1) + # triton.compile's front half, stopped after ``names[last]``. + env_vars = api.get_cache_invalidating_env_vars() + key = api.get_cache_key(source, backend, options, env_vars) + if stages_hook is not None: + # What triton.compile appends for a custom pipeline (it asks the + # hook again, as here). + key += stages_hook()[0] + 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] = {} + if _SOURCE_STAGE in stages: + asm[_SOURCE_STAGE] = str(module) + for name in names[: last + 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 28f38301a..fd0e24642 100644 --- a/tilelens/core/trace.py +++ b/tilelens/core/trace.py @@ -1,18 +1,89 @@ -from copy import deepcopy +import copy +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 .config import config as cfg, untested_triton_version 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 launches: list[Launch] = [] +# The clients trace() takes by name, each built with its defaults. +_NAMED_CLIENTS: dict[str, type[Client]] = { + "sanitizer": Sanitizer, + "profiler": Profiler, + "race_detector": RaceDetector, + "tracer": Tracer, +} + + +def _named_client_type(name: str) -> type[Client]: + try: + return _NAMED_CLIENTS[name.lower()] + except KeyError: + raise ValueError(f"Unknown client: {name}") from None + + +def _without_warmup(kwargs: dict[str, Any]) -> dict[str, Any]: + # Launch kwargs carry warmup=False; the warmup entry points set their own. + 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) + if not closure: + return fn + cells = [] + for cell in closure: + try: + value = cell.cell_contents + except ValueError: # empty cell + value = None + cells.append(types.CellType(new) if value is old else cell) + if all(a is b for a, b in zip(cells, closure)): + return fn + rebound = types.FunctionType( + fn.__code__, fn.__globals__, fn.__name__, fn.__defaults__, tuple(cells) + ) + rebound.__kwdefaults__ = fn.__kwdefaults__ + return rebound + + +def _refers_to(fn: Any, obj: Any) -> bool: + """Whether ``fn`` is bound to ``obj`` or holds it in a closure cell.""" + if getattr(fn, "__self__", None) is obj: + return True + for cell in getattr(fn, "__closure__", None) or (): + try: + if cell.cell_contents is obj: + return True + except ValueError: # empty cell + pass + return False + class TraceInterface: def __init__(self, client: str | Client) -> None: @@ -22,27 +93,47 @@ def __init__(self, client: str | Client) -> None: @staticmethod def _normalize_client(client: str | Client) -> Client: if isinstance(client, str): - name = client.lower() - if name == "sanitizer": - return Sanitizer() - if name == "profiler": - return Profiler() - if name == "race_detector": - return RaceDetector() - if name == "tracer": - return Tracer() - raise ValueError(f"Unknown client: {client}") + return _named_client_type(client)() elif isinstance(client, Client): return client else: raise TypeError(f"Expected str or Client, got {type(client)}") def add_client(self, new_client: str | Client) -> None: + # A name asks for that kind of client with its defaults: one already + # in the trace serves it, and none of the caller's settings are lost. + if isinstance(new_client, str): + name = _named_client_type(new_client).NAME + if self.client_manager.get_client(name) is not None: + return 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 + + def _interpreter_wanted(self) -> bool: + # With no compiled kernel to launch, the interpreted run stands in for + # the launch if an interpreting client needs it or no IR client asked + # to skip the launch. + mgr = self.client_manager + return bool(mgr.interpreting_clients()) or mgr.launch_policy() == "run" class LaunchInterface: @@ -84,24 +175,137 @@ def dummy_benchmarker(fn, quantiles): return (1.0, 1.0, 1.0) def _interpreter_runner(self, runner: Any, interpreted_fn: Any) -> Any: - if self._is_autotuner(runner): - runner.fn = interpreted_fn - # Kernel Cache: replace the benchmark with a dummy to skip performance testing. - runner._do_bench = self.dummy_benchmarker - return runner - if self._is_heuristics(runner): - runner.fn = interpreted_fn - return runner - return interpreted_fn + return self._rebuild_runner(runner, interpreted_fn, interpreted=True) def _warmup_runner(self, runner: Any, jit_fn: Any | None) -> Any | None: - if not (self._is_autotuner(runner) or self._is_heuristics(runner)): - return jit_fn if jit_fn is None: return None - warmup_runner = deepcopy(runner) - warmup_runner.fn = jit_fn - return warmup_runner + return self._rebuild_runner(runner, jit_fn, interpreted=False) + + 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 (D23). + """ + 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. ``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, 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 3.6 and 3.8 + # word 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: + # Per-run state written by Autotuner.run/_bench/warmup lives on the + # copy, so a trace never changes what the user's autotuner picks. + layer.cache = {} + layer.nargs = None + # A disk-cache hit would skip benchmarking (hiding configs from IR + # clients), and interpreter timings must never be persisted. + layer.cache_results = False + # Triton's own reset_to_zero/restore_value hooks close over the + # Autotuner they were built for; point them at the copy. A hook still + # tied to the original afterwards would write its state onto the + # user's autotuner, so refuse instead. + for name in ("pre_hook", "post_hook"): + if getattr(layer, f"user_defined_{name}", False): + continue + hook = _rebind_closure(getattr(layer, name), original, layer) + if _refers_to(hook, original): + raise RuntimeError( + f"cannot trace {original!r}: Triton's default Autotuner " + f"{name} no longer closes over the Autotuner as in Triton " + "3.6 and 3.8, so the traced copy could not be isolated from " + "it (untested Triton version)" + ) + setattr(layer, name, hook) + if interpreted: + # Kernel Cache: replace the benchmark with a dummy to skip performance testing. + layer._do_bench = self.dummy_benchmarker + # do_bench is a cached_property; drop a value the original cached. + layer.__dict__.pop("do_bench", None) + + @staticmethod + def _heuristics_warmup(layer: Any) -> Callable: + # Heuristics inherits KernelInterface.warmup, which calls + # run(warmup=True) and so bypasses fn.warmup, where patch_warmup + # collects the warmup votes. Warm up like Autotuner.warmup does + # instead: fill in the heuristic kwargs as Heuristics.run does, then + # call fn.warmup. + def warmup(*args, **kwargs): + for name, heur in layer.values.items(): + kwargs[name] = heur({**dict(zip(layer.arg_names, args)), **kwargs}) + return layer.fn.warmup(*args, **kwargs) + + return warmup def _copy_callable_attrs( self, @@ -161,7 +365,11 @@ 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 IR compiles (on the host, D25) and real launches, whose + # compiles 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 @@ -172,40 +380,356 @@ def unpack_kernel( self._copy_callable_attrs(runner, self.base_fn, src_fallback=self.jit_fn) def run(self, *args, **kwargs): - with self.client_manager.patch_warmup(self.jit_fn): - if self.warmup_runner: - self.warmup_runner.warmup(*args, **kwargs) + mgr = self.client_manager + has_ir = bool(mgr.ir_clients()) + # IR mode runs only on a tested Triton release (D10b). + capture = ( + has_ir and self.jit_fn is not None and untested_triton_version() is 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 not capture: + # No compiled kernel for the IR clients (call.capture is + # False): TRITON_INTERPRET / an InterpretedFunction runner + # (no JITFunction), or an untested Triton release. Nothing + # binds the call either (D28 holds where IR mode runs): the + # JIT's binder is private API the release gate keeps off an + # untested release, reached only through the autotune layers + # that add the configs' arguments, so a call that does not + # bind returns None like any other launch here. + if self._interpreter_wanted(): + return self._run_interpreted(*args, **kwargs) + self.finalize() + return None + if mgr.interpreting_clients(): + # Mixed trace (D4b): 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 (D28), 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) + ret = self._run_compiled(*args, **kwargs) + self.finalize() + return ret + + def _real_compile_window(self): + return _unwrapped_trace_globals(self.base_fn) + + def _compile_for_ir(self, args, kwargs): + """Compile every (pruned) config through the IR runner's warmup; the + capture host-compiles each call for every IR target (D25) 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 of the + first IR target (None if it failed to compile). + + A config that fails to compile never fails the launch (D27), 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 (D28, 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, compile_only=True, 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. Every config is compiled for the IR + clients first, so what they see never depends on the autotune cache + or on benchmark timing (D3); the real launch follows unless an IR + client declared LAUNCH="skip". + + A skipped launch 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 of a skipped launch needs a GPU. A config that failed to + compile for the IR target fails neither kind of launch (D27): the IR + clients get it through compile_failed, and a real launch ("run") + compiles its own kernel for the device, which decides as it would + untraced. 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 + real 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 (D28). + + The compiles are host compiles and never enter JITFunction.run, so + the user's pre_run_hooks fire only for real launches (under "run"): + once per real call, as untraced, benchmark calls included. Only the + real launch needs a device; it compiles its own kernel through the + JIT. + """ + ret, window = self._compile_for_ir(args, kwargs) + refused = [e for e in window.failures if isinstance(e, LanguagePatchedError)] + if refused: + raise refused[0] + if self.client_manager.launch_policy() == "skip": + return None if self._autotuned(self.ir_runner) else ret + try: + with ( + self._real_compile_window(), + self.client_manager.ir_capture( + self.jit_fn, real_args=_untraced_call_args + ), + ): + return self.ir_runner.run(*args, **kwargs) + finally: + self._drop_autotuner_args(self.ir_runner) + + 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 def __call__(self, *args, **kwargs): - # When a traced JIT function is called from within another JIT function, - # we need to execute the underlying function directly - - # check that client sets match for calling and called functions + # A traced JIT function called from inside a traced kernel's + # interpreted run executes its interpreted function directly. from .frontend import triton as triton_frontend outer_client_manager = triton_frontend.frontend.current_client_manager() - if outer_client_manager is not None: - outer_clients = set(outer_client_manager.clients) - inner_clients = set(self.client_manager.clients) - if outer_clients != inner_clients: - raise RuntimeError( - "nested traced calls require matching clients; " - f"outer={outer_clients}, inner={inner_clients}" - ) + if outer_client_manager is None: + # Outside an interpreted launch this is a real compile that + # reached the trace as a plain Python callable, through a path + # _untraced_call_args does not map. Running the interpreter here + # would patch triton.language for every later compile. + raise TypeError( + f"{self.__name__} is a tilelens-traced Triton function called " + "outside a traced launch's interpreter, e.g. by a real compile " + "that reached it through an argument; pass its JITFunction " + f"({self.__name__}.jit_fn) there instead." + ) + # Only interpreting clients take part in the interpreted run (D4b). + 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; " + f"outer={outer_clients}, inner={inner_clients}" + ) return self.interpreted_fn(*args, **kwargs) def warmup(self, *args, **kwargs): - with self.client_manager.patch_warmup(self.jit_fn): - if self.warmup_runner: - self.warmup_runner.warmup(*args, **kwargs) + return self._voted_warmup(*args, **kwargs) + + def _voted_warmup(self, *args, **kwargs): + # The pre/post_warmup vote: a real compile only if some client asks. + if not self.warmup_runner: + return None + with self.client_manager.patch_warmup( + self.jit_fn, + compile_context=self._real_compile_window, + real_args=_untraced_call_args, + ): + 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: Triton's + Autotuner follows ``.fn`` from its own down to the JITFunction it + tunes (its disk cache key; ``knobs.autotuning.listener``, which + Triton 3.8 calls from a real launch's autotuning). + """ + + 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( + jit_fn: Any, args: tuple, kwargs: dict[str, Any] +) -> tuple[tuple, dict[str, Any]]: + """One JITFunction.run / warmup call's arguments as a compile (on the + host or the device) must see them: each TritonTrace passed as an + argument (also inside a tuple), or bound as a parameter default, + replaced by its JITFunction (the default passed explicitly). Triton's code generator treats any other callee as + plain Python and would call TritonTrace.__call__. + """ + + def real(value: Any) -> Any: + if isinstance(value, TritonTrace) and value.jit_fn is not None: + return value.jit_fn + if isinstance(value, tuple): + items = [real(item) for item in value] + if all(new is old for new, old in zip(items, value)): + return value + # A namedtuple is rebuilt from fields, a plain tuple from items. + return type(value)(*items) if hasattr(value, "_fields") else tuple(items) + return value + + real_args = tuple(real(arg) for arg in args) + real_kwargs = {name: real(value) for name, value in kwargs.items()} + signature = getattr(jit_fn, "signature", None) + if isinstance(signature, inspect.Signature): + for index, (name, param) in enumerate(signature.parameters.items()): + if index < len(real_args) or name in real_kwargs: + continue + if real(param.default) is not param.default: + real_kwargs[name] = real(param.default) + return real_args, real_kwargs + + +def _code_names(code: types.CodeType) -> set[str]: + names = set(code.co_names) + for const in code.co_consts: + if isinstance(const, types.CodeType): + names |= _code_names(const) + return names + + +def _is_triton_internal(value: Any) -> bool: + # Triton's own modules and stdlib functions hold no user traces. + module: Any = getattr(value, "__module__", None) + if isinstance(value, types.ModuleType): + module = value.__name__ + return isinstance(module, str) and module.split(".")[0] == "triton" + + +def _traced_references( + base_fn: Callable | None, +) -> list[tuple[dict, str, "TritonTrace"]]: + """Every (namespace, name, trace) binding a real compile of ``base_fn`` + resolves to a TritonTrace. + + Triton resolves a callee from the caller's globals and then through + module attributes (``helpers.fn``, ``pkg.api.fn``). The walk follows the + same paths, filtered by the names each function's code mentions + (co_names), and continues into every referenced JIT function's own code, + so helpers of helpers are found and unrelated bindings are left alone. + Triton also resolves parameter default expressions in the globals; their + names are not in co_names, so a global bound to a traced default counts + too. + """ + from triton import JITFunction + + found: list[tuple[dict, str, TritonTrace]] = [] + bindings: set[tuple[int, str]] = set() + visited_fns: set[int] = set() + pending: list[Any] = [base_fn] + while pending: + fn = pending.pop() + code = getattr(fn, "__code__", None) + fn_globals = getattr(fn, "__globals__", None) + if not isinstance(code, types.CodeType) or not isinstance(fn_globals, dict): + continue + if id(fn) in visited_fns: + continue + visited_fns.add(id(fn)) + names = _code_names(code) + traced_defaults = [ + d + for d in getattr(fn, "__defaults__", None) or () + if isinstance(d, TritonTrace) + ] + if traced_defaults: + names |= { + name + for name, value in fn_globals.items() + if any(value is default for default in traced_defaults) + } + namespaces = [fn_globals] + visited_namespaces = {id(fn_globals)} + while namespaces: + namespace = namespaces.pop() + for name in names: + value = namespace.get(name) + if isinstance(value, TritonTrace): + if ( + value.jit_fn is not None + and (id(namespace), name) not in bindings + ): + bindings.add((id(namespace), name)) + found.append((namespace, name, value)) + pending.append(value.base_fn) + elif isinstance(value, JITFunction) and not _is_triton_internal(value): + pending.append(value.fn) + elif isinstance(value, types.ModuleType) and not _is_triton_internal( + value + ): + module_dict = getattr(value, "__dict__", None) + if ( + isinstance(module_dict, dict) + and id(module_dict) not in visited_namespaces + ): + visited_namespaces.add(id(module_dict)) + namespaces.append(module_dict) + return found + + +@contextmanager +def _unwrapped_trace_globals(base_fn: Callable | None = None): + """Temporarily bind each TritonTrace a real compile of ``base_fn`` would + reach back to its JITFunction. + + Under the CLI wrappers every ``@triton.jit`` function, device functions + included, becomes a TritonTrace. Triton's dependency walker and code + generator only accept JITCallables as callees ("Unsupported function + referenced"); the interpreter tolerates the wrapper through + ``TritonTrace.__call__``, a real compile does not. Only the bindings the + kernel's code can reach are swapped (see _traced_references), and each is + restored on exit unless it was rebound in the meantime. Callees reached + through closure variables are not covered (Triton's dependency walker + rejects them); callees passed as arguments are mapped per call instead + (_untraced_call_args). + """ + swapped: list[tuple[dict, str, TritonTrace]] = [] + try: + for namespace, name, trace in _traced_references(base_fn): + if namespace.get(name) is trace: + namespace[name] = trace.jit_fn + swapped.append((namespace, name, trace)) + yield + finally: + for namespace, name, trace in reversed(swapped): + if namespace.get(name) is trace.jit_fn: + namespace[name] = trace class NKITrace(LaunchInterface, TraceInterface): @@ -273,23 +797,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._interpreter_wanted(): + # Only IR clients, and one asked to skip the launch: there is + # no compiled kernel here, 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): @@ -331,17 +861,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._interpreter_wanted(): + # Only IR clients, and one asked to skip the launch: there is + # no compiled kernel here, so nothing runs. + self.finalize() + return None + with self.client_manager.patch_run(self.base_fn, frontend_name="gluon"): + try: + ret = self.runner.run( + *args, + **kwargs, + client_manager=self.client_manager, + ) + finally: + self.client_manager.post_run_callback(self.base_fn) + self.finalize() + return ret def __call__(self, *args, **kwargs): return self.fn(*args, **kwargs) diff --git a/tilelens/core/trace_io.py b/tilelens/core/trace_io.py index c95723c9c..aa233be65 100644 --- a/tilelens/core/trace_io.py +++ b/tilelens/core/trace_io.py @@ -10,6 +10,8 @@ from ..clients.profiler import data as profiler_data from ..clients.sanitizer import data as sanitizer_data +from ..ir import launch as ir_launch +from ..ir import verdict as ir_verdict from ..utils import traceback_utils from . import data as trace_data from .data import Launch, TensorSnapshot @@ -22,7 +24,16 @@ ArrayMap = dict[str, np.ndarray] _TRACE_CLASSES = { f"{cls.__module__}:{cls.__qualname__}": cls - for module in (trace_data, profiler_data, sanitizer_data, traceback_utils) + for module in ( + trace_data, + profiler_data, + sanitizer_data, + traceback_utils, + # Pure data, like verdict (neither imports Triton): TensorFacts is + # registered in its own right, not only as sanitizer_data's import. + ir_launch, + ir_verdict, + ) for cls in vars(module).values() if isinstance(cls, type) and is_dataclass(cls) } diff --git a/tilelens/ir/__init__.py b/tilelens/ir/__init__.py new file mode 100644 index 000000000..6bcba1c6a --- /dev/null +++ b/tilelens/ir/__init__.py @@ -0,0 +1,41 @@ +"""Compiled-IR layer: read Triton-compiled kernels (TTIR) for the IR-mode clients. + +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/_mlir_walk.py b/tilelens/ir/_mlir_walk.py new file mode 100644 index 000000000..e80632fba --- /dev/null +++ b/tilelens/ir/_mlir_walk.py @@ -0,0 +1,2666 @@ +"""Walk a printed TTIR module: MLIR bindings for structure, the text for attributes. + +Private to ``tilelens.ir``; only ``ttir_reader.py`` imports it. This is the one +module that touches Triton's private MLIR bindings (``triton._C.libtriton.ir``), +and every lifetime hazard of those bindings stays here (D10a, amendment D7+). + +``walk_module(text)`` reads the same bytes twice: + +1. The bindings parse the text (``ir.parse_mlir_module``, TTIR dialects only): + op names, operand / result values, full type strings, region / block + nesting, block arguments, value locs, and the few attributes the 3.6 + getters can read (``get_str_attr`` / ``get_bool_attr`` / + ``get_flat_symbol_ref_attr``). +2. A line scan of the text rebuilds the same op tree and recovers what the + bindings keep opaque: integer and enum attributes, constant values, cf + successors, and the locs of zero-result ops. + +The two trees are zipped in pre-order and must agree on every op: name, +result / operand / region / block / block-argument counts, every SSA edge +(each printed use resolves, through region-scoped name tables, to the value +the bindings report for that operand slot), every value loc, every +bindings-readable attribute, the type-constrained integer attributes, the cf +successor arities and the printer's ``// pred:`` comments. Text-only enums +must lie in a closed per-op vocabulary, other recovered attributes must have +their expected Python type, and a module that prints locs must print one on +every op (zero-result op locs have no second source). The only tolerated +differences are the printer's own elisions: a trailing zero-operand +``scf.yield`` of an ``scf.for`` / ``scf.if`` block (also when that yield is +the block's only op and the region prints as ``{ }``), unprinted empty +trailing regions, and the default-valued attributes listed in the +release's ``Printer.defaults``. Anything else raises +``MisalignedModule``; a text the MLIR parser rejects raises +``ModuleParseError`` with the parser's own diagnostic. + +The result is frozen pure-Python data: ``Module`` holds ``Op`` / ``Block`` / +``Value`` / ``Func`` records, which compare by value, hash, pickle and +deep-copy. Values are dense local ints (indices into ``Module.values``), +never the bindings' ``value.id()`` pointers. No binding object ever leaves +``_bind_walk``. + +Binding lifetime (spike condition 2): the context is pinned on the module +(``mod.context = ctx``), everything is extracted inside one function, the +module's body block is erased, the module is dropped before the context, and +nothing derived from either is returned. No binding erases the module op +itself, so each parse still leaks that (empty) op; results are cached by +sha256 of the text (condition 4). The parse window redirects the process's +fd 2 under a lock; a forked child gets both back. + +The text vocabulary is per Triton minor release: ``PRINTERS`` maps a +release ("3.6", "3.8") to its ``Printer`` table (what the reader needs, +the leading keywords and their closed vocabularies, the elided defaults, +the printed operand orders, the attributes the bindings can read, what the +dialect's types allow), each audited against that release's printer. The +installed Triton's release selects the table (``printer()``); a release +without one fails closed (``UnknownTritonRelease``), never borrowing +another release's table. The syntax recognizers below are shared by every +table's release (audited per release too); a release whose printer spells +a construct differently needs its own row. ``Module.release`` records the +table a module was read with. +""" + +from __future__ import annotations + +import collections +import dataclasses +import hashlib +import os +import re +import sys +import tempfile +import threading +import types +from dataclasses import dataclass +from typing import Any, Callable, Iterator, Mapping, Sequence + +# ─────────────────────────── records ─────────────────────────── + + +class _FrozenMap(Mapping[str, Any]): + """The read-only mapping behind the records' mapping fields. Unlike + ``MappingProxyType`` it hashes (its values are plain hashable data) and + pickles, so the frozen records hash, pickle and deep-copy too.""" + + __slots__ = ("_d",) + + def __init__(self, items: Mapping[str, Any]) -> None: + self._d = dict(items) + + def __getitem__(self, key: str) -> Any: + return self._d[key] + + def __iter__(self) -> Iterator[str]: + return iter(self._d) + + def __len__(self) -> int: + return len(self._d) + + def __hash__(self) -> int: + return hash(frozenset(self._d.items())) + + def __repr__(self) -> str: + return repr(self._d) + + def __reduce__(self) -> tuple[Any, ...]: + return (_FrozenMap, (self._d,)) + + +@dataclass(frozen=True) +class SourceLoc: + file: str + line: int + col: int + + +@dataclass(frozen=True) +class Value: + index: int # position in Module.values + type: str # full type text: "tensor<64x!tt.ptr>", "i32", ... + op: int | None # defining op index (an op result), else None + block: int | None # owning block index (a block argument), else None + position: int # result number, or argument number + # NameLoc name of the value's own loc (the Python variable / parameter + # name), None when unnamed. Printed SSA names are never exposed: the + # printer sanitises and uniquifies them. + name: str | None + + +@dataclass(frozen=True) +class Block: + index: int # position in Module.blocks + op: int # owning op index + region: int # region number within the owning op + position: int # block number within the region + label: str | None # printed label ("^bb1"); None for an unlabeled entry block + args: tuple[int, ...] # value indices + arg_types: tuple[str, ...] + arg_names: tuple[str | None, ...] # NameLoc names (see Value.name) + ops: tuple[int, ...] # op indices in program order + + +@dataclass(frozen=True) +class Op: + index: int # pre-order position in Module.ops (the text order) + name: str # "tt.load", "scf.for", "builtin.module", ... + operands: tuple[int, ...] # value indices, in ODS operand order + operand_types: tuple[str, ...] + results: tuple[int, ...] # value indices + result_types: tuple[str, ...] + # Recovered attributes (read-only). Integer attrs are ints, enums their + # printed keyword ("slt", "acq_rel", "ieee"); program-id axes are ints; + # arith.constant "value" is an int (the signed value, as the printer + # prints signless integers) / bool, ("float", literal) for a float, or + # ("dense", literal) for a non-splat dense (then "splat" is False; a + # dense splat has "splat" True and a scalar "value"); cf ops carry + # "successors" as destination block indices. + attrs: Mapping[str, Any] + regions: tuple[tuple[int, ...], ...] # block indices, per walked region + # block indices from the module body down to the parent block; () for + # the module op itself + path: tuple[int, ...] + position: int # position within the parent block + line_no: int | None # header line (1-based); None for an elided terminator + end_line: int | None # closing line of a region op, else == line_no + loc: SourceLoc | None # the op's own site (callee frame of a callsite loc) + callers: tuple[SourceLoc, ...] # callsite chain, innermost caller first + loc_name: str | None # NameLoc label of the op's loc, if any + implicit: bool = False # a terminator the printer elided + + @property + def block(self) -> int | None: + return self.path[-1] if self.path else None + + +@dataclass(frozen=True) +class FuncArg: + index: int + value: int # value index of the entry-block argument + type: str + name: str | None # NameLoc name: the Python parameter name + attrs: Mapping[str, Any] # printed argument attributes (tt.divisibility, ...) + + +@dataclass(frozen=True) +class Func: + op: int # the tt.func op index + sym_name: str + visibility: str + args: tuple[FuncArg, ...] # () for a body-less declaration + + +@dataclass(frozen=True) +class Module: + ops: tuple[Op, ...] + blocks: tuple[Block, ...] + values: tuple[Value, ...] + funcs: tuple[Func, ...] # tt.func ops in text order + stats: Mapping[str, int] # what the alignment checked (tests / bulk tool) + release: str # the Triton release whose Printer table read the text ("3.6") + + +class MisalignedModule(Exception): + """The text scan and the bindings disagree (or the text holds a construct + the text layer cannot read faithfully). ``problems`` lists every mismatch + found, ``line_no`` is the first text line involved (None if unknown).""" + + def __init__(self, problems: Sequence[str], line_no: int | None = None) -> None: + self.problems = tuple(problems) or ("misaligned module",) + self.line_no = line_no + super().__init__(self.problems[0]) + + +class ModuleParseError(Exception): + """The MLIR parser rejected the text (``diagnostic`` is its own message, + with the temporary parse path replaced by ````), the parse input + could not be created, or the text holds a construct the release's parser + cannot be handed (``Printer.block_pointer_types``).""" + + def __init__(self, diagnostic: str, line_no: int | None = None) -> None: + self.diagnostic = diagnostic + self.line_no = line_no + super().__init__(diagnostic) + + +class UnknownTritonRelease(Exception): + """The installed Triton's minor release has no ``Printer`` table: the + text layer does not know that release's printer, so nothing is read + (``release`` is the minor release, ``version`` the full version).""" + + def __init__(self, release: str, version: str) -> None: + self.release = release + self.version = version + known = ", ".join(f"{r}.x" for r in PRINTERS) + self.message = ( + f"the TTIR walk layer has no printer table for Triton {version} " + f"(it reads the printers of Triton {known}); a release is added " + "to tilelens.ir._mlir_walk.PRINTERS after auditing its printer" + ) + super().__init__(self.message) + + +# ─────────────────────────── types ─────────────────────────── + + +@dataclass(frozen=True) +class TypeInfo: + """A printed TTIR type split into shape and element (D9 widths).""" + + text: str + shape: tuple[int, ...] # () for scalars + elem: str # element type text: "i32", "f16", "!tt.ptr", ... + int_bits: int | None # iN element -> N (signless); index -> 64 + float_bits: int | None + pointee: str | None # element pointer -> pointee type text + pointee_bits: int | None + block_ptr: bool # !tt.ptr> + + +_FLOAT_BITS = { + "f64": 64, "f32": 32, "f16": 16, "bf16": 16, "tf32": 32, + "f8E4M3FN": 8, "f8E5M2": 8, "f8E4M3FNUZ": 8, "f8E5M2FNUZ": 8, + "f8E4M3B11FNUZ": 8, "f8E8M0FNU": 8, "f4E2M1FN": 4, +} # fmt: skip +_RE_TENSOR = re.compile(r"^tensor<((?:\d+x)*)(.*)>$") +_RE_INT = re.compile(r"^i(\d+)$") +_RE_PTR = re.compile(r"^!tt\.ptr<(.*?)(?:, \d+)?>$") + + +def _scalar_bits(t: str) -> int | None: + m = _RE_INT.match(t) + if m: + return int(m.group(1)) + return _FLOAT_BITS.get(t) + + +def parse_type(text: str) -> TypeInfo: + shape: tuple[int, ...] = () + elem = text + m = _RE_TENSOR.match(text) + if m: + shape = tuple(int(d) for d in m.group(1).split("x") if d) + elem = m.group(2) + im = _RE_INT.match(elem) + pm = _RE_PTR.match(elem) + pointee = pm.group(1) if pm else None + int_bits = int(im.group(1)) if im else (64 if elem == "index" else None) + return TypeInfo( + text=text, + shape=shape, + elem=elem, + int_bits=int_bits, + float_bits=_FLOAT_BITS.get(elem), + pointee=pointee, + pointee_bits=_scalar_bits(pointee) if pointee else None, + block_ptr=bool(pointee and pointee.startswith("tensor<")), + ) + + +# ─────────────────────────── strings and brackets ─────────────────────────── + + +def _string_end(s: str, start: int) -> int: + """Index of the quote closing the string literal that opens at ``start``.""" + i = start + 1 + while i < len(s): + c = s[i] + if c == "\\": + i += 2 + continue + if c == '"': + return i + i += 1 + raise ValueError("unterminated string literal") + + +_HEX = frozenset("0123456789abcdefABCDEF") +_SIMPLE_ESCAPES = {"\\": 0x5C, '"': 0x22, "n": 0x0A, "t": 0x09} + + +def _unescape(body: str) -> str: + """Decode an MLIR string literal body. The printer escapes every + non-printable or non-ASCII byte as ``\\XX``; the escapes of one character + are the bytes of its UTF-8 encoding, so collect bytes and decode once.""" + out = bytearray() + i = 0 + while i < len(body): + c = body[i] + if c != "\\": + out += c.encode("utf-8") + i += 1 + continue + nxt = body[i + 1 : i + 2] + if nxt in _SIMPLE_ESCAPES: + out.append(_SIMPLE_ESCAPES[nxt]) + i += 2 + elif len(body) >= i + 3 and body[i + 1] in _HEX and body[i + 2] in _HEX: + out.append(int(body[i + 1 : i + 3], 16)) + i += 3 + else: + raise ValueError(f"bad string escape {body[i : i + 3]!r}") + try: + return out.decode("utf-8") + except UnicodeDecodeError as e: + raise ValueError(f"string literal is not UTF-8: {e}") from None + + +def _mask(line: str) -> tuple[str, list[tuple[int, int, str]]]: + """Replace every string literal's body by '_' (same length, so indices into + the masked line index the raw line too). Returns the masked line and the + literals as (open quote index, close quote index, decoded text).""" + out: list[str] = [] + lits: list[tuple[int, int, str]] = [] + i = 0 + while i < len(line): + c = line[i] + if c == '"': + j = _string_end(line, i) + lits.append((i, j, _unescape(line[i + 1 : j]))) + out.append('"' + "_" * (j - i - 1) + '"') + i = j + 1 + continue + out.append(c) + i += 1 + return "".join(out), lits + + +_OPEN = {"(": ")", "[": "]", "{": "}", "<": ">"} + + +def _match_close(s: str, i: int) -> int: + """Index of the bracket closing ``s[i]`` in masked text ('->' is no '>').""" + stack = [_OPEN[s[i]]] + j = i + 1 + while j < len(s): + c = s[j] + if c in _OPEN: + stack.append(_OPEN[c]) + elif c == ">" and s[j - 1] == "-": + pass + elif c in ")]}>": + if c != stack[-1]: + raise ValueError(f"bracket mismatch at column {j + 1}") + stack.pop() + if not stack: + return j + j += 1 + raise ValueError("unclosed bracket") + + +def _split_top(s: str, sep: str = ",") -> list[tuple[int, str]]: + """Split masked text on ``sep`` at bracket depth 0. Returns (start index of + the stripped part in ``s``, stripped part) for every non-empty part.""" + parts: list[tuple[int, str]] = [] + depth = 0 + begin = 0 + for i, c in enumerate(s): + if c in "([{<": + depth += 1 + elif c in ")]}" or (c == ">" and (i == 0 or s[i - 1] != "-")): + depth -= 1 + elif c == sep and depth == 0: + parts.append((begin, s[begin:i])) + begin = i + 1 + parts.append((begin, s[begin:])) + out = [] + for start, p in parts: + if p.strip(): + out.append((start + len(p) - len(p.lstrip()), p.strip())) + return out + + +def _trailing_loc(masked: str) -> tuple[int, int] | None: + """(start, end) of the ``loc(...)`` that ends the masked text, if any.""" + s = masked.rstrip() + if not s.endswith(")"): + return None + depth = 0 + for j in range(len(s) - 1, -1, -1): + c = s[j] + if c == ")": + depth += 1 + elif c == "(": + depth -= 1 + if depth == 0: + if s[max(0, j - 3) : j] == "loc" and ( + j < 4 or not (s[j - 4].isalnum() or s[j - 4] in "_.$") + ): + return j - 3, len(s) + return None + return None + + +def _blank_dicts(masked: str) -> str: + """Masked text with every top-level ``{...}`` replaced by spaces.""" + out = list(masked) + i = 0 + while i < len(masked): + if masked[i] == "{": + j = _match_close(masked, i) + out[i : j + 1] = " " * (j + 1 - i) + i = j + i += 1 + return "".join(out) + + +# ─────────────────────────── locs ─────────────────────────── +# One grammar for both sources: the text's `loc(#locN)` trailers resolved +# through the `#locN = loc(...)` table, and the bindings' fully inlined +# `str(value.get_loc())`. Both normalise to nested tuples compared with ==. + +_RE_LOC_ALIAS_DEF = re.compile(r"^(#loc\d*)\s*=\s*loc\((.*)\)\s*$") +_RE_LOC_ALIAS = re.compile(r"#loc\d*") +_RE_FILE_POS = re.compile(r":(\d+):(\d+)") +_LOC_DEPTH_LIMIT = 128 + + +class _LocParser: + def __init__(self, aliases: Mapping[str, str], parse_path: str) -> None: + self._aliases = aliases # "#loc12" -> inner text of loc(...) + self._parse_path = parse_path + self._alias_memo: dict[str, tuple] = {} + self._text_memo: dict[str, tuple | None] = {} + + def parse(self, loc_text: str) -> tuple | None: + """Normalised tree of a ``loc(...)`` text; None for a loc the parser + invented (it names the temporary parse path: the text printed none).""" + got = self._text_memo.get(loc_text, _MISSING) + if got is not _MISSING: + return got # type: ignore[return-value] + s = loc_text.strip() + if not (s.startswith("loc(") and s.endswith(")")): + raise ValueError(f"not a loc: {loc_text[:80]!r}") + tree: tuple | None = self._full(s[4:-1], 0) + if _names_path(tree, self._parse_path): # type: ignore[arg-type] + tree = None + self._text_memo[loc_text] = tree + return tree + + def _full(self, s: str, depth: int) -> tuple: + tree, rest = self._expr(s.strip(), depth) + if rest.strip(): + raise ValueError(f"trailing loc text: {rest[:60]!r}") + return tree + + def _expr(self, s: str, depth: int) -> tuple[tuple, str]: + if depth > _LOC_DEPTH_LIMIT: + raise ValueError("loc nesting too deep") + if s.startswith("#loc"): + m = _RE_LOC_ALIAS.match(s) + assert m is not None + name = m.group(0) + if name not in self._alias_memo: + if name not in self._aliases: + raise ValueError(f"undefined loc alias {name}") + self._alias_memo[name] = ("pending",) + self._alias_memo[name] = self._full(self._aliases[name], depth + 1) + elif self._alias_memo[name] == ("pending",): + raise ValueError(f"cyclic loc alias {name}") + return self._alias_memo[name], s[m.end() :] + if s.startswith("unknown"): + return ("unknown",), s[len("unknown") :] + if s.startswith("callsite("): + callee, rest = self._expr(s[len("callsite(") :].lstrip(), depth + 1) + rest = rest.lstrip() + if not rest.startswith("at "): + raise ValueError("callsite loc without 'at'") + caller, rest = self._expr(rest[3:].lstrip(), depth + 1) + rest = rest.lstrip() + if not rest.startswith(")"): + raise ValueError("unclosed callsite loc") + return ("callsite", callee, caller), rest[1:] + if s.startswith("fused"): + rest = s[len("fused") :] + meta: str | None = None + if rest.startswith("<"): + close = _match_close(_mask(rest)[0], 0) + meta = rest[1:close] + rest = rest[close + 1 :] + if not rest.startswith("["): + raise ValueError("bad fused loc") + parts = [] + rest = rest[1:].lstrip() + while not rest.startswith("]"): + p, rest = self._expr(rest, depth + 1) + parts.append(p) + rest = rest.lstrip() + if rest.startswith(","): + rest = rest[1:].lstrip() + elif not rest.startswith("]"): + raise ValueError("bad fused loc list") + return ("fused", meta, tuple(parts)), rest[1:] + if s.startswith('"'): + end = _string_end(s, 0) + text = _unescape(s[1:end]) + rest = s[end + 1 :] + m = _RE_FILE_POS.match(rest) + if m: + rest = rest[m.end() :] + if rest.lstrip().startswith("to"): + raise ValueError("file range locs are not supported") + return ("file", text, int(m.group(1)), int(m.group(2))), rest + if rest.startswith("("): + child, rest = self._expr(rest[1:], depth + 1) + rest = rest.lstrip() + if not rest.startswith(")"): + raise ValueError("unclosed name loc") + return ("name", text, child), rest[1:] + return ("name", text, None), rest + raise ValueError(f"unrecognized loc: {s[:60]!r}") + + +_MISSING = object() + + +def _names_path(tree: tuple, path: str) -> bool: + """Does any file loc inside ``tree`` name ``path``? (iterative)""" + stack = [tree] + while stack: + t = stack.pop() + if t is None: + continue + kind = t[0] + if kind == "file": + if t[1] == path: + return True + elif kind == "name": + stack.append(t[2]) + elif kind == "callsite": + stack += [t[1], t[2]] + elif kind == "fused": + stack += list(t[2]) + return False + + +def _loc_site( + tree: tuple | None, +) -> tuple[SourceLoc | None, tuple[SourceLoc, ...], str | None]: + """(site, callers, name) of a normalised loc. The callee frame of a + callsite is the op's site (a memory op belongs to the callee, #361 + _LocTable); the caller chain follows, innermost first. Recursion depth is + bounded by the parser's ``_LOC_DEPTH_LIMIT``.""" + if tree is None: + return None, (), None + kind = tree[0] + if kind == "file": + return SourceLoc(tree[1], tree[2], tree[3]), (), None + if kind == "name": + site, callers, _ = _loc_site(tree[2]) + return site, callers, tree[1] + if kind == "callsite": + site, callers, name = _loc_site(tree[1]) + csite, ccallers, _ = _loc_site(tree[2]) + return site, callers + ((csite,) if csite else ()) + ccallers, name + if kind == "fused": + for part in tree[2]: + got = _loc_site(part) + if got[0] is not None: + return got + return None, (), None + + +# ─────────────────────────── text scan ─────────────────────────── + +_RE_RESULTS = re.compile( + r"^((?:%[-\w.$]+(?::\d+)?)(?:\s*,\s*%[-\w.$]+(?::\d+)?)*)\s*=\s*" +) +_RE_USE = re.compile(r"%[-\w.$]+(?:#\d+)?") +_RE_OPNAME = re.compile(r"^[A-Za-z_][\w$.]*") +_RE_LABEL = re.compile(r"^\^[-\w.$]+") +_RE_SUCC = re.compile(r"\^[-\w.$]+") +_RE_WORD_BEFORE = re.compile(r"([A-Za-z_]\w*)\s*$") +_RE_DEF_EQ = re.compile(r"\s*=(?!=)") +_RE_BARE_ASSIGN = re.compile( + r"(? None: + super().__init__(msg) + self.line_no = line_no + self.msg = msg + + +@dataclass +class _ArgDef: + """A block argument printed in a label, a func header or an op header.""" + + name: str # "%x" + type: str | None # printed type text (labels / func headers) + loc: str | None # printed "loc(...)" text + attrs: dict[str, Any] # printed argument attrs (func headers) + + +@dataclass +class _TBlock: + label: str | None + args: list[_ArgDef] + ops: list["_TOp"] + line_no: int + # the printer's predecessor comment: labels (a multiset); None if absent + preds: tuple[str, ...] | None = None + + +@dataclass +class _TOp: + name: str + line_no: int + raw: str # stripped raw header line (comment removed) + masked: str # masked header line (same length) + lits: list[tuple[int, int, str]] + name_end: int # index just past the op name + hdr_end: int # end of the header operands/attrs (before loc / region opener) + result_names: list[str] + n_results: int + opens_region: bool + regions: list[list[_TBlock]] + close_line: int | None = None + close_raw: str = "" + close_masked: str = "" + close_lits: list[tuple[int, int, str]] | None = None + + +@dataclass +class _TextTree: + root: _TOp + aliases: dict[str, str] + pred_comments: bool # any block label carries a predecessor comment + + +def _parse_results(prefix: str) -> tuple[list[str], int]: + names: list[str] = [] + n = 0 + for tok in prefix.split(","): + tok = tok.strip() + if ":" in tok: + base, k = tok.split(":") + names += [f"{base}#{i}" for i in range(int(k))] + n += int(k) + else: + names.append(tok) + n += 1 + return names, n + + +def _arg_def( + part_raw: str, part_masked: str, lits_rel: list[tuple[int, int, str]] +) -> _ArgDef: + """``%x: type {attrs} loc(...)`` (masked part + raw part, same indices).""" + colon = part_masked.find(":") + if not part_masked.startswith("%") or colon < 0: + raise ValueError(f"bad block argument {part_raw[:60]!r}") + name = part_masked[:colon].strip() + rest_m = part_masked[colon + 1 :] + rest_r = part_raw[colon + 1 :] + loc = None + span = _trailing_loc(rest_m) + if span is not None: + loc = rest_r[span[0] : span[1]] + rest_m, rest_r = rest_m[: span[0]], rest_r[: span[0]] + attrs: dict[str, Any] = {} + brace = rest_m.find("{") + if brace >= 0: + close = _match_close(rest_m, brace) + attrs = _parse_dict( + rest_m, + brace, + close, + [(a - colon - 1, b - colon - 1, t) for a, b, t in lits_rel], + ) + if rest_m[close + 1 :].strip(): + raise ValueError(f"text after argument attrs: {part_raw[:60]!r}") + rest_r = rest_r[:brace] + return _ArgDef(name, rest_r.strip(), loc, attrs) + + +def _arg_list( + raw: str, masked: str, lits: list[tuple[int, int, str]], open_idx: int +) -> tuple[list[_ArgDef], int]: + """Parse the parenthesised argument list opening at ``open_idx``; returns + the defs and the index of the closing paren.""" + close = _match_close(masked, open_idx) + inner_m = masked[open_idx + 1 : close] + defs = [] + for start, part in _split_top(inner_m): + a = open_idx + 1 + start + rel = [(x - a, y - a, t) for x, y, t in lits if a <= x < a + len(part)] + defs.append(_arg_def(raw[a : a + len(part)], part, rel)) + return defs, close + + +def _pred_comment(comment: str, line_no: int) -> tuple[str, ...]: + m = _RE_PREDS.match(comment.strip()) + if m is None: + raise _TextError(line_no, f"unrecognized block comment {comment[:60]!r}") + if m.group(1): + return (m.group(1),) + if m.group(2): + preds = tuple(p.strip() for p in m.group(3).split(",")) + if len(preds) != int(m.group(2)) or not all( + _RE_LABEL.fullmatch(p) for p in preds + ): + raise _TextError(line_no, f"malformed predecessor comment {comment[:60]!r}") + return preds + return () + + +def _parse_op_line( + ln: int, raw: str, masked: str, lits: list[tuple[int, int, str]] +) -> _TOp: + rm = _RE_RESULTS.match(masked) + result_names: list[str] = [] + n_results = 0 + start = 0 + if rm: + result_names, n_results = _parse_results(rm.group(1)) + start = rm.end() + body = masked[start:] + if body.startswith('"'): # generic form: "tt.reduce"(...) + end = _string_end(raw, start) + name = next((t for a, _b, t in lits if a == start), "") + name_end = end + 1 + else: + nm = _RE_OPNAME.match(body) + if nm is None: + raise _TextError(ln, f"cannot read an op name: {raw[:80]!r}") + name = nm.group(0) + name_end = start + nm.end() + if "." not in name: + name = f"builtin.{name}" + opens = masked.endswith("{") + hdr_end = len(masked) + if opens: + hdr_end = masked.rfind("{") + else: + span = _trailing_loc(masked) + if span is not None: + hdr_end = span[0] + return _TOp( + name, + ln, + raw, + masked, + lits, + name_end, + hdr_end, + result_names, + n_results, + opens, + [], + ) + + +def _scan_text(text: str) -> _TextTree: + """Rebuild the op tree from printed lines. Returns a synthetic file-level + op whose single region holds the top-level ops, and the ``#loc`` table.""" + root = _TOp("", 0, "", "", [], 0, 0, [], 0, True, [[]]) + stack = [root] + aliases: dict[str, str] = {} + pred_comments = False + for ln, raw_line in enumerate(text.splitlines(), start=1): + raw = raw_line.strip() + if not raw or raw.startswith("//"): + continue + try: + masked, lits = _mask(raw) + except ValueError as e: + raise _TextError(ln, str(e)) from None + comment = None + cut = masked.find("//") + if cut >= 0: + comment = raw[cut + 2 :].strip() + raw, masked = raw[:cut].rstrip(), masked[:cut].rstrip() + lits = [x for x in lits if x[1] < cut] + top = stack[-1] + if masked.startswith("#"): + m = _RE_LOC_ALIAS_DEF.match(raw) + if len(stack) != 1 or m is None: + raise _TextError( + ln, + f"attribute alias {raw[:60]!r}: only #loc aliases are TTIR (TTGIR input?)", + ) + aliases[m.group(1)] = m.group(2) + continue + if comment is not None and not masked.startswith("^"): + raise _TextError(ln, f"unexpected comment {comment[:60]!r}") + try: + if masked.startswith("^"): + if top is root: + raise _TextError(ln, "block label outside any region") + lm = _RE_LABEL.match(masked) + if lm is None: + raise _TextError(ln, f"bad block label {raw[:60]!r}") + after = masked[lm.end() :] + args: list[_ArgDef] = [] + if after.startswith("("): + args, close = _arg_list(raw, masked, lits, lm.end()) + after = masked[close + 1 :] + if after.strip() != ":": + raise _TextError(ln, f"bad block label {raw[:60]!r}") + preds = None + if comment is not None: + preds = _pred_comment(comment, ln) + pred_comments = True + top.regions[-1].append(_TBlock(lm.group(0), args, [], ln, preds)) + continue + if masked.startswith("}"): + if top is root: + raise _TextError(ln, "unbalanced '}'") + rest = masked[1:].lstrip() + if rest.endswith("{"): # "} else {", "} do {", "}, {" + top.regions.append([]) + continue + if rest.startswith(")"): # generic op closer "}) ..." + rest = rest[1:].lstrip() + off = len(masked) - len(rest) + top.close_line = ln + top.close_raw = raw[off:] + top.close_masked = rest + top.close_lits = [(a - off, b - off, t) for a, b, t in lits if a >= off] + stack.pop() + continue + op = _parse_op_line(ln, raw, masked, lits) + except ValueError as e: + raise _TextError(ln, str(e)) from None + region = top.regions[-1] + if not region: + region.append(_TBlock(None, [], [], ln)) + region[-1].ops.append(op) + if op.opens_region: + op.regions = [[]] + stack.append(op) + if len(stack) != 1: + raise _TextError( + stack[-1].line_no, + f"region opened at line {stack[-1].line_no} is never closed", + ) + return _TextTree(root, aliases, pred_comments) + + +# ─────────────────────────── attribute recovery (text) ─────────────────────────── + + +def _attr_value(v: str, v_start: int, lits: list[tuple[int, int, str]]) -> Any: + if v.startswith('"'): + return next((t for a, _b, t in lits if a == v_start), None) + m = re.fullmatch(r"(-?\d+)(?:\s*:\s*(?:i\d+|index))?", v) + if m: + return int(m.group(1)) + if v in ("true", "false"): + return v == "true" + m = re.fullmatch(r"array", v) + if m: + return tuple(int(x) for x in (m.group(1) or "").split(",") if x.strip()) + m = re.fullmatch(r"dense<(-?\d+)>\s*:\s*tensor<.*>", v) + if m: + return ("splat", int(m.group(1))) + return v # anything else stays raw text + + +def _parse_dict( + masked: str, open_idx: int, close_idx: int, lits: list[tuple[int, int, str]] +) -> dict[str, Any]: + d: dict[str, Any] = {} + inner = masked[open_idx + 1 : close_idx] + for start, part in _split_top(inner): + p_start = open_idx + 1 + start + eq = _split_top(part, "=") + if len(eq) == 2: + (_k_off, k), (v_off, v) = eq + key = k + if key.startswith('"'): + key = next((t for a, _b, t in lits if a == p_start), key.strip('"')) + d[key] = _attr_value(v, p_start + v_off, lits) + elif len(eq) == 1: + d[part] = True # unit attribute + else: + raise ValueError(f"bad attribute {part[:60]!r}") + return d + + +def _dicts( + masked: str, lits, start: int, stop: int +) -> list[tuple[int, dict[str, Any]]]: + """Every ``{...}`` / ``<{...}>`` attribute dict in ``masked[start:stop]`` + with its paren depth (0 = op attributes).""" + out = [] + depth = 0 + i = start + while i < stop: + c = masked[i] + if c == "(": + depth += 1 + elif c == ")": + depth -= 1 + elif c == "{": + close = _match_close(masked, i) + out.append((depth, _parse_dict(masked, i, close, lits))) + i = close + i += 1 + return out + + +def _leading_keywords(masked: str, start: int, stop: int) -> list[str]: + """Bare tokens between the op name and the first operand / attr / type: + ``arith.cmpi slt, %a`` -> ['slt']; ``tt.atomic_rmw fadd, relaxed, gpu, %p`` + -> ['fadd', 'relaxed', 'gpu']; ``tt.get_program_id x : i32`` -> ['x'].""" + s = masked[start:stop] + m = re.match(r"\s*((?:[A-Za-z_]\w*\s*,\s*)*[A-Za-z_]\w*)(?=\s*(?:,|:|$))", s) + if not m: + return [] + return [t.strip() for t in m.group(1).split(",")] + + +_RE_INT_LIT = re.compile(r"-?\d+") +_RE_HEX_LIT = re.compile(r"0x[0-9A-Fa-f]+") +_RE_FLOAT_LIT = re.compile( + r"[-+]?(?:\d+\.?\d*(?:[eE][-+]?\d+)?|\.\d+(?:[eE][-+]?\d+)?|inf|nan)" +) + + +def _scalar_literal(lit: str, elem: TypeInfo) -> Any: + """A scalar constant literal read against its element type.""" + if lit in ("true", "false"): + if elem.int_bits != 1: + raise ValueError(f"bool literal {lit} of type {elem.elem}") + return lit == "true" + if elem.int_bits is not None: + if _RE_INT_LIT.fullmatch(lit): + return int(lit) + if _RE_HEX_LIT.fullmatch(lit): + return int(lit, 16) + raise ValueError(f"integer constant {lit!r}") + if elem.float_bits is not None: + if _RE_FLOAT_LIT.fullmatch(lit) or _RE_HEX_LIT.fullmatch(lit): + return ("float", lit) + raise ValueError(f"float constant {lit!r}") + raise ValueError(f"constant of unsupported element type {elem.elem}") + + +def _constant_attrs(t: _TOp, result_type: str | None) -> dict[str, Any]: + """``arith.constant [{attrs}] [: ]`` (raw text keeps a + ``dense<"0x...">`` blob intact).""" + s = t.masked[t.name_end : t.hdr_end] + off = t.name_end + lead = len(s) - len(s.lstrip()) + if s.lstrip().startswith("{"): # a leading attr-dict precedes the value + close = _match_close(s, lead) + off += close + 1 + s = s[close + 1 :] + parts = _split_top(s, ":") + if not parts: + raise ValueError("arith.constant without a value") + p_off, p_masked = parts[0] + lit = t.raw[off + p_off : off + p_off + len(p_masked)] + if result_type is None: + raise ValueError("arith.constant without a result") + ty = parse_type(result_type) + elem = parse_type(ty.elem) + a: dict[str, Any] = {"literal": lit} + if lit.startswith("dense<") and lit.endswith(">"): + inner = lit[6:-1].strip() + if not ty.shape and ty.text == ty.elem: + raise ValueError(f"dense constant of scalar type {ty.text}") + if inner.startswith(("[", '"')) or not inner: + a["value"] = ("dense", inner) + a["splat"] = False + else: + a["value"] = _scalar_literal(inner, elem) + a["splat"] = True + else: + if ty.shape: + raise ValueError(f"scalar literal {lit!r} of tensor type {ty.text}") + a["value"] = _scalar_literal(lit, elem) + return a + + +def _symbol(t: _TOp) -> str: + m = re.search(r"@", t.masked[t.name_end : t.hdr_end]) + if m is None: + raise ValueError("no symbol") + at = t.name_end + m.start() + if t.masked[at + 1 : at + 2] == '"': + return next(x for a, _b, x in t.lits if a == at + 1) + sm = re.match(r"@([\w$.-]+)", t.raw[at:]) + if sm is None: + raise ValueError("bad symbol") + return sm.group(1) + + +def _first_string(t: _TOp) -> str | None: + return next((x for a, _b, x in t.lits if t.name_end <= a < t.hdr_end), None) + + +_AXES = {"x": 0, "y": 1, "z": 2} +_RESHAPE_KEYWORDS = frozenset({"allow_reorder", "efficient_layout"}) +_V = r"%[-\w.$]+" # a printed value name +_U = _V + r"(?:#\d+)?" # a use (result #k of a multi-result op) +# `scf.for [unsigned] %iv = %lb to %ub step %s +# [iter_args(%a = %init, ...) -> (T, ...)] [: T]` +_RE_SCF_FOR = re.compile( + rf"\s*(unsigned\s+)?{_V}\s*=\s*{_U}\s+to\s+{_U}\s+step\s+{_U}" + rf"(?:\s+iter_args\(\s*{_V}\s*=\s*{_U}(?:\s*,\s*{_V}\s*=\s*{_U})*\s*\)" + r"\s*->\s*\(.+\))?(?:\s*:\s*[A-Za-z_]\w*)?\s*" +) + + +def _text_attrs( + t: _TOp, result_types: Sequence[str], printer: "Printer" +) -> tuple[dict[str, Any], dict[str, Any]]: + """(attrs, func header args info) recovered from the op's own text: + header and, for a region op, its closing line.""" + name = t.name + a: dict[str, Any] = {} + for depth, d in _dicts(t.masked, t.lits, t.name_end, t.hdr_end): + if depth == 0: + a.update(d) + if t.close_masked: + stop = len(t.close_masked) + span = _trailing_loc(t.close_masked) + if span is not None: + stop = span[0] + for depth, d in _dicts(t.close_masked, t.close_lits or [], 0, stop): + if depth == 0: + a.update(d) + extra: dict[str, Any] = {} + keys = printer.keywords.get(name) + if keys is not None: + kw = _leading_keywords(t.masked, t.name_end, t.hdr_end) + if len(kw) != len(keys): + raise ValueError( + f"expected {len(keys)} leading keyword(s) {keys}, got {kw}" + ) + a.update(zip(keys, kw)) + if name == "arith.constant": + a.update(_constant_attrs(t, result_types[0] if result_types else None)) + elif name == "tt.dot": + # `%a, %b, %c(, inputPrecision = X)?`: any other bare assignment is a + # misread, not an elided default + bare = _blank_dicts(t.masked[t.name_end : t.hdr_end]) + assigned = _RE_BARE_ASSIGN.findall(bare) + if assigned not in ( + [], + [("inputPrecision", assigned[0][1] if assigned else "")], + ): + raise ValueError(f"unrecognized tt.dot assignments {assigned}") + if assigned: + a["inputPrecision"] = assigned[0][1] + elif name == "tt.reshape": + bare = _blank_dicts(t.masked[t.name_end : t.hdr_end]) + m = re.match(r"\s*%[-\w.$]+(?:#\d+)?((?:\s+[A-Za-z_]\w*)*)\s*(?::|$)", bare) + if m is None: + raise ValueError("unrecognized tt.reshape syntax") + for word in m.group(1).split(): + if word not in _RESHAPE_KEYWORDS: + raise ValueError(f"unknown tt.reshape keyword {word!r}") + a[word] = True + elif name == "scf.for": + # the one header keyword is `unsigned` (unsignedCmp: the loop compares + # its bounds unsigned, which changes the trip count); any other + # header shape is a misread + m = _RE_SCF_FOR.fullmatch(t.masked, t.name_end, t.hdr_end) + if m is None: + raise ValueError("unrecognized scf.for header") + if "unsignedCmp" in a: + raise ValueError("scf.for prints unsignedCmp in its attribute dict") + a["unsignedCmp"] = m.group(1) is not None + elif name == "tt.elementwise_inline_asm": + s = _first_string(t) + if s is not None: + a["asm_string"] = s + elif name == "tt.call": + a["callee"] = _symbol(t) + elif name == "tt.func": + a["sym_name"] = _symbol(t) + vm = re.match(r"\s*(public|private|nested)\b", t.masked[t.name_end :]) + if vm: + a["visibility"] = vm.group(1) + # the parameter list: `@name(%a: T {attrs} loc(..), ...)` + m = re.search(r"@", t.masked[t.name_end : t.hdr_end]) + assert m is not None + at = t.name_end + m.start() + sym_end = ( + _string_end(t.masked, at + 1) + 1 + if t.masked[at + 1 : at + 2] == '"' + else at + 1 + ) + while sym_end < t.hdr_end and ( + t.masked[sym_end].isalnum() or t.masked[sym_end] in "_$.-" + ): + sym_end += 1 + if t.masked[sym_end : sym_end + 1] != "(": + raise ValueError("tt.func without a parameter list") + close = _match_close(t.masked, sym_end) + if "%" in t.masked[sym_end:close]: + extra["args"], _ = _arg_list(t.raw, t.masked, t.lits, sym_end) + else: + extra["args"] = [] + elif name == "tt.print": + s = _first_string(t) + if s is not None: + a["prefix"] = s + elif name == "tt.assert": + s = _first_string(t) + if s is not None: + a["message"] = s + elif name in ("cf.br", "cf.cond_br"): + extra["successors"] = _successor_groups(t) + if name not in ("cf.br", "cf.cond_br") and "^" in t.masked[t.name_end : t.hdr_end]: + raise ValueError(f"successor syntax on {name} is not supported") + return a, extra + + +def _successor_groups(t: _TOp) -> list[tuple[str, int]]: + """(label, number of printed operands) per successor, in order.""" + s = t.masked + out = [] + for m in _RE_SUCC.finditer(s, t.name_end, t.hdr_end): + n = 0 + j = m.end() + if s[j : j + 1] == "(": + close = _match_close(s, j) + n = len(_RE_USE.findall(s, j, close)) + out.append((m.group(0), n)) + return out + + +def _header_uses_and_defs(t: _TOp) -> tuple[list[str], list[str], list[str]]: + """(uses, the word printed right before each use, defs) in the header. + Defs are block arguments printed in the header: ``%x: type`` (func args) + and ``%iv = ...`` / ``iter_args(%a = %init)`` / ``scf.while (%a = %init)``.""" + s = t.masked + uses: list[str] = [] + before: list[str] = [] + defs: list[str] = [] + for m in _RE_USE.finditer(s, t.name_end, t.hdr_end): + if s.startswith(":", m.end(), t.hdr_end) or _RE_DEF_EQ.match( + s, m.end(), t.hdr_end + ): + defs.append(m.group(0)) + else: + uses.append(m.group(0)) + w = _RE_WORD_BEFORE.search(s, max(t.name_end, m.start() - 32), m.start()) + before.append(w.group(1) if w else "") + return uses, before, defs + + +# ─────────────────────────── per-version tables ─────────────────────────── +# One Printer per Triton minor release, keyed by release in PRINTERS. A +# table states what its release's printer prints; it is audited against +# that release (the goldens, the conformance corpus, and a bulk walk of +# real compiled kernels, tools/ir_bulk_conformance.py), never inferred from +# another release. + + +def _order_descriptor_store(uses: list[str], _before: list[str]) -> list[str]: + # `%desc[%i, %j], %src` -> ODS (desc, src, indices...) + return uses[:1] + uses[-1:] + uses[1:-1] if len(uses) >= 2 else uses + + +def _order_dot_scaled(uses: list[str], before: list[str]) -> list[str]: + # `%a scale %as, %b scale %bs, %c` -> ODS (a, b, c, a_scale?, b_scale?) + main = [u for u, w in zip(uses, before) if w != "scale"] + scales = [u for u, w in zip(uses, before) if w == "scale"] + return main + scales + + +def _frozen(d: Mapping[Any, Any]) -> Mapping[Any, Any]: + return types.MappingProxyType(dict(d)) + + +@dataclass(frozen=True, eq=False) +class Printer: + """What the text layer knows of one Triton minor release's TTIR printer. + + ``needed``: what the reader consumes, per op; every key must be + recovered. ``keywords``: the leading enum keywords each custom syntax + prints, in order. ``vocab``: the closed vocabularies of the text-only + enums (any other spelling, or a generic-form integer, is a + misalignment). ``defaults``: attributes the custom printer omits while + they hold their default value. ``attr_types``: value types of the + recovered non-enum attributes (a garbled value that falls back to raw + text must not pass as recovered). ``bind_attrs``: attributes the + release's getters can read, per op, as (getter, name), read on the + bindings side and cross-checked against the text; ``bind_ints`` maps a + keyword the text prints to the integer the ``int`` getter reads for it. + ``printed_order``: custom syntaxes whose printed operand order is not + the ODS operand order (the SSA-edge check fails closed on the others). + ``elides_yield``: ops whose custom printer drops a trailing + zero-operand scf.yield. ``block_pointer_types``: the dialect has + ``!tt.ptr>``; where it has not, a text holding one is + refused before the bindings parse it.""" + + release: str + needed: Mapping[str, tuple[str, ...]] + keywords: Mapping[str, tuple[str, ...]] + vocab: Mapping[tuple[str, str], frozenset[str]] + defaults: Mapping[str, Mapping[str, Any]] + attr_types: Mapping[tuple[str, str], type] + bind_attrs: Mapping[str, tuple[tuple[str, str], ...]] + bind_ints: Mapping[tuple[str, str], Mapping[str, int]] + printed_order: Mapping[str, Callable[[list[str], list[str]], list[str]]] + elides_yield: frozenset[str] + block_pointer_types: bool + + def extend(self, release: str, **changes: Any) -> "Printer": + """This table with ``changes`` merged in: a mapping field's entries + are added to (or replace) this table's, any other field replaced.""" + merged: dict[str, Any] = {} + for key, value in changes.items(): + base = getattr(self, key) + merged[key] = ( + _frozen({**base, **value}) if isinstance(base, Mapping) else value + ) + return dataclasses.replace(self, release=release, **merged) + + +_SEM = frozenset({"relaxed", "acquire", "release", "acq_rel"}) +_SCOPE = frozenset({"gpu", "cta", "sys"}) +_AXIS_WORDS = frozenset(_AXES) + +# Triton 3.6 (the D10a spike, its review corpora and the #361 goldens). +_PRINTER_3_6 = Printer( + release="3.6", + needed=_frozen( + { + "arith.constant": ("value",), + "tt.make_range": ("start", "end"), + "tt.get_program_id": ("axis",), + "tt.get_num_programs": ("axis",), + "arith.cmpi": ("predicate",), + "arith.cmpf": ("predicate",), + "tt.expand_dims": ("axis",), + "tt.reduce": ("axis",), + "tt.scan": ("axis", "reverse"), + "tt.atomic_rmw": ("rmw_op", "sem", "scope"), + "tt.atomic_cas": ("sem", "scope"), + "tt.elementwise_inline_asm": ( + "asm_string", + "constraints", + "pure", + "packed_element", + ), + "tt.call": ("callee",), + "tt.func": ("sym_name", "visibility", "noinline"), + "tt.load": ("isVolatile",), + "tt.trans": ("order",), + "tt.reshape": ("allow_reorder",), + "tt.dot": ("inputPrecision", "maxNumImpreciseAcc"), + "tt.print": ("prefix",), + "scf.for": ("unsignedCmp",), + "cf.br": ("successors",), + "cf.cond_br": ("successors",), + } + ), + keywords=_frozen( + { + "arith.cmpi": ("predicate",), + "arith.cmpf": ("predicate",), + "tt.atomic_rmw": ("rmw_op", "sem", "scope"), + "tt.atomic_cas": ("sem", "scope"), + "tt.get_program_id": ("axis",), + "tt.get_num_programs": ("axis",), + "tt.descriptor_reduce": ("kind",), + } + ), + vocab=_frozen( + { + ("arith.cmpi", "predicate"): frozenset( + {"eq", "ne", "slt", "sle", "sgt", "sge", "ult", "ule", "ugt", "uge"} + ), + ("arith.cmpf", "predicate"): frozenset( + { + "false", + "oeq", + "ogt", + "oge", + "olt", + "ole", + "one", + "ord", + "ueq", + "ugt", + "uge", + "ult", + "ule", + "une", + "uno", + "true", + } + ), # fmt: skip + ("tt.atomic_rmw", "rmw_op"): frozenset( + { + "and", + "or", + "xor", + "add", + "fadd", + "max", + "min", + "umax", + "umin", + "exch", + } + ), + ("tt.atomic_rmw", "sem"): _SEM, + ("tt.atomic_rmw", "scope"): _SCOPE, + ("tt.atomic_cas", "sem"): _SEM, + ("tt.atomic_cas", "scope"): _SCOPE, + ("tt.get_program_id", "axis"): _AXIS_WORDS, + ("tt.get_num_programs", "axis"): _AXIS_WORDS, + ("tt.dot", "inputPrecision"): frozenset( + {"tf32", "tf32x3", "ieee", "bf16x3", "bf16x6"} + ), + ("tt.descriptor_reduce", "kind"): frozenset( + {"add", "min", "max", "inc", "dec", "and", "or", "xor"} + ), + ("tt.func", "visibility"): frozenset({"public", "private", "nested"}), + } + ), + defaults=_frozen( + { + "tt.dot": _frozen({"inputPrecision": "ieee", "maxNumImpreciseAcc": 0}), + "tt.load": _frozen({"isVolatile": False}), + "tt.reshape": _frozen({"allow_reorder": False, "efficient_layout": False}), + "tt.func": _frozen({"visibility": "public"}), + } + ), + attr_types=_frozen( + { + ("tt.make_range", "start"): int, + ("tt.make_range", "end"): int, + ("tt.expand_dims", "axis"): int, + ("tt.reduce", "axis"): int, + ("tt.scan", "axis"): int, + ("tt.scan", "reverse"): bool, + ("tt.elementwise_inline_asm", "asm_string"): str, + ("tt.elementwise_inline_asm", "constraints"): str, + ("tt.elementwise_inline_asm", "pure"): bool, + ("tt.elementwise_inline_asm", "packed_element"): int, + ("tt.call", "callee"): str, + ("tt.func", "sym_name"): str, + ("tt.func", "noinline"): bool, + ("tt.load", "isVolatile"): bool, + ("tt.trans", "order"): tuple, + ("tt.reshape", "allow_reorder"): bool, + ("scf.for", "unsignedCmp"): bool, + ("tt.dot", "maxNumImpreciseAcc"): int, + ("tt.print", "prefix"): str, + ("tt.print", "hex"): bool, + ("tt.print", "isSigned"): tuple, + ("tt.assert", "message"): str, + } + ), + # the 3.6 getters: get_str_attr / get_bool_attr / get_flat_symbol_ref_attr + # (integer and enum attributes are opaque to them) + bind_attrs=_frozen( + { + "tt.func": ( + ("str", "sym_name"), + ("str", "sym_visibility"), + ("bool", "noinline"), + ), + "tt.call": (("sym", "callee"),), + "tt.elementwise_inline_asm": ( + ("str", "asm_string"), + ("str", "constraints"), + ("bool", "pure"), + ), + "tt.load": (("bool", "isVolatile"),), + "arith.constant": (("bool", "value"),), + "tt.scan": (("bool", "reverse"),), + "tt.print": (("str", "prefix"), ("bool", "hex")), + "tt.assert": (("str", "message"),), + } + ), + bind_ints=_frozen({}), + printed_order=_frozen( + { + # audited against TritonOps.td + "tt.descriptor_store": _order_descriptor_store, + "tt.descriptor_reduce": _order_descriptor_store, + "tt.dot_scaled": _order_dot_scaled, + } + ), + elides_yield=frozenset({"scf.for", "scf.if"}), + block_pointer_types=True, +) + +# Triton 3.8 (audited on 3.8.0: 365 host-compiled kernels of the +# conformance and soundness corpora and the golden generators, extra and +# descriptor kernels for sm89 / sm90 / sm100, the goldens regenerated under +# 3.8 and 1127 Triton-cache texts, all walked with 0 misalignments; printed +# attributes, elided defaults, enum keywords and operand orders are 3.6's, +# the goldens pin identically). tl.debug_barrier() prints `ttg.barrier +# all` instead of `gpu.barrier`: its one attribute, addrSpace, is a bit +# enum printed as one keyword per set of bits (`all`, or single flags; +# a combination prints `local|global_read`, which the one-keyword syntax +# does not read: misaligned), and the 3.8 get_int_attr reads its bits. The +# dialect has no block-pointer types (tt.make_tensor_ptr / tt.advance are +# gone, tl.make_block_ptr lowers to pointer arithmetic), and its parser +# aborts the process on a `!tt.ptr>` instead of reporting an +# error. The tensordesc type prints `!tt.tensordesc<32x32xf16>` (3.6: +# `!tt.tensordesc>`); types are compared as printed, so +# that takes no table entry. +_ADDR_SPACE = _frozen( + { + "none": 0, + "local": 1, + "global_read": 2, + "global_write": 4, + "tensor_read": 8, + "tensor_write": 16, + "all": 31, + } +) +_PRINTER_3_8 = _PRINTER_3_6.extend( + "3.8", + # addrSpace is not read by the TTIR reader (the barrier is inert there); + # recovered for the consumers that order memory (a race detector) + needed={"ttg.barrier": ("addrSpace",)}, + keywords={"ttg.barrier": ("addrSpace",)}, + vocab={("ttg.barrier", "addrSpace"): frozenset(_ADDR_SPACE)}, + bind_attrs={"ttg.barrier": (("int", "addrSpace"),)}, + bind_ints={("ttg.barrier", "addrSpace"): _ADDR_SPACE}, + block_pointer_types=False, +) + +PRINTERS: Mapping[str, Printer] = _frozen( + {p.release: p for p in (_PRINTER_3_6, _PRINTER_3_8)} +) + + +def triton_release() -> tuple[str, str]: + """(minor release, full version) of the installed Triton, e.g. + ("3.6", "3.6.0"): the release is the version's first two components, + as the D10b gate reads it.""" + import triton + + version = str(triton.__version__) + return ".".join(version.split(".")[:2]), version + + +def printer(release: str | None = None) -> Printer: + """The Printer table of ``release`` (default: the installed Triton's); + raises ``UnknownTritonRelease`` for a release without one.""" + version = release + if release is None: + release, version = triton_release() + table = PRINTERS.get(release) + if table is None: + raise UnknownTritonRelease(release, version or release) + return table + + +_GETTERS = { + "str": "get_str_attr", + "bool": "get_bool_attr", + "sym": "get_flat_symbol_ref_attr", + "int": "get_int_attr", +} +_GETTER_TYPES: dict[str, type] = {"str": str, "bool": bool, "sym": str, "int": int} +_BIND_KEY = {"sym_visibility": "visibility"} + + +def _is_exactly(value: Any, want: type) -> bool: + """``isinstance`` without bool passing as int; a tuple holds ints.""" + if isinstance(value, bool) != (want is bool) or not isinstance(value, want): + return False + return want is not tuple or all( + isinstance(x, int) and not isinstance(x, bool) + for x in value # type: ignore[attr-defined] + ) + + +# A block-pointer type anywhere in a text: `ptr>`, `!tt>>`, whitespace, +# newlines or `//` comments in between), and any type alias definition, which +# could hide the pointee. The _GAP forms read the raw text (a comment runs to +# the end of its line), the others the code view of _screen_view. +_GAP = r"(?:\s|//[^\n]*(?![^\n]))*" +_RE_BLOCK_PTR_TYPE = re.compile(r"\bptr\s*<\s*tensor\b") +_RE_BLOCK_PTR_TYPE_GAP = re.compile(rf"\bptr{_GAP}<{_GAP}tensor\b") +_RE_TYPE_ALIAS_DEF = re.compile(r"^\s*(![-\w.$]+)\s*=", re.M) +_RE_TYPE_ALIAS_DEF_GAP = re.compile(rf"^\s*![-\w.$]+{_GAP}=", re.M) + + +def _code_line(line: str) -> str: + """``line`` as the parser's tokens see it, same length: string literal + bodies masked ('_'), a ``//`` comment outside strings blanked to the end + of the line. After an unterminated string the rest stays raw (the + parser stops there; reading it can only add a refusal).""" + out: list[str] = [] + i = 0 + while i < len(line): + c = line[i] + if c == '"': + try: + j = _string_end(line, i) + except ValueError: + out.append(line[i:]) + break + out.append('"' + "_" * (j - i - 1) + '"') + i = j + 1 + continue + if line.startswith("//", i): + out.append(" " * (len(line) - i)) + break + out.append(c) + i += 1 + return "".join(out) + + +def _screen_view(text: str) -> str: + """The code view of ``text`` (``_code_line`` per line), offsets kept, so + a match's line is its count of newlines before it plus one.""" + return "\n".join(_code_line(line) for line in text.split("\n")) + + +def _screen(text: str, table: Printer) -> None: + """Refuse, before the bindings see it, a text the release's parser + cannot be handed: without block-pointer types (3.8) the parser aborts + the whole process on one (an assertion in PointerType::get), so such a + text never reaches it. The whole text is searched, so a type split over + lines or around a comment is found too; string literals are not read as + types, comments are skipped.""" + if table.block_pointer_types or not ( + _RE_BLOCK_PTR_TYPE_GAP.search(text) or _RE_TYPE_ALIAS_DEF_GAP.search(text) + ): + return + view = _screen_view(text) + found = [] + m = _RE_BLOCK_PTR_TYPE.search(view) + if m is not None: + found.append((m.start(), "a block-pointer type (!tt.ptr>)")) + m = _RE_TYPE_ALIAS_DEF.search(view) + if m is not None: + found.append( + (m.start(1), "a type alias definition (it may name a block-pointer type)") + ) + if not found: + return + at, what = min(found) + ln = view.count("\n", 0, at) + 1 + raise ModuleParseError( + f"line {ln}: {what}: the TTIR of Triton {table.release} has no " + "block pointers, and its parser aborts the process on one; the " + "text is refused before parsing", + ln, + ) + + +def _type_checks( + name: str, attrs: Mapping[str, Any], opnd_types, res_types +) -> list[str] | None: + """Cross-check text-only integer attributes against the bindings' types, + the one independent channel for them. None = no check applies.""" + if ( + name == "tt.make_range" + and isinstance(attrs.get("start"), int) + and isinstance(attrs.get("end"), int) + ): + shape = parse_type(res_types[0]).shape + if shape != (attrs["end"] - attrs["start"],): + return [ + f"make_range [{attrs['start']}, {attrs['end']}) vs result shape {shape}" + ] + return [] + if name == "tt.expand_dims" and isinstance(attrs.get("axis"), int): + src, dst = parse_type(opnd_types[0]).shape, parse_type(res_types[0]).shape + ax = attrs["axis"] + if not 0 <= ax <= len(src) or dst != src[:ax] + (1,) + src[ax:]: + return [f"expand_dims axis {ax}: {src} -> {dst}"] + return [] + if name in ("tt.reduce", "tt.scan") and isinstance(attrs.get("axis"), int): + src = parse_type(opnd_types[0]).shape + ax = attrs["axis"] + want = src if name == "tt.scan" else src[:ax] + src[ax + 1 :] + got = parse_type(res_types[0]).shape + if not (0 <= ax < len(src)) or got != want: + return [f"{name} axis {ax}: {src} -> {got}"] + return [] + if name == "tt.trans" and isinstance(attrs.get("order"), tuple): + src, dst = parse_type(opnd_types[0]).shape, parse_type(res_types[0]).shape + order = attrs["order"] + if sorted(order) != list(range(len(src))) or dst != tuple( + src[i] for i in order + ): + return [f"trans order {order}: {src} -> {dst}"] + return [] + if name == "arith.constant" and "value" in attrs: + t = parse_type(res_types[0]) + v = attrs["value"] + if isinstance(v, bool) or not isinstance(v, int): + return [] # kind vs element type is checked by _scalar_literal + bits = parse_type(t.elem).int_bits + assert bits is not None + # the printer prints i1 as true / false and every wider signless + # integer as signed: `4294967295 : i32` parses, but MLIR holds -1 + if bits == 1: + return [f"constant {v}: the printer prints {t.text} as true / false"] + lo, hi = -(1 << (bits - 1)), (1 << (bits - 1)) - 1 + if lo <= v <= hi: + return [] + return [f"constant {v} is outside the printed (signed) range of {t.text}"] + return None + + +# ─────────────────────────── bindings side ─────────────────────────── + + +@dataclass +class _BOp: + name: str + operands: tuple[int, ...] # raw value ids (valid within one parse only) + operand_types: tuple[str, ...] + results: tuple[int, ...] + result_types: tuple[str, ...] + result_locs: tuple[str, ...] + region_ids: tuple[int, ...] + region_sizes: tuple[int, ...] + block_id: int | None + attrs: dict[str, Any] + + +@dataclass +class _BBlock: + region_id: int + args: tuple[int, ...] + arg_types: tuple[str, ...] + arg_locs: tuple[str, ...] + ops: list[int] # indices into the walk list + + +@dataclass +class _BindWalk: + ops: list[_BOp] + blocks: dict[int, _BBlock] + region_blocks: dict[int, list[int]] # raw region id -> raw block ids in order + path: str # the parse path (parser-invented locs name it) + + +def _write_all(fd: int, data: bytes) -> None: + view = memoryview(data) + while view: + n = os.write(fd, view) + view = view[n:] + + +class _ParseInput: + """The text as a file path for ``parse_mlir_module``: an anonymous memfd + via ``/proc/self/fd`` when available, else a temporary file.""" + + def __init__(self, data: bytes) -> None: + self.path: str | None = None + self._fd: int | None = None + self._tmp: str | None = None + memfd_create = getattr(os, "memfd_create", None) + if memfd_create is not None: + try: + fd = memfd_create("tilelens-ttir", getattr(os, "MFD_CLOEXEC", 0)) + except OSError: + fd = None + if fd is not None: + path = f"/proc/self/fd/{fd}" + try: + _write_all(fd, data) + ok = os.path.exists(path) + except OSError: + ok = False + if ok: + self._fd, self.path = fd, path + return + os.close(fd) + try: + fd, tmp = tempfile.mkstemp(prefix="tilelens-", suffix=".ttir") + except OSError as e: + raise ModuleParseError( + f"cannot create a temporary file for the MLIR parser: {e}" + ) from None + try: + try: + _write_all(fd, data) + finally: + os.close(fd) + except OSError as e: + os.unlink(tmp) + raise ModuleParseError(f"cannot write the MLIR parser input: {e}") from None + self._tmp = self.path = tmp + + def close(self) -> None: + if self._fd is not None: + os.close(self._fd) + self._fd = None + if self._tmp is not None: + try: + os.unlink(self._tmp) + except OSError: + pass + self._tmp = None + + +# (saved fd 2, capture buffer) while a capture redirects fd 2; set and +# cleared under _PARSE_LOCK, read by the fork handler (_after_fork_in_child) +_REDIRECT: tuple[int, int] | None = None + + +class _StderrCapture: + """Redirect fd 2 (where the C++ parser prints its diagnostic) into an + anonymous file for the duration of a ``with`` block; only used under + ``_PARSE_LOCK``. Everything written to fd 2 in that window lands in + ``data``: on a failed parse that is the diagnostic, on a successful one + it belongs to someone else (parser warnings, other threads), so + ``replay()`` passes it on to the real fd 2.""" + + def __init__(self) -> None: + self._buf: int | None = None + self._saved: int | None = None + self.data = b"" + self.text = "" + + def __enter__(self) -> "_StderrCapture": + global _REDIRECT + try: + memfd_create = getattr(os, "memfd_create", None) + if memfd_create is not None: + buf = memfd_create("tilelens-diag", getattr(os, "MFD_CLOEXEC", 0)) + else: + with tempfile.TemporaryFile() as f: + buf = os.dup(f.fileno()) + except OSError: + return self # no capture: the diagnostic stays on stderr + try: + sys.stderr.flush() + except (AttributeError, OSError, ValueError): + pass + try: + saved = os.dup(2) + except OSError: + os.close(buf) + return self + # published before fd 2 moves (and cleared after it is back), so a + # fork at any point of the window can restore fd 2 in the child + _REDIRECT = (saved, buf) + try: + os.dup2(buf, 2) + except OSError: + _REDIRECT = None + os.close(saved) + os.close(buf) + return self + self._buf, self._saved = buf, saved + return self + + def __exit__(self, *exc: object) -> None: + global _REDIRECT + if self._saved is not None: + os.dup2(self._saved, 2) + _REDIRECT = None + os.close(self._saved) + self._saved = None + if self._buf is not None: + try: + os.lseek(self._buf, 0, os.SEEK_SET) + chunks = [] + while chunk := os.read(self._buf, 1 << 16): + chunks.append(chunk) + self.data = b"".join(chunks) + self.text = self.data.decode("utf-8", "replace") + finally: + os.close(self._buf) + self._buf = None + + def replay(self) -> None: + if self.data: + try: + _write_all(2, self.data) + except OSError: + pass + + +_PARSE_LOCK = threading.Lock() # fd-2 redirection is process-wide + + +def _bind_walk(data: bytes, table: Printer) -> _BindWalk: + """Parse ``data`` with the bindings and flatten the module into + pure-Python records, reading the attributes ``table`` names. The context + is pinned on the module for the module's whole life, the module is + dropped before the context, and no binding object survives the call + (dropping the context first segfaults). After a successful parse, + whatever else reached fd 2 in the window is passed on.""" + from triton._C.libtriton import ir # TTIR dialects only: no backend (TTGIR) loading + + source = _ParseInput(data) + try: + path = source.path + assert path is not None + ctx = ir.context() + try: + ir.load_dialects(ctx) + capture = _StderrCapture() + mod = None + with _PARSE_LOCK, capture: + try: + mod = ir.parse_mlir_module(path, ctx) + except RuntimeError: + pass + if mod is None: + raise _parse_error(capture.text, path) + capture.replay() + try: + walk, failure = _pinned_extract(mod, ctx, table) + finally: + del mod # the module before its context + finally: + del ctx + finally: + source.close() + if walk is None: + raise MisalignedModule([f"bindings walk failed: {failure}"]) + ops, blocks, region_blocks = walk + return _BindWalk(ops, blocks, region_blocks, path) + + +def _pinned_extract(mod, ctx, table: Printer) -> tuple[Any, str | None]: + """Pin ``ctx`` on ``mod`` (proton's pattern: the module keeps its context + alive), then extract. Returns (records, None) or (None, reason); an + exception's traceback, whose frames hold binding objects, dies here while + the module and its context are still alive.""" + try: + mod.context = ctx + except Exception as e: # noqa: BLE001 + return ( + None, + f"cannot pin the MLIR context on the module: {type(e).__name__}: {e}", + ) + body: list[Any] = [] + try: + walk = _extract(mod, body, table) + except Exception as e: # noqa: BLE001 (bindings drift: an unexpected getter result) + return None, f"{type(e).__name__}: {e}" + # No binding erases a module, so each parse would leak its whole op tree + # (~20 KiB for a 12 KiB text); erasing the body block frees all but the + # empty module op. Safe here: extraction is over, ctx is pinned, and the + # body block is the one binding object still alive (dropped at once). + if body: + try: + body.pop().erase() + except Exception: # noqa: BLE001 (bindings drift: keep the bounded leak) + pass + return walk, None + + +_RE_DIAG_POS = re.compile(r'loc\("":(\d+):\d+\)') + + +def _parse_error(diag: str, path: str) -> ModuleParseError: + diag = diag.replace(path, "").strip() or "the MLIR parser rejected the text" + m = _RE_DIAG_POS.search(diag) + return ModuleParseError(diag, int(m.group(1)) if m else None) + + +def _extract( + mod, body: list[Any], table: Printer +) -> tuple[list[_BOp], dict[int, _BBlock], dict[int, list[int]]]: + """Copy the walked module into pure-Python records (post-order walk: + ops of a block arrive in program order, blocks of a region in order). + The module's body block (the one binding object kept) goes to ``body``, + for ``_pinned_extract`` to erase.""" + ops: list[_BOp] = [] + blocks: dict[int, _BBlock] = {} + region_blocks: dict[int, list[int]] = {} + + def cb(op) -> None: + blk = op.get_block() + bid = None + if blk is not None: + bid = blk.id() + rec = blocks.get(bid) + if rec is None: + args = [blk.get_argument(i) for i in range(blk.get_num_arguments())] + parent = blk.get_parent() + rid = parent.id() + if parent.get_parent_region() is None: + body.append(blk) # the module's body block + rec = blocks[bid] = _BBlock( + rid, + tuple(x.id() for x in args), + tuple(str(x.get_type()) for x in args), + tuple(str(x.get_loc()) for x in args), + [], + ) + region_blocks.setdefault(rid, []).append(bid) + rec.ops.append(len(ops)) + name = op.get_name() + opnds = [op.get_operand(i) for i in range(op.get_num_operands())] + res = [op.get_result(i) for i in range(op.get_num_results())] + regs = [op.get_region(i) for i in range(op.get_num_regions())] + attrs: dict[str, Any] = {} + for kind, aname in table.bind_attrs.get(name, ()): + getter = getattr(op, _GETTERS[kind], None) + if getter is None: # the table names a getter these bindings lack + raise TypeError( + f"{name}.{aname}: the bindings have no {_GETTERS[kind]} " + f"(the Triton {table.release} table reads it)" + ) + v = getter(aname) + if v is not None: + if not _is_exactly(v, _GETTER_TYPES[kind]): + raise TypeError( + f"{name}.{aname}: getter returned {type(v).__name__}" + ) + attrs[aname] = v + ops.append( + _BOp( + name, + tuple(x.id() for x in opnds), + tuple(str(x.get_type()) for x in opnds), + tuple(x.id() for x in res), + tuple(str(x.get_type()) for x in res), + tuple(str(x.get_loc()) for x in res), + tuple(x.id() for x in regs), + tuple(x.size() for x in regs), + bid, + attrs, + ) + ) + + mod.walk(cb) + return ops, blocks, region_blocks + + +# ─────────────────────────── alignment ─────────────────────────── + + +@dataclass +class _OpB: # mutable op builder + name: str + operands_raw: tuple[int, ...] + operand_types: tuple[str, ...] + results: tuple[int, ...] + result_types: tuple[str, ...] + attrs: dict[str, Any] + path: tuple[int, ...] + position: int + line_no: int | None + end_line: int | None + loc: SourceLoc | None + callers: tuple[SourceLoc, ...] + loc_name: str | None + implicit: bool + regions: list[tuple[int, ...]] + successor_labels: list[tuple[str, int]] | None = None + func_args: list[_ArgDef] | None = None + + +@dataclass +class _BlockB: + op: int + region: int + position: int + label: str | None + name: str # the printer's name: label, or ^bb0 for an unlabeled entry + args: tuple[int, ...] + arg_types: tuple[str, ...] + arg_names: tuple[str | None, ...] + ops: list[int] + path: tuple[int, ...] # path of the ops inside this block + preds: tuple[str, ...] | None + line_no: int + + +# what Module.stats counts +_STATS = ( + "ops", "implicit_ops", "blocks", "values", "funcs", "ssa_edges", "result_locs", "arg_locs", + "bind_attrs", "type_checks", "needed_attrs", "cf_edges", "pred_checks", +) # fmt: skip + + +class _Aligner: + def __init__(self, tree: _TextTree, bw: _BindWalk, table: Printer) -> None: + self.tree = tree + self.bw = bw + self.table = table + self.locp = _LocParser(tree.aliases, bw.path) + self.problems: list[tuple[int | None, str]] = [] + self.ops: list[_OpB] = [] + self.blocks: list[_BlockB] = [] + self.values: list[list[Any]] = [] # [type, op, block, position, name] + self.vmap: dict[int, int] = {} # raw value id -> value index + self.uses: list[ + tuple[int, list[str], list[str], tuple[dict[str, int], ...]] + ] = [] + self.region_labels: dict[tuple[int, int], dict[str, int]] = {} + # op lines with / without a printed trailing loc: the printer (debug + # info on) gives every op one, so a module mixing both is misread + self.loc_lines: list[int] = [] + self.no_loc_lines: list[int] = [] + self.stats: collections.Counter[str] = collections.Counter( + dict.fromkeys(_STATS, 0) + ) + + def bad(self, line: int | None, msg: str) -> None: + self.problems.append((line, f"line {line}: {msg}" if line is not None else msg)) + + # ── values ── + def new_value( + self, + raw: int, + type_: str, + op: int | None, + block: int | None, + pos: int, + loc: str, + ) -> int: + if raw in self.vmap: + raise ValueError("a value appears twice in the walk") + idx = len(self.values) + name = None + tree = self.locp.parse(loc) + if tree is not None: + name = _loc_site(tree)[2] + self.values.append([type_, op, block, pos, name]) + self.vmap[raw] = idx + return idx + + def text_loc(self, t: _TOp) -> tuple | None: + masked, raw = ( + (t.close_masked, t.close_raw) if t.opens_region else (t.masked, t.raw) + ) + span = _trailing_loc(masked) + (self.loc_lines if span is not None else self.no_loc_lines).append(t.line_no) + if span is None: + return None + return self.locp.parse(raw[span[0] : span[1]]) + + # ── traversal ── + def run(self) -> Module: + root_w = [i for i, b in enumerate(self.bw.ops) if b.block_id is None] + top = self.tree.root.regions[0][0].ops if self.tree.root.regions[0] else [] + if len(root_w) != 1 or len(top) != 1: + self.bad( + None, + f"expected one top-level op: text has {len(top)}, bindings {len(root_w)}", + ) + raise self.failure() + # task: ("op", text op, walk index, parent block index | None, position, scopes) + # or ("implicit", walk index, parent block index, position) + stack: list[tuple] = [("op", top[0], root_w[0], None, 0, ({},))] + while stack: + task = stack.pop() + if task[0] == "op": + children = self.visit(*task[1:]) + else: + children = [] + self.implicit(*task[1:]) + stack.extend(reversed(children)) + if len(self.ops) != len(self.bw.ops): + self.bad( + None, f"{len(self.ops)} aligned ops != {len(self.bw.ops)} walked ops" + ) + if self.loc_lines: + for line in self.no_loc_lines: + self.bad(line, "op prints no loc while the module prints locs") + if not self.problems: + self.check_uses() + self.check_successors() + if self.problems: + raise self.failure() + return self.freeze() + + def failure(self) -> MisalignedModule: + first = next((ln for ln, _ in self.problems if ln is not None), None) + return MisalignedModule([m for _, m in self.problems], first) + + def implicit(self, wi: int, block: int, pos: int) -> None: + b = self.bw.ops[wi] + idx = len(self.ops) + self.stats["implicit_ops"] += 1 + self.ops.append( + _OpB( + b.name, + (), + (), + (), + (), + {}, + self.blocks[block].path, + pos, + None, + None, + None, + (), + None, + True, + [], + ) + ) + self.blocks[block].ops.append(idx) + + def visit( + self, + t: _TOp, + wi: int, + block: int | None, + pos: int, + scopes: tuple[dict[str, int], ...], + ) -> list: + b = self.bw.ops[wi] + line = t.line_no + idx = len(self.ops) + path = self.blocks[block].path if block is not None else () + if block is not None: + self.blocks[block].ops.append(idx) + structural_ok = True + if t.name != b.name: + self.bad(line, f"text op {t.name!r} != walked op {b.name!r}") + structural_ok = False + if t.n_results != len(b.results): + self.bad( + line, + f"{t.name}: {t.n_results} printed results != {len(b.results)} walked", + ) + structural_ok = False + uses, before, hdefs = _header_uses_and_defs(t) + if len(uses) != len(b.operands): + self.bad( + line, + f"{t.name}: {len(uses)} printed operands != {len(b.operands)} walked", + ) + # printed regions: trailing empty regions may be omitted + if len(t.regions) > len(b.region_ids) or any( + b.region_sizes[i] for i in range(len(t.regions), len(b.region_ids)) + ): + self.bad( + line, + f"{t.name}: {len(t.regions)} printed regions vs walked sizes {list(b.region_sizes)}", + ) + structural_ok = False + # locs + try: + tloc = self.text_loc(t) + for k, s in enumerate(b.result_locs): + tree = self.locp.parse(s) + if tree is not None: + self.stats["result_locs"] += 1 + if tree != tloc: + self.bad( + line, + f"{t.name}: result {k} loc differs: text {tloc} vs bindings {tree}", + ) + except ValueError as e: + self.bad(line, f"{t.name}: unreadable loc: {e}") + tloc = None + site, callers, lname = _loc_site(tloc) + # attributes + attrs: dict[str, Any] = {} + extra: dict[str, Any] = {} + if structural_ok: + try: + attrs, extra = _text_attrs(t, b.result_types, self.table) + except ValueError as e: + self.bad(line, f"{t.name}: {e}") + self.check_attrs(b, attrs, line) + rec = _OpB( + b.name, + b.operands, + b.operand_types, + (), + b.result_types, + attrs, + path, + pos, + line, + t.close_line if t.opens_region else line, + site, + callers, + lname, + False, + [], + successor_labels=extra.get("successors"), + func_args=extra.get("args"), + ) + self.ops.append(rec) + rec.results = tuple( + self.new_value(raw, b.result_types[k], idx, None, k, b.result_locs[k]) + for k, raw in enumerate(b.results) + ) + if len(t.result_names) == len(b.results): + for name, raw in zip(t.result_names, b.results): + scopes[-1][name] = raw + self.uses.append((idx, uses, before, scopes)) + if ( + b.name == "tt.func" + and rec.func_args is not None + and [d.name for d in rec.func_args] != hdefs + ): + self.bad( + line, + f"tt.func: parameter list {[d.name for d in rec.func_args]} vs header defs {hdefs}", + ) + if not structural_ok: + rec.regions = [() for _ in b.region_ids] + return [] + if hdefs and not b.region_ids: + self.bad(line, f"{t.name}: header defines {hdefs} but the op has no region") + children: list[tuple] = [] + for ri, rid in enumerate(b.region_ids): + children += self.region(t, b, idx, ri, rid, hdefs, rec, scopes) + return children + + def region( + self, t: _TOp, b: _BOp, idx: int, ri: int, rid: int, hdefs, rec: _OpB, scopes + ) -> list: + line = t.line_no + wblocks = self.bw.region_blocks.get(rid, []) + if len(wblocks) != b.region_sizes[ri]: + self.bad(line, f"{t.name}: region {ri} has a block without ops") + tblocks = t.regions[ri] if ri < len(t.regions) else [] + if not tblocks and len(wblocks) == 1 and b.name in self.table.elides_yield: + only = self.bw.blocks[wblocks[0]].ops + tail = self.bw.ops[only[-1]] if len(only) == 1 else None + if tail is not None and tail.name == "scf.yield" and not tail.operands: + # `{ }`: the region's one block held only the elided yield + tblocks = [_TBlock(None, [], [], line)] + if len(tblocks) != len(wblocks): + self.bad( + line, + f"{t.name}: region {ri}: {len(tblocks)} printed blocks != {len(wblocks)} walked", + ) + rec.regions.append(()) + return [] + region_scope: dict[str, int] = {} + inner = scopes + (region_scope,) + labels: dict[str, int] = {} + self.region_labels[(idx, ri)] = labels + keys = [] + children: list[tuple] = [] + for bi, (tb, rbid) in enumerate(zip(tblocks, wblocks)): + bb = self.bw.blocks[rbid] + bidx = len(self.blocks) + keys.append(bidx) + name = tb.label if tb.label is not None else "^bb0" + if tb.label is None and bi != 0: + self.bad(tb.line_no, f"{t.name}: unlabeled block {ri}.{bi}") + if name in labels: + self.bad(tb.line_no, f"{t.name}: duplicate block label {name}") + labels[name] = bidx + blk = _BlockB( + idx, + ri, + bi, + tb.label, + name, + (), + bb.arg_types, + (), + [], + rec.path + (bidx,), + tb.preds, + tb.line_no, + ) + self.blocks.append(blk) + blk.args = tuple( + self.new_value(raw, bb.arg_types[ai], None, bidx, ai, bb.arg_locs[ai]) + for ai, raw in enumerate(bb.args) + ) + blk.arg_names = tuple(self.values[v][4] for v in blk.args) + # printed arguments: the label's, or the op header's for the + # implicit entry block of region 0 (func args, iv + iter_args, + # scf.while inits) + if tb.label is not None: + printed = tb.args + elif bi == 0 and ri == 0: + printed = ( + rec.func_args + if rec.func_args is not None + else [_ArgDef(n, None, None, {}) for n in hdefs] + ) + else: + printed = [] + if len(printed) != len(bb.args): + self.bad( + tb.line_no, + f"{t.name}: block {ri}.{bi}: {len(printed)} printed args != {len(bb.args)} walked", + ) + else: + for ai, (d, raw) in enumerate(zip(printed, bb.args)): + region_scope[d.name] = raw + self.check_arg(d, bb, ai, tb.line_no, t.name) + # ops, with the one tolerated gap: an elided trailing scf.yield + n_text, n_walk = len(tb.ops), len(bb.ops) + elided = False + if n_text != n_walk: + tail = self.bw.ops[bb.ops[-1]] if bb.ops else None + elided = ( + n_text == n_walk - 1 + and tail is not None + and tail.name == "scf.yield" + and not tail.operands + and b.name in self.table.elides_yield + ) + if not elided: + self.bad( + tb.line_no, + f"{t.name}: block {ri}.{bi}: {n_text} printed ops != {n_walk} walked", + ) + continue + for p, (tchild, wchild) in enumerate(zip(tb.ops, bb.ops)): + children.append(("op", tchild, wchild, bidx, p, inner)) + if elided: + children.append(("implicit", bb.ops[-1], bidx, n_text)) + rec.regions.append(tuple(keys)) + return children + + def check_arg( + self, d: _ArgDef, bb: _BBlock, ai: int, line: int, owner: str + ) -> None: + if d.type is not None and d.type != bb.arg_types[ai]: + self.bad( + line, + f"{owner}: argument {d.name}: printed type {d.type!r} != {bb.arg_types[ai]!r}", + ) + try: + bind = self.locp.parse(bb.arg_locs[ai]) + text = self.locp.parse(d.loc) if d.loc is not None else None + except ValueError as e: + self.bad(line, f"{owner}: argument {d.name}: unreadable loc: {e}") + return + if bind is not None or text is not None: + self.stats["arg_locs"] += 1 + if bind != text: + self.bad( + line, + f"{owner}: argument {d.name}: loc differs: text {text} vs bindings {bind}", + ) + + def check_attrs(self, b: _BOp, attrs: dict[str, Any], line: int) -> None: + name = b.name + table = self.table + for key, value in list(attrs.items()): + vocab = table.vocab.get((name, key)) + if vocab is not None and not (isinstance(value, str) and value in vocab): + self.bad(line, f"{name}: {key} {value!r} outside the closed vocabulary") + for key, default in table.defaults.get(name, {}).items(): + attrs.setdefault(key, default) + for key, value in attrs.items(): + want = table.attr_types.get((name, key)) + if want is not None and not _is_exactly(value, want): + self.bad(line, f"{name}: {key} {value!r} is not a {want.__name__}") + if ( + name in ("tt.get_program_id", "tt.get_num_programs") + and attrs.get("axis") in _AXES + ): + attrs["axis"] = _AXES[attrs["axis"]] + for key, bv in b.attrs.items(): + self.stats["bind_attrs"] += 1 + akey = _BIND_KEY.get(key, key) + tv = attrs.get(akey) + ints = table.bind_ints.get((name, key)) + if ints is not None: # a keyword the bindings read as its integer + tv = ints.get(tv) if isinstance(tv, str) else None + if tv != bv or not _is_exactly(tv, type(bv)): + self.bad( + line, + f"{name}: attr {key}: text {attrs.get(akey)!r} vs bindings {bv!r}", + ) + try: + tc = _type_checks(name, attrs, b.operand_types, b.result_types) + except (IndexError, ValueError) as e: + tc = [f"type check failed: {e}"] + if tc is not None: + self.stats["type_checks"] += 1 + for msg in tc: + self.bad(line, f"{name}: {msg}") + if name in ("cf.br", "cf.cond_br"): + return # successors are checked once every block is known + for key in table.needed.get(name, ()): + self.stats["needed_attrs"] += 1 + if attrs.get(key) is None: + self.bad(line, f"{name}: attribute {key!r} not recovered") + + def check_uses(self) -> None: + for idx, uses, before, scopes in self.uses: + rec = self.ops[idx] + order = self.table.printed_order.get(rec.name) + if order is not None: + uses = order(uses, before) + got = [] + for u in uses: + got.append(next((d[u] for d in reversed(scopes) if u in d), None)) + self.stats["ssa_edges"] += len(uses) + if tuple(got) != rec.operands_raw: + unresolved = [u for u, v in zip(uses, got) if v is None] + what = ( + f"unresolved {unresolved}" + if unresolved + else "operand values differ" + ) + if not unresolved and sorted(got) == sorted(rec.operands_raw): # type: ignore[type-var] + what = "operand ORDER differs (same multiset)" + self.bad(rec.line_no, f"{rec.name}: SSA edges: {what} ({uses})") + + def check_successors(self) -> None: + got: dict[int, list[str]] = collections.defaultdict(list) + for rec in self.ops: + groups = rec.successor_labels + if groups is None: + continue + parent = self.blocks[rec.path[-1]] + labels = self.region_labels[(parent.op, parent.region)] + succ = [] + n_dest_args = 0 + for label, n_printed in groups: + dest = labels.get(label) + if dest is None: + self.bad(rec.line_no, f"{rec.name}: unknown successor {label}") + return + n = len(self.blocks[dest].args) + if n_printed != n: + self.bad( + rec.line_no, + f"{rec.name}: {label}: {n_printed} successor operands != {n} block args", + ) + n_dest_args += n + succ.append(dest) + got[dest].append(parent.name) + own = 1 if rec.name == "cf.cond_br" else 0 + if len(succ) != (2 if rec.name == "cf.cond_br" else 1): + self.bad(rec.line_no, f"{rec.name}: {len(succ)} successors") + if len(rec.operands_raw) != own + n_dest_args: + self.bad( + rec.line_no, + f"{rec.name}: {len(rec.operands_raw)} operands != {own} + {n_dest_args} successor args", + ) + self.stats["cf_edges"] += len(succ) + rec.attrs["successors"] = tuple(succ) + self.stats["needed_attrs"] += 1 + # the printer's predecessor comments: a second, printer-computed CFG + for bidx, blk in enumerate(self.blocks): + want = blk.preds + if want is None: + if ( + self.tree.pred_comments + and blk.label is not None + and blk.position != 0 + ): + self.bad(blk.line_no, f"block {blk.name}: no predecessor comment") + if blk.position == 0 and got.get(bidx): + self.bad( + blk.line_no, + f"entry block {blk.name} has predecessors {got[bidx]}", + ) + continue + self.stats["pred_checks"] += 1 + if collections.Counter(want) != collections.Counter(got.get(bidx, [])): + self.bad( + blk.line_no, + f"block {blk.name}: printer preds {sorted(want)} != {sorted(got.get(bidx, []))}", + ) + + def freeze(self) -> Module: + values = tuple(Value(i, *v) for i, v in enumerate(self.values)) + blocks = tuple( + Block( + i, + b.op, + b.region, + b.position, + b.label, + b.args, + b.arg_types, + b.arg_names, + tuple(b.ops), + ) + for i, b in enumerate(self.blocks) + ) + ops = [] + funcs = [] + for i, r in enumerate(self.ops): + operands = tuple(self.vmap[v] for v in r.operands_raw) + ops.append( + Op( + i, + r.name, + operands, + r.operand_types, + r.results, + r.result_types, + _FrozenMap(r.attrs), + tuple(r.regions), + r.path, + r.position, + r.line_no, + r.end_line, + r.loc, + r.callers, + r.loc_name, + r.implicit, + ) + ) + if r.name == "tt.func": + args: tuple[FuncArg, ...] = () + if r.regions and r.regions[0]: + entry = blocks[r.regions[0][0]] + printed = r.func_args or [] + args = tuple( + FuncArg( + k, + v, + entry.arg_types[k], + entry.arg_names[k], + _FrozenMap(printed[k].attrs if k < len(printed) else {}), + ) + for k, v in enumerate(entry.args) + ) + funcs.append(Func(i, r.attrs["sym_name"], r.attrs["visibility"], args)) + self.stats.update( + ops=len(ops), blocks=len(blocks), values=len(values), funcs=len(funcs) + ) + return Module( + tuple(ops), + blocks, + values, + tuple(funcs), + _FrozenMap(self.stats), + self.table.release, + ) + + +# ─────────────────────────── entry point ─────────────────────────── + + +def _walk( + text: str, scan_text: str | None = None, *, table: Printer | None = None +) -> Module: + """Uncached walk with ``table`` (default: the installed Triton's + ``printer()``). ``scan_text`` (tests only) feeds the text layer a + different string than the bindings parse, to prove the checks catch a + text layer that mis-reads the module.""" + if table is None: + table = printer() + try: + data = text.encode("utf-8") + except UnicodeEncodeError as e: + raise ModuleParseError(f"the text is not encodable as UTF-8: {e}") from None + _screen(text, table) + bw = _bind_walk(data, table) + try: + tree = _scan_text(text if scan_text is None else scan_text) + except _TextError as e: + raise MisalignedModule( + [f"line {e.line_no}: {e.msg}" if e.line_no else e.msg], e.line_no + ) from None + try: + return _Aligner(tree, bw, table).run() + except MisalignedModule: + raise + except (ValueError, KeyError, IndexError, AssertionError, TypeError) as e: + # a text shape the aligner does not model: fail closed + raise MisalignedModule([f"aligner: {type(e).__name__}: {e}"]) from None + + +_CACHE_SIZE = 32 +_CACHE: collections.OrderedDict[ + tuple[bytes, str], Module | tuple[tuple[str, ...], int | None] +] = collections.OrderedDict() +_CACHE_LOCK = threading.Lock() + + +def walk_module(text: str) -> Module: + """Walk one printed TTIR module (see the module docstring) with the + installed Triton's ``Printer`` table. + + Raises ``UnknownTritonRelease`` when that release has no table, + ``MisalignedModule`` when the text layer and the bindings disagree and + ``ModuleParseError`` when the MLIR parser rejects the text (or the text + holds a construct the release's parser cannot be handed). Results (and + misalignments) are cached by sha256 of the text and the release, so a + text is parsed once while it stays among the last ``_CACHE_SIZE`` + distinct texts. + """ + table = printer() + key = ( + hashlib.sha256(text.encode("utf-8", "surrogatepass")).digest(), + table.release, + ) + with _CACHE_LOCK: + hit = _CACHE.get(key) + if hit is not None: + _CACHE.move_to_end(key) + if hit is None: + try: + hit = _walk(text, table=table) + except MisalignedModule as e: + hit = (e.problems, e.line_no) + with _CACHE_LOCK: + _CACHE[key] = hit + while len(_CACHE) > _CACHE_SIZE: + _CACHE.popitem(last=False) + if isinstance(hit, Module): + return hit + raise MisalignedModule(*hit) + + +def _after_fork_in_child() -> None: + """A fork while another thread holds ``_PARSE_LOCK`` / ``_CACHE_LOCK`` + leaves the child a lock nobody releases (its next walk would hang), and a + fork inside a parse window leaves the child's fd 2 in the capture buffer. + The child gets fresh locks and its fd 2 back; the capture's two fds stay + open (their owner is the parent's thread, which does not run here).""" + global _PARSE_LOCK, _CACHE_LOCK, _REDIRECT + _PARSE_LOCK = threading.Lock() + _CACHE_LOCK = threading.Lock() + redirect, _REDIRECT = _REDIRECT, None + if redirect is not None: + try: + os.dup2(redirect[0], 2) + except OSError: + pass + + +if hasattr(os, "register_at_fork"): + os.register_at_fork(after_in_child=_after_fork_in_child) diff --git a/tilelens/ir/capture.py b/tilelens/ir/capture.py new file mode 100644 index 000000000..3d7955bfd --- /dev/null +++ b/tilelens/ir/capture.py @@ -0,0 +1,282 @@ +"""Per-launch compiled artifacts, and a content-addressed parse cache (the L2 layer). + +``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 + + +def _triton_version() -> str: + import triton + + return triton.__version__ + + +_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, options and Triton version. + + ``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 the first lookup, not with this module. 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())), + _triton_version(), + ) + cached = self._outcomes.get(key) + except Exception as exc: + return ParseOutcome(error=_describe(exc)) + if cached is not None: + return cached + try: + outcome = ParseOutcome(graph=reader(text, **options)) + except Exception as exc: + refusal = self._refusal if self._refusal is not None else _default_refusal() + if refusal is None or not isinstance(exc, refusal): + return ParseOutcome(error=_describe(exc)) + outcome = ParseOutcome(refusal=_without_frames(exc, outer)) + self._outcomes[key] = outcome + return outcome diff --git a/tilelens/ir/client.py b/tilelens/ir/client.py new file mode 100644 index 000000000..26083197e --- /dev/null +++ b/tilelens/ir/client.py @@ -0,0 +1,147 @@ +"""The IR client base (the L5 layer): 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: + +1. the D10b version gate: outside the tested Triton window (and without + ``TILELENS_IR_ALLOW_UNTESTED_TRITON=1``), ``on_refusal`` gets a + ``Refusal`` of kind ``"untested-triton-version"``; +2. otherwise ``analyze_launch(log)`` returns the reports and the verdict, + also for a launch nothing could be captured for (see analyze_launch); +3. an ``Exception`` from either 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, D26; the configured default otherwise) and +implements the three 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.config import TESTED_TRITON_VERSIONS, untested_triton_version +from ..core.data import Op +from .capture import ArtifactLog +from .verdict import IRVerdict, Refusal + + +class IRClient(Client): + NEEDS_INTERPRETER: ClassVar[bool] = False + + def __init__(self) -> None: + super().__init__() + self.artifacts = ArtifactLog(self.IR_STAGES) + # Compatibility view of the last finalized launch's verdict. + 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; the other cause, an untested Triton, is refused by + the version gate first). 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``.""" + + @abstractmethod + def on_refusal(self, refusal: Refusal) -> IRVerdict: + """The verdict for a launch the base refused to analyze (kind + ``"untested-triton-version"``).""" + + # ── 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: + version = untested_triton_version() + if version is not None: + tested = ", ".join(f"{v}.x" for v in TESTED_TRITON_VERSIONS) + return [], self.on_refusal( + Refusal( + kind="untested-triton-version", + message=( + f"IR mode is tested on Triton {tested}, not {version}; " + "set TILELENS_IR_ALLOW_UNTESTED_TRITON=1 to run it anyway" + ), + ) + ) + reports, verdict = self.analyze_launch(self.artifacts) + return list(reports), verdict + except Exception as exc: + return [], self.on_analysis_error(exc) + + # ── inert interpreter path: the core calls none of these for an IR + # client except the warmup vote, which declines (IR compiles go through + # ir_capture) ── + + def pre_run_callback(self, fn: Callable) -> bool: + return False + + def post_run_callback(self, fn: Callable) -> bool: + return False + + def arg_callback(self, name: str, arg: Any, arg_cvt: Any) -> None: + pass + + def grid_callback(self, grid: tuple[int, ...]) -> None: + pass + + def grid_idx_callback(self, grid_idx: tuple[int, ...]) -> None: + pass + + def register_op_callback( + self, op_type: type[Op], *args: Any, **kwargs: Any + ) -> OpCallbacks: + return OpCallbacks() + + def register_for_loop_callback(self) -> ForLoopCallbacks: + return ForLoopCallbacks() + + def pre_warmup_callback(self, jit_fn: Callable, *args: Any, **kwargs: Any) -> bool: + return False + + def post_warmup_callback(self, jit_fn: Callable, ret: Any) -> None: + pass diff --git a/tilelens/ir/launch.py b/tilelens/ir/launch.py new file mode 100644 index 000000000..1727b319b --- /dev/null +++ b/tilelens/ir/launch.py @@ -0,0 +1,227 @@ +"""What one traced launch bound its kernel parameters to (the L3 layer). + +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/lowering.py b/tilelens/ir/lowering.py new file mode 100644 index 000000000..00311a803 --- /dev/null +++ b/tilelens/ir/lowering.py @@ -0,0 +1,291 @@ +"""Term -> Z3 lowering shared by the compiled-mode clients (D13). + +The one reading of the TTIR reader's term algebra (``ttir_reader``'s +``Term``) as Z3 integers and booleans, for every client that queries an +``AccessGraph`` with Z3. Mechanism only: the operator semantics live here, +every leaf's meaning comes from the client. A client's :class:`TermLeaves` +says what a scalar argument, a program id, the grid, a lane, a loop's +iteration, an atomic observation and an unmodeled value are, and a client +that cannot model a leaf raises its own exception from it, which propagates +unchanged. The lowering itself raises only on a malformed graph, a bug and +never a limit of the model: ``TypeError`` or ``ValueError`` for a term or op +outside the algebra, or for a ``LoopVar`` (or an ``IterArgOffset`` whose +IterArgInfo names no loop) in a graph without a loop, and ``IndexError`` for +an ``arg_id`` outside ``iter_args``. + +Operator semantics, the reader's integer model (see ``ttir_reader``): + +* ``+ - *`` are unbounded Int arithmetic; ``//`` and ``%`` truncate toward + zero (``arith.divsi`` / ``remsi``: the remainder takes the dividend's + sign); ``min`` / ``max`` pick an operand; +* the unsigned ops (``u//``, ``u%``, ``umin``, ``umax``) and predicates + (``ult``, ...) read as their signed twins, and an ``IntCast`` as its + operand: exact only where the access's ``width_obligations`` hold, which + the client discharges (a zero divisor is the client's concern too); +* an i1 value is a Bool in a boolean position (a mask or path, a Select's + condition, ``and`` / ``or`` / negation) and 0/1 in an integer one, and a + Select whose arms differ in sort is an Int; +* ``LoopVar`` is ``lower + k * step`` and ``IterArgOffset`` is ``offset0 + + k * delta``, with ``k`` the client's iteration of that loop. + +Z3 context: the lowering never creates one. Constants are made in +``leaves.ctx`` (None: Z3's main context) and every other term in its +operands' context, so all of a client's terms share the context of its +leaves (the compiled sanitizer's is a Context per check). + +Terms can be deeper than Python's recursion limit and their generated +``==`` / ``hash`` recurse, so the walk is iterative and memoized by term +identity. +""" + +from __future__ import annotations + +from collections.abc import Callable, Sequence +from typing import Any, Protocol + +from z3 import And, ArithRef, BoolRef, Context, If, IntVal, Or, is_bool +from z3 import Not as Z3Not + +from .ttir_reader import ( + AccessGraph, + Arange, + Bin, + BoolBin, + Cmp, + Const, + DataDep, + IntCast, + IterArgOffset, + LoopInfo, + LoopVar, + Not, + NumPrograms, + Observed, + Param, + Pid, + Select, +) + + +class TermLeaves(Protocol): + """A client's meaning of every leaf of the algebra (there are no + defaults): each method returns a Z3 Int in ``ctx``, or raises the + client's own refusal.""" + + @property + def ctx(self) -> Context | None: + """The Z3 context of the client's terms (None: the main one).""" + + def param(self, t: Param) -> ArithRef: + """A scalar kernel argument.""" + + def pid(self, t: Pid) -> ArithRef: + """The program id along ``t.axis``.""" + + def num_programs(self, t: NumPrograms) -> ArithRef: + """The grid size along ``t.axis``.""" + + def arange(self, t: Arange) -> ArithRef: + """``t``'s value at the lane a query reads (the reader's contract: + key a lane by ``(dim, end - start)``, see ``Arange``).""" + + def iteration(self, loop_ssa: str) -> ArithRef: + """The iteration index ``k`` of the loop ``loop_ssa``.""" + + def observed(self, t: Observed) -> ArithRef: + """The old value the atomic at ``t.access_index`` observed.""" + + def data_dep(self, t: DataDep) -> ArithRef: + """A value the reader could not model (``t.why`` says which).""" + + +def as_bool(e: Any) -> BoolRef: + """An i1 value in a boolean position: i1 constants (e.g. the dense + mask of an unmasked atomic, Const(1)) lower to Int.""" + return e if is_bool(e) else e != 0 + + +def as_int(e: Any) -> ArithRef: + """An i1 value in an integer position (an extui of a compare, ...).""" + return If(e, IntVal(1, e.ctx), IntVal(0, e.ctx)) if is_bool(e) else e + + +def _trunc_div(a: ArithRef, b: ArithRef) -> ArithRef: + """arith.divsi rounds toward zero, but Z3's Int ``/`` is Euclidean + (floor for a positive divisor): they disagree on negative dividends. + Divide the magnitudes, where the two agree, and re-apply the sign.""" + aa = If(a >= 0, a, -a) + ab = If(b >= 0, b, -b) + q = aa / ab + return If((a >= 0) == (b >= 0), q, -q) + + +def _divrem_ir(op: str, a: ArithRef, b: ArithRef) -> ArithRef: + """The IR's quotient (``op`` ``//`` or ``u//``) or remainder (``%`` or + ``u%``): truncating, so the remainder has the dividend's sign.""" + q = _trunc_div(a, b) + return q if op in ("//", "u//") else a - b * q + + +def _bin(op: str, a: ArithRef, b: ArithRef) -> ArithRef: + # The unsigned ops read their operands unsigned; their width + # obligations make both non-negative, where they equal the signed ones. + if op == "+": + return a + b + if op == "-": + return a - b + if op == "*": + return a * b + if op in ("//", "u//", "%", "u%"): + return _divrem_ir(op, a, b) + if op in ("min", "umin"): + return If(a <= b, a, b) + if op in ("max", "umax"): + return If(a >= b, a, b) + raise ValueError(f"unknown integer op {op!r}") + + +# Unsigned predicates read their operands unsigned; their width obligations +# make both non-negative, where they equal the signed ones. +_SIGNED_TWIN = {"ult": "slt", "ule": "sle", "ugt": "sgt", "uge": "sge"} + + +def _cmp(pred: str, a: Any, b: Any) -> BoolRef: + if is_bool(a) or is_bool(b): # i1 operands, as 0/1 + a, b = as_int(a), as_int(b) + p = _SIGNED_TWIN.get(pred, pred) + if p == "slt": + return a < b + if p == "sle": + return a <= b + if p == "sgt": + return a > b + if p == "sge": + return a >= b + if p == "eq": + return a == b + if p == "ne": + return a != b + raise ValueError(f"unknown cmpi predicate {pred!r}") + + +def _loop(graph: AccessGraph, t: object) -> LoopInfo: + if graph.loop is None: + raise ValueError( + f"kernel {graph.kernel_name!r}: a loop term ({type(t).__name__}) " + "without a loop" + ) + return graph.loop + + +def children(t: object, graph: AccessGraph) -> tuple: + """The terms ``t``'s value is computed from: its operands; for a + loop-carried pointer's offset, its IterArgInfo's ``offset0`` and + ``delta``; for the induction variable, the loop's ``lower`` and + ``step``. A DataDep has none: its ``keep`` is not its value (the + reader's ``observed_indices`` is the relation that reaches it).""" + if isinstance(t, (Bin, Cmp, BoolBin)): + return (t.a, t.b) + if isinstance(t, Select): + return (t.cond, t.t, t.f) + if isinstance(t, Not): + return (t.a,) + if isinstance(t, IntCast): + return (t.x,) + if isinstance(t, IterArgOffset): + info = graph.iter_args[t.arg_id] + return (info.offset0, info.delta) + if isinstance(t, LoopVar): + loop = _loop(graph, t) + return (loop.lower, loop.step) + return () + + +def fold( + root: object, + graph: AccessGraph, + apply: Callable[[object, Sequence[Any]], Any], + memo: dict[int, tuple[object, Any]], +) -> Any: + """``apply(t, the values of children(t))`` at ``root``, children first, + each term once for ``memo`` (id(term) -> (term, value): holding the + term keeps its id unique). Iterative post-order.""" + stack: list[tuple[object, bool]] = [(root, False)] + while stack: + t, ready = stack.pop() + if id(t) in memo: + continue + kids = children(t, graph) + if kids and not ready: + stack.append((t, True)) + stack.extend((k, False) for k in reversed(kids) if id(k) not in memo) + continue + memo[id(t)] = (t, apply(t, [memo[id(k)][1] for k in kids])) + return memo[id(root)][1] + + +class Lowerer: + """The lowering of one family of a client's queries: bound to one graph + and one leaves object (so one Z3 context), with one memo.""" + + def __init__(self, graph: AccessGraph, leaves: TermLeaves) -> None: + self.graph = graph + self.leaves = leaves + self._memo: dict[int, tuple[object, Any]] = {} + + def lower(self, term: object) -> Any: + """``term`` as a Z3 expression: an Int, or a Bool for a compare, + ``and`` / ``or`` / negation, and a Select of Bools.""" + return fold(term, self.graph, self._apply, self._memo) + + def value(self, term: object) -> ArithRef: + return as_int(self.lower(term)) + + def cond(self, term: object) -> BoolRef: + return as_bool(self.lower(term)) + + def _apply(self, t: object, kids: Sequence[Any]) -> Any: + leaves = self.leaves + if isinstance(t, Const): + return IntVal(t.value, leaves.ctx) + if isinstance(t, Param): + return leaves.param(t) + if isinstance(t, Pid): + return leaves.pid(t) + if isinstance(t, NumPrograms): + return leaves.num_programs(t) + if isinstance(t, Arange): + return leaves.arange(t) + if isinstance(t, LoopVar): + k = leaves.iteration(t.loop_ssa) + return as_int(kids[0]) + k * as_int(kids[1]) + if isinstance(t, IterArgOffset): + loop_ssa = self.graph.iter_args[t.arg_id].loop_ssa + k = leaves.iteration(loop_ssa or _loop(self.graph, t).loop_ssa) + return as_int(kids[0]) + k * as_int(kids[1]) + if isinstance(t, Bin): + return _bin(t.op, as_int(kids[0]), as_int(kids[1])) + if isinstance(t, Cmp): + return _cmp(t.pred, kids[0], kids[1]) + if isinstance(t, BoolBin): + a, b = as_bool(kids[0]), as_bool(kids[1]) + if t.op == "and": + return And(a, b) + if t.op == "or": + return Or(a, b) + raise ValueError(f"unknown boolean op {t.op!r}") + if isinstance(t, Select): + a, b = kids[1], kids[2] + if is_bool(a) != is_bool(b): + a, b = as_int(a), as_int(b) + return If(as_bool(kids[0]), a, b) + if isinstance(t, Not): + return Z3Not(as_bool(kids[0])) + if isinstance(t, IntCast): + # Its value is the operand's while its width obligation holds. + return as_int(kids[0]) + if isinstance(t, Observed): + return leaves.observed(t) + if isinstance(t, DataDep): + return leaves.data_dep(t) + raise TypeError(f"unknown term {type(t).__name__}") diff --git a/tilelens/ir/ttir_reader.py b/tilelens/ir/ttir_reader.py new file mode 100644 index 000000000..2a5035d5f --- /dev/null +++ b/tilelens/ir/ttir_reader.py @@ -0,0 +1,1756 @@ +"""TTIR reader shared by the compiled-mode clients. + +Reads the pre-optimization Triton IR (TTIR) of one kernel specialization +into an ``AccessGraph``: the kernel's function arguments, every global +memory access (``tt.load`` / ``tt.store`` / ``tt.atomic_rmw`` / +``tt.atomic_cas``) as an *element offset* expression relative to a base +pointer argument, the mask guarding it, and the loop structure. Scalar +arguments (``n_elements``, ``M``, strides, ...) stay symbolic (``Param`` +nodes) and are substituted with concrete launch values later; +``tl.constexpr`` values are already folded into TTIR constants. + +Why TTIR (not TTGIR): element addressing is cleanest here, before +layouts/pipelining add noise, and TTIR has no indirect loads unless the +kernel itself gathers — the data-dependent case, marked with ``DataDep``. + +This module is mechanism-only: it reads and flags (``DataDep`` markers, +``guarded`` accesses, width obligations, ``UnsupportedTTIR``); what to do +about a flagged or unsupported kernel is the policy of each client that +consumes the graph. It either represents the IR faithfully or raises an +``UnsupportedTTIR`` whose ``kind`` says what it cannot represent. + +Structure comes from ``_mlir_walk`` (the MLIR bindings plus the aligned +text layer): the reader walks its op tree over regions and blocks, and its +environment is keyed by the walk's value indices, never by printed SSA +names. Ported from the #361 regex reader (``parse_ttir(multipath=False)``; +the layout below stays diffable with it), with the audit's soundness fixes +built in: loop-variant and swapped pointer advances, ``tt.call``, +graph-aware observation walkers, integer widths and casts (D9), and inline +asm that is impure or handed an address. + +Address model: ``tt.addptr(base, off)`` accumulates an ELEMENT offset; the +byte address is ``base.data_ptr() + offset * elem_size``. An access is OOB +iff, for some program id / arange lane / loop iteration with its mask true, +the element offset escapes ``[0, numel)`` of its base tensor. + +Integer model (D9): every integer term denotes the IR value's SIGNED +reading as an unbounded integer (i1 terms are booleans, 0/1). That reading +is exact while the ``width_obligations`` of the access hold (they include +the loop's own increment); a consumer that evaluates terms with unbounded +integers must discharge (or report) them. ``Param`` values are the IR's +signed reading of the argument too. Division or remainder by zero is +undefined in the IR rather than a width condition, and so is an scf.for +step that is not positive: no obligation excludes them, that is the +consumer's call. + +Terms are frozen dataclasses with the generated ``==`` / ``hash`` / +``repr``, which recurse: on a term deeper than Python's recursion limit +(``kernel_deep_chain`` has more than 1000 levels) they, and +``copy.deepcopy``, raise ``RecursionError``; only ``pickle`` and this +module's own walkers are iterative. Consumers of such graphs key memo +tables by identity. +""" + +from __future__ import annotations + +import re +from dataclasses import dataclass, field, replace +from enum import Enum +from typing import Callable, Iterable, Iterator, NoReturn, Sequence + +from ._mlir_walk import ( + MisalignedModule, + Module, + ModuleParseError, + Op, + SourceLoc, + UnknownTritonRelease, + parse_type, + walk_module, +) + + +class TTIRKind(str, Enum): + """What the reader cannot represent. Only representational limits live + here; refusals a client makes about a graph it did receive (a + data-dependent mask, a CAS value, ...) are that client's own kinds.""" + + INDIRECT_ADDRESS = "indirect-address" + DATA_DEPENDENT_BOUND = "data-dependent-bound" + NESTED_LOOP = "nested-loop" + CONTROL_FLOW = "control-flow" + BLOCK_POINTER = "block-pointer" + OUT_OF_VOCABULARY = "out-of-vocabulary" + CALL = "call" + LOOP_VARIANT_ADVANCE = "loop-variant-advance" + INLINE_ASM = "inline-asm" + READER_MISALIGNMENT = "reader-misalignment" + UNPARSABLE = "unparsable" + # the installed Triton's release has no table (walk layer or reader): + # its TTIR is not read at all + UNTESTED_TRITON_VERSION = "untested-triton-version" + OTHER = "other" + + def __str__(self) -> str: + return self.value + + def __format__(self, spec: str) -> str: + return format(self.value, spec) + + +class UnsupportedTTIR(Exception): + """Raised for constructs outside the compiled-mode model (indirect or + data-dependent addressing, block pointers, nested loops, calls, ...). + + ``kind`` (a :class:`TTIRKind`) is the machine-readable class of the + limitation, ``message`` the human-readable detail, ``line_no`` the + refused op's line in the TTIR text and ``loc`` its user-source location + (None when unknown). Clients read these fields; ``str(exc)`` is the + message alone. + """ + + def __init__( + self, + kind: TTIRKind | str, + message: str, + *, + line_no: int | None = None, + loc: SourceLoc | None = None, + ) -> None: + super().__init__(message) + self.kind = TTIRKind(kind) + self.message = message + self.line_no = line_no + self.loc = loc + + def __reduce__(self): + return ( + type(self), + (self.kind, self.message), + {"line_no": self.line_no, "loc": self.loc}, + ) + + +# ─────────────────────────── address-expression terms ─────────────────────────── +# A small lazily-evaluated tree. Leaves that are only known at launch time +# (scalar kernel args) are Param nodes; pid / arange / loop variables become +# free variables with range constraints in a client's query. +# +# Bin, Cmp and IntCast also record the op they came from (``line_no``, +# ``loc``) for width obligations; those two fields take no part in +# equality, hashing or repr, so equal expressions compare equal wherever +# they were computed. + + +@dataclass(frozen=True) +class Const: + value: int + + +@dataclass(frozen=True) +class Pid: + axis: int # 0=x, 1=y, 2=z + + +@dataclass(frozen=True) +class NumPrograms: + """``tt.get_num_programs axis`` — the launch grid size along ``axis``. + Uniform across program instances, but it PARAMETERIZES the kernel's + behavior by the grid, so parsing one records the axis in ``pid_axes``: + a verdict must stay symbolic along that dim.""" + + axis: int + + +@dataclass(frozen=True) +class Arange: + ssa: str # unique per make_range site (the walk's result value index) + start: int + end: int + # Which tensor dimension this lane index varies along. -1 = 1D / not yet + # placed; set and kept current by expand_dims. Consumers key a lane + # variable by (dim, end - start), not by make_range site: every tensor + # one access combines has the access's shape, so all aranges along one + # dim with one extent index the SAME position there (tl.arange(0, 16) + + # tl.arange(16, 32) is 2i + 16, not i + j + 16), while a single + # make_range reused for several dimensions of a tile (triton does this) + # is one independent variable per dimension, or the modeled footprint + # would collapse to the diagonal. + dim: int = -1 + + +@dataclass(frozen=True) +class Param: + name: str # scalar kernel argument, substituted per launch + + +@dataclass(frozen=True) +class IterArgOffset: + """The element-offset contribution of a loop-carried pointer at the + current iteration: ``offset0 + k * delta`` (resolved from + ``graph.iter_args[arg_id]`` at eval time).""" + + arg_id: int + + +@dataclass(frozen=True) +class LoopVar: + """The scf.for induction variable; a free variable over the iterations + that run (e.g. it appears in masks like ``K - k*BLOCK_K``).""" + + loop_ssa: str + + +# Integer ops (Bin.op): the arith op each spelling reads, signed first. +# "//" and "%" truncate toward zero (divsi / remsi); the "u"-prefixed ops +# read their operands unsigned (divui / remui / minui / maxui). +_BIN_OPS = { + "arith.addi": "+", + "arith.subi": "-", + "arith.muli": "*", + "arith.divsi": "//", + "arith.remsi": "%", + "arith.minsi": "min", + "arith.maxsi": "max", + "arith.divui": "u//", + "arith.remui": "u%", + "arith.minui": "umin", + "arith.maxui": "umax", +} +UNSIGNED_BIN_OPS = frozenset({"u//", "u%", "umin", "umax"}) +UNSIGNED_PREDICATES = frozenset({"ult", "ule", "ugt", "uge"}) +_SIGNED_PREDICATES = frozenset({"slt", "sle", "sgt", "sge"}) + + +@dataclass(frozen=True) +class Bin: + op: str # + - * // % min max, or u// u% umin umax (see _BIN_OPS) + a: "Term" + b: "Term" + # Width of the integer result (the IR type). None for the element-offset + # sum a ``tt.addptr`` accumulates, which is address arithmetic, not an + # IR integer. + bits: int | None = None + line_no: int | None = field(default=None, compare=False, repr=False) + loc: SourceLoc | None = field(default=None, compare=False, repr=False) + + +@dataclass(frozen=True) +class Cmp: + pred: str # eq/ne, slt/sle/sgt/sge, ult/ule/ugt/uge (unsigned reads) + a: "Term" + b: "Term" + bits: int | None = None # operand width + line_no: int | None = field(default=None, compare=False, repr=False) + loc: SourceLoc | None = field(default=None, compare=False, repr=False) + + +@dataclass(frozen=True) +class BoolBin: + op: str # and / or + a: "Term" + b: "Term" + + +@dataclass(frozen=True) +class Select: + cond: "Term" + t: "Term" + f: "Term" + + +@dataclass(frozen=True) +class Not: + """Boolean negation — the path condition of an scf.if else-region.""" + + a: "Term" + + +@dataclass(frozen=True) +class IntCast: + """``arith.trunci`` / ``extsi`` / ``extui`` of ``x`` from ``src_bits`` to + ``dst_bits`` (D9: a cast is never a value passthrough). Its value is + ``x`` exactly when the cast's width obligation holds (trunci: ``x`` + fits the destination; extui: ``x`` is non-negative); extsi always + preserves the signed reading. ``extsi`` from i1 (true -> -1) is read as + ``0 - extui(x)`` and never appears as an IntCast.""" + + kind: str # "trunci" | "extsi" | "extui" + src_bits: int + dst_bits: int + x: "Term" + line_no: int | None = field(default=None, compare=False, repr=False) + loc: SourceLoc | None = field(default=None, compare=False, repr=False) + + +# Sentinel for a value loaded from memory (tt.load result) or computed from +# loaded data (arith.*f, tt.dot, ...). If one ever reaches an address it +# means data-dependent addressing -> unsupported; in a mask it is dropped +# (``mask_dropped``), under an scf.if it leaves the branch ``guarded``. +@dataclass(frozen=True) +class DataDep: + why: str = "value derived from loaded data" + # For a boolean ``and`` with one unmodelable operand: the modelable + # conjunct(s). The true value implies ``keep``, so a consumer may use + # ``keep`` as a sound over-approximation of such a mask. + keep: "Term | None" = None + + +@dataclass(frozen=True) +class Observed: + """The OLD value observed by the atomic at ``graph.accesses[access_index]``: + a fresh per-program-instance symbol, NOT a function of other leaves. The + reader binds an INTEGER-typed ``tt.atomic_rmw`` / ``tt.atomic_cas`` + result to this instead of ``DataDep`` so downstream masks and branch + conditions stay modelable; float-typed atomic results keep the DataDep + fallback. What an observation means (a free variable, a modeled value, + a refusal in an address) is each consumer's policy; find them with the + graph-aware :func:`mentions_observed` / :func:`observed_indices`. + + A tensor atomic observes one old value per lane, and the symbol stands + for the value at the lane the surrounding term is read at. It has no + lane placement of its own, so the reader never lets two lanes of one + tensor observation meet: an ``expand_dims`` of a term holding one + degrades to DataDep.""" + + access_index: int + + +Term = ( + Const + | Pid + | NumPrograms + | Arange + | Param + | IterArgOffset + | LoopVar + | Bin + | Cmp + | BoolBin + | Select + | Not + | IntCast + | DataDep + | Observed +) + + +def _children(t: object) -> tuple: + if isinstance(t, (Bin, Cmp, BoolBin)): + return (t.a, t.b) + if isinstance(t, Select): + return (t.cond, t.t, t.f) + if isinstance(t, (Not,)): + return (t.a,) + if isinstance(t, IntCast): + return (t.x,) + if isinstance(t, DataDep) and t.keep is not None: + return (t.keep,) + return () + + +def _nodes( + roots: Iterable[object], iter_args: "Sequence[IterArgInfo] | None" +) -> Iterator[object]: + """Every node reachable from ``roots`` in pre-order, each once (by + identity), descending ``DataDep.keep``; with ``iter_args``, an + IterArgOffset also reaches its IterArgInfo's ``offset0`` and ``delta``. + Iterative: terms can be deeper than Python's recursion limit.""" + seen: set[int] = set() + stack = [r for r in reversed(list(roots)) if r is not None] + while stack: + t = stack.pop() + if id(t) in seen: + continue + seen.add(id(t)) + yield t + if isinstance(t, IterArgOffset): + if iter_args is not None: + info = iter_args[t.arg_id] + stack += [info.delta, info.offset0] + continue + stack.extend(reversed(_children(t))) + + +def mentions_observed(term: object, graph: "AccessGraph") -> bool: + """True when ``term`` reaches an :class:`Observed` leaf, through the + graph's loop-carried pointers (an IterArgOffset's ``offset0`` and + ``delta``) and ``DataDep.keep`` included.""" + return any(isinstance(n, Observed) for n in _nodes((term,), graph.iter_args)) + + +def observed_indices(term: object, graph: "AccessGraph") -> frozenset[int]: + """Access indices of every :class:`Observed` leaf ``term`` reaches (see + :func:`mentions_observed`).""" + return frozenset( + n.access_index + for n in _nodes((term,), graph.iter_args) + if isinstance(n, Observed) + ) + + +# DataDep is also the generic unknown-value top (loop accumulators, +# unmodeled ops, ...). Only these ``why`` prefixes mean the value truly +# derives from MEMORY CONTENTS — refusals classify just those as +# indirection; the rest are modeling gaps and keep the default kind. +_MEMORY_WHYS = ( + "loaded value", + "atomic result", + "arith over loaded data", + "cmpi over loaded data", + "select over loaded data", + "bool op over loaded data", +) + + +def _from_memory(v: object) -> bool: + return isinstance(v, DataDep) and v.why.startswith(_MEMORY_WHYS) + + +@dataclass(frozen=True) +class PtrValue: + """A pointer-typed value: base argument + accumulated element offset (a + single lane's offset; arange/loop free vars cover all lanes and + iterations in the query).""" + + base_param: str + offset: Term + + +# ─────────────────────────── graph structures ─────────────────────────── + + +@dataclass(frozen=True) +class FuncArg: + name: str # the Python parameter name (NameLoc), else "arg" + is_ptr: bool + elem_bits: int # for ptr args: pointee width; 0 for scalars + # Float-typed pointee (f*/bf*): atomic results on it stay DataDep. + elem_float: bool = False + # For integer scalars: the IR width (i32 -> 32, i1 -> 1); 0 otherwise. + int_bits: int = 0 + + +@dataclass(frozen=True) +class AtomicInfo: + """Atomicity metadata for ``tt.atomic_rmw`` / ``tt.atomic_cas`` accesses.""" + + rmw_op: str | None # "fadd", "max", "exch", ... ; None for CAS + sem: str # memory semantic: "acq_rel", "relaxed", ... + scope: str # sync scope: "gpu", "cta", "sys" + + +@dataclass(frozen=True) +class AccessEvent: + kind: str # "load" | "store" | "atomic_rmw" | "atomic_cas" + base_param: str + offset: Term + mask: Term | None # None = unconditional access + elem_bits: int + loc: SourceLoc | None + line_no: int + # True when some enclosing scf.if condition could NOT be modeled (it + # derives from loaded data). The access is then checked as if + # unconditional: UNSAT stays a sound proof, but a SAT model may sit in a + # branch the launch never takes. Modeled conditions ride in ``path`` + # instead and do not set this flag. + guarded: bool = False + # Conjunction of the MODELED enclosing branch conditions, with + # else-regions negated (Not). The access executes iff path ∧ mask. + path: Term | None = None + # True when the access sits inside the scf.for body: it executes once + # per iteration — and NOT AT ALL when the launch's trip count is zero, + # which consumers must model (a zero-trip loop has no footprint). + in_loop: bool = False + # Present iff kind is atomic_*: an atomic is a read AND a write of its + # footprint (RMW). + atomic: AtomicInfo | None = None + # True when the printed mask operand derived from loaded data and was + # over-approximated as FREE (mask=None): dropping a constraint only + # widens the modeled footprint, so UNSAT stays a sound proof — but a SAT + # model may pick a lane the real mask disables. + mask_dropped: bool = False + # For atomics: the VALUE operand (tt.atomic_rmw val / tt.atomic_cas val) + # as a Term, or None when it is not modelable (loaded data). + atomic_val: "Term | None" = None + # For tt.atomic_cas only: the compare operand. + atomic_cmp: "Term | None" = None + # Float-typed pointee of the accessed pointer. + elem_float: bool = False + + @property + def is_read(self) -> bool: + return self.kind != "store" + + @property + def is_write(self) -> bool: + return self.kind != "load" + + +@dataclass(frozen=True) +class IterArgInfo: + """A loop-carried pointer: ``offset0 + k * delta`` at iteration k. A + tile of one expanded (``tt.expand_dims``) inside the loop is an entry of + its own, the same pointer with ``offset0`` and ``delta`` expanded alike, + so its lanes sit at their positions in the expanded shape.""" + + arg_id: int + base_param: str + offset0: Term + delta: Term # per-iteration element advance (loop-invariant) + loop_ssa: str = "" # the scf.for this iter_arg belongs to (LoopInfo.loop_ssa) + + +@dataclass(frozen=True) +class LoopInfo: + loop_ssa: str + induction_var: str + lower: Term + upper: Term + step: Term + # Width of the induction variable, and whether the loop compares its + # bounds unsigned (``scf.for unsigned``: the bounds' width obligations + # then require them non-negative). + bits: int | None = None + unsigned: bool = False + line_no: int | None = None + loc: SourceLoc | None = None + + +@dataclass(frozen=True) +class AccessGraph: + kernel_name: str + func_args: tuple[FuncArg, ...] + accesses: tuple[AccessEvent, ...] + loop: LoopInfo | None + # Loop-carried pointers, indexed by arg_id (``iter_args[k].arg_id == k``), + # expanded tiles of them included (see IterArgInfo). + iter_args: tuple[IterArgInfo, ...] = () + # Every pid axis with a parsed tt.get_program_id / tt.get_num_programs — + # recorded at PARSE time, before any DataDep swallowing. Consumers + # deciding grid coverage must use THIS set, not the axes that happen to + # survive into modeled address/mask terms. + pid_axes: frozenset[int] = frozenset() + + def __post_init__(self) -> None: + for name in ("func_args", "accesses", "iter_args"): + object.__setattr__(self, name, tuple(getattr(self, name))) + object.__setattr__(self, "pid_axes", frozenset(self.pid_axes)) + + def arg(self, name: str) -> FuncArg | None: + for a in self.func_args: + if a.name == name: + return a + return None + + +# ─────────────────────────── width obligations (D9) ─────────────────────────── + + +@dataclass(frozen=True) +class WidthObligation: + """``term`` must fit the width it is read at: with ``signed``, + ``-2**(bits-1) <= term < 2**(bits-1)``; otherwise ``0 <= term < 2**bits``. + ``line_no`` / ``loc`` are the op that imposes it; ``role`` is the part + of the access it comes from (see :func:`width_obligations`).""" + + term: Term + bits: int + signed: bool + line_no: int | None + loc: SourceLoc | None + role: str # "loop" | "path" | "mask" | "offset" + + +def width_obligations( + graph: AccessGraph, access: AccessEvent +) -> tuple[WidthObligation, ...]: + """The conditions under which the unbounded-integer reading of + ``access`` (its offset, mask and path, loop-carried pointers resolved, + and the loop of an access in the loop) equals the IR's fixed-width + arithmetic: + + * every integer ``Bin`` result fits its width, signed; + * the quotient of a signed remainder fits too (``INT_MIN % -1`` is + undefined in the IR); + * the operands of an unsigned op or predicate are non-negative; + * a ``trunci`` operand fits the destination width (to i1: is 0 or 1); + * an ``extui`` operand is non-negative; + * the bounds of an ``unsigned`` loop are non-negative; + * the loop's increment does not wrap: ``upper - 1 + step``, which + bounds the last iterate plus ``step``, fits the induction variable + (sufficient, not necessary). + + ``role`` says where an obligation comes from: ``"loop"`` (the loop's + bounds and increment), ``"path"``, ``"mask"`` or ``"offset"``. A term + node several parts share is listed once, under the first role in that + order, and the order is also the discipline for discharging them: an + ``offset`` obligation only matters where the access executes (path and + mask hold), a ``mask`` one only where the path holds, and ``path`` and + ``loop`` ones hold unconditionally (the increment's only when the loop + runs at least once). Both arms of a ``Select`` are listed, without its + condition, so checking an arm unconditionally may report an overflow in + the arm the launch does not take. + + Mechanism only: listed role by role in walk order, each (term node, + width, signedness) once; nothing is evaluated.""" + loop = graph.loop if access.in_loop else None + out: list[WidthObligation] = [] + seen: set[tuple[int, int, bool]] = set() + + def need(term: object, bits: int, signed: bool, site: object, role: str) -> None: + key = (id(term), bits, signed) + if key not in seen: + seen.add(key) + out.append( + WidthObligation( + term, # type: ignore[arg-type] + bits, + signed, + getattr(site, "line_no", None), + getattr(site, "loc", None), + role, + ) + ) + + groups: list[tuple[str, tuple[object, ...]]] = [] + if loop is not None: + if loop.bits is not None: + if loop.unsigned: + for bound in (loop.lower, loop.upper, loop.step): + need(bound, loop.bits, False, loop, "loop") + latch = Bin( + "+", + Bin("-", loop.upper, Const(1), loop.bits, loop.line_no, loop.loc), + loop.step, + loop.bits, + loop.line_no, + loop.loc, + ) + need(latch, loop.bits, not loop.unsigned, loop, "loop") + groups.append(("loop", (loop.lower, loop.upper, loop.step))) + groups += [ + ("path", (access.path,)), + ("mask", (access.mask,)), + ("offset", (access.offset,)), + ] + for role, roots in groups: + for n in _nodes(roots, graph.iter_args): + if isinstance(n, Bin) and n.bits is not None: + need(n, n.bits, True, n, role) + if n.op in UNSIGNED_BIN_OPS: + need(n.a, n.bits, False, n, role) + need(n.b, n.bits, False, n, role) + elif n.op == "%": + quotient = Bin("//", n.a, n.b, n.bits, n.line_no, n.loc) + need(quotient, n.bits, True, n, role) + elif isinstance(n, Cmp) and n.pred in UNSIGNED_PREDICATES and n.bits: + need(n.a, n.bits, False, n, role) + need(n.b, n.bits, False, n, role) + elif isinstance(n, IntCast): + if n.kind == "trunci": + need(n.x, n.dst_bits, n.dst_bits > 1, n, role) + elif n.kind == "extui" and n.src_bits > 1: + need(n.x, n.src_bits, False, n, role) + return tuple(out) + + +# ─────────────────────────── the reader ─────────────────────────── + + +def parse_ttir(text: str) -> AccessGraph: + """Read one TTIR module into an AccessGraph (single-path model: at most + one scf.for, structured scf.if only). + + Raises :class:`UnsupportedTTIR` for anything the graph cannot represent: + TTIR of a Triton release the walk layer or the reader has no table for + (``UNTESTED_TRITON_VERSION``), a text the walk layer cannot align + (``READER_MISALIGNMENT``) or the MLIR parser rejects (``UNPARSABLE``), + indirect addressing, block pointers, + pointers outside global memory, nested/while loops, unstructured + control flow, calls, loop-variant addresses, inline asm that is impure + or handed an address, or any op outside the address vocabulary that + could touch memory. + """ + try: + module = walk_module(text) + except UnknownTritonRelease as e: + raise UnsupportedTTIR(TTIRKind.UNTESTED_TRITON_VERSION, e.message) from e + except MisalignedModule as e: + raise UnsupportedTTIR( + TTIRKind.READER_MISALIGNMENT, "; ".join(e.problems), line_no=e.line_no + ) from e + except ModuleParseError as e: + raise UnsupportedTTIR( + TTIRKind.UNPARSABLE, e.diagnostic, line_no=e.line_no + ) from e + return _Builder(module).build() + + +# Memory ops outside the modeled vocabulary (TMA descriptors, ...: targets +# without native TMA lower descriptors to pointer math before TTIR is +# printed; sm90+ and, on 3.8, hip TDM targets such as gfx1250 keep them). +_MEMORY_PREFIXES = ("tt.descriptor_", "tt.experimental_") +_ACCESS_OPS = frozenset({"tt.load", "tt.store", "tt.atomic_rmw", "tt.atomic_cas"}) +# Dialects whose ops the reader reads as data or control; from any other +# dialect only a release's inert ops (below) are accepted. +_DIALECTS = frozenset({"tt", "arith", "math", "scf", "cf", "ub"}) + + +@dataclass(frozen=True) +class _Vocabulary: + """What the reader knows of one Triton minor release's TTIR ops, keyed + by the release whose walk-layer table read the module + (``Module.release``); a release without one refuses.""" + + # result-free ops without memory effects the reader skips (a barrier + # orders memory, it addresses none) + inert: frozenset[str] + # the dialect's block-pointer ops (refused as BLOCK_POINTER) + block_pointer_ops: frozenset[str] + + +_COMMON_INERT = frozenset( + { + "tt.return", + "scf.yield", + "tt.reduce.return", + "tt.scan.return", + "tt.print", + "tt.assert", + "llvm.intr.assume", + } +) +_VOCABULARIES: dict[str, _Vocabulary] = { + # tl.debug_barrier() is gpu.barrier + "3.6": _Vocabulary( + inert=_COMMON_INERT | {"gpu.barrier"}, + block_pointer_ops=frozenset({"tt.make_tensor_ptr", "tt.advance"}), + ), + # tl.debug_barrier() is `ttg.barrier all` (TritonSemantic.debug_barrier -> + # builder.create_barrier), which lowers exactly as 3.6's gpu.barrier did + # (cuda:89: llvm.nvvm.barrier.cta.sync.aligned.all(0), `bar.sync 0`): + # the one ttg op accepted, not the dialect. gpu.barrier, which the 3.8 + # frontend never emits, is not (its bindings still parse it). No + # block-pointer ops: tl.make_block_ptr lowers to pointer arithmetic in + # the frontend. + "3.8": _Vocabulary( + inert=_COMMON_INERT | {"ttg.barrier"}, + block_pointer_ops=frozenset(), + ), +} +_NO_VOCABULARY = _Vocabulary(inert=frozenset(), block_pointer_ops=frozenset()) +# Ops that hand the reader values it cannot see into (their operands must +# carry no address, not even as an integer). +_OPAQUE = frozenset({"tt.elementwise_inline_asm", "tt.extern_elementwise"}) +# Structured-region nesting the reader follows (deeper refuses, instead of +# exhausting Python's recursion limit). +_MAX_DEPTH = 200 +_LOOP_SSA = "%loop" # the single-path model has at most one loop +# The DataDep reason of a value carried by the loop that is not a pointer: +# in an address it makes the address loop-variant. +_LOOP_CARRIED = "loop accumulator" +_RE_ADDR_SPACE = re.compile(r", (\d+)>$") + + +def _addr_space(type_text: str) -> int: + """Address space of a (tensor of) ``!tt.ptr`` type; 1 (global memory) + when unprinted, as the printer elides it.""" + m = _RE_ADDR_SPACE.search(parse_type(type_text).elem) + return int(m.group(1)) if m else 1 + + +@dataclass +class _IfFrame: + """Walker state for one scf.if region being read.""" + + cond: "Term | None" # modeled condition; None -> accesses stay `guarded` + branch: str = "then" + + +class _ForFrame: + """Marks the scf.for body being read.""" + + +def _pointee_bits(pointee: str | None) -> int: + if pointee is None: + return 0 + if pointee.startswith("!tt.ptr<"): + return 64 # a pointer to pointers + return parse_type(pointee).int_bits or parse_type(pointee).float_bits or 0 + + +def _is_float(type_text: str | None) -> bool: + return type_text is not None and parse_type(type_text).float_bits is not None + + +class _Builder: + def __init__(self, module: Module) -> None: + self.m = module + # the ops of the release whose table read the module (build() + # refuses a release without one) + vocab = _VOCABULARIES.get(module.release) + self.release_known = vocab is not None + self.vocab = vocab if vocab is not None else _NO_VOCABULARY + # value index -> Term (int/bool), PtrValue, or DataDep + self.env: dict[int, object] = {} + self.func_args: list[FuncArg] = [] + self.accesses: list[AccessEvent] = [] + self.iter_args: list[IterArgInfo] = [] + self.loop: LoopInfo | None = None + self.loops_opened = 0 + self.pid_axes: set[int] = set() + self.frames: list[_IfFrame | _ForFrame] = [] + self.depth = 0 + # (arg_id, axis) -> the arg_id of that iter_arg's tile expanded at + # axis inside the loop (its delta is filled in when the loop closes) + self.expanded: dict[tuple[int, int], int] = {} + # access indices of the tensor atomics (one observation per lane) + self.lane_observed: set[int] = set() + + # ── helpers ── + def refuse(self, kind: TTIRKind, message: str, op: Op | None) -> NoReturn: + raise UnsupportedTTIR( + kind, + message, + line_no=op.line_no if op is not None else None, + loc=op.loc if op is not None else None, + ) + + def val(self, v: int) -> object: + got = self.env.get(v) + if got is None: + # A value the reader never bound (defined in a region it does not + # read): be conservative. + return DataDep(f"unresolved value {v}") + return got + + def bind(self, op: Op, value: object) -> None: + for r in op.results: + self.env[r] = value + + def as_term(self, v: object, ctx: str, op: Op) -> Term: + if isinstance(v, DataDep): + self.refuse(TTIRKind.OTHER, f"{ctx}: data-dependent ({v.why})", op) + if isinstance(v, PtrValue): + self.refuse(TTIRKind.OTHER, f"{ctx}: pointer used as integer", op) + return v # type: ignore[return-value] + + def int_bits(self, v: int) -> int | None: + return parse_type(parse_type(self.m.values[v].type).elem).int_bits + + def branch_state(self) -> tuple[bool, Term | None, bool]: + """(guarded, path, in_loop) for an access under the open frames: + ``guarded`` if any enclosing condition is unmodeled; ``path`` is the + conjunction of the modeled ones (else-regions negated), outermost + first; ``in_loop`` when the scf.for body encloses the access.""" + guarded = False + path: Term | None = None + in_loop = False + for f in self.frames: + if isinstance(f, _ForFrame): + in_loop = True + continue + if f.cond is None: + guarded = True + continue + c: Term = f.cond if f.branch == "then" else Not(f.cond) + path = c if path is None else BoolBin("and", path, c) + return guarded, path, in_loop + + def arg(self, name: str) -> FuncArg | None: + return next((a for a in self.func_args if a.name == name), None) + + def global_pointers(self, types: Iterable[str], op: Op) -> None: + """Element offsets are into global memory: refuse a pointer into any + other address space (shared memory, ...).""" + for t in types: + if "!tt.ptr<" in t and _addr_space(t) != 1: + self.refuse( + TTIRKind.OUT_OF_VOCABULARY, + f"pointer type {t} is not in global memory (address space 1)", + op, + ) + + # ── the module ── + def build(self) -> AccessGraph: + m = self.m + if not self.release_known: + known = ", ".join(f"{r}.x" for r in _VOCABULARIES) + self.refuse( + TTIRKind.UNTESTED_TRITON_VERSION, + f"the TTIR reader has no op vocabulary for Triton {m.release} " + f"(it reads the TTIR of Triton {known})", + None, + ) + if not m.funcs: + self.refuse(TTIRKind.OTHER, "no tt.func found (not TTIR?)", None) + for op in m.ops: + if len(op.path) == 1 and op.name != "tt.func": + self.refuse( + TTIRKind.OUT_OF_VOCABULARY, + f"{op.name} at module level is not TTIR", + op, + ) + # A call's callee runs in its own frame with its own arguments; the + # graph has no call model, so any call, and any second function, + # refuses (the callee name comes from the symbol, quoted or not). + calls = [op for op in m.ops if op.name == "tt.call"] + if calls: + first = calls[0] + self.refuse( + TTIRKind.CALL, + f"tt.call to {first.attrs.get('callee')!r}: calls are not modeled", + first, + ) + if len(m.funcs) > 1: + extra = m.funcs[1] + self.refuse( + TTIRKind.CALL, + f"{len(m.funcs)} functions in the module (tt.func " + f"{extra.sym_name!r}): calls are not modeled", + m.ops[extra.op], + ) + func = m.funcs[0] + fop = m.ops[func.op] + if not fop.regions or not fop.regions[0]: + self.refuse(TTIRKind.OTHER, f"tt.func {func.sym_name!r} has no body", fop) + for fa in func.args: + self.bind_arg(fa, fop) + body = fop.regions[0] + self.block(body[0]) + if len(body) > 1: + self.refuse( + TTIRKind.CONTROL_FLOW, + f"tt.func {func.sym_name!r} has {len(body)} blocks", + fop, + ) + return AccessGraph( + kernel_name=func.sym_name, + func_args=tuple(self.func_args), + accesses=tuple(self.accesses), + loop=self.loop, + iter_args=tuple(self.iter_args), + pid_axes=frozenset(self.pid_axes), + ) + + def bind_arg(self, fa, fop: Op) -> None: + ti = parse_type(fa.type) + name = fa.name if fa.name is not None else f"arg{fa.index}" + if self.arg(name) is not None: + self.refuse(TTIRKind.OTHER, f"two parameters named {name!r}", fop) + self.global_pointers((fa.type,), fop) + is_ptr = ti.pointee is not None and not ti.shape + int_bits = ti.int_bits if not is_ptr and not ti.shape else None + self.func_args.append( + FuncArg( + name=name, + is_ptr=is_ptr, + elem_bits=_pointee_bits(ti.pointee) if is_ptr else 0, + elem_float=_is_float(ti.pointee) if is_ptr else False, + int_bits=int_bits or 0, + ) + ) + # Pointer args seed addptr chains; integer args are Param leaves. + if is_ptr and not ti.block_ptr: + self.env[fa.value] = PtrValue(name, Const(0)) + elif int_bits is not None: + self.env[fa.value] = Param(name) + else: + self.env[fa.value] = DataDep(f"{fa.type} argument") + + def block(self, bidx: int) -> None: + self.depth += 1 + try: + for oi in self.m.blocks[bidx].ops: + self.visit(self.m.ops[oi]) + finally: + self.depth -= 1 + + def visit(self, op: Op) -> None: + if self.depth > _MAX_DEPTH: + self.refuse( + TTIRKind.OTHER, + f"structured regions nested deeper than {_MAX_DEPTH}", + op, + ) + name = op.name + if name in self.vocab.block_pointer_ops or any( + "!tt.ptr None: + # Parse-time record (see AccessGraph.pid_axes): the read counts even + # if this value never survives into a modeled term. + self.pid_axes.add(op.attrs["axis"]) + self.bind(op, Pid(op.attrs["axis"])) + + def op_num_programs(self, op: Op) -> None: + self.pid_axes.add(op.attrs["axis"]) + self.bind(op, NumPrograms(op.attrs["axis"])) + + def op_make_range(self, op: Op) -> None: + self.bind(op, Arange(f"%v{op.results[0]}", op.attrs["start"], op.attrs["end"])) + + def op_constant(self, op: Op) -> None: + v = op.attrs["value"] + if isinstance(v, bool): + # i1 constants (e.g. the dense mask of an unmasked atomic). + # Const(0/1) in a boolean position is coerced by the evaluator. + self.bind(op, Const(1 if v else 0)) + elif isinstance(v, int): + self.bind(op, Const(v)) + else: + self.bind(op, DataDep("float/array constant")) + + def op_passthrough(self, op: Op) -> None: + # tt.splat: replicate a scalar / seed a pointer tile; + # tt.broadcast: a shape change, value passthrough + self.bind(op, self.val(op.operands[0])) + + def op_expand_dims(self, op: Op) -> None: + v = self.val(op.operands[0]) + axis = op.attrs["axis"] + root = v.offset if isinstance(v, PtrValue) else v + if any( + isinstance(n, Observed) and n.access_index in self.lane_observed + for n in _nodes((root,), self.iter_args) + ): + # Observed has no lane placement: two lanes of one tensor + # observation must not meet in a term (see Observed). + self.bind(op, DataDep("atomic result expanded across lanes")) + return + self.bind(op, _expand_dims(v, axis, lambda t: self.expanded_iter_arg(t, axis))) + + def expanded_iter_arg(self, t: IterArgOffset, axis: int) -> IterArgOffset: + """The loop-carried pointer ``t`` with its lanes re-placed by an + expand_dims at ``axis``: an iter_args entry of its own (IterArgInfo), + one per (pointer, axis), whose delta the loop fills in on closing.""" + key = (t.arg_id, axis) + aid = self.expanded.get(key) + if aid is None: + src = self.iter_args[t.arg_id] + aid = len(self.iter_args) + offset0 = _expand_dims(src.offset0, axis) + self.iter_args.append( + IterArgInfo(aid, src.base_param, offset0, Const(0), _LOOP_SSA) # type: ignore[arg-type] + ) + self.expanded[key] = aid + return IterArgOffset(aid) + + def op_cast(self, op: Op) -> None: + x = self.val(op.operands[0]) + if isinstance(x, (DataDep, PtrValue)): + self.bind(op, x) + return + kind = op.name.split(".", 1)[1] + src = self.int_bits(op.operands[0]) + dst = self.int_bits(op.results[0]) + assert src is not None and dst is not None + term: Term = x # type: ignore[assignment] + if kind == "extsi" and src == 1: + # sign-extending an i1 maps true to -1 + ext = IntCast("extui", 1, dst, term, op.line_no, op.loc) + self.bind(op, Bin("-", Const(0), ext, dst, op.line_no, op.loc)) + return + self.bind(op, IntCast(kind, src, dst, term, op.line_no, op.loc)) + + def op_addptr(self, op: Op) -> None: + base, off = self.val(op.operands[0]), self.val(op.operands[1]) + if not isinstance(base, PtrValue): + self.refuse( + _address_kind(base), + f"addptr base is not a pointer{_why(base)}", + op, + ) + if isinstance(off, DataDep): + # A value in an address chain that cannot be modeled: a free + # address makes the query meaningless. Only offsets truly + # derived from MEMORY CONTENTS classify as indirection. + kind = _address_kind(off) + what = ( + "carried by the loop" + if kind is TTIRKind.LOOP_VARIANT_ADVANCE + else "data-dependent" + ) + self.refuse(kind, f"addptr offset: {what} ({off.why})", op) + off_t = self.as_term(off, "addptr offset", op) + self.bind( + op, + PtrValue( + base.base_param, # type: ignore[union-attr] + Bin("+", base.offset, off_t, None, op.line_no, op.loc), # type: ignore[union-attr] + ), + ) + + def op_bin(self, op: Op) -> None: + a, b = self.val(op.operands[0]), self.val(op.operands[1]) + bits = self.int_bits(op.results[0]) + if isinstance(a, DataDep) or isinstance(b, DataDep): + self.bind(op, _data_dep((a, b), "arith over loaded data")) + elif bits == 1: + self.bind(op, DataDep(f"{op.name} on i1")) + else: + self.bind( + op, + Bin( + _BIN_OPS[op.name], + self.as_term(a, "arith", op), + self.as_term(b, "arith", op), + bits, + op.line_no, + op.loc, + ), + ) + + def op_cmpi(self, op: Op) -> None: + a, b = self.val(op.operands[0]), self.val(op.operands[1]) + pred = op.attrs["predicate"] + bits = self.int_bits(op.operands[0]) + if isinstance(a, DataDep) or isinstance(b, DataDep): + self.bind(op, _data_dep((a, b), "cmpi over loaded data")) + elif bits == 1 and pred in _SIGNED_PREDICATES: + # a signed i1 reads true as -1; the boolean model reads it as 1 + self.bind(op, DataDep("signed comparison of i1 values")) + else: + self.bind( + op, + Cmp( + pred, + self.as_term(a, "cmpi", op), + self.as_term(b, "cmpi", op), + bits, + op.line_no, + op.loc, + ), + ) + + def op_boolbin(self, op: Op) -> None: + if self.int_bits(op.results[0]) != 1: + # Wide-int andi/ori is BITWISE arithmetic, not boolean logic; + # modeling it as And/Or would silently corrupt address math. + # Degrade to DataDep so an address use fails closed. + self.bind( + op, + DataDep(f"bitwise {op.name} on non-i1 type {op.result_types[0]}"), + ) + return + a, b = self.val(op.operands[0]), self.val(op.operands[1]) + is_and = op.name == "arith.andi" + if isinstance(a, DataDep) or isinstance(b, DataDep): + keep: Term | None = None + if is_and: + # ``modelable ∧ unmodelable`` implies ``modelable``: remember + # the modelable conjunct(s) so a mask can keep them. + parts: list[Term] = [] + for x in (a, b): + if isinstance(x, DataDep): + if x.keep is not None: + parts.append(x.keep) + elif not isinstance(x, PtrValue): + parts.append(x) # type: ignore[arg-type] + for part in parts: + keep = part if keep is None else BoolBin("and", keep, part) + why = _data_dep((a, b), "bool op over loaded data").why + self.bind(op, DataDep(why, keep=keep)) + return + self.bind( + op, + BoolBin( + "and" if is_and else "or", + self.as_term(a, "bool", op), + self.as_term(b, "bool", op), + ), + ) + + def op_select(self, op: Op) -> None: + c, t, f = (self.val(v) for v in op.operands) + self.bind(op, self.merge(c, t, f, "select")) + + def merge(self, c: object, t: object, f: object, what: str) -> object: + """``c ? t : f`` for an arith.select or an scf.if result: pointers + of one base select their offsets; anything unmodelable is DataDep.""" + if ( + isinstance(c, (DataDep, PtrValue)) + or isinstance(t, DataDep) + or (isinstance(f, DataDep)) + ): + return _data_dep((c, t, f), "select over loaded data") + if isinstance(t, PtrValue) and isinstance(f, PtrValue): + if t.base_param != f.base_param: + return DataDep(f"{what} of pointers with different bases") + return PtrValue(t.base_param, Select(c, t.offset, f.offset)) # type: ignore[arg-type] + if isinstance(t, PtrValue) or isinstance(f, PtrValue): + return DataDep(f"{what} of a pointer and an integer") + return Select(c, t, f) # type: ignore[arg-type] + + def op_bitcast(self, op: Op) -> None: + # A pointer cast that keeps the element width keeps element offsets + # (atomic_max on floats casts f32 -> i32 pointers); any other + # bitcast reinterprets data. + src = parse_type(op.operand_types[0]) + dst = parse_type(op.result_types[0]) + x = self.val(op.operands[0]) + if src.pointee is not None and dst.pointee is not None: + if _pointee_bits(src.pointee) == _pointee_bits(dst.pointee): + self.bind(op, x) + else: + self.bind(op, DataDep("pointer bitcast changes the element width")) + return + self.bind(op, _data_dep((x,), f"unmodeled op {op.name} at line {op.line_no}")) + + # ── accesses ── + def access( + self, + op: Op, + kind: str, + mask_v: int | None, + atomic: AtomicInfo | None = None, + atomic_val: Term | None = None, + atomic_cmp: Term | None = None, + ) -> None: + ptr_v = op.operands[0] + ptr = self.val(ptr_v) + if not isinstance(ptr, PtrValue): + self.refuse( + _address_kind(ptr), + f"{kind} of a non-pointer value{_why(ptr)}", + op, + ) + pointee = parse_type(parse_type(self.m.values[ptr_v].type).elem).pointee + elem_bits = _pointee_bits(pointee) + base = self.arg(ptr.base_param) # type: ignore[union-attr] + if base is None or base.elem_bits != elem_bits: + self.refuse( + TTIRKind.OTHER, + f"{kind} element width {elem_bits} differs from its base " + f"{ptr.base_param!r}", # type: ignore[union-attr] + op, + ) + mask: Term | None = None + mask_dropped = False + if mask_v is not None: + mv = self.val(mask_v) + if isinstance(mv, DataDep): + # Mask derived from loaded data: over-approximate it as free + # (any lane may be active) instead of failing the kernel. + # See AccessEvent.mask_dropped for the soundness discipline. + mask_dropped = True + elif isinstance(mv, PtrValue): + self.refuse(TTIRKind.OTHER, "pointer as mask", op) + else: + mask = mv # type: ignore[assignment] + guarded, path, in_loop = self.branch_state() + assert op.line_no is not None + self.accesses.append( + AccessEvent( + kind=kind, + base_param=ptr.base_param, # type: ignore[union-attr] + offset=ptr.offset, # type: ignore[union-attr] + mask=mask, + elem_bits=elem_bits, + loc=op.loc, + line_no=op.line_no, + guarded=guarded, + path=path, + in_loop=in_loop, + atomic=atomic, + mask_dropped=mask_dropped, + atomic_val=atomic_val, + atomic_cmp=atomic_cmp, + elem_float=_is_float(pointee), + ) + ) + + def mask_operand(self, op: Op, index: int) -> int | None: + if len(op.operands) <= index: + return None + v = op.operands[index] + if self.int_bits(v) != 1: + self.refuse(TTIRKind.OTHER, f"{op.name} operand {index} is not a mask", op) + return v + + def observed_binding(self, op: Op) -> object: + """The value of the just-recorded atomic's result: Observed for an + integer-typed result, DataDep otherwise (floats stay outside the Int + model).""" + if op.results and self.int_bits(op.results[0]) is not None: + index = len(self.accesses) - 1 + if parse_type(self.m.values[op.results[0]].type).shape: + self.lane_observed.add(index) + return Observed(index) + return DataDep("atomic result") + + def op_load(self, op: Op) -> None: + # ODS operands (ptr, mask?, other?). Both are optional, so the + # generic form spells which are present in operandSegmentSizes (the + # custom form prints a lone second operand as the mask). + segments = op.attrs.get("operandSegmentSizes") + if segments is None: + mask_v = self.mask_operand(op, 1) + elif ( + len(segments) != 3 or segments[0] != 1 or sum(segments) != len(op.operands) + ): + self.refuse(TTIRKind.OTHER, f"tt.load operand segments {segments}", op) + else: + mask_v = self.mask_operand(op, 1) if segments[1] else None + self.access(op, "load", mask_v) + self.bind(op, DataDep("loaded value")) + + def op_store(self, op: Op) -> None: + # ODS operands (ptr, value, mask?) + self.access(op, "store", self.mask_operand(op, 2)) + + def op_atomic_rmw(self, op: Op) -> None: + # ODS operands (ptr, val, mask?) + self.access( + op, + "atomic_rmw", + self.mask_operand(op, 2), + atomic=AtomicInfo(op.attrs["rmw_op"], op.attrs["sem"], op.attrs["scope"]), + atomic_val=_operand_term(self.val(op.operands[1])), + ) + self.bind(op, self.observed_binding(op)) + + def op_atomic_cas(self, op: Op) -> None: + # ODS operands (ptr, cmp, val); CAS has no mask: unconditional footprint + self.access( + op, + "atomic_cas", + None, + atomic=AtomicInfo(None, op.attrs["sem"], op.attrs["scope"]), + atomic_val=_operand_term(self.val(op.operands[2])), + atomic_cmp=_operand_term(self.val(op.operands[1])), + ) + self.bind(op, self.observed_binding(op)) + + # ── structured control flow ── + def op_for(self, op: Op) -> None: + # The single-path model has one induction variable: a second loop + # (sequential or nested) cannot be represented, and a loop under an + # scf.if runs a branch-dependent iteration count — a control-flow + # limitation, not one more induction variable. + if self.loops_opened or self.frames: + self.refuse( + TTIRKind.CONTROL_FLOW + if any(isinstance(f, _IfFrame) for f in self.frames) + else TTIRKind.NESTED_LOOP, + "multiple/nested loops", + op, + ) + self.loops_opened += 1 + bounds: dict[str, Term] = {} + for label, v in zip(("lower", "upper", "step"), op.operands[:3]): + bv = self.val(v) + if isinstance(bv, DataDep): + # The CSR shape: for k in range(loaded_start, loaded_end). + self.refuse( + TTIRKind.DATA_DEPENDENT_BOUND + if _from_memory(bv) + else TTIRKind.OTHER, + f"loop {label} bound: data-dependent ({bv.why})", + op, + ) + if any(isinstance(n, Observed) for n in _nodes((bv,), None)): + # A trip count driven by an atomic observation is a dynamic + # work-fetch loop. + self.refuse( + TTIRKind.DATA_DEPENDENT_BOUND, + f"loop {label} bound depends on an atomic observation", + op, + ) + bounds[label] = self.as_term(bv, f"loop {label}", op) + body = self.m.blocks[op.regions[0][0]] + iv, carried = body.args[0], body.args[1:] + iv_name = self.m.values[iv].name + self.env[iv] = LoopVar(_LOOP_SSA) + # Pointer iter_args become IterArgOffset; the rest are accumulators. + ptr_args: list[tuple[int, int]] = [] # (iter_arg position, arg_id) + for k, (arg_v, init_v) in enumerate(zip(carried, op.operands[3:])): + init = self.val(init_v) + if isinstance(init, PtrValue): + aid = len(self.iter_args) + self.iter_args.append( + IterArgInfo(aid, init.base_param, init.offset, Const(0), _LOOP_SSA) + ) + self.env[arg_v] = PtrValue(init.base_param, IterArgOffset(aid)) + ptr_args.append((k, aid)) + elif isinstance(init, DataDep): + # the iter_arg's value from the second iteration on is the + # yield's: the init's modelable conjuncts (keep) do not carry + self.env[arg_v] = DataDep(init.why) + else: + self.env[arg_v] = DataDep(_LOOP_CARRIED) + self.frames.append(_ForFrame()) + self.block(op.regions[0][0]) + self.frames.pop() + yield_op = self.m.ops[body.ops[-1]] + for k, aid in ptr_args: + delta = self.loop_delta(self.val(yield_op.operands[k]), aid, yield_op) + self.iter_args[aid] = replace(self.iter_args[aid], delta=delta) + # Expanded tiles advance by their source's delta, expanded alike (a + # source precedes the tiles expanded from it). + for (src, axis), aid in self.expanded.items(): + source = self.iter_args[src].delta + if any( + isinstance(n, Observed) and n.access_index in self.lane_observed + for n in _nodes((source,), None) + ): + self.refuse( + TTIRKind.OTHER, + f"loop-carried pointer {src} is expanded but advances by a " + "per-lane atomic result", + yield_op, + ) + expanded = _expand_dims(source, axis) + self.iter_args[aid] = replace(self.iter_args[aid], delta=expanded) # type: ignore[arg-type] + self.loop = LoopInfo( + loop_ssa=_LOOP_SSA, + induction_var=f"%{iv_name}" if iv_name else f"%v{iv}", + lower=bounds["lower"], + upper=bounds["upper"], + step=bounds["step"], + bits=self.int_bits(iv), + unsigned=bool(op.attrs["unsignedCmp"]), + line_no=op.line_no, + loc=op.loc, + ) + self.bind(op, DataDep("loop result")) + + def loop_delta(self, y: object, aid: int, op: Op) -> Term: + """The loop-invariant advance of loop-carried pointer ``aid`` from + its yielded value, which must be ``IterArgOffset(aid) + delta`` on the + same base: the model reads iteration k's offset as + ``offset0 + k * delta``.""" + info = self.iter_args[aid] + if not isinstance(y, PtrValue) or y.base_param != info.base_param: + self.refuse( + TTIRKind.LOOP_VARIANT_ADVANCE, + f"loop-carried pointer {aid} (base {info.base_param!r}) is not " + "advanced from itself", + op, + ) + delta = _loop_delta(y.offset, aid) # type: ignore[union-attr] + if delta is None: + self.refuse( + TTIRKind.LOOP_VARIANT_ADVANCE, + f"loop-carried pointer {aid} is not advanced by addptr from " + "its own previous value", + op, + ) + for n in _nodes((delta,), None): + if ( + isinstance(n, (LoopVar, IterArgOffset)) + or isinstance(n, Observed) + and self.accesses[n.access_index].in_loop + ): + self.refuse( + TTIRKind.LOOP_VARIANT_ADVANCE, + f"loop-carried pointer {aid} advances by a loop-variant amount", + op, + ) + return delta # type: ignore[return-value] + + def op_if(self, op: Op) -> None: + cv = self.val(op.operands[0]) + # A pointer can't be a condition; loaded data (DataDep) can't be + # modeled -> the region stays pessimistically ``guarded``. + cond = None if isinstance(cv, (DataDep, PtrValue)) else cv + frame = _IfFrame(cond) # type: ignore[arg-type] + yields: list[list[object]] = [] + for ri, region in enumerate(op.regions): + if not region: + continue + if len(region) != 1: + self.op_control_flow(op) + frame.branch = "then" if ri == 0 else "else" + self.frames.append(frame) + self.block(region[0]) + self.frames.pop() + y = self.m.ops[self.m.blocks[region[0]].ops[-1]] + yields.append([self.val(v) for v in y.operands]) + if op.results and len(yields) != 2: + self.refuse(TTIRKind.OTHER, "scf.if with results but no else region", op) + for k, r in enumerate(op.results): + if cond is None: + self.env[r] = _data_dep((cv,), "select over loaded data") + else: + self.env[r] = self.merge(cond, yields[0][k], yields[1][k], "scf.if") + + # ── everything else ── + def op_control_flow(self, op: Op) -> None: + self.refuse(TTIRKind.CONTROL_FLOW, f"control flow {op.name} is unsupported", op) + + def op_call(self, op: Op) -> None: + self.refuse(TTIRKind.CALL, "calls are not modeled", op) + + def opaque_hazard(self, op: Op) -> str | None: + """Why an opaque op (inline asm, an extern function) may touch + memory the graph cannot see, or None: it is impure, or an operand + hands it an address, as a pointer or as an integer made from one.""" + if op.attrs.get("pure") is not True: + return "side effects" + if any(parse_type(t).pointee is not None for t in op.operand_types): + return "a pointer operand" + if self.from_pointer(op.operands): + return "an operand computed from a pointer (tt.ptr_to_int)" + return None + + def from_pointer(self, values: Iterable[int]) -> bool: + """True when a ``tt.ptr_to_int`` result flows into one of ``values`` + (through any op, region results and loop-carried values included; + loaded data carries no address).""" + m = self.m + stack = list(values) + seen: set[int] = set() + while stack: + v = stack.pop() + if v in seen: + continue + seen.add(v) + value = m.values[v] + if value.op is not None: + src = m.ops[value.op] + if src.name == "tt.ptr_to_int": + return True + if src.name in _ACCESS_OPS: + continue + else: + assert value.block is not None + src = m.ops[m.blocks[value.block].op] + if src.name == "tt.func": + continue + stack += src.operands + for region in src.regions: + for b in region: + ops = m.blocks[b].ops + if ops: + stack += m.ops[ops[-1]].operands + return False + + def op_inline_asm(self, op: Op) -> None: + # Inline asm is opaque: an impure one may touch memory the graph + # cannot see, and an address operand hands it an address. + hazard = self.opaque_hazard(op) + if hazard is not None: + self.refuse(TTIRKind.INLINE_ASM, f"inline asm with {hazard}", op) + self.bind(op, _data_dep(map(self.val, op.operands), "unmodeled inline asm")) + + def op_extern(self, op: Op) -> None: + hazard = self.opaque_hazard(op) + if hazard is not None: + self.refuse( + TTIRKind.OUT_OF_VOCABULARY, + f"{op.name} {op.attrs.get('symbol')!r} with {hazard}", + op, + ) + self.op_other(op) + + def op_other(self, op: Op) -> None: + name = op.name + if name.startswith(_MEMORY_PREFIXES): + self.refuse(TTIRKind.OUT_OF_VOCABULARY, f"unsupported memory op {name}", op) + if name.split(".", 1)[0] not in _DIALECTS: + self.refuse(TTIRKind.OUT_OF_VOCABULARY, f"op {name} is not TTIR", op) + if name.startswith(("scf.", "cf.")): + self.op_control_flow(op) + if op.regions: + self.pure_regions(op) + elif not op.results: + # a result-free op the reader does not know may write memory + self.refuse(TTIRKind.OUT_OF_VOCABULARY, f"unmodeled op {name}", op) + # Plain data (floats, dots, reductions, ...): memory-derived exactly + # when an operand is. + self.bind( + op, + _data_dep( + map(self.val, op.operands), f"unmodeled op {name} at line {op.line_no}" + ), + ) + + def pure_regions(self, op: Op) -> None: + """The regions of an op the reader does not follow (tt.reduce / + tt.scan combine bodies) compute values only: refuse any memory, + call or opaque op inside.""" + stack = [b for region in op.regions for b in region] + while stack: + for oi in self.m.blocks[stack.pop()].ops: + inner = self.m.ops[oi] + name = inner.name + if ( + name in _ACCESS_OPS + or name == "tt.call" + or name.startswith(_MEMORY_PREFIXES) + or ( + name.split(".", 1)[0] not in _DIALECTS + and name not in self.vocab.inert + ) + or (name in _OPAQUE and self.opaque_hazard(inner) is not None) + or ( + not inner.results + and not inner.regions + and name not in self.vocab.inert + and name != "scf.condition" + ) + ): + self.refuse( + TTIRKind.OUT_OF_VOCABULARY, + f"{name} inside the region of {op.name}", + inner, + ) + stack.extend(b for region in inner.regions for b in region) + + +def _data_dep(operands: Iterable[object], why: str) -> DataDep: + """The DataDep of an op the reader cannot model over ``operands``: a + memory-derived reason when any operand derives from memory (``why`` + itself when it is one), else the first unmodelable operand's own reason + (a modeling gap stays a modeling gap), else ``why``.""" + first: DataDep | None = None + for x in operands: + if _from_memory(x): + return DataDep( + why + if why.startswith(_MEMORY_WHYS) + else f"arith over loaded data ({why})" + ) + if first is None and isinstance(x, DataDep): + first = x + return DataDep(first.why) if first is not None else DataDep(why) + + +def _why(v: object) -> str: + return f" ({v.why})" if isinstance(v, DataDep) else "" + + +def _address_kind(v: object) -> TTIRKind: + """The refusal kind of an unmodelable value in an address: loaded data + is indirection, a loop-carried integer a loop-variant address.""" + if _from_memory(v): + return TTIRKind.INDIRECT_ADDRESS + if isinstance(v, DataDep) and v.why == _LOOP_CARRIED: + return TTIRKind.LOOP_VARIANT_ADVANCE + return TTIRKind.OTHER + + +def _operand_term(v: object) -> Term | None: + """An atomic cmp/val operand as a Term, or None when unmodelable.""" + return None if isinstance(v, (DataDep, PtrValue)) else v # type: ignore[return-value] + + +def _with_children(t: object, kids: tuple) -> object: + if isinstance(t, (Bin, Cmp, BoolBin)): + return replace(t, a=kids[0], b=kids[1]) + if isinstance(t, Select): + return replace(t, cond=kids[0], t=kids[1], f=kids[2]) + if isinstance(t, Not): + return replace(t, a=kids[0]) + if isinstance(t, IntCast): + return replace(t, x=kids[0]) + if isinstance(t, DataDep): + return replace(t, keep=kids[0]) + raise TypeError(type(t).__name__) + + +def _expand_dims( + v: object, + axis: int, + iter_arg: Callable[[IterArgOffset], object] | None = None, +) -> object: + """``tt.expand_dims`` inserts a size-1 dimension at ``axis``: every + Arange lane index moves to its position in the new shape (a 1D range + sits at 0; a dimension at or after ``axis`` shifts up by one). Pointer + tiles follow like integer tiles, so an address and a mask expanded the + same way keep sharing their lane variables; ``iter_arg`` maps a + loop-carried pointer's IterArgOffset to its expanded tile's. Iterative, + sharing-preserving (terms can be deeper than the recursion limit).""" + if isinstance(v, PtrValue): + return PtrValue(v.base_param, _expand_dims(v.offset, axis, iter_arg)) # type: ignore[arg-type] + memo: dict[int, object] = {} + stack: list[tuple[object, bool]] = [(v, False)] + while stack: + t, ready = stack.pop() + if id(t) in memo: + continue + kids = _children(t) + if isinstance(t, Arange): + pos = 0 if t.dim < 0 else t.dim + memo[id(t)] = replace(t, dim=pos + 1 if pos >= axis else pos) + elif isinstance(t, IterArgOffset) and iter_arg is not None: + memo[id(t)] = iter_arg(t) + elif not kids: + memo[id(t)] = t + elif not ready: + stack.append((t, True)) + stack.extend((k, False) for k in kids if id(k) not in memo) + else: + new = tuple(memo[id(k)] for k in kids) + memo[id(t)] = ( + t if all(n is k for n, k in zip(new, kids)) else _with_children(t, new) + ) + return memo[id(v)] + + +def _loop_delta(offset: Term, arg_id: int) -> Term | None: + """From a yielded pointer offset ``IterArgOffset(arg_id) + d1 + ... + dn`` + (a chain of addptr sums, any association), pull out the delta + ``d1 + ... + dn`` (``Const(0)`` for none); None for any other shape.""" + found = 0 + rest: list[Term] = [] + stack: list[Term] = [offset] + while stack: + t = stack.pop() + if isinstance(t, Bin) and t.op == "+" and t.bits is None: + stack += [t.b, t.a] + elif isinstance(t, IterArgOffset) and t.arg_id == arg_id: + found += 1 + else: + rest.append(t) + if found != 1: + return None + if not rest: + return Const(0) + delta = rest[0] + for t in rest[1:]: + delta = Bin("+", delta, t) + return delta + + +_HANDLERS = { + "tt.get_program_id": _Builder.op_program_id, + "tt.get_num_programs": _Builder.op_num_programs, + "tt.make_range": _Builder.op_make_range, + "arith.constant": _Builder.op_constant, + "tt.splat": _Builder.op_passthrough, + "tt.broadcast": _Builder.op_passthrough, + "tt.expand_dims": _Builder.op_expand_dims, + "arith.extsi": _Builder.op_cast, + "arith.extui": _Builder.op_cast, + "arith.trunci": _Builder.op_cast, + "tt.addptr": _Builder.op_addptr, + **{name: _Builder.op_bin for name in _BIN_OPS}, + "arith.cmpi": _Builder.op_cmpi, + "arith.andi": _Builder.op_boolbin, + "arith.ori": _Builder.op_boolbin, + "arith.select": _Builder.op_select, + "tt.bitcast": _Builder.op_bitcast, + "tt.load": _Builder.op_load, + "tt.store": _Builder.op_store, + "tt.atomic_rmw": _Builder.op_atomic_rmw, + "tt.atomic_cas": _Builder.op_atomic_cas, + "scf.for": _Builder.op_for, + "scf.if": _Builder.op_if, + "tt.call": _Builder.op_call, + "tt.elementwise_inline_asm": _Builder.op_inline_asm, + "tt.extern_elementwise": _Builder.op_extern, +} diff --git a/tilelens/ir/verdict.py b/tilelens/ir/verdict.py new file mode 100644 index 000000000..552101116 --- /dev/null +++ b/tilelens/ir/verdict.py @@ -0,0 +1,132 @@ +"""The structured record an IR client adds to ``Launch.records`` (D5). + +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 (D20), 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 (e.g. the IR client's version gate, + "untested-triton-version", the kind the TTIR reader also gives a release + it has no table for).""" + + 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 (D3).""" + + # 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) diff --git a/tilelens/wrapper.py b/tilelens/wrapper.py index d2759da85..910bc59f8 100644 --- a/tilelens/wrapper.py +++ b/tilelens/wrapper.py @@ -14,6 +14,17 @@ PROFILER_COMMAND = "tile-profiler" RACE_DETECTOR_COMMAND = "tile-race" +# tile-sanitizer's flag for the compiled sanitizer, given before the script +# name (D14). +COMPILE_FLAG = "--compile" +# Printed (to stderr) once when the flag is given: the kernels do not run (D2). +COMPILE_NOTE = ( + f"[{COMPILE_FLAG}] kernel launches are compiled and checked, not run: " + "their outputs are never written, so the script sees them unchanged, and " + "a launch whose arguments it computes from them is checked with those " + "values" +) + # Former Triton-Viz command names, still installed as aliases. LEGACY_COMMANDS = { SANITIZER_COMMAND: "triton-sanitizer", @@ -36,6 +47,16 @@ def sanitizer_wrapper(kernel, *, frontend: str = "triton"): return tracer(kernel) +def compiled_sanitizer_wrapper(kernel, *, frontend: str = "triton"): + # Checks each launch against the kernel's compiled TTIR instead of + # interpreting it; the kernel does not run (Sanitizer(compile=True)). + tracer = tilelens.trace( + client=Sanitizer(compile=True, abort_on_error=True), + frontend=frontend, + ) + return tracer(kernel) + + def profiler_wrapper(kernel, *, frontend: str = "triton"): tracer = tilelens.trace(client=Profiler(), frontend=frontend) return tracer(kernel) @@ -78,9 +99,11 @@ def _decorator(f): return _patched_autotune -def _apply_wrapper(wrapper_func, command_name, usage_msg): +def _apply_wrapper(wrapper_func, command_name, usage_msg, compile_wrapper=None): """ Generic function to apply a wrapper to triton.jit and run the user script. + A command with a ``compile_wrapper`` uses it instead when its first + argument is COMPILE_FLAG. """ legacy_command = LEGACY_COMMANDS[command_name] if os.path.basename(sys.argv[0]) not in (command_name, legacy_command): @@ -91,6 +114,11 @@ def _apply_wrapper(wrapper_func, command_name, usage_msg): cfg.cli_active = True + if compile_wrapper is not None and sys.argv[1:2] == [COMPILE_FLAG]: + del sys.argv[1] + wrapper_func = compile_wrapper + print(COMPILE_NOTE, file=sys.stderr) + # Patch Triton kernels with the Triton frontend. _patched_jit = create_patched_jit( wrapper_func, @@ -149,12 +177,21 @@ def _apply_wrapper(wrapper_func, command_name, usage_msg): def apply_sanitizer(): """ - Apply the sanitizer wrapper to triton.jit and run the user script. + Apply the sanitizer wrapper to triton.jit and run the user script; with + ``--compile`` before the script, the compiled sanitizer's. """ _apply_wrapper( sanitizer_wrapper, SANITIZER_COMMAND, - f"Usage: {SANITIZER_COMMAND} [args...]", + f"Usage: {SANITIZER_COMMAND} [{COMPILE_FLAG}] [args...]\n" + f" {COMPILE_FLAG} check each launch against the compiled kernel " + "(Sanitizer(compile=True)) instead of interpreting it; kernels are not " + "run, so their outputs are never written. An 'ok' holds for the " + "arguments each launch was called with. Kernels are compiled on the " + "host (no GPU needed) for TILELENS_IR_TARGET, by default cuda:89 " + "(e.g. cuda:90, hip:gfx942); a config that fails to compile for it is " + "reported as not checked, and the script goes on.", + compile_wrapper=compiled_sanitizer_wrapper, ) diff --git a/tools/ir_bulk_conformance.py b/tools/ir_bulk_conformance.py new file mode 100644 index 000000000..02d71e50e --- /dev/null +++ b/tools/ir_bulk_conformance.py @@ -0,0 +1,257 @@ +"""Bulk-run the TTIR walker over a directory of ``.ttir`` files and summarise. + +A D10b conformance aid, not a test: every file the MLIR parser accepts must +walk with 0 misalignments. Point it at a Triton cache (the default, +``~/.triton/cache``) to cover real compiled kernels, not only the curated +goldens. Files are only read. + + python tools/ir_bulk_conformance.py [ROOT ...] [--jobs N] [--any-version] + [--jsonl OUT] [--limit N] [--show N] + +The walk uses the installed Triton's printer table +(``tilelens.ir._mlir_walk.PRINTERS``); a release without one is reported and +nothing is walked. By default only cache entries whose metadata names the +installed Triton version are walked (the table is that release's printer); +files with no metadata at all (a plain directory of .ttir files) are always +walked. Texts are deduplicated by sha256. Walks run in worker subprocesses, +so a crash inside the MLIR bindings loses one chunk, which the summary +reports. +""" + +from __future__ import annotations + +import argparse +import collections +import concurrent.futures +import glob +import hashlib +import json +import os +import re +import subprocess +import sys +import time + +REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + + +def _cache_version(ttir_path: str) -> str | None | bool: + """The Triton version recorded next to a cache entry: a version string, + None when the metadata has no version, False when there is no metadata.""" + found = False + for j in glob.glob(os.path.join(os.path.dirname(ttir_path), "*.json")): + if os.path.basename(j).startswith("__grp__"): + continue + found = True + try: + with open(j) as f: + meta = json.load(f) + except (OSError, ValueError): + continue + if isinstance(meta, dict) and meta.get("triton_version"): + return str(meta["triton_version"]) + return None if found else False + + +def collect( + roots: list[str], version: str | None +) -> tuple[list[str], collections.Counter]: + counts: collections.Counter = collections.Counter() + seen: set[bytes] = set() + paths: list[str] = [] + for root in roots: + for dirpath, _dirs, files in os.walk(root): + for fn in sorted(files): + if not fn.endswith(".ttir"): + continue + p = os.path.join(dirpath, fn) + counts["files"] += 1 + if version is not None: + v = _cache_version(p) + if v is not False and v != version: + counts["skipped: other or unrecorded Triton version"] += 1 + continue + try: + with open(p, "rb") as f: + h = hashlib.sha256(f.read()).digest() + except OSError: + counts["unreadable"] += 1 + continue + if h in seen: + counts["duplicate texts"] += 1 + continue + seen.add(h) + paths.append(p) + return paths, counts + + +def worker() -> None: + """Walk each path read from stdin; print one JSON line per file.""" + import warnings + + warnings.filterwarnings("ignore") + sys.path.insert(0, REPO) + from tilelens.ir import _mlir_walk as W + + for p in sys.stdin.read().splitlines(): + if not p: + continue + row: dict = {"path": p} + t0 = time.perf_counter() + try: + with open(p, encoding="utf-8") as f: + text = f.read() + m = W._walk(text) # uncached: every file is distinct anyway + row.update(status="ok", stats=dict(m.stats)) + except W.MisalignedModule as e: + row.update( + status="misaligned", problems=list(e.problems[:20]), line_no=e.line_no + ) + except W.ModuleParseError as e: + row.update( + status="parse-error", problems=[e.diagnostic[:300]], line_no=e.line_no + ) + except W.UnknownTritonRelease as e: + row.update(status="unknown-release", problems=[e.message]) + except Exception as e: # noqa: BLE001 (an escape is a walker bug: report it) + row.update( + status="exception", problems=[f"{type(e).__name__}: {str(e)[:300]}"] + ) + row["ms"] = round((time.perf_counter() - t0) * 1e3, 3) + print(json.dumps(row), flush=True) + + +def run_chunk(chunk: list[str], timeout: float) -> list[dict]: + cmd = [sys.executable, os.path.abspath(__file__), "--worker"] + try: + r = subprocess.run( + cmd, input="\n".join(chunk), capture_output=True, text=True, timeout=timeout + ) + out, rc, err = r.stdout, str(r.returncode), r.stderr + except subprocess.TimeoutExpired as e: + out = e.stdout.decode() if isinstance(e.stdout, bytes) else (e.stdout or "") + rc, err = "timeout", "" + rows = [json.loads(line) for line in out.splitlines() if line.startswith("{")] + done = {r["path"] for r in rows} + for p in chunk: + if p not in done: + rows.append( + { + "path": p, + "status": "worker-lost", + "problems": [f"worker rc={rc}: {err.strip()[-300:]}"], + } + ) + return rows + + +def _category(problem: str) -> str: + q = re.sub(r"^line \d+: ", "", problem) + q = re.sub(r"\(\[.*|\(\['.*", "(...)", q) + q = re.sub(r"'[^']*'|\"[^\"]*\"", "'…'", q) + q = re.sub(r"\d+", "N", q) + return q[:160] + + +def summarize(rows: list[dict], counts: collections.Counter, show: int) -> None: + by = collections.Counter(r["status"] for r in rows) + print("input:", dict(counts)) + print("walked:", len(rows), dict(by)) + parsed = by["ok"] + by["misaligned"] + if parsed: + rate = by["misaligned"] / parsed + print( + f"misalignment rate: {by['misaligned']}/{parsed} parsed texts = {rate:.4%}" + ) + for status in ( + "misaligned", + "parse-error", + "unknown-release", + "exception", + "worker-lost", + ): + bad = [r for r in rows if r["status"] == status] + if not bad: + continue + cats: collections.Counter[str] = collections.Counter() + example: dict[str, str] = {} + for r in bad: + c = _category(r["problems"][0]) if r.get("problems") else "?" + cats[c] += 1 + example.setdefault(c, r["path"]) + print(f"\n-- {status}: {len(bad)} files, first problem by category") + for c, n in cats.most_common(show): + print(f"{n:7d} {c}\n e.g. {example[c]}") + tot: collections.Counter = collections.Counter() + for r in rows: + if r["status"] == "ok": + tot.update(r["stats"]) + if tot: + print("\nchecked over aligned texts:", dict(tot)) + ms = sorted(r["ms"] for r in rows if "ms" in r) + if ms: + q = lambda f: ms[min(len(ms) - 1, int(f * len(ms)))] # noqa: E731 + print( + f"walk time ms: median {q(0.5):.2f} p99 {q(0.99):.2f} max {ms[-1]:.2f} total {sum(ms) / 1e3:.1f}s" + ) + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__.split("\n")[0]) + ap.add_argument("roots", nargs="*", default=[os.path.expanduser("~/.triton/cache")]) + ap.add_argument("--jobs", type=int, default=min(16, os.cpu_count() or 1)) + ap.add_argument( + "--chunk", type=int, default=200, help="files per worker subprocess" + ) + ap.add_argument( + "--timeout", type=float, default=900.0, help="seconds per worker subprocess" + ) + ap.add_argument( + "--any-version", + action="store_true", + help="walk cache entries of every Triton version", + ) + ap.add_argument("--limit", type=int, default=0, help="walk at most N texts") + ap.add_argument("--jsonl", help="write one JSON row per walked file here") + ap.add_argument("--show", type=int, default=15, help="categories listed per status") + ap.add_argument("--worker", action="store_true", help=argparse.SUPPRESS) + args = ap.parse_args() + if args.worker: + worker() + return 0 + sys.path.insert(0, REPO) + from tilelens.ir import _mlir_walk as W + + try: + table = W.printer() + except W.UnknownTritonRelease as e: + print(e.message, file=sys.stderr) + return 2 + version = None + if not args.any_version: + import triton + + version = triton.__version__ + paths, counts = collect(args.roots, version) + if args.limit: + paths = paths[: args.limit] + print( + f"{len(paths)} distinct TTIR texts (Triton {version or 'any version'}; " + f"the Triton {table.release} printer table)", + file=sys.stderr, + ) + chunks = [paths[i : i + args.chunk] for i in range(0, len(paths), args.chunk)] + rows: list[dict] = [] + with concurrent.futures.ThreadPoolExecutor(max_workers=max(1, args.jobs)) as ex: + for got in ex.map(lambda c: run_chunk(c, args.timeout), chunks): + rows += got + if args.jsonl: + with open(args.jsonl, "w") as f: + for r in rows: + f.write(json.dumps(r) + "\n") + summarize(rows, counts, args.show) + return 0 if all(r["status"] in ("ok", "parse-error") for r in rows) else 1 + + +if __name__ == "__main__": + sys.exit(main())