Coverage for src/bioimageio/spec/model/v0_4.py: 91%

596 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 

4from typing import ( 

5 TYPE_CHECKING, 

6 Any, 

7 Callable, 

8 ClassVar, 

9 Dict, 

10 List, 

11 Literal, 

12 Sequence, 

13 Union, 

14 cast, 

15) 

16 

17import numpy as np 

18from annotated_types import Ge, Interval, MaxLen, MinLen, MultipleOf 

19from numpy.typing import NDArray 

20from pydantic import ( 

21 AllowInfNan, 

22 Discriminator, 

23 Field, 

24 RootModel, 

25 SerializationInfo, 

26 SerializerFunctionWrapHandler, 

27 StringConstraints, 

28 TypeAdapter, 

29 ValidationInfo, 

30 WrapSerializer, 

31 field_validator, 

32 model_validator, 

33) 

34from typing_extensions import Annotated, Self, assert_never, get_args 

35 

36from .._internal.common_nodes import ( 

37 KwargsNode, 

38 Node, 

39 NodeWithExplicitlySetFields, 

40) 

41from .._internal.constants import SHA256_HINT 

42from .._internal.field_validation import validate_unique_entries 

43from .._internal.field_warning import issue_warning, warn 

44from .._internal.io import BioimageioYamlContent, WithSuffix 

45from .._internal.io import FileDescr as FileDescr 

46from .._internal.io_basics import Sha256 as Sha256 

47from .._internal.io_packaging import FileSource_package, include_in_package 

48from .._internal.io_utils import load_array 

49from .._internal.packaging_context import packaging_context_var 

50from .._internal.types import Datetime as Datetime 

51from .._internal.types import FileSource, LowerCaseIdentifier 

52from .._internal.types import Identifier as Identifier 

53from .._internal.types import LicenseId as LicenseId 

54from .._internal.types import NotEmpty as NotEmpty 

55from .._internal.url import HttpUrl as HttpUrl 

56from .._internal.validated_string_with_inner_node import ValidatedStringWithInnerNode 

57from .._internal.validator_annotations import AfterValidator, RestrictCharacters 

58from .._internal.version_type import Version as Version 

59from .._internal.warning_levels import ALERT, INFO 

60from ..dataset.v0_2 import VALID_COVER_IMAGE_EXTENSIONS as VALID_COVER_IMAGE_EXTENSIONS 

61from ..dataset.v0_2 import DatasetDescr as DatasetDescr 

62from ..dataset.v0_2 import LinkedDataset as LinkedDataset 

63from ..generic.v0_2 import AttachmentsDescr as AttachmentsDescr 

64from ..generic.v0_2 import Author as Author 

65from ..generic.v0_2 import BadgeDescr as BadgeDescr 

66from ..generic.v0_2 import CiteEntry as CiteEntry 

67from ..generic.v0_2 import Doi as Doi 

68from ..generic.v0_2 import GenericModelDescrBase 

69from ..generic.v0_2 import LinkedResource as LinkedResource 

70from ..generic.v0_2 import Maintainer as Maintainer 

71from ..generic.v0_2 import OrcidId as OrcidId 

72from ..generic.v0_2 import RelativeFilePath as RelativeFilePath 

73from ..generic.v0_2 import ResourceId as ResourceId 

74from ..generic.v0_2 import Uploader as Uploader 

75from ._v0_4_converter import convert_from_older_format 

76 

77 

78class ModelId(ResourceId): 

79 pass 

80 

81 

82AxesStr = Annotated[ 

83 str, RestrictCharacters("bitczyx"), AfterValidator(validate_unique_entries) 

84] 

85AxesInCZYX = Annotated[ 

86 str, RestrictCharacters("czyx"), AfterValidator(validate_unique_entries) 

87] 

88 

89PostprocessingName = Literal[ 

90 "binarize", 

91 "clip", 

92 "scale_linear", 

93 "sigmoid", 

94 "zero_mean_unit_variance", 

95 "scale_range", 

96 "scale_mean_variance", 

97] 

98PreprocessingName = Literal[ 

99 "binarize", 

100 "clip", 

101 "scale_linear", 

102 "sigmoid", 

103 "zero_mean_unit_variance", 

104 "scale_range", 

105] 

106 

107 

108class TensorName(LowerCaseIdentifier): 

109 pass 

110 

111 

112class CallableFromDepencencyNode(Node): 

113 _submodule_adapter: ClassVar[TypeAdapter[Identifier]] = TypeAdapter(Identifier) 

114 

115 module_name: str 

116 """The Python module that implements **callable_name**.""" 

117 

118 @field_validator("module_name", mode="after") 

119 def _check_submodules(cls, module_name: str) -> str: 

120 for submod in module_name.split("."): 

121 _ = cls._submodule_adapter.validate_python(submod) 

122 

123 return module_name 

124 

125 callable_name: Identifier 

126 """The callable Python identifier implemented in module **module_name**.""" 

127 

128 

129class CallableFromDepencency(ValidatedStringWithInnerNode[CallableFromDepencencyNode]): 

130 _inner_node_class = CallableFromDepencencyNode 

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

132 Annotated[ 

133 str, 

134 StringConstraints(strip_whitespace=True, pattern=r"^.+\..+$"), 

135 ] 

136 ] 

137 

138 @classmethod 

139 def _get_data(cls, valid_string_data: str): 

140 *mods, callname = valid_string_data.split(".") 

141 return {"module_name": ".".join(mods), "callable_name": callname} 

142 

143 @property 

144 def module_name(self): 

145 """The Python module that implements **callable_name**.""" 

146 return self._inner_node.module_name 

147 

148 @property 

149 def callable_name(self): 

150 """The callable Python identifier implemented in module **module_name**.""" 

151 return self._inner_node.callable_name 

152 

153 

154class CallableFromFileNode(Node): 

155 source_file: Annotated[ 

156 RelativeFilePath | HttpUrl, 

157 Field(union_mode="left_to_right"), 

158 include_in_package, 

159 ] 

160 """The Python source file that implements **callable_name**.""" 

161 callable_name: Identifier 

162 """The callable Python identifier implemented in **source_file**.""" 

163 

164 

165class CallableFromFile(ValidatedStringWithInnerNode[CallableFromFileNode]): 

