Supermix v80: 15.3M-parameter physics and arithmetic solver, solver-verified corpus
fd32dda verified | """Load supermix-v74 and ask it a question. | |
| python example_usage.py "A body of mass 47 kg has an acceleration of 6 m/s^2. What is the force?" | |
| The model answers arithmetic in the format it was trained on, so the question | |
| is normalised first -- see the "Prompt format matters" section of the model | |
| card. `prompt_normaliser` prints what it actually sent, and `answer_check` | |
| independently re-derives the answer so a wrong reply is reported as wrong | |
| rather than presented as fact. | |
| """ | |
| from __future__ import annotations | |
| import sys | |
| from pathlib import Path | |
| HERE = Path(__file__).resolve().parent | |
| sys.path.insert(0, str(HERE / "src")) | |
| import answer_check # noqa: E402 | |
| import prompt_normaliser # noqa: E402 | |
| from train_mimomix_talk import generate_reply, load_talk_checkpoint # noqa: E402 | |
| def main() -> int: | |
| question = " ".join(sys.argv[1:]) or "A body of mass 47 kg has an acceleration of 6 m/s^2. What is the force?" | |
| model, tokenizer, _ = load_talk_checkpoint(str(HERE / "supermix_v80.pt")) | |
| model.eval() | |
| rewritten = prompt_normaliser.normalise(question) | |
| if rewritten.changed: | |
| print(f"asked as : {rewritten.prompt} ({rewritten.rule})") | |
| result = generate_reply(model, tokenizer, rewritten.prompt, max_new_tokens=64) | |
| reply = result["reply"] if isinstance(result, dict) else str(result) | |
| print(f"reply : {reply}") | |
| verdict = answer_check.check(rewritten.prompt, reply) | |
| if verdict is None: | |
| print("check : NOT CHECKED (not a question this can verify)") | |
| elif verdict.correct: | |
| print(f"check : CORRECT ({verdict.expected})") | |
| else: | |
| print(f"check : WRONG (answered {verdict.predicted}, expected {verdict.expected})") | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |