#!/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())