Coverage for src/bioimageio/spec/model/v0_5.py: 72%

1764 statements  

« prev     ^ index     » next       coverage.py v7.16.1, created at 2026-09-15 12:09 +0000

1from __future__ import annotations 

2 

3import collections.abc 

4import re 

5import string 

6import warnings 

7from abc import ABC 

8from copy import deepcopy 

9from functools import partial 

10from itertools import chain 

11from math import ceil 

12from pathlib import Path, PurePosixPath 

13from tempfile import mkdtemp 

14from textwrap import dedent 

15from typing import ( 

16 TYPE_CHECKING, 

17 Any, 

18 Callable, 

19 ClassVar, 

20 Dict, 

21 Generic, 

22 List, 

23 Literal, 

24 Mapping, 

25 NamedTuple, 

26 Optional, 

27 Sequence, 

28 TypeVar, 

29 Union, 

30 cast, 

31 overload, 

32) 

33 

34import numpy as np 

35from annotated_types import Ge, Gt, Interval, MaxLen, MinLen, Predicate 

36from imageio.v3 import imread, imwrite # pyright: ignore[reportUnknownVariableType] 

37from loguru import logger 

38from numpy.typing import NDArray 

39from pydantic import ( 

40 AfterValidator, 

41 Discriminator, 

42 Field, 

43 RootModel, 

44 SerializationInfo, 

45 SerializerFunctionWrapHandler, 

46 StrictInt, 

47 Tag, 

48 ValidationInfo, 

49 WrapSerializer, 

50 field_validator, 

51 model_serializer, 

52 model_validator, 

53) 

54from typing_extensions import Annotated, Self, TypeAlias, assert_never, get_args 

55 

56from .._internal.common_nodes import ( 

57 InvalidDescr, 

58 KwargsNode, 

59 Node, 

60 NodeWithExplicitlySetFields, 

61) 

62from .._internal.constants import DTYPE_LIMITS 

63from .._internal.field_warning import issue_warning, warn 

64from .._internal.io import BioimageioYamlContent as BioimageioYamlContent 

65from .._internal.io import FileDescr as FileDescr 

66from .._internal.io import ( 

67 FileSource, 

68 WithSuffix, 

69 YamlValue, 

70 extract_file_name, 

71 get_reader, 

72 wo_special_file_name, 

73) 

74from .._internal.io_basics import Sha256 as Sha256 

75from .._internal.io_packaging import ( 

76 FileDescr_package, 

77 package_file_descr_serializer, 

78) 

79from .._internal.io_utils import load_array, open_bioimageio_yaml 

80from .._internal.node_converter import Converter 

81from .._internal.type_guards import is_dict, is_sequence 

82from .._internal.types import ( 

83 FAIR, 

84 LowerCaseIdentifier, 

85 LowerCaseIdentifierAnno, 

86 validate_identifier, 

87 validate_is_not_keyword, 

88) 

89from .._internal.types import ( 

90 AbsoluteTolerance as AbsoluteTolerance, 

91) 

92from .._internal.types import Datetime as Datetime 

93from .._internal.types import Identifier as Identifier 

94from .._internal.types import ( 

95 MismatchedElementsPerMillion as MismatchedElementsPerMillion, 

96) 

97from .._internal.types import NotEmpty as NotEmpty 

98from .._internal.types import ( 

99 RelativeTolerance as RelativeTolerance, 

100) 

101from .._internal.types import SiUnit as SiUnit 

102from .._internal.url import HttpUrl as HttpUrl 

103from .._internal.utils import try_all_raise_last 

104from .._internal.validation_context import get_validation_context 

105from .._internal.validator_annotations import RestrictCharacters 

106from .._internal.version_type import Version as Version 

107from .._internal.warning_levels import INFO 

108from ..dataset.v0_2 import DatasetDescr as DatasetDescr02 

109from ..dataset.v0_2 import LinkedDataset as LinkedDataset02 

110from ..dataset.v0_3 import DatasetDescr as DatasetDescr 

111from ..dataset.v0_3 import DatasetId as DatasetId 

112from ..dataset.v0_3 import LinkedDataset as LinkedDataset 

113from ..dataset.v0_3 import Uploader as Uploader 

114from ..generic._v0_3_converter import convert_plain_covers_and_docs_and_icon 

115from ..generic.v0_3 import ( 

116 VALID_COVER_IMAGE_EXTENSIONS as VALID_COVER_IMAGE_EXTENSIONS, 

117) 

118from ..generic.v0_3 import Author as Author 

119from ..generic.v0_3 import BadgeDescr as BadgeDescr 

120from ..generic.v0_3 import CiteEntry as CiteEntry 

121from ..generic.v0_3 import DeprecatedLicenseId as DeprecatedLicenseId 

122from ..generic.v0_3 import Doi as Doi 

123from ..generic.v0_3 import ( 

124 FileDescr_documentation, 

125 GenericModelDescrBase, 

126 LinkedResourceBase, 

127 _author_conv, # pyright: ignore[reportPrivateUsage] 

128 _maintainer_conv, # pyright: ignore[reportPrivateUsage] 

129) 

130from ..generic.v0_3 import LicenseId as LicenseId 

131from ..generic.v0_3 import LinkedResource as LinkedResource 

132from ..generic.v0_3 import Maintainer as Maintainer 

133from ..generic.v0_3 import OrcidId as OrcidId 

134from ..generic.v0_3 import RelativeFilePath as RelativeFilePath 

135from ..generic.v0_3 import ResourceId as ResourceId 

136from .v0_4 import Author as _Author_v0_4 

137from .v0_4 import BinarizeDescr as _BinarizeDescr_v0_4 

138from .v0_4 import CallableFromDepencency as CallableFromDepencency 

139from .v0_4 import CallableFromDepencency as _CallableFromDepencency_v0_4 

140from .v0_4 import CallableFromFile as _CallableFromFile_v0_4 

141from .v0_4 import ClipDescr as _ClipDescr_v0_4 

142from .v0_4 import ImplicitOutputShape as _ImplicitOutputShape_v0_4 

143from .v0_4 import InputTensorDescr as _InputTensorDescr_v0_4 

144from .v0_4 import KnownRunMode as KnownRunMode 

145from .v0_4 import ModelDescr as _ModelDescr04 

146from .v0_4 import ModelDescr as _ModelDescr_v0_4 

147from .v0_4 import OutputTensorDescr as _OutputTensorDescr_v0_4 

148from .v0_4 import ParameterizedInputShape as _ParameterizedInputShape_v0_4 

149from .v0_4 import PostprocessingDescr as _PostprocessingDescr_v0_4 

150from .v0_4 import PreprocessingDescr as _PreprocessingDescr_v0_4 

151from .v0_4 import RunMode as RunMode 

152from .v0_4 import ScaleLinearDescr as _ScaleLinearDescr_v0_4 

153from .v0_4 import ScaleMeanVarianceDescr as _ScaleMeanVarianceDescr_v0_4 

154from .v0_4 import ScaleRangeDescr as _ScaleRangeDescr_v0_4 

155from .v0_4 import SigmoidDescr as _SigmoidDescr_v0_4 

156from .v0_4 import TensorName as _TensorName_v0_4 

157from .v0_4 import ZeroMeanUnitVarianceDescr as _ZeroMeanUnitVarianceDescr_v0_4 

158from .v0_4 import package_weights 

159 

160SpaceUnit = Literal[ 

161 "attometer", 

162 "angstrom", 

163 "centimeter", 

164 "decimeter", 

165 "exameter", 

166 "femtometer", 

167 "foot", 

168 "gigameter", 

169 "hectometer", 

170 "inch", 

171 "kilometer", 

172 "megameter", 

173 "meter", 

174 "micrometer", 

175 "mile", 

176 "millimeter", 

177 "nanometer", 

178 "parsec", 

179 "petameter", 

180 "picometer", 

181 "terameter", 

182 "yard", 

183 "yoctometer", 

184 "yottameter", 

185 "zeptometer", 

186 "zettameter", 

187] 

188"""Space unit compatible to the [OME-Zarr axes specification 0.5](https://ngff.openmicroscopy.org/0.5/#axes-md)""" 

189 

190TimeUnit = Literal[ 

191 "attosecond", 

192 "centisecond", 

193 "day", 

194 "decisecond", 

195 "exasecond", 

196 "femtosecond", 

197 "gigasecond", 

198 "hectosecond", 

199 "hour", 

200 "kilosecond", 

201 "megasecond", 

202 "microsecond", 

203 "millisecond", 

204 "minute", 

205 "nanosecond", 

206 "petasecond", 

207 "picosecond", 

208 "second", 

209 "terasecond", 

210 "yoctosecond", 

211 "yottasecond", 

212 "zeptosecond", 

213 "zettasecond", 

214] 

215"""Time unit compatible to the [OME-Zarr axes specification 0.5](https://ngff.openmicroscopy.org/0.5/#axes-md)""" 

216 

217AxisType = Literal["batch", "channel", "index", "time", "space"] 

218 

219_AXIS_TYPE_MAP: Mapping[str, AxisType] = { 

220 "b": "batch", 

221 "t": "time", 

222 "i": "index", 

223 "c": "channel", 

224 "x": "space", 

225 "y": "space", 

226 "z": "space", 

227} 

228 

229_AXIS_ID_MAP = { 

230 "b": "batch", 

231 "t": "time", 

232 "i": "index", 

233 "c": "channel", 

234 "s": "channel", 

235} 

236 

237WeightsFormat = Literal[ 

238 "keras_hdf5", 

239 "keras_v3", 

240 "onnx", 

241 "pytorch_state_dict", 

242 "tensorflow_js", 

243 "tensorflow_saved_model_bundle", 

244 "torchscript", 

245] 

246 

247 

248class TensorId(LowerCaseIdentifier): 

249 root_model: ClassVar[type[RootModel[Any]]] = RootModel[ 

250 Annotated[LowerCaseIdentifierAnno, MaxLen(32)] 

251 ] 

252 

253 

254def _normalize_axis_id(a: str): 

255 b = str(a).lower() 

256 normalized = _AXIS_ID_MAP.get(b, b) 

257 if a != normalized: 

258 logger.opt(depth=3).debug( 

259 "Normalized axis id from '{}' to '{}'.", a, normalized 

260 ) 

261 return normalized 

262 

263 

264class AxisId(LowerCaseIdentifier): 

265 root_model: ClassVar[type[RootModel[Any]]] = RootModel[ 

266 Annotated[ 

267 NotEmpty[str], 

268 AfterValidator(_normalize_axis_id), 

269 MaxLen(16), 

270 AfterValidator(validate_identifier), 

271 AfterValidator(validate_is_not_keyword), 

272 ] 

273 ] 

274 

275 

276def _is_batch(a: str) -> bool: 

277 return str(a) == "batch" 

278 

279 

280def _is_not_batch(a: str) -> bool: 

281 return not _is_batch(a) 

282 

283 

284NonBatchAxisId = Annotated[AxisId, Predicate(_is_not_batch)] 

285 

286PreprocessingId = Literal[ 

287 "binarize", 

288 "clip", 

289 "ensure_dtype", 

290 "fixed_zero_mean_unit_variance", 

291 "scale_linear", 

292 "scale_range", 

293 "sigmoid", 

294 "softmax", 

295] 

296PostprocessingId = Literal[ 

297 "binarize", 

298 "clip", 

299 "custom", 

300 "ensure_dtype", 

301 "fixed_zero_mean_unit_variance", 

302 "scale_linear", 

303 "scale_mean_variance", 

304 "scale_range", 

305 "sigmoid", 

306 "softmax", 

307 "zero_mean_unit_variance", 

308] 

309 

310 

311SAME_AS_TYPE = "<same as type>" 

312 

313 

314ParameterizedSize_N: TypeAlias = int 

315""" 

316Annotates an integer to calculate a concrete axis size from a `ParameterizedSize`. 

317""" 

318 

319 

320class ParameterizedSize(Node): 

321 """Describes a range of valid tensor axis sizes as `size = min + n*step`. 

322 

323 - **min** and **step** are given by the model description. 

324 - All blocksize paramters n = 0,1,2,... yield a valid `size`. 

325 - A greater blocksize paramter n = 0,1,2,... results in a greater **size**. 

326 This allows to adjust the axis size more generically. 

327 """ 

328 

329 N: ClassVar[type[int]] = ParameterizedSize_N 

330 """Positive integer to parameterize this axis""" 

331 

332 min: Annotated[int, Gt(0)] 

333 step: Annotated[int, Gt(0)] 

334 

335 def validate_size(self, size: int, msg_prefix: str = "") -> int: 

336 if size < self.min: 

337 raise ValueError( 

338 f"{msg_prefix}size {size} < {self.min} (minimum axis size)" 

339 ) 

340 if (size - self.min) % self.step != 0: 

341 raise ValueError( 

342 f"{msg_prefix}size {size} is not parameterized by `min + n*step` =" 

343 + f" `{self.min} + n*{self.step}`" 

344 ) 

345 

346 return size 

347 

348 def get_size(self, n: ParameterizedSize_N) -> int: 

349 return self.min + self.step * n 

350 

351 def get_n(self, s: int) -> ParameterizedSize_N: 

352 """return smallest n parameterizing a size greater or equal than `s`""" 

353 return ceil((s - self.min) / self.step) 

354 

355 

356class DataDependentSize(Node): 

357 min: Annotated[int, Gt(0)] = 1 

358 max: Annotated[int | None, Gt(1)] = None 

359 

360 @model_validator(mode="after") 

361 def _validate_max_gt_min(self): 

362 if self.max is not None and self.min >= self.max: 

363 raise ValueError(f"expected `min` < `max`, but got {self.min}, {self.max}") 

364 

365 return self 

366 

367 def validate_size(self, size: int, msg_prefix: str = "") -> int: 

368 if size < self.min: 

369 raise ValueError(f"{msg_prefix}size {size} < {self.min}") 

370 

371 if self.max is not None and size > self.max: 

372 raise ValueError(f"{msg_prefix}size {size} > {self.max}") 

373 

374 return size 

375 

376 

377class SizeReference(Node): 

378 """A tensor axis size (extent in pixels/frames) defined in relation to a reference axis. 

379 

380 `axis.size = reference.size * reference.scale / axis.scale + offset` 

381 

382 Note: 

383 1. The axis and the referenced axis need to have the same unit (or no unit). 

384 2. Batch axes may not be referenced. 

385 3. Fractions are rounded down. 

386 4. If the reference axis is `concatenable` the referencing axis is assumed to be 

387 `concatenable` as well with the same block order. 

388 

389 Example: 

390 An unisotropic input image of w*h=100*49 pixels depicts a phsical space of 200*196mm². 

391 Let's assume that we want to express the image height h in relation to its width w 

392 instead of only accepting input images of exactly 100*49 pixels 

393 (for example to express a range of valid image shapes by parametrizing w, see `ParameterizedSize`). 

394 

395 >>> w = SpaceInputAxis(id=AxisId("w"), size=100, unit="millimeter", scale=2) 

396 >>> h = SpaceInputAxis( 

397 ... id=AxisId("h"), 

398 ... size=SizeReference(tensor_id=TensorId("input"), axis_id=AxisId("w"), offset=-1), 

399 ... unit="millimeter", 

400 ... scale=4, 

401 ... ) 

402 >>> print(h.size.get_size(h, w)) 

403 49 

404 

405 ⇒ h = w * w.scale / h.scale + offset = 100 * 2mm / 4mm - 1 = 49 

406 """ 

407 

408 tensor_id: TensorId 

409 """tensor id of the reference axis""" 

410 

411 axis_id: AxisId 

412 """axis id of the reference axis""" 

413 

414 offset: StrictInt = 0 

415 

416 def get_size( 

417 self, 

418 axis: ChannelAxis 

419 | IndexInputAxis 

420 | IndexOutputAxis 

421 | TimeInputAxis 

422 | SpaceInputAxis 

423 | TimeOutputAxis 

424 | TimeOutputAxisWithHalo 

425 | SpaceOutputAxis 

426 | SpaceOutputAxisWithHalo, 

427 ref_axis: ChannelAxis 

428 | IndexInputAxis 

429 | IndexOutputAxis 

430 | TimeInputAxis 

431 | SpaceInputAxis 

432 | TimeOutputAxis 

433 | TimeOutputAxisWithHalo 

434 | SpaceOutputAxis 

435 | SpaceOutputAxisWithHalo, 

436 n: ParameterizedSize_N = 0, 

437 ref_size: int | None = None, 

438 ): 

439 """Compute the concrete size for a given axis and its reference axis. 

440 

441 Args: 

442 axis: The axis this [SizeReference][] is the size of. 

443 ref_axis: The reference axis to compute the size from. 

444 n: If the **ref_axis** is parameterized (of type `ParameterizedSize`) 

445 and no fixed **ref_size** is given, 

446 **n** is used to compute the size of the parameterized **ref_axis**. 

447 ref_size: Overwrite the reference size instead of deriving it from 

448 **ref_axis** 

449 (**ref_axis.scale** is still used; any given **n** is ignored). 

450 """ 

451 assert axis.size == self, ( 

452 "Given `axis.size` is not defined by this `SizeReference`" 

453 ) 

454 

455 assert ref_axis.id == self.axis_id, ( 

456 f"Expected `ref_axis.id` to be {self.axis_id}, but got {ref_axis.id}." 

457 ) 

458 

459 assert axis.unit == ref_axis.unit, ( 

460 "`SizeReference` requires `axis` and `ref_axis` to have the same `unit`," 

461 f" but {axis.unit}!={ref_axis.unit}" 

462 ) 

463 if ref_size is None: 

464 if isinstance(ref_axis.size, (int, float)): 

465 ref_size = ref_axis.size 

466 elif isinstance(ref_axis.size, ParameterizedSize): 

467 ref_size = ref_axis.size.get_size(n) 

468 elif isinstance(ref_axis.size, DataDependentSize): 

469 raise ValueError( 

470 "Reference axis referenced in `SizeReference` may not be a `DataDependentSize`." 

471 ) 

472 elif isinstance(ref_axis.size, SizeReference): 

473 raise ValueError( 

474 "Reference axis referenced in `SizeReference` may not be sized by a" 

475 + " `SizeReference` itself." 

476 ) 

477 else: 

478 assert_never(ref_axis.size) 

479 

480 return int(ref_size * ref_axis.scale / axis.scale + self.offset) 

481 

482 @staticmethod 

483 def _get_unit( 

484 axis: ChannelAxis 

485 | IndexInputAxis 

486 | IndexOutputAxis 

487 | TimeInputAxis 

488 | SpaceInputAxis 

489 | TimeOutputAxis 

490 | TimeOutputAxisWithHalo 

491 | SpaceOutputAxis 

492 | SpaceOutputAxisWithHalo, 

493 ): 

494 return axis.unit 

495 

496 

497class AxisBase(NodeWithExplicitlySetFields): 

498 id: AxisId 

499 """An axis id unique across all axes of one tensor.""" 

500 

501 description: Annotated[str, MaxLen(128)] = "" 

502 """A short description of this axis beyond its type and id.""" 

503 

504 

505class WithHalo(Node): 

506 halo: Annotated[int, Ge(1)] 

507 """The halo should be cropped from the output tensor to avoid boundary effects. 

508 It is to be cropped from both sides, i.e. `size_after_crop = size - 2 * halo`. 

509 To document a halo that is already cropped by the model use `size.offset` instead.""" 

510 

511 size: Annotated[ 

512 SizeReference, 

513 Field(examples=[{"tensor_id": "t", "axis_id": "a", "offset": 5}]), 

514 ] 

515 """reference to another axis with an optional offset (see [SizeReference][])""" 

516 

517 

518BATCH_AXIS_ID = AxisId("batch") 

519CHANNEL_AXIS_ID = AxisId("channel") 

520DEFAULT_SPACE_AXIS_ID = AxisId("x") 

521DEFAULT_INDEX_AXIS_ID = AxisId("index") 

522DEFAULT_TIME_AXIS_ID = AxisId("time") 

523 

524 

525class BatchAxis(AxisBase): 

526 implemented_type: ClassVar[Literal["batch"]] = "batch" 

527 if TYPE_CHECKING: 

528 type: Literal["batch"] = "batch" 

529 else: 

530 type: Literal["batch"] 

531 

532 id: Annotated[AxisId, Predicate(_is_batch)] = BATCH_AXIS_ID 

533 size: Literal[1] | None = None 

534 """The batch size may be fixed to 1, 

535 otherwise (the default) it may be chosen arbitrarily depending on available memory""" 

536 

537 @property 

538 def scale(self): 

539 return 1.0 

540 

541 @property 

542 def concatenable(self): 

543 return True 

544 

545 @property 

546 def unit(self): 

547 return None 

548 

549 

550class ChannelAxis(AxisBase): 

551 implemented_type: ClassVar[Literal["channel"]] = "channel" 

552 if TYPE_CHECKING: 

553 type: Literal["channel"] = "channel" 

554 else: 

555 type: Literal["channel"] 

556 

557 id: NonBatchAxisId = CHANNEL_AXIS_ID 

558 

559 channel_names: NotEmpty[list[str]] 

560 """Name/label for each channel. The number of channels is given by `len(channel_names)`.""" 

561 

562 @property 

563 def size(self) -> int: 

564 return len(self.channel_names) 

565 

566 @property 

567 def concatenable(self): 

568 return False 

569 

570 @property 

571 def scale(self) -> float: 

572 return 1.0 

573 

574 @property 

575 def unit(self): 

576 return None 

577 

578 

579class _WithInputAxisSize(Node): 

580 size: Annotated[ 

581 Annotated[int, Gt(0)] | ParameterizedSize | SizeReference, 

582 Field( 

583 examples=[ 

584 10, 

585 ParameterizedSize(min=32, step=16).model_dump(mode="json"), 

586 {"tensor_id": "t", "axis_id": "a", "offset": 5}, 

587 ] 

588 ), 

589 ] 

590 """The size/length of this axis can be specified as 

591 - fixed integer 

592 - parameterized series of valid sizes ([ParameterizedSize][]) 

593 - reference to another axis with an optional offset ([SizeReference][]) 

594 """ 

595 

596 

597class IndexAxisBase(AxisBase): 

598 implemented_type: ClassVar[Literal["index"]] = "index" 

599 if TYPE_CHECKING: 

600 type: Literal["index"] = "index" 

601 else: 

602 type: Literal["index"] 

603 

604 id: NonBatchAxisId = DEFAULT_INDEX_AXIS_ID 

605 

606 @property 

607 def scale(self) -> float: 

608 return 1.0 

609 

610 @property 

611 def unit(self): 

612 return None 

613 

614 

615class IndexInputAxis(IndexAxisBase, _WithInputAxisSize): 

616 concatenable: bool = False 

617 """If a model has a `concatenable` input axis, it can be processed blockwise, 

618 splitting a longer sample axis into blocks matching its input tensor description. 

619 Output axes are concatenable if they have a [SizeReference][] to a concatenable 

620 input axis. 

621 """ 

622 

623 

624class IndexOutputAxis(IndexAxisBase): 

625 size: Annotated[ 

626 Annotated[int, Gt(0)] | SizeReference | DataDependentSize, 

627 Field(examples=[10, {"tensor_id": "t", "axis_id": "a", "offset": 5}]), 

628 ] 

629 """The size/length of this axis can be specified as 

630 - fixed integer 

631 - reference to another axis with an optional offset ([SizeReference][]) 

632 - data dependent size using [DataDependentSize][] (size is only known after model inference) 

633 """ 

634 

635 

636class TimeAxisBase(AxisBase): 

637 implemented_type: ClassVar[Literal["time"]] = "time" 

638 if TYPE_CHECKING: 

639 type: Literal["time"] = "time" 

640 else: 

641 type: Literal["time"] 

642 

643 id: NonBatchAxisId = DEFAULT_TIME_AXIS_ID 

644 unit: TimeUnit | None = None 

645 scale: Annotated[float, Gt(0)] = 1.0 

646 

647 

648class TimeInputAxis(TimeAxisBase, _WithInputAxisSize): 

649 concatenable: bool = False 

650 """If a model has a `concatenable` input axis, it can be processed blockwise, 

651 splitting a longer sample axis into blocks matching its input tensor description. 

652 Output axes are concatenable if they have a [SizeReference][] to a concatenable 

653 input axis. 

654 """ 

