DesonDai commited on
Commit
212e9d7
·
verified ·
1 Parent(s): 8231345

Add files using upload-large-folder tool

Browse files
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()