mirror of
https://github.com/macaodha/batdetect2.git
synced 2026-08-22 03:00:10 +02:00
Add vertical mean block
This commit is contained in:
parent
0431ef2679
commit
575e0b3d28
@ -474,6 +474,91 @@ class VerticalConv(Block):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class VerticalMeanConfig(BaseConfig):
|
||||||
|
"""Configuration for a ``VerticalMean`` block.
|
||||||
|
|
||||||
|
Attributes
|
||||||
|
----------
|
||||||
|
name : str
|
||||||
|
Discriminator field; always ``"VerticalMean"``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
name: Literal["VerticalMean"] = "VerticalMean"
|
||||||
|
"""Discriminator field indicating the block type."""
|
||||||
|
|
||||||
|
channels: int
|
||||||
|
"""Number of output channels."""
|
||||||
|
|
||||||
|
|
||||||
|
class VerticalMean(Block):
|
||||||
|
"""Mean pooling block operating along the height dimension.
|
||||||
|
|
||||||
|
Applies a 2D mean pooling operation, followed by a 2D convolution,
|
||||||
|
followed by a batch normalization and ReLU activation.
|
||||||
|
|
||||||
|
Sequence: Mean Pool -> Conv -> BN -> ReLU.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
in_channels : int
|
||||||
|
Number of channels in the input tensor.
|
||||||
|
out_channels : int
|
||||||
|
Number of output channels after the mean pooling.
|
||||||
|
input_height : int
|
||||||
|
The height (H dimension) of the input tensor. The convolutional kernel
|
||||||
|
will be sized `(1, 1)`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int,
|
||||||
|
out_channels: int,
|
||||||
|
input_height: int,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.in_channels = in_channels
|
||||||
|
self.out_channels = out_channels
|
||||||
|
self.input_height = input_height
|
||||||
|
self.conv = nn.Conv2d(
|
||||||
|
in_channels,
|
||||||
|
out_channels,
|
||||||
|
kernel_size=(1, 1),
|
||||||
|
padding=0,
|
||||||
|
)
|
||||||
|
self.batch_norm = nn.BatchNorm2d(out_channels)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Apply avg pooling -> Conv -> BN -> ReLU.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
x : torch.Tensor
|
||||||
|
Input tensor, shape `(B, C_in, H, W)`.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
torch.Tensor
|
||||||
|
Output tensor, shape `(B, C_out, 1, W)`.
|
||||||
|
"""
|
||||||
|
x = x.mean(dim=2, keepdim=True)
|
||||||
|
x = self.conv(x)
|
||||||
|
return F.relu(self.batch_norm(x), inplace=True)
|
||||||
|
|
||||||
|
@block_registry.register(VerticalMeanConfig)
|
||||||
|
@staticmethod
|
||||||
|
def from_config(
|
||||||
|
config: VerticalMeanConfig,
|
||||||
|
input_channels: int,
|
||||||
|
input_height: int,
|
||||||
|
):
|
||||||
|
return VerticalMean(
|
||||||
|
in_channels=input_channels,
|
||||||
|
out_channels=config.channels,
|
||||||
|
input_height=input_height,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class FreqCoordConvDownConfig(BaseConfig):
|
class FreqCoordConvDownConfig(BaseConfig):
|
||||||
"""Configuration for a FreqCoordConvDownBlock."""
|
"""Configuration for a FreqCoordConvDownBlock."""
|
||||||
|
|
||||||
|
|||||||
@ -29,6 +29,9 @@ from batdetect2.models.blocks import (
|
|||||||
Block,
|
Block,
|
||||||
SelfAttentionConfig,
|
SelfAttentionConfig,
|
||||||
VerticalConv,
|
VerticalConv,
|
||||||
|
VerticalConvConfig,
|
||||||
|
VerticalMean,
|
||||||
|
VerticalMeanConfig,
|
||||||
build_layer,
|
build_layer,
|
||||||
)
|
)
|
||||||
from batdetect2.models.types import BottleneckProtocol
|
from batdetect2.models.types import BottleneckProtocol
|
||||||
@ -94,6 +97,7 @@ class Bottleneck(Block):
|
|||||||
in_channels: int,
|
in_channels: int,
|
||||||
out_channels: int,
|
out_channels: int,
|
||||||
bottleneck_channels: int | None = None,
|
bottleneck_channels: int | None = None,
|
||||||
|
frequency_aggregator: Block | None = None,
|
||||||
layers: List[torch.nn.Module] | None = None,
|
layers: List[torch.nn.Module] | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Initialise the Bottleneck layer.
|
"""Initialise the Bottleneck layer.
|
||||||
@ -125,12 +129,15 @@ class Bottleneck(Block):
|
|||||||
)
|
)
|
||||||
self.layers = nn.ModuleList(layers or [])
|
self.layers = nn.ModuleList(layers or [])
|
||||||
|
|
||||||
self.conv_vert = VerticalConv(
|
if frequency_aggregator is None:
|
||||||
|
frequency_aggregator = VerticalConv(
|
||||||
in_channels=in_channels,
|
in_channels=in_channels,
|
||||||
out_channels=self.bottleneck_channels,
|
out_channels=self.bottleneck_channels,
|
||||||
input_height=input_height,
|
input_height=input_height,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self.conv_vert = frequency_aggregator
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
"""Process the encoder's bottleneck features.
|
"""Process the encoder's bottleneck features.
|
||||||
|
|
||||||
@ -166,6 +173,13 @@ BottleneckLayerConfig = Annotated[
|
|||||||
"""Type alias for the discriminated union of block configs usable in the Bottleneck."""
|
"""Type alias for the discriminated union of block configs usable in the Bottleneck."""
|
||||||
|
|
||||||
|
|
||||||
|
FrequencyAggregationLayerConfig = Annotated[
|
||||||
|
(VerticalConvConfig | VerticalMeanConfig),
|
||||||
|
Field(discriminator="name"),
|
||||||
|
]
|
||||||
|
"""Type alias for the discriminated union of block configs usable in the FrequencyAggregation."""
|
||||||
|
|
||||||
|
|
||||||
class BottleneckConfig(BaseConfig):
|
class BottleneckConfig(BaseConfig):
|
||||||
"""Configuration for the bottleneck component.
|
"""Configuration for the bottleneck component.
|
||||||
|
|
||||||
@ -182,17 +196,71 @@ class BottleneckConfig(BaseConfig):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
channels: int
|
channels: int
|
||||||
|
frequency_aggregation: FrequencyAggregationLayerConfig = Field(
|
||||||
|
default_factory=lambda: VerticalConvConfig(channels=256)
|
||||||
|
)
|
||||||
layers: List[BottleneckLayerConfig] = Field(default_factory=list)
|
layers: List[BottleneckLayerConfig] = Field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
DEFAULT_BOTTLENECK_CONFIG: BottleneckConfig = BottleneckConfig(
|
DEFAULT_BOTTLENECK_CONFIG: BottleneckConfig = BottleneckConfig(
|
||||||
channels=256,
|
channels=256,
|
||||||
|
frequency_aggregation=VerticalConvConfig(channels=256),
|
||||||
layers=[
|
layers=[
|
||||||
SelfAttentionConfig(attention_channels=256),
|
SelfAttentionConfig(attention_channels=256),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def build_frequency_aggregation(
|
||||||
|
input_height: int,
|
||||||
|
in_channels: int,
|
||||||
|
config: FrequencyAggregationLayerConfig,
|
||||||
|
) -> Block:
|
||||||
|
"""Build a block for aggregating frequency information.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
input_height : int
|
||||||
|
Height (number of frequency bins) of the input tensor from the
|
||||||
|
encoder. Must be positive.
|
||||||
|
in_channels : int
|
||||||
|
Number of channels in the input tensor from the encoder. Must be
|
||||||
|
positive.
|
||||||
|
config : FrequencyAggregationLayerConfig, optional
|
||||||
|
Configuration specifying the output channel count and any
|
||||||
|
additional layers. Uses ``VerticalConvConfig`` if ``None``.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
Block
|
||||||
|
An initialised ``VerticalConv`` module.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
AssertionError
|
||||||
|
If any configured layer changes the height of the feature map
|
||||||
|
(bottleneck layers must preserve height so that it can be restored
|
||||||
|
by repetition).
|
||||||
|
"""
|
||||||
|
if config.name == "VerticalConv":
|
||||||
|
return VerticalConv(
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=config.channels,
|
||||||
|
input_height=input_height,
|
||||||
|
)
|
||||||
|
|
||||||
|
if config.name == "VerticalMean":
|
||||||
|
return VerticalMean(
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=config.channels,
|
||||||
|
input_height=input_height,
|
||||||
|
)
|
||||||
|
|
||||||
|
raise NotImplementedError(
|
||||||
|
f"Unknown frequency aggregation layer: {config.name}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_bottleneck(
|
def build_bottleneck(
|
||||||
input_height: int,
|
input_height: int,
|
||||||
in_channels: int,
|
in_channels: int,
|
||||||
@ -250,9 +318,16 @@ def build_bottleneck(
|
|||||||
)
|
)
|
||||||
layers.append(layer)
|
layers.append(layer)
|
||||||
|
|
||||||
|
frequency_aggregator = build_frequency_aggregation(
|
||||||
|
input_height=input_height,
|
||||||
|
in_channels=current_channels,
|
||||||
|
config=config.frequency_aggregation,
|
||||||
|
)
|
||||||
|
|
||||||
return Bottleneck(
|
return Bottleneck(
|
||||||
input_height=input_height,
|
input_height=input_height,
|
||||||
in_channels=in_channels,
|
in_channels=in_channels,
|
||||||
out_channels=config.channels,
|
out_channels=config.channels,
|
||||||
|
frequency_aggregator=frequency_aggregator,
|
||||||
layers=layers,
|
layers=layers,
|
||||||
)
|
)
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user