Download tools/reference_full_token.py from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 6.96 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/reference_full_token.py
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/tools/reference_full_token.py
-
curl -L -o reference_full_token.py https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/reference_full_token.py
6.96 kB
| #!/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()) | |