File size: 11,419 Bytes
6fe482b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
"""
Publication-grade visualisations for the retinal QC world.

All figures share a restrained clinical palette (deep slate ink on warm paper,
teal / amber / coral accents) so the tool reads like a journal figure rather
than a dashboard. Every function returns an RGB numpy array for Gradio.
"""

from __future__ import annotations
import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from matplotlib.patches import Circle
import cv2
import io

# ----------------------------------------------------------------- palette / style
INK = "#12232e"
PAPER = "#faf8f4"
TEAL = "#1f7a8c"
DEEP = "#0b3c49"
AMBER = "#d69e2e"
CORAL = "#e05252"
GREEN = "#1f9d61"
MUTED = "#6b7c85"
GRID = "#dfe4e2"

VERDICT_COLORS = {"PASS": GREEN, "ACCEPTABLE": AMBER, "FAIL": CORAL}

plt.rcParams.update({
    "figure.facecolor": PAPER, "axes.facecolor": PAPER,
    "savefig.facecolor": PAPER, "text.color": INK,
    "axes.edgecolor": "#c8d0cd", "axes.labelcolor": INK,
    "xtick.color": MUTED, "ytick.color": MUTED,
    "font.family": "DejaVu Sans", "font.size": 10,
    "axes.grid": True, "grid.color": GRID, "grid.linewidth": 0.7,
})


def _fig_to_rgb(fig, dpi=130):
    buf = io.BytesIO()
    fig.savefig(buf, format="png", dpi=dpi, bbox_inches="tight")
    plt.close(fig)
    buf.seek(0)
    arr = np.array(matplotlib.image.imread(buf) * 255, np.uint8)
    return arr[..., :3]


# --------------------------------------------------------------------- score card
def score_gauge(summary):
    """A circular composite-score gauge with the verdict."""
    fig, ax = plt.subplots(figsize=(3.2, 3.2))
    ax.set_xlim(-1.3, 1.3); ax.set_ylim(-1.3, 1.3); ax.axis("off"); ax.set_aspect("equal")
    score = summary["composite"]
    color = VERDICT_COLORS[summary["verdict"]]
    # track
    theta = np.linspace(np.pi * 1.25, -np.pi * 0.25, 200)
    ax.plot(np.cos(theta), np.sin(theta), color=GRID, lw=14, solid_capstyle="round")
    frac = score / 100
    theta2 = np.linspace(np.pi * 1.25, np.pi * 1.25 - 1.5 * np.pi * frac, 200)
    ax.plot(np.cos(theta2), np.sin(theta2), color=color, lw=14, solid_capstyle="round")
    ax.text(0, 0.12, f"{score:.0f}", ha="center", va="center", fontsize=40,
            fontweight="bold", color=INK)
    ax.text(0, -0.32, "/ 100", ha="center", va="center", fontsize=12, color=MUTED)
    ax.text(0, -0.72, summary["verdict"], ha="center", va="center", fontsize=17,
            fontweight="bold", color=color)
    ax.text(0, -1.02, f'({summary["band"]})', ha="center", va="center",
            fontsize=10, color=MUTED)
    return _fig_to_rgb(fig)


def metric_radar(metrics):
    """Radar chart of the per-axis 0-1 scores."""
    names = [m["name"].replace(" / ", "/\n").replace(" (", "\n(") for m in metrics]
    vals = [m["score"] for m in metrics]
    N = len(vals)
    ang = np.linspace(0, 2 * np.pi, N, endpoint=False).tolist()
    vals2 = vals + vals[:1]; ang2 = ang + ang[:1]
    fig, ax = plt.subplots(figsize=(5.2, 5.2), subplot_kw=dict(polar=True))
    ax.set_facecolor(PAPER)
    ax.plot(ang2, vals2, color=TEAL, lw=2)
    ax.fill(ang2, vals2, color=TEAL, alpha=0.22)
    # pass / acceptable rings
    ax.plot(np.linspace(0, 2*np.pi, 100), [0.66]*100, color=GREEN, ls="--", lw=1, alpha=0.6)
    ax.plot(np.linspace(0, 2*np.pi, 100), [0.40]*100, color=AMBER, ls="--", lw=1, alpha=0.6)
    ax.set_xticks(ang); ax.set_xticklabels(names, fontsize=8)
    ax.set_yticks([0.4, 0.66, 1.0]); ax.set_yticklabels(["0.4", "0.66", "1.0"], fontsize=7)
    ax.set_ylim(0, 1)
    ax.set_title("Per-axis quality profile", color=INK, fontsize=12, pad=18)
    return _fig_to_rgb(fig)