166 _inner_node_class = CallableFromFileNode 

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

168 Annotated[ 

169 str, 

170 StringConstraints(strip_whitespace=True, pattern=r"^.+:.+$"), 

171 ] 

172 ] 

173 

174 @classmethod 

175 def _get_data(cls, valid_string_data: str): 

176 *file_parts, callname = valid_string_data.split(":") 

177 return {"source_file": ":".join(file_parts), "callable_name": callname} 

178 

179 @property 

180 def source_file(self): 

181 """The Python source file that implements **callable_name**.""" 

182 return self._inner_node.source_file 

183 

184 @property 

185 def callable_name(self): 

186 """The callable Python identifier implemented in **source_file**.""" 

187 return self._inner_node.callable_name 

188 

189 

190CustomCallable = Annotated[ 

191 Union[CallableFromFile, CallableFromDepencency], Field(union_mode="left_to_right") 

192] 

193 

194 

195class DependenciesNode(Node): 

196 manager: Annotated[NotEmpty[str], Field(examples=["conda", "maven", "pip"])] 

197 """Dependency manager""" 

198 

199 file: FileSource_package 

200 """Dependency file""" 

201 

202 

203class Dependencies(ValidatedStringWithInnerNode[DependenciesNode]): 

204 _inner_node_class = DependenciesNode 

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

206 Annotated[ 

207 str, 

208 StringConstraints(strip_whitespace=True, pattern=r"^.+:.+$"), 

209 ] 

210 ] 

211 

212 @classmethod 

213 def _get_data(cls, valid_string_data: str): 

214 manager, *file_parts = valid_string_data.split(":") 

215 return {"manager": manager, "file": ":".join(file_parts)} 

216 

217 @property 

218 def manager(self): 

219 """Dependency manager""" 

220 return self._inner_node.manager 

221 

222 @property 

223 def file(self): 

224 """Dependency file""" 

225 return self._inner_node.file 

226 

227 

228WeightsFormat = Literal[ 

229 "keras_hdf5", 

230 "onnx", 

231 "pytorch_state_dict", 

232 "tensorflow_js", 

233 "tensorflow_saved_model_bundle", 

234 "torchscript", 

235] 

236 

237 

238class WeightsEntryDescrBase(FileDescr): 

239 type: ClassVar[WeightsFormat] 

240 weights_format_name: ClassVar[str] # human readable 

241 

242 source: FileSource_package 

243 """The weights file.""" 

244 

245 attachments: Annotated[ 

246 AttachmentsDescr | None, 

247 warn(None, "Weights entry depends on additional attachments.", ALERT), 

248 ] = None 

249 """Attachments that are specific to this weights entry.""" 

250 

251 authors: list[Author] | None = None 

252 """Authors 

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

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

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

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

257 """ 

258 

259 dependencies: Annotated[ 

260 Dependencies | None, 

261 warn( 

262 None, 

263 "Custom dependencies ({value}) specified. Avoid this whenever possible " 

264 + "to allow execution in a wider range of software environments.", 

265 ), 

266 Field( 

267 examples=[ 

268 "conda:environment.yaml", 

269 "maven:./pom.xml", 

270 "pip:./requirements.txt", 

271 ] 

272 ), 

273 ] = None 

274 """Dependency manager and dependency file, specified as `<dependency manager>:<relative file path>`.""" 

275 

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

277 None 

278 ) 

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

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

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

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

283 need to have this field.""" 

284 

285 @model_validator(mode="after") 

286 def check_parent_is_not_self(self) -> Self: 

287 if self.type == self.parent: 

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

289 

290 return self 

291 

292 

293class KerasHdf5WeightsDescr(WeightsEntryDescrBase): 

294 type: ClassVar[WeightsFormat] = "keras_hdf5" 

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

296 tensorflow_version: Version | None = None 

297 """TensorFlow version used to create these weights""" 

298 

299 @field_validator("tensorflow_version", mode="after") 

300 @classmethod 

301 def _tfv(cls, value: Any): 

302 if value is None: 

303 issue_warning( 

304 "missing. Please specify the TensorFlow version" 

305 + " these weights were created with.", 

306 value=value, 

307 severity=ALERT, 

308 field="tensorflow_version", 

309 ) 

310 return value 

311 

312 

313class OnnxWeightsDescr(WeightsEntryDescrBase): 

314 type: ClassVar[WeightsFormat] = "onnx" 

315 weights_format_name: ClassVar[str] = "ONNX" 

316 opset_version: Annotated[int, Ge(7)] | None = None 

317 """ONNX opset version""" 

318 

319 @field_validator("opset_version", mode="after") 

320 @classmethod 

321 def _ov(cls, value: Any): 

322 if value is None: 

323 issue_warning( 

324 "Missing ONNX opset version (aka ONNX opset number). " 

325 + "Please specify the ONNX opset version these weights were created" 

326 + " with.", 

327 value=value, 

328 severity=ALERT, 

329 field="opset_version", 

330 ) 

331 return value 

332 

333 

334class PytorchStateDictWeightsDescr(WeightsEntryDescrBase): 

335 type: ClassVar[WeightsFormat] = "pytorch_state_dict" 

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

337 architecture: CustomCallable = Field( 

338 examples=["my_function.py:MyNetworkClass", "my_module.submodule.get_my_model"] 

339 ) 

340 """callable returning a torch.nn.Module instance. 

341 Local implementation: `<relative path to file>:<identifier of implementation within the file>`. 

342 Implementation in a dependency: `<dependency-package>.<[dependency-module]>.<identifier>`.""" 

343 

344 architecture_sha256: Annotated[ 

345 Sha256 | None, 

346 Field( 

347 description=( 

348 "The SHA256 of the architecture source file, if the architecture is not" 

349 " defined in a module listed in `dependencies`\n" 

350 ) 

351 + SHA256_HINT, 

352 ), 

353 ] = None 

354 """The SHA256 of the architecture source file, 

