pre-refactor 041426
This commit is contained in:
@@ -0,0 +1,693 @@
|
||||
"""
|
||||
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()
|
||||
Reference in New Issue
Block a user