Premchan369/Q-TensorFormer
2185
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 