forked from tinygrad/tinygrad
qwen 3.6
This commit is contained in:
@@ -4,10 +4,10 @@ from tinygrad.dtype import AddrSpace, dtypes
|
||||
from tinygrad.helpers import GlobalCounters, Context
|
||||
import math
|
||||
|
||||
BLOCK_M, BLOCK_N = 64, 64
|
||||
BLOCK_M, BLOCK_N = 32, 32
|
||||
WARP_SIZE = 32
|
||||
WMMA_M, WMMA_N, WMMA_K = 16, 16, 16
|
||||
WAVES_M, WAVES_N = 4, 1
|
||||
WAVES_M, WAVES_N = 2, 2
|
||||
LANES_PER_WAVE_M, LANES_PER_WAVE_N = 2, 16
|
||||
WMMA_ACC = WMMA_M // LANES_PER_WAVE_M
|
||||
THREADS_PER_BLOCK = WARP_SIZE * WAVES_M * WAVES_N
|
||||
@@ -49,7 +49,8 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, causal:bool, valid_kv_len:i
|
||||
f"D={D} must be divisible by WMMA_K={WMMA_K} and LANES_PER_WAVE_N={LANES_PER_WAVE_N}"
|
||||
assert BLOCK_M % (WAVES_M * WMMA_M) == 0 and BLOCK_N % LANES_PER_WAVE_N == 0
|
||||
TM = BLOCK_M // (WAVES_M * LANES_PER_WAVE_M)
|
||||
TN = BLOCK_N // (WAVES_N * LANES_PER_WAVE_N)
|
||||
# Each N wave computes the same score tile, then owns a disjoint slice of D for P@V.
|
||||
TN = BLOCK_N // LANES_PER_WAVE_N
|
||||
TD = D // (WAVES_N * LANES_PER_WAVE_N)
|
||||
SCALE = 1.0 / math.sqrt(D)
|
||||
|
||||
@@ -69,7 +70,8 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, causal:bool, valid_kv_len:i
|
||||
|
||||
# LDS allocation: slot 0 = Q then P (shared), slot 1 = K then V
|
||||
# TODO: the memory planner should be able to find this reuse
|
||||
ELEMS_PER_THREAD = BLOCK_M * D // THREADS_PER_BLOCK
|
||||
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]
|
||||
|
||||
@@ -89,11 +91,11 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, causal:bool, valid_kv_len:i
|
||||
|
||||
# load Q + K into LDS (Q reloaded each iteration since P overwrites slot 0)
|
||||
Q_lds = QP_lds[:, :D]
|
||||
Q_store = Q_lds.after(n_tile).reshape(THREADS_PER_BLOCK, ELEMS_PER_THREAD)[tid].store(
|
||||
q.reshape(THREADS_PER_BLOCK, ELEMS_PER_THREAD)[tid])
|
||||
load_k = UOp.range(ELEMS_PER_THREAD, 90, AxisType.LOOP)
|
||||
K_store = KV_lds.reshape(THREADS_PER_BLOCK, ELEMS_PER_THREAD)[tid, load_k].store(
|
||||
k.reshape(physical_n*D)[n_tile*BLOCK_N*D + tid*ELEMS_PER_THREAD + load_k]).end(load_k)
|
||||
Q_store = Q_lds.after(n_tile).reshape(THREADS_PER_BLOCK, Q_ELEMS_PER_THREAD)[tid].store(
|
||||
q.reshape(THREADS_PER_BLOCK, Q_ELEMS_PER_THREAD)[tid])
|
||||
load_k = UOp.range(KV_ELEMS_PER_THREAD, 90, AxisType.LOOP)
|
||||
K_store = KV_lds.reshape(THREADS_PER_BLOCK, KV_ELEMS_PER_THREAD)[tid, load_k].store(
|
||||
k.reshape(physical_n*D)[n_tile*BLOCK_N*D + tid*KV_ELEMS_PER_THREAD + load_k]).end(load_k)
|
||||
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)
|
||||
@@ -106,7 +108,7 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, causal:bool, valid_kv_len:i
|
||||
tn1 = UOp.range(TN, 201, AxisType.LOOP)
|
||||
S_frag = S_reg.reshape(TM // WMMA_ACC, WMMA_ACC, TN).permute(0, 2, 1)[tm1, tn1]
|
||||
q_frag = Q_lds.reshape(WAVES_M, TM // WMMA_ACC, WMMA_M, D // WMMA_K, WMMA_K)[wave_m, tm1, lane_n, k_qk]
|
||||
k_frag = KV_lds_k.reshape(WAVES_N, TN, WMMA_N, D // WMMA_K, WMMA_K)[wave_n, tn1, lane_n, k_qk]
|
||||
k_frag = KV_lds_k.reshape(TN, WMMA_N, D // WMMA_K, WMMA_K)[tn1, lane_n, k_qk]
|
||||
qk = UOp.wmma(q_frag, k_frag, S_frag.after(k_qk), *WMMA_ARG)
|
||||
qk_done = S_frag.store(qk).end(tm1, tn1).end(k_qk)
|
||||
S_reg = S_reg.after(qk_done)
|
||||
@@ -144,9 +146,9 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, causal:bool, valid_kv_len:i
|
||||
p_sum = p_local.after(p_local[ri_ws].store(warp_reduce_sum(p_local[ri_ws], lane)).end(ri_ws))
|
||||
|
||||
# write P = exp(S - m_ij) to P_lds (reuses slot 0, Q no longer needed)
|
||||
P_lds = QP_lds[:, :BLOCK_N]
|
||||
P_write = P_lds.reshape(WAVES_M, TM, LANES_PER_WAVE_M, 1, WAVES_N, TN, LANES_PER_WAVE_N, 1)
|
||||
P_write = P_write.permute((0, 4, 2, 6, 1, 3, 5, 7)).reshape(THREADS_PER_BLOCK, TM, TN)
|
||||
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 --
|
||||
@@ -171,9 +173,9 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, causal:bool, valid_kv_len:i
|
||||
# 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]
|
||||
V_copy = V_lds.after(qk_done).permute(1, 0)
|
||||
load_v = UOp.range(ELEMS_PER_THREAD, 390, AxisType.LOOP)
|
||||
V_store = V_copy.reshape(THREADS_PER_BLOCK, ELEMS_PER_THREAD)[tid, load_v].store(
|
||||
v.reshape(physical_n*D)[n_tile*BLOCK_N*D + tid*ELEMS_PER_THREAD + load_v]).end(load_v)
|
||||
load_v = UOp.range(KV_ELEMS_PER_THREAD, 390, AxisType.LOOP)
|
||||
V_store = V_copy.reshape(THREADS_PER_BLOCK, KV_ELEMS_PER_THREAD)[tid, load_v].store(
|
||||
v.reshape(physical_n*D)[n_tile*BLOCK_N*D + tid*KV_ELEMS_PER_THREAD + load_v]).end(load_v)
|
||||
pv_barrier = UOp.barrier(UOp.group(P_store, V_store))
|
||||
P_lds = P_lds.after(pv_barrier)
|
||||
V_lds = V_lds.after(pv_barrier)
|
||||
@@ -185,7 +187,7 @@ def _amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp, causal:bool, valid_kv_len:i
|
||||
tm2 = UOp.range(TM // WMMA_ACC, 401, AxisType.LOOP)
|
||||
tn2 = UOp.range(TD, 402, AxisType.LOOP)
|
||||
pv_frag = pv_acc.reshape(TM // WMMA_ACC, WMMA_ACC, TD).permute(0, 2, 1)[tm2, tn2]
|
||||
p_frag = P_lds.reshape(WAVES_M, TM // WMMA_ACC, WMMA_M, BLOCK_N // WMMA_K, WMMA_K)[wave_m, tm2, lane_n, k_pv]
|
||||
p_frag = P_lds[wave_n].reshape(WAVES_M, TM // WMMA_ACC, WMMA_M, BLOCK_N // WMMA_K, WMMA_K)[wave_m, tm2, lane_n, k_pv]
|
||||
v_frag = V_lds.reshape(WAVES_N, TD, WMMA_N, BLOCK_N // WMMA_K, WMMA_K)[wave_n, tn2, lane_n, k_pv]
|
||||
pv = UOp.wmma(p_frag, v_frag, pv_frag.after(k_pv), *WMMA_ARG)
|
||||
pv_done = pv_frag.store(pv).end(tm2, tn2).end(k_pv)
|
||||
|
||||
@@ -17,6 +17,7 @@ class TestLLMServer(unittest.TestCase):
|
||||
cls.mock_tok.is_end = Mock(side_effect=lambda tid: tid in (999,))
|
||||
|
||||
cls.mock_model = Mock()
|
||||
cls.mock_model.max_context = 4
|
||||
cls.mock_model.generate = Mock(side_effect=lambda ids, **kwargs: iter([300, 301, 999]))
|
||||
cls.mock_model.get_start_pos = Mock(return_value=0)
|
||||
|
||||
@@ -129,6 +130,16 @@ class TestLLMServer(unittest.TestCase):
|
||||
self.assertIsNotNone(resp.usage.prompt_tokens)
|
||||
self.assertIsNotNone(resp.usage.completion_tokens)
|
||||
|
||||
def test_context_length_error(self):
|
||||
from openai import BadRequestError
|
||||
self.mock_tok.encode.return_value = [200, 201, 202, 203]
|
||||
try:
|
||||
with self.assertRaises(BadRequestError) as err:
|
||||
self.client.chat.completions.create(model="test-model", messages=[{"role":"user", "content":"too long"}])
|
||||
self.assertEqual(err.exception.code, "context_length_exceeded")
|
||||
finally:
|
||||
self.mock_tok.encode.return_value = [200, 201, 202]
|
||||
|
||||
def test_max_tokens_streaming(self):
|
||||
self.mock_model.generate = Mock(side_effect=lambda ids, **kwargs: iter([300, 301, 302, 303, 999]))
|
||||
stream = self.client.chat.completions.create(
|
||||
@@ -170,6 +181,7 @@ class TestLLMToolCalls(unittest.TestCase):
|
||||
cls.mock_tok.is_end = Mock(return_value=False)
|
||||
|
||||
cls.mock_model = Mock()
|
||||
cls.mock_model.max_context = 4
|
||||
cls.mock_model.get_start_pos = Mock(return_value=0)
|
||||
|
||||
from tinygrad.llm.serve import LLMServer
|
||||
|
||||
@@ -39,7 +39,6 @@ def graph_split_rewrite(linear:UOp, max_batch_size:int=0) -> UOp:
|
||||
if len(current_batch) <= 1 and not getenv("GRAPH_ONE_KERNEL"): new_src.extend(current_batch)
|
||||
else:
|
||||
new_src.append(create_graph_call(current_batch))
|
||||
max_batch_size *= 2
|
||||
if DEBUG >= 2: print(f"JIT GRAPHing batch with {len(current_batch)} kernels")
|
||||
current_batch, current_batch_devs = [], []
|
||||
|
||||
|
||||
+5
-5
@@ -130,7 +130,7 @@ class FallbackTemplate:
|
||||
if self.tok.preset == 'glm4': return ""
|
||||
if self.tok.preset == 'tekken': return "[/INST]"
|
||||
return self.tok.decode([self.tok.eos_id])
|
||||
def render(self, messages:list[dict], tools=None, add_generation_prompt:bool=True) -> str:
|
||||
def render(self, messages:list[dict], tools=None, add_generation_prompt:bool=True, enable_thinking:bool=False) -> str:
|
||||
out = self.tok.decode([] if self.tok.bos_id is None else [self.tok.bos_id]) + ("<sop>" if self.tok.preset == 'glm4' else "")
|
||||
for msg in messages:
|
||||
out += self.role(msg["role"])
|
||||
@@ -152,7 +152,7 @@ def main():
|
||||
parser.add_argument("--max_context", type=int, default=4096, help="Max Context Length")
|
||||
parser.add_argument("--serve", nargs='?', type=int, const=8000, metavar="PORT", help="Run OpenAI compatible API (optional port, default 8000)")
|
||||
parser.add_argument("--warmup", action="store_true", help="warmup the JIT")
|
||||
parser.add_argument("--beam", type=int, help="Kernel optimization beam width (serving default: 2)")
|
||||
parser.add_argument("--beam", type=int, help="Kernel optimization beam width")
|
||||
parser.add_argument("--benchmark", nargs='?', type=int, const=20, metavar="COUNT", help="Benchmark tok/s (optional count, default 20)")
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -160,8 +160,8 @@ def main():
|
||||
model, kv = Transformer.from_gguf(fetch(models.get(args.model, args.model)), args.max_context)
|
||||
model_name = kv.get('general.name') or kv.get('general.basename') or args.model
|
||||
file_sizes = [y.nbytes() for y in UOp.sink(*[x.uop for x in nn.state.get_parameters(model)]).toposort() if y.op is Ops.BUFFER]
|
||||
print(f"using model \"{model_name}\" with {sum(file_sizes):,} bytes and {sum(x.numel() for x in nn.state.get_parameters(model)):,} params, "
|
||||
f"max context {args.max_context} on {nn.state.get_parameters(model)[0].device}")
|
||||
print(f"using model \"{model_name}\" with {sum(file_sizes):,} bytes and {model.parameter_count:,} params, "
|
||||
f"max context {model.max_context} on {nn.state.get_parameters(model)[0].device}")
|
||||
|
||||
# get tokenizer
|
||||
tok = SimpleTokenizer.from_gguf_kv(kv)
|
||||
@@ -182,7 +182,7 @@ def main():
|
||||
|
||||
# warmup the JIT
|
||||
if args.warmup or args.serve:
|
||||
beam = args.beam if args.beam is not None else BEAM.value or (2 if args.serve else 0)
|
||||
beam = args.beam if args.beam is not None else BEAM.value
|
||||
print(f"warming serving JITs with BEAM={beam}")
|
||||
with Context(DEBUG=max(DEBUG.value, 1), BEAM=beam):
|
||||
model.warmup()
|
||||
|
||||
+23
-8
@@ -1,7 +1,8 @@
|
||||
import functools, io, pathlib, re, struct
|
||||
import functools, io, pathlib, re, struct, weakref
|
||||
from typing import Any, Callable
|
||||
|
||||
from tinygrad.tensor import Tensor
|
||||
from tinygrad.uop.ops import UOp
|
||||
from tinygrad.dtype import dtypes
|
||||
from tinygrad.helpers import prod, round_up
|
||||
from tinygrad.nn.state import TensorIO
|
||||
@@ -20,7 +21,14 @@ _GGML_NATIVE = {0: dtypes.float32, 1: dtypes.float16, 24: dtypes.int8, 25: dtype
|
||||
_GGML_QUANT = {2:(32,18), 3:(32,20), 6:(32,22), 7:(32,24), 8:(32,34),
|
||||
12:(256,144), 13:(256,176), 14:(256,210), 18:(256,98), 21:(256,110), 22:(256,82), 23:(256,136), 39:(32,17), 41:(128,18)}
|
||||
|
||||
def ggml_data_to_tensor(t: Tensor, n: int, ggml_type: int) -> Tensor:
|
||||
_quantized_tensors:weakref.WeakKeyDictionary[UOp, tuple[UOp, int]] = weakref.WeakKeyDictionary()
|
||||
|
||||
def get_ggml_quantization(tensor:Tensor) -> tuple[Tensor, int]|None:
|
||||
if (meta:=_quantized_tensors.get(tensor.uop)) is None: return None
|
||||
packed, ggml_type = meta
|
||||
return Tensor(packed), ggml_type
|
||||
|
||||
def ggml_data_to_tensor(t: Tensor, n: int, ggml_type: int, contiguous:bool=True) -> Tensor:
|
||||
"""
|
||||
Converts ggml tensor data to a tinygrad tensor.
|
||||
|
||||
@@ -35,14 +43,14 @@ def ggml_data_to_tensor(t: Tensor, n: int, ggml_type: int) -> Tensor:
|
||||
if (dtype := _GGML_NATIVE.get(ggml_type)) is not None:
|
||||
return t[:dtype.itemsize * n].contiguous().bitcast(dtype)
|
||||
|
||||
def q_to_uint8(t: Tensor, b: int) -> Tensor:
|
||||
# TODO: rewrite with arange?
|
||||
shift_tensor, bitmask = Tensor.stack(*[ Tensor(2**(i*b), device=t.device, dtype=t.dtype) for i in range(8//b) ]), 0xff >> (8 - b)
|
||||
return t.unsqueeze(-1).expand((*t.shape,8//b)).div(shift_tensor, rounding_mode="trunc").bitwise_and(bitmask).transpose(-1, -2).flatten(-2)
|
||||
def q_to_uint8(t:Tensor, b:int) -> Tensor:
|
||||
shift_tensor, bitmask = Tensor.stack(*[Tensor(2**(i*b), device=t.device, dtype=t.dtype) for i in range(8//b)]), 0xff >> (8-b)
|
||||
return t.unsqueeze(-1).expand((*t.shape, 8//b)).div(shift_tensor, rounding_mode="trunc").bitwise_and(bitmask).transpose(-1, -2).flatten(-2)
|
||||
|
||||
if (nelements_nbytes := _GGML_QUANT.get(ggml_type)) is not None:
|
||||
from tinygrad.runtime.autogen import ggml_common as _ggml
|
||||
blocks = t[:(n//nelements_nbytes[0])*nelements_nbytes[1]].reshape((-1, nelements_nbytes[1])).contiguous()
|
||||
blocks = t[:(n//nelements_nbytes[0])*nelements_nbytes[1]].reshape((-1, nelements_nbytes[1]))
|
||||
if contiguous: blocks = blocks.contiguous()
|
||||
if ggml_type == 2: return (q_to_uint8(blocks[:,2:], 4).bitcast(dtypes.int8) - 8) * blocks[:,:2].bitcast(dtypes.float16).cast(dtypes.float32)
|
||||
if ggml_type == 3:
|
||||
d, m = (blocks[:,s:s+2].bitcast(dtypes.float16).cast(dtypes.float32) for s in [ 0, 2 ])
|
||||
@@ -146,7 +154,14 @@ def _gguf_parse(tensor: Tensor) -> tuple[dict, dict[str, Tensor]]:
|
||||
alignment, pos = kv_data.get("general.alignment", 32), r.tell()
|
||||
data_start = round_up(pos, alignment)
|
||||
|
||||
state_dict = {name: ggml_data_to_tensor(tensor[data_start + off:], prod(dims), typ).reshape(*reversed(dims)) for name, dims, typ, off in t_infos}
|
||||
state_dict = {}
|
||||
for name, dims, typ, off in t_infos:
|
||||
n, shape = prod(dims), tuple(reversed(dims))
|
||||
decoded = ggml_data_to_tensor(data:=tensor[data_start + off:], n, typ).reshape(*shape)
|
||||
if typ in _GGML_QUANT:
|
||||
block_size, type_size = _GGML_QUANT[typ]
|
||||
_quantized_tensors[decoded.uop] = (data[:n//block_size*type_size].uop, typ)
|
||||
state_dict[name] = decoded
|
||||
return kv_data, state_dict
|
||||
|
||||
def _gguf_split_paths(path: pathlib.Path, kv: dict) -> list[pathlib.Path]:
|
||||
|
||||
+441
-98
@@ -1,9 +1,156 @@
|
||||
from __future__ import annotations
|
||||
import functools, itertools, pathlib
|
||||
from dataclasses import dataclass, replace
|
||||
from tinygrad import Device, Tensor, nn, UOp, TinyJit, getenv, function
|
||||
from tinygrad.llm.gguf import gguf_load
|
||||
from tinygrad.uop.ops import resolve
|
||||
from tinygrad import Device, Tensor, nn, UOp, TinyJit, getenv, function, dtypes
|
||||
from tinygrad.dtype import AddrSpace
|
||||
from tinygrad.llm.gguf import get_ggml_quantization, ggml_data_to_tensor, gguf_load, _GGML_QUANT
|
||||
from tinygrad.uop.ops import resolve, Ops, KernelInfo, AxisType
|
||||
|
||||
def _q8_kernel(quant:UOp, scale:UOp, x:UOp, in_features:int) -> 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.REDUCE)
|
||||
amax = UOp.placeholder((1,), dtypes.float32, 0, addrspace=AddrSpace.REG)
|
||||
amax = amax.after(token, group)[0].set(0.0)
|
||||
amax = amax[0].set(amax.after(lane)[0].maximum(x[token * in_features + group * 32 + lane].cast(dtypes.float32).abs()), end=lane)
|
||||
d = (amax[0] / 127).maximum(1e-8)
|
||||
stores = [scale[token, group].store(d)]
|
||||
for word_idx in range(8):
|
||||
word = UOp.const(dtypes.uint32, 0)
|
||||
for byte_idx in range(4):
|
||||
value = (x[token * in_features + group * 32 + word_idx * 4 + byte_idx].cast(dtypes.float32) / d).round().maximum(-127).minimum(127)
|
||||
byte = value.cast(dtypes.int8).bitcast(dtypes.uint8).cast(dtypes.uint32)
|
||||
word = word | (byte << (8 * byte_idx))
|
||||
stores.append(quant[token, group, word_idx].store(word))
|
||||
return UOp.group(*stores).end(token, group).sink(arg=KernelInfo(name="q8_quantize", opts_to_apply=()))
|
||||
|
||||
def _q8_quantize(x:Tensor, tokens:int, in_features:int) -> tuple[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)
|
||||
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 _amd_dp4a(a:UOp, b:UOp, c:UOp) -> UOp:
|
||||
return UOp(Ops.CUSTOMI, dtypes.int32, (a.cast(dtypes.int32), b.cast(dtypes.int32), c),
|
||||
arg="__builtin_amdgcn_sudot4(true, {}, true, {}, {}, false)")
|
||||
|
||||
def _amd_wave_sum(value:UOp, lane:UOp, lane_count:int) -> UOp:
|
||||
assert lane_count in (8, 16, 32)
|
||||
for offset in (16, 8, 4, 2, 1)[{32:0, 16:1, 8:2}[lane_count]:]:
|
||||
value = value + UOp(Ops.CUSTOM, dtypes.float32, (((lane ^ offset) * 4).cast(dtypes.int32), value),
|
||||
arg="__builtin_bit_cast(float, __builtin_amdgcn_ds_bpermute({0}, __builtin_bit_cast(int, {1})))")
|
||||
return value
|
||||
|
||||
def _q8_linear_kernel(out:UOp, raw:UOp, xq:UOp, xd:UOp, out_features:int, in_features:int, raw_offset:int=0) -> UOp:
|
||||
token_tile = 4 if out.shape[0] % 4 == 0 else 1
|
||||
token_block, output = UOp.range(out.shape[0] // token_tile, 0), UOp.range(out_features, 1)
|
||||
tokens = tuple(token_block * token_tile + i for i in range(token_tile))
|
||||
group_count, lane_count = in_features // 32, min(32, in_features // 32)
|
||||
lane = UOp.range(lane_count, 2, axis_type=AxisType.LOCAL)
|
||||
|
||||
def group_dot(group:UOp) -> list[UOp]:
|
||||
block = output * group_count + group
|
||||
base, odd = raw_offset + block * 8 + block // 2, (block & 1).ne(0)
|
||||
dots = [UOp.const(dtypes.int32, 0)] * token_tile
|
||||
for word_idx in range(8):
|
||||
# Q8_0 blocks are 34 bytes, so their two-byte scale makes alternate blocks word-aligned. Read aligned u32s
|
||||
# directly; the other blocks need only two adjacent words instead of four individual byte loads.
|
||||
word = odd.where(raw[base + 1 + word_idx], (raw[base + word_idx] >> 16) | (raw[base + 1 + word_idx] << 16))
|
||||
dots = [_amd_dp4a(word, xq[token, group, word_idx], dot) for token,dot in zip(tokens, dots)]
|
||||
dbits = odd.where(raw[base] >> 16, raw[base] & 0xffff).cast(dtypes.uint16)
|
||||
return [dot.cast(dtypes.float32) * xd[token, group] * dbits.bitcast(dtypes.float16).float() for token,dot in zip(tokens, dots)]
|
||||
|
||||
values = [UOp.const(dtypes.float32, 0)] * token_tile
|
||||
for offset in range(0, group_count, lane_count):
|
||||
dots = group_dot((lane + offset).valid(lane + offset < group_count))
|
||||
values = [value + dot for value,dot in zip(values, dots)]
|
||||
totals = [_amd_wave_sum(value, lane, lane_count) for value in values]
|
||||
stores = [out[token.valid(lane.eq(0)), output].store(total.cast(out.dtype)) for token,total in zip(tokens, totals)]
|
||||
return UOp.group(*stores).end(token_block, output, lane).sink(
|
||||
arg=KernelInfo(name="linear_q8", opts_to_apply=()))
|
||||
|
||||
class Linear(nn.Linear):
|
||||
def __init__(self, in_features:int, out_features:int, bias=True):
|
||||
super().__init__(in_features, out_features, bias)
|
||||
self.in_features, self.out_features = in_features, out_features
|
||||
self.ggml_type:int|None = None
|
||||
def set_quantized(self, packed:Tensor, ggml_type:int):
|
||||
self.weight, self.ggml_type = packed.flatten(), ggml_type
|
||||
def prepare(self, x:Tensor) -> tuple[Tensor, Tensor]|None:
|
||||
return _q8_quantize(x, int(x.numel()) // self.in_features, self.in_features) \
|
||||
if self.ggml_type == 8 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 == 8 and str(self.weight.device).startswith("AMD"):
|
||||
tokens = int(x.numel()) // self.in_features
|
||||
xq, xd = prepared if prepared is not None else _q8_quantize(x, tokens, self.in_features)
|
||||
out = Tensor.empty(tokens, self.out_features, dtype=dtypes.float32, device=x.device)
|
||||
raw, raw_offset = self.weight.uop, 0
|
||||
while raw.op in (Ops.BITCAST, Ops.RESHAPE): raw = raw.src[0]
|
||||
while raw.op is Ops.SHRINK:
|
||||
raw_offset += raw.src[1].arg * raw.dtype.itemsize
|
||||
raw = raw.src[0]
|
||||
assert raw_offset % 4 == 0 and raw.dtype == dtypes.uint8
|
||||
srcs = (out.uop, raw, xq.uop, xd.uop)
|
||||
params = [UOp.placeholder_like(src, slot=i) for i,src in enumerate(srcs)]
|
||||
params[1] = params[1].replace(dtype=dtypes.uint32, src=(params[1].src[0] * raw.dtype.itemsize // 4,),
|
||||
arg=replace(params[1].arg, dtype=dtypes.uint32))
|
||||
kernel = _q8_linear_kernel(params[0], params[1], params[2], params[3], self.out_features, self.in_features, raw_offset // 4).call(*srcs)
|
||||
out = Tensor(srcs[0].after(kernel)).reshape(*x.shape[:-1], self.out_features)
|
||||
return out if self.bias is None else out + self.bias
|
||||
return super().__call__(x)
|
||||
|
||||
def _packed_expert_kernel(out:UOp, raw:UOp, sel:UOp, xq:UOp, xd:UOp, lut:UOp,
|
||||
out_features:int, in_features:int, ggml_type:int, routes_per_input:int) -> UOp:
|
||||
route, output = UOp.range(out.shape[0], 0), UOp.range(out_features, 1)
|
||||
group_count, lane_count = in_features // 32, min(32, in_features // 32)
|
||||
lane = UOp.range(lane_count, 2, axis_type=AxisType.LOCAL)
|
||||
expert, xidx = sel[route].cast(dtypes.index), route // routes_per_input
|
||||
type_size = _GGML_QUANT[ggml_type][1]
|
||||
expert_size = out_features * in_features // 256 * type_size
|
||||
|
||||
def group_dot(group:UOp) -> UOp:
|
||||
block, subgroup = group // 8, group % 8
|
||||
base = expert * expert_size + output * (in_features // 256 * type_size) + block * type_size
|
||||
dot = UOp.const(dtypes.int32, 0)
|
||||
if ggml_type == 21: # IQ3_S
|
||||
for word_idx in range(8):
|
||||
qi = raw[base + 2 + subgroup * 8 + word_idx].cast(dtypes.uint16) + \
|
||||
(((raw[base + 66 + subgroup] >> word_idx) & 1).cast(dtypes.uint16) << 8)
|
||||
word, signs = UOp.const(dtypes.uint32, 0), raw[base + 74 + subgroup * 4 + word_idx // 2]
|
||||
for byte_idx in range(4):
|
||||
sign = ((signs >> (word_idx % 2 * 4 + byte_idx)) & 1).ne(0).where(-1, 1).cast(dtypes.int8)
|
||||
byte = (lut[qi.cast(dtypes.index) * 4 + byte_idx] * sign).cast(dtypes.int8).bitcast(dtypes.uint8).cast(dtypes.uint32)
|
||||
word = word | (byte << (8 * byte_idx))
|
||||
dot = _amd_dp4a(word, xq[xidx, group, word_idx], dot)
|
||||
scale_shift = (4 * (subgroup % 2)).cast(dtypes.uint8)
|
||||
scale = 1 + 2 * ((raw[base + 106 + subgroup // 2] >> scale_shift) & 15).cast(dtypes.float32)
|
||||
else: # IQ4_XS
|
||||
for word_idx in range(8):
|
||||
word = UOp.const(dtypes.uint32, 0)
|
||||
for byte_idx in range(4):
|
||||
qbyte = raw[base + 8 + subgroup * 16 + (word_idx % 4) * 4 + byte_idx]
|
||||
q = (qbyte >> (4 * (word_idx // 4))) & 15
|
||||
byte = lut[q.cast(dtypes.index)].cast(dtypes.int8).bitcast(dtypes.uint8).cast(dtypes.uint32)
|
||||
word = word | (byte << (8 * byte_idx))
|
||||
dot = _amd_dp4a(word, xq[xidx, group, word_idx], dot)
|
||||
low = (raw[base + 4 + subgroup // 2] >> (4 * (subgroup % 2)).cast(dtypes.uint8)) & 15
|
||||
high_word = raw[base + 2].cast(dtypes.uint16) | (raw[base + 3].cast(dtypes.uint16) << 8)
|
||||
scale = ((low.cast(dtypes.uint16) | (((high_word >> (2 * subgroup).cast(dtypes.uint16)) & 3) << 4)).cast(dtypes.uint8).
|
||||
bitcast(dtypes.int8)-32).float()
|
||||
dbits = raw[base].cast(dtypes.uint16) | (raw[base + 1].cast(dtypes.uint16) << 8)
|
||||
return dot.cast(dtypes.float32) * xd[xidx, group] * dbits.bitcast(dtypes.float16).float() * scale
|
||||
|
||||
value = sum((group_dot((lane + offset).valid(lane + offset < group_count)) for offset in range(0, group_count, lane_count)),
|
||||
UOp.const(dtypes.float32, 0))
|
||||
total = _amd_wave_sum(value, lane, lane_count)
|
||||
return out[route.valid(lane.eq(0)), output].store(total.cast(out.dtype)).end(route, output, lane).sink(
|
||||
arg=KernelInfo(name=f"expert_q8_{ggml_type}", opts_to_apply=()))
|
||||
|
||||
@functools.cache
|
||||
def _expert_lut(device:str, ggml_type:int) -> Tensor:
|
||||
from tinygrad.runtime.autogen import ggml_common
|
||||
values = [((word >> (8 * i)) & 0xff) for word in ggml_common.iq3s_grid for i in range(4)] if ggml_type == 21 else ggml_common.kvalues_iq4nl
|
||||
return Tensor(values, dtype=dtypes.int8, device=device).contiguous().realize()
|
||||
|
||||
@functools.cache
|
||||
def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0, device:str|None=None) -> Tensor:
|
||||
@@ -14,10 +161,34 @@ def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0, device:str|
|
||||
class ExpertWeights:
|
||||
"""Like nn.Linear but with num_experts dimension. Weight shape: (num_experts, out_features, in_features)."""
|
||||
def __init__(self, num_experts:int, in_features:int, out_features:int):
|
||||
self.num_experts, self.in_features, self.out_features = num_experts, in_features, out_features
|
||||
self.weight = Tensor.zeros(num_experts, out_features, in_features)
|
||||
def __call__(self, sel:Tensor, x:Tensor) -> Tensor:
|
||||
self.ggml_type:int|None = None
|
||||
def set_quantized(self, weight:Tensor, packed:Tensor, ggml_type:int):
|
||||
assert weight.shape == (self.num_experts, self.out_features, self.in_features)
|
||||
self.weight, self.ggml_type = packed.flatten(), ggml_type
|
||||
def prepare(self, x:Tensor) -> tuple[Tensor, Tensor]:
|
||||
return _q8_quantize(x, int(x.numel()) // self.in_features, self.in_features)
|
||||
def __call__(self, sel:Tensor, x:Tensor, prepared:tuple[Tensor, Tensor]|None=None) -> Tensor:
|
||||
# sel: (B, T, k), x: (B, T, 1, in) or (B, T, k, in) -> output: (B, T, k, out)
|
||||
return (x.unsqueeze(-2) @ self.weight[sel].transpose(-1, -2)).contiguous().squeeze(-2)
|
||||
ggml_type = self.ggml_type
|
||||
if ggml_type in (21, 23) and str(self.weight.device).startswith("AMD"):
|
||||
input_count = int(x.numel()) // self.in_features
|
||||
routes_per_input = int(sel.numel()) // input_count
|
||||
xq, xd = prepared if prepared is not None else self.prepare(x)
|
||||
flat_sel = sel if len(sel.shape) == 1 else sel.flatten().clone()
|
||||
out = Tensor.empty(int(sel.numel()), self.out_features, dtype=dtypes.float32, device=x.device)
|
||||
out = Tensor.custom_kernel(out, self.weight, flat_sel, xq, xd, _expert_lut(str(x.device), ggml_type),
|
||||
fxn=lambda out,raw,sel,xq,xd,lut:_packed_expert_kernel(out, raw, sel, xq, xd, lut, self.out_features,
|
||||
self.in_features, ggml_type, routes_per_input))[0]
|
||||
return out if len(sel.shape) == 1 else out.reshape(*sel.shape, self.out_features)
|
||||
if self.ggml_type is None: weight = self.weight[sel]
|
||||
else:
|
||||
packed = self.weight.reshape(self.num_experts, -1)[sel].flatten()
|
||||
weight = ggml_data_to_tensor(packed, int(sel.numel()) * self.out_features * self.in_features,
|
||||
self.ggml_type, contiguous=False).reshape(*sel.shape, self.out_features, self.in_features)
|
||||
if getenv("HALF", 1): weight = weight.cast('float16')
|
||||
return (x.unsqueeze(-2) @ weight.transpose(-1, -2)).contiguous().squeeze(-2)
|
||||
|
||||
def apply_rope(x:Tensor, freqs_cis:Tensor) -> Tensor:
|
||||
assert x.shape[-1] % 2 == 0
|
||||
@@ -33,6 +204,40 @@ def pairwise_topk(x: Tensor, k: int) -> tuple[Tensor, Tensor]:
|
||||
sel = x.const_like(0).scatter(-1, cmp.sum(axis=-1).cast('int32'), vals)[:,:,n-k:].cast('int32')
|
||||
return x.gather(-1, sel), sel
|
||||
|
||||
def _inverse_unit_lower_kernel(out:UOp, x:UOp, n:int) -> UOp:
|
||||
outer_count = 1
|
||||
for dim in out.shape[:-2]:
|
||||
assert isinstance(dim, int)
|
||||
outer_count *= dim
|
||||
outer, lane = UOp.range(outer_count, 0), UOp.range(n, 1, axis_type=AxisType.LOCAL)
|
||||
raw = UOp.placeholder((n*n,), x.dtype, 0, addrspace=AddrSpace.LOCAL)
|
||||
solved = UOp.placeholder((n*n,), x.dtype, 1, addrspace=AddrSpace.LOCAL)
|
||||
ready = UOp.group(*(raw[row*n+lane].store(x.flatten()[outer*n*n+row*n+lane]) for row in range(n))).barrier()
|
||||
for row in range(n):
|
||||
base, previous = raw.after(ready), solved.after(ready)
|
||||
value = base[row*n+lane] + sum((base[row*n+i] * previous[i*n+lane] for i in range(row)), UOp.const(x.dtype, 0))
|
||||
ready = solved.after(ready)[row*n+lane].store((lane < row).where(value, UOp.const(x.dtype, 0))).barrier()
|
||||
result = solved.after(ready)
|
||||
stores = [out.flatten()[outer*n*n+row*n+lane].store(lane.eq(row).where(UOp.const(x.dtype, 1), result[row*n+lane]))
|
||||
for row in range(n)]
|
||||
return UOp.group(*stores).end(outer, lane).sink(arg=KernelInfo(name="inverse_unit_lower", opts_to_apply=()))
|
||||
|
||||
def inverse_unit_lower(x:Tensor) -> Tensor:
|
||||
"""Reference-ordered inverse of I-x for a strictly lower-triangular x."""
|
||||
n = x.shape[-1]
|
||||
assert isinstance(n, int)
|
||||
if n == 64 and str(x.device).startswith("AMD"):
|
||||
out = Tensor.empty(*x.shape, dtype=x.dtype, device=x.device)
|
||||
return Tensor.custom_kernel(out, x, fxn=lambda out,x:_inverse_unit_lower_kernel(out, x, n))[0]
|
||||
rows = [x[..., 0, :].const_like(0)]
|
||||
for i in range(1, n):
|
||||
prefix = x[..., i, :i]
|
||||
previous = Tensor.stack(*rows, dim=-2)[..., :, :i]
|
||||
rows.append((prefix + (prefix.unsqueeze(-1) * previous).sum(-2)).pad((0, n-i)))
|
||||
return Tensor.stack(*rows, dim=-2) + Tensor.eye(n, dtype=x.dtype)
|
||||
|
||||
def l2norm(x:Tensor) -> Tensor: return x * (x.square().sum(-1, keepdim=True) + 1e-6).rsqrt()
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SSMConfig:
|
||||
conv_kernel: int
|
||||
@@ -75,6 +280,7 @@ class TransformerConfig:
|
||||
class FFNBlock:
|
||||
def __init__(self, config:TransformerConfig):
|
||||
self.config = config
|
||||
self.pending_state:tuple[Tensor, Tensor]|None = None
|
||||
|
||||
# --- RMSNorms --------------------------------------------------------
|
||||
self.attn_norm = nn.RMSNorm(config.dim, config.norm_eps)
|
||||
@@ -82,25 +288,26 @@ class FFNBlock:
|
||||
|
||||
# --- feed-forward (MoE or dense) -------------------------------------
|
||||
if config.num_experts > 0:
|
||||
self.ffn_gate_inp = nn.Linear(config.dim, config.num_experts, bias=False) # router
|
||||
self.ffn_gate_inp = Linear(config.dim, config.num_experts, bias=False) # router
|
||||
if config.expert_bias: self.exp_probs_b = {"bias": Tensor.zeros(config.num_experts)}
|
||||
self.ffn_gate_exps = ExpertWeights(config.num_experts, config.dim, config.hidden_dim)
|
||||
self.ffn_up_exps = ExpertWeights(config.num_experts, config.dim, config.hidden_dim)
|
||||
self.ffn_down_exps = ExpertWeights(config.num_experts, config.hidden_dim, config.dim)
|
||||
if config.shared_expert_dim > 0:
|
||||
self.ffn_gate_shexp = nn.Linear(config.dim, config.shared_expert_dim, bias=False)
|
||||
self.ffn_up_shexp = nn.Linear(config.dim, config.shared_expert_dim, bias=False)
|
||||
self.ffn_down_shexp = nn.Linear(config.shared_expert_dim, config.dim, bias=False)
|
||||
self.ffn_gate_shexp = Linear(config.dim, config.shared_expert_dim, bias=False)
|
||||
self.ffn_up_shexp = Linear(config.dim, config.shared_expert_dim, bias=False)
|
||||
self.ffn_down_shexp = Linear(config.shared_expert_dim, config.dim, bias=False)
|
||||
if config.shared_expert_gate: self.ffn_gate_inp_shexp = {"weight": Tensor.zeros(config.dim)}
|
||||
else:
|
||||
self.ffn_gate = nn.Linear(config.dim, config.hidden_dim, bias=False)
|
||||
self.ffn_up = nn.Linear(config.dim, config.hidden_dim, bias=False)
|
||||
self.ffn_down = nn.Linear(config.hidden_dim, config.dim, bias=False)
|
||||
self.ffn_gate = Linear(config.dim, config.hidden_dim, bias=False)
|
||||
self.ffn_up = Linear(config.dim, config.hidden_dim, bias=False)
|
||||
self.ffn_down = Linear(config.hidden_dim, config.dim, bias=False)
|
||||
|
||||
def _feed_forward(self, x:Tensor) -> Tensor:
|
||||
if hasattr(self, 'ffn_gate_exps'):
|
||||
h = x.unsqueeze(2) # (B, T, 1, D) - add expert dim for broadcasting
|
||||
logits = self.ffn_gate_inp(x)
|
||||
prepared = self.ffn_gate_exps.prepare(h) if self.ffn_gate_exps.ggml_type in (21, 23) and str(h.device).startswith("AMD") else None
|
||||
logits = self.ffn_gate_inp(x, prepared)
|
||||
if hasattr(self, 'exp_probs_b'):
|
||||
probs = logits.sigmoid()
|
||||
_, sel = pairwise_topk(probs + self.exp_probs_b["bias"], self.config.num_experts_per_tok)
|
||||
@@ -110,31 +317,48 @@ class FFNBlock:
|
||||
vals, sel = pairwise_topk(logits, self.config.num_experts_per_tok)
|
||||
probs = vals.softmax(-1) if self.config.norm_topk_prob else logits.softmax(-1).gather(-1, sel)
|
||||
probs = probs * self.config.routed_scaling_factor
|
||||
x_down = self.ffn_down_exps(sel, (self.ffn_gate_exps(sel, h).silu() * self.ffn_up_exps(sel, h)).contiguous()) # (B, T, k, D)
|
||||
if prepared is not None:
|
||||
flat_sel = sel.flatten().clone()
|
||||
gate, up = self.ffn_gate_exps(flat_sel, h, prepared), self.ffn_up_exps(flat_sel, h, prepared)
|
||||
x_down = self.ffn_down_exps(flat_sel, (gate.silu() * up).contiguous()).reshape(*sel.shape, self.config.dim)
|
||||
else: x_down = self.ffn_down_exps(sel, (self.ffn_gate_exps(sel, h).silu() * self.ffn_up_exps(sel, h)).contiguous())
|
||||
out = (x_down * probs.unsqueeze(-1)).sum(axis=2) # (B, T, D)
|
||||
if hasattr(self, 'ffn_gate_shexp'):
|
||||
shexp = self.ffn_down_shexp(self.ffn_gate_shexp(x).silu().contiguous() * self.ffn_up_shexp(x))
|
||||
shexp = self.ffn_down_shexp(self.ffn_gate_shexp(x, prepared).silu().contiguous() * self.ffn_up_shexp(x, prepared))
|
||||
if hasattr(self, 'ffn_gate_inp_shexp'): shexp = shexp * (x * self.ffn_gate_inp_shexp["weight"]).sum(axis=-1, keepdim=True).sigmoid()
|
||||
out = out + shexp
|
||||
return out
|
||||
# TODO: remove the need for this contiguous
|
||||
return self.ffn_down(self.ffn_gate(x).silu().contiguous() * self.ffn_up(x))
|
||||
prepared = self.ffn_gate.prepare(x)
|
||||
return self.ffn_down(self.ffn_gate(x, prepared).silu().contiguous() * self.ffn_up(x, prepared))
|
||||
|
||||
# given the token-prefix match, return how much cached state this block can still reuse
|
||||
def _reusable_prefix_len(self, prefix_len:int, cached_len:int) -> int: return prefix_len
|
||||
# return writes that reset this block's state after a cache mismatch
|
||||
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|None=None) -> 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) -> Tensor: raise NotImplementedError
|
||||
|
||||
def __call__(self, x: Tensor, start_pos: int|UOp, use_flash:bool=False, kv_len:int|None=None):
|
||||
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)
|
||||
if hasattr(self, 'ssm_a'):
|
||||
self.pending_state = None
|
||||
@function(precompile=True, allow_implicit=True)
|
||||
def _run_stateful(x:Tensor, start_pos:int|UOp, valid_len:int|UOp|None):
|
||||
h = x + self._attention(self.attn_norm(x), start_pos, use_flash, kv_len, valid_len)
|
||||
out = (h + self._feed_forward(self.ffn_norm(h))).contiguous()
|
||||
assert self.pending_state is not None
|
||||
return out, *self.pending_state
|
||||
out, conv_state, recurrent_state = _run_stateful(x, start_pos, valid_len)
|
||||
stores = (getattr(self, "conv_state").uop.store(conv_state.uop), getattr(self, "recurrent_state").uop.store(recurrent_state.uop))
|
||||
state = getattr(self, "recurrent_state").uop.after(*stores)
|
||||
return Tensor(out.uop.after(state))
|
||||
# we pass in the weights implicitly so we unpack the GGUF on the fly
|
||||
@function(precompile=True, allow_implicit=True)
|
||||
def _run(x:Tensor, start_pos:int|UOp):
|
||||
h = x + self._attention(self.attn_norm(x), start_pos, use_flash, kv_len)
|
||||
return (h + self._feed_forward(self.ffn_norm(h))).contiguous()
|
||||
return _run(x, start_pos)
|
||||
return function(precompile=True, allow_implicit=True)(_run)(x, start_pos)
|
||||
|
||||
class TransformerBlock(FFNBlock):
|
||||
def __init__(self, config:TransformerConfig):
|
||||
@@ -144,14 +368,16 @@ class TransformerBlock(FFNBlock):
|
||||
# --- attention projections (all linear, bias-free) ------------------
|
||||
q_proj_out = config.head_dim * config.n_heads * (2 if config.attn_output_gate else 1)
|
||||
kv_proj_out = config.head_dim * config.n_kv_heads
|
||||
self.attn_q = nn.Linear(config.dim, q_proj_out, bias=config.qkv_bias)
|
||||
self.attn_k = nn.Linear(config.dim, kv_proj_out, bias=config.qkv_bias)
|
||||
self.attn_v = nn.Linear(config.dim, kv_proj_out, bias=config.qkv_bias)
|
||||
self.attn_output = nn.Linear(config.head_dim * config.n_heads, config.dim, bias=False)
|
||||
self.attn_q = Linear(config.dim, q_proj_out, bias=config.qkv_bias)
|
||||
self.attn_k = Linear(config.dim, kv_proj_out, bias=config.qkv_bias)
|
||||
self.attn_v = Linear(config.dim, kv_proj_out, bias=config.qkv_bias)
|
||||
self.attn_output = Linear(config.head_dim * config.n_heads, config.dim, bias=False)
|
||||
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|None=None) -> Tensor:
|
||||
q, k, v = self.attn_q(x), self.attn_k(x), self.attn_v(x)
|
||||
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) -> Tensor:
|
||||
prepared = self.attn_q.prepare(x)
|
||||
q, k, v = self.attn_q(x, prepared), self.attn_k(x, prepared), self.attn_v(x, prepared)
|
||||
if self.config.qk_norm and self.config.qk_norm != self.config.head_dim: q, k = self.attn_q_norm(q), self.attn_k_norm(k)
|
||||
|
||||
B, T, _ = x.shape
|
||||
@@ -179,22 +405,32 @@ class TransformerBlock(FFNBlock):
|
||||
|
||||
# NOTE: this mask is causal_lower_right, not the causal_upper_left generated by is_casual = True
|
||||
# TODO: this if statement should be removed and it shouldn't generate extra kernels
|
||||
if use_flash:
|
||||
flash_decode = resolve(T == 1) and kv_len is not None and str(x.device).startswith("AMD") and self.config.head_dim == 256
|
||||
if flash_decode:
|
||||
decode_len = self.config.max_context
|
||||
decode_pos = (start_pos.unbind()[0] if isinstance(start_pos, UOp) else start_pos) + 1
|
||||
decode_mask = (Tensor.arange(decode_len) < Tensor(decode_pos)) \
|
||||
.where(0.0, float("-inf")).reshape(1, 1, 1, decode_len)
|
||||
attn = q.scaled_dot_product_attention(assigned_kv[0, :, :, :decode_len], assigned_kv[1, :, :, :decode_len],
|
||||
attn_mask=decode_mask, enable_gqa=True)
|
||||
elif use_flash:
|
||||
from extra.gemm.amd_flash_attention import amd_flash_attention_causal_cached
|
||||
flash_start_pos = start_pos.unbind()[0] if isinstance(start_pos, UOp) else start_pos
|
||||
valid_kv_len = flash_start_pos + T
|
||||
valid_kv_len = ((start_pos.unbind()[0] + 1) if isinstance(start_pos, UOp) else start_pos + 1) if flash_decode else \
|
||||
(start_pos.unbind()[0] if isinstance(start_pos, UOp) else start_pos) + T
|
||||
q_flat = q.half().reshape(B*self.config.n_heads, T, self.config.head_dim)
|
||||
out = Tensor.empty(B*self.config.n_heads, T, self.config.head_dim, dtype="float32", device=x.device)
|
||||
out = Tensor.empty(B*self.config.n_heads, q_flat.shape[1], self.config.head_dim, dtype="float32", device=x.device)
|
||||
attn = Tensor.custom_kernel(out, q_flat, assigned_kv,
|
||||
fxn=functools.partial(amd_flash_attention_causal_cached, valid_kv_len=valid_kv_len))[0].reshape(B, self.config.n_heads, T, -1)
|
||||
fxn=functools.partial(amd_flash_attention_causal_cached, valid_kv_len=valid_kv_len))[0] \
|
||||
.reshape(B, self.config.n_heads, q_flat.shape[1], -1)
|
||||
else:
|
||||
mask:Tensor|None
|
||||
if kv_len is not None:
|
||||
mask = Tensor.full((1, 1, 1, kv_len), float("-inf"), dtype=x.dtype, buffer=False).triu(start_pos+1)
|
||||
mask = None if resolve(T == 1) and self.config.ssm is not None else \
|
||||
Tensor.full((1, 1, 1, kv_len), float("-inf"), dtype=x.dtype, buffer=False).triu(start_pos+1)
|
||||
else:
|
||||
mask = Tensor.full((1, 1, T, start_pos+T), float("-inf"), dtype=x.dtype, buffer=False).triu(start_pos+1) \
|
||||
if resolve(T != 1) else None
|
||||
attn = q.half().scaled_dot_product_attention(k, v, attn_mask=mask, enable_gqa=True) # (B,H,T,Hd)
|
||||
attn = q.float().scaled_dot_product_attention(k.float(), v.float(), attn_mask=mask, enable_gqa=True) # (B,H,T,Hd)
|
||||
attn = attn.transpose(1, 2).reshape(B, T, -1) # back to (B,T,D)
|
||||
return self.attn_output(attn if not self.config.attn_output_gate else (attn * gate.sigmoid()))
|
||||
|
||||
@@ -203,27 +439,28 @@ class TransformerBlock(FFNBlock):
|
||||
# TODO: how is the dtype of this determined?
|
||||
# Decode uses fixed-size KV buckets. Unwritten entries must be zero: masking happens after QK, so values left
|
||||
# uninitialized by Tensor.empty can inject NaNs before the mask is applied.
|
||||
self.cache_kv = Tensor.zeros(2, x.shape[0], self.config.n_kv_heads, self.config.max_context, self.config.head_dim,
|
||||
self.cache_kv = Tensor.zeros(2, x.shape[0], self.config.n_kv_heads, self.config.max_context+192, self.config.head_dim,
|
||||
dtype="float16", device=x.device).contiguous()
|
||||
self.freqs_cis = precompute_freqs_cis(self.config.rope_dim, self.config.max_context, self.config.rope_theta, device=x.device)
|
||||
self.freqs_cis = precompute_freqs_cis(self.config.rope_dim, self.config.max_context+192, self.config.rope_theta, device=x.device)
|
||||
|
||||
class MLATransformerBlock(FFNBlock):
|
||||
def __init__(self, config:TransformerConfig):
|
||||
super().__init__(config)
|
||||
qk_nope_head_dim = config.head_dim - config.rope_dim
|
||||
if config.q_lora_rank > 0:
|
||||
self.attn_q_a = nn.Linear(config.dim, config.q_lora_rank, bias=False)
|
||||
self.attn_q_a = Linear(config.dim, config.q_lora_rank, bias=False)
|
||||
self.attn_q_a_norm = nn.RMSNorm(config.q_lora_rank, config.norm_eps)
|
||||
self.attn_q_b = nn.Linear(config.q_lora_rank, config.n_heads * config.head_dim, bias=False)
|
||||
self.attn_q_b = Linear(config.q_lora_rank, config.n_heads * config.head_dim, bias=False)
|
||||
else:
|
||||
self.attn_q = nn.Linear(config.dim, config.n_heads * config.head_dim, bias=False)
|
||||
self.attn_kv_a_mqa = nn.Linear(config.dim, config.kv_lora_rank + config.rope_dim, bias=False)
|
||||
self.attn_q = Linear(config.dim, config.n_heads * config.head_dim, bias=False)
|
||||
self.attn_kv_a_mqa = Linear(config.dim, config.kv_lora_rank + config.rope_dim, bias=False)
|
||||
self.attn_kv_a_norm = nn.RMSNorm(config.kv_lora_rank, config.norm_eps)
|
||||
self.attn_k_b = {"weight": Tensor.zeros(config.n_heads, config.kv_lora_rank, qk_nope_head_dim)}
|
||||
self.attn_v_b = {"weight": Tensor.zeros(config.n_heads, config.v_head_dim, config.kv_lora_rank)}
|
||||
self.attn_output = nn.Linear(config.n_heads * config.v_head_dim, config.dim, bias=False)
|
||||
self.attn_output = 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|None=None) -> Tensor:
|
||||
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) -> 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)
|
||||
@@ -251,8 +488,8 @@ class MLATransformerBlock(FFNBlock):
|
||||
|
||||
def _init_state(self, x:Tensor):
|
||||
if not hasattr(self, "cache_k"):
|
||||
self.cache_k = Tensor.empty(x.shape[0], 1, self.config.max_context, self.config.kv_lora_rank + self.config.rope_dim, device=x.device)
|
||||
self.freqs_cis = precompute_freqs_cis(self.config.rope_dim, self.config.max_context, self.config.rope_theta, device=x.device)
|
||||
self.cache_k = Tensor.empty(x.shape[0], 1, self.config.max_context+192, self.config.kv_lora_rank + self.config.rope_dim, device=x.device)
|
||||
self.freqs_cis = precompute_freqs_cis(self.config.rope_dim, self.config.max_context+192, self.config.rope_theta, device=x.device)
|
||||
|
||||
class GatedDeltaNetBlock(FFNBlock):
|
||||
def __init__(self, config:TransformerConfig, ssm:SSMConfig):
|
||||
@@ -261,44 +498,82 @@ class GatedDeltaNetBlock(FFNBlock):
|
||||
assert self.num_v_heads % self.num_k_heads == 0
|
||||
self.head_v_dim, self.ssm_conv_kernel = ssm.inner_size // ssm.time_step_rank, ssm.conv_kernel
|
||||
self.conv_channels, self.q_dim = ssm.inner_size + 2*ssm.group_count*ssm.state_size, ssm.state_size*ssm.group_count
|
||||
self.attn_qkv, self.attn_gate = nn.Linear(config.dim, self.conv_channels, bias=False), nn.Linear(config.dim, ssm.inner_size, bias=False)
|
||||
self.ssm_alpha, self.ssm_beta = nn.Linear(config.dim, self.num_v_heads, bias=False), nn.Linear(config.dim, self.num_v_heads, bias=False)
|
||||
self.attn_qkv, self.attn_gate = Linear(config.dim, self.conv_channels, bias=False), Linear(config.dim, ssm.inner_size, bias=False)
|
||||
self.ssm_alpha, self.ssm_beta = Linear(config.dim, self.num_v_heads, bias=False), Linear(config.dim, self.num_v_heads, bias=False)
|
||||
self.ssm_conv1d = {"weight": Tensor.zeros(self.conv_channels, self.ssm_conv_kernel)}
|
||||
self.ssm_dt = {"bias": Tensor.zeros(self.num_v_heads)}
|
||||
self.ssm_a = Tensor.zeros(self.num_v_heads)
|
||||
self.ssm_norm, self.ssm_out = nn.RMSNorm(self.head_v_dim, config.norm_eps), nn.Linear(ssm.inner_size, config.dim, bias=False)
|
||||
self.ssm_norm, self.ssm_out = nn.RMSNorm(self.head_v_dim, config.norm_eps), Linear(ssm.inner_size, config.dim, bias=False)
|
||||
|
||||
def _attention(self, x:Tensor, start_pos:int|UOp, use_flash:bool=False, kv_len:int|None=None) -> Tensor:
|
||||
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) -> Tensor:
|
||||
B, T, _ = x.shape
|
||||
assert T == 1, "GatedDeltaNetBlock currently only supports T=1"
|
||||
conv_state, initial_state = self.conv_state, self.recurrent_state
|
||||
|
||||
# input processing
|
||||
if T == 1:
|
||||
x = x.half()
|
||||
prepared = self.attn_gate.prepare(x)
|
||||
out_gate = self.attn_gate(x, prepared).reshape(B, 1, self.num_v_heads, self.head_v_dim)
|
||||
beta = self.ssm_beta(x, prepared).sigmoid().reshape(B, self.num_v_heads, 1, 1)
|
||||
alpha = ((self.ssm_alpha(x, prepared).float() + self.ssm_dt["bias"]).softplus() * self.ssm_a).reshape(B, self.num_v_heads, 1, 1).exp()
|
||||
conv_window = conv_state.cat(self.attn_qkv(x, prepared), dim=1)
|
||||
conv_out = (conv_window * self.ssm_conv1d["weight"].T.unsqueeze(0)).sum(1).silu()
|
||||
q, k, v = conv_out.split([self.q_dim, self.q_dim, self.conv_channels - 2*self.q_dim], dim=-1)
|
||||
q = l2norm(q.reshape(B, self.num_k_heads, self.head_k_dim)).repeat(1, self.num_v_heads//self.num_k_heads, 1)
|
||||
k = l2norm(k.reshape(B, self.num_k_heads, self.head_k_dim)).repeat(1, self.num_v_heads//self.num_k_heads, 1)
|
||||
v = v.reshape(B, self.num_v_heads, self.head_v_dim)
|
||||
q, k, v = q.mul(self.head_k_dim**-0.5).unsqueeze(-1), k.unsqueeze(-1), v.unsqueeze(-1)
|
||||
recurrent_state = initial_state * alpha
|
||||
recurrent_state = recurrent_state + ((v - recurrent_state@k) * beta)@k.transpose(-1, -2)
|
||||
self.pending_state = (conv_window[:, 1:, :].cast(self.conv_state.dtype).contiguous(),
|
||||
recurrent_state.cast(self.recurrent_state.dtype).contiguous())
|
||||
core_attn_out = self.ssm_norm((recurrent_state@q).squeeze(-1).reshape(B, 1, self.num_v_heads, self.head_v_dim))
|
||||
return self.ssm_out((core_attn_out * out_gate.silu()).reshape(B, 1, -1).cast(x.dtype))
|
||||
|
||||
# Batched projections and causal depthwise convolution.
|
||||
x = x.half()
|
||||
out_gate = self.attn_gate(x).reshape(B, 1, self.num_v_heads, self.head_v_dim)
|
||||
beta = self.ssm_beta(x).sigmoid().reshape(B, self.num_v_heads, 1, 1)
|
||||
alpha = ((self.ssm_alpha(x).float() + self.ssm_dt["bias"]).softplus() * self.ssm_a).reshape(B, self.num_v_heads, 1, 1).exp()
|
||||
|
||||
# qkv conv
|
||||
conv_window = self.conv_state.cat(self.attn_qkv(x), dim=1)
|
||||
conv_out = (conv_window * self.ssm_conv1d["weight"].T.unsqueeze(0)).sum(1).silu()
|
||||
prepared = self.attn_gate.prepare(x)
|
||||
out_gate = self.attn_gate(x, prepared).reshape(B, T, self.num_v_heads, self.head_v_dim)
|
||||
beta = self.ssm_beta(x, prepared).sigmoid().reshape(B, T, self.num_v_heads)
|
||||
log_alpha = ((self.ssm_alpha(x, prepared).float() + self.ssm_dt["bias"]).softplus() * self.ssm_a).reshape(B, T, self.num_v_heads)
|
||||
if valid_len is not None:
|
||||
active = (Tensor.arange(T) < Tensor(valid_len)).reshape(1, T, 1)
|
||||
beta, log_alpha = beta * active, log_alpha * active
|
||||
conv_window = conv_state.cat(self.attn_qkv(x, prepared), dim=1)
|
||||
conv_out = functools.reduce(lambda a,b: a+b,
|
||||
(conv_window[:, i:i+T] * self.ssm_conv1d["weight"][:, i] for i in range(self.ssm_conv_kernel))).silu()
|
||||
q, k, v = conv_out.split([self.q_dim, self.q_dim, self.conv_channels - 2*self.q_dim], dim=-1)
|
||||
q = q.reshape(B, self.num_k_heads, self.head_k_dim).normalize(dim=-1).repeat(1, self.num_v_heads//self.num_k_heads, 1)
|
||||
k = k.reshape(B, self.num_k_heads, self.head_k_dim).normalize(dim=-1).repeat(1, self.num_v_heads//self.num_k_heads, 1)
|
||||
v = v.reshape(B, self.num_v_heads, self.head_v_dim)
|
||||
q, k, v = q.mul(self.head_k_dim**-0.5).unsqueeze(-1), k.unsqueeze(-1), v.unsqueeze(-1)
|
||||
q = l2norm(q.reshape(B, T, self.num_k_heads, self.head_k_dim)).repeat(1, 1, self.num_v_heads//self.num_k_heads, 1)
|
||||
k = l2norm(k.reshape(B, T, self.num_k_heads, self.head_k_dim)).repeat(1, 1, self.num_v_heads//self.num_k_heads, 1)
|
||||
v = v.reshape(B, T, self.num_v_heads, self.head_v_dim)
|
||||
|
||||
# recurrent
|
||||
recurrent_state = self.recurrent_state * alpha
|
||||
recurrent_state = recurrent_state + ((v - recurrent_state@k) * beta)@k.transpose(-1, -2)
|
||||
# Chunked gated delta rule. The strictly-lower update is the triangular solve from the reference implementation.
|
||||
q, k, v, beta, log_alpha = [z.transpose(1, 2).float() for z in (q, k, v, beta, log_alpha)]
|
||||
q = q * self.head_k_dim**-0.5
|
||||
state = initial_state.transpose(-1, -2).float()
|
||||
core_chunks = []
|
||||
for start in range(0, T, 64):
|
||||
qc, kc, vc, bc, gc = q[:,:,start:start+64], k[:,:,start:start+64], v[:,:,start:start+64], \
|
||||
beta[:,:,start:start+64], log_alpha[:,:,start:start+64]
|
||||
chunk_len = qc.shape[2]
|
||||
g = (gc @ Tensor.ones(chunk_len, chunk_len, dtype=gc.dtype).tril().T).contiguous()
|
||||
decay = (g.unsqueeze(-1) - g.unsqueeze(-2)).exp().tril().contiguous()
|
||||
base = (-(kc * bc.unsqueeze(-1) @ kc.transpose(-1, -2) * decay).tril(-1)).contiguous()
|
||||
attn = inverse_unit_lower(base)
|
||||
value = attn @ (vc * bc.unsqueeze(-1))
|
||||
k_cumdecay = attn @ (kc * bc.unsqueeze(-1) * g.exp().unsqueeze(-1))
|
||||
value = value - k_cumdecay @ state
|
||||
core_chunks.append((qc * g.exp().unsqueeze(-1)) @ state + (qc @ kc.transpose(-1, -2) * decay) @ value)
|
||||
state = state * g[..., -1, None, None].exp() + \
|
||||
(kc * (g[..., -1, None] - g).exp().unsqueeze(-1)).transpose(-1, -2) @ value
|
||||
core_attn_out = functools.reduce(lambda a,b: a.cat(b, dim=2), core_chunks)
|
||||
core_attn_out = self.ssm_norm(core_attn_out.transpose(1, 2))
|
||||
out = self.ssm_out((core_attn_out * out_gate.silu()).reshape(B, T, -1).cast(x.dtype)).contiguous()
|
||||
|
||||
# store the updated state
|
||||
conv_state_store = self.conv_state.uop.store(conv_window[:, 1:, :].cast(self.conv_state.dtype).uop)
|
||||
recurrent_state_store = self.recurrent_state.uop.store(recurrent_state.cast(self.recurrent_state.dtype).uop)
|
||||
recurrent_state = Tensor(self.recurrent_state.uop.after(recurrent_state_store, conv_state_store))
|
||||
|
||||
# output
|
||||
core_attn_out = self.ssm_norm((recurrent_state@q).squeeze(-1).reshape(B, 1, self.num_v_heads, self.head_v_dim))
|
||||
return self.ssm_out((core_attn_out * out_gate.silu()).reshape(B, 1, -1).cast(x.dtype))
|
||||
state_pos = T if valid_len is None else valid_len
|
||||
self.pending_state = (conv_window[:, state_pos:state_pos+self.ssm_conv_kernel-1, :].cast(self.conv_state.dtype).contiguous(),
|
||||
state.transpose(-1, -2).cast(self.recurrent_state.dtype).contiguous())
|
||||
return out
|
||||
|
||||
# recurrent state can't be partially reused after divergence, force a full rebuild
|
||||
def _state_reset_ops(self):
|
||||
@@ -320,32 +595,52 @@ class Transformer:
|
||||
block_cls(dense_config if i < config.leading_dense_blocks else config) for i in range(config.num_blocks)]
|
||||
self.token_embd = nn.Embedding(config.vocab_size, config.dim)
|
||||
self.output_norm = nn.RMSNorm(config.dim, config.norm_eps)
|
||||
self.output = nn.Linear(config.dim, config.vocab_size, bias=False)
|
||||
self.output = Linear(config.dim, config.vocab_size, bias=False)
|
||||
self.max_context = config.max_context
|
||||
self.parameter_count = 0
|
||||
self.has_recurrent_block = any(isinstance(b, GatedDeltaNetBlock) for b in self.blk)
|
||||
self._cached_tokens: list[int] = []
|
||||
# we specialize the JIT for prefill and rollout
|
||||
self.prefill_jit = TinyJit(self.forward)
|
||||
self.flash_prefill_jit = TinyJit(functools.partial(self.forward, use_flash=True))
|
||||
self.sample_prefill_jit = TinyJit(functools.partial(self.forward, sample=True))
|
||||
self.rollout_jits:dict[int, TinyJit] = {}
|
||||
self.sample_rollout_jits:dict[int, TinyJit] = {}
|
||||
|
||||
def forward(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor, use_flash:bool=False, kv_len:int|None=None) -> Tensor:
|
||||
def forward(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor, use_flash:bool=False, kv_len:int|UOp|None=None,
|
||||
valid_len:int|UOp|None=None, sample:bool=False) -> Tensor:
|
||||
x = self.token_embd(tokens).float() # (B, T, D)
|
||||
for block in self.blk: x = block(x, start_pos, use_flash, kv_len)
|
||||
logits = self.output(self.output_norm(x))[:, -1, :]
|
||||
for block in self.blk:
|
||||
x = block(x, start_pos, use_flash, kv_len, valid_len)
|
||||
last = x[:, -1:] if valid_len is None else x[:, valid_len-1:valid_len]
|
||||
logits = self.output(self.output_norm(last))[:, -1, :]
|
||||
# Gumbel-max trick: argmax(logits/temp - log(-log(uniform))) is equivalent to sampling from softmax(logits/temp)
|
||||
return (logits / temperature.maximum(1e-12) - (Tensor.rand_like(logits).maximum(1e-12).log().neg()).log()).argmax(-1, keepdim=True)
|
||||
if not sample: return logits.argmax(-1, keepdim=True)
|
||||
return (logits / temperature - (Tensor.rand_like(logits).maximum(1e-12).log().neg()).log()).argmax(-1, keepdim=True)
|
||||
|
||||
def __call__(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor, use_flash:bool=False) -> Tensor:
|
||||
def forward_recurrent_decode(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor, valid_len:int|UOp|None=None,
|
||||
sample:bool=False) -> Tensor:
|
||||
return self.forward(tokens, start_pos, temperature, kv_len=start_pos+1, valid_len=valid_len, sample=sample)
|
||||
|
||||
def __call__(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor, use_flash:bool=False,
|
||||
valid_len:int|UOp|None=None, sample:bool=False) -> Tensor:
|
||||
jit_kwargs = {"valid_len":valid_len}
|
||||
if resolve(tokens.shape[1] == 1):
|
||||
pos = start_pos.unbind()[1] if isinstance(start_pos, UOp) else start_pos
|
||||
min_bucket = max(1, getenv("DECODE_BUCKET", 256))
|
||||
kv_len = min(self.max_context, max(min_bucket, 1 << pos.bit_length()))
|
||||
if kv_len not in self.rollout_jits: self.rollout_jits[kv_len] = TinyJit(functools.partial(self.forward, kv_len=kv_len))
|
||||
jit = self.rollout_jits[kv_len]
|
||||
if self.has_recurrent_block:
|
||||
key = 0
|
||||
else:
|
||||
min_bucket = max(1, getenv("DECODE_BUCKET", 256))
|
||||
kv_len = key = min(self.max_context, max(min_bucket, 1 << pos.bit_length()))
|
||||
rollout_jits = self.sample_rollout_jits if sample else self.rollout_jits
|
||||
if key not in rollout_jits:
|
||||
rollout_jits[key] = TinyJit(functools.partial(self.forward_recurrent_decode, sample=sample) if self.has_recurrent_block else
|
||||
functools.partial(self.forward, kv_len=kv_len, sample=sample))
|
||||
jit = rollout_jits[key]
|
||||
else:
|
||||
jit = self.flash_prefill_jit if use_flash else self.prefill_jit
|
||||
return jit(tokens.contiguous(), start_pos, temperature)
|
||||
jit = self.sample_prefill_jit if sample else self.flash_prefill_jit if use_flash else self.prefill_jit
|
||||
ret = jit(tokens.contiguous(), start_pos, temperature, **jit_kwargs)
|
||||
return ret[0] if isinstance(ret, tuple) else ret
|
||||
|
||||
@staticmethod
|
||||
def from_gguf(gguf:Tensor|str|pathlib.Path, max_context:int|None=None,
|
||||
@@ -353,9 +648,6 @@ class Transformer:
|
||||
# TODO: remove the need for copy to default device
|
||||
kv, state_dict = gguf_load(gguf.to(None).realize() if isinstance(gguf, Tensor) else gguf)
|
||||
|
||||
# all state items should be float16, not float32
|
||||
state_dict = {k:v.cast('float16') if getenv("HALF", 1) else v for k,v in state_dict.items()}
|
||||
|
||||
# some models like Llama 3.2 don't have an output.weight, they just tie to the token_embd.weight
|
||||
if 'output.weight' not in state_dict: state_dict['output.weight'] = state_dict['token_embd.weight']
|
||||
|
||||
@@ -409,7 +701,29 @@ class Transformer:
|
||||
qkv_bias='blk.0.attn_q.bias' in state_dict,
|
||||
expert_bias=f"blk.{kv.get(f'{arch}.leading_dense_block_count', 0)}.exp_probs_b.bias" in state_dict)
|
||||
model = Transformer(config)
|
||||
model.parameter_count = sum(int(weight.numel()) for weight in state_dict.values())
|
||||
packed_weights:set[str] = set()
|
||||
def resolve_owner(path:list[str]):
|
||||
obj = model
|
||||
for part in path: obj = obj[int(part)] if isinstance(obj, list) else getattr(obj, part)
|
||||
return obj
|
||||
for name, weight in state_dict.items():
|
||||
parts = name.split('.')
|
||||
quantization = get_ggml_quantization(weight)
|
||||
if quantization is not None and quantization[1] == 8 and parts[-1] == "weight" and isinstance(owner:=resolve_owner(parts[:-1]), Linear):
|
||||
owner.set_quantized(*quantization)
|
||||
state_dict[name], packed_weights = owner.weight, packed_weights | {name}
|
||||
elif len(parts) == 4 and parts[0] == "blk" and parts[2].endswith("_exps") and parts[3] == "weight" and quantization is not None:
|
||||
expert_weights = getattr(model.blk[int(parts[1])], parts[2])
|
||||
expert_weights.set_quantized(weight, *quantization)
|
||||
state_dict[name], packed_weights = expert_weights.weight, packed_weights | {name}
|
||||
|
||||
state_dict = {k:v if k in packed_weights else v.cast('float16') if getenv("HALF", 1) else v for k,v in state_dict.items()}
|
||||
nn.state.load_state_dict(model, state_dict, verbose=False, consume=True, realize=False) # NOTE: rope_freqs.weight (32,) is unused
|
||||
expert_types = {getattr(block, name).ggml_type for block in model.blk if hasattr(block, "ffn_gate_exps")
|
||||
for name in ("ffn_gate_exps", "ffn_down_exps")}
|
||||
for ggml_type in expert_types:
|
||||
if ggml_type in (21, 23) and str(model.token_embd.weight.device).startswith("AMD"): _expert_lut(str(model.token_embd.weight.device), ggml_type)
|
||||
# NOTE: without this contiguous, it unpacks the weights from the model every time. we shouldn't need this, but for now it's faster
|
||||
if realize:
|
||||
for s in (params:=nn.state.get_parameters(model)): s.replace(s.contiguous())
|
||||
@@ -421,10 +735,28 @@ class Transformer:
|
||||
return min(block._reusable_prefix_len(prefix_len, len(self._cached_tokens)) for block in self.blk)
|
||||
|
||||
def warmup(self, chunk_size:int=256):
|
||||
direct_capture = not self.has_recurrent_block and all(isinstance(block, TransformerBlock) for block in self.blk)
|
||||
if direct_capture:
|
||||
device = str(self.token_embd.weight.device)
|
||||
direct_capture = device.startswith("AMD") and Device[device].renderer.target.arch.startswith("gfx11")
|
||||
|
||||
# Capture both prefill JITs. Different first tokens prevent the second pass from reusing the first pass's KV cache.
|
||||
warm_len = min(1 if self.has_recurrent_block else chunk_size * 2, self.max_context - 1)
|
||||
recurrent_chunk = min(chunk_size, 192)
|
||||
warm_len = min(recurrent_chunk * 3 + 1 if self.has_recurrent_block else chunk_size * 2, self.max_context - 1)
|
||||
if warm_len > 0:
|
||||
for salt in range(2): next(self.generate([salt] + [0] * (warm_len - 1), chunk_size=chunk_size))
|
||||
if direct_capture:
|
||||
x = Tensor.zeros(1, 1, self.blk[0].config.dim)
|
||||
for block in self.blk: block._init_state(x)
|
||||
Tensor.realize(*[state for block in self.blk for state in (getattr(block, "cache_kv"), getattr(block, "freqs_cis"))])
|
||||
self.prefill_jit.cnt = self.flash_prefill_jit.cnt = 1
|
||||
next(self.generate([0] * warm_len, chunk_size=chunk_size))
|
||||
elif self.has_recurrent_block:
|
||||
warm = self.generate([0] * warm_len, chunk_size=chunk_size)
|
||||
next(warm)
|
||||
next(warm)
|
||||
next(warm)
|
||||
else:
|
||||
for salt in range(2): next(self.generate([salt] + [0] * (warm_len - 1), chunk_size=chunk_size))
|
||||
|
||||
# Rollout uses fixed power-of-two KV shapes. Capture every shape up front so requests never pay a JIT transition.
|
||||
if not self.has_recurrent_block:
|
||||
@@ -435,38 +767,49 @@ class Transformer:
|
||||
bucket_positions.setdefault(bucket, pos)
|
||||
v_start_pos = UOp.variable("start_pos", 0, self.max_context-1)
|
||||
token, temperature = Tensor([[0]], dtype="int32"), Tensor([0.0])
|
||||
for _, pos in sorted(bucket_positions.items()):
|
||||
for _ in range(2): self(token, v_start_pos.bind(pos), temperature).realize()
|
||||
for bucket, pos in sorted(bucket_positions.items()):
|
||||
if direct_capture:
|
||||
self.rollout_jits[bucket] = TinyJit(functools.partial(self.forward, kv_len=bucket))
|
||||
self.rollout_jits[bucket].cnt = 1
|
||||
for _ in range(1 if direct_capture else 2):
|
||||
result = self(token, v_start_pos.bind(pos), temperature)
|
||||
assert isinstance(result, Tensor)
|
||||
result.realize()
|
||||
|
||||
if resets := [r for block in self.blk for r in block._state_reset_ops()]: Tensor.realize(*resets)
|
||||
self._cached_tokens = []
|
||||
|
||||
def generate(self, tokens:list[int], chunk_size:int=256, temperature:float=0.0):
|
||||
if self.has_recurrent_block: chunk_size = 1
|
||||
if self.has_recurrent_block: chunk_size = min(chunk_size, 192)
|
||||
v_start_pos = UOp.variable("start_pos", 0, self.max_context-1)
|
||||
v_toks = UOp.variable("toks", 1, chunk_size)
|
||||
# TODO: use UOp.variable for temperature once float variables are supported
|
||||
temp = Tensor([temperature])
|
||||
# assign all input tokens once, then slice from start_pos for the model call
|
||||
t = Tensor(tokens + [0] * (self.max_context - len(tokens)), dtype="int32").reshape(1, self.max_context)
|
||||
t = Tensor(tokens + [0] * (self.max_context + chunk_size - len(tokens)), dtype="int32").reshape(1, self.max_context + chunk_size)
|
||||
# recompute start_pos from what's currently valid in the caches
|
||||
start_pos = self.get_start_pos(tokens)
|
||||
if start_pos < len(self._cached_tokens) and (resets := [r for b in self.blk for r in b._state_reset_ops()]): Tensor.realize(*resets)
|
||||
out, prompt_len = None, len(tokens)
|
||||
while len(tokens) < self.max_context:
|
||||
remaining = len(tokens) - start_pos
|
||||
can_flash = bool(getenv("AMD_FLASH_ATTENTION", 1)) and start_pos > 0 and remaining >= chunk_size and chunk_size % 64 == 0 and \
|
||||
not self.has_recurrent_block
|
||||
can_flash = bool(getenv("AMD_FLASH_ATTENTION", 1)) and start_pos > 0 and remaining >= chunk_size and chunk_size % 64 == 0
|
||||
if can_flash:
|
||||
device = str(getattr(self.blk[0], "cache_kv").device)
|
||||
device = str(self.token_embd.weight.device)
|
||||
can_flash = device.startswith("AMD") and Device[device].renderer.target.arch.startswith("gfx11")
|
||||
use_flash = can_flash and start_pos % 64 == 0
|
||||
sp = v_start_pos.bind(start_pos)
|
||||
# The flash kernel requires its cached prefix to start on a 64-token tile. Cache reuse can resume at any
|
||||
# token, so process one short generic chunk to reach the next tile boundary before entering flash prefill.
|
||||
nt = chunk_size if use_flash else v_toks.bind(min(64 - start_pos % 64, remaining) if can_flash else min(chunk_size, remaining))
|
||||
actual_nt = min(chunk_size, remaining)
|
||||
nt = chunk_size if use_flash or self.has_recurrent_block and start_pos < prompt_len else 1 if self.has_recurrent_block else \
|
||||
v_toks.bind(min(64 - start_pos % 64, remaining) if can_flash else actual_nt)
|
||||
inp = t[:, sp:sp+nt] if start_pos < prompt_len or out is None else out
|
||||
out = (self(inp, sp, temp, use_flash=True) if use_flash else self(inp, sp, temp)).realize()
|
||||
start_pos += nt if isinstance(nt, int) else nt.val
|
||||
valid_len = v_toks.bind(actual_nt) if self.has_recurrent_block and nt == chunk_size else None
|
||||
result = self(inp, sp, temp, use_flash=True, valid_len=valid_len, sample=temperature > 0) if use_flash else \
|
||||
self(inp, sp, temp, valid_len=valid_len, sample=temperature > 0)
|
||||
out = result.realize()
|
||||
start_pos += actual_nt if self.has_recurrent_block else nt if isinstance(nt, int) else nt.val
|
||||
# chunked prefill: keep processing until all prompt tokens are consumed
|
||||
if start_pos < len(tokens): continue
|
||||
tokens.append(int(out.item()))
|
||||
|
||||
+10
-3
@@ -67,6 +67,7 @@ class Handler(HTTPRequestHandler):
|
||||
else: self.send_data((pathlib.Path(__file__).parent / "chat.html").read_bytes(), content_type="text/html")
|
||||
def run_model(self, ids:list[int], model_name:str, include_usage=False, max_tokens:int|None=None, temperature:float=0.0):
|
||||
model, tok = self.server.model, self.server.tok
|
||||
prompt_tokens = len(ids)
|
||||
cache_start_pos = model.get_start_pos(ids)
|
||||
stderr_log(f"in:{colored(f'{cache_start_pos:5d}', 'green')} +{len(ids)-cache_start_pos:5d} {colored('--', 'BLACK')} ")
|
||||
tmpl = {"id":f"chatcmpl-{uuid.uuid4().hex[:24]}", "object":"chat.completion.chunk", "created":int(time.time()), "model":model_name}
|
||||
@@ -78,7 +79,7 @@ class Handler(HTTPRequestHandler):
|
||||
dec = tok.stream_decoder()
|
||||
router = StreamRouter()
|
||||
for next_id in model.generate(ids, temperature=temperature):
|
||||
if len(out) == 0: stderr_log(f"prefill:{(len(ids)-cache_start_pos)/((pt:=time.perf_counter())-st):4.0f} tok/s {colored('--', 'BLACK')} ")
|
||||
if len(out) == 0: stderr_log(f"prefill:{(prompt_tokens-cache_start_pos)/((pt:=time.perf_counter())-st):4.0f} tok/s {colored('--', 'BLACK')} ")
|
||||
if tok.is_end(next_id): break
|
||||
out.append(next_id)
|
||||
for field, delta in router.route(dec(next_id)): yield chunk({field:delta})
|
||||
@@ -100,7 +101,8 @@ class Handler(HTTPRequestHandler):
|
||||
if finish_reason == "stop": finish_reason = "tool_calls"
|
||||
yield {"choices": [{"index":0, "delta":{},"finish_reason":finish_reason}], **tmpl}
|
||||
if include_usage:
|
||||
yield {"choices": [], "usage": {"prompt_tokens": len(ids), "completion_tokens": len(out), "total_tokens": len(ids) + len(out)}, **tmpl}
|
||||
yield {"choices": [], "usage": {"prompt_tokens": prompt_tokens, "completion_tokens": len(out),
|
||||
"total_tokens": prompt_tokens + len(out)}, **tmpl}
|
||||
et = time.perf_counter()
|
||||
stderr_log(f"gen:{len(out)/(et-pt) if len(out) > 1 else 0:4.0f} tok/s {colored('--', 'BLACK')} "
|
||||
f"out:{len(out):5d} {colored('--', 'BLACK')} total:{et-st:6.2f}s\n")
|
||||
@@ -114,9 +116,14 @@ class Handler(HTTPRequestHandler):
|
||||
if self.path == "/v1/chat/completions":
|
||||
# render and tokenize
|
||||
normalize_messages(body["messages"])
|
||||
rendered = self.server.template.render(messages=body["messages"], tools=body.get("tools"), add_generation_prompt=True)
|
||||
rendered = self.server.template.render(messages=body["messages"], tools=body.get("tools"), add_generation_prompt=True,
|
||||
enable_thinking=body.get("enable_thinking", False))
|
||||
ids: list[int] = self.server.tok.encode(rendered)
|
||||
stderr_log(f"prep:{(time.perf_counter()-request_st)*1e3:5.0f} ms {colored('--', 'BLACK')} ")
|
||||
if len(ids) >= self.server.model.max_context:
|
||||
return self.send_data(json.dumps({"error":{"message":f"prompt has {len(ids)} tokens, but the model context is "
|
||||
f"{self.server.model.max_context}", "type":"invalid_request_error", "param":"messages", "code":"context_length_exceeded"}}).encode(),
|
||||
status_code=400)
|
||||
|
||||
# reply
|
||||
max_tokens = body.get("max_completion_tokens") or body.get("max_tokens")
|
||||
|
||||
Reference in New Issue
Block a user