File size: 8,026 Bytes
f915d94
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
"""LLM explanation layer: turns the calculator's structured record into plain prose.

The LLM never computes anything. It receives the record text, is told to reuse only
those numbers, and its answer is then audited by a deterministic grounding check.
"""
import re

import calculator as C

MODEL_ID = "Qwen/Qwen3-0.6B"
MODEL_REVISION = "c1899de289a04d12100db370d81485cdf75e47ca"  # pinned, same as class notebook 05
MAX_NEW_TOKENS = 200

SYSTEM_PROMPT = (
    "You explain the result of an engineering calculation to someone who is not a control "
    "engineer. Use only numbers that appear in the record, written exactly as they appear there. "
    "Never calculate new numbers. In 4 or 5 plain sentences, describe what the system does after "
    "the step (rise, overshoot, oscillation, settling), say what the overshoot and settling time "
    "mean for this output, and state the result of every design check. For a failed check, give "
    "the requirement the record lists for passing it. If no design checks were requested, say so. "
    "No lists, headings, or formulas."
)

# Worked examples come from real calculator runs so their numbers are exact
_EX_A = C.compute("mass-spring-damper", mass=2, damping=8, stiffness=200, force=20,
                  band_name="2%", os_limit=20, ts_limit=3)
_EX_A_ANSWER = (
    "After the 20 N step, the mass heads toward its final displacement of 100 mm, but with a "
    "damping ratio of only 0.2 it overshoots to a peak of 152.7 mm at 0.3206 s, a 52.66 % overshoot. "
    "It then oscillates with a period of 0.6413 s as the swings die away, and it stays within 2% of "
    "the final position after 1.96 s. The overshoot check fails because 52.66 % is above the 20 % "
    "limit; the damping ratio would need to be at least 0.4559 to pass. The settling check passes, "
    "since 1.96 s is under the 3 s limit."
)
_EX_B = C.compute("normalised", wn=5, zeta=0.7, amplitude=1, gain=1, band_name="5%")
_EX_B_ANSWER = (
    "After the step, the output rises toward its final value of 1 and overshoots only slightly, "
    "peaking at 1.046 at 0.8798 s, a 4.599 % overshoot. Because that overshoot is smaller than the "
    "5% band, the output enters the band at 0.58 s and never leaves it, so it counts as settled "
    "before it even reaches its peak. That is why the textbook estimate of 0.8571 s is noticeably "
    "longer than the exact settling time. No design checks were requested."
)
FEW_SHOT = [
    {"role": "user", "content": C.to_text(_EX_A)},
    {"role": "assistant", "content": _EX_A_ANSWER},
    {"role": "user", "content": C.to_text(_EX_B)},
    {"role": "assistant", "content": _EX_B_ANSWER},
]

_tokenizer = None
_model = None


def load_model():
    """Load Qwen3-0.6B once on CPU; later calls reuse the cached objects."""
    global _tokenizer, _model
    if _model is None:
        import torch
        from transformers import AutoModelForCausalLM, AutoTokenizer
        _tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, revision=MODEL_REVISION)
        _model = AutoModelForCausalLM.from_pretrained(
            MODEL_ID, revision=MODEL_REVISION, dtype=torch.float32
        ).eval()  # float32 on CPU avoids half-precision kernels a free Space lacks
    return _tokenizer, _model


def build_messages(record):
    """System rules, two worked examples, then the current record text."""
    return [{"role": "system", "content": SYSTEM_PROMPT}, *FEW_SHOT,
            {"role": "user", "content": C.to_text(record)}]


