Piko-9b / tests /test_text_generation.py
Dexy2's picture
Rewrite model card around verified evidence; correct misattributed benchmarks and config path leak
0810902 verified
Raw
History Blame Contribute Delete
4.21 kB
"""Text generation behaviour.
Every test here loads weights, so the whole module is marked slow.
"""
from __future__ import annotations
import pytest
pytestmark = pytest.mark.slow
def generate(loaded_model, prompt: str, max_new_tokens: int = 64) -> str:
import torch
model, processor = loaded_model
inputs = processor.apply_chat_template(
[{"role": "user", "content": [{"type": "text", "text": prompt}]}],
add_generation_prompt=True,
tokenize=True,
return_dict=True,
return_tensors="pt",
).to(model.device)
with torch.inference_mode():
output = model.generate(**inputs, max_new_tokens=max_new_tokens, do_sample=False)
return processor.decode(
output[0][inputs["input_ids"].shape[1] :], skip_special_tokens=True
).strip()
def visible(text: str) -> str:
return text.rsplit("</think>", 1)[-1].strip() if "</think>" in text else text
def test_generation_is_not_degenerate(loaded_model) -> None:
"""Guard against the CPU-offload failure mode: a single repeated token.
With device_map='auto' and offload, this model emits '!!!!!!...'. That is the
single most likely way a user's deployment silently breaks.
"""
text = generate(loaded_model, "What is the capital of France?")
assert text, "model produced no output"
distinct = set(text.replace(" ", ""))
assert len(distinct) > 5, f"degenerate output, only {distinct!r} — check device placement"
def test_answers_a_simple_factual_question(loaded_model) -> None:
assert "paris" in visible(generate(loaded_model, "What is the capital of France?")).lower()
def test_arithmetic(loaded_model) -> None:
assert "391" in generate(loaded_model, "What is 17 * 23? Number only.", 48)
def test_greedy_decoding_is_deterministic(loaded_model) -> None:
prompt = "Name three primary colours."
assert generate(loaded_model, prompt, 40) == generate(loaded_model, prompt, 40)
def test_multi_turn_context_is_carried(loaded_model) -> None:
import torch
model, processor = loaded_model
messages = [
{"role": "user", "content": [{"type": "text", "text": "My favourite number is 47."}]},
{"role": "assistant", "content": [{"type": "text", "text": "Noted."}]},
{
"role": "user",
"content": [{"type": "text", "text": "Double my favourite number. Number only."}],
},
]
inputs = processor.apply_chat_template(
messages,
add_generation_prompt=True,
tokenize=True,
return_dict=True,
return_tensors="pt",
).to(model.device)
with torch.inference_mode():
output = model.generate(**inputs, max_new_tokens=48, do_sample=False)
text = processor.decode(output[0][inputs["input_ids"].shape[1] :], skip_special_tokens=True)
assert "94" in text
def test_batch_generation_matches_single(loaded_model) -> None:
"""Left padding must not change the answer for the shorter prompt."""
import torch
model, processor = loaded_model
prompts = ["What is 2 + 2? Number only.", "Name the capital of Japan. Name only."]
texts = [
processor.apply_chat_template(
[{"role": "user", "content": [{"type": "text", "text": p}]}],
add_generation_prompt=True,
tokenize=False,
)
for p in prompts
]
inputs = processor(text=texts, return_tensors="pt", padding=True).to(model.device)
with torch.inference_mode():
output = model.generate(**inputs, max_new_tokens=40, do_sample=False)
decoded = [
processor.decode(row[inputs["input_ids"].shape[1] :], skip_special_tokens=True)
for row in output
]
assert "4" in decoded[0]
assert "tokyo" in decoded[1].lower()
def test_long_input_is_accepted(loaded_model) -> None:
"""A needle at 8K tokens should at minimum not crash the model."""
needle = "The maintenance code for the north pump is QF-8812."
filler = "Routine log entry: all systems nominal. " * 700
prompt = f"{filler}\n{needle}\n{filler}\n\nWhat is the maintenance code for the north pump?"
text = generate(loaded_model, prompt, 40)
assert text, "no output for long input"