655 

656 

657class SpaceAxisBase(AxisBase): 

658 implemented_type: ClassVar[Literal["space"]] = "space" 

659 if TYPE_CHECKING: 

660 type: Literal["space"] = "space" 

661 else: 

662 type: Literal["space"] 

663 

664 id: Annotated[NonBatchAxisId, Field(examples=["x", "y", "z"])] = ( 

665 DEFAULT_SPACE_AXIS_ID 

666 ) 

667 unit: SpaceUnit | None = None 

668 scale: Annotated[float, Gt(0)] = 1.0 

669 

670 

671class SpaceInputAxis(SpaceAxisBase, _WithInputAxisSize): 

672 concatenable: bool = False 

673 """If a model has a `concatenable` input axis, it can be processed blockwise, 

674 splitting a longer sample axis into blocks matching its input tensor description. 

675 Output axes are concatenable if they have a [SizeReference][] to a concatenable 

676 input axis. 

677 """ 

678 

679 

680INPUT_AXIS_TYPES = ( 

681 BatchAxis, 

682 ChannelAxis, 

683 IndexInputAxis, 

684 TimeInputAxis, 

685 SpaceInputAxis, 

686) 

687"""intended for isinstance comparisons in py<3.10""" 

688 

689_InputAxisUnion = Union[ 

690 BatchAxis, ChannelAxis, IndexInputAxis, TimeInputAxis, SpaceInputAxis 

691] 

692InputAxis = Annotated[_InputAxisUnion, Discriminator("type")] 

693 

694 

695class _WithOutputAxisSize(Node): 

696 size: Annotated[ 

697 Annotated[int, Gt(0)] | SizeReference, 

698 Field(examples=[10, {"tensor_id": "t", "axis_id": "a", "offset": 5}]), 

699 ] 

700 """The size/length of this axis can be specified as 

701 - fixed integer 

702 - reference to another axis with an optional offset (see [SizeReference][]) 

703 """ 

704 

705 

706class TimeOutputAxis(TimeAxisBase, _WithOutputAxisSize): 

707 pass 

708 

709 

710class TimeOutputAxisWithHalo(TimeAxisBase, WithHalo): 

711 pass 

712 

713 

714def _get_halo_axis_discriminator_value(v: Any) -> Literal["with_halo", "wo_halo"]: 

715 if isinstance(v, dict): 

716 return "with_halo" if "halo" in v else "wo_halo" 

717 else: 

718 return "with_halo" if hasattr(v, "halo") else "wo_halo" 

719 

720 

721_TimeOutputAxisUnion = Annotated[ 

722 Union[ 

723 Annotated[TimeOutputAxis, Tag("wo_halo")], 

724 Annotated[TimeOutputAxisWithHalo, Tag("with_halo")], 

725 ], 

726 Discriminator(_get_halo_axis_discriminator_value), 

727] 

728 

729 

730class SpaceOutputAxis(SpaceAxisBase, _WithOutputAxisSize): 

731 pass 

732 

733 

734class SpaceOutputAxisWithHalo(SpaceAxisBase, WithHalo): 

735 pass 

736 

737 

738_SpaceOutputAxisUnion = Annotated[ 

739 Union[ 

740 Annotated[SpaceOutputAxis, Tag("wo_halo")], 

741 Annotated[SpaceOutputAxisWithHalo, Tag("with_halo")], 

742 ], 

743 Discriminator(_get_halo_axis_discriminator_value), 

744] 

745 

746 

747_OutputAxisUnion = Union[ 

748 BatchAxis, ChannelAxis, IndexOutputAxis, _TimeOutputAxisUnion, _SpaceOutputAxisUnion 

749] 

750OutputAxis = Annotated[_OutputAxisUnion, Discriminator("type")] 

751 

752OUTPUT_AXIS_TYPES = ( 

753 BatchAxis, 

754 ChannelAxis, 

755 IndexOutputAxis, 

756 TimeOutputAxis, 

757 TimeOutputAxisWithHalo, 

758 SpaceOutputAxis, 

759 SpaceOutputAxisWithHalo, 

760) 

761"""intended for isinstance comparisons in py<3.10""" 

762 

763 

764AnyAxis = Union[InputAxis, OutputAxis] 

765 

766ANY_AXIS_TYPES = INPUT_AXIS_TYPES + OUTPUT_AXIS_TYPES 

767"""intended for isinstance comparisons in py<3.10""" 

768 

769TVs = Union[ 

770 NotEmpty[List[int]], 

771 NotEmpty[List[float]], 

772 NotEmpty[List[bool]], 

773 NotEmpty[List[str]], 

774] 

775 

776 

777NominalOrOrdinalDType = Literal[ 

778 "float32", 

779 "float64", 

780 "uint8", 

781 "int8", 

782 "uint16", 

783 "int16", 

784 "uint32", 

785 "int32", 

786 "uint64", 

787 "int64", 

788 "bool", 

789] 

790 

791 

792class NominalOrOrdinalDataDescr(Node): 

793 values: TVs 

794 """A fixed set of nominal or an ascending sequence of ordinal values. 

795 In this case `data.type` is required to be an unsigend integer type, e.g. 'uint8'. 

796 String `values` are interpreted as labels for tensor values 0, ..., N. 

797 Note: as YAML 1.2 does not natively support a "set" datatype, 

798 nominal values should be given as a sequence (aka list/array) as well. 

799 """ 

800 

801 type: Annotated[ 

802 NominalOrOrdinalDType, 

803 Field( 

804 examples=[ 

805 "float32", 

806 "uint8", 

807 "uint16", 

808 "int64", 

809 "bool", 

810 ], 

811 ), 

812 ] = "uint8" 

813 

814 @model_validator(mode="after") 

815 def _validate_values_match_type( 

816 self, 

817 ) -> Self: 

818 incompatible: list[Any] = [] 

819 for v in self.values: 

820 if self.type == "bool": 

821 if not isinstance(v, bool): 

822 incompatible.append(v) 

823 elif self.type in DTYPE_LIMITS: 

824 if ( 

825 isinstance(v, (int, float)) 

826 and ( 

827 v < DTYPE_LIMITS[self.type].min 

828 or v > DTYPE_LIMITS[self.type].max 

829 ) 

830 or (isinstance(v, str) and "uint" not in self.type) 

831 or (isinstance(v, float) and "int" in self.type) 

832 ): 

833 incompatible.append(v) 

834 else: 

835 incompatible.append(v) 

836 

837 if len(incompatible) == 5: 

838 incompatible.append("...") 

839 break 

840 

841 if incompatible: 

842 raise ValueError( 

843 f"data type '{self.type}' incompatible with values {incompatible}" 

844 ) 

845 

846 return self 

847 

848 unit: Literal["arbitrary unit"] | SiUnit | None = None 

849 

850 @property 

851 def range(self): 

852 if isinstance(self.values[0], str): 

853 return 0, len(self.values) - 1 

854 else: 

855 return min(self.values), max(self.values) 

856 

857 

858IntervalOrRatioDType = Literal[ 

859 "float32", 

860 "float64", 

861 "uint8", 

862 "int8", 

863 "uint16", 

864 "int16", 

865 "uint32", 

866 "int32", 

867 "uint64", 

868 "int64", 

869] 

870 

871 

872class IntervalOrRatioDataDescr(Node): 

873 type: Annotated[ # TODO: rename to dtype 

874 IntervalOrRatioDType, 

875 Field( 

876 examples=["float32", "float64", "uint8", "uint16"], 

877 ), 

878 ] = "float32" 

879 range: tuple[float | None, float | None] = ( 

880 None, 

881 None, 

882 ) 

883 """Tuple `(minimum, maximum)` specifying the allowed range of the data in this tensor. 

884 `None` corresponds to min/max of what can be expressed by **type**.""" 

885 unit: Literal["arbitrary unit"] | SiUnit = "arbitrary unit" 

886 scale: float = 1.0 

887 """Scale for data on an interval (or ratio) scale.""" 

888 offset: float | None = None 

889 """Offset for data on a ratio scale.""" 

890 

891 @model_validator(mode="before") 

892 def _replace_inf(cls, data: Any): 

893 if is_dict(data) and "range" in data and is_sequence(data["range"]): 

894 forbidden = ( 

895 "inf", 

896 "-inf", 

897 ".inf", 

898 "-.inf", 

899 float("inf"), 

900 float("-inf"), 

901 ) 

902 if any(v in forbidden for v in data["range"]): 

903 issue_warning("replaced 'inf' value", value=data["range"]) 

904 

905 data["range"] = tuple( 

906 (None if v in forbidden else v) for v in data["range"] 

907 ) 

908 

909 return data 

910 

911 

912TensorDataDescr = Union[NominalOrOrdinalDataDescr, IntervalOrRatioDataDescr] 

913 

914 

915class BinarizeKwargs(KwargsNode): 

916 """key word arguments for [BinarizeDescr][]""" 

917 

918 threshold: float 

919 """The fixed threshold""" 

920 

921 

922class BinarizeAlongAxisKwargs(KwargsNode): 

923 """key word arguments for [BinarizeDescr][]""" 

924 

925 threshold: NotEmpty[list[float]] 

926 """The fixed threshold values along `axis`""" 

927 

928 axis: Annotated[NonBatchAxisId, Field(examples=["channel"])] 

929 """The `threshold` axis""" 

930 

931 

932class BinarizeDescr(NodeWithExplicitlySetFields): 

933 """Binarize the tensor with a fixed threshold. 

934 

935 Values above [BinarizeKwargs.threshold][]/[BinarizeAlongAxisKwargs.threshold][] 

936 will be set to one, values below the threshold to zero. 

937 

938 Examples: 

939 - in YAML 

940 ```yaml 

941 postprocessing: 

942 - id: binarize 

943 kwargs: 

944 axis: 'channel' 

945 threshold: [0.25, 0.5, 0.75] 

946 ``` 

947 - in Python: 

948 

949 >>> postprocessing = [BinarizeDescr( 

950 ... kwargs=BinarizeAlongAxisKwargs( 

951 ... axis=AxisId('channel'), 

952 ... threshold=[0.25, 0.5, 0.75], 

953 ... ) 

954 ... )] 

955 """ 

956 

957 implemented_id: ClassVar[Literal["binarize"]] = "binarize" 

958 if TYPE_CHECKING: 

959 id: Literal["binarize"] = "binarize" 

960 else: 

961 id: Literal["binarize"] 

962 kwargs: BinarizeKwargs | BinarizeAlongAxisKwargs 

963 

964 

965class ClipKwargs(KwargsNode): 

966 """key word arguments for [ClipDescr][]""" 

967 

968 min: float | None = None 

969 """Minimum value for clipping. 

970 

971 Exclusive with [min_percentile][] 

972 """ 

973 min_percentile: Annotated[float, Interval(ge=0, lt=100)] | None = None 

974 """Minimum percentile for clipping. 

975 

976 Exclusive with [min][]. 

977 

978 In range [0, 100). 

979 """ 

980 

981 max: float | None = None 

982 """Maximum value for clipping. 

983 

984 Exclusive with `max_percentile`. 

985 """ 

986 max_percentile: Annotated[float, Interval(gt=1, le=100)] | None = None 

987 """Maximum percentile for clipping. 

988 

989 Exclusive with `max`. 

990 

991 In range (1, 100]. 

992 """ 

993 

994 axes: Annotated[Sequence[AxisId] | None, Field(examples=[("batch", "x", "y")])] = ( 

995 None 

996 ) 

997 """The subset of axes to determine percentiles jointly, 

998 

999 i.e. axes to reduce to compute min/max from `min_percentile`/`max_percentile`. 

1000 For example to clip 'batch', 'x' and 'y' jointly in a tensor ('batch', 'channel', 'y', 'x') 

1001 resulting in a tensor of equal shape with clipped values per channel, specify `axes=('batch', 'x', 'y')`. 

1002 To clip samples independently, leave out the 'batch' axis. 

1003 

1004 Only valid if `min_percentile` and/or `max_percentile` are set. 

1005 

1006 Default: Compute percentiles over all axes jointly.""" 

1007 

1008 @model_validator(mode="after") 

1009 def _validate(self) -> Self: 

1010 if (self.min is not None) and (self.min_percentile is not None): 

1011 raise ValueError( 

1012 "Only one of `min` and `min_percentile` may be set, not both." 

1013 ) 

1014 if (self.max is not None) and (self.max_percentile is not None): 

1015 raise ValueError( 

1016 "Only one of `max` and `max_percentile` may be set, not both." 

1017 ) 

1018 if ( 

1019 self.min is None 

1020 and self.min_percentile is None 

1021 and self.max is None 

1022 and self.max_percentile is None 

1023 ): 

1024 raise ValueError( 

1025 "At least one of `min`, `min_percentile`, `max`, or `max_percentile` must be set." 

1026 ) 

1027 

1028 if ( 

1029 self.axes is not None 

1030 and self.min_percentile is None 

1031 and self.max_percentile is None 

1032 ): 

1033 raise ValueError( 

1034 "If `axes` is set, at least one of `min_percentile` or `max_percentile` must be set." 

1035 ) 

1036 

1037 return self 

1038 

1039 

1040class ClipDescr(NodeWithExplicitlySetFields): 

1041 """Set tensor values below min to min and above max to max. 

1042 

1043 See `ScaleRangeDescr` for examples. 

1044 """ 

1045 

1046 implemented_id: ClassVar[Literal["clip"]] = "clip" 

1047 if TYPE_CHECKING: 

1048 id: Literal["clip"] = "clip" 

1049 else: 

1050 id: Literal["clip"] 

1051 

1052 kwargs: ClipKwargs 

1053 

1054 

1055class EnsureDtypeKwargs(KwargsNode): 

1056 """key word arguments for [EnsureDtypeDescr][]""" 

1057 

1058 dtype: Literal[ 

1059 "float32", 

1060 "float64", 

1061 "uint8", 

1062 "int8", 

1063 "uint16", 

1064 "int16", 

1065 "uint32", 

1066 "int32", 

1067 "uint64", 

1068 "int64", 

1069 "bool", 

1070 ] 

1071 

1072 

1073class EnsureDtypeDescr(NodeWithExplicitlySetFields): 

1074 """Cast the tensor data type to `EnsureDtypeKwargs.dtype` (if not matching). 

1075 

1076 This can for example be used to ensure the inner neural network model gets a 

1077 different input tensor data type than the fully described bioimage.io model does. 

1078 

1079 Examples: 

1080 The described bioimage.io model (incl. preprocessing) accepts any 

1081 float32-compatible tensor, normalizes it with percentiles and clipping and then 

1082 casts it to uint8, which is what the neural network in this example expects. 

1083 - in YAML 

1084 ```yaml 

1085 inputs: 

1086 - data: 

1087 type: float32 # described bioimage.io model is compatible with any float32 input tensor 

1088 preprocessing: 

1089 - id: scale_range 

1090 kwargs: 

1091 axes: ['y', 'x'] 

1092 max_percentile: 99.8 

1093 min_percentile: 5.0 

1094 - id: clip 

1095 kwargs: 

1096 min: 0.0 

1097 max: 1.0 

1098 - id: ensure_dtype # the neural network of the model requires uint8 

1099 kwargs: 

1100 dtype: uint8 

1101 ``` 

1102 - in Python: 

1103 >>> preprocessing = [ 

1104 ... ScaleRangeDescr( 

1105 ... kwargs=ScaleRangeKwargs( 

1106 ... axes= (AxisId('y'), AxisId('x')), 

1107 ... max_percentile= 99.8, 

1108 ... min_percentile= 5.0, 

1109 ... ) 

1110 ... ), 

1111 ... ClipDescr(kwargs=ClipKwargs(min=0.0, max=1.0)), 

1112 ... EnsureDtypeDescr(kwargs=EnsureDtypeKwargs(dtype="uint8")), 

1113 ... ] 

1114 """ 

1115 

1116 implemented_id: ClassVar[Literal["ensure_dtype"]] = "ensure_dtype" 

1117 if TYPE_CHECKING: 

1118 id: Literal["ensure_dtype"] = "ensure_dtype" 

1119 else: 

1120 id: Literal["ensure_dtype"] 

1121 

1122 kwargs: EnsureDtypeKwargs 

1123 

1124 

1125class ScaleLinearKwargs(KwargsNode): 

1126 """Key word arguments for [ScaleLinearDescr][]""" 

1127 

1128 gain: float = 1.0 

1129 """multiplicative factor""" 

1130 

1131 offset: float = 0.0 

1132 """additive term""" 

1133 

1134 @model_validator(mode="after") 

1135 def _validate(self) -> Self: 

1136 if self.gain == 1.0 and self.offset == 0.0: 

1137 raise ValueError( 

1138 "Redundant linear scaling not allowd. Set `gain` != 1.0 and/or `offset`" 

1139 + " != 0.0." 

1140 ) 

1141 

1142 return self 

1143 

1144 

1145class ScaleLinearAlongAxisKwargs(KwargsNode): 

1146 """Key word arguments for [ScaleLinearDescr][]""" 

1147 

1148 axis: Annotated[NonBatchAxisId, Field(examples=["channel"])] 

1149 """The axis of gain and offset values.""" 

1150 

1151 gain: float | NotEmpty[list[float]] = 1.0 

1152 """multiplicative factor""" 

1153 

1154 offset: float | NotEmpty[list[float]] = 0.0 

1155 """additive term""" 

1156 

1157 @model_validator(mode="after") 

1158 def _validate(self) -> Self: 

1159 if isinstance(self.gain, list): 

1160 if isinstance(self.offset, list): 

1161 if len(self.gain) != len(self.offset): 

1162 raise ValueError( 

1163 f"Size of `gain` ({len(self.gain)}) and `offset` ({len(self.offset)}) must match." 

1164 ) 

1165 else: 

1166 self.offset = [float(self.offset)] * len(self.gain) 

1167 elif isinstance(self.offset, list): 

1168 self.gain = [float(self.gain)] * len(self.offset) 

1169 else: 

1170 raise ValueError( 

1171 "Do not specify an `axis` for scalar gain and offset values." 

1172 ) 

1173 

1174 if all(g == 1.0 for g in self.gain) and all(off == 0.0 for off in self.offset): 

1175 raise ValueError( 

1176 "Redundant linear scaling not allowd. Set `gain` != 1.0 and/or `offset`" 

1177 + " != 0.0." 

1178 ) 

1179 

1180 return self 

1181 

1182 

1183class ScaleLinearDescr(NodeWithExplicitlySetFields): 

1184 """Fixed linear scaling. 

1185 

1186 Examples: 

1187 1. Scale with scalar gain and offset 

1188 - in YAML 

1189 ```yaml 

1190 preprocessing: 

1191 - id: scale_linear 

1192 kwargs: 

1193 gain: 2.0 

1194 offset: 3.0 

1195 ``` 

1196 - in Python: 

1197 

1198 >>> preprocessing = [ 

1199 ... ScaleLinearDescr(kwargs=ScaleLinearKwargs(gain= 2.0, offset=3.0)) 

1200 ... ] 

1201 

1202 2. Independent scaling along an axis 

1203 - in YAML 

1204 ```yaml 

1205 preprocessing: 

1206 - id: scale_linear 

1207 kwargs: 

1208 axis: 'channel' 

1209 gain: [1.0, 2.0, 3.0] 

1210 ``` 

1211 - in Python: 

1212 

1213 >>> preprocessing = [ 

1214 ... ScaleLinearDescr( 

1215 ... kwargs=ScaleLinearAlongAxisKwargs( 

1216 ... axis=AxisId("channel"), 

1217 ... gain=[1.0, 2.0, 3.0], 

1218 ... ) 

1219 ... ) 

1220 ... ] 

1221 

1222 """ 

1223 

1224 implemented_id: ClassVar[Literal["scale_linear"]] = "scale_linear" 

1225 if TYPE_CHECKING: 

1226 id: Literal["scale_linear"] = "scale_linear" 

1227 else: 

1228 id: Literal["scale_linear"] 

1229 kwargs: ScaleLinearKwargs | ScaleLinearAlongAxisKwargs 

1230 

1231 

1232class SigmoidDescr(NodeWithExplicitlySetFields): 

1233 """The logistic sigmoid function, a.k.a. expit function. 

1234 

1235 Examples: 

1236 - in YAML 

1237 ```yaml 

1238 postprocessing: 

1239 - id: sigmoid 

1240 ``` 

1241 - in Python: 

1242 

1243 >>> postprocessing = [SigmoidDescr()] 

1244 """ 

1245 

1246 implemented_id: ClassVar[Literal["sigmoid"]] = "sigmoid" 

1247 if TYPE_CHECKING: 

1248 id: Literal["sigmoid"] = "sigmoid" 

1249 else: 

1250 id: Literal["sigmoid"] 

1251 

1252 @property 

1253 def kwargs(self) -> KwargsNode: 

1254 """empty kwargs""" 

1255 return KwargsNode() 

1256 

1257 

1258class SoftmaxKwargs(KwargsNode): 

1259 """key word arguments for [SoftmaxDescr][]""" 

1260 

1261 axis: Annotated[NonBatchAxisId, Field(examples=["channel"])] = CHANNEL_AXIS_ID 

1262 """The axis to apply the softmax function along. 

1263 Note: 

1264 Defaults to 'channel' axis 

1265 (which may not exist, in which case 

1266 a different axis id has to be specified). 

1267 """ 

1268 

1269 

1270class SoftmaxDescr(NodeWithExplicitlySetFields): 

1271 """The softmax function. 

1272 

1273 Examples: 

1274 - in YAML 

1275 ```yaml 

1276 postprocessing: 

1277 - id: softmax 

1278 kwargs: 

1279 axis: channel 

1280 ``` 

1281 - in Python: 

1282 

1283 >>> postprocessing = [SoftmaxDescr(kwargs=SoftmaxKwargs(axis=AxisId("channel")))] 

1284 """ 

1285 

1286 implemented_id: ClassVar[Literal["softmax"]] = "softmax" 

1287 if TYPE_CHECKING: 

1288 id: Literal["softmax"] = "softmax" 

1289 else: 

1290 id: Literal["softmax"] 

1291 

1292 kwargs: SoftmaxKwargs = Field(default_factory=SoftmaxKwargs.model_construct) 

1293 

1294 

1295class _StardistPostprocessingKwargsBase(KwargsNode): 

1296 """key word arguments for [StardistPostprocessingDescr][]""" 

1297 

1298 prob_threshold: float 

1299 """The probability threshold for object candidate selection.""" 

1300 

1301 nms_threshold: float 

1302 """The IoU threshold for non-maximum suppression.""" 

1303 

1304 n_rays: int 

1305 """Number of radial lines (rays) cast from the center of an object to its boundary.""" 

1306 

1307 

1308class StardistPostprocessingKwargs2D(_StardistPostprocessingKwargsBase): 

1309 grid: tuple[int, int] 

1310 """Grid size of network predictions.""" 

1311 

1312 b: int | tuple[tuple[int, int], tuple[int, int]] 

1313 """Border region in which object probability is set to zero.""" 

1314 

1315 

1316class StardistPostprocessingKwargs3D(_StardistPostprocessingKwargsBase): 

1317 grid: tuple[int, int, int] 

1318 """Grid size of network predictions.""" 

1319 

1320 b: int | tuple[tuple[int, int], tuple[int, int], tuple[int, int]] 

1321 """Border region in which object probability is set to zero.""" 

1322 

1323 anisotropy: tuple[float, float, float] 

1324 """Anisotropy factors for 3D star-convex polyhedra, i.e. the physical pixel size along each spatial axis.""" 

1325 

1326 overlap_label: int | None = None 

1327 """Optional label to apply to any area of overlapping predicted objects.""" 

1328 

1329 

1330class StardistPostprocessingDescr(NodeWithExplicitlySetFields): 

