Glyph / model.py
ThingsAI's picture
Upload Glyph multi-task model (epoch 3, avg acc 88.0%)
e3d43ae verified
Raw
History Blame Contribute Delete
3.92 kB
"""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)))}