File size: 15,879 Bytes
1f2c4ac
55fe6e9
1f2c4ac
934571c
 
3816943
e048678
 
1f2c4ac
4e02a60
1f2c4ac
 
 
 
 
 
 
 
3816943
 
 
55fe6e9
1f2c4ac
 
 
 
e048678
1f2c4ac
 
 
3816943
4e02a60
55fe6e9
cf61a4d
1f2c4ac
 
55fe6e9
1f2c4ac
 
3816943
1f2c4ac
 
 
 
 
 
 
3816943
4e02a60
1f2c4ac
 
 
4e02a60
1f2c4ac
4e02a60
55fe6e9
4e02a60
1f2c4ac
 
 
 
 
4e02a60
3816943
 
1f2c4ac
 
 
 
 
 
 
 
 
 
 
 
55fe6e9
 
1f2c4ac
 
 
 
e048678
1f2c4ac
 
 
 
 
 
 
 
e048678
934571c
 
e048678
1f2c4ac
 
 
 
e048678
934571c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e048678
934571c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e048678
934571c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e048678
934571c
 
 
e048678
934571c
 
e048678
 
 
1f2c4ac
55fe6e9
1f2c4ac
 
 
 
 
 
 
 
 
55fe6e9
 
 
 
 
 
 
1f2c4ac
55fe6e9
1f2c4ac
55fe6e9
 
1f2c4ac
 
 
 
 
 
 
e048678
1f2c4ac
 
e048678
934571c
1f2c4ac
 
 
e048678
934571c
e048678
 
1f2c4ac
 
 
 
 
55fe6e9
934571c
 
 
 
 
 
1f2c4ac
e048678
934571c
 
 
 
 
 
 
 
1f2c4ac
 
934571c
1f2c4ac
 
 
55fe6e9
1f2c4ac
934571c
1f2c4ac
 
934571c
 
 
 
1f2c4ac
 
 
 
 
 
 
 
 
 
 
e048678
934571c
e048678
934571c
 
 
 
e048678
934571c
e048678
934571c
e048678
934571c
 
 
 
 
e048678
 
 
 
934571c
e048678
 
934571c
 
1f2c4ac
 
 
 
 
 
e048678
 
934571c
 
 
 
1f2c4ac
 
 
e048678
 
1f2c4ac
 
 
e048678
 
 
934571c
 
1f2c4ac
 
 
934571c
e048678
934571c
 
 
 
 
 
 
e048678
 
1f2c4ac
 
 
 
55fe6e9
 
 
1f2c4ac
 
55fe6e9
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
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
"""
ECG -> ECHO Screening (V3.2)
Predicts ejection fraction (regression) and significant RWMA (binary screen)
from a 12-lead ECG image, gives a numerical model-confidence score, and a
cardiologist-referral recommendation.

Research / feasibility prototype -- decision support only, NOT a diagnostic device.
Model: 5-fold ensemble, EfficientNet-B3 (PTB-XL pretrained), multi-task heads.
"""

import os
import numpy as np
import cv2
from PIL import Image
import streamlit as st
import torch
import torch.nn as nn
import timm
import albumentations as A
from albumentations.pytorch import ToTensorV2

MODEL_PATH = "ensemble_clinical_v3_2.pth"
BACKBONE   = "tf_efficientnet_b3.ns_jft_in1k"
IMG_SIZE   = 384
EDGE_CROP  = 0.05
DEVICE     = torch.device("cpu")
VALIDATED_EF_MAE = 9.0   # cross-validated EF mean-absolute-error (EF points)

st.set_page_config(page_title="ECG -> ECHO Screening", page_icon=":anatomical_heart:", layout="wide")


