"""CORTEX model definition — the trainable Frankenstein-Labs architecture. This module defines the model the CORTEX training pipeline builds and trains. It is a standard decoder-only Transformer with: * RMSNorm (pre-norm) * rotary position embeddings (RoPE) * grouped-query attention (GQA) * a SwiGLU feed-forward network Nothing here loads, mutates or depends on the distributed 1.65T checkpoint that this repository also hosts. Weights produced from this module are initialised from scratch and are owned by Frankenstein-Labs. """ from __future__ import annotations import json import math from dataclasses import asdict, dataclass from pathlib import Path import torch import torch.nn as nn import torch.nn.functional as F __all__ = ["CortexConfig", "CortexForCausalLM", "count_parameters"] @dataclass class CortexConfig: """Hyper-parameters of a trainable CORTEX model.""" vocab_size: int = 129280 hidden_size: int = 768 num_hidden_layers: int = 12 num_attention_heads: int = 12 num_key_value_heads: int = 4 intermediate_size: int = 2048 hidden_act: str = "silu" max_position_embeddings: int = 2048 rope_theta: float = 10000.0 rms_norm_eps: float = 1e-5 attention_bias: bool = False attention_dropout: float = 0.0 tie_word_embeddings: bool = True initializer_range: float = 0.02 bos_token_id: int = 0 eos_token_id: int = 1 pad_token_id: int = 2 model_name: str = "cortex-dev-1" @property def head_dim(self) -> int: if self.hidden_size % self.num_attention_heads: raise ValueError( f"hidden_size ({self.hidden_size}) must be divisible by " f"num_attention_heads ({self.num_attention_heads})" ) return self.hidden_size // self.num_attention_heads def validate(self) -> None: if self.num_attention_heads % self.num_key_value_heads: raise ValueError( f"num_attention_heads ({self.num_attention_heads}) must be a multiple of " f"num_key_value_heads ({self.num_key_value_heads})" ) if self.hidden_size % 2: raise ValueError("hidden_size must be even so RoPE can split the head dimension") if self.head_dim % 2: raise ValueError("head_dim must be even for RoPE") if self.vocab_size <= 0: raise ValueError("vocab_size must be positive") if self.num_hidden_layers <= 0: raise ValueError("num_hidden_layers must be positive") @classmethod def from_json(cls, path: str | Path) -> "CortexConfig": raw = json.loads(Path(path).read_text(encoding="utf-8")) fields = set(cls.__dataclass_fields__) return cls(**{k: v for k, v in raw.items() if k in fields}) def to_json(self, path: str | Path) -> None: Path(path).write_text( json.dumps(asdict(self), indent=2, ensure_ascii=False) + "\n", encoding="utf-8" ) def parameter_breakdown(self) -> dict: """Analytical parameter count, including norms and (tied) embeddings.""" h = self.hidden_size kv = self.num_key_value_heads * self.head_dim q = self.num_attention_heads * self.head_dim attn = h * q + h * kv * 2 + q * h if self.attention_bias: attn += q + kv * 2 mlp = 3 * h * self.intermediate_size norms = 2 * h per_layer = attn + mlp + norms embeds = self.vocab_size * h if self.tie_word_embeddings else self.vocab_size * h * 2 total = embeds + per_layer * self.num_hidden_layers + h return { "embedding_and_head_shared": self.vocab_size * h, "per_layer_attention": attn, "per_layer_mlp": mlp, "per_layer_norms": norms, "total_per_layer": per_layer, "total_estimated": total, } class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-5): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: dtype = x.dtype x = x.float() x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) return x.to(dtype) * self.weight def build_rope_cache(head_dim: int, max_position_embeddings: int, theta: float, device, dtype): """Precompute cos/sin tables of shape [max_pos, head_dim].""" inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim)) positions = torch.arange(max_position_embeddings, device=device).float() freqs = torch.outer(positions, inv_freq) emb = torch.cat((freqs, freqs), dim=-1) return emb.cos().to(dtype), emb.sin().to(dtype) def rotate_half(x: torch.Tensor) -> torch.Tensor: half = x.shape[-1] // 2 x1, x2 = x[..., :half], x[..., half:] return torch.cat((-x2, x1), dim=-1) def apply_rope(q, k, cos, sin): cos = cos.unsqueeze(0).unsqueeze(0) sin = sin.unsqueeze(0).unsqueeze(0) return q * cos + rotate_half(q) * sin, k * cos + rotate_half(k) * sin class CortexAttention(nn.Module): def __init__(self, config: CortexConfig): super().__init__() self.config = config self.num_heads = config.num_attention_heads self.num_kv_heads = config.num_key_value_heads self.head_dim = config.head_dim self.num_kv_groups = self.num_heads // self.num_kv_heads self.q_proj = nn.Linear( config.hidden_size, self.num_heads * self.head_dim, bias=config.attention_bias ) self.k_proj = nn.Linear( config.hidden_size, self.num_kv_heads * self.head_dim, bias=config.attention_bias ) self.v_proj = nn.Linear( config.hidden_size, self.num_kv_heads * self.head_dim, bias=config.attention_bias ) self.o_proj = nn.Linear( self.num_heads * self.head_dim, config.hidden_size, bias=config.attention_bias ) self.attention_dropout = config.attention_dropout def forward(self, hidden_states, cos, sin, attention_mask=None): batch, seq, _ = hidden_states.shape q = self.q_proj(hidden_states).view(batch, seq, self.num_heads, self.head_dim) k = self.k_proj(hidden_states).view(batch, seq, self.num_kv_heads, self.head_dim) v = self.v_proj(hidden_states).view(batch, seq, self.num_kv_heads, self.head_dim) q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2) q, k = apply_rope(q, k, cos, sin) if self.num_kv_groups > 1: k = k.repeat_interleave(self.num_kv_groups, dim=1) v = v.repeat_interleave(self.num_kv_groups, dim=1) dropout_p = self.attention_dropout if self.training else 0.0 attn = F.scaled_dot_product_attention( q, k, v, attn_mask=attention_mask, dropout_p=dropout_p, is_causal=attention_mask is None, ) attn = attn.transpose(1, 2).reshape(batch, seq, self.num_heads * self.head_dim) return self.o_proj(attn) class CortexMLP(nn.Module): def __init__(self, config: CortexConfig): 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) self.act = nn.SiLU() if config.hidden_act == "silu" else nn.GELU() def forward(self, x): return self.down_proj(self.act(self.gate_proj(x)) * self.up_proj(x)) class CortexDecoderLayer(nn.Module): def __init__(self, config: CortexConfig): super().__init__() self.self_attn = CortexAttention(config) self.mlp = CortexMLP(config) self.input_layernorm = RMSNorm(config.hidden_size, config.rms_norm_eps) self.post_attention_layernorm = RMSNorm(config.hidden_size, config.rms_norm_eps) def forward(self, hidden_states, cos, sin, attention_mask=None): residual = hidden_states hidden_states = self.self_attn( self.input_layernorm(hidden_states), cos, sin, attention_mask ) hidden_states = residual + hidden_states residual = hidden_states hidden_states = self.mlp(self.post_attention_layernorm(hidden_states)) return residual + hidden_states class CortexForCausalLM(nn.Module): """Decoder-only causal language model, the CORTEX training target.""" def __init__(self, config: CortexConfig): super().__init__() config.validate() self.config = config self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) self.layers = nn.ModuleList( [CortexDecoderLayer(config) for _ in range(config.num_hidden_layers)] ) self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) if config.tie_word_embeddings: self.lm_head.weight = self.embed_tokens.weight self.apply(self._init_weights) for name, param in self.named_parameters(): if name.endswith(("o_proj.weight", "down_proj.weight")): nn.init.normal_( param, mean=0.0, std=config.initializer_range / math.sqrt(2 * config.num_hidden_layers), ) self._rope_cache = None def _init_weights(self, module): if isinstance(module, nn.Linear): nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range) def _rope(self, seq_len, device, dtype): if ( self._rope_cache is None or self._rope_cache[0].shape[0] < seq_len or self._rope_cache[0].device != device ): self._rope_cache = build_rope_cache( self.config.head_dim, self.config.max_position_embeddings, self.config.rope_theta, device, dtype, ) cos, sin = self._rope_cache return cos[:seq_len], sin[:seq_len] def forward(self, input_ids, attention_mask=None, labels=None): batch, seq = input_ids.shape hidden_states = self.embed_tokens(input_ids) cos, sin = self._rope(seq, input_ids.device, hidden_states.dtype) causal = None if attention_mask is not None: causal = torch.tril( torch.ones(seq, seq, dtype=torch.bool, device=input_ids.device) )[None, None, :, :] causal = causal & attention_mask[:, None, None, :].bool() for layer in self.layers: hidden_states = layer(hidden_states, cos, sin, causal) logits = self.lm_head(self.norm(hidden_states)) loss = None if labels is not None: loss = self._loss(logits, labels) return {"logits": logits, "loss": loss} @staticmethod def _loss(logits: torch.Tensor, labels: torch.Tensor) -> torch.Tensor: """Cross-entropy computed in chunks over the vocabulary. A full fp32 logits tensor of [batch*seq, vocab] is several hundred megabytes for the CORTEX vocabulary, so the softmax is done in slices. """ shift_logits = logits[:, :-1] shift_labels = labels[:, 1:] total, count = 0.0, 0 chunk = 8192 flat_logits = shift_logits.reshape(-1, shift_logits.size(-1)) flat_labels = shift_labels.reshape(-1) for start in range(0, flat_labels.size(0), chunk): sl = flat_logits[start:start + chunk].float() lb = flat_labels[start:start + chunk] valid = lb.ne(-100) if not valid.any(): continue total = total + F.cross_entropy(sl, lb, ignore_index=-100, reduction="sum") count += int(valid.sum()) if count == 0: return torch.zeros((), device=logits.device, requires_grad=True) return total / count @torch.no_grad() def generate(self, input_ids, max_new_tokens=32, temperature=1.0, eos_token_id=None): self.eval() for _ in range(max_new_tokens): ctx = input_ids[:, -self.config.max_position_embeddings:] logits = self.forward(ctx)["logits"][:, -1, :] / max(temperature, 1e-5) probs = torch.softmax(logits.float(), dim=-1) next_id = torch.multinomial(probs, num_samples=1) input_ids = torch.cat([input_ids, next_id], dim=1) if eos_token_id is not None and bool((next_id == eos_token_id).all()): break return input_ids def count_parameters(model: nn.Module) -> int: """Count unique parameters, so tied weights are not double counted.""" seen, total = set(), 0 for param in model.parameters(): if id(param) in seen: continue seen.add(id(param)) total += param.numel() return total