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