File size: 7,236 Bytes
af2b273
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
70fb2d5
af2b273
 
 
 
 
5d4272a
af2b273
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a482d12
 
 
 
 
 
 
 
 
 
 
 
 
729ee5e
 
 
 
 
a482d12
 
 
 
 
 
9218b1c
5d4272a
 
 
 
 
 
9218b1c
 
af2b273
 
 
9218b1c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
70fb2d5
 
 
 
 
9218b1c
 
 
 
 
 
5d4272a
 
 
af2b273
 
 
 
 
 
 
 
70fb2d5
 
9218b1c
 
af2b273
 
 
 
 
 
5d4272a
 
af2b273
5d4272a
af2b273
 
 
 
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
"""RADAR abdominal-CT ZeroGPU Space.

Reuse map (all inference logic is vendored, not reimplemented):
- `RADAR_inference/inference_demo.py::initialize` builds the RADAR model and
  loads `ckpt/checkpoint_radar_pretrain.pth` (strict=False), then `.cuda()`.
- `RADAR_inference/inference_demo.py::evaluate` runs the full single-volume
  pipeline: MONAI resample to 1x1x5mm -> HU clip [-300,400] -> min-max norm ->
  non-zero ROI crop -> pad (96,256,384) -> sliding-window forward with
  `inference_demo.RADAR.forward_test_win` -> per-organ center-crop second
  pass -> CSV of mean positive-class scores for 146 organ_finding pairs.
- `RADAR_inference/inference_demo.py::DataFolder` owns every preprocessing
  transform. This file adds no transforms and no report prose: it bridges
  Gradio upload -> temp dir -> evaluate -> (Label, Dataframe of raw scores).

Space layout notes:
- `MODEL_ROOT`/`CONFIGS_ROOT` must be absolute before importing the vendored
  module (it reads them at import time). Weights arrive preloaded at build
  time (`preload_from_hub`, same HF cache `snapshot_download` reads) and are
  symlinked into `ckpt/`; prompt embeddings (`infer_text_embedding_radar.pt`,
  340KB) ship in git because they are absent from the HuggingFace repo.
"""

import os
import re
import shutil
import tempfile

ROOT = os.path.dirname(os.path.abspath(__file__))
CKPT_DIR = os.path.join(ROOT, "ckpt")
os.environ.setdefault("MODEL_ROOT", CKPT_DIR)
os.environ.setdefault("CONFIGS_ROOT", CKPT_DIR)

import sys
sys.path.insert(0, os.path.join(ROOT, "RADAR_inference"))  # vendored absolute imports (dynamic_network_architectures) resolve from here
from intake import HEADER_SUFFIXES, VOLUME_SUFFIXES, _series_to_nifti, _stage_uploads, _to_nifti  # noqa: E402
import pandas as pd  # noqa: E402
import spaces  # noqa: E402
import gradio as gr  # noqa: E402
from huggingface_hub import snapshot_download  # noqa: E402
from inference_demo import initialize, evaluate  # noqa: E402
from backend.inference_backend import VastAIBackend, ZeroGPUBackend, get_backend_kind  # noqa: E402

REPO_ID = "radar-generalist/RADAR"
CHECKPOINT_NAME = "checkpoint_radar_pretrain.pth"


def _ensure_ckpts() -> None:
    snap = snapshot_download(
        REPO_ID,
        allow_patterns=[CHECKPOINT_NAME, "bert-base-chinese/*"],
    )
    targets = [CHECKPOINT_NAME, "bert-base-chinese"]
    for name in targets:
        src = os.path.join(snap, name)
        dst = os.path.join(CKPT_DIR, name)
        if not os.path.exists(src):
            raise RuntimeError(f"{name} absent from {REPO_ID} snapshot {snap}")
        if os.path.lexists(dst):
            continue
        os.symlink(src, dst)
    ckpt_path = os.path.join(CKPT_DIR, CHECKPOINT_NAME)
    if not os.path.exists(ckpt_path):
        raise RuntimeError(
            f"{CHECKPOINT_NAME} missing after download; "
            "check Space logs/network and re-run."
        )


_ensure_ckpts()
pad_func, model = initialize()  # module level; .cuda() here is intentional (ZeroGPU pattern)


def _english_name(column: str) -> str:
    m = re.search(r"\((.+)\)\s*$", column)
    return m.group(1) if m else column


