"""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("", 1)[-1].strip() if "" 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"