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