-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
106 lines (94 loc) · 4.41 KB
/
Copy pathmain.py
File metadata and controls
106 lines (94 loc) · 4.41 KB
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
#!/usr/bin/env python
"""SNET pipeline entry point.
Usage
-----
python main.py --stage 1 --mode train
python main.py --stage 2 --mode infer
python main.py --stage 3 --mode train
python main.py --stage 123 --mode infer # runs all 3 stages sequentially
"""
import argparse
import sys
def main():
parser = argparse.ArgumentParser(description="SNET 3-Stage Spinal Segmentation Pipeline")
parser.add_argument("--stage", type=int, choices=[1, 2, 3, 123], default=123,
help="Which stage to run (1, 2, 3, or 123 for the full pipeline)")
parser.add_argument("--mode", type=str, choices=["train", "infer"], default="train",
help="Run training or inference")
parser.add_argument("--csv", type=str, default=None, help="Override CSV path")
parser.add_argument("--checkpoint", type=str, default=None, help="Override checkpoint path")
parser.add_argument("--device", type=str, default=None, help="Override device (cuda/cpu)")
args = parser.parse_args()
stages = [1, 2, 3] if args.stage == 123 else [args.stage]
for stage in stages:
print(f"\n{'='*60}")
print(f" Stage {stage} — {args.mode.upper()}")
print(f"{'='*60}\n")
if stage == 1:
if args.mode == "train":
from snet.stage1.train import train
if args.csv:
from snet.config import STAGE1_CFG
STAGE1_CFG["csv_path"] = args.csv
if args.checkpoint:
from snet.config import STAGE1_CFG
STAGE1_CFG["checkpoint_path"] = args.checkpoint
if args.device:
from snet.config import STAGE1_CFG
STAGE1_CFG["device"] = args.device
train()
else:
from snet.stage1.infer import run_inference
results = run_inference(csv_path=args.csv, checkpoint_path=args.checkpoint,
device=args.device)
if results:
import pandas as pd
print(pd.DataFrame(results).to_string(index=False))
elif stage == 2:
if args.mode == "train":
from snet.stage2.train import train
if args.csv:
from snet.config import STAGE2_CFG
STAGE2_CFG["csv_path"] = args.csv
if args.checkpoint:
from snet.config import STAGE2_CFG
STAGE2_CFG["checkpoint_path"] = args.checkpoint
if args.device:
from snet.config import STAGE2_CFG
STAGE2_CFG["device"] = args.device
train()
else:
from snet.stage2.infer import run_inference
results = run_inference(csv_path=args.csv, checkpoint_path=args.checkpoint,
device=args.device)
if results:
import pandas as pd
print(pd.DataFrame(results).to_string(index=False))
elif stage == 3:
if args.mode == "train":
from snet.stage3.train import train
if args.csv:
from snet.config import STAGE3_CFG
STAGE3_CFG["csv_path"] = args.csv
if args.checkpoint:
from snet.config import STAGE3_CFG
STAGE3_CFG["checkpoint_path"] = args.checkpoint
if args.device:
from snet.config import STAGE3_CFG
STAGE3_CFG["device"] = args.device
train()
else:
from snet.stage3.infer import run_inference
results = run_inference(csv_path=args.csv, checkpoint_path=args.checkpoint,
device=args.device)
if results:
from snet.eval import compute_dice, compute_hausdorff
for r in results:
if r["gt_bin"].sum() > 0:
dice = compute_dice(r["pred_bin"], r["gt_bin"])
hd = compute_hausdorff(r["pred_bin"], r["gt_bin"])
print(f" {r['subject']} | label {r['label']:3d} | "
f"Dice={dice:.4f} | HD95={hd:.2f} mm")
print(f"\nPipeline finished.")
if __name__ == "__main__":
main()