forked from tinygrad/tinygrad
/raid/models/Qwen3.6-35B-A3B-UD-IQ4_XS.gguf
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
+319
-14
@@ -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)
|
||||
|
||||
+13
-9
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user