Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 0 additions & 6 deletions src/osekit/core/base_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,6 @@ def __init__(
"""Instantiate a Dataset object from the Data objects."""
self.data = data
self._name = name
self._has_default_name = name is None
self._suffix = suffix
self._folder = folder

Expand Down Expand Up @@ -101,11 +100,6 @@ def suffix(self) -> str:
def suffix(self, suffix: str | None) -> None:
self._suffix = suffix

@property
def has_default_name(self) -> bool:
"""Return ``True`` if the dataset has a default name, ``False`` if it has a given name."""
return self._has_default_name

@property
def begin(self) -> Timestamp:
"""Begin of the first data object."""
Expand Down
3 changes: 2 additions & 1 deletion src/osekit/core/json_serializer.py
Original file line number Diff line number Diff line change
Expand Up @@ -124,7 +124,8 @@ def serialize_json(path: Path, serialized_dict: dict) -> None:
Dictionary to be serialized.

"""
path.parent.mkdir(parents=True, exist_ok=True)
if not (parent_folder := path.parent).exists():
parent_folder.mkdir(parents=True)
set_path_reference(
serialized_dict=serialized_dict,
root_path=path.parent,
Expand Down
78 changes: 55 additions & 23 deletions src/osekit/public/project.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
from osekit.utils.core import (
file_indexes_per_batch,
get_umask,
locked,
)
from osekit.utils.path import move_tree, ensure_within_base

Expand Down Expand Up @@ -465,12 +466,41 @@ def run(

self.write_json()

@staticmethod
def _reserve_folder(folder: Path) -> None:
"""Create a target folder in which the transform outputs will be exported.

A ``FileExistsError`` is raised if the target folder already exists.
This could happen if a transform with the same name is beeing run
from another process.

Parameters
----------
folder: Path
Folder in which the transform output files will be exported.

"""
try:
folder.mkdir(parents=True, exist_ok=False)
except FileExistsError as e:
msg = (
f"Target folder {folder} already exists.\n"
f"It might mean that another process already ran a transform that"
f"exports in this folder.\n"
f"Change the current transform name or use the"
f"Project.delete_transform_with_outputs() or"
f"Project.rename_transform_with_outputs() method."
)
raise FileExistsError(msg) from e

def _add_audio_dataset(
self,
ads: AudioDataset,
transform_name: str,
) -> None:
ads.folder = self._get_audio_dataset_subpath(ads=ads)
self._reserve_folder(folder=ads.folder)

self.outputs[ads.name] = {
"class": type(ads).__name__,
"transform": transform_name,
Expand All @@ -482,16 +512,7 @@ def _get_audio_dataset_subpath(
self,
ads: AudioDataset,
) -> Path:
return (
self.folder
/ self.SUBFOLDERS["data"]
/ "audio"
/ (
f"{round(ads.data_duration.total_seconds())}_{round(ads.sample_rate)}"
if ads.has_default_name
else ads.name
)
)
return self.folder / self.SUBFOLDERS["data"] / "audio" / ads.name

def export(
self,
Expand Down Expand Up @@ -621,6 +642,7 @@ def _add_spectro_dataset(
transform_name: str,
) -> None:
sds.folder = self._get_spectro_dataset_subpath(sds=sds)
self._reserve_folder(folder=sds.folder)
self.outputs[sds.name] = {
"class": type(sds).__name__,
"dataset": sds,
Expand All @@ -632,15 +654,7 @@ def _get_spectro_dataset_subpath(
self,
sds: SpectroDataset | LTASDataset,
) -> Path:
ads_folder = Path(
f"{round(sds.data_duration.total_seconds())}_{round(sds.fft.fs)}",
)
fft_folder = f"{sds.fft.mfft}_{sds.fft.win.shape[0]}_{sds.fft.hop}_linear"
return (
self.folder
/ self.SUBFOLDERS["processed"]
/ (ads_folder / fft_folder if sds.has_default_name else sds.name)
)
return self.folder / self.SUBFOLDERS["processed"] / sds.name

def _sort_dataset(self, dataset: type[DatasetChild]) -> None:
if type(dataset) is AudioDataset:
Expand Down Expand Up @@ -678,7 +692,7 @@ def _delete_output(self, output_dataset_name: str) -> None:

afm.close()
shutil.rmtree(str(output_to_remove.folder))
self.write_json()
self.write_json(output_to_skip=output_to_remove.name)

def get_output_by_transform_name(
self,
Expand Down Expand Up @@ -820,7 +834,7 @@ def to_dict(self) -> dict:
if isinstance(dataset["dataset"], Path)
else str(dataset["dataset"].folder / f"{name}.json"),
}
for name, dataset in self.outputs.items()
for name, dataset in sorted(self.outputs.items(), key=lambda kv: kv[0])
},
"instrument": (
None if self.instrument is None else self.instrument.to_dict()
Expand Down Expand Up @@ -863,10 +877,28 @@ def from_dict(cls, dictionary: dict) -> Project:
outputs=outputs,
)

def write_json(self, folder: Path | None = None) -> None:
def write_json(
self,
folder: Path | None = None,
output_to_skip: str | None = None,
) -> None:
"""Write a serialized Project to a JSON file."""
folder = folder if folder is not None else self.folder
serialize_json(folder / "project.json", self.to_dict())
json_file = folder / "project.json"

@locked(lock_file=folder / "project.lock")
def _write() -> None:
dictionary = self.to_dict()
if json_file.exists():
# Update outputs in case there are unexisting keys in the dictionary.
existing_outputs = deserialize_json(path=json_file).get("outputs", {})
if output_to_skip and output_to_skip in existing_outputs:
existing_outputs.pop(output_to_skip)
dictionary["outputs"] |= existing_outputs

serialize_json(folder / "project.json", dictionary)

_write()

@classmethod
def from_json(cls, file: Path) -> Project:
Expand Down
9 changes: 3 additions & 6 deletions src/osekit/public/transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ class Transform:
def __init__(
self,
output_type: OutputType,
name: str,
begin: Timestamp | None = None,
end: Timestamp | None = None,
data_duration: Timedelta | None = None,
Expand All @@ -74,7 +75,6 @@ def __init__(
sample_rate: float | None = None,
normalization: Normalization = Normalization.RAW,
butter: Butterworth | None = None,
name: str | None = None,
subtype: str | None = None,
fft: ShortTimeFFT | None = None,
v_lim: tuple[float, float] | None = None,
Expand All @@ -89,6 +89,8 @@ def __init__(
output_type: OutputType
The type of transform to run.
See ``OutputType`` docstring for more info.
name: str | None
Name of the transform dataset.
begin: Timestamp | None
The begin of the transform dataset.
Defaulted to the begin of the original dataset.
Expand Down Expand Up @@ -121,11 +123,6 @@ def __init__(
The type of normalization to apply to the audio data.
butter: Butterworth | None
Butterworth filter to apply to the audio data.
name: str | None
Name of the transform dataset.
Defaulted as the begin timestamp of the transform dataset.
If both audio and spectro outputs are selected, the audio
transform dataset name will be suffixed with ``"_audio"``.
subtype: str | None
Subtype of the written audio files as provided by the soundfile module.
Defaulted as the default ``16-bit PCM`` for ``wav`` audio files.
Expand Down
55 changes: 55 additions & 0 deletions tests/test_core_api_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -796,6 +796,61 @@ def test_base_dataset_from_folder(
assert np.array_equal(sorted(f.path for f in data.files), sorted(expected[1]))


@pytest.mark.parametrize(
("data_durations", "expected_duration"),
[
pytest.param(
[Timedelta(seconds=10)],
Timedelta(seconds=10),
id="only_one_data",
),
pytest.param(
[
Timedelta(seconds=10),
Timedelta(seconds=10),
Timedelta(seconds=10),
],
Timedelta(seconds=10),
id="all_data_have_the_same_duration",
),
pytest.param(
[
Timedelta(seconds=15),
Timedelta(seconds=15),
Timedelta(seconds=10),
],
Timedelta(seconds=15),
id="dataset_duration_is_most_frequent_one",
),
pytest.param(
[
Timedelta(seconds=10),
Timedelta(seconds=20),
Timedelta(seconds=15),
],
Timedelta(seconds=20),
id="only_one_of_each_duration_takes_the_longest",
),
],
)
def test_base_dataset_data_duration(
data_durations: list[Timedelta], expected_duration: Timedelta
) -> None:
files = []
for data_duration in data_durations:
df = DummyFile(
path=Path(),
begin=Timestamp("1994-09-27 00:00:00")
+ Timedelta(seconds=sum(f.duration.total_seconds() for f in files)),
)
df.end = df.begin + data_duration
files.append(df)

assert (
DummyDataset.from_files(files, mode="files").data_duration == expected_duration
)


@pytest.mark.parametrize(
"destination_folder",
[
Expand Down
Loading