olmo3-1b-cubit-fp8 / modeling_cubit.py
arch-space's picture
Upload OLMo3-1B Cubit FP8 step 75080
60f465f verified
Raw
History Blame Contribute Delete
23.9 kB
"""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,
)