# coding=utf-8 # AETHER-30B-11Attn — Auto-extracted from new_attentions.py import math from typing import Optional, Tuple import torch import torch.nn as nn import torch.nn.functional as F # ============================================================================= # 1. MLA (Multi-head Latent Attention) — DeepSeek-V2/V3 # ============================================================================= class MLAAttention(nn.Module): """Multi-head Latent Attention (low-rank Q/KV decomposition). DeepSeek-V2 paper: Q/KV을 low-rank로 분해 → KV cache 1/14 절감. Q: hidden → q_lora_rank → split (rope + nope) → multi-head KV: hidden → kv_lora_rank → split (rope + nope) → K (rope+nope), V (nope) """ def __init__(self, config, layer_idx: int = 0): super().__init__() self.config = config self.layer_idx = layer_idx self.hidden_size = config.hidden_size self.num_heads = config.num_attention_heads self.q_lora_rank = getattr(config, "mla_q_lora_rank", 768) self.kv_lora_rank = getattr(config, "mla_kv_lora_rank", 256) self.qk_rope_head_dim = getattr(config, "mla_qk_rope_head_dim", 64) self.qk_nope_head_dim = getattr(config, "mla_qk_nope_head_dim", 64) self.v_head_dim = getattr(config, "mla_v_head_dim", 128) qk_head_dim = self.qk_rope_head_dim + self.qk_nope_head_dim # Q low-rank self.q_a_proj = nn.Linear(self.hidden_size, self.q_lora_rank, bias=False) self.q_a_norm = nn.LayerNorm(self.q_lora_rank, eps=1e-5) self.q_b_proj = nn.Linear(self.q_lora_rank, self.num_heads * qk_head_dim, bias=False) # KV low-rank (compressed) self.kv_a_proj = nn.Linear(self.hidden_size, self.kv_lora_rank + self.qk_rope_head_dim, bias=False) self.kv_a_norm = nn.LayerNorm(self.kv_lora_rank, eps=1e-5) self.kv_b_proj = nn.Linear( self.kv_lora_rank, self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), bias=False, ) # Output self.o_proj = nn.Linear(self.num_heads * self.v_head_dim, self.hidden_size, bias=False) def forward(self, hidden_states, attention_mask=None, position_ids=None, past_key_value=None, use_cache=False, **kwargs): bsz, q_len, _ = hidden_states.shape qk_head_dim = self.qk_rope_head_dim + self.qk_nope_head_dim # Q q = self.q_a_proj(hidden_states) q = self.q_a_norm(q) q = self.q_b_proj(q).view(bsz, q_len, self.num_heads, qk_head_dim).transpose(1, 2) # KV kv = self.kv_a_proj(hidden_states) kv_compressed, k_rope = kv.split([self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) kv_compressed = self.kv_a_norm(kv_compressed) kv_b = self.kv_b_proj(kv_compressed).view( bsz, q_len, self.num_heads, self.qk_nope_head_dim + self.v_head_dim ).transpose(1, 2) k_nope, v = kv_b.split([self.qk_nope_head_dim, self.v_head_dim], dim=-1) # Combine k (broadcast k_rope to all heads) k_rope_expanded = k_rope.view(bsz, q_len, 1, self.qk_rope_head_dim).expand( -1, -1, self.num_heads, -1 ).transpose(1, 2) k = torch.cat([k_nope, k_rope_expanded], dim=-1) # (bsz, n_h, seq, qk_head_dim) # SDPA attn_out = F.scaled_dot_product_attention( q, k, v, attn_mask=(attention_mask.to(q.dtype) if attention_mask is not None else None), dropout_p=0.0, is_causal=(attention_mask is None and q_len > 1), ) attn_out = attn_out.transpose(1, 2).contiguous().view(bsz, q_len, -1) return self.o_proj(attn_out), past_key_value # ============================================================================= # 2. Mamba2 — State Space Model (placeholder simplification) # ============================================================================= class Mamba2Attention(nn.Module): """Mamba2 SSM placeholder (simplified S6 selective scan). True Mamba2: selective_scan_fn from mamba_ssm package. Here: simplified GRU-style mixing as placeholder. """ def __init__(self, config, layer_idx: int = 0): super().__init__() self.config = config self.layer_idx = layer_idx self.hidden_size = config.hidden_size self.d_state = getattr(config, "mamba2_d_state", 128) self.d_conv = getattr(config, "mamba2_d_conv", 4) self.expand = getattr(config, "mamba2_expand", 2) self.d_inner = self.expand * self.hidden_size self.in_proj = nn.Linear(self.hidden_size, 2 * self.d_inner, bias=False) self.conv1d = nn.Conv1d(self.d_inner, self.d_inner, kernel_size=self.d_conv, groups=self.d_inner, padding=self.d_conv - 1, bias=False) # SSM params (simplified) self.A_log = nn.Parameter(torch.zeros(self.d_inner, self.d_state)) self.D = nn.Parameter(torch.zeros(self.d_inner)) self.dt_proj = nn.Linear(self.d_inner, self.d_inner, bias=True) self.B_proj = nn.Linear(self.d_inner, self.d_state, bias=False) self.C_proj = nn.Linear(self.d_inner, self.d_state, bias=False) self.out_proj = nn.Linear(self.d_inner, self.hidden_size, bias=False) self.norm = nn.LayerNorm(self.d_inner, eps=1e-5) def forward(self, hidden_states, attention_mask=None, position_ids=None, past_key_value=None, use_cache=False, **kwargs): bsz, seq, _ = hidden_states.shape # Project xz = self.in_proj(hidden_states) x, z = xz.chunk(2, dim=-1) # (bsz, seq, d_inner) each # 1D conv (causal) x = x.transpose(1, 2) # (bsz, d_inner, seq) x = self.conv1d(x)[:, :, :seq] # causal trim x = F.silu(x).transpose(1, 2) # (bsz, seq, d_inner) # Simplified SSM: scan as cumulative recurrent # For placeholder, we use the simple gated linear approximation dt = F.softplus(self.dt_proj(x)) # (bsz, seq, d_inner) out = x * dt + self.D[None, None, :] * x # ultra-simplified out = self.norm(out) out = out * F.silu(z) # gating out = self.out_proj(out) return out, past_key_value # ============================================================================= # 3. GDN (Gated Delta Network) — RWKV/RetNet style __all__ = ['Mamba2Attention']