"""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, }