Skip to content

Commit 5d141ef

Browse files
committed
docs(triton-route): PtrFwd carry_index 6->0 + other-paths GEMM perf (no regression, NO-GO on §P1 bwd parity)
1 parent b2caa77 commit 5d141ef

1 file changed

Lines changed: 133 additions & 0 deletions

File tree

docs/TRITON-ROUTE-PTRFWD.md

Lines changed: 133 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,133 @@
1+
# Triton-route loop-carried PtrState forwarding (`carry_index` gather elimination)
2+
3+
Status as of commit `15949e5a` (branch `merge/upstream-codegen-reorg`, repo
4+
`/Volumes/external/sources/tilelang`). All GPU execution measured on **gb10**
5+
(NVIDIA GB10, sm_121, aarch64-linux), source `/home/dave/source/tilelang`,
6+
venv `/home/dave/cppmega-venv`. NO local Metal / Apple-GPU dispatch was used.
7+
8+
---
9+
10+
## 1. What changed this run
11+
12+
Three frontend files (parent `40201225` -> HEAD `15949e5a`, commits `54b72972`
13+
+ `15949e5a`):
14+
15+
| File | Change |
16+
|---|---|
17+
| `poc/triton_frontend/op_emitters/control.py` | Forward loop-carried PtrState across `scf.for`: each iter_arg/block-arg inherits the PtrState of its INIT operand; the per-iteration `tt.addptr` advance is folded so in-loop `tt.load`s on the carried block-args emit strided `make_tptr` tile-loads instead of the scalar `carry_index` gather. Plus FIX 1 (canonical per-axis `blockIdx` cache; sorted x,y,z thread_extent). |
18+
| `poc/triton_frontend/op_emitters/memory.py` | FIX 2: `_ptrstate_flat_index` no longer re-multiplies the already-flat per-axis `offsets` by the axis stride (`stride^2` double-count removed); `flat += off; flat += lv*stride`. |
19+
| `poc/triton_frontend/op_mapping.py` | `map_tt_program_id` caches one Var per program-id axis (`ctx.program_id_axis_var`) so a repeated axis read returns the SAME binding -> one `launch_thread` (kills the duplicate `gridDim_0_1`). |
20+
21+
---
22+
23+
## 2. carry_index gather elimination (the headline)
24+
25+
The `carry_index` scalar masked-GATHER (one repeated base address per GEMM
26+
operand tile -> structural permutation) is the defect being killed.
27+
28+
| Kernel (Tri-Dao bwd) | carry_index BEFORE | carry_index AFTER |
29+
|---|---|---|
30+
| dstates, dc, dx, state_db, state_dx, ddAcs_stable | 6 each | **0 each** |
31+
32+
Verified on the regenerated routed PrimFunc `/tmp/pf__chunk_scan_bwd_dstates.json`
33+
(`grep -c carry_index` -> 0). The in-loop tile_loads carry the real
34+
per-program 2D strided `make_tptr` recovered by the aarch64 C++ PtrAnalysis shim.
35+
36+
Additionally FIX 1 dropped the duplicate launch dim: routed dstates kernel
37+
`nparams` went **35 -> 34** (`gridDim_0_1` gone), and FIX 2 made the placement
38+
correct (every grid block writes its native `dprev` slot).
39+
40+
---
41+
42+
## 3. Parity (HONEST — NO-GO on §P1 prod)
43+
44+
* **Placement permutation: ELIMINATED.** With FIX 1 + FIX 2 the grid no longer
45+
collapses to block (0,0,0); the prior-run io.pt config (nc8 nh8 = 64 grid
46+
blocks) wrote 64/64 distinct 4096-tiles, total nz 262144/262144 (vs
47+
4096/131072 = one block before).
48+
* **In-tile value parity: STILL FAILS.** dstates small config MAXDIFF
49+
≈ 2.41e2, `allclose(1e-3)=False`. A SEPARATE in-tile GEMM value error remains
50+
(distinct from the now-fixed placement permutation and the now-zero gather).
51+
* **§P1 production (S4096, c64, g8, H112, P64, N64, bs1, grid (1,64,112)):
52+
routed kernel SEGFAULTS at CUDA launch** (measured this run with the correct
53+
34-arg call via `/tmp/time_routed34.py` and `/tmp/parity_prod_dstates34.py`).
54+
This is the residual NO-GO blocker; native ref = 1.14 ms/kernel.
55+
56+
Per RULE #1 no routed ms is reported as a PASS for dstates: the kernel does not
57+
yet produce a correct result at the prod config, so a timing number there would
58+
be a degenerate/garbage-output number. Reported as NO-GO, not fabricated.
59+
60+
---
61+
62+
## 4. Routed ms (parity-PASS kernels only)
63+
64+
**None.** No bwd kernel reaches `allclose(1e-3)` parity at the prod config this
65+
run (dstates in-tile values wrong + §P1 segfault; the other 6 unverified
66+
downstream of the same in-tile GEMM defect). Therefore, per the rule "routed ms
67+
only for parity-PASS", there is no routed-ms table to report. Native baseline
68+
for reference: `_chunk_scan_bwd_dstates` = **1.14 ms/kernel** @ §P1.
69+
70+
Residual bottleneck quantified: §P1 routed launch segfault in the multi-K-trip
71+
in-tile GEMM (cs=64, BK=32 -> 2 K-trips), plus the redundant whole-block
72+
prologue (168 loops) noted previously.
73+
74+
---
75+
76+
## 5. Other-paths GEMM perf — BEFORE vs AFTER (MANDATORY, delivered)
77+
78+
UNCONDITIONAL user ask: confirm this run's frontend edits did NOT regress the
79+
OTHER frontend GEMM paths. Measured on gb10 by capturing TTIR (subprocess),
80+
routing via `from_ttir(target="cuda")`, in fresh isolated processes.
81+
82+
* **BEFORE** = the 3 frontend files checked out at parent `40201225`
83+
(`git checkout 40201225 -- control.py memory.py op_mapping.py`, pyc purged).
84+
* **AFTER** = HEAD `15949e5a`.
85+
* `route_ms` = min over 8 reps, then min over 3 isolated-process trials (the
86+
routing path is exactly where the edited code executes; `route_ms` is the
87+
honest frontend-cost metric for these paths).
88+
89+
| Kernel | scf.for | carry_index B / A | route_ms BEFORE (min) | route_ms AFTER (min) | Verdict |
90+
|---|---|---|---|---|---|
91+
| `fla_dot_exp2` (dot + exp2) | 0 | 0 / 0 | 73.86 | 75.68 | NO REGRESSION (~+2.5%, machine noise; route 73.8/74.3/99.7 B vs 75.7/76.2/97.6 A) |
92+
| `matmul` (K-loop GEMM) | 1 | 14 / 14 | 112.30 | 115.24 | NO REGRESSION (~+2.6%, noise; route 116/112/113 B vs 119/115/115 A) |
93+
| `dot_reduce_atomic` (dot+rowsum+atomic) | 0 | n/a | EmitError `arith.constant true` | StructuralEqual `(32,)` (later stage) | NO REGRESSION — AFTER fails LATER (the `arith.constant true` emit gap was fixed; now reaches a shape check) |
94+
| `dot_reduce_atomic_trans_b` | 0 | n/a | EmitError `arith.constant true` | StructuralEqual `(32,)` (later stage) | NO REGRESSION — same as above |
95+
96+
GPU-exec of these 4 standalone kernels is blocked by PRE-EXISTING tilelang
97+
codegen issues present identically BEFORE and AFTER (matmul = CUDA segfault at
98+
launch on the 32-tile config; fla_dot_exp2 = `m_warp*n_warp==num_warps` check
99+
`1*1 != 4`). These are not introduced by this run's edits (verified by running
100+
the identical BEFORE/AFTER call). Hence exec_ms is not reportable for these
101+
paths in either state; the comparable frontend metric is `route_ms` + the
102+
`carry_index` count, both of which show no regression.
103+
104+
**Honest attribution note on matmul:** the loop-carried PtrState forwarding did
105+
NOT eliminate the `carry_index` gather on the simple `matmul` K-loop
106+
(`carry_index=14` BEFORE and AFTER). The forwarding implemented this run is
107+
scoped to the Tri-Dao bwd carried-ptr pattern; the generic
108+
`a_ptrs += BK*stride` / `b_ptrs += BK*stride` carry in `matmul` still falls to
109+
the scalar gather. This is a real, disclosed limitation — not a regression, and
110+
not claimed as fixed.
111+
112+
---
113+
114+
## 6. Full-7 routed-bwd sum / vs native / vs path_c
115+
116+
Not computable as a parity-PASS aggregate: the routed bwd is NO-GO this run
117+
(§P1 segfault + in-tile value error). path_c v1 reference chain = 905 ms;
118+
native full-bwd reference per-kernel = 1.14 ms/kernel for dstates. No fabricated
119+
routed sum is reported.
120+
121+
---
122+
123+
## 7. GO / NO-GO
124+
125+
**NO-GO for routed Tri-Dao bwd parity.** Achieved this run: carry_index gather
126+
6 -> 0 on all 6 routable kernels; placement permutation eliminated; duplicate
127+
launch-dim fixed (nparams 35->34); `arith.constant true` emit gap fixed.
128+
Remaining: a SEPARATE in-tile GEMM value error and a §P1 prod launch segfault —
129+
the routed bwd is not yet numerically correct or runnable at prod scale.
130+
131+
**GO on the mandatory other-paths regression check:** fla_dot_exp2 and matmul
132+
route with NO regression (route_ms within machine noise, carry_index unchanged);
133+
dot_reduce paths fail later (improved), not earlier. No GEMM path regressed.

0 commit comments

Comments
 (0)