355 if the architecture is not defined in a module listed in `dependencies`""" 

356 

357 @model_validator(mode="after") 

358 def check_architecture_sha256(self) -> Self: 

359 if isinstance(self.architecture, CallableFromFile): 

360 if self.architecture_sha256 is None: 

361 raise ValueError( 

362 "Missing required `architecture_sha256` for `architecture` with" 

363 + " source file." 

364 ) 

365 elif self.architecture_sha256 is not None: 

366 raise ValueError( 

367 "Got `architecture_sha256` for architecture that does not have a source" 

368 + " file." 

369 ) 

370 

371 return self 

372 

373 kwargs: dict[str, Any] = Field( 

374 default_factory=cast(Callable[[], Dict[str, Any]], dict) 

375 ) 

376 """key word arguments for the `architecture` callable""" 

377 

378 pytorch_version: Version | None = None 

379 """Version of the PyTorch library used. 

380 If `depencencies` is specified it should include pytorch and the verison has to match. 

381 (`dependencies` overrules `pytorch_version`)""" 

382 

383 @field_validator("pytorch_version", mode="after") 

384 @classmethod 

385 def _ptv(cls, value: Any): 

386 if value is None: 

387 issue_warning( 

388 "missing. Please specify the PyTorch version these" 

389 + " PyTorch state dict weights were created with.", 

390 value=value, 

391 severity=ALERT, 

392 field="pytorch_version", 

393 ) 

394 return value 

395 

396 

397class TorchscriptWeightsDescr(WeightsEntryDescrBase): 

398 type: ClassVar[WeightsFormat] = "torchscript" 

399 weights_format_name: ClassVar[str] = "TorchScript" 

400 pytorch_version: Version | None = None 

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

402 

403 @field_validator("pytorch_version", mode="after") 

404 @classmethod 

405 def _ptv(cls, value: Any): 

406 if value is None: 

407 issue_warning( 

408 "missing. Please specify the PyTorch version these" 

409 + " Torchscript weights were created with.", 

410 value=value, 

411 severity=ALERT, 

412 field="pytorch_version", 

413 ) 

414 return value 

415 

416 

417class TensorflowJsWeightsDescr(WeightsEntryDescrBase): 

418 type: ClassVar[WeightsFormat] = "tensorflow_js" 

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

420 tensorflow_version: Version | None = None 

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

422 

423 @field_validator("tensorflow_version", mode="after") 

424 @classmethod 

425 def _tfv(cls, value: Any): 

426 if value is None: 

427 issue_warning( 

428 "missing. Please specify the TensorFlow version" 

429 + " these TensorflowJs weights were created with.", 

430 value=value, 

431 severity=ALERT, 

432 field="tensorflow_version", 

433 ) 

434 return value 

435 

436 source: FileSource_package 

437 """The multi-file weights. 

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

439 

440 

441class TensorflowSavedModelBundleWeightsDescr(WeightsEntryDescrBase): 

442 type: ClassVar[WeightsFormat] = "tensorflow_saved_model_bundle" 

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

444 tensorflow_version: Version | None = None 

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

446 

447 @field_validator("tensorflow_version", mode="after") 

448 @classmethod 

449 def _tfv(cls, value: Any): 

450 if value is None: 

451 issue_warning( 

452 "missing. Please specify the TensorFlow version" 

453 + " these Tensorflow saved model bundle weights were created with.", 

454 value=value, 

455 severity=ALERT, 

456 field="tensorflow_version", 

457 ) 

458 return value 

459 

460 

461class WeightsDescr(Node): 

462 keras_hdf5: KerasHdf5WeightsDescr | None = None 

463 onnx: OnnxWeightsDescr | None = None 

464 pytorch_state_dict: PytorchStateDictWeightsDescr | None = None 

465 tensorflow_js: TensorflowJsWeightsDescr | None = None 

466 tensorflow_saved_model_bundle: TensorflowSavedModelBundleWeightsDescr | None = None 

467 torchscript: TorchscriptWeightsDescr | None = None 

468 

469 @model_validator(mode="after") 

470 def check_one_entry(self) -> Self: 

471 if all( 

472 entry is None 

473 for entry in [ 

474 self.keras_hdf5, 

475 self.onnx, 

476 self.pytorch_state_dict, 

477 self.tensorflow_js, 

478 self.tensorflow_saved_model_bundle, 

479 self.torchscript, 

480 ] 

481 ): 

482 raise ValueError("Missing weights entry") 

483 

484 return self 

485 

486 def __getitem__( 

487 self, 

488 key: WeightsFormat, 

489 ): 

490 if key == "keras_hdf5": 

491 ret = self.keras_hdf5 

492 elif key == "onnx": 

493 ret = self.onnx 

494 elif key == "pytorch_state_dict": 

495 ret = self.pytorch_state_dict 

496 elif key == "tensorflow_js": 

497 ret = self.tensorflow_js 

498 elif key == "tensorflow_saved_model_bundle": 

499 ret = self.tensorflow_saved_model_bundle 

500 elif key == "torchscript": 

501 ret = self.torchscript 

502 else: 

503 raise KeyError(key) 

504 

505 if ret is None: 

506 raise KeyError(key) 

507 

508 return ret 

509 

510 @property 

511 def available_formats(self): 

512 return { 

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

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

515 **( 

516 {} 

517 if self.pytorch_state_dict is None 

518 else {"pytorch_state_dict": self.pytorch_state_dict} 

519 ), 

520 **( 

521 {} 

522 if self.tensorflow_js is None 

523 else {"tensorflow_js": self.tensorflow_js} 

524 ), 

525 **( 

526 {} 

527 if self.tensorflow_saved_model_bundle is None 

528 else { 

529 "tensorflow_saved_model_bundle": self.tensorflow_saved_model_bundle 

530 } 

531 ), 

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

533 } 

534 

535 @property 

536 def missing_formats(self): 

537 return { 

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

539 } 

540 

541 

542class ParameterizedInputShape(Node): 

543 """A sequence of valid shapes given by `shape_k = min + k * step for k in {0, 1, ...}`.""" 

544 

545 min: NotEmpty[list[int]] 

546 """The minimum input shape""" 

547 

548 step: NotEmpty[list[int]] 

549 """The minimum shape change""" 

550 

551 def __len__(self) -> int: 

552 return len(self.min) 

553 

554 @model_validator(mode="after") 

555 def matching_lengths(self) -> Self: 

556 if len(self.min) != len(self.step): 

557 raise ValueError("`min` and `step` required to have the same length") 

558 

559 return self 

560 

561 

562class ImplicitOutputShape(Node): 

