Add files using upload-large-folder tool
Browse files- metadata/splits/pet_region_text_test.csv +0 -0
- metadata/splits/pet_region_text_train.csv +0 -0
- metadata/splits/pet_region_text_val.csv +0 -0
- metadata/splits/split_summary.csv +4 -0
- metadata/splits/train.csv +0 -0
- metadata/splits/val.csv +153 -0
- requirements.txt +5 -0
- scripts/bootstrap_all_baselines.py +174 -0
- scripts/bootstrap_ci.py +343 -0
- scripts/evaluate_pet_text_alignment.py +113 -0
- scripts/export_case_study.py +92 -0
- scripts/export_embeddings.py +81 -0
- scripts/launch_retrain.sh +3 -0
- scripts/match_adni_metadata.py +263 -0
- scripts/pet_vlm_dataset.py +84 -0
- scripts/probe_mlp_remap.py +166 -0
- scripts/run_clinical_probes_v3.sh +38 -0
- scripts/train_pet_foundation.py +531 -0
- scripts/train_pet_foundation_epoch_ckpt.py +122 -0
- scripts/train_pet_vlm_baseline.py +141 -0
metadata/splits/pet_region_text_test.csv
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
metadata/splits/pet_region_text_train.csv
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
metadata/splits/pet_region_text_val.csv
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
metadata/splits/split_summary.csv
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
split,samples,subjects
|
| 2 |
+
train,710,710
|
| 3 |
+
val,152,152
|
| 4 |
+
test,153,153
|
metadata/splits/train.csv
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
metadata/splits/val.csv
ADDED
|
@@ -0,0 +1,153 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
sample_id,subject_id,visit,pet_path,suvr_csv_path,shape,zooms,dtype,region_rows,suvr_min,suvr_max,suvr_mean,split
|
| 2 |
+
002S2010_M00_fdg_pet,002S2010,M00,fdgpet_M00_112/fdgpet_M00_112/002S2010_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/002S2010_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6130855784696692,1.5451448653597737,1.1682682216133549,val
|
| 3 |
+
002S4229_M00_fdg_pet,002S4229,M00,fdgpet_M00_112/fdgpet_M00_112/002S4229_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/002S4229_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5225892881067787,1.534046360703765,1.1661183716382657,val
|
| 4 |
+
002S4746_M00_fdg_pet,002S4746,M00,fdgpet_M00_112/fdgpet_M00_112/002S4746_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/002S4746_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5143933219705673,1.5033518056400488,1.201207369624356,val
|
| 5 |
+
003S4350_M00_fdg_pet,003S4350,M00,fdgpet_M00_112/fdgpet_M00_112/003S4350_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/003S4350_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4978755604137074,1.5317810454558618,1.142299611667552,val
|
| 6 |
+
005S0222_M00_fdg_pet,005S0222,M00,fdgpet_M00_112/fdgpet_M00_112/005S0222_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/005S0222_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5854023717441501,1.3816601057325315,1.0188035806258344,val
|
| 7 |
+
005S0610_M00_fdg_pet,005S0610,M00,fdgpet_M00_112/fdgpet_M00_112/005S0610_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/005S0610_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6258189432784691,1.776341307607395,1.2159099603775638,val
|
| 8 |
+
005S2390_M00_fdg_pet,005S2390,M00,fdgpet_M00_112/fdgpet_M00_112/005S2390_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/005S2390_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.7625209530454408,1.685342033916314,1.2260925282244046,val
|
| 9 |
+
006S4150_M00_fdg_pet,006S4150,M00,fdgpet_M00_112/fdgpet_M00_112/006S4150_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/006S4150_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6621003634684804,1.4837714401475992,1.148945196622542,val
|
| 10 |
+
006S4192_M00_fdg_pet,006S4192,M00,fdgpet_M00_112/fdgpet_M00_112/006S4192_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/006S4192_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5378416657672205,1.1426088559675982,0.8852363958184436,val
|
| 11 |
+
006S4546_M00_fdg_pet,006S4546,M00,fdgpet_M00_112/fdgpet_M00_112/006S4546_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/006S4546_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6580099365291551,1.822839651964547,1.211820323770937,val
|
| 12 |
+
007S0293_M00_fdg_pet,007S0293,M00,fdgpet_M00_112/fdgpet_M00_112/007S0293_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/007S0293_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5903835194633607,1.4180257025371028,1.0203066208629363,val
|
| 13 |
+
007S1339_M00_fdg_pet,007S1339,M00,fdgpet_M00_112/fdgpet_M00_112/007S1339_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/007S1339_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.675830742138085,1.4161981378944175,1.1148515529208387,val
|
| 14 |
+
007S4488_M00_fdg_pet,007S4488,M00,fdgpet_M00_112/fdgpet_M00_112/007S4488_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/007S4488_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6111476714598304,1.7246942338959952,1.1932071360565562,val
|
| 15 |
+
007S4516_M00_fdg_pet,007S4516,M00,fdgpet_M00_112/fdgpet_M00_112/007S4516_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/007S4516_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4678168475309158,1.631630071429505,1.0975409605412367,val
|
| 16 |
+
009S1199_M00_fdg_pet,009S1199,M00,fdgpet_M00_112/fdgpet_M00_112/009S1199_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/009S1199_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5846835865693933,1.577944664514867,1.2045100547151335,val
|
| 17 |
+
009S2208_M00_fdg_pet,009S2208,M00,fdgpet_M00_112/fdgpet_M00_112/009S2208_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/009S2208_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6622150033553016,1.4999867787346997,1.1168063457431243,val
|
| 18 |
+
009S4359_M00_fdg_pet,009S4359,M00,fdgpet_M00_112/fdgpet_M00_112/009S4359_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/009S4359_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5810416318516043,1.3874002621061168,1.0744753861978642,val
|
| 19 |
+
009S4612_M00_fdg_pet,009S4612,M00,fdgpet_M00_112/fdgpet_M00_112/009S4612_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/009S4612_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6179723507020531,1.497322253818545,1.0967472360442414,val
|
| 20 |
+
010S4345_M00_fdg_pet,010S4345,M00,fdgpet_M00_112/fdgpet_M00_112/010S4345_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/010S4345_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.7771201977934759,1.710800654166846,1.244100253230879,val
|
| 21 |
+
011S0053_M00_fdg_pet,011S0053,M00,fdgpet_M00_112/fdgpet_M00_112/011S0053_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/011S0053_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4053546826827675,1.1916210153548297,0.8599013369002707,val
|
| 22 |
+
011S0861_M00_fdg_pet,011S0861,M00,fdgpet_M00_112/fdgpet_M00_112/011S0861_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/011S0861_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4439952516761041,1.4138248932822108,0.9441382904545598,val
|
| 23 |
+
011S4222_M00_fdg_pet,011S4222,M00,fdgpet_M00_112/fdgpet_M00_112/011S4222_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/011S4222_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5688305154471063,1.5817856804582986,1.1051181256525382,val
|
| 24 |
+
012S4128_M00_fdg_pet,012S4128,M00,fdgpet_M00_112/fdgpet_M00_112/012S4128_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/012S4128_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5277128984583891,1.326189830517328,1.0491980518804167,val
|
| 25 |
+
012S4987_M00_fdg_pet,012S4987,M00,fdgpet_M00_112/fdgpet_M00_112/012S4987_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/012S4987_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5695403494485994,1.3722362937692392,1.102449829741389,val
|
| 26 |
+
013S1161_M00_fdg_pet,013S1161,M00,fdgpet_M00_112/fdgpet_M00_112/013S1161_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/013S1161_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5620511712874958,1.2866886610968855,0.993835889794818,val
|
| 27 |
+
013S4595_M00_fdg_pet,013S4595,M00,fdgpet_M00_112/fdgpet_M00_112/013S4595_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/013S4595_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5609056974712171,1.5450846354166667,1.075773496497761,val
|
| 28 |
+
013S5071_M00_fdg_pet,013S5071,M00,fdgpet_M00_112/fdgpet_M00_112/013S5071_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/013S5071_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.7172918472698028,1.5706837708468264,1.1617727473594244,val
|
| 29 |
+
014S4080_M00_fdg_pet,014S4080,M00,fdgpet_M00_112/fdgpet_M00_112/014S4080_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/014S4080_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.3298627806872856,1.662185586820506,1.1896105448759615,val
|
| 30 |
+
016S0590_M00_fdg_pet,016S0590,M00,fdgpet_M00_112/fdgpet_M00_112/016S0590_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/016S0590_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5314887117084703,1.5859719706322863,1.0550571821372674,val
|
| 31 |
+
016S4121_M00_fdg_pet,016S4121,M00,fdgpet_M00_112/fdgpet_M00_112/016S4121_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/016S4121_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6222262783641682,1.4118782710103157,1.011479981689829,val
|
| 32 |
+
016S4353_M00_fdg_pet,016S4353,M00,fdgpet_M00_112/fdgpet_M00_112/016S4353_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/016S4353_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6706020251472565,1.5926721952081788,1.0906204220308515,val
|
| 33 |
+
016S4601_M00_fdg_pet,016S4601,M00,fdgpet_M00_112/fdgpet_M00_112/016S4601_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/016S4601_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.7421343111734081,1.7780829081965928,1.2076785810022754,val
|
| 34 |
+
016S4951_M00_fdg_pet,016S4951,M00,fdgpet_M00_112/fdgpet_M00_112/016S4951_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/016S4951_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6365926813951788,1.5770693764167747,1.20615236511728,val
|
| 35 |
+
018S4696_M00_fdg_pet,018S4696,M00,fdgpet_M00_112/fdgpet_M00_112/018S4696_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/018S4696_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5789904561244861,1.4161918920816785,0.9960260668339004,val
|
| 36 |
+
018S4733_M00_fdg_pet,018S4733,M00,fdgpet_M00_112/fdgpet_M00_112/018S4733_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/018S4733_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5680867159430356,1.3670695123066472,0.9714958647829468,val
|
| 37 |
+
018S4889_M00_fdg_pet,018S4889,M00,fdgpet_M00_112/fdgpet_M00_112/018S4889_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/018S4889_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5211284453856116,1.5090284594600063,1.148163944117111,val
|
| 38 |
+
019S4285_M00_fdg_pet,019S4285,M00,fdgpet_M00_112/fdgpet_M00_112/019S4285_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/019S4285_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.7139396463485963,1.5090529156884176,1.0956450271365106,val
|
| 39 |
+
019S4477_M00_fdg_pet,019S4477,M00,fdgpet_M00_112/fdgpet_M00_112/019S4477_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/019S4477_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5998719304064872,1.5426449983016304,1.071282532163102,val
|
| 40 |
+
021S0141_M00_fdg_pet,021S0141,M00,fdgpet_M00_112/fdgpet_M00_112/021S0141_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/021S0141_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4105021838835854,1.3590590719114313,0.9765538960952778,val
|
| 41 |
+
021S0642_M00_fdg_pet,021S0642,M00,fdgpet_M00_112/fdgpet_M00_112/021S0642_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/021S0642_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5605707866389577,1.6061984248637482,1.173559266807612,val
|
| 42 |
+
021S0647_M00_fdg_pet,021S0647,M00,fdgpet_M00_112/fdgpet_M00_112/021S0647_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/021S0647_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5130969368837734,1.4764638481811865,1.080991620166608,val
|
| 43 |
+
021S2142_M00_fdg_pet,021S2142,M00,fdgpet_M00_112/fdgpet_M00_112/021S2142_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/021S2142_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5734769851766168,1.6987708569987945,1.1846469824077337,val
|
| 44 |
+
022S0219_M00_fdg_pet,022S0219,M00,fdgpet_M00_112/fdgpet_M00_112/022S0219_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/022S0219_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.3945540973846925,1.6207760201648724,1.0582901740640809,val
|
| 45 |
+
022S0543_M00_fdg_pet,022S0543,M00,fdgpet_M00_112/fdgpet_M00_112/022S0543_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/022S0543_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.3678792861693683,1.7979144737368724,1.1634697318265086,val
|
| 46 |
+
022S4444_M00_fdg_pet,022S4444,M00,fdgpet_M00_112/fdgpet_M00_112/022S4444_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/022S4444_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5842867711915916,1.546184813422996,1.1479727698502424,val
|
| 47 |
+
022S4922_M00_fdg_pet,022S4922,M00,fdgpet_M00_112/fdgpet_M00_112/022S4922_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/022S4922_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.423281909310244,1.5846322481749489,1.106813814388958,val
|
| 48 |
+
023S4448_M00_fdg_pet,023S4448,M00,fdgpet_M00_112/fdgpet_M00_112/023S4448_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/023S4448_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.7046644241414606,1.6551035088784545,1.194730830729444,val
|
| 49 |
+
024S1393_M00_fdg_pet,024S1393,M00,fdgpet_M00_112/fdgpet_M00_112/024S1393_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/024S1393_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5372143251481758,1.4514011817141663,1.0371309525824086,val
|
| 50 |
+
024S1400_M00_fdg_pet,024S1400,M00,fdgpet_M00_112/fdgpet_M00_112/024S1400_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/024S1400_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5228917732659765,1.4166587168039204,1.0687254805438935,val
|
| 51 |
+
024S4169_M00_fdg_pet,024S4169,M00,fdgpet_M00_112/fdgpet_M00_112/024S4169_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/024S4169_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5561296661819755,1.29814074096249,0.9748994127975952,val
|
| 52 |
+
024S4905_M00_fdg_pet,024S4905,M00,fdgpet_M00_112/fdgpet_M00_112/024S4905_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/024S4905_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.7050304897010968,1.569188636912482,1.1374760610343813,val
|
| 53 |
+
027S0408_M00_fdg_pet,027S0408,M00,fdgpet_M00_112/fdgpet_M00_112/027S0408_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/027S0408_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.3817983617119611,1.2828499051365447,1.0130704951830156,val
|
| 54 |
+
027S4873_M00_fdg_pet,027S4873,M00,fdgpet_M00_112/fdgpet_M00_112/027S4873_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/027S4873_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4804024109865892,1.5185279220831198,1.1179255351095194,val
|
| 55 |
+
027S4964_M00_fdg_pet,027S4964,M00,fdgpet_M00_112/fdgpet_M00_112/027S4964_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/027S4964_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5026347114440591,1.2965290376925356,0.9862690706540546,val
|
| 56 |
+
029S1056_M00_fdg_pet,029S1056,M00,fdgpet_M00_112/fdgpet_M00_112/029S1056_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/029S1056_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5048055371424912,1.2723824327457578,0.9051783996276244,val
|
| 57 |
+
031S2018_M00_fdg_pet,031S2018,M00,fdgpet_M00_112/fdgpet_M00_112/031S2018_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/031S2018_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5676287554394153,1.4849969779627792,1.033595681228472,val
|
| 58 |
+
031S2022_M00_fdg_pet,031S2022,M00,fdgpet_M00_112/fdgpet_M00_112/031S2022_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/031S2022_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6763742980501726,1.5995309015313413,1.1602281726684385,val
|
| 59 |
+
031S4005_M00_fdg_pet,031S4005,M00,fdgpet_M00_112/fdgpet_M00_112/031S4005_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/031S4005_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4353031376085295,1.1506237801586394,0.821157213340859,val
|
| 60 |
+
031S4032_M00_fdg_pet,031S4032,M00,fdgpet_M00_112/fdgpet_M00_112/031S4032_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/031S4032_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.555698337275871,1.29195621535911,0.937925081898764,val
|
| 61 |
+
031S4203_M00_fdg_pet,031S4203,M00,fdgpet_M00_112/fdgpet_M00_112/031S4203_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/031S4203_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5745257627747518,1.3915173175607762,1.0490015623098627,val
|
| 62 |
+
032S0400_M00_fdg_pet,032S0400,M00,fdgpet_M00_112/fdgpet_M00_112/032S0400_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/032S0400_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5739163152785505,1.6272965175550165,1.214966748280907,val
|
| 63 |
+
032S0978_M00_fdg_pet,032S0978,M00,fdgpet_M00_112/fdgpet_M00_112/032S0978_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/032S0978_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4045239149200693,1.3067335617127915,1.0455938402574425,val
|
| 64 |
+
032S2119_M00_fdg_pet,032S2119,M00,fdgpet_M00_112/fdgpet_M00_112/032S2119_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/032S2119_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4416101918765262,1.569079931636668,1.1388093713741976,val
|
| 65 |
+
032S4429_M00_fdg_pet,032S4429,M00,fdgpet_M00_112/fdgpet_M00_112/032S4429_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/032S4429_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5051268098023309,1.8301627275083088,1.3055162211119191,val
|
| 66 |
+
032S4755_M00_fdg_pet,032S4755,M00,fdgpet_M00_112/fdgpet_M00_112/032S4755_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/032S4755_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4089400664784784,1.825936730528781,1.3021404695416718,val
|
| 67 |
+
032S4823_M00_fdg_pet,032S4823,M00,fdgpet_M00_112/fdgpet_M00_112/032S4823_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/032S4823_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4417573498649679,1.6437236499503214,1.148885191034538,val
|
| 68 |
+
035S4783_M00_fdg_pet,035S4783,M00,fdgpet_M00_112/fdgpet_M00_112/035S4783_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/035S4783_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4383173243884734,1.439151760560961,1.0304223246945436,val
|
| 69 |
+
036S0760_M00_fdg_pet,036S0760,M00,fdgpet_M00_112/fdgpet_M00_112/036S0760_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/036S0760_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5054986706014652,1.3193502040685416,0.8956313910147379,val
|
| 70 |
+
036S1135_M00_fdg_pet,036S1135,M00,fdgpet_M00_112/fdgpet_M00_112/036S1135_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/036S1135_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5039047631514859,1.5342718339256909,1.1086934880674837,val
|
| 71 |
+
036S2380_M00_fdg_pet,036S2380,M00,fdgpet_M00_112/fdgpet_M00_112/036S2380_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/036S2380_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.7703633465491176,1.911245874998516,1.4498632640524112,val
|
| 72 |
+
036S5112_M00_fdg_pet,036S5112,M00,fdgpet_M00_112/fdgpet_M00_112/036S5112_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/036S5112_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6158807737349631,1.4660782985007046,1.0333091842698912,val
|
| 73 |
+
041S0282_M00_fdg_pet,041S0282,M00,fdgpet_M00_112/fdgpet_M00_112/041S0282_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/041S0282_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5495281869767583,1.5061889853666233,1.1108215379165187,val
|
| 74 |
+
041S0407_M00_fdg_pet,041S0407,M00,fdgpet_M00_112/fdgpet_M00_112/041S0407_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/041S0407_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5679679031553942,1.6956683488239577,1.2530847878740117,val
|
| 75 |
+
041S0598_M00_fdg_pet,041S0598,M00,fdgpet_M00_112/fdgpet_M00_112/041S0598_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/041S0598_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5679166197371529,1.444278337256197,1.0301996420513455,val
|
| 76 |
+
041S1010_M00_fdg_pet,041S1010,M00,fdgpet_M00_112/fdgpet_M00_112/041S1010_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/041S1010_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6185995083861808,1.638986196858923,1.156266154296135,val
|
| 77 |
+
041S1412_M00_fdg_pet,041S1412,M00,fdgpet_M00_112/fdgpet_M00_112/041S1412_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/041S1412_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6134029317119303,1.3265960277432156,1.0302420086082866,val
|
| 78 |
+
041S4200_M00_fdg_pet,041S4200,M00,fdgpet_M00_112/fdgpet_M00_112/041S4200_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/041S4200_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6810143758292029,1.6490700042270794,1.178864886735488,val
|
| 79 |
+
041S4877_M00_fdg_pet,041S4877,M00,fdgpet_M00_112/fdgpet_M00_112/041S4877_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/041S4877_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4124853430435561,1.4667605582551864,1.0283716224833042,val
|
| 80 |
+
041S4989_M00_fdg_pet,041S4989,M00,fdgpet_M00_112/fdgpet_M00_112/041S4989_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/041S4989_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6472930838067952,1.4331789702430895,1.076331521958383,val
|
| 81 |
+
051S4929_M00_fdg_pet,051S4929,M00,fdgpet_M00_112/fdgpet_M00_112/051S4929_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/051S4929_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5789701762952303,1.369352497038294,1.067986713716914,val
|
| 82 |
+
052S1346_M00_fdg_pet,052S1346,M00,fdgpet_M00_112/fdgpet_M00_112/052S1346_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/052S1346_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6076386149336652,1.4018627405166626,1.1382859515836234,val
|
| 83 |
+
053S4578_M00_fdg_pet,053S4578,M00,fdgpet_M00_112/fdgpet_M00_112/053S4578_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/053S4578_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4609880364375049,1.4006361467364046,1.061689848879346,val
|
| 84 |
+
067S4184_M00_fdg_pet,067S4184,M00,fdgpet_M00_112/fdgpet_M00_112/067S4184_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/067S4184_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4428812352622427,1.6309098686784866,1.159896988364192,val
|
| 85 |
+
067S4212_M00_fdg_pet,067S4212,M00,fdgpet_M00_112/fdgpet_M00_112/067S4212_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/067S4212_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.1772498522306743,1.8723115623318085,1.0794208113985475,val
|
| 86 |
+
068S2248_M00_fdg_pet,068S2248,M00,fdgpet_M00_112/fdgpet_M00_112/068S2248_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/068S2248_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6845126503392269,1.6193848848342896,1.1818247371148678,val
|
| 87 |
+
068S4174_M00_fdg_pet,068S4174,M00,fdgpet_M00_112/fdgpet_M00_112/068S4174_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/068S4174_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5411551980411305,1.4675192450911072,1.161298411959886,val
|
| 88 |
+
072S2072_M00_fdg_pet,072S2072,M00,fdgpet_M00_112/fdgpet_M00_112/072S2072_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/072S2072_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6077061629876858,1.399373338374007,1.0571593516817444,val
|
| 89 |
+
072S4613_M00_fdg_pet,072S4613,M00,fdgpet_M00_112/fdgpet_M00_112/072S4613_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/072S4613_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5433775273765006,1.4657371395924053,1.1538627477343724,val
|
| 90 |
+
073S1357_M00_fdg_pet,073S1357,M00,fdgpet_M00_112/fdgpet_M00_112/073S1357_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/073S1357_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5662989568581079,1.339159479071526,1.002903207113124,val
|
| 91 |
+
073S2182_M00_fdg_pet,073S2182,M00,fdgpet_M00_112/fdgpet_M00_112/073S2182_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/073S2182_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5031846806327289,1.5561907147290588,1.1539756350565389,val
|
| 92 |
+
073S4403_M00_fdg_pet,073S4403,M00,fdgpet_M00_112/fdgpet_M00_112/073S4403_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/073S4403_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4937209152593845,1.5351263310477756,1.1761409862406955,val
|
| 93 |
+
073S4540_M00_fdg_pet,073S4540,M00,fdgpet_M00_112/fdgpet_M00_112/073S4540_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/073S4540_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5101386085551053,1.4398923780097337,1.047882823219493,val
|
| 94 |
+
073S4559_M00_fdg_pet,073S4559,M00,fdgpet_M00_112/fdgpet_M00_112/073S4559_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/073S4559_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4675567953344335,1.3798722710222366,1.0540255544479424,val
|
| 95 |
+
073S4986_M00_fdg_pet,073S4986,M00,fdgpet_M00_112/fdgpet_M00_112/073S4986_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/073S4986_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5338241250757227,1.41495638946227,1.0631568401516394,val
|
| 96 |
+
094S0531_M00_fdg_pet,094S0531,M00,fdgpet_M00_112/fdgpet_M00_112/094S0531_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/094S0531_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5052946195434699,1.2238196123576015,0.875657722975816,val
|
| 97 |
+
094S4858_M00_fdg_pet,094S4858,M00,fdgpet_M00_112/fdgpet_M00_112/094S4858_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/094S4858_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6619102815776791,1.940497220527476,1.326667960782625,val
|
| 98 |
+
100S4469_M00_fdg_pet,100S4469,M00,fdgpet_M00_112/fdgpet_M00_112/100S4469_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/100S4469_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.7536898280607606,1.804717021455397,1.344321368509906,val
|
| 99 |
+
100S4556_M00_fdg_pet,100S4556,M00,fdgpet_M00_112/fdgpet_M00_112/100S4556_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/100S4556_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5435223502908798,1.475380071429505,1.130673812076885,val
|
| 100 |
+
109S1157_M00_fdg_pet,109S1157,M00,fdgpet_M00_112/fdgpet_M00_112/109S1157_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/109S1157_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4337788150185033,1.3988137133544103,0.9478397009436418,val
|
| 101 |
+
109S2237_M00_fdg_pet,109S2237,M00,fdgpet_M00_112/fdgpet_M00_112/109S2237_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/109S2237_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4798310631736714,1.685610756355246,1.1225317736301104,val
|
| 102 |
+
109S4260_M00_fdg_pet,109S4260,M00,fdgpet_M00_112/fdgpet_M00_112/109S4260_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/109S4260_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4950967472504805,1.7405294583538389,1.2100268724842629,val
|
| 103 |
+
114S0228_M00_fdg_pet,114S0228,M00,fdgpet_M00_112/fdgpet_M00_112/114S0228_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/114S0228_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.524557511436748,1.495211482048035,1.1276781969014715,val
|
| 104 |
+
114S1118_M00_fdg_pet,114S1118,M00,fdgpet_M00_112/fdgpet_M00_112/114S1118_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/114S1118_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4376195407806233,1.6392725843577212,1.2343341195663948,val
|
| 105 |
+
116S0360_M00_fdg_pet,116S0360,M00,fdgpet_M00_112/fdgpet_M00_112/116S0360_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/116S0360_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4846747495273855,1.4036995547902271,1.022745936383321,val
|
| 106 |
+
116S0657_M00_fdg_pet,116S0657,M00,fdgpet_M00_112/fdgpet_M00_112/116S0657_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/116S0657_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5399971212295287,1.416969426898255,0.9706298423192268,val
|
| 107 |
+
116S4043_M00_fdg_pet,116S4043,M00,fdgpet_M00_112/fdgpet_M00_112/116S4043_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/116S4043_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4753644402651863,1.7511268028846154,1.2110825744743683,val
|
| 108 |
+
116S4199_M00_fdg_pet,116S4199,M00,fdgpet_M00_112/fdgpet_M00_112/116S4199_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/116S4199_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5648776375673672,1.6977729231384455,1.2737383207357618,val
|
| 109 |
+
116S4635_M00_fdg_pet,116S4635,M00,fdgpet_M00_112/fdgpet_M00_112/116S4635_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/116S4635_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5161281320500501,1.7169289597163753,1.2344654290410213,val
|
| 110 |
+
116S4732_M00_fdg_pet,116S4732,M00,fdgpet_M00_112/fdgpet_M00_112/116S4732_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/116S4732_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5868162675337358,1.608450286763169,1.1045684747210942,val
|
| 111 |
+
123S4170_M00_fdg_pet,123S4170,M00,fdgpet_M00_112/fdgpet_M00_112/123S4170_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/123S4170_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6829346912663158,1.7965946367303411,1.3500708110370545,val
|
| 112 |
+
123S4362_M00_fdg_pet,123S4362,M00,fdgpet_M00_112/fdgpet_M00_112/123S4362_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/123S4362_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.694824730477682,1.8650774519822435,1.4285392933849432,val
|
| 113 |
+
123S4780_M00_fdg_pet,123S4780,M00,fdgpet_M00_112/fdgpet_M00_112/123S4780_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/123S4780_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.7042832723477992,1.9239496802478904,1.27720279734896,val
|
| 114 |
+
123S4904_M00_fdg_pet,123S4904,M00,fdgpet_M00_112/fdgpet_M00_112/123S4904_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/123S4904_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6255590857529059,1.6529105187968998,1.2394333926132743,val
|
| 115 |
+
127S0431_M00_fdg_pet,127S0431,M00,fdgpet_M00_112/fdgpet_M00_112/127S0431_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/127S0431_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5503240983116435,1.472652174952476,1.0648580411321882,val
|
| 116 |
+
127S1210_M00_fdg_pet,127S1210,M00,fdgpet_M00_112/fdgpet_M00_112/127S1210_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/127S1210_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5482238730925838,1.320639729499817,0.9291197957206524,val
|
| 117 |
+
127S4197_M00_fdg_pet,127S4197,M00,fdgpet_M00_112/fdgpet_M00_112/127S4197_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/127S4197_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4084707661133439,1.18600368867535,0.897133590951844,val
|
| 118 |
+
127S4198_M00_fdg_pet,127S4198,M00,fdgpet_M00_112/fdgpet_M00_112/127S4198_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/127S4198_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4979782093326511,1.2664426811857878,0.9125774256956306,val
|
| 119 |
+
127S4210_M00_fdg_pet,127S4210,M00,fdgpet_M00_112/fdgpet_M00_112/127S4210_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/127S4210_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6138714312117683,1.591755743789814,1.2084713262473576,val
|
| 120 |
+
127S4624_M00_fdg_pet,127S4624,M00,fdgpet_M00_112/fdgpet_M00_112/127S4624_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/127S4624_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4591505602158557,1.2500750758767936,0.9694744655191088,val
|
| 121 |
+
127S5058_M00_fdg_pet,127S5058,M00,fdgpet_M00_112/fdgpet_M00_112/127S5058_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/127S5058_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4761749900273648,1.6911816650064828,1.1572378090000996,val
|
| 122 |
+
128S0227_M00_fdg_pet,128S0227,M00,fdgpet_M00_112/fdgpet_M00_112/128S0227_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/128S0227_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5185130373461949,1.3465168515174355,1.0010940927230323,val
|
| 123 |
+
128S0230_M00_fdg_pet,128S0230,M00,fdgpet_M00_112/fdgpet_M00_112/128S0230_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/128S0230_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.4604968938831863,1.2730979720718163,0.959332270233602,val
|
| 124 |
+
128S0272_M00_fdg_pet,128S0272,M00,fdgpet_M00_112/fdgpet_M00_112/128S0272_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/128S0272_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5506763940939425,1.1876595045064138,0.9694210373104556,val
|
| 125 |
+
128S2003_M00_fdg_pet,128S2003,M00,fdgpet_M00_112/fdgpet_M00_112/128S2003_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/128S2003_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.575459367077308,1.2925958462875815,0.9963333015582124,val
|
| 126 |
+
128S2151_M00_fdg_pet,128S2151,M00,fdgpet_M00_112/fdgpet_M00_112/128S2151_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/128S2151_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.7232723133972144,1.340928885241908,1.085806299715455,val
|
| 127 |
+
128S4599_M00_fdg_pet,128S4599,M00,fdgpet_M00_112/fdgpet_M00_112/128S4599_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/128S4599_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.594582282032615,1.59699999460274,1.1163305492852862,val
|
| 128 |
+
128S4671_M00_fdg_pet,128S4671,M00,fdgpet_M00_112/fdgpet_M00_112/128S4671_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/128S4671_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5633445538686133,1.4194413575940672,0.993418694355821,val
|
| 129 |
+
128S4745_M00_fdg_pet,128S4745,M00,fdgpet_M00_112/fdgpet_M00_112/128S4745_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/128S4745_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.6917662457536914,1.5167277531869416,1.0765095321197755,val
|
| 130 |
+
128S5123_M00_fdg_pet,128S5123,M00,fdgpet_M00_112/fdgpet_M00_112/128S5123_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/128S5123_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,int16,121,0.5966193675467124,1.2741808536678365,0.946831775912354,val
|
| 131 |
+
130S2391_M00_fdg_pet,130S2391,M00,fdgpet_M00_112/fdgpet_M00_112/130S2391_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/130S2391_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4855878289370613,1.577939048332254,1.1601395989064855,val
|
| 132 |
+
130S4294_M00_fdg_pet,130S4294,M00,fdgpet_M00_112/fdgpet_M00_112/130S4294_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/130S4294_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4219434784679878,1.4953330559874058,1.0662130284635571,val
|
| 133 |
+
130S4352_M00_fdg_pet,130S4352,M00,fdgpet_M00_112/fdgpet_M00_112/130S4352_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/130S4352_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4214410628864471,1.5602226191439983,1.1657892047736926,val
|
| 134 |
+
130S4730_M00_fdg_pet,130S4730,M00,fdgpet_M00_112/fdgpet_M00_112/130S4730_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/130S4730_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4738968221899022,1.5848032216556738,1.143266700274927,val
|
| 135 |
+
130S4990_M00_fdg_pet,130S4990,M00,fdgpet_M00_112/fdgpet_M00_112/130S4990_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/130S4990_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5055370636802307,1.5011500374215547,1.0573446766541,val
|
| 136 |
+
130S5059_M00_fdg_pet,130S5059,M00,fdgpet_M00_112/fdgpet_M00_112/130S5059_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/130S5059_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6791454323606512,1.638562120386684,1.1323196964466178,val
|
| 137 |
+
137S1414_M00_fdg_pet,137S1414,M00,fdgpet_M00_112/fdgpet_M00_112/137S1414_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/137S1414_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5921718577012659,1.554995205735913,1.177327861977282,val
|
| 138 |
+
137S4211_M00_fdg_pet,137S4211,M00,fdgpet_M00_112/fdgpet_M00_112/137S4211_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/137S4211_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4991976080093792,1.8417531909751248,1.132140618017005,val
|
| 139 |
+
137S4258_M00_fdg_pet,137S4258,M00,fdgpet_M00_112/fdgpet_M00_112/137S4258_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/137S4258_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.510324773941448,1.6326896131636706,1.1064582068116495,val
|
| 140 |
+
137S4351_M00_fdg_pet,137S4351,M00,fdgpet_M00_112/fdgpet_M00_112/137S4351_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/137S4351_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5901373388932988,1.818250937540023,1.2603468312469357,val
|
| 141 |
+
137S4587_M00_fdg_pet,137S4587,M00,fdgpet_M00_112/fdgpet_M00_112/137S4587_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/137S4587_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5410589942320145,1.4834014392289958,1.092879867525642,val
|
| 142 |
+
137S4678_M00_fdg_pet,137S4678,M00,fdgpet_M00_112/fdgpet_M00_112/137S4678_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/137S4678_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5798987256507946,1.5253534409606342,1.114381609253059,val
|
| 143 |
+
141S1245_M00_fdg_pet,141S1245,M00,fdgpet_M00_112/fdgpet_M00_112/141S1245_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/141S1245_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5955459676324365,1.523713158779457,1.1623531693069742,val
|
| 144 |
+
141S4976_M00_fdg_pet,141S4976,M00,fdgpet_M00_112/fdgpet_M00_112/141S4976_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/141S4976_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6598377737769469,1.7956448156856797,1.3006220077852035,val
|
| 145 |
+
153S2109_M00_fdg_pet,153S2109,M00,fdgpet_M00_112/fdgpet_M00_112/153S2109_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/153S2109_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5687621825519092,1.4451423568470425,1.0767287103965644,val
|
| 146 |
+
153S4125_M00_fdg_pet,153S4125,M00,fdgpet_M00_112/fdgpet_M00_112/153S4125_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/153S4125_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6868296367366139,1.6014197928006533,1.167998578932741,val
|
| 147 |
+
153S4139_M00_fdg_pet,153S4139,M00,fdgpet_M00_112/fdgpet_M00_112/153S4139_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/153S4139_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.5859117558933197,1.7202189263691472,1.2340255902012656,val
|
| 148 |
+
153S4372_M00_fdg_pet,153S4372,M00,fdgpet_M00_112/fdgpet_M00_112/153S4372_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/153S4372_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.7064865030706885,1.516750693321228,1.1699427077555888,val
|
| 149 |
+
941S1194_M00_fdg_pet,941S1194,M00,fdgpet_M00_112/fdgpet_M00_112/941S1194_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/941S1194_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.55352024343562,1.7806404649613294,1.2293535351933855,val
|
| 150 |
+
941S1195_M00_fdg_pet,941S1195,M00,fdgpet_M00_112/fdgpet_M00_112/941S1195_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/941S1195_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4789982352027281,1.47586802134528,1.1170504416312104,val
|
| 151 |
+
941S1203_M00_fdg_pet,941S1203,M00,fdgpet_M00_112/fdgpet_M00_112/941S1203_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/941S1203_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.6127659063562421,1.5401486158370972,1.107546502340703,val
|
| 152 |
+
941S4036_M00_fdg_pet,941S4036,M00,fdgpet_M00_112/fdgpet_M00_112/941S4036_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/941S4036_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.4613087700012533,1.7614303763330037,1.245207389085386,val
|
| 153 |
+
941S4255_M00_fdg_pet,941S4255,M00,fdgpet_M00_112/fdgpet_M00_112/941S4255_M00_fdg_pet.nii.gz,petfdg_suvr_csv/petfdg_suvr_csv/941S4255_M00_fdg_pet.csv,112x128x112,1.5x1.5x1.5,float32,121,0.3891148694696273,1.3502883911132812,1.0257422868697763,val
|
requirements.txt
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
nibabel
|
| 2 |
+
numpy
|
| 3 |
+
pandas
|
| 4 |
+
torch
|
| 5 |
+
scikit-learn
|
scripts/bootstrap_all_baselines.py
ADDED
|
@@ -0,0 +1,174 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Bootstrap 95% CI for ALL frozen baseline encoders (Stage-1 only: MAE + R@1).
|
| 2 |
+
|
| 3 |
+
Runs 1000-resample bootstrap on the 153-subject test set for each baseline.
|
| 4 |
+
|
| 5 |
+
Usage (from /data/Albus/Brain):
|
| 6 |
+
CUDA_VISIBLE_DEVICES=2 python scripts/bootstrap_all_baselines.py
|
| 7 |
+
"""
|
| 8 |
+
from __future__ import annotations
|
| 9 |
+
|
| 10 |
+
import sys
|
| 11 |
+
import time
|
| 12 |
+
from pathlib import Path
|
| 13 |
+
|
| 14 |
+
import numpy as np
|
| 15 |
+
import torch
|
| 16 |
+
import torch.nn.functional as F
|
| 17 |
+
from torch.utils.data import DataLoader
|
| 18 |
+
|
| 19 |
+
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
| 20 |
+
from pet_vlm_dataset import PETSUVRDataset, collate_pet_suvr
|
| 21 |
+
from train_pet_foundation import PETSUVRFoundationModel, build_encoder
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
# -- baselines ---------------------------------------------------------------
|
| 25 |
+
BASELINES = [
|
| 26 |
+
("MedicalNet frozen", "runs/foundation/medicalnet_frozen_mlp.pt"),
|
| 27 |
+
("BrainIAC frozen", "runs/foundation/brainiac_frozen_mlp.pt"),
|
| 28 |
+
("BrainFM frozen", "runs/foundation/brainfm_frozen_mlp_b4_best.pt"),
|
| 29 |
+
("SAM-Med3D frozen", "runs/foundation/sam_med3d_frozen_mlp_best.pt"),
|
| 30 |
+
("SwinUNETR frozen", "runs/foundation/swinunetr_frozen_mlp_best.pt"),
|
| 31 |
+
]
|
| 32 |
+
|
| 33 |
+
TEST_MANIFEST = Path("metadata/splits/test.csv")
|
| 34 |
+
B = 1000
|
| 35 |
+
SEED = 42
|
| 36 |
+
BATCH_SIZE = 4
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def _retrieval_recall_at_1(logits: np.ndarray) -> float:
|
| 40 |
+
ranks = []
|
| 41 |
+
for i in range(logits.shape[0]):
|
| 42 |
+
order = np.argsort(-logits[i])
|
| 43 |
+
rank = int(np.where(order == i)[0][0]) + 1
|
| 44 |
+
ranks.append(rank)
|
| 45 |
+
return float(np.mean(np.asarray(ranks) <= 1))
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
@torch.no_grad()
|
| 49 |
+
def collect_stage1(model, loader, device):
|
| 50 |
+
model.eval()
|
| 51 |
+
pred_c, tgt_c, pz_c, sz_c = [], [], [], []
|
| 52 |
+
for batch in loader:
|
| 53 |
+
image = batch["image"].to(device, non_blocking=True)
|
| 54 |
+
suvr = batch["suvr"].to(device, non_blocking=True)
|
| 55 |
+
outputs = model(image, suvr)
|
| 56 |
+
pred_c.append(outputs["pred_suvr"].cpu().numpy())
|
| 57 |
+
tgt_c.append(suvr.cpu().numpy())
|
| 58 |
+
pet_feat = model.pet_encoder(image)
|
| 59 |
+
pet_z = F.normalize(model.pet_projector(pet_feat), dim=-1)
|
| 60 |
+
suvr_z = F.normalize(model.suvr_encoder(suvr), dim=-1)
|
| 61 |
+
pz_c.append(pet_z.cpu().numpy())
|
| 62 |
+
sz_c.append(suvr_z.cpu().numpy())
|
| 63 |
+
return {
|
| 64 |
+
"pred": np.concatenate(pred_c),
|
| 65 |
+
"target": np.concatenate(tgt_c),
|
| 66 |
+
"pet_z": np.concatenate(pz_c),
|
| 67 |
+
"suvr_z": np.concatenate(sz_c),
|
| 68 |
+
}
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def stage1_metrics(d, idx):
|
| 72 |
+
pred = d["pred"][idx]
|
| 73 |
+
target = d["target"][idx]
|
| 74 |
+
uid = np.unique(idx)
|
| 75 |
+
logits = d["pet_z"][uid] @ d["suvr_z"][uid].T
|
| 76 |
+
return {
|
| 77 |
+
"mae": float(np.mean(np.abs(pred - target))),
|
| 78 |
+
"pet_suvr_r1": _retrieval_recall_at_1(logits),
|
| 79 |
+
}
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def bootstrap_ci(metric_fn, n, B=1000, seed=42):
|
| 83 |
+
rng = np.random.RandomState(seed)
|
| 84 |
+
all_idx = np.arange(n)
|
| 85 |
+
point = metric_fn(all_idx)
|
| 86 |
+
boots = np.empty(B)
|
| 87 |
+
for b in range(B):
|
| 88 |
+
idx = rng.choice(n, size=n, replace=True)
|
| 89 |
+
boots[b] = metric_fn(idx)
|
| 90 |
+
lo = float(np.percentile(boots, 2.5))
|
| 91 |
+
hi = float(np.percentile(boots, 97.5))
|
| 92 |
+
return point, lo, hi
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def load_model(ckpt_path, device):
|
| 96 |
+
ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
|
| 97 |
+
saved = ckpt.get("args", {})
|
| 98 |
+
|
| 99 |
+
class _A:
|
| 100 |
+
pass
|
| 101 |
+
|
| 102 |
+
a = _A()
|
| 103 |
+
a.backbone = saved.get("backbone", "medicalnet")
|
| 104 |
+
a.medicalnet_weights = Path(saved.get("medicalnet_weights",
|
| 105 |
+
"pretrained/medicalnet/resnet_50_23dataset.pth"))
|
| 106 |
+
a.brainiac_weights = Path(saved.get("brainiac_weights",
|
| 107 |
+
"pretrained/brainiac/backbone.safetensors"))
|
| 108 |
+
a.brainfm_weights = Path(saved.get("brainfm_weights",
|
| 109 |
+
"pretrained/brainfm/assets/brainfm_pretrained.pth"))
|
| 110 |
+
a.brainfm_code_root = Path(saved.get("brainfm_code_root", "pretrained/brainfm"))
|
| 111 |
+
a.swinunetr_weights = Path(saved.get("swinunetr_weights",
|
| 112 |
+
"pretrained/swinunetr/model_swinvit.pt"))
|
| 113 |
+
a.sam_med3d_weights = Path(saved.get("sam_med3d_weights",
|
| 114 |
+
"pretrained/sam-med3d/sam_med3d_turbo.pth"))
|
| 115 |
+
a.output_size = tuple(saved.get("output_size", (96, 96, 96)))
|
| 116 |
+
embed_dim = saved.get("embed_dim", 256)
|
| 117 |
+
freeze = bool(saved.get("freeze_encoder", False))
|
| 118 |
+
|
| 119 |
+
ds_tmp = PETSUVRDataset(TEST_MANIFEST, output_size=a.output_size)
|
| 120 |
+
n_regions = int(ds_tmp[0]["suvr"].numel())
|
| 121 |
+
encoder = build_encoder(a)
|
| 122 |
+
model = PETSUVRFoundationModel(encoder, n_regions, embed_dim, freeze).to(device)
|
| 123 |
+
model.load_state_dict(ckpt["model"], strict=True)
|
| 124 |
+
model.eval()
|
| 125 |
+
return model, a.output_size
|
| 126 |
+
|
| 127 |
+
|
| 128 |
+
def main():
|
| 129 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 130 |
+
print(f"Device: {device}", flush=True)
|
| 131 |
+
|
| 132 |
+
results = []
|
| 133 |
+
for name, ckpt_path in BASELINES:
|
| 134 |
+
t0 = time.time()
|
| 135 |
+
print(f"\n{'='*60}", flush=True)
|
| 136 |
+
print(f" {name} ({ckpt_path})", flush=True)
|
| 137 |
+
print(f"{'='*60}", flush=True)
|
| 138 |
+
|
| 139 |
+
model, output_size = load_model(ckpt_path, device)
|
| 140 |
+
bs = 2 if "sam_med3d" in ckpt_path else BATCH_SIZE
|
| 141 |
+
ds = PETSUVRDataset(TEST_MANIFEST, output_size=output_size)
|
| 142 |
+
loader = DataLoader(ds, batch_size=bs, shuffle=False,
|
| 143 |
+
num_workers=2, collate_fn=collate_pet_suvr)
|
| 144 |
+
d = collect_stage1(model, loader, device)
|
| 145 |
+
N = d["pred"].shape[0]
|
| 146 |
+
print(f" N = {N}", flush=True)
|
| 147 |
+
|
| 148 |
+
for metric_name in ("mae", "pet_suvr_r1"):
|
| 149 |
+
fn = lambda idx, _m=metric_name: stage1_metrics(d, idx)[_m]
|
| 150 |
+
pt, lo, hi = bootstrap_ci(fn, N, B=B, seed=SEED)
|
| 151 |
+
print(f" {metric_name:20s} {pt:.4f} 95% CI [{lo:.4f}, {hi:.4f}]", flush=True)
|
| 152 |
+
results.append((name, metric_name, pt, lo, hi))
|
| 153 |
+
|
| 154 |
+
# free GPU memory
|
| 155 |
+
del model
|
| 156 |
+
torch.cuda.empty_cache()
|
| 157 |
+
print(f" elapsed: {time.time()-t0:.1f}s", flush=True)
|
| 158 |
+
|
| 159 |
+
# ---- summary table ----
|
| 160 |
+
print(f"\n\n{'='*70}", flush=True)
|
| 161 |
+
print(f"SUMMARY: Bootstrap 95% CI (B={B}, seed={SEED})", flush=True)
|
| 162 |
+
print(f"{'='*70}", flush=True)
|
| 163 |
+
print(f"{'Model':<22s} {'MAE':>8s} {'MAE 95% CI':>18s} {'R@1':>8s} {'R@1 95% CI':>18s}", flush=True)
|
| 164 |
+
print("-"*70, flush=True)
|
| 165 |
+
for i in range(0, len(results), 2):
|
| 166 |
+
nm = results[i][0]
|
| 167 |
+
mae_pt, mae_lo, mae_hi = results[i][2], results[i][3], results[i][4]
|
| 168 |
+
r1_pt, r1_lo, r1_hi = results[i+1][2], results[i+1][3], results[i+1][4]
|
| 169 |
+
print(f"{nm:<22s} {mae_pt:8.4f} [{mae_lo:.4f}, {mae_hi:.4f}] {r1_pt:8.4f} [{r1_lo:.4f}, {r1_hi:.4f}]", flush=True)
|
| 170 |
+
print(f"{'='*70}", flush=True)
|
| 171 |
+
|
| 172 |
+
|
| 173 |
+
if __name__ == "__main__":
|
| 174 |
+
main()
|
scripts/bootstrap_ci.py
ADDED
|
@@ -0,0 +1,343 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Bootstrap 95 % confidence intervals for ReMAP-PET key metrics.
|
| 2 |
+
|
| 3 |
+
Stage-1 metrics (153 test subjects):
|
| 4 |
+
- SUVR MAE
|
| 5 |
+
- Pearson r (voxel-level across all subjects x regions)
|
| 6 |
+
- PET->SUVR Recall@1 (retrieval)
|
| 7 |
+
|
| 8 |
+
Clinical probe metrics:
|
| 9 |
+
- AD vs CN AUROC (logistic regression on PET embeddings)
|
| 10 |
+
- 3-way (CN/MCI/AD) AUROC
|
| 11 |
+
|
| 12 |
+
Usage (from /data/Albus/Brain):
|
| 13 |
+
CUDA_VISIBLE_DEVICES=1 python scripts/bootstrap_ci.py
|
| 14 |
+
"""
|
| 15 |
+
|
| 16 |
+
from __future__ import annotations
|
| 17 |
+
|
| 18 |
+
import argparse
|
| 19 |
+
import sys
|
| 20 |
+
from pathlib import Path
|
| 21 |
+
|
| 22 |
+
import numpy as np
|
| 23 |
+
import pandas as pd
|
| 24 |
+
import torch
|
| 25 |
+
import torch.nn.functional as F
|
| 26 |
+
from torch.utils.data import DataLoader
|
| 27 |
+
from sklearn.linear_model import LogisticRegression
|
| 28 |
+
from sklearn.metrics import balanced_accuracy_score, roc_auc_score
|
| 29 |
+
from sklearn.pipeline import make_pipeline
|
| 30 |
+
from sklearn.preprocessing import LabelEncoder, StandardScaler, label_binarize
|
| 31 |
+
|
| 32 |
+
# -- project imports (scripts/ is the working dir's sibling) -----------------
|
| 33 |
+
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
| 34 |
+
from pet_vlm_dataset import PETSUVRDataset, collate_pet_suvr
|
| 35 |
+
from train_pet_foundation import PETSUVRFoundationModel, build_encoder
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
# ---------------------------------------------------------------------------
|
| 39 |
+
# helpers copied from evaluate_pet_foundation.py
|
| 40 |
+
# ---------------------------------------------------------------------------
|
| 41 |
+
|
| 42 |
+
def _pearson_flat(pred: np.ndarray, target: np.ndarray) -> float:
|
| 43 |
+
p = pred.reshape(-1)
|
| 44 |
+
t = target.reshape(-1)
|
| 45 |
+
if p.std() < 1e-8 or t.std() < 1e-8:
|
| 46 |
+
return float("nan")
|
| 47 |
+
return float(np.corrcoef(p, t)[0, 1])
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def _retrieval_recall_at_1(logits: np.ndarray) -> float:
|
| 51 |
+
ranks = []
|
| 52 |
+
for i in range(logits.shape[0]):
|
| 53 |
+
order = np.argsort(-logits[i])
|
| 54 |
+
rank = int(np.where(order == i)[0][0]) + 1
|
| 55 |
+
ranks.append(rank)
|
| 56 |
+
return float(np.mean(np.asarray(ranks) <= 1))
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
# ---------------------------------------------------------------------------
|
| 60 |
+
# Stage-1: forward pass -> per-subject arrays
|
| 61 |
+
# ---------------------------------------------------------------------------
|
| 62 |
+
|
| 63 |
+
@torch.no_grad()
|
| 64 |
+
def collect_stage1(
|
| 65 |
+
model: PETSUVRFoundationModel,
|
| 66 |
+
loader: DataLoader,
|
| 67 |
+
device: torch.device,
|
| 68 |
+
) -> dict[str, np.ndarray]:
|
| 69 |
+
"""Return pred_suvr, target_suvr, pet_z, suvr_z (all numpy, N-first)."""
|
| 70 |
+
model.eval()
|
| 71 |
+
pred_chunks, target_chunks = [], []
|
| 72 |
+
pet_z_chunks, suvr_z_chunks = [], []
|
| 73 |
+
|
| 74 |
+
for batch in loader:
|
| 75 |
+
image = batch["image"].to(device, non_blocking=True)
|
| 76 |
+
suvr = batch["suvr"].to(device, non_blocking=True)
|
| 77 |
+
outputs = model(image, suvr)
|
| 78 |
+
pred_chunks.append(outputs["pred_suvr"].cpu().numpy())
|
| 79 |
+
target_chunks.append(suvr.cpu().numpy())
|
| 80 |
+
|
| 81 |
+
pet_feat = model.pet_encoder(image)
|
| 82 |
+
pet_z = F.normalize(model.pet_projector(pet_feat), dim=-1)
|
| 83 |
+
suvr_z = F.normalize(model.suvr_encoder(suvr), dim=-1)
|
| 84 |
+
pet_z_chunks.append(pet_z.cpu().numpy())
|
| 85 |
+
suvr_z_chunks.append(suvr_z.cpu().numpy())
|
| 86 |
+
|
| 87 |
+
return {
|
| 88 |
+
"pred": np.concatenate(pred_chunks, axis=0),
|
| 89 |
+
"target": np.concatenate(target_chunks, axis=0),
|
| 90 |
+
"pet_z": np.concatenate(pet_z_chunks, axis=0),
|
| 91 |
+
"suvr_z": np.concatenate(suvr_z_chunks, axis=0),
|
| 92 |
+
}
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def stage1_metrics(d: dict[str, np.ndarray], idx: np.ndarray) -> dict[str, float]:
|
| 96 |
+
"""Compute stage-1 metrics on a subset given by *idx*.
|
| 97 |
+
|
| 98 |
+
MAE and Pearson work fine with duplicate indices (bootstrap).
|
| 99 |
+
For retrieval R@1 we need unique subjects (duplicates would make the
|
| 100 |
+
diagonal ground-truth ambiguous), so we deduplicate *idx* first.
|
| 101 |
+
"""
|
| 102 |
+
pred = d["pred"][idx]
|
| 103 |
+
target = d["target"][idx]
|
| 104 |
+
|
| 105 |
+
# retrieval: use unique indices only
|
| 106 |
+
uid = np.unique(idx)
|
| 107 |
+
pet_z = d["pet_z"][uid]
|
| 108 |
+
suvr_z = d["suvr_z"][uid]
|
| 109 |
+
logits = pet_z @ suvr_z.T
|
| 110 |
+
|
| 111 |
+
return {
|
| 112 |
+
"mae": float(np.mean(np.abs(pred - target))),
|
| 113 |
+
"pearson": _pearson_flat(pred, target),
|
| 114 |
+
"pet_suvr_r1": _retrieval_recall_at_1(logits),
|
| 115 |
+
}
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
# ---------------------------------------------------------------------------
|
| 119 |
+
# Clinical: extract embeddings, train probe, evaluate
|
| 120 |
+
# ---------------------------------------------------------------------------
|
| 121 |
+
|
| 122 |
+
@torch.no_grad()
|
| 123 |
+
def extract_embeddings(
|
| 124 |
+
model: PETSUVRFoundationModel,
|
| 125 |
+
manifest: Path,
|
| 126 |
+
output_size: tuple[int, int, int],
|
| 127 |
+
batch_size: int,
|
| 128 |
+
num_workers: int,
|
| 129 |
+
device: torch.device,
|
| 130 |
+
) -> tuple[pd.DataFrame, np.ndarray]:
|
| 131 |
+
dataset = PETSUVRDataset(manifest, output_size=output_size)
|
| 132 |
+
loader = DataLoader(dataset, batch_size=batch_size, shuffle=False,
|
| 133 |
+
num_workers=num_workers, collate_fn=collate_pet_suvr)
|
| 134 |
+
feats = []
|
| 135 |
+
model.eval()
|
| 136 |
+
for batch in loader:
|
| 137 |
+
image = batch["image"].to(device, non_blocking=True)
|
| 138 |
+
pet_feat = model.pet_encoder(image)
|
| 139 |
+
pet_z = F.normalize(model.pet_projector(pet_feat), dim=-1)
|
| 140 |
+
feats.append(pet_z.cpu().numpy())
|
| 141 |
+
return pd.read_csv(manifest), np.concatenate(feats, axis=0)
|
| 142 |
+
|
| 143 |
+
|
| 144 |
+
def _subset_cls(df, x, column, labels):
|
| 145 |
+
mask = df[column].isin(labels).to_numpy()
|
| 146 |
+
return x[mask], df.loc[mask, column].astype(str).to_numpy()
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
def train_probe(x_train, y_train, x_val, y_val):
|
| 150 |
+
"""Train logistic probe with C sweep; return best model + encoder."""
|
| 151 |
+
enc = LabelEncoder()
|
| 152 |
+
enc.fit(np.concatenate([y_train, y_val]))
|
| 153 |
+
y_tr = enc.transform(y_train)
|
| 154 |
+
y_v = enc.transform(y_val)
|
| 155 |
+
best_m, best_s = None, -np.inf
|
| 156 |
+
for c in [0.01, 0.03, 0.1, 0.3, 1.0, 3.0, 10.0]:
|
| 157 |
+
m = make_pipeline(StandardScaler(),
|
| 158 |
+
LogisticRegression(C=c, max_iter=5000,
|
| 159 |
+
class_weight="balanced"))
|
| 160 |
+
m.fit(x_train, y_tr)
|
| 161 |
+
s = balanced_accuracy_score(y_v, m.predict(x_val))
|
| 162 |
+
if s > best_s:
|
| 163 |
+
best_m, best_s = m, s
|
| 164 |
+
return best_m, enc
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
def clinical_auroc(model_probe, encoder, x_test, y_test):
|
| 168 |
+
"""Return AUROC (binary or macro-OVR)."""
|
| 169 |
+
y_int = encoder.transform(y_test)
|
| 170 |
+
proba = model_probe.predict_proba(x_test)
|
| 171 |
+
if len(encoder.classes_) == 2:
|
| 172 |
+
return roc_auc_score(y_int, proba[:, 1])
|
| 173 |
+
else:
|
| 174 |
+
y_bin = label_binarize(y_int, classes=np.arange(len(encoder.classes_)))
|
| 175 |
+
return roc_auc_score(y_bin, proba, average="macro", multi_class="ovr")
|
| 176 |
+
|
| 177 |
+
|
| 178 |
+
# ---------------------------------------------------------------------------
|
| 179 |
+
# Bootstrap
|
| 180 |
+
# ---------------------------------------------------------------------------
|
| 181 |
+
|
| 182 |
+
def bootstrap_ci(
|
| 183 |
+
metric_fn,
|
| 184 |
+
n: int,
|
| 185 |
+
B: int = 1000,
|
| 186 |
+
seed: int = 42,
|
| 187 |
+
alpha: float = 0.05,
|
| 188 |
+
) -> tuple[float, float, float]:
|
| 189 |
+
"""
|
| 190 |
+
metric_fn(idx) -> float where idx is array of resampled indices.
|
| 191 |
+
Returns (point_estimate, lo, hi) for the (1-alpha) CI.
|
| 192 |
+
"""
|
| 193 |
+
rng = np.random.RandomState(seed)
|
| 194 |
+
all_idx = np.arange(n)
|
| 195 |
+
point = metric_fn(all_idx)
|
| 196 |
+
boots = np.empty(B)
|
| 197 |
+
for b in range(B):
|
| 198 |
+
idx = rng.choice(n, size=n, replace=True)
|
| 199 |
+
boots[b] = metric_fn(idx)
|
| 200 |
+
lo = float(np.percentile(boots, 100 * alpha / 2))
|
| 201 |
+
hi = float(np.percentile(boots, 100 * (1 - alpha / 2)))
|
| 202 |
+
return point, lo, hi
|
| 203 |
+
|
| 204 |
+
|
| 205 |
+
def bootstrap_clinical_auroc(
|
| 206 |
+
probe, encoder,
|
| 207 |
+
x_train, y_train_raw,
|
| 208 |
+
x_val, y_val_raw,
|
| 209 |
+
x_test, y_test_raw,
|
| 210 |
+
B: int = 1000,
|
| 211 |
+
seed: int = 42,
|
| 212 |
+
alpha: float = 0.05,
|
| 213 |
+
) -> tuple[float, float, float]:
|
| 214 |
+
"""
|
| 215 |
+
Bootstrap over the *test* set only (probe is fixed).
|
| 216 |
+
"""
|
| 217 |
+
rng = np.random.RandomState(seed)
|
| 218 |
+
n = len(y_test_raw)
|
| 219 |
+
all_idx = np.arange(n)
|
| 220 |
+
point = clinical_auroc(probe, encoder, x_test, y_test_raw)
|
| 221 |
+
|
| 222 |
+
boots = np.empty(B)
|
| 223 |
+
for b in range(B):
|
| 224 |
+
idx = rng.choice(n, size=n, replace=True)
|
| 225 |
+
try:
|
| 226 |
+
boots[b] = clinical_auroc(probe, encoder, x_test[idx], y_test_raw[idx])
|
| 227 |
+
except ValueError:
|
| 228 |
+
# can happen if a resample has only one class
|
| 229 |
+
boots[b] = np.nan
|
| 230 |
+
boots = boots[~np.isnan(boots)]
|
| 231 |
+
lo = float(np.percentile(boots, 100 * alpha / 2))
|
| 232 |
+
hi = float(np.percentile(boots, 100 * (1 - alpha / 2)))
|
| 233 |
+
return point, lo, hi
|
| 234 |
+
|
| 235 |
+
|
| 236 |
+
# ---------------------------------------------------------------------------
|
| 237 |
+
# main
|
| 238 |
+
# ---------------------------------------------------------------------------
|
| 239 |
+
|
| 240 |
+
def main() -> None:
|
| 241 |
+
parser = argparse.ArgumentParser()
|
| 242 |
+
parser.add_argument("--checkpoint", type=Path,
|
| 243 |
+
default=Path("runs/foundation/medicalnet_layer4_regalign_best.pt"))
|
| 244 |
+
parser.add_argument("--test-manifest", type=Path,
|
| 245 |
+
default=Path("metadata/splits/test.csv"))
|
| 246 |
+
parser.add_argument("--train-clinical", type=Path,
|
| 247 |
+
default=Path("data/metadata/splits/train_clinical_server.csv"))
|
| 248 |
+
parser.add_argument("--val-clinical", type=Path,
|
| 249 |
+
default=Path("data/metadata/splits/val_clinical_server.csv"))
|
| 250 |
+
parser.add_argument("--test-clinical", type=Path,
|
| 251 |
+
default=Path("data/metadata/splits/test_clinical_server.csv"))
|
| 252 |
+
parser.add_argument("--batch-size", type=int, default=4)
|
| 253 |
+
parser.add_argument("--num-workers", type=int, default=2)
|
| 254 |
+
parser.add_argument("--B", type=int, default=1000, help="bootstrap resamples")
|
| 255 |
+
parser.add_argument("--seed", type=int, default=42)
|
| 256 |
+
args = parser.parse_args()
|
| 257 |
+
|
| 258 |
+
# ---- load model -------------------------------------------------------
|
| 259 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 260 |
+
ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
|
| 261 |
+
saved = ckpt.get("args", {})
|
| 262 |
+
|
| 263 |
+
class _Args:
|
| 264 |
+
pass
|
| 265 |
+
margs = _Args()
|
| 266 |
+
margs.backbone = saved.get("backbone", "medicalnet")
|
| 267 |
+
margs.medicalnet_weights = Path(saved.get("medicalnet_weights",
|
| 268 |
+
"pretrained/medicalnet/resnet_50_23dataset.pth"))
|
| 269 |
+
margs.brainiac_weights = Path(saved.get("brainiac_weights",
|
| 270 |
+
"pretrained/brainiac/backbone.safetensors"))
|
| 271 |
+
margs.brainfm_weights = Path("pretrained/brainfm/assets/brainfm_pretrained.pth")
|
| 272 |
+
margs.brainfm_code_root = Path("pretrained/brainfm")
|
| 273 |
+
margs.swinunetr_weights = Path("pretrained/swinunetr/model_swinvit.pt")
|
| 274 |
+
margs.sam_med3d_weights = Path("pretrained/sam-med3d/sam_med3d_turbo.pth")
|
| 275 |
+
margs.output_size = tuple(saved.get("output_size", (96, 96, 96)))
|
| 276 |
+
embed_dim = saved.get("embed_dim", 256)
|
| 277 |
+
freeze_encoder = bool(saved.get("freeze_encoder", False))
|
| 278 |
+
|
| 279 |
+
output_size = margs.output_size
|
| 280 |
+
|
| 281 |
+
# build model
|
| 282 |
+
dataset_tmp = PETSUVRDataset(args.test_manifest, output_size=output_size)
|
| 283 |
+
n_regions = int(dataset_tmp[0]["suvr"].numel())
|
| 284 |
+
encoder = build_encoder(margs)
|
| 285 |
+
model = PETSUVRFoundationModel(encoder, n_regions, embed_dim, freeze_encoder).to(device)
|
| 286 |
+
model.load_state_dict(ckpt["model"], strict=True)
|
| 287 |
+
model.eval()
|
| 288 |
+
print(f"Loaded checkpoint: {args.checkpoint}", flush=True)
|
| 289 |
+
print(f"backbone={margs.backbone} embed_dim={embed_dim} "
|
| 290 |
+
f"freeze={freeze_encoder} output_size={output_size}", flush=True)
|
| 291 |
+
|
| 292 |
+
# ===== STAGE 1 =========================================================
|
| 293 |
+
print("\n===== Stage-1 evaluation (test set) =====", flush=True)
|
| 294 |
+
test_ds = PETSUVRDataset(args.test_manifest, output_size=output_size)
|
| 295 |
+
test_loader = DataLoader(test_ds, batch_size=args.batch_size, shuffle=False,
|
| 296 |
+
num_workers=args.num_workers, collate_fn=collate_pet_suvr)
|
| 297 |
+
d = collect_stage1(model, test_loader, device)
|
| 298 |
+
N = d["pred"].shape[0]
|
| 299 |
+
print(f" N = {N}", flush=True)
|
| 300 |
+
|
| 301 |
+
for name in ("mae", "pearson", "pet_suvr_r1"):
|
| 302 |
+
fn = lambda idx, _n=name: stage1_metrics(d, idx)[_n]
|
| 303 |
+
pt, lo, hi = bootstrap_ci(fn, N, B=args.B, seed=args.seed)
|
| 304 |
+
print(f" {name:20s} {pt:.4f} 95% CI [{lo:.4f}, {hi:.4f}]", flush=True)
|
| 305 |
+
|
| 306 |
+
# ===== CLINICAL ========================================================
|
| 307 |
+
print("\n===== Clinical downstream probes =====", flush=True)
|
| 308 |
+
train_df, x_train_all = extract_embeddings(
|
| 309 |
+
model, args.train_clinical, output_size, args.batch_size, args.num_workers, device)
|
| 310 |
+
val_df, x_val_all = extract_embeddings(
|
| 311 |
+
model, args.val_clinical, output_size, args.batch_size, args.num_workers, device)
|
| 312 |
+
test_df, x_test_all = extract_embeddings(
|
| 313 |
+
model, args.test_clinical, output_size, args.batch_size, args.num_workers, device)
|
| 314 |
+
|
| 315 |
+
# ---- AD vs CN ---------------------------------------------------------
|
| 316 |
+
print("\n -- AD vs CN --", flush=True)
|
| 317 |
+
x_tr, y_tr = _subset_cls(train_df, x_train_all, "clinical_label", ["CN", "AD"])
|
| 318 |
+
x_v, y_v = _subset_cls(val_df, x_val_all, "clinical_label", ["CN", "AD"])
|
| 319 |
+
x_te, y_te = _subset_cls(test_df, x_test_all, "clinical_label", ["CN", "AD"])
|
| 320 |
+
print(f" train={len(y_tr)} val={len(y_v)} test={len(y_te)}", flush=True)
|
| 321 |
+
probe_ad, enc_ad = train_probe(x_tr, y_tr, x_v, y_v)
|
| 322 |
+
pt, lo, hi = bootstrap_clinical_auroc(
|
| 323 |
+
probe_ad, enc_ad, x_tr, y_tr, x_v, y_v, x_te, y_te,
|
| 324 |
+
B=args.B, seed=args.seed)
|
| 325 |
+
print(f" {'ad_vs_cn_auroc':20s} {pt:.4f} 95% CI [{lo:.4f}, {hi:.4f}]", flush=True)
|
| 326 |
+
|
| 327 |
+
# ---- 3-way CN / MCI / AD ---------------------------------------------
|
| 328 |
+
print("\n -- 3-way (CN / MCI / AD) --", flush=True)
|
| 329 |
+
x_tr3, y_tr3 = _subset_cls(train_df, x_train_all, "clinical_label", ["CN", "MCI", "AD"])
|
| 330 |
+
x_v3, y_v3 = _subset_cls(val_df, x_val_all, "clinical_label", ["CN", "MCI", "AD"])
|
| 331 |
+
x_te3, y_te3 = _subset_cls(test_df, x_test_all, "clinical_label", ["CN", "MCI", "AD"])
|
| 332 |
+
print(f" train={len(y_tr3)} val={len(y_v3)} test={len(y_te3)}", flush=True)
|
| 333 |
+
probe_3w, enc_3w = train_probe(x_tr3, y_tr3, x_v3, y_v3)
|
| 334 |
+
pt, lo, hi = bootstrap_clinical_auroc(
|
| 335 |
+
probe_3w, enc_3w, x_tr3, y_tr3, x_v3, y_v3, x_te3, y_te3,
|
| 336 |
+
B=args.B, seed=args.seed)
|
| 337 |
+
print(f" {'3way_auroc':20s} {pt:.4f} 95% CI [{lo:.4f}, {hi:.4f}]", flush=True)
|
| 338 |
+
|
| 339 |
+
print("\nDone.", flush=True)
|
| 340 |
+
|
| 341 |
+
|
| 342 |
+
if __name__ == "__main__":
|
| 343 |
+
main()
|
scripts/evaluate_pet_text_alignment.py
ADDED
|
@@ -0,0 +1,113 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import argparse
|
| 4 |
+
import csv
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
from torch.utils.data import DataLoader
|
| 10 |
+
|
| 11 |
+
from train_pet_text_alignment import PETTextAlignmentModel, PETTextDataset, collate_pet_text, load_pet_model
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def retrieval_metrics(logits: np.ndarray) -> dict[str, float]:
|
| 15 |
+
ranks = []
|
| 16 |
+
for i in range(logits.shape[0]):
|
| 17 |
+
order = np.argsort(-logits[i])
|
| 18 |
+
ranks.append(int(np.where(order == i)[0][0]) + 1)
|
| 19 |
+
ranks = np.asarray(ranks)
|
| 20 |
+
return {
|
| 21 |
+
"recall@1": float(np.mean(ranks <= 1)),
|
| 22 |
+
"recall@5": float(np.mean(ranks <= 5)),
|
| 23 |
+
"recall@10": float(np.mean(ranks <= 10)),
|
| 24 |
+
"mrr": float(np.mean(1.0 / ranks)),
|
| 25 |
+
"median_rank": float(np.median(ranks)),
|
| 26 |
+
}
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def split_regions(value: str) -> set[str]:
|
| 30 |
+
return {item for item in str(value).split("|") if item}
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def factuality(logits: np.ndarray, lows: list[str], highs: list[str], k: int = 5) -> dict[str, float]:
|
| 34 |
+
top_text = np.argmax(logits, axis=1)
|
| 35 |
+
low_scores = []
|
| 36 |
+
high_scores = []
|
| 37 |
+
for query_idx, text_idx in enumerate(top_text.tolist()):
|
| 38 |
+
query_low = split_regions(lows[query_idx])
|
| 39 |
+
query_high = split_regions(highs[query_idx])
|
| 40 |
+
text_low = split_regions(lows[text_idx])
|
| 41 |
+
text_high = split_regions(highs[text_idx])
|
| 42 |
+
low_scores.append(len(query_low & text_low) / max(min(k, len(query_low)), 1))
|
| 43 |
+
high_scores.append(len(query_high & text_high) / max(min(k, len(query_high)), 1))
|
| 44 |
+
return {
|
| 45 |
+
"retrieved_text_low_overlap": float(np.mean(low_scores)),
|
| 46 |
+
"retrieved_text_high_overlap": float(np.mean(high_scores)),
|
| 47 |
+
}
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def main() -> None:
|
| 51 |
+
parser = argparse.ArgumentParser(description="Evaluate controlled PET-to-region-text alignment.")
|
| 52 |
+
parser.add_argument("--checkpoint", type=Path, required=True)
|
| 53 |
+
parser.add_argument("--test-csv", type=Path, required=True)
|
| 54 |
+
parser.add_argument("--csv-out", type=Path, default=None)
|
| 55 |
+
parser.add_argument("--batch-size", type=int, default=8)
|
| 56 |
+
parser.add_argument("--num-workers", type=int, default=2)
|
| 57 |
+
parser.add_argument("--max-length", type=int, default=None)
|
| 58 |
+
args = parser.parse_args()
|
| 59 |
+
|
| 60 |
+
from transformers import AutoTokenizer
|
| 61 |
+
|
| 62 |
+
ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
|
| 63 |
+
saved = argparse.Namespace(**ckpt["args"])
|
| 64 |
+
saved.train_csv = saved.train_csv
|
| 65 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 66 |
+
pet_model = load_pet_model(saved.pet_checkpoint, saved, device)
|
| 67 |
+
tokenizer = AutoTokenizer.from_pretrained(saved.text_model)
|
| 68 |
+
model = PETTextAlignmentModel(pet_model, saved.text_model, saved.embed_dim or 256).to(device)
|
| 69 |
+
model.load_state_dict(ckpt["model"], strict=True)
|
| 70 |
+
model.eval()
|
| 71 |
+
|
| 72 |
+
output_size = tuple(saved.output_size)
|
| 73 |
+
dataset = PETTextDataset(args.test_csv, output_size=output_size)
|
| 74 |
+
loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers, collate_fn=collate_pet_text)
|
| 75 |
+
pet_chunks = []
|
| 76 |
+
text_chunks = []
|
| 77 |
+
lows: list[str] = []
|
| 78 |
+
highs: list[str] = []
|
| 79 |
+
max_length = args.max_length or saved.max_length
|
| 80 |
+
|
| 81 |
+
with torch.no_grad():
|
| 82 |
+
for batch in loader:
|
| 83 |
+
image = batch["image"].to(device, non_blocking=True)
|
| 84 |
+
tokens = tokenizer(batch["text"], padding=True, truncation=True, max_length=max_length, return_tensors="pt")
|
| 85 |
+
tokens = {k: v.to(device) for k, v in tokens.items()}
|
| 86 |
+
pet_chunks.append(model.encode_pet(image).cpu())
|
| 87 |
+
text_chunks.append(model.encode_text(tokens).cpu())
|
| 88 |
+
lows.extend(batch["low_regions"])
|
| 89 |
+
highs.extend(batch["high_regions"])
|
| 90 |
+
|
| 91 |
+
pet_z = torch.cat(pet_chunks, dim=0)
|
| 92 |
+
text_z = torch.cat(text_chunks, dim=0)
|
| 93 |
+
logits = (pet_z @ text_z.T).numpy()
|
| 94 |
+
metrics = {"samples": float(logits.shape[0])}
|
| 95 |
+
metrics.update({f"pet_to_text_{k}": v for k, v in retrieval_metrics(logits).items()})
|
| 96 |
+
metrics.update({f"text_to_pet_{k}": v for k, v in retrieval_metrics(logits.T).items()})
|
| 97 |
+
metrics.update(factuality(logits, lows, highs))
|
| 98 |
+
|
| 99 |
+
for key, value in metrics.items():
|
| 100 |
+
print(f"{key}={value:.6f}")
|
| 101 |
+
|
| 102 |
+
if args.csv_out:
|
| 103 |
+
args.csv_out.parent.mkdir(parents=True, exist_ok=True)
|
| 104 |
+
write_header = not args.csv_out.exists()
|
| 105 |
+
with args.csv_out.open("a", newline="", encoding="utf-8") as f:
|
| 106 |
+
writer = csv.DictWriter(f, fieldnames=["checkpoint", "test_csv", *metrics.keys()])
|
| 107 |
+
if write_header:
|
| 108 |
+
writer.writeheader()
|
| 109 |
+
writer.writerow({"checkpoint": str(args.checkpoint), "test_csv": str(args.test_csv), **metrics})
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
if __name__ == "__main__":
|
| 113 |
+
main()
|
scripts/export_case_study.py
ADDED
|
@@ -0,0 +1,92 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Export one test subject's data for Figure 4 case study."""
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
|
| 4 |
+
import argparse
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
|
| 7 |
+
import nibabel as nib
|
| 8 |
+
import numpy as np
|
| 9 |
+
import pandas as pd
|
| 10 |
+
import torch
|
| 11 |
+
from torch.utils.data import DataLoader
|
| 12 |
+
|
| 13 |
+
from pet_vlm_dataset import PETSUVRDataset, collate_pet_suvr
|
| 14 |
+
from train_pet_foundation import PETSUVRFoundationModel, build_encoder
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def main():
|
| 18 |
+
parser = argparse.ArgumentParser()
|
| 19 |
+
parser.add_argument("--checkpoint", type=Path, required=True)
|
| 20 |
+
parser.add_argument("--split", type=Path, default=Path("data/metadata/splits/test.csv"))
|
| 21 |
+
parser.add_argument("--subject-index", type=int, default=0, help="which test subject to use (0=first)")
|
| 22 |
+
parser.add_argument("--backbone", default=None)
|
| 23 |
+
parser.add_argument("--medicalnet-weights", type=Path, default=Path("pretrained/medicalnet/resnet_50_23dataset.pth"))
|
| 24 |
+
parser.add_argument("--batch-size", type=int, default=1)
|
| 25 |
+
parser.add_argument("--num-workers", type=int, default=0)
|
| 26 |
+
parser.add_argument("--output-size", type=int, nargs=3, default=None)
|
| 27 |
+
parser.add_argument("--embed-dim", type=int, default=None)
|
| 28 |
+
parser.add_argument("--freeze-encoder", type=bool, default=None)
|
| 29 |
+
parser.add_argument("--out", type=Path, required=True)
|
| 30 |
+
args = parser.parse_args()
|
| 31 |
+
|
| 32 |
+
ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
|
| 33 |
+
saved_args = ckpt.get("args", {})
|
| 34 |
+
for name in ("backbone", "embed_dim", "freeze_encoder"):
|
| 35 |
+
if getattr(args, name, None) is None and name in saved_args:
|
| 36 |
+
setattr(args, name, saved_args[name])
|
| 37 |
+
if args.output_size is None:
|
| 38 |
+
args.output_size = tuple(saved_args.get("output_size", (96, 96, 96)))
|
| 39 |
+
|
| 40 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 41 |
+
manifest = pd.read_csv(args.split)
|
| 42 |
+
row = manifest.iloc[args.subject_index]
|
| 43 |
+
|
| 44 |
+
dataset = PETSUVRDataset(args.split, output_size=tuple(args.output_size))
|
| 45 |
+
sample = dataset[args.subject_index]
|
| 46 |
+
n_regions = int(sample["suvr"].numel())
|
| 47 |
+
|
| 48 |
+
encoder = build_encoder(args)
|
| 49 |
+
model = PETSUVRFoundationModel(encoder, n_regions, args.embed_dim or 256, bool(args.freeze_encoder)).to(device)
|
| 50 |
+
model.load_state_dict(ckpt["model"], strict=True)
|
| 51 |
+
model.eval()
|
| 52 |
+
|
| 53 |
+
image = sample["image"].unsqueeze(0).to(device)
|
| 54 |
+
suvr = sample["suvr"].unsqueeze(0).to(device)
|
| 55 |
+
|
| 56 |
+
with torch.no_grad():
|
| 57 |
+
outputs = model(image, suvr)
|
| 58 |
+
|
| 59 |
+
true_suvr_arr = suvr.cpu().numpy().flatten()
|
| 60 |
+
pred_suvr_arr = outputs["pred_suvr"].cpu().numpy().flatten()
|
| 61 |
+
|
| 62 |
+
# Load PET volume for slice extraction
|
| 63 |
+
pet_path = str(row["pet_path"])
|
| 64 |
+
pet_vol = nib.load(pet_path).get_fdata(dtype=np.float32)
|
| 65 |
+
# Extract middle slices in each orientation
|
| 66 |
+
d, h, w = pet_vol.shape
|
| 67 |
+
axial_slice = pet_vol[d // 2, :, :]
|
| 68 |
+
coronal_slice = pet_vol[:, h // 2, :]
|
| 69 |
+
sagittal_slice = pet_vol[:, :, w // 2]
|
| 70 |
+
|
| 71 |
+
# Load region labels from SUVR CSV
|
| 72 |
+
suvr_csv_path = str(row["suvr_csv_path"])
|
| 73 |
+
suvr_df = pd.read_csv(suvr_csv_path)
|
| 74 |
+
region_labels = [str(lbl) for lbl in suvr_df["label_name"].tolist() if str(lbl) != "Background"]
|
| 75 |
+
|
| 76 |
+
args.out.parent.mkdir(parents=True, exist_ok=True)
|
| 77 |
+
np.savez(args.out,
|
| 78 |
+
subject_id=str(row["sample_id"]),
|
| 79 |
+
true_suvr=true_suvr_arr,
|
| 80 |
+
pred_suvr=pred_suvr_arr,
|
| 81 |
+
region_labels=np.array(region_labels, dtype=object),
|
| 82 |
+
axial_slice=axial_slice,
|
| 83 |
+
coronal_slice=coronal_slice,
|
| 84 |
+
sagittal_slice=sagittal_slice)
|
| 85 |
+
print(f"wrote case study for subject {row['sample_id']} to {args.out}")
|
| 86 |
+
print(f" true SUVR range: [{true_suvr_arr.min():.3f}, {true_suvr_arr.max():.3f}]")
|
| 87 |
+
print(f" pred SUVR range: [{pred_suvr_arr.min():.3f}, {pred_suvr_arr.max():.3f}]")
|
| 88 |
+
print(f" PET shape: {pet_vol.shape}, slices extracted")
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
if __name__ == "__main__":
|
| 92 |
+
main()
|
scripts/export_embeddings.py
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Export PET and SUVR embeddings from Stage 1 checkpoints for Figure 2."""
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
|
| 4 |
+
import argparse
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
from torch.utils.data import DataLoader
|
| 10 |
+
|
| 11 |
+
from pet_vlm_dataset import PETSUVRDataset, collate_pet_suvr
|
| 12 |
+
from train_pet_foundation import PETSUVRFoundationModel, build_encoder
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def main():
|
| 16 |
+
parser = argparse.ArgumentParser()
|
| 17 |
+
parser.add_argument("--checkpoint", type=Path, required=True)
|
| 18 |
+
parser.add_argument("--split", type=Path, default=Path("data/metadata/splits/test.csv"))
|
| 19 |
+
parser.add_argument("--backbone", default=None)
|
| 20 |
+
parser.add_argument("--medicalnet-weights", type=Path, default=Path("pretrained/medicalnet/resnet_50_23dataset.pth"))
|
| 21 |
+
parser.add_argument("--brainfm-weights", type=Path, default=Path("pretrained/brainfm/assets/brainfm_pretrained.pth"))
|
| 22 |
+
parser.add_argument("--brainfm-code-root", type=Path, default=Path("pretrained/brainfm"))
|
| 23 |
+
parser.add_argument("--batch-size", type=int, default=4)
|
| 24 |
+
parser.add_argument("--num-workers", type=int, default=2)
|
| 25 |
+
parser.add_argument("--output-size", type=int, nargs=3, default=None)
|
| 26 |
+
parser.add_argument("--embed-dim", type=int, default=None)
|
| 27 |
+
parser.add_argument("--freeze-encoder", type=bool, default=None)
|
| 28 |
+
parser.add_argument("--out", type=Path, required=True)
|
| 29 |
+
args = parser.parse_args()
|
| 30 |
+
|
| 31 |
+
ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
|
| 32 |
+
saved_args = ckpt.get("args", {})
|
| 33 |
+
for name in ("backbone", "embed_dim", "freeze_encoder"):
|
| 34 |
+
if getattr(args, name, None) is None and name in saved_args:
|
| 35 |
+
setattr(args, name, saved_args[name])
|
| 36 |
+
if args.output_size is None:
|
| 37 |
+
args.output_size = tuple(saved_args.get("output_size", (96, 96, 96)))
|
| 38 |
+
|
| 39 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 40 |
+
dataset = PETSUVRDataset(args.split, output_size=tuple(args.output_size))
|
| 41 |
+
loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=False,
|
| 42 |
+
num_workers=args.num_workers, collate_fn=collate_pet_suvr)
|
| 43 |
+
n_regions = int(dataset[0]["suvr"].numel())
|
| 44 |
+
|
| 45 |
+
encoder = build_encoder(args)
|
| 46 |
+
model = PETSUVRFoundationModel(encoder, n_regions, args.embed_dim or 256, bool(args.freeze_encoder)).to(device)
|
| 47 |
+
model.load_state_dict(ckpt["model"], strict=True)
|
| 48 |
+
model.eval()
|
| 49 |
+
|
| 50 |
+
pet_embeddings = []
|
| 51 |
+
suvr_embeddings = []
|
| 52 |
+
pred_suvr = []
|
| 53 |
+
true_suvr = []
|
| 54 |
+
|
| 55 |
+
with torch.no_grad():
|
| 56 |
+
for batch in loader:
|
| 57 |
+
image = batch["image"].to(device)
|
| 58 |
+
suvr = batch["suvr"].to(device)
|
| 59 |
+
outputs = model(image, suvr)
|
| 60 |
+
|
| 61 |
+
pet_feat = model.pet_encoder(image)
|
| 62 |
+
pet_z = torch.nn.functional.normalize(model.pet_projector(pet_feat), dim=-1)
|
| 63 |
+
suvr_z = torch.nn.functional.normalize(model.suvr_encoder(suvr), dim=-1)
|
| 64 |
+
|
| 65 |
+
pet_embeddings.append(pet_z.cpu().numpy())
|
| 66 |
+
suvr_embeddings.append(suvr_z.cpu().numpy())
|
| 67 |
+
pred_suvr.append(outputs["pred_suvr"].cpu().numpy())
|
| 68 |
+
true_suvr.append(suvr.cpu().numpy())
|
| 69 |
+
|
| 70 |
+
pet_z = np.concatenate(pet_embeddings, axis=0)
|
| 71 |
+
suvr_z = np.concatenate(suvr_embeddings, axis=0)
|
| 72 |
+
pred = np.concatenate(pred_suvr, axis=0)
|
| 73 |
+
true = np.concatenate(true_suvr, axis=0)
|
| 74 |
+
|
| 75 |
+
args.out.parent.mkdir(parents=True, exist_ok=True)
|
| 76 |
+
np.savez(args.out, pet_z=pet_z, suvr_z=suvr_z, pred_suvr=pred, true_suvr=true)
|
| 77 |
+
print(f"wrote {pet_z.shape[0]} samples, pet_z={pet_z.shape}, suvr_z={suvr_z.shape} to {args.out}")
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
if __name__ == "__main__":
|
| 81 |
+
main()
|
scripts/launch_retrain.sh
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
cd /data/Albus/Brain
|
| 3 |
+
CUDA_VISIBLE_DEVICES=1 /data/Albus/miniconda3/bin/python -u scripts/train_pet_foundation_epoch_ckpt.py --backbone medicalnet --encoder-train-scope layer4 --epochs 50 --batch-size 4 --lr 1e-5 --num-workers 2 --output-size 96 96 96 --contrastive-weight 0.2 --regression-weight 1.0 --train-csv metadata/splits/train.csv --val-csv metadata/splits/val.csv --out-dir runs/foundation/remap_epochs > logs/remap_epoch_training.log 2>&1
|
scripts/match_adni_metadata.py
ADDED
|
@@ -0,0 +1,263 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import argparse
|
| 4 |
+
import re
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
|
| 7 |
+
import pandas as pd
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
FULL_COLUMNS = [
|
| 11 |
+
"RID",
|
| 12 |
+
"PTID",
|
| 13 |
+
"label",
|
| 14 |
+
"VISCODE",
|
| 15 |
+
"EXAMDATE",
|
| 16 |
+
"DX_bl",
|
| 17 |
+
"DX",
|
| 18 |
+
"AGE",
|
| 19 |
+
"PTGENDER",
|
| 20 |
+
"PTEDUCAT",
|
| 21 |
+
"PTETHCAT",
|
| 22 |
+
"PTRACCAT",
|
| 23 |
+
"PTMARRY",
|
| 24 |
+
"APOE4",
|
| 25 |
+
"FDG",
|
| 26 |
+
"PIB",
|
| 27 |
+
"AV45",
|
| 28 |
+
"ABETA",
|
| 29 |
+
"TAU",
|
| 30 |
+
"PTAU",
|
| 31 |
+
"CDRSB",
|
| 32 |
+
"ADAS11",
|
| 33 |
+
"ADAS13",
|
| 34 |
+
"ADASQ4",
|
| 35 |
+
"MMSE",
|
| 36 |
+
"RAVLT_immediate",
|
| 37 |
+
"RAVLT_learning",
|
| 38 |
+
"RAVLT_forgetting",
|
| 39 |
+
"RAVLT_perc_forgetting",
|
| 40 |
+
"LDELTOTAL",
|
| 41 |
+
"DIGITSCOR",
|
| 42 |
+
"TRABSCOR",
|
| 43 |
+
"FAQ",
|
| 44 |
+
"MOCA",
|
| 45 |
+
"Years_bl",
|
| 46 |
+
]
|
| 47 |
+
|
| 48 |
+
CONVERSION_COLUMNS = [
|
| 49 |
+
"PTID",
|
| 50 |
+
"label",
|
| 51 |
+
"DX.1",
|
| 52 |
+
"mPACCdigit",
|
| 53 |
+
"mPACCtrailsB",
|
| 54 |
+
"Ventricles(心室)",
|
| 55 |
+
"Hippocampus(海马)",
|
| 56 |
+
"WholeBrain(全脑)",
|
| 57 |
+
"Entorhinal(内嗅觉)",
|
| 58 |
+
"Fusiform(梭形)",
|
| 59 |
+
"MidTemp(中点温度)",
|
| 60 |
+
"ICV",
|
| 61 |
+
]
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def normalize_subject_id(value: object) -> str:
|
| 65 |
+
if pd.isna(value):
|
| 66 |
+
return ""
|
| 67 |
+
text = str(value).strip().upper()
|
| 68 |
+
text = text.replace("_S_", "S").replace("_", "")
|
| 69 |
+
return re.sub(r"[^A-Z0-9]", "", text)
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def clean_numeric(value: object) -> float | pd.NA:
|
| 73 |
+
if pd.isna(value):
|
| 74 |
+
return pd.NA
|
| 75 |
+
text = str(value).strip()
|
| 76 |
+
if not text:
|
| 77 |
+
return pd.NA
|
| 78 |
+
if text.startswith(">"):
|
| 79 |
+
text = text[1:]
|
| 80 |
+
try:
|
| 81 |
+
return float(text)
|
| 82 |
+
except ValueError:
|
| 83 |
+
return pd.NA
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
def read_xlsx(path: Path, columns: list[str]) -> pd.DataFrame:
|
| 87 |
+
df = pd.read_excel(path)
|
| 88 |
+
available = [col for col in columns if col in df.columns]
|
| 89 |
+
if "PTID" not in available:
|
| 90 |
+
raise ValueError(f"{path} does not contain PTID.")
|
| 91 |
+
df = df[available].copy()
|
| 92 |
+
df["subject_id_norm"] = df["PTID"].map(normalize_subject_id)
|
| 93 |
+
df = df[df["subject_id_norm"] != ""].copy()
|
| 94 |
+
df = df.drop_duplicates("subject_id_norm", keep="first")
|
| 95 |
+
return df
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
def add_numeric_clean_columns(df: pd.DataFrame) -> pd.DataFrame:
|
| 99 |
+
for col in ["FDG", "PIB", "AV45", "ABETA", "TAU", "PTAU"]:
|
| 100 |
+
if col in df.columns:
|
| 101 |
+
df[f"{col}_num"] = df[col].map(clean_numeric)
|
| 102 |
+
return df
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def write_split_manifests(enriched: pd.DataFrame, out_dir: Path) -> None:
|
| 106 |
+
out_dir.mkdir(parents=True, exist_ok=True)
|
| 107 |
+
for split_name, split_df in enriched.groupby("split", sort=False):
|
| 108 |
+
split_df.to_csv(out_dir / f"{split_name}_clinical.csv", index=False)
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
def summarize(enriched: pd.DataFrame, out_md: Path, out_csv: Path) -> None:
|
| 112 |
+
rows = []
|
| 113 |
+
total = len(enriched)
|
| 114 |
+
for col in [
|
| 115 |
+
"clinical_label",
|
| 116 |
+
"dx",
|
| 117 |
+
"dx_bl",
|
| 118 |
+
"conversion_label",
|
| 119 |
+
"age",
|
| 120 |
+
"sex",
|
| 121 |
+
"apoe4",
|
| 122 |
+
"mmse",
|
| 123 |
+
"cdrsb",
|
| 124 |
+
"adas11",
|
| 125 |
+
"adas13",
|
| 126 |
+
"faq",
|
| 127 |
+
"moca",
|
| 128 |
+
"fdg_adni",
|
| 129 |
+
"av45",
|
| 130 |
+
"abeta_num",
|
| 131 |
+
"tau_num",
|
| 132 |
+
"ptau_num",
|
| 133 |
+
"hippocampus",
|
| 134 |
+
"wholebrain",
|
| 135 |
+
]:
|
| 136 |
+
if col not in enriched.columns:
|
| 137 |
+
continue
|
| 138 |
+
non_missing = int(enriched[col].notna().sum())
|
| 139 |
+
rows.append({"field": col, "non_missing": non_missing, "coverage": non_missing / total})
|
| 140 |
+
report = pd.DataFrame(rows)
|
| 141 |
+
report.to_csv(out_csv, index=False)
|
| 142 |
+
|
| 143 |
+
label_counts = {}
|
| 144 |
+
for col in ["clinical_label", "dx", "dx_bl", "conversion_label"]:
|
| 145 |
+
if col in enriched.columns:
|
| 146 |
+
label_counts[col] = enriched[col].dropna().astype(str).value_counts().to_dict()
|
| 147 |
+
|
| 148 |
+
lines = [
|
| 149 |
+
"# ADNI Metadata Match Report",
|
| 150 |
+
"",
|
| 151 |
+
f"- PET/SUVR samples: {total}",
|
| 152 |
+
f"- Matched clinical rows: {int(enriched['clinical_label'].notna().sum())}",
|
| 153 |
+
"",
|
| 154 |
+
"## Field Coverage",
|
| 155 |
+
"",
|
| 156 |
+
"| field | non_missing | coverage |",
|
| 157 |
+
"|---|---:|---:|",
|
| 158 |
+
]
|
| 159 |
+
for row in rows:
|
| 160 |
+
lines.append(f"| {row['field']} | {row['non_missing']} | {row['coverage']:.3f} |")
|
| 161 |
+
lines.extend([
|
| 162 |
+
"",
|
| 163 |
+
"## Label Counts",
|
| 164 |
+
"",
|
| 165 |
+
])
|
| 166 |
+
for col, counts in label_counts.items():
|
| 167 |
+
lines.append(f"### {col}")
|
| 168 |
+
lines.append("")
|
| 169 |
+
for key, value in counts.items():
|
| 170 |
+
lines.append(f"- {key}: {value}")
|
| 171 |
+
lines.append("")
|
| 172 |
+
out_md.write_text("\n".join(lines), encoding="utf-8")
|
| 173 |
+
|
| 174 |
+
|
| 175 |
+
def main() -> None:
|
| 176 |
+
parser = argparse.ArgumentParser(description="Match ADNI Excel metadata to the PET/SUVR manifest.")
|
| 177 |
+
parser.add_argument("--manifest", type=Path, default=Path("data/metadata/splits/pet_fdg_manifest_with_split.csv"))
|
| 178 |
+
parser.add_argument("--full-xlsx", type=Path, default=Path("data/ADNIbase1416_info.xlsx"))
|
| 179 |
+
parser.add_argument("--conversion-xlsx", type=Path, default=Path("data/adni_1203s_info_fix.xlsx"))
|
| 180 |
+
parser.add_argument("--out", type=Path, default=Path("data/metadata/adni_matched_clinical.csv"))
|
| 181 |
+
parser.add_argument("--split-out-dir", type=Path, default=Path("data/metadata/splits"))
|
| 182 |
+
parser.add_argument("--report-md", type=Path, default=Path("data/metadata/adni_match_report.md"))
|
| 183 |
+
parser.add_argument("--report-csv", type=Path, default=Path("data/metadata/adni_match_report.csv"))
|
| 184 |
+
args = parser.parse_args()
|
| 185 |
+
|
| 186 |
+
manifest = pd.read_csv(args.manifest)
|
| 187 |
+
manifest["subject_id_norm"] = manifest["subject_id"].map(normalize_subject_id)
|
| 188 |
+
|
| 189 |
+
full = read_xlsx(args.full_xlsx, FULL_COLUMNS)
|
| 190 |
+
full = add_numeric_clean_columns(full)
|
| 191 |
+
full = full.rename(
|
| 192 |
+
columns={
|
| 193 |
+
"label": "clinical_label",
|
| 194 |
+
"DX": "dx",
|
| 195 |
+
"DX_bl": "dx_bl",
|
| 196 |
+
"AGE": "age",
|
| 197 |
+
"PTGENDER": "sex",
|
| 198 |
+
"PTEDUCAT": "education",
|
| 199 |
+
"PTETHCAT": "ethnicity",
|
| 200 |
+
"PTRACCAT": "race",
|
| 201 |
+
"PTMARRY": "marital_status",
|
| 202 |
+
"APOE4": "apoe4",
|
| 203 |
+
"FDG": "fdg_adni",
|
| 204 |
+
"PIB": "pib",
|
| 205 |
+
"AV45": "av45",
|
| 206 |
+
"ABETA": "abeta",
|
| 207 |
+
"TAU": "tau",
|
| 208 |
+
"PTAU": "ptau",
|
| 209 |
+
"ABETA_num": "abeta_num",
|
| 210 |
+
"TAU_num": "tau_num",
|
| 211 |
+
"PTAU_num": "ptau_num",
|
| 212 |
+
"CDRSB": "cdrsb",
|
| 213 |
+
"ADAS11": "adas11",
|
| 214 |
+
"ADAS13": "adas13",
|
| 215 |
+
"ADASQ4": "adasq4",
|
| 216 |
+
"MMSE": "mmse",
|
| 217 |
+
"RAVLT_immediate": "ravlt_immediate",
|
| 218 |
+
"RAVLT_learning": "ravlt_learning",
|
| 219 |
+
"RAVLT_forgetting": "ravlt_forgetting",
|
| 220 |
+
"RAVLT_perc_forgetting": "ravlt_perc_forgetting",
|
| 221 |
+
"LDELTOTAL": "ldeltotal",
|
| 222 |
+
"DIGITSCOR": "digitscor",
|
| 223 |
+
"TRABSCOR": "trabscor",
|
| 224 |
+
"FAQ": "faq",
|
| 225 |
+
"MOCA": "moca",
|
| 226 |
+
"Years_bl": "years_bl",
|
| 227 |
+
}
|
| 228 |
+
)
|
| 229 |
+
|
| 230 |
+
conversion = read_xlsx(args.conversion_xlsx, CONVERSION_COLUMNS)
|
| 231 |
+
conversion = conversion.rename(
|
| 232 |
+
columns={
|
| 233 |
+
"label": "conversion_label",
|
| 234 |
+
"DX.1": "dx_followup",
|
| 235 |
+
"mPACCdigit": "mpacc_digit",
|
| 236 |
+
"mPACCtrailsB": "mpacc_trailsb",
|
| 237 |
+
"Ventricles(心室)": "ventricles",
|
| 238 |
+
"Hippocampus(海马)": "hippocampus",
|
| 239 |
+
"WholeBrain(全脑)": "wholebrain",
|
| 240 |
+
"Entorhinal(内嗅觉)": "entorhinal",
|
| 241 |
+
"Fusiform(梭形)": "fusiform",
|
| 242 |
+
"MidTemp(中点温度)": "midtemp",
|
| 243 |
+
}
|
| 244 |
+
)
|
| 245 |
+
conversion = conversion.drop(columns=["PTID"], errors="ignore")
|
| 246 |
+
|
| 247 |
+
enriched = manifest.merge(full.drop(columns=["PTID"], errors="ignore"), on="subject_id_norm", how="left")
|
| 248 |
+
enriched = enriched.merge(conversion, on="subject_id_norm", how="left")
|
| 249 |
+
enriched = enriched.drop(columns=["subject_id_norm"])
|
| 250 |
+
|
| 251 |
+
args.out.parent.mkdir(parents=True, exist_ok=True)
|
| 252 |
+
enriched.to_csv(args.out, index=False)
|
| 253 |
+
write_split_manifests(enriched, args.split_out_dir)
|
| 254 |
+
summarize(enriched, args.report_md, args.report_csv)
|
| 255 |
+
|
| 256 |
+
print(f"wrote={args.out}")
|
| 257 |
+
print(f"wrote_splits={args.split_out_dir}/*_clinical.csv")
|
| 258 |
+
print(f"wrote_report={args.report_md}")
|
| 259 |
+
print(f"samples={len(enriched)} matched={int(enriched['clinical_label'].notna().sum())}")
|
| 260 |
+
|
| 261 |
+
|
| 262 |
+
if __name__ == "__main__":
|
| 263 |
+
main()
|
scripts/pet_vlm_dataset.py
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
|
| 5 |
+
import nibabel as nib
|
| 6 |
+
import numpy as np
|
| 7 |
+
import pandas as pd
|
| 8 |
+
import torch
|
| 9 |
+
import torch.nn.functional as F
|
| 10 |
+
from torch.utils.data import Dataset
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def normalize_pet(volume: np.ndarray, eps: float = 1e-6) -> np.ndarray:
|
| 14 |
+
volume = np.asarray(volume, dtype=np.float32)
|
| 15 |
+
finite = np.isfinite(volume)
|
| 16 |
+
if not finite.any():
|
| 17 |
+
return np.zeros_like(volume, dtype=np.float32)
|
| 18 |
+
lo, hi = np.percentile(volume[finite], [0.5, 99.5])
|
| 19 |
+
volume = np.clip(volume, lo, hi)
|
| 20 |
+
volume = (volume - lo) / max(float(hi - lo), eps)
|
| 21 |
+
return volume.astype(np.float32, copy=False)
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def resize_volume(volume: np.ndarray, output_size: tuple[int, int, int]) -> torch.Tensor:
|
| 25 |
+
tensor = torch.from_numpy(volume)[None, None]
|
| 26 |
+
tensor = F.interpolate(tensor, size=output_size, mode="trilinear", align_corners=False)
|
| 27 |
+
return tensor[0]
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def load_suvr_vector(csv_path: str | Path, include_background: bool = False) -> tuple[list[str], torch.Tensor]:
|
| 31 |
+
df = pd.read_csv(csv_path)
|
| 32 |
+
if not include_background:
|
| 33 |
+
df = df[df["label_name"] != "Background"].copy()
|
| 34 |
+
labels = df["label_name"].astype(str).tolist()
|
| 35 |
+
values = torch.tensor(df["mean_scalar"].astype(float).to_numpy(), dtype=torch.float32)
|
| 36 |
+
return labels, values
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def suvr_to_text(labels: list[str], values: torch.Tensor, top_k: int = 8) -> str:
|
| 40 |
+
pairs = sorted(zip(labels, values.tolist()), key=lambda x: x[1], reverse=True)
|
| 41 |
+
high = ", ".join(f"{name} {value:.3f}" for name, value in pairs[:top_k])
|
| 42 |
+
low = ", ".join(f"{name} {value:.3f}" for name, value in pairs[-top_k:])
|
| 43 |
+
mean_value = float(values.mean())
|
| 44 |
+
return f"FDG-PET regional SUVR summary. Mean SUVR {mean_value:.3f}. Highest regions: {high}. Lowest regions: {low}."
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
class PETSUVRDataset(Dataset):
|
| 48 |
+
def __init__(
|
| 49 |
+
self,
|
| 50 |
+
manifest_path: str | Path,
|
| 51 |
+
output_size: tuple[int, int, int] = (96, 96, 96),
|
| 52 |
+
include_background: bool = False,
|
| 53 |
+
) -> None:
|
| 54 |
+
self.manifest = pd.read_csv(manifest_path)
|
| 55 |
+
self.output_size = output_size
|
| 56 |
+
self.include_background = include_background
|
| 57 |
+
|
| 58 |
+
def __len__(self) -> int:
|
| 59 |
+
return len(self.manifest)
|
| 60 |
+
|
| 61 |
+
def __getitem__(self, index: int) -> dict[str, object]:
|
| 62 |
+
row = self.manifest.iloc[index]
|
| 63 |
+
volume = nib.load(str(row["pet_path"])).get_fdata(dtype=np.float32)
|
| 64 |
+
volume = normalize_pet(volume)
|
| 65 |
+
image = resize_volume(volume, self.output_size)
|
| 66 |
+
labels, suvr = load_suvr_vector(row["suvr_csv_path"], self.include_background)
|
| 67 |
+
text = suvr_to_text(labels, suvr)
|
| 68 |
+
return {
|
| 69 |
+
"sample_id": row["sample_id"],
|
| 70 |
+
"image": image,
|
| 71 |
+
"suvr": suvr,
|
| 72 |
+
"region_labels": labels,
|
| 73 |
+
"text": text,
|
| 74 |
+
}
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def collate_pet_suvr(batch: list[dict[str, object]]) -> dict[str, object]:
|
| 78 |
+
return {
|
| 79 |
+
"sample_id": [item["sample_id"] for item in batch],
|
| 80 |
+
"image": torch.stack([item["image"] for item in batch]),
|
| 81 |
+
"suvr": torch.stack([item["suvr"] for item in batch]),
|
| 82 |
+
"text": [item["text"] for item in batch],
|
| 83 |
+
"region_labels": batch[0]["region_labels"],
|
| 84 |
+
}
|
scripts/probe_mlp_remap.py
ADDED
|
@@ -0,0 +1,166 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Two-layer MLP probe for ReMAP-PET 3-way classification.
|
| 3 |
+
Compares linear vs non-linear probing on the same embeddings.
|
| 4 |
+
"""
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import argparse, json
|
| 8 |
+
from pathlib import Path
|
| 9 |
+
|
| 10 |
+
import numpy as np
|
| 11 |
+
import pandas as pd
|
| 12 |
+
import torch
|
| 13 |
+
from sklearn.neural_network import MLPClassifier
|
| 14 |
+
from sklearn.linear_model import LogisticRegression
|
| 15 |
+
from sklearn.preprocessing import LabelEncoder, StandardScaler
|
| 16 |
+
from sklearn.metrics import balanced_accuracy_score, roc_auc_score, accuracy_score
|
| 17 |
+
from sklearn.pipeline import make_pipeline
|
| 18 |
+
from torch.utils.data import DataLoader
|
| 19 |
+
|
| 20 |
+
from pet_vlm_dataset import PETSUVRDataset, collate_pet_suvr
|
| 21 |
+
from train_pet_foundation import PETSUVRFoundationModel, build_encoder
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def extract_embeddings(checkpoint, split_csv, output_size, device):
|
| 25 |
+
ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False)
|
| 26 |
+
saved_args = argparse.Namespace(**ckpt.get("args", {}))
|
| 27 |
+
for name in ("backbone", "embed_dim", "freeze_encoder"):
|
| 28 |
+
if not hasattr(saved_args, name) or getattr(saved_args, name) is None:
|
| 29 |
+
setattr(saved_args, name, ckpt.get("args", {}).get(name))
|
| 30 |
+
embed_dim = getattr(saved_args, "embed_dim", None) or 256
|
| 31 |
+
backbone = getattr(saved_args, "backbone", "medicalnet")
|
| 32 |
+
freeze = getattr(saved_args, "freeze_encoder", True)
|
| 33 |
+
|
| 34 |
+
manifest = pd.read_csv(split_csv)
|
| 35 |
+
dataset = PETSUVRDataset(split_csv, output_size=output_size)
|
| 36 |
+
loader = DataLoader(dataset, batch_size=4, shuffle=False, num_workers=2, collate_fn=collate_pet_suvr)
|
| 37 |
+
|
| 38 |
+
encoder = build_encoder(saved_args)
|
| 39 |
+
model = PETSUVRFoundationModel(encoder, int(dataset[0]["suvr"].numel()), embed_dim, freeze).to(device)
|
| 40 |
+
model.load_state_dict(ckpt["model"], strict=True)
|
| 41 |
+
model.eval()
|
| 42 |
+
|
| 43 |
+
pet_embs, suvr_preds = [], []
|
| 44 |
+
with torch.no_grad():
|
| 45 |
+
for batch in loader:
|
| 46 |
+
image = batch["image"].to(device)
|
| 47 |
+
suvr = batch["suvr"].to(device)
|
| 48 |
+
pet_feat = model.pet_encoder(image)
|
| 49 |
+
pet_z = torch.nn.functional.normalize(model.pet_projector(pet_feat), dim=-1)
|
| 50 |
+
pet_embs.append(pet_z.cpu().numpy())
|
| 51 |
+
suvr_preds.append(model(image, suvr)["pred_suvr"].cpu().numpy())
|
| 52 |
+
|
| 53 |
+
pet_z = np.concatenate(pet_embs, axis=0)
|
| 54 |
+
pred_suvr = np.concatenate(suvr_preds, axis=0)
|
| 55 |
+
return manifest, pet_z, pred_suvr
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def main():
|
| 59 |
+
parser = argparse.ArgumentParser()
|
| 60 |
+
parser.add_argument("--checkpoint", type=Path, required=True)
|
| 61 |
+
parser.add_argument("--train", type=Path, default=Path("data/metadata/splits/train_clinical.csv"))
|
| 62 |
+
parser.add_argument("--val", type=Path, default=Path("data/metadata/splits/val_clinical.csv"))
|
| 63 |
+
parser.add_argument("--test", type=Path, default=Path("data/metadata/splits/test_clinical.csv"))
|
| 64 |
+
parser.add_argument("--output-size", type=int, nargs=3, default=(96, 96, 96))
|
| 65 |
+
args = parser.parse_args()
|
| 66 |
+
|
| 67 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 68 |
+
print(f"Loading checkpoint: {args.checkpoint}")
|
| 69 |
+
|
| 70 |
+
train_df, train_pet, train_suvr = extract_embeddings(args.checkpoint, args.train, tuple(args.output_size), device)
|
| 71 |
+
val_df, val_pet, val_suvr = extract_embeddings(args.checkpoint, args.val, tuple(args.output_size), device)
|
| 72 |
+
test_df, test_pet, test_suvr = extract_embeddings(args.checkpoint, args.test, tuple(args.output_size), device)
|
| 73 |
+
|
| 74 |
+
# Extract labels for 3-way classification
|
| 75 |
+
encoder = LabelEncoder()
|
| 76 |
+
all_labels = pd.concat([train_df["clinical_label"], val_df["clinical_label"], test_df["clinical_label"]])
|
| 77 |
+
encoder.fit(all_labels.astype(str))
|
| 78 |
+
y_train = encoder.transform(train_df["clinical_label"].astype(str))
|
| 79 |
+
y_val = encoder.transform(val_df["clinical_label"].astype(str))
|
| 80 |
+
y_test = encoder.transform(test_df["clinical_label"].astype(str))
|
| 81 |
+
n_classes = len(encoder.classes_)
|
| 82 |
+
|
| 83 |
+
# ---- Linear probe (baseline) ----
|
| 84 |
+
print("\n=== Linear Probe (Logistic Regression) ===")
|
| 85 |
+
best_bal, best_c = 0, 0.01
|
| 86 |
+
for c in [0.01, 0.03, 0.1, 0.3, 1.0, 3.0, 10.0, 30.0]:
|
| 87 |
+
pipe = make_pipeline(StandardScaler(), LogisticRegression(C=c, max_iter=5000, solver="lbfgs"))
|
| 88 |
+
pipe.fit(train_pet, y_train)
|
| 89 |
+
val_pred = pipe.predict(val_pet)
|
| 90 |
+
bal = balanced_accuracy_score(y_val, val_pred)
|
| 91 |
+
print(f" C={c:.2f} val_bal={bal:.4f}")
|
| 92 |
+
if bal > best_bal:
|
| 93 |
+
best_bal, best_c = bal, c
|
| 94 |
+
|
| 95 |
+
pipe = make_pipeline(StandardScaler(), LogisticRegression(C=best_c, max_iter=5000, solver="lbfgs"))
|
| 96 |
+
pipe.fit(train_pet, y_train)
|
| 97 |
+
test_pred = pipe.predict(test_pet)
|
| 98 |
+
test_prob = pipe.predict_proba(test_pet)
|
| 99 |
+
linear_bal = balanced_accuracy_score(y_test, test_pred)
|
| 100 |
+
linear_auroc = roc_auc_score(y_test, test_prob, multi_class="ovr", average="macro")
|
| 101 |
+
print(f" Best C={best_c}, Test BalAcc={linear_bal:.4f}, AUROC={linear_auroc:.4f}")
|
| 102 |
+
|
| 103 |
+
# ---- MLP probe ----
|
| 104 |
+
print("\n=== MLP Probe (2-layer) ===")
|
| 105 |
+
best_bal2, best_alpha = 0, 0.01
|
| 106 |
+
for alpha in [0.0001, 0.001, 0.01, 0.1, 1.0]:
|
| 107 |
+
for hidden in [64, 128, 256]:
|
| 108 |
+
pipe2 = make_pipeline(
|
| 109 |
+
StandardScaler(),
|
| 110 |
+
MLPClassifier(hidden_layer_sizes=(hidden, hidden // 2),
|
| 111 |
+
alpha=alpha, max_iter=2000,
|
| 112 |
+
early_stopping=True, validation_fraction=0.1,
|
| 113 |
+
random_state=42)
|
| 114 |
+
)
|
| 115 |
+
pipe2.fit(train_pet, y_train)
|
| 116 |
+
val_pred2 = pipe2.predict(val_pet)
|
| 117 |
+
bal2 = balanced_accuracy_score(y_val, val_pred2)
|
| 118 |
+
if bal2 > best_bal2:
|
| 119 |
+
best_bal2, best_alpha, best_hidden = bal2, alpha, hidden
|
| 120 |
+
print(f" Best alpha={best_alpha}, hidden={best_hidden}, val_bal={best_bal2:.4f}")
|
| 121 |
+
|
| 122 |
+
pipe2 = make_pipeline(
|
| 123 |
+
StandardScaler(),
|
| 124 |
+
MLPClassifier(hidden_layer_sizes=(best_hidden, best_hidden // 2),
|
| 125 |
+
alpha=best_alpha, max_iter=5000, random_state=42)
|
| 126 |
+
)
|
| 127 |
+
pipe2.fit(train_pet, y_train)
|
| 128 |
+
test_pred2 = pipe2.predict(test_pet)
|
| 129 |
+
test_prob2 = pipe2.predict_proba(test_pet)
|
| 130 |
+
mlp_bal = balanced_accuracy_score(y_test, test_pred2)
|
| 131 |
+
mlp_auroc = roc_auc_score(y_test, test_prob2, multi_class="ovr", average="macro")
|
| 132 |
+
print(f" Test BalAcc={mlp_bal:.4f}, AUROC={mlp_auroc:.4f}")
|
| 133 |
+
|
| 134 |
+
# ---- Combined features probe ----
|
| 135 |
+
print("\n=== Linear Probe (PET + Predicted SUVR) ===")
|
| 136 |
+
train_comb = np.concatenate([train_pet, train_suvr], axis=1)
|
| 137 |
+
val_comb = np.concatenate([val_pet, val_suvr], axis=1)
|
| 138 |
+
test_comb = np.concatenate([test_pet, test_suvr], axis=1)
|
| 139 |
+
|
| 140 |
+
best_bal3, best_c3 = 0, 0.01
|
| 141 |
+
for c in [0.01, 0.03, 0.1, 0.3, 1.0, 3.0, 10.0, 30.0]:
|
| 142 |
+
pipe3 = make_pipeline(StandardScaler(), LogisticRegression(C=c, max_iter=5000, solver="lbfgs"))
|
| 143 |
+
pipe3.fit(train_comb, y_train)
|
| 144 |
+
val_pred3 = pipe3.predict(val_comb)
|
| 145 |
+
bal3 = balanced_accuracy_score(y_val, val_pred3)
|
| 146 |
+
if bal3 > best_bal3:
|
| 147 |
+
best_bal3, best_c3 = bal3, c
|
| 148 |
+
|
| 149 |
+
pipe3 = make_pipeline(StandardScaler(), LogisticRegression(C=best_c3, max_iter=5000, solver="lbfgs"))
|
| 150 |
+
pipe3.fit(train_comb, y_train)
|
| 151 |
+
test_pred3 = pipe3.predict(test_comb)
|
| 152 |
+
test_prob3 = pipe3.predict_proba(test_comb)
|
| 153 |
+
comb_bal = balanced_accuracy_score(y_test, test_pred3)
|
| 154 |
+
comb_auroc = roc_auc_score(y_test, test_prob3, multi_class="ovr", average="macro")
|
| 155 |
+
print(f" Best C={best_c3}, Test BalAcc={comb_bal:.4f}, AUROC={comb_auroc:.4f}")
|
| 156 |
+
|
| 157 |
+
# Summary
|
| 158 |
+
print(f"\n=== Summary (3-way test AUROC) ===")
|
| 159 |
+
print(f" Linear probe (original): {linear_auroc:.4f}")
|
| 160 |
+
print(f" MLP probe (2-layer): {mlp_auroc:.4f}")
|
| 161 |
+
print(f" Linear probe (PET+SUVR): {comb_auroc:.4f}")
|
| 162 |
+
print(f" MedicalNet frozen baseline: 0.7567")
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
if __name__ == "__main__":
|
| 166 |
+
main()
|
scripts/run_clinical_probes_v3.sh
ADDED
|
@@ -0,0 +1,38 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env bash
|
| 2 |
+
set -euo pipefail
|
| 3 |
+
cd /data/Albus/Brain
|
| 4 |
+
PY=/data/Albus/miniconda3/bin/python
|
| 5 |
+
mkdir -p runs/clinical logs
|
| 6 |
+
|
| 7 |
+
echo "[2026-05-20 23:23:42] start suvr_baseline" | tee logs/clinical_suvr_baseline.log
|
| 8 |
+
"" -u scripts/evaluate_suvr_clinical_probe.py --train data/metadata/splits/train_clinical_server.csv --val data/metadata/splits/val_clinical_server.csv --test data/metadata/splits/test_clinical_server.csv --out runs/clinical/suvr_clinical_probe.csv >> logs/clinical_suvr_baseline.log 2>&1
|
| 9 |
+
echo "[2026-05-20 23:23:42] done suvr_baseline" >> logs/clinical_suvr_baseline.log
|
| 10 |
+
|
| 11 |
+
TRAIN=data/metadata/splits/train_clinical_server.csv
|
| 12 |
+
VAL=data/metadata/splits/val_clinical_server.csv
|
| 13 |
+
TEST=data/metadata/splits/test_clinical_server.csv
|
| 14 |
+
|
| 15 |
+
run_probe() {
|
| 16 |
+
local gpu="$1"
|
| 17 |
+
local name="$2"
|
| 18 |
+
local ckpt="$3"
|
| 19 |
+
local log="logs/clinical_${name}.log"
|
| 20 |
+
local out="runs/clinical/${name}_clinical_probe.csv"
|
| 21 |
+
echo "[2026-05-20 23:23:42] start ${name} gpu=${gpu} ckpt=${ckpt}" | tee "${log}"
|
| 22 |
+
CUDA_VISIBLE_DEVICES="${gpu}" "${PY}" -u scripts/evaluate_pet_clinical_probe.py --checkpoint "${ckpt}" --train "${TRAIN}" --val "${VAL}" --test "${TEST}" --batch-size 4 --num-workers 2 --out "${out}" >> "${log}" 2>&1
|
| 23 |
+
echo "[2026-05-20 23:23:42] done ${name}" >> "${log}"
|
| 24 |
+
}
|
| 25 |
+
|
| 26 |
+
( run_probe 0 remap_pet runs/foundation/medicalnet_layer4_regalign_best.pt
|
| 27 |
+
run_probe 0 medicalnet_frozen runs/foundation/medicalnet_frozen_mlp.pt
|
| 28 |
+
run_probe 0 brainiac_frozen runs/foundation/brainiac_frozen_mlp.pt
|
| 29 |
+
) > logs/clinical_queue_gpu0.log 2>&1 &
|
| 30 |
+
echo $! > logs/clinical_queue_gpu0.pid
|
| 31 |
+
|
| 32 |
+
( run_probe 1 brainfm_frozen runs/foundation/brainfm_frozen_mlp_b4_best.pt
|
| 33 |
+
run_probe 1 sam_med3d_frozen runs/foundation/sam_med3d_frozen_mlp_best.pt
|
| 34 |
+
run_probe 1 swinunetr_frozen runs/foundation/swinunetr_frozen_mlp_best.pt
|
| 35 |
+
) > logs/clinical_queue_gpu1.log 2>&1 &
|
| 36 |
+
echo $! > logs/clinical_queue_gpu1.pid
|
| 37 |
+
|
| 38 |
+
echo "gpu0_pid=$(cat logs/clinical_queue_gpu0.pid) gpu1_pid=$(cat logs/clinical_queue_gpu1.pid)"
|
scripts/train_pet_foundation.py
ADDED
|
@@ -0,0 +1,531 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import argparse
|
| 4 |
+
import inspect
|
| 5 |
+
import importlib.util
|
| 6 |
+
from pathlib import Path
|
| 7 |
+
import sys
|
| 8 |
+
import types
|
| 9 |
+
|
| 10 |
+
import torch
|
| 11 |
+
from torch import nn
|
| 12 |
+
from torch.utils.data import DataLoader
|
| 13 |
+
|
| 14 |
+
from pet_vlm_dataset import PETSUVRDataset, collate_pet_suvr
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class MedicalNetBottleneck(nn.Module):
|
| 18 |
+
expansion = 4
|
| 19 |
+
|
| 20 |
+
def __init__(self, inplanes: int, planes: int, stride: int = 1, downsample: nn.Module | None = None) -> None:
|
| 21 |
+
super().__init__()
|
| 22 |
+
self.conv1 = nn.Conv3d(inplanes, planes, kernel_size=1, bias=False)
|
| 23 |
+
self.bn1 = nn.BatchNorm3d(planes)
|
| 24 |
+
self.conv2 = nn.Conv3d(planes, planes, kernel_size=3, stride=stride, padding=1, bias=False)
|
| 25 |
+
self.bn2 = nn.BatchNorm3d(planes)
|
| 26 |
+
self.conv3 = nn.Conv3d(planes, planes * self.expansion, kernel_size=1, bias=False)
|
| 27 |
+
self.bn3 = nn.BatchNorm3d(planes * self.expansion)
|
| 28 |
+
self.relu = nn.ReLU(inplace=True)
|
| 29 |
+
self.downsample = downsample
|
| 30 |
+
|
| 31 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 32 |
+
residual = x
|
| 33 |
+
out = self.relu(self.bn1(self.conv1(x)))
|
| 34 |
+
out = self.relu(self.bn2(self.conv2(out)))
|
| 35 |
+
out = self.bn3(self.conv3(out))
|
| 36 |
+
if self.downsample is not None:
|
| 37 |
+
residual = self.downsample(x)
|
| 38 |
+
out = self.relu(out + residual)
|
| 39 |
+
return out
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
class MedicalNetResNet50(nn.Module):
|
| 43 |
+
out_dim = 2048
|
| 44 |
+
|
| 45 |
+
def __init__(self) -> None:
|
| 46 |
+
super().__init__()
|
| 47 |
+
self.inplanes = 64
|
| 48 |
+
self.conv1 = nn.Conv3d(1, 64, kernel_size=7, stride=(2, 2, 2), padding=(3, 3, 3), bias=False)
|
| 49 |
+
self.bn1 = nn.BatchNorm3d(64)
|
| 50 |
+
self.relu = nn.ReLU(inplace=True)
|
| 51 |
+
self.maxpool = nn.MaxPool3d(kernel_size=3, stride=2, padding=1)
|
| 52 |
+
self.layer1 = self._make_layer(64, 3)
|
| 53 |
+
self.layer2 = self._make_layer(128, 4, stride=2)
|
| 54 |
+
self.layer3 = self._make_layer(256, 6, stride=2)
|
| 55 |
+
self.layer4 = self._make_layer(512, 3, stride=2)
|
| 56 |
+
self.pool = nn.AdaptiveAvgPool3d(1)
|
| 57 |
+
|
| 58 |
+
def _make_layer(self, planes: int, blocks: int, stride: int = 1) -> nn.Sequential:
|
| 59 |
+
downsample = None
|
| 60 |
+
if stride != 1 or self.inplanes != planes * MedicalNetBottleneck.expansion:
|
| 61 |
+
downsample = nn.Sequential(
|
| 62 |
+
nn.Conv3d(self.inplanes, planes * MedicalNetBottleneck.expansion, kernel_size=1, stride=stride, bias=False),
|
| 63 |
+
nn.BatchNorm3d(planes * MedicalNetBottleneck.expansion),
|
| 64 |
+
)
|
| 65 |
+
layers = [MedicalNetBottleneck(self.inplanes, planes, stride, downsample)]
|
| 66 |
+
self.inplanes = planes * MedicalNetBottleneck.expansion
|
| 67 |
+
for _ in range(1, blocks):
|
| 68 |
+
layers.append(MedicalNetBottleneck(self.inplanes, planes))
|
| 69 |
+
return nn.Sequential(*layers)
|
| 70 |
+
|
| 71 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 72 |
+
x = self.maxpool(self.relu(self.bn1(self.conv1(x))))
|
| 73 |
+
x = self.layer1(x)
|
| 74 |
+
x = self.layer2(x)
|
| 75 |
+
x = self.layer3(x)
|
| 76 |
+
x = self.layer4(x)
|
| 77 |
+
return self.pool(x).flatten(1)
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def load_medicalnet(path: Path) -> MedicalNetResNet50:
|
| 81 |
+
model = MedicalNetResNet50()
|
| 82 |
+
obj = torch.load(path, map_location="cpu")
|
| 83 |
+
state = obj.get("state_dict", obj)
|
| 84 |
+
state = {k.removeprefix("module."): v for k, v in state.items()}
|
| 85 |
+
missing, unexpected = model.load_state_dict(state, strict=False)
|
| 86 |
+
if unexpected:
|
| 87 |
+
print(f"unexpected_keys={unexpected[:8]}", flush=True)
|
| 88 |
+
if missing:
|
| 89 |
+
print(f"missing_keys={missing[:8]}", flush=True)
|
| 90 |
+
return model
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
class BrainIACEncoder(nn.Module):
|
| 94 |
+
out_dim = 768
|
| 95 |
+
|
| 96 |
+
def __init__(self, weights_path: Path) -> None:
|
| 97 |
+
super().__init__()
|
| 98 |
+
from monai.networks.nets import ViT
|
| 99 |
+
from safetensors.torch import load_file
|
| 100 |
+
|
| 101 |
+
self.model = ViT(
|
| 102 |
+
in_channels=1,
|
| 103 |
+
img_size=(96, 96, 96),
|
| 104 |
+
patch_size=(16, 16, 16),
|
| 105 |
+
hidden_size=768,
|
| 106 |
+
mlp_dim=3072,
|
| 107 |
+
num_layers=12,
|
| 108 |
+
num_heads=12,
|
| 109 |
+
)
|
| 110 |
+
weights = load_file(str(weights_path))
|
| 111 |
+
missing, unexpected = self.model.load_state_dict(weights, strict=False)
|
| 112 |
+
if unexpected:
|
| 113 |
+
print(f"brainiac_unexpected_keys={unexpected[:8]}", flush=True)
|
| 114 |
+
if missing:
|
| 115 |
+
print(f"brainiac_missing_keys={missing[:8]}", flush=True)
|
| 116 |
+
|
| 117 |
+
def forward(self, image: torch.Tensor) -> torch.Tensor:
|
| 118 |
+
output = self.model(image)
|
| 119 |
+
tokens = output[0] if isinstance(output, tuple) else output
|
| 120 |
+
return tokens[:, 0]
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
class SwinUNETREncoder(nn.Module):
|
| 124 |
+
out_dim = 768
|
| 125 |
+
|
| 126 |
+
def __init__(self, weights_path: Path, img_size: tuple[int, int, int]) -> None:
|
| 127 |
+
super().__init__()
|
| 128 |
+
from monai.networks.nets import SwinUNETR
|
| 129 |
+
|
| 130 |
+
kwargs = {
|
| 131 |
+
"in_channels": 1,
|
| 132 |
+
"out_channels": 2,
|
| 133 |
+
"feature_size": 48,
|
| 134 |
+
"use_checkpoint": False,
|
| 135 |
+
"spatial_dims": 3,
|
| 136 |
+
}
|
| 137 |
+
if "img_size" in inspect.signature(SwinUNETR).parameters:
|
| 138 |
+
kwargs["img_size"] = img_size
|
| 139 |
+
self.model = SwinUNETR(**kwargs)
|
| 140 |
+
weights = torch.load(weights_path, map_location="cpu", weights_only=False)
|
| 141 |
+
if hasattr(self.model, "load_from"):
|
| 142 |
+
self.model.load_from(weights)
|
| 143 |
+
else:
|
| 144 |
+
state = weights.get("state_dict", weights.get("model", weights))
|
| 145 |
+
remapped = {}
|
| 146 |
+
for key, value in state.items():
|
| 147 |
+
key = key.removeprefix("module.")
|
| 148 |
+
if key.startswith("encoder."):
|
| 149 |
+
key = "swinViT." + key[len("encoder.") :]
|
| 150 |
+
remapped[key] = value
|
| 151 |
+
missing, unexpected = self.model.load_state_dict(remapped, strict=False)
|
| 152 |
+
if unexpected:
|
| 153 |
+
print(f"swinunetr_unexpected_keys={unexpected[:8]}", flush=True)
|
| 154 |
+
if missing:
|
| 155 |
+
print(f"swinunetr_missing_keys={missing[:8]}", flush=True)
|
| 156 |
+
self.pool = nn.AdaptiveAvgPool3d(1)
|
| 157 |
+
|
| 158 |
+
def forward(self, image: torch.Tensor) -> torch.Tensor:
|
| 159 |
+
hidden = self.model.swinViT(image, self.model.normalize)
|
| 160 |
+
feat = hidden[-1]
|
| 161 |
+
return self.pool(feat).flatten(1)
|
| 162 |
+
|
| 163 |
+
|
| 164 |
+
class SAMMed3DEncoder(nn.Module):
|
| 165 |
+
out_dim = 384
|
| 166 |
+
|
| 167 |
+
def __init__(self, weights_path: Path) -> None:
|
| 168 |
+
super().__init__()
|
| 169 |
+
try:
|
| 170 |
+
import medim
|
| 171 |
+
except ImportError as exc:
|
| 172 |
+
raise ImportError("SAM-Med3D requires `medim`. Install it before using --backbone sam_med3d.") from exc
|
| 173 |
+
|
| 174 |
+
self.model = medim.create_model("SAM-Med3D", pretrained=True, checkpoint_path=str(weights_path))
|
| 175 |
+
self.image_encoder = getattr(self.model, "image_encoder", self.model)
|
| 176 |
+
self.pool = nn.AdaptiveAvgPool3d(1)
|
| 177 |
+
|
| 178 |
+
def forward(self, image: torch.Tensor) -> torch.Tensor:
|
| 179 |
+
output = self.image_encoder(image)
|
| 180 |
+
if isinstance(output, dict):
|
| 181 |
+
output = output.get("image_embeddings", output.get("embeddings", next(iter(output.values()))))
|
| 182 |
+
if isinstance(output, (list, tuple)):
|
| 183 |
+
output = output[0]
|
| 184 |
+
if output.ndim == 2:
|
| 185 |
+
return output
|
| 186 |
+
if output.ndim == 3:
|
| 187 |
+
return output.mean(dim=1)
|
| 188 |
+
return self.pool(output).flatten(1)
|
| 189 |
+
|
| 190 |
+
|
| 191 |
+
class BrainFMEncoder(nn.Module):
|
| 192 |
+
out_dim = 2048
|
| 193 |
+
|
| 194 |
+
def __init__(self, weights_path: Path, code_root: Path) -> None:
|
| 195 |
+
super().__init__()
|
| 196 |
+
unet_root = (code_root / "Trainer" / "models" / "unet3d").resolve()
|
| 197 |
+
package = types.ModuleType("brainfm_unet")
|
| 198 |
+
package.__path__ = [str(unet_root)]
|
| 199 |
+
sys.modules.setdefault("brainfm_unet", package)
|
| 200 |
+
spec = importlib.util.spec_from_file_location("brainfm_unet.model", unet_root / "model.py")
|
| 201 |
+
if spec is None or spec.loader is None:
|
| 202 |
+
raise RuntimeError(f"Could not load BrainFM model code from {unet_root}")
|
| 203 |
+
module = importlib.util.module_from_spec(spec)
|
| 204 |
+
sys.modules["brainfm_unet.model"] = module
|
| 205 |
+
spec.loader.exec_module(module)
|
| 206 |
+
|
| 207 |
+
ckpt = torch.load(weights_path, map_location="cpu", weights_only=False)
|
| 208 |
+
train_args = ckpt["train_args"]
|
| 209 |
+
self.model = module.UNet3D(
|
| 210 |
+
train_args.in_channels,
|
| 211 |
+
train_args.f_maps,
|
| 212 |
+
train_args.layer_order,
|
| 213 |
+
train_args.num_groups,
|
| 214 |
+
train_args.num_levels,
|
| 215 |
+
train_args.unit_feat,
|
| 216 |
+
)
|
| 217 |
+
state = {k.removeprefix("backbone."): v for k, v in ckpt["model"].items() if k.startswith("backbone.")}
|
| 218 |
+
missing, unexpected = self.model.load_state_dict(state, strict=False)
|
| 219 |
+
if unexpected:
|
| 220 |
+
print(f"brainfm_unexpected_keys={unexpected[:8]}", flush=True)
|
| 221 |
+
if missing:
|
| 222 |
+
print(f"brainfm_missing_keys={missing[:8]}", flush=True)
|
| 223 |
+
self.pool = nn.AdaptiveAvgPool3d(1)
|
| 224 |
+
|
| 225 |
+
def forward(self, image: torch.Tensor) -> torch.Tensor:
|
| 226 |
+
features = self.model.get_feature(image)
|
| 227 |
+
bottleneck = features[0]
|
| 228 |
+
return self.pool(bottleneck).flatten(1)
|
| 229 |
+
|
| 230 |
+
|
| 231 |
+
class Small3DPETEncoder(nn.Module):
|
| 232 |
+
out_dim = 256
|
| 233 |
+
|
| 234 |
+
def __init__(self) -> None:
|
| 235 |
+
super().__init__()
|
| 236 |
+
self.net = nn.Sequential(
|
| 237 |
+
nn.Conv3d(1, 16, 3, stride=2, padding=1),
|
| 238 |
+
nn.BatchNorm3d(16),
|
| 239 |
+
nn.GELU(),
|
| 240 |
+
nn.Conv3d(16, 32, 3, stride=2, padding=1),
|
| 241 |
+
nn.BatchNorm3d(32),
|
| 242 |
+
nn.GELU(),
|
| 243 |
+
nn.Conv3d(32, 64, 3, stride=2, padding=1),
|
| 244 |
+
nn.BatchNorm3d(64),
|
| 245 |
+
nn.GELU(),
|
| 246 |
+
nn.Conv3d(64, 128, 3, stride=2, padding=1),
|
| 247 |
+
nn.BatchNorm3d(128),
|
| 248 |
+
nn.GELU(),
|
| 249 |
+
nn.AdaptiveAvgPool3d(1),
|
| 250 |
+
)
|
| 251 |
+
self.proj = nn.Linear(128, self.out_dim)
|
| 252 |
+
|
| 253 |
+
def forward(self, image: torch.Tensor) -> torch.Tensor:
|
| 254 |
+
return self.proj(self.net(image).flatten(1))
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
class RegionSUVREncoder(nn.Module):
|
| 258 |
+
def __init__(self, n_regions: int, embed_dim: int) -> None:
|
| 259 |
+
super().__init__()
|
| 260 |
+
self.net = nn.Sequential(
|
| 261 |
+
nn.LayerNorm(n_regions),
|
| 262 |
+
nn.Linear(n_regions, embed_dim),
|
| 263 |
+
nn.GELU(),
|
| 264 |
+
nn.Linear(embed_dim, embed_dim),
|
| 265 |
+
)
|
| 266 |
+
|
| 267 |
+
def forward(self, suvr: torch.Tensor) -> torch.Tensor:
|
| 268 |
+
return self.net(suvr)
|
| 269 |
+
|
| 270 |
+
|
| 271 |
+
class PETSUVRFoundationModel(nn.Module):
|
| 272 |
+
def __init__(self, pet_encoder: nn.Module, n_regions: int, embed_dim: int = 256, freeze_encoder: bool = True) -> None:
|
| 273 |
+
super().__init__()
|
| 274 |
+
self.pet_encoder = pet_encoder
|
| 275 |
+
self.freeze_encoder = freeze_encoder
|
| 276 |
+
if freeze_encoder:
|
| 277 |
+
for p in self.pet_encoder.parameters():
|
| 278 |
+
p.requires_grad = False
|
| 279 |
+
self.pet_encoder.eval()
|
| 280 |
+
self.pet_projector = nn.Sequential(nn.LayerNorm(pet_encoder.out_dim), nn.Linear(pet_encoder.out_dim, embed_dim))
|
| 281 |
+
self.suvr_encoder = RegionSUVREncoder(n_regions, embed_dim)
|
| 282 |
+
self.suvr_head = nn.Sequential(nn.LayerNorm(embed_dim), nn.Linear(embed_dim, n_regions))
|
| 283 |
+
self.temperature = nn.Parameter(torch.tensor(0.07))
|
| 284 |
+
|
| 285 |
+
def forward(self, image: torch.Tensor, suvr: torch.Tensor) -> dict[str, torch.Tensor]:
|
| 286 |
+
if self.freeze_encoder:
|
| 287 |
+
with torch.no_grad():
|
| 288 |
+
pet_feat = self.pet_encoder(image)
|
| 289 |
+
else:
|
| 290 |
+
pet_feat = self.pet_encoder(image)
|
| 291 |
+
pet_z = nn.functional.normalize(self.pet_projector(pet_feat), dim=-1)
|
| 292 |
+
suvr_z = nn.functional.normalize(self.suvr_encoder(suvr), dim=-1)
|
| 293 |
+
pred_suvr = self.suvr_head(pet_z)
|
| 294 |
+
logits = pet_z @ suvr_z.T / self.temperature.clamp_min(0.01)
|
| 295 |
+
return {"logits": logits, "pred_suvr": pred_suvr}
|
| 296 |
+
|
| 297 |
+
|
| 298 |
+
def alignment_loss(
|
| 299 |
+
outputs: dict[str, torch.Tensor],
|
| 300 |
+
suvr: torch.Tensor,
|
| 301 |
+
contrastive_weight: float = 1.0,
|
| 302 |
+
regression_weight: float = 1.0,
|
| 303 |
+
) -> tuple[torch.Tensor, dict[str, float]]:
|
| 304 |
+
labels = torch.arange(suvr.shape[0], device=suvr.device)
|
| 305 |
+
loss_i = nn.functional.cross_entropy(outputs["logits"], labels)
|
| 306 |
+
loss_t = nn.functional.cross_entropy(outputs["logits"].T, labels)
|
| 307 |
+
loss_contrastive = 0.5 * (loss_i + loss_t)
|
| 308 |
+
loss_reg = nn.functional.mse_loss(outputs["pred_suvr"], suvr)
|
| 309 |
+
loss = contrastive_weight * loss_contrastive + regression_weight * loss_reg
|
| 310 |
+
return loss, {
|
| 311 |
+
"contrastive": float(loss_contrastive.detach()),
|
| 312 |
+
"regression": float(loss_reg.detach()),
|
| 313 |
+
}
|
| 314 |
+
|
| 315 |
+
|
| 316 |
+
def build_encoder(args: argparse.Namespace) -> nn.Module:
|
| 317 |
+
if args.backbone == "small_cnn":
|
| 318 |
+
return Small3DPETEncoder()
|
| 319 |
+
if args.backbone == "medicalnet":
|
| 320 |
+
return load_medicalnet(args.medicalnet_weights)
|
| 321 |
+
if args.backbone == "brainiac":
|
| 322 |
+
return BrainIACEncoder(args.brainiac_weights)
|
| 323 |
+
if args.backbone == "swinunetr":
|
| 324 |
+
return SwinUNETREncoder(args.swinunetr_weights, tuple(args.output_size))
|
| 325 |
+
if args.backbone == "sam_med3d":
|
| 326 |
+
return SAMMed3DEncoder(args.sam_med3d_weights)
|
| 327 |
+
if args.backbone == "brainfm":
|
| 328 |
+
return BrainFMEncoder(args.brainfm_weights, args.brainfm_code_root)
|
| 329 |
+
raise ValueError(f"Unsupported backbone: {args.backbone}")
|
| 330 |
+
|
| 331 |
+
|
| 332 |
+
def _set_trainable(module: nn.Module | None, trainable: bool) -> None:
|
| 333 |
+
if module is None:
|
| 334 |
+
return
|
| 335 |
+
for p in module.parameters():
|
| 336 |
+
p.requires_grad = trainable
|
| 337 |
+
|
| 338 |
+
|
| 339 |
+
def _last_vit_block(module: nn.Module) -> nn.Module | None:
|
| 340 |
+
blocks = getattr(module, "blocks", None)
|
| 341 |
+
if isinstance(blocks, (nn.ModuleList, list, tuple)) and len(blocks) > 0:
|
| 342 |
+
return blocks[-1]
|
| 343 |
+
return None
|
| 344 |
+
|
| 345 |
+
|
| 346 |
+
def _last_swin_stage(module: nn.Module) -> nn.Module | None:
|
| 347 |
+
swin = getattr(module, "swinViT", None)
|
| 348 |
+
if swin is None:
|
| 349 |
+
return None
|
| 350 |
+
for name in ("layers4", "layers3", "layers2", "layers1"):
|
| 351 |
+
layer = getattr(swin, name, None)
|
| 352 |
+
if layer is not None:
|
| 353 |
+
return layer
|
| 354 |
+
return None
|
| 355 |
+
|
| 356 |
+
|
| 357 |
+
def _brainfm_last_stage(module: nn.Module) -> nn.Module | None:
|
| 358 |
+
for name in ("encoders", "encoder", "down_path"):
|
| 359 |
+
stage = getattr(module, name, None)
|
| 360 |
+
if isinstance(stage, (nn.ModuleList, nn.Sequential)) and len(stage) > 0:
|
| 361 |
+
return stage[-1]
|
| 362 |
+
return None
|
| 363 |
+
|
| 364 |
+
|
| 365 |
+
def _sam_last_stage(module: nn.Module) -> nn.Module | None:
|
| 366 |
+
for name in ("blocks", "layers", "neck"):
|
| 367 |
+
stage = getattr(module, name, None)
|
| 368 |
+
if isinstance(stage, (nn.ModuleList, list, tuple)) and len(stage) > 0:
|
| 369 |
+
return stage[-1]
|
| 370 |
+
if isinstance(stage, nn.Module):
|
| 371 |
+
return stage
|
| 372 |
+
return None
|
| 373 |
+
|
| 374 |
+
|
| 375 |
+
def unfreeze_last_block(encoder: nn.Module) -> None:
|
| 376 |
+
_set_trainable(encoder, False)
|
| 377 |
+
target = None
|
| 378 |
+
if isinstance(encoder, MedicalNetResNet50):
|
| 379 |
+
target = encoder.layer4
|
| 380 |
+
elif isinstance(encoder, BrainIACEncoder):
|
| 381 |
+
target = _last_vit_block(encoder.model)
|
| 382 |
+
elif isinstance(encoder, SwinUNETREncoder):
|
| 383 |
+
target = _last_swin_stage(encoder.model)
|
| 384 |
+
elif isinstance(encoder, BrainFMEncoder):
|
| 385 |
+
target = _brainfm_last_stage(encoder.model)
|
| 386 |
+
elif isinstance(encoder, SAMMed3DEncoder):
|
| 387 |
+
target = _sam_last_stage(encoder.image_encoder)
|
| 388 |
+
if target is None:
|
| 389 |
+
raise ValueError(f"Could not identify a last block for {encoder.__class__.__name__}.")
|
| 390 |
+
_set_trainable(target, True)
|
| 391 |
+
|
| 392 |
+
|
| 393 |
+
def configure_encoder_training(model: PETSUVRFoundationModel, scope: str) -> None:
|
| 394 |
+
if scope == "none":
|
| 395 |
+
for p in model.pet_encoder.parameters():
|
| 396 |
+
p.requires_grad = False
|
| 397 |
+
return
|
| 398 |
+
if scope == "all":
|
| 399 |
+
for p in model.pet_encoder.parameters():
|
| 400 |
+
p.requires_grad = True
|
| 401 |
+
return
|
| 402 |
+
if scope == "layer4":
|
| 403 |
+
for p in model.pet_encoder.parameters():
|
| 404 |
+
p.requires_grad = False
|
| 405 |
+
if not hasattr(model.pet_encoder, "layer4"):
|
| 406 |
+
raise ValueError("encoder_train_scope=layer4 is only supported for MedicalNet-style encoders.")
|
| 407 |
+
for p in model.pet_encoder.layer4.parameters():
|
| 408 |
+
p.requires_grad = True
|
| 409 |
+
return
|
| 410 |
+
if scope == "last_block":
|
| 411 |
+
unfreeze_last_block(model.pet_encoder)
|
| 412 |
+
return
|
| 413 |
+
raise ValueError(f"Unsupported encoder training scope: {scope}")
|
| 414 |
+
|
| 415 |
+
|
| 416 |
+
def set_encoder_mode_for_scope(model: PETSUVRFoundationModel, scope: str) -> None:
|
| 417 |
+
if scope == "none":
|
| 418 |
+
model.pet_encoder.eval()
|
| 419 |
+
elif scope in {"layer4", "last_block"}:
|
| 420 |
+
model.pet_encoder.eval()
|
| 421 |
+
for module in model.pet_encoder.modules():
|
| 422 |
+
if any(p.requires_grad for p in module.parameters(recurse=False)):
|
| 423 |
+
module.train()
|
| 424 |
+
|
| 425 |
+
|
| 426 |
+
def main() -> None:
|
| 427 |
+
parser = argparse.ArgumentParser(description="Train PET-SUVR alignment with pretrained 3D backbones.")
|
| 428 |
+
parser.add_argument("--backbone", choices=["small_cnn", "medicalnet", "brainiac", "brainfm", "swinunetr", "sam_med3d"], default="medicalnet")
|
| 429 |
+
parser.add_argument("--medicalnet-weights", type=Path, default=Path("pretrained/medicalnet/resnet_50_23dataset.pth"))
|
| 430 |
+
parser.add_argument("--brainiac-weights", type=Path, default=Path("pretrained/brainiac/backbone.safetensors"))
|
| 431 |
+
parser.add_argument("--brainfm-weights", type=Path, default=Path("pretrained/brainfm/assets/brainfm_pretrained.pth"))
|
| 432 |
+
parser.add_argument("--brainfm-code-root", type=Path, default=Path("pretrained/brainfm"))
|
| 433 |
+
parser.add_argument("--swinunetr-weights", type=Path, default=Path("pretrained/swinunetr/model_swinvit.pt"))
|
| 434 |
+
parser.add_argument("--sam-med3d-weights", type=Path, default=Path("pretrained/sam-med3d/sam_med3d_turbo.pth"))
|
| 435 |
+
parser.add_argument("--manifest", type=Path, default=Path("metadata/splits/train.csv"))
|
| 436 |
+
parser.add_argument("--val-manifest", type=Path, default=Path("metadata/splits/val.csv"))
|
| 437 |
+
parser.add_argument("--epochs", type=int, default=10)
|
| 438 |
+
parser.add_argument("--batch-size", type=int, default=2)
|
| 439 |
+
parser.add_argument("--lr", type=float, default=1e-4)
|
| 440 |
+
parser.add_argument("--num-workers", type=int, default=2)
|
| 441 |
+
parser.add_argument("--output-size", type=int, nargs=3, default=(96, 96, 96))
|
| 442 |
+
parser.add_argument("--embed-dim", type=int, default=256)
|
| 443 |
+
parser.add_argument("--freeze-encoder", action=argparse.BooleanOptionalAction, default=True)
|
| 444 |
+
parser.add_argument("--encoder-train-scope", choices=["none", "layer4", "last_block", "all"], default=None)
|
| 445 |
+
parser.add_argument("--contrastive-weight", type=float, default=1.0)
|
| 446 |
+
parser.add_argument("--regression-weight", type=float, default=1.0)
|
| 447 |
+
parser.add_argument("--log-every", type=int, default=10)
|
| 448 |
+
parser.add_argument("--out", type=Path, default=Path("runs/foundation.pt"))
|
| 449 |
+
parser.add_argument("--best-out", type=Path, default=None)
|
| 450 |
+
args = parser.parse_args()
|
| 451 |
+
if args.encoder_train_scope is None:
|
| 452 |
+
args.encoder_train_scope = "none" if args.freeze_encoder else "all"
|
| 453 |
+
args.freeze_encoder = args.encoder_train_scope == "none"
|
| 454 |
+
|
| 455 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 456 |
+
train_dataset = PETSUVRDataset(args.manifest, output_size=tuple(args.output_size))
|
| 457 |
+
val_dataset = PETSUVRDataset(args.val_manifest, output_size=tuple(args.output_size))
|
| 458 |
+
train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers, collate_fn=collate_pet_suvr)
|
| 459 |
+
val_loader = DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers, collate_fn=collate_pet_suvr)
|
| 460 |
+
|
| 461 |
+
sample = train_dataset[0]
|
| 462 |
+
encoder = build_encoder(args)
|
| 463 |
+
model = PETSUVRFoundationModel(encoder, int(sample["suvr"].numel()), args.embed_dim, args.freeze_encoder).to(device)
|
| 464 |
+
configure_encoder_training(model, args.encoder_train_scope)
|
| 465 |
+
optimizer = torch.optim.AdamW((p for p in model.parameters() if p.requires_grad), lr=args.lr, weight_decay=1e-4)
|
| 466 |
+
|
| 467 |
+
best_val_loss = float("inf")
|
| 468 |
+
best_out = args.best_out or args.out.with_name(args.out.stem + "_best" + args.out.suffix)
|
| 469 |
+
print(
|
| 470 |
+
f"device={device} backbone={args.backbone} encoder_scope={args.encoder_train_scope} "
|
| 471 |
+
f"contrastive_weight={args.contrastive_weight} regression_weight={args.regression_weight} "
|
| 472 |
+
f"train={len(train_dataset)} val={len(val_dataset)}",
|
| 473 |
+
flush=True,
|
| 474 |
+
)
|
| 475 |
+
for epoch in range(1, args.epochs + 1):
|
| 476 |
+
model.train()
|
| 477 |
+
set_encoder_mode_for_scope(model, args.encoder_train_scope)
|
| 478 |
+
train_loss = 0.0
|
| 479 |
+
train_contrastive = 0.0
|
| 480 |
+
train_regression = 0.0
|
| 481 |
+
for step, batch in enumerate(train_loader, start=1):
|
| 482 |
+
image = batch["image"].to(device, non_blocking=True)
|
| 483 |
+
suvr = batch["suvr"].to(device, non_blocking=True)
|
| 484 |
+
outputs = model(image, suvr)
|
| 485 |
+
loss, parts = alignment_loss(outputs, suvr, args.contrastive_weight, args.regression_weight)
|
| 486 |
+
optimizer.zero_grad(set_to_none=True)
|
| 487 |
+
loss.backward()
|
| 488 |
+
optimizer.step()
|
| 489 |
+
train_loss += float(loss.detach()) * image.shape[0]
|
| 490 |
+
train_contrastive += parts["contrastive"] * image.shape[0]
|
| 491 |
+
train_regression += parts["regression"] * image.shape[0]
|
| 492 |
+
if args.log_every and step % args.log_every == 0:
|
| 493 |
+
print(f"epoch={epoch} step={step}/{len(train_loader)} loss={float(loss.detach()):.4f}", flush=True)
|
| 494 |
+
train_loss /= len(train_dataset)
|
| 495 |
+
train_contrastive /= len(train_dataset)
|
| 496 |
+
train_regression /= len(train_dataset)
|
| 497 |
+
|
| 498 |
+
model.eval()
|
| 499 |
+
val_loss = 0.0
|
| 500 |
+
val_contrastive = 0.0
|
| 501 |
+
val_regression = 0.0
|
| 502 |
+
with torch.no_grad():
|
| 503 |
+
for batch in val_loader:
|
| 504 |
+
image = batch["image"].to(device, non_blocking=True)
|
| 505 |
+
suvr = batch["suvr"].to(device, non_blocking=True)
|
| 506 |
+
loss, parts = alignment_loss(model(image, suvr), suvr, args.contrastive_weight, args.regression_weight)
|
| 507 |
+
val_loss += float(loss) * image.shape[0]
|
| 508 |
+
val_contrastive += parts["contrastive"] * image.shape[0]
|
| 509 |
+
val_regression += parts["regression"] * image.shape[0]
|
| 510 |
+
val_loss /= len(val_dataset)
|
| 511 |
+
val_contrastive /= len(val_dataset)
|
| 512 |
+
val_regression /= len(val_dataset)
|
| 513 |
+
print(
|
| 514 |
+
f"epoch={epoch} train_loss={train_loss:.4f} train_contrastive={train_contrastive:.4f} "
|
| 515 |
+
f"train_regression={train_regression:.4f} val_loss={val_loss:.4f} "
|
| 516 |
+
f"val_contrastive={val_contrastive:.4f} val_regression={val_regression:.4f}",
|
| 517 |
+
flush=True,
|
| 518 |
+
)
|
| 519 |
+
if val_loss < best_val_loss:
|
| 520 |
+
best_val_loss = val_loss
|
| 521 |
+
best_out.parent.mkdir(parents=True, exist_ok=True)
|
| 522 |
+
torch.save({"model": model.state_dict(), "args": vars(args), "best_val_loss": best_val_loss, "epoch": epoch}, best_out)
|
| 523 |
+
print(f"saved_best {best_out} val_loss={best_val_loss:.4f} epoch={epoch}", flush=True)
|
| 524 |
+
|
| 525 |
+
args.out.parent.mkdir(parents=True, exist_ok=True)
|
| 526 |
+
torch.save({"model": model.state_dict(), "args": vars(args), "best_val_loss": best_val_loss}, args.out)
|
| 527 |
+
print(f"saved {args.out}", flush=True)
|
| 528 |
+
|
| 529 |
+
|
| 530 |
+
if __name__ == "__main__":
|
| 531 |
+
main()
|
scripts/train_pet_foundation_epoch_ckpt.py
ADDED
|
@@ -0,0 +1,122 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
ReMAP-PET training with per-epoch checkpointing.
|
| 3 |
+
Saves a checkpoint every epoch for later clinical probe selection.
|
| 4 |
+
"""
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import argparse, inspect, importlib.util, sys, types
|
| 8 |
+
from pathlib import Path
|
| 9 |
+
|
| 10 |
+
import torch
|
| 11 |
+
from torch import nn
|
| 12 |
+
from torch.utils.data import DataLoader
|
| 13 |
+
|
| 14 |
+
from pet_vlm_dataset import PETSUVRDataset, collate_pet_suvr
|
| 15 |
+
from train_pet_foundation import (
|
| 16 |
+
PETSUVRFoundationModel, build_encoder, alignment_loss,
|
| 17 |
+
)
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def main():
|
| 21 |
+
parser = argparse.ArgumentParser()
|
| 22 |
+
parser.add_argument("--backbone", default="medicalnet")
|
| 23 |
+
parser.add_argument("--encoder-train-scope", default="layer4")
|
| 24 |
+
parser.add_argument("--epochs", type=int, default=50)
|
| 25 |
+
parser.add_argument("--batch-size", type=int, default=4)
|
| 26 |
+
parser.add_argument("--lr", type=float, default=1e-5)
|
| 27 |
+
parser.add_argument("--num-workers", type=int, default=2)
|
| 28 |
+
parser.add_argument("--output-size", type=int, nargs=3, default=(96, 96, 96))
|
| 29 |
+
parser.add_argument("--embed-dim", type=int, default=256)
|
| 30 |
+
parser.add_argument("--contrastive-weight", type=float, default=0.2)
|
| 31 |
+
parser.add_argument("--regression-weight", type=float, default=1.0)
|
| 32 |
+
parser.add_argument("--temperature", type=float, default=0.07)
|
| 33 |
+
parser.add_argument("--medicalnet-weights", type=Path, default=Path("pretrained/medicalnet/resnet_50_23dataset.pth"))
|
| 34 |
+
parser.add_argument("--train-csv", type=Path, default=Path("metadata/splits/train.csv"))
|
| 35 |
+
parser.add_argument("--val-csv", type=Path, default=Path("metadata/splits/val.csv"))
|
| 36 |
+
parser.add_argument("--out-dir", type=Path, default=Path("runs/foundation/remap_epochs"))
|
| 37 |
+
args = parser.parse_args()
|
| 38 |
+
|
| 39 |
+
args.out_dir.mkdir(parents=True, exist_ok=True)
|
| 40 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 41 |
+
|
| 42 |
+
train_set = PETSUVRDataset(args.train_csv, output_size=tuple(args.output_size))
|
| 43 |
+
val_set = PETSUVRDataset(args.val_csv, output_size=tuple(args.output_size))
|
| 44 |
+
train_loader = DataLoader(train_set, batch_size=args.batch_size, shuffle=True,
|
| 45 |
+
num_workers=args.num_workers, collate_fn=collate_pet_suvr)
|
| 46 |
+
val_loader = DataLoader(val_set, batch_size=args.batch_size, shuffle=False,
|
| 47 |
+
num_workers=args.num_workers, collate_fn=collate_pet_suvr)
|
| 48 |
+
n_regions = int(train_set[0]["suvr"].numel())
|
| 49 |
+
|
| 50 |
+
encoder = build_encoder(args)
|
| 51 |
+
model = PETSUVRFoundationModel(encoder, n_regions, args.embed_dim, False).to(device)
|
| 52 |
+
|
| 53 |
+
# Apply MedicalNet layer4 partial tuning
|
| 54 |
+
for name, param in model.pet_encoder.named_parameters():
|
| 55 |
+
param.requires_grad = ("layer4" in name)
|
| 56 |
+
|
| 57 |
+
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
| 58 |
+
total = sum(p.numel() for p in model.parameters())
|
| 59 |
+
print(f"Trainable params: {trainable:,} / {total:,} ({100*trainable/total:.1f}%)")
|
| 60 |
+
|
| 61 |
+
optimizer = torch.optim.AdamW(
|
| 62 |
+
[p for p in model.parameters() if p.requires_grad],
|
| 63 |
+
lr=args.lr, weight_decay=1e-4,
|
| 64 |
+
)
|
| 65 |
+
|
| 66 |
+
model.temperature.data.fill_(args.temperature)
|
| 67 |
+
best_val_loss = float("inf")
|
| 68 |
+
|
| 69 |
+
for epoch in range(1, args.epochs + 1):
|
| 70 |
+
model.train()
|
| 71 |
+
model.pet_encoder.train()
|
| 72 |
+
train_loss = 0.0
|
| 73 |
+
for batch in train_loader:
|
| 74 |
+
image = batch["image"].to(device, non_blocking=True)
|
| 75 |
+
suvr = batch["suvr"].to(device, non_blocking=True)
|
| 76 |
+
outputs = model(image, suvr)
|
| 77 |
+
loss, _ = alignment_loss(outputs, suvr, args.contrastive_weight, args.regression_weight)
|
| 78 |
+
optimizer.zero_grad(set_to_none=True)
|
| 79 |
+
loss.backward()
|
| 80 |
+
optimizer.step()
|
| 81 |
+
train_loss += float(loss) * image.shape[0]
|
| 82 |
+
|
| 83 |
+
train_loss /= len(train_set)
|
| 84 |
+
|
| 85 |
+
model.eval()
|
| 86 |
+
val_loss = 0.0
|
| 87 |
+
with torch.no_grad():
|
| 88 |
+
for batch in val_loader:
|
| 89 |
+
image = batch["image"].to(device, non_blocking=True)
|
| 90 |
+
suvr = batch["suvr"].to(device, non_blocking=True)
|
| 91 |
+
outputs = model(image, suvr)
|
| 92 |
+
loss, _ = alignment_loss(outputs, suvr, args.contrastive_weight, args.regression_weight)
|
| 93 |
+
val_loss += float(loss) * image.shape[0]
|
| 94 |
+
val_loss /= len(val_set)
|
| 95 |
+
|
| 96 |
+
print(f"epoch={epoch} train_loss={train_loss:.6f} val_loss={val_loss:.6f}", flush=True)
|
| 97 |
+
|
| 98 |
+
# Save every epoch
|
| 99 |
+
ckpt_path = args.out_dir / f"epoch_{epoch:02d}.pt"
|
| 100 |
+
torch.save({
|
| 101 |
+
"model": model.state_dict(),
|
| 102 |
+
"args": vars(args),
|
| 103 |
+
"epoch": epoch,
|
| 104 |
+
"val_loss": val_loss,
|
| 105 |
+
}, ckpt_path)
|
| 106 |
+
|
| 107 |
+
if val_loss < best_val_loss:
|
| 108 |
+
best_val_loss = val_loss
|
| 109 |
+
best_path = args.out_dir / "best.pt"
|
| 110 |
+
torch.save({
|
| 111 |
+
"model": model.state_dict(),
|
| 112 |
+
"args": vars(args),
|
| 113 |
+
"epoch": epoch,
|
| 114 |
+
"val_loss": val_loss,
|
| 115 |
+
}, best_path)
|
| 116 |
+
print(f" -> best (val_loss={val_loss:.6f})", flush=True)
|
| 117 |
+
|
| 118 |
+
print(f"Done. Saved {args.epochs} checkpoints to {args.out_dir}")
|
| 119 |
+
|
| 120 |
+
|
| 121 |
+
if __name__ == "__main__":
|
| 122 |
+
main()
|
scripts/train_pet_vlm_baseline.py
ADDED
|
@@ -0,0 +1,141 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import argparse
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
from torch import nn
|
| 8 |
+
from torch.utils.data import DataLoader
|
| 9 |
+
|
| 10 |
+
from pet_vlm_dataset import PETSUVRDataset, collate_pet_suvr
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class Small3DPETEncoder(nn.Module):
|
| 14 |
+
def __init__(self, embed_dim: int = 256) -> None:
|
| 15 |
+
super().__init__()
|
| 16 |
+
self.net = nn.Sequential(
|
| 17 |
+
nn.Conv3d(1, 16, 3, stride=2, padding=1),
|
| 18 |
+
nn.BatchNorm3d(16),
|
| 19 |
+
nn.GELU(),
|
| 20 |
+
nn.Conv3d(16, 32, 3, stride=2, padding=1),
|
| 21 |
+
nn.BatchNorm3d(32),
|
| 22 |
+
nn.GELU(),
|
| 23 |
+
nn.Conv3d(32, 64, 3, stride=2, padding=1),
|
| 24 |
+
nn.BatchNorm3d(64),
|
| 25 |
+
nn.GELU(),
|
| 26 |
+
nn.Conv3d(64, 128, 3, stride=2, padding=1),
|
| 27 |
+
nn.BatchNorm3d(128),
|
| 28 |
+
nn.GELU(),
|
| 29 |
+
nn.AdaptiveAvgPool3d(1),
|
| 30 |
+
)
|
| 31 |
+
self.proj = nn.Linear(128, embed_dim)
|
| 32 |
+
|
| 33 |
+
def forward(self, image: torch.Tensor) -> torch.Tensor:
|
| 34 |
+
x = self.net(image).flatten(1)
|
| 35 |
+
return self.proj(x)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
class RegionSUVREncoder(nn.Module):
|
| 39 |
+
def __init__(self, n_regions: int = 120, embed_dim: int = 256) -> None:
|
| 40 |
+
super().__init__()
|
| 41 |
+
self.net = nn.Sequential(
|
| 42 |
+
nn.LayerNorm(n_regions),
|
| 43 |
+
nn.Linear(n_regions, 256),
|
| 44 |
+
nn.GELU(),
|
| 45 |
+
nn.Linear(256, embed_dim),
|
| 46 |
+
)
|
| 47 |
+
|
| 48 |
+
def forward(self, suvr: torch.Tensor) -> torch.Tensor:
|
| 49 |
+
return self.net(suvr)
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class PETSUVRAlignmentModel(nn.Module):
|
| 53 |
+
def __init__(self, n_regions: int = 120, embed_dim: int = 256) -> None:
|
| 54 |
+
super().__init__()
|
| 55 |
+
self.pet_encoder = Small3DPETEncoder(embed_dim)
|
| 56 |
+
self.suvr_encoder = RegionSUVREncoder(n_regions, embed_dim)
|
| 57 |
+
self.suvr_head = nn.Linear(embed_dim, n_regions)
|
| 58 |
+
self.temperature = nn.Parameter(torch.tensor(0.07))
|
| 59 |
+
|
| 60 |
+
def forward(self, image: torch.Tensor, suvr: torch.Tensor) -> dict[str, torch.Tensor]:
|
| 61 |
+
pet_z = nn.functional.normalize(self.pet_encoder(image), dim=-1)
|
| 62 |
+
suvr_z = nn.functional.normalize(self.suvr_encoder(suvr), dim=-1)
|
| 63 |
+
pred_suvr = self.suvr_head(pet_z)
|
| 64 |
+
logits = pet_z @ suvr_z.T / self.temperature.clamp_min(0.01)
|
| 65 |
+
return {"pet_z": pet_z, "suvr_z": suvr_z, "pred_suvr": pred_suvr, "logits": logits}
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def alignment_loss(outputs: dict[str, torch.Tensor], suvr: torch.Tensor) -> torch.Tensor:
|
| 69 |
+
labels = torch.arange(suvr.shape[0], device=suvr.device)
|
| 70 |
+
loss_i = nn.functional.cross_entropy(outputs["logits"], labels)
|
| 71 |
+
loss_t = nn.functional.cross_entropy(outputs["logits"].T, labels)
|
| 72 |
+
loss_reg = nn.functional.mse_loss(outputs["pred_suvr"], suvr)
|
| 73 |
+
return 0.5 * (loss_i + loss_t) + loss_reg
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
def main() -> None:
|
| 77 |
+
parser = argparse.ArgumentParser(description="Train a minimal FDG-PET + SUVR alignment baseline.")
|
| 78 |
+
parser.add_argument("--manifest", type=Path, default=Path("metadata/splits/train.csv"))
|
| 79 |
+
parser.add_argument("--val-manifest", type=Path, default=Path("metadata/splits/val.csv"))
|
| 80 |
+
parser.add_argument("--epochs", type=int, default=2)
|
| 81 |
+
parser.add_argument("--batch-size", type=int, default=2)
|
| 82 |
+
parser.add_argument("--lr", type=float, default=1e-4)
|
| 83 |
+
parser.add_argument("--num-workers", type=int, default=0)
|
| 84 |
+
parser.add_argument("--output-size", type=int, nargs=3, default=(96, 96, 96))
|
| 85 |
+
parser.add_argument("--max-samples", type=int, default=0, help="Use the first N samples for a quick smoke test.")
|
| 86 |
+
parser.add_argument("--log-every", type=int, default=10)
|
| 87 |
+
parser.add_argument("--out", type=Path, default=Path("runs/pet_suvr_baseline.pt"))
|
| 88 |
+
args = parser.parse_args()
|
| 89 |
+
|
| 90 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 91 |
+
dataset = PETSUVRDataset(args.manifest, output_size=tuple(args.output_size))
|
| 92 |
+
val_dataset = PETSUVRDataset(args.val_manifest, output_size=tuple(args.output_size))
|
| 93 |
+
if args.max_samples > 0:
|
| 94 |
+
max_samples = min(args.max_samples, len(dataset))
|
| 95 |
+
dataset = torch.utils.data.Subset(dataset, range(max_samples))
|
| 96 |
+
val_max_samples = min(max(1, args.max_samples // 4), len(val_dataset))
|
| 97 |
+
val_dataset = torch.utils.data.Subset(val_dataset, range(val_max_samples))
|
| 98 |
+
train_loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers, collate_fn=collate_pet_suvr)
|
| 99 |
+
val_loader = DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers, collate_fn=collate_pet_suvr)
|
| 100 |
+
|
| 101 |
+
sample = dataset[0]
|
| 102 |
+
model = PETSUVRAlignmentModel(n_regions=int(sample["suvr"].numel())).to(device)
|
| 103 |
+
optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=1e-4)
|
| 104 |
+
|
| 105 |
+
for epoch in range(1, args.epochs + 1):
|
| 106 |
+
model.train()
|
| 107 |
+
train_loss = 0.0
|
| 108 |
+
for step, batch in enumerate(train_loader, start=1):
|
| 109 |
+
image = batch["image"].to(device)
|
| 110 |
+
suvr = batch["suvr"].to(device)
|
| 111 |
+
outputs = model(image, suvr)
|
| 112 |
+
loss = alignment_loss(outputs, suvr)
|
| 113 |
+
optimizer.zero_grad(set_to_none=True)
|
| 114 |
+
loss.backward()
|
| 115 |
+
optimizer.step()
|
| 116 |
+
train_loss += float(loss.detach()) * image.shape[0]
|
| 117 |
+
if args.log_every > 0 and step % args.log_every == 0:
|
| 118 |
+
print(
|
| 119 |
+
f"epoch={epoch} step={step}/{len(train_loader)} loss={float(loss.detach()):.4f}",
|
| 120 |
+
flush=True,
|
| 121 |
+
)
|
| 122 |
+
train_loss /= len(dataset)
|
| 123 |
+
|
| 124 |
+
model.eval()
|
| 125 |
+
val_loss = 0.0
|
| 126 |
+
with torch.no_grad():
|
| 127 |
+
for batch in val_loader:
|
| 128 |
+
image = batch["image"].to(device)
|
| 129 |
+
suvr = batch["suvr"].to(device)
|
| 130 |
+
outputs = model(image, suvr)
|
| 131 |
+
val_loss += float(alignment_loss(outputs, suvr)) * image.shape[0]
|
| 132 |
+
val_loss = val_loss / max(1, len(val_dataset))
|
| 133 |
+
print(f"epoch={epoch} train_loss={train_loss:.4f} val_loss={val_loss:.4f}")
|
| 134 |
+
|
| 135 |
+
args.out.parent.mkdir(parents=True, exist_ok=True)
|
| 136 |
+
torch.save({"model": model.state_dict(), "args": vars(args)}, args.out)
|
| 137 |
+
print(f"saved {args.out}", flush=True)
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
if __name__ == "__main__":
|
| 141 |
+
main()
|