"""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", ]