lavlab-shell 0.2.0__tar.gz → 0.3.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 (42) hide show
  1. lavlab_shell-0.3.0/.github/workflows/build.yml +54 -0
  2. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/.github/workflows/publish.yml +11 -8
  3. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/.github/workflows/pytest.yml +11 -3
  4. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/Dockerfile +4 -2
  5. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/PKG-INFO +2 -1
  6. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/pyproject.toml +11 -12
  7. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/src/shell/__about__.py +1 -1
  8. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/src/shell/cli.py +0 -9
  9. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/src/shell/infer_omero_wsi.py +0 -10
  10. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/src/shell/infer_wsi.py +116 -29
  11. lavlab_shell-0.3.0/src/shell/inference.py +193 -0
  12. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/src/shell/model.py +3 -3
  13. lavlab_shell-0.3.0/src/shell/preprocessing.py +21 -0
  14. lavlab_shell-0.3.0/src/shell/transforms.py +648 -0
  15. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/src/shell/weights/model_v1.pth +0 -0
  16. lavlab_shell-0.2.0/.github/workflows/build.yml +0 -22
  17. lavlab_shell-0.2.0/src/shell/inference.py +0 -126
  18. lavlab_shell-0.2.0/src/shell/preprocessing.py +0 -337
  19. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/.devcontainer/Dockerfile +0 -0
  20. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/.editorconfig +0 -0
  21. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/.github/dependabot.yml +0 -0
  22. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/.github/workflows/lint.yml +0 -0
  23. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/.gitignore +0 -0
  24. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/.pre-commit-config.yaml +0 -0
  25. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/CONTRIBUTING.md +0 -0
  26. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/LICENSE.txt +0 -0
  27. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/Makefile +0 -0
  28. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/README.md +0 -0
  29. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/docs/api.md +0 -0
  30. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/docs/index.md +0 -0
  31. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/mkdocs.yml +0 -0
  32. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/requirements/requirements-docs.txt +0 -0
  33. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/requirements/requirements-lint.txt +0 -0
  34. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/requirements/requirements-test.txt +0 -0
  35. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/requirements/requirements-types.txt +0 -0
  36. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/requirements.txt +0 -0
  37. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/src/shell/__init__.py +0 -0
  38. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/src/shell/benchmark.py +0 -0
  39. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/src/shell/py.typed +0 -0
  40. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/tests/__init__.py +0 -0
  41. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/tests/conftest.py +0 -0
  42. {lavlab_shell-0.2.0 → lavlab_shell-0.3.0}/tests/test_shell.py +0 -0
