From 2403e0ad10648fb1b3f6ac9e444f08429fad5764 Mon Sep 17 00:00:00 2001 From: Nick Riasanovsky Date: Mon, 5 Oct 2026 19:01:14 -0700 Subject: [PATCH] Replace num_stages=0 with num_stages=1 (#12) Summary: Pull Request resolved: https://github.com/facebookresearch/ads_model_kernel_library/pull/12 X-link: https://github.com/facebookexperimental/CUTracer/pull/231 `num_stages=0` has been removed from TLX (and the AMD path is updated to match), so every remaining use needs to move to `num_stages=1`. This diff is a mechanical sweep of fbsource for all the places a `0` could reach a Triton kernel. Three mechanisms are covered: 1. `triton.Config(..., num_stages=0)` -- the kernel-level autotune param. This is the bulk of the change: the TLX idiom of passing `num_stages=0` to suppress the compiler's auto-pipelining on top of manual TLX pipelining. Covers `hammer/v2`, `hammer/v3`, `ads_mkl/ops/tlx`, `simplicial_attention`, `tritonbench` fusionbench, and the TLX tutorials. 2. `tl.range(..., num_stages=0)` -- the loop-level `tt.num_stages` attribute, same intent, different knob. 3. Autotune sweeps that generate a 0 -- e.g. `for num_stages in [0, 1, 2]` in `fbr/flash/triton`, the HIP-only `[0, 1]` sweeps in the ragged HSTU attention scripts, `openfold_triton`, and FlagGems. The `0` is dropped rather than rewritten to a duplicate `1` (`[0, 1, 2]` -> `[1, 2]`, not `[1, 1, 2]`). TLX docs and agent skills that instruct readers/agents to emit `num_stages=0` (`RecGenHITL` language profiles, `kperfagent` TLX prompt skills, `ace` kernel_info) are updated too, so regenerated kernels do not reintroduce it. Deliberately not touched: - `not_allowed_values = {0}` guards in `ads_mkl/.../hardware.py` and the `gem` kernel library -- these already strip `0` from AMD sweeps and stay correct. - TileLang kernels (`tilelang.jit` / `T.Kernel`) -- `num_stages` there is a different DSL's parameter, not Triton's. - `torch/_inductor/select_algorithm.py` -- a dummy sentinel object, never launched. - `llama4x`/`mslk` `triton_splitk.py` -- already at `num_stages = 1`; only a stale `TODO` comment mentions `0`. - `fbcode/gem/next_gen` -- an ACL'd path, so it lives in the child diff. Phabricator rejects diffs that touch both ACL'd and non-ACL'd paths (https://fburl.com/no-mixed-paths). - `third-party/triton/stable` and `fbcode/triton_mtia/third_party/triton` -- reverted per request; these vendored TLX tutorial copies should be updated via the upstream sync instead. - Generated snapshots under `pyper_models/*/archive/` and `minimal_viable_ai/p4p/*/cloned_files/`, and unrelated `num_stages` concepts (dataswarm, shardmanager, eval pipelines). Note: D114743914 (`[Triton] [Addmm] Fix num_stages range`) may overlap; expect to prune anything already covered by other diffs. Differential Revision: D114751669 --- gdpa_megakernel/src/tlx_gdpa_megakernel.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gdpa_megakernel/src/tlx_gdpa_megakernel.py b/gdpa_megakernel/src/tlx_gdpa_megakernel.py index a367ed2..208df70 100644 --- a/gdpa_megakernel/src/tlx_gdpa_megakernel.py +++ b/gdpa_megakernel/src/tlx_gdpa_megakernel.py @@ -7102,7 +7102,7 @@ def gdpa_backward_tlx( qlen, start_n, BLOCK_M1, BLOCK_N1, WINDOW_SIZE ) - for i in tl.range(0, num_steps, 1, num_stages=0): + for i in tl.range(0, num_steps, 1, num_stages=1): q_buf_id, q_phase = _get_bufidx_phase( accum_cnt_inner + i, NUM_BUFFERS_Q )