cherry-alpha / modeling_cherry_alpha.py
Banaxi-Tech's picture
Initialize Cherry Alpha
426ab3f verified
Raw History Blame Contribute Delete
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",
]