mirror of
https://github.com/macaodha/batdetect2.git
synced 2026-08-22 03:00:10 +02:00
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
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:
commit
d7896c6d75
@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@ -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
|
||||||
|
}
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user