@@ -0,0 +1,54 @@
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
32
+
33
+ - uses: actions/checkout@v4
34
+
35
+ - name: Build Docker image (hatch stage)
36
+ uses: docker/build-push-action@v4
37
+ with:
38
+ context: .
39
+ file: ./Dockerfile
40
+ target: hatch
41
+ tags: shell-ci:latest
42
+
43
+ - name: Run build inside Docker container
44
+ run: |
45
+ docker run --rm \
46
+ -v "${{ github.workspace }}:/app" \
47
+ -w /app \
48
+ shell-ci:latest build
49
+
50
+ - name: Upload wheel
51
+ uses: actions/upload-artifact@v4
52
+ with:
53
+ name: wheel
54
+ path: dist/*.whl
@@ -17,17 +17,20 @@ jobs:
17
17
  - name: Checkout
18
18
  uses: actions/checkout@v4
19
19
 
20
- - name: Set up Python 3.12
21
- uses: actions/setup-python@v5
20
+ - name: Build Docker image (hatch stage)
21
+ uses: docker/build-push-action@v4
22
22
  with:
23
- python-version: "3.12"
23
+ context: .
24
+ file: ./Dockerfile
25
+ target: hatch
26
+ tags: shell-ci:latest
24
27
 
25
- - name: Install build tools
28
+ - name: Build wheel and sdist inside Docker
26
29
  run: |
27
- python -m pip install --upgrade pip setuptools wheel hatch
28
-
29
- - name: Build wheel and sdist
30
- run: hatch build
30
+ docker run --rm \
31
+ -v "${{ github.workspace }}:/app" \
32
+ -w /app \
33
+ shell-ci:latest build
31
34
 
32
35
  - name: Publish to PyPI
33
36
  uses: pypa/gh-action-pypi-publish@release/v1
@@ -10,11 +10,19 @@ jobs:
10
10
  steps:
11
11
  - uses: actions/checkout@v4
12
12
 
13
- - name: Build Docker image
14
- run: docker build --target hatch -t shell:hatch .
13
+ - name: Set up Python
14
+ uses: actions/setup-python@v4
15
+ with:
16
+ python-version: '3.12'
17
+
18
+ - name: Install test dependencies
19
+ run: |
20
+ python -m pip install --upgrade pip
21
+ python -m pip install pytest pytest-cov coverage[toml]>=6.2
15
22
 
16
23
  - name: Run pytest and generate coverage report
17
- run: docker run --rm -e HATCH_ENV=test -v "${{ github.workspace }}:/app" shell:hatch cov
24
+ run: |
25
+ pytest --cov=src --cov-report=xml:coverage.xml
18
26
 
19
27
  - name: Upload coverage artifact
20
28
  uses: actions/upload-artifact@v4
@@ -22,7 +22,9 @@ COPY src/ src/
22
22
  RUN chown -R vscode:vscode /app
23
23
 
24
24
  FROM base AS hatch
25
- RUN pip3 install --no-cache-dir hatch uv
25
+ RUN pip3 install --no-cache-dir --upgrade pip setuptools wheel
26
+ RUN pip3 install --no-cache-dir 'virtualenv<21'
27
+ RUN pip3 install --no-cache-dir --upgrade hatch hatch-uv
26
28
  ENV HATCH_ENV=default
27
29
  ENTRYPOINT ["hatch", "run"]
28
30
 
@@ -32,7 +34,7 @@ COPY requirements.txt ./
32
34
  COPY tests/ tests/
33
35
  COPY docs/ docs/
34
36
  COPY mkdocs.yml ./
35
- RUN pip3 install --no-cache-dir hatch \
37
+ RUN pip3 install --no-cache-dir hatch hatch-uv \
36
38
  && hatch build \
37
39
  && pip3 install --no-cache-dir $(find /app -name 'requirement*.txt' -exec echo -n '-r {} ' \;)
38
40
  USER vscode
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: lavlab-shell
3
- Version: 0.2.0
3
+ Version: 0.3.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
@@ -21,6 +21,7 @@ Requires-Dist: monai
21
21
  Requires-Dist: numpy
22
22
  Requires-Dist: openslide-python
23
23
  Requires-Dist: pyvips
24
+ Requires-Dist: scikit-image
24
25
  Requires-Dist: scipy
25
26
  Requires-Dist: torch
26
27
  Provides-Extra: omero
@@ -33,6 +33,7 @@ dependencies = [
33
33
  "pyvips",
34
34
  "openslide-python",
35
35
  "macenko-pca",
36
+ "scikit-image",
36
37
  ]
37
38
 
38
39
  [project.optional-dependencies]
@@ -53,27 +54,25 @@ path = "src/shell/__about__.py"
53
54
  [tool.hatch.build.targets.wheel]
54
55
  packages = ["src/shell"]
55
56
 
56
- [tool.hatch.env]
57
- requires = ["hatch-pip-compile"]
58
-
59
57
  [tool.hatch.envs.default]
60
- type = "pip-compile"
61
- pip-compile-resolver = "uv"
62
- pip-compile-hashes = false
63
- pip-compile-args = ["--allow-unsafe", "--universal"]
64
- dependencies = ["setuptools==82.0.0"]
58
+ dependencies = ["setuptools>=82.0.0"]
65
59
 
66
60
  [tool.hatch.envs.default.scripts]
67
61
  build = "hatch build && chmod -R 777 dist/*"
68
62
 
69
63
  [tool.hatch.envs.test]
70
- dependencies = ["pytest", "pytest-cov", "coverage[toml]>=6.2", "setuptools==82.0.0"]
64
+ dependencies = [
65
+ "pytest",
66
+ "pytest-cov",
67
+ "coverage[toml]>=6.2",
68
+ "setuptools>=82.0.0",
69
+ ]
71
70
  [tool.hatch.envs.test.scripts]
72
71
  test = "pytest {args:tests}"
73
72
  cov = "pytest --cov=src --cov-report=term-missing --cov-report=xml {args:tests}"
74
73
 
75
74
  [tool.hatch.envs.lint]
76
- dependencies = ["ruff>=0.4.0", "setuptools==82.0.0"]
75
+ dependencies = ["ruff>=0.4.0", "setuptools>=82.0.0"]
77
76
  [tool.hatch.envs.lint.scripts]
78
77
  check = "ruff check src tests"
79
78
  format = "ruff format src tests"
@@ -82,7 +81,7 @@ fix = "ruff check --fix src tests"
82
81
  all = ["format", "fix", "check"]
83
82
 
84
83
  [tool.hatch.envs.types]
85
- dependencies = ["mypy>=1.0.0", "setuptools==82.0.0"]
84
+ dependencies = ["mypy>=1.0.0", "setuptools>=82.0.0"]
86
85
  [tool.hatch.envs.types.scripts]
87
86
  check = "mypy --install-types --non-interactive {args:src/shell tests}"
88
87
 
@@ -95,7 +94,7 @@ dependencies = [
95
94
  "mkdocs-minify-plugin",
96
95
  "mkdocs-material-extensions",
97
96
  "mkdocs-git-revision-date-localized-plugin",
98
- "setuptools==82.0.0",
97
+ "setuptools>=82.0.0",
99
98
  ]
100
99
  [tool.hatch.envs.docs.scripts]
101
100
  build-docs = "mkdocs build"
@@ -3,4 +3,4 @@
3
3
  # SPDX-License-Identifier: MIT
4
4
  """Version information for shell."""
5
5
 
6
- __version__ = "0.2.0"
6
+ __version__ = "0.3.0"
@@ -98,12 +98,6 @@ def build_parser() -> argparse.ArgumentParser:
98
98
  default=None,
99
99
  help="Optional: also save the intermediate EHO image.",
100
100
  )
101
- infer_p.add_argument(
102
- "--stain-downsample",
103
- type=int,
104
- default=4,
105
- help="Downsample factor for stain parameter estimation.",
106
- )
107
101
  infer_p.add_argument(
108
102
  "--device",
109
103
  type=str,
@@ -138,7 +132,6 @@ def build_parser() -> argparse.ArgumentParser:
138
132
  default=1.0,
139
133
  help="Desired output resolution (um/pixel)",
140
134
  )
141
- omero_p.add_argument("--max-dim", type=int, default=None)
142
135
  omero_p.add_argument(
143
136
  "--model-path",
144
137
  type=str,
@@ -283,7 +276,6 @@ def main(argv: list[str] | None = None) -> int:
283
276
  target_mpp=args.target_mpp,
284
277
  mpp=args.mpp,
285
278
  save_eho=args.save_eho,
286
- stain_downsample=args.stain_downsample,
287
279
  device=args.device,
288
280
  )
289
281
 
@@ -311,7 +303,6 @@ def main(argv: list[str] | None = None) -> int:
311
303
  model_version=args.model_version,
312
304
  output_path=args.output,
313
305
  target_mpp=args.target_mpp,
314
- max_dim=args.max_dim,
315
306
  group_id=args.group_id,
316
307
  save_eho=args.save_eho,
317
308
  no_tissue_crop=args.no_tissue_crop,
@@ -1095,7 +1095,6 @@ def infer_omero_wsi(
1095
1095
  model_path: str | None = None,
1096
1096
  output_path: str,
1097
1097
  target_mpp: float = 1.0,
1098
- max_dim: int | None = None,
1099
1098
  group_id: int | None = None,
1100
1099
  save_eho: str | None = None,
1101
1100
  no_tissue_crop: bool = False,
@@ -1146,8 +1145,6 @@ def infer_omero_wsi(
1146
1145
  Where to save the prediction image (format from extension).
1147
1146
  target_mpp : float
1148
1147
  Desired microns-per-pixel.
1149
- max_dim : int or None
1150
- Clamp the longest output dimension.
1151
1148
  group_id : int or None
1152
1149
  OMERO group to switch to.
1153
1150
  save_eho : str or None
@@ -1267,13 +1264,6 @@ def infer_omero_wsi(
1267
1264
  out_w = max(1, round(best_lsz_x * sx))
1268
1265
  out_h = max(1, round(best_lsz_y * sy))
1269
1266
 
1270
- if max_dim is not None and (out_w > max_dim or out_h > max_dim):
1271
- clamp = min(max_dim / out_w, max_dim / out_h)
1272
- out_w = max(1, round(out_w * clamp))
1273
- out_h = max(1, round(out_h * clamp))
1274
- sx *= clamp
1275
- sy *= clamp
1276
-
1277
1267
  # Native tile size at best level
1278
1268
  probe2 = conn.c.sf.createRawPixelsStore() # type: ignore[union-attr]
1279
1269
  try:
@@ -23,9 +23,9 @@ import gc
23
23
  import logging
24
24
  import os
25
25
  import warnings
26
+ from types import ModuleType
26
27
 
27
28
  import numpy as np
28
- import openslide
29
29
 
30
30
  # shell.inference and shell.model load torch at their module level.
31
31
  # pyvips must come *after* them so PyTorch initialises its thread-pool
@@ -34,17 +34,16 @@ import openslide
34
34
  # own the same OpenMP/GCD thread infrastructure.
35
35
  from shell.inference import run_inference
36
36
  from shell.model import build_model
37
- from shell.preprocessing import (
38
- apply_eho_chunked,
39
- detect_background,
40
- estimate_stain_params,
41
- )
37
+ from shell.transforms import EHOd, TissueMaskd
42
38
 
43
39
  # pyvips intentionally after torch-loading shell imports above (macOS safety)
44
40
  import pyvips # isort: skip
45
41
 
46
42
  log = logging.getLogger(__name__)
47
43
 
44
+ _OPENSLIDE_MODULE: ModuleType | None = None
45
+ _OPENSLIDE_IMPORT_FAILED: bool = False
46
+
48
47
  # ---------------------------------------------------------------------------
49
48
  # Default parameters
50
49
  # ---------------------------------------------------------------------------
@@ -54,12 +53,34 @@ TARGET_MPP: float = 1.0
54
53
  # ---------------------------------------------------------------------------
55
54
  # Helpers
56
55
  # ---------------------------------------------------------------------------
56
+ def _get_openslide() -> ModuleType | None:
57
+ """Import and cache the ``openslide`` module lazily.
58
+
59
+ Returns ``None`` when openslide is unavailable.
60
+ """
61
+ global _OPENSLIDE_MODULE, _OPENSLIDE_IMPORT_FAILED
62
+ if _OPENSLIDE_IMPORT_FAILED:
63
+ return None
64
+ if _OPENSLIDE_MODULE is None:
65
+ try:
66
+ import openslide
67
+ except Exception:
68
+ _OPENSLIDE_IMPORT_FAILED = True
69
+ return None
70
+ _OPENSLIDE_MODULE = openslide
71
+ return _OPENSLIDE_MODULE
72
+
73
+
57
74
  def _read_mpp_from_openslide(image_path: str) -> tuple[float, float] | None:
58
75
  """Try to extract um/px from OpenSlide metadata.
59
76
 
60
77
  Returns ``(mpp_x, mpp_y)`` or ``None`` if the format is unsupported
61
78
  or the metadata is missing.
62
79
  """
80
+ openslide = _get_openslide()
81
+ if openslide is None:
82
+ return None
83
+
63
84
  try:
64
85
  slide = openslide.OpenSlide(image_path)
65
86
  except (
@@ -98,6 +119,15 @@ def _load_image(
98
119
  log.info("pyvips could not open %s; falling back to OpenSlide.", image_path)
99
120
 
100
121
  # --- attempt 2: openslide ---
122
+ openslide = _get_openslide()
123
+ if openslide is None:
124
+ msg = (
125
+ f"pyvips could not open '{image_path}', and OpenSlide is not available. "
126
+ "Install openslide-python (and OpenSlide runtime) or use a format "
127
+ "supported by pyvips."
128
+ )
129
+ raise ValueError(msg)
130
+
101
131
  try:
102
132
  slide = openslide.OpenSlide(image_path)
103
133
  dims = slide.dimensions # (width, height)
@@ -129,6 +159,47 @@ def _vips_to_rgb_numpy(vips_img: pyvips.Image) -> np.ndarray:
129
159
  return vips_img.numpy()
130
160
 
131
161
 
162
+ def _read_image_size(image_path: str) -> tuple[int, int]:
163
+ """Return ``(height, width)`` for *image_path*."""
164
+ try:
165
+ vips_img = pyvips.Image.new_from_file(image_path, access="sequential")
166
+ return int(vips_img.height), int(vips_img.width)
167
+ except pyvips.Error:
168
+ pass
169
+
170
+ openslide = _get_openslide()
171
+ if openslide is None:
172
+ msg = f"Could not determine image size for '{image_path}'."
173
+ raise ValueError(msg)
174
+
175
+ try:
176
+ slide = openslide.OpenSlide(image_path)
177
+ width, height = slide.dimensions
178
+ slide.close()
179
+ return int(height), int(width)
180
+ except (
181
+ openslide.OpenSlideUnsupportedFormatError,
182
+ openslide.OpenSlideError,
183
+ ) as exc:
184
+ msg = f"Could not determine image size for '{image_path}'."
185
+ raise ValueError(msg) from exc
186
+
187
+
188
+ def _resize_label_map_nearest(
189
+ label_map: np.ndarray,
190
+ out_h: int,
191
+ out_w: int,
192
+ ) -> np.ndarray:
193
+ """Resize a label map to ``(out_h, out_w)`` using nearest-neighbour."""
194
+ in_h, in_w = label_map.shape[:2]
195
+ if in_h == out_h and in_w == out_w:
196
+ return label_map
197
+
198
+ y_idx = np.clip((np.arange(out_h) * in_h / out_h).astype(np.int64), 0, in_h - 1)
199
+ x_idx = np.clip((np.arange(out_w) * in_w / out_w).astype(np.int64), 0, in_w - 1)
200
+ return label_map[y_idx[:, None], x_idx[None, :]]
201
+
202
+
132
203
  # ---------------------------------------------------------------------------
133
204
  # Public API
134
205
  # ---------------------------------------------------------------------------
@@ -137,19 +208,23 @@ def preprocess_wsi(
137
208
  *,
138
209
  target_mpp: float = TARGET_MPP,
139
210
  mpp: float | None = None,
140
- stain_downsample: int = 4,
141
- ) -> np.ndarray:
211
+ ) -> tuple[np.ndarray, np.ndarray]:
142
212
  """Read a raw RGB image, scale to *target_mpp*, and produce an EHO image.
143
213
 
214
+ Uses the MONAI transform pipeline (``TissueMaskd`` → ``EHOd``) from
215
+ :mod:`shell.transforms`.
216
+
144
217
  :param image_path: path to an RGB image (TIFF, PNG, JPEG, etc.).
145
218
  :param target_mpp: desired microns-per-pixel.
146
219
  :param mpp: manual override for the source image um/px. When
147
220
  ``None`` the value is read from slide metadata; if metadata is
148
221
  unavailable (e.g. plain PNG) a warning is emitted and scaling is
149
222
  skipped (the image is assumed to already be at *target_mpp*).
150
- :param stain_downsample: downsample factor for stain parameter estimation.
151
- :return: (H, W, 3) uint8 EHO image.
223
+ :return: tuple of (H, W, 3) uint8 EHO image and (H, W) bool tissue mask.
152
224
  """
225
+ from monai.data import MetaTensor
226
+ from monai.transforms import Compose
227
+
153
228
  # 1. Determine MPP
154
229
  if mpp is not None:
155
230
  mpp_x = mpp_y = float(mpp)
@@ -194,20 +269,26 @@ def preprocess_wsi(
194
269
  image_np = _vips_to_rgb_numpy(vips_tmp)
195
270
  del vips_tmp
196
271
 
197
- # 4. Estimate stain parameters on a down-sampled copy
198
- ds = stain_downsample
199
- rgb_small = image_np[::ds, ::ds]
200
- bg_small = detect_background(rgb_small)
201
- sp = estimate_stain_params(rgb_small, bg_mask=bg_small)
202
- del rgb_small, bg_small
272
+ # 4. Run MONAI transform pipeline: TissueMask → EHO
273
+ pipeline = Compose(
274
+ [
275
+ TissueMaskd(keys=["image"]),
276
+ EHOd(
277
+ keys=["image"],
278
+ tissue_mask_keys=["image_tissue_mask"],
279
+ ),
280
+ ]
281
+ )
282
+ data = pipeline({"image": MetaTensor(image_np)})
283
+ del image_np
203
284
  gc.collect()
204
285
 
205
- # 5. Apply EHO colour transform (chunked)
206
- eho = apply_eho_chunked(image_np, **sp)
207
- del image_np, sp
208
- gc.collect()
286
+ # EHOd outputs (3, H, W) MetaTensor — convert back to (H, W, 3) uint8
287
+ eho = data["image"].numpy().transpose(1, 2, 0).astype(np.uint8)
288
+ tissue_mask = data["image_tissue_mask"].numpy().squeeze() > 0
289
+ del data
209
290
 
210
- return eho
291
+ return eho, tissue_mask
211
292
 
212
293
 
213
294
  def infer_wsi(
@@ -219,7 +300,6 @@ def infer_wsi(
219
300
  target_mpp: float = TARGET_MPP,
220
301
  mpp: float | None = None,
221
302
  save_eho: str | None = None,
222
- stain_downsample: int = 4,
223
303
  device: str = "auto",
224
304
  ) -> np.ndarray:
225
305
  """Full pipeline: preprocess -> model -> label image.
@@ -234,9 +314,8 @@ def infer_wsi(
234
314
  :param mpp: manual source um/px override. See
235
315
  :func:`preprocess_wsi` for details.
236
316
  :param save_eho: optional path to save the intermediate EHO image.
237
- :param stain_downsample: downsample factor for stain estimation.
238
317
  :param device: ``"auto"``, ``"cpu"``, or ``"cuda"``.
239
- :return: (H, W) uint8 label map.
318
+ :return: (H, W) uint8 label map at the original input resolution.
240
319
  """
241
320
  import torch
242
321
 
@@ -249,11 +328,10 @@ def infer_wsi(
249
328
  device = "cpu"
250
329
 
251
330
  # 1. Preprocess
252
- eho = preprocess_wsi(
331
+ eho, tissue_mask = preprocess_wsi(
253
332
  input_path,
254
333
  target_mpp=target_mpp,
255
334
  mpp=mpp,
256
- stain_downsample=stain_downsample,
257
335
  )
258
336
 
259
337
  if save_eho:
@@ -262,11 +340,20 @@ def infer_wsi(
262
340
 
263
341
  # 2. Load model + inference
264
342
  model = build_model(model_path, device, model_version=model_version)
265
- label_map = run_inference(eho, model, device)
266
- del eho, model
343
+ label_map = run_inference(
344
+ eho,
345
+ model,
346
+ device,
347
+ tissue_mask=tissue_mask,
348
+ )
349
+ del eho, model, tissue_mask
267
350
  gc.collect()
268
351
 
269
- # 3. Save
352
+ # 3. Always return/save at original input resolution.
353
+ input_h, input_w = _read_image_size(input_path)
354
+ label_map = _resize_label_map_nearest(label_map, input_h, input_w)
355
+
356
+ # 4. Save
270
357
  os.makedirs(os.path.dirname(output_path) or ".", exist_ok=True)
271
358
  pyvips.Image.new_from_array(label_map).write_to_file(output_path)
272
359