"""CircuitNet outputs -> devices, wiring runs, circuits. Pipeline (all wiring work happens on the stride-2 output grid): 1. symbols peaks > per-class threshold -> boxes; same-place duplicates across classes suppressed 2. cut wire mask zeroed inside every device box, so each drawn run becomes its own stroke 3. thin Zhang-Suen skeleton 4. graph endpoints (crossing number 1) and junction clusters (crossing number >= 3), branches between 5. repair junctions joined by a short bridge merge (an X usually thins into two Ts); short free spurs pruned 6. crossings at a 4-way junction the two straightest continuations pair up: runs that cross on paper without a dot stay separate circuits. 3-way and 5+ junctions join everything (conservative). 7. attach stroke ends within ATTACH of a symbol box land on that symbol; an end at an arrowhead is a home run 8. circuits union-find over devices joined by runs """ from __future__ import annotations import math import numpy as np CLASSES = [ "receptacle", "gfci_receptacle", "switch", "switch_3way", "ceiling_fixture", "downlight", "troffer", "strip_light", "exit_sign", "junction_box", "panelboard", "data_outlet", "homerun_arrow", ] ARROW = CLASSES.index("homerun_arrow") VIRTUAL = len(CLASSES) # not a model channel: a device inferred from where a drawn run ends UNWIRED = ("data_outlet", "panelboard") def cname(c): return CLASSES[c] if c < len(CLASSES) else "unrecognized" STRIDE = 2 SIZE_REF = 16.0 CFG = { "heat_thr": 0.4, # tuned on real dev sheets 2026-09-27 (was 0.3): fewer weak false symbols "wire_thr": 0.5, "cut_pad": 1, # grid px added around a symbol box before cutting the wire mask "attach": 8.0, # grid px: a stroke end this close to a box lands on it (tuned on real dev sheets; was 6) "bridge": 4, # grid px: junction clusters joined by a branch this short are one junction "spur": 6, # grid px: free-ended branches shorter than this are drawing noise (hash marks, tails) "dir_len": 7, # grid px along a branch used to measure its direction at a junction "pair_cos": -0.5, # a 4-way / 3-way pairing is trusted only if the pair is at least this opposite "t3": "pair", # 3-way junction away from symbols: "pair" joins the straight-through arms, "all" joins all "through_cos": -0.8, # arms this opposite at a junction on a symbol are a run crossing past it "arrow_touch": 1.5, # grid px: an end this close to an arrowhead is a home-run tip regardless of symbols "arrow_reach": 16, # grid px: an arrowhead no stroke reached belongs to the nearest symbol this close "short_gap": 8, # grid px: side-by-side symbols this close are tested for a run straight across the gap "virtual": True, # stroke ends that reach no detected symbol become unrecognised devices "virtual_cluster": 10, # grid px: free ends this close are the same unrecognised symbol (runs arrive from sides) "virtual_border": 3, # grid px: free ends this close to the image edge are runs leaving the view, not devices "arrow_tip": 8, # grid px: a free end this close to an arrowhead is its tip, never an unrecognised device } # ------------------------------------------------------------------------------------------------ symbols def detect(peaks, size, cfg=CFG): """peaks [K,H,W] (already 3x3-NMS'd), size [2,H,W] -> list of detections in input px.""" K, H, W = peaks.shape dets = [] ks, ys, xs = np.nonzero(peaks > cfg["heat_thr"]) for k, y, x in zip(ks.tolist(), ys.tolist(), xs.tolist()): cx, cy = (x + 0.5) * STRIDE, (y + 0.5) * STRIDE w = SIZE_REF * math.exp(float(size[0, y, x])) h = SIZE_REF * math.exp(float(size[1, y, x])) dets.append({"cls": k, "score": float(peaks[k, y, x]), "cx": cx, "cy": cy, "box": [cx - w / 2, cy - h / 2, cx + w / 2, cy + h / 2]}) return suppress(dets) def suppress(dets): dets = sorted(dets, key=lambda d: (-d["score"], d["cy"], d["cx"], d["cls"])) keep = [] for d in dets: dup = False for k in keep: if (d["cls"] == ARROW) != (k["cls"] == ARROW): continue b = k["box"] r = 0.35 * min(b[2] - b[0], b[3] - b[1], d["box"][2] - d["box"][0], d["box"][3] - d["box"][1]) if math.hypot(d["cx"] - k["cx"], d["cy"] - k["cy"]) < max(3.0, r): dup = True break if not dup: keep.append(d) for i, d in enumerate(keep): d["id"] = i return keep # ------------------------------------------------------------------------------------------------ skeleton def zhang_suen(img): """img: uint8 {0,1} [H,W]. Returns the 8-connected skeleton (standard two-subiteration Zhang-Suen).""" a = np.pad(img.astype(np.uint8), 1) while True: changed = False for step in (0, 1): P2, P3, P4 = a[:-2, 1:-1], a[:-2, 2:], a[1:-1, 2:] P5, P6, P7 = a[2:, 2:], a[2:, 1:-1], a[2:, :-2] P8, P9 = a[1:-1, :-2], a[:-2, :-2] nb = [P2, P3, P4, P5, P6, P7, P8, P9] B = sum(n.astype(np.int16) for n in nb) A = sum(((nb[i] == 0) & (nb[(i + 1) % 8] == 1)).astype(np.int16) for i in range(8)) c = a[1:-1, 1:-1] == 1 if step == 0: m = c & (B >= 2) & (B <= 6) & (A == 1) & ((P2 * P4 * P6) == 0) & ((P4 * P6 * P8) == 0) else: m = c & (B >= 2) & (B <= 6) & (A == 1) & ((P2 * P4 * P8) == 0) & ((P2 * P6 * P8) == 0) if m.any(): a[1:-1, 1:-1][m] = 0 changed = True if not changed: return a[1:-1, 1:-1] # neighbour order P2..P9 (clockwise from north) as (dy, dx) NB = [(-1, 0), (-1, 1), (0, 1), (1, 1), (1, 0), (1, -1), (0, -1), (-1, -1)] def crossing(sk, y, x, H, W): v = [1 if 0 <= y + dy < H and 0 <= x + dx < W and sk[y + dy, x + dx] else 0 for dy, dx in NB] return sum(1 for i in range(8) if v[i] == 0 and v[(i + 1) % 8] == 1), sum(v) # ------------------------------------------------------------------------------------------------ graph def trace(wire, dets, cfg=CFG): H, W = wire.shape m = (wire > cfg["wire_thr"]).astype(np.uint8) for d in dets: if d["cls"] == ARROW: continue x0, y0, x1, y1 = (v / STRIDE - 0.5 for v in d["box"]) p = cfg["cut_pad"] ya, yb = max(0, math.floor(y0 - p)), min(H, math.ceil(y1 + p) + 1) xa, xb = max(0, math.floor(x0 - p)), min(W, math.ceil(x1 + p) + 1) if yb > ya and xb > xa: m[ya:yb, xa:xb] = 0 sk = zhang_suen(m) ys, xs = np.nonzero(sk) kind = {} for y, x in zip(ys.tolist(), xs.tolist()): t, n = crossing(sk, y, x, H, W) if n == 0: continue kind[(y, x)] = "end" if t == 1 and n <= 2 else "junc" if t >= 3 else "body" # junction clusters (8-connected) node_of = {} nodes = [] # {"px": [...], "kind": "end"|"junc"} for p, k in sorted(kind.items()): if k == "body" or p in node_of: continue nid = len(nodes) stack, px = [p], [] node_of[p] = nid while stack: q = stack.pop() px.append(q) if k == "end": break for dy, dx in NB: r = (q[0] + dy, q[1] + dx) if kind.get(r) == "junc" and r not in node_of: node_of[r] = nid stack.append(r) nodes.append({"px": sorted(px), "kind": k}) # branches branches = [] # {"a": node, "b": node|None, "px": [...]} seen = set() for nid, nd in enumerate(nodes): for p in nd["px"]: for dy, dx in NB: q = (p[0] + dy, p[1] + dx) if q not in kind: continue if q in node_of: other = node_of[q] if other > nid and not any(b["a"] == nid and b["b"] == other and len(b["px"]) == 2 for b in branches): branches.append({"a": nid, "b": other, "px": [p, q]}) continue if q in seen: continue path, prev, cur, end = [p, q], p, q, None seen.add(q) while True: nxt, hit = None, None for oy, ox in (NB[0], NB[2], NB[4], NB[6], NB[1], NB[3], NB[5], NB[7]): # 4-neighbours first r = (cur[0] + oy, cur[1] + ox) if r == prev or r not in kind: continue if r in node_of: if node_of[r] != nid or len(path) > 2: hit = r break continue if r not in seen and nxt is None: nxt = r if hit is not None: path.append(hit) end = node_of[hit] break if nxt is None: break seen.add(nxt) path.append(nxt) prev, cur = cur, nxt branches.append({"a": nid, "b": end, "px": path}) for b in branches: b["len"] = sum(math.hypot(b["px"][i + 1][0] - b["px"][i][0], b["px"][i + 1][1] - b["px"][i][1]) for i in range(len(b["px"]) - 1)) return sk, nodes, branches def box_dist(py, px, box): """Distance in grid px from a grid point to a detection box (0 inside).""" x0, y0, x1, y1 = (v / STRIDE - 0.5 for v in box) dx = max(x0 - px, 0.0, px - x1) dy = max(y0 - py, 0.0, py - y1) return math.hypot(dx, dy) class DSU: def __init__(self, n): self.p = list(range(n)) def find(self, x): while self.p[x] != x: self.p[x] = self.p[self.p[x]] x = self.p[x] return x def union(self, a, b): a, b = self.find(a), self.find(b) if a != b: self.p[max(a, b)] = min(a, b) def decode(peaks, size, wire, cfg=CFG, dets=None): """Returns {"devices", "runs", "circuits", "unwired"}; dets may be passed in (oracle evaluation).""" if dets is None: dets = detect(peaks, size, cfg) dets = [dict(d) for d in dets] sk, nodes, branches = trace(wire, dets, cfg) devs = [d for d in dets if d["cls"] != ARROW] arrows = [d for d in dets if d["cls"] == ARROW] # what each endpoint node lands on term = {} for nid, nd in enumerate(nodes): if nd["kind"] != "end": continue y, x = nd["px"][0] best, bd = None, cfg["attach"] for d in devs: dd = box_dist(y, x, d["box"]) if dd <= bd: best, bd = ("dev", d["id"]), dd # arrowheads: the nearer target wins, except an end practically touching an arrowhead is always the home # run's tip (tips often stop beside another device) for a in arrows: dd = box_dist(y, x, a["box"]) if dd <= bd or dd <= cfg["arrow_touch"]: best, bd = ("arrow", a["id"]), min(dd, bd) term[nid] = best # merge junction clusters joined by a short bridge nd_dsu = DSU(len(nodes)) alive = [True] * len(branches) for i, b in enumerate(branches): if b["b"] is not None and b["a"] != b["b"] and nodes[b["a"]]["kind"] == "junc" and \ nodes[b["b"]]["kind"] == "junc" and b["len"] <= cfg["bridge"]: nd_dsu.union(b["a"], b["b"]) alive[i] = False # prune short spurs whose free end lands on nothing for i, b in enumerate(branches): if not alive[i] or b["len"] >= cfg["spur"]: continue ends = [b["a"], b["b"]] free = [e for e in ends if e is not None and nodes[e]["kind"] == "end" and term.get(e) is None] other = [e for e in ends if e is not None and nodes[e]["kind"] == "junc"] if free and other: alive[i] = False elif b["b"] is None or (free and len(free) == 2): alive[i] = False # isolated scrap # a drawn run always ends on a symbol: stroke ends that reached no detected symbol become unrecognised # devices (one per cluster of ends — runs reach an undetected fixture from several sides) if cfg["virtual"]: H, W = wire.shape bd_ = cfg["virtual_border"] free_ends = sorted({e for i, b in enumerate(branches) if alive[i] for e in (b["a"], b["b"]) if e is not None and nodes[e]["kind"] == "end" and term.get(e) is None}) free_ends = [e for e in free_ends if bd_ <= nodes[e]["px"][0][0] < H - bd_ and bd_ <= nodes[e]["px"][0][1] < W - bd_] # an end just short of an arrowhead is that home run's tip (thinning stops before the filled triangle) tips = [] for e in free_ends: y, x = nodes[e]["px"][0] best, bd = None, cfg["arrow_tip"] for a in arrows: dd = box_dist(y, x, a["box"]) if dd <= bd: best, bd = a["id"], dd if best is not None: term[e] = ("arrow", best) tips.append(e) free_ends = [e for e in free_ends if e not in tips] vd = DSU(len(free_ends)) for i in range(len(free_ends)): for j in range(i + 1, len(free_ends)): (y1, x1), (y2, x2) = nodes[free_ends[i]]["px"][0], nodes[free_ends[j]]["px"][0] if math.hypot(y1 - y2, x1 - x2) <= cfg["virtual_cluster"]: vd.union(i, j) vgroups = {} for i in range(len(free_ends)): vgroups.setdefault(vd.find(i), []).append(i) for root in sorted(vgroups): mem = vgroups[root] cy = (sum(nodes[free_ends[i]]["px"][0][0] for i in mem) / len(mem) + 0.5) * STRIDE cx = (sum(nodes[free_ends[i]]["px"][0][1] for i in mem) / len(mem) + 0.5) * STRIDE d = {"id": len(dets), "cls": VIRTUAL, "score": 0.0, "cx": cx, "cy": cy, "box": [cx - 6, cy - 6, cx + 6, cy + 6]} dets.append(d) devs.append(d) for i in mem: term[free_ends[i]] = ("dev", d["id"]) # a junction sitting on a symbol is where several runs land on it, not a place where they join each other jterm = {} for n, nd in enumerate(nodes): if nd["kind"] != "junc": continue j = nd_dsu.find(n) for (y, x) in nd["px"]: for d in devs: dd = box_dist(y, x, d["box"]) if dd <= cfg["attach"] and (j not in jterm or dd < jterm[j][1] or (dd == jterm[j][1] and d["id"] < jterm[j][0])): jterm[j] = (d["id"], dd) # link branches through junctions br_dsu = DSU(len(branches)) at = {} for i, b in enumerate(branches): if not alive[i]: continue for endpos, e in ((0, b["a"]), (1, b["b"])): if e is not None and nodes[e]["kind"] == "junc": at.setdefault(nd_dsu.find(e), []).append((i, endpos)) # at a junction on a symbol, a dead-straight pair of arms is another run crossing right beside it (two runs # arriving from opposite sides cannot touch: the symbol's cut-out lies between them); the rest land on it passthru = set() for j, lst in sorted(at.items()): if j not in jterm or len(lst) < 3: continue vec = junction_dirs(j, lst, nodes, nd_dsu, branches, cfg) pairs = sorted((vec[u][0] * vec[v][0] + vec[u][1] * vec[v][1], u, v) for u in range(len(lst)) for v in range(u + 1, len(lst))) if pairs[0][0] <= cfg["through_cos"]: u, v = pairs[0][1], pairs[0][2] br_dsu.union(lst[u][0], lst[v][0]) passthru.add((lst[u][0], j)) passthru.add((lst[v][0], j)) for j, lst in sorted(at.items()): if j in jterm: continue if len(lst) == 3 and cfg["t3"] == "pair": # an X whose fourth arm was lost (thinning, pruning): join only the straight-through pair vec = junction_dirs(j, lst, nodes, nd_dsu, branches, cfg) pairs = sorted(((vec[u][0] * vec[v][0] + vec[u][1] * vec[v][1], u, v) for u, v in ((0, 1), (0, 2), (1, 2)))) if pairs[0][0] <= cfg["pair_cos"]: br_dsu.union(lst[pairs[0][1]][0], lst[pairs[0][2]][0]) continue if len(lst) == 4: vec = junction_dirs(j, lst, nodes, nd_dsu, branches, cfg) best = None for pairing in (((0, 1), (2, 3)), ((0, 2), (1, 3)), ((0, 3), (1, 2))): c1 = vec[pairing[0][0]][0] * vec[pairing[0][1]][0] + vec[pairing[0][0]][1] * vec[pairing[0][1]][1] c2 = vec[pairing[1][0]][0] * vec[pairing[1][1]][0] + vec[pairing[1][0]][1] * vec[pairing[1][1]][1] score = max(c1, c2) if best is None or score < best[0]: best = (score, pairing) if best[0] <= cfg["pair_cos"]: for u, v in best[1]: br_dsu.union(lst[u][0], lst[v][0]) continue for i, _ in lst[1:]: br_dsu.union(lst[0][0], i) # collect stroke groups groups = {} for i, b in enumerate(branches): if not alive[i]: continue g = groups.setdefault(br_dsu.find(i), {"devs": [], "arrows": [], "free": 0, "len": 0.0, "branches": []}) g["len"] += b["len"] g["branches"].append(i) for e in (b["a"], b["b"]): if e is not None and nodes[e]["kind"] == "junc" and nd_dsu.find(e) in jterm: if (i, nd_dsu.find(e)) in passthru: continue dv = jterm[nd_dsu.find(e)][0] if dv not in g["devs"]: g["devs"].append(dv) continue if e is None or nodes[e]["kind"] != "end": continue t = term.get(e) if t is None: g["free"] += 1 elif t[0] == "dev" and t[1] not in g["devs"]: g["devs"].append(t[1]) elif t[0] == "arrow" and t[1] not in g["arrows"]: g["arrows"].append(t[1]) byid = {d["id"]: d for d in dets} runs = [] homerun_devs = {} for gk in sorted(groups): g = groups[gk] ds = sorted(g["devs"]) if not ds: continue pts = [[(p[1] + 0.5) * STRIDE, (p[0] + 0.5) * STRIDE] for i in g["branches"] for p in branches[i]["px"]] if g["arrows"]: # the home run leaves from the device nearest the arrowhead a = byid[g["arrows"][0]] # a detected symbol beats an inferred one: an inferred device next to an arrow is usually a leader bend hd = min(ds, key=lambda i: (byid[i]["cls"] == VIRTUAL, math.hypot(byid[i]["cx"] - a["cx"], byid[i]["cy"] - a["cy"]), i)) homerun_devs[hd] = a["id"] if len(ds) == 1: runs.append({"devices": [hd], "homerun": True, "arrow": a["id"], "length_px": g["len"] * STRIDE, "points": pts}) continue order = nn_chain(ds, byid) for u, v in zip(order, order[1:]): runs.append({"devices": sorted([u, v]), "homerun": False, "length_px": g["len"] * STRIDE / max(1, len(order) - 1), "points": pts}) # arrowheads no stroke claimed: short home runs mostly vanish once the symbol is cut out and thinning # retracts both ends. Such an arrow belongs to the nearest wireable symbol within arrow_reach. used = {a for a in homerun_devs.values()} for a in arrows: if a["id"] in used: continue ax0, ay0, ax1, ay1 = (v / STRIDE - 0.5 for v in a["box"]) best, bd = None, cfg["arrow_reach"] for d in devs: if cname(d["cls"]) in UNWIRED: continue dd = min(box_dist(ay0, ax0, d["box"]), box_dist(ay0, ax1, d["box"]), box_dist(ay1, ax0, d["box"]), box_dist(ay1, ax1, d["box"]), box_dist((ay0 + ay1) / 2, (ax0 + ax1) / 2, d["box"])) if dd < bd or (dd == bd and best is not None and d["id"] < best): best, bd = d["id"], dd if best is not None and best not in homerun_devs: homerun_devs[best] = a["id"] runs.append({"devices": [best], "homerun": True, "arrow": a["id"], "length_px": bd * STRIDE, "points": []}) # neighbours joined by a run too short to survive the cut: wire straight across the gap in the uncut mask joined = {tuple(r["devices"]) for r in runs if len(r["devices"]) == 2} wireable = [d for d in devs if cname(d["cls"]) not in UNWIRED] for ai in range(len(wireable)): for bi in range(ai + 1, len(wireable)): A, B = wireable[ai], wireable[bi] key = tuple(sorted((A["id"], B["id"]))) if key in joined: continue gap = short_gap(A["box"], B["box"]) if gap is None or gap > cfg["short_gap"] * STRIDE: continue (x0, y0), (x1, y1) = closest_points(A["box"], B["box"]) n = max(3, int(math.hypot(x1 - x0, y1 - y0) / STRIDE) + 1) hits = 0 for k in range(n): t = (k + 0.5) / n gx = math.floor((x0 + (x1 - x0) * t) / STRIDE) # == round-half-up(v/STRIDE - 0.5), as in JS gy = math.floor((y0 + (y1 - y0) * t) / STRIDE) if 0 <= gy < wire.shape[0] and 0 <= gx < wire.shape[1] and wire[gy, gx] > cfg["wire_thr"]: hits += 1 if hits >= 0.8 * n: runs.append({"devices": list(key), "homerun": False, "length_px": gap, "points": [[x0, y0], [x1, y1]]}) joined.add(key) dsu = DSU(len(dets)) for r in runs: if len(r["devices"]) == 2: dsu.union(*r["devices"]) wired = {i for r in runs for i in r["devices"]} circ = {} for i in sorted(wired): circ.setdefault(dsu.find(i), []).append(i) circuits = [] for root in sorted(circ): mem = circ[root] circuits.append({ "devices": mem, "homeruns": sorted(homerun_devs[i] for i in mem if i in homerun_devs), "length_px": sum(r["length_px"] for r in runs if r["devices"][0] in mem), }) unwired = [d["id"] for d in devs if d["id"] not in wired and cname(d["cls"]) not in UNWIRED] return {"devices": [{k: d[k] for k in ("id", "cls", "score", "box")} for d in dets], "runs": runs, "circuits": circuits, "unwired": unwired, "homeruns": [{"device": i, "arrow": homerun_devs[i]} for i in sorted(homerun_devs)]} def short_gap(a, b): """Edge-to-edge distance between two boxes (px), None if they overlap on neither axis' projection.""" dx = max(b[0] - a[2], a[0] - b[2], 0.0) dy = max(b[1] - a[3], a[1] - b[3], 0.0) if dx > 0 and dy > 0: return None # diagonal neighbours: a short straight run would not join them edge to edge return max(dx, dy) def closest_points(a, b): """Facing points on two boxes that sit side by side (the gap a short run would bridge).""" if max(b[0] - a[2], a[0] - b[2]) > 0: # side by side horizontally y = (max(a[1], b[1]) + min(a[3], b[3])) / 2 return ((a[2], y), (b[0], y)) if b[0] >= a[2] else ((a[0], y), (b[2], y)) x = (max(a[0], b[0]) + min(a[2], b[2])) / 2 return ((x, a[3]), (x, b[1])) if b[1] >= a[3] else ((x, a[1]), (x, b[3])) def junction_dirs(j, lst, nodes, nd_dsu, branches, cfg): """Unit direction of each branch leaving junction cluster j, measured dir_len px out from its centroid.""" members = [n for n in range(len(nodes)) if nodes[n]["kind"] == "junc" and nd_dsu.find(n) == j] pts = [p for n in members for p in nodes[n]["px"]] cy = sum(p[0] for p in pts) / len(pts) cx = sum(p[1] for p in pts) / len(pts) vec = [] for i, endpos in lst: px = branches[i]["px"] if endpos == 0 else branches[i]["px"][::-1] q = px[min(len(px) - 1, cfg["dir_len"])] vy, vx = q[0] - cy, q[1] - cx n = math.hypot(vy, vx) or 1.0 vec.append((vy / n, vx / n)) return vec def nn_chain(ids, byid): ids = sorted(ids) start = min(ids, key=lambda i: (byid[i]["cx"] + byid[i]["cy"], i)) path, left = [start], set(ids) - {start} while left: c = byid[path[-1]] nxt = min(left, key=lambda i: (math.hypot(byid[i]["cx"] - c["cx"], byid[i]["cy"] - c["cy"]), i)) path.append(nxt) left.remove(nxt) return path