diff --git a/test/backend/test_amd_flash_attention.py b/test/backend/test_amd_flash_attention.py index c7e9fa543c..43c1b28946 100644 --- a/test/backend/test_amd_flash_attention.py +++ b/test/backend/test_amd_flash_attention.py @@ -7,10 +7,10 @@ from tinygrad.llm.kernels.amd import amd_flash_attention_decode @unittest.skipUnless(Device.DEFAULT.startswith("AMD"), "AMD flash attention required") class TestAMDFlashAttention(unittest.TestCase): - def _test_decode(self, max_kv_len:int, valid_kv_len:int): + def _test_decode(self, max_kv_len:int, valid_kv_len:int, n_heads:int=16, n_kv_heads:int=2): rng = np.random.default_rng(1) - q_np = rng.standard_normal((1, 16, 1, 256)).astype(np.float16) - kv_np = rng.standard_normal((2, 1, 2, max_kv_len, 256)).astype(np.float16) + q_np = rng.standard_normal((1, n_heads, 1, 256)).astype(np.float16) + kv_np = rng.standard_normal((2, 1, n_kv_heads, max_kv_len, 256)).astype(np.float16) q, kv = Tensor(q_np).realize(), Tensor(kv_np).realize() @TinyJit @@ -21,17 +21,19 @@ class TestAMDFlashAttention(unittest.TestCase): assert out is not None q_ref = q_np[0, :, 0].astype(np.float32) k_ref, v_ref = kv_np[:, 0, :, :valid_kv_len].astype(np.float32) - expected = np.empty((16, 256), dtype=np.float32) - for head in range(16): - scores = q_ref[head] @ k_ref[head // 8].T / np.sqrt(256) + expected = np.empty((n_heads, 256), dtype=np.float32) + for head in range(n_heads): + scores = q_ref[head] @ k_ref[head // (n_heads // n_kv_heads)].T / np.sqrt(256) probs = np.exp(scores - scores.max()) - expected[head] = probs @ v_ref[head // 8] / probs.sum() + expected[head] = probs @ v_ref[head // (n_heads // n_kv_heads)] / probs.sum() self.assertTrue(np.isfinite(out).all()) np.testing.assert_allclose(out[0, :, 0], expected, rtol=2e-3, atol=2e-3) def test_short_decode_is_finite_and_matches_reference(self): self._test_decode(8192, 25) + def test_six_query_heads_per_kv_head(self): self._test_decode(8192, 25, n_heads=12, n_kv_heads=2) + def test_hierarchical_decode_matches_reference(self): self._test_decode(16384, 4097) diff --git a/test/unit/test_llm_quant_cpu.py b/test/unit/test_llm_quant_cpu.py index df67ba06ca..dcaa5c37fb 100644 --- a/test/unit/test_llm_quant_cpu.py +++ b/test/unit/test_llm_quant_cpu.py @@ -34,6 +34,7 @@ def random_packed(rng:np.random.Generator, ggml_type:int, elements:int) -> np.nd blocks = rng.integers(0, 256, size=(elements // block_size, type_size), dtype=np.uint8) scales = rng.uniform(0.001, 0.02, size=len(blocks)).astype(np.float16).view(np.uint8).reshape(-1, 2) blocks[:, :2] = scales + if ggml_type in (12, 13): blocks[:, 2:4] = scales if ggml_type == 14: blocks[:, -2:] = scales return blocks.flatten() @@ -42,14 +43,17 @@ def random_packed(rng:np.random.Generator, ggml_type:int, elements:int) -> np.nd class TestLLMQuantAMD(unittest.TestCase): def test_packed_linear_offset_matches_reference(self): rng = np.random.default_rng(32) - for ggml_type,in_features in ((8, 256), (14, 256)): - raw, out_features = random_packed(rng, ggml_type, 64 * in_features), 64 - weight = ggml_data_to_tensor(Tensor(raw), out_features * in_features, ggml_type).numpy().reshape(out_features, in_features) - storage = Tensor(np.concatenate((np.zeros(68, dtype=np.uint8), raw)), dtype=dtypes.uint8, device="AMD").realize() - layer = Linear(in_features, out_features, bias=False) - layer.set_quantized(storage[68:], ggml_type) - x = rng.standard_normal((1, in_features), dtype=np.float32) - np.testing.assert_allclose(layer(Tensor(x, device="AMD")).numpy(), q8_activation(x) @ weight.T, rtol=1e-5, atol=5e-4) + for ggml_type,in_features in ((8, 256), (12, 256), (13, 256), (14, 256), (23, 256)): + for tokens in ((1, 16, 32) if ggml_type == 23 else (1, 16) if ggml_type in (12, 13, 14) else (1,)): + raw, out_features = random_packed(rng, ggml_type, 64 * in_features), 64 + weight = ggml_data_to_tensor(Tensor(raw), out_features * in_features, ggml_type).numpy().reshape(out_features, in_features) + storage = Tensor(np.concatenate((np.zeros(68, dtype=np.uint8), raw)), dtype=dtypes.uint8, device="AMD").realize() + layer = Linear(in_features, out_features, bias=False) + layer.set_quantized(storage[68:], ggml_type) + x = rng.standard_normal((tokens, in_features), dtype=np.float32) + expected = x.astype(np.float16).astype(np.float32) @ weight.astype(np.float16).astype(np.float32).T \ + if ggml_type in (12, 13, 23) and tokens > 1 else q8_activation(x) @ weight.T + np.testing.assert_allclose(layer(Tensor(x, device="AMD")).numpy(), expected, rtol=1e-5, atol=2e-3) def test_iq3_expert_prefill_and_decode_match_reference(self): rng = np.random.default_rng(31) diff --git a/tinygrad/llm/kernels/amd.py b/tinygrad/llm/kernels/amd.py index 758eb280ef..9e80be2646 100644 --- a/tinygrad/llm/kernels/amd.py +++ b/tinygrad/llm/kernels/amd.py @@ -55,14 +55,15 @@ def _amd_flash_attention_decode_partial(out:UOp, stats:UOp, q:UOp, cache_kv:UOp, # Each wave owns two GQA query heads while the workgroup shares one KV head. This keeps per-wave register pressure low # and lets the cache coalesce the identical KV stream instead of launching a second workgroup for the same KV head. heads_per_wave = 2 - decode_waves = DECODE_HEAD_TILE // heads_per_wave + head_tile = min(DECODE_HEAD_TILE, G) + assert G % head_tile == 0 and head_tile % heads_per_wave == 0 + decode_waves = head_tile // heads_per_wave decode_group = 4 if max_kv_len <= 8192 else 2 - assert G % DECODE_HEAD_TILE == 0 - block_bhkv = UOp.range(B*H_KV*(G//DECODE_HEAD_TILE), 0, AxisType.GLOBAL) + block_bhkv = UOp.range(B*H_KV*(G//head_tile), 0, AxisType.GLOBAL) block_n = UOp.range((valid_kv_len+CHUNK-1)//CHUNK, 1, AxisType.GLOBAL) lane, wave = UOp.range(WARP_SIZE, 2, AxisType.LOCAL), UOp.range(decode_waves, 3, AxisType.LOCAL) - head_group = block_bhkv % (G//DECODE_HEAD_TILE) - bhkv = block_bhkv // (G//DECODE_HEAD_TILE) + head_group = block_bhkv % (G//head_tile) + bhkv = block_bhkv // (G//head_tile) b, kv_head = bhkv // H_KV, bhkv % H_KV dims = tuple(lane + i*WARP_SIZE for i in range(DV)) @@ -79,7 +80,7 @@ def _amd_flash_attention_decode_partial(out:UOp, stats:UOp, q:UOp, cache_kv:UOp, vvals = tuple(tuple(cache_kv[1, b, kv_head, key, d].float() for d in dims) for key in keys) updates = [] for head in range(heads_per_wave): - q_head = kv_head*G + head_group*DECODE_HEAD_TILE + wave*heads_per_wave + head + q_head = kv_head*G + head_group*head_tile + wave*heads_per_wave + head scores = tuple(wave_reduce_sum(sum((q[b, q_head, 0, d].float()*k for d,k in zip(dims, key_kvals)), UOp.const(dtypes.float, 0)), lane + wave*WARP_SIZE) / math.sqrt(D) for key_kvals in kvals) prev_acc, prev_max, prev_sum = acc.after(offset)[head], row_max.after(offset)[head], row_sum.after(offset)[head] @@ -94,7 +95,7 @@ def _amd_flash_attention_decode_partial(out:UOp, stats:UOp, q:UOp, cache_kv:UOp, stores = [] for head in range(heads_per_wave): - q_head = kv_head*G + head_group*DECODE_HEAD_TILE + wave*heads_per_wave + head + q_head = kv_head*G + head_group*head_tile + wave*heads_per_wave + head stores += [out[b, q_head, block_n, d].store(acc[head, i]) for i,d in enumerate(dims)] stores += [stats[b, q_head.valid(lane.eq(0)), block_n, 0].store(row_max[head]), stats[b, q_head.valid(lane.eq(0)), block_n, 1].store(row_sum[head])] @@ -395,7 +396,7 @@ def _amd_wave_sum(value:UOp, lane:UOp, lane_count:int, wave:UOp|None=None) -> UO return value @functools.cache -def _q8_kernel(quant:UOp, scale:UOp, x:UOp, in_features:int) -> UOp: +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) @@ -413,6 +414,12 @@ def _q8_kernel(quant:UOp, scale:UOp, x:UOp, in_features:int) -> UOp: byte = value.cast(dtypes.int8).bitcast(dtypes.uint8).cast(dtypes.uint32) word = word | (byte << (8 * byte_idx)) 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)) + for offset in (4, 2, 1): + qsum = qsum + UOp(Ops.CUSTOM, dtypes.int32, (((lane ^ offset) * 4).int(), qsum), + arg="__builtin_amdgcn_ds_bpermute({0}, {1})") + stores.append(group_sum[token.valid(lane.eq(0)), group].store(qsum)) return UOp.group(*stores).end(token, group, lane).sink(arg=KernelInfo(name="q8_quantize", opts_to_apply=())) def q8_quantize(x:Tensor, tokens:int, in_features:int) -> tuple[Tensor, Tensor]: @@ -421,6 +428,13 @@ def q8_quantize(x:Tensor, tokens:int, in_features:int) -> tuple[Tensor, Tensor]: return tuple(Tensor.custom_kernel(quant, scale, x, fxn=lambda quant,scale,x:_q8_kernel(quant, scale, x, in_features))[:2]) # type: ignore[return-value] +def q8_quantize_sum(x:Tensor, tokens:int, in_features:int) -> tuple[Tensor, Tensor, Tensor]: + quant = Tensor.empty(tokens, in_features // 32, 8, dtype=dtypes.uint32, device=x.device) + scale = Tensor.empty(tokens, in_features // 32, dtype=dtypes.float32, device=x.device) + group_sum = Tensor.empty(tokens, in_features // 32, dtype=dtypes.int32, device=x.device) + 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 _q8_silu_mul_kernel(quant:UOp, scale:UOp, gate:UOp, up:UOp, in_features:int) -> UOp: gate, up = gate.flatten(), up.flatten() @@ -538,6 +552,203 @@ def _q8_linear_wmma_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, i stores = [out[token, output].store(acc.after(update)[i]) for acc,tile_tokens in zip(accs, tokens) for i,token in enumerate(tile_tokens)] return UOp.group(*stores).end(token_block, output_block, lane).sink(arg=KernelInfo(name="linear_q8_wmma", opts_to_apply=())) +@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 + 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] + + 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 + 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 + 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) + 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] + 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 + token_tile = 32 if out.shape[0] % 32 == 0 else 16 + token_block, output_block = UOp.range(out.shape[0] // token_tile, 0), UOp.range(out_features // 16, 1) + lane = UOp.range(32, 2, 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 + output = output_block*16 + physical_col + 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 + accs = tuple(UOp.placeholder((8,), dtypes.float32, slot=tile, addrspace=AddrSpace.REG) for tile in range(token_tile // 16)) + accs = tuple(acc.after(acc.store(acc.const_like(0))) for acc in accs) + group = UOp.range(group_count, 3, AxisType.REDUCE) + block, subgroup = group // 8, group % 8 + 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) + d = (scales & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float() + dmin = (scales >> 16).cast(dtypes.uint16).bitcast(dtypes.float16).float() + wmma_accs = list(accs) + for half in range(2): + 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)] + if ggml_type == 13: + qwords = [word | (((load_word(base + 16 + half*16 + i*4) >> subgroup.cast(dtypes.uint32)) & 0x01010101) << 4) + for i,word in enumerate(qwords)] + bfrag = UOp.stack(*(((word >> (byte*8) & 255).float()*d*scale-dmin*minimum).cast(dtypes.float16) + for word in qwords for byte in range(4))) + for tile,input_token in enumerate(input_tokens): + afrag = UOp.stack(*(x[input_token, group*32 + half*16 + i].cast(dtypes.float16) for i in range(16))) + previous = accs[tile].after(group) if half == 0 else wmma_accs[tile] + wmma_accs[tile] = UOp.wmma(afrag, bfrag, previous, (16, 16, 16), 'AMD', 32) + update = UOp.group(*(acc.store(value) for acc,value in zip(accs, wmma_accs))).end(group) + logical_values = [] + for acc in accs: + vals = tuple(acc.after(update)[i].load() for i in range(8)) + swapped = tuple(UOp(Ops.CUSTOM, dtypes.float32, (value,), + arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_swizzle(__builtin_bit_cast(int, {0}), 50688))") for value in vals) + low = physical_half.eq(0) + logical_values.append((low.where(vals[0], swapped[4]), low.where(swapped[0], vals[4]), + low.where(vals[1], swapped[5]), low.where(swapped[1], vals[5]), + low.where(vals[2], swapped[6]), low.where(swapped[2], vals[6]), + low.where(vals[3], swapped[7]), low.where(swapped[3], vals[7]))) + stores = [out[token, output].store(value) for tile_tokens,logical in zip(tokens, logical_values) + for token,value in zip(tile_tokens, logical)] + return UOp.group(*stores).end(token_block, output_block, lane).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, + 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 + 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] + + def group_dot(group:UOp, output:UOp) -> UOp: + block, subgroup = group // 8, group % 8 + base = raw_offset + output * output_size + block * type_size + 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) + 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() + 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] + 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 dequant(base:UOp, subgroup:UOp, half:int) -> 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) + 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)) + 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) + word = low | (high << 16) + values += [(((word >> (byte*8)) & 255).cast(dtypes.uint8).bitcast(dtypes.int8).float()*scale).cast(dtypes.float16) + for byte in range(4)] + return tuple(values) + + token_tile, output_tiles = (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) + 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)) + 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 + 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) + 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): + bfrag = UOp.stack(*dequant(raw_offset + output*output_size + block*type_size, subgroup, 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) + update = UOp.group(*(acc.store(value) for output_accs,output_values in zip(accs, wmma_accs) + for acc,value in zip(output_accs, output_values))).end(group) + + def logical_values(acc:UOp) -> tuple[UOp, ...]: + vals = tuple(acc.after(update)[i].load() for i in range(8)) + swapped = tuple(UOp(Ops.CUSTOM, dtypes.float32, (value,), + arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_swizzle(__builtin_bit_cast(int, {0}), 50688))") for value in vals) + low = physical_half.eq(0) + return (low.where(vals[0], swapped[4]), low.where(swapped[0], vals[4]), + low.where(vals[1], swapped[5]), low.where(swapped[1], vals[5]), + low.where(vals[2], swapped[6]), low.where(swapped[2], vals[6]), + 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( + 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) @@ -571,12 +782,104 @@ def _q6_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, in_fea 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=())) -def q8_linear(layer:Linear, x:Tensor, prepared:tuple[Tensor, Tensor]|None=None) -> Tensor: +@functools.cache +def _q6_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) + 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 + token_tile = 32 if out.shape[0] % 32 == 0 else 16 + token_block, output_block = UOp.range(out.shape[0] // token_tile, 0), UOp.range(out_features // 16, 1) + lane = UOp.range(32, 2, 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 + output = output_block * 16 + physical_col + 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[14][1] + output_size = in_features // 256 * type_size + accs = tuple(UOp.placeholder((8,), dtypes.float32, slot=tile, addrspace=AddrSpace.REG) for tile in range(token_tile // 16)) + accs = tuple(acc.after(acc.store(acc.const_like(0))) for acc in accs) + group = UOp.range(group_count, 3, AxisType.REDUCE) + block, subgroup = group // 8, group % 8 + base = raw_offset + output * output_size + block * type_size + half_accs:list[list[UOp]] = [] + for half in range(2): + kbase = half * 16 + physical_half * 8 + ql_base = base + (subgroup // 4) * 64 + (subgroup % 2) * 32 + kbase + qh_base = base + 128 + (subgroup // 4) * 32 + kbase + ql_words, qh_words = (tuple(load_word(addr + i*4) for i in range(2)) for addr in (ql_base, qh_base)) + values = tuple(((((ql >> (byte*8)) >> (((subgroup//2) & 1)*4).cast(dtypes.uint32)) & 15) | + ((((qh >> (byte*8)) >> ((subgroup & 3)*2).cast(dtypes.uint32)) & 3) << 4)).cast(dtypes.int8) - 32 + for ql,qh in zip(ql_words, qh_words) for byte in range(4)) + swapped = tuple(UOp(Ops.CUSTOM, dtypes.int32, (value.int(),), arg="__builtin_amdgcn_ds_swizzle({0}, 16415)") for value in values) + bfrag = UOp.stack(*(value.cast(dtypes.int8) for value in (*values, *swapped))) + tile_accs = [] + for input_token in input_tokens: + awords = tuple(xq[input_token, group, kbase // 4 + i].load() for i in range(2)) + swapped_words = tuple(UOp(Ops.CUSTOM, dtypes.uint32, (word,), arg="__builtin_amdgcn_ds_swizzle({0}, 16415)") for word in awords) + afrag = UOp.stack(*(((word >> (byte*8)) & 255).cast(dtypes.uint8).bitcast(dtypes.int8) + for word in (*awords, *swapped_words) for byte in range(4))) + tile_accs.append(UOp.wmma(afrag, bfrag, UOp.const(dtypes.int32, 0).broadcast(8), (16, 16, 16), 'AMD', 32)) + half_accs.append(tile_accs) + + def logical_values(raw_acc:UOp) -> tuple[UOp, ...]: + vals = tuple(raw_acc[i] for i in range(8)) + swapped = tuple(UOp(Ops.CUSTOM, dtypes.int32, (value,), arg="__builtin_amdgcn_ds_swizzle({0}, 50688)") for value in vals) + low = physical_half.eq(0) + return (low.where(vals[0], swapped[4]), low.where(swapped[0], vals[4]), + low.where(vals[1], swapped[5]), low.where(swapped[1], vals[5]), + low.where(vals[2], swapped[6]), low.where(swapped[2], vals[6]), + low.where(vals[3], swapped[7]), low.where(swapped[3], vals[7])) + + logical = [[logical_values(half_accs[half][tile]) for half in range(2)] for tile in range(len(accs))] + scales = tuple(load_byte(base + 192 + subgroup*2 + half).cast(dtypes.uint8).bitcast(dtypes.int8).float() for half in range(2)) + d = (load_word(base + 208) & 0xffff).cast(dtypes.uint16).bitcast(dtypes.float16).float() + update = UOp.group(*(acc.after(group)[i].store(acc.after(group)[i] + + sum((logical[tile][half][i].float() * scales[half] for half in range(2)), UOp.const(dtypes.float32, 0)) * d * xd[token, group]) + for tile,(acc,tile_tokens) in enumerate(zip(accs, tokens)) for i,token in enumerate(tile_tokens))).end(group) + stores = [out[token, output].store(acc.after(update)[i]) for acc,tile_tokens in zip(accs, tokens) for i,token in enumerate(tile_tokens)] + return UOp.group(*stores).end(token_block, output_block, lane).sink(arg=KernelInfo(name="linear_q6_wmma", opts_to_apply=())) + +def q8_linear(layer:Linear, x:Tensor, prepared:tuple[Tensor, ...]|None=None) -> Tensor: tokens = int(x.numel()) // layer.in_features - xq, xd = prepared if prepared is not None else q8_quantize(x, tokens, layer.in_features) + if layer.ggml_type in (12, 13): + xq, xd, xsum = prepared if prepared is not None and len(prepared) == 3 else q8_quantize_sum(x, tokens, layer.in_features) + else: + xq, xd = prepared[:2] if prepared is not None else q8_quantize(x, tokens, layer.in_features) out = Tensor.empty(tokens, layer.out_features, dtype=dtypes.float32, device=x.device) 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): + if tokens % 16 == 0 and layer.out_features % 16 == 0: + qk_srcs = (out.uop, layer._raw_uop, x.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) + 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) + 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) + 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: + lut = expert_lut(str(x.device), 23) + iq4_srcs = ((out.uop, layer._raw_uop, x.contiguous().uop, lut.uop, layer._raw_offset_uop) + if tokens % 16 == 0 and layer.out_features % 16 == 0 else + (out.uop, layer._raw_uop, xq.uop, xd.uop, lut.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) + 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) params = [UOp.placeholder_like(src, slot=i) for i,src in enumerate(srcs)] if layer.ggml_type == 8: @@ -590,8 +893,10 @@ def q8_linear(layer:Linear, x:Tensor, prepared:tuple[Tensor, Tensor]|None=None) layer.in_features, params[4][0]).call(*srcs) else: assert layer.ggml_type == 14 - kernel = _q6_linear_kernel(params[0], params[1], params[2], params[3], layer.out_features, - layer.in_features, params[4][0] * 4).call(*srcs) + kernel = (_q6_linear_wmma_kernel(params[0], params[1], params[2], params[3], layer.out_features, + layer.in_features, params[4][0] * 4) if tokens % 16 == 0 and layer.out_features % 16 == 0 else + _q6_linear_kernel(params[0], params[1], params[2], params[3], layer.out_features, + layer.in_features, params[4][0] * 4)).call(*srcs) out = Tensor(srcs[0].after(kernel)).reshape(*x.shape[:-1], layer.out_features) return out if layer.bias is None else out + layer.bias @@ -625,14 +930,14 @@ def _q8_linear_pair_kernel(out0:UOp, out1:UOp, raw0:UOp, raw1:UOp, xq:UOp, xd:UO [out1[token.valid(lane.eq(0)), output1].store(total) for token,total in zip(tokens, totals1)] return UOp.group(*stores).end(token_block, output, lane).sink(arg=KernelInfo(name="linear_q8_pair", opts_to_apply=())) -def q8_linear_pair(first:Linear, second:Linear, x:Tensor, prepared:tuple[Tensor, Tensor]) -> tuple[Tensor, Tensor]: +def q8_linear_pair(first:Linear, second:Linear, x:Tensor, prepared:tuple[Tensor, ...]) -> tuple[Tensor, Tensor]: assert first.ggml_type == second.ggml_type == 8 and first.in_features == second.in_features if (tokens := int(x.numel()) // first.in_features) % 16 == 0: return first(x, prepared), second(x, prepared) if first._raw_uop is None: first._prepare_packed() if second._raw_uop is None: second._prepare_packed() assert first._raw_uop is not None and first._raw_offset_uop is not None assert second._raw_uop is not None and second._raw_offset_uop is not None - xq, xd = prepared + xq, xd = prepared[0], prepared[1] out0 = Tensor.empty(tokens, first.out_features, dtype=dtypes.float32, device=x.device) out1 = Tensor.empty(tokens, second.out_features, dtype=dtypes.float32, device=x.device) srcs = (out0.uop, out1.uop, first._raw_uop, second._raw_uop, xq.uop, xd.uop, first._raw_offset_uop, second._raw_offset_uop) diff --git a/tinygrad/llm/model.py b/tinygrad/llm/model.py index 6437a1011f..35c43ed8f6 100644 --- a/tinygrad/llm/model.py +++ b/tinygrad/llm/model.py @@ -33,11 +33,13 @@ 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, Tensor]|None: + def prepare(self, x:Tensor) -> tuple[Tensor, ...]|None: + if self.ggml_type in (12, 13) 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) and str(self.weight.device).startswith("AMD") else None - def __call__(self, x:Tensor, prepared:tuple[Tensor, Tensor]|None=None) -> Tensor: - if self.ggml_type in (8, 14) and str(self.weight.device).startswith("AMD") and (self.ggml_type == 8 or int(x.numel()) == self.in_features): + if self.ggml_type in (8, 14, 23) and str(self.weight.device).startswith("AMD") else None + def __call__(self, x:Tensor, prepared:tuple[Tensor, ...]|None=None) -> Tensor: + if self.ggml_type in (8, 12, 13, 14, 23) and str(self.weight.device).startswith("AMD"): return llm_amd.q8_linear(self, x, prepared) if llm_cpu.SUPPORTED and self.ggml_type in (8, 14) and str(self.weight.device).startswith("CPU") and \ x.dtype in (dtypes.float16, dtypes.float32): @@ -267,12 +269,12 @@ class FFNBlock: out = out + shexp return out # TODO: remove the need for this contiguous - prepared = self.ffn_gate.prepare(x) - if 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, prepared) + dense_prepared = 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, prepared), self.ffn_up(x, prepared) + else: gate, up = self.ffn_gate(x, dense_prepared), self.ffn_up(x, dense_prepared) return self.ffn_down(gate.silu().contiguous() * up) def _normalized_feed_forward(self, x:Tensor) -> Tensor: @@ -766,6 +768,7 @@ class Transformer: packed_linears:list[Linear] = [] 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) def resolve_owner(path:list[str]): obj = model for part in path: obj = obj[int(part)] if isinstance(obj, list) else getattr(obj, part) @@ -773,7 +776,8 @@ 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 (8, 14) and parts[-1] == "weight" and isinstance(owner:=resolve_owner(parts[:-1]), Linear): + 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.set_quantized(*quantization) packed_linears.append(owner) if quantization[1] == 8 and str(load_device).startswith("CPU") and getenv("CPU_Q8_REPACK", 1):