class ECGNetV3(nn.Module):
    def __init__(self, dropout=0.45, drop_path=0.1):
        super().__init__()
        self.backbone = timm.create_model(
            BACKBONE, pretrained=False, num_classes=0,
            global_pool="avg", drop_rate=0.2, drop_path_rate=drop_path,
        )
        feat = self.backbone.num_features
        self.neck = nn.Sequential(
            nn.Linear(feat, 512), nn.LayerNorm(512), nn.GELU(), nn.Dropout(dropout),
            nn.Linear(512, 256), nn.LayerNorm(256), nn.GELU(), nn.Dropout(dropout * 0.7),
            nn.Linear(256, 128), nn.LayerNorm(128), nn.GELU(), nn.Dropout(dropout * 0.5),
        )
        self.ef_reg   = nn.Linear(128, 1)
        self.rwma_cls = nn.Linear(128, 2)

    def forward(self, x):
        z = self.neck(self.backbone(x))
        return {"ef_norm": self.ef_reg(z).squeeze(-1),
                "rwma_logits": self.rwma_cls(z)}


def preprocess_image(img_bgr):
    if EDGE_CROP > 0:
        H, W = img_bgr.shape[:2]; c = EDGE_CROP
        img_bgr = img_bgr[int(H*c):int(H*(1-c)), int(W*c):int(W*(1-c))]
    gray  = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2GRAY)
    clahe = cv2.createCLAHE(clipLimit=2.5, tileGridSize=(8, 8))
    enh   = clahe.apply(gray)
    kernel = np.array([[0, -1, 0], [-1, 5, -1], [0, -1, 0]])
    sharp = cv2.filter2D(enh, -1, kernel)
    sharp = cv2.normalize(sharp, None, 0, 255, cv2.NORM_MINMAX)
    return cv2.cvtColor(sharp, cv2.COLOR_GRAY2RGB)


val_tf = A.Compose([
    A.Resize(IMG_SIZE, IMG_SIZE),
    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
    ToTensorV2(),
])


@st.cache_resource(show_spinner=False)
def load_models():
    ckpt = torch.load(MODEL_PATH, map_location="cpu", weights_only=False)
    models = []
    for s in ckpt["fold_models"]:
        m = ECGNetV3().to(DEVICE); m.load_state_dict(s); m.eval()
        models.append(m)
    return models, ckpt.get("overall", {})


# ----------------------------------------------------------------------------- predict
def predict(img_rgb, models, sig_threshold=0.5):
    tensor = val_tf(image=img_rgb)["image"].unsqueeze(0).to(DEVICE)
    ef_vals, sig_probs = [], []
    with torch.no_grad():
        for m in models:
            out = m(tensor)
            ef_vals.append(float(out["ef_norm"].cpu()) * 100)
            sig_probs.append(float(torch.softmax(out["rwma_logits"], -1)[0, 1].cpu()))
    ef_vals = np.array(ef_vals); sig_probs = np.array(sig_probs)
    ef_mean, ef_std  = float(ef_vals.mean()),  float(ef_vals.std())
    sig_p,   sig_std = float(sig_probs.mean()), float(sig_probs.std())

    if   ef_mean >= 50: ef_sev = "Normal"
    elif ef_mean >= 40: ef_sev = "Mildly reduced"
    elif ef_mean >= 30: ef_sev = "Moderately reduced"
    else:               ef_sev = "Severely reduced"

    return {"ef_mean": ef_mean, "ef_std": ef_std,
            "ef_low":  round(max(0, ef_mean - 1.96 * ef_std), 1),
            "ef_high": round(min(100, ef_mean + 1.96 * ef_std), 1),
            "ef_sev":  ef_sev, "sig_p": sig_p, "sig_std": sig_std,
            "sig_flag": sig_p >= sig_threshold}


# ----------------------------------------------------------------------------- confidence (numerical %)
def ef_confidence_pct(ef_std):
    # Tighter agreement across the 5 fold-models -> higher confidence.
    return int(np.clip(round(100 - ef_std * 7.0), 40, 99))

def rwma_confidence_pct(p, sig_std):
    # Decisiveness (distance from 0.5) penalised by inter-model disagreement.
    decisiveness = max(p, 1 - p) * 100
    return int(np.clip(round(decisiveness - sig_std * 100), 40, 99))

