palimpseste-max / tests /test_long_context.py
thefinalboss's picture
Upload tests/test_long_context.py with huggingface_hub
87c9f03 verified
Raw
History Blame Contribute Delete
7.13 kB
"""Tests for the hierarchical (million-token) context window."""
import pytest
import numpy as np
from palimseste.hierarchical_context import (
HierarchicalContext,
HierarchicalContextConfig,
)
from palimseste.hv import HV, similarity, random_hv
from palimseste.learner import Encoder
from palimseste.tokenizer import CharTokenizer
from palimseste.lm import PalimpsesteForCausalLM, PalimpsesteConfig, PRESETS
# ----------------------------------------------------------- fixtures
@pytest.fixture
def encoder():
return Encoder(D=2000)
@pytest.fixture
def tokenizer(encoder):
tok = CharTokenizer(encoder=encoder)
tok.build_vocab("abcdefghijklmnopqrstuvwxyz .,!?0123456789")
return tok
@pytest.fixture
def hc(encoder):
return HierarchicalContext(
D=2000, encoder=encoder,
config=HierarchicalContextConfig(chunk_size=10, local_window=8, top_k_chunks=3),
)
# ----------------------------------------------------------- basic functionality
class TestHierarchicalContextBasic:
def test_empty_state(self, hc):
"""Empty context returns a valid HV."""
state = hc.get_state()
assert isinstance(state, HV)
assert state.D == 2000
def test_ingest_increments_tokens(self, hc, tokenizer):
ids = tokenizer.encode("hello world", add_bos=True, add_eos=True)
hc.ingest(ids, tokenizer)
assert hc.n_tokens == len(ids)
assert hc.n_local_tokens <= 8
def test_chunking(self, hc, tokenizer):
"""Tokens are split into chunks of chunk_size."""
ids = tokenizer.encode("aaaaaaaaaa bbbbbbbbbb cccccccccc", add_bos=True, add_eos=True)
hc.ingest(ids, tokenizer)
hc.flush()
# Should have at least 1 chunk
assert hc.n_chunks >= 1
assert hc.n_tokens == len(ids)
def test_stats(self, hc, tokenizer):
ids = tokenizer.encode("test", add_bos=True, add_eos=True)
hc.ingest(ids, tokenizer)
stats = hc.stats()
assert "n_chunks" in stats
assert "n_tokens" in stats
assert "max_context_tokens" in stats
assert stats["max_context_tokens"] == 4096 * 10
# ----------------------------------------------------------- state quality
class TestHierarchicalContextState:
def test_same_query_same_state(self, hc, tokenizer):
"""Same input + same query produces identical state."""
ids = tokenizer.encode("hello world test context", add_bos=True, add_eos=True)
hc.ingest(ids, tokenizer)
s1 = hc.get_state()
s2 = hc.get_state()
assert similarity(s1, s2) == pytest.approx(1.0, abs=0.01)
def test_different_context_different_state(self, encoder, tokenizer):
"""Different contexts produce different states."""
hc1 = HierarchicalContext(D=2000, encoder=encoder)
hc2 = HierarchicalContext(D=2000, encoder=encoder)
ids1 = tokenizer.encode("alpha beta gamma delta", add_bos=True, add_eos=True)
ids2 = tokenizer.encode("one two three four", add_bos=True, add_eos=True)
hc1.ingest(ids1, tokenizer)
hc2.ingest(ids2, tokenizer)
s1 = hc1.get_state()
s2 = hc2.get_state()
# Should NOT be identical
assert similarity(s1, s2) < 0.95
def test_state_is_valid_hv(self, hc, tokenizer):
ids = tokenizer.encode("test", add_bos=True, add_eos=True)
hc.ingest(ids, tokenizer)
state = hc.get_state()
assert isinstance(state, HV)
assert state.D == 2000
assert state.bits is not None
# ----------------------------------------------------------- scale
class TestHierarchicalContextScale:
def test_large_context(self, encoder, tokenizer):
"""100 chunks of 100 tokens = 10K tokens."""
hc = HierarchicalContext(
D=2000, encoder=encoder,
config=HierarchicalContextConfig(chunk_size=100, local_window=64),
)
valid = list(range(tokenizer.vocab_size))
ids = np.random.default_rng(42).integers(0, tokenizer.vocab_size, size=10000).tolist()
hc.ingest(ids, tokenizer)
hc.flush()
assert hc.n_chunks == 100
assert hc.n_tokens == 10000
def test_chunk_pruning(self, encoder, tokenizer):
"""Old chunks are pruned when max_chunks is exceeded."""
hc = HierarchicalContext(
D=2000, encoder=encoder,
config=HierarchicalContextConfig(chunk_size=5, local_window=5, max_chunks=3),
)
ids = list(range(tokenizer.vocab_size)) * 10 # lots of tokens
hc.ingest(ids[:100], tokenizer)
hc.flush()
assert hc.n_chunks <= 3
def test_retrieval_at_scale(self, encoder, tokenizer):
"""get_state works quickly even with many chunks."""
hc = HierarchicalContext(
D=2000, encoder=encoder,
config=HierarchicalContextConfig(chunk_size=50, local_window=32, top_k_chunks=5),
)
rng = np.random.default_rng(0)
ids = rng.integers(0, tokenizer.vocab_size, size=5000).tolist()
hc.ingest(ids, tokenizer)
hc.flush()
assert hc.n_chunks > 50
state = hc.get_state()
assert isinstance(state, HV)
# ----------------------------------------------------------- LM integration
class TestLongContextIntegration:
def test_preset_1b_long(self):
"""The 1b-long preset has context_window=1M."""
cfg = PRESETS["1b-long"]
assert cfg.context_window == 1_000_000
assert cfg.D == 100_000
def test_enable_long_context(self):
"""enable_long_context activates hierarchical context."""
lm = PalimpsesteForCausalLM(
config=PalimpsesteConfig(D=2000, context_window=1_000_000, kernel_radius=100)
)
lm.build_tokenizer("hello world test context")
assert lm.long_context_stats() is None
lm.enable_long_context(chunk_size=10, local_window=8)
assert lm.long_context_stats() is not None
def test_state_hv_uses_hier(self):
"""_state_hv uses hierarchical context when enabled."""
lm = PalimpsesteForCausalLM(
config=PalimpsesteConfig(D=2000, context_window=1_000_000, kernel_radius=100)
)
lm.build_tokenizer("hello world test context")
lm.enable_long_context(chunk_size=10, local_window=8)
ids = lm.tokenizer.encode("hello world", add_bos=True, add_eos=True)
hv = lm._state_hv(ids)
assert isinstance(hv, HV)
stats = lm.long_context_stats()
assert stats["n_tokens"] > 0
def test_respond_with_long_context(self):
"""respond() works with long context enabled."""
lm = PalimpsesteForCausalLM(
config=PalimpsesteConfig(D=2000, context_window=1_000_000, kernel_radius=100)
)
lm.build_tokenizer("hello world test context system memory hypervector")
lm.enable_long_context(chunk_size=10, local_window=8, top_k_chunks=3)
resp = lm.respond("hello", max_new_tokens=5)
assert isinstance(resp, str)
stats = lm.long_context_stats()
assert stats["n_tokens"] > 0