constructelligence's picture
circuits-0.3 lite model: ONNX weights, decoder, example, model card
9e41b00 verified
Raw History Blame Contribute Delete
4.43 kB
#!/usr/bin/env python3
"""Read the circuits off an electrical floor plan with onnxruntime alone (no torch).
pip install onnxruntime numpy pillow
python predict.py plan.png [--symbol-px 16] [--out circuits.png] [--json circuits.json]
The model was trained on sheets where a duplex receptacle is ~12-24 px across. --symbol-px is the size of a
receptacle circle in YOUR image; the sheet is rescaled so symbols land in that range before inference.
"""
from __future__ import annotations
import argparse
import json
import pathlib
import numpy as np
import onnxruntime as ort
from PIL import Image, ImageDraw
import decode as D
HERE = pathlib.Path(__file__).resolve().parent
TILE, OVERLAP = 512, 64
TRAIN_SYMBOL_PX = 16.0
def run_sheet(sess, gray: np.ndarray):
"""Tile a grayscale sheet (0..255, HxW) through the model and stitch the stride-2 outputs by centre crop."""
H, W = gray.shape
gh, gw = (H + 1) // 2, (W + 1) // 2
peaks = np.zeros((len(D.CLASSES), gh, gw), np.float32)
size = np.zeros((2, gh, gw), np.float32)
wire = np.zeros((gh, gw), np.float32)
step = TILE - 2 * OVERLAP
for y0 in range(-OVERLAP, max(1, H - OVERLAP), step):
for x0 in range(-OVERLAP, max(1, W - OVERLAP), step):
tile = np.full((TILE, TILE), 255, np.float32)
ys, xs = max(0, y0), max(0, x0)
ye, xe = min(H, y0 + TILE), min(W, x0 + TILE)
tile[ys - y0:ye - y0, xs - x0:xe - x0] = gray[ys:ye, xs:xe]
p, s, w = sess.run(None, {"image": tile[None, None]})
# keep the centre (step x step) of each tile, in grid units
cy0, cx0 = (y0 + OVERLAP) // 2, (x0 + OVERLAP) // 2
n = step // 2
o = OVERLAP // 2
ya, xa = max(0, cy0), max(0, cx0)
yb, xb = min(gh, cy0 + n), min(gw, cx0 + n)
if yb <= ya or xb <= xa:
continue
sl = (slice(o + ya - cy0, o + yb - cy0), slice(o + xa - cx0, o + xb - cx0))
peaks[:, ya:yb, xa:xb] = p[0][:, sl[0], sl[1]]
size[:, ya:yb, xa:xb] = s[0][:, sl[0], sl[1]]
wire[ya:yb, xa:xb] = w[0, 0][sl]
return peaks, size, wire
def main():
ap = argparse.ArgumentParser()
ap.add_argument("image")
ap.add_argument("--model", default=str(HERE / "circuits.onnx"))
ap.add_argument("--symbol-px", type=float, default=20, help="receptacle circle diameter in the input (20 suits a 150-dpi PDF render)")
ap.add_argument("--out")
ap.add_argument("--json")
a = ap.parse_args()
img = Image.open(a.image).convert("L")
scale = TRAIN_SYMBOL_PX / a.symbol_px
if abs(scale - 1) > 0.05:
img = img.resize((round(img.width * scale), round(img.height * scale)), Image.LANCZOS)
sess = ort.InferenceSession(a.model, providers=["CPUExecutionProvider"])
peaks, size, wire = run_sheet(sess, np.asarray(img, np.float32))
res = D.decode(peaks, size, wire)
inv = 1 / scale
for d in res["devices"]:
d["box"] = [round(v * inv, 1) for v in d["box"]]
d["class"] = D.cname(d["cls"])
for c in res["circuits"]:
c["length_px"] = round(c["length_px"] * inv, 1)
for r in res["runs"]:
r.pop("points", None)
for k, c in enumerate(res["circuits"]):
kinds = [res["devices"][i]["class"] for i in c["devices"]]
print(f"circuit {k + 1}: {len(c['devices'])} devices ({', '.join(sorted(set(kinds)))}), "
f"home run: {'yes' if c['homeruns'] else 'no'}, ~{c['length_px']:.0f} px of wiring")
if res["unwired"]:
print(f"{len(res['unwired'])} powered symbols with no wiring found")
if a.json:
pathlib.Path(a.json).write_text(json.dumps(res, indent=1))
if a.out:
draw_result(Image.open(a.image).convert("RGB"), res).save(a.out)
PALETTE = [(31, 119, 180), (214, 39, 40), (44, 160, 44), (148, 103, 189), (255, 127, 14), (23, 190, 207),
(227, 119, 194), (140, 86, 75)]
def draw_result(img, res):
d = ImageDraw.Draw(img)
cid = {i: k for k, c in enumerate(res["circuits"]) for i in c["devices"]}
for dv in res["devices"]:
col = PALETTE[cid[dv["id"]] % len(PALETTE)] if dv["id"] in cid else (130, 130, 130)
d.rectangle(dv["box"], outline=col, width=3)
if dv["id"] in cid:
d.text((dv["box"][0], dv["box"][1] - 11), str(cid[dv["id"]] + 1), fill=col)
return img
if __name__ == "__main__":
main()