Merge pull request #74 from macaodha/fix/raw-format-stack
Some checks are pending
CI / Checks (push) Waiting to run
CI / Tests (Python ${{ matrix.python-version }}) (3.10) (push) Waiting to run
CI / Tests (Python ${{ matrix.python-version }}) (3.11) (push) Waiting to run
CI / Tests (Python ${{ matrix.python-version }}) (3.12) (push) Waiting to run
Docs Pages / Build Docs (push) Waiting to run
Docs Pages / Deploy Docs (push) Blocked by required conditions

Fix raw formatter loading for empty detections and large outputs
This commit is contained in:
Santiago Martinez Balvanera 2026-08-08 10:59:47 +01:00 committed by GitHub
commit d7896c6d75
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 200 additions and 40 deletions

View File

@ -1,4 +1,6 @@
import json
from collections import defaultdict from collections import defaultdict
from multiprocessing import Pool
from pathlib import Path from pathlib import Path
from typing import List, Literal, Sequence from typing import List, Literal, Sequence
from uuid import UUID, uuid4 from uuid import UUID, uuid4
@ -8,6 +10,7 @@ import xarray as xr
from loguru import logger from loguru import logger
from soundevent import data from soundevent import data
from soundevent.geometry import compute_bounds from soundevent.geometry import compute_bounds
from tqdm import tqdm
from batdetect2.core import BaseConfig from batdetect2.core import BaseConfig
from batdetect2.outputs.formats.base import ( from batdetect2.outputs.formats.base import (
@ -25,6 +28,8 @@ class RawOutputConfig(BaseConfig):
include_class_scores: bool = True include_class_scores: bool = True
include_features: bool = True include_features: bool = True
include_geometry: bool = True include_geometry: bool = True
n_jobs: int = 1
show_progress: bool = False
class RawFormatter(OutputFormatterProtocol[ClipDetections]): class RawFormatter(OutputFormatterProtocol[ClipDetections]):
@ -35,12 +40,19 @@ class RawFormatter(OutputFormatterProtocol[ClipDetections]):
include_features: bool = True, include_features: bool = True,
include_geometry: bool = True, include_geometry: bool = True,
parse_full_geometry: bool = False, parse_full_geometry: bool = False,
n_jobs: int = 1,
show_progress: bool = False,
): ):
self.targets = targets self.targets = targets
self.include_class_scores = include_class_scores self.include_class_scores = include_class_scores
self.include_features = include_features self.include_features = include_features
self.include_geometry = include_geometry self.include_geometry = include_geometry
self.parse_full_geometry = parse_full_geometry self.parse_full_geometry = parse_full_geometry
self.n_jobs = n_jobs
self.show_progress = show_progress
if n_jobs < 1:
raise ValueError("n_jobs must be >= 1")
def format( def format(
self, self,
@ -68,15 +80,40 @@ class RawFormatter(OutputFormatterProtocol[ClipDetections]):
def load(self, path: data.PathLike) -> List[ClipDetections]: def load(self, path: data.PathLike) -> List[ClipDetections]:
path = Path(path) path = Path(path)
files = list(path.glob("*.nc")) files = list(path.glob("*.nc"))
predictions: List[ClipDetections] = []
for filepath in files: if self.n_jobs == 1:
logger.debug(f"Loading clip predictions {filepath}") return self._load_sequential(files)
clip_data = xr.load_dataset(filepath)
prediction = self.pred_from_xr(clip_data)
predictions.append(prediction)
return predictions return self._load_parallel(files)
def _load_sequential(
self, files: Sequence[data.PathLike]
) -> List[ClipDetections]:
iterable = files
if self.show_progress:
iterable = tqdm(files, total=len(files))
return [self.load_single_file(filepath) for filepath in iterable]
def _load_parallel(
self, files: Sequence[data.PathLike]
) -> List[ClipDetections]:
with Pool(self.n_jobs) as pool:
if not self.show_progress:
return pool.map(self.load_single_file, files)
return list(
tqdm(
pool.imap(self.load_single_file, files),
total=len(files),
)
)
def load_single_file(self, filepath: data.PathLike) -> ClipDetections:
logger.debug(f"Loading clip predictions {filepath}")
clip_data = xr.load_dataset(filepath)
return self.pred_from_xr(clip_data)
def pred_to_xr( def pred_to_xr(
self, self,
@ -140,7 +177,7 @@ class RawFormatter(OutputFormatterProtocol[ClipDetections]):
"clip_id": str(clip.uuid), "clip_id": str(clip.uuid),
} }
if self.include_class_scores: if self.include_class_scores and values["class_scores"]:
class_scores = np.stack(values["class_scores"], axis=0) class_scores = np.stack(values["class_scores"], axis=0)
data_vars["class_scores"] = ( data_vars["class_scores"] = (
["detection", "classes"], ["detection", "classes"],
@ -148,7 +185,7 @@ class RawFormatter(OutputFormatterProtocol[ClipDetections]):
) )
coords["classes"] = ("classes", self.targets.class_names) coords["classes"] = ("classes", self.targets.class_names)
if self.include_features: if self.include_features and values["features"]:
features = np.stack(values["features"], axis=0) features = np.stack(values["features"], axis=0)
data_vars["features"] = (["detection", "feature"], features) data_vars["features"] = (["detection", "feature"], features)
coords["feature"] = ("feature", np.arange(num_features)) coords["feature"] = ("feature", np.arange(num_features))
@ -167,59 +204,81 @@ class RawFormatter(OutputFormatterProtocol[ClipDetections]):
def pred_from_xr(self, dataset: xr.Dataset) -> ClipDetections: def pred_from_xr(self, dataset: xr.Dataset) -> ClipDetections:
clip_data = dataset clip_data = dataset
recording = data.Recording.model_validate_json( recording = data.Recording.model_validate(
clip_data.attrs["recording"] json.loads(clip_data.attrs["recording"])
) )
clip_id = clip_data.clip_id.item() clip_id = clip_data.clip_id.item()
clip = data.Clip( clip = data.Clip.model_construct(
recording=recording, recording=recording,
uuid=UUID(clip_id), uuid=UUID(clip_id),
start_time=clip_data.clip_start, start_time=float(clip_data.clip_start),
end_time=clip_data.clip_end, end_time=float(clip_data.clip_end),
) )
sound_events = [] sound_events = []
for detection in clip_data.coords["detection"]: num_detections = len(clip_data.coords["detection"])
detection_data = clip_data.sel(detection=detection)
score = detection_data.score.item()
if "geometry" in clip_data and self.parse_full_geometry: scores = clip_data.score.data
geometry = data.geometry_validate( start_times = clip_data.start_time.data
detection_data.geometry.item() end_times = clip_data.end_time.data
) low_freqs = clip_data.low_freq.data
high_freqs = clip_data.high_freq.data
top_class_scores = clip_data.top_class_score.data
top_class = clip_data.top_class.data
num_classes = len(self.targets.class_names)
class_map = dict(
zip(
self.targets.class_names,
range(num_classes),
strict=True,
)
)
geometries = None
if self.parse_full_geometry and "geometry" in clip_data:
geometries = clip_data.geometry.data
class_scores = None
if "class_scores" in clip_data:
class_scores = clip_data.class_scores.data
features = None
if "features" in clip_data:
features = clip_data.features.data
for index in range(num_detections):
score = scores[index]
if geometries is not None:
geometry = data.geometry_validate(geometries[index])
else: else:
start_time = detection_data.start_time.item() start_time = start_times[index]
end_time = detection_data.end_time.item() end_time = end_times[index]
low_freq = detection_data.low_freq.item() low_freq = low_freqs[index]
high_freq = detection_data.high_freq.item() high_freq = high_freqs[index]
geometry = data.BoundingBox.model_construct( geometry = data.BoundingBox.model_construct(
coordinates=[start_time, low_freq, end_time, high_freq] coordinates=[start_time, low_freq, end_time, high_freq]
) )
if "class_scores" in detection_data: if class_scores is not None:
class_scores = detection_data.class_scores.data class_score = class_scores[index]
else: else:
class_scores = np.zeros(len(self.targets.class_names)) class_score = np.zeros(num_classes)
class_index = self.targets.class_names.index( class_index = class_map[top_class[index]]
detection_data.top_class.item() class_score[class_index] = top_class_scores[index]
)
class_scores[class_index] = (
detection_data.top_class_score.item()
)
if "features" in detection_data: feats = features[index] if features is not None else np.zeros(0)
features = detection_data.features.data
else:
features = np.zeros(0)
sound_events.append( sound_events.append(
Detection( Detection(
geometry=geometry, geometry=geometry,
detection_score=score, detection_score=score,
class_scores=class_scores, class_scores=class_score,
features=features, features=feats,
) )
) )
@ -236,4 +295,6 @@ class RawFormatter(OutputFormatterProtocol[ClipDetections]):
include_class_scores=config.include_class_scores, include_class_scores=config.include_class_scores,
include_features=config.include_features, include_features=config.include_features,
include_geometry=config.include_geometry, include_geometry=config.include_geometry,
n_jobs=config.n_jobs,
show_progress=config.show_progress,
) )

