| """Glyph model definition — standalone, no dependencies beyond PyTorch.""" |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
|
|
| class TrigramHashEmbedding(nn.Module): |
| def __init__(self, n_buckets=8192, d_embed=64, prime=31): |
| super().__init__() |
| self.n_buckets, self.prime = n_buckets, prime |
| self.embed = nn.Embedding(n_buckets, d_embed) |
| def forward(self, x): |
| xp = F.pad(x.long(), (2, 0), value=0) |
| h = (xp[:, :-2] * self.prime * self.prime + xp[:, 1:-1] * self.prime + xp[:, 2:]) % self.n_buckets |
| return self.embed(h) |
|
|
|
|
| class RoPE(nn.Module): |
| def __init__(self, head_dim, max_len=1024, theta=10000.0): |
| super().__init__() |
| inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2).float() / head_dim)) |
| self.register_buffer("inv_freq", inv_freq, persistent=False) |
| self._build(max_len) |
| def _build(self, seq_len): |
| t = torch.arange(seq_len, device=self.inv_freq.device).float() |
| freqs = torch.outer(t, self.inv_freq); emb = torch.cat([freqs, freqs], dim=-1) |
| self.register_buffer("cos", emb.cos()[None, None], persistent=False) |
| self.register_buffer("sin", emb.sin()[None, None], persistent=False); self._max = seq_len |
| @staticmethod |
| def _rotate(x): |
| x1, x2 = x.chunk(2, dim=-1); return torch.cat([-x2, x1], dim=-1) |
| def forward(self, q, k): |
| T = q.size(2) |
| if T > self._max: self._build(T) |
| c, s = self.cos[:,:,:T], self.sin[:,:,:T] |
| return q*c + self._rotate(q)*s, k*c + self._rotate(k)*s |
|
|
|
|
| class Attention(nn.Module): |
| def __init__(self, d_model, n_heads): |
| super().__init__() |
| self.n_heads, self.head_dim = n_heads, d_model // n_heads |
| self.qkv = nn.Linear(d_model, 3 * d_model); self.out = nn.Linear(d_model, d_model) |
| self.rope = RoPE(self.head_dim); self.norm = nn.LayerNorm(d_model) |
| def forward(self, x): |
| res = x; x = self.norm(x); B, T, C = x.shape |
| qkv = self.qkv(x).view(B, T, 3, self.n_heads, self.head_dim) |
| q, k, v = qkv.permute(2, 0, 3, 1, 4); q, k = self.rope(q, k) |
| out = F.scaled_dot_product_attention(q, k, v, is_causal=False) |
| return res + self.out(out.transpose(1, 2).contiguous().view(B, T, C)) |
|
|
|
|
| class ConvBlock(nn.Module): |
| def __init__(self, d_in, d_out, kernel=3): |
| super().__init__() |
| self.conv = nn.Conv1d(d_in, d_out, kernel, padding=kernel // 2) |
| self.bn = nn.BatchNorm1d(d_out) |
| self.residual = nn.Conv1d(d_in, d_out, 1) if d_in != d_out else nn.Identity() |
| def forward(self, x): return F.gelu(self.bn(self.conv(x))) + self.residual(x) |
|
|
|
|
| class MultiTaskLID(nn.Module): |
| """ |
| Glyph: Multi-task byte-level text classifier. |
| ~4M shared parameters + per-task classification heads. |
| """ |
| def __init__(self, task_configs, max_len=512, d_byte=64, d_tri=64, |
| n_buckets=8192, d_model=384, n_conv=4, n_attn=2, n_heads=6, dropout=0.0): |
| super().__init__() |
| self.max_len = max_len; self.d_model = d_model |
| self.byte_embed = nn.Embedding(256, d_byte) |
| self.tri_embed = TrigramHashEmbedding(n_buckets, d_tri) |
| self.proj = nn.Linear(d_byte + d_tri, d_model) |
| self.convs = nn.ModuleList([ConvBlock(d_model, d_model, 3) for _ in range(n_conv)]) |
| self.attns = nn.ModuleList([Attention(d_model, n_heads) for _ in range(n_attn)]) |
| self.drop = nn.Dropout(dropout); self.norm = nn.LayerNorm(d_model) |
| self.heads = nn.ModuleDict({t: nn.Linear(d_model, n) for t, n in task_configs.items()}) |
|
|
| def forward(self, x, task): |
| h = self.proj(torch.cat([self.byte_embed(x), self.tri_embed(x)], dim=-1)) |
| h = h.transpose(1, 2) |
| for conv in self.convs: h = conv(h) |
| h = h.transpose(1, 2) |
| for attn in self.attns: h = attn(h) |
| return {"logits": self.heads[task](self.drop(self.norm(h).mean(dim=1)))} |
|
|