Coverage for src/bioimageio/core/_op_base.py: 97%
36 statements
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-08 15:59 +0000
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-08 15:59 +0000
1from abc import ABC, abstractmethod
2from dataclasses import dataclass
3from typing import Collection, Generic, Union
5from typing_extensions import TypeVar, assert_never
7from ._restore_batch_multi_index import restore_batch_multi_index
8from .axis import PerAxis
9from .block import Block
10from .common import MemberId
11from .sample import Sample, SampleBlock, SampleBlockWithOrigin
12from .stat_measures import (
13 Measure,
14 Stat,
15)
16from .tensor import Tensor
18SampleT = TypeVar("SampleT", bound=Union[Sample, SampleBlock, SampleBlockWithOrigin])
21@dataclass
22class Operator(Generic[SampleT], ABC):
23 """Base class for all operators."""
25 @abstractmethod
26 def __call__(self, sample: SampleT) -> None: ...
28 @property
29 @abstractmethod
30 def required_measures(self) -> Collection[Measure]: ...
33@dataclass
34class SamplewiseOperator(Operator[Sample]):
35 """Base class for operators that can only be applied to whole samples."""
38@dataclass
39class BlockwiseOperator(Operator[Union[Sample, SampleBlock]]):
40 """Base class for operators that can be applied to whole sample or blockwise."""
43@dataclass
44class SimpleOperator(BlockwiseOperator):
45 """Convenience base class for blockwise operators with a single input and single output."""
47 input: MemberId
48 output: MemberId
50 @abstractmethod
51 def get_output_shape(self, input_shape: PerAxis[int]) -> PerAxis[int]: ...
53 def __call__(self, sample: Union[Sample, SampleBlock]) -> None:
54 if self.input not in sample.members:
55 return # TODO: raise?
57 input_tensor = sample.members[self.input]
58 output_tensor = self._apply(input_tensor, sample.stat)
59 output_tensor = restore_batch_multi_index(
60 {self.input: input_tensor}, {self.output: output_tensor}
61 )[self.output]
62 assert output_tensor is not None
63 if self.output in sample.members:
64 assert (
65 sample.members[self.output].tagged_shape == output_tensor.tagged_shape
66 )
68 if isinstance(sample, Sample):
69 sample.members[self.output] = output_tensor
70 elif isinstance(sample, SampleBlock):
71 b = sample.blocks[self.input]
72 sample.blocks[self.output] = Block(
73 sample_shape=self.get_output_shape(sample.sample_shape[self.input]),
74 data=output_tensor,
75 inner_slice=b.inner_slice,
76 halo=b.halo,
77 block_index=b.block_index,
78 blocks_in_sample=b.blocks_in_sample,
79 )
80 else:
81 assert_never(sample)
83 @abstractmethod
84 def _apply(self, x: Tensor, stat: Stat) -> Tensor: ...