lavlab-shell 0.3.2__tar.gz → 0.5.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- lavlab_shell-0.5.0/.github/workflows/build.yml +31 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.github/workflows/publish.yml +1 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.github/workflows/pytest.yml +6 -6
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/PKG-INFO +4 -2
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/pyproject.toml +24 -3
- lavlab_shell-0.5.0/scripts/export_onnx.py +90 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/__about__.py +1 -1
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/__init__.py +8 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/cli.py +65 -4
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/infer_omero_wsi.py +59 -12
- lavlab_shell-0.5.0/src/shell/infer_wsi.py +816 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/inference.py +44 -34
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/model.py +102 -41
- lavlab_shell-0.5.0/src/shell/post_process.py +1360 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/transforms.py +7 -3
- lavlab_shell-0.5.0/src/shell/weights/model_v1.onnx +0 -0
- lavlab_shell-0.5.0/src/shell/weights/model_v1.onnx.data +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/weights/model_v1.pth +0 -0
- lavlab_shell-0.5.0/tests/test_shell.py +68 -0
- lavlab_shell-0.3.2/.github/workflows/build.yml +0 -54
- lavlab_shell-0.3.2/src/shell/infer_wsi.py +0 -373
- lavlab_shell-0.3.2/tests/test_shell.py +0 -16
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.devcontainer/Dockerfile +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.editorconfig +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.github/dependabot.yml +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.github/workflows/lint.yml +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.gitignore +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.pre-commit-config.yaml +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/CONTRIBUTING.md +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/Dockerfile +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/LICENSE.txt +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/Makefile +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/README.md +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/docs/api.md +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/docs/index.md +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/mkdocs.yml +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/requirements/requirements-docs.txt +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/requirements/requirements-lint.txt +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/requirements/requirements-test.txt +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/requirements/requirements-types.txt +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/requirements.txt +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/benchmark.py +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/py.typed +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/tests/__init__.py +0 -0
- {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/tests/conftest.py +0 -0
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
name: Build Wheel using Hatch
|
|
2
|
+
|
|
3
|
+
on:
|
|
4
|
+
push:
|
|
5
|
+
branches:
|
|
6
|
+
- main
|
|
7
|
+
pull_request:
|
|
8
|
+
|
|
9
|
+
jobs:
|
|
10
|
+
build-wheel:
|
|
11
|
+
runs-on: ubuntu-latest
|
|
12
|
+
steps:
|
|
13
|
+
- uses: actions/checkout@v4
|
|
14
|
+
- name: Set up Python
|
|
15
|
+
uses: actions/setup-python@v4
|
|
16
|
+
with:
|
|
17
|
+
python-version: '3.12'
|
|
18
|
+
|
|
19
|
+
- name: Install build tools
|
|
20
|
+
run: |
|
|
21
|
+
python -m pip install --upgrade pip setuptools wheel
|
|
22
|
+
python -m pip install "virtualenv==20.23.1"
|
|
23
|
+
python -m pip install --upgrade hatch
|
|
24
|
+
|
|
25
|
+
- name: Build wheel with hatch
|
|
26
|
+
run: hatch build
|
|
27
|
+
- name: Upload wheel
|
|
28
|
+
uses: actions/upload-artifact@v4
|
|
29
|
+
with:
|
|
30
|
+
name: wheel
|
|
31
|
+
path: dist/*.whl
|
|
@@ -7,6 +7,8 @@ on:
|
|
|
7
7
|
jobs:
|
|
8
8
|
test:
|
|
9
9
|
runs-on: ubuntu-latest
|
|
10
|
+
env:
|
|
11
|
+
CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}
|
|
10
12
|
steps:
|
|
11
13
|
- uses: actions/checkout@v4
|
|
12
14
|
|
|
@@ -18,11 +20,11 @@ jobs:
|
|
|
18
20
|
- name: Install test dependencies
|
|
19
21
|
run: |
|
|
20
22
|
python -m pip install --upgrade pip
|
|
21
|
-
python -m pip install
|
|
23
|
+
python -m pip install hatch
|
|
22
24
|
|
|
23
25
|
- name: Run pytest and generate coverage report
|
|
24
26
|
run: |
|
|
25
|
-
|
|
27
|
+
hatch run test:cov --cov-report=xml:coverage.xml
|
|
26
28
|
|
|
27
29
|
- name: Upload coverage artifact
|
|
28
30
|
uses: actions/upload-artifact@v4
|
|
@@ -31,14 +33,12 @@ jobs:
|
|
|
31
33
|
path: ${{ github.workspace }}/coverage.xml
|
|
32
34
|
|
|
33
35
|
- name: Upload coverage report to Codecov
|
|
34
|
-
if:
|
|
36
|
+
if: env.CODECOV_TOKEN != ''
|
|
35
37
|
uses: codecov/codecov-action@v4
|
|
36
38
|
with:
|
|
37
39
|
fail_ci_if_error: true
|
|
38
40
|
files: ${{ github.workspace }}/coverage.xml
|
|
39
|
-
env:
|
|
40
|
-
CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}
|
|
41
41
|
|
|
42
42
|
- name: Skip Codecov upload (no token)
|
|
43
|
-
if:
|
|
43
|
+
if: env.CODECOV_TOKEN == ''
|
|
44
44
|
run: echo "Skipping Codecov — CODECOV_TOKEN not set."
|
|
@@ -1,6 +1,6 @@
|
|
|
1
|
-
Metadata-Version: 2.
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
2
|
Name: lavlab-shell
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.5.0
|
|
4
4
|
Summary: SHELL Highlights Epithelium and Lumen Locations — whole-slide H&E segmentation
|
|
5
5
|
Project-URL: Documentation, https://github.com/laviolette-lab/shell#readme
|
|
6
6
|
Project-URL: Issues, https://github.com/laviolette-lab/shell/issues
|
|
@@ -19,6 +19,8 @@ Requires-Python: <3.13,>=3.10
|
|
|
19
19
|
Requires-Dist: macenko-pca
|
|
20
20
|
Requires-Dist: monai
|
|
21
21
|
Requires-Dist: numpy
|
|
22
|
+
Requires-Dist: onnxruntime-gpu>=1.17.0; platform_system == 'Linux' and platform_machine != 'arm64'
|
|
23
|
+
Requires-Dist: onnxruntime>=1.17.0; platform_system != 'Linux' or platform_machine == 'arm64'
|
|
22
24
|
Requires-Dist: openslide-python
|
|
23
25
|
Requires-Dist: pyvips
|
|
24
26
|
Requires-Dist: scikit-image
|
|
@@ -1,5 +1,5 @@
|
|
|
1
1
|
[build-system]
|
|
2
|
-
requires = ["hatchling"]
|
|
2
|
+
requires = ["hatch", "hatchling"]
|
|
3
3
|
build-backend = "hatchling.build"
|
|
4
4
|
|
|
5
5
|
[project]
|
|
@@ -34,6 +34,8 @@ dependencies = [
|
|
|
34
34
|
"openslide-python",
|
|
35
35
|
"macenko-pca",
|
|
36
36
|
"scikit-image",
|
|
37
|
+
"onnxruntime-gpu>=1.17.0; platform_system == 'Linux' and platform_machine != 'arm64'",
|
|
38
|
+
"onnxruntime>=1.17.0; platform_system != 'Linux' or platform_machine == 'arm64'",
|
|
37
39
|
]
|
|
38
40
|
|
|
39
41
|
[project.optional-dependencies]
|
|
@@ -53,8 +55,11 @@ path = "src/shell/__about__.py"
|
|
|
53
55
|
|
|
54
56
|
[tool.hatch.build.targets.wheel]
|
|
55
57
|
packages = ["src/shell"]
|
|
58
|
+
exclude = ["src/shell/weights/*.pth"]
|
|
59
|
+
artifacts = ["src/shell/weights/*.onnx"]
|
|
56
60
|
|
|
57
61
|
[tool.hatch.envs.default]
|
|
62
|
+
installer = "uv"
|
|
58
63
|
dependencies = ["setuptools>=82.0.0"]
|
|
59
64
|
|
|
60
65
|
[tool.hatch.envs.default.scripts]
|
|
@@ -100,6 +105,16 @@ dependencies = [
|
|
|
100
105
|
build-docs = "mkdocs build"
|
|
101
106
|
serve-docs = "mkdocs serve"
|
|
102
107
|
|
|
108
|
+
[tool.hatch.envs.export]
|
|
109
|
+
dependencies = [
|
|
110
|
+
"onnx>=1.16.0",
|
|
111
|
+
"onnxruntime>=1.17.0",
|
|
112
|
+
"onnxscript>=0.1.0",
|
|
113
|
+
"setuptools>=82.0.0",
|
|
114
|
+
]
|
|
115
|
+
[tool.hatch.envs.export.scripts]
|
|
116
|
+
export-model = "python scripts/export_onnx.py"
|
|
117
|
+
|
|
103
118
|
[[tool.hatch.envs.all.matrix]]
|
|
104
119
|
python = ["3.10", "3.11", "3.12"]
|
|
105
120
|
|
|
@@ -145,8 +160,11 @@ ignore = []
|
|
|
145
160
|
# Scientific code uses uppercase variable names (Io, W, H) and import aliases
|
|
146
161
|
# (F as torch.nn.functional, BlitzGateway as _BG) that violate pep8-naming.
|
|
147
162
|
"src/shell/preprocessing.py" = ["N803", "N806"]
|
|
148
|
-
"src/shell/inference.py" = ["N812"]
|
|
149
|
-
"src/shell/infer_omero_wsi.py" = ["N814", "F401"]
|
|
163
|
+
"src/shell/inference.py" = ["N812", "E501"]
|
|
164
|
+
"src/shell/infer_omero_wsi.py" = ["N814", "F401", "N806", "E501"]
|
|
165
|
+
"src/shell/infer_wsi.py" = ["N812", "N806", "E501", "T201", "F841", "F821"]
|
|
166
|
+
"src/shell/post_process.py" = ["RUF002", "RUF003", "N806", "B905", "T201", "E501"]
|
|
167
|
+
"src/shell/transforms.py" = ["N803", "N806", "RUF002", "E501"]
|
|
150
168
|
|
|
151
169
|
[tool.ruff.lint.isort]
|
|
152
170
|
known-first-party = ["shell"]
|
|
@@ -157,4 +175,7 @@ docstring-code-format = true
|
|
|
157
175
|
[tool.pytest.ini_options]
|
|
158
176
|
testpaths = ["tests"]
|
|
159
177
|
addopts = ["-ra", "--strict-markers", "--strict-config"]
|
|
178
|
+
filterwarnings = [
|
|
179
|
+
'ignore:.*torch.*jit.*interface.*:FutureWarning',
|
|
180
|
+
]
|
|
160
181
|
xfail_strict = true
|
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""Export the bundled SegResNetVAE checkpoint to ONNX."""
|
|
3
|
+
|
|
4
|
+
from __future__ import annotations
|
|
5
|
+
|
|
6
|
+
import argparse
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
|
|
9
|
+
import torch
|
|
10
|
+
|
|
11
|
+
from shell.model import LATEST_MODEL, MODEL_REGISTRY, _resolve_bundled_weights
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def _build_model_for_export(device: torch.device) -> torch.nn.Module:
|
|
15
|
+
"""Create the reference SegResNetVAE model in eval mode."""
|
|
16
|
+
from monai.networks.nets import SegResNetVAE
|
|
17
|
+
|
|
18
|
+
model = SegResNetVAE(
|
|
19
|
+
spatial_dims=2,
|
|
20
|
+
init_filters=16,
|
|
21
|
+
in_channels=3,
|
|
22
|
+
out_channels=3,
|
|
23
|
+
dropout_prob=0.2,
|
|
24
|
+
norm=("GROUP", {"num_groups": 8}),
|
|
25
|
+
act=("MISH", {"inplace": True}),
|
|
26
|
+
input_image_size=(320, 320),
|
|
27
|
+
vae_nz=256,
|
|
28
|
+
vae_estimate_std=True,
|
|
29
|
+
).to(device)
|
|
30
|
+
model.eval()
|
|
31
|
+
return model
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def export_onnx(
|
|
35
|
+
checkpoint: str | None = None,
|
|
36
|
+
*,
|
|
37
|
+
version: str | None = None,
|
|
38
|
+
output: str | None = None,
|
|
39
|
+
device: str = "cpu",
|
|
40
|
+
) -> Path:
|
|
41
|
+
"""Export a bundled checkpoint to ONNX and return the output path."""
|
|
42
|
+
if checkpoint is None:
|
|
43
|
+
if version is None:
|
|
44
|
+
version = LATEST_MODEL
|
|
45
|
+
checkpoint = str(_resolve_bundled_weights(version))
|
|
46
|
+
|
|
47
|
+
dev = torch.device(device)
|
|
48
|
+
model = _build_model_for_export(dev)
|
|
49
|
+
state = torch.load(checkpoint, map_location=dev, weights_only=True)
|
|
50
|
+
model.load_state_dict(state)
|
|
51
|
+
model.eval()
|
|
52
|
+
|
|
53
|
+
if output is None:
|
|
54
|
+
target = Path(checkpoint).with_suffix(".onnx")
|
|
55
|
+
else:
|
|
56
|
+
target = Path(output)
|
|
57
|
+
target.parent.mkdir(parents=True, exist_ok=True)
|
|
58
|
+
|
|
59
|
+
dummy = torch.randn(1, 3, 320, 320, device=dev)
|
|
60
|
+
torch.onnx.export(
|
|
61
|
+
model,
|
|
62
|
+
dummy,
|
|
63
|
+
str(target),
|
|
64
|
+
export_params=True,
|
|
65
|
+
opset_version=17,
|
|
66
|
+
do_constant_folding=True,
|
|
67
|
+
input_names=["input"],
|
|
68
|
+
output_names=["logits"],
|
|
69
|
+
dynamic_axes={
|
|
70
|
+
"input": {0: "batch_size"},
|
|
71
|
+
"logits": {0: "batch_size"},
|
|
72
|
+
},
|
|
73
|
+
)
|
|
74
|
+
return target
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def main() -> None:
|
|
78
|
+
parser = argparse.ArgumentParser(description="Export the latest SHELL checkpoint to ONNX.")
|
|
79
|
+
parser.add_argument("--checkpoint", type=str, default=None, help="Path to a .pth checkpoint to export.")
|
|
80
|
+
parser.add_argument("--version", type=str, default=None, help="Bundled model version to export (e.g. v1).")
|
|
81
|
+
parser.add_argument("--output", type=str, default=None, help="Destination .onnx path.")
|
|
82
|
+
parser.add_argument("--device", type=str, default="cpu", help="Device for export (cpu/cuda/mps).")
|
|
83
|
+
args = parser.parse_args()
|
|
84
|
+
|
|
85
|
+
out = export_onnx(args.checkpoint, version=args.version, output=args.output, device=args.device)
|
|
86
|
+
print(f"Exported ONNX model: {out}")
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
if __name__ == "__main__":
|
|
90
|
+
main()
|
|
@@ -9,6 +9,14 @@ or simple checks does not eagerly import heavy runtime dependencies
|
|
|
9
9
|
(like ``torch``, ``monai``, ``pyvips``, or ``omero``).
|
|
10
10
|
"""
|
|
11
11
|
|
|
12
|
+
import warnings
|
|
13
|
+
|
|
12
14
|
from .__about__ import __version__
|
|
13
15
|
|
|
16
|
+
warnings.filterwarnings(
|
|
17
|
+
"ignore",
|
|
18
|
+
message=r".*torch.*jit.*interface.*",
|
|
19
|
+
category=FutureWarning,
|
|
20
|
+
)
|
|
21
|
+
|
|
14
22
|
__all__ = ["__version__"]
|
|
@@ -105,7 +105,7 @@ def build_parser() -> argparse.ArgumentParser:
|
|
|
105
105
|
infer_p.add_argument(
|
|
106
106
|
"--target-mpp",
|
|
107
107
|
type=float,
|
|
108
|
-
default=
|
|
108
|
+
default=2.0,
|
|
109
109
|
help="Desired output resolution (um/pixel).",
|
|
110
110
|
)
|
|
111
111
|
infer_p.add_argument(
|
|
@@ -125,6 +125,52 @@ def build_parser() -> argparse.ArgumentParser:
|
|
|
125
125
|
default=None,
|
|
126
126
|
help="Optional: also save the intermediate EHO image.",
|
|
127
127
|
)
|
|
128
|
+
infer_p.add_argument(
|
|
129
|
+
"--save-raw",
|
|
130
|
+
type=str,
|
|
131
|
+
default=None,
|
|
132
|
+
help=(
|
|
133
|
+
"Optional: save raw model predictions as a 3-band uint8 image "
|
|
134
|
+
"(band 0 = inner / lumen, band 1 = outer / epithelium, "
|
|
135
|
+
"band 2 = background), scaled to 0/255. "
|
|
136
|
+
"Saved at the original input resolution. "
|
|
137
|
+
"Useful for debugging or re-running post-processing offline."
|
|
138
|
+
),
|
|
139
|
+
)
|
|
140
|
+
infer_p.add_argument(
|
|
141
|
+
"--profile",
|
|
142
|
+
type=str,
|
|
143
|
+
default="best_effort",
|
|
144
|
+
choices=["best_effort", "precise", "sensitive"],
|
|
145
|
+
help=(
|
|
146
|
+
"Post-processing filter profile. "
|
|
147
|
+
"'best_effort' (default) balances sensitivity and precision; "
|
|
148
|
+
"'precise' is more conservative; 'sensitive' keeps more predictions."
|
|
149
|
+
),
|
|
150
|
+
)
|
|
151
|
+
infer_p.add_argument(
|
|
152
|
+
"--mode",
|
|
153
|
+
type=str,
|
|
154
|
+
default="wsi",
|
|
155
|
+
choices=["wsi", "biopsy", "tile"],
|
|
156
|
+
help=(
|
|
157
|
+
"Post-processing mode. "
|
|
158
|
+
"'wsi' (default): full pipeline with tissue restriction and "
|
|
159
|
+
"urethra detection. 'biopsy': tissue restriction but no urethra "
|
|
160
|
+
"detection. 'tile': no tissue mask/urethra; reflect-pads "
|
|
161
|
+
"predictions before morphological ops. Use 'tile' for individual "
|
|
162
|
+
"image tiles that lack surrounding context."
|
|
163
|
+
),
|
|
164
|
+
)
|
|
165
|
+
infer_p.add_argument(
|
|
166
|
+
"--tile-pad",
|
|
167
|
+
type=int,
|
|
168
|
+
default=None,
|
|
169
|
+
help=(
|
|
170
|
+
"Reflect-padding in pixels applied on each side in tile mode. "
|
|
171
|
+
"Defaults to 50%% of the shorter output dimension when not set."
|
|
172
|
+
),
|
|
173
|
+
)
|
|
128
174
|
infer_p.add_argument(
|
|
129
175
|
"--device",
|
|
130
176
|
type=str,
|
|
@@ -156,7 +202,7 @@ def build_parser() -> argparse.ArgumentParser:
|
|
|
156
202
|
omero_p.add_argument(
|
|
157
203
|
"--target-mpp",
|
|
158
204
|
type=float,
|
|
159
|
-
default=
|
|
205
|
+
default=2.0,
|
|
160
206
|
help="Desired output resolution (um/pixel)",
|
|
161
207
|
)
|
|
162
208
|
omero_p.add_argument(
|
|
@@ -252,6 +298,17 @@ def build_parser() -> argparse.ArgumentParser:
|
|
|
252
298
|
"sessions. Default 4."
|
|
253
299
|
),
|
|
254
300
|
)
|
|
301
|
+
omero_p.add_argument(
|
|
302
|
+
"--profile",
|
|
303
|
+
type=str,
|
|
304
|
+
default="best_effort",
|
|
305
|
+
choices=["best_effort", "precise", "sensitive"],
|
|
306
|
+
help=(
|
|
307
|
+
"Post-processing filter profile. "
|
|
308
|
+
"'best_effort' (default) balances sensitivity and precision; "
|
|
309
|
+
"'precise' is more conservative; 'sensitive' keeps more predictions."
|
|
310
|
+
),
|
|
311
|
+
)
|
|
255
312
|
|
|
256
313
|
return parser
|
|
257
314
|
|
|
@@ -303,6 +360,10 @@ def main(argv: list[str] | None = None) -> int:
|
|
|
303
360
|
target_mpp=args.target_mpp,
|
|
304
361
|
mpp=args.mpp,
|
|
305
362
|
save_eho=args.save_eho,
|
|
363
|
+
save_raw=args.save_raw,
|
|
364
|
+
profile=args.profile,
|
|
365
|
+
mode=args.mode,
|
|
366
|
+
tile_pad=args.tile_pad,
|
|
306
367
|
device=args.device,
|
|
307
368
|
)
|
|
308
369
|
|
|
@@ -351,13 +412,13 @@ def main(argv: list[str] | None = None) -> int:
|
|
|
351
412
|
min_tissue_frac=args.min_tissue_frac,
|
|
352
413
|
prefetch_depth=args.prefetch_depth,
|
|
353
414
|
num_fetch_workers=args.fetch_workers,
|
|
415
|
+
profile=args.profile,
|
|
354
416
|
)
|
|
355
417
|
log.debug("DIAG: infer_omero_wsi returned normally")
|
|
356
|
-
except BaseException
|
|
418
|
+
except BaseException:
|
|
357
419
|
# Catch and print Python-level exceptions; note segfaults (native crashes)
|
|
358
420
|
# will not be caught here, but these diagnostics will show progress up to
|
|
359
421
|
# the crash point.
|
|
360
|
-
import sys
|
|
361
422
|
import traceback
|
|
362
423
|
|
|
363
424
|
log.debug("DIAG: infer_omero_wsi raised an exception:")
|
|
@@ -67,7 +67,7 @@ def _ensure_omero() -> None:
|
|
|
67
67
|
if _OMERO_LOADED:
|
|
68
68
|
return
|
|
69
69
|
try:
|
|
70
|
-
import omero
|
|
70
|
+
import omero
|
|
71
71
|
from omero.gateway import BlitzGateway as _BG
|
|
72
72
|
from omero.model import enums as omero_enums
|
|
73
73
|
|
|
@@ -244,7 +244,7 @@ def create_omero_connection(
|
|
|
244
244
|
# If we created a raw omero.client earlier (websocket path) attach it
|
|
245
245
|
# to the BlitzGateway instance so it can be closed explicitly later.
|
|
246
246
|
try:
|
|
247
|
-
|
|
247
|
+
conn._omero_ws_client = client
|
|
248
248
|
except Exception:
|
|
249
249
|
pass
|
|
250
250
|
try:
|
|
@@ -1016,7 +1016,7 @@ def _precompute_tile_info(
|
|
|
1016
1016
|
|
|
1017
1017
|
def _fetch_worker(
|
|
1018
1018
|
worker_id: int,
|
|
1019
|
-
conn:
|
|
1019
|
+
conn: BlitzGateway,
|
|
1020
1020
|
work_queue: queue.Queue[Any],
|
|
1021
1021
|
result_queue: queue.Queue[Any],
|
|
1022
1022
|
stop_event: threading.Event,
|
|
@@ -1175,7 +1175,7 @@ def _parallel_fetch_coordinator(
|
|
|
1175
1175
|
conn = create_omero_connection(host, port, username, password)
|
|
1176
1176
|
conns.append(conn)
|
|
1177
1177
|
logger.debug("Created OMERO connection for worker %d", i)
|
|
1178
|
-
except Exception
|
|
1178
|
+
except Exception:
|
|
1179
1179
|
# Clean up any already-created connections and re-raise.
|
|
1180
1180
|
for c in conns:
|
|
1181
1181
|
try:
|
|
@@ -1292,6 +1292,7 @@ def infer_omero_wsi(
|
|
|
1292
1292
|
min_tissue_frac: float = 0.01,
|
|
1293
1293
|
prefetch_depth: int = 8,
|
|
1294
1294
|
num_fetch_workers: int = 4,
|
|
1295
|
+
profile: str = "best_effort",
|
|
1295
1296
|
) -> np.ndarray:
|
|
1296
1297
|
"""Tile-based OMERO inference pipeline.
|
|
1297
1298
|
|
|
@@ -1633,6 +1634,19 @@ def infer_omero_wsi(
|
|
|
1633
1634
|
)
|
|
1634
1635
|
tiles_tissue = len(tile_schedule)
|
|
1635
1636
|
|
|
1637
|
+
# Upscale tissue mask to full resolution for post-processing.
|
|
1638
|
+
th_h, th_w = tissue_mask.shape
|
|
1639
|
+
y_idx = np.clip(
|
|
1640
|
+
(np.arange(out_h) * th_h / out_h).astype(np.int64),
|
|
1641
|
+
0,
|
|
1642
|
+
th_h - 1,
|
|
1643
|
+
)
|
|
1644
|
+
x_idx = np.clip(
|
|
1645
|
+
(np.arange(out_w) * th_w / out_w).astype(np.int64),
|
|
1646
|
+
0,
|
|
1647
|
+
th_w - 1,
|
|
1648
|
+
)
|
|
1649
|
+
tissue_mask_full = tissue_mask[y_idx[:, None], x_idx[None, :]]
|
|
1636
1650
|
del tissue_mask
|
|
1637
1651
|
gc.collect()
|
|
1638
1652
|
|
|
@@ -1691,6 +1705,8 @@ def infer_omero_wsi(
|
|
|
1691
1705
|
tile_overlap=tile_overlap,
|
|
1692
1706
|
roi_size=inference_tile_size,
|
|
1693
1707
|
sw_overlap=sw_overlap,
|
|
1708
|
+
tissue_mask_full=tissue_mask_full,
|
|
1709
|
+
profile=profile,
|
|
1694
1710
|
num_fetch_workers=num_fetch_workers,
|
|
1695
1711
|
)
|
|
1696
1712
|
|
|
@@ -1739,6 +1755,8 @@ def _run_pipeline(
|
|
|
1739
1755
|
tile_overlap: int,
|
|
1740
1756
|
roi_size: int,
|
|
1741
1757
|
sw_overlap: float,
|
|
1758
|
+
tissue_mask_full: np.ndarray | None = None,
|
|
1759
|
+
profile: str = "best_effort",
|
|
1742
1760
|
num_fetch_workers: int = 4,
|
|
1743
1761
|
) -> np.ndarray:
|
|
1744
1762
|
"""Producer/consumer pipeline: parallel fetch+EHO → inference in main.
|
|
@@ -1820,7 +1838,10 @@ def _run_pipeline(
|
|
|
1820
1838
|
except Exception:
|
|
1821
1839
|
logger.exception("Failed to start OMERO fetch thread; continuing.")
|
|
1822
1840
|
|
|
1823
|
-
|
|
1841
|
+
# Allocate canvases for inner/outer predictions and hematoxylin.
|
|
1842
|
+
inner_pred = np.zeros((out_h, out_w), dtype=bool)
|
|
1843
|
+
outer_pred = np.zeros((out_h, out_w), dtype=bool)
|
|
1844
|
+
hematoxylin_full = np.zeros((out_h, out_w), dtype=np.uint8)
|
|
1824
1845
|
|
|
1825
1846
|
eho_canvas: np.ndarray | None = None
|
|
1826
1847
|
if save_eho:
|
|
@@ -1840,25 +1861,34 @@ def _run_pipeline(
|
|
|
1840
1861
|
if eho_canvas is not None:
|
|
1841
1862
|
eho_canvas[oy0 : oy0 + oh, ox0 : ox0 + ow] = tile_eho
|
|
1842
1863
|
|
|
1843
|
-
|
|
1864
|
+
# Store hematoxylin channel (ch1) for nuclei segmentation.
|
|
1865
|
+
hematoxylin_full[oy0 : oy0 + oh, ox0 : ox0 + ow] = tile_eho[:, :, 1]
|
|
1866
|
+
|
|
1867
|
+
inner_tile, outer_tile = run_inference(
|
|
1844
1868
|
tile_eho,
|
|
1845
1869
|
model,
|
|
1846
1870
|
device,
|
|
1847
1871
|
roi_size=(roi_size, roi_size),
|
|
1848
1872
|
overlap=sw_overlap,
|
|
1873
|
+
return_raw=True,
|
|
1849
1874
|
)
|
|
1850
1875
|
del tile_eho
|
|
1851
1876
|
|
|
1852
1877
|
# Write only the centre-cropped keep region to avoid seams.
|
|
1853
|
-
ph, pw =
|
|
1878
|
+
ph, pw = inner_tile.shape[:2]
|
|
1854
1879
|
ky1 = min(keep_y0 + keep_h, ph)
|
|
1855
1880
|
kx1 = min(keep_x0 + keep_w, pw)
|
|
1856
|
-
kept = tile_pred[keep_y0:ky1, keep_x0:kx1]
|
|
1857
1881
|
out_y0 = oy0 + keep_y0
|
|
1858
1882
|
out_x0 = ox0 + keep_x0
|
|
1859
|
-
ch
|
|
1860
|
-
|
|
1861
|
-
|
|
1883
|
+
ch = ky1 - keep_y0
|
|
1884
|
+
cw = kx1 - keep_x0
|
|
1885
|
+
inner_pred[out_y0 : out_y0 + ch, out_x0 : out_x0 + cw] = inner_tile[
|
|
1886
|
+
keep_y0:ky1, keep_x0:kx1
|
|
1887
|
+
]
|
|
1888
|
+
outer_pred[out_y0 : out_y0 + ch, out_x0 : out_x0 + cw] = outer_tile[
|
|
1889
|
+
keep_y0:ky1, keep_x0:kx1
|
|
1890
|
+
]
|
|
1891
|
+
del inner_tile, outer_tile
|
|
1862
1892
|
|
|
1863
1893
|
processed += 1
|
|
1864
1894
|
if processed % 25 == 0:
|
|
@@ -1876,7 +1906,24 @@ def _run_pipeline(
|
|
|
1876
1906
|
if torch.cuda.is_available():
|
|
1877
1907
|
torch.cuda.empty_cache()
|
|
1878
1908
|
|
|
1879
|
-
|
|
1909
|
+
# ── Post-processing ──────────────────────────────────────────
|
|
1910
|
+
from shell.post_process import post_process
|
|
1911
|
+
|
|
1912
|
+
logger.info("Running post-processing (profile=%s) …", profile)
|
|
1913
|
+
label_map = post_process(
|
|
1914
|
+
inner_pred,
|
|
1915
|
+
outer_pred,
|
|
1916
|
+
tissue_mask_full
|
|
1917
|
+
if tissue_mask_full is not None
|
|
1918
|
+
else np.ones((out_h, out_w), dtype=bool),
|
|
1919
|
+
hematoxylin_full,
|
|
1920
|
+
profile_name=profile,
|
|
1921
|
+
verbose=True,
|
|
1922
|
+
)
|
|
1923
|
+
del inner_pred, outer_pred, hematoxylin_full, tissue_mask_full
|
|
1924
|
+
gc.collect()
|
|
1925
|
+
|
|
1926
|
+
return _save_results(label_map, eho_canvas, save_eho, output_path)
|
|
1880
1927
|
|
|
1881
1928
|
finally:
|
|
1882
1929
|
fetch_thread.join(timeout=30)
|