diff --git a/pixi.lock b/pixi.lock index 63a4d91..fd342c2 100644 --- a/pixi.lock +++ b/pixi.lock @@ -6,8 +6,6 @@ environments: - url: https://conda.anaconda.org/bioconda/ indexes: - https://pypi.org/simple - options: - pypi-prerelease-mode: if-necessary-or-explicit packages: linux-64: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda @@ -603,6 +601,8 @@ environments: - conda: https://conda.anaconda.org/conda-forge/linux-64/zstandard-0.25.0-py312h5253ce2_1.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/zstd-1.5.7-hb78ec9c_6.conda - pypi: https://files.pythonhosted.org/packages/6f/4b/817da308fa1170da07ef01259585887a3bbb6ab80700b3e61ce4967301ec/dask_image-2025.11.0-py3-none-any.whl + - pypi: https://files.pythonhosted.org/packages/2f/e3/07a70c9b83e5446528dd5d222fdca88ab30efcb07094536a864546c0096f/dynamic_network_architectures-0.4.4-py3-none-any.whl + - pypi: https://files.pythonhosted.org/packages/2a/09/f8d8f8f31e4483c10a906437b4ce31bdf3d6d417b73fe33f1a8b59e34228/einops-0.8.2-py3-none-any.whl - pypi: https://files.pythonhosted.org/packages/9e/e9/1a19e42cd43cc1365e127db6aae85e1c671da1d9a5d746f4d34a50edb577/h5py-3.16.0-cp312-cp312-manylinux_2_28_x86_64.whl - pypi: https://files.pythonhosted.org/packages/ed/2f/1046d464ad1db29a4f6c70ba4e19b39baa8a6542c719eaa4e765108f07f1/hdf5plugin-6.0.0-py3-none-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl - pypi: https://files.pythonhosted.org/packages/c0/47/7dbdec41b18f4c9aa90240c8804482a8c4ada6db3d2844e287695cd42109/imaris_ims_zarr-0.2.0-py3-none-any.whl @@ -623,9 +623,12 @@ environments: - pypi: https://files.pythonhosted.org/packages/f6/17/7468d6cb45b0c537b235bb910812bb248c581a50ffc5eeaeba342f6fdf2d/pylibtiff-0.7.0.tar.gz - pypi: https://files.pythonhosted.org/packages/2f/43/d7e2b9ad768c07b5473bea3ac7db9ca4d995c09399cbea3d4df1c0bd4955/rangehttpserver-1.4.0-py2.py3-none-any.whl - pypi: https://files.pythonhosted.org/packages/0b/35/1cceccc5fcb50fa2ed53e2aa278cd032f3902682a73e763fb1ac3be8e6fa/rich_argparse-1.8.0-py3-none-any.whl + - pypi: https://files.pythonhosted.org/packages/28/50/f203ff3a3ddfe19308efc83c5a3a29ed02bf786732ec35e68bf9162f3365/safetensors-0.8.0-cp310-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl - pypi: https://files.pythonhosted.org/packages/c3/5d/5e017f59bbc06af5b9c78331036478b2ddb371d7ff810c559c36e5389904/simpleitk-2.5.4-cp311-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl - pypi: https://files.pythonhosted.org/packages/e8/ae/fa6cd331b364ad2bbc31652d025f5747d89cbb75576733dfdf8efe3e4d62/slicerator-1.1.0-py3-none-any.whl - pypi: https://files.pythonhosted.org/packages/d1/14/7688e984cdd0be1438779825640b943574b89946ed868d76497b3cffb3d5/tensorstore-0.1.83-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl + - pypi: https://files.pythonhosted.org/packages/d6/14/fc04d491527b774ec7479897f5861959209de1480e4c4cd32ed098ff8bea/timm-1.0.22-py3-none-any.whl + - pypi: https://files.pythonhosted.org/packages/17/8b/155f99042f9319bd7759536779b2a5b67cbd4f89c380854670850f89a2f4/torchvision-0.22.1-cp312-cp312-manylinux_2_28_x86_64.whl - pypi: git+https://github.com/khanlab/vesselfm?rev=504baab#504baabb253c2b6bf7d5a4037a93c2f590c9da51 - pypi: https://files.pythonhosted.org/packages/c5/c1/9c64ed6c2d152b56482bb66297bfc03442d25449bfa91da7bf2e58a80f01/wasmtime-44.0.0-py3-none-manylinux1_x86_64.whl - pypi: https://files.pythonhosted.org/packages/bd/4d/5e4df18906ddeb6e64af6712e72056c73ee926d402c0a5fd16feb72f6fe9/zarrnii-0.21.2-py3-none-any.whl @@ -636,8 +639,6 @@ environments: - url: https://conda.anaconda.org/bioconda/ indexes: - https://pypi.org/simple - options: - pypi-prerelease-mode: if-necessary-or-explicit packages: linux-64: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda @@ -1410,6 +1411,8 @@ environments: - conda: https://conda.anaconda.org/conda-forge/linux-64/zstandard-0.25.0-py312h5253ce2_1.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/zstd-1.5.7-hb78ec9c_6.conda - pypi: https://files.pythonhosted.org/packages/6f/4b/817da308fa1170da07ef01259585887a3bbb6ab80700b3e61ce4967301ec/dask_image-2025.11.0-py3-none-any.whl + - pypi: https://files.pythonhosted.org/packages/d8/ae/b8a29e6dd366327530390464bf7eee0e3b3bcb03deea98b1a6951edb68ba/dynamic_network_architectures-0.4.2.tar.gz + - pypi: https://files.pythonhosted.org/packages/2a/09/f8d8f8f31e4483c10a906437b4ce31bdf3d6d417b73fe33f1a8b59e34228/einops-0.8.2-py3-none-any.whl - pypi: https://files.pythonhosted.org/packages/ed/2f/1046d464ad1db29a4f6c70ba4e19b39baa8a6542c719eaa4e765108f07f1/hdf5plugin-6.0.0-py3-none-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl - pypi: https://files.pythonhosted.org/packages/c0/47/7dbdec41b18f4c9aa90240c8804482a8c4ada6db3d2844e287695cd42109/imaris_ims_zarr-0.2.0-py3-none-any.whl - pypi: https://files.pythonhosted.org/packages/aa/57/0a3499479cb19d7e4b7fc38b2ba15c0ea20e1e88cb2b634b75dd3e9ff8b8/itk_core-5.4.6-cp311-abi3-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl @@ -1440,8 +1443,6 @@ environments: - url: https://conda.anaconda.org/bioconda/ indexes: - https://pypi.org/simple - options: - pypi-prerelease-mode: if-necessary-or-explicit packages: linux-64: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-20_gnu.conda @@ -1740,8 +1741,6 @@ environments: - url: https://conda.anaconda.org/bioconda/ indexes: - https://pypi.org/simple - options: - pypi-prerelease-mode: if-necessary-or-explicit packages: linux-64: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda @@ -2365,6 +2364,8 @@ environments: - conda: https://conda.anaconda.org/conda-forge/linux-64/zstandard-0.25.0-py312h5253ce2_1.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/zstd-1.5.7-hb78ec9c_6.conda - pypi: https://files.pythonhosted.org/packages/6f/4b/817da308fa1170da07ef01259585887a3bbb6ab80700b3e61ce4967301ec/dask_image-2025.11.0-py3-none-any.whl + - pypi: https://files.pythonhosted.org/packages/2f/e3/07a70c9b83e5446528dd5d222fdca88ab30efcb07094536a864546c0096f/dynamic_network_architectures-0.4.4-py3-none-any.whl + - pypi: https://files.pythonhosted.org/packages/2a/09/f8d8f8f31e4483c10a906437b4ce31bdf3d6d417b73fe33f1a8b59e34228/einops-0.8.2-py3-none-any.whl - pypi: https://files.pythonhosted.org/packages/9e/e9/1a19e42cd43cc1365e127db6aae85e1c671da1d9a5d746f4d34a50edb577/h5py-3.16.0-cp312-cp312-manylinux_2_28_x86_64.whl - pypi: https://files.pythonhosted.org/packages/ed/2f/1046d464ad1db29a4f6c70ba4e19b39baa8a6542c719eaa4e765108f07f1/hdf5plugin-6.0.0-py3-none-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl - pypi: https://files.pythonhosted.org/packages/c0/47/7dbdec41b18f4c9aa90240c8804482a8c4ada6db3d2844e287695cd42109/imaris_ims_zarr-0.2.0-py3-none-any.whl @@ -2385,9 +2386,12 @@ environments: - pypi: https://files.pythonhosted.org/packages/f6/17/7468d6cb45b0c537b235bb910812bb248c581a50ffc5eeaeba342f6fdf2d/pylibtiff-0.7.0.tar.gz - pypi: https://files.pythonhosted.org/packages/2f/43/d7e2b9ad768c07b5473bea3ac7db9ca4d995c09399cbea3d4df1c0bd4955/rangehttpserver-1.4.0-py2.py3-none-any.whl - pypi: https://files.pythonhosted.org/packages/0b/35/1cceccc5fcb50fa2ed53e2aa278cd032f3902682a73e763fb1ac3be8e6fa/rich_argparse-1.8.0-py3-none-any.whl + - pypi: https://files.pythonhosted.org/packages/28/50/f203ff3a3ddfe19308efc83c5a3a29ed02bf786732ec35e68bf9162f3365/safetensors-0.8.0-cp310-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl - pypi: https://files.pythonhosted.org/packages/c3/5d/5e017f59bbc06af5b9c78331036478b2ddb371d7ff810c559c36e5389904/simpleitk-2.5.4-cp311-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl - pypi: https://files.pythonhosted.org/packages/e8/ae/fa6cd331b364ad2bbc31652d025f5747d89cbb75576733dfdf8efe3e4d62/slicerator-1.1.0-py3-none-any.whl - pypi: https://files.pythonhosted.org/packages/d1/14/7688e984cdd0be1438779825640b943574b89946ed868d76497b3cffb3d5/tensorstore-0.1.83-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl + - pypi: https://files.pythonhosted.org/packages/d6/14/fc04d491527b774ec7479897f5861959209de1480e4c4cd32ed098ff8bea/timm-1.0.22-py3-none-any.whl + - pypi: https://files.pythonhosted.org/packages/17/8b/155f99042f9319bd7759536779b2a5b67cbd4f89c380854670850f89a2f4/torchvision-0.22.1-cp312-cp312-manylinux_2_28_x86_64.whl - pypi: git+https://github.com/khanlab/vesselfm?rev=504baab#504baabb253c2b6bf7d5a4037a93c2f590c9da51 - pypi: https://files.pythonhosted.org/packages/c5/c1/9c64ed6c2d152b56482bb66297bfc03442d25449bfa91da7bf2e58a80f01/wasmtime-44.0.0-py3-none-manylinux1_x86_64.whl - pypi: https://files.pythonhosted.org/packages/bd/4d/5e4df18906ddeb6e64af6712e72056c73ee926d402c0a5fd16feb72f6fe9/zarrnii-0.21.2-py3-none-any.whl @@ -4745,6 +4749,25 @@ packages: - pkg:pypi/dunamai?source=hash-mapping size: 30866 timestamp: 1775597880526 +- pypi: https://files.pythonhosted.org/packages/d8/ae/b8a29e6dd366327530390464bf7eee0e3b3bcb03deea98b1a6951edb68ba/dynamic_network_architectures-0.4.2.tar.gz + name: dynamic-network-architectures + version: 0.4.2 + sha256: f0f4cfa26e99e7c1cc44d3ac48ce802f109c403d109c4e4d813f682fa9d0ab84 + requires_dist: + - torch>=1.6.0a0 + - numpy + - timm + - einops +- pypi: https://files.pythonhosted.org/packages/2f/e3/07a70c9b83e5446528dd5d222fdca88ab30efcb07094536a864546c0096f/dynamic_network_architectures-0.4.4-py3-none-any.whl + name: dynamic-network-architectures + version: 0.4.4 + sha256: 7c51bc6518d5c3ee5449b81e00b4e140bb5fd26d035744c01538cdb47b55751d + requires_dist: + - torch>=1.6.0a0 + - numpy + - timm<1.0.23 + - einops + requires_python: '>=3.8' - conda: https://conda.anaconda.org/conda-forge/noarch/eido-0.2.5-pyhd8ed1ab_0.conda sha256: 6972edd5d9dffccd0afcbec077ef2f3a51abc6cc70060bcc818d02988b70983d md5: 679c2b63523ef75652e5690a39e0f0a2 @@ -4768,6 +4791,11 @@ packages: purls: [] size: 1311364 timestamp: 1773744756289 +- pypi: https://files.pythonhosted.org/packages/2a/09/f8d8f8f31e4483c10a906437b4ce31bdf3d6d417b73fe33f1a8b59e34228/einops-0.8.2-py3-none-any.whl + name: einops + version: 0.8.2 + sha256: 54058201ac7087911181bfec4af6091bb59380360f069276601256a76af08193 + requires_python: '>=3.9' - conda: https://conda.anaconda.org/conda-forge/noarch/email-validator-2.3.0-pyhd8ed1ab_0.conda sha256: c37320864c35ef996b0e02e289df6ee89582d6c8e233e18dc9983375803c46bb md5: 3bc0ac31178387e8ed34094d9481bfe8 @@ -13062,6 +13090,48 @@ packages: - pkg:pypi/s3transfer?source=hash-mapping size: 65987 timestamp: 1757487748738 +- pypi: https://files.pythonhosted.org/packages/28/50/f203ff3a3ddfe19308efc83c5a3a29ed02bf786732ec35e68bf9162f3365/safetensors-0.8.0-cp310-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl + name: safetensors + version: 0.8.0 + sha256: fd6f3f93c9a0a7cc2788ee63fb763353d4bd2e89b0751bc78fcf7dda00bea774 + requires_dist: + - safetensors[torch] ; extra == 'all' + - safetensors[numpy] ; extra == 'all' + - safetensors[jax] ; extra == 'all' + - safetensors[paddlepaddle] ; extra == 'all' + - safetensors[convert] ; extra == 'all' + - safetensors[quality] ; extra == 'all' + - safetensors[testing] ; extra == 'all' + - safetensors[torch] ; extra == 'convert' + - huggingface-hub>=1.4 ; extra == 'convert' + - safetensors[all] ; extra == 'dev' + - safetensors[pinned-tf] ; extra == 'dev' + - safetensors[numpy] ; extra == 'jax' + - flax>=0.6.3 ; extra == 'jax' + - jax>=0.3.25 ; extra == 'jax' + - jaxlib>=0.3.25 ; extra == 'jax' + - mlx>=0.0.9 ; extra == 'mlx' + - numpy>=1.24.6 ; extra == 'numpy' + - safetensors[numpy] ; extra == 'paddlepaddle' + - paddlepaddle>=2.4.1 ; extra == 'paddlepaddle' + - safetensors[numpy] ; extra == 'pinned-tf' + - tensorflow==2.18.0 ; extra == 'pinned-tf' + - ruff ; extra == 'quality' + - safetensors[numpy] ; extra == 'tensorflow' + - tensorflow>=2.11.0 ; extra == 'tensorflow' + - safetensors[numpy] ; extra == 'testing' + - h5py>=3.7.0 ; extra == 'testing' + - setuptools-rust>=1.12.0 ; extra == 'testing' + - pytest>=9.0 ; extra == 'testing' + - pytest-benchmark>=5.2 ; extra == 'testing' + - hypothesis>=6.70.2 ; extra == 'testing' + - fsspec>=2024.6.0 ; extra == 'testing' + - s3fs>=2024.6.0 ; extra == 'testing' + - safetensors[numpy] ; extra == 'tf-nightly' + - tf-nightly ; extra == 'tf-nightly' + - safetensors[numpy] ; extra == 'torch' + - torch>=2.4 ; extra == 'torch' + requires_python: '>=3.10' - conda: https://conda.anaconda.org/conda-forge/linux-64/safetensors-0.7.0-py312h0ccc70a_0.conda sha256: 59886c90dbfe388eee79c78256125ef2346a6d18058cf9769758c35f2e836a8c md5: c11270acf79465903f1508314c3e86fd @@ -13918,8 +13988,9 @@ packages: - pypi: ./ name: spimquant version: 0.1.0 - sha256: 822584aa22bfc1737bd444f82c39ef79c6df0f880a359ab332a1a98155899064 + sha256: f51bae3a4dc97ad20fa034c702e598b9becd0971f6d09d50f79df5f8b72c8d1d requires_python: '>=3.11' + editable: true - conda: https://conda.anaconda.org/conda-forge/linux-64/spirv-tools-2025.5-hb700be7_0.conda sha256: 7547142ab1352132adf98d555ed955badd96c9f277cbd054ae52f7edd6cf6cb8 md5: 058d5f16eaa3018be91aa3508df00d7c @@ -14268,6 +14339,17 @@ packages: - pkg:pypi/tifffile?source=hash-mapping size: 209916 timestamp: 1777805970265 +- pypi: https://files.pythonhosted.org/packages/d6/14/fc04d491527b774ec7479897f5861959209de1480e4c4cd32ed098ff8bea/timm-1.0.22-py3-none-any.whl + name: timm + version: 1.0.22 + sha256: 888981753e65cbaacfc07494370138b1700a27b1f0af587f4f9b47bc024161d0 + requires_dist: + - torch + - torchvision + - pyyaml + - huggingface-hub + - safetensors + requires_python: '>=3.8' - conda: https://conda.anaconda.org/conda-forge/noarch/timm-1.0.27-pyhcf101f3_0.conda sha256: 77c69e2ea6d040200fd21e9b82889e9eb1fd4cade6b4eaa877b442af636e0654 md5: 59a0c3968e59e6ce52855ef9bc9ff421 @@ -14397,6 +14479,17 @@ packages: - pkg:pypi/torchmetrics?source=hash-mapping size: 398600 timestamp: 1773827733172 +- pypi: https://files.pythonhosted.org/packages/17/8b/155f99042f9319bd7759536779b2a5b67cbd4f89c380854670850f89a2f4/torchvision-0.22.1-cp312-cp312-manylinux_2_28_x86_64.whl + name: torchvision + version: 0.22.1 + sha256: 699c2d70d33951187f6ed910ea05720b9b4aaac1dcc1135f53162ce7d42481d3 + requires_dist: + - numpy + - torch==2.7.1 + - pillow>=5.3.0,!=8.3.* + - gdown>=4.7.3 ; extra == 'gdown' + - scipy ; extra == 'scipy' + requires_python: '>=3.9' - conda: https://conda.anaconda.org/conda-forge/linux-64/torchvision-0.24.0-cpu_py312_h56cfa8b_0.conda sha256: a22214010fd1f3398ad7ea8311e0d76b0f7750590b502b0a73f57c14ca806595 md5: feb0c7bde000f25617408266492d2d2a diff --git a/pyproject.toml b/pyproject.toml index f60ae60..31b28ff 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -79,6 +79,13 @@ features = ["runtime"] [tool.pixi.environments.gpu] features = ["runtime", "gpu"] +# Scoped to `runtime` rather than the workspace-global pypi-dependencies: it is a +# torch package, and putting it at workspace level pulls a PyPI torch + CUDA stack +# into the lint-only `dev-only` environment, which has no other torch dependency. +[tool.pixi.feature.runtime.pypi-dependencies] +dynamic-network-architectures = ">=0.4" + + [tool.pixi.feature.gpu.system-requirements] cuda = "12.0" diff --git a/spimquant/config/snakebids.yml b/spimquant/config/snakebids.yml index f518ff2..b4025b6 100644 --- a/spimquant/config/snakebids.yml +++ b/spimquant/config/snakebids.yml @@ -194,6 +194,18 @@ parse_args: ``threshold = expm1(mean_n + k * sigma_n)``. The float multiplier k may be encoded with 'p' as a decimal separator to keep BIDS-compatible filenames, e.g. ``gmm+n2k2p5`` → n=2, k=2.5. + + ``lantern`` – 5-fold LANTERN deep-learning ensemble for amyloid-beta + plaques. Unlike the histogram-based methods above it reads raw + intensities (not the bias-field-corrected image) and runs at + ``plaque_level``; the mask at ``segmentation_level`` is that prediction + upsampled. Requires a GPU. + + Also unlike them, it is stain-specific: it runs only on the first stain + from ``stains_for_plaques``. The remaining stains in ``stains_for_seg`` + use the other methods given here, or -- if ``lantern`` was the only one -- + this option's own default, so they behave exactly as if it had not been + passed. default: - gmm+n3k1 nargs: '+' @@ -447,6 +459,12 @@ templates: models: vesselfm: https://huggingface.co/bwittmann/vesselFM/resolve/main/vesselFM_base.pt + lantern_abeta: + fold0: https://huggingface.co/apooladi/lantern-ki3-abeta/resolve/main/fold0/seg_model.pt + fold1: https://huggingface.co/apooladi/lantern-ki3-abeta/resolve/main/fold1/seg_model.pt + fold2: https://huggingface.co/apooladi/lantern-ki3-abeta/resolve/main/fold2/seg_model.pt + fold3: https://huggingface.co/apooladi/lantern-ki3-abeta/resolve/main/fold3/seg_model.pt + fold4: https://huggingface.co/apooladi/lantern-ki3-abeta/resolve/main/fold4/seg_model.pt @@ -524,6 +542,31 @@ vessel_seg_method: vesselfm vessel_metrics: - fieldfrac +# LANTERN deep-learning plaque segmentation (opt in with --seg_method lantern). +# The ensemble is scale-sensitive and was built for plaque_level; the mask at +# segmentation_level is that prediction upsampled, not native inference. +plaque_seg_method: lantern +# The LANTERN ensemble is trained on amyloid-beta only, so unlike the +# intensity-based methods it is not applied to every stain in stains_for_seg. +# The first of these found is the one it runs on; the remaining stains are +# segmented exactly as they would be if --seg_method had never been passed. +stains_for_plaques: + - abeta + - Abeta + - BetaAmyloid +plaque_level: 1 +plaque_n_folds: 5 +plaque_vote_threshold: 3 +plaque_tile: 128 +plaque_stride: 64 +plaque_batch_size: 8 +plaque_chunk: 512 +# GPUs per inference job. The ensemble is replicated onto each card and dask +# blocks are handed to whichever card is free, so this scales one subject's +# wall-clock rather than throughput across subjects. Clamped at runtime to the +# number of visible devices. +plaque_n_gpus: 2 + coloc_seg_metrics: - overlapratio diff --git a/spimquant/workflow/Snakefile b/spimquant/workflow/Snakefile index 08230ec..eee1ebb 100644 --- a/spimquant/workflow/Snakefile +++ b/spimquant/workflow/Snakefile @@ -175,6 +175,113 @@ else: do_vessels = False +# The LANTERN plaque ensemble is opt-in via --seg_method, but it is trained on +# amyloid-beta only, so it is resolved to a single stain the same way vessels are +# rather than being fanned out over every stain in stains_for_seg. +do_plaques = do_seg and config["plaque_seg_method"] in config["seg_method"] + +stain_for_plaques = None + +if do_plaques: + for stain in config["stains_for_plaques"]: + if stain in stains_for_seg: + stain_for_plaques = stain + break + + if stain_for_plaques is None: + raise ValueError( + f"--seg_method {config['plaque_seg_method']} was requested but none of " + f"{config['stains_for_plaques']} is in stains_for_seg ({stains_for_seg}). " + f"The LANTERN ensemble only segments amyloid-beta." + ) + +# Segmentation is a set of (method, stain) pairs, not a method x stain cross +# product: the generic intensity-based methods run on every stain, while the plaque +# method runs only on stain_for_plaques. Everything below exists to keep those two +# groups from contaminating each other. +# +# Generic methods, i.e. everything the user asked for except the plaque method. +generic_seg_methods = [ + m for m in config["seg_method"] if m != config["plaque_seg_method"] +] + +if generic_seg_methods: + # The user asked for generic methods explicitly, so they apply everywhere -- + # including the plaque stain, which is how `--seg_method gmm+n3k1 lantern` + # produces a like-for-like comparison on the same channel. + seg_methods = generic_seg_methods + stains_for_seg_methods = stains_for_seg +else: + # The plaque method was the ONLY one requested. The remaining stains must behave + # exactly as if --seg_method had never been passed -- but the plaque stain must + # NOT, or it would be segmented twice, once by LANTERN and once by the default + # GMM. The default is read back off the CLI definition rather than restated here + # so it cannot drift out of sync if the actual default ever changes. + seg_methods = list(config["parse_args"]["--seg_method"]["default"]) + stains_for_seg_methods = [s for s in stains_for_seg if s != stain_for_plaques] + + +def stains_for_desc(desc): + """Stains that segmentation method `desc` produces a mask for.""" + if do_plaques and desc == config["plaque_seg_method"]: + return [stain_for_plaques] + return stains_for_seg_methods + + +# Every method that produces a mask, plaque method included. +# +# Methods left with no stains are dropped. This is not hypothetical: requesting +# only the plaque method on a single-stain dataset removes that stain from +# stains_for_seg_methods and leaves the fallback generic method covering nothing. +# Keeping it here would schedule cross-stain aggregations (mergedsegstats, +# aggregated regionprops) with an empty input list, which fail at runtime rather +# than at parse time. +all_seg_methods = [ + m + for m in seg_methods + ([config["plaque_seg_method"]] if do_plaques else []) + if stains_for_desc(m) +] + + +# Methods that segment two or more stains, and so can be colocalized across them. +# The plaque method is single-stain by construction, so it drops out here. +coloc_seg_methods = [m for m in all_seg_methods if len(stains_for_desc(m)) > 1] + +# Recomputed now that the plaque stain is known: colocalization is only meaningful +# if some method segments two or more stains. Requesting the plaque method alone on +# a single-stain dataset leaves nothing to colocalize. +do_coloc = do_seg and len(coloc_seg_methods) > 0 + + +def seg_expand(paths, **wildcards): + """Expand `paths` over every (desc, stain) segmentation pair. + + Replaces `inputs["spim"].expand(..., desc=seg_methods, stain=stains_for_seg)`. + Expanding each method against its own stains and concatenating gives the plaque + method the same per-subject downstream products as every other method + (featuremaps, fieldfrac, counts, regionprops, segstats, QC) without the generic + methods leaking onto the plaque stain. + + With no plaque method requested this returns exactly what the cross product did. + """ + out = [] + for desc in all_seg_methods: + stains = stains_for_desc(desc) + if stains: + out += inputs["spim"].expand(paths, desc=[desc], stain=stains, **wildcards) + return out + + +def seg_expand_plain(paths, **wildcards): + """`seg_expand` for group-level paths, which carry no subject wildcards.""" + out = [] + for desc in all_seg_methods: + stains = stains_for_desc(desc) + if stains: + out += expand(paths, desc=[desc], stain=stains, **wildcards) + return out + + # atlas segmentations to use if config["atlas_segs"] is None: @@ -540,9 +647,54 @@ rule all_vessels: ), -rule all_segment: +rule all_plaques: input: inputs["spim"].expand( + bids_oz_in( + root=root, + datatype="seg", + stain="{stain}", + level="{level}", + desc=config["plaque_seg_method"], + suffix="probseg.{ext}", + **inputs["spim"].wildcards, + ), + level=config["plaque_level"], + stain=stain_for_plaques, + ), + inputs["spim"].expand( + bids_oz_in( + root=root, + datatype="seg", + stain="{stain}", + level="{level}", + desc=config["plaque_seg_method"], + suffix="mask.{ext}", + **inputs["spim"].wildcards, + ), + level=config["segmentation_level"], + stain=stain_for_plaques, + ), + inputs["spim"].expand( + bids( + root=root, + datatype="seg", + stain="{stain}", + level="{level}", + desc=config["plaque_seg_method"], + space="{template}", + suffix="fieldfrac.nii.gz", + **inputs["spim"].wildcards, + ), + level=config["registration_level"], + template=config["template"], + stain=stain_for_plaques, + ), + + +rule all_segment: + input: + seg_expand( bids( root=root, datatype="featuremap", @@ -553,10 +705,8 @@ rule all_segment: **inputs["spim"].wildcards, ), seg=atlas_segs, - desc=config["seg_method"], template=config["template"], suffix=config["seg_metrics"], - stain=stains_for_seg, ), inputs["spim"].expand( bids( @@ -572,7 +722,7 @@ rule all_segment: level=config["registration_level"], template=config["template"], ), - inputs["spim"].expand( + seg_expand( bids( root=root, datatype="seg", @@ -583,12 +733,10 @@ rule all_segment: suffix="fieldfrac.nii.gz", **inputs["spim"].wildcards, ), - stain=stains_for_seg, level=config["registration_level"], - desc=config["seg_method"], template=config["template"], ), - inputs["spim"].expand( + seg_expand( bids( root=root, datatype="seg", @@ -598,11 +746,9 @@ rule all_segment: suffix="counts.nii.gz", **inputs["spim"].wildcards, ), - stain=stains_for_seg, level=config["registration_level"], - desc=config["seg_method"], ), - inputs["spim"].expand( + seg_expand( bids( root=root, datatype="seg", @@ -612,9 +758,7 @@ rule all_segment: suffix="counts.nii.gz", **inputs["spim"].wildcards, ), - stain=stains_for_seg, template=config["template"], - desc=config["seg_method"], ), expand( bids( @@ -642,7 +786,7 @@ rule all_segment_coloc: **inputs["spim"].wildcards, ), seg=atlas_segs, - desc=config["seg_method"], + desc=coloc_seg_methods, template=config["template"], suffix=config["coloc_seg_metrics"], ), @@ -655,7 +799,7 @@ rule all_segment_coloc: suffix="coloccounts.nii.gz", **inputs["spim"].wildcards, ), - desc=config["seg_method"], + desc=coloc_seg_methods, template=config["template"], ), @@ -706,7 +850,7 @@ rule all_spim_patches: level=config["segmentation_level"], desc=["raw", "corrected" + config["correction_method"]], ), - inputs["spim"].expand( + seg_expand( bids( root=root, datatype="seg", @@ -718,10 +862,8 @@ rule all_spim_patches: suffix="mask.patches", **inputs["spim"].wildcards, ), - stain=stains_for_seg, seg=patch_atlas_segs, template=config["template"], - desc=config["seg_method"], level=config["segmentation_level"], ), @@ -759,13 +901,13 @@ rule all_group_stats: suffix="groupstats.{ext}", ), seg=atlas_segs, - desc=config["seg_method"], + desc=all_seg_methods, template=config["template"], pairwise_contrast=pairwise_contrast_labels, ext=["png", "tsv"], ), # Per-contrast NIfTI stat maps for each seg metric - expand( + seg_expand_plain( bids( root=group_stats_root, seg="{seg}", @@ -776,15 +918,13 @@ rule all_group_stats: suffix="{stat}.nii.gz", ), seg=atlas_segs, - desc=config["seg_method"], template=config["template"], - stain=stains_for_seg, metric=config["seg_metrics"], stat=config["stats_maps"], pairwise_contrast=pairwise_contrast_labels, ), # All-subjects count maps (not contrast-specific) - expand( + seg_expand_plain( bids( root=group_stats_root, level="{level}", @@ -792,10 +932,8 @@ rule all_group_stats: desc="{desc}", suffix="{stain}+count.nii.gz", ), - desc=config["seg_method"], template=config["template"], level=range(4), - stain=stains_for_seg, ), @@ -813,7 +951,7 @@ rule all_group_stats_coloc: suffix="{stat}.nii.gz", ), seg=atlas_segs, - desc=config["seg_method"], + desc=coloc_seg_methods, template=config["template"], metric=config["coloc_seg_metrics"], stat=config["stats_maps"], @@ -828,7 +966,7 @@ rule all_group_stats_coloc: desc="{desc}", suffix="coloccount.nii.gz", ), - desc=config["seg_method"], + desc=coloc_seg_methods, template=config["template"], level=range(4), ), @@ -849,7 +987,7 @@ rule all_qc: stain=stains, ), # Segmentation overview figures (per stain, per seg method) - inputs["spim"].expand( + seg_expand( bids( root=root, datatype="qc", @@ -858,13 +996,11 @@ rule all_qc: suffix="segslices.png", **inputs["spim"].wildcards, ), - stain=stains_for_seg, - desc=config["seg_method"], ) if do_seg else [], # Segmentation ROI zoom montage (per stain, per atlas, per seg method) - inputs["spim"].expand( + seg_expand( bids( root=root, datatype="qc", @@ -877,13 +1013,11 @@ rule all_qc: ), seg=atlas_segs, template=config["template"], - stain=stains_for_seg, - desc=config["seg_method"], ) if do_seg else [], # Segmentation ROI zoom montage - no mask overlay - inputs["spim"].expand( + seg_expand( bids( root=root, datatype="qc", @@ -896,8 +1030,6 @@ rule all_qc: ), seg=atlas_segs, template=config["template"], - stain=stains_for_seg, - desc=config["seg_method"], ) if do_seg else [], @@ -955,7 +1087,7 @@ rule all_qc: if do_vessels else [], # Z-profile QC (per stain, per seg method) - inputs["spim"].expand( + seg_expand( bids( root=root, datatype="qc", @@ -964,13 +1096,11 @@ rule all_qc: suffix="zprofile.png", **inputs["spim"].wildcards, ), - stain=stains_for_seg, - desc=config["seg_method"], ) if do_seg else [], # Object-level statistics (per stain, per seg method) - inputs["spim"].expand( + seg_expand( bids( root=root, datatype="qc", @@ -979,8 +1109,6 @@ rule all_qc: suffix="objectstats.png", **inputs["spim"].wildcards, ), - stain=stains_for_seg, - desc=config["seg_method"], ) if do_seg else [], @@ -997,12 +1125,12 @@ rule all_qc: ), seg=atlas_segs, template=config["template"], - desc=config["seg_method"], + desc=all_seg_methods, ) if do_seg else [], # Instance-centred animated GIFs (per stain, per atlas seg, per seg method) - inputs["spim"].expand( + seg_expand( [ bids( root=root, @@ -1027,8 +1155,6 @@ rule all_qc: ], seg=atlas_segs, template=config["template"], - stain=stains_for_seg, - desc=config["seg_method"], ) if do_seg and config["qc_instance_seg"] else [], @@ -1056,7 +1182,7 @@ rule all_qc: ], seg=atlas_segs, template=config["template"], - desc=config["seg_method"], + desc=coloc_seg_methods, ) if do_coloc and config["qc_instance_seg"] else [], @@ -1077,12 +1203,12 @@ rule all_tabular_sidecars: **inputs["spim"].wildcards, ), seg=atlas_segs, - desc=config["seg_method"], + desc=all_seg_methods, template=config["template"], ) if do_seg else [], - inputs["spim"].expand( + seg_expand( bids( root=root, datatype="tabular", @@ -1095,9 +1221,7 @@ rule all_tabular_sidecars: **inputs["spim"].wildcards, ), seg=atlas_segs, - desc=config["seg_method"], template=config["template"], - stain=stains_for_seg, ) if do_seg else [], @@ -1111,7 +1235,7 @@ rule all_tabular_sidecars: suffix="regionprops.json", **inputs["spim"].wildcards, ), - desc=config["seg_method"], + desc=all_seg_methods, template=config["template"], ) if do_seg @@ -1125,7 +1249,7 @@ rule all_tabular_sidecars: suffix="coloc.json", **inputs["spim"].wildcards, ), - desc=config["seg_method"], + desc=coloc_seg_methods, template=config["template"], ) if do_coloc @@ -1146,7 +1270,7 @@ rule all_group_stats_tabular_sidecars: suffix="groupstats.json", ), seg=atlas_segs, - desc=config["seg_method"], + desc=all_seg_methods, template=config["template"], pairwise_contrast=pairwise_contrast_labels, ) @@ -1162,7 +1286,7 @@ rule all_group_stats_tabular_sidecars: suffix="allsubjects.json", ), seg=atlas_segs, - desc=config["seg_method"], + desc=all_seg_methods, template=config["template"], ) if do_seg @@ -1175,7 +1299,7 @@ rule all_group_stats_tabular_sidecars: desc="{desc}", suffix="regionprops.json", ), - desc=config["seg_method"], + desc=all_seg_methods, template=config["template"], ) if do_seg @@ -1187,7 +1311,7 @@ rule all_group_stats_tabular_sidecars: desc="{desc}", suffix="coloc.json", ), - desc=config["seg_method"], + desc=coloc_seg_methods, template=config["template"], ) if do_coloc @@ -1207,7 +1331,7 @@ rule all_group_merged_tsvs: suffix="allsubjects.tsv", ), seg=atlas_segs, - desc=config["seg_method"], + desc=all_seg_methods, template=config["template"], ) if do_seg @@ -1219,6 +1343,7 @@ rule all_participant: input: rules.all_register.input, rules.all_vessels.input if do_vessels else [], + rules.all_plaques.input if do_plaques else [], rules.all_segment.input if do_seg else [], rules.all_mri_reg.input if config["register_to_mri"] else [], rules.all_segment_coloc.input if do_coloc else [], @@ -1240,6 +1365,7 @@ include: "rules/import.smk" include: "rules/masking.smk" include: "rules/templatereg.smk" include: "rules/vessels.smk" +include: "rules/plaques.smk" include: "rules/segmentation.smk" include: "rules/counts.smk" include: "rules/fieldfrac.smk" diff --git a/spimquant/workflow/rules/plaques.smk b/spimquant/workflow/rules/plaques.smk new file mode 100644 index 0000000..f4ac85d --- /dev/null +++ b/spimquant/workflow/rules/plaques.smk @@ -0,0 +1,135 @@ +rule import_lantern_abeta_fold: + """Download one fold of the LANTERN Abeta plaque ensemble.""" + input: + model=lambda wildcards: storage( + config["models"]["lantern_abeta"][f"fold{wildcards.fold}"] + ), + output: + "resources/models/lantern-ki3-abeta/fold{fold}/seg_model.pt", + wildcard_constraints: + fold="[0-9]+", + localrule: True + shell: + "cp {input} {output}" + + +rule run_lantern_plaques: + """Segment Abeta plaques with the 5-fold LANTERN ensemble. + + Runs at config['plaque_level'] on RAW intensities -- the model was not + trained on N4-corrected data, so this bypasses the bias-field chain that + the GMM and Otsu methods depend on. + + Tiles at 128^3 with stride 64 and combines overlapping tiles by max over + votes. The low-res brain mask restricts inference to blocks that touch + tissue, which is what makes whole-brain 5-fold inference tractable. + + Writes votes/n_folds rather than a binary mask, so the vote threshold can be + changed downstream without re-running inference. + """ + input: + spim=inputs["spim"].path, + models=expand( + "resources/models/lantern-ki3-abeta/fold{fold}/seg_model.pt", + fold=range(config["plaque_n_folds"]), + ), + mask=bids( + root=root, + datatype="micr", + stain=stain_for_reg, + level=config["correction_level"], + desc="brain", + suffix="mask.nii.gz", + **inputs["spim"].wildcards, + ), + params: + zarrnii_kwargs=zarrnii_in_kwargs, + tile=config["plaque_tile"], + stride=config["plaque_stride"], + batch_size=config["plaque_batch_size"], + chunk=config["plaque_chunk"], + n_gpus=config["plaque_n_gpus"], + output: + probseg=bids_oz_out( + root=root, + datatype="seg", + stain="{stain}", + level="{level}", + desc=config["plaque_seg_method"], + suffix="probseg.{ext}", + **inputs["spim"].wildcards, + ), + wildcard_constraints: + # Trained on amyloid-beta only. Constraining the stain per-rule (never + # globally -- a module-level constraint in an included smk applies to the + # whole workflow) means a target asking for a lantern mask on any other + # stain fails at DAG build instead of quietly burning a GPU on it. + stain=stain_for_plaques or "^$", + threads: 8 * config["plaque_n_gpus"] + resources: + gpu=config["plaque_n_gpus"], + # Suppress the executor's default --ntasks-per-gpu=1. With --gpus=N that + # asks SLURM for N tasks, so srun launches the whole job N times, each task + # pinned to one GPU: the ensemble runs N times over the same volume, every + # copy sees a single device, and they race to write the same output store. + # Setting this to 0 skips the flag entirely (submit_string.py:82), leaving + # one task that sees all N GPUs -- which is what DevicePool expects. + tasks_per_gpu=0, + cpus_per_gpu=8, + mem_mb=256000, + disk_mb=2097152, + runtime=lambda wildcards: max( + 1, + int( + 2880.0 + / (3.0 ** float(wildcards.level)) + / float(config["plaque_n_gpus"]) + ), + ), + # stride-64 tiling x 5 folds, split across GPUs + script: + "../scripts/lantern_plaques.py" + + +rule binarize_lantern_plaques: + """Threshold the LANTERN vote map and upsample it to the segmentation level. + + Emits 0/100 at config['segmentation_level'], matching the segmentation.smk + mask convention, so the standard fieldfrac / regionprops / counts / segstats + chain consumes it unchanged. + """ + input: + probseg=bids_oz_in( + root=root, + datatype="seg", + stain="{stain}", + level=config["plaque_level"], + desc=config["plaque_seg_method"], + suffix="probseg.{ext}", + **inputs["spim"].wildcards, + ), + spim=inputs["spim"].path, + params: + zarrnii_kwargs=zarrnii_in_kwargs, + n_folds=config["plaque_n_folds"], + vote_threshold=config["plaque_vote_threshold"], + target_level=config["segmentation_level"], + output: + mask=bids_oz_out( + root=root, + datatype="seg", + stain="{stain}", + level=config["segmentation_level"], + desc=config["plaque_seg_method"], + suffix="mask.{ext}", + **inputs["spim"].wildcards, + ), + wildcard_constraints: + stain=stain_for_plaques or "^$", + threads: 32 + resources: + mem_mb=64000, + disk_mb=2097152, + runtime=180, + script: + "../scripts/binarize_lantern_plaques.py" diff --git a/spimquant/workflow/rules/regionprops.smk b/spimquant/workflow/rules/regionprops.smk index d3594d6..78293fd 100644 --- a/spimquant/workflow/rules/regionprops.smk +++ b/spimquant/workflow/rules/regionprops.smk @@ -68,24 +68,35 @@ rule transform_regionprops_to_template: "../scripts/transform_regionprops_to_template.py" +def get_regionprops_parquets_to_aggregate(wildcards): + """Per-stain regionprops for the stains this method actually segmented. + + Cannot be resolved at parse time: a single-stain method such as the plaque + ensemble aggregates over only its own stain, not all of stains_for_seg. + """ + paths = expand( + bids( + root=root, + datatype="tabular", + stain="{stain}", + desc="{desc}", + space="{template}", + suffix="regionprops.parquet", + **inputs["spim"].wildcards, + ), + stain=stains_for_desc(wildcards.desc), + allow_missing=True, + ) + return [path.format(**wildcards) for path in paths] + + rule aggregate_regionprops_across_stains: """Aggregate transformed regionprops across stains.""" input: - regionprops_parquets=expand( - bids( - root=root, - datatype="tabular", - stain="{stain}", - desc="{desc}", - space="{template}", - suffix="regionprops.parquet", - **inputs["spim"].wildcards, - ), - stain=stains_for_seg, - allow_missing=True, - ), + regionprops_parquets=get_regionprops_parquets_to_aggregate, params: - stains=stains_for_seg, + # must stay aligned with regionprops_parquets, which is per-method + stains=lambda wildcards: stains_for_desc(wildcards.desc), output: regionprops_aggregated_parquet=bids( root=root, diff --git a/spimquant/workflow/rules/segstats.smk b/spimquant/workflow/rules/segstats.smk index 39e016b..f9e7afc 100644 --- a/spimquant/workflow/rules/segstats.smk +++ b/spimquant/workflow/rules/segstats.smk @@ -149,6 +149,33 @@ rule merge_into_segstats_tsv: "../scripts/merge_into_segstats_tsv.py" +def get_coloc_reference_fieldfrac(wildcards): + """Fieldfrac table used only to recover ROI volume, so count becomes density. + + Colocalization is cross-stain, so every stain this method segmented carries the + same ROI volume column and the first one is an arbitrary but stable choice. + Which stains a method segments depends on the method -- the plaque method covers + only its own stain -- so this cannot be resolved at parse time. + """ + stains = stains_for_desc(wildcards.desc) + if not stains: + raise ValueError( + f"--seg_method {wildcards.desc} segments no stains, so colocalization " + f"statistics cannot be computed for it" + ) + return bids( + root=root, + datatype="tabular", + seg="{seg}", + from_="{template}", + stain=stains[0], + level=config["registration_level"], + desc="{desc}", + suffix="fieldfracstats.tsv", + **inputs["spim"].wildcards, + ).format(**wildcards) + + rule merge_into_colocsegstats_tsv: """ also includes fieldfracstats.tsv to obtain the volume to turn count into density""" input: @@ -170,17 +197,7 @@ rule merge_into_colocsegstats_tsv: suffix="coloccountstats.tsv", **inputs["spim"].wildcards, ), - fieldfrac_tsv=bids( - root=root, - datatype="tabular", - seg="{seg}", - from_="{template}", - stain=stains_for_seg[0], - level=config["registration_level"], - desc="{desc}", - suffix="fieldfracstats.tsv", - **inputs["spim"].wildcards, - ), + fieldfrac_tsv=get_coloc_reference_fieldfrac, params: columns_to_drop=["fieldfrac"], output: @@ -203,44 +220,54 @@ rule merge_into_colocsegstats_tsv: "../scripts/merge_into_segstats_tsv.py" -def get_coloc_tsv_input_kwargs(): - """return coloc_tsv only if we have multiple stains to segment""" - if len(stains_for_seg) == 1: - return {} - else: - return { - "coloc_tsv": bids( - root=root, - datatype="tabular", - seg="{seg}", - from_="{template}", - desc="{desc}", - suffix="colocsegstats.tsv", - **inputs["spim"].wildcards, - ) - } +def get_coloc_tsv_input(wildcards): + """Colocalization table, only for methods that segment two or more stains. + + Single-stain methods (the plaque ensemble) have nothing to colocalize, so this + resolves to an empty list and the merge falls back to the per-stain tables. + """ + if wildcards.desc not in coloc_seg_methods: + return [] + return bids( + root=root, + datatype="tabular", + seg="{seg}", + from_="{template}", + desc="{desc}", + suffix="colocsegstats.tsv", + **inputs["spim"].wildcards, + ).format(**wildcards) + + +def get_indiv_segstats_tsvs(wildcards): + """Per-stain segstats tables for the stains this method actually segmented.""" + paths = expand( + bids( + root=root, + datatype="tabular", + seg="{seg}", + from_="{template}", + stain="{stain}", + level=config["registration_level"], + desc="{desc}", + suffix="segstats.tsv", + **inputs["spim"].wildcards, + ), + stain=stains_for_desc(wildcards.desc), + allow_missing=True, + ) + # Input functions are handed to snakemake verbatim, so the remaining wildcards + # have to be substituted here rather than left for the usual resolution pass. + return [path.format(**wildcards) for path in paths] rule merge_indiv_and_coloc_segstats_tsv: input: - **get_coloc_tsv_input_kwargs(), - indiv_tsvs=expand( - bids( - root=root, - datatype="tabular", - seg="{seg}", - from_="{template}", - stain="{stain}", - level=config["registration_level"], - desc="{desc}", - suffix="segstats.tsv", - **inputs["spim"].wildcards, - ), - stain=stains_for_seg, - allow_missing=True, - ), + coloc_tsv=get_coloc_tsv_input, + indiv_tsvs=get_indiv_segstats_tsvs, params: - stains=stains_for_seg, + # must stay aligned with indiv_tsvs, which is per-method + stains=lambda wildcards: stains_for_desc(wildcards.desc), output: merged_tsv=bids( root=root, diff --git a/spimquant/workflow/scripts/binarize_lantern_plaques.py b/spimquant/workflow/scripts/binarize_lantern_plaques.py new file mode 100644 index 0000000..83b9115 --- /dev/null +++ b/spimquant/workflow/scripts/binarize_lantern_plaques.py @@ -0,0 +1,85 @@ +"""Threshold the LANTERN vote map and upsample it to the segmentation level. + +The ensemble runs at ``plaque_level`` because it is scale-sensitive and was built +for that grid; the segmentation level output is therefore the upsampled +prediction, NOT native inference at that level. + +Upsampling is exact replication (``da.repeat``) rather than interpolation, and +the per-axis factors are derived from the shape ratio so this stays correct when +the pyramid downsamples only x and y. + +Output is valued 0/100 rather than 0/1, matching ``gmmthresh`` / ``multiotsu`` / +``threshold`` in segmentation.smk, so that mean-pool downsampling in ``fieldfrac`` +yields a percentage directly. +""" + +import dask.array as da +import numpy as np +from dask.diagnostics import ProgressBar +from zarrnii import ZarrNii + + +def upsample_nearest(arr, target_shape): + """Replicate voxels to `target_shape`, then trim/pad to land on it exactly.""" + if len(target_shape) != arr.ndim: + raise ValueError(f"rank mismatch: {arr.shape} vs target {target_shape}") + + out = arr + for axis, (have, want) in enumerate(zip(arr.shape, target_shape)): + factor = max(1, int(round(want / have))) + if factor > 1: + out = da.repeat(out, factor, axis=axis) + + out = out[tuple(slice(0, min(h, w)) for h, w in zip(out.shape, target_shape))] + pad = [(0, w - h) for h, w in zip(out.shape, target_shape)] + if any(after for _, after in pad): + out = da.pad(out, pad, mode="edge") + return out + + +def main(): + n_folds = int(snakemake.params.n_folds) + vote_threshold = int(snakemake.params.vote_threshold) + + if not 1 <= vote_threshold <= n_folds: + raise ValueError( + f"plaque_vote_threshold={vote_threshold} is out of range for an " + f"{n_folds}-fold ensemble; it must be between 1 and {n_folds}" + ) + + probseg = ZarrNii.from_file(snakemake.input.probseg, level=0) + ref = ZarrNii.from_file( + snakemake.input.spim, + level=int(snakemake.params.target_level), + channel_labels=[snakemake.wildcards.stain], + **snakemake.params.zarrnii_kwargs, + ) + + # probseg holds votes/n_folds. Compare against the MIDPOINT between adjacent + # attainable values rather than the exact fraction: 2.5/5 = 0.5 sits safely + # between 0.4 and 0.6, so no float representation error can flip a voxel at the + # boundary the way `>= 3/5` could. + binary = probseg.data >= (vote_threshold - 0.5) / n_folds + + upsampled = upsample_nearest(binary, ref.data.shape) + + znimg_mask = ref.copy() + znimg_mask.data = (upsampled * 100).astype(np.uint8) + + print( + f"votes>={vote_threshold} of {n_folds} | " + f"{probseg.data.shape} -> {ref.data.shape}", + flush=True, + ) + + with ProgressBar(): + znimg_mask.to_ome_zarr( + snakemake.output.mask, + max_layer=5, + match_scale_factors_from=snakemake.input.spim, + **snakemake.config["zarrnii_out_kwargs"], + ) + + +if __name__ == "__main__": + main() diff --git a/spimquant/workflow/scripts/lantern_model.py b/spimquant/workflow/scripts/lantern_model.py new file mode 100644 index 0000000..8ab61bc --- /dev/null +++ b/spimquant/workflow/scripts/lantern_model.py @@ -0,0 +1,73 @@ +"""ResEncL U-Net architecture for the LANTERN Abeta plaque segmentation ensemble. + +Vendored from the LANTERN repository (``lantern/encoders/resenc_l.py`` and +``lantern/models/unet_seg.py``) so that SPIMquant can load the published +checkpoints without depending on that repository. The definitions must stay +byte-compatible with the checkpoints on HuggingFace +(``apooladi/lantern-ki3-abeta``) -- changing any architectural constant here +will make ``load_state_dict`` fail. + +The encoder is identical to nnssl's ``architecture_registry.get_res_enc_l()``; +the decoder is nnU-Net's ``UNetDecoder`` over that encoder's skips. +""" + +from __future__ import annotations + +import torch +import torch.nn as nn +from dynamic_network_architectures.building_blocks.residual_encoders import ( + ResidualEncoder, +) +from dynamic_network_architectures.building_blocks.unet_decoder import UNetDecoder + + +def build_resenc_l(num_input_channels: int = 1) -> ResidualEncoder: + """Build the ResEncL encoder. + + Architecture: 6 stages, features [32, 64, 128, 256, 320, 320], + strides [1,1,1] then [2,2,2] x 5, blocks per stage [1, 3, 4, 6, 6, 6], + InstanceNorm3d (eps=1e-5, affine=True), LeakyReLU(inplace=True). + """ + n_stages = 6 + encoder = ResidualEncoder( + input_channels=num_input_channels, + n_stages=n_stages, + features_per_stage=[32, 64, 128, 256, 320, 320], + conv_op=nn.Conv3d, + kernel_sizes=[[3, 3, 3]] * n_stages, + strides=[[1, 1, 1]] + [[2, 2, 2]] * 5, + n_blocks_per_stage=[1, 3, 4, 6, 6, 6], + conv_bias=True, + norm_op=nn.InstanceNorm3d, + norm_op_kwargs={"eps": 1e-5, "affine": True}, + nonlin=nn.LeakyReLU, + nonlin_kwargs={"inplace": True}, + return_skips=True, + disable_default_stem=False, + stem_channels=None, + ) + return encoder + + +class ResEncLUNet(nn.Module): + """ResEncL encoder + UNetDecoder, 2-class (background / plaque) logits.""" + + def __init__( + self, + num_classes: int = 2, + num_input_channels: int = 1, + n_conv_per_stage_decoder: int = 1, + deep_supervision: bool = False, + ) -> None: + super().__init__() + self.encoder = build_resenc_l(num_input_channels) + n_dec_stages = len(self.encoder.output_channels) - 1 + self.decoder = UNetDecoder( + self.encoder, + num_classes, + [n_conv_per_stage_decoder] * n_dec_stages, + deep_supervision=deep_supervision, + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.decoder(self.encoder(x)) diff --git a/spimquant/workflow/scripts/lantern_plaques.py b/spimquant/workflow/scripts/lantern_plaques.py new file mode 100644 index 0000000..0eae48e --- /dev/null +++ b/spimquant/workflow/scripts/lantern_plaques.py @@ -0,0 +1,356 @@ +"""5-fold LANTERN ensemble Abeta plaque segmentation. + +Runs the published LANTERN ensemble (``apooladi/lantern-ki3-abeta``) over the +Abeta channel and writes the per-fold agreement map -- the fraction of folds +voting foreground, so ``{0, .2, .4, .6, .8, 1}`` for five folds -- as an +OME-Zarr. Binarization is deliberately left to a downstream rule so the vote +threshold can be changed without re-running inference. + +Deviations from ``vesselfm.py``, and why: + +- ``ZarrNii.segment()`` cannot be used. It dispatches through ``da.map_blocks`` + with disjoint blocks, no block position and a hardcoded ``uint8`` output, so it + can express neither the stride-64 tile overlap the model requires nor + brain-mask-based block skipping. ``da.map_overlap`` is driven directly here. + +- The model reads RAW intensities. It was never trained on N4-corrected or + otherwise rescaled data, so this path takes ``input.spim`` directly rather than + the ``desc-corrected*`` store the GMM/Otsu rules consume. + +Inference contract (from the model card, must match training exactly): + tiles of 128^3 at stride 64, per-tile 0.5/99.5 percentile clip then z-score, + softmax over 2 classes, foreground where class-1 probability >= 0.5 per fold, + overlapping tiles combined by MAX over votes. +""" + +import contextlib +import queue + +import dask.array as da +import numpy as np +import torch +from dask.diagnostics import ProgressBar +from zarrnii import ZarrNii + +from dask_setup import get_dask_client +from lantern_model import ResEncLUNet + + +def tile_origins(extent, tile, stride): + """Tile origins covering [0, extent), with the last pulled back to hit the edge.""" + if extent <= tile: + return [0] + origins = list(range(0, extent - tile + 1, stride)) + if origins[-1] + tile < extent: + origins.append(extent - tile) + return origins + + +def normalize_tiles(x): + """Per-tile percentile clip then z-score, matching LANTERN's training transform.""" + out = torch.empty_like(x) + quantiles = torch.tensor([0.005, 0.995], device=x.device, dtype=x.dtype) + for i in range(x.shape[0]): + lo, hi = torch.quantile(x[i].reshape(-1), quantiles) + clipped = torch.clamp(x[i], lo, hi) + std = clipped.std() + out[i] = (clipped - clipped.mean()) / (std if std > 1e-8 else 1.0) + return out + + +def predict_volume(vol, nets, device, tile, stride, batch_size): + """Tiled 5-fold vote map for one 3D block, as uint8 in [0, len(nets)]. + + The caller owns `device` exclusively for the duration of this call (see + DevicePool), so no additional locking is needed here. + """ + orig_shape = vol.shape + pad = [max(0, tile - n) for n in orig_shape] + if any(pad): + # Blocks at the volume border can be thinner than one tile; pad, then crop back. + vol = np.pad(vol, [(0, p) for p in pad], mode="edge") + + votes = np.zeros(vol.shape, dtype=np.uint8) + coords = [ + (z, y, x) + for z in tile_origins(vol.shape[0], tile, stride) + for y in tile_origins(vol.shape[1], tile, stride) + for x in tile_origins(vol.shape[2], tile, stride) + ] + + for start in range(0, len(coords), batch_size): + batch = coords[start : start + batch_size] + arr = np.stack( + [vol[z : z + tile, y : y + tile, x : x + tile] for z, y, x in batch] + ) + tensor = torch.from_numpy(arr).to(device) + tensor = normalize_tiles(tensor)[:, None] + batch_votes = torch.zeros( + (len(batch), tile, tile, tile), dtype=torch.uint8, device=device + ) + with ( + torch.no_grad(), + torch.autocast( + device.type, dtype=torch.bfloat16, enabled=device.type == "cuda" + ), + ): + for net in nets: + prob = torch.softmax(net(tensor).float(), dim=1)[:, 1] + batch_votes += (prob >= 0.5).to(torch.uint8) + batch_votes = batch_votes.cpu().numpy() + + for j, (z, y, x) in enumerate(batch): + region = votes[z : z + tile, y : y + tile, x : x + tile] + # MAX, not sum: a plaque clipped at one tile edge is recovered by the + # tile that contains it whole. + # + # Folds are summed WITHIN a tile (above) and tiles are combined with max + # (here). Those two steps do not commute -- max-then-sum would score a + # voxel higher when different folds find it from different tiles -- but + # this order is the one LANTERN's own inference uses + # (scripts/infer_wholebrain.py), and plaque_vote_threshold was calibrated + # under it. Swapping them would silently change what "3 of 5" means. + np.maximum(region, batch_votes[j], out=region) + + return votes[: orig_shape[0], : orig_shape[1], : orig_shape[2]] + + +def load_folds(model_paths, device): + """Instantiate every fold on `device`, once, in the main process. + + ~2 GB of fp32 weights for a 5-fold ensemble, so replicating it per GPU is + cheap. DevicePool then guarantees only one dask block uses a given device at + a time, which bounds peak VRAM without needing a lock inside the forward pass. + + ``weights_only=True`` because these checkpoints are downloaded over the + network; they carry only tensors and plain dicts, so nothing is lost. + """ + nets = [] + for path in model_paths: + ckpt = torch.load(path, map_location="cpu", weights_only=True) + if ckpt.get("model_type") not in (None, "ResEncLUNet"): + raise ValueError( + f"{path} declares model_type={ckpt['model_type']!r}, " + "but this script only builds ResEncLUNet" + ) + net = ResEncLUNet(num_classes=2, num_input_channels=1) + net.load_state_dict(ckpt["state_dict"]) + nets.append(net.eval().to(device)) + return nets + + +def resolve_devices(requested): + """The devices to run on, clamped to what is actually visible.""" + if not torch.cuda.is_available(): + print("WARNING: no CUDA device visible; falling back to CPU", flush=True) + return [torch.device("cpu")] + available = torch.cuda.device_count() + if requested > available: + print( + f"WARNING: {requested} GPUs requested but only {available} visible; " + f"using {available}", + flush=True, + ) + return [torch.device(f"cuda:{i}") for i in range(max(1, min(requested, available)))] + + +class DevicePool: + """Lends out one GPU at a time, with the full ensemble resident on each. + + The ensemble is replicated per device (~2 GB of fp32 weights for 5 folds), so + each GPU works independently on a whole dask block. Handing out devices from + a queue rather than striping blocks by index balances the load for free: + blocks vary enormously in cost because most of a frame is background, so a + static assignment would leave one card idle -- the same imbalance LANTERN's + own two-GPU split had to correct for by splitting on cumulative tile count. + """ + + def __init__(self, devices, model_paths): + self._free = queue.Queue() + self.nets = {} + for device in devices: + self.nets[device] = load_folds(model_paths, device) + self._free.put(device) + + @contextlib.contextmanager + def acquire(self): + device = self._free.get() + try: + yield device, self.nets[device] + finally: + self._free.put(device) + + +def build_brain_block_mask(mask_path, image_shape, chunks): + """Per-block "does this block touch brain" flags, aligned with the image chunking. + + The brain mask lives at a much coarser pyramid level, so rather than + interpolating it up to the inference grid it is held in memory and indexed by + each block's array location -- the same approach LANTERN's whole-brain + inference uses to skip background tiles. + """ + mask_znimg = ZarrNii.from_nifti(mask_path, axes_order="ZYX") + mask_arr = np.asarray(mask_znimg.data).squeeze() > 0 + if mask_arr.ndim != 3: + raise ValueError(f"expected a 3D brain mask, got shape {mask_arr.shape}") + + scales = [m / i for m, i in zip(mask_arr.shape, image_shape[-3:])] + + def block_has_brain(block_info=None): + shape = block_info[None]["chunk-shape"] + location = block_info[None]["array-location"] + index = [] + for axis, (lo, hi) in enumerate(location[-3:]): + i0 = int(np.floor(lo * scales[axis])) + i1 = int(np.ceil(hi * scales[axis])) + i0 = min(max(i0, 0), mask_arr.shape[axis] - 1) + i1 = min(max(i1, i0 + 1), mask_arr.shape[axis]) + index.append(slice(i0, i1)) + present = bool(mask_arr[index[0], index[1], index[2]].any()) + return np.broadcast_to(np.bool_(present), shape) + + return da.map_blocks( + block_has_brain, + chunks=chunks, + dtype=bool, + meta=np.array([], dtype=bool), + ) + + +def plan_chunks(shape, block, halo): + """Chunk sizes for `shape` at `block`, with no chunk shorter than `halo`. + + A plain rechunk leaves a remainder chunk that can be far shorter than the + halo (z=1030 at block 512 leaves a 6-voxel tail), and dask refuses to overlap + a chunk smaller than the depth. A short tail is absorbed into its + predecessor instead. Every chunk start stays a multiple of `block`, so the + per-block tile grid stays aligned with the global one. + """ + chunks = [(shape[0],)] + for extent in shape[1:]: + if extent <= block: + chunks.append((extent,)) + continue + full, remainder = divmod(extent, block) + if remainder == 0: + chunks.append((block,) * full) + elif remainder < halo: + chunks.append((block,) * (full - 1) + (block + remainder,)) + else: + chunks.append((block,) * full + (remainder,)) + return tuple(chunks) + + +def vote_map(data, brain, predict, tile, stride): + """Lazily map `predict` over `data`, skipping blocks that carry no brain. + + The halo is `tile - stride`, so every voxel in a block's core is covered by + at least one tile lying wholly inside the block that produced it -- which is + what makes the stride-64 overlap correct across block boundaries and not just + within them. `brain` is halo-expanded alongside `data`, so the skip test is + conservative: a block is processed if its core *or* its halo touches brain. + + An axis that fits in a single chunk gets no halo: there is no block boundary + to bridge, and dask rejects a depth larger than the chunk. + """ + halo = tile - stride + + depth = {} + for axis, axis_chunks in enumerate(data.chunks): + if axis == 0 or len(axis_chunks) == 1: + depth[axis] = 0 + continue + if min(axis_chunks) < halo: + raise ValueError( + f"axis {axis} has a chunk of {min(axis_chunks)} voxels, shorter " + f"than the {halo}-voxel halo; use plan_chunks() to lay out the " + f"chunks before calling vote_map()" + ) + depth[axis] = halo + + def segment_block(image_block, brain_block): + if not brain_block.any(): + return np.zeros(image_block.shape, dtype=np.uint8) + out = np.empty(image_block.shape, dtype=np.uint8) + for c in range(image_block.shape[0]): + out[c] = predict(image_block[c].astype(np.float32)) + return out + + return da.map_overlap( + segment_block, + data, + brain, + depth=depth, + boundary="none", + trim=True, + allow_rechunk=False, + dtype=np.uint8, + meta=np.array([], dtype=np.uint8), + ) + + +def main(): + tile = int(snakemake.params.tile) + stride = int(snakemake.params.stride) + batch_size = int(snakemake.params.batch_size) + chunk = int(snakemake.params.chunk) + n_gpus = int(snakemake.params.n_gpus) + + devices = resolve_devices(n_gpus) + + with get_dask_client("threads", snakemake.threads): + znimg = ZarrNii.from_file( + snakemake.input.spim, + level=int(snakemake.wildcards.level), + channel_labels=[snakemake.wildcards.stain], + **snakemake.params.zarrnii_kwargs, + ) + + if znimg.data.ndim != 4 or znimg.data.shape[0] != 1: + raise ValueError( + f"expected a single-channel (c,z,y,x) image, got shape {znimg.data.shape}" + ) + data = znimg.data.rechunk(plan_chunks(znimg.data.shape, chunk, tile - stride)) + + brain = build_brain_block_mask(snakemake.input.mask, data.shape, data.chunks) + + pool = DevicePool(devices, snakemake.input.models) + n_folds = len(pool.nets[devices[0]]) + print( + f"loaded {n_folds} folds on each of {[str(d) for d in devices]}; " + f"grid {data.shape} chunk {chunk} halo {tile - stride} " + f"tile {tile} stride {stride} batch {batch_size}", + flush=True, + ) + + def predict(vol): + with pool.acquire() as (device, nets): + return predict_volume(vol, nets, device, tile, stride, batch_size) + + votes = vote_map(data, brain, predict, tile, stride) + + # The fraction of folds voting foreground: float32 in {0, .2, .4, .6, .8, 1} + # for a 5-fold ensemble. Same form as LANTERN's own whole-brain probmask. + # + # Deliberately a probability rather than the raw count or a 0-100 rescaling: + # it is directly readable as model confidence in a viewer, and independent of + # how many folds the ensemble happens to have. The 0/100 mask convention does + # not apply -- that exists so fieldfrac can mean-pool a *mask* into a + # percentage, and nothing but binarize_lantern_plaques reads this store. + # + # float32 costs ~4 bytes/voxel, so a level-1 whole brain is ~100 GB raw; it + # compresses hard, being overwhelmingly zero. + znimg_prob = znimg.copy() + znimg_prob.data = votes.astype(np.float32) / float(n_folds) + + with ProgressBar(): + znimg_prob.to_ome_zarr( + snakemake.output.probseg, + max_layer=5, + match_scale_factors_from=snakemake.input.spim, + **snakemake.config["zarrnii_out_kwargs"], + ) + + +if __name__ == "__main__": + main() diff --git a/spimquant/workflow/scripts/merge_indiv_and_coloc_segstats_tsv.py b/spimquant/workflow/scripts/merge_indiv_and_coloc_segstats_tsv.py index 2aee731..12924fd 100644 --- a/spimquant/workflow/scripts/merge_indiv_and_coloc_segstats_tsv.py +++ b/spimquant/workflow/scripts/merge_indiv_and_coloc_segstats_tsv.py @@ -10,7 +10,9 @@ import pandas as pd indiv_files = snakemake.input.indiv_tsvs -coloc_file = getattr(snakemake.input, "coloc_tsv", None) +# Single-stain methods resolve coloc_tsv to an empty list rather than dropping the +# key, so an empty value has to be treated the same as a missing one. +coloc_file = getattr(snakemake.input, "coloc_tsv", None) or None output_file = snakemake.output.merged_tsv stains = snakemake.params.stains # list aligned to indiv_files