feat: audit K3 open model artifacts

This commit is contained in:
wuyang
2026-07-29 12:44:37 +08:00
parent c3d332680c
commit be2b291a04
13 changed files with 2387 additions and 29 deletions
+388
View File
@@ -0,0 +1,388 @@
#!/usr/bin/env python3
"""Build a small, auditable snapshot from Kimi K3's public model artifacts.
The script intentionally does not download a checkpoint. It consumes:
1. the public config and safetensors index;
2. safetensors JSON headers fetched with HTTP Range;
3. two small byte ranges containing one KDA parameter prefix and one MoE
router prefix;
4. a local checkout of the official FlashKDA repository.
Raw model bytes stay local. The generated JSON contains only aggregate
statistics, public shapes, revisions, checksums, and a clearly labelled
synthetic-input router stress probe.
"""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import platform
import re
import subprocess
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
import torch
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--config", type=Path, required=True)
parser.add_argument("--index", type=Path, required=True)
parser.add_argument("--hf-model", type=Path, required=True)
parser.add_argument("--kda-slice", type=Path, required=True)
parser.add_argument("--router-prefix", type=Path, required=True)
parser.add_argument("--mla-header", type=Path, required=True)
parser.add_argument("--vision-header", type=Path, required=True)
parser.add_argument("--flashkda-dir", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--synthetic-tokens", type=int, default=2048)
parser.add_argument("--seed", type=int, default=20260729)
parser.add_argument("--captured-at", default=None)
return parser.parse_args()
def read_json(path: Path) -> dict[str, Any]:
return json.loads(path.read_text())
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for block in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def summarize(values: torch.Tensor) -> dict[str, Any]:
flat = values.detach().float().flatten().cpu()
points = torch.tensor([0, 0.01, 0.1, 0.25, 0.5, 0.75, 0.9, 0.99, 1])
quantiles = torch.quantile(flat, points).tolist()
labels = ["min", "p01", "p10", "p25", "p50", "p75", "p90", "p99", "max"]
return {
"count": flat.numel(),
"mean": flat.mean().item(),
"std": flat.std().item(),
"quantiles": dict(zip(labels, quantiles, strict=True)),
}
def tensor_fact(header: dict[str, Any], name: str) -> dict[str, Any]:
entry = header[name]
return {
"name": name,
"dtype": entry["dtype"],
"shape": entry["shape"],
"bytes": entry["data_offsets"][1] - entry["data_offsets"][0],
}
def parse_benchmark(path: Path, heads: int = 96) -> dict[str, Any]:
text = path.read_text()
section = text.split(f"### `T=8192`, `H={heads}`, `D=128`", 1)[1].split("###", 1)[0]
rows = {}
for label, flash, chunk, speedup, gdn, gdn_speedup in re.findall(
r"\| ([^|]+?) \| ([0-9.]+) \| ([0-9.]+) \| ([0-9.]+)× \| ([0-9.]+) \| ([0-9.]+)× \|",
section,
):
rows[label.strip()] = {
"flash_kda_ms": float(flash),
"fla_chunk_kda_ms": float(chunk),
"speedup_vs_chunk_kda": float(speedup),
"fla_chunk_gdn_ms": float(gdn),
"speedup_vs_gdn": float(gdn_speedup),
}
return {"sequence": 8192, "heads": heads, "dimension": 128, "rows": rows}
def load_metrics(load: torch.Tensor) -> dict[str, Any]:
mean = load.mean()
ordered = load.sort().values
count = load.numel()
indices = torch.arange(1, count + 1, device=load.device, dtype=torch.float32)
gini = ((2 * indices - count - 1) * ordered).sum() / (count * ordered.sum())
quantiles = torch.quantile(
load,
torch.tensor([0, 0.1, 0.25, 0.5, 0.75, 0.9, 0.99, 1], device=load.device),
).tolist()
labels = ["min", "p10", "p25", "p50", "p75", "p90", "p99", "max"]
return {
"mean": mean.item(),
"std": load.std().item(),
"cv": (load.std() / mean).item(),
"gini": gini.item(),
"zero_experts": int((load == 0).sum()),
"quantiles": dict(zip(labels, quantiles, strict=True)),
}
def main() -> None:
args = parse_args()
config = read_json(args.config)
text = config["text_config"]
index = read_json(args.index)
hf_model = read_json(args.hf_model)
mla_header = read_json(args.mla_header)
vision_header = read_json(args.vision_header)
names = list(index["weight_map"])
shard_files = [
item
for item in hf_model["siblings"]
if re.fullmatch(r"model-\d+-of-\d+\.safetensors", item["rfilename"])
]
shard_sizes = [item["size"] for item in shard_files]
kda_raw = args.kda_slice.read_bytes()
if len(kda_raw) != 49_664:
raise ValueError(f"unexpected KDA slice length: {len(kda_raw)}")
a_log = torch.frombuffer(bytearray(kda_raw[:512]), dtype=torch.float32).clone()
dt_bias = torch.frombuffer(bytearray(kda_raw[512:]), dtype=torch.float32).clone().view(96, 128)
router_raw = args.router_prefix.read_bytes()
if len(router_raw) != 13_488_640:
raise ValueError(f"unexpected router prefix length: {len(router_raw)}")
correction_bias = torch.frombuffer(
bytearray(router_raw[:3584]), dtype=torch.float32
).clone()
router_weight = torch.frombuffer(
bytearray(router_raw[643_584:13_488_640]), dtype=torch.bfloat16
).clone().view(896, 7168)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.manual_seed(args.seed)
if device.type == "cuda":
torch.cuda.manual_seed_all(args.seed)
router_device = router_weight.to(device)
bias_device = correction_bias.to(device)
synthetic = torch.randn(
args.synthetic_tokens, 7168, device=device, dtype=torch.float32
)
synthetic *= torch.rsqrt(synthetic.square().mean(-1, keepdim=True) + 1e-6)
scores = torch.sigmoid(synthetic.to(torch.bfloat16) @ router_device.T).float()
unbiased_ids = scores.topk(16, dim=-1).indices
biased_ids = (scores + bias_device).topk(16, dim=-1).indices
def loads(ids: torch.Tensor) -> torch.Tensor:
return torch.bincount(ids.flatten(), minlength=896).float()
unbiased_load = loads(unbiased_ids)
biased_load = loads(biased_ids)
overlap = torch.tensor(
[
len(set(unbiased_ids[row].tolist()) & set(biased_ids[row].tolist()))
for row in range(args.synthetic_tokens)
],
dtype=torch.float32,
)
router_norms = router_weight.float().norm(dim=1)
# This is deliberately a hypothesis probe, not a canonical forward pass:
# the checkpoint stores A_log[128], while public code/API expect A_log[H=96].
channelwise_log_decay = -5.0 * torch.sigmoid(torch.exp(a_log).view(1, 128) * dt_bias)
channelwise_retention = torch.exp(channelwise_log_decay)
per_expert_bytes = 3 * (5_505_024 + 344_064)
all_routed_expert_bytes = per_expert_bytes * 896 * 92
kda_layers = text["linear_attn_config"]["kda_layers"]
mla_layers = text["linear_attn_config"]["full_attn_layers"]
flash_revision = subprocess.check_output(
["git", "-C", str(args.flashkda_dir), "rev-parse", "HEAD"],
text=True,
).strip()
captured_at = args.captured_at or datetime.now(timezone.utc).isoformat()
result = {
"schema_version": 1,
"captured_at": captured_at,
"evidence_boundary": {
"checkpoint_forward_run": False,
"raw_weights_committed": False,
"router_inputs": "deterministic synthetic RMS-normalized vectors, not token hidden states",
"kda_retention_probe": "noncanonical channel-wise interpretation used only to expose the A_log shape ambiguity",
},
"provenance": {
"huggingface_model": "moonshotai/Kimi-K3",
"huggingface_revision": hf_model["sha"],
"flashkda_revision": flash_revision,
"sha256": {
"config": sha256(args.config),
"index": sha256(args.index),
"kda_slice": sha256(args.kda_slice),
"router_prefix": sha256(args.router_prefix),
},
},
"checkpoint": {
"tensor_data_bytes": index["metadata"]["total_size"],
"tensor_data_tb": index["metadata"]["total_size"] / 1e12,
"tensor_data_tib": index["metadata"]["total_size"] / 2**40,
"shards": len(shard_files),
"shard_file_bytes": {
"sum": sum(shard_sizes),
"min": min(shard_sizes),
"max": max(shard_sizes),
"mean": sum(shard_sizes) / len(shard_sizes),
},
"tensor_entries": len(names),
"tensor_counts": {
"expert_packed": sum(
bool(re.search(r"experts\.\d+\.w[123]\.weight_packed$", name))
for name in names
),
"expert_scales": sum(
bool(re.search(r"experts\.\d+\.w[123]\.weight_scale$", name))
for name in names
),
"router_weight": sum(
name.endswith("block_sparse_moe.gate.weight") for name in names
),
"router_correction_bias": sum(
name.endswith("gate.e_score_correction_bias") for name in names
),
"attnres_proj": sum(
bool(re.search(r"(_res_proj|output_attn_res_proj)\.weight$", name))
for name in names
),
"attnres_norm": sum(
bool(re.search(r"(_res_norm|output_attn_res_norm)\.weight$", name))
for name in names
),
"kda_a_log": sum(name.endswith("self_attn.A_log") for name in names),
"kda_dt_bias": sum(name.endswith("self_attn.dt_bias") for name in names),
"vision": sum(name.startswith("vision_tower.") for name in names),
"projector": sum(name.startswith("mm_projector.") for name in names),
},
"derived_routed_expert_bytes": all_routed_expert_bytes,
"derived_routed_expert_share": all_routed_expert_bytes
/ index["metadata"]["total_size"],
},
"configuration": {
"layers": text["num_hidden_layers"],
"dense_layers": text["first_k_dense_replace"],
"hidden": text["hidden_size"],
"vocabulary": text["vocab_size"],
"context": text["max_position_embeddings"],
"kda_layers": kda_layers,
"mla_layers": mla_layers,
"heads": text["num_attention_heads"],
"head_dim": text["linear_attn_config"]["head_dim"],
"attnres_block": text["attn_res_block_size"],
"experts": text["num_experts"],
"active_experts": text["num_experts_per_token"],
"shared_experts": text["num_shared_experts"],
"latent_width": text["routed_expert_hidden_size"],
"expert_intermediate": text["moe_intermediate_size"],
"situ_beta": text["activation_situ_beta"],
"situ_linear_beta": text["activation_situ_linear_beta"],
"mla_nope": text["mla_use_nope"],
"mla_output_gate": text["mla_use_output_gate"],
},
"tensor_examples": {
"mla_layer_4": [
tensor_fact(
mla_header,
"language_model.model.layers.3.self_attn.kv_a_proj_with_mqa.weight",
),
tensor_fact(
mla_header,
"language_model.model.layers.3.self_attn.kv_b_proj.weight",
),
tensor_fact(
mla_header,
"language_model.model.layers.3.self_attn.q_a_proj.weight",
),
tensor_fact(
mla_header,
"language_model.model.layers.3.self_attn.q_b_proj.weight",
),
tensor_fact(
mla_header,
"language_model.model.layers.3.self_attn.g_proj.weight",
),
],
"vision": [
tensor_fact(vision_header, "vision_tower.patch_embed.proj.weight"),
tensor_fact(vision_header, "vision_tower.patch_embed.pos_emb.weight"),
tensor_fact(vision_header, "vision_tower.encoder.blocks.0.wqkv.weight"),
tensor_fact(vision_header, "vision_tower.encoder.blocks.26.wqkv.weight"),
tensor_fact(vision_header, "vision_tower.encoder.final_layernorm.weight"),
],
"routed_expert_0": [
{"name": "w1.weight_packed", "dtype": "U8", "shape": [3072, 1792], "bytes": 5_505_024},
{"name": "w1.weight_scale", "dtype": "U8", "shape": [3072, 112], "bytes": 344_064},
{"name": "w2.weight_packed", "dtype": "U8", "shape": [3584, 1536], "bytes": 5_505_024},
{"name": "w2.weight_scale", "dtype": "U8", "shape": [3584, 96], "bytes": 344_064},
{"name": "w3.weight_packed", "dtype": "U8", "shape": [3072, 1792], "bytes": 5_505_024},
{"name": "w3.weight_scale", "dtype": "U8", "shape": [3072, 112], "bytes": 344_064},
],
},
"parameter_audit": {
"a_log_checkpoint_shape": [128],
"a_log_public_code_shape": [96],
"dt_bias_shape": [96, 128],
"beta_projection_shape": [96, 7168],
"status": "observed shape inconsistency; runtime meaning unresolved",
"a_log": summarize(a_log),
"a_rate_exp": summarize(torch.exp(a_log)),
"dt_bias": summarize(dt_bias),
"channelwise_hypothesis": {
"log_decay": summarize(channelwise_log_decay),
"one_step_retention": summarize(channelwise_retention),
"retention_after_64_steps": summarize(channelwise_retention.pow(64)),
},
"router_correction_bias": summarize(correction_bias),
"router_row_l2": summarize(router_norms),
"router_bias_norm_correlation": torch.corrcoef(
torch.stack([correction_bias, router_norms])
)[0, 1].item(),
},
"router_stress_probe": {
"seed": args.seed,
"synthetic_tokens": args.synthetic_tokens,
"hidden_rms": synthetic.square().mean().sqrt().item(),
"without_correction_bias": load_metrics(unbiased_load),
"with_correction_bias": load_metrics(biased_load),
"membership_overlap_mean": overlap.mean().item(),
"tokens_changed": int((overlap < 16).sum()),
"changed_fraction": (overlap < 16).float().mean().item(),
"mean_replacements_per_token": (16 - overlap).mean().item(),
},
"flashkda": {
"supported_architectures": ["90a", "100a", "103a", "120a"],
"requirements": {"cuda": ">=12.9", "pytorch": ">=2.4", "gpu": "SM90+"},
"official_benchmarks": {
"h20": parse_benchmark(args.flashkda_dir / "BENCHMARK_H20.md"),
"gb200": parse_benchmark(args.flashkda_dir / "BENCHMARK_GB200.md"),
},
"local_environment": {
"python": platform.python_version(),
"torch": torch.__version__,
"torch_cuda": torch.version.cuda,
"gpu": torch.cuda.get_device_name(0) if torch.cuda.is_available() else None,
"capability": list(torch.cuda.get_device_capability(0))
if torch.cuda.is_available()
else None,
"libc": list(platform.libc_ver()),
},
"local_build": {
"status": "blocked_before_kernel execution",
"attempt_1": "system g++ 15 exceeds CUDA 12.8 host compiler range",
"attempt_2": "temporary g++ 13 reaches nvcc, then CUDA 12.8 headers conflict with current glibc math declarations",
"interpretation": "GPU architecture is listed by the repository, but the local CUDA 12.8 stack is below the official CUDA 12.9 requirement",
},
},
}
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n")
print(json.dumps({"output": str(args.output), "bytes": args.output.stat().st_size}))
if __name__ == "__main__":
main()