1331 """Stardist postprocessing including non-maximum suppression and converting polygon representations to instance labels 

1332 

1333 as described in: 

1334 - Uwe Schmidt, Martin Weigert, Coleman Broaddus, and Gene Myers. 

1335 [*Cell Detection with Star-convex Polygons*](https://arxiv.org/abs/1806.03535). 

1336 International Conference on Medical Image Computing and Computer-Assisted Intervention (MICCAI), Granada, Spain, September 2018. 

1337 - Martin Weigert, Uwe Schmidt, Robert Haase, Ko Sugawara, and Gene Myers. 

1338 [*Star-convex Polyhedra for 3D Object Detection and Segmentation in Microscopy*](http://openaccess.thecvf.com/content_WACV_2020/papers/Weigert_Star-convex_Polyhedra_for_3D_Object_Detection_and_Segmentation_in_Microscopy_WACV_2020_paper.pdf). 

1339 The IEEE Winter Conference on Applications of Computer Vision (WACV), Snowmass Village, Colorado, March 2020. 

1340 

1341 Note: Only available if the `stardist` package is installed. 

1342 """ 

1343 

1344 implemented_id: ClassVar[Literal["stardist_postprocessing"]] = ( 

1345 "stardist_postprocessing" 

1346 ) 

1347 if TYPE_CHECKING: 

1348 id: Literal["stardist_postprocessing"] = "stardist_postprocessing" 

1349 else: 

1350 id: Literal["stardist_postprocessing"] 

1351 

1352 kwargs: StardistPostprocessingKwargs2D | StardistPostprocessingKwargs3D 

1353 

1354 

1355class CellposeFlowDynamicsKwargs(KwargsNode): 

1356 """key word arguments for [CellposeFlowDynamicsDescr][]""" 

1357 

1358 cellprob_threshold: float 

1359 flow_threshold: float 

1360 do_3D: bool 

1361 min_size: int = 15 

1362 """Minimum size of objects to keep, in pixels. Default is 15, which is the default in Cellpose. Set to 0 to disable filtering by size.""" 

1363 output_dtype: Literal["uint16", "uint32"] = "uint16" 

1364 

1365 

1366class CellposeFlowDynamicsDescr(NodeWithExplicitlySetFields): 

1367 """Cellpose flow dynamics postprocessing as described in: 

1368 - Carsen Stringer and Marius Pachitariu. [*Cellpose: a generalist algorithm for cellular segmentation*](https://www.nature.com/articles/s41592-020-01018-x). Nature Methods, 2021. 

1369 

1370 Note: Only available if the `cellpose` package is installed. 

1371 """ 

1372 

1373 implemented_id: ClassVar[Literal["cellpose_flow_dynamics"]] = ( 

1374 "cellpose_flow_dynamics" 

1375 ) 

1376 if TYPE_CHECKING: 

1377 id: Literal["cellpose_flow_dynamics"] = "cellpose_flow_dynamics" 

1378 else: 

1379 id: Literal["cellpose_flow_dynamics"] 

1380 

1381 kwargs: CellposeFlowDynamicsKwargs 

1382 

1383 

1384class CustomProcessingDescr(NodeWithExplicitlySetFields, FileDescr): 

1385 """Custom (post)processing op — source file shipped inline with the model. 

1386 

1387 Supports (post)processing that cannot be expressed by the built-in named 

1388 operations (watershed, connected components, etc.) 

1389 using a simple Python callable interface. 

1390 

1391 The op is implemented in a ``.py`` file packaged alongside the model weights. 

1392 Two styles are supported: 

1393 

1394 *Callable class* — kwargs go to ``__init__``, tensors arrive in ``__call__``: 

1395 

1396 .. code-block:: python 

1397 

1398 # my_postprocess.py 

1399 import numpy as np 

1400 

1401 class my_postprocess: 

1402 def __init__(self, threshold: float = 0.5) -> None: 

1403 self.threshold = threshold 

1404 def __call__(self, *arrays: np.ndarray) -> np.ndarray: 

1405 # arrays = model output tensors in rdf.yaml declaration order 

1406 return (arrays[0] > self.threshold).astype(np.uint8) 

1407 

1408 *Factory function* — alternative closure style, identical runtime behaviour: 

1409 

1410 .. code-block:: python 

1411 

1412 # my_postprocess.py 

1413 import numpy as np 

1414 

1415 def my_postprocess(threshold: float = 0.5): 

1416 def run(*arrays: np.ndarray) -> np.ndarray: 

1417 return (arrays[0] > threshold).astype(np.uint8) 

1418 return run 

1419 

1420 Reference it in ``rdf.yaml`` with the source file included in the package: 

1421 

1422 .. code-block:: yaml 

1423 

1424 postprocessing: 

1425 - id: custom 

1426 callable: my_postprocess # class or function name in source 

1427 source: my_postprocess.py # packaged alongside weights 

1428 sha256: <hash> # sha256 of the source file 

1429 kwargs: # forwarded to __init__ / factory 

1430 threshold: 0.5 

1431 

1432 **Security:** source files are SHA-256 verified before execution. 

1433 Execution requires explicit opt-in in bioimageio.core and curator 

1434 review before Zoo publication. 

1435 """ 

1436 

1437 implemented_id: ClassVar[Literal["custom"]] = "custom" 

1438 if TYPE_CHECKING: 

1439 id: Literal["custom"] = "custom" 

1440 else: 

1441 id: Literal["custom"] 

1442 

1443 callable: Annotated[ 

1444 str, 

1445 Field(examples=["my_postprocess_factory", "MyPostprocessClass"]), 

1446 ] 

1447 """Name of the callable class or factory function defined in ``source``. 

1448 

1449 At runtime: ``op = callable(**kwargs)``, then ``result = op(*output_tensors)`` 

1450 per image. Both a class with ``__call__`` and a factory function returning 

1451 a callable satisfy this protocol.""" 

1452 

1453 source: Annotated[FileSource, AfterValidator(wo_special_file_name)] 

1454 """Python source file (included when packaging the model).""" 

1455 

1456 kwargs: dict[str, YamlValue] = Field( 

1457 default_factory=cast(Callable[[], Dict[str, YamlValue]], dict) 

1458 ) 

1459 """Keyword arguments forwarded to the callable (``__init__`` or factory).""" 

1460 

1461 @model_serializer(mode="wrap", when_used="unless-none") 

1462 def _serialize( 

1463 self, nxt: SerializerFunctionWrapHandler, info: SerializationInfo 

1464 ) -> dict[str, YamlValue]: 

1465 return package_file_descr_serializer(self, nxt, info) 

1466 

1467 

1468class FixedZeroMeanUnitVarianceKwargs(KwargsNode): 

1469 """key word arguments for [FixedZeroMeanUnitVarianceDescr][]""" 

1470 

1471 mean: float 

1472 """The mean value to normalize with.""" 

1473 

1474 std: Annotated[float, Ge(1e-6)] 

1475 """The standard deviation value to normalize with.""" 

1476 

1477 

1478class FixedZeroMeanUnitVarianceAlongAxisKwargs(KwargsNode): 

1479 """key word arguments for [FixedZeroMeanUnitVarianceDescr][]""" 

1480 

1481 mean: NotEmpty[list[float]] 

1482 """The mean value(s) to normalize with.""" 

1483 

1484 std: NotEmpty[list[Annotated[float, Ge(1e-6)]]] 

1485 """The standard deviation value(s) to normalize with. 

1486 Size must match `mean` values.""" 

1487 

1488 axis: Annotated[NonBatchAxisId, Field(examples=["channel", "index"])] 

1489 """The axis of the mean/std values to normalize each entry along that dimension 

1490 separately.""" 

1491 

1492 @model_validator(mode="after") 

1493 def _mean_and_std_match(self) -> Self: 

1494 if len(self.mean) != len(self.std): 

1495 raise ValueError( 

1496 f"Size of `mean` ({len(self.mean)}) and `std` ({len(self.std)})" 

1497 + " must match." 

1498 ) 

1499 

1500 return self 

1501 

1502 

1503class FixedZeroMeanUnitVarianceDescr(NodeWithExplicitlySetFields): 

1504 """Subtract a given mean and divide by the standard deviation. 

1505 

1506 Normalize with fixed, precomputed values for 

1507 `FixedZeroMeanUnitVarianceKwargs.mean` and `FixedZeroMeanUnitVarianceKwargs.std` 

1508 Use `FixedZeroMeanUnitVarianceAlongAxisKwargs` for independent scaling along given 

1509 axes. 

1510 

1511 Examples: 

1512 1. scalar value for whole tensor 

1513 - in YAML 

1514 ```yaml 

1515 preprocessing: 

1516 - id: fixed_zero_mean_unit_variance 

1517 kwargs: 

1518 mean: 103.5 

1519 std: 13.7 

1520 ``` 

1521 - in Python 

1522 >>> preprocessing = [FixedZeroMeanUnitVarianceDescr( 

1523 ... kwargs=FixedZeroMeanUnitVarianceKwargs(mean=103.5, std=13.7) 

1524 ... )] 

1525 

1526 2. independently along an axis 

1527 - in YAML 

1528 ```yaml 

1529 preprocessing: 

1530 - id: fixed_zero_mean_unit_variance 

1531 kwargs: 

1532 axis: channel 

1533 mean: [101.5, 102.5, 103.5] 

1534 std: [11.7, 12.7, 13.7] 

1535 ``` 

1536 - in Python 

1537 >>> preprocessing = [FixedZeroMeanUnitVarianceDescr( 

1538 ... kwargs=FixedZeroMeanUnitVarianceAlongAxisKwargs( 

1539 ... axis=AxisId("channel"), 

1540 ... mean=[101.5, 102.5, 103.5], 

1541 ... std=[11.7, 12.7, 13.7], 

1542 ... ) 

1543 ... )] 

1544 """ 

1545 

1546 implemented_id: ClassVar[Literal["fixed_zero_mean_unit_variance"]] = ( 

1547 "fixed_zero_mean_unit_variance" 

1548 ) 

1549 if TYPE_CHECKING: 

1550 id: Literal["fixed_zero_mean_unit_variance"] = "fixed_zero_mean_unit_variance" 

1551 else: 

1552 id: Literal["fixed_zero_mean_unit_variance"] 

1553 

1554 kwargs: FixedZeroMeanUnitVarianceKwargs | FixedZeroMeanUnitVarianceAlongAxisKwargs 

1555 

1556 

1557class ZeroMeanUnitVarianceKwargs(KwargsNode): 

1558 """key word arguments for [ZeroMeanUnitVarianceDescr][]""" 

1559 

1560 axes: Annotated[Sequence[AxisId] | None, Field(examples=[("batch", "x", "y")])] = ( 

1561 None 

1562 ) 

1563 """The subset of axes to normalize jointly, i.e. axes to reduce to compute mean/std. 

1564 For example to normalize 'batch', 'x' and 'y' jointly in a tensor ('batch', 'channel', 'y', 'x') 

1565 resulting in a tensor of equal shape normalized per channel, specify `axes=('batch', 'x', 'y')`. 

1566 To normalize each sample independently leave out the 'batch' axis. 

1567 Default: Scale all axes jointly.""" 

1568 

1569 eps: Annotated[float, Interval(gt=0, le=0.1)] = 1e-6 

1570 """epsilon for numeric stability: `out = (tensor - mean) / (std + eps)`.""" 

1571 

1572 

1573class ZeroMeanUnitVarianceDescr(NodeWithExplicitlySetFields): 

1574 """Subtract mean and divide by variance. 

1575 

1576 Examples: 

1577 Subtract tensor mean and variance 

1578 - in YAML 

1579 ```yaml 

1580 preprocessing: 

1581 - id: zero_mean_unit_variance 

1582 ``` 

1583 - in Python 

1584 >>> preprocessing = [ZeroMeanUnitVarianceDescr()] 

1585 """ 

1586 

1587 implemented_id: ClassVar[Literal["zero_mean_unit_variance"]] = ( 

1588 "zero_mean_unit_variance" 

1589 ) 

1590 if TYPE_CHECKING: 

1591 id: Literal["zero_mean_unit_variance"] = "zero_mean_unit_variance" 

1592 else: 

1593 id: Literal["zero_mean_unit_variance"] 

1594 

1595 kwargs: ZeroMeanUnitVarianceKwargs = Field( 

1596 default_factory=ZeroMeanUnitVarianceKwargs.model_construct 

1597 ) 

1598 

1599 

1600class ScaleRangeKwargs(KwargsNode): 

1601 """key word arguments for [ScaleRangeDescr][] 

1602 

1603 For `min_percentile`=0.0 (the default) and `max_percentile`=100 (the default) 

1604 this processing step normalizes data to the [0, 1] intervall. 

1605 For other percentiles the normalized values will partially be outside the [0, 1] 

1606 intervall. Use `ScaleRange` followed by `ClipDescr` if you want to limit the 

1607 normalized values to a range. 

1608 """ 

1609 

1610 axes: Annotated[Sequence[AxisId] | None, Field(examples=[("batch", "x", "y")])] = ( 

1611 None 

1612 ) 

1613 """The subset of axes to normalize jointly, i.e. axes to reduce to compute the min/max percentile value. 

1614 For example to normalize 'batch', 'x' and 'y' jointly in a tensor ('batch', 'channel', 'y', 'x') 

1615 resulting in a tensor of equal shape normalized per channel, specify `axes=('batch', 'x', 'y')`. 

1616 To normalize samples independently, leave out the "batch" axis. 

1617 Default: Scale all axes jointly.""" 

1618 

1619 min_percentile: Annotated[float, Interval(ge=0, lt=100)] = 0.0 

1620 """The lower percentile used to determine the value to align with zero.""" 

1621 

1622 max_percentile: Annotated[float, Interval(gt=1, le=100)] = 100.0 

1623 """The upper percentile used to determine the value to align with one. 

1624 Has to be bigger than `min_percentile`. 

1625 The range is 1 to 100 instead of 0 to 100 to avoid mistakenly 

1626 accepting percentiles specified in the range 0.0 to 1.0.""" 

1627 

1628 eps: Annotated[float, Interval(gt=0, le=0.1)] = 1e-6 

1629 """Epsilon for numeric stability. 

1630 `out = (tensor - v_lower) / (v_upper - v_lower + eps)`; 

1631 with `v_lower,v_upper` values at the respective percentiles.""" 

1632 

1633 reference_tensor: TensorId | None = None 

1634 """ID of the unprocessed input tensor to compute the percentiles from. 

1635 Default: The tensor itself. 

1636 """ 

1637 

1638 @field_validator("max_percentile", mode="after") 

1639 @classmethod 

1640 def min_smaller_max(cls, value: float, info: ValidationInfo) -> float: 

1641 if (min_p := info.data["min_percentile"]) >= value: 

1642 raise ValueError(f"min_percentile {min_p} >= max_percentile {value}") 

1643 

1644 return value 

1645 

1646 

1647class ScaleRangeDescr(NodeWithExplicitlySetFields): 

1648 """Scale with percentiles. 

1649 

1650 Examples: 

1651 1. Scale linearly to map 5th percentile to 0 and 99.8th percentile to 1.0 

1652 - in YAML 

1653 ```yaml 

1654 preprocessing: 

1655 - id: scale_range 

1656 kwargs: 

1657 axes: ['y', 'x'] 

1658 max_percentile: 99.8 

1659 min_percentile: 5.0 

1660 ``` 

1661 - in Python 

1662 

1663 >>> preprocessing = [ 

1664 ... ScaleRangeDescr( 

1665 ... kwargs=ScaleRangeKwargs( 

1666 ... axes= (AxisId('y'), AxisId('x')), 

1667 ... max_percentile= 99.8, 

1668 ... min_percentile= 5.0, 

1669 ... ) 

1670 ... ) 

1671 ... ] 

1672 

1673 2. Combine the above scaling with additional clipping to clip values outside the range given by the percentiles. 

1674 - in YAML 

1675 ```yaml 

1676 preprocessing: 

1677 - id: scale_range 

1678 kwargs: 

1679 axes: ['y', 'x'] 

1680 max_percentile: 99.8 

1681 min_percentile: 5.0 

1682 - id: clip 

1683 kwargs: 

1684 min: 0.0 

1685 max: 1.0 

1686 ``` 

1687 - in Python 

1688 

1689 >>> preprocessing = [ 

1690 ... ScaleRangeDescr( 

1691 ... kwargs=ScaleRangeKwargs( 

1692 ... axes= (AxisId('y'), AxisId('x')), 

1693 ... max_percentile= 99.8, 

1694 ... min_percentile= 5.0, 

1695 ... ) 

1696 ... ), 

1697 ... ClipDescr( 

1698 ... kwargs=ClipKwargs( 

1699 ... min=0.0, 

1700 ... max=1.0, 

1701 ... ) 

1702 ... ), 

1703 ... ] 

1704 

1705 """ 

1706 

1707 implemented_id: ClassVar[Literal["scale_range"]] = "scale_range" 

1708 if TYPE_CHECKING: 

1709 id: Literal["scale_range"] = "scale_range" 

1710 else: 

1711 id: Literal["scale_range"] 

1712 kwargs: ScaleRangeKwargs = Field(default_factory=ScaleRangeKwargs.model_construct) 

1713 

1714 

1715class ScaleMeanVarianceKwargs(KwargsNode): 

1716 """key word arguments for [ScaleMeanVarianceKwargs][]""" 

1717 

1718 reference_tensor: TensorId 

1719 """ID of unprocessed input tensor to match.""" 

1720 

1721 axes: Annotated[Sequence[AxisId] | None, Field(examples=[("batch", "x", "y")])] = ( 

1722 None 

1723 ) 

1724 """The subset of axes to normalize jointly, i.e. axes to reduce to compute mean/std. 

1725 For example to normalize 'batch', 'x' and 'y' jointly in a tensor ('batch', 'channel', 'y', 'x') 

1726 resulting in a tensor of equal shape normalized per channel, specify `axes=('batch', 'x', 'y')`. 

1727 To normalize samples independently, leave out the 'batch' axis. 

1728 Default: Scale all axes jointly.""" 

1729 

1730 eps: Annotated[float, Interval(gt=0, le=0.1)] = 1e-6 

1731 """Epsilon for numeric stability: 

1732 `out = (tensor - mean) / (std + eps) * (ref_std + eps) + ref_mean.`""" 

1733 

1734 

1735class ScaleMeanVarianceDescr(NodeWithExplicitlySetFields): 

1736 """Scale a tensor's data distribution to match another tensor's mean/std. 

1737 `out = (tensor - mean) / (std + eps) * (ref_std + eps) + ref_mean.` 

1738 """ 

1739 

1740 implemented_id: ClassVar[Literal["scale_mean_variance"]] = "scale_mean_variance" 

1741 if TYPE_CHECKING: 

1742 id: Literal["scale_mean_variance"] = "scale_mean_variance" 

1743 else: 

1744 id: Literal["scale_mean_variance"] 

1745 kwargs: ScaleMeanVarianceKwargs 

1746 

1747 

1748PreprocessingDescr = Annotated[ 

1749 Union[ 

1750 BinarizeDescr, 

1751 ClipDescr, 

1752 EnsureDtypeDescr, 

1753 FixedZeroMeanUnitVarianceDescr, 

1754 ScaleLinearDescr, 

1755 ScaleRangeDescr, 

1756 SigmoidDescr, 

1757 SoftmaxDescr, 

1758 ZeroMeanUnitVarianceDescr, 

1759 ], 

1760 Discriminator("id"), 

1761] 

1762PostprocessingDescr = Annotated[ 

1763 Union[ 

1764 BinarizeDescr, 

1765 CellposeFlowDynamicsDescr, 

1766 ClipDescr, 

1767 CustomProcessingDescr, 

1768 EnsureDtypeDescr, 

1769 FixedZeroMeanUnitVarianceDescr, 

1770 ScaleLinearDescr, 

1771 ScaleMeanVarianceDescr, 

1772 ScaleRangeDescr, 

1773 SigmoidDescr, 

1774 SoftmaxDescr, 

1775 StardistPostprocessingDescr, 

1776 ZeroMeanUnitVarianceDescr, 

1777 ], 

1778 Discriminator("id"), 

1779] 

1780 

1781IO_AxisT = TypeVar("IO_AxisT", InputAxis, OutputAxis) 

1782 

1783 

1784class TensorDescrBase(Node, Generic[IO_AxisT]): 

1785 id: TensorId 

1786 """Tensor id. No duplicates are allowed.""" 

1787 

1788 description: Annotated[str, MaxLen(128)] = "" 

1789 """free text description""" 

1790 

1791 axes: NotEmpty[Sequence[IO_AxisT]] 

1792 """tensor axes""" 

1793 

1794 @property 

1795 def shape(self): 

1796 return tuple(a.size for a in self.axes) 

1797 

1798 @field_validator("axes", mode="after", check_fields=False) 

1799 @classmethod 

1800 def _validate_axes(cls, axes: Sequence[AnyAxis]) -> Sequence[AnyAxis]: 

1801 batch_axes = [a for a in axes if a.type == "batch"] 

1802 if len(batch_axes) > 1: 

1803 raise ValueError( 

1804 f"Only one batch axis (per tensor) allowed, but got {batch_axes}" 

1805 ) 

1806 

1807 seen_ids: set[AxisId] = set() 

1808 duplicate_axes_ids: set[AxisId] = set() 

1809 for a in axes: 

1810 (duplicate_axes_ids if a.id in seen_ids else seen_ids).add(a.id) 

1811 

1812 if duplicate_axes_ids: 

1813 raise ValueError(f"Duplicate axis ids: {duplicate_axes_ids}") 

1814 

1815 return axes 

1816 

1817 test_tensor: FAIR[FileDescr_package | None] = None 

1818 """An example tensor to use for testing. 

1819 Using the model with the test input tensors is expected to yield the test output tensors. 

1820 Each test tensor has be a an ndarray in the 

1821 [numpy.lib file format](https://numpy.org/doc/stable/reference/generated/numpy.lib.format.html#module-numpy.lib.format). 

1822 The file extension must be '.npy'.""" 

1823 

1824 sample_tensor: FAIR[FileDescr_package | None] = None 

1825 """A sample tensor to illustrate a possible input/output for the model, 

1826 The sample image primarily serves to inform a human user about an example use case 

1827 and is typically stored as .hdf5, .png or .tiff. 

1828 It has to be readable by the [imageio library](https://imageio.readthedocs.io/en/stable/formats/index.html#supported-formats) 

1829 (numpy's `.npy` format is not supported). 

1830 The image dimensionality has to match the number of axes specified in this tensor description. 

1831 """ 

1832 

1833 @model_validator(mode="after") 

1834 def _validate_sample_tensor(self) -> Self: 

1835 if self.sample_tensor is None or not get_validation_context().perform_io_checks: 

1836 return self 

1837 

1838 reader = get_reader(self.sample_tensor.source, sha256=self.sample_tensor.sha256) 

1839 tensor: NDArray[Any] = imread( # pyright: ignore[reportUnknownVariableType] 

1840 reader.read(), 

1841 extension=PurePosixPath(reader.original_file_name).suffix, 

1842 ) 

1843 n_dims = len(tensor.squeeze().shape) 

1844 n_dims_min = n_dims_max = len(self.axes) 

1845 

1846 for a in self.axes: 

1847 if isinstance(a, BatchAxis): 

1848 n_dims_min -= 1 

1849 elif isinstance(a.size, int): 

1850 if a.size == 1: 

1851 n_dims_min -= 1 

1852 elif isinstance(a.size, (ParameterizedSize, DataDependentSize)): 

1853 if a.size.min == 1: 

1854 n_dims_min -= 1 

1855 elif isinstance(a.size, SizeReference): 

1856 if a.size.offset < 2: 

1857 # size reference may result in singleton axis 

1858 n_dims_min -= 1 

1859 else: 

1860 assert_never(a.size) 

1861 

1862 n_dims_min = max(0, n_dims_min) 

1863 if n_dims < n_dims_min or n_dims > n_dims_max: 

1864 raise ValueError( 

1865 f"Expected sample tensor to have {n_dims_min} to" 

1866 + f" {n_dims_max} dimensions, but found {n_dims} (shape: {tensor.shape})." 

1867 ) 

1868 

1869 return self 

1870 

1871 data: TensorDataDescr | NotEmpty[Sequence[TensorDataDescr]] = ( 

1872 IntervalOrRatioDataDescr() 

1873 ) 

1874 """Description of the tensor's data values, optionally per channel. 

1875 If specified per channel, the data `type` needs to match across channels.""" 

1876 

