From 70ed0a70e006f53ef1efebe33cc8f862bcede744 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Mon, 3 Aug 2026 23:47:39 +0000 Subject: [PATCH] llm: minimize recurrent prefill integration --- test/unit/test_attention.py | 5 +- tinygrad/llm/cli.py | 5 +- tinygrad/llm/kernels/__init__.py | 4 + tinygrad/llm/model.py | 196 +++++++++++-------------------- 4 files changed, 79 insertions(+), 131 deletions(-) diff --git a/test/unit/test_attention.py b/test/unit/test_attention.py index be4497c585..250f6c9197 100644 --- a/test/unit/test_attention.py +++ b/test/unit/test_attention.py @@ -78,9 +78,8 @@ class TestGatedDeltaNetBlock(unittest.TestCase): def _run_attention(self, block:GatedDeltaNetBlock, x:Tensor, start_pos:int): x_norm = block.attn_norm(x) block._init_state(x_norm) - out, conv_state = block._attention(x_norm, start_pos) - Tensor.realize(out, block.conv_state.assign(conv_state)) - return out.numpy() + ret = block._attention(x_norm, start_pos) + return (block._update_state(*ret) if isinstance(ret, tuple) else ret).realize().numpy() def _cache_views(self, block:GatedDeltaNetBlock) -> tuple[np.ndarray, np.ndarray]: if hasattr(block, 'conv_state'): diff --git a/tinygrad/llm/cli.py b/tinygrad/llm/cli.py index 666b89912f..d0865e26ba 100644 --- a/tinygrad/llm/cli.py +++ b/tinygrad/llm/cli.py @@ -3,7 +3,7 @@ import sys, argparse, codecs, itertools, typing, re, unicodedata, json, time from typing import TYPE_CHECKING from tinygrad import nn from tinygrad.uop.ops import UOp, Ops -from tinygrad.helpers import partition, Timing, GlobalCounters, fetch, profile_marker, getenv +from tinygrad.helpers import partition, DEBUG, Timing, GlobalCounters, Context, fetch, profile_marker, getenv from tinygrad.llm.model import Transformer if TYPE_CHECKING: import jinja2 @@ -163,7 +163,8 @@ def main(): except ImportError: print("warning: jinja2 is not installed, the model's chat template is disabled") # warmup the JIT - if args.warmup or args.serve: model.warmup() + if args.warmup or args.serve: + with Context(DEBUG=max(DEBUG.value, 1)): model.warmup() # start server if args.serve: LLMServer(('', args.serve), model, model_name, tok, template).serve_forever() diff --git a/tinygrad/llm/kernels/__init__.py b/tinygrad/llm/kernels/__init__.py index 8a2c905b57..3e8f57fa28 100644 --- a/tinygrad/llm/kernels/__init__.py +++ b/tinygrad/llm/kernels/__init__.py @@ -44,6 +44,10 @@ def _prepare_quantized_weights(model:Any, state_dict:dict[str, Tensor]) -> list[ return packed def load_state_dict(model:Any, state_dict:dict[str, Tensor]): + for key in nn.state.get_state_dict(model): + if key.endswith(".ssm_beta_alpha.weight") and key not in state_dict: + prefix = key.removesuffix("beta_alpha.weight") + state_dict[key] = state_dict.pop(prefix+"beta.weight").cat(state_dict.pop(prefix+"alpha.weight"), dim=0).contiguous() packed = _prepare_quantized_weights(model, state_dict) nn.state.load_state_dict(model, state_dict, verbose=False, consume=True, realize=False) if packed: Tensor.realize(*(offset for _,offset in packed)) diff --git a/tinygrad/llm/model.py b/tinygrad/llm/model.py index 3aee6f532d..7dad162056 100644 --- a/tinygrad/llm/model.py +++ b/tinygrad/llm/model.py @@ -4,7 +4,6 @@ from dataclasses import dataclass, replace from tinygrad import Tensor, nn, UOp, TinyJit, getenv, function, Context from tinygrad.llm.kernels import Linear, cached_attention, gated_delta_prefill, load_state_dict, make_attention_cache from tinygrad.llm.gguf import gguf_load -from tinygrad.helpers import DEBUG from tinygrad.uop.ops import resolve @functools.cache @@ -128,16 +127,20 @@ class FFNBlock: # return writes that reset this block's state after a cache mismatch def _state_reset_ops(self) -> list[Tensor]: return [] def _init_state(self, x:Tensor): raise NotImplementedError - def _attention(self, x:Tensor, start_pos:int|UOp) -> Tensor: raise NotImplementedError + def _attention(self, x:Tensor, start_pos:int|UOp, **kwargs) -> Tensor|tuple[Tensor, ...]: raise NotImplementedError + def _update_state(self, out:Tensor, *state:Tensor) -> Tensor: return out - def __call__(self, x: Tensor, start_pos: int|UOp): + def __call__(self, x: Tensor, start_pos: int|UOp, **kwargs): self._init_state(x) # we pass in the weights implicitly so we unpack the GGUF on the fly @function(precompile=True, allow_implicit=True) def _run(x:Tensor, start_pos:int|UOp): - h = x + self._attention(self.attn_norm(x), start_pos) - return (h + self._feed_forward(self.ffn_norm(h))).contiguous() - return _run(x, start_pos) + attn = self._attention(self.attn_norm(x), start_pos, **kwargs) + h = x + (attn[0] if isinstance(attn, tuple) else attn) + out = (h + self._feed_forward(self.ffn_norm(h))).contiguous() + return (out, *attn[1:]) if isinstance(attn, tuple) else out + ret = _run(x, start_pos) + return self._update_state(*ret) if isinstance(ret, tuple) else ret class TransformerBlock(FFNBlock): def __init__(self, config:TransformerConfig): @@ -153,7 +156,7 @@ class TransformerBlock(FFNBlock): self.attn_output = Linear(config.head_dim * config.n_heads, config.dim, bias=False) if config.qk_norm: self.attn_q_norm, self.attn_k_norm = nn.RMSNorm(config.qk_norm, config.norm_eps), nn.RMSNorm(config.qk_norm, config.norm_eps) - def _attention(self, x:Tensor, start_pos:int|UOp) -> Tensor: + def _attention(self, x:Tensor, start_pos:int|UOp, **kwargs) -> Tensor: q, k, v = self.attn_q(x), self.attn_k(x), self.attn_v(x) if self.config.qk_norm and self.config.qk_norm != self.config.head_dim: q, k = self.attn_q_norm(q), self.attn_k_norm(k) @@ -197,7 +200,7 @@ class MLATransformerBlock(FFNBlock): self.attn_v_b = {"weight": Tensor.zeros(config.n_heads, config.v_head_dim, config.kv_lora_rank)} self.attn_output = Linear(config.n_heads * config.v_head_dim, config.dim, bias=False) - def _attention(self, x:Tensor, start_pos:int|UOp) -> Tensor: + def _attention(self, x:Tensor, start_pos:int|UOp, **kwargs) -> Tensor: B, T, _ = x.shape q_nope_head_dim = self.config.head_dim - self.config.rope_dim q_proj = self.attn_q_b(self.attn_q_a_norm(self.attn_q_a(x))) if self.config.q_lora_rank > 0 else self.attn_q(x) @@ -248,78 +251,47 @@ class GatedDeltaNetBlock(FFNBlock): self.ssm_a = Tensor.zeros(self.num_v_heads, 1) if ssm.kda else Tensor.zeros(self.num_v_heads) self.ssm_norm, self.ssm_out = nn.RMSNorm(self.head_v_dim, config.norm_eps), Linear(ssm.inner_size, config.dim, bias=False) - def __call__(self, x:Tensor, start_pos:int|UOp, valid_len:int|UOp|None=None): - if not hasattr(self, 'attn_gate'): return super().__call__(x, start_pos) - self._init_state(x) - @function(precompile=True, allow_implicit=True) - def _run(x:Tensor, start_pos:int|UOp, valid_len:int|UOp|None): - attn, conv_state = self._attention(self.attn_norm(x), start_pos, valid_len) - h = x + attn - return (h + self._feed_forward(self.ffn_norm(h))).contiguous(), conv_state - out, conv_state = _run(x, start_pos, valid_len) - return Tensor(out.uop.after(self.conv_state.uop.after(self.conv_state.uop.store(conv_state.uop)))) - - def _attention(self, x:Tensor, start_pos:int|UOp, valid_len:int|UOp|None=None): + def _attention(self, x:Tensor, start_pos:int|UOp, valid_len:int|UOp|None=None, **kwargs): B, T, _ = x.shape - if hasattr(self, "attn_gate"): return self._attention_qwen(x, valid_len) - assert T == 1, "GatedDeltaNetBlock currently only supports T=1" - - # input processing - x = x.half() - out_gate = self.ssm_g_b(self.ssm_g_a(x)) - out_gate = out_gate.reshape(B, 1, self.num_v_heads, self.head_v_dim) - beta = self.ssm_beta(x).sigmoid().reshape(B, self.num_v_heads, 1, 1) - alpha = self.ssm_f_b(self.ssm_f_a(x)) - alpha = ((alpha.float() + self.ssm_dt["bias"]).softplus().reshape(B, self.num_v_heads, -1) * - self.ssm_a.reshape(1, self.num_v_heads, -1)).exp().unsqueeze(-2) - - # qkv conv + is_kda, x = hasattr(self, "ssm_g_a"), x.half() + out_gate = (self.ssm_g_b(self.ssm_g_a(x)) if is_kda else self.attn_gate(x)).reshape(B, T, self.num_v_heads, self.head_v_dim) + beta, alpha = (self.ssm_beta(x), self.ssm_f_b(self.ssm_f_a(x))) if is_kda else \ + self.ssm_beta_alpha(x).split(self.num_v_heads, dim=-1) conv_window = self.conv_state.cat(self.attn_qkv(x), dim=1) - conv_out = (conv_window * self.ssm_conv1d["weight"].T.unsqueeze(0)).sum(1).silu() + conv_out = ((conv_window * self.ssm_conv1d["weight"].T.unsqueeze(0)).sum(1) if is_kda else functools.reduce(lambda a,b: a+b, + (conv_window[:, i:i+T] * self.ssm_conv1d["weight"][:, i] for i in range(self.ssm_conv_kernel)))).silu() q, k, v = conv_out.split([self.q_dim, self.q_dim, self.conv_channels - 2*self.q_dim], dim=-1) - q = q.reshape(B, self.num_k_heads, self.head_k_dim).normalize(dim=-1).repeat(1, self.num_v_heads//self.num_k_heads, 1) - k = k.reshape(B, self.num_k_heads, self.head_k_dim).normalize(dim=-1).repeat(1, self.num_v_heads//self.num_k_heads, 1) - v = v.reshape(B, self.num_v_heads, self.head_v_dim) - q, k, v = q.mul(self.head_k_dim**-0.5).unsqueeze(-1), k.unsqueeze(-1), v.unsqueeze(-1) - - # recurrent - recurrent_state = self.recurrent_state * alpha - recurrent_state = recurrent_state + ((v - recurrent_state@k) * beta)@k.transpose(-1, -2) - - # store the updated state - conv_state_store = self.conv_state.uop.store(conv_window[:, 1:, :].cast(self.conv_state.dtype).uop) - recurrent_state_store = self.recurrent_state.uop.store(recurrent_state.cast(self.recurrent_state.dtype).uop) - recurrent_state = Tensor(self.recurrent_state.uop.after(recurrent_state_store, conv_state_store)) - - # output - core_attn_out = self.ssm_norm((recurrent_state@q).squeeze(-1).reshape(B, 1, self.num_v_heads, self.head_v_dim)) - return self.ssm_out((core_attn_out * out_gate.sigmoid()).reshape(B, 1, -1).cast(x.dtype)) - - def _attention_qwen(self, x:Tensor, valid_len:int|UOp|None): - B, T, _ = x.shape - x = x.half() - out_gate = self.attn_gate(x).reshape(B, T, self.num_v_heads, self.head_v_dim) - beta, alpha = self.ssm_beta_alpha(x).split(self.num_v_heads, dim=-1) - conv_window = self.conv_state.cat(self.attn_qkv(x), dim=1) - conv_out = functools.reduce(lambda a,b: a+b, - (conv_window[:, i:i+T] * self.ssm_conv1d["weight"][:, i] for i in range(self.ssm_conv_kernel))).silu() - q, k, v = conv_out.split([self.q_dim, self.q_dim, self.conv_channels - 2*self.q_dim], dim=-1) - q = q.reshape(B, T, self.num_k_heads, self.head_k_dim).normalize(dim=-1, eps=1e-6).repeat( - 1, 1, self.num_v_heads//self.num_k_heads, 1) - k = k.reshape(B, T, self.num_k_heads, self.head_k_dim).normalize(dim=-1, eps=1e-6).repeat( - 1, 1, self.num_v_heads//self.num_k_heads, 1) + q, k = (z.reshape(B, T, self.num_k_heads, self.head_k_dim).normalize(dim=-1, eps=1e-12 if is_kda else 1e-6).repeat( + 1, 1, self.num_v_heads//self.num_k_heads, 1) for z in (q, k)) v = v.reshape(B, T, self.num_v_heads, self.head_v_dim) - beta, log_alpha = beta.sigmoid().reshape(B, T, self.num_v_heads), \ - ((alpha.float() + self.ssm_dt["bias"]).softplus() * self.ssm_a).reshape(B, T, self.num_v_heads) - if valid_len is not None: - active = (Tensor.arange(T).to(x.device) < Tensor(valid_len, device=x.device)).reshape(1, T, 1) - beta, log_alpha = beta * active, log_alpha * active - q, k, v, beta, log_alpha = [z.transpose(1, 2).float() for z in (q, k, v, beta, log_alpha)] - core = gated_delta_prefill(q * self.head_k_dim**-0.5, k, v, beta, log_alpha.exp(), self.recurrent_state) - out = self.ssm_out((self.ssm_norm(core.transpose(1, 2)) * out_gate.silu()).reshape(B, T, -1).cast(x.dtype)).contiguous() - state_pos = T if valid_len is None else valid_len + state_pos:int|UOp + if is_kda: + assert T == 1, "channel-wise gated delta prefill is not supported" + beta = beta.sigmoid().reshape(B, self.num_v_heads, 1, 1) + alpha = ((alpha.float() + self.ssm_dt["bias"]).softplus().reshape(B, self.num_v_heads, -1) * + self.ssm_a.reshape(1, self.num_v_heads, -1)).exp().unsqueeze(-2) + q, k, v = q[:, 0].mul(self.head_k_dim**-0.5).unsqueeze(-1), k[:, 0].unsqueeze(-1), v[:, 0].unsqueeze(-1) + recurrent_state = self.recurrent_state * alpha + recurrent_state = recurrent_state + ((v - recurrent_state@k) * beta)@k.transpose(-1, -2) + recurrent_state = Tensor(self.recurrent_state.uop.after( + self.recurrent_state.uop.store(recurrent_state.cast(self.recurrent_state.dtype).uop))) + core, gate, state_pos = (recurrent_state@q).squeeze(-1).reshape(B, 1, self.num_v_heads, self.head_v_dim), out_gate.sigmoid(), 1 + else: + beta, log_alpha = beta.sigmoid().reshape(B, T, self.num_v_heads), \ + ((alpha.float() + self.ssm_dt["bias"]).softplus() * self.ssm_a).reshape(B, T, self.num_v_heads) + if valid_len is not None: + active = (Tensor.arange(T).to(x.device) < Tensor(valid_len, device=x.device)).reshape(1, T, 1) + beta, log_alpha = beta * active, log_alpha * active + q, k, v, beta, log_alpha = [z.transpose(1, 2).float() for z in (q, k, v, beta, log_alpha)] + core = gated_delta_prefill(q * self.head_k_dim**-0.5, k, v, beta, log_alpha.exp(), self.recurrent_state).transpose(1, 2) + gate, state_pos = out_gate.silu(), T if valid_len is None else valid_len + out = self.ssm_out((self.ssm_norm(core) * gate).reshape(B, T, -1).cast(x.dtype)).contiguous() conv_state = conv_window[:, state_pos:state_pos+self.ssm_conv_kernel-1].cast(self.conv_state.dtype).contiguous() - return out, conv_state + return Tensor(out.uop.after(self.conv_state.uop.after(self.conv_state.uop.store(conv_state.uop)))) if is_kda else (out, conv_state) + + def _update_state(self, out:Tensor, *state:Tensor) -> Tensor: + conv_state, = state + return Tensor(out.uop.after(self.conv_state.uop.after(self.conv_state.uop.store(conv_state.uop)))) # recurrent state can't be partially reused after divergence, force a full rebuild def _state_reset_ops(self): @@ -345,18 +317,19 @@ class Transformer: self.output = Linear(config.dim, config.vocab_size, bias=False) self.max_context = config.max_context self.has_recurrent_block = any(isinstance(b, GatedDeltaNetBlock) for b in self.blk) + self.has_recurrent_prefill = config.ssm is not None and not config.ssm.kda self._cached_tokens: list[int] = [] # we specialize the JIT for prefill and rollout self.prefill_jit = TinyJit(self.forward) - self.rollout_jit = TinyJit(self.forward_recurrent_decode if self.has_recurrent_block else self.forward) + self.rollout_jit = TinyJit(self.forward_recurrent_decode if self.has_recurrent_prefill else self.forward) def forward(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor, valid_len:int|UOp|None=None) -> Tensor: x = self.token_embd(tokens).float() # (B, T, D) - for block in self.blk: x = block(x, start_pos, valid_len) if isinstance(block, GatedDeltaNetBlock) else block(x, start_pos) + for block in self.blk: x = block(x, start_pos, valid_len=valid_len) if isinstance(block, GatedDeltaNetBlock) else block(x, start_pos) last = x[:, tokens.shape[1]-1:tokens.shape[1]] if valid_len is None else x[:, valid_len-1:valid_len] logits = self.output(self.output_norm(last))[:, -1, :] # Gumbel-max trick: argmax(logits/temp - log(-log(uniform))) is equivalent to sampling from softmax(logits/temp) - if self.has_recurrent_block: return logits.argmax(-1, keepdim=True) + if self.has_recurrent_prefill: return logits.argmax(-1, keepdim=True) return (logits / temperature.maximum(1e-12) - (Tensor.rand_like(logits).maximum(1e-12).log().neg()).log()).argmax(-1, keepdim=True) def forward_recurrent_decode(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor) -> Tensor: @@ -388,9 +361,6 @@ class Transformer: if arch in ('qwen35', 'qwen35moe'): ssm = SSMConfig(**{k: kv[f'{arch}.ssm.{k}'] for k in ('conv_kernel','state_size','group_count','time_step_rank','inner_size')}) ssm_layers = tuple((i+1) % kv[f'{arch}.full_attention_interval'] != 0 for i in range(kv[f'{arch}.block_count'])) - for i,is_ssm in enumerate(ssm_layers): - if is_ssm: state_dict[f"blk.{i}.ssm_beta_alpha.weight"] = state_dict.pop(f"blk.{i}.ssm_beta.weight").cat( - state_dict.pop(f"blk.{i}.ssm_alpha.weight"), dim=0).contiguous() elif arch == 'kimi-linear': ssm_layers = tuple(x == 0 for x in n_kv_heads) n_kv_heads = max(n_kv_heads) @@ -458,15 +428,13 @@ class Transformer: return min(block._reusable_prefix_len(prefix_len, len(self._cached_tokens)) for block in self.blk) def warmup(self, chunk_size:int=256): - if not self.has_recurrent_block: - with Context(DEBUG=max(DEBUG.value, 1)): - for _ in range(2): list(zip(range(2), self.generate([0]))) + if not self.has_recurrent_prefill: + for _ in range(2): list(zip(range(2), self.generate([0]))) return x = Tensor.zeros(1, 1, self.blk[0].config.dim, device=self.token_embd.weight.device) for block in self.blk: block._init_state(x) - states = [getattr(block, name) for block in self.blk - for name in ("cache_kv", "cache_kv_scale", "freqs_cis", "conv_state", "recurrent_state") if hasattr(block, name)] - Tensor.realize(*states) + Tensor.realize(*(getattr(block, name) for block in self.blk + for name in ("cache_kv", "cache_kv_scale", "freqs_cis", "conv_state", "recurrent_state") if hasattr(block, name))) self.prefill_jit.cnt = self.rollout_jit.cnt = 1 warm = self.generate([0] * min(chunk_size, 256, self.max_context-1), chunk_size=chunk_size) with Context(JIT_BATCH_SIZE=getenv("PREFILL_JIT_BATCH_SIZE", 512)): next(warm) @@ -476,57 +444,33 @@ class Transformer: self._cached_tokens = [] def generate(self, tokens:list[int], chunk_size:int|None=None, temperature:float=0.0): - if self.has_recurrent_block: - yield from self._generate_recurrent(tokens, min(chunk_size or 256, 256), temperature) - return - chunk_size = chunk_size or 32 + chunk_size = min(chunk_size or (256 if self.has_recurrent_prefill else 32), 256) + if self.has_recurrent_block and not self.has_recurrent_prefill: chunk_size = 1 v_start_pos = UOp.variable("start_pos", 0, self.max_context-1) v_toks = UOp.variable("toks", 1, chunk_size) # TODO: use UOp.variable for temperature once float variables are supported - temp = Tensor([temperature]) + device = self.token_embd.weight.device + temp = Tensor([temperature], device=device) # assign all input tokens once, then slice from start_pos for the model call - t = Tensor(tokens + [0] * (self.max_context - len(tokens)), dtype="int32").reshape(1, self.max_context) + t = Tensor(tokens + [0] * (self.max_context + chunk_size - len(tokens)), dtype="int32", device=device).reshape(1, self.max_context+chunk_size) # recompute start_pos from what's currently valid in the caches start_pos = self.get_start_pos(tokens) + decode_resume = self.has_recurrent_prefill and bool(self._cached_tokens) and start_pos > 0 if start_pos < len(self._cached_tokens) and (resets := [r for b in self.blk for r in b._state_reset_ops()]): Tensor.realize(*resets) out, prompt_len = None, len(tokens) while len(tokens) < self.max_context: - n_toks = min(chunk_size, len(tokens) - start_pos) - sp, nt = v_start_pos.bind(start_pos), v_toks.bind(n_toks) - out = self(t[:, sp:sp+nt] if start_pos < prompt_len or out is None else out, sp, temp).realize() + padded_prefill = self.has_recurrent_prefill and start_pos < prompt_len and not decode_resume + n_toks = min(1 if decode_resume and start_pos < prompt_len else chunk_size, len(tokens)-start_pos) + sp = v_start_pos.bind(start_pos) + nt = chunk_size if padded_prefill else n_toks if self.has_recurrent_block else v_toks.bind(n_toks) + if decode_resume and start_pos < prompt_len and out is not None: + inp = out.assign(Tensor([[tokens[start_pos]]], dtype="int32", device=device)).realize() + else: inp = t[:, sp:sp+nt] if start_pos < prompt_len or out is None else out + valid_len = v_toks.bind(n_toks) if padded_prefill else None + out = self(inp, sp, temp, valid_len).realize() if self.has_recurrent_prefill else self(inp, sp, temp).realize() start_pos += n_toks # chunked prefill: keep processing until all prompt tokens are consumed if start_pos < len(tokens): continue tokens.append(int(out.item())) self._cached_tokens = tokens[:-1] yield tokens[-1] - - def _generate_recurrent(self, tokens:list[int], chunk_size:int, temperature:float): - start_pos = self.get_start_pos(tokens) - decode_resume = bool(self._cached_tokens) and start_pos > 0 - v_start_pos = UOp.variable("start_pos", 0, self.max_context-1) - v_toks = UOp.variable("toks", 1, chunk_size) - # TODO: use UOp.variable for temperature once float variables are supported - device = self.token_embd.weight.device - temp = Tensor([temperature], device=device) - t = Tensor(tokens + [0] * (self.max_context + chunk_size - len(tokens)), dtype="int32", device=device).reshape(1, self.max_context + chunk_size) - if start_pos < len(self._cached_tokens) and (resets := [r for b in self.blk for r in b._state_reset_ops()]): Tensor.realize(*resets) - out, prompt_len = None, len(tokens) - while len(tokens) < self.max_context: - recurrent_prefill = start_pos < prompt_len and not decode_resume - sp = v_start_pos.bind(start_pos) - actual_nt = min(1 if decode_resume and start_pos < prompt_len else chunk_size, len(tokens)-start_pos) - nt = chunk_size if recurrent_prefill else 1 - if decode_resume and start_pos < prompt_len and out is not None: - inp = out.assign(Tensor([[tokens[start_pos]]], dtype="int32", device=device)).realize() - elif start_pos < prompt_len or out is None: - inp = t[:, sp:sp+nt] - else: inp = out - valid_len = v_toks.bind(actual_nt) if recurrent_prefill else None - out = self(inp, sp, temp, valid_len=valid_len).realize() - start_pos += actual_nt - # chunked prefill: keep processing until all prompt tokens are consumed - if start_pos < len(tokens): continue - tokens.append(int(out.item())) - self._cached_tokens = tokens[:-1] - yield tokens[-1]