563 """Output tensor shape depending on an input tensor shape. 

564 `shape(output_tensor) = shape(input_tensor) * scale + 2 * offset`""" 

565 

566 reference_tensor: TensorName 

567 """Name of the reference tensor.""" 

568 

569 scale: NotEmpty[list[float | None]] 

570 """output_pix/input_pix for each dimension. 

571 'null' values indicate new dimensions, whose length is defined by 2*`offset`""" 

572 

573 offset: NotEmpty[list[int | Annotated[float, MultipleOf(0.5)]]] 

574 """Position of origin wrt to input.""" 

575 

576 def __len__(self) -> int: 

577 return len(self.scale) 

578 

579 @model_validator(mode="after") 

580 def matching_lengths(self) -> Self: 

581 if len(self.scale) != len(self.offset): 

582 raise ValueError( 

583 f"scale {self.scale} has to have same length as offset {self.offset}!" 

584 ) 

585 # if we have an expanded dimension, make sure that it's offet is not zero 

586 for sc, off in zip(self.scale, self.offset): 

587 if sc is None and not off: 

588 raise ValueError("`offset` must not be zero if `scale` is none/zero") 

589 

590 return self 

591 

592 

593class TensorDescrBase(Node): 

594 name: TensorName 

595 """Tensor name. No duplicates are allowed.""" 

596 

597 description: str = "" 

598 

599 axes: AxesStr 

600 """Axes identifying characters. Same length and order as the axes in `shape`. 

601 | axis | description | 

602 | --- | --- | 

603 | b | batch (groups multiple samples) | 

604 | i | instance/index/element | 

605 | t | time | 

606 | c | channel | 

607 | z | spatial dimension z | 

608 | y | spatial dimension y | 

609 | x | spatial dimension x | 

610 """ 

611 

612 data_range: ( 

613 tuple[Annotated[float, AllowInfNan(True)], Annotated[float, AllowInfNan(True)]] 

614 | None 

615 ) = None 

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

617 If not specified, the full data range that can be expressed in `data_type` is allowed.""" 

618 

619 

620class BinarizeKwargs(KwargsNode): 

621 """key word arguments for `BinarizeDescr`""" 

622 

623 threshold: float 

624 """The fixed threshold""" 

625 

626 

627class BinarizeDescr(NodeWithExplicitlySetFields): 

628 """BinarizeDescr the tensor with a fixed `BinarizeKwargs.threshold`. 

629 Values above the threshold will be set to one, values below the threshold to zero. 

630 """ 

631 

632 implemented_name: ClassVar[Literal["binarize"]] = "binarize" 

633 if TYPE_CHECKING: 

634 name: Literal["binarize"] = "binarize" 

635 else: 

636 name: Literal["binarize"] 

637 

638 kwargs: BinarizeKwargs 

639 

640 

641class ClipKwargs(KwargsNode): 

642 """key word arguments for `ClipDescr`""" 

643 

644 min: float 

645 """minimum value for clipping""" 

646 max: float 

647 """maximum value for clipping""" 

648 

649 

650class ClipDescr(NodeWithExplicitlySetFields): 

651 """Clip tensor values to a range. 

652 

653 Set tensor values below `ClipKwargs.min` to `ClipKwargs.min` 

654 and above `ClipKwargs.max` to `ClipKwargs.max`. 

655 """ 

656 

657 implemented_name: ClassVar[Literal["clip"]] = "clip" 

658 if TYPE_CHECKING: 

659 name: Literal["clip"] = "clip" 

660 else: 

661 name: Literal["clip"] 

662 

663 kwargs: ClipKwargs 

664 

665 

666class ScaleLinearKwargs(KwargsNode): 

667 """key word arguments for `ScaleLinearDescr`""" 

668 

669 axes: Annotated[AxesInCZYX | None, Field(examples=["xy"])] = None 

670 """The subset of axes to scale jointly. 

671 For example xy to scale the two image axes for 2d data jointly.""" 

672 

673 gain: float | list[float] = 1.0 

674 """multiplicative factor""" 

675 

676 offset: float | list[float] = 0.0 

677 """additive term""" 

678 

679 @model_validator(mode="after") 

680 def either_gain_or_offset(self) -> Self: 

681 if ( 

682 self.gain == 1.0 

683 or isinstance(self.gain, list) 

684 and all(g == 1.0 for g in self.gain) 

685 ) and ( 

686 self.offset == 0.0 

687 or isinstance(self.offset, list) 

688 and all(off == 0.0 for off in self.offset) 

689 ): 

690 raise ValueError( 

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

692 + " 0.0." 

693 ) 

694 

695 return self 

696 

697 

698class ScaleLinearDescr(NodeWithExplicitlySetFields): 

699 """Fixed linear scaling.""" 

700 

701 implemented_name: ClassVar[Literal["scale_linear"]] = "scale_linear" 

702 if TYPE_CHECKING: 

703 name: Literal["scale_linear"] = "scale_linear" 

704 else: 

705 name: Literal["scale_linear"] 

706 

707 kwargs: ScaleLinearKwargs 

708 

709 

710class SigmoidDescr(NodeWithExplicitlySetFields): 

711 """The logistic sigmoid funciton, a.k.a. expit function.""" 

712 

713 implemented_name: ClassVar[Literal["sigmoid"]] = "sigmoid" 

714 if TYPE_CHECKING: 

715 name: Literal["sigmoid"] = "sigmoid" 

716 else: 

717 name: Literal["sigmoid"] 

718 

719 @property 

720 def kwargs(self) -> KwargsNode: 

721 """empty kwargs""" 

722 return KwargsNode() 

723 

724 

725class ZeroMeanUnitVarianceKwargs(KwargsNode): 

726 """key word arguments for `ZeroMeanUnitVarianceDescr`""" 

727 

728 mode: Literal["fixed", "per_dataset", "per_sample"] = "fixed" 

729 """Mode for computing mean and variance. 

730 | mode | description | 

731 | ----------- | ------------------------------------ | 

732 | fixed | Fixed values for mean and variance | 

733 | per_dataset | Compute for the entire dataset | 

734 | per_sample | Compute for each sample individually | 

735 """ 

736 axes: Annotated[AxesInCZYX, Field(examples=["xy"])] 

737 """The subset of axes to normalize jointly. 