1877 @property 

1878 def dtype( 

1879 self, 

1880 ) -> Literal[ 

1881 "float32", 

1882 "float64", 

1883 "uint8", 

1884 "int8", 

1885 "uint16", 

1886 "int16", 

1887 "uint32", 

1888 "int32", 

1889 "uint64", 

1890 "int64", 

1891 "bool", 

1892 ]: 

1893 """dtype as specified under `data.type` or `data[i].type`""" 

1894 if isinstance(self.data, collections.abc.Sequence): 

1895 return self.data[0].type 

1896 else: 

1897 return self.data.type 

1898 

1899 @field_validator("data", mode="after") 

1900 @classmethod 

1901 def _check_data_type_across_channels( 

1902 cls, value: TensorDataDescr | NotEmpty[Sequence[TensorDataDescr]] 

1903 ) -> TensorDataDescr | NotEmpty[Sequence[TensorDataDescr]]: 

1904 if not isinstance(value, list): 

1905 return value 

1906 

1907 dtypes = {t.type for t in value} 

1908 if len(dtypes) > 1: 

1909 raise ValueError( 

1910 "Tensor data descriptions per channel need to agree in their data" 

1911 + f" `type`, but found {dtypes}." 

1912 ) 

1913 

1914 return value 

1915 

1916 @model_validator(mode="after") 

1917 def _check_data_matches_channelaxis(self) -> Self: 

1918 if not isinstance(self.data, (list, tuple)): 

1919 return self 

1920 

1921 for a in self.axes: 

1922 if isinstance(a, ChannelAxis): 

1923 size = a.size 

1924 assert isinstance(size, int) 

1925 break 

1926 else: 

1927 return self 

1928 

1929 if len(self.data) != size: 

1930 raise ValueError( 

1931 f"Got tensor data descriptions for {len(self.data)} channels, but" 

1932 + f" '{a.id}' axis has size {size}." 

1933 ) 

1934 

1935 return self 

1936 

1937 def get_axis_sizes_for_array(self, array: NDArray[Any]) -> dict[AxisId, int]: 

1938 if len(array.shape) != len(self.axes): 

1939 raise ValueError( 

1940 f"Dimension mismatch: array shape {array.shape} (#{len(array.shape)})" 

1941 + f" incompatible with {len(self.axes)} axes." 

1942 ) 

1943 return {a.id: array.shape[i] for i, a in enumerate(self.axes)} 

1944 

1945 

1946class ConstantPadding(Node): 

1947 mode: Literal["constant"] = "constant" 

1948 value: int | float = 0 

1949 

1950 

1951class EdgePadding(Node): 

1952 mode: Literal["edge"] = "edge" 

1953 

1954 

1955class ReflectPadding(Node): 

1956 mode: Literal["reflect"] = "reflect" 

1957 

1958 

1959class SymmetricPadding(Node): 

1960 mode: Literal["symmetric"] = "symmetric" 

1961 

1962 

1963Padding = Union[ConstantPadding, EdgePadding, ReflectPadding, SymmetricPadding] 

1964 

1965 

1966class ModelId(ResourceId): 

1967 pass 

1968 

1969 

1970class InputTensorDescr(TensorDescrBase[InputAxis]): 

1971 id: TensorId = TensorId("input") 

1972 """Input tensor id. 

1973 No duplicates are allowed across all inputs and outputs.""" 

1974 

1975 output_of: ModelId | None = None 

1976 """If this input tensor is the output of another model, specify the model id here. 

1977 This model's input id must match the output id of the referenced model. 

1978 """ 

1979 

1980 @model_validator(mode="after") 

1981 def _validate_output_of(self) -> Self: 

1982 if self.output_of is None: 

1983 return self 

1984 

1985 try: 

1986 with get_validation_context().replace(perform_io_checks=False): 

1987 opened_ref_model = open_bioimageio_yaml(self.output_of) 

1988 format_version = opened_ref_model.content["format_version"] 

1989 assert isinstance(format_version, str) 

1990 if format_version.startswith("0.4"): 

1991 ref_model = _ModelDescr04.model_validate(opened_ref_model.content) 

1992 else: 

1993 ref_model = ModelDescr.model_validate(opened_ref_model.content) 

1994 except Exception as e: 

1995 raise ValueError( 

1996 f"Failed to load model '{self.output_of}' referenced under output_of: {e}" 

1997 ) 

1998 

1999 try: 

2000 ref_model_outputs = { 

2001 t.id if isinstance(t, OutputTensorDescr) else TensorId(t.name) 

2002 for t in ref_model.outputs 

2003 } 

2004 except Exception as e: 

2005 raise ValueError( 

2006 f"Failed to read output IDs of model '{self.output_of}' referenced under output_of: {e}" 

2007 ) 

2008 

2009 if self.id not in ref_model_outputs: 

2010 raise ValueError( 

2011 f"Input tensor '{self.id}' is specified as output of model '{self.output_of}', " 

2012 + f"but that model's outputs are {ref_model_outputs}." 

2013 ) 

2014 return self 

2015 

2016 optional: bool = False 

2017 """indicates that this tensor may be `None`""" 

2018 

2019 pad: Padding | None = None 

2020 """Explicitly specify how to pad this input tensor. 

2021 

2022 Use `axes[i].pad` to specify padding width. 

2023 

2024 Note: 

2025 Non-blockwise sample prediction only applies padding for axes with a `pad` specification. 

2026 """ 

2027 

2028 preprocessing: list[PreprocessingDescr] = Field( 

2029 default_factory=cast(Callable[[], List[PreprocessingDescr]], list) 

2030 ) 

2031 """Description of how this input should be preprocessed. 

2032 

2033 notes: 

2034 - If preprocessing does not start with an 'ensure_dtype' entry, it is added 

2035 to ensure an input tensor's data type matches the input tensor's data description. 

2036 - If preprocessing does not end with an 'ensure_dtype' or 'binarize' entry, an 

2037 'ensure_dtype' step is added to ensure preprocessing steps are not unintentionally 

2038 changing the data type. 

2039 """ 

2040 

2041 @model_validator(mode="after") 

2042 def _validate_preprocessing_kwargs(self) -> Self: 

2043 axes_ids = [a.id for a in self.axes] 

2044 for p in self.preprocessing: 

2045 kwargs_axes: Sequence[Any] | None = p.kwargs.get("axes") 

2046 if kwargs_axes is None: 

2047 continue 

2048 

2049 if not isinstance(kwargs_axes, collections.abc.Sequence): 

2050 raise ValueError( 

2051 f"Expected `preprocessing.i.kwargs.axes` to be a sequence, but got {type(kwargs_axes)}" 

2052 ) 

2053 

2054 if any(a not in axes_ids for a in kwargs_axes): 

2055 raise ValueError( 

2056 "`preprocessing.i.kwargs.axes` needs to be subset of axes ids" 

2057 ) 

2058 

2059 if isinstance(self.data, (NominalOrOrdinalDataDescr, IntervalOrRatioDataDescr)): 

2060 dtype = self.data.type 

2061 else: 

2062 dtype = self.data[0].type 

2063 

2064 # ensure `preprocessing` begins with `EnsureDtypeDescr` 

2065 if not self.preprocessing or not isinstance( 

2066 self.preprocessing[0], EnsureDtypeDescr 

2067 ): 

2068 self.preprocessing.insert( 

2069 0, EnsureDtypeDescr(kwargs=EnsureDtypeKwargs(dtype=dtype)) 

2070 ) 

2071 

2072 # ensure `preprocessing` ends with `EnsureDtypeDescr` or `BinarizeDescr` 

2073 if not isinstance(self.preprocessing[-1], (EnsureDtypeDescr, BinarizeDescr)): 

2074 self.preprocessing.append( 

2075 EnsureDtypeDescr(kwargs=EnsureDtypeKwargs(dtype=dtype)) 

2076 ) 

2077 

2078 return self 

2079 

2080 

2081def convert_axes( 

2082 axes: str, 

2083 *, 

2084 shape: Sequence[int] | _ParameterizedInputShape_v0_4 | _ImplicitOutputShape_v0_4, 

2085 tensor_type: Literal["input", "output"], 

2086 halo: Sequence[int] | None, 

2087 size_refs: Mapping[_TensorName_v0_4, Mapping[str, int]], 

2088): 

2089 ret: list[AnyAxis] = [] 

2090 for i, a in enumerate(axes): 

2091 axis_type = _AXIS_TYPE_MAP.get(a, a) 

2092 if axis_type == "batch": 

2093 ret.append(BatchAxis()) 

2094 continue 

2095 

2096 scale = 1.0 

2097 if isinstance(shape, _ParameterizedInputShape_v0_4): 

2098 if shape.step[i] == 0: 

2099 size = shape.min[i] 

2100 else: 

2101 size = ParameterizedSize(min=shape.min[i], step=shape.step[i]) 

2102 elif isinstance(shape, _ImplicitOutputShape_v0_4): 

2103 ref_t = str(shape.reference_tensor) 

2104 if ref_t.count(".") == 1: 

2105 t_id, orig_a_id = ref_t.split(".") 

2106 else: 

2107 t_id = ref_t 

2108 orig_a_id = a 

2109 

2110 a_id = _AXIS_ID_MAP.get(orig_a_id, a) 

2111 if not (orig_scale := shape.scale[i]): 

2112 # old way to insert a new axis dimension 

2113 size = int(2 * shape.offset[i]) 

2114 else: 

2115 scale = 1 / orig_scale 

2116 if axis_type in ("channel", "index"): 

2117 # these axes no longer have a scale 

2118 offset_from_scale = orig_scale * size_refs.get( 

2119 _TensorName_v0_4(t_id), {} 

2120 ).get(orig_a_id, 0) 

2121 else: 

2122 offset_from_scale = 0 

2123 size = SizeReference( 

2124 tensor_id=TensorId(t_id), 

2125 axis_id=AxisId(a_id), 

2126 offset=int(offset_from_scale + 2 * shape.offset[i]), 

2127 ) 

2128 else: 

2129 size = shape[i] 

2130 

2131 if axis_type == "time": 

2132 if tensor_type == "input": 

2133 ret.append(TimeInputAxis(size=size, scale=scale)) 

2134 else: 

2135 assert not isinstance(size, ParameterizedSize) 

2136 if halo is None: 

2137 ret.append(TimeOutputAxis(size=size, scale=scale)) 

2138 else: 

2139 assert not isinstance(size, int) 

2140 ret.append( 

2141 TimeOutputAxisWithHalo(size=size, scale=scale, halo=halo[i]) 

2142 ) 

2143 

2144 elif axis_type == "index": 

2145 if tensor_type == "input": 

2146 ret.append(IndexInputAxis(size=size)) 

2147 else: 

2148 if isinstance(size, ParameterizedSize): 

2149 size = DataDependentSize(min=size.min) 

2150 

2151 ret.append(IndexOutputAxis(size=size)) 

2152 elif axis_type == "channel": 

2153 assert not isinstance(size, ParameterizedSize) 

2154 if isinstance(size, SizeReference): 

2155 warnings.warn( 

2156 "Conversion of channel size from an implicit output shape may be" 

2157 + " wrong" 

2158 ) 

2159 ret.append( 

2160 ChannelAxis( 

2161 channel_names=[f"channel{i}" for i in range(size.offset)] 

2162 ) 

2163 ) 

2164 else: 

2165 ret.append( 

2166 ChannelAxis(channel_names=[f"channel{i}" for i in range(size)]) 

2167 ) 

2168 elif axis_type == "space": 

2169 if tensor_type == "input": 

2170 ret.append(SpaceInputAxis(id=AxisId(a), size=size, scale=scale)) 

2171 else: 

2172 assert not isinstance(size, ParameterizedSize) 

2173 if halo is None or halo[i] == 0: 

2174 ret.append(SpaceOutputAxis(id=AxisId(a), size=size, scale=scale)) 

2175 elif isinstance(size, int): 

2176 raise NotImplementedError( 

2177 f"output axis with halo and fixed size (here {size}) not allowed" 

2178 ) 

2179 else: 

2180 ret.append( 

2181 SpaceOutputAxisWithHalo( 

2182 id=AxisId(a), size=size, scale=scale, halo=halo[i] 

2183 ) 

2184 ) 

2185 

2186 return ret 

2187 

2188 

2189def _axes_letters_to_ids( 

2190 axes: str | None, 

2191) -> list[AxisId] | None: 

2192 if axes is None: 

2193 return None 

2194 

2195 return [AxisId(a) for a in axes] 

2196 

2197 

2198def _get_complement_v04_axis( 

2199 tensor_axes: Sequence[str], axes: Sequence[str] | None 

2200) -> AxisId | None: 

2201 if axes is None: 

2202 return None 

2203 

2204 non_complement_axes = set(axes) | {"b"} 

2205 complement_axes = [a for a in tensor_axes if a not in non_complement_axes] 

2206 if len(complement_axes) > 1: 

2207 raise ValueError( 

2208 f"Expected none or a single complement axis, but axes '{axes}' " 

2209 + f"for tensor dims '{tensor_axes}' leave '{complement_axes}'." 

2210 ) 

2211 

2212 return None if not complement_axes else AxisId(complement_axes[0]) 

2213 

2214 

2215def _convert_proc( 

2216 p: _PreprocessingDescr_v0_4 | _PostprocessingDescr_v0_4, 

2217 tensor_axes: Sequence[str], 

2218) -> PreprocessingDescr | PostprocessingDescr: 

2219 if isinstance(p, _BinarizeDescr_v0_4): 

2220 return BinarizeDescr(kwargs=BinarizeKwargs(threshold=p.kwargs.threshold)) 

2221 elif isinstance(p, _ClipDescr_v0_4): 

2222 return ClipDescr(kwargs=ClipKwargs(min=p.kwargs.min, max=p.kwargs.max)) 

2223 elif isinstance(p, _SigmoidDescr_v0_4): 

2224 return SigmoidDescr() 

2225 elif isinstance(p, _ScaleLinearDescr_v0_4): 

2226 axes = _axes_letters_to_ids(p.kwargs.axes) 

2227 if p.kwargs.axes is None: 

2228 axis = None 

2229 else: 

2230 axis = _get_complement_v04_axis(tensor_axes, p.kwargs.axes) 

2231 

2232 if axis is None: 

2233 assert not isinstance(p.kwargs.gain, list) 

2234 assert not isinstance(p.kwargs.offset, list) 

2235 kwargs = ScaleLinearKwargs(gain=p.kwargs.gain, offset=p.kwargs.offset) 

2236 else: 

2237 kwargs = ScaleLinearAlongAxisKwargs( 

2238 axis=axis, gain=p.kwargs.gain, offset=p.kwargs.offset 

2239 ) 

2240 return ScaleLinearDescr(kwargs=kwargs) 

2241 elif isinstance(p, _ScaleMeanVarianceDescr_v0_4): 

2242 return ScaleMeanVarianceDescr( 

2243 kwargs=ScaleMeanVarianceKwargs( 

2244 axes=_axes_letters_to_ids(p.kwargs.axes), 

2245 reference_tensor=TensorId(str(p.kwargs.reference_tensor)), 

2246 eps=p.kwargs.eps, 

2247 ) 

2248 ) 

2249 elif isinstance(p, _ZeroMeanUnitVarianceDescr_v0_4): 

2250 if p.kwargs.mode == "fixed": 

2251 mean = p.kwargs.mean 

2252 std = p.kwargs.std 

2253 assert mean is not None 

2254 assert std is not None 

2255 

2256 axis = _get_complement_v04_axis(tensor_axes, p.kwargs.axes) 

2257 

2258 if axis is None: 

2259 if isinstance(mean, list): 

2260 raise ValueError("Expected single float value for mean, not <list>") 

2261 if isinstance(std, list): 

2262 raise ValueError("Expected single float value for std, not <list>") 

2263 return FixedZeroMeanUnitVarianceDescr( 

2264 kwargs=FixedZeroMeanUnitVarianceKwargs.model_construct( 

2265 mean=mean, 

2266 std=std, 

2267 ) 

2268 ) 

2269 else: 

2270 if not isinstance(mean, list): 

2271 mean = [float(mean)] 

2272 if not isinstance(std, list): 

2273 std = [float(std)] 

2274 

2275 return FixedZeroMeanUnitVarianceDescr( 

2276 kwargs=FixedZeroMeanUnitVarianceAlongAxisKwargs( 

2277 axis=axis, mean=mean, std=std 

2278 ) 

2279 ) 

2280 

2281 else: 

2282 axes = _axes_letters_to_ids(p.kwargs.axes) or [] 

2283 if p.kwargs.mode == "per_dataset": 

2284 axes = [AxisId("batch")] + axes 

2285 if not axes: 

2286 axes = None 

2287 return ZeroMeanUnitVarianceDescr( 

2288 kwargs=ZeroMeanUnitVarianceKwargs(axes=axes, eps=p.kwargs.eps) 

2289 ) 

2290 

2291 elif isinstance(p, _ScaleRangeDescr_v0_4): 

2292 return ScaleRangeDescr( 

2293 kwargs=ScaleRangeKwargs( 

2294 axes=_axes_letters_to_ids(p.kwargs.axes), 

2295 min_percentile=p.kwargs.min_percentile, 

2296 max_percentile=p.kwargs.max_percentile, 

2297 eps=p.kwargs.eps, 

2298 ) 

2299 ) 

2300 else: 

2301 assert_never(p) 

2302 

2303 

2304class _InputTensorConv( 

2305 Converter[ 

2306 _InputTensorDescr_v0_4, 

2307 InputTensorDescr, 

2308 FileSource, 

2309 Optional[FileSource], 

2310 Mapping[_TensorName_v0_4, Mapping[str, int]], 

2311 ] 

2312): 

2313 def _convert( 

2314 self, 

2315 src: _InputTensorDescr_v0_4, 

2316 tgt: type[InputTensorDescr | dict[str, Any]], 

2317 test_tensor: FileSource, 

2318 sample_tensor: FileSource | None, 

2319 size_refs: Mapping[_TensorName_v0_4, Mapping[str, int]], 

2320 ) -> InputTensorDescr | dict[str, Any]: 

2321 axes: list[InputAxis] = convert_axes( # pyright: ignore[reportAssignmentType] 

2322 src.axes, 

2323 shape=src.shape, 

2324 tensor_type="input", 

2325 halo=None, 

2326 size_refs=size_refs, 

2327 ) 

2328 prep: list[PreprocessingDescr] = [] 

2329 for p in src.preprocessing: 

2330 cp = _convert_proc(p, src.axes) 

2331 assert not isinstance( 

2332 cp, 

2333 ( 

2334 CellposeFlowDynamicsDescr, 

2335 CustomProcessingDescr, 

2336 ScaleMeanVarianceDescr, 

2337 StardistPostprocessingDescr, 

2338 ), 

2339 ) 

2340 prep.append(cp) 

2341 

2342 prep.append(EnsureDtypeDescr(kwargs=EnsureDtypeKwargs(dtype="float32"))) 

2343 

2344 return tgt( 

2345 axes=axes, 

2346 id=TensorId(str(src.name)), 

2347 test_tensor=FileDescr(source=test_tensor), 

2348 sample_tensor=( 

2349 None if sample_tensor is None else FileDescr(source=sample_tensor) 

2350 ), 

2351 data={"type": src.data_type}, # pyright: ignore[reportArgumentType] 

2352 preprocessing=prep, 

2353 ) 

2354 

2355 

2356_input_tensor_conv = _InputTensorConv(_InputTensorDescr_v0_4, InputTensorDescr) 

2357 

2358 

2359class OutputTensorDescr(TensorDescrBase[OutputAxis]): 

2360 id: TensorId = TensorId("output") 

2361 """Output tensor id. 

2362 No duplicates are allowed across all inputs and outputs.""" 

2363 

2364 postprocessing: list[PostprocessingDescr] = Field( 

2365 default_factory=cast(Callable[[], List[PostprocessingDescr]], list) 

2366 ) 

2367 """Description of how this output should be postprocessed. 

2368 

2369 note: `postprocessing` always ends with an 'ensure_dtype' operation. 

2370 If not given this is added to cast to this tensor's `data.type`. 

2371 """ 

2372 

2373 @model_validator(mode="after") 

2374 def _validate_postprocessing_kwargs(self) -> Self: 

2375 axes_ids = [a.id for a in self.axes] 

2376 for p in self.postprocessing: 

2377 kwargs_axes = p.kwargs.get("axes") 

2378 if kwargs_axes is None: 

2379 continue 

2380 

2381 if not isinstance(kwargs_axes, collections.abc.Sequence): 

2382 raise ValueError( 

2383 f"expected `axes` sequence, but got {type(kwargs_axes)}" 

2384 ) 

2385 

2386 kwargs_axes_seq: Sequence[Any] = cast(Sequence[Any], kwargs_axes) 

2387 if any(a not in axes_ids for a in kwargs_axes_seq): 

2388 raise ValueError("`kwargs.axes` needs to be subset of axes ids") 

2389 

2390 if isinstance(self.data, (NominalOrOrdinalDataDescr, IntervalOrRatioDataDescr)): 

2391 dtype = self.data.type 

2392 else: 

2393 dtype = self.data[0].type 

2394 

2395 # ensure `postprocessing` ends with `EnsureDtypeDescr` or `BinarizeDescr` 

2396 if not self.postprocessing or not isinstance( 

2397 self.postprocessing[-1], (EnsureDtypeDescr, BinarizeDescr) 

2398 ): 

2399 self.postprocessing.append( 

2400 EnsureDtypeDescr(kwargs=EnsureDtypeKwargs(dtype=dtype)) 

2401 ) 

2402 return self 

2403 

2404 

2405class _OutputTensorConv( 

2406 Converter[ 

2407 _OutputTensorDescr_v0_4, 

2408 OutputTensorDescr, 

2409 FileSource, 

2410 Optional[FileSource], 

2411 Mapping[_TensorName_v0_4, Mapping[str, int]], 

2412 ] 

2413): 

2414 def _convert( 

2415 self, 

2416 src: _OutputTensorDescr_v0_4, 

2417 tgt: type[OutputTensorDescr | dict[str, Any]], 

2418 test_tensor: FileSource, 

2419 sample_tensor: FileSource | None, 

2420 size_refs: Mapping[_TensorName_v0_4, Mapping[str, int]], 

2421 ) -> OutputTensorDescr | dict[str, Any]: 

2422 # TODO: split convert_axes into convert_output_axes and convert_input_axes 

2423 axes: list[OutputAxis] = convert_axes( # pyright: ignore[reportAssignmentType] 

2424 src.axes, 

2425 shape=src.shape, 

2426 tensor_type="output", 

2427 halo=src.halo, 

2428 size_refs=size_refs, 

2429 ) 

2430 data_descr: dict[str, Any] = {"type": src.data_type} 

2431 if data_descr["type"] == "bool": 

2432 data_descr["values"] = [False, True] 

2433 

2434 return tgt( 

2435 axes=axes, 

2436 id=TensorId(str(src.name)), 

2437 test_tensor=FileDescr(source=test_tensor), 

2438 sample_tensor=( 

2439 None if sample_tensor is None else FileDescr(source=sample_tensor) 

2440 ), 

2441 data=data_descr, # pyright: ignore[reportArgumentType] 

2442 postprocessing=[_convert_proc(p, src.axes) for p in src.postprocessing], 

2443 ) 

2444 

2445 

2446_output_tensor_conv = _OutputTensorConv(_OutputTensorDescr_v0_4, OutputTensorDescr) 

2447 

2448 

2449TensorDescr = Union[InputTensorDescr, OutputTensorDescr] 

2450 

2451 

2452def get_halos( 

2453 tensors: Mapping[TensorId, TensorDescr], 

2454 /, 

2455) -> dict[TensorId, dict[AxisId, tuple[int, int]]]: 

2456 """Get all input and output halos from tensor descriptions. 

2457 

2458 Note: 

2459 - Input halos are to be padded 

2460 - Output halos are to be cropped 

2461 """ 

2462 halos: dict[TensorId, dict[AxisId, tuple[int, int]]] = {} 

2463 for descr in tensors.values(): 

2464 if isinstance(descr, InputTensorDescr): 

2465 continue 

2466 for axis in descr.axes: 

