Download explainer.py from jackstev/step-response-explainer: direct link, hf CLI and curl.
- Browser
- Download file 8.03 kB
-
https://huggingface.co/spaces/jackstev/step-response-explainer/resolve/main/explainer.py
- Command line
-
hf download hf://spaces/jackstev/step-response-explainer/explainer.py
-
curl -L -o explainer.py https://huggingface.co/spaces/jackstev/step-response-explainer/resolve/main/explainer.py
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, | |
| } | |