mirror of
https://github.com/macaodha/batdetect2.git
synced 2026-08-22 03:00:10 +02:00
fix: address bottleneck config review feedback
This commit is contained in:
parent
cf49c6e3da
commit
ac805be169
@ -150,8 +150,8 @@ class SelfAttentionConfig(BaseConfig):
|
|||||||
attention_channels : int
|
attention_channels : int
|
||||||
Dimensionality of the query, key, and value projections.
|
Dimensionality of the query, key, and value projections.
|
||||||
temperature : float
|
temperature : float
|
||||||
Scaling factor applied to the weighted values before the final
|
Divisor applied together with ``attention_channels`` when scaling
|
||||||
linear projection. Defaults to ``1``.
|
dot-product attention logits. Defaults to ``1``.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
name: Literal["SelfAttention"] = "SelfAttention"
|
name: Literal["SelfAttention"] = "SelfAttention"
|
||||||
@ -321,8 +321,8 @@ class EfficientSelfAttentionConfig(BaseConfig):
|
|||||||
attention_channels : int
|
attention_channels : int
|
||||||
Dimensionality of the query, key, and value projections.
|
Dimensionality of the query, key, and value projections.
|
||||||
temperature : float
|
temperature : float
|
||||||
Scaling factor applied to the weighted values before the final
|
Divisor applied together with ``attention_channels`` when scaling
|
||||||
linear projection. Defaults to ``1``.
|
dot-product attention logits. Defaults to ``1``.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
name: Literal["EfficientSelfAttention"] = "EfficientSelfAttention"
|
name: Literal["EfficientSelfAttention"] = "EfficientSelfAttention"
|
||||||
|
|||||||
@ -21,7 +21,7 @@ This module provides:
|
|||||||
from typing import Annotated, List
|
from typing import Annotated, List
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from pydantic import Field
|
from pydantic import Field, model_validator
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
from batdetect2.core.configs import BaseConfig
|
from batdetect2.core.configs import BaseConfig
|
||||||
@ -197,13 +197,19 @@ class BottleneckConfig(BaseConfig):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
channels: int
|
channels: int
|
||||||
frequency_aggregation: FrequencyAggregationLayerConfig = Field(
|
frequency_aggregation: FrequencyAggregationLayerConfig | None = None
|
||||||
default_factory=lambda data: VerticalConvConfig(
|
|
||||||
channels=data["channels"]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
layers: List[BottleneckLayerConfig] = Field(default_factory=list)
|
layers: List[BottleneckLayerConfig] = Field(default_factory=list)
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def set_default_frequency_aggregation(self) -> "BottleneckConfig":
|
||||||
|
"""Default frequency aggregation to the bottleneck channel count."""
|
||||||
|
if self.frequency_aggregation is None:
|
||||||
|
self.frequency_aggregation = VerticalConvConfig(
|
||||||
|
channels=self.channels
|
||||||
|
)
|
||||||
|
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
DEFAULT_BOTTLENECK_CONFIG: BottleneckConfig = BottleneckConfig(
|
DEFAULT_BOTTLENECK_CONFIG: BottleneckConfig = BottleneckConfig(
|
||||||
channels=256,
|
channels=256,
|
||||||
@ -302,11 +308,14 @@ def build_bottleneck(
|
|||||||
by repetition).
|
by repetition).
|
||||||
"""
|
"""
|
||||||
config = config or DEFAULT_BOTTLENECK_CONFIG
|
config = config or DEFAULT_BOTTLENECK_CONFIG
|
||||||
|
frequency_aggregation = config.frequency_aggregation
|
||||||
|
if frequency_aggregation is None:
|
||||||
|
raise ValueError("frequency_aggregation must be configured.")
|
||||||
|
|
||||||
frequency_aggregator = build_frequency_aggregation(
|
frequency_aggregator = build_frequency_aggregation(
|
||||||
input_height=input_height,
|
input_height=input_height,
|
||||||
in_channels=in_channels,
|
in_channels=in_channels,
|
||||||
config=config.frequency_aggregation,
|
config=frequency_aggregation,
|
||||||
)
|
)
|
||||||
|
|
||||||
current_channels = frequency_aggregator.out_channels
|
current_channels = frequency_aggregator.out_channels
|
||||||
|
|||||||
@ -82,4 +82,8 @@ def test_bundled_checkpoint_loads_with_current_config_schema() -> None:
|
|||||||
|
|
||||||
assert model.class_names
|
assert model.class_names
|
||||||
assert isinstance(configs.model.architecture, UNetBackboneConfig)
|
assert isinstance(configs.model.architecture, UNetBackboneConfig)
|
||||||
assert configs.model.architecture.bottleneck.frequency_aggregation.name
|
frequency_aggregation = (
|
||||||
|
configs.model.architecture.bottleneck.frequency_aggregation
|
||||||
|
)
|
||||||
|
assert frequency_aggregation is not None
|
||||||
|
assert frequency_aggregation.name
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user