2467 if not isinstance(axis, WithHalo): 

2468 continue 

2469 

2470 ref_scale = next( 

2471 a 

2472 for a in tensors[axis.size.tensor_id].axes 

2473 if a.id == axis.size.axis_id 

2474 ).scale 

2475 

2476 # set output halo (to be cropped) 

2477 halos.setdefault(descr.id, {})[axis.id] = (axis.halo, axis.halo) 

2478 # set input halo (to be padded) 

2479 pad_width = int(axis.halo / axis.scale * ref_scale) 

2480 halos.setdefault(axis.size.tensor_id, {})[axis.size.axis_id] = ( 

2481 pad_width, 

2482 pad_width, 

2483 ) 

2484 

2485 return halos 

2486 

2487 

2488def validate_tensors( 

2489 tensors: Mapping[TensorId, tuple[TensorDescr, NDArray[Any] | None]], 

2490 tensor_origin: Literal[ 

2491 "source", "test_tensor" 

2492 ] = "source", # for more precise error messages 

2493 *, 

2494 pad_inputs: bool | Literal["allow"] = True, 

2495 crop_outputs: bool | Literal["allow"] = True, 

2496): 

2497 """Validate all inputs (and optionally output tensors) against their tensor descriptions. 

2498 

2499 Args: 

2500 tensors: Mapping of tensor id to a tuple of tensor description and optional numpy array. 

2501 tensor_origin: String to use in error messages to indicate the origin of the tensors being validated. 

2502 pad_inputs: Wether to apply/allow padding of inputs before shape comparison 

2503 crop_outputs: Wether to apply/allow cropping of outputs before shape comparison. 

2504 """ 

2505 all_tensor_axes: dict[TensorId, dict[AxisId, tuple[AnyAxis, int | None]]] = {} 

2506 

2507 def e_msg_location(d: TensorDescr): 

2508 return f"{'inputs' if isinstance(d, InputTensorDescr) else 'outputs'}[{d.id}]" 

2509 

2510 for descr, array in tensors.values(): 

2511 if array is None: 

2512 axis_sizes = {a.id: None for a in descr.axes} 

2513 else: 

2514 try: 

2515 axis_sizes = descr.get_axis_sizes_for_array(array) 

2516 except ValueError as e: 

2517 raise ValueError(f"{e_msg_location(descr)} {e}") 

2518 

2519 all_tensor_axes[descr.id] = {a.id: (a, axis_sizes[a.id]) for a in descr.axes} 

2520 

2521 # get halos to be padded/cropped to validate against halo-adjusted sizes 

2522 io_halos = get_halos({k: v[0] for k, v in tensors.items()}) 

2523 

2524 for descr, array in tensors.values(): 

2525 if array is None: 

2526 continue 

2527 

2528 if descr.dtype in ("float32", "float64"): 

2529 invalid_test_tensor_dtype = array.dtype.name not in ( 

2530 "float32", 

2531 "float64", 

2532 "uint8", 

2533 "int8", 

2534 "uint16", 

2535 "int16", 

2536 "uint32", 

2537 "int32", 

2538 "uint64", 

2539 "int64", 

2540 ) 

2541 else: 

2542 invalid_test_tensor_dtype = array.dtype.name != descr.dtype 

2543 

2544 if invalid_test_tensor_dtype: 

2545 raise ValueError( 

2546 f"{tensor_origin} data type '{array.dtype.name}' does not" 

2547 + f" match described {e_msg_location(descr)}.dtype '{descr.dtype}'" 

2548 ) 

2549 

2550 if array.min() > -1e-4 and array.max() < 1e-4: 

2551 raise ValueError( 

2552 "Output values are too small for reliable testing." 

2553 + f" Values <-1e5 or >=1e5 must be present in {tensor_origin}" 

2554 ) 

2555 

2556 for a in descr.axes: 

2557 actual_size = all_tensor_axes[descr.id][a.id][1] 

2558 

2559 if actual_size is None: 

2560 continue 

2561 

2562 if a.size is None: 

2563 continue 

2564 

2565 # add padding width to actual tensor size 

2566 total_axis_halo = sum(io_halos.get(descr.id, {}).get(a.id, (0, 0))) 

2567 if isinstance(descr, InputTensorDescr): 

2568 # pad input halos 

2569 actual_size_with_halo = actual_size + total_axis_halo 

2570 if pad_inputs is True: 

2571 check_sizes = {actual_size_with_halo} 

2572 size_hint = " (after padding input halo)" 

2573 elif pad_inputs == "allow": 

2574 check_sizes = {actual_size, actual_size_with_halo} 

2575 size_hint = " (with or without padding input halo)" 

2576 elif pad_inputs is False: 

2577 check_sizes = {actual_size} 

2578 size_hint = "" 

2579 else: 

2580 assert_never(pad_inputs) 

2581 

2582 elif isinstance(descr, OutputTensorDescr): 

2583 # crop output halos 

2584 actual_size_with_halo = max(0, actual_size - total_axis_halo) 

2585 if crop_outputs is True: 

2586 check_sizes = {actual_size_with_halo} 

2587 size_hint = " (after cropping output halo)" 

2588 elif crop_outputs == "allow": 

2589 check_sizes = {actual_size, actual_size_with_halo} 

2590 size_hint = " (with or without cropping output halo)" 

2591 elif crop_outputs is False: 

2592 check_sizes = {actual_size} 

2593 size_hint = "" 

2594 else: 

2595 assert_never(crop_outputs) 

2596 else: 

2597 assert_never(descr) 

2598 

2599 del actual_size # make sure we explicitly use unchanged or halo-adjusted size from here on 

2600 

2601 if isinstance(a.size, int): 

2602 if a.size not in check_sizes: 

2603 raise ValueError( 

2604 f"{e_msg_location(descr)}.axes[{a.id}]: {tensor_origin} axis " 

2605 + f"has incompatible size {check_sizes}{size_hint}, expected {a.size}" 

2606 ) 

2607 elif isinstance(a.size, (ParameterizedSize, DataDependentSize)): 

2608 _ = try_all_raise_last( 

2609 (partial(a.size.validate_size, s) for s in check_sizes), 

2610 f"{e_msg_location(descr)}.axes[{a.id}]: {tensor_origin} axis ", 

2611 ) 

2612 elif isinstance(a.size, SizeReference): 

2613 ref_tensor_axes = all_tensor_axes.get(a.size.tensor_id) 

2614 if ref_tensor_axes is None: 

2615 raise ValueError( 

2616 f"{e_msg_location(descr)}.axes[{a.id}].size.tensor_id: Unknown tensor" 

2617 + f" reference '{a.size.tensor_id}', available: {list(all_tensor_axes)}" 

2618 ) 

2619 

2620 ref_axis, ref_size = ref_tensor_axes.get(a.size.axis_id, (None, None)) 

2621 if ref_axis is None or ref_size is None: 

2622 raise ValueError( 

2623 f"{e_msg_location(descr)}.axes[{a.id}].size.axis_id: Unknown tensor axis" 

2624 + f" reference '{a.size.tensor_id}.{a.size.axis_id}, available: {list(ref_tensor_axes)}" 

2625 ) 

2626 

2627 if a.unit != ref_axis.unit: 

2628 raise ValueError( 

2629 f"{e_msg_location(descr)}.axes[{a.id}].size: `SizeReference` requires" 

2630 + " axis and reference axis to have the same `unit`, but" 

2631 + f" {a.unit}!={ref_axis.unit}" 

2632 ) 

2633 

2634 if ( 

2635 expected_size := ( 

2636 ref_size * ref_axis.scale / a.scale + a.size.offset 

2637 ) 

2638 ) not in check_sizes: 

2639 raise ValueError( 

2640 f"{e_msg_location(descr)}.{tensor_origin}: axis '{a.id}' of size" 

2641 + f" {check_sizes} invalid for referenced size {ref_size};" 

2642 + f" expected {expected_size}" 

2643 ) 

2644 else: 

2645 assert_never(a.size) 

2646 

2647 

2648FileDescr_dependencies = Annotated[ 

2649 FileDescr_package, 

2650 WithSuffix((".yaml", ".yml"), case_sensitive=True), 

2651 Field(examples=[{"source": "environment.yaml"}]), 

2652] 

2653 

2654 

2655class _ArchitectureCallableDescr(Node): 

2656 callable: Annotated[Identifier, Field(examples=["MyNetworkClass", "get_my_model"])] 

2657 """Identifier of the callable that returns a torch.nn.Module instance.""" 

2658 

2659 kwargs: dict[str, YamlValue] = Field( 

2660 default_factory=cast(Callable[[], Dict[str, YamlValue]], dict) 

2661 ) 

2662 """key word arguments for the `callable`""" 

2663 

2664 

2665class ArchitectureFromFileDescr(_ArchitectureCallableDescr, FileDescr): 

2666 source: Annotated[FileSource, AfterValidator(wo_special_file_name)] 

2667 """Architecture source file""" 

2668 

2669 @model_serializer(mode="wrap", when_used="unless-none") 

2670 def _serialize(self, nxt: SerializerFunctionWrapHandler, info: SerializationInfo): 

2671 return package_file_descr_serializer(self, nxt, info) 

2672 

2673 

2674class ArchitectureFromLibraryDescr(_ArchitectureCallableDescr): 

2675 import_from: str 

2676 """Where to import the callable from, i.e. `from <import_from> import <callable>`""" 

2677 

2678 

2679class _ArchFileConv( 

2680 Converter[ 

2681 _CallableFromFile_v0_4, 

2682 ArchitectureFromFileDescr, 

2683 Optional[Sha256], 

2684 Dict[str, Any], 

2685 ] 

2686): 

2687 def _convert( 

2688 self, 

2689 src: _CallableFromFile_v0_4, 

2690 tgt: type[ArchitectureFromFileDescr | dict[str, Any]], 

2691 sha256: Sha256 | None, 

2692 kwargs: dict[str, Any], 

2693 ) -> ArchitectureFromFileDescr | dict[str, Any]: 

2694 if src.startswith("http") and src.count(":") == 2: 

2695 http, source, callable_ = src.split(":") 

2696 source = f"{http}:{source}" 

2697 elif not src.startswith("http") and src.count(":") == 1: 

2698 source, callable_ = src.split(":") 

2699 else: 

2700 source = str(src) 

2701 callable_ = str(src) 

2702 return tgt( 

2703 callable=Identifier(callable_), 

2704 source=cast(FileSource, source), 

2705 sha256=sha256, 

2706 kwargs=kwargs, 

2707 ) 

2708 

2709 

2710_arch_file_conv = _ArchFileConv(_CallableFromFile_v0_4, ArchitectureFromFileDescr) 

2711 

2712 

2713class _ArchLibConv( 

2714 Converter[ 

2715 _CallableFromDepencency_v0_4, ArchitectureFromLibraryDescr, Dict[str, Any] 

2716 ] 

2717): 

2718 def _convert( 

2719 self, 

2720 src: _CallableFromDepencency_v0_4, 

2721 tgt: type[ArchitectureFromLibraryDescr | dict[str, Any]], 

2722 kwargs: dict[str, Any], 

2723 ) -> ArchitectureFromLibraryDescr | dict[str, Any]: 

2724 *mods, callable_ = src.split(".") 

2725 import_from = ".".join(mods) 

2726 return tgt( 

2727 import_from=import_from, callable=Identifier(callable_), kwargs=kwargs 

2728 ) 

2729 

2730 

2731_arch_lib_conv = _ArchLibConv( 

2732 _CallableFromDepencency_v0_4, ArchitectureFromLibraryDescr 

2733) 

2734 

2735 

2736class WeightsEntryDescrBase(FileDescr): 

2737 type: ClassVar[WeightsFormat] 

2738 weights_format_name: ClassVar[str] # human readable 

2739 

2740 source: Annotated[FileSource, AfterValidator(wo_special_file_name)] 

2741 """Source of the weights file.""" 

2742 

2743 authors: list[Author] | None = None 

2744 """Authors 

2745 Either the person(s) that have trained this model resulting in the original weights file. 

2746 (If this is the initial weights entry, i.e. it does not have a `parent`) 

2747 Or the person(s) who have converted the weights to this weights format. 

2748 (If this is a child weight, i.e. it has a `parent` field) 

2749 """ 

2750 

2751 parent: Annotated[WeightsFormat | None, Field(examples=["pytorch_state_dict"])] = ( 

2752 None 

2753 ) 

2754 """The source weights these weights were converted from. 

2755 For example, if a model's weights were converted from the `pytorch_state_dict` format to `torchscript`, 

2756 The `pytorch_state_dict` weights entry has no `parent` and is the parent of the `torchscript` weights. 

2757 All weight entries except one (the initial set of weights resulting from training the model), 

2758 need to have this field.""" 

2759 

2760 comment: str = "" 

2761 """A comment about this weights entry, for example how these weights were created.""" 

2762 

2763 @model_validator(mode="after") 

2764 def _validate(self) -> Self: 

2765 if self.type == self.parent: 

2766 raise ValueError("Weights entry can't be it's own parent.") 

2767 

2768 return self 

2769 

2770 @model_serializer(mode="wrap", when_used="unless-none") 

2771 def _serialize(self, nxt: SerializerFunctionWrapHandler, info: SerializationInfo): 

2772 return package_file_descr_serializer(self, nxt, info) 

2773 

2774 

2775class KerasHdf5WeightsDescr(WeightsEntryDescrBase): 

2776 type: ClassVar[WeightsFormat] = "keras_hdf5" 

2777 weights_format_name: ClassVar[str] = "Keras HDF5" 

2778 tensorflow_version: Version 

2779 """TensorFlow version used to create these weights.""" 

2780 

2781 

2782class KerasV3WeightsDescr(WeightsEntryDescrBase): 

2783 type: ClassVar[WeightsFormat] = "keras_v3" 

2784 weights_format_name: ClassVar[str] = "Keras v3" 

2785 keras_version: Annotated[Version, Ge(Version(3))] 

2786 """Keras version used to create these weights.""" 

2787 backend: tuple[Literal["tensorflow", "jax", "torch"], Version] 

2788 """Keras backend used to create these weights.""" 

2789 source: Annotated[ 

2790 FileSource, 

2791 AfterValidator(wo_special_file_name), 

2792 WithSuffix(".keras", case_sensitive=True), 

2793 ] 

2794 """Source of the .keras weights file.""" 

2795 

2796 

2797FileDescr_external_data = Annotated[ 

2798 FileDescr_package, 

2799 WithSuffix(".data", case_sensitive=True), 

2800 Field(examples=[{"source": "weights.onnx.data"}]), 

2801] 

2802 

2803 

2804class OnnxWeightsDescr(WeightsEntryDescrBase): 

2805 type: ClassVar[WeightsFormat] = "onnx" 

2806 weights_format_name: ClassVar[str] = "ONNX" 

2807 opset_version: Annotated[int, Ge(7)] 

2808 """ONNX opset version""" 

2809 

2810 external_data: FileDescr_external_data | None = None 

2811 """Source of the external ONNX data file holding the weights. 

2812 (If present **source** holds the ONNX architecture without weights).""" 

2813 

2814 @model_validator(mode="after") 

2815 def _validate_external_data_unique_file_name(self) -> Self: 

2816 if self.external_data is not None and ( 

2817 extract_file_name(self.source) 

2818 == extract_file_name(self.external_data.source) 

2819 ): 

2820 raise ValueError( 

2821 f"ONNX `external_data` file name '{extract_file_name(self.external_data.source)}'" 

2822 + " must be different from ONNX `source` file name." 

2823 ) 

2824 

2825 return self 

2826 

2827 

2828class PytorchStateDictWeightsDescr(WeightsEntryDescrBase): 

2829 type: ClassVar[WeightsFormat] = "pytorch_state_dict" 

2830 weights_format_name: ClassVar[str] = "Pytorch State Dict" 

2831 architecture: ArchitectureFromFileDescr | ArchitectureFromLibraryDescr 

2832 pytorch_version: Version 

2833 """Version of the PyTorch library used. 

2834 If `architecture.depencencies` is specified it has to include pytorch and any version pinning has to be compatible. 

2835 """ 

2836 dependencies: FileDescr_dependencies | None = None 

2837 """Custom depencies beyond pytorch described in a Conda environment file. 

2838 Allows to specify custom dependencies, see conda docs: 

2839 - [Exporting an environment file across platforms](https://conda.io/projects/conda/en/latest/user-guide/tasks/manage-environments.html#exporting-an-environment-file-across-platforms) 

2840 - [Creating an environment file manually](https://conda.io/projects/conda/en/latest/user-guide/tasks/manage-environments.html#creating-an-environment-file-manually) 

2841 

2842 The conda environment file should include pytorch and any version pinning has to be compatible with 

2843 **pytorch_version**. 

2844 """ 

2845 strict: bool = True 

2846 """Whether to allow missing or unexpected keys or to be strict about the architecture matching the state dict weights.""" 

2847 

2848 

2849class TensorflowJsWeightsDescr(WeightsEntryDescrBase): 

2850 type: ClassVar[WeightsFormat] = "tensorflow_js" 

2851 weights_format_name: ClassVar[str] = "Tensorflow.js" 

2852 tensorflow_version: Version 

2853 """Version of the TensorFlow library used.""" 

2854 

2855 source: Annotated[FileSource, AfterValidator(wo_special_file_name)] 

2856 """The multi-file weights. 

2857 All required files/folders should be a zip archive.""" 

2858 

2859 

2860class TensorflowSavedModelBundleWeightsDescr(WeightsEntryDescrBase): 

2861 type: ClassVar[WeightsFormat] = "tensorflow_saved_model_bundle" 

2862 weights_format_name: ClassVar[str] = "Tensorflow Saved Model" 

2863 tensorflow_version: Version 

2864 """Version of the TensorFlow library used.""" 

2865 

2866 dependencies: FileDescr_dependencies | None = None 

2867 """Custom dependencies beyond tensorflow. 

2868 Should include tensorflow and any version pinning has to be compatible with **tensorflow_version**.""" 

2869 

2870 source: Annotated[FileSource, AfterValidator(wo_special_file_name)] 

2871 """The multi-file weights. 

2872 All required files/folders should be a zip archive.""" 

2873 

2874 

2875class TorchscriptWeightsDescr(WeightsEntryDescrBase): 

2876 type: ClassVar[WeightsFormat] = "torchscript" 

2877 weights_format_name: ClassVar[str] = "TorchScript" 

2878 pytorch_version: Version 

2879 """Version of the PyTorch library used.""" 

2880 

2881 

2882SpecificWeightsDescr = Union[ 

2883 KerasHdf5WeightsDescr, 

2884 KerasV3WeightsDescr, 

2885 OnnxWeightsDescr, 

2886 PytorchStateDictWeightsDescr, 

2887 TensorflowJsWeightsDescr, 

2888 TensorflowSavedModelBundleWeightsDescr, 

2889 TorchscriptWeightsDescr, 

2890] 

2891 

2892 

2893class WeightsDescr(Node): 

2894 keras_hdf5: KerasHdf5WeightsDescr | None = None 

2895 keras_v3: KerasV3WeightsDescr | None = None 

2896 onnx: OnnxWeightsDescr | None = None 

2897 pytorch_state_dict: PytorchStateDictWeightsDescr | None = None 

2898 tensorflow_js: TensorflowJsWeightsDescr | None = None 

2899 tensorflow_saved_model_bundle: TensorflowSavedModelBundleWeightsDescr | None = None 

2900 torchscript: TorchscriptWeightsDescr | None = None 

2901 

2902 @model_validator(mode="after") 

2903 def check_entries(self) -> Self: 

2904 entries = {wtype for wtype, entry in self if entry is not None} 

2905 

2906 if not entries: 

2907 raise ValueError("Missing weights entry") 

2908 

2909 entries_wo_parent = { 

2910 wtype 

2911 for wtype, entry in self 

2912 if entry is not None and hasattr(entry, "parent") and entry.parent is None 

2913 } 

2914 if len(entries_wo_parent) != 1: 

2915 issue_warning( 

2916 "Exactly one weights entry may not specify the `parent` field (got" 

2917 + " {value}). That entry is considered the original set of model weights." 

2918 + " Other weight formats are created through conversion of the orignal or" 

2919 + " already converted weights. They have to reference the weights format" 

2920 + " they were converted from as their `parent`.", 

2921 value=len(entries_wo_parent), 

2922 field="weights", 

2923 ) 

2924 

2925 for wtype, entry in self: 

2926 if entry is None: 

2927 continue 

2928 

2929 assert hasattr(entry, "type") 

2930 assert hasattr(entry, "parent") 

2931 assert wtype == entry.type 

2932 if ( 

2933 entry.parent is not None and entry.parent not in entries 

2934 ): # self reference checked for `parent` field 

2935 raise ValueError( 

2936 f"`weights.{wtype}.parent={entry.parent} not in specified weight" 

2937 + f" formats: {entries}" 

2938 ) 

2939 

2940 return self 

2941 

2942 def __getitem__( 

2943 self, 

2944 key: WeightsFormat, 

2945 ): 

2946 if key == "keras_hdf5": 

2947 ret = self.keras_hdf5 

2948 elif key == "keras_v3": 

2949 ret = self.keras_v3 

2950 elif key == "onnx": 

2951 ret = self.onnx 

2952 elif key == "pytorch_state_dict": 

2953 ret = self.pytorch_state_dict 

2954 elif key == "tensorflow_js": 

2955 ret = self.tensorflow_js 

2956 elif key == "tensorflow_saved_model_bundle": 

2957 ret = self.tensorflow_saved_model_bundle 

2958 elif key == "torchscript": 

2959 ret = self.torchscript 

2960 else: 

2961 raise KeyError(key) 

2962 

2963 if ret is None: 

2964 raise KeyError(key) 

2965 

2966 return ret 

2967 

2968 @overload 

2969 def __setitem__( 

2970 self, key: Literal["keras_hdf5"], value: KerasHdf5WeightsDescr | None 

2971 ) -> None: ... 

2972 @overload 

2973 def __setitem__( 

2974 self, key: Literal["keras_v3"], value: KerasV3WeightsDescr | None 

2975 ) -> None: ... 

2976 @overload 

2977 def __setitem__( 

2978 self, key: Literal["onnx"], value: OnnxWeightsDescr | None 

2979 ) -> None: ... 

2980 @overload 

2981 def __setitem__( 

2982 self, 

2983 key: Literal["pytorch_state_dict"], 

2984 value: PytorchStateDictWeightsDescr | None, 

2985 ) -> None: ... 

2986 @overload 

2987 def __setitem__( 

2988 self, key: Literal["tensorflow_js"], value: TensorflowJsWeightsDescr | None 

2989 ) -> None: ... 

2990 @overload 

2991 def __setitem__( 

2992 self, 

2993 key: Literal["tensorflow_saved_model_bundle"], 

2994 value: TensorflowSavedModelBundleWeightsDescr | None, 

2995 ) -> None: ... 

2996 @overload 

2997 def __setitem__( 

2998 self, key: Literal["torchscript"], value: TorchscriptWeightsDescr | None 

2999 ) -> None: ... 

3000 

3001 def __setitem__( 

3002 self, 

3003 key: WeightsFormat, 

3004 value: SpecificWeightsDescr | None, 

3005 ): 

3006 if key == "keras_hdf5": 

3007 if value is not None and not isinstance(value, KerasHdf5WeightsDescr): 

3008 raise TypeError( 

3009 f"Expected KerasHdf5WeightsDescr or None for key 'keras_hdf5', got {type(value)}" 

3010 ) 

3011 self.keras_hdf5 = value 

3012 elif key == "keras_v3": 

3013 if value is not None and not isinstance(value, KerasV3WeightsDescr): 

3014 raise TypeError( 

3015 f"Expected KerasV3WeightsDescr or None for key 'keras_v3', got {type(value)}" 

3016 ) 

3017 self.keras_v3 = value 

3018 elif key == "onnx": 

3019 if value is not None and not isinstance(value, OnnxWeightsDescr): 

3020 raise TypeError( 

3021 f"Expected OnnxWeightsDescr or None for key 'onnx', got {type(value)}" 

3022 ) 

3023 self.onnx = value 

3024 elif key == "pytorch_state_dict": 