738 For example `xy` to normalize the two image axes for 2d data jointly.""" 

739 

740 mean: Annotated[ 

741 float | NotEmpty[list[float]] | None, Field(examples=[(1.1, 2.2, 3.3)]) 

742 ] = None 

743 """The mean value(s) to use for `mode: fixed`. 

744 For example `[1.1, 2.2, 3.3]` in the case of a 3 channel image with `axes: xy`.""" 

745 # todo: check if means match input axes (for mode 'fixed') 

746 

747 std: Annotated[ 

748 float | NotEmpty[list[float]] | None, Field(examples=[(0.1, 0.2, 0.3)]) 

749 ] = None 

750 """The standard deviation values to use for `mode: fixed`. Analogous to mean.""" 

751 

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

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

754 

755 @model_validator(mode="after") 

756 def mean_and_std_match_mode(self) -> Self: 

757 if self.mode == "fixed" and (self.mean is None or self.std is None): 

758 raise ValueError("`mean` and `std` are required for `mode: fixed`.") 

759 elif self.mode != "fixed" and (self.mean is not None or self.std is not None): 

760 raise ValueError(f"`mean` and `std` not allowed for `mode: {self.mode}`") 

761 

762 return self 

763 

764 

765class ZeroMeanUnitVarianceDescr(NodeWithExplicitlySetFields): 

766 """Subtract mean and divide by variance.""" 

767 

768 implemented_name: ClassVar[Literal["zero_mean_unit_variance"]] = ( 

769 "zero_mean_unit_variance" 

770 ) 

771 if TYPE_CHECKING: 

772 name: Literal["zero_mean_unit_variance"] = "zero_mean_unit_variance" 

773 else: 

774 name: Literal["zero_mean_unit_variance"] 

775 

776 kwargs: ZeroMeanUnitVarianceKwargs 

777 

778 

779class ScaleRangeKwargs(KwargsNode): 

780 """key word arguments for `ScaleRangeDescr` 

781 

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

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

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

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

786 normalized values to a range. 

787 """ 

788 

789 mode: Literal["per_dataset", "per_sample"] 

790 """Mode for computing percentiles. 

791 | mode | description | 

792 | ----------- | ------------------------------------ | 

793 | per_dataset | compute for the entire dataset | 

794 | per_sample | compute for each sample individually | 

795 """ 

796 axes: Annotated[AxesInCZYX, Field(examples=["xy"])] 

797 """The subset of axes to normalize jointly. 

798 For example xy to normalize the two image axes for 2d data jointly.""" 

799 

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

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

802 

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

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

805 Has to be bigger than `min_percentile`. 

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

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

808 

809 @model_validator(mode="after") 

810 def min_smaller_max(self, info: ValidationInfo) -> Self: 

811 if self.min_percentile >= self.max_percentile: 

812 raise ValueError( 

813 f"min_percentile {self.min_percentile} >= max_percentile" 

814 + f" {self.max_percentile}" 

815 ) 

816 

817 return self 

818 

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

820 """Epsilon for numeric stability. 

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

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

823 

824 reference_tensor: TensorName | None = None 

825 """Tensor name to compute the percentiles from. Default: The tensor itself. 

826 For any tensor in `inputs` only input tensor references are allowed. 

827 For a tensor in `outputs` only input tensor refereences are allowed if `mode: per_dataset`""" 

828 

829 

830class ScaleRangeDescr(NodeWithExplicitlySetFields): 

831 """Scale with percentiles.""" 

832 

833 implemented_name: ClassVar[Literal["scale_range"]] = "scale_range" 

834 if TYPE_CHECKING: 

835 name: Literal["scale_range"] = "scale_range" 

836 else: 

837 name: Literal["scale_range"] 

838 

839 kwargs: ScaleRangeKwargs 

840 

841 

842class ScaleMeanVarianceKwargs(KwargsNode): 

843 """key word arguments for `ScaleMeanVarianceDescr`""" 

844 

845 mode: Literal["per_dataset", "per_sample"] 

846 """Mode for computing mean and variance. 

847 | mode | description | 

848 | ----------- | ------------------------------------ | 

849 | per_dataset | Compute for the entire dataset | 

850 | per_sample | Compute for each sample individually | 

851 """ 

852 

853 reference_tensor: TensorName 

854 """Name of tensor to match.""" 

855 

856 axes: Annotated[AxesInCZYX | None, Field(examples=["xy"])] = None 

857 """The subset of axes to scale jointly. 

858 For example xy to normalize the two image axes for 2d data jointly. 

859 Default: scale all non-batch axes jointly.""" 

860 

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

862 """Epsilon for numeric stability: 

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

864 

865 

866class ScaleMeanVarianceDescr(NodeWithExplicitlySetFields): 

867 """Scale the tensor s.t. its mean and variance match a reference tensor.""" 

868 

869 implemented_name: ClassVar[Literal["scale_mean_variance"]] = "scale_mean_variance" 

870 if TYPE_CHECKING: 

871 name: Literal["scale_mean_variance"] = "scale_mean_variance" 

872 else: 

873 name: Literal["scale_mean_variance"] 

874 

875 kwargs: ScaleMeanVarianceKwargs 

876 

877 

878PreprocessingDescr = Annotated[ 

879 Union[ 

880 BinarizeDescr, 

881 ClipDescr, 

882 ScaleLinearDescr, 

883 SigmoidDescr, 

884 ZeroMeanUnitVarianceDescr, 

885 ScaleRangeDescr, 

886 ], 

887 Discriminator("name"), 

888] 

889PostprocessingDescr = Annotated[ 

890 Union[ 

891 BinarizeDescr, 

892 ClipDescr, 

893 ScaleLinearDescr, 

894 SigmoidDescr, 

895 ZeroMeanUnitVarianceDescr, 

896 ScaleRangeDescr, 

897 ScaleMeanVarianceDescr, 

898 ], 

899 Discriminator("name"), 

900] 

901 

902 

903class InputTensorDescr(TensorDescrBase): 

904 data_type: Literal["float32", "uint8", "uint16"] 

905 """For now an input tensor is expected to be given as `float32`. 

906 The data flow in bioimage.io models is explained 