def metric_bars(metrics):
    """Horizontal bar chart of per-axis scores, coloured by status."""
    metrics = sorted(metrics, key=lambda m: m["score"])
    names = [m["name"] for m in metrics]
    vals = [m["score"] for m in metrics]
    cols = [GREEN if m["score"] >= 0.66 else AMBER if m["score"] >= 0.40 else CORAL
            for m in metrics]
    fig, ax = plt.subplots(figsize=(6.4, 4.3))
    y = np.arange(len(names))
    ax.barh(y, vals, color=cols, height=0.66, edgecolor="white")
    for yi, m in zip(y, metrics):
        ax.text(min(m["score"] + 0.02, 0.92), yi,
                f'{m["value"]:.1f} {m["unit"]}', va="center", fontsize=7.5, color=MUTED)
    ax.set_yticks(y); ax.set_yticklabels(names, fontsize=9)
    ax.axvline(0.66, color=GREEN, ls="--", lw=1, alpha=0.6)
    ax.axvline(0.40, color=AMBER, ls="--", lw=1, alpha=0.6)
    ax.set_xlim(0, 1); ax.set_xlabel("axis score (0-1)")
    ax.set_title("Quality axes  -  worst first", fontsize=12, loc="left")
    ax.grid(axis="y", alpha=0)
    return _fig_to_rgb(fig)


# --------------------------------------------------------------------- cohort views
def cohort_distribution(df):
    """Histogram of composite scores + PASS/ACCEPTABLE/FAIL bar."""
    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9.2, 3.6),
                                 gridspec_kw=dict(width_ratios=[2.2, 1]))
    a1.hist(df["composite"], bins=20, range=(0, 100), color=TEAL, alpha=0.85,
            edgecolor="white")
    a1.axvline(70, color=GREEN, ls="--", lw=1.2); a1.axvline(45, color=AMBER, ls="--", lw=1.2)
    a1.set_xlabel("composite quality score"); a1.set_ylabel("images")
    a1.set_title("Cohort quality distribution", fontsize=12, loc="left")
    counts = df["verdict"].value_counts()
    order = ["PASS", "ACCEPTABLE", "FAIL"]
    vals = [int(counts.get(k, 0)) for k in order]
    a2.bar(order, vals, color=[VERDICT_COLORS[k] for k in order], edgecolor="white")
    for i, v in enumerate(vals):
        a2.text(i, v, str(v), ha="center", va="bottom", fontsize=11, fontweight="bold")
    a2.set_title("Verdicts", fontsize=12, loc="left"); a2.grid(axis="x", alpha=0)
    a2.tick_params(axis="x", labelrotation=25)
    return _fig_to_rgb(fig)


def axis_heatmap(df, metric_names):
    """Image x axis score heatmap for the cohort (small-multiple-free overview)."""
    M = df[metric_names].values.T
    fig, ax = plt.subplots(figsize=(max(6, 0.28 * len(df)), 4.2))
    im = ax.imshow(M, aspect="auto", cmap="RdYlGn", vmin=0, vmax=1)
    ax.set_yticks(range(len(metric_names)))
    ax.set_yticklabels(metric_names, fontsize=8)
    ax.set_xlabel("image index")
    ax.set_title("Per-axis scores across the cohort", fontsize=12, loc="left")
    fig.colorbar(im, ax=ax, fraction=0.02, pad=0.01, label="score")
    ax.grid(False)
    return _fig_to_rgb(fig)


def metric_correlation(df, metric_names):
    C = np.corrcoef(df[metric_names].values.T)
    fig, ax = plt.subplots(figsize=(5.6, 5.0))
    im = ax.imshow(C, cmap="RdBu_r", vmin=-1, vmax=1)
    ax.set_xticks(range(len(metric_names))); ax.set_yticks(range(len(metric_names)))
    ax.set_xticklabels(metric_names, rotation=55, ha="right", fontsize=7.5)
    ax.set_yticklabels(metric_names, fontsize=7.5)
    ax.set_title("Quality-axis correlation", fontsize=12, loc="left")
    fig.colorbar(im, ax=ax, fraction=0.046, pad=0.02); ax.grid(False)
    return _fig_to_rgb(fig)


# --------------------------------------------------------------------- batch views
def batch_scatter(emb, batches, name, title_suffix=""):
    fig, ax = plt.subplots(figsize=(5.6, 4.8))
    batches = np.asarray(batches)
    cmap = plt.cm.get_cmap("tab10")
    for i, g in enumerate(np.unique(batches)):
        sel = batches == g
        ax.scatter(emb[sel, 0], emb[sel, 1], s=42, color=cmap(i % 10),
                   edgecolor="white", lw=0.6, label=str(g), alpha=0.9)
    ax.set_xlabel(f"{name}-1"); ax.set_ylabel(f"{name}-2")
    ax.legend(title="batch", frameon=False, fontsize=8)
    ax.set_title(f"Embedding by batch{title_suffix}", fontsize=12, loc="left")
    return _fig_to_rgb(fig)


