| """Tests for creative generation (mixture logits).""" |
|
|
| import pytest |
| import numpy as np |
|
|
| from palimseste.lm import PalimpsesteForCausalLM, PalimpsesteConfig |
| from palimseste.creative import CreativeGenerator, CreativeResult |
|
|
|
|
| def _build_model(D=5000, ctx=128, radius=100): |
| cfg = PalimpsesteConfig(D=D, context_window=ctx, kernel_radius=radius, temperature=0.0) |
| lm = PalimpsesteForCausalLM(config=cfg) |
| pairs = [ |
| ("hello", "hi i am palimpseste"), |
| ("who are you", "i am palimpseste a hypervectorial cortex"), |
| ("what is python", "python is a programming language"), |
| ("what is java", "java is a programming language"), |
| ("what is gravity", "gravity is a force that attracts masses"), |
| ("what is magnetism", "magnetism is a force from magnets"), |
| ("what is ruby", "ruby is a programming language"), |
| ("what is rust", "rust is a programming language"), |
| ] |
| lm.build_tokenizer("".join(q + a for q, a in pairs)) |
| lm.train_on_qa_pairs(pairs) |
| return lm |
|
|
|
|
| class TestCreativeGenerator: |
| def test_mixture_logits_returns_array(self): |
| lm = _build_model() |
| gen = CreativeGenerator(lm=lm, top_k=3) |
| ids = lm.tokenizer.encode("what is python", add_bos=True, add_eos=True) + [1] |
| logits = gen.mixture_logits(ids) |
| assert isinstance(logits, np.ndarray) |
| assert len(logits) == lm.config.vocab_size |
|
|
| def test_generate_returns_result(self): |
| lm = _build_model() |
| gen = CreativeGenerator(lm=lm, top_k=3, diversity=0.2) |
| result = gen.generate("what is python", max_new_tokens=20, temperature=0.0) |
| assert isinstance(result, CreativeResult) |
| assert isinstance(result.text, str) |
|
|
| def test_respond_returns_string(self): |
| lm = _build_model() |
| gen = CreativeGenerator(lm=lm, top_k=3) |
| text = gen.respond("hello", max_new_tokens=20) |
| assert isinstance(text, str) |
|
|
| def test_novel_tokens_tracked(self): |
| lm = _build_model() |
| gen = CreativeGenerator(lm=lm, top_k=3, diversity=0.5) |
| result = gen.generate("what is python", max_new_tokens=20, temperature=0.0) |
| assert result.novel_tokens >= 0 |
|
|
| def test_temperature_changes_output(self): |
| lm = _build_model() |
| gen = CreativeGenerator(lm=lm, top_k=3, diversity=0.3) |
| r1 = gen.generate("what is gravity", max_new_tokens=20, temperature=0.0, seed=42) |
| r2 = gen.generate("what is gravity", max_new_tokens=20, temperature=0.5, seed=99) |
| |
| assert isinstance(r1.text, str) |
| assert isinstance(r2.text, str) |
|
|
| def test_ngram_blocking(self): |
| lm = _build_model() |
| gen = CreativeGenerator(lm=lm, top_k=3) |
| result = gen.generate("hello", max_new_tokens=30, temperature=0.0, ngram_block=3) |
| assert isinstance(result, CreativeResult) |
|
|
| def test_diversity_parameter(self): |
| lm = _build_model() |
| gen_low = CreativeGenerator(lm=lm, top_k=3, diversity=0.0) |
| gen_high = CreativeGenerator(lm=lm, top_k=3, diversity=0.8) |
| |
| r1 = gen_low.generate("hello", max_new_tokens=10) |
| r2 = gen_high.generate("hello", max_new_tokens=10) |
| assert isinstance(r1.text, str) |
| assert isinstance(r2.text, str) |
|
|
| def test_empty_context(self): |
| lm = _build_model() |
| gen = CreativeGenerator(lm=lm, top_k=3) |
| |
| result = gen.generate("x", max_new_tokens=5) |
| assert isinstance(result, CreativeResult) |
|
|
| def test_known_question_works(self): |
| lm = _build_model() |
| gen = CreativeGenerator(lm=lm, top_k=3, diversity=0.2) |
| result = gen.generate("what is python", max_new_tokens=20, temperature=0.0) |
| assert len(result.text) > 0 |
|
|
| def test_mixture_logits_different_from_standard(self): |
| """Mixture logits should differ from standard single-HV logits.""" |
| lm = _build_model() |
| gen = CreativeGenerator(lm=lm, top_k=5, diversity=0.3) |
| ids = lm.tokenizer.encode("what is gravity", add_bos=True, add_eos=True) + [1] |
|
|
| standard = lm._logits(ids) |
| mixture = gen.mixture_logits(ids) |
|
|
| |
| assert not np.allclose(standard, mixture) |
|
|