"""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))