Cortex-ai / model /cortex_model.py
openhands
openhands
feat(cortex): add CORTEX training pipeline and model audit
6fbe100
Raw History Blame Contribute Delete
13.2 kB
"""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