File size: 5,002 Bytes
d73d9e7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
"""TinyChat-5M: 5.1M param transformer. d=192, 6L, 3H, head_dim=64, SwiGLU 4x, RoPE, RMSNorm, tied embeddings."""
import torch
import torch.nn as nn
import torch.nn.functional as F


class RMSNorm(nn.Module):
    def __init__(self, d, eps=1e-6):
        super().__init__()
        self.weight = nn.Parameter(torch.ones(d))
        self.eps = eps

    def forward(self, x):
        rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
        return x * rms * self.weight


def rotate_half(x):
    x1, x2 = x.chunk(2, dim=-1)
    return torch.cat([-x2, x1], dim=-1)


class Attention(nn.Module):
    def __init__(self, d, n_heads, head_dim):
        super().__init__()
        self.n_heads = n_heads
        self.head_dim = head_dim
        self.wq = nn.Linear(d, n_heads * head_dim, bias=False)
        self.wk = nn.Linear(d, n_heads * head_dim, bias=False)
        self.wv = nn.Linear(d, n_heads * head_dim, bias=False)
        self.wo = nn.Linear(n_heads * head_dim, d, bias=False)

    def forward(self, x, cos, sin):
        B, S, _ = x.shape
        q = self.wq(x).view(B, S, self.n_heads, self.head_dim).transpose(1, 2)
        k = self.wk(x).view(B, S, self.n_heads, self.head_dim).transpose(1, 2)
        v = self.wv(x).view(B, S, self.n_heads, self.head_dim).transpose(1, 2)
        cos_f = cos[:S].unsqueeze(0).unsqueeze(0)
        sin_f = sin[:S].unsqueeze(0).unsqueeze(0)
        q = (q * cos_f) + (rotate_half(q) * sin_f)
        k = (k * cos_f) + (rotate_half(k) * sin_f)
        scale = self.head_dim ** -0.5
        scores = (q @ k.transpose(-2, -1)) * scale
        mask = torch.triu(torch.ones(S, S, device=x.device, dtype=torch.bool), diagonal=1)
        scores = scores.masked_fill(mask, float('-inf'))
        attn = F.softmax(scores, dim=-1)
        out = (attn @ v).transpose(1, 2).reshape(B, S, -1)
        return self.wo(out)


class FFN(nn.Module):
    def __init__(self, d, mult):
        super().__init__()
        hidden = d * mult
        self.w1 = nn.Linear(d, hidden, bias=False)
        self.w2 = nn.Linear(hidden, d, bias=False)
        self.w3 = nn.Linear(d, hidden, bias=False)

    def forward(self, x):
        return self.w2(F.silu(self.w1(x)) * self.w3(x))


class Block(nn.Module):
    def __init__(self, d, n_heads, head_dim, ffn_mult):
        super().__init__()
        self.norm1 = RMSNorm(d)
        self.attn = Attention(d, n_heads, head_dim)
        self.norm2 = RMSNorm(d)
        self.ffn = FFN(d, ffn_mult)

    def forward(self, x, cos, sin):
        x = x + self.attn(self.norm1(x), cos, sin)
        x = x + self.ffn(self.norm2(x))
        return x


class TinyLM(nn.Module):
    def __init__(self, vocab=4096, d=192, n_layers=6, n_heads=3, head_dim=64, ffn_mult=4):
        super().__init__()
        self.tok_emb = nn.Embedding(vocab, d)
        self.layers = nn.ModuleList([Block(d, n_heads, head_dim, ffn_mult) for _ in range(n_layers)])
        self.norm_f = RMSNorm(d)
        self.head = nn.Linear(d, vocab, bias=False)
        # Tied embeddings
        self.head.weight = self.tok_emb.weight
        # RoPE
        freqs = 1.0 / (10000 ** (torch.arange(0, head_dim, 2).float() / head_dim))
        t = torch.arange(2048).float()
        outer = torch.outer(t, freqs)
        angles = torch.cat([outer, outer], dim=-1)
        self.register_buffer("cos", angles.cos())
        self.register_buffer("sin", angles.sin())

    def forward(self, x, y=None):
        h = self.tok_emb(x)
        for layer in self.layers:
            h = layer(h, self.cos, self.sin)
        h = self.norm_f(h)
        logits = self.head(h)
        if y is not None:
            return F.cross_entropy(logits.reshape(-1, logits.size(-1)), y.reshape(-1))
        return logits


def load_model(path="model.pt", device="cpu"):
    model = TinyLM()
    ckpt = torch.load(path, map_location=device, weights_only=True)
    model.load_state_dict(ckpt["model"])
    model.eval()
    return model


def generate(model, tokenizer, prompt, max_new=100, temperature=0.8, rep_penalty=1.5):
    ids = tokenizer.encode(prompt, add_special_tokens=False).ids
    ids = torch.tensor([ids])
    with torch.no_grad():
        for _ in range(max_new):
            inp = ids[:, -512:]
            logits = model(inp)[0, -1]
            for tid in ids[0].unique():
                if logits[tid] > 0:
                    logits[tid] /= rep_penalty
                else:
                    logits[tid] *= rep_penalty
            logits = logits / temperature
            probs = F.softmax(logits, dim=-1)
            next_id = torch.multinomial(probs, 1).item()
            ids = torch.cat([ids, torch.tensor([[next_id]])], dim=1)
    return tokenizer.decode(ids[0].tolist(), skip_special_tokens=True)


if __name__ == "__main__":
    from tokenizers import Tokenizer
    tok = Tokenizer.from_file("tokenizer.json")
    model = load_model()
    print(f"Params: {sum(p.numel() for p in model.parameters()):,}")
    print(generate(model, tok, "[INST] Hello! [/INST]", max_new=100))