def generate(messages, max_tokens=MAX_NEW_TOKENS):
    """Greedy decoding so the same record always yields the same explanation."""
    import torch
    from transformers import GenerationConfig
    tokenizer, model = load_model()
    inputs = tokenizer.apply_chat_template(
        messages, tokenize=True, add_generation_prompt=True,
        enable_thinking=False, return_dict=True, return_tensors="pt",
    )
    config = GenerationConfig(
        max_new_tokens=max_tokens, do_sample=False,
        eos_token_id=tokenizer.eos_token_id, pad_token_id=tokenizer.eos_token_id,
    )
    model.generation_config = config  # stop repo sampling defaults overriding greedy
    with torch.inference_mode():
        output = model.generate(**inputs, generation_config=config)
    new_tokens = output[0, inputs["input_ids"].shape[-1]:]  # drop the prompt tokens
    text = tokenizer.decode(new_tokens, skip_special_tokens=True).strip()
    if new_tokens[-1].item() != tokenizer.eos_token_id:
        text += " [Explanation hit its length limit and may be incomplete.]"
    return text


def template_explanation(record):
    """Deterministic fallback used only if the language model cannot run."""
    m = record["metrics"]
    unit = f" {record['output_unit']}" if record["output_unit"] else ""
    parts = [
        f"The {record['output_name']} rises toward its final value of {C.fmt(m['final']['value'])}{unit}, "
        f"peaks at {C.fmt(m['peak_value']['value'])}{unit} after {C.fmt(m['peak_time']['value'])} s "
        f"({C.fmt(m['overshoot']['value'])} % overshoot), and stays inside the "
        f"{record['settling_band']} band after {C.fmt(m['ts_exact']['value'])} s."
    ]
    for ch in record["checks"]:
        req = ch["requirement"]
        verdict = "passes" if ch["result"] == "PASS" else "fails"
        parts.append(f"The {ch['name'].lower()} check {verdict} ({C.fmt(ch['value'])} {ch['unit']} "
                     f"against a limit of {C.fmt(ch['limit'])} {ch['unit']}).")
        if ch["result"] == "FAIL":
            parts.append(f"To pass, the {req['label'].lower()} is {C.fmt(req['value'])} {req['unit']}.".replace(" .", "."))
    if not record["checks"]:
        parts.append("No design checks were requested.")
    return " ".join(parts)


def explain(record):
    """Return (explanation, source) where source says which path produced it."""
    try:
        return generate(build_messages(record)), f"{MODEL_ID} (greedy, few-shot)"
    except Exception as exc:  # model download or runtime failure should not break the app
        return template_explanation(record), f"template fallback ({type(exc).__name__})"


_NUM = re.compile(r"\d+(?:\.\d+)?")
# Numbers that describe the method rather than this result (2% band, 10% to 90%, ...)
_METHOD_NUMBERS = {0.0, 1.0, 2.0, 5.0, 10.0, 90.0, 100.0}


def _record_numbers(record):
    """Every number printed in the record text, the only numbers the LLM may use."""
    return {float(x) for x in _NUM.findall(C.to_text(record))} | _METHOD_NUMBERS


def _supported(token, allowed):
    """A number is supported if it equals a record number rounded to its own precision."""
    value = float(token)
    decimals = len(token.split(".")[1]) if "." in token else 0
    tol = 0.5 * 10 ** (-decimals) + 1e-9  # half a unit in the last written digit
    return any(abs(value - a) <= tol for a in allowed)


def grounding_check(text, record):
    """Audit the explanation: unsupported numbers and pass/fail contradictions."""
    allowed = _record_numbers(record)
    tokens = _NUM.findall(text)
    unsupported = sorted({t for t in tokens if not _supported(t, allowed)}, key=float)

    contradictions = []
    sentences = re.split(r"(?<=[.!?])\s+", text.lower())
    keywords = {"Overshoot": "overshoot", "Settling time": "settl"}
    for ch in record["checks"]:
        key = keywords[ch["name"]]
        for s in sentences:
            if key not in s or "check" not in s:
                continue  # only judge sentences that talk about this check
            says_pass = "pass" in s and "fail" not in s
            says_fail = "fail" in s and "pass" not in s
            if (ch["result"] == "PASS" and says_fail) or (ch["result"] == "FAIL" and says_pass):
                contradictions.append(f"{ch['name']} is {ch['result']} in the record, the text disagrees")

    return {
        "numbers_found": len(tokens),
        "unsupported": unsupported,
        "contradictions": contradictions,
        "grounded": not unsupported and not contradictions,
    }