Download tests/test_long_context.py from thefinalboss/palimpseste-max: direct link, hf CLI and curl.
- Browser
- Download file 7.13 kB
-
https://huggingface.co/thefinalboss/palimpseste-max/resolve/main/tests/test_long_context.py
- Command line
-
hf download hf://thefinalboss/palimpseste-max/tests/test_long_context.py
-
curl -L -o test_long_context.py https://huggingface.co/thefinalboss/palimpseste-max/resolve/main/tests/test_long_context.py
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 | |
| def encoder(): | |
| return Encoder(D=2000) | |
| def tokenizer(encoder): | |
| tok = CharTokenizer(encoder=encoder) | |
| tok.build_vocab("abcdefghijklmnopqrstuvwxyz .,!?0123456789") | |
| return tok | |
| 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 | |