Download src/oxmini/model.py from Shivam3002/OxMini: direct link, hf CLI and curl.
- Browser
- Download file 8.02 kB
-
https://huggingface.co/Shivam3002/OxMini/resolve/main/src/oxmini/model.py
- Command line
-
hf download hf://Shivam3002/OxMini/src/oxmini/model.py
-
curl -L -o model.py https://huggingface.co/Shivam3002/OxMini/resolve/main/src/oxmini/model.py
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 | |
| 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 | |
| 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) | |
| 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) | |
| 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 | |
| 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) | |