""" Retriever500M - Decoder-only transformer built from scratch. Architecture (LLaMA-style): - vocab_size: 32,000 - d_model: 1,280 - n_layers: 23 - n_heads: 20 - d_ff: 3,456 (SwiGLU, 2/3 * 4 * d_model) - RoPE positional encoding - RMSNorm (no biases) - Tied input/output embeddings - Total parameters: ~497M """ import math from dataclasses import dataclass import torch import torch.nn as nn import torch.nn.functional as F @dataclass class ModelConfig: vocab_size: int = 32_000 d_model: int = 1_280 n_layers: int = 23 n_heads: int = 20 d_ff: int = 3_456 max_seq_len: int = 1_024 rope_theta: float = 10_000.0 rope_pct: float = 0.25 # fraction of d_model per head used for RoPE dropout: float = 0.0 tie_embeddings: bool = True def __post_init__(self): assert self.d_model % self.n_heads == 0 self.d_head = self.d_model // self.n_heads # 64 class RMSNorm(nn.Module): """RMSNorm with optional bias (no bias by default, LLaMA-style).""" def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.weight = nn.Parameter(torch.ones(dim)) self.eps = eps def forward(self, x: torch.Tensor) -> torch.Tensor: # Compute in float32 for stability, then cast back orig_dtype = x.dtype x = x.float() rms = x.pow(2).mean(dim=-1, keepdim=True) x = x * torch.rsqrt(rms + self.eps) x = x.to(orig_dtype) return x * self.weight def precompute_rope_frequencies( d_head: int, max_seq_len: int, theta: float = 10_000.0, device: torch.device | None = None, ) -> torch.Tensor: """Precompute RoPE frequency table. Returns tensor of shape (max_seq_len, d_head // 2) with complex frequencies (cos, sin interleaved is handled in apply_rope). """ inv_freq = 1.0 / (theta ** (torch.arange(0, d_head, 2, device=device).float() / d_head)) positions = torch.arange(max_seq_len, device=device).float() freqs = torch.outer(positions, inv_freq) # (seq, d_head//2) return freqs def apply_rope( x: torch.Tensor, freqs: torch.Tensor, ) -> torch.Tensor: """Apply rotary position embeddings to tensor x. x: (batch, n_heads, seq, d_head) freqs: (seq, d_head // 2) """ seq_len = x.shape[2] d_head = x.shape[-1] freqs = freqs[:seq_len] # (seq, d_head//2) cos = freqs.cos() sin = freqs.sin() # Interleave cos/sin to match the rotate_half pattern # x is split into two halves: x1 = x[..., :d//2], x2 = x[..., d//2:] x1 = x[..., : d_head // 2] x2 = x[..., d_head // 2 :] # Broadcast cos/sin: (1, 1, seq, d_head//2) cos = cos.unsqueeze(0).unsqueeze(0) sin = sin.unsqueeze(0).unsqueeze(0) rotated = torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1) return rotated class Attention(nn.Module): """Multi-head self-attention with RoPE, no biases, causal masking.""" def __init__(self, config: ModelConfig): super().__init__() self.n_heads = config.n_heads self.d_head = config.d_head self.d_model = config.d_model self.scale = 1.0 / math.sqrt(self.d_head) # Fused QKV projection self.qkv = nn.Linear(config.d_model, 3 * config.d_model, bias=False) self.o_proj = nn.Linear(config.d_model, config.d_model, bias=False) self.dropout = config.dropout def forward( self, x: torch.Tensor, rope_freqs: torch.Tensor, mask: torch.Tensor | None = None, ) -> torch.Tensor: B, T, C = x.shape qkv = self.qkv(x) # (B, T, 3*C) q, k, v = qkv.chunk(3, dim=-1) # Reshape to (B, n_heads, T, d_head) q = q.view(B, T, self.n_heads, self.d_head).transpose(1, 2) k = k.view(B, T, self.n_heads, self.d_head).transpose(1, 2) v = v.view(B, T, self.n_heads, self.d_head).transpose(1, 2) # Apply RoPE to Q and K q = apply_rope(q, rope_freqs) k = apply_rope(k, rope_freqs) # Use PyTorch's scaled_dot_product_attention (uses Flash Attention on CUDA) if mask is not None: # mask: (1, 1, T, T) additive mask attn_mask = mask else: attn_mask = None out = F.scaled_dot_product_attention( q, k, v, attn_mask=attn_mask, dropout_p=self.dropout if self.training else 0.0, is_causal=(mask is None), ) # (B, n_heads, T, d_head) -> (B, T, C) out = out.transpose(1, 2).contiguous().view(B, T, C) return self.o_proj(out) class SwiGLU(nn.Module): """SwiGLU feed-forward network: (xW_gate * SiLU(xW_up)) * W_down.""" def __init__(self, config: ModelConfig): super().__init__() self.w_gate = nn.Linear(config.d_model, config.d_ff, bias=False) self.w_up = nn.Linear(config.d_model, config.d_ff, bias=False) self.w_down = nn.Linear(config.d_ff, config.d_model, bias=False) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x)) class TransformerBlock(nn.Module): """One transformer decoder block: pre-norm attention + pre-norm FFN.""" def __init__(self, config: ModelConfig): super().__init__() self.norm1 = RMSNorm(config.d_model) self.attn = Attention(config) self.norm2 = RMSNorm(config.d_model) self.ffn = SwiGLU(config) def forward( self, x: torch.Tensor, rope_freqs: torch.Tensor, mask: torch.Tensor | None = None, ) -> torch.Tensor: x = x + self.attn(self.norm1(x), rope_freqs, mask) x = x + self.ffn(self.norm2(x)) return x class Retriever500M(nn.Module): """Full decoder-only transformer model.""" def __init__(self, config: ModelConfig): super().__init__() self.config = config # Token embedding (tied with output head) self.token_embedding = nn.Embedding(config.vocab_size, config.d_model) # Transformer blocks self.layers = nn.ModuleList([ TransformerBlock(config) for _ in range(config.n_layers) ]) # Final norm self.norm_f = RMSNorm(config.d_model) # Output projection (tied with embedding) if config.tie_embeddings: self.lm_head = None # use token_embedding weight else: self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False) # Precompute RoPE frequencies (registered as buffer, moved with .to()) freqs = precompute_rope_frequencies( config.d_head, config.max_seq_len, config.rope_theta, ) self.register_buffer("rope_freqs", freqs, persistent=False) # Causal mask buffer mask = torch.full( (1, 1, config.max_seq_len, config.max_seq_len), float("-inf"), ) mask = torch.triu(mask, diagonal=1) self.register_buffer("causal_mask", mask, persistent=False) # Initialize weights self.apply(self._init_weights) def _init_weights(self, module: nn.Module): if isinstance(module, nn.Linear): nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) def get_output_weight(self): """Return the weight matrix for the output projection.""" if self.config.tie_embeddings: return self.token_embedding.weight return self.lm_head.weight def forward( self, input_ids: torch.Tensor, targets: torch.Tensor | None = None, use_checkpoint: bool = False, ) -> dict: B, T = input_ids.shape # Token embeddings x = self.token_embedding(input_ids) # (B, T, d_model) # Get RoPE freqs and causal mask for current sequence length rope_freqs = self.rope_freqs[:T] mask = self.causal_mask[:, :, :T, :T] # Transformer blocks (with optional gradient checkpointing) for layer in self.layers: if use_checkpoint and self.training: # Gradient checkpointing: recompute activations during backward x = torch.utils.checkpoint.checkpoint( layer, x, rope_freqs, mask, use_reentrant=False, ) else: x = layer(x, rope_freqs, mask) x = self.norm_f(x) # Output logits logits = F.linear(x, self.get_output_weight()) # (B, T, vocab_size) loss = None if targets is not None: loss = F.cross_entropy( logits.view(-1, logits.size(-1)), targets.view(-1), ignore_index=-100, ) return {"logits": logits, "loss": loss} @torch.no_grad() def generate( self, input_ids: torch.Tensor, max_new_tokens: int = 128, temperature: float = 1.0, top_k: int | None = None, eos_token_id: int | None = None, ) -> torch.Tensor: """Simple autoregressive generation.""" self.eval() for _ in range(max_new_tokens): # Crop context if it exceeds max_seq_len idx_cond = input_ids if input_ids.size(1) <= self.config.max_seq_len else \ input_ids[:, -self.config.max_seq_len:] logits = self(idx_cond)["logits"] logits = logits[:, -1, :] / max(temperature, 1e-6) if top_k is not None: v, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < v[:, [-1]]] = float("-inf") probs = F.softmax(logits, dim=-1) next_token = torch.multinomial(probs, num_samples=1) input_ids = torch.cat([input_ids, next_token], dim=1) if eos_token_id is not None and next_token.item() == eos_token_id: break return input_ids def count_parameters(self) -> int: """Count total trainable parameters.""" return sum(p.numel() for p in self.parameters() if p.requires_grad) def build_model(config: ModelConfig | None = None) -> Retriever500M: """Build the Retriever500M model.""" if config is None: config = ModelConfig() model = Retriever500M(config) return model if __name__ == "__main__": config = ModelConfig() model = build_model(config) total_params = model.count_parameters() print(f"Model: Retriever500M") print(f" d_model: {config.d_model}") print(f" n_layers: {config.n_layers}") print(f" n_heads: {config.n_heads}") print(f" d_ff: {config.d_ff}") print(f" d_head: {config.d_head}") print(f" vocab_size: {config.vocab_size}") print(f" max_seq_len: {config.max_seq_len}") print(f" Total parameters: {total_params:,} ({total_params / 1e6:.1f}M)") # Quick forward pass test device = "cuda" if torch.cuda.is_available() else "cpu" model = model.to(device) model.eval() input_ids = torch.randint(0, config.vocab_size, (2, 64), device=device) with torch.no_grad(): out = model(input_ids) print(f" Output logits shape: {out['logits'].shape}") print(" Forward pass OK.")