step-response-explainer / explainer.py
jackstev's picture
Deploy HW3 explainer app
f915d94 verified
Raw History Blame Contribute Delete
8.03 kB
"""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,
}