Download model.py from X-Orange/Data_Quality_Scoring_Model: direct link, hf CLI and curl.
- Browser
- Download file 10.7 kB
-
https://huggingface.co/X-Orange/Data_Quality_Scoring_Model/resolve/main/model.py
- Command line
-
hf download hf://X-Orange/Data_Quality_Scoring_Model/model.py
-
curl -L -o model.py https://huggingface.co/X-Orange/Data_Quality_Scoring_Model/resolve/main/model.py
10.7 kB
| import math | |
| from typing import Optional | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| # YaRN Rotary Position Embedding | |
| 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 | |
| # Scaled Dot-Product Attention (with GQA support) | |
| 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 | |
| # Grouped Query Attention | |
| 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 | |
| # Gated GELU Feed-Forward Network | |
| 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) | |
| # Transformer Decoder Layer | |
| 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 | |
| # Decoder with Dense Layer Connections | |
| 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 | |
| # Classifier | |
| 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 | |