View File

@ -61,3 +61,102 @@ def test_roundtrip(
).all() ).all()
assert (recovered_prediction.features == detection.features).all() assert (recovered_prediction.features == detection.features).all()
assert recovered_prediction.geometry == detection.geometry assert recovered_prediction.geometry == detection.geometry
def test_roundtrip_recovers_recording_metadata(
sample_formatter,
create_recording,
create_clip,
sample_targets: TargetProtocol,
tmp_path: Path,
):
recording = create_recording(
tags=[data.Tag(key="source", value="test-recorder")],
duration=2,
samplerate=384_000,
time_expansion=10,
)
clip = create_clip(recording=recording, start_time=0.25, end_time=0.75)
detection = Detection(
geometry=data.BoundingBox(
coordinates=[0.3, 45_000, 0.4, 70_000],
),
detection_score=0.5,
class_scores=np.ones(len(sample_targets.class_names)),
features=np.ones(32),
)
prediction = ClipDetections(clip=clip, detections=[detection])
path = tmp_path / "predictions"
sample_formatter.save(predictions=[prediction], path=path)
recovered = sample_formatter.load(path=path)
assert len(recovered) == 1
assert recovered[0].clip.recording.model_dump(mode="json") == (
recording.model_dump(mode="json")
)
def test_roundtrip_empty_detections(
sample_formatter,
clip: data.Clip,
tmp_path: Path,
):
prediction = ClipDetections(clip=clip, detections=[])
path = tmp_path / "predictions"
sample_formatter.save(predictions=[prediction], path=path)
recovered = sample_formatter.load(path=path)
assert len(recovered) == 1
assert recovered[0].detections == []
assert recovered[0].clip.uuid == prediction.clip.uuid
assert recovered[0].clip.start_time == prediction.clip.start_time
assert recovered[0].clip.end_time == prediction.clip.end_time
def test_roundtrip_loads_with_multiprocessing(
clip: data.Clip,
sample_targets: TargetProtocol,
tmp_path: Path,
):
save_formatter = build_output_formatter(
config=RawOutputConfig(),
targets=sample_targets,
)
load_formatter = build_output_formatter(
config=RawOutputConfig(n_jobs=2),
targets=sample_targets,
)
predictions = [
ClipDetections(
clip=data.Clip(
recording=clip.recording,
start_time=index,
end_time=index + 0.5,
),
detections=[
Detection(
geometry=data.BoundingBox(
coordinates=[index, 45_000, index + 0.1, 70_000],
),
detection_score=0.5,
class_scores=np.ones(len(sample_targets.class_names)),
features=np.ones(32),
)
],
)
for index in range(2)
]
path = tmp_path / "predictions"
save_formatter.save(predictions=predictions, path=path)
recovered = load_formatter.load(path=path)
assert len(recovered) == len(predictions)
assert {item.clip.uuid for item in recovered} == {
item.clip.uuid for item in predictions
}