llm: prepare Kimi K3 and accelerate recurrent prefill

This commit is contained in:
2026-08-10 02:03:14 +00:00
parent 6edb5f9698
commit 77e5be99bc
15 changed files with 956 additions and 80 deletions
+109
View File
@@ -0,0 +1,109 @@
# Kimi K3 on 8× MI350X
This branch targets text generation from the official `moonshotai/Kimi-K3` checkpoint. It intentionally ignores the vision tower and multimodal projector. The checkpoint remains in its official 96-shard format; no conversion or second 1.56 TB copy is required.
The checked TP8 layout consumes 196.78 GB (183.27 GiB) of text weights per GPU. The compressed MLA cache adds 28.99 GB (27 GiB) per GPU at the full 1,048,576-token context, leaving approximately 62.23 GB of each nominal 288 GB MI350X for execution buffers and allocator overhead. Start much smaller.
## Before renting the machine
- Reserve at least 1.7 TB of local model storage. More headroom is preferable for download caches and logs.
- The host should have roughly 3 TB RAM, in line with AMD's MI350X platform guidance. The loader itself is streaming and must not need checkpoint-sized RAM.
- Use a recent kernel/ROCm stack supported by the host vendor, although tinygrad uses its own AMD userspace driver when `DEV=AMD`.
- Clone this exact commit/branch and keep the official checkpoint directory separate from the repository.
Download on a machine with the storage bandwidth and network allocation intended for the run:
```sh
hf download moonshotai/Kimi-K3 --local-dir /models/Kimi-K3
python examples/kimi_k3_prepare.py /models/Kimi-K3 --context 4096
```
For a metadata-only preflight, place the official `config.json` and `model.safetensors.index.json` in a directory and run:
```sh
python examples/kimi_k3_prepare.py /models/Kimi-K3-metadata --metadata-only
```
## Hardware admission checks
Do these before loading weights. Stop if any device is missing or reports a different architecture.
```sh
lspci -d 1002:75a0
amd-smi list
DEV=AMD DEBUG=2 python - <<'PY'
from tinygrad import Device
for i in range(8):
dev = Device[f"AMD:{i}"]
print(i, dev.arch)
PY
```
Expected architecture: `gfx950` on all eight devices. Then run the small TP8 graph tests:
```sh
python -m pytest test/unit/test_llm_k3.py test/null/test_kimi_k3.py -q -n12
DEV=NULL:HIP:gfx950 NULL_ALLOW_COPYOUT=1 python -m pytest \
test/unit/test_llm_k3.py::TestKimiK3::test_chunked_recurrent_generate -q -n1
DEV=AMD python examples/kimi_k3_smoke.py --devices 8
```
The last two commands are deliberately small. They compile CDNA4 kernels and then exercise the complete TP8 topology without loading the checkpoint.
## First official load
Start at a short context so cache allocation and compilation are bounded. The loader reads disk-backed safetensors, TP-shards every destination before realizing it, and drops each source shard/projection immediately afterward.
```sh
/usr/bin/time -v env DEV=AMD DEBUG=1 python -m tinygrad.llm.cli \
--model /models/Kimi-K3 --devices 8 --max_context 128 </dev/null 2>&1 | tee kimi-k3-load.log
```
Watch host RAM, swap, HBM, temperatures, and XGMI traffic from a second terminal. Do not start with a one-million-token cache. If loading fails, preserve the first exception and the last loader progress line; do not retry with a larger host-side cache.
## Correctness and performance sequence
1. Load with context 128 and generate one token.
2. Repeat a fixed prompt twice and confirm token-for-token deterministic greedy output.
3. Compare the first several greedy tokens against the official Transformers implementation at temperature zero.
4. Benchmark decode only after two warm-up tokens.
5. Benchmark prefill at 128, 512, 2K, and 8K tokens. Increase context only while HBM and compile time remain healthy.
6. Use `VIZ=1` plus `python -m tinygrad.viz.cli` to inspect kernels; use `VIZ=2` only for short SQTT captures because it adds overhead.
Example decode benchmark:
```sh
DEV=AMD DEBUG=1 python -m tinygrad.llm.cli --model /models/Kimi-K3 \
--devices 8 --max_context 4096 --warmup --benchmark 20
```
## Known hardware-only gate
The correctness path expands only selected MXFP4 expert weights and emulates MXFP8 activation quantization. tinygrad has gfx950/CDNA4 BF16 and FP8 matrix-core support, but this branch does not yet have a hardware-validated fused native MXFP4×MXFP8 expert GEMM. Expect the first run to be a correctness bring-up, not production throughput. Capture profiles on MI350X before changing the representation: native FP4 work cannot be validated faithfully on the available gfx1100 cards.
Recurrent prefill is fused and defaults to eight-token chunks. The portable custom kernel compiles for gfx950, but the wave-parallel version is deliberately restricted to gfx11 because it uses RDNA3 wave32 swizzles. A gfx950-tuned wave64/MFMA recurrent kernel remains a performance task for the rented machine; do not enable the gfx11 kernel on CDNA without rewriting its lane reduction and validating every recurrent-state transition.
The official checkpoint also contains MoonViT-V2 and multimodal projector weights. They are skipped by the text loader. Image input remains a separate implementation and validation task.
## Local TP4 performance baseline
The pre-rental benchmark uses the converted `Kimi-Linear-48B-A3B-Instruct-MXFP4-v2` checkpoint on four gfx1100 GPUs. It is a useful regression test for the KDA/MLA/MoE text path, not a projection of K3 throughput on MI350X.
```sh
DEV=AMD python extra/benchmark_kimi.py \
/home/tiny/models/Kimi-Linear-48B-A3B-Instruct-MXFP4-v2 \
--devices 4 --max-context 128 --prompt-tokens 32 --decode-tokens 8 --chunk-size 8
```
Results from 2026-08-10:
- load: 310.280s for 29.27 GB
- cold prefill including compilation: 47.622s
- captured prefill: 1.020s, 31.38 tok/s
- steady prefill replay: 0.898s, 35.64 tok/s
- steady decode replay: 23.03 tok/s, 43.43 ms/token
- peak host RSS: 901,192 KiB; swap: 0
The same resident model measured 19.60 tok/s with tokenwise prefill, so the selected eight-token chunk is 1.82× faster. A 32-token chunk fell to 3.03 tok/s because the current selected-expert path expands packed weights per token; larger batches multiply that temporary dequantization work. Chunk 8 is therefore a conservative default until native grouped MXFP4×MXFP8 expert GEMM exists.
Different chunk sizes can choose a different final token because their matrix kernels use different floating-point reduction orders. Each measured shape was repeatable between cold and captured execution. For official K3 validation, compare logits/tokens against the reference at one fixed chunk size and greedy settings rather than requiring bitwise agreement between performance shapes.
+27
View File
@@ -0,0 +1,27 @@
#!/usr/bin/env python3
"""Cheap preflight for an official moonshotai/Kimi-K3 checkout. Does not load model weights."""
import argparse, json, pathlib, shutil
from tinygrad.llm.kimi_k3 import KIMI_K3_TP8_BYTES_PER_GPU, audit_kimi_k3_checkpoint
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("model_dir", type=pathlib.Path)
parser.add_argument("--metadata-only", action="store_true", help="permit absent weight shards")
parser.add_argument("--context", type=int, default=4096, help="context length used for the memory estimate")
args = parser.parse_args()
stats = audit_kimi_k3_checkpoint(args.model_dir, require_shards=not args.metadata_only)
if not 1 <= args.context <= 1_048_576: raise ValueError("--context must be between 1 and 1048576")
# K3 has 24 MLA layers. Each token stores the 512-value compressed latent plus 64 RoPE values in BF16.
per_gpu_weights = KIMI_K3_TP8_BYTES_PER_GPU
mla_cache = 24 * args.context * (512 + 64) * 2
hbm = 288_000_000_000
print(json.dumps(stats, indent=2))
print(f"exact text weights/GPU under this TP8 layout: {per_gpu_weights/1e9:.2f} GB ({per_gpu_weights/2**30:.2f} GiB)")
print(f"replicated MLA cache/GPU at {args.context:,} tokens: {mla_cache/1e9:.2f} GB ({mla_cache/2**30:.2f} GiB)")
print(f"nominal MI350X headroom before runtime buffers: {(hbm-per_gpu_weights-mla_cache)/1e9:.2f} GB")
if not args.metadata_only:
usage = shutil.disk_usage(args.model_dir)
print(f"filesystem free space: {usage.free/1e9:.2f} GB")
if __name__ == "__main__": main()
+27
View File
@@ -0,0 +1,27 @@
#!/usr/bin/env python3
"""Run a reduced, architecture-complete K3 prefill/decode on tensor-parallel devices."""
import argparse, time
from tinygrad import Tensor, Device, dtypes, nn
from tinygrad.llm.kimi_k3 import _shard_kimi_k3, kimi_k3_smoke_config
from tinygrad.llm.model import Transformer
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--devices", type=int, default=8)
args = parser.parse_args()
if args.devices not in (1, 2, 4, 8): raise ValueError("the K3 admission smoke test supports 1, 2, 4, or 8 devices")
devices = tuple(f"AMD:{i}" for i in range(args.devices))
model = Transformer(kimi_k3_smoke_config())
for name,value in nn.state.get_state_dict(model).items():
fill = 127 if name.endswith("weight_scale") else 0
dtype = value.dtype if value.dtype is dtypes.uint8 else dtypes.bfloat16
value.replace(Tensor.full(value.shape, fill, dtype=dtype, device="CPU"))
_shard_kimi_k3(model, devices)
temperature = Tensor([0.0], device=devices)
for label,tokens,start in (("prefill", [[1, 2]], 0), ("decode", [[3]], 2), ("decode replay", [[4]], 3)):
begin = time.perf_counter()
out = model(Tensor(tokens, dtype=dtypes.int32, device=devices), start, temperature).realize()
for device in devices: Device[device].synchronize()
print(f"{label}: shape={out.shape}, {time.perf_counter()-begin:.3f}s")
if __name__ == "__main__": main()
+70
View File
@@ -0,0 +1,70 @@
#!/usr/bin/env python3
"""Benchmark Kimi-Linear load, prefill, and decode on its TP4 checkpoint."""
import argparse, resource, time
from tinygrad import Device, TinyJit
from tinygrad.llm.kimi import load_kimi
def sync(devices:int) -> None:
for i in range(devices): Device[f"AMD:{i}"].synchronize()
def timed_next(gen, devices:int) -> tuple[int, float]:
begin = time.perf_counter()
token = next(gen)
sync(devices)
return token, time.perf_counter()-begin
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("model", help="converted Kimi-Linear-48B-A3B MXFP4-v2 directory")
parser.add_argument("--devices", type=int, default=4)
parser.add_argument("--max-context", type=int, default=128)
parser.add_argument("--prompt-tokens", type=int, default=32)
parser.add_argument("--decode-tokens", type=int, default=8)
parser.add_argument("--chunk-size", type=int, default=32)
parser.add_argument("--sweep-chunks", help="comma-separated prefill chunk sizes; uses the fastest for decode")
args = parser.parse_args()
if args.prompt_tokens < 1 or args.prompt_tokens + args.decode_tokens + 1 > args.max_context:
raise ValueError("prompt and decode tokens must fit within --max-context")
begin = time.perf_counter()
model = load_kimi(args.model, max_context=args.max_context, devices=args.devices)
sync(args.devices)
print(f"load: {time.perf_counter()-begin:.3f}s", flush=True)
prompt = [1] + [1000+i%1000 for i in range(args.prompt_tokens-1)]
chunks = [int(x) for x in args.sweep_chunks.split(",")] if args.sweep_chunks else [args.chunk_size]
if any(x < 1 or x > args.prompt_tokens for x in chunks): raise ValueError("prefill chunks must be between 1 and --prompt-tokens")
timings:list[tuple[float, int]] = []
prefill_jits:dict[int, TinyJit] = {}
for chunk in chunks:
# Recurrent prefill has a static token dimension. Give each swept shape its own capture;
# the rollout JIT remains shared and independently benchmarks chunk 1/decode.
if chunk != 1: model.prefill_jit = TinyJit(model.forward)
cold = model.generate(prompt.copy(), chunk_size=chunk)
first, cold_prefill = timed_next(cold, args.devices)
print(f"chunk {chunk}: cold prefill {cold_prefill:.3f}s, token={first}", flush=True)
warm = model.generate(prompt.copy(), chunk_size=chunk)
warm_first, prefill = timed_next(warm, args.devices)
if first != warm_first: raise RuntimeError(f"chunk {chunk} is not repeatable: cold={first}, warm={warm_first}")
timings.append((prefill, chunk))
if chunk != 1: prefill_jits[chunk] = model.prefill_jit
print(f"chunk {chunk}: prefill {prefill:.3f}s ({args.prompt_tokens/prefill:.3f} tok/s), token={first}", flush=True)
prefill, best_chunk = min(timings)
if best_chunk != 1: model.prefill_jit = prefill_jits[best_chunk]
warm = model.generate(prompt.copy(), chunk_size=best_chunk)
first, replay_prefill = timed_next(warm, args.devices)
_, cold_decode = timed_next(warm, args.devices)
_, capture_decode = timed_next(warm, args.devices)
print(f"selected chunk: {best_chunk}; prefill replay {replay_prefill:.3f}s "
f"({args.prompt_tokens/replay_prefill:.3f} tok/s), token={first}", flush=True)
print(f"cold decode: {cold_decode:.3f}s", flush=True)
print(f"capture decode: {capture_decode:.3f}s", flush=True)
begin = time.perf_counter()
output = [next(warm) for _ in range(args.decode_tokens)]
sync(args.devices)
decode = time.perf_counter()-begin
print(f"decode: {decode:.3f}s ({args.decode_tokens/decode:.3f} tok/s, {decode/args.decode_tokens*1e3:.3f} ms/tok), output={output}", flush=True)
print(f"peak RSS: {resource.getrusage(resource.RUSAGE_SELF).ru_maxrss/1024:.1f} MiB", flush=True)
if __name__ == "__main__": main()
+29
View File
@@ -0,0 +1,29 @@
import unittest
from tinygrad import Tensor, dtypes, nn
from tinygrad.llm.kimi_k3 import _shard_kimi_k3
from test.unit.test_llm_k3 import small_k3_config
from tinygrad.llm.model import Transformer
class TestKimiK3TP8(unittest.TestCase):
@staticmethod
def _model():
model = Transformer(small_k3_config())
for name,value in nn.state.get_state_dict(model).items():
fill = 127 if name.endswith("weight_scale") else 0
dtype = value.dtype if value.dtype is dtypes.uint8 else dtypes.bfloat16
value.replace(Tensor.full(value.shape, fill, dtype=dtype, device="NULL"))
_shard_kimi_k3(model, tuple(f"NULL:{i}" for i in range(8)))
return model
def test_prefill_decode_and_jit_replay(self):
devices = tuple(f"NULL:{i}" for i in range(8))
model = self._model()
temperature = Tensor([0.0], device=devices)
self.assertEqual(model(Tensor([[1, 2]], dtype=dtypes.int32, device=devices), 0, temperature).realize().shape, (1, 1))
model(Tensor([[1, 2]], dtype=dtypes.int32, device=devices), 0, temperature).realize()
self.assertEqual(model(Tensor([[3]], dtype=dtypes.int32, device=devices), 2, temperature).realize().shape, (1, 1))
model(Tensor([[4]], dtype=dtypes.int32, device=devices), 3, temperature).realize()
self.assertEqual(model.blk[0].recurrent_state.uop.axis, 1)
self.assertEqual(model.blk[1].cache_k.dtype, dtypes.bfloat16)
if __name__ == "__main__": unittest.main()
+26 -1
View File
@@ -5,7 +5,7 @@ from tinygrad import Tensor, dtypes, nn
from tinygrad.llm.kimi import _shard_kimi
from tinygrad.llm.model import (
GatedDeltaNetBlock, SSMConfig, TransformerBlock, TransformerConfig,
apply_rope as apply_rope_new, l2norm, precompute_freqs_cis, pairwise_topk,
apply_rope as apply_rope_new, iterative_topk, l2norm, precompute_freqs_cis, pairwise_topk,
)
def apply_rope(x:Tensor, start_pos:int):
@@ -197,6 +197,23 @@ class TestGatedDeltaNetBlock(unittest.TestCase):
alpha = np.exp(-self._softplus_np(np.arange(1, 5)).reshape(1, 2, 1, 2))
np.testing.assert_allclose(block.recurrent_state.numpy(), initial_state.numpy() * alpha, rtol=1e-5, atol=1e-5)
def test_kda_safe_gate_decay(self):
config = self._make_config(n_heads=2, kda_full_rank_gate=True, kda_gate_lower_bound=-5.0,
ssm=SSMConfig(conv_kernel=2, state_size=2, group_count=2, time_step_rank=2, inner_size=4, kda=True))
block, x = GatedDeltaNetBlock(config, config.ssm), Tensor([[[1., 2., 0., 0.]]])
block.ssm_f_a.weight = Tensor([[1., 0., 0., 0.], [0., 1., 0., 0.]])
block.ssm_f_b.weight = Tensor([[1., 0.], [0., 1.], [1., 1.], [2., 1.]])
block.ssm_dt["bias"] = Tensor.zeros(4)
block.ssm_a = Tensor([[-2.], [-3.]]) # stores -exp(A_log)
block._init_state(x)
initial_state = Tensor.arange(8, dtype=dtypes.float32).reshape(1, 2, 2, 2)
block.recurrent_state.assign(initial_state).realize()
block._attention(x, 0).realize()
gate_logits = np.arange(1, 5, dtype=np.float32).reshape(1, 2, 2)
exp_a = np.array([2., 3.], dtype=np.float32).reshape(1, 2, 1)
alpha = np.exp(-5.0 / (1.0 + np.exp(-(exp_a * gate_logits)))).reshape(1, 2, 1, 2)
np.testing.assert_allclose(block.recurrent_state.numpy(), initial_state.numpy() * alpha, rtol=2e-5, atol=2e-5)
def test_kda_chunked_prefill_matches_decode(self):
config = self._make_config(max_context=4, n_heads=2,
ssm=SSMConfig(conv_kernel=2, state_size=2, group_count=2, time_step_rank=2, inner_size=4, kda=True), kda_split_qkv=True)
@@ -267,5 +284,13 @@ class TestPairwiseTopk(unittest.TestCase):
self.assertEqual(set(sel.numpy()[b, t].tolist()), expected)
np.testing.assert_allclose(vals.numpy()[b, t], data[b, t][sel.numpy()[b, t]])
def test_iterative_matches_numpy(self):
rng = np.random.default_rng(42)
data = rng.standard_normal((2, 3, 896), dtype=np.float32)
vals, sel = iterative_topk(Tensor(data), 16)
expected = np.argsort(-data, axis=-1, stable=True)[..., :16]
np.testing.assert_equal(sel.numpy(), expected)
np.testing.assert_allclose(vals.numpy(), np.take_along_axis(data, expected, axis=-1))
if __name__ == '__main__':
unittest.main()
+30
View File
@@ -0,0 +1,30 @@
import unittest
from tinygrad.llm.cli import KimiK3Template
from tinygrad.llm.serve import StreamRouter
class TestKimiK3Template(unittest.TestCase):
def test_simple_text_chat(self):
template = KimiK3Template()
got = template.render([{"role":"system", "content":"Be concise."}, {"role":"user", "content":"Hello"}])
self.assertTrue(got.startswith('<|open|>message role="system" type="thinking-effort"<|sep|>'))
self.assertIn('<|open|>message role="user"<|sep|>Hello<|close|>message<|sep|><|end_of_msg|>', got)
self.assertTrue(got.endswith('<|open|>message role="assistant"<|sep|><|open|>think<|sep|>'))
def test_preserves_assistant_thinking(self):
got = KimiK3Template().render([{"role":"assistant", "reasoning_content":"why", "content":"answer"}], add_generation_prompt=False)
self.assertIn('<|open|>think<|sep|>why<|close|>think<|sep|>', got)
self.assertIn('<|open|>response<|sep|>answer<|close|>response<|sep|>', got)
def test_rejects_unimplemented_modalities(self):
with self.assertRaisesRegex(ValueError, "text-only"):
KimiK3Template().render([{"role":"user", "content":[{"type":"image", "url":"x"}]}])
with self.assertRaisesRegex(ValueError, "tool rendering"):
KimiK3Template().render([{"role":"user", "content":"x"}], tools=[{"type":"function"}])
def test_xtml_stream_router(self):
router, routed = StreamRouter(reasoning=True, xtml=True), []
for piece in ("rea", "son<|close|>thi", "nk<|sep|><|open|>response<|sep|>ans", "wer<|close|>response<|sep|>"):
routed.extend(router.route(piece))
self.assertEqual(routed, [("reasoning_content", "rea"), ("reasoning_content", "son"), ("content", "ans"), ("content", "wer")])
if __name__ == "__main__": unittest.main()
+87
View File
@@ -0,0 +1,87 @@
import unittest
from dataclasses import replace
import numpy as np
from tinygrad import Tensor, dtypes, nn
from tinygrad.llm.kimi_k3 import KIMI_K3_FULL_ATTN_LAYERS, KIMI_K3_SSM_LAYERS, KIMI_K3_TEXT_SIZE, KIMI_K3_TP8_BYTES_PER_GPU, \
_layer_sources, _shard_kimi_k3, _validate_config, kimi_k3_config, kimi_k3_smoke_config
from tinygrad.llm.model import FFNBlock, Transformer
def small_k3_config(max_context:int=4): return replace(kimi_k3_smoke_config(max_context), num_experts=8)
class TestKimiK3(unittest.TestCase):
def test_official_config(self):
c = kimi_k3_config(1_048_576)
self.assertEqual((c.num_blocks, c.dim, c.n_heads, c.num_experts, c.num_experts_per_tok), (93, 7168, 96, 896, 16))
self.assertEqual((sum(KIMI_K3_SSM_LAYERS), len(KIMI_K3_FULL_ATTN_LAYERS)), (69, 24))
self.assertEqual(KIMI_K3_FULL_ATTN_LAYERS, (*range(3, 93, 4), 92))
self.assertEqual((c.routed_expert_dim, c.hidden_dim, c.shared_expert_dim), (3584, 3072, 6144))
self.assertTrue(c.route_weights_uncorrected and c.kda_full_rank_gate and c.attn_output_gate)
self.assertEqual((c.activation_situ_beta, c.activation_situ_linear_beta, c.kda_gate_lower_bound), (4.0, 25.0, -5.0))
def test_config_rejects_wrong_checkpoint(self):
with self.assertRaisesRegex(ValueError, "not the supported official"):
_validate_config({"model_type":"kimi_linear", "hidden_size":2304})
def test_official_mapping_covers_model(self):
model = Transformer(kimi_k3_config(1))
state = nn.state.get_state_dict(model)
targets = {"token_embd.weight", "output_norm.weight", "output.weight", "output_attn_res_norm.weight", "output_attn_res_proj.weight"}
for i,is_kda in enumerate(KIMI_K3_SSM_LAYERS):
for target in _layer_sources(i, is_kda).values(): targets.update(target.split("|"))
if i:
for name in ("ffn_gate_exps.weight", "ffn_gate_exps.weight_scale", "ffn_up_exps.weight", "ffn_up_exps.weight_scale",
"ffn_down_exps.weight", "ffn_down_exps.weight_scale"): targets.add(f"blk.{i}.{name}")
self.assertEqual(targets, set(state))
self.assertEqual(state["blk.1.ffn_gate_exps.weight"].shape, (896, 3072, 1792))
self.assertEqual(state["blk.1.ffn_gate_exps.weight_scale"].shape, (896, 3072, 112))
_shard_kimi_k3(model, tuple(f"NULL:{i}" for i in range(8)))
total, per_gpu = 0, 0
for name,value in state.items():
dtype = dtypes.uint8 if name.endswith(("weight_scale", "_exps.weight")) else \
dtypes.float32 if name.endswith(("ssm_a", "ssm_dt.bias")) else dtypes.bfloat16
size = value.numel() * dtype.itemsize
total += size
per_gpu += size if value.uop.axis is None else size//8
self.assertEqual((total, per_gpu), (KIMI_K3_TEXT_SIZE, KIMI_K3_TP8_BYTES_PER_GPU))
def test_situ_matches_reference(self):
block = FFNBlock(small_k3_config())
gate, up = Tensor([[-8., -1., 0., 3.]]), Tensor([[-30., -2., 5., 40.]])
got = block._activation(gate, up).numpy()
g, u = gate.numpy().astype(np.float32), up.numpy().astype(np.float32)
expected = (4*np.tanh(g/4)/(1+np.exp(-g))) * (25*np.tanh(u/25))
np.testing.assert_allclose(got, expected, rtol=1e-5, atol=1e-5)
def test_attention_residual_matches_reference(self):
block = FFNBlock(small_k3_config())
block.attn_res_norm.weight.assign([1.0+i/16 for i in range(32)])
block.attn_res_proj.weight.assign([[(-1.0)**i/8 for i in range(32)]])
prefix, residual = Tensor.arange(64).reshape(2, 32).float()/16, Tensor.arange(128).reshape(2, 2, 32).float()/32
got = block._apply_attn_res(prefix, residual, block.attn_res_proj, block.attn_res_norm).numpy()
v = np.concatenate((residual.numpy(), prefix.numpy()[:, None]), axis=1).astype(np.float32)
k = v / np.sqrt(np.mean(v*v, axis=-1, keepdims=True) + 1e-5)
scores = np.sum(k * block.attn_res_norm.weight.numpy() * block.attn_res_proj.weight.numpy()[0], axis=-1)
probs = np.exp(scores-scores.max(axis=-1, keepdims=True))
probs /= probs.sum(axis=-1, keepdims=True)
expected = np.matmul(probs[:, None], v).squeeze(1)
np.testing.assert_allclose(got, expected, rtol=1e-5, atol=1e-5)
def test_tp8_schema(self):
model = Transformer(small_k3_config())
_shard_kimi_k3(model, tuple(f"NULL:{i}" for i in range(8)))
state = nn.state.get_state_dict(model)
for name,axis in (("token_embd.weight",0), ("blk.1.ffn_gate_exps.weight",1), ("blk.1.ffn_down_exps.weight_scale",2),
("blk.1.ffn_routed_down.weight",1), ("blk.0.ssm_g_full.weight",0), ("blk.1.attn_q_b.weight",0)):
self.assertEqual(state[name].uop.axis, axis, name)
self.assertIsNone(state["blk.1.attn_res_norm.weight"].uop.axis)
self.assertIsNone(state["blk.1.ffn_routed_norm.weight"].uop.axis)
def test_chunked_recurrent_generate(self):
model = Transformer(small_k3_config(max_context=8))
for name,value in nn.state.get_state_dict(model).items():
fill = 127 if name.endswith("weight_scale") else 0
value.replace(Tensor.full(value.shape, fill, dtype=value.dtype if value.dtype is dtypes.uint8 else dtypes.bfloat16, device="PYTHON"))
for _ in range(3): self.assertIsInstance(next(model.generate([1, 2, 3, 4], chunk_size=2)), int)
self.assertEqual(model._cached_tokens[:4], [1, 2, 3, 4])
if __name__ == "__main__": unittest.main()
+75 -11
View File
@@ -127,7 +127,48 @@ class FallbackTemplate:
out += self.end_turn()
return out + self.role("assistant") if add_generation_prompt else out
from tinygrad.llm.serve import LLMServer
class KimiK3Template:
"""Official K3 XTML envelope for text-only system/user/assistant conversations."""
OPEN, CLOSE, SEP, END = "<|open|>", "<|close|>", "<|sep|>", "<|end_of_msg|>"
def _open(self, tag:str, attrs:tuple[tuple[str, str], ...]=()) -> str:
escaped = ((k, str(v).replace("&", "&amp;").replace('"', "&quot;")) for k,v in attrs)
return self.OPEN + tag + "".join(f' {k}="{v}"' for k,v in escaped) + self.SEP
def _close(self, tag:str) -> str: return self.CLOSE + tag + self.SEP
def _message(self, role:str, content:str, name:str|None=None) -> str:
attrs = (("role", role),) + (() if name is None else (("name", name),))
return self._open("message", attrs) + content + self._close("message") + self.END
@staticmethod
def _content(message:dict) -> str:
content = message.get("content")
if content is None: return ""
if isinstance(content, str): return content
if isinstance(content, list):
if any(part.get("type") != "text" for part in content): raise ValueError("Kimi K3 native loader is text-only; image content is not implemented")
return "".join(part["text"] for part in content)
raise ValueError(f"unsupported Kimi K3 content type {type(content).__name__}")
def render(self, messages:list[dict], tools=None, add_generation_prompt:bool=True, preserve_thinking:bool=False, **kwargs) -> str:
if tools or any(m.get("role") == "tool" or m.get("tool_calls") for m in messages):
raise ValueError("Kimi K3 XTML tool rendering is not implemented in the native text loader")
effort = kwargs.get("thinking_effort", "max")
if effort not in ("low", "high", "max"): raise ValueError(f"invalid Kimi K3 thinking_effort {effort!r}")
body = "`thinking_effort` guides on how much to think in your thinking channel (not including the response channel), " \
"supported values include `low`, `medium`, `high`, and `max`.\n" \
f"Now the system is invoked with `thinking_effort={effort}`."
out = self._open("message", (("role", "system"), ("type", "thinking-effort"))) + body + self._close("message") + self.END
for message in messages:
role = message["role"]
if role in ("user", "system"):
out += self._message(role, self._content(message), message.get("name"))
elif role == "assistant":
reasoning = message.get("reasoning_content") or message.get("reasoning") or ""
content = self._open("think") + str(reasoning) + self._close("think")
content += self._open("response") + self._content(message) + self._close("response")
out += self._message(role, content, message.get("name"))
else: raise ValueError(f"unsupported Kimi K3 role {role!r}")
if add_generation_prompt: out += self._open("message", (("role", "assistant"),)) + self._open("think")
return out
from tinygrad.llm.serve import LLMServer, StreamRouter
def main():
parser = argparse.ArgumentParser()
@@ -137,13 +178,25 @@ def main():
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("--benchmark", nargs='?', type=int, const=20, metavar="COUNT", help="Benchmark tok/s (optional count, default 20)")
parser.add_argument("--devices", type=int, default=1, help="Tensor-parallel device count (Kimi MXFP4 requires 4)")
parser.add_argument("--devices", type=int, default=1, help="Tensor-parallel device count (Kimi-Linear requires 4, Kimi K3 requires 8)")
args = parser.parse_args()
# load the model
model_path = pathlib.Path(args.model)
kv:dict[str, typing.Any]
if model_path.is_dir() and (model_path / "tinygrad-kimi.json").exists():
is_k3 = False
if model_path.is_dir() and (model_path / "config.json").exists():
raw_config = json.loads((model_path / "config.json").read_text())
is_k3 = raw_config.get("model_type") == "kimi_k3"
if is_k3:
from tinygrad.llm.kimi_k3 import load_kimi_k3, load_kimi_tokenizer_data
model, kv = load_kimi_k3(model_path, args.max_context, args.devices), {}
normal, special, bos, eos = load_kimi_tokenizer_data(model_path)
tok = SimpleTokenizer(normal, special, "kimi-k2", bos_id=bos, eos_id=eos, eot_id=eos)
model_name = "Kimi-K3"
tok_cfg = json.loads((model_path / "tokenizer_config.json").read_text())
ct = tok_cfg.get("chat_template")
elif model_path.is_dir() and (model_path / "tinygrad-kimi.json").exists():
from tinygrad.llm.kimi import load_kimi, load_kimi_tokenizer_data
model, kv = load_kimi(model_path, args.max_context, args.devices), {}
normal, special, bos, eos = load_kimi_tokenizer_data(model_path)
@@ -161,7 +214,7 @@ def main():
f"max context {args.max_context} on {nn.state.get_parameters(model)[0].device}")
# use the model's chat template if jinja2 is available (enables model-specific formatting)
template: jinja2.Template|FallbackTemplate = FallbackTemplate(tok)
template: jinja2.Template|FallbackTemplate|KimiK3Template = KimiK3Template() if is_k3 else FallbackTemplate(tok)
if ct is not None:
try:
import jinja2
@@ -201,15 +254,26 @@ def main():
while 1:
try: messages.append({"role":"user", "content":input('>>> ')})
except EOFError: break
ids = tok.encode(template.render(messages=messages, add_generation_prompt=True))
reply, dec = "", tok.stream_decoder()
rendered = template.render(messages=messages, add_generation_prompt=True)
ids = tok.encode(rendered)
reply, reasoning_reply, dec = "", "", tok.stream_decoder()
xtml = rendered.rstrip().endswith("<|open|>think<|sep|>")
router = StreamRouter(reasoning=xtml or rendered.rstrip().endswith("<think>"), xtml=xtml)
for next_id in model.generate(ids):
if tok.is_end(next_id):
sys.stdout.write(dec() + "\n\n")
for field,text in router.route(dec(), final=True):
if field == "content": reply += text
elif field == "reasoning_content": reasoning_reply += text
sys.stdout.write(text)
sys.stdout.write("\n\n")
break
reply += (piece := dec(next_id))
sys.stdout.write(piece)
sys.stdout.flush()
messages.append({"role":"assistant", "content":reply})
for field,text in router.route(dec(next_id)):
if field == "content": reply += text
elif field == "reasoning_content": reasoning_reply += text
sys.stdout.write(text)
sys.stdout.flush()
assistant = {"role":"assistant", "content":reply}
if reasoning_reply: assistant["reasoning_content"] = reasoning_reply
messages.append(assistant)
if __name__ == "__main__": main()
+51
View File
@@ -0,0 +1,51 @@
import functools
from typing import cast
from tinygrad import Tensor, UOp, Device, Context, dtypes
from tinygrad.dtype import AddrSpace
from tinygrad.uop.ops import AxisType, KernelInfo
def amd_custom_kernels_supported(device:str|tuple[str, ...]|None) -> bool:
"""The hand-written wave32 kernel is intentionally limited to RDNA3/gfx11."""
if device is None: return False
device = device[0] if isinstance(device, tuple) else device
with Context(ALLOW_DEVICE_USAGE=1):
return (target:=getattr(Device[device], "target", None)) is not None and target[0] == 11
@functools.cache
def _gated_delta_prefill_kernel(core:UOp, next_state:UOp, q:UOp, k:UOp, v:UOp, beta:UOp, alpha:UOp, state:UOp, kq:UOp) -> UOp:
batch, heads, tokens, value_dim = cast(tuple[int, int, int, int], core.shape)
key_dim, alpha_dim = cast(int, q.shape[-1]), cast(int, alpha.shape[-1]) if len(alpha.shape) == 4 else 1
core, v = (x.reshape(batch*heads, tokens, value_dim) for x in (core, v))
q, k = (x.reshape(batch*heads, tokens, key_dim) for x in (q, k))
beta, kq = (x.reshape(batch*heads, tokens) for x in (beta, kq))
alpha = alpha.reshape(batch*heads, tokens, alpha_dim)
state, next_state = (x.reshape(batch*heads, value_dim, key_dim) for x in (state, next_state))
bh, row, cols = UOp.range(batch*heads, 0, AxisType.GLOBAL), UOp.range(value_dim, 2), tuple(range(key_dim))
current = UOp.placeholder((key_dim,), dtypes.float32, slot=0, addrspace=AddrSpace.REG)
current = current.after(UOp.group(*(current[col].store(state[bh, row, col].float()) for col in cols)))
token = UOp.range(tokens, 1, AxisType.REDUCE)
previous = tuple(current.after(token)[col].load() for col in cols)
keys, queries = (tuple(x[bh, token, col].load() for col in cols) for x in (k, q))
av = tuple(alpha[bh, token, col if alpha_dim > 1 else 0].load() for col in cols)
bv = beta[bh, token].load()
state_k = sum((x*a*y for x,a,y in zip(previous, av, keys)), UOp.const(0, dtypes.float32))
state_q = sum((x*a*y for x,a,y in zip(previous, av, queries)), UOp.const(0, dtypes.float32))
delta = (v[bh, token, row].load() - state_k) * bv
step = UOp.group(core[bh, token, row].store(state_q + delta*kq[bh, token]),
*(current[col].store(x*a + delta*y) for col,x,a,y in zip(cols, previous, av, keys))).end(token)
stores = (next_state[bh, row, col].store(current.after(step)[col].load().cast(next_state.dtype)) for col in cols)
return UOp.group(*stores).end(row, bh).sink(arg=KernelInfo(name="gated_delta_prefill", opts_to_apply=()))
def gated_delta_prefill(q:Tensor, k:Tensor, v:Tensor, beta:Tensor, alpha:Tensor, state:Tensor) -> tuple[Tensor, Tensor]:
batch, heads, tokens, key_dim = q.shape
value_dim = v.shape[-1]
assert q.shape == k.shape and v.shape[:3] == q.shape[:3] and beta.shape == (batch, heads, tokens)
assert alpha.shape in ((batch, heads, tokens), (batch, heads, tokens, key_dim))
assert state.shape == (batch, heads, value_dim, key_dim)
kernel = _gated_delta_prefill_kernel
if amd_custom_kernels_supported(q.device) and key_dim % 32 == 0 and value_dim % 4 == 0:
from tinygrad.llm.kernels.amd import _gated_delta_prefill_kernel as kernel
core, next_state, kq = Tensor.empty_like(v), Tensor.empty_like(state), (q*k).sum(-1).contiguous()
result = Tensor.custom_kernel(core, next_state, q.contiguous(), k.contiguous(), v.contiguous(), beta.contiguous(), alpha.contiguous(), state, kq,
fxn=kernel)
return result[0], result[1]
+50
View File
@@ -0,0 +1,50 @@
from __future__ import annotations
import functools
from typing import cast
from tinygrad import UOp
from tinygrad.uop.ops import AxisType, KernelInfo, Ops
from tinygrad.dtype import AddrSpace, dtypes
def warp_reduce(val:UOp, full_wave:bool=False) -> UOp:
for offset in ((16, 8, 4, 2, 1) if full_wave else (8, 4, 2, 1)):
if val.op is Ops.INDEX and val.addrspace == AddrSpace.REG: val = val.load()
other = UOp(Ops.CUSTOM, dtypes.float, (val,), arg=
f"__builtin_bit_cast(float, __builtin_amdgcn_ds_swizzle(__builtin_bit_cast(int, {{0}}), {0x1f | offset<<10}))")
val = val + other
return val
@functools.cache
def _gated_delta_prefill_kernel(core:UOp, next_state:UOp, q:UOp, k:UOp, v:UOp, beta:UOp, alpha:UOp, state:UOp, kq:UOp) -> UOp:
batch, heads, tokens, value_dim, row_tile = *core.shape, 4
key_dim, alpha_dim = q.shape[-1], alpha.shape[-1] if len(alpha.shape) == 4 else 1
assert all(isinstance(x, int) for x in (batch, heads, tokens, value_dim, key_dim)) and key_dim % 32 == 0 and value_dim % row_tile == 0
batch, heads, tokens, value_dim, key_dim = cast(tuple[int, int, int, int, int], (batch, heads, tokens, value_dim, key_dim))
core, v = (x.reshape(batch*heads, tokens, value_dim) for x in (core, v))
q, k = (x.reshape(batch*heads, tokens, key_dim) for x in (q, k))
beta, kq = (x.reshape(batch*heads, tokens) for x in (beta, kq))
alpha = alpha.reshape(batch*heads, tokens, alpha_dim)
state, next_state = (x.reshape(batch*heads, value_dim, key_dim) for x in (state, next_state))
bh_row, lane = UOp.range(batch*heads*value_dim//row_tile, 0), UOp.range(32, 1, axis_type=AxisType.LOCAL)
bh, row_base = bh_row // (value_dim//row_tile), (bh_row % (value_dim//row_tile))*row_tile
rows = tuple(row_base+i for i in range(row_tile))
cols = tuple(lane + i*32 for i in range(key_dim//32))
current = UOp.placeholder((row_tile*key_dim//32,), dtypes.float32, slot=0, addrspace=AddrSpace.REG)
current = current.after(current.store(UOp.stack(*(state[bh, row, col].float() for row in rows for col in cols))))
token = UOp.range(tokens, 2, AxisType.REDUCE)
keys = tuple(k[bh, token, col].load() for col in cols)
queries = tuple(q[bh, token, col].load() for col in cols)
updates:list[UOp] = []
stores:list[UOp] = []
for row_idx,row in enumerate(rows):
previous = tuple(current.after(token)[row_idx*key_dim//32+i].load() for i in range(key_dim//32))
av = tuple(alpha[bh, token, col if alpha_dim > 1 else 0].load() for col in cols)
bv = beta[bh, token].load()
state_k = warp_reduce(sum((x*a*y for x,a,y in zip(previous, av, keys)), UOp.const(0, dtypes.float32)), full_wave=True)
state_q = warp_reduce(sum((x*a*y for x,a,y in zip(previous, av, queries)), UOp.const(0, dtypes.float32)), full_wave=True)
delta = (v[bh, token, row].load() - state_k) * bv
updates += [x*a + delta*y for x,a,y in zip(previous, av, keys)]
stores.append(core[bh, token, row.valid(lane.eq(0))].store(state_q + delta*kq[bh, token]))
step = UOp.group(*stores, current.store(UOp.stack(*updates))).end(token)
state_stores = (next_state[bh, row, col].store(current.after(step)[row_idx*key_dim//32+i].load().cast(next_state.dtype))
for row_idx,row in enumerate(rows) for i,col in enumerate(cols))
return UOp.group(*state_stores).end(lane, bh_row).sink(arg=KernelInfo(name="gated_delta_prefill", opts_to_apply=()))
+1
View File
@@ -15,6 +15,7 @@ def kimi_config(max_context:int, expert_mxfp4:bool=True) -> TransformerConfig:
num_experts=256, num_experts_per_tok=8, norm_topk_prob=True, shared_expert_dim=1024, leading_dense_blocks=1,
dense_hidden_dim=9216, routed_scaling_factor=2.446, expert_bias=True, max_context=max_context, expert_mxfp4=expert_mxfp4,
shared_expert_gate=False, bf16_activations=True, kda_split_qkv=True,
recurrent_prefill_chunked=True, recurrent_prefill_chunk_size=8,
ssm=SSMConfig(conv_kernel=4, state_size=128, group_count=32, time_step_rank=32, inner_size=4096, kda=True),
ssm_layers=KIMI_SSM_LAYERS)
+200
View File
@@ -0,0 +1,200 @@
from __future__ import annotations
import gc, json, pathlib
from dataclasses import replace
from collections import defaultdict
from typing import Callable
from tinygrad import Tensor, Device, nn
from tinygrad.nn.state import safe_load
from tinygrad.llm.kimi import load_kimi_tokenizer_data
from tinygrad.llm.model import SSMConfig, Transformer, TransformerConfig
KIMI_K3_TOTAL_SIZE = 1_560_860_324_864
KIMI_K3_TEXT_SIZE = 1_559_945_066_624
KIMI_K3_TP8_BYTES_PER_GPU = 196_781_639_152
KIMI_K3_SHARDS = 96
KIMI_K3_EXPERTS = 896
KIMI_K3_LAYERS = 93
KIMI_K3_FULL_ATTN_LAYERS = (*range(3, KIMI_K3_LAYERS, 4), 92)
KIMI_K3_SSM_LAYERS = tuple(i not in KIMI_K3_FULL_ATTN_LAYERS for i in range(KIMI_K3_LAYERS))
def kimi_k3_config(max_context:int) -> TransformerConfig:
"""Official Kimi K3 text-tower configuration (zero-based full-attention layers)."""
return TransformerConfig(num_blocks=93, dim=7168, hidden_dim=3072, n_heads=96, n_kv_heads=96, norm_eps=1e-5,
vocab_size=163840, head_dim=192, rope_theta=10000.0, rope_dim=64, v_head_dim=128, max_context=max_context,
q_lora_rank=1536, kv_lora_rank=512, num_experts=896, num_experts_per_tok=16, norm_topk_prob=True,
shared_expert_dim=6144, leading_dense_blocks=1, dense_hidden_dim=33792, routed_scaling_factor=1.0,
expert_bias=True, expert_mxfp4=True, bf16_activations=True, kda_split_qkv=True,
ssm=SSMConfig(conv_kernel=4, state_size=128, group_count=96, time_step_rank=96, inner_size=12288, kda=True),
ssm_layers=KIMI_K3_SSM_LAYERS, shared_expert_gate=False, attn_output_gate=True,
activation_situ_beta=4.0, activation_situ_linear_beta=25.0, routed_expert_dim=3584, latent_moe_norm=True,
route_weights_uncorrected=True, attn_res_block_size=12, kda_full_rank_gate=True, kda_gate_lower_bound=-5.0,
recurrent_prefill_chunked=True, recurrent_prefill_chunk_size=8)
def kimi_k3_smoke_config(max_context:int=4) -> TransformerConfig:
"""Reduced K3 with every architectural feature retained for cheap compile/hardware admission tests."""
return replace(kimi_k3_config(max_context), num_blocks=2, dim=32, hidden_dim=256, n_heads=8, n_kv_heads=8,
vocab_size=64, head_dim=8, rope_dim=4, v_head_dim=4, q_lora_rank=16, kv_lora_rank=8, num_experts=512,
num_experts_per_tok=2, shared_expert_dim=32, dense_hidden_dim=64, routed_expert_dim=32,
ssm=SSMConfig(4, 4, 8, 8, 32, True), ssm_layers=(True, False), attn_res_block_size=1)
def _shard_kimi_k3(model:Transformer, devices:tuple[str, ...]) -> None:
"""Tensor parallel layout for K3. The official dimensions are divisible by TP8."""
if len(devices) not in (1, 2, 4, 8): raise ValueError(f"Kimi K3 tensor parallelism requires 1, 2, 4, or 8 devices, got {len(devices)}")
for name, value in nn.state.get_state_dict(model).items():
axis = None
if name in ("token_embd.weight", "output.weight"): axis = 0
elif ".ffn_gate_exps.weight" in name or ".ffn_up_exps.weight" in name: axis = 1
elif ".ffn_gate_exps.weight_scale" in name or ".ffn_up_exps.weight_scale" in name: axis = 1
elif ".ffn_down_exps.weight" in name or ".ffn_down_exps.weight_scale" in name: axis = 2
elif name.endswith((".ffn_gate.weight", ".ffn_up.weight", ".ffn_gate_shexp.weight", ".ffn_up_shexp.weight")): axis = 0
elif name.endswith((".ffn_down.weight", ".ffn_down_shexp.weight", ".ffn_routed_down.weight", ".ffn_routed_up.weight",
".attn_output.weight", ".ssm_out.weight")): axis = 1
elif name.endswith((".attn_q_b.weight", ".attn_k_b.weight", ".attn_v_b.weight", ".attn_gate.weight",
".attn_q.weight", ".attn_k.weight", ".attn_v.weight", ".ssm_f_b.weight", ".ssm_g_full.weight", ".ssm_beta.weight")): axis = 0
elif name.endswith((".ssm_q_conv1d.weight", ".ssm_k_conv1d.weight", ".ssm_v_conv1d.weight", ".ssm_a", ".ssm_dt.bias")): axis = 0
value.shard_(devices, axis=axis)
def _validate_config(config:dict) -> None:
text = config.get("text_config", config)
expected = {"model_type":"kimi_linear", "hidden_size":7168, "num_hidden_layers":93, "num_attention_heads":96,
"vocab_size":163840, "intermediate_size":33792, "num_experts":896, "num_experts_per_token":16,
"moe_intermediate_size":3072, "num_shared_experts":2, "q_lora_rank":1536, "kv_lora_rank":512,
"qk_nope_head_dim":128, "qk_rope_head_dim":64, "v_head_dim":128, "routed_expert_hidden_size":3584,
"attn_res_block_size":12, "hidden_act":"situ", "mla_use_nope":True, "mla_use_output_gate":True,
"activation_situ_beta":4.0, "activation_situ_linear_beta":25.0, "latent_moe_use_norm":True,
"moe_renormalize":True, "first_k_dense_replace":1, "num_expert_group":1, "topk_group":1}
bad = {k:(text.get(k), v) for k,v in expected.items() if text.get(k) != v}
linear = text.get("linear_attn_config", {})
linear_expected = {"head_dim":128, "num_heads":96, "short_conv_kernel_size":4, "use_full_rank_gate":True,
"gate_lower_bound":-5.0, "full_attn_layers":[i+1 for i in KIMI_K3_FULL_ATTN_LAYERS],
"kda_layers":[i+1 for i,x in enumerate(KIMI_K3_SSM_LAYERS) if x]}
bad.update({f"linear_attn_config.{k}":(linear.get(k), v) for k,v in linear_expected.items() if linear.get(k) != v})
quant = text.get("quantization_config", {})
if quant.get("format") != "mxfp4-pack-quantized": bad["quantization_config.format"] = (quant.get("format"), "mxfp4-pack-quantized")
if bad: raise ValueError(f"not the supported official Kimi K3 checkpoint: {bad}")
def audit_kimi_k3_checkpoint(model_dir:str|pathlib.Path, require_shards:bool=True) -> dict[str, int]:
"""Validate checkpoint metadata only. This never opens weight data and is safe on small hosts."""
root = pathlib.Path(model_dir)
_validate_config(json.loads((root / "config.json").read_text()))
index = json.loads((root / "model.safetensors.index.json").read_text())
weight_map, total = index.get("weight_map", {}), index.get("metadata", {}).get("total_size")
language = [k for k in weight_map if k.startswith("language_model.")]
experts = [k for k in language if ".block_sparse_moe.experts." in k]
missing_files = {fn for fn in weight_map.values() if not (root / fn).is_file()}
if total != KIMI_K3_TOTAL_SIZE: raise ValueError(f"unexpected checkpoint size {total}, expected {KIMI_K3_TOTAL_SIZE}")
if len(set(weight_map.values())) != KIMI_K3_SHARDS: raise ValueError("official Kimi K3 must contain 96 safetensor shards")
if len(experts) != 92 * KIMI_K3_EXPERTS * 3 * 2: raise ValueError(f"unexpected routed-expert tensor count {len(experts)}")
if require_shards and missing_files: raise FileNotFoundError(f"missing {len(missing_files)} checkpoint shards, first: {sorted(missing_files)[0]}")
return {"tensors":len(weight_map), "language_tensors":len(language), "expert_tensors":len(experts),
"shards":len(set(weight_map.values())), "missing_shards":len(missing_files), "total_size":total}
def _layer_sources(i:int, is_kda:bool) -> dict[str, str]:
src, dst = f"language_model.model.layers.{i}.", f"blk.{i}."
out = {
src+"input_layernorm.weight":dst+"attn_norm.weight", src+"post_attention_layernorm.weight":dst+"ffn_norm.weight",
src+"self_attention_res_norm.weight":dst+"attn_res_norm.weight", src+"self_attention_res_proj.weight":dst+"attn_res_proj.weight",
src+"mlp_res_norm.weight":dst+"mlp_res_norm.weight", src+"mlp_res_proj.weight":dst+"mlp_res_proj.weight",
}
if is_kda:
for a,b in (("q_proj","attn_q"),("k_proj","attn_k"),("v_proj","attn_v"),("g_proj","ssm_g_full"),
("f_a_proj","ssm_f_a"),("f_b_proj","ssm_f_b"),("b_proj","ssm_beta"),("o_proj","ssm_out")):
out[src+f"self_attn.{a}.weight"] = dst+b+".weight"
for a,b in (("q_conv1d","ssm_q_conv1d"),("k_conv1d","ssm_k_conv1d"),("v_conv1d","ssm_v_conv1d")):
out[src+f"self_attn.{a}.weight"] = dst+b+".weight"
out[src+"self_attn.o_norm.weight"], out[src+"self_attn.dt_bias"], out[src+"self_attn.A_log"] = \
dst+"ssm_norm.weight", dst+"ssm_dt.bias", dst+"ssm_a"
else:
for a,b in (("q_a_proj","attn_q_a"),("q_a_layernorm","attn_q_a_norm"),("q_b_proj","attn_q_b"),
("kv_a_proj_with_mqa","attn_kv_a_mqa"),("kv_a_layernorm","attn_kv_a_norm"),
("g_proj","attn_gate"),("o_proj","attn_output")):
out[src+f"self_attn.{a}.weight"] = dst+b+".weight"
# kv_b_proj is split into head-wise K and V tensors while loading.
out[src+"self_attn.kv_b_proj.weight"] = dst+"attn_k_b.weight|"+dst+"attn_v_b.weight"
if i == 0:
for a,b in (("gate_proj","ffn_gate"),("up_proj","ffn_up"),("down_proj","ffn_down")): out[src+f"mlp.{a}.weight"] = dst+b+".weight"
else:
base = src+"block_sparse_moe."
out[base+"gate.weight"], out[base+"gate.e_score_correction_bias"] = dst+"ffn_gate_inp.weight", dst+"exp_probs_b.bias"
for a,b in (("gate_proj","ffn_gate_shexp"),("up_proj","ffn_up_shexp"),("down_proj","ffn_down_shexp"),
("routed_expert_down_proj","ffn_routed_down"),("routed_expert_up_proj","ffn_routed_up"),
("routed_expert_norm","ffn_routed_norm")):
out[base+(f"shared_experts.{a}.weight" if a.endswith("_proj") and not a.startswith("routed_") else a+".weight")] = dst+b+".weight"
return out
def _replace(dst:Tensor, src:Tensor) -> None:
if dst.shape != src.shape: raise ValueError(f"shape mismatch: expected {dst.shape}, got {src.shape}")
dst.replace(src if isinstance(src.device, tuple) else src.shard(dst.device, dst.uop.axis) if isinstance(dst.device, tuple) else src.to(dst.device))
dst.realize()
def _load_nonexperts(root:pathlib.Path, weight_map:dict[str, str], model:Transformer, progress:Callable[[str], None]) -> set[str]:
model_state, mappings = nn.state.get_state_dict(model), {
"language_model.model.embed_tokens.weight":"token_embd.weight", "language_model.model.norm.weight":"output_norm.weight",
"language_model.lm_head.weight":"output.weight", "language_model.model.output_attn_res_norm.weight":"output_attn_res_norm.weight",
"language_model.model.output_attn_res_proj.weight":"output_attn_res_proj.weight"}
for i,is_kda in enumerate(KIMI_K3_SSM_LAYERS): mappings.update(_layer_sources(i, is_kda))
by_file:dict[str, list[str]] = defaultdict(list)
for source in mappings:
if source not in weight_map: raise ValueError(f"missing Kimi K3 tensor {source}")
by_file[weight_map[source]].append(source)
consumed:set[str] = set()
for filename, sources in sorted(by_file.items()):
progress(f"loading non-expert tensors from {filename}")
shard = safe_load(root / filename)
for source in sources:
value, targets = shard[source], mappings[source].split("|")
if source.endswith("A_log"): value = -value.float().exp().reshape(96, 1)
if source.endswith("conv1d.weight"): value = value.squeeze(1)
if source.endswith("kv_b_proj.weight"):
value = value.reshape(96, 256, 512)
values:tuple[Tensor, ...] = (value[:, :128].transpose(1, 2), value[:, 128:])
else: values = (value,)
for target,tensor in zip(targets, values): _replace(model_state[target], tensor)
consumed.add(source)
del shard
gc.collect()
return consumed
def _load_experts(root:pathlib.Path, weight_map:dict[str, str], model:Transformer, progress:Callable[[str], None]) -> set[str]:
model_state, consumed = nn.state.get_state_dict(model), set()
for i in range(1, KIMI_K3_LAYERS):
for wid,dst_name in (("w1","ffn_gate_exps"),("w3","ffn_up_exps"),("w2","ffn_down_exps")):
base = f"language_model.model.layers.{i}.block_sparse_moe.experts"
packed_keys = [f"{base}.{e}.{wid}.weight_packed" for e in range(KIMI_K3_EXPERTS)]
scale_keys = [f"{base}.{e}.{wid}.weight_scale" for e in range(KIMI_K3_EXPERTS)]
files = sorted({weight_map[k] for k in packed_keys+scale_keys})
progress(f"loading layer {i}/92 {wid} routed experts from {', '.join(files)}")
shards = {fn:safe_load(root / fn) for fn in files}
_replace(model_state[f"blk.{i}.{dst_name}.weight"], Tensor.stack(*(shards[weight_map[k]][k] for k in packed_keys)))
_replace(model_state[f"blk.{i}.{dst_name}.weight_scale"], Tensor.stack(*(shards[weight_map[k]][k] for k in scale_keys)))
consumed.update(packed_keys+scale_keys)
del shards
gc.collect()
return consumed
def load_kimi_k3(model_dir:str|pathlib.Path, max_context:int=4096, devices:int=8,
progress:Callable[[str], None]=print) -> Transformer:
"""Load the official native K3 checkpoint without ever materializing it in host RAM.
Safetensor shards remain disk-backed, each destination is TP-sharded before transfer, and source
mappings are discarded after every file/projection. Vision tensors are intentionally ignored.
"""
root = pathlib.Path(model_dir)
_validate_config(json.loads((root / "config.json").read_text()))
index = json.loads((root / "model.safetensors.index.json").read_text())
weight_map = index["weight_map"]
if devices != 8: raise ValueError("official Kimi K3 currently requires --devices 8")
if index.get("metadata", {}).get("total_size") != KIMI_K3_TOTAL_SIZE or len(set(weight_map.values())) != KIMI_K3_SHARDS:
raise ValueError("checkpoint index does not match the official 96-shard Kimi K3 release")
missing_files = {fn for fn in weight_map.values() if not (root / fn).is_file()}
if missing_files: raise FileNotFoundError(f"missing {len(missing_files)} checkpoint shards, first: {sorted(missing_files)[0]}")
model = Transformer(kimi_k3_config(max_context))
_shard_kimi_k3(model, tuple(f"{Device.DEFAULT}:{i}" for i in range(devices)))
consumed = _load_nonexperts(root, weight_map, model, progress)
consumed.update(_load_experts(root, weight_map, model, progress))
unused_language = {k for k in weight_map if k.startswith("language_model.")} - consumed
if unused_language: raise ValueError(f"unmapped language tensors: {sorted(unused_language)[:20]}")
return model
__all__ = ["KIMI_K3_FULL_ATTN_LAYERS", "KIMI_K3_SSM_LAYERS", "KIMI_K3_TEXT_SIZE", "KIMI_K3_TP8_BYTES_PER_GPU",
"audit_kimi_k3_checkpoint", "kimi_k3_config", "kimi_k3_smoke_config", "load_kimi_k3", "load_kimi_tokenizer_data"]
+157 -63
View File
@@ -5,6 +5,7 @@ from tinygrad import Tensor, nn, UOp, TinyJit, getenv, function, dtypes
from tinygrad.nn import Linear
from tinygrad.llm.gguf import gguf_load
from tinygrad.llm.quant import dequantize_mxfp4, quantize_dequantize_mxfp8
from tinygrad.llm.kernels import gated_delta_prefill
from tinygrad.uop.ops import resolve
@functools.cache
@@ -54,6 +55,16 @@ 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 iterative_topk(x:Tensor, k:int) -> tuple[Tensor, Tensor]:
"""O(k*N) top-k for very wide MoE routers, with stable first-index tie breaking."""
work, values, indices = x, [], []
for _ in range(k):
sel = work.argmax(-1, keepdim=True)
values.append(x.gather(-1, sel))
indices.append(sel)
work = work.scatter(-1, sel, x.dtype.min)
return values[0].cat(*values[1:], dim=-1), indices[0].cat(*indices[1:], dim=-1)
@dataclass(frozen=True)
class SSMConfig:
conv_kernel: int
@@ -96,6 +107,17 @@ class TransformerConfig:
expert_mxfp4: bool = False
bf16_activations: bool = False
kda_split_qkv: bool = False
# Kimi K3 extensions. Defaults preserve all existing model behavior.
activation_situ_beta: float = 0.0
activation_situ_linear_beta: float = 0.0
routed_expert_dim: int = 0
latent_moe_norm: bool = False
route_weights_uncorrected: bool = False
attn_res_block_size: int = 0
kda_full_rank_gate: bool = False
kda_gate_lower_bound: float = 0.0
recurrent_prefill_chunked: bool = False
recurrent_prefill_chunk_size: int = 0
class FFNBlock:
def __init__(self, config:TransformerConfig):
@@ -110,9 +132,14 @@ class FFNBlock:
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)}
expert_cls = MXFP4ExpertWeights if config.expert_mxfp4 else ExpertWeights
self.ffn_gate_exps = expert_cls(config.num_experts, config.dim, config.hidden_dim)
self.ffn_up_exps = expert_cls(config.num_experts, config.dim, config.hidden_dim)
self.ffn_down_exps = expert_cls(config.num_experts, config.hidden_dim, config.dim)
expert_dim = config.routed_expert_dim or config.dim
self.ffn_gate_exps = expert_cls(config.num_experts, expert_dim, config.hidden_dim)
self.ffn_up_exps = expert_cls(config.num_experts, expert_dim, config.hidden_dim)
self.ffn_down_exps = expert_cls(config.num_experts, config.hidden_dim, expert_dim)
if config.routed_expert_dim:
self.ffn_routed_down = Linear(config.dim, expert_dim, bias=False)
self.ffn_routed_up = Linear(expert_dim, config.dim, bias=False)
if config.latent_moe_norm: self.ffn_routed_norm = nn.RMSNorm(expert_dim, config.norm_eps)
if config.shared_expert_dim > 0:
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)
@@ -123,31 +150,47 @@ class FFNBlock:
self.ffn_up = Linear(config.dim, config.hidden_dim, bias=False)
self.ffn_down = Linear(config.hidden_dim, config.dim, bias=False)
if config.attn_res_block_size:
self.attn_res_norm, self.mlp_res_norm = nn.RMSNorm(config.dim, config.norm_eps), nn.RMSNorm(config.dim, config.norm_eps)
self.attn_res_proj, self.mlp_res_proj = Linear(config.dim, 1, bias=False), Linear(config.dim, 1, bias=False)
def _activation(self, gate:Tensor, up:Tensor) -> Tensor:
if not self.config.activation_situ_beta: return gate.silu() * up
gate32, up32, beta = gate.float(), up.float(), self.config.activation_situ_beta
gate32 = beta * (gate32 / beta).tanh() * gate32.sigmoid()
if (linear_beta := self.config.activation_situ_linear_beta): up32 = linear_beta * (up32 / linear_beta).tanh()
return (gate32 * up32).cast(gate.dtype)
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
# Kimi computes router logits in FP32 even though the residual stream and weights are BF16.
logits = x.float().linear(self.ffn_gate_inp.weight.float().transpose()) if self.config.bf16_activations else self.ffn_gate_inp(x)
if hasattr(self, 'exp_probs_b'):
# Kimi's reference mutates the sigmoid-score view when adding the correction bias,
# so the corrected values determine both selection and the normalized route weights.
scores = logits.sigmoid() + self.exp_probs_b["bias"]
_, sel = pairwise_topk(scores, self.config.num_experts_per_tok)
probs = scores.gather(-1, sel)
scores = logits.sigmoid()
topk = iterative_topk if self.config.num_experts >= 512 else pairwise_topk
_, sel = topk(scores + self.exp_probs_b["bias"], self.config.num_experts_per_tok)
# Kimi-Linear-48B's older reference weights corrected scores. K3 selects with the correction
# but gathers the uncorrected sigmoid scores, so keep this an explicit compatibility switch.
probs = (scores if self.config.route_weights_uncorrected else scores + self.exp_probs_b["bias"]).gather(-1, sel)
if self.config.norm_topk_prob: probs = probs / (probs.sum(axis=-1, keepdim=True) + 1e-20)
else:
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 hasattr(self, 'ffn_routed_down'): h = self.ffn_routed_down(x).unsqueeze(2)
x_down = self.ffn_down_exps(sel, self._activation(self.ffn_gate_exps(sel, h), self.ffn_up_exps(sel, h)).contiguous())
out = (x_down * probs.unsqueeze(-1)).sum(axis=2).cast(x_down.dtype) # (B, T, D)
if hasattr(self, 'ffn_routed_up'):
if hasattr(self, 'ffn_routed_norm'): out = self.ffn_routed_norm(out)
out = self.ffn_routed_up(out)
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._activation(self.ffn_gate_shexp(x), self.ffn_up_shexp(x)).contiguous())
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))
return self.ffn_down(self._activation(self.ffn_gate(x), self.ffn_up(x)).contiguous())
# 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
@@ -170,6 +213,32 @@ class FFNBlock:
return (h + self._feed_forward(self.ffn_norm(h))).contiguous()
return _run(x, start_pos)
@staticmethod
def _apply_attn_res(prefix_sum:Tensor, block_residual:Tensor, proj:Linear, norm:nn.RMSNorm) -> Tensor:
# Both inputs are flattened over B*T. Scoring is intentionally FP32, matching K3 eager inference.
v = block_residual.cat(prefix_sum.unsqueeze(1), dim=1)
vf = v.float()
k = vf * (vf.square().mean(axis=-1, keepdim=True) + norm.eps).rsqrt()
assert norm.weight is not None
scores = (k * (norm.weight.float() * proj.weight.squeeze(0).float())).sum(axis=-1)
return (scores.softmax(-1).unsqueeze(1) @ vf).squeeze(1).cast(v.dtype)
def attn_residual(self, x:Tensor, start_pos:int|UOp, block_residual:Tensor, layer_idx:int) -> tuple[Tensor, Tensor]:
self._init_state(x)
shape, prefix_sum = x.shape, x
prefix:Tensor|None = prefix_sum
if block_residual.shape[1]: x = self._apply_attn_res(x.reshape(-1, shape[-1]), block_residual,
self.attn_res_proj, self.attn_res_norm).reshape(shape)
if layer_idx % self.config.attn_res_block_size == 0:
block_residual = block_residual.cat(prefix_sum.reshape(-1, shape[-1]).unsqueeze(1), dim=1)
prefix = None
attn = self._attention(self.attn_norm(x), start_pos)
prefix = attn if prefix is None else prefix + attn
x = self._apply_attn_res(prefix.reshape(-1, shape[-1]), block_residual,
self.mlp_res_proj, self.mlp_res_norm).reshape(shape)
mlp = self._feed_forward(self.ffn_norm(x))
return (prefix + mlp).contiguous(), block_residual
class TransformerBlock(FFNBlock):
def __init__(self, config:TransformerConfig):
super().__init__(config)
@@ -211,7 +280,9 @@ 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
mask = Tensor.full((1, 1, T, start_pos+T), float("-inf"), dtype=x.dtype, buffer=False).triu(start_pos+1) \
# Build the static T×T causal corner on-device, then prepend the unmasked cached prefix.
# A broadcast const with symbolic width otherwise defaults to CPU in multi-device graphs.
mask = Tensor.full((1, 1, T, T), float("-inf"), dtype=x.dtype, device=x.device).triu(1).pad(((0, 0),)*3+((start_pos, 0),)) \
if resolve(T != 1) else None
attn = q.scaled_dot_product_attention(k, v, attn_mask=mask, enable_gqa=True) # (B,H,T,Hd)
attn = attn.transpose(1, 2).reshape(B, T, -1) # back to (B,T,D)
@@ -239,6 +310,7 @@ class MLATransformerBlock(FFNBlock):
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 = Linear(config.n_heads * config.v_head_dim, config.dim, bias=False)
if config.attn_output_gate: self.attn_gate = Linear(config.dim, config.n_heads * config.v_head_dim, bias=False)
def _attention(self, x:Tensor, start_pos:int|UOp) -> Tensor:
B, T, _ = x.shape
@@ -258,13 +330,14 @@ class MLATransformerBlock(FFNBlock):
k = Tensor(self.cache_k.uop.after(self.cache_k[:, :, start_pos:start_pos+T, :].uop.store(k_store.uop)))[:, :, 0:start_pos+T, :]
v = k[..., :self.config.kv_lora_rank]
mask = Tensor.full((1, 1, T, start_pos+T), float("-inf"), dtype=x.dtype, buffer=False).triu(start_pos+1) \
mask = Tensor.full((1, 1, T, T), float("-inf"), dtype=x.dtype, device=x.device).triu(1).pad(((0, 0),)*3+((start_pos, 0),)) \
if resolve(T != 1) else None
attn = q @ k.transpose(-1, -2) * (1.0 / self.config.head_dim ** 0.5)
if mask is not None: attn = attn + mask
# Match eager Kimi MLA: normalize attention scores in FP32, then return to the query dtype.
attn = attn.softmax(-1, dtype=dtypes.float32).cast(q.dtype)
attn = ((attn @ v) @ self.attn_v_b["weight"].transpose(-1, -2)).transpose(1, 2).reshape(B, T, -1)
if hasattr(self, "attn_gate"): attn = attn * self.attn_gate(x).sigmoid()
return self.attn_output(attn)
def _init_state(self, x:Tensor):
@@ -290,7 +363,8 @@ class GatedDeltaNetBlock(FFNBlock):
self.attn_qkv = Linear(config.dim, self.conv_channels, bias=False)
self.ssm_conv1d = {"weight": Tensor.zeros(self.conv_channels, self.ssm_conv_kernel)}
if ssm.kda:
self.ssm_g_a, self.ssm_g_b = Linear(config.dim, self.head_v_dim, bias=False), Linear(self.head_v_dim, ssm.inner_size, bias=False)
if config.kda_full_rank_gate: self.ssm_g_full = Linear(config.dim, ssm.inner_size, bias=False)
else: self.ssm_g_a, self.ssm_g_b = Linear(config.dim, self.head_v_dim, bias=False), Linear(self.head_v_dim, ssm.inner_size, bias=False)
self.ssm_f_a, self.ssm_f_b = Linear(config.dim, self.head_k_dim, bias=False), Linear(self.head_k_dim, ssm.inner_size, bias=False)
else:
self.attn_gate = Linear(config.dim, ssm.inner_size, bias=False)
@@ -306,12 +380,12 @@ class GatedDeltaNetBlock(FFNBlock):
# input processing
# Kimi-Linear is a BF16 model. Qwen 3.5 GGDN checkpoints historically use FP16 here.
x = x.cast(dtypes.bfloat16) if self.config.ssm and self.config.ssm.kda else x.half()
out_gate = self.ssm_g_b(self.ssm_g_a(x)) if hasattr(self, "ssm_g_a") else self.attn_gate(x)
out_gate = self.ssm_g_full(x) if hasattr(self, "ssm_g_full") else self.ssm_g_b(self.ssm_g_a(x)) if hasattr(self, "ssm_g_a") else self.attn_gate(x)
beta_logits = self.ssm_beta(x)
alpha_logits = self.ssm_f_b(self.ssm_f_a(x)) if hasattr(self, "ssm_f_a") else self.ssm_alpha(x)
# Causal depthwise Q/K/V convolution. Keeping the recurrence explicit gives exactly the
# same cache transition for prefill and decode (and supports arbitrary static T).
# Causal depthwise Q/K/V convolution. All tokens are projected together, then the recurrent
# update is fused into one kernel so prefill doesn't build a Python-unrolled graph.
split_qkv = hasattr(self, "attn_q")
if split_qkv:
projected_q, projected_k, projected_v = self.attn_q(x), self.attn_k(x), self.attn_v(x)
@@ -319,39 +393,40 @@ class GatedDeltaNetBlock(FFNBlock):
# overwrite their buffers before earlier outputs in a multi-token lazy graph consume them.
conv_state_q, conv_state_k, conv_state_v = self.conv_state_q.clone(), self.conv_state_k.clone(), self.conv_state_v.clone()
else: projected, conv_state = self.attn_qkv(x), self.conv_state
recurrent_state, outputs = self.recurrent_state.clone() if split_qkv else self.recurrent_state, []
for t in range(T):
if split_qkv:
conv_window_q, conv_window_k = conv_state_q.cat(projected_q[:, t:t+1], dim=1), conv_state_k.cat(projected_k[:, t:t+1], dim=1)
conv_window_v = conv_state_v.cat(projected_v[:, t:t+1], dim=1)
q = (conv_window_q * self.ssm_q_conv1d["weight"].T.unsqueeze(0)).sum(1).silu()
k = (conv_window_k * self.ssm_k_conv1d["weight"].T.unsqueeze(0)).sum(1).silu()
v = (conv_window_v * self.ssm_v_conv1d["weight"].T.unsqueeze(0)).sum(1).silu()
else:
conv_window = conv_state.cat(projected[:, t:t+1], 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, k = q.reshape(B, self.num_k_heads, self.head_k_dim), k.reshape(B, self.num_k_heads, self.head_k_dim)
q, k = (l2norm(q), l2norm(k)) if self.config.ssm and self.config.ssm.kda else (q.normalize(dim=-1), k.normalize(dim=-1))
q, k = q.repeat(1, self.num_v_heads//self.num_k_heads, 1), k.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)
def causal_conv(projected:Tensor, state:Tensor, weight:Tensor) -> tuple[Tensor, Tensor]:
window = state.cat(projected, dim=1)
out = functools.reduce(lambda a,b: a+b, (window[:, i:i+T] * weight[:, i] for i in range(self.ssm_conv_kernel))).silu()
return out, window[:, T:T+self.ssm_conv_kernel-1]
if split_qkv:
q, conv_state_q = causal_conv(projected_q, conv_state_q, self.ssm_q_conv1d["weight"])
k, conv_state_k = causal_conv(projected_k, conv_state_k, self.ssm_k_conv1d["weight"])
v, conv_state_v = causal_conv(projected_v, conv_state_v, self.ssm_v_conv1d["weight"])
else:
conv_out, conv_state = causal_conv(projected, conv_state, self.ssm_conv1d["weight"])
q, k, v = conv_out.split([self.q_dim, self.q_dim, self.conv_channels - 2*self.q_dim], dim=-1)
alpha = ((alpha_logits[:, t:t+1].float() + self.ssm_dt["bias"]).softplus().reshape(B, self.num_v_heads, -1) *
self.ssm_a.reshape(1, self.num_v_heads, -1)).exp().unsqueeze(-2)
beta = beta_logits[:, t:t+1]
beta = (beta.float() if self.config.ssm and self.config.ssm.kda else beta).sigmoid().reshape(B, self.num_v_heads, 1, 1)
recurrent_state = recurrent_state * alpha
recurrent_state = recurrent_state + ((v - recurrent_state@k) * beta)@k.transpose(-1, -2)
if t != T-1:
core_input = (recurrent_state@q).squeeze(-1).reshape(B, 1, self.num_v_heads, self.head_v_dim)
core_attn_out = self.ssm_norm(core_input.cast(x.dtype) if self.config.ssm and self.config.ssm.kda else core_input)
gate = out_gate[:, t:t+1].reshape(B, 1, self.num_v_heads, self.head_v_dim)
gate = gate.float().sigmoid().cast(core_attn_out.dtype) if hasattr(self, "ssm_g_a") else gate.silu()
outputs.append((core_attn_out * gate).reshape(B, 1, -1))
if split_qkv:
conv_state_q, conv_state_k, conv_state_v = conv_window_q[:, 1:, :], conv_window_k[:, 1:, :], conv_window_v[:, 1:, :]
else: conv_state = conv_window[:, 1:, :]
q, k = q.reshape(B, T, self.num_k_heads, self.head_k_dim), k.reshape(B, T, self.num_k_heads, self.head_k_dim)
q, k = (l2norm(q), l2norm(k)) if self.config.ssm and self.config.ssm.kda else (q.normalize(dim=-1), k.normalize(dim=-1))
q = q.repeat(1, 1, self.num_v_heads//self.num_k_heads, 1).transpose(1, 2).float() * self.head_k_dim**-0.5
k = k.repeat(1, 1, self.num_v_heads//self.num_k_heads, 1).transpose(1, 2).float()
v = v.reshape(B, T, self.num_v_heads, self.head_v_dim).transpose(1, 2).float()
beta = (beta_logits.float() if self.config.ssm and self.config.ssm.kda else beta_logits).sigmoid().transpose(1, 2)
gate_logits = (alpha_logits.float() + self.ssm_dt["bias"]).reshape(B, T, self.num_v_heads, -1)
if self.config.kda_gate_lower_bound:
log_alpha = self.config.kda_gate_lower_bound * ((-self.ssm_a).reshape(1, 1, self.num_v_heads, -1) * gate_logits).sigmoid()
else: log_alpha = gate_logits.softplus() * self.ssm_a.reshape(1, 1, self.num_v_heads, -1)
alpha = log_alpha.squeeze(-1).transpose(1, 2).exp() if log_alpha.shape[-1] == 1 else log_alpha.permute(0, 2, 1, 3).exp()
if T == 1:
# Keep decode on the small elementwise graph. The fused prefill kernel writes a temporary
# recurrent matrix, which is worthwhile for multiple tokens but needlessly copies state at T=1.
decay = alpha if len(alpha.shape) == 4 else alpha.unsqueeze(-1)
recurrent_state = self.recurrent_state * decay
k1, q1 = k[:, :, 0].unsqueeze(-1), q[:, :, 0].unsqueeze(-1)
recurrent_state = recurrent_state + ((v[:, :, 0].unsqueeze(-1) - recurrent_state@k1) * beta[:, :, 0].reshape(B, self.num_v_heads, 1, 1)) @ \
k1.transpose(-1, -2)
core = (recurrent_state @ q1).squeeze(-1).unsqueeze(2)
else: core, recurrent_state = gated_delta_prefill(q, k, v, beta, alpha, self.recurrent_state)
core = core.transpose(1, 2)
# Store each cache with its own AFTER. Multi-device lowering handles one sharded STORE per
# AFTER; grouping these effects under one cache silently drops stores on the other shards.
@@ -362,15 +437,10 @@ class GatedDeltaNetBlock(FFNBlock):
self.conv_state_v.assign(conv_state_v.cast(self.conv_state_v.dtype))]
else: state_updates = [self.conv_state.assign(conv_state.cast(self.conv_state.dtype))]
state_updates.append(self.recurrent_state.assign(recurrent_state.cast(self.recurrent_state.dtype)))
# Use the computed state for the final output. Re-reading a just-stored MULTI buffer loses
# shard-local values. The output and independent cache assignments are realized together below.
core_input = (recurrent_state@q).squeeze(-1).reshape(B, 1, self.num_v_heads, self.head_v_dim)
core_attn_out = self.ssm_norm(core_input.cast(x.dtype) if self.config.ssm and self.config.ssm.kda else core_input)
gate = out_gate[:, -1:].reshape(B, 1, self.num_v_heads, self.head_v_dim)
core_attn_out = self.ssm_norm(core.cast(x.dtype) if self.config.ssm and self.config.ssm.kda else core)
gate = out_gate.reshape(B, T, self.num_v_heads, self.head_v_dim)
gate = gate.float().sigmoid().cast(core_attn_out.dtype) if hasattr(self, "ssm_g_a") else gate.silu()
outputs.append((core_attn_out * gate).reshape(B, 1, -1))
out = outputs[0].cat(*outputs[1:], dim=1) if len(outputs) > 1 else outputs[0]
out = (core_attn_out * gate).reshape(B, T, -1)
ret = self.ssm_out(out.cast(x.dtype))
return ret.realize(*state_updates)
@@ -409,6 +479,9 @@ class Transformer:
self.token_embd = nn.Embedding(config.vocab_size, config.dim)
self.output_norm = nn.RMSNorm(config.dim, config.norm_eps)
self.output = Linear(config.dim, config.vocab_size, bias=False)
if config.attn_res_block_size:
self.output_attn_res_norm = nn.RMSNorm(config.dim, config.norm_eps)
self.output_attn_res_proj = Linear(config.dim, 1, bias=False)
self.max_context = config.max_context
self.has_recurrent_block = any(isinstance(b, GatedDeltaNetBlock) for b in self.blk)
self._cached_tokens: list[int] = []
@@ -417,20 +490,28 @@ class Transformer:
self.rollout_jit = TinyJit(self.forward)
def forward(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor) -> Tensor:
if len(tokens.shape) == 1: tokens = tokens.reshape(1, -1)
x = self.token_embd(tokens).cast(dtypes.bfloat16) if self.config.bf16_activations else self.token_embd(tokens).float()
for block in self.blk:
x = block(x, start_pos)
block_residual = Tensor.zeros(x.shape[0]*x.shape[1], 0, x.shape[2], device=x.device, dtype=x.dtype) \
if self.config.attn_res_block_size else None
for i, block in enumerate(self.blk):
if block_residual is not None: x, block_residual = block.attn_residual(x, start_pos, block_residual, i)
else: x = block(x, start_pos)
# Tensor indexing lowers selected experts through a fused one-hot reduction. Keeping all 26
# of those high-level graphs alive until the final output is scheduled exhausts host memory.
# A realization boundary lowers one block at a time; TinyJit still captures and memory-plans
# the resulting schedules for rollout replay.
if self.config.expert_mxfp4: x.realize()
if block_residual is not None:
x = FFNBlock._apply_attn_res(x.reshape(-1, x.shape[-1]), block_residual,
self.output_attn_res_proj, self.output_attn_res_norm).reshape(x.shape)
logits = self.output(self.output_norm(x))[:, -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)
def __call__(self, tokens:Tensor, start_pos:int|UOp, temperature:Tensor) -> Tensor:
return (self.prefill_jit if resolve(tokens.shape[1] != 1) else self.rollout_jit)(tokens.contiguous(), start_pos, temperature)
token_count = tokens.numel()
return (self.prefill_jit if resolve(token_count != 1) else self.rollout_jit)(tokens.flatten().contiguous(), start_pos, temperature)
@staticmethod
def from_gguf(gguf:Tensor|str|pathlib.Path, max_context:int|None=None,
@@ -523,7 +604,10 @@ class Transformer:
return min(block._reusable_prefix_len(prefix_len, len(self._cached_tokens)) for block in self.blk)
def generate(self, tokens:list[int], chunk_size:int=32, temperature:float=0.0):
if self.has_recurrent_block: chunk_size = 1
chunked_recurrent = self.has_recurrent_block and self.config.recurrent_prefill_chunked
if chunked_recurrent and self.config.recurrent_prefill_chunk_size:
chunk_size = min(chunk_size, self.config.recurrent_prefill_chunk_size)
if self.has_recurrent_block and not chunked_recurrent: chunk_size = 1
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
@@ -537,8 +621,18 @@ class Transformer:
out, prompt_len = None, len(tokens)
while len(tokens) < self.max_context:
n_toks = min(chunk_size, len(tokens) - start_pos)
sp, nt = v_start_pos.bind(start_pos), v_toks.bind(n_toks)
out = self(t[:, sp:sp+nt] if start_pos < prompt_len or out is None else out, sp, temp).realize()
# Recurrent blocks execute an explicit recurrence over T. Give them a static chunk length so
# Python constructs the recurrence once per encountered size; decode remains the T=1 JIT.
if chunked_recurrent and n_toks != 1:
# Token count is static for the recurrent kernel, but cache position must remain a runtime
# variable so repeated chunks do not replay MLA stores at the capture position.
sp = v_start_pos.bind(start_pos)
model_input = t[:, start_pos:start_pos+n_toks] if start_pos < prompt_len or out is None else out
else:
sp = v_start_pos.bind(start_pos)
nt = v_toks.bind(n_toks)
model_input = t[:, sp:sp+nt] if start_pos < prompt_len or out is None else out
out = self(model_input, sp, temp).realize()
start_pos += n_toks
# chunked prefill: keep processing until all prompt tokens are consumed
if start_pos < len(tokens): continue
+17 -5
View File
@@ -34,9 +34,10 @@ def normalize_messages(messages:list[dict]) -> None:
class StreamRouter:
# routes streamed output text to (field, text) deltas, keeping tool_call regions in .buf for the final parse
def __init__(self, reasoning:bool=False):
def __init__(self, reasoning:bool=False, xtml:bool=False):
self.buf = ""
self.mode = "reasoning" if reasoning else "undecided" # output inside a think block is sent as reasoning_content
self.xtml = xtml
def split(self, tag:str, final:bool) -> tuple[str, bool]:
# split buf on the first full tag, holding back a partial tag at the end unless final
if tag in self.buf:
@@ -51,10 +52,20 @@ class StreamRouter:
if not final and len(self.buf) < len("<think>") and "<think>".startswith(self.buf): return
self.mode, self.buf = ("reasoning", self.buf[len("<think>"):]) if self.buf.startswith("<think>") else ("content", self.buf)
if self.mode == "reasoning":
emit, done = self.split("</think>", final)
emit, done = self.split("<|close|>think<|sep|>" if self.xtml else "</think>", final)
if emit: yield "reasoning_content", emit
if not done: return
self.mode = "content_open" if self.xtml else "content"
if self.mode == "content_open":
_, found = self.split("<|open|>response<|sep|>", final)
if not found: return
self.mode = "content"
if self.mode == "done": return
if self.xtml and self.mode == "content":
emit, found = self.split("<|close|>response<|sep|>", final)
if emit: yield "content", emit
if found: self.mode = "done"
return
if self.mode == "tool": return
emit, found = self.split("<tool_call>", final)
if emit: yield "content", emit
@@ -67,7 +78,7 @@ class Handler(HTTPRequestHandler):
if self.path == "/v1/models": self.send_data(json.dumps({"object":"list","data":[{"id":self.server.model_name,"object":"model"}]}).encode())
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,
reasoning:bool=False):
reasoning:bool=False, xtml:bool=False):
model, tok = self.server.model, self.server.tok
prompt_tokens = len(ids)
cache_start_pos = model.get_start_pos(ids)
@@ -78,7 +89,7 @@ class Handler(HTTPRequestHandler):
finish_reason = "stop"
st = pt = time.perf_counter()
dec = tok.stream_decoder()
router = StreamRouter(reasoning)
router = StreamRouter(reasoning, xtml)
def log_stats(interrupted:bool=False):
et = time.perf_counter()
total = f"total:{et-st:6.2f}s"
@@ -139,9 +150,10 @@ class Handler(HTTPRequestHandler):
# reply
max_tokens = body.get("max_completion_tokens") or body.get("max_tokens")
xtml = rendered.rstrip().endswith("<|open|>think<|sep|>")
chunks = self.run_model(ids, body["model"], not body.get("stream") or body.get("stream_options",{}).get("include_usage", False),
max_tokens=max_tokens, temperature=float(body.get("temperature", 0.0)),
reasoning=rendered.rstrip().endswith("<think>"))
reasoning=xtml or rendered.rstrip().endswith("<think>"), xtml=xtml)
if body.get("stream"): self.stream_json(chunks)
else:
out, reasoning, tool_calls, finish_reason = [], [], [], "stop"