File size: 16,103 Bytes
005bed0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8b19a38
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
005bed0
 
 
 
8b19a38
005bed0
 
 
 
 
 
 
 
8b19a38
005bed0
 
 
 
 
 
 
 
8b19a38
005bed0
 
 
 
 
 
 
 
8b19a38
005bed0
 
 
 
 
 
 
 
8b19a38
005bed0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
398190f
 
 
 
 
 
 
 
005bed0
 
 
 
 
 
 
 
 
 
398190f
005bed0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
import os
os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"
os.environ["OMP_NUM_THREADS"] = "1"
import sys
import torch
torch.set_num_threads(1)


import torch.nn.functional as F
import numpy as np
try:
    import gradio as gr
    HAS_GRADIO = True
except ImportError:
    gr = None
    HAS_GRADIO = False

from PIL import Image, ImageDraw, ImageFont
from pathlib import Path

os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"

# Add repository root and model subdirectories to path
WORKSPACE_ROOT = Path(__file__).resolve().parent.parent
if str(WORKSPACE_ROOT) not in sys.path:
    sys.path.insert(0, str(WORKSPACE_ROOT))
m2_path = WORKSPACE_ROOT / "models_suite/model2_choroidalyzer"
if str(m2_path) not in sys.path:
    sys.path.insert(0, str(m2_path))


# Import models from models_suite
from models_suite.model1_oct5k_layers.unet_layers import RetinalLayersUNet
from models_suite.model2_choroidalyzer.choroidalyze.model import UNet as ChoroidalyzerUNet

from models_suite.model3_hrf_dme.hrf_aunet import HRFAttentionUNet as HRF_AttentionUNet

from models_suite.model4_oimhs_hole_cysts.oimhs_unet import OIMHSUNet
from models_suite.model5_oct5k_detection.detector import OCTPathologyDetector, OCT5K_DETECTION_CLASSES

try:
    import spaces
    IS_HF_SPACE = True
except ImportError:
    spaces = None
    IS_HF_SPACE = False

# Device configuration:
# Local runs: use Mac GPU (mps/cuda) if available
# Deployed HF Space: ConvNeXtV2 uses @spaces.GPU for ZeroGPU, Segmentation models use CPU only
if IS_HF_SPACE:
    SEGMENT_DEVICE = torch.device("cpu")
elif torch.cuda.is_available():
    SEGMENT_DEVICE = torch.device("cuda")
elif torch.backends.mps.is_available():
    SEGMENT_DEVICE = torch.device("mps")
else:
    SEGMENT_DEVICE = torch.device("cpu")

DEVICE = SEGMENT_DEVICE
print(f"Initializing OCT Analyzer Suite | Deployed HF Space: {IS_HF_SPACE} | Segmentation Device: {SEGMENT_DEVICE}")

try:
    import numpy._core.multiarray as np_multiarray
    torch.serialization.add_safe_globals([np_multiarray.scalar])
except Exception:
    pass

def safe_load_ckpt(cp_path, device):
    try:
        return torch.load(cp_path, map_location=device, weights_only=True)
    except Exception:
        return torch.load(cp_path, map_location=device, weights_only=False)




def get_checkpoint_file(local_rel_path: str, bucket_key: str) -> Path:
    target_path = WORKSPACE_ROOT / local_rel_path
    if target_path.exists():
        return target_path
    
    target_path.parent.mkdir(parents=True, exist_ok=True)
    token = os.getenv("HF_TOKEN")
    
    try:
        from huggingface_hub import hf_hub_download
        print(f"Downloading weight from HF Bucket: segmentation/{bucket_key}...", flush=True)
        cached_path = hf_hub_download(
            repo_id="NMundhra/OCT-Image-Classifier-Model-storage",
            filename=f"segmentation/{bucket_key}",
            repo_type="dataset",
            token=token
        )
        import shutil
        shutil.copy(cached_path, target_path)
        print(f"Successfully cached {bucket_key} to {target_path}", flush=True)
    except Exception as e:
        print(f"Warning: Could not download segmentation/{bucket_key} from bucket: {e}", flush=True)
        
    return target_path


