{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# \ud83e\uddee 3-digit-basic-calc \u2014 3-digit arithmetic with a 1.6M-parameter transformer\n", "\n", "A from-scratch transformer that does `+ \u2212 \u00d7 \u00f7` on 3-digit numbers by writing out\n", "the algorithm step by step. This notebook loads the model and runs it.\n", "\n", "Model: [`vmal/3-digit-basic-calc`](https://huggingface.co/vmal/3-digit-basic-calc)\n" ] }, { "cell_type": "code", "metadata": {}, "execution_count": null, "outputs": [], "source": [ "%pip install -q transformers==5.3.0 torch safetensors\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Load the model\n" ] }, { "cell_type": "code", "metadata": {}, "execution_count": null, "outputs": [], "source": [ "from transformers import AutoModelForCausalLM, AutoTokenizer\n", "\n", "REPO = \"vmal/3-digit-basic-calc\"\n", "model = AutoModelForCausalLM.from_pretrained(REPO, trust_remote_code=True).eval()\n", "tok = AutoTokenizer.from_pretrained(REPO, trust_remote_code=True)\n", "print(\"parameters:\", sum(p.numel() for p in model.parameters()))\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Solve some problems\n", "\n", "`model.solve(tokenizer, expression)` returns the human-readable answer.\n" ] }, { "cell_type": "code", "metadata": {}, "execution_count": null, "outputs": [], "source": [ "for e in [\"842/37\", \"213*145\", \"999*999\", \"-500+500\", \"3/31\", \"12/0\"]:\n", " print(f\"{e:>10} = {model.solve(tok, e)}\")\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## See the model's reasoning\n", "\n", "Pass `return_trace=True` to get the raw scratchpad the model generates \u2014\n", "each ``/``/`` is one local computation.\n" ] }, { "cell_type": "code", "metadata": {}, "execution_count": null, "outputs": [], "source": [ "answer, trace = model.solve(tok, \"842/37\", return_trace=True)\n", "print(\"answer:\", answer)\n", "print()\n", "for c in ['
','','','','','','','','','']:\n", " trace = trace.replace(c, '\\n'+c+' ')\n", "print(trace.strip())\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Quick in-range behavior check on random problems\n", "\n", "This small smoke test samples the supported operand range. It does not load\n", "the training prompts, so it should not be described as a leakage-controlled\n", "unseen benchmark; the card reports that benchmark separately.\n" ] }, { "cell_type": "code", "metadata": {}, "execution_count": null, "outputs": [], "source": [ "import random\n", "random.seed(0)\n", "def truth(a, b, op):\n", " if op == '+': return str(a + b)\n", " if op == '-': return str(a - b)\n", " if op == '*': return str(a * b)\n", " if b == 0: return 'NAN'\n", " # Exact integer round-half-up to three decimals (no float/banker's rounding).\n", " denominator = abs(b)\n", " scaled, remainder = divmod(abs(a) * 1000, denominator)\n", " scaled += int(2 * remainder >= denominator)\n", " integer, fraction = divmod(scaled, 1000)\n", " magnitude = str(integer)\n", " if fraction:\n", " magnitude += '.' + f'{fraction:03d}'.rstrip('0')\n", " negative = (a < 0) != (b < 0)\n", " return '-' + magnitude if negative and scaled else magnitude\n", "good=n=0\n", "for _ in range(30):\n", " a=random.randint(-999,999); b=random.randint(-999,999); op=random.choice('+-*/')\n", " if op=='/' and b==0: b=7\n", " got=model.solve(tok, f'{a}{op}{b}'); exp=truth(a,b,op)\n", " good+= (got==exp); n+=1\n", "print(f'{good}/{n} correct on random in-range problems')\n" ] } ], "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python" } }, "nbformat": 4, "nbformat_minor": 5 }