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)
|