File size: 3,921 Bytes
e3d43ae
5434506
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e3d43ae
5434506
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
"""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)))}