# Initialize model instances and load checkpoints
def load_suite():
    print("Loading M1...", flush=True)
    m1 = RetinalLayersUNet(in_channels=1, num_classes=6)
    cp1 = get_checkpoint_file("models_suite/model1_oct5k_layers/checkpoints/best_model.pth", "model1_oct5k_layers.pth")
    if cp1.exists():
        ckpt = safe_load_ckpt(cp1, DEVICE)
        m1.load_state_dict(ckpt["model_state_dict"] if "model_state_dict" in ckpt else ckpt)
    m1.to(DEVICE).eval()
    print("M1 Done.", flush=True)

    print("Loading M2...", flush=True)
    m2 = ChoroidalyzerUNet(in_channels=1, out_channels=3, depth=7, channels='8_doublemax-64', up_type='conv_then_interpolate', extra_out_conv=True)
    cp2 = get_checkpoint_file("models_suite/model2_choroidalyzer/checkpoints/best_model.pth", "model2_choroidalyzer.pth")
    if cp2.exists():
        ckpt = safe_load_ckpt(cp2, DEVICE)
        m2.load_state_dict(ckpt["model_state_dict"] if "model_state_dict" in ckpt else ckpt)
    m2.to(DEVICE).eval()
    print("M2 Done.", flush=True)

    print("Loading M3...", flush=True)
    m3 = HRF_AttentionUNet(n_channels=3, n_classes=1)
    cp3 = get_checkpoint_file("models_suite/model3_hrf_dme/checkpoints/best_model.pth", "model3_hrf_dme.pth")
    if cp3.exists():
        ckpt = safe_load_ckpt(cp3, DEVICE)
        m3.load_state_dict(ckpt["model_state_dict"] if "model_state_dict" in ckpt else ckpt)
    m3.to(DEVICE).eval()
    print("M3 Done.", flush=True)

    print("Loading M4...", flush=True)
    m4 = OIMHSUNet(in_channels=1, num_classes=5)
    cp4 = get_checkpoint_file("models_suite/model4_oimhs_hole_cysts/checkpoints/best_model.pth", "model4_oimhs_hole_cysts.pth")
    if cp4.exists():
        ckpt = safe_load_ckpt(cp4, DEVICE)
        m4.load_state_dict(ckpt["model_state_dict"] if "model_state_dict" in ckpt else ckpt)
    m4.to(DEVICE).eval()
    print("M4 Done.", flush=True)

    print("Loading M5...", flush=True)
    m5 = OCTPathologyDetector(num_classes=10)
    cp5 = get_checkpoint_file("models_suite/model5_oct5k_detection/checkpoints/best_model.pth", "model5_oct5k_detection.pth")
    if cp5.exists():
        ckpt = safe_load_ckpt(cp5, DEVICE)
        m5.load_state_dict(ckpt["model_state_dict"] if "model_state_dict" in ckpt else ckpt)
    m5.to(DEVICE).eval()
    print("M5 Done.", flush=True)

    return m1, m2, m3, m4, m5




model1, model2, model3, model4, model5 = load_suite()
print("OCT Analyser 5-Model Suite Loaded Successfully!", flush=True)


# Preprocessing helpers
def preprocess_image(image: Image.Image, target_size=(256, 256)):
    gray = image.convert("L").resize(target_size)
    arr = np.array(gray, dtype=np.float32) / 255.0
    tensor = torch.from_numpy(arr).unsqueeze(0).unsqueeze(0) # (1, 1, H, W)
    return gray, tensor

def create_pure_mask(orig_img: Image.Image, mask: np.ndarray, num_classes: int, cmap: list) -> Image.Image:
    h, w = mask.shape
    mask_rgba = np.zeros((h, w, 4), dtype=np.uint8)
    for c in range(1, num_classes):
        color = cmap[c % len(cmap)]
        mask_rgba[mask == c] = [color[0], color[1], color[2], 255]
    return Image.fromarray(mask_rgba, mode="RGBA").resize(orig_img.size, Image.NEAREST)

def create_segmentation_overlay(orig_img: Image.Image, mask: np.ndarray, num_classes: int, cmap: list):
    orig_rgb = orig_img.convert("RGB")
    h, w = mask.shape
    overlay = np.zeros((h, w, 3), dtype=np.uint8)
    for c in range(1, num_classes):
        color = cmap[c % len(cmap)]
        overlay[mask == c] = color

    overlay_img = Image.fromarray(overlay).resize(orig_rgb.size)
    blended = Image.blend(orig_rgb, overlay_img, alpha=0.45)
    return blended, Image.fromarray(overlay).resize(orig_rgb.size)

def preprocess_rgb_image(image: Image.Image, target_size=(256, 256)):
    rgb = image.convert("RGB").resize(target_size)
    arr = np.array(rgb, dtype=np.float32).transpose(2, 0, 1) / 255.0  # (3, H, W)
    tensor = torch.from_numpy(arr).unsqueeze(0)  # (1, 3, H, W)
    return rgb, tensor

COLOR_MAP = [
    [0, 0, 0],       # 0: BG
    [255, 50, 50],   # 1: Red
    [50, 255, 50],   # 2: Green
    [50, 50, 255],   # 3: Blue
    [255, 255, 50],  # 4: Yellow
    [255, 50, 255],  # 5: Magenta
    [50, 255, 255]   # 6: Cyan
]

