File size: 3,225 Bytes
92953bb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
"""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)
        # Should be very similar (coherence dominant)
        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