def conf_color(pct, strong_cut):
    if pct >= strong_cut:      return "#2e7d32"   # strong
    if pct >= strong_cut - 15: return "#e8730c"   # moderate
    return "#c62828"                              # low


# ----------------------------------------------------------------------------- assessment + referral matrix
def assess(r, threshold, strong_cut):
    ef, sig = r["ef_mean"], r["sig_p"]
    ef_c = ef_confidence_pct(r["ef_std"])
    rw_c = rwma_confidence_pct(sig, r["sig_std"])

    ef_abn = ef < 50
    rw_abn = sig >= threshold
    abnormal = ef_abn or rw_abn

    # Overall confidence: for an ABNORMAL call, confidence that something is wrong
    # = the strongest abnormal finding. For a NORMAL call, we must be confident on
    # BOTH fronts, so take the weaker (min).
    if abnormal:
        confs = ([ef_c] if ef_abn else []) + ([rw_c] if rw_abn else [])
        overall = max(confs)
    else:
        overall = min(ef_c, rw_c)
    strong = overall >= strong_cut

    reasons = []
    if ef < 40:   reasons.append(f"Predicted EF {ef:.0f}% β€” moderately-to-severely reduced systolic function")
    elif ef < 50: reasons.append(f"Predicted EF {ef:.0f}% β€” mildly reduced systolic function")
    if sig >= 0.60:        reasons.append(f"High probability of significant wall-motion abnormality ({sig:.0%})")
    elif sig >= threshold: reasons.append(f"Possible significant wall-motion abnormality ({sig:.0%})")

    # ---- referral decision matrix (confidence x result) ----
    if abnormal:
        priority = (ef < 40) or (sig >= 0.60)
        level = "CARDIOLOGIST REFERRAL NEEDED" + (" β€” PRIORITY" if priority else "")
        color = "#c62828"
        if not strong:
            reasons.append(f"Model confidence is low ({overall}%) β€” refer and correlate clinically")
    elif strong:
        level = "NO REFERRAL NEEDED β€” AI SCREENING SUFFICIENT"
        color = "#2e7d32"
        reasons.append(f"Predicted EF {ef:.0f}% (normal) and low RWMA probability ({sig:.0%}), with high model confidence ({overall}%)")
    else:
        level = "CLINICAL CORRELATION ADVISED"
        color = "#f9a825"
        reasons.append(f"Result appears normal, but model confidence is low ({overall}%) β€” do not clear on AI alone; clinician review advised")

    return {"ef_c": ef_c, "rw_c": rw_c, "overall": overall, "strong": strong,
            "abnormal": abnormal, "level": level, "color": color, "reasons": reasons}


# ----------------------------------------------------------------------------- grad-cam
def grad_cam(model, img_rgb, mode="ef"):
    from pytorch_grad_cam import GradCAM, LayerCAM
    from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget

    class WrapEF(nn.Module):
        def __init__(s, m): super().__init__(); s.m = m
        def forward(s, x): return s.m(x)["ef_norm"].unsqueeze(1)
    class WrapRWMA(nn.Module):
        def __init__(s, m): super().__init__(); s.m = m
        def forward(s, x): return s.m(x)["rwma_logits"][:, 1:2]

    if mode == "ef":
        wrapper = WrapEF(model); cam_cls = LayerCAM
        target_layers = [model.backbone.blocks[-2], model.backbone.blocks[-1]]
    else:
        wrapper = WrapRWMA(model); cam_cls = GradCAM
        target_layers = [model.backbone.blocks[-1]]

    tensor = val_tf(image=img_rgb)["image"].unsqueeze(0).to(DEVICE)
    with cam_cls(model=wrapper, target_layers=target_layers) as cam:
        g = cam(input_tensor=tensor, targets=[ClassifierOutputTarget(0)])[0]
    g = np.clip(g, 0, None)
    if g.max() > 0: g = g / g.max()
    H, W = img_rgb.shape[:2]
    g = cv2.resize(g, (W, H))
    heat = cv2.applyColorMap(np.uint8(255 * g), cv2.COLORMAP_JET)
    heat = cv2.cvtColor(heat, cv2.COLOR_BGR2RGB)
    return cv2.addWeighted(img_rgb.astype(np.uint8), 0.55, heat, 0.45, 0)


