File size: 6,788 Bytes
de1f81a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9e244b4
de1f81a
9e244b4
 
 
 
de1f81a
 
 
9e244b4
 
 
 
 
 
 
 
 
 
 
de1f81a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9e244b4
de1f81a
 
 
 
 
 
 
 
 
 
 
 
 
 
9e244b4
 
 
 
 
de1f81a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9e244b4
de1f81a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9e244b4
de1f81a
 
 
 
9e244b4
 
de1f81a
 
 
 
 
 
9e244b4
 
de1f81a
9e244b4
de1f81a
 
9e244b4
de1f81a
 
9e244b4
 
 
 
 
 
 
 
de1f81a
 
 
 
 
 
 
9e244b4
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
import os
import numpy as np
import tensorflow as tf
import gradio as gr
from PIL import Image

# =========================
# CONFIG
# =========================
MODEL_PATH = "best_model.keras"

PATCH_SIZE = 256
INFER_STRIDE = 192
BATCH_SIZE = 8

USE_GRAYSCALE = False
CHANNELS = 3

TITLE = "Research Demo: Urine Sediment Virtual Phase Contrast Conversion"
DESCRIPTION = (
    "This Space demonstrates a research prototype that converts a brightfield urine "
    "sediment microscopy image into a virtual phase contrast-style image. The output "
    "is intended only to illustrate the model concept and should not be interpreted "
    "as a clinical result."
)

DISCLAIMER = (
    "⚠️ Demonstration only — not for clinical use. "
    "This model is a research prototype and is not validated as a medical device. "
    "It must not be used for diagnosis, patient care, treatment decisions, laboratory reporting, "
    "or any clinical decision-making. The generated image may be inaccurate, may suppress subtle "
    "features, or may fail to represent clinically important urinary sediment findings, especially "
    "findings outside the training/evaluation data. Upload only de-identified, non-patient-care images."
)

FOOTER = (
    "For research and educational demonstration only. Always rely on validated laboratory methods "
    "and qualified professional review for clinical interpretation."
)

# =========================
# LOAD MODEL
# =========================
if not os.path.exists(MODEL_PATH):
    raise FileNotFoundError(f"Model file not found: {MODEL_PATH}")

model = tf.keras.models.load_model(MODEL_PATH, compile=False)


# =========================
# IMAGE CONVERSION
# =========================
def pil_to_model_input(img: Image.Image, grayscale=False):
    if grayscale:
        img = img.convert("L")
        arr = np.array(img, dtype=np.float32) / 255.0
        arr = arr[..., None]
    else:
        img = img.convert("RGB")
        arr = np.array(img, dtype=np.float32) / 255.0
    return arr


def array_to_pil(arr, grayscale=False):
    arr = np.clip(arr, 0.0, 1.0)

    if grayscale:
        if arr.ndim == 3 and arr.shape[-1] == 1:
            arr = arr[..., 0]
        arr_uint8 = (arr * 255).astype(np.uint8)
        return Image.fromarray(arr_uint8, mode="L")

    if arr.ndim == 3 and arr.shape[-1] == 1:
        arr = np.repeat(arr, 3, axis=-1)
    arr_uint8 = (arr * 255).astype(np.uint8)
    return Image.fromarray(arr_uint8, mode="RGB")


def make_side_by_side(left_pil: Image.Image, right_pil: Image.Image):
    left_rgb = left_pil.convert("RGB")
    right_rgb = right_pil.convert("RGB")

    w1, h1 = left_rgb.size
    w2, h2 = right_rgb.size

    canvas = Image.new("RGB", (w1 + w2, max(h1, h2)), color=(255, 255, 255))
    canvas.paste(left_rgb, (0, 0))
    canvas.paste(right_rgb, (w1, 0))
    return canvas


# =========================
# TILING / STITCHING
# =========================
def pad_to_min_size(img, min_h, min_w):
    h, w = img.shape[:2]
    pad_h = max(0, min_h - h)
    pad_w = max(0, min_w - w)

    if pad_h == 0 and pad_w == 0:
        return img

    return np.pad(img, ((0, pad_h), (0, pad_w), (0, 0)), mode="reflect")


