File size: 4,268 Bytes
3fd1a35
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
#!/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())