drvikasgaur commited on
Commit
f5911f2
·
verified ·
1 Parent(s): 82b29d7

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +1134 -0
app.py ADDED
@@ -0,0 +1,1134 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # appgradiofinal3_radio.py
2
+ # Gradio — TBNet + Lung U-Net Auto Mask + Grad-CAM + RADIO
3
+ # + SAFER PHONE MODE + MASK POST-PROCESSING + MASK SANITY FAILSAFE
4
+ # + 3-STATE CONSENSUS (LOW / INDET / SCREEN+)
5
+ #
6
+ # Run:
7
+ # python appgradiofinal3_radio.py
8
+ #
9
+ # Requirements:
10
+ # pip install gradio timm torchvision opencv-python pillow transformers einops
11
+
12
+ import os
13
+ import cv2
14
+ import numpy as np
15
+
16
+ import torch
17
+ import torch.nn as nn
18
+ import timm
19
+ import gradio as gr
20
+
21
+ from torchvision import transforms
22
+ from typing import List, Tuple, Dict, Any, Optional
23
+
24
+ # RADIO deps (same env as TBNet)
25
+ from transformers import AutoModel, CLIPImageProcessor
26
+ from einops import rearrange
27
+ from PIL import Image
28
+
29
+
30
+ # ============================================================
31
+ # USER CONFIG
32
+ # ============================================================
33
+
34
+ # ---- Default TB/Lung weights ----
35
+ DEFAULT_TB_WEIGHTS = "weights/tb_best.pt"
36
+ DEFAULT_LUNG_WEIGHTS = "weights/lung_unet.pt"
37
+
38
+ # ---- RADIO config (same env as TB) ----
39
+ RADIO_HF_REPO = "nvidia/C-RADIOv4-SO400M"
40
+ RADIO_REVISION = "c0457f5dc26ca145f954cd4fc5bb6114e5705ad8"
41
+
42
+ RADIO_RAW_HEAD_PATH = "weights/radio_best_raw.pt"
43
+ RADIO_MASKED_HEAD_PATH = "weights/radio_best_masked.pt"
44
+
45
+ RADIO_IMG_SIZE = 320
46
+ RADIO_PATCH_SIZE = 16
47
+ RADIO_THR_SCREEN = 0.05
48
+ RADIO_THR_RED = 0.23
49
+ RADIO_MASKED_MIN_COV = 0.15
50
+ RADIO_GATE_DEFAULT = 0.21
51
+
52
+ # ---- Consensus logic thresholds ----
53
+ TBNET_SCREEN_THR = 0.30
54
+ TBNET_MARGIN = 0.03 # 3% margin around threshold → INDET zone
55
+
56
+ RADIO_SCREEN_THR = RADIO_THR_SCREEN
57
+ RADIO_MARGIN = 0.02 # 2% margin around radio screen threshold
58
+
59
+ # ---- Mask fail-safes ----
60
+ FAIL_COV = 0.10 # <10% -> segmentation fail
61
+ WARN_COV = 0.18 # <18% -> warn
62
+ # if mask looks like a single lung / cropped, we fail-safe and do not output TB score
63
+ FAILSAFE_ON_BAD_MASK = True
64
+
65
+ # ---- Device policy ----
66
+ FORCE_CPU = True # set True if you want TB+RADIO to always run CPU
67
+ DEVICE = torch.device("cpu" if FORCE_CPU else ("cuda" if torch.cuda.is_available() else "cpu"))
68
+
69
+
70
+ # ============================================================
71
+ # CLINICAL DISCLAIMER / REPORT TEXT
72
+ # ============================================================
73
+ CLINICAL_DISCLAIMER = """
74
+ ⚠️ IMPORTANT CLINICAL NOTICE (Decision Support Only)
75
+ This AI system is for **research/decision support** and is NOT a diagnostic device.
76
+ It may NOT reliably detect early/subtle tuberculosis, including **MILIARY TB**,
77
+ which can appear near-normal or subtle on chest X-ray (especially on phone photos / WhatsApp images).
78
+
79
+ If clinical suspicion exists (fever, weight loss, immunosuppression, known exposure),
80
+ recommend **CBNAAT / GeneXpert**, sputum studies, and/or **CT chest** regardless of AI output.
81
+ """
82
+
83
+ REPORT_LABELS = {
84
+ "GREEN": {"title": "LIKELY NORMAL", "summary": "No radiographic features suggestive of pulmonary tuberculosis detected by AI."},
85
+ "YELLOW": {"title": "RADIOLOGIST INTERPRETATION RECOMMENDED", "summary": "This AI output is indeterminate / not definitive.A qualified radiologist review is recommended to confirm findings and correlate clinically."},
86
+ "RED": {"title": "LIKELY TB", "summary": "AI detected focal lung patterns commonly associated with pulmonary tuberculosis. This is not a diagnosis; microbiological confirmation is required."},
87
+ }
88
+
89
+ CLINICAL_GUIDANCE = (
90
+ "If clinical suspicion for tuberculosis exists, further evaluation "
91
+ "(e.g., CBNAAT / GeneXpert, sputum studies, CT chest) is recommended "
92
+ "regardless of AI output."
93
+ )
94
+
95
+
96
+ # ============================================================
97
+ # LUNG U-NET (INFERENCE)
98
+ # ============================================================
99
+ class DoubleConv(nn.Module):
100
+ def __init__(self, in_c, out_c):
101
+ super().__init__()
102
+ self.net = nn.Sequential(
103
+ nn.Conv2d(in_c, out_c, 3, padding=1),
104
+ nn.BatchNorm2d(out_c),
105
+ nn.ReLU(inplace=True),
106
+ nn.Conv2d(out_c, out_c, 3, padding=1),
107
+ nn.BatchNorm2d(out_c),
108
+ nn.ReLU(inplace=True),
109
+ )
110
+ def forward(self, x): return self.net(x)
111
+
112
+ class LungUNet(nn.Module):
113
+ def __init__(self):
114
+ super().__init__()
115
+ self.d1 = DoubleConv(1, 64)
116
+ self.d2 = DoubleConv(64, 128)
117
+ self.d3 = DoubleConv(128, 256)
118
+ self.d4 = DoubleConv(256, 512)
119
+ self.pool = nn.MaxPool2d(2)
120
+ self.mid = DoubleConv(512, 1024)
121
+ self.u4 = nn.ConvTranspose2d(1024, 512, 2, 2)
122
+ self.u3 = nn.ConvTranspose2d(512, 256, 2, 2)
123
+ self.u2 = nn.ConvTranspose2d(256, 128, 2, 2)
124
+ self.u1 = nn.ConvTranspose2d(128, 64, 2, 2)
125
+ self.c4 = DoubleConv(1024, 512)
126
+ self.c3 = DoubleConv(512, 256)
127
+ self.c2 = DoubleConv(256, 128)
128
+ self.c1 = DoubleConv(128, 64)
129
+ self.out = nn.Conv2d(64, 1, 1)
130
+
131
+ def forward(self, x):
132
+ d1 = self.d1(x)
133
+ d2 = self.d2(self.pool(d1))
134
+ d3 = self.d3(self.pool(d2))
135
+ d4 = self.d4(self.pool(d3))
136
+ m = self.mid(self.pool(d4))
137
+ x = self.c4(torch.cat([self.u4(m), d4], 1))
138
+ x = self.c3(torch.cat([self.u3(x), d3], 1))
139
+ x = self.c2(torch.cat([self.u2(x), d2], 1))
140
+ x = self.c1(torch.cat([self.u1(x), d1], 1))
141
+ return self.out(x)
142
+
143
+
144
+ # ============================================================
145
+ # TB MODEL + GRAD-CAM
146
+ # ============================================================
147
+ class TBNet(nn.Module):
148
+ def __init__(self, backbone="efficientnet_b0"):
149
+ super().__init__()
150
+ self.backbone = timm.create_model(backbone, pretrained=False, num_classes=0, global_pool="avg")
151
+ self.fc = nn.Linear(self.backbone.num_features, 1)
152
+ def forward(self, x): return self.fc(self.backbone(x)).view(-1)
153
+
154
+ def load_tb_weights(model: nn.Module, ckpt_path: str, device: torch.device):
155
+ sd = torch.load(ckpt_path, map_location=device)
156
+ model.load_state_dict(sd, strict=True)
157
+
158
+ class GradCAM:
159
+ def __init__(self, model: nn.Module, target_layer: nn.Module):
160
+ self.model = model
161
+ self.activ = None
162
+ self.grad = None
163
+ target_layer.register_forward_hook(self._fwd)
164
+ target_layer.register_full_backward_hook(self._bwd)
165
+
166
+ def _fwd(self, _, __, out): self.activ = out
167
+ def _bwd(self, _, grad_in, grad_out): self.grad = grad_out[0]
168
+
169
+ def generate(self, x: torch.Tensor) -> Tuple[np.ndarray, float, float]:
170
+ with torch.enable_grad():
171
+ self.model.zero_grad(set_to_none=True)
172
+ logits = self.model(x)
173
+ score = logits[0]
174
+ score.backward()
175
+
176
+ A = self.activ[0]
177
+ G = self.grad[0]
178
+ w = G.mean(dim=(1, 2))
179
+ cam = (w[:, None, None] * A).sum(dim=0)
180
+ cam = torch.relu(cam)
181
+ cam = cam - cam.min()
182
+ cam = cam / (cam.max() + 1e-8)
183
+
184
+ logit = float(logits.detach().cpu()[0].item())
185
+ prob = float(torch.sigmoid(logits.detach().cpu())[0].item())
186
+
187
+ return cam.detach().cpu().numpy(), prob, logit
188
+
189
+
190
+ # ============================================================
191
+ # PREPROCESS HELPERS + QUALITY
192
+ # ============================================================
193
+ def preprocess_for_lung_unet(gray_u8: np.ndarray) -> torch.Tensor:
194
+ g = gray_u8.astype(np.float32)
195
+ g = cv2.resize(g, (256, 256), interpolation=cv2.INTER_AREA)
196
+ lo, hi = np.percentile(g, (1, 99))
197
+ g = np.clip(g, lo, hi)
198
+ g = (g - lo) / (hi - lo + 1e-8)
199
+ return torch.from_numpy(g).unsqueeze(0).unsqueeze(0).float()
200
+
201
+ def tb_training_preprocess(gray_u8: np.ndarray) -> np.ndarray:
202
+ gray = gray_u8.astype(np.float32)
203
+ lo, hi = np.percentile(gray, (1, 99))
204
+ gray = np.clip(gray, lo, hi)
205
+ gray = (gray - lo) / (hi - lo + 1e-8)
206
+ return gray
207
+
208
+ def laplacian_sharpness(gray_u8: np.ndarray) -> float:
209
+ g = cv2.resize(gray_u8, (512, 512), interpolation=cv2.INTER_AREA)
210
+ g = cv2.GaussianBlur(g, (3, 3), 0)
211
+ return float(cv2.Laplacian(g, cv2.CV_64F).var())
212
+
213
+ def exposure_scores(gray_u8: np.ndarray) -> Tuple[float, float]:
214
+ lo = float((gray_u8 < 10).mean())
215
+ hi = float((gray_u8 > 245).mean())
216
+ return lo, hi
217
+
218
+ def border_fraction(gray_u8: np.ndarray) -> float:
219
+ h, w = gray_u8.shape
220
+ b = max(5, int(0.06 * min(h, w)))
221
+ top = gray_u8[:b, :]
222
+ bot = gray_u8[-b:, :]
223
+ left = gray_u8[:, :b]
224
+ right = gray_u8[:, -b:]
225
+ def frac_border(x): return float(((x < 15) | (x > 240)).mean())
226
+ return float(np.mean([frac_border(top), frac_border(bot), frac_border(left), frac_border(right)]))
227
+
228
+ def phone_quality_report(gray_u8: np.ndarray) -> Tuple[float, List[str]]:
229
+ warnings: List[str] = []
230
+ h, w = gray_u8.shape
231
+
232
+ score = 100.0
233
+
234
+ if min(h, w) < 400:
235
+ warnings.append("Low resolution (may reduce detection reliability).")
236
+ score -= 8
237
+
238
+ sharp = laplacian_sharpness(gray_u8)
239
+ lo_clip, hi_clip = exposure_scores(gray_u8)
240
+ border = border_fraction(gray_u8)
241
+
242
+ likely_phone = (border > 0.35) or (lo_clip > 0.10) or (hi_clip > 0.05)
243
+
244
+ if likely_phone:
245
+ if sharp < 40:
246
+ score -= 25; warnings.append("Blurry / motion blur detected (phone capture).")
247
+ elif sharp < 80:
248
+ score -= 12; warnings.append("Slight blur detected.")
249
+ else:
250
+ if sharp < 30:
251
+ score -= 8; warnings.append("Low fine-detail / mild blur (digital CXR or downsample).")
252
+
253
+ if hi_clip > 0.05:
254
+ score -= 15; warnings.append("Overexposed highlights (washed out areas).")
255
+ if lo_clip > 0.10:
256
+ score -= 12; warnings.append("Underexposed shadows (very dark areas).")
257
+
258
+ if border > 0.55:
259
+ score -= 18; warnings.append("Large border/margins detected (screenshot/phone framing).")
260
+ elif border > 0.35:
261
+ score -= 10; warnings.append("Some border/margins detected.")
262
+
263
+ return float(np.clip(score, 0, 100)), warnings
264
+
265
+ def auto_border_crop(gray_u8: np.ndarray) -> np.ndarray:
266
+ g = gray_u8.copy()
267
+ g_blur = cv2.GaussianBlur(g, (5, 5), 0)
268
+ _, th = cv2.threshold(g_blur, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)
269
+ if th.mean() > 127: th = 255 - th
270
+
271
+ k = max(3, int(0.01 * min(g.shape)))
272
+ kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))
273
+ th = cv2.morphologyEx(th, cv2.MORPH_CLOSE, kernel, iterations=2)
274
+
275
+ contours, _ = cv2.findContours(th, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
276
+ if not contours: return gray_u8
277
+
278
+ c = max(contours, key=cv2.contourArea)
279
+ x, y, w, h = cv2.boundingRect(c)
280
+ H, W = gray_u8.shape
281
+ if w * h < 0.20 * (H * W): return gray_u8
282
+
283
+ pad = int(0.03 * min(H, W))
284
+ x1 = max(0, x - pad); y1 = max(0, y - pad)
285
+ x2 = min(W, x + w + pad); y2 = min(H, y + h + pad)
286
+ return gray_u8[y1:y2, x1:x2]
287
+
288
+ def apply_clahe(gray_u8: np.ndarray) -> np.ndarray:
289
+ clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))
290
+ return clahe.apply(gray_u8)
291
+
292
+ def phone_preprocess(gray_u8: np.ndarray) -> np.ndarray:
293
+ """
294
+ Safer phone preprocessing:
295
+ - only crop if border artifacts suggest a framed/screenshot input
296
+ - only apply CLAHE if underexposed or low sharpness
297
+ - crop sanity check to avoid destroying clean digital CXRs
298
+ """
299
+ sharp = laplacian_sharpness(gray_u8)
300
+ lo_clip, _hi_clip = exposure_scores(gray_u8)
301
+ border = border_fraction(gray_u8)
302
+
303
+ g = gray_u8
304
+
305
+ if border > 0.35:
306
+ cropped = auto_border_crop(g)
307
+ if cropped.size >= 0.70 * g.size:
308
+ g = cropped
309
+
310
+ if lo_clip > 0.10 or sharp < 80:
311
+ g = apply_clahe(g)
312
+
313
+ return g
314
+
315
+ def cam_entropy(cam: np.ndarray) -> float:
316
+ cam = cam.astype(np.float32)
317
+ cam = cam / (cam.sum() + 1e-8)
318
+ return float(-np.sum(cam * np.log(cam + 1e-8)))
319
+
320
+ def detect_diffuse_risk(prob_tb: float, cam_up: np.ndarray, quality_score: float) -> bool:
321
+ if quality_score < 55:
322
+ return False
323
+
324
+ # Only apply diffuse-risk heuristic in the "near-threshold negative" zone
325
+ if prob_tb < 0.05:
326
+ return False
327
+
328
+ ent = cam_entropy(cam_up)
329
+ return (prob_tb < TBNET_SCREEN_THR) and (ent > 6.5)
330
+
331
+
332
+ def confidence_band(prob_tb: float, quality_score: float, diffuse: bool):
333
+ if prob_tb < 0.05 and quality_score >= 45 and not diffuse:
334
+ return ("GREEN", "LIKELY NORMAL — very low AI TB signal (image quality suboptimal)")
335
+ if quality_score < 55:
336
+ return ("YELLOW", "NO DEFINITE TB FEATURES — low image quality, treat as indeterminate")
337
+ if diffuse:
338
+ return ("YELLOW", "NO DEFINITE TB FEATURES — non-focal/diffuse attention pattern")
339
+ if prob_tb >= TBNET_SCREEN_THR:
340
+ return ("YELLOW", "NO DEFINITE TB FEATURES — AI confidence limited")
341
+ return ("GREEN", "LIKELY NORMAL — no strong AI signal for TB")
342
+
343
+ def make_mask_overlay(gray_u8: np.ndarray, mask_u8: np.ndarray) -> np.ndarray:
344
+ base = cv2.cvtColor(gray_u8, cv2.COLOR_GRAY2RGB)
345
+ mask_color = cv2.applyColorMap((mask_u8 * 255).astype(np.uint8), cv2.COLORMAP_JET)
346
+ return cv2.addWeighted(base, 0.75, mask_color, 0.25, 0)
347
+
348
+ def fill_holes(binary_u8: np.ndarray) -> np.ndarray:
349
+ m = (binary_u8 * 255).astype(np.uint8)
350
+ h, w = m.shape
351
+ flood = m.copy()
352
+ mask = np.zeros((h+2, w+2), np.uint8)
353
+ cv2.floodFill(flood, mask, (0, 0), 255)
354
+ holes = cv2.bitwise_not(flood)
355
+ filled = cv2.bitwise_or(m, holes)
356
+ return (filled > 0).astype(np.uint8)
357
+
358
+ def keep_top_k_components(binary_u8: np.ndarray, k: int = 2) -> np.ndarray:
359
+ m = (binary_u8 > 0).astype(np.uint8)
360
+ n, labels = cv2.connectedComponents(m)
361
+ if n <= 1:
362
+ return m
363
+ areas = []
364
+ for i in range(1, n):
365
+ areas.append((i, int((labels == i).sum())))
366
+ areas.sort(key=lambda x: x[1], reverse=True)
367
+ keep_ids = set([i for i, _ in areas[:k]])
368
+ out = np.zeros_like(m)
369
+ for i in keep_ids:
370
+ out[labels == i] = 1
371
+ return out
372
+
373
+ def mask_sanity_warnings(mask_full_u8: np.ndarray) -> List[str]:
374
+ m = (mask_full_u8 > 0).astype(np.uint8)
375
+ n, labels = cv2.connectedComponents(m)
376
+ warns = []
377
+
378
+ if n <= 2:
379
+ warns.append("Only one lung component detected (possible crop/segmentation failure).")
380
+ return warns
381
+
382
+ areas = []
383
+ for i in range(1, n):
384
+ areas.append(int((labels == i).sum()))
385
+ areas.sort(reverse=True)
386
+ total = int(m.sum())
387
+ top1 = areas[0]
388
+ top2 = areas[1] if len(areas) > 1 else 0
389
+
390
+ if total > 0 and top1 / total > 0.80:
391
+ warns.append("Mask dominated by a single component (likely one lung / cropped view).")
392
+
393
+ border = np.concatenate([m[0, :], m[-1, :], m[:, 0], m[:, -1]])
394
+ if border.mean() > 0.05:
395
+ warns.append("Lung mask touches image border (possible cropped/non-standard CXR).")
396
+
397
+ if total > 0 and (top1 + top2) / total < 0.90:
398
+ warns.append("Significant mask fragmentation/holes (post-processing may be insufficient).")
399
+
400
+ return warns
401
+
402
+ def recommendation_for_band(band: Optional[str]) -> str:
403
+ if band in (None, "YELLOW"):
404
+ return "✅ Recommendation: Radiologist interpretation recommended (AI result is indeterminate / not definitive)."
405
+ if band == "RED":
406
+ return "✅ Recommendation: Urgent clinician/radiologist review + microbiological confirmation (CBNAAT/GeneXpert, sputum)."
407
+ return "✅ Recommendation: If symptoms/risk factors exist, clinician/radiologist correlation is still advised."
408
+
409
+
410
+ # ============================================================
411
+ # CONSENSUS LOGIC (TBNet vs RADIO) — 3-state
412
+ # ============================================================
413
+ def tbnet_state(tb_prob: float, tb_band: str) -> str:
414
+ if tb_band == "RED":
415
+ return "TB+"
416
+ if tb_prob >= TBNET_SCREEN_THR:
417
+ return "SCREEN+"
418
+ return "LOW"
419
+
420
+ def radio_state_from_prob(radio_prob: float) -> str:
421
+ if radio_prob >= RADIO_THR_RED:
422
+ return "TB+"
423
+ if radio_prob >= RADIO_THR_SCREEN:
424
+ return "SCREEN+"
425
+ return "LOW"
426
+
427
+ def build_consensus(
428
+ tb_prob: Optional[float],
429
+ tb_band: Optional[str],
430
+ radio_raw: Optional[float],
431
+ radio_masked: Optional[float],
432
+ radio_band: Optional[str] = None
433
+ ) -> Tuple[str, str]:
434
+
435
+ if tb_prob is None or tb_band is None:
436
+ return ("N/A", "TBNet unavailable (lung segmentation failed / fail-safe).")
437
+
438
+ # PRIMARY = masked if available else raw
439
+ if radio_masked is not None:
440
+ radio_primary = radio_masked
441
+ radio_used = "MASKED"
442
+ else:
443
+ radio_primary = radio_raw
444
+ radio_used = "RAW"
445
+
446
+ if radio_primary is None:
447
+ return ("TBNet only", f"RADIO unavailable → TBNet={tb_prob:.4f} (band={tb_band}).")
448
+
449
+ t = tbnet_state(tb_prob, tb_band)
450
+ r = radio_state_from_prob(radio_primary)
451
+
452
+ rb = f" (RADIO band={radio_band})" if radio_band else ""
453
+
454
+ if t == r:
455
+ return (
456
+ f"AGREE: {t}",
457
+ f"Both: {t}. TBNet={tb_prob:.4f}, RADIO({radio_used})={radio_primary:.4f}{rb}."
458
+ )
459
+
460
+ if (t in ("SCREEN+", "TB+") and r == "LOW") or (r in ("SCREEN+", "TB+") and t == "LOW"):
461
+ return (
462
+ "DISAGREE",
463
+ f"Strong disagreement: TBNet={t} (band={tb_band}) vs RADIO={r} ({radio_used})={radio_primary:.4f}{rb}."
464
+ )
465
+
466
+ return (
467
+ "MIXED/INDET",
468
+ f"Mixed/uncertain: TBNet={t} (band={tb_band}) vs RADIO={r} ({radio_used})={radio_primary:.4f}{rb}."
469
+ )
470
+
471
+
472
+ # ============================================================
473
+ # TB + LUNG MODEL BUNDLE (cached)
474
+ # ============================================================
475
+ class ModelBundle:
476
+ def __init__(self):
477
+ self.device = DEVICE
478
+ self.tb = None
479
+ self.cammer = None
480
+ self.lung = None
481
+ self.tb_path = None
482
+ self.lung_path = None
483
+ self.backbone = "efficientnet_b0"
484
+
485
+ self.tfm = transforms.Compose([
486
+ transforms.ToPILImage(),
487
+ transforms.Resize((224, 224)),
488
+ transforms.ToTensor(),
489
+ transforms.Normalize(mean=[0.485, 0.456, 0.406],
490
+ std=[0.229, 0.224, 0.225]),
491
+ ])
492
+
493
+ def load(self, tb_weights: str, lung_weights: str, backbone: str = "efficientnet_b0"):
494
+ if (self.tb_path != tb_weights) or (self.tb is None) or (self.cammer is None) or (self.backbone != backbone):
495
+ tb = TBNet(backbone=backbone).to(self.device)
496
+ load_tb_weights(tb, tb_weights, self.device)
497
+ tb.eval()
498
+ self.tb = tb
499
+ self.cammer = GradCAM(tb, tb.backbone.conv_head)
500
+ self.tb_path = tb_weights
501
+ self.backbone = backbone
502
+
503
+ if (self.lung_path != lung_weights) or (self.lung is None):
504
+ lung = LungUNet().to(self.device)
505
+ lung.load_state_dict(torch.load(lung_weights, map_location=self.device))
506
+ lung.eval()
507
+ self.lung = lung
508
+ self.lung_path = lung_weights
509
+
510
+ BUNDLE = ModelBundle()
511
+
512
+
513
+ # ============================================================
514
+ # RADIO BUNDLE (cached)
515
+ # ============================================================
516
+ class RadioMLPHead(nn.Module):
517
+ def __init__(self, dim: int, hidden: int = 512, dropout: float = 0.2):
518
+ super().__init__()
519
+ self.net = nn.Sequential(
520
+ nn.LayerNorm(dim),
521
+ nn.Dropout(dropout),
522
+ nn.Linear(dim, hidden),
523
+ nn.GELU(),
524
+ nn.Dropout(dropout),
525
+ nn.Linear(hidden, 1),
526
+ )
527
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
528
+ return self.net(x).squeeze(1)
529
+
530
+ class RadioBundle:
531
+ def __init__(self):
532
+ self.loaded = False
533
+ self.processor = None
534
+ self.radio = None
535
+ self.raw_head = None
536
+ self.masked_head = None
537
+ self.summary_dim = None
538
+ self.device_str = None
539
+
540
+ def load(self, device: torch.device):
541
+ dev_str = str(device)
542
+ if self.loaded and self.device_str == dev_str:
543
+ return
544
+
545
+ if not os.path.exists(RADIO_RAW_HEAD_PATH):
546
+ raise FileNotFoundError(f"RADIO raw head not found: {RADIO_RAW_HEAD_PATH}")
547
+ if not os.path.exists(RADIO_MASKED_HEAD_PATH):
548
+ raise FileNotFoundError(f"RADIO masked head not found: {RADIO_MASKED_HEAD_PATH}")
549
+
550
+ self.processor = CLIPImageProcessor.from_pretrained(RADIO_HF_REPO, revision=RADIO_REVISION)
551
+ dtype = torch.float16 if device.type == "cuda" else torch.float32
552
+
553
+ self.radio = AutoModel.from_pretrained(
554
+ RADIO_HF_REPO,
555
+ revision=RADIO_REVISION,
556
+ trust_remote_code=True,
557
+ dtype=dtype,
558
+ ).eval().to(device)
559
+
560
+ with torch.no_grad():
561
+ dummy = torch.zeros((1, 3, RADIO_IMG_SIZE, RADIO_IMG_SIZE), device=device, dtype=dtype)
562
+ summary, _ = self.radio(dummy)
563
+ self.summary_dim = int(summary.shape[-1])
564
+
565
+ def _load_head(path: str) -> nn.Module:
566
+ ckpt = torch.load(path, map_location="cpu")
567
+ dim = int(ckpt.get("dim", self.summary_dim))
568
+ head = RadioMLPHead(dim=dim).to(device).eval()
569
+ head.load_state_dict(ckpt["head_state"], strict=True)
570
+ return head
571
+
572
+ self.raw_head = _load_head(RADIO_RAW_HEAD_PATH)
573
+ self.masked_head = _load_head(RADIO_MASKED_HEAD_PATH)
574
+
575
+ self.device_str = dev_str
576
+ self.loaded = True
577
+
578
+ RADIO_BUNDLE = RadioBundle()
579
+
580
+ def radio_heatmap_from_spatial(spatial_tokens: torch.Tensor, in_h: int, in_w: int, patch_size: int = 16) -> np.ndarray:
581
+ ht = in_h // patch_size
582
+ wt = in_w // patch_size
583
+ feat = rearrange(spatial_tokens, "b (h w) d -> b d h w", h=ht, w=wt)
584
+ energy = torch.sqrt(torch.clamp((feat ** 2).sum(dim=1), min=1e-8))[0]
585
+ energy = (energy - energy.min()) / (energy.max() - energy.min() + 1e-8)
586
+ hm = energy.detach().float().cpu().numpy().astype(np.float32)
587
+ hm_img = Image.fromarray((hm * 255).astype(np.uint8)).resize((in_w, in_h), resample=Image.BILINEAR)
588
+ return np.array(hm_img, dtype=np.float32) / 255.0
589
+
590
+ def radio_overlay_heatmap(rgb_u8: np.ndarray, heatmap01: np.ndarray, alpha: float = 0.35) -> np.ndarray:
591
+ img = rgb_u8.astype(np.float32) / 255.0
592
+ hm = np.clip(heatmap01, 0, 1).astype(np.float32)
593
+ out = img.copy()
594
+ out[..., 0] = np.clip(out[..., 0] * (1 - alpha) + hm * alpha, 0, 1)
595
+ return (out * 255).astype(np.uint8)
596
+
597
+ @torch.inference_mode()
598
+ def radio_predict_from_arrays(gray_vis_u8: np.ndarray,
599
+ lung_mask_u8: np.ndarray,
600
+ coverage: float,
601
+ device: torch.device,
602
+ gate_threshold: float) -> Dict[str, Any]:
603
+ RADIO_BUNDLE.load(device=device)
604
+ dtype = torch.float16 if device.type == "cuda" else torch.float32
605
+
606
+ # ---------- RAW ----------
607
+ raw_rgb = cv2.cvtColor(gray_vis_u8, cv2.COLOR_GRAY2RGB)
608
+ px = RADIO_BUNDLE.processor(
609
+ images=Image.fromarray(raw_rgb),
610
+ return_tensors="pt",
611
+ do_resize=True,
612
+ size={"shortest_edge": RADIO_IMG_SIZE},
613
+ do_center_crop=True,
614
+ ).pixel_values.to(device).to(dtype)
615
+
616
+ summary, spatial = RADIO_BUNDLE.radio(px)
617
+ logit_raw = RADIO_BUNDLE.raw_head(summary)
618
+ prob_raw = float(torch.sigmoid(logit_raw)[0].item())
619
+
620
+ hm_raw = radio_heatmap_from_spatial(spatial, px.shape[-2], px.shape[-1], RADIO_PATCH_SIZE)
621
+ raw_overlay = radio_overlay_heatmap(
622
+ cv2.resize(raw_rgb, (px.shape[-1], px.shape[-2])),
623
+ hm_raw,
624
+ alpha=0.35
625
+ )
626
+
627
+ # ---------- MASKED (optional) ----------
628
+ masked_prob = None
629
+ masked_overlay = None
630
+ masked_ran = False
631
+
632
+ if lung_mask_u8 is not None and coverage >= RADIO_MASKED_MIN_COV and coverage >= gate_threshold:
633
+ masked_ran = True
634
+ masked_u8 = (gray_vis_u8 * lung_mask_u8).astype(np.uint8)
635
+ masked_rgb = cv2.cvtColor(masked_u8, cv2.COLOR_GRAY2RGB)
636
+
637
+ pxm = RADIO_BUNDLE.processor(
638
+ images=Image.fromarray(masked_rgb),
639
+ return_tensors="pt",
640
+ do_resize=True,
641
+ size={"shortest_edge": RADIO_IMG_SIZE},
642
+ do_center_crop=True,
643
+ ).pixel_values.to(device).to(dtype)
644
+
645
+ summary_m, spatial_m = RADIO_BUNDLE.radio(pxm)
646
+ logit_m = RADIO_BUNDLE.masked_head(summary_m)
647
+ masked_prob = float(torch.sigmoid(logit_m)[0].item())
648
+
649
+ hm_m = radio_heatmap_from_spatial(spatial_m, pxm.shape[-2], pxm.shape[-1], RADIO_PATCH_SIZE)
650
+ masked_overlay = radio_overlay_heatmap(
651
+ cv2.resize(masked_rgb, (pxm.shape[-1], pxm.shape[-2])),
652
+ hm_m,
653
+ alpha=0.35
654
+ )
655
+
656
+ # ---------- PRIMARY = masked if available else raw ----------
657
+ prob_primary = masked_prob if masked_prob is not None else prob_raw
658
+
659
+ if prob_primary >= RADIO_THR_RED:
660
+ band = "RED"
661
+ pred = "LIKELY TB (RADIO)"
662
+ elif prob_primary >= RADIO_THR_SCREEN:
663
+ band = "YELLOW"
664
+ pred = "SCREEN-POSITIVE / INDETERMINATE (RADIO)"
665
+ else:
666
+ band = "GREEN"
667
+ pred = "LOW TB LIKELIHOOD (RADIO)"
668
+
669
+ return {
670
+ "prob_raw": prob_raw,
671
+ "prob_primary": prob_primary,
672
+ "pred": pred,
673
+ "band": band,
674
+ "raw_overlay": raw_overlay,
675
+ "masked_prob": masked_prob,
676
+ "masked_overlay": masked_overlay,
677
+ "masked_ran": masked_ran,
678
+ "gate_threshold": float(gate_threshold),
679
+ }
680
+
681
+
682
+ # ============================================================
683
+ # TB CORE ANALYSIS
684
+ # ============================================================
685
+ def analyze_one_image(
686
+ img_bgr: np.ndarray,
687
+ tb_weights: str,
688
+ lung_weights: str,
689
+ backbone: str,
690
+ threshold: float,
691
+ phone_mode: bool,
692
+ img_size: int = 224,
693
+ fail_cov: float = FAIL_COV,
694
+ warn_cov: float = WARN_COV,
695
+ ) -> Dict[str, Any]:
696
+
697
+ BUNDLE.load(tb_weights, lung_weights, backbone)
698
+ device = BUNDLE.device
699
+
700
+ gray = img_bgr if img_bgr.ndim == 2 else cv2.cvtColor(img_bgr, cv2.COLOR_BGR2GRAY)
701
+ q_score, q_warn = phone_quality_report(gray)
702
+
703
+ gray_vis = phone_preprocess(gray) if phone_mode else gray
704
+ if gray_vis.dtype != np.uint8:
705
+ gray_vis = np.clip(gray_vis, 0, 255).astype(np.uint8)
706
+
707
+ with torch.no_grad():
708
+ x_lung = preprocess_for_lung_unet(gray_vis).to(device)
709
+ mask_logits = BUNDLE.lung(x_lung)
710
+ mask256 = torch.sigmoid(mask_logits)[0, 0].cpu().numpy()
711
+
712
+ mask256_bin = (mask256 > 0.5).astype(np.uint8)
713
+
714
+ # post-process: keep 2 lungs, close, fill holes
715
+ mask256_bin = keep_top_k_components(mask256_bin, k=2)
716
+ k = max(3, int(0.02 * 256))
717
+ kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))
718
+ mask256_bin = cv2.morphologyEx(mask256_bin, cv2.MORPH_CLOSE, kernel, iterations=1)
719
+ mask256_bin = fill_holes(mask256_bin)
720
+
721
+ coverage = float(mask256_bin.mean())
722
+ mask_full = cv2.resize(mask256_bin, (gray_vis.shape[1], gray_vis.shape[0]), interpolation=cv2.INTER_NEAREST)
723
+
724
+ # fail-safe if coverage too low
725
+ if coverage < fail_cov:
726
+ overlay_rgb = cv2.cvtColor(cv2.resize(gray_vis, (img_size, img_size)), cv2.COLOR_GRAY2RGB)
727
+ return {
728
+ "prob": None,
729
+ "logit": None,
730
+ "pred": "INDETERMINATE",
731
+ "band": "YELLOW",
732
+ "band_text": "Lung segmentation failed. AI TB assessment cannot be performed reliably on this image.",
733
+ "quality_score": float(q_score),
734
+ "diffuse_risk": False,
735
+ "warnings": (
736
+ ["Lung segmentation failed (<10% lung area).", f"Lung coverage: {coverage*100:.1f}%"]
737
+ + (["Phone/WhatsApp mode enabled; artifacts possible."] if phone_mode else [])
738
+ + q_warn
739
+ ),
740
+ "lung_coverage": coverage,
741
+ "orig_gray": gray,
742
+ "vis_gray": gray_vis,
743
+ "masked_gray": None,
744
+ "proc_gray": None,
745
+ "lung_mask": mask_full,
746
+ "mask_overlay": make_mask_overlay(gray_vis, mask_full),
747
+ "overlay": overlay_rgb,
748
+ "overlay_clean": overlay_rgb,
749
+ }
750
+
751
+ # fail-safe if mask looks like single lung / cropped
752
+ sanity = mask_sanity_warnings(mask_full.astype(np.uint8))
753
+ if FAILSAFE_ON_BAD_MASK and sanity:
754
+ overlay_rgb = cv2.cvtColor(cv2.resize(gray_vis, (img_size, img_size)), cv2.COLOR_GRAY2RGB)
755
+ return {
756
+ "prob": None,
757
+ "logit": None,
758
+ "pred": "INDETERMINATE",
759
+ "band": "YELLOW",
760
+ "band_text": "Non-standard/cropped view or unreliable lung segmentation. TB scoring disabled (fail-safe).",
761
+ "quality_score": float(q_score),
762
+ "diffuse_risk": False,
763
+ "warnings": (
764
+ sanity
765
+ + [f"Lung coverage: {coverage*100:.1f}%"]
766
+ + (["Phone/WhatsApp mode enabled; artifacts possible."] if phone_mode else [])
767
+ + q_warn
768
+ ),
769
+ "lung_coverage": coverage,
770
+ "orig_gray": gray,
771
+ "vis_gray": gray_vis,
772
+ "masked_gray": None,
773
+ "proc_gray": None,
774
+ "lung_mask": mask_full,
775
+ "mask_overlay": make_mask_overlay(gray_vis, mask_full),
776
+ "overlay": overlay_rgb,
777
+ "overlay_clean": overlay_rgb,
778
+ }
779
+
780
+ masked = (gray_vis * mask_full).astype(np.uint8)
781
+ masked_f01 = tb_training_preprocess(masked)
782
+ masked_u8 = (masked_f01 * 255).astype(np.uint8)
783
+
784
+ masked_u8_rs = cv2.resize(masked_u8, (img_size, img_size), interpolation=cv2.INTER_AREA)
785
+ rgb = cv2.cvtColor(masked_u8_rs, cv2.COLOR_GRAY2RGB)
786
+ x = BUNDLE.tfm(rgb).unsqueeze(0).to(device)
787
+
788
+ cam, prob_tb, logit = BUNDLE.cammer.generate(x)
789
+
790
+ cam_u8 = (np.clip(cam, 0, 1) * 255).astype(np.uint8)
791
+ cam_u8 = cv2.resize(cam_u8, (img_size, img_size), interpolation=cv2.INTER_CUBIC)
792
+ cam_up = cam_u8.astype(np.float32) / 255.0
793
+
794
+ diffuse = detect_diffuse_risk(prob_tb, cam_up, q_score)
795
+ band_base, _ = confidence_band(prob_tb, q_score, diffuse)
796
+
797
+ allow_red = (prob_tb >= 0.70 and q_score >= 55 and not diffuse and coverage >= warn_cov)
798
+ band = "RED" if allow_red else band_base
799
+
800
+ pred = REPORT_LABELS[band]["title"]
801
+ band_text = REPORT_LABELS[band]["summary"]
802
+
803
+ abnormal_non_tb = (prob_tb >= 0.60 and q_score < 55 and band != "RED")
804
+ if abnormal_non_tb:
805
+ band_text = (
806
+ "Significant abnormal lung findings detected. "
807
+ "Findings are non-specific and not characteristic of pulmonary tuberculosis. "
808
+ "Image quality may affect AI reliability."
809
+ )
810
+
811
+ heat = cv2.applyColorMap((cam_up * 255).astype(np.uint8), cv2.COLORMAP_JET)
812
+ overlay_clean = cv2.addWeighted(rgb, 0.65, heat, 0.35, 0)
813
+
814
+ overlay_annotated = overlay_clean.copy()
815
+ text1 = f"{band}: {pred}"
816
+ text2 = f"TB prob={prob_tb:.3f} | Quality={q_score:.0f}/100 | Lung coverage={coverage*100:.1f}%"
817
+ cv2.putText(overlay_annotated, text1, (8, 20), cv2.FONT_HERSHEY_SIMPLEX, 0.52, (255, 255, 255), 2)
818
+ cv2.putText(overlay_annotated, text1, (8, 20), cv2.FONT_HERSHEY_SIMPLEX, 0.52, (0, 0, 0), 1)
819
+ cv2.putText(overlay_annotated, text2, (8, 42), cv2.FONT_HERSHEY_SIMPLEX, 0.50, (255, 255, 255), 2)
820
+ cv2.putText(overlay_annotated, text2, (8, 42), cv2.FONT_HERSHEY_SIMPLEX, 0.50, (0, 0, 0), 1)
821
+
822
+ warnings = []
823
+ if phone_mode: warnings.append("Phone/WhatsApp mode enabled; artifacts possible.")
824
+ if q_score < 55: warnings.append("Suboptimal image quality limits AI reliability.")
825
+ if coverage < warn_cov: warnings.append(f"Partial lung segmentation ({coverage*100:.1f}% coverage).")
826
+ if diffuse: warnings.append("Diffuse, non-focal AI attention pattern; TB-specific features not identified.")
827
+ if abnormal_non_tb: warnings.append("Abnormal lung findings detected; pattern not specific for tuberculosis.")
828
+ warnings.extend(q_warn)
829
+
830
+ return {
831
+ "prob": float(prob_tb),
832
+ "logit": float(logit),
833
+ "pred": pred,
834
+ "band": band,
835
+ "band_text": band_text,
836
+ "quality_score": float(q_score),
837
+ "diffuse_risk": bool(diffuse),
838
+ "warnings": warnings,
839
+ "lung_coverage": coverage,
840
+ "orig_gray": gray,
841
+ "vis_gray": gray_vis,
842
+ "masked_gray": masked,
843
+ "proc_gray": masked_u8_rs,
844
+ "lung_mask": mask_full,
845
+ "mask_overlay": make_mask_overlay(gray_vis, mask_full),
846
+ "overlay": overlay_annotated,
847
+ "overlay_clean": overlay_clean,
848
+ }
849
+
850
+
851
+ # ============================================================
852
+ # GRADIO CALLBACK
853
+ # ============================================================
854
+ def run_analysis(
855
+ files: List[gr.File],
856
+ tb_weights: str,
857
+ lung_weights: str,
858
+ backbone: str,
859
+ threshold: float,
860
+ phone_mode: bool,
861
+ use_radio: bool,
862
+ radio_gate: float,
863
+ ):
864
+ if not files:
865
+ return [], None, None, None, "Please upload at least one image."
866
+
867
+ if not os.path.exists(tb_weights):
868
+ return [], None, None, None, f"TB weights not found: {tb_weights}"
869
+ if not os.path.exists(lung_weights):
870
+ return [], None, None, None, f"Lung U-Net weights not found: {lung_weights}"
871
+
872
+ rows = []
873
+ gallery_items = []
874
+ details_md = []
875
+
876
+ for f in files:
877
+ path = f.name if hasattr(f, "name") else str(f)
878
+ name = os.path.basename(path)
879
+
880
+ img = cv2.imread(path, cv2.IMREAD_COLOR)
881
+ if img is None:
882
+ rows.append([name, "", "SKIP", "", "Unreadable image", "", "", "", "", ""])
883
+ continue
884
+
885
+ out = analyze_one_image(
886
+ img_bgr=img,
887
+ tb_weights=tb_weights,
888
+ lung_weights=lung_weights,
889
+ backbone=backbone,
890
+ threshold=threshold,
891
+ phone_mode=phone_mode,
892
+ img_size=224,
893
+ )
894
+
895
+ # -------------------------
896
+ # RADIO (optional)
897
+ # -------------------------
898
+ radio_text = "RADIO disabled."
899
+ radio_raw_overlay = None
900
+ radio_masked_overlay = None
901
+ radio_raw_val: Optional[float] = None
902
+ radio_masked_val: Optional[float] = None
903
+ radio_band: Optional[str] = None
904
+
905
+ radio_raw_str = ""
906
+ radio_masked_str = ""
907
+
908
+ if use_radio and out["prob"] is not None:
909
+ try:
910
+ r = radio_predict_from_arrays(
911
+ gray_vis_u8=out["vis_gray"],
912
+ lung_mask_u8=out["lung_mask"].astype(np.uint8),
913
+ coverage=float(out["lung_coverage"]),
914
+ device=BUNDLE.device,
915
+ gate_threshold=float(radio_gate),
916
+ )
917
+
918
+ radio_raw_val = float(r["prob_raw"])
919
+ radio_masked_val = None if r["masked_prob"] is None else float(r["masked_prob"])
920
+ radio_band = str(r["band"])
921
+
922
+ radio_raw_str = f"{radio_raw_val:.4f}"
923
+ radio_masked_str = "" if radio_masked_val is None else f"{radio_masked_val:.4f}"
924
+
925
+ radio_text = (
926
+ f"**RADIO:** {r['pred']} | RAW={radio_raw_val:.4f}"
927
+ + (f" | MASKED={radio_masked_val:.4f}" if radio_masked_val is not None else "")
928
+ + f" | Band={radio_band}"
929
+ )
930
+ radio_raw_overlay = r["raw_overlay"]
931
+ radio_masked_overlay = r["masked_overlay"]
932
+ except Exception as e:
933
+ radio_text = f"RADIO error: {type(e).__name__}: {e}"
934
+ radio_raw_str = ""
935
+ radio_masked_str = ""
936
+ radio_raw_val = None
937
+ radio_masked_val = None
938
+ radio_band = None
939
+
940
+ # -------------------------
941
+ # Consensus
942
+ # -------------------------
943
+ consensus_label, consensus_detail = build_consensus(
944
+ tb_prob=out["prob"],
945
+ tb_band=out["band"],
946
+ radio_raw=radio_raw_val,
947
+ radio_masked=radio_masked_val,
948
+ radio_band=radio_band,
949
+ )
950
+
951
+ # -------------------------
952
+ # Table row
953
+ # -------------------------
954
+ prob_str = "" if out["prob"] is None else f"{out['prob']:.4f}"
955
+ cov_str = f"{out.get('lung_coverage', 0.0) * 100:.1f}%"
956
+
957
+ rows.append([
958
+ name,
959
+ prob_str,
960
+ out["pred"],
961
+ out["band"],
962
+ out["band_text"],
963
+ f"{out['quality_score']:.0f}",
964
+ cov_str,
965
+ radio_raw_str,
966
+ radio_masked_str,
967
+ consensus_label,
968
+ ])
969
+
970
+ # -------------------------
971
+ # Visual outputs
972
+ # -------------------------
973
+ orig_rgb = cv2.cvtColor(cv2.resize(out["orig_gray"], (512, 512)), cv2.COLOR_GRAY2RGB)
974
+ vis_rgb = cv2.cvtColor(cv2.resize(out["vis_gray"], (512, 512)), cv2.COLOR_GRAY2RGB)
975
+ mask_overlay = cv2.resize(out["mask_overlay"], (512, 512))
976
+ overlay_big = cv2.resize(out["overlay"], (512, 512))
977
+
978
+ gallery_items.append((orig_rgb, f"{name} • ORIGINAL"))
979
+ gallery_items.append((vis_rgb, f"{name} • PHONE-PROC" if phone_mode else f"{name} • INPUT"))
980
+ gallery_items.append((mask_overlay, f"{name} • Lung mask overlay"))
981
+
982
+ if out["proc_gray"] is not None:
983
+ proc_rgb = cv2.cvtColor(cv2.resize(out["proc_gray"], (512, 512)), cv2.COLOR_GRAY2RGB)
984
+ gallery_items.append((proc_rgb, f"{name} • Masked model input (224x224)"))
985
+
986
+ gallery_items.append((overlay_big, f"{name} • Grad-CAM overlay (TBNet)"))
987
+
988
+ if radio_raw_overlay is not None:
989
+ gallery_items.append((cv2.resize(radio_raw_overlay, (512, 512)), f"{name} • RADIO RAW heatmap"))
990
+ if radio_masked_overlay is not None:
991
+ gallery_items.append((cv2.resize(radio_masked_overlay, (512, 512)), f"{name} • RADIO MASKED heatmap"))
992
+
993
+ # -------------------------
994
+ # Details panel
995
+ # -------------------------
996
+ warn_txt = "\n".join([f"- {w}" for w in out["warnings"]]) if out["warnings"] else "- None"
997
+ tb_line = "N/A (segmentation failed / fail-safe)" if out["prob"] is None else f"{out['prob']:.4f}"
998
+ rec_line = recommendation_for_band(out.get("band"))
999
+
1000
+ details_md.append(
1001
+ f"""### {name}
1002
+
1003
+ **AI Assessment (TBNet):** **{out['pred']}**
1004
+ {rec_line}
1005
+
1006
+ **TB Probability (screening model):** {tb_line}
1007
+
1008
+ **Interpretation**
1009
+ {out['band_text']}
1010
+
1011
+ **Image Quality:** {out['quality_score']:.0f}/100
1012
+ **Lung Mask Coverage:** {out.get('lung_coverage', 0.0) * 100:.1f}%
1013
+ **AI Attention Pattern (TBNet):** {"Diffuse / non-focal" if out["diffuse_risk"] else "Focal / localized"}
1014
+
1015
+ **Warnings**
1016
+ {warn_txt}
1017
+
1018
+ **RADIO Output**
1019
+ {radio_text}
1020
+
1021
+ **Final consensus (TBNet vs RADIO):** **{consensus_label}**
1022
+ - {consensus_detail}
1023
+
1024
+ **Clinical Guidance**
1025
+ {CLINICAL_GUIDANCE}
1026
+
1027
+ ---
1028
+ """
1029
+ )
1030
+
1031
+ return rows, gallery_items, "\n".join(details_md), CLINICAL_DISCLAIMER, "Done."
1032
+
1033
+
1034
+ # ============================================================
1035
+ # UI
1036
+ # ============================================================
1037
+ def build_ui():
1038
+ css = """
1039
+ .title {font-size: 28px; font-weight: 800; margin-bottom: 6px;}
1040
+ .subtitle {font-size: 14px; opacity: 0.85; margin-bottom: 14px;}
1041
+ .warnbox {border-left: 6px solid #f59e0b; padding: 10px 12px; background: rgba(245,158,11,0.08); border-radius: 10px;}
1042
+ """
1043
+
1044
+ with gr.Blocks(title="TB X-ray Assistant (TBNet + RADIO)", css=css) as demo:
1045
+ gr.Markdown('<div class="title">TB X-ray Assistant (Auto Lung Mask • Research Use)</div>')
1046
+ gr.Markdown('<div class="subtitle">Lung U-Net masking → EfficientNet TBNet + Grad-CAM • Optional RADIO (C-RADIOv4 + heads) • 3-state consensus</div>')
1047
+
1048
+ with gr.Row():
1049
+ with gr.Column(scale=1):
1050
+ gr.Markdown("#### Model settings")
1051
+
1052
+ tb_weights = gr.Textbox(label="TB Weights (.pt)", value=DEFAULT_TB_WEIGHTS)
1053
+ lung_weights = gr.Textbox(label="Lung U-Net Weights (.pt)", value=DEFAULT_LUNG_WEIGHTS)
1054
+
1055
+ backbone = gr.Dropdown(choices=["efficientnet_b0"], value="efficientnet_b0", label="Backbone")
1056
+
1057
+ threshold = gr.Slider(0.01, 0.99, value=TBNET_SCREEN_THR, step=0.01,
1058
+ label=f"Reference threshold (TBNet screen+) = {TBNET_SCREEN_THR:.2f}")
1059
+
1060
+ phone_mode = gr.Checkbox(value=False,
1061
+ label="Phone/WhatsApp Mode (SAFE: conditional crop + conditional CLAHE)")
1062
+
1063
+ # RADIO
1064
+ use_radio = gr.Checkbox(value=False, label="Enable RADIO layer (C-RADIOv4 + heads)")
1065
+ radio_gate = gr.Slider(0.10, 0.40, value=RADIO_GATE_DEFAULT, step=0.01,
1066
+ label="RADIO masked gate (run masked head if lung coverage ≥ gate)")
1067
+
1068
+ gr.Markdown(
1069
+ '<div class="warnbox"><b>Fail-safe:</b> If lung segmentation is too small or looks like a cropped/single-lung mask, TB scoring is disabled to avoid false positives.</div>'
1070
+ )
1071
+
1072
+ gr.Markdown(
1073
+ f"<div class='subtitle'>Device for TB+RADIO: <b>{DEVICE}</b> (set FORCE_CPU=True to force CPU)</div>"
1074
+ )
1075
+
1076
+ with gr.Column(scale=2):
1077
+ gr.Markdown("#### Upload images")
1078
+ files = gr.Files(label="Upload one or multiple X-ray images", file_types=[".png", ".jpg", ".jpeg", ".bmp"])
1079
+ run_btn = gr.Button("Run Analysis", variant="primary")
1080
+ status = gr.Textbox(label="Status", value="Ready.", interactive=False)
1081
+
1082
+ with gr.Row():
1083
+ table = gr.Dataframe(
1084
+ headers=[
1085
+ "Image",
1086
+ "TB Probability",
1087
+ "AI Assessment",
1088
+ "Band",
1089
+ "Band meaning",
1090
+ "Quality",
1091
+ "LungCov",
1092
+ "RADIO RAW",
1093
+ "RADIO MASKED",
1094
+ "CONSENSUS",
1095
+ ],
1096
+ datatype=["str","str","str","str","str","str","str","str","str","str"],
1097
+ interactive=False,
1098
+ label="Results"
1099
+ )
1100
+
1101
+ with gr.Row():
1102
+ gallery = gr.Gallery(
1103
+ label="Visual outputs (Original • Input • Mask • Masked • Grad-CAM • RADIO)",
1104
+ columns=3,
1105
+ height=560
1106
+ )
1107
+
1108
+ with gr.Row():
1109
+ with gr.Column(scale=1):
1110
+ disclaimer_box = gr.Markdown(CLINICAL_DISCLAIMER)
1111
+ with gr.Column(scale=2):
1112
+ details = gr.Markdown("")
1113
+
1114
+ run_btn.click(
1115
+ fn=run_analysis,
1116
+ inputs=[
1117
+ files,
1118
+ tb_weights,
1119
+ lung_weights,
1120
+ backbone,
1121
+ threshold,
1122
+ phone_mode,
1123
+ use_radio,
1124
+ radio_gate,
1125
+ ],
1126
+ outputs=[table, gallery, details, disclaimer_box, status]
1127
+ )
1128
+
1129
+ return demo
1130
+
1131
+
1132
+ if __name__ == "__main__":
1133
+ demo = build_ui()
1134
+ demo.launch(server_name="0.0.0.0", server_port=7860, show_error=True)