mathcore / mathcore.py
not-abdo's picture
Upload 7 files
9906a3d verified
Raw
History Blame Contribute Delete
15.9 kB
"""MathCore — local CPU inference (single file).
python mathcore.py --chat # terminal REPL
python mathcore.py --service --port 8000 # FastAPI service
python mathcore.py --export-onnx # ckpt -> mathcore.onnx (needs torch, once)
Backends: --backend auto|torch|onnx (auto = onnx if mathcore.onnx exists, else torch)
جهاز ضعيف؟ اعمل export مرة واحدة (على Kaggle مثلاً)، انقل mathcore.onnx،
وشغّل بـ onnxruntime + numpy بس — من غير PyTorch خالص.
pip (torch backend): torch
pip (onnx backend): onnxruntime numpy
pip (service): fastapi uvicorn
"""
import argparse
import os
import re
import sys
import time
from dataclasses import dataclass, field
import numpy as np
# ---------------------------------------------------------------- vocab
TOK_PLUS, TOK_MINUS, TOK_Q, TOK_A, TOK_PAD, TOK_ANS = 10, 11, 12, 13, 14, 15
ANS_PAD = 10
# ---- configs: نفس أسماء تدريب Kaggle عشان unpickle بتاع الـ ckpt يلاقيها ----
@dataclass
class DataConfig:
max_digits: int = 15
ood_min_digits: int = 16
ood_max_digits: int = 24
abacus_max: int = 128
len_buckets: tuple = ((1, 3), (4, 6), (7, 10), (11, 15))
carry_qs: tuple = (0.1, 0.5, 0.9)
p_edge: float = 0.05
p_family: float = 0.25
val_mod: int = 1000
val_lt: int = 10
val_set_size: int = 512
@dataclass
class ModelConfig:
vocab_size: int = 16
d_model: int = 640
n_heads: int = 10
d_ff: int = 1728
n_prelude: int = 2
n_core: int = 4
n_coda: int = 2
r_steps: int = 3
dropout: float = 0.0
ans_vocab: int = 11
@dataclass
class TrainConfig:
steps: int = 20000
batch_size: int = 1024
lr: float = 3e-4
@dataclass
class RunConfig:
data: DataConfig = field(default_factory=DataConfig)
model: ModelConfig = field(default_factory=ModelConfig)
train: TrainConfig = field(default_factory=TrainConfig)
# ---------------------------------------------------------------- rendering (numpy)
def int_digits_lsd(n: int):
"""Non-negative int -> LSD-first digit list."""
return [int(c) for c in reversed(str(n))]
def build_inputs(a: int, b: int, op: str, W: int):
"""Fixed layout: [Q] a-field(W) [op] b-field(W) [A] slots(W+1). batch=1."""
da, db = int_digits_lsd(a), int_digits_lsd(b)
if len(da) > W or len(db) > W:
raise ValueError(f"operand exceeds {W} digits")
S = 3 * W + 4
slots = W + 1
tokens = np.full((1, S), TOK_PAD, dtype=np.int64)
abacus = np.zeros((1, S), dtype=np.int64)
role = np.zeros((1, S), dtype=np.int64)
def put_number(d, start, r):
l = len(d)
for j in range(l): # MSD-first, left aligned
tokens[0, start + j] = d[l - 1 - j]
abacus[0, start + j] = l - j
role[0, start + j] = r
tokens[0, 0] = TOK_Q
put_number(da, 1, 1)
tokens[0, 1 + W] = TOK_PLUS if op == "+" else TOK_MINUS
put_number(db, 2 + W, 2)
tokens[0, 2 + 2 * W] = TOK_A
s0 = 3 + 2 * W
tokens[0, s0:s0 + slots] = TOK_ANS
abacus[0, s0:s0 + slots] = np.arange(1, slots + 1)
role[0, s0:s0 + slots] = 3
pad_mask = tokens != TOK_PAD
return tokens, abacus, role, pad_mask, slots
def decode(logits: np.ndarray):
"""[1, M, 11] -> (value:int, min confidence over used slots)."""
z = logits[0]
e = np.exp(z - z.max(axis=-1, keepdims=True))
p = e / e.sum(axis=-1, keepdims=True)
pred = z.argmax(axis=-1)
digits = []
used = []
for j, d in enumerate(pred):
used.append(p[j, d])
if d == ANS_PAD:
break
digits.append(int(d))
if not digits:
return 0, float(min(used)) if used else 0.0
val = int("".join(str(d) for d in reversed(digits)))
return val, float(min(used))
# ---------------------------------------------------------------- torch stack (lazy)
def _torch_stack():
import torch
import torch.nn as nn
import torch.nn.functional as F
class RMSNorm(nn.Module):
def __init__(self, d, eps=1e-6):
super().__init__()
self.w = nn.Parameter(torch.ones(d)); self.eps = eps
def forward(self, x):
return self.w * x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
class Attention(nn.Module):
def __init__(self, d, h):
super().__init__()
self.h, self.hd = h, d // h
self.qkv = nn.Linear(d, 3 * d, bias=False)
self.o = nn.Linear(d, d, bias=False)
self.qn = RMSNorm(self.hd); self.kn = RMSNorm(self.hd)
def forward(self, x, add_mask):
B, S, D = x.shape
q, k, v = self.qkv(x).chunk(3, dim=-1)
q = self.qn(q.view(B, S, self.h, self.hd)).transpose(1, 2)
k = self.kn(k.view(B, S, self.h, self.hd)).transpose(1, 2)
v = v.view(B, S, self.h, self.hd).transpose(1, 2)
y = F.scaled_dot_product_attention(q, k, v, attn_mask=add_mask)
return self.o(y.transpose(1, 2).reshape(B, S, D))
class SwiGLU(nn.Module):
def __init__(self, d, dff):
super().__init__()
self.w1 = nn.Linear(d, dff, bias=False)
self.w2 = nn.Linear(d, dff, bias=False)
self.w3 = nn.Linear(dff, d, bias=False)
def forward(self, x):
return self.w3(F.silu(self.w1(x)) * self.w2(x))
class Block(nn.Module):
def __init__(self, d, h, dff):
super().__init__()
self.n1, self.n2 = RMSNorm(d), RMSNorm(d)
self.attn = Attention(d, h); self.mlp = SwiGLU(d, dff)
def forward(self, x, m):
x = x + self.attn(self.n1(x), m)
return x + self.mlp(self.n2(x))
class MathCore(nn.Module):
def __init__(self, mcfg, abacus_max):
super().__init__()
d, h, dff = mcfg.d_model, mcfg.n_heads, mcfg.d_ff
self.tok = nn.Embedding(mcfg.vocab_size, d)
self.abacus = nn.Embedding(abacus_max + 1, d, padding_idx=0)
self.role = nn.Embedding(4, d, padding_idx=0)
self.prelude = nn.ModuleList(Block(d, h, dff) for _ in range(mcfg.n_prelude))
self.core = nn.ModuleList(Block(d, h, dff) for _ in range(mcfg.n_core))
self.coda = nn.ModuleList(Block(d, h, dff) for _ in range(mcfg.n_coda))
self.inject = RMSNorm(d); self.out_norm = RMSNorm(d)
self.ans_head = nn.Linear(d, mcfg.ans_vocab, bias=False)
self.carry_head = nn.Linear(d, 2, bias=False)
self.r_default = mcfg.r_steps
def forward(self, tokens, abacus, role, pad_mask, ans_slots, r=None):
r = r or self.r_default
# additive float mask (ONNX-friendly, math-equivalent to bool mask)
m = (~pad_mask)[:, None, None, :].float() * -1e9
e = self.tok(tokens) + self.abacus(abacus) + self.role(role)
for blk in self.prelude:
e = blk(e, m)
s = torch.zeros_like(e)
outs = []
for _ in range(r):
s = self.inject(s + e)
for blk in self.core:
s = blk(s, m)
h = s
for blk in self.coda:
h = blk(h, m)
outs.append(self.ans_head(self.out_norm(h[:, -ans_slots:])))
return outs, None
return torch, MathCore
# ---------------------------------------------------------------- backends
def _register_cfg_classes():
"""الـ ckpt بيعمل pickle للكلاسات دي تحت __main__ (كذا اتحفظ في Kaggle) —
نسجلها هناك عشان التحميل يشتغل سواء الملف اتشغّل مباشرة أو اتعمله import."""
import __main__ as _m
for cls in (DataConfig, ModelConfig, TrainConfig, RunConfig):
if not hasattr(_m, cls.__name__):
setattr(_m, cls.__name__, cls)
class TorchBackend:
name = "torch"
def __init__(self, ckpt, r):
torch, MathCore = _torch_stack()
self.torch = torch
_register_cfg_classes()
ck = torch.load(ckpt, map_location="cpu", weights_only=False)
cfg = ck["cfg"]
self.model = MathCore(cfg.model, cfg.data.abacus_max)
self.model.load_state_dict(ck["model"])
self.model.eval()
self.r = r
n = sum(p.numel() for p in self.model.parameters())
print(f"[torch] loaded {ckpt} | {n/1e6:.1f}M params | R={r}")
def infer(self, tokens, abacus, role, pad_mask, slots):
t = self.torch
with t.inference_mode():
outs, _ = self.model(t.from_numpy(tokens), t.from_numpy(abacus),
t.from_numpy(role), t.from_numpy(pad_mask), slots, r=self.r)
return outs[-1].numpy()
class OnnxBackend:
name = "onnx"
def __init__(self, path):
import onnxruntime as ort
so = ort.SessionOptions()
self.sess = ort.InferenceSession(path, so, providers=["CPUExecutionProvider"])
self.W = int(self.sess.get_modelmeta().custom_metadata_map.get("W", "30"))
print(f"[onnx] loaded {path} | W={self.W}")
def infer(self, tokens, abacus, role, pad_mask, slots):
return self.sess.run(["logits"], {"tokens": tokens, "abacus": abacus,
"role": role, "pad_mask": pad_mask})[0]
def export_onnx(ckpt, out, W, r):
torch, MathCore = _torch_stack()
import torch.nn as nn
_register_cfg_classes()
ck = torch.load(ckpt, map_location="cpu", weights_only=False)
cfg = ck["cfg"]
core = MathCore(cfg.model, cfg.data.abacus_max)
core.load_state_dict(ck["model"]); core.eval()
slots = W + 1
class Wrap(nn.Module):
def __init__(self):
super().__init__(); self.m = core
def forward(self, tokens, abacus, role, pad_mask):
outs, _ = self.m(tokens, abacus, role, pad_mask, slots, r=r)
return outs[-1]
t, a, ro, pm, _ = build_inputs(123, 45, "+", W)
args = (torch.from_numpy(t), torch.from_numpy(a), torch.from_numpy(ro), torch.from_numpy(pm))
torch.onnx.export(Wrap(), args, out, opset_version=17, dynamo=False,
input_names=["tokens", "abacus", "role", "pad_mask"],
output_names=["logits"])
import onnx
m = onnx.load(out)
meta = m.metadata_props.add(); meta.key = "W"; meta.value = str(W)
onnx.save(m, out)
print(f"[export] {out} | W={W} (operands up to {W} digits) | R={r} baked in")
# ---------------------------------------------------------------- expression engine
class Solver:
def __init__(self, backend, W):
self.be, self.W = backend, W
self.calls = 0
def _binop(self, a, b, op):
"""a,b >= 0; for '-' requires a >= b. Returns (value, conf, matches_python)."""
inputs = build_inputs(a, b, op, self.W)
logits = self.be.infer(*inputs)
val, conf = decode(logits)
self.calls += 1
truth = a + b if op == "+" else a - b
return val, conf, (val == truth)
def _signed_add(self, x, y, steps):
"""Signed x + y via unsigned model calls with sign bookkeeping."""
if x >= 0 and y >= 0:
v, c, ok = self._binop(x, y, "+"); steps.append((f"{x}+{y}", v, c, ok)); return v
if x < 0 and y < 0:
v, c, ok = self._binop(-x, -y, "+"); steps.append((f"{-x}+{-y}", v, c, ok)); return -v
pos, neg = (x, y) if x >= 0 else (y, x)
n = -neg
if pos >= n:
v, c, ok = self._binop(pos, n, "-"); steps.append((f"{pos}-{n}", v, c, ok)); return v
v, c, ok = self._binop(n, pos, "-"); steps.append((f"{n}-{pos}", v, c, ok)); return -v
def eval(self, expr):
s = expr.replace(" ", "")
if not re.fullmatch(r"-?\d+([+-]\d+)*", s):
raise ValueError("صيغة غير مفهومة — المدعوم: أرقام و + و - (مثال: 999+1000-20+5)")
toks = re.findall(r"\d+|[+-]", s)
i = 0
sign = 1
if toks[0] == "-":
sign, i = -1, 1
acc = sign * int(toks[i]); i += 1
steps = []
self.calls = 0
while i < len(toks):
op, num = toks[i], int(toks[i + 1]); i += 2
term = num if op == "+" else -num
if abs(acc) > 10 ** self.W - 1 or num > 10 ** self.W - 1:
raise ValueError(f"قيمة تعدت حد الموديل ({self.W} خانة)")
acc = self._signed_add(acc, term, steps)
verified = all(ok for (_, _, _, ok) in steps) if steps else True
min_conf = min((c for (_, _, c, _) in steps), default=1.0)
return acc, steps, verified, min_conf
# ---------------------------------------------------------------- modes
def chat(solver):
print("MathCore Infinite Chat")
print("Examples:\n1+2-3+10+5\n999+1000-20+5\nexit\n")
while True:
try:
expr = input("You: ").strip()
except (EOFError, KeyboardInterrupt):
print(); break
if expr.lower() in ("exit", "quit", ""):
if expr.lower() in ("exit", "quit"): break
continue
try:
t0 = time.time()
val, steps, verified, conf = solver.eval(expr)
ms = (time.time() - t0) * 1000
mark = "✓" if verified else "✗ (اختلف عن الحساب المؤكد!)"
print(f"Bot: {val} {mark} [{len(steps)} ops, {ms:.0f}ms, conf {conf:.3f}]")
except ValueError as e:
print(f"Bot: {e}")
def service(solver, port):
try:
from fastapi import FastAPI, HTTPException
import uvicorn
except ImportError:
sys.exit("pip install fastapi uvicorn")
app = FastAPI(title="MathCore", version="0.2")
@app.get("/health")
def health():
return {"status": "ok", "backend": solver.be.name, "max_digits": solver.W}
@app.get("/solve")
def solve(expr: str):
try:
t0 = time.time()
val, steps, verified, conf = solver.eval(expr)
return {"expr": expr, "result": str(val), "verified": verified,
"confidence": round(conf, 4), "model_calls": len(steps),
"steps": [{"op": s, "out": v, "conf": round(c, 4), "ok": ok}
for (s, v, c, ok) in steps],
"ms": round((time.time() - t0) * 1000, 1)}
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
uvicorn.run(app, host="0.0.0.0", port=port)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--chat", action="store_true")
ap.add_argument("--service", action="store_true")
ap.add_argument("--export-onnx", action="store_true")
ap.add_argument("--ckpt", default="mathcore_ckpt.pt")
ap.add_argument("--onnx-path", default="mathcore.onnx")
ap.add_argument("--backend", default="auto", choices=["auto", "torch", "onnx"])
ap.add_argument("--r", type=int, default=2, help="recurrence steps (2 = best OOD/speed)")
ap.add_argument("--max-digits", type=int, default=30)
ap.add_argument("--port", type=int, default=8000)
args = ap.parse_args()
if args.export_onnx:
export_onnx(args.ckpt, args.onnx_path, args.max_digits, args.r)
return
if args.backend == "auto":
args.backend = "onnx" if os.path.exists(args.onnx_path) else "torch"
if args.backend == "onnx":
be = OnnxBackend(args.onnx_path)
W = be.W
else:
be = TorchBackend(args.ckpt, args.r)
W = args.max_digits
solver = Solver(be, W)
if args.service:
service(solver, args.port)
else:
chat(solver)
if __name__ == "__main__":
main()