Download predict.py from constructelligence/electrical-circuit-connectivity: direct link, hf CLI and curl.
- Browser
- Download file 4.43 kB
-
https://huggingface.co/constructelligence/electrical-circuit-connectivity/resolve/main/predict.py
- Command line
-
hf download hf://constructelligence/electrical-circuit-connectivity/predict.py
-
curl -L -o predict.py https://huggingface.co/constructelligence/electrical-circuit-connectivity/resolve/main/predict.py
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() | |