Ling-3.0-tiny-RKNN / tools /reference_full_token.py
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
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())