add fast exact-shape Kimi K3 benchmark

This commit is contained in:
2026-08-10 06:13:24 -07:00
parent f53f0e7e79
commit 2b1b8c22a9
2 changed files with 95 additions and 0 deletions
+11
View File
@@ -49,6 +49,15 @@ 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.
For performance iteration, use the exact-width fake-weight harness before another official load:
```sh
DEV=AMD python extra/benchmark_kimi_k3_fake.py --mode attention --iterations 20
DEV=AMD python extra/benchmark_kimi_k3_fake.py --mode block --iterations 20
```
It retains K3's 7,168-wide residual stream, 12,288-wide KDA state, 96 heads, 128×128 recurrent matrices, TP8 layouts, top-k 16 routing, packed MXFP4 expert shapes, collectives, and decode JIT, but uses one layer and 16 fake experts. Fake attention weights initialize in about 0.9 seconds and the full block in about 3 seconds. The retained path measured 0.630 ms per fake attention layer and 1.367 ms per complete fake block, projecting about 7.87 tok/s across 93 identical blocks versus 6.25 tok/s for the official heterogeneous model. Treat this as a candidate admission benchmark, not a correctness substitute for official weights.
## 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.
@@ -115,6 +124,8 @@ A final load-first experiment increased the disk-to-HBM io_uring queue depth fro
Two direct packed-expert MFMA prototypes were also rejected after that load. A fused gate/up kernel was about 29% faster in isolation at the TP8-local shape, and a routed-down kernel which combined projection, probability weighting, and route reduction measured 1.45 ms versus 2.42 ms in isolation. End-to-end, however, stable replay produced `[198, 2338, 2127, 148297]`, prefill measured 38.87 tok/s, and decode measured 6.263 tok/s. That is indistinguishable from the retained 38.65/6.25 tok/s path while changing floating-point reduction order, so neither kernel was retained.
A whole-core KDA decode experiment fused convolution, Q/K normalization, channel decay, recurrence, RMS normalization, output gating, and four persistent state updates. Its raw kernel replayed in about 109 microseconds per local KDA layer and matched a one-step synthetic reference within `9.77e-4` output and `8.13e-4` state maximum error. The exact-width fake-layer gate caught that it was slower than the retained attention path (0.665 versus 0.633 ms/layer). The already-running official validation was stopped after its first invalid greedy sequence, `[198, 163840, 163840, 163840]`, where 163840 is outside the checkpoint's vocabulary. The kernel was rejected and removed.
| Maximum context | Load | Short-prompt replay | Result |
|---:|---:|---:|---|
| 128 | 489.59s | 14.32s | stable 8-token replay |
+84
View File
@@ -0,0 +1,84 @@
#!/usr/bin/env python3
"""Fast exact-shape K3 KDA/layer benchmark using bounded fake weights instead of the 1.56 TB checkpoint."""
from __future__ import annotations
import argparse, statistics, time
from dataclasses import replace
from tinygrad import Device, Tensor, TinyJit, dtypes, nn
from tinygrad.helpers import profile_marker
from tinygrad.llm.kimi_k3 import kimi_k3_config
from tinygrad.llm.model import GatedDeltaNetBlock
def tp_axis(name:str) -> int|None:
if "ffn_gate_exps.weight" in name or "ffn_up_exps.weight" in name: return 1
if "ffn_gate_exps.weight_scale" in name or "ffn_up_exps.weight_scale" in name: return 1
if "ffn_down_exps.weight" in name or "ffn_down_exps.weight_scale" in name: return 2
if name.endswith(("ffn_gate_shexp.weight", "ffn_up_shexp.weight")): return 0
if name.endswith(("ffn_down_shexp.weight", "ffn_routed_down.weight", "ffn_routed_up.weight", "ssm_out.weight")): return 1
if name.endswith(("attn_q.weight", "attn_k.weight", "attn_v.weight", "ssm_g_full.weight", "ssm_f_b.weight", "ssm_beta.weight")): return 0
if name.endswith(("ssm_q_conv1d.weight", "ssm_k_conv1d.weight", "ssm_v_conv1d.weight", "ssm_dt.bias")): return 0
return None
def fake_value(name:str) -> tuple[int|float, object]:
if name.endswith("weight_scale"): return 120, dtypes.uint8
if name.endswith("_exps.weight"): return 0x11, dtypes.uint8
if name.endswith("ssm_a"): return -0.1, dtypes.float32
if name.endswith("ssm_dt.bias"): return 0.1, dtypes.float32
if "conv1d.weight" in name: return 0.1, dtypes.float32
if name.endswith("exp_probs_b.bias"): return 0.0, dtypes.float32
if name.endswith("norm.weight"): return 1.0, dtypes.bfloat16
return 0.001, dtypes.bfloat16
def fake_tp_tensor(shape:tuple[int, ...], value:int|float, dtype, devices:tuple[str, ...], axis:int|None) -> Tensor:
if axis is not None and shape[axis] % len(devices): raise ValueError(f"shape {shape} is not TP{len(devices)} divisible on axis {axis}")
source = Tensor.full(shape, value, dtype=dtype, device=devices[0]).clone().realize()
return source.shard(devices, axis=axis).realize()
def sync(devices:tuple[str, ...]) -> None:
for device in devices: Device[device].synchronize()
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--devices", type=int, default=8)
parser.add_argument("--mode", choices=("attention", "block"), default="attention")
parser.add_argument("--iterations", type=int, default=20)
args = parser.parse_args()
devices = tuple(f"AMD:{i}" for i in range(args.devices))
# One exact-width KDA layer, but only 16 fake routed experts. This retains top-k 16 and every
# official per-GPU matrix/state shape while keeping fake expert storage below 300 MB per layer.
config = replace(kimi_k3_config(4), num_blocks=1, num_experts=16, num_experts_per_tok=16, ssm_layers=(True,),
attn_res_block_size=0)
block = GatedDeltaNetBlock(config, config.ssm)
begin = time.perf_counter()
for name,tensor in nn.state.get_state_dict(block).items():
if args.mode == "attention" and name.startswith(("ffn_", "exp_probs_")): continue
value, dtype = fake_value(name)
tensor.replace(fake_tp_tensor(tuple(int(x) for x in tensor.shape), value, dtype, devices, tp_axis(name)))
sync(devices)
print(f"fake weights: {time.perf_counter()-begin:.3f}s", flush=True)
x_source = (((Tensor.arange(config.dim, dtype=dtypes.float32).reshape(1, 1, config.dim) % 31) / 31) \
.cast(dtypes.bfloat16).to(devices[0])).clone().realize()
x = x_source.shard(devices, axis=None).realize()
block._init_state(x)
# Use direct buffer-backed state shards. The production path reaches this form after prefill;
# the fake harness begins immediately at decode and must not feed lazy clone graphs to TinyJit.
for state,axis in ((block.conv_state_q, 2), (block.conv_state_k, 2), (block.conv_state_v, 2), (block.recurrent_state, 1)):
state.replace(Tensor.zeros(*state.shape, dtype=state.dtype, device=devices[0]).shard(devices, axis=axis).realize())
@TinyJit
def run(inp:Tensor) -> Tensor:
if args.mode == "attention": return block._attention(block.attn_norm(inp), 0).realize()
return block(inp, 0).realize()
# uncaptured, capture, then replay only
run(x); sync(devices)
run(x); sync(devices)
samples:list[float] = []
profile_marker(f"fake K3 {args.mode} start")
for _ in range(args.iterations):
begin = time.perf_counter(); out = run(x); sync(devices); samples.append((time.perf_counter()-begin)*1e3)
profile_marker(f"fake K3 {args.mode} end")
print(f"{args.mode}: median={statistics.median(samples):.3f} ms/layer, min={min(samples):.3f} ms/layer, "
f"projected_93_layer_rate={1000/(statistics.median(samples)*93):.3f} tok/s, finite={out.float().isfinite().all().item()}")
if __name__ == "__main__": main()