| import math |
| import torch |
| import torch.nn.functional as F |
| from torch import nn |
|
|
| from core.norm import RMSNorm |
| from core.rope import apply_rotary_pos_emb, repeat_kv |
|
|
|
|
| class Attention(nn.Module): |
| def __init__(self, config: "LMConfig"): |
| super().__init__() |
| self.num_key_value_heads = config.num_attention_heads if config.num_key_value_heads is None else config.num_key_value_heads |
| self.n_local_heads = config.num_attention_heads |
| self.n_local_kv_heads = self.num_key_value_heads |
| self.n_rep = self.n_local_heads // self.n_local_kv_heads |
| self.head_dim = config.head_dim |
| self.is_causal = True |
| self.q_proj = nn.Linear(config.hidden_size, config.num_attention_heads * self.head_dim, bias=False) |
| self.k_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False) |
| self.v_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False) |
| self.o_proj = nn.Linear(config.num_attention_heads * self.head_dim, config.hidden_size, bias=False) |
| self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps) |
| self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps) |
| self.attn_dropout = nn.Dropout(config.dropout) |
| self.resid_dropout = nn.Dropout(config.dropout) |
| self.dropout = config.dropout |
| self.flash = hasattr(torch.nn.functional, 'scaled_dot_product_attention') and config.flash_attn |
|
|
| def forward(self, x, position_embeddings, past_key_value=None, use_cache=False, attention_mask=None): |
| bsz, seq_len, _ = x.shape |
| xq, xk, xv = self.q_proj(x), self.k_proj(x), self.v_proj(x) |
| xq = xq.view(bsz, seq_len, self.n_local_heads, self.head_dim) |
| xk = xk.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim) |
| xv = xv.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim) |
| xq, xk = self.q_norm(xq), self.k_norm(xk) |
| cos, sin = position_embeddings |
| xq, xk = apply_rotary_pos_emb(xq, xk, cos, sin) |
| if past_key_value is not None: |
| xk = torch.cat([past_key_value[0], xk], dim=1) |
| xv = torch.cat([past_key_value[1], xv], dim=1) |
| past_kv = (xk, xv) if use_cache else None |
| xq, xk, xv = (xq.transpose(1, 2), repeat_kv(xk, self.n_rep).transpose(1, 2), repeat_kv(xv, self.n_rep).transpose(1, 2)) |
| if self.flash and (seq_len > 1) and (not self.is_causal or past_key_value is None) and (attention_mask is None or torch.all(attention_mask == 1)): |
| output = F.scaled_dot_product_attention(xq, xk, xv, dropout_p=self.dropout if self.training else 0.0, is_causal=self.is_causal) |
| else: |
| scores = (xq @ xk.transpose(-2, -1)) / math.sqrt(self.head_dim) |
| if self.is_causal: |
| scores[:, :, :, -seq_len:] += torch.full((seq_len, seq_len), float("-inf"), device=scores.device).triu(1) |
| if attention_mask is not None: |
| scores += (1.0 - attention_mask.unsqueeze(1).unsqueeze(2)) * -1e9 |
| output = self.attn_dropout(F.softmax(scores.float(), dim=-1).type_as(xq)) @ xv |
| output = output.transpose(1, 2).reshape(bsz, seq_len, -1) |
| output = self.resid_dropout(self.o_proj(output)) |
| return output, past_kv |
|
|