forked from tinygrad/tinygrad
llm: minimize recurrent prefill integration
This commit is contained in:
@@ -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'):
|
||||
|
||||
+3
-2
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
+70
-126
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user