diff --git a/test/unit/test_attention.py b/test/unit/test_attention.py index 622712cc6e..8203caf0d9 100644 --- a/test/unit/test_attention.py +++ b/test/unit/test_attention.py @@ -1,6 +1,6 @@ import unittest import numpy as np -from tinygrad import Tensor, dtypes +from tinygrad import Tensor, dtypes, nn from tinygrad.llm.model import ( GatedDeltaNetBlock, SSMConfig, TransformerBlock, TransformerConfig, apply_rope as apply_rope_new, precompute_freqs_cis, pairwise_topk, @@ -45,10 +45,10 @@ class TestGatedDeltaNetBlock(unittest.TestCase): return Tensor.linspace(start, stop, int(np.prod(shape)), dtype=dtypes.float32).reshape(*shape) def _make_config(self, **kwargs): - return TransformerConfig(**({"num_blocks":1, "dim":4, "hidden_dim":8, "n_heads":1, "n_kv_heads":1, - "norm_eps":1e-5, "vocab_size":32, "head_dim":4, "rope_theta":10000.0, - "rope_dim":4, "v_head_dim":4, "max_context":4, "ssm_layers":(True,), - "ssm":SSMConfig(conv_kernel=2, state_size=2, group_count=1, time_step_rank=1, inner_size=2)} | kwargs)) + return TransformerConfig(**({"num_blocks":1, "dim":32, "hidden_dim":64, "n_heads":1, "n_kv_heads":1, + "norm_eps":1e-5, "vocab_size":32, "head_dim":32, "rope_theta":10000.0, + "rope_dim":32, "v_head_dim":32, "max_context":4, "ssm_layers":(True,), + "ssm":SSMConfig(conv_kernel=2, state_size=32, group_count=1, time_step_rank=1, inner_size=32)} | kwargs)) def _make_block(self, config:TransformerConfig) -> GatedDeltaNetBlock: block = GatedDeltaNetBlock(config, config.ssm) @@ -79,6 +79,10 @@ class TestGatedDeltaNetBlock(unittest.TestCase): recurrent_state = cache[:, conv_flat:].reshape(cache.shape[0], block.num_v_heads, block.head_v_dim, block.head_v_dim) return conv_state, recurrent_state + def _reset_state(self, block:GatedDeltaNetBlock): + Tensor.realize(block.conv_state.assign(block.conv_state.const_like(0)), + block.recurrent_state.assign(block.recurrent_state.const_like(0))) + def _linear_np(self, x:np.ndarray, weight:np.ndarray) -> np.ndarray: return x.astype(np.float32) @ weight.T.astype(np.float32) @@ -86,7 +90,7 @@ class TestGatedDeltaNetBlock(unittest.TestCase): x_float = x.astype(np.float32) return (x_float / np.sqrt((x_float * x_float).mean(axis=-1, keepdims=True) + eps)) * weight.astype(np.float32) - def _normalize_np(self, x:np.ndarray, eps:float=1e-12) -> np.ndarray: + def _normalize_np(self, x:np.ndarray, eps:float=1e-6) -> np.ndarray: return x / np.maximum(np.sqrt((x * x).sum(axis=-1, keepdims=True)), eps) def _softplus_np(self, x:np.ndarray) -> np.ndarray: @@ -148,6 +152,12 @@ class TestGatedDeltaNetBlock(unittest.TestCase): x = Tensor.linspace(-1.0, 1.0, 3 * config.dim, dtype=dtypes.float32).reshape(1, 3, config.dim) expected_outs, expected_conv, expected_recurrent = self._naive_attention(block, x) + out = self._run_attention(block, x, 0) + conv_state, recurrent_state = self._cache_views(block) + np.testing.assert_allclose(out, np.concatenate(expected_outs, axis=1), rtol=1e-3, atol=1e-3) + np.testing.assert_allclose(conv_state, expected_conv[-1], rtol=1e-3, atol=1e-3) + np.testing.assert_allclose(recurrent_state, expected_recurrent[-1], rtol=1e-3, atol=1e-3) + self._reset_state(block) for step in range(x.shape[1]): out = self._run_attention(block, x[:, step:step+1], step) @@ -163,7 +173,7 @@ class TestGatedDeltaNetBlock(unittest.TestCase): prompt = Tensor.linspace(0.75, -0.75, 2 * config.dim, dtype=dtypes.float32).reshape(1, 2, config.dim) for i in range(warmup.shape[1]): self._run_attention(block, warmup[:, i:i+1], i) - Tensor.realize(*block._state_reset_ops()) + self._reset_state(block) expected_outs, expected_conv, expected_recurrent = self._naive_attention(block, prompt) for step in range(prompt.shape[1]): @@ -177,18 +187,64 @@ class TestGatedDeltaNetBlock(unittest.TestCase): err_msg=f"GatedDeltaNet reset recurrent cache mismatch at step {step}") def test_kda_channel_decay(self): - config = self._make_config(n_heads=2, ssm=SSMConfig(conv_kernel=2, state_size=2, group_count=2, time_step_rank=2, inner_size=4, kda=True)) - block, x = GatedDeltaNetBlock(config, config.ssm), Tensor([[[1., 2., 0., 0.]]]) - # f_b(f_a(x)) = [1, 2, 3, 4] + config = self._make_config(dim=4, hidden_dim=8, n_heads=2, head_dim=4, rope_dim=4, v_head_dim=4, + ssm=SSMConfig(conv_kernel=2, state_size=2, group_count=2, time_step_rank=2, inner_size=4, kda=True)) + block, x = GatedDeltaNetBlock(config, config.ssm), Tensor([[[1., 2., 0., 0.], [2., 1., 0., 0.]]]) block.ssm_f_a.weight = Tensor([[1., 0., 0., 0.], [0., 1., 0., 0.]]) block.ssm_f_b.weight = Tensor([[1., 0.], [0., 1.], [1., 1.], [2., 1.]]) block._init_state(x) initial_state = Tensor.arange(8, dtype=dtypes.float32).reshape(1, 2, 2, 2) block.recurrent_state.assign(initial_state).realize() block.ssm_a = Tensor([[-1.], [-1.]]) - block._attention(x, 0).realize() - alpha = np.exp(-self._softplus_np(np.arange(1, 5)).reshape(1, 2, 1, 2)) - np.testing.assert_allclose(block.recurrent_state.numpy(), initial_state.numpy() * alpha, rtol=1e-5, atol=1e-5) + block._attention(x, x.shape[1]).realize() + alpha = np.exp(-self._softplus_np(np.array([[1, 2, 3, 4], [2, 1, 3, 5]])).reshape(2, 2, 2)).prod(0) + np.testing.assert_allclose(block.recurrent_state.numpy(), initial_state.numpy() * alpha[..., None], rtol=1e-5, atol=1e-5) + + def test_kda_prefill_matches_decode(self): + config = self._make_config(ssm=SSMConfig(conv_kernel=2, state_size=32, group_count=1, time_step_rank=1, inner_size=32, kda=True)) + block = GatedDeltaNetBlock(config, config.ssm) + for p in nn.state.get_parameters(block): + p.replace(self._tensor_linspace(-0.05, 0.05, p.shape) if len(p.shape) > 1 else self._tensor_linspace(0.05, 0.1, p.shape)) + x = self._tensor_linspace(-0.5, 0.5, (1, 3, config.dim)) + prefill = self._run_attention(block, x, 0) + prefill_conv, prefill_recurrent = self._cache_views(block) + self._reset_state(block) + decode = np.concatenate([self._run_attention(block, x[:, i:i+1], i) for i in range(3)], axis=1) + decode_conv, decode_recurrent = self._cache_views(block) + np.testing.assert_allclose(prefill, decode, rtol=1e-3, atol=1e-3) + np.testing.assert_allclose(prefill_conv, decode_conv, rtol=1e-3, atol=1e-3) + np.testing.assert_allclose(prefill_recurrent, decode_recurrent, rtol=1e-3, atol=1e-3) + + def test_varied_chunk_sizes_match_decode(self): + for kda in (False, True): + ssm = SSMConfig(conv_kernel=2, state_size=32, group_count=1, time_step_rank=1, inner_size=32, kda=kda) + config = self._make_config(ssm=ssm) + if kda: + block = GatedDeltaNetBlock(config, config.ssm) + for p in nn.state.get_parameters(block): + p.replace(self._tensor_linspace(-0.05, 0.05, p.shape) if len(p.shape) > 1 else self._tensor_linspace(0.05, 0.1, p.shape)) + else: block = self._make_block(config) + x = self._tensor_linspace(-0.5, 0.5, (1, 4, config.dim)) + decode = np.concatenate([self._run_attention(block, x[:, i:i+1], i) for i in range(4)], axis=1) + decode_conv, decode_recurrent = self._cache_views(block) + for chunking in ([4], [2, 2], [1, 3], [3, 1], [2, 1, 1]): + self._reset_state(block) + outs, start = [], 0 + for size in chunking: + outs.append(self._run_attention(block, x[:, start:start+size], start)) + start += size + chunked_conv, chunked_recurrent = self._cache_views(block) + np.testing.assert_allclose(np.concatenate(outs, axis=1), decode, rtol=1e-3, atol=1e-3, err_msg=f"{kda=} {chunking=}") + np.testing.assert_allclose(chunked_conv, decode_conv, rtol=1e-3, atol=1e-3, err_msg=f"{kda=} {chunking=}") + np.testing.assert_allclose(chunked_recurrent, decode_recurrent, rtol=1e-3, atol=1e-3, err_msg=f"{kda=} {chunking=}") + + def test_start_zero_resets_realized_state(self): + config, x = self._make_config(max_context=3), self._tensor_linspace(-1, 1, (1, 3, 32)) + block = self._make_block(config) + self._run_attention(block, x, 0) + restarted = self._run_attention(block, x[:, :2], 0) + fresh = self._run_attention(self._make_block(config), x[:, :2], 0) + np.testing.assert_allclose(restarted, fresh, rtol=1e-3, atol=1e-3) class TestPairwiseTopk(unittest.TestCase): def test_basic_topk(self): diff --git a/tinygrad/llm/model.py b/tinygrad/llm/model.py index 7d2034a018..0d167f4313 100644 --- a/tinygrad/llm/model.py +++ b/tinygrad/llm/model.py @@ -138,8 +138,6 @@ class FFNBlock: # given the token-prefix match, return how much cached state this block can still reuse def _reusable_prefix_len(self, prefix_len:int, cached_len:int) -> int: return prefix_len - # 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 @@ -274,45 +272,65 @@ class GatedDeltaNetBlock(FFNBlock): def _attention(self, x:Tensor, start_pos:int|UOp) -> Tensor: B, T, _ = x.shape - assert T == 1, "GatedDeltaNetBlock currently only supports T=1" + # bind ints to a variable so the reset flag stays a runtime value (it toggles when generation restarts at position 0) + start_pos = start_pos if isinstance(start_pos, UOp) else UOp.variable("start_pos", 0, self.config.max_context-1).bind(start_pos) + initial = Tensor(start_pos).eq(0) is_kda = hasattr(self, "ssm_g_a") + symbolic = isinstance(T, UOp) + T_pad = x.max_shape[1] # symbolic chunks are padded to their max size: one graph serves every size # input processing x = x.half() out_gate = self.ssm_g_b(self.ssm_g_a(x)) if is_kda else self.attn_gate(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) + out_gate = out_gate.reshape(B, T, self.num_v_heads, self.head_v_dim) + beta = self.ssm_beta(x).sigmoid().reshape(B, T, self.num_v_heads) alpha = self.ssm_f_b(self.ssm_f_a(x)) if is_kda else self.ssm_alpha(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) + log_alpha = ((alpha.float() + self.ssm_dt["bias"]).softplus().reshape(B, T, self.num_v_heads, -1) * + self.ssm_a.reshape(self.num_v_heads, -1)) - # qkv conv - 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() + # qkv conv, conv_state is reset when starting from position 0 + conv_state = initial.where(0, self.conv_state) + # assemble the conv window in a static-size buffer: [conv_state | qkv rows | zero-pad]. + # padded steps are exact no-ops: beta=0 (delta rule off), log_alpha=0 (decay 1 after exp) + win = Tensor.zeros(B, self.ssm_conv_kernel-1 + T_pad, self.conv_channels).uop + win = win.after(win[:, :self.ssm_conv_kernel-1].store(conv_state.cast(win.dtype).uop)) + win = win.after(win[:, self.ssm_conv_kernel-1:self.ssm_conv_kernel-1+T].store(self.attn_qkv(x).cast(win.dtype).uop)) + conv_window = Tensor(win) + # the last conv_kernel-1 columns of the window become the next conv state + conv_state_store = self.conv_state.uop.store(conv_window[:, T:T+self.ssm_conv_kernel-1].cast(self.conv_state.dtype).uop) + + conv_out = functools.reduce(lambda a,b: a+b, + (conv_window[:, i:i+T_pad] * self.ssm_conv1d["weight"][:, i] for i in range(self.ssm_conv_kernel))).silu() + if symbolic: + out_gate = out_gate.pad_to((B, T_pad, self.num_v_heads, self.head_v_dim)) + beta, log_alpha = beta.pad_to((B, T_pad, self.num_v_heads)), log_alpha.pad_to((B, T_pad, *log_alpha.shape[2:])) 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) + qk_eps = 1e-12 if is_kda else 1e-6 + q, k = (z.reshape(B, T_pad, self.num_k_heads, self.head_k_dim).normalize(dim=-1, eps=qk_eps) + .repeat(1, 1, self.num_v_heads//self.num_k_heads, 1) for z in (q, k)) + v = v.reshape(B, T_pad, self.num_v_heads, self.head_v_dim) + # layout the per-step operands to broadcast against the (B, H, V, K) state + q, k, v, beta = (z.transpose(1, 2).float() for z in (q, k, v, beta)) + q, k, v, beta = q.unsqueeze(-2) * self.head_k_dim**-0.5, k.unsqueeze(-2), v.unsqueeze(-1), beta.unsqueeze(-1).unsqueeze(-1) + alpha = log_alpha.transpose(1, 2).exp().unsqueeze(-1) # per-channel decay for kda, per-head otherwise (B, H, T, V|1, 1) - # recurrent - recurrent_state = self.recurrent_state * alpha - recurrent_state = recurrent_state + ((v - recurrent_state@k) * beta)@k.transpose(-1, -2) + # recurrent: scan over the (padded) tokens, updating the recurrent state. collect the per-step outputs + state = Tensor(self.recurrent_state.uop.after(conv_state_store)).float() # carry the conv write into this graph + state = initial.where(0, state) + outs = [] + for t in range(T_pad): + s1 = state * alpha[:, :, t] # decay the state + delta = (v[:, :, t] - (s1*k[:, :, t]).sum(-1, keepdim=True)) * beta[:, :, t] # the delta rule update + state = s1 + delta * k[:, :, t] + outs.append((state * q[:, :, t]).sum(-1)) - # 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)) + # store the updated recurrent state in place, then read the stacked outputs after the write + core = Tensor(outs[0].stack(*outs[1:], dim=1).contiguous().uop.after(self.recurrent_state.uop.store(state.cast(self.recurrent_state.dtype).uop))) - # output - core_attn_out = self.ssm_norm((recurrent_state@q).squeeze(-1).reshape(B, 1, self.num_v_heads, self.head_v_dim)) - out_gate = out_gate.sigmoid() if is_kda else out_gate.silu() - return self.ssm_out((core_attn_out * out_gate).reshape(B, 1, -1).cast(x.dtype)) - - # recurrent state can't be partially reused after divergence, force a full rebuild - def _state_reset_ops(self): - return [self.conv_state.assign(self.conv_state.const_like(0)), - self.recurrent_state.assign(self.recurrent_state.const_like(0))] if hasattr(self, "conv_state") else [] + # output; undo the padding before the output projection + z = (self.ssm_norm(core) * (out_gate.sigmoid() if is_kda else out_gate.silu())).cast(x.dtype).contiguous() + if symbolic: z = z[:, :T] + return self.ssm_out(z.reshape(B, T, -1)) def _init_state(self, x): if not hasattr(self, "conv_state"): @@ -453,7 +471,6 @@ class Transformer: t = Tensor(tokens + [0] * (self.max_context - len(tokens)), dtype="int32").reshape(1, self.max_context) # recompute start_pos from what's currently valid in the caches start_pos = self.get_start_pos(tokens) - 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)