# coding=utf-8 """ Native Sparse Attention (NSA) — 3-branch design. ⚠️ 동생 작성 placeholder (사용자님 본 PC 정식 버전 도착 시 교체) Branches: 1. Compress: average over blocks → KV cache reduction 2. Select: top-k token selection per query 3. Sliding: local window attention 4. Gate: learnable combination of 3 branches """ import math from typing import Optional, Tuple import torch import torch.nn as nn import torch.nn.functional as F class NSAAttention(nn.Module): """3-branch Native Sparse Attention. For testing purposes. Production version 사용자님 본 PC. """ 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.num_kv_heads = getattr(config, "num_key_value_heads", config.num_attention_heads) self.head_dim = config.head_dim self.num_kv_groups = self.num_heads // self.num_kv_heads self.compress_block = getattr(config, "compress_block_size", 16) self.sliding_window = getattr(config, "sliding_window_size", 256) self.select_topk = getattr(config, "nsa_select_topk", 16) # Shared Q/K/V projections (3 branches share base projections) self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=False) self.k_proj = nn.Linear(self.hidden_size, self.num_kv_heads * self.head_dim, bias=False) self.v_proj = nn.Linear(self.hidden_size, self.num_kv_heads * self.head_dim, bias=False) self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False) # Compress branch: pooling MLP self.compress_mlp = nn.Linear(self.head_dim, self.head_dim, bias=False) # Select branch: scoring head (project to scalar score per token) self.select_score = nn.Linear(self.head_dim, 1, bias=False) # 3-branch gate (initialized at 0.5 each, will be normalized) self.gate_logit = nn.Parameter(torch.zeros(3)) # sigmoid → ~0.5 def _repeat_kv(self, x: torch.Tensor) -> torch.Tensor: """GQA: repeat KV heads to match Q heads.""" if self.num_kv_groups == 1: return x bsz, n_kv, seq, dim = x.shape return (x[:, :, None, :, :] .expand(bsz, n_kv, self.num_kv_groups, seq, dim) .reshape(bsz, n_kv * self.num_kv_groups, seq, dim)) def _compress_branch(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor: """Compress branch: average KV in blocks, then full attention on compressed.""" bsz, n_h, seq, dim = q.shape if seq < self.compress_block * 2: return self._full_attention(q, k, v) n_blocks = seq // self.compress_block # Compress K, V by averaging blocks (only complete blocks) usable = n_blocks * self.compress_block k_comp = k[:, :, :usable, :].view(bsz, n_h, n_blocks, self.compress_block, dim).mean(dim=3) v_comp = v[:, :, :usable, :].view(bsz, n_h, n_blocks, self.compress_block, dim).mean(dim=3) # Apply MLP to compressed K k_comp = self.compress_mlp(k_comp) # Full attention against compressed return F.scaled_dot_product_attention(q, k_comp, v_comp, dropout_p=0.0, is_causal=False) def _select_branch(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor: """Select branch: top-k token selection per query position.""" bsz, n_h, seq, dim = q.shape if seq <= self.select_topk: return self._full_attention(q, k, v) # Score each KV position via select_score on K scores = self.select_score(k).squeeze(-1) # (bsz, n_h, seq) topk = min(self.select_topk, seq) _, top_idx = scores.topk(topk, dim=-1) # (bsz, n_h, topk) # Gather top-k K, V (causal mask: only attend to past) idx_expand = top_idx.unsqueeze(-1).expand(-1, -1, -1, dim) k_sel = k.gather(2, idx_expand) v_sel = v.gather(2, idx_expand) return F.scaled_dot_product_attention(q, k_sel, v_sel, dropout_p=0.0, is_causal=False) def _sliding_branch(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor: """Sliding window: only attend to local window.""" bsz, n_h, seq, dim = q.shape if seq <= self.sliding_window: return self._full_attention(q, k, v) # Build sliding window mask idx = torch.arange(seq, device=q.device) rel = idx.unsqueeze(0) - idx.unsqueeze(1) # (seq, seq) mask = (rel >= 0) & (rel <= self.sliding_window) # local + causal attn_mask = torch.where(mask, 0.0, float("-inf")).unsqueeze(0).unsqueeze(0).to(q.dtype) return F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=0.0) def _full_attention(self, q, k, v): return F.scaled_dot_product_attention(q, k, v, dropout_p=0.0, is_causal=True) def forward( self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_value=None, use_cache: bool = False, **kwargs, ) -> Tuple[torch.Tensor, Optional[object]]: bsz, q_len, _ = hidden_states.size() q = self.q_proj(hidden_states).view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2) k = self.k_proj(hidden_states).view(bsz, q_len, self.num_kv_heads, self.head_dim).transpose(1, 2) v = self.v_proj(hidden_states).view(bsz, q_len, self.num_kv_heads, self.head_dim).transpose(1, 2) # GQA: repeat KV to match Q heads k = self._repeat_kv(k) v = self._repeat_kv(v) # 3 branches out_compress = self._compress_branch(q, k, v) out_select = self._select_branch(q, k, v) out_sliding = self._sliding_branch(q, k, v) # Gate (softmax over 3 branches) gate = F.softmax(self.gate_logit, dim=0) out = (gate[0] * out_compress + gate[1] * out_select + gate[2] * out_sliding) # Reshape and project out = out.transpose(1, 2).contiguous().view(bsz, q_len, -1) return self.o_proj(out), past_key_value __all__ = ["NSAAttention"]