def batch_before_after(emb_b, emb_a, batches, name):
    fig, axes = plt.subplots(1, 2, figsize=(10.2, 4.5), sharex=False)
    batches = np.asarray(batches); cmap = plt.cm.get_cmap("tab10")
    for ax, emb, ttl in [(axes[0], emb_b, "Before correction"),
                         (axes[1], emb_a, "After correction")]:
        for i, g in enumerate(np.unique(batches)):
            sel = batches == g
            ax.scatter(emb[sel, 0], emb[sel, 1], s=40, color=cmap(i % 10),
                       edgecolor="white", lw=0.6, label=str(g), alpha=0.9)
        ax.set_title(ttl, fontsize=12, loc="left")
        ax.set_xlabel(f"{name}-1"); ax.set_ylabel(f"{name}-2")
    axes[1].legend(title="batch", frameon=False, fontsize=8)
    return _fig_to_rgb(fig)


# --------------------------------------------------------------------- image panels
def av_overlay(rgb, vessels):
    """Overlay arteries (red) and veins (blue) from a deep A/V segmentation."""
    out = rgb.copy().astype(np.float32)
    a = vessels.get("artery"); v = vessels.get("vein")
    if a is None or v is None:
        return vessel_overlay(rgb, vessels)
    am = np.clip(a, 0, 1)[..., None]; vm = np.clip(v, 0, 1)[..., None]
    red = np.array([230, 60, 60], np.float32)
    blue = np.array([70, 130, 235], np.float32)
    out = out * (1 - 0.65 * am) + red * (0.65 * am)
    out = out * (1 - 0.65 * vm) + blue * (0.65 * vm)
    return np.clip(out, 0, 255).astype(np.uint8)


def vessel_overlay(rgb, vessels):
    """Overlay the vessel probability map (teal) and skeleton on the fundus."""
    out = rgb.copy().astype(np.float32)
    prob = vessels["prob_map"]
    heat = np.zeros_like(out)
    heat[..., 0] = 31; heat[..., 1] = 122; heat[..., 2] = 140     # teal
    a = np.clip(prob / (prob.max() + 1e-6), 0, 1)[..., None]
    out = out * (1 - 0.6 * a) + heat * (0.6 * a)
    sk = vessels["skeleton"]
    out[sk] = [230, 210, 60]                                       # amber skeleton
    return np.clip(out, 0, 255).astype(np.uint8)


def vessel_stats_panel(vessels):
    """Bar chart of structural vessel descriptors."""
    keys = ["density", "mean_vesselness", "skeleton_length", "fractal_dimension"]
    labels = ["Vessel density", "Mean vesselness", "Skeleton length", "Fractal dim"]
    vals = [vessels[k] for k in keys]
    fig, ax = plt.subplots(figsize=(5.2, 3.2))
    ax.barh(labels, vals, color=[TEAL, DEEP, GREEN, AMBER], edgecolor="white")
    for i, v in enumerate(vals):
        ax.text(v, i, f" {v:.3f}", va="center", fontsize=9, color=MUTED)
    ax.set_title(f"Vascular structure  ({vessels['backend']} backend)",
                 fontsize=12, loc="left")
    ax.grid(axis="y", alpha=0)
    return _fig_to_rgb(fig)


def quality_panel(thumbs, labels, verdicts, ncol=4, thumb=190):
    """Montage of image thumbnails with a coloured quality frame and label."""
    n = len(thumbs)
    if n == 0:
        return np.full((200, 400, 3), 250, np.uint8)
    nrow = int(np.ceil(n / ncol))
    pad, lab_h = 14, 30
    cell = thumb + 2 * pad + lab_h
    W, H = ncol * cell, nrow * cell
    canvas = np.full((H, W, 3), 248, np.uint8)
    for i, (t, lab, vd) in enumerate(zip(thumbs, labels, verdicts)):
        r, c = divmod(i, ncol)
        y0, x0 = r * cell, c * cell
        img = cv2.resize(t, (thumb, thumb))
        col = tuple(int(VERDICT_COLORS[vd].lstrip("#")[k:k+2], 16) for k in (0, 2, 4))
        fy, fx = y0 + pad, x0 + pad
        cv2.rectangle(canvas, (fx - 4, fy - 4), (fx + thumb + 4, fy + thumb + 4), col, 3)
        canvas[fy:fy + thumb, fx:fx + thumb] = img
        cv2.putText(canvas, lab[:22], (fx - 2, fy + thumb + 20),
                    cv2.FONT_HERSHEY_SIMPLEX, 0.44, (30, 40, 45), 1, cv2.LINE_AA)
    return canvas