SeaWolf-AI's picture
fix: make model loadable via AutoModelForCausalLM (add auto_map, flatten module files to repo root, inherit GenerationMixin)
01b945d verified
Raw
History Blame Contribute Delete
6.58 kB
# 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']