#!/usr/bin/env python3 """Export official-BF16 logits for one isolated Ling-3.0-tiny token.""" from __future__ import annotations import argparse import json from pathlib import Path import torch from safetensors import safe_open HIDDEN = 1536 HEADS = 16 HEAD_DIM = 128 EXPERTS = 128 EXPERTS_PER_TOKEN = 8 GROUPS = 8 SELECTED_GROUPS = 4 ROUTED_SCALE = 2.5 class Source: def __init__(self, root: Path): index = json.loads((root / "model.safetensors.index.json").read_text()) self.root = root self.weight_map = index["weight_map"] self.handles = { shard: safe_open(root / shard, framework="pt", device="cpu") for shard in sorted(set(self.weight_map.values())) } def get(self, name: str) -> torch.Tensor: return self.handles[self.weight_map[name]].get_tensor(name) def rms_norm(value: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: normalized = value.float() * torch.rsqrt(value.float().square().mean() + 1.0e-6) return (weight * normalized.to(value.dtype)).to(torch.bfloat16) def linear(value: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: return torch.mv(weight.float(), value.float()).to(torch.bfloat16) def mlp(source: Source, prefix: str, value: torch.Tensor) -> torch.Tensor: gate = linear(value, source.get(prefix + ".gate_proj.weight")) up = linear(value, source.get(prefix + ".up_proj.weight")) activated = torch.nn.functional.silu(gate.float()).to(torch.bfloat16) * up return linear(activated, source.get(prefix + ".down_proj.weight")) def select_route(logits: torch.Tensor, bias: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: scores = torch.sigmoid(logits.float()) routing_scores = scores + bias.float() group_scores = routing_scores.view(GROUPS, -1).topk(2, dim=-1).values.sum(dim=-1) selected_groups = group_scores.topk(SELECTED_GROUPS, sorted=False).indices mask = torch.zeros(GROUPS, dtype=torch.bool) mask[selected_groups] = True masked = routing_scores.masked_fill(~mask[:, None].expand(GROUPS, EXPERTS // GROUPS).reshape(-1), -torch.inf) selected = masked.topk(EXPERTS_PER_TOKEN, sorted=False).indices weights = scores[selected] weights = weights / weights.sum() * ROUTED_SCALE return selected, weights def first_kda(source: Source, prefix: str, value: torch.Tensor) -> torch.Tensor: def conv(name: str) -> torch.Tensor: projected = linear(value, source.get(prefix + f".{name}_proj.weight")) weight = source.get(prefix + f".{name}_conv1d.weight")[:, 0, -1] return torch.nn.functional.silu(projected.float() * weight.float()).to(torch.bfloat16) query = conv("q").view(HEADS, HEAD_DIM).float() key = conv("k").view(HEADS, HEAD_DIM).float() projected_value = conv("v").view(HEADS, HEAD_DIM).float() query *= torch.rsqrt(query.square().sum(-1, keepdim=True) + 1.0e-6) key *= torch.rsqrt(key.square().sum(-1, keepdim=True) + 1.0e-6) beta = torch.sigmoid(linear(value, source.get(prefix + ".b_proj.weight")).float()) recurrence = beta[:, None] * projected_value * (query * key).sum(-1, keepdim=True) / (HEAD_DIM**0.5) recurrence = recurrence.to(torch.bfloat16) normalized = rms_norm(recurrence, source.get(prefix + ".o_norm.weight")) gate = torch.sigmoid(linear(value, source.get(prefix + ".g_proj.weight")).float()).to(torch.bfloat16) return linear((normalized * gate.view(HEADS, HEAD_DIM)).flatten(), source.get(prefix + ".o_proj.weight")) def first_mla(source: Source, prefix: str, value: torch.Tensor) -> torch.Tensor: compressed = linear(value, source.get(prefix + ".kv_a_proj_with_mqa.weight")) latent = rms_norm(compressed[:512], source.get(prefix + ".kv_a_layernorm.weight")) key_value = linear(latent, source.get(prefix + ".kv_b_proj.weight")).view(HEADS, 256) attention = key_value[:, 128:].contiguous() gate = torch.sigmoid(linear(value, source.get(prefix + ".g_proj.weight")).float()).to(torch.bfloat16) return linear((attention * gate[:, None]).flatten(), source.get(prefix + ".dense.weight")) def main() -> int: parser = argparse.ArgumentParser() parser.add_argument("--source", type=Path, required=True) parser.add_argument("--token", type=int, default=34355) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--routes", type=Path) args = parser.parse_args() source = Source(args.source.resolve()) hidden = source.get("model.word_embeddings.weight")[args.token].contiguous() routes: list[dict[str, object]] = [] with torch.inference_mode(): for layer in range(24): root = f"model.layers.{layer}" normalized = rms_norm(hidden, source.get(root + ".input_layernorm.weight")) attention = first_mla(source, root + ".attention", normalized) if (layer + 1) % 4 == 0 else first_kda(source, root + ".attention", normalized) hidden = (hidden + attention).to(torch.bfloat16) normalized = rms_norm(hidden, source.get(root + ".post_attention_layernorm.weight")) if layer == 0: feed_forward = mlp(source, root + ".mlp", normalized) else: gate_prefix = root + ".mlp.gate" router_logits = torch.mv(source.get(gate_prefix + ".weight").float(), normalized.float()) selected, weights = select_route(router_logits, source.get(gate_prefix + ".expert_bias")) routed = torch.zeros(HIDDEN, dtype=torch.float32) for expert, weight in zip(selected.tolist(), weights.tolist()): routed += mlp(source, root + f".mlp.experts.{expert}", normalized).float() * weight routed = routed.to(torch.bfloat16) shared = mlp(source, root + ".mlp.shared_experts", normalized) feed_forward = (routed + shared).to(torch.bfloat16) routes.append({ "layer": layer, "experts": selected.tolist(), "weights": weights.tolist(), }) hidden = (hidden + feed_forward).to(torch.bfloat16) print(f"layer={layer} rms={hidden.float().square().mean().sqrt().item():.9f}", flush=True) normalized = rms_norm(hidden, source.get("model.norm.weight")) logits = linear(normalized, source.get("lm_head.weight")).float() args.output.parent.mkdir(parents=True, exist_ok=True) logits.numpy().tofile(args.output) if args.routes is not None: args.routes.write_text(json.dumps(routes, indent=2) + "\n") top = logits.topk(10) print(f"token={args.token} output={args.output} elements={logits.numel()}") print("top_ids=" + ",".join(str(value) for value in top.indices.tolist())) print("top_logits=" + ",".join(f"{value:.8f}" for value in top.values.tolist())) return 0 if __name__ == "__main__": raise SystemExit(main())