File size: 4,212 Bytes
0810902
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
106
107
108
109
110
111
112
113
114
115
116
117
118
"""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"