"""Hugging Face Transformers model for the Cubit KRR token mixer.""" from __future__ import annotations import math from typing import Optional, Union import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from transformers.generation import GenerationMixin from transformers.modeling_outputs import BaseModelOutput, CausalLMOutput from .configuration_cubit import CubitConfig def _batch_matmul(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor: """Multiply ``[B, H, M, K]`` by ``[B, H, K, N]``.""" batch_size, n_heads, _, inner_dim = left.shape if right.shape[:2] != (batch_size, n_heads) or right.shape[2] != inner_dim: raise RuntimeError("incompatible batched matrix multiplication shapes") return torch.matmul(left, right) def _score_mask( *, batch_size: int, q_start: int, q_end: int, k_start: int, k_end: int, device: torch.device, window_size: Optional[int], ) -> torch.Tensor: """Build one causal (and optionally sliding-window) mask tile.""" query_positions = torch.arange(q_start, q_end, device=device) key_positions = torch.arange(k_start, k_end, device=device) mask = key_positions[None, :] <= query_positions[:, None] if window_size is not None: mask = mask & ( key_positions[None, :] >= query_positions[:, None] - (window_size - 1) ) return mask.view(1, 1, q_end - q_start, k_end - k_start).expand( batch_size, 1, -1, -1 ) def _kernel_weights( reference_queries: torch.Tensor, reference_keys: torch.Tensor, *, q_start: int, k_start: int, window_size: Optional[int], ) -> tuple[torch.Tensor, torch.Tensor]: """Return one causal Cubit kernel tile and its row log-normalizers.""" q_end = q_start + reference_queries.shape[2] k_end = k_start + reference_keys.shape[2] scores = _batch_matmul(reference_queries, reference_keys.transpose(-2, -1)) mask = _score_mask( batch_size=reference_queries.shape[0], q_start=q_start, q_end=q_end, k_start=k_start, k_end=k_end, device=reference_queries.device, window_size=window_size, ) scores = scores.masked_fill(~mask, -torch.inf) logsumexp = torch.logsumexp(scores, dim=-1) weights = torch.exp(scores - logsumexp.unsqueeze(-1)).masked_fill(~mask, 0.0) return weights, logsumexp def _first_key_for_block(q_start: int, window_size: Optional[int]) -> int: if window_size is None: return 0 return max(0, q_start - window_size + 1) def streaming_causal_krr_solve( reference_queries: torch.Tensor, reference_keys: torch.Tensor, rhs: torch.Tensor, regularization: torch.Tensor, *, window_size: Optional[int], block_size: int, ) -> torch.Tensor: """Run the exact block-streaming forward solve used by the OLMo Core Cubit model.""" reference_queries = reference_queries.contiguous() reference_keys = reference_keys.contiguous() rhs = rhs.contiguous() batch_size, n_heads, seq_len, _ = rhs.shape solution = torch.empty_like(rhs) for q_start in range(0, seq_len, block_size): q_end = min(q_start + block_size, seq_len) key_start = _first_key_for_block(q_start, window_size) weights, _ = _kernel_weights( reference_queries[:, :, q_start:q_end], reference_keys[:, :, key_start:q_end], q_start=q_start, k_start=key_start, window_size=window_size, ) previous_end = q_start - key_start residual = rhs[:, :, q_start:q_end] if previous_end > 0: residual = residual - _batch_matmul( weights[:, :, :, :previous_end], solution[:, :, key_start:q_start], ) diagonal_block = weights[:, :, :, previous_end:] block_len = q_end - q_start identity = torch.eye( block_len, dtype=diagonal_block.dtype, device=diagonal_block.device, ).view(1, 1, block_len, block_len) diagonal_block = diagonal_block + regularization.view(1, n_heads, 1, 1) * identity solution[:, :, q_start:q_end] = torch.linalg.solve_triangular( diagonal_block, residual, upper=False, ) return solution class CubitRMSNorm(nn.Module): """RMSNorm with the same fp32 calculation and cast order as OLMo Core.""" def __init__(self, hidden_size: int, eps: float) -> None: super().__init__() self.weight = nn.Parameter(torch.ones(hidden_size)) self.variance_epsilon = eps def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: input_dtype = hidden_states.dtype hidden_states = hidden_states.float() variance = hidden_states.pow(2).mean(-1, keepdim=True) hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) # OLMo Core applies the affine weight while the normalized activation is # still fp32, and only then casts the result back to the input dtype. # Reversing these two operations compounds BF16 rounding error. hidden_states = self.weight.type_as(hidden_states) * hidden_states return hidden_states.to(input_dtype) class CubitRotaryEmbedding(nn.Module): """Full-precision RoPE matching OLMo Core's real-valued implementation.""" def __init__(self, config: CubitConfig) -> None: super().__init__() self.dim = config.head_dim self.theta = config.rope_theta @staticmethod def _rotate_half(hidden_states: torch.Tensor) -> torch.Tensor: batch_size, seq_len, n_heads, head_dim = hidden_states.shape hidden_states = hidden_states.view(batch_size, seq_len, n_heads, 2, head_dim // 2) first, second = hidden_states.unbind(dim=-2) return torch.cat((-second, first), dim=-1) def _get_sin_cos( self, seq_len: int, device: torch.device, ) -> tuple[torch.Tensor, torch.Tensor]: with torch.autocast(device.type, enabled=False): inv_freq = 1.0 / ( self.theta ** (torch.arange(0, self.dim, 2, device=device, dtype=torch.float) / self.dim) ) positions = torch.arange(seq_len, device=device, dtype=torch.float) frequencies = torch.einsum("i , j -> i j", positions, inv_freq) embeddings = torch.cat((frequencies, frequencies), dim=-1) return embeddings.sin(), embeddings.cos() def forward( self, hidden_states: torch.Tensor, position_ids: torch.Tensor, ) -> torch.Tensor: input_dtype = hidden_states.dtype hidden_states = hidden_states.float() max_position = int(position_ids.max().item()) + 1 pos_sin, pos_cos = self._get_sin_cos(max_position, hidden_states.device) pos_sin = pos_sin[position_ids].unsqueeze(2).type_as(hidden_states) pos_cos = pos_cos[position_ids].unsqueeze(2).type_as(hidden_states) with torch.autocast(hidden_states.device.type, enabled=False): hidden_states = (hidden_states * pos_cos) + ( self._rotate_half(hidden_states) * pos_sin ) return hidden_states.to(input_dtype) class CubitMLP(nn.Module): """SwiGLU feed-forward layer.""" def __init__(self, config: CubitConfig) -> None: super().__init__() self.gate_proj = nn.Linear( config.hidden_size, config.intermediate_size, bias=config.mlp_bias ) self.up_proj = nn.Linear( config.hidden_size, config.intermediate_size, bias=config.mlp_bias ) self.down_proj = nn.Linear( config.intermediate_size, config.hidden_size, bias=config.mlp_bias ) 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 CubitAttention(nn.Module): """Causal kernel-ridge-regression token mixer followed by output attention.""" def __init__(self, config: CubitConfig, layer_idx: int) -> None: super().__init__() self.config = config self.layer_idx = layer_idx self.n_heads = config.num_attention_heads self.head_dim = config.head_dim self.hidden_size = config.hidden_size self.window_size = ( config.sliding_window if config.layer_types[layer_idx] == "sliding_attention" else None ) self.share_reference = config.share_reference self.q_proj = nn.Linear( config.hidden_size, config.hidden_size, bias=config.attention_bias ) self.k_proj = nn.Linear( config.hidden_size, config.hidden_size, bias=config.attention_bias ) self.v_proj = nn.Linear( config.hidden_size, config.hidden_size, bias=config.attention_bias ) self.o_proj = nn.Linear( config.hidden_size, config.hidden_size, bias=config.attention_bias ) self.r_proj = ( None if config.share_reference else nn.Linear(config.hidden_size, config.hidden_size, bias=config.attention_bias) ) self.lrr_proj = nn.Linear( config.hidden_size, config.num_attention_heads, bias=config.attention_bias ) self.q_norm = CubitRMSNorm(config.hidden_size, config.rms_norm_eps) self.k_norm = CubitRMSNorm(config.hidden_size, config.rms_norm_eps) self.lrr_lower = nn.Parameter(torch.full((self.n_heads,), 0.5)) self.lrr_range = nn.Parameter(torch.full((self.n_heads,), 1.5)) self.reference_scale = nn.Parameter(torch.ones(self.n_heads)) self.log_regularization = nn.Parameter( torch.full((self.n_heads,), math.log(1e-10)) ) self.rotary_emb = CubitRotaryEmbedding(config) def _build_attention_mask( self, batch_size: int, seq_len: int, device: torch.device, attention_mask: Optional[torch.Tensor], ) -> torch.Tensor: positions = torch.arange(seq_len, device=device) query_pos = positions[:, None] key_pos = positions[None, :] mask = key_pos <= query_pos if self.window_size is not None: mask = mask & (key_pos >= query_pos - (self.window_size - 1)) mask = mask.unsqueeze(0).expand(batch_size, -1, -1) if attention_mask is not None: if attention_mask.ndim == 2: key_mask = attention_mask.to(device=device, dtype=torch.bool) mask = mask & key_mask[:, None, :] elif attention_mask.ndim == 4: layer_mask = attention_mask[:, 0] allowed = layer_mask if layer_mask.dtype == torch.bool else layer_mask >= 0 mask = mask & allowed.to(device=device) else: raise ValueError("attention_mask must have rank 2 or 4") # Keep fully padded query rows finite. Their outputs are ignored by standard LM scoring. empty_rows = ~mask.any(dim=-1) if empty_rows.any(): diagonal = torch.eye(seq_len, dtype=torch.bool, device=device).unsqueeze(0) mask = mask | (empty_rows.unsqueeze(-1) & diagonal) return mask @staticmethod def _masked_softmax(scores: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: return torch.softmax(scores.masked_fill(~mask[:, None, :, :], -torch.inf), dim=-1) def _dense_krr_solve( self, reference_queries: torch.Tensor, reference_keys: torch.Tensor, rhs: torch.Tensor, regularization: torch.Tensor, mask: torch.Tensor, ) -> torch.Tensor: seq_len = rhs.shape[2] inverse_sigma = self._masked_softmax( reference_queries @ reference_keys.transpose(-2, -1), mask ) identity = torch.eye(seq_len, device=rhs.device, dtype=torch.float32) inverse_sigma = inverse_sigma + regularization.view( 1, self.n_heads, 1, 1 ) * identity.view(1, 1, seq_len, seq_len) return torch.linalg.solve_triangular(inverse_sigma, rhs, upper=False) def forward( self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.Tensor] = None, **kwargs, ) -> torch.Tensor: if kwargs.get("past_key_values") is not None or kwargs.get("past_key_value") is not None: raise NotImplementedError("Cubit v1 does not support KV caching") batch_size, seq_len, _ = hidden_states.shape if position_ids is None: position_ids = torch.arange(seq_len, device=hidden_states.device).unsqueeze(0) position_ids = position_ids.expand(batch_size, -1) q = self.q_norm(self.q_proj(hidden_states)) k = self.k_norm(self.k_proj(hidden_states)) v = self.v_proj(hidden_states) projected_r = None if self.share_reference else self.r_proj(hidden_states) q = q.view(batch_size, seq_len, self.n_heads, self.head_dim) k = k.view(batch_size, seq_len, self.n_heads, self.head_dim) v = v.view(batch_size, seq_len, self.n_heads, self.head_dim) if self.share_reference: r = k else: assert projected_r is not None r = projected_r.view(batch_size, seq_len, self.n_heads, self.head_dim) reference_scale = self.reference_scale.float().view(1, 1, self.n_heads, 1) normalized_r = F.normalize( r.float(), p=2, dim=-1, eps=self.config.reference_norm_eps ) * reference_scale q = self.rotary_emb(q, position_ids) k = self.rotary_emb(k, position_ids) r = self.rotary_emb(r, position_ids) normalized_r = self.rotary_emb(normalized_r, position_ids) r_heads = r.float().transpose(1, 2) normalized_r_heads = normalized_r.float().transpose(1, 2) lrr_logits = self.lrr_proj(hidden_states).float().transpose(1, 2).unsqueeze(-1) lrr = self.lrr_lower.float().view(1, self.n_heads, 1, 1) lrr = lrr + self.lrr_range.float().view(1, self.n_heads, 1, 1) * torch.sigmoid( lrr_logits ) rhs = lrr * v.float().transpose(1, 2) regularization = self.log_regularization.float().exp() mask = self._build_attention_mask( batch_size, seq_len, hidden_states.device, attention_mask ) use_streaming = self.config.krr_implementation == "streaming" if attention_mask is not None and attention_mask.ndim == 2: use_streaming = use_streaming and bool(attention_mask.to(torch.bool).all()) if use_streaming: solution = streaming_causal_krr_solve( r_heads, normalized_r_heads, rhs, regularization, window_size=self.window_size, block_size=self.config.krr_block_size, ) else: solution = self._dense_krr_solve( r_heads, normalized_r_heads, rhs, regularization, mask ) solution = solution.transpose(1, 2).to(q.dtype).contiguous() scores = torch.einsum("bthd,bshd->bhts", q.float(), k.float()) * ( self.head_dim**-0.5 ) weights = self._masked_softmax(scores, mask) output = torch.einsum("bhts,bshd->bthd", weights, solution.float()).to(q.dtype) return self.o_proj(output.reshape(batch_size, seq_len, -1)) class CubitDecoderLayer(nn.Module): """OLMo 3 reordered-norm decoder block.""" def __init__(self, config: CubitConfig, layer_idx: int) -> None: super().__init__() self.self_attn = CubitAttention(config, layer_idx) self.mlp = CubitMLP(config) self.post_attention_layernorm = CubitRMSNorm(config.hidden_size, config.rms_norm_eps) self.post_feedforward_layernorm = CubitRMSNorm(config.hidden_size, config.rms_norm_eps) def forward( self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.Tensor] = None, **kwargs, ) -> torch.Tensor: hidden_states = hidden_states + self.post_attention_layernorm( self.self_attn( hidden_states, attention_mask=attention_mask, position_ids=position_ids, **kwargs, ) ) return hidden_states + self.post_feedforward_layernorm(self.mlp(hidden_states)) class CubitPreTrainedModel(PreTrainedModel): """Base class for Cubit Transformers models.""" config_class = CubitConfig base_model_prefix = "model" supports_gradient_checkpointing = False _no_split_modules = ["CubitDecoderLayer"] _supports_flash_attn = False _supports_sdpa = False _supports_cache_class = False def _init_weights(self, module: nn.Module) -> None: if isinstance(module, nn.Linear): module.weight.data.normal_(mean=0.0, std=self.config.initializer_range) if module.bias is not None: module.bias.data.zero_() elif isinstance(module, nn.Embedding): module.weight.data.normal_(mean=0.0, std=self.config.initializer_range) if module.padding_idx is not None: module.weight.data[module.padding_idx].zero_() class CubitModel(CubitPreTrainedModel): """Bare Cubit decoder model.""" def __init__(self, config: CubitConfig) -> None: super().__init__(config) self.padding_idx = config.pad_token_id self.embed_tokens = nn.Embedding( config.vocab_size, config.hidden_size, padding_idx=self.padding_idx ) self.layers = nn.ModuleList( [CubitDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)] ) self.norm = CubitRMSNorm(config.hidden_size, config.rms_norm_eps) self.post_init() def get_input_embeddings(self) -> nn.Embedding: return self.embed_tokens def set_input_embeddings(self, value: nn.Embedding) -> None: self.embed_tokens = value def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, inputs_embeds: Optional[torch.FloatTensor] = None, use_cache: Optional[bool] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, **kwargs, ) -> Union[tuple, BaseModelOutput]: if (input_ids is None) == (inputs_embeds is None): raise ValueError("provide exactly one of input_ids or inputs_embeds") if use_cache: raise NotImplementedError("Cubit v1 does not support KV caching") if output_attentions: raise NotImplementedError("Cubit v1 does not return attention matrices") output_hidden_states = ( self.config.output_hidden_states if output_hidden_states is None else output_hidden_states ) return_dict = self.config.use_return_dict if return_dict is None else return_dict hidden_states = self.embed_tokens(input_ids) if inputs_embeds is None else inputs_embeds batch_size, seq_len, _ = hidden_states.shape if position_ids is None: if attention_mask is not None and attention_mask.ndim == 2: position_ids = attention_mask.long().cumsum(-1) - 1 position_ids.masked_fill_(attention_mask == 0, 0) else: position_ids = torch.arange(seq_len, device=hidden_states.device).unsqueeze(0) position_ids = position_ids.expand(batch_size, -1) collected_hidden_states = () if output_hidden_states else None for decoder_layer in self.layers: if collected_hidden_states is not None: collected_hidden_states += (hidden_states,) hidden_states = decoder_layer( hidden_states, attention_mask=attention_mask, position_ids=position_ids, **kwargs, ) hidden_states = self.norm(hidden_states) if collected_hidden_states is not None: collected_hidden_states += (hidden_states,) if not return_dict: return tuple(value for value in (hidden_states, collected_hidden_states) if value is not None) return BaseModelOutput( last_hidden_state=hidden_states, hidden_states=collected_hidden_states, attentions=None, ) class CubitForCausalLM(CubitPreTrainedModel, GenerationMixin): """Cubit decoder with a causal language-modeling head.""" _tied_weights_keys: list[str] = [] def __init__(self, config: CubitConfig) -> None: super().__init__(config) self.model = CubitModel(config) self.vocab_size = config.vocab_size self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.post_init() def get_input_embeddings(self) -> nn.Embedding: return self.model.embed_tokens def set_input_embeddings(self, value: nn.Embedding) -> None: self.model.embed_tokens = value def get_output_embeddings(self) -> nn.Linear: return self.lm_head def set_output_embeddings(self, value: nn.Linear) -> None: self.lm_head = value def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, inputs_embeds: Optional[torch.FloatTensor] = None, labels: Optional[torch.LongTensor] = None, use_cache: Optional[bool] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, logits_to_keep: Union[int, torch.Tensor] = 0, **kwargs, ) -> Union[tuple, CausalLMOutput]: return_dict = self.config.use_return_dict if return_dict is None else return_dict outputs = self.model( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, inputs_embeds=inputs_embeds, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=True, **kwargs, ) hidden_states = outputs.last_hidden_state if isinstance(logits_to_keep, int): if logits_to_keep: hidden_states = hidden_states[:, -logits_to_keep:, :] else: hidden_states = hidden_states.gather( 1, logits_to_keep.unsqueeze(-1).expand(-1, -1, hidden_states.size(-1)) ) logits = self.lm_head(hidden_states) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous().float() shift_labels = labels[..., 1:].contiguous().to(shift_logits.device) loss = F.cross_entropy( shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1) ) if not return_dict: values = (logits, outputs.hidden_states, outputs.attentions) return ((loss,) + values) if loss is not None else values return CausalLMOutput( loss=loss, logits=logits, hidden_states=outputs.hidden_states, attentions=outputs.attentions, )