File size: 4,387 Bytes
101d2a1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
"""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)