| import math
|
| from typing import Optional
|
|
|
| import torch
|
| import torch.nn as nn
|
| import torch.nn.functional as F
|
|
|
|
|
|
|
| class YaRNRoPE(nn.Module):
|
| def __init__(
|
| self,
|
| head_dim: int,
|
| original_max_seq_len: int = 4096,
|
| factor: float = 1.0,
|
| base: float = 10000.0,
|
| beta_fast: int = 32,
|
| beta_slow: int = 1,
|
| ):
|
| super().__init__()
|
| self.head_dim = head_dim
|
| self.original_max_seq_len = original_max_seq_len
|
| self.factor = factor
|
|
|
| if factor > 1.0:
|
| self.attention_factor = math.log(factor) * 0.1 + 1.0
|
| t = torch.arange(head_dim // 2)
|
| inv_freq = 1.0 / (base ** (2 * t.float() / head_dim))
|
| wavelength = 2 * math.pi / inv_freq
|
| low_freq_wavelen = original_max_seq_len / beta_slow
|
| high_freq_wavelen = original_max_seq_len / beta_fast
|
| ratio = (wavelength - high_freq_wavelen) / (low_freq_wavelen - high_freq_wavelen)
|
| ratio = torch.clamp(ratio, 0.0, 1.0)
|
| scale = 1 - ratio + ratio * factor
|
| inv_freq = inv_freq / scale
|
| else:
|
| self.attention_factor = 1.0
|
| inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2).float() / head_dim))
|
|
|
| self.register_buffer("inv_freq", inv_freq)
|
| self._set_cos_sin_cache(int(original_max_seq_len * factor))
|
|
|
| def _set_cos_sin_cache(self, seq_len: int):
|
| t = torch.arange(seq_len, device=self.inv_freq.device)
|
| freqs = torch.outer(t, self.inv_freq)
|
| emb = torch.cat((freqs, freqs), dim=-1)
|
| self.register_buffer("cos_cached", emb.cos()[None, None, :, :], persistent=False)
|
| self.register_buffer("sin_cached", emb.sin()[None, None, :, :], persistent=False)
|
| self.max_seq_len_cached = seq_len
|
|
|
| def forward(self, x: torch.Tensor, seq_len: Optional[int] = None):
|
| if seq_len is None:
|
| seq_len = x.shape[-2]
|
| if seq_len > self.max_seq_len_cached:
|
| self._set_cos_sin_cache(seq_len)
|
|
|
| cos = self.cos_cached[:, :, :seq_len, :]
|
| sin = self.sin_cached[:, :, :seq_len, :]
|
|
|
| x1, x2 = x[..., ::2], x[..., 1::2]
|
| rotated = torch.stack(
|
| [
|
| x1 * cos[..., ::2] - x2 * sin[..., ::2],
|
| x1 * sin[..., ::2] + x2 * cos[..., ::2],
|
| ],
|
| dim=-1,
|
| ).flatten(-2)
|
|
|
| return rotated * self.attention_factor
|
|
|
|
|
|
|
| def scaled_dot_product_attention(
|
| query: torch.Tensor,
|
| key: torch.Tensor,
|
| value: torch.Tensor,
|
| attention_mask: Optional[torch.Tensor] = None,
|
| dropout: float = 0.0,
|
| is_causal: bool = False,
|
| scale: Optional[float] = None,
|
| enable_gqa: bool = False,
|
| ) -> torch.Tensor:
|
| B, Hq, L, E = query.shape
|
| _, Hkv, S, _ = key.shape
|
|
|
| if enable_gqa and Hq != Hkv:
|
| assert Hq % Hkv == 0
|
| n_rep = Hq // Hkv
|
| key = key.unsqueeze(2).repeat(1, 1, n_rep, 1, 1).flatten(1, 2)
|
| value = value.unsqueeze(2).repeat(1, 1, n_rep, 1, 1).flatten(1, 2)
|
| if scale is None:
|
| scale = E ** -0.5
|
|
|
| scores = torch.matmul(query, key.transpose(-2, -1)) * scale
|
|
|
| if is_causal and attention_mask is not None:
|
| raise RuntimeError("is_causal and attention_mask cannot be set at the same time")
|
| if is_causal:
|
| causal_mask = torch.triu(torch.ones(L, S, dtype=torch.bool, device=query.device), diagonal=1)
|
| scores = scores.masked_fill(causal_mask, float("-inf"))
|
| if attention_mask is not None:
|
| if attention_mask.dtype == torch.bool:
|
| scores = scores.masked_fill(~attention_mask, float("-inf"))
|
| else:
|
| scores = scores + attention_mask
|
|
|
| attn_weights = F.softmax(scores, dim=-1)
|
|
|
| if dropout > 0.0:
|
| attn_weights = F.dropout(attn_weights, p=dropout, training=True)
|
|
|
| output = torch.matmul(attn_weights, value)
|
| return output
|
|
|
|
|
|
|
| class GroupedQueryAttention(nn.Module):
|
| def __init__(
|
| self,
|
| hidden_size: int,
|
| num_heads: int,
|
| num_key_value_heads: int,
|
| head_dim: int,
|
| max_seq_len: int,
|
| ):
|
| super().__init__()
|
| self.hidden_size = hidden_size
|
| self.num_heads = num_heads
|
| self.num_key_value_heads = num_key_value_heads
|
| self.head_dim = head_dim
|
| self.max_seq_len = max_seq_len
|
|
|
| self.q_proj = nn.Linear(hidden_size, head_dim * num_heads, bias=False)
|
| self.k_proj = nn.Linear(hidden_size, head_dim * num_key_value_heads, bias=False)
|
| self.v_proj = nn.Linear(hidden_size, head_dim * num_key_value_heads, bias=False)
|
| self.out_proj = nn.Linear(num_heads * head_dim, hidden_size, bias=False)
|
|
|
| self.rope = YaRNRoPE(
|
| head_dim=head_dim,
|
| original_max_seq_len=max_seq_len,
|
| factor=16.0,
|
| )
|
|
|
| def forward(self, query, key, value):
|
| B, L_q, _ = query.size()
|
| _, L_kv, _ = key.size()
|
|
|
| q = self.q_proj(query).view(B, L_q, self.num_heads, self.head_dim).transpose(1, 2)
|
| k = self.k_proj(key).view(B, L_kv, self.num_key_value_heads, self.head_dim).transpose(1, 2)
|
| v = self.v_proj(value).view(B, L_kv, self.num_key_value_heads, self.head_dim).transpose(1, 2)
|
|
|
| q_embed = self.rope(q)
|
| k_embed = self.rope(k)
|
| q_embed, k_embed = q_embed.to(q.dtype), k_embed.to(k.dtype)
|
|
|
| attn_output = scaled_dot_product_attention(
|
| q_embed, k_embed, v,
|
| attention_mask=torch.ones(L_q, L_kv, dtype=torch.bool, device=q_embed.device),
|
| dropout=0.0,
|
| is_causal=False,
|
| enable_gqa=True,
|
| )
|
|
|
| context = attn_output.transpose(1, 2).contiguous().view(B, L_q, self.num_heads * self.head_dim)
|
| output = self.out_proj(context)
|
| return output
|
|
|
|
|
|
|
| class GEGLU(nn.Module):
|
| def __init__(self, hidden_size: int, intermediate_size: Optional[int] = None):
|
| super().__init__()
|
| if intermediate_size is None:
|
| intermediate_size = int(8 / 3 * hidden_size)
|
| self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
|
| self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
|
| self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
|
|
|
| def forward(self, x):
|
| gate = F.gelu(self.gate_proj(x))
|
| value = self.up_proj(x)
|
| hidden = gate * value
|
| return self.down_proj(hidden)
|
|
|
|
|
|
|
| class TransformerDecoderLayer(nn.Module):
|
| def __init__(
|
| self,
|
| hidden_size: int,
|
| num_heads: int,
|
| num_key_value_heads: int,
|
| intermediate_size: int,
|
| head_dim: int,
|
| max_seq_len: int,
|
| dropout: float,
|
| ):
|
| super().__init__()
|
| self.self_attn = GroupedQueryAttention(
|
| hidden_size, num_heads, num_key_value_heads, head_dim, max_seq_len
|
| )
|
| self.ffn = GEGLU(hidden_size, intermediate_size)
|
| self.input_layernorm = nn.RMSNorm(hidden_size)
|
| self.post_attention_layernorm = nn.RMSNorm(hidden_size)
|
| self.dropout = nn.Dropout(dropout)
|
|
|
| def forward(self, hidden_states):
|
| residual = hidden_states
|
| hidden_states = self.input_layernorm(hidden_states)
|
| attn_output = self.dropout(self.self_attn(hidden_states, hidden_states, hidden_states))
|
| hidden_states = residual + attn_output
|
|
|
| residual = hidden_states
|
| hidden_states = self.post_attention_layernorm(hidden_states)
|
| ffn_output = self.dropout(self.ffn(hidden_states))
|
| hidden_states = residual + ffn_output
|
|
|
| return hidden_states
|
|
|
|
|
|
|
| class TransformerDecoder(nn.Module):
|
| def __init__(
|
| self,
|
| hidden_size: int,
|
| num_heads: int,
|
| num_key_value_heads: int,
|
| intermediate_size: int,
|
| head_dim: int,
|
| num_layers: int,
|
| max_seq_len: int,
|
| dropout: float,
|
| ):
|
| super().__init__()
|
| self.num_layers = num_layers
|
| self.layers = nn.ModuleList([
|
| TransformerDecoderLayer(
|
| hidden_size, num_heads, num_key_value_heads,
|
| intermediate_size, head_dim, max_seq_len, dropout
|
| )
|
| for _ in range(num_layers)
|
| ])
|
|
|
| mask = torch.tril(torch.ones(num_layers, num_layers), diagonal=-1)
|
| self.register_buffer("layer_weight_mask", mask)
|
| self.layer_raw_weights = nn.Parameter(torch.randn(num_layers, num_layers) / 10)
|
|
|
| def forward(self, hidden_states):
|
| history = []
|
| for idx_layer, layer in enumerate(self.layers):
|
| layer_output = layer(hidden_states)
|
|
|
| if history:
|
| raw_weights = self.layer_raw_weights[idx_layer, :idx_layer]
|
| masked_weights = raw_weights * self.layer_weight_mask[idx_layer, :idx_layer]
|
| weights = F.softmax(masked_weights, dim=0)
|
| hist_stack = torch.stack(history, dim=0)
|
| residual = torch.einsum("lbtd,l->btd", hist_stack, weights)
|
| hidden_states = layer_output + residual
|
| else:
|
| hidden_states = layer_output
|
|
|
| history.append(hidden_states)
|
|
|
| return hidden_states
|
|
|
|
|
|
|
| class Classifier(nn.Module):
|
| def __init__(
|
| self,
|
| hidden_size: int,
|
| num_heads: int,
|
| num_key_value_heads: int,
|
| intermediate_size: int,
|
| head_dim: int,
|
| vocab_size: int,
|
| num_layers: int,
|
| max_seq_len: int,
|
| dropout: float,
|
| ):
|
| super().__init__()
|
| self.token_embedding = nn.Embedding(vocab_size, hidden_size)
|
| self.decoder = TransformerDecoder(
|
| hidden_size, num_heads, num_key_value_heads,
|
| intermediate_size, head_dim, num_layers, max_seq_len, dropout
|
| )
|
| self.final_layernorm = nn.RMSNorm(hidden_size)
|
| self.lm_head = nn.Linear(hidden_size, 6, bias=False)
|
|
|
| def forward(self, input_ids):
|
| hidden_states = self.token_embedding(input_ids)
|
| hidden_states = self.decoder(hidden_states)
|
| hidden_states = self.final_layernorm(hidden_states)
|
| logits = self.lm_head(hidden_states).mean(-2)
|
| return logits
|
|
|