3025 if value is not None and not isinstance( 

3026 value, PytorchStateDictWeightsDescr 

3027 ): 

3028 raise TypeError( 

3029 f"Expected PytorchStateDictWeightsDescr or None for key 'pytorch_state_dict', got {type(value)}" 

3030 ) 

3031 self.pytorch_state_dict = value 

3032 elif key == "tensorflow_js": 

3033 if value is not None and not isinstance(value, TensorflowJsWeightsDescr): 

3034 raise TypeError( 

3035 f"Expected TensorflowJsWeightsDescr or None for key 'tensorflow_js', got {type(value)}" 

3036 ) 

3037 self.tensorflow_js = value 

3038 elif key == "tensorflow_saved_model_bundle": 

3039 if value is not None and not isinstance( 

3040 value, TensorflowSavedModelBundleWeightsDescr 

3041 ): 

3042 raise TypeError( 

3043 f"Expected TensorflowSavedModelBundleWeightsDescr or None for key 'tensorflow_saved_model_bundle', got {type(value)}" 

3044 ) 

3045 self.tensorflow_saved_model_bundle = value 

3046 elif key == "torchscript": 

3047 if value is not None and not isinstance(value, TorchscriptWeightsDescr): 

3048 raise TypeError( 

3049 f"Expected TorchscriptWeightsDescr or None for key 'torchscript', got {type(value)}" 

3050 ) 

3051 self.torchscript = value 

3052 else: 

3053 raise KeyError(key) 

3054 

3055 @property 

3056 def available_formats(self) -> dict[WeightsFormat, SpecificWeightsDescr]: 

3057 return { 

3058 **({} if self.keras_hdf5 is None else {"keras_hdf5": self.keras_hdf5}), 

3059 **({} if self.keras_v3 is None else {"keras_v3": self.keras_v3}), 

3060 **({} if self.onnx is None else {"onnx": self.onnx}), 

3061 **( 

3062 {} 

3063 if self.pytorch_state_dict is None 

3064 else {"pytorch_state_dict": self.pytorch_state_dict} 

3065 ), 

3066 **( 

3067 {} 

3068 if self.tensorflow_js is None 

3069 else {"tensorflow_js": self.tensorflow_js} 

3070 ), 

3071 **( 

3072 {} 

3073 if self.tensorflow_saved_model_bundle is None 

3074 else { 

3075 "tensorflow_saved_model_bundle": self.tensorflow_saved_model_bundle 

3076 } 

3077 ), 

3078 **({} if self.torchscript is None else {"torchscript": self.torchscript}), 

3079 } 

3080 

3081 @property 

3082 def missing_formats(self) -> set[WeightsFormat]: 

3083 return { 

3084 wf for wf in get_args(WeightsFormat) if wf not in self.available_formats 

3085 } 

3086 

3087 

3088class LinkedModel(LinkedResourceBase): 

3089 """Reference to a bioimage.io model.""" 

3090 

3091 id: ModelId 

3092 """A valid model `id` from the bioimage.io collection.""" 

3093 

3094 

3095class _DataDepSize(NamedTuple): 

3096 min: StrictInt 

3097 max: StrictInt | None 

3098 

3099 

3100class _AxisSizes(NamedTuple): 

3101 """the lenghts of all axes of model inputs and outputs""" 

3102 

3103 inputs: dict[tuple[TensorId, AxisId], int] 

3104 outputs: dict[tuple[TensorId, AxisId], int | _DataDepSize] 

3105 

3106 

3107class _TensorSizes(NamedTuple): 

3108 """_AxisSizes as nested dicts""" 

3109 

3110 inputs: dict[TensorId, dict[AxisId, int]] 

3111 outputs: dict[TensorId, dict[AxisId, int | _DataDepSize]] 

3112 

3113 

3114_RT = TypeVar("_RT", RelativeTolerance, Literal[0]) 

3115_AT = TypeVar("_AT", AbsoluteTolerance, Literal[0]) 

3116 

3117 

3118class _ReproducibilityToleranceBase(Node, ABC, Generic[_RT, _AT], extra="allow"): 

3119 """Describes what small numerical differences -- if any -- may be tolerated 

3120 in the generated output when executing in different environments. 

3121 

3122 A tensor element *output* is considered mismatched to the **test_tensor** if 

3123 abs(*output* - **test_tensor**) > **absolute_tolerance** + **relative_tolerance** * abs(**test_tensor**). 

3124 (Internally we call [numpy.testing.assert_allclose](https://numpy.org/doc/stable/reference/generated/numpy.testing.assert_allclose.html).) 

3125 

3126 Motivation: 

3127 For testing we can request the respective deep learning frameworks to be as 

3128 reproducible as possible by setting seeds and chosing deterministic algorithms, 

3129 but differences in operating systems, available hardware and installed drivers 

3130 may still lead to numerical differences. 

3131 """ 

3132 

3133 mismatched_elements_per_million: MismatchedElementsPerMillion = 100 

3134 """Maximum number of mismatched elements/pixels per million to tolerate.""" 

3135 

3136 weights_formats: Sequence[WeightsFormat] = () 

3137 """Limits the weights formats these details apply to.""" 

3138 

3139 relative_tolerance: _RT 

3140 absolute_tolerance: _AT 

3141 

3142 @model_serializer(mode="wrap") 

3143 def serialize( 

3144 self, 

3145 handler: SerializerFunctionWrapHandler, 

3146 ) -> dict[str, Any]: 

3147 """explicitly serialize tolerance values 

3148 

3149 (even if they are default values and `exclude_defaults=True`). 

3150 """ 

3151 data = handler(self) 

3152 data["relative_tolerance"] = self.relative_tolerance 

3153 data["absolute_tolerance"] = self.absolute_tolerance 

3154 data["mismatched_elements_per_million"] = self.mismatched_elements_per_million 

3155 return data 

3156 

3157 

3158class ReproducibilityTolerance( 

3159 _ReproducibilityToleranceBase[RelativeTolerance, AbsoluteTolerance] 

3160): 

3161 output_ids: Sequence[TensorId] = () 

3162 """Limits the output tensor IDs these reproducibility details apply to. 

3163 

3164 If empty, the reproducibility details apply to all outputs.""" 

3165 

3166 relative_tolerance: RelativeTolerance = 1e-3 

3167 """Maximum relative tolerance of reproduced test tensor. 

3168 Needs to be 0 for discrete (integer and boolean) outputs. 

3169 """ 

3170 

3171 absolute_tolerance: AbsoluteTolerance = 1e-3 

3172 """Maximum absolute tolerance of reproduced test tensor. 

3173 

3174 The `absolute_tolerance` may not be greater than 1% of the maximum absolute value of the test tensor 

3175 (this serves to prevent false positives in the presence of very small numbers). 

3176 Needs to be 0 for discrete (integer and boolean) outputs. 

3177 """ 

3178 

3179 

3180class DiscreteReproducibilityTolerance( 

3181 _ReproducibilityToleranceBase[Literal[0], Literal[0]] 

3182): 

3183 output_ids: Literal["any_discrete_output"] = "any_discrete_output" 

3184 """Limits the output tensor IDs to all outputs of integer or boolean data type.""" 

3185 

3186 relative_tolerance: Literal[0] = 0 

3187 """Relative tolerance does not apply to discrete outputs""" 

3188 

3189 absolute_tolerance: Literal[0] = 0 

3190 """Absolute tolerance does not apply to discrete outputs""" 

3191 

3192 mismatched_elements_per_million: MismatchedElementsPerMillion = 1000 

3193 """Maximum number of mismatched elements/pixels per million to tolerate. 

3194 

3195 Note: Increased default mismatched elements per million compared to 

3196 [`ReproducibilityTolerance`][ReproducibilityTolerance]/[`ContinuousReproducibilityTolerance`][ContinuousReproducibilityTolerance], 

3197 since a single pixel mismatch can be more significant than for continuous outputs with numeric tolerance.""" 

3198 

3199 

3200class ContinuousReproducibilityTolerance( 

3201 _ReproducibilityToleranceBase[RelativeTolerance, AbsoluteTolerance] 

3202): 

3203 output_ids: Literal["any_continuous_output"] = "any_continuous_output" 

3204 """Limits the output tensor IDs to all outputs of float data type.""" 

3205 

3206 relative_tolerance: RelativeTolerance = 1e-3 

3207 """Maximum relative tolerance of reproduced test tensor. 

3208 Needs to be 0 for discrete (integer and boolean) outputs. 

3209 """ 

3210 

3211 absolute_tolerance: AbsoluteTolerance = 1e-3 

3212 """Maximum absolute tolerance of reproduced test tensor. 

3213 

3214 The `absolute_tolerance` may not be greater than 1% of the maximum absolute value of the test tensor 

3215 (this serves to prevent false positives in the presence of very small numbers). 

3216 Needs to be 0 for discrete (integer and boolean) outputs. 

3217 """ 

3218 

3219 

3220class BiasRisksLimitations(Node, extra="allow"): 

3221 """Known biases, risks, technical limitations, and recommendations for model use.""" 

3222 

3223 known_biases: str = dedent("""\ 

3224 In general bioimage models may suffer from biases caused by: 

3225 

3226 - Imaging protocol dependencies 

3227 - Use of a specific cell type 

3228 - Species-specific training data limitations 

3229 

3230 """) 

3231 """Biases in training data or model behavior.""" 

3232 

3233 risks: str = dedent("""\ 

3234 Common risks in bioimage analysis include: 

3235 

3236 - Erroneously assuming generalization to unseen experimental conditions 

3237 - Trusting (overconfident) model outputs without validation 

3238 - Misinterpretation of results 

3239 

3240 """) 

3241 """Potential risks in the context of bioimage analysis.""" 

3242 

3243 limitations: str | None = None 

3244 """Technical limitations and failure modes.""" 

3245 

3246 recommendations: str = "Users (both direct and downstream) should be made aware of the risks, biases and limitations of the model." 

3247 """Mitigation strategies regarding `known_biases`, `risks`, and `limitations`, as well as applicable best practices. 

3248 

3249 Consider: 

3250 - How to use a validation dataset? 

3251 - How to manually validate? 

3252 - Feasibility of domain adaptation for different experimental setups? 

3253 

3254 """ 

3255 

3256 def format_md(self) -> str: 

3257 if self.limitations is None: 

3258 limitations_header = "" 

3259 else: 

3260 limitations_header = "## Limitations\n\n" 

3261 

3262 return f"""# Bias, Risks, and Limitations 

3263 

3264{self.known_biases} 

3265 

3266{self.risks} 

3267 

3268{limitations_header}{self.limitations or ""} 

3269 

3270## Recommendations 

3271 

3272{self.recommendations} 

3273 

3274""" 

3275 

3276 

3277class TrainingDetails(Node, extra="allow"): 

3278 training_preprocessing: str | None = None 

3279 """Detailed image preprocessing steps during model training: 

3280 

3281 Mention: 

3282 - *Normalization methods* 

3283 - *Augmentation strategies* 

3284 - *Resizing/resampling procedures* 

3285 - *Artifact handling* 

3286 

3287 """ 

3288 

3289 training_epochs: float | None = None 

3290 """Number of training epochs.""" 

3291 

3292 training_batch_size: float | None = None 

3293 """Batch size used in training.""" 

3294 

3295 initial_learning_rate: float | None = None 

3296 """Initial learning rate used in training.""" 

3297 

3298 learning_rate_schedule: str | None = None 

3299 """Learning rate schedule used in training.""" 

3300 

3301 loss_function: str | None = None 

3302 """Loss function used in training, e.g. nn.MSELoss.""" 

3303 

3304 loss_function_kwargs: dict[str, YamlValue] = Field( 

3305 default_factory=cast(Callable[[], Dict[str, YamlValue]], dict) 

3306 ) 

3307 """key word arguments for the `loss_function`""" 

3308 

3309 optimizer: str | None = None 

3310 """optimizer, e.g. torch.optim.Adam""" 

3311 

3312 optimizer_kwargs: dict[str, YamlValue] = Field( 

3313 default_factory=cast(Callable[[], Dict[str, YamlValue]], dict) 

3314 ) 

3315 """key word arguments for the `optimizer`""" 

3316 

3317 regularization: str | None = None 

3318 """Regularization techniques used during training, e.g. drop-out or weight decay.""" 

3319 

3320 training_duration: float | None = None 

3321 """Total training duration in hours.""" 

3322 

3323 

3324class Evaluation(Node, extra="allow"): 

3325 model_id: ModelId | None = None 

3326 """Model being evaluated.""" 

3327 

3328 dataset_id: DatasetId 

3329 """Dataset used for evaluation.""" 

3330 

3331 dataset_source: HttpUrl 

3332 """Source of the dataset.""" 

3333 

3334 dataset_role: Literal["train", "validation", "test", "independent", "unknown"] 

3335 """Role of the dataset used for evaluation. 

3336 

3337 - `train`: dataset was (part of) the training data 

3338 - `validation`: dataset was (part of) the validation data used during training, e.g. used for model selection or hyperparameter tuning 

3339 - `test`: dataset was (part of) the designated test data; not used during training or validation, but acquired from the same source/distribution as training data 

3340 - `independent`: dataset is entirely independent test data; not used during training or validation, and acquired from a different source/distribution than training data 

3341 - `unknown`: role of the dataset is unknown; choose this if you are not certain if (a subset) of the data was seen by the model during training. 

3342 """ 

3343 

3344 sample_count: int 

3345 """Number of evaluated samples.""" 

3346 

3347 evaluation_factors: list[Annotated[str, MaxLen(16)]] 

3348 """(Abbreviations of) each evaluation factor. 

3349 

3350 Evaluation factors are criteria along which model performance is evaluated, e.g. different image conditions 

3351 like 'low SNR', 'high cell density', or different biological conditions like 'cell type A', 'cell type B'. 

3352 An 'overall' factor may be included to summarize performance across all conditions. 

3353 """ 

3354 

3355 evaluation_factors_long: list[str] 

3356 """Descriptions (long form) of each evaluation factor.""" 

3357 

3358 metrics: list[Annotated[str, MaxLen(16)]] 

3359 """(Abbreviations of) metrics used for evaluation.""" 

3360 

3361 metrics_long: list[str] 

3362 """Description of each metric used.""" 

3363 

3364 @model_validator(mode="after") 

3365 def _validate_list_lengths(self) -> Self: 

3366 if len(self.evaluation_factors) != len(self.evaluation_factors_long): 

3367 raise ValueError( 

3368 "`evaluation_factors` and `evaluation_factors_long` must have the same length" 

3369 ) 

3370 

3371 if len(self.metrics) != len(self.metrics_long): 

3372 raise ValueError("`metrics` and `metrics_long` must have the same length") 

3373 

3374 if len(self.results) != len(self.metrics): 

3375 raise ValueError("`results` must have the same number of rows as `metrics`") 

3376 

3377 for row in self.results: 

3378 if len(row) != len(self.evaluation_factors): 

3379 raise ValueError( 

3380 "`results` must have the same number of columns (in every row) as `evaluation_factors`" 

3381 ) 

3382 

3383 return self 

3384 

3385 results: list[list[str | float | int]] 

3386 """Results for each metric (rows; outer list) and each evaluation factor (columns; inner list).""" 

3387 

3388 results_summary: str | None = None 

3389 """Interpretation of results for general audience. 

3390 

3391 Consider: 

3392 - Overall model performance 

3393 - Comparison to existing methods 

3394 - Limitations and areas for improvement 

3395 

3396""" 

3397 

3398 def format_md(self): 

3399 results_header = ["Metric"] + self.evaluation_factors 

3400 results_table_cells = [results_header, ["---"] * len(results_header)] + [ 

3401 [metric] + [str(r) for r in row] 

3402 for metric, row in zip(self.metrics, self.results) 

3403 ] 

3404 

3405 results_table = "".join( 

3406 "| " + " | ".join(row) + " |\n" for row in results_table_cells 

3407 ) 

3408 factors = "".join( 

3409 f"\n - {ef}: {efl}" 

3410 for ef, efl in zip(self.evaluation_factors, self.evaluation_factors_long) 

3411 ) 

3412 metrics = "".join( 

3413 f"\n - {em}: {eml}" for em, eml in zip(self.metrics, self.metrics_long) 

3414 ) 

3415 

3416 return f"""## Testing Data, Factors & Metrics 

3417 

3418Evaluation of {self.model_id or "this"} model on the {self.dataset_id} dataset (dataset role: {self.dataset_role}). 

3419 

3420### Testing Data 

3421 

3422- **Source:** [{self.dataset_id}]({self.dataset_source}) 

3423- **Size:** {self.sample_count} evaluated samples 

3424 

3425### Factors 

3426{factors} 

3427 

3428### Metrics 

3429{metrics} 

3430 

3431## Results 

3432 

3433### Quantitative Results 

3434 

3435{results_table} 

3436 

3437### Summary 

3438 

3439{self.results_summary or "missing"} 

3440 

3441""" 

3442 

3443 

3444class EnvironmentalImpact(Node, extra="allow"): 

3445 """Environmental considerations for model training and deployment. 

3446 

3447 Carbon emissions can be estimated using the [Machine Learning Impact calculator](https://mlco2.github.io/impact#compute) presented in [Lacoste et al. (2019)](https://arxiv.org/abs/1910.09700). 

3448 """ 

3449 

3450 hardware_type: str | None = None 

3451 """GPU/CPU specifications""" 

3452 

3453 hours_used: float | None = None 

3454 """Total compute hours""" 

3455 

3456 cloud_provider: str | None = None 

3457 """If applicable""" 

3458 

3459 compute_region: str | None = None 

3460 """Geographic location""" 

3461 

3462 co2_emitted: float | None = None 

3463 """kg CO2 equivalent 

3464 

3465 Carbon emissions can be estimated using the [Machine Learning Impact calculator](https://mlco2.github.io/impact#compute) presented in [Lacoste et al. (2019)](https://arxiv.org/abs/1910.09700). 

3466 """ 

3467 

3468 def format_md(self): 

3469 """Filled Markdown template section following [Hugging Face Model Card Template](https://huggingface.co/docs/hub/en/model-card-annotated).""" 

3470 if self == self.__class__(): 

3471 return "" 

3472 

3473 ret = "# Environmental Impact\n\n" 

3474 if self.hardware_type is not None: 

3475 ret += f"- **Hardware Type:** {self.hardware_type}\n" 

3476 if self.hours_used is not None: 

3477 ret += f"- **Hours used:** {self.hours_used}\n" 

3478 if self.cloud_provider is not None: 

3479 ret += f"- **Cloud Provider:** {self.cloud_provider}\n" 

3480 if self.compute_region is not None: 

3481 ret += f"- **Compute Region:** {self.compute_region}\n" 

3482 if self.co2_emitted is not None: 

3483 ret += f"- **Carbon Emitted:** {self.co2_emitted} kg CO2e\n" 

3484 

3485 return ret + "\n" 

3486 

3487 

3488class BioimageioConfig(Node, extra="allow"): 

3489 reproducibility_tolerance: Sequence[ 

3490 ReproducibilityTolerance 

3491 | ContinuousReproducibilityTolerance 

3492 | DiscreteReproducibilityTolerance 

3493 ] = Field( 

3494 default_factory=lambda: ( 

3495 ContinuousReproducibilityTolerance.model_construct(), 

3496 DiscreteReproducibilityTolerance.model_construct(), 

3497 ) 

3498 ) 

3499 """Tolerances to allow when reproducing the model's test outputs 

3500 from the model's test inputs. 

3501 Only the first entry matching tensor id and weights format is considered. 

3502 For defaults see [`ContinuousReproducibilityTolerance`][ContinuousReproducibilityTolerance]/[`DiscreteReproducibilityTolerance`][DiscreteReproducibilityTolerance] 

3503 """ 

3504 

3505 @model_serializer(mode="wrap") 

3506 def _serialize( 

3507 self, 

3508 handler: SerializerFunctionWrapHandler, 

3509 ) -> dict[str, Any]: 

3510 """explicitly serialize reproducibility_tolerance 

3511 

3512 (even if they are default values and `exclude_defaults=True`). 

3513 """ 

3514 data = handler(self) 

3515 data["reproducibility_tolerance"] = self.reproducibility_tolerance 

3516 return data 

3517 

3518 funded_by: str | None = None 

3519 """Funding agency, grant number if applicable""" 

3520 

3521 architecture_type: Annotated[str, MaxLen(32)] | None = ( 

3522 None # TODO: add to differentiated tags 

3523 ) 

3524 """Model architecture type, e.g., 3D U-Net, ResNet, transformer""" 

3525 

3526 architecture_description: str | None = None 

3527 """Text description of model architecture.""" 

3528 

3529 modality: str | None = None # TODO: add to differentiated tags 

3530 """Input modality, e.g., fluorescence microscopy, electron microscopy""" 

3531 

3532 target_structure: list[str] = Field( # TODO: add to differentiated tags 

3533 default_factory=cast(Callable[[], List[str]], list) 

3534 ) 

3535 """Biological structure(s) the model is designed to analyze, e.g., nuclei, mitochondria, cells""" 

3536 

3537 task: str | None = None # TODO: add to differentiated tags 

3538 """Bioimage-specific task type, e.g., segmentation, classification, detection, denoising""" 

3539 

3540 new_version: ModelId | None = None 

3541 """A new version of this model exists with a different model id.""" 

3542 

3543 out_of_scope_use: str | None = None 

3544 """Describe how the model may be misused in bioimage analysis contexts and what users should **not** do with the model.""" 

3545 

3546 bias_risks_limitations: BiasRisksLimitations = Field( 

3547 default_factory=BiasRisksLimitations.model_construct 

3548 ) 

3549 """Description of known bias, risks, and technical limitations for in-scope model use.""" 

3550 

3551 model_parameter_count: int | None = None 

3552 """Total number of model parameters.""" 

3553 

3554 training: TrainingDetails = Field(default_factory=TrainingDetails.model_construct) 

3555 """Details on how the model was trained.""" 

3556 

3557 inference_time: str | None = None 

3558 """Average inference time per image/tile. Specify hardware and image size. Multiple examples can be given.""" 

3559 

3560 memory_requirements_inference: str | None = None 

3561 """GPU memory needed for inference. Multiple examples with different image size can be given.""" 

3562 

3563 memory_requirements_training: str | None = None 

3564 """GPU memory needed for training. Multiple examples with different image/batch sizes can be given.""" 

3565 

3566 evaluations: list[Evaluation] = Field( 

3567 default_factory=cast(Callable[[], List[Evaluation]], list) 

3568 ) 

3569 """Quantitative model evaluations. 

3570 

3571 Note: 

3572 At the moment we recommend to include only a single test dataset 

3573 (with evaluation factors that may mark subsets of the dataset) 

3574 to avoid confusion and make the presentation of results cleaner. 

3575 """ 

3576 

3577 environmental_impact: EnvironmentalImpact = Field( 

3578 default_factory=EnvironmentalImpact.model_construct 

3579 ) 

3580 """Environmental considerations for model training and deployment""" 

3581 

3582 

3583class Config(Node, extra="allow"): 

3584 bioimageio: BioimageioConfig = Field( 

3585 default_factory=BioimageioConfig.model_construct 

3586 ) 

3587 stardist: YamlValue = None 

3588 

3589 

3590class ModelDescr(GenericModelDescrBase): 

3591 """Specification of the fields used in a bioimage.io-compliant RDF to describe AI models with pretrained weights. 

3592 These fields are typically stored in a YAML file which we call a model resource description file (model RDF). 

3593 """ 

3594 

3595 implemented_format_version: ClassVar[Literal["0.5.14"]] = "0.5.14" 

3596 if TYPE_CHECKING: 

3597 format_version: Literal["0.5.14"] = "0.5.14" 

3598 else: 

3599 format_version: Literal["0.5.14"] 

3600 """Version of the bioimage.io model description specification used. 

3601 When creating a new model always use the latest micro/patch version described here. 

3602 The `format_version` is important for any consumer software to understand how to parse the fields. 

3603 """ 

3604 

3605 implemented_type: ClassVar[Literal["model"]] = "model" 

3606 if TYPE_CHECKING: 

3607 type: Literal["model"] = "model" 

3608 else: 

3609 type: Literal["model"] 

3610 """Specialized resource type 'model'""" 

