From b550665e7edbeaee700b3c0c7821de04b377511a Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Tue, 29 Sep 2026 14:44:46 -0400 Subject: [PATCH 1/5] [FEAT] Add a host-compiled IR layer for clients Lets a client analyze a kernel's compiled IR instead of interpreting it, without a GPU. - core: clients declare NEEDS_INTERPRETER, IR_STAGES and LAUNCH; IR clients get launch events from a capture around jit_fn.run (every autotune/heuristics config, deduplicated by binding), take no part in op/loop patching or the pre_run vote, and conflicting launch preferences are refused at registration. Each launch gets its own Launch, client finalize is isolated, and runner chains are rebuilt instead of mutating the user's Autotuner/Heuristics. - host compile: TTIR through the JIT's own binder and specialization, compiled on the host for a configurable target (default cuda:89, TILELENS_IR_TARGET); no driver or device access on the IR path. - tilelens.ir: a TTIR reader that walks the MLIR bindings and reads the attributes they cannot expose from the aligned text; the AccessGraph term model with bit widths and width obligations; capture, launch binding, IRClient base and IRVerdict records (saved by tilelens.save). - Tested on Triton 3.6 and 3.8; other releases are refused unless TILELENS_IR_ALLOW_UNTESTED_TRITON is set, and IR tests skip there. - Tests: reader conformance suite (static footprint vs Triton's interpreter), golden TTIR per release, lifecycle and host-compile tests, and tools/ir_bulk_conformance.py. --- .codespellrc | 5 + tests/conformance/__init__.py | 0 tests/conformance/_corpus.py | 1753 ++++++ tests/conformance/_interp_footprint.py | 334 ++ tests/conformance/_static_footprint.py | 603 ++ tests/conformance/_ttir_capture.py | 215 + tests/conformance/conftest.py | 79 + tests/conformance/test_reader_conformance.py | 486 ++ tests/conftest.py | 121 + tests/end_to_end/test_ir_client.py | 301 + tests/end_to_end/test_ir_smoke.py | 417 ++ tests/golden/ir/expected.json | 5038 +++++++++++++++++ tests/golden/ir/expected_3.8.json | 684 +++ tests/golden/ir/generate_reader_ttir.py | 109 + tests/golden/ir/generate_ttir.py | 406 ++ tests/golden/ir/reader_kernels.py | 248 + .../ir/reader_ttir/expand_iterarg_3d.ttir | 94 + .../ir/reader_ttir/expand_iterarg_mask.ttir | 62 + .../ir/reader_ttir/int_iterarg_offset.ttir | 31 + tests/golden/ir/reader_ttir/iv_wrap.ttir | 20 + .../ir/reader_ttir/loop_observed_advance.ttir | 29 + .../ir/reader_ttir/loop_two_step_advance.ttir | 29 + .../golden/ir/reader_ttir/observed_lanes.ttir | 45 + .../ir/reader_ttir/p1_variant_delta.ttir | 30 + tests/golden/ir/reader_ttir/p2_swap.ttir | 27 + .../ir/reader_ttir/p3_call_formals.ttir | 30 + .../ir/reader_ttir/p3_call_guarded.ttir | 31 + .../golden/ir/reader_ttir/p3_call_offset.ttir | 30 + .../ir/reader_ttir/p4_observed_delta.ttir | 38 + .../ir/reader_ttir/p4_observed_direct.ttir | 30 + .../ir/reader_ttir/p4_observed_loop.ttir | 41 + .../ir/reader_ttir/pure_asm_int_addr.ttir | 22 + tests/golden/ir/reader_ttir/rv_i32_wrap.ttir | 23 + .../ir/reader_ttir/rv_inline_asm_store.ttir | 20 + .../ir/reader_ttir/rv_trunci_alias.ttir | 27 + .../ir/reader_ttir/tile3d_shared_arange.ttir | 46 + .../golden/ir/reader_ttir/unsigned_index.ttir | 26 + .../golden/ir/reader_ttir/where_pointer.ttir | 31 + .../ir/reader_ttir_3.8/expand_iterarg_3d.ttir | 78 + .../reader_ttir_3.8/expand_iterarg_mask.ttir | 54 + .../reader_ttir_3.8/int_iterarg_offset.ttir | 29 + tests/golden/ir/reader_ttir_3.8/iv_wrap.ttir | 19 + .../loop_observed_advance.ttir | 27 + .../loop_two_step_advance.ttir | 27 + .../ir/reader_ttir_3.8/observed_lanes.ttir | 43 + .../ir/reader_ttir_3.8/p1_variant_delta.ttir | 28 + tests/golden/ir/reader_ttir_3.8/p2_swap.ttir | 25 + .../ir/reader_ttir_3.8/p3_call_formals.ttir | 28 + .../ir/reader_ttir_3.8/p3_call_guarded.ttir | 29 + .../ir/reader_ttir_3.8/p3_call_offset.ttir | 28 + .../ir/reader_ttir_3.8/p4_observed_delta.ttir | 36 + .../reader_ttir_3.8/p4_observed_direct.ttir | 29 + .../ir/reader_ttir_3.8/p4_observed_loop.ttir | 39 + .../ir/reader_ttir_3.8/pure_asm_int_addr.ttir | 21 + .../ir/reader_ttir_3.8/rv_i32_wrap.ttir | 22 + .../reader_ttir_3.8/rv_inline_asm_store.ttir | 19 + .../ir/reader_ttir_3.8/rv_trunci_alias.ttir | 24 + .../reader_ttir_3.8/tile3d_shared_arange.ttir | 37 + .../ir/reader_ttir_3.8/unsigned_index.ttir | 25 + .../ir/reader_ttir_3.8/where_pointer.ttir | 30 + tests/golden/ir/ttir/adv_cf_blockargs.ttir | 86 + tests/golden/ir/ttir/adv_consts.ttir | 98 + tests/golden/ir/ttir/adv_descs.ttir | 36 + tests/golden/ir/ttir/adv_hinted.ttir | 41 + tests/golden/ir/ttir/adv_multi_func.ttir | 83 + tests/golden/ir/ttir/adv_multi_result.ttir | 105 + tests/golden/ir/ttir/adv_names.ttir | 37 + tests/golden/ir/ttir/adv_nest3.ttir | 196 + tests/golden/ir/ttir/adv_reduce3.ttir | 218 + tests/golden/ir/ttir/adv_views.ttir | 76 + tests/golden/ir/ttir/adv_while_nested.ttir | 63 + tests/golden/ir/ttir/adv_zero_result.ttir | 59 + tests/golden/ir/ttir/crafted_attr_dicts.ttir | 18 + tests/golden/ir/ttir/crafted_deep_nest.ttir | 43 + .../golden/ir/ttir/crafted_empty_bodies.ttir | 24 + tests/golden/ir/ttir/crafted_empty_else.ttir | 10 + tests/golden/ir/ttir/crafted_empty_for.ttir | 10 + tests/golden/ir/ttir/crafted_fwd_ref_cf.ttir | 12 + .../golden/ir/ttir/crafted_generic_form.ttir | 11 + tests/golden/ir/ttir/crafted_locs.ttir | 16 + tests/golden/ir/ttir/crafted_odd_names.ttir | 12 + tests/golden/ir/ttir/crafted_same_dest.ttir | 15 + .../ir/ttir/crafted_symbols_strings.ttir | 16 + .../ir/ttir/crafted_unicode_strings.ttir | 10 + tests/golden/ir/ttir/golden_add_sm80.ttir | 51 + tests/golden/ir/ttir/golden_add_sm90.ttir | 51 + .../ir/ttir/golden_atomic_fmax_sm80.ttir | 53 + .../ir/ttir/golden_atomic_fmax_sm90.ttir | 53 + tests/golden/ir/ttir/golden_atomic_sm80.ttir | 48 + tests/golden/ir/ttir/golden_atomic_sm90.ttir | 48 + tests/golden/ir/ttir/golden_cas_sm80.ttir | 16 + tests/golden/ir/ttir/golden_cas_sm90.ttir | 16 + .../ttir/golden_early_return_loaded_sm80.ttir | 56 + .../ir/ttir/golden_early_return_pid_sm80.ttir | 48 + tests/golden/ir/ttir/golden_gather_sm80.ttir | 51 + tests/golden/ir/ttir/golden_gather_sm90.ttir | 51 + .../ir/ttir/golden_grid_stride_sm80.ttir | 41 + .../ir/ttir/golden_guard_then_loop_sm80.ttir | 55 + .../ir/ttir/golden_if_else_load_sm80.ttir | 60 + .../ir/ttir/golden_if_else_load_sm90.ttir | 60 + .../ir/ttir/golden_if_else_offset_sm80.ttir | 40 + .../ir/ttir/golden_if_else_offset_sm90.ttir | 40 + .../ir/ttir/golden_loop_under_if_sm80.ttir | 52 + .../ir/ttir/golden_matmul_bp_s3_sm80.ttir | 168 + .../ir/ttir/golden_matmul_bp_s3_sm90.ttir | 168 + .../golden/ir/ttir/golden_matmul_s1_sm80.ttir | 172 + .../golden/ir/ttir/golden_matmul_s1_sm90.ttir | 172 + .../golden/ir/ttir/golden_matmul_s3_sm80.ttir | 172 + .../golden/ir/ttir/golden_matmul_s3_sm90.ttir | 172 + .../ir/ttir/golden_matmul_tma_s1_sm90.ttir | 78 + .../ir/ttir/golden_matmul_tma_s3_sm90.ttir | 78 + .../ir/ttir/golden_matmul_tma_ws_s3_sm90.ttir | 78 + .../ttir/golden_nested_guard_merge_sm80.ttir | 59 + .../ir/ttir/golden_nested_loops_sm80.ttir | 45 + .../ir/ttir/golden_pid_branch_sm80.ttir | 47 + .../ir/ttir/golden_pid_branch_sm90.ttir | 47 + .../ir/ttir/golden_sequential_loops_sm80.ttir | 68 + tests/golden/ir/ttir/golden_tile2d_sm80.ttir | 91 + tests/golden/ir/ttir/golden_tile2d_sm90.ttir | 91 + tests/golden/ir/ttir/kernel_deep_chain.ttir | 1221 ++++ .../golden/ir/ttir/kernel_dot_precisions.ttir | 55 + tests/golden/ir/ttir/kernel_dot_scaled.ttir | 112 + tests/golden/ir/ttir/kernel_eps_consts.ttir | 39 + tests/golden/ir/ttir/kernel_unicode_msgs.ttir | 30 + tests/golden/ir/ttir/nat_dead_if.ttir | 36 + tests/golden/ir/ttir/nat_empty_loop.ttir | 21 + tests/golden/ir/ttir/nat_empty_then.ttir | 27 + .../golden/ir/ttir/nat_hint_arange_const.ttir | 21 + .../golden/ir/ttir/nat_hint_scalar_const.ttir | 27 + tests/golden/ir/ttir/nat_k_uni.ttir | 21 + tests/golden/ir/ttir/nat_uni_params.ttir | 22 + tests/golden/ir/ttir/spike_atomics.ttir | 59 + tests/golden/ir/ttir/spike_casts.ttir | 63 + tests/golden/ir/ttir/spike_dot.ttir | 59 + tests/golden/ir/ttir/spike_early_return.ttir | 46 + .../ir/ttir/spike_early_return_loop.ttir | 58 + .../ir/ttir/spike_for_ptr_iterargs.ttir | 71 + tests/golden/ir/ttir/spike_i64_index.ttir | 51 + tests/golden/ir/ttir/spike_if_yield.ttir | 77 + tests/golden/ir/ttir/spike_inline_asm.ttir | 27 + tests/golden/ir/ttir/spike_misc.ttir | 54 + tests/golden/ir/ttir/spike_nested_for.ttir | 61 + tests/golden/ir/ttir/spike_noinline_call.ttir | 84 + tests/golden/ir/ttir/spike_reduce_scan.ttir | 117 + tests/golden/ir/ttir/spike_spin_while.ttir | 47 + tests/golden/ir/ttir/spike_tile2d_i64.ttir | 85 + tests/golden/ir/ttir_3.8/adv_descs.ttir | 35 + tests/golden/ir/ttir_3.8/adv_zero_result.ttir | 56 + .../ttir_3.8/golden_matmul_tma_s1_sm90.ttir | 76 + .../ttir_3.8/golden_matmul_tma_s3_sm90.ttir | 76 + .../golden_matmul_tma_ws_s3_sm90.ttir | 76 + .../golden/ir/ttir_3.8/kernel_deep_chain.ttir | 1218 ++++ .../ir/ttir_3.8/kernel_dot_precisions.ttir | 50 + .../golden/ir/ttir_3.8/kernel_dot_scaled.ttir | 98 + .../golden/ir/ttir_3.8/kernel_eps_consts.ttir | 35 + .../ir/ttir_3.8/kernel_unicode_msgs.ttir | 29 + tests/golden/ir/ttir_3.8/spike_misc.ttir | 51 + tests/unit/ir/__init__.py | 0 tests/unit/ir/_goldens.py | 89 + tests/unit/ir/_oracle_ttir_reader_361.py | 2134 +++++++ tests/unit/ir/test_host_compile.py | 1302 +++++ tests/unit/ir/test_ir_capture.py | 855 +++ tests/unit/ir/test_mlir_walk.py | 1847 ++++++ tests/unit/ir/test_ttir_reader.py | 1442 +++++ tests/unit/test_ir_lifecycle.py | 2428 ++++++++ tilelens/core/client.py | 1026 +++- tilelens/core/config.py | 50 + tilelens/core/host_compile.py | 1113 ++++ tilelens/core/trace.py | 654 ++- tilelens/core/trace_io.py | 13 +- tilelens/ir/__init__.py | 41 + tilelens/ir/_mlir_walk.py | 2666 +++++++++ tilelens/ir/capture.py | 282 + tilelens/ir/client.py | 147 + tilelens/ir/launch.py | 227 + tilelens/ir/ttir_reader.py | 1756 ++++++ tilelens/ir/verdict.py | 132 + tools/ir_bulk_conformance.py | 257 + 178 files changed, 38875 insertions(+), 126 deletions(-) create mode 100644 .codespellrc create mode 100644 tests/conformance/__init__.py create mode 100644 tests/conformance/_corpus.py create mode 100644 tests/conformance/_interp_footprint.py create mode 100644 tests/conformance/_static_footprint.py create mode 100644 tests/conformance/_ttir_capture.py create mode 100644 tests/conformance/conftest.py create mode 100644 tests/conformance/test_reader_conformance.py create mode 100644 tests/end_to_end/test_ir_client.py create mode 100644 tests/end_to_end/test_ir_smoke.py create mode 100644 tests/golden/ir/expected.json create mode 100644 tests/golden/ir/expected_3.8.json create mode 100644 tests/golden/ir/generate_reader_ttir.py create mode 100644 tests/golden/ir/generate_ttir.py create mode 100644 tests/golden/ir/reader_kernels.py create mode 100644 tests/golden/ir/reader_ttir/expand_iterarg_3d.ttir create mode 100644 tests/golden/ir/reader_ttir/expand_iterarg_mask.ttir create mode 100644 tests/golden/ir/reader_ttir/int_iterarg_offset.ttir create mode 100644 tests/golden/ir/reader_ttir/iv_wrap.ttir create mode 100644 tests/golden/ir/reader_ttir/loop_observed_advance.ttir create mode 100644 tests/golden/ir/reader_ttir/loop_two_step_advance.ttir create mode 100644 tests/golden/ir/reader_ttir/observed_lanes.ttir create mode 100644 tests/golden/ir/reader_ttir/p1_variant_delta.ttir create mode 100644 tests/golden/ir/reader_ttir/p2_swap.ttir create mode 100644 tests/golden/ir/reader_ttir/p3_call_formals.ttir create mode 100644 tests/golden/ir/reader_ttir/p3_call_guarded.ttir create mode 100644 tests/golden/ir/reader_ttir/p3_call_offset.ttir create mode 100644 tests/golden/ir/reader_ttir/p4_observed_delta.ttir create mode 100644 tests/golden/ir/reader_ttir/p4_observed_direct.ttir create mode 100644 tests/golden/ir/reader_ttir/p4_observed_loop.ttir create mode 100644 tests/golden/ir/reader_ttir/pure_asm_int_addr.ttir create mode 100644 tests/golden/ir/reader_ttir/rv_i32_wrap.ttir create mode 100644 tests/golden/ir/reader_ttir/rv_inline_asm_store.ttir create mode 100644 tests/golden/ir/reader_ttir/rv_trunci_alias.ttir create mode 100644 tests/golden/ir/reader_ttir/tile3d_shared_arange.ttir create mode 100644 tests/golden/ir/reader_ttir/unsigned_index.ttir create mode 100644 tests/golden/ir/reader_ttir/where_pointer.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/expand_iterarg_3d.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/expand_iterarg_mask.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/int_iterarg_offset.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/iv_wrap.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/loop_observed_advance.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/loop_two_step_advance.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/observed_lanes.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/p1_variant_delta.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/p2_swap.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/p3_call_formals.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/p3_call_guarded.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/p3_call_offset.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/p4_observed_delta.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/p4_observed_direct.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/p4_observed_loop.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/pure_asm_int_addr.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/rv_i32_wrap.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/rv_inline_asm_store.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/rv_trunci_alias.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/tile3d_shared_arange.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/unsigned_index.ttir create mode 100644 tests/golden/ir/reader_ttir_3.8/where_pointer.ttir create mode 100644 tests/golden/ir/ttir/adv_cf_blockargs.ttir create mode 100644 tests/golden/ir/ttir/adv_consts.ttir create mode 100644 tests/golden/ir/ttir/adv_descs.ttir create mode 100644 tests/golden/ir/ttir/adv_hinted.ttir create mode 100644 tests/golden/ir/ttir/adv_multi_func.ttir create mode 100644 tests/golden/ir/ttir/adv_multi_result.ttir create mode 100644 tests/golden/ir/ttir/adv_names.ttir create mode 100644 tests/golden/ir/ttir/adv_nest3.ttir create mode 100644 tests/golden/ir/ttir/adv_reduce3.ttir create mode 100644 tests/golden/ir/ttir/adv_views.ttir create mode 100644 tests/golden/ir/ttir/adv_while_nested.ttir create mode 100644 tests/golden/ir/ttir/adv_zero_result.ttir create mode 100644 tests/golden/ir/ttir/crafted_attr_dicts.ttir create mode 100644 tests/golden/ir/ttir/crafted_deep_nest.ttir create mode 100644 tests/golden/ir/ttir/crafted_empty_bodies.ttir create mode 100644 tests/golden/ir/ttir/crafted_empty_else.ttir create mode 100644 tests/golden/ir/ttir/crafted_empty_for.ttir create mode 100644 tests/golden/ir/ttir/crafted_fwd_ref_cf.ttir create mode 100644 tests/golden/ir/ttir/crafted_generic_form.ttir create mode 100644 tests/golden/ir/ttir/crafted_locs.ttir create mode 100644 tests/golden/ir/ttir/crafted_odd_names.ttir create mode 100644 tests/golden/ir/ttir/crafted_same_dest.ttir create mode 100644 tests/golden/ir/ttir/crafted_symbols_strings.ttir create mode 100644 tests/golden/ir/ttir/crafted_unicode_strings.ttir create mode 100644 tests/golden/ir/ttir/golden_add_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_add_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_atomic_fmax_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_atomic_fmax_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_atomic_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_atomic_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_cas_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_cas_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_early_return_loaded_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_early_return_pid_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_gather_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_gather_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_grid_stride_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_guard_then_loop_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_if_else_load_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_if_else_load_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_if_else_offset_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_if_else_offset_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_loop_under_if_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_matmul_bp_s3_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_matmul_bp_s3_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_matmul_s1_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_matmul_s1_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_matmul_s3_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_matmul_s3_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_matmul_tma_s1_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_matmul_tma_s3_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_matmul_tma_ws_s3_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_nested_guard_merge_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_nested_loops_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_pid_branch_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_pid_branch_sm90.ttir create mode 100644 tests/golden/ir/ttir/golden_sequential_loops_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_tile2d_sm80.ttir create mode 100644 tests/golden/ir/ttir/golden_tile2d_sm90.ttir create mode 100644 tests/golden/ir/ttir/kernel_deep_chain.ttir create mode 100644 tests/golden/ir/ttir/kernel_dot_precisions.ttir create mode 100644 tests/golden/ir/ttir/kernel_dot_scaled.ttir create mode 100644 tests/golden/ir/ttir/kernel_eps_consts.ttir create mode 100644 tests/golden/ir/ttir/kernel_unicode_msgs.ttir create mode 100644 tests/golden/ir/ttir/nat_dead_if.ttir create mode 100644 tests/golden/ir/ttir/nat_empty_loop.ttir create mode 100644 tests/golden/ir/ttir/nat_empty_then.ttir create mode 100644 tests/golden/ir/ttir/nat_hint_arange_const.ttir create mode 100644 tests/golden/ir/ttir/nat_hint_scalar_const.ttir create mode 100644 tests/golden/ir/ttir/nat_k_uni.ttir create mode 100644 tests/golden/ir/ttir/nat_uni_params.ttir create mode 100644 tests/golden/ir/ttir/spike_atomics.ttir create mode 100644 tests/golden/ir/ttir/spike_casts.ttir create mode 100644 tests/golden/ir/ttir/spike_dot.ttir create mode 100644 tests/golden/ir/ttir/spike_early_return.ttir create mode 100644 tests/golden/ir/ttir/spike_early_return_loop.ttir create mode 100644 tests/golden/ir/ttir/spike_for_ptr_iterargs.ttir create mode 100644 tests/golden/ir/ttir/spike_i64_index.ttir create mode 100644 tests/golden/ir/ttir/spike_if_yield.ttir create mode 100644 tests/golden/ir/ttir/spike_inline_asm.ttir create mode 100644 tests/golden/ir/ttir/spike_misc.ttir create mode 100644 tests/golden/ir/ttir/spike_nested_for.ttir create mode 100644 tests/golden/ir/ttir/spike_noinline_call.ttir create mode 100644 tests/golden/ir/ttir/spike_reduce_scan.ttir create mode 100644 tests/golden/ir/ttir/spike_spin_while.ttir create mode 100644 tests/golden/ir/ttir/spike_tile2d_i64.ttir create mode 100644 tests/golden/ir/ttir_3.8/adv_descs.ttir create mode 100644 tests/golden/ir/ttir_3.8/adv_zero_result.ttir create mode 100644 tests/golden/ir/ttir_3.8/golden_matmul_tma_s1_sm90.ttir create mode 100644 tests/golden/ir/ttir_3.8/golden_matmul_tma_s3_sm90.ttir create mode 100644 tests/golden/ir/ttir_3.8/golden_matmul_tma_ws_s3_sm90.ttir create mode 100644 tests/golden/ir/ttir_3.8/kernel_deep_chain.ttir create mode 100644 tests/golden/ir/ttir_3.8/kernel_dot_precisions.ttir create mode 100644 tests/golden/ir/ttir_3.8/kernel_dot_scaled.ttir create mode 100644 tests/golden/ir/ttir_3.8/kernel_eps_consts.ttir create mode 100644 tests/golden/ir/ttir_3.8/kernel_unicode_msgs.ttir create mode 100644 tests/golden/ir/ttir_3.8/spike_misc.ttir create mode 100644 tests/unit/ir/__init__.py create mode 100644 tests/unit/ir/_goldens.py create mode 100644 tests/unit/ir/_oracle_ttir_reader_361.py create mode 100644 tests/unit/ir/test_host_compile.py create mode 100644 tests/unit/ir/test_ir_capture.py create mode 100644 tests/unit/ir/test_mlir_walk.py create mode 100644 tests/unit/ir/test_ttir_reader.py create mode 100644 tests/unit/test_ir_lifecycle.py create mode 100644 tilelens/core/host_compile.py create mode 100644 tilelens/ir/__init__.py create mode 100644 tilelens/ir/_mlir_walk.py create mode 100644 tilelens/ir/capture.py create mode 100644 tilelens/ir/client.py create mode 100644 tilelens/ir/launch.py create mode 100644 tilelens/ir/ttir_reader.py create mode 100644 tilelens/ir/verdict.py create mode 100644 tools/ir_bulk_conformance.py 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/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_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_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/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_mlir_walk.py b/tests/unit/ir/test_mlir_walk.py new file mode 100644 index 000000000..8edf48687 --- /dev/null +++ b/tests/unit/ir/test_mlir_walk.py @@ -0,0 +1,1847 @@ +"""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"') + + +@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 path + 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 _GENERATOR_LOC.sub('loc("G"', got) == _GENERATOR_LOC.sub( + 'loc("G"', 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/test_ir_lifecycle.py b/tests/unit/test_ir_lifecycle.py new file mode 100644 index 000000000..172fefc7b --- /dev/null +++ b/tests/unit/test_ir_lifecycle.py @@ -0,0 +1,2428 @@ +"""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 list(manager.clients) == ["ir_skip", "eager"] + + with pytest.raises(RuntimeError, match="LAUNCH='skip'"): + ClientManager([_RunIRClient(), _SkipIRClient()]) + + +def test_launch_conflict_check_sees_the_resulting_ir_clients(): + # A same-NAME client replaces the one it would otherwise conflict with. + class _RunInSkipSlot(_IRClient): + NAME = "ir_skip" + LAUNCH = "run" + + manager = ClientManager([_SkipIRClient()]) + manager.add_clients([_RunInSkipSlot()]) + assert [type(c) for c in manager.clients.values()] == [_RunInSkipSlot] + assert manager.launch_policy() == "run" + + # 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_duplicate_rule_and_accepts_indifferent(): + first = _SkipIRClient() + manager = ClientManager([first, _IndifferentIRClient(), _EagerClient()]) + manager.add_clients([_SkipIRClient()]) + + assert list(manager.clients) == ["ir_skip", "ir_indifferent", "eager"] + assert manager.clients["ir_skip"] is first + + +def test_add_clients_rejects_unknown_launch_value(): + class _BadIRClient(_IRClient): + NAME = "ir_bad" + LAUNCH = "maybe" + + with pytest.raises(ValueError, match="LAUNCH must be one of"): + ClientManager([_BadIRClient()]) + + +def test_trace_decorator_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 list(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.clients["ir_run"].events[0].launched is False + assert manager.clients["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_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.clients["ir_skip"].launch_calls + assert call.jit_fn is None and call.capture is False and call.grid == (2,) + + log = [] + traced = trace_cls.__new__(trace_cls) + TraceInterface.__init__(traced, _BrokenBegin(log)) + with pytest.raises(KeyError): + traced[(2,)](torch.zeros(4)) + assert log == ["begin", ("abort", KeyError)] + + +# ======== 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/tilelens/core/client.py b/tilelens/core/client.py index 940dbf506..4ddad941a 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,151 @@ 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): 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. 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 +239,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,6 +323,167 @@ 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] = {} @@ -123,7 +491,36 @@ def __init__(self, clients: list[Client] | None = None): self.add_clients(clients) self.launch = Launch() self._lock = threading.Lock() + # Compiles every IR client's kernels on the host (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: @@ -134,68 +531,560 @@ def get_client(self, name: str) -> Client | None: return self.clients.get(name) def add_clients(self, new_clients_list: list[Client]) -> None: + # Validate the whole resulting set before inserting anything, so a + # rejected composition leaves the manager unchanged. + additions: dict[str, Client] = {} for new_client in new_clients_list: duplicate = any( isinstance(existing_client, new_client.__class__) - for existing_client in self.clients.values() + for existing_client in (*self.clients.values(), *additions.values()) ) if not duplicate: - self.clients[new_client.NAME] = new_client + additions[new_client.NAME] = new_client + # A same-NAME addition replaces the existing client, so check the set + # that will result, not the one before replacement. + self._check_launch_preferences(list({**self.clients, **additions}.values())) + self.clients.update(additions) + + @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." + ) + + def interpreting_clients(self) -> list[Client]: + return [c for c in self.clients.values() if c.NEEDS_INTERPRETER] + + def ir_clients(self) -> list[Client]: + return [c for c in self.clients.values() if not c.NEEDS_INTERPRETER] + + def 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.values(): + begun.append(client) + client.begin_launch(call) + except BaseException as exc: + self._abort_clients(begun, exc) + raise + + def _claim_launch(self) -> None: + thread = threading.get_ident() + with self._owner_lock: + if self._launch_owner not in (None, thread): + raise RuntimeError( + "this trace is already running a launch on another host " + "thread; concurrent launches of one traced kernel are not " + "supported." + ) + self._launch_owner = thread + + def _release_launch(self) -> None: + with self._owner_lock: + if self._launch_owner == threading.get_ident(): + self._launch_owner = None + + def abort_launch(self, exc: BaseException) -> None: + """Deliver ``exc`` to every client's abort_launch. + + Nothing is sent once finalize has started: every client is finalized + by then, and each launch ends in either finalize or abort. Nor is + anything sent while another host thread's launch is in flight (ours + has ended already). A failing hook never replaces ``exc``, which the + caller re-raises; the failure is attached to it as a note. Only a + hook's KeyboardInterrupt or SystemExit is raised, after every client + got the abort. + """ + thread = threading.get_ident() + with self._owner_lock: + if self._launch_owner not in (None, thread): + return + if self._finalize_started: + self._launch_owner = None + return + # Held until every client got the abort, so no other thread's + # begin_launch resets the state in between. + self._launch_owner = thread + self._abort_clients(list(self.clients.values()), exc) + + def _abort_clients(self, clients: list[Client], exc: BaseException) -> None: + interrupt: BaseException | None = None + try: + for client in clients: + try: + client.abort_launch(exc) + except Exception as hook_exc: + message = ( + f"{type(client).__name__}.abort_launch raised " + f"{type(hook_exc).__name__}: {hook_exc}" + ) + if hasattr(exc, "add_note"): + exc.add_note(message) + else: # Python 3.10 + warnings.warn(message, RuntimeWarning, stacklevel=3) + except BaseException as hook_exc: + if interrupt is None: + interrupt = hook_exc + finally: + self._release_pinned() + self._release_launch() + if interrupt is not None: + raise interrupt from exc @contextmanager - def patch_warmup(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.values() + ] + 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.values(): + 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 +1093,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.values(): + try: + # client may introduce tensors not declared in kernel args (e.g. tracer recording a tensor allocation) + self.launch.tensors.update(getattr(client, "tensors", []) or []) + self.launch.records += client.finalize() + except BaseException as exc: + if first_exc is None: + first_exc = exc + if first_exc is not None: + raise first_exc + finally: + self._release_pinned() + self._release_launch() def arg_callback(self, name, arg, arg_cvt): with self._lock_context(): if hasattr(arg, "data_ptr"): self.launch.tensors.add(arg) - for client in self.clients.values(): + for client in self.interpreting_clients(): client.arg_callback(name, arg, arg_cvt) def grid_callback(self, grid: tuple[int]): with self._lock_context(): self.launch.grid = grid - for client in self.clients.values(): + for client in self.interpreting_clients(): client.grid_callback(grid) def grid_idx_callback(self, grid_idx: tuple[int, ...]): with self._lock_context(): - for client in self.clients.values(): + for client in self.interpreting_clients(): client.grid_idx_callback(grid_idx) # --- For-loop callback management --- diff --git a/tilelens/core/config.py b/tilelens/core/config.py index 9163d60bd..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..b7de78e94 100644 --- a/tilelens/core/trace.py +++ b/tilelens/core/trace.py @@ -1,12 +1,15 @@ -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 @@ -14,6 +17,59 @@ launches: list[Launch] = [] +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: self.client_manager = ClientManager() @@ -41,8 +97,31 @@ def add_client(self, new_client: str | Client) -> None: self.client_manager.add_clients([self._normalize_client(new_client)]) def finalize(self): + # Take the Launch first: once finalize ends the launch, another host + # thread may begin the next one on this manager. + launch = self.client_manager.launch self.client_manager.finalize() - launches.append(self.client_manager.launch) + launches.append(launch) + + @contextmanager + def _launch_scope(self, call: LaunchCall): + """begin_launch, then abort_launch if the launch raises. A refused or + failing begin_launch cleans up after itself and is not aborted: the + refusal must not reach the clients of another thread's launch.""" + mgr = self.client_manager + mgr.begin_launch(call) + try: + yield + except BaseException as exc: + mgr.abort_launch(exc) + raise + + 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 +163,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 +353,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 +368,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 +785,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 +849,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/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/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()) From 616cb5449c67597c1598fcf8d99e84b4ccd288a9 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 2 Oct 2026 21:18:59 -0400 Subject: [PATCH 2/5] [FIX] Compare regenerated goldens without Triton's install path The golden regeneration test took the generator's path out of the locs but not Triton's: a golden whose kernel calls into Triton's own sources (tl.cdiv, tl.zeros, ...) names the directory Triton is installed at, so it never regenerated byte for byte on another machine (CI failed on golden_matmul_tma_s1_sm90 under Triton 3.8). Take both paths out before comparing. --- tests/unit/ir/test_mlir_walk.py | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/tests/unit/ir/test_mlir_walk.py b/tests/unit/ir/test_mlir_walk.py index 8edf48687..b53080318 100644 --- a/tests/unit/ir/test_mlir_walk.py +++ b/tests/unit/ir/test_mlir_walk.py @@ -155,6 +155,14 @@ def _real_compiles_available() -> bool: # 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( @@ -163,9 +171,10 @@ def _real_compiles_available() -> bool: ) 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 path - aside): its locs name the kernels' lines, so a line added above them - shows here, not as goldens that no longer regenerate.""" + 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. @@ -192,9 +201,7 @@ def test_the_kernel_goldens_regenerate_byte_for_byte(monkeypatch, tmp_path): for name, spec in todo.items(): want = (out / f"{name}.ttir").read_text(encoding="utf-8") got = gen.ttir(spec) - assert _GENERATOR_LOC.sub('loc("G"', got) == _GENERATOR_LOC.sub( - 'loc("G"', want - ), name + assert _portable(got) == _portable(want), name @pytest.mark.parametrize("label", PINNED) From 6d331a9776285b8278c861a8f2618414a6ad293e Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Tue, 29 Sep 2026 14:44:58 -0400 Subject: [PATCH 3/5] [FEAT] Add a compiled mode to the sanitizer Sanitizer(compile=True), or tile-sanitizer --compile, checks a launch statically on the host-compiled TTIR instead of interpreting it. The kernel is not launched, so no GPU is needed and outputs are not written. - Every autotune config is checked with Z3 against the launch's scalar arguments, grid and tensor layouts; a launch is ok only if every config is proven in bounds. - Findings: out-of-bounds (against the view's element footprint, gaps of strided views included), integer-overflow on address, mask, branch and loop values, and division-by-zero, each with a witness. - Anything that cannot be modeled is reported as unsupported with a typed refusal; a solver unknown or timeout never counts as a proof. A kernel that fails to compile for the target is unsupported and the program continues; a call that does not bind raises as untraced. - Each launch records an IRVerdict with per-config verdicts in Launch.records. --- README.md | 97 + tests/end_to_end/test_compiled_sanitizer.py | 1862 +++++++++++++++++ tests/end_to_end/test_host_compile.py | 377 ++++ .../end_to_end/test_ir_lifecycle_compiled.py | 1090 ++++++++++ tests/unit/ir/test_verdict_io.py | 305 +++ tests/unit/sanitizer_compiled/__init__.py | 0 tests/unit/sanitizer_compiled/test_client.py | 1061 ++++++++++ tests/unit/sanitizer_compiled/test_oob.py | 1235 +++++++++++ tests/unit/test_ir_version_gate.py | 452 ++++ tests/unit/test_wrapper.py | 82 + tilelens/clients/__init__.py | 8 + .../clients/sanitizer/compiled/__init__.py | 42 + tilelens/clients/sanitizer/compiled/client.py | 776 +++++++ tilelens/clients/sanitizer/compiled/oob.py | 1248 +++++++++++ tilelens/clients/sanitizer/data.py | 55 +- tilelens/clients/sanitizer/sanitizer.py | 31 +- tilelens/wrapper.py | 43 +- 17 files changed, 8754 insertions(+), 10 deletions(-) create mode 100644 tests/end_to_end/test_compiled_sanitizer.py create mode 100644 tests/end_to_end/test_host_compile.py create mode 100644 tests/end_to_end/test_ir_lifecycle_compiled.py create mode 100644 tests/unit/ir/test_verdict_io.py create mode 100644 tests/unit/sanitizer_compiled/__init__.py create mode 100644 tests/unit/sanitizer_compiled/test_client.py create mode 100644 tests/unit/sanitizer_compiled/test_oob.py create mode 100644 tests/unit/test_ir_version_gate.py create mode 100644 tilelens/clients/sanitizer/compiled/__init__.py create mode 100644 tilelens/clients/sanitizer/compiled/client.py create mode 100644 tilelens/clients/sanitizer/compiled/oob.py diff --git a/README.md b/README.md index 0efc0527e..0951d246f 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,97 @@ 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. +- 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/end_to_end/test_compiled_sanitizer.py b/tests/end_to_end/test_compiled_sanitizer.py new file mode 100644 index 000000000..ed9083370 --- /dev/null +++ b/tests/end_to_end/test_compiled_sanitizer.py @@ -0,0 +1,1862 @@ +"""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 + + +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_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_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/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..9ff569742 --- /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 set(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..ef67db839 --- /dev/null +++ b/tests/unit/sanitizer_compiled/test_oob.py @@ -0,0 +1,1235 @@ +"""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 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")) + + +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 + + +_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_version_gate.py b/tests/unit/test_ir_version_gate.py new file mode 100644 index 000000000..8eba829a6 --- /dev/null +++ b/tests/unit/test_ir_version_gate.py @@ -0,0 +1,452 @@ +"""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_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..b31c062a6 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(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..360399ab6 --- /dev/null +++ b/tilelens/clients/sanitizer/compiled/oob.py @@ -0,0 +1,1248 @@ +"""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 can be deeper than Python's recursion limit and their generated +``==`` / ``hash`` recurse, so lowering is iterative and memoized by term +identity. +""" + +from __future__ import annotations + +from collections.abc import Hashable, Iterable, Iterator, Mapping, Sequence +from dataclasses import dataclass, replace +from enum import Enum +from typing import Any, Literal + +from z3 import ( + And, + ArithRef, + BoolRef, + BoolVal, + Context, + Exists, + If, + Implies, + Int, + IntVal, + ModelRef, + Or, + Solver, + Sum, + is_bool, + is_false, + is_int_value, + sat, + simplify, + unknown, + unsat, +) +from z3 import Not as Z3Not + +from ....ir.launch import LaunchBinding, TensorFacts +from ....ir.ttir_reader import ( + AccessEvent, + AccessGraph, + Arange, + Bin, + BoolBin, + Cmp, + Const, + DataDep, + IntCast, + IterArgOffset, + LoopInfo, + LoopVar, + Not, + 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 _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 _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//"): + return _trunc_div(a, b) + if op in ("%", "u%"): + # arith.remsi: the remainder carries the dividend's sign + return a - b * _trunc_div(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_PRED = {"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) + table = { + "slt": a < b, "sle": a <= b, "sgt": a > b, + "sge": a >= b, "eq": a == b, "ne": a != b, + } # fmt: skip + try: + return table[_SIGNED_PRED.get(pred, pred)] + except KeyError: + raise ValueError(f"unknown cmpi predicate {pred!r}") from None + + +def _kids(t: object, graph: AccessGraph) -> tuple: + """The terms ``t`` is computed from; a loop-carried pointer's offset + is computed from its IterArgInfo's ``offset0`` and ``delta``.""" + 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) + return () + + +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(_kids(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 lowering memo.""" + + 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 lane). + self.positions: dict[tuple[int, int], ArithRef] = {} + # witness name -> the arange's value at the lane + self.lanes: dict[str, ArithRef] = {} + self.iteration: ArithRef | None = None + self._observed: dict[int, ArithRef] = {} + # id(term) -> (term, lowered): holding the term keeps its id unique. + self._memo: dict[int, tuple[object, Any]] = {} + + # ── leaves ── + + def param(self, name: str) -> int: + try: + value = self.binding.params[name] + except KeyError: + raise _Refused( + SanitizerKind.MISSING_BINDING, + f"scalar argument {name!r} has no launch binding" + + _binding_error(self.binding), + ) from None + arg = self.graph.arg(name) + return _signed(value, arg.int_bits if arg is not None else 0) + + def lane(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 observed(self, index: int) -> ArithRef: + """An atomic observation: a free value of the atomic's width.""" + 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 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.iteration 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.iteration = k + return self.iteration + + 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 _as_int(self.lower(term)) + + def cond(self, term: object) -> BoolRef: + return _as_bool(self.lower(term)) + + def lower(self, root: object) -> Any: + """``root`` as a Z3 expression (Int, or Bool for a compare).""" + memo = self._memo + stack: list[tuple[object, bool]] = [(root, False)] + while stack: + t, ready = stack.pop() + if id(t) in memo: + continue + kids = _kids(t, self.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, self._apply(t, [memo[id(k)][1] for k in kids])) + return memo[id(root)][1] + + def _apply(self, t: object, kids: Sequence[Any]) -> Any: + if isinstance(t, Const): + return IntVal(t.value, self.ctx) + if isinstance(t, Param): + return IntVal(self.param(t.name), self.ctx) + if isinstance(t, Pid): + return self.pids[t.axis] + if isinstance(t, NumPrograms): + return IntVal(self.grid[t.axis], self.ctx) + if isinstance(t, Arange): + return self.lane(t) + if isinstance(t, LoopVar): + lower, _upper, step = self.bounds() + return lower + self.loop_iteration() * step + if isinstance(t, IterArgOffset): + return _as_int(kids[0]) + self.loop_iteration() * _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]) + return And(a, b) if t.op == "and" else Or(a, b) + 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 self.observed(t.access_index) + if isinstance(t, DataDep): + raise _Refused( + SanitizerKind.UNMODELED_VALUE, f"an unmodeled value ({t.why})" + ) + raise TypeError(f"unknown term {type(t).__name__}") + + +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(_kids(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 _kids(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 _kids(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: + refusal = r + 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.iteration is not None: + out[str(env.iteration)] = val(env.iteration) + 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/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, ) From eafc0a0e392455428e067477c168b285c1aad996 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Wed, 30 Sep 2026 07:07:38 -0400 Subject: [PATCH 4/5] [REFACTOR] Share the Term to Z3 lowering between compiled clients The compiled sanitizer's evaluator becomes tilelens.ir.lowering, the one reading of the TTIR reader's term algebra as Z3 terms, so the compiled race detector can sit on the same semantics instead of a second copy. - lowering.py is mechanism only: operator semantics (truncating division, unsigned ops and predicates as their signed twins under the width obligations, IntCast as its operand, i1 coercion, loop terms) live there; every leaf (scalar arguments, pids, grid, lanes, the loop iteration, observations, unmodeled values) comes from the client's TermLeaves, which also raises the client's own refusals. - It never creates a Z3 context: constants use the leaves' context and every other term its operands', so a client keeps all of its terms in one context (per check for the sanitizer). - The walk is iterative and memoized by term identity, so terms deeper than the recursion limit lower. - The sanitizer keeps its policy (obligation findings, refusals, solving, timeouts). Its leaves reach the lowering through a weak proxy and a refused loop keeps a never-raised copy of its refusal, so a check's Z3 context is freed when the check returns instead of by the cyclic GC on another thread (concurrent checks could hang or crash). - Its walks stop at the induction variable, as before, so a loop bound read only through a Select arm keeps that arm's guard. No verdict, finding or witness changes: the sanitizer's results match the previous evaluator on the golden texts and the differential corpus on Triton 3.6 and 3.8. Hand-built graphs outside the reader's invariants (e.g. an iter arg naming another loop) now raise instead of being misread. --- tests/unit/ir/test_lowering.py | 417 +++++++++++++++++++++ tests/unit/sanitizer_compiled/test_oob.py | 102 +++++ tests/unit/test_ir_version_gate.py | 1 + tilelens/clients/sanitizer/compiled/oob.py | 249 ++++-------- tilelens/ir/lowering.py | 291 ++++++++++++++ 5 files changed, 889 insertions(+), 171 deletions(-) create mode 100644 tests/unit/ir/test_lowering.py create mode 100644 tilelens/ir/lowering.py 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/sanitizer_compiled/test_oob.py b/tests/unit/sanitizer_compiled/test_oob.py index ef67db839..b27bc83c2 100644 --- a/tests/unit/sanitizer_compiled/test_oob.py +++ b/tests/unit/sanitizer_compiled/test_oob.py @@ -9,6 +9,7 @@ from __future__ import annotations +import gc import pickle import subprocess import sys @@ -947,6 +948,67 @@ def test_a_value_also_read_directly_is_checked_unguarded(): 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.""" @@ -1168,6 +1230,46 @@ def test_a_real_ctrl_c_stops_a_hard_query(): 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 diff --git a/tests/unit/test_ir_version_gate.py b/tests/unit/test_ir_version_gate.py index 8eba829a6..b7dfebb55 100644 --- a/tests/unit/test_ir_version_gate.py +++ b/tests/unit/test_ir_version_gate.py @@ -391,6 +391,7 @@ 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", diff --git a/tilelens/clients/sanitizer/compiled/oob.py b/tilelens/clients/sanitizer/compiled/oob.py index 360399ab6..a49986f24 100644 --- a/tilelens/clients/sanitizer/compiled/oob.py +++ b/tilelens/clients/sanitizer/compiled/oob.py @@ -61,17 +61,20 @@ Ctrl+C that Z3 caught during a query is raised as ``KeyboardInterrupt``, never taken for an unknown. -Terms can be deeper than Python's recursion limit and their generated -``==`` / ``hash`` recurse, so lowering is iterative and memoized by term -identity. +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 +from typing import Any, Literal, NoReturn from z3 import ( And, @@ -80,7 +83,6 @@ BoolVal, Context, Exists, - If, Implies, Int, IntVal, @@ -88,7 +90,6 @@ Or, Solver, Sum, - is_bool, is_false, is_int_value, sat, @@ -99,20 +100,15 @@ 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, - BoolBin, - Cmp, - Const, DataDep, - IntCast, - IterArgOffset, LoopInfo, LoopVar, - Not, NumPrograms, Observed, Param, @@ -312,81 +308,13 @@ def at(self, line_no: int | None, loc: Any) -> _Refused: return self -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 _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//"): - return _trunc_div(a, b) - if op in ("%", "u%"): - # arith.remsi: the remainder carries the dividend's sign - return a - b * _trunc_div(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_PRED = {"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) - table = { - "slt": a < b, "sle": a <= b, "sgt": a > b, - "sge": a >= b, "eq": a == b, "ne": a != b, - } # fmt: skip - try: - return table[_SIGNED_PRED.get(pred, pred)] - except KeyError: - raise ValueError(f"unknown cmpi predicate {pred!r}") from None - - -def _kids(t: object, graph: AccessGraph) -> tuple: - """The terms ``t`` is computed from; a loop-carried pointer's offset - is computed from its IterArgInfo's ``offset0`` and ``delta``.""" - 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) - return () +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( @@ -401,7 +329,7 @@ def _walk( continue seen.add(id(t)) yield t - stack.extend(reversed(_kids(t, graph))) + stack.extend(reversed(_operands(t, graph))) _DIVISIONS = frozenset({"//", "%", "u//", "u%"}) @@ -435,7 +363,8 @@ def _lane_name(t: Arange) -> str: class _Env: """The Z3 variables of one family of queries (one access, or the - loop's own checks), their range premises, and the lowering memo.""" + loop's own checks), their range premises, and the family's lowering: + the env is its leaves (``tilelens.ir.lowering.TermLeaves``).""" def __init__( self, @@ -453,30 +382,41 @@ def __init__( 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 lane). + # (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.iteration: ArithRef | None = None + self.k: ArithRef | None = None # the loop's iteration index, once used self._observed: dict[int, ArithRef] = {} - # id(term) -> (term, lowered): holding the term keeps its id unique. - self._memo: dict[int, tuple[object, Any]] = {} + # 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, name: str) -> int: - try: - value = self.binding.params[name] - except KeyError: + 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 {name!r} has no launch binding" + f"scalar argument {t.name!r} has no launch binding" + _binding_error(self.binding), - ) from None - arg = self.graph.arg(name) - return _signed(value, arg.int_bits if arg is not None else 0) + ) + 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 lane(self, t: Arange) -> ArithRef: + 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).""" @@ -491,8 +431,20 @@ def lane(self, t: Arange) -> ArithRef: self.lanes.setdefault(_lane_name(t), value) return value - def observed(self, index: int) -> ArithRef: + 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) @@ -503,6 +455,11 @@ def observed(self, index: int) -> ArithRef: 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) @@ -511,13 +468,13 @@ 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.iteration is None: + 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.iteration = k - return self.iteration + self.k = k + return self.k def _loop(self) -> LoopInfo: loop = self.graph.loop @@ -530,67 +487,10 @@ def _loop(self) -> LoopInfo: # ── terms ── def value(self, term: object) -> ArithRef: - return _as_int(self.lower(term)) + return self.lowering.value(term) def cond(self, term: object) -> BoolRef: - return _as_bool(self.lower(term)) - - def lower(self, root: object) -> Any: - """``root`` as a Z3 expression (Int, or Bool for a compare).""" - memo = self._memo - stack: list[tuple[object, bool]] = [(root, False)] - while stack: - t, ready = stack.pop() - if id(t) in memo: - continue - kids = _kids(t, self.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, self._apply(t, [memo[id(k)][1] for k in kids])) - return memo[id(root)][1] - - def _apply(self, t: object, kids: Sequence[Any]) -> Any: - if isinstance(t, Const): - return IntVal(t.value, self.ctx) - if isinstance(t, Param): - return IntVal(self.param(t.name), self.ctx) - if isinstance(t, Pid): - return self.pids[t.axis] - if isinstance(t, NumPrograms): - return IntVal(self.grid[t.axis], self.ctx) - if isinstance(t, Arange): - return self.lane(t) - if isinstance(t, LoopVar): - lower, _upper, step = self.bounds() - return lower + self.loop_iteration() * step - if isinstance(t, IterArgOffset): - return _as_int(kids[0]) + self.loop_iteration() * _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]) - return And(a, b) if t.op == "and" else Or(a, b) - 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 self.observed(t.access_index) - if isinstance(t, DataDep): - raise _Refused( - SanitizerKind.UNMODELED_VALUE, f"an unmodeled value ({t.why})" - ) - raise TypeError(f"unknown term {type(t).__name__}") + return self.lowering.cond(term) class _Dag: @@ -617,7 +517,9 @@ def __init__( continue stack.append((t, True)) stack.extend( - (k, False) for k in reversed(_kids(t, graph)) if id(k) not in self.order + (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: @@ -637,7 +539,7 @@ def __init__( 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 _kids(t, graph)] + edges = [(k, guard) for k in _operands(t, graph)] for kid, kid_guard in edges: if kid_guard is None: direct.add(id(kid)) @@ -652,7 +554,9 @@ def rank(self, term: object) -> float: if index is not None: return index kids = [ - self.order[id(k)] for k in _kids(term, self.graph) if id(k) in self.order + 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) @@ -875,7 +779,10 @@ def check_loop(self, first: int, grid: tuple[int, int, int]) -> _Refused | None: _lower, _upper, step = env.bounds() refusal = self.step_refusal(env, step) except _Refused as r: - refusal = 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) @@ -1226,8 +1133,8 @@ def val(v: ArithRef) -> int: 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.iteration is not None: - out[str(env.iteration)] = val(env.iteration) + 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]: 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__}") From 64aed2e0aeb3206a40aa219ca7f9ba00f5fafc83 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 2 Oct 2026 23:12:44 -0400 Subject: [PATCH 5/5] [FIX] Keep every client a trace is given ClientManager.add_clients dropped a client silently when one of its class was already in the trace, and a client of another class with the same NAME replaced the first one. So two Sanitizer(compile=True) instances for different targets left only the first: the second never compiled, its last_status stayed None, and nothing said so. A client passed to a trace is now in it afterwards, or add_clients raises: - adding the same object again changes nothing; - IR clients are all kept, so a trace can hold one instance of a class per target, each compiled for its own target with its own verdict; - a second interpreting client of one NAME raises ValueError, since one interpreted run serves one client per name; - a client named by string ("sanitizer") is a no-op when the trace already has one of that name, as no settings of the caller's are lost. ClientManager.clients is now a list in trace order, and get_client returns the first client of a NAME. --- README.md | 6 +- tests/end_to_end/test_compiled_sanitizer.py | 20 +++++ tests/end_to_end/test_core.py | 24 +++--- tests/end_to_end/test_race_detector.py | 2 +- tests/unit/sanitizer_compiled/test_client.py | 2 +- tests/unit/test_ir_lifecycle.py | 80 +++++++++++++++----- tests/unit/test_wrapper.py | 2 +- tilelens/core/client.py | 68 +++++++++++------ tilelens/core/trace.py | 32 +++++--- 9 files changed, 172 insertions(+), 64 deletions(-) diff --git a/README.md b/README.md index 0951d246f..bd52af756 100644 --- a/README.md +++ b/README.md @@ -238,7 +238,11 @@ def kernel(x_ptr, n, BLOCK: tl.constexpr): `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. + 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 diff --git a/tests/end_to_end/test_compiled_sanitizer.py b/tests/end_to_end/test_compiled_sanitizer.py index ed9083370..69991b5c6 100644 --- a/tests/end_to_end/test_compiled_sanitizer.py +++ b/tests/end_to_end/test_compiled_sanitizer.py @@ -1577,6 +1577,26 @@ 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.""" 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_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/unit/sanitizer_compiled/test_client.py b/tests/unit/sanitizer_compiled/test_client.py index 9ff569742..ada321411 100644 --- a/tests/unit/sanitizer_compiled/test_client.py +++ b/tests/unit/sanitizer_compiled/test_client.py @@ -285,7 +285,7 @@ def test_declarations_and_composition(): # 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 set(manager.clients) == {"compiled_sanitizer", "sanitizer"} + 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()]) diff --git a/tests/unit/test_ir_lifecycle.py b/tests/unit/test_ir_lifecycle.py index 172fefc7b..1d4900694 100644 --- a/tests/unit/test_ir_lifecycle.py +++ b/tests/unit/test_ir_lifecycle.py @@ -465,22 +465,24 @@ def test_add_clients_rejects_skip_run_conflict_before_inserting(): manager.add_clients([_IndifferentIRClient(), _RunIRClient()]) # Nothing from the rejected batch was inserted. - assert list(manager.clients) == ["ir_skip", "eager"] + 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_the_resulting_ir_clients(): - # A same-NAME client replaces the one it would otherwise conflict with. +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" - manager = ClientManager([_SkipIRClient()]) - manager.add_clients([_RunInSkipSlot()]) - assert [type(c) for c in manager.clients.values()] == [_RunInSkipSlot] - assert manager.launch_policy() == "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): @@ -492,13 +494,34 @@ class _EagerSkip(_EagerClient): assert manager.launch_policy() == "run" -def test_add_clients_keeps_duplicate_rule_and_accepts_indifferent(): - first = _SkipIRClient() - manager = ClientManager([first, _IndifferentIRClient(), _EagerClient()]) - manager.add_clients([_SkipIRClient()]) +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 list(manager.clients) == ["ir_skip", "ir_indifferent", "eager"] - assert manager.clients["ir_skip"] is first + 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(): @@ -516,7 +539,7 @@ def test_trace_decorator_rejects_conflicting_launch_preferences(): with pytest.raises(RuntimeError, match="cannot share one trace"): tilelens.trace(_RunIRClient())(traced) - assert list(traced.client_manager.clients) == ["ir_skip"] + assert [c.NAME for c in traced.client_manager.clients] == ["ir_skip"] def test_client_partition_and_launch_policy(): @@ -989,8 +1012,8 @@ def test_ir_capture_warmup_call_never_launches(): jit_fn.run(torch.zeros(4), 4, grid=None, warmup=True) assert log == ["compile", "before", "after"] - assert manager.clients["ir_run"].events[0].launched is False - assert manager.clients["ir_run"].events[0].resolved_grid is None + 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(): @@ -1286,6 +1309,29 @@ def test_each_target_compiles_once_and_reaches_only_its_clients(_fake_host_compi 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 @@ -2283,7 +2329,7 @@ def begin_launch(self, call): # 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.clients["ir_skip"].launch_calls + (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 = [] diff --git a/tests/unit/test_wrapper.py b/tests/unit/test_wrapper.py index b31c062a6..ab5d3e3d8 100644 --- a/tests/unit/test_wrapper.py +++ b/tests/unit/test_wrapper.py @@ -295,7 +295,7 @@ def kernel(x_ptr): pass -print(sorted(kernel.client_manager.clients), sys.argv[1:]) +print(sorted(c.NAME for c in kernel.client_manager.clients), sys.argv[1:]) """ diff --git a/tilelens/core/client.py b/tilelens/core/client.py index 4ddad941a..bee44f790 100644 --- a/tilelens/core/client.py +++ b/tilelens/core/client.py @@ -143,6 +143,9 @@ class CompileGroup: 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 @@ -163,8 +166,9 @@ class Client(ABC): # 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. Core compiles once per - # distinct target and gives each client only its own target's events. + # "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: @@ -486,7 +490,8 @@ def gate(*args, **kwargs): 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() @@ -528,23 +533,42 @@ 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. - additions: dict[str, Client] = {} + 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(), *additions.values()) - ) - if not duplicate: - additions[new_client.NAME] = new_client - # A same-NAME addition replaces the existing client, so check the set - # that will result, not the one before replacement. - self._check_launch_preferences(list({**self.clients, **additions}.values())) - self.clients.update(additions) + 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: @@ -569,10 +593,10 @@ def _check_launch_preferences(clients: list[Client]) -> None: ) def interpreting_clients(self) -> list[Client]: - return [c for c in self.clients.values() if c.NEEDS_INTERPRETER] + return [c for c in self.clients if c.NEEDS_INTERPRETER] def ir_clients(self) -> list[Client]: - return [c for c in self.clients.values() if not c.NEEDS_INTERPRETER] + 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 @@ -626,7 +650,7 @@ def begin_launch(self, call: LaunchCall) -> None: self._reset_launch_state() begun: list[Client] = [] try: - for client in self.clients.values(): + for client in self.clients: begun.append(client) client.begin_launch(call) except BaseException as exc: @@ -670,7 +694,7 @@ def abort_launch(self, exc: BaseException) -> None: # Held until every client got the abort, so no other thread's # begin_launch resets the state in between. self._launch_owner = thread - self._abort_clients(list(self.clients.values()), exc) + self._abort_clients(list(self.clients), exc) def _abort_clients(self, clients: list[Client], exc: BaseException) -> None: interrupt: BaseException | None = None @@ -747,7 +771,7 @@ def _warmup_by_vote( # not short-circuit on the first True. votes = [ client.pre_warmup_callback(jit_fn, *args, **kwargs) - for client in self.clients.values() + for client in self.clients ] if not any(votes): return None @@ -757,7 +781,7 @@ def _warmup_by_vote( args, kwargs = real_args(jit_fn, args, kwargs) with compile_context(): ret = warmup(*args, **kwargs) - for client in self.clients.values(): + for client in self.clients: client.post_warmup_callback(jit_fn, ret) return ret @@ -1112,7 +1136,7 @@ def finalize(self) -> None: # 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.values(): + 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 []) diff --git a/tilelens/core/trace.py b/tilelens/core/trace.py index b7de78e94..fd0e24642 100644 --- a/tilelens/core/trace.py +++ b/tilelens/core/trace.py @@ -16,6 +16,21 @@ 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. @@ -78,22 +93,19 @@ 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):