SeaWolf-AI's picture
fix: make model loadable via AutoModelForCausalLM (add auto_map, flatten module files to repo root, inherit GenerationMixin)
0957e22 verified
Raw
History Blame Contribute Delete
3.99 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)
__all__ = ['MLAAttention']