Ling-3.0-tiny-RKNN / tools /reference_layer0.py
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
4.27 kB
#!/usr/bin/env python3
"""Export a deterministic first-token layer-0 reference without FLA/CUDA."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import torch
from safetensors import safe_open
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:
dtype = value.dtype
normalized = value.float() * torch.rsqrt(
value.float().square().mean(dim=-1, keepdim=True) + 1.0e-6
)
return weight * normalized.to(dtype)
def linear(value: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
return torch.mv(weight.float(), value.float()).to(torch.bfloat16)
def conv_first(value: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
result = value.float() * weight[:, 0, -1].float()
return torch.nn.functional.silu(result).to(torch.bfloat16)
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)
args = parser.parse_args()
source = Source(args.source.resolve())
prefix = "model.layers.0"
attention = prefix + ".attention"
mlp = prefix + ".mlp"
hidden = source.get("model.word_embeddings.weight")[args.token].contiguous()
normalized = rms_norm(hidden, source.get(prefix + ".input_layernorm.weight"))
q = conv_first(linear(normalized, source.get(attention + ".q_proj.weight")), source.get(attention + ".q_conv1d.weight"))
k = conv_first(linear(normalized, source.get(attention + ".k_proj.weight")), source.get(attention + ".k_conv1d.weight"))
v = conv_first(linear(normalized, source.get(attention + ".v_proj.weight")), source.get(attention + ".v_conv1d.weight"))
f = linear(normalized, source.get(attention + ".f_proj.weight"))
gate = linear(normalized, source.get(attention + ".g_proj.weight"))
beta = torch.sigmoid(linear(normalized, source.get(attention + ".b_proj.weight")).float())
q = q.view(16, 128).float()
k = k.view(16, 128).float()
v = v.view(16, 128).float()
q = q * torch.rsqrt(q.square().sum(-1, keepdim=True) + 1.0e-6)
k = k * torch.rsqrt(k.square().sum(-1, keepdim=True) + 1.0e-6)
# With a zero recurrent state, decay does not affect the first KDA output.
_decay = -5.0 * torch.sigmoid(
torch.exp(source.get(attention + ".A_log")).view(16, 1) *
(f.float().view(16, 128) + source.get(attention + ".dt_bias").float().view(16, 128))
)
recurrence = beta.view(16, 1) * v * (q * k).sum(-1, keepdim=True) / (128.0**0.5)
recurrence = recurrence.to(torch.bfloat16)
normalized_output = rms_norm(
recurrence,
source.get(attention + ".o_norm.weight"),
)
attention_vector = normalized_output * torch.sigmoid(gate.view(16, 128).float()).to(torch.bfloat16)
attention_result = linear(attention_vector.flatten(), source.get(attention + ".o_proj.weight"))
hidden = hidden + attention_result
normalized = rms_norm(hidden, source.get(prefix + ".post_attention_layernorm.weight"))
gate_value = linear(normalized, source.get(mlp + ".gate_proj.weight"))
up_value = linear(normalized, source.get(mlp + ".up_proj.weight"))
ffn = torch.nn.functional.silu(gate_value.float()).to(torch.bfloat16) * up_value
result = hidden + linear(ffn, source.get(mlp + ".down_proj.weight"))
args.output.parent.mkdir(parents=True, exist_ok=True)
result.float().numpy().tofile(args.output)
print(f"token={args.token} output={args.output} elements={result.numel()}")
print(f"mean={result.float().mean().item():.9f} rms={result.float().square().mean().sqrt().item():.9f}")
return 0
if __name__ == "__main__":
raise SystemExit(main())