From c21a552f3d2be0f9dc8281015e33498c19dc4e38 Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Mon, 3 Aug 2026 18:23:14 -0700 Subject: [PATCH] llm: bugfixes + warmup (#17384) --- extra/gemm/amd_flash_attention.py | 203 ------------------------------ test/unit/test_llm_server.py | 8 ++ tinygrad/llm/cli.py | 4 +- tinygrad/llm/model.py | 5 +- 4 files changed, 13 insertions(+), 207 deletions(-) delete mode 100644 extra/gemm/amd_flash_attention.py diff --git a/extra/gemm/amd_flash_attention.py b/extra/gemm/amd_flash_attention.py deleted file mode 100644 index d0896d5d78..0000000000 --- a/extra/gemm/amd_flash_attention.py +++ /dev/null @@ -1,203 +0,0 @@ -from tinygrad import Tensor, UOp, getenv -from tinygrad.uop.ops import AxisType, KernelInfo, Ops -from tinygrad.dtype import AddrSpace, dtypes -from tinygrad.helpers import DEBUG, GlobalCounters, Context -import math - -BLOCK_M, BLOCK_N = 64, 64 -WARP_SIZE = 32 -WMMA_M, WMMA_N, WMMA_K = 16, 16, 16 -WAVES_M, WAVES_N = 4, 1 -LANES_PER_WAVE_M, LANES_PER_WAVE_N = 2, 16 -WMMA_ACC = WMMA_M // LANES_PER_WAVE_M -THREADS_PER_BLOCK = WARP_SIZE * WAVES_M * WAVES_N -LDS_PAD = 4 # pad LDS rows to reduce bank conflicts - -WMMA_ARG = (WMMA_M, WMMA_N, WMMA_K), 'AMD', 32 -LOG2E = math.log2(math.e) - -def warp_shfl_xor(val, offset, lane): - """Read val from lane ^ offset using ds_bpermute.""" - idx = ((lane ^ offset) * 4).cast(dtypes.int) - if val.op is Ops.INDEX and val.addrspace == AddrSpace.REG: val = val.load() - return UOp(Ops.CUSTOM, dtypes.float, (idx, val), - arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_bpermute({0}, __builtin_bit_cast(int, {1})))") - -def warp_reduce_max(val, lane): - """Tree reduce MAX across LANES_PER_WAVE_N=16 lanes.""" - for offset in [8, 4, 2, 1]: - val = UOp(Ops.MAX, dtypes.float, (val, warp_shfl_xor(val, offset, lane))) - return val - -def warp_reduce_sum(val, lane): - """Tree reduce SUM across LANES_PER_WAVE_N=16 lanes.""" - for offset in [8, 4, 2, 1]: - val = val + warp_shfl_xor(val, offset, lane) - return val - -def amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp) -> UOp: - # inputs are (B*H, N, D) - BH, N, D = q.shape - assert N % BLOCK_M == 0 and N % BLOCK_N == 0, f"N={N} must be divisible by BLOCK_M={BLOCK_M} and BLOCK_N={BLOCK_N}" - assert D % WMMA_K == 0 and D % LANES_PER_WAVE_N == 0, f"D={D} must be divisible by WMMA_K={WMMA_K} and LANES_PER_WAVE_N={LANES_PER_WAVE_N}" - assert BLOCK_M % (WAVES_M * WMMA_M) == 0 and BLOCK_N % LANES_PER_WAVE_N == 0 - TM = BLOCK_M // (WAVES_M * LANES_PER_WAVE_M) - TN = BLOCK_N // (WAVES_N * LANES_PER_WAVE_N) - TD = D // (WAVES_N * LANES_PER_WAVE_N) - SCALE = 1.0 / math.sqrt(D) - - block_bh = UOp.range(BH, 0, AxisType.GLOBAL) - block_m = UOp.range(N // BLOCK_M, 1, AxisType.GLOBAL) - - q = q.reshape(BH, N//BLOCK_M, BLOCK_M, D)[block_bh, block_m] - k = k.reshape(BH, N//BLOCK_N, BLOCK_N, D)[block_bh] - v = v.reshape(BH, N//BLOCK_N, BLOCK_N, D)[block_bh] - o = o.reshape(BH, N//BLOCK_M, BLOCK_M, D)[block_bh, block_m] - - wave_m = UOp.range(WAVES_M, 2, AxisType.LOCAL) - wave_n = UOp.range(WAVES_N, 3, AxisType.LOCAL) - lane = UOp.range(WARP_SIZE, -1, AxisType.WARP) - tid = (wave_m * WAVES_N + wave_n) * WARP_SIZE + lane - lane_m = lane // LANES_PER_WAVE_N - lane_n = lane % LANES_PER_WAVE_N - - # LDS allocation: slot 0 = Q then P (shared), slot 1 = K then V - # TODO: the memory planner should be able to find this reuse - ELEMS_PER_THREAD = BLOCK_M * D // THREADS_PER_BLOCK - QP_lds = UOp.placeholder((BLOCK_M, D + LDS_PAD), dtypes.half, slot=0, addrspace=AddrSpace.LOCAL) - KV_lds = UOp.placeholder((BLOCK_N, D + LDS_PAD), dtypes.half, slot=1, addrspace=AddrSpace.LOCAL)[:, :D] - - # register state - acc = UOp.placeholder((TM, TD), dtypes.float, slot=2, addrspace=AddrSpace.REG) - m_i = UOp.placeholder((TM,), dtypes.float, slot=3, addrspace=AddrSpace.REG) - l_i = UOp.placeholder((TM,), dtypes.float, slot=4, addrspace=AddrSpace.REG) - acc = acc.after(acc.store(acc.const_like(0))) - m_i = m_i.after(m_i.store(m_i.const_like(-math.inf))) - l_i = l_i.after(l_i.store(l_i.const_like(0))) - - # ====== KV tile loop ====== - n_tile = UOp.range(N // BLOCK_N, 100, AxisType.REDUCE) - - # load Q + K into LDS (Q reloaded each iteration since P overwrites slot 0) - Q_lds = QP_lds[:, :D] - Q_store = Q_lds.after(n_tile).reshape(THREADS_PER_BLOCK, ELEMS_PER_THREAD)[tid].store( - q.reshape(THREADS_PER_BLOCK, ELEMS_PER_THREAD)[tid]) - K_store = KV_lds.reshape(THREADS_PER_BLOCK, ELEMS_PER_THREAD)[tid].store( - k[n_tile].reshape(THREADS_PER_BLOCK, ELEMS_PER_THREAD)[tid]) - # NOTE: no explicit barrier needed, the AFTER on the LOCAL buffers implies it in late codegen - Q_lds = Q_lds.after(UOp.group(Q_store, K_store)) - KV_lds_k = KV_lds.after(UOp.group(Q_store, K_store)) - - # -- S = Q @ K^T via WMMA (re-init each n_tile) -- - S_reg = UOp.placeholder((TM, TN), dtypes.float, slot=6, addrspace=AddrSpace.REG) - S_reg = S_reg.after(S_reg.after(n_tile).store(S_reg.const_like(0))) - k_qk = UOp.range(D // WMMA_K, 101, AxisType.REDUCE) - tm1 = UOp.range(TM // WMMA_ACC, 200) - tn1 = UOp.range(TN, 201) - S_frag = S_reg.reshape(TM // WMMA_ACC, WMMA_ACC, TN).permute(0, 2, 1)[tm1, tn1] - q_frag = Q_lds.reshape(WAVES_M, TM // WMMA_ACC, WMMA_M, D // WMMA_K, WMMA_K)[wave_m, tm1, lane_n, k_qk] - k_frag = KV_lds_k.reshape(WAVES_N, TN, WMMA_N, D // WMMA_K, WMMA_K)[wave_n, tn1, lane_n, k_qk] - qk = UOp.wmma(q_frag, k_frag, S_frag.after(k_qk), *WMMA_ARG) - qk_done = S_frag.store(qk).end(tm1, tn1).end(k_qk) - S_reg = S_reg.after(qk_done) - - # -- softmax in registers with warp shuffles -- - S_reg = S_reg.after(S_reg.store(S_reg * SCALE)) - - # per-thread local row max over TN=4 elements, then warp reduce across 16 lanes - m_ij = UOp.placeholder((TM,), dtypes.float, slot=7, addrspace=AddrSpace.REG) - m_ij = m_ij.after(m_ij.after(n_tile).store(m_ij.const_like(-math.inf))) - rm2 = UOp.range(TN, 261, AxisType.REDUCE) - m_ij = m_ij.after(m_ij.store(m_ij.after(rm2).maximum(S_reg[:, rm2])).end(rm2)) - # warp reduce max (in-place) - ri_w = UOp.range(TM, 270) - m_ij = m_ij.after(m_ij[ri_w].store(warp_reduce_max(m_ij[ri_w], lane)).end(ri_w)) - - # compute P = exp(S - m_ij) in S_reg - S_reg = S_reg.after(S_reg.store(((S_reg - m_ij.reshape(TM, 1).expand(TM, TN)) * LOG2E).exp2())) - - p_local = UOp.placeholder((TM,), dtypes.float, slot=8, addrspace=AddrSpace.REG) - p_local = p_local.after(p_local.after(n_tile).store(p_local.const_like(0))) - rp2 = UOp.range(TN, 291, AxisType.REDUCE) - p_local = p_local.after(p_local.store(p_local.after(rp2) + S_reg[:, rp2]).end(rp2)) - ri_ws = UOp.range(TM, 295) - p_sum = p_local.after(p_local[ri_ws].store(warp_reduce_sum(p_local[ri_ws], lane)).end(ri_ws)) - - # write P = exp(S - m_ij) to P_lds (reuses slot 0, Q no longer needed) - P_lds = QP_lds[:, :BLOCK_N] - P_write = P_lds.reshape(WAVES_M, TM // WMMA_ACC, WMMA_ACC, LANES_PER_WAVE_M, WAVES_N, TN, LANES_PER_WAVE_N) - P_write = P_write.permute((0, 4, 3, 6, 1, 2, 5)).reshape(THREADS_PER_BLOCK, TM, TN) - P_store = P_write[tid].store(S_reg.cast(dtypes.half)) - - # -- online softmax correction -- - ri4 = UOp.range(TM, 330) - m_new_val = m_i[ri4].maximum(m_ij[ri4]) - alpha_val = ((m_i[ri4] - m_new_val) * LOG2E).exp2() - beta_val = ((m_ij[ri4] - m_new_val) * LOG2E).exp2() - rj4 = UOp.range(TD, 331) - correction = UOp.group( - acc[ri4, rj4].store(alpha_val * acc[ri4, rj4]).end(rj4), - l_i[ri4].store(alpha_val * l_i[ri4] + beta_val * p_sum[ri4]), - m_i[ri4].store(m_new_val), - ).end(ri4) - acc = acc.after(correction) - l_i = l_i.after(correction) - m_i = m_i.after(correction) - - # load V into KV_lds (must wait for QK WMMA to finish reading K from KV_lds) - V_store = KV_lds.after(qk_done).reshape(THREADS_PER_BLOCK, ELEMS_PER_THREAD)[tid].store( - v[n_tile].reshape(THREADS_PER_BLOCK, ELEMS_PER_THREAD)[tid]) - # NOTE: no explicit barrier needed, the AFTER on the LOCAL buffers implies it in late codegen - P_lds = P_lds.after(UOp.group(P_store, V_store)) - KV_lds_v = KV_lds.after(UOp.group(P_store, V_store)) - - # -- acc += P @ V via WMMA -- - k_pv = UOp.range(BLOCK_N // WMMA_K, 400, AxisType.REDUCE) - tm2 = UOp.range(TM // WMMA_ACC, 401) - tn2 = UOp.range(TD, 402) - acc_frag = acc.reshape(TM // WMMA_ACC, WMMA_ACC, TD).permute(0, 2, 1)[tm2, tn2] - p_frag = P_lds.reshape(WAVES_M, TM // WMMA_ACC, WMMA_M, BLOCK_N // WMMA_K, WMMA_K)[wave_m, tm2, lane_n, k_pv] - v_frag = KV_lds_v.reshape(WAVES_N, TD, WMMA_N, BLOCK_N // WMMA_K, WMMA_K)[wave_n, tn2, lane_n, k_pv] - pv = UOp.wmma(p_frag, v_frag, acc_frag.after(k_pv), *WMMA_ARG) - - # end KV tile loop - n_tile_end = acc_frag.store(pv).end(tm2, tn2).end(k_pv).end(n_tile) - acc = acc.after(n_tile_end) - l_i = l_i.after(n_tile_end) - m_i = m_i.after(n_tile_end) - - # normalize: acc /= l_i - acc = acc.after(acc.store(acc * (1 / l_i).reshape(TM, 1).expand(TM, TD))) - - # store output - o = o.reshape(WAVES_M, TM // WMMA_ACC, WMMA_ACC, LANES_PER_WAVE_M, WAVES_N, TD, LANES_PER_WAVE_N) - o = o.permute((0, 4, 3, 6, 1, 2, 5)).reshape(THREADS_PER_BLOCK, TM, TD) - return o[tid].store(acc).end(wave_m, wave_n, lane).end(block_m, block_bh).sink(arg=KernelInfo(opts_to_apply=())) - -if __name__ == "__main__": - B, H, N, D = getenv("B", 1), getenv("H", 32), getenv("N", 1024), getenv("D", 64) - q = Tensor.rand(B, H, N, D).cast(dtypes.half) - k = Tensor.rand(B, H, N, D).cast(dtypes.half) - v = Tensor.rand(B, H, N, D).cast(dtypes.half) - o = Tensor.empty(B, H, N, D, dtype=dtypes.float) - with Context(DEBUG=0): Tensor.realize(q, k, v) - - q_flat, k_flat, v_flat, o_flat = q.reshape(B*H, N, D), k.reshape(B*H, N, D), v.reshape(B*H, N, D), o.reshape(B*H, N, D) - NUM_RUNS = getenv("CNT", 5) - ets = [] - with Context(DEBUG=2): - for _ in range(NUM_RUNS): - GlobalCounters.reset() - tst = Tensor.custom_kernel(o_flat, q_flat, k_flat, v_flat, fxn=amd_flash_attention)[0].realize() - ets.append(GlobalCounters.time_sum_s) - print(f"best time: {min(ets)*1e3:.2f}ms") - - if getenv("VERIFY", 1): - with Context(DEBUG=0): - ref = q.float().scaled_dot_product_attention(k.float(), v.float()).reshape(B*H, N, D).realize() - err = (ref - tst).square().mean().item() - print(f"mean squared error {err}") - if err > 1e-2: - raise RuntimeError("flash attention is wrong!") - else: - print("flash attention is correct!") diff --git a/test/unit/test_llm_server.py b/test/unit/test_llm_server.py index 6e10f62924..4eee0e5225 100644 --- a/test/unit/test_llm_server.py +++ b/test/unit/test_llm_server.py @@ -11,6 +11,14 @@ V_START_POS = UOp.variable("start_pos", 0, TEST_CONFIG.max_context-1) V_TOKS = UOp.variable("toks", 1, 32) # 32 is the default chunk_size in generate class TestTransformerGenerate(unittest.TestCase): + def test_warmup(self): + model, calls = Transformer(TEST_CONFIG), [] + def generate(tokens): + calls.append(tokens) + yield from (1, 2) + with patch.object(model, "generate", generate): model.warmup() + self.assertEqual(calls, [[0], [0]]) + def test_first_recurrent_generate_before_state_init(self): model = Transformer(TEST_CONFIG) model.has_recurrent_block = True diff --git a/tinygrad/llm/cli.py b/tinygrad/llm/cli.py index cba67a4a1c..d0865e26ba 100644 --- a/tinygrad/llm/cli.py +++ b/tinygrad/llm/cli.py @@ -164,9 +164,7 @@ def main(): # warmup the JIT if args.warmup or args.serve: - # run 2 tokens through the model twice to capture the JIT before serving - with Context(DEBUG=max(DEBUG.value, 1)): - for _ in range(2): list(zip(range(2), model.generate([0]))) + 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/model.py b/tinygrad/llm/model.py index 3fab067602..1f4feb9698 100644 --- a/tinygrad/llm/model.py +++ b/tinygrad/llm/model.py @@ -303,7 +303,7 @@ class GatedDeltaNetBlock(FFNBlock): def _init_state(self, x): if not hasattr(self, "conv_state"): self.conv_state = Tensor.zeros(x.shape[0], self.ssm_conv_kernel-1, self.conv_channels, device=x.device).clone() - self.recurrent_state = Tensor.zeros(x.shape[0], self.num_v_heads, self.head_v_dim, self.head_v_dim, device=x.device).clone() + self.recurrent_state = Tensor.zeros(x.shape[0], self.num_v_heads, self.head_v_dim, self.head_k_dim, device=x.device).clone() class Transformer: def __init__(self, config:TransformerConfig): @@ -416,6 +416,9 @@ class Transformer: Tensor.realize(*params) return model, kv + def warmup(self): + for _ in range(2): list(zip(range(2), self.generate([0]))) + def get_start_pos(self, tokens:list[int]) -> int: prefix_len = sum(1 for _ in itertools.takewhile(lambda ab: ab[0] == ab[1], zip(tokens[:-1], self._cached_tokens))) return min(block._reusable_prefix_len(prefix_len, len(self._cached_tokens)) for block in self.blk)