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.
Files changed (45) hide show
  1. lavlab_shell-0.5.0/.github/workflows/build.yml +31 -0
  2. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.github/workflows/publish.yml +1 -0
  3. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.github/workflows/pytest.yml +6 -6
  4. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/PKG-INFO +4 -2
  5. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/pyproject.toml +24 -3
  6. lavlab_shell-0.5.0/scripts/export_onnx.py +90 -0
  7. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/__about__.py +1 -1
  8. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/__init__.py +8 -0
  9. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/cli.py +65 -4
  10. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/infer_omero_wsi.py +59 -12
  11. lavlab_shell-0.5.0/src/shell/infer_wsi.py +816 -0
  12. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/inference.py +44 -34
  13. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/model.py +102 -41
  14. lavlab_shell-0.5.0/src/shell/post_process.py +1360 -0
  15. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/transforms.py +7 -3
  16. lavlab_shell-0.5.0/src/shell/weights/model_v1.onnx +0 -0
  17. lavlab_shell-0.5.0/src/shell/weights/model_v1.onnx.data +0 -0
  18. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/weights/model_v1.pth +0 -0
  19. lavlab_shell-0.5.0/tests/test_shell.py +68 -0
  20. lavlab_shell-0.3.2/.github/workflows/build.yml +0 -54
  21. lavlab_shell-0.3.2/src/shell/infer_wsi.py +0 -373
  22. lavlab_shell-0.3.2/tests/test_shell.py +0 -16
  23. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.devcontainer/Dockerfile +0 -0
  24. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.editorconfig +0 -0
  25. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.github/dependabot.yml +0 -0
  26. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.github/workflows/lint.yml +0 -0
  27. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.gitignore +0 -0
  28. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/.pre-commit-config.yaml +0 -0
  29. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/CONTRIBUTING.md +0 -0
  30. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/Dockerfile +0 -0
  31. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/LICENSE.txt +0 -0
  32. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/Makefile +0 -0
  33. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/README.md +0 -0
  34. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/docs/api.md +0 -0
  35. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/docs/index.md +0 -0
  36. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/mkdocs.yml +0 -0
  37. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/requirements/requirements-docs.txt +0 -0
  38. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/requirements/requirements-lint.txt +0 -0
  39. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/requirements/requirements-test.txt +0 -0
  40. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/requirements/requirements-types.txt +0 -0
  41. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/requirements.txt +0 -0
  42. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/benchmark.py +0 -0
  43. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/src/shell/py.typed +0 -0
  44. {lavlab_shell-0.3.2 → lavlab_shell-0.5.0}/tests/__init__.py +0 -0
  45. {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
@@ -28,6 +28,7 @@ jobs:
28
28
  - name: Build wheel and sdist inside Docker
29
29
  run: |
30
30
  docker run --rm \
31
+ --entrypoint hatch \
31
32
  -v "${{ github.workspace }}:/app" \
32
33
  -w /app \
33
34
  shell-ci:latest build
@@ -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 pytest pytest-cov coverage[toml]>=6.2
23
+ python -m pip install hatch
22
24
 
23
25
  - name: Run pytest and generate coverage report
24
26
  run: |
25
- pytest --cov=src --cov-report=xml:coverage.xml
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: ${{ secrets.CODECOV_TOKEN != '' }}
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: ${{ secrets.CODECOV_TOKEN == '' }}
43
+ if: env.CODECOV_TOKEN == ''
44
44
  run: echo "Skipping Codecov — CODECOV_TOKEN not set."
@@ -1,6 +1,6 @@
1
- Metadata-Version: 2.4
1
+ Metadata-Version: 2.5
2
2
  Name: lavlab-shell
3
- Version: 0.3.2
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()
@@ -3,4 +3,4 @@
3
3
  # SPDX-License-Identifier: MIT
4
4
  """Version information for shell."""
5
5
 
6
- __version__ = "0.3.2"
6
+ __version__ = "0.5.0"
@@ -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=1.0,
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=1.0,
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 as e:
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 # noqa: F401 (import to ensure runtime available)
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
- setattr(conn, "_omero_ws_client", client)
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: "BlitzGateway",
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 as exc:
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
- pred = np.zeros((out_h, out_w), dtype=np.uint8)
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
- tile_pred = run_inference(
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 = tile_pred.shape[:2]
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, cw = kept.shape[:2]
1860
- pred[out_y0 : out_y0 + ch, out_x0 : out_x0 + cw] = kept
1861
- del tile_pred, kept
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
- return _save_results(pred, eho_canvas, save_eho, output_path)
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)