Download modeling_cherry_alpha.py from Banaxi-Tech/cherry-alpha: direct link, hf CLI and curl.
- Browser
- Download file 17.8 kB
-
https://huggingface.co/Banaxi-Tech/cherry-alpha/resolve/main/modeling_cherry_alpha.py
- Command line
-
hf download hf://Banaxi-Tech/cherry-alpha/modeling_cherry_alpha.py
-
curl -L -o modeling_cherry_alpha.py https://huggingface.co/Banaxi-Tech/cherry-alpha/resolve/main/modeling_cherry_alpha.py
17.8 kB
| """Cherry Alpha causal language model.""" | |
| from __future__ import annotations | |
| import math | |
| from typing import Optional | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from torch.utils.checkpoint import checkpoint | |
| from transformers import PreTrainedModel | |
| from transformers.generation.utils import GenerationMixin | |
| from transformers.modeling_outputs import CausalLMOutputWithPast | |
| try: | |
| from .configuration_cherry_alpha import CherryAlphaConfig | |
| except ImportError: # Standalone Hugging Face Jobs import downloaded files. | |
| from configuration_cherry_alpha import CherryAlphaConfig | |
| class CherryAlphaRMSNorm(nn.Module): | |
| def __init__(self, width: int, eps: float): | |
| super().__init__() | |
| self.weight = nn.Parameter(torch.ones(width)) | |
| self.eps = eps | |
| def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: | |
| source_dtype = hidden_states.dtype | |
| normalized = hidden_states.float() | |
| normalized = normalized * torch.rsqrt( | |
| normalized.square().mean(dim=-1, keepdim=True) + self.eps | |
| ) | |
| return (normalized * self.weight.float()).to(source_dtype) | |
| def rotate_half(hidden_states: torch.Tensor) -> torch.Tensor: | |
| first, second = hidden_states.chunk(2, dim=-1) | |
| return torch.cat((-second, first), dim=-1) | |
| class CherryAlphaRotaryEmbedding(nn.Module): | |
| def __init__(self, config: CherryAlphaConfig): | |
| super().__init__() | |
| inverse = 1.0 / ( | |
| config.rope_theta | |
| ** (torch.arange(0, config.head_dim, 2).float() / config.head_dim) | |
| ) | |
| positions = torch.arange(config.max_position_embeddings).float() | |
| frequencies = torch.outer(positions, inverse) | |
| angles = torch.cat((frequencies, frequencies), dim=-1) | |
| # These deterministic caches must be checkpointed. Hugging Face can | |
| # instantiate remote-code models under a low-memory/meta-device | |
| # context; non-persistent RoPE buffers can then remain invalid after | |
| # loading and make every continuation log-probability NaN. This is the | |
| # same export/load issue fixed in the original partial-Edu model. | |
| self.register_buffer("cosine", angles.cos(), persistent=True) | |
| self.register_buffer("sine", angles.sin(), persistent=True) | |
| def forward( | |
| self, | |
| sequence_length: int, | |
| device: torch.device, | |
| dtype: torch.dtype, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| cosine = self.cosine[:sequence_length].to(device=device, dtype=dtype) | |
| sine = self.sine[:sequence_length].to(device=device, dtype=dtype) | |
| return cosine[None, None, :, :], sine[None, None, :, :] | |
| class CherryAlphaAttention(nn.Module): | |
| def __init__(self, config: CherryAlphaConfig): | |
| super().__init__() | |
| self.num_heads = config.num_attention_heads | |
| self.num_kv_heads = config.num_key_value_heads | |
| self.head_dim = config.head_dim | |
| self.kv_repeats = self.num_heads // self.num_kv_heads | |
| hidden = config.hidden_size | |
| self.q_proj = nn.Linear(hidden, self.num_heads * self.head_dim, bias=False) | |
| self.k_proj = nn.Linear(hidden, self.num_kv_heads * self.head_dim, bias=False) | |
| self.v_proj = nn.Linear(hidden, self.num_kv_heads * self.head_dim, bias=False) | |
| self.o_proj = nn.Linear(self.num_heads * self.head_dim, hidden, bias=False) | |
| self.q_norm = CherryAlphaRMSNorm(self.head_dim, config.rms_norm_eps) | |
| self.k_norm = CherryAlphaRMSNorm(self.head_dim, config.rms_norm_eps) | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| cosine: torch.Tensor, | |
| sine: torch.Tensor, | |
| attention_mask: torch.Tensor | None, | |
| ) -> torch.Tensor: | |
| batch, length, _ = hidden_states.shape | |
| query = self.q_proj(hidden_states).view( | |
| batch, length, self.num_heads, self.head_dim | |
| ).transpose(1, 2) | |
| key = self.k_proj(hidden_states).view( | |
| batch, length, self.num_kv_heads, self.head_dim | |
| ).transpose(1, 2) | |
| value = self.v_proj(hidden_states).view( | |
| batch, length, self.num_kv_heads, self.head_dim | |
| ).transpose(1, 2) | |
| query = self.q_norm(query) | |
| key = self.k_norm(key) | |
| query = query * cosine + rotate_half(query) * sine | |
| key = key * cosine + rotate_half(key) * sine | |
| repeated_key = key.repeat_interleave(self.kv_repeats, dim=1) | |
| repeated_value = value.repeat_interleave(self.kv_repeats, dim=1) | |
| if attention_mask is None or bool(torch.all(attention_mask)): | |
| attended = F.scaled_dot_product_attention( | |
| query, | |
| repeated_key, | |
| repeated_value, | |
| is_causal=True, | |
| ) | |
| else: | |
| # Do not combine a broadcast key-padding mask with is_causal=True. | |
| # Some fused CUDA SDPA kernels have produced non-finite outputs for | |
| # that combination during batched continuation scoring. Build the | |
| # complete allowed-attention mask explicitly for padded inference. | |
| causal_mask = torch.ones( | |
| (length, length), | |
| dtype=torch.bool, | |
| device=hidden_states.device, | |
| ).tril_() | |
| key_mask = attention_mask.to(dtype=torch.bool)[:, None, None, :] | |
| allowed_mask = causal_mask[None, None, :, :] & key_mask | |
| attended = F.scaled_dot_product_attention( | |
| query, | |
| repeated_key, | |
| repeated_value, | |
| attn_mask=allowed_mask, | |
| is_causal=False, | |
| ) | |
| # XSA removes the component parallel to the current-token value vector. | |
| grouped = attended.view( | |
| batch, | |
| self.num_kv_heads, | |
| self.kv_repeats, | |
| length, | |
| self.head_dim, | |
| ) | |
| current_value = value.unsqueeze(2) | |
| numerator = (grouped.float() * current_value.float()).sum( | |
| dim=-1, | |
| keepdim=True, | |
| ) | |
| denominator = current_value.float().square().sum( | |
| dim=-1, | |
| keepdim=True, | |
| ).clamp_min(1e-6) | |
| grouped = grouped - (numerator / denominator).to(grouped.dtype) * current_value | |
| attended = grouped.reshape(batch, self.num_heads, length, self.head_dim) | |
| attended = attended.transpose(1, 2).contiguous().view(batch, length, -1) | |
| return self.o_proj(attended) | |
| class CherryAlphaRefreshGate(nn.Module): | |
| def __init__(self, config: CherryAlphaConfig): | |
| super().__init__() | |
| hidden = config.hidden_size | |
| self.kernel_size = config.refresh_kernel_size | |
| self.attention_norm = CherryAlphaRMSNorm(hidden, config.rms_norm_eps) | |
| self.embedding_norm = CherryAlphaRMSNorm(hidden, config.rms_norm_eps) | |
| self.output_norm = CherryAlphaRMSNorm(hidden, config.rms_norm_eps) | |
| self.gate_proj = nn.Linear(hidden, hidden, bias=False) | |
| self.value_proj = nn.Linear(hidden, hidden, bias=False) | |
| self.out_proj = nn.Linear(hidden, hidden, bias=False) | |
| # Stored as a matrix so stock torch.optim.Muon can optimize it. | |
| self.depthwise_kernel = nn.Parameter( | |
| torch.empty(hidden, config.refresh_kernel_size) | |
| ) | |
| self.alpha = nn.Parameter(torch.zeros(())) | |
| nn.init.normal_(self.depthwise_kernel, mean=0.0, std=config.initializer_range) | |
| def forward( | |
| self, | |
| attention_output: torch.Tensor, | |
| original_embedding: torch.Tensor, | |
| ) -> torch.Tensor: | |
| signal = self.attention_norm(attention_output.detach()) | |
| convolution = F.conv1d( | |
| F.pad(signal.transpose(1, 2), (self.kernel_size - 1, 0)), | |
| self.depthwise_kernel.unsqueeze(1), | |
| groups=signal.shape[-1], | |
| ).transpose(1, 2) | |
| gate = self.gate_proj(signal) + convolution | |
| value = self.value_proj(self.embedding_norm(original_embedding)) | |
| refreshed = self.output_norm(self.out_proj(F.silu(gate) * value)) | |
| return self.alpha * refreshed | |
| class CherryAlphaMLP(nn.Module): | |
| def __init__(self, config: CherryAlphaConfig): | |
| super().__init__() | |
| self.gate_proj = nn.Linear( | |
| config.hidden_size, | |
| config.intermediate_size, | |
| bias=False, | |
| ) | |
| self.up_proj = nn.Linear( | |
| config.hidden_size, | |
| config.intermediate_size, | |
| bias=False, | |
| ) | |
| self.down_proj = nn.Linear( | |
| config.intermediate_size, | |
| config.hidden_size, | |
| bias=False, | |
| ) | |
| def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: | |
| return self.down_proj( | |
| F.silu(self.gate_proj(hidden_states)) * self.up_proj(hidden_states) | |
| ) | |
| class CherryAlphaBlock(nn.Module): | |
| def __init__(self, config: CherryAlphaConfig): | |
| super().__init__() | |
| self.attention_norm = CherryAlphaRMSNorm( | |
| config.hidden_size, | |
| config.rms_norm_eps, | |
| ) | |
| self.attention = CherryAlphaAttention(config) | |
| self.refresh = CherryAlphaRefreshGate(config) | |
| self.mlp_norm = CherryAlphaRMSNorm(config.hidden_size, config.rms_norm_eps) | |
| self.mlp = CherryAlphaMLP(config) | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| original_embedding: torch.Tensor, | |
| cosine: torch.Tensor, | |
| sine: torch.Tensor, | |
| attention_mask: torch.Tensor | None, | |
| ) -> torch.Tensor: | |
| attention_output = self.attention( | |
| self.attention_norm(hidden_states), | |
| cosine, | |
| sine, | |
| attention_mask, | |
| ) | |
| hidden_states = hidden_states + attention_output | |
| hidden_states = hidden_states + self.refresh( | |
| attention_output, | |
| original_embedding, | |
| ) | |
| return hidden_states + self.mlp(self.mlp_norm(hidden_states)) | |
| class CherryAlphaNGramEmbedding(nn.Module): | |
| """Forty-million-parameter hashed causal trigram embedding.""" | |
| def __init__(self, config: CherryAlphaConfig): | |
| super().__init__() | |
| self.buckets = config.ngram_buckets | |
| self.table = nn.Embedding(config.ngram_buckets, config.hidden_size) | |
| self.scale = nn.Parameter(torch.ones(())) | |
| def forward(self, input_ids: torch.Tensor) -> torch.Tensor: | |
| previous_1 = F.pad(input_ids[:, :-1], (1, 0), value=0) | |
| previous_2 = F.pad(input_ids[:, :-2], (2, 0), value=0) | |
| hashes = ((previous_2 * 2053 + previous_1) * 2053 + input_ids) % self.buckets | |
| return self.scale * self.table(hashes) | |
| class CherryAlphaPreTrainedModel(PreTrainedModel): | |
| config_class = CherryAlphaConfig | |
| base_model_prefix = "transformer" | |
| supports_gradient_checkpointing = True | |
| _no_split_modules = ["CherryAlphaBlock"] | |
| _supports_sdpa = True | |
| def _init_weights(self, module: nn.Module): | |
| if isinstance(module, (nn.Linear, nn.Embedding)): | |
| nn.init.normal_( | |
| module.weight, | |
| mean=0.0, | |
| std=self.config.initializer_range, | |
| ) | |
| elif isinstance(module, CherryAlphaRMSNorm): | |
| nn.init.ones_(module.weight) | |
| class CherryAlphaForCausalLM(CherryAlphaPreTrainedModel, GenerationMixin): | |
| _tied_weights_keys = {"lm_head.weight": "transformer.wte.weight"} | |
| def __init__(self, config: CherryAlphaConfig): | |
| super().__init__(config) | |
| self.transformer = nn.ModuleDict( | |
| { | |
| "wte": nn.Embedding(config.vocab_size, config.hidden_size), | |
| "ngram": CherryAlphaNGramEmbedding(config), | |
| "h": nn.ModuleList( | |
| CherryAlphaBlock(config) | |
| for _ in range(config.num_hidden_layers) | |
| ), | |
| "ln_f": CherryAlphaRMSNorm( | |
| config.hidden_size, | |
| config.rms_norm_eps, | |
| ), | |
| } | |
| ) | |
| self.rotary_embedding = CherryAlphaRotaryEmbedding(config) | |
| self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) | |
| self.embedding_scale = math.sqrt(config.hidden_size) | |
| self.gradient_checkpointing = config.gradient_checkpointing | |
| self.post_init() | |
| if config.tie_word_embeddings: | |
| self.tie_weights() | |
| def get_input_embeddings(self): | |
| return self.transformer["wte"] | |
| def set_input_embeddings(self, value): | |
| self.transformer["wte"] = value | |
| def get_output_embeddings(self): | |
| return self.lm_head | |
| def set_output_embeddings(self, value): | |
| self.lm_head = value | |
| def prepare_inputs_for_generation( | |
| self, | |
| input_ids, | |
| attention_mask=None, | |
| **kwargs, | |
| ): | |
| return { | |
| "input_ids": input_ids, | |
| "attention_mask": attention_mask, | |
| "use_cache": False, | |
| } | |
| def _run_block( | |
| self, | |
| block: CherryAlphaBlock, | |
| hidden_states: torch.Tensor, | |
| original_embedding: torch.Tensor, | |
| cosine: torch.Tensor, | |
| sine: torch.Tensor, | |
| attention_mask: torch.Tensor | None, | |
| ) -> torch.Tensor: | |
| if self.training and self.gradient_checkpointing: | |
| return checkpoint( | |
| block, | |
| hidden_states, | |
| original_embedding, | |
| cosine, | |
| sine, | |
| attention_mask, | |
| use_reentrant=False, | |
| ) | |
| return block( | |
| hidden_states, | |
| original_embedding, | |
| cosine, | |
| sine, | |
| attention_mask, | |
| ) | |
| def hidden_states( | |
| self, | |
| input_ids: torch.LongTensor, | |
| attention_mask: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| if input_ids.ndim != 2: | |
| raise ValueError("input_ids must have shape [batch, sequence]") | |
| if input_ids.shape[1] > self.config.max_position_embeddings: | |
| raise ValueError("Sequence exceeds the 4,096-token context") | |
| original_embedding = ( | |
| self.transformer["wte"](input_ids) * self.embedding_scale | |
| ) | |
| ngram = self.transformer["ngram"](input_ids) | |
| cosine, sine = self.rotary_embedding( | |
| input_ids.shape[1], | |
| input_ids.device, | |
| original_embedding.dtype, | |
| ) | |
| hidden = original_embedding | |
| for layer_index in self.config.loop_schedule: | |
| # Match the original partial-Edu design: inject the hashed trigram | |
| # stream before both visits to the shared physical L2 block. | |
| if layer_index == 1: | |
| hidden = hidden + ngram | |
| hidden = self._run_block( | |
| self.transformer["h"][layer_index], | |
| hidden, | |
| original_embedding, | |
| cosine, | |
| sine, | |
| attention_mask, | |
| ) | |
| return self.transformer["ln_f"](hidden) | |
| def _chunked_training_losses( | |
| self, | |
| hidden_states: torch.Tensor, | |
| shifted_labels: torch.Tensor, | |
| z_loss_coefficient: float, | |
| chunk_tokens: int, | |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| hidden = hidden_states.reshape(-1, hidden_states.shape[-1]) | |
| labels = shifted_labels.reshape(-1) | |
| if hidden.shape[0] != labels.shape[0]: | |
| raise ValueError("shifted_labels must match the input shape") | |
| ce_sum = hidden.new_zeros((), dtype=torch.float32) | |
| z_sum = hidden.new_zeros((), dtype=torch.float32) | |
| for start in range(0, labels.numel(), chunk_tokens): | |
| stop = min(start + chunk_tokens, labels.numel()) | |
| chunk_labels = labels[start:stop] | |
| logits = F.linear(hidden[start:stop], self.lm_head.weight).float() | |
| ce_sum = ce_sum + F.cross_entropy( | |
| logits, | |
| chunk_labels, | |
| ignore_index=-100, | |
| reduction="sum", | |
| ) | |
| if z_loss_coefficient: | |
| z_sum = z_sum + torch.logsumexp(logits, dim=-1).square().sum() | |
| ce_loss = ce_sum / labels.numel() | |
| z_loss = z_sum / labels.numel() if z_loss_coefficient else z_sum | |
| loss = ce_loss + z_loss_coefficient * z_loss | |
| return loss, ce_loss.detach(), z_loss.detach() | |
| def forward( | |
| self, | |
| input_ids: torch.LongTensor, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| labels: Optional[torch.LongTensor] = None, | |
| shifted_labels: Optional[torch.LongTensor] = None, | |
| return_training_losses: bool = False, | |
| z_loss_coefficient: float = 0.0, | |
| loss_chunk_tokens: int = 16384, | |
| use_cache: Optional[bool] = None, | |
| **kwargs, | |
| ): | |
| hidden = self.hidden_states(input_ids, attention_mask=attention_mask) | |
| if return_training_losses: | |
| if shifted_labels is None: | |
| raise ValueError("shifted_labels are required for training losses") | |
| return self._chunked_training_losses( | |
| hidden, | |
| shifted_labels, | |
| z_loss_coefficient, | |
| loss_chunk_tokens, | |
| ) | |
| logits = self.lm_head(hidden).float() | |
| loss = None | |
| if labels is not None: | |
| loss = F.cross_entropy( | |
| logits[:, :-1].reshape(-1, logits.shape[-1]), | |
| labels[:, 1:].reshape(-1), | |
| ignore_index=-100, | |
| ) | |
| return CausalLMOutputWithPast( | |
| loss=loss, | |
| logits=logits, | |
| past_key_values=None, | |
| ) | |
| CherryAlphaForCausalLM.register_for_auto_class("AutoModelForCausalLM") | |
| __all__ = [ | |
| "CherryAlphaConfig", | |
| "CherryAlphaForCausalLM", | |
| "CherryAlphaPreTrainedModel", | |
| ] | |