ensemble / palimseste /ngram_model.py
thefinalboss's picture
Initial release: ENSEMBLE training-free AI — compressed .exp experts + Kuramoto brain
1f71c7d verified
Raw
History Blame Contribute Delete
10.5 kB
"""PALIMPSESTE — Token transition model for free generation.
The core limitation: _logits() always picks the nearest ATOMIC token to the
retrieved HV. This means the model can only REPEAT stored sequences.
This module adds a TOKEN TRANSITION MODEL: it learns which tokens follow
which token sequences (markov chain at the sub-word level). During generation,
it blends:
1. HV retrieval (what the associative memory says)
2. Transition probability (what usually follows this sequence)
The blend ratio controls creativity vs accuracy. With high transition weight,
the model generates NOVEL text by following learned patterns rather than
echoing exact stored answers.
This is how PALIMPSESTE learns to GENERATE, not just RETRIEVE.
"""
from __future__ import annotations
import numpy as np
from collections import defaultdict, Counter
from dataclasses import dataclass, field
from typing import Optional
__all__ = ["TransitionModel", "TransitionGenerator", "TransitionResult"]
@dataclass
class TransitionResult:
"""Result of transition-based generation."""
text: str
n_tokens: int
retrieval_tokens: int # tokens from HV retrieval
transition_tokens: int # tokens from transition model
novel_ratio: float
class TransitionModel:
"""Markov chain transition model over BPE tokens.
Learns P(next_token | last_N_tokens) from training data. This captures
local patterns like:
- "is a" → likely "language" or "programming"
- "the capital of" → likely a proper noun
- "def " → likely a function name
Combined with HV retrieval, this enables FREE GENERATION where the model
can produce text it has never seen verbatim.
"""
def __init__(self, order: int = 2) -> None:
"""
Parameters
----------
order : int
Markov order (2 = bigram context, 3 = trigram context).
Higher = more context but sparser data.
"""
self.order = order
# transitions[context_tuple] = Counter({token_id: count})
self._transitions: dict[tuple, Counter] = defaultdict(Counter)
# Backoff: unigram counts for smoothing
self._unigrams: Counter = Counter()
self._total_tokens: int = 0
self._vocab_size: int = 0
def train(self, token_sequences: list[list[int]]) -> None:
"""Learn transitions from token sequences.
Parameters
----------
token_sequences : list[list[int]]
Each inner list is a sequence of token IDs (e.g., from BPE encoding).
"""
for seq in token_sequences:
for i, tok in enumerate(seq):
self._unigrams[tok] += 1
self._total_tokens += 1
# Record n-gram transitions
for o in range(1, self.order + 1):
if i >= o:
context = tuple(seq[i-o:i])
self._transitions[context][tok] += 1
self._vocab_size = len(self._unigrams)
def predict(self, context: list[int], top_k: int = 10) -> dict[int, float]:
"""Predict next token probabilities given context.
Returns dict {token_id: probability} for top-k candidates.
Uses backoff: try highest order first, fall back to lower orders.
"""
result: dict[int, float] = {}
# Try each order from highest to lowest (backoff)
for o in range(min(self.order, len(context)), 0, -1):
ctx = tuple(context[-o:]) if len(context) >= o else tuple(context)
if ctx in self._transitions:
counts = self._transitions[ctx]
total = sum(counts.values())
for tok, count in counts.most_common(top_k):
prob = count / total
# Blend with lower order (interpolation)
result[tok] = result.get(tok, 0) + prob * 0.7
break # use highest available order
# Backoff to unigram if nothing found
if not result:
for tok, count in self._unigrams.most_common(top_k):
result[tok] = count / self._total_tokens * 0.1
return result
def get_logit_boost(self, context: list[int], vocab_size: int) -> np.ndarray:
"""Get a logit boost vector for blending with HV retrieval.
Returns a vector of shape (vocab_size,) where higher values
indicate tokens that are likely to follow the context.
"""
boost = np.zeros(vocab_size, dtype=np.float64)
probs = self.predict(context, top_k=20)
for tok, prob in probs.items():
if 0 <= tok < vocab_size:
boost[tok] = prob
return boost
@property
def n_contexts(self) -> int:
return len(self._transitions)
@property
def vocab_size(self) -> int:
return self._vocab_size
class TransitionGenerator:
"""Generate text by blending HV retrieval with token transitions.
The generation process at each step:
1. Compute HV retrieval logits (standard _logits)
2. Compute transition boost (markov chain prediction)
3. Blend: final_logits = retrieval * retrieval_weight + transition * transition_weight
4. Apply temperature sampling
5. Apply anti-repetition
With transition_weight > 0, the model can generate NOVEL sequences
that follow learned patterns but were never stored verbatim.
"""
def __init__(
self,
lm,
transition_model: TransitionModel,
retrieval_weight: float = 0.6,
transition_weight: float = 0.4,
) -> None:
self.lm = lm
self.tm = transition_model
self.r_weight = retrieval_weight
self.t_weight = transition_weight
def generate(
self,
prompt: str,
max_new_tokens: int = 100,
temperature: float = 0.3,
ngram_block: int = 3,
seed: int | None = None,
) -> TransitionResult:
"""Generate text by blending retrieval + transitions.
Parameters
----------
prompt : str
Input text.
max_new_tokens : int
Max tokens to generate.
temperature : float
Sampling temperature (0 = greedy, >0 = sample).
For creative generation, use 0.2-0.5.
ngram_block : int
Block repetition of n-grams of this length.
seed : int | None
RNG seed.
"""
from .lm import BOS, EOS, PAD
assert self.lm.tokenizer is not None
tok = self.lm.tokenizer
rng = np.random.default_rng(seed) if seed is not None else np.random.default_rng()
ids = tok.encode(prompt, add_bos=True, add_eos=True) + [BOS]
generated: list[int] = []
retrieval_count = 0
transition_count = 0
V = getattr(tok, 'vocab_size_actual', None) or self.lm.config.vocab_size
for step in range(max_new_tokens):
# 1. HV retrieval logits
retrieval_logits = self.lm._logits(ids)
if retrieval_logits.max() == 0.0 and retrieval_logits.min() == 0.0:
break # cold start
# 2. Transition boost
transition_boost = self.tm.get_logit_boost(generated or ids[-5:], V)
# 3. The retrieval logits ARE the generation distribution.
# For creative generation, we use temperature sampling on
# retrieval logits directly — this naturally produces variation
# because the softmax over token similarities has a long tail.
# The transition model is used ONLY to break ties and provide
# a gentle nudge toward locally coherent sequences.
if retrieval_logits.max() > retrieval_logits.min():
retrieval_logits = (retrieval_logits - retrieval_logits.min()) / (
retrieval_logits.max() - retrieval_logits.min()
)
# 4. Pure temperature-scaled retrieval (the creative engine)
# With temp > 0, the model samples from the distribution
# instead of taking argmax. This produces NOVEL sequences.
# The transition model adds a tiny bias to keep output coherent.
coherent_bias = np.zeros_like(retrieval_logits)
if transition_boost.max() > 0:
coherent_bias = transition_boost / transition_boost.max()
blended = retrieval_logits * (1.0 - self.t_weight * 0.3) + coherent_bias * (self.t_weight * 0.3)
# 5. N-gram blocking
if ngram_block > 0 and len(generated) >= ngram_block:
recent = tuple(generated[-(ngram_block - 1):]) if ngram_block > 1 else ()
for i in range(len(generated) - ngram_block + 1):
ng = tuple(generated[i:i + ngram_block - 1]) if ngram_block > 1 else ()
if ng == recent and i + ngram_block - 1 < len(generated):
blended[generated[i + ngram_block - 1]] *= 0.1
# 6. Select token
retrieval_choice = int(np.argmax(retrieval_logits))
if temperature <= 1e-6:
next_id = int(np.argmax(blended))
else:
z = blended / max(temperature, 0.01)
z -= z.max()
e = np.exp(z)
probs = e / e.sum()
probs[PAD] = 0.0
s = probs.sum()
if s <= 0:
break
probs /= s
next_id = int(rng.choice(len(probs), p=probs))
# Track source
if next_id == retrieval_choice:
retrieval_count += 1
else:
transition_count += 1
generated.append(next_id)
ids.append(next_id)
if next_id == EOS:
break
text = tok.decode(generated)
total = max(len(generated), 1)
return TransitionResult(
text=text,
n_tokens=len(generated),
retrieval_tokens=retrieval_count,
transition_tokens=transition_count,
novel_ratio=transition_count / total,
)
def respond(self, question: str, max_new_tokens: int = 100,
temperature: float = 0.3, seed: int | None = None) -> str:
"""Respond to a question."""
result = self.generate(question, max_new_tokens=max_new_tokens,
temperature=temperature, seed=seed)
return result.text