Skip to content

Commit 146d238

Browse files
committed
Refactor AbstractBlock and Stage classes to improve output channel handling and streamline output dimension calculations
1 parent 6ba51f4 commit 146d238

2 files changed

Lines changed: 6 additions & 7 deletions

File tree

‎src/virtual_stain_flow/models/blocks.py‎

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -49,7 +49,7 @@ class AbstractBlock(ABC, nn.Module):
4949
def __init__(
5050
self,
5151
in_channels: int,
52-
out_channels: int,
52+
out_channels: Optional[int] = None,
5353
num_units: int = 1,
5454
**kwargs: dict
5555
):
@@ -63,6 +63,8 @@ def __init__(
6363
if in_channels <= 0:
6464
raise ValueError("Expected in_channels to be positive, "
6565
f"got {in_channels}")
66+
if out_channels is None:
67+
out_channels = in_channels
6668
if not isinstance(out_channels, int):
6769
raise TypeError("Expected out_channels to be int, "
6870
f"got {type(out_channels).__name__}")
@@ -101,10 +103,9 @@ def num_units(self) -> int:
101103
# These 2 below should be overriden to reflect the actual spatial dimension
102104
# changes the block applies. By default they indicate spatial preserving
103105
# blocks, i.e. the height and width of the input tensor remain unchanged.
104-
@property
105106
def out_h(self, in_h: int) -> int:
106107
return in_h
107-
@property
108+
108109
def out_w(self, in_w: int) -> int:
109110
return in_w
110111

‎src/virtual_stain_flow/models/stages.py‎

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -173,18 +173,16 @@ def skip_channels(self) -> int:
173173
def out_channels(self) -> int:
174174
return self._out_channels
175175

176-
@property
177176
def out_h(self, in_h: int) -> int:
178177
_out_h = in_h
179-
for block in self.blocks:
178+
for block in [self.in_block, self.comp_block]:
180179
if isinstance(block, Conv2DDownBlock):
181180
_out_h = block.out_h(_out_h)
182181
return _out_h
183182

184-
@property
185183
def out_w(self, in_w: int) -> int:
186184
_out_w = in_w
187-
for block in self.blocks:
185+
for block in [self.in_block, self.comp_block]:
188186
if isinstance(block, Conv2DDownBlock):
189187
_out_w = block.out_w(_out_w)
190188
return _out_w

0 commit comments

Comments
 (0)