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