# Inference handlers for each model
def predict_model1(image):
    if image is None: return None, "Please upload an OCT scan image."
    gray, tensor = preprocess_image(image, (256, 256))
    with torch.no_grad():
        logits = model1(tensor.to(DEVICE))
        preds = torch.argmax(logits, dim=1).cpu().numpy()[0]
    
    overlay = create_segmentation_overlay(image, preds, 6, COLOR_MAP)
    layer_names = ["Background", "ILM -> OPL", "OPL -> IS-OS", "IS-OS -> IBRPE", "IBRPE -> OBRPE", "Choroid & Below"]
    counts = {layer_names[c]: int((preds == c).sum()) for c in range(6)}
    info = f"Retinal Layer Segmentation Complete.\nPixel Breakdown:\n" + "\n".join([f"• {k}: {v} px" for k, v in counts.items()])
    return overlay, info

def predict_model2(image):
    if image is None: return None, "Please upload an OCT scan image."
    gray, tensor = preprocess_image(image, (256, 256))
    with torch.no_grad():
        logits = model2(tensor.to(DEVICE))
        probs = torch.sigmoid(logits).cpu().numpy()[0, 0]
        mask = (probs > 0.5).astype(np.uint8)
    
    overlay = create_segmentation_overlay(image, mask, 2, [[0,0,0], [0, 220, 255]])
    choroid_area = int(mask.sum())
    mean_thickness = float(mask.sum(axis=0).mean())
    info = f"Choroid Region Analysis:\n• Choroid Area: {choroid_area} px\n• Est. Mean Thickness: {mean_thickness:.2f} px"
    return overlay, info

def predict_model3(image):
    if image is None: return None, "Please upload an OCT scan image."
    rgb, tensor = preprocess_rgb_image(image, (256, 256))
    with torch.no_grad():
        logits = model3(tensor.to(DEVICE))
        probs = torch.sigmoid(logits).cpu().numpy()[0, 0]
        mask = (probs > 0.5).astype(np.uint8)
    
    overlay = create_segmentation_overlay(image, mask, 2, [[0,0,0], [255, 0, 120]])
    fluid_px = int(mask.sum())
    info = f"High-Resolution HRF DME Fluid & Lesion Attention Analysis:\n• Pathological Fluid / Lesion Region: {fluid_px} px"
    return overlay, info


def predict_model4(image):
    if image is None: return None, "Please upload an OCT scan image."
    gray, tensor = preprocess_image(image, (256, 256))
    with torch.no_grad():
        logits = model4(tensor.to(DEVICE))
        preds = torch.argmax(logits, dim=1).cpu().numpy()[0]
    
    overlay = create_segmentation_overlay(image, preds, 5, COLOR_MAP)
    classes = ["Background", "Macular Hole", "Choroid", "Retina", "Intraretinal Cysts (IRC)"]
    counts = {classes[c]: int((preds == c).sum()) for c in range(5)}
    info = "OIMHS Pathology Analysis:\n" + "\n".join([f"• {k}: {v} px" for k, v in counts.items()])
    return overlay, info

def predict_model5(image, score_threshold=0.5):
    if image is None: return None, "Please upload an OCT scan image."
    orig_rgb = image.convert("RGB")
    gray, tensor = preprocess_image(image, (256, 256))
    # Faster R-CNN expects 3-channel input
    input_tensor = tensor.repeat(1, 3, 1, 1).to(DEVICE)
    
    with torch.no_grad():
        outputs = model5([input_tensor[0]])[0]

    
    boxes = outputs["boxes"].cpu().numpy()
    labels = outputs["labels"].cpu().numpy()
    scores = outputs["scores"].cpu().numpy()

    # Scale boxes back to original image size
    orig_w, orig_h = orig_rgb.size
    scale_x = orig_w / 256.0
    scale_y = orig_h / 256.0

    draw_img = orig_rgb.copy()
    draw = ImageDraw.Draw(draw_img)

    detected_items = []
    for box, label, score in zip(boxes, labels, scores):
        if score >= score_threshold:
            x1, y1, x2, y2 = box[0] * scale_x, box[1] * scale_y, box[2] * scale_x, box[3] * scale_y
            cls_name = OCT5K_DETECTION_CLASSES[label] if label < len(OCT5K_DETECTION_CLASSES) else f"Class {label}"
            draw.rectangle([x1, y1, x2, y2], outline="red", width=3)
            draw.text((x1 + 4, max(0, y1 - 12)), f"{cls_name} {score:.2f}", fill="yellow")
            detected_items.append(f"• {cls_name}: Conf {score:.2%} at [{int(x1)}, {int(y1)}, {int(x2)}, {int(y2)}]")

    info = f"OCT Pathology Object Detector ({len(detected_items)} objects detected above threshold {score_threshold:.2f}):\n"
    if detected_items:
        info += "\n".join(detected_items)
    else:
        info += "No objects detected above threshold."
        
    return draw_img, info

