""" Basic Multi-headed Latent Attention (MLA). Simple implementation without KV cache. """ import torch import torch.nn as nn import torch.nn.functional as F import math from .rope import RotaryEmbedding, apply_rotary class MemoryOptimizedMLA(nn.Module): """ Basic MLA: Project to latent space, apply multi-head attention, project back. Numerically stable implementation with proper normalization. """ def __init__(self, config): super().__init__() self.config = config self.n_heads = config.n_heads self.d_head = config.d_kv_comp // config.n_heads self.d_rope = config.d_rope # Improved scaling: use sqrt(d_head) with a small epsilon for numerical stability self.scale = 1.0 / math.sqrt(max(self.d_head, 1.0)) # Layer normalization before projections for stability self.norm_latent = nn.LayerNorm(config.d_model) # Projections self.to_latent = nn.Linear(config.d_model, config.d_kv_comp, bias=False) # Q/K/V from latent self.q_proj = nn.Linear(config.d_kv_comp, config.d_kv_comp, bias=False) self.k_proj = nn.Linear(config.d_kv_comp, config.d_kv_comp, bias=False) self.v_proj = nn.Linear(config.d_kv_comp, config.d_kv_comp, bias=False) # RoPE self.rotary = RotaryEmbedding(config.d_rope) # Output self.out_proj = nn.Linear(config.d_kv_comp, config.d_model, bias=False) self.attn_dropout = nn.Dropout(config.dropout) self.resid_dropout = nn.Dropout(config.dropout) def forward(self, x, mask=None): """ Args: x: (batch_size, seq_len, d_model) mask: (batch_size, seq_len) or (batch_size, 1, seq_len, seq_len), optional Returns: out: (batch_size, seq_len, d_model) """ batch_size, seq_len, _ = x.shape # Normalize input before projection to prevent activation explosion x_norm = self.norm_latent(x) # Project to latent space latent = self.to_latent(x_norm) # Generate Q/K/V q = self.q_proj(latent) k = self.k_proj(latent) v = self.v_proj(latent) # Reshape for multi-head attention: (batch_size, seq_len, d_kv_comp) -> (batch_size, n_heads, seq_len, d_head) q = q.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2) k = k.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2) v = v.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2) # Normalize Q and K for stable attention (standard practice in modern attention mechanisms) q = F.normalize(q, dim=-1, p=2) k = F.normalize(k, dim=-1, p=2) # Apply RoPE if self.d_rope > 0: rotary_emb = self.rotary(seq_len, x.device) cos = torch.cos(rotary_emb).unsqueeze(0).unsqueeze(0) sin = torch.sin(rotary_emb).unsqueeze(0).unsqueeze(0) q_rot = apply_rotary(q[..., :self.d_rope], cos, sin) k_rot = apply_rotary(k[..., :self.d_rope], cos, sin) q = torch.cat([q_rot, q[..., self.d_rope:]], dim=-1) k = torch.cat([k_rot, k[..., self.d_rope:]], dim=-1) # Attention computation with numerical stability # Scale before matmul to prevent overflow attn_scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale # Clamp attention scores to prevent inf/-inf in softmax attn_scores = torch.clamp(attn_scores, min=-20.0, max=20.0) if mask is not None: attn_scores = attn_scores.masked_fill(mask == 0, float('-inf')) # Numerically stable softmax attn_weights = F.softmax(attn_scores, dim=-1) # Check for NaN and print warning if torch.isnan(attn_weights).any(): print(f"WARNING: NaN detected in attention weights! " f"attn_scores min={attn_scores.min():.4f}, max={attn_scores.max():.4f}, " f"attn_weights min={attn_weights.min():.4f}, max={attn_weights.max():.4f}") attn_weights = self.attn_dropout(attn_weights) # Apply attention to values out = torch.matmul(attn_weights, v) # Reshape back: (batch_size, n_heads, seq_len, d_head) -> (batch_size, seq_len, d_kv_comp) out = out.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) # Project back to model dimension out = self.out_proj(out) out = self.resid_dropout(out) return out