From 5ae4ccb0fd97d2a903c07e52dea7949f429d1f42 Mon Sep 17 00:00:00 2001 From: DavidLi <71075143+davidli1515@users.noreply.github.com> Date: Tue, 8 Sep 2026 04:38:01 -0400 Subject: [PATCH] [Bench] Amortize the GDN ReplaySSM decode latency over the flush cycle (#38399) Co-authored-by: Claude Opus 5 --- .../attention/bench_gdn_replayssm_decode.py | 74 ++++++++++++++++--- 1 file changed, 63 insertions(+), 11 deletions(-) diff --git a/benchmark/kernels/attention/bench_gdn_replayssm_decode.py b/benchmark/kernels/attention/bench_gdn_replayssm_decode.py index c27ef70c4..b4659095a 100644 --- a/benchmark/kernels/attention/bench_gdn_replayssm_decode.py +++ b/benchmark/kernels/attention/bench_gdn_replayssm_decode.py @@ -13,6 +13,13 @@ ReplaySSM kernel reads S every step but writes it only 1-in-L steps, plus a small ring append (d:[HV,V], k:[H,K], g:[HV] per step). The amortized state traffic ratio is reported per L. +Because the write happens only 1-in-L steps, the latency has to be amortized the +same way. The kernel branches on the caller's `write_pos` cursor and folds the +ring into the checkpoint only at `write_pos == L-1`, so both phases are timed and +combined: a cycle is L-1 non-flush steps and one flush step, and the flush step +costs roughly 2.3x a non-flush one. Timing a single pinned cursor position would +report whichever phase was pinned as if it were the per-step cost. + Run:: python benchmark/kernels/attention/bench_gdn_replayssm_decode.py @@ -62,7 +69,24 @@ def _state_bytes_per_step(B, HV, K, V, L, dtype): return packed, replay -def _bench_cfg(B, H, HV, K, V, Ls, dtype, device, num_slots=None, warmup=25, rep=100): +def _amortize(t_nonflush, t_flush, L): + """Mean per-step latency over one full L-phase cycle. + + `write_pos` walks 0, 1, ... L-1 and wraps, so a steady-state cycle is L-1 + non-flush steps and one flush step. The non-flush steps are not identical -- + each replays one more ring entry than the last -- so the midpoint phase is + sampled as their representative rather than phase 0, which replays nothing + and is the cheapest step in the cycle. Against a full per-phase sweep this + is accurate to <1%, where using phase 0 understates the mean by ~3%. + """ + if L == 1: + return t_flush + return ((L - 1) * t_nonflush + t_flush) / L + + +def _bench_cfg( + B, H, HV, K, V, Ls, dtype, device, nk=2, num_slots=None, warmup=25, rep=100 +): num_slots = num_slots or B mixed_qkv, a, b, A_log, dt_bias = _make_static(B, H, HV, K, V, dtype, device) scale = K**-0.5 @@ -94,9 +118,8 @@ def _bench_cfg(B, H, HV, K, V, Ls, dtype, device, num_slots=None, warmup=25, rep d_cache = torch.zeros(num_slots, HV, L, V, device=device, dtype=dtype) k_cache = torch.zeros(num_slots, H, L, K, device=device, dtype=dtype) g_cache = torch.zeros(num_slots, HV, L, device=device, dtype=torch.float32) - write_pos = torch.zeros(B, device=device, dtype=torch.int32) + write_pos = torch.empty(B, device=device, dtype=torch.int32) rout = mixed_qkv.new_empty(B, 1, HV, V) - nk = 1 if L == 1 else 2 def run_replay(): fused_recurrent_gdn_replayssm_decode( @@ -117,9 +140,29 @@ def _bench_cfg(B, H, HV, K, V, Ls, dtype, device, num_slots=None, warmup=25, rep nk=nk, ) - t_replay = triton.testing.do_bench(run_replay, warmup=warmup, rep=rep) + def time_at(pos): + write_pos.fill_(pos) + return triton.testing.do_bench(run_replay, warmup=warmup, rep=rep) + + # The kernel branches on write_pos: it appends to the ring on every step + # but folds the ring into the checkpoint only when write_pos == L-1. Both + # phases have to be timed, or the reported cost is the cheap one repeated + # (see `_amortize`). L == 1 is all flush, so it has no non-flush phase. + t_flush = time_at(L - 1) + t_nonflush = time_at((L - 1) // 2) if L > 1 else None + t_replay = _amortize(t_nonflush, t_flush, L) + packed_bytes, replay_bytes = _state_bytes_per_step(B, HV, K, V, L, dtype) - rows.append((L, t_replay, t_packed / t_replay, replay_bytes / packed_bytes)) + rows.append( + ( + L, + t_replay, + t_nonflush, + t_flush, + t_packed / t_replay, + replay_bytes / packed_bytes, + ) + ) return t_packed, rows @@ -132,6 +175,9 @@ def main(): parser.add_argument("--batch-sizes", type=int, nargs="+", default=[1, 16, 64, 256]) parser.add_argument("--ls", type=int, nargs="+", default=[1, 8, 16]) parser.add_argument("--dtype", choices=["bf16", "fp16", "fp32"], default="bf16") + parser.add_argument( + "--nk", type=int, default=2, help="K-chunks the reconstruction walks" + ) args = parser.parse_args() if not torch.cuda.is_available(): @@ -145,20 +191,26 @@ def main(): print( f"GDN ReplaySSM decode microbench HV={args.hv} H={args.h} " - f"K={args.k} V={args.v} dtype={args.dtype}\n" - "per-step latency (ms); speedup = packed/replay; " - "state-traffic = replay/packed (lower is better)" + f"K={args.k} V={args.v} dtype={args.dtype} nk={args.nk}\n" + "replay = mean per-step latency over one L-phase cycle (ms), split into " + "its\nnon-flush and flush steps; speedup = packed/replay; " + "state-traffic = replay/packed\n(lower is better)" ) for B in args.batch_sizes: t_packed, rows = _bench_cfg( - B, args.h, args.hv, args.k, args.v, args.ls, dtype, device + B, args.h, args.hv, args.k, args.v, args.ls, dtype, device, nk=args.nk ) print(f"\nB={B:<4d} packed={t_packed:.4f} ms") - for L, t_replay, speedup, traffic_ratio in rows: + for L, t_replay, t_nonflush, t_flush, speedup, traffic_ratio in rows: + split = ( + f"(flush only)" + if t_nonflush is None + else f"(non-flush {t_nonflush:.4f}, flush {t_flush:.4f})" + ) print( f" L={L:<3d} replay={t_replay:.4f} ms " f"speedup={speedup:5.2f}x " - f"state-traffic={traffic_ratio:5.2f}x" + f"state-traffic={traffic_ratio:5.2f}x {split}" )