# Build Gradio UI for HF Space
if HAS_GRADIO:
    with gr.Blocks(title="OCT Analyser 5-Model Microservice Suite") as demo:
        gr.Markdown("# 👁️ OCT Analyser 5-Model Suite (Hugging Face API Services Deployment)")
        gr.Markdown("Comprehensive Optical Coherence Tomography (OCT) Deep Learning Suite providing 5 API services for segmentation, choroid analysis, fluid/lesion quantification, macular hole/cyst detection, and 9-class biomarker object detection.")

        with gr.Tabs():
            with gr.TabItem("Model 1: Retinal Layers U-Net"):
                gr.Markdown("### 6-Class Retinal Layer Segmentation U-Net (OCT5K Benchmark)")
                with gr.Row():
                    with gr.Column():
                        img1 = gr.Image(type="pil", label="Input OCT Scan")
                        btn1 = gr.Button("Segment Retinal Layers", variant="primary")
                    with gr.Column():
                        out_img1 = gr.Image(type="pil", label="6-Layer Segmentation Overlay")
                        txt1 = gr.Textbox(label="Layer Metrics", lines=7)
                btn1.click(predict_model1, inputs=img1, outputs=[out_img1, txt1], api_name="predict_model1")

            with gr.TabItem("Model 2: Choroidalyzer U-Net"):
                gr.Markdown("### Choroid Region & Thickness Quantification U-Net")
                with gr.Row():
                    with gr.Column():
                        img2 = gr.Image(type="pil", label="Input OCT Scan")
                        btn2 = gr.Button("Analyze Choroid Region", variant="primary")
                    with gr.Column():
                        out_img2 = gr.Image(type="pil", label="Choroid Mask Overlay")
                        txt2 = gr.Textbox(label="Choroid Metrics", lines=5)
                btn2.click(predict_model2, inputs=img2, outputs=[out_img2, txt2], api_name="predict_model2")

            with gr.TabItem("Model 3: HRF Attention U-Net"):
                gr.Markdown("### High-Resolution Fluid & Lesion Attention U-Net (HRF DME/AMD)")
                with gr.Row():
                    with gr.Column():
                        img3 = gr.Image(type="pil", label="Input OCT Scan")
                        btn3 = gr.Button("Segment Fluid & Lesions", variant="primary")
                    with gr.Column():
                        out_img3 = gr.Image(type="pil", label="Fluid & Lesion Mask Overlay")
                        txt3 = gr.Textbox(label="Fluid & Lesion Metrics", lines=5)
                btn3.click(predict_model3, inputs=img3, outputs=[out_img3, txt3], api_name="predict_model3")

            with gr.TabItem("Model 4: OIMHS Hole & Cyst U-Net"):
                gr.Markdown("### Macular Hole & Intraretinal Cyst (IRC) U-Net (OIMHS)")
                with gr.Row():
                    with gr.Column():
                        img4 = gr.Image(type="pil", label="Input OCT Scan")
                        btn4 = gr.Button("Detect Hole & Cysts", variant="primary")
                    with gr.Column():
                        out_img4 = gr.Image(type="pil", label="Hole & Cyst Segmentation Overlay")
                        txt4 = gr.Textbox(label="Pathology Metrics", lines=6)
                btn4.click(predict_model4, inputs=img4, outputs=[out_img4, txt4], api_name="predict_model4")

            with gr.TabItem("Model 5: OCT Pathology Detector"):
                gr.Markdown("### Faster R-CNN 9-Class Biomarker Object Detector")
                with gr.Row():
                    with gr.Column():
                        img5 = gr.Image(type="pil", label="Input OCT Scan")
                        thresh = gr.Slider(minimum=0.1, maximum=0.9, value=0.5, step=0.05, label="Confidence Threshold")
                        btn5 = gr.Button("Detect Biomarker Bounding Boxes", variant="primary")
                    with gr.Column():
                        out_img5 = gr.Image(type="pil", label="Biomarker Bounding Box Detections")
                        txt5 = gr.Textbox(label="Detection Results", lines=8)
                btn5.click(predict_model5, inputs=[img5, thresh], outputs=[out_img5, txt5], api_name="predict_model5")

if __name__ == "__main__":
    if HAS_GRADIO:
        demo.launch(server_name="0.0.0.0", server_port=7860)
    else:
        print("Gradio is not installed locally. Run pip install gradio to launch the server web interface.")