CoolFace
Modelpublic

Premchan369/Q-TensorFormer

sourceHugging Faceapache-2.0updated 10d agoView on Hugging Face
2likes185downloads
baselines.py234 linesDownload Raw Back to src
1"""2Baseline implementations for fair comparison.3 4Baselines:5  1. Standard Transformer: Dense MLP FFN, no TT, no quantum.6  2. Distilled: Smaller transformer trained with KD.7  3. Pruned: Magnitude-based structured pruning.8  4. TT-Only: Tensor network FFN without quantum or adaptive rank.9"""10 11import torch12import torch.nn as nn13import torch.nn.functional as F14import math15from typing import Optional16 17 18class StandardTransformer(nn.Module):19    """20    Basic transformer decoder (GPT-style) with dense MLP FFN.21 22    Reference baseline — matches Q-TensorFormer architecture23    exactly except for TT decomposition and quantum layers.24    """25 26    def __init__(self, vocab_size: int = 10000, d_model: int = 128,27                 n_heads: int = 4, n_layers: int = 2, ff_mult: int = 4,28                 max_seq_len: int = 128, dropout: float = 0.1):29        super().__init__()30        self.d_model = d_model31        self.config = type("config", (), {32            "d_model": d_model, "n_heads": n_heads, "n_layers": n_layers,33            "ff_multiplier": ff_mult, "max_seq_len": max_seq_len,34            "vocab_size": vocab_size, "dropout": dropout,35        })()36 37        self.embedding = nn.Embedding(vocab_size, d_model)38        self.pos_encoding = _PositionalEncoding(d_model, max_seq_len, dropout)39 40        self.blocks = nn.ModuleList([41            _StandardBlock(d_model, n_heads, ff_mult, dropout, max_seq_len)42            for _ in range(n_layers)43        ])44 45        self.ln_f = nn.LayerNorm(d_model)46        self.lm_head = nn.Linear(d_model, vocab_size, bias=False)47        self.lm_head.weight = self.embedding.weight48 49    def forward(self, input_ids, attention_mask=None, return_stats=False):50        x = self.embedding(input_ids)51        x = self.pos_encoding(x)52 53        for block in self.blocks:54            x = block(x, mask=attention_mask)55 56        x = self.ln_f(x)57        logits = self.lm_head(x)58 59        if return_stats:60            return logits, []61        return logits62 63    @property64    def total_params(self) -> int:65        return sum(p.numel() for p in self.parameters())66 67 68class DistilledTransformer(nn.Module):69    """70    Smaller transformer trained via knowledge distillation.71 72    Designed to match Q-TensorFormer parameter counts.73    """74 75    def __init__(self, vocab_size: int = 10000, d_model: int = 96,76                 n_heads: int = 4, n_layers: int = 2, ff_mult: int = 3,77                 max_seq_len: int = 128, dropout: float = 0.1):78        super().__init__()79        self.d_model = d_model80        self.config = type("config", (), {81            "d_model": d_model, "n_heads": n_heads, "n_layers": n_layers,82            "ff_multiplier": ff_mult, "max_seq_len": max_seq_len,83            "vocab_size": vocab_size, "dropout": dropout,84        })()85 86        self.embedding = nn.Embedding(vocab_size, d_model)87        self.pos_encoding = _PositionalEncoding(d_model, max_seq_len, dropout)88 89        self.blocks = nn.ModuleList([90            _StandardBlock(d_model, n_heads, ff_mult, dropout, max_seq_len)91            for _ in range(n_layers)92        ])93 94        self.ln_f = nn.LayerNorm(d_model)95        self.lm_head = nn.Linear(d_model, vocab_size, bias=False)96        self.lm_head.weight = self.embedding.weight97 98    def forward(self, input_ids, attention_mask=None, return_stats=False):99        x = self.embedding(input_ids)100        x = self.pos_encoding(x)101 102        for block in self.blocks:103            x = block(x, mask=attention_mask)104 105        x = self.ln_f(x)106        logits = self.lm_head(x)107 108        if return_stats:109            return logits, []110        return logits111 112    @property113    def total_params(self) -> int:114        return sum(p.numel() for p in self.parameters())115 116 117class PrunedTransformer(nn.Module):118    """119    Magnitude-pruned standard transformer.120 121    Prunes FFN weights globally to match Q-TensorFormer parameter count.122    Applies structured pruning (zeroing channels) for efficiency.123    """124 125    def __init__(self, base_model: StandardTransformer,126                 prune_ratio: float = 0.5):127        super().__init__()128        self.base = base_model129        self.prune_ratio = prune_ratio130        self.config = base_model.config131        self._prune()132 133    def _prune(self):134        """Apply structured magnitude pruning to FFN layers."""135        all_weights = []136        for block in self.base.blocks:137            for weight in [block.ffn[0].weight, block.ffn[2].weight]:138                all_weights.append(weight.flatten())139 140        # Compute global threshold141        flat = torch.cat(all_weights)142        k = int(len(flat) * self.prune_ratio)143        threshold = torch.topk(flat.abs(), k, largest=False).values[-1]144 145        # Apply structured pruning (zero rows/cols)146        for block in self.base.blocks:147            for layer in [block.ffn[0], block.ffn[2]]:148                mask = (layer.weight.abs() > threshold).float()149                # Zero small rows entirely150                row_norms = mask.sum(dim=1)151                dead_rows = row_norms < layer.weight.size(1) * 0.1152                mask[dead_rows] = 0153                layer.weight.data *= mask154 155    def forward(self, *args, **kwargs):156        return self.base(*args, **kwargs)157 158    @property159    def total_params(self) -> int:160        return sum(p.numel() for p in self.parameters())161 162 163class _StandardBlock(nn.Module):164    """Standard transformer decoder block."""165 166    def __init__(self, d_model, n_heads, ff_mult, dropout, max_seq_len):167        super().__init__()168        self.ln1 = nn.LayerNorm(d_model)169        self.attn = _CausalAttention(d_model, n_heads, dropout, max_seq_len)170        self.ln2 = nn.LayerNorm(d_model)171        self.ffn = nn.Sequential(172            nn.Linear(d_model, d_model * ff_mult),173            nn.GELU(),174            nn.Linear(d_model * ff_mult, d_model),175            nn.Dropout(dropout),176        )177        self.dropout = nn.Dropout(dropout)178 179    def forward(self, x, mask=None):180        x = x + self.dropout(self.attn(self.ln1(x), mask=mask))181        x = x + self.ffn(self.ln2(x))182        return x183 184 185class _CausalAttention(nn.Module):186    """Causal multi-head attention."""187 188    def __init__(self, d_model, n_heads, dropout, max_seq_len):189        super().__init__()190        assert d_model % n_heads == 0191        self.n_heads = n_heads192        self.head_dim = d_model // n_heads193        self.scale = math.sqrt(self.head_dim)194 195        self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)196        self.out_proj = nn.Linear(d_model, d_model, bias=False)197        self.dropout = nn.Dropout(dropout)198 199        self.max_seq_len = max_seq_len200 201    def forward(self, x, mask=None):202        B, T, C = x.shape203        qkv = self.qkv(x).reshape(B, T, 3, self.n_heads, self.head_dim)204        q, k, v = qkv.unbind(dim=2)205        q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)206 207        attn = (q @ k.transpose(-2, -1)) / self.scale208        causal = torch.triu(torch.ones(T, T, device=x.device) * float("-inf"), diagonal=1)209        attn = attn + causal210 211        if mask is not None:212            attn = attn + mask.unsqueeze(1).unsqueeze(2) * float("-inf")213 214        attn = F.softmax(attn, dim=-1)215        attn = self.dropout(attn)216 217        out = (attn @ v).transpose(1, 2).reshape(B, T, C)218        return self.out_proj(out)219 220 221class _PositionalEncoding(nn.Module):222    def __init__(self, d_model, max_len, dropout):223        super().__init__()224        self.dropout = nn.Dropout(dropout)225        pe = torch.zeros(max_len, d_model)226        pos = torch.arange(max_len).unsqueeze(1).float()227        div = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))228        pe[:, 0::2] = torch.sin(pos * div)229        pe[:, 1::2] = torch.cos(pos * div)230        self.register_buffer("pe", pe.unsqueeze(0))231 232    def forward(self, x):233        return self.dropout(x + self.pe[:, :x.size(1)])234