3611 

3612 id: ModelId | None = None 

3613 """bioimage.io-wide unique resource identifier 

3614 assigned by bioimage.io; version **un**specific.""" 

3615 

3616 authors: FAIR[list[Author]] = Field( 

3617 default_factory=cast(Callable[[], List[Author]], list) 

3618 ) 

3619 """The authors are the creators of the model RDF and the primary points of contact.""" 

3620 

3621 documentation: FAIR[FileDescr_documentation | None] = None 

3622 """Additional model documentation. 

3623 The recommended documentation source file name is `README.md`. An `.md` suffix is mandatory. 

3624 The documentation should include a '#[#] Validation' (sub)section 

3625 with details on how to quantitatively validate the model on unseen data.""" 

3626 

3627 @field_validator("documentation", mode="after") 

3628 @classmethod 

3629 def _validate_documentation(cls, value: FileDescr | None) -> FileDescr | None: 

3630 if not get_validation_context().perform_io_checks or value is None: 

3631 return value 

3632 

3633 doc_reader = get_reader(value) 

3634 doc_content = doc_reader.read().decode(encoding="utf-8") 

3635 if not re.search("#.*[vV]alidation", doc_content): 

3636 issue_warning( 

3637 "No '# Validation' (sub)section found in {value}.", 

3638 value=value, 

3639 field="documentation", 

3640 ) 

3641 

3642 return value 

3643 

3644 inputs: NotEmpty[Sequence[InputTensorDescr]] 

3645 """Describes the input tensors expected by this model.""" 

3646 

3647 @field_validator("inputs", mode="after") 

3648 @classmethod 

3649 def _validate_input_axes( 

3650 cls, inputs: Sequence[InputTensorDescr] 

3651 ) -> Sequence[InputTensorDescr]: 

3652 input_size_refs = cls._get_axes_with_independent_size(inputs) 

3653 

3654 for i, ipt in enumerate(inputs): 

3655 valid_independent_refs: dict[ 

3656 tuple[TensorId, AxisId], 

3657 tuple[TensorDescr, AnyAxis, int | ParameterizedSize], 

3658 ] = { 

3659 **{ 

3660 (ipt.id, a.id): (ipt, a, a.size) 

3661 for a in ipt.axes 

3662 if not isinstance(a, BatchAxis) 

3663 and isinstance(a.size, (int, ParameterizedSize)) 

3664 }, 

3665 **input_size_refs, 

3666 } 

3667 for a, ax in enumerate(ipt.axes): 

3668 cls._validate_axis( 

3669 "inputs", 

3670 i=i, 

3671 tensor_id=ipt.id, 

3672 a=a, 

3673 axis=ax, 

3674 valid_independent_refs=valid_independent_refs, 

3675 ) 

3676 return inputs 

3677 

3678 @staticmethod 

3679 def _validate_axis( 

3680 field_name: str, 

3681 i: int, 

3682 tensor_id: TensorId, 

3683 a: int, 

3684 axis: AnyAxis, 

3685 valid_independent_refs: dict[ 

3686 tuple[TensorId, AxisId], 

3687 tuple[TensorDescr, AnyAxis, int | ParameterizedSize], 

3688 ], 

3689 ): 

3690 if isinstance(axis, BatchAxis) or isinstance( 

3691 axis.size, (int, ParameterizedSize, DataDependentSize) 

3692 ): 

3693 return 

3694 elif not isinstance(axis.size, SizeReference): 

3695 assert_never(axis.size) 

3696 

3697 # validate axis.size SizeReference 

3698 ref = (axis.size.tensor_id, axis.size.axis_id) 

3699 if ref not in valid_independent_refs: 

3700 raise ValueError( 

3701 "Invalid tensor axis reference at" 

3702 + f" {field_name}[{i}].axes[{a}].size: {axis.size}." 

3703 ) 

3704 if ref == (tensor_id, axis.id): 

3705 raise ValueError( 

3706 "Self-referencing not allowed for" 

3707 + f" {field_name}[{i}].axes[{a}].size: {axis.size}" 

3708 ) 

3709 if axis.type == "channel": 

3710 if valid_independent_refs[ref][1].type != "channel": 

3711 raise ValueError( 

3712 "A channel axis' size may only reference another fixed size" 

3713 + " channel axis." 

3714 ) 

3715 if isinstance(axis.channel_names, str) and "{i}" in axis.channel_names: 

3716 ref_size = valid_independent_refs[ref][2] 

3717 assert isinstance(ref_size, int), ( 

3718 "channel axis ref (another channel axis) has to specify fixed" 

3719 + " size" 

3720 ) 

3721 generated_channel_names = [ 

3722 axis.channel_names.format(i=i) for i in range(1, ref_size + 1) 

3723 ] 

3724 axis.channel_names = generated_channel_names 

3725 

3726 if (ax_unit := getattr(axis, "unit", None)) != ( 

3727 ref_unit := getattr(valid_independent_refs[ref][1], "unit", None) 

3728 ): 

3729 raise ValueError( 

3730 "The units of an axis and its reference axis need to match, but" 

3731 + f" '{ax_unit}' != '{ref_unit}'." 

3732 ) 

3733 ref_axis = valid_independent_refs[ref][1] 

3734 if isinstance(ref_axis, BatchAxis): 

3735 raise ValueError( 

3736 f"Invalid reference axis '{ref_axis.id}' for {tensor_id}.{axis.id}" 

3737 + " (a batch axis is not allowed as reference)." 

3738 ) 

3739 

3740 if isinstance(axis, WithHalo): 

3741 min_size = axis.size.get_size(axis, ref_axis, n=0) 

3742 if (min_size - 2 * axis.halo) < 1: 

3743 raise ValueError( 

3744 f"axis {axis.id} with minimum size {min_size} is too small for halo" 

3745 + f" {axis.halo}." 

3746 ) 

3747 

3748 ref_halo = axis.halo * axis.scale / ref_axis.scale 

3749 if ref_halo != int(ref_halo): 

3750 raise ValueError( 

3751 f"Inferred halo for {'.'.join(ref)} is not an integer ({ref_halo} =" 

3752 + f" {tensor_id}.{axis.id}.halo {axis.halo}" 

3753 + f" * {tensor_id}.{axis.id}.scale {axis.scale}" 

3754 + f" / {'.'.join(ref)}.scale {ref_axis.scale})." 

3755 ) 

3756 

3757 def validate_input_tensors( 

3758 self, 

3759 sources: Sequence[NDArray[Any]] | Mapping[TensorId, NDArray[Any] | None], 

3760 *, 

3761 pad_inputs: bool | Literal["allow"] = True, 

3762 crop_outputs: bool | Literal["allow"] = True, 

3763 ) -> Mapping[TensorId, NDArray[Any] | None]: 

3764 """Check if the given input tensors match the model's input tensor descriptions. 

3765 This includes checks of tensor shapes and dtypes, but not of the actual values. 

3766 """ 

3767 if not isinstance(sources, collections.abc.Mapping): 

3768 sources = {descr.id: tensor for descr, tensor in zip(self.inputs, sources)} 

3769 

3770 tensors = { 

3771 **{descr.id: (descr, sources.get(descr.id)) for descr in self.inputs}, 

3772 **{ # outputs are required for halo 

3773 descr.id: (descr, None) for descr in self.outputs 

3774 }, 

3775 } 

3776 validate_tensors(tensors, pad_inputs=pad_inputs, crop_outputs=crop_outputs) 

3777 

3778 return sources 

3779 

3780 @model_validator(mode="after") 

3781 def _validate_test_tensors(self) -> Self: 

3782 if not get_validation_context().perform_io_checks: 

3783 return self 

3784 

3785 test_inputs = { 

3786 descr.id: ( 

3787 descr, 

3788 None if descr.test_tensor is None else load_array(descr.test_tensor), 

3789 ) 

3790 for descr in self.inputs 

3791 } 

3792 test_outputs = { 

3793 descr.id: ( 

3794 descr, 

3795 None if descr.test_tensor is None else load_array(descr.test_tensor), 

3796 ) 

3797 for descr in self.outputs 

3798 } 

3799 

3800 validate_tensors( 

3801 {**test_inputs, **test_outputs}, 

3802 tensor_origin="test_tensor", 

3803 pad_inputs="allow", 

3804 crop_outputs="allow", 

3805 ) 

3806 

3807 seen_output_ids: set[TensorId] = ( 

3808 set() 

3809 ) # reproducibility tolerance is only applicable to the first matching output tensor id 

3810 for rep_tol in self.config.bioimageio.reproducibility_tolerance: 

3811 if not rep_tol.output_ids: 

3812 out_arrays: dict[TensorId, NDArray[Any] | None] = { 

3813 k: v[1] for k, v in test_outputs.items() if k not in seen_output_ids 

3814 } 

3815 else: 

3816 out_arrays = {} 

3817 for k, v in test_outputs.items(): 

3818 if k in seen_output_ids: 

3819 continue 

3820 

3821 if rep_tol.output_ids == "any_discrete_output": 

3822 if not ( 

3823 v[0].dtype.startswith("int") 

3824 or v[0].dtype.startswith("uint") 

3825 or v[0].dtype.startswith("bool") 

3826 ): 

3827 continue 

3828 elif rep_tol.output_ids == "any_continuous_output": 

3829 if not v[0].dtype.startswith("float"): 

3830 continue 

3831 elif k not in rep_tol.output_ids: 

3832 continue 

3833 

3834 out_arrays[k] = v[1] 

3835 

3836 for out_id, array in out_arrays.items(): 

3837 seen_output_ids.add(out_id) 

3838 if array is None: 

3839 continue 

3840 

3841 if rep_tol.absolute_tolerance > (max_test_value := array.max()) * 0.01: 

3842 raise ValueError( 

3843 "config.bioimageio.reproducibility_tolerance.absolute_tolerance=" 

3844 + f"{rep_tol.absolute_tolerance} > 0.01*{max_test_value}" 

3845 + f" (1% of the maximum value of the test tensor '{out_id}')" 

3846 ) 

3847 

3848 return self 

3849 

3850 @model_validator(mode="after") 

3851 def _validate_tensor_references_in_proc_kwargs(self, info: ValidationInfo) -> Self: 

3852 ipt_refs = {t.id for t in self.inputs} 

3853 missing_refs = [ 

3854 k["reference_tensor"] 

3855 for k in [p.kwargs for ipt in self.inputs for p in ipt.preprocessing] 

3856 + [p.kwargs for out in self.outputs for p in out.postprocessing] 

3857 if "reference_tensor" in k 

3858 and k["reference_tensor"] is not None 

3859 and k["reference_tensor"] not in ipt_refs 

3860 ] 

3861 

3862 if missing_refs: 

3863 raise ValueError( 

3864 f"`reference_tensor`s {missing_refs} not found. Valid input tensor" 

3865 + f" references are: {ipt_refs}." 

3866 ) 

3867 

3868 return self 

3869 

3870 name: Annotated[ 

3871 str, 

3872 RestrictCharacters(string.ascii_letters + string.digits + "_+- ()"), 

3873 MinLen(5), 

3874 MaxLen(128), 

3875 warn(MaxLen(64), "Name longer than 64 characters.", INFO), 

3876 ] 

3877 """A human-readable name of this model. 

3878 It should be no longer than 64 characters 

3879 and may only contain letter, number, underscore, minus, parentheses and spaces. 

3880 We recommend to chose a name that refers to the model's task and image modality. 

3881 """ 

3882 

3883 outputs: NotEmpty[Sequence[OutputTensorDescr]] 

3884 """Describes the output tensors.""" 

3885 

3886 @field_validator("outputs", mode="after") 

3887 @classmethod 

3888 def _validate_tensor_ids( 

3889 cls, outputs: Sequence[OutputTensorDescr], info: ValidationInfo 

3890 ) -> Sequence[OutputTensorDescr]: 

3891 tensor_ids = [ 

3892 t.id for t in info.data.get("inputs", []) + info.data.get("outputs", []) 

3893 ] 

3894 duplicate_tensor_ids: list[str] = [] 

3895 seen: set[str] = set() 

3896 for t in tensor_ids: 

3897 if t in seen: 

3898 duplicate_tensor_ids.append(t) 

3899 

3900 seen.add(t) 

3901 

3902 if duplicate_tensor_ids: 

3903 raise ValueError(f"Duplicate tensor ids: {duplicate_tensor_ids}") 

3904 

3905 return outputs 

3906 

3907 @staticmethod 

3908 def _get_axes_with_parameterized_size( 

3909 io: Sequence[InputTensorDescr] | Sequence[OutputTensorDescr], 

3910 ): 

3911 return { 

3912 f"{t.id}.{a.id}": (t, a, a.size) 

3913 for t in io 

3914 for a in t.axes 

3915 if not isinstance(a, BatchAxis) and isinstance(a.size, ParameterizedSize) 

3916 } 

3917 

3918 @staticmethod 

3919 def _get_axes_with_independent_size( 

3920 io: Sequence[InputTensorDescr] | Sequence[OutputTensorDescr], 

3921 ): 

3922 return { 

3923 (t.id, a.id): (t, a, a.size) 

3924 for t in io 

3925 for a in t.axes 

3926 if not isinstance(a, BatchAxis) 

3927 and isinstance(a.size, (int, ParameterizedSize)) 

3928 } 

3929 

3930 @field_validator("outputs", mode="after") 

3931 @classmethod 

3932 def _validate_output_axes( 

3933 cls, outputs: list[OutputTensorDescr], info: ValidationInfo 

3934 ) -> list[OutputTensorDescr]: 

3935 input_size_refs = cls._get_axes_with_independent_size( 

3936 info.data.get("inputs", []) 

3937 ) 

3938 output_size_refs = cls._get_axes_with_independent_size(outputs) 

3939 

3940 for i, out in enumerate(outputs): 

3941 valid_independent_refs: dict[ 

3942 tuple[TensorId, AxisId], 

3943 tuple[TensorDescr, AnyAxis, int | ParameterizedSize], 

3944 ] = { 

3945 **{ 

3946 (out.id, a.id): (out, a, a.size) 

3947 for a in out.axes 

3948 if not isinstance(a, BatchAxis) 

3949 and isinstance(a.size, (int, ParameterizedSize)) 

3950 }, 

3951 **input_size_refs, 

3952 **output_size_refs, 

3953 } 

3954 for a, ax in enumerate(out.axes): 

3955 cls._validate_axis( 

3956 "outputs", 

3957 i, 

3958 out.id, 

3959 a, 

3960 ax, 

3961 valid_independent_refs=valid_independent_refs, 

3962 ) 

3963 

3964 return outputs 

3965 

3966 packaged_by: list[Author] = Field( 

3967 default_factory=cast(Callable[[], List[Author]], list) 

3968 ) 

3969 """The persons that have packaged and uploaded this model. 

3970 Only required if those persons differ from the `authors`.""" 

3971 

3972 parent: LinkedModel | None = None 

3973 """The model from which this model is derived, e.g. by fine-tuning the weights.""" 

3974 

3975 @model_validator(mode="after") 

3976 def _validate_parent_is_not_self(self) -> Self: 

3977 if self.parent is not None and self.parent.id == self.id: 

3978 raise ValueError("A model description may not reference itself as parent.") 

3979 

3980 return self 

3981 

3982 run_mode: Annotated[ 

3983 RunMode | None, 

3984 warn(None, "Run mode '{value}' has limited support across consumer softwares."), 

3985 ] = None 

3986 """Custom run mode for this model: for more complex prediction procedures like test time 

3987 data augmentation that currently cannot be expressed in the specification. 

3988 No standard run modes are defined yet.""" 

3989 

3990 timestamp: Datetime = Field(default_factory=Datetime.now) 

3991 """Timestamp in [ISO 8601](#https://en.wikipedia.org/wiki/ISO_8601) format 

3992 with a few restrictions listed [here](https://docs.python.org/3/library/datetime.html#datetime.datetime.fromisoformat). 

3993 (In Python a datetime object is valid, too).""" 

3994 

3995 training_data: Annotated[ 

3996 None | LinkedDataset | DatasetDescr | DatasetDescr02, 

3997 Field(union_mode="left_to_right"), 

3998 ] = None 

3999 """The dataset used to train this model""" 

4000 

4001 weights: Annotated[WeightsDescr, WrapSerializer(package_weights)] 

4002 """The weights for this model. 

4003 Weights can be given for different formats, but should otherwise be equivalent. 

4004 The available weight formats determine which consumers can use this model.""" 

4005 

4006 config: Config = Field(default_factory=Config.model_construct) 

4007 

4008 @model_validator(mode="after") 

4009 def _validate_reproducibility_tolerance_output_ids(self) -> Self: 

4010 seen_output_ids: set[TensorId] = set() 

4011 for rep_tol in self.config.bioimageio.reproducibility_tolerance: 

4012 if not rep_tol.output_ids: 

4013 continue 

4014 

4015 if rep_tol.output_ids == "any_discrete_output": 

4016 applicable_output_ids = [ 

4017 t.id 

4018 for t in self.outputs 

4019 if t.dtype.startswith("int") 

4020 or t.dtype.startswith("uint") 

4021 or t.dtype.startswith("bool") 

4022 ] 

4023 elif rep_tol.output_ids == "any_continuous_output": 

4024 applicable_output_ids = [ 

4025 t.id for t in self.outputs if t.dtype.startswith("float") 

4026 ] 

4027 else: 

4028 applicable_output_ids = rep_tol.output_ids 

4029 

4030 for out_id in applicable_output_ids: 

4031 if out_id in seen_output_ids: 

4032 continue 

4033 else: 

4034 seen_output_ids.add(out_id) 

4035 

4036 discrete_values = next( 

4037 ( 

4038 t.dtype.startswith("int") 

4039 or t.dtype.startswith("uint") 

4040 or t.dtype.startswith("bool") 

4041 for t in self.outputs 

4042 if t.id == out_id 

4043 ), 

4044 None, 

4045 ) 

4046 if discrete_values is None: 

4047 raise ValueError( 

4048 f"config.bioimageio.reproducibility_tolerance.output_ids contains '{out_id}', which is not a valid output tensor id." 

4049 ) 

4050 if discrete_values and rep_tol.relative_tolerance != 0: 

4051 raise ValueError( 

4052 f"config.bioimageio.reproducibility_tolerance.output_ids contains '{out_id}', which is a discrete output tensor, but config.bioimageio.reproducibility_tolerance.relative_tolerance={rep_tol.relative_tolerance} is not 0." 

4053 ) 

4054 if discrete_values and rep_tol.absolute_tolerance != 0: 

4055 raise ValueError( 

4056 f"config.bioimageio.reproducibility_tolerance.output_ids contains '{out_id}', which is a discrete output tensor, but config.bioimageio.reproducibility_tolerance.absolute_tolerance={rep_tol.absolute_tolerance} is not 0." 

4057 ) 

4058 

4059 return self 

4060 

4061 @model_validator(mode="after") 

4062 def _add_default_cover(self) -> Self: 

4063 if not get_validation_context().perform_io_checks or self.covers: 

4064 return self 

4065 

4066 try: 

4067 generated_covers = generate_covers( 

4068 [ 

4069 (t, load_array(t.test_tensor)) 

4070 for t in self.inputs 

4071 if t.test_tensor is not None 

4072 ], 

4073 [ 

4074 (t, load_array(t.test_tensor)) 

4075 for t in self.outputs 

4076 if t.test_tensor is not None 

4077 ], 

4078 ) 

4079 except Exception as e: 

4080 issue_warning( 

4081 "Failed to generate cover image(s): {e}", 

4082 value=self.covers, 

4083 msg_context={"e": e}, 

4084 field="covers", 

4085 ) 

4086 else: 

4087 self.covers.extend(generated_covers) 

4088 

4089 return self 

4090 

4091 def get_input_test_arrays(self) -> list[NDArray[Any]]: 

4092 return self._get_test_arrays(self.inputs) 

4093 

4094 def get_output_test_arrays(self) -> list[NDArray[Any]]: 

4095 return self._get_test_arrays(self.outputs) 

4096 

4097 @staticmethod 

4098 def _get_test_arrays( 

4099 io_descr: Sequence[InputTensorDescr] | Sequence[OutputTensorDescr], 

4100 ): 

4101 ts: list[FileDescr] = [] 

4102 for d in io_descr: 

4103 if d.test_tensor is None: 

4104 raise ValueError( 

4105 f"Failed to get test arrays: description of '{d.id}' is missing a `test_tensor`." 

4106 ) 

4107 ts.append(d.test_tensor) 

4108 

4109 data = [load_array(t) for t in ts] 

4110 assert all(isinstance(d, np.ndarray) for d in data) 

4111 return data 

4112 

4113 @staticmethod 

4114 def get_batch_size(tensor_sizes: Mapping[TensorId, Mapping[AxisId, int]]) -> int: 

4115 batch_size = 1 

4116 tensor_with_batchsize: TensorId | None = None 

4117 for tid in tensor_sizes: 

4118 for aid, s in tensor_sizes[tid].items(): 

4119 if aid != BATCH_AXIS_ID or s == 1 or s == batch_size: 

4120 continue 

4121 

4122 if batch_size != 1: 

4123 assert tensor_with_batchsize is not None 

4124 raise ValueError( 

4125 f"batch size mismatch for tensors '{tensor_with_batchsize}' ({batch_size}) and '{tid}' ({s})" 

4126 ) 

4127 

4128 batch_size = s 

4129 tensor_with_batchsize = tid 

4130 

4131 return batch_size 

4132 

4133 def get_output_tensor_sizes( 

4134 self, input_sizes: Mapping[TensorId, Mapping[AxisId, int]] 

4135 ) -> dict[TensorId, dict[AxisId, int | _DataDepSize]]: 

4136 """Returns the tensor output sizes for given **input_sizes**. 

4137 Only if **input_sizes** has a valid input shape, the tensor output size is exact. 

4138 Otherwise it might be larger than the actual (valid) output""" 

4139 batch_size = self.get_batch_size(input_sizes) 

4140 ns = self.get_ns(input_sizes) 

4141 

4142 tensor_sizes = self.get_tensor_sizes(ns, batch_size=batch_size) 

4143 return tensor_sizes.outputs 

4144 

4145 def get_ns(self, input_sizes: Mapping[TensorId, Mapping[AxisId, int]]): 

4146 """get parameter `n` for each parameterized axis 

4147 such that the valid input size is >= the given input size""" 

4148 ret: dict[tuple[TensorId, AxisId], ParameterizedSize_N] = {} 

4149 axes = {t.id: {a.id: a for a in t.axes} for t in self.inputs} 

4150 for tid in input_sizes: 

4151 for aid, s in input_sizes[tid].items(): 

4152 size_descr = axes[tid][aid].size 

4153 if isinstance(size_descr, ParameterizedSize): 

4154 ret[(tid, aid)] = size_descr.get_n(s) 

4155 elif size_descr is None or isinstance(size_descr, (int, SizeReference)): 

4156 pass 

4157 else: 

4158 assert_never(size_descr) 

4159 

4160 return ret 

4161 

4162 def get_tensor_sizes( 

4163 self, 

4164 ns: Mapping[tuple[TensorId, AxisId], ParameterizedSize_N], 

4165 batch_size: int, 

4166 max_input_shape: Mapping[TensorId, Mapping[AxisId, int]] | None = None, 

4167 ) -> _TensorSizes: 

4168 max_axis_sizes: dict[tuple[TensorId, AxisId], int] = {} 

4169 for m, this_max_axis_sizes in (max_input_shape or {}).items(): 

4170 for a, s in this_max_axis_sizes.items(): 

4171 max_axis_sizes[(m, a)] = s 

4172 

4173 axis_sizes = self.get_axis_sizes( 

4174 ns, batch_size=batch_size, max_input_shape=max_axis_sizes 

4175 ) 

4176 return _TensorSizes( 

4177 { 

4178 t: { 

4179 aa: axis_sizes.inputs[(tt, aa)] 

4180 for tt, aa in axis_sizes.inputs 

4181 if tt == t 

4182 } 

4183 for t in {tt for tt, _ in axis_sizes.inputs} 

4184 }, 

4185 { 

4186 t: { 

4187 aa: axis_sizes.outputs[(tt, aa)] 

4188 for tt, aa in axis_sizes.outputs 

4189 if tt == t 

4190 } 

4191 for t in {tt for tt, _ in axis_sizes.outputs} 

4192 }, 

4193 ) 

