OxMini / src /oxmini /model.py
Shivam3002's picture
Publish trained OxMini checkpoint and measured model card
46144df verified
Raw
History Blame Contribute Delete
8.02 kB
"""OxMini hybrid KDA-lite/MLA-lite causal language model."""
from __future__ import annotations
from dataclasses import dataclass
import json
from pathlib import Path
from typing import Any
import torch
from torch import nn
from torch.nn import functional as F
from .attention_kda import KDALiteAttention
from .attention_mla import MLALiteAttention
from .config import OxMiniConfig
from .layers import RMSNorm, SwiGLU
from .mhc import MHCResidual, StreamCollapse
@dataclass
class CausalLMOutput:
logits: torch.Tensor
loss: torch.Tensor | None = None
class OxMiniBlock(nn.Module):
def __init__(self, config: OxMiniConfig, attention_type: str) -> None:
super().__init__()
self.use_mhc = config.use_mhc
self.norm_attn = RMSNorm(config.n_embd, config.rms_norm_eps)
self.norm_mlp = RMSNorm(config.n_embd, config.rms_norm_eps)
if attention_type == "kda":
self.attention = KDALiteAttention(
config.n_embd, config.n_head, config.dropout, config.bias
)
elif attention_type == "mla":
self.attention = MLALiteAttention(
config.n_embd,
config.n_head,
config.mla_latent_dim,
config.dropout,
config.bias,
)
else:
raise ValueError(f"unknown attention type: {attention_type}")
self.mlp = SwiGLU(
config.n_embd,
config.n_embd * config.ffn_multiplier,
config.dropout,
config.bias,
)
if self.use_mhc:
self.attn_residual = MHCResidual(config.hc_streams, config.use_sinkhorn_mhc)
self.mlp_residual = MHCResidual(config.hc_streams, config.use_sinkhorn_mhc)
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.use_mhc:
# Attention and MLP each have an independent routing matrix, just as
# a standard pre-norm block has two independent residual additions.
x = self.attn_residual(x, lambda value: self.attention(self.norm_attn(value)))
return self.mlp_residual(x, lambda value: self.mlp(self.norm_mlp(value)))
x = x + self.attention(self.norm_attn(x))
return x + self.mlp(self.norm_mlp(x))
class OxMiniForCausalLM(nn.Module):
def __init__(self, config: OxMiniConfig) -> None:
super().__init__()
self.config = config
self.token_embedding = nn.Embedding(config.vocab_size, config.n_embd)
self.blocks = nn.ModuleList(
[OxMiniBlock(config, attention_type) for attention_type in config.layer_types]
)
self.collapse = StreamCollapse(config.hc_streams) if config.use_mhc else nn.Identity()
self.final_norm = RMSNorm(config.n_embd, config.rms_norm_eps)
self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
self.apply(self._init_weights)
if config.tie_embeddings:
self.lm_head.weight = self.token_embedding.weight
@staticmethod
def _init_weights(module: nn.Module) -> None:
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
@property
def num_parameters(self) -> int:
return sum(parameter.numel() for parameter in self.parameters())
def forward(
self,
input_ids: torch.Tensor,
targets: torch.Tensor | None = None,
) -> CausalLMOutput:
if input_ids.ndim != 2:
raise ValueError("input_ids must have shape [batch, sequence]")
if input_ids.shape[1] > self.config.block_size:
raise ValueError(
f"sequence length {input_ids.shape[1]} exceeds block_size {self.config.block_size}"
)
x = self.token_embedding(input_ids)
if self.config.use_mhc:
# Broadcast, do not concatenate: every stream starts as the same
# token representation and subsequently diverges through learned
# per-sublayer post-routing coefficients.
x = x.unsqueeze(2).expand(-1, -1, self.config.hc_streams, -1)
for block in self.blocks:
x = block(x)
x = self.collapse(x)
logits = self.lm_head(self.final_norm(x))
loss = None
if targets is not None:
# Every position predicts the next character supplied in ``targets``;
# data batching performs the one-token shift before this call.
loss = F.cross_entropy(logits.reshape(-1, logits.shape[-1]), targets.reshape(-1))
return CausalLMOutput(logits=logits, loss=loss)
@torch.no_grad()
def generate(
self,
input_ids: torch.Tensor,
max_new_tokens: int,
temperature: float = 1.0,
top_k: int | None = None,
generator: torch.Generator | None = None,
) -> torch.Tensor:
if input_ids.ndim != 2 or input_ids.shape[1] == 0:
raise ValueError("input_ids must be a non-empty [batch, sequence] tensor")
was_training = self.training
self.eval()
generated = input_ids
for _ in range(max_new_tokens):
# This correctness-first implementation recomputes the cropped
# context each step. It does not claim a production KV/state cache.
context = generated[:, -self.config.block_size :]
logits = self(context).logits[:, -1, :]
if not torch.isfinite(logits).all():
raise FloatingPointError("non-finite logits encountered during generation")
if temperature <= 0:
next_token = logits.argmax(dim=-1, keepdim=True)
else:
logits = logits / temperature
if top_k is not None:
k = min(top_k, logits.shape[-1])
cutoff = torch.topk(logits, k).values[:, [-1]]
logits = logits.masked_fill(logits < cutoff, float("-inf"))
probabilities = torch.softmax(logits, dim=-1)
next_token = torch.multinomial(probabilities, 1, generator=generator)
generated = torch.cat((generated, next_token), dim=1)
if was_training:
self.train()
return generated
def save_pretrained(self, directory: str | Path) -> Path:
from safetensors.torch import save_file
directory = Path(directory)
directory.mkdir(parents=True, exist_ok=True)
values: dict[str, Any] = self.config.to_dict()
values.update({"architectures": [self.__class__.__name__], "model_type": "oxmini"})
with (directory / "config.json").open("w", encoding="utf-8") as handle:
json.dump(values, handle, indent=2, sort_keys=True)
handle.write("\n")
# Clone tied tensors so safetensors sees independent storage for both
# state-dict keys while preserving strict-load compatibility.
state = {
key: value.detach().cpu().clone().contiguous()
for key, value in self.state_dict().items()
}
save_file(state, str(directory / "pytorch_model.safetensors"))
return directory
@classmethod
def from_pretrained(
cls,
model_id_or_path: str | Path,
map_location: str | torch.device = "cpu",
revision: str | None = None,
) -> "OxMiniForCausalLM":
from safetensors.torch import load_file
path = Path(model_id_or_path)
if not path.exists():
from huggingface_hub import snapshot_download
path = Path(snapshot_download(str(model_id_or_path), revision=revision))
config = OxMiniConfig.from_file(path / "config.json")
model = cls(config)
state = load_file(str(path / "pytorch_model.safetensors"), device=str(map_location))
model.load_state_dict(state)
return model.to(map_location)