Add files using upload-large-folder tool
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +6 -0
- fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/code/_session_history.ipynb +1664 -0
- fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/config.yaml +49 -0
- fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/output.log +5 -0
- fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/wandb-metadata.json +144 -0
- fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/wandb-summary.json +1 -0
- fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug-core.log +12 -0
- fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug-internal.log +25 -0
- fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug.log +59 -0
- fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/run-HCPflat_raw_83810.wandb +0 -0
- fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/tmp/code/_session_history.ipynb +1664 -0
- fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/code/src/HCP_downstream_finetune.py +587 -0
- fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/output.log +2 -0
- fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/requirements.txt +198 -0
- fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/wandb-metadata.json +131 -0
- fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug-core.log +14 -0
- fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug-internal.log +11 -0
- fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug.log +26 -0
- fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/run-NSDflat_large_gsrFalse__HCP_FT_83810.wandb +0 -0
- fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/files/output.log +0 -0
- fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/files/requirements.txt +198 -0
- fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/files/wandb-metadata.json +144 -0
- fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug-core.log +12 -0
- fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug-internal.log +12 -0
- fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug.log +39 -0
- fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/run-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f.wandb +0 -0
- fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/logs/debug-internal.log +11 -0
- fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/logs/debug.log +25 -0
- fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/run-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3.wandb +3 -0
- fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/code/src/HCP_downstream_finetune.py +596 -0
- fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/output.log +59 -0
- fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/requirements.txt +198 -0
- fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/wandb-metadata.json +131 -0
- fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug-core.log +7 -0
- fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug-internal.log +11 -0
- fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug.log +25 -0
- fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/run-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532.wandb +3 -0
- fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/code/src/HCP_downstream_finetune.py +597 -0
- fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/output.log +23 -0
- fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/requirements.txt +198 -0
- fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/wandb-metadata.json +131 -0
- fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug-core.log +7 -0
- fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug-internal.log +11 -0
- fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug.log +25 -0
- fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/run-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427.wandb +3 -0
- fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/code/src/HCP_downstream_finetune.py +597 -0
- fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/output.log +0 -0
- fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/requirements.txt +198 -0
- fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/wandb-metadata.json +131 -0
- fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/logs/debug-core.log +21 -0
.gitattributes
CHANGED
|
@@ -5035,3 +5035,9 @@ fMRI-foundation-model/src/wandb/run-20241024_030627-HCPflat_large_gsrFalse__HCP_
|
|
| 5035 |
fMRI-foundation-model/src/wandb/run-20241024_160745-HCPflat_large_gsrFalse__HCP_FT_f1d2455a-d8d7-4d9d-9f91-58748aebb427/run-HCPflat_large_gsrFalse__HCP_FT_f1d2455a-d8d7-4d9d-9f91-58748aebb427.wandb filter=lfs diff=lfs merge=lfs -text
|
| 5036 |
fMRI-foundation-model/src/wandb/run-20241127_125303-HCPflat_raw_beta_trial_type_a0e1e642-966f-441b-9bc7-974dce26cba1/run-HCPflat_raw_beta_trial_type_a0e1e642-966f-441b-9bc7-974dce26cba1.wandb filter=lfs diff=lfs merge=lfs -text
|
| 5037 |
fMRI-foundation-model/src/wandb/run-20241127_023946-NSDflat_large_gsrFalse__beta_age_HCPFT_260e5584-8c2b-4e13-a5ed-11dcdc2a522f/run-NSDflat_large_gsrFalse__beta_age_HCPFT_260e5584-8c2b-4e13-a5ed-11dcdc2a522f.wandb filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5035 |
fMRI-foundation-model/src/wandb/run-20241024_160745-HCPflat_large_gsrFalse__HCP_FT_f1d2455a-d8d7-4d9d-9f91-58748aebb427/run-HCPflat_large_gsrFalse__HCP_FT_f1d2455a-d8d7-4d9d-9f91-58748aebb427.wandb filter=lfs diff=lfs merge=lfs -text
|
| 5036 |
fMRI-foundation-model/src/wandb/run-20241127_125303-HCPflat_raw_beta_trial_type_a0e1e642-966f-441b-9bc7-974dce26cba1/run-HCPflat_raw_beta_trial_type_a0e1e642-966f-441b-9bc7-974dce26cba1.wandb filter=lfs diff=lfs merge=lfs -text
|
| 5037 |
fMRI-foundation-model/src/wandb/run-20241127_023946-NSDflat_large_gsrFalse__beta_age_HCPFT_260e5584-8c2b-4e13-a5ed-11dcdc2a522f/run-NSDflat_large_gsrFalse__beta_age_HCPFT_260e5584-8c2b-4e13-a5ed-11dcdc2a522f.wandb filter=lfs diff=lfs merge=lfs -text
|
| 5038 |
+
fMRI-foundation-model/src/wandb/run-20241126_143710-HCPflat_large_gsrFalse__HCP_FT_79bf330c-a53f-43d5-86dd-b4bb676b9b78/run-HCPflat_large_gsrFalse__HCP_FT_79bf330c-a53f-43d5-86dd-b4bb676b9b78.wandb filter=lfs diff=lfs merge=lfs -text
|
| 5039 |
+
fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/run-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532.wandb filter=lfs diff=lfs merge=lfs -text
|
| 5040 |
+
fMRI-foundation-model/src/wandb/run-20241127_013141-HCPflat_raw_beta_age_18a4fb68-2904-438d-a472-1c5e5f991d72/run-HCPflat_raw_beta_age_18a4fb68-2904-438d-a472-1c5e5f991d72.wandb filter=lfs diff=lfs merge=lfs -text
|
| 5041 |
+
fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/run-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a.wandb filter=lfs diff=lfs merge=lfs -text
|
| 5042 |
+
fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/run-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427.wandb filter=lfs diff=lfs merge=lfs -text
|
| 5043 |
+
fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/run-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3.wandb filter=lfs diff=lfs merge=lfs -text
|
fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/code/_session_history.ipynb
ADDED
|
@@ -0,0 +1,1664 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"cells": [
|
| 3 |
+
{
|
| 4 |
+
"cell_type": "code",
|
| 5 |
+
"execution_count": 1,
|
| 6 |
+
"id": "406d87ac",
|
| 7 |
+
"metadata": {},
|
| 8 |
+
"outputs": [],
|
| 9 |
+
"source": [
|
| 10 |
+
"# Import packages and setup gpu configuration.\n",
|
| 11 |
+
"# This code block shouldnt need to be adjusted!\n",
|
| 12 |
+
"import os\n",
|
| 13 |
+
"import sys\n",
|
| 14 |
+
"import json\n",
|
| 15 |
+
"import yaml\n",
|
| 16 |
+
"import numpy as np\n",
|
| 17 |
+
"import copy\n",
|
| 18 |
+
"import math\n",
|
| 19 |
+
"import time\n",
|
| 20 |
+
"import random\n",
|
| 21 |
+
"from tqdm.auto import tqdm\n",
|
| 22 |
+
"import webdataset as wds\n",
|
| 23 |
+
"import matplotlib.pyplot as plt\n",
|
| 24 |
+
"\n",
|
| 25 |
+
"import torch\n",
|
| 26 |
+
"import torch.nn as nn\n",
|
| 27 |
+
"from torchvision import transforms\n",
|
| 28 |
+
"import utils\n",
|
| 29 |
+
"from mae_utils.flat_models import *\n",
|
| 30 |
+
"import h5py\n",
|
| 31 |
+
"\n",
|
| 32 |
+
"# tf32 data type is faster than standard float32\n",
|
| 33 |
+
"torch.backends.cuda.matmul.allow_tf32 = True\n",
|
| 34 |
+
"# following fixes a Conv3D CUDNN_NOT_SUPPORTED error\n",
|
| 35 |
+
"torch.backends.cudnn.benchmark = True\n",
|
| 36 |
+
"\n",
|
| 37 |
+
"# ## MODEL TO LOAD ##\n",
|
| 38 |
+
"model_name = \"HCPflat_large_gsrFalse_\"\n",
|
| 39 |
+
"parquet_folder = \"epoch99\"\n",
|
| 40 |
+
"\n",
|
| 41 |
+
"# outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 42 |
+
"outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 43 |
+
"\n",
|
| 44 |
+
"print(\"outdir\", outdir)\n",
|
| 45 |
+
"# Load previous config.yaml if available\n",
|
| 46 |
+
"if os.path.exists(f\"{outdir}/config.yaml\"):\n",
|
| 47 |
+
" config = yaml.load(open(f\"{outdir}/config.yaml\", 'r'), Loader=yaml.FullLoader)\n",
|
| 48 |
+
" print(f\"Loaded config.yaml from ckpt folder {outdir}\")\n",
|
| 49 |
+
" # create global variables from the config\n",
|
| 50 |
+
" print(\"\\n__CONFIG__\")\n",
|
| 51 |
+
" for attribute_name in config.keys():\n",
|
| 52 |
+
" print(f\"{attribute_name} = {config[attribute_name]}\")\n",
|
| 53 |
+
" globals()[attribute_name] = config[f'{attribute_name}']\n",
|
| 54 |
+
" print(\"\\n\")\n",
|
| 55 |
+
"\n",
|
| 56 |
+
"world_size = os.getenv('WORLD_SIZE')\n",
|
| 57 |
+
"if world_size is None: \n",
|
| 58 |
+
" world_size = 1\n",
|
| 59 |
+
"else:\n",
|
| 60 |
+
" world_size = int(world_size)\n",
|
| 61 |
+
"print(f\"WORLD_SIZE={world_size}\")\n",
|
| 62 |
+
"\n",
|
| 63 |
+
"if utils.is_interactive():\n",
|
| 64 |
+
" # Following allows you to change functions in models.py or utils.py and \n",
|
| 65 |
+
" # have this notebook automatically update with your revisions\n",
|
| 66 |
+
" %load_ext autoreload\n",
|
| 67 |
+
" %autoreload 2\n",
|
| 68 |
+
"\n",
|
| 69 |
+
"batch_size = probe_batch_size\n",
|
| 70 |
+
"num_epochs = probe_num_epochs\n",
|
| 71 |
+
"\n",
|
| 72 |
+
"data_type = torch.float32 # change depending on your mixed_precision\n",
|
| 73 |
+
"global_batch_size = batch_size * world_size\n",
|
| 74 |
+
"\n",
|
| 75 |
+
"device = torch.device('cuda')\n",
|
| 76 |
+
"\n",
|
| 77 |
+
"hcp_flat_path = \"/weka/proj-medarc/shared/HCP-Flat\"\n",
|
| 78 |
+
"# seed = 42\n",
|
| 79 |
+
"# num_frames = 16\n",
|
| 80 |
+
"# gsr = False\n",
|
| 81 |
+
"# num_workers = 10\n",
|
| 82 |
+
"# batch_size = 128\n",
|
| 83 |
+
"\n",
|
| 84 |
+
"print(\"PID of this process =\",os.getpid())\n",
|
| 85 |
+
"utils.seed_everything(seed)"
|
| 86 |
+
]
|
| 87 |
+
},
|
| 88 |
+
{
|
| 89 |
+
"cell_type": "code",
|
| 90 |
+
"execution_count": 2,
|
| 91 |
+
"id": "3f6365eb",
|
| 92 |
+
"metadata": {},
|
| 93 |
+
"outputs": [],
|
| 94 |
+
"source": [
|
| 95 |
+
"# Import packages and setup gpu configuration.\n",
|
| 96 |
+
"# This code block shouldnt need to be adjusted!\n",
|
| 97 |
+
"import os\n",
|
| 98 |
+
"import sys\n",
|
| 99 |
+
"import json\n",
|
| 100 |
+
"import yaml\n",
|
| 101 |
+
"import numpy as np\n",
|
| 102 |
+
"import copy\n",
|
| 103 |
+
"import math\n",
|
| 104 |
+
"import time\n",
|
| 105 |
+
"import random\n",
|
| 106 |
+
"from tqdm.auto import tqdm\n",
|
| 107 |
+
"import webdataset as wds\n",
|
| 108 |
+
"import matplotlib.pyplot as plt\n",
|
| 109 |
+
"\n",
|
| 110 |
+
"import torch\n",
|
| 111 |
+
"import torch.nn as nn\n",
|
| 112 |
+
"from torchvision import transforms\n",
|
| 113 |
+
"import utils\n",
|
| 114 |
+
"from mae_utils.flat_models import *\n",
|
| 115 |
+
"import h5py\n",
|
| 116 |
+
"\n",
|
| 117 |
+
"# tf32 data type is faster than standard float32\n",
|
| 118 |
+
"torch.backends.cuda.matmul.allow_tf32 = True\n",
|
| 119 |
+
"# following fixes a Conv3D CUDNN_NOT_SUPPORTED error\n",
|
| 120 |
+
"torch.backends.cudnn.benchmark = True\n",
|
| 121 |
+
"\n",
|
| 122 |
+
"# ## MODEL TO LOAD ##\n",
|
| 123 |
+
"model_name = \"HCPflat_large_gsrFalse_\"\n",
|
| 124 |
+
"parquet_folder = \"epoch99\"\n",
|
| 125 |
+
"\n",
|
| 126 |
+
"# outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 127 |
+
"outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 128 |
+
"\n",
|
| 129 |
+
"print(\"outdir\", outdir)\n",
|
| 130 |
+
"# Load previous config.yaml if available\n",
|
| 131 |
+
"if os.path.exists(f\"{outdir}/config.yaml\"):\n",
|
| 132 |
+
" config = yaml.load(open(f\"{outdir}/config.yaml\", 'r'), Loader=yaml.FullLoader)\n",
|
| 133 |
+
" print(f\"Loaded config.yaml from ckpt folder {outdir}\")\n",
|
| 134 |
+
" # create global variables from the config\n",
|
| 135 |
+
" print(\"\\n__CONFIG__\")\n",
|
| 136 |
+
" for attribute_name in config.keys():\n",
|
| 137 |
+
" print(f\"{attribute_name} = {config[attribute_name]}\")\n",
|
| 138 |
+
" globals()[attribute_name] = config[f'{attribute_name}']\n",
|
| 139 |
+
" print(\"\\n\")\n",
|
| 140 |
+
"\n",
|
| 141 |
+
"world_size = os.getenv('WORLD_SIZE')\n",
|
| 142 |
+
"if world_size is None: \n",
|
| 143 |
+
" world_size = 1\n",
|
| 144 |
+
"else:\n",
|
| 145 |
+
" world_size = int(world_size)\n",
|
| 146 |
+
"print(f\"WORLD_SIZE={world_size}\")\n",
|
| 147 |
+
"\n",
|
| 148 |
+
"if utils.is_interactive():\n",
|
| 149 |
+
" # Following allows you to change functions in models.py or utils.py and \n",
|
| 150 |
+
" # have this notebook automatically update with your revisions\n",
|
| 151 |
+
" %load_ext autoreload\n",
|
| 152 |
+
" %autoreload 2\n",
|
| 153 |
+
"\n",
|
| 154 |
+
"batch_size = probe_batch_size\n",
|
| 155 |
+
"num_epochs = probe_num_epochs\n",
|
| 156 |
+
"\n",
|
| 157 |
+
"data_type = torch.float32 # change depending on your mixed_precision\n",
|
| 158 |
+
"global_batch_size = batch_size * world_size\n",
|
| 159 |
+
"\n",
|
| 160 |
+
"device = torch.device('cuda')\n",
|
| 161 |
+
"\n",
|
| 162 |
+
"hcp_flat_path = \"/weka/proj-medarc/shared/HCP-Flat\"\n",
|
| 163 |
+
"# seed = 42\n",
|
| 164 |
+
"# num_frames = 16\n",
|
| 165 |
+
"# gsr = False\n",
|
| 166 |
+
"# num_workers = 10\n",
|
| 167 |
+
"# batch_size = 128\n",
|
| 168 |
+
"\n",
|
| 169 |
+
"print(\"PID of this process =\",os.getpid())\n",
|
| 170 |
+
"utils.seed_everything(seed)"
|
| 171 |
+
]
|
| 172 |
+
},
|
| 173 |
+
{
|
| 174 |
+
"cell_type": "code",
|
| 175 |
+
"execution_count": 3,
|
| 176 |
+
"id": "b96ed0fa",
|
| 177 |
+
"metadata": {},
|
| 178 |
+
"outputs": [],
|
| 179 |
+
"source": [
|
| 180 |
+
"from mae_utils.flat import load_hcp_flat_mask\n",
|
| 181 |
+
"from mae_utils.flat import create_hcp_flat\n",
|
| 182 |
+
"from mae_utils.flat import batch_unmask\n",
|
| 183 |
+
"import mae_utils.visualize as vis\n",
|
| 184 |
+
"\n",
|
| 185 |
+
"flat_mask = load_hcp_flat_mask(hcp_flat_path)\n",
|
| 186 |
+
"\n",
|
| 187 |
+
"model = flat_models.mae_vit_large_fmri(\n",
|
| 188 |
+
" patch_size=patch_size,\n",
|
| 189 |
+
" decoder_embed_dim=decoder_embed_dim,\n",
|
| 190 |
+
" t_patch_size=t_patch_size,\n",
|
| 191 |
+
" pred_t_dim=pred_t_dim,\n",
|
| 192 |
+
" decoder_depth=4,\n",
|
| 193 |
+
" cls_embed=cls_embed,\n",
|
| 194 |
+
" norm_pix_loss=norm_pix_loss,\n",
|
| 195 |
+
" no_qkv_bias=no_qkv_bias,\n",
|
| 196 |
+
" sep_pos_embed=sep_pos_embed,\n",
|
| 197 |
+
" trunc_init=trunc_init,\n",
|
| 198 |
+
" pct_masks_to_decode=pct_masks_to_decode,\n",
|
| 199 |
+
" img_mask=flat_mask,\n",
|
| 200 |
+
")"
|
| 201 |
+
]
|
| 202 |
+
},
|
| 203 |
+
{
|
| 204 |
+
"cell_type": "code",
|
| 205 |
+
"execution_count": 4,
|
| 206 |
+
"id": "2344601f",
|
| 207 |
+
"metadata": {},
|
| 208 |
+
"outputs": [],
|
| 209 |
+
"source": [
|
| 210 |
+
"from mae_utils.flat import load_hcp_flat_mask\n",
|
| 211 |
+
"from mae_utils.flat import create_hcp_flat\n",
|
| 212 |
+
"from mae_utils.flat import batch_unmask\n",
|
| 213 |
+
"import mae_utils.visualize as vis\n",
|
| 214 |
+
"\n",
|
| 215 |
+
"flat_mask = load_hcp_flat_mask(hcp_flat_path)\n",
|
| 216 |
+
"\n",
|
| 217 |
+
"model = flat_models.mae_vit_large_fmri(\n",
|
| 218 |
+
" patch_size=patch_size,\n",
|
| 219 |
+
" decoder_embed_dim=decoder_embed_dim,\n",
|
| 220 |
+
" t_patch_size=t_patch_size,\n",
|
| 221 |
+
" pred_t_dim=pred_t_dim,\n",
|
| 222 |
+
" decoder_depth=4,\n",
|
| 223 |
+
" cls_embed=cls_embed,\n",
|
| 224 |
+
" norm_pix_loss=norm_pix_loss,\n",
|
| 225 |
+
" no_qkv_bias=no_qkv_bias,\n",
|
| 226 |
+
" sep_pos_embed=sep_pos_embed,\n",
|
| 227 |
+
" trunc_init=trunc_init,\n",
|
| 228 |
+
" pct_masks_to_decode=pct_masks_to_decode,\n",
|
| 229 |
+
" img_mask=flat_mask,\n",
|
| 230 |
+
")"
|
| 231 |
+
]
|
| 232 |
+
},
|
| 233 |
+
{
|
| 234 |
+
"cell_type": "code",
|
| 235 |
+
"execution_count": 5,
|
| 236 |
+
"id": "33cf89e4",
|
| 237 |
+
"metadata": {},
|
| 238 |
+
"outputs": [],
|
| 239 |
+
"source": [
|
| 240 |
+
"# Import packages and setup gpu configuration.\n",
|
| 241 |
+
"# This code block shouldnt need to be adjusted!\n",
|
| 242 |
+
"import os\n",
|
| 243 |
+
"import sys\n",
|
| 244 |
+
"import json\n",
|
| 245 |
+
"import yaml\n",
|
| 246 |
+
"import numpy as np\n",
|
| 247 |
+
"import copy\n",
|
| 248 |
+
"import math\n",
|
| 249 |
+
"import time\n",
|
| 250 |
+
"import random\n",
|
| 251 |
+
"from tqdm.auto import tqdm\n",
|
| 252 |
+
"import webdataset as wds\n",
|
| 253 |
+
"import matplotlib.pyplot as plt\n",
|
| 254 |
+
"\n",
|
| 255 |
+
"import torch\n",
|
| 256 |
+
"import torch.nn as nn\n",
|
| 257 |
+
"from torchvision import transforms\n",
|
| 258 |
+
"import utils\n",
|
| 259 |
+
"from mae_utils.flat_models import *\n",
|
| 260 |
+
"import h5py\n",
|
| 261 |
+
"from mae_utils import flat_models\n",
|
| 262 |
+
"\n",
|
| 263 |
+
"# tf32 data type is faster than standard float32\n",
|
| 264 |
+
"torch.backends.cuda.matmul.allow_tf32 = True\n",
|
| 265 |
+
"# following fixes a Conv3D CUDNN_NOT_SUPPORTED error\n",
|
| 266 |
+
"torch.backends.cudnn.benchmark = True\n",
|
| 267 |
+
"\n",
|
| 268 |
+
"# ## MODEL TO LOAD ##\n",
|
| 269 |
+
"model_name = \"HCPflat_large_gsrFalse_\"\n",
|
| 270 |
+
"parquet_folder = \"epoch99\"\n",
|
| 271 |
+
"\n",
|
| 272 |
+
"# outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 273 |
+
"outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 274 |
+
"\n",
|
| 275 |
+
"print(\"outdir\", outdir)\n",
|
| 276 |
+
"# Load previous config.yaml if available\n",
|
| 277 |
+
"if os.path.exists(f\"{outdir}/config.yaml\"):\n",
|
| 278 |
+
" config = yaml.load(open(f\"{outdir}/config.yaml\", 'r'), Loader=yaml.FullLoader)\n",
|
| 279 |
+
" print(f\"Loaded config.yaml from ckpt folder {outdir}\")\n",
|
| 280 |
+
" # create global variables from the config\n",
|
| 281 |
+
" print(\"\\n__CONFIG__\")\n",
|
| 282 |
+
" for attribute_name in config.keys():\n",
|
| 283 |
+
" print(f\"{attribute_name} = {config[attribute_name]}\")\n",
|
| 284 |
+
" globals()[attribute_name] = config[f'{attribute_name}']\n",
|
| 285 |
+
" print(\"\\n\")\n",
|
| 286 |
+
"\n",
|
| 287 |
+
"world_size = os.getenv('WORLD_SIZE')\n",
|
| 288 |
+
"if world_size is None: \n",
|
| 289 |
+
" world_size = 1\n",
|
| 290 |
+
"else:\n",
|
| 291 |
+
" world_size = int(world_size)\n",
|
| 292 |
+
"print(f\"WORLD_SIZE={world_size}\")\n",
|
| 293 |
+
"\n",
|
| 294 |
+
"if utils.is_interactive():\n",
|
| 295 |
+
" # Following allows you to change functions in models.py or utils.py and \n",
|
| 296 |
+
" # have this notebook automatically update with your revisions\n",
|
| 297 |
+
" %load_ext autoreload\n",
|
| 298 |
+
" %autoreload 2\n",
|
| 299 |
+
"\n",
|
| 300 |
+
"batch_size = probe_batch_size\n",
|
| 301 |
+
"num_epochs = probe_num_epochs\n",
|
| 302 |
+
"\n",
|
| 303 |
+
"data_type = torch.float32 # change depending on your mixed_precision\n",
|
| 304 |
+
"global_batch_size = batch_size * world_size\n",
|
| 305 |
+
"\n",
|
| 306 |
+
"device = torch.device('cuda')\n",
|
| 307 |
+
"\n",
|
| 308 |
+
"hcp_flat_path = \"/weka/proj-medarc/shared/HCP-Flat\"\n",
|
| 309 |
+
"# seed = 42\n",
|
| 310 |
+
"# num_frames = 16\n",
|
| 311 |
+
"# gsr = False\n",
|
| 312 |
+
"# num_workers = 10\n",
|
| 313 |
+
"# batch_size = 128\n",
|
| 314 |
+
"\n",
|
| 315 |
+
"print(\"PID of this process =\",os.getpid())\n",
|
| 316 |
+
"utils.seed_everything(seed)"
|
| 317 |
+
]
|
| 318 |
+
},
|
| 319 |
+
{
|
| 320 |
+
"cell_type": "code",
|
| 321 |
+
"execution_count": 6,
|
| 322 |
+
"id": "bc2281a4",
|
| 323 |
+
"metadata": {},
|
| 324 |
+
"outputs": [],
|
| 325 |
+
"source": [
|
| 326 |
+
"from mae_utils.flat import load_hcp_flat_mask\n",
|
| 327 |
+
"from mae_utils.flat import create_hcp_flat\n",
|
| 328 |
+
"from mae_utils.flat import batch_unmask\n",
|
| 329 |
+
"import mae_utils.visualize as vis\n",
|
| 330 |
+
"\n",
|
| 331 |
+
"flat_mask = load_hcp_flat_mask(hcp_flat_path)\n",
|
| 332 |
+
"\n",
|
| 333 |
+
"model = flat_models.mae_vit_large_fmri(\n",
|
| 334 |
+
" patch_size=patch_size,\n",
|
| 335 |
+
" decoder_embed_dim=decoder_embed_dim,\n",
|
| 336 |
+
" t_patch_size=t_patch_size,\n",
|
| 337 |
+
" pred_t_dim=pred_t_dim,\n",
|
| 338 |
+
" decoder_depth=4,\n",
|
| 339 |
+
" cls_embed=cls_embed,\n",
|
| 340 |
+
" norm_pix_loss=norm_pix_loss,\n",
|
| 341 |
+
" no_qkv_bias=no_qkv_bias,\n",
|
| 342 |
+
" sep_pos_embed=sep_pos_embed,\n",
|
| 343 |
+
" trunc_init=trunc_init,\n",
|
| 344 |
+
" pct_masks_to_decode=pct_masks_to_decode,\n",
|
| 345 |
+
" img_mask=flat_mask,\n",
|
| 346 |
+
")"
|
| 347 |
+
]
|
| 348 |
+
},
|
| 349 |
+
{
|
| 350 |
+
"cell_type": "code",
|
| 351 |
+
"execution_count": 7,
|
| 352 |
+
"id": "acfaeaad",
|
| 353 |
+
"metadata": {},
|
| 354 |
+
"outputs": [],
|
| 355 |
+
"source": [
|
| 356 |
+
"checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')]\n",
|
| 357 |
+
"\n",
|
| 358 |
+
"if utils.is_interactive():\n",
|
| 359 |
+
" latest_checkpoint = \"epoch99.pth\"\n",
|
| 360 |
+
"else:\n",
|
| 361 |
+
" latest_checkpoint = sys.argv[2] \n",
|
| 362 |
+
"print(f\"latest_checkpoint: {latest_checkpoint}\")\n",
|
| 363 |
+
"\n",
|
| 364 |
+
"# Load the checkpoint\n",
|
| 365 |
+
"checkpoint_path = os.path.join(outdir, latest_checkpoint)\n",
|
| 366 |
+
"\n",
|
| 367 |
+
"state = torch.load(checkpoint_path)\n",
|
| 368 |
+
"model.load_state_dict(state[\"model_state_dict\"], strict=False)\n",
|
| 369 |
+
"model.to(device)\n",
|
| 370 |
+
"model.eval()\n",
|
| 371 |
+
"\n",
|
| 372 |
+
"print(f\"\\nLoaded checkpoint {latest_checkpoint} from {outdir}\\n\")"
|
| 373 |
+
]
|
| 374 |
+
},
|
| 375 |
+
{
|
| 376 |
+
"cell_type": "code",
|
| 377 |
+
"execution_count": 8,
|
| 378 |
+
"id": "8ffbe1b4",
|
| 379 |
+
"metadata": {},
|
| 380 |
+
"outputs": [],
|
| 381 |
+
"source": [
|
| 382 |
+
"f_train = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/train_hcp.hdf5', 'r')\n",
|
| 383 |
+
"flatmaps_train = f_train['flatmaps']\n",
|
| 384 |
+
"\n",
|
| 385 |
+
"f_test = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/test_hcp.hdf5', 'r')\n",
|
| 386 |
+
"flatmaps_test = f_test['flatmaps']\n",
|
| 387 |
+
"\n",
|
| 388 |
+
"metadata_train = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_train_HCP.npy', allow_pickle=True)\n",
|
| 389 |
+
"metadata_test = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_test_HCP.npy', allow_pickle=True)"
|
| 390 |
+
]
|
| 391 |
+
},
|
| 392 |
+
{
|
| 393 |
+
"cell_type": "code",
|
| 394 |
+
"execution_count": 9,
|
| 395 |
+
"id": "767e8c90",
|
| 396 |
+
"metadata": {},
|
| 397 |
+
"outputs": [],
|
| 398 |
+
"source": [
|
| 399 |
+
"from mae_utils.flat import load_hcp_flat_mask\n",
|
| 400 |
+
"from mae_utils.flat import create_hcp_flat\n",
|
| 401 |
+
"from mae_utils.flat import batch_unmask\n",
|
| 402 |
+
"import mae_utils.visualize as vis\n",
|
| 403 |
+
"\n",
|
| 404 |
+
"flat_mask = load_hcp_flat_mask(hcp_flat_path)\n",
|
| 405 |
+
"\n",
|
| 406 |
+
"mae_model = flat_models.mae_vit_large_fmri(\n",
|
| 407 |
+
" patch_size=patch_size,\n",
|
| 408 |
+
" decoder_embed_dim=decoder_embed_dim,\n",
|
| 409 |
+
" t_patch_size=t_patch_size,\n",
|
| 410 |
+
" pred_t_dim=pred_t_dim,\n",
|
| 411 |
+
" decoder_depth=4,\n",
|
| 412 |
+
" cls_embed=cls_embed,\n",
|
| 413 |
+
" norm_pix_loss=norm_pix_loss,\n",
|
| 414 |
+
" no_qkv_bias=no_qkv_bias,\n",
|
| 415 |
+
" sep_pos_embed=sep_pos_embed,\n",
|
| 416 |
+
" trunc_init=trunc_init,\n",
|
| 417 |
+
" pct_masks_to_decode=pct_masks_to_decode,\n",
|
| 418 |
+
" img_mask=flat_mask,\n",
|
| 419 |
+
")"
|
| 420 |
+
]
|
| 421 |
+
},
|
| 422 |
+
{
|
| 423 |
+
"cell_type": "code",
|
| 424 |
+
"execution_count": 10,
|
| 425 |
+
"id": "fa59d7d5",
|
| 426 |
+
"metadata": {},
|
| 427 |
+
"outputs": [],
|
| 428 |
+
"source": [
|
| 429 |
+
"checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')]\n",
|
| 430 |
+
"\n",
|
| 431 |
+
"if utils.is_interactive():\n",
|
| 432 |
+
" latest_checkpoint = \"epoch99.pth\"\n",
|
| 433 |
+
"else:\n",
|
| 434 |
+
" latest_checkpoint = sys.argv[2] \n",
|
| 435 |
+
"print(f\"latest_checkpoint: {latest_checkpoint}\")\n",
|
| 436 |
+
"\n",
|
| 437 |
+
"# Load the checkpoint\n",
|
| 438 |
+
"checkpoint_path = os.path.join(outdir, latest_checkpoint)\n",
|
| 439 |
+
"\n",
|
| 440 |
+
"state = torch.load(checkpoint_path)\n",
|
| 441 |
+
"mae_model.load_state_dict(state[\"model_state_dict\"], strict=False)\n",
|
| 442 |
+
"mae_model.to(device)\n",
|
| 443 |
+
"\n",
|
| 444 |
+
"print(f\"\\nLoaded checkpoint {latest_checkpoint} from {outdir}\\n\")"
|
| 445 |
+
]
|
| 446 |
+
},
|
| 447 |
+
{
|
| 448 |
+
"cell_type": "code",
|
| 449 |
+
"execution_count": 11,
|
| 450 |
+
"id": "cea19e70",
|
| 451 |
+
"metadata": {},
|
| 452 |
+
"outputs": [],
|
| 453 |
+
"source": [
|
| 454 |
+
"f_train = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/train_hcp.hdf5', 'r')\n",
|
| 455 |
+
"flatmaps_train = f_train['flatmaps']\n",
|
| 456 |
+
"\n",
|
| 457 |
+
"f_test = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/test_hcp.hdf5', 'r')\n",
|
| 458 |
+
"flatmaps_test = f_test['flatmaps']\n",
|
| 459 |
+
"\n",
|
| 460 |
+
"metadata_train = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_train_HCP.npy', allow_pickle=True)\n",
|
| 461 |
+
"metadata_test = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_test_HCP.npy', allow_pickle=True)"
|
| 462 |
+
]
|
| 463 |
+
},
|
| 464 |
+
{
|
| 465 |
+
"cell_type": "code",
|
| 466 |
+
"execution_count": 12,
|
| 467 |
+
"id": "f1f85737",
|
| 468 |
+
"metadata": {},
|
| 469 |
+
"outputs": [],
|
| 470 |
+
"source": [
|
| 471 |
+
"from torch.utils.data import Dataset, DataLoader\n",
|
| 472 |
+
"\n",
|
| 473 |
+
"class HCPFlatDataset(Dataset):\n",
|
| 474 |
+
" def __init__(self, flatmaps, metadata):\n",
|
| 475 |
+
" self.flatmaps = flatmaps\n",
|
| 476 |
+
" self.metadata = metadata\n",
|
| 477 |
+
"\n",
|
| 478 |
+
" def __len__(self):\n",
|
| 479 |
+
" return len(self.metadata)\n",
|
| 480 |
+
"\n",
|
| 481 |
+
" def __getitem__(self, idx):\n",
|
| 482 |
+
" return self.flatmaps[idx], json.loads(self.metadata[idx])\n",
|
| 483 |
+
"\n",
|
| 484 |
+
"# Loading to cpu for faster training, this can take several minutes. Remove this [:] if you want to move one at the time.\n",
|
| 485 |
+
"train_dataset = HCPFlatDataset(flatmaps_train, metadata_train)\n",
|
| 486 |
+
"train_dl = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)\n",
|
| 487 |
+
"\n",
|
| 488 |
+
"test_dataset = HCPFlatDataset(flatmaps_test, metadata_test)\n",
|
| 489 |
+
"test_dl = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)"
|
| 490 |
+
]
|
| 491 |
+
},
|
| 492 |
+
{
|
| 493 |
+
"cell_type": "code",
|
| 494 |
+
"execution_count": 13,
|
| 495 |
+
"id": "f80c3176",
|
| 496 |
+
"metadata": {},
|
| 497 |
+
"outputs": [],
|
| 498 |
+
"source": [
|
| 499 |
+
"for i in train_dl:\n",
|
| 500 |
+
" break"
|
| 501 |
+
]
|
| 502 |
+
},
|
| 503 |
+
{
|
| 504 |
+
"cell_type": "code",
|
| 505 |
+
"execution_count": 14,
|
| 506 |
+
"id": "eb364065",
|
| 507 |
+
"metadata": {},
|
| 508 |
+
"outputs": [],
|
| 509 |
+
"source": [
|
| 510 |
+
"for i in test_dl:\n",
|
| 511 |
+
" break"
|
| 512 |
+
]
|
| 513 |
+
},
|
| 514 |
+
{
|
| 515 |
+
"cell_type": "code",
|
| 516 |
+
"execution_count": 15,
|
| 517 |
+
"id": "d7581eea",
|
| 518 |
+
"metadata": {},
|
| 519 |
+
"outputs": [
|
| 520 |
+
{
|
| 521 |
+
"name": "stdout",
|
| 522 |
+
"output_type": "stream",
|
| 523 |
+
"text": [
|
| 524 |
+
"[tensor([[[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 525 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 526 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 527 |
+
" ...,\n",
|
| 528 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 529 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 530 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 531 |
+
" \n",
|
| 532 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 533 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 534 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 535 |
+
" ...,\n",
|
| 536 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 537 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 538 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 539 |
+
" \n",
|
| 540 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 541 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 542 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 543 |
+
" ...,\n",
|
| 544 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 545 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 546 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 547 |
+
" \n",
|
| 548 |
+
" ...,\n",
|
| 549 |
+
" \n",
|
| 550 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 551 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 552 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 553 |
+
" ...,\n",
|
| 554 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 555 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 556 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 557 |
+
" \n",
|
| 558 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 559 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 560 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 561 |
+
" ...,\n",
|
| 562 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 563 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 564 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 565 |
+
" \n",
|
| 566 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 567 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 568 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 569 |
+
" ...,\n",
|
| 570 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 571 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 572 |
+
" [0., 0., 0., ..., 0., 0., 0.]]],\n",
|
| 573 |
+
" \n",
|
| 574 |
+
" \n",
|
| 575 |
+
" [[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 576 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 577 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 578 |
+
" ...,\n",
|
| 579 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 580 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 581 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 582 |
+
" \n",
|
| 583 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 584 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 585 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 586 |
+
" ...,\n",
|
| 587 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 588 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 589 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 590 |
+
" \n",
|
| 591 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 592 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 593 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 594 |
+
" ...,\n",
|
| 595 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 596 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 597 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 598 |
+
" \n",
|
| 599 |
+
" ...,\n",
|
| 600 |
+
" \n",
|
| 601 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 602 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 603 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 604 |
+
" ...,\n",
|
| 605 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 606 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 607 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 608 |
+
" \n",
|
| 609 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 610 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 611 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 612 |
+
" ...,\n",
|
| 613 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 614 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 615 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 616 |
+
" \n",
|
| 617 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 618 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 619 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 620 |
+
" ...,\n",
|
| 621 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 622 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 623 |
+
" [0., 0., 0., ..., 0., 0., 0.]]],\n",
|
| 624 |
+
" \n",
|
| 625 |
+
" \n",
|
| 626 |
+
" [[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 627 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 628 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 629 |
+
" ...,\n",
|
| 630 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 631 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 632 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 633 |
+
" \n",
|
| 634 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 635 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 636 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 637 |
+
" ...,\n",
|
| 638 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 639 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 640 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 641 |
+
" \n",
|
| 642 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 643 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 644 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 645 |
+
" ...,\n",
|
| 646 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 647 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 648 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 649 |
+
" \n",
|
| 650 |
+
" ...,\n",
|
| 651 |
+
" \n",
|
| 652 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 653 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 654 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 655 |
+
" ...,\n",
|
| 656 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 657 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 658 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 659 |
+
" \n",
|
| 660 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 661 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 662 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 663 |
+
" ...,\n",
|
| 664 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 665 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 666 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 667 |
+
" \n",
|
| 668 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 669 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 670 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 671 |
+
" ...,\n",
|
| 672 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 673 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 674 |
+
" [0., 0., 0., ..., 0., 0., 0.]]],\n",
|
| 675 |
+
" \n",
|
| 676 |
+
" \n",
|
| 677 |
+
" ...,\n",
|
| 678 |
+
" \n",
|
| 679 |
+
" \n",
|
| 680 |
+
" [[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 681 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 682 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 683 |
+
" ...,\n",
|
| 684 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 685 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 686 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 687 |
+
" \n",
|
| 688 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 689 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 690 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 691 |
+
" ...,\n",
|
| 692 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 693 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 694 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 695 |
+
" \n",
|
| 696 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 697 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 698 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 699 |
+
" ...,\n",
|
| 700 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 701 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 702 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 703 |
+
" \n",
|
| 704 |
+
" ...,\n",
|
| 705 |
+
" \n",
|
| 706 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 707 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 708 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 709 |
+
" ...,\n",
|
| 710 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 711 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 712 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 713 |
+
" \n",
|
| 714 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 715 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 716 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 717 |
+
" ...,\n",
|
| 718 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 719 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 720 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 721 |
+
" \n",
|
| 722 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 723 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 724 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 725 |
+
" ...,\n",
|
| 726 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 727 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 728 |
+
" [0., 0., 0., ..., 0., 0., 0.]]],\n",
|
| 729 |
+
" \n",
|
| 730 |
+
" \n",
|
| 731 |
+
" [[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 732 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 733 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 734 |
+
" ...,\n",
|
| 735 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 736 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 737 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 738 |
+
" \n",
|
| 739 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 740 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 741 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 742 |
+
" ...,\n",
|
| 743 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 744 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 745 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 746 |
+
" \n",
|
| 747 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 748 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 749 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 750 |
+
" ...,\n",
|
| 751 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 752 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 753 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 754 |
+
" \n",
|
| 755 |
+
" ...,\n",
|
| 756 |
+
" \n",
|
| 757 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 758 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 759 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 760 |
+
" ...,\n",
|
| 761 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 762 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 763 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 764 |
+
" \n",
|
| 765 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 766 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 767 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 768 |
+
" ...,\n",
|
| 769 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 770 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 771 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 772 |
+
" \n",
|
| 773 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 774 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 775 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 776 |
+
" ...,\n",
|
| 777 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 778 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 779 |
+
" [0., 0., 0., ..., 0., 0., 0.]]],\n",
|
| 780 |
+
" \n",
|
| 781 |
+
" \n",
|
| 782 |
+
" [[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 783 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 784 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 785 |
+
" ...,\n",
|
| 786 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 787 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 788 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 789 |
+
" \n",
|
| 790 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 791 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 792 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 793 |
+
" ...,\n",
|
| 794 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 795 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 796 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 797 |
+
" \n",
|
| 798 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 799 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 800 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 801 |
+
" ...,\n",
|
| 802 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 803 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 804 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 805 |
+
" \n",
|
| 806 |
+
" ...,\n",
|
| 807 |
+
" \n",
|
| 808 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 809 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 810 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 811 |
+
" ...,\n",
|
| 812 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 813 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 814 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 815 |
+
" \n",
|
| 816 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 817 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 818 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 819 |
+
" ...,\n",
|
| 820 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 821 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 822 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 823 |
+
" \n",
|
| 824 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 825 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 826 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 827 |
+
" ...,\n",
|
| 828 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 829 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 830 |
+
" [0., 0., 0., ..., 0., 0., 0.]]]], dtype=torch.float16),\n",
|
| 831 |
+
" {'key': ['sub-102109_mod-tfMRI_task-WM_mag-3T_dir-LR',\n",
|
| 832 |
+
" 'sub-731140_mod-tfMRI_task-WM_mag-3T_dir-RL',\n",
|
| 833 |
+
" 'sub-149539_mod-tfMRI_task-MOTOR_mag-3T_dir-RL',\n",
|
| 834 |
+
" 'sub-376247_mod-tfMRI_task-WM_mag-3T_dir-RL',\n",
|
| 835 |
+
" 'sub-164939_mod-tfMRI_task-EMOTION_mag-3T_dir-LR',\n",
|
| 836 |
+
" 'sub-198653_mod-tfMRI_task-SOCIAL_mag-3T_dir-RL',\n",
|
| 837 |
+
" 'sub-210011_mod-tfMRI_task-WM_mag-3T_dir-RL',\n",
|
| 838 |
+
" 'sub-356948_mod-tfMRI_task-EMOTION_mag-3T_dir-LR'],\n",
|
| 839 |
+
" 'sub': ['102109',\n",
|
| 840 |
+
" '731140',\n",
|
| 841 |
+
" '149539',\n",
|
| 842 |
+
" '376247',\n",
|
| 843 |
+
" '164939',\n",
|
| 844 |
+
" '198653',\n",
|
| 845 |
+
" '210011',\n",
|
| 846 |
+
" '356948'],\n",
|
| 847 |
+
" 'mod': ['tfMRI',\n",
|
| 848 |
+
" 'tfMRI',\n",
|
| 849 |
+
" 'tfMRI',\n",
|
| 850 |
+
" 'tfMRI',\n",
|
| 851 |
+
" 'tfMRI',\n",
|
| 852 |
+
" 'tfMRI',\n",
|
| 853 |
+
" 'tfMRI',\n",
|
| 854 |
+
" 'tfMRI'],\n",
|
| 855 |
+
" 'task': ['WM', 'WM', 'MOTOR', 'WM', 'EMOTION', 'SOCIAL', 'WM', 'EMOTION'],\n",
|
| 856 |
+
" 'mag': ['3T', '3T', '3T', '3T', '3T', '3T', '3T', '3T'],\n",
|
| 857 |
+
" 'dir': ['LR', 'RL', 'RL', 'RL', 'LR', 'RL', 'RL', 'LR'],\n",
|
| 858 |
+
" 'start': tensor([15, 15, 19, 15, 19, 15, 15, 19]),\n",
|
| 859 |
+
" 'trial_type': ['2bk_tools',\n",
|
| 860 |
+
" '2bk_body',\n",
|
| 861 |
+
" 'lh',\n",
|
| 862 |
+
" '2bk_body',\n",
|
| 863 |
+
" 'neut',\n",
|
| 864 |
+
" 'mental',\n",
|
| 865 |
+
" '2bk_body',\n",
|
| 866 |
+
" 'neut']}]"
|
| 867 |
+
]
|
| 868 |
+
}
|
| 869 |
+
],
|
| 870 |
+
"source": [
|
| 871 |
+
"i"
|
| 872 |
+
]
|
| 873 |
+
},
|
| 874 |
+
{
|
| 875 |
+
"cell_type": "code",
|
| 876 |
+
"execution_count": 16,
|
| 877 |
+
"id": "ae8a1da7",
|
| 878 |
+
"metadata": {},
|
| 879 |
+
"outputs": [
|
| 880 |
+
{
|
| 881 |
+
"name": "stdout",
|
| 882 |
+
"output_type": "stream",
|
| 883 |
+
"text": [
|
| 884 |
+
"torch.Size([8, 16, 144, 320])"
|
| 885 |
+
]
|
| 886 |
+
}
|
| 887 |
+
],
|
| 888 |
+
"source": [
|
| 889 |
+
"i[0].shape"
|
| 890 |
+
]
|
| 891 |
+
},
|
| 892 |
+
{
|
| 893 |
+
"cell_type": "code",
|
| 894 |
+
"execution_count": 17,
|
| 895 |
+
"id": "ce3f32de",
|
| 896 |
+
"metadata": {},
|
| 897 |
+
"outputs": [],
|
| 898 |
+
"source": [
|
| 899 |
+
"mae_model(torch.randn(1,16,144,320)).shape"
|
| 900 |
+
]
|
| 901 |
+
},
|
| 902 |
+
{
|
| 903 |
+
"cell_type": "code",
|
| 904 |
+
"execution_count": 18,
|
| 905 |
+
"id": "ff571604",
|
| 906 |
+
"metadata": {},
|
| 907 |
+
"outputs": [],
|
| 908 |
+
"source": [
|
| 909 |
+
"global_pool"
|
| 910 |
+
]
|
| 911 |
+
},
|
| 912 |
+
{
|
| 913 |
+
"cell_type": "code",
|
| 914 |
+
"execution_count": 19,
|
| 915 |
+
"id": "348c52bc",
|
| 916 |
+
"metadata": {},
|
| 917 |
+
"outputs": [],
|
| 918 |
+
"source": [
|
| 919 |
+
"if os.getenv('global_pool') == \"False\":\n",
|
| 920 |
+
" global_pool = False\n",
|
| 921 |
+
"else:\n",
|
| 922 |
+
" global_pool = True\n",
|
| 923 |
+
"print(f\"global_pool = {global_pool}\")\n",
|
| 924 |
+
"\n",
|
| 925 |
+
"try:\n",
|
| 926 |
+
" gsr\n",
|
| 927 |
+
"except:\n",
|
| 928 |
+
" gsr = True\n",
|
| 929 |
+
" print(\"set gsr to True\")\n",
|
| 930 |
+
"print(f\"gsr = {gsr}\")"
|
| 931 |
+
]
|
| 932 |
+
},
|
| 933 |
+
{
|
| 934 |
+
"cell_type": "code",
|
| 935 |
+
"execution_count": 20,
|
| 936 |
+
"id": "82ccddb7",
|
| 937 |
+
"metadata": {},
|
| 938 |
+
"outputs": [],
|
| 939 |
+
"source": [
|
| 940 |
+
"mae_model(torch.randn(1,16,144,320),global_pool=global_pool, forward_features = True).shape"
|
| 941 |
+
]
|
| 942 |
+
},
|
| 943 |
+
{
|
| 944 |
+
"cell_type": "code",
|
| 945 |
+
"execution_count": 21,
|
| 946 |
+
"id": "856a0186",
|
| 947 |
+
"metadata": {},
|
| 948 |
+
"outputs": [],
|
| 949 |
+
"source": [
|
| 950 |
+
"mae_model(torch.randn(1,1,16,144,320),global_pool=global_pool, forward_features = True).shape"
|
| 951 |
+
]
|
| 952 |
+
},
|
| 953 |
+
{
|
| 954 |
+
"cell_type": "code",
|
| 955 |
+
"execution_count": 22,
|
| 956 |
+
"id": "33e1c8cd",
|
| 957 |
+
"metadata": {},
|
| 958 |
+
"outputs": [
|
| 959 |
+
{
|
| 960 |
+
"name": "stdout",
|
| 961 |
+
"output_type": "stream",
|
| 962 |
+
"text": [
|
| 963 |
+
"torch.Size([1, 1024])"
|
| 964 |
+
]
|
| 965 |
+
}
|
| 966 |
+
],
|
| 967 |
+
"source": [
|
| 968 |
+
"mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape"
|
| 969 |
+
]
|
| 970 |
+
},
|
| 971 |
+
{
|
| 972 |
+
"cell_type": "code",
|
| 973 |
+
"execution_count": 23,
|
| 974 |
+
"id": "c56d3fc0",
|
| 975 |
+
"metadata": {},
|
| 976 |
+
"outputs": [
|
| 977 |
+
{
|
| 978 |
+
"name": "stdout",
|
| 979 |
+
"output_type": "stream",
|
| 980 |
+
"text": [
|
| 981 |
+
"torch.Size([1, 1024])"
|
| 982 |
+
]
|
| 983 |
+
}
|
| 984 |
+
],
|
| 985 |
+
"source": [
|
| 986 |
+
"mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape"
|
| 987 |
+
]
|
| 988 |
+
},
|
| 989 |
+
{
|
| 990 |
+
"cell_type": "code",
|
| 991 |
+
"execution_count": 24,
|
| 992 |
+
"id": "d6395f12",
|
| 993 |
+
"metadata": {},
|
| 994 |
+
"outputs": [],
|
| 995 |
+
"source": [
|
| 996 |
+
"mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:].sum()"
|
| 997 |
+
]
|
| 998 |
+
},
|
| 999 |
+
{
|
| 1000 |
+
"cell_type": "code",
|
| 1001 |
+
"execution_count": 25,
|
| 1002 |
+
"id": "2dca71b2",
|
| 1003 |
+
"metadata": {},
|
| 1004 |
+
"outputs": [],
|
| 1005 |
+
"source": [
|
| 1006 |
+
"list(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:]).sum()"
|
| 1007 |
+
]
|
| 1008 |
+
},
|
| 1009 |
+
{
|
| 1010 |
+
"cell_type": "code",
|
| 1011 |
+
"execution_count": 26,
|
| 1012 |
+
"id": "a25dd030",
|
| 1013 |
+
"metadata": {},
|
| 1014 |
+
"outputs": [
|
| 1015 |
+
{
|
| 1016 |
+
"name": "stdout",
|
| 1017 |
+
"output_type": "stream",
|
| 1018 |
+
"text": [
|
| 1019 |
+
"1024"
|
| 1020 |
+
]
|
| 1021 |
+
}
|
| 1022 |
+
],
|
| 1023 |
+
"source": [
|
| 1024 |
+
"sum(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])"
|
| 1025 |
+
]
|
| 1026 |
+
},
|
| 1027 |
+
{
|
| 1028 |
+
"cell_type": "code",
|
| 1029 |
+
"execution_count": 27,
|
| 1030 |
+
"id": "edd2edd2",
|
| 1031 |
+
"metadata": {},
|
| 1032 |
+
"outputs": [
|
| 1033 |
+
{
|
| 1034 |
+
"name": "stdout",
|
| 1035 |
+
"output_type": "stream",
|
| 1036 |
+
"text": [
|
| 1037 |
+
"np.int64(1024)"
|
| 1038 |
+
]
|
| 1039 |
+
}
|
| 1040 |
+
],
|
| 1041 |
+
"source": [
|
| 1042 |
+
"np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])"
|
| 1043 |
+
]
|
| 1044 |
+
},
|
| 1045 |
+
{
|
| 1046 |
+
"cell_type": "code",
|
| 1047 |
+
"execution_count": 28,
|
| 1048 |
+
"id": "7c934b0f",
|
| 1049 |
+
"metadata": {},
|
| 1050 |
+
"outputs": [],
|
| 1051 |
+
"source": [
|
| 1052 |
+
"class LinearClassifier(nn.Module):\n",
|
| 1053 |
+
" def __init__(self, input_dim, num_classes):\n",
|
| 1054 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1055 |
+
" self.linear = nn.Linear(input_dim, num_classes)\n",
|
| 1056 |
+
" \n",
|
| 1057 |
+
" def forward(self, x):\n",
|
| 1058 |
+
" # Flatten the input except for the batch dimension\n",
|
| 1059 |
+
" x = x.view(x.size(0), -1)\n",
|
| 1060 |
+
" out = self.linear(x)\n",
|
| 1061 |
+
" return out # Raw logits\n",
|
| 1062 |
+
"\n",
|
| 1063 |
+
"# Determine the input dimension from a single sample\n",
|
| 1064 |
+
"# Assuming images are of shape [1, 16, 144, 320]\n",
|
| 1065 |
+
"input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])\n",
|
| 1066 |
+
"print(f\"Input dimension: {input_dim}\")"
|
| 1067 |
+
]
|
| 1068 |
+
},
|
| 1069 |
+
{
|
| 1070 |
+
"cell_type": "code",
|
| 1071 |
+
"execution_count": 29,
|
| 1072 |
+
"id": "bd0fcfa7",
|
| 1073 |
+
"metadata": {},
|
| 1074 |
+
"outputs": [],
|
| 1075 |
+
"source": [
|
| 1076 |
+
"class FullModel(nn.Module):\n",
|
| 1077 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1078 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1079 |
+
" self.lc_model = lc_model\n",
|
| 1080 |
+
" self.mae_model = mae_model\n",
|
| 1081 |
+
" \n",
|
| 1082 |
+
" \n",
|
| 1083 |
+
" def forward(self, x, gsr):\n",
|
| 1084 |
+
" x = self.mae_model(x, global_pool=global_pool, forward_features = True)\n",
|
| 1085 |
+
" x = self.lc_model(x)\n",
|
| 1086 |
+
" return x"
|
| 1087 |
+
]
|
| 1088 |
+
},
|
| 1089 |
+
{
|
| 1090 |
+
"cell_type": "code",
|
| 1091 |
+
"execution_count": 30,
|
| 1092 |
+
"id": "25b06ed9",
|
| 1093 |
+
"metadata": {},
|
| 1094 |
+
"outputs": [],
|
| 1095 |
+
"source": [
|
| 1096 |
+
"class LinearClassifier(nn.Module):\n",
|
| 1097 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1098 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1099 |
+
" self.lc_model = lc_model\n",
|
| 1100 |
+
" \n",
|
| 1101 |
+
" \n",
|
| 1102 |
+
" def forward(self, x):\n",
|
| 1103 |
+
" # Flatten the input except for the batch dimension\n",
|
| 1104 |
+
" x = x.view(x.size(0), -1)\n",
|
| 1105 |
+
" out = self.linear(x)\n",
|
| 1106 |
+
" return out # Raw logits\n",
|
| 1107 |
+
"\n",
|
| 1108 |
+
"# Determine the input dimension from a single sample\n",
|
| 1109 |
+
"# Assuming images are of shape [1, 16, 144, 320]\n",
|
| 1110 |
+
"input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])\n",
|
| 1111 |
+
"print(f\"Input dimension: {input_dim}\")"
|
| 1112 |
+
]
|
| 1113 |
+
},
|
| 1114 |
+
{
|
| 1115 |
+
"cell_type": "code",
|
| 1116 |
+
"execution_count": 31,
|
| 1117 |
+
"id": "97f9bbc3",
|
| 1118 |
+
"metadata": {},
|
| 1119 |
+
"outputs": [],
|
| 1120 |
+
"source": [
|
| 1121 |
+
"# Initialize the model\n",
|
| 1122 |
+
"lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)\n",
|
| 1123 |
+
"\n",
|
| 1124 |
+
"model = FullModel(lc_model, mae_model)\n",
|
| 1125 |
+
"\n",
|
| 1126 |
+
"# Move the model to the GPU\n",
|
| 1127 |
+
"model.to(device)\n",
|
| 1128 |
+
"\n",
|
| 1129 |
+
"# Define loss function\n",
|
| 1130 |
+
"criterion = nn.CrossEntropyLoss()\n",
|
| 1131 |
+
"\n",
|
| 1132 |
+
"# Define optimizer with L2 regularization (weight_decay)\n",
|
| 1133 |
+
"learning_rate = 1e-4\n",
|
| 1134 |
+
"weight_decay = 1e-5 # Adjust based on your needs\n",
|
| 1135 |
+
"optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n",
|
| 1136 |
+
"num_epochs = 20 # Adjust as needed"
|
| 1137 |
+
]
|
| 1138 |
+
},
|
| 1139 |
+
{
|
| 1140 |
+
"cell_type": "code",
|
| 1141 |
+
"execution_count": 32,
|
| 1142 |
+
"id": "313b529e",
|
| 1143 |
+
"metadata": {},
|
| 1144 |
+
"outputs": [],
|
| 1145 |
+
"source": [
|
| 1146 |
+
"from sklearn.preprocessing import LabelEncoder\n",
|
| 1147 |
+
"\n",
|
| 1148 |
+
"INCLUDE_CONDS = {\n",
|
| 1149 |
+
" \"fear\",\n",
|
| 1150 |
+
" \"neut\",\n",
|
| 1151 |
+
" \"math\",\n",
|
| 1152 |
+
" \"story\",\n",
|
| 1153 |
+
" \"lf\",\n",
|
| 1154 |
+
" \"lh\",\n",
|
| 1155 |
+
" \"rf\",\n",
|
| 1156 |
+
" \"rh\",\n",
|
| 1157 |
+
" \"t\",\n",
|
| 1158 |
+
" \"match\",\n",
|
| 1159 |
+
" \"relation\",\n",
|
| 1160 |
+
" \"mental\",\n",
|
| 1161 |
+
" \"rnd\",\n",
|
| 1162 |
+
" \"0bk_body\",\n",
|
| 1163 |
+
" \"2bk_body\",\n",
|
| 1164 |
+
" \"0bk_faces\",\n",
|
| 1165 |
+
" \"2bk_faces\",\n",
|
| 1166 |
+
" \"0bk_places\",\n",
|
| 1167 |
+
" \"2bk_places\",\n",
|
| 1168 |
+
" \"0bk_tools\",\n",
|
| 1169 |
+
" \"2bk_tools\",\n",
|
| 1170 |
+
"}\n",
|
| 1171 |
+
"\n",
|
| 1172 |
+
"# test_data = []\n",
|
| 1173 |
+
"\n",
|
| 1174 |
+
"# # Iterate over the DataLoader with a progress bar\n",
|
| 1175 |
+
"# for sample in tqdm(train_dl, desc=\"Processing samples\"):\n",
|
| 1176 |
+
"# x = sample['image']\n",
|
| 1177 |
+
"# y = sample['meta']['trial_type']\n",
|
| 1178 |
+
"# key = sample['meta']['key']\n",
|
| 1179 |
+
"# print(x.shape, y, key)\n",
|
| 1180 |
+
"# break\n",
|
| 1181 |
+
"# Initialize the label encoder\n",
|
| 1182 |
+
"label_encoder = LabelEncoder()\n",
|
| 1183 |
+
"label_encoder.fit(sorted(INCLUDE_CONDS)) # Ensure consistent ordering\n",
|
| 1184 |
+
"\n",
|
| 1185 |
+
"num_classes = len(label_encoder.classes_)\n",
|
| 1186 |
+
"print(f\"Number of classes: {num_classes}\")"
|
| 1187 |
+
]
|
| 1188 |
+
},
|
| 1189 |
+
{
|
| 1190 |
+
"cell_type": "code",
|
| 1191 |
+
"execution_count": 33,
|
| 1192 |
+
"id": "c5f58f30",
|
| 1193 |
+
"metadata": {},
|
| 1194 |
+
"outputs": [],
|
| 1195 |
+
"source": [
|
| 1196 |
+
"f_train = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/train_hcp.hdf5', 'r')\n",
|
| 1197 |
+
"flatmaps_train = f_train['flatmaps']\n",
|
| 1198 |
+
"\n",
|
| 1199 |
+
"f_test = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/test_hcp.hdf5', 'r')\n",
|
| 1200 |
+
"flatmaps_test = f_test['flatmaps']\n",
|
| 1201 |
+
"\n",
|
| 1202 |
+
"metadata_train = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_train_HCP.npy', allow_pickle=True)\n",
|
| 1203 |
+
"metadata_test = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_test_HCP.npy', allow_pickle=True)"
|
| 1204 |
+
]
|
| 1205 |
+
},
|
| 1206 |
+
{
|
| 1207 |
+
"cell_type": "code",
|
| 1208 |
+
"execution_count": 34,
|
| 1209 |
+
"id": "928abadb",
|
| 1210 |
+
"metadata": {},
|
| 1211 |
+
"outputs": [],
|
| 1212 |
+
"source": [
|
| 1213 |
+
"from torch.utils.data import Dataset, DataLoader\n",
|
| 1214 |
+
"\n",
|
| 1215 |
+
"class HCPFlatDataset(Dataset):\n",
|
| 1216 |
+
" def __init__(self, flatmaps, metadata):\n",
|
| 1217 |
+
" self.flatmaps = flatmaps\n",
|
| 1218 |
+
" self.metadata = metadata\n",
|
| 1219 |
+
"\n",
|
| 1220 |
+
" def __len__(self):\n",
|
| 1221 |
+
" return len(self.metadata)\n",
|
| 1222 |
+
"\n",
|
| 1223 |
+
" def __getitem__(self, idx):\n",
|
| 1224 |
+
" return self.flatmaps[idx], json.loads(self.metadata[idx])\n",
|
| 1225 |
+
"\n",
|
| 1226 |
+
"# Loading to cpu for faster training, this can take several minutes. Remove this [:] if you want to move one at the time.\n",
|
| 1227 |
+
"train_dataset = HCPFlatDataset(flatmaps_train, metadata_train)\n",
|
| 1228 |
+
"train_dl = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)\n",
|
| 1229 |
+
"\n",
|
| 1230 |
+
"test_dataset = HCPFlatDataset(flatmaps_test, metadata_test)\n",
|
| 1231 |
+
"test_dl = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)"
|
| 1232 |
+
]
|
| 1233 |
+
},
|
| 1234 |
+
{
|
| 1235 |
+
"cell_type": "code",
|
| 1236 |
+
"execution_count": 35,
|
| 1237 |
+
"id": "ec69248a",
|
| 1238 |
+
"metadata": {},
|
| 1239 |
+
"outputs": [],
|
| 1240 |
+
"source": [
|
| 1241 |
+
"from mae_utils.flat import load_hcp_flat_mask\n",
|
| 1242 |
+
"from mae_utils.flat import create_hcp_flat\n",
|
| 1243 |
+
"from mae_utils.flat import batch_unmask\n",
|
| 1244 |
+
"import mae_utils.visualize as vis\n",
|
| 1245 |
+
"\n",
|
| 1246 |
+
"flat_mask = load_hcp_flat_mask(hcp_flat_path)\n",
|
| 1247 |
+
"\n",
|
| 1248 |
+
"mae_model = flat_models.mae_vit_large_fmri(\n",
|
| 1249 |
+
" patch_size=patch_size,\n",
|
| 1250 |
+
" decoder_embed_dim=decoder_embed_dim,\n",
|
| 1251 |
+
" t_patch_size=t_patch_size,\n",
|
| 1252 |
+
" pred_t_dim=pred_t_dim,\n",
|
| 1253 |
+
" decoder_depth=4,\n",
|
| 1254 |
+
" cls_embed=cls_embed,\n",
|
| 1255 |
+
" norm_pix_loss=norm_pix_loss,\n",
|
| 1256 |
+
" no_qkv_bias=no_qkv_bias,\n",
|
| 1257 |
+
" sep_pos_embed=sep_pos_embed,\n",
|
| 1258 |
+
" trunc_init=trunc_init,\n",
|
| 1259 |
+
" pct_masks_to_decode=pct_masks_to_decode,\n",
|
| 1260 |
+
" img_mask=flat_mask,\n",
|
| 1261 |
+
")"
|
| 1262 |
+
]
|
| 1263 |
+
},
|
| 1264 |
+
{
|
| 1265 |
+
"cell_type": "code",
|
| 1266 |
+
"execution_count": 36,
|
| 1267 |
+
"id": "4e5045c3",
|
| 1268 |
+
"metadata": {},
|
| 1269 |
+
"outputs": [],
|
| 1270 |
+
"source": [
|
| 1271 |
+
"checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')]\n",
|
| 1272 |
+
"\n",
|
| 1273 |
+
"if utils.is_interactive():\n",
|
| 1274 |
+
" latest_checkpoint = \"epoch99.pth\"\n",
|
| 1275 |
+
"else:\n",
|
| 1276 |
+
" latest_checkpoint = sys.argv[2] \n",
|
| 1277 |
+
"print(f\"latest_checkpoint: {latest_checkpoint}\")\n",
|
| 1278 |
+
"\n",
|
| 1279 |
+
"# Load the checkpoint\n",
|
| 1280 |
+
"checkpoint_path = os.path.join(outdir, latest_checkpoint)\n",
|
| 1281 |
+
"\n",
|
| 1282 |
+
"state = torch.load(checkpoint_path)\n",
|
| 1283 |
+
"mae_model.load_state_dict(state[\"model_state_dict\"], strict=False)\n",
|
| 1284 |
+
"mae_model.to(device)\n",
|
| 1285 |
+
"\n",
|
| 1286 |
+
"print(f\"\\nLoaded checkpoint {latest_checkpoint} from {outdir}\\n\")"
|
| 1287 |
+
]
|
| 1288 |
+
},
|
| 1289 |
+
{
|
| 1290 |
+
"cell_type": "code",
|
| 1291 |
+
"execution_count": 37,
|
| 1292 |
+
"id": "0173f847",
|
| 1293 |
+
"metadata": {},
|
| 1294 |
+
"outputs": [],
|
| 1295 |
+
"source": [
|
| 1296 |
+
"class FullModel(nn.Module):\n",
|
| 1297 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1298 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1299 |
+
" self.lc_model = lc_model\n",
|
| 1300 |
+
" self.mae_model = mae_model\n",
|
| 1301 |
+
" \n",
|
| 1302 |
+
" \n",
|
| 1303 |
+
" def forward(self, x, gsr):\n",
|
| 1304 |
+
" x = self.mae_model(x, global_pool=global_pool, forward_features = True)\n",
|
| 1305 |
+
" x = self.lc_model(x)\n",
|
| 1306 |
+
" return x"
|
| 1307 |
+
]
|
| 1308 |
+
},
|
| 1309 |
+
{
|
| 1310 |
+
"cell_type": "code",
|
| 1311 |
+
"execution_count": 38,
|
| 1312 |
+
"id": "551a5976",
|
| 1313 |
+
"metadata": {},
|
| 1314 |
+
"outputs": [],
|
| 1315 |
+
"source": [
|
| 1316 |
+
"class LinearClassifier(nn.Module):\n",
|
| 1317 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1318 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1319 |
+
" self.lc_model = lc_model\n",
|
| 1320 |
+
" \n",
|
| 1321 |
+
" \n",
|
| 1322 |
+
" def forward(self, x):\n",
|
| 1323 |
+
" # Flatten the input except for the batch dimension\n",
|
| 1324 |
+
" x = x.view(x.size(0), -1)\n",
|
| 1325 |
+
" out = self.linear(x)\n",
|
| 1326 |
+
" return out # Raw logits\n",
|
| 1327 |
+
"\n",
|
| 1328 |
+
"# Determine the input dimension from a single sample\n",
|
| 1329 |
+
"# Assuming images are of shape [1, 16, 144, 320]\n",
|
| 1330 |
+
"input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])\n",
|
| 1331 |
+
"print(f\"Input dimension: {input_dim}\")"
|
| 1332 |
+
]
|
| 1333 |
+
},
|
| 1334 |
+
{
|
| 1335 |
+
"cell_type": "code",
|
| 1336 |
+
"execution_count": 39,
|
| 1337 |
+
"id": "a21df922",
|
| 1338 |
+
"metadata": {},
|
| 1339 |
+
"outputs": [],
|
| 1340 |
+
"source": [
|
| 1341 |
+
"# Initialize the model\n",
|
| 1342 |
+
"lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)\n",
|
| 1343 |
+
"\n",
|
| 1344 |
+
"model = FullModel(lc_model, mae_model)\n",
|
| 1345 |
+
"\n",
|
| 1346 |
+
"# Move the model to the GPU\n",
|
| 1347 |
+
"model.to(device)\n",
|
| 1348 |
+
"\n",
|
| 1349 |
+
"# Define loss function\n",
|
| 1350 |
+
"criterion = nn.CrossEntropyLoss()\n",
|
| 1351 |
+
"\n",
|
| 1352 |
+
"# Define optimizer with L2 regularization (weight_decay)\n",
|
| 1353 |
+
"learning_rate = 1e-4\n",
|
| 1354 |
+
"weight_decay = 1e-5 # Adjust based on your needs\n",
|
| 1355 |
+
"optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n",
|
| 1356 |
+
"num_epochs = 20 # Adjust as needed"
|
| 1357 |
+
]
|
| 1358 |
+
},
|
| 1359 |
+
{
|
| 1360 |
+
"cell_type": "code",
|
| 1361 |
+
"execution_count": 40,
|
| 1362 |
+
"id": "7b262408",
|
| 1363 |
+
"metadata": {},
|
| 1364 |
+
"outputs": [],
|
| 1365 |
+
"source": [
|
| 1366 |
+
"class LinearClassifier(nn.Module):\n",
|
| 1367 |
+
" def __init__(self, input_dim, num_classes):\n",
|
| 1368 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1369 |
+
" self.linear = nn.Linear(input_dim, num_classes)\n",
|
| 1370 |
+
" \n",
|
| 1371 |
+
" def forward(self, x):\n",
|
| 1372 |
+
" # Flatten the input except for the batch dimension\n",
|
| 1373 |
+
" x = x.view(x.size(0), -1)\n",
|
| 1374 |
+
" out = self.linear(x)\n",
|
| 1375 |
+
" return out # Raw logits\n",
|
| 1376 |
+
"\n",
|
| 1377 |
+
"# Determine the input dimension from a single sample\n",
|
| 1378 |
+
"# Assuming images are of shape [1, 16, 144, 320]\n",
|
| 1379 |
+
"input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])\n",
|
| 1380 |
+
"print(f\"Input dimension: {input_dim}\")"
|
| 1381 |
+
]
|
| 1382 |
+
},
|
| 1383 |
+
{
|
| 1384 |
+
"cell_type": "code",
|
| 1385 |
+
"execution_count": 41,
|
| 1386 |
+
"id": "7bff8f73",
|
| 1387 |
+
"metadata": {},
|
| 1388 |
+
"outputs": [],
|
| 1389 |
+
"source": [
|
| 1390 |
+
"class FullModel(nn.Module):\n",
|
| 1391 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1392 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1393 |
+
" self.lc_model = lc_model\n",
|
| 1394 |
+
" self.mae_model = mae_model\n",
|
| 1395 |
+
" \n",
|
| 1396 |
+
" \n",
|
| 1397 |
+
" def forward(self, x, gsr):\n",
|
| 1398 |
+
" x = self.mae_model(x, global_pool=global_pool, forward_features = True)\n",
|
| 1399 |
+
" x = self.lc_model(x)\n",
|
| 1400 |
+
" return x"
|
| 1401 |
+
]
|
| 1402 |
+
},
|
| 1403 |
+
{
|
| 1404 |
+
"cell_type": "code",
|
| 1405 |
+
"execution_count": 42,
|
| 1406 |
+
"id": "697fc651",
|
| 1407 |
+
"metadata": {},
|
| 1408 |
+
"outputs": [],
|
| 1409 |
+
"source": [
|
| 1410 |
+
"# Initialize the model\n",
|
| 1411 |
+
"lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)\n",
|
| 1412 |
+
"\n",
|
| 1413 |
+
"model = FullModel(lc_model, mae_model)\n",
|
| 1414 |
+
"\n",
|
| 1415 |
+
"# Move the model to the GPU\n",
|
| 1416 |
+
"model.to(device)\n",
|
| 1417 |
+
"\n",
|
| 1418 |
+
"# Define loss function\n",
|
| 1419 |
+
"criterion = nn.CrossEntropyLoss()\n",
|
| 1420 |
+
"\n",
|
| 1421 |
+
"# Define optimizer with L2 regularization (weight_decay)\n",
|
| 1422 |
+
"learning_rate = 1e-4\n",
|
| 1423 |
+
"weight_decay = 1e-5 # Adjust based on your needs\n",
|
| 1424 |
+
"optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n",
|
| 1425 |
+
"num_epochs = 20 # Adjust as needed"
|
| 1426 |
+
]
|
| 1427 |
+
},
|
| 1428 |
+
{
|
| 1429 |
+
"cell_type": "code",
|
| 1430 |
+
"execution_count": 43,
|
| 1431 |
+
"id": "99570156",
|
| 1432 |
+
"metadata": {},
|
| 1433 |
+
"outputs": [],
|
| 1434 |
+
"source": [
|
| 1435 |
+
"class FullModel(nn.Module):\n",
|
| 1436 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1437 |
+
" super(FullModel, self).__init__()\n",
|
| 1438 |
+
" self.lc_model = lc_model\n",
|
| 1439 |
+
" self.mae_model = mae_model\n",
|
| 1440 |
+
" \n",
|
| 1441 |
+
" \n",
|
| 1442 |
+
" def forward(self, x, gsr):\n",
|
| 1443 |
+
" x = self.mae_model(x, global_pool=global_pool, forward_features = True)\n",
|
| 1444 |
+
" x = self.lc_model(x)\n",
|
| 1445 |
+
" return x"
|
| 1446 |
+
]
|
| 1447 |
+
},
|
| 1448 |
+
{
|
| 1449 |
+
"cell_type": "code",
|
| 1450 |
+
"execution_count": 44,
|
| 1451 |
+
"id": "8e5c0e3e",
|
| 1452 |
+
"metadata": {},
|
| 1453 |
+
"outputs": [],
|
| 1454 |
+
"source": [
|
| 1455 |
+
"# Initialize the model\n",
|
| 1456 |
+
"lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)\n",
|
| 1457 |
+
"\n",
|
| 1458 |
+
"model = FullModel(lc_model, mae_model)\n",
|
| 1459 |
+
"\n",
|
| 1460 |
+
"# Move the model to the GPU\n",
|
| 1461 |
+
"model.to(device)\n",
|
| 1462 |
+
"\n",
|
| 1463 |
+
"# Define loss function\n",
|
| 1464 |
+
"criterion = nn.CrossEntropyLoss()\n",
|
| 1465 |
+
"\n",
|
| 1466 |
+
"# Define optimizer with L2 regularization (weight_decay)\n",
|
| 1467 |
+
"learning_rate = 1e-4\n",
|
| 1468 |
+
"weight_decay = 1e-5 # Adjust based on your needs\n",
|
| 1469 |
+
"optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n",
|
| 1470 |
+
"num_epochs = 20 # Adjust as needed"
|
| 1471 |
+
]
|
| 1472 |
+
},
|
| 1473 |
+
{
|
| 1474 |
+
"cell_type": "code",
|
| 1475 |
+
"execution_count": 45,
|
| 1476 |
+
"id": "dddffb31",
|
| 1477 |
+
"metadata": {},
|
| 1478 |
+
"outputs": [
|
| 1479 |
+
{
|
| 1480 |
+
"data": {
|
| 1481 |
+
"text/html": [
|
| 1482 |
+
"Tracking run with wandb version 0.18.3"
|
| 1483 |
+
],
|
| 1484 |
+
"text/plain": [
|
| 1485 |
+
"<IPython.core.display.HTML object>"
|
| 1486 |
+
]
|
| 1487 |
+
},
|
| 1488 |
+
"metadata": {},
|
| 1489 |
+
"output_type": "display_data"
|
| 1490 |
+
},
|
| 1491 |
+
{
|
| 1492 |
+
"data": {
|
| 1493 |
+
"text/html": [
|
| 1494 |
+
"Run data is saved locally in <code>/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810</code>"
|
| 1495 |
+
],
|
| 1496 |
+
"text/plain": [
|
| 1497 |
+
"<IPython.core.display.HTML object>"
|
| 1498 |
+
]
|
| 1499 |
+
},
|
| 1500 |
+
"metadata": {},
|
| 1501 |
+
"output_type": "display_data"
|
| 1502 |
+
},
|
| 1503 |
+
{
|
| 1504 |
+
"data": {
|
| 1505 |
+
"text/html": [
|
| 1506 |
+
"Syncing run <strong><a href='https://stability.wandb.io/ckadirt/fMRI-foundation-model/runs/HCPflat_raw_83810' target=\"_blank\">HCPflat_raw</a></strong> to <a href='https://stability.wandb.io/ckadirt/fMRI-foundation-model' target=\"_blank\">Weights & Biases</a> (<a href='https://wandb.me/run' target=\"_blank\">docs</a>)<br/>"
|
| 1507 |
+
],
|
| 1508 |
+
"text/plain": [
|
| 1509 |
+
"<IPython.core.display.HTML object>"
|
| 1510 |
+
]
|
| 1511 |
+
},
|
| 1512 |
+
"metadata": {},
|
| 1513 |
+
"output_type": "display_data"
|
| 1514 |
+
},
|
| 1515 |
+
{
|
| 1516 |
+
"data": {
|
| 1517 |
+
"text/html": [
|
| 1518 |
+
" View project at <a href='https://stability.wandb.io/ckadirt/fMRI-foundation-model' target=\"_blank\">https://stability.wandb.io/ckadirt/fMRI-foundation-model</a>"
|
| 1519 |
+
],
|
| 1520 |
+
"text/plain": [
|
| 1521 |
+
"<IPython.core.display.HTML object>"
|
| 1522 |
+
]
|
| 1523 |
+
},
|
| 1524 |
+
"metadata": {},
|
| 1525 |
+
"output_type": "display_data"
|
| 1526 |
+
},
|
| 1527 |
+
{
|
| 1528 |
+
"data": {
|
| 1529 |
+
"text/html": [
|
| 1530 |
+
" View run at <a href='https://stability.wandb.io/ckadirt/fMRI-foundation-model/runs/HCPflat_raw_83810' target=\"_blank\">https://stability.wandb.io/ckadirt/fMRI-foundation-model/runs/HCPflat_raw_83810</a>"
|
| 1531 |
+
],
|
| 1532 |
+
"text/plain": [
|
| 1533 |
+
"<IPython.core.display.HTML object>"
|
| 1534 |
+
]
|
| 1535 |
+
},
|
| 1536 |
+
"metadata": {},
|
| 1537 |
+
"output_type": "display_data"
|
| 1538 |
+
}
|
| 1539 |
+
],
|
| 1540 |
+
"source": [
|
| 1541 |
+
"import wandb\n",
|
| 1542 |
+
"\n",
|
| 1543 |
+
"if utils.is_interactive():\n",
|
| 1544 |
+
" print(\"Running in interactive notebook. Disabling W&B and ckpt saving.\")\n",
|
| 1545 |
+
" wandb_log = True #False\n",
|
| 1546 |
+
" save_ckpt = True #False\n",
|
| 1547 |
+
"\n",
|
| 1548 |
+
"if wandb_log:\n",
|
| 1549 |
+
" wandb_project = 'fMRI-foundation-model'\n",
|
| 1550 |
+
" wandb_config = {\n",
|
| 1551 |
+
" \"model_name\": \"HCPflat_raw\",\n",
|
| 1552 |
+
" \"batch_size\": batch_size,\n",
|
| 1553 |
+
" \"learning_rate\": learning_rate,\n",
|
| 1554 |
+
" \"weight_decay\": weight_decay,\n",
|
| 1555 |
+
" \"num_epochs\": num_epochs,\n",
|
| 1556 |
+
" \"seed\": seed,\n",
|
| 1557 |
+
" }\n",
|
| 1558 |
+
" print(\"wandb_config:\\n\", wandb_config)\n",
|
| 1559 |
+
" random_id = random.randint(0, 100000)\n",
|
| 1560 |
+
" print(\"wandb_id:\", \"HCPflat_raw\" + f\"_{random_id}\")\n",
|
| 1561 |
+
" wandb.init(\n",
|
| 1562 |
+
" id=\"HCPflat_raw\" + f\"_{random_id}\",\n",
|
| 1563 |
+
" project=wandb_project,\n",
|
| 1564 |
+
" name=\"HCPflat_raw\",\n",
|
| 1565 |
+
" config=wandb_config,\n",
|
| 1566 |
+
" resume=\"allow\",\n",
|
| 1567 |
+
" )"
|
| 1568 |
+
]
|
| 1569 |
+
},
|
| 1570 |
+
{
|
| 1571 |
+
"cell_type": "code",
|
| 1572 |
+
"execution_count": 46,
|
| 1573 |
+
"id": "f8a61d67",
|
| 1574 |
+
"metadata": {},
|
| 1575 |
+
"outputs": [],
|
| 1576 |
+
"source": [
|
| 1577 |
+
"import wandb\n",
|
| 1578 |
+
"\n",
|
| 1579 |
+
"if utils.is_interactive():\n",
|
| 1580 |
+
" print(\"Running in interactive notebook. Disabling W&B and ckpt saving.\")\n",
|
| 1581 |
+
" wandb_log = False\n",
|
| 1582 |
+
" save_ckpt = False\n",
|
| 1583 |
+
"\n",
|
| 1584 |
+
"if wandb_log:\n",
|
| 1585 |
+
" wandb_project = 'fMRI-foundation-model'\n",
|
| 1586 |
+
" wandb_config = {\n",
|
| 1587 |
+
" \"model_name\": model_name+'_HCP_FT',\n",
|
| 1588 |
+
" \"batch_size\": batch_size,\n",
|
| 1589 |
+
" \"learning_rate\": learning_rate,\n",
|
| 1590 |
+
" \"weight_decay\": weight_decay,\n",
|
| 1591 |
+
" \"num_epochs\": num_epochs,\n",
|
| 1592 |
+
" \"seed\": seed,\n",
|
| 1593 |
+
" }\n",
|
| 1594 |
+
" print(\"wandb_config:\\n\", wandb_config)\n",
|
| 1595 |
+
" random_id = random.randint(0, 100000)\n",
|
| 1596 |
+
" print(\"wandb_id:\", \"HCPflat_raw\" + f\"_{random_id}\")\n",
|
| 1597 |
+
" wandb.init(\n",
|
| 1598 |
+
" id=model_name+'_HCP_FT' + f\"_{random_id}\",\n",
|
| 1599 |
+
" project=wandb_project,\n",
|
| 1600 |
+
" name=model_name+'_HCP_FT',\n",
|
| 1601 |
+
" config=wandb_config,\n",
|
| 1602 |
+
" resume=\"allow\",\n",
|
| 1603 |
+
" )"
|
| 1604 |
+
]
|
| 1605 |
+
},
|
| 1606 |
+
{
|
| 1607 |
+
"cell_type": "code",
|
| 1608 |
+
"execution_count": 47,
|
| 1609 |
+
"id": "71ac3c4f",
|
| 1610 |
+
"metadata": {},
|
| 1611 |
+
"outputs": [],
|
| 1612 |
+
"source": [
|
| 1613 |
+
"import wandb\n",
|
| 1614 |
+
"\n",
|
| 1615 |
+
"if utils.is_interactive():\n",
|
| 1616 |
+
" print(\"Running in interactive notebook. Disabling W&B and ckpt saving.\")\n",
|
| 1617 |
+
" wandb_log = True\n",
|
| 1618 |
+
" save_ckpt = True\n",
|
| 1619 |
+
"\n",
|
| 1620 |
+
"if wandb_log:\n",
|
| 1621 |
+
" wandb_project = 'fMRI-foundation-model'\n",
|
| 1622 |
+
" wandb_config = {\n",
|
| 1623 |
+
" \"model_name\": model_name+'_HCP_FT',\n",
|
| 1624 |
+
" \"batch_size\": batch_size,\n",
|
| 1625 |
+
" \"learning_rate\": learning_rate,\n",
|
| 1626 |
+
" \"weight_decay\": weight_decay,\n",
|
| 1627 |
+
" \"num_epochs\": num_epochs,\n",
|
| 1628 |
+
" \"seed\": seed,\n",
|
| 1629 |
+
" }\n",
|
| 1630 |
+
" print(\"wandb_config:\\n\", wandb_config)\n",
|
| 1631 |
+
" random_id = random.randint(0, 100000)\n",
|
| 1632 |
+
" print(\"wandb_id:\", \"HCPflat_raw\" + f\"_{random_id}\")\n",
|
| 1633 |
+
" wandb.init(\n",
|
| 1634 |
+
" id=model_name+'_HCP_FT' + f\"_{random_id}\",\n",
|
| 1635 |
+
" project=wandb_project,\n",
|
| 1636 |
+
" name=model_name+'_HCP_FT',\n",
|
| 1637 |
+
" config=wandb_config,\n",
|
| 1638 |
+
" resume=\"allow\",\n",
|
| 1639 |
+
" )"
|
| 1640 |
+
]
|
| 1641 |
+
}
|
| 1642 |
+
],
|
| 1643 |
+
"metadata": {
|
| 1644 |
+
"kernelspec": {
|
| 1645 |
+
"display_name": "Python 3",
|
| 1646 |
+
"language": "python",
|
| 1647 |
+
"name": "python3"
|
| 1648 |
+
},
|
| 1649 |
+
"language_info": {
|
| 1650 |
+
"codemirror_mode": {
|
| 1651 |
+
"name": "ipython",
|
| 1652 |
+
"version": 3
|
| 1653 |
+
},
|
| 1654 |
+
"file_extension": ".py",
|
| 1655 |
+
"mimetype": "text/x-python",
|
| 1656 |
+
"name": "python",
|
| 1657 |
+
"nbconvert_exporter": "python",
|
| 1658 |
+
"pygments_lexer": "ipython3",
|
| 1659 |
+
"version": "3.11.10"
|
| 1660 |
+
}
|
| 1661 |
+
},
|
| 1662 |
+
"nbformat": 4,
|
| 1663 |
+
"nbformat_minor": 5
|
| 1664 |
+
}
|
fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/config.yaml
ADDED
|
@@ -0,0 +1,49 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
_wandb:
|
| 2 |
+
value:
|
| 3 |
+
cli_version: 0.18.3
|
| 4 |
+
m: []
|
| 5 |
+
python_version: 3.11.10
|
| 6 |
+
session_history: code/_session_history.ipynb
|
| 7 |
+
t:
|
| 8 |
+
"1":
|
| 9 |
+
- 1
|
| 10 |
+
- 5
|
| 11 |
+
- 41
|
| 12 |
+
- 49
|
| 13 |
+
- 53
|
| 14 |
+
- 55
|
| 15 |
+
- 63
|
| 16 |
+
"2":
|
| 17 |
+
- 1
|
| 18 |
+
- 5
|
| 19 |
+
- 41
|
| 20 |
+
- 49
|
| 21 |
+
- 53
|
| 22 |
+
- 55
|
| 23 |
+
- 63
|
| 24 |
+
"3":
|
| 25 |
+
- 2
|
| 26 |
+
- 13
|
| 27 |
+
- 14
|
| 28 |
+
- 16
|
| 29 |
+
- 23
|
| 30 |
+
- 55
|
| 31 |
+
"4": 3.11.10
|
| 32 |
+
"5": 0.18.3
|
| 33 |
+
"8":
|
| 34 |
+
- 1
|
| 35 |
+
- 5
|
| 36 |
+
"12": 0.18.3
|
| 37 |
+
"13": linux-x86_64
|
| 38 |
+
batch_size:
|
| 39 |
+
value: 8
|
| 40 |
+
learning_rate:
|
| 41 |
+
value: 0.0001
|
| 42 |
+
model_name:
|
| 43 |
+
value: HCPflat_raw
|
| 44 |
+
num_epochs:
|
| 45 |
+
value: 20
|
| 46 |
+
seed:
|
| 47 |
+
value: 42
|
| 48 |
+
weight_decay:
|
| 49 |
+
value: 1e-05
|
fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/output.log
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Running in interactive notebook. Disabling W&B and ckpt saving.
|
| 2 |
+
Running in interactive notebook. Disabling W&B and ckpt saving.
|
| 3 |
+
wandb_config:
|
| 4 |
+
{'model_name': 'HCPflat_large_gsrFalse__HCP_FT', 'batch_size': 8, 'learning_rate': 0.0001, 'weight_decay': 1e-05, 'num_epochs': 20, 'seed': 42}
|
| 5 |
+
wandb_id: HCPflat_raw_14592
|
fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/wandb-metadata.json
ADDED
|
@@ -0,0 +1,144 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"os": "Linux-5.15.0-1058-aws-x86_64-with-glibc2.31",
|
| 3 |
+
"python": "3.11.10",
|
| 4 |
+
"startedAt": "2024-10-23T03:22:15.095152Z",
|
| 5 |
+
"program": "ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.ipynb",
|
| 6 |
+
"git": {
|
| 7 |
+
"remote": "https://github.com/MedARC-AI/fMRI-foundation-model",
|
| 8 |
+
"commit": "b1ba684ae7a5cc4155cc046b0abe613de09bf700"
|
| 9 |
+
},
|
| 10 |
+
"email": "torrico.villanueva.cesar.kadir@gmail.com",
|
| 11 |
+
"root": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src",
|
| 12 |
+
"host": "ip-10-0-160-143",
|
| 13 |
+
"username": "ckadirt",
|
| 14 |
+
"executable": "/admin/home-ckadirt/foundation_env/bin/python",
|
| 15 |
+
"cpu_count": 96,
|
| 16 |
+
"cpu_count_logical": 192,
|
| 17 |
+
"gpu": "[NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3]",
|
| 18 |
+
"gpu_count": 8,
|
| 19 |
+
"disk": {
|
| 20 |
+
"/": {
|
| 21 |
+
"total": "249555763200",
|
| 22 |
+
"used": "184990625792"
|
| 23 |
+
}
|
| 24 |
+
},
|
| 25 |
+
"memory": {
|
| 26 |
+
"total": "2147443429376"
|
| 27 |
+
},
|
| 28 |
+
"cpu": {
|
| 29 |
+
"count": 96,
|
| 30 |
+
"countLogical": 192
|
| 31 |
+
},
|
| 32 |
+
"gpu_nvidia": [
|
| 33 |
+
{
|
| 34 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 35 |
+
"memoryTotal": "85520809984",
|
| 36 |
+
"cudaCores": 16896,
|
| 37 |
+
"architecture": "Hopper"
|
| 38 |
+
},
|
| 39 |
+
{
|
| 40 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 41 |
+
"memoryTotal": "85520809984",
|
| 42 |
+
"cudaCores": 16896,
|
| 43 |
+
"architecture": "Hopper"
|
| 44 |
+
},
|
| 45 |
+
{
|
| 46 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 47 |
+
"memoryTotal": "85520809984",
|
| 48 |
+
"cudaCores": 16896,
|
| 49 |
+
"architecture": "Hopper"
|
| 50 |
+
},
|
| 51 |
+
{
|
| 52 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 53 |
+
"memoryTotal": "85520809984",
|
| 54 |
+
"cudaCores": 16896,
|
| 55 |
+
"architecture": "Hopper"
|
| 56 |
+
},
|
| 57 |
+
{
|
| 58 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 59 |
+
"memoryTotal": "85520809984",
|
| 60 |
+
"cudaCores": 16896,
|
| 61 |
+
"architecture": "Hopper"
|
| 62 |
+
},
|
| 63 |
+
{
|
| 64 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 65 |
+
"memoryTotal": "85520809984",
|
| 66 |
+
"cudaCores": 16896,
|
| 67 |
+
"architecture": "Hopper"
|
| 68 |
+
},
|
| 69 |
+
{
|
| 70 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 71 |
+
"memoryTotal": "85520809984",
|
| 72 |
+
"cudaCores": 16896,
|
| 73 |
+
"architecture": "Hopper"
|
| 74 |
+
},
|
| 75 |
+
{
|
| 76 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 77 |
+
"memoryTotal": "85520809984",
|
| 78 |
+
"cudaCores": 16896,
|
| 79 |
+
"architecture": "Hopper"
|
| 80 |
+
}
|
| 81 |
+
],
|
| 82 |
+
"slurm": {
|
| 83 |
+
"cluster_name": "sagemaker2",
|
| 84 |
+
"conf": "/opt/slurm/etc/slurm.conf",
|
| 85 |
+
"cpu_bind": "quiet,mask_cpu:0x00000000000000FFC000000000000000000000FFC0000000",
|
| 86 |
+
"cpu_bind_list": "0x00000000000000FFC000000000000000000000FFC0000000",
|
| 87 |
+
"cpu_bind_type": "mask_cpu:",
|
| 88 |
+
"cpu_bind_verbose": "quiet",
|
| 89 |
+
"cpus_on_node": "20",
|
| 90 |
+
"gpus": "1",
|
| 91 |
+
"gpus_on_node": "1",
|
| 92 |
+
"gtids": "0",
|
| 93 |
+
"job_account": "fmri",
|
| 94 |
+
"job_cpus_per_node": "20",
|
| 95 |
+
"job_end_time": "1729702756",
|
| 96 |
+
"job_gid": "1879800513",
|
| 97 |
+
"job_group": "Domain Users",
|
| 98 |
+
"job_id": "528040",
|
| 99 |
+
"job_name": "bash",
|
| 100 |
+
"job_nodelist": "ip-10-0-160-143",
|
| 101 |
+
"job_num_nodes": "1",
|
| 102 |
+
"job_partition": "p5",
|
| 103 |
+
"job_qos": "idle",
|
| 104 |
+
"job_start_time": "1729648756",
|
| 105 |
+
"job_uid": "1879804696",
|
| 106 |
+
"job_user": "ckadirt",
|
| 107 |
+
"jobid": "528040",
|
| 108 |
+
"launch_node_ipaddr": "172.17.12.61",
|
| 109 |
+
"localid": "0",
|
| 110 |
+
"mpi_type": "pmix_v3",
|
| 111 |
+
"nnodes": "1",
|
| 112 |
+
"nodeid": "0",
|
| 113 |
+
"nodelist": "ip-10-0-160-143",
|
| 114 |
+
"nprocs": "1",
|
| 115 |
+
"ntasks": "1",
|
| 116 |
+
"pmix_mapping_serv": "(vector,(0,1,1))",
|
| 117 |
+
"pmixp_abort_agent_port": "34923",
|
| 118 |
+
"prio_process": "0",
|
| 119 |
+
"procid": "0",
|
| 120 |
+
"pty_port": "45733",
|
| 121 |
+
"pty_win_col": "199",
|
| 122 |
+
"pty_win_row": "17",
|
| 123 |
+
"script_context": "prolog_task",
|
| 124 |
+
"srun_comm_host": "172.17.12.61",
|
| 125 |
+
"srun_comm_port": "39353",
|
| 126 |
+
"step_gpus": "3",
|
| 127 |
+
"step_id": "0",
|
| 128 |
+
"step_launcher_port": "39353",
|
| 129 |
+
"step_nodelist": "ip-10-0-160-143",
|
| 130 |
+
"step_num_nodes": "1",
|
| 131 |
+
"step_num_tasks": "1",
|
| 132 |
+
"step_tasks_per_node": "1",
|
| 133 |
+
"stepid": "0",
|
| 134 |
+
"submit_dir": "/weka/proj-fmri",
|
| 135 |
+
"submit_host": "ip-172-17-12-61",
|
| 136 |
+
"task_pid": "1032669",
|
| 137 |
+
"tasks_per_node": "1",
|
| 138 |
+
"topology_addr": "ip-10-0-160-143",
|
| 139 |
+
"topology_addr_pattern": "node",
|
| 140 |
+
"umask": "0022",
|
| 141 |
+
"working_cluster": "sagemaker2:ip-172-17-63-161:6817:9984:109"
|
| 142 |
+
},
|
| 143 |
+
"cudaVersion": "12.2"
|
| 144 |
+
}
|
fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/wandb-summary.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"_wandb":{"runtime":1}}
|
fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug-core.log
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-23T03:22:14.402317432Z","level":"INFO","msg":"started logging, with flags","port-filename":"/tmp/tmps2dvy9x1/port-1034100.txt","pid":1034100,"debug":false,"disable-analytics":false}
|
| 2 |
+
{"time":"2024-10-23T03:22:14.402858522Z","level":"INFO","msg":"FeatureState","shutdownOnParentExitEnabled":false}
|
| 3 |
+
{"time":"2024-10-23T03:22:14.407755671Z","level":"INFO","msg":"Will exit if parent process dies.","ppid":1034100}
|
| 4 |
+
{"time":"2024-10-23T03:22:14.40773793Z","level":"INFO","msg":"server is running","addr":{"IP":"127.0.0.1","Port":41495,"Zone":""}}
|
| 5 |
+
{"time":"2024-10-23T03:22:14.557294663Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"127.0.0.1:58122"}
|
| 6 |
+
{"time":"2024-10-23T03:22:15.098249959Z","level":"INFO","msg":"handleInformInit: received","streamId":"HCPflat_raw_83810","id":"127.0.0.1:58122"}
|
| 7 |
+
{"time":"2024-10-23T03:22:15.159214287Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"HCPflat_raw_83810","id":"127.0.0.1:58122"}
|
| 8 |
+
{"time":"2024-10-23T03:24:02.210685377Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"HCPflat_raw_83810","id":"127.0.0.1:58122"}
|
| 9 |
+
{"time":"2024-10-23T03:24:02.210981773Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"HCPflat_raw_83810","id":"127.0.0.1:58122"}
|
| 10 |
+
{"time":"2024-10-23T03:24:02.366411684Z","level":"INFO","msg":"handleInformInit: received","streamId":"HCPflat_large_gsrFalse__HCP_FT_14592","id":"127.0.0.1:58122"}
|
| 11 |
+
{"time":"2024-10-23T03:24:02.542190524Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"HCPflat_large_gsrFalse__HCP_FT_14592","id":"127.0.0.1:58122"}
|
| 12 |
+
{"time":"2024-10-23T03:36:47.830235978Z","level":"INFO","msg":"Parent process exited, terminating service process."}
|
fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug-internal.log
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-23T03:22:15.107200519Z","level":"INFO","msg":"using version","core version":"0.18.3"}
|
| 2 |
+
{"time":"2024-10-23T03:22:15.107222919Z","level":"INFO","msg":"created symlink","path":"/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug-core.log"}
|
| 3 |
+
{"time":"2024-10-23T03:22:15.12212287Z","level":"ERROR","msg":"dialing: google: could not find default credentials. See https://cloud.google.com/docs/authentication/external/set-up-adc for more information"}
|
| 4 |
+
{"time":"2024-10-23T03:22:15.159184506Z","level":"INFO","msg":"created new stream","id":"HCPflat_raw_83810"}
|
| 5 |
+
{"time":"2024-10-23T03:22:15.159208527Z","level":"INFO","msg":"stream: started","id":"HCPflat_raw_83810"}
|
| 6 |
+
{"time":"2024-10-23T03:22:15.159221567Z","level":"INFO","msg":"writer: Do: started","stream_id":{"value":"HCPflat_raw_83810"}}
|
| 7 |
+
{"time":"2024-10-23T03:22:15.159246878Z","level":"INFO","msg":"sender: started","stream_id":{"value":"HCPflat_raw_83810"}}
|
| 8 |
+
{"time":"2024-10-23T03:22:15.159265338Z","level":"INFO","msg":"handler: started","stream_id":{"value":"HCPflat_raw_83810"}}
|
| 9 |
+
{"time":"2024-10-23T03:22:15.639713975Z","level":"INFO","msg":"wandb-core","!BADKEY":null}
|
| 10 |
+
{"time":"2024-10-23T03:22:15.65038308Z","level":"INFO","msg":"Starting system monitor"}
|
| 11 |
+
{"time":"2024-10-23T03:22:15.650435701Z","level":"WARN","msg":"handleCodeSave: program relative path is empty"}
|
| 12 |
+
{"time":"2024-10-23T03:22:15.652277478Z","level":"ERROR","msg":"git repo not found","error":"repository does not exist"}
|
| 13 |
+
{"time":"2024-10-23T03:22:16.318300603Z","level":"INFO","msg":"Pausing system monitor"}
|
| 14 |
+
{"time":"2024-10-23T03:23:43.467389909Z","level":"INFO","msg":"Resuming system monitor"}
|
| 15 |
+
{"time":"2024-10-23T03:23:43.474183646Z","level":"INFO","msg":"Pausing system monitor"}
|
| 16 |
+
{"time":"2024-10-23T03:23:54.344053787Z","level":"INFO","msg":"Resuming system monitor"}
|
| 17 |
+
{"time":"2024-10-23T03:23:55.025154766Z","level":"INFO","msg":"Stopping system monitor"}
|
| 18 |
+
{"time":"2024-10-23T03:23:55.038944374Z","level":"INFO","msg":"Stopped system monitor"}
|
| 19 |
+
{"time":"2024-10-23T03:24:00.028953872Z","level":"ERROR","msg":"monitor: gpu: timeout waiting for process to exit"}
|
| 20 |
+
{"time":"2024-10-23T03:24:00.893610579Z","level":"INFO","msg":"fileTransfer: Close: file transfer manager closed"}
|
| 21 |
+
{"time":"2024-10-23T03:24:02.210783329Z","level":"INFO","msg":"stream: closing","id":"HCPflat_raw_83810"}
|
| 22 |
+
{"time":"2024-10-23T03:24:02.21081877Z","level":"INFO","msg":"handler: closed","stream_id":{"value":"HCPflat_raw_83810"}}
|
| 23 |
+
{"time":"2024-10-23T03:24:02.21083897Z","level":"INFO","msg":"writer: Close: closed","stream_id":{"value":"HCPflat_raw_83810"}}
|
| 24 |
+
{"time":"2024-10-23T03:24:02.21085142Z","level":"INFO","msg":"sender: closed","stream_id":{"value":"HCPflat_raw_83810"}}
|
| 25 |
+
{"time":"2024-10-23T03:24:02.210972013Z","level":"INFO","msg":"stream: closed","id":"HCPflat_raw_83810"}
|
fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug.log
ADDED
|
@@ -0,0 +1,59 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
2024-10-23 03:22:15,080 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Current SDK version is 0.18.3
|
| 2 |
+
2024-10-23 03:22:15,081 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Configure stats pid to 1034100
|
| 3 |
+
2024-10-23 03:22:15,081 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Loading settings from /admin/home-ckadirt/.config/wandb/settings
|
| 4 |
+
2024-10-23 03:22:15,081 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Loading settings from /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/settings
|
| 5 |
+
2024-10-23 03:22:15,081 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Loading settings from environment variables: {}
|
| 6 |
+
2024-10-23 03:22:15,081 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Applying setup settings: {'mode': None, '_disable_service': None}
|
| 7 |
+
2024-10-23 03:22:15,081 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Inferring run settings from compute environment: {'program': '<python with no main file>'}
|
| 8 |
+
2024-10-23 03:22:15,081 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Applying login settings: {}
|
| 9 |
+
2024-10-23 03:22:15,081 INFO MainThread:1034100 [wandb_init.py:_log_setup():532] Logging user logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug.log
|
| 10 |
+
2024-10-23 03:22:15,082 INFO MainThread:1034100 [wandb_init.py:_log_setup():533] Logging internal logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug-internal.log
|
| 11 |
+
2024-10-23 03:22:15,082 INFO MainThread:1034100 [wandb_init.py:_jupyter_setup():478] configuring jupyter hooks <wandb.sdk.wandb_init._WandbInit object at 0x7f0957b13950>
|
| 12 |
+
2024-10-23 03:22:15,083 INFO MainThread:1034100 [wandb_init.py:init():617] calling init triggers
|
| 13 |
+
2024-10-23 03:22:15,083 INFO MainThread:1034100 [wandb_init.py:init():624] wandb.init called with sweep_config: {}
|
| 14 |
+
config: {'model_name': 'HCPflat_raw', 'batch_size': 8, 'learning_rate': 0.0001, 'weight_decay': 1e-05, 'num_epochs': 20, 'seed': 42}
|
| 15 |
+
2024-10-23 03:22:15,083 INFO MainThread:1034100 [wandb_init.py:init():667] starting backend
|
| 16 |
+
2024-10-23 03:22:15,083 INFO MainThread:1034100 [wandb_init.py:init():671] sending inform_init request
|
| 17 |
+
2024-10-23 03:22:15,093 INFO MainThread:1034100 [backend.py:_multiprocessing_setup():104] multiprocessing start_methods=fork,spawn,forkserver, using: spawn
|
| 18 |
+
2024-10-23 03:22:15,094 INFO MainThread:1034100 [wandb_init.py:init():684] backend started and connected
|
| 19 |
+
2024-10-23 03:22:15,118 INFO MainThread:1034100 [wandb_run.py:_label_probe_notebook():1346] probe notebook
|
| 20 |
+
2024-10-23 03:22:15,124 INFO MainThread:1034100 [wandb_run.py:_label_probe_notebook():1356] Unable to probe notebook: 'NoneType' object has no attribute 'get'
|
| 21 |
+
2024-10-23 03:22:15,124 INFO MainThread:1034100 [wandb_init.py:init():779] updated telemetry
|
| 22 |
+
2024-10-23 03:22:15,219 INFO MainThread:1034100 [wandb_init.py:init():812] communicating run to backend with 90.0 second timeout
|
| 23 |
+
2024-10-23 03:22:15,624 INFO MainThread:1034100 [wandb_init.py:init():863] starting run threads in backend
|
| 24 |
+
2024-10-23 03:22:16,273 INFO MainThread:1034100 [wandb_run.py:_console_start():2465] atexit reg
|
| 25 |
+
2024-10-23 03:22:16,273 INFO MainThread:1034100 [wandb_run.py:_redirect():2313] redirect: wrap_raw
|
| 26 |
+
2024-10-23 03:22:16,273 INFO MainThread:1034100 [wandb_run.py:_redirect():2378] Wrapping output streams.
|
| 27 |
+
2024-10-23 03:22:16,273 INFO MainThread:1034100 [wandb_run.py:_redirect():2403] Redirects installed.
|
| 28 |
+
2024-10-23 03:22:16,289 INFO MainThread:1034100 [wandb_init.py:init():907] run started, returning control to user process
|
| 29 |
+
2024-10-23 03:22:16,295 INFO MainThread:1034100 [jupyter.py:_save_ipynb():398] looking for notebook: ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.ipynb
|
| 30 |
+
2024-10-23 03:22:16,296 INFO MainThread:1034100 [wandb_init.py:_pause_backend():443] pausing backend
|
| 31 |
+
2024-10-23 03:23:43,467 INFO MainThread:1034100 [wandb_init.py:_resume_backend():448] resuming backend
|
| 32 |
+
2024-10-23 03:23:43,469 INFO MainThread:1034100 [jupyter.py:_save_ipynb():398] looking for notebook: ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.ipynb
|
| 33 |
+
2024-10-23 03:23:43,474 INFO MainThread:1034100 [wandb_init.py:_pause_backend():443] pausing backend
|
| 34 |
+
2024-10-23 03:23:54,343 INFO MainThread:1034100 [wandb_init.py:_resume_backend():448] resuming backend
|
| 35 |
+
2024-10-23 03:23:54,855 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Current SDK version is 0.18.3
|
| 36 |
+
2024-10-23 03:23:54,855 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Configure stats pid to 1034100
|
| 37 |
+
2024-10-23 03:23:54,856 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Loading settings from /admin/home-ckadirt/.config/wandb/settings
|
| 38 |
+
2024-10-23 03:23:54,856 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Loading settings from /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/settings
|
| 39 |
+
2024-10-23 03:23:54,856 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Loading settings from environment variables: {}
|
| 40 |
+
2024-10-23 03:23:54,856 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Applying setup settings: {'mode': None, '_disable_service': None}
|
| 41 |
+
2024-10-23 03:23:54,856 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Inferring run settings from compute environment: {'program': '<python with no main file>'}
|
| 42 |
+
2024-10-23 03:23:54,856 INFO MainThread:1034100 [wandb_setup.py:_flush():79] Applying login settings: {}
|
| 43 |
+
2024-10-23 03:23:54,860 INFO MainThread:1034100 [wandb_init.py:_log_setup():532] Logging user logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_032354-HCPflat_large_gsrFalse__HCP_FT_14592/logs/debug.log
|
| 44 |
+
2024-10-23 03:23:54,865 INFO MainThread:1034100 [wandb_init.py:_log_setup():533] Logging internal logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_032354-HCPflat_large_gsrFalse__HCP_FT_14592/logs/debug-internal.log
|
| 45 |
+
2024-10-23 03:23:54,865 INFO MainThread:1034100 [wandb_init.py:init():617] calling init triggers
|
| 46 |
+
2024-10-23 03:23:54,865 INFO MainThread:1034100 [wandb_init.py:init():624] wandb.init called with sweep_config: {}
|
| 47 |
+
config: {'model_name': 'HCPflat_large_gsrFalse__HCP_FT', 'batch_size': 8, 'learning_rate': 0.0001, 'weight_decay': 1e-05, 'num_epochs': 20, 'seed': 42}
|
| 48 |
+
2024-10-23 03:23:54,865 INFO MainThread:1034100 [wandb_init.py:init():642] re-initializing run, found existing run on stack: HCPflat_raw_83810
|
| 49 |
+
2024-10-23 03:23:54,871 INFO MainThread:1034100 [wandb_run.py:_finish():2164] finishing run ckadirt/fMRI-foundation-model/HCPflat_raw_83810
|
| 50 |
+
2024-10-23 03:23:54,912 INFO MainThread:1034100 [jupyter.py:save_history():488] saving 47 cells to _session_history.ipynb
|
| 51 |
+
2024-10-23 03:23:54,913 INFO MainThread:1034100 [wandb_run.py:_config_callback():1394] config_cb ('_wandb', 'session_history') code/_session_history.ipynb None
|
| 52 |
+
2024-10-23 03:23:55,013 INFO MainThread:1034100 [jupyter.py:_save_ipynb():398] looking for notebook: ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.ipynb
|
| 53 |
+
2024-10-23 03:23:55,017 INFO MainThread:1034100 [wandb_init.py:_jupyter_teardown():460] cleaning up jupyter logic
|
| 54 |
+
2024-10-23 03:23:55,017 INFO MainThread:1034100 [wandb_run.py:_atexit_cleanup():2428] got exitcode: 0
|
| 55 |
+
2024-10-23 03:23:55,024 INFO MainThread:1034100 [wandb_run.py:_restore():2410] restore
|
| 56 |
+
2024-10-23 03:23:55,024 INFO MainThread:1034100 [wandb_run.py:_restore():2416] restore done
|
| 57 |
+
2024-10-23 03:24:02,162 INFO MainThread:1034100 [wandb_run.py:_footer_history_summary_info():4049] rendering history
|
| 58 |
+
2024-10-23 03:24:02,162 INFO MainThread:1034100 [wandb_run.py:_footer_history_summary_info():4081] rendering summary
|
| 59 |
+
2024-10-23 03:24:02,204 INFO MainThread:1034100 [wandb_run.py:_footer_sync_info():4008] logging synced files
|
fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/run-HCPflat_raw_83810.wandb
ADDED
|
Binary file (2.39 kB). View file
|
|
|
fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/tmp/code/_session_history.ipynb
ADDED
|
@@ -0,0 +1,1664 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"cells": [
|
| 3 |
+
{
|
| 4 |
+
"cell_type": "code",
|
| 5 |
+
"execution_count": 1,
|
| 6 |
+
"id": "406d87ac",
|
| 7 |
+
"metadata": {},
|
| 8 |
+
"outputs": [],
|
| 9 |
+
"source": [
|
| 10 |
+
"# Import packages and setup gpu configuration.\n",
|
| 11 |
+
"# This code block shouldnt need to be adjusted!\n",
|
| 12 |
+
"import os\n",
|
| 13 |
+
"import sys\n",
|
| 14 |
+
"import json\n",
|
| 15 |
+
"import yaml\n",
|
| 16 |
+
"import numpy as np\n",
|
| 17 |
+
"import copy\n",
|
| 18 |
+
"import math\n",
|
| 19 |
+
"import time\n",
|
| 20 |
+
"import random\n",
|
| 21 |
+
"from tqdm.auto import tqdm\n",
|
| 22 |
+
"import webdataset as wds\n",
|
| 23 |
+
"import matplotlib.pyplot as plt\n",
|
| 24 |
+
"\n",
|
| 25 |
+
"import torch\n",
|
| 26 |
+
"import torch.nn as nn\n",
|
| 27 |
+
"from torchvision import transforms\n",
|
| 28 |
+
"import utils\n",
|
| 29 |
+
"from mae_utils.flat_models import *\n",
|
| 30 |
+
"import h5py\n",
|
| 31 |
+
"\n",
|
| 32 |
+
"# tf32 data type is faster than standard float32\n",
|
| 33 |
+
"torch.backends.cuda.matmul.allow_tf32 = True\n",
|
| 34 |
+
"# following fixes a Conv3D CUDNN_NOT_SUPPORTED error\n",
|
| 35 |
+
"torch.backends.cudnn.benchmark = True\n",
|
| 36 |
+
"\n",
|
| 37 |
+
"# ## MODEL TO LOAD ##\n",
|
| 38 |
+
"model_name = \"HCPflat_large_gsrFalse_\"\n",
|
| 39 |
+
"parquet_folder = \"epoch99\"\n",
|
| 40 |
+
"\n",
|
| 41 |
+
"# outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 42 |
+
"outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 43 |
+
"\n",
|
| 44 |
+
"print(\"outdir\", outdir)\n",
|
| 45 |
+
"# Load previous config.yaml if available\n",
|
| 46 |
+
"if os.path.exists(f\"{outdir}/config.yaml\"):\n",
|
| 47 |
+
" config = yaml.load(open(f\"{outdir}/config.yaml\", 'r'), Loader=yaml.FullLoader)\n",
|
| 48 |
+
" print(f\"Loaded config.yaml from ckpt folder {outdir}\")\n",
|
| 49 |
+
" # create global variables from the config\n",
|
| 50 |
+
" print(\"\\n__CONFIG__\")\n",
|
| 51 |
+
" for attribute_name in config.keys():\n",
|
| 52 |
+
" print(f\"{attribute_name} = {config[attribute_name]}\")\n",
|
| 53 |
+
" globals()[attribute_name] = config[f'{attribute_name}']\n",
|
| 54 |
+
" print(\"\\n\")\n",
|
| 55 |
+
"\n",
|
| 56 |
+
"world_size = os.getenv('WORLD_SIZE')\n",
|
| 57 |
+
"if world_size is None: \n",
|
| 58 |
+
" world_size = 1\n",
|
| 59 |
+
"else:\n",
|
| 60 |
+
" world_size = int(world_size)\n",
|
| 61 |
+
"print(f\"WORLD_SIZE={world_size}\")\n",
|
| 62 |
+
"\n",
|
| 63 |
+
"if utils.is_interactive():\n",
|
| 64 |
+
" # Following allows you to change functions in models.py or utils.py and \n",
|
| 65 |
+
" # have this notebook automatically update with your revisions\n",
|
| 66 |
+
" %load_ext autoreload\n",
|
| 67 |
+
" %autoreload 2\n",
|
| 68 |
+
"\n",
|
| 69 |
+
"batch_size = probe_batch_size\n",
|
| 70 |
+
"num_epochs = probe_num_epochs\n",
|
| 71 |
+
"\n",
|
| 72 |
+
"data_type = torch.float32 # change depending on your mixed_precision\n",
|
| 73 |
+
"global_batch_size = batch_size * world_size\n",
|
| 74 |
+
"\n",
|
| 75 |
+
"device = torch.device('cuda')\n",
|
| 76 |
+
"\n",
|
| 77 |
+
"hcp_flat_path = \"/weka/proj-medarc/shared/HCP-Flat\"\n",
|
| 78 |
+
"# seed = 42\n",
|
| 79 |
+
"# num_frames = 16\n",
|
| 80 |
+
"# gsr = False\n",
|
| 81 |
+
"# num_workers = 10\n",
|
| 82 |
+
"# batch_size = 128\n",
|
| 83 |
+
"\n",
|
| 84 |
+
"print(\"PID of this process =\",os.getpid())\n",
|
| 85 |
+
"utils.seed_everything(seed)"
|
| 86 |
+
]
|
| 87 |
+
},
|
| 88 |
+
{
|
| 89 |
+
"cell_type": "code",
|
| 90 |
+
"execution_count": 2,
|
| 91 |
+
"id": "3f6365eb",
|
| 92 |
+
"metadata": {},
|
| 93 |
+
"outputs": [],
|
| 94 |
+
"source": [
|
| 95 |
+
"# Import packages and setup gpu configuration.\n",
|
| 96 |
+
"# This code block shouldnt need to be adjusted!\n",
|
| 97 |
+
"import os\n",
|
| 98 |
+
"import sys\n",
|
| 99 |
+
"import json\n",
|
| 100 |
+
"import yaml\n",
|
| 101 |
+
"import numpy as np\n",
|
| 102 |
+
"import copy\n",
|
| 103 |
+
"import math\n",
|
| 104 |
+
"import time\n",
|
| 105 |
+
"import random\n",
|
| 106 |
+
"from tqdm.auto import tqdm\n",
|
| 107 |
+
"import webdataset as wds\n",
|
| 108 |
+
"import matplotlib.pyplot as plt\n",
|
| 109 |
+
"\n",
|
| 110 |
+
"import torch\n",
|
| 111 |
+
"import torch.nn as nn\n",
|
| 112 |
+
"from torchvision import transforms\n",
|
| 113 |
+
"import utils\n",
|
| 114 |
+
"from mae_utils.flat_models import *\n",
|
| 115 |
+
"import h5py\n",
|
| 116 |
+
"\n",
|
| 117 |
+
"# tf32 data type is faster than standard float32\n",
|
| 118 |
+
"torch.backends.cuda.matmul.allow_tf32 = True\n",
|
| 119 |
+
"# following fixes a Conv3D CUDNN_NOT_SUPPORTED error\n",
|
| 120 |
+
"torch.backends.cudnn.benchmark = True\n",
|
| 121 |
+
"\n",
|
| 122 |
+
"# ## MODEL TO LOAD ##\n",
|
| 123 |
+
"model_name = \"HCPflat_large_gsrFalse_\"\n",
|
| 124 |
+
"parquet_folder = \"epoch99\"\n",
|
| 125 |
+
"\n",
|
| 126 |
+
"# outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 127 |
+
"outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 128 |
+
"\n",
|
| 129 |
+
"print(\"outdir\", outdir)\n",
|
| 130 |
+
"# Load previous config.yaml if available\n",
|
| 131 |
+
"if os.path.exists(f\"{outdir}/config.yaml\"):\n",
|
| 132 |
+
" config = yaml.load(open(f\"{outdir}/config.yaml\", 'r'), Loader=yaml.FullLoader)\n",
|
| 133 |
+
" print(f\"Loaded config.yaml from ckpt folder {outdir}\")\n",
|
| 134 |
+
" # create global variables from the config\n",
|
| 135 |
+
" print(\"\\n__CONFIG__\")\n",
|
| 136 |
+
" for attribute_name in config.keys():\n",
|
| 137 |
+
" print(f\"{attribute_name} = {config[attribute_name]}\")\n",
|
| 138 |
+
" globals()[attribute_name] = config[f'{attribute_name}']\n",
|
| 139 |
+
" print(\"\\n\")\n",
|
| 140 |
+
"\n",
|
| 141 |
+
"world_size = os.getenv('WORLD_SIZE')\n",
|
| 142 |
+
"if world_size is None: \n",
|
| 143 |
+
" world_size = 1\n",
|
| 144 |
+
"else:\n",
|
| 145 |
+
" world_size = int(world_size)\n",
|
| 146 |
+
"print(f\"WORLD_SIZE={world_size}\")\n",
|
| 147 |
+
"\n",
|
| 148 |
+
"if utils.is_interactive():\n",
|
| 149 |
+
" # Following allows you to change functions in models.py or utils.py and \n",
|
| 150 |
+
" # have this notebook automatically update with your revisions\n",
|
| 151 |
+
" %load_ext autoreload\n",
|
| 152 |
+
" %autoreload 2\n",
|
| 153 |
+
"\n",
|
| 154 |
+
"batch_size = probe_batch_size\n",
|
| 155 |
+
"num_epochs = probe_num_epochs\n",
|
| 156 |
+
"\n",
|
| 157 |
+
"data_type = torch.float32 # change depending on your mixed_precision\n",
|
| 158 |
+
"global_batch_size = batch_size * world_size\n",
|
| 159 |
+
"\n",
|
| 160 |
+
"device = torch.device('cuda')\n",
|
| 161 |
+
"\n",
|
| 162 |
+
"hcp_flat_path = \"/weka/proj-medarc/shared/HCP-Flat\"\n",
|
| 163 |
+
"# seed = 42\n",
|
| 164 |
+
"# num_frames = 16\n",
|
| 165 |
+
"# gsr = False\n",
|
| 166 |
+
"# num_workers = 10\n",
|
| 167 |
+
"# batch_size = 128\n",
|
| 168 |
+
"\n",
|
| 169 |
+
"print(\"PID of this process =\",os.getpid())\n",
|
| 170 |
+
"utils.seed_everything(seed)"
|
| 171 |
+
]
|
| 172 |
+
},
|
| 173 |
+
{
|
| 174 |
+
"cell_type": "code",
|
| 175 |
+
"execution_count": 3,
|
| 176 |
+
"id": "b96ed0fa",
|
| 177 |
+
"metadata": {},
|
| 178 |
+
"outputs": [],
|
| 179 |
+
"source": [
|
| 180 |
+
"from mae_utils.flat import load_hcp_flat_mask\n",
|
| 181 |
+
"from mae_utils.flat import create_hcp_flat\n",
|
| 182 |
+
"from mae_utils.flat import batch_unmask\n",
|
| 183 |
+
"import mae_utils.visualize as vis\n",
|
| 184 |
+
"\n",
|
| 185 |
+
"flat_mask = load_hcp_flat_mask(hcp_flat_path)\n",
|
| 186 |
+
"\n",
|
| 187 |
+
"model = flat_models.mae_vit_large_fmri(\n",
|
| 188 |
+
" patch_size=patch_size,\n",
|
| 189 |
+
" decoder_embed_dim=decoder_embed_dim,\n",
|
| 190 |
+
" t_patch_size=t_patch_size,\n",
|
| 191 |
+
" pred_t_dim=pred_t_dim,\n",
|
| 192 |
+
" decoder_depth=4,\n",
|
| 193 |
+
" cls_embed=cls_embed,\n",
|
| 194 |
+
" norm_pix_loss=norm_pix_loss,\n",
|
| 195 |
+
" no_qkv_bias=no_qkv_bias,\n",
|
| 196 |
+
" sep_pos_embed=sep_pos_embed,\n",
|
| 197 |
+
" trunc_init=trunc_init,\n",
|
| 198 |
+
" pct_masks_to_decode=pct_masks_to_decode,\n",
|
| 199 |
+
" img_mask=flat_mask,\n",
|
| 200 |
+
")"
|
| 201 |
+
]
|
| 202 |
+
},
|
| 203 |
+
{
|
| 204 |
+
"cell_type": "code",
|
| 205 |
+
"execution_count": 4,
|
| 206 |
+
"id": "2344601f",
|
| 207 |
+
"metadata": {},
|
| 208 |
+
"outputs": [],
|
| 209 |
+
"source": [
|
| 210 |
+
"from mae_utils.flat import load_hcp_flat_mask\n",
|
| 211 |
+
"from mae_utils.flat import create_hcp_flat\n",
|
| 212 |
+
"from mae_utils.flat import batch_unmask\n",
|
| 213 |
+
"import mae_utils.visualize as vis\n",
|
| 214 |
+
"\n",
|
| 215 |
+
"flat_mask = load_hcp_flat_mask(hcp_flat_path)\n",
|
| 216 |
+
"\n",
|
| 217 |
+
"model = flat_models.mae_vit_large_fmri(\n",
|
| 218 |
+
" patch_size=patch_size,\n",
|
| 219 |
+
" decoder_embed_dim=decoder_embed_dim,\n",
|
| 220 |
+
" t_patch_size=t_patch_size,\n",
|
| 221 |
+
" pred_t_dim=pred_t_dim,\n",
|
| 222 |
+
" decoder_depth=4,\n",
|
| 223 |
+
" cls_embed=cls_embed,\n",
|
| 224 |
+
" norm_pix_loss=norm_pix_loss,\n",
|
| 225 |
+
" no_qkv_bias=no_qkv_bias,\n",
|
| 226 |
+
" sep_pos_embed=sep_pos_embed,\n",
|
| 227 |
+
" trunc_init=trunc_init,\n",
|
| 228 |
+
" pct_masks_to_decode=pct_masks_to_decode,\n",
|
| 229 |
+
" img_mask=flat_mask,\n",
|
| 230 |
+
")"
|
| 231 |
+
]
|
| 232 |
+
},
|
| 233 |
+
{
|
| 234 |
+
"cell_type": "code",
|
| 235 |
+
"execution_count": 5,
|
| 236 |
+
"id": "33cf89e4",
|
| 237 |
+
"metadata": {},
|
| 238 |
+
"outputs": [],
|
| 239 |
+
"source": [
|
| 240 |
+
"# Import packages and setup gpu configuration.\n",
|
| 241 |
+
"# This code block shouldnt need to be adjusted!\n",
|
| 242 |
+
"import os\n",
|
| 243 |
+
"import sys\n",
|
| 244 |
+
"import json\n",
|
| 245 |
+
"import yaml\n",
|
| 246 |
+
"import numpy as np\n",
|
| 247 |
+
"import copy\n",
|
| 248 |
+
"import math\n",
|
| 249 |
+
"import time\n",
|
| 250 |
+
"import random\n",
|
| 251 |
+
"from tqdm.auto import tqdm\n",
|
| 252 |
+
"import webdataset as wds\n",
|
| 253 |
+
"import matplotlib.pyplot as plt\n",
|
| 254 |
+
"\n",
|
| 255 |
+
"import torch\n",
|
| 256 |
+
"import torch.nn as nn\n",
|
| 257 |
+
"from torchvision import transforms\n",
|
| 258 |
+
"import utils\n",
|
| 259 |
+
"from mae_utils.flat_models import *\n",
|
| 260 |
+
"import h5py\n",
|
| 261 |
+
"from mae_utils import flat_models\n",
|
| 262 |
+
"\n",
|
| 263 |
+
"# tf32 data type is faster than standard float32\n",
|
| 264 |
+
"torch.backends.cuda.matmul.allow_tf32 = True\n",
|
| 265 |
+
"# following fixes a Conv3D CUDNN_NOT_SUPPORTED error\n",
|
| 266 |
+
"torch.backends.cudnn.benchmark = True\n",
|
| 267 |
+
"\n",
|
| 268 |
+
"# ## MODEL TO LOAD ##\n",
|
| 269 |
+
"model_name = \"HCPflat_large_gsrFalse_\"\n",
|
| 270 |
+
"parquet_folder = \"epoch99\"\n",
|
| 271 |
+
"\n",
|
| 272 |
+
"# outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 273 |
+
"outdir = os.path.abspath(f'checkpoints/{model_name}')\n",
|
| 274 |
+
"\n",
|
| 275 |
+
"print(\"outdir\", outdir)\n",
|
| 276 |
+
"# Load previous config.yaml if available\n",
|
| 277 |
+
"if os.path.exists(f\"{outdir}/config.yaml\"):\n",
|
| 278 |
+
" config = yaml.load(open(f\"{outdir}/config.yaml\", 'r'), Loader=yaml.FullLoader)\n",
|
| 279 |
+
" print(f\"Loaded config.yaml from ckpt folder {outdir}\")\n",
|
| 280 |
+
" # create global variables from the config\n",
|
| 281 |
+
" print(\"\\n__CONFIG__\")\n",
|
| 282 |
+
" for attribute_name in config.keys():\n",
|
| 283 |
+
" print(f\"{attribute_name} = {config[attribute_name]}\")\n",
|
| 284 |
+
" globals()[attribute_name] = config[f'{attribute_name}']\n",
|
| 285 |
+
" print(\"\\n\")\n",
|
| 286 |
+
"\n",
|
| 287 |
+
"world_size = os.getenv('WORLD_SIZE')\n",
|
| 288 |
+
"if world_size is None: \n",
|
| 289 |
+
" world_size = 1\n",
|
| 290 |
+
"else:\n",
|
| 291 |
+
" world_size = int(world_size)\n",
|
| 292 |
+
"print(f\"WORLD_SIZE={world_size}\")\n",
|
| 293 |
+
"\n",
|
| 294 |
+
"if utils.is_interactive():\n",
|
| 295 |
+
" # Following allows you to change functions in models.py or utils.py and \n",
|
| 296 |
+
" # have this notebook automatically update with your revisions\n",
|
| 297 |
+
" %load_ext autoreload\n",
|
| 298 |
+
" %autoreload 2\n",
|
| 299 |
+
"\n",
|
| 300 |
+
"batch_size = probe_batch_size\n",
|
| 301 |
+
"num_epochs = probe_num_epochs\n",
|
| 302 |
+
"\n",
|
| 303 |
+
"data_type = torch.float32 # change depending on your mixed_precision\n",
|
| 304 |
+
"global_batch_size = batch_size * world_size\n",
|
| 305 |
+
"\n",
|
| 306 |
+
"device = torch.device('cuda')\n",
|
| 307 |
+
"\n",
|
| 308 |
+
"hcp_flat_path = \"/weka/proj-medarc/shared/HCP-Flat\"\n",
|
| 309 |
+
"# seed = 42\n",
|
| 310 |
+
"# num_frames = 16\n",
|
| 311 |
+
"# gsr = False\n",
|
| 312 |
+
"# num_workers = 10\n",
|
| 313 |
+
"# batch_size = 128\n",
|
| 314 |
+
"\n",
|
| 315 |
+
"print(\"PID of this process =\",os.getpid())\n",
|
| 316 |
+
"utils.seed_everything(seed)"
|
| 317 |
+
]
|
| 318 |
+
},
|
| 319 |
+
{
|
| 320 |
+
"cell_type": "code",
|
| 321 |
+
"execution_count": 6,
|
| 322 |
+
"id": "bc2281a4",
|
| 323 |
+
"metadata": {},
|
| 324 |
+
"outputs": [],
|
| 325 |
+
"source": [
|
| 326 |
+
"from mae_utils.flat import load_hcp_flat_mask\n",
|
| 327 |
+
"from mae_utils.flat import create_hcp_flat\n",
|
| 328 |
+
"from mae_utils.flat import batch_unmask\n",
|
| 329 |
+
"import mae_utils.visualize as vis\n",
|
| 330 |
+
"\n",
|
| 331 |
+
"flat_mask = load_hcp_flat_mask(hcp_flat_path)\n",
|
| 332 |
+
"\n",
|
| 333 |
+
"model = flat_models.mae_vit_large_fmri(\n",
|
| 334 |
+
" patch_size=patch_size,\n",
|
| 335 |
+
" decoder_embed_dim=decoder_embed_dim,\n",
|
| 336 |
+
" t_patch_size=t_patch_size,\n",
|
| 337 |
+
" pred_t_dim=pred_t_dim,\n",
|
| 338 |
+
" decoder_depth=4,\n",
|
| 339 |
+
" cls_embed=cls_embed,\n",
|
| 340 |
+
" norm_pix_loss=norm_pix_loss,\n",
|
| 341 |
+
" no_qkv_bias=no_qkv_bias,\n",
|
| 342 |
+
" sep_pos_embed=sep_pos_embed,\n",
|
| 343 |
+
" trunc_init=trunc_init,\n",
|
| 344 |
+
" pct_masks_to_decode=pct_masks_to_decode,\n",
|
| 345 |
+
" img_mask=flat_mask,\n",
|
| 346 |
+
")"
|
| 347 |
+
]
|
| 348 |
+
},
|
| 349 |
+
{
|
| 350 |
+
"cell_type": "code",
|
| 351 |
+
"execution_count": 7,
|
| 352 |
+
"id": "acfaeaad",
|
| 353 |
+
"metadata": {},
|
| 354 |
+
"outputs": [],
|
| 355 |
+
"source": [
|
| 356 |
+
"checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')]\n",
|
| 357 |
+
"\n",
|
| 358 |
+
"if utils.is_interactive():\n",
|
| 359 |
+
" latest_checkpoint = \"epoch99.pth\"\n",
|
| 360 |
+
"else:\n",
|
| 361 |
+
" latest_checkpoint = sys.argv[2] \n",
|
| 362 |
+
"print(f\"latest_checkpoint: {latest_checkpoint}\")\n",
|
| 363 |
+
"\n",
|
| 364 |
+
"# Load the checkpoint\n",
|
| 365 |
+
"checkpoint_path = os.path.join(outdir, latest_checkpoint)\n",
|
| 366 |
+
"\n",
|
| 367 |
+
"state = torch.load(checkpoint_path)\n",
|
| 368 |
+
"model.load_state_dict(state[\"model_state_dict\"], strict=False)\n",
|
| 369 |
+
"model.to(device)\n",
|
| 370 |
+
"model.eval()\n",
|
| 371 |
+
"\n",
|
| 372 |
+
"print(f\"\\nLoaded checkpoint {latest_checkpoint} from {outdir}\\n\")"
|
| 373 |
+
]
|
| 374 |
+
},
|
| 375 |
+
{
|
| 376 |
+
"cell_type": "code",
|
| 377 |
+
"execution_count": 8,
|
| 378 |
+
"id": "8ffbe1b4",
|
| 379 |
+
"metadata": {},
|
| 380 |
+
"outputs": [],
|
| 381 |
+
"source": [
|
| 382 |
+
"f_train = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/train_hcp.hdf5', 'r')\n",
|
| 383 |
+
"flatmaps_train = f_train['flatmaps']\n",
|
| 384 |
+
"\n",
|
| 385 |
+
"f_test = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/test_hcp.hdf5', 'r')\n",
|
| 386 |
+
"flatmaps_test = f_test['flatmaps']\n",
|
| 387 |
+
"\n",
|
| 388 |
+
"metadata_train = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_train_HCP.npy', allow_pickle=True)\n",
|
| 389 |
+
"metadata_test = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_test_HCP.npy', allow_pickle=True)"
|
| 390 |
+
]
|
| 391 |
+
},
|
| 392 |
+
{
|
| 393 |
+
"cell_type": "code",
|
| 394 |
+
"execution_count": 9,
|
| 395 |
+
"id": "767e8c90",
|
| 396 |
+
"metadata": {},
|
| 397 |
+
"outputs": [],
|
| 398 |
+
"source": [
|
| 399 |
+
"from mae_utils.flat import load_hcp_flat_mask\n",
|
| 400 |
+
"from mae_utils.flat import create_hcp_flat\n",
|
| 401 |
+
"from mae_utils.flat import batch_unmask\n",
|
| 402 |
+
"import mae_utils.visualize as vis\n",
|
| 403 |
+
"\n",
|
| 404 |
+
"flat_mask = load_hcp_flat_mask(hcp_flat_path)\n",
|
| 405 |
+
"\n",
|
| 406 |
+
"mae_model = flat_models.mae_vit_large_fmri(\n",
|
| 407 |
+
" patch_size=patch_size,\n",
|
| 408 |
+
" decoder_embed_dim=decoder_embed_dim,\n",
|
| 409 |
+
" t_patch_size=t_patch_size,\n",
|
| 410 |
+
" pred_t_dim=pred_t_dim,\n",
|
| 411 |
+
" decoder_depth=4,\n",
|
| 412 |
+
" cls_embed=cls_embed,\n",
|
| 413 |
+
" norm_pix_loss=norm_pix_loss,\n",
|
| 414 |
+
" no_qkv_bias=no_qkv_bias,\n",
|
| 415 |
+
" sep_pos_embed=sep_pos_embed,\n",
|
| 416 |
+
" trunc_init=trunc_init,\n",
|
| 417 |
+
" pct_masks_to_decode=pct_masks_to_decode,\n",
|
| 418 |
+
" img_mask=flat_mask,\n",
|
| 419 |
+
")"
|
| 420 |
+
]
|
| 421 |
+
},
|
| 422 |
+
{
|
| 423 |
+
"cell_type": "code",
|
| 424 |
+
"execution_count": 10,
|
| 425 |
+
"id": "fa59d7d5",
|
| 426 |
+
"metadata": {},
|
| 427 |
+
"outputs": [],
|
| 428 |
+
"source": [
|
| 429 |
+
"checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')]\n",
|
| 430 |
+
"\n",
|
| 431 |
+
"if utils.is_interactive():\n",
|
| 432 |
+
" latest_checkpoint = \"epoch99.pth\"\n",
|
| 433 |
+
"else:\n",
|
| 434 |
+
" latest_checkpoint = sys.argv[2] \n",
|
| 435 |
+
"print(f\"latest_checkpoint: {latest_checkpoint}\")\n",
|
| 436 |
+
"\n",
|
| 437 |
+
"# Load the checkpoint\n",
|
| 438 |
+
"checkpoint_path = os.path.join(outdir, latest_checkpoint)\n",
|
| 439 |
+
"\n",
|
| 440 |
+
"state = torch.load(checkpoint_path)\n",
|
| 441 |
+
"mae_model.load_state_dict(state[\"model_state_dict\"], strict=False)\n",
|
| 442 |
+
"mae_model.to(device)\n",
|
| 443 |
+
"\n",
|
| 444 |
+
"print(f\"\\nLoaded checkpoint {latest_checkpoint} from {outdir}\\n\")"
|
| 445 |
+
]
|
| 446 |
+
},
|
| 447 |
+
{
|
| 448 |
+
"cell_type": "code",
|
| 449 |
+
"execution_count": 11,
|
| 450 |
+
"id": "cea19e70",
|
| 451 |
+
"metadata": {},
|
| 452 |
+
"outputs": [],
|
| 453 |
+
"source": [
|
| 454 |
+
"f_train = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/train_hcp.hdf5', 'r')\n",
|
| 455 |
+
"flatmaps_train = f_train['flatmaps']\n",
|
| 456 |
+
"\n",
|
| 457 |
+
"f_test = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/test_hcp.hdf5', 'r')\n",
|
| 458 |
+
"flatmaps_test = f_test['flatmaps']\n",
|
| 459 |
+
"\n",
|
| 460 |
+
"metadata_train = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_train_HCP.npy', allow_pickle=True)\n",
|
| 461 |
+
"metadata_test = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_test_HCP.npy', allow_pickle=True)"
|
| 462 |
+
]
|
| 463 |
+
},
|
| 464 |
+
{
|
| 465 |
+
"cell_type": "code",
|
| 466 |
+
"execution_count": 12,
|
| 467 |
+
"id": "f1f85737",
|
| 468 |
+
"metadata": {},
|
| 469 |
+
"outputs": [],
|
| 470 |
+
"source": [
|
| 471 |
+
"from torch.utils.data import Dataset, DataLoader\n",
|
| 472 |
+
"\n",
|
| 473 |
+
"class HCPFlatDataset(Dataset):\n",
|
| 474 |
+
" def __init__(self, flatmaps, metadata):\n",
|
| 475 |
+
" self.flatmaps = flatmaps\n",
|
| 476 |
+
" self.metadata = metadata\n",
|
| 477 |
+
"\n",
|
| 478 |
+
" def __len__(self):\n",
|
| 479 |
+
" return len(self.metadata)\n",
|
| 480 |
+
"\n",
|
| 481 |
+
" def __getitem__(self, idx):\n",
|
| 482 |
+
" return self.flatmaps[idx], json.loads(self.metadata[idx])\n",
|
| 483 |
+
"\n",
|
| 484 |
+
"# Loading to cpu for faster training, this can take several minutes. Remove this [:] if you want to move one at the time.\n",
|
| 485 |
+
"train_dataset = HCPFlatDataset(flatmaps_train, metadata_train)\n",
|
| 486 |
+
"train_dl = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)\n",
|
| 487 |
+
"\n",
|
| 488 |
+
"test_dataset = HCPFlatDataset(flatmaps_test, metadata_test)\n",
|
| 489 |
+
"test_dl = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)"
|
| 490 |
+
]
|
| 491 |
+
},
|
| 492 |
+
{
|
| 493 |
+
"cell_type": "code",
|
| 494 |
+
"execution_count": 13,
|
| 495 |
+
"id": "f80c3176",
|
| 496 |
+
"metadata": {},
|
| 497 |
+
"outputs": [],
|
| 498 |
+
"source": [
|
| 499 |
+
"for i in train_dl:\n",
|
| 500 |
+
" break"
|
| 501 |
+
]
|
| 502 |
+
},
|
| 503 |
+
{
|
| 504 |
+
"cell_type": "code",
|
| 505 |
+
"execution_count": 14,
|
| 506 |
+
"id": "eb364065",
|
| 507 |
+
"metadata": {},
|
| 508 |
+
"outputs": [],
|
| 509 |
+
"source": [
|
| 510 |
+
"for i in test_dl:\n",
|
| 511 |
+
" break"
|
| 512 |
+
]
|
| 513 |
+
},
|
| 514 |
+
{
|
| 515 |
+
"cell_type": "code",
|
| 516 |
+
"execution_count": 15,
|
| 517 |
+
"id": "d7581eea",
|
| 518 |
+
"metadata": {},
|
| 519 |
+
"outputs": [
|
| 520 |
+
{
|
| 521 |
+
"name": "stdout",
|
| 522 |
+
"output_type": "stream",
|
| 523 |
+
"text": [
|
| 524 |
+
"[tensor([[[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 525 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 526 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 527 |
+
" ...,\n",
|
| 528 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 529 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 530 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 531 |
+
" \n",
|
| 532 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 533 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 534 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 535 |
+
" ...,\n",
|
| 536 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 537 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 538 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 539 |
+
" \n",
|
| 540 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 541 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 542 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 543 |
+
" ...,\n",
|
| 544 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 545 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 546 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 547 |
+
" \n",
|
| 548 |
+
" ...,\n",
|
| 549 |
+
" \n",
|
| 550 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 551 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 552 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 553 |
+
" ...,\n",
|
| 554 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 555 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 556 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 557 |
+
" \n",
|
| 558 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 559 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 560 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 561 |
+
" ...,\n",
|
| 562 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 563 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 564 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 565 |
+
" \n",
|
| 566 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 567 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 568 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 569 |
+
" ...,\n",
|
| 570 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 571 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 572 |
+
" [0., 0., 0., ..., 0., 0., 0.]]],\n",
|
| 573 |
+
" \n",
|
| 574 |
+
" \n",
|
| 575 |
+
" [[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 576 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 577 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 578 |
+
" ...,\n",
|
| 579 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 580 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 581 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 582 |
+
" \n",
|
| 583 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 584 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 585 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 586 |
+
" ...,\n",
|
| 587 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 588 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 589 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 590 |
+
" \n",
|
| 591 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 592 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 593 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 594 |
+
" ...,\n",
|
| 595 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 596 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 597 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 598 |
+
" \n",
|
| 599 |
+
" ...,\n",
|
| 600 |
+
" \n",
|
| 601 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 602 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 603 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 604 |
+
" ...,\n",
|
| 605 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 606 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 607 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 608 |
+
" \n",
|
| 609 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 610 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 611 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 612 |
+
" ...,\n",
|
| 613 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 614 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 615 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 616 |
+
" \n",
|
| 617 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 618 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 619 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 620 |
+
" ...,\n",
|
| 621 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 622 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 623 |
+
" [0., 0., 0., ..., 0., 0., 0.]]],\n",
|
| 624 |
+
" \n",
|
| 625 |
+
" \n",
|
| 626 |
+
" [[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 627 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 628 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 629 |
+
" ...,\n",
|
| 630 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 631 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 632 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 633 |
+
" \n",
|
| 634 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 635 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 636 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 637 |
+
" ...,\n",
|
| 638 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 639 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 640 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 641 |
+
" \n",
|
| 642 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 643 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 644 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 645 |
+
" ...,\n",
|
| 646 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 647 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 648 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 649 |
+
" \n",
|
| 650 |
+
" ...,\n",
|
| 651 |
+
" \n",
|
| 652 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 653 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 654 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 655 |
+
" ...,\n",
|
| 656 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 657 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 658 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 659 |
+
" \n",
|
| 660 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 661 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 662 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 663 |
+
" ...,\n",
|
| 664 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 665 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 666 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 667 |
+
" \n",
|
| 668 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 669 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 670 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 671 |
+
" ...,\n",
|
| 672 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 673 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 674 |
+
" [0., 0., 0., ..., 0., 0., 0.]]],\n",
|
| 675 |
+
" \n",
|
| 676 |
+
" \n",
|
| 677 |
+
" ...,\n",
|
| 678 |
+
" \n",
|
| 679 |
+
" \n",
|
| 680 |
+
" [[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 681 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 682 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 683 |
+
" ...,\n",
|
| 684 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 685 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 686 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 687 |
+
" \n",
|
| 688 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 689 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 690 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 691 |
+
" ...,\n",
|
| 692 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 693 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 694 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 695 |
+
" \n",
|
| 696 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 697 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 698 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 699 |
+
" ...,\n",
|
| 700 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 701 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 702 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 703 |
+
" \n",
|
| 704 |
+
" ...,\n",
|
| 705 |
+
" \n",
|
| 706 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 707 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 708 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 709 |
+
" ...,\n",
|
| 710 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 711 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 712 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 713 |
+
" \n",
|
| 714 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 715 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 716 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 717 |
+
" ...,\n",
|
| 718 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 719 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 720 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 721 |
+
" \n",
|
| 722 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 723 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 724 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 725 |
+
" ...,\n",
|
| 726 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 727 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 728 |
+
" [0., 0., 0., ..., 0., 0., 0.]]],\n",
|
| 729 |
+
" \n",
|
| 730 |
+
" \n",
|
| 731 |
+
" [[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 732 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 733 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 734 |
+
" ...,\n",
|
| 735 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 736 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 737 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 738 |
+
" \n",
|
| 739 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 740 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 741 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 742 |
+
" ...,\n",
|
| 743 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 744 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 745 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 746 |
+
" \n",
|
| 747 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 748 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 749 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 750 |
+
" ...,\n",
|
| 751 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 752 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 753 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 754 |
+
" \n",
|
| 755 |
+
" ...,\n",
|
| 756 |
+
" \n",
|
| 757 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 758 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 759 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 760 |
+
" ...,\n",
|
| 761 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 762 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 763 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 764 |
+
" \n",
|
| 765 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 766 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 767 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 768 |
+
" ...,\n",
|
| 769 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 770 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 771 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 772 |
+
" \n",
|
| 773 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 774 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 775 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 776 |
+
" ...,\n",
|
| 777 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 778 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 779 |
+
" [0., 0., 0., ..., 0., 0., 0.]]],\n",
|
| 780 |
+
" \n",
|
| 781 |
+
" \n",
|
| 782 |
+
" [[[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 783 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 784 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 785 |
+
" ...,\n",
|
| 786 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 787 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 788 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 789 |
+
" \n",
|
| 790 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 791 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 792 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 793 |
+
" ...,\n",
|
| 794 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 795 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 796 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 797 |
+
" \n",
|
| 798 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 799 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 800 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 801 |
+
" ...,\n",
|
| 802 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 803 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 804 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 805 |
+
" \n",
|
| 806 |
+
" ...,\n",
|
| 807 |
+
" \n",
|
| 808 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 809 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 810 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 811 |
+
" ...,\n",
|
| 812 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 813 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 814 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 815 |
+
" \n",
|
| 816 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 817 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 818 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 819 |
+
" ...,\n",
|
| 820 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 821 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 822 |
+
" [0., 0., 0., ..., 0., 0., 0.]],\n",
|
| 823 |
+
" \n",
|
| 824 |
+
" [[0., 0., 0., ..., 0., 0., 0.],\n",
|
| 825 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 826 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 827 |
+
" ...,\n",
|
| 828 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 829 |
+
" [0., 0., 0., ..., 0., 0., 0.],\n",
|
| 830 |
+
" [0., 0., 0., ..., 0., 0., 0.]]]], dtype=torch.float16),\n",
|
| 831 |
+
" {'key': ['sub-102109_mod-tfMRI_task-WM_mag-3T_dir-LR',\n",
|
| 832 |
+
" 'sub-731140_mod-tfMRI_task-WM_mag-3T_dir-RL',\n",
|
| 833 |
+
" 'sub-149539_mod-tfMRI_task-MOTOR_mag-3T_dir-RL',\n",
|
| 834 |
+
" 'sub-376247_mod-tfMRI_task-WM_mag-3T_dir-RL',\n",
|
| 835 |
+
" 'sub-164939_mod-tfMRI_task-EMOTION_mag-3T_dir-LR',\n",
|
| 836 |
+
" 'sub-198653_mod-tfMRI_task-SOCIAL_mag-3T_dir-RL',\n",
|
| 837 |
+
" 'sub-210011_mod-tfMRI_task-WM_mag-3T_dir-RL',\n",
|
| 838 |
+
" 'sub-356948_mod-tfMRI_task-EMOTION_mag-3T_dir-LR'],\n",
|
| 839 |
+
" 'sub': ['102109',\n",
|
| 840 |
+
" '731140',\n",
|
| 841 |
+
" '149539',\n",
|
| 842 |
+
" '376247',\n",
|
| 843 |
+
" '164939',\n",
|
| 844 |
+
" '198653',\n",
|
| 845 |
+
" '210011',\n",
|
| 846 |
+
" '356948'],\n",
|
| 847 |
+
" 'mod': ['tfMRI',\n",
|
| 848 |
+
" 'tfMRI',\n",
|
| 849 |
+
" 'tfMRI',\n",
|
| 850 |
+
" 'tfMRI',\n",
|
| 851 |
+
" 'tfMRI',\n",
|
| 852 |
+
" 'tfMRI',\n",
|
| 853 |
+
" 'tfMRI',\n",
|
| 854 |
+
" 'tfMRI'],\n",
|
| 855 |
+
" 'task': ['WM', 'WM', 'MOTOR', 'WM', 'EMOTION', 'SOCIAL', 'WM', 'EMOTION'],\n",
|
| 856 |
+
" 'mag': ['3T', '3T', '3T', '3T', '3T', '3T', '3T', '3T'],\n",
|
| 857 |
+
" 'dir': ['LR', 'RL', 'RL', 'RL', 'LR', 'RL', 'RL', 'LR'],\n",
|
| 858 |
+
" 'start': tensor([15, 15, 19, 15, 19, 15, 15, 19]),\n",
|
| 859 |
+
" 'trial_type': ['2bk_tools',\n",
|
| 860 |
+
" '2bk_body',\n",
|
| 861 |
+
" 'lh',\n",
|
| 862 |
+
" '2bk_body',\n",
|
| 863 |
+
" 'neut',\n",
|
| 864 |
+
" 'mental',\n",
|
| 865 |
+
" '2bk_body',\n",
|
| 866 |
+
" 'neut']}]"
|
| 867 |
+
]
|
| 868 |
+
}
|
| 869 |
+
],
|
| 870 |
+
"source": [
|
| 871 |
+
"i"
|
| 872 |
+
]
|
| 873 |
+
},
|
| 874 |
+
{
|
| 875 |
+
"cell_type": "code",
|
| 876 |
+
"execution_count": 16,
|
| 877 |
+
"id": "ae8a1da7",
|
| 878 |
+
"metadata": {},
|
| 879 |
+
"outputs": [
|
| 880 |
+
{
|
| 881 |
+
"name": "stdout",
|
| 882 |
+
"output_type": "stream",
|
| 883 |
+
"text": [
|
| 884 |
+
"torch.Size([8, 16, 144, 320])"
|
| 885 |
+
]
|
| 886 |
+
}
|
| 887 |
+
],
|
| 888 |
+
"source": [
|
| 889 |
+
"i[0].shape"
|
| 890 |
+
]
|
| 891 |
+
},
|
| 892 |
+
{
|
| 893 |
+
"cell_type": "code",
|
| 894 |
+
"execution_count": 17,
|
| 895 |
+
"id": "ce3f32de",
|
| 896 |
+
"metadata": {},
|
| 897 |
+
"outputs": [],
|
| 898 |
+
"source": [
|
| 899 |
+
"mae_model(torch.randn(1,16,144,320)).shape"
|
| 900 |
+
]
|
| 901 |
+
},
|
| 902 |
+
{
|
| 903 |
+
"cell_type": "code",
|
| 904 |
+
"execution_count": 18,
|
| 905 |
+
"id": "ff571604",
|
| 906 |
+
"metadata": {},
|
| 907 |
+
"outputs": [],
|
| 908 |
+
"source": [
|
| 909 |
+
"global_pool"
|
| 910 |
+
]
|
| 911 |
+
},
|
| 912 |
+
{
|
| 913 |
+
"cell_type": "code",
|
| 914 |
+
"execution_count": 19,
|
| 915 |
+
"id": "348c52bc",
|
| 916 |
+
"metadata": {},
|
| 917 |
+
"outputs": [],
|
| 918 |
+
"source": [
|
| 919 |
+
"if os.getenv('global_pool') == \"False\":\n",
|
| 920 |
+
" global_pool = False\n",
|
| 921 |
+
"else:\n",
|
| 922 |
+
" global_pool = True\n",
|
| 923 |
+
"print(f\"global_pool = {global_pool}\")\n",
|
| 924 |
+
"\n",
|
| 925 |
+
"try:\n",
|
| 926 |
+
" gsr\n",
|
| 927 |
+
"except:\n",
|
| 928 |
+
" gsr = True\n",
|
| 929 |
+
" print(\"set gsr to True\")\n",
|
| 930 |
+
"print(f\"gsr = {gsr}\")"
|
| 931 |
+
]
|
| 932 |
+
},
|
| 933 |
+
{
|
| 934 |
+
"cell_type": "code",
|
| 935 |
+
"execution_count": 20,
|
| 936 |
+
"id": "82ccddb7",
|
| 937 |
+
"metadata": {},
|
| 938 |
+
"outputs": [],
|
| 939 |
+
"source": [
|
| 940 |
+
"mae_model(torch.randn(1,16,144,320),global_pool=global_pool, forward_features = True).shape"
|
| 941 |
+
]
|
| 942 |
+
},
|
| 943 |
+
{
|
| 944 |
+
"cell_type": "code",
|
| 945 |
+
"execution_count": 21,
|
| 946 |
+
"id": "856a0186",
|
| 947 |
+
"metadata": {},
|
| 948 |
+
"outputs": [],
|
| 949 |
+
"source": [
|
| 950 |
+
"mae_model(torch.randn(1,1,16,144,320),global_pool=global_pool, forward_features = True).shape"
|
| 951 |
+
]
|
| 952 |
+
},
|
| 953 |
+
{
|
| 954 |
+
"cell_type": "code",
|
| 955 |
+
"execution_count": 22,
|
| 956 |
+
"id": "33e1c8cd",
|
| 957 |
+
"metadata": {},
|
| 958 |
+
"outputs": [
|
| 959 |
+
{
|
| 960 |
+
"name": "stdout",
|
| 961 |
+
"output_type": "stream",
|
| 962 |
+
"text": [
|
| 963 |
+
"torch.Size([1, 1024])"
|
| 964 |
+
]
|
| 965 |
+
}
|
| 966 |
+
],
|
| 967 |
+
"source": [
|
| 968 |
+
"mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape"
|
| 969 |
+
]
|
| 970 |
+
},
|
| 971 |
+
{
|
| 972 |
+
"cell_type": "code",
|
| 973 |
+
"execution_count": 23,
|
| 974 |
+
"id": "c56d3fc0",
|
| 975 |
+
"metadata": {},
|
| 976 |
+
"outputs": [
|
| 977 |
+
{
|
| 978 |
+
"name": "stdout",
|
| 979 |
+
"output_type": "stream",
|
| 980 |
+
"text": [
|
| 981 |
+
"torch.Size([1, 1024])"
|
| 982 |
+
]
|
| 983 |
+
}
|
| 984 |
+
],
|
| 985 |
+
"source": [
|
| 986 |
+
"mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape"
|
| 987 |
+
]
|
| 988 |
+
},
|
| 989 |
+
{
|
| 990 |
+
"cell_type": "code",
|
| 991 |
+
"execution_count": 24,
|
| 992 |
+
"id": "d6395f12",
|
| 993 |
+
"metadata": {},
|
| 994 |
+
"outputs": [],
|
| 995 |
+
"source": [
|
| 996 |
+
"mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:].sum()"
|
| 997 |
+
]
|
| 998 |
+
},
|
| 999 |
+
{
|
| 1000 |
+
"cell_type": "code",
|
| 1001 |
+
"execution_count": 25,
|
| 1002 |
+
"id": "2dca71b2",
|
| 1003 |
+
"metadata": {},
|
| 1004 |
+
"outputs": [],
|
| 1005 |
+
"source": [
|
| 1006 |
+
"list(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:]).sum()"
|
| 1007 |
+
]
|
| 1008 |
+
},
|
| 1009 |
+
{
|
| 1010 |
+
"cell_type": "code",
|
| 1011 |
+
"execution_count": 26,
|
| 1012 |
+
"id": "a25dd030",
|
| 1013 |
+
"metadata": {},
|
| 1014 |
+
"outputs": [
|
| 1015 |
+
{
|
| 1016 |
+
"name": "stdout",
|
| 1017 |
+
"output_type": "stream",
|
| 1018 |
+
"text": [
|
| 1019 |
+
"1024"
|
| 1020 |
+
]
|
| 1021 |
+
}
|
| 1022 |
+
],
|
| 1023 |
+
"source": [
|
| 1024 |
+
"sum(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])"
|
| 1025 |
+
]
|
| 1026 |
+
},
|
| 1027 |
+
{
|
| 1028 |
+
"cell_type": "code",
|
| 1029 |
+
"execution_count": 27,
|
| 1030 |
+
"id": "edd2edd2",
|
| 1031 |
+
"metadata": {},
|
| 1032 |
+
"outputs": [
|
| 1033 |
+
{
|
| 1034 |
+
"name": "stdout",
|
| 1035 |
+
"output_type": "stream",
|
| 1036 |
+
"text": [
|
| 1037 |
+
"np.int64(1024)"
|
| 1038 |
+
]
|
| 1039 |
+
}
|
| 1040 |
+
],
|
| 1041 |
+
"source": [
|
| 1042 |
+
"np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])"
|
| 1043 |
+
]
|
| 1044 |
+
},
|
| 1045 |
+
{
|
| 1046 |
+
"cell_type": "code",
|
| 1047 |
+
"execution_count": 28,
|
| 1048 |
+
"id": "7c934b0f",
|
| 1049 |
+
"metadata": {},
|
| 1050 |
+
"outputs": [],
|
| 1051 |
+
"source": [
|
| 1052 |
+
"class LinearClassifier(nn.Module):\n",
|
| 1053 |
+
" def __init__(self, input_dim, num_classes):\n",
|
| 1054 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1055 |
+
" self.linear = nn.Linear(input_dim, num_classes)\n",
|
| 1056 |
+
" \n",
|
| 1057 |
+
" def forward(self, x):\n",
|
| 1058 |
+
" # Flatten the input except for the batch dimension\n",
|
| 1059 |
+
" x = x.view(x.size(0), -1)\n",
|
| 1060 |
+
" out = self.linear(x)\n",
|
| 1061 |
+
" return out # Raw logits\n",
|
| 1062 |
+
"\n",
|
| 1063 |
+
"# Determine the input dimension from a single sample\n",
|
| 1064 |
+
"# Assuming images are of shape [1, 16, 144, 320]\n",
|
| 1065 |
+
"input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])\n",
|
| 1066 |
+
"print(f\"Input dimension: {input_dim}\")"
|
| 1067 |
+
]
|
| 1068 |
+
},
|
| 1069 |
+
{
|
| 1070 |
+
"cell_type": "code",
|
| 1071 |
+
"execution_count": 29,
|
| 1072 |
+
"id": "bd0fcfa7",
|
| 1073 |
+
"metadata": {},
|
| 1074 |
+
"outputs": [],
|
| 1075 |
+
"source": [
|
| 1076 |
+
"class FullModel(nn.Module):\n",
|
| 1077 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1078 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1079 |
+
" self.lc_model = lc_model\n",
|
| 1080 |
+
" self.mae_model = mae_model\n",
|
| 1081 |
+
" \n",
|
| 1082 |
+
" \n",
|
| 1083 |
+
" def forward(self, x, gsr):\n",
|
| 1084 |
+
" x = self.mae_model(x, global_pool=global_pool, forward_features = True)\n",
|
| 1085 |
+
" x = self.lc_model(x)\n",
|
| 1086 |
+
" return x"
|
| 1087 |
+
]
|
| 1088 |
+
},
|
| 1089 |
+
{
|
| 1090 |
+
"cell_type": "code",
|
| 1091 |
+
"execution_count": 30,
|
| 1092 |
+
"id": "25b06ed9",
|
| 1093 |
+
"metadata": {},
|
| 1094 |
+
"outputs": [],
|
| 1095 |
+
"source": [
|
| 1096 |
+
"class LinearClassifier(nn.Module):\n",
|
| 1097 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1098 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1099 |
+
" self.lc_model = lc_model\n",
|
| 1100 |
+
" \n",
|
| 1101 |
+
" \n",
|
| 1102 |
+
" def forward(self, x):\n",
|
| 1103 |
+
" # Flatten the input except for the batch dimension\n",
|
| 1104 |
+
" x = x.view(x.size(0), -1)\n",
|
| 1105 |
+
" out = self.linear(x)\n",
|
| 1106 |
+
" return out # Raw logits\n",
|
| 1107 |
+
"\n",
|
| 1108 |
+
"# Determine the input dimension from a single sample\n",
|
| 1109 |
+
"# Assuming images are of shape [1, 16, 144, 320]\n",
|
| 1110 |
+
"input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])\n",
|
| 1111 |
+
"print(f\"Input dimension: {input_dim}\")"
|
| 1112 |
+
]
|
| 1113 |
+
},
|
| 1114 |
+
{
|
| 1115 |
+
"cell_type": "code",
|
| 1116 |
+
"execution_count": 31,
|
| 1117 |
+
"id": "97f9bbc3",
|
| 1118 |
+
"metadata": {},
|
| 1119 |
+
"outputs": [],
|
| 1120 |
+
"source": [
|
| 1121 |
+
"# Initialize the model\n",
|
| 1122 |
+
"lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)\n",
|
| 1123 |
+
"\n",
|
| 1124 |
+
"model = FullModel(lc_model, mae_model)\n",
|
| 1125 |
+
"\n",
|
| 1126 |
+
"# Move the model to the GPU\n",
|
| 1127 |
+
"model.to(device)\n",
|
| 1128 |
+
"\n",
|
| 1129 |
+
"# Define loss function\n",
|
| 1130 |
+
"criterion = nn.CrossEntropyLoss()\n",
|
| 1131 |
+
"\n",
|
| 1132 |
+
"# Define optimizer with L2 regularization (weight_decay)\n",
|
| 1133 |
+
"learning_rate = 1e-4\n",
|
| 1134 |
+
"weight_decay = 1e-5 # Adjust based on your needs\n",
|
| 1135 |
+
"optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n",
|
| 1136 |
+
"num_epochs = 20 # Adjust as needed"
|
| 1137 |
+
]
|
| 1138 |
+
},
|
| 1139 |
+
{
|
| 1140 |
+
"cell_type": "code",
|
| 1141 |
+
"execution_count": 32,
|
| 1142 |
+
"id": "313b529e",
|
| 1143 |
+
"metadata": {},
|
| 1144 |
+
"outputs": [],
|
| 1145 |
+
"source": [
|
| 1146 |
+
"from sklearn.preprocessing import LabelEncoder\n",
|
| 1147 |
+
"\n",
|
| 1148 |
+
"INCLUDE_CONDS = {\n",
|
| 1149 |
+
" \"fear\",\n",
|
| 1150 |
+
" \"neut\",\n",
|
| 1151 |
+
" \"math\",\n",
|
| 1152 |
+
" \"story\",\n",
|
| 1153 |
+
" \"lf\",\n",
|
| 1154 |
+
" \"lh\",\n",
|
| 1155 |
+
" \"rf\",\n",
|
| 1156 |
+
" \"rh\",\n",
|
| 1157 |
+
" \"t\",\n",
|
| 1158 |
+
" \"match\",\n",
|
| 1159 |
+
" \"relation\",\n",
|
| 1160 |
+
" \"mental\",\n",
|
| 1161 |
+
" \"rnd\",\n",
|
| 1162 |
+
" \"0bk_body\",\n",
|
| 1163 |
+
" \"2bk_body\",\n",
|
| 1164 |
+
" \"0bk_faces\",\n",
|
| 1165 |
+
" \"2bk_faces\",\n",
|
| 1166 |
+
" \"0bk_places\",\n",
|
| 1167 |
+
" \"2bk_places\",\n",
|
| 1168 |
+
" \"0bk_tools\",\n",
|
| 1169 |
+
" \"2bk_tools\",\n",
|
| 1170 |
+
"}\n",
|
| 1171 |
+
"\n",
|
| 1172 |
+
"# test_data = []\n",
|
| 1173 |
+
"\n",
|
| 1174 |
+
"# # Iterate over the DataLoader with a progress bar\n",
|
| 1175 |
+
"# for sample in tqdm(train_dl, desc=\"Processing samples\"):\n",
|
| 1176 |
+
"# x = sample['image']\n",
|
| 1177 |
+
"# y = sample['meta']['trial_type']\n",
|
| 1178 |
+
"# key = sample['meta']['key']\n",
|
| 1179 |
+
"# print(x.shape, y, key)\n",
|
| 1180 |
+
"# break\n",
|
| 1181 |
+
"# Initialize the label encoder\n",
|
| 1182 |
+
"label_encoder = LabelEncoder()\n",
|
| 1183 |
+
"label_encoder.fit(sorted(INCLUDE_CONDS)) # Ensure consistent ordering\n",
|
| 1184 |
+
"\n",
|
| 1185 |
+
"num_classes = len(label_encoder.classes_)\n",
|
| 1186 |
+
"print(f\"Number of classes: {num_classes}\")"
|
| 1187 |
+
]
|
| 1188 |
+
},
|
| 1189 |
+
{
|
| 1190 |
+
"cell_type": "code",
|
| 1191 |
+
"execution_count": 33,
|
| 1192 |
+
"id": "c5f58f30",
|
| 1193 |
+
"metadata": {},
|
| 1194 |
+
"outputs": [],
|
| 1195 |
+
"source": [
|
| 1196 |
+
"f_train = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/train_hcp.hdf5', 'r')\n",
|
| 1197 |
+
"flatmaps_train = f_train['flatmaps']\n",
|
| 1198 |
+
"\n",
|
| 1199 |
+
"f_test = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/test_hcp.hdf5', 'r')\n",
|
| 1200 |
+
"flatmaps_test = f_test['flatmaps']\n",
|
| 1201 |
+
"\n",
|
| 1202 |
+
"metadata_train = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_train_HCP.npy', allow_pickle=True)\n",
|
| 1203 |
+
"metadata_test = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_test_HCP.npy', allow_pickle=True)"
|
| 1204 |
+
]
|
| 1205 |
+
},
|
| 1206 |
+
{
|
| 1207 |
+
"cell_type": "code",
|
| 1208 |
+
"execution_count": 34,
|
| 1209 |
+
"id": "928abadb",
|
| 1210 |
+
"metadata": {},
|
| 1211 |
+
"outputs": [],
|
| 1212 |
+
"source": [
|
| 1213 |
+
"from torch.utils.data import Dataset, DataLoader\n",
|
| 1214 |
+
"\n",
|
| 1215 |
+
"class HCPFlatDataset(Dataset):\n",
|
| 1216 |
+
" def __init__(self, flatmaps, metadata):\n",
|
| 1217 |
+
" self.flatmaps = flatmaps\n",
|
| 1218 |
+
" self.metadata = metadata\n",
|
| 1219 |
+
"\n",
|
| 1220 |
+
" def __len__(self):\n",
|
| 1221 |
+
" return len(self.metadata)\n",
|
| 1222 |
+
"\n",
|
| 1223 |
+
" def __getitem__(self, idx):\n",
|
| 1224 |
+
" return self.flatmaps[idx], json.loads(self.metadata[idx])\n",
|
| 1225 |
+
"\n",
|
| 1226 |
+
"# Loading to cpu for faster training, this can take several minutes. Remove this [:] if you want to move one at the time.\n",
|
| 1227 |
+
"train_dataset = HCPFlatDataset(flatmaps_train, metadata_train)\n",
|
| 1228 |
+
"train_dl = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)\n",
|
| 1229 |
+
"\n",
|
| 1230 |
+
"test_dataset = HCPFlatDataset(flatmaps_test, metadata_test)\n",
|
| 1231 |
+
"test_dl = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)"
|
| 1232 |
+
]
|
| 1233 |
+
},
|
| 1234 |
+
{
|
| 1235 |
+
"cell_type": "code",
|
| 1236 |
+
"execution_count": 35,
|
| 1237 |
+
"id": "ec69248a",
|
| 1238 |
+
"metadata": {},
|
| 1239 |
+
"outputs": [],
|
| 1240 |
+
"source": [
|
| 1241 |
+
"from mae_utils.flat import load_hcp_flat_mask\n",
|
| 1242 |
+
"from mae_utils.flat import create_hcp_flat\n",
|
| 1243 |
+
"from mae_utils.flat import batch_unmask\n",
|
| 1244 |
+
"import mae_utils.visualize as vis\n",
|
| 1245 |
+
"\n",
|
| 1246 |
+
"flat_mask = load_hcp_flat_mask(hcp_flat_path)\n",
|
| 1247 |
+
"\n",
|
| 1248 |
+
"mae_model = flat_models.mae_vit_large_fmri(\n",
|
| 1249 |
+
" patch_size=patch_size,\n",
|
| 1250 |
+
" decoder_embed_dim=decoder_embed_dim,\n",
|
| 1251 |
+
" t_patch_size=t_patch_size,\n",
|
| 1252 |
+
" pred_t_dim=pred_t_dim,\n",
|
| 1253 |
+
" decoder_depth=4,\n",
|
| 1254 |
+
" cls_embed=cls_embed,\n",
|
| 1255 |
+
" norm_pix_loss=norm_pix_loss,\n",
|
| 1256 |
+
" no_qkv_bias=no_qkv_bias,\n",
|
| 1257 |
+
" sep_pos_embed=sep_pos_embed,\n",
|
| 1258 |
+
" trunc_init=trunc_init,\n",
|
| 1259 |
+
" pct_masks_to_decode=pct_masks_to_decode,\n",
|
| 1260 |
+
" img_mask=flat_mask,\n",
|
| 1261 |
+
")"
|
| 1262 |
+
]
|
| 1263 |
+
},
|
| 1264 |
+
{
|
| 1265 |
+
"cell_type": "code",
|
| 1266 |
+
"execution_count": 36,
|
| 1267 |
+
"id": "4e5045c3",
|
| 1268 |
+
"metadata": {},
|
| 1269 |
+
"outputs": [],
|
| 1270 |
+
"source": [
|
| 1271 |
+
"checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')]\n",
|
| 1272 |
+
"\n",
|
| 1273 |
+
"if utils.is_interactive():\n",
|
| 1274 |
+
" latest_checkpoint = \"epoch99.pth\"\n",
|
| 1275 |
+
"else:\n",
|
| 1276 |
+
" latest_checkpoint = sys.argv[2] \n",
|
| 1277 |
+
"print(f\"latest_checkpoint: {latest_checkpoint}\")\n",
|
| 1278 |
+
"\n",
|
| 1279 |
+
"# Load the checkpoint\n",
|
| 1280 |
+
"checkpoint_path = os.path.join(outdir, latest_checkpoint)\n",
|
| 1281 |
+
"\n",
|
| 1282 |
+
"state = torch.load(checkpoint_path)\n",
|
| 1283 |
+
"mae_model.load_state_dict(state[\"model_state_dict\"], strict=False)\n",
|
| 1284 |
+
"mae_model.to(device)\n",
|
| 1285 |
+
"\n",
|
| 1286 |
+
"print(f\"\\nLoaded checkpoint {latest_checkpoint} from {outdir}\\n\")"
|
| 1287 |
+
]
|
| 1288 |
+
},
|
| 1289 |
+
{
|
| 1290 |
+
"cell_type": "code",
|
| 1291 |
+
"execution_count": 37,
|
| 1292 |
+
"id": "0173f847",
|
| 1293 |
+
"metadata": {},
|
| 1294 |
+
"outputs": [],
|
| 1295 |
+
"source": [
|
| 1296 |
+
"class FullModel(nn.Module):\n",
|
| 1297 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1298 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1299 |
+
" self.lc_model = lc_model\n",
|
| 1300 |
+
" self.mae_model = mae_model\n",
|
| 1301 |
+
" \n",
|
| 1302 |
+
" \n",
|
| 1303 |
+
" def forward(self, x, gsr):\n",
|
| 1304 |
+
" x = self.mae_model(x, global_pool=global_pool, forward_features = True)\n",
|
| 1305 |
+
" x = self.lc_model(x)\n",
|
| 1306 |
+
" return x"
|
| 1307 |
+
]
|
| 1308 |
+
},
|
| 1309 |
+
{
|
| 1310 |
+
"cell_type": "code",
|
| 1311 |
+
"execution_count": 38,
|
| 1312 |
+
"id": "551a5976",
|
| 1313 |
+
"metadata": {},
|
| 1314 |
+
"outputs": [],
|
| 1315 |
+
"source": [
|
| 1316 |
+
"class LinearClassifier(nn.Module):\n",
|
| 1317 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1318 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1319 |
+
" self.lc_model = lc_model\n",
|
| 1320 |
+
" \n",
|
| 1321 |
+
" \n",
|
| 1322 |
+
" def forward(self, x):\n",
|
| 1323 |
+
" # Flatten the input except for the batch dimension\n",
|
| 1324 |
+
" x = x.view(x.size(0), -1)\n",
|
| 1325 |
+
" out = self.linear(x)\n",
|
| 1326 |
+
" return out # Raw logits\n",
|
| 1327 |
+
"\n",
|
| 1328 |
+
"# Determine the input dimension from a single sample\n",
|
| 1329 |
+
"# Assuming images are of shape [1, 16, 144, 320]\n",
|
| 1330 |
+
"input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])\n",
|
| 1331 |
+
"print(f\"Input dimension: {input_dim}\")"
|
| 1332 |
+
]
|
| 1333 |
+
},
|
| 1334 |
+
{
|
| 1335 |
+
"cell_type": "code",
|
| 1336 |
+
"execution_count": 39,
|
| 1337 |
+
"id": "a21df922",
|
| 1338 |
+
"metadata": {},
|
| 1339 |
+
"outputs": [],
|
| 1340 |
+
"source": [
|
| 1341 |
+
"# Initialize the model\n",
|
| 1342 |
+
"lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)\n",
|
| 1343 |
+
"\n",
|
| 1344 |
+
"model = FullModel(lc_model, mae_model)\n",
|
| 1345 |
+
"\n",
|
| 1346 |
+
"# Move the model to the GPU\n",
|
| 1347 |
+
"model.to(device)\n",
|
| 1348 |
+
"\n",
|
| 1349 |
+
"# Define loss function\n",
|
| 1350 |
+
"criterion = nn.CrossEntropyLoss()\n",
|
| 1351 |
+
"\n",
|
| 1352 |
+
"# Define optimizer with L2 regularization (weight_decay)\n",
|
| 1353 |
+
"learning_rate = 1e-4\n",
|
| 1354 |
+
"weight_decay = 1e-5 # Adjust based on your needs\n",
|
| 1355 |
+
"optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n",
|
| 1356 |
+
"num_epochs = 20 # Adjust as needed"
|
| 1357 |
+
]
|
| 1358 |
+
},
|
| 1359 |
+
{
|
| 1360 |
+
"cell_type": "code",
|
| 1361 |
+
"execution_count": 40,
|
| 1362 |
+
"id": "7b262408",
|
| 1363 |
+
"metadata": {},
|
| 1364 |
+
"outputs": [],
|
| 1365 |
+
"source": [
|
| 1366 |
+
"class LinearClassifier(nn.Module):\n",
|
| 1367 |
+
" def __init__(self, input_dim, num_classes):\n",
|
| 1368 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1369 |
+
" self.linear = nn.Linear(input_dim, num_classes)\n",
|
| 1370 |
+
" \n",
|
| 1371 |
+
" def forward(self, x):\n",
|
| 1372 |
+
" # Flatten the input except for the batch dimension\n",
|
| 1373 |
+
" x = x.view(x.size(0), -1)\n",
|
| 1374 |
+
" out = self.linear(x)\n",
|
| 1375 |
+
" return out # Raw logits\n",
|
| 1376 |
+
"\n",
|
| 1377 |
+
"# Determine the input dimension from a single sample\n",
|
| 1378 |
+
"# Assuming images are of shape [1, 16, 144, 320]\n",
|
| 1379 |
+
"input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])\n",
|
| 1380 |
+
"print(f\"Input dimension: {input_dim}\")"
|
| 1381 |
+
]
|
| 1382 |
+
},
|
| 1383 |
+
{
|
| 1384 |
+
"cell_type": "code",
|
| 1385 |
+
"execution_count": 41,
|
| 1386 |
+
"id": "7bff8f73",
|
| 1387 |
+
"metadata": {},
|
| 1388 |
+
"outputs": [],
|
| 1389 |
+
"source": [
|
| 1390 |
+
"class FullModel(nn.Module):\n",
|
| 1391 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1392 |
+
" super(LinearClassifier, self).__init__()\n",
|
| 1393 |
+
" self.lc_model = lc_model\n",
|
| 1394 |
+
" self.mae_model = mae_model\n",
|
| 1395 |
+
" \n",
|
| 1396 |
+
" \n",
|
| 1397 |
+
" def forward(self, x, gsr):\n",
|
| 1398 |
+
" x = self.mae_model(x, global_pool=global_pool, forward_features = True)\n",
|
| 1399 |
+
" x = self.lc_model(x)\n",
|
| 1400 |
+
" return x"
|
| 1401 |
+
]
|
| 1402 |
+
},
|
| 1403 |
+
{
|
| 1404 |
+
"cell_type": "code",
|
| 1405 |
+
"execution_count": 42,
|
| 1406 |
+
"id": "697fc651",
|
| 1407 |
+
"metadata": {},
|
| 1408 |
+
"outputs": [],
|
| 1409 |
+
"source": [
|
| 1410 |
+
"# Initialize the model\n",
|
| 1411 |
+
"lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)\n",
|
| 1412 |
+
"\n",
|
| 1413 |
+
"model = FullModel(lc_model, mae_model)\n",
|
| 1414 |
+
"\n",
|
| 1415 |
+
"# Move the model to the GPU\n",
|
| 1416 |
+
"model.to(device)\n",
|
| 1417 |
+
"\n",
|
| 1418 |
+
"# Define loss function\n",
|
| 1419 |
+
"criterion = nn.CrossEntropyLoss()\n",
|
| 1420 |
+
"\n",
|
| 1421 |
+
"# Define optimizer with L2 regularization (weight_decay)\n",
|
| 1422 |
+
"learning_rate = 1e-4\n",
|
| 1423 |
+
"weight_decay = 1e-5 # Adjust based on your needs\n",
|
| 1424 |
+
"optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n",
|
| 1425 |
+
"num_epochs = 20 # Adjust as needed"
|
| 1426 |
+
]
|
| 1427 |
+
},
|
| 1428 |
+
{
|
| 1429 |
+
"cell_type": "code",
|
| 1430 |
+
"execution_count": 43,
|
| 1431 |
+
"id": "99570156",
|
| 1432 |
+
"metadata": {},
|
| 1433 |
+
"outputs": [],
|
| 1434 |
+
"source": [
|
| 1435 |
+
"class FullModel(nn.Module):\n",
|
| 1436 |
+
" def __init__(self, lc_model, mae_model):\n",
|
| 1437 |
+
" super(FullModel, self).__init__()\n",
|
| 1438 |
+
" self.lc_model = lc_model\n",
|
| 1439 |
+
" self.mae_model = mae_model\n",
|
| 1440 |
+
" \n",
|
| 1441 |
+
" \n",
|
| 1442 |
+
" def forward(self, x, gsr):\n",
|
| 1443 |
+
" x = self.mae_model(x, global_pool=global_pool, forward_features = True)\n",
|
| 1444 |
+
" x = self.lc_model(x)\n",
|
| 1445 |
+
" return x"
|
| 1446 |
+
]
|
| 1447 |
+
},
|
| 1448 |
+
{
|
| 1449 |
+
"cell_type": "code",
|
| 1450 |
+
"execution_count": 44,
|
| 1451 |
+
"id": "8e5c0e3e",
|
| 1452 |
+
"metadata": {},
|
| 1453 |
+
"outputs": [],
|
| 1454 |
+
"source": [
|
| 1455 |
+
"# Initialize the model\n",
|
| 1456 |
+
"lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)\n",
|
| 1457 |
+
"\n",
|
| 1458 |
+
"model = FullModel(lc_model, mae_model)\n",
|
| 1459 |
+
"\n",
|
| 1460 |
+
"# Move the model to the GPU\n",
|
| 1461 |
+
"model.to(device)\n",
|
| 1462 |
+
"\n",
|
| 1463 |
+
"# Define loss function\n",
|
| 1464 |
+
"criterion = nn.CrossEntropyLoss()\n",
|
| 1465 |
+
"\n",
|
| 1466 |
+
"# Define optimizer with L2 regularization (weight_decay)\n",
|
| 1467 |
+
"learning_rate = 1e-4\n",
|
| 1468 |
+
"weight_decay = 1e-5 # Adjust based on your needs\n",
|
| 1469 |
+
"optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n",
|
| 1470 |
+
"num_epochs = 20 # Adjust as needed"
|
| 1471 |
+
]
|
| 1472 |
+
},
|
| 1473 |
+
{
|
| 1474 |
+
"cell_type": "code",
|
| 1475 |
+
"execution_count": 45,
|
| 1476 |
+
"id": "dddffb31",
|
| 1477 |
+
"metadata": {},
|
| 1478 |
+
"outputs": [
|
| 1479 |
+
{
|
| 1480 |
+
"data": {
|
| 1481 |
+
"text/html": [
|
| 1482 |
+
"Tracking run with wandb version 0.18.3"
|
| 1483 |
+
],
|
| 1484 |
+
"text/plain": [
|
| 1485 |
+
"<IPython.core.display.HTML object>"
|
| 1486 |
+
]
|
| 1487 |
+
},
|
| 1488 |
+
"metadata": {},
|
| 1489 |
+
"output_type": "display_data"
|
| 1490 |
+
},
|
| 1491 |
+
{
|
| 1492 |
+
"data": {
|
| 1493 |
+
"text/html": [
|
| 1494 |
+
"Run data is saved locally in <code>/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810</code>"
|
| 1495 |
+
],
|
| 1496 |
+
"text/plain": [
|
| 1497 |
+
"<IPython.core.display.HTML object>"
|
| 1498 |
+
]
|
| 1499 |
+
},
|
| 1500 |
+
"metadata": {},
|
| 1501 |
+
"output_type": "display_data"
|
| 1502 |
+
},
|
| 1503 |
+
{
|
| 1504 |
+
"data": {
|
| 1505 |
+
"text/html": [
|
| 1506 |
+
"Syncing run <strong><a href='https://stability.wandb.io/ckadirt/fMRI-foundation-model/runs/HCPflat_raw_83810' target=\"_blank\">HCPflat_raw</a></strong> to <a href='https://stability.wandb.io/ckadirt/fMRI-foundation-model' target=\"_blank\">Weights & Biases</a> (<a href='https://wandb.me/run' target=\"_blank\">docs</a>)<br/>"
|
| 1507 |
+
],
|
| 1508 |
+
"text/plain": [
|
| 1509 |
+
"<IPython.core.display.HTML object>"
|
| 1510 |
+
]
|
| 1511 |
+
},
|
| 1512 |
+
"metadata": {},
|
| 1513 |
+
"output_type": "display_data"
|
| 1514 |
+
},
|
| 1515 |
+
{
|
| 1516 |
+
"data": {
|
| 1517 |
+
"text/html": [
|
| 1518 |
+
" View project at <a href='https://stability.wandb.io/ckadirt/fMRI-foundation-model' target=\"_blank\">https://stability.wandb.io/ckadirt/fMRI-foundation-model</a>"
|
| 1519 |
+
],
|
| 1520 |
+
"text/plain": [
|
| 1521 |
+
"<IPython.core.display.HTML object>"
|
| 1522 |
+
]
|
| 1523 |
+
},
|
| 1524 |
+
"metadata": {},
|
| 1525 |
+
"output_type": "display_data"
|
| 1526 |
+
},
|
| 1527 |
+
{
|
| 1528 |
+
"data": {
|
| 1529 |
+
"text/html": [
|
| 1530 |
+
" View run at <a href='https://stability.wandb.io/ckadirt/fMRI-foundation-model/runs/HCPflat_raw_83810' target=\"_blank\">https://stability.wandb.io/ckadirt/fMRI-foundation-model/runs/HCPflat_raw_83810</a>"
|
| 1531 |
+
],
|
| 1532 |
+
"text/plain": [
|
| 1533 |
+
"<IPython.core.display.HTML object>"
|
| 1534 |
+
]
|
| 1535 |
+
},
|
| 1536 |
+
"metadata": {},
|
| 1537 |
+
"output_type": "display_data"
|
| 1538 |
+
}
|
| 1539 |
+
],
|
| 1540 |
+
"source": [
|
| 1541 |
+
"import wandb\n",
|
| 1542 |
+
"\n",
|
| 1543 |
+
"if utils.is_interactive():\n",
|
| 1544 |
+
" print(\"Running in interactive notebook. Disabling W&B and ckpt saving.\")\n",
|
| 1545 |
+
" wandb_log = True #False\n",
|
| 1546 |
+
" save_ckpt = True #False\n",
|
| 1547 |
+
"\n",
|
| 1548 |
+
"if wandb_log:\n",
|
| 1549 |
+
" wandb_project = 'fMRI-foundation-model'\n",
|
| 1550 |
+
" wandb_config = {\n",
|
| 1551 |
+
" \"model_name\": \"HCPflat_raw\",\n",
|
| 1552 |
+
" \"batch_size\": batch_size,\n",
|
| 1553 |
+
" \"learning_rate\": learning_rate,\n",
|
| 1554 |
+
" \"weight_decay\": weight_decay,\n",
|
| 1555 |
+
" \"num_epochs\": num_epochs,\n",
|
| 1556 |
+
" \"seed\": seed,\n",
|
| 1557 |
+
" }\n",
|
| 1558 |
+
" print(\"wandb_config:\\n\", wandb_config)\n",
|
| 1559 |
+
" random_id = random.randint(0, 100000)\n",
|
| 1560 |
+
" print(\"wandb_id:\", \"HCPflat_raw\" + f\"_{random_id}\")\n",
|
| 1561 |
+
" wandb.init(\n",
|
| 1562 |
+
" id=\"HCPflat_raw\" + f\"_{random_id}\",\n",
|
| 1563 |
+
" project=wandb_project,\n",
|
| 1564 |
+
" name=\"HCPflat_raw\",\n",
|
| 1565 |
+
" config=wandb_config,\n",
|
| 1566 |
+
" resume=\"allow\",\n",
|
| 1567 |
+
" )"
|
| 1568 |
+
]
|
| 1569 |
+
},
|
| 1570 |
+
{
|
| 1571 |
+
"cell_type": "code",
|
| 1572 |
+
"execution_count": 46,
|
| 1573 |
+
"id": "f8a61d67",
|
| 1574 |
+
"metadata": {},
|
| 1575 |
+
"outputs": [],
|
| 1576 |
+
"source": [
|
| 1577 |
+
"import wandb\n",
|
| 1578 |
+
"\n",
|
| 1579 |
+
"if utils.is_interactive():\n",
|
| 1580 |
+
" print(\"Running in interactive notebook. Disabling W&B and ckpt saving.\")\n",
|
| 1581 |
+
" wandb_log = False\n",
|
| 1582 |
+
" save_ckpt = False\n",
|
| 1583 |
+
"\n",
|
| 1584 |
+
"if wandb_log:\n",
|
| 1585 |
+
" wandb_project = 'fMRI-foundation-model'\n",
|
| 1586 |
+
" wandb_config = {\n",
|
| 1587 |
+
" \"model_name\": model_name+'_HCP_FT',\n",
|
| 1588 |
+
" \"batch_size\": batch_size,\n",
|
| 1589 |
+
" \"learning_rate\": learning_rate,\n",
|
| 1590 |
+
" \"weight_decay\": weight_decay,\n",
|
| 1591 |
+
" \"num_epochs\": num_epochs,\n",
|
| 1592 |
+
" \"seed\": seed,\n",
|
| 1593 |
+
" }\n",
|
| 1594 |
+
" print(\"wandb_config:\\n\", wandb_config)\n",
|
| 1595 |
+
" random_id = random.randint(0, 100000)\n",
|
| 1596 |
+
" print(\"wandb_id:\", \"HCPflat_raw\" + f\"_{random_id}\")\n",
|
| 1597 |
+
" wandb.init(\n",
|
| 1598 |
+
" id=model_name+'_HCP_FT' + f\"_{random_id}\",\n",
|
| 1599 |
+
" project=wandb_project,\n",
|
| 1600 |
+
" name=model_name+'_HCP_FT',\n",
|
| 1601 |
+
" config=wandb_config,\n",
|
| 1602 |
+
" resume=\"allow\",\n",
|
| 1603 |
+
" )"
|
| 1604 |
+
]
|
| 1605 |
+
},
|
| 1606 |
+
{
|
| 1607 |
+
"cell_type": "code",
|
| 1608 |
+
"execution_count": 47,
|
| 1609 |
+
"id": "71ac3c4f",
|
| 1610 |
+
"metadata": {},
|
| 1611 |
+
"outputs": [],
|
| 1612 |
+
"source": [
|
| 1613 |
+
"import wandb\n",
|
| 1614 |
+
"\n",
|
| 1615 |
+
"if utils.is_interactive():\n",
|
| 1616 |
+
" print(\"Running in interactive notebook. Disabling W&B and ckpt saving.\")\n",
|
| 1617 |
+
" wandb_log = True\n",
|
| 1618 |
+
" save_ckpt = True\n",
|
| 1619 |
+
"\n",
|
| 1620 |
+
"if wandb_log:\n",
|
| 1621 |
+
" wandb_project = 'fMRI-foundation-model'\n",
|
| 1622 |
+
" wandb_config = {\n",
|
| 1623 |
+
" \"model_name\": model_name+'_HCP_FT',\n",
|
| 1624 |
+
" \"batch_size\": batch_size,\n",
|
| 1625 |
+
" \"learning_rate\": learning_rate,\n",
|
| 1626 |
+
" \"weight_decay\": weight_decay,\n",
|
| 1627 |
+
" \"num_epochs\": num_epochs,\n",
|
| 1628 |
+
" \"seed\": seed,\n",
|
| 1629 |
+
" }\n",
|
| 1630 |
+
" print(\"wandb_config:\\n\", wandb_config)\n",
|
| 1631 |
+
" random_id = random.randint(0, 100000)\n",
|
| 1632 |
+
" print(\"wandb_id:\", \"HCPflat_raw\" + f\"_{random_id}\")\n",
|
| 1633 |
+
" wandb.init(\n",
|
| 1634 |
+
" id=model_name+'_HCP_FT' + f\"_{random_id}\",\n",
|
| 1635 |
+
" project=wandb_project,\n",
|
| 1636 |
+
" name=model_name+'_HCP_FT',\n",
|
| 1637 |
+
" config=wandb_config,\n",
|
| 1638 |
+
" resume=\"allow\",\n",
|
| 1639 |
+
" )"
|
| 1640 |
+
]
|
| 1641 |
+
}
|
| 1642 |
+
],
|
| 1643 |
+
"metadata": {
|
| 1644 |
+
"kernelspec": {
|
| 1645 |
+
"display_name": "Python 3",
|
| 1646 |
+
"language": "python",
|
| 1647 |
+
"name": "python3"
|
| 1648 |
+
},
|
| 1649 |
+
"language_info": {
|
| 1650 |
+
"codemirror_mode": {
|
| 1651 |
+
"name": "ipython",
|
| 1652 |
+
"version": 3
|
| 1653 |
+
},
|
| 1654 |
+
"file_extension": ".py",
|
| 1655 |
+
"mimetype": "text/x-python",
|
| 1656 |
+
"name": "python",
|
| 1657 |
+
"nbconvert_exporter": "python",
|
| 1658 |
+
"pygments_lexer": "ipython3",
|
| 1659 |
+
"version": "3.11.10"
|
| 1660 |
+
}
|
| 1661 |
+
},
|
| 1662 |
+
"nbformat": 4,
|
| 1663 |
+
"nbformat_minor": 5
|
| 1664 |
+
}
|
fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/code/src/HCP_downstream_finetune.py
ADDED
|
@@ -0,0 +1,587 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
# coding: utf-8
|
| 3 |
+
|
| 4 |
+
# In[1]:
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
# Import packages and setup gpu configuration.
|
| 8 |
+
# This code block shouldnt need to be adjusted!
|
| 9 |
+
import os
|
| 10 |
+
import sys
|
| 11 |
+
import json
|
| 12 |
+
import yaml
|
| 13 |
+
import numpy as np
|
| 14 |
+
import copy
|
| 15 |
+
import math
|
| 16 |
+
import time
|
| 17 |
+
import random
|
| 18 |
+
from tqdm.auto import tqdm
|
| 19 |
+
import webdataset as wds
|
| 20 |
+
import matplotlib.pyplot as plt
|
| 21 |
+
|
| 22 |
+
import torch
|
| 23 |
+
import torch.nn as nn
|
| 24 |
+
from torchvision import transforms
|
| 25 |
+
import utils
|
| 26 |
+
from mae_utils.flat_models import *
|
| 27 |
+
import h5py
|
| 28 |
+
from mae_utils import flat_models
|
| 29 |
+
|
| 30 |
+
# tf32 data type is faster than standard float32
|
| 31 |
+
torch.backends.cuda.matmul.allow_tf32 = True
|
| 32 |
+
# following fixes a Conv3D CUDNN_NOT_SUPPORTED error
|
| 33 |
+
torch.backends.cudnn.benchmark = True
|
| 34 |
+
|
| 35 |
+
# ## MODEL TO LOAD ##
|
| 36 |
+
if utils.is_interactive():
|
| 37 |
+
model_name = "HCPflat_large_gsrFalse_"
|
| 38 |
+
else:
|
| 39 |
+
model_name = sys.argv[1]
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
# outdir = os.path.abspath(f'checkpoints/{model_name}')
|
| 43 |
+
outdir = os.path.abspath(f'checkpoints/{model_name}')
|
| 44 |
+
|
| 45 |
+
print("outdir", outdir)
|
| 46 |
+
# Load previous config.yaml if available
|
| 47 |
+
if os.path.exists(f"{outdir}/config.yaml"):
|
| 48 |
+
config = yaml.load(open(f"{outdir}/config.yaml", 'r'), Loader=yaml.FullLoader)
|
| 49 |
+
print(f"Loaded config.yaml from ckpt folder {outdir}")
|
| 50 |
+
# create global variables from the config
|
| 51 |
+
print("\n__CONFIG__")
|
| 52 |
+
for attribute_name in config.keys():
|
| 53 |
+
print(f"{attribute_name} = {config[attribute_name]}")
|
| 54 |
+
globals()[attribute_name] = config[f'{attribute_name}']
|
| 55 |
+
print("\n")
|
| 56 |
+
|
| 57 |
+
world_size = os.getenv('WORLD_SIZE')
|
| 58 |
+
if world_size is None:
|
| 59 |
+
world_size = 1
|
| 60 |
+
else:
|
| 61 |
+
world_size = int(world_size)
|
| 62 |
+
print(f"WORLD_SIZE={world_size}")
|
| 63 |
+
|
| 64 |
+
if utils.is_interactive():
|
| 65 |
+
# Following allows you to change functions in models.py or utils.py and
|
| 66 |
+
# have this notebook automatically update with your revisions
|
| 67 |
+
get_ipython().run_line_magic('load_ext', 'autoreload')
|
| 68 |
+
get_ipython().run_line_magic('autoreload', '2')
|
| 69 |
+
|
| 70 |
+
batch_size = probe_batch_size
|
| 71 |
+
num_epochs = probe_num_epochs
|
| 72 |
+
|
| 73 |
+
data_type = torch.float32 # change depending on your mixed_precision
|
| 74 |
+
global_batch_size = batch_size * world_size
|
| 75 |
+
|
| 76 |
+
device = torch.device('cuda')
|
| 77 |
+
|
| 78 |
+
hcp_flat_path = "/weka/proj-medarc/shared/HCP-Flat"
|
| 79 |
+
# seed = 42
|
| 80 |
+
# num_frames = 16
|
| 81 |
+
# gsr = False
|
| 82 |
+
# num_workers = 10
|
| 83 |
+
# batch_size = 128
|
| 84 |
+
|
| 85 |
+
print("PID of this process =",os.getpid())
|
| 86 |
+
utils.seed_everything(seed)
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
# In[2]:
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
if os.getenv('global_pool') == "False":
|
| 93 |
+
global_pool = False
|
| 94 |
+
else:
|
| 95 |
+
global_pool = True
|
| 96 |
+
print(f"global_pool = {global_pool}")
|
| 97 |
+
|
| 98 |
+
try:
|
| 99 |
+
gsr
|
| 100 |
+
except:
|
| 101 |
+
gsr = True
|
| 102 |
+
print("set gsr to True")
|
| 103 |
+
print(f"gsr = {gsr}")
|
| 104 |
+
|
| 105 |
+
|
| 106 |
+
# In[3]:
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
#### UNCOMMENT THIS TO SAVE THE HCP-FLAT IN HDF5 FORMAT
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
# from torch.utils.data import default_collate
|
| 113 |
+
# from mae_utils.flat import load_hcp_flat_mask
|
| 114 |
+
# from mae_utils.flat import create_hcp_flat
|
| 115 |
+
# from mae_utils.flat import batch_unmask
|
| 116 |
+
# import mae_utils.visualize as vis
|
| 117 |
+
|
| 118 |
+
|
| 119 |
+
# batch_size = 26
|
| 120 |
+
# print(f"changed batch_size to {batch_size}")
|
| 121 |
+
|
| 122 |
+
# ## Test ##
|
| 123 |
+
# datasets_to_include = "HCP"
|
| 124 |
+
# assert "HCP" in datasets_to_include
|
| 125 |
+
# test_dataset = create_hcp_flat(root=hcp_flat_path,
|
| 126 |
+
# clip_mode="event", frames=num_frames, shuffle=False, gsr=gsr, sub_list = 'test')
|
| 127 |
+
# test_dl = wds.WebLoader(
|
| 128 |
+
# test_dataset.batched(batch_size, partial=False, collation_fn=default_collate),
|
| 129 |
+
# batch_size=None,
|
| 130 |
+
# shuffle=False,
|
| 131 |
+
# num_workers=num_workers,
|
| 132 |
+
# pin_memory=True,
|
| 133 |
+
# )
|
| 134 |
+
|
| 135 |
+
# ## Train ##
|
| 136 |
+
# assert "HCP" in datasets_to_include
|
| 137 |
+
# train_dataset = create_hcp_flat(root=hcp_flat_path,
|
| 138 |
+
# clip_mode="event", frames=num_frames, shuffle=False, gsr=gsr, sub_list = 'train')
|
| 139 |
+
# train_dl = wds.WebLoader(
|
| 140 |
+
# train_dataset.batched(batch_size, partial=False, collation_fn=default_collate),
|
| 141 |
+
# batch_size=None,
|
| 142 |
+
# shuffle=False,
|
| 143 |
+
# num_workers=num_workers,
|
| 144 |
+
# pin_memory=True,
|
| 145 |
+
# )
|
| 146 |
+
|
| 147 |
+
# def flatten_meta(meta_dict):
|
| 148 |
+
# """
|
| 149 |
+
# Flatten the meta dictionary by:
|
| 150 |
+
# - Replacing single-item lists with the item itself.
|
| 151 |
+
# - Converting tensors to scalar numbers.
|
| 152 |
+
# """
|
| 153 |
+
# flattened = {}
|
| 154 |
+
# for key, value in meta_dict.items():
|
| 155 |
+
# if isinstance(value, list):
|
| 156 |
+
# if len(value) == 1:
|
| 157 |
+
# flattened[key] = value[0] # Replace list with its single item
|
| 158 |
+
# else:
|
| 159 |
+
# flattened[key] = value # Keep as is if multiple items
|
| 160 |
+
# elif isinstance(value, torch.Tensor):
|
| 161 |
+
# # Convert tensor to scalar
|
| 162 |
+
# if value.numel() == 1:
|
| 163 |
+
# flattened[key] = value.item()
|
| 164 |
+
# else:
|
| 165 |
+
# flattened[key] = value.tolist() # Convert multi-element tensor to list
|
| 166 |
+
# else:
|
| 167 |
+
# flattened[key] = value # Keep the value as is
|
| 168 |
+
# return flattened
|
| 169 |
+
|
| 170 |
+
# import h5py
|
| 171 |
+
# meta_array = np.array([], dtype=object)
|
| 172 |
+
# # Open an HDF5 file in write mode
|
| 173 |
+
# with h5py.File('train_hcp.hdf5', 'w') as h5f:
|
| 174 |
+
# flatmaps_dset = None
|
| 175 |
+
|
| 176 |
+
# total_samples = 0
|
| 177 |
+
|
| 178 |
+
# for i, batch in tqdm(enumerate(train_dl), total = 120000):
|
| 179 |
+
# images = batch['image'][0]
|
| 180 |
+
# meta = batch['meta']
|
| 181 |
+
# batch_size = images.shape[0]
|
| 182 |
+
# meta_serializable = meta.copy()
|
| 183 |
+
|
| 184 |
+
|
| 185 |
+
# # Step 2: Serialize the dictionary to a JSON string
|
| 186 |
+
# meta_str = json.dumps(flatten_meta(meta_serializable), indent=4)
|
| 187 |
+
# meta_array = np.append(meta_array, meta_str)
|
| 188 |
+
# if flatmaps_dset is None:
|
| 189 |
+
# # Initialize datasets with unlimited (None) maxshape along the first axis
|
| 190 |
+
# flatmaps_shape = (0,) + images.shape[1:]
|
| 191 |
+
# flatmaps_maxshape = (None,) + images.shape[1:]
|
| 192 |
+
|
| 193 |
+
# flatmaps_dset = h5f.create_dataset(
|
| 194 |
+
# 'flatmaps',
|
| 195 |
+
# shape=flatmaps_shape,
|
| 196 |
+
# maxshape=flatmaps_maxshape,
|
| 197 |
+
# dtype=np.float16,
|
| 198 |
+
# chunks=True # Enable chunking for efficient resizing
|
| 199 |
+
# )
|
| 200 |
+
|
| 201 |
+
# # Resize datasets to accommodate new data
|
| 202 |
+
# flatmaps_dset.resize(total_samples + batch_size, axis=0)
|
| 203 |
+
|
| 204 |
+
# # Write data to the datasets
|
| 205 |
+
# flatmaps_dset[total_samples:total_samples + batch_size] = images.numpy().astype(np.float16)
|
| 206 |
+
|
| 207 |
+
# total_samples += batch_size
|
| 208 |
+
|
| 209 |
+
# print(f"Processed {total_samples} samples")
|
| 210 |
+
# np.save('metadata_test_HCP.npy', meta_array)
|
| 211 |
+
|
| 212 |
+
|
| 213 |
+
# import h5py
|
| 214 |
+
# meta_array = np.array([], dtype=object)
|
| 215 |
+
# # Open an HDF5 file in write mode
|
| 216 |
+
# with h5py.File('test_hcp.hdf5', 'w') as h5f:
|
| 217 |
+
# flatmaps_dset = None
|
| 218 |
+
|
| 219 |
+
# total_samples = 0
|
| 220 |
+
|
| 221 |
+
# for i, batch in tqdm(enumerate(test_dl), total = 12000):
|
| 222 |
+
# images = batch['image'][0]
|
| 223 |
+
# meta = batch['meta']
|
| 224 |
+
# batch_size = images.shape[0]
|
| 225 |
+
# meta_serializable = meta.copy()
|
| 226 |
+
|
| 227 |
+
|
| 228 |
+
# # Step 2: Serialize the dictionary to a JSON string
|
| 229 |
+
# meta_str = json.dumps(flatten_meta(meta_serializable), indent=4)
|
| 230 |
+
# meta_array = np.append(meta_array, meta_str)
|
| 231 |
+
# if flatmaps_dset is None:
|
| 232 |
+
# # Initialize datasets with unlimited (None) maxshape along the first axis
|
| 233 |
+
# flatmaps_shape = (0,) + images.shape[1:]
|
| 234 |
+
# flatmaps_maxshape = (None,) + images.shape[1:]
|
| 235 |
+
|
| 236 |
+
# flatmaps_dset = h5f.create_dataset(
|
| 237 |
+
# 'flatmaps',
|
| 238 |
+
# shape=flatmaps_shape,
|
| 239 |
+
# maxshape=flatmaps_maxshape,
|
| 240 |
+
# dtype=np.float16,
|
| 241 |
+
# chunks=True # Enable chunking for efficient resizing
|
| 242 |
+
# )
|
| 243 |
+
|
| 244 |
+
# # Resize datasets to accommodate new data
|
| 245 |
+
# flatmaps_dset.resize(total_samples + batch_size, axis=0)
|
| 246 |
+
|
| 247 |
+
# # Write data to the datasets
|
| 248 |
+
# flatmaps_dset[total_samples:total_samples + batch_size] = images.numpy().astype(np.float16)
|
| 249 |
+
|
| 250 |
+
# total_samples += batch_size
|
| 251 |
+
|
| 252 |
+
# print(f"Processed {total_samples} samples")
|
| 253 |
+
# np.save('metadata_train_HCP.npy', meta_array)
|
| 254 |
+
|
| 255 |
+
|
| 256 |
+
# ### Preparing data
|
| 257 |
+
|
| 258 |
+
# In[4]:
|
| 259 |
+
|
| 260 |
+
|
| 261 |
+
from sklearn.preprocessing import LabelEncoder
|
| 262 |
+
|
| 263 |
+
INCLUDE_CONDS = {
|
| 264 |
+
"fear",
|
| 265 |
+
"neut",
|
| 266 |
+
"math",
|
| 267 |
+
"story",
|
| 268 |
+
"lf",
|
| 269 |
+
"lh",
|
| 270 |
+
"rf",
|
| 271 |
+
"rh",
|
| 272 |
+
"t",
|
| 273 |
+
"match",
|
| 274 |
+
"relation",
|
| 275 |
+
"mental",
|
| 276 |
+
"rnd",
|
| 277 |
+
"0bk_body",
|
| 278 |
+
"2bk_body",
|
| 279 |
+
"0bk_faces",
|
| 280 |
+
"2bk_faces",
|
| 281 |
+
"0bk_places",
|
| 282 |
+
"2bk_places",
|
| 283 |
+
"0bk_tools",
|
| 284 |
+
"2bk_tools",
|
| 285 |
+
}
|
| 286 |
+
|
| 287 |
+
# test_data = []
|
| 288 |
+
|
| 289 |
+
# # Iterate over the DataLoader with a progress bar
|
| 290 |
+
# for sample in tqdm(train_dl, desc="Processing samples"):
|
| 291 |
+
# x = sample['image']
|
| 292 |
+
# y = sample['meta']['trial_type']
|
| 293 |
+
# key = sample['meta']['key']
|
| 294 |
+
# print(x.shape, y, key)
|
| 295 |
+
# break
|
| 296 |
+
# Initialize the label encoder
|
| 297 |
+
label_encoder = LabelEncoder()
|
| 298 |
+
label_encoder.fit(sorted(INCLUDE_CONDS)) # Ensure consistent ordering
|
| 299 |
+
|
| 300 |
+
num_classes = len(label_encoder.classes_)
|
| 301 |
+
print(f"Number of classes: {num_classes}")
|
| 302 |
+
|
| 303 |
+
|
| 304 |
+
# In[5]:
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
f_train = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/train_hcp.hdf5', 'r')
|
| 308 |
+
flatmaps_train = f_train['flatmaps']
|
| 309 |
+
|
| 310 |
+
f_test = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/test_hcp.hdf5', 'r')
|
| 311 |
+
flatmaps_test = f_test['flatmaps']
|
| 312 |
+
|
| 313 |
+
metadata_train = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_train_HCP.npy', allow_pickle=True)
|
| 314 |
+
metadata_test = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_test_HCP.npy', allow_pickle=True)
|
| 315 |
+
|
| 316 |
+
|
| 317 |
+
# In[6]:
|
| 318 |
+
|
| 319 |
+
|
| 320 |
+
from torch.utils.data import Dataset, DataLoader
|
| 321 |
+
|
| 322 |
+
class HCPFlatDataset(Dataset):
|
| 323 |
+
def __init__(self, flatmaps, metadata):
|
| 324 |
+
self.flatmaps = flatmaps
|
| 325 |
+
self.metadata = metadata
|
| 326 |
+
|
| 327 |
+
def __len__(self):
|
| 328 |
+
return len(self.metadata)
|
| 329 |
+
|
| 330 |
+
def __getitem__(self, idx):
|
| 331 |
+
return self.flatmaps[idx], json.loads(self.metadata[idx])
|
| 332 |
+
print("Moving datasets to ram")
|
| 333 |
+
# Loading to cpu for faster training, this can take several minutes. Remove this [:] if you want to move one at the time.
|
| 334 |
+
train_dataset = HCPFlatDataset(flatmaps_train, metadata_train)
|
| 335 |
+
train_dl = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)
|
| 336 |
+
|
| 337 |
+
test_dataset = HCPFlatDataset(flatmaps_test, metadata_test)
|
| 338 |
+
test_dl = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
|
| 339 |
+
print("Datasets ready")
|
| 340 |
+
|
| 341 |
+
|
| 342 |
+
# ### Creating and loading Model
|
| 343 |
+
|
| 344 |
+
# In[7]:
|
| 345 |
+
|
| 346 |
+
|
| 347 |
+
from mae_utils.flat import load_hcp_flat_mask
|
| 348 |
+
from mae_utils.flat import create_hcp_flat
|
| 349 |
+
from mae_utils.flat import batch_unmask
|
| 350 |
+
import mae_utils.visualize as vis
|
| 351 |
+
|
| 352 |
+
flat_mask = load_hcp_flat_mask(hcp_flat_path)
|
| 353 |
+
|
| 354 |
+
mae_model = flat_models.mae_vit_large_fmri(
|
| 355 |
+
patch_size=patch_size,
|
| 356 |
+
decoder_embed_dim=decoder_embed_dim,
|
| 357 |
+
t_patch_size=t_patch_size,
|
| 358 |
+
pred_t_dim=pred_t_dim,
|
| 359 |
+
decoder_depth=4,
|
| 360 |
+
cls_embed=cls_embed,
|
| 361 |
+
norm_pix_loss=norm_pix_loss,
|
| 362 |
+
no_qkv_bias=no_qkv_bias,
|
| 363 |
+
sep_pos_embed=sep_pos_embed,
|
| 364 |
+
trunc_init=trunc_init,
|
| 365 |
+
pct_masks_to_decode=pct_masks_to_decode,
|
| 366 |
+
img_mask=flat_mask,
|
| 367 |
+
)
|
| 368 |
+
|
| 369 |
+
|
| 370 |
+
# In[8]:
|
| 371 |
+
|
| 372 |
+
|
| 373 |
+
checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')]
|
| 374 |
+
|
| 375 |
+
if utils.is_interactive():
|
| 376 |
+
latest_checkpoint = "epoch99.pth"
|
| 377 |
+
else:
|
| 378 |
+
latest_checkpoint = sys.argv[2]
|
| 379 |
+
print(f"latest_checkpoint: {latest_checkpoint}")
|
| 380 |
+
|
| 381 |
+
# Load the checkpoint
|
| 382 |
+
checkpoint_path = os.path.join(outdir, latest_checkpoint)
|
| 383 |
+
|
| 384 |
+
state = torch.load(checkpoint_path)
|
| 385 |
+
mae_model.load_state_dict(state["model_state_dict"], strict=False)
|
| 386 |
+
mae_model.to(device)
|
| 387 |
+
|
| 388 |
+
print(f"\nLoaded checkpoint {latest_checkpoint} from {outdir}\n")
|
| 389 |
+
|
| 390 |
+
|
| 391 |
+
# In[9]:
|
| 392 |
+
|
| 393 |
+
|
| 394 |
+
class LinearClassifier(nn.Module):
|
| 395 |
+
def __init__(self, input_dim, num_classes):
|
| 396 |
+
super(LinearClassifier, self).__init__()
|
| 397 |
+
self.linear = nn.Linear(input_dim, num_classes)
|
| 398 |
+
|
| 399 |
+
def forward(self, x):
|
| 400 |
+
# Flatten the input except for the batch dimension
|
| 401 |
+
x = x.view(x.size(0), -1)
|
| 402 |
+
out = self.linear(x)
|
| 403 |
+
return out # Raw logits
|
| 404 |
+
|
| 405 |
+
# Determine the input dimension from a single sample
|
| 406 |
+
# Assuming images are of shape [1, 16, 144, 320]
|
| 407 |
+
input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])
|
| 408 |
+
print(f"Input dimension: {input_dim}")
|
| 409 |
+
|
| 410 |
+
|
| 411 |
+
# In[10]:
|
| 412 |
+
|
| 413 |
+
|
| 414 |
+
class FullModel(nn.Module):
|
| 415 |
+
def __init__(self, lc_model, mae_model):
|
| 416 |
+
super(FullModel, self).__init__()
|
| 417 |
+
self.lc_model = lc_model
|
| 418 |
+
self.mae_model = mae_model
|
| 419 |
+
|
| 420 |
+
|
| 421 |
+
def forward(self, x, gsr):
|
| 422 |
+
x = self.mae_model(x, global_pool=global_pool, forward_features = True)
|
| 423 |
+
x = self.lc_model(x)
|
| 424 |
+
return x
|
| 425 |
+
|
| 426 |
+
|
| 427 |
+
# In[11]:
|
| 428 |
+
|
| 429 |
+
|
| 430 |
+
# Initialize the model
|
| 431 |
+
lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)
|
| 432 |
+
|
| 433 |
+
model = FullModel(lc_model, mae_model)
|
| 434 |
+
|
| 435 |
+
# Move the model to the GPU
|
| 436 |
+
model.to(device)
|
| 437 |
+
|
| 438 |
+
# Define loss function
|
| 439 |
+
criterion = nn.CrossEntropyLoss()
|
| 440 |
+
|
| 441 |
+
# Define optimizer with L2 regularization (weight_decay)
|
| 442 |
+
learning_rate = 1e-4
|
| 443 |
+
weight_decay = 1e-5 # Adjust based on your needs
|
| 444 |
+
optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)
|
| 445 |
+
num_epochs = 20 # Adjust as needed
|
| 446 |
+
|
| 447 |
+
|
| 448 |
+
# ### Data
|
| 449 |
+
|
| 450 |
+
# In[12]:
|
| 451 |
+
|
| 452 |
+
|
| 453 |
+
import wandb
|
| 454 |
+
|
| 455 |
+
if utils.is_interactive():
|
| 456 |
+
print("Running in interactive notebook. Disabling W&B and ckpt saving.")
|
| 457 |
+
wandb_log = True
|
| 458 |
+
save_ckpt = True
|
| 459 |
+
|
| 460 |
+
if wandb_log:
|
| 461 |
+
wandb_project = 'fMRI-foundation-model'
|
| 462 |
+
wandb_config = {
|
| 463 |
+
"model_name": model_name+'_HCP_FT',
|
| 464 |
+
"batch_size": batch_size,
|
| 465 |
+
"learning_rate": learning_rate,
|
| 466 |
+
"weight_decay": weight_decay,
|
| 467 |
+
"num_epochs": num_epochs,
|
| 468 |
+
"seed": seed,
|
| 469 |
+
}
|
| 470 |
+
print("wandb_config:\n", wandb_config)
|
| 471 |
+
random_id = random.randint(0, 100000)
|
| 472 |
+
print("wandb_id:", "HCPflat_raw" + f"_{random_id}")
|
| 473 |
+
wandb.init(
|
| 474 |
+
id=model_name+'_HCP_FT' + f"_{random_id}",
|
| 475 |
+
project=wandb_project,
|
| 476 |
+
name=model_name+'_HCP_FT',
|
| 477 |
+
config=wandb_config,
|
| 478 |
+
resume="allow",
|
| 479 |
+
)
|
| 480 |
+
|
| 481 |
+
|
| 482 |
+
# In[13]:
|
| 483 |
+
|
| 484 |
+
|
| 485 |
+
for epoch in range(num_epochs):
|
| 486 |
+
running_train_loss = 0.0
|
| 487 |
+
correct_train = 0
|
| 488 |
+
total_train = 0
|
| 489 |
+
step = 0
|
| 490 |
+
|
| 491 |
+
# with torch.amp.autocast(device_type='cuda'):
|
| 492 |
+
# Training Phase
|
| 493 |
+
model.train()
|
| 494 |
+
for batch in tqdm(train_dl, desc=f"Epoch {epoch+1}/{num_epochs} - Training"):
|
| 495 |
+
optimizer.zero_grad()
|
| 496 |
+
images = batch[0].to(device).float().unsqueeze(1) #fix this # Shape: [batch_size, 1, 16, 144, 320]
|
| 497 |
+
labels = batch[1]['trial_type'] # List of labels
|
| 498 |
+
|
| 499 |
+
encoded_labels = label_encoder.transform(labels)
|
| 500 |
+
encoded_labels = torch.tensor(encoded_labels, dtype=torch.long).to(device) # Shape: [batch_size]
|
| 501 |
+
|
| 502 |
+
# Forward pass
|
| 503 |
+
outputs = model(images, gsr=gsr) # Shape: [num_train_samples, num_classes]
|
| 504 |
+
|
| 505 |
+
# Compute loss
|
| 506 |
+
loss = criterion(outputs, encoded_labels)
|
| 507 |
+
|
| 508 |
+
# Backward pass and optimization
|
| 509 |
+
loss.backward()
|
| 510 |
+
optimizer.step()
|
| 511 |
+
|
| 512 |
+
# Accumulate loss
|
| 513 |
+
running_train_loss += loss.item() * images.size(0)
|
| 514 |
+
|
| 515 |
+
|
| 516 |
+
# Calculate accuracy
|
| 517 |
+
_, predicted = torch.max(outputs, 1)
|
| 518 |
+
|
| 519 |
+
correct_train += (predicted == encoded_labels).sum().item()
|
| 520 |
+
total_train += encoded_labels.size(0)
|
| 521 |
+
|
| 522 |
+
step = step + 1
|
| 523 |
+
if step % 100 == 0:
|
| 524 |
+
print(f"Step [{step}/{len(train_dl)}] - Training Loss: {loss.item():.4f} - Training Accuracy: {100 * correct_train / total_train:.2f}%")
|
| 525 |
+
# thth
|
| 526 |
+
|
| 527 |
+
epoch_train_loss = running_train_loss / total_train if total_train > 0 else 0.0
|
| 528 |
+
train_accuracy = 100 * correct_train / total_train if total_train > 0 else 0.0
|
| 529 |
+
|
| 530 |
+
# Validation Phase
|
| 531 |
+
model.eval()
|
| 532 |
+
running_val_loss = 0.0
|
| 533 |
+
correct_val = 0
|
| 534 |
+
total_val = 0
|
| 535 |
+
|
| 536 |
+
with torch.no_grad():
|
| 537 |
+
for batch in tqdm(test_dl, desc=f"Epoch {epoch+1}/{num_epochs} - Validation"):
|
| 538 |
+
|
| 539 |
+
images = batch[0].to(device).float().unsqueeze(1) #fix this
|
| 540 |
+
labels = batch[1]['trial_type']
|
| 541 |
+
|
| 542 |
+
# Encode labels to integer indices
|
| 543 |
+
encoded_labels = label_encoder.transform(labels)
|
| 544 |
+
encoded_labels = torch.tensor(encoded_labels, dtype=torch.long).to(device)
|
| 545 |
+
|
| 546 |
+
|
| 547 |
+
# Forward pass
|
| 548 |
+
outputs = model(images, gsr=gsr)
|
| 549 |
+
|
| 550 |
+
# Compute loss
|
| 551 |
+
loss = criterion(outputs, encoded_labels)
|
| 552 |
+
|
| 553 |
+
# Accumulate loss
|
| 554 |
+
running_val_loss += loss.item() * images.size(0)
|
| 555 |
+
|
| 556 |
+
# Calculate accuracy
|
| 557 |
+
_, predicted = torch.max(outputs, 1)
|
| 558 |
+
correct_val += (predicted == encoded_labels).sum().item()
|
| 559 |
+
total_val += encoded_labels.size(0)
|
| 560 |
+
|
| 561 |
+
|
| 562 |
+
|
| 563 |
+
epoch_val_loss = running_val_loss / total_val if total_val > 0 else 0.0
|
| 564 |
+
val_accuracy = 100 * correct_val / total_val if total_val > 0 else 0.0
|
| 565 |
+
|
| 566 |
+
print(f"Epoch [{epoch+1}/{num_epochs}] "
|
| 567 |
+
f"- Training Loss: {epoch_train_loss:.4f}, Training Accuracy: {train_accuracy:.2f}% "
|
| 568 |
+
f"- Validation Loss: {epoch_val_loss:.4f}, Validation Accuracy: {val_accuracy:.2f}%")
|
| 569 |
+
|
| 570 |
+
if wandb_log:
|
| 571 |
+
wandb.log({
|
| 572 |
+
"epoch_train_loss": epoch_train_loss,
|
| 573 |
+
"epoch_val_loss": epoch_val_loss,
|
| 574 |
+
"train_accuracy": train_accuracy,
|
| 575 |
+
"val_accuracy": val_accuracy,
|
| 576 |
+
})
|
| 577 |
+
if save_ckpt:
|
| 578 |
+
outdir = os.path.abspath(f'checkpoints/{model_name+"HCP_FT"}')
|
| 579 |
+
os.makedirs(outdir, exist_ok=True)
|
| 580 |
+
print("outdir", outdir)
|
| 581 |
+
# Save model and config
|
| 582 |
+
torch.save(model.state_dict(), f"{outdir}/model.pth")
|
| 583 |
+
with open(f"{outdir}/config.yaml", 'w') as f:
|
| 584 |
+
yaml.dump(wandb_config, f)
|
| 585 |
+
print(f"Saved model and config to {outdir}")
|
| 586 |
+
|
| 587 |
+
|
fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/output.log
ADDED
|
@@ -0,0 +1,2 @@
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Epoch 1/20 - Training: 1%| | 106/13913 [00:56<1:19:52, 2.88it/s]
|
| 2 |
+
Step [100/13913] - Training Loss: 2.6343 - Training Accuracy: 10.25%
|
fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/requirements.txt
ADDED
|
@@ -0,0 +1,198 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
protobuf==5.28.2
|
| 2 |
+
imageio==2.35.1
|
| 3 |
+
MarkupSafe==3.0.0
|
| 4 |
+
regex==2024.9.11
|
| 5 |
+
matplotlib==3.9.2
|
| 6 |
+
notebook==7.2.2
|
| 7 |
+
debugpy==1.8.6
|
| 8 |
+
aiosignal==1.3.1
|
| 9 |
+
jupyter_core==5.7.2
|
| 10 |
+
torchaudio==2.4.1+cu121
|
| 11 |
+
python-json-logger==2.0.7
|
| 12 |
+
six==1.16.0
|
| 13 |
+
scikit-image==0.24.0
|
| 14 |
+
types-python-dateutil==2.9.0.20241003
|
| 15 |
+
PyYAML==6.0.2
|
| 16 |
+
httpcore==1.0.6
|
| 17 |
+
clip==1.0
|
| 18 |
+
babel==2.16.0
|
| 19 |
+
webcolors==24.8.0
|
| 20 |
+
omegaconf==2.3.0
|
| 21 |
+
webencodings==0.5.1
|
| 22 |
+
kiwisolver==1.4.7
|
| 23 |
+
uri-template==1.3.0
|
| 24 |
+
diffusers==0.23.0
|
| 25 |
+
idna==3.10
|
| 26 |
+
fsspec==2024.9.0
|
| 27 |
+
parso==0.8.4
|
| 28 |
+
setuptools==65.5.0
|
| 29 |
+
tornado==6.4.1
|
| 30 |
+
webdataset==0.2.100
|
| 31 |
+
decord==0.6.0
|
| 32 |
+
nvidia-curand-cu12==10.3.2.106
|
| 33 |
+
ipykernel==6.29.5
|
| 34 |
+
jupyter==1.1.1
|
| 35 |
+
pexpect==4.9.0
|
| 36 |
+
kornia_rs==0.1.5
|
| 37 |
+
iopath==0.1.10
|
| 38 |
+
async-lru==2.0.4
|
| 39 |
+
future==1.0.0
|
| 40 |
+
torchvision==0.19.1+cu121
|
| 41 |
+
botocore==1.34.162
|
| 42 |
+
cycler==0.12.1
|
| 43 |
+
tzdata==2024.2
|
| 44 |
+
jupyter_server_terminals==0.5.3
|
| 45 |
+
click==8.1.7
|
| 46 |
+
einops==0.8.0
|
| 47 |
+
pyzmq==26.2.0
|
| 48 |
+
jupyter_client==8.6.3
|
| 49 |
+
nbconvert==7.16.4
|
| 50 |
+
scikit-learn==1.5.2
|
| 51 |
+
executing==2.1.0
|
| 52 |
+
asttokens==2.4.1
|
| 53 |
+
docker-pycreds==0.4.0
|
| 54 |
+
matplotlib-inline==0.1.7
|
| 55 |
+
overrides==7.7.0
|
| 56 |
+
websocket-client==1.8.0
|
| 57 |
+
nbformat==5.10.4
|
| 58 |
+
elbow==0.1.1
|
| 59 |
+
contourpy==1.3.0
|
| 60 |
+
nvidia-cudnn-cu12==9.1.0.70
|
| 61 |
+
transformers==4.44.2
|
| 62 |
+
gitdb==4.0.11
|
| 63 |
+
jupyterlab_nvdashboard==0.11.0
|
| 64 |
+
lazy_loader==0.4
|
| 65 |
+
jsonpointer==3.0.0
|
| 66 |
+
notebook_shim==0.2.4
|
| 67 |
+
nvidia-nccl-cu12==2.20.5
|
| 68 |
+
ffmpeg-python==0.2.0
|
| 69 |
+
triton==3.0.0
|
| 70 |
+
mistune==3.0.2
|
| 71 |
+
python-dateutil==2.9.0.post0
|
| 72 |
+
beautifulsoup4==4.12.3
|
| 73 |
+
nbclient==0.10.0
|
| 74 |
+
h5py==3.12.1
|
| 75 |
+
ftfy==6.2.3
|
| 76 |
+
zipp==3.20.2
|
| 77 |
+
ptyprocess==0.7.0
|
| 78 |
+
huggingface-hub==0.25.1
|
| 79 |
+
pytz==2024.2
|
| 80 |
+
jupyterlab_pygments==0.3.0
|
| 81 |
+
nvidia-cublas-cu12==12.1.3.1
|
| 82 |
+
pandocfilters==1.5.1
|
| 83 |
+
Jinja2==3.1.4
|
| 84 |
+
arrow==1.3.0
|
| 85 |
+
rpds-py==0.20.0
|
| 86 |
+
jupyter_server==2.14.2
|
| 87 |
+
simplejson==3.19.3
|
| 88 |
+
networkx==3.3
|
| 89 |
+
packaging==24.1
|
| 90 |
+
traitlets==5.14.3
|
| 91 |
+
pandas==2.2.3
|
| 92 |
+
xformers==0.0.22.post7
|
| 93 |
+
lightning-utilities==0.11.7
|
| 94 |
+
tifffile==2024.9.20
|
| 95 |
+
nvidia-cuda-cupti-cu12==12.1.105
|
| 96 |
+
mpmath==1.3.0
|
| 97 |
+
GitPython==3.1.43
|
| 98 |
+
scipy==1.14.1
|
| 99 |
+
jsonschema==4.23.0
|
| 100 |
+
prompt_toolkit==3.0.48
|
| 101 |
+
s3transfer==0.10.2
|
| 102 |
+
multidict==6.1.0
|
| 103 |
+
bleach==6.1.0
|
| 104 |
+
sentry-sdk==2.15.0
|
| 105 |
+
nibabel==5.2.1
|
| 106 |
+
accelerate==1.0.0
|
| 107 |
+
pyarrow==17.0.0
|
| 108 |
+
threadpoolctl==3.5.0
|
| 109 |
+
attrs==24.2.0
|
| 110 |
+
rfc3986-validator==0.1.1
|
| 111 |
+
nvidia-cuda-runtime-cu12==12.1.105
|
| 112 |
+
ipywidgets==8.1.5
|
| 113 |
+
frozenlist==1.4.1
|
| 114 |
+
pycparser==2.22
|
| 115 |
+
jupyterlab_server==2.27.3
|
| 116 |
+
nvidia-cuda-nvrtc-cu12==12.1.105
|
| 117 |
+
yarl==1.13.1
|
| 118 |
+
setproctitle==1.3.3
|
| 119 |
+
isoduration==20.11.0
|
| 120 |
+
Pygments==2.18.0
|
| 121 |
+
jedi==0.19.1
|
| 122 |
+
boto3==1.34.57
|
| 123 |
+
tokenizers==0.19.1
|
| 124 |
+
referencing==0.35.1
|
| 125 |
+
rfc3339-validator==0.1.4
|
| 126 |
+
pillow==10.4.0
|
| 127 |
+
jupyterlab==4.2.5
|
| 128 |
+
stack-data==0.6.3
|
| 129 |
+
h11==0.14.0
|
| 130 |
+
anyio==4.6.0
|
| 131 |
+
nilearn==0.10.4
|
| 132 |
+
nvidia-cusolver-cu12==11.4.5.107
|
| 133 |
+
tinycss2==1.3.0
|
| 134 |
+
defusedxml==0.7.1
|
| 135 |
+
argon2-cffi-bindings==21.2.0
|
| 136 |
+
soupsieve==2.6
|
| 137 |
+
nest-asyncio==1.6.0
|
| 138 |
+
torchmetrics==1.3.0.post0
|
| 139 |
+
tqdm==4.66.5
|
| 140 |
+
cffi==1.17.1
|
| 141 |
+
charset-normalizer==3.3.2
|
| 142 |
+
jsonschema-specifications==2023.12.1
|
| 143 |
+
decorator==5.1.1
|
| 144 |
+
open_clip_torch==2.26.1
|
| 145 |
+
jupyter-events==0.10.0
|
| 146 |
+
smart-open==7.0.5
|
| 147 |
+
antlr4-python3-runtime==4.9.3
|
| 148 |
+
prometheus_client==0.21.0
|
| 149 |
+
kornia==0.7.3
|
| 150 |
+
typing_extensions==4.12.2
|
| 151 |
+
sniffio==1.3.1
|
| 152 |
+
joblib==1.4.2
|
| 153 |
+
comm==0.2.2
|
| 154 |
+
aiohappyeyeballs==2.4.3
|
| 155 |
+
numpy==2.1.2
|
| 156 |
+
braceexpand==0.1.7
|
| 157 |
+
certifi==2024.8.30
|
| 158 |
+
psutil==6.0.0
|
| 159 |
+
pyparsing==3.1.4
|
| 160 |
+
pure_eval==0.2.3
|
| 161 |
+
nvidia-cusparse-cu12==12.1.0.106
|
| 162 |
+
wandb==0.18.3
|
| 163 |
+
urllib3==2.2.3
|
| 164 |
+
smmap==5.0.1
|
| 165 |
+
platformdirs==4.3.6
|
| 166 |
+
torch==2.4.1+cu121
|
| 167 |
+
requests==2.32.3
|
| 168 |
+
json5==0.9.25
|
| 169 |
+
nvidia-nvjitlink-cu12==12.6.77
|
| 170 |
+
jupyterlab_widgets==3.0.13
|
| 171 |
+
lxml==5.3.0
|
| 172 |
+
httpx==0.27.2
|
| 173 |
+
opencv-python==4.6.0.66
|
| 174 |
+
portalocker==2.10.1
|
| 175 |
+
pytorch-lightning==2.0.1
|
| 176 |
+
sympy==1.13.3
|
| 177 |
+
wcwidth==0.2.13
|
| 178 |
+
jmespath==1.0.1
|
| 179 |
+
fqdn==1.5.1
|
| 180 |
+
pynvml==11.5.3
|
| 181 |
+
pip==24.0
|
| 182 |
+
wrapt==1.16.0
|
| 183 |
+
aiohttp==3.10.9
|
| 184 |
+
filelock==3.16.1
|
| 185 |
+
fonttools==4.54.1
|
| 186 |
+
fastjsonschema==2.20.0
|
| 187 |
+
jupyter-console==6.6.3
|
| 188 |
+
widgetsnbextension==4.0.13
|
| 189 |
+
timm==1.0.9
|
| 190 |
+
nvidia-cufft-cu12==11.0.2.54
|
| 191 |
+
ipython==8.28.0
|
| 192 |
+
nvidia-nvtx-cu12==12.1.105
|
| 193 |
+
jupyter-lsp==2.2.5
|
| 194 |
+
safetensors==0.4.5
|
| 195 |
+
terminado==0.18.1
|
| 196 |
+
argon2-cffi==23.1.0
|
| 197 |
+
Send2Trash==1.8.3
|
| 198 |
+
importlib_metadata==8.5.0
|
fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/wandb-metadata.json
ADDED
|
@@ -0,0 +1,131 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"os": "Linux-5.15.0-1058-aws-x86_64-with-glibc2.31",
|
| 3 |
+
"python": "3.11.10",
|
| 4 |
+
"startedAt": "2024-10-23T04:08:30.389906Z",
|
| 5 |
+
"args": [
|
| 6 |
+
"NSDflat_large_gsrFalse_",
|
| 7 |
+
"epoch99.pth"
|
| 8 |
+
],
|
| 9 |
+
"program": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py",
|
| 10 |
+
"codePath": "src/HCP_downstream_finetune.py",
|
| 11 |
+
"git": {
|
| 12 |
+
"remote": "https://github.com/MedARC-AI/fMRI-foundation-model",
|
| 13 |
+
"commit": "b1ba684ae7a5cc4155cc046b0abe613de09bf700"
|
| 14 |
+
},
|
| 15 |
+
"email": "torrico.villanueva.cesar.kadir@gmail.com",
|
| 16 |
+
"root": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src",
|
| 17 |
+
"host": "ip-10-0-139-117",
|
| 18 |
+
"username": "ckadirt",
|
| 19 |
+
"executable": "/admin/home-ckadirt/foundation_env/bin/python",
|
| 20 |
+
"codePathLocal": "HCP_downstream_finetune.py",
|
| 21 |
+
"cpu_count": 96,
|
| 22 |
+
"cpu_count_logical": 192,
|
| 23 |
+
"gpu": "[NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3]",
|
| 24 |
+
"gpu_count": 8,
|
| 25 |
+
"disk": {
|
| 26 |
+
"/": {
|
| 27 |
+
"total": "249555763200",
|
| 28 |
+
"used": "181767704576"
|
| 29 |
+
}
|
| 30 |
+
},
|
| 31 |
+
"memory": {
|
| 32 |
+
"total": "2147443380224"
|
| 33 |
+
},
|
| 34 |
+
"cpu": {
|
| 35 |
+
"count": 96,
|
| 36 |
+
"countLogical": 192
|
| 37 |
+
},
|
| 38 |
+
"gpu_nvidia": [
|
| 39 |
+
{
|
| 40 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 41 |
+
"memoryTotal": "85520809984",
|
| 42 |
+
"cudaCores": 16896,
|
| 43 |
+
"architecture": "Hopper"
|
| 44 |
+
},
|
| 45 |
+
{
|
| 46 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 47 |
+
"memoryTotal": "85520809984",
|
| 48 |
+
"cudaCores": 16896,
|
| 49 |
+
"architecture": "Hopper"
|
| 50 |
+
},
|
| 51 |
+
{
|
| 52 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 53 |
+
"memoryTotal": "85520809984",
|
| 54 |
+
"cudaCores": 16896,
|
| 55 |
+
"architecture": "Hopper"
|
| 56 |
+
},
|
| 57 |
+
{
|
| 58 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 59 |
+
"memoryTotal": "85520809984",
|
| 60 |
+
"cudaCores": 16896,
|
| 61 |
+
"architecture": "Hopper"
|
| 62 |
+
},
|
| 63 |
+
{
|
| 64 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 65 |
+
"memoryTotal": "85520809984",
|
| 66 |
+
"cudaCores": 16896,
|
| 67 |
+
"architecture": "Hopper"
|
| 68 |
+
},
|
| 69 |
+
{
|
| 70 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 71 |
+
"memoryTotal": "85520809984",
|
| 72 |
+
"cudaCores": 16896,
|
| 73 |
+
"architecture": "Hopper"
|
| 74 |
+
},
|
| 75 |
+
{
|
| 76 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 77 |
+
"memoryTotal": "85520809984",
|
| 78 |
+
"cudaCores": 16896,
|
| 79 |
+
"architecture": "Hopper"
|
| 80 |
+
},
|
| 81 |
+
{
|
| 82 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 83 |
+
"memoryTotal": "85520809984",
|
| 84 |
+
"cudaCores": 16896,
|
| 85 |
+
"architecture": "Hopper"
|
| 86 |
+
}
|
| 87 |
+
],
|
| 88 |
+
"slurm": {
|
| 89 |
+
"cluster_name": "sagemaker2",
|
| 90 |
+
"conf": "/opt/slurm/etc/slurm.conf",
|
| 91 |
+
"cpus_on_node": "20",
|
| 92 |
+
"gpus_on_node": "1",
|
| 93 |
+
"gpus_per_task": "1",
|
| 94 |
+
"gtids": "0",
|
| 95 |
+
"job_account": "fmri",
|
| 96 |
+
"job_cpus_per_node": "20",
|
| 97 |
+
"job_end_time": "1729699688",
|
| 98 |
+
"job_gid": "1879800513",
|
| 99 |
+
"job_gpus": "4",
|
| 100 |
+
"job_id": "528150",
|
| 101 |
+
"job_name": "finetuneHCP",
|
| 102 |
+
"job_nodelist": "ip-10-0-139-117",
|
| 103 |
+
"job_num_nodes": "1",
|
| 104 |
+
"job_partition": "p5",
|
| 105 |
+
"job_qos": "idle",
|
| 106 |
+
"job_start_time": "1729656488",
|
| 107 |
+
"job_uid": "1879804696",
|
| 108 |
+
"job_user": "ckadirt",
|
| 109 |
+
"jobid": "528150",
|
| 110 |
+
"localid": "0",
|
| 111 |
+
"mem_per_cpu": "11500",
|
| 112 |
+
"nnodes": "1",
|
| 113 |
+
"node_aliases": "(null)",
|
| 114 |
+
"nodeid": "0",
|
| 115 |
+
"nodelist": "ip-10-0-139-117",
|
| 116 |
+
"nprocs": "1",
|
| 117 |
+
"ntasks": "1",
|
| 118 |
+
"ntasks_per_node": "1",
|
| 119 |
+
"prio_process": "0",
|
| 120 |
+
"procid": "0",
|
| 121 |
+
"script_context": "prolog_task",
|
| 122 |
+
"submit_dir": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src",
|
| 123 |
+
"submit_host": "ip-172-17-12-61",
|
| 124 |
+
"task_pid": "3329662",
|
| 125 |
+
"tasks_per_node": "1",
|
| 126 |
+
"topology_addr": "ip-10-0-139-117",
|
| 127 |
+
"topology_addr_pattern": "node",
|
| 128 |
+
"working_cluster": "sagemaker2:ip-172-17-63-161:6817:9984:109"
|
| 129 |
+
},
|
| 130 |
+
"cudaVersion": "12.2"
|
| 131 |
+
}
|
fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug-core.log
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-23T04:08:29.759432555Z","level":"INFO","msg":"started logging, with flags","port-filename":"/tmp/tmpv0zbfmyf/port-3329596.txt","pid":3329596,"debug":false,"disable-analytics":false}
|
| 2 |
+
{"time":"2024-10-23T04:08:29.759434395Z","level":"INFO","msg":"started logging, with flags","port-filename":"/tmp/tmpdgtkqv8m/port-3329693.txt","pid":3329693,"debug":false,"disable-analytics":false}
|
| 3 |
+
{"time":"2024-10-23T04:08:29.75965549Z","level":"INFO","msg":"FeatureState","shutdownOnParentExitEnabled":false}
|
| 4 |
+
{"time":"2024-10-23T04:08:29.75967947Z","level":"INFO","msg":"FeatureState","shutdownOnParentExitEnabled":false}
|
| 5 |
+
{"time":"2024-10-23T04:08:29.762594934Z","level":"INFO","msg":"Will exit if parent process dies.","ppid":3329693}
|
| 6 |
+
{"time":"2024-10-23T04:08:29.762592144Z","level":"INFO","msg":"server is running","addr":{"IP":"127.0.0.1","Port":43169,"Zone":""}}
|
| 7 |
+
{"time":"2024-10-23T04:08:29.764024661Z","level":"INFO","msg":"Will exit if parent process dies.","ppid":3329596}
|
| 8 |
+
{"time":"2024-10-23T04:08:29.764057981Z","level":"INFO","msg":"server is running","addr":{"IP":"127.0.0.1","Port":45635,"Zone":""}}
|
| 9 |
+
{"time":"2024-10-23T04:08:29.893047528Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"127.0.0.1:44072"}
|
| 10 |
+
{"time":"2024-10-23T04:08:29.893048808Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"127.0.0.1:60086"}
|
| 11 |
+
{"time":"2024-10-23T04:08:30.348153132Z","level":"INFO","msg":"handleInformInit: received","streamId":"HCPflat_large_gsrFalse__HCP_FT_83810","id":"127.0.0.1:60086"}
|
| 12 |
+
{"time":"2024-10-23T04:08:30.390297655Z","level":"INFO","msg":"handleInformInit: received","streamId":"NSDflat_large_gsrFalse__HCP_FT_83810","id":"127.0.0.1:44072"}
|
| 13 |
+
{"time":"2024-10-23T04:08:30.392846482Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"HCPflat_large_gsrFalse__HCP_FT_83810","id":"127.0.0.1:60086"}
|
| 14 |
+
{"time":"2024-10-23T04:08:30.434661309Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"NSDflat_large_gsrFalse__HCP_FT_83810","id":"127.0.0.1:44072"}
|
fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug-internal.log
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-23T04:08:30.398164322Z","level":"INFO","msg":"using version","core version":"0.18.3"}
|
| 2 |
+
{"time":"2024-10-23T04:08:30.398178952Z","level":"INFO","msg":"created symlink","path":"/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug-core.log"}
|
| 3 |
+
{"time":"2024-10-23T04:08:30.400590046Z","level":"ERROR","msg":"dialing: google: could not find default credentials. See https://cloud.google.com/docs/authentication/external/set-up-adc for more information"}
|
| 4 |
+
{"time":"2024-10-23T04:08:30.434629119Z","level":"INFO","msg":"created new stream","id":"NSDflat_large_gsrFalse__HCP_FT_83810"}
|
| 5 |
+
{"time":"2024-10-23T04:08:30.434655859Z","level":"INFO","msg":"stream: started","id":"NSDflat_large_gsrFalse__HCP_FT_83810"}
|
| 6 |
+
{"time":"2024-10-23T04:08:30.434669539Z","level":"INFO","msg":"writer: Do: started","stream_id":{"value":"NSDflat_large_gsrFalse__HCP_FT_83810"}}
|
| 7 |
+
{"time":"2024-10-23T04:08:30.43470082Z","level":"INFO","msg":"handler: started","stream_id":{"value":"NSDflat_large_gsrFalse__HCP_FT_83810"}}
|
| 8 |
+
{"time":"2024-10-23T04:08:30.4347016Z","level":"INFO","msg":"sender: started","stream_id":{"value":"NSDflat_large_gsrFalse__HCP_FT_83810"}}
|
| 9 |
+
{"time":"2024-10-23T04:08:30.880621614Z","level":"INFO","msg":"wandb-core","!BADKEY":null}
|
| 10 |
+
{"time":"2024-10-23T04:08:30.884818882Z","level":"INFO","msg":"Starting system monitor"}
|
| 11 |
+
{"time":"2024-10-23T04:08:30.912744151Z","level":"ERROR","msg":"git repo not found","error":"repository does not exist"}
|
fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug.log
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
2024-10-23 04:08:30,377 INFO MainThread:3329693 [wandb_setup.py:_flush():79] Current SDK version is 0.18.3
|
| 2 |
+
2024-10-23 04:08:30,377 INFO MainThread:3329693 [wandb_setup.py:_flush():79] Configure stats pid to 3329693
|
| 3 |
+
2024-10-23 04:08:30,377 INFO MainThread:3329693 [wandb_setup.py:_flush():79] Loading settings from /admin/home-ckadirt/.config/wandb/settings
|
| 4 |
+
2024-10-23 04:08:30,377 INFO MainThread:3329693 [wandb_setup.py:_flush():79] Loading settings from /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/settings
|
| 5 |
+
2024-10-23 04:08:30,377 INFO MainThread:3329693 [wandb_setup.py:_flush():79] Loading settings from environment variables: {}
|
| 6 |
+
2024-10-23 04:08:30,377 INFO MainThread:3329693 [wandb_setup.py:_flush():79] Applying setup settings: {'mode': None, '_disable_service': None}
|
| 7 |
+
2024-10-23 04:08:30,377 INFO MainThread:3329693 [wandb_setup.py:_flush():79] Inferring run settings from compute environment: {'program_relpath': 'src/HCP_downstream_finetune.py', 'program_abspath': '/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py', 'program': '/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py'}
|
| 8 |
+
2024-10-23 04:08:30,377 INFO MainThread:3329693 [wandb_setup.py:_flush():79] Applying login settings: {}
|
| 9 |
+
2024-10-23 04:08:30,379 INFO MainThread:3329693 [wandb_init.py:_log_setup():532] Logging user logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug.log
|
| 10 |
+
2024-10-23 04:08:30,380 INFO MainThread:3329693 [wandb_init.py:_log_setup():533] Logging internal logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug-internal.log
|
| 11 |
+
2024-10-23 04:08:30,380 INFO MainThread:3329693 [wandb_init.py:init():617] calling init triggers
|
| 12 |
+
2024-10-23 04:08:30,380 INFO MainThread:3329693 [wandb_init.py:init():624] wandb.init called with sweep_config: {}
|
| 13 |
+
config: {'model_name': 'NSDflat_large_gsrFalse__HCP_FT', 'batch_size': 8, 'learning_rate': 0.0001, 'weight_decay': 1e-05, 'num_epochs': 20, 'seed': 42}
|
| 14 |
+
2024-10-23 04:08:30,380 INFO MainThread:3329693 [wandb_init.py:init():667] starting backend
|
| 15 |
+
2024-10-23 04:08:30,380 INFO MainThread:3329693 [wandb_init.py:init():671] sending inform_init request
|
| 16 |
+
2024-10-23 04:08:30,389 INFO MainThread:3329693 [backend.py:_multiprocessing_setup():104] multiprocessing start_methods=fork,spawn,forkserver, using: spawn
|
| 17 |
+
2024-10-23 04:08:30,389 INFO MainThread:3329693 [wandb_init.py:init():684] backend started and connected
|
| 18 |
+
2024-10-23 04:08:30,409 INFO MainThread:3329693 [wandb_init.py:init():779] updated telemetry
|
| 19 |
+
2024-10-23 04:08:30,435 INFO MainThread:3329693 [wandb_init.py:init():812] communicating run to backend with 90.0 second timeout
|
| 20 |
+
2024-10-23 04:08:30,824 INFO MainThread:3329693 [wandb_init.py:init():855] run resumed
|
| 21 |
+
2024-10-23 04:08:30,866 INFO MainThread:3329693 [wandb_init.py:init():863] starting run threads in backend
|
| 22 |
+
2024-10-23 04:08:31,304 INFO MainThread:3329693 [wandb_run.py:_console_start():2465] atexit reg
|
| 23 |
+
2024-10-23 04:08:31,304 INFO MainThread:3329693 [wandb_run.py:_redirect():2313] redirect: wrap_raw
|
| 24 |
+
2024-10-23 04:08:31,304 INFO MainThread:3329693 [wandb_run.py:_redirect():2378] Wrapping output streams.
|
| 25 |
+
2024-10-23 04:08:31,304 INFO MainThread:3329693 [wandb_run.py:_redirect():2403] Redirects installed.
|
| 26 |
+
2024-10-23 04:08:31,310 INFO MainThread:3329693 [wandb_init.py:init():907] run started, returning control to user process
|
fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/run-NSDflat_large_gsrFalse__HCP_FT_83810.wandb
ADDED
|
Binary file (65.5 kB). View file
|
|
|
fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/files/output.log
ADDED
|
File without changes
|
fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/files/requirements.txt
ADDED
|
@@ -0,0 +1,198 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
protobuf==5.28.2
|
| 2 |
+
imageio==2.35.1
|
| 3 |
+
MarkupSafe==3.0.0
|
| 4 |
+
regex==2024.9.11
|
| 5 |
+
matplotlib==3.9.2
|
| 6 |
+
notebook==7.2.2
|
| 7 |
+
debugpy==1.8.6
|
| 8 |
+
aiosignal==1.3.1
|
| 9 |
+
jupyter_core==5.7.2
|
| 10 |
+
torchaudio==2.4.1+cu121
|
| 11 |
+
python-json-logger==2.0.7
|
| 12 |
+
six==1.16.0
|
| 13 |
+
scikit-image==0.24.0
|
| 14 |
+
types-python-dateutil==2.9.0.20241003
|
| 15 |
+
PyYAML==6.0.2
|
| 16 |
+
httpcore==1.0.6
|
| 17 |
+
clip==1.0
|
| 18 |
+
babel==2.16.0
|
| 19 |
+
webcolors==24.8.0
|
| 20 |
+
omegaconf==2.3.0
|
| 21 |
+
webencodings==0.5.1
|
| 22 |
+
kiwisolver==1.4.7
|
| 23 |
+
uri-template==1.3.0
|
| 24 |
+
diffusers==0.23.0
|
| 25 |
+
idna==3.10
|
| 26 |
+
fsspec==2024.9.0
|
| 27 |
+
parso==0.8.4
|
| 28 |
+
setuptools==65.5.0
|
| 29 |
+
tornado==6.4.1
|
| 30 |
+
webdataset==0.2.100
|
| 31 |
+
decord==0.6.0
|
| 32 |
+
nvidia-curand-cu12==10.3.2.106
|
| 33 |
+
ipykernel==6.29.5
|
| 34 |
+
jupyter==1.1.1
|
| 35 |
+
pexpect==4.9.0
|
| 36 |
+
kornia_rs==0.1.5
|
| 37 |
+
iopath==0.1.10
|
| 38 |
+
async-lru==2.0.4
|
| 39 |
+
future==1.0.0
|
| 40 |
+
torchvision==0.19.1+cu121
|
| 41 |
+
botocore==1.34.162
|
| 42 |
+
cycler==0.12.1
|
| 43 |
+
tzdata==2024.2
|
| 44 |
+
jupyter_server_terminals==0.5.3
|
| 45 |
+
click==8.1.7
|
| 46 |
+
einops==0.8.0
|
| 47 |
+
pyzmq==26.2.0
|
| 48 |
+
jupyter_client==8.6.3
|
| 49 |
+
nbconvert==7.16.4
|
| 50 |
+
scikit-learn==1.5.2
|
| 51 |
+
executing==2.1.0
|
| 52 |
+
asttokens==2.4.1
|
| 53 |
+
docker-pycreds==0.4.0
|
| 54 |
+
matplotlib-inline==0.1.7
|
| 55 |
+
overrides==7.7.0
|
| 56 |
+
websocket-client==1.8.0
|
| 57 |
+
nbformat==5.10.4
|
| 58 |
+
elbow==0.1.1
|
| 59 |
+
contourpy==1.3.0
|
| 60 |
+
nvidia-cudnn-cu12==9.1.0.70
|
| 61 |
+
transformers==4.44.2
|
| 62 |
+
gitdb==4.0.11
|
| 63 |
+
jupyterlab_nvdashboard==0.11.0
|
| 64 |
+
lazy_loader==0.4
|
| 65 |
+
jsonpointer==3.0.0
|
| 66 |
+
notebook_shim==0.2.4
|
| 67 |
+
nvidia-nccl-cu12==2.20.5
|
| 68 |
+
ffmpeg-python==0.2.0
|
| 69 |
+
triton==3.0.0
|
| 70 |
+
mistune==3.0.2
|
| 71 |
+
python-dateutil==2.9.0.post0
|
| 72 |
+
beautifulsoup4==4.12.3
|
| 73 |
+
nbclient==0.10.0
|
| 74 |
+
h5py==3.12.1
|
| 75 |
+
ftfy==6.2.3
|
| 76 |
+
zipp==3.20.2
|
| 77 |
+
ptyprocess==0.7.0
|
| 78 |
+
huggingface-hub==0.25.1
|
| 79 |
+
pytz==2024.2
|
| 80 |
+
jupyterlab_pygments==0.3.0
|
| 81 |
+
nvidia-cublas-cu12==12.1.3.1
|
| 82 |
+
pandocfilters==1.5.1
|
| 83 |
+
Jinja2==3.1.4
|
| 84 |
+
arrow==1.3.0
|
| 85 |
+
rpds-py==0.20.0
|
| 86 |
+
jupyter_server==2.14.2
|
| 87 |
+
simplejson==3.19.3
|
| 88 |
+
networkx==3.3
|
| 89 |
+
packaging==24.1
|
| 90 |
+
traitlets==5.14.3
|
| 91 |
+
pandas==2.2.3
|
| 92 |
+
xformers==0.0.22.post7
|
| 93 |
+
lightning-utilities==0.11.7
|
| 94 |
+
tifffile==2024.9.20
|
| 95 |
+
nvidia-cuda-cupti-cu12==12.1.105
|
| 96 |
+
mpmath==1.3.0
|
| 97 |
+
GitPython==3.1.43
|
| 98 |
+
scipy==1.14.1
|
| 99 |
+
jsonschema==4.23.0
|
| 100 |
+
prompt_toolkit==3.0.48
|
| 101 |
+
s3transfer==0.10.2
|
| 102 |
+
multidict==6.1.0
|
| 103 |
+
bleach==6.1.0
|
| 104 |
+
sentry-sdk==2.15.0
|
| 105 |
+
nibabel==5.2.1
|
| 106 |
+
accelerate==1.0.0
|
| 107 |
+
pyarrow==17.0.0
|
| 108 |
+
threadpoolctl==3.5.0
|
| 109 |
+
attrs==24.2.0
|
| 110 |
+
rfc3986-validator==0.1.1
|
| 111 |
+
nvidia-cuda-runtime-cu12==12.1.105
|
| 112 |
+
ipywidgets==8.1.5
|
| 113 |
+
frozenlist==1.4.1
|
| 114 |
+
pycparser==2.22
|
| 115 |
+
jupyterlab_server==2.27.3
|
| 116 |
+
nvidia-cuda-nvrtc-cu12==12.1.105
|
| 117 |
+
yarl==1.13.1
|
| 118 |
+
setproctitle==1.3.3
|
| 119 |
+
isoduration==20.11.0
|
| 120 |
+
Pygments==2.18.0
|
| 121 |
+
jedi==0.19.1
|
| 122 |
+
boto3==1.34.57
|
| 123 |
+
tokenizers==0.19.1
|
| 124 |
+
referencing==0.35.1
|
| 125 |
+
rfc3339-validator==0.1.4
|
| 126 |
+
pillow==10.4.0
|
| 127 |
+
jupyterlab==4.2.5
|
| 128 |
+
stack-data==0.6.3
|
| 129 |
+
h11==0.14.0
|
| 130 |
+
anyio==4.6.0
|
| 131 |
+
nilearn==0.10.4
|
| 132 |
+
nvidia-cusolver-cu12==11.4.5.107
|
| 133 |
+
tinycss2==1.3.0
|
| 134 |
+
defusedxml==0.7.1
|
| 135 |
+
argon2-cffi-bindings==21.2.0
|
| 136 |
+
soupsieve==2.6
|
| 137 |
+
nest-asyncio==1.6.0
|
| 138 |
+
torchmetrics==1.3.0.post0
|
| 139 |
+
tqdm==4.66.5
|
| 140 |
+
cffi==1.17.1
|
| 141 |
+
charset-normalizer==3.3.2
|
| 142 |
+
jsonschema-specifications==2023.12.1
|
| 143 |
+
decorator==5.1.1
|
| 144 |
+
open_clip_torch==2.26.1
|
| 145 |
+
jupyter-events==0.10.0
|
| 146 |
+
smart-open==7.0.5
|
| 147 |
+
antlr4-python3-runtime==4.9.3
|
| 148 |
+
prometheus_client==0.21.0
|
| 149 |
+
kornia==0.7.3
|
| 150 |
+
typing_extensions==4.12.2
|
| 151 |
+
sniffio==1.3.1
|
| 152 |
+
joblib==1.4.2
|
| 153 |
+
comm==0.2.2
|
| 154 |
+
aiohappyeyeballs==2.4.3
|
| 155 |
+
numpy==2.1.2
|
| 156 |
+
braceexpand==0.1.7
|
| 157 |
+
certifi==2024.8.30
|
| 158 |
+
psutil==6.0.0
|
| 159 |
+
pyparsing==3.1.4
|
| 160 |
+
pure_eval==0.2.3
|
| 161 |
+
nvidia-cusparse-cu12==12.1.0.106
|
| 162 |
+
wandb==0.18.3
|
| 163 |
+
urllib3==2.2.3
|
| 164 |
+
smmap==5.0.1
|
| 165 |
+
platformdirs==4.3.6
|
| 166 |
+
torch==2.4.1+cu121
|
| 167 |
+
requests==2.32.3
|
| 168 |
+
json5==0.9.25
|
| 169 |
+
nvidia-nvjitlink-cu12==12.6.77
|
| 170 |
+
jupyterlab_widgets==3.0.13
|
| 171 |
+
lxml==5.3.0
|
| 172 |
+
httpx==0.27.2
|
| 173 |
+
opencv-python==4.6.0.66
|
| 174 |
+
portalocker==2.10.1
|
| 175 |
+
pytorch-lightning==2.0.1
|
| 176 |
+
sympy==1.13.3
|
| 177 |
+
wcwidth==0.2.13
|
| 178 |
+
jmespath==1.0.1
|
| 179 |
+
fqdn==1.5.1
|
| 180 |
+
pynvml==11.5.3
|
| 181 |
+
pip==24.0
|
| 182 |
+
wrapt==1.16.0
|
| 183 |
+
aiohttp==3.10.9
|
| 184 |
+
filelock==3.16.1
|
| 185 |
+
fonttools==4.54.1
|
| 186 |
+
fastjsonschema==2.20.0
|
| 187 |
+
jupyter-console==6.6.3
|
| 188 |
+
widgetsnbextension==4.0.13
|
| 189 |
+
timm==1.0.9
|
| 190 |
+
nvidia-cufft-cu12==11.0.2.54
|
| 191 |
+
ipython==8.28.0
|
| 192 |
+
nvidia-nvtx-cu12==12.1.105
|
| 193 |
+
jupyter-lsp==2.2.5
|
| 194 |
+
safetensors==0.4.5
|
| 195 |
+
terminado==0.18.1
|
| 196 |
+
argon2-cffi==23.1.0
|
| 197 |
+
Send2Trash==1.8.3
|
| 198 |
+
importlib_metadata==8.5.0
|
fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/files/wandb-metadata.json
ADDED
|
@@ -0,0 +1,144 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"os": "Linux-5.15.0-1058-aws-x86_64-with-glibc2.31",
|
| 3 |
+
"python": "3.11.10",
|
| 4 |
+
"startedAt": "2024-10-23T04:11:25.180317Z",
|
| 5 |
+
"program": "ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.ipynb",
|
| 6 |
+
"git": {
|
| 7 |
+
"remote": "https://github.com/MedARC-AI/fMRI-foundation-model",
|
| 8 |
+
"commit": "b1ba684ae7a5cc4155cc046b0abe613de09bf700"
|
| 9 |
+
},
|
| 10 |
+
"email": "torrico.villanueva.cesar.kadir@gmail.com",
|
| 11 |
+
"root": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src",
|
| 12 |
+
"host": "ip-10-0-160-143",
|
| 13 |
+
"username": "ckadirt",
|
| 14 |
+
"executable": "/admin/home-ckadirt/foundation_env/bin/python",
|
| 15 |
+
"cpu_count": 96,
|
| 16 |
+
"cpu_count_logical": 192,
|
| 17 |
+
"gpu": "[NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3]",
|
| 18 |
+
"gpu_count": 8,
|
| 19 |
+
"disk": {
|
| 20 |
+
"/": {
|
| 21 |
+
"total": "249555763200",
|
| 22 |
+
"used": "185027031040"
|
| 23 |
+
}
|
| 24 |
+
},
|
| 25 |
+
"memory": {
|
| 26 |
+
"total": "2147443429376"
|
| 27 |
+
},
|
| 28 |
+
"cpu": {
|
| 29 |
+
"count": 96,
|
| 30 |
+
"countLogical": 192
|
| 31 |
+
},
|
| 32 |
+
"gpu_nvidia": [
|
| 33 |
+
{
|
| 34 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 35 |
+
"memoryTotal": "85520809984",
|
| 36 |
+
"cudaCores": 16896,
|
| 37 |
+
"architecture": "Hopper"
|
| 38 |
+
},
|
| 39 |
+
{
|
| 40 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 41 |
+
"memoryTotal": "85520809984",
|
| 42 |
+
"cudaCores": 16896,
|
| 43 |
+
"architecture": "Hopper"
|
| 44 |
+
},
|
| 45 |
+
{
|
| 46 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 47 |
+
"memoryTotal": "85520809984",
|
| 48 |
+
"cudaCores": 16896,
|
| 49 |
+
"architecture": "Hopper"
|
| 50 |
+
},
|
| 51 |
+
{
|
| 52 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 53 |
+
"memoryTotal": "85520809984",
|
| 54 |
+
"cudaCores": 16896,
|
| 55 |
+
"architecture": "Hopper"
|
| 56 |
+
},
|
| 57 |
+
{
|
| 58 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 59 |
+
"memoryTotal": "85520809984",
|
| 60 |
+
"cudaCores": 16896,
|
| 61 |
+
"architecture": "Hopper"
|
| 62 |
+
},
|
| 63 |
+
{
|
| 64 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 65 |
+
"memoryTotal": "85520809984",
|
| 66 |
+
"cudaCores": 16896,
|
| 67 |
+
"architecture": "Hopper"
|
| 68 |
+
},
|
| 69 |
+
{
|
| 70 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 71 |
+
"memoryTotal": "85520809984",
|
| 72 |
+
"cudaCores": 16896,
|
| 73 |
+
"architecture": "Hopper"
|
| 74 |
+
},
|
| 75 |
+
{
|
| 76 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 77 |
+
"memoryTotal": "85520809984",
|
| 78 |
+
"cudaCores": 16896,
|
| 79 |
+
"architecture": "Hopper"
|
| 80 |
+
}
|
| 81 |
+
],
|
| 82 |
+
"slurm": {
|
| 83 |
+
"cluster_name": "sagemaker2",
|
| 84 |
+
"conf": "/opt/slurm/etc/slurm.conf",
|
| 85 |
+
"cpu_bind": "quiet,mask_cpu:0x00000000000000FFC000000000000000000000FFC0000000",
|
| 86 |
+
"cpu_bind_list": "0x00000000000000FFC000000000000000000000FFC0000000",
|
| 87 |
+
"cpu_bind_type": "mask_cpu:",
|
| 88 |
+
"cpu_bind_verbose": "quiet",
|
| 89 |
+
"cpus_on_node": "20",
|
| 90 |
+
"gpus": "1",
|
| 91 |
+
"gpus_on_node": "1",
|
| 92 |
+
"gtids": "0",
|
| 93 |
+
"job_account": "fmri",
|
| 94 |
+
"job_cpus_per_node": "20",
|
| 95 |
+
"job_end_time": "1729702756",
|
| 96 |
+
"job_gid": "1879800513",
|
| 97 |
+
"job_group": "Domain Users",
|
| 98 |
+
"job_id": "528040",
|
| 99 |
+
"job_name": "bash",
|
| 100 |
+
"job_nodelist": "ip-10-0-160-143",
|
| 101 |
+
"job_num_nodes": "1",
|
| 102 |
+
"job_partition": "p5",
|
| 103 |
+
"job_qos": "idle",
|
| 104 |
+
"job_start_time": "1729648756",
|
| 105 |
+
"job_uid": "1879804696",
|
| 106 |
+
"job_user": "ckadirt",
|
| 107 |
+
"jobid": "528040",
|
| 108 |
+
"launch_node_ipaddr": "172.17.12.61",
|
| 109 |
+
"localid": "0",
|
| 110 |
+
"mpi_type": "pmix_v3",
|
| 111 |
+
"nnodes": "1",
|
| 112 |
+
"nodeid": "0",
|
| 113 |
+
"nodelist": "ip-10-0-160-143",
|
| 114 |
+
"nprocs": "1",
|
| 115 |
+
"ntasks": "1",
|
| 116 |
+
"pmix_mapping_serv": "(vector,(0,1,1))",
|
| 117 |
+
"pmixp_abort_agent_port": "34923",
|
| 118 |
+
"prio_process": "0",
|
| 119 |
+
"procid": "0",
|
| 120 |
+
"pty_port": "45733",
|
| 121 |
+
"pty_win_col": "199",
|
| 122 |
+
"pty_win_row": "17",
|
| 123 |
+
"script_context": "prolog_task",
|
| 124 |
+
"srun_comm_host": "172.17.12.61",
|
| 125 |
+
"srun_comm_port": "39353",
|
| 126 |
+
"step_gpus": "3",
|
| 127 |
+
"step_id": "0",
|
| 128 |
+
"step_launcher_port": "39353",
|
| 129 |
+
"step_nodelist": "ip-10-0-160-143",
|
| 130 |
+
"step_num_nodes": "1",
|
| 131 |
+
"step_num_tasks": "1",
|
| 132 |
+
"step_tasks_per_node": "1",
|
| 133 |
+
"stepid": "0",
|
| 134 |
+
"submit_dir": "/weka/proj-fmri",
|
| 135 |
+
"submit_host": "ip-172-17-12-61",
|
| 136 |
+
"task_pid": "1032669",
|
| 137 |
+
"tasks_per_node": "1",
|
| 138 |
+
"topology_addr": "ip-10-0-160-143",
|
| 139 |
+
"topology_addr_pattern": "node",
|
| 140 |
+
"umask": "0022",
|
| 141 |
+
"working_cluster": "sagemaker2:ip-172-17-63-161:6817:9984:109"
|
| 142 |
+
},
|
| 143 |
+
"cudaVersion": "12.2"
|
| 144 |
+
}
|
fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug-core.log
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-23T04:00:15.891404575Z","level":"INFO","msg":"started logging, with flags","port-filename":"/tmp/tmp623bhuv_/port-1087257.txt","pid":1087257,"debug":false,"disable-analytics":false}
|
| 2 |
+
{"time":"2024-10-23T04:00:15.891695701Z","level":"INFO","msg":"FeatureState","shutdownOnParentExitEnabled":false}
|
| 3 |
+
{"time":"2024-10-23T04:00:15.894505677Z","level":"INFO","msg":"Will exit if parent process dies.","ppid":1087257}
|
| 4 |
+
{"time":"2024-10-23T04:00:15.894459316Z","level":"INFO","msg":"server is running","addr":{"IP":"127.0.0.1","Port":40341,"Zone":""}}
|
| 5 |
+
{"time":"2024-10-23T04:00:16.078034023Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"127.0.0.1:58780"}
|
| 6 |
+
{"time":"2024-10-23T04:00:16.909578833Z","level":"INFO","msg":"handleInformInit: received","streamId":"HCPflat_large_gsrFalse__HCP_FT_83810","id":"127.0.0.1:58780"}
|
| 7 |
+
{"time":"2024-10-23T04:00:17.021711364Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"HCPflat_large_gsrFalse__HCP_FT_83810","id":"127.0.0.1:58780"}
|
| 8 |
+
{"time":"2024-10-23T04:11:25.031512552Z","level":"INFO","msg":"handleInformFinish: finish message received","streamId":"HCPflat_large_gsrFalse__HCP_FT_83810","id":"127.0.0.1:58780"}
|
| 9 |
+
{"time":"2024-10-23T04:11:25.032442271Z","level":"INFO","msg":"handleInformFinish: stream closed","streamId":"HCPflat_large_gsrFalse__HCP_FT_83810","id":"127.0.0.1:58780"}
|
| 10 |
+
{"time":"2024-10-23T04:11:25.179399522Z","level":"INFO","msg":"handleInformInit: received","streamId":"HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f","id":"127.0.0.1:58780"}
|
| 11 |
+
{"time":"2024-10-23T04:11:25.24570696Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f","id":"127.0.0.1:58780"}
|
| 12 |
+
{"time":"2024-10-23T04:19:30.473127709Z","level":"INFO","msg":"Parent process exited, terminating service process."}
|
fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug-internal.log
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-23T04:11:25.205972524Z","level":"INFO","msg":"using version","core version":"0.18.3"}
|
| 2 |
+
{"time":"2024-10-23T04:11:25.205987165Z","level":"INFO","msg":"created symlink","path":"/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug-core.log"}
|
| 3 |
+
{"time":"2024-10-23T04:11:25.206996565Z","level":"ERROR","msg":"dialing: google: could not find default credentials. See https://cloud.google.com/docs/authentication/external/set-up-adc for more information"}
|
| 4 |
+
{"time":"2024-10-23T04:11:25.245647669Z","level":"INFO","msg":"created new stream","id":"HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f"}
|
| 5 |
+
{"time":"2024-10-23T04:11:25.24569935Z","level":"INFO","msg":"stream: started","id":"HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f"}
|
| 6 |
+
{"time":"2024-10-23T04:11:25.24572727Z","level":"INFO","msg":"handler: started","stream_id":{"value":"HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f"}}
|
| 7 |
+
{"time":"2024-10-23T04:11:25.24571807Z","level":"INFO","msg":"sender: started","stream_id":{"value":"HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f"}}
|
| 8 |
+
{"time":"2024-10-23T04:11:25.24571519Z","level":"INFO","msg":"writer: Do: started","stream_id":{"value":"HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f"}}
|
| 9 |
+
{"time":"2024-10-23T04:11:25.87412901Z","level":"INFO","msg":"wandb-core","!BADKEY":null}
|
| 10 |
+
{"time":"2024-10-23T04:11:25.884775353Z","level":"INFO","msg":"Starting system monitor"}
|
| 11 |
+
{"time":"2024-10-23T04:11:25.884793743Z","level":"WARN","msg":"handleCodeSave: program relative path is empty"}
|
| 12 |
+
{"time":"2024-10-23T04:11:25.886547389Z","level":"ERROR","msg":"git repo not found","error":"repository does not exist"}
|
fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug.log
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
2024-10-23 04:11:22,311 INFO MainThread:1087257 [wandb_setup.py:_flush():79] Current SDK version is 0.18.3
|
| 2 |
+
2024-10-23 04:11:22,311 INFO MainThread:1087257 [wandb_setup.py:_flush():79] Configure stats pid to 1087257
|
| 3 |
+
2024-10-23 04:11:22,311 INFO MainThread:1087257 [wandb_setup.py:_flush():79] Loading settings from /admin/home-ckadirt/.config/wandb/settings
|
| 4 |
+
2024-10-23 04:11:22,311 INFO MainThread:1087257 [wandb_setup.py:_flush():79] Loading settings from /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/settings
|
| 5 |
+
2024-10-23 04:11:22,311 INFO MainThread:1087257 [wandb_setup.py:_flush():79] Loading settings from environment variables: {}
|
| 6 |
+
2024-10-23 04:11:22,312 INFO MainThread:1087257 [wandb_setup.py:_flush():79] Applying setup settings: {'mode': None, '_disable_service': None}
|
| 7 |
+
2024-10-23 04:11:22,312 INFO MainThread:1087257 [wandb_setup.py:_flush():79] Inferring run settings from compute environment: {'program': '<python with no main file>'}
|
| 8 |
+
2024-10-23 04:11:22,312 INFO MainThread:1087257 [wandb_setup.py:_flush():79] Applying login settings: {}
|
| 9 |
+
2024-10-23 04:11:22,312 INFO MainThread:1087257 [wandb_init.py:_log_setup():532] Logging user logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug.log
|
| 10 |
+
2024-10-23 04:11:22,313 INFO MainThread:1087257 [wandb_init.py:_log_setup():533] Logging internal logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug-internal.log
|
| 11 |
+
2024-10-23 04:11:22,313 INFO MainThread:1087257 [wandb_init.py:init():617] calling init triggers
|
| 12 |
+
2024-10-23 04:11:22,313 INFO MainThread:1087257 [wandb_init.py:init():624] wandb.init called with sweep_config: {}
|
| 13 |
+
config: {'model_name': 'HCPflat_large_gsrFalse__HCP_FT', 'batch_size': 8, 'learning_rate': 0.0001, 'weight_decay': 1e-05, 'num_epochs': 20, 'seed': 42}
|
| 14 |
+
2024-10-23 04:11:22,313 INFO MainThread:1087257 [wandb_init.py:init():642] re-initializing run, found existing run on stack: HCPflat_large_gsrFalse__HCP_FT_83810
|
| 15 |
+
2024-10-23 04:11:22,315 INFO MainThread:1087257 [wandb_run.py:_finish():2164] finishing run ckadirt/fMRI-foundation-model/HCPflat_large_gsrFalse__HCP_FT_83810
|
| 16 |
+
2024-10-23 04:11:22,357 INFO MainThread:1087257 [jupyter.py:save_history():488] saving 17 cells to _session_history.ipynb
|
| 17 |
+
2024-10-23 04:11:22,358 INFO MainThread:1087257 [wandb_run.py:_config_callback():1394] config_cb ('_wandb', 'session_history') code/_session_history.ipynb None
|
| 18 |
+
2024-10-23 04:11:22,410 INFO MainThread:1087257 [jupyter.py:_save_ipynb():398] looking for notebook: ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.ipynb
|
| 19 |
+
2024-10-23 04:11:22,410 INFO MainThread:1087257 [wandb_init.py:_jupyter_teardown():460] cleaning up jupyter logic
|
| 20 |
+
2024-10-23 04:11:22,410 INFO MainThread:1087257 [wandb_run.py:_atexit_cleanup():2428] got exitcode: 0
|
| 21 |
+
2024-10-23 04:11:22,412 INFO MainThread:1087257 [wandb_run.py:_restore():2410] restore
|
| 22 |
+
2024-10-23 04:11:22,413 INFO MainThread:1087257 [wandb_run.py:_restore():2416] restore done
|
| 23 |
+
2024-10-23 04:11:25,001 INFO MainThread:1087257 [wandb_run.py:_footer_history_summary_info():4049] rendering history
|
| 24 |
+
2024-10-23 04:11:25,001 INFO MainThread:1087257 [wandb_run.py:_footer_history_summary_info():4081] rendering summary
|
| 25 |
+
2024-10-23 04:11:25,028 INFO MainThread:1087257 [wandb_run.py:_footer_sync_info():4008] logging synced files
|
| 26 |
+
2024-10-23 04:11:25,164 INFO MainThread:1087257 [wandb_init.py:init():667] starting backend
|
| 27 |
+
2024-10-23 04:11:25,164 INFO MainThread:1087257 [wandb_init.py:init():671] sending inform_init request
|
| 28 |
+
2024-10-23 04:11:25,179 INFO MainThread:1087257 [backend.py:_multiprocessing_setup():104] multiprocessing start_methods=fork,spawn,forkserver, using: spawn
|
| 29 |
+
2024-10-23 04:11:25,179 INFO MainThread:1087257 [wandb_init.py:init():684] backend started and connected
|
| 30 |
+
2024-10-23 04:11:25,203 INFO MainThread:1087257 [wandb_run.py:_label_probe_notebook():1346] probe notebook
|
| 31 |
+
2024-10-23 04:11:25,204 INFO MainThread:1087257 [wandb_run.py:_label_probe_notebook():1356] Unable to probe notebook: 'NoneType' object has no attribute 'get'
|
| 32 |
+
2024-10-23 04:11:25,204 INFO MainThread:1087257 [wandb_init.py:init():779] updated telemetry
|
| 33 |
+
2024-10-23 04:11:25,319 INFO MainThread:1087257 [wandb_init.py:init():812] communicating run to backend with 90.0 second timeout
|
| 34 |
+
2024-10-23 04:11:25,855 INFO MainThread:1087257 [wandb_init.py:init():863] starting run threads in backend
|
| 35 |
+
2024-10-23 04:11:26,356 INFO MainThread:1087257 [wandb_run.py:_console_start():2465] atexit reg
|
| 36 |
+
2024-10-23 04:11:26,356 INFO MainThread:1087257 [wandb_run.py:_redirect():2313] redirect: wrap_raw
|
| 37 |
+
2024-10-23 04:11:26,357 INFO MainThread:1087257 [wandb_run.py:_redirect():2378] Wrapping output streams.
|
| 38 |
+
2024-10-23 04:11:26,357 INFO MainThread:1087257 [wandb_run.py:_redirect():2403] Redirects installed.
|
| 39 |
+
2024-10-23 04:11:26,357 INFO MainThread:1087257 [wandb_init.py:init():907] run started, returning control to user process
|
fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/run-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f.wandb
ADDED
|
Binary file (557 kB). View file
|
|
|
fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/logs/debug-internal.log
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-23T04:12:23.854291574Z","level":"INFO","msg":"using version","core version":"0.18.3"}
|
| 2 |
+
{"time":"2024-10-23T04:12:23.854310795Z","level":"INFO","msg":"created symlink","path":"/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/logs/debug-core.log"}
|
| 3 |
+
{"time":"2024-10-23T04:12:23.856970424Z","level":"ERROR","msg":"dialing: google: could not find default credentials. See https://cloud.google.com/docs/authentication/external/set-up-adc for more information"}
|
| 4 |
+
{"time":"2024-10-23T04:12:23.904377013Z","level":"INFO","msg":"created new stream","id":"NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3"}
|
| 5 |
+
{"time":"2024-10-23T04:12:23.904404853Z","level":"INFO","msg":"stream: started","id":"NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3"}
|
| 6 |
+
{"time":"2024-10-23T04:12:23.904415433Z","level":"INFO","msg":"sender: started","stream_id":{"value":"NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3"}}
|
| 7 |
+
{"time":"2024-10-23T04:12:23.904437214Z","level":"INFO","msg":"writer: Do: started","stream_id":{"value":"NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3"}}
|
| 8 |
+
{"time":"2024-10-23T04:12:23.904419533Z","level":"INFO","msg":"handler: started","stream_id":{"value":"NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3"}}
|
| 9 |
+
{"time":"2024-10-23T04:12:24.41246663Z","level":"INFO","msg":"wandb-core","!BADKEY":null}
|
| 10 |
+
{"time":"2024-10-23T04:12:24.425346089Z","level":"INFO","msg":"Starting system monitor"}
|
| 11 |
+
{"time":"2024-10-23T04:12:24.491097327Z","level":"ERROR","msg":"git repo not found","error":"repository does not exist"}
|
fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/logs/debug.log
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
2024-10-23 04:12:23,820 INFO MainThread:3333420 [wandb_setup.py:_flush():79] Current SDK version is 0.18.3
|
| 2 |
+
2024-10-23 04:12:23,820 INFO MainThread:3333420 [wandb_setup.py:_flush():79] Configure stats pid to 3333420
|
| 3 |
+
2024-10-23 04:12:23,820 INFO MainThread:3333420 [wandb_setup.py:_flush():79] Loading settings from /admin/home-ckadirt/.config/wandb/settings
|
| 4 |
+
2024-10-23 04:12:23,820 INFO MainThread:3333420 [wandb_setup.py:_flush():79] Loading settings from /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/settings
|
| 5 |
+
2024-10-23 04:12:23,820 INFO MainThread:3333420 [wandb_setup.py:_flush():79] Loading settings from environment variables: {}
|
| 6 |
+
2024-10-23 04:12:23,820 INFO MainThread:3333420 [wandb_setup.py:_flush():79] Applying setup settings: {'mode': None, '_disable_service': None}
|
| 7 |
+
2024-10-23 04:12:23,820 INFO MainThread:3333420 [wandb_setup.py:_flush():79] Inferring run settings from compute environment: {'program_relpath': 'src/HCP_downstream_finetune.py', 'program_abspath': '/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py', 'program': '/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py'}
|
| 8 |
+
2024-10-23 04:12:23,820 INFO MainThread:3333420 [wandb_setup.py:_flush():79] Applying login settings: {}
|
| 9 |
+
2024-10-23 04:12:23,821 INFO MainThread:3333420 [wandb_init.py:_log_setup():532] Logging user logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/logs/debug.log
|
| 10 |
+
2024-10-23 04:12:23,822 INFO MainThread:3333420 [wandb_init.py:_log_setup():533] Logging internal logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/logs/debug-internal.log
|
| 11 |
+
2024-10-23 04:12:23,822 INFO MainThread:3333420 [wandb_init.py:init():617] calling init triggers
|
| 12 |
+
2024-10-23 04:12:23,822 INFO MainThread:3333420 [wandb_init.py:init():624] wandb.init called with sweep_config: {}
|
| 13 |
+
config: {'model_name': 'NSDflat_large_gsrFalse__HCP_FT', 'batch_size': 8, 'learning_rate': 0.0001, 'weight_decay': 1e-05, 'num_epochs': 20, 'seed': 42}
|
| 14 |
+
2024-10-23 04:12:23,822 INFO MainThread:3333420 [wandb_init.py:init():667] starting backend
|
| 15 |
+
2024-10-23 04:12:23,822 INFO MainThread:3333420 [wandb_init.py:init():671] sending inform_init request
|
| 16 |
+
2024-10-23 04:12:23,831 INFO MainThread:3333420 [backend.py:_multiprocessing_setup():104] multiprocessing start_methods=fork,spawn,forkserver, using: spawn
|
| 17 |
+
2024-10-23 04:12:23,831 INFO MainThread:3333420 [wandb_init.py:init():684] backend started and connected
|
| 18 |
+
2024-10-23 04:12:23,852 INFO MainThread:3333420 [wandb_init.py:init():779] updated telemetry
|
| 19 |
+
2024-10-23 04:12:23,895 INFO MainThread:3333420 [wandb_init.py:init():812] communicating run to backend with 90.0 second timeout
|
| 20 |
+
2024-10-23 04:12:24,396 INFO MainThread:3333420 [wandb_init.py:init():863] starting run threads in backend
|
| 21 |
+
2024-10-23 04:12:24,755 INFO MainThread:3333420 [wandb_run.py:_console_start():2465] atexit reg
|
| 22 |
+
2024-10-23 04:12:24,755 INFO MainThread:3333420 [wandb_run.py:_redirect():2313] redirect: wrap_raw
|
| 23 |
+
2024-10-23 04:12:24,755 INFO MainThread:3333420 [wandb_run.py:_redirect():2378] Wrapping output streams.
|
| 24 |
+
2024-10-23 04:12:24,756 INFO MainThread:3333420 [wandb_run.py:_redirect():2403] Redirects installed.
|
| 25 |
+
2024-10-23 04:12:24,758 INFO MainThread:3333420 [wandb_init.py:init():907] run started, returning control to user process
|
fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/run-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3.wandb
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:4b94ade7deadb51db63431f075745b200d705982cf2fd53012bc0d3a272293b0
|
| 3 |
+
size 3604480
|
fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/code/src/HCP_downstream_finetune.py
ADDED
|
@@ -0,0 +1,596 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
# coding: utf-8
|
| 3 |
+
|
| 4 |
+
# In[1]:
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
# Import packages and setup gpu configuration.
|
| 8 |
+
# This code block shouldnt need to be adjusted!
|
| 9 |
+
import os
|
| 10 |
+
import sys
|
| 11 |
+
import json
|
| 12 |
+
import yaml
|
| 13 |
+
import numpy as np
|
| 14 |
+
import copy
|
| 15 |
+
import math
|
| 16 |
+
import time
|
| 17 |
+
import random
|
| 18 |
+
from tqdm.auto import tqdm
|
| 19 |
+
import webdataset as wds
|
| 20 |
+
import matplotlib.pyplot as plt
|
| 21 |
+
|
| 22 |
+
import torch
|
| 23 |
+
import torch.nn as nn
|
| 24 |
+
from torchvision import transforms
|
| 25 |
+
import utils
|
| 26 |
+
from mae_utils.flat_models import *
|
| 27 |
+
import h5py
|
| 28 |
+
from mae_utils import flat_models
|
| 29 |
+
|
| 30 |
+
# tf32 data type is faster than standard float32
|
| 31 |
+
torch.backends.cuda.matmul.allow_tf32 = True
|
| 32 |
+
# following fixes a Conv3D CUDNN_NOT_SUPPORTED error
|
| 33 |
+
torch.backends.cudnn.benchmark = True
|
| 34 |
+
|
| 35 |
+
# ## MODEL TO LOAD ##
|
| 36 |
+
if utils.is_interactive():
|
| 37 |
+
model_name = "HCPflat_large_gsrFalse_"
|
| 38 |
+
else:
|
| 39 |
+
model_name = sys.argv[1]
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
# outdir = os.path.abspath(f'checkpoints/{model_name}')
|
| 43 |
+
outdir = os.path.abspath(f'checkpoints/{model_name}')
|
| 44 |
+
|
| 45 |
+
print("outdir", outdir)
|
| 46 |
+
# Load previous config.yaml if available
|
| 47 |
+
if os.path.exists(f"{outdir}/config.yaml"):
|
| 48 |
+
config = yaml.load(open(f"{outdir}/config.yaml", 'r'), Loader=yaml.FullLoader)
|
| 49 |
+
print(f"Loaded config.yaml from ckpt folder {outdir}")
|
| 50 |
+
# create global variables from the config
|
| 51 |
+
print("\n__CONFIG__")
|
| 52 |
+
for attribute_name in config.keys():
|
| 53 |
+
print(f"{attribute_name} = {config[attribute_name]}")
|
| 54 |
+
globals()[attribute_name] = config[f'{attribute_name}']
|
| 55 |
+
print("\n")
|
| 56 |
+
|
| 57 |
+
world_size = os.getenv('WORLD_SIZE')
|
| 58 |
+
if world_size is None:
|
| 59 |
+
world_size = 1
|
| 60 |
+
else:
|
| 61 |
+
world_size = int(world_size)
|
| 62 |
+
print(f"WORLD_SIZE={world_size}")
|
| 63 |
+
|
| 64 |
+
if utils.is_interactive():
|
| 65 |
+
# Following allows you to change functions in models.py or utils.py and
|
| 66 |
+
# have this notebook automatically update with your revisions
|
| 67 |
+
get_ipython().run_line_magic('load_ext', 'autoreload')
|
| 68 |
+
get_ipython().run_line_magic('autoreload', '2')
|
| 69 |
+
|
| 70 |
+
batch_size = probe_batch_size
|
| 71 |
+
num_epochs = probe_num_epochs
|
| 72 |
+
|
| 73 |
+
data_type = torch.float32 # change depending on your mixed_precision
|
| 74 |
+
global_batch_size = batch_size * world_size
|
| 75 |
+
|
| 76 |
+
device = torch.device('cuda')
|
| 77 |
+
|
| 78 |
+
hcp_flat_path = "/weka/proj-medarc/shared/HCP-Flat"
|
| 79 |
+
# seed = 42
|
| 80 |
+
# num_frames = 16
|
| 81 |
+
# gsr = False
|
| 82 |
+
# num_workers = 10
|
| 83 |
+
# batch_size = 128
|
| 84 |
+
|
| 85 |
+
print("PID of this process =",os.getpid())
|
| 86 |
+
utils.seed_everything(seed)
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
# In[2]:
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
if os.getenv('global_pool') == "False":
|
| 93 |
+
global_pool = False
|
| 94 |
+
else:
|
| 95 |
+
global_pool = True
|
| 96 |
+
print(f"global_pool = {global_pool}")
|
| 97 |
+
|
| 98 |
+
try:
|
| 99 |
+
gsr
|
| 100 |
+
except:
|
| 101 |
+
gsr = True
|
| 102 |
+
print("set gsr to True")
|
| 103 |
+
print(f"gsr = {gsr}")
|
| 104 |
+
|
| 105 |
+
|
| 106 |
+
# In[3]:
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
#### UNCOMMENT THIS TO SAVE THE HCP-FLAT IN HDF5 FORMAT
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
# from torch.utils.data import default_collate
|
| 113 |
+
# from mae_utils.flat import load_hcp_flat_mask
|
| 114 |
+
# from mae_utils.flat import create_hcp_flat
|
| 115 |
+
# from mae_utils.flat import batch_unmask
|
| 116 |
+
# import mae_utils.visualize as vis
|
| 117 |
+
|
| 118 |
+
|
| 119 |
+
# batch_size = 26
|
| 120 |
+
# print(f"changed batch_size to {batch_size}")
|
| 121 |
+
|
| 122 |
+
# ## Test ##
|
| 123 |
+
# datasets_to_include = "HCP"
|
| 124 |
+
# assert "HCP" in datasets_to_include
|
| 125 |
+
# test_dataset = create_hcp_flat(root=hcp_flat_path,
|
| 126 |
+
# clip_mode="event", frames=num_frames, shuffle=False, gsr=gsr, sub_list = 'test')
|
| 127 |
+
# test_dl = wds.WebLoader(
|
| 128 |
+
# test_dataset.batched(batch_size, partial=False, collation_fn=default_collate),
|
| 129 |
+
# batch_size=None,
|
| 130 |
+
# shuffle=False,
|
| 131 |
+
# num_workers=num_workers,
|
| 132 |
+
# pin_memory=True,
|
| 133 |
+
# )
|
| 134 |
+
|
| 135 |
+
# ## Train ##
|
| 136 |
+
# assert "HCP" in datasets_to_include
|
| 137 |
+
# train_dataset = create_hcp_flat(root=hcp_flat_path,
|
| 138 |
+
# clip_mode="event", frames=num_frames, shuffle=False, gsr=gsr, sub_list = 'train')
|
| 139 |
+
# train_dl = wds.WebLoader(
|
| 140 |
+
# train_dataset.batched(batch_size, partial=False, collation_fn=default_collate),
|
| 141 |
+
# batch_size=None,
|
| 142 |
+
# shuffle=False,
|
| 143 |
+
# num_workers=num_workers,
|
| 144 |
+
# pin_memory=True,
|
| 145 |
+
# )
|
| 146 |
+
|
| 147 |
+
# def flatten_meta(meta_dict):
|
| 148 |
+
# """
|
| 149 |
+
# Flatten the meta dictionary by:
|
| 150 |
+
# - Replacing single-item lists with the item itself.
|
| 151 |
+
# - Converting tensors to scalar numbers.
|
| 152 |
+
# """
|
| 153 |
+
# flattened = {}
|
| 154 |
+
# for key, value in meta_dict.items():
|
| 155 |
+
# if isinstance(value, list):
|
| 156 |
+
# if len(value) == 1:
|
| 157 |
+
# flattened[key] = value[0] # Replace list with its single item
|
| 158 |
+
# else:
|
| 159 |
+
# flattened[key] = value # Keep as is if multiple items
|
| 160 |
+
# elif isinstance(value, torch.Tensor):
|
| 161 |
+
# # Convert tensor to scalar
|
| 162 |
+
# if value.numel() == 1:
|
| 163 |
+
# flattened[key] = value.item()
|
| 164 |
+
# else:
|
| 165 |
+
# flattened[key] = value.tolist() # Convert multi-element tensor to list
|
| 166 |
+
# else:
|
| 167 |
+
# flattened[key] = value # Keep the value as is
|
| 168 |
+
# return flattened
|
| 169 |
+
|
| 170 |
+
# import h5py
|
| 171 |
+
# meta_array = np.array([], dtype=object)
|
| 172 |
+
# # Open an HDF5 file in write mode
|
| 173 |
+
# with h5py.File('train_hcp.hdf5', 'w') as h5f:
|
| 174 |
+
# flatmaps_dset = None
|
| 175 |
+
|
| 176 |
+
# total_samples = 0
|
| 177 |
+
|
| 178 |
+
# for i, batch in tqdm(enumerate(train_dl), total = 120000):
|
| 179 |
+
# images = batch['image'][0]
|
| 180 |
+
# meta = batch['meta']
|
| 181 |
+
# batch_size = images.shape[0]
|
| 182 |
+
# meta_serializable = meta.copy()
|
| 183 |
+
|
| 184 |
+
|
| 185 |
+
# # Step 2: Serialize the dictionary to a JSON string
|
| 186 |
+
# meta_str = json.dumps(flatten_meta(meta_serializable), indent=4)
|
| 187 |
+
# meta_array = np.append(meta_array, meta_str)
|
| 188 |
+
# if flatmaps_dset is None:
|
| 189 |
+
# # Initialize datasets with unlimited (None) maxshape along the first axis
|
| 190 |
+
# flatmaps_shape = (0,) + images.shape[1:]
|
| 191 |
+
# flatmaps_maxshape = (None,) + images.shape[1:]
|
| 192 |
+
|
| 193 |
+
# flatmaps_dset = h5f.create_dataset(
|
| 194 |
+
# 'flatmaps',
|
| 195 |
+
# shape=flatmaps_shape,
|
| 196 |
+
# maxshape=flatmaps_maxshape,
|
| 197 |
+
# dtype=np.float16,
|
| 198 |
+
# chunks=True # Enable chunking for efficient resizing
|
| 199 |
+
# )
|
| 200 |
+
|
| 201 |
+
# # Resize datasets to accommodate new data
|
| 202 |
+
# flatmaps_dset.resize(total_samples + batch_size, axis=0)
|
| 203 |
+
|
| 204 |
+
# # Write data to the datasets
|
| 205 |
+
# flatmaps_dset[total_samples:total_samples + batch_size] = images.numpy().astype(np.float16)
|
| 206 |
+
|
| 207 |
+
# total_samples += batch_size
|
| 208 |
+
|
| 209 |
+
# print(f"Processed {total_samples} samples")
|
| 210 |
+
# np.save('metadata_test_HCP.npy', meta_array)
|
| 211 |
+
|
| 212 |
+
|
| 213 |
+
# import h5py
|
| 214 |
+
# meta_array = np.array([], dtype=object)
|
| 215 |
+
# # Open an HDF5 file in write mode
|
| 216 |
+
# with h5py.File('test_hcp.hdf5', 'w') as h5f:
|
| 217 |
+
# flatmaps_dset = None
|
| 218 |
+
|
| 219 |
+
# total_samples = 0
|
| 220 |
+
|
| 221 |
+
# for i, batch in tqdm(enumerate(test_dl), total = 12000):
|
| 222 |
+
# images = batch['image'][0]
|
| 223 |
+
# meta = batch['meta']
|
| 224 |
+
# batch_size = images.shape[0]
|
| 225 |
+
# meta_serializable = meta.copy()
|
| 226 |
+
|
| 227 |
+
|
| 228 |
+
# # Step 2: Serialize the dictionary to a JSON string
|
| 229 |
+
# meta_str = json.dumps(flatten_meta(meta_serializable), indent=4)
|
| 230 |
+
# meta_array = np.append(meta_array, meta_str)
|
| 231 |
+
# if flatmaps_dset is None:
|
| 232 |
+
# # Initialize datasets with unlimited (None) maxshape along the first axis
|
| 233 |
+
# flatmaps_shape = (0,) + images.shape[1:]
|
| 234 |
+
# flatmaps_maxshape = (None,) + images.shape[1:]
|
| 235 |
+
|
| 236 |
+
# flatmaps_dset = h5f.create_dataset(
|
| 237 |
+
# 'flatmaps',
|
| 238 |
+
# shape=flatmaps_shape,
|
| 239 |
+
# maxshape=flatmaps_maxshape,
|
| 240 |
+
# dtype=np.float16,
|
| 241 |
+
# chunks=True # Enable chunking for efficient resizing
|
| 242 |
+
# )
|
| 243 |
+
|
| 244 |
+
# # Resize datasets to accommodate new data
|
| 245 |
+
# flatmaps_dset.resize(total_samples + batch_size, axis=0)
|
| 246 |
+
|
| 247 |
+
# # Write data to the datasets
|
| 248 |
+
# flatmaps_dset[total_samples:total_samples + batch_size] = images.numpy().astype(np.float16)
|
| 249 |
+
|
| 250 |
+
# total_samples += batch_size
|
| 251 |
+
|
| 252 |
+
# print(f"Processed {total_samples} samples")
|
| 253 |
+
# np.save('metadata_train_HCP.npy', meta_array)
|
| 254 |
+
|
| 255 |
+
|
| 256 |
+
# ### Preparing data
|
| 257 |
+
|
| 258 |
+
# In[4]:
|
| 259 |
+
|
| 260 |
+
|
| 261 |
+
from sklearn.preprocessing import LabelEncoder
|
| 262 |
+
|
| 263 |
+
INCLUDE_CONDS = {
|
| 264 |
+
"fear",
|
| 265 |
+
"neut",
|
| 266 |
+
"math",
|
| 267 |
+
"story",
|
| 268 |
+
"lf",
|
| 269 |
+
"lh",
|
| 270 |
+
"rf",
|
| 271 |
+
"rh",
|
| 272 |
+
"t",
|
| 273 |
+
"match",
|
| 274 |
+
"relation",
|
| 275 |
+
"mental",
|
| 276 |
+
"rnd",
|
| 277 |
+
"0bk_body",
|
| 278 |
+
"2bk_body",
|
| 279 |
+
"0bk_faces",
|
| 280 |
+
"2bk_faces",
|
| 281 |
+
"0bk_places",
|
| 282 |
+
"2bk_places",
|
| 283 |
+
"0bk_tools",
|
| 284 |
+
"2bk_tools",
|
| 285 |
+
}
|
| 286 |
+
|
| 287 |
+
# test_data = []
|
| 288 |
+
|
| 289 |
+
# # Iterate over the DataLoader with a progress bar
|
| 290 |
+
# for sample in tqdm(train_dl, desc="Processing samples"):
|
| 291 |
+
# x = sample['image']
|
| 292 |
+
# y = sample['meta']['trial_type']
|
| 293 |
+
# key = sample['meta']['key']
|
| 294 |
+
# print(x.shape, y, key)
|
| 295 |
+
# break
|
| 296 |
+
# Initialize the label encoder
|
| 297 |
+
label_encoder = LabelEncoder()
|
| 298 |
+
label_encoder.fit(sorted(INCLUDE_CONDS)) # Ensure consistent ordering
|
| 299 |
+
|
| 300 |
+
num_classes = len(label_encoder.classes_)
|
| 301 |
+
print(f"Number of classes: {num_classes}")
|
| 302 |
+
|
| 303 |
+
|
| 304 |
+
# In[5]:
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
f_train = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/train_hcp.hdf5', 'r')
|
| 308 |
+
flatmaps_train = f_train['flatmaps']
|
| 309 |
+
|
| 310 |
+
f_test = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/test_hcp.hdf5', 'r')
|
| 311 |
+
flatmaps_test = f_test['flatmaps']
|
| 312 |
+
|
| 313 |
+
metadata_train = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_train_HCP.npy', allow_pickle=True)
|
| 314 |
+
metadata_test = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_test_HCP.npy', allow_pickle=True)
|
| 315 |
+
|
| 316 |
+
|
| 317 |
+
# In[6]:
|
| 318 |
+
|
| 319 |
+
|
| 320 |
+
from torch.utils.data import Dataset, DataLoader
|
| 321 |
+
|
| 322 |
+
class HCPFlatDataset(Dataset):
|
| 323 |
+
def __init__(self, flatmaps, metadata):
|
| 324 |
+
self.flatmaps = flatmaps
|
| 325 |
+
self.metadata = metadata
|
| 326 |
+
|
| 327 |
+
def __len__(self):
|
| 328 |
+
return len(self.metadata)
|
| 329 |
+
|
| 330 |
+
def __getitem__(self, idx):
|
| 331 |
+
return self.flatmaps[idx], json.loads(self.metadata[idx])
|
| 332 |
+
print("Moving datasets to ram")
|
| 333 |
+
# Loading to cpu for faster training, this can take several minutes. Remove this [:] if you want to move one at the time.
|
| 334 |
+
train_dataset = HCPFlatDataset(flatmaps_train, metadata_train)
|
| 335 |
+
train_dl = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)
|
| 336 |
+
|
| 337 |
+
test_dataset = HCPFlatDataset(flatmaps_test, metadata_test)
|
| 338 |
+
test_dl = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
|
| 339 |
+
print("Datasets ready")
|
| 340 |
+
|
| 341 |
+
|
| 342 |
+
# ### Creating and loading Model
|
| 343 |
+
|
| 344 |
+
# In[7]:
|
| 345 |
+
|
| 346 |
+
|
| 347 |
+
from mae_utils.flat import load_hcp_flat_mask
|
| 348 |
+
from mae_utils.flat import create_hcp_flat
|
| 349 |
+
from mae_utils.flat import batch_unmask
|
| 350 |
+
import mae_utils.visualize as vis
|
| 351 |
+
|
| 352 |
+
flat_mask = load_hcp_flat_mask(hcp_flat_path)
|
| 353 |
+
|
| 354 |
+
mae_model = flat_models.mae_vit_large_fmri(
|
| 355 |
+
patch_size=patch_size,
|
| 356 |
+
decoder_embed_dim=decoder_embed_dim,
|
| 357 |
+
t_patch_size=t_patch_size,
|
| 358 |
+
pred_t_dim=pred_t_dim,
|
| 359 |
+
decoder_depth=4,
|
| 360 |
+
cls_embed=cls_embed,
|
| 361 |
+
norm_pix_loss=norm_pix_loss,
|
| 362 |
+
no_qkv_bias=no_qkv_bias,
|
| 363 |
+
sep_pos_embed=sep_pos_embed,
|
| 364 |
+
trunc_init=trunc_init,
|
| 365 |
+
pct_masks_to_decode=pct_masks_to_decode,
|
| 366 |
+
img_mask=flat_mask,
|
| 367 |
+
)
|
| 368 |
+
|
| 369 |
+
|
| 370 |
+
# In[8]:
|
| 371 |
+
|
| 372 |
+
|
| 373 |
+
checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')]
|
| 374 |
+
|
| 375 |
+
if utils.is_interactive():
|
| 376 |
+
latest_checkpoint = "epoch99.pth"
|
| 377 |
+
else:
|
| 378 |
+
latest_checkpoint = sys.argv[2]
|
| 379 |
+
print(f"latest_checkpoint: {latest_checkpoint}")
|
| 380 |
+
|
| 381 |
+
# Load the checkpoint
|
| 382 |
+
checkpoint_path = os.path.join(outdir, latest_checkpoint)
|
| 383 |
+
|
| 384 |
+
state = torch.load(checkpoint_path)
|
| 385 |
+
mae_model.load_state_dict(state["model_state_dict"], strict=False)
|
| 386 |
+
mae_model.to(device)
|
| 387 |
+
|
| 388 |
+
print(f"\nLoaded checkpoint {latest_checkpoint} from {outdir}\n")
|
| 389 |
+
|
| 390 |
+
|
| 391 |
+
# In[9]:
|
| 392 |
+
|
| 393 |
+
|
| 394 |
+
class LinearClassifier(nn.Module):
|
| 395 |
+
def __init__(self, input_dim, num_classes):
|
| 396 |
+
super(LinearClassifier, self).__init__()
|
| 397 |
+
self.linear = nn.Linear(input_dim, num_classes)
|
| 398 |
+
|
| 399 |
+
def forward(self, x):
|
| 400 |
+
# Flatten the input except for the batch dimension
|
| 401 |
+
x = x.view(x.size(0), -1)
|
| 402 |
+
out = self.linear(x)
|
| 403 |
+
return out # Raw logits
|
| 404 |
+
|
| 405 |
+
# Determine the input dimension from a single sample
|
| 406 |
+
# Assuming images are of shape [1, 16, 144, 320]
|
| 407 |
+
input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])
|
| 408 |
+
print(f"Input dimension: {input_dim}")
|
| 409 |
+
|
| 410 |
+
|
| 411 |
+
# In[10]:
|
| 412 |
+
|
| 413 |
+
|
| 414 |
+
class FullModel(nn.Module):
|
| 415 |
+
def __init__(self, lc_model, mae_model):
|
| 416 |
+
super(FullModel, self).__init__()
|
| 417 |
+
self.lc_model = lc_model
|
| 418 |
+
self.mae_model = mae_model
|
| 419 |
+
|
| 420 |
+
|
| 421 |
+
def forward(self, x, gsr):
|
| 422 |
+
x = self.mae_model(x, global_pool=global_pool, forward_features = True)
|
| 423 |
+
x = self.lc_model(x)
|
| 424 |
+
return x
|
| 425 |
+
|
| 426 |
+
|
| 427 |
+
# In[11]:
|
| 428 |
+
|
| 429 |
+
|
| 430 |
+
# Initialize the model
|
| 431 |
+
lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)
|
| 432 |
+
|
| 433 |
+
model = FullModel(lc_model, mae_model)
|
| 434 |
+
|
| 435 |
+
# Move the model to the GPU
|
| 436 |
+
model.to(device)
|
| 437 |
+
|
| 438 |
+
# Define loss function
|
| 439 |
+
criterion = nn.CrossEntropyLoss()
|
| 440 |
+
|
| 441 |
+
# Define optimizer with L2 regularization (weight_decay)
|
| 442 |
+
learning_rate = 1e-4
|
| 443 |
+
weight_decay = 1e-5 # Adjust based on your needs
|
| 444 |
+
optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)
|
| 445 |
+
num_epochs = 20 # Adjust as needed
|
| 446 |
+
|
| 447 |
+
|
| 448 |
+
# ### Data
|
| 449 |
+
|
| 450 |
+
# In[16]:
|
| 451 |
+
|
| 452 |
+
|
| 453 |
+
import uuid
|
| 454 |
+
|
| 455 |
+
myuuid = uuid.uuid4()
|
| 456 |
+
str(myuuid)
|
| 457 |
+
|
| 458 |
+
|
| 459 |
+
# In[17]:
|
| 460 |
+
|
| 461 |
+
|
| 462 |
+
import wandb
|
| 463 |
+
|
| 464 |
+
if utils.is_interactive():
|
| 465 |
+
print("Running in interactive notebook. Disabling W&B and ckpt saving.")
|
| 466 |
+
wandb_log = True
|
| 467 |
+
save_ckpt = True
|
| 468 |
+
|
| 469 |
+
if wandb_log:
|
| 470 |
+
wandb_project = 'fMRI-foundation-model'
|
| 471 |
+
wandb_config = {
|
| 472 |
+
"model_name": model_name+'_HCP_FT',
|
| 473 |
+
"batch_size": batch_size,
|
| 474 |
+
"learning_rate": learning_rate,
|
| 475 |
+
"weight_decay": weight_decay,
|
| 476 |
+
"num_epochs": num_epochs,
|
| 477 |
+
"seed": seed,
|
| 478 |
+
}
|
| 479 |
+
print("wandb_config:\n", wandb_config)
|
| 480 |
+
random_id = str(uuid.uuid4())
|
| 481 |
+
print("wandb_id:", "HCPflat_raw" + f"_{random_id}")
|
| 482 |
+
wandb.init(
|
| 483 |
+
id=model_name+'_HCP_FT' + f"_{random_id}",
|
| 484 |
+
project=wandb_project,
|
| 485 |
+
name=model_name+'_HCP_FT',
|
| 486 |
+
config=wandb_config,
|
| 487 |
+
resume="allow",
|
| 488 |
+
)
|
| 489 |
+
|
| 490 |
+
|
| 491 |
+
# In[13]:
|
| 492 |
+
|
| 493 |
+
|
| 494 |
+
for epoch in range(num_epochs):
|
| 495 |
+
running_train_loss = 0.0
|
| 496 |
+
correct_train = 0
|
| 497 |
+
total_train = 0
|
| 498 |
+
step = 0
|
| 499 |
+
|
| 500 |
+
# with torch.amp.autocast(device_type='cuda'):
|
| 501 |
+
# Training Phase
|
| 502 |
+
model.train()
|
| 503 |
+
for batch in tqdm(train_dl, desc=f"Epoch {epoch+1}/{num_epochs} - Training"):
|
| 504 |
+
optimizer.zero_grad()
|
| 505 |
+
images = batch[0].to(device).float().unsqueeze(1) #fix this # Shape: [batch_size, 1, 16, 144, 320]
|
| 506 |
+
labels = batch[1]['trial_type'] # List of labels
|
| 507 |
+
|
| 508 |
+
encoded_labels = label_encoder.transform(labels)
|
| 509 |
+
encoded_labels = torch.tensor(encoded_labels, dtype=torch.long).to(device) # Shape: [batch_size]
|
| 510 |
+
|
| 511 |
+
# Forward pass
|
| 512 |
+
outputs = model(images, gsr=gsr) # Shape: [num_train_samples, num_classes]
|
| 513 |
+
|
| 514 |
+
# Compute loss
|
| 515 |
+
loss = criterion(outputs, encoded_labels)
|
| 516 |
+
|
| 517 |
+
# Backward pass and optimization
|
| 518 |
+
loss.backward()
|
| 519 |
+
optimizer.step()
|
| 520 |
+
|
| 521 |
+
# Accumulate loss
|
| 522 |
+
running_train_loss += loss.item() * images.size(0)
|
| 523 |
+
|
| 524 |
+
|
| 525 |
+
# Calculate accuracy
|
| 526 |
+
_, predicted = torch.max(outputs, 1)
|
| 527 |
+
|
| 528 |
+
correct_train += (predicted == encoded_labels).sum().item()
|
| 529 |
+
total_train += encoded_labels.size(0)
|
| 530 |
+
|
| 531 |
+
step = step + 1
|
| 532 |
+
if step % 100 == 0:
|
| 533 |
+
print(f"Step [{step}/{len(train_dl)}] - Training Loss: {loss.item():.4f} - Training Accuracy: {100 * correct_train / total_train:.2f}%")
|
| 534 |
+
# thth
|
| 535 |
+
|
| 536 |
+
epoch_train_loss = running_train_loss / total_train if total_train > 0 else 0.0
|
| 537 |
+
train_accuracy = 100 * correct_train / total_train if total_train > 0 else 0.0
|
| 538 |
+
|
| 539 |
+
# Validation Phase
|
| 540 |
+
model.eval()
|
| 541 |
+
running_val_loss = 0.0
|
| 542 |
+
correct_val = 0
|
| 543 |
+
total_val = 0
|
| 544 |
+
|
| 545 |
+
with torch.no_grad():
|
| 546 |
+
for batch in tqdm(test_dl, desc=f"Epoch {epoch+1}/{num_epochs} - Validation"):
|
| 547 |
+
|
| 548 |
+
images = batch[0].to(device).float().unsqueeze(1) #fix this
|
| 549 |
+
labels = batch[1]['trial_type']
|
| 550 |
+
|
| 551 |
+
# Encode labels to integer indices
|
| 552 |
+
encoded_labels = label_encoder.transform(labels)
|
| 553 |
+
encoded_labels = torch.tensor(encoded_labels, dtype=torch.long).to(device)
|
| 554 |
+
|
| 555 |
+
|
| 556 |
+
# Forward pass
|
| 557 |
+
outputs = model(images, gsr=gsr)
|
| 558 |
+
|
| 559 |
+
# Compute loss
|
| 560 |
+
loss = criterion(outputs, encoded_labels)
|
| 561 |
+
|
| 562 |
+
# Accumulate loss
|
| 563 |
+
running_val_loss += loss.item() * images.size(0)
|
| 564 |
+
|
| 565 |
+
# Calculate accuracy
|
| 566 |
+
_, predicted = torch.max(outputs, 1)
|
| 567 |
+
correct_val += (predicted == encoded_labels).sum().item()
|
| 568 |
+
total_val += encoded_labels.size(0)
|
| 569 |
+
|
| 570 |
+
|
| 571 |
+
|
| 572 |
+
epoch_val_loss = running_val_loss / total_val if total_val > 0 else 0.0
|
| 573 |
+
val_accuracy = 100 * correct_val / total_val if total_val > 0 else 0.0
|
| 574 |
+
|
| 575 |
+
print(f"Epoch [{epoch+1}/{num_epochs}] "
|
| 576 |
+
f"- Training Loss: {epoch_train_loss:.4f}, Training Accuracy: {train_accuracy:.2f}% "
|
| 577 |
+
f"- Validation Loss: {epoch_val_loss:.4f}, Validation Accuracy: {val_accuracy:.2f}%")
|
| 578 |
+
|
| 579 |
+
if wandb_log:
|
| 580 |
+
wandb.log({
|
| 581 |
+
"epoch_train_loss": epoch_train_loss,
|
| 582 |
+
"epoch_val_loss": epoch_val_loss,
|
| 583 |
+
"train_accuracy": train_accuracy,
|
| 584 |
+
"val_accuracy": val_accuracy,
|
| 585 |
+
})
|
| 586 |
+
if save_ckpt:
|
| 587 |
+
outdir = os.path.abspath(f'checkpoints/{model_name+"HCP_FT"}')
|
| 588 |
+
os.makedirs(outdir, exist_ok=True)
|
| 589 |
+
print("outdir", outdir)
|
| 590 |
+
# Save model and config
|
| 591 |
+
torch.save(model.state_dict(), f"{outdir}/model.pth")
|
| 592 |
+
with open(f"{outdir}/config.yaml", 'w') as f:
|
| 593 |
+
yaml.dump(wandb_config, f)
|
| 594 |
+
print(f"Saved model and config to {outdir}")
|
| 595 |
+
|
| 596 |
+
|
fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/output.log
ADDED
|
@@ -0,0 +1,59 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Epoch 1/20 - Training: 23%|██▎ | 3199/13913 [18:32<1:01:49, 2.89it/s]
|
| 2 |
+
Step [100/13913] - Training Loss: 1.9649 - Training Accuracy: 59.38%
|
| 3 |
+
Step [200/13913] - Training Loss: 1.0904 - Training Accuracy: 71.19%
|
| 4 |
+
Step [300/13913] - Training Loss: 0.0775 - Training Accuracy: 75.50%
|
| 5 |
+
Step [400/13913] - Training Loss: 0.6052 - Training Accuracy: 78.84%
|
| 6 |
+
Step [500/13913] - Training Loss: 0.0226 - Training Accuracy: 80.95%
|
| 7 |
+
Step [600/13913] - Training Loss: 0.2728 - Training Accuracy: 82.54%
|
| 8 |
+
Step [700/13913] - Training Loss: 0.1662 - Training Accuracy: 83.70%
|
| 9 |
+
Step [800/13913] - Training Loss: 0.0385 - Training Accuracy: 84.89%
|
| 10 |
+
Step [900/13913] - Training Loss: 0.2377 - Training Accuracy: 85.61%
|
| 11 |
+
Step [1000/13913] - Training Loss: 0.8172 - Training Accuracy: 85.83%
|
| 12 |
+
Step [1100/13913] - Training Loss: 0.2276 - Training Accuracy: 86.68%
|
| 13 |
+
Step [1200/13913] - Training Loss: 0.0118 - Training Accuracy: 87.36%
|
| 14 |
+
Step [1300/13913] - Training Loss: 1.0419 - Training Accuracy: 87.68%
|
| 15 |
+
Step [1400/13913] - Training Loss: 0.8943 - Training Accuracy: 87.96%
|
| 16 |
+
Step [1500/13913] - Training Loss: 0.2801 - Training Accuracy: 88.23%
|
| 17 |
+
Step [1600/13913] - Training Loss: 0.6734 - Training Accuracy: 88.58%
|
| 18 |
+
Step [1700/13913] - Training Loss: 0.6202 - Training Accuracy: 88.88%
|
| 19 |
+
Step [1800/13913] - Training Loss: 0.0159 - Training Accuracy: 89.12%
|
| 20 |
+
Step [1900/13913] - Training Loss: 0.0682 - Training Accuracy: 89.37%
|
| 21 |
+
Step [2000/13913] - Training Loss: 0.3378 - Training Accuracy: 89.52%
|
| 22 |
+
Step [2100/13913] - Training Loss: 0.0509 - Training Accuracy: 89.77%
|
| 23 |
+
Step [2200/13913] - Training Loss: 0.1161 - Training Accuracy: 89.99%
|
| 24 |
+
Step [2300/13913] - Training Loss: 0.0025 - Training Accuracy: 90.12%
|
| 25 |
+
Step [2400/13913] - Training Loss: 0.5385 - Training Accuracy: 90.25%
|
| 26 |
+
Step [2500/13913] - Training Loss: 0.0003 - Training Accuracy: 90.47%
|
| 27 |
+
Step [2600/13913] - Training Loss: 0.0962 - Training Accuracy: 90.52%
|
| 28 |
+
Step [2700/13913] - Training Loss: 0.0480 - Training Accuracy: 90.62%
|
| 29 |
+
Step [2800/13913] - Training Loss: 0.0004 - Training Accuracy: 90.71%
|
| 30 |
+
Step [2900/13913] - Training Loss: 0.3049 - Training Accuracy: 90.84%
|
| 31 |
+
Step [3000/13913] - Training Loss: 0.0339 - Training Accuracy: 90.95%
|
| 32 |
+
Step [3100/13913] - Training Loss: 0.5572 - Training Accuracy: 91.02%
|
| 33 |
+
Step [3200/13913] - Training Loss: 0.0673 - Training Accuracy: 91.11%
|
| 34 |
+
Step [3300/13913] - Training Loss: 0.0018 - Training Accuracy: 91.21%
|
| 35 |
+
Step [3400/13913] - Training Loss: 0.0048 - Training Accuracy: 91.28%
|
| 36 |
+
Step [3500/13913] - Training Loss: 1.2120 - Training Accuracy: 91.33%
|
| 37 |
+
Step [3600/13913] - Training Loss: 0.0542 - Training Accuracy: 91.37%
|
| 38 |
+
Step [3700/13913] - Training Loss: 0.0016 - Training Accuracy: 91.43%
|
| 39 |
+
Step [3800/13913] - Training Loss: 0.0582 - Training Accuracy: 91.52%
|
| 40 |
+
Step [3900/13913] - Training Loss: 0.0960 - Training Accuracy: 91.60%
|
| 41 |
+
Step [4000/13913] - Training Loss: 0.0020 - Training Accuracy: 91.68%
|
| 42 |
+
Step [4100/13913] - Training Loss: 0.0043 - Training Accuracy: 91.76%
|
| 43 |
+
Step [4200/13913] - Training Loss: 0.0029 - Training Accuracy: 91.77%
|
| 44 |
+
Step [4300/13913] - Training Loss: 0.0107 - Training Accuracy: 91.83%
|
| 45 |
+
Step [4400/13913] - Training Loss: 0.1122 - Training Accuracy: 91.86%
|
| 46 |
+
Step [4500/13913] - Training Loss: 0.1595 - Training Accuracy: 91.92%
|
| 47 |
+
Step [4600/13913] - Training Loss: 0.0453 - Training Accuracy: 91.97%
|
| 48 |
+
Step [4700/13913] - Training Loss: 0.2770 - Training Accuracy: 92.05%
|
| 49 |
+
Step [4800/13913] - Training Loss: 0.0057 - Training Accuracy: 92.09%
|
| 50 |
+
Step [4900/13913] - Training Loss: 0.0120 - Training Accuracy: 92.16%
|
| 51 |
+
Step [5000/13913] - Training Loss: 0.0235 - Training Accuracy: 92.24%
|
| 52 |
+
Step [5100/13913] - Training Loss: 0.3907 - Training Accuracy: 92.32%
|
| 53 |
+
Step [5200/13913] - Training Loss: 0.4558 - Training Accuracy: 92.34%
|
| 54 |
+
Step [5300/13913] - Training Loss: 0.0051 - Training Accuracy: 92.40%
|
| 55 |
+
Step [5400/13913] - Training Loss: 0.0017 - Training Accuracy: 92.46%
|
| 56 |
+
Step [5500/13913] - Training Loss: 1.3554 - Training Accuracy: 92.50%
|
| 57 |
+
Step [5600/13913] - Training Loss: 0.0617 - Training Accuracy: 92.55%
|
| 58 |
+
Step [5700/13913] - Training Loss: 0.2618 - Training Accuracy: 92.60%
|
| 59 |
+
Step [5800/13913] - Training Loss: 0.0192 - Training Accuracy: 92.60%
|
fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/requirements.txt
ADDED
|
@@ -0,0 +1,198 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
protobuf==5.28.2
|
| 2 |
+
imageio==2.35.1
|
| 3 |
+
MarkupSafe==3.0.0
|
| 4 |
+
regex==2024.9.11
|
| 5 |
+
matplotlib==3.9.2
|
| 6 |
+
notebook==7.2.2
|
| 7 |
+
debugpy==1.8.6
|
| 8 |
+
aiosignal==1.3.1
|
| 9 |
+
jupyter_core==5.7.2
|
| 10 |
+
torchaudio==2.4.1+cu121
|
| 11 |
+
python-json-logger==2.0.7
|
| 12 |
+
six==1.16.0
|
| 13 |
+
scikit-image==0.24.0
|
| 14 |
+
types-python-dateutil==2.9.0.20241003
|
| 15 |
+
PyYAML==6.0.2
|
| 16 |
+
httpcore==1.0.6
|
| 17 |
+
clip==1.0
|
| 18 |
+
babel==2.16.0
|
| 19 |
+
webcolors==24.8.0
|
| 20 |
+
omegaconf==2.3.0
|
| 21 |
+
webencodings==0.5.1
|
| 22 |
+
kiwisolver==1.4.7
|
| 23 |
+
uri-template==1.3.0
|
| 24 |
+
diffusers==0.23.0
|
| 25 |
+
idna==3.10
|
| 26 |
+
fsspec==2024.9.0
|
| 27 |
+
parso==0.8.4
|
| 28 |
+
setuptools==65.5.0
|
| 29 |
+
tornado==6.4.1
|
| 30 |
+
webdataset==0.2.100
|
| 31 |
+
decord==0.6.0
|
| 32 |
+
nvidia-curand-cu12==10.3.2.106
|
| 33 |
+
ipykernel==6.29.5
|
| 34 |
+
jupyter==1.1.1
|
| 35 |
+
pexpect==4.9.0
|
| 36 |
+
kornia_rs==0.1.5
|
| 37 |
+
iopath==0.1.10
|
| 38 |
+
async-lru==2.0.4
|
| 39 |
+
future==1.0.0
|
| 40 |
+
torchvision==0.19.1+cu121
|
| 41 |
+
botocore==1.34.162
|
| 42 |
+
cycler==0.12.1
|
| 43 |
+
tzdata==2024.2
|
| 44 |
+
jupyter_server_terminals==0.5.3
|
| 45 |
+
click==8.1.7
|
| 46 |
+
einops==0.8.0
|
| 47 |
+
pyzmq==26.2.0
|
| 48 |
+
jupyter_client==8.6.3
|
| 49 |
+
nbconvert==7.16.4
|
| 50 |
+
scikit-learn==1.5.2
|
| 51 |
+
executing==2.1.0
|
| 52 |
+
asttokens==2.4.1
|
| 53 |
+
docker-pycreds==0.4.0
|
| 54 |
+
matplotlib-inline==0.1.7
|
| 55 |
+
overrides==7.7.0
|
| 56 |
+
websocket-client==1.8.0
|
| 57 |
+
nbformat==5.10.4
|
| 58 |
+
elbow==0.1.1
|
| 59 |
+
contourpy==1.3.0
|
| 60 |
+
nvidia-cudnn-cu12==9.1.0.70
|
| 61 |
+
transformers==4.44.2
|
| 62 |
+
gitdb==4.0.11
|
| 63 |
+
jupyterlab_nvdashboard==0.11.0
|
| 64 |
+
lazy_loader==0.4
|
| 65 |
+
jsonpointer==3.0.0
|
| 66 |
+
notebook_shim==0.2.4
|
| 67 |
+
nvidia-nccl-cu12==2.20.5
|
| 68 |
+
ffmpeg-python==0.2.0
|
| 69 |
+
triton==3.0.0
|
| 70 |
+
mistune==3.0.2
|
| 71 |
+
python-dateutil==2.9.0.post0
|
| 72 |
+
beautifulsoup4==4.12.3
|
| 73 |
+
nbclient==0.10.0
|
| 74 |
+
h5py==3.12.1
|
| 75 |
+
ftfy==6.2.3
|
| 76 |
+
zipp==3.20.2
|
| 77 |
+
ptyprocess==0.7.0
|
| 78 |
+
huggingface-hub==0.25.1
|
| 79 |
+
pytz==2024.2
|
| 80 |
+
jupyterlab_pygments==0.3.0
|
| 81 |
+
nvidia-cublas-cu12==12.1.3.1
|
| 82 |
+
pandocfilters==1.5.1
|
| 83 |
+
Jinja2==3.1.4
|
| 84 |
+
arrow==1.3.0
|
| 85 |
+
rpds-py==0.20.0
|
| 86 |
+
jupyter_server==2.14.2
|
| 87 |
+
simplejson==3.19.3
|
| 88 |
+
networkx==3.3
|
| 89 |
+
packaging==24.1
|
| 90 |
+
traitlets==5.14.3
|
| 91 |
+
pandas==2.2.3
|
| 92 |
+
xformers==0.0.22.post7
|
| 93 |
+
lightning-utilities==0.11.7
|
| 94 |
+
tifffile==2024.9.20
|
| 95 |
+
nvidia-cuda-cupti-cu12==12.1.105
|
| 96 |
+
mpmath==1.3.0
|
| 97 |
+
GitPython==3.1.43
|
| 98 |
+
scipy==1.14.1
|
| 99 |
+
jsonschema==4.23.0
|
| 100 |
+
prompt_toolkit==3.0.48
|
| 101 |
+
s3transfer==0.10.2
|
| 102 |
+
multidict==6.1.0
|
| 103 |
+
bleach==6.1.0
|
| 104 |
+
sentry-sdk==2.15.0
|
| 105 |
+
nibabel==5.2.1
|
| 106 |
+
accelerate==1.0.0
|
| 107 |
+
pyarrow==17.0.0
|
| 108 |
+
threadpoolctl==3.5.0
|
| 109 |
+
attrs==24.2.0
|
| 110 |
+
rfc3986-validator==0.1.1
|
| 111 |
+
nvidia-cuda-runtime-cu12==12.1.105
|
| 112 |
+
ipywidgets==8.1.5
|
| 113 |
+
frozenlist==1.4.1
|
| 114 |
+
pycparser==2.22
|
| 115 |
+
jupyterlab_server==2.27.3
|
| 116 |
+
nvidia-cuda-nvrtc-cu12==12.1.105
|
| 117 |
+
yarl==1.13.1
|
| 118 |
+
setproctitle==1.3.3
|
| 119 |
+
isoduration==20.11.0
|
| 120 |
+
Pygments==2.18.0
|
| 121 |
+
jedi==0.19.1
|
| 122 |
+
boto3==1.34.57
|
| 123 |
+
tokenizers==0.19.1
|
| 124 |
+
referencing==0.35.1
|
| 125 |
+
rfc3339-validator==0.1.4
|
| 126 |
+
pillow==10.4.0
|
| 127 |
+
jupyterlab==4.2.5
|
| 128 |
+
stack-data==0.6.3
|
| 129 |
+
h11==0.14.0
|
| 130 |
+
anyio==4.6.0
|
| 131 |
+
nilearn==0.10.4
|
| 132 |
+
nvidia-cusolver-cu12==11.4.5.107
|
| 133 |
+
tinycss2==1.3.0
|
| 134 |
+
defusedxml==0.7.1
|
| 135 |
+
argon2-cffi-bindings==21.2.0
|
| 136 |
+
soupsieve==2.6
|
| 137 |
+
nest-asyncio==1.6.0
|
| 138 |
+
torchmetrics==1.3.0.post0
|
| 139 |
+
tqdm==4.66.5
|
| 140 |
+
cffi==1.17.1
|
| 141 |
+
charset-normalizer==3.3.2
|
| 142 |
+
jsonschema-specifications==2023.12.1
|
| 143 |
+
decorator==5.1.1
|
| 144 |
+
open_clip_torch==2.26.1
|
| 145 |
+
jupyter-events==0.10.0
|
| 146 |
+
smart-open==7.0.5
|
| 147 |
+
antlr4-python3-runtime==4.9.3
|
| 148 |
+
prometheus_client==0.21.0
|
| 149 |
+
kornia==0.7.3
|
| 150 |
+
typing_extensions==4.12.2
|
| 151 |
+
sniffio==1.3.1
|
| 152 |
+
joblib==1.4.2
|
| 153 |
+
comm==0.2.2
|
| 154 |
+
aiohappyeyeballs==2.4.3
|
| 155 |
+
numpy==2.1.2
|
| 156 |
+
braceexpand==0.1.7
|
| 157 |
+
certifi==2024.8.30
|
| 158 |
+
psutil==6.0.0
|
| 159 |
+
pyparsing==3.1.4
|
| 160 |
+
pure_eval==0.2.3
|
| 161 |
+
nvidia-cusparse-cu12==12.1.0.106
|
| 162 |
+
wandb==0.18.3
|
| 163 |
+
urllib3==2.2.3
|
| 164 |
+
smmap==5.0.1
|
| 165 |
+
platformdirs==4.3.6
|
| 166 |
+
torch==2.4.1+cu121
|
| 167 |
+
requests==2.32.3
|
| 168 |
+
json5==0.9.25
|
| 169 |
+
nvidia-nvjitlink-cu12==12.6.77
|
| 170 |
+
jupyterlab_widgets==3.0.13
|
| 171 |
+
lxml==5.3.0
|
| 172 |
+
httpx==0.27.2
|
| 173 |
+
opencv-python==4.6.0.66
|
| 174 |
+
portalocker==2.10.1
|
| 175 |
+
pytorch-lightning==2.0.1
|
| 176 |
+
sympy==1.13.3
|
| 177 |
+
wcwidth==0.2.13
|
| 178 |
+
jmespath==1.0.1
|
| 179 |
+
fqdn==1.5.1
|
| 180 |
+
pynvml==11.5.3
|
| 181 |
+
pip==24.0
|
| 182 |
+
wrapt==1.16.0
|
| 183 |
+
aiohttp==3.10.9
|
| 184 |
+
filelock==3.16.1
|
| 185 |
+
fonttools==4.54.1
|
| 186 |
+
fastjsonschema==2.20.0
|
| 187 |
+
jupyter-console==6.6.3
|
| 188 |
+
widgetsnbextension==4.0.13
|
| 189 |
+
timm==1.0.9
|
| 190 |
+
nvidia-cufft-cu12==11.0.2.54
|
| 191 |
+
ipython==8.28.0
|
| 192 |
+
nvidia-nvtx-cu12==12.1.105
|
| 193 |
+
jupyter-lsp==2.2.5
|
| 194 |
+
safetensors==0.4.5
|
| 195 |
+
terminado==0.18.1
|
| 196 |
+
argon2-cffi==23.1.0
|
| 197 |
+
Send2Trash==1.8.3
|
| 198 |
+
importlib_metadata==8.5.0
|
fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/wandb-metadata.json
ADDED
|
@@ -0,0 +1,131 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"os": "Linux-5.15.0-1058-aws-x86_64-with-glibc2.31",
|
| 3 |
+
"python": "3.11.10",
|
| 4 |
+
"startedAt": "2024-10-23T04:13:25.994897Z",
|
| 5 |
+
"args": [
|
| 6 |
+
"HCPflat_large_gsrFalse_",
|
| 7 |
+
"epoch99.pth"
|
| 8 |
+
],
|
| 9 |
+
"program": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py",
|
| 10 |
+
"codePath": "src/HCP_downstream_finetune.py",
|
| 11 |
+
"git": {
|
| 12 |
+
"remote": "https://github.com/MedARC-AI/fMRI-foundation-model",
|
| 13 |
+
"commit": "b1ba684ae7a5cc4155cc046b0abe613de09bf700"
|
| 14 |
+
},
|
| 15 |
+
"email": "torrico.villanueva.cesar.kadir@gmail.com",
|
| 16 |
+
"root": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src",
|
| 17 |
+
"host": "ip-10-0-139-117",
|
| 18 |
+
"username": "ckadirt",
|
| 19 |
+
"executable": "/admin/home-ckadirt/foundation_env/bin/python",
|
| 20 |
+
"codePathLocal": "HCP_downstream_finetune.py",
|
| 21 |
+
"cpu_count": 96,
|
| 22 |
+
"cpu_count_logical": 192,
|
| 23 |
+
"gpu": "[NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3]",
|
| 24 |
+
"gpu_count": 8,
|
| 25 |
+
"disk": {
|
| 26 |
+
"/": {
|
| 27 |
+
"total": "249555763200",
|
| 28 |
+
"used": "181820596224"
|
| 29 |
+
}
|
| 30 |
+
},
|
| 31 |
+
"memory": {
|
| 32 |
+
"total": "2147443380224"
|
| 33 |
+
},
|
| 34 |
+
"cpu": {
|
| 35 |
+
"count": 96,
|
| 36 |
+
"countLogical": 192
|
| 37 |
+
},
|
| 38 |
+
"gpu_nvidia": [
|
| 39 |
+
{
|
| 40 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 41 |
+
"memoryTotal": "85520809984",
|
| 42 |
+
"cudaCores": 16896,
|
| 43 |
+
"architecture": "Hopper"
|
| 44 |
+
},
|
| 45 |
+
{
|
| 46 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 47 |
+
"memoryTotal": "85520809984",
|
| 48 |
+
"cudaCores": 16896,
|
| 49 |
+
"architecture": "Hopper"
|
| 50 |
+
},
|
| 51 |
+
{
|
| 52 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 53 |
+
"memoryTotal": "85520809984",
|
| 54 |
+
"cudaCores": 16896,
|
| 55 |
+
"architecture": "Hopper"
|
| 56 |
+
},
|
| 57 |
+
{
|
| 58 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 59 |
+
"memoryTotal": "85520809984",
|
| 60 |
+
"cudaCores": 16896,
|
| 61 |
+
"architecture": "Hopper"
|
| 62 |
+
},
|
| 63 |
+
{
|
| 64 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 65 |
+
"memoryTotal": "85520809984",
|
| 66 |
+
"cudaCores": 16896,
|
| 67 |
+
"architecture": "Hopper"
|
| 68 |
+
},
|
| 69 |
+
{
|
| 70 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 71 |
+
"memoryTotal": "85520809984",
|
| 72 |
+
"cudaCores": 16896,
|
| 73 |
+
"architecture": "Hopper"
|
| 74 |
+
},
|
| 75 |
+
{
|
| 76 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 77 |
+
"memoryTotal": "85520809984",
|
| 78 |
+
"cudaCores": 16896,
|
| 79 |
+
"architecture": "Hopper"
|
| 80 |
+
},
|
| 81 |
+
{
|
| 82 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 83 |
+
"memoryTotal": "85520809984",
|
| 84 |
+
"cudaCores": 16896,
|
| 85 |
+
"architecture": "Hopper"
|
| 86 |
+
}
|
| 87 |
+
],
|
| 88 |
+
"slurm": {
|
| 89 |
+
"cluster_name": "sagemaker2",
|
| 90 |
+
"conf": "/opt/slurm/etc/slurm.conf",
|
| 91 |
+
"cpus_on_node": "20",
|
| 92 |
+
"gpus_on_node": "1",
|
| 93 |
+
"gpus_per_task": "1",
|
| 94 |
+
"gtids": "0",
|
| 95 |
+
"job_account": "fmri",
|
| 96 |
+
"job_cpus_per_node": "20",
|
| 97 |
+
"job_end_time": "1729699988",
|
| 98 |
+
"job_gid": "1879800513",
|
| 99 |
+
"job_gpus": "4",
|
| 100 |
+
"job_id": "528152",
|
| 101 |
+
"job_name": "finetuneHCP",
|
| 102 |
+
"job_nodelist": "ip-10-0-139-117",
|
| 103 |
+
"job_num_nodes": "1",
|
| 104 |
+
"job_partition": "p5",
|
| 105 |
+
"job_qos": "idle",
|
| 106 |
+
"job_start_time": "1729656788",
|
| 107 |
+
"job_uid": "1879804696",
|
| 108 |
+
"job_user": "ckadirt",
|
| 109 |
+
"jobid": "528152",
|
| 110 |
+
"localid": "0",
|
| 111 |
+
"mem_per_cpu": "11500",
|
| 112 |
+
"nnodes": "1",
|
| 113 |
+
"node_aliases": "(null)",
|
| 114 |
+
"nodeid": "0",
|
| 115 |
+
"nodelist": "ip-10-0-139-117",
|
| 116 |
+
"nprocs": "1",
|
| 117 |
+
"ntasks": "1",
|
| 118 |
+
"ntasks_per_node": "1",
|
| 119 |
+
"prio_process": "0",
|
| 120 |
+
"procid": "0",
|
| 121 |
+
"script_context": "prolog_task",
|
| 122 |
+
"submit_dir": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src",
|
| 123 |
+
"submit_host": "ip-172-17-12-61",
|
| 124 |
+
"task_pid": "3335147",
|
| 125 |
+
"tasks_per_node": "1",
|
| 126 |
+
"topology_addr": "ip-10-0-139-117",
|
| 127 |
+
"topology_addr_pattern": "node",
|
| 128 |
+
"working_cluster": "sagemaker2:ip-172-17-63-161:6817:9984:109"
|
| 129 |
+
},
|
| 130 |
+
"cudaVersion": "12.2"
|
| 131 |
+
}
|
fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug-core.log
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-23T04:13:25.195730514Z","level":"INFO","msg":"started logging, with flags","port-filename":"/tmp/tmp2j4o8616/port-3335224.txt","pid":3335224,"debug":false,"disable-analytics":false}
|
| 2 |
+
{"time":"2024-10-23T04:13:25.196029969Z","level":"INFO","msg":"FeatureState","shutdownOnParentExitEnabled":false}
|
| 3 |
+
{"time":"2024-10-23T04:13:25.199363081Z","level":"INFO","msg":"Will exit if parent process dies.","ppid":3335224}
|
| 4 |
+
{"time":"2024-10-23T04:13:25.199355901Z","level":"INFO","msg":"server is running","addr":{"IP":"127.0.0.1","Port":44407,"Zone":""}}
|
| 5 |
+
{"time":"2024-10-23T04:13:25.388406125Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"127.0.0.1:59886"}
|
| 6 |
+
{"time":"2024-10-23T04:13:25.995203792Z","level":"INFO","msg":"handleInformInit: received","streamId":"HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532","id":"127.0.0.1:59886"}
|
| 7 |
+
{"time":"2024-10-23T04:13:26.086124707Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532","id":"127.0.0.1:59886"}
|
fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug-internal.log
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-23T04:13:26.025058265Z","level":"INFO","msg":"using version","core version":"0.18.3"}
|
| 2 |
+
{"time":"2024-10-23T04:13:26.025081085Z","level":"INFO","msg":"created symlink","path":"/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug-core.log"}
|
| 3 |
+
{"time":"2024-10-23T04:13:26.027044192Z","level":"ERROR","msg":"dialing: google: could not find default credentials. See https://cloud.google.com/docs/authentication/external/set-up-adc for more information"}
|
| 4 |
+
{"time":"2024-10-23T04:13:26.086056755Z","level":"INFO","msg":"created new stream","id":"HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532"}
|
| 5 |
+
{"time":"2024-10-23T04:13:26.086114916Z","level":"INFO","msg":"stream: started","id":"HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532"}
|
| 6 |
+
{"time":"2024-10-23T04:13:26.086145987Z","level":"INFO","msg":"sender: started","stream_id":{"value":"HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532"}}
|
| 7 |
+
{"time":"2024-10-23T04:13:26.086149037Z","level":"INFO","msg":"handler: started","stream_id":{"value":"HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532"}}
|
| 8 |
+
{"time":"2024-10-23T04:13:26.086128047Z","level":"INFO","msg":"writer: Do: started","stream_id":{"value":"HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532"}}
|
| 9 |
+
{"time":"2024-10-23T04:13:26.629699111Z","level":"INFO","msg":"wandb-core","!BADKEY":null}
|
| 10 |
+
{"time":"2024-10-23T04:13:26.639782248Z","level":"INFO","msg":"Starting system monitor"}
|
| 11 |
+
{"time":"2024-10-23T04:13:26.752242533Z","level":"ERROR","msg":"git repo not found","error":"repository does not exist"}
|
fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug.log
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
2024-10-23 04:13:25,982 INFO MainThread:3335224 [wandb_setup.py:_flush():79] Current SDK version is 0.18.3
|
| 2 |
+
2024-10-23 04:13:25,982 INFO MainThread:3335224 [wandb_setup.py:_flush():79] Configure stats pid to 3335224
|
| 3 |
+
2024-10-23 04:13:25,982 INFO MainThread:3335224 [wandb_setup.py:_flush():79] Loading settings from /admin/home-ckadirt/.config/wandb/settings
|
| 4 |
+
2024-10-23 04:13:25,982 INFO MainThread:3335224 [wandb_setup.py:_flush():79] Loading settings from /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/settings
|
| 5 |
+
2024-10-23 04:13:25,982 INFO MainThread:3335224 [wandb_setup.py:_flush():79] Loading settings from environment variables: {}
|
| 6 |
+
2024-10-23 04:13:25,982 INFO MainThread:3335224 [wandb_setup.py:_flush():79] Applying setup settings: {'mode': None, '_disable_service': None}
|
| 7 |
+
2024-10-23 04:13:25,982 INFO MainThread:3335224 [wandb_setup.py:_flush():79] Inferring run settings from compute environment: {'program_relpath': 'src/HCP_downstream_finetune.py', 'program_abspath': '/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py', 'program': '/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py'}
|
| 8 |
+
2024-10-23 04:13:25,982 INFO MainThread:3335224 [wandb_setup.py:_flush():79] Applying login settings: {}
|
| 9 |
+
2024-10-23 04:13:25,983 INFO MainThread:3335224 [wandb_init.py:_log_setup():532] Logging user logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug.log
|
| 10 |
+
2024-10-23 04:13:25,984 INFO MainThread:3335224 [wandb_init.py:_log_setup():533] Logging internal logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug-internal.log
|
| 11 |
+
2024-10-23 04:13:25,984 INFO MainThread:3335224 [wandb_init.py:init():617] calling init triggers
|
| 12 |
+
2024-10-23 04:13:25,984 INFO MainThread:3335224 [wandb_init.py:init():624] wandb.init called with sweep_config: {}
|
| 13 |
+
config: {'model_name': 'HCPflat_large_gsrFalse__HCP_FT', 'batch_size': 8, 'learning_rate': 0.0001, 'weight_decay': 1e-05, 'num_epochs': 20, 'seed': 42}
|
| 14 |
+
2024-10-23 04:13:25,984 INFO MainThread:3335224 [wandb_init.py:init():667] starting backend
|
| 15 |
+
2024-10-23 04:13:25,984 INFO MainThread:3335224 [wandb_init.py:init():671] sending inform_init request
|
| 16 |
+
2024-10-23 04:13:25,994 INFO MainThread:3335224 [backend.py:_multiprocessing_setup():104] multiprocessing start_methods=fork,spawn,forkserver, using: spawn
|
| 17 |
+
2024-10-23 04:13:25,994 INFO MainThread:3335224 [wandb_init.py:init():684] backend started and connected
|
| 18 |
+
2024-10-23 04:13:26,013 INFO MainThread:3335224 [wandb_init.py:init():779] updated telemetry
|
| 19 |
+
2024-10-23 04:13:26,066 INFO MainThread:3335224 [wandb_init.py:init():812] communicating run to backend with 90.0 second timeout
|
| 20 |
+
2024-10-23 04:13:26,593 INFO MainThread:3335224 [wandb_init.py:init():863] starting run threads in backend
|
| 21 |
+
2024-10-23 04:13:27,015 INFO MainThread:3335224 [wandb_run.py:_console_start():2465] atexit reg
|
| 22 |
+
2024-10-23 04:13:27,015 INFO MainThread:3335224 [wandb_run.py:_redirect():2313] redirect: wrap_raw
|
| 23 |
+
2024-10-23 04:13:27,015 INFO MainThread:3335224 [wandb_run.py:_redirect():2378] Wrapping output streams.
|
| 24 |
+
2024-10-23 04:13:27,015 INFO MainThread:3335224 [wandb_run.py:_redirect():2403] Redirects installed.
|
| 25 |
+
2024-10-23 04:13:27,018 INFO MainThread:3335224 [wandb_init.py:init():907] run started, returning control to user process
|
fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/run-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532.wandb
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:c0adae22e19034c2b45367dc668038ee90f898ca7caf78be9bc50b8504115501
|
| 3 |
+
size 3506176
|
fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/code/src/HCP_downstream_finetune.py
ADDED
|
@@ -0,0 +1,597 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
# coding: utf-8
|
| 3 |
+
|
| 4 |
+
# In[1]:
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
# Import packages and setup gpu configuration.
|
| 8 |
+
# This code block shouldnt need to be adjusted!
|
| 9 |
+
import os
|
| 10 |
+
import sys
|
| 11 |
+
import json
|
| 12 |
+
import yaml
|
| 13 |
+
import numpy as np
|
| 14 |
+
import copy
|
| 15 |
+
import math
|
| 16 |
+
import time
|
| 17 |
+
import random
|
| 18 |
+
from tqdm.auto import tqdm
|
| 19 |
+
import webdataset as wds
|
| 20 |
+
import matplotlib.pyplot as plt
|
| 21 |
+
|
| 22 |
+
import torch
|
| 23 |
+
import torch.nn as nn
|
| 24 |
+
from torchvision import transforms
|
| 25 |
+
import utils
|
| 26 |
+
from mae_utils.flat_models import *
|
| 27 |
+
import h5py
|
| 28 |
+
from mae_utils import flat_models
|
| 29 |
+
|
| 30 |
+
# tf32 data type is faster than standard float32
|
| 31 |
+
torch.backends.cuda.matmul.allow_tf32 = True
|
| 32 |
+
# following fixes a Conv3D CUDNN_NOT_SUPPORTED error
|
| 33 |
+
torch.backends.cudnn.benchmark = True
|
| 34 |
+
|
| 35 |
+
# ## MODEL TO LOAD ##
|
| 36 |
+
if utils.is_interactive():
|
| 37 |
+
model_name = "HCPflat_large_gsrFalse_"
|
| 38 |
+
else:
|
| 39 |
+
model_name = sys.argv[1]
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
# outdir = os.path.abspath(f'checkpoints/{model_name}')
|
| 43 |
+
outdir = os.path.abspath(f'checkpoints/{model_name}')
|
| 44 |
+
|
| 45 |
+
print("outdir", outdir)
|
| 46 |
+
# Load previous config.yaml if available
|
| 47 |
+
if os.path.exists(f"{outdir}/config.yaml"):
|
| 48 |
+
config = yaml.load(open(f"{outdir}/config.yaml", 'r'), Loader=yaml.FullLoader)
|
| 49 |
+
print(f"Loaded config.yaml from ckpt folder {outdir}")
|
| 50 |
+
# create global variables from the config
|
| 51 |
+
print("\n__CONFIG__")
|
| 52 |
+
for attribute_name in config.keys():
|
| 53 |
+
print(f"{attribute_name} = {config[attribute_name]}")
|
| 54 |
+
globals()[attribute_name] = config[f'{attribute_name}']
|
| 55 |
+
print("\n")
|
| 56 |
+
|
| 57 |
+
world_size = os.getenv('WORLD_SIZE')
|
| 58 |
+
if world_size is None:
|
| 59 |
+
world_size = 1
|
| 60 |
+
else:
|
| 61 |
+
world_size = int(world_size)
|
| 62 |
+
print(f"WORLD_SIZE={world_size}")
|
| 63 |
+
|
| 64 |
+
if utils.is_interactive():
|
| 65 |
+
# Following allows you to change functions in models.py or utils.py and
|
| 66 |
+
# have this notebook automatically update with your revisions
|
| 67 |
+
get_ipython().run_line_magic('load_ext', 'autoreload')
|
| 68 |
+
get_ipython().run_line_magic('autoreload', '2')
|
| 69 |
+
|
| 70 |
+
batch_size = probe_batch_size
|
| 71 |
+
num_epochs = probe_num_epochs
|
| 72 |
+
|
| 73 |
+
data_type = torch.float32 # change depending on your mixed_precision
|
| 74 |
+
global_batch_size = batch_size * world_size
|
| 75 |
+
|
| 76 |
+
device = torch.device('cuda')
|
| 77 |
+
|
| 78 |
+
hcp_flat_path = "/weka/proj-medarc/shared/HCP-Flat"
|
| 79 |
+
# seed = 42
|
| 80 |
+
# num_frames = 16
|
| 81 |
+
# gsr = False
|
| 82 |
+
# num_workers = 10
|
| 83 |
+
# batch_size = 128
|
| 84 |
+
save_ckpt = True
|
| 85 |
+
wandb_log = True
|
| 86 |
+
print("PID of this process =",os.getpid())
|
| 87 |
+
utils.seed_everything(seed)
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
# In[2]:
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
if os.getenv('global_pool') == "False":
|
| 94 |
+
global_pool = False
|
| 95 |
+
else:
|
| 96 |
+
global_pool = True
|
| 97 |
+
print(f"global_pool = {global_pool}")
|
| 98 |
+
|
| 99 |
+
try:
|
| 100 |
+
gsr
|
| 101 |
+
except:
|
| 102 |
+
gsr = True
|
| 103 |
+
print("set gsr to True")
|
| 104 |
+
print(f"gsr = {gsr}")
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
# In[3]:
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
#### UNCOMMENT THIS TO SAVE THE HCP-FLAT IN HDF5 FORMAT
|
| 111 |
+
|
| 112 |
+
|
| 113 |
+
# from torch.utils.data import default_collate
|
| 114 |
+
# from mae_utils.flat import load_hcp_flat_mask
|
| 115 |
+
# from mae_utils.flat import create_hcp_flat
|
| 116 |
+
# from mae_utils.flat import batch_unmask
|
| 117 |
+
# import mae_utils.visualize as vis
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
# batch_size = 26
|
| 121 |
+
# print(f"changed batch_size to {batch_size}")
|
| 122 |
+
|
| 123 |
+
# ## Test ##
|
| 124 |
+
# datasets_to_include = "HCP"
|
| 125 |
+
# assert "HCP" in datasets_to_include
|
| 126 |
+
# test_dataset = create_hcp_flat(root=hcp_flat_path,
|
| 127 |
+
# clip_mode="event", frames=num_frames, shuffle=False, gsr=gsr, sub_list = 'test')
|
| 128 |
+
# test_dl = wds.WebLoader(
|
| 129 |
+
# test_dataset.batched(batch_size, partial=False, collation_fn=default_collate),
|
| 130 |
+
# batch_size=None,
|
| 131 |
+
# shuffle=False,
|
| 132 |
+
# num_workers=num_workers,
|
| 133 |
+
# pin_memory=True,
|
| 134 |
+
# )
|
| 135 |
+
|
| 136 |
+
# ## Train ##
|
| 137 |
+
# assert "HCP" in datasets_to_include
|
| 138 |
+
# train_dataset = create_hcp_flat(root=hcp_flat_path,
|
| 139 |
+
# clip_mode="event", frames=num_frames, shuffle=False, gsr=gsr, sub_list = 'train')
|
| 140 |
+
# train_dl = wds.WebLoader(
|
| 141 |
+
# train_dataset.batched(batch_size, partial=False, collation_fn=default_collate),
|
| 142 |
+
# batch_size=None,
|
| 143 |
+
# shuffle=False,
|
| 144 |
+
# num_workers=num_workers,
|
| 145 |
+
# pin_memory=True,
|
| 146 |
+
# )
|
| 147 |
+
|
| 148 |
+
# def flatten_meta(meta_dict):
|
| 149 |
+
# """
|
| 150 |
+
# Flatten the meta dictionary by:
|
| 151 |
+
# - Replacing single-item lists with the item itself.
|
| 152 |
+
# - Converting tensors to scalar numbers.
|
| 153 |
+
# """
|
| 154 |
+
# flattened = {}
|
| 155 |
+
# for key, value in meta_dict.items():
|
| 156 |
+
# if isinstance(value, list):
|
| 157 |
+
# if len(value) == 1:
|
| 158 |
+
# flattened[key] = value[0] # Replace list with its single item
|
| 159 |
+
# else:
|
| 160 |
+
# flattened[key] = value # Keep as is if multiple items
|
| 161 |
+
# elif isinstance(value, torch.Tensor):
|
| 162 |
+
# # Convert tensor to scalar
|
| 163 |
+
# if value.numel() == 1:
|
| 164 |
+
# flattened[key] = value.item()
|
| 165 |
+
# else:
|
| 166 |
+
# flattened[key] = value.tolist() # Convert multi-element tensor to list
|
| 167 |
+
# else:
|
| 168 |
+
# flattened[key] = value # Keep the value as is
|
| 169 |
+
# return flattened
|
| 170 |
+
|
| 171 |
+
# import h5py
|
| 172 |
+
# meta_array = np.array([], dtype=object)
|
| 173 |
+
# # Open an HDF5 file in write mode
|
| 174 |
+
# with h5py.File('train_hcp.hdf5', 'w') as h5f:
|
| 175 |
+
# flatmaps_dset = None
|
| 176 |
+
|
| 177 |
+
# total_samples = 0
|
| 178 |
+
|
| 179 |
+
# for i, batch in tqdm(enumerate(train_dl), total = 120000):
|
| 180 |
+
# images = batch['image'][0]
|
| 181 |
+
# meta = batch['meta']
|
| 182 |
+
# batch_size = images.shape[0]
|
| 183 |
+
# meta_serializable = meta.copy()
|
| 184 |
+
|
| 185 |
+
|
| 186 |
+
# # Step 2: Serialize the dictionary to a JSON string
|
| 187 |
+
# meta_str = json.dumps(flatten_meta(meta_serializable), indent=4)
|
| 188 |
+
# meta_array = np.append(meta_array, meta_str)
|
| 189 |
+
# if flatmaps_dset is None:
|
| 190 |
+
# # Initialize datasets with unlimited (None) maxshape along the first axis
|
| 191 |
+
# flatmaps_shape = (0,) + images.shape[1:]
|
| 192 |
+
# flatmaps_maxshape = (None,) + images.shape[1:]
|
| 193 |
+
|
| 194 |
+
# flatmaps_dset = h5f.create_dataset(
|
| 195 |
+
# 'flatmaps',
|
| 196 |
+
# shape=flatmaps_shape,
|
| 197 |
+
# maxshape=flatmaps_maxshape,
|
| 198 |
+
# dtype=np.float16,
|
| 199 |
+
# chunks=True # Enable chunking for efficient resizing
|
| 200 |
+
# )
|
| 201 |
+
|
| 202 |
+
# # Resize datasets to accommodate new data
|
| 203 |
+
# flatmaps_dset.resize(total_samples + batch_size, axis=0)
|
| 204 |
+
|
| 205 |
+
# # Write data to the datasets
|
| 206 |
+
# flatmaps_dset[total_samples:total_samples + batch_size] = images.numpy().astype(np.float16)
|
| 207 |
+
|
| 208 |
+
# total_samples += batch_size
|
| 209 |
+
|
| 210 |
+
# print(f"Processed {total_samples} samples")
|
| 211 |
+
# np.save('metadata_test_HCP.npy', meta_array)
|
| 212 |
+
|
| 213 |
+
|
| 214 |
+
# import h5py
|
| 215 |
+
# meta_array = np.array([], dtype=object)
|
| 216 |
+
# # Open an HDF5 file in write mode
|
| 217 |
+
# with h5py.File('test_hcp.hdf5', 'w') as h5f:
|
| 218 |
+
# flatmaps_dset = None
|
| 219 |
+
|
| 220 |
+
# total_samples = 0
|
| 221 |
+
|
| 222 |
+
# for i, batch in tqdm(enumerate(test_dl), total = 12000):
|
| 223 |
+
# images = batch['image'][0]
|
| 224 |
+
# meta = batch['meta']
|
| 225 |
+
# batch_size = images.shape[0]
|
| 226 |
+
# meta_serializable = meta.copy()
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
# # Step 2: Serialize the dictionary to a JSON string
|
| 230 |
+
# meta_str = json.dumps(flatten_meta(meta_serializable), indent=4)
|
| 231 |
+
# meta_array = np.append(meta_array, meta_str)
|
| 232 |
+
# if flatmaps_dset is None:
|
| 233 |
+
# # Initialize datasets with unlimited (None) maxshape along the first axis
|
| 234 |
+
# flatmaps_shape = (0,) + images.shape[1:]
|
| 235 |
+
# flatmaps_maxshape = (None,) + images.shape[1:]
|
| 236 |
+
|
| 237 |
+
# flatmaps_dset = h5f.create_dataset(
|
| 238 |
+
# 'flatmaps',
|
| 239 |
+
# shape=flatmaps_shape,
|
| 240 |
+
# maxshape=flatmaps_maxshape,
|
| 241 |
+
# dtype=np.float16,
|
| 242 |
+
# chunks=True # Enable chunking for efficient resizing
|
| 243 |
+
# )
|
| 244 |
+
|
| 245 |
+
# # Resize datasets to accommodate new data
|
| 246 |
+
# flatmaps_dset.resize(total_samples + batch_size, axis=0)
|
| 247 |
+
|
| 248 |
+
# # Write data to the datasets
|
| 249 |
+
# flatmaps_dset[total_samples:total_samples + batch_size] = images.numpy().astype(np.float16)
|
| 250 |
+
|
| 251 |
+
# total_samples += batch_size
|
| 252 |
+
|
| 253 |
+
# print(f"Processed {total_samples} samples")
|
| 254 |
+
# np.save('metadata_train_HCP.npy', meta_array)
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
# ### Preparing data
|
| 258 |
+
|
| 259 |
+
# In[4]:
|
| 260 |
+
|
| 261 |
+
|
| 262 |
+
from sklearn.preprocessing import LabelEncoder
|
| 263 |
+
|
| 264 |
+
INCLUDE_CONDS = {
|
| 265 |
+
"fear",
|
| 266 |
+
"neut",
|
| 267 |
+
"math",
|
| 268 |
+
"story",
|
| 269 |
+
"lf",
|
| 270 |
+
"lh",
|
| 271 |
+
"rf",
|
| 272 |
+
"rh",
|
| 273 |
+
"t",
|
| 274 |
+
"match",
|
| 275 |
+
"relation",
|
| 276 |
+
"mental",
|
| 277 |
+
"rnd",
|
| 278 |
+
"0bk_body",
|
| 279 |
+
"2bk_body",
|
| 280 |
+
"0bk_faces",
|
| 281 |
+
"2bk_faces",
|
| 282 |
+
"0bk_places",
|
| 283 |
+
"2bk_places",
|
| 284 |
+
"0bk_tools",
|
| 285 |
+
"2bk_tools",
|
| 286 |
+
}
|
| 287 |
+
|
| 288 |
+
# test_data = []
|
| 289 |
+
|
| 290 |
+
# # Iterate over the DataLoader with a progress bar
|
| 291 |
+
# for sample in tqdm(train_dl, desc="Processing samples"):
|
| 292 |
+
# x = sample['image']
|
| 293 |
+
# y = sample['meta']['trial_type']
|
| 294 |
+
# key = sample['meta']['key']
|
| 295 |
+
# print(x.shape, y, key)
|
| 296 |
+
# break
|
| 297 |
+
# Initialize the label encoder
|
| 298 |
+
label_encoder = LabelEncoder()
|
| 299 |
+
label_encoder.fit(sorted(INCLUDE_CONDS)) # Ensure consistent ordering
|
| 300 |
+
|
| 301 |
+
num_classes = len(label_encoder.classes_)
|
| 302 |
+
print(f"Number of classes: {num_classes}")
|
| 303 |
+
|
| 304 |
+
|
| 305 |
+
# In[5]:
|
| 306 |
+
|
| 307 |
+
|
| 308 |
+
f_train = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/train_hcp.hdf5', 'r')
|
| 309 |
+
flatmaps_train = f_train['flatmaps']
|
| 310 |
+
|
| 311 |
+
f_test = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/test_hcp.hdf5', 'r')
|
| 312 |
+
flatmaps_test = f_test['flatmaps']
|
| 313 |
+
|
| 314 |
+
metadata_train = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_train_HCP.npy', allow_pickle=True)
|
| 315 |
+
metadata_test = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_test_HCP.npy', allow_pickle=True)
|
| 316 |
+
|
| 317 |
+
|
| 318 |
+
# In[6]:
|
| 319 |
+
|
| 320 |
+
|
| 321 |
+
from torch.utils.data import Dataset, DataLoader
|
| 322 |
+
|
| 323 |
+
class HCPFlatDataset(Dataset):
|
| 324 |
+
def __init__(self, flatmaps, metadata):
|
| 325 |
+
self.flatmaps = flatmaps
|
| 326 |
+
self.metadata = metadata
|
| 327 |
+
|
| 328 |
+
def __len__(self):
|
| 329 |
+
return len(self.metadata)
|
| 330 |
+
|
| 331 |
+
def __getitem__(self, idx):
|
| 332 |
+
return self.flatmaps[idx], json.loads(self.metadata[idx])
|
| 333 |
+
print("Moving datasets to ram")
|
| 334 |
+
# Loading to cpu for faster training, this can take several minutes. Remove this [:] if you want to move one at the time.
|
| 335 |
+
train_dataset = HCPFlatDataset(flatmaps_train, metadata_train)
|
| 336 |
+
train_dl = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)
|
| 337 |
+
|
| 338 |
+
test_dataset = HCPFlatDataset(flatmaps_test, metadata_test)
|
| 339 |
+
test_dl = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
|
| 340 |
+
print("Datasets ready")
|
| 341 |
+
|
| 342 |
+
|
| 343 |
+
# ### Creating and loading Model
|
| 344 |
+
|
| 345 |
+
# In[7]:
|
| 346 |
+
|
| 347 |
+
|
| 348 |
+
from mae_utils.flat import load_hcp_flat_mask
|
| 349 |
+
from mae_utils.flat import create_hcp_flat
|
| 350 |
+
from mae_utils.flat import batch_unmask
|
| 351 |
+
import mae_utils.visualize as vis
|
| 352 |
+
|
| 353 |
+
flat_mask = load_hcp_flat_mask(hcp_flat_path)
|
| 354 |
+
|
| 355 |
+
mae_model = flat_models.mae_vit_large_fmri(
|
| 356 |
+
patch_size=patch_size,
|
| 357 |
+
decoder_embed_dim=decoder_embed_dim,
|
| 358 |
+
t_patch_size=t_patch_size,
|
| 359 |
+
pred_t_dim=pred_t_dim,
|
| 360 |
+
decoder_depth=4,
|
| 361 |
+
cls_embed=cls_embed,
|
| 362 |
+
norm_pix_loss=norm_pix_loss,
|
| 363 |
+
no_qkv_bias=no_qkv_bias,
|
| 364 |
+
sep_pos_embed=sep_pos_embed,
|
| 365 |
+
trunc_init=trunc_init,
|
| 366 |
+
pct_masks_to_decode=pct_masks_to_decode,
|
| 367 |
+
img_mask=flat_mask,
|
| 368 |
+
)
|
| 369 |
+
|
| 370 |
+
|
| 371 |
+
# In[8]:
|
| 372 |
+
|
| 373 |
+
|
| 374 |
+
checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')]
|
| 375 |
+
|
| 376 |
+
if utils.is_interactive():
|
| 377 |
+
latest_checkpoint = "epoch99.pth"
|
| 378 |
+
else:
|
| 379 |
+
latest_checkpoint = sys.argv[2]
|
| 380 |
+
print(f"latest_checkpoint: {latest_checkpoint}")
|
| 381 |
+
|
| 382 |
+
# Load the checkpoint
|
| 383 |
+
checkpoint_path = os.path.join(outdir, latest_checkpoint)
|
| 384 |
+
|
| 385 |
+
state = torch.load(checkpoint_path)
|
| 386 |
+
mae_model.load_state_dict(state["model_state_dict"], strict=False)
|
| 387 |
+
mae_model.to(device)
|
| 388 |
+
|
| 389 |
+
print(f"\nLoaded checkpoint {latest_checkpoint} from {outdir}\n")
|
| 390 |
+
|
| 391 |
+
|
| 392 |
+
# In[9]:
|
| 393 |
+
|
| 394 |
+
|
| 395 |
+
class LinearClassifier(nn.Module):
|
| 396 |
+
def __init__(self, input_dim, num_classes):
|
| 397 |
+
super(LinearClassifier, self).__init__()
|
| 398 |
+
self.linear = nn.Linear(input_dim, num_classes)
|
| 399 |
+
|
| 400 |
+
def forward(self, x):
|
| 401 |
+
# Flatten the input except for the batch dimension
|
| 402 |
+
x = x.view(x.size(0), -1)
|
| 403 |
+
out = self.linear(x)
|
| 404 |
+
return out # Raw logits
|
| 405 |
+
|
| 406 |
+
# Determine the input dimension from a single sample
|
| 407 |
+
# Assuming images are of shape [1, 16, 144, 320]
|
| 408 |
+
input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])
|
| 409 |
+
print(f"Input dimension: {input_dim}")
|
| 410 |
+
|
| 411 |
+
|
| 412 |
+
# In[10]:
|
| 413 |
+
|
| 414 |
+
|
| 415 |
+
class FullModel(nn.Module):
|
| 416 |
+
def __init__(self, lc_model, mae_model):
|
| 417 |
+
super(FullModel, self).__init__()
|
| 418 |
+
self.lc_model = lc_model
|
| 419 |
+
self.mae_model = mae_model
|
| 420 |
+
|
| 421 |
+
|
| 422 |
+
def forward(self, x, gsr):
|
| 423 |
+
x = self.mae_model(x, global_pool=global_pool, forward_features = True)
|
| 424 |
+
x = self.lc_model(x)
|
| 425 |
+
return x
|
| 426 |
+
|
| 427 |
+
|
| 428 |
+
# In[11]:
|
| 429 |
+
|
| 430 |
+
|
| 431 |
+
# Initialize the model
|
| 432 |
+
lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)
|
| 433 |
+
|
| 434 |
+
model = FullModel(lc_model, mae_model)
|
| 435 |
+
|
| 436 |
+
# Move the model to the GPU
|
| 437 |
+
model.to(device)
|
| 438 |
+
|
| 439 |
+
# Define loss function
|
| 440 |
+
criterion = nn.CrossEntropyLoss()
|
| 441 |
+
|
| 442 |
+
# Define optimizer with L2 regularization (weight_decay)
|
| 443 |
+
learning_rate = 1e-4
|
| 444 |
+
weight_decay = 1e-5 # Adjust based on your needs
|
| 445 |
+
optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)
|
| 446 |
+
num_epochs = 20 # Adjust as needed
|
| 447 |
+
|
| 448 |
+
|
| 449 |
+
# ### Data
|
| 450 |
+
|
| 451 |
+
# In[16]:
|
| 452 |
+
|
| 453 |
+
|
| 454 |
+
import uuid
|
| 455 |
+
|
| 456 |
+
myuuid = uuid.uuid4()
|
| 457 |
+
str(myuuid)
|
| 458 |
+
|
| 459 |
+
|
| 460 |
+
# In[17]:
|
| 461 |
+
|
| 462 |
+
|
| 463 |
+
import wandb
|
| 464 |
+
|
| 465 |
+
if utils.is_interactive():
|
| 466 |
+
print("Running in interactive notebook. Disabling W&B and ckpt saving.")
|
| 467 |
+
wandb_log = True
|
| 468 |
+
save_ckpt = True
|
| 469 |
+
|
| 470 |
+
if wandb_log:
|
| 471 |
+
wandb_project = 'fMRI-foundation-model'
|
| 472 |
+
wandb_config = {
|
| 473 |
+
"model_name": model_name+'_HCP_FT',
|
| 474 |
+
"batch_size": batch_size,
|
| 475 |
+
"learning_rate": learning_rate,
|
| 476 |
+
"weight_decay": weight_decay,
|
| 477 |
+
"num_epochs": num_epochs,
|
| 478 |
+
"seed": seed,
|
| 479 |
+
}
|
| 480 |
+
print("wandb_config:\n", wandb_config)
|
| 481 |
+
random_id = str(uuid.uuid4())
|
| 482 |
+
print("wandb_id:", "HCPflat_raw" + f"_{random_id}")
|
| 483 |
+
wandb.init(
|
| 484 |
+
id=model_name+'_HCP_FT' + f"_{random_id}",
|
| 485 |
+
project=wandb_project,
|
| 486 |
+
name=model_name+'_HCP_FT',
|
| 487 |
+
config=wandb_config,
|
| 488 |
+
resume="allow",
|
| 489 |
+
)
|
| 490 |
+
|
| 491 |
+
|
| 492 |
+
# In[13]:
|
| 493 |
+
|
| 494 |
+
|
| 495 |
+
for epoch in range(num_epochs):
|
| 496 |
+
running_train_loss = 0.0
|
| 497 |
+
correct_train = 0
|
| 498 |
+
total_train = 0
|
| 499 |
+
step = 0
|
| 500 |
+
|
| 501 |
+
# with torch.amp.autocast(device_type='cuda'):
|
| 502 |
+
# Training Phase
|
| 503 |
+
model.train()
|
| 504 |
+
for batch in tqdm(train_dl, desc=f"Epoch {epoch+1}/{num_epochs} - Training"):
|
| 505 |
+
optimizer.zero_grad()
|
| 506 |
+
images = batch[0].to(device).float().unsqueeze(1) #fix this # Shape: [batch_size, 1, 16, 144, 320]
|
| 507 |
+
labels = batch[1]['trial_type'] # List of labels
|
| 508 |
+
|
| 509 |
+
encoded_labels = label_encoder.transform(labels)
|
| 510 |
+
encoded_labels = torch.tensor(encoded_labels, dtype=torch.long).to(device) # Shape: [batch_size]
|
| 511 |
+
|
| 512 |
+
# Forward pass
|
| 513 |
+
outputs = model(images, gsr=gsr) # Shape: [num_train_samples, num_classes]
|
| 514 |
+
|
| 515 |
+
# Compute loss
|
| 516 |
+
loss = criterion(outputs, encoded_labels)
|
| 517 |
+
|
| 518 |
+
# Backward pass and optimization
|
| 519 |
+
loss.backward()
|
| 520 |
+
optimizer.step()
|
| 521 |
+
|
| 522 |
+
# Accumulate loss
|
| 523 |
+
running_train_loss += loss.item() * images.size(0)
|
| 524 |
+
|
| 525 |
+
|
| 526 |
+
# Calculate accuracy
|
| 527 |
+
_, predicted = torch.max(outputs, 1)
|
| 528 |
+
|
| 529 |
+
correct_train += (predicted == encoded_labels).sum().item()
|
| 530 |
+
total_train += encoded_labels.size(0)
|
| 531 |
+
|
| 532 |
+
step = step + 1
|
| 533 |
+
if step % 100 == 0:
|
| 534 |
+
print(f"Step [{step}/{len(train_dl)}] - Training Loss: {loss.item():.4f} - Training Accuracy: {100 * correct_train / total_train:.2f}%")
|
| 535 |
+
# thth
|
| 536 |
+
|
| 537 |
+
epoch_train_loss = running_train_loss / total_train if total_train > 0 else 0.0
|
| 538 |
+
train_accuracy = 100 * correct_train / total_train if total_train > 0 else 0.0
|
| 539 |
+
|
| 540 |
+
# Validation Phase
|
| 541 |
+
model.eval()
|
| 542 |
+
running_val_loss = 0.0
|
| 543 |
+
correct_val = 0
|
| 544 |
+
total_val = 0
|
| 545 |
+
|
| 546 |
+
with torch.no_grad():
|
| 547 |
+
for batch in tqdm(test_dl, desc=f"Epoch {epoch+1}/{num_epochs} - Validation"):
|
| 548 |
+
|
| 549 |
+
images = batch[0].to(device).float().unsqueeze(1) #fix this
|
| 550 |
+
labels = batch[1]['trial_type']
|
| 551 |
+
|
| 552 |
+
# Encode labels to integer indices
|
| 553 |
+
encoded_labels = label_encoder.transform(labels)
|
| 554 |
+
encoded_labels = torch.tensor(encoded_labels, dtype=torch.long).to(device)
|
| 555 |
+
|
| 556 |
+
|
| 557 |
+
# Forward pass
|
| 558 |
+
outputs = model(images, gsr=gsr)
|
| 559 |
+
|
| 560 |
+
# Compute loss
|
| 561 |
+
loss = criterion(outputs, encoded_labels)
|
| 562 |
+
|
| 563 |
+
# Accumulate loss
|
| 564 |
+
running_val_loss += loss.item() * images.size(0)
|
| 565 |
+
|
| 566 |
+
# Calculate accuracy
|
| 567 |
+
_, predicted = torch.max(outputs, 1)
|
| 568 |
+
correct_val += (predicted == encoded_labels).sum().item()
|
| 569 |
+
total_val += encoded_labels.size(0)
|
| 570 |
+
|
| 571 |
+
|
| 572 |
+
|
| 573 |
+
epoch_val_loss = running_val_loss / total_val if total_val > 0 else 0.0
|
| 574 |
+
val_accuracy = 100 * correct_val / total_val if total_val > 0 else 0.0
|
| 575 |
+
|
| 576 |
+
print(f"Epoch [{epoch+1}/{num_epochs}] "
|
| 577 |
+
f"- Training Loss: {epoch_train_loss:.4f}, Training Accuracy: {train_accuracy:.2f}% "
|
| 578 |
+
f"- Validation Loss: {epoch_val_loss:.4f}, Validation Accuracy: {val_accuracy:.2f}%")
|
| 579 |
+
|
| 580 |
+
if wandb_log:
|
| 581 |
+
wandb.log({
|
| 582 |
+
"epoch_train_loss": epoch_train_loss,
|
| 583 |
+
"epoch_val_loss": epoch_val_loss,
|
| 584 |
+
"train_accuracy": train_accuracy,
|
| 585 |
+
"val_accuracy": val_accuracy,
|
| 586 |
+
})
|
| 587 |
+
if save_ckpt:
|
| 588 |
+
outdir = os.path.abspath(f'checkpoints/{model_name+"HCP_FT"}')
|
| 589 |
+
os.makedirs(outdir, exist_ok=True)
|
| 590 |
+
print("outdir", outdir)
|
| 591 |
+
# Save model and config
|
| 592 |
+
torch.save(model.state_dict(), f"{outdir}/model.pth")
|
| 593 |
+
with open(f"{outdir}/config.yaml", 'w') as f:
|
| 594 |
+
yaml.dump(wandb_config, f)
|
| 595 |
+
print(f"Saved model and config to {outdir}")
|
| 596 |
+
|
| 597 |
+
|
fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/output.log
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Epoch 1/20 - Training: 16%|█▌ | 2209/13913 [13:24<1:07:17, 2.90it/s]
|
| 2 |
+
Step [100/13913] - Training Loss: 1.9649 - Training Accuracy: 59.38%
|
| 3 |
+
Step [200/13913] - Training Loss: 1.0904 - Training Accuracy: 71.19%
|
| 4 |
+
Step [300/13913] - Training Loss: 0.0775 - Training Accuracy: 75.50%
|
| 5 |
+
Step [400/13913] - Training Loss: 0.6052 - Training Accuracy: 78.84%
|
| 6 |
+
Step [500/13913] - Training Loss: 0.0226 - Training Accuracy: 80.95%
|
| 7 |
+
Step [600/13913] - Training Loss: 0.2728 - Training Accuracy: 82.54%
|
| 8 |
+
Step [700/13913] - Training Loss: 0.1662 - Training Accuracy: 83.70%
|
| 9 |
+
Step [800/13913] - Training Loss: 0.0385 - Training Accuracy: 84.89%
|
| 10 |
+
Step [900/13913] - Training Loss: 0.2377 - Training Accuracy: 85.61%
|
| 11 |
+
Step [1000/13913] - Training Loss: 0.8172 - Training Accuracy: 85.83%
|
| 12 |
+
Step [1100/13913] - Training Loss: 0.2276 - Training Accuracy: 86.68%
|
| 13 |
+
Step [1200/13913] - Training Loss: 0.0118 - Training Accuracy: 87.36%
|
| 14 |
+
Step [1300/13913] - Training Loss: 1.0419 - Training Accuracy: 87.68%
|
| 15 |
+
Step [1400/13913] - Training Loss: 0.8943 - Training Accuracy: 87.96%
|
| 16 |
+
Step [1500/13913] - Training Loss: 0.2801 - Training Accuracy: 88.23%
|
| 17 |
+
Step [1600/13913] - Training Loss: 0.6734 - Training Accuracy: 88.58%
|
| 18 |
+
Step [1700/13913] - Training Loss: 0.6202 - Training Accuracy: 88.88%
|
| 19 |
+
Step [1800/13913] - Training Loss: 0.0159 - Training Accuracy: 89.12%
|
| 20 |
+
Step [1900/13913] - Training Loss: 0.0682 - Training Accuracy: 89.37%
|
| 21 |
+
Step [2000/13913] - Training Loss: 0.3378 - Training Accuracy: 89.52%
|
| 22 |
+
Step [2100/13913] - Training Loss: 0.0509 - Training Accuracy: 89.77%
|
| 23 |
+
Step [2200/13913] - Training Loss: 0.1161 - Training Accuracy: 89.99%
|
fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/requirements.txt
ADDED
|
@@ -0,0 +1,198 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
protobuf==5.28.2
|
| 2 |
+
imageio==2.35.1
|
| 3 |
+
MarkupSafe==3.0.0
|
| 4 |
+
regex==2024.9.11
|
| 5 |
+
matplotlib==3.9.2
|
| 6 |
+
notebook==7.2.2
|
| 7 |
+
debugpy==1.8.6
|
| 8 |
+
aiosignal==1.3.1
|
| 9 |
+
jupyter_core==5.7.2
|
| 10 |
+
torchaudio==2.4.1+cu121
|
| 11 |
+
python-json-logger==2.0.7
|
| 12 |
+
six==1.16.0
|
| 13 |
+
scikit-image==0.24.0
|
| 14 |
+
types-python-dateutil==2.9.0.20241003
|
| 15 |
+
PyYAML==6.0.2
|
| 16 |
+
httpcore==1.0.6
|
| 17 |
+
clip==1.0
|
| 18 |
+
babel==2.16.0
|
| 19 |
+
webcolors==24.8.0
|
| 20 |
+
omegaconf==2.3.0
|
| 21 |
+
webencodings==0.5.1
|
| 22 |
+
kiwisolver==1.4.7
|
| 23 |
+
uri-template==1.3.0
|
| 24 |
+
diffusers==0.23.0
|
| 25 |
+
idna==3.10
|
| 26 |
+
fsspec==2024.9.0
|
| 27 |
+
parso==0.8.4
|
| 28 |
+
setuptools==65.5.0
|
| 29 |
+
tornado==6.4.1
|
| 30 |
+
webdataset==0.2.100
|
| 31 |
+
decord==0.6.0
|
| 32 |
+
nvidia-curand-cu12==10.3.2.106
|
| 33 |
+
ipykernel==6.29.5
|
| 34 |
+
jupyter==1.1.1
|
| 35 |
+
pexpect==4.9.0
|
| 36 |
+
kornia_rs==0.1.5
|
| 37 |
+
iopath==0.1.10
|
| 38 |
+
async-lru==2.0.4
|
| 39 |
+
future==1.0.0
|
| 40 |
+
torchvision==0.19.1+cu121
|
| 41 |
+
botocore==1.34.162
|
| 42 |
+
cycler==0.12.1
|
| 43 |
+
tzdata==2024.2
|
| 44 |
+
jupyter_server_terminals==0.5.3
|
| 45 |
+
click==8.1.7
|
| 46 |
+
einops==0.8.0
|
| 47 |
+
pyzmq==26.2.0
|
| 48 |
+
jupyter_client==8.6.3
|
| 49 |
+
nbconvert==7.16.4
|
| 50 |
+
scikit-learn==1.5.2
|
| 51 |
+
executing==2.1.0
|
| 52 |
+
asttokens==2.4.1
|
| 53 |
+
docker-pycreds==0.4.0
|
| 54 |
+
matplotlib-inline==0.1.7
|
| 55 |
+
overrides==7.7.0
|
| 56 |
+
websocket-client==1.8.0
|
| 57 |
+
nbformat==5.10.4
|
| 58 |
+
elbow==0.1.1
|
| 59 |
+
contourpy==1.3.0
|
| 60 |
+
nvidia-cudnn-cu12==9.1.0.70
|
| 61 |
+
transformers==4.44.2
|
| 62 |
+
gitdb==4.0.11
|
| 63 |
+
jupyterlab_nvdashboard==0.11.0
|
| 64 |
+
lazy_loader==0.4
|
| 65 |
+
jsonpointer==3.0.0
|
| 66 |
+
notebook_shim==0.2.4
|
| 67 |
+
nvidia-nccl-cu12==2.20.5
|
| 68 |
+
ffmpeg-python==0.2.0
|
| 69 |
+
triton==3.0.0
|
| 70 |
+
mistune==3.0.2
|
| 71 |
+
python-dateutil==2.9.0.post0
|
| 72 |
+
beautifulsoup4==4.12.3
|
| 73 |
+
nbclient==0.10.0
|
| 74 |
+
h5py==3.12.1
|
| 75 |
+
ftfy==6.2.3
|
| 76 |
+
zipp==3.20.2
|
| 77 |
+
ptyprocess==0.7.0
|
| 78 |
+
huggingface-hub==0.25.1
|
| 79 |
+
pytz==2024.2
|
| 80 |
+
jupyterlab_pygments==0.3.0
|
| 81 |
+
nvidia-cublas-cu12==12.1.3.1
|
| 82 |
+
pandocfilters==1.5.1
|
| 83 |
+
Jinja2==3.1.4
|
| 84 |
+
arrow==1.3.0
|
| 85 |
+
rpds-py==0.20.0
|
| 86 |
+
jupyter_server==2.14.2
|
| 87 |
+
simplejson==3.19.3
|
| 88 |
+
networkx==3.3
|
| 89 |
+
packaging==24.1
|
| 90 |
+
traitlets==5.14.3
|
| 91 |
+
pandas==2.2.3
|
| 92 |
+
xformers==0.0.22.post7
|
| 93 |
+
lightning-utilities==0.11.7
|
| 94 |
+
tifffile==2024.9.20
|
| 95 |
+
nvidia-cuda-cupti-cu12==12.1.105
|
| 96 |
+
mpmath==1.3.0
|
| 97 |
+
GitPython==3.1.43
|
| 98 |
+
scipy==1.14.1
|
| 99 |
+
jsonschema==4.23.0
|
| 100 |
+
prompt_toolkit==3.0.48
|
| 101 |
+
s3transfer==0.10.2
|
| 102 |
+
multidict==6.1.0
|
| 103 |
+
bleach==6.1.0
|
| 104 |
+
sentry-sdk==2.15.0
|
| 105 |
+
nibabel==5.2.1
|
| 106 |
+
accelerate==1.0.0
|
| 107 |
+
pyarrow==17.0.0
|
| 108 |
+
threadpoolctl==3.5.0
|
| 109 |
+
attrs==24.2.0
|
| 110 |
+
rfc3986-validator==0.1.1
|
| 111 |
+
nvidia-cuda-runtime-cu12==12.1.105
|
| 112 |
+
ipywidgets==8.1.5
|
| 113 |
+
frozenlist==1.4.1
|
| 114 |
+
pycparser==2.22
|
| 115 |
+
jupyterlab_server==2.27.3
|
| 116 |
+
nvidia-cuda-nvrtc-cu12==12.1.105
|
| 117 |
+
yarl==1.13.1
|
| 118 |
+
setproctitle==1.3.3
|
| 119 |
+
isoduration==20.11.0
|
| 120 |
+
Pygments==2.18.0
|
| 121 |
+
jedi==0.19.1
|
| 122 |
+
boto3==1.34.57
|
| 123 |
+
tokenizers==0.19.1
|
| 124 |
+
referencing==0.35.1
|
| 125 |
+
rfc3339-validator==0.1.4
|
| 126 |
+
pillow==10.4.0
|
| 127 |
+
jupyterlab==4.2.5
|
| 128 |
+
stack-data==0.6.3
|
| 129 |
+
h11==0.14.0
|
| 130 |
+
anyio==4.6.0
|
| 131 |
+
nilearn==0.10.4
|
| 132 |
+
nvidia-cusolver-cu12==11.4.5.107
|
| 133 |
+
tinycss2==1.3.0
|
| 134 |
+
defusedxml==0.7.1
|
| 135 |
+
argon2-cffi-bindings==21.2.0
|
| 136 |
+
soupsieve==2.6
|
| 137 |
+
nest-asyncio==1.6.0
|
| 138 |
+
torchmetrics==1.3.0.post0
|
| 139 |
+
tqdm==4.66.5
|
| 140 |
+
cffi==1.17.1
|
| 141 |
+
charset-normalizer==3.3.2
|
| 142 |
+
jsonschema-specifications==2023.12.1
|
| 143 |
+
decorator==5.1.1
|
| 144 |
+
open_clip_torch==2.26.1
|
| 145 |
+
jupyter-events==0.10.0
|
| 146 |
+
smart-open==7.0.5
|
| 147 |
+
antlr4-python3-runtime==4.9.3
|
| 148 |
+
prometheus_client==0.21.0
|
| 149 |
+
kornia==0.7.3
|
| 150 |
+
typing_extensions==4.12.2
|
| 151 |
+
sniffio==1.3.1
|
| 152 |
+
joblib==1.4.2
|
| 153 |
+
comm==0.2.2
|
| 154 |
+
aiohappyeyeballs==2.4.3
|
| 155 |
+
numpy==2.1.2
|
| 156 |
+
braceexpand==0.1.7
|
| 157 |
+
certifi==2024.8.30
|
| 158 |
+
psutil==6.0.0
|
| 159 |
+
pyparsing==3.1.4
|
| 160 |
+
pure_eval==0.2.3
|
| 161 |
+
nvidia-cusparse-cu12==12.1.0.106
|
| 162 |
+
wandb==0.18.3
|
| 163 |
+
urllib3==2.2.3
|
| 164 |
+
smmap==5.0.1
|
| 165 |
+
platformdirs==4.3.6
|
| 166 |
+
torch==2.4.1+cu121
|
| 167 |
+
requests==2.32.3
|
| 168 |
+
json5==0.9.25
|
| 169 |
+
nvidia-nvjitlink-cu12==12.6.77
|
| 170 |
+
jupyterlab_widgets==3.0.13
|
| 171 |
+
lxml==5.3.0
|
| 172 |
+
httpx==0.27.2
|
| 173 |
+
opencv-python==4.6.0.66
|
| 174 |
+
portalocker==2.10.1
|
| 175 |
+
pytorch-lightning==2.0.1
|
| 176 |
+
sympy==1.13.3
|
| 177 |
+
wcwidth==0.2.13
|
| 178 |
+
jmespath==1.0.1
|
| 179 |
+
fqdn==1.5.1
|
| 180 |
+
pynvml==11.5.3
|
| 181 |
+
pip==24.0
|
| 182 |
+
wrapt==1.16.0
|
| 183 |
+
aiohttp==3.10.9
|
| 184 |
+
filelock==3.16.1
|
| 185 |
+
fonttools==4.54.1
|
| 186 |
+
fastjsonschema==2.20.0
|
| 187 |
+
jupyter-console==6.6.3
|
| 188 |
+
widgetsnbextension==4.0.13
|
| 189 |
+
timm==1.0.9
|
| 190 |
+
nvidia-cufft-cu12==11.0.2.54
|
| 191 |
+
ipython==8.28.0
|
| 192 |
+
nvidia-nvtx-cu12==12.1.105
|
| 193 |
+
jupyter-lsp==2.2.5
|
| 194 |
+
safetensors==0.4.5
|
| 195 |
+
terminado==0.18.1
|
| 196 |
+
argon2-cffi==23.1.0
|
| 197 |
+
Send2Trash==1.8.3
|
| 198 |
+
importlib_metadata==8.5.0
|
fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/wandb-metadata.json
ADDED
|
@@ -0,0 +1,131 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"os": "Linux-5.15.0-1058-aws-x86_64-with-glibc2.31",
|
| 3 |
+
"python": "3.11.10",
|
| 4 |
+
"startedAt": "2024-10-24T02:16:25.309620Z",
|
| 5 |
+
"args": [
|
| 6 |
+
"HCPflat_large_gsrFalse_",
|
| 7 |
+
"epoch99.pth"
|
| 8 |
+
],
|
| 9 |
+
"program": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py",
|
| 10 |
+
"codePath": "src/HCP_downstream_finetune.py",
|
| 11 |
+
"git": {
|
| 12 |
+
"remote": "https://github.com/MedARC-AI/fMRI-foundation-model",
|
| 13 |
+
"commit": "cf8214d4ebe437188b68b4ee5a34c5211a810db0"
|
| 14 |
+
},
|
| 15 |
+
"email": "torrico.villanueva.cesar.kadir@gmail.com",
|
| 16 |
+
"root": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src",
|
| 17 |
+
"host": "ip-10-0-181-106",
|
| 18 |
+
"username": "ckadirt",
|
| 19 |
+
"executable": "/admin/home-ckadirt/foundation_env/bin/python",
|
| 20 |
+
"codePathLocal": "HCP_downstream_finetune.py",
|
| 21 |
+
"cpu_count": 96,
|
| 22 |
+
"cpu_count_logical": 192,
|
| 23 |
+
"gpu": "[NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3]",
|
| 24 |
+
"gpu_count": 8,
|
| 25 |
+
"disk": {
|
| 26 |
+
"/": {
|
| 27 |
+
"total": "249555763200",
|
| 28 |
+
"used": "184254554112"
|
| 29 |
+
}
|
| 30 |
+
},
|
| 31 |
+
"memory": {
|
| 32 |
+
"total": "2147443404800"
|
| 33 |
+
},
|
| 34 |
+
"cpu": {
|
| 35 |
+
"count": 96,
|
| 36 |
+
"countLogical": 192
|
| 37 |
+
},
|
| 38 |
+
"gpu_nvidia": [
|
| 39 |
+
{
|
| 40 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 41 |
+
"memoryTotal": "85520809984",
|
| 42 |
+
"cudaCores": 16896,
|
| 43 |
+
"architecture": "Hopper"
|
| 44 |
+
},
|
| 45 |
+
{
|
| 46 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 47 |
+
"memoryTotal": "85520809984",
|
| 48 |
+
"cudaCores": 16896,
|
| 49 |
+
"architecture": "Hopper"
|
| 50 |
+
},
|
| 51 |
+
{
|
| 52 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 53 |
+
"memoryTotal": "85520809984",
|
| 54 |
+
"cudaCores": 16896,
|
| 55 |
+
"architecture": "Hopper"
|
| 56 |
+
},
|
| 57 |
+
{
|
| 58 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 59 |
+
"memoryTotal": "85520809984",
|
| 60 |
+
"cudaCores": 16896,
|
| 61 |
+
"architecture": "Hopper"
|
| 62 |
+
},
|
| 63 |
+
{
|
| 64 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 65 |
+
"memoryTotal": "85520809984",
|
| 66 |
+
"cudaCores": 16896,
|
| 67 |
+
"architecture": "Hopper"
|
| 68 |
+
},
|
| 69 |
+
{
|
| 70 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 71 |
+
"memoryTotal": "85520809984",
|
| 72 |
+
"cudaCores": 16896,
|
| 73 |
+
"architecture": "Hopper"
|
| 74 |
+
},
|
| 75 |
+
{
|
| 76 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 77 |
+
"memoryTotal": "85520809984",
|
| 78 |
+
"cudaCores": 16896,
|
| 79 |
+
"architecture": "Hopper"
|
| 80 |
+
},
|
| 81 |
+
{
|
| 82 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 83 |
+
"memoryTotal": "85520809984",
|
| 84 |
+
"cudaCores": 16896,
|
| 85 |
+
"architecture": "Hopper"
|
| 86 |
+
}
|
| 87 |
+
],
|
| 88 |
+
"slurm": {
|
| 89 |
+
"cluster_name": "sagemaker2",
|
| 90 |
+
"conf": "/opt/slurm/etc/slurm.conf",
|
| 91 |
+
"cpus_on_node": "20",
|
| 92 |
+
"gpus_on_node": "1",
|
| 93 |
+
"gpus_per_task": "1",
|
| 94 |
+
"gtids": "0",
|
| 95 |
+
"job_account": "fmri",
|
| 96 |
+
"job_cpus_per_node": "20",
|
| 97 |
+
"job_end_time": "1729779337",
|
| 98 |
+
"job_gid": "1879800513",
|
| 99 |
+
"job_gpus": "0",
|
| 100 |
+
"job_id": "528653",
|
| 101 |
+
"job_name": "finetuneHCP",
|
| 102 |
+
"job_nodelist": "ip-10-0-181-106",
|
| 103 |
+
"job_num_nodes": "1",
|
| 104 |
+
"job_partition": "p5",
|
| 105 |
+
"job_qos": "normal",
|
| 106 |
+
"job_start_time": "1729736136",
|
| 107 |
+
"job_uid": "1879804696",
|
| 108 |
+
"job_user": "ckadirt",
|
| 109 |
+
"jobid": "528653",
|
| 110 |
+
"localid": "0",
|
| 111 |
+
"mem_per_cpu": "11500",
|
| 112 |
+
"nnodes": "1",
|
| 113 |
+
"node_aliases": "(null)",
|
| 114 |
+
"nodeid": "0",
|
| 115 |
+
"nodelist": "ip-10-0-181-106",
|
| 116 |
+
"nprocs": "1",
|
| 117 |
+
"ntasks": "1",
|
| 118 |
+
"ntasks_per_node": "1",
|
| 119 |
+
"prio_process": "0",
|
| 120 |
+
"procid": "0",
|
| 121 |
+
"script_context": "prolog_task",
|
| 122 |
+
"submit_dir": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src",
|
| 123 |
+
"submit_host": "ip-172-17-12-61",
|
| 124 |
+
"task_pid": "2312183",
|
| 125 |
+
"tasks_per_node": "1",
|
| 126 |
+
"topology_addr": "ip-10-0-181-106",
|
| 127 |
+
"topology_addr_pattern": "node",
|
| 128 |
+
"working_cluster": "sagemaker2:ip-172-17-63-161:6817:9984:109"
|
| 129 |
+
},
|
| 130 |
+
"cudaVersion": "12.2"
|
| 131 |
+
}
|
fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug-core.log
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-24T02:16:24.2815064Z","level":"INFO","msg":"started logging, with flags","port-filename":"/tmp/tmpbnequ20j/port-2312212.txt","pid":2312212,"debug":false,"disable-analytics":false}
|
| 2 |
+
{"time":"2024-10-24T02:16:24.281779825Z","level":"INFO","msg":"FeatureState","shutdownOnParentExitEnabled":false}
|
| 3 |
+
{"time":"2024-10-24T02:16:24.287252623Z","level":"INFO","msg":"Will exit if parent process dies.","ppid":2312212}
|
| 4 |
+
{"time":"2024-10-24T02:16:24.287238992Z","level":"INFO","msg":"server is running","addr":{"IP":"127.0.0.1","Port":43197,"Zone":""}}
|
| 5 |
+
{"time":"2024-10-24T02:16:24.297514377Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"127.0.0.1:48772"}
|
| 6 |
+
{"time":"2024-10-24T02:16:25.309558909Z","level":"INFO","msg":"handleInformInit: received","streamId":"HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427","id":"127.0.0.1:48772"}
|
| 7 |
+
{"time":"2024-10-24T02:16:25.471977517Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427","id":"127.0.0.1:48772"}
|
fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug-internal.log
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-24T02:16:25.343111738Z","level":"INFO","msg":"using version","core version":"0.18.3"}
|
| 2 |
+
{"time":"2024-10-24T02:16:25.343149488Z","level":"INFO","msg":"created symlink","path":"/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug-core.log"}
|
| 3 |
+
{"time":"2024-10-24T02:16:25.354611723Z","level":"ERROR","msg":"dialing: google: could not find default credentials. See https://cloud.google.com/docs/authentication/external/set-up-adc for more information"}
|
| 4 |
+
{"time":"2024-10-24T02:16:25.471914526Z","level":"INFO","msg":"created new stream","id":"HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427"}
|
| 5 |
+
{"time":"2024-10-24T02:16:25.471970027Z","level":"INFO","msg":"stream: started","id":"HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427"}
|
| 6 |
+
{"time":"2024-10-24T02:16:25.472001617Z","level":"INFO","msg":"handler: started","stream_id":{"value":"HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427"}}
|
| 7 |
+
{"time":"2024-10-24T02:16:25.472008128Z","level":"INFO","msg":"writer: Do: started","stream_id":{"value":"HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427"}}
|
| 8 |
+
{"time":"2024-10-24T02:16:25.472045188Z","level":"INFO","msg":"sender: started","stream_id":{"value":"HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427"}}
|
| 9 |
+
{"time":"2024-10-24T02:16:26.180235591Z","level":"INFO","msg":"wandb-core","!BADKEY":null}
|
| 10 |
+
{"time":"2024-10-24T02:16:26.191656334Z","level":"INFO","msg":"Starting system monitor"}
|
| 11 |
+
{"time":"2024-10-24T02:16:26.243500377Z","level":"ERROR","msg":"git repo not found","error":"repository does not exist"}
|
fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug.log
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
2024-10-24 02:16:25,238 INFO MainThread:2312212 [wandb_setup.py:_flush():79] Current SDK version is 0.18.3
|
| 2 |
+
2024-10-24 02:16:25,239 INFO MainThread:2312212 [wandb_setup.py:_flush():79] Configure stats pid to 2312212
|
| 3 |
+
2024-10-24 02:16:25,239 INFO MainThread:2312212 [wandb_setup.py:_flush():79] Loading settings from /admin/home-ckadirt/.config/wandb/settings
|
| 4 |
+
2024-10-24 02:16:25,239 INFO MainThread:2312212 [wandb_setup.py:_flush():79] Loading settings from /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/settings
|
| 5 |
+
2024-10-24 02:16:25,239 INFO MainThread:2312212 [wandb_setup.py:_flush():79] Loading settings from environment variables: {}
|
| 6 |
+
2024-10-24 02:16:25,239 INFO MainThread:2312212 [wandb_setup.py:_flush():79] Applying setup settings: {'mode': None, '_disable_service': None}
|
| 7 |
+
2024-10-24 02:16:25,239 INFO MainThread:2312212 [wandb_setup.py:_flush():79] Inferring run settings from compute environment: {'program_relpath': 'src/HCP_downstream_finetune.py', 'program_abspath': '/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py', 'program': '/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py'}
|
| 8 |
+
2024-10-24 02:16:25,239 INFO MainThread:2312212 [wandb_setup.py:_flush():79] Applying login settings: {}
|
| 9 |
+
2024-10-24 02:16:25,242 INFO MainThread:2312212 [wandb_init.py:_log_setup():532] Logging user logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug.log
|
| 10 |
+
2024-10-24 02:16:25,246 INFO MainThread:2312212 [wandb_init.py:_log_setup():533] Logging internal logs to /weka/proj-fmri/ckadirt/fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug-internal.log
|
| 11 |
+
2024-10-24 02:16:25,246 INFO MainThread:2312212 [wandb_init.py:init():617] calling init triggers
|
| 12 |
+
2024-10-24 02:16:25,246 INFO MainThread:2312212 [wandb_init.py:init():624] wandb.init called with sweep_config: {}
|
| 13 |
+
config: {'model_name': 'HCPflat_large_gsrFalse__HCP_FT', 'batch_size': 8, 'learning_rate': 0.0001, 'weight_decay': 1e-05, 'num_epochs': 20, 'seed': 42}
|
| 14 |
+
2024-10-24 02:16:25,246 INFO MainThread:2312212 [wandb_init.py:init():667] starting backend
|
| 15 |
+
2024-10-24 02:16:25,246 INFO MainThread:2312212 [wandb_init.py:init():671] sending inform_init request
|
| 16 |
+
2024-10-24 02:16:25,305 INFO MainThread:2312212 [backend.py:_multiprocessing_setup():104] multiprocessing start_methods=fork,spawn,forkserver, using: spawn
|
| 17 |
+
2024-10-24 02:16:25,305 INFO MainThread:2312212 [wandb_init.py:init():684] backend started and connected
|
| 18 |
+
2024-10-24 02:16:25,385 INFO MainThread:2312212 [wandb_init.py:init():779] updated telemetry
|
| 19 |
+
2024-10-24 02:16:25,618 INFO MainThread:2312212 [wandb_init.py:init():812] communicating run to backend with 90.0 second timeout
|
| 20 |
+
2024-10-24 02:16:26,156 INFO MainThread:2312212 [wandb_init.py:init():863] starting run threads in backend
|
| 21 |
+
2024-10-24 02:16:26,723 INFO MainThread:2312212 [wandb_run.py:_console_start():2465] atexit reg
|
| 22 |
+
2024-10-24 02:16:26,723 INFO MainThread:2312212 [wandb_run.py:_redirect():2313] redirect: wrap_raw
|
| 23 |
+
2024-10-24 02:16:26,723 INFO MainThread:2312212 [wandb_run.py:_redirect():2378] Wrapping output streams.
|
| 24 |
+
2024-10-24 02:16:26,723 INFO MainThread:2312212 [wandb_run.py:_redirect():2403] Redirects installed.
|
| 25 |
+
2024-10-24 02:16:26,731 INFO MainThread:2312212 [wandb_init.py:init():907] run started, returning control to user process
|
fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/run-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427.wandb
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:fae4667e9bd6c4dd0115f39374a34a0cd6bde90440ba75d82671c57fa7cd83d1
|
| 3 |
+
size 1343488
|
fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/code/src/HCP_downstream_finetune.py
ADDED
|
@@ -0,0 +1,597 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
# coding: utf-8
|
| 3 |
+
|
| 4 |
+
# In[1]:
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
# Import packages and setup gpu configuration.
|
| 8 |
+
# This code block shouldnt need to be adjusted!
|
| 9 |
+
import os
|
| 10 |
+
import sys
|
| 11 |
+
import json
|
| 12 |
+
import yaml
|
| 13 |
+
import numpy as np
|
| 14 |
+
import copy
|
| 15 |
+
import math
|
| 16 |
+
import time
|
| 17 |
+
import random
|
| 18 |
+
from tqdm.auto import tqdm
|
| 19 |
+
import webdataset as wds
|
| 20 |
+
import matplotlib.pyplot as plt
|
| 21 |
+
|
| 22 |
+
import torch
|
| 23 |
+
import torch.nn as nn
|
| 24 |
+
from torchvision import transforms
|
| 25 |
+
import utils
|
| 26 |
+
from mae_utils.flat_models import *
|
| 27 |
+
import h5py
|
| 28 |
+
from mae_utils import flat_models
|
| 29 |
+
|
| 30 |
+
# tf32 data type is faster than standard float32
|
| 31 |
+
torch.backends.cuda.matmul.allow_tf32 = True
|
| 32 |
+
# following fixes a Conv3D CUDNN_NOT_SUPPORTED error
|
| 33 |
+
torch.backends.cudnn.benchmark = True
|
| 34 |
+
|
| 35 |
+
# ## MODEL TO LOAD ##
|
| 36 |
+
if utils.is_interactive():
|
| 37 |
+
model_name = "HCPflat_large_gsrFalse_"
|
| 38 |
+
else:
|
| 39 |
+
model_name = sys.argv[1]
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
# outdir = os.path.abspath(f'checkpoints/{model_name}')
|
| 43 |
+
outdir = os.path.abspath(f'checkpoints/{model_name}')
|
| 44 |
+
|
| 45 |
+
print("outdir", outdir)
|
| 46 |
+
# Load previous config.yaml if available
|
| 47 |
+
if os.path.exists(f"{outdir}/config.yaml"):
|
| 48 |
+
config = yaml.load(open(f"{outdir}/config.yaml", 'r'), Loader=yaml.FullLoader)
|
| 49 |
+
print(f"Loaded config.yaml from ckpt folder {outdir}")
|
| 50 |
+
# create global variables from the config
|
| 51 |
+
print("\n__CONFIG__")
|
| 52 |
+
for attribute_name in config.keys():
|
| 53 |
+
print(f"{attribute_name} = {config[attribute_name]}")
|
| 54 |
+
globals()[attribute_name] = config[f'{attribute_name}']
|
| 55 |
+
print("\n")
|
| 56 |
+
|
| 57 |
+
world_size = os.getenv('WORLD_SIZE')
|
| 58 |
+
if world_size is None:
|
| 59 |
+
world_size = 1
|
| 60 |
+
else:
|
| 61 |
+
world_size = int(world_size)
|
| 62 |
+
print(f"WORLD_SIZE={world_size}")
|
| 63 |
+
|
| 64 |
+
if utils.is_interactive():
|
| 65 |
+
# Following allows you to change functions in models.py or utils.py and
|
| 66 |
+
# have this notebook automatically update with your revisions
|
| 67 |
+
get_ipython().run_line_magic('load_ext', 'autoreload')
|
| 68 |
+
get_ipython().run_line_magic('autoreload', '2')
|
| 69 |
+
|
| 70 |
+
batch_size = probe_batch_size
|
| 71 |
+
num_epochs = probe_num_epochs
|
| 72 |
+
|
| 73 |
+
data_type = torch.float32 # change depending on your mixed_precision
|
| 74 |
+
global_batch_size = batch_size * world_size
|
| 75 |
+
|
| 76 |
+
device = torch.device('cuda')
|
| 77 |
+
|
| 78 |
+
hcp_flat_path = "/weka/proj-medarc/shared/HCP-Flat"
|
| 79 |
+
# seed = 42
|
| 80 |
+
# num_frames = 16
|
| 81 |
+
# gsr = False
|
| 82 |
+
# num_workers = 10
|
| 83 |
+
# batch_size = 128
|
| 84 |
+
save_ckpt = True
|
| 85 |
+
wandb_log = True
|
| 86 |
+
print("PID of this process =",os.getpid())
|
| 87 |
+
utils.seed_everything(seed)
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
# In[2]:
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
if os.getenv('global_pool') == "False":
|
| 94 |
+
global_pool = False
|
| 95 |
+
else:
|
| 96 |
+
global_pool = True
|
| 97 |
+
print(f"global_pool = {global_pool}")
|
| 98 |
+
|
| 99 |
+
try:
|
| 100 |
+
gsr
|
| 101 |
+
except:
|
| 102 |
+
gsr = True
|
| 103 |
+
print("set gsr to True")
|
| 104 |
+
print(f"gsr = {gsr}")
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
# In[3]:
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
#### UNCOMMENT THIS TO SAVE THE HCP-FLAT IN HDF5 FORMAT
|
| 111 |
+
|
| 112 |
+
|
| 113 |
+
# from torch.utils.data import default_collate
|
| 114 |
+
# from mae_utils.flat import load_hcp_flat_mask
|
| 115 |
+
# from mae_utils.flat import create_hcp_flat
|
| 116 |
+
# from mae_utils.flat import batch_unmask
|
| 117 |
+
# import mae_utils.visualize as vis
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
# batch_size = 26
|
| 121 |
+
# print(f"changed batch_size to {batch_size}")
|
| 122 |
+
|
| 123 |
+
# ## Test ##
|
| 124 |
+
# datasets_to_include = "HCP"
|
| 125 |
+
# assert "HCP" in datasets_to_include
|
| 126 |
+
# test_dataset = create_hcp_flat(root=hcp_flat_path,
|
| 127 |
+
# clip_mode="event", frames=num_frames, shuffle=False, gsr=gsr, sub_list = 'test')
|
| 128 |
+
# test_dl = wds.WebLoader(
|
| 129 |
+
# test_dataset.batched(batch_size, partial=False, collation_fn=default_collate),
|
| 130 |
+
# batch_size=None,
|
| 131 |
+
# shuffle=False,
|
| 132 |
+
# num_workers=num_workers,
|
| 133 |
+
# pin_memory=True,
|
| 134 |
+
# )
|
| 135 |
+
|
| 136 |
+
# ## Train ##
|
| 137 |
+
# assert "HCP" in datasets_to_include
|
| 138 |
+
# train_dataset = create_hcp_flat(root=hcp_flat_path,
|
| 139 |
+
# clip_mode="event", frames=num_frames, shuffle=False, gsr=gsr, sub_list = 'train')
|
| 140 |
+
# train_dl = wds.WebLoader(
|
| 141 |
+
# train_dataset.batched(batch_size, partial=False, collation_fn=default_collate),
|
| 142 |
+
# batch_size=None,
|
| 143 |
+
# shuffle=False,
|
| 144 |
+
# num_workers=num_workers,
|
| 145 |
+
# pin_memory=True,
|
| 146 |
+
# )
|
| 147 |
+
|
| 148 |
+
# def flatten_meta(meta_dict):
|
| 149 |
+
# """
|
| 150 |
+
# Flatten the meta dictionary by:
|
| 151 |
+
# - Replacing single-item lists with the item itself.
|
| 152 |
+
# - Converting tensors to scalar numbers.
|
| 153 |
+
# """
|
| 154 |
+
# flattened = {}
|
| 155 |
+
# for key, value in meta_dict.items():
|
| 156 |
+
# if isinstance(value, list):
|
| 157 |
+
# if len(value) == 1:
|
| 158 |
+
# flattened[key] = value[0] # Replace list with its single item
|
| 159 |
+
# else:
|
| 160 |
+
# flattened[key] = value # Keep as is if multiple items
|
| 161 |
+
# elif isinstance(value, torch.Tensor):
|
| 162 |
+
# # Convert tensor to scalar
|
| 163 |
+
# if value.numel() == 1:
|
| 164 |
+
# flattened[key] = value.item()
|
| 165 |
+
# else:
|
| 166 |
+
# flattened[key] = value.tolist() # Convert multi-element tensor to list
|
| 167 |
+
# else:
|
| 168 |
+
# flattened[key] = value # Keep the value as is
|
| 169 |
+
# return flattened
|
| 170 |
+
|
| 171 |
+
# import h5py
|
| 172 |
+
# meta_array = np.array([], dtype=object)
|
| 173 |
+
# # Open an HDF5 file in write mode
|
| 174 |
+
# with h5py.File('train_hcp.hdf5', 'w') as h5f:
|
| 175 |
+
# flatmaps_dset = None
|
| 176 |
+
|
| 177 |
+
# total_samples = 0
|
| 178 |
+
|
| 179 |
+
# for i, batch in tqdm(enumerate(train_dl), total = 120000):
|
| 180 |
+
# images = batch['image'][0]
|
| 181 |
+
# meta = batch['meta']
|
| 182 |
+
# batch_size = images.shape[0]
|
| 183 |
+
# meta_serializable = meta.copy()
|
| 184 |
+
|
| 185 |
+
|
| 186 |
+
# # Step 2: Serialize the dictionary to a JSON string
|
| 187 |
+
# meta_str = json.dumps(flatten_meta(meta_serializable), indent=4)
|
| 188 |
+
# meta_array = np.append(meta_array, meta_str)
|
| 189 |
+
# if flatmaps_dset is None:
|
| 190 |
+
# # Initialize datasets with unlimited (None) maxshape along the first axis
|
| 191 |
+
# flatmaps_shape = (0,) + images.shape[1:]
|
| 192 |
+
# flatmaps_maxshape = (None,) + images.shape[1:]
|
| 193 |
+
|
| 194 |
+
# flatmaps_dset = h5f.create_dataset(
|
| 195 |
+
# 'flatmaps',
|
| 196 |
+
# shape=flatmaps_shape,
|
| 197 |
+
# maxshape=flatmaps_maxshape,
|
| 198 |
+
# dtype=np.float16,
|
| 199 |
+
# chunks=True # Enable chunking for efficient resizing
|
| 200 |
+
# )
|
| 201 |
+
|
| 202 |
+
# # Resize datasets to accommodate new data
|
| 203 |
+
# flatmaps_dset.resize(total_samples + batch_size, axis=0)
|
| 204 |
+
|
| 205 |
+
# # Write data to the datasets
|
| 206 |
+
# flatmaps_dset[total_samples:total_samples + batch_size] = images.numpy().astype(np.float16)
|
| 207 |
+
|
| 208 |
+
# total_samples += batch_size
|
| 209 |
+
|
| 210 |
+
# print(f"Processed {total_samples} samples")
|
| 211 |
+
# np.save('metadata_test_HCP.npy', meta_array)
|
| 212 |
+
|
| 213 |
+
|
| 214 |
+
# import h5py
|
| 215 |
+
# meta_array = np.array([], dtype=object)
|
| 216 |
+
# # Open an HDF5 file in write mode
|
| 217 |
+
# with h5py.File('test_hcp.hdf5', 'w') as h5f:
|
| 218 |
+
# flatmaps_dset = None
|
| 219 |
+
|
| 220 |
+
# total_samples = 0
|
| 221 |
+
|
| 222 |
+
# for i, batch in tqdm(enumerate(test_dl), total = 12000):
|
| 223 |
+
# images = batch['image'][0]
|
| 224 |
+
# meta = batch['meta']
|
| 225 |
+
# batch_size = images.shape[0]
|
| 226 |
+
# meta_serializable = meta.copy()
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
# # Step 2: Serialize the dictionary to a JSON string
|
| 230 |
+
# meta_str = json.dumps(flatten_meta(meta_serializable), indent=4)
|
| 231 |
+
# meta_array = np.append(meta_array, meta_str)
|
| 232 |
+
# if flatmaps_dset is None:
|
| 233 |
+
# # Initialize datasets with unlimited (None) maxshape along the first axis
|
| 234 |
+
# flatmaps_shape = (0,) + images.shape[1:]
|
| 235 |
+
# flatmaps_maxshape = (None,) + images.shape[1:]
|
| 236 |
+
|
| 237 |
+
# flatmaps_dset = h5f.create_dataset(
|
| 238 |
+
# 'flatmaps',
|
| 239 |
+
# shape=flatmaps_shape,
|
| 240 |
+
# maxshape=flatmaps_maxshape,
|
| 241 |
+
# dtype=np.float16,
|
| 242 |
+
# chunks=True # Enable chunking for efficient resizing
|
| 243 |
+
# )
|
| 244 |
+
|
| 245 |
+
# # Resize datasets to accommodate new data
|
| 246 |
+
# flatmaps_dset.resize(total_samples + batch_size, axis=0)
|
| 247 |
+
|
| 248 |
+
# # Write data to the datasets
|
| 249 |
+
# flatmaps_dset[total_samples:total_samples + batch_size] = images.numpy().astype(np.float16)
|
| 250 |
+
|
| 251 |
+
# total_samples += batch_size
|
| 252 |
+
|
| 253 |
+
# print(f"Processed {total_samples} samples")
|
| 254 |
+
# np.save('metadata_train_HCP.npy', meta_array)
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
# ### Preparing data
|
| 258 |
+
|
| 259 |
+
# In[4]:
|
| 260 |
+
|
| 261 |
+
|
| 262 |
+
from sklearn.preprocessing import LabelEncoder
|
| 263 |
+
|
| 264 |
+
INCLUDE_CONDS = {
|
| 265 |
+
"fear",
|
| 266 |
+
"neut",
|
| 267 |
+
"math",
|
| 268 |
+
"story",
|
| 269 |
+
"lf",
|
| 270 |
+
"lh",
|
| 271 |
+
"rf",
|
| 272 |
+
"rh",
|
| 273 |
+
"t",
|
| 274 |
+
"match",
|
| 275 |
+
"relation",
|
| 276 |
+
"mental",
|
| 277 |
+
"rnd",
|
| 278 |
+
"0bk_body",
|
| 279 |
+
"2bk_body",
|
| 280 |
+
"0bk_faces",
|
| 281 |
+
"2bk_faces",
|
| 282 |
+
"0bk_places",
|
| 283 |
+
"2bk_places",
|
| 284 |
+
"0bk_tools",
|
| 285 |
+
"2bk_tools",
|
| 286 |
+
}
|
| 287 |
+
|
| 288 |
+
# test_data = []
|
| 289 |
+
|
| 290 |
+
# # Iterate over the DataLoader with a progress bar
|
| 291 |
+
# for sample in tqdm(train_dl, desc="Processing samples"):
|
| 292 |
+
# x = sample['image']
|
| 293 |
+
# y = sample['meta']['trial_type']
|
| 294 |
+
# key = sample['meta']['key']
|
| 295 |
+
# print(x.shape, y, key)
|
| 296 |
+
# break
|
| 297 |
+
# Initialize the label encoder
|
| 298 |
+
label_encoder = LabelEncoder()
|
| 299 |
+
label_encoder.fit(sorted(INCLUDE_CONDS)) # Ensure consistent ordering
|
| 300 |
+
|
| 301 |
+
num_classes = len(label_encoder.classes_)
|
| 302 |
+
print(f"Number of classes: {num_classes}")
|
| 303 |
+
|
| 304 |
+
|
| 305 |
+
# In[5]:
|
| 306 |
+
|
| 307 |
+
|
| 308 |
+
f_train = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/train_hcp.hdf5', 'r')
|
| 309 |
+
flatmaps_train = f_train['flatmaps']
|
| 310 |
+
|
| 311 |
+
f_test = h5py.File('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/test_hcp.hdf5', 'r')
|
| 312 |
+
flatmaps_test = f_test['flatmaps']
|
| 313 |
+
|
| 314 |
+
metadata_train = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_train_HCP.npy', allow_pickle=True)
|
| 315 |
+
metadata_test = np.load('/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/metadata_test_HCP.npy', allow_pickle=True)
|
| 316 |
+
|
| 317 |
+
|
| 318 |
+
# In[6]:
|
| 319 |
+
|
| 320 |
+
|
| 321 |
+
from torch.utils.data import Dataset, DataLoader
|
| 322 |
+
|
| 323 |
+
class HCPFlatDataset(Dataset):
|
| 324 |
+
def __init__(self, flatmaps, metadata):
|
| 325 |
+
self.flatmaps = flatmaps
|
| 326 |
+
self.metadata = metadata
|
| 327 |
+
|
| 328 |
+
def __len__(self):
|
| 329 |
+
return len(self.metadata)
|
| 330 |
+
|
| 331 |
+
def __getitem__(self, idx):
|
| 332 |
+
return self.flatmaps[idx], json.loads(self.metadata[idx])
|
| 333 |
+
print("Moving datasets to ram")
|
| 334 |
+
# Loading to cpu for faster training, this can take several minutes. Remove this [:] if you want to move one at the time.
|
| 335 |
+
train_dataset = HCPFlatDataset(flatmaps_train, metadata_train)
|
| 336 |
+
train_dl = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)
|
| 337 |
+
|
| 338 |
+
test_dataset = HCPFlatDataset(flatmaps_test, metadata_test)
|
| 339 |
+
test_dl = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
|
| 340 |
+
print("Datasets ready")
|
| 341 |
+
|
| 342 |
+
|
| 343 |
+
# ### Creating and loading Model
|
| 344 |
+
|
| 345 |
+
# In[7]:
|
| 346 |
+
|
| 347 |
+
|
| 348 |
+
from mae_utils.flat import load_hcp_flat_mask
|
| 349 |
+
from mae_utils.flat import create_hcp_flat
|
| 350 |
+
from mae_utils.flat import batch_unmask
|
| 351 |
+
import mae_utils.visualize as vis
|
| 352 |
+
|
| 353 |
+
flat_mask = load_hcp_flat_mask(hcp_flat_path)
|
| 354 |
+
|
| 355 |
+
mae_model = flat_models.mae_vit_large_fmri(
|
| 356 |
+
patch_size=patch_size,
|
| 357 |
+
decoder_embed_dim=decoder_embed_dim,
|
| 358 |
+
t_patch_size=t_patch_size,
|
| 359 |
+
pred_t_dim=pred_t_dim,
|
| 360 |
+
decoder_depth=4,
|
| 361 |
+
cls_embed=cls_embed,
|
| 362 |
+
norm_pix_loss=norm_pix_loss,
|
| 363 |
+
no_qkv_bias=no_qkv_bias,
|
| 364 |
+
sep_pos_embed=sep_pos_embed,
|
| 365 |
+
trunc_init=trunc_init,
|
| 366 |
+
pct_masks_to_decode=pct_masks_to_decode,
|
| 367 |
+
img_mask=flat_mask,
|
| 368 |
+
)
|
| 369 |
+
|
| 370 |
+
|
| 371 |
+
# In[8]:
|
| 372 |
+
|
| 373 |
+
|
| 374 |
+
checkpoint_files = [f for f in os.listdir(outdir) if f.endswith('.pth')]
|
| 375 |
+
|
| 376 |
+
if utils.is_interactive():
|
| 377 |
+
latest_checkpoint = "epoch99.pth"
|
| 378 |
+
else:
|
| 379 |
+
latest_checkpoint = sys.argv[2]
|
| 380 |
+
print(f"latest_checkpoint: {latest_checkpoint}")
|
| 381 |
+
|
| 382 |
+
# Load the checkpoint
|
| 383 |
+
checkpoint_path = os.path.join(outdir, latest_checkpoint)
|
| 384 |
+
|
| 385 |
+
state = torch.load(checkpoint_path)
|
| 386 |
+
mae_model.load_state_dict(state["model_state_dict"], strict=False)
|
| 387 |
+
mae_model.to(device)
|
| 388 |
+
|
| 389 |
+
print(f"\nLoaded checkpoint {latest_checkpoint} from {outdir}\n")
|
| 390 |
+
|
| 391 |
+
|
| 392 |
+
# In[9]:
|
| 393 |
+
|
| 394 |
+
|
| 395 |
+
class LinearClassifier(nn.Module):
|
| 396 |
+
def __init__(self, input_dim, num_classes):
|
| 397 |
+
super(LinearClassifier, self).__init__()
|
| 398 |
+
self.linear = nn.Linear(input_dim, num_classes)
|
| 399 |
+
|
| 400 |
+
def forward(self, x):
|
| 401 |
+
# Flatten the input except for the batch dimension
|
| 402 |
+
x = x.view(x.size(0), -1)
|
| 403 |
+
out = self.linear(x)
|
| 404 |
+
return out # Raw logits
|
| 405 |
+
|
| 406 |
+
# Determine the input dimension from a single sample
|
| 407 |
+
# Assuming images are of shape [1, 16, 144, 320]
|
| 408 |
+
input_dim = np.prod(mae_model(torch.randn(1,1,16,144,320).to(device),global_pool=global_pool, forward_features = True).shape[1:])
|
| 409 |
+
print(f"Input dimension: {input_dim}")
|
| 410 |
+
|
| 411 |
+
|
| 412 |
+
# In[10]:
|
| 413 |
+
|
| 414 |
+
|
| 415 |
+
class FullModel(nn.Module):
|
| 416 |
+
def __init__(self, lc_model, mae_model):
|
| 417 |
+
super(FullModel, self).__init__()
|
| 418 |
+
self.lc_model = lc_model
|
| 419 |
+
self.mae_model = mae_model
|
| 420 |
+
|
| 421 |
+
|
| 422 |
+
def forward(self, x, gsr):
|
| 423 |
+
x = self.mae_model(x, global_pool=global_pool, forward_features = True)
|
| 424 |
+
x = self.lc_model(x)
|
| 425 |
+
return x
|
| 426 |
+
|
| 427 |
+
|
| 428 |
+
# In[11]:
|
| 429 |
+
|
| 430 |
+
|
| 431 |
+
# Initialize the model
|
| 432 |
+
lc_model = LinearClassifier(input_dim=input_dim, num_classes=num_classes)
|
| 433 |
+
|
| 434 |
+
model = FullModel(lc_model, mae_model)
|
| 435 |
+
|
| 436 |
+
# Move the model to the GPU
|
| 437 |
+
model.to(device)
|
| 438 |
+
|
| 439 |
+
# Define loss function
|
| 440 |
+
criterion = nn.CrossEntropyLoss()
|
| 441 |
+
|
| 442 |
+
# Define optimizer with L2 regularization (weight_decay)
|
| 443 |
+
learning_rate = 1e-4
|
| 444 |
+
weight_decay = 1e-5 # Adjust based on your needs
|
| 445 |
+
optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)
|
| 446 |
+
num_epochs = 20 # Adjust as needed
|
| 447 |
+
|
| 448 |
+
|
| 449 |
+
# ### Data
|
| 450 |
+
|
| 451 |
+
# In[16]:
|
| 452 |
+
|
| 453 |
+
|
| 454 |
+
import uuid
|
| 455 |
+
|
| 456 |
+
myuuid = uuid.uuid4()
|
| 457 |
+
str(myuuid)
|
| 458 |
+
|
| 459 |
+
|
| 460 |
+
# In[17]:
|
| 461 |
+
|
| 462 |
+
|
| 463 |
+
import wandb
|
| 464 |
+
|
| 465 |
+
if utils.is_interactive():
|
| 466 |
+
print("Running in interactive notebook. Disabling W&B and ckpt saving.")
|
| 467 |
+
wandb_log = True
|
| 468 |
+
save_ckpt = True
|
| 469 |
+
|
| 470 |
+
if wandb_log:
|
| 471 |
+
wandb_project = 'fMRI-foundation-model'
|
| 472 |
+
wandb_config = {
|
| 473 |
+
"model_name": model_name+'_HCP_FT',
|
| 474 |
+
"batch_size": batch_size,
|
| 475 |
+
"learning_rate": learning_rate,
|
| 476 |
+
"weight_decay": weight_decay,
|
| 477 |
+
"num_epochs": num_epochs,
|
| 478 |
+
"seed": seed,
|
| 479 |
+
}
|
| 480 |
+
print("wandb_config:\n", wandb_config)
|
| 481 |
+
random_id = str(uuid.uuid4())
|
| 482 |
+
print("wandb_id:", "HCPflat_raw" + f"_{random_id}")
|
| 483 |
+
wandb.init(
|
| 484 |
+
id=model_name+'_HCP_FT' + f"_{random_id}",
|
| 485 |
+
project=wandb_project,
|
| 486 |
+
name=model_name+'_HCP_FT',
|
| 487 |
+
config=wandb_config,
|
| 488 |
+
resume="allow",
|
| 489 |
+
)
|
| 490 |
+
|
| 491 |
+
|
| 492 |
+
# In[13]:
|
| 493 |
+
|
| 494 |
+
|
| 495 |
+
for epoch in range(num_epochs):
|
| 496 |
+
running_train_loss = 0.0
|
| 497 |
+
correct_train = 0
|
| 498 |
+
total_train = 0
|
| 499 |
+
step = 0
|
| 500 |
+
|
| 501 |
+
# with torch.amp.autocast(device_type='cuda'):
|
| 502 |
+
# Training Phase
|
| 503 |
+
model.train()
|
| 504 |
+
for batch in tqdm(train_dl, desc=f"Epoch {epoch+1}/{num_epochs} - Training"):
|
| 505 |
+
optimizer.zero_grad()
|
| 506 |
+
images = batch[0].to(device).float().unsqueeze(1) #fix this # Shape: [batch_size, 1, 16, 144, 320]
|
| 507 |
+
labels = batch[1]['trial_type'] # List of labels
|
| 508 |
+
|
| 509 |
+
encoded_labels = label_encoder.transform(labels)
|
| 510 |
+
encoded_labels = torch.tensor(encoded_labels, dtype=torch.long).to(device) # Shape: [batch_size]
|
| 511 |
+
|
| 512 |
+
# Forward pass
|
| 513 |
+
outputs = model(images, gsr=gsr) # Shape: [num_train_samples, num_classes]
|
| 514 |
+
|
| 515 |
+
# Compute loss
|
| 516 |
+
loss = criterion(outputs, encoded_labels)
|
| 517 |
+
|
| 518 |
+
# Backward pass and optimization
|
| 519 |
+
loss.backward()
|
| 520 |
+
optimizer.step()
|
| 521 |
+
|
| 522 |
+
# Accumulate loss
|
| 523 |
+
running_train_loss += loss.item() * images.size(0)
|
| 524 |
+
|
| 525 |
+
|
| 526 |
+
# Calculate accuracy
|
| 527 |
+
_, predicted = torch.max(outputs, 1)
|
| 528 |
+
|
| 529 |
+
correct_train += (predicted == encoded_labels).sum().item()
|
| 530 |
+
total_train += encoded_labels.size(0)
|
| 531 |
+
|
| 532 |
+
step = step + 1
|
| 533 |
+
if step % 100 == 0:
|
| 534 |
+
print(f"Step [{step}/{len(train_dl)}] - Training Loss: {loss.item():.4f} - Training Accuracy: {100 * correct_train / total_train:.2f}%")
|
| 535 |
+
# thth
|
| 536 |
+
|
| 537 |
+
epoch_train_loss = running_train_loss / total_train if total_train > 0 else 0.0
|
| 538 |
+
train_accuracy = 100 * correct_train / total_train if total_train > 0 else 0.0
|
| 539 |
+
|
| 540 |
+
# Validation Phase
|
| 541 |
+
model.eval()
|
| 542 |
+
running_val_loss = 0.0
|
| 543 |
+
correct_val = 0
|
| 544 |
+
total_val = 0
|
| 545 |
+
|
| 546 |
+
with torch.no_grad():
|
| 547 |
+
for batch in tqdm(test_dl, desc=f"Epoch {epoch+1}/{num_epochs} - Validation"):
|
| 548 |
+
|
| 549 |
+
images = batch[0].to(device).float().unsqueeze(1) #fix this
|
| 550 |
+
labels = batch[1]['trial_type']
|
| 551 |
+
|
| 552 |
+
# Encode labels to integer indices
|
| 553 |
+
encoded_labels = label_encoder.transform(labels)
|
| 554 |
+
encoded_labels = torch.tensor(encoded_labels, dtype=torch.long).to(device)
|
| 555 |
+
|
| 556 |
+
|
| 557 |
+
# Forward pass
|
| 558 |
+
outputs = model(images, gsr=gsr)
|
| 559 |
+
|
| 560 |
+
# Compute loss
|
| 561 |
+
loss = criterion(outputs, encoded_labels)
|
| 562 |
+
|
| 563 |
+
# Accumulate loss
|
| 564 |
+
running_val_loss += loss.item() * images.size(0)
|
| 565 |
+
|
| 566 |
+
# Calculate accuracy
|
| 567 |
+
_, predicted = torch.max(outputs, 1)
|
| 568 |
+
correct_val += (predicted == encoded_labels).sum().item()
|
| 569 |
+
total_val += encoded_labels.size(0)
|
| 570 |
+
|
| 571 |
+
|
| 572 |
+
|
| 573 |
+
epoch_val_loss = running_val_loss / total_val if total_val > 0 else 0.0
|
| 574 |
+
val_accuracy = 100 * correct_val / total_val if total_val > 0 else 0.0
|
| 575 |
+
|
| 576 |
+
print(f"Epoch [{epoch+1}/{num_epochs}] "
|
| 577 |
+
f"- Training Loss: {epoch_train_loss:.4f}, Training Accuracy: {train_accuracy:.2f}% "
|
| 578 |
+
f"- Validation Loss: {epoch_val_loss:.4f}, Validation Accuracy: {val_accuracy:.2f}%")
|
| 579 |
+
|
| 580 |
+
if wandb_log:
|
| 581 |
+
wandb.log({
|
| 582 |
+
"epoch_train_loss": epoch_train_loss,
|
| 583 |
+
"epoch_val_loss": epoch_val_loss,
|
| 584 |
+
"train_accuracy": train_accuracy,
|
| 585 |
+
"val_accuracy": val_accuracy,
|
| 586 |
+
})
|
| 587 |
+
if save_ckpt:
|
| 588 |
+
outdir = os.path.abspath(f'checkpoints/{model_name+"HCP_FT"}')
|
| 589 |
+
os.makedirs(outdir, exist_ok=True)
|
| 590 |
+
print("outdir", outdir)
|
| 591 |
+
# Save model and config
|
| 592 |
+
torch.save(model.state_dict(), f"{outdir}/model.pth")
|
| 593 |
+
with open(f"{outdir}/config.yaml", 'w') as f:
|
| 594 |
+
yaml.dump(wandb_config, f)
|
| 595 |
+
print(f"Saved model and config to {outdir}")
|
| 596 |
+
|
| 597 |
+
|
fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/output.log
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/requirements.txt
ADDED
|
@@ -0,0 +1,198 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
protobuf==5.28.2
|
| 2 |
+
imageio==2.35.1
|
| 3 |
+
MarkupSafe==3.0.0
|
| 4 |
+
regex==2024.9.11
|
| 5 |
+
matplotlib==3.9.2
|
| 6 |
+
notebook==7.2.2
|
| 7 |
+
debugpy==1.8.6
|
| 8 |
+
aiosignal==1.3.1
|
| 9 |
+
jupyter_core==5.7.2
|
| 10 |
+
torchaudio==2.4.1+cu121
|
| 11 |
+
python-json-logger==2.0.7
|
| 12 |
+
six==1.16.0
|
| 13 |
+
scikit-image==0.24.0
|
| 14 |
+
types-python-dateutil==2.9.0.20241003
|
| 15 |
+
PyYAML==6.0.2
|
| 16 |
+
httpcore==1.0.6
|
| 17 |
+
clip==1.0
|
| 18 |
+
babel==2.16.0
|
| 19 |
+
webcolors==24.8.0
|
| 20 |
+
omegaconf==2.3.0
|
| 21 |
+
webencodings==0.5.1
|
| 22 |
+
kiwisolver==1.4.7
|
| 23 |
+
uri-template==1.3.0
|
| 24 |
+
diffusers==0.23.0
|
| 25 |
+
idna==3.10
|
| 26 |
+
fsspec==2024.9.0
|
| 27 |
+
parso==0.8.4
|
| 28 |
+
setuptools==65.5.0
|
| 29 |
+
tornado==6.4.1
|
| 30 |
+
webdataset==0.2.100
|
| 31 |
+
decord==0.6.0
|
| 32 |
+
nvidia-curand-cu12==10.3.2.106
|
| 33 |
+
ipykernel==6.29.5
|
| 34 |
+
jupyter==1.1.1
|
| 35 |
+
pexpect==4.9.0
|
| 36 |
+
kornia_rs==0.1.5
|
| 37 |
+
iopath==0.1.10
|
| 38 |
+
async-lru==2.0.4
|
| 39 |
+
future==1.0.0
|
| 40 |
+
torchvision==0.19.1+cu121
|
| 41 |
+
botocore==1.34.162
|
| 42 |
+
cycler==0.12.1
|
| 43 |
+
tzdata==2024.2
|
| 44 |
+
jupyter_server_terminals==0.5.3
|
| 45 |
+
click==8.1.7
|
| 46 |
+
einops==0.8.0
|
| 47 |
+
pyzmq==26.2.0
|
| 48 |
+
jupyter_client==8.6.3
|
| 49 |
+
nbconvert==7.16.4
|
| 50 |
+
scikit-learn==1.5.2
|
| 51 |
+
executing==2.1.0
|
| 52 |
+
asttokens==2.4.1
|
| 53 |
+
docker-pycreds==0.4.0
|
| 54 |
+
matplotlib-inline==0.1.7
|
| 55 |
+
overrides==7.7.0
|
| 56 |
+
websocket-client==1.8.0
|
| 57 |
+
nbformat==5.10.4
|
| 58 |
+
elbow==0.1.1
|
| 59 |
+
contourpy==1.3.0
|
| 60 |
+
nvidia-cudnn-cu12==9.1.0.70
|
| 61 |
+
transformers==4.44.2
|
| 62 |
+
gitdb==4.0.11
|
| 63 |
+
jupyterlab_nvdashboard==0.11.0
|
| 64 |
+
lazy_loader==0.4
|
| 65 |
+
jsonpointer==3.0.0
|
| 66 |
+
notebook_shim==0.2.4
|
| 67 |
+
nvidia-nccl-cu12==2.20.5
|
| 68 |
+
ffmpeg-python==0.2.0
|
| 69 |
+
triton==3.0.0
|
| 70 |
+
mistune==3.0.2
|
| 71 |
+
python-dateutil==2.9.0.post0
|
| 72 |
+
beautifulsoup4==4.12.3
|
| 73 |
+
nbclient==0.10.0
|
| 74 |
+
h5py==3.12.1
|
| 75 |
+
ftfy==6.2.3
|
| 76 |
+
zipp==3.20.2
|
| 77 |
+
ptyprocess==0.7.0
|
| 78 |
+
huggingface-hub==0.25.1
|
| 79 |
+
pytz==2024.2
|
| 80 |
+
jupyterlab_pygments==0.3.0
|
| 81 |
+
nvidia-cublas-cu12==12.1.3.1
|
| 82 |
+
pandocfilters==1.5.1
|
| 83 |
+
Jinja2==3.1.4
|
| 84 |
+
arrow==1.3.0
|
| 85 |
+
rpds-py==0.20.0
|
| 86 |
+
jupyter_server==2.14.2
|
| 87 |
+
simplejson==3.19.3
|
| 88 |
+
networkx==3.3
|
| 89 |
+
packaging==24.1
|
| 90 |
+
traitlets==5.14.3
|
| 91 |
+
pandas==2.2.3
|
| 92 |
+
xformers==0.0.22.post7
|
| 93 |
+
lightning-utilities==0.11.7
|
| 94 |
+
tifffile==2024.9.20
|
| 95 |
+
nvidia-cuda-cupti-cu12==12.1.105
|
| 96 |
+
mpmath==1.3.0
|
| 97 |
+
GitPython==3.1.43
|
| 98 |
+
scipy==1.14.1
|
| 99 |
+
jsonschema==4.23.0
|
| 100 |
+
prompt_toolkit==3.0.48
|
| 101 |
+
s3transfer==0.10.2
|
| 102 |
+
multidict==6.1.0
|
| 103 |
+
bleach==6.1.0
|
| 104 |
+
sentry-sdk==2.15.0
|
| 105 |
+
nibabel==5.2.1
|
| 106 |
+
accelerate==1.0.0
|
| 107 |
+
pyarrow==17.0.0
|
| 108 |
+
threadpoolctl==3.5.0
|
| 109 |
+
attrs==24.2.0
|
| 110 |
+
rfc3986-validator==0.1.1
|
| 111 |
+
nvidia-cuda-runtime-cu12==12.1.105
|
| 112 |
+
ipywidgets==8.1.5
|
| 113 |
+
frozenlist==1.4.1
|
| 114 |
+
pycparser==2.22
|
| 115 |
+
jupyterlab_server==2.27.3
|
| 116 |
+
nvidia-cuda-nvrtc-cu12==12.1.105
|
| 117 |
+
yarl==1.13.1
|
| 118 |
+
setproctitle==1.3.3
|
| 119 |
+
isoduration==20.11.0
|
| 120 |
+
Pygments==2.18.0
|
| 121 |
+
jedi==0.19.1
|
| 122 |
+
boto3==1.34.57
|
| 123 |
+
tokenizers==0.19.1
|
| 124 |
+
referencing==0.35.1
|
| 125 |
+
rfc3339-validator==0.1.4
|
| 126 |
+
pillow==10.4.0
|
| 127 |
+
jupyterlab==4.2.5
|
| 128 |
+
stack-data==0.6.3
|
| 129 |
+
h11==0.14.0
|
| 130 |
+
anyio==4.6.0
|
| 131 |
+
nilearn==0.10.4
|
| 132 |
+
nvidia-cusolver-cu12==11.4.5.107
|
| 133 |
+
tinycss2==1.3.0
|
| 134 |
+
defusedxml==0.7.1
|
| 135 |
+
argon2-cffi-bindings==21.2.0
|
| 136 |
+
soupsieve==2.6
|
| 137 |
+
nest-asyncio==1.6.0
|
| 138 |
+
torchmetrics==1.3.0.post0
|
| 139 |
+
tqdm==4.66.5
|
| 140 |
+
cffi==1.17.1
|
| 141 |
+
charset-normalizer==3.3.2
|
| 142 |
+
jsonschema-specifications==2023.12.1
|
| 143 |
+
decorator==5.1.1
|
| 144 |
+
open_clip_torch==2.26.1
|
| 145 |
+
jupyter-events==0.10.0
|
| 146 |
+
smart-open==7.0.5
|
| 147 |
+
antlr4-python3-runtime==4.9.3
|
| 148 |
+
prometheus_client==0.21.0
|
| 149 |
+
kornia==0.7.3
|
| 150 |
+
typing_extensions==4.12.2
|
| 151 |
+
sniffio==1.3.1
|
| 152 |
+
joblib==1.4.2
|
| 153 |
+
comm==0.2.2
|
| 154 |
+
aiohappyeyeballs==2.4.3
|
| 155 |
+
numpy==2.1.2
|
| 156 |
+
braceexpand==0.1.7
|
| 157 |
+
certifi==2024.8.30
|
| 158 |
+
psutil==6.0.0
|
| 159 |
+
pyparsing==3.1.4
|
| 160 |
+
pure_eval==0.2.3
|
| 161 |
+
nvidia-cusparse-cu12==12.1.0.106
|
| 162 |
+
wandb==0.18.3
|
| 163 |
+
urllib3==2.2.3
|
| 164 |
+
smmap==5.0.1
|
| 165 |
+
platformdirs==4.3.6
|
| 166 |
+
torch==2.4.1+cu121
|
| 167 |
+
requests==2.32.3
|
| 168 |
+
json5==0.9.25
|
| 169 |
+
nvidia-nvjitlink-cu12==12.6.77
|
| 170 |
+
jupyterlab_widgets==3.0.13
|
| 171 |
+
lxml==5.3.0
|
| 172 |
+
httpx==0.27.2
|
| 173 |
+
opencv-python==4.6.0.66
|
| 174 |
+
portalocker==2.10.1
|
| 175 |
+
pytorch-lightning==2.0.1
|
| 176 |
+
sympy==1.13.3
|
| 177 |
+
wcwidth==0.2.13
|
| 178 |
+
jmespath==1.0.1
|
| 179 |
+
fqdn==1.5.1
|
| 180 |
+
pynvml==11.5.3
|
| 181 |
+
pip==24.0
|
| 182 |
+
wrapt==1.16.0
|
| 183 |
+
aiohttp==3.10.9
|
| 184 |
+
filelock==3.16.1
|
| 185 |
+
fonttools==4.54.1
|
| 186 |
+
fastjsonschema==2.20.0
|
| 187 |
+
jupyter-console==6.6.3
|
| 188 |
+
widgetsnbextension==4.0.13
|
| 189 |
+
timm==1.0.9
|
| 190 |
+
nvidia-cufft-cu12==11.0.2.54
|
| 191 |
+
ipython==8.28.0
|
| 192 |
+
nvidia-nvtx-cu12==12.1.105
|
| 193 |
+
jupyter-lsp==2.2.5
|
| 194 |
+
safetensors==0.4.5
|
| 195 |
+
terminado==0.18.1
|
| 196 |
+
argon2-cffi==23.1.0
|
| 197 |
+
Send2Trash==1.8.3
|
| 198 |
+
importlib_metadata==8.5.0
|
fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/wandb-metadata.json
ADDED
|
@@ -0,0 +1,131 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"os": "Linux-5.15.0-1058-aws-x86_64-with-glibc2.31",
|
| 3 |
+
"python": "3.11.10",
|
| 4 |
+
"startedAt": "2024-10-24T21:38:35.093209Z",
|
| 5 |
+
"args": [
|
| 6 |
+
"HCPflat_large_gsrFalse_",
|
| 7 |
+
"epoch99.pth"
|
| 8 |
+
],
|
| 9 |
+
"program": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src/HCP_downstream_finetune.py",
|
| 10 |
+
"codePath": "src/HCP_downstream_finetune.py",
|
| 11 |
+
"git": {
|
| 12 |
+
"remote": "https://github.com/MedARC-AI/fMRI-foundation-model",
|
| 13 |
+
"commit": "cf8214d4ebe437188b68b4ee5a34c5211a810db0"
|
| 14 |
+
},
|
| 15 |
+
"email": "torrico.villanueva.cesar.kadir@gmail.com",
|
| 16 |
+
"root": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src",
|
| 17 |
+
"host": "ip-10-0-152-216",
|
| 18 |
+
"username": "ckadirt",
|
| 19 |
+
"executable": "/admin/home-ckadirt/foundation_env/bin/python",
|
| 20 |
+
"codePathLocal": "HCP_downstream_finetune.py",
|
| 21 |
+
"cpu_count": 96,
|
| 22 |
+
"cpu_count_logical": 192,
|
| 23 |
+
"gpu": "[NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3, NVIDIA H100 80GB HBM3]",
|
| 24 |
+
"gpu_count": 8,
|
| 25 |
+
"disk": {
|
| 26 |
+
"/": {
|
| 27 |
+
"total": "249555763200",
|
| 28 |
+
"used": "182052564992"
|
| 29 |
+
}
|
| 30 |
+
},
|
| 31 |
+
"memory": {
|
| 32 |
+
"total": "2147443372032"
|
| 33 |
+
},
|
| 34 |
+
"cpu": {
|
| 35 |
+
"count": 96,
|
| 36 |
+
"countLogical": 192
|
| 37 |
+
},
|
| 38 |
+
"gpu_nvidia": [
|
| 39 |
+
{
|
| 40 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 41 |
+
"memoryTotal": "85520809984",
|
| 42 |
+
"cudaCores": 16896,
|
| 43 |
+
"architecture": "Hopper"
|
| 44 |
+
},
|
| 45 |
+
{
|
| 46 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 47 |
+
"memoryTotal": "85520809984",
|
| 48 |
+
"cudaCores": 16896,
|
| 49 |
+
"architecture": "Hopper"
|
| 50 |
+
},
|
| 51 |
+
{
|
| 52 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 53 |
+
"memoryTotal": "85520809984",
|
| 54 |
+
"cudaCores": 16896,
|
| 55 |
+
"architecture": "Hopper"
|
| 56 |
+
},
|
| 57 |
+
{
|
| 58 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 59 |
+
"memoryTotal": "85520809984",
|
| 60 |
+
"cudaCores": 16896,
|
| 61 |
+
"architecture": "Hopper"
|
| 62 |
+
},
|
| 63 |
+
{
|
| 64 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 65 |
+
"memoryTotal": "85520809984",
|
| 66 |
+
"cudaCores": 16896,
|
| 67 |
+
"architecture": "Hopper"
|
| 68 |
+
},
|
| 69 |
+
{
|
| 70 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 71 |
+
"memoryTotal": "85520809984",
|
| 72 |
+
"cudaCores": 16896,
|
| 73 |
+
"architecture": "Hopper"
|
| 74 |
+
},
|
| 75 |
+
{
|
| 76 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 77 |
+
"memoryTotal": "85520809984",
|
| 78 |
+
"cudaCores": 16896,
|
| 79 |
+
"architecture": "Hopper"
|
| 80 |
+
},
|
| 81 |
+
{
|
| 82 |
+
"name": "NVIDIA H100 80GB HBM3",
|
| 83 |
+
"memoryTotal": "85520809984",
|
| 84 |
+
"cudaCores": 16896,
|
| 85 |
+
"architecture": "Hopper"
|
| 86 |
+
}
|
| 87 |
+
],
|
| 88 |
+
"slurm": {
|
| 89 |
+
"cluster_name": "sagemaker2",
|
| 90 |
+
"conf": "/opt/slurm/etc/slurm.conf",
|
| 91 |
+
"cpus_on_node": "20",
|
| 92 |
+
"gpus_on_node": "1",
|
| 93 |
+
"gpus_per_task": "1",
|
| 94 |
+
"gtids": "0",
|
| 95 |
+
"job_account": "fmri",
|
| 96 |
+
"job_cpus_per_node": "20",
|
| 97 |
+
"job_end_time": "1729921061",
|
| 98 |
+
"job_gid": "1879800513",
|
| 99 |
+
"job_gpus": "1",
|
| 100 |
+
"job_id": "529188",
|
| 101 |
+
"job_name": "finetuneHCP",
|
| 102 |
+
"job_nodelist": "ip-10-0-152-216",
|
| 103 |
+
"job_num_nodes": "1",
|
| 104 |
+
"job_partition": "p5",
|
| 105 |
+
"job_qos": "normal",
|
| 106 |
+
"job_start_time": "1729805861",
|
| 107 |
+
"job_uid": "1879804696",
|
| 108 |
+
"job_user": "ckadirt",
|
| 109 |
+
"jobid": "529188",
|
| 110 |
+
"localid": "0",
|
| 111 |
+
"mem_per_cpu": "11500",
|
| 112 |
+
"nnodes": "1",
|
| 113 |
+
"node_aliases": "(null)",
|
| 114 |
+
"nodeid": "0",
|
| 115 |
+
"nodelist": "ip-10-0-152-216",
|
| 116 |
+
"nprocs": "1",
|
| 117 |
+
"ntasks": "1",
|
| 118 |
+
"ntasks_per_node": "1",
|
| 119 |
+
"prio_process": "0",
|
| 120 |
+
"procid": "0",
|
| 121 |
+
"script_context": "prolog_task",
|
| 122 |
+
"submit_dir": "/weka/proj-fmri/ckadirt/fMRI-foundation-model/src",
|
| 123 |
+
"submit_host": "ip-172-17-12-61",
|
| 124 |
+
"task_pid": "337193",
|
| 125 |
+
"tasks_per_node": "1",
|
| 126 |
+
"topology_addr": "ip-10-0-152-216",
|
| 127 |
+
"topology_addr_pattern": "node",
|
| 128 |
+
"working_cluster": "sagemaker2:ip-172-17-63-161:6817:9984:109"
|
| 129 |
+
},
|
| 130 |
+
"cudaVersion": "12.2"
|
| 131 |
+
}
|
fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/logs/debug-core.log
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{"time":"2024-10-24T21:38:34.366295935Z","level":"INFO","msg":"started logging, with flags","port-filename":"/tmp/tmpwn5d9jvk/port-337263.txt","pid":337263,"debug":false,"disable-analytics":false}
|
| 2 |
+
{"time":"2024-10-24T21:38:34.366306785Z","level":"INFO","msg":"started logging, with flags","port-filename":"/tmp/tmp_63dqr1q/port-337587.txt","pid":337587,"debug":false,"disable-analytics":false}
|
| 3 |
+
{"time":"2024-10-24T21:38:34.366577281Z","level":"INFO","msg":"FeatureState","shutdownOnParentExitEnabled":false}
|
| 4 |
+
{"time":"2024-10-24T21:38:34.366603112Z","level":"INFO","msg":"FeatureState","shutdownOnParentExitEnabled":false}
|
| 5 |
+
{"time":"2024-10-24T21:38:34.372814525Z","level":"INFO","msg":"Will exit if parent process dies.","ppid":337587}
|
| 6 |
+
{"time":"2024-10-24T21:38:34.372817685Z","level":"INFO","msg":"server is running","addr":{"IP":"127.0.0.1","Port":43157,"Zone":""}}
|
| 7 |
+
{"time":"2024-10-24T21:38:34.373811796Z","level":"INFO","msg":"Will exit if parent process dies.","ppid":337263}
|
| 8 |
+
{"time":"2024-10-24T21:38:34.373817646Z","level":"INFO","msg":"server is running","addr":{"IP":"127.0.0.1","Port":44789,"Zone":""}}
|
| 9 |
+
{"time":"2024-10-24T21:38:34.528156376Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"127.0.0.1:38070"}
|
| 10 |
+
{"time":"2024-10-24T21:38:34.528235848Z","level":"INFO","msg":"connection: ManageConnectionData: new connection created","id":"127.0.0.1:51210"}
|
| 11 |
+
{"time":"2024-10-24T21:38:35.06181568Z","level":"INFO","msg":"handleInformInit: received","streamId":"HCPflat_large_gsrFalse__HCP_FT_cdc7b70b-b155-4483-907d-91148a52e1d7","id":"127.0.0.1:38070"}
|
| 12 |
+
{"time":"2024-10-24T21:38:35.093628892Z","level":"INFO","msg":"handleInformInit: received","streamId":"HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a","id":"127.0.0.1:51210"}
|
| 13 |
+
{"time":"2024-10-24T21:38:35.140878185Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"HCPflat_large_gsrFalse__HCP_FT_cdc7b70b-b155-4483-907d-91148a52e1d7","id":"127.0.0.1:38070"}
|
| 14 |
+
{"time":"2024-10-24T21:38:35.145483394Z","level":"INFO","msg":"handleInformInit: stream started","streamId":"HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a","id":"127.0.0.1:51210"}
|
| 15 |
+
{"time":"2024-10-26T02:19:36.419933572Z","level":"INFO","msg":"handleInformTeardown: server teardown initiated","id":"127.0.0.1:38070"}
|
| 16 |
+
{"time":"2024-10-26T02:19:36.421381993Z","level":"INFO","msg":"server is shutting down"}
|
| 17 |
+
{"time":"2024-10-26T02:19:36.421375983Z","level":"INFO","msg":"connection: Close: initiating connection closure","id":"127.0.0.1:38070"}
|
| 18 |
+
{"time":"2024-10-26T02:19:36.421592698Z","level":"INFO","msg":"connection: Close: connection successfully closed","id":"127.0.0.1:38070"}
|
| 19 |
+
{"time":"2024-10-26T02:19:38.264827227Z","level":"INFO","msg":"handleInformTeardown: server shutdown complete","id":"127.0.0.1:38070"}
|
| 20 |
+
{"time":"2024-10-26T02:19:38.264913688Z","level":"INFO","msg":"connection: ManageConnectionData: connection closed","id":"127.0.0.1:38070"}
|
| 21 |
+
{"time":"2024-10-26T02:19:38.264935459Z","level":"INFO","msg":"server is closed"}
|