""" Publication-quality architecture diagrams for HyperTower. Generates: architecture_single_tower.png — single-eye image-only tower architecture_hypertower.png — single-eye image + clinical fusion architecture_ensemble.png — bilateral ensemble (two HyperTowers + average) architecture_fused_head.png — bilateral ensemble + learned head Usage: python -m v3.scripts.output_analysis.plot_architecture python -m v3.scripts.output_analysis.plot_architecture --out figures/ """ from __future__ import annotations import argparse from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt from matplotlib.patches import FancyBboxPatch, FancyArrowPatch import matplotlib.patheffects as pe # ── Colour palette ──────────────────────────────────────────────────────────── C_IMG = "#4e8d3a" # green — image / CNN C_MD = "#4c72b0" # blue — clinical / MLP C_BRIDGE = "#c44e52" # red — bridge / fusion C_EMB = "#2a9d8f" # teal — embedding vectors (z) C_OUT = "#8c6bb1" # purple — output nodes C_HEAD = "#d4a017" # gold — learned head / average C_INPUT = "#a0a0a0" # grey — raw input nodes C_BG = "#e8e8e8" C_ARROW = "#444444" FONT = "DejaVu Sans" # ── Low-level primitives ────────────────────────────────────────────────────── def _box( ax, cx, cy, w, h, color, text="", fontsize=9, text_color="white", bold=False, alpha=0.92, radius=0.12, lw=1.5, ): """Rounded rectangle centered at (cx, cy).""" patch = FancyBboxPatch( (cx - w / 2, cy - h / 2), w, h, boxstyle=f"round,pad=0,rounding_size={radius}", facecolor=color, edgecolor="white", linewidth=lw, alpha=alpha, zorder=3, transform=ax.transData, ) ax.add_patch(patch) if text: ax.text( cx, cy, text, ha="center", va="center", fontsize=fontsize, color=text_color, fontweight="bold" if bold else "normal", fontfamily=FONT, zorder=4, ) return patch def _arrow(ax, x0, y0, x1, y1, lw=1.6, color=C_ARROW, style="->", rad=0.0): ax.annotate( "", xy=(x1, y1), xytext=(x0, y0), arrowprops=dict( arrowstyle=style, color=color, lw=lw, connectionstyle=f"arc3,rad={rad}", ), zorder=2, ) def _text( ax, x, y, s, fontsize=8, color="#333333", ha="center", va="center", bold=False ): ax.text( x, y, s, ha=ha, va=va, fontsize=fontsize, color=color, fontfamily=FONT, fontweight="bold" if bold else "normal", zorder=5, ) def _bracket( ax, x, y0, y1, text="", fontsize=8.5, color="#888888", pad=0.15, lw=1.4, badge_color=None, ): """Vertical C-bracket on the right side. If badge_color is set, the label is drawn as white text on a filled badge.""" mid = (y0 + y1) / 2 ax.plot( [x, x + pad, x + pad, x], [y1, y1, y0, y0], color=color, lw=lw, solid_capstyle="round", zorder=2, ) if text: if badge_color: ax.text( x + pad * 1.4, mid, text, ha="left", va="center", fontsize=fontsize, color="white", fontfamily=FONT, fontweight="bold", zorder=6, bbox=dict( facecolor=badge_color, edgecolor="none", pad=3.5, boxstyle="round,pad=0.3", ), ) else: ax.text( x + pad * 1.4, mid, text, ha="left", va="center", fontsize=fontsize, color=color, fontfamily=FONT, style="italic", ) def _setup(fig, ax, w, h, title): ax.set_xlim(0, w) ax.set_ylim(0, h) ax.axis("off") ax.set_facecolor(C_BG) fig.patch.set_facecolor(C_BG) if title: ax.set_title( title, fontsize=12, fontweight="bold", fontfamily=FONT, pad=10, color="#222222", ) # ── Reusable sub-blocks ─────────────────────────────────────────────────────── def _draw_cnn_block(ax, x_center, y, w=2.0, h=0.75): """Three-layer CNN block with labels: Conv Layers → Conv Layers → GAP.""" labels = ["Conv\nLayers", "Conv\nLayers", "GAP"] sub_w = [w * 0.42, w * 0.30, w * 0.22] sub_h = [h, h * 0.82, h * 0.65] alphas = [0.82, 0.74, 0.66] fsizes = [8.0, 7.5, 7.5] gap = (w - sum(sub_w)) / 2 xs = [ x_center - w / 2 + sub_w[0] / 2, x_center - w / 2 + sub_w[0] + gap + sub_w[1] / 2, x_center - w / 2 + sub_w[0] + gap + sub_w[1] + gap + sub_w[2] / 2, ] for i, (sx, sw, sh, lbl, alp, fs) in enumerate( zip(xs, sub_w, sub_h, labels, alphas, fsizes) ): _box(ax, sx, y, sw, sh, C_IMG, lbl, fontsize=fs, alpha=alp, radius=0.08) if i < 2: _arrow( ax, sx + sw / 2, y, xs[i + 1] - sub_w[i + 1] / 2, y, lw=1.2, style="-|>" ) return xs[-1] + sub_w[-1] / 2 def _draw_mlp_block(ax, x_center, y, w=1.4, h=0.65): """Two-layer MLP block: FC(128) → FC(128) (hidden_dim=128 both layers).""" labels = ["FC\n(128)", "FC\n(128)"] w0, w1 = w * 0.55, w * 0.45 gap = w - w0 - w1 x0 = x_center - w / 2 + w0 / 2 x1 = x0 + w0 / 2 + gap + w1 / 2 _box(ax, x0, y, w0, h, C_MD, labels[0], fontsize=8.0, alpha=0.82, radius=0.08) _arrow(ax, x0 + w0 / 2, y, x1 - w1 / 2, y, lw=1.2, style="-|>") _box( ax, x1, y, w1, h * 0.88, C_MD, labels[1], fontsize=7.5, alpha=0.72, radius=0.08 ) return x1 + w1 / 2 def _draw_embedding(ax, x, y, w=0.40, h=0.75, label="z\n(emb)"): _box(ax, x + w / 2, y, w, h, C_EMB, label, fontsize=8, bold=True, radius=0.08) return x + w def _draw_output(ax, x, y, dy=0.45, classes=("Glaucoma", "Normal")): """Stacked output class boxes, connected from (x, y) via arrows.""" n = len(classes) bw = 1.10 bh = 0.38 gap = 0.08 total = n * bh + (n - 1) * gap y_top = y + total / 2 - bh / 2 for i, cls in enumerate(classes): cy = y_top - i * (bh + gap) _box(ax, x + bw / 2, cy, bw, bh, C_OUT, cls, fontsize=8.5, radius=0.08) _arrow(ax, x, y, x, cy, lw=1.1, style="-|>", rad=0.0) _text(ax, x + bw / 2, y - total / 2 - 0.20, "Softmax", fontsize=7.5, color=C_OUT) def _draw_compact_ht(ax, x_left, y_img, y_md, eye_label): """Compact HyperTower block: Image+Clinical boxes → Bridge. Returns (x_right_of_bridge, y_bridge_center). """ bw_img = 1.40 bh_img = 0.72 bw_md = 1.20 bh_md = 0.62 bw_br = 0.72 cy_br = (y_img + y_md) / 2 bh_br = abs(y_img - y_md) * 0.60 # Image box: CNN Backbone _box( ax, x_left + bw_img / 2, y_img, bw_img, bh_img, C_IMG, f"{eye_label}\nCNN Backbone", fontsize=8.5, radius=0.08, ) # MD box: Clinical MLP _box( ax, x_left + bw_md / 2, y_md, bw_md, bh_md, C_MD, f"{eye_label}\nClinical MLP", fontsize=8.5, radius=0.08, ) # Arrows to bridge br_x = x_left + max(bw_img, bw_md) + 0.60 _arrow(ax, x_left + bw_img, y_img, br_x - bw_br / 2, cy_br, lw=1.2, style="-|>") _arrow(ax, x_left + bw_md, y_md, br_x - bw_br / 2, cy_br, lw=1.2, style="-|>") # Bridge label kept simple — detail lives in the hypertower diagram _box( ax, br_x, cy_br, bw_br, max(bh_br, 0.70), C_BRIDGE, "Bridge", fontsize=8.0, radius=0.08, ) return br_x + bw_br / 2, cy_br # ── Figure 1: Single Tower ──────────────────────────────────────────────────── def make_single_tower(out_dir: Path): W, H = 9.0, 3.2 fig, ax = plt.subplots(figsize=(W, H)) _setup(fig, ax, W, H, "Single Tower") cy = H / 2 # Input _box( ax, 0.75, cy, 0.95, 0.60, C_INPUT, "Fundus\nImage", fontsize=8.5, radius=0.08, alpha=0.75, text_color="#333", ) _arrow(ax, 1.22, cy, 1.60, cy) # CNN Backbone cnn_x_right = _draw_cnn_block(ax, x_center=3.10, y=cy, w=2.80, h=0.78) _text(ax, 3.10, cy - 0.68, "CNN Backbone", fontsize=8.5, color=C_IMG, bold=True) _arrow(ax, 1.60, cy, 1.73, cy, lw=1.4, style="-|>") # Embedding emb_x_right = _draw_embedding(ax, x=cnn_x_right + 0.28, y=cy, w=0.48, h=0.78) _arrow(ax, cnn_x_right, cy, cnn_x_right + 0.28, cy, lw=1.4, style="-|>") # Classifier _arrow(ax, emb_x_right, cy, emb_x_right + 0.25, cy, lw=1.4, style="-|>") _draw_output(ax, emb_x_right + 0.25, cy) path = out_dir / "architecture_single_tower.png" fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) print(f" Saved: {path}") # ── Figure 2: HyperTower (single eye) ──────────────────────────────────────── def make_hypertower(out_dir: Path): W, H = 11.0, 5.5 fig, ax = plt.subplots(figsize=(W, H)) _setup(fig, ax, W, H, "HyperTower — Single Eye") y_img = 3.70 y_md = 1.60 # ── Image tower ──────────────────────────────────────────────── _box( ax, 0.80, y_img, 1.00, 0.60, C_INPUT, "Fundus\nImage", fontsize=8.5, radius=0.08, alpha=0.75, text_color="#333", ) _arrow(ax, 1.30, y_img, 1.85, y_img) cnn_x_r = _draw_cnn_block(ax, x_center=3.50, y=y_img, w=2.80, h=0.75) _text(ax, 3.50, y_img - 0.65, "CNN Backbone", fontsize=8, color=C_IMG, bold=True) _arrow(ax, 1.85, y_img, 1.98, y_img, lw=1.4, style="-|>") emb_img_x = _draw_embedding(ax, x=cnn_x_r + 0.30, y=y_img, w=0.65, h=0.75) _arrow(ax, cnn_x_r, y_img, cnn_x_r + 0.30, y_img, lw=1.4, style="-|>") _text( ax, (1.30 + emb_img_x) / 2, y_img + 0.65, "Image Tower", fontsize=9, color=C_IMG, bold=True, ) # ── Clinical tower ───────────────────────────────────────────── _box( ax, 0.80, y_md, 1.00, 0.55, C_INPUT, "Clinical\nData", fontsize=8.5, radius=0.08, alpha=0.75, text_color="#333", ) _arrow(ax, 1.30, y_md, 1.65, y_md) mlp_x_r = _draw_mlp_block(ax, x_center=2.90, y=y_md, w=1.60, h=0.65) _arrow(ax, 1.65, y_md, 1.74, y_md, lw=1.4, style="-|>") emb_md_x = _draw_embedding(ax, x=mlp_x_r + 0.30, y=y_md, w=0.65, h=0.65) _arrow(ax, mlp_x_r, y_md, mlp_x_r + 0.30, y_md, lw=1.4, style="-|>") _text( ax, (1.30 + emb_md_x) / 2, y_md - 0.60, "Clinical Tower", fontsize=9, color=C_MD, bold=True, ) # ── Bridge ───────────────────────────────────────────────────── br_x = max(emb_img_x, emb_md_x) + 0.80 cy_br = (y_img + y_md) / 2 bh_br = abs(y_img - y_md) * 0.55 _arrow(ax, emb_img_x, y_img, br_x - 0.40, cy_br, lw=1.4, style="-|>") _arrow(ax, emb_md_x, y_md, br_x - 0.40, cy_br, lw=1.4, style="-|>") _box( ax, br_x, cy_br, 1.40, max(bh_br, 1.35), C_BRIDGE, "Bridge\nFC(img→256)\nFC(md→256)\n⊙ Hadamard\n→ ReLU→FC(2)", fontsize=8, radius=0.10, ) # ── Output ───────────────────────────────────────────────────── out_x = br_x + 0.65 + 0.40 _arrow(ax, br_x + 0.65, cy_br, out_x, cy_br, lw=1.4, style="-|>") _draw_output(ax, out_x, cy_br) # ── Bracket (right of output nodes; output bw=1.10 so right edge = out_x+1.10) _bracket( ax, x=out_x + 1.25, y0=y_md - 0.50, y1=y_img + 0.50, text="HyperTower", fontsize=9, pad=0.22, color="#555", badge_color="#555", ) path = out_dir / "architecture_hypertower.png" fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) print(f" Saved: {path}") # ── Figure 3: Bilateral Ensemble ───────────────────────────────────────────── def make_ensemble(out_dir: Path): W, H = 10.5, 7.5 fig, ax = plt.subplots(figsize=(W, H)) _setup(fig, ax, W, H, "Bilateral Ensemble HyperTower") x_left = 3.0 inp_cx = 1.85 inp_w = 0.90 inp_h = 0.55 # OD (top) od_y_img, od_y_md = 5.90, 4.60 br_od_x, cy_od = _draw_compact_ht( ax, x_left=x_left, y_img=od_y_img, y_md=od_y_md, eye_label="OD" ) _text(ax, 0.45, (od_y_img + od_y_md) / 2, "OD\n(Right Eye)", fontsize=9, color="#444", bold=True) _box(ax, inp_cx, od_y_img, inp_w, inp_h, C_INPUT, "Fundus\nImage", fontsize=8, radius=0.08, alpha=0.75, text_color="#333") _box(ax, inp_cx, od_y_md, inp_w, inp_h, C_INPUT, "Clinical\nData", fontsize=8, radius=0.08, alpha=0.75, text_color="#333") _arrow(ax, inp_cx + inp_w / 2, od_y_img, x_left, od_y_img, lw=1.2, style="-|>") _arrow(ax, inp_cx + inp_w / 2, od_y_md, x_left, od_y_md, lw=1.2, style="-|>") # OS (bottom) os_y_img, os_y_md = 2.80, 1.50 br_os_x, cy_os = _draw_compact_ht( ax, x_left=x_left, y_img=os_y_img, y_md=os_y_md, eye_label="OS" ) _text(ax, 0.45, (os_y_img + os_y_md) / 2, "OS\n(Left Eye)", fontsize=9, color="#444", bold=True) _box(ax, inp_cx, os_y_img, inp_w, inp_h, C_INPUT, "Fundus\nImage", fontsize=8, radius=0.08, alpha=0.75, text_color="#333") _box(ax, inp_cx, os_y_md, inp_w, inp_h, C_INPUT, "Clinical\nData", fontsize=8, radius=0.08, alpha=0.75, text_color="#333") _arrow(ax, inp_cx + inp_w / 2, os_y_img, x_left, os_y_img, lw=1.2, style="-|>") _arrow(ax, inp_cx + inp_w / 2, os_y_md, x_left, os_y_md, lw=1.2, style="-|>") # Average node avg_x = max(br_od_x, br_os_x) + 1.20 avg_y = (cy_od + cy_os) / 2 avg_size = 0.90 _arrow(ax, br_od_x, cy_od, avg_x - avg_size / 2, avg_y, lw=1.4, style="-|>") _arrow(ax, br_os_x, cy_os, avg_x - avg_size / 2, avg_y, lw=1.4, style="-|>") _box( ax, avg_x, avg_y, avg_size, avg_size, C_HEAD, "Average", fontsize=10, bold=True, radius=0.10, ) # Output out_x = avg_x + avg_size / 2 + 0.50 _arrow(ax, avg_x + avg_size / 2, avg_y, out_x, avg_y, lw=1.5, style="-|>") _draw_output(ax, out_x, avg_y) # Side brackets — white text on badge _bracket( ax, x=br_od_x + 0.10, y0=od_y_md - 0.45, y1=od_y_img + 0.45, text="OD HyperTower", fontsize=8.5, pad=0.20, color="#555", badge_color="#555", ) _bracket( ax, x=br_os_x + 0.10, y0=os_y_md - 0.45, y1=os_y_img + 0.45, text="OS HyperTower", fontsize=8.5, pad=0.20, color="#555", badge_color="#555", ) path = out_dir / "architecture_ensemble.png" fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) print(f" Saved: {path}") # ── Figure 4: Fused Head ────────────────────────────────────────────────────── def make_fused_head(out_dir: Path): W, H = 10.5, 7.5 fig, ax = plt.subplots(figsize=(W, H)) _setup(fig, ax, W, H, "Fused Head Bilateral HyperTower") x_left = 3.0 inp_cx = 1.85 inp_w = 0.90 inp_h = 0.55 # OD (top) — same layout as ensemble od_y_img, od_y_md = 5.90, 4.60 br_od_x, cy_od = _draw_compact_ht( ax, x_left=x_left, y_img=od_y_img, y_md=od_y_md, eye_label="OD" ) _text(ax, 0.45, (od_y_img + od_y_md) / 2, "OD\n(Right Eye)", fontsize=9, color="#444", bold=True) _box(ax, inp_cx, od_y_img, inp_w, inp_h, C_INPUT, "Fundus\nImage", fontsize=8, radius=0.08, alpha=0.75, text_color="#333") _box(ax, inp_cx, od_y_md, inp_w, inp_h, C_INPUT, "Clinical\nData", fontsize=8, radius=0.08, alpha=0.75, text_color="#333") _arrow(ax, inp_cx + inp_w / 2, od_y_img, x_left, od_y_img, lw=1.2, style="-|>") _arrow(ax, inp_cx + inp_w / 2, od_y_md, x_left, od_y_md, lw=1.2, style="-|>") # OS (bottom) os_y_img, os_y_md = 2.80, 1.50 br_os_x, cy_os = _draw_compact_ht( ax, x_left=x_left, y_img=os_y_img, y_md=os_y_md, eye_label="OS" ) _text(ax, 0.45, (os_y_img + os_y_md) / 2, "OS\n(Left Eye)", fontsize=9, color="#444", bold=True) _box(ax, inp_cx, os_y_img, inp_w, inp_h, C_INPUT, "Fundus\nImage", fontsize=8, radius=0.08, alpha=0.75, text_color="#333") _box(ax, inp_cx, os_y_md, inp_w, inp_h, C_INPUT, "Clinical\nData", fontsize=8, radius=0.08, alpha=0.75, text_color="#333") _arrow(ax, inp_cx + inp_w / 2, os_y_img, x_left, os_y_img, lw=1.2, style="-|>") _arrow(ax, inp_cx + inp_w / 2, os_y_md, x_left, os_y_md, lw=1.2, style="-|>") _bracket( ax, x=br_od_x + -0.17, y0=od_y_md - 0.45, y1=od_y_img + 0.45, text="OD HyperTower", fontsize=8.5, pad=0.20, color="#555", badge_color="#555", ) _bracket( ax, x=br_os_x + -0.17, y0=os_y_md - 0.45, y1=os_y_img + 0.45, text="OS HyperTower", fontsize=8.5, pad=0.20, color="#555", badge_color="#555", ) # Fused Head box with logit MLP detail avg_y = (cy_od + cy_os) / 2 head_x = max(br_od_x, br_os_x) + 2.20 head_w = 1.80 head_h = 1.20 _arrow(ax, br_od_x + 0.05, cy_od, head_x - head_w / 2, avg_y, lw=1.4, style="-|>") _arrow(ax, br_os_x + 0.05, cy_os, head_x - head_w / 2, avg_y, lw=1.4, style="-|>") _box( ax, head_x, avg_y, head_w, head_h, C_HEAD, "Fused Head\ncat(l_OD, l_OS)\n→ FC(64) → logits", fontsize=8.5, bold=False, radius=0.10, ) # Output nodes + softmax out_x = head_x + head_w / 2 + 0.50 _arrow(ax, head_x + head_w / 2, avg_y, out_x, avg_y, lw=1.5, style="-|>") _draw_output(ax, out_x, avg_y) path = out_dir / "architecture_fused_head.png" fig.savefig(path, dpi=180, bbox_inches="tight") plt.close(fig) print(f" Saved: {path}") # ── Main ────────────────────────────────────────────────────────────────────── def main(): ap = argparse.ArgumentParser( description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter ) ap.add_argument( "--out", type=Path, default=Path(__file__).resolve().parents[3] / "v3" / "figures", help="Output directory", ) args = ap.parse_args() args.out.mkdir(parents=True, exist_ok=True) print("Generating architecture diagrams...") make_single_tower(args.out) make_hypertower(args.out) make_ensemble(args.out) make_fused_head(args.out) print("Done.") if __name__ == "__main__": main()