4194 

4195 def get_axis_sizes( 

4196 self, 

4197 ns: Mapping[tuple[TensorId, AxisId], ParameterizedSize_N], 

4198 batch_size: int | None = None, 

4199 *, 

4200 max_input_shape: Mapping[tuple[TensorId, AxisId], int] | None = None, 

4201 ) -> _AxisSizes: 

4202 """Determine input and output block shape for scale factors **ns** 

4203 of parameterized input sizes. 

4204 

4205 Args: 

4206 ns: Scale factor `n` for each axis (keyed by (tensor_id, axis_id)) 

4207 that is parameterized as `size = min + n * step`. 

4208 batch_size: The desired size of the batch dimension. 

4209 If given **batch_size** overwrites any batch size present in 

4210 **max_input_shape**. Default 1. 

4211 max_input_shape: Limits the derived block shapes. 

4212 Each axis for which the input size, parameterized by `n`, is larger 

4213 than **max_input_shape** is set to the minimal value `n_min` for which 

4214 this is still true. 

4215 Use this for small input samples or large values of **ns**. 

4216 Or simply whenever you know the full input shape. 

4217 

4218 Returns: 

4219 Resolved axis sizes for model inputs and outputs. 

4220 """ 

4221 max_input_shape = max_input_shape or {} 

4222 if batch_size is None: 

4223 for (_t_id, a_id), s in max_input_shape.items(): 

4224 if a_id == BATCH_AXIS_ID: 

4225 batch_size = s 

4226 break 

4227 else: 

4228 batch_size = 1 

4229 

4230 all_axes = { 

4231 t.id: {a.id: a for a in t.axes} for t in chain(self.inputs, self.outputs) 

4232 } 

4233 

4234 inputs: dict[tuple[TensorId, AxisId], int] = {} 

4235 outputs: dict[tuple[TensorId, AxisId], int | _DataDepSize] = {} 

4236 

4237 def get_axis_size(a: InputAxis | OutputAxis): 

4238 if isinstance(a, BatchAxis): 

4239 if (t_descr.id, a.id) in ns: 

4240 logger.warning( 

4241 "Ignoring unexpected size increment factor (n) for batch axis" 

4242 + " of tensor '{}'.", 

4243 t_descr.id, 

4244 ) 

4245 return batch_size 

4246 elif isinstance(a.size, int): 

4247 if (t_descr.id, a.id) in ns: 

4248 logger.warning( 

4249 "Ignoring unexpected size increment factor (n) for fixed size" 

4250 + " axis '{}' of tensor '{}'.", 

4251 a.id, 

4252 t_descr.id, 

4253 ) 

4254 return a.size 

4255 elif isinstance(a.size, ParameterizedSize): 

4256 if (t_descr.id, a.id) not in ns: 

4257 raise ValueError( 

4258 "Size increment factor (n) missing for parametrized axis" 

4259 + f" '{a.id}' of tensor '{t_descr.id}'." 

4260 ) 

4261 n = ns[(t_descr.id, a.id)] 

4262 s_max = max_input_shape.get((t_descr.id, a.id)) 

4263 if s_max is not None: 

4264 n = min(n, a.size.get_n(s_max)) 

4265 

4266 return a.size.get_size(n) 

4267 

4268 elif isinstance(a.size, SizeReference): 

4269 if (t_descr.id, a.id) in ns: 

4270 logger.warning( 

4271 "Ignoring unexpected size increment factor (n) for axis '{}'" 

4272 + " of tensor '{}' with size reference.", 

4273 a.id, 

4274 t_descr.id, 

4275 ) 

4276 assert not isinstance(a, BatchAxis) 

4277 ref_axis = all_axes[a.size.tensor_id][a.size.axis_id] 

4278 assert not isinstance(ref_axis, BatchAxis) 

4279 ref_key = (a.size.tensor_id, a.size.axis_id) 

4280 ref_size = inputs.get(ref_key, outputs.get(ref_key)) 

4281 assert ref_size is not None, ref_key 

4282 assert not isinstance(ref_size, _DataDepSize), ref_key 

4283 return a.size.get_size( 

4284 axis=a, 

4285 ref_axis=ref_axis, 

4286 ref_size=ref_size, 

4287 ) 

4288 elif isinstance(a.size, DataDependentSize): 

4289 if (t_descr.id, a.id) in ns: 

4290 logger.warning( 

4291 "Ignoring unexpected increment factor (n) for data dependent" 

4292 + " size axis '{}' of tensor '{}'.", 

4293 a.id, 

4294 t_descr.id, 

4295 ) 

4296 return _DataDepSize(a.size.min, a.size.max) 

4297 else: 

4298 assert_never(a.size) 

4299 

4300 # first resolve all , but the `SizeReference` input sizes 

4301 for t_descr in self.inputs: 

4302 for a in t_descr.axes: 

4303 if not isinstance(a.size, SizeReference): 

4304 s = get_axis_size(a) 

4305 assert not isinstance(s, _DataDepSize) 

4306 inputs[t_descr.id, a.id] = s 

4307 

4308 # resolve all other input axis sizes 

4309 for t_descr in self.inputs: 

4310 for a in t_descr.axes: 

4311 if isinstance(a.size, SizeReference): 

4312 s = get_axis_size(a) 

4313 assert not isinstance(s, _DataDepSize) 

4314 inputs[t_descr.id, a.id] = s 

4315 

4316 # resolve all output axis sizes 

4317 for t_descr in self.outputs: 

4318 for a in t_descr.axes: 

4319 assert not isinstance(a.size, ParameterizedSize) 

4320 s = get_axis_size(a) 

4321 outputs[t_descr.id, a.id] = s 

4322 

4323 return _AxisSizes(inputs=inputs, outputs=outputs) 

4324 

4325 @model_validator(mode="before") 

4326 @classmethod 

4327 def _convert(cls, data: dict[str, Any]) -> dict[str, Any]: 

4328 cls.convert_from_old_format_wo_validation(data) 

4329 return data 

4330 

4331 @classmethod 

4332 def convert_from_old_format_wo_validation(cls, data: dict[str, Any]) -> None: 

4333 """Convert metadata following an older format version to this classes' format 

4334 without validating the result. 

4335 """ 

4336 if ( 

4337 data.get("type") == "model" 

4338 and isinstance(fv := data.get("format_version"), str) 

4339 and fv.count(".") == 2 

4340 ): 

4341 fv_parts = fv.split(".") 

4342 if any(not p.isdigit() for p in fv_parts): 

4343 return 

4344 

4345 fv_tuple = tuple(map(int, fv_parts)) 

4346 

4347 assert cls.implemented_format_version_tuple[0:2] == (0, 5) 

4348 if fv_tuple[:2] in ((0, 3), (0, 4)): 

4349 m04 = _ModelDescr_v0_4.load(data) 

4350 if isinstance(m04, InvalidDescr): 

4351 try: 

4352 updated = _model_conv.convert_as_dict( 

4353 m04 # pyright: ignore[reportArgumentType] 

4354 ) 

4355 except Exception as e: 

4356 logger.error( 

4357 "Failed to convert from invalid model 0.4 description." 

4358 + f"\nerror: {e}" 

4359 + "\nProceeding with model 0.5 validation without conversion." 

4360 ) 

4361 updated = None 

4362 else: 

4363 updated = _model_conv.convert_as_dict(m04) 

4364 

4365 if updated is not None: 

4366 data.clear() 

4367 data.update(updated) 

4368 

4369 elif fv_tuple[:2] == (0, 5): 

4370 # bump patch version 

4371 data["format_version"] = cls.implemented_format_version 

4372 

4373 if fv_tuple[:2] in ((0, 3), (0, 4)) or ( 

4374 fv_tuple[:2] == (0, 5) and fv_tuple[2] < 11 

4375 ): 

4376 convert_plain_covers_and_docs_and_icon(data) 

4377 

4378 

4379class _ModelConv(Converter[_ModelDescr_v0_4, ModelDescr]): 

4380 def _convert( 

4381 self, src: _ModelDescr_v0_4, tgt: type[ModelDescr | dict[str, Any]] 

4382 ) -> ModelDescr | dict[str, Any]: 

4383 name = "".join( 

4384 c if c in string.ascii_letters + string.digits + "_+- ()" else " " 

4385 for c in src.name 

4386 ) 

4387 

4388 def conv_authors(auths: Sequence[_Author_v0_4] | None): 

4389 conv = ( 

4390 _author_conv.convert if TYPE_CHECKING else _author_conv.convert_as_dict 

4391 ) 

4392 return None if auths is None else [conv(a) for a in auths] 

4393 

4394 if TYPE_CHECKING: 

4395 arch_file_conv = _arch_file_conv.convert 

4396 arch_lib_conv = _arch_lib_conv.convert 

4397 else: 

4398 arch_file_conv = _arch_file_conv.convert_as_dict 

4399 arch_lib_conv = _arch_lib_conv.convert_as_dict 

4400 

4401 input_size_refs = { 

4402 ipt.name: { 

4403 a: s 

4404 for a, s in zip( 

4405 ipt.axes, 

4406 ( 

4407 ipt.shape.min 

4408 if isinstance(ipt.shape, _ParameterizedInputShape_v0_4) 

4409 else ipt.shape 

4410 ), 

4411 ) 

4412 } 

4413 for ipt in src.inputs 

4414 if ipt.shape 

4415 } 

4416 output_size_refs = { 

4417 **{ 

4418 out.name: {a: s for a, s in zip(out.axes, out.shape)} 

4419 for out in src.outputs 

4420 if not isinstance(out.shape, _ImplicitOutputShape_v0_4) 

4421 }, 

4422 **input_size_refs, 

4423 } 

4424 

4425 return tgt( 

4426 attachments=( 

4427 [] 

4428 if src.attachments is None 

4429 else [FileDescr(source=f) for f in src.attachments.files] 

4430 ), 

4431 authors=[_author_conv.convert_as_dict(a) for a in src.authors], # pyright: ignore[reportArgumentType] 

4432 cite=[{"text": c.text, "doi": c.doi, "url": c.url} for c in src.cite], # pyright: ignore[reportArgumentType] 

4433 config=src.config, # pyright: ignore[reportArgumentType] 

4434 covers=[{"source": c} for c in src.covers], # pyright: ignore[reportArgumentType] 

4435 description=src.description, 

4436 documentation={"source": src.documentation} if src.documentation else None, # pyright: ignore[reportArgumentType] 

4437 format_version="0.5.14", 

4438 git_repo=src.git_repo, # pyright: ignore[reportArgumentType] 

4439 icon={"source": src.icon} if src.icon else None, # pyright: ignore[reportArgumentType] 

4440 id=None if src.id is None else ModelId(src.id), 

4441 id_emoji=src.id_emoji, 

4442 license=src.license, # type: ignore 

4443 links=src.links, 

4444 maintainers=[_maintainer_conv.convert_as_dict(m) for m in src.maintainers], # pyright: ignore[reportArgumentType] 

4445 name=name, 

4446 tags=src.tags, 

4447 type=src.type, 

4448 uploader=src.uploader, 

4449 version=src.version, 

4450 inputs=[ # pyright: ignore[reportArgumentType] 

4451 _input_tensor_conv.convert_as_dict(ipt, tt, st, input_size_refs) 

4452 for ipt, tt, st in zip( 

4453 src.inputs, 

4454 src.test_inputs, 

4455 src.sample_inputs or [None] * len(src.test_inputs), 

4456 ) 

4457 ], 

4458 outputs=[ # pyright: ignore[reportArgumentType] 

4459 _output_tensor_conv.convert_as_dict(out, tt, st, output_size_refs) 

4460 for out, tt, st in zip( 

4461 src.outputs, 

4462 src.test_outputs, 

4463 src.sample_outputs or [None] * len(src.test_outputs), 

4464 ) 

4465 ], 

4466 parent=( 

4467 None 

4468 if src.parent is None 

4469 else LinkedModel( 

4470 id=ModelId( 

4471 str(src.parent.id) 

4472 + ( 

4473 "" 

4474 if src.parent.version_number is None 

4475 else f"/{src.parent.version_number}" 

4476 ) 

4477 ) 

4478 ) 

4479 ), 

4480 training_data=( 

4481 None 

4482 if src.training_data is None 

4483 else ( 

4484 LinkedDataset( 

4485 id=DatasetId( 

4486 str(src.training_data.id) 

4487 + ( 

4488 "" 

4489 if src.training_data.version_number is None 

4490 else f"/{src.training_data.version_number}" 

4491 ) 

4492 ) 

4493 ) 

4494 if isinstance(src.training_data, LinkedDataset02) 

4495 else src.training_data 

4496 ) 

4497 ), 

4498 packaged_by=[_author_conv.convert_as_dict(a) for a in src.packaged_by], # pyright: ignore[reportArgumentType] 

4499 run_mode=src.run_mode, 

4500 timestamp=src.timestamp, 

4501 weights=(WeightsDescr if TYPE_CHECKING else dict)( 

4502 keras_hdf5=(w := src.weights.keras_hdf5) 

4503 and (KerasHdf5WeightsDescr if TYPE_CHECKING else dict)( 

4504 authors=conv_authors(w.authors), 

4505 source=w.source, 

4506 tensorflow_version=w.tensorflow_version or Version("1.15"), 

4507 parent=w.parent, 

4508 ), 

4509 onnx=(w := src.weights.onnx) 

4510 and (OnnxWeightsDescr if TYPE_CHECKING else dict)( 

4511 source=w.source, 

4512 authors=conv_authors(w.authors), 

4513 parent=w.parent, 

4514 opset_version=w.opset_version or 15, 

4515 ), 

4516 pytorch_state_dict=(w := src.weights.pytorch_state_dict) 

4517 and (PytorchStateDictWeightsDescr if TYPE_CHECKING else dict)( 

4518 source=w.source, 

4519 authors=conv_authors(w.authors), 

4520 parent=w.parent, 

4521 architecture=( 

4522 arch_file_conv( 

4523 w.architecture, 

4524 w.architecture_sha256, 

4525 w.kwargs, 

4526 ) 

4527 if isinstance(w.architecture, _CallableFromFile_v0_4) 

4528 else arch_lib_conv(w.architecture, w.kwargs) 

4529 ), 

4530 pytorch_version=w.pytorch_version or Version("1.10"), 

4531 dependencies=( 

4532 None 

4533 if w.dependencies is None 

4534 else (FileDescr if TYPE_CHECKING else dict)( 

4535 source=cast( 

4536 FileSource, 

4537 str(deps := w.dependencies)[ 

4538 ( 

4539 len("conda:") 

4540 if str(deps).startswith("conda:") 

4541 else 0 

4542 ) : 

4543 ], 

4544 ) 

4545 ) 

4546 ), 

4547 ), 

4548 tensorflow_js=(w := src.weights.tensorflow_js) 

4549 and (TensorflowJsWeightsDescr if TYPE_CHECKING else dict)( 

4550 source=w.source, 

4551 authors=conv_authors(w.authors), 

4552 parent=w.parent, 

4553 tensorflow_version=w.tensorflow_version or Version("1.15"), 

4554 ), 

4555 tensorflow_saved_model_bundle=( 

4556 w := src.weights.tensorflow_saved_model_bundle 

4557 ) 

4558 and (TensorflowSavedModelBundleWeightsDescr if TYPE_CHECKING else dict)( 

4559 authors=conv_authors(w.authors), 

4560 parent=w.parent, 

4561 source=w.source, 

4562 tensorflow_version=w.tensorflow_version or Version("1.15"), 

4563 dependencies=( 

4564 None 

4565 if w.dependencies is None 

4566 else (FileDescr if TYPE_CHECKING else dict)( 

4567 source=cast( 

4568 FileSource, 

4569 ( 

4570 str(w.dependencies)[len("conda:") :] 

4571 if str(w.dependencies).startswith("conda:") 

4572 else str(w.dependencies) 

4573 ), 

4574 ) 

4575 ) 

4576 ), 

4577 ), 

4578 torchscript=(w := src.weights.torchscript) 

4579 and (TorchscriptWeightsDescr if TYPE_CHECKING else dict)( 

4580 source=w.source, 

4581 authors=conv_authors(w.authors), 

4582 parent=w.parent, 

4583 pytorch_version=w.pytorch_version or Version("1.10"), 

4584 ), 

4585 ), 

4586 ) 

4587 

4588 

4589_model_conv = _ModelConv(_ModelDescr_v0_4, ModelDescr) 

4590 

4591 

4592# create better cover images for 3d data and non-image outputs 

4593def generate_covers( 

4594 inputs: Sequence[tuple[InputTensorDescr, NDArray[Any]]], 

4595 outputs: Sequence[tuple[OutputTensorDescr, NDArray[Any]]], 

4596) -> list[FileDescr]: 

4597 def squeeze( 

4598 data: NDArray[Any], axes: Sequence[AnyAxis] 

4599 ) -> tuple[NDArray[Any], list[AnyAxis]]: 

4600 """apply numpy.ndarray.squeeze while keeping track of the axis descriptions remaining""" 

4601 if data.ndim != len(axes): 

4602 raise ValueError( 

4603 f"tensor shape {data.shape} does not match described axes" 

4604 + f" {[a.id for a in axes]}" 

4605 ) 

4606 

4607 axes = [deepcopy(a) for a, s in zip(axes, data.shape) if s != 1] 

4608 return data.squeeze(), axes 

4609 

4610 def normalize( 

4611 data: NDArray[Any], axis: tuple[int, ...] | None, eps: float = 1e-7 

4612 ) -> NDArray[np.float32]: 

4613 data = data.astype("float32") 

4614 data -= data.min(axis=axis, keepdims=True) 

4615 data /= data.max(axis=axis, keepdims=True) + eps 

4616 return data 

4617 

4618 def to_2d_image(data: NDArray[Any], axes: Sequence[AnyAxis]): 

4619 original_shape = data.shape 

4620 original_axes = list(axes) 

4621 data, axes = squeeze(data, axes) 

4622 

4623 # take slice fom any batch or index axis if needed 

4624 # and convert the first channel axis and take a slice from any additional channel axes 

4625 slices: tuple[slice, ...] = () 

4626 ndim = data.ndim 

4627 ndim_need = 3 if any(isinstance(a, ChannelAxis) for a in axes) else 2 

4628 has_c_axis = False 

4629 for i, a in enumerate(axes): 

4630 s = data.shape[i] 

4631 assert s > 1 

4632 if ( 

4633 isinstance(a, (BatchAxis, IndexInputAxis, IndexOutputAxis)) 

4634 and ndim > ndim_need 

4635 ): 

4636 data = data[slices + (slice(s // 2 - 1, s // 2),)] 

4637 ndim -= 1 

4638 elif isinstance(a, ChannelAxis): 

4639 if has_c_axis: 

4640 # second channel axis 

4641 data = data[slices + (slice(0, 1),)] 

4642 ndim -= 1 

4643 else: 

4644 has_c_axis = True 

4645 if s == 2: 

4646 # visualize two channels with cyan and magenta 

4647 data = np.concatenate( 

4648 [ 

4649 data[slices + (slice(1, 2),)], 

4650 data[slices + (slice(0, 1),)], 

4651 ( 

4652 data[slices + (slice(0, 1),)] 

4653 + data[slices + (slice(1, 2),)] 

4654 ) 

4655 / 2, # TODO: take maximum instead? 

4656 ], 

4657 axis=i, 

4658 ) 

4659 elif data.shape[i] == 3: 

4660 pass # visualize 3 channels as RGB 

4661 else: 

4662 # visualize first 3 channels as RGB 

4663 data = data[slices + (slice(3),)] 

4664 

4665 assert data.shape[i] == 3 

4666 

4667 slices += (slice(None),) 

4668 

4669 data, axes = squeeze(data, axes) 

4670 assert len(axes) == ndim 

4671 # take slice from z axis if needed 

4672 slices = () 

4673 if ndim > ndim_need: 

4674 for i, a in enumerate(axes): 

4675 s = data.shape[i] 

4676 if a.id == AxisId("z"): 

4677 data = data[slices + (slice(s // 2 - 1, s // 2),)] 

4678 data, axes = squeeze(data, axes) 

4679 ndim -= 1 

4680 break 

4681 

4682 slices += (slice(None),) 

4683 

4684 # take slice from any space or time axis 

4685 slices = () 

4686 

4687 for i, a in enumerate(axes): 

4688 if ndim <= ndim_need: 

4689 break 

4690 

4691 s = data.shape[i] 

4692 assert s > 1 

4693 if isinstance( 

4694 a, (SpaceInputAxis, SpaceOutputAxis, TimeInputAxis, TimeOutputAxis) 

4695 ): 

4696 data = data[slices + (slice(s // 2 - 1, s // 2),)] 

4697 ndim -= 1 

4698 

4699 slices += (slice(None),) 

4700 

4701 del slices 

4702 data, axes = squeeze(data, axes) 

4703 assert len(axes) == ndim 

4704 

4705 if (has_c_axis and ndim != 3) or (not has_c_axis and ndim != 2): 

4706 raise ValueError( 

4707 f"Failed to construct cover image from shape {original_shape} with axes {[a.id for a in original_axes]}." 

4708 ) 

4709 

4710 if not has_c_axis: 

4711 assert ndim == 2 

4712 data = np.repeat(data[:, :, None], 3, axis=2) 

4713 axes.append(ChannelAxis(channel_names=list("RGB"))) 

4714 ndim += 1 

4715 

4716 assert ndim == 3 

4717 

4718 # transpose axis order such that longest axis comes first... 

4719 axis_order: list[int] = [int(i) for i in np.argsort(list(data.shape))] 

4720 axis_order.reverse() 

4721 # ... and channel axis is last 

4722 c = next(i for i in range(3) if isinstance(axes[i], ChannelAxis)) 

4723 axis_order.append(axis_order.pop(c)) 

4724 axes = [axes[ao] for ao in axis_order] 

4725 data = data.transpose(axis_order) 

4726 

4727 # h, w = data.shape[:2] 

4728 # if h / w in (1.0 or 2.0): 

4729 # pass 

4730 # elif h / w < 2: 

4731 # TODO: enforce 2:1 or 1:1 aspect ratio for generated cover images 

4732 

4733 norm_along = ( 

4734 tuple(i for i, a in enumerate(axes) if a.type in ("space", "time")) or None 

4735 ) 

4736 # normalize the data and map to 8 bit 

4737 data = normalize(data, norm_along) 

4738 data = (data * 255).astype("uint8") 

4739 

4740 return data 

4741 

4742 def create_diagonal_split_image(im0: NDArray[Any], im1: NDArray[Any]): 

4743 assert im0.dtype == im1.dtype == np.uint8 

4744 assert im0.shape == im1.shape 

4745 assert im0.ndim == 3 

4746 N, M, C = im0.shape 

4747 assert C == 3 

4748 out = np.ones((N, M, C), dtype="uint8") 

4749 for c in range(C): 

4750 outc = np.tril(im0[..., c]) 

4751 mask = outc == 0 

4752 outc[mask] = np.triu(im1[..., c])[mask] 

4753 out[..., c] = outc 

4754 

4755 return out 

4756 

4757 if not inputs: 

4758 raise ValueError("Missing test input tensor for cover generation.") 

4759 

4760 if not outputs: 

4761 raise ValueError("Missing test output tensor for cover generation.") 

4762 

4763 ipt_descr, ipt = inputs[0] 

4764 out_descr, out = outputs[0] 

4765 

4766 ipt_img = to_2d_image(ipt, ipt_descr.axes) 

4767 out_img = to_2d_image(out, out_descr.axes) 

4768 

4769 cover_folder = Path(mkdtemp()) 

4770 if ipt_img.shape == out_img.shape: 

4771 covers = [cover_folder / "cover.png"] 

4772 imwrite(covers[0], create_diagonal_split_image(ipt_img, out_img)) 

4773 else: 

4774 covers = [cover_folder / "input.png", cover_folder / "output.png"] 

4775 imwrite(covers[0], ipt_img) 

4776 imwrite(covers[1], out_img) 

4777 

4778 return [FileDescr(source=c) for c in covers]