| """Tests for coherent generator.""" |
| import pytest |
| import numpy as np |
|
|
| from palimseste.lm import PalimpsesteForCausalLM, PalimpsesteConfig |
| from palimseste.kuramoto import KuramotoAttractor |
| from palimseste.ngram_model import TransitionModel |
| from palimseste.coherent import CoherentGenerator, CoherentResult |
|
|
|
|
| def _build(D=3000, radius=200): |
| cfg = PalimpsesteConfig(D=D, context_window=64, kernel_radius=radius, temperature=0.0) |
| lm = PalimpsesteForCausalLM(config=cfg) |
| pairs = [ |
| ("what is python", "python is a programming language"), |
| ("what is gravity", "gravity is a force that attracts masses"), |
| ("who are you", "i am palimpseste"), |
| ("hello", "hi i am palimpseste"), |
| ] |
| lm.build_tokenizer("".join(q+a for q,a in pairs)) |
| lm.train_on_qa_pairs(pairs) |
| tok = lm.tokenizer |
| seqs = [tok.encode(q+' '+a, add_bos=True, add_eos=True) for q,a in pairs] |
| tm = TransitionModel(order=2) |
| tm.train(seqs) |
| return lm, tm |
|
|
|
|
| class TestCoherentGenerator: |
| def test_generate_returns_result(self): |
| lm, tm = _build() |
| gen = CoherentGenerator(lm=lm, kuramoto=KuramotoAttractor(n_iterations=10), transition_model=tm) |
| result = gen.generate("what is python", max_new_tokens=15) |
| assert isinstance(result, CoherentResult) |
|
|
| def test_respond_returns_string(self): |
| lm, tm = _build() |
| gen = CoherentGenerator(lm=lm, transition_model=tm) |
| text = gen.respond("hello", max_new_tokens=15) |
| assert isinstance(text, str) |
|
|
| def test_coherence_score(self): |
| lm, tm = _build() |
| gen = CoherentGenerator(lm=lm, transition_model=tm) |
| result = gen.generate("what is python", max_new_tokens=15, temperature=0.0) |
| assert 0.0 <= result.coherence_score <= 1.0 |
|
|
| def test_without_kuramoto(self): |
| lm, tm = _build() |
| gen = CoherentGenerator(lm=lm, kuramoto=None, transition_model=tm) |
| result = gen.generate("hello", max_new_tokens=10) |
| assert isinstance(result.text, str) |
|
|
| def test_without_transition(self): |
| lm, _ = _build() |
| gen = CoherentGenerator(lm=lm, kuramoto=KuramotoAttractor(n_iterations=10), transition_model=None) |
| result = gen.generate("hello", max_new_tokens=10) |
| assert isinstance(result.text, str) |
|
|
| def test_without_both(self): |
| lm, _ = _build() |
| gen = CoherentGenerator(lm=lm, kuramoto=None, transition_model=None) |
| result = gen.generate("hello", max_new_tokens=10) |
| assert isinstance(result.text, str) |
|
|
| def test_greedy_coherent(self): |
| """At temp=0, output should match standard retrieval closely.""" |
| lm, tm = _build() |
| gen = CoherentGenerator(lm=lm, transition_model=tm) |
| standard = lm.respond("what is python", max_new_tokens=15) |
| coherent = gen.respond("what is python", max_new_tokens=15, temperature=0.0) |
| |
| assert len(coherent) > 0 or len(standard) > 0 |
|
|
| def test_confidence_tracked(self): |
| lm, tm = _build() |
| gen = CoherentGenerator(lm=lm, transition_model=tm) |
| result = gen.generate("what is python", max_new_tokens=10) |
| assert result.avg_confidence >= 0.0 |
|
|