907 [in this diagram.](https://docs.google.com/drawings/d/1FTw8-Rn6a6nXdkZ_SkMumtcjvur9mtIhRqLwnKqZNHM/edit).""" 

908 

909 shape: Annotated[ 

910 Sequence[int] | ParameterizedInputShape, 

911 Field( 

912 examples=[(1, 512, 512, 1), {"min": (1, 64, 64, 1), "step": (0, 32, 32, 0)}] 

913 ), 

914 ] 

915 """Specification of input tensor shape.""" 

916 

917 preprocessing: list[PreprocessingDescr] = Field( 

918 default_factory=cast( # TODO: (py>3.8) use list[PreprocessingDesr] 

919 Callable[[], List[PreprocessingDescr]], list 

920 ) 

921 ) 

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

923 

924 @model_validator(mode="after") 

925 def zero_batch_step_and_one_batch_size(self) -> Self: 

926 bidx = self.axes.find("b") 

927 if bidx == -1: 

928 return self 

929 

930 if isinstance(self.shape, ParameterizedInputShape): 

931 step = self.shape.step 

932 shape = self.shape.min 

933 if step[bidx] != 0: 

934 raise ValueError( 

935 "Input shape step has to be zero in the batch dimension (the batch" 

936 + " dimension can always be increased, but `step` should specify how" 

937 + " to increase the minimal shape to find the largest single batch" 

938 + " shape)" 

939 ) 

940 else: 

941 shape = self.shape 

942 

943 if shape[bidx] != 1: 

944 raise ValueError("Input shape has to be 1 in the batch dimension b.") 

945 

946 return self 

947 

948 @model_validator(mode="after") 

949 def validate_preprocessing_kwargs(self) -> Self: 

950 for p in self.preprocessing: 

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

952 if isinstance(kwargs_axes, str) and any( 

953 a not in self.axes for a in kwargs_axes 

954 ): 

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

956 

957 return self 

958 

959 

960class OutputTensorDescr(TensorDescrBase): 

961 data_type: Literal[ 

962 "float32", 

963 "float64", 

964 "uint8", 

965 "int8", 

966 "uint16", 

967 "int16", 

968 "uint32", 

969 "int32", 

970 "uint64", 

971 "int64", 

972 "bool", 

973 ] 

974 """Data type. 

975 The data flow in bioimage.io models is explained 

976 [in this diagram.](https://docs.google.com/drawings/d/1FTw8-Rn6a6nXdkZ_SkMumtcjvur9mtIhRqLwnKqZNHM/edit).""" 

977 

978 @property 

979 def dtype(self): 

980 """alias for `data_type`""" 

981 return self.data_type 

982 

983 shape: Sequence[int] | ImplicitOutputShape 

984 """Output tensor shape.""" 

985 

986 halo: Sequence[int] | None = None 

987 """The `halo` that should be cropped from the output tensor to avoid boundary effects. 

988 The `halo` is to be cropped from both sides, i.e. `shape_after_crop = shape - 2 * halo`. 

989 To document a `halo` that is already cropped by the model `shape.offset` has to be used instead.""" 

990 

991 postprocessing: list[PostprocessingDescr] = Field( 

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

993 ) 

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

995 

996 @model_validator(mode="after") 

997 def matching_halo_length(self) -> Self: 

998 if self.halo and len(self.halo) != len(self.shape): 

999 raise ValueError( 

1000 f"halo {self.halo} has to have same length as shape {self.shape}!" 

1001 ) 

1002 

1003 return self 

1004 

1005 @model_validator(mode="after") 

1006 def validate_postprocessing_kwargs(self) -> Self: 

1007 for p in self.postprocessing: 

1008 kwargs_axes = p.kwargs.get("axes", "") 

1009 if not isinstance(kwargs_axes, str): 

1010 raise ValueError(f"Expected {kwargs_axes} to be a string") 

1011 

1012 if any(a not in self.axes for a in kwargs_axes): 

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

1014 

1015 return self 

1016 

1017 

1018KnownRunMode = Literal["deepimagej"] 

1019 

1020 

1021class RunMode(Node): 

1022 name: Annotated[ 

1023 KnownRunMode | str, warn(KnownRunMode, "Unknown run mode '{value}'.") 

1024 ] 

1025 """Run mode name""" 

1026 

1027 kwargs: dict[str, Any] = Field( 

1028 default_factory=cast(Callable[[], Dict[str, Any]], dict) 

1029 ) 

1030 """Run mode specific key word arguments""" 

1031 

1032 

1033class LinkedModel(Node): 

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

1035 

1036 id: Annotated[ModelId, Field(examples=["affable-shark", "ambitious-sloth"])] 

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

1038 

1039 version_number: int | None = None 

1040 """version number (n-th published version, not the semantic version) of linked model""" 

1041 

1042 

1043def package_weights( 

1044 value: Node, # Union[v0_4.WeightsDescr, v0_5.WeightsDescr] 

1045 handler: SerializerFunctionWrapHandler, 

1046 info: SerializationInfo, 

1047): 

1048 ctxt = packaging_context_var.get() 

1049 if ctxt is not None and ctxt.weights_priority_order is not None: 

1050 for wf in ctxt.weights_priority_order: 

1051 w = getattr(value, wf, None) 

1052 if w is not None: 

1053 break 

1054 else: 

1055 raise ValueError( 

1056 "None of the weight formats in `weights_priority_order`" 

1057 + f" ({ctxt.weights_priority_order}) is present in the given model." 

1058 ) 

1059 

1060 assert isinstance(w, Node), type(w) 

1061 # construct WeightsDescr with new single weight format entry 

1062 new_w = w.model_construct(**{k: v for k, v in w if k != "parent"}) 

1063 value = value.model_construct(None, **{wf: new_w}) 

1064 

1065 return handler( 

1066 value, 

1067 info, # pyright: ignore[reportArgumentType] # taken from pydantic docs 

1068 ) 

1069 

1070 

1071class ModelDescr(GenericModelDescrBase): 

1072 """Specification of the fields used in a bioimage.io-compliant RDF that describes AI models with pretrained weights. 

1073 

1074 These fields are typically stored in a YAML file which we call a model resource description file (model RDF). 

1075 """ 

1076 

1077 implemented_format_version: ClassVar[Literal["0.4.10"]] = "0.4.10" 

1078 if TYPE_CHECKING: 

1079 format_version: Literal["0.4.10"] = "0.4.10" 

1080 else: 

1081 format_version: Literal["0.4.10"] 

1082 """Version of the bioimage.io model description specification used. 

1083 When creating a new model always use the latest micro/patch version described here. 

1084 The `format_version` is important for any consumer software to understand how to parse the fields. 

1085 """ 

1086 

1087 implemented_type: ClassVar[Literal["model"]] = "model" 

1088 if TYPE_CHECKING: 

1089 type: Literal["model"] = "model" 

1090 else: 

1091 type: Literal["model"] 

1092 """Specialized resource type 'model'""" 

1093 

1094 id: ModelId | None = None 

1095 """bioimage.io-wide unique resource identifier 

1096 assigned by bioimage.io; version **un**specific.""" 

1097 

1098 authors: NotEmpty[ # pyright: ignore[reportGeneralTypeIssues] # make mandatory 

1099 list[Author] 

1100 ] 

1101 """The authors are the creators of the model RDF and the primary points of contact.""" 

1102 

1103 documentation: Annotated[ 

1104 FileSource_package, 

1105 Field( 

1106 examples=[ 

1107 "https://raw.githubusercontent.com/bioimage-io/spec-bioimage-io/main/example_descriptions/models/unet2d_nuclei_broad/README.md", 

1108 "README.md", 

1109 ], 

1110 ), 

1111 ] 

1112 """URL or relative path to a markdown file with additional documentation. 

1113 The recommended documentation file name is `README.md`. An `.md` suffix is mandatory. 

1114 The documentation should include a '[#[#]]# Validation' (sub)section 

1115 with details on how to quantitatively validate the model on unseen data.""" 

1116 

1117 inputs: NotEmpty[list[InputTensorDescr]] 

1118 """Describes the input tensors expected by this model.""" 

1119 

1120 license: Annotated[ 

1121 LicenseId | str, 

1122 warn(LicenseId, "Unknown license id '{value}'."), 

1123 Field(examples=["CC0-1.0", "MIT", "BSD-2-Clause"]), 

1124 ] 

1125 """A [SPDX license identifier](https://spdx.org/licenses/). 

1126 We do notsupport custom license beyond the SPDX license list, if you need that please 

1127 [open a GitHub issue](https://github.com/bioimage-io/spec-bioimage-io/issues/new/choose 

1128 ) to discuss your intentions with the community.""" 

1129 

1130 name: Annotated[ 

1131 str, 

1132 MinLen(1), 

1133 warn(MinLen(5), "Name shorter than 5 characters.", INFO), 

1134 warn(MaxLen(64), "Name longer than 64 characters.", INFO), 

1135 ] 

1136 """A human-readable name of this model. 

1137 It should be no longer than 64 characters and only contain letter, number, underscore, minus or space characters.""" 

1138 

1139 outputs: NotEmpty[list[OutputTensorDescr]] 

1140 """Describes the output tensors.""" 

1141 

1142 @field_validator("inputs", "outputs") 

1143 @classmethod 

1144 def unique_tensor_descr_names( 

1145 cls, value: Sequence[InputTensorDescr | OutputTensorDescr] 

1146 ) -> Sequence[InputTensorDescr | OutputTensorDescr]: 

1147 unique_names = {str(v.name) for v in value} 

1148 if len(unique_names) != len(value): 

1149 raise ValueError("Duplicate tensor descriptor names") 

1150 

1151 return value 

1152 

1153 @model_validator(mode="after") 

1154 def unique_io_names(self) -> Self: 

1155 unique_names = {str(ss.name) for s in (self.inputs, self.outputs) for ss in s} 

1156 if len(unique_names) != (len(self.inputs) + len(self.outputs)): 

1157 raise ValueError("Duplicate tensor descriptor names across inputs/outputs") 

1158 

1159 return self 

1160 

1161 @model_validator(mode="after") 

1162 def minimum_shape2valid_output(self) -> Self: 

1163 tensors_by_name: dict[TensorName, InputTensorDescr | OutputTensorDescr] = { 

1164 t.name: t for t in self.inputs + self.outputs 

1165 } 

1166 

1167 for out in self.outputs: 

1168 if isinstance(out.shape, ImplicitOutputShape): 

1169 ndim_ref = len(tensors_by_name[out.shape.reference_tensor].shape) 

1170 ndim_out_ref = len( 

1171 [scale for scale in out.shape.scale if scale is not None] 

1172 ) 

1173 if ndim_ref != ndim_out_ref: 

1174 expanded_dim_note = ( 

1175 " Note that expanded dimensions (`scale`: null) are not" 

1176 + f" counted for {out.name}'sdimensionality here." 

1177 if None in out.shape.scale 

1178 else "" 

1179 ) 

1180 raise ValueError( 

1181 f"Referenced tensor '{out.shape.reference_tensor}' with" 

1182 + f" {ndim_ref} dimensions does not match output tensor" 

1183 + f" '{out.name}' with" 

1184 + f" {ndim_out_ref} dimensions.{expanded_dim_note}" 

1185 ) 

1186 

1187 min_out_shape = self._get_min_shape(out, tensors_by_name) 

1188 if out.halo: 

1189 halo = out.halo 

1190 halo_msg = f" for halo {out.halo}" 

1191 else: 

1192 halo = [0] * len(min_out_shape) 

1193 halo_msg = "" 

1194 

1195 if any(s - 2 * h < 1 for s, h in zip(min_out_shape, halo)): 

1196 raise ValueError( 

1197 f"Minimal shape {min_out_shape} of output {out.name} is too" 

1198 + f" small{halo_msg}." 

1199 ) 

1200 

1201 return self 

1202 

1203 @classmethod 

1204 def _get_min_shape( 

1205 cls, 

1206 t: InputTensorDescr | OutputTensorDescr, 

1207 tensors_by_name: dict[TensorName, InputTensorDescr | OutputTensorDescr], 

1208 ) -> Sequence[int]: 

1209 """output with subtracted halo has to result in meaningful output even for the minimal input 

1210 see https://github.com/bioimage-io/spec-bioimage-io/issues/392 

1211 """ 

1212 if isinstance(t.shape, collections.abc.Sequence): 

1213 return t.shape 

1214 elif isinstance(t.shape, ParameterizedInputShape): 

1215 return t.shape.min 

1216 elif isinstance(t.shape, ImplicitOutputShape): 

1217 pass 

1218 else: 

1219 assert_never(t.shape) 

1220 

1221 ref_shape = cls._get_min_shape( 

1222 tensors_by_name[t.shape.reference_tensor], tensors_by_name 

1223 ) 

1224 

1225 if None not in t.shape.scale: 

1226 scale: Sequence[float, ...] = t.shape.scale # type: ignore 

1227 else: 

1228 expanded_dims = [idx for idx, sc in enumerate(t.shape.scale) if sc is None] 

1229 new_ref_shape: list[int] = [] 

1230 for idx in range(len(t.shape.scale)): 

1231 ref_idx = idx - sum(int(exp < idx) for exp in expanded_dims) 

1232 new_ref_shape.append(1 if idx in expanded_dims else ref_shape[ref_idx]) 

1233 

1234 ref_shape = new_ref_shape 

1235 assert len(ref_shape) == len(t.shape.scale) 

1236 scale = [0.0 if sc is None else sc for sc in t.shape.scale] 

1237 

1238 offset = t.shape.offset 

1239 assert len(offset) == len(scale) 

1240 return [int(rs * s + 2 * off) for rs, s, off in zip(ref_shape, scale, offset)] 

1241 

1242 @model_validator(mode="after") 

1243 def validate_tensor_references_in_inputs(self) -> Self: 

1244 for t in self.inputs: 

1245 for proc in t.preprocessing: 

1246 if "reference_tensor" not in proc.kwargs: 

1247 continue 

1248 

1249 ref_tensor = proc.kwargs["reference_tensor"] 

1250 if ref_tensor is not None and str(ref_tensor) not in { 

1251 str(t.name) for t in self.inputs 

1252 }: 

1253 raise ValueError(f"'{ref_tensor}' not found in inputs") 

1254 

1255 if ref_tensor == t.name: 

1256 raise ValueError( 

1257 f"invalid self reference for preprocessing of tensor {t.name}" 

1258 ) 

1259 

1260 return self 

1261 

1262 @model_validator(mode="after") 

1263 def validate_tensor_references_in_outputs(self) -> Self: 

1264 for t in self.outputs: 

1265 for proc in t.postprocessing: 

1266 if "reference_tensor" not in proc.kwargs: 

1267 continue 

1268 ref_tensor = proc.kwargs["reference_tensor"] 

1269 if ref_tensor is not None and str(ref_tensor) not in { 

1270 str(t.name) for t in self.inputs 

1271 }: 

1272 raise ValueError(f"{ref_tensor} not found in inputs") 

1273 

1274 return self 

1275 

1276 packaged_by: list[Author] = Field( 

1277 default_factory=cast(Callable[[], List[Author]], list) 

1278 ) 

1279 """The persons that have packaged and uploaded this model. 

1280 Only required if those persons differ from the `authors`.""" 

1281 

1282 parent: LinkedModel | None = None 

1283 """The model from which this model is derived, e.g. by fine-tuning the weights.""" 

1284 

1285 @field_validator("parent", mode="before") 

1286 @classmethod 

1287 def ignore_url_parent(cls, parent: Any): 

1288 if isinstance(parent, dict): 

1289 return None 

1290 

1291 else: 

1292 return parent 

1293 

1294 run_mode: RunMode | None = None 

1295 """Custom run mode for this model: for more complex prediction procedures like test time 

1296 data augmentation that currently cannot be expressed in the specification. 

1297 No standard run modes are defined yet.""" 

1298 

1299 sample_inputs: list[FileSource_package] = Field( 

1300 default_factory=cast(Callable[[], List[FileSource]], list) 

1301 ) 

1302 """URLs/relative paths to sample inputs to illustrate possible inputs for the model, 

1303 for example stored as PNG or TIFF images. 

1304 The sample files primarily serve to inform a human user about an example use case""" 

1305 

1306 sample_outputs: list[FileSource_package] = Field( 

1307 default_factory=cast(Callable[[], List[FileSource]], list) 

1308 ) 

1309 """URLs/relative paths to sample outputs corresponding to the `sample_inputs`.""" 

1310 

1311 test_inputs: NotEmpty[ 

1312 list[Annotated[FileSource_package, WithSuffix(".npy", case_sensitive=True)]] 

1313 ] 

1314 """Test input tensors compatible with the `inputs` description for a **single test case**. 

1315 This means if your model has more than one input, you should provide one URL/relative path for each input. 

1316 Each test input should be a file with an ndarray in 

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

1318 The extension must be '.npy'.""" 

1319 

1320 test_outputs: NotEmpty[ 

1321 list[Annotated[FileSource_package, WithSuffix(".npy", case_sensitive=True)]] 

1322 ] 

1323 """Analog to `test_inputs`.""" 

1324 

1325 timestamp: Datetime 

1326 """Timestamp in [ISO 8601](#https://en.wikipedia.org/wiki/ISO_8601) format 

1327 with a few restrictions listed [here](https://docs.python.org/3/library/datetime.html#datetime.datetime.fromisoformat).""" 

1328 

1329 training_data: LinkedDataset | DatasetDescr | None = None 

1330 """The dataset used to train this model""" 

1331 

1332 weights: Annotated[WeightsDescr, WrapSerializer(package_weights)] 

1333 """The weights for this model. 

1334 Weights can be given for different formats, but should otherwise be equivalent. 

1335 The available weight formats determine which consumers can use this model.""" 

1336 

1337 @model_validator(mode="before") 

1338 @classmethod 

1339 def _convert_from_older_format( 

1340 cls, data: BioimageioYamlContent, / 

1341 ) -> BioimageioYamlContent: 

1342 convert_from_older_format(data) 

1343 return data 

1344 

1345 def get_input_test_arrays(self) -> list[NDArray[Any]]: 

1346 data = [load_array(ipt) for ipt in self.test_inputs] 

1347 assert all(isinstance(d, np.ndarray) for d in data) 

1348 return data 

1349 

1350 def get_output_test_arrays(self) -> list[NDArray[Any]]: 

1351 data = [load_array(out) for out in self.test_outputs] 

1352 assert all(isinstance(d, np.ndarray) for d in data) 

1353 return data