diff --git a/.gitattributes b/.gitattributes index a6344aac8c09253b3b630fb776ae94478aa0275b..6fe834d6dbb03ffa11f6c82f1249c70cb181a96d 100644 --- a/.gitattributes +++ b/.gitattributes @@ -33,3 +33,45 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text *.zip filter=lfs diff=lfs merge=lfs -text *.zst filter=lfs diff=lfs merge=lfs -text *tfevents* filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00096.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00095.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00099.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00098.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00097.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00000.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00001.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00002.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00004.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00003.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00005.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00007.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00006.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00008.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00011.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00009.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00010.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00012.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00013.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00014.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00015.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00016.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00017.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00018.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00019.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00020.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00021.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00023.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00024.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00022.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00026.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00025.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00027.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00029.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00028.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00032.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00030.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00031.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00033.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00034.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00035.png filter=lfs diff=lfs merge=lfs -text +checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00037.png filter=lfs diff=lfs merge=lfs -text diff --git a/checkpoints/pretrain_full_90_10_h100/config.yaml b/checkpoints/pretrain_full_90_10_h100/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..15eecd8cdf133469a7ebd7300c4ff78213732ee0 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/config.yaml @@ -0,0 +1,69 @@ +name: pretrain_full_90_10_h100 +notes: null +output_dir: checkpoints/pretrain_full_90_10_h100 +img_size: +- 208 +- 240 +- 208 +patch_size: 8 +mask_ratio: 0.8 +pred_mask_ratio: null +model: mae_vit_large +model_kwargs: + decoding: attn + target_norm: none + no_decode_pos: false + mask_drop_scale: false + class_token: true + reg_tokens: 0 + no_embed_class: false + decoder_depth: 4 + drop_path_rate: 0.0 +datasets: + fomo_train: + url: datasets/FOMO300/wds/shard.{000000..001020}.tar + samples_per_epoch: 153000 + shuffle: true + buffer_size: 8000 + drop_last: true + fomo_val: + url: datasets/FOMO300/wds/shard.{001021..001134}.tar + samples_per_epoch: 17000 + shuffle: false + buffer_size: 1000 + drop_last: true +train_dataset: fomo_train +eval_datasets: +- fomo_val +num_workers: 16 +prefetch_factor: 8 +presend_cuda: false +epochs: 100 +batch_size: 64 +accum_iter: 1 +base_lr: 0.001 +min_lr: 1.0e-06 +warmup_epochs: 10 +weight_decay: 0.05 +betas: +- 0.9 +- 0.95 +clip_grad: 1.0 +amp: true +amp_dtype: bfloat16 +compile: false +ckpt: null +resume: false +auto_resume: true +start_epoch: 0 +checkpoint_period: 5 +max_checkpoints: 5 +r2_sync: null +device: cuda +distributed: false +seed: 7338 +eval_seed: 7338 +debug: false +wandb: true +wandb_entity: null +wandb_project: smri-fm diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00000.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00000.png new file mode 100644 index 0000000000000000000000000000000000000000..0f9c50aac31581a0efcbcf54c9922c6cd23b5426 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00000.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:049cf8178d9e2156e6abcd30309a2c0cb24bf1de8e30965867a973f66d90fb3e +size 167792 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00001.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00001.png new file mode 100644 index 0000000000000000000000000000000000000000..65b2c518cad226fd2b9f5ed8f58649a5f1b463e0 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00001.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93b0e03d74dd349ee07ad64dd9a22daa7ffc3d2faf196dbbf285529013224c91 +size 171060 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00002.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00002.png new file mode 100644 index 0000000000000000000000000000000000000000..bd2147aa3319be011d2b2fbbdc636d23765c872d --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00002.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a19c10fd992b23cd2477ac1ba5c0955ecb825eb04a8141139c95fef66525639c +size 155823 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00003.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00003.png new file mode 100644 index 0000000000000000000000000000000000000000..c7afdf94efb2cb8188e9e72dea49961bd245a422 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00003.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aec6e39208f27baed70a07a7e71fe37b7f780afc7a0fc02cce8b19a5fbf49c06 +size 146495 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00004.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00004.png new file mode 100644 index 0000000000000000000000000000000000000000..453e79e2100104cd7d73e6eef2e05dcee34a1bf4 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00004.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8fb6368a4e9397a0cf77c73aa1a097138daa6946626f7bbc9685f858b85ba171 +size 165604 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00005.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00005.png new file mode 100644 index 0000000000000000000000000000000000000000..455061b0ff8ba9dcc3568b09333bf4087ab08991 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00005.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bd96ed3939119b141c698f95ef9d1d5930d103d2ed981580186d78cf8b4217fa +size 186031 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00006.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00006.png new file mode 100644 index 0000000000000000000000000000000000000000..ee73966ea913061175b3145035eacee8fc8625a0 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00006.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:63884ae4f9f68e71842e79733379740dcd2e6966bb94ce391297c603812f88aa +size 157307 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00007.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00007.png new file mode 100644 index 0000000000000000000000000000000000000000..555b92f361b0c1e41a80d4126fc37a542516116e --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00007.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb0078b35e62a86af3522cad31e933a947612e5ec73c23b2169f1776e6ecf596 +size 156527 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00008.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00008.png new file mode 100644 index 0000000000000000000000000000000000000000..79e95daaa606ce5b66522710ce14c5de3c0a5f27 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00008.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20e69a2a66963a4556a59769d149a2da28d712dc0ad082a8e9d53494ec273627 +size 169808 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00009.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00009.png new file mode 100644 index 0000000000000000000000000000000000000000..40b91c7dfce65d52cb65a11ea116a6dad7fc84a3 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00009.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:626e9eee4f3ec6e6f983e59a3cd208c7055cae83e0654920b513b896639cc4b3 +size 195079 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00010.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00010.png new file mode 100644 index 0000000000000000000000000000000000000000..583aeb117e52678d0ec844bbbd22607750abbaa1 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00010.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:52530ffdfa0a5f3e819b26a2060814b6046e9860377fcd8849fe84ca52d7f14b +size 129017 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00011.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00011.png new file mode 100644 index 0000000000000000000000000000000000000000..f625a9389c5d4868cec166e4987aa2701bb02a48 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00011.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:badbe2ca7c350002e21ac4db097b30291942d89695de24c57f2ea7cad86e7ed4 +size 174192 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00012.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00012.png new file mode 100644 index 0000000000000000000000000000000000000000..6278c048127f301e17df811073cf4e8a84b16aec --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00012.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a1bd8f797e6851097c8cc451eb12eab2b492085625300fe90cd6e511f1280581 +size 212840 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00013.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00013.png new file mode 100644 index 0000000000000000000000000000000000000000..70c395b0ad9ddb47458d0770a573cf2fd4396abe --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00013.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bd7d579bf99a7507eca6a7405038a99d058150012f8a84dd4001c0c8817bf28c +size 161091 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00014.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00014.png new file mode 100644 index 0000000000000000000000000000000000000000..1e425811bd6d204e4e0e8e8ca91fe29b87c55be1 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00014.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e27d6fb6a3fd96c36e5500ba330957c5939422df641bb97e0f952013eb3f85e9 +size 174680 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00015.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00015.png new file mode 100644 index 0000000000000000000000000000000000000000..0820a819a5500691cc38ca3b4144188f3b9a9909 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00015.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d6ec4cd6bc053f81bcbb281d95069430852c6307b38439b5e2fdf26e3d5082b2 +size 179326 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00016.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00016.png new file mode 100644 index 0000000000000000000000000000000000000000..78b281392a19434e961afe2abe86ce56a91008c4 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00016.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:960834d50638ff587f94a221bd0f861e26067c21aa74370800c42e99442eb526 +size 171394 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00017.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00017.png new file mode 100644 index 0000000000000000000000000000000000000000..ced91161399247d3e6bc680e11b863393fea554c --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00017.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:323c2cb77512bbe35fca84809da55d284a38c06febd86ed21f97fcd41282a4bf +size 179727 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00018.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00018.png new file mode 100644 index 0000000000000000000000000000000000000000..9b00478f9978eb50c29e07bde401defb889c3d55 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00018.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:421cd0dba7f54db47f5d12a5e112d8de5cd4ae4f4f11387faab8a66c392ca248 +size 172953 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00019.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00019.png new file mode 100644 index 0000000000000000000000000000000000000000..a378668eff019fd3df870dcca9417dc7cf716544 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00019.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4459a00be9bc4a36168a0effcfff6231091a3065c4846f1d991ecf2bf9a85ac3 +size 188384 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00020.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00020.png new file mode 100644 index 0000000000000000000000000000000000000000..608e3536179dbc54776006ae0afb9a651f681b52 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00020.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:212c328d27cae6cabface81892e00d6d569316ddc652a6156617aea69ec44661 +size 186359 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00021.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00021.png new file mode 100644 index 0000000000000000000000000000000000000000..07de620c1f482a3a3c81dfe33746c404538ce547 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00021.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3cc10385580c13cabf1e49dcf26978f10151ce1da77a4fd78694919591194475 +size 191647 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00022.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00022.png new file mode 100644 index 0000000000000000000000000000000000000000..e538e156004fbc3652b3e5e3eba82f985d067d97 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00022.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2ac767a2dc077a38fd035546e9f3f92ede4a54200a5615014b4b6581591d9448 +size 155783 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00023.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00023.png new file mode 100644 index 0000000000000000000000000000000000000000..be635311320f5f090e219771df01273f953d8950 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00023.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:766bc14165574913b46130a535113a9734d010c45a2a920b08dfa88f340dae6e +size 161609 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00024.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00024.png new file mode 100644 index 0000000000000000000000000000000000000000..6c3736677c706c1fa252464ef306c851bc2c71ad --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00024.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dc9e9935055243a930c62ce9748dfad6a31f1356edbf3575740a0514acc22702 +size 183978 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00025.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00025.png new file mode 100644 index 0000000000000000000000000000000000000000..ac3a7170f596c5f73dae22920495f48d8e40902f --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00025.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ac256ad0ad9e27e24af7016fb1446d32ae8318b5cb184cef67ebbef2884abe9b +size 153735 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00026.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00026.png new file mode 100644 index 0000000000000000000000000000000000000000..e73f3a3384e2e1df3550ec979f17d06bb201371b --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00026.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8fd791a441713393e6826adf429369b12746c9ac6f230683732de50a0ebcb7c6 +size 188305 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00027.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00027.png new file mode 100644 index 0000000000000000000000000000000000000000..6986a9914dc8d6e7fc3f38af1cae7993989e329d --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00027.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:948d69578f96217cbb20723e191d0ed8b9d6e22d36e5392c121fc28aabf8f1c3 +size 170647 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00028.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00028.png new file mode 100644 index 0000000000000000000000000000000000000000..0d9050d17d51b81067f74997a14f6d399e9d8e73 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00028.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c1708de62264b8d5c0b602cdb68bc6bfee15ebd250ca2a62a294d8c1bdbe5c10 +size 147605 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00029.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00029.png new file mode 100644 index 0000000000000000000000000000000000000000..bfc30724efadae3da804e15cfc7805e238d80acb --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00029.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:97881a5f9c8936fa9d165ed216601a6f60401f358802a6ce04bd11b4815acc1c +size 208860 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00030.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00030.png new file mode 100644 index 0000000000000000000000000000000000000000..9e3731705f5945e52f4edbec6369e6c6d1747acf --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00030.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:823b3dc511af86fc4faaaf9b6e742c4542e5135ee03a7961ec184d2b84d0c3df +size 167251 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00031.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00031.png new file mode 100644 index 0000000000000000000000000000000000000000..6ec86e848e59f186c4d3c0f08d085e2a8df1f82c --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00031.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8135344dcad2e9b43ac32a2e23754327ca65b1a1360aa017cc3c8cdc97ef1ca9 +size 171183 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00032.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00032.png new file mode 100644 index 0000000000000000000000000000000000000000..2cd02e9c622728e1b67e4fe8fe5803d32687a203 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00032.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8ce9e81c38e7f2ac954a8b14b7f4976eda6fcda7155a4828170d174835a73355 +size 150804 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00033.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00033.png new file mode 100644 index 0000000000000000000000000000000000000000..2328540ad9ba9039ac8d33c7c803d9d845e271bd --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00033.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5cdc7be4fbda272a3c7f9edd4faba6b52e7d8eea1c90cc52490fbc85bf60b5cd +size 172334 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00034.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00034.png new file mode 100644 index 0000000000000000000000000000000000000000..e15e3e93078e8a8386233df7ccab8a3b1c5a577c --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00034.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8f712bcd47e0212994bcfffbf49b2e18c5c49769ac8f9e39d09f4723248e5eef +size 174517 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00035.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00035.png new file mode 100644 index 0000000000000000000000000000000000000000..0e0ea3d580ab4f90e6c5afc4b6768c89c88567a9 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00035.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0421b22ba16fb23f21f11f338875cb8b55831ae7dbb28ebe903d091b57d71abc +size 171551 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00037.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00037.png new file mode 100644 index 0000000000000000000000000000000000000000..ebb461edbbafb35ddbb390c6966894585daf37b1 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00037.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5979cda6e1ee43acd9ac309552c12270315c99417e6699383596eaa23a1d1c74 +size 188014 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00095.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00095.png new file mode 100644 index 0000000000000000000000000000000000000000..868efe0604e3508c4f9a03511de4e4045533f3c2 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00095.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a840e88a256264ee300d9919553895d75b0eff24635e7605c90d1704d2602d9a +size 160911 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00096.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00096.png new file mode 100644 index 0000000000000000000000000000000000000000..164f8f15883131b2806b13fdce62673f9325dbb1 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00096.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:39f6c0c8215fff05a1cd93275ff5b9a865dfd802cfdea831444a2087ad200540 +size 220150 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00097.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00097.png new file mode 100644 index 0000000000000000000000000000000000000000..0e296cdb933ca513d35632b6fcb1c1c3c92f45e6 --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00097.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1b2f9ad4760703a016c487a673626d35eb7dc24c95888e7204ac4d3b23ebb712 +size 145983 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00098.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00098.png new file mode 100644 index 0000000000000000000000000000000000000000..1d2355562834f869e681ab28ca18dd90dbb697ba --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00098.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1a42d0f98da08c0faf6babd70a619a73629f4b10d486f87a6db103ca54e3b331 +size 236221 diff --git a/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00099.png b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00099.png new file mode 100644 index 0000000000000000000000000000000000000000..dfd62ced2fdbc1c47b79f01466ccad1485091cac --- /dev/null +++ b/checkpoints/pretrain_full_90_10_h100/eval__fomo_val__mask_pred__00099.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:33cc7b0fbb9637feb22cfa6b47156a808deee594a05e9e2b6df4bbbe7f605f43 +size 158882 diff --git a/finetune/fomo_tune_baseline/output/task1/build/Apptainer.def b/finetune/fomo_tune_baseline/output/task1/build/Apptainer.def new file mode 100644 index 0000000000000000000000000000000000000000..98a8bdfc1fefdbffdaed0dedbc7166a3da28b472 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/Apptainer.def @@ -0,0 +1,31 @@ +Bootstrap: docker +From: python:3.11-slim + +# NOT buildable where it sits: the %files paths below are relative to the build cwd, which is the +# staging dir `build.py` writes. Build it with `python -m fomo_tune.build `, not by +# pointing apptainer at this file. +# +# Versions are pinned to the training environment: numpy, scikit-learn and joblib because they +# unpickle `head.joblib`, torch because that is what the checkpoint was written by. + +%files + fomo_tune /app/fomo_tune + smri_mae /app/smri_mae + model /app/model + predict.py /app/predict.py + +%post + pip install --no-cache-dir \ + torch==2.8.0 \ + numpy==2.4.6 \ + nibabel==5.4.2 \ + einops==0.8.2 \ + jaxtyping==0.3.10 \ + timm==1.0.27 \ + huggingface-hub==0.36.2 \ + scikit-learn==1.8.0 \ + joblib==1.5.3 \ + omegaconf==2.3.0 + +%runscript + exec python /app/predict.py "$@" diff --git a/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/README.md b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/README.md new file mode 100644 index 0000000000000000000000000000000000000000..aa7a6cdadfe314f9d1b0fc2d48ef67b85032d3c5 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/README.md @@ -0,0 +1,239 @@ +# fomo_tune + +The five FOMO26 challenge tasks, one script each, tuned independently. + +This is a spinoff of `nanobrain.eval`, which scored every backbone on every task through one fixed +probe. That was the right shape for a benchmark and the wrong shape for a competition: here we care +about one backbone (sMRI MAE) and five scores, and each task wants a different method. **Nothing +here imports `nanobrain.eval`, and it should stay that way** — this package may be shared with +people who won't get the eval suite. + +## Layout + +| File | | +|---|---| +| `datasets.py` | core, **frozen**. One `load_fomo_task()` per task, streaming the challenge zips into an HF dataset. Raw niftis, no resampling — the backbone transform does that. | +| `backbone.py` | core, **frozen**. `load_backbone(ckpt_path) -> (SmriMaeBackbone, SmriMaeTransform)`. Frozen sMRI MAE encoder; the transform canonicalizes to RAS, rescales to 1mm, fits to the pretraining shape, z-scores in a mean-threshold brain mask. | +| `utils.py` | core. `set_seed`, `git_sha`, `setup_logging`. | +| `main_task.py` | shell. One task, end to end. Task 1 is the worked example; copy it. | +| `build.py` + `Apptainer.def` | shell. Package a run dir into the challenge `.sif`. Shared by every task. | + +`datasets.py` and `backbone.py` are settled and their caches are warm. Treat them as read-only: +new work goes in `main_task.py`. If one of them genuinely needs to change, that is a +conversation first, because it invalidates every score already recorded. + +## The pattern + +`main_task1.py` is in three sections, and the split is the point of the whole design. + +**`Task1Method` — the part we tune.** Features, head, hyperparameters, anything that might move +the score. Its interface is: + +```python +method.fit(rows) # rows are dataset records: subject, label, images +method.predict(images) # -> the challenge's output for one subject +method.save(model_dir) # config.yaml + head.joblib +Task1Method.load(model_dir, **overrides) +``` + +**The protocol — fixed.** Pool out-of-fold predictions over all subjects, bootstrap subjects for +the CI. No repeats, no stratification; the bootstrap is the only variance estimate. Splitting is +per-task but fixed within a task — leave-one-out where n is tiny (task 1, n=21), **20-fold** where +it isn't (tasks 3 and 5), which is close enough to LOO without paying for 494 refits. Once a task's +scheme is set, hold it or scores stop being comparable across iterations. That is also why +`cross_validate` seeds its shuffle at 0 rather than from `cfg.seed`: the folds are part of the +protocol, so tuning the run's seed must not silently redraw them. + +**Two entrypoints.** `train` runs the protocol then fits a head on all subjects and saves it; +`predict` is the challenge CLI. Both go through `Method.predict`, which is why every fold +exercises the code the submission will run. + +That last point is the load-bearing one. `predict` is not a wrapper written at packaging time — it +is the same call cross-validation already made once per held-out subject. When you add a task, keep +that property. + +```bash +uv run python -m fomo_tune.main_task1 train modalities=[dwi_b1000,flair] name=task1_dwi_flair +uv run python -m fomo_tune.main_task1 predict --model-dir output/fomo_tune/task1_dwi/model \ + --adc adc.nii.gz --dwi dwi.nii.gz --flair flair.nii.gz --output prob.txt +``` + +`train` takes omegaconf dotlist overrides against the `Config` dataclass at the top of the file. +It writes `config.yaml`, `log.txt`, `metrics.json`, and `model/` into `{output_root}/{name}/`. + +## Status + +Tasks 1, 5 and 3 are drafted and verified. Task 1 is also packaged — its container passes the +challenge validator; 5 and 3 have not been built yet. **Tasks 2 and 4 are tabled** — both are +segmentation, both need `predict` to emit a nifti on the input grid, and neither is worth opening +until the classification and regression tasks are settled. + +All three on `vitl_fomo300`, one H100, wall being the cross-validation loop: + +| run | result | wall | +|---|---|---| +| `task1_dwi`, dwi_b1000, n=21, LOO | AUROC **0.990** [0.944, 1.000] | 25s | +| `task5_t1w`, t1w, n=48, 20-fold | AUROC **0.984** [0.953, 1.000] | 73s | +| `task3_t1w`, t1w, n=494, 20-fold | r **0.962** [0.956, 0.968], MAE **3.71y** [3.45, 3.97] | 260s | + +**Task 3's row is one fold-seed stale.** It was measured before `cross_validate` froze its shuffle +at 0, so it is a 20-fold run with `random_state=4466`. Task 1 (LOO) and task 5 are current. The +re-run is cheap — 260s on a GPU — it just has not been done. Expect a shift of the same order task +5 saw when its folds moved (0.948 → 0.984, i.e. inside the CI but not negligible). + +Task 1's earlier probe sweep got 0.954 [0.861, 1.000] on the same checkpoint +(`experiments/eval_global_0728`), so it roughly reproduces — the gap is LOO vs 5×5 stratified CV, +one interpolation instead of two, and a head selected on AUROC instead of balanced accuracy. + +Two checks worth repeating per task — `.claude/scratch/verify_task1.py` and +`.claude/scratch/verify_task35.py ` do both: +- features are **bit-identical** whether the nifti comes from the HF dataset wrapper or from + `nib.load` off disk, so CV numbers transfer to the container +- the `predict` CLI agrees with the in-process method + +## What changes per task + +Counts and modalities, read from the local zips: + +| Task | n | Inputs | Output | Split | Notes | +|---|---|---|---|---|---| +| 1 infarct | 21 | adc, dwi_b1000, flair (+t2s/swi) | probability | LOO | done | +| 5 polymicrogyria | 48 | t1w | probability | 20-fold | done | +| 3 brain age | 494 | t1w | age in years | 20-fold | done — RidgeCV head, scored by **Pearson r and MAE**, each with its own bootstrap CI | +| 2 meningioma | 23 | dwi_b1000, flair (+t2s/swi) | mask, input grid | — | tabled | +| 4 trigeminal | 40 | t2w | mask, labels 1=nerve 2=vessel | — | tabled | + +Tasks 5 and 3 diverge from task 1 only where that table says. `cross_validate` over a shuffled +`KFold` replaces `leave_one_out`; both take one modality, so `features` loses the +concat-over-modalities loop and `Config` loses `modalities`; the challenge CLI flag is `--t1` for +both, and it is `--t1` for task 3 too even though the file in the zip is `t1w.nii.gz`. + +Task 3 is the first regression, so its `score` loops over the two metrics rather than returning +one, and it drops task 1's guard against bootstrap resamples with fewer than two distinct labels. +The analogous degenerate case for regression is a resample with no spread in `y`, where Pearson r +is undefined rather than merely unstable — at n=494 it does not happen. + +When tasks 2 and 4 come back: `predict` must write a nifti on the input's grid, and the method +needs localized features rather than a pooled vector — `backbone.forward` returns `patch_coords` +in world mm for exactly that. Task 4's label order (1=nerve, 2=vessel) is still a guess and needs +confirming against the challenge data before per-class numbers mean anything. + +## Gotchas + +**Raw niftis are on disk** at `data/fomo_eval/Task_/preprocessed//ses-01/`, which is the +easy way to exercise `predict` on a real file rather than one written out of the dataset: + +```bash +uv run python -m fomo_tune.main_task1 predict \ + --model-dir output/fomo_tune/task1_dwi/model \ + --adc data/fomo_eval/Task_1/preprocessed/sub-20/ses-01/adc.nii.gz \ + --dwi data/fomo_eval/Task_1/preprocessed/sub-20/ses-01/dwi_b1000.nii.gz \ + --flair data/fomo_eval/Task_1/preprocessed/sub-20/ses-01/flair.nii.gz \ + --output /tmp/prob.txt +``` + +Task 5 breaks the naming: `Task_5/preprocessed/sub_01/ses_01/t1.nii.gz` — underscores throughout, +and `t1` not `t1w`. `datasets.py` already handles it; anything you write by hand won't. + +```bash +uv run python -m fomo_tune.main_task5 predict --model-dir output/fomo_tune/task5_t1w/model \ + --t1 data/fomo_eval/Task_5/preprocessed/sub_01/ses_01/t1.nii.gz --output /tmp/prob.txt +``` + +**Volumes are wildly anisotropic.** Task 1's DWI is 0.46×0.46×**5.6**mm, so the transform +upsamples z by 5.6× to reach 1mm iso. Nothing is wrong, but don't read the 1mm grid as real +resolution. + +**The backbone never saw skull or neck.** Pretraining used a SynthSeg brain mask; the transform +substitutes a mean-intensity threshold, which keeps both. Known fidelity gap — see +`.claude/memory/smri-mae-preprocessing-gap.md`. + +**Probabilities are not calibrated.** `LogisticRegressionCV` on ~20 samples × 1024 features shrinks +hard; task 1's out-of-fold probabilities all land in 0.48–0.52 with near-perfect ranking. Fine for +AUROC, which is what the challenge scores, but don't read them as probabilities. Task 5's do span +0–1, which is n=48 rather than n=21 and not evidence of calibration. + +**n is tiny, so the CI is the result.** Task 1's is ~0.06 wide at the top of the range. Most tuning +deltas you chase will be inside it. `.claude/NOTES.md` thread 1 has the longer argument. + +**GPUs need an allocation** — the login node has no driver. See the `gpu-session` skill. + +## Submission + +`build.py` packages a run dir into the `.sif` the challenge wants. One command, taking the run dir +the shipped head was saved into: + +```bash +uv run python -m fomo_tune.build output/fomo_tune/task1_dwi +``` + +It stages `/app`, then builds from there: + +``` +/app/predict.py # shim: calls fomo_tune.main_task predict +/app/model/config.yaml # from the run dir +/app/model/head.joblib # from the run dir +/app/model/backbone.pth # stripped checkpoint, --ckpt-path points here +``` + +**Both `build.py` and `Apptainer.def` are shared across tasks**, which they can be because nothing +in staging or in the dependency list is task-specific. The one thing that does vary is the module +the shim imports, and that comes from `task` in the run's saved config — so a run dir knows which +task it belongs to, and `build.py` never needs telling. + +`predict.py` is **generated at build time** rather than checked in. It is eight lines whose whole +meaning is the container layout staged around it, so there is nowhere outside a container to run +it. This does not weaken the point above about `predict` not being written at packaging time: the +logic still lives in `main_task.py`, exercised once per fold, and the shim only picks the +subcommand and two paths. + +**`Apptainer.def` is not buildable where it sits.** Its `%files` paths are relative to the build +cwd, which is the staging dir. Pointing `apptainer build` at it in the repo fails confusingly; go +through `build.py`. + +The run dir deliberately does *not* carry backbone weights — that checkpoint is 3.9G and would be +copied on every run. `--ckpt-path` overrides what `config.yaml` recorded, so the saved config stays +a faithful record of what trained rather than being rewritten at build time. + +**The staged checkpoint is stripped to `model` and `args`**, which is 3.9G → 1.3G because the rest +is optimizer state inference never reads. `load_backbone` needs no change for this, and `predict` +gives a bit-identical probability either way (0.524739 on `sub-20`, checked on GPU). + +**The base image is `python:3.11-slim`, not a CUDA image.** The PyPI torch wheel *is* the cu128 +build and vendors the whole CUDA userspace as `nvidia-*` packages, so all the container needs from +the host is the driver, which `--nv`/`--nvccli` binds in. That keeps the SIF at 5.3G (4.0G of +image, 1.3G of weights) against roughly double for `pytorch/pytorch` and far more for NGC. +Versions are pinned to the training environment +mostly so `head.joblib` unpickles against the numpy/sklearn that wrote it. + +Apptainer caches the bootstrap layers but **always re-runs `%post`**, so every build re-downloads +~3G of wheels. If that gets annoying, bake a deps-only base SIF and `Bootstrap: localimage` off it. + +### Validating + +`third_party/container-validator` is the challenge's own validator, test niftis included: + +```bash +python third_party/container-validator/container_validator/validate.py \ + --task task1 --sif output/fomo_tune/task1_dwi/task1.sif +``` + +It runs `python /app/predict.py --flair /input/… --adc … --dwi … --swi … --output /output/.txt` +inside an `apptainer instance` with `/input`, `/output` and `/tmp` bound — which is exactly the +shim's contract, so nothing in `predict.py` is guessing at the interface. + +One thing it does that is easy to miss: it takes GPU via `--nvccli` rather than `--nv`, and one of +its tests runs `nvidia-smi -L` **inside** the container. `python:3.11-slim` ships no `nvidia-smi`, +so that test passes only because `--nvccli` injects the host one — a CUDA base image would hide +that dependency rather than remove it. + +**The `task1_dwi` container passes all 20 validator tests**, and `predict` inside it returns +0.524739 on `sub-20`, identical to the same call outside the container. So the packaging is +verified end to end, not just built. + +**Run it on a compute node with apptainer, which as of 2026-08-11 means `n-6`** — `salloc +--nodelist=n-6`. The other nodes fail the validator's preflight. The login node has apptainer but +no driver, and +`predict` there dies inside `can_use_cudnn_attention` — the jagged-SDPA path reaches into CUDA even +when the tensors are on CPU, so a driver-less host fails at the forward pass rather than falling +back. That is the CPU gap worth remembering; it is not a container problem. diff --git a/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/backbone.py b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/backbone.py new file mode 100644 index 0000000000000000000000000000000000000000..4a1275a32e0f9da0963043225e8b3aefcbe8edc6 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/backbone.py @@ -0,0 +1,153 @@ +import inspect + +import nibabel as nib +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange +from torch import Tensor + +import smri_mae.model_mae as models_mae + + +class SmriMaeBackbone(nn.Module): + grid_coords: Tensor + + def __init__(self, encoder: models_mae.MaskedEncoder): + super().__init__() + self.encoder = encoder + self.img_size = self.encoder.patchify.img_size + + grid_size = self.encoder.patchify.grid_size + patch_size = np.array(self.encoder.patchify.patch_size) + grid_coords = rearrange(np.indices(grid_size), "c x y z -> (x y z) c") + grid_coords = grid_coords * patch_size + (patch_size - 1) / 2 + grid_coords = torch.as_tensor(grid_coords, dtype=torch.float32) + self.register_buffer("grid_coords", grid_coords) + + def forward(self, batch: dict[str, Tensor]) -> dict[str, Tensor]: + images = batch["image"] + mask = batch["mask"] + affine = batch["affine"] + + B, C, X, Y, Z = images.shape + assert (X, Y, Z) == self.img_size, f"expected {self.img_size}, got {(X, Y, Z)}" + + _, _, patch_embeds, _, patch_ids, token_mask = self.encoder(images, mask=mask) + + # [B, L, 3] world xyz coords of embeddings + patch_coords = self.grid_coords[patch_ids, :] + rot = affine[:, :3, :3] + trans = affine[:, :3, 3] + patch_coords = patch_coords @ rot.transpose(1, 2) + trans[:, None, :] + + return { + "patch_embeds": patch_embeds, + "patch_ids": patch_ids, + "token_mask": token_mask, + "patch_coords": patch_coords, + } + + +class SmriMaeTransform: + def __init__( + self, + img_size: tuple[int, int, int] = (208, 240, 208), + spacing: tuple[float, float, float] = (1.0, 1.0, 1.0), + ): + self.img_size = img_size + self.spacing = spacing + + def __call__(self, img: nib.Nifti1Image) -> dict[str, Tensor]: + # repack image to handle incomplete hf Nifti interface + img = nib.Nifti1Image(img.dataobj, img.affine, img.header) + img = nib.as_closest_canonical(img) + + data = torch.from_numpy(np.ascontiguousarray(img.get_fdata(dtype=np.float32))) + affine = np.asarray(img.affine) + + spacing = img.header.get_zooms() + if max(abs(s - s_) for s, s_ in zip(spacing, self.spacing)) > 0.05: + data, affine = rescale(data, affine, spacing, self.spacing) + + data, affine = fit_to_shape(data, affine, target_shape=self.img_size) + + # mean threshold, not the SynthSeg mask used in pretraining, so skull and neck stay in + mask = data > data.mean() + brain = data[mask] + mean = brain.mean() + # population std (correction=0) to match the pretraining normalization + std = brain.std(correction=0).clamp_min(1e-6) + data = torch.where(mask, (data - mean) / std, 0.0) + + return { + "image": data.unsqueeze(0), + "mask": mask.unsqueeze(0), + "affine": torch.as_tensor(affine, dtype=torch.float32), + } + + +def rescale( + x: torch.Tensor, + affine: np.ndarray, + spacing: tuple[float, ...], + target_spacing: tuple[float, ...] = (1.0, 1.0, 1.0), +) -> tuple[torch.Tensor, np.ndarray]: + scales = tuple([current / target for current, target in zip(spacing, target_spacing)]) + resampled = F.interpolate(x[None, None], scale_factor=scales, mode="trilinear").squeeze(0, 1) + + # align_corners=False reads output voxel j from input voxel (j + 0.5) / scale - 0.5 + scale = np.asarray(scales, dtype=float) + step = np.diag([*(1 / scale), 1.0]) + step[:3, 3] = 0.5 / scale - 0.5 + return resampled, affine @ step + + +def fit_to_shape( + x: torch.Tensor, affine: np.ndarray, target_shape: tuple[int, ...] +) -> tuple[torch.Tensor, np.ndarray]: + """Centre the volume in `target_shape`, padding the short axes and cropping the long ones.""" + pads = [target - size for size, target in zip(x.shape, target_shape)] + padding = [side for pad in reversed(pads) for side in (pad // 2, pad - pad // 2)] + + # a crop is a negative pad, so output voxel k came from input voxel k - pad // 2 either way + step = np.eye(4) + step[:3, 3] = [-(pad // 2) for pad in pads] + return F.pad(x, padding), affine @ step + + +def resolve_ckpt(ckpt_path: str) -> str: + """A local path for a checkpoint, downloading it if it is an hf://// URI.""" + from huggingface_hub import hf_hub_download + + if ckpt_path.startswith("hf://"): + org, repo, *rest = ckpt_path.removeprefix("hf://").split("/") + return hf_hub_download(f"{org}/{repo}", "/".join(rest)) + + return ckpt_path + + +def load_backbone(ckpt_path: str) -> tuple[SmriMaeBackbone, SmriMaeTransform]: + path = resolve_ckpt(ckpt_path) + ckpt = torch.load(path, map_location="cpu", weights_only=True, mmap=True) + args = ckpt["args"] + + model_fn = models_mae.__dict__[args["model"]] + model: models_mae.MaskedAutoencoderViT = model_fn( + img_size=args["img_size"], + in_chans=args.get("in_chans", 1), + patch_size=args["patch_size"], + # older checkpoints carry training flags the current model_mae no longer takes + **filter_kwargs(models_mae.MaskedAutoencoderViT, args.get("model_kwargs") or {}), + ) + model.load_state_dict(ckpt["model"]) + backbone = SmriMaeBackbone(model.encoder) + transform = SmriMaeTransform(img_size=args["img_size"]) + return backbone, transform + + +def filter_kwargs(func, kwargs): + signature = inspect.signature(func) + kwargs = {k: v for k, v in kwargs.items() if k in signature.parameters} + return kwargs diff --git a/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/datasets.py b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/datasets.py new file mode 100644 index 0000000000000000000000000000000000000000..64c2ac2da798086e88cbb335601ae850c2dad249 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/datasets.py @@ -0,0 +1,205 @@ +import os +import shutil +import tempfile +import zipfile +from collections.abc import Generator +from contextlib import contextmanager +from pathlib import Path + +import fsspec +from datasets import Dataset, Features, Nifti, Value + +FOMO_EVAL_BASE_URL = os.getenv( + "FOMO_EVAL_BASE_URL", + "https://sid.erda.dk/share_redirect/fmeuvo1EdF", +) +FOMO_EVAL_TASK5_URL = os.getenv( + "FOMO_EVAL_TASK5_URL", + "https://huggingface.co/datasets/medarc/smri-fm/resolve/main/fomo_eval/Task_5.zip", +) + + +@contextmanager +def open_zip(url: str) -> Generator[zipfile.ZipFile, None, None]: + """Open a task zip, copying a remote url to a temp file first.""" + with tempfile.TemporaryDirectory() as tmp: + local = Path(url) + if not local.exists(): + local = Path(tmp) / "task.zip" + with fsspec.open(url) as src, local.open("wb") as dst: + shutil.copyfileobj(src, dst) + with zipfile.ZipFile(local) as zf: + yield zf + + +def subject_ids(zf: zipfile.ZipFile) -> list[str]: + return sorted({name.split("/")[2] for name in zf.namelist() if name.endswith(".nii.gz")}) + + +# ---- Task 1: acute infarct (classification; positives also carry a lesion mask) -------- + + +def load_fomo_task1() -> Dataset: + # No 4th modality: it is swi on 16 subjects and t2s on the other 5. + suffixes = ("adc", "dwi_b1000", "flair") + features = Features( + { + "subject": Value("string"), + "label": Value("int32"), + **{suffix: Nifti() for suffix in suffixes}, + } + ) + dataset = Dataset.from_generator( + _fomo_task1_generator, + features=features, + gen_kwargs={"suffixes": suffixes}, + writer_batch_size=16, + ) + return dataset + + +def _fomo_task1_generator(suffixes: tuple[str, ...]): + url = f"{FOMO_EVAL_BASE_URL}/Task_1.zip" + with open_zip(url) as zf: + for sub in subject_ids(zf): + label = int(zf.read(f"Task_1/labels/{sub}/ses-01/label.txt").strip()) + sample = {"subject": sub, "label": label} + for suffix in suffixes: + name = f"Task_1/preprocessed/{sub}/ses-01/{suffix}.nii.gz" + sample[suffix] = {"path": None, "bytes": zf.read(name)} + yield sample + + +# ---- Task 2: meningioma segmentation --------------------------------------------------- + + +def load_fomo_task2() -> Dataset: + # No 4th modality: it is t2s on 15 subjects and swi on the other 8. + suffixes = ("dwi_b1000", "flair") + features = Features( + { + "subject": Value("string"), + **{suffix: Nifti() for suffix in suffixes}, + "seg": Nifti(), + } + ) + dataset = Dataset.from_generator( + _fomo_task2_generator, + features=features, + gen_kwargs={"suffixes": suffixes}, + writer_batch_size=16, + ) + return dataset + + +def _fomo_task2_generator(suffixes: tuple[str, ...]): + url = f"{FOMO_EVAL_BASE_URL}/Task_2.zip" + with open_zip(url) as zf: + for sub in subject_ids(zf): + sample = {"subject": sub} + for suffix in suffixes: + name = f"Task_2/preprocessed/{sub}/ses-01/{suffix}.nii.gz" + sample[suffix] = {"path": None, "bytes": zf.read(name)} + # Seg is on the image grid (shapes match) but its affine differs by up to 0.03mm. + name = f"Task_2/labels/{sub}/ses-01/seg.nii.gz" + sample["seg"] = {"path": None, "bytes": zf.read(name)} + yield sample + + +# ---- Task 3: brain age regression ------------------------------------------------------ + + +def load_fomo_task3() -> Dataset: + features = Features( + { + "subject": Value("string"), + "age": Value("int32"), + "t1w": Nifti(), + } + ) + dataset = Dataset.from_generator( + _fomo_task3_generator, + features=features, + writer_batch_size=16, + ) + return dataset + + +def _fomo_task3_generator(): + url = f"{FOMO_EVAL_BASE_URL}/Task_3.zip" + with open_zip(url) as zf: + for sub in subject_ids(zf): + age = int(zf.read(f"Task_3/labels/{sub}/ses-01/labels.txt").strip()) + image_gz = zf.read(f"Task_3/preprocessed/{sub}/ses-01/t1w.nii.gz") + sample = { + "subject": sub, + "age": age, + "t1w": {"path": None, "bytes": image_gz}, + } + yield sample + + +# ---- Task 4: trigeminal nerve/vessel segmentation -------------------------------------- + + +def load_fomo_task4() -> Dataset: + # Volumes are uncropped 0.5mm near-iso, ~360x512x512; crop before feeding a model. + features = Features( + { + "subject": Value("string"), + "t2w": Nifti(), + "seg": Nifti(), + } + ) + dataset = Dataset.from_generator( + _fomo_task4_generator, + features=features, + writer_batch_size=16, + ) + return dataset + + +def _fomo_task4_generator(): + url = f"{FOMO_EVAL_BASE_URL}/Task_4.zip" + with open_zip(url) as zf: + for sub in subject_ids(zf): + image_gz = zf.read(f"Task_4/preprocessed/{sub}/ses-01/t2w.nii.gz") + seg_gz = zf.read(f"Task_4/labels/{sub}/ses-01/seg.nii.gz") + sample = { + "subject": sub, + "t2w": {"path": None, "bytes": image_gz}, + "seg": {"path": None, "bytes": seg_gz}, + } + yield sample + + +# ---- Task 5: polymicrogyria classification --------------------------------------------- + + +def load_fomo_task5() -> Dataset: + features = Features( + { + "subject": Value("string"), + "label": Value("int32"), + "t1w": Nifti(), + } + ) + dataset = Dataset.from_generator( + _fomo_task5_generator, + features=features, + writer_batch_size=16, + ) + return dataset + + +def _fomo_task5_generator(): + with open_zip(FOMO_EVAL_TASK5_URL) as zf: + for sub in subject_ids(zf): + label = int(zf.read(f"Task_5/labels/{sub}/ses_01/labels.txt").strip()) + image_gz = zf.read(f"Task_5/preprocessed/{sub}/ses_01/t1.nii.gz") + sample = { + "subject": sub, + "label": label, + "t1w": {"path": None, "bytes": image_gz}, + } + yield sample diff --git a/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/main_task1.py b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/main_task1.py new file mode 100644 index 0000000000000000000000000000000000000000..bfbc47d49ceeed6fc43dc52a0a007d933ff86d4b --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/main_task1.py @@ -0,0 +1,253 @@ +"""FOMO task 1: acute infarct classification, scored by AUROC as the challenge scores it. + +`Task1Method` is the part we tune -- features, head, hyperparameters. The protocol below it is +fixed so scores stay comparable across iterations: leave one subject out, pool the out-of-fold +predictions, bootstrap subjects for the CI. + +`train` runs that protocol then fits and saves a head; `predict` is the challenge contract, +modality paths in and one probability out. Both go through `Task1Method.predict`, so every fold +exercises the path the submission will run. +""" + +import argparse +import json +import logging +import time +from dataclasses import dataclass, field +from pathlib import Path + +import joblib +import nibabel as nib +import numpy as np +import torch +from omegaconf import OmegaConf +from sklearn.linear_model import LogisticRegressionCV +from sklearn.metrics import roc_auc_score +from sklearn.pipeline import make_pipeline +from sklearn.preprocessing import StandardScaler + +from fomo_tune.backbone import load_backbone +from fomo_tune.utils import git_sha, set_seed, setup_logging + +logger = logging.getLogger("fomo_tune") + +Images = dict[str, nib.Nifti1Image] + + +@dataclass +class Config: + task: str = "task1" + ckpt_path: str = ( + "/data/mihir-stuff/smri-pretrained/pretrain_full_90_10_h100/checkpoint-last.pth" + ) + modalities: list[str] = field(default_factory=lambda: ["dwi_b1000"]) + output_root: str = "output/fomo_tune" + name: str = "task1" + device: str = "cuda" + seed: int = 4466 + + +# ---- method: the part we tune ----------------------------------------------------------- + + +class Task1Method: + """Frozen sMRI MAE, mean-pooled tokens per modality concatenated, logistic head.""" + + def __init__(self, cfg: Config): + self.cfg = cfg + self.backbone, self.transform = load_backbone(cfg.ckpt_path) + self.device = torch.device(cfg.device) + self.backbone.to(self.device).eval().requires_grad_(False) + self.modalities = list(cfg.modalities) + self.cache: dict[str, np.ndarray] = {} + self.head = None + + @torch.inference_mode() + def features(self, images: Images) -> np.ndarray: + """(D,) per subject. A pure function of the images, so training and inference agree.""" + pooled = [] + for modality in self.modalities: + sample = self.transform(images[modality]) + batch = {key: value[None].to(self.device) for key, value in sample.items()} + + with torch.autocast("cuda", torch.bfloat16, enabled=self.device.type == "cuda"): + out = self.backbone(batch) + + patch_embeds = out["patch_embeds"] + token_mask = out["token_mask"].bool().unsqueeze(-1) + embed = (patch_embeds * token_mask).sum(dim=1) / token_mask.sum(dim=1) + pooled.append(embed[0].float().cpu()) + + return torch.cat(pooled).numpy() + + def cached_features(self, row: dict) -> np.ndarray: + if row["subject"] not in self.cache: + self.cache[row["subject"]] = self.features(row) + return self.cache[row["subject"]] + + def fit(self, rows: list[dict]) -> None: + X = np.stack([self.cached_features(row) for row in rows]) + y = np.array([row["label"] for row in rows]) + + clf = LogisticRegressionCV( + Cs=10, + class_weight="balanced", + scoring="roc_auc", + max_iter=1000, + l1_ratios=(0,), + use_legacy_attributes=False, + ) + self.head = make_pipeline(StandardScaler(), clf) + self.head.fit(X, y) + self.positive = list(self.head.classes_).index(1) + + def predict(self, images: Images) -> float: + """Positive-class probability. Indexes `classes_` rather than assuming column 1, which + would silently score the wrong class if the label order differed.""" + X = self.features(images)[None] + probs = self.head.predict_proba(X)[0] + return float(probs[self.positive]) + + def save(self, model_dir: Path) -> None: + """Everything `load` needs but the backbone weights, which stay wherever `ckpt_path` + points -- a few hundred KB, so a run saves one without copying a 3.7G checkpoint.""" + model_dir.mkdir(parents=True, exist_ok=True) + OmegaConf.save(self.cfg, model_dir / "config.yaml") + joblib.dump({"head": self.head, "positive": self.positive}, model_dir / "head.joblib") + + @classmethod + def load(cls, model_dir: Path, **overrides) -> "Task1Method": + """Rebuild a fitted method from `save`. Overrides are Config fields, for what differs + between here and the container -- the backbone path, the device.""" + cfg = OmegaConf.merge( + OmegaConf.structured(Config), OmegaConf.load(model_dir / "config.yaml"), overrides + ) + method = cls(cfg) + state = joblib.load(model_dir / "head.joblib") + method.head, method.positive = state["head"], state["positive"] + return method + + +# ---- protocol: the part we hold fixed --------------------------------------------------- + +# Every image the task ships. The method picks which of them it wants, as at inference, where +# the challenge hands over all four modalities whether or not a model uses them. +IMAGE_COLS = ("adc", "dwi_b1000", "flair") + + +def leave_one_out(rows: list[dict], method: Task1Method) -> tuple[np.ndarray, np.ndarray]: + """Out-of-fold score for every subject, each predicted by a head fit on the other n-1.""" + y = np.array([row["label"] for row in rows]) + oof = np.zeros(len(rows), dtype=float) + start = time.perf_counter() + for held_out, row in enumerate(rows): + method.fit([r for r in rows if r["subject"] != row["subject"]]) + oof[held_out] = method.predict({key: row[key] for key in IMAGE_COLS}) + logger.info( + f"fold {held_out + 1}/{len(rows)} {row['subject']} " + f"y={y[held_out]} p={oof[held_out]:.3f} ({time.perf_counter() - start:.0f}s)" + ) + return y, oof + + +def score( + y: np.ndarray, oof: np.ndarray, seed: int = 0, n_boot: int = 2000, alpha: float = 0.05 +) -> dict: + """AUROC, the challenge metric, plus a percentile CI resampling subjects with replacement.""" + rng = np.random.default_rng(seed) + samples = [] + for _ in range(n_boot): + rows = rng.integers(0, len(y), size=len(y)) + if len(np.unique(y[rows])) < 2: + continue + samples.append(roc_auc_score(y[rows], oof[rows])) + + low, high = np.percentile(samples, [100 * alpha / 2, 100 * (1 - alpha / 2)]) + return { + "auroc": float(roc_auc_score(y, oof)), + "auroc_ci_low": float(low), + "auroc_ci_high": float(high), + } + + +# ---- entrypoints ------------------------------------------------------------------------ + + +def train(args: argparse.Namespace) -> None: + # imported here, not at the top, so the container needs no dataset stack to run `predict` + from fomo_tune.datasets import load_fomo_task1 + + cfg = OmegaConf.merge(OmegaConf.structured(Config), OmegaConf.from_dotlist(args.overrides)) + run_dir = Path(cfg.output_root) / cfg.name + run_dir.mkdir(parents=True, exist_ok=True) + + setup_logging(run_dir) + set_seed(cfg.seed) + logger.info(f"run {cfg.name} (git {git_sha()})") + logger.info(f"config:\n{OmegaConf.to_yaml(cfg).rstrip()}") + OmegaConf.save(cfg, run_dir / "config.yaml") + + # decoded once: leave-one-out revisits every subject n times, and the niftis are small + rows = list(load_fomo_task1()) + logger.info(f"dataset: {len(rows)} subjects, {sum(r['label'] for r in rows)} positive") + + method = Task1Method(cfg) + start = time.perf_counter() + y, oof = leave_one_out(rows, method) + run_time = time.perf_counter() - start + summary = score(y, oof) + + # the shipped head sees all n subjects, so it is not any of the models scored above + method.fit(rows) + method.save(run_dir / "model") + + record = {"name": cfg.name, **summary, "run_time": round(run_time, 1)} + (run_dir / "metrics.json").write_text(json.dumps(record) + "\n") + scores = " ".join(f"{k}={v:.4f}" for k, v in summary.items()) + logger.info(f"result: {scores} ({run_time:.0f}s)") + + +def predict(args: argparse.Namespace) -> None: + """The challenge contract: modality paths in, one probability written to `--output`. + + `/app/predict.py` in the container is a shim over this, so what scores the submission is the + code leave-one-out already ran, not something generated at build time. + """ + overrides = {"device": args.device} + if args.ckpt_path: + overrides["ckpt_path"] = args.ckpt_path + method = Task1Method.load(args.model_dir, **overrides) + + # every image the challenge hands over, as in `leave_one_out`; the method takes what it uses + paths = {"adc": args.adc, "dwi_b1000": args.dwi, "flair": args.flair} + probability = method.predict({key: nib.load(path) for key, path in paths.items()}) + + args.output.write_text(f"{probability:.6f}\n") + + +def main() -> None: + parser = argparse.ArgumentParser() + modes = parser.add_subparsers(required=True) + + train_parser = modes.add_parser("train", help="leave-one-out over the task, then fit and save") + train_parser.add_argument("overrides", nargs="*", help="config overrides, e.g. device=cpu") + train_parser.set_defaults(run=train) + + predict_parser = modes.add_parser("predict", help="one subject, one probability") + for flag in ("--flair", "--adc", "--dwi"): + predict_parser.add_argument(flag, type=Path, required=True) + # accepted and ignored: the 4th modality is swi on some subjects and t2s on others + for flag in ("--t2s", "--swi"): + predict_parser.add_argument(flag, type=Path) + predict_parser.add_argument("--output", type=Path, required=True) + predict_parser.add_argument("--model-dir", type=Path, default=Path("/app/model")) + predict_parser.add_argument("--ckpt-path", help="overrides the trained config's backbone path") + predict_parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") + predict_parser.set_defaults(run=predict) + + args = parser.parse_args() + args.run(args) + + +if __name__ == "__main__": + main() diff --git a/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/main_task3.py b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/main_task3.py new file mode 100644 index 0000000000000000000000000000000000000000..02db5413c10a0d44eb248e7f2739a436c45806c2 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/main_task3.py @@ -0,0 +1,241 @@ +"""FOMO task 3: brain age regression, scored by Pearson r and MAE as the challenge scores it. + +`Task3Method` is the part we tune -- features, head, hyperparameters. The protocol below it is +fixed so scores stay comparable across iterations: 20-fold over the 494 subjects, pool the +out-of-fold predictions, bootstrap subjects for the CI. + +`train` runs that protocol then fits and saves a head; `predict` is the challenge contract, one t1 +path in and one age out. Both go through `Task3Method.predict`, so every fold exercises the path +the submission will run. +""" + +import argparse +import json +import logging +import time +from dataclasses import dataclass +from pathlib import Path + +import joblib +import nibabel as nib +import numpy as np +import torch +from omegaconf import OmegaConf +from sklearn.linear_model import RidgeCV +from sklearn.model_selection import KFold +from sklearn.pipeline import make_pipeline +from sklearn.preprocessing import StandardScaler + +from fomo_tune.backbone import load_backbone +from fomo_tune.utils import git_sha, set_seed, setup_logging + +logger = logging.getLogger("fomo_tune") + +Images = dict[str, nib.Nifti1Image] + + +@dataclass +class Config: + task: str = "task3" + ckpt_path: str = ( + "/data/mihir-stuff/smri-pretrained/pretrain_full_90_10_h100/checkpoint-last.pth" + ) + output_root: str = "output/fomo_tune" + name: str = "task3" + device: str = "cuda" + seed: int = 4466 + + +# ---- method: the part we tune ----------------------------------------------------------- + + +class Task3Method: + """Frozen sMRI MAE, mean-pooled tokens over the t1w, ridge head.""" + + def __init__(self, cfg: Config): + self.cfg = cfg + self.backbone, self.transform = load_backbone(cfg.ckpt_path) + self.device = torch.device(cfg.device) + self.backbone.to(self.device).eval().requires_grad_(False) + self.cache: dict[str, np.ndarray] = {} + self.head = None + + @torch.inference_mode() + def features(self, images: Images) -> np.ndarray: + """(D,) per subject. A pure function of the images, so training and inference agree.""" + sample = self.transform(images["t1w"]) + batch = {key: value[None].to(self.device) for key, value in sample.items()} + + with torch.autocast("cuda", torch.bfloat16, enabled=self.device.type == "cuda"): + out = self.backbone(batch) + + patch_embeds = out["patch_embeds"] + token_mask = out["token_mask"].bool().unsqueeze(-1) + embed = (patch_embeds * token_mask).sum(dim=1) / token_mask.sum(dim=1) + return embed[0].float().cpu().numpy() + + def cached_features(self, row: dict) -> np.ndarray: + if row["subject"] not in self.cache: + self.cache[row["subject"]] = self.features(row) + return self.cache[row["subject"]] + + def fit(self, rows: list[dict]) -> None: + X = np.stack([self.cached_features(row) for row in rows]) + y = np.array([row["age"] for row in rows], dtype=float) + + # RidgeCV picks alpha by its own efficient leave-one-out, so the fold's own split is + # never touched by model selection + self.head = make_pipeline(StandardScaler(), RidgeCV(alphas=np.logspace(-3, 6, 19))) + self.head.fit(X, y) + + def predict(self, images: Images) -> float: + """Age in years.""" + X = self.features(images)[None] + return float(self.head.predict(X)[0]) + + def save(self, model_dir: Path) -> None: + """Everything `load` needs but the backbone weights, which stay wherever `ckpt_path` + points -- a few hundred KB, so a run saves one without copying a 3.7G checkpoint.""" + model_dir.mkdir(parents=True, exist_ok=True) + OmegaConf.save(self.cfg, model_dir / "config.yaml") + joblib.dump(self.head, model_dir / "head.joblib") + + @classmethod + def load(cls, model_dir: Path, **overrides) -> "Task3Method": + """Rebuild a fitted method from `save`. Overrides are Config fields, for what differs + between here and the container -- the backbone path, the device.""" + cfg = OmegaConf.merge( + OmegaConf.structured(Config), OmegaConf.load(model_dir / "config.yaml"), overrides + ) + method = cls(cfg) + method.head = joblib.load(model_dir / "head.joblib") + return method + + +# ---- protocol: the part we hold fixed --------------------------------------------------- + +# Every image the task ships. The method picks which of them it wants, as at inference, where the +# challenge hands over the modalities whether or not a model uses them. +IMAGE_COLS = ("t1w",) + + +def cross_validate( + rows: list[dict], method: Task3Method, seed: int = 0, n_folds: int = 20 +) -> tuple[np.ndarray, np.ndarray]: + """Out-of-fold age for every subject, each predicted by a head fit on the other folds.""" + y = np.array([row["age"] for row in rows], dtype=float) + oof = np.zeros(len(rows), dtype=float) + folds = KFold(n_splits=n_folds, shuffle=True, random_state=seed) + start = time.perf_counter() + for fold, (train, test) in enumerate(folds.split(rows)): + method.fit([rows[i] for i in train]) + for i in test: + oof[i] = method.predict({key: rows[i][key] for key in IMAGE_COLS}) + logger.info( + f"fold {fold + 1}/{n_folds} n={len(test)} mae={np.abs(y[test] - oof[test]).mean():.2f} " + f"({time.perf_counter() - start:.0f}s)" + ) + return y, oof + + +def metrics(y: np.ndarray, oof: np.ndarray) -> dict: + return { + "pearson_r": float(np.corrcoef(y, oof)[0, 1]), + "mae": float(np.abs(y - oof).mean()), + } + + +def score( + y: np.ndarray, oof: np.ndarray, seed: int = 0, n_boot: int = 2000, alpha: float = 0.05 +) -> dict: + """Both challenge metrics, each with a percentile CI resampling subjects with replacement.""" + rng = np.random.default_rng(seed) + resamples = rng.integers(0, len(y), size=(n_boot, len(y))) + + summary = {} + for name, point in metrics(y, oof).items(): + samples = [metrics(y[rows], oof[rows])[name] for rows in resamples] + low, high = np.percentile(samples, [100 * alpha / 2, 100 * (1 - alpha / 2)]) + summary[name] = point + summary[f"{name}_ci_low"] = float(low) + summary[f"{name}_ci_high"] = float(high) + return summary + + +# ---- entrypoints ------------------------------------------------------------------------ + + +def train(args: argparse.Namespace) -> None: + # imported here, not at the top, so the container needs no dataset stack to run `predict` + from fomo_tune.datasets import load_fomo_task3 + + cfg = OmegaConf.merge(OmegaConf.structured(Config), OmegaConf.from_dotlist(args.overrides)) + run_dir = Path(cfg.output_root) / cfg.name + run_dir.mkdir(parents=True, exist_ok=True) + + setup_logging(run_dir) + set_seed(cfg.seed) + logger.info(f"run {cfg.name} (git {git_sha()})") + logger.info(f"config:\n{OmegaConf.to_yaml(cfg).rstrip()}") + OmegaConf.save(cfg, run_dir / "config.yaml") + + rows = list(load_fomo_task3()) + ages = np.array([row["age"] for row in rows]) + logger.info( + f"dataset: {len(rows)} subjects, age {ages.min()}-{ages.max()} mean {ages.mean():.1f}" + ) + + method = Task3Method(cfg) + start = time.perf_counter() + y, oof = cross_validate(rows, method) + run_time = time.perf_counter() - start + summary = score(y, oof) + + # the shipped head sees all n subjects, so it is not any of the models scored above + method.fit(rows) + method.save(run_dir / "model") + + record = {"name": cfg.name, **summary, "run_time": round(run_time, 1)} + (run_dir / "metrics.json").write_text(json.dumps(record) + "\n") + scores = " ".join(f"{k}={v:.4f}" for k, v in summary.items()) + logger.info(f"result: {scores} ({run_time:.0f}s)") + + +def predict(args: argparse.Namespace) -> None: + """The challenge contract: a t1 path in, one age written to `--output`. + + `/app/predict.py` in the container is a shim over this, so what scores the submission is the + code cross-validation already ran, not something generated at build time. + """ + overrides = {"device": args.device} + if args.ckpt_path: + overrides["ckpt_path"] = args.ckpt_path + method = Task3Method.load(args.model_dir, **overrides) + + age = method.predict({"t1w": nib.load(args.t1)}) + + args.output.write_text(f"{age:.6f}\n") + + +def main() -> None: + parser = argparse.ArgumentParser() + modes = parser.add_subparsers(required=True) + + train_parser = modes.add_parser("train", help="cross-validate over the task, then fit and save") + train_parser.add_argument("overrides", nargs="*", help="config overrides, e.g. device=cpu") + train_parser.set_defaults(run=train) + + predict_parser = modes.add_parser("predict", help="one subject, one age in years") + predict_parser.add_argument("--t1", type=Path, required=True) + predict_parser.add_argument("--output", type=Path, required=True) + predict_parser.add_argument("--model-dir", type=Path, default=Path("/app/model")) + predict_parser.add_argument("--ckpt-path", help="overrides the trained config's backbone path") + predict_parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") + predict_parser.set_defaults(run=predict) + + args = parser.parse_args() + args.run(args) + + +if __name__ == "__main__": + main() diff --git a/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/main_task5.py b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/main_task5.py new file mode 100644 index 0000000000000000000000000000000000000000..d8cccbecf130497c483b3a4b7665e835ca428873 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/main_task5.py @@ -0,0 +1,245 @@ +"""FOMO task 5: polymicrogyria classification, scored by AUROC as the challenge scores it. + +`Task5Method` is the part we tune -- features, head, hyperparameters. The protocol below it is +fixed so scores stay comparable across iterations: 20-fold over the 48 subjects, pool the +out-of-fold predictions, bootstrap subjects for the CI. + +`train` runs that protocol then fits and saves a head; `predict` is the challenge contract, one t1 +path in and one probability out. Both go through `Task5Method.predict`, so every fold exercises +the path the submission will run. +""" + +import argparse +import json +import logging +import time +from dataclasses import dataclass +from pathlib import Path + +import joblib +import nibabel as nib +import numpy as np +import torch +from omegaconf import OmegaConf +from sklearn.linear_model import LogisticRegressionCV +from sklearn.metrics import roc_auc_score +from sklearn.model_selection import KFold +from sklearn.pipeline import make_pipeline +from sklearn.preprocessing import StandardScaler + +from fomo_tune.backbone import load_backbone +from fomo_tune.utils import git_sha, set_seed, setup_logging + +logger = logging.getLogger("fomo_tune") + +Images = dict[str, nib.Nifti1Image] + + +@dataclass +class Config: + task: str = "task5" + ckpt_path: str = ( + "/data/mihir-stuff/smri-pretrained/pretrain_full_90_10_h100/checkpoint-last.pth" + ) + output_root: str = "output/fomo_tune" + name: str = "task5" + device: str = "cuda" + seed: int = 4466 + + +# ---- method: the part we tune ----------------------------------------------------------- + + +class Task5Method: + """Frozen sMRI MAE, mean-pooled tokens over the t1w, logistic head.""" + + def __init__(self, cfg: Config): + self.cfg = cfg + self.backbone, self.transform = load_backbone(cfg.ckpt_path) + self.device = torch.device(cfg.device) + self.backbone.to(self.device).eval().requires_grad_(False) + self.cache: dict[str, np.ndarray] = {} + self.head = None + + @torch.inference_mode() + def features(self, images: Images) -> np.ndarray: + """(D,) per subject. A pure function of the images, so training and inference agree.""" + sample = self.transform(images["t1w"]) + batch = {key: value[None].to(self.device) for key, value in sample.items()} + + with torch.autocast("cuda", torch.bfloat16, enabled=self.device.type == "cuda"): + out = self.backbone(batch) + + patch_embeds = out["patch_embeds"] + token_mask = out["token_mask"].bool().unsqueeze(-1) + embed = (patch_embeds * token_mask).sum(dim=1) / token_mask.sum(dim=1) + return embed[0].float().cpu().numpy() + + def cached_features(self, row: dict) -> np.ndarray: + if row["subject"] not in self.cache: + self.cache[row["subject"]] = self.features(row) + return self.cache[row["subject"]] + + def fit(self, rows: list[dict]) -> None: + X = np.stack([self.cached_features(row) for row in rows]) + y = np.array([row["label"] for row in rows]) + + clf = LogisticRegressionCV( + Cs=10, + class_weight="balanced", + scoring="roc_auc", + max_iter=1000, + l1_ratios=(0,), + use_legacy_attributes=False, + ) + self.head = make_pipeline(StandardScaler(), clf) + self.head.fit(X, y) + self.positive = list(self.head.classes_).index(1) + + def predict(self, images: Images) -> float: + """Positive-class probability. Indexes `classes_` rather than assuming column 1, which + would silently score the wrong class if the label order differed.""" + X = self.features(images)[None] + probs = self.head.predict_proba(X)[0] + return float(probs[self.positive]) + + def save(self, model_dir: Path) -> None: + """Everything `load` needs but the backbone weights, which stay wherever `ckpt_path` + points -- a few hundred KB, so a run saves one without copying a 3.7G checkpoint.""" + model_dir.mkdir(parents=True, exist_ok=True) + OmegaConf.save(self.cfg, model_dir / "config.yaml") + joblib.dump({"head": self.head, "positive": self.positive}, model_dir / "head.joblib") + + @classmethod + def load(cls, model_dir: Path, **overrides) -> "Task5Method": + """Rebuild a fitted method from `save`. Overrides are Config fields, for what differs + between here and the container -- the backbone path, the device.""" + cfg = OmegaConf.merge( + OmegaConf.structured(Config), OmegaConf.load(model_dir / "config.yaml"), overrides + ) + method = cls(cfg) + state = joblib.load(model_dir / "head.joblib") + method.head, method.positive = state["head"], state["positive"] + return method + + +# ---- protocol: the part we hold fixed --------------------------------------------------- + +# Every image the task ships. The method picks which of them it wants, as at inference, where the +# challenge hands over the modalities whether or not a model uses them. +IMAGE_COLS = ("t1w",) + + +def cross_validate( + rows: list[dict], method: Task5Method, seed: int = 0, n_folds: int = 20 +) -> tuple[np.ndarray, np.ndarray]: + """Out-of-fold score for every subject, each predicted by a head fit on the other folds.""" + y = np.array([row["label"] for row in rows]) + oof = np.zeros(len(rows), dtype=float) + folds = KFold(n_splits=n_folds, shuffle=True, random_state=seed) + start = time.perf_counter() + for fold, (train, test) in enumerate(folds.split(rows)): + method.fit([rows[i] for i in train]) + for i in test: + oof[i] = method.predict({key: rows[i][key] for key in IMAGE_COLS}) + logger.info( + f"fold {fold + 1}/{n_folds} n={len(test)} y={y[test]} " + f"p={np.round(oof[test], 3)} ({time.perf_counter() - start:.0f}s)" + ) + return y, oof + + +def score( + y: np.ndarray, oof: np.ndarray, seed: int = 0, n_boot: int = 2000, alpha: float = 0.05 +) -> dict: + """AUROC, the challenge metric, plus a percentile CI resampling subjects with replacement.""" + rng = np.random.default_rng(seed) + samples = [] + for _ in range(n_boot): + rows = rng.integers(0, len(y), size=len(y)) + if len(np.unique(y[rows])) < 2: + continue + samples.append(roc_auc_score(y[rows], oof[rows])) + + low, high = np.percentile(samples, [100 * alpha / 2, 100 * (1 - alpha / 2)]) + return { + "auroc": float(roc_auc_score(y, oof)), + "auroc_ci_low": float(low), + "auroc_ci_high": float(high), + } + + +# ---- entrypoints ------------------------------------------------------------------------ + + +def train(args: argparse.Namespace) -> None: + # imported here, not at the top, so the container needs no dataset stack to run `predict` + from fomo_tune.datasets import load_fomo_task5 + + cfg = OmegaConf.merge(OmegaConf.structured(Config), OmegaConf.from_dotlist(args.overrides)) + run_dir = Path(cfg.output_root) / cfg.name + run_dir.mkdir(parents=True, exist_ok=True) + + setup_logging(run_dir) + set_seed(cfg.seed) + logger.info(f"run {cfg.name} (git {git_sha()})") + logger.info(f"config:\n{OmegaConf.to_yaml(cfg).rstrip()}") + OmegaConf.save(cfg, run_dir / "config.yaml") + + rows = list(load_fomo_task5()) + logger.info(f"dataset: {len(rows)} subjects, {sum(r['label'] for r in rows)} positive") + + method = Task5Method(cfg) + start = time.perf_counter() + y, oof = cross_validate(rows, method) + run_time = time.perf_counter() - start + summary = score(y, oof) + + # the shipped head sees all n subjects, so it is not any of the models scored above + method.fit(rows) + method.save(run_dir / "model") + + record = {"name": cfg.name, **summary, "run_time": round(run_time, 1)} + (run_dir / "metrics.json").write_text(json.dumps(record) + "\n") + scores = " ".join(f"{k}={v:.4f}" for k, v in summary.items()) + logger.info(f"result: {scores} ({run_time:.0f}s)") + + +def predict(args: argparse.Namespace) -> None: + """The challenge contract: a t1 path in, one probability written to `--output`. + + `/app/predict.py` in the container is a shim over this, so what scores the submission is the + code cross-validation already ran, not something generated at build time. + """ + overrides = {"device": args.device} + if args.ckpt_path: + overrides["ckpt_path"] = args.ckpt_path + method = Task5Method.load(args.model_dir, **overrides) + + probability = method.predict({"t1w": nib.load(args.t1)}) + + args.output.write_text(f"{probability:.6f}\n") + + +def main() -> None: + parser = argparse.ArgumentParser() + modes = parser.add_subparsers(required=True) + + train_parser = modes.add_parser("train", help="cross-validate over the task, then fit and save") + train_parser.add_argument("overrides", nargs="*", help="config overrides, e.g. device=cpu") + train_parser.set_defaults(run=train) + + predict_parser = modes.add_parser("predict", help="one subject, one probability") + predict_parser.add_argument("--t1", type=Path, required=True) + predict_parser.add_argument("--output", type=Path, required=True) + predict_parser.add_argument("--model-dir", type=Path, default=Path("/app/model")) + predict_parser.add_argument("--ckpt-path", help="overrides the trained config's backbone path") + predict_parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") + predict_parser.set_defaults(run=predict) + + args = parser.parse_args() + args.run(args) + + +if __name__ == "__main__": + main() diff --git a/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/utils.py b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..57a5c6d0396d8811e763254a987b331cf41ee1f1 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/fomo_tune/utils.py @@ -0,0 +1,33 @@ +import logging +import random +import subprocess +import sys +from pathlib import Path + +import numpy as np +import torch + +logger = logging.getLogger("fomo_tune") + + +def set_seed(seed: int) -> None: + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + + +def git_sha() -> str: + kwargs = dict(cwd=Path(__file__).parent, capture_output=True, text=True, check=True) + sha = subprocess.run(["git", "rev-parse", "--short", "HEAD"], **kwargs).stdout.strip() + dirty = subprocess.run(["git", "status", "--porcelain", "-uno"], **kwargs).stdout.strip() + return f"{sha}-dirty" if dirty else sha + + +def setup_logging(run_dir: Path) -> None: + handlers = [logging.StreamHandler(sys.stdout), logging.FileHandler(run_dir / "log.txt")] + logger.setLevel(logging.INFO) + logger.handlers.clear() + for handler in handlers: + handler.setFormatter(logging.Formatter("%(asctime)s %(message)s", datefmt="%H:%M:%S")) + logger.addHandler(handler) + logger.propagate = False diff --git a/finetune/fomo_tune_baseline/output/task1/build/model/backbone.pth b/finetune/fomo_tune_baseline/output/task1/build/model/backbone.pth new file mode 100644 index 0000000000000000000000000000000000000000..793b42f5cea0b055534c29cd18c21f3200873219 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/model/backbone.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aacf2582e9464e7fd4bd3d5ed41a5b9aa770e5b04a24b2a82c1c4ba5fb5f6e7a +size 1389681124 diff --git a/finetune/fomo_tune_baseline/output/task1/build/model/config.yaml b/finetune/fomo_tune_baseline/output/task1/build/model/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..5dca875152bb8d3adf4bf758bd4681bf19e1d83e --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/model/config.yaml @@ -0,0 +1,8 @@ +task: task1 +ckpt_path: /data/mihir-stuff/smri-pretrained/pretrain_full_90_10_h100/checkpoint-last.pth +modalities: +- dwi_b1000 +output_root: experiments/fomo_tune_baseline/output +name: task1 +device: cuda +seed: 4466 diff --git a/finetune/fomo_tune_baseline/output/task1/build/model/head.joblib b/finetune/fomo_tune_baseline/output/task1/build/model/head.joblib new file mode 100644 index 0000000000000000000000000000000000000000..90ed37cfc3dfd4a13e1f754990808e15c9b1f7e0 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/model/head.joblib @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f4a4826f5ea9f3f74279b798caf6a5970a3c63b02b70c4d0828ee5325d29e0dd +size 445215 diff --git a/finetune/fomo_tune_baseline/output/task1/build/predict.py b/finetune/fomo_tune_baseline/output/task1/build/predict.py new file mode 100644 index 0000000000000000000000000000000000000000..6ecbf03e1ba0f3fba34984bd334f8d44befdc132 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/predict.py @@ -0,0 +1,16 @@ + +import sys + +from fomo_tune.main_task1 import main + +sys.argv = [ + sys.argv[0], + "predict", + *sys.argv[1:], + "--model-dir", + "/app/model", + "--ckpt-path", + "/app/model/backbone.pth", +] + +main() diff --git a/finetune/fomo_tune_baseline/output/task1/build/smri_mae/config/default_pretrain.yaml b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/config/default_pretrain.yaml new file mode 100644 index 0000000000000000000000000000000000000000..c16290b1e2c24e4d206f7e32a7b9c2ee5266067e --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/config/default_pretrain.yaml @@ -0,0 +1,98 @@ +# Name of the run. Used for output directory suffix and wandb. +name: pretrain + +# Description of the run. Goes in wandb notes. +notes: null + +# Root output directory. +# The run writes to checkpoints/ when name is set. +output_dir: checkpoints + +# Standard 3D structural MRI volume size. +img_size: [208, 240, 208] +patch_size: 8 + +# Masking. +mask_ratio: 0.80 +pred_mask_ratio: null +pad_to_multiple: 32 + +# Model. +model: mae_vit_large +model_kwargs: + # target normalization: null/none, global, slice, or patch. + target_norm: none + + no_decode_pos: false + mask_drop_scale: false + + class_token: true + reg_tokens: 0 + no_embed_class: false + + decoder_depth: 4 + drop_path_rate: 0.0 + +# Datasets. +datasets: + fomo_train: + url: datasets/FOMO_with_dwi/shard.{000000..001620}.tar + samples_per_epoch: 243200 + shuffle: true + buffer_size: 8000 + drop_last: true + + fomo_val: + url: datasets/FOMO_with_dwi/shard.{001621..001800}.tar + samples_per_epoch: 26880 + shuffle: false + buffer_size: 1000 + drop_last: true + +train_dataset: fomo_train +eval_datasets: + - fomo_val + +# Data loader. +num_workers: 4 +prefetch_factor: 2 +presend_cuda: true + +# Optimization. +epochs: 100 +batch_size: 64 +accum_iter: 1 + +base_lr: 0.001 +min_lr: 1e-6 +warmup_epochs: 10 +weight_decay: 0.05 +betas: [0.9, 0.95] +clip_grad: 1.0 + +amp: true +amp_dtype: bfloat16 + +# Checkpointing. +ckpt: null +resume: false +auto_resume: true +start_epoch: 0 +checkpoint_period: 10 +max_checkpoints: 5 + +# Evaluation. +eval_period: 10 + +# Sync checkpoints to an R2 bucket using the AWS CLI. Set to a URL to enable. +r2_sync: null + +device: cuda +distributed: false +seed: 7338 +eval_seed: 7338 +debug: false + +wandb: false +wandb_entity: null +wandb_project: smri-fm diff --git a/finetune/fomo_tune_baseline/output/task1/build/smri_mae/main_pretrain.py b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/main_pretrain.py new file mode 100644 index 0000000000000000000000000000000000000000..5470551d5b56bf4653fad4c2a23888c417c0c6e8 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/main_pretrain.py @@ -0,0 +1,486 @@ +# Copyright (c) Sophont, Inc +# This source code is licensed under the Apache License, Version 2.0 +# +# References: +# deit: https://github.com/facebookresearch/deit/blob/main/main.py +# capi: https://github.com/facebookresearch/capi/blob/main/train_capi.py + +import argparse +import datetime +import json +import math +import random +import subprocess +import time +from contextlib import nullcontext +from functools import partial +from itertools import islice +from pathlib import Path +from typing import Iterable, Sequence + +import torch +import torch.nn as nn +import wandb +import webdataset as wds +from omegaconf import DictConfig, OmegaConf +from PIL import Image + +from matplotlib import pyplot as plt +from torch import Tensor + +import data.mri_data as mri_data +import smri_mae.model_mae as models_mae +import smri_mae.utils as ut +import smri_mae.visualization as vis + +DEFAULT_CONFIG = Path(__file__).parent / "config/default_pretrain.yaml" + +MODELS_DICT = models_mae.__dict__ + + +def main(args: DictConfig): + # setup + ut.init_distributed_mode(args) + global_rank = ut.get_rank() + is_master = global_rank == 0 + world_size = ut.get_world_size() + device = torch.device(args.device) + ut.configure_flash_sdpa() + ut.random_seed(args.seed, rank=global_rank) + + if args.name and not args.output_dir.endswith(args.name): + args.output_dir = f"{args.output_dir}/{args.name}" + output_dir = Path(args.output_dir) + + if is_master: + output_dir.mkdir(parents=True, exist_ok=True) + out_cfg_path = output_dir / "config.yaml" + if out_cfg_path.exists(): + prev_cfg = OmegaConf.load(out_cfg_path) + assert args == prev_cfg, "current config doesn't match previous config" + else: + OmegaConf.save(args, out_cfg_path) + + if args.wandb: + wandb.init( + entity=args.wandb_entity, + project=args.wandb_project, + name=args.name, + notes=args.notes, + config=OmegaConf.to_container(args), + ) + + ut.setup_for_distributed(log_path=output_dir / "log.txt") + + print("pretraining 3D ViTMAE") + print(f"start: {datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')}") + print(f"cwd: {Path.cwd()}") + print(ut.get_sha()) + print("config:", OmegaConf.to_yaml(args), sep="\n") + + # data loaders + train_loader, eval_loaders = create_data_loaders(args) + + # model + model = MODELS_DICT[args.model]( + img_size=args.img_size, + in_chans=args.get("in_chans", 1), + patch_size=args.patch_size, + **(args.get("model_kwargs") or {}), + ) + model.to(device) + print("model:", model, sep="\n") + num_params = sum(p.numel() for p in model.parameters() if p.requires_grad) + print(f"num params: {num_params / 1e6:.1f}M") + + model_without_ddp = model + if args.distributed: + model = torch.nn.parallel.DistributedDataParallel( + model, + device_ids=[args.gpu], + gradient_as_bucket_view=True, + ) + model_without_ddp = model.module + + # optimizer + total_batch_size = args.batch_size * args.accum_iter * world_size + print( + f"total batch size: {total_batch_size} = " + f"{args.batch_size} bs per gpu x {args.accum_iter} accum x {world_size} gpus" + ) + + if not args.get("lr"): + args.lr = args.base_lr * total_batch_size / 256 + print(f"lr: {args.lr:.2e} = {args.base_lr:.2e} x {total_batch_size} / 256") + else: + print(f"lr: {args.lr:.2e}") + + param_groups = ut.get_param_groups(model) + ut.update_lr(param_groups, args.lr) + ut.update_wd(param_groups, args.weight_decay) + # cast or else it corrupts the checkpoint + betas = tuple(args.betas) if args.betas is not None else None + optimizer = torch.optim.AdamW(param_groups, betas=betas, fused=True) + + epoch_num_batches = len(train_loader) + steps_per_epoch = math.ceil(epoch_num_batches / args.accum_iter) + total_steps = args.epochs * steps_per_epoch + warmup_steps = args.warmup_epochs * steps_per_epoch + lr_schedule = ut.WarmupThenCosine( + base_value=args.lr, + final_value=args.min_lr, + total_iters=total_steps, + warmup_iters=warmup_steps, + ) + print(f"full schedule: epochs = {args.epochs} (steps = {total_steps})") + print(f"warmup: epochs = {args.warmup_epochs} (steps = {warmup_steps})") + + # loss scaling not needed for bfloat16 (according to timm) + if args.amp and args.amp_dtype != "bfloat16": + loss_scaler = torch.GradScaler(device.type) + else: + loss_scaler = None + + # load checkpoint/resume training + ut.load_model(args, model_without_ddp, optimizer, loss_scaler) + + print(f"start training for {args.epochs} epochs") + start_time = time.monotonic() + for epoch in range(args.start_epoch, args.epochs): + train_stats = train_one_epoch( + args, + model, + train_loader, + optimizer, + loss_scaler, + lr_schedule, + epoch, + device, + ) + eval_stats = {} + eval_plots = {} + eval_period = args.get("eval_period", 1) + if eval_period and (epoch % eval_period == 0 or epoch == args.epochs - 1): + for name, loader in eval_loaders.items(): + stats, plots = evaluate( + args, + model, + loader, + epoch, + device, + eval_name=name, + ) + eval_stats.update(stats) + eval_plots.update(plots) + + merged_stats = {"epoch": epoch, **train_stats, **eval_stats} + if is_master: + with (output_dir / "log.json").open("a") as f: + print(json.dumps(merged_stats), file=f) + + for plot_name, img in eval_plots.items(): + plot_name = plot_name.replace("/", "__") + img.save(output_dir / f"{plot_name}__{epoch:05d}.png") + + ut.save_model(args, epoch, model_without_ddp, optimizer, loss_scaler) + sync_checkpoints_to_r2(args, output_dir) + + if args.distributed: + torch.distributed.destroy_process_group() + + total_time = time.monotonic() - start_time + print(f"done! training time: {datetime.timedelta(seconds=int(total_time))}") + + +def create_data_loaders(args: DictConfig): + data_loaders = {} + dataset_names = [args.train_dataset] + args.eval_datasets + + for dataset_name in dataset_names: + dataset_config = args.datasets[dataset_name].copy() + drop_last = dataset_config.pop("drop_last") + is_train = dataset_name == args.train_dataset + + print(f"loading dataset: {dataset_name}\n\n{OmegaConf.to_yaml(dataset_config)}") + shuffle = dataset_config["shuffle"] + samples_per_epoch = dataset_config.pop("samples_per_epoch") + dataset = mri_data.make_sparse_wds_dataset( + dataset_config["url"], + shuffle=shuffle, + buffer_size=dataset_config["buffer_size"], + ) + num_workers = int(args.num_workers) + loader_kwargs = { + "batch_size": args.batch_size, + "collate_fn": partial(mri_data.collate, include_meta=not is_train), + "shuffle": False, + "num_workers": num_workers, + "persistent_workers": num_workers > 0, + "pin_memory": True, + "drop_last": drop_last, + "prefetch_factor": args.prefetch_factor, + } + loader = wds.WebLoader(dataset, **loader_kwargs) + num_batches = samples_per_epoch // (ut.get_world_size() * args.batch_size) + loader = loader.with_epoch(num_batches) + loader = loader.with_length(num_batches, silent=True) + + data_loaders[dataset_name] = loader + + train_loader = data_loaders.pop(args.train_dataset) + return train_loader, data_loaders + + +def sync_checkpoints_to_r2(args: DictConfig, output_dir: Path) -> None: + r2_sync_url = args.get("r2_sync") + if not r2_sync_url or not ut.is_main_process(): + return + + cmd = ["aws", "s3", "sync", str(output_dir), str(r2_sync_url), "--profile", "r2"] + print(f"syncing checkpoints to R2: {output_dir} -> {r2_sync_url}") + subprocess.run(cmd, check=True) + + +def train_one_epoch( + args: DictConfig, + model: nn.Module, + data_loader: Iterable, + optimizer: torch.optim.Optimizer, + loss_scaler: torch.GradScaler | None, + lr_schedule: Sequence[float], + epoch: int, + device: torch.device, +): + model.train() + + metric_logger = ut.MetricLogger(delimiter=" ") + metric_logger.add_meter("lr", ut.SmoothedValue(window_size=1, fmt="{value:.6f}")) + metric_logger.add_meter("grad", ut.SmoothedValue()) + header = f"Train: [{epoch}]" + log_wandb = args.wandb and ut.is_main_process() + + epoch_num_batches = len(data_loader) + steps_per_epoch = math.ceil(epoch_num_batches / args.accum_iter) + + print_freq = args.get("print_freq", 100) if not args.debug else 1 + num_batches = epoch_num_batches if not args.debug else 10 + amp_dtype = getattr(torch, args.amp_dtype) + use_cuda = device.type == "cuda" + if use_cuda and args.presend_cuda: + data_loader = ut.pre_send_to_cuda_wrapper( + data_loader, device, dtype_map={torch.float16: amp_dtype} + ) + + optimizer.zero_grad() + + for batch_idx, batch in enumerate( + metric_logger.log_every(data_loader, print_freq, header, total_steps=num_batches) + ): + if use_cuda and not args.presend_cuda: + batch = ut.send_data(batch, device, dtype_map={torch.float16: amp_dtype}) + + batch_step = batch_idx + 1 + log_step = batch_step % print_freq == 0 or batch_step == num_batches + update_in_epoch = batch_idx // args.accum_iter + group_size = min(args.accum_iter, num_batches - update_in_epoch * args.accum_iter) + need_update = batch_step % args.accum_iter == 0 or batch_step == num_batches + global_step = epoch * steps_per_epoch + update_in_epoch + lr = lr_schedule[global_step] + if need_update: + ut.update_lr(optimizer.param_groups, lr) + + images, img_mask = mri_data.densify_sparse_image_batch( + batch["image_values"], + batch["img_mask"], + (int(args.get("in_chans", 1)), *args.img_size), + dtype=amp_dtype, + ) + + sync_context = model.no_sync() if args.distributed and not need_update else nullcontext() + with sync_context: + with torch.autocast(device_type=device.type, dtype=amp_dtype, enabled=args.amp): + loss = model( + images, + img_mask=img_mask, + mask_ratio=args.mask_ratio, + pred_mask_ratio=args.pred_mask_ratio, + pad_to_multiple=args.pad_to_multiple, + with_state=False, + ) + + loss_for_log = loss.detach() + torch._assert_async(torch.isfinite(loss_for_log), "non-finite loss") + + grad_norm = ut.backward_step( + loss / group_size, + optimizer, + scaler=loss_scaler, + need_update=need_update, + max_norm=args.clip_grad, + ) + + if need_update and log_step: + loss_value = loss_for_log.item() + grad_norm_value = grad_norm.item() + metric_logger.update(loss=loss_value, lr=lr, grad=grad_norm_value) + if log_wandb: + wandb.log( + { + "train/loss": loss_value, + "train/lr": lr, + "train/grad": grad_norm_value, + }, + step=int(1000 * (epoch + batch_step / epoch_num_batches)), + ) + + # gather the stats from all processes + metric_logger.synchronize_between_processes() + print("Averaged stats:", metric_logger) + return {f"train/{k}": meter.global_avg for k, meter in metric_logger.meters.items()} + + +@torch.inference_mode() +def evaluate( + args: DictConfig, + model: nn.Module, + data_loader: Iterable, + epoch: int, + device: torch.device, + eval_name: str, +): + model.eval() + + metric_logger = ut.MetricLogger(delimiter=" ") + header = f"Eval ({eval_name}): [{epoch}]" + is_master = ut.is_main_process() + log_wandb = args.wandb and is_master + + epoch_num_batches = len(data_loader) + if epoch_num_batches <= 0: + raise ValueError(f"eval loader {eval_name!r} has zero batches") + + print_freq = args.get("print_freq", 100) if not args.debug else 1 + num_batches = epoch_num_batches if not args.debug else 10 + num_batches = min(num_batches, epoch_num_batches) + eval_seed = int(args.get("eval_seed", args.seed)) + ut.get_rank() + example_step = random.Random(eval_seed).randint(1, num_batches) + amp_dtype = getattr(torch, args.amp_dtype) + use_cuda = device.type == "cuda" + rng_state = ut.capture_rng_state() + torch.set_rng_state(torch.Generator().manual_seed(eval_seed).get_state()) + if use_cuda: + torch.cuda.manual_seed(eval_seed) + if use_cuda and args.presend_cuda: + data_loader = ut.pre_send_to_cuda_wrapper( + data_loader, device, dtype_map={torch.float16: amp_dtype} + ) + + eval_batches = islice(data_loader, num_batches) + for batch_idx, batch in enumerate( + metric_logger.log_every(eval_batches, print_freq, header, total_steps=num_batches) + ): + if use_cuda and not args.presend_cuda: + batch = ut.send_data(batch, device, dtype_map={torch.float16: amp_dtype}) + + batch_step = batch_idx + 1 + + images, img_mask = mri_data.densify_sparse_image_batch( + batch["image_values"], + batch["img_mask"], + (int(args.get("in_chans", 1)), *args.img_size), + dtype=amp_dtype, + ) + + with torch.autocast(device_type=device.type, dtype=amp_dtype, enabled=args.amp): + loss, state = model( + images, + img_mask=img_mask, + mask_ratio=args.mask_ratio, + pred_mask_ratio=args.pred_mask_ratio, + pad_to_multiple=args.pad_to_multiple, + ) + + loss_value = loss.detach().float().item() + finite = torch.tensor(int(math.isfinite(loss_value)), dtype=torch.int32, device=device) + if args.distributed: + torch.distributed.all_reduce(finite, op=torch.distributed.ReduceOp.MIN) + if not finite.item(): + raise RuntimeError("non-finite validation loss detected") + metric_logger.meters["loss"].update(loss_value, n=int(batch["img_mask"].shape[0])) + + if is_master and batch_step == example_step: + example_batch = {"image": images[:1], "img_mask": img_mask[:1]} + if "meta" in batch: + example_batch["meta"] = batch["meta"][:1] + example_state = { + "pred_images": state["pred_images"][:1], + "pred_mask": state["pred_mask"][:1], + } + example_data = { + "batch": ut.send_data(example_batch, "cpu"), + "state": ut.send_data(example_state, "cpu"), + } + + # gather the stats from all processes + metric_logger.synchronize_between_processes() + print(f"Averaged stats ({eval_name}):", metric_logger) + stats = {f"eval/{eval_name}/{k}": meter.global_avg for k, meter in metric_logger.meters.items()} + + plots = {} + if is_master: + print(f"Making plots ({eval_name}): example={example_step}") + plots = make_plots(args, **example_data) + plots = {f"eval/{eval_name}/{k}": img for k, img in plots.items()} + + if log_wandb: + wandb.log(stats, step=1000 * (epoch + 1)) + wandb.log( + {k: wandb.Image(img, caption=f"example={example_step}") for k, img in plots.items()}, + step=1000 * (epoch + 1), + ) + ut.restore_rng_state(rng_state) + return stats, plots + + +def make_plots( + args: DictConfig, + batch: dict[str, Tensor], + state: dict[str, Tensor], +) -> dict[str, Image.Image]: + fig_kwargs = args.get("fig_kwargs", {}) + + images = batch["image"] + img_mask = batch.get("img_mask") + if img_mask is not None: + img_mask = img_mask.expand_as(images) + + raw_mean, raw_std = vis.raw_stats_from_batch(batch) + + plots = {} + mask_pred_fig = vis.plot_mask_pred( + target=images, + pred=state["pred_images"], + pred_mask=state["pred_mask"], + img_mask=img_mask, + patch_size=args.patch_size, + raw_mean=raw_mean, + raw_std=raw_std, + **ut.filter_kwargs(vis.plot_mask_pred, fig_kwargs), + ) + plots["mask_pred"] = vis.fig2pil(mask_pred_fig) + plt.close(mask_pred_fig) + + return plots + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--cfg-path", type=str, default=None) + parser.add_argument("--overrides", type=str, default=None, nargs="+") + args = parser.parse_args() + cfg = OmegaConf.load(DEFAULT_CONFIG) + if args.cfg_path: + cfg = OmegaConf.unsafe_merge(cfg, OmegaConf.load(args.cfg_path)) + if args.overrides: + cfg = OmegaConf.unsafe_merge(cfg, OmegaConf.from_dotlist(args.overrides)) + main(cfg) diff --git a/finetune/fomo_tune_baseline/output/task1/build/smri_mae/masking.py b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/masking.py new file mode 100644 index 0000000000000000000000000000000000000000..28f9d53cbbd8d32a3923f9b1a6655c92cd526bf7 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/masking.py @@ -0,0 +1,80 @@ +import torch +from jaxtyping import Float, Int +from torch import Tensor + + +def pad_patch_mask( + patch_mask: Float[Tensor, "B N"], + mask_ratio: float, + shuffle: bool = False, + generator: torch.Generator | None = None, + pad_to_multiple: int | None = None, +) -> tuple[Float[Tensor, "B N"], Int[Tensor, "B L"], Tensor]: + """ + Select each row's own mask-ratio count, then pad ids to the batch max length. + + Returns: + - selected patch mask [B, N] + - padded selected patch ids [B, Lpad] + - token mask [B, Lpad], true for real ids and false for padding + """ + if not 0.0 <= mask_ratio <= 1.0: + raise ValueError(f"mask_ratio must be in [0, 1], got {mask_ratio}") + + B, N = patch_mask.shape + device = patch_mask.device + patch_mask = patch_mask.to(dtype=torch.bool) + + valid_counts = patch_mask.sum(dim=1) + num_keep = torch.floor(valid_counts.to(torch.float64) * (1.0 - mask_ratio)).to(torch.long) + if not shuffle: + selected = patch_mask & (patch_mask.cumsum(dim=1) <= num_keep.unsqueeze(1)) + padded_ids, token_mask = patch_ids_from_mask( + selected, + pad_to_multiple=pad_to_multiple, + ) + return selected, padded_ids, token_mask + + # One masked sort directly produces random valid IDs. The previous + # shuffle/select/inverse-shuffle path required two full argsorts plus a + # dynamic nonzero/scatter solely to recover the same selected set. + noise = torch.rand(B, N, generator=generator, device=device) + noise.masked_fill_(~patch_mask, torch.inf) + shuffled_ids = torch.argsort(noise, dim=1) + + max_count = int(num_keep.max().item()) + if pad_to_multiple is not None: + if pad_to_multiple <= 0: + raise ValueError(f"pad_to_multiple must be positive, got {pad_to_multiple}") + max_count = (max_count + pad_to_multiple - 1) // pad_to_multiple * pad_to_multiple + padded_ids = shuffled_ids[:, :max_count] + token_mask = torch.arange(max_count, device=device).unsqueeze(0) < num_keep.unsqueeze(1) + selected = torch.zeros_like(patch_mask).scatter_(1, padded_ids, token_mask) + return selected, padded_ids, token_mask + + +def patch_ids_from_mask( + patch_mask: Tensor, + pad_to_multiple: int | None = None, +) -> tuple[Int[Tensor, "B L"], Tensor]: + """Return optionally rounded patch IDs and their token-validity mask.""" + if pad_to_multiple is not None and pad_to_multiple <= 0: + raise ValueError(f"pad_to_multiple must be positive, got {pad_to_multiple}") + + patch_mask = patch_mask.to(dtype=torch.bool) + B, N = patch_mask.shape + device = patch_mask.device + counts = patch_mask.sum(dim=1) + max_count = int(counts.max().item()) + if pad_to_multiple is not None: + max_count = (max_count + pad_to_multiple - 1) // pad_to_multiple * pad_to_multiple + + patch_ids = torch.zeros((B, max_count), dtype=torch.long, device=device) + token_mask = torch.arange(max_count, device=device).unsqueeze(0) < counts.unsqueeze(1) + if max_count == 0: + return patch_ids, token_mask + + batch_ids, selected_ids = patch_mask.nonzero(as_tuple=True) + slot_ids = patch_mask.cumsum(dim=1)[batch_ids, selected_ids].to(torch.long) - 1 + patch_ids[batch_ids, slot_ids] = selected_ids + return patch_ids, token_mask diff --git a/finetune/fomo_tune_baseline/output/task1/build/smri_mae/model_mae.py b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/model_mae.py new file mode 100644 index 0000000000000000000000000000000000000000..f9a4d884784e53f55ae557c46dd4afa2eb70e7d3 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/model_mae.py @@ -0,0 +1,916 @@ +# Copyright (c) Sophont, Inc +# This source code is licensed under the Apache License, Version 2.0 +# +# References: +# capi: https://github.com/facebookresearch/capi/blob/main/model.py +# timm: https://github.com/huggingface/pytorch-image-models/blob/v1.0.20/timm/models/vision_transformer.py + +""" +From-scratch re-implementation of the original MAE model. + +MaskedEncoder: standard ViT with masking +MaskedDecoder: standard self-attention MAE decoder +MaskedAutoEncoderViT: full MAE model for 3D structural MRI volumes +""" + +from collections.abc import Sequence +from typing import Literal + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch import Tensor +from torch.utils.checkpoint import checkpoint +from huggingface_hub import PyTorchModelHubMixin +from jaxtyping import Float, Int + +from .modules import ( + AbsolutePosEmbed, + Block, + LayerNorm, + Normalize, + JaggedBatch, + Patchify3D, + SeparablePosEmbed, + SinCosPosEmbed3D, + unpack_tokens, +) +from .masking import pad_patch_mask + + +class MaskedEncoder(nn.Module): + """ + Masked transformer encoder. + """ + + def __init__( + self, + patchify: nn.Module, + patch_embed: nn.Module, + pos_embed: nn.Module, + depth: int = 12, + embed_dim: int = 768, + num_heads: int = 12, + qkv_bias: bool = True, + proj_bias: bool = True, + mlp_ratio: int | float = 4, + class_token: bool = True, + reg_tokens: int = 0, + no_embed_class: bool = False, + final_norm: bool = True, + drop_path_rate: float = 0.0, + mask_drop_scale: bool = False, + ): + super().__init__() + self.num_prefix_tokens = int(class_token) + reg_tokens + self.num_reg_tokens = reg_tokens + self.has_class_token = class_token + self.no_embed_class = no_embed_class + + # scale inputs by 1 / observed rate (like dropout) + self.mask_drop_scale = mask_drop_scale + + # inject tokenization modules, so that the encoder doesn't specifically need to + # know how the data are tokenized, while still implementing a complete + # self-contained model. + self.patchify = patchify + self.patch_embed = patch_embed + self.pos_embed = pos_embed + + R = reg_tokens + self.cls_token = nn.Parameter(torch.empty(1, 1, embed_dim)) if class_token else None + self.reg_token = nn.Parameter(torch.empty(1, R, embed_dim)) if reg_tokens else None + + if not no_embed_class: + self.cls_token_pos = nn.Parameter(torch.empty(1, 1, embed_dim)) if class_token else None + self.reg_token_pos = nn.Parameter(torch.empty(1, R, embed_dim)) if reg_tokens else None + else: + self.cls_token_pos = self.reg_token_pos = None + + # stochastic depth decay rule + dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)] + + self.blocks = nn.ModuleList( + [ + Block( + dim=embed_dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + mlp_ratio=mlp_ratio, + drop_path=dpr[ii], + ) + for ii in range(depth) + ] + ) + + self.norm = LayerNorm(embed_dim) if final_norm else nn.Identity() + + self.reset_parameters() + + def extra_repr(self): + return ( + f"class_token={self.has_class_token}, reg_tokens={self.num_reg_tokens}, " + f"no_embed_class={self.no_embed_class}, mask_drop_scale={self.mask_drop_scale}" + ) + + def reset_parameters(self) -> None: + for p in [self.cls_token, self.cls_token_pos, self.reg_token, self.reg_token_pos]: + if p is not None: + nn.init.trunc_normal_(p, std=0.02) + + def cat_tokens(self, x: Tensor) -> Tensor: + # prepend cls and reg tokens with optional learned position embedding + # the cls and reg pos embedding is ofc redundant, but included in many other + # implementations. + B, _, _ = x.shape + + to_cat = [] + if self.has_class_token: + cls_token = self.cls_token + if not self.no_embed_class: + cls_token = cls_token + self.cls_token_pos + to_cat.append(cls_token.expand(B, -1, -1)) + + if self.num_reg_tokens: + reg_token = self.reg_token + if not self.no_embed_class: + reg_token = reg_token + self.reg_token_pos + to_cat.append(reg_token.expand(B, -1, -1)) + + if to_cat: + x = torch.cat(to_cat + [x], dim=1) + return x + + def cat_token_mask(self, token_mask: Tensor, batch_size: int) -> Tensor: + if self.num_prefix_tokens: + prefix_mask = torch.ones( + (batch_size, self.num_prefix_tokens), + dtype=torch.bool, + device=token_mask.device, + ) + token_mask = torch.cat([prefix_mask, token_mask], dim=1) + return token_mask + + def chunk_tokens(self, x: Tensor) -> tuple[Tensor | None, Tensor | None, Tensor]: + cls_offset = int(self.has_class_token) + cls = x[:, :cls_offset] if self.has_class_token else None + if self.num_reg_tokens: + reg = x[:, cls_offset : self.num_prefix_tokens, :] + else: + reg = None + patch = x[:, self.num_prefix_tokens :, :] + return cls, reg, patch + + def forward( + self, + x: Tensor, + mask: Tensor | None = None, + mask_ratio: float | None = None, + pad_to_multiple: int | None = None, + ) -> tuple[ + Float[Tensor, "B 1 D"] | None, + Float[Tensor, "B R D"] | None, + Float[Tensor, "B L D"], + Tensor | None, + Int[Tensor, "B L"] | None, + Tensor | None, + ]: + """ + x: input data shape [B, C, D, H, W] + mask: visible mask, 1 = visible, 0 = invisible. broadcastable shape + mask_ratio: mask ratio for uniform random masking + + returns: + - cls_embeds: [B, 1, D] + - reg_embeds: [B, R, D] + - patch_embeds: [B, L, D], where L is the number of visible patches + - mask: observed mask, 1 = observed, 0 = unobserved. same shape as input + - mask_ids: indices of visible patches [B L] + - token_mask: valid token mask for padded per-sample masking [B L] + """ + # apply mask to the input + if mask is not None: + mask = mask.to(device=x.device, dtype=torch.bool).expand_as(x) + x = x.masked_fill(~mask, 0) + + # patchify input + x = self.patchify(x) + B, N, P = x.shape + + # patchify mask and apply dropout style scaling + if mask is not None: + mask_patches = self.patchify(mask) + patch_num_obs = mask_patches.sum(dim=-1) + patch_mask = patch_num_obs > 0 + if self.mask_drop_scale: + patch_num_obs = patch_num_obs.to(x.dtype) + x = x * (P / patch_num_obs.unsqueeze(-1).clamp(min=1.0)) + elif mask_ratio is not None: + patch_mask = torch.ones((B, N), dtype=torch.bool, device=x.device) + mask_patches = patch_mask.unsqueeze(-1).expand(-1, -1, P) + else: + patch_mask = mask_patches = None + + # patch and position embed + x = self.patch_embed(x) + x = self.pos_embed(x) + + if mask is not None or mask_ratio is not None: + mask_ratio = 0.0 if mask_ratio is None else mask_ratio + patch_mask, mask_ids, token_mask = pad_patch_mask( + patch_mask, + mask_ratio=mask_ratio, + shuffle=mask_ratio > 0, + pad_to_multiple=pad_to_multiple, + ) + + mask_patches = mask_patches & patch_mask.unsqueeze(-1) + mask = self.patchify.unpatchify(mask_patches) + x = x.gather(1, mask_ids.unsqueeze(-1).expand(-1, -1, x.shape[-1])) + else: + mask_ids = None + token_mask = None + + cls_embeds, reg_embeds, patch_embeds = self.forward_patch_embeds( + x, + token_mask=token_mask, + ) + return cls_embeds, reg_embeds, patch_embeds, mask, mask_ids, token_mask + + def forward_patch_embeds( + self, + x: Float[Tensor, "B L D"], + token_mask: Tensor | None = None, + ) -> tuple[ + Float[Tensor, "B 1 D"] | None, + Float[Tensor, "B R D"] | None, + Float[Tensor, "B L D"], + ]: + B = x.shape[0] + if token_mask is None: + token_mask = torch.ones(x.shape[:2], dtype=torch.bool, device=x.device) + x = self.cat_tokens(x) + token_mask = self.cat_token_mask(token_mask, B) + jagged_batch = JaggedBatch.from_mask(token_mask) + x = x[token_mask] + for block in self.blocks: + x = block(x, jagged_batch=jagged_batch) + x = self.norm(x) + x = unpack_tokens(x, token_mask) + + cls_embeds, reg_embeds, patch_embeds = self.chunk_tokens(x) + return cls_embeds, reg_embeds, patch_embeds + + def forward_visible_ids( + self, + x: Tensor, + visible_ids: Int[Tensor, "B L"], + img_mask: Tensor | None = None, + ) -> tuple[ + Float[Tensor, "B 1 D"] | None, + Float[Tensor, "B R D"] | None, + Float[Tensor, "B L D"], + ]: + if img_mask is not None: + img_mask = img_mask.to(device=x.device, dtype=torch.bool).expand_as(x) + x = x.masked_fill(~img_mask, 0) + + x = self.patchify(x) + if self.mask_drop_scale and img_mask is not None: + mask_patches = self.patchify(img_mask) + patch_num_obs = mask_patches.sum(dim=-1).to(x.dtype) + x = x * (self.patchify.patch_dim / patch_num_obs.unsqueeze(-1).clamp(min=1.0)) + x = self.patch_embed(x) + x = self.pos_embed(x) + visible_ids = visible_ids.to(device=x.device) + x = x.gather(1, visible_ids.unsqueeze(-1).expand(-1, -1, x.shape[-1])) + return self.forward_patch_embeds(x) + + def forward_embedding( + self, + x: Tensor, + mask: Tensor | None = None, + mask_ratio: float | None = None, + ): + cls_embeds, reg_embeds, patch_embeds, *_ = self.forward( + x, + mask=mask, + mask_ratio=mask_ratio, + ) + return cls_embeds, reg_embeds, patch_embeds + + +class MaskedDecoder(nn.Module): + """Self-attention MAE decoder supporting sparse subset decoding via pred_ids.""" + + def __init__( + self, + pos_embed: nn.Module, + head: nn.Module | None = None, + input_dim: int | None = None, + depth: int = 12, + embed_dim: int = 768, + num_heads: int = 12, + qkv_bias: bool = True, + proj_bias: bool = True, + mlp_ratio: int | float = 4, + class_token: bool = True, + no_embed_class: bool = False, + final_norm: bool = True, + ): + super().__init__() + input_dim = embed_dim if input_dim is None else input_dim + self.has_class_token = class_token + self.no_embed_class = no_embed_class + + self.cls_token = nn.Parameter(torch.empty(1, 1, embed_dim)) if class_token else None + self.cls_token_pos = ( + nn.Parameter(torch.empty(1, 1, embed_dim)) + if class_token and not no_embed_class + else None + ) + self.mask_token = nn.Parameter(torch.empty(1, 1, embed_dim)) + + # decoder position embedding, encodes query position information into masks + self.pos_embed = pos_embed + + self.proj = nn.Identity() if input_dim == embed_dim else nn.Linear(input_dim, embed_dim) + + self.blocks = nn.ModuleList( + [ + Block( + dim=embed_dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + mlp_ratio=mlp_ratio, + ) + for _ in range(depth) + ] + ) + + self.norm = LayerNorm(embed_dim) if final_norm else nn.Identity() + + # optional injected prediction head + self.head = nn.Identity() if head is None else head + + self.reset_parameters() + + def extra_repr(self): + return f"class_token={self.has_class_token}, no_embed_class={self.no_embed_class}" + + def reset_parameters(self) -> None: + # official mae initializes decoder cls token to zeros + # although perhaps this was an oversight + if self.cls_token is not None: + nn.init.zeros_(self.cls_token) + if self.cls_token_pos is not None: + nn.init.trunc_normal_(self.cls_token_pos, std=0.02) + nn.init.trunc_normal_(self.mask_token, std=0.02) + + def cat_tokens(self, x: Tensor) -> Tensor: + if not self.has_class_token: + return x + cls_token = self.cls_token + if not self.no_embed_class: + cls_token = cls_token + self.cls_token_pos + return torch.cat([cls_token.expand(x.shape[0], -1, -1), x], dim=1) + + def cat_token_mask(self, token_mask: Tensor, batch_size: int) -> Tensor: + if self.has_class_token: + cls_mask = torch.ones( + (batch_size, 1), + dtype=torch.bool, + device=token_mask.device, + ) + token_mask = torch.cat([cls_mask, token_mask], dim=1) + return token_mask + + def chunk_tokens(self, x: Tensor) -> tuple[Tensor | None, Tensor]: + cls_offset = int(self.has_class_token) + cls = x[:, :cls_offset] if self.has_class_token else None + patch = x[:, cls_offset:, :] + return cls, patch + + def forward( + self, + embeds: Float[Tensor, "B L D"], + embed_ids: Int[Tensor, "B L"] | None = None, + pred_ids: Int[Tensor, "B Q"] | None = None, + embed_token_mask: Tensor | None = None, + pred_token_mask: Tensor | None = None, + packed_output: bool = False, + ) -> Float[Tensor, "B Q P"] | Float[Tensor, "T P"]: + """ + embeds: input patch embeddings. + embed_ids: optional patch indices for input embeddings. If not provided, no + position will be added to the embeddings. + pred_ids: patch indices of query mask positions. If None, decode *all* patches. + + returns: + - pred [B, Q, P] where Q is the number of prediction patches and P is the output + dimension + """ + B, L, _ = embeds.shape + + Q = self.pos_embed.num_patches if pred_ids is None else pred_ids.shape[1] + mask = self.mask_token.expand(B, Q, -1) + mask = self.pos_embed(mask, pos_ids=pred_ids) + + embeds = self.proj(embeds) + + if embed_ids is not None: + embeds = self.pos_embed(embeds, pos_ids=embed_ids) + if embed_token_mask is None: + embed_token_mask = torch.ones((B, L), dtype=torch.bool, device=embeds.device) + if pred_token_mask is None: + pred_token_mask = torch.ones((B, Q), dtype=torch.bool, device=embeds.device) + x = torch.cat([embeds, mask], dim=1) + token_mask = torch.cat([embed_token_mask, pred_token_mask], dim=1) + + x = self.cat_tokens(x) + token_mask = self.cat_token_mask(token_mask, B) + jagged_batch = JaggedBatch.from_mask(token_mask) + x = x[token_mask] + # Keep headroom for rare maximum-length PSP batches. + checkpoint_start = max(0, len(self.blocks) - 2) + for block_index, block in enumerate(self.blocks): + if self.training and torch.is_grad_enabled() and block_index >= checkpoint_start: + x = checkpoint(block, x, jagged_batch, use_reentrant=False) + else: + x = block(x, jagged_batch=jagged_batch) + + x = self.norm(x) + if packed_output: + pred_offset = int(self.has_class_token) + L + prediction_mask = F.pad(pred_token_mask, (pred_offset, 0)) + return self.head(x[prediction_mask[token_mask]]) + + x = unpack_tokens(x, token_mask) + _, x = self.chunk_tokens(x) + + pred = x[:, L:] + pred = pred.masked_fill(~pred_token_mask.unsqueeze(-1), 0) + pred = self.head(pred) + return pred + + +class MaskedAutoencoderViT(nn.Module, PyTorchModelHubMixin): + def __init__( + self, + img_size: int | tuple[int, int, int] = (208, 240, 208), + patch_size: int | tuple[int, int, int] = (16, 16, 16), + in_chans: int = 1, + depth: int = 12, + embed_dim: int = 768, + num_heads: int = 12, + decoder_depth: int = 4, + decoder_embed_dim: int | None = 512, + decoder_num_heads: int | None = 16, # default from mae, head dim = 32 + qkv_bias: bool = True, + proj_bias: bool = True, + mlp_ratio: int | float = 4, + class_token: bool = True, + reg_tokens: int = 0, + no_embed_class: bool = False, + drop_path_rate: float = 0.0, + mask_drop_scale: bool = False, + no_decode_pos: bool = False, + pos_embed: Literal["abs", "sep", "sincos"] = "sincos", + target_norm: Literal["none", "global", "slice", "patch"] | None = None, + ): + super().__init__() + img_size = _to_3d_tuple(img_size, "img_size") + patch_size = _to_3d_tuple(patch_size, "patch_size") + + self.no_decode_pos = no_decode_pos # don't pos encode embeddings in decoder + + # patchify reshapes input into sequence of flattened patches, shape [B, N, P] + ndim = 3 + patchify = Patchify3D(img_size, patch_size, in_chans=in_chans) + + # linear patch embedding P -> D + patch_embed = nn.Linear(patchify.patch_dim, embed_dim) + + # position embedding + # separable position embedding decouples the first spatial axis from the + # others. Fixed sin/cos embeddings are the default for sMRI volumes. + if pos_embed == "sincos": + pos_embed_layer = SinCosPosEmbed3D + else: + pos_embed_layer = {"abs": AbsolutePosEmbed, "sep": SeparablePosEmbed}[pos_embed] + pos_embed = pos_embed_layer(embed_dim, patchify.grid_size) + + # encoder. for inference, this model can be extracted and used like a regular vit + self.encoder = MaskedEncoder( + patchify=patchify, + patch_embed=patch_embed, + pos_embed=pos_embed, + depth=depth, + embed_dim=embed_dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + mlp_ratio=mlp_ratio, + class_token=class_token, + reg_tokens=reg_tokens, + no_embed_class=no_embed_class, + drop_path_rate=drop_path_rate, + mask_drop_scale=mask_drop_scale, + ) + + self.pred_patchify = patchify + + # fall back to encoder architecture width + decoder_embed_dim = decoder_embed_dim or embed_dim + decoder_num_heads = decoder_num_heads or num_heads + + decoder_pos_embed = pos_embed_layer(decoder_embed_dim, self.pred_patchify.grid_size) + # we might want to try tying the weights of the prediction head to the patch + # embedding at some point. + decoder_head = nn.Linear(decoder_embed_dim, self.pred_patchify.patch_dim) + + self.decoder = MaskedDecoder( + pos_embed=decoder_pos_embed, + head=decoder_head, + input_dim=embed_dim, + depth=decoder_depth, + embed_dim=decoder_embed_dim, + num_heads=decoder_num_heads, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + mlp_ratio=mlp_ratio, + class_token=class_token, + no_embed_class=no_embed_class, + ) + + # mae style target normalization + # dim is relative to an unflattened embedding tensor of shape [B, *grid_size, D] + if target_norm not in {"none", None}: + norm_dim = { + "global": tuple(range(1, ndim + 2)), # full sequence + "slice": tuple(range(2, ndim + 2)), # each depth slice along first dim + "patch": -1, # normalize each patch independently (mae pix norm loss) + }[target_norm] + self.target_norm = Normalize(self.pred_patchify.grid_size, dim=norm_dim) + else: + self.target_norm = None + + self.init_weights() + + def extra_repr(self): + return f"no_decode_pos={self.no_decode_pos}" + + def init_weights(self): + self.apply(_init_weights) + + def prepare_targets(self, images: Tensor, img_mask: Tensor | None): + """ + images: [B, C, D, H, W] + img_mask: mask of valid data. only used for computing correct normalization + stats. same shape as images. + """ + targets_patches = self.pred_patchify(images) # [B, N, P] + + # target normalization + if self.target_norm is not None: + # full image data mask used for normalization stats only + if img_mask is not None: + img_mask_patches = self.pred_patchify(img_mask) + else: + img_mask_patches = None + targets_patches, *targets_stats = self.target_norm(targets_patches, img_mask_patches) + else: + targets_stats = None + + return targets_patches, targets_stats + + def prepare_masks( + self, + img_mask: Tensor, + visible_mask: Tensor | None, + pred_mask: Tensor | None, + device: torch.device, + ): + img_mask = img_mask.to(device=device, dtype=torch.bool) + + if visible_mask is None: + visible_mask = img_mask + else: + visible_mask = img_mask & visible_mask.to(device=device, dtype=torch.bool) + + if pred_mask is None: + pred_mask = img_mask + else: + pred_mask = img_mask & pred_mask.to(device=device, dtype=torch.bool) + + return img_mask, visible_mask, pred_mask + + def prepare_pred_mask( + self, + visible_mask: Tensor, + pred_mask: Tensor | None = None, + pred_mask_ratio: float | None = None, + pad_to_multiple: int | None = None, + ): + """ + prepare prediction mask by removing visible content + visible_mask: [B, C, D, H, W], 1 = visible, 0 = invisible + pred_mask: same shape, 1 = predict, 0 = don't predict + """ + if pred_mask is None: + pred_mask = torch.ones_like(visible_mask) + + pred_mask = pred_mask & ~visible_mask + + pred_mask_patches = self.pred_patchify(pred_mask) + pred_patch_mask = pred_mask_patches.any(dim=-1) + # Optionally subsample the prediction candidates. + mask_ratio = 0.0 if pred_mask_ratio is None else pred_mask_ratio + pred_patch_mask, pred_ids, pred_token_mask = pad_patch_mask( + pred_patch_mask, + mask_ratio=mask_ratio, + # With per-sample padding every candidate is retained when the ratio + # is zero, so randomizing their order is pure overhead. + shuffle=mask_ratio > 0, + pad_to_multiple=pad_to_multiple, + ) + pred_mask_patches = pred_mask_patches & pred_patch_mask.unsqueeze(-1) + return pred_mask_patches, pred_ids, pred_token_mask + + def forward_decoder( + self, + patch_embeds: Float[Tensor, "B L D"], + visible_ids: Int[Tensor, "B L"], + pred_ids: Int[Tensor, "B Q"] | None, + visible_token_mask: Tensor | None = None, + pred_token_mask: Tensor | None = None, + packed_output: bool = False, + ) -> Float[Tensor, "B Q P"] | Float[Tensor, "T P"]: + return self.decoder( + patch_embeds, + embed_ids=None if self.no_decode_pos else visible_ids, + pred_ids=pred_ids, + embed_token_mask=visible_token_mask, + pred_token_mask=pred_token_mask, + packed_output=packed_output, + ) + + def forward_loss( + self, + preds: Float[Tensor, "T P"], + targets_patches: Float[Tensor, "B N P"], + pred_mask_patches: Float[Tensor, "B N P"], + pred_ids: Int[Tensor, "B Q"], + pred_token_mask: Tensor, + ) -> Tensor: + """Average valid-voxel MSE within each scan, then average across scans.""" + batch_ids, slot_ids = pred_token_mask.nonzero(as_tuple=True) + patch_ids = pred_ids[batch_ids, slot_ids] + targets = targets_patches[batch_ids, patch_ids] + voxel_mask = pred_mask_patches[batch_ids, patch_ids] + + patch_errors = ((preds - targets) ** 2 * voxel_mask).sum(dim=1) + patch_voxels = voxel_mask.sum(dim=1).to(dtype=patch_errors.dtype) + batch_size = targets_patches.shape[0] + scan_errors = patch_errors.new_zeros(batch_size).scatter_add_(0, batch_ids, patch_errors) + scan_voxels = patch_voxels.new_zeros(batch_size).scatter_add_(0, batch_ids, patch_voxels) + return (scan_errors / scan_voxels).mean() + + @torch.no_grad() + def forward_pred_images( + self, + preds: Float[Tensor, "B Q P"], + pred_ids: Int[Tensor, "B Q"], + pred_token_mask: Tensor | None = None, + img_mask: Tensor | None = None, + targets_stats: tuple[Tensor, Tensor] | None = None, + ) -> Tensor: + B, _, P = preds.shape + N = self.pred_patchify.num_patches + if pred_token_mask is not None: + preds = preds.masked_fill(~pred_token_mask.unsqueeze(-1), 0) + + preds = torch.zeros((B, N, P), dtype=preds.dtype, device=preds.device).scatter_add_( + 1, pred_ids.unsqueeze(-1).expand(-1, -1, P), preds + ) + + if targets_stats is not None: + targets_mean, targets_std = targets_stats + preds = preds * targets_std + targets_mean + + pred_images = self.pred_patchify.unpatchify(preds) + if img_mask is not None: + pred_images = pred_images.masked_fill(~img_mask, 0) + return pred_images + + def forward( + self, + images: Tensor, + img_mask: Tensor, + mask_ratio: float, + pred_mask_ratio: float | None = None, + pad_to_multiple: int | None = None, + with_state: bool = True, + ) -> Tensor | tuple[Tensor, dict]: + img_mask, visible_mask, pred_mask = self.prepare_masks( + img_mask, + None, + None, + device=images.device, + ) + targets_patches, targets_stats = self.prepare_targets(images, img_mask) + + ( + cls_embeds, + reg_embeds, + patch_embeds, + visible_mask, + visible_ids, + visible_token_mask, + ) = self.encoder( + images, + mask=visible_mask, + mask_ratio=mask_ratio, + pad_to_multiple=pad_to_multiple, + ) + + pred_mask_patches, pred_ids, pred_token_mask = self.prepare_pred_mask( + visible_mask, + pred_mask=pred_mask, + pred_mask_ratio=pred_mask_ratio, + pad_to_multiple=pad_to_multiple, + ) + + preds = self.forward_decoder( + patch_embeds, + visible_ids, + pred_ids, + visible_token_mask=visible_token_mask, + pred_token_mask=pred_token_mask, + packed_output=not with_state, + ) + + loss_preds = preds if not with_state else preds[pred_token_mask] + loss = self.forward_loss( + loss_preds, + targets_patches, + pred_mask_patches, + pred_ids, + pred_token_mask, + ) + + if not with_state: + return loss + + pred_mask = self.pred_patchify.unpatchify(pred_mask_patches) + pred_images = self.forward_pred_images( + preds, + pred_ids, + pred_token_mask=pred_token_mask, + img_mask=img_mask, + targets_stats=targets_stats, + ) + + state = { + "targets_patches": targets_patches, + "targets_stats": targets_stats, + "patch_embeds": patch_embeds, + "cls_embeds": cls_embeds, + "reg_embeds": reg_embeds, + "visible_mask": visible_mask, + "visible_ids": visible_ids, + "visible_token_mask": visible_token_mask, + "pred_mask": pred_mask, + "pred_ids": pred_ids, + "pred_token_mask": pred_token_mask, + "preds": preds, + "pred_images": pred_images, + } + return loss, state + + def forward_embedding( + self, + x: Tensor, + mask: Tensor | None = None, + mask_ratio: float | None = None, + ): + return self.encoder.forward_embedding(x, mask=mask, mask_ratio=mask_ratio) + + +class MaskedViT(MaskedEncoder, PyTorchModelHubMixin): + def __init__( + self, + img_size: int | tuple[int, int, int] = (208, 240, 208), + in_chans: int = 1, + patch_size: int | tuple[int, int, int] = (16, 16, 16), + depth: int = 12, + embed_dim: int = 768, + num_heads: int = 12, + qkv_bias: bool = True, + proj_bias: bool = True, + mlp_ratio: int | float = 4, + class_token: bool = True, + reg_tokens: int = 0, + no_embed_class: bool = False, + final_norm: bool = True, + drop_path_rate: float = 0.0, + mask_drop_scale: bool = False, + pos_embed: Literal["abs", "sep", "sincos"] = "sincos", + ): + img_size = _to_3d_tuple(img_size, "img_size") + patch_size = _to_3d_tuple(patch_size, "patch_size") + + patchify = Patchify3D(img_size, patch_size, in_chans=in_chans) + patch_embed = nn.Linear(patchify.patch_dim, embed_dim) + if pos_embed == "sincos": + pos_embed_layer = SinCosPosEmbed3D + else: + pos_embed_layer = {"abs": AbsolutePosEmbed, "sep": SeparablePosEmbed}[pos_embed] + pos_embed = pos_embed_layer(embed_dim, patchify.grid_size) + + super().__init__( + patchify=patchify, + patch_embed=patch_embed, + pos_embed=pos_embed, + depth=depth, + embed_dim=embed_dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + mlp_ratio=mlp_ratio, + class_token=class_token, + reg_tokens=reg_tokens, + no_embed_class=no_embed_class, + final_norm=final_norm, + drop_path_rate=drop_path_rate, + mask_drop_scale=mask_drop_scale, + ) + + self.init_weights() + + def init_weights(self): + self.apply(_init_weights) + + +def _to_3d_tuple(value: int | Sequence[int], name: str) -> tuple[int, int, int]: + if isinstance(value, int): + return (value, value, value) + if len(value) != 3: + raise ValueError(f"{name} must have exactly 3 spatial dimensions, got {tuple(value)}") + return tuple(int(item) for item in value) + + +# JAX ViT xavier uniform init +# https://github.com/facebookresearch/capi/blob/main/model.py +def _init_weights(m: nn.Module) -> None: + if isinstance(m, nn.Linear): + nn.init.xavier_uniform_(m.weight) + if m.bias is not None: + nn.init.constant_(m.bias, 0) + elif isinstance(m, nn.LayerNorm) and m.elementwise_affine: + nn.init.constant_(m.weight, 1.0) + if m.bias is not None: + nn.init.constant_(m.bias, 0) + + +def _create_vit(**kwargs): + model = MaskedViT(**kwargs) + return model + + +def _create_mae_vit(**kwargs): + model = MaskedAutoencoderViT(**kwargs) + return model + + +def mae_vit_small(**kwargs): + model_args = dict(embed_dim=384, depth=12, num_heads=6) + return _create_mae_vit(**model_args, **kwargs) + + +def mae_vit_base(**kwargs): + model_args = dict(embed_dim=768, depth=12, num_heads=12) + return _create_mae_vit(**model_args, **kwargs) + + +def mae_vit_large(**kwargs): + model_args = dict(embed_dim=1024, depth=24, num_heads=16) + return _create_mae_vit(**model_args, **kwargs) + + +def mae_vit_huge(**kwargs): + model_args = dict(embed_dim=1280, depth=32, num_heads=16) + return _create_mae_vit(**model_args, **kwargs) + + +# "patch embed" baseline model, depth 0 ViT (hah) +def patch_embed_small(**kwargs): + model_args = dict(embed_dim=384, depth=0) + return _create_vit(**model_args, **kwargs) + + +def patch_embed_base(**kwargs): + model_args = dict(embed_dim=768, depth=0) + return _create_vit(**model_args, **kwargs) diff --git a/finetune/fomo_tune_baseline/output/task1/build/smri_mae/modules.py b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/modules.py new file mode 100644 index 0000000000000000000000000000000000000000..828aecb1bd57182f1690feb4e943395db1bdd515 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/modules.py @@ -0,0 +1,453 @@ +# This source code is licensed under the Apache License, Version 2.0 +# +# References: +# capi: https://github.com/facebookresearch/capi/blob/main/model.py +# timm: https://github.com/huggingface/pytorch-image-models/blob/v1.0.20/timm/models/vision_transformer.py +# vjepa2: https://github.com/facebookresearch/vjepa2/blob/main/src/models/utils/pos_embs.py + +import math +from functools import partial +from typing import NamedTuple, Type + +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch import Tensor +from einops import rearrange +from jaxtyping import Float, Int +from timm.layers import DropPath, to_3tuple + +Layer = Type[nn.Module] + + +class JaggedBatch(NamedTuple): + """Sequence boundaries and cached launch metadata for jagged attention.""" + + offsets: Tensor + max_seqlen: int + + @classmethod + def from_mask(cls, mask: Tensor) -> "JaggedBatch": + mask = mask.to(dtype=torch.bool) + counts = mask.sum(dim=1) + return cls( + offsets=F.pad(counts.cumsum(dim=0), (1, 0)), + max_seqlen=mask.shape[1], + ) + + def as_nested(self, tokens: Tensor) -> Tensor: + # Cached conservative bounds avoid min/max reductions and GPU-to-CPU + # synchronization when Flash SDPA inspects the jagged sequence lengths. + return torch.nested.nested_tensor_from_jagged( + tokens, + self.offsets, + min_seqlen=1, + max_seqlen=self.max_seqlen, + ).transpose(1, 2) + + +def unpack_tokens(tokens: Tensor, token_mask: Tensor) -> Tensor: + """Restore packed values to a padded batch, filling invalid slots with zero.""" + output = tokens.new_zeros((*token_mask.shape, *tokens.shape[1:])) + return output.index_put((token_mask,), tokens) + + +def jagged_scaled_dot_product_attention( + query: Tensor, + key: Tensor, + value: Tensor, + jagged_batch: JaggedBatch, +) -> Tensor: + """Run SDPA on a packed batch of variable-length sequences.""" + output_jagged = F.scaled_dot_product_attention( + jagged_batch.as_nested(query), + jagged_batch.as_nested(key), + jagged_batch.as_nested(value), + ) + return output_jagged.transpose(1, 2).values() + + +class Attention(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + qkv_bias: bool = False, + proj_bias: bool = False, + ) -> None: + super().__init__() + self.num_heads = num_heads + self.head_dim = dim // num_heads + self.qkv = nn.Linear(dim, 3 * dim, bias=qkv_bias) + self.proj = nn.Linear(dim, dim, bias=proj_bias) + + def extra_repr(self): + return f"num_heads={self.num_heads}" + + def forward( + self, + x: Float[Tensor, "L D"], + jagged_batch: JaggedBatch, + ) -> Float[Tensor, "L D"]: + L, D = x.shape + h, dh = self.num_heads, self.head_dim + + qkv = self.qkv(x).reshape(L, 3, h, dh) + q, k, v = qkv.unbind(1) + + x = jagged_scaled_dot_product_attention( + q, + k, + v, + jagged_batch=jagged_batch, + ) + x = x.reshape(L, D) + x = self.proj(x) + return x + + +class Mlp(nn.Module): + def __init__( + self, + dim: int, + mlp_ratio: int | float = 4, + bias: bool = False, + ) -> None: + super().__init__() + hidden_features = int(dim * mlp_ratio) + self.fc1 = nn.Linear(dim, hidden_features, bias=bias) + self.act = nn.GELU() + self.fc2 = nn.Linear(hidden_features, dim, bias=bias) + + def forward(self, x: Float[Tensor, "... D"]) -> Float[Tensor, "... D"]: + x = self.fc1(x) + x = self.act(x) + x = self.fc2(x) + return x + + +# timm default eps=1e-6 +LayerNorm = partial(nn.LayerNorm, eps=1e-6) + + +class Block(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + qkv_bias: bool = False, + proj_bias: bool = False, + mlp_ratio: int | float = 4, + drop_path: float = 0.0, + norm_layer: Layer = LayerNorm, + ) -> None: + super().__init__() + self.norm1 = norm_layer(dim) + self.attn = Attention( + dim=dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + ) + self.drop_path1 = DropPath(drop_path) if drop_path > 0 else nn.Identity() + + self.norm2 = norm_layer(dim) + self.mlp = Mlp( + dim=dim, + mlp_ratio=mlp_ratio, + bias=proj_bias, + ) + self.drop_path2 = DropPath(drop_path) if drop_path > 0 else nn.Identity() + + def forward( + self, + x: Float[Tensor, "L D"], + jagged_batch: JaggedBatch, + ) -> Float[Tensor, "L D"]: + x = x + self.drop_path1( + self.attn( + self.norm1(x), + jagged_batch=jagged_batch, + ) + ) + x = x + self.drop_path2(self.mlp(self.norm2(x))) + return x + + +# Patching and position embedding modules + + +class Patchify3D(nn.Module): + def __init__( + self, + img_size: int | tuple[int, int, int], + patch_size: int | tuple[int, int, int], + in_chans: int = 3, + ) -> None: + super().__init__() + self.img_size = to_3tuple(img_size) + self.patch_size = to_3tuple(patch_size) + self.in_chans = in_chans + + T, H, W = self.img_size + p_t, p_h, p_w = self.patch_size + if T % p_t or H % p_h or W % p_w: + raise ValueError( + f"img_size {self.img_size} must be divisible by patch_size {self.patch_size}" + ) + self.grid_size = (T // p_t, H // p_h, W // p_w) + self.num_patches = math.prod(self.grid_size) + self.patch_dim = in_chans * math.prod(self.patch_size) + + def forward(self, x: Float[Tensor, "B C T H W"]) -> Float[Tensor, "B N P"]: + x = patchify3d(x, self.patch_size) + return x + + def unpatchify(self, x: Float[Tensor, "B N P"]) -> Float[Tensor, "B C T H W"]: + x = unpatchify3d(x, patch_size=self.patch_size, img_size=self.img_size) + return x + + def extra_repr(self): + return f"{self.img_size}, {self.patch_size}, in_chans={self.in_chans}" + + +def patchify3d(x: Tensor, patch_size: tuple[int, int, int]) -> Tensor: + p_t, p_h, p_w = to_3tuple(patch_size) + B, C, T, H, W = x.shape + x = rearrange(x, "b c (t u) (h p) (w q) -> b (t h w) (c u p q)", u=p_t, p=p_h, q=p_w) + return x + + +def unpatchify3d( + x: Tensor, + patch_size: tuple[int, int, int], + img_size: tuple[int, int, int], +) -> Tensor: + B, N, P = x.shape + p_t, p_h, p_w = to_3tuple(patch_size) + T, H, W = to_3tuple(img_size) + x = rearrange( + x, + "b (t h w) (c u p q) -> b c (t u) (h p) (w q)", + t=T // p_t, + h=H // p_h, + w=W // p_w, + u=p_t, + p=p_h, + q=p_w, + ) + return x + + +class AbsolutePosEmbed(nn.Module): + def __init__(self, embed_dim: int, grid_size: tuple[int, ...]) -> None: + super().__init__() + self.embed_dim = embed_dim + self.grid_size = grid_size + self.num_patches = math.prod(grid_size) + + self.weight = nn.Parameter(torch.empty(self.num_patches, embed_dim)) + self.reset_parameters() + + def reset_parameters(self): + nn.init.trunc_normal_(self.weight, std=0.02) + + def forward( + self, + x: Float[Tensor, "B L D"], + pos_ids: Int[Tensor, "B L"] | None = None, + ) -> Float[Tensor, "B L D"]: + x = apply_pos_embed(x, self.weight, pos_ids=pos_ids) + return x + + def extra_repr(self): + return f"{self.embed_dim}, {self.grid_size}" + + +class SeparablePosEmbed(nn.Module): + def __init__(self, embed_dim: int, grid_size: tuple[int, ...]) -> None: + super().__init__() + self.embed_dim = embed_dim + self.grid_size = grid_size + self.num_patches = math.prod(grid_size) + + N_t, *grid_size_spatial = grid_size + N_s = math.prod(grid_size_spatial) + self.weight_spatial = nn.Parameter(torch.empty(1, N_s, embed_dim)) + self.weight_temporal = nn.Parameter(torch.empty(N_t, 1, embed_dim)) + self.reset_parameters() + + def reset_parameters(self): + nn.init.trunc_normal_(self.weight_spatial, std=0.02) + nn.init.trunc_normal_(self.weight_temporal, std=0.02) + + def forward( + self, + x: Float[Tensor, "B L D"], + pos_ids: Int[Tensor, "B L"] | None = None, + ) -> Float[Tensor, "B L D"]: + B, N, D = x.shape + weight = (self.weight_temporal + self.weight_spatial).flatten(0, 1) # [N, D] + x = apply_pos_embed(x, weight, pos_ids=pos_ids) + return x + + def extra_repr(self): + return f"{self.embed_dim}, {self.grid_size}" + + +class SinCosPosEmbed3D(nn.Module): + def __init__(self, embed_dim: int, grid_size: tuple[int, int, int]) -> None: + super().__init__() + self.embed_dim = embed_dim + self.grid_size = grid_size + self.num_patches = math.prod(grid_size) + + N_t, N_h, N_w = grid_size + weight = get_3d_sincos_pos_embed( + embed_dim=embed_dim, + grid_size=(N_h, N_w), + grid_depth=N_t, + uniform_power=True, + ) + self.weight = nn.Parameter(torch.from_numpy(weight).float(), requires_grad=False) + + def forward( + self, + x: Float[Tensor, "B L D"], + pos_ids: Int[Tensor, "B L"] | None = None, + ) -> Float[Tensor, "B L D"]: + x = apply_pos_embed(x, self.weight, pos_ids=pos_ids) + return x + + def extra_repr(self): + return f"{self.embed_dim}, {self.grid_size}" + + +# sincos pos embed utils from vjepa2, but fixed the confusing meshgrid indexing + + +def get_3d_sincos_pos_embed(embed_dim, grid_size, grid_depth, cls_token=False, uniform_power=False): + """ + grid_size: tuple of int of the grid height and width + grid_depth: int of the grid depth + returns: + pos_embed: [grid_depth*grid_height*grid_width, embed_dim] (w/o cls_token) + or [1+grid_depth*grid_height*grid_width, embed_dim] (w/ cls_token) + """ + grid_d = np.arange(grid_depth, dtype=float) + grid_h = np.arange(grid_size[0], dtype=float) + grid_w = np.arange(grid_size[1], dtype=float) + grid_d, grid_h, grid_w = np.meshgrid(grid_d, grid_h, grid_w, indexing="ij") + + if not uniform_power: + h_embed_dim = embed_dim // 4 + w_embed_dim = embed_dim // 4 + d_embed_dim = embed_dim // 2 + else: + h_embed_dim = w_embed_dim = d_embed_dim = int(np.ceil(embed_dim / 6) * 2) + + emb_h = get_1d_sincos_pos_embed_from_grid(h_embed_dim, grid_h) # (T*H*W, D1) + emb_w = get_1d_sincos_pos_embed_from_grid(w_embed_dim, grid_w) # (T*H*W, D2) + emb_d = get_1d_sincos_pos_embed_from_grid(d_embed_dim, grid_d) # (T*H*W, D3) + pos_embed = np.concatenate([emb_d, emb_h, emb_w], axis=1) + pos_embed = pos_embed[:, :embed_dim] + if cls_token: + pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0) + return pos_embed + + +def get_1d_sincos_pos_embed_from_grid(embed_dim, pos): + """ + embed_dim: output dimension for each position + pos: a list of positions to be encoded: size (M,) + returns: (M, D) + """ + assert embed_dim % 2 == 0 + omega = np.arange(embed_dim // 2, dtype=float) + omega /= embed_dim / 2.0 + omega = 1.0 / 10000**omega # (D/2,) + + pos = pos.reshape(-1) # (M,) + out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product + + emb_sin = np.sin(out) # (M, D/2) + emb_cos = np.cos(out) # (M, D/2) + + emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D) + return emb + + +def apply_pos_embed( + x: Float[Tensor, "B L D"], + weight: Float[Tensor, "N D"], + pos_ids: Int[Tensor, "B L"] | None = None, +) -> Float[Tensor, "B L D"]: + B, L, D = x.shape + weight = weight.expand(B, -1, -1) + if pos_ids is not None: + weight = weight.gather(1, pos_ids.unsqueeze(-1).expand(-1, -1, D)) + x = x + weight + return x + + +# (masked) normalization used for MAE target normalization + + +class Normalize(nn.Module): + def __init__( + self, + grid_size: tuple[int, ...], + dim: int | tuple[int, ...] | None = -1, + eps: float = 1e-6, + ) -> None: + super().__init__() + self.grid_size = grid_size + self.dim = dim + self.eps = eps + + def forward(self, x: Tensor, mask: Tensor | None = None) -> tuple[Tensor, Tensor, Tensor]: + """ + Normalize input sequence along dim(s) after reshaping to grid. + Returns tuple of (x, mean, std). + """ + B, N, D = x.shape + x = x.reshape((B, *self.grid_size, D)) + if mask is not None: + mask = mask.reshape((B, *self.grid_size, D)) + x, mean, std = masked_normalize(x, mask, dim=self.dim, eps=self.eps) + else: + x, mean, std = normalize(x, dim=self.dim, eps=self.eps) + mean = mean.expand_as(x).reshape(B, N, D) + std = std.expand_as(x).reshape(B, N, D) + x = x.reshape(B, N, D) + return x, mean, std + + def extra_repr(self): + return f"{self.grid_size}, dim={self.dim}" + + +def masked_normalize( + x: Tensor, + mask: Tensor, + dim: int | tuple[int, ...] | None = -1, + eps: float = 1e-6, +) -> tuple[Tensor, Tensor, Tensor]: + num_obs = mask.sum(dim=dim, keepdim=True).clamp(min=1) + mean = (mask * x).sum(dim=dim, keepdim=True) / num_obs + var = (mask * (x - mean) ** 2).sum(dim=dim, keepdim=True) / num_obs + std = (var + eps) ** 0.5 + x = mask * (x - mean) / std + return x, mean, std + + +def normalize( + x: Tensor, + dim: int | tuple[int, ...] | None = -1, + eps: float = 1e-6, +) -> tuple[Tensor, Tensor, Tensor]: + mean = x.mean(dim=dim, keepdim=True) + var = torch.var(x, dim=dim, keepdim=True, unbiased=False) + std = (var + eps) ** 0.5 + x = (x - mean) / std + return x, mean, std diff --git a/finetune/fomo_tune_baseline/output/task1/build/smri_mae/utils.py b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..0f6274f813429552c9ddfccfe4033ae677b52342 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/utils.py @@ -0,0 +1,581 @@ +# Copyright (c) Sophont, Inc +# This source code is licensed under the Apache License, Version 2.0 +# +# References: +# deit: https://github.com/facebookresearch/deit/blob/main/utils.py +# beit3: https://github.com/microsoft/unilm/blob/master/beit3/utils.py +# capi: https://github.com/facebookresearch/capi/blob/main/utils.py +# dinov2: https://github.com/facebookresearch/dinov2/blob/main/dinov2/utils/param_groups.py +# timm: https://github.com/huggingface/pytorch-image-models/blob/main/timm/utils/cuda.py +# dino: https://github.com/facebookresearch/dino/blob/main/utils.py + +import datetime +import inspect +import math +import os +import random +import subprocess +import time +from collections import defaultdict, deque +from omegaconf import OmegaConf +from pathlib import Path + +import numpy as np +import torch +import torch.distributed as dist +import torch.nn as nn +from torch import Tensor +from torch.amp import GradScaler +from torch.optim import Optimizer + + +# these very useful utils copied from deit with only minor changes +# thanks to the original authors, wherever you are + + +def configure_flash_sdpa() -> None: + """Use Flash Attention exclusively for CUDA SDPA.""" + torch.backends.cuda.enable_flash_sdp(True) + torch.backends.cuda.enable_mem_efficient_sdp(False) + torch.backends.cuda.enable_math_sdp(False) + torch.backends.cuda.enable_cudnn_sdp(False) + print("SDPA backend: flash") + + +class SmoothedValue: + """Track a series of values and provide access to smoothed values over a + window or the global series average. + """ + + def __init__(self, window_size=20, fmt=None): + if fmt is None: + fmt = "{median:.4f} ({global_avg:.4f})" + self.deque = deque(maxlen=window_size) + self.total = 0.0 + self.count = 0 + self.fmt = fmt + + def update(self, value, n=1): + value = float(value) + if math.isfinite(value): + self.deque.append(value) + self.count += n + self.total += value * n + + def synchronize_between_processes(self): + """ + Warning: does not synchronize the deque! + """ + if not is_dist_avail_and_initialized(): + return + t = torch.tensor([self.count, self.total], dtype=torch.float64, device="cuda") + dist.barrier() + dist.all_reduce(t) + t = t.tolist() + self.count = int(t[0]) + self.total = t[1] + + @property + def median(self): + if not self.count: + return float("nan") + d = torch.tensor(list(self.deque)) + return d.median().item() + + @property + def avg(self): + if not self.count: + return float("nan") + d = torch.tensor(list(self.deque), dtype=torch.float32) + return d.mean().item() + + @property + def global_avg(self): + if not self.count: + return float("nan") + return self.total / self.count + + @property + def max(self): + if not self.count: + return float("nan") + return max(self.deque) + + @property + def value(self): + if not self.count: + return float("nan") + return self.deque[-1] + + def __str__(self): + return self.fmt.format( + median=self.median, + avg=self.avg, + global_avg=self.global_avg, + max=self.max, + value=self.value, + ) + + +class MetricLogger: + def __init__(self, delimiter="\t"): + self.meters = defaultdict(SmoothedValue) + self.delimiter = delimiter + + def update(self, **kwargs): + for k, v in kwargs.items(): + if v is None: + continue + if isinstance(v, (torch.Tensor, np.generic)): + v = v.item() + assert isinstance(v, (float, int)) + self.meters[k].update(v) + + def __getattr__(self, attr): + if attr in self.meters: + return self.meters[attr] + if attr in self.__dict__: + return self.__dict__[attr] + raise AttributeError("'{}' object has no attribute '{}'".format(type(self).__name__, attr)) + + def __str__(self): + loss_str = [] + for name, meter in self.meters.items(): + loss_str.append("{}: {}".format(name, str(meter))) + return self.delimiter.join(loss_str) + + def synchronize_between_processes(self): + for meter in self.meters.values(): + meter.synchronize_between_processes() + + def add_meter(self, name, meter): + self.meters[name] = meter + + def log_every(self, iterable, print_freq, header=None, total_steps=None): + i = 0 + total_steps = total_steps or len(iterable) + if not header: + header = "" + start_time = time.time() + end = time.time() + iter_time = SmoothedValue(fmt="{avg:.4f}") + data_time = SmoothedValue(fmt="{avg:.4f}") + space_fmt = ":" + str(len(str(total_steps))) + "d" + log_msg = [ + header, + "[{0" + space_fmt + "}/{1}]", + "eta: {eta}", + "{meters}", + "time: {time}", + "data: {data}", + ] + if torch.cuda.is_available(): + log_msg.append("max mem: {memory:.0f}") + log_msg = self.delimiter.join(log_msg) + MB = 1024.0 * 1024.0 + for obj in iterable: + if i >= total_steps: + break + data_time.update(time.time() - end) + yield obj + iter_time.update(time.time() - end) + if i % print_freq == 0 or i == total_steps - 1: + eta_seconds = iter_time.global_avg * (total_steps - i) + eta_string = str(datetime.timedelta(seconds=int(eta_seconds))) + if torch.cuda.is_available(): + print( + log_msg.format( + i, + total_steps, + eta=eta_string, + meters=str(self), + time=str(iter_time), + data=str(data_time), + memory=torch.cuda.max_memory_allocated() / MB, + ) + ) + + else: + print( + log_msg.format( + i, + total_steps, + eta=eta_string, + meters=str(self), + time=str(iter_time), + data=str(data_time), + ) + ) + i += 1 + end = time.time() + total_time = time.time() - start_time + total_time_str = str(datetime.timedelta(seconds=int(total_time))) + print( + "{} Total time: {} ({:.4f} s / it)".format( + header, total_time_str, total_time / total_steps + ) + ) + + +def setup_for_distributed(log_path=None): + """ + This function disables printing when not in master process + """ + import builtins as __builtin__ + + builtin_print = __builtin__.print + + is_master = is_main_process() + + def print(*args, **kwargs): + force = kwargs.pop("force", False) + if is_master or force: + builtin_print(*args, **kwargs) + # tee to log file + if log_path and "file" not in kwargs: + with open(log_path, "a") as f: + builtin_print(*args, file=f, **kwargs) + + __builtin__.print = print + + +def is_dist_avail_and_initialized(): + if not dist.is_available(): + return False + if not dist.is_initialized(): + return False + return True + + +def get_world_size(): + if not is_dist_avail_and_initialized(): + return 1 + return dist.get_world_size() + + +def get_rank(): + if not is_dist_avail_and_initialized(): + return 0 + return dist.get_rank() + + +def is_main_process(): + return get_rank() == 0 + + +def save_on_master(obj, path, *args, **kwargs): + if is_main_process(): + path = Path(path) + tmp_path = path.with_name(f".{path.name}.tmp-{os.getpid()}") + try: + torch.save(obj, tmp_path, *args, **kwargs) + os.replace(tmp_path, path) + except Exception: + tmp_path.unlink(missing_ok=True) + raise + + +def init_distributed_mode(args): + # removed slurm block, can add if we use slurm + if "RANK" in os.environ and "WORLD_SIZE" in os.environ: + args.rank = int(os.environ["RANK"]) + args.world_size = int(os.environ["WORLD_SIZE"]) + args.gpu = int(os.environ["LOCAL_RANK"]) + else: + args.distributed = False + return + + args.distributed = True + + torch.cuda.set_device(args.gpu) + args.dist_backend = "nccl" + print(f"| distributed init (rank {args.rank})") + torch.distributed.init_process_group( + backend=args.dist_backend, + world_size=args.world_size, + rank=args.rank, + device_id=args.gpu, + ) + torch.distributed.barrier() + + +# checkpoint saving utils adapted from beit3 + + +def capture_rng_state() -> dict: + state = {"torch": torch.get_rng_state()} + if torch.cuda.is_available(): + state["cuda"] = torch.cuda.get_rng_state() + return state + + +def restore_rng_state(state: dict) -> None: + torch.set_rng_state(state["torch"]) + if "cuda" in state and torch.cuda.is_available(): + torch.cuda.set_rng_state(state["cuda"]) + + +def _all_rank_rng_states() -> list[dict]: + local_state = capture_rng_state() + if not is_dist_avail_and_initialized(): + return [local_state] + states = [None] * get_world_size() + dist.all_gather_object(states, local_state) + return states + + +def save_model(args, epoch, model_without_ddp, optimizer, loss_scaler): + output_dir = Path(args.output_dir) + checkpoint_path = output_dir / f"checkpoint-{epoch:05d}.pth" + last_checkpoint_path = output_dir / "checkpoint-last.pth" + if epoch % args.checkpoint_period != 0 and epoch != args.epochs - 1: + return + + to_save = { + "model": model_without_ddp.state_dict(), + "optimizer": optimizer.state_dict(), + "epoch": epoch, + "scaler": None if loss_scaler is None else loss_scaler.state_dict(), + "args": OmegaConf.to_container(args), + "rng_states": _all_rank_rng_states(), + } + + print(f"saving checkpoint {last_checkpoint_path}") + save_on_master(to_save, last_checkpoint_path) + print(f"saving checkpoint {checkpoint_path}") + save_on_master(to_save, checkpoint_path) + + if args.max_checkpoints and is_main_process(): + all_checkpoints = sorted(output_dir.glob("checkpoint-[0-9]*.pth")) + del_count = max(0, len(all_checkpoints) - args.max_checkpoints) + for checkpoint_path in all_checkpoints[:del_count]: + print(f"removing checkpoint {checkpoint_path}") + checkpoint_path.unlink() + + +def load_model(args, model_without_ddp, optimizer, loss_scaler): + auto_resume = getattr(args, "auto_resume", True) + output_dir = Path(args.output_dir) + + last_checkpoint_path = output_dir / "checkpoint-last.pth" + if auto_resume and last_checkpoint_path.exists(): + args.ckpt = str(last_checkpoint_path) + args.resume = True + + if args.ckpt: + ckpt = torch.load(args.ckpt, map_location="cpu", weights_only=True) + model_without_ddp.load_state_dict(ckpt["model"]) + print(f"loaded model state from checkpoint {args.ckpt}") + + if args.resume: + optimizer.load_state_dict(ckpt["optimizer"]) + if loss_scaler is not None: + loss_scaler.load_state_dict(ckpt["scaler"]) + args.start_epoch = ckpt["epoch"] + 1 + rng_states = ckpt.get("rng_states") + if rng_states is not None: + if len(rng_states) != get_world_size(): + raise ValueError( + "checkpoint RNG state world size does not match current world size: " + f"{len(rng_states)} != {get_world_size()}" + ) + restore_rng_state(rng_states[get_rank()]) + print(f"restored RNG state for rank {get_rank()}") + print(f"loaded optimizer state, resuming training from {args.start_epoch}") + + +# optimization utils + + +# from capi +class WarmupThenCosine: + def __init__( + self, + base_value: float, + final_value: float, + total_iters: int, + warmup_iters: int = 0, + start_warmup_value: float = 0.0, + freeze_iters: int = 0, + truncate_cos: float = 1.0, + ): + super().__init__() + self.final_value = final_value + self.total_iters = total_iters + + freeze_schedule = np.zeros(freeze_iters) + + warmup_schedule = np.linspace(start_warmup_value, base_value, warmup_iters) + + iters = np.arange(total_iters - warmup_iters - freeze_iters) + schedule = final_value + 0.5 * (base_value - final_value) * ( + 1 + np.cos(np.pi * truncate_cos * iters / len(iters)) + ) + self.schedule = np.concatenate((freeze_schedule, warmup_schedule, schedule)) + assert len(self.schedule) == self.total_iters + + def __getitem__(self, it: int) -> float: + if it >= self.total_iters: + return self.final_value + # cast to float or else it can corrupt the checkpoint + return float(self.schedule[it]) + + +# adapted from timm backward logic +# https://github.com/huggingface/pytorch-image-models/blob/main/timm/utils/cuda.py +def backward_step( + loss: Tensor, + optimizer: Optimizer, + scaler: GradScaler = None, + need_update: bool = True, + max_norm: float | None = None, +) -> Tensor | None: + if scaler is not None: + scaler.scale(loss).backward() + else: + loss.backward() + + if need_update: + if scaler is not None: + scaler.unscale_(optimizer) + + total_norm = clip_grad(optimizer, max_norm) + + if scaler is not None: + scaler.step(optimizer) + scaler.update() + else: + optimizer.step() + optimizer.zero_grad() + else: + total_norm = None + return total_norm + + +def clip_grad(optimizer: Optimizer, max_norm: float | None = None) -> Tensor: + params = [p for group in optimizer.param_groups for p in group["params"]] + if max_norm: + total_norm = nn.utils.clip_grad_norm_(params, max_norm, error_if_nonfinite=False) + else: + grads = [p.grad for p in params if p.grad is not None] + total_norm = nn.utils.get_total_norm(grads, error_if_nonfinite=False) + torch._assert_async(torch.isfinite(total_norm), "non-finite gradient norm") + return total_norm + + +# from dinov2 with some minor changes +def get_param_groups(model, patch_embed_lr_mult=1.0): + # no lr decay, we could add this later if needed + all_params = [] + + for name, param in model.named_parameters(): + if not param.requires_grad: + continue + d = {"param": param, "lr_multiplier": 1.0, "wd_multiplier": 1.0, "name": name} + + if name.endswith(".bias") or "norm" in name or "gamma" in name: + d["wd_multiplier"] = 0.0 + + if "patch_embed" in name: + d["lr_multiplier"] = d["lr_multiplier"] * patch_embed_lr_mult + + all_params.append(d) + + param_groups = _fuse_param_groups(all_params) + return param_groups + + +def _fuse_param_groups(all_param_groups): + fused_param_groups = defaultdict(lambda: {"params": []}) + for d in all_param_groups: + keys = sorted(set(d.keys()) - {"param", "name"}) + identifier = "_".join(f"{k}{d[k]}" for k in keys) + for k in keys: + fused_param_groups[identifier][k] = d[k] + fused_param_groups[identifier]["params"].append(d["param"]) + + param_groups = list(fused_param_groups.values()) + return param_groups + + +def update_lr(param_groups, lr: float): + for group in param_groups: + group["lr"] = lr * group["lr_multiplier"] + + +def update_wd(param_groups, weight_decay: float | None = None): + for group in param_groups: + group["weight_decay"] = weight_decay * group["wd_multiplier"] + + +# moving data to cuda utils copied from capi +# added device argument + + +def send_data(x, device=None, dtype_map=None): + if device is None: + device = torch.device("cuda") + else: + device = torch.device(device) + + if isinstance(x, torch.Tensor): + dtype = dtype_map.get(x.dtype) if dtype_map else None + return x.to(device=device, dtype=dtype, non_blocking=True) + if isinstance(x, dict): + return {k: send_data(v, device=device, dtype_map=dtype_map) for k, v in x.items()} + if isinstance(x, list): + return [send_data(v, device=device, dtype_map=dtype_map) for v in x] + return x + + +def pre_send_to_cuda_wrapper(generator, device=None, dtype_map=None): + """From apex""" + data = None + stream = torch.cuda.Stream(device) + for next_data in generator: + with torch.cuda.stream(stream): + next_data = send_data(next_data, device=device, dtype_map=dtype_map) + if data is not None: + yield data + torch.cuda.current_stream(device).wait_stream(stream) + data = next_data + if data is not None: + yield data + + +# other misc utils + + +# from dino +def get_sha(): + cwd = os.path.dirname(os.path.abspath(__file__)) + + def _run(command): + return subprocess.check_output(command, cwd=cwd).decode("ascii").strip() + + sha = "N/A" + diff = "clean" + branch = "N/A" + try: + sha = _run(["git", "rev-parse", "HEAD"]) + diff = _run(["git", "diff-index", "HEAD"]) + diff = "has uncommitted changes" if diff else "clean" + branch = _run(["git", "rev-parse", "--abbrev-ref", "HEAD"]) + except Exception: + pass + message = f"sha: {sha}, status: {diff}, branch: {branch}" + return message + + +# from timm +def random_seed(seed=42, rank=0): + torch.manual_seed(seed + rank) + np.random.seed(seed + rank) + random.seed(seed + rank) + + +# mine :) +def filter_kwargs(func, kwargs): + sigature = inspect.signature(func) + kwargs = {k: v for k, v in kwargs.items() if k in sigature.parameters} + return kwargs diff --git a/finetune/fomo_tune_baseline/output/task1/build/smri_mae/visualization.py b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/visualization.py new file mode 100644 index 0000000000000000000000000000000000000000..f721fe06e93dc1a2dfb4a679069029bdf55afc6b --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/build/smri_mae/visualization.py @@ -0,0 +1,400 @@ +from collections.abc import Mapping +from io import BytesIO + +import torch + +from matplotlib import patches +from matplotlib import pyplot as plt +from PIL import Image +from torch import Tensor + +VIEW_NAMES = { + "sagittal": "Sagittal", + "saggital": "Sagittal", + "axial": "Axial", + "coronal": "Coronal", +} + + +def fig2pil(fig) -> Image.Image: + buffer = BytesIO() + fig.savefig(buffer, format="png", dpi=fig.dpi, facecolor=fig.get_facecolor()) + buffer.seek(0) + image = Image.open(buffer).convert("RGB") + buffer.close() + return image + + +def raw_stats_from_batch(batch: dict) -> tuple[Tensor | None, Tensor | None]: + metas = batch.get("meta") + if not metas: + return None, None + + means = [] + stds = [] + for meta in metas: + try: + mean = meta["raw_mean"] + std = meta["raw_std"] + except (KeyError, TypeError): + return None, None + if mean in ("", None) or std in ("", None): + return None, None + means.append(float(mean)) + stds.append(float(std)) + return torch.tensor(means), torch.tensor(stds) + + +def plot_mask_pred( + target: Tensor, + pred: Tensor, + pred_mask: Tensor | None = None, + img_mask: Tensor | None = None, + sample_idx: int = 0, + channel_idx: int = 0, + slice_idx: int | Mapping[str, int] | None = None, + patch_size: int | tuple[int, int, int] = 16, + views: tuple[str, ...] = ("sagittal", "axial", "coronal"), + cmap: str = "gray", + figsize: tuple[float, float] | None = None, + mask_style: str = "blank", + raw_mean: float | Tensor | None = None, + raw_std: float | Tensor | None = None, +): + target_vol = _select_volume(target, sample_idx=sample_idx, channel_idx=channel_idx) + pred_vol = _select_volume(pred, sample_idx=sample_idx, channel_idx=channel_idx) + if raw_mean is not None and raw_std is not None: + raw_mean = _select_scalar(raw_mean, sample_idx=sample_idx) + raw_std = _select_scalar(raw_std, sample_idx=sample_idx) + target_vol = target_vol * raw_std + raw_mean + pred_vol = pred_vol * raw_std + raw_mean + pred_mask_vol = ( + torch.zeros_like(target_vol) + if pred_mask is None + else _select_volume(pred_mask, sample_idx=sample_idx, channel_idx=channel_idx) > 0 + ) + img_mask_vol = None + if img_mask is not None: + img_mask_vol = _select_volume(img_mask, sample_idx=sample_idx, channel_idx=channel_idx) > 0 + + composite_vol = _prediction_composite(target_vol, pred_vol, pred_mask_vol) + vmin, vmax = _intensity_limits(target_vol, img_mask_vol) + + patch_size = _as_3tuple(patch_size) + view_items = [] + for view in views: + view_key = view.lower() + if view_key not in VIEW_NAMES: + raise ValueError(f"unknown MRI view {view!r}; expected one of {tuple(VIEW_NAMES)}") + target_slice = _extract_view_slice(target_vol, view_key, slice_idx) + composite_slice = _extract_view_slice(composite_vol, view_key, slice_idx) + mask_slice = _extract_view_slice(pred_mask_vol.float(), view_key, slice_idx) > 0 + img_mask_slice = None + if img_mask_vol is not None: + img_mask_slice = _extract_view_slice(img_mask_vol.float(), view_key, slice_idx) > 0 + view_items.append( + { + "key": view_key, + "title": VIEW_NAMES[view_key], + "target": _masked_input_display(target_slice, mask_slice, img_mask_slice, vmin), + "composite": _apply_display_mask(composite_slice, img_mask_slice, vmin), + "actual": _apply_display_mask(target_slice, img_mask_slice, vmin), + "mask": mask_slice, + "img_mask": img_mask_slice, + "patch_rc": _view_patch_size(view_key, patch_size), + } + ) + + _crop_view_items(view_items) + + fig, axes, layout = _make_figure_canvas(view_items, figsize=figsize) + for item, x in zip(view_items, layout["col_centers"]): + fig.text( + x, + layout["title_y"], + item["title"], + ha="center", + va="center", + color="#f8fafc", + fontsize=7, + ) + for label, y in zip(("Masked", "Pred", "Actual"), layout["row_centers"]): + fig.text(layout["label_x"], y, label, ha="right", va="center", color="#cbd5e1", fontsize=6) + + for item, top_ax, middle_ax, bottom_ax in zip(view_items, axes[0], axes[1], axes[2]): + top_ax.imshow( + item["target"], + cmap=cmap, + vmin=vmin, + vmax=vmax, + interpolation="nearest", + origin="upper", + ) + if mask_style == "boxes": + _draw_patch_boxes(top_ax, item["mask"], item["patch_rc"]) + elif mask_style != "blank": + raise ValueError("mask_style must be 'blank' or 'boxes'") + _style_axis(top_ax) + + middle_ax.imshow( + item["composite"], + cmap=cmap, + vmin=vmin, + vmax=vmax, + interpolation="nearest", + origin="upper", + ) + _style_axis(middle_ax) + + bottom_ax.imshow( + item["actual"], + cmap=cmap, + vmin=vmin, + vmax=vmax, + interpolation="nearest", + origin="upper", + ) + _style_axis(bottom_ax) + return fig + + +def _select_volume( + x: Tensor, + sample_idx: int = 0, + channel_idx: int = 0, +) -> Tensor: + x = x.detach().float().cpu() + if x.ndim == 5: + return x[sample_idx, channel_idx] + if x.ndim == 4: + return x[sample_idx] + if x.ndim == 3: + return x + raise ValueError(f"expected a 3D volume tensor, got shape {tuple(x.shape)}") + + +def _select_scalar(value: float | Tensor, sample_idx: int = 0) -> float: + if isinstance(value, Tensor): + value = value.detach().float().cpu() + if value.ndim > 0: + value = value.reshape(-1)[sample_idx] + return float(value) + return float(value) + + +def _prediction_composite(target: Tensor, pred: Tensor, pred_mask: Tensor) -> Tensor: + pred_mask = pred_mask.to(dtype=target.dtype) + return target * (1 - pred_mask) + pred * pred_mask + + +def _extract_view_slice( + volume: Tensor, + view: str, + slice_idx: int | Mapping[str, int] | None = None, +) -> Tensor: + if isinstance(slice_idx, Mapping): + slice_idx = slice_idx.get(view) + + if view in ("sagittal", "saggital"): + idx = _resolve_slice_idx(volume.shape[0], slice_idx) + return volume[idx, :, :].transpose(0, 1).flip(0) + if view == "axial": + idx = _resolve_slice_idx(volume.shape[2], slice_idx) + return volume[:, :, idx].transpose(0, 1).flip(0) + if view == "coronal": + idx = _resolve_slice_idx(volume.shape[1], slice_idx) + return volume[:, idx, :].transpose(0, 1).flip(0) + raise ValueError(f"unknown MRI view {view!r}") + + +def _resolve_slice_idx(size: int, slice_idx: int | None) -> int: + idx = size // 2 if slice_idx is None else int(slice_idx) + if idx < 0: + idx += size + if idx < 0 or idx >= size: + raise IndexError(f"slice index {idx} is out of bounds for axis with size {size}") + return idx + + +def _intensity_limits(volume: Tensor, mask: Tensor | None = None) -> tuple[float, float]: + values = volume[mask] if mask is not None and mask.any() else volume.flatten() + values = values[torch.isfinite(values)] + if values.numel() == 0: + return 0.0, 1.0 + if values.numel() < 32: + vmin = values.min() + vmax = values.max() + else: + vmin, vmax = torch.quantile(values, torch.tensor([0.005, 0.995])) + if torch.isclose(vmin, vmax): + delta = max(abs(float(vmin)) * 0.05, 1.0) + return float(vmin) - delta, float(vmax) + delta + return float(vmin), float(vmax) + + +def _apply_display_mask(image: Tensor, mask: Tensor | None, fill_value: float) -> Tensor: + if mask is None: + return image + return torch.where(mask, image, torch.full_like(image, fill_value)) + + +def _masked_input_display( + image: Tensor, + pred_mask: Tensor, + img_mask: Tensor | None, + fill_value: float, +) -> Tensor: + display = torch.where(pred_mask, torch.full_like(image, fill_value), image) + return _apply_display_mask(display, img_mask, fill_value) + + +def _crop_view_items(view_items: list[dict]) -> None: + for item in view_items: + mask = item["img_mask"] + if mask is None: + mask = item["actual"] != item["actual"].min() + row_slice, col_slice = _content_crop(mask, item["patch_rc"]) + for key in ("target", "composite", "actual", "mask"): + item[key] = item[key][row_slice, col_slice] + if item["img_mask"] is not None: + item["img_mask"] = item["img_mask"][row_slice, col_slice] + + +def _content_crop(mask: Tensor, patch_size: tuple[int, int]) -> tuple[slice, slice]: + mask = mask.detach().cpu() > 0 + if not mask.any(): + return slice(None), slice(None) + + rows, cols = mask.nonzero(as_tuple=True) + patch_h, patch_w = patch_size + height, width = mask.shape + row0 = max((int(rows.min()) // patch_h - 1) * patch_h, 0) + row1 = min((int(rows.max()) // patch_h + 2) * patch_h, height) + col0 = max((int(cols.min()) // patch_w - 1) * patch_w, 0) + col1 = min((int(cols.max()) // patch_w + 2) * patch_w, width) + return slice(row0, row1), slice(col0, col1) + + +def _as_3tuple(value: int | tuple[int, int, int]) -> tuple[int, int, int]: + if isinstance(value, int): + return (value, value, value) + if len(value) != 3: + raise ValueError(f"expected a 3-tuple patch size, got {value!r}") + return tuple(int(v) for v in value) + + +def _view_patch_size(view: str, patch_size: tuple[int, int, int]) -> tuple[int, int]: + p_x, p_y, p_z = patch_size + if view in ("sagittal", "saggital"): + return p_z, p_y + if view == "axial": + return p_y, p_x + if view == "coronal": + return p_z, p_x + raise ValueError(f"unknown MRI view {view!r}") + + +def _make_figure_canvas( + view_items: list[dict[str, Tensor | str | tuple[int, int]]], + figsize: tuple[float, float] | None = None, +): + dpi = 160 + left = 58 + right = 6 + top = 18 + bottom = 8 + row_gap = 14 + col_gap = 8 + widths = [int(item["target"].shape[1]) for item in view_items] + heights = [int(item["target"].shape[0]) for item in view_items] + row_h = max(heights) + num_rows = 3 + fig_w = left + right + sum(widths) + col_gap * (len(widths) - 1) + fig_h = top + bottom + row_h * num_rows + row_gap * (num_rows - 1) + + scale = 1.35 + if figsize is not None: + requested_w = figsize[0] * dpi + requested_h = figsize[1] * dpi + scale = max(requested_w / fig_w, requested_h / fig_h) + figsize = (fig_w * scale / dpi, fig_h * scale / dpi) + fig = plt.figure(figsize=figsize, dpi=dpi, facecolor="#0b0f14") + + axes = [[] for _ in range(num_rows)] + col_centers = [] + x = left + for width, height in zip(widths, heights): + col_centers.append((x + width / 2) / fig_w) + ys = [ + bottom + (num_rows - row - 1) * (row_h + row_gap) + (row_h - height) / 2 + for row in range(num_rows) + ] + for row, y in enumerate(ys): + axes[row].append( + fig.add_axes( + [ + x / fig_w, + y / fig_h, + width / fig_w, + height / fig_h, + ], + facecolor="black", + ) + ) + x += width + col_gap + row_centers = [ + (bottom + (num_rows - row - 1) * (row_h + row_gap) + row_h / 2) / fig_h + for row in range(num_rows) + ] + layout = { + "col_centers": col_centers, + "row_centers": row_centers, + "label_x": (left - 8) / fig_w, + "title_y": (fig_h - top / 2) / fig_h, + } + + return fig, axes, layout + + +def _style_axis(ax) -> None: + ax.set_xticks([]) + ax.set_yticks([]) + for spine in ax.spines.values(): + spine.set_visible(False) + + +def _draw_patch_boxes( + ax, + mask: Tensor, + patch_size: tuple[int, int], + color: str = "#facc15", +) -> None: + for col, row, width, height in _patch_rectangles(mask, patch_size): + ax.add_patch( + patches.Rectangle( + (col - 0.5, row - 0.5), + width, + height, + fill=False, + edgecolor=color, + linewidth=0.75, + alpha=0.95, + ) + ) + + +def _patch_rectangles( + mask: Tensor, + patch_size: tuple[int, int], +) -> list[tuple[int, int, int, int]]: + mask = mask.detach().cpu() > 0 + patch_h, patch_w = patch_size + height, width = mask.shape + rectangles = [] + for row in range(0, height, patch_h): + box_h = min(patch_h, height - row) + for col in range(0, width, patch_w): + box_w = min(patch_w, width - col) + if mask[row : row + box_h, col : col + box_w].any(): + rectangles.append((col, row, box_w, box_h)) + return rectangles diff --git a/finetune/fomo_tune_baseline/output/task1/model/config.yaml b/finetune/fomo_tune_baseline/output/task1/model/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..5dca875152bb8d3adf4bf758bd4681bf19e1d83e --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/model/config.yaml @@ -0,0 +1,8 @@ +task: task1 +ckpt_path: /data/mihir-stuff/smri-pretrained/pretrain_full_90_10_h100/checkpoint-last.pth +modalities: +- dwi_b1000 +output_root: experiments/fomo_tune_baseline/output +name: task1 +device: cuda +seed: 4466 diff --git a/finetune/fomo_tune_baseline/output/task1/model/head.joblib b/finetune/fomo_tune_baseline/output/task1/model/head.joblib new file mode 100644 index 0000000000000000000000000000000000000000..90ed37cfc3dfd4a13e1f754990808e15c9b1f7e0 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task1/model/head.joblib @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f4a4826f5ea9f3f74279b798caf6a5970a3c63b02b70c4d0828ee5325d29e0dd +size 445215 diff --git a/finetune/fomo_tune_baseline/output/task3/build/model/backbone.pth b/finetune/fomo_tune_baseline/output/task3/build/model/backbone.pth new file mode 100644 index 0000000000000000000000000000000000000000..793b42f5cea0b055534c29cd18c21f3200873219 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task3/build/model/backbone.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aacf2582e9464e7fd4bd3d5ed41a5b9aa770e5b04a24b2a82c1c4ba5fb5f6e7a +size 1389681124 diff --git a/finetune/fomo_tune_baseline/output/task3/build/model/head.joblib b/finetune/fomo_tune_baseline/output/task3/build/model/head.joblib new file mode 100644 index 0000000000000000000000000000000000000000..7ade23f9dc38a75d598a1a8dfeb4930ff2d4b1a9 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task3/build/model/head.joblib @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e263a95e347fbaf86b4edc44717e55891711bf37a7a1309c7f04753af46a4dbf +size 29906 diff --git a/finetune/fomo_tune_baseline/output/task3/build/smri_mae/config/default_pretrain.yaml b/finetune/fomo_tune_baseline/output/task3/build/smri_mae/config/default_pretrain.yaml new file mode 100644 index 0000000000000000000000000000000000000000..c16290b1e2c24e4d206f7e32a7b9c2ee5266067e --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task3/build/smri_mae/config/default_pretrain.yaml @@ -0,0 +1,98 @@ +# Name of the run. Used for output directory suffix and wandb. +name: pretrain + +# Description of the run. Goes in wandb notes. +notes: null + +# Root output directory. +# The run writes to checkpoints/ when name is set. +output_dir: checkpoints + +# Standard 3D structural MRI volume size. +img_size: [208, 240, 208] +patch_size: 8 + +# Masking. +mask_ratio: 0.80 +pred_mask_ratio: null +pad_to_multiple: 32 + +# Model. +model: mae_vit_large +model_kwargs: + # target normalization: null/none, global, slice, or patch. + target_norm: none + + no_decode_pos: false + mask_drop_scale: false + + class_token: true + reg_tokens: 0 + no_embed_class: false + + decoder_depth: 4 + drop_path_rate: 0.0 + +# Datasets. +datasets: + fomo_train: + url: datasets/FOMO_with_dwi/shard.{000000..001620}.tar + samples_per_epoch: 243200 + shuffle: true + buffer_size: 8000 + drop_last: true + + fomo_val: + url: datasets/FOMO_with_dwi/shard.{001621..001800}.tar + samples_per_epoch: 26880 + shuffle: false + buffer_size: 1000 + drop_last: true + +train_dataset: fomo_train +eval_datasets: + - fomo_val + +# Data loader. +num_workers: 4 +prefetch_factor: 2 +presend_cuda: true + +# Optimization. +epochs: 100 +batch_size: 64 +accum_iter: 1 + +base_lr: 0.001 +min_lr: 1e-6 +warmup_epochs: 10 +weight_decay: 0.05 +betas: [0.9, 0.95] +clip_grad: 1.0 + +amp: true +amp_dtype: bfloat16 + +# Checkpointing. +ckpt: null +resume: false +auto_resume: true +start_epoch: 0 +checkpoint_period: 10 +max_checkpoints: 5 + +# Evaluation. +eval_period: 10 + +# Sync checkpoints to an R2 bucket using the AWS CLI. Set to a URL to enable. +r2_sync: null + +device: cuda +distributed: false +seed: 7338 +eval_seed: 7338 +debug: false + +wandb: false +wandb_entity: null +wandb_project: smri-fm diff --git a/finetune/fomo_tune_baseline/output/task3/build/smri_mae/masking.py b/finetune/fomo_tune_baseline/output/task3/build/smri_mae/masking.py new file mode 100644 index 0000000000000000000000000000000000000000..28f9d53cbbd8d32a3923f9b1a6655c92cd526bf7 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task3/build/smri_mae/masking.py @@ -0,0 +1,80 @@ +import torch +from jaxtyping import Float, Int +from torch import Tensor + + +def pad_patch_mask( + patch_mask: Float[Tensor, "B N"], + mask_ratio: float, + shuffle: bool = False, + generator: torch.Generator | None = None, + pad_to_multiple: int | None = None, +) -> tuple[Float[Tensor, "B N"], Int[Tensor, "B L"], Tensor]: + """ + Select each row's own mask-ratio count, then pad ids to the batch max length. + + Returns: + - selected patch mask [B, N] + - padded selected patch ids [B, Lpad] + - token mask [B, Lpad], true for real ids and false for padding + """ + if not 0.0 <= mask_ratio <= 1.0: + raise ValueError(f"mask_ratio must be in [0, 1], got {mask_ratio}") + + B, N = patch_mask.shape + device = patch_mask.device + patch_mask = patch_mask.to(dtype=torch.bool) + + valid_counts = patch_mask.sum(dim=1) + num_keep = torch.floor(valid_counts.to(torch.float64) * (1.0 - mask_ratio)).to(torch.long) + if not shuffle: + selected = patch_mask & (patch_mask.cumsum(dim=1) <= num_keep.unsqueeze(1)) + padded_ids, token_mask = patch_ids_from_mask( + selected, + pad_to_multiple=pad_to_multiple, + ) + return selected, padded_ids, token_mask + + # One masked sort directly produces random valid IDs. The previous + # shuffle/select/inverse-shuffle path required two full argsorts plus a + # dynamic nonzero/scatter solely to recover the same selected set. + noise = torch.rand(B, N, generator=generator, device=device) + noise.masked_fill_(~patch_mask, torch.inf) + shuffled_ids = torch.argsort(noise, dim=1) + + max_count = int(num_keep.max().item()) + if pad_to_multiple is not None: + if pad_to_multiple <= 0: + raise ValueError(f"pad_to_multiple must be positive, got {pad_to_multiple}") + max_count = (max_count + pad_to_multiple - 1) // pad_to_multiple * pad_to_multiple + padded_ids = shuffled_ids[:, :max_count] + token_mask = torch.arange(max_count, device=device).unsqueeze(0) < num_keep.unsqueeze(1) + selected = torch.zeros_like(patch_mask).scatter_(1, padded_ids, token_mask) + return selected, padded_ids, token_mask + + +def patch_ids_from_mask( + patch_mask: Tensor, + pad_to_multiple: int | None = None, +) -> tuple[Int[Tensor, "B L"], Tensor]: + """Return optionally rounded patch IDs and their token-validity mask.""" + if pad_to_multiple is not None and pad_to_multiple <= 0: + raise ValueError(f"pad_to_multiple must be positive, got {pad_to_multiple}") + + patch_mask = patch_mask.to(dtype=torch.bool) + B, N = patch_mask.shape + device = patch_mask.device + counts = patch_mask.sum(dim=1) + max_count = int(counts.max().item()) + if pad_to_multiple is not None: + max_count = (max_count + pad_to_multiple - 1) // pad_to_multiple * pad_to_multiple + + patch_ids = torch.zeros((B, max_count), dtype=torch.long, device=device) + token_mask = torch.arange(max_count, device=device).unsqueeze(0) < counts.unsqueeze(1) + if max_count == 0: + return patch_ids, token_mask + + batch_ids, selected_ids = patch_mask.nonzero(as_tuple=True) + slot_ids = patch_mask.cumsum(dim=1)[batch_ids, selected_ids].to(torch.long) - 1 + patch_ids[batch_ids, slot_ids] = selected_ids + return patch_ids, token_mask diff --git a/finetune/fomo_tune_baseline/output/task3/build/smri_mae/modules.py b/finetune/fomo_tune_baseline/output/task3/build/smri_mae/modules.py new file mode 100644 index 0000000000000000000000000000000000000000..828aecb1bd57182f1690feb4e943395db1bdd515 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task3/build/smri_mae/modules.py @@ -0,0 +1,453 @@ +# This source code is licensed under the Apache License, Version 2.0 +# +# References: +# capi: https://github.com/facebookresearch/capi/blob/main/model.py +# timm: https://github.com/huggingface/pytorch-image-models/blob/v1.0.20/timm/models/vision_transformer.py +# vjepa2: https://github.com/facebookresearch/vjepa2/blob/main/src/models/utils/pos_embs.py + +import math +from functools import partial +from typing import NamedTuple, Type + +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch import Tensor +from einops import rearrange +from jaxtyping import Float, Int +from timm.layers import DropPath, to_3tuple + +Layer = Type[nn.Module] + + +class JaggedBatch(NamedTuple): + """Sequence boundaries and cached launch metadata for jagged attention.""" + + offsets: Tensor + max_seqlen: int + + @classmethod + def from_mask(cls, mask: Tensor) -> "JaggedBatch": + mask = mask.to(dtype=torch.bool) + counts = mask.sum(dim=1) + return cls( + offsets=F.pad(counts.cumsum(dim=0), (1, 0)), + max_seqlen=mask.shape[1], + ) + + def as_nested(self, tokens: Tensor) -> Tensor: + # Cached conservative bounds avoid min/max reductions and GPU-to-CPU + # synchronization when Flash SDPA inspects the jagged sequence lengths. + return torch.nested.nested_tensor_from_jagged( + tokens, + self.offsets, + min_seqlen=1, + max_seqlen=self.max_seqlen, + ).transpose(1, 2) + + +def unpack_tokens(tokens: Tensor, token_mask: Tensor) -> Tensor: + """Restore packed values to a padded batch, filling invalid slots with zero.""" + output = tokens.new_zeros((*token_mask.shape, *tokens.shape[1:])) + return output.index_put((token_mask,), tokens) + + +def jagged_scaled_dot_product_attention( + query: Tensor, + key: Tensor, + value: Tensor, + jagged_batch: JaggedBatch, +) -> Tensor: + """Run SDPA on a packed batch of variable-length sequences.""" + output_jagged = F.scaled_dot_product_attention( + jagged_batch.as_nested(query), + jagged_batch.as_nested(key), + jagged_batch.as_nested(value), + ) + return output_jagged.transpose(1, 2).values() + + +class Attention(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + qkv_bias: bool = False, + proj_bias: bool = False, + ) -> None: + super().__init__() + self.num_heads = num_heads + self.head_dim = dim // num_heads + self.qkv = nn.Linear(dim, 3 * dim, bias=qkv_bias) + self.proj = nn.Linear(dim, dim, bias=proj_bias) + + def extra_repr(self): + return f"num_heads={self.num_heads}" + + def forward( + self, + x: Float[Tensor, "L D"], + jagged_batch: JaggedBatch, + ) -> Float[Tensor, "L D"]: + L, D = x.shape + h, dh = self.num_heads, self.head_dim + + qkv = self.qkv(x).reshape(L, 3, h, dh) + q, k, v = qkv.unbind(1) + + x = jagged_scaled_dot_product_attention( + q, + k, + v, + jagged_batch=jagged_batch, + ) + x = x.reshape(L, D) + x = self.proj(x) + return x + + +class Mlp(nn.Module): + def __init__( + self, + dim: int, + mlp_ratio: int | float = 4, + bias: bool = False, + ) -> None: + super().__init__() + hidden_features = int(dim * mlp_ratio) + self.fc1 = nn.Linear(dim, hidden_features, bias=bias) + self.act = nn.GELU() + self.fc2 = nn.Linear(hidden_features, dim, bias=bias) + + def forward(self, x: Float[Tensor, "... D"]) -> Float[Tensor, "... D"]: + x = self.fc1(x) + x = self.act(x) + x = self.fc2(x) + return x + + +# timm default eps=1e-6 +LayerNorm = partial(nn.LayerNorm, eps=1e-6) + + +class Block(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + qkv_bias: bool = False, + proj_bias: bool = False, + mlp_ratio: int | float = 4, + drop_path: float = 0.0, + norm_layer: Layer = LayerNorm, + ) -> None: + super().__init__() + self.norm1 = norm_layer(dim) + self.attn = Attention( + dim=dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + ) + self.drop_path1 = DropPath(drop_path) if drop_path > 0 else nn.Identity() + + self.norm2 = norm_layer(dim) + self.mlp = Mlp( + dim=dim, + mlp_ratio=mlp_ratio, + bias=proj_bias, + ) + self.drop_path2 = DropPath(drop_path) if drop_path > 0 else nn.Identity() + + def forward( + self, + x: Float[Tensor, "L D"], + jagged_batch: JaggedBatch, + ) -> Float[Tensor, "L D"]: + x = x + self.drop_path1( + self.attn( + self.norm1(x), + jagged_batch=jagged_batch, + ) + ) + x = x + self.drop_path2(self.mlp(self.norm2(x))) + return x + + +# Patching and position embedding modules + + +class Patchify3D(nn.Module): + def __init__( + self, + img_size: int | tuple[int, int, int], + patch_size: int | tuple[int, int, int], + in_chans: int = 3, + ) -> None: + super().__init__() + self.img_size = to_3tuple(img_size) + self.patch_size = to_3tuple(patch_size) + self.in_chans = in_chans + + T, H, W = self.img_size + p_t, p_h, p_w = self.patch_size + if T % p_t or H % p_h or W % p_w: + raise ValueError( + f"img_size {self.img_size} must be divisible by patch_size {self.patch_size}" + ) + self.grid_size = (T // p_t, H // p_h, W // p_w) + self.num_patches = math.prod(self.grid_size) + self.patch_dim = in_chans * math.prod(self.patch_size) + + def forward(self, x: Float[Tensor, "B C T H W"]) -> Float[Tensor, "B N P"]: + x = patchify3d(x, self.patch_size) + return x + + def unpatchify(self, x: Float[Tensor, "B N P"]) -> Float[Tensor, "B C T H W"]: + x = unpatchify3d(x, patch_size=self.patch_size, img_size=self.img_size) + return x + + def extra_repr(self): + return f"{self.img_size}, {self.patch_size}, in_chans={self.in_chans}" + + +def patchify3d(x: Tensor, patch_size: tuple[int, int, int]) -> Tensor: + p_t, p_h, p_w = to_3tuple(patch_size) + B, C, T, H, W = x.shape + x = rearrange(x, "b c (t u) (h p) (w q) -> b (t h w) (c u p q)", u=p_t, p=p_h, q=p_w) + return x + + +def unpatchify3d( + x: Tensor, + patch_size: tuple[int, int, int], + img_size: tuple[int, int, int], +) -> Tensor: + B, N, P = x.shape + p_t, p_h, p_w = to_3tuple(patch_size) + T, H, W = to_3tuple(img_size) + x = rearrange( + x, + "b (t h w) (c u p q) -> b c (t u) (h p) (w q)", + t=T // p_t, + h=H // p_h, + w=W // p_w, + u=p_t, + p=p_h, + q=p_w, + ) + return x + + +class AbsolutePosEmbed(nn.Module): + def __init__(self, embed_dim: int, grid_size: tuple[int, ...]) -> None: + super().__init__() + self.embed_dim = embed_dim + self.grid_size = grid_size + self.num_patches = math.prod(grid_size) + + self.weight = nn.Parameter(torch.empty(self.num_patches, embed_dim)) + self.reset_parameters() + + def reset_parameters(self): + nn.init.trunc_normal_(self.weight, std=0.02) + + def forward( + self, + x: Float[Tensor, "B L D"], + pos_ids: Int[Tensor, "B L"] | None = None, + ) -> Float[Tensor, "B L D"]: + x = apply_pos_embed(x, self.weight, pos_ids=pos_ids) + return x + + def extra_repr(self): + return f"{self.embed_dim}, {self.grid_size}" + + +class SeparablePosEmbed(nn.Module): + def __init__(self, embed_dim: int, grid_size: tuple[int, ...]) -> None: + super().__init__() + self.embed_dim = embed_dim + self.grid_size = grid_size + self.num_patches = math.prod(grid_size) + + N_t, *grid_size_spatial = grid_size + N_s = math.prod(grid_size_spatial) + self.weight_spatial = nn.Parameter(torch.empty(1, N_s, embed_dim)) + self.weight_temporal = nn.Parameter(torch.empty(N_t, 1, embed_dim)) + self.reset_parameters() + + def reset_parameters(self): + nn.init.trunc_normal_(self.weight_spatial, std=0.02) + nn.init.trunc_normal_(self.weight_temporal, std=0.02) + + def forward( + self, + x: Float[Tensor, "B L D"], + pos_ids: Int[Tensor, "B L"] | None = None, + ) -> Float[Tensor, "B L D"]: + B, N, D = x.shape + weight = (self.weight_temporal + self.weight_spatial).flatten(0, 1) # [N, D] + x = apply_pos_embed(x, weight, pos_ids=pos_ids) + return x + + def extra_repr(self): + return f"{self.embed_dim}, {self.grid_size}" + + +class SinCosPosEmbed3D(nn.Module): + def __init__(self, embed_dim: int, grid_size: tuple[int, int, int]) -> None: + super().__init__() + self.embed_dim = embed_dim + self.grid_size = grid_size + self.num_patches = math.prod(grid_size) + + N_t, N_h, N_w = grid_size + weight = get_3d_sincos_pos_embed( + embed_dim=embed_dim, + grid_size=(N_h, N_w), + grid_depth=N_t, + uniform_power=True, + ) + self.weight = nn.Parameter(torch.from_numpy(weight).float(), requires_grad=False) + + def forward( + self, + x: Float[Tensor, "B L D"], + pos_ids: Int[Tensor, "B L"] | None = None, + ) -> Float[Tensor, "B L D"]: + x = apply_pos_embed(x, self.weight, pos_ids=pos_ids) + return x + + def extra_repr(self): + return f"{self.embed_dim}, {self.grid_size}" + + +# sincos pos embed utils from vjepa2, but fixed the confusing meshgrid indexing + + +def get_3d_sincos_pos_embed(embed_dim, grid_size, grid_depth, cls_token=False, uniform_power=False): + """ + grid_size: tuple of int of the grid height and width + grid_depth: int of the grid depth + returns: + pos_embed: [grid_depth*grid_height*grid_width, embed_dim] (w/o cls_token) + or [1+grid_depth*grid_height*grid_width, embed_dim] (w/ cls_token) + """ + grid_d = np.arange(grid_depth, dtype=float) + grid_h = np.arange(grid_size[0], dtype=float) + grid_w = np.arange(grid_size[1], dtype=float) + grid_d, grid_h, grid_w = np.meshgrid(grid_d, grid_h, grid_w, indexing="ij") + + if not uniform_power: + h_embed_dim = embed_dim // 4 + w_embed_dim = embed_dim // 4 + d_embed_dim = embed_dim // 2 + else: + h_embed_dim = w_embed_dim = d_embed_dim = int(np.ceil(embed_dim / 6) * 2) + + emb_h = get_1d_sincos_pos_embed_from_grid(h_embed_dim, grid_h) # (T*H*W, D1) + emb_w = get_1d_sincos_pos_embed_from_grid(w_embed_dim, grid_w) # (T*H*W, D2) + emb_d = get_1d_sincos_pos_embed_from_grid(d_embed_dim, grid_d) # (T*H*W, D3) + pos_embed = np.concatenate([emb_d, emb_h, emb_w], axis=1) + pos_embed = pos_embed[:, :embed_dim] + if cls_token: + pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0) + return pos_embed + + +def get_1d_sincos_pos_embed_from_grid(embed_dim, pos): + """ + embed_dim: output dimension for each position + pos: a list of positions to be encoded: size (M,) + returns: (M, D) + """ + assert embed_dim % 2 == 0 + omega = np.arange(embed_dim // 2, dtype=float) + omega /= embed_dim / 2.0 + omega = 1.0 / 10000**omega # (D/2,) + + pos = pos.reshape(-1) # (M,) + out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product + + emb_sin = np.sin(out) # (M, D/2) + emb_cos = np.cos(out) # (M, D/2) + + emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D) + return emb + + +def apply_pos_embed( + x: Float[Tensor, "B L D"], + weight: Float[Tensor, "N D"], + pos_ids: Int[Tensor, "B L"] | None = None, +) -> Float[Tensor, "B L D"]: + B, L, D = x.shape + weight = weight.expand(B, -1, -1) + if pos_ids is not None: + weight = weight.gather(1, pos_ids.unsqueeze(-1).expand(-1, -1, D)) + x = x + weight + return x + + +# (masked) normalization used for MAE target normalization + + +class Normalize(nn.Module): + def __init__( + self, + grid_size: tuple[int, ...], + dim: int | tuple[int, ...] | None = -1, + eps: float = 1e-6, + ) -> None: + super().__init__() + self.grid_size = grid_size + self.dim = dim + self.eps = eps + + def forward(self, x: Tensor, mask: Tensor | None = None) -> tuple[Tensor, Tensor, Tensor]: + """ + Normalize input sequence along dim(s) after reshaping to grid. + Returns tuple of (x, mean, std). + """ + B, N, D = x.shape + x = x.reshape((B, *self.grid_size, D)) + if mask is not None: + mask = mask.reshape((B, *self.grid_size, D)) + x, mean, std = masked_normalize(x, mask, dim=self.dim, eps=self.eps) + else: + x, mean, std = normalize(x, dim=self.dim, eps=self.eps) + mean = mean.expand_as(x).reshape(B, N, D) + std = std.expand_as(x).reshape(B, N, D) + x = x.reshape(B, N, D) + return x, mean, std + + def extra_repr(self): + return f"{self.grid_size}, dim={self.dim}" + + +def masked_normalize( + x: Tensor, + mask: Tensor, + dim: int | tuple[int, ...] | None = -1, + eps: float = 1e-6, +) -> tuple[Tensor, Tensor, Tensor]: + num_obs = mask.sum(dim=dim, keepdim=True).clamp(min=1) + mean = (mask * x).sum(dim=dim, keepdim=True) / num_obs + var = (mask * (x - mean) ** 2).sum(dim=dim, keepdim=True) / num_obs + std = (var + eps) ** 0.5 + x = mask * (x - mean) / std + return x, mean, std + + +def normalize( + x: Tensor, + dim: int | tuple[int, ...] | None = -1, + eps: float = 1e-6, +) -> tuple[Tensor, Tensor, Tensor]: + mean = x.mean(dim=dim, keepdim=True) + var = torch.var(x, dim=dim, keepdim=True, unbiased=False) + std = (var + eps) ** 0.5 + x = (x - mean) / std + return x, mean, std diff --git a/finetune/fomo_tune_baseline/output/task3/build/smri_mae/utils.py b/finetune/fomo_tune_baseline/output/task3/build/smri_mae/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..0f6274f813429552c9ddfccfe4033ae677b52342 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task3/build/smri_mae/utils.py @@ -0,0 +1,581 @@ +# Copyright (c) Sophont, Inc +# This source code is licensed under the Apache License, Version 2.0 +# +# References: +# deit: https://github.com/facebookresearch/deit/blob/main/utils.py +# beit3: https://github.com/microsoft/unilm/blob/master/beit3/utils.py +# capi: https://github.com/facebookresearch/capi/blob/main/utils.py +# dinov2: https://github.com/facebookresearch/dinov2/blob/main/dinov2/utils/param_groups.py +# timm: https://github.com/huggingface/pytorch-image-models/blob/main/timm/utils/cuda.py +# dino: https://github.com/facebookresearch/dino/blob/main/utils.py + +import datetime +import inspect +import math +import os +import random +import subprocess +import time +from collections import defaultdict, deque +from omegaconf import OmegaConf +from pathlib import Path + +import numpy as np +import torch +import torch.distributed as dist +import torch.nn as nn +from torch import Tensor +from torch.amp import GradScaler +from torch.optim import Optimizer + + +# these very useful utils copied from deit with only minor changes +# thanks to the original authors, wherever you are + + +def configure_flash_sdpa() -> None: + """Use Flash Attention exclusively for CUDA SDPA.""" + torch.backends.cuda.enable_flash_sdp(True) + torch.backends.cuda.enable_mem_efficient_sdp(False) + torch.backends.cuda.enable_math_sdp(False) + torch.backends.cuda.enable_cudnn_sdp(False) + print("SDPA backend: flash") + + +class SmoothedValue: + """Track a series of values and provide access to smoothed values over a + window or the global series average. + """ + + def __init__(self, window_size=20, fmt=None): + if fmt is None: + fmt = "{median:.4f} ({global_avg:.4f})" + self.deque = deque(maxlen=window_size) + self.total = 0.0 + self.count = 0 + self.fmt = fmt + + def update(self, value, n=1): + value = float(value) + if math.isfinite(value): + self.deque.append(value) + self.count += n + self.total += value * n + + def synchronize_between_processes(self): + """ + Warning: does not synchronize the deque! + """ + if not is_dist_avail_and_initialized(): + return + t = torch.tensor([self.count, self.total], dtype=torch.float64, device="cuda") + dist.barrier() + dist.all_reduce(t) + t = t.tolist() + self.count = int(t[0]) + self.total = t[1] + + @property + def median(self): + if not self.count: + return float("nan") + d = torch.tensor(list(self.deque)) + return d.median().item() + + @property + def avg(self): + if not self.count: + return float("nan") + d = torch.tensor(list(self.deque), dtype=torch.float32) + return d.mean().item() + + @property + def global_avg(self): + if not self.count: + return float("nan") + return self.total / self.count + + @property + def max(self): + if not self.count: + return float("nan") + return max(self.deque) + + @property + def value(self): + if not self.count: + return float("nan") + return self.deque[-1] + + def __str__(self): + return self.fmt.format( + median=self.median, + avg=self.avg, + global_avg=self.global_avg, + max=self.max, + value=self.value, + ) + + +class MetricLogger: + def __init__(self, delimiter="\t"): + self.meters = defaultdict(SmoothedValue) + self.delimiter = delimiter + + def update(self, **kwargs): + for k, v in kwargs.items(): + if v is None: + continue + if isinstance(v, (torch.Tensor, np.generic)): + v = v.item() + assert isinstance(v, (float, int)) + self.meters[k].update(v) + + def __getattr__(self, attr): + if attr in self.meters: + return self.meters[attr] + if attr in self.__dict__: + return self.__dict__[attr] + raise AttributeError("'{}' object has no attribute '{}'".format(type(self).__name__, attr)) + + def __str__(self): + loss_str = [] + for name, meter in self.meters.items(): + loss_str.append("{}: {}".format(name, str(meter))) + return self.delimiter.join(loss_str) + + def synchronize_between_processes(self): + for meter in self.meters.values(): + meter.synchronize_between_processes() + + def add_meter(self, name, meter): + self.meters[name] = meter + + def log_every(self, iterable, print_freq, header=None, total_steps=None): + i = 0 + total_steps = total_steps or len(iterable) + if not header: + header = "" + start_time = time.time() + end = time.time() + iter_time = SmoothedValue(fmt="{avg:.4f}") + data_time = SmoothedValue(fmt="{avg:.4f}") + space_fmt = ":" + str(len(str(total_steps))) + "d" + log_msg = [ + header, + "[{0" + space_fmt + "}/{1}]", + "eta: {eta}", + "{meters}", + "time: {time}", + "data: {data}", + ] + if torch.cuda.is_available(): + log_msg.append("max mem: {memory:.0f}") + log_msg = self.delimiter.join(log_msg) + MB = 1024.0 * 1024.0 + for obj in iterable: + if i >= total_steps: + break + data_time.update(time.time() - end) + yield obj + iter_time.update(time.time() - end) + if i % print_freq == 0 or i == total_steps - 1: + eta_seconds = iter_time.global_avg * (total_steps - i) + eta_string = str(datetime.timedelta(seconds=int(eta_seconds))) + if torch.cuda.is_available(): + print( + log_msg.format( + i, + total_steps, + eta=eta_string, + meters=str(self), + time=str(iter_time), + data=str(data_time), + memory=torch.cuda.max_memory_allocated() / MB, + ) + ) + + else: + print( + log_msg.format( + i, + total_steps, + eta=eta_string, + meters=str(self), + time=str(iter_time), + data=str(data_time), + ) + ) + i += 1 + end = time.time() + total_time = time.time() - start_time + total_time_str = str(datetime.timedelta(seconds=int(total_time))) + print( + "{} Total time: {} ({:.4f} s / it)".format( + header, total_time_str, total_time / total_steps + ) + ) + + +def setup_for_distributed(log_path=None): + """ + This function disables printing when not in master process + """ + import builtins as __builtin__ + + builtin_print = __builtin__.print + + is_master = is_main_process() + + def print(*args, **kwargs): + force = kwargs.pop("force", False) + if is_master or force: + builtin_print(*args, **kwargs) + # tee to log file + if log_path and "file" not in kwargs: + with open(log_path, "a") as f: + builtin_print(*args, file=f, **kwargs) + + __builtin__.print = print + + +def is_dist_avail_and_initialized(): + if not dist.is_available(): + return False + if not dist.is_initialized(): + return False + return True + + +def get_world_size(): + if not is_dist_avail_and_initialized(): + return 1 + return dist.get_world_size() + + +def get_rank(): + if not is_dist_avail_and_initialized(): + return 0 + return dist.get_rank() + + +def is_main_process(): + return get_rank() == 0 + + +def save_on_master(obj, path, *args, **kwargs): + if is_main_process(): + path = Path(path) + tmp_path = path.with_name(f".{path.name}.tmp-{os.getpid()}") + try: + torch.save(obj, tmp_path, *args, **kwargs) + os.replace(tmp_path, path) + except Exception: + tmp_path.unlink(missing_ok=True) + raise + + +def init_distributed_mode(args): + # removed slurm block, can add if we use slurm + if "RANK" in os.environ and "WORLD_SIZE" in os.environ: + args.rank = int(os.environ["RANK"]) + args.world_size = int(os.environ["WORLD_SIZE"]) + args.gpu = int(os.environ["LOCAL_RANK"]) + else: + args.distributed = False + return + + args.distributed = True + + torch.cuda.set_device(args.gpu) + args.dist_backend = "nccl" + print(f"| distributed init (rank {args.rank})") + torch.distributed.init_process_group( + backend=args.dist_backend, + world_size=args.world_size, + rank=args.rank, + device_id=args.gpu, + ) + torch.distributed.barrier() + + +# checkpoint saving utils adapted from beit3 + + +def capture_rng_state() -> dict: + state = {"torch": torch.get_rng_state()} + if torch.cuda.is_available(): + state["cuda"] = torch.cuda.get_rng_state() + return state + + +def restore_rng_state(state: dict) -> None: + torch.set_rng_state(state["torch"]) + if "cuda" in state and torch.cuda.is_available(): + torch.cuda.set_rng_state(state["cuda"]) + + +def _all_rank_rng_states() -> list[dict]: + local_state = capture_rng_state() + if not is_dist_avail_and_initialized(): + return [local_state] + states = [None] * get_world_size() + dist.all_gather_object(states, local_state) + return states + + +def save_model(args, epoch, model_without_ddp, optimizer, loss_scaler): + output_dir = Path(args.output_dir) + checkpoint_path = output_dir / f"checkpoint-{epoch:05d}.pth" + last_checkpoint_path = output_dir / "checkpoint-last.pth" + if epoch % args.checkpoint_period != 0 and epoch != args.epochs - 1: + return + + to_save = { + "model": model_without_ddp.state_dict(), + "optimizer": optimizer.state_dict(), + "epoch": epoch, + "scaler": None if loss_scaler is None else loss_scaler.state_dict(), + "args": OmegaConf.to_container(args), + "rng_states": _all_rank_rng_states(), + } + + print(f"saving checkpoint {last_checkpoint_path}") + save_on_master(to_save, last_checkpoint_path) + print(f"saving checkpoint {checkpoint_path}") + save_on_master(to_save, checkpoint_path) + + if args.max_checkpoints and is_main_process(): + all_checkpoints = sorted(output_dir.glob("checkpoint-[0-9]*.pth")) + del_count = max(0, len(all_checkpoints) - args.max_checkpoints) + for checkpoint_path in all_checkpoints[:del_count]: + print(f"removing checkpoint {checkpoint_path}") + checkpoint_path.unlink() + + +def load_model(args, model_without_ddp, optimizer, loss_scaler): + auto_resume = getattr(args, "auto_resume", True) + output_dir = Path(args.output_dir) + + last_checkpoint_path = output_dir / "checkpoint-last.pth" + if auto_resume and last_checkpoint_path.exists(): + args.ckpt = str(last_checkpoint_path) + args.resume = True + + if args.ckpt: + ckpt = torch.load(args.ckpt, map_location="cpu", weights_only=True) + model_without_ddp.load_state_dict(ckpt["model"]) + print(f"loaded model state from checkpoint {args.ckpt}") + + if args.resume: + optimizer.load_state_dict(ckpt["optimizer"]) + if loss_scaler is not None: + loss_scaler.load_state_dict(ckpt["scaler"]) + args.start_epoch = ckpt["epoch"] + 1 + rng_states = ckpt.get("rng_states") + if rng_states is not None: + if len(rng_states) != get_world_size(): + raise ValueError( + "checkpoint RNG state world size does not match current world size: " + f"{len(rng_states)} != {get_world_size()}" + ) + restore_rng_state(rng_states[get_rank()]) + print(f"restored RNG state for rank {get_rank()}") + print(f"loaded optimizer state, resuming training from {args.start_epoch}") + + +# optimization utils + + +# from capi +class WarmupThenCosine: + def __init__( + self, + base_value: float, + final_value: float, + total_iters: int, + warmup_iters: int = 0, + start_warmup_value: float = 0.0, + freeze_iters: int = 0, + truncate_cos: float = 1.0, + ): + super().__init__() + self.final_value = final_value + self.total_iters = total_iters + + freeze_schedule = np.zeros(freeze_iters) + + warmup_schedule = np.linspace(start_warmup_value, base_value, warmup_iters) + + iters = np.arange(total_iters - warmup_iters - freeze_iters) + schedule = final_value + 0.5 * (base_value - final_value) * ( + 1 + np.cos(np.pi * truncate_cos * iters / len(iters)) + ) + self.schedule = np.concatenate((freeze_schedule, warmup_schedule, schedule)) + assert len(self.schedule) == self.total_iters + + def __getitem__(self, it: int) -> float: + if it >= self.total_iters: + return self.final_value + # cast to float or else it can corrupt the checkpoint + return float(self.schedule[it]) + + +# adapted from timm backward logic +# https://github.com/huggingface/pytorch-image-models/blob/main/timm/utils/cuda.py +def backward_step( + loss: Tensor, + optimizer: Optimizer, + scaler: GradScaler = None, + need_update: bool = True, + max_norm: float | None = None, +) -> Tensor | None: + if scaler is not None: + scaler.scale(loss).backward() + else: + loss.backward() + + if need_update: + if scaler is not None: + scaler.unscale_(optimizer) + + total_norm = clip_grad(optimizer, max_norm) + + if scaler is not None: + scaler.step(optimizer) + scaler.update() + else: + optimizer.step() + optimizer.zero_grad() + else: + total_norm = None + return total_norm + + +def clip_grad(optimizer: Optimizer, max_norm: float | None = None) -> Tensor: + params = [p for group in optimizer.param_groups for p in group["params"]] + if max_norm: + total_norm = nn.utils.clip_grad_norm_(params, max_norm, error_if_nonfinite=False) + else: + grads = [p.grad for p in params if p.grad is not None] + total_norm = nn.utils.get_total_norm(grads, error_if_nonfinite=False) + torch._assert_async(torch.isfinite(total_norm), "non-finite gradient norm") + return total_norm + + +# from dinov2 with some minor changes +def get_param_groups(model, patch_embed_lr_mult=1.0): + # no lr decay, we could add this later if needed + all_params = [] + + for name, param in model.named_parameters(): + if not param.requires_grad: + continue + d = {"param": param, "lr_multiplier": 1.0, "wd_multiplier": 1.0, "name": name} + + if name.endswith(".bias") or "norm" in name or "gamma" in name: + d["wd_multiplier"] = 0.0 + + if "patch_embed" in name: + d["lr_multiplier"] = d["lr_multiplier"] * patch_embed_lr_mult + + all_params.append(d) + + param_groups = _fuse_param_groups(all_params) + return param_groups + + +def _fuse_param_groups(all_param_groups): + fused_param_groups = defaultdict(lambda: {"params": []}) + for d in all_param_groups: + keys = sorted(set(d.keys()) - {"param", "name"}) + identifier = "_".join(f"{k}{d[k]}" for k in keys) + for k in keys: + fused_param_groups[identifier][k] = d[k] + fused_param_groups[identifier]["params"].append(d["param"]) + + param_groups = list(fused_param_groups.values()) + return param_groups + + +def update_lr(param_groups, lr: float): + for group in param_groups: + group["lr"] = lr * group["lr_multiplier"] + + +def update_wd(param_groups, weight_decay: float | None = None): + for group in param_groups: + group["weight_decay"] = weight_decay * group["wd_multiplier"] + + +# moving data to cuda utils copied from capi +# added device argument + + +def send_data(x, device=None, dtype_map=None): + if device is None: + device = torch.device("cuda") + else: + device = torch.device(device) + + if isinstance(x, torch.Tensor): + dtype = dtype_map.get(x.dtype) if dtype_map else None + return x.to(device=device, dtype=dtype, non_blocking=True) + if isinstance(x, dict): + return {k: send_data(v, device=device, dtype_map=dtype_map) for k, v in x.items()} + if isinstance(x, list): + return [send_data(v, device=device, dtype_map=dtype_map) for v in x] + return x + + +def pre_send_to_cuda_wrapper(generator, device=None, dtype_map=None): + """From apex""" + data = None + stream = torch.cuda.Stream(device) + for next_data in generator: + with torch.cuda.stream(stream): + next_data = send_data(next_data, device=device, dtype_map=dtype_map) + if data is not None: + yield data + torch.cuda.current_stream(device).wait_stream(stream) + data = next_data + if data is not None: + yield data + + +# other misc utils + + +# from dino +def get_sha(): + cwd = os.path.dirname(os.path.abspath(__file__)) + + def _run(command): + return subprocess.check_output(command, cwd=cwd).decode("ascii").strip() + + sha = "N/A" + diff = "clean" + branch = "N/A" + try: + sha = _run(["git", "rev-parse", "HEAD"]) + diff = _run(["git", "diff-index", "HEAD"]) + diff = "has uncommitted changes" if diff else "clean" + branch = _run(["git", "rev-parse", "--abbrev-ref", "HEAD"]) + except Exception: + pass + message = f"sha: {sha}, status: {diff}, branch: {branch}" + return message + + +# from timm +def random_seed(seed=42, rank=0): + torch.manual_seed(seed + rank) + np.random.seed(seed + rank) + random.seed(seed + rank) + + +# mine :) +def filter_kwargs(func, kwargs): + sigature = inspect.signature(func) + kwargs = {k: v for k, v in kwargs.items() if k in sigature.parameters} + return kwargs diff --git a/finetune/fomo_tune_baseline/output/task3/build/smri_mae/visualization.py b/finetune/fomo_tune_baseline/output/task3/build/smri_mae/visualization.py new file mode 100644 index 0000000000000000000000000000000000000000..f721fe06e93dc1a2dfb4a679069029bdf55afc6b --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task3/build/smri_mae/visualization.py @@ -0,0 +1,400 @@ +from collections.abc import Mapping +from io import BytesIO + +import torch + +from matplotlib import patches +from matplotlib import pyplot as plt +from PIL import Image +from torch import Tensor + +VIEW_NAMES = { + "sagittal": "Sagittal", + "saggital": "Sagittal", + "axial": "Axial", + "coronal": "Coronal", +} + + +def fig2pil(fig) -> Image.Image: + buffer = BytesIO() + fig.savefig(buffer, format="png", dpi=fig.dpi, facecolor=fig.get_facecolor()) + buffer.seek(0) + image = Image.open(buffer).convert("RGB") + buffer.close() + return image + + +def raw_stats_from_batch(batch: dict) -> tuple[Tensor | None, Tensor | None]: + metas = batch.get("meta") + if not metas: + return None, None + + means = [] + stds = [] + for meta in metas: + try: + mean = meta["raw_mean"] + std = meta["raw_std"] + except (KeyError, TypeError): + return None, None + if mean in ("", None) or std in ("", None): + return None, None + means.append(float(mean)) + stds.append(float(std)) + return torch.tensor(means), torch.tensor(stds) + + +def plot_mask_pred( + target: Tensor, + pred: Tensor, + pred_mask: Tensor | None = None, + img_mask: Tensor | None = None, + sample_idx: int = 0, + channel_idx: int = 0, + slice_idx: int | Mapping[str, int] | None = None, + patch_size: int | tuple[int, int, int] = 16, + views: tuple[str, ...] = ("sagittal", "axial", "coronal"), + cmap: str = "gray", + figsize: tuple[float, float] | None = None, + mask_style: str = "blank", + raw_mean: float | Tensor | None = None, + raw_std: float | Tensor | None = None, +): + target_vol = _select_volume(target, sample_idx=sample_idx, channel_idx=channel_idx) + pred_vol = _select_volume(pred, sample_idx=sample_idx, channel_idx=channel_idx) + if raw_mean is not None and raw_std is not None: + raw_mean = _select_scalar(raw_mean, sample_idx=sample_idx) + raw_std = _select_scalar(raw_std, sample_idx=sample_idx) + target_vol = target_vol * raw_std + raw_mean + pred_vol = pred_vol * raw_std + raw_mean + pred_mask_vol = ( + torch.zeros_like(target_vol) + if pred_mask is None + else _select_volume(pred_mask, sample_idx=sample_idx, channel_idx=channel_idx) > 0 + ) + img_mask_vol = None + if img_mask is not None: + img_mask_vol = _select_volume(img_mask, sample_idx=sample_idx, channel_idx=channel_idx) > 0 + + composite_vol = _prediction_composite(target_vol, pred_vol, pred_mask_vol) + vmin, vmax = _intensity_limits(target_vol, img_mask_vol) + + patch_size = _as_3tuple(patch_size) + view_items = [] + for view in views: + view_key = view.lower() + if view_key not in VIEW_NAMES: + raise ValueError(f"unknown MRI view {view!r}; expected one of {tuple(VIEW_NAMES)}") + target_slice = _extract_view_slice(target_vol, view_key, slice_idx) + composite_slice = _extract_view_slice(composite_vol, view_key, slice_idx) + mask_slice = _extract_view_slice(pred_mask_vol.float(), view_key, slice_idx) > 0 + img_mask_slice = None + if img_mask_vol is not None: + img_mask_slice = _extract_view_slice(img_mask_vol.float(), view_key, slice_idx) > 0 + view_items.append( + { + "key": view_key, + "title": VIEW_NAMES[view_key], + "target": _masked_input_display(target_slice, mask_slice, img_mask_slice, vmin), + "composite": _apply_display_mask(composite_slice, img_mask_slice, vmin), + "actual": _apply_display_mask(target_slice, img_mask_slice, vmin), + "mask": mask_slice, + "img_mask": img_mask_slice, + "patch_rc": _view_patch_size(view_key, patch_size), + } + ) + + _crop_view_items(view_items) + + fig, axes, layout = _make_figure_canvas(view_items, figsize=figsize) + for item, x in zip(view_items, layout["col_centers"]): + fig.text( + x, + layout["title_y"], + item["title"], + ha="center", + va="center", + color="#f8fafc", + fontsize=7, + ) + for label, y in zip(("Masked", "Pred", "Actual"), layout["row_centers"]): + fig.text(layout["label_x"], y, label, ha="right", va="center", color="#cbd5e1", fontsize=6) + + for item, top_ax, middle_ax, bottom_ax in zip(view_items, axes[0], axes[1], axes[2]): + top_ax.imshow( + item["target"], + cmap=cmap, + vmin=vmin, + vmax=vmax, + interpolation="nearest", + origin="upper", + ) + if mask_style == "boxes": + _draw_patch_boxes(top_ax, item["mask"], item["patch_rc"]) + elif mask_style != "blank": + raise ValueError("mask_style must be 'blank' or 'boxes'") + _style_axis(top_ax) + + middle_ax.imshow( + item["composite"], + cmap=cmap, + vmin=vmin, + vmax=vmax, + interpolation="nearest", + origin="upper", + ) + _style_axis(middle_ax) + + bottom_ax.imshow( + item["actual"], + cmap=cmap, + vmin=vmin, + vmax=vmax, + interpolation="nearest", + origin="upper", + ) + _style_axis(bottom_ax) + return fig + + +def _select_volume( + x: Tensor, + sample_idx: int = 0, + channel_idx: int = 0, +) -> Tensor: + x = x.detach().float().cpu() + if x.ndim == 5: + return x[sample_idx, channel_idx] + if x.ndim == 4: + return x[sample_idx] + if x.ndim == 3: + return x + raise ValueError(f"expected a 3D volume tensor, got shape {tuple(x.shape)}") + + +def _select_scalar(value: float | Tensor, sample_idx: int = 0) -> float: + if isinstance(value, Tensor): + value = value.detach().float().cpu() + if value.ndim > 0: + value = value.reshape(-1)[sample_idx] + return float(value) + return float(value) + + +def _prediction_composite(target: Tensor, pred: Tensor, pred_mask: Tensor) -> Tensor: + pred_mask = pred_mask.to(dtype=target.dtype) + return target * (1 - pred_mask) + pred * pred_mask + + +def _extract_view_slice( + volume: Tensor, + view: str, + slice_idx: int | Mapping[str, int] | None = None, +) -> Tensor: + if isinstance(slice_idx, Mapping): + slice_idx = slice_idx.get(view) + + if view in ("sagittal", "saggital"): + idx = _resolve_slice_idx(volume.shape[0], slice_idx) + return volume[idx, :, :].transpose(0, 1).flip(0) + if view == "axial": + idx = _resolve_slice_idx(volume.shape[2], slice_idx) + return volume[:, :, idx].transpose(0, 1).flip(0) + if view == "coronal": + idx = _resolve_slice_idx(volume.shape[1], slice_idx) + return volume[:, idx, :].transpose(0, 1).flip(0) + raise ValueError(f"unknown MRI view {view!r}") + + +def _resolve_slice_idx(size: int, slice_idx: int | None) -> int: + idx = size // 2 if slice_idx is None else int(slice_idx) + if idx < 0: + idx += size + if idx < 0 or idx >= size: + raise IndexError(f"slice index {idx} is out of bounds for axis with size {size}") + return idx + + +def _intensity_limits(volume: Tensor, mask: Tensor | None = None) -> tuple[float, float]: + values = volume[mask] if mask is not None and mask.any() else volume.flatten() + values = values[torch.isfinite(values)] + if values.numel() == 0: + return 0.0, 1.0 + if values.numel() < 32: + vmin = values.min() + vmax = values.max() + else: + vmin, vmax = torch.quantile(values, torch.tensor([0.005, 0.995])) + if torch.isclose(vmin, vmax): + delta = max(abs(float(vmin)) * 0.05, 1.0) + return float(vmin) - delta, float(vmax) + delta + return float(vmin), float(vmax) + + +def _apply_display_mask(image: Tensor, mask: Tensor | None, fill_value: float) -> Tensor: + if mask is None: + return image + return torch.where(mask, image, torch.full_like(image, fill_value)) + + +def _masked_input_display( + image: Tensor, + pred_mask: Tensor, + img_mask: Tensor | None, + fill_value: float, +) -> Tensor: + display = torch.where(pred_mask, torch.full_like(image, fill_value), image) + return _apply_display_mask(display, img_mask, fill_value) + + +def _crop_view_items(view_items: list[dict]) -> None: + for item in view_items: + mask = item["img_mask"] + if mask is None: + mask = item["actual"] != item["actual"].min() + row_slice, col_slice = _content_crop(mask, item["patch_rc"]) + for key in ("target", "composite", "actual", "mask"): + item[key] = item[key][row_slice, col_slice] + if item["img_mask"] is not None: + item["img_mask"] = item["img_mask"][row_slice, col_slice] + + +def _content_crop(mask: Tensor, patch_size: tuple[int, int]) -> tuple[slice, slice]: + mask = mask.detach().cpu() > 0 + if not mask.any(): + return slice(None), slice(None) + + rows, cols = mask.nonzero(as_tuple=True) + patch_h, patch_w = patch_size + height, width = mask.shape + row0 = max((int(rows.min()) // patch_h - 1) * patch_h, 0) + row1 = min((int(rows.max()) // patch_h + 2) * patch_h, height) + col0 = max((int(cols.min()) // patch_w - 1) * patch_w, 0) + col1 = min((int(cols.max()) // patch_w + 2) * patch_w, width) + return slice(row0, row1), slice(col0, col1) + + +def _as_3tuple(value: int | tuple[int, int, int]) -> tuple[int, int, int]: + if isinstance(value, int): + return (value, value, value) + if len(value) != 3: + raise ValueError(f"expected a 3-tuple patch size, got {value!r}") + return tuple(int(v) for v in value) + + +def _view_patch_size(view: str, patch_size: tuple[int, int, int]) -> tuple[int, int]: + p_x, p_y, p_z = patch_size + if view in ("sagittal", "saggital"): + return p_z, p_y + if view == "axial": + return p_y, p_x + if view == "coronal": + return p_z, p_x + raise ValueError(f"unknown MRI view {view!r}") + + +def _make_figure_canvas( + view_items: list[dict[str, Tensor | str | tuple[int, int]]], + figsize: tuple[float, float] | None = None, +): + dpi = 160 + left = 58 + right = 6 + top = 18 + bottom = 8 + row_gap = 14 + col_gap = 8 + widths = [int(item["target"].shape[1]) for item in view_items] + heights = [int(item["target"].shape[0]) for item in view_items] + row_h = max(heights) + num_rows = 3 + fig_w = left + right + sum(widths) + col_gap * (len(widths) - 1) + fig_h = top + bottom + row_h * num_rows + row_gap * (num_rows - 1) + + scale = 1.35 + if figsize is not None: + requested_w = figsize[0] * dpi + requested_h = figsize[1] * dpi + scale = max(requested_w / fig_w, requested_h / fig_h) + figsize = (fig_w * scale / dpi, fig_h * scale / dpi) + fig = plt.figure(figsize=figsize, dpi=dpi, facecolor="#0b0f14") + + axes = [[] for _ in range(num_rows)] + col_centers = [] + x = left + for width, height in zip(widths, heights): + col_centers.append((x + width / 2) / fig_w) + ys = [ + bottom + (num_rows - row - 1) * (row_h + row_gap) + (row_h - height) / 2 + for row in range(num_rows) + ] + for row, y in enumerate(ys): + axes[row].append( + fig.add_axes( + [ + x / fig_w, + y / fig_h, + width / fig_w, + height / fig_h, + ], + facecolor="black", + ) + ) + x += width + col_gap + row_centers = [ + (bottom + (num_rows - row - 1) * (row_h + row_gap) + row_h / 2) / fig_h + for row in range(num_rows) + ] + layout = { + "col_centers": col_centers, + "row_centers": row_centers, + "label_x": (left - 8) / fig_w, + "title_y": (fig_h - top / 2) / fig_h, + } + + return fig, axes, layout + + +def _style_axis(ax) -> None: + ax.set_xticks([]) + ax.set_yticks([]) + for spine in ax.spines.values(): + spine.set_visible(False) + + +def _draw_patch_boxes( + ax, + mask: Tensor, + patch_size: tuple[int, int], + color: str = "#facc15", +) -> None: + for col, row, width, height in _patch_rectangles(mask, patch_size): + ax.add_patch( + patches.Rectangle( + (col - 0.5, row - 0.5), + width, + height, + fill=False, + edgecolor=color, + linewidth=0.75, + alpha=0.95, + ) + ) + + +def _patch_rectangles( + mask: Tensor, + patch_size: tuple[int, int], +) -> list[tuple[int, int, int, int]]: + mask = mask.detach().cpu() > 0 + patch_h, patch_w = patch_size + height, width = mask.shape + rectangles = [] + for row in range(0, height, patch_h): + box_h = min(patch_h, height - row) + for col in range(0, width, patch_w): + box_w = min(patch_w, width - col) + if mask[row : row + box_h, col : col + box_w].any(): + rectangles.append((col, row, box_w, box_h)) + return rectangles diff --git a/finetune/fomo_tune_baseline/output/task3/model/head.joblib b/finetune/fomo_tune_baseline/output/task3/model/head.joblib new file mode 100644 index 0000000000000000000000000000000000000000..7ade23f9dc38a75d598a1a8dfeb4930ff2d4b1a9 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task3/model/head.joblib @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e263a95e347fbaf86b4edc44717e55891711bf37a7a1309c7f04753af46a4dbf +size 29906 diff --git a/finetune/fomo_tune_baseline/output/task5/build/model/backbone.pth b/finetune/fomo_tune_baseline/output/task5/build/model/backbone.pth new file mode 100644 index 0000000000000000000000000000000000000000..793b42f5cea0b055534c29cd18c21f3200873219 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task5/build/model/backbone.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aacf2582e9464e7fd4bd3d5ed41a5b9aa770e5b04a24b2a82c1c4ba5fb5f6e7a +size 1389681124 diff --git a/finetune/fomo_tune_baseline/output/task5/build/model/head.joblib b/finetune/fomo_tune_baseline/output/task5/build/model/head.joblib new file mode 100644 index 0000000000000000000000000000000000000000..fb645916a1c0161712332085d845e2a4bd8f5014 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task5/build/model/head.joblib @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dc3c55c52bd9306b8385abe2232837d7ff0181320f93c487db2eaea27a895c1c +size 445215 diff --git a/finetune/fomo_tune_baseline/output/task5/model/head.joblib b/finetune/fomo_tune_baseline/output/task5/model/head.joblib new file mode 100644 index 0000000000000000000000000000000000000000..fb645916a1c0161712332085d845e2a4bd8f5014 --- /dev/null +++ b/finetune/fomo_tune_baseline/output/task5/model/head.joblib @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dc3c55c52bd9306b8385abe2232837d7ff0181320f93c487db2eaea27a895c1c +size 445215