# ----------------------------------------------------------------------------- UI
st.markdown(
    "<h1 style='margin-bottom:0'>ECG -> ECHO Screening</h1>"
    "<p style='color:#666;margin-top:4px'>Estimates ejection fraction, screens for significant "
    "wall-motion abnormality, and gives a confidence-scored referral recommendation from a 12-lead ECG image.</p>",
    unsafe_allow_html=True,
)
st.warning(
    "**Decision support only β€” NOT a medical device.** Research/feasibility model trained on 500 "
    "ECG-echo pairs from a single center. The referral suggestion and confidence score are aids for a "
    "clinician; they do not replace echocardiography or physician judgment. The final decision rests "
    "with the treating doctor."
)

with st.sidebar:
    st.header("Settings")
    threshold = st.slider(
        "RWMA referral threshold", 0.20, 0.70, 0.40, 0.05,
        help="Lower = more sensitive (flags more cases). Screening favors higher sensitivity.",
    )
    strong_cut = st.slider(
        "Strong-confidence cutoff (%)", 50, 90, 70, 5,
        help="At or above this, model confidence is treated as 'strong'. A normal result with strong "
             "confidence is cleared as 'AI sufficient'; below it, clinician review is advised.",
    )
    st.caption("A clear 12-lead ECG image (phone photo or scan) works best.")
    with st.expander("Referral logic"):
        st.markdown(
            "| Confidence | Result | Recommendation |\n|---|---|---|\n"
            "| Strong | Normal | No referral β€” AI sufficient |\n"
            "| Strong | Abnormal | Cardiologist referral |\n"
            "| Low | Abnormal | Cardiologist referral |\n"
            "| Low | Normal | Clinical correlation advised |"
        )

try:
    models, overall_metrics = load_models()
    model_ok = True
except Exception as e:
    model_ok = False
    st.error(f"Could not load model file `{MODEL_PATH}`. Make sure it is uploaded to this Space.\n\n{e}")

if model_ok and overall_metrics:
    with st.expander("Model performance (cross-validated, n=500)"):
        c1, c2, c3, c4 = st.columns(4)
        c1.metric("EF severe AUROC", f"{overall_metrics.get('ef_severe_auroc', float('nan')):.2f}")
        c2.metric("EF MAE", f"{overall_metrics.get('ef_mae', float('nan')):.1f}%")
        c3.metric("RWMA AUROC", f"{overall_metrics.get('rwma_auroc', float('nan')):.2f}")
        c4.metric("EF within +/-10%", f"{overall_metrics.get('ef_within_10', float('nan')):.0f}%")

uploaded = st.file_uploader("Upload a 12-lead ECG image", type=["jpg", "jpeg", "png"])

