import math from typing import Optional import torch import torch.nn as nn import torch.nn.functional as F # YaRN Rotary Position Embedding class YaRNRoPE(nn.Module): def __init__( self, head_dim: int, original_max_seq_len: int = 4096, factor: float = 1.0, base: float = 10000.0, beta_fast: int = 32, beta_slow: int = 1, ): super().__init__() self.head_dim = head_dim self.original_max_seq_len = original_max_seq_len self.factor = factor if factor > 1.0: self.attention_factor = math.log(factor) * 0.1 + 1.0 t = torch.arange(head_dim // 2) inv_freq = 1.0 / (base ** (2 * t.float() / head_dim)) wavelength = 2 * math.pi / inv_freq low_freq_wavelen = original_max_seq_len / beta_slow high_freq_wavelen = original_max_seq_len / beta_fast ratio = (wavelength - high_freq_wavelen) / (low_freq_wavelen - high_freq_wavelen) ratio = torch.clamp(ratio, 0.0, 1.0) scale = 1 - ratio + ratio * factor inv_freq = inv_freq / scale else: self.attention_factor = 1.0 inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2).float() / head_dim)) self.register_buffer("inv_freq", inv_freq) self._set_cos_sin_cache(int(original_max_seq_len * factor)) def _set_cos_sin_cache(self, seq_len: int): t = torch.arange(seq_len, device=self.inv_freq.device) freqs = torch.outer(t, self.inv_freq) emb = torch.cat((freqs, freqs), dim=-1) self.register_buffer("cos_cached", emb.cos()[None, None, :, :], persistent=False) self.register_buffer("sin_cached", emb.sin()[None, None, :, :], persistent=False) self.max_seq_len_cached = seq_len def forward(self, x: torch.Tensor, seq_len: Optional[int] = None): if seq_len is None: seq_len = x.shape[-2] if seq_len > self.max_seq_len_cached: self._set_cos_sin_cache(seq_len) cos = self.cos_cached[:, :, :seq_len, :] sin = self.sin_cached[:, :, :seq_len, :] x1, x2 = x[..., ::2], x[..., 1::2] rotated = torch.stack( [ x1 * cos[..., ::2] - x2 * sin[..., ::2], x1 * sin[..., ::2] + x2 * cos[..., ::2], ], dim=-1, ).flatten(-2) return rotated * self.attention_factor # Scaled Dot-Product Attention (with GQA support) def scaled_dot_product_attention( query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, dropout: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, ) -> torch.Tensor: B, Hq, L, E = query.shape _, Hkv, S, _ = key.shape if enable_gqa and Hq != Hkv: assert Hq % Hkv == 0 n_rep = Hq // Hkv key = key.unsqueeze(2).repeat(1, 1, n_rep, 1, 1).flatten(1, 2) value = value.unsqueeze(2).repeat(1, 1, n_rep, 1, 1).flatten(1, 2) if scale is None: scale = E ** -0.5 scores = torch.matmul(query, key.transpose(-2, -1)) * scale if is_causal and attention_mask is not None: raise RuntimeError("is_causal and attention_mask cannot be set at the same time") if is_causal: causal_mask = torch.triu(torch.ones(L, S, dtype=torch.bool, device=query.device), diagonal=1) scores = scores.masked_fill(causal_mask, float("-inf")) if attention_mask is not None: if attention_mask.dtype == torch.bool: scores = scores.masked_fill(~attention_mask, float("-inf")) else: scores = scores + attention_mask attn_weights = F.softmax(scores, dim=-1) if dropout > 0.0: attn_weights = F.dropout(attn_weights, p=dropout, training=True) output = torch.matmul(attn_weights, value) return output # Grouped Query Attention class GroupedQueryAttention(nn.Module): def __init__( self, hidden_size: int, num_heads: int, num_key_value_heads: int, head_dim: int, max_seq_len: int, ): super().__init__() self.hidden_size = hidden_size self.num_heads = num_heads self.num_key_value_heads = num_key_value_heads self.head_dim = head_dim self.max_seq_len = max_seq_len self.q_proj = nn.Linear(hidden_size, head_dim * num_heads, bias=False) self.k_proj = nn.Linear(hidden_size, head_dim * num_key_value_heads, bias=False) self.v_proj = nn.Linear(hidden_size, head_dim * num_key_value_heads, bias=False) self.out_proj = nn.Linear(num_heads * head_dim, hidden_size, bias=False) self.rope = YaRNRoPE( head_dim=head_dim, original_max_seq_len=max_seq_len, factor=16.0, ) def forward(self, query, key, value): B, L_q, _ = query.size() _, L_kv, _ = key.size() q = self.q_proj(query).view(B, L_q, self.num_heads, self.head_dim).transpose(1, 2) k = self.k_proj(key).view(B, L_kv, self.num_key_value_heads, self.head_dim).transpose(1, 2) v = self.v_proj(value).view(B, L_kv, self.num_key_value_heads, self.head_dim).transpose(1, 2) q_embed = self.rope(q) k_embed = self.rope(k) q_embed, k_embed = q_embed.to(q.dtype), k_embed.to(k.dtype) attn_output = scaled_dot_product_attention( q_embed, k_embed, v, attention_mask=torch.ones(L_q, L_kv, dtype=torch.bool, device=q_embed.device), dropout=0.0, is_causal=False, enable_gqa=True, ) context = attn_output.transpose(1, 2).contiguous().view(B, L_q, self.num_heads * self.head_dim) output = self.out_proj(context) return output # Gated GELU Feed-Forward Network class GEGLU(nn.Module): def __init__(self, hidden_size: int, intermediate_size: Optional[int] = None): super().__init__() if intermediate_size is None: intermediate_size = int(8 / 3 * hidden_size) self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False) self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False) self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) def forward(self, x): gate = F.gelu(self.gate_proj(x)) value = self.up_proj(x) hidden = gate * value return self.down_proj(hidden) # Transformer Decoder Layer class TransformerDecoderLayer(nn.Module): def __init__( self, hidden_size: int, num_heads: int, num_key_value_heads: int, intermediate_size: int, head_dim: int, max_seq_len: int, dropout: float, ): super().__init__() self.self_attn = GroupedQueryAttention( hidden_size, num_heads, num_key_value_heads, head_dim, max_seq_len ) self.ffn = GEGLU(hidden_size, intermediate_size) self.input_layernorm = nn.RMSNorm(hidden_size) self.post_attention_layernorm = nn.RMSNorm(hidden_size) self.dropout = nn.Dropout(dropout) def forward(self, hidden_states): residual = hidden_states hidden_states = self.input_layernorm(hidden_states) attn_output = self.dropout(self.self_attn(hidden_states, hidden_states, hidden_states)) hidden_states = residual + attn_output residual = hidden_states hidden_states = self.post_attention_layernorm(hidden_states) ffn_output = self.dropout(self.ffn(hidden_states)) hidden_states = residual + ffn_output return hidden_states # Decoder with Dense Layer Connections class TransformerDecoder(nn.Module): def __init__( self, hidden_size: int, num_heads: int, num_key_value_heads: int, intermediate_size: int, head_dim: int, num_layers: int, max_seq_len: int, dropout: float, ): super().__init__() self.num_layers = num_layers self.layers = nn.ModuleList([ TransformerDecoderLayer( hidden_size, num_heads, num_key_value_heads, intermediate_size, head_dim, max_seq_len, dropout ) for _ in range(num_layers) ]) mask = torch.tril(torch.ones(num_layers, num_layers), diagonal=-1) self.register_buffer("layer_weight_mask", mask) self.layer_raw_weights = nn.Parameter(torch.randn(num_layers, num_layers) / 10) def forward(self, hidden_states): history = [] for idx_layer, layer in enumerate(self.layers): layer_output = layer(hidden_states) if history: raw_weights = self.layer_raw_weights[idx_layer, :idx_layer] masked_weights = raw_weights * self.layer_weight_mask[idx_layer, :idx_layer] weights = F.softmax(masked_weights, dim=0) hist_stack = torch.stack(history, dim=0) residual = torch.einsum("lbtd,l->btd", hist_stack, weights) hidden_states = layer_output + residual else: hidden_states = layer_output history.append(hidden_states) return hidden_states # Classifier class Classifier(nn.Module): def __init__( self, hidden_size: int, num_heads: int, num_key_value_heads: int, intermediate_size: int, head_dim: int, vocab_size: int, num_layers: int, max_seq_len: int, dropout: float, ): super().__init__() self.token_embedding = nn.Embedding(vocab_size, hidden_size) self.decoder = TransformerDecoder( hidden_size, num_heads, num_key_value_heads, intermediate_size, head_dim, num_layers, max_seq_len, dropout ) self.final_layernorm = nn.RMSNorm(hidden_size) self.lm_head = nn.Linear(hidden_size, 6, bias=False) def forward(self, input_ids): hidden_states = self.token_embedding(input_ids) hidden_states = self.decoder(hidden_states) hidden_states = self.final_layernorm(hidden_states) logits = self.lm_head(hidden_states).mean(-2) return logits