From 0565053a21792699334ad52aad17cbaf05e223a3 Mon Sep 17 00:00:00 2001 From: Ajit Mistry Date: Mon, 24 Aug 2026 18:57:19 +0200 Subject: [PATCH] [Kimi-K3] Enable KDA projection fusion for DP attention --- benchmark/kernels/kimi_k3/README.md | 71 ++++ .../kimi_k3/bench_kda_dep16_projections.py | 240 +++++++++++++ .../kernels/kimi_k3/compare_nsys_slices.py | 329 ++++++++++++++++++ python/sglang/srt/models/kimi_k3.py | 33 +- .../unit/models/test_kimi_k3_bfa_overlap.py | 60 ++++ 5 files changed, 727 insertions(+), 6 deletions(-) create mode 100644 benchmark/kernels/kimi_k3/README.md create mode 100755 benchmark/kernels/kimi_k3/bench_kda_dep16_projections.py create mode 100755 benchmark/kernels/kimi_k3/compare_nsys_slices.py diff --git a/benchmark/kernels/kimi_k3/README.md b/benchmark/kernels/kimi_k3/README.md new file mode 100644 index 000000000000..7772d5776b9a --- /dev/null +++ b/benchmark/kernels/kimi_k3/README.md @@ -0,0 +1,71 @@ +# Kimi-K3 kernel benchmarks + +## DEP16 KDA projection prologue + +`bench_kda_dep16_projections.py` measures the current unfused BF16 projection +chain used by Kimi-K3 KDA layers under full DP attention (`tp_size=dp_size=16`, +so `attn_tp_size=1`): + +```text +qkv: [M, 7168] x [7168, 36864] +beta: [M, 7168] x [7168, 96] +f_a: [M, 7168] x [7168, 128] +f_b: [M, 128] x [ 128, 12288] +gate: [M, 7168] x [7168, 12288] +``` + +The benchmark captures the complete five-GEMM chain in a CUDA graph. It can +capture multiple independent weight sets and divides the measured replay time +by that rotation count. Rotating weights avoids measuring an unrealistically +warm weight cache. + +At the serving campaign's global decode batch of 512, DEP16 normally has 32 +local tokens per attention-DP rank, so `M=32` is the primary baseline point. +A useful sweep is: + +```bash +python3 benchmark/kernels/kimi_k3/bench_kda_dep16_projections.py \ + --m 1,2,4,8,16,32,64,128 \ + --rotations 2 --warmup-replays 10 \ + --batches 7 --replays-per-batch 40 \ + --output baseline.json +``` + +For a narrow Nsight Systems capture of ten `M=32` graph replays: + +```bash +nsys profile \ + --trace=cuda,nvtx \ + --capture-range=cudaProfilerApi --capture-range-end=stop \ + --cuda-graph-trace=node --sample=none --cpuctxsw=none \ + --output=dep16-projections-m32 \ + python3 benchmark/kernels/kimi_k3/bench_kda_dep16_projections.py \ + --m 32 --rotations 2 \ + --warmup-replays 2 --batches 1 --replays-per-batch 2 \ + --nvtx-replays 10 +``` + +This is a focused GPU-step benchmark. Any production change still requires a +matched serving and Nsight A/B with the real checkpoint. + +### Baseline + +Job `517390` ran on an NVIDIA GB300 (SM103), CUDA 13.0, PyTorch +`2.11.0+cu130`, from source revision `1ef7882a5b76` using two weight rotations. +CUDA-graph replay medians were: + +| Local tokens (M) | Unfused chain | +|---:|---:| +| 1 | 123.769 us | +| 2 | 124.242 us | +| 4 | 124.535 us | +| 8 | 124.848 us | +| 16 | 125.365 us | +| **32** | **129.364 us** | +| 64 | 127.372 us | +| 128 | 129.783 us | + +The ten-replay M=32 Nsight capture contained, per chain, one wide QKV GEMM, +one small-K `f_b` GEMM, three split-K GEMMs, and three split-K reducers. This +confirms that the microbenchmark reproduces the serving kernel pattern targeted +by the fusion work. diff --git a/benchmark/kernels/kimi_k3/bench_kda_dep16_projections.py b/benchmark/kernels/kimi_k3/bench_kda_dep16_projections.py new file mode 100755 index 000000000000..ec16ddea013b --- /dev/null +++ b/benchmark/kernels/kimi_k3/bench_kda_dep16_projections.py @@ -0,0 +1,240 @@ +#!/usr/bin/env python3 +"""Benchmark Kimi-K3's unfused DEP16 KDA projection prologue. + +This reproduces the five BF16 GEMMs selected when full DP attention makes +``attn_tp_size=1`` while the global TP/EP size is 16:: + + qkv = qkv_proj(x) # 7168 -> 36864 + beta = b_proj(x) # 7168 -> 96 + fa = f_a_proj(x) # 7168 -> 128 + forget = f_b_proj(fa) # 128 -> 12288 + gate = g_proj(x) # 7168 -> 12288 + +The primary result is CUDA-graph replay time for the complete chain. Multiple +weight rotations can be captured to prevent an unrealistically warm weight +cache; reported latency is divided by the rotation count. +""" + +from __future__ import annotations + +import argparse +import json +import os +import statistics +import subprocess +import time +from dataclasses import asdict, dataclass +from pathlib import Path + +import torch +import torch.nn.functional as F + +HIDDEN = 7168 +NUM_HEADS = 96 +HEAD_DIM = 128 +PROJECTION = NUM_HEADS * HEAD_DIM +SHAPES = { + "qkv": (3 * PROJECTION, HIDDEN), + "beta": (NUM_HEADS, HIDDEN), + "f_a": (HEAD_DIM, HIDDEN), + "f_b": (PROJECTION, HEAD_DIM), + "gate": (PROJECTION, HIDDEN), +} + + +@dataclass +class Result: + m: int + rotations: int + warmup_replays: int + batches: int + replays_per_batch: int + batch_us: list[float] + median_us: float + min_us: float + max_us: float + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--m", + default="32", + help="Comma-separated local token counts. DEP16 at global batch 512 uses M=32.", + ) + parser.add_argument("--rotations", type=int, default=2) + parser.add_argument("--warmup-replays", type=int, default=10) + parser.add_argument("--batches", type=int, default=5) + parser.add_argument("--replays-per-batch", type=int, default=40) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--output", type=Path) + parser.add_argument( + "--nvtx-replays", + type=int, + default=0, + help="After timing, replay this many times in an NVTX range for nsys.", + ) + return parser.parse_args() + + +def git_revision() -> str | None: + try: + return subprocess.check_output( + ["git", "rev-parse", "HEAD"], text=True, stderr=subprocess.DEVNULL + ).strip() + except (OSError, subprocess.CalledProcessError): + return None + + +def allocate_weights(rotations: int) -> list[dict[str, torch.Tensor]]: + weights = [] + for _ in range(rotations): + current = { + name: torch.empty(shape, device="cuda", dtype=torch.bfloat16) + for name, shape in SHAPES.items() + } + # Materialize every page before timing. Values do not affect GEMM dispatch. + for weight in current.values(): + weight.fill_(0.01) + weights.append(current) + torch.cuda.synchronize() + return weights + + +def unfused_chain( + x: torch.Tensor, weights: dict[str, torch.Tensor] +) -> tuple[torch.Tensor, ...]: + qkv = F.linear(x, weights["qkv"]) + beta = F.linear(x, weights["beta"]) + f_a = F.linear(x, weights["f_a"]) + forget = F.linear(f_a, weights["f_b"]) + gate = F.linear(x, weights["gate"]) + return qkv, beta, forget, gate + + +def capture_graph( + x: torch.Tensor, weight_sets: list[dict[str, torch.Tensor]] +) -> tuple[torch.cuda.CUDAGraph, tuple[torch.Tensor, ...]]: + side_stream = torch.cuda.Stream() + side_stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side_stream): + for _ in range(3): + for weights in weight_sets: + outputs = unfused_chain(x, weights) + torch.cuda.current_stream().wait_stream(side_stream) + torch.cuda.synchronize() + + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for weights in weight_sets: + outputs = unfused_chain(x, weights) + return graph, outputs + + +def benchmark_m(args: argparse.Namespace, m: int, weight_sets) -> Result: + x = torch.randn((m, HIDDEN), device="cuda", dtype=torch.bfloat16) + graph, outputs = capture_graph(x, weight_sets) + + # Keep graph outputs alive and make accidental removal obvious. + assert len(outputs) == 4 and outputs[0].shape == (m, 3 * PROJECTION) + + for _ in range(args.warmup_replays): + graph.replay() + torch.cuda.synchronize() + + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + samples = [] + for _ in range(args.batches): + start.record() + for _ in range(args.replays_per_batch): + graph.replay() + end.record() + end.synchronize() + per_chain_us = ( + start.elapsed_time(end) * 1000.0 / args.replays_per_batch / args.rotations + ) + samples.append(per_chain_us) + + if args.nvtx_replays: + # This also provides a narrow capture range for: + # nsys profile --capture-range=cudaProfilerApi ... + cudart = torch.cuda.cudart() + cudart.cudaProfilerStart() + torch.cuda.nvtx.range_push(f"k3_dep16_unfused_m{m}") + for _ in range(args.nvtx_replays): + graph.replay() + torch.cuda.nvtx.range_pop() + torch.cuda.synchronize() + cudart.cudaProfilerStop() + + return Result( + m=m, + rotations=args.rotations, + warmup_replays=args.warmup_replays, + batches=args.batches, + replays_per_batch=args.replays_per_batch, + batch_us=samples, + median_us=statistics.median(samples), + min_us=min(samples), + max_us=max(samples), + ) + + +def main() -> None: + args = parse_args() + if not torch.cuda.is_available(): + raise RuntimeError("CUDA is required") + if args.rotations < 1: + raise ValueError("--rotations must be positive") + + ms = [int(value) for value in args.m.split(",")] + if any(m <= 0 for m in ms): + raise ValueError("all M values must be positive") + + torch.manual_seed(args.seed) + torch.cuda.manual_seed_all(args.seed) + torch.set_grad_enabled(False) + + device = torch.cuda.current_device() + props = torch.cuda.get_device_properties(device) + metadata = { + "benchmark": "kimi_k3_kda_dep16_unfused_projections", + "timestamp": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), + "git_revision": git_revision() or os.environ.get("K3_BENCH_GIT_REV"), + "torch_version": torch.__version__, + "cuda_version": torch.version.cuda, + "device": props.name, + "device_capability": list(torch.cuda.get_device_capability(device)), + "device_memory_bytes": props.total_memory, + "environment": { + key: os.environ.get(key) + for key in ("CUDA_VISIBLE_DEVICES", "NVIDIA_VISIBLE_DEVICES") + }, + "shapes": {name: list(shape) for name, shape in SHAPES.items()}, + } + + print(json.dumps(metadata, indent=2), flush=True) + print(f"Allocating {args.rotations} rotating weight sets...", flush=True) + weight_sets = allocate_weights(args.rotations) + + results = [] + for m in ms: + result = benchmark_m(args, m, weight_sets) + results.append(result) + print( + f"M={m:4d}: median={result.median_us:9.3f} us " + f"min={result.min_us:9.3f} us max={result.max_us:9.3f} us " + f"samples={','.join(f'{v:.3f}' for v in result.batch_us)}", + flush=True, + ) + + document = {**metadata, "results": [asdict(result) for result in results]} + if args.output: + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps(document, indent=2) + "\n") + print(f"Wrote {args.output}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/benchmark/kernels/kimi_k3/compare_nsys_slices.py b/benchmark/kernels/kimi_k3/compare_nsys_slices.py new file mode 100755 index 000000000000..a5c1de7038e2 --- /dev/null +++ b/benchmark/kernels/kimi_k3/compare_nsys_slices.py @@ -0,0 +1,329 @@ +#!/usr/bin/env python3 +"""Crop and align matching CUDA-kernel slices from two Nsight SQLite exports. + +The output uses the Chrome/Perfetto trace-event format. Independent captures +are placed in separate process tracks and normalized to the selected marker's +start, allowing their GPU timelines to be viewed at the same horizontal scale. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import sqlite3 +import statistics +from dataclasses import dataclass +from pathlib import Path + + +@dataclass(frozen=True) +class Kernel: + start_ns: int + end_ns: int + stream_id: int + short_name: str + full_name: str + + +@dataclass(frozen=True) +class Slice: + step: int + start_ns: int + end_ns: int + marker_starts_ns: tuple[int, ...] + kernels: tuple[Kernel, ...] + + @property + def wall_us(self) -> float: + return (self.end_ns - self.start_ns) / 1000.0 + + @property + def summed_us(self) -> float: + return sum(k.end_ns - k.start_ns for k in self.kernels) / 1000.0 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("baseline", type=Path) + parser.add_argument("candidate", type=Path) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--marker", default="sm100_fp8_fp4_mega_moe_impl") + parser.add_argument("--markers-per-step", type=int, default=92) + parser.add_argument("--start-marker", type=int, default=83) + parser.add_argument("--end-marker", type=int, default=87) + return parser.parse_args() + + +def load_slices( + path: Path, + marker_name: str, + markers_per_step: int, + start_marker: int, + end_marker: int, +) -> list[Slice]: + connection = sqlite3.connect(path) + marker_starts = [ + row[0] + for row in connection.execute( + """ + SELECT k.start + FROM CUPTI_ACTIVITY_KIND_KERNEL AS k + JOIN StringIds AS s ON s.id = k.shortName + WHERE k.deviceId = 0 AND s.value = ? + ORDER BY k.start + """, + (marker_name,), + ) + ] + if not marker_starts or len(marker_starts) % markers_per_step: + raise RuntimeError( + f"{path}: found {len(marker_starts)} markers; expected a positive " + f"multiple of {markers_per_step}" + ) + if not 0 <= start_marker < end_marker < markers_per_step: + raise ValueError("marker indices must satisfy 0 <= start < end < markers/step") + + slices = [] + for step in range(len(marker_starts) // markers_per_step): + base = step * markers_per_step + selected_markers = tuple( + marker_starts[base + index] for index in range(start_marker, end_marker + 1) + ) + start_ns, end_ns = selected_markers[0], selected_markers[-1] + kernels = tuple( + Kernel(*row) + for row in connection.execute( + """ + SELECT k.start, k.end, k.streamId, short.value, full.value + FROM CUPTI_ACTIVITY_KIND_KERNEL AS k + JOIN StringIds AS short ON short.id = k.shortName + JOIN StringIds AS full ON full.id = k.demangledName + WHERE k.deviceId = 0 AND k.start >= ? AND k.start < ? + ORDER BY k.start + """, + (start_ns, end_ns), + ) + ) + slices.append( + Slice( + step=step, + start_ns=start_ns, + end_ns=end_ns, + marker_starts_ns=selected_markers, + kernels=kernels, + ) + ) + connection.close() + return slices + + +def representative(slices: list[Slice]) -> Slice: + median = statistics.median(item.wall_us for item in slices) + return min(slices, key=lambda item: (abs(item.wall_us - median), item.step)) + + +def category(name: str) -> str: + if name.startswith("nvjet") or "splitKreduce" in name: + return "projection" + if "mega_moe" in name or "fused_front_epilogue" in name: + return "moe" + if "recurrent_kda" in name or "causal_conv1d" in name: + return "kda" + if "cutlass_split_kv" in name or "mla" in name: + return "mla" + if "attn_res" in name: + return "residual" + return "other" + + +def assign_overlap_lanes(kernels: tuple[Kernel, ...]) -> list[tuple[Kernel, int]]: + """Color overlapping intervals so each emitted Perfetto track is linear.""" + assignments = [] + lane_ends: dict[int, list[int]] = {} + for kernel in kernels: + ends = lane_ends.setdefault(kernel.stream_id, []) + lane = next( + (index for index, end_ns in enumerate(ends) if end_ns <= kernel.start_ns), + len(ends), + ) + if lane == len(ends): + ends.append(kernel.end_ns) + else: + ends[lane] = kernel.end_ns + assignments.append((kernel, lane)) + return assignments + + +def trace_events(label: str, item: Slice, pid: int, start_marker: int) -> list[dict]: + events: list[dict] = [ + { + "ph": "M", + "name": "process_name", + "pid": pid, + "tid": 0, + "args": {"name": f"{label} (capture step {item.step})"}, + }, + { + "ph": "M", + "name": "thread_name", + "pid": pid, + "tid": 0, + "args": {"name": "Layer intervals"}, + }, + ] + + assignments = assign_overlap_lanes(item.kernels) + lanes = sorted({(kernel.stream_id, lane) for kernel, lane in assignments}) + lane_tids = {lane: index + 1 for index, lane in enumerate(lanes)} + for stream_id, lane in lanes: + suffix = "" if lane == 0 else f" — overlap lane {lane}" + events.append( + { + "ph": "M", + "name": "thread_name", + "pid": pid, + "tid": lane_tids[(stream_id, lane)], + "args": {"name": f"CUDA stream {stream_id}{suffix}"}, + } + ) + + for offset, (begin, end) in enumerate( + zip(item.marker_starts_ns, item.marker_starts_ns[1:]) + ): + begin_us = (begin - item.start_ns) / 1000.0 + end_us = (end - item.start_ns) / 1000.0 + # Leave a 1 ns visual gap. Otherwise floating-point addition of one + # interval's ts+dur can round just beyond the next interval's ts and + # trigger Perfetto's overlapping-complete-event importer warning. + display_duration_us = max(0.0, end_us - begin_us - 0.001) + events.append( + { + "ph": "X", + "name": f"marker {start_marker + offset} → {start_marker + offset + 1}", + "cat": "layer interval", + "pid": pid, + "tid": 0, + "ts": begin_us, + "dur": display_duration_us, + "args": {"wall_us": (end - begin) / 1000.0}, + } + ) + + for kernel, lane in assignments: + events.append( + { + "ph": "X", + "name": kernel.short_name, + "cat": category(kernel.short_name), + "pid": pid, + "tid": lane_tids[(kernel.stream_id, lane)], + "ts": (kernel.start_ns - item.start_ns) / 1000.0, + "dur": (kernel.end_ns - kernel.start_ns) / 1000.0, + "args": { + "full_name": kernel.full_name, + "cuda_stream": kernel.stream_id, + "overlap_lane": lane, + "absolute_start_ns": kernel.start_ns, + "duration_us": (kernel.end_ns - kernel.start_ns) / 1000.0, + }, + } + ) + return events + + +def write_trace( + path: Path, + entries: list[tuple[str, Slice]], + start_marker: int, + end_marker: int, +) -> None: + events = [] + for pid, (label, item) in enumerate(entries, start=1): + events.extend(trace_events(label, item, pid, start_marker)) + document = { + "displayTimeUnit": "us", + "traceEvents": events, + "otherData": { + "selection": f"markers {start_marker} through {end_marker}", + "normalization": "each process track starts at its selected marker", + }, + } + path.write_text(json.dumps(document, separators=(",", ":")) + "\n") + + +def write_samples(path: Path, variants: list[tuple[str, list[Slice]]]) -> None: + with path.open("w", newline="") as output: + writer = csv.writer(output) + writer.writerow(["variant", "step", "wall_us", "summed_kernel_us", "kernels"]) + for label, slices in variants: + for item in slices: + writer.writerow( + [label, item.step, item.wall_us, item.summed_us, len(item.kernels)] + ) + + +def main() -> None: + args = parse_args() + args.output_dir.mkdir(parents=True, exist_ok=True) + baseline = load_slices( + args.baseline, + args.marker, + args.markers_per_step, + args.start_marker, + args.end_marker, + ) + candidate = load_slices( + args.candidate, + args.marker, + args.markers_per_step, + args.start_marker, + args.end_marker, + ) + baseline_rep = representative(baseline) + candidate_rep = representative(candidate) + + write_trace( + args.output_dir / "four-layer-baseline.perfetto.json", + [("Baseline", baseline_rep)], + args.start_marker, + args.end_marker, + ) + write_trace( + args.output_dir / "four-layer-fused.perfetto.json", + [("Fused", candidate_rep)], + args.start_marker, + args.end_marker, + ) + write_trace( + args.output_dir / "four-layer-comparison.perfetto.json", + [("Baseline", baseline_rep), ("Fused", candidate_rep)], + args.start_marker, + args.end_marker, + ) + write_trace( + args.output_dir / "four-layer-all-steps.perfetto.json", + [(f"Baseline step {item.step}", item) for item in baseline] + + [(f"Fused step {item.step}", item) for item in candidate], + args.start_marker, + args.end_marker, + ) + write_samples( + args.output_dir / "four-layer-samples.csv", + [("baseline", baseline), ("fused", candidate)], + ) + + for label, slices, selected in ( + ("baseline", baseline, baseline_rep), + ("fused", candidate, candidate_rep), + ): + print( + f"{label}: median={statistics.median(x.wall_us for x in slices):.3f} us; " + f"representative step={selected.step}, wall={selected.wall_us:.3f} us, " + f"kernels={len(selected.kernels)}" + ) + print(f"Wrote comparison files to {args.output_dir}") + + +if __name__ == "__main__": + main() diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index afe3d9c30f0c..1420aeef4f6d 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -1383,6 +1383,22 @@ def forward( return out.view(num_tokens, hidden_size) +def _should_fuse_kda_projections( + *, + use_full_rank_gate: bool, + quant_config: Optional[QuantizationConfig], + tp_size: int, + attn_tp_size: int, +) -> bool: + """Whether the checkpoint projections can use a supported fused layout.""" + if use_full_rank_gate: + # K3's merged QKVG linear takes explicit attention-TP rank/size, so it + # also supports DP attention where attention weights are replicated. + return True + # The low-rank repeated/batched layout still follows the global TP group. + return quant_config is None and attn_tp_size == tp_size + + class KimiK3DeltaAttention(nn.Module): """KDA attention; optional full-rank gate.""" @@ -1439,10 +1455,15 @@ def __init__( quant_config, f"{prefix}.b_proj" ) - # The fused path hardcodes tp_size sharding, so require attn_tp == tp. - # Full-rank K3 also fuses mixed block-FP8 attention projections. - self.do_fuse_qkvbfg = self.attn_tp_size == self.tp_size and ( - quant_config is None or self.use_full_rank_gate + # The full-rank K3 layout passes the attention-TP rank/size explicitly + # to every head-sharded projection, including under DP attention, and + # supports the mixed block-FP8 attention projections. The low-rank + # repeated/batched layout still requires matching attention/global TP. + self.do_fuse_qkvbfg = _should_fuse_kda_projections( + use_full_rank_gate=self.use_full_rank_gate, + quant_config=quant_config, + tp_size=self.tp_size, + attn_tp_size=self.attn_tp_size, ) if self.do_fuse_qkvbfg and self.use_full_rank_gate: @@ -1465,8 +1486,8 @@ def __init__( prefix=f"{prefix}.fused_qkvg_proj", ) self.split_sizes = [ - 3 * projection_size // self.tp_size, - projection_size // self.tp_size, + 3 * projection_size // self.attn_tp_size, + projection_size // self.attn_tp_size, ] self.b_proj = ColumnParallelLinear( self.hidden_size, diff --git a/test/registered/unit/models/test_kimi_k3_bfa_overlap.py b/test/registered/unit/models/test_kimi_k3_bfa_overlap.py index db856e293fe8..ff80bb3b5217 100644 --- a/test/registered/unit/models/test_kimi_k3_bfa_overlap.py +++ b/test/registered/unit/models/test_kimi_k3_bfa_overlap.py @@ -11,6 +11,7 @@ from sglang.srt.models.kimi_k3 import ( KimiK3DeltaAttention, _get_k3_dense_weight, + _should_fuse_kda_projections, ) from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import CustomTestCase @@ -41,6 +42,7 @@ def fused_qkvg_proj(x): owner = SimpleNamespace( use_full_rank_gate=True, + _qkvg_w=qkvg_w, _bfa_w=_randn(_BFA_W_ROWS, _H).contiguous(), _bfa_f_b_w=_randn(1536, _N_FA).contiguous(), _bfa_fa_size=_N_FA, @@ -58,12 +60,70 @@ def _run(owner, x): return [t.clone() for t in out] +class TestKimiK3ProjectionFusionPolicy(unittest.TestCase): + def test_full_rank_gate_supports_dp_attention(self): + self.assertTrue( + _should_fuse_kda_projections( + use_full_rank_gate=True, + quant_config=object(), + tp_size=16, + attn_tp_size=1, + ) + ) + + def test_low_rank_layout_still_requires_matching_tp(self): + self.assertTrue( + _should_fuse_kda_projections( + use_full_rank_gate=False, + quant_config=None, + tp_size=16, + attn_tp_size=16, + ) + ) + self.assertFalse( + _should_fuse_kda_projections( + use_full_rank_gate=False, + quant_config=None, + tp_size=16, + attn_tp_size=1, + ) + ) + + class TestKimiK3BfaOverlap(CustomTestCase): @classmethod def setUpClass(cls): if not torch.cuda.is_available(): raise unittest.SkipTest("CUDA is not available") + def test_fused_projections_match_unfused_math(self): + owner = _make_owner(with_stream=False) + qkv_w, gate_w = torch.split(owner._qkvg_w, owner.split_sizes) + fa_w = owner._bfa_w[: owner._bfa_fa_size] + beta_w = owner._bfa_w[ + owner._bfa_fa_size : owner._bfa_fa_size + owner._bfa_b_size + ] + + for num_tokens in (1, 32): + with self.subTest(num_tokens=num_tokens): + x = torch.randn(num_tokens, _H, device="cuda", dtype=torch.bfloat16) + fused = _run(owner, x) + fa = torch.nn.functional.linear(x, fa_w) + unfused = ( + torch.nn.functional.linear(x, qkv_w), + torch.nn.functional.linear(x, beta_w), + torch.nn.functional.linear(fa, owner._bfa_f_b_w), + torch.nn.functional.linear(x, gate_w), + ) + for got, ref, name in zip( + fused, unfused, ("qkv", "beta", "forget_gate", "gate") + ): + # The merged and standalone BF16 GEMMs may select different + # accumulation schedules; allow roughly one output ULP. + torch.testing.assert_close( + got, ref, rtol=1e-2, atol=2e-2, msg=lambda msg: f"{name}: {msg}" + ) + def test_capture_replay_matches_serial(self): torch.manual_seed(0) for T in (1, 4, 12):