slide2vec 5.1.0__tar.gz → 5.2.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.
- {slide2vec-5.1.0 → slide2vec-5.2.0}/PKG-INFO +1 -1
- {slide2vec-5.1.0 → slide2vec-5.2.0}/pyproject.toml +2 -2
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/__init__.py +1 -1
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/api.py +11 -1
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/artifacts.py +22 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/configs/default.yaml +1 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/__init__.py +4 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/conch.py +2 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/gigapath.py +1 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/hibou.py +2 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/hoptimus.py +3 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/lunit.py +1 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/midnight.py +1 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/musk.py +1 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/phikon.py +2 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/prost40m.py +1 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/uni.py +2 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/virchow.py +2 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/registry.py +43 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/inference.py +16 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/dense_regions.py +35 -6
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/embedding.py +12 -6
- slide2vec-5.2.0/slide2vec/runtime/model_settings.py +97 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/patient_pipeline.py +9 -2
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/serialization.py +3 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec.egg-info/PKG-INFO +1 -1
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec.egg-info/SOURCES.txt +1 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_dense_regions.py +43 -0
- slide2vec-5.2.0/tests/test_patch_size_metadata.py +191 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_regression_core.py +164 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_regression_inference.py +5 -0
- slide2vec-5.1.0/slide2vec/runtime/model_settings.py +0 -48
- {slide2vec-5.1.0 → slide2vec-5.2.0}/LICENSE +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/README.md +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/setup.cfg +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/__main__.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/cli.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/configs/__init__.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/configs/resources.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/data/__init__.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/data/dataset.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/data/tile_reader.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/data/tile_store.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/distributed/__init__.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/distributed/direct_embed_worker.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/distributed/pipeline_worker.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/base.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/__init__.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/moozy/__init__.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/moozy/blocks.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/moozy/case.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/moozy/loading.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/moozy/slide.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/moozy/types.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/prism.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/models/titan.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/encoders/validation.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/progress.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/__init__.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/artifacts_collect.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/batching.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/cpu_budget.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/dense_sliding.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/distributed.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/distributed_stage.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/embedding_persist.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/embedding_pipeline.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/hierarchical.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/manifest.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/persist_callbacks.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/persistence.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/process_list.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/progress_bridge.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/registry.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/slide_encode.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/tiling.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/tiling_pipeline.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/types.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/runtime/worker_io.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/utils/__init__.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/utils/config.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/utils/coordinates.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/utils/log_utils.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/utils/tiling_io.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec/utils/utils.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec.egg-info/dependency_links.txt +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec.egg-info/entry_points.txt +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec.egg-info/not-zip-safe +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec.egg-info/requires.txt +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/slide2vec.egg-info/top_level.txt +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_architecture_runtime_split.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_attention_extraction.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_dense_extraction.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_dense_sliding.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_encoder_registry.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_hs2p_package_cutover.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_output_consistency.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_progress.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_regression_models.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_runtime_batching.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_tile_store.py +0 -0
- {slide2vec-5.1.0 → slide2vec-5.2.0}/tests/test_tiling_pipeline.py +0 -0
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "slide2vec"
|
|
7
|
-
version = "5.
|
|
7
|
+
version = "5.2.0"
|
|
8
8
|
description = "Embedding of whole slide images with Foundation Models"
|
|
9
9
|
readme = "README.md"
|
|
10
10
|
requires-python = ">=3.10"
|
|
@@ -167,7 +167,7 @@ no_implicit_reexport = true
|
|
|
167
167
|
max-line-length = 160
|
|
168
168
|
|
|
169
169
|
[tool.bumpver]
|
|
170
|
-
current_version = "5.
|
|
170
|
+
current_version = "5.2.0"
|
|
171
171
|
version_pattern = "MAJOR.MINOR.PATCH"
|
|
172
172
|
commit = false # We do version bumping in CI, not as a commit
|
|
173
173
|
tag = false # Git tag already exists — we don't auto-tag
|
|
@@ -21,7 +21,11 @@ from slide2vec.encoders.registry import (
|
|
|
21
21
|
resolve_preprocessing_defaults,
|
|
22
22
|
)
|
|
23
23
|
from slide2vec.encoders.validation import validate_encoder_config
|
|
24
|
-
from slide2vec.runtime.model_settings import
|
|
24
|
+
from slide2vec.runtime.model_settings import (
|
|
25
|
+
canonicalize_model_name,
|
|
26
|
+
normalize_output_dtype,
|
|
27
|
+
normalize_precision_name,
|
|
28
|
+
)
|
|
25
29
|
from slide2vec.progress import emit_progress
|
|
26
30
|
from slide2vec.runtime.types import LoadedModel
|
|
27
31
|
from slide2vec.utils.utils import cpu_worker_limit, slurm_cpu_limit
|
|
@@ -226,6 +230,10 @@ class ExecutionOptions:
|
|
|
226
230
|
#: Forward-pass dtype — ``"fp16"``, ``"bf16"``, ``"fp32"``,
|
|
227
231
|
#: or ``None`` (auto-determined from the model preset).
|
|
228
232
|
precision: str | None = None
|
|
233
|
+
#: On-disk feature dtype — ``"fp16"``, ``"fp32"``, or ``None`` to follow
|
|
234
|
+
#: :attr:`precision` (fp16 → fp16, else fp32). Applies to tile, slide, hierarchical,
|
|
235
|
+
#: and patient artifacts; ``"bf16"`` is rejected (numpy has no bfloat16).
|
|
236
|
+
output_dtype: str | None = None
|
|
229
237
|
#: DataLoader prefetch queue depth per worker (default ``4``).
|
|
230
238
|
prefetch_factor: int = 4
|
|
231
239
|
#: Persist tile embeddings to disk when running a slide-level model.
|
|
@@ -253,6 +261,7 @@ class ExecutionOptions:
|
|
|
253
261
|
),
|
|
254
262
|
num_gpus=1 if run_on_cpu else (int(configured_num_gpus) if configured_num_gpus is not None else None),
|
|
255
263
|
precision="fp32" if run_on_cpu else requested_precision,
|
|
264
|
+
output_dtype=getattr(cfg.speed, "output_dtype", None),
|
|
256
265
|
prefetch_factor=prefetch_factor,
|
|
257
266
|
save_tile_embeddings=bool(cfg.model.save_tile_embeddings),
|
|
258
267
|
save_slide_embeddings=bool(cfg.model.save_slide_embeddings),
|
|
@@ -263,6 +272,7 @@ class ExecutionOptions:
|
|
|
263
272
|
resolved_num_gpus = _default_num_gpus() if self.num_gpus is None else self.num_gpus
|
|
264
273
|
object.__setattr__(self, "num_gpus", resolved_num_gpus)
|
|
265
274
|
object.__setattr__(self, "precision", normalize_precision_name(self.precision))
|
|
275
|
+
object.__setattr__(self, "output_dtype", normalize_output_dtype(self.output_dtype))
|
|
266
276
|
if resolved_num_gpus < 1:
|
|
267
277
|
raise ValueError("ExecutionOptions.num_gpus must be at least 1")
|
|
268
278
|
if self.prefetch_factor < 1:
|
|
@@ -7,6 +7,8 @@ import numpy as np
|
|
|
7
7
|
import torch
|
|
8
8
|
from hs2p.fileops import is_flattened_annotation
|
|
9
9
|
|
|
10
|
+
from slide2vec.runtime.model_settings import output_torch_dtype
|
|
11
|
+
|
|
10
12
|
|
|
11
13
|
@dataclass(frozen=True, kw_only=True)
|
|
12
14
|
class TileEmbeddingArtifact:
|
|
@@ -74,6 +76,26 @@ def _validate_output_format(output_format: str) -> str:
|
|
|
74
76
|
return normalized
|
|
75
77
|
|
|
76
78
|
|
|
79
|
+
_OUTPUT_NUMPY_DTYPE = {"fp16": np.float16, "fp32": np.float32}
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def cast_feature_dtype(data: Any, precision: str) -> Any:
|
|
83
|
+
"""Cast features to the on-disk ``precision`` (``"fp16"`` / ``"fp32"``), keeping their kind.
|
|
84
|
+
|
|
85
|
+
Torch tensors are cast via ``.to`` and arrays via ``astype``; ``None`` (no features)
|
|
86
|
+
passes through. This is what makes the pooled tile/slide/hierarchical/patient artifacts
|
|
87
|
+
land in a deterministic dtype, mirroring the dense path's ``output_dtype``. The precision
|
|
88
|
+
is resolved upstream by :func:`slide2vec.runtime.model_settings.resolve_output_precision`,
|
|
89
|
+
so only ``"fp16"`` / ``"fp32"`` reach here.
|
|
90
|
+
"""
|
|
91
|
+
if data is None:
|
|
92
|
+
return data
|
|
93
|
+
torch_dtype = output_torch_dtype(precision) # validates precision (shared string→dtype map)
|
|
94
|
+
if torch.is_tensor(data):
|
|
95
|
+
return data.to(torch_dtype)
|
|
96
|
+
return np.asarray(data).astype(_OUTPUT_NUMPY_DTYPE[precision], copy=False)
|
|
97
|
+
|
|
98
|
+
|
|
77
99
|
def _ensure_array(data: Any) -> np.ndarray:
|
|
78
100
|
if isinstance(data, np.ndarray):
|
|
79
101
|
return data
|
|
@@ -88,6 +88,7 @@ tiling:
|
|
|
88
88
|
|
|
89
89
|
speed:
|
|
90
90
|
precision: # model inference precision ["fp32", "fp16", "bf16"]; if not set, determined automatically based on model recommendations
|
|
91
|
+
output_dtype: # on-disk feature dtype ["fp16", "fp32"]; if not set, follows speed.precision (fp16 -> fp16, otherwise fp32)
|
|
91
92
|
num_dataloader_workers: # number of DataLoader worker processes per GPU rank for reading tiles during embedding; defaults to auto (job CPU budget split across GPUs, except cuCIM on-the-fly uses per-GPU budget // speed.num_cucim_workers)
|
|
92
93
|
num_gpus: # number of GPUs to use for feature extraction; defaults to all available GPUs
|
|
93
94
|
num_preprocessing_workers: # number of workers for hs2p tiling (WSI reading, JPEG encoding, tar writing); defaults to the runtime CPU budget capped at 64
|
|
@@ -16,8 +16,10 @@ from slide2vec.encoders.base import (
|
|
|
16
16
|
)
|
|
17
17
|
from slide2vec.encoders.registry import (
|
|
18
18
|
encoder_registry,
|
|
19
|
+
normalize_patch_size,
|
|
19
20
|
register_encoder,
|
|
20
21
|
resolve_encoder_output,
|
|
22
|
+
resolve_patch_size,
|
|
21
23
|
resolve_preprocessing_requirements,
|
|
22
24
|
resolve_tile_dependency_output,
|
|
23
25
|
)
|
|
@@ -35,7 +37,9 @@ __all__ = [
|
|
|
35
37
|
"resolve_recommended_dynamic_img_size",
|
|
36
38
|
"resolve_requested_output_variant",
|
|
37
39
|
"encoder_registry",
|
|
40
|
+
"normalize_patch_size",
|
|
38
41
|
"register_encoder",
|
|
42
|
+
"resolve_patch_size",
|
|
39
43
|
"resolve_preprocessing_requirements",
|
|
40
44
|
"resolve_encoder_output",
|
|
41
45
|
"resolve_tile_dependency_output",
|
|
@@ -78,6 +78,7 @@ def _encode_trunk_dense(*, trunk, batch: Tensor, encoder_name: str) -> Tensor:
|
|
|
78
78
|
output_variants={"default": {"encode_dim": 512}},
|
|
79
79
|
default_output_variant="default",
|
|
80
80
|
input_size=448,
|
|
81
|
+
patch_size=16,
|
|
81
82
|
supported_spacing_um=0.5,
|
|
82
83
|
precision="fp32",
|
|
83
84
|
source="MahmoodLab/conch",
|
|
@@ -163,6 +164,7 @@ class CONCH(TileEncoder):
|
|
|
163
164
|
output_variants={"default": {"encode_dim": 768}},
|
|
164
165
|
default_output_variant="default",
|
|
165
166
|
input_size=448,
|
|
167
|
+
patch_size=16,
|
|
166
168
|
supported_spacing_um=0.5,
|
|
167
169
|
precision="fp16",
|
|
168
170
|
source="MahmoodLab/TITAN",
|
|
@@ -151,6 +151,7 @@ class _HibouBase(TileEncoder):
|
|
|
151
151
|
output_variants={"default": {"encode_dim": 768}},
|
|
152
152
|
default_output_variant="default",
|
|
153
153
|
input_size=224,
|
|
154
|
+
patch_size=14,
|
|
154
155
|
supported_spacing_um=0.5,
|
|
155
156
|
precision="fp16",
|
|
156
157
|
source="histai/hibou-b",
|
|
@@ -167,6 +168,7 @@ class HibouB(_HibouBase):
|
|
|
167
168
|
output_variants={"default": {"encode_dim": 1024}},
|
|
168
169
|
default_output_variant="default",
|
|
169
170
|
input_size=224,
|
|
171
|
+
patch_size=14,
|
|
170
172
|
supported_spacing_um=0.5,
|
|
171
173
|
precision="fp16",
|
|
172
174
|
source="histai/hibou-L",
|
|
@@ -72,6 +72,7 @@ class _HOptimusBase(TimmTileEncoder):
|
|
|
72
72
|
output_variants={"default": {"encode_dim": 1536}},
|
|
73
73
|
default_output_variant="default",
|
|
74
74
|
input_size=224,
|
|
75
|
+
patch_size=14,
|
|
75
76
|
supported_spacing_um=0.5,
|
|
76
77
|
precision="fp16",
|
|
77
78
|
source="bioptimus/H-optimus-0",
|
|
@@ -108,6 +109,7 @@ class HOptimus0(TimmTileEncoder):
|
|
|
108
109
|
output_variants={"default": {"encode_dim": 1536}},
|
|
109
110
|
default_output_variant="default",
|
|
110
111
|
input_size=224,
|
|
112
|
+
patch_size=14,
|
|
111
113
|
supported_spacing_um=0.5,
|
|
112
114
|
precision="fp16",
|
|
113
115
|
source="bioptimus/H-optimus-1",
|
|
@@ -145,6 +147,7 @@ class HOptimus1(TimmTileEncoder):
|
|
|
145
147
|
},
|
|
146
148
|
default_output_variant="cls_patch_mean",
|
|
147
149
|
input_size=224,
|
|
150
|
+
patch_size=14,
|
|
148
151
|
supported_spacing_um=0.5,
|
|
149
152
|
precision="fp16",
|
|
150
153
|
source="bioptimus/H0-mini",
|
|
@@ -9,6 +9,7 @@ from slide2vec.encoders.registry import register_encoder
|
|
|
9
9
|
output_variants={"default": {"encode_dim": 384}},
|
|
10
10
|
default_output_variant="default",
|
|
11
11
|
input_size=224,
|
|
12
|
+
patch_size=8,
|
|
12
13
|
supported_spacing_um=0.5,
|
|
13
14
|
precision="fp32",
|
|
14
15
|
source="1aurent/vit_small_patch8_224.lunit_dino",
|
|
@@ -27,6 +27,7 @@ from slide2vec.encoders.registry import register_encoder
|
|
|
27
27
|
output_variants={"default": {"encode_dim": 3072}},
|
|
28
28
|
default_output_variant="default",
|
|
29
29
|
input_size=224,
|
|
30
|
+
patch_size=14,
|
|
30
31
|
supported_spacing_um=[0.25, 0.5, 1.0, 2.0],
|
|
31
32
|
precision="fp16",
|
|
32
33
|
source="kaiko-ai/midnight",
|
|
@@ -148,6 +148,7 @@ class _PhikonBase(TileEncoder):
|
|
|
148
148
|
output_variants={"default": {"encode_dim": 768}},
|
|
149
149
|
default_output_variant="default",
|
|
150
150
|
input_size=224,
|
|
151
|
+
patch_size=16,
|
|
151
152
|
supported_spacing_um=0.5,
|
|
152
153
|
precision="fp32",
|
|
153
154
|
source="owkin/phikon",
|
|
@@ -165,6 +166,7 @@ class Phikon(_PhikonBase):
|
|
|
165
166
|
output_variants={"default": {"encode_dim": 1024}},
|
|
166
167
|
default_output_variant="default",
|
|
167
168
|
input_size=224,
|
|
169
|
+
patch_size=16,
|
|
168
170
|
supported_spacing_um=0.5,
|
|
169
171
|
precision="fp32",
|
|
170
172
|
source="owkin/phikon-v2",
|
|
@@ -13,6 +13,7 @@ from slide2vec.encoders.registry import register_encoder
|
|
|
13
13
|
output_variants={"default": {"encode_dim": 1024}},
|
|
14
14
|
default_output_variant="default",
|
|
15
15
|
input_size=224,
|
|
16
|
+
patch_size=16,
|
|
16
17
|
supported_spacing_um=0.5,
|
|
17
18
|
precision="fp16",
|
|
18
19
|
source="MahmoodLab/UNI",
|
|
@@ -32,6 +33,7 @@ class UNI(TimmTileEncoder):
|
|
|
32
33
|
output_variants={"default": {"encode_dim": 1536}},
|
|
33
34
|
default_output_variant="default",
|
|
34
35
|
input_size=224,
|
|
36
|
+
patch_size=14,
|
|
35
37
|
supported_spacing_um=0.5,
|
|
36
38
|
precision="fp16",
|
|
37
39
|
source="MahmoodLab/UNI2-h",
|
|
@@ -50,6 +50,7 @@ class _VirchowBase(TimmTileEncoder):
|
|
|
50
50
|
},
|
|
51
51
|
default_output_variant="cls_patch_mean",
|
|
52
52
|
input_size=224,
|
|
53
|
+
patch_size=14,
|
|
53
54
|
supported_spacing_um=0.5,
|
|
54
55
|
precision="fp16",
|
|
55
56
|
source="paige-ai/Virchow",
|
|
@@ -67,6 +68,7 @@ class Virchow(_VirchowBase):
|
|
|
67
68
|
},
|
|
68
69
|
default_output_variant="cls_patch_mean",
|
|
69
70
|
input_size=224,
|
|
71
|
+
patch_size=14,
|
|
70
72
|
supported_spacing_um=[0.25, 0.5, 1.0, 2.0],
|
|
71
73
|
precision="fp16",
|
|
72
74
|
source="paige-ai/Virchow2",
|
|
@@ -37,6 +37,7 @@ def register_encoder(
|
|
|
37
37
|
output_variants: dict[str, dict[str, Any]],
|
|
38
38
|
default_output_variant: str,
|
|
39
39
|
input_size: int | None = None,
|
|
40
|
+
patch_size: int | tuple[int, int] | None = None,
|
|
40
41
|
level: str = "tile",
|
|
41
42
|
tile_encoder: str | None = None,
|
|
42
43
|
tile_encoder_output_variant: str | None = None,
|
|
@@ -51,6 +52,12 @@ def register_encoder(
|
|
|
51
52
|
output_variants: Supported named encoder outputs with concrete metadata.
|
|
52
53
|
default_output_variant: Default output variant name.
|
|
53
54
|
input_size: Recommended encoder input image size in pixels.
|
|
55
|
+
patch_size: Backbone patch size, as ``int`` (square) or ``(patch_h,
|
|
56
|
+
patch_w)``. Optional: only dense-capable ViT tile encoders have a
|
|
57
|
+
meaningful patch grid. Declared statically so the dense token grid /
|
|
58
|
+
cache key can be resolved via :func:`resolve_patch_size` WITHOUT
|
|
59
|
+
instantiating the (multi-GB) encoder; the model-load path asserts this
|
|
60
|
+
static value still equals the loaded model's runtime ``patch_size``.
|
|
54
61
|
level: Encoder output level ("tile" or "slide").
|
|
55
62
|
tile_encoder: Registered tile encoder dependency for slide-level models.
|
|
56
63
|
tile_encoder_output_variant: Fixed tile-encoder output variant for slide models.
|
|
@@ -67,6 +74,7 @@ def register_encoder(
|
|
|
67
74
|
"default_output_variant": default_output_variant,
|
|
68
75
|
"level": level,
|
|
69
76
|
"input_size": input_size,
|
|
77
|
+
"patch_size": patch_size,
|
|
70
78
|
"tile_encoder": tile_encoder,
|
|
71
79
|
"tile_encoder_output_variant": tile_encoder_output_variant,
|
|
72
80
|
"supported_spacing_um": supported_spacing_um,
|
|
@@ -76,6 +84,41 @@ def register_encoder(
|
|
|
76
84
|
return encoder_registry.register_decorator(name, metadata=metadata)
|
|
77
85
|
|
|
78
86
|
|
|
87
|
+
def normalize_patch_size(value: int | tuple[int, int]) -> tuple[int, int]:
|
|
88
|
+
"""Normalize a patch size to a ``(patch_h, patch_w)`` int tuple.
|
|
89
|
+
|
|
90
|
+
This is the SAME representation the runtime ``encoder.patch_size`` instance
|
|
91
|
+
property returns, so static and runtime values compare/serialize identically
|
|
92
|
+
(a downstream dense cache key depends on this byte-for-byte equality).
|
|
93
|
+
"""
|
|
94
|
+
if isinstance(value, int):
|
|
95
|
+
return (value, value)
|
|
96
|
+
patch_h, patch_w = value
|
|
97
|
+
return (int(patch_h), int(patch_w))
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def resolve_patch_size(
|
|
101
|
+
encoder_name: str,
|
|
102
|
+
metadata: dict[str, Any] | None = None,
|
|
103
|
+
) -> tuple[int, int]:
|
|
104
|
+
"""Resolve an encoder's static patch size WITHOUT constructing the model.
|
|
105
|
+
|
|
106
|
+
Reads the ``patch_size`` declared on the encoder's ``@register_encoder`` and
|
|
107
|
+
normalizes it to a ``(patch_h, patch_w)`` int tuple (matching the runtime
|
|
108
|
+
``encoder.patch_size``). Raises a clear error for encoders that do not declare
|
|
109
|
+
one (non-dense encoders) rather than returning a wrong value.
|
|
110
|
+
"""
|
|
111
|
+
info = metadata if metadata is not None else encoder_registry.info(encoder_name)
|
|
112
|
+
patch = info.get("patch_size")
|
|
113
|
+
if patch is None:
|
|
114
|
+
raise ValueError(
|
|
115
|
+
f"Encoder '{encoder_name}' does not declare a patch_size in its registry "
|
|
116
|
+
"metadata. patch_size is only defined for dense-capable ViT tile encoders; "
|
|
117
|
+
"non-dense encoders have no recoverable patch grid."
|
|
118
|
+
)
|
|
119
|
+
return normalize_patch_size(patch)
|
|
120
|
+
|
|
121
|
+
|
|
79
122
|
def resolve_preprocessing_requirements(
|
|
80
123
|
encoder_name: str,
|
|
81
124
|
metadata: dict[str, Any] | None = None,
|
|
@@ -59,6 +59,7 @@ from slide2vec.artifacts import (
|
|
|
59
59
|
from slide2vec.encoders.registry import (
|
|
60
60
|
encoder_registry,
|
|
61
61
|
resolve_encoder_output,
|
|
62
|
+
resolve_patch_size,
|
|
62
63
|
resolve_preprocessing_defaults,
|
|
63
64
|
)
|
|
64
65
|
from slide2vec.runtime.model_settings import canonicalize_model_name
|
|
@@ -121,6 +122,21 @@ def load_model(
|
|
|
121
122
|
extra_kwargs["allow_non_recommended_settings"] = allow_non_recommended_settings
|
|
122
123
|
encoder = encoder_cls(output_variant=output_variant, **extra_kwargs)
|
|
123
124
|
|
|
125
|
+
# Drift guard: the static patch_size declared on @register_encoder is read
|
|
126
|
+
# (without loading the model) to resolve the dense cache key. Assert it still
|
|
127
|
+
# equals the loaded model's runtime patch_size so the static metadata can never
|
|
128
|
+
# silently diverge from the architecture and corrupt that key.
|
|
129
|
+
if info.get("patch_size") is not None:
|
|
130
|
+
static_patch_size = resolve_patch_size(name, metadata=info)
|
|
131
|
+
runtime_patch_size = encoder.patch_size
|
|
132
|
+
if tuple(runtime_patch_size) != static_patch_size:
|
|
133
|
+
raise ValueError(
|
|
134
|
+
f"Encoder '{name}' declares a static patch_size {static_patch_size} "
|
|
135
|
+
f"(registry metadata) but the loaded model reports "
|
|
136
|
+
f"{tuple(runtime_patch_size)}. The static value seeds the dense cache "
|
|
137
|
+
"key; fix the @register_encoder patch_size to match the architecture."
|
|
138
|
+
)
|
|
139
|
+
|
|
124
140
|
tile_encoder = None
|
|
125
141
|
if resolved_level == "tile":
|
|
126
142
|
transforms = encoder.get_transform()
|
|
@@ -40,8 +40,30 @@ import torch.nn.functional as F
|
|
|
40
40
|
from PIL import Image
|
|
41
41
|
|
|
42
42
|
from slide2vec.runtime.dense_sliding import encode_dense_sliding
|
|
43
|
+
from slide2vec.runtime.model_settings import output_torch_dtype, resolve_output_precision
|
|
43
44
|
from slide2vec.runtime.slide_encode import slide_encode_autocast_ctx
|
|
44
45
|
|
|
46
|
+
|
|
47
|
+
def _resolve_output_dtype(output_dtype: "torch.dtype | None", precision: str) -> "torch.dtype":
|
|
48
|
+
"""Resolve the dtype emitted grids are materialized in.
|
|
49
|
+
|
|
50
|
+
Defaults (``output_dtype is None``) to the model's compute precision — fp16 runs
|
|
51
|
+
yield fp16 grids — so the engine no longer force-upcasts everything to float32.
|
|
52
|
+
numpy has no bfloat16, so a bf16 compute precision widens to float32 (its lossless
|
|
53
|
+
container) and an *explicit* ``torch.bfloat16`` request is rejected: the grid is
|
|
54
|
+
materialized via ``.numpy()`` and bfloat16 cannot cross that boundary.
|
|
55
|
+
"""
|
|
56
|
+
if output_dtype is None:
|
|
57
|
+
# Shared rule with the pooled write path: fp16 compute -> fp16, else fp32.
|
|
58
|
+
return output_torch_dtype(resolve_output_precision(None, precision))
|
|
59
|
+
if output_dtype == torch.bfloat16:
|
|
60
|
+
raise ValueError(
|
|
61
|
+
"output_dtype=torch.bfloat16 cannot be materialized as a numpy grid; "
|
|
62
|
+
"request torch.float16 or torch.float32"
|
|
63
|
+
)
|
|
64
|
+
return output_dtype
|
|
65
|
+
|
|
66
|
+
|
|
45
67
|
_PAD_MODES = {"reflect", "constant", "zero", "replicate"}
|
|
46
68
|
|
|
47
69
|
|
|
@@ -161,14 +183,16 @@ def iter_regions_dense(
|
|
|
161
183
|
attention_include_registers: bool = False,
|
|
162
184
|
batch_size: int = 1,
|
|
163
185
|
precision: str = "fp32",
|
|
186
|
+
output_dtype: "torch.dtype | None" = None,
|
|
164
187
|
dense_transform: Callable | None = None,
|
|
165
188
|
) -> Iterator[np.ndarray]:
|
|
166
189
|
"""Stream slide regions at ``coordinates`` into dense grids, one per coordinate.
|
|
167
190
|
|
|
168
|
-
Yields one ``(d, grid_h, grid_w)``
|
|
169
|
-
|
|
170
|
-
|
|
171
|
-
|
|
191
|
+
Yields one ``(d, grid_h, grid_w)`` grid per coordinate, in coordinate order, in the
|
|
192
|
+
model's compute ``precision`` by default (fp16 runs yield fp16 grids; see
|
|
193
|
+
``output_dtype``). Regions are read and encoded one ``batch_size`` chunk at a time, so
|
|
194
|
+
resident host memory is bounded by ``batch_size`` rather than by a slide's ROI count
|
|
195
|
+
(the loop holds at most one batch of grids resident — no per-slide accumulation).
|
|
172
196
|
|
|
173
197
|
Injectable core: takes a constructed dense-capable ``model`` (with
|
|
174
198
|
``encode_tiles_dense`` / ``encode_tiles_attention`` / ``patch_size`` /
|
|
@@ -194,8 +218,12 @@ def iter_regions_dense(
|
|
|
194
218
|
sliding is internal to extraction.
|
|
195
219
|
overlap: fractional window overlap in ``[0, 1)`` for the sliding path (ignored when
|
|
196
220
|
``window_size is None``); the stride is ``window * (1 - overlap)``.
|
|
221
|
+
output_dtype: torch dtype the grids are materialized in. ``None`` (default) follows
|
|
222
|
+
the compute ``precision`` — fp16 → fp16, fp32 → fp32, bf16 → fp32 (numpy has no
|
|
223
|
+
bfloat16). Pass e.g. ``torch.float32`` to force a lossless cache regardless of
|
|
224
|
+
precision; an explicit ``torch.bfloat16`` is rejected (cannot cross ``.numpy()``).
|
|
197
225
|
|
|
198
|
-
Yields
|
|
226
|
+
Yields grids in coordinate order in ``output_dtype``; empty ``coordinates`` yields nothing.
|
|
199
227
|
``feature_kind`` selects ``encode_tiles_dense`` (patch grid) vs
|
|
200
228
|
``encode_tiles_attention`` (CLS-attention grid); both produce a ``(C, gh, gw)`` grid and
|
|
201
229
|
share this path. Each yielded grid is a standalone contiguous copy, so it does not pin
|
|
@@ -203,6 +231,7 @@ def iter_regions_dense(
|
|
|
203
231
|
"""
|
|
204
232
|
if pad_mode not in _PAD_MODES:
|
|
205
233
|
raise ValueError(f"unsupported pad_mode {pad_mode!r}; expected one of {sorted(_PAD_MODES)}")
|
|
234
|
+
resolved_output_dtype = _resolve_output_dtype(output_dtype, precision)
|
|
206
235
|
geometry = compute_dense_geometry(target_size=target_size, patch_size=model.patch_size)
|
|
207
236
|
if dense_transform is None:
|
|
208
237
|
dense_transform = model.get_dense_transform()
|
|
@@ -262,7 +291,7 @@ def iter_regions_dense(
|
|
|
262
291
|
raise ValueError(
|
|
263
292
|
f"{feature_kind} encode returned a {out.ndim}-D tensor; expected (B, d, gh, gw)."
|
|
264
293
|
)
|
|
265
|
-
batch_np = out.detach().
|
|
294
|
+
batch_np = out.detach().to(resolved_output_dtype).cpu().numpy()
|
|
266
295
|
for i in range(batch_np.shape[0]):
|
|
267
296
|
# Standalone C-contiguous copy: a per-row view would pin the whole
|
|
268
297
|
# batch alive (the blended sliding output is contiguous, so a view of
|
|
@@ -10,11 +10,13 @@ from slide2vec.artifacts import (
|
|
|
10
10
|
HierarchicalEmbeddingArtifact,
|
|
11
11
|
SlideEmbeddingArtifact,
|
|
12
12
|
TileEmbeddingArtifact,
|
|
13
|
+
cast_feature_dtype,
|
|
13
14
|
write_hierarchical_embeddings,
|
|
14
15
|
write_slide_embeddings,
|
|
15
16
|
write_tile_embeddings,
|
|
16
17
|
)
|
|
17
18
|
from slide2vec.runtime.hierarchical import resolve_hierarchical_geometry
|
|
19
|
+
from slide2vec.runtime.model_settings import resolve_output_precision
|
|
18
20
|
|
|
19
21
|
|
|
20
22
|
def tiling_result_annotation(tiling_result) -> str | None:
|
|
@@ -114,12 +116,14 @@ def write_tile_embedding_artifact(
|
|
|
114
116
|
) -> TileEmbeddingArtifact:
|
|
115
117
|
if execution.output_dir is None:
|
|
116
118
|
raise ValueError("ExecutionOptions.output_dir is required to persist tile embeddings")
|
|
119
|
+
precision = resolve_output_precision(execution.output_dtype, execution.precision)
|
|
120
|
+
features = cast_feature_dtype(features, precision)
|
|
117
121
|
return write_tile_embeddings(
|
|
118
122
|
sample_id,
|
|
119
123
|
features,
|
|
120
124
|
output_dir=execution.output_dir,
|
|
121
125
|
output_format=execution.output_format,
|
|
122
|
-
metadata=metadata,
|
|
126
|
+
metadata={**metadata, "feature_dtype": precision},
|
|
123
127
|
tile_index=np.arange(_num_rows(features), dtype=np.int64),
|
|
124
128
|
annotation=annotation,
|
|
125
129
|
)
|
|
@@ -142,13 +146,14 @@ def write_slide_embedding_artifact(
|
|
|
142
146
|
) -> SlideEmbeddingArtifact:
|
|
143
147
|
if execution.output_dir is None:
|
|
144
148
|
raise ValueError("ExecutionOptions.output_dir is required to persist slide embeddings")
|
|
149
|
+
precision = resolve_output_precision(execution.output_dtype, execution.precision)
|
|
145
150
|
return write_slide_embeddings(
|
|
146
151
|
sample_id,
|
|
147
|
-
embedding,
|
|
152
|
+
cast_feature_dtype(embedding, precision),
|
|
148
153
|
output_dir=execution.output_dir,
|
|
149
154
|
output_format=execution.output_format,
|
|
150
|
-
metadata=metadata,
|
|
151
|
-
latents=latents,
|
|
155
|
+
metadata={**metadata, "feature_dtype": precision},
|
|
156
|
+
latents=cast_feature_dtype(latents, precision),
|
|
152
157
|
annotation=annotation,
|
|
153
158
|
)
|
|
154
159
|
|
|
@@ -163,11 +168,12 @@ def write_hierarchical_embedding_artifact(
|
|
|
163
168
|
) -> HierarchicalEmbeddingArtifact:
|
|
164
169
|
if execution.output_dir is None:
|
|
165
170
|
raise ValueError("ExecutionOptions.output_dir is required to persist hierarchical embeddings")
|
|
171
|
+
precision = resolve_output_precision(execution.output_dtype, execution.precision)
|
|
166
172
|
return write_hierarchical_embeddings(
|
|
167
173
|
sample_id,
|
|
168
|
-
features,
|
|
174
|
+
cast_feature_dtype(features, precision),
|
|
169
175
|
output_dir=execution.output_dir,
|
|
170
176
|
output_format=execution.output_format,
|
|
171
|
-
metadata=metadata,
|
|
177
|
+
metadata={**metadata, "feature_dtype": precision},
|
|
172
178
|
annotation=annotation,
|
|
173
179
|
)
|
|
@@ -0,0 +1,97 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
from typing import Any
|
|
3
|
+
|
|
4
|
+
logger = logging.getLogger("slide2vec")
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
PRECISION_ALIASES = {
|
|
8
|
+
"fp32": "fp32",
|
|
9
|
+
"float32": "fp32",
|
|
10
|
+
"32": "fp32",
|
|
11
|
+
"fp16": "fp16",
|
|
12
|
+
"float16": "fp16",
|
|
13
|
+
"16": "fp16",
|
|
14
|
+
"half": "fp16",
|
|
15
|
+
"bf16": "bf16",
|
|
16
|
+
"bfloat16": "bf16",
|
|
17
|
+
}
|
|
18
|
+
|
|
19
|
+
MODEL_NAME_ALIASES = {
|
|
20
|
+
"conch-v1.5": "conchv15",
|
|
21
|
+
"conch_v15": "conchv15",
|
|
22
|
+
"conchv1.5": "conchv15",
|
|
23
|
+
"conchv1_5": "conchv15",
|
|
24
|
+
"phikon-v2": "phikonv2",
|
|
25
|
+
"hibou-b": "hibou-b",
|
|
26
|
+
"hibou-l": "hibou-l",
|
|
27
|
+
"h-optimus-0-mini": "h0-mini",
|
|
28
|
+
"prov-gigapath": "gigapath",
|
|
29
|
+
"prov-gigapath-tile": "gigapath",
|
|
30
|
+
"prov-gigapath-slide": "gigapath-slide",
|
|
31
|
+
"kaiko-midnight": "midnight",
|
|
32
|
+
}
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def normalize_precision_name(value: Any) -> str | None:
|
|
36
|
+
if value is None:
|
|
37
|
+
return None
|
|
38
|
+
normalized = str(value).strip().lower()
|
|
39
|
+
if normalized not in PRECISION_ALIASES:
|
|
40
|
+
supported = ", ".join(sorted(PRECISION_ALIASES))
|
|
41
|
+
raise ValueError(f"Unsupported precision {value!r}. Expected one of: {supported}")
|
|
42
|
+
return PRECISION_ALIASES[normalized]
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def normalize_output_dtype(value: Any) -> str | None:
|
|
46
|
+
"""Normalize a requested on-disk feature dtype to ``"fp16"`` / ``"fp32"`` (or ``None``).
|
|
47
|
+
|
|
48
|
+
``None`` means *follow the compute precision* (see :func:`resolve_output_precision`).
|
|
49
|
+
``bf16`` is rejected: tile/slide/hierarchical/patient artifacts serialize through numpy,
|
|
50
|
+
which has no bfloat16 — the same boundary the dense path guards.
|
|
51
|
+
"""
|
|
52
|
+
normalized = normalize_precision_name(value)
|
|
53
|
+
if normalized == "bf16":
|
|
54
|
+
raise ValueError(
|
|
55
|
+
"Unsupported output dtype 'bf16'. Feature artifacts serialize through numpy "
|
|
56
|
+
"(no bfloat16); choose 'fp16' or 'fp32', or leave unset to follow precision."
|
|
57
|
+
)
|
|
58
|
+
return normalized
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def resolve_output_precision(output_dtype: Any, compute_precision: Any) -> str:
|
|
62
|
+
"""Resolve the concrete on-disk feature precision (``"fp16"`` or ``"fp32"``).
|
|
63
|
+
|
|
64
|
+
``output_dtype is None`` follows ``compute_precision``: an fp16 forward keeps fp16
|
|
65
|
+
features, while bf16 / fp32 (and an unset or unknown compute precision) widen to fp32 —
|
|
66
|
+
fp32 is bf16's lossless container and the only float dtype a numpy artifact can hold.
|
|
67
|
+
A non-null ``output_dtype`` is honored verbatim (after :func:`normalize_output_dtype`).
|
|
68
|
+
This is the single source of truth shared by the pooled write path and the dense
|
|
69
|
+
``iter_regions_dense`` path.
|
|
70
|
+
"""
|
|
71
|
+
normalized = normalize_output_dtype(output_dtype)
|
|
72
|
+
if normalized is not None:
|
|
73
|
+
return normalized
|
|
74
|
+
return "fp16" if normalize_precision_name(compute_precision) == "fp16" else "fp32"
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def output_torch_dtype(precision: str):
|
|
78
|
+
"""The torch dtype an on-disk feature artifact materializes in for ``precision``.
|
|
79
|
+
|
|
80
|
+
``precision`` is a value returned by :func:`resolve_output_precision` (``"fp16"`` or
|
|
81
|
+
``"fp32"``); anything else is a programming error. This is the single string→dtype
|
|
82
|
+
mapping shared by the pooled write path (:func:`slide2vec.artifacts.cast_feature_dtype`)
|
|
83
|
+
and the dense ``iter_regions_dense`` path, so both agree on the materialized dtype.
|
|
84
|
+
torch is imported lazily to keep this module importable without it.
|
|
85
|
+
"""
|
|
86
|
+
import torch
|
|
87
|
+
|
|
88
|
+
mapping = {"fp16": torch.float16, "fp32": torch.float32}
|
|
89
|
+
if precision not in mapping:
|
|
90
|
+
raise ValueError(f"Unsupported output precision {precision!r}; expected 'fp16' or 'fp32'.")
|
|
91
|
+
return mapping[precision]
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def canonicalize_model_name(name: str) -> str:
|
|
95
|
+
normalized = name.strip().lower()
|
|
96
|
+
return MODEL_NAME_ALIASES.get(normalized, normalized)
|
|
97
|
+
|