Download modeling_sing_probe.py from inclusionAI/MiniMax-M2.7-singprobe: direct link, hf CLI and curl.
- Browser
- Download file 6.07 kB
-
https://huggingface.co/inclusionAI/MiniMax-M2.7-singprobe/resolve/main/modeling_sing_probe.py
- Command line
-
hf download hf://inclusionAI/MiniMax-M2.7-singprobe/modeling_sing_probe.py
-
curl -L -o modeling_sing_probe.py https://huggingface.co/inclusionAI/MiniMax-M2.7-singprobe/resolve/main/modeling_sing_probe.py
6.07 kB
| from typing import Any | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import PreTrainedModel | |
| from transformers.activations import ACT2FN | |
| from transformers.modeling_outputs import TokenClassifierOutput | |
| from .configuration_sing_probe import SingProbeAttnConfig, SingProbeMlpConfig | |
| class SingProbePreTrainedModel(PreTrainedModel): | |
| base_model_prefix = "" | |
| main_input_name = "hidden_states" | |
| def _init_weights(self, module: nn.Module) -> None: | |
| if isinstance(module, nn.Linear): | |
| nn.init.normal_(module.weight, mean=0.0, std=0.02) | |
| if module.bias is not None: | |
| nn.init.zeros_(module.bias) | |
| elif isinstance(module, nn.RMSNorm): | |
| nn.init.ones_(module.weight) | |
| def _validate_input(config: Any, hidden_states: torch.Tensor) -> None: | |
| if hidden_states.shape[-1] != config.input_size: | |
| raise ValueError( | |
| f"input feature size {hidden_states.shape[-1]} does not match " | |
| f"hidden_size * number of tapped layers ({config.input_size})" | |
| ) | |
| class SingProbeMlpModel(SingProbePreTrainedModel): | |
| config_class = SingProbeMlpConfig | |
| def __init__(self, config: SingProbeMlpConfig) -> None: | |
| super().__init__(config) | |
| self.fc1 = nn.Linear(config.input_size, config.intermediate_size) | |
| self.fc2 = nn.Linear(config.intermediate_size, config.num_labels) | |
| self.act_fn = ACT2FN[config.hidden_act] | |
| self.post_init() | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| return_dict: bool | None = None, | |
| **_: Any, | |
| ) -> TokenClassifierOutput | tuple[torch.Tensor]: | |
| self._validate_input(self.config, hidden_states) | |
| hidden_states = hidden_states.to(self.fc1.weight.dtype) | |
| hidden_states = self.fc1(hidden_states) | |
| hidden_states = self.act_fn(hidden_states) | |
| logits = self.fc2(hidden_states) | |
| if return_dict is False: | |
| return (logits,) | |
| return TokenClassifierOutput(logits=logits) | |
| class SingProbeAttnModel(SingProbePreTrainedModel): | |
| config_class = SingProbeAttnConfig | |
| def __init__(self, config: SingProbeAttnConfig) -> None: | |
| super().__init__(config) | |
| if config.num_attention_heads < 1 or config.head_dim < 1: | |
| raise ValueError("num_attention_heads and head_dim must be positive") | |
| if config.sliding_window is not None and config.sliding_window <= 0: | |
| raise ValueError("sliding_window must be positive") | |
| self.num_attention_heads = config.num_attention_heads | |
| self.head_dim = config.head_dim | |
| self.projection_size = self.num_attention_heads * self.head_dim | |
| self.proj_q = nn.Linear(config.input_size, self.projection_size, bias=False) | |
| self.proj_k = nn.Linear(config.input_size, self.head_dim, bias=False) | |
| self.proj_v = nn.Linear(config.input_size, self.head_dim, bias=False) | |
| self.o_proj = nn.Linear(self.projection_size, self.projection_size, bias=False) | |
| self.norm = nn.RMSNorm(self.projection_size, eps=1e-6) | |
| self.classifier = nn.Linear(self.projection_size, config.num_labels) | |
| self.post_init() | |
| def _sliding_window_attention( | |
| self, | |
| query: torch.Tensor, | |
| key: torch.Tensor, | |
| value: torch.Tensor, | |
| ) -> torch.Tensor: | |
| seq_len = query.shape[2] | |
| window = self.config.sliding_window | |
| assert window is not None | |
| outputs = [] | |
| for start in range(0, seq_len, window): | |
| end = min(start + window, seq_len) | |
| key_start = max(0, start - window + 1) | |
| query_positions = torch.arange(start, end, device=query.device) | |
| key_positions = torch.arange(key_start, end, device=query.device) | |
| relative_positions = query_positions[:, None] - key_positions | |
| attention_mask = (relative_positions >= 0) & (relative_positions < window) | |
| attention_mask = attention_mask[None, None] | |
| query_block = query[:, :, start:end] | |
| key_block = key[:, :, key_start:end] | |
| value_block = value[:, :, key_start:end] | |
| attention_output = F.scaled_dot_product_attention( | |
| query_block, | |
| key_block, | |
| value_block, | |
| attn_mask=attention_mask, | |
| enable_gqa=self.num_attention_heads > 1, | |
| ) | |
| outputs.append(attention_output) | |
| return torch.cat(outputs, dim=2) | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| return_dict: bool | None = None, | |
| **_: Any, | |
| ) -> TokenClassifierOutput | tuple[torch.Tensor]: | |
| self._validate_input(self.config, hidden_states) | |
| hidden_states = hidden_states.to(self.classifier.weight.dtype) | |
| batch_size, seq_len, _ = hidden_states.shape | |
| query_features = self.proj_q(hidden_states) | |
| query_shape = (batch_size, seq_len, self.num_attention_heads, self.head_dim) | |
| kv_shape = (batch_size, seq_len, 1, self.head_dim) | |
| query = query_features.view(query_shape).transpose(1, 2) | |
| key = self.proj_k(hidden_states).view(kv_shape).transpose(1, 2) | |
| value = self.proj_v(hidden_states).view(kv_shape).transpose(1, 2) | |
| if self.config.sliding_window and self.config.sliding_window < seq_len: | |
| context = self._sliding_window_attention(query, key, value) | |
| else: | |
| context = F.scaled_dot_product_attention( | |
| query=query, | |
| key=key, | |
| value=value, | |
| is_causal=True, | |
| enable_gqa=self.num_attention_heads > 1, | |
| ) | |
| context = context.transpose(1, 2).contiguous() | |
| context = context.view(batch_size, seq_len, self.projection_size) | |
| hidden_states = self.o_proj(context) + query_features | |
| logits = self.classifier(self.norm(hidden_states)) | |
| if return_dict is False: | |
| return (logits,) | |
| return TokenClassifierOutput(logits=logits) | |