File size: 3,960 Bytes
d5b9248
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
830c033
d5b9248
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Tiny CED model architecture required to load model.safetensors."""

import torch
import torch.nn.functional as F
from torch import nn


class RMSNorm(nn.Module):
    def __init__(self, dim):
        super().__init__()
        self.weight = nn.Parameter(torch.ones(dim))

    def forward(self, x):
        z = x.float()
        return (
            z
            * torch.rsqrt(z.square().mean(-1, keepdim=True) + 1e-5)
            * self.weight
        ).to(x.dtype)


def rope(x, pos):
    d = x.shape[-1]
    freq = 10000 ** (-torch.arange(0, d, 2, device=x.device) / d)
    angle = pos[:, None, :, None].float() * freq
    a, b = x.float()[..., ::2], x.float()[..., 1::2]
    c, s = angle.cos(), angle.sin()
    return torch.stack((a * c - b * s, a * s + b * c), dim=-1).flatten(-2).to(x.dtype)


class Attention(nn.Module):
    def __init__(self, dim, heads, window):
        super().__init__()
        self.heads = heads
        self.window = window
        for name in ["wq", "wkg", "wvg", "wkl", "wvl", "wo"]:
            setattr(self, name, nn.Linear(dim, dim, bias=False))

    def forward(self, u, source, pos, valid):
        batch, length, dim = u.shape

        def split(z):
            return z.view(batch, length, self.heads, dim // self.heads).transpose(1, 2)

        q = rope(split(self.wq(u)), pos)
        kg, vg = rope(split(self.wkg(source)), pos), split(self.wvg(source))
        kl, vl = rope(split(self.wkl(u)), pos), split(self.wvl(u))
        distance = pos[:, :, None] - pos[:, None, :]
        global_mask = (distance >= 0) & valid[:, :, None] & valid[:, None, :]
        local_mask = global_mask & (distance < self.window)
        mask = torch.cat((global_mask, local_mask), -1)[:, None]
        k, v = torch.cat((kg, kl), 2), torch.cat((vg, vl), 2)
        out = F.scaled_dot_product_attention(
            q,
            k,
            v,
            attn_mask=mask,
            dropout_p=0,
            is_causal=False,
        )
        return self.wo(out.transpose(1, 2).contiguous().view(batch, length, dim))


class Block(nn.Module):
    def __init__(self, dim, heads, ff, window):
        super().__init__()
        self.attn_norm, self.ffn_norm = RMSNorm(dim), RMSNorm(dim)
        self.attn = Attention(dim, heads, window)
        self.gate, self.up = (nn.Linear(dim, ff, bias=False) for _ in range(2))
        self.down = nn.Linear(ff, dim, bias=False)

    def forward(self, h, e, pos, valid):
        u = self.attn_norm(h)
        h = h + self.attn(u, u if e is None else e, pos, valid)
        z = self.ffn_norm(h)
        return h + self.down(F.silu(self.gate(z)) * self.up(z))


class CED(nn.Module):
    def __init__(self, vocab=8192, dim=384, heads=6, ff=1024, layers=4, window=64):
        super().__init__()
        assert dim % heads == 0 and (dim // heads) % 2 == 0
        self.embedding = nn.Embedding(vocab, dim)
        self.encoder = nn.ModuleList(
            [Block(dim, heads, ff, window) for _ in range(layers)]
        )
        self.decoder = nn.ModuleList(
            [Block(dim, heads, ff, window) for _ in range(layers)]
        )
        self.encoder_norm, self.final_norm = RMSNorm(dim), RMSNorm(dim)
        self.head = nn.Linear(dim, vocab, bias=False)
        self.head.weight = self.embedding.weight
        for parameter in self.parameters():
            if parameter.ndim > 1:
                nn.init.normal_(parameter, std=0.02)

    def forward(self, ids, positions=None, valid=None):
        if positions is None:
            positions = torch.arange(ids.shape[1], device=ids.device)[None].expand_as(ids)
        if valid is None:
            valid = torch.ones_like(ids, dtype=torch.bool)
        h = self.embedding(ids)
        for layer in self.encoder:
            h = layer(h, None, positions, valid)
        e = self.encoder_norm(h)
        h = e
        for layer in self.decoder:
            h = layer(h, e, positions, valid)
        return self.head(self.final_norm(h))