From 475ad15f2829dbd8afa53d596af01e981fcdc538 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Sun, 2 Aug 2026 00:36:02 +0000 Subject: [PATCH] 27b 935/47 --- test/null/test_hcq_iface.py | 17 + test/unit/test_llm_quant_cpu.py | 103 +++++- tinygrad/llm/kernels/amd.py | 497 +++++++++++++++++++++-------- tinygrad/llm/model.py | 133 +++++--- tinygrad/runtime/support/system.py | 3 +- 5 files changed, 585 insertions(+), 168 deletions(-) diff --git a/test/null/test_hcq_iface.py b/test/null/test_hcq_iface.py index 8af6353a0e..16ea69d96e 100644 --- a/test/null/test_hcq_iface.py +++ b/test/null/test_hcq_iface.py @@ -1,9 +1,26 @@ import unittest, array, time from tinygrad.helpers import mv_address from tinygrad.runtime.support.hcq import MMIOInterface +from tinygrad.runtime.support.memory import VirtMapping +from tinygrad.runtime.support.system import PCIIfaceBase from tinygrad.runtime.support.usb import USBMMIOInterface from test.mockgpu.usb import MockUSB +class TestPCIIface(unittest.TestCase): + def test_sysmem_mapping_respects_uncached(self): + class MM: + def alloc_vaddr(self, size, align): return 0x10000 + def map_range(self, vaddr, size, paddrs, aspace, uncached=False, snooped=False): + return VirtMapping(vaddr, size, paddrs, aspace, uncached, snooped) + class PCI: + def bar_info(self, bar): return 0, 256 << 20 + def alloc_sysmem(self, size, **kwargs): return memoryview(bytearray(size)), [0x20000] + iface = PCIIfaceBase.__new__(PCIIfaceBase) + iface.dev, iface.vram_bar, iface.pci_dev = None, 0, PCI() + iface.dev_impl = type("DevImpl", (), {"mm": MM()})() + for uncached in (False, True): + with self.subTest(uncached=uncached): self.assertEqual(iface.alloc(4096, host=True, uncached=uncached).meta.mapping.uncached, uncached) + class TestHCQIface(unittest.TestCase): def setUp(self): self.size = 4 << 10 diff --git a/test/unit/test_llm_quant_cpu.py b/test/unit/test_llm_quant_cpu.py index 4b24d57b61..6f77cb370b 100644 --- a/test/unit/test_llm_quant_cpu.py +++ b/test/unit/test_llm_quant_cpu.py @@ -3,6 +3,7 @@ import numpy as np from tinygrad import Device, Tensor, TinyJit, UOp, dtypes, nn from tinygrad.llm.gguf import _GGML_QUANT, ggml_data_to_tensor +from tinygrad.llm.kernels import amd as llm_amd from tinygrad.llm.kernels.cpu import (attention_decode, attention_prefill, causal_conv_silu, expert_pair, expert_silu, expert_weighted_sum, f16_linear, f16_matvec, gated_delta, gated_delta_prefill, gated_delta_q8, gdn_qkv, iq3_repack, moe_ffn, q6_argmax, q8_batched_pair, @@ -11,7 +12,7 @@ from tinygrad.llm.kernels.cpu import (attention_decode, attention_prefill, causa silu, silu_mul, uop_attention_prefill, uop_f16_matvec, uop_linear, uop_moe_ffn, uop_q8_linear_pair, uop_q8_prequant_linear, uop_expert_silu_weighted, weighted_sum) from tinygrad.llm.kernels.cpu import _dot_bytes_ptr, _dot_nibbles_ptr -from tinygrad.llm.model import biased_sigmoid_topk, pairwise_topk, ExpertWeights, FFNBlock, Linear, Transformer, TransformerConfig +from tinygrad.llm.model import biased_sigmoid_topk, pairwise_topk, Embedding, ExpertWeights, FFNBlock, Linear, Transformer, TransformerConfig from tinygrad.uop.ops import KernelInfo @@ -41,6 +42,106 @@ def random_packed(rng:np.random.Generator, ggml_type:int, elements:int) -> np.nd @unittest.skipUnless(Device.DEFAULT == "AMD", "requires DEV=AMD") class TestLLMQuantAMD(unittest.TestCase): + @staticmethod + def assert_q8_equal(result:tuple[Tensor, Tensor, Tensor], expected:np.ndarray): + grouped = expected.reshape(expected.shape[0], -1, 32) + scale = np.maximum(np.max(np.abs(grouped), axis=-1) / 127, 1e-8) + quant = np.clip(np.rint(grouped / scale[..., None]), -127, 127).astype(np.int8) + np.testing.assert_equal(result[0].numpy().view(np.int8).reshape(grouped.shape), quant) + np.testing.assert_allclose(result[1].numpy(), scale, rtol=1e-6, atol=1e-8) + np.testing.assert_equal(result[2].numpy(), quant.astype(np.int32).sum(-1)) + + def test_gated_delta_decode_matches_reference(self): + rng = np.random.default_rng(39) + batch, heads, dim = 1, 2, 128 + q, k, v = [rng.standard_normal((batch, heads, dim), dtype=np.float32) for _ in range(3)] + beta, alpha = rng.random((batch, heads), dtype=np.float32), rng.uniform(0.9, 1, (batch, heads)).astype(np.float32) + state = rng.standard_normal((batch, heads, dim, dim), dtype=np.float32).astype(np.float16) + state_k, state_q = np.einsum("bhij,bhj->bhi", state, k), np.einsum("bhij,bhj->bhi", state, q) + delta = (v - state_k * alpha[..., None]) * beta[..., None] + expected_core = state_q * alpha[..., None] + delta * np.sum(k*q, axis=-1)[..., None] + expected_state = (state * alpha[..., None, None] + delta[..., None] * k[..., None, :]).astype(np.float16) + core, next_state = llm_amd.gated_delta_decode( + *(Tensor(x, device="AMD") for x in (q, k, v, beta, alpha)), Tensor(state, device="AMD")) + Tensor.realize(core, next_state) + np.testing.assert_allclose(core.numpy(), expected_core, rtol=2e-4, atol=1e-3) + np.testing.assert_allclose(next_state.numpy(), expected_state, rtol=1e-3, atol=4e-3) + + def test_f16_matvec_matches_reference(self): + rng, in_features, out_features = np.random.default_rng(38), 512, 13 + x = rng.standard_normal((1, in_features), dtype=np.float32).astype(np.float16) + weight = rng.standard_normal((out_features, in_features), dtype=np.float32).astype(np.float16) + got = llm_amd.f16_matvec(Tensor(x, device="AMD"), Tensor(weight, device="AMD")).numpy() + np.testing.assert_allclose(got, x.astype(np.float32) @ weight.astype(np.float32).T, rtol=2e-5, atol=2e-4) + + def test_fused_rmsnorm_quantization_matches_reference(self): + rng, eps = np.random.default_rng(37), 1e-6 + x = rng.standard_normal((1, 256), dtype=np.float32) + weight = rng.standard_normal((256,), dtype=np.float32).astype(np.float16) + normalized = x / np.sqrt(np.mean(x*x, axis=-1, keepdims=True) + eps) * weight + self.assert_q8_equal(llm_amd.q8_rmsnorm(Tensor(x, device="AMD"), Tensor(weight, device="AMD"), eps), normalized) + + core = rng.standard_normal((1, 2, 128), dtype=np.float32) + gate = rng.standard_normal((1, 1, 2, 128), dtype=np.float32) + head_norm = core / np.sqrt(np.mean(core*core, axis=-1, keepdims=True) + eps) * weight[:128] + expected = head_norm * gate.reshape(1, 2, 128) / (1 + np.exp(-gate.reshape(1, 2, 128))) + self.assert_q8_equal(llm_amd.q8_gated_rmsnorm(*(Tensor(v, device="AMD") for v in (core, gate, weight[:128])), eps), + expected.reshape(1, -1)) + + def test_gated_delta_prefill_matches_sequential_reference(self): + rng = np.random.default_rng(36) + batch, heads, tokens, dim = 1, 2, 5, 128 + q, k, v = [rng.standard_normal((batch, heads, tokens, dim), dtype=np.float32) for _ in range(3)] + q = q / np.linalg.norm(q, axis=-1, keepdims=True) / np.float32(np.sqrt(dim)) + k = k / np.linalg.norm(k, axis=-1, keepdims=True) + beta, alpha = rng.random((batch, heads, tokens), dtype=np.float32), rng.uniform(0.9, 1, (batch, heads, tokens)).astype(np.float32) + state = rng.standard_normal((batch, heads, dim, dim), dtype=np.float32).astype(np.float16) + expected_core, expected_state = np.empty_like(q), state.astype(np.float32) + for token in range(tokens): + state_k = np.einsum("bhij,bhj->bhi", expected_state, k[:, :, token]) + state_q = np.einsum("bhij,bhj->bhi", expected_state, q[:, :, token]) + delta = (v[:, :, token] - state_k * alpha[:, :, token, None]) * beta[:, :, token, None] + expected_core[:, :, token] = state_q * alpha[:, :, token, None] + delta * np.sum(k[:, :, token] * q[:, :, token], axis=-1)[..., None] + expected_state = expected_state * alpha[:, :, token, None, None] + delta[..., None] * k[:, :, token, None, :] + core, next_state = llm_amd.gated_delta_prefill( + *(Tensor(x, device="AMD") for x in (q, k, v, beta, alpha)), Tensor(state, device="AMD")) + np.testing.assert_allclose(core.numpy(), expected_core, rtol=2e-4, atol=1e-3) + np.testing.assert_allclose(next_state.numpy(), expected_state.astype(np.float16), rtol=2e-4, atol=1e-3) + + def test_q8_quantize_matches_reference(self): + rng, tokens, in_features = np.random.default_rng(35), 3, 256 + x = rng.standard_normal((tokens, in_features), dtype=np.float32) + grouped = x.reshape(tokens, -1, 32) + expected_scale = np.maximum(np.max(np.abs(grouped), axis=-1) / 127, 1e-8) + expected_quant = np.clip(np.rint(grouped / expected_scale[..., None]), -127, 127).astype(np.int8) + quant, scale, group_sum = llm_amd.q8_quantize_sum(Tensor(x, device="AMD"), tokens, in_features) + np.testing.assert_equal(quant.numpy().view(np.int8).reshape(grouped.shape), expected_quant) + np.testing.assert_allclose(scale.numpy(), expected_scale, rtol=1e-7, atol=0) + np.testing.assert_equal(group_sum.numpy(), expected_quant.astype(np.int32).sum(-1)) + + def test_q4_embedding_matches_reference(self): + rng, vocab_size, embed_size = np.random.default_rng(34), 16, 256 + raw = random_packed(rng, 12, vocab_size * embed_size) + expected = ggml_data_to_tensor(Tensor(raw), vocab_size * embed_size, 12).reshape(vocab_size, embed_size).half() + storage = Tensor(np.concatenate((np.zeros(68, dtype=np.uint8), raw)), dtype=dtypes.uint8, device="AMD").realize() + embedding = Embedding(vocab_size, embed_size) + embedding.set_quantized(storage[68:], 12) + idx = np.array([[7, 1, 15], [0, 4, 7]], dtype=np.int32) + np.testing.assert_equal(embedding(Tensor(idx, device="AMD")).numpy(), expected.numpy()[idx]) + + def test_iq4_lut_is_ready_for_jit_capture(self): + rng, in_features, out_features = np.random.default_rng(33), 256, 16 + raw = random_packed(rng, 23, out_features * in_features) + weight = ggml_data_to_tensor(Tensor(raw), out_features * in_features, 23).numpy().reshape(out_features, in_features) + llm_amd.iq4_half_lut.cache_clear() + layer = Linear(in_features, out_features, bias=False) + layer.set_quantized(Tensor(raw, dtype=dtypes.uint8, device="AMD").realize(), 23) + @TinyJit + def run(x:Tensor): return layer(x).realize() + x = rng.standard_normal((16, in_features), dtype=np.float32) + expected = x.astype(np.float16).astype(np.float32) @ weight.astype(np.float16).astype(np.float32).T + for _ in range(2): np.testing.assert_allclose(run(Tensor(x, device="AMD")).numpy(), expected, rtol=1e-5, atol=2e-3) + def test_packed_linear_offset_matches_reference(self): rng = np.random.default_rng(32) for ggml_type,in_features in ((8, 256), (12, 256), (13, 256), (14, 256), (23, 256)): diff --git a/tinygrad/llm/kernels/amd.py b/tinygrad/llm/kernels/amd.py index 6b53d00ab9..46448d38d9 100644 --- a/tinygrad/llm/kernels/amd.py +++ b/tinygrad/llm/kernels/amd.py @@ -7,7 +7,7 @@ from tinygrad.uop.ops import AxisType, KernelInfo, Ops from tinygrad.dtype import AddrSpace, dtypes from tinygrad.llm.gguf import _GGML_QUANT if TYPE_CHECKING: - from tinygrad.llm.model import ExpertWeights, Linear + from tinygrad.llm.model import Embedding, ExpertWeights, Linear BLOCK_M, BLOCK_N = 32, 32 DECODE_HEAD_TILE = 8 @@ -371,6 +371,26 @@ def _amd_dp4a(a:UOp, b:UOp, c:UOp) -> UOp: return UOp(Ops.CUSTOMI, dtypes.int32, (a.int(), b.int(), c), arg="__builtin_amdgcn_sudot4(true, {}, true, {}, {}, false)") +def _amd_byte_perm(a:UOp, b:UOp, selectors:UOp) -> UOp: + return UOp(Ops.CUSTOMI, dtypes.uint32, (a.cast(dtypes.uint32), b.cast(dtypes.uint32), selectors.cast(dtypes.uint32)), + arg="__builtin_amdgcn_perm({}, {}, {})") + +def _amd_vector_load(ptr:UOp, lanes:int) -> UOp: + assert ptr.op is Ops.INDEX + buf, coords = ptr.src[0], ptr.src[1:] + index = sum((coord*math.prod(buf.shape[i+1:]) for i,coord in enumerate(coords)), UOp.const(dtypes.weakint, 0)) + return UOp(Ops.SHRINK, src=(buf.flatten(), index, UOp.const(dtypes.weakint, lanes))).load(dtype=ptr.dtype) + +def _amd_stream_load(ptr:UOp) -> UOp: + assert ptr.op is Ops.INDEX + return UOp(Ops.CUSTOMI, ptr.dtype, (ptr,), arg="__builtin_nontemporal_load({0})") + +def _iq4_bytes(packed:UOp, shift:int) -> UOp: + selectors = (packed >> shift) & 0x0f0f0f0f + low = _amd_byte_perm(UOp.const(dtypes.uint32, 0xf6eaddcf), UOp.const(dtypes.uint32, 0xbfad9881), selectors) + high = _amd_byte_perm(UOp.const(dtypes.uint32, 0x71594535), UOp.const(dtypes.uint32, 0x26190d01), selectors & 0x07070707) + return _amd_byte_perm(high, low, 0x03020100 | ((selectors & 0x08080808) >> 1)) + def _amd_wave_sum(value:UOp, lane:UOp, lane_count:int, wave:UOp|None=None) -> UOp: assert lane_count in (8, 16, 32) for offset in (16, 8, 4, 2, 1)[{32:0, 16:1, 8:2}[lane_count]:]: @@ -379,24 +399,29 @@ def _amd_wave_sum(value:UOp, lane:UOp, lane_count:int, wave:UOp|None=None) -> UO arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_bpermute({0}, __builtin_bit_cast(int, {1})))") return value +def _amd_wave_max(value:UOp, lane:UOp) -> UOp: + for offset in (16, 8, 4, 2, 1): + value = value.maximum(UOp(Ops.CUSTOM, dtypes.float32, (((lane ^ offset) * 4).int(), value), + arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_bpermute({0}, __builtin_bit_cast(int, {1})))")) + return value + +def _q8_pack(raw_value:UOp, lane:UOp) -> tuple[UOp, UOp]: + d = (_amd_wave_max(raw_value.abs(), lane) / 127).maximum(1e-8) + quantized = (raw_value / d).round().maximum(-127).minimum(127).cast(dtypes.int8).bitcast(dtypes.uint8).cast(dtypes.uint32) + word = UOp.const(dtypes.uint32, 0) + for byte_idx in range(4): + byte = UOp(Ops.CUSTOM, dtypes.uint32, (((lane * 4 + byte_idx) * 4).int(), quantized), + arg="__builtin_amdgcn_ds_bpermute({0}, {1})") + word = word | (byte << (8 * byte_idx)) + return d, word + @functools.cache def _q8_kernel(quant:UOp, scale:UOp, x:UOp, in_features:int, group_sum:UOp|None=None) -> UOp: x = x.flatten() token, group = UOp.range(quant.shape[0], 0), UOp.range(in_features // 32, 1) lane = UOp.range(32, 2, axis_type=AxisType.LOCAL) raw_value = x[token * in_features + group * 32 + lane].load().float() - amax = raw_value.abs() - for offset in (16, 8, 4, 2, 1): - amax = amax.maximum(UOp(Ops.CUSTOM, dtypes.float32, (((lane ^ offset) * 4).int(), amax), - arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_bpermute({0}, __builtin_bit_cast(int, {1})))")) - d = (amax / 127).maximum(1e-8) - word = UOp.const(dtypes.uint32, 0) - for byte_idx in range(4): - source_lane = lane * 4 + byte_idx - value = (UOp(Ops.CUSTOM, dtypes.float32, ((source_lane * 4).int(), raw_value), - arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_bpermute({0}, __builtin_bit_cast(int, {1})))") / d).round().maximum(-127).minimum(127) - byte = value.cast(dtypes.int8).bitcast(dtypes.uint8).cast(dtypes.uint32) - word = word | (byte << (8 * byte_idx)) + d, word = _q8_pack(raw_value, lane) stores = [scale[token.valid(lane.eq(0)), group].store(d), quant[token, group, lane.valid(lane < 8)].store(word)] if group_sum is not None: qsum = _amd_dp4a(UOp.const(dtypes.uint32, 0x01010101), word, UOp.const(dtypes.int32, 0)) @@ -419,6 +444,207 @@ def q8_quantize_sum(x:Tensor, tokens:int, in_features:int) -> tuple[Tensor, Tens return tuple(Tensor.custom_kernel(quant, scale, group_sum, x, fxn=lambda quant,scale,group_sum,x:_q8_kernel(quant, scale, x, in_features, group_sum))[:3]) # type: ignore[return-value] +@functools.cache +def _gated_delta_prefill_kernel(core:UOp, next_state:UOp, q:UOp, k:UOp, v:UOp, beta:UOp, alpha:UOp, state:UOp, + kq:UOp) -> UOp: + batch, heads, tokens, dim = core.shape + row_tile = 4 + assert all(isinstance(x, int) for x in (batch, heads, tokens, dim)) and dim % 32 == 0 and dim % row_tile == 0 + batch, heads, tokens, dim = cast(tuple[int, int, int, int], (batch, heads, tokens, dim)) + core, q, k, v = (x.reshape(batch*heads, tokens, dim) for x in (core, q, k, v)) + beta, alpha, kq = (x.reshape(batch*heads, tokens) for x in (beta, alpha, kq)) + state, next_state = (x.reshape(batch*heads, dim, dim) for x in (state, next_state)) + bh_row, lane = UOp.range(batch*heads*dim//row_tile, 0), UOp.range(32, 1, axis_type=AxisType.LOCAL) + bh, row_base = bh_row // (dim//row_tile), (bh_row % (dim//row_tile))*row_tile + rows = tuple(row_base+i for i in range(row_tile)) + cols = tuple(lane + i*32 for i in range(dim//32)) + current = UOp.placeholder((row_tile*dim//32,), dtypes.float32, slot=0, addrspace=AddrSpace.REG) + current = current.after(current.store(UOp.stack(*(state[bh, row, col].float() for row in rows for col in cols)))) + token = UOp.range(tokens, 2, AxisType.REDUCE) + keys = tuple(k[bh, token, col].load() for col in cols) + queries = tuple(q[bh, token, col].load() for col in cols) + av, bv = alpha[bh, token].load(), beta[bh, token].load() + updates, stores = [], [] + for row_idx,row in enumerate(rows): + previous = tuple(current.after(token)[row_idx*dim//32+i].load() for i in range(dim//32)) + state_k = _amd_wave_sum(sum((x*y for x,y in zip(previous, keys)), UOp.const(dtypes.float32, 0)), lane, 32) + state_q = _amd_wave_sum(sum((x*y for x,y in zip(previous, queries)), UOp.const(dtypes.float32, 0)), lane, 32) + delta = (v[bh, token, row].load() - state_k*av) * bv + updates += [x*av + delta*y for x,y in zip(previous, keys)] + stores.append(core[bh, token, row.valid(lane.eq(0))].store(state_q*av + delta*kq[bh, token])) + step = UOp.group(*stores, current.store(UOp.stack(*updates))).end(token) + return UOp.group(*(next_state[bh, row, col].store(current.after(step)[row_idx*dim//32+i].load().cast(next_state.dtype)) + for row_idx,row in enumerate(rows) for i,col in enumerate(cols))).end(lane, bh_row).sink( + arg=KernelInfo(name="gated_delta_prefill", opts_to_apply=())) + +def gated_delta_prefill(q:Tensor, k:Tensor, v:Tensor, beta:Tensor, alpha:Tensor, state:Tensor) -> tuple[Tensor, Tensor]: + batch, heads, tokens, dim = q.shape + assert q.shape == k.shape == v.shape and beta.shape == alpha.shape == (batch, heads, tokens) and \ + state.shape == (batch, heads, dim, dim) + core, next_state = Tensor.empty_like(q), Tensor.empty_like(state) + kq = (q*k).sum(-1).contiguous() + srcs = (core.uop, next_state.uop, q.contiguous().uop, k.contiguous().uop, v.contiguous().uop, + beta.contiguous().uop, alpha.contiguous().uop, state.uop, kq.uop) + params = [UOp.placeholder_like(src, slot=i) for i,src in enumerate(srcs)] + out = _gated_delta_prefill_kernel(*params).call(*srcs) + return Tensor(srcs[0].after(out)), Tensor(srcs[1].after(out)) + +@functools.cache +def _gated_delta_decode_kernel(core:UOp, state:UOp, q:UOp, k:UOp, v:UOp, beta:UOp, alpha:UOp) -> UOp: + batch, heads, dim = core.shape + assert q.shape == k.shape == v.shape == core.shape and beta.shape == alpha.shape == (batch, heads) and \ + state.shape == (batch, heads, dim, dim) and dim % 32 == 0 + bh_row, lane = UOp.range(batch*heads*dim, 0), UOp.range(16, 1, AxisType.LOCAL) + bh, row = bh_row // dim, bh_row % dim + cols = tuple(lane*8 + i for i in range(8)) + state_values = tuple(state.reshape(batch*heads, dim, dim)[bh, row, col].float() for col in cols) + k_values = tuple(k.reshape(batch*heads, dim)[bh, col].float() for col in cols) + q_values = tuple(q.reshape(batch*heads, dim)[bh, col].float() for col in cols) + partials = UOp.placeholder((3, 16), dtypes.float32, slot=0, addrspace=AddrSpace.LOCAL) + partial_values = tuple(sum(values, UOp.const(dtypes.float32, 0)) for values in + ((s*kv for s,kv in zip(state_values, k_values)), (s*qv for s,qv in zip(state_values, q_values)), + (kv*qv for kv,qv in zip(k_values, q_values)))) + partials = partials.after(UOp.barrier(UOp.group(*(partials[i, lane].store(value) for i,value in enumerate(partial_values))))) + state_k, state_q, kq = (sum((partials[i, source_lane] for source_lane in range(16)), UOp.const(dtypes.float32, 0)) for i in range(3)) + av, bv = alpha.flatten()[bh].float(), beta.flatten()[bh].float() + delta = (v.reshape(batch*heads, dim)[bh, row].float() - state_k*av) * bv + stores = [state.reshape(batch*heads, dim, dim)[bh, row, col].store((s*av + delta*kv).cast(state.dtype)) + for col,s,kv in zip(cols, state_values, k_values)] + stores.append(core.reshape(batch*heads, dim)[bh, row.valid(lane.eq(0))].store(state_q*av + delta*kq)) + return UOp.group(*stores).end(lane, bh_row).sink(arg=KernelInfo(name="gated_delta_decode", opts_to_apply=())) + +def gated_delta_decode(q:Tensor, k:Tensor, v:Tensor, beta:Tensor, alpha:Tensor, state:Tensor) -> tuple[Tensor, Tensor]: + core = Tensor.empty_like(v) + srcs = (core, state, q, k, v, beta, alpha) + return tuple(Tensor.custom_kernel(*srcs, fxn=_gated_delta_decode_kernel)[:2]) # type: ignore[return-value] + +@functools.cache +def _q8_gated_rmsnorm_kernel(quant:UOp, scale:UOp, group_sum:UOp, core:UOp, gate:UOp, weight:UOp, eps:float) -> UOp: + batch, heads, dim = core.shape + assert gate.shape == (batch, 1, heads, dim) and weight.shape == (dim,) and dim % 32 == 0 + group, lane = UOp.range(batch*heads*dim//32, 0), UOp.range(32, 1, AxisType.LOCAL) + bh, subgroup = group // (dim//32), group % (dim//32) + head_values = tuple(core.reshape(batch*heads, dim)[bh, lane + i*32].float() for i in range(dim//32)) + sum_squared = _amd_wave_sum(sum((x*x for x in head_values), UOp.const(dtypes.float32, 0)), lane, 32) + idx = subgroup*32 + lane + normalized = core.reshape(batch*heads, dim)[bh, idx].float() * (sum_squared/dim + eps).rsqrt() * weight[idx].float() + gate_value = gate.reshape(batch*heads, dim)[bh, idx].float() + d, word = _q8_pack(normalized * gate_value * gate_value.sigmoid(), lane) + qsum = _amd_dp4a(UOp.const(dtypes.uint32, 0x01010101), word, UOp.const(dtypes.int32, 0)) + qsum = _amd_wave_sum(qsum.float(), lane, 8).cast(dtypes.int32) + stores = [scale.flatten()[group.valid(lane.eq(0))].store(d), + quant.reshape(batch*heads*dim//32, 8)[group, lane.valid(lane < 8)].store(word), + group_sum.flatten()[group.valid(lane.eq(0))].store(qsum)] + return UOp.group(*stores).end(lane, group).sink(arg=KernelInfo(name="q8_gated_rmsnorm", opts_to_apply=())) + +def q8_gated_rmsnorm(core:Tensor, gate:Tensor, weight:Tensor, eps:float) -> tuple[Tensor, Tensor, Tensor]: + batch, heads, dim = core.shape + quant = Tensor.empty(batch, heads*dim//32, 8, dtype=dtypes.uint32, device=core.device) + scale = Tensor.empty(batch, heads*dim//32, dtype=dtypes.float32, device=core.device) + group_sum = Tensor.empty(batch, heads*dim//32, dtype=dtypes.int32, device=core.device) + return tuple(Tensor.custom_kernel(quant, scale, group_sum, core, gate, weight, fxn=functools.partial( + _q8_gated_rmsnorm_kernel, eps=eps))[:3]) # type: ignore[return-value] + +@functools.cache +def _q8_rmsnorm_kernel(quant:UOp, scale:UOp, group_sum:UOp, x:UOp, weight:UOp, eps:float, + normalized_out:UOp|None=None, round_half:bool=False) -> UOp: + tokens, dim = x.shape + groups = dim//32 + group, lane = UOp.range(tokens*groups, 0), UOp.range(32, 1, AxisType.LOCAL) + token, token_group = group//groups, group%groups + acc = UOp.placeholder((1,), dtypes.float32, slot=0, addrspace=AddrSpace.REG) + acc = acc.after(acc.store(acc.const_like(0))) + chunk = UOp.range(groups, 2, AxisType.REDUCE) + value = x[token, chunk*32+lane].float() + update = acc.store(acc.after(chunk) + value*value).end(chunk) + sum_squared = _amd_wave_sum(acc.after(update)[0].load(), lane, 32) + idx = token_group*32+lane + normalized = x[token, idx].float() * (sum_squared/dim + eps).rsqrt() * weight[idx].float() + if round_half: normalized = UOp(Ops.CUSTOMI, dtypes.float32, (normalized,), arg="((float)((_Float16)({0})))") + d, word = _q8_pack(normalized, lane) + qsum = _amd_dp4a(UOp.const(dtypes.uint32, 0x01010101), word, UOp.const(dtypes.int32, 0)) + qsum = _amd_wave_sum(qsum.float(), lane, 8).cast(dtypes.int32) + stores = [scale[token.valid(lane.eq(0)), token_group].store(d), quant[token, token_group, lane.valid(lane < 8)].store(word), + group_sum[token.valid(lane.eq(0)), token_group].store(qsum)] + if normalized_out is not None: stores.append(normalized_out[token, idx].store(normalized)) + return UOp.group(*stores).end(lane, group).sink(arg=KernelInfo(name="q8_rmsnorm", opts_to_apply=())) + +@functools.cache +def _q8_rmsnorm_output_kernel(normalized:UOp, quant:UOp, scale:UOp, group_sum:UOp, x:UOp, weight:UOp, eps:float) -> UOp: + return _q8_rmsnorm_kernel(quant, scale, group_sum, x, weight, eps, normalized, round_half=True) + +def q8_rmsnorm(x:Tensor, weight:Tensor, eps:float) -> tuple[Tensor, Tensor, Tensor]: + flat, dim = x.reshape(-1, x.shape[-1]), x.shape[-1] + quant = Tensor.empty(flat.shape[0], dim//32, 8, dtype=dtypes.uint32, device=x.device) + scale = Tensor.empty(flat.shape[0], dim//32, dtype=dtypes.float32, device=x.device) + group_sum = Tensor.empty(flat.shape[0], dim//32, dtype=dtypes.int32, device=x.device) + return tuple(Tensor.custom_kernel(quant, scale, group_sum, flat, weight, + fxn=functools.partial(_q8_rmsnorm_kernel, eps=eps))[:3]) # type: ignore[return-value] + +def rmsnorm_q8(x:Tensor, weight:Tensor, eps:float) -> tuple[Tensor, tuple[Tensor, Tensor, Tensor]]: + flat, dim = x.reshape(-1, x.shape[-1]), x.shape[-1] + normalized = Tensor.empty_like(flat) + quant = Tensor.empty(flat.shape[0], dim//32, 8, dtype=dtypes.uint32, device=x.device) + scale = Tensor.empty(flat.shape[0], dim//32, dtype=dtypes.float32, device=x.device) + group_sum = Tensor.empty(flat.shape[0], dim//32, dtype=dtypes.int32, device=x.device) + normalized, quant, scale, group_sum = Tensor.custom_kernel(normalized, quant, scale, group_sum, flat, weight, + fxn=functools.partial(_q8_rmsnorm_output_kernel, eps=eps))[:4] + return normalized.reshape(x.shape), (quant, scale, group_sum) + +@functools.cache +def _f16_matvec_kernel(out:UOp, x:UOp, weight:UOp) -> UOp: + tokens, in_features = x.shape + out_features = out.shape[1] + waves = 8 + assert tokens == 1 and weight.shape == (out_features, in_features) and in_features % (waves*32) == 0 + output, lane, wave = UOp.range(out_features, 0), UOp.range(32, 1, AxisType.LOCAL), UOp.range(waves, 2, AxisType.LOCAL) + indices = tuple(wave*32 + lane + i*waves*32 for i in range(in_features//(waves*32))) + partial = sum((x[0, idx].float()*weight[output, idx].float() for idx in indices), UOp.const(dtypes.float32, 0)) + partial = _amd_wave_sum(partial, lane, 32, wave) + scratch = UOp.placeholder((waves,), dtypes.float32, slot=0, addrspace=AddrSpace.LOCAL) + scratch = scratch.after(UOp.barrier(scratch[wave.valid(lane.eq(0))].store(partial))) + total = sum((scratch[i] for i in range(waves)), UOp.const(dtypes.float32, 0)) + return out[0, output.valid((wave.eq(0)) & (lane.eq(0)))].store(total).end(lane, wave, output).sink( + arg=KernelInfo(name="f16_matvec", opts_to_apply=())) + +def f16_matvec(x:Tensor, weight:Tensor) -> Tensor: + out = Tensor.empty(1, weight.shape[0], dtype=dtypes.float32, device=x.device) + return Tensor.custom_kernel(out, x.reshape(1, -1), weight, fxn=_f16_matvec_kernel)[0] + +@functools.cache +def _q4_embedding_kernel(out:UOp, raw:UOp, idx:UOp, raw_offset:UOp, embed_size:int) -> UOp: + token_count = math.prod(idx.shape) + out, idx, raw_offset = out.reshape(token_count, embed_size), idx.flatten(), raw_offset.cast(dtypes.uint64) + token, group = UOp.range(token_count, 0), UOp.range(embed_size // 32, 1) + lane = UOp.range(32, 2, axis_type=AxisType.LOCAL) + block, subgroup = group // 8, group % 8 + type_words, row_words = _GGML_QUANT[12][1] // 4, embed_size // 256 * (_GGML_QUANT[12][1] // 4) + base = raw_offset + idx[token].cast(dtypes.uint64) * row_words + block * type_words + def load_byte(byte_offset:UOp) -> UOp: + return (raw[base + byte_offset // 4] >> ((byte_offset & 3) * 8).cast(dtypes.uint32)) & 255 + scale = (subgroup < 4).where(load_byte(4 + subgroup) & 63, + (load_byte(8 + subgroup) & 15) | ((load_byte(subgroup) >> 6) << 4)).float() + minimum = (subgroup < 4).where(load_byte(8 + subgroup) & 63, + (load_byte(8 + subgroup) >> 4) | ((load_byte(4 + subgroup) >> 6) << 4)).float() + packed = raw[base + 4 + (subgroup // 2) * 8 + lane // 4] + q = ((packed >> ((subgroup & 1) * 4).cast(dtypes.uint32)) >> ((lane % 4) * 8).cast(dtypes.uint32)) & 15 + scales = raw[base] + d = (scales & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float() + dmin = (scales >> 16).cast(dtypes.uint16).bitcast(dtypes.float16).float() + value = (q.float() * d * scale - dmin * minimum).cast(dtypes.float16) + return out[token, group*32+lane].store(value).end(token, group, lane).sink( + arg=KernelInfo(name="embedding_q4_k", opts_to_apply=())) + +def q4_embedding(layer:Embedding, idx:Tensor) -> Tensor: + if layer._raw_uop is None: layer._prepare_packed() + assert layer._raw_uop is not None and layer._raw_offset_uop is not None + out = Tensor.empty(*idx.shape, layer.embed_size, dtype=dtypes.float16, device=idx.device) + raw = layer._raw_uop.bitcast(dtypes.uint32) + srcs = (out.uop, raw, idx.contiguous().uop, layer._raw_offset_uop) + params = [UOp.placeholder_like(src, slot=i) for i,src in enumerate(srcs)] + kernel = _q4_embedding_kernel(params[0], params[1], params[2], params[3], layer.embed_size).call(*srcs) + return Tensor(srcs[0].after(kernel)) + @functools.cache def _q8_silu_mul_kernel(quant:UOp, scale:UOp, gate:UOp, up:UOp, in_features:int) -> UOp: gate, up = gate.flatten(), up.flatten() @@ -428,18 +654,7 @@ def _q8_silu_mul_kernel(quant:UOp, scale:UOp, gate:UOp, up:UOp, in_features:int) x = gate[idx].float() return x * x.sigmoid() * up[idx].float() raw_value = value(token * in_features + group * 32 + lane) - amax = raw_value.abs() - for offset in (16, 8, 4, 2, 1): - amax = amax.maximum(UOp(Ops.CUSTOM, dtypes.float32, (((lane ^ offset) * 4).int(), amax), - arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_bpermute({0}, __builtin_bit_cast(int, {1})))")) - d = (amax / 127).maximum(1e-8) - word = UOp.const(dtypes.uint32, 0) - for byte_idx in range(4): - source_lane = lane * 4 + byte_idx - packed_value = UOp(Ops.CUSTOM, dtypes.float32, ((source_lane * 4).int(), raw_value), - arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_bpermute({0}, __builtin_bit_cast(int, {1})))") - byte = (packed_value / d).round().maximum(-127).minimum(127).cast(dtypes.int8).bitcast(dtypes.uint8).cast(dtypes.uint32) - word = word | (byte << (8 * byte_idx)) + d, word = _q8_pack(raw_value, lane) stores = [scale[token.valid(lane.eq(0)), group].store(d), quant[token, group, lane.valid(lane < 8)].store(word)] return UOp.group(*stores).end(token, group, lane).sink(arg=KernelInfo(name="q8_silu_mul", opts_to_apply=())) @@ -485,7 +700,6 @@ def _q8_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, in_fea @functools.cache def _q8_linear_wmma_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, in_features:int, raw_offset:UOp) -> UOp: - raw = raw.replace(dtype=dtypes.uint32, src=(raw.src[0] * raw.dtype.itemsize // 4,), arg=replace(raw.arg, dtype=dtypes.uint32)) raw_offset = raw_offset.cast(dtypes.uint64) def load_word(byte_offset:UOp) -> UOp: word_index, half_aligned = byte_offset // 4, (byte_offset & 2).ne(0) @@ -539,84 +753,87 @@ def _q8_linear_wmma_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, i @functools.cache def _qk_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, xsum:UOp, out_features:int, in_features:int, ggml_type:int, raw_offset:int|UOp=0) -> UOp: - raw = raw.replace(dtype=dtypes.uint32, src=(raw.src[0] * raw.dtype.itemsize // 4,), arg=replace(raw.arg, dtype=dtypes.uint32)) if isinstance(raw_offset, UOp): raw_offset = raw_offset.cast(dtypes.uint64) - def load_word(byte_offset:UOp) -> UOp: return raw[byte_offset // 4] - def load_byte(byte_offset:UOp) -> UOp: return (load_word(byte_offset) >> ((byte_offset & 3) * 8).cast(dtypes.uint32)) & 255 + def load_byte(base:UOp, byte_offset:UOp) -> UOp: + return (raw[base + byte_offset // 4] >> ((byte_offset & 3) * 8).cast(dtypes.uint32)) & 255 output_tile = 1 output_block, lane = UOp.range(out_features // output_tile, 0), UOp.range(32, 1, axis_type=AxisType.LOCAL) outputs, group_count = tuple(output_block * output_tile + i for i in range(output_tile)), in_features // 32 - type_size, output_size = _GGML_QUANT[ggml_type][1], in_features // 256 * _GGML_QUANT[ggml_type][1] + type_words, output_words = _GGML_QUANT[ggml_type][1] // 4, in_features // 256 * _GGML_QUANT[ggml_type][1] // 4 def group_dot(group:UOp, output:UOp) -> UOp: block, subgroup = group // 8, group % 8 - base = raw_offset + output * output_size + block * type_size - qs_base = base + (48 if ggml_type == 13 else 16) + (subgroup // 2) * 32 + base = raw_offset + output * output_words + block * type_words + qs_base = base + (12 if ggml_type == 13 else 4) + (subgroup // 2) * 8 + xwords = _amd_vector_load(xq[0, group, 0], 8) dot = UOp.const(dtypes.int32, 0) for word_idx in range(8): - word = (load_word(qs_base + word_idx * 4) >> ((subgroup & 1) * 4).cast(dtypes.uint32)) & 0x0f0f0f0f + word = (raw[qs_base + word_idx] >> ((subgroup & 1) * 4).cast(dtypes.uint32)) & 0x0f0f0f0f if ggml_type == 13: - word = word | (((load_word(base + 16 + word_idx * 4) >> subgroup.cast(dtypes.uint32)) & 0x01010101) << 4) - dot = _amd_dp4a(word, xq[0, group, word_idx], dot) - scale = (subgroup < 4).where(load_byte(base + 4 + subgroup) & 63, - (load_byte(base + 8 + subgroup) & 15) | ((load_byte(base + subgroup) >> 6) << 4)) - minimum = (subgroup < 4).where(load_byte(base + 8 + subgroup) & 63, - (load_byte(base + 8 + subgroup) >> 4) | ((load_byte(base + 4 + subgroup) >> 6) << 4)) - scales = load_word(base) + word = word | (((raw[base + 4 + word_idx] >> subgroup.cast(dtypes.uint32)) & 0x01010101) << 4) + dot = _amd_dp4a(word, xwords[word_idx], dot) + scale = (subgroup < 4).where(load_byte(base, 4 + subgroup) & 63, + (load_byte(base, 8 + subgroup) & 15) | ((load_byte(base, subgroup) >> 6) << 4)) + minimum = (subgroup < 4).where(load_byte(base, 8 + subgroup) & 63, + (load_byte(base, 8 + subgroup) >> 4) | ((load_byte(base, 4 + subgroup) >> 6) << 4)) + scales = raw[base] dbits, dminbits = (scales & 0xffff).cast(dtypes.uint16), (scales >> 16).cast(dtypes.uint16) return (dot.float() * dbits.bitcast(dtypes.float16).float() * scale.float() - xsum[0, group].float() * dminbits.bitcast(dtypes.float16).float() * minimum.float()) * xd[0, group] - totals = [_amd_wave_sum(sum((group_dot((lane + offset).valid(lane + offset < group_count), output) - for offset in range(0, group_count, 32)), UOp.const(dtypes.float32, 0)), lane, 32) for output in outputs] + accs = tuple(UOp.placeholder((1,), dtypes.float32, slot=i, addrspace=AddrSpace.REG) for i in range(output_tile)) + accs = tuple(acc.after(acc.store(acc.const_like(0))) for acc in accs) + chunk = UOp.range((group_count + 31) // 32, 2, AxisType.REDUCE) + update = UOp.group(*(acc.store(acc.after(chunk) + group_dot((lane + chunk*32).valid(lane + chunk*32 < group_count), output)) + for acc,output in zip(accs, outputs))).end(chunk) + totals = [_amd_wave_sum(acc.after(update)[0].load(), lane, 32) for acc in accs] stores = [out[0, output.valid(lane.eq(0))].store(total.cast(out.dtype)) for output,total in zip(outputs, totals)] return UOp.group(*stores).end(output_block, lane).sink(arg=KernelInfo(name=f"linear_q{ggml_type}", opts_to_apply=())) @functools.cache def _qk_linear_f16_wmma_kernel(out:UOp, raw:UOp, x:UOp, out_features:int, in_features:int, ggml_type:int, raw_offset:UOp) -> UOp: - raw = raw.replace(dtype=dtypes.uint32, src=(raw.src[0] * raw.dtype.itemsize // 4,), arg=replace(raw.arg, dtype=dtypes.uint32)) x = x.reshape(out.shape[0], in_features) raw_offset = raw_offset.cast(dtypes.uint64) - def load_word(byte_offset:UOp) -> UOp: - word_index, half_aligned = byte_offset // 4, (byte_offset & 2).ne(0) - return half_aligned.where((raw[word_index] >> 16) | (raw[word_index + 1] << 16), raw[word_index]) - def load_byte(byte_offset:UOp) -> UOp: return (raw[byte_offset // 4] >> ((byte_offset & 3)*8).cast(dtypes.uint32)) & 255 + def load_byte(base:UOp, byte_offset:UOp) -> UOp: + return (raw[base + byte_offset // 4] >> ((byte_offset & 3)*8).cast(dtypes.uint32)) & 255 token_tile, output_tiles = (64, 1) if out_features <= 1024 and out.shape[0] % 64 == 0 else \ (64, 2) if out.shape[0] % 64 == 0 else (32 if out.shape[0] % 32 == 0 else 16, 2) - token_block, output_block = UOp.range(out.shape[0] // token_tile, 0), UOp.range(out_features // (16*output_tiles), 1) - lane = UOp.range(32, 2, axis_type=AxisType.LOCAL) + output_waves = 2 if out_features % (32*output_tiles) == 0 else 1 + token_block = UOp.range(out.shape[0] // token_tile, 0) + output_block = UOp.range(out_features // (16*output_tiles*output_waves), 1) + lane, wave = UOp.range(32, 2, axis_type=AxisType.LOCAL), UOp.range(output_waves, 3, axis_type=AxisType.LOCAL) hw_lane = UOp(Ops.CUSTOM, dtypes.int32, (lane.int(),), arg="__builtin_amdgcn_mbcnt_lo(-1, 0)").cast(dtypes.weakint) physical_col, physical_half = hw_lane % 16, hw_lane // 16 - outputs = tuple(output_block*(16*output_tiles) + output_tile*16 + physical_col for output_tile in range(output_tiles)) + outputs = tuple((output_block*output_waves+wave)*(16*output_tiles) + output_tile*16 + physical_col for output_tile in range(output_tiles)) input_tokens = tuple(token_block*token_tile + tile*16 + physical_col for tile in range(token_tile // 16)) tokens = tuple(tuple(token_block*token_tile + tile*16 + physical_half*8 + i for i in range(8)) for tile in range(token_tile // 16)) - group_count, type_size = in_features // 32, _GGML_QUANT[ggml_type][1] - output_size = in_features // 256 * type_size + group_count, type_words = in_features // 32, _GGML_QUANT[ggml_type][1] // 4 + output_words = in_features // 256 * type_words accs = tuple(tuple(UOp.placeholder((8,), dtypes.float32, slot=output_tile*(token_tile//16)+tile, addrspace=AddrSpace.REG) for tile in range(token_tile // 16)) for output_tile in range(output_tiles)) accs = tuple(tuple(acc.after(acc.store(acc.const_like(0))) for acc in output_accs) for output_accs in accs) - group = UOp.range(group_count, 3, AxisType.REDUCE) + group = UOp.range(group_count, 4, AxisType.REDUCE) block, subgroup = group // 8, group % 8 wmma_accs = [list(output_accs) for output_accs in accs] for half in range(2): afrags = tuple(UOp.stack(*(x[input_token, group*32 + half*16 + i].cast(dtypes.float16) for i in range(16))) for input_token in input_tokens) for output_tile,output in enumerate(outputs): - base = raw_offset + output*output_size + block*type_size - scale = (subgroup < 4).where(load_byte(base + 4 + subgroup) & 63, - (load_byte(base + 8 + subgroup) & 15) | ((load_byte(base + subgroup) >> 6) << 4)).float() - minimum = (subgroup < 4).where(load_byte(base + 8 + subgroup) & 63, - (load_byte(base + 8 + subgroup) >> 4) | ((load_byte(base + 4 + subgroup) >> 6) << 4)).float() - scales = load_word(base) + base = raw_offset + output*output_words + block*type_words + scale = (subgroup < 4).where(load_byte(base, 4 + subgroup) & 63, + (load_byte(base, 8 + subgroup) & 15) | ((load_byte(base, subgroup) >> 6) << 4)).float() + minimum = (subgroup < 4).where(load_byte(base, 8 + subgroup) & 63, + (load_byte(base, 8 + subgroup) >> 4) | ((load_byte(base, 4 + subgroup) >> 6) << 4)).float() + scales = raw[base] d = (scales & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float() dmin = (scales >> 16).cast(dtypes.uint16).bitcast(dtypes.float16).float() weight_scale, weight_min = d*scale, dmin*minimum - qs_base = base + (48 if ggml_type == 13 else 16) + (subgroup // 2)*32 + half*16 - qwords = [((load_word(qs_base + i*4) >> ((subgroup & 1)*4).cast(dtypes.uint32)) & 0x0f0f0f0f) for i in range(4)] + qs_base = base + (12 if ggml_type == 13 else 4) + (subgroup // 2)*8 + half*4 + qwords = [((raw[qs_base + i] >> ((subgroup & 1)*4).cast(dtypes.uint32)) & 0x0f0f0f0f) for i in range(4)] if ggml_type == 13: - qwords = [word | (((load_word(base + 16 + half*16 + i*4) >> subgroup.cast(dtypes.uint32)) & 0x01010101) << 4) + qwords = [word | (((raw[base + 4 + half*4 + i] >> subgroup.cast(dtypes.uint32)) & 0x01010101) << 4) for i,word in enumerate(qwords)] bfrag = UOp.stack(*(((word >> (byte*8) & 255).float()*weight_scale-weight_min).cast(dtypes.float16) for word in qwords for byte in range(4))) @@ -640,89 +857,108 @@ def _qk_linear_f16_wmma_kernel(out:UOp, raw:UOp, x:UOp, out_features:int, in_fea logical_values.append(output_values) stores = [out[token, output].store(value) for output,output_values in zip(outputs, logical_values) for tile_tokens,logical in zip(tokens, output_values) for token,value in zip(tile_tokens, logical)] - return UOp.group(*stores).end(token_block, output_block, lane).sink( + return UOp.group(*stores).end(token_block, output_block, lane, wave).sink( arg=KernelInfo(name=f"linear_q{ggml_type}_f16_wmma", opts_to_apply=())) @functools.cache -def _iq4_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, lut:UOp, out_features:int, in_features:int, +def _iq4_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, in_features:int, raw_offset:int|UOp=0) -> UOp: - raw = raw.replace(dtype=dtypes.uint32, src=(raw.src[0] * raw.dtype.itemsize // 4,), arg=replace(raw.arg, dtype=dtypes.uint32)) if isinstance(raw_offset, UOp): raw_offset = raw_offset.cast(dtypes.uint64) - def load_word(byte_offset:UOp) -> UOp: return raw[byte_offset // 4] - def load_byte(byte_offset:UOp) -> UOp: return (load_word(byte_offset) >> ((byte_offset & 3) * 8).cast(dtypes.uint32)) & 255 + def load_byte(base:UOp, byte_offset:UOp) -> UOp: + return (raw[base + byte_offset // 4] >> ((byte_offset & 3) * 8).cast(dtypes.uint32)) & 255 output_tile = 1 output_block, lane = UOp.range(out_features // output_tile, 0), UOp.range(32, 1, axis_type=AxisType.LOCAL) outputs, group_count = tuple(output_block * output_tile + i for i in range(output_tile)), in_features // 32 - type_size, output_size = _GGML_QUANT[23][1], in_features // 256 * _GGML_QUANT[23][1] + type_words, output_words = _GGML_QUANT[23][1] // 4, in_features // 256 * _GGML_QUANT[23][1] // 4 def group_dot(group:UOp, output:UOp) -> UOp: block, subgroup = group // 8, group % 8 - base = raw_offset + output * output_size + block * type_size + base = raw_offset + output * output_words + block * type_words + xwords = _amd_vector_load(xq[0, group, 0], 8) dot = UOp.const(dtypes.int32, 0) for word_idx in range(8): - packed, shift = load_word(base + 8 + subgroup*16 + (word_idx % 4)*4), 4 * (word_idx // 4) - nibbles = tuple((packed >> (8*i + shift)) & 15 for i in range(4)) - low = lut[(nibbles[0] | (nibbles[1] << 4)).cast(dtypes.weakint)].cast(dtypes.uint32) - high = lut[(nibbles[2] | (nibbles[3] << 4)).cast(dtypes.weakint)].cast(dtypes.uint32) - dot = _amd_dp4a(low | (high << 16), xq[0, group, word_idx], dot) - low_byte = load_byte(base + 4 + subgroup // 2) + packed, shift = _amd_stream_load(raw[base + 2 + subgroup*4 + word_idx % 4]), 4 * (word_idx // 4) + dot = _amd_dp4a(_iq4_bytes(packed, shift), xwords[word_idx], dot) + low_byte = load_byte(base, 4 + subgroup // 2) scale = ((low_byte >> (4*(subgroup % 2)).cast(dtypes.uint32)) & 15) | \ - ((((load_word(base) >> 16) >> (2*subgroup).cast(dtypes.uint32)) & 3) << 4) - d = (load_word(base) & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float() + ((((raw[base] >> 16) >> (2*subgroup).cast(dtypes.uint32)) & 3) << 4) + d = (raw[base] & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float() return dot.float() * xd[0, group] * d * (scale.cast(dtypes.uint8).bitcast(dtypes.int8)-32).float() - totals = [_amd_wave_sum(sum((group_dot((lane + offset).valid(lane + offset < group_count), output) - for offset in range(0, group_count, 32)), UOp.const(dtypes.float32, 0)), lane, 32) for output in outputs] + accs = tuple(UOp.placeholder((1,), dtypes.float32, slot=i, addrspace=AddrSpace.REG) for i in range(output_tile)) + accs = tuple(acc.after(acc.store(acc.const_like(0))) for acc in accs) + chunk = UOp.range((group_count + 31) // 32, 2, AxisType.REDUCE) + update = UOp.group(*(acc.store(acc.after(chunk) + group_dot((lane + chunk*32).valid(lane + chunk*32 < group_count), output)) + for acc,output in zip(accs, outputs))).end(chunk) + totals = [_amd_wave_sum(acc.after(update)[0].load(), lane, 32) for acc in accs] stores = [out[0, output.valid(lane.eq(0))].store(total.cast(out.dtype)) for output,total in zip(outputs, totals)] return UOp.group(*stores).end(output_block, lane).sink(arg=KernelInfo(name="linear_iq4_xs", opts_to_apply=())) @functools.cache def _iq4_linear_f16_wmma_kernel(out:UOp, raw:UOp, x:UOp, lut:UOp, out_features:int, in_features:int, raw_offset:UOp) -> UOp: - raw = raw.replace(dtype=dtypes.uint32, src=(raw.src[0] * raw.dtype.itemsize // 4,), arg=replace(raw.arg, dtype=dtypes.uint32)) x = x.reshape(out.shape[0], in_features) raw_offset = raw_offset.cast(dtypes.uint64) - def load_word(byte_offset:UOp) -> UOp: return raw[byte_offset // 4] - def load_byte(byte_offset:UOp) -> UOp: return (load_word(byte_offset) >> ((byte_offset & 3) * 8).cast(dtypes.uint32)) & 255 + def load_byte(base:UOp, byte_offset:UOp) -> UOp: + return (raw[base + byte_offset // 4] >> ((byte_offset & 3) * 8).cast(dtypes.uint32)) & 255 - def dequant(base:UOp, subgroup:UOp, half:int) -> tuple[UOp, ...]: - low_byte = load_byte(base + 4 + subgroup // 2) + def dequant(base:UOp, subgroup:UOp) -> tuple[tuple[UOp, ...], tuple[UOp, ...]]: + low_byte = load_byte(base, 4 + subgroup // 2) scale_bits = ((low_byte >> (4*(subgroup % 2)).cast(dtypes.uint32)) & 15) | \ - ((((load_word(base) >> 16) >> (2*subgroup).cast(dtypes.uint32)) & 3) << 4) + ((((raw[base] >> 16) >> (2*subgroup).cast(dtypes.uint32)) & 3) << 4) scale = (scale_bits.cast(dtypes.uint8).bitcast(dtypes.int8)-32).float() * \ - (load_word(base) & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float() - values = [] - for packed in (load_word(base + 8 + subgroup*16 + i*4) for i in range(4)): - nibbles = tuple((packed >> (8*i + 4*half)) & 15 for i in range(4)) - pairs = tuple(lut[(nibbles[i] | (nibbles[i+1] << 4)).cast(dtypes.weakint)] for i in (0, 2)) - values += [(((pair >> (i*16)) & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float()*scale).cast(dtypes.float16) - for pair in pairs for i in range(2)] - return tuple(values) + (raw[base] & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float() + if out_features <= 6144: + pairs = tuple(lut[((raw[base + 2 + subgroup*4 + word] >> (byte*8)) & 255).cast(dtypes.weakint)] + for word in range(4) for byte in range(4)) + return tuple(tuple((((pair >> (half*16)) & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float()*scale).cast(dtypes.float16) + for pair in pairs) for half in range(2)) # type: ignore[return-value] + halves = [] + for half in range(2): + values = [] + for packed in (raw[base + 2 + subgroup*4 + i] for i in range(4)): + nibbles = tuple((packed >> (8*i + 4*half)) & 15 for i in range(4)) + pairs = tuple(lut[(nibbles[i] | (nibbles[i+1] << 4)).cast(dtypes.weakint)] for i in (0, 2)) + values += [(((pair >> (i*16)) & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float()*scale).cast(dtypes.float16) + for pair in pairs for i in range(2)] + halves.append(tuple(values)) + return tuple(halves) # type: ignore[return-value] - token_tile = 32 if out_features <= 1024 and out.shape[0] % 32 == 0 else 64 if out_features <= 6144 and out.shape[0] % 64 == 0 else \ + token_tile = 32 if out_features <= 1024 and out.shape[0] % 32 == 0 else \ + 128 if out_features == 5120 and in_features > 8192 and out.shape[0] % 128 == 0 else \ + 64 if out_features <= 6144 and out.shape[0] % 64 == 0 else \ 128 if out.shape[0] % 128 == 0 else \ 64 if out.shape[0] % 64 == 0 else 32 if out.shape[0] % 32 == 0 else 16 output_tiles = 1 if out_features <= 1024 else 2 if out_features <= 6144 else 1 if out_features < 8192 else 2 - token_block, output_block = UOp.range(out.shape[0] // token_tile, 0), UOp.range(out_features // (16*output_tiles), 1) - lane = UOp.range(32, 2, axis_type=AxisType.LOCAL) + output_waves = 2 if out_features % (32*output_tiles) == 0 else 1 + assert out_features % (16*output_tiles*output_waves) == 0 + token_block = UOp.range(out.shape[0] // token_tile, 0) + output_block = UOp.range(out_features // (16*output_tiles*output_waves), 1) + lane, wave = UOp.range(32, 2, axis_type=AxisType.LOCAL), UOp.range(output_waves, 3, axis_type=AxisType.LOCAL) hw_lane = UOp(Ops.CUSTOM, dtypes.int32, (lane.int(),), arg="__builtin_amdgcn_mbcnt_lo(-1, 0)").cast(dtypes.weakint) + local_lut = UOp.placeholder((256,), dtypes.uint32, slot=32, addrspace=AddrSpace.LOCAL) + tid = wave*32+lane + lut_items = 256 // (32*output_waves) + lut_ready = UOp.group(*(local_lut[tid*lut_items+i].store(lut[tid*lut_items+i]) for i in range(lut_items))).barrier() + lut = local_lut.after(lut_ready) physical_col, physical_half = hw_lane % 16, hw_lane // 16 - outputs = tuple(output_block*(16*output_tiles) + output_tile*16 + physical_col for output_tile in range(output_tiles)) + outputs = tuple((output_block*output_waves+wave)*(16*output_tiles) + output_tile*16 + physical_col for output_tile in range(output_tiles)) input_tokens = tuple(token_block*token_tile + tile*16 + physical_col for tile in range(token_tile // 16)) tokens = tuple(tuple(token_block*token_tile + tile*16 + physical_half*8 + i for i in range(8)) for tile in range(token_tile // 16)) - type_size = _GGML_QUANT[23][1] - output_size = in_features // 256 * type_size + type_words = _GGML_QUANT[23][1] // 4 + output_words = in_features // 256 * type_words accs = tuple(tuple(UOp.placeholder((8,), dtypes.float32, slot=output_tile*(token_tile//16)+tile, addrspace=AddrSpace.REG) for tile in range(token_tile // 16)) for output_tile in range(output_tiles)) accs = tuple(tuple(acc.after(acc.store(acc.const_like(0))) for acc in output_accs) for output_accs in accs) - group = UOp.range(in_features // 32, 3, AxisType.REDUCE) + group = UOp.range(in_features // 32, 4, AxisType.REDUCE) block, subgroup = group // 8, group % 8 + weights = tuple(dequant(raw_offset + output*output_words + block*type_words, subgroup) for output in outputs) wmma_accs = [list(output_accs) for output_accs in accs] for half in range(2): afrags = tuple(UOp.stack(*(x[input_token, group*32 + half*16 + i].cast(dtypes.float16) for i in range(16))) for input_token in input_tokens) - for output_tile,output in enumerate(outputs): - bfrag = UOp.stack(*dequant(raw_offset + output*output_size + block*type_size, subgroup, half)) + for output_tile,weight in enumerate(weights): + bfrag = UOp.stack(*weight[half]) for tile,afrag in enumerate(afrags): previous = accs[output_tile][tile].after(group) if half == 0 else wmma_accs[output_tile][tile] wmma_accs[output_tile][tile] = UOp.wmma(afrag, bfrag, previous, (16, 16, 16), 'AMD', 32) @@ -740,19 +976,26 @@ def _iq4_linear_f16_wmma_kernel(out:UOp, raw:UOp, x:UOp, lut:UOp, out_features:i low.where(vals[3], swapped[7]), low.where(swapped[3], vals[7])) stores = [out[token, output].store(value) for output,output_accs in zip(outputs, accs) for tile_tokens,acc in zip(tokens, output_accs) for token,value in zip(tile_tokens, logical_values(acc))] - return UOp.group(*stores).end(token_block, output_block, lane).sink( + return UOp.group(*stores).end(token_block, output_block, lane, wave).sink( arg=KernelInfo(name="linear_iq4_xs_f16_wmma", opts_to_apply=())) @functools.cache def _q6_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, in_features:int, raw_offset:int|UOp=0) -> UOp: if isinstance(raw_offset, UOp): raw_offset = raw_offset.cast(dtypes.uint64) - output_tile = 2 + output_tile = 1 output_block, lane = UOp.range(out_features // output_tile, 0), UOp.range(32, 1, axis_type=AxisType.LOCAL) - outputs, group_count = tuple(output_block * output_tile + i for i in range(output_tile)), in_features // 32 - type_size, output_size = _GGML_QUANT[14][1], in_features // 256 * _GGML_QUANT[14][1] + outputs = tuple(output_block * output_tile + i for i in range(output_tile)) + group_count, type_size = in_features // 32, _GGML_QUANT[14][1] + output_size = in_features // 256 * type_size - def group_dot(group:UOp, output:UOp) -> UOp: - block, subgroup = group // 8, group % 8 + accs = tuple(UOp.placeholder((1,), dtypes.float32, slot=i, addrspace=AddrSpace.REG) for i in range(output_tile)) + accs = tuple(acc.after(acc.store(acc.const_like(0))) for acc in accs) + chunk = UOp.range((group_count + 31) // 32, 2, AxisType.REDUCE) + group = (lane + chunk * 32).valid(lane + chunk * 32 < group_count) + block, subgroup = group // 8, group % 8 + xwords = _amd_vector_load(xq[0, group, 0], 8) + updates = [] + for acc,output in zip(accs, outputs): base = raw_offset + output * output_size + block * type_size dots = [UOp.const(dtypes.int32, 0), UOp.const(dtypes.int32, 0)] for word_idx in range(8): @@ -766,13 +1009,13 @@ def _q6_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, in_fea high = (high_byte >> ((within // 32) * 2).cast(dtypes.uint8)) & 3 q = (low | (high << 4)).cast(dtypes.uint8).bitcast(dtypes.int8) - 32 word = word | (q.cast(dtypes.int8).bitcast(dtypes.uint8).cast(dtypes.uint32) << (8 * byte_idx)) - dots[word_idx // 4] = _amd_dp4a(word, xq[0, group, word_idx], dots[word_idx // 4]) + dots[word_idx // 4] = _amd_dp4a(word, xwords[word_idx], dots[word_idx // 4]) scales = [raw[base + 192 + subgroup * 2 + i].cast(dtypes.uint8).bitcast(dtypes.int8).float() for i in range(2)] dbits = raw[base + 208].cast(dtypes.uint16) | (raw[base + 209].cast(dtypes.uint16) << 8) - return (dots[0].float() * scales[0] + dots[1].float() * scales[1]) * xd[0, group] * dbits.bitcast(dtypes.float16).float() - - totals = [_amd_wave_sum(sum((group_dot((lane + offset).valid(lane + offset < group_count), output) - for offset in range(0, group_count, 32)), UOp.const(dtypes.float32, 0)), lane, 32) for output in outputs] + value = (dots[0].float() * scales[0] + dots[1].float() * scales[1]) * xd[0, group] * dbits.bitcast(dtypes.float16).float() + updates.append(acc.store(acc.after(chunk) + value)) + update = UOp.group(*updates).end(chunk) + totals = [_amd_wave_sum(acc.after(update)[0].load(), lane, 32) for acc in accs] stores = [out[0, output.valid(lane.eq(0))].store(total.cast(out.dtype)) for output,total in zip(outputs, totals)] return UOp.group(*stores).end(output_block, lane).sink(arg=KernelInfo(name="linear_q6", opts_to_apply=())) @@ -849,40 +1092,42 @@ def q8_linear(layer:Linear, x:Tensor, prepared:tuple[Tensor, ...]|None=None) -> if layer._raw_uop is None: layer._prepare_packed() assert layer._raw_uop is not None and layer._raw_offset_uop is not None if layer.ggml_type in (12, 13): + raw_words = layer._raw_uop.bitcast(dtypes.uint32) if tokens % 16 == 0 and layer.out_features % 16 == 0: - qk_srcs = (out.uop, layer._raw_uop, x.cast(dtypes.float16).contiguous().uop, layer._raw_offset_uop) + qk_srcs = (out.uop, raw_words, x.cast(dtypes.float16).contiguous().uop, layer._raw_offset_uop) params = [UOp.placeholder_like(src, slot=i) for i,src in enumerate(qk_srcs)] kernel = _qk_linear_f16_wmma_kernel(params[0], params[1], params[2], layer.out_features, - layer.in_features, layer.ggml_type, params[3][0]*4).call(*qk_srcs) + layer.in_features, layer.ggml_type, params[3][0]).call(*qk_srcs) out = Tensor(qk_srcs[0].after(kernel)).reshape(*x.shape[:-1], layer.out_features) return out if layer.bias is None else out + layer.bias - qk_decode_srcs = (out.uop, layer._raw_uop, xq.uop, xd.uop, xsum.uop, layer._raw_offset_uop) + qk_decode_srcs = (out.uop, raw_words, xq.uop, xd.uop, xsum.uop, layer._raw_offset_uop) params = [UOp.placeholder_like(src, slot=i) for i,src in enumerate(qk_decode_srcs)] kernel = _qk_linear_kernel(params[0], params[1], params[2], params[3], params[4], layer.out_features, - layer.in_features, layer.ggml_type, params[5][0] * 4).call(*qk_decode_srcs) + layer.in_features, layer.ggml_type, params[5][0]).call(*qk_decode_srcs) out = Tensor(qk_decode_srcs[0].after(kernel)).reshape(*x.shape[:-1], layer.out_features) return out if layer.bias is None else out + layer.bias if layer.ggml_type == 23: use_wmma = tokens % 16 == 0 and layer.out_features % 16 == 0 - lut = iq4_half_lut(str(x.device)) if use_wmma else expert_lut(str(x.device), 23) - iq4_srcs = ((out.uop, layer._raw_uop, x.cast(dtypes.float16).contiguous().uop, lut.uop, layer._raw_offset_uop) if use_wmma else - (out.uop, layer._raw_uop, xq.uop, xd.uop, lut.uop, layer._raw_offset_uop)) + raw_words = layer._raw_uop.bitcast(dtypes.uint32) + if use_wmma: + lut = iq4_half_lut(str(x.device)) + iq4_srcs = (out.uop, raw_words, x.cast(dtypes.float16).contiguous().uop, lut.uop, layer._raw_offset_uop) + else: iq4_srcs = (out.uop, raw_words, xq.uop, xd.uop, layer._raw_offset_uop) params = [UOp.placeholder_like(src, slot=i) for i,src in enumerate(iq4_srcs)] kernel = (_iq4_linear_f16_wmma_kernel(params[0], params[1], params[2], params[3], layer.out_features, - layer.in_features, params[4][0] * 4) - if len(iq4_srcs) == 5 else _iq4_linear_kernel(params[0], params[1], params[2], params[3], params[4], - layer.out_features, layer.in_features, params[5][0] * 4)).call(*iq4_srcs) + layer.in_features, params[4][0]) + if use_wmma else _iq4_linear_kernel(params[0], params[1], params[2], params[3], layer.out_features, + layer.in_features, params[4][0])).call(*iq4_srcs) out = Tensor(iq4_srcs[0].after(kernel)).reshape(*x.shape[:-1], layer.out_features) return out if layer.bias is None else out + layer.bias - srcs = (out.uop, layer._raw_uop, xq.uop, xd.uop, layer._raw_offset_uop) + raw = layer._raw_uop.bitcast(dtypes.uint32) if layer.ggml_type == 8 else layer._raw_uop + srcs = (out.uop, raw, xq.uop, xd.uop, layer._raw_offset_uop) params = [UOp.placeholder_like(src, slot=i) for i,src in enumerate(srcs)] if layer.ggml_type == 8: if tokens % 16 == 0 and layer.out_features % 16 == 0: kernel = _q8_linear_wmma_kernel(params[0], params[1], params[2], params[3], layer.out_features, layer.in_features, params[4][0] * 4).call(*srcs) else: - params[1] = params[1].replace(dtype=dtypes.uint32, src=(params[1].src[0] * layer._raw_uop.dtype.itemsize // 4,), - arg=replace(params[1].arg, dtype=dtypes.uint32)) kernel = _q8_linear_kernel(params[0], params[1], params[2], params[3], layer.out_features, layer.in_features, params[4][0]).call(*srcs) else: diff --git a/tinygrad/llm/model.py b/tinygrad/llm/model.py index c23a4f2c3d..008e012d2e 100644 --- a/tinygrad/llm/model.py +++ b/tinygrad/llm/model.py @@ -7,21 +7,18 @@ from tinygrad.llm.kernels import amd as llm_amd, cpu as llm_cpu, generic as llm_ from tinygrad.llm.gguf import get_ggml_quantization, ggml_data_to_tensor, gguf_load from tinygrad.uop.ops import resolve, Ops -class Linear(nn.Linear): - def __init__(self, in_features:int, out_features:int, bias=True): - # GGUF loading replaces every LLM weight. Lazy zeros avoid constructing hundreds of random-init graphs first, - # while keeping directly-created test models deterministic and valid. - self.weight = Tensor.zeros(out_features, in_features) - self.bias = Tensor.zeros(out_features) if bias else None - self.in_features, self.out_features = in_features, out_features +class PackedWeight: + weight:Tensor + def _init_packed(self): self.ggml_type:int|None = None - self.cpu_repacked:Tensor|None = None self._raw_uop:UOp|None = None self._raw_offset_uop:UOp|None = None self._raw_offset_words:int|None = None def set_quantized(self, packed:Tensor, ggml_type:int): self.weight, self.ggml_type = packed.flatten(), ggml_type self._raw_uop = self._raw_offset_uop = self._raw_offset_words = None + # IQ4 prefill uses a device LUT. Build it with the weights, before this linear can be captured by TinyJit. + if ggml_type == 23 and str(packed.device).startswith("AMD"): llm_amd.iq4_half_lut(str(packed.device)) def _packed_offset(self) -> Tensor: raw, raw_offset = self.weight.uop, 0 while raw.op in (Ops.BITCAST, Ops.RESHAPE): raw = raw.src[0] @@ -34,8 +31,19 @@ class Linear(nn.Linear): return Tensor([self._raw_offset_words], dtype=dtypes.uint64, device=self.weight.device) def _prepare_packed(self): self._raw_offset_uop = self._packed_offset().realize().uop - def prepare(self, x:Tensor) -> tuple[Tensor, ...]|None: - if self.ggml_type in (12, 13) and str(self.weight.device).startswith("AMD"): + +class Linear(nn.Linear, PackedWeight): + def __init__(self, in_features:int, out_features:int, bias=True): + # GGUF loading replaces every LLM weight. Lazy zeros avoid constructing hundreds of random-init graphs first, + # while keeping directly-created test models deterministic and valid. + self.weight = Tensor.zeros(out_features, in_features) + self.bias = Tensor.zeros(out_features) if bias else None + self.in_features, self.out_features = in_features, out_features + self._init_packed() + self.cpu_repacked:Tensor|None = None + def prepare(self, x:Tensor, with_sum:bool=False) -> tuple[Tensor, ...]|None: + if (with_sum or self.ggml_type in (12, 13)) and self.ggml_type in (8, 12, 13, 14, 23) and \ + str(self.weight.device).startswith("AMD"): return llm_amd.q8_quantize_sum(x, int(x.numel()) // self.in_features, self.in_features) return llm_amd.q8_quantize(x, int(x.numel()) // self.in_features, self.in_features) \ if self.ggml_type in (8, 14, 23) and str(self.weight.device).startswith("AMD") else None @@ -52,6 +60,15 @@ class Linear(nn.Linear): return x.linear(weight.transpose(), self.bias) return super().__call__(x) +class Embedding(nn.Embedding, PackedWeight): + def __init__(self, vocab_size:int, embed_size:int): + self.weight = Tensor.zeros(vocab_size, embed_size) + self.vocab_size, self.embed_size = vocab_size, embed_size + self._init_packed() + def __call__(self, idx:Tensor) -> Tensor: + if self.ggml_type == 12 and str(self.weight.device).startswith("AMD"): return llm_amd.q4_embedding(self, idx) + return super().__call__(idx) + @functools.cache def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0, device:str|None=None) -> Tensor: freqs = 1.0 / (theta ** (Tensor.arange(0, dim, 2).to(device)[:(dim // 2)] / dim)) @@ -275,20 +292,22 @@ class FFNBlock: if hasattr(self, 'ffn_gate_inp_shexp'): shexp = shexp * llm_cpu.shared_gate(x, self.ffn_gate_inp_shexp["weight"]) out = out + shexp return out - # TODO: remove the need for this contiguous - dense_prepared = self.ffn_gate.prepare(x) + dense_prepared = llm_amd.q8_rmsnorm(x, input_norm.weight, input_norm.eps) if input_norm is not None and \ + str(x.device).startswith("AMD") and input_norm.weight is not None else self.ffn_gate.prepare(x) if dense_prepared is not None and self.ffn_gate.ggml_type == self.ffn_up.ggml_type == 8: gate, up = llm_amd.q8_linear_pair(self.ffn_gate, self.ffn_up, x, dense_prepared) elif self.ffn_gate.ggml_type == self.ffn_up.ggml_type == 8 and str(x.device).startswith("CPU"): gate, up = (llm_cpu.q8_linear_pair if int(x.numel()) == self.config.dim else llm_cpu.q8_batched_pair)(self.ffn_gate, self.ffn_up, x) else: gate, up = self.ffn_gate(x, dense_prepared), self.ffn_up(x, dense_prepared) - return self.ffn_down(gate.silu().contiguous() * up) + return self.ffn_down(gate.silu() * up) def _normalized_feed_forward(self, x:Tensor) -> Tensor: - fuse_norm = hasattr(self, "ffn_gate_exps") and x.shape[-2] == 1 and str(x.device).startswith("CPU") and \ - x.dtype == dtypes.float32 and self.ffn_norm.weight is not None and self.ffn_norm.weight.dtype == dtypes.float16 and \ - self.ffn_gate_inp.ggml_type is None and self.ffn_gate_inp.bias is None and self.ffn_gate_inp.weight.dtype == dtypes.float16 - return self._feed_forward(x, self.ffn_norm) if fuse_norm else self._feed_forward(llm_cpu.rmsnorm(self.ffn_norm, x)) + amd_fuse = x.shape[-2] == 1 and str(x.device).startswith("AMD") and not hasattr(self, "ffn_gate_exps") and \ + self.ffn_norm.weight is not None and self.ffn_gate.ggml_type in (8, 12, 13, 14, 23) + cpu_fuse = hasattr(self, "ffn_gate_exps") and x.shape[-2] == 1 and str(x.device).startswith("CPU") and x.dtype == dtypes.float32 and \ + self.ffn_norm.weight is not None and self.ffn_norm.weight.dtype == dtypes.float16 and self.ffn_gate_inp.ggml_type is None and \ + self.ffn_gate_inp.bias is None and self.ffn_gate_inp.weight.dtype == dtypes.float16 + return self._feed_forward(x, self.ffn_norm) if amd_fuse or cpu_fuse else self._feed_forward(llm_cpu.rmsnorm(self.ffn_norm, x)) # 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 @@ -305,12 +324,13 @@ class FFNBlock: self.pending_recurrent_inplace = False @function(precompile=True, allow_implicit=True) def _run_stateful(x:Tensor, start_pos:int|UOp, valid_len:int|UOp|None): - fuse_attn_norm = x.shape[-2] == 1 and \ - str(x.device).startswith("CPU") and x.dtype == dtypes.float32 and \ + amd_fuse = x.shape[-2] == 1 and str(x.device).startswith("AMD") and self.attn_norm.weight is not None + cpu_fuse = x.shape[-2] == 1 and str(x.device).startswith("CPU") and x.dtype == dtypes.float32 and \ self.attn_norm.weight is not None and self.attn_norm.weight.dtype == dtypes.float16 and \ getattr(self, "attn_gate").ggml_type == getattr(self, "attn_qkv").ggml_type == 8 and \ getattr(self, "attn_gate").cpu_repacked is not None and getattr(self, "attn_qkv").cpu_repacked is not None and \ getattr(self, "ssm_beta_alpha_weight", None) is not None and getattr(self, "ssm_beta_alpha_weight").dtype == dtypes.float16 + fuse_attn_norm = amd_fuse or cpu_fuse attention_input = x if fuse_attn_norm else llm_cpu.rmsnorm(self.attn_norm, x) h = x + self._attention(attention_input, start_pos, use_flash, kv_len, valid_len, self.attn_norm if fuse_attn_norm else None) @@ -325,7 +345,9 @@ class FFNBlock: return Tensor(out.uop.after(state)) # we pass in the weights implicitly so we unpack the GGUF on the fly def _run(x:Tensor, start_pos:int|UOp): - h = x + self._attention(llm_cpu.rmsnorm(self.attn_norm, x), start_pos, use_flash, kv_len) + fuse_attn_norm = x.shape[-2] == 1 and str(x.device).startswith("AMD") and self.attn_norm.weight is not None + h = x + self._attention(x if fuse_attn_norm else llm_cpu.rmsnorm(self.attn_norm, x), start_pos, use_flash, kv_len, + input_norm=self.attn_norm if fuse_attn_norm else None) return (h + self._normalized_feed_forward(h)).contiguous() return function(precompile=True, allow_implicit=True)(_run)(x, start_pos) @@ -346,8 +368,11 @@ class TransformerBlock(FFNBlock): def _attention(self, x:Tensor, start_pos:int|UOp, use_flash:bool=False, kv_len:int|UOp|None=None, valid_len:int|UOp|None=None, input_norm:nn.RMSNorm|None=None) -> Tensor: - assert input_norm is None - prepared = self.attn_q.prepare(x) + prepared:tuple[Tensor, ...]|None + if input_norm is not None: + assert str(x.device).startswith("AMD") and input_norm.weight is not None + prepared = llm_amd.q8_rmsnorm(x, input_norm.weight, input_norm.eps) + else: prepared = self.attn_q.prepare(x, any(layer.ggml_type in (12, 13) for layer in (self.attn_q, self.attn_k, self.attn_v))) q = self.attn_q(x, prepared) if prepared is not None and self.attn_k.ggml_type == self.attn_v.ggml_type == 8: k, v = llm_amd.q8_linear_pair(self.attn_k, self.attn_v, x, prepared) @@ -509,10 +534,15 @@ class GatedDeltaNetBlock(FFNBlock): valid_len:int|UOp|None=None, input_norm:nn.RMSNorm|None=None) -> Tensor: B, T, _ = x.shape conv_state, initial_state = self.conv_state, self.recurrent_state + prepared:tuple[Tensor, ...]|None if T == 1: if input_norm is None: x = x.half() - prepared = None if input_norm is not None else self.attn_gate.prepare(x) + if input_norm is not None and str(x.device).startswith("AMD"): + assert input_norm.weight is not None + x, prepared = llm_amd.rmsnorm_q8(x, input_norm.weight, input_norm.eps) + input_norm = None + else: prepared = None if input_norm is not None else self.attn_gate.prepare(x, self.attn_qkv.ggml_type in (12, 13)) beta_alpha = None if input_norm is not None: assert self.ssm_beta_alpha_weight is not None @@ -527,7 +557,10 @@ class GatedDeltaNetBlock(FFNBlock): out_gate, qkv = llm_cpu.q8_linear_pair(self.attn_gate, self.attn_qkv, x) else: out_gate, qkv = self.attn_gate(x, prepared), self.attn_qkv(x, prepared) if beta_alpha is not None: beta, alpha = beta_alpha.split(self.num_v_heads, dim=-1) - elif self.ssm_beta_alpha_weight is not None: beta, alpha = (x @ self.ssm_beta_alpha_weight.T).split(self.num_v_heads, dim=-1) + elif self.ssm_beta_alpha_weight is not None: + beta_alpha = llm_amd.f16_matvec(x, self.ssm_beta_alpha_weight) if str(x.device).startswith("AMD") else \ + x @ self.ssm_beta_alpha_weight.T + beta, alpha = beta_alpha.reshape(B, 1, -1).split(self.num_v_heads, dim=-1) elif str(x.device).startswith("CPU") and self.ssm_beta.ggml_type == self.ssm_alpha.ggml_type == 8: beta, alpha = llm_cpu.q8_linear_pair(self.ssm_beta, self.ssm_alpha, x) else: beta, alpha = self.ssm_beta(x, prepared), self.ssm_alpha(x, prepared) @@ -564,19 +597,28 @@ class GatedDeltaNetBlock(FFNBlock): reshape(B, 1, -1).cast(x.dtype)) beta, alpha = beta.reshape(B, self.num_v_heads, 1, 1), alpha.reshape(B, self.num_v_heads, 1, 1) q, k, v = q.unsqueeze(-1), k.unsqueeze(-1), v.unsqueeze(-1) - state_dots = initial_state @ k.cat(q, dim=-1) - state_k, state_q = state_dots[..., :1] * alpha, state_dots[..., 1:] * alpha - delta = (v - state_k) * beta - recurrent_state = initial_state * alpha + delta @ k.transpose(-1, -2) + if str(x.device).startswith("AMD"): + core, recurrent_state = llm_amd.gated_delta_decode(q.squeeze(-1), k.squeeze(-1), v.squeeze(-1), + beta.squeeze(-1).squeeze(-1), alpha.squeeze(-1).squeeze(-1), initial_state) + self.pending_recurrent_inplace = True + core = core.unsqueeze(-1) + else: + state_dots = initial_state @ k.cat(q, dim=-1) + state_k, state_q = state_dots[..., :1] * alpha, state_dots[..., 1:] * alpha + delta = (v - state_k) * beta + recurrent_state = initial_state * alpha + delta @ k.transpose(-1, -2) + core = state_q + delta * (k.transpose(-1, -2) @ q) self.pending_state = (conv_window[:, 1:, :].cast(self.conv_state.dtype).contiguous(), recurrent_state.cast(self.recurrent_state.dtype).contiguous()) - core = state_q + delta * (k.transpose(-1, -2) @ q) + if str(x.device).startswith("AMD") and self.ssm_out.ggml_type in (12, 13) and self.ssm_norm.weight is not None: + prepared_out = llm_amd.q8_gated_rmsnorm(core.squeeze(-1), out_gate, self.ssm_norm.weight, self.ssm_norm.eps) + return self.ssm_out(core.reshape(B, 1, -1), prepared_out) core_attn_out = llm_cpu.rmsnorm(self.ssm_norm, core.squeeze(-1).reshape(B, 1, self.num_v_heads, self.head_v_dim)) return self.ssm_out((core_attn_out * out_gate.silu()).reshape(B, 1, -1).cast(x.dtype)) # Batched projections and causal depthwise convolution. x = x.half() - prepared = self.attn_gate.prepare(x) + prepared = self.attn_gate.prepare(x, self.attn_qkv.ggml_type in (12, 13)) if prepared is not None and self.attn_gate.ggml_type == self.attn_qkv.ggml_type == 8: out_gate, qkv = llm_amd.q8_linear_pair(self.attn_gate, self.attn_qkv, x, prepared) elif str(x.device).startswith("CPU") and self.attn_gate.ggml_type == self.attn_qkv.ggml_type == 8: @@ -627,6 +669,14 @@ class GatedDeltaNetBlock(FFNBlock): self.pending_state = (conv_window[:, state_pos:state_pos+self.ssm_conv_kernel-1, :].cast(self.conv_state.dtype).contiguous(), recurrent_state) return out + if str(x.device).startswith("AMD"): + core_attn_out, recurrent_state = llm_amd.gated_delta_prefill(q, k, v, beta, log_alpha.exp(), initial_state) + core_attn_out = llm_cpu.rmsnorm(self.ssm_norm, core_attn_out.transpose(1, 2)) + out = self.ssm_out((core_attn_out * out_gate.silu()).reshape(B, T, -1).cast(x.dtype)).contiguous() + state_pos = T if valid_len is None else valid_len + self.pending_state = (conv_window[:, state_pos:state_pos+self.ssm_conv_kernel-1, :].cast(self.conv_state.dtype).contiguous(), + recurrent_state) + return out state = initial_state.transpose(-1, -2).float() core_chunks = [] for start in range(0, T, 64): @@ -675,7 +725,7 @@ class Transformer: # A full Q8 cache for 262k Qwen leaves no graph workspace on a 24 GB card. Keep two fixed layers host-mapped; # this avoids runtime cache growth while leaving the other attention layers resident in VRAM. for block in [block for block in self.blk if isinstance(block, TransformerBlock)][-2:]: block.kv_cache_host = True - self.token_embd = nn.Embedding(config.vocab_size, config.dim) + self.token_embd = Embedding(config.vocab_size, config.dim) self.output_norm = nn.RMSNorm(config.dim, config.norm_eps) self.output = Linear(config.dim, config.vocab_size, bias=False) self.max_context = config.max_context @@ -709,7 +759,7 @@ class Transformer: def forward_recurrent_decode(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor, decode_len:int, valid_len:int|UOp|None=None, sample:bool=False) -> Tensor: - return self.forward(tokens, start_pos, temperature, kv_len=decode_len, valid_len=valid_len, sample=sample) + return tokens.assign(self.forward(tokens, start_pos, temperature, kv_len=decode_len, valid_len=valid_len, sample=sample)) def __call__(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor, use_flash:bool=False, valid_len:int|UOp|None=None, sample:bool=False) -> Tensor: @@ -799,7 +849,7 @@ class Transformer: (use_physical_cores:=getattr(Device[str(load_device)], "use_physical_cores", None)) is not None: use_physical_cores() for param in nn.state.get_parameters(model): param.replace(param.to(load_device)) packed_weights:set[str] = set() - packed_linears:list[Linear] = [] + packed_layers:list[PackedWeight] = [] cpu_q8_weights:list[tuple[Linear, Tensor]] = [] cpu_iq3_weights:list[tuple[ExpertWeights, Tensor]] = [] packed_linear_types = (8, 12, 13, 14, 23) if str(load_device).startswith("AMD") else (8, 14) @@ -810,11 +860,14 @@ class Transformer: for name, weight in state_dict.items(): parts = name.split('.') quantization = get_ggml_quantization(weight) - if quantization is not None and quantization[1] in packed_linear_types and parts[-1] == "weight" and \ - isinstance(owner:=resolve_owner(parts[:-1]), Linear): + owner = resolve_owner(parts[:-1]) if parts[-1] == "weight" else None + packed = quantization is not None and (isinstance(owner, Linear) and quantization[1] in packed_linear_types or + isinstance(owner, Embedding) and quantization[1] == 12 and str(load_device).startswith("AMD")) + if packed: + assert quantization is not None and isinstance(owner, PackedWeight) owner.set_quantized(*quantization) - packed_linears.append(owner) - if quantization[1] == 8 and str(load_device).startswith("CPU") and getenv("CPU_Q8_REPACK", 1): + packed_layers.append(owner) + if isinstance(owner, Linear) and quantization[1] == 8 and str(load_device).startswith("CPU") and getenv("CPU_Q8_REPACK", 1): cpu_q8_weights.append((owner, llm_cpu.q8_repack(owner.weight, owner.out_features, owner.in_features))) state_dict[name], packed_weights = owner.weight, packed_weights | {name} elif len(parts) == 4 and parts[0] == "blk" and parts[2].endswith("_exps") and parts[3] == "weight" and quantization is not None: @@ -848,9 +901,9 @@ class Transformer: Tensor.realize(*recurrent_weights, *cpu_conv_weights, *cpu_router_weights, *(weight for _,weight in cpu_q8_weights), *(weight for _,weight in cpu_iq3_weights)) # Custom kernels need the shared GGUF buffer and byte offset before function tracing disables device access. - packed_offsets = [linear._packed_offset() for linear in packed_linears] + packed_offsets = [layer._packed_offset() for layer in packed_layers] if packed_offsets: Tensor.realize(*packed_offsets) - for linear,offset in zip(packed_linears, packed_offsets): linear._raw_offset_uop = offset.uop + for layer,offset in zip(packed_layers, packed_offsets): layer._raw_offset_uop = offset.uop expert_types = {getattr(block, name).ggml_type for block in model.blk if hasattr(block, "ffn_gate_exps") for name in ("ffn_gate_exps", "ffn_down_exps")} model_device = str(model.token_embd.weight.device) @@ -943,7 +996,7 @@ class Transformer: warm = self.generate([0] * warm_len, chunk_size=chunk_size) prefill_batch = getenv("PREFILL_JIT_BATCH_SIZE", 16 if str(device).startswith("CPU") else 128) with Context(JIT_BATCH_SIZE=prefill_batch): next(warm) - next(warm) + with Context(JIT_BATCH_SIZE=0): next(warm) self._warming_up = False else: for salt in range(2): next(self.generate([salt] + [0] * (warm_len - 1), chunk_size=chunk_size)) diff --git a/tinygrad/runtime/support/system.py b/tinygrad/runtime/support/system.py index 44e5606186..75ae90d19e 100644 --- a/tinygrad/runtime/support/system.py +++ b/tinygrad/runtime/support/system.py @@ -269,7 +269,8 @@ class PCIIfaceBase: if should_use_sysmem: vaddr = self.dev_impl.mm.alloc_vaddr(size:=round_up(size, mmap.PAGESIZE), align=mmap.PAGESIZE) memview, paddrs = self.pci_dev.alloc_sysmem(size, vaddr=vaddr, contiguous=contiguous) - mapping = self.dev_impl.mm.map_range(vaddr, size, [(paddr, 0x1000) for paddr in paddrs], aspace=AddrSpace.SYS, snooped=True, uncached=True) + mapping = self.dev_impl.mm.map_range(vaddr, size, [(paddr, 0x1000) for paddr in paddrs], aspace=AddrSpace.SYS, + snooped=True, uncached=uncached) return HCQBuffer(vaddr, size, meta=PCIAllocationMeta(mapping, has_cpu_mapping=True, hMemory=paddrs[0]), view=memview, owner=self.dev) mapping = self.dev_impl.mm.valloc(size:=round_up(size, 0x1000), uncached=uncached, contiguous=cpu_access)