mirror of
https://github.com/macaodha/batdetect2.git
synced 2026-08-22 03:00:10 +02:00
Compare commits
No commits in common. "9e2697458d3b3b03c30ccd7e49ed5409c8ac330d" and "d7896c6d75b8cbea85d5826ca465ed6703136aa6" have entirely different histories.
9e2697458d
...
d7896c6d75
10
README.md
10
README.md
@ -6,7 +6,7 @@ Code for detecting and classifying bat echolocation calls in high-frequency
|
|||||||
audio recordings.
|
audio recordings.
|
||||||
|
|
||||||
> [!WARNING]
|
> [!WARNING]
|
||||||
> `batdetect2` 2.0.0b3 is out.
|
> `batdetect2` 2.0.0b2 is out.
|
||||||
> This is a beta release and we are gathering user feedback.
|
> This is a beta release and we are gathering user feedback.
|
||||||
> If you run into issues or have feedback on the new workflows, please use the
|
> If you run into issues or have feedback on the new workflows, please use the
|
||||||
> GitHub issues page to let us know.
|
> GitHub issues page to let us know.
|
||||||
@ -68,13 +68,13 @@ can try the following:
|
|||||||
### Installing BatDetect2
|
### Installing BatDetect2
|
||||||
|
|
||||||
> [!NOTE]
|
> [!NOTE]
|
||||||
> `2.0.0b3` is a pre-release on PyPI.
|
> `2.0.0b2` is a pre-release on PyPI.
|
||||||
> You may need to request it explicitly by version, for example:
|
> You may need to request it explicitly by version, for example:
|
||||||
>
|
>
|
||||||
> ```bash
|
> ```bash
|
||||||
> uvx --from batdetect2==2.0.0b3 batdetect2
|
> uvx --from batdetect2==2.0.0b2 batdetect2
|
||||||
> uv tool install batdetect2==2.0.0b3
|
> uv tool install batdetect2==2.0.0b2
|
||||||
> pip install batdetect2==2.0.0b3
|
> pip install batdetect2==2.0.0b2
|
||||||
> ```
|
> ```
|
||||||
|
|
||||||
If you have `uv` installed (if not, we recommend it; follow the instructions
|
If you have `uv` installed (if not, we recommend it; follow the instructions
|
||||||
|
|||||||
@ -15,9 +15,6 @@ Defined in `batdetect2.api_v2`.
|
|||||||
- `BatDetect2API.from_config(model_config=..., targets_config=..., ...)`
|
- `BatDetect2API.from_config(model_config=..., targets_config=..., ...)`
|
||||||
- build a full model stack from config objects.
|
- build a full model stack from config objects.
|
||||||
|
|
||||||
Both constructors accept `compile_model=True` to compile the detector after the
|
|
||||||
API is built.
|
|
||||||
|
|
||||||
## Common tasks
|
## Common tasks
|
||||||
|
|
||||||
- Load a checkpoint and run prediction on one file.
|
- Load a checkpoint and run prediction on one file.
|
||||||
@ -25,8 +22,6 @@ API is built.
|
|||||||
- Save predictions in one of the supported output formats.
|
- Save predictions in one of the supported output formats.
|
||||||
- Evaluate a model on labelled data.
|
- Evaluate a model on labelled data.
|
||||||
- Fine-tune an existing checkpoint on new targets.
|
- Fine-tune an existing checkpoint on new targets.
|
||||||
- Compile the detector explicitly with `BatDetect2API.compile()` when you want
|
|
||||||
to opt into PyTorch runtime compilation from Python.
|
|
||||||
|
|
||||||
## Generated reference
|
## Generated reference
|
||||||
|
|
||||||
|
|||||||
@ -7,8 +7,6 @@ Defined in `batdetect2.inference.config`.
|
|||||||
|
|
||||||
## Top-level fields
|
## Top-level fields
|
||||||
|
|
||||||
- `compile_model`
|
|
||||||
- compile the detector before batch prediction. This is off by default.
|
|
||||||
- `loader`
|
- `loader`
|
||||||
- data-loader settings for inference.
|
- data-loader settings for inference.
|
||||||
- `clipping`
|
- `clipping`
|
||||||
@ -36,19 +34,8 @@ Override `InferenceConfig` when:
|
|||||||
|
|
||||||
- long recordings need different clipping behavior,
|
- long recordings need different clipping behavior,
|
||||||
- you want to tune batch size for your hardware,
|
- you want to tune batch size for your hardware,
|
||||||
- you want to opt into runtime model compilation for repeated predictions,
|
|
||||||
- you need reproducible prediction settings across runs.
|
- you need reproducible prediction settings across runs.
|
||||||
|
|
||||||
## Runtime compilation
|
|
||||||
|
|
||||||
Set `compile_model: true` to compile the detector before batch inference. This
|
|
||||||
can help when you run repeated predictions with stable input shapes. For a
|
|
||||||
single short run, the compile step can cost more time than it saves.
|
|
||||||
|
|
||||||
In Python, you can also compile explicitly with `BatDetect2API.compile()` or by
|
|
||||||
passing `compile_model=True` to `BatDetect2API.from_checkpoint(...)` or
|
|
||||||
`BatDetect2API.from_config(...)`.
|
|
||||||
|
|
||||||
## Related pages
|
## Related pages
|
||||||
|
|
||||||
- Tune inference clipping:
|
- Tune inference clipping:
|
||||||
|
|||||||
@ -7,8 +7,6 @@ Defined in `batdetect2.train.config`.
|
|||||||
|
|
||||||
## Top-level fields
|
## Top-level fields
|
||||||
|
|
||||||
- `compile_model`
|
|
||||||
- compile the detector before training starts. This is off by default.
|
|
||||||
- `train_loader`
|
- `train_loader`
|
||||||
- training data loading and clipping settings.
|
- training data loading and clipping settings.
|
||||||
- `val_loader`
|
- `val_loader`
|
||||||
@ -35,18 +33,10 @@ Use `TrainingConfig` when you want to change things like:
|
|||||||
- batch size,
|
- batch size,
|
||||||
- augmentation,
|
- augmentation,
|
||||||
- optimiser and scheduler settings,
|
- optimiser and scheduler settings,
|
||||||
- runtime options such as model compilation,
|
|
||||||
- number of epochs,
|
- number of epochs,
|
||||||
- validation frequency,
|
- validation frequency,
|
||||||
- checkpoint behaviour.
|
- checkpoint behaviour.
|
||||||
|
|
||||||
## Runtime options
|
|
||||||
|
|
||||||
Use `compile_model: true` to call `torch.compile` on the detector used during
|
|
||||||
training. This can help on longer runs with stable tensor shapes, but it may be
|
|
||||||
slower for short CPU-only experiments because PyTorch has to compile the graph
|
|
||||||
before it can reuse it.
|
|
||||||
|
|
||||||
Example files live under `example_data/configs/`, including
|
Example files live under `example_data/configs/`, including
|
||||||
`example_data/configs/training.yaml`.
|
`example_data/configs/training.yaml`.
|
||||||
|
|
||||||
|
|||||||
2
justfile
2
justfile
@ -136,7 +136,7 @@ clean: clean-build clean-pyc clean-test clean-docs
|
|||||||
|
|
||||||
# Train on example data.
|
# Train on example data.
|
||||||
example-train OPTIONS="":
|
example-train OPTIONS="":
|
||||||
uv run batdetect2 -v train \
|
uv run batdetect2 train \
|
||||||
--val-dataset example_data/dataset.yaml \
|
--val-dataset example_data/dataset.yaml \
|
||||||
--base-dir . \
|
--base-dir . \
|
||||||
--targets example_data/targets.yaml \
|
--targets example_data/targets.yaml \
|
||||||
|
|||||||
@ -1,6 +1,6 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "batdetect2"
|
name = "batdetect2"
|
||||||
version = "2.0.0b3"
|
version = "2.0.0b2"
|
||||||
description = "Deep learning model for detecting and classifying bat echolocation calls in high frequency audio recordings."
|
description = "Deep learning model for detecting and classifying bat echolocation calls in high frequency audio recordings."
|
||||||
authors = [
|
authors = [
|
||||||
{ "name" = "Oisin Mac Aodha", "email" = "oisin.macaodha@ed.ac.uk" },
|
{ "name" = "Oisin Mac Aodha", "email" = "oisin.macaodha@ed.ac.uk" },
|
||||||
@ -11,9 +11,10 @@ dependencies = [
|
|||||||
"deepmerge>=2.0",
|
"deepmerge>=2.0",
|
||||||
"hydra-core>=1.3.2",
|
"hydra-core>=1.3.2",
|
||||||
"librosa>=0.10.1",
|
"librosa>=0.10.1",
|
||||||
"lightning==2.5.0",
|
"lightning[extra]==2.5.0",
|
||||||
"loguru>=0.7.3",
|
"loguru>=0.7.3",
|
||||||
"matplotlib>=3.7.1",
|
"matplotlib>=3.7.1",
|
||||||
|
"netcdf4>=1.6.5",
|
||||||
"numpy>=1.23.5",
|
"numpy>=1.23.5",
|
||||||
"pandas>=1.5.3",
|
"pandas>=1.5.3",
|
||||||
"pydantic>=2.0.0",
|
"pydantic>=2.0.0",
|
||||||
@ -23,10 +24,9 @@ dependencies = [
|
|||||||
"seaborn>=0.13.2",
|
"seaborn>=0.13.2",
|
||||||
"soundevent[audio,geometry,plot]>=2.10.0",
|
"soundevent[audio,geometry,plot]>=2.10.0",
|
||||||
"soundfile>=0.12.1",
|
"soundfile>=0.12.1",
|
||||||
"tabulate>=0.10.0",
|
"tensorboard>=2.16.2",
|
||||||
"torch>=2.0.0",
|
"torch>=2.0.0",
|
||||||
"torchaudio>=2.0.0",
|
"torchaudio>=2.0.0",
|
||||||
"tqdm>=4.70.0",
|
|
||||||
"xarray>=2024.0.0",
|
"xarray>=2024.0.0",
|
||||||
]
|
]
|
||||||
requires-python = ">=3.10,<3.14"
|
requires-python = ">=3.10,<3.14"
|
||||||
@ -88,7 +88,6 @@ dev = [
|
|||||||
"deepdiff>=8.6.1",
|
"deepdiff>=8.6.1",
|
||||||
"pytest-xdist[psutil]>=3.8.0",
|
"pytest-xdist[psutil]>=3.8.0",
|
||||||
]
|
]
|
||||||
tensorboard = ["tensorboard>=2.16.2"]
|
|
||||||
dvclive = ["dvclive>=3.48.2"]
|
dvclive = ["dvclive>=3.48.2"]
|
||||||
mlflow = ["mlflow>=3.1.1"]
|
mlflow = ["mlflow>=3.1.1"]
|
||||||
gradio = [
|
gradio = [
|
||||||
|
|||||||
@ -8,7 +8,6 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from lightning.pytorch.loggers import Logger
|
|
||||||
from soundevent import data
|
from soundevent import data
|
||||||
|
|
||||||
from batdetect2.audio import AudioConfig, AudioLoader
|
from batdetect2.audio import AudioConfig, AudioLoader
|
||||||
@ -153,19 +152,6 @@ class BatDetect2API:
|
|||||||
|
|
||||||
self.model.eval()
|
self.model.eval()
|
||||||
|
|
||||||
def compile(self) -> "BatDetect2API":
|
|
||||||
"""Compile the detector path used by inference.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
BatDetect2API
|
|
||||||
This API instance with the detector compiled.
|
|
||||||
"""
|
|
||||||
from batdetect2.models import compile_model
|
|
||||||
|
|
||||||
compile_model(self.model)
|
|
||||||
return self
|
|
||||||
|
|
||||||
def load_annotations(
|
def load_annotations(
|
||||||
self,
|
self,
|
||||||
path: data.PathLike,
|
path: data.PathLike,
|
||||||
@ -206,7 +192,6 @@ class BatDetect2API:
|
|||||||
train_config: TrainingConfig | None = None,
|
train_config: TrainingConfig | None = None,
|
||||||
logger_config: LoggerConfig | None = None,
|
logger_config: LoggerConfig | None = None,
|
||||||
logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (),
|
logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (),
|
||||||
train_logger: Logger | None = None,
|
|
||||||
):
|
):
|
||||||
"""Train the current model on a set of annotations.
|
"""Train the current model on a set of annotations.
|
||||||
|
|
||||||
@ -241,9 +226,6 @@ class BatDetect2API:
|
|||||||
Training logger config override.
|
Training logger config override.
|
||||||
logging_callbacks : Sequence[LoggingCallback[TrainLoggingContext]], optional
|
logging_callbacks : Sequence[LoggingCallback[TrainLoggingContext]], optional
|
||||||
Extra logging callbacks to run during training setup.
|
Extra logging callbacks to run during training setup.
|
||||||
train_logger : Logger | None, optional
|
|
||||||
Pre-built Lightning logger to use for training. If omitted, one is
|
|
||||||
built from ``logger_config``.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
@ -273,7 +255,6 @@ class BatDetect2API:
|
|||||||
audio_config=audio_config or self.audio_config,
|
audio_config=audio_config or self.audio_config,
|
||||||
logger_config=logger_config or self.logging_config.train,
|
logger_config=logger_config or self.logging_config.train,
|
||||||
logging_callbacks=logging_callbacks,
|
logging_callbacks=logging_callbacks,
|
||||||
train_logger=train_logger,
|
|
||||||
)
|
)
|
||||||
self.model.eval()
|
self.model.eval()
|
||||||
return self
|
return self
|
||||||
@ -998,7 +979,6 @@ class BatDetect2API:
|
|||||||
inference_config: InferenceConfig | None = None,
|
inference_config: InferenceConfig | None = None,
|
||||||
outputs_config: OutputsConfig | None = None,
|
outputs_config: OutputsConfig | None = None,
|
||||||
logging_config: AppLoggingConfig | None = None,
|
logging_config: AppLoggingConfig | None = None,
|
||||||
compile_model: bool = False,
|
|
||||||
) -> "BatDetect2API":
|
) -> "BatDetect2API":
|
||||||
"""Build an API instance from config objects.
|
"""Build an API instance from config objects.
|
||||||
|
|
||||||
@ -1024,8 +1004,6 @@ class BatDetect2API:
|
|||||||
Output config. If omitted, the default outputs config is used.
|
Output config. If omitted, the default outputs config is used.
|
||||||
logging_config : AppLoggingConfig | None, optional
|
logging_config : AppLoggingConfig | None, optional
|
||||||
Logging config. If omitted, the default logging config is used.
|
Logging config. If omitted, the default logging config is used.
|
||||||
compile_model : bool, optional
|
|
||||||
If ``True``, compile the detector path after building the API.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
@ -1108,7 +1086,7 @@ class BatDetect2API:
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
api = cls(
|
return cls(
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
audio_config=audio_config,
|
audio_config=audio_config,
|
||||||
train_config=train_config,
|
train_config=train_config,
|
||||||
@ -1127,11 +1105,6 @@ class BatDetect2API:
|
|||||||
output_transform=output_transform,
|
output_transform=output_transform,
|
||||||
)
|
)
|
||||||
|
|
||||||
if compile_model:
|
|
||||||
api.compile()
|
|
||||||
|
|
||||||
return api
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_checkpoint(
|
def from_checkpoint(
|
||||||
cls,
|
cls,
|
||||||
@ -1142,7 +1115,6 @@ class BatDetect2API:
|
|||||||
inference_config: InferenceConfig | None = None,
|
inference_config: InferenceConfig | None = None,
|
||||||
outputs_config: OutputsConfig | None = None,
|
outputs_config: OutputsConfig | None = None,
|
||||||
logging_config: AppLoggingConfig | None = None,
|
logging_config: AppLoggingConfig | None = None,
|
||||||
compile_model: bool = False,
|
|
||||||
) -> "BatDetect2API":
|
) -> "BatDetect2API":
|
||||||
"""Build an API instance from a saved checkpoint.
|
"""Build an API instance from a saved checkpoint.
|
||||||
|
|
||||||
@ -1163,8 +1135,6 @@ class BatDetect2API:
|
|||||||
Output config override.
|
Output config override.
|
||||||
logging_config : AppLoggingConfig | None, optional
|
logging_config : AppLoggingConfig | None, optional
|
||||||
Logging config override.
|
Logging config override.
|
||||||
compile_model : bool, optional
|
|
||||||
If ``True``, compile the detector path after building the API.
|
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
@ -1248,7 +1218,7 @@ class BatDetect2API:
|
|||||||
transform=output_transform,
|
transform=output_transform,
|
||||||
)
|
)
|
||||||
|
|
||||||
api = cls(
|
return cls(
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
audio_config=audio_config,
|
audio_config=audio_config,
|
||||||
train_config=train_config,
|
train_config=train_config,
|
||||||
@ -1267,11 +1237,6 @@ class BatDetect2API:
|
|||||||
output_transform=output_transform,
|
output_transform=output_transform,
|
||||||
)
|
)
|
||||||
|
|
||||||
if compile_model:
|
|
||||||
api.compile()
|
|
||||||
|
|
||||||
return api
|
|
||||||
|
|
||||||
def _set_trainable_parameters(
|
def _set_trainable_parameters(
|
||||||
self,
|
self,
|
||||||
trainable: Literal["all", "heads", "classifier_head", "size_head"],
|
trainable: Literal["all", "heads", "classifier_head", "size_head"],
|
||||||
|
|||||||
@ -10,7 +10,6 @@ from batdetect2.inference.clips import get_clips_from_files
|
|||||||
from batdetect2.inference.config import InferenceConfig
|
from batdetect2.inference.config import InferenceConfig
|
||||||
from batdetect2.inference.dataset import build_inference_loader
|
from batdetect2.inference.dataset import build_inference_loader
|
||||||
from batdetect2.inference.lightning import InferenceModule
|
from batdetect2.inference.lightning import InferenceModule
|
||||||
from batdetect2.models import compile_model
|
|
||||||
from batdetect2.models.types import ModelProtocol
|
from batdetect2.models.types import ModelProtocol
|
||||||
from batdetect2.outputs import (
|
from batdetect2.outputs import (
|
||||||
OutputsConfig,
|
OutputsConfig,
|
||||||
@ -72,9 +71,6 @@ def run_batch_inference(
|
|||||||
batch_size=batch_size,
|
batch_size=batch_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
if inference_config.compile_model:
|
|
||||||
compile_model(model)
|
|
||||||
|
|
||||||
module = InferenceModule(
|
module = InferenceModule(
|
||||||
model,
|
model,
|
||||||
output_transform=output_transform,
|
output_transform=output_transform,
|
||||||
|
|||||||
@ -15,7 +15,6 @@ class ClipingConfig(BaseConfig):
|
|||||||
|
|
||||||
|
|
||||||
class InferenceConfig(BaseConfig):
|
class InferenceConfig(BaseConfig):
|
||||||
compile_model: bool = False
|
|
||||||
loader: InferenceLoaderConfig = Field(
|
loader: InferenceLoaderConfig = Field(
|
||||||
default_factory=InferenceLoaderConfig
|
default_factory=InferenceLoaderConfig
|
||||||
)
|
)
|
||||||
|
|||||||
@ -201,14 +201,7 @@ def create_tensorboard_logger(
|
|||||||
experiment_name: str | None = None,
|
experiment_name: str | None = None,
|
||||||
run_name: str | None = None,
|
run_name: str | None = None,
|
||||||
) -> Logger:
|
) -> Logger:
|
||||||
try:
|
|
||||||
from lightning.pytorch.loggers import TensorBoardLogger
|
from lightning.pytorch.loggers import TensorBoardLogger
|
||||||
except ImportError as error:
|
|
||||||
raise ValueError(
|
|
||||||
"TensorBoard is not installed and cannot be used for logging. "
|
|
||||||
"Make sure you have it installed by running `pip install tensorboard` "
|
|
||||||
"or `uv add tensorboard`"
|
|
||||||
) from error
|
|
||||||
|
|
||||||
if log_dir is None:
|
if log_dir is None:
|
||||||
log_dir = Path(config.log_dir)
|
log_dir = Path(config.log_dir)
|
||||||
|
|||||||
@ -100,7 +100,6 @@ __all__ = [
|
|||||||
"ModelConfig",
|
"ModelConfig",
|
||||||
"build_model",
|
"build_model",
|
||||||
"build_model_with_new_targets",
|
"build_model_with_new_targets",
|
||||||
"compile_model",
|
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
@ -320,15 +319,3 @@ def build_model_with_new_targets(
|
|||||||
dimension_names=roi_mapper.dimension_names,
|
dimension_names=roi_mapper.dimension_names,
|
||||||
config=model.get_config(),
|
config=model.get_config(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def compile_model(model: ModelProtocol) -> ModelProtocol:
|
|
||||||
"""Compile the detector path used by training and inference."""
|
|
||||||
if not isinstance(model.detector, torch.nn.Module):
|
|
||||||
raise TypeError("Detector must be a torch.nn.Module to compile.")
|
|
||||||
|
|
||||||
if getattr(model.detector, "_compiled_call_impl", None) is not None:
|
|
||||||
return model
|
|
||||||
|
|
||||||
model.detector.compile()
|
|
||||||
return model
|
|
||||||
|
|||||||
@ -1,5 +1,5 @@
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import List, Literal, Sequence, TypedDict
|
from typing import List, Literal, Sequence
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@ -26,11 +26,6 @@ class ParquetOutputConfig(BaseConfig):
|
|||||||
include_geometry: bool = True
|
include_geometry: bool = True
|
||||||
|
|
||||||
|
|
||||||
class ClipInfo(TypedDict):
|
|
||||||
clip: data.Clip
|
|
||||||
preds: list[Detection]
|
|
||||||
|
|
||||||
|
|
||||||
class ParquetFormatter(OutputFormatterProtocol[ClipDetections]):
|
class ParquetFormatter(OutputFormatterProtocol[ClipDetections]):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@ -125,7 +120,7 @@ class ParquetFormatter(OutputFormatterProtocol[ClipDetections]):
|
|||||||
else:
|
else:
|
||||||
df = pd.read_parquet(path)
|
df = pd.read_parquet(path)
|
||||||
|
|
||||||
predictions_by_clip: dict[UUID, ClipInfo] = {}
|
predictions_by_clip = {}
|
||||||
|
|
||||||
for _, row in df.iterrows():
|
for _, row in df.iterrows():
|
||||||
clip_uuid = row["clip_uuid"]
|
clip_uuid = row["clip_uuid"]
|
||||||
|
|||||||
@ -185,25 +185,25 @@ def plot_clip_evaluation(
|
|||||||
label="found GT",
|
label="found GT",
|
||||||
edgecolor=gt_color,
|
edgecolor=gt_color,
|
||||||
facecolor="none" if not fill else gt_color,
|
facecolor="none" if not fill else gt_color,
|
||||||
linestyle=gt_linestyle, # type: ignore
|
linestyle=gt_linestyle,
|
||||||
),
|
),
|
||||||
patches.Patch(
|
patches.Patch(
|
||||||
label="missed GT",
|
label="missed GT",
|
||||||
edgecolor=missed_gt_color,
|
edgecolor=missed_gt_color,
|
||||||
facecolor="none" if not fill else missed_gt_color,
|
facecolor="none" if not fill else missed_gt_color,
|
||||||
linestyle=missed_gt_linestyle, # type: ignore
|
linestyle=missed_gt_linestyle,
|
||||||
),
|
),
|
||||||
patches.Patch(
|
patches.Patch(
|
||||||
label="true Det",
|
label="true Det",
|
||||||
edgecolor=true_pred_color,
|
edgecolor=true_pred_color,
|
||||||
facecolor="none" if not fill else true_pred_color,
|
facecolor="none" if not fill else true_pred_color,
|
||||||
linestyle=true_pred_linestyle, # type: ignore
|
linestyle=true_pred_linestyle,
|
||||||
),
|
),
|
||||||
patches.Patch(
|
patches.Patch(
|
||||||
label="false Det",
|
label="false Det",
|
||||||
edgecolor=false_pred_color,
|
edgecolor=false_pred_color,
|
||||||
facecolor="none" if not fill else false_pred_color,
|
facecolor="none" if not fill else false_pred_color,
|
||||||
linestyle=false_pred_linestyle, # type: ignore
|
linestyle=false_pred_linestyle,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|||||||
@ -1,9 +1,9 @@
|
|||||||
"""Plot heatmaps."""
|
"""Plot heatmaps."""
|
||||||
|
|
||||||
import matplotlib.pyplot as plt
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from matplotlib import axes, patches
|
from matplotlib import axes, patches
|
||||||
|
from matplotlib.cm import get_cmap
|
||||||
from matplotlib.colors import Colormap, LinearSegmentedColormap, to_rgba
|
from matplotlib.colors import Colormap, LinearSegmentedColormap, to_rgba
|
||||||
|
|
||||||
from batdetect2.plotting.common import create_ax
|
from batdetect2.plotting.common import create_ax
|
||||||
@ -80,7 +80,7 @@ def plot_classification_heatmap(
|
|||||||
raise ValueError("Inconsistent number of class names")
|
raise ValueError("Inconsistent number of class names")
|
||||||
|
|
||||||
if not isinstance(cmap, Colormap):
|
if not isinstance(cmap, Colormap):
|
||||||
cmap = plt.get_cmap(cmap)
|
cmap = get_cmap(cmap)
|
||||||
|
|
||||||
handles = []
|
handles = []
|
||||||
|
|
||||||
|
|||||||
@ -39,7 +39,6 @@ class PLTrainerConfig(BaseConfig):
|
|||||||
|
|
||||||
|
|
||||||
class TrainingConfig(BaseConfig):
|
class TrainingConfig(BaseConfig):
|
||||||
compile_model: bool = False
|
|
||||||
train_loader: TrainLoaderConfig = Field(default_factory=TrainLoaderConfig)
|
train_loader: TrainLoaderConfig = Field(default_factory=TrainLoaderConfig)
|
||||||
val_loader: ValLoaderConfig = Field(default_factory=ValLoaderConfig)
|
val_loader: ValLoaderConfig = Field(default_factory=ValLoaderConfig)
|
||||||
optimizer: OptimizerConfig = Field(default_factory=AdamOptimizerConfig)
|
optimizer: OptimizerConfig = Field(default_factory=AdamOptimizerConfig)
|
||||||
|
|||||||
@ -23,6 +23,21 @@ __all__ = [
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class CosineAnnealingSchedulerConfig(BaseConfig):
|
||||||
|
"""Configuration for ``CosineAnnealingLR``.
|
||||||
|
|
||||||
|
Attributes
|
||||||
|
----------
|
||||||
|
name : Literal["cosine_annealing"]
|
||||||
|
Discriminator field used by the scheduler registry.
|
||||||
|
t_max : int
|
||||||
|
Number of epochs to complete one cosine cycle.
|
||||||
|
"""
|
||||||
|
|
||||||
|
name: Literal["cosine_annealing"] = "cosine_annealing"
|
||||||
|
t_max: int = 200
|
||||||
|
|
||||||
|
|
||||||
scheduler_registry: Registry[LRScheduler, [Optimizer]] = Registry("scheduler")
|
scheduler_registry: Registry[LRScheduler, [Optimizer]] = Registry("scheduler")
|
||||||
|
|
||||||
|
|
||||||
@ -38,24 +53,6 @@ class SchedulerImportConfig(ImportConfig):
|
|||||||
name: Literal["import"] = "import"
|
name: Literal["import"] = "import"
|
||||||
|
|
||||||
|
|
||||||
class CosineAnnealingSchedulerConfig(BaseConfig):
|
|
||||||
"""Configuration for ``CosineAnnealingLR``.
|
|
||||||
|
|
||||||
Attributes
|
|
||||||
----------
|
|
||||||
name : Literal["cosine_annealing"]
|
|
||||||
Discriminator field used by the scheduler registry.
|
|
||||||
t_max : int
|
|
||||||
Number of epochs to complete one cosine cycle.
|
|
||||||
eta_min : float, optional
|
|
||||||
Minimum learning rate. Defaults to 0.
|
|
||||||
"""
|
|
||||||
|
|
||||||
name: Literal["cosine_annealing"] = "cosine_annealing"
|
|
||||||
t_max: int = 200
|
|
||||||
eta_min: float = 0
|
|
||||||
|
|
||||||
|
|
||||||
@scheduler_registry.register(CosineAnnealingSchedulerConfig)
|
@scheduler_registry.register(CosineAnnealingSchedulerConfig)
|
||||||
def build_cosine_scheduler(
|
def build_cosine_scheduler(
|
||||||
config: CosineAnnealingSchedulerConfig,
|
config: CosineAnnealingSchedulerConfig,
|
||||||
@ -66,11 +63,7 @@ def build_cosine_scheduler(
|
|||||||
``t_max`` is interpreted in epochs because Lightning steps the scheduler
|
``t_max`` is interpreted in epochs because Lightning steps the scheduler
|
||||||
once per epoch when ``interval="epoch"`` is used.
|
once per epoch when ``interval="epoch"`` is used.
|
||||||
"""
|
"""
|
||||||
return CosineAnnealingLR(
|
return CosineAnnealingLR(optimizer, T_max=config.t_max)
|
||||||
optimizer,
|
|
||||||
T_max=config.t_max,
|
|
||||||
eta_min=config.eta_min,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
SchedulerConfig = Annotated[
|
SchedulerConfig = Annotated[
|
||||||
|
|||||||
@ -15,7 +15,7 @@ from batdetect2.logging import (
|
|||||||
LoggingCallback,
|
LoggingCallback,
|
||||||
build_logger,
|
build_logger,
|
||||||
)
|
)
|
||||||
from batdetect2.models import ModelConfig, build_model, compile_model
|
from batdetect2.models import ModelConfig, build_model
|
||||||
from batdetect2.models.types import ModelProtocol
|
from batdetect2.models.types import ModelProtocol
|
||||||
from batdetect2.preprocess import PreprocessorProtocol, build_preprocessor
|
from batdetect2.preprocess import PreprocessorProtocol, build_preprocessor
|
||||||
from batdetect2.targets import (
|
from batdetect2.targets import (
|
||||||
@ -71,7 +71,6 @@ def run_train(
|
|||||||
run_name: str | None = None,
|
run_name: str | None = None,
|
||||||
seed: int | None = None,
|
seed: int | None = None,
|
||||||
logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (),
|
logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (),
|
||||||
train_logger: Logger | None = None,
|
|
||||||
):
|
):
|
||||||
if seed is not None:
|
if seed is not None:
|
||||||
seed_everything(seed)
|
seed_everything(seed)
|
||||||
@ -165,7 +164,7 @@ def run_train(
|
|||||||
roi_mapper=roi_mapper,
|
roi_mapper=roi_mapper,
|
||||||
)
|
)
|
||||||
|
|
||||||
train_logger = train_logger or build_logger(
|
train_logger = build_logger(
|
||||||
logger_config or CSVLoggerConfig(),
|
logger_config or CSVLoggerConfig(),
|
||||||
log_dir=log_dir,
|
log_dir=log_dir,
|
||||||
experiment_name=experiment_name,
|
experiment_name=experiment_name,
|
||||||
@ -207,10 +206,6 @@ def run_train(
|
|||||||
run_name=run_name,
|
run_name=run_name,
|
||||||
)
|
)
|
||||||
|
|
||||||
if train_config.compile_model:
|
|
||||||
logger.info("Compiling detector...")
|
|
||||||
compile_model(module.model)
|
|
||||||
|
|
||||||
logger.info("Starting main training loop...")
|
logger.info("Starting main training loop...")
|
||||||
trainer.fit(
|
trainer.fit(
|
||||||
module,
|
module,
|
||||||
|
|||||||
@ -1,7 +1,6 @@
|
|||||||
import uuid
|
import uuid
|
||||||
from dataclasses import dataclass
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Callable, List, Optional, cast
|
from typing import Callable, List, Optional
|
||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
|
|
||||||
import lightning as L
|
import lightning as L
|
||||||
@ -16,7 +15,6 @@ from batdetect2.audio.clips import build_clipper
|
|||||||
from batdetect2.audio.types import AudioLoader, ClipperProtocol
|
from batdetect2.audio.types import AudioLoader, ClipperProtocol
|
||||||
from batdetect2.data import DatasetConfig, load_dataset
|
from batdetect2.data import DatasetConfig, load_dataset
|
||||||
from batdetect2.data.annotations.batdetect2 import BatDetect2FilesAnnotations
|
from batdetect2.data.annotations.batdetect2 import BatDetect2FilesAnnotations
|
||||||
from batdetect2.models.types import ModelProtocol
|
|
||||||
from batdetect2.preprocess import build_preprocessor
|
from batdetect2.preprocess import build_preprocessor
|
||||||
from batdetect2.preprocess.types import PreprocessorProtocol
|
from batdetect2.preprocess.types import PreprocessorProtocol
|
||||||
from batdetect2.targets import (
|
from batdetect2.targets import (
|
||||||
@ -33,12 +31,6 @@ from batdetect2.train.lightning import build_training_module
|
|||||||
from batdetect2.train.types import ClipLabeller
|
from batdetect2.train.types import ClipLabeller
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class DetectorCompileRecorder:
|
|
||||||
compile_count: int = 0
|
|
||||||
call_count: int = 0
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def example_data_dir() -> Path:
|
def example_data_dir() -> Path:
|
||||||
pkg_dir = Path(__file__).parent.parent
|
pkg_dir = Path(__file__).parent.parent
|
||||||
@ -164,15 +156,12 @@ def generate_whistle(tmp_path: Path):
|
|||||||
|
|
||||||
offset = int((time - duration / 2) * samplerate)
|
offset = int((time - duration / 2) * samplerate)
|
||||||
t = np.linspace(-duration / 2, duration / 2, frames, endpoint=False)
|
t = np.linspace(-duration / 2, duration / 2, frames, endpoint=False)
|
||||||
pulse = np.asarray(
|
data = signal.gausspulse(
|
||||||
signal.gausspulse(
|
|
||||||
t,
|
t,
|
||||||
fc=frequency,
|
fc=frequency,
|
||||||
bw=2 / (frequency * whistle_duration),
|
bw=2 / (frequency * whistle_duration),
|
||||||
),
|
|
||||||
dtype=np.float64,
|
|
||||||
)
|
)
|
||||||
wave = (np.roll(pulse, offset) * np.iinfo(np.int16).max).astype(
|
wave = (np.roll(data, offset) * np.iinfo(np.int16).max).astype(
|
||||||
np.int16
|
np.int16
|
||||||
)
|
)
|
||||||
sf.write(str(path), wave, samplerate, subtype="PCM_16")
|
sf.write(str(path), wave, samplerate, subtype="PCM_16")
|
||||||
@ -374,30 +363,6 @@ def sample_audio_loader() -> AudioLoader:
|
|||||||
return build_audio_loader()
|
return build_audio_loader()
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def record_detector_compilation(
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> Callable[[ModelProtocol], DetectorCompileRecorder]:
|
|
||||||
def factory(model: ModelProtocol) -> DetectorCompileRecorder:
|
|
||||||
recorder = DetectorCompileRecorder()
|
|
||||||
detector = cast(Any, model.detector)
|
|
||||||
original_call_impl = detector._call_impl
|
|
||||||
|
|
||||||
def compile_detector() -> None:
|
|
||||||
recorder.compile_count += 1
|
|
||||||
|
|
||||||
def compiled_call(*args, **kwargs):
|
|
||||||
recorder.call_count += 1
|
|
||||||
return original_call_impl(*args, **kwargs)
|
|
||||||
|
|
||||||
detector._compiled_call_impl = compiled_call
|
|
||||||
|
|
||||||
monkeypatch.setattr(detector, "compile", compile_detector)
|
|
||||||
return recorder
|
|
||||||
|
|
||||||
return factory
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def bat_tag() -> data.Tag:
|
def bat_tag() -> data.Tag:
|
||||||
return data.Tag(key="class", value="bat")
|
return data.Tag(key="class", value="bat")
|
||||||
|
|||||||
@ -153,61 +153,6 @@ def test_process_spectrogram_rejects_batched_input(
|
|||||||
api_v2.process_spectrogram(spec)
|
api_v2.process_spectrogram(spec)
|
||||||
|
|
||||||
|
|
||||||
def test_user_can_compile_api_detector(
|
|
||||||
api_v2: BatDetect2API,
|
|
||||||
example_audio_files: list[Path],
|
|
||||||
record_detector_compilation,
|
|
||||||
) -> None:
|
|
||||||
recorder = record_detector_compilation(api_v2.model)
|
|
||||||
audio = api_v2.load_audio(example_audio_files[0])
|
|
||||||
spec = api_v2.generate_spectrogram(audio)
|
|
||||||
|
|
||||||
api_v2.compile()
|
|
||||||
api_v2.compile()
|
|
||||||
api_v2.process_spectrogram(spec)
|
|
||||||
|
|
||||||
assert recorder.compile_count == 1
|
|
||||||
assert recorder.call_count == 1
|
|
||||||
|
|
||||||
|
|
||||||
def test_api_from_config_compiles_detector_when_requested(
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> None:
|
|
||||||
compiled_models = []
|
|
||||||
|
|
||||||
def compile_model(model):
|
|
||||||
compiled_models.append(model)
|
|
||||||
return model
|
|
||||||
|
|
||||||
monkeypatch.setattr("batdetect2.models.compile_model", compile_model)
|
|
||||||
|
|
||||||
api = BatDetect2API.from_config(
|
|
||||||
compile_model=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert compiled_models == [api.model]
|
|
||||||
|
|
||||||
|
|
||||||
def test_api_from_checkpoint_compiles_detector_when_requested(
|
|
||||||
tiny_checkpoint_path: Path,
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
) -> None:
|
|
||||||
compiled_models = []
|
|
||||||
|
|
||||||
def compile_model(model):
|
|
||||||
compiled_models.append(model)
|
|
||||||
return model
|
|
||||||
|
|
||||||
monkeypatch.setattr("batdetect2.models.compile_model", compile_model)
|
|
||||||
|
|
||||||
api = BatDetect2API.from_checkpoint(
|
|
||||||
tiny_checkpoint_path,
|
|
||||||
compile_model=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert compiled_models == [api.model]
|
|
||||||
|
|
||||||
|
|
||||||
def test_user_can_read_top_class_and_other_class_scores(
|
def test_user_can_read_top_class_and_other_class_scores(
|
||||||
api_v2: BatDetect2API,
|
api_v2: BatDetect2API,
|
||||||
example_audio_files: list[Path],
|
example_audio_files: list[Path],
|
||||||
|
|||||||
@ -3,8 +3,6 @@ from pathlib import Path
|
|||||||
import pytest
|
import pytest
|
||||||
from soundevent import data
|
from soundevent import data
|
||||||
|
|
||||||
from batdetect2.api_v2 import BatDetect2API
|
|
||||||
from batdetect2.inference import InferenceConfig
|
|
||||||
from batdetect2.inference.batch import run_batch_inference
|
from batdetect2.inference.batch import run_batch_inference
|
||||||
from batdetect2.targets import build_roi_mapping, build_targets
|
from batdetect2.targets import build_roi_mapping, build_targets
|
||||||
from batdetect2.train import load_model_from_checkpoint
|
from batdetect2.train import load_model_from_checkpoint
|
||||||
@ -55,54 +53,3 @@ def test_run_batch_inference_matches_single_clip_inference(
|
|||||||
strict=True,
|
strict=True,
|
||||||
):
|
):
|
||||||
assert_clip_detections_equal(batched, single)
|
assert_clip_detections_equal(batched, single)
|
||||||
|
|
||||||
|
|
||||||
def test_run_batch_inference_compiles_detector_when_config_requests_compile(
|
|
||||||
example_annotations: list[data.ClipAnnotation],
|
|
||||||
record_detector_compilation,
|
|
||||||
) -> None:
|
|
||||||
api = BatDetect2API.from_config()
|
|
||||||
recorder = record_detector_compilation(api.model)
|
|
||||||
|
|
||||||
predictions = run_batch_inference(
|
|
||||||
api.model,
|
|
||||||
[example_annotations[0].clip],
|
|
||||||
targets=api.targets,
|
|
||||||
roi_mapper=api.roi_mapper,
|
|
||||||
audio_loader=api.audio_loader,
|
|
||||||
preprocessor=api.preprocessor,
|
|
||||||
output_transform=api.output_transform,
|
|
||||||
inference_config=InferenceConfig(compile_model=True),
|
|
||||||
batch_size=1,
|
|
||||||
num_workers=0,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert predictions
|
|
||||||
assert recorder.compile_count == 1
|
|
||||||
assert recorder.call_count > 0
|
|
||||||
|
|
||||||
|
|
||||||
def test_run_batch_inference_does_not_recompile_compiled_detector(
|
|
||||||
example_annotations: list[data.ClipAnnotation],
|
|
||||||
record_detector_compilation,
|
|
||||||
) -> None:
|
|
||||||
api = BatDetect2API.from_config()
|
|
||||||
recorder = record_detector_compilation(api.model)
|
|
||||||
api.compile()
|
|
||||||
|
|
||||||
predictions = run_batch_inference(
|
|
||||||
api.model,
|
|
||||||
[example_annotations[0].clip],
|
|
||||||
targets=api.targets,
|
|
||||||
roi_mapper=api.roi_mapper,
|
|
||||||
audio_loader=api.audio_loader,
|
|
||||||
preprocessor=api.preprocessor,
|
|
||||||
output_transform=api.output_transform,
|
|
||||||
inference_config=InferenceConfig(compile_model=True),
|
|
||||||
batch_size=1,
|
|
||||||
num_workers=0,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert predictions
|
|
||||||
assert recorder.compile_count == 1
|
|
||||||
assert recorder.call_count > 0
|
|
||||||
|
|||||||
@ -49,16 +49,6 @@ def build_default_module(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_fast_train_config() -> TrainingConfig:
|
|
||||||
train_config = TrainingConfig()
|
|
||||||
train_config.trainer.limit_train_batches = 1
|
|
||||||
train_config.trainer.limit_val_batches = 1
|
|
||||||
train_config.trainer.log_every_n_steps = 1
|
|
||||||
train_config.train_loader.batch_size = 1
|
|
||||||
train_config.train_loader.augmentations.enabled = False
|
|
||||||
return train_config
|
|
||||||
|
|
||||||
|
|
||||||
def test_can_initialize_default_module():
|
def test_can_initialize_default_module():
|
||||||
module = build_default_module()
|
module = build_default_module()
|
||||||
assert isinstance(module, L.LightningModule)
|
assert isinstance(module, L.LightningModule)
|
||||||
@ -281,7 +271,19 @@ def test_train_smoke_produces_loadable_checkpoint(
|
|||||||
sample_audio_loader: AudioLoader,
|
sample_audio_loader: AudioLoader,
|
||||||
):
|
):
|
||||||
# Given
|
# Given
|
||||||
train_config = build_fast_train_config()
|
train_config = TrainingConfig.model_validate(
|
||||||
|
{
|
||||||
|
"trainer": {
|
||||||
|
"limit_train_batches": 1,
|
||||||
|
"limit_val_batches": 1,
|
||||||
|
"log_every_n_steps": 1,
|
||||||
|
},
|
||||||
|
"train_loader": {
|
||||||
|
"batch_size": 1,
|
||||||
|
"augmentations": {"enabled": False},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
# When
|
# When
|
||||||
run_train(
|
run_train(
|
||||||
@ -308,47 +310,6 @@ def test_train_smoke_produces_loadable_checkpoint(
|
|||||||
assert outputs is not None
|
assert outputs is not None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.slow
|
|
||||||
def test_run_train_compiles_detector_when_train_config_requests_compile(
|
|
||||||
tmp_path: Path,
|
|
||||||
example_annotations: list[data.ClipAnnotation],
|
|
||||||
record_detector_compilation,
|
|
||||||
) -> None:
|
|
||||||
targets_config = TargetConfig()
|
|
||||||
targets = build_targets(targets_config)
|
|
||||||
roi_mapper = build_roi_mapping(targets_config.roi)
|
|
||||||
model = build_model(
|
|
||||||
ModelConfig(),
|
|
||||||
class_names=targets.class_names,
|
|
||||||
dimension_names=roi_mapper.dimension_names,
|
|
||||||
)
|
|
||||||
train_config = build_fast_train_config()
|
|
||||||
train_config.compile_model = True
|
|
||||||
recorder = record_detector_compilation(model)
|
|
||||||
|
|
||||||
module = run_train(
|
|
||||||
train_annotations=example_annotations[:1],
|
|
||||||
val_annotations=example_annotations[:1],
|
|
||||||
model=model,
|
|
||||||
targets=targets,
|
|
||||||
roi_mapper=roi_mapper,
|
|
||||||
targets_config=targets_config,
|
|
||||||
train_config=train_config,
|
|
||||||
num_epochs=1,
|
|
||||||
train_workers=0,
|
|
||||||
val_workers=0,
|
|
||||||
checkpoint_dir=tmp_path / "checkpoints",
|
|
||||||
log_dir=tmp_path / "logs",
|
|
||||||
seed=0,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert (
|
|
||||||
getattr(module.model.detector, "_compiled_call_impl", None) is not None
|
|
||||||
)
|
|
||||||
assert recorder.compile_count == 1
|
|
||||||
assert recorder.call_count > 0
|
|
||||||
|
|
||||||
|
|
||||||
def test_build_training_module_uses_provided_model() -> None:
|
def test_build_training_module_uses_provided_model() -> None:
|
||||||
targets = build_targets(TargetConfig())
|
targets = build_targets(TargetConfig())
|
||||||
roi_mapper = build_roi_mapping(TargetConfig().roi)
|
roi_mapper = build_roi_mapping(TargetConfig().roi)
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user