if uploaded and model_ok:
    file_bytes = np.frombuffer(uploaded.read(), np.uint8)
    img_bgr = cv2.imdecode(file_bytes, cv2.IMREAD_COLOR)
    if img_bgr is None:
        st.error("Could not read that image. Try a different file.")
    else:
        proc = preprocess_image(img_bgr)
        with st.spinner("Running 5-model ensemble..."):
            r = predict(proc, models, sig_threshold=threshold)
        a = assess(r, threshold, strong_cut)

        # ---- referral recommendation + overall confidence (top) ----
        conf_tag = "STRONG" if a["strong"] else "LOW"
        conf_tag_color = "#2e7d32" if a["strong"] else "#c62828"
        reason_html = "".join(f"<li style='margin:2px 0'>{x}</li>" for x in a["reasons"])
        st.markdown(
            f"<div style='border-left:8px solid {a['color']};background:#fafafa;border-radius:10px;"
            f"padding:16px 20px;margin:6px 0 14px 0'>"
            f"<div style='display:flex;justify-content:space-between;align-items:center'>"
            f"<div style='color:#888;font-size:13px;letter-spacing:1px'>SCREENING RECOMMENDATION</div>"
            f"<div style='font-size:15px;color:#555'>Model confidence: "
            f"<b style='color:{conf_tag_color};font-size:22px'>{a['overall']}%</b> "
            f"<span style='color:{conf_tag_color};font-weight:700'>({conf_tag})</span></div>"
            f"</div>"
            f"<div style='font-size:25px;font-weight:800;color:{a['color']};margin:6px 0 8px 0'>{a['level']}</div>"
            f"<ul style='margin:0 0 0 18px;color:#333;font-size:14px'>{reason_html}</ul>"
            f"</div>", unsafe_allow_html=True,
        )

        # ---- EF + RWMA detail cards with numerical confidence ----
        ef = r["ef_mean"]
        ef_color = "#2e7d32" if ef >= 50 else "#f9a825" if ef >= 40 else "#e65100" if ef >= 30 else "#c62828"
        efc_col = conf_color(a["ef_c"], strong_cut)
        rwc_col = conf_color(a["rw_c"], strong_cut)

        col1, col2 = st.columns(2)
        with col1:
            st.markdown(
                f"<div style='border:1px solid #ddd;border-radius:12px;padding:18px'>"
                f"<div style='color:#888;font-size:14px'>EJECTION FRACTION</div>"
                f"<div style='font-size:42px;font-weight:700;color:{ef_color}'>{ef:.1f}%</div>"
                f"<div style='color:#555'>{r['ef_sev']}</div>"
                f"<div style='margin-top:8px;font-size:14px;color:#444'>Confidence: "
                f"<b style='color:{efc_col};font-size:17px'>{a['ef_c']}%</b></div>"
                f"<div style='font-size:12px;color:#999;margin-top:3px'>5-model range {r['ef_low']}-{r['ef_high']}%"
                f" Β· validated typical error Β±{VALIDATED_EF_MAE:.0f} pts</div>"
                f"</div>", unsafe_allow_html=True,
            )
        with col2:
            flag = r["sig_flag"]; box = "#c62828" if flag else "#2e7d32"
            label = "SIGNIFICANT" if flag else "Non-significant"
            st.markdown(
                f"<div style='border:1px solid #ddd;border-radius:12px;padding:18px'>"
                f"<div style='color:#888;font-size:14px'>WALL-MOTION ABNORMALITY</div>"
                f"<div style='font-size:30px;font-weight:700;color:{box};margin-top:4px'>{label}</div>"
                f"<div style='color:#555;margin-top:6px'>Significant probability: "
                f"<b>{r['sig_p']:.0%}</b> (threshold {threshold:.0%})</div>"
                f"<div style='margin-top:8px;font-size:14px;color:#444'>Confidence: "
                f"<b style='color:{rwc_col};font-size:17px'>{a['rw_c']}%</b></div>"
                f"</div>", unsafe_allow_html=True,
            )

        with st.expander("How the confidence score is computed"):
            st.markdown(
                "- The **confidence score (%)** reflects how strongly the 5 ensemble models agree on this "
                "ECG β€” for EF, how tightly their predictions cluster; for RWMA, how decisive and consistent "
                "their vote is.\n"
                "- **Overall confidence** is the weakest of the two when the result is normal (we must be sure "
                "on both fronts to clear a patient) and the strongest abnormal finding when something looks wrong.\n"
                "- It measures **model agreement, not guaranteed accuracy.** A high score is reassuring but never "
                "a substitute for clinical judgment; a low score is itself a reason to involve a clinician."
            )

        st.divider()
        st.subheader("Where the model is looking (Grad-CAM)")
        st.caption("Heatmaps should fall on the ECG waveforms, not borders or text.")
        t1, t2, t3 = st.tabs(["Preprocessed input", "EF attention", "RWMA attention"])
        with t1: st.image(proc, use_column_width=True)
        with t2: st.image(grad_cam(models[0], proc, "ef"), use_column_width=True)
        with t3: st.image(grad_cam(models[0], proc, "rwma"), use_column_width=True)

elif not uploaded:
    st.info("Upload an ECG image to run the pipeline.")