File size: 4,433 Bytes
9e41b00
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
#!/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()