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

1from abc import ABC, abstractmethod 

2from dataclasses import dataclass 

3from typing import Collection, Generic, Union 

4 

5from typing_extensions import TypeVar, assert_never 

6 

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 

17 

18SampleT = TypeVar("SampleT", bound=Union[Sample, SampleBlock, SampleBlockWithOrigin]) 

19 

20 

21@dataclass 

22class Operator(Generic[SampleT], ABC): 

23 """Base class for all operators.""" 

24 

25 @abstractmethod 

26 def __call__(self, sample: SampleT) -> None: ... 

27 

28 @property 

29 @abstractmethod 

30 def required_measures(self) -> Collection[Measure]: ... 

31 

32 

33@dataclass 

34class SamplewiseOperator(Operator[Sample]): 

35 """Base class for operators that can only be applied to whole samples.""" 

36 

37 

38@dataclass 

39class BlockwiseOperator(Operator[Union[Sample, SampleBlock]]): 

40 """Base class for operators that can be applied to whole sample or blockwise.""" 

41 

42 

43@dataclass 

44class SimpleOperator(BlockwiseOperator): 

45 """Convenience base class for blockwise operators with a single input and single output.""" 

46 

47 input: MemberId 

48 output: MemberId 

49 

50 @abstractmethod 

51 def get_output_shape(self, input_shape: PerAxis[int]) -> PerAxis[int]: ... 

52 

53 def __call__(self, sample: Union[Sample, SampleBlock]) -> None: 

54 if self.input not in sample.members: 

55 return # TODO: raise? 

56 

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 ) 

67 

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) 

82 

83 @abstractmethod 

84 def _apply(self, x: Tensor, stat: Stat) -> Tensor: ...