"""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) # Different seeds + temperatures may produce different outputs 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) # Both should produce valid results 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) # Should handle short context gracefully 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) # They should NOT be identical (mixture uses per-candidate info) assert not np.allclose(standard, mixture)