@spaces.GPU(duration=90)
def _score_case(tmpdir, outdir):
    """GPU step only: run the vendored chain on the staged case and parse its CSV."""
    try:
        evaluate(pad_func, model, tmpdir, outdir, "space")
    except (OSError, ValueError, RuntimeError) as exc:
        raise gr.Error(f"could not process this volume: {exc}")
    csv_path = os.path.join(outdir, "RADAR_infer_results_space.csv")
    df = pd.read_csv(csv_path, encoding="utf-8-sig")
    if df.empty:
        raise gr.Error("model skipped this volume (check dimensions/spacing)")
    row = df.iloc[0]
    scores = {
        _english_name(col): float(row[col])
        for col in df.columns[1:]
        # NaN sorts arbitrarily under reverse=True and scrambles ranking —
        # keep finite scores only (phantom/empty cells come back NaN).
        if pd.notna(row[col]) and str(row[col]).strip() != ""
    }
    ranked = sorted(scores.items(), key=lambda kv: kv[1], reverse=True)
    table = pd.DataFrame(ranked, columns=["Finding", "Score"])
    return scores, table


def diagnose(ct_files):
    """Gradio entry (CPU): stage uploads, then score on the active backend."""
    kind = get_backend_kind()
    if kind not in ("zerogpu", "vastai"):
        raise gr.Error("unknown INFERENCE_BACKEND")
    if kind == "vastai":
        VastAIBackend().ensure_available()  # fail before mkdtemp/staging
    items = ct_files if isinstance(ct_files, list) else [ct_files]
    paths = [p if isinstance(p, str) else p.name for p in items]
    tmpdir = tempfile.mkdtemp(prefix="radar_case_")
    outdir = tempfile.mkdtemp(prefix="radar_out_")
    try:
        singles = [p for p in paths if p.lower().endswith(VOLUME_SUFFIXES)]
        rest = [p for p in paths if p not in singles]
        if singles and rest:
            raise gr.Error("upload either one volume file or one DICOM series, not both")
        if len(singles) > 1:
            raise gr.Error("upload a single volume file")
        if singles:
            src = _to_nifti(singles[0], tmpdir)
            fname = os.path.basename(src)
            if src != os.path.join(tmpdir, fname):
                if not fname.endswith((".nii", ".nii.gz")):
                    fname += ".nii.gz"
                shutil.copy(src, os.path.join(tmpdir, fname))
        else:
            if len(paths) == 1 and paths[0].lower().endswith(".dcm"):
                raise gr.Error("a single .dcm is one slice, not a volume: upload the full series")
            headers = [p for p in paths if p.lower().endswith(HEADER_SUFFIXES)]
            if headers:
                raise gr.Error(
                    f"{os.path.basename(headers[0])} needs its raw pair: convert to .nii.gz/.nrrd/.mha first"
                )
            staged = tempfile.mkdtemp(prefix="radar_dcm_")
            try:
                _stage_uploads(paths, staged)
                _series_to_nifti(staged, tmpdir)
            finally:
                shutil.rmtree(staged, ignore_errors=True)
        if kind == "vastai":
            return VastAIBackend().score_case(tmpdir, outdir)
        return ZeroGPUBackend(_score_case).score_case(tmpdir, outdir)
    finally:
        shutil.rmtree(tmpdir, ignore_errors=True)
        shutil.rmtree(outdir, ignore_errors=True)


demo = gr.Interface(
    fn=diagnose,
    inputs=gr.File(
        # No file_types whitelist: PACS DICOM exports are often extensionless, and any
        # whitelist would reject them client-side. diagnose() validates server-side.
        label="Abdominal CT (volume file or DICOM series)",
        file_count="multiple",
    ),
    outputs=[gr.Label(label="Top findings", num_top_classes=10), gr.Dataframe(label="All finding scores")],
    title="RADAR Abdominal CT",
    description=(
        "Expert-level generalist AI for contrast-enhanced abdominal CT "
        "(18 structures, 146 findings). Non-commercial research demo "
        "(CC BY-NC-SA 4.0); assistance tool, not a diagnosis. "
        "Full UI: https://radar-ct.pages.dev — this Space remains as API."
    ),
    api_name="diagnose",
)

if __name__ == "__main__":
    demo.launch()