def compute_start_positions(length, patch_size, stride):
    if length <= patch_size:
        return [0]

    starts = list(range(0, length - patch_size + 1, stride))
    if starts[-1] != length - patch_size:
        starts.append(length - patch_size)
    return starts


def make_blend_weight(patch_size):
    w = np.hanning(patch_size).astype(np.float32)
    w = np.outer(w, w).astype(np.float32)
    w = np.maximum(w, 1e-3)
    return w[..., None]


def predict_full_image(model, bright_img, patch_size=256, stride=192, channels=3, batch_size=8):
    orig_h, orig_w = bright_img.shape[:2]
    bright_pad = pad_to_min_size(bright_img, patch_size, patch_size)

    H, W = bright_pad.shape[:2]
    ys = compute_start_positions(H, patch_size, stride)
    xs = compute_start_positions(W, patch_size, stride)

    weight = make_blend_weight(patch_size)
    pred_sum = np.zeros((H, W, channels), dtype=np.float32)
    weight_sum = np.zeros((H, W, 1), dtype=np.float32)

    patches = []
    coords = []

    for y in ys:
        for x in xs:
            patch = bright_pad[y:y + patch_size, x:x + patch_size, :]
            patches.append(patch)
            coords.append((y, x))

    patches = np.asarray(patches, dtype=np.float32)
    preds = model.predict(patches, batch_size=batch_size, verbose=0)

    for pred, (y, x) in zip(preds, coords):
        pred_sum[y:y + patch_size, x:x + patch_size, :] += pred * weight
        weight_sum[y:y + patch_size, x:x + patch_size, :] += weight

    stitched = pred_sum / np.clip(weight_sum, 1e-6, None)
    stitched = stitched[:orig_h, :orig_w, :]
    stitched = np.clip(stitched, 0.0, 1.0)

    return stitched


# =========================
# INFERENCE FUNCTION
# =========================
def run_inference(input_image: Image.Image):
    if input_image is None:
        raise gr.Error("Please upload a de-identified brightfield microscopy image for demonstration.")

    original_display = input_image.convert("RGB")
    x = pil_to_model_input(input_image, grayscale=USE_GRAYSCALE)

    y_pred = predict_full_image(
        model=model,
        bright_img=x,
        patch_size=PATCH_SIZE,
        stride=INFER_STRIDE,
        channels=CHANNELS,
        batch_size=BATCH_SIZE,
    )

    pred_pil = array_to_pil(y_pred, grayscale=USE_GRAYSCALE)
    compare_pil = make_side_by_side(original_display, pred_pil)

    return pred_pil, compare_pil


# =========================
# UI
# =========================
with gr.Blocks(title=TITLE) as demo:
    gr.Markdown(f"# {TITLE}")
    gr.Markdown(DESCRIPTION)
    gr.Markdown(
        f"""
<div style="padding: 14px; border: 2px solid #d9534f; background-color: #fff3f3;
border-radius: 8px; color: #8a1f11; font-weight: 700; line-height: 1.45;">
{DISCLAIMER}
</div>
"""
    )

    with gr.Row():
        input_image = gr.Image(type="pil", label="Input brightfield image for demonstration")
        pred_image = gr.Image(type="pil", label="Demonstration output: virtual phase contrast-style image")

    compare_image = gr.Image(type="pil", label="Side-by-side demonstration comparison")

    with gr.Row():
        run_button = gr.Button("Generate demonstration output")
        clear_button = gr.ClearButton([input_image, pred_image, compare_image])

    gr.Markdown(
        f"""
<div style="font-size: 0.95em; color: #555; padding-top: 8px;">
{FOOTER}
</div>
"""
    )

    run_button.click(
        fn=run_inference,
        inputs=input_image,
        outputs=[pred_image, compare_image],
    )

if __name__ == "__main__":
    demo.launch(server_name="0.0.0.0", server_port=int(os.environ.get("PORT", 7860)))