"""03_decode_generation.py -- demonstrate prefill + N decode steps. Prefills a small context, then runs N single-token decode steps via `mixer.step(u, cache)`. Verifies the cache persists correctly across steps by comparing NKI-accelerated decode outputs against eager on identical weights. Run: python 03_decode_generation.py """ import time import torch from _loader import load_kernel def main(): print("=" * 60) print("Prefill + Decode: 128 consecutive decode steps") print("=" * 60) mamba3 = load_kernel() torch.manual_seed(0) # SISO mixer -- decode step uses NKI kernel when use_nki_ssd=True mixer_nki = mamba3.NeuronMamba3Mixer( d_model=1024, d_state=128, headdim=64, chunk_size=64, mimo_rank=1, use_nki_ssd=True, ).to("neuron") torch.manual_seed(0) mixer_ref = mamba3.NeuronMamba3Mixer( d_model=1024, d_state=128, headdim=64, chunk_size=64, mimo_rank=1, use_nki_ssd=False, ).to("neuron") mixer_ref.load_state_dict(mixer_nki.state_dict()) torch.manual_seed(100) prefill_len = 64 n_decode = 128 u_prefill = torch.randn(1, prefill_len, 1024, device="neuron") tokens = [torch.randn(1, 1, 1024, device="neuron") for _ in range(n_decode)] print(f"\nprefill_len={prefill_len}, decode_steps={n_decode}") # Prefill on both paths y_pre_nki, cache_nki = mixer_nki(u_prefill) y_pre_ref, cache_ref = mixer_ref(u_prefill) try: torch.neuron.synchronize() except Exception: pass diff = (y_pre_nki.cpu() - y_pre_ref.cpu()).abs().max().item() cos = torch.dot(y_pre_nki.cpu().float().reshape(-1), y_pre_ref.cpu().float().reshape(-1)).item() cos /= (y_pre_nki.norm().item() * y_pre_ref.norm().item() + 1e-12) print(f"prefill: cos_sim={cos:.6f} max_abs={diff:.3e}") # Decode steps print(f"\nrunning {n_decode} decode steps...") n_pass = 0 times = [] for i, tok in enumerate(tokens): # NKI path t0 = time.perf_counter() y_nki, cache_nki = mixer_nki.step(tok, cache_nki) try: torch.neuron.synchronize() except Exception: pass elapsed = (time.perf_counter() - t0) * 1000 times.append(elapsed) # Reference y_ref, cache_ref = mixer_ref.step(tok, cache_ref) try: torch.neuron.synchronize() except Exception: pass y_n = y_nki.cpu().float().reshape(-1) y_r = y_ref.cpu().float().reshape(-1) cos = torch.dot(y_n, y_r).item() / (y_n.norm().item() * y_r.norm().item() + 1e-12) if cos > 0.999: n_pass += 1 warm_times = sorted(times[3:]) if warm_times: med = warm_times[len(warm_times) // 2] p95 = warm_times[int(len(warm_times) * 0.95)] else: med = p95 = 0.0 print(f"\ndecode steps passing cos_sim > 0.999: {n_pass}/{n_decode}") print(f"per-step latency (warm): median={med:.2f} ms, p95={p95:.2f} ms") if n_pass == n_decode: print("\n=== ALL DECODE STEPS PASS ===") else: print(f"\n=== FAILED: {n_decode - n_pass}/{n_decode} steps below tolerance ===") if __name__ == "__main__": main()