File size: 4,654 Bytes
745873a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Pulmo — end-to-end two-stage inference.

    raw CT volume  ->  Stage 1 (find nodule centres)
                   ->  Stage 2 (per-candidate diagnosis + concept explanation)

`analyze_scan` returns one dict per detected nodule with its location, the
detection / malignancy probabilities, the segmentation mask of the central
slice, and the concept-bottleneck explanation of the malignancy decision.

Requires: torch, numpy, scipy, huggingface_hub
    pip install torch numpy scipy huggingface_hub
"""

import numpy as np
import torch
from huggingface_hub import hf_hub_download

from modeling import (
    load_stage1, load_stage2, find_candidates,
    crop_stage2_input, explain_malignancy,
)

REPO_ID = "ariyul/Pulmo"
STAGE1_FILE = "stage1_detector_v2.pth"
STAGE2_FILE = "student_2p5d_best.pth"


def load_pipeline(device=None):
    """Download both checkpoints from the Hub and return (stage1, stage2, device)."""
    device = device or ("cuda" if torch.cuda.is_available() else "cpu")
    s1_path = hf_hub_download(REPO_ID, STAGE1_FILE)
    s2_path = hf_hub_download(REPO_ID, STAGE2_FILE)
    stage1 = load_stage1(s1_path, device=device)
    stage2 = load_stage2(s2_path, device=device)
    return stage1, stage2, device


@torch.no_grad()
def analyze_scan(volume, spacing, stage1, stage2, device="cpu",
                 peak_thresh=0.1, det_thresh=0.5, max_candidates=None):
    """Run the full Pulmo pipeline on one CT volume.

    Args:
        volume       : (Z, Y, X) numpy array of raw HU values.
        spacing      : (sz, sy, sx) voxel spacing in mm, [z, y, x] order.
        stage1/stage2: loaded models (see load_pipeline).
        peak_thresh  : Stage-1 heatmap threshold (lower -> more candidates;
                       Stage 2 filters false positives).
        det_thresh   : keep a candidate only if Stage-2 detection prob >= this.
        max_candidates: optionally cap how many Stage-1 candidates to characterise.

    Returns:
        list of dicts (one per kept nodule), sorted by malignancy probability:
          {
            'location_voxel'   : (z, y, x),
            'detection_prob'   : float,
            'malignancy_prob'  : float,
            'prediction'       : 'MALIGNANT' | 'BENIGN',
            'segmentation'     : (64, 64) float mask of the central slice,
            'concepts'         : {name: value, ...},
            'top_reasons'      : [(concept, value, contribution), ...],  # most malignancy-driving first
          }
    """
    cands = find_candidates(stage1, volume, spacing, device=device,
                            peak_thresh=peak_thresh)
    if max_candidates is not None:
        cands = cands[:max_candidates]

    findings = []
    for (z, y, x) in cands:
        xin = crop_stage2_input(volume, (z, y, x)).to(device)
        out = stage2(xin)
        det_p = torch.softmax(out["detection"][0], 0)[1].item()
        if det_p < det_thresh:                 # Stage 2 rejects this candidate
            continue
        mal_p = torch.softmax(out["malignancy"][0], 0)[1].item()
        seg = torch.sigmoid(out["segmentation"][0, 0]).cpu().numpy()
        concepts = out["concepts"][0].cpu().numpy()
        from modeling import CONCEPT_NAMES
        findings.append({
            "location_voxel": (int(z), int(y), int(x)),
            "detection_prob": det_p,
            "malignancy_prob": mal_p,
            "prediction": "MALIGNANT" if mal_p >= 0.5 else "BENIGN",
            "segmentation": seg,
            "concepts": {n: float(v) for n, v in zip(CONCEPT_NAMES, concepts)},
            "top_reasons": explain_malignancy(stage2, out)[:3],
        })

    findings.sort(key=lambda f: -f["malignancy_prob"])
    return findings


def main():
    stage1, stage2, device = load_pipeline()
    print(f"Pipeline loaded on {device}.")

    # --- Replace this with a real CT volume + its voxel spacing ---
    # volume: (Z, Y, X) raw HU; spacing: (sz, sy, sx) mm in [z, y, x] order.
    dummy_volume = np.random.randint(-1000, 200, size=(120, 512, 512)).astype(np.int16)
    dummy_spacing = (1.25, 0.7, 0.7)

    findings = analyze_scan(dummy_volume, dummy_spacing, stage1, stage2, device=device)
    print(f"\n{len(findings)} nodule(s) kept after Stage-2 filtering:\n")
    for i, f in enumerate(findings, 1):
        z, y, x = f["location_voxel"]
        print(f"[{i}] voxel (z={z}, y={y}, x={x})  "
              f"det={f['detection_prob']:.2f}  "
              f"malignancy={f['malignancy_prob']:.2f}  -> {f['prediction']}")
        print("    reasons:", ", ".join(
            f"{name}({contrib:+.2f})" for name, _val, contrib in f["top_reasons"]))


if __name__ == "__main__":
    main()