from __future__ import annotations """Persistent CUDA-Graph greedy decoder for the all-recurrent GDN24 runtime. The core idea is specific to a fully recurrent model: the cache has fixed tensor addresses and token position is implicit in state evolution. We can therefore capture multiple autoregressive steps into one CUDA Graph and let the graph feed its own argmax token into the next step. """ from dataclasses import dataclass import time from typing import Any import torch @dataclass class _LayerSnapshot: conv: dict[int, torch.Tensor] recurrent: dict[int, torch.Tensor] has_previous_state: dict[int, bool] conv_initialized: dict[int, bool] recurrent_initialized: dict[int, bool] @dataclass class CacheSnapshot: seen_tokens: int layers: list[_LayerSnapshot] def snapshot_cache(cache) -> CacheSnapshot: layers: list[_LayerSnapshot] = [] for layer in cache.layers: conv = {int(i): t.detach().clone() for i, t in layer.conv_states.items() if t is not None} recurrent = { int(i): t.detach().clone() for i, t in layer.recurrent_states.items() if t is not None } layers.append( _LayerSnapshot( conv=conv, recurrent=recurrent, has_previous_state=dict(layer.has_previous_state), conv_initialized=dict(layer.is_conv_states_initialized), recurrent_initialized=dict(layer.is_recurrent_states_initialized), ) ) return CacheSnapshot(seen_tokens=int(cache.seen_tokens), layers=layers) def restore_cache_(cache, snap: CacheSnapshot) -> None: if len(cache.layers) != len(snap.layers): raise ValueError("cache topology changed while restoring snapshot") for layer, state in zip(cache.layers, snap.layers): for i, src in state.conv.items(): dst = layer.conv_states[i] if dst is None or dst.shape != src.shape: raise ValueError(f"conv state storage changed at state {i}") dst.copy_(src) for i, src in state.recurrent.items(): dst = layer.recurrent_states[i] if dst is None or dst.shape != src.shape: raise ValueError(f"recurrent state storage changed at state {i}") dst.copy_(src) layer.has_previous_state.update(state.has_previous_state) layer.is_conv_states_initialized.update(state.conv_initialized) layer.is_recurrent_states_initialized.update(state.recurrent_initialized) cache.seen_tokens = int(snap.seen_tokens) class SuperTurboGraphDecoder: """Persistent, self-feeding greedy CUDA Graph decoder. A single replay advances `block_size` recurrent decode steps. The graph owns one persistent recurrent cache, so it can be reused across prompts: `reset()` preserves all CUDA addresses, prefill writes the new prompt state, then graph replay continues autoregressively from that state. """ def __init__( self, model, *, block_size: int = 8, warmup_steps: int = 2, graph_pool=None, ): if not torch.cuda.is_available(): raise RuntimeError("MAX TURBO CUDA Graph decoding requires CUDA") if block_size < 1: raise ValueError("block_size must be >= 1") self.model = model.eval() self.block_size = int(block_size) self.warmup_steps = int(warmup_steps) self.graph_pool = graph_pool self.device = model.get_input_embeddings().weight.device if self.device.type != "cuda": raise RuntimeError(f"model must be on CUDA, got {self.device}") self.cache = self.model.make_recurrent_cache() self.static_token = torch.zeros((1, 1), dtype=torch.long, device=self.device) self.output_tokens = torch.empty((1, self.block_size), dtype=torch.long, device=self.device) self.graph: torch.cuda.CUDAGraph | None = None self.capture_seconds: float | None = None self.capture_error: str | None = None self.capture_allocated_delta_mib: float = 0.0 self.capture_reserved_delta_mib: float = 0.0 self._captured = False self._capture_attempted = False @torch.inference_mode() def _greedy_step(self) -> torch.Tensor: # Private clock flag is consumed by Qwen35GDN24Model and never reaches GDN kernels. return self.model.greedy_step( self.static_token, self.cache, advance_cache_clock=False, ) @torch.inference_mode() def _unrolled_block(self) -> None: for i in range(self.block_size): next_token = self._greedy_step() self.output_tokens[:, i : i + 1].copy_(next_token) self.static_token.copy_(next_token) @torch.inference_mode() def capture(self) -> bool: if self._captured: return True if self._capture_attempted: return False self._capture_attempted = True try: # Initialize every native LinearAttentionLayer cache with stable addresses. self.cache.reset() dummy = torch.zeros((1, 1), dtype=torch.long, device=self.device) first = self.model( input_ids=dummy, past_key_values=self.cache, use_cache=True, logits_to_keep=1, ).logits[:, -1, :].argmax(dim=-1, keepdim=True) self.static_token.copy_(first) torch.cuda.synchronize() baseline = snapshot_cache(self.cache) baseline_token = self.static_token.detach().clone() # Warm single-token recurrent kernels and Triton fusions before capture. for _ in range(max(self.warmup_steps, 0)): self._unrolled_block() torch.cuda.synchronize() restore_cache_(self.cache, baseline) self.static_token.copy_(baseline_token) graph = torch.cuda.CUDAGraph() torch.cuda.synchronize() alloc_before = torch.cuda.memory_allocated(self.device) reserve_before = torch.cuda.memory_reserved(self.device) t0 = time.perf_counter() graph_kwargs = {} if self.graph_pool is None else {"pool": self.graph_pool} with torch.cuda.graph(graph, **graph_kwargs): self._unrolled_block() torch.cuda.synchronize() self.capture_seconds = time.perf_counter() - t0 self.capture_allocated_delta_mib = max(0, torch.cuda.memory_allocated(self.device) - alloc_before) / 2**20 self.capture_reserved_delta_mib = max(0, torch.cuda.memory_reserved(self.device) - reserve_before) / 2**20 # Capture executes once. Restore the exact pre-capture numerical state # without replacing any tensor objects/addresses recorded by the graph. restore_cache_(self.cache, baseline) self.static_token.copy_(baseline_token) if hasattr(graph, "instantiate"): try: graph.instantiate() except Exception: pass self.graph = graph self._captured = True self.cache.reset() return True except Exception as exc: self.capture_error = f"{type(exc).__name__}: {exc}" self.graph = None self._captured = False try: self.cache.reset() except Exception: pass return False @torch.inference_mode() def reset(self) -> None: self.cache.reset() @torch.inference_mode() def prefill(self, input_ids: torch.Tensor) -> tuple[torch.Tensor, Any]: """Prefill into the persistent graph cache and return the first next token.""" self.cache.reset() out = self.model( input_ids=input_ids, past_key_values=self.cache, use_cache=True, logits_to_keep=1, ) token = torch.argmax(out.logits[:, -1, :], dim=-1, keepdim=True) return token, out @torch.inference_mode() def decode_forwards(self, first_token: torch.Tensor, steps: int) -> torch.Tensor: """Run `steps` recurrent forward passes, matching the benchmark's decode metric. The returned tensor contains the last generated token. For the fast path, `steps` should be divisible by `block_size`; a short eager tail handles any remainder. """ steps = int(steps) if steps < 0: raise ValueError("steps must be >= 0") if steps == 0: return first_token if not self._captured: return self._decode_eager(first_token, steps) self.static_token.copy_(first_token) full_blocks, remainder = divmod(steps, self.block_size) for _ in range(full_blocks): self.graph.replay() # The graph deliberately does not update the Python token clock. self.cache.advance(full_blocks * self.block_size) if remainder: token = self.static_token for _ in range(remainder): out = self.model( input_ids=token, past_key_values=self.cache, use_cache=True, logits_to_keep=1, ) token = torch.argmax(out.logits[:, -1, :], dim=-1, keepdim=True) self.static_token.copy_(token) return self.static_token @torch.inference_mode() def _decode_eager(self, first_token: torch.Tensor, steps: int) -> torch.Tensor: token = first_token for _ in range(steps): out = self.model( input_ids=token, past_key_values=self.cache, use_cache=True, logits_to_keep=1, ) token = torch.argmax(out.logits[:, -1, :], dim=-1, keepdim=True) return token @torch.inference_mode() def decode_tokens(self, first_token: torch.Tensor, steps: int) -> torch.Tensor: """Return every greedy token produced by ``steps`` recurrent forwards. Unlike :meth:`decode_forwards`, this materializes the token sequence and is intended for user generation, quality checks, and UNI MAX verification. The persistent recurrent cache remains graph-address-stable. """ steps = int(steps) if steps < 0: raise ValueError("steps must be >= 0") if steps == 0: return torch.empty((first_token.shape[0], 0), dtype=torch.long, device=first_token.device) self.static_token.copy_(first_token) pieces: list[torch.Tensor] = [] if self._captured: full_blocks, remainder = divmod(steps, self.block_size) for _ in range(full_blocks): self.graph.replay() pieces.append(self.output_tokens.detach().clone()) if full_blocks: self.cache.advance(full_blocks * self.block_size) token = self.static_token else: remainder = steps token = first_token for _ in range(remainder): token = self.model.greedy_step(token, self.cache) pieces.append(token.detach().clone()) if pieces: self.static_token.copy_(pieces[-1][:, -1:]) return torch.cat(pieces, dim=1) @torch.inference_mode() def replay_block_tokens(self, first_token: torch.Tensor) -> torch.Tensor: """Replay exactly one captured block and return its generated tokens.""" if not self._captured: if not self.capture(): raise RuntimeError(f"CUDA Graph capture failed: {self.capture_error}") self.static_token.copy_(first_token) self.graph.replay() self.cache.advance(self.block_size) return self.output_tokens.detach().clone() @torch.inference_mode() def validate_against_eager(self, first_token: torch.Tensor) -> bool: """Bit-exact greedy-token check for one captured block.""" if not self._captured and not self.capture(): raise RuntimeError(f"CUDA Graph capture failed: {self.capture_error}") snap = snapshot_cache(self.cache) start = first_token.detach().clone() try: token = start eager_tokens = [] for _ in range(self.block_size): token = self.model.greedy_step( token, self.cache, advance_cache_clock=False ) eager_tokens.append(token.detach().clone()) eager = torch.cat(eager_tokens, dim=1) restore_cache_(self.cache, snap) self.static_token.copy_(start) self.graph.replay() torch.cuda.synchronize() graphed = self.output_tokens.detach().clone() if not torch.equal(eager, graphed): raise RuntimeError( f"CUDA Graph greedy mismatch: eager={eager.tolist()} graph={graphed.tolist()}" ) return True finally: restore_cache_(self.cache, snap) self.static_token.copy_(start) @property def ready(self) -> bool: return self._captured @dataclass(frozen=True) class GraphTuneResult: block_size: int decode_tok_s: float capture_seconds: float graph_allocated_mib: float = 0.0 graph_reserved_mib: float = 0.0 @torch.inference_mode() def autotune_graph_decoder( model, *, candidates: tuple[int, ...] = (4, 8, 16), warmup_replays: int = 2, timed_replays: int = 6, ) -> tuple[SuperTurboGraphDecoder, list[GraphTuneResult]]: """Capture and benchmark several self-feeding graph block sizes. The search measures steady-state graph replay only; graph capture time is reported separately and excluded. Every candidate is numerically checked against eager greedy decoding before it can win. """ if not candidates: raise ValueError("at least one graph block candidate is required") results: list[GraphTuneResult] = [] best: SuperTurboGraphDecoder | None = None best_speed = -1.0 failures: list[str] = [] shared_pool = torch.cuda.graph_pool_handle() if hasattr(torch.cuda, "graph_pool_handle") else None for raw_block in candidates: block = int(raw_block) if block < 1: failures.append(f"x{block}: invalid block size") continue engine = SuperTurboGraphDecoder(model, block_size=block, warmup_steps=2, graph_pool=shared_pool) if not engine.capture(): failures.append(f"x{block}: {engine.capture_error}") continue # Start from a valid persistent recurrent state, then verify exact greedy # token agreement before any timing result is accepted. seed = torch.zeros((1, 8), dtype=torch.long, device=engine.device) first, _ = engine.prefill(seed) engine.validate_against_eager(first) engine.reset() first, _ = engine.prefill(seed) engine.static_token.copy_(first) for _ in range(max(int(warmup_replays), 0)): engine.graph.replay() torch.cuda.synchronize() t0 = time.perf_counter() for _ in range(max(int(timed_replays), 1)): engine.graph.replay() torch.cuda.synchronize() dt = time.perf_counter() - t0 forwards = block * max(int(timed_replays), 1) speed = forwards / dt result = GraphTuneResult( block_size=block, decode_tok_s=float(speed), capture_seconds=float(engine.capture_seconds or 0.0), graph_allocated_mib=float(engine.capture_allocated_delta_mib), graph_reserved_mib=float(engine.capture_reserved_delta_mib), ) results.append(result) if speed > best_speed: best_speed = speed previous_best = best best = engine if previous_best is not None: del previous_best torch.cuda.empty_cache() else: # Drop non-winning graph pools as early as possible. del engine torch.cuda.empty_cache() if best is None: reason = "; ".join(failures) if failures else "no valid candidates" raise RuntimeError(f"MAX TURBO graph autotune failed: {reason}") best.reset() return best, results # MAX TURBO public name; keep the v3 class name as a compatibility alias. MaxTurboGraphDecoder = SuperTurboGraphDecoder