ckadirt commited on
Commit
b02f774
·
verified ·
1 Parent(s): d11d0fb

Add files using upload-large-folder tool

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +6 -0
  2. fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/code/_session_history.ipynb +1664 -0
  3. fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/config.yaml +49 -0
  4. fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/output.log +5 -0
  5. fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/wandb-metadata.json +144 -0
  6. fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/files/wandb-summary.json +1 -0
  7. fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug-core.log +12 -0
  8. fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug-internal.log +25 -0
  9. fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/logs/debug.log +59 -0
  10. fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/run-HCPflat_raw_83810.wandb +0 -0
  11. fMRI-foundation-model/src/wandb/run-20241023_032214-HCPflat_raw_83810/tmp/code/_session_history.ipynb +1664 -0
  12. fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/code/src/HCP_downstream_finetune.py +587 -0
  13. fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/output.log +2 -0
  14. fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/requirements.txt +198 -0
  15. fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/files/wandb-metadata.json +131 -0
  16. fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug-core.log +14 -0
  17. fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug-internal.log +11 -0
  18. fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/logs/debug.log +26 -0
  19. fMRI-foundation-model/src/wandb/run-20241023_040830-NSDflat_large_gsrFalse__HCP_FT_83810/run-NSDflat_large_gsrFalse__HCP_FT_83810.wandb +0 -0
  20. fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/files/output.log +0 -0
  21. fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/files/requirements.txt +198 -0
  22. fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/files/wandb-metadata.json +144 -0
  23. fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug-core.log +12 -0
  24. fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug-internal.log +12 -0
  25. fMRI-foundation-model/src/wandb/run-20241023_041122-HCPflat_large_gsrFalse__HCP_FT_7ee35929-85c0-47da-91ab-dac2776c444f/logs/debug.log +39 -0
  26. 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
  27. fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/logs/debug-internal.log +11 -0
  28. fMRI-foundation-model/src/wandb/run-20241023_041223-NSDflat_large_gsrFalse__HCP_FT_9df42194-aaea-4534-8c08-2364d19544b3/logs/debug.log +25 -0
  29. 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
  30. 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
  31. fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/output.log +59 -0
  32. fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/requirements.txt +198 -0
  33. fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/files/wandb-metadata.json +131 -0
  34. fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug-core.log +7 -0
  35. fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug-internal.log +11 -0
  36. fMRI-foundation-model/src/wandb/run-20241023_041325-HCPflat_large_gsrFalse__HCP_FT_62df940a-b42c-469f-b82e-ec86df778532/logs/debug.log +25 -0
  37. 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
  38. 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
  39. fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/output.log +23 -0
  40. fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/requirements.txt +198 -0
  41. fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/files/wandb-metadata.json +131 -0
  42. fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug-core.log +7 -0
  43. fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug-internal.log +11 -0
  44. fMRI-foundation-model/src/wandb/run-20241024_021624-HCPflat_large_gsrFalse__HCP_FT_55946e57-be97-4417-b0d9-e535f9bdc427/logs/debug.log +25 -0
  45. 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
  46. 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
  47. fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/output.log +0 -0
  48. fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/requirements.txt +198 -0
  49. fMRI-foundation-model/src/wandb/run-20241024_213834-HCPflat_large_gsrFalse__HCP_FT_d07870e8-3c53-420c-b734-f32a1c5c3e5a/files/wandb-metadata.json +131 -0
  50. 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"}