Download model.py from Ne30Charm/tiny-ced: direct link, hf CLI and curl.
- Browser
- Download file 3.96 kB
-
https://huggingface.co/Ne30Charm/tiny-ced/resolve/main/model.py
- Command line
-
hf download hf://Ne30Charm/tiny-ced/model.py
-
curl -L -o model.py https://huggingface.co/Ne30Charm/tiny-ced/resolve/main/model.py
3.96 kB
| """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)) | |