palimpseste-max / tests /test_creative.py
thefinalboss's picture
Upload tests/test_creative.py with huggingface_hub
101d2a1 verified
Raw
History Blame Contribute Delete
4.39 kB
"""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)