Sol-Lassi / keystone_mlx /model.py
j0no12's picture
Publish Sol Lassi 600K Base
063093a verified
Raw History Blame Contribute Delete
16.9 kB
"""Keystone implemented directly with MLX for Apple Metal execution.
Keystone is an experimental, non-Transformer language-model hypothesis. A
causal read/write chronicle summarizes the prefix into a compact bank of dense
states; one shared dense gated processor then refines each token state several
times. The repeated processor intentionally trades extra compute for stored
parameter efficiency. It is not a Llama implementation and has not yet been
validated at meaningful language-model scale.
This module deliberately depends only on ``mlx.core`` and ``mlx.nn``. It has
no PyTorch fallback: a missing or inaccessible MLX Metal runtime is a setup
error, not an invitation to silently run a different backend.
"""
from __future__ import annotations
from dataclasses import asdict, dataclass
import math
from typing import Any
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_flatten
import numpy as np
@dataclass(frozen=True)
class KeystoneConfig:
"""Immutable architecture specification for the 1.95M-parameter model.
Args:
vocab_size: Number of token IDs accepted by the tied dense interface.
width: Width of token states and every chronicle slot.
memory_slots: Independently updated chronicle states.
processor_width: Hidden width of the shared gated processor.
refinement_steps: Number of repeated processor applications.
Side effects:
None. Invalid dimensions raise ``ValueError`` during construction.
"""
vocab_size: int = 4_096
width: int = 192
memory_slots: int = 17
processor_width: int = 1_500
refinement_steps: int = 4
def __post_init__(self) -> None:
if self.vocab_size < 2:
raise ValueError("vocab_size must be at least two")
if self.width < 2:
raise ValueError("width must be at least two")
if self.memory_slots < 2:
raise ValueError("memory_slots must be at least two")
if self.processor_width < self.width:
raise ValueError("processor_width must be at least width")
if self.refinement_steps < 1:
raise ValueError("refinement_steps must be positive")
# Every published configuration retains the same architecture. The vocabulary
# choice is an explicit stored-parameter allocation decision, not compression:
# a 2K interface buys more dense chronicle/processor capacity at the same 1M
# budget, while the 4K controls preserve the earlier allocation.
KEYSTONE_2M_CONFIG = KeystoneConfig()
KEYSTONE_1M_4K_CONFIG = KeystoneConfig(width=144, memory_slots=17, processor_width=556, refinement_steps=4)
KEYSTONE_1M_2K_CONFIG = KeystoneConfig(
vocab_size=2_048,
width=192,
memory_slots=17,
processor_width=532,
refinement_steps=4,
)
# Compatibility alias for callers that referred to the old generic 1M label.
# New training commands must use explicit `1m-2k` or `1m-4k` selection.
KEYSTONE_1M_CONFIG = KEYSTONE_1M_2K_CONFIG
DEFAULT_CONFIG = KEYSTONE_1M_2K_CONFIG
def config_for_size(model_size: str) -> KeystoneConfig:
"""Select one of the explicitly budgeted Keystone architecture variants.
Args:
model_size: One explicit published budget label: ``"1m-2k"``,
``"1m-4k"``, or ``"2m-4k"``.
Returns:
Immutable configuration for that exact stored-parameter budget.
Side effects:
None. An unsupported label raises ``ValueError`` rather than choosing a
nearby architecture silently.
"""
configurations = {
"1m-2k": KEYSTONE_1M_2K_CONFIG,
"1m-4k": KEYSTONE_1M_4K_CONFIG,
"2m-4k": KEYSTONE_2M_CONFIG,
}
try:
return configurations[model_size]
except KeyError as error:
raise ValueError(f"unknown Keystone model size: {model_size!r}") from error
def model_size_for_config(config: KeystoneConfig) -> str:
"""Return the published budget label for an exact Keystone configuration."""
for size, candidate in (
("1m-2k", KEYSTONE_1M_2K_CONFIG),
("1m-4k", KEYSTONE_1M_4K_CONFIG),
("2m-4k", KEYSTONE_2M_CONFIG),
):
if config == candidate:
return size
raise ValueError("configuration is not one of the published Keystone budget selections")
class CenteredUnitNorm(nn.Module):
"""Learned affine mean-and-variance normalization over the final axis."""
def __init__(self, width: int, epsilon: float = 1e-5) -> None:
super().__init__()
self.scale = mx.ones((width,))
self.shift = mx.zeros((width,))
self.epsilon = epsilon
def __call__(self, values: mx.array) -> mx.array:
"""Normalize vectors while preserving every batch and sequence axis.
Args:
values: Array whose final dimension is this norm's configured width.
Returns:
Centered, variance-scaled, affine transformed values.
Side effects:
None.
"""
mean = mx.mean(values, axis=-1, keepdims=True)
variance = mx.mean(mx.square(values - mean), axis=-1, keepdims=True)
return (values - mean) * mx.rsqrt(variance + self.epsilon) * self.scale + self.shift
class DenseTiedTokens(nn.Module):
"""One full-rank table, used exactly for both token input and output."""
def __init__(self, vocab_size: int, width: int) -> None:
super().__init__()
self.table = mx.random.normal(shape=(vocab_size, width)) * 0.02
def encode(self, token_ids: mx.array) -> mx.array:
"""Look up uncompressed token vectors from the sole vocabulary table."""
return self.table[token_ids]
def decode(self, hidden: mx.array) -> mx.array:
"""Produce logits using the exact transpose of the input table."""
return hidden @ self.table.T
def deterministic_coordinates(length: int, width: int) -> mx.array:
"""Build a fixed sinusoidal coordinate signal with no learned parameters.
Args:
length: Number of sequence positions.
width: Number of state channels.
Returns:
A float32 array shaped ``[length, width]``.
Side effects:
Allocates a temporary MLX array but does not examine token values.
"""
if length < 1:
raise ValueError("length must be positive")
half = (width + 1) // 2
positions = mx.expand_dims(mx.arange(length, dtype=mx.float32), axis=1)
frequencies = mx.exp(
mx.arange(half, dtype=mx.float32) * (-math.log(20_000.0) / max(half - 1, 1))
)
coordinates = mx.concatenate((mx.sin(positions * frequencies), mx.cos(positions * frequencies)), axis=-1)
return coordinates[:, :width]
class ChronicleReadWrite(nn.Module):
"""Prefix-only dense read/write state bank, evaluated left to right."""
def __init__(self, width: int, memory_slots: int) -> None:
super().__init__()
self.memory_slots = memory_slots
self.reader_address = nn.Linear(width, width, bias=False)
self.reader_output = nn.Linear(width, width, bias=False)
self.writer_address = nn.Linear(width, width, bias=False)
self.writer_proposal = nn.Linear(2 * width, width, bias=False)
self.initial_slots = mx.random.normal(shape=(memory_slots, width)) * 0.01
def __call__(self, hidden: mx.array) -> mx.array:
"""Read prefix state, then soft-write the current token into every slot.
Args:
hidden: Input states shaped ``[batch, sequence, width]``.
Returns:
Prefix-only context shaped like ``hidden``. Output at position ``i``
cannot depend on tokens after ``i``.
Side effects:
Allocates ephemeral recurrent slot arrays; stored initial slots are
never mutated in place.
"""
if hidden.ndim != 3:
raise ValueError("hidden must have shape [batch, sequence, width]")
batch, length, width = hidden.shape
if width != self.initial_slots.shape[1]:
raise ValueError("hidden width does not match chronicle slot width")
slots = mx.broadcast_to(mx.expand_dims(self.initial_slots, axis=0), (batch, self.memory_slots, width))
contexts: list[mx.array] = []
scale = width**-0.5
for position in range(length):
current = hidden[:, position, :]
read_query = self.reader_address(current)
read_weights = mx.softmax(mx.sum(mx.expand_dims(read_query, 1) * slots, axis=-1) * scale, axis=-1)
retrieved = mx.sum(mx.expand_dims(read_weights, -1) * slots, axis=1)
context = self.reader_output(retrieved)
contexts.append(context)
write_query = self.writer_address(current)
write_weights = mx.softmax(mx.sum(mx.expand_dims(write_query, 1) * slots, axis=-1) * scale, axis=-1)
proposal = mx.tanh(self.writer_proposal(mx.concatenate((current, context), axis=-1)))
slots = slots + mx.expand_dims(write_weights, -1) * (mx.expand_dims(proposal, 1) - slots)
return mx.stack(contexts, axis=1)
class SharedKeystoneProcessor(nn.Module):
"""Full-dense gated processor whose matrices are reused across refinements."""
def __init__(self, width: int, processor_width: int) -> None:
super().__init__()
self.expand = nn.Linear(width, 2 * processor_width, bias=False)
self.contract = nn.Linear(processor_width, width, bias=False)
def __call__(self, state: mx.array) -> mx.array:
"""Run one gated dense transform without storing a second processor."""
content, gate = mx.split(self.expand(state), 2, axis=-1)
return self.contract(nn.silu(content) * mx.sigmoid(gate))
class KeystoneLM(nn.Module):
"""Custom causal LM: chronicle state followed by repeated shared refinement."""
def __init__(self, config: KeystoneConfig = DEFAULT_CONFIG) -> None:
super().__init__()
self.config = config
self.tokens = DenseTiedTokens(config.vocab_size, config.width)
self.ingress = nn.Linear(config.width, config.width, bias=False)
self.read_norm = CenteredUnitNorm(config.width)
self.chronicle = ChronicleReadWrite(config.width, config.memory_slots)
self.processor_norm = CenteredUnitNorm(config.width)
self.processor = SharedKeystoneProcessor(config.width, config.processor_width)
self.refinement_gate = nn.Linear(2 * config.width, config.width, bias=False)
self.step_codes = mx.zeros((config.refinement_steps, config.width))
self.final_norm = CenteredUnitNorm(config.width)
def __call__(self, token_ids: mx.array) -> mx.array:
"""Compute full-vocabulary causal next-token logits.
Args:
token_ids: Integer IDs shaped ``[batch, sequence]`` in the
configured vocabulary range.
Returns:
A ``[batch, sequence, vocab_size]`` logit tensor. All sequence
dependence is created by the explicit prefix-only chronicle scan.
Side effects:
Allocates coordinate, chronicle, and refinement intermediate arrays.
"""
if token_ids.ndim != 2:
raise ValueError("token_ids must have shape [batch, sequence]")
if token_ids.shape[1] < 1:
raise ValueError("token_ids must include at least one position")
hidden = self.tokens.encode(token_ids)
coordinates = deterministic_coordinates(token_ids.shape[1], self.config.width)
hidden = self.ingress(hidden + mx.expand_dims(coordinates, axis=0))
context = self.chronicle(self.read_norm(hidden))
state = hidden + context
for step in range(self.config.refinement_steps):
coded_state = state + context + self.step_codes[step][None, None, :]
candidate = self.processor(self.processor_norm(coded_state))
gate = mx.sigmoid(self.refinement_gate(mx.concatenate((state, candidate), axis=-1)))
state = state + gate * (candidate - state)
return self.tokens.decode(self.final_norm(state))
def parameter_accounting(config: KeystoneConfig = DEFAULT_CONFIG) -> dict[str, int]:
"""Return the exact stored-parameter budget, counting tied tokens once.
Args:
config: Architecture specification to count.
Returns:
Counts by subsystem plus ``total``. No activation or optimizer memory
appears here because this is deliberately a stored-parameter budget.
Side effects:
None.
"""
vocabulary = config.vocab_size * config.width
ingress = config.width * config.width
chronicle_reader = 2 * config.width * config.width
chronicle_writer = 3 * config.width * config.width
chronicle_slots = config.memory_slots * config.width
processor = 3 * config.width * config.processor_width
refinement_gate = 2 * config.width * config.width
step_codes = config.refinement_steps * config.width
normalization = 3 * 2 * config.width
total = sum((
vocabulary, ingress, chronicle_reader, chronicle_writer, chronicle_slots,
processor, refinement_gate, step_codes, normalization,
))
return {
"dense_tied_vocabulary_interface": vocabulary,
"dense_ingress": ingress,
"chronicle_reader": chronicle_reader,
"chronicle_writer": chronicle_writer,
"chronicle_initial_slots": chronicle_slots,
"shared_dense_processor": processor,
"shared_refinement_gate": refinement_gate,
"refinement_step_codes": step_codes,
"normalization": normalization,
"total": total,
}
def parameter_count(model: nn.Module) -> int:
"""Count trainable MLX scalar parameters, including each tied array once."""
return sum(int(parameter.size) for _, parameter in tree_flatten(model.trainable_parameters()))
def future_token_causality_error(
model: KeystoneLM,
sequence_length: int = 11,
batch_size: int = 2,
seed: int = 23,
) -> float:
"""Return the largest changed prefix logit after deterministic suffix edits.
Args:
model: Keystone model under test.
sequence_length: Input length, at least three.
batch_size: Number of independent sequences.
seed: NumPy seed used only for this test.
Returns:
Maximum absolute prefix difference. A correct deterministic causal model
should return exactly zero to its evaluated floating-point precision.
Side effects:
Executes two MLX inference graphs and synchronizes their outputs.
"""
if sequence_length < 3 or batch_size < 1:
raise ValueError("sequence_length must be at least three and batch_size must be positive")
generator = np.random.default_rng(seed)
values = generator.integers(0, model.config.vocab_size, size=(batch_size, sequence_length), dtype=np.int32)
boundary = sequence_length // 2
mutated = values.copy()
mutated[:, boundary:] = (mutated[:, boundary:] + 1) % model.config.vocab_size
original_logits = model(mx.array(values))
changed_logits = model(mx.array(mutated))
mx.eval(original_logits, changed_logits)
return float(mx.max(mx.abs(original_logits[:, :boundary] - changed_logits[:, :boundary])))
def self_check(config: KeystoneConfig = DEFAULT_CONFIG) -> dict[str, Any]:
"""Validate accounting, output shape, and causal-prefix invariance.
Returns:
A JSON-serializable diagnostic record.
Side effects:
Instantiates the full model and evaluates two Metal inference passes.
It therefore requires a usable MLX device.
"""
model = KeystoneLM(config)
formula = parameter_accounting(config)
observed = parameter_count(model)
if observed != formula["total"]:
raise RuntimeError(f"formula counted {formula['total']}, model stores {observed}")
known_budgets = {
KEYSTONE_1M_2K_CONFIG: 999_744,
KEYSTONE_1M_4K_CONFIG: 999_792,
KEYSTONE_2M_CONFIG: 1_950_528,
}
expected_budget = known_budgets.get(config)
if expected_budget is None:
raise RuntimeError("self_check only accepts a published Keystone budget configuration")
if observed != expected_budget:
raise RuntimeError(f"model has {observed} parameters; expected published budget {expected_budget}")
probe = model(mx.zeros((2, 11), dtype=mx.int32))
mx.eval(probe)
expected_shape = (2, 11, config.vocab_size)
if tuple(probe.shape) != expected_shape:
raise RuntimeError(f"expected {expected_shape}, received {tuple(probe.shape)}")
causal_error = future_token_causality_error(model)
if causal_error != 0.0:
raise RuntimeError(f"future-token leakage: {causal_error}")
return {
"architecture": f"Keystone-{model_size_for_config(config).upper()}",
"config": asdict(config),
**formula,
"causal_forward_shape": list(probe.shape),
"future_token_causality_error": causal_error,
}