llm: simplify AMD Qwen kernels

This commit is contained in:
2026-08-02 14:17:43 +00:00
parent fb53442c17
commit f92ac85fbd
2 changed files with 62 additions and 117 deletions
+55 -111
View File
@@ -8,18 +8,13 @@ from tinygrad.llm.gguf import _GGML_QUANT
if TYPE_CHECKING:
from tinygrad.llm.model import Embedding, Linear
BLOCK_M, BLOCK_N = 32, 32
DECODE_HEAD_TILE = 8
WARP_SIZE = 32
BLOCK_M, BLOCK_N, DECODE_HEAD_TILE, WARP_SIZE = 32, 32, 8, 32
WMMA_M, WMMA_N, WMMA_K = 16, 16, 16
WAVES_M, WAVES_N = 2, 2
LANES_PER_WAVE_M, LANES_PER_WAVE_N = 2, 16
WAVES_M, WAVES_N, LANES_PER_WAVE_M, LANES_PER_WAVE_N = 2, 2, 2, 16
WMMA_ACC = WMMA_M // LANES_PER_WAVE_M
THREADS_PER_BLOCK = WARP_SIZE * WAVES_M * WAVES_N
LDS_PAD = 4 # pad LDS rows to reduce bank conflicts
WMMA_ARG = (WMMA_M, WMMA_N, WMMA_K), 'AMD', 32
LOG2E = math.log2(math.e)
WMMA_ARG, LOG2E = ((WMMA_M, WMMA_N, WMMA_K), 'AMD', 32), math.log2(math.e)
def warp_shfl_xor(val, offset, lane):
"""Read val from lane ^ offset using ds_bpermute."""
@@ -28,20 +23,10 @@ def warp_shfl_xor(val, offset, lane):
return UOp(Ops.CUSTOM, dtypes.float, (idx, val),
arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_bpermute({0}, __builtin_bit_cast(int, {1})))")
def warp_reduce_max(val, lane):
"""Tree reduce MAX across LANES_PER_WAVE_N=16 lanes."""
for offset in [8, 4, 2, 1]:
val = val.maximum(warp_shfl_xor(val, offset, lane))
return val
def warp_reduce_sum(val, lane):
"""Tree reduce SUM across LANES_PER_WAVE_N=16 lanes."""
for offset in [8, 4, 2, 1]:
val = val + warp_shfl_xor(val, offset, lane)
return val
def wave_reduce_sum(val, lane):
for offset in [16, 8, 4, 2, 1]: val = val + warp_shfl_xor(val, offset, lane)
def warp_reduce(val:UOp, lane:UOp, maximum:bool=False, full_wave:bool=False) -> UOp:
for offset in ([16, 8, 4, 2, 1] if full_wave else [8, 4, 2, 1]):
other = warp_shfl_xor(val, offset, lane)
val = val.maximum(other) if maximum else val + other
return val
@functools.cache
@@ -67,13 +52,11 @@ def _amd_flash_attention_decode_partial(out:UOp, stats:UOp, q:UOp, cache_kv:UOp,
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))
acc = UOp.placeholder((heads_per_wave, DV), dtypes.float, slot=0, addrspace=AddrSpace.REG)
row_max = UOp.placeholder((heads_per_wave,), dtypes.float, slot=1, addrspace=AddrSpace.REG)
row_sum = UOp.placeholder((heads_per_wave,), dtypes.float, slot=2, addrspace=AddrSpace.REG)
init = UOp.group(acc.store(acc.const_like(0)), row_max.store(row_max.const_like(-math.inf)), row_sum.store(row_sum.const_like(0)))
acc, row_max, row_sum = acc.after(init), row_max.after(init), row_sum.after(init)
groups_per_chunk = CHUNK // decode_group
offset = UOp.range(((valid_chunks+group_count-1)//group_count)*groups_per_chunk, 100, AxisType.REDUCE)
chunk = block_n + (offset // groups_per_chunk) * group_count
@@ -86,8 +69,8 @@ def _amd_flash_attention_decode_partial(out:UOp, stats:UOp, q:UOp, cache_kv:UOp,
updates = []
for head in range(heads_per_wave):
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(0, dtypes.float)), lane + wave*WARP_SIZE) / math.sqrt(D) for key_kvals in kvals)
scores = tuple(warp_reduce(sum((q[b, q_head, 0, d].float()*k for d,k in zip(dims, key_kvals)),
UOp.const(0, dtypes.float)), lane + wave*WARP_SIZE, full_wave=True) / 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]
new_max = prev_max
for is_valid, score in zip(valid, scores): new_max = new_max.maximum(is_valid.where(score, UOp.const(-math.inf, dtypes.float)))
@@ -97,7 +80,6 @@ def _amd_flash_attention_decode_partial(out:UOp, stats:UOp, q:UOp, cache_kv:UOp,
row_sum[head].store(prev_sum*alpha + sum(betas, UOp.const(0, dtypes.float))), row_max[head].store(new_max)]
update = UOp.group(*updates).end(offset)
acc, row_max, row_sum = acc.after(update), row_max.after(update), row_sum.after(update)
stores = []
for head in range(heads_per_wave):
q_head = kv_head*G + head_group*head_tile + wave*heads_per_wave + head
@@ -114,13 +96,11 @@ def _amd_flash_attention_decode_reduce(out:UOp, partial:UOp, stats:UOp, valid_ch
block_bh, lane = UOp.range(B*H, 0, AxisType.GLOBAL), UOp.range(WARP_SIZE, 1, AxisType.LOCAL)
b, head = block_bh // H, block_bh % H
dims = tuple(lane + i*WARP_SIZE for i in range(DV))
row_max = UOp.placeholder((1,), dtypes.float, slot=0, addrspace=AddrSpace.REG)
row_max = row_max.after(row_max.store(row_max.const_like(-math.inf)))
chunk_max = UOp.range(valid_chunks, 100, AxisType.REDUCE)
max_done = row_max.store(row_max.after(chunk_max).maximum(stats[b, head, chunk_max, 0])).end(chunk_max)
row_max = row_max.after(max_done)
numerator = UOp.placeholder((DV,), dtypes.float, slot=1, addrspace=AddrSpace.REG)
denominator = UOp.placeholder((1,), dtypes.float, slot=2, addrspace=AddrSpace.REG)
init = UOp.group(numerator.store(numerator.const_like(0)), denominator.store(denominator.const_like(0)))
@@ -175,29 +155,24 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, kv_scale:UOp, valid_kv_len:
TN = BLOCK_N // LANES_PER_WAVE_N
TD = D // (WAVES_N * LANES_PER_WAVE_N)
SCALE = 1.0 / math.sqrt(D)
block_bh = UOp.range(BH, 0, AxisType.GLOBAL)
block_m = UOp.range(M // BLOCK_M, 1, AxisType.GLOBAL)
q = q.reshape(BH, M//BLOCK_M, BLOCK_M, D)[block_bh, block_m]
kv_head = block_bh // gqa_group
k, v = k[kv_head], v[kv_head]
o = o.reshape(BH, M//BLOCK_M, BLOCK_M, D)[block_bh, block_m]
wave_m = UOp.range(WAVES_M, 2, AxisType.LOCAL)
wave_n = UOp.range(WAVES_N, 3, AxisType.LOCAL)
lane = UOp.range(WARP_SIZE, -1, AxisType.WARP)
tid = (wave_m * WAVES_N + wave_n) * WARP_SIZE + lane
lane_m = lane // LANES_PER_WAVE_N
lane_n = lane % LANES_PER_WAVE_N
# LDS allocation: slot 0 = Q then P (shared), slot 1 = K then V
# TODO: the memory planner should be able to find this reuse
Q_ELEMS_PER_THREAD = BLOCK_M * D // THREADS_PER_BLOCK
KV_ELEMS_PER_THREAD = BLOCK_N * D // THREADS_PER_BLOCK
QP_lds = UOp.placeholder((BLOCK_M, D + LDS_PAD), dtypes.half, slot=0, addrspace=AddrSpace.LOCAL)
KV_lds = UOp.placeholder((BLOCK_N, D + LDS_PAD), dtypes.half, slot=1, addrspace=AddrSpace.LOCAL)[:, :D]
# register state
acc = UOp.placeholder((TM, TD), dtypes.float, slot=2, addrspace=AddrSpace.REG)
m_i = UOp.placeholder((TM,), dtypes.float, slot=3, addrspace=AddrSpace.REG)
@@ -205,13 +180,11 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, kv_scale:UOp, valid_kv_len:
acc = acc.after(acc.store(acc.const_like(0)))
m_i = m_i.after(m_i.store(m_i.const_like(-math.inf)))
l_i = l_i.after(l_i.store(l_i.const_like(0)))
# ====== KV tile loop ======
# Causal blocks never need KV tiles strictly to their right. Besides saving work, this avoids an all
# -inf tile, whose online-softmax update would otherwise contain -inf - -inf.
n_tiles = (N - M + (block_m + 1) * BLOCK_M + BLOCK_N - 1) // BLOCK_N
n_tile = UOp.range(n_tiles, 100, AxisType.REDUCE)
# load Q + K into LDS (Q reloaded each iteration since P overwrites slot 0)
Q_lds = QP_lds[:, :D]
Q_store = Q_lds.after(n_tile).reshape(THREADS_PER_BLOCK, Q_ELEMS_PER_THREAD)[tid].store(
@@ -224,7 +197,6 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, kv_scale:UOp, valid_kv_len:
qk_load_barrier = UOp.barrier(UOp.group(Q_store, K_store))
Q_lds = Q_lds.after(qk_load_barrier)
KV_lds_k = KV_lds.after(qk_load_barrier)
# -- S = Q @ K^T via WMMA (re-init each n_tile) --
S_reg = UOp.placeholder((TM, TN), dtypes.float, slot=6, addrspace=AddrSpace.REG)
S_reg = S_reg.after(S_reg.after(n_tile).store(S_reg.const_like(0)))
@@ -237,10 +209,8 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, kv_scale:UOp, valid_kv_len:
qk = UOp.wmma(q_frag, k_frag, S_frag.after(k_qk), *WMMA_ARG)
qk_done = S_frag.store(qk).end(tm1, tn1).end(k_qk)
S_reg = S_reg.after(qk_done)
# -- softmax in registers with warp shuffles --
S_reg = S_reg.after(S_reg.store(S_reg * SCALE))
# WMMA accumulator ownership: each lane owns an 8x4 fragment of the 64x64 score tile.
# q is aligned to the right of k, matching PyTorch's causal_lower_right mask.
rm, rn = UOp.range(TM, 250, AxisType.WEAK), UOp.range(TN, 251, AxisType.WEAK)
@@ -249,7 +219,6 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, kv_scale:UOp, valid_kv_len:
valid = k_idx <= q_idx
if key_limit is not None: valid = valid & (k_idx < key_limit)
S_reg = S_reg.after(S_reg[rm, rn].store(valid.where(S_reg[rm, rn], S_reg[rm, rn].const_like(-math.inf))).end(rm, rn))
# per-thread local row max over TN=4 elements, then warp reduce across 16 lanes
m_ij = UOp.placeholder((TM,), dtypes.float, slot=7, addrspace=AddrSpace.REG)
m_ij = m_ij.after(m_ij.after(n_tile).store(m_ij.const_like(-math.inf)))
@@ -257,24 +226,20 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, kv_scale:UOp, valid_kv_len:
m_ij = m_ij.after(m_ij.store(m_ij.after(rm2).maximum(S_reg[:, rm2])).end(rm2))
# warp reduce max (in-place)
ri_w = UOp.range(TM, 270)
m_ij = m_ij.after(m_ij[ri_w].store(warp_reduce_max(m_ij[ri_w], lane)).end(ri_w))
m_ij = m_ij.after(m_ij[ri_w].store(warp_reduce(m_ij[ri_w], lane, maximum=True)).end(ri_w))
# compute P = exp(S - m_ij) in S_reg
S_reg = S_reg.after(S_reg.store(((S_reg - m_ij.reshape(TM, 1).expand(TM, TN)) * LOG2E).exp2()))
p_local = UOp.placeholder((TM,), dtypes.float, slot=8, addrspace=AddrSpace.REG)
p_local = p_local.after(p_local.after(n_tile).store(p_local.const_like(0)))
ri_ws = UOp.range(TM, 295, AxisType.WEAK)
# Reduce contiguous 16-key groups independently, matching the ordinary softmax reduction tree.
p_sum = p_local.after(p_local[ri_ws].store(
sum((warp_reduce_sum(S_reg[ri_ws, rn], lane) for rn in range(TN)), S_reg.const_like(0))).end(ri_ws))
sum((warp_reduce(S_reg[ri_ws, rn], lane) for rn in range(TN)), S_reg.const_like(0))).end(ri_ws))
# Store softmax weights in half for the WMMA P@V product; accumulation remains float.
P_lds = QP_lds.flatten()[:WAVES_N * BLOCK_M * BLOCK_N].reshape(WAVES_N, BLOCK_M, BLOCK_N)
P_write = P_lds.reshape(WAVES_N, WAVES_M, TM, LANES_PER_WAVE_M, 1, TN, LANES_PER_WAVE_N, 1)
P_write = P_write.permute((1, 0, 3, 6, 2, 4, 5, 7)).reshape(THREADS_PER_BLOCK, TM, TN)
P_store = P_write[tid].store(S_reg.cast(dtypes.half))
# -- online softmax correction --
beta_i = UOp.placeholder((TM,), dtypes.float, slot=9, addrspace=AddrSpace.REG)
ri4 = UOp.range(TM, 330, AxisType.WEAK)
@@ -292,7 +257,6 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, kv_scale:UOp, valid_kv_len:
l_i = l_i.after(correction)
m_i = m_i.after(correction)
beta_i = beta_i.after(correction)
# Load V transposed into LDS: PV's B operand is logically (D, BLOCK_N), while global V is (BLOCK_N, D).
# It reuses K's slot and must wait for QK WMMA to finish reading that slot.
V_lds = UOp.placeholder((D, BLOCK_N + LDS_PAD), dtypes.half, slot=1, addrspace=AddrSpace.LOCAL)[:, :BLOCK_N]
@@ -305,7 +269,6 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, kv_scale:UOp, valid_kv_len:
pv_barrier = UOp.barrier(UOp.group(P_store, V_store))
P_lds = P_lds.after(pv_barrier)
V_lds = V_lds.after(pv_barrier)
# -- acc += beta * (P @ V) via WMMA --
pv_acc = UOp.placeholder((TM, TD), dtypes.float, slot=10, addrspace=AddrSpace.REG)
pv_acc = pv_acc.after(pv_acc.after(n_tile).store(pv_acc.const_like(0))).after(pv_barrier)
@@ -318,20 +281,16 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, kv_scale:UOp, valid_kv_len:
pv = UOp.wmma(p_frag, v_frag, pv_frag.after(k_pv), *WMMA_ARG)
pv_done = pv_frag.store(pv).end(tm2, tn2).end(k_pv)
pv_acc = pv_acc.after(pv_done)
ri5 = UOp.range(TM, 410, AxisType.WEAK)
rj5 = UOp.range(TD, 411, AxisType.WEAK)
accumulate = acc[ri5, rj5].store(acc[ri5, rj5] + beta_i[ri5] * pv_acc[ri5, rj5]).end(ri5, rj5)
# end KV tile loop
n_tile_end = accumulate.barrier().end(n_tile)
acc = acc.after(n_tile_end)
l_i = l_i.after(n_tile_end)
m_i = m_i.after(n_tile_end)
# normalize: acc /= l_i
acc = acc.after(acc.store(acc * (1 / l_i).reshape(TM, 1).expand(TM, TD)))
# store output
o = o.reshape(WAVES_M, TM, LANES_PER_WAVE_M, 1, WAVES_N, TD, LANES_PER_WAVE_N, 1)
o = o.permute((0, 4, 2, 6, 1, 3, 5, 7)).reshape(THREADS_PER_BLOCK, TM, TD)
@@ -513,12 +472,10 @@ def _qk_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, xsum:UOp, out_features:i
if isinstance(raw_offset, UOp): raw_offset = raw_offset.cast(dtypes.uint64)
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
output, lane = UOp.range(out_features, 0), UOp.range(32, 1, axis_type=AxisType.LOCAL)
group_count = in_features // 32
type_words, output_words = _GGML_QUANT[13][1] // 4, in_features // 256 * _GGML_QUANT[13][1] // 4
def group_dot(group:UOp, output:UOp) -> UOp:
def group_dot(group:UOp) -> UOp:
block, subgroup = group // 8, group % 8
base = raw_offset + output * output_words + block * type_words
qs_base = base + 12 + (subgroup // 2) * 8
@@ -536,15 +493,14 @@ def _qk_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, xsum:UOp, out_features:i
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]
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)
acc = UOp.placeholder((1,), dtypes.float32, slot=0, addrspace=AddrSpace.REG)
acc = acc.after(acc.store(acc.const_like(0)))
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_q5_k", opts_to_apply=()))
group = (lane + chunk*32).valid(lane + chunk*32 < group_count)
update = acc.store(acc.after(chunk) + group_dot(group)).end(chunk)
total = _amd_wave_sum(acc.after(update)[0].load(), lane, 32)
return out[0, output.valid(lane.eq(0))].store(total.cast(out.dtype)).end(output, lane).sink(
arg=KernelInfo(name="linear_q5_k", opts_to_apply=()))
@functools.cache
def _qk_linear_f16_wmma_kernel(out:UOp, raw:UOp, x:UOp, out_features:int, in_features:int,
@@ -621,12 +577,10 @@ def _iq4_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, in_fe
if isinstance(raw_offset, UOp): raw_offset = raw_offset.cast(dtypes.uint64)
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
output, lane = UOp.range(out_features, 0), UOp.range(32, 1, axis_type=AxisType.LOCAL)
group_count = in_features // 32
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:
def group_dot(group:UOp) -> UOp:
block, subgroup = group // 8, group % 8
base = raw_offset + output * output_words + block * type_words
xwords = _amd_vector_load(xq[0, group, 0], 8)
@@ -639,15 +593,14 @@ def _iq4_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, in_fe
((((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()
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)
acc = UOp.placeholder((1,), dtypes.float32, slot=0, addrspace=AddrSpace.REG)
acc = acc.after(acc.store(acc.const_like(0)))
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=()))
group = (lane + chunk*32).valid(lane + chunk*32 < group_count)
update = acc.store(acc.after(chunk) + group_dot(group)).end(chunk)
total = _amd_wave_sum(acc.after(update)[0].load(), lane, 32)
return out[0, output.valid(lane.eq(0))].store(total.cast(out.dtype)).end(output, 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,
@@ -656,7 +609,6 @@ def _iq4_linear_f16_wmma_kernel(out:UOp, raw:UOp, x:UOp, lut:UOp, out_features:i
raw_offset = raw_offset.cast(dtypes.uint64)
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) -> tuple[tuple[UOp, ...], tuple[UOp, ...]]:
low_byte = load_byte(base, 4 + subgroup // 2)
scale_bits = ((low_byte >> (4*(subgroup % 2)).cast(dtypes.uint32)) & 15) | \
@@ -678,7 +630,6 @@ def _iq4_linear_f16_wmma_kernel(out:UOp, raw:UOp, x:UOp, lut:UOp, out_features:i
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 \
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 \
@@ -719,7 +670,6 @@ def _iq4_linear_f16_wmma_kernel(out:UOp, raw:UOp, x:UOp, lut:UOp, out_features:i
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,),
@@ -737,42 +687,36 @@ def _iq4_linear_f16_wmma_kernel(out:UOp, raw:UOp, x:UOp, lut:UOp, out_features:i
@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 = 1
output_block, lane = UOp.range(out_features // output_tile, 0), UOp.range(32, 1, axis_type=AxisType.LOCAL)
outputs = tuple(output_block * output_tile + i for i in range(output_tile))
output, lane = UOp.range(out_features, 0), UOp.range(32, 1, axis_type=AxisType.LOCAL)
group_count, type_size = in_features // 32, _GGML_QUANT[14][1]
output_size = in_features // 256 * type_size
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)
acc = UOp.placeholder((1,), dtypes.float32, slot=0, addrspace=AddrSpace.REG)
acc = acc.after(acc.store(acc.const_like(0)))
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(0, dtypes.int32), UOp.const(0, dtypes.int32)]
for word_idx in range(8):
word = UOp.const(0, dtypes.uint32)
for byte_idx in range(4):
pos = subgroup * 32 + word_idx * 4 + byte_idx
within = pos % 128
low_byte = raw[base + (pos // 128) * 64 + within % 64]
low = (low_byte >> ((within // 64) * 4).cast(dtypes.uint8)) & 15
high_byte = raw[base + 128 + (pos // 128) * 32 + within % 32]
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, 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)
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=()))
base = raw_offset + output * output_size + block * type_size
dots = [UOp.const(0, dtypes.int32), UOp.const(0, dtypes.int32)]
for word_idx in range(8):
word = UOp.const(0, dtypes.uint32)
for byte_idx in range(4):
pos = subgroup * 32 + word_idx * 4 + byte_idx
within = pos % 128
low_byte = raw[base + (pos // 128) * 64 + within % 64]
low = (low_byte >> ((within // 64) * 4).cast(dtypes.uint8)) & 15
high_byte = raw[base + 128 + (pos // 128) * 32 + within % 32]
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, 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)
value = (dots[0].float() * scales[0] + dots[1].float() * scales[1]) * xd[0, group] * dbits.bitcast(dtypes.float16).float()
update = acc.store(acc.after(chunk) + value).end(chunk)
total = _amd_wave_sum(acc.after(update)[0].load(), lane, 32)
return out[0, output.valid(lane.eq(0))].store(total.cast(out.dtype)).end(output, lane).sink(
arg=KernelInfo(name="linear_q6", opts_to_apply=()))
def q8_linear(layer:Linear, x:Tensor, prepared:tuple[Tensor, ...]|None=None) -> Tensor:
tokens = int(x.numel()) // layer.in_features
+7 -6
View File
@@ -176,7 +176,7 @@ class FFNBlock:
def _state_reset_ops(self) -> list[Tensor]: return []
def _init_state(self, x:Tensor): raise NotImplementedError
def _attention(self, x:Tensor, start_pos:int|UOp, use_flash:bool=False, kv_len:int|UOp|None=None,
valid_len:int|UOp|None=None, input_norm:nn.RMSNorm|None=None) -> Tensor: raise NotImplementedError
valid_len:int|UOp|None=None) -> Tensor: raise NotImplementedError
def __call__(self, x:Tensor, start_pos:int|UOp, use_flash:bool=False, kv_len:int|UOp|None=None, valid_len:int|UOp|None=None):
self._init_state(x)
@@ -215,7 +215,7 @@ class TransformerBlock(FFNBlock):
if config.qk_norm: self.attn_q_norm, self.attn_k_norm = nn.RMSNorm(config.qk_norm, config.norm_eps), nn.RMSNorm(config.qk_norm, config.norm_eps)
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:
valid_len:int|UOp|None=None) -> Tensor:
prepared:tuple[Tensor, ...]|None
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)
@@ -317,7 +317,7 @@ class MLATransformerBlock(FFNBlock):
self.attn_output = nn.Linear(config.n_heads * config.v_head_dim, config.dim, bias=False)
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:
valid_len:int|UOp|None=None) -> Tensor:
B, T, _ = x.shape
q_nope_head_dim = self.config.head_dim - self.config.rope_dim
q_proj = self.attn_q_b(self.attn_q_a_norm(self.attn_q_a(x))) if self.config.q_lora_rank > 0 else self.attn_q(x)
@@ -377,7 +377,7 @@ class GatedDeltaNetBlock(FFNBlock):
return Tensor(self.recurrent_state.uop.after(*stores))
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:
valid_len:int|UOp|None=None) -> Tensor:
B, T, _ = x.shape
conv_state, initial_state = self.conv_state, self.recurrent_state
if hasattr(self, "ssm_g_a"):
@@ -403,7 +403,7 @@ class GatedDeltaNetBlock(FFNBlock):
gate = out_gate.reshape(B, 1, self.num_v_heads, self.head_v_dim).sigmoid()
return self.ssm_out((core * gate).reshape(B, 1, -1).cast(x.dtype))
if T == 1:
if input_norm is None: x = x.half()
x = x.half()
prepared = self.attn_gate.prepare(x, self.attn_qkv.ggml_type in (12, 13))
out_gate, qkv = self.attn_gate(x, prepared), self.attn_qkv(x, prepared)
if self.ssm_beta_alpha_weight is not None:
@@ -482,7 +482,8 @@ class Transformer:
dense_config = replace(config, num_experts=0, num_experts_per_tok=0, shared_expert_dim=0, hidden_dim=config.dense_hidden_dim or config.hidden_dim)
if config.ssm: config = replace(config, qk_norm=config.head_dim)
block_cls = MLATransformerBlock if config.kv_lora_rank > 0 else TransformerBlock
self.blk:list[FFNBlock] = [GatedDeltaNetBlock(config, config.ssm) if config.ssm and config.ssm_layers[i] else
self.blk:list[FFNBlock] = [GatedDeltaNetBlock(dense_config if i < config.leading_dense_blocks else config, config.ssm)
if config.ssm and config.ssm_layers[i] else
block_cls(dense_config if i < config.leading_dense_blocks else config) for i in range(config.num_blocks)]
if config.max_context > 8192:
# A full Q8 cache for 262k Qwen leaves no graph workspace on a 24 GB card. Keep two fixed layers host-mapped;