context-calibrated-beats 1.0.1__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- context_calibrated_beats-1.0.1/API.py +1497 -0
- context_calibrated_beats-1.0.1/CCB.py +4001 -0
- context_calibrated_beats-1.0.1/CHANGELOG.md +22 -0
- context_calibrated_beats-1.0.1/INSTRUCTIONS.md +347 -0
- context_calibrated_beats-1.0.1/INSTRUCTIONS_CN.md +293 -0
- context_calibrated_beats-1.0.1/LICENSE +21 -0
- context_calibrated_beats-1.0.1/MANIFEST.in +6 -0
- context_calibrated_beats-1.0.1/PKG-INFO +183 -0
- context_calibrated_beats-1.0.1/README.md +153 -0
- context_calibrated_beats-1.0.1/README_CN.md +139 -0
- context_calibrated_beats-1.0.1/THIRD_PARTY_NOTICES.md +24 -0
- context_calibrated_beats-1.0.1/context_calibrated_beats/__init__.py +6 -0
- context_calibrated_beats-1.0.1/context_calibrated_beats/__main__.py +6 -0
- context_calibrated_beats-1.0.1/context_calibrated_beats/api.py +6 -0
- context_calibrated_beats-1.0.1/context_calibrated_beats/cli.py +11 -0
- context_calibrated_beats-1.0.1/context_calibrated_beats.egg-info/PKG-INFO +183 -0
- context_calibrated_beats-1.0.1/context_calibrated_beats.egg-info/SOURCES.txt +24 -0
- context_calibrated_beats-1.0.1/context_calibrated_beats.egg-info/dependency_links.txt +1 -0
- context_calibrated_beats-1.0.1/context_calibrated_beats.egg-info/entry_points.txt +2 -0
- context_calibrated_beats-1.0.1/context_calibrated_beats.egg-info/requires.txt +7 -0
- context_calibrated_beats-1.0.1/context_calibrated_beats.egg-info/top_level.txt +3 -0
- context_calibrated_beats-1.0.1/pyproject.toml +49 -0
- context_calibrated_beats-1.0.1/requirements-bpm.txt +7 -0
- context_calibrated_beats-1.0.1/setup.cfg +4 -0
- context_calibrated_beats-1.0.1/tests/test_api.py +481 -0
- context_calibrated_beats-1.0.1/tests/test_pipeline.py +966 -0
|
@@ -0,0 +1,1497 @@
|
|
|
1
|
+
"""Small Python API for the CCB production pipeline.
|
|
2
|
+
|
|
3
|
+
Example:
|
|
4
|
+
|
|
5
|
+
from context_calibrated_beats import run, set_click_gain, set_music_gain
|
|
6
|
+
|
|
7
|
+
set_music_gain(0.25)
|
|
8
|
+
set_click_gain(0.75)
|
|
9
|
+
result = run("song.mp3")
|
|
10
|
+
print(result.beats_csv)
|
|
11
|
+
print(result.report["result"]["dominant_bpm"])
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import argparse
|
|
17
|
+
import csv
|
|
18
|
+
import json
|
|
19
|
+
import math
|
|
20
|
+
import os
|
|
21
|
+
import shutil
|
|
22
|
+
import tempfile
|
|
23
|
+
from collections.abc import Iterable
|
|
24
|
+
from dataclasses import dataclass
|
|
25
|
+
from datetime import datetime
|
|
26
|
+
from pathlib import Path
|
|
27
|
+
from typing import Any
|
|
28
|
+
|
|
29
|
+
import CCB as ccb
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
API_VERSION = "1.0"
|
|
33
|
+
__version__ = ccb.CCB_VERSION
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class CCBError(Exception):
|
|
37
|
+
"""Base class for errors intentionally exposed by the CCB API."""
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class InvalidArgumentError(CCBError, ValueError):
|
|
41
|
+
"""A public API argument has an invalid value."""
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class ResourceNotFoundError(CCBError, FileNotFoundError):
|
|
45
|
+
"""An input audio file or previously generated result is unavailable."""
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
class ResultStateError(CCBError, ValueError):
|
|
49
|
+
"""Stored result files are incomplete, inconsistent, or belong elsewhere."""
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
class EditConflictError(InvalidArgumentError):
|
|
53
|
+
"""A requested manual edit conflicts with an existing beat or range."""
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
class ItemNotFoundError(CCBError, KeyError):
|
|
57
|
+
"""A requested beat or NO_BEAT identifier does not exist."""
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
_music_gain = ccb.DEFAULT_MUSIC_GAIN
|
|
61
|
+
_click_gain = ccb.DEFAULT_CLICK_GAIN
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def set_music_gain(value: float) -> None:
|
|
65
|
+
"""Set the process-wide music gain used by subsequent ``run`` calls."""
|
|
66
|
+
global _music_gain
|
|
67
|
+
try:
|
|
68
|
+
music_gain, _ = ccb.validate_mix_gains(value, _click_gain)
|
|
69
|
+
except ValueError as exc:
|
|
70
|
+
raise InvalidArgumentError(str(exc)) from exc
|
|
71
|
+
_music_gain = music_gain
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def set_click_gain(value: float) -> None:
|
|
75
|
+
"""Set the process-wide click gain used by subsequent ``run`` calls."""
|
|
76
|
+
global _click_gain
|
|
77
|
+
try:
|
|
78
|
+
_, click_gain = ccb.validate_mix_gains(_music_gain, value)
|
|
79
|
+
except ValueError as exc:
|
|
80
|
+
raise InvalidArgumentError(str(exc)) from exc
|
|
81
|
+
_click_gain = click_gain
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
@dataclass(frozen=True)
|
|
85
|
+
class RunResult:
|
|
86
|
+
"""Paths and report produced by one successful CCB run."""
|
|
87
|
+
|
|
88
|
+
audio: Path
|
|
89
|
+
output_directory: Path
|
|
90
|
+
beats_csv: Path
|
|
91
|
+
click_wav: Path | None
|
|
92
|
+
overview_png: Path
|
|
93
|
+
segments_csv: Path
|
|
94
|
+
report_json: Path
|
|
95
|
+
report: dict[str, Any]
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
@dataclass(frozen=True)
|
|
99
|
+
class NoBeatRange:
|
|
100
|
+
"""One time-ordered NO_BEAT range from a song's segments file."""
|
|
101
|
+
|
|
102
|
+
segment_id: int
|
|
103
|
+
start_seconds: float
|
|
104
|
+
end_seconds: float
|
|
105
|
+
source: str
|
|
106
|
+
note: str
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
@dataclass(frozen=True)
|
|
110
|
+
class Beat:
|
|
111
|
+
"""One beat from the current time-ordered beat grid."""
|
|
112
|
+
|
|
113
|
+
beat_id: int
|
|
114
|
+
time_seconds: float
|
|
115
|
+
is_downbeat: bool
|
|
116
|
+
reliability_class: str
|
|
117
|
+
reliability_reason: str
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
@dataclass(frozen=True)
|
|
121
|
+
class CacheEntry:
|
|
122
|
+
"""One time-ordered local Beat This! inference cache."""
|
|
123
|
+
|
|
124
|
+
cache_directory: Path
|
|
125
|
+
audio_path: Path | None
|
|
126
|
+
song_name: str
|
|
127
|
+
size_bytes: int
|
|
128
|
+
last_used_at: datetime
|
|
129
|
+
status: str
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
@dataclass(frozen=True)
|
|
133
|
+
class CacheCleanupResult:
|
|
134
|
+
"""Summary of a cache-pruning operation."""
|
|
135
|
+
|
|
136
|
+
deleted: tuple[CacheEntry, ...]
|
|
137
|
+
kept: tuple[CacheEntry, ...]
|
|
138
|
+
freed_bytes: int
|
|
139
|
+
dry_run: bool
|
|
140
|
+
|
|
141
|
+
@property
|
|
142
|
+
def deleted_count(self) -> int:
|
|
143
|
+
return len(self.deleted)
|
|
144
|
+
|
|
145
|
+
@property
|
|
146
|
+
def kept_count(self) -> int:
|
|
147
|
+
return len(self.kept)
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
@dataclass(frozen=True)
|
|
151
|
+
class ManualBeatAddition:
|
|
152
|
+
time_seconds: float
|
|
153
|
+
is_downbeat: bool
|
|
154
|
+
note: str
|
|
155
|
+
|
|
156
|
+
|
|
157
|
+
@dataclass(frozen=True)
|
|
158
|
+
class ManualBeatAdjustment:
|
|
159
|
+
original_time_seconds: float
|
|
160
|
+
new_time_seconds: float
|
|
161
|
+
is_downbeat: bool
|
|
162
|
+
note: str
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
@dataclass(frozen=True)
|
|
166
|
+
class ManualBeatDeletion:
|
|
167
|
+
original_time_seconds: float
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
@dataclass(frozen=True)
|
|
171
|
+
class ManualBeatEdits:
|
|
172
|
+
added: tuple[ManualBeatAddition, ...]
|
|
173
|
+
adjusted: tuple[ManualBeatAdjustment, ...]
|
|
174
|
+
deleted: tuple[ManualBeatDeletion, ...]
|
|
175
|
+
warnings: tuple[str, ...]
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
@dataclass(frozen=True)
|
|
179
|
+
class ReviewRange:
|
|
180
|
+
start_seconds: float
|
|
181
|
+
end_seconds: float
|
|
182
|
+
classification: str
|
|
183
|
+
reliability_score: float
|
|
184
|
+
reason: str
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
@dataclass(frozen=True)
|
|
188
|
+
class SongInfo:
|
|
189
|
+
audio_path: Path
|
|
190
|
+
output_directory: Path
|
|
191
|
+
duration_seconds: float | None
|
|
192
|
+
result_status: str
|
|
193
|
+
cache_status: str
|
|
194
|
+
beat_count: int
|
|
195
|
+
downbeat_count: int
|
|
196
|
+
dominant_bpm: float | None
|
|
197
|
+
manual_added: int
|
|
198
|
+
manual_adjusted: int
|
|
199
|
+
manual_deleted: int
|
|
200
|
+
no_beat_count: int
|
|
201
|
+
last_processed_at: datetime | None
|
|
202
|
+
last_cache_used_at: datetime | None
|
|
203
|
+
|
|
204
|
+
|
|
205
|
+
@dataclass(frozen=True)
|
|
206
|
+
class ValidationResult:
|
|
207
|
+
valid: bool
|
|
208
|
+
errors: tuple[str, ...]
|
|
209
|
+
warnings: tuple[str, ...]
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def _cache_roots(
|
|
213
|
+
output_dir: str | Path,
|
|
214
|
+
cache_dir: str | Path | None,
|
|
215
|
+
) -> list[Path]:
|
|
216
|
+
if cache_dir is not None:
|
|
217
|
+
candidates = [Path(cache_dir).expanduser()]
|
|
218
|
+
elif os.environ.get("CCB_CACHE_DIR"):
|
|
219
|
+
candidates = [ccb.default_cache_root()]
|
|
220
|
+
else:
|
|
221
|
+
candidates = [
|
|
222
|
+
ccb.default_cache_root(),
|
|
223
|
+
ccb.legacy_cache_root(Path(output_dir).expanduser()),
|
|
224
|
+
]
|
|
225
|
+
roots: list[Path] = []
|
|
226
|
+
seen: set[Path] = set()
|
|
227
|
+
for candidate in candidates:
|
|
228
|
+
resolved = candidate.resolve()
|
|
229
|
+
if resolved not in seen:
|
|
230
|
+
roots.append(resolved)
|
|
231
|
+
seen.add(resolved)
|
|
232
|
+
return roots
|
|
233
|
+
|
|
234
|
+
|
|
235
|
+
def _cache_entry(path: Path) -> CacheEntry:
|
|
236
|
+
required = {
|
|
237
|
+
"metadata": path / "inference.csv",
|
|
238
|
+
"frames": path / "frames.csv",
|
|
239
|
+
"fused": path / "fused_beats.csv",
|
|
240
|
+
}
|
|
241
|
+
audio_path: Path | None = None
|
|
242
|
+
status = "INCOMPLETE"
|
|
243
|
+
if required["metadata"].is_file():
|
|
244
|
+
try:
|
|
245
|
+
metadata = ccb.read_inference_metadata_csv(required["metadata"])
|
|
246
|
+
audio_path = Path(metadata.audio_path).expanduser().resolve()
|
|
247
|
+
if all(item.is_file() for item in required.values()):
|
|
248
|
+
status = "READY" if audio_path.is_file() else "SOURCE_MISSING"
|
|
249
|
+
except (OSError, ValueError, TypeError, KeyError):
|
|
250
|
+
status = "INVALID"
|
|
251
|
+
size_bytes = sum(
|
|
252
|
+
item.stat().st_size
|
|
253
|
+
for item in path.rglob("*")
|
|
254
|
+
if item.is_file()
|
|
255
|
+
)
|
|
256
|
+
return CacheEntry(
|
|
257
|
+
cache_directory=path.resolve(),
|
|
258
|
+
audio_path=audio_path,
|
|
259
|
+
song_name=(audio_path.stem if audio_path is not None else path.name),
|
|
260
|
+
size_bytes=size_bytes,
|
|
261
|
+
last_used_at=datetime.fromtimestamp(path.stat().st_mtime).astimezone(),
|
|
262
|
+
status=status,
|
|
263
|
+
)
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def list_caches(
|
|
267
|
+
*,
|
|
268
|
+
output_dir: str | Path = "results",
|
|
269
|
+
cache_dir: str | Path | None = None,
|
|
270
|
+
) -> list[CacheEntry]:
|
|
271
|
+
"""List system/custom and compatible legacy caches, newest first."""
|
|
272
|
+
entries: list[CacheEntry] = []
|
|
273
|
+
for root in _cache_roots(output_dir, cache_dir):
|
|
274
|
+
if not root.is_dir():
|
|
275
|
+
continue
|
|
276
|
+
for path in root.iterdir():
|
|
277
|
+
if path.is_dir():
|
|
278
|
+
entries.append(_cache_entry(path))
|
|
279
|
+
return sorted(
|
|
280
|
+
entries,
|
|
281
|
+
key=lambda item: (-item.last_used_at.timestamp(), item.song_name.casefold()),
|
|
282
|
+
)
|
|
283
|
+
|
|
284
|
+
|
|
285
|
+
def prune_caches(
|
|
286
|
+
keep: int | Iterable[str | Path] | None = None,
|
|
287
|
+
*,
|
|
288
|
+
output_dir: str | Path = "results",
|
|
289
|
+
cache_dir: str | Path | None = None,
|
|
290
|
+
dry_run: bool = False,
|
|
291
|
+
) -> CacheCleanupResult:
|
|
292
|
+
"""Delete all caches except the newest N or caches for selected audio files."""
|
|
293
|
+
if not isinstance(dry_run, bool):
|
|
294
|
+
raise TypeError("dry_run must be a boolean")
|
|
295
|
+
entries = list_caches(output_dir=output_dir, cache_dir=cache_dir)
|
|
296
|
+
if keep is None:
|
|
297
|
+
kept: list[CacheEntry] = []
|
|
298
|
+
elif isinstance(keep, int) and not isinstance(keep, bool):
|
|
299
|
+
if keep < 0:
|
|
300
|
+
raise InvalidArgumentError("keep must be non-negative")
|
|
301
|
+
kept = entries[:keep]
|
|
302
|
+
else:
|
|
303
|
+
if isinstance(keep, (str, bytes, Path)) or not isinstance(keep, Iterable):
|
|
304
|
+
raise TypeError("keep must be an integer, a file list, or None")
|
|
305
|
+
keep_items = list(keep)
|
|
306
|
+
if any(not isinstance(item, (str, Path)) for item in keep_items):
|
|
307
|
+
raise TypeError("every item in keep must be a string or Path")
|
|
308
|
+
keep_paths = {Path(item).expanduser().resolve() for item in keep_items}
|
|
309
|
+
kept = [item for item in entries if item.audio_path in keep_paths]
|
|
310
|
+
kept_directories = {item.cache_directory for item in kept}
|
|
311
|
+
deleted = [
|
|
312
|
+
item for item in entries if item.cache_directory not in kept_directories
|
|
313
|
+
]
|
|
314
|
+
if not dry_run:
|
|
315
|
+
allowed_roots = {
|
|
316
|
+
root.resolve() for root in _cache_roots(output_dir, cache_dir)
|
|
317
|
+
}
|
|
318
|
+
for item in deleted:
|
|
319
|
+
target = item.cache_directory.resolve()
|
|
320
|
+
if target.parent not in allowed_roots or target == target.parent:
|
|
321
|
+
raise ResultStateError(
|
|
322
|
+
f"Refusing to delete unsafe cache path: {target}"
|
|
323
|
+
)
|
|
324
|
+
shutil.rmtree(target)
|
|
325
|
+
return CacheCleanupResult(
|
|
326
|
+
deleted=tuple(deleted),
|
|
327
|
+
kept=tuple(kept),
|
|
328
|
+
freed_bytes=sum(item.size_bytes for item in deleted),
|
|
329
|
+
dry_run=dry_run,
|
|
330
|
+
)
|
|
331
|
+
|
|
332
|
+
|
|
333
|
+
def _result_paths(
|
|
334
|
+
file_name: str | Path,
|
|
335
|
+
output_dir: str | Path,
|
|
336
|
+
) -> tuple[Path, Path, Path, Path]:
|
|
337
|
+
audio_path = Path(file_name).expanduser()
|
|
338
|
+
if not audio_path.is_file():
|
|
339
|
+
raise ResourceNotFoundError(f"Audio file does not exist: {audio_path}")
|
|
340
|
+
result_dir = ccb.result_directory(
|
|
341
|
+
audio_path, Path(output_dir).expanduser()
|
|
342
|
+
)
|
|
343
|
+
return (
|
|
344
|
+
audio_path,
|
|
345
|
+
result_dir / "beats.csv",
|
|
346
|
+
result_dir / "segments.csv",
|
|
347
|
+
result_dir / "report.json",
|
|
348
|
+
)
|
|
349
|
+
|
|
350
|
+
|
|
351
|
+
def _no_beat_paths(
|
|
352
|
+
file_name: str | Path,
|
|
353
|
+
output_dir: str | Path,
|
|
354
|
+
) -> tuple[Path, Path, Path]:
|
|
355
|
+
audio_path, _, segments_path, report_path = _result_paths(file_name, output_dir)
|
|
356
|
+
if not segments_path.is_file() or not report_path.is_file():
|
|
357
|
+
raise ResourceNotFoundError(
|
|
358
|
+
"NO_BEAT settings do not exist yet; run API.run() for this audio first"
|
|
359
|
+
)
|
|
360
|
+
with report_path.open(encoding="utf-8") as handle:
|
|
361
|
+
report = json.load(handle)
|
|
362
|
+
report_audio = report.get("audio")
|
|
363
|
+
if not report_audio or Path(report_audio).resolve() != audio_path.resolve():
|
|
364
|
+
raise ResultStateError(
|
|
365
|
+
"The existing result belongs to a different audio file with the same name"
|
|
366
|
+
)
|
|
367
|
+
return audio_path, segments_path, report_path
|
|
368
|
+
|
|
369
|
+
|
|
370
|
+
def _load_beat_state(
|
|
371
|
+
file_name: str | Path,
|
|
372
|
+
output_dir: str | Path,
|
|
373
|
+
) -> tuple[Path, Path, dict[str, Any], list[dict[str, str]]]:
|
|
374
|
+
audio_path, beats_path, _, report_path = _result_paths(file_name, output_dir)
|
|
375
|
+
if not beats_path.is_file() or not report_path.is_file():
|
|
376
|
+
raise ResourceNotFoundError(
|
|
377
|
+
"Beat settings do not exist yet; run API.run() for this audio first"
|
|
378
|
+
)
|
|
379
|
+
with report_path.open(encoding="utf-8") as handle:
|
|
380
|
+
report = json.load(handle)
|
|
381
|
+
if Path(report.get("audio", "")).resolve() != audio_path.resolve():
|
|
382
|
+
raise ResultStateError(
|
|
383
|
+
"The existing result belongs to a different audio file with the same name"
|
|
384
|
+
)
|
|
385
|
+
with beats_path.open(newline="", encoding="utf-8-sig") as handle:
|
|
386
|
+
rows = list(csv.DictReader(handle))
|
|
387
|
+
rows.sort(key=lambda row: float(row["beat_time_seconds"]))
|
|
388
|
+
return beats_path, report_path, report, rows
|
|
389
|
+
|
|
390
|
+
|
|
391
|
+
def _as_beat(row: dict[str, str]) -> Beat:
|
|
392
|
+
return Beat(
|
|
393
|
+
beat_id=int(row["beat_index"]),
|
|
394
|
+
time_seconds=float(row["beat_time_seconds"]),
|
|
395
|
+
is_downbeat=ccb._parse_bool(row.get("is_downbeat", "0")),
|
|
396
|
+
reliability_class=row.get("reliability_class", ""),
|
|
397
|
+
reliability_reason=row.get("reliability_reason", ""),
|
|
398
|
+
)
|
|
399
|
+
|
|
400
|
+
|
|
401
|
+
def _validate_beat_time(value: float, duration: float) -> float:
|
|
402
|
+
try:
|
|
403
|
+
time_seconds = float(value)
|
|
404
|
+
except (TypeError, ValueError) as exc:
|
|
405
|
+
raise InvalidArgumentError("Beat time must be a finite number") from exc
|
|
406
|
+
if not math.isfinite(time_seconds):
|
|
407
|
+
raise InvalidArgumentError("Beat time must be a finite number")
|
|
408
|
+
if not 0.0 <= time_seconds <= duration + 1e-9:
|
|
409
|
+
raise InvalidArgumentError(
|
|
410
|
+
f"Beat time must satisfy 0 <= time <= {duration:.6f}"
|
|
411
|
+
)
|
|
412
|
+
return min(time_seconds, duration)
|
|
413
|
+
|
|
414
|
+
|
|
415
|
+
def _validate_beat_collision(
|
|
416
|
+
time_seconds: float,
|
|
417
|
+
rows: list[dict[str, str]],
|
|
418
|
+
*,
|
|
419
|
+
ignored_row: dict[str, str] | None = None,
|
|
420
|
+
) -> None:
|
|
421
|
+
for row in rows:
|
|
422
|
+
if row is ignored_row:
|
|
423
|
+
continue
|
|
424
|
+
existing = float(row["beat_time_seconds"])
|
|
425
|
+
if abs(existing - time_seconds) <= ccb.MANUAL_BEAT_COLLISION_TOLERANCE:
|
|
426
|
+
raise EditConflictError(
|
|
427
|
+
f"Beat time overlaps existing beat {row['beat_index']} at {existing:.9f}s"
|
|
428
|
+
)
|
|
429
|
+
|
|
430
|
+
|
|
431
|
+
def _write_json_atomic(path: Path, value: dict[str, Any]) -> None:
|
|
432
|
+
temporary_path: Path | None = None
|
|
433
|
+
try:
|
|
434
|
+
with tempfile.NamedTemporaryFile(
|
|
435
|
+
mode="w",
|
|
436
|
+
encoding="utf-8",
|
|
437
|
+
dir=path.parent,
|
|
438
|
+
prefix=f".{path.name}.",
|
|
439
|
+
suffix=".tmp",
|
|
440
|
+
delete=False,
|
|
441
|
+
) as handle:
|
|
442
|
+
temporary_path = Path(handle.name)
|
|
443
|
+
json.dump(value, handle, ensure_ascii=False, indent=2)
|
|
444
|
+
temporary_path.replace(path)
|
|
445
|
+
finally:
|
|
446
|
+
if temporary_path is not None:
|
|
447
|
+
temporary_path.unlink(missing_ok=True)
|
|
448
|
+
|
|
449
|
+
|
|
450
|
+
def _write_beats_atomic(path: Path, rows: list[dict[str, str]]) -> None:
|
|
451
|
+
fieldnames = [
|
|
452
|
+
"beat_index", "sample_index", "beat_time_seconds", "is_downbeat",
|
|
453
|
+
"activity_segment_id", "is_no_beat", "interval_midpoint_seconds",
|
|
454
|
+
"raw_local_bpm", "smoothed_local_bpm", "reliability_class",
|
|
455
|
+
"reliability_score", "reliability_reason",
|
|
456
|
+
]
|
|
457
|
+
rows.sort(key=lambda row: float(row["beat_time_seconds"]))
|
|
458
|
+
sample_rate_candidates: list[float] = []
|
|
459
|
+
for row in rows:
|
|
460
|
+
time_seconds = float(row["beat_time_seconds"])
|
|
461
|
+
if (
|
|
462
|
+
time_seconds > 0
|
|
463
|
+
and row.get("sample_index", "")
|
|
464
|
+
and row.get("reliability_class") != "MANUAL_EDIT"
|
|
465
|
+
):
|
|
466
|
+
sample_rate_candidates.append(float(row["sample_index"]) / time_seconds)
|
|
467
|
+
sample_rate = (
|
|
468
|
+
round(float(ccb.np.median(sample_rate_candidates)))
|
|
469
|
+
if sample_rate_candidates else 22050.0
|
|
470
|
+
)
|
|
471
|
+
activity_ranges: list[tuple[float, float]] = []
|
|
472
|
+
segments_path = path.parent / "segments.csv"
|
|
473
|
+
if segments_path.is_file():
|
|
474
|
+
with segments_path.open(newline="", encoding="utf-8-sig") as handle:
|
|
475
|
+
segment_rows = list(csv.DictReader(handle))
|
|
476
|
+
duration = max(
|
|
477
|
+
(float(row["end_seconds"]) for row in segment_rows), default=0.0
|
|
478
|
+
)
|
|
479
|
+
activity_ranges = ccb.active_ranges(
|
|
480
|
+
ccb.read_segments_csv(segments_path, duration), duration
|
|
481
|
+
)
|
|
482
|
+
beat_times = ccb.np.asarray(
|
|
483
|
+
[float(row["beat_time_seconds"]) for row in rows], dtype=float
|
|
484
|
+
)
|
|
485
|
+
midpoints, raw_bpm, smooth_bpm = ccb.local_tempo(beat_times)
|
|
486
|
+
temporary_path: Path | None = None
|
|
487
|
+
try:
|
|
488
|
+
with tempfile.NamedTemporaryFile(
|
|
489
|
+
mode="w", newline="", encoding="utf-8-sig", dir=path.parent,
|
|
490
|
+
prefix=f".{path.name}.", suffix=".tmp", delete=False,
|
|
491
|
+
) as handle:
|
|
492
|
+
temporary_path = Path(handle.name)
|
|
493
|
+
writer = csv.DictWriter(handle, fieldnames=fieldnames)
|
|
494
|
+
writer.writeheader()
|
|
495
|
+
for index, row in enumerate(rows):
|
|
496
|
+
time_seconds = float(row["beat_time_seconds"])
|
|
497
|
+
activity_id = next(
|
|
498
|
+
(
|
|
499
|
+
range_index
|
|
500
|
+
for range_index, (start, end) in enumerate(activity_ranges)
|
|
501
|
+
if start <= time_seconds <= end
|
|
502
|
+
),
|
|
503
|
+
-1,
|
|
504
|
+
)
|
|
505
|
+
row["beat_index"] = str(index + 1)
|
|
506
|
+
row["sample_index"] = str(int(round(time_seconds * sample_rate)))
|
|
507
|
+
row["beat_time_seconds"] = f"{time_seconds:.9f}"
|
|
508
|
+
row["activity_segment_id"] = str(activity_id)
|
|
509
|
+
row["is_no_beat"] = str(int(activity_id < 0))
|
|
510
|
+
previous_activity = (
|
|
511
|
+
int(rows[index - 1].get("activity_segment_id", -1))
|
|
512
|
+
if index else -1
|
|
513
|
+
)
|
|
514
|
+
if index == 0 or activity_id < 0 or activity_id != previous_activity:
|
|
515
|
+
row["interval_midpoint_seconds"] = ""
|
|
516
|
+
row["raw_local_bpm"] = ""
|
|
517
|
+
row["smoothed_local_bpm"] = ""
|
|
518
|
+
else:
|
|
519
|
+
row["interval_midpoint_seconds"] = f"{midpoints[index - 1]:.9f}"
|
|
520
|
+
row["raw_local_bpm"] = f"{raw_bpm[index - 1]:.4f}"
|
|
521
|
+
row["smoothed_local_bpm"] = f"{smooth_bpm[index - 1]:.4f}"
|
|
522
|
+
writer.writerow({name: row.get(name, "") for name in fieldnames})
|
|
523
|
+
temporary_path.replace(path)
|
|
524
|
+
finally:
|
|
525
|
+
if temporary_path is not None:
|
|
526
|
+
temporary_path.unlink(missing_ok=True)
|
|
527
|
+
|
|
528
|
+
|
|
529
|
+
def _manual_edits(report: dict[str, Any]) -> dict[str, list[dict]]:
|
|
530
|
+
return ccb.normalize_manual_beat_edits(report.get("manual_edits"))
|
|
531
|
+
|
|
532
|
+
|
|
533
|
+
def _find_edit(items: list[dict], key: str, time_seconds: float) -> dict | None:
|
|
534
|
+
return next(
|
|
535
|
+
(
|
|
536
|
+
item for item in items
|
|
537
|
+
if abs(float(item.get(key, float("inf"))) - time_seconds)
|
|
538
|
+
<= ccb.MANUAL_BEAT_COLLISION_TOLERANCE
|
|
539
|
+
),
|
|
540
|
+
None,
|
|
541
|
+
)
|
|
542
|
+
|
|
543
|
+
|
|
544
|
+
def _load_segments(
|
|
545
|
+
file_name: str | Path,
|
|
546
|
+
output_dir: str | Path,
|
|
547
|
+
) -> tuple[Path, list[ccb.ActivitySegment], float]:
|
|
548
|
+
_, segments_path, report_path = _no_beat_paths(file_name, output_dir)
|
|
549
|
+
with report_path.open(encoding="utf-8") as handle:
|
|
550
|
+
report = json.load(handle)
|
|
551
|
+
with segments_path.open(newline="", encoding="utf-8-sig") as handle:
|
|
552
|
+
raw_rows = list(csv.DictReader(handle))
|
|
553
|
+
active_ends = [
|
|
554
|
+
float(row["end_seconds"])
|
|
555
|
+
for row in raw_rows
|
|
556
|
+
if not ccb._parse_bool(row.get("no_beat", "0"))
|
|
557
|
+
]
|
|
558
|
+
duration = (
|
|
559
|
+
max(active_ends)
|
|
560
|
+
if active_ends
|
|
561
|
+
else float(report["duration_seconds"])
|
|
562
|
+
)
|
|
563
|
+
return segments_path, ccb.read_segments_csv(segments_path, duration), duration
|
|
564
|
+
|
|
565
|
+
|
|
566
|
+
def _ordered_no_beat_segments(
|
|
567
|
+
segments: list[ccb.ActivitySegment],
|
|
568
|
+
) -> list[ccb.ActivitySegment]:
|
|
569
|
+
ordered = sorted(
|
|
570
|
+
(item for item in segments if item.no_beat),
|
|
571
|
+
key=lambda item: (item.start_seconds, item.end_seconds),
|
|
572
|
+
)
|
|
573
|
+
for segment_id, item in enumerate(ordered, start=1):
|
|
574
|
+
item.segment_id = segment_id
|
|
575
|
+
return ordered
|
|
576
|
+
|
|
577
|
+
|
|
578
|
+
def _as_no_beat_range(item: ccb.ActivitySegment) -> NoBeatRange:
|
|
579
|
+
return NoBeatRange(
|
|
580
|
+
segment_id=item.segment_id,
|
|
581
|
+
start_seconds=item.start_seconds,
|
|
582
|
+
end_seconds=item.end_seconds,
|
|
583
|
+
source=item.source,
|
|
584
|
+
note=item.note,
|
|
585
|
+
)
|
|
586
|
+
|
|
587
|
+
|
|
588
|
+
def _validate_range(
|
|
589
|
+
start_seconds: float,
|
|
590
|
+
end_seconds: float,
|
|
591
|
+
duration: float,
|
|
592
|
+
) -> tuple[float, float]:
|
|
593
|
+
try:
|
|
594
|
+
start = float(start_seconds)
|
|
595
|
+
end = float(end_seconds)
|
|
596
|
+
except (TypeError, ValueError) as exc:
|
|
597
|
+
raise InvalidArgumentError(
|
|
598
|
+
"NO_BEAT start and end must be finite numbers"
|
|
599
|
+
) from exc
|
|
600
|
+
if not math.isfinite(start) or not math.isfinite(end):
|
|
601
|
+
raise InvalidArgumentError("NO_BEAT start and end must be finite numbers")
|
|
602
|
+
if not 0.0 <= start < end <= duration + 1e-9:
|
|
603
|
+
raise InvalidArgumentError(
|
|
604
|
+
f"NO_BEAT range must satisfy 0 <= start < end <= {duration:.6f}"
|
|
605
|
+
)
|
|
606
|
+
return start, min(end, duration)
|
|
607
|
+
|
|
608
|
+
|
|
609
|
+
def _validate_no_overlap(
|
|
610
|
+
target: ccb.ActivitySegment,
|
|
611
|
+
no_beat_segments: list[ccb.ActivitySegment],
|
|
612
|
+
) -> None:
|
|
613
|
+
for item in no_beat_segments:
|
|
614
|
+
if item is target:
|
|
615
|
+
continue
|
|
616
|
+
if target.start_seconds < item.end_seconds and target.end_seconds > item.start_seconds:
|
|
617
|
+
raise EditConflictError(
|
|
618
|
+
"NO_BEAT range overlaps segment "
|
|
619
|
+
f"{item.segment_id}: {item.start_seconds:.6f}-{item.end_seconds:.6f}"
|
|
620
|
+
)
|
|
621
|
+
|
|
622
|
+
|
|
623
|
+
def _write_segments_atomic(
|
|
624
|
+
path: Path,
|
|
625
|
+
segments: list[ccb.ActivitySegment],
|
|
626
|
+
) -> None:
|
|
627
|
+
active = sorted(
|
|
628
|
+
(item for item in segments if not item.no_beat),
|
|
629
|
+
key=lambda item: (item.start_seconds, item.end_seconds),
|
|
630
|
+
)
|
|
631
|
+
for item in active:
|
|
632
|
+
if item.source == "default" and item.start_seconds == 0.0:
|
|
633
|
+
item.segment_id = 0
|
|
634
|
+
blocked = _ordered_no_beat_segments(segments)
|
|
635
|
+
temporary_path: Path | None = None
|
|
636
|
+
try:
|
|
637
|
+
with tempfile.NamedTemporaryFile(
|
|
638
|
+
mode="w",
|
|
639
|
+
newline="",
|
|
640
|
+
encoding="utf-8-sig",
|
|
641
|
+
dir=path.parent,
|
|
642
|
+
prefix=f".{path.name}.",
|
|
643
|
+
suffix=".tmp",
|
|
644
|
+
delete=False,
|
|
645
|
+
) as handle:
|
|
646
|
+
temporary_path = Path(handle.name)
|
|
647
|
+
writer = csv.writer(handle)
|
|
648
|
+
writer.writerow(
|
|
649
|
+
[
|
|
650
|
+
"segment_id",
|
|
651
|
+
"start_seconds",
|
|
652
|
+
"end_seconds",
|
|
653
|
+
"no_beat",
|
|
654
|
+
"source",
|
|
655
|
+
"note",
|
|
656
|
+
]
|
|
657
|
+
)
|
|
658
|
+
for item in [*active, *blocked]:
|
|
659
|
+
writer.writerow(
|
|
660
|
+
[
|
|
661
|
+
item.segment_id,
|
|
662
|
+
f"{item.start_seconds:.9f}",
|
|
663
|
+
f"{item.end_seconds:.9f}",
|
|
664
|
+
int(item.no_beat),
|
|
665
|
+
item.source,
|
|
666
|
+
item.note,
|
|
667
|
+
]
|
|
668
|
+
)
|
|
669
|
+
temporary_path.replace(path)
|
|
670
|
+
finally:
|
|
671
|
+
if temporary_path is not None:
|
|
672
|
+
temporary_path.unlink(missing_ok=True)
|
|
673
|
+
|
|
674
|
+
|
|
675
|
+
def list_no_beat_ranges(
|
|
676
|
+
file_name: str | Path,
|
|
677
|
+
*,
|
|
678
|
+
output_dir: str | Path = "results",
|
|
679
|
+
) -> list[NoBeatRange]:
|
|
680
|
+
"""List all NO_BEAT ranges in chronological ID order."""
|
|
681
|
+
_, segments, _ = _load_segments(file_name, output_dir)
|
|
682
|
+
return [_as_no_beat_range(item) for item in _ordered_no_beat_segments(segments)]
|
|
683
|
+
|
|
684
|
+
|
|
685
|
+
def create_no_beat_range(
|
|
686
|
+
file_name: str | Path,
|
|
687
|
+
start_seconds: float,
|
|
688
|
+
end_seconds: float,
|
|
689
|
+
*,
|
|
690
|
+
output_dir: str | Path = "results",
|
|
691
|
+
note: str = "",
|
|
692
|
+
) -> NoBeatRange:
|
|
693
|
+
"""Create a user NO_BEAT range and renumber all ranges by time."""
|
|
694
|
+
if not isinstance(note, str):
|
|
695
|
+
raise TypeError("note must be a string")
|
|
696
|
+
segments_path, segments, duration = _load_segments(file_name, output_dir)
|
|
697
|
+
start, end = _validate_range(start_seconds, end_seconds, duration)
|
|
698
|
+
target = ccb.ActivitySegment(
|
|
699
|
+
segment_id=0,
|
|
700
|
+
start_seconds=start,
|
|
701
|
+
end_seconds=end,
|
|
702
|
+
no_beat=True,
|
|
703
|
+
source="user",
|
|
704
|
+
note=note.strip(),
|
|
705
|
+
)
|
|
706
|
+
no_beat_segments = _ordered_no_beat_segments(segments)
|
|
707
|
+
_validate_no_overlap(target, no_beat_segments)
|
|
708
|
+
segments.append(target)
|
|
709
|
+
_write_segments_atomic(segments_path, segments)
|
|
710
|
+
return _as_no_beat_range(target)
|
|
711
|
+
|
|
712
|
+
|
|
713
|
+
def update_no_beat_range(
|
|
714
|
+
file_name: str | Path,
|
|
715
|
+
segment_id: int,
|
|
716
|
+
*,
|
|
717
|
+
start_seconds: float | None = None,
|
|
718
|
+
end_seconds: float | None = None,
|
|
719
|
+
note: str | None = None,
|
|
720
|
+
output_dir: str | Path = "results",
|
|
721
|
+
) -> NoBeatRange:
|
|
722
|
+
"""Update one user range, then return its new chronological ID."""
|
|
723
|
+
if not isinstance(segment_id, int) or isinstance(segment_id, bool):
|
|
724
|
+
raise TypeError("segment_id must be an integer")
|
|
725
|
+
if note is not None and not isinstance(note, str):
|
|
726
|
+
raise TypeError("note must be a string or None")
|
|
727
|
+
segments_path, segments, duration = _load_segments(file_name, output_dir)
|
|
728
|
+
no_beat_segments = _ordered_no_beat_segments(segments)
|
|
729
|
+
target = next(
|
|
730
|
+
(item for item in no_beat_segments if item.segment_id == segment_id),
|
|
731
|
+
None,
|
|
732
|
+
)
|
|
733
|
+
if target is None:
|
|
734
|
+
raise ItemNotFoundError(
|
|
735
|
+
f"NO_BEAT segment does not exist: {segment_id}"
|
|
736
|
+
)
|
|
737
|
+
if target.source != "user":
|
|
738
|
+
raise PermissionError(f"NO_BEAT segment {segment_id} is not user-managed")
|
|
739
|
+
start, end = _validate_range(
|
|
740
|
+
target.start_seconds if start_seconds is None else start_seconds,
|
|
741
|
+
target.end_seconds if end_seconds is None else end_seconds,
|
|
742
|
+
duration,
|
|
743
|
+
)
|
|
744
|
+
target.start_seconds = start
|
|
745
|
+
target.end_seconds = end
|
|
746
|
+
if note is not None:
|
|
747
|
+
target.note = note.strip()
|
|
748
|
+
_validate_no_overlap(target, no_beat_segments)
|
|
749
|
+
_write_segments_atomic(segments_path, segments)
|
|
750
|
+
return _as_no_beat_range(target)
|
|
751
|
+
|
|
752
|
+
|
|
753
|
+
def delete_no_beat_range(
|
|
754
|
+
file_name: str | Path,
|
|
755
|
+
segment_id: int,
|
|
756
|
+
*,
|
|
757
|
+
output_dir: str | Path = "results",
|
|
758
|
+
) -> None:
|
|
759
|
+
"""Delete one user range and renumber the remaining ranges by time."""
|
|
760
|
+
if not isinstance(segment_id, int) or isinstance(segment_id, bool):
|
|
761
|
+
raise TypeError("segment_id must be an integer")
|
|
762
|
+
segments_path, segments, _ = _load_segments(file_name, output_dir)
|
|
763
|
+
no_beat_segments = _ordered_no_beat_segments(segments)
|
|
764
|
+
target = next(
|
|
765
|
+
(item for item in no_beat_segments if item.segment_id == segment_id),
|
|
766
|
+
None,
|
|
767
|
+
)
|
|
768
|
+
if target is None:
|
|
769
|
+
raise ItemNotFoundError(
|
|
770
|
+
f"NO_BEAT segment does not exist: {segment_id}"
|
|
771
|
+
)
|
|
772
|
+
if target.source != "user":
|
|
773
|
+
raise PermissionError(f"NO_BEAT segment {segment_id} is not user-managed")
|
|
774
|
+
segments.remove(target)
|
|
775
|
+
_write_segments_atomic(segments_path, segments)
|
|
776
|
+
|
|
777
|
+
|
|
778
|
+
def clear_no_beat_ranges(
|
|
779
|
+
file_name: str | Path,
|
|
780
|
+
*,
|
|
781
|
+
output_dir: str | Path = "results",
|
|
782
|
+
) -> int:
|
|
783
|
+
"""Delete every user-managed NO_BEAT range and return the count."""
|
|
784
|
+
segments_path, segments, _ = _load_segments(file_name, output_dir)
|
|
785
|
+
kept = [
|
|
786
|
+
item
|
|
787
|
+
for item in segments
|
|
788
|
+
if not (item.no_beat and item.source == "user")
|
|
789
|
+
]
|
|
790
|
+
deleted = len(segments) - len(kept)
|
|
791
|
+
if deleted:
|
|
792
|
+
_write_segments_atomic(segments_path, kept)
|
|
793
|
+
return deleted
|
|
794
|
+
|
|
795
|
+
|
|
796
|
+
def get_result(
|
|
797
|
+
file_name: str | Path,
|
|
798
|
+
*,
|
|
799
|
+
output_dir: str | Path = "results",
|
|
800
|
+
) -> RunResult:
|
|
801
|
+
"""Read one existing result without running CCB or touching its cache."""
|
|
802
|
+
audio_path, beats_path, segments_path, report_path = _result_paths(
|
|
803
|
+
file_name, output_dir
|
|
804
|
+
)
|
|
805
|
+
if not report_path.is_file():
|
|
806
|
+
raise ResourceNotFoundError(
|
|
807
|
+
"Result does not exist yet; run API.run() for this audio first"
|
|
808
|
+
)
|
|
809
|
+
with report_path.open(encoding="utf-8") as handle:
|
|
810
|
+
report = json.load(handle)
|
|
811
|
+
if Path(report.get("audio", "")).resolve() != audio_path.resolve():
|
|
812
|
+
raise ResultStateError(
|
|
813
|
+
"The existing result belongs to a different audio file with the same name"
|
|
814
|
+
)
|
|
815
|
+
result_dir = report_path.parent
|
|
816
|
+
result_section = report.get("result", {})
|
|
817
|
+
overview_path = result_dir / result_section.get("overview_png", "overview.png")
|
|
818
|
+
click_name = result_section.get("click_wav")
|
|
819
|
+
click_path = None if not click_name else result_dir / click_name
|
|
820
|
+
required = [beats_path, segments_path, overview_path]
|
|
821
|
+
if click_path is not None:
|
|
822
|
+
required.append(click_path)
|
|
823
|
+
missing = [path.name for path in required if not path.is_file()]
|
|
824
|
+
if missing:
|
|
825
|
+
raise ResourceNotFoundError(
|
|
826
|
+
"Existing result is incomplete; missing: " + ", ".join(missing)
|
|
827
|
+
)
|
|
828
|
+
return RunResult(
|
|
829
|
+
audio=audio_path.resolve(),
|
|
830
|
+
output_directory=result_dir.resolve(),
|
|
831
|
+
beats_csv=beats_path.resolve(),
|
|
832
|
+
click_wav=None if click_path is None else click_path.resolve(),
|
|
833
|
+
overview_png=overview_path.resolve(),
|
|
834
|
+
segments_csv=segments_path.resolve(),
|
|
835
|
+
report_json=report_path.resolve(),
|
|
836
|
+
report=report,
|
|
837
|
+
)
|
|
838
|
+
|
|
839
|
+
|
|
840
|
+
def get_manual_beat_edits(
|
|
841
|
+
file_name: str | Path,
|
|
842
|
+
*,
|
|
843
|
+
output_dir: str | Path = "results",
|
|
844
|
+
) -> ManualBeatEdits:
|
|
845
|
+
"""Read persisted manual beat operations and matching warnings."""
|
|
846
|
+
report = get_result(file_name, output_dir=output_dir).report
|
|
847
|
+
edits = _manual_edits(report)
|
|
848
|
+
return ManualBeatEdits(
|
|
849
|
+
added=tuple(
|
|
850
|
+
ManualBeatAddition(
|
|
851
|
+
time_seconds=float(item["time_seconds"]),
|
|
852
|
+
is_downbeat=bool(item.get("is_downbeat", False)),
|
|
853
|
+
note=str(item.get("note", "")),
|
|
854
|
+
)
|
|
855
|
+
for item in edits["added"]
|
|
856
|
+
if "time_seconds" in item
|
|
857
|
+
),
|
|
858
|
+
adjusted=tuple(
|
|
859
|
+
ManualBeatAdjustment(
|
|
860
|
+
original_time_seconds=float(item["original_time_seconds"]),
|
|
861
|
+
new_time_seconds=float(item["new_time_seconds"]),
|
|
862
|
+
is_downbeat=bool(item.get("is_downbeat", False)),
|
|
863
|
+
note=str(item.get("note", "")),
|
|
864
|
+
)
|
|
865
|
+
for item in edits["adjusted"]
|
|
866
|
+
if "original_time_seconds" in item and "new_time_seconds" in item
|
|
867
|
+
),
|
|
868
|
+
deleted=tuple(
|
|
869
|
+
ManualBeatDeletion(
|
|
870
|
+
original_time_seconds=float(item["original_time_seconds"])
|
|
871
|
+
)
|
|
872
|
+
for item in edits["deleted"]
|
|
873
|
+
if "original_time_seconds" in item
|
|
874
|
+
),
|
|
875
|
+
warnings=tuple(str(item) for item in report.get("manual_edit_warnings", [])),
|
|
876
|
+
)
|
|
877
|
+
|
|
878
|
+
|
|
879
|
+
def get_review_ranges(
|
|
880
|
+
file_name: str | Path,
|
|
881
|
+
*,
|
|
882
|
+
output_dir: str | Path = "results",
|
|
883
|
+
classifications: Iterable[str] | None = None,
|
|
884
|
+
) -> list[ReviewRange]:
|
|
885
|
+
"""Read the report's recommended human-review ranges."""
|
|
886
|
+
selected: set[str] | None = None
|
|
887
|
+
if classifications is not None:
|
|
888
|
+
if isinstance(classifications, (str, bytes)):
|
|
889
|
+
raise TypeError("classifications must be a collection of strings")
|
|
890
|
+
values = list(classifications)
|
|
891
|
+
if any(not isinstance(item, str) for item in values):
|
|
892
|
+
raise TypeError("every classification must be a string")
|
|
893
|
+
selected = set(values)
|
|
894
|
+
report = get_result(file_name, output_dir=output_dir).report
|
|
895
|
+
values = report.get("reliability", {}).get("recommended_review_ranges", [])
|
|
896
|
+
return [
|
|
897
|
+
ReviewRange(
|
|
898
|
+
start_seconds=float(item["start_seconds"]),
|
|
899
|
+
end_seconds=float(item["end_seconds"]),
|
|
900
|
+
classification=str(item["classification"]),
|
|
901
|
+
reliability_score=float(item["reliability_score"]),
|
|
902
|
+
reason=str(item.get("reason", "")),
|
|
903
|
+
)
|
|
904
|
+
for item in values
|
|
905
|
+
if isinstance(item, dict)
|
|
906
|
+
and (selected is None or item.get("classification") in selected)
|
|
907
|
+
]
|
|
908
|
+
|
|
909
|
+
|
|
910
|
+
def inspect_song(
|
|
911
|
+
file_name: str | Path,
|
|
912
|
+
*,
|
|
913
|
+
output_dir: str | Path = "results",
|
|
914
|
+
cache_dir: str | Path | None = None,
|
|
915
|
+
) -> SongInfo:
|
|
916
|
+
"""Summarize result, cache, manual-edit, and NO_BEAT state."""
|
|
917
|
+
audio_path, beats_path, segments_path, report_path = _result_paths(
|
|
918
|
+
file_name, output_dir
|
|
919
|
+
)
|
|
920
|
+
result_dir = report_path.parent
|
|
921
|
+
report: dict[str, Any] = {}
|
|
922
|
+
result_status = "MISSING"
|
|
923
|
+
if report_path.is_file():
|
|
924
|
+
try:
|
|
925
|
+
with report_path.open(encoding="utf-8") as handle:
|
|
926
|
+
report = json.load(handle)
|
|
927
|
+
result_section = report.get("result", {})
|
|
928
|
+
overview = result_dir / result_section.get("overview_png", "overview.png")
|
|
929
|
+
click_name = result_section.get("click_wav")
|
|
930
|
+
required = [beats_path, segments_path, overview]
|
|
931
|
+
if click_name:
|
|
932
|
+
required.append(result_dir / click_name)
|
|
933
|
+
result_status = (
|
|
934
|
+
"READY"
|
|
935
|
+
if Path(report.get("audio", "")).resolve() == audio_path.resolve()
|
|
936
|
+
and all(path.is_file() for path in required)
|
|
937
|
+
else "INCOMPLETE"
|
|
938
|
+
)
|
|
939
|
+
except (OSError, ValueError, TypeError, json.JSONDecodeError):
|
|
940
|
+
result_status = "INVALID"
|
|
941
|
+
report = {}
|
|
942
|
+
matching_caches = [
|
|
943
|
+
item
|
|
944
|
+
for item in list_caches(output_dir=output_dir, cache_dir=cache_dir)
|
|
945
|
+
if item.audio_path == audio_path.resolve()
|
|
946
|
+
]
|
|
947
|
+
cache = matching_caches[0] if matching_caches else None
|
|
948
|
+
edits = _manual_edits(report)
|
|
949
|
+
no_beat_count = 0
|
|
950
|
+
if segments_path.is_file() and report.get("duration_seconds") is not None:
|
|
951
|
+
try:
|
|
952
|
+
no_beat_count = len(
|
|
953
|
+
ccb.no_beat_ranges(
|
|
954
|
+
ccb.read_segments_csv(
|
|
955
|
+
segments_path, float(report["duration_seconds"])
|
|
956
|
+
),
|
|
957
|
+
float(report["duration_seconds"]),
|
|
958
|
+
)
|
|
959
|
+
)
|
|
960
|
+
except (OSError, ValueError, TypeError, KeyError):
|
|
961
|
+
pass
|
|
962
|
+
result_section = report.get("result", {})
|
|
963
|
+
return SongInfo(
|
|
964
|
+
audio_path=audio_path.resolve(),
|
|
965
|
+
output_directory=result_dir.resolve(),
|
|
966
|
+
duration_seconds=(
|
|
967
|
+
float(report["duration_seconds"])
|
|
968
|
+
if report.get("duration_seconds") is not None else None
|
|
969
|
+
),
|
|
970
|
+
result_status=result_status,
|
|
971
|
+
cache_status="MISSING" if cache is None else cache.status,
|
|
972
|
+
beat_count=int(result_section.get("beat_count", 0)),
|
|
973
|
+
downbeat_count=int(result_section.get("downbeat_count", 0)),
|
|
974
|
+
dominant_bpm=(
|
|
975
|
+
float(result_section["dominant_bpm"])
|
|
976
|
+
if result_section.get("dominant_bpm") is not None else None
|
|
977
|
+
),
|
|
978
|
+
manual_added=len(edits["added"]),
|
|
979
|
+
manual_adjusted=len(edits["adjusted"]),
|
|
980
|
+
manual_deleted=len(edits["deleted"]),
|
|
981
|
+
no_beat_count=no_beat_count,
|
|
982
|
+
last_processed_at=(
|
|
983
|
+
datetime.fromtimestamp(report_path.stat().st_mtime).astimezone()
|
|
984
|
+
if report_path.is_file() else None
|
|
985
|
+
),
|
|
986
|
+
last_cache_used_at=None if cache is None else cache.last_used_at,
|
|
987
|
+
)
|
|
988
|
+
|
|
989
|
+
|
|
990
|
+
def list_beats(
|
|
991
|
+
file_name: str | Path,
|
|
992
|
+
*,
|
|
993
|
+
output_dir: str | Path = "results",
|
|
994
|
+
start_seconds: float | None = None,
|
|
995
|
+
end_seconds: float | None = None,
|
|
996
|
+
reliability_class: str | None = None,
|
|
997
|
+
manual_only: bool = False,
|
|
998
|
+
) -> list[Beat]:
|
|
999
|
+
"""List current beats, optionally filtered without changing their IDs."""
|
|
1000
|
+
if reliability_class is not None and not isinstance(reliability_class, str):
|
|
1001
|
+
raise TypeError("reliability_class must be a string or None")
|
|
1002
|
+
if not isinstance(manual_only, bool):
|
|
1003
|
+
raise TypeError("manual_only must be a boolean")
|
|
1004
|
+
_, _, report, rows = _load_beat_state(file_name, output_dir)
|
|
1005
|
+
duration = float(report["duration_seconds"])
|
|
1006
|
+
start = (
|
|
1007
|
+
0.0 if start_seconds is None
|
|
1008
|
+
else _validate_beat_time(start_seconds, duration)
|
|
1009
|
+
)
|
|
1010
|
+
end = (
|
|
1011
|
+
duration if end_seconds is None
|
|
1012
|
+
else _validate_beat_time(end_seconds, duration)
|
|
1013
|
+
)
|
|
1014
|
+
if start > end:
|
|
1015
|
+
raise InvalidArgumentError("start_seconds must not exceed end_seconds")
|
|
1016
|
+
for index, row in enumerate(rows, start=1):
|
|
1017
|
+
row["beat_index"] = str(index)
|
|
1018
|
+
beats = [_as_beat(row) for row in rows]
|
|
1019
|
+
return [
|
|
1020
|
+
item
|
|
1021
|
+
for item in beats
|
|
1022
|
+
if start <= item.time_seconds <= end
|
|
1023
|
+
and (
|
|
1024
|
+
reliability_class is None
|
|
1025
|
+
or item.reliability_class == reliability_class
|
|
1026
|
+
)
|
|
1027
|
+
and (not manual_only or item.reliability_class == "MANUAL_EDIT")
|
|
1028
|
+
]
|
|
1029
|
+
|
|
1030
|
+
|
|
1031
|
+
def validate_result(
|
|
1032
|
+
file_name: str | Path,
|
|
1033
|
+
*,
|
|
1034
|
+
output_dir: str | Path = "results",
|
|
1035
|
+
cache_dir: str | Path | None = None,
|
|
1036
|
+
) -> ValidationResult:
|
|
1037
|
+
"""Validate an existing result without modifying files or cache state."""
|
|
1038
|
+
errors: list[str] = []
|
|
1039
|
+
warnings: list[str] = []
|
|
1040
|
+
audio_path, beats_path, segments_path, report_path = _result_paths(
|
|
1041
|
+
file_name, output_dir
|
|
1042
|
+
)
|
|
1043
|
+
if not report_path.is_file():
|
|
1044
|
+
return ValidationResult(False, ("report.json is missing",), ())
|
|
1045
|
+
try:
|
|
1046
|
+
with report_path.open(encoding="utf-8") as handle:
|
|
1047
|
+
report = json.load(handle)
|
|
1048
|
+
except (OSError, ValueError, TypeError, json.JSONDecodeError) as exc:
|
|
1049
|
+
return ValidationResult(False, (f"report.json is invalid: {exc}",), ())
|
|
1050
|
+
if Path(report.get("audio", "")).resolve() != audio_path.resolve():
|
|
1051
|
+
errors.append("report.json belongs to a different audio file")
|
|
1052
|
+
duration = report.get("duration_seconds")
|
|
1053
|
+
try:
|
|
1054
|
+
duration_seconds = float(duration)
|
|
1055
|
+
except (TypeError, ValueError):
|
|
1056
|
+
duration_seconds = math.nan
|
|
1057
|
+
errors.append("duration_seconds is missing or invalid")
|
|
1058
|
+
result_section = report.get("result", {})
|
|
1059
|
+
result_dir = report_path.parent
|
|
1060
|
+
expected_paths = {
|
|
1061
|
+
"beats.csv": beats_path,
|
|
1062
|
+
"segments.csv": segments_path,
|
|
1063
|
+
"overview.png": result_dir / result_section.get(
|
|
1064
|
+
"overview_png", "overview.png"
|
|
1065
|
+
),
|
|
1066
|
+
}
|
|
1067
|
+
click_name = result_section.get("click_wav")
|
|
1068
|
+
if click_name:
|
|
1069
|
+
expected_paths["click.wav"] = result_dir / click_name
|
|
1070
|
+
for label, path in expected_paths.items():
|
|
1071
|
+
if not path.is_file():
|
|
1072
|
+
errors.append(f"{label} is missing")
|
|
1073
|
+
|
|
1074
|
+
rows: list[dict[str, str]] = []
|
|
1075
|
+
if beats_path.is_file():
|
|
1076
|
+
try:
|
|
1077
|
+
with beats_path.open(newline="", encoding="utf-8-sig") as handle:
|
|
1078
|
+
rows = list(csv.DictReader(handle))
|
|
1079
|
+
indices = [int(row["beat_index"]) for row in rows]
|
|
1080
|
+
if indices != list(range(1, len(rows) + 1)):
|
|
1081
|
+
errors.append("beat IDs are not consecutive")
|
|
1082
|
+
times = [float(row["beat_time_seconds"]) for row in rows]
|
|
1083
|
+
if any(not math.isfinite(value) for value in times):
|
|
1084
|
+
errors.append("beat timestamp is not finite")
|
|
1085
|
+
if any(
|
|
1086
|
+
right - left <= ccb.MANUAL_BEAT_COLLISION_TOLERANCE
|
|
1087
|
+
for left, right in zip(times, times[1:])
|
|
1088
|
+
):
|
|
1089
|
+
errors.append("beat timestamps overlap or are not strictly increasing")
|
|
1090
|
+
if math.isfinite(duration_seconds) and any(
|
|
1091
|
+
value < 0 or value > duration_seconds + 1e-9 for value in times
|
|
1092
|
+
):
|
|
1093
|
+
errors.append("beat timestamp falls outside the audio duration")
|
|
1094
|
+
if result_section.get("beat_count") != len(rows):
|
|
1095
|
+
errors.append("report beat_count does not match beats.csv")
|
|
1096
|
+
downbeat_count = sum(
|
|
1097
|
+
ccb._parse_bool(row.get("is_downbeat", "0")) for row in rows
|
|
1098
|
+
)
|
|
1099
|
+
if result_section.get("downbeat_count") != downbeat_count:
|
|
1100
|
+
errors.append("report downbeat_count does not match beats.csv")
|
|
1101
|
+
except (OSError, ValueError, TypeError, KeyError) as exc:
|
|
1102
|
+
errors.append(f"beats.csv is invalid: {exc}")
|
|
1103
|
+
|
|
1104
|
+
blocked_ranges: list[tuple[float, float]] = []
|
|
1105
|
+
if segments_path.is_file() and math.isfinite(duration_seconds):
|
|
1106
|
+
try:
|
|
1107
|
+
segments = ccb.read_segments_csv(segments_path, duration_seconds)
|
|
1108
|
+
blocked_ranges = ccb.no_beat_ranges(segments, duration_seconds)
|
|
1109
|
+
ordered = sorted(blocked_ranges)
|
|
1110
|
+
if any(
|
|
1111
|
+
right_start < left_end
|
|
1112
|
+
for (_, left_end), (right_start, _) in zip(ordered, ordered[1:])
|
|
1113
|
+
):
|
|
1114
|
+
errors.append("NO_BEAT ranges overlap")
|
|
1115
|
+
for row in rows:
|
|
1116
|
+
time_seconds = float(row["beat_time_seconds"])
|
|
1117
|
+
if any(start <= time_seconds <= end for start, end in blocked_ranges):
|
|
1118
|
+
errors.append(
|
|
1119
|
+
f"beat {row['beat_index']} falls inside a NO_BEAT range"
|
|
1120
|
+
)
|
|
1121
|
+
break
|
|
1122
|
+
except (OSError, ValueError, TypeError, KeyError) as exc:
|
|
1123
|
+
errors.append(f"segments.csv is invalid: {exc}")
|
|
1124
|
+
|
|
1125
|
+
edits = _manual_edits(report)
|
|
1126
|
+
visible_times = [float(row["beat_time_seconds"]) for row in rows]
|
|
1127
|
+
manual_times = [
|
|
1128
|
+
float(row["beat_time_seconds"])
|
|
1129
|
+
for row in rows
|
|
1130
|
+
if row.get("reliability_class") == "MANUAL_EDIT"
|
|
1131
|
+
]
|
|
1132
|
+
for item in edits["added"]:
|
|
1133
|
+
if "time_seconds" not in item:
|
|
1134
|
+
errors.append("manual added edit is missing time_seconds")
|
|
1135
|
+
continue
|
|
1136
|
+
target = float(item["time_seconds"])
|
|
1137
|
+
hidden = any(start <= target <= end for start, end in blocked_ranges)
|
|
1138
|
+
if not hidden and not any(
|
|
1139
|
+
abs(value - target) <= ccb.MANUAL_BEAT_COLLISION_TOLERANCE
|
|
1140
|
+
for value in manual_times
|
|
1141
|
+
):
|
|
1142
|
+
errors.append(f"manual added beat is missing at {target:.9f}s")
|
|
1143
|
+
for item in edits["adjusted"]:
|
|
1144
|
+
if "new_time_seconds" not in item:
|
|
1145
|
+
errors.append("manual adjusted edit is missing new_time_seconds")
|
|
1146
|
+
continue
|
|
1147
|
+
target = float(item["new_time_seconds"])
|
|
1148
|
+
hidden = any(start <= target <= end for start, end in blocked_ranges)
|
|
1149
|
+
if not hidden and not any(
|
|
1150
|
+
abs(value - target) <= ccb.MANUAL_BEAT_COLLISION_TOLERANCE
|
|
1151
|
+
for value in manual_times
|
|
1152
|
+
):
|
|
1153
|
+
errors.append(f"manual adjusted beat is missing at {target:.9f}s")
|
|
1154
|
+
replacement_times = {
|
|
1155
|
+
float(item[key])
|
|
1156
|
+
for category, key in (("added", "time_seconds"), ("adjusted", "new_time_seconds"))
|
|
1157
|
+
for item in edits[category]
|
|
1158
|
+
if key in item
|
|
1159
|
+
}
|
|
1160
|
+
for item in edits["deleted"]:
|
|
1161
|
+
if "original_time_seconds" not in item:
|
|
1162
|
+
errors.append("manual deleted edit is missing original_time_seconds")
|
|
1163
|
+
continue
|
|
1164
|
+
target = float(item["original_time_seconds"])
|
|
1165
|
+
intentionally_replaced = any(
|
|
1166
|
+
abs(value - target) <= ccb.MANUAL_BEAT_COLLISION_TOLERANCE
|
|
1167
|
+
for value in replacement_times
|
|
1168
|
+
)
|
|
1169
|
+
if not intentionally_replaced and any(
|
|
1170
|
+
abs(value - target) <= ccb.MANUAL_BEAT_COLLISION_TOLERANCE
|
|
1171
|
+
for value in visible_times
|
|
1172
|
+
):
|
|
1173
|
+
errors.append(f"manually deleted beat reappeared at {target:.9f}s")
|
|
1174
|
+
warnings.extend(str(item) for item in report.get("manual_edit_warnings", []))
|
|
1175
|
+
|
|
1176
|
+
matching_caches = [
|
|
1177
|
+
item
|
|
1178
|
+
for item in list_caches(output_dir=output_dir, cache_dir=cache_dir)
|
|
1179
|
+
if item.audio_path == audio_path.resolve()
|
|
1180
|
+
]
|
|
1181
|
+
if not matching_caches:
|
|
1182
|
+
warnings.append("inference cache is missing")
|
|
1183
|
+
elif matching_caches[0].status != "READY":
|
|
1184
|
+
warnings.append(f"inference cache status is {matching_caches[0].status}")
|
|
1185
|
+
return ValidationResult(not errors, tuple(errors), tuple(warnings))
|
|
1186
|
+
|
|
1187
|
+
|
|
1188
|
+
def create_beat(
|
|
1189
|
+
file_name: str | Path,
|
|
1190
|
+
time_seconds: float,
|
|
1191
|
+
*,
|
|
1192
|
+
is_downbeat: bool = False,
|
|
1193
|
+
output_dir: str | Path = "results",
|
|
1194
|
+
note: str = "",
|
|
1195
|
+
) -> Beat:
|
|
1196
|
+
"""Add one manual beat and renumber all beats chronologically."""
|
|
1197
|
+
if not isinstance(is_downbeat, bool):
|
|
1198
|
+
raise TypeError("is_downbeat must be a boolean")
|
|
1199
|
+
if not isinstance(note, str):
|
|
1200
|
+
raise TypeError("note must be a string")
|
|
1201
|
+
beats_path, report_path, report, rows = _load_beat_state(file_name, output_dir)
|
|
1202
|
+
value = _validate_beat_time(time_seconds, float(report["duration_seconds"]))
|
|
1203
|
+
_validate_beat_collision(value, rows)
|
|
1204
|
+
edits = _manual_edits(report)
|
|
1205
|
+
edits["added"].append(
|
|
1206
|
+
{"time_seconds": value, "is_downbeat": is_downbeat, "note": note.strip()}
|
|
1207
|
+
)
|
|
1208
|
+
target = {name: "" for name in (
|
|
1209
|
+
"beat_index", "sample_index", "beat_time_seconds", "is_downbeat",
|
|
1210
|
+
"activity_segment_id", "is_no_beat", "interval_midpoint_seconds",
|
|
1211
|
+
"raw_local_bpm", "smoothed_local_bpm", "reliability_class",
|
|
1212
|
+
"reliability_score", "reliability_reason",
|
|
1213
|
+
)}
|
|
1214
|
+
target.update(
|
|
1215
|
+
beat_time_seconds=f"{value:.9f}", is_downbeat=str(int(is_downbeat)),
|
|
1216
|
+
activity_segment_id="0", is_no_beat="0",
|
|
1217
|
+
reliability_class="MANUAL_EDIT", reliability_score="1.000000",
|
|
1218
|
+
reliability_reason="manual_added",
|
|
1219
|
+
)
|
|
1220
|
+
rows.append(target)
|
|
1221
|
+
report["manual_edits"] = edits
|
|
1222
|
+
report.setdefault("result", {})["beat_count"] = len(rows)
|
|
1223
|
+
report["result"]["downbeat_count"] = sum(
|
|
1224
|
+
ccb._parse_bool(row.get("is_downbeat", "0")) for row in rows
|
|
1225
|
+
)
|
|
1226
|
+
_write_json_atomic(report_path, report)
|
|
1227
|
+
_write_beats_atomic(beats_path, rows)
|
|
1228
|
+
return next(
|
|
1229
|
+
item for item in list_beats(file_name, output_dir=output_dir)
|
|
1230
|
+
if abs(item.time_seconds - value) <= ccb.MANUAL_BEAT_COLLISION_TOLERANCE
|
|
1231
|
+
)
|
|
1232
|
+
|
|
1233
|
+
|
|
1234
|
+
def update_beat(
|
|
1235
|
+
file_name: str | Path,
|
|
1236
|
+
beat_id: int,
|
|
1237
|
+
*,
|
|
1238
|
+
time_seconds: float | None = None,
|
|
1239
|
+
is_downbeat: bool | None = None,
|
|
1240
|
+
output_dir: str | Path = "results",
|
|
1241
|
+
note: str | None = None,
|
|
1242
|
+
) -> Beat:
|
|
1243
|
+
"""Adjust one beat and return its new chronological ID."""
|
|
1244
|
+
if not isinstance(beat_id, int) or isinstance(beat_id, bool):
|
|
1245
|
+
raise TypeError("beat_id must be an integer")
|
|
1246
|
+
if is_downbeat is not None and not isinstance(is_downbeat, bool):
|
|
1247
|
+
raise TypeError("is_downbeat must be a boolean or None")
|
|
1248
|
+
if note is not None and not isinstance(note, str):
|
|
1249
|
+
raise TypeError("note must be a string or None")
|
|
1250
|
+
beats_path, report_path, report, rows = _load_beat_state(file_name, output_dir)
|
|
1251
|
+
if not 1 <= beat_id <= len(rows):
|
|
1252
|
+
raise ItemNotFoundError(f"Beat does not exist: {beat_id}")
|
|
1253
|
+
target = rows[beat_id - 1]
|
|
1254
|
+
old_time = float(target["beat_time_seconds"])
|
|
1255
|
+
new_time = _validate_beat_time(
|
|
1256
|
+
old_time if time_seconds is None else time_seconds,
|
|
1257
|
+
float(report["duration_seconds"]),
|
|
1258
|
+
)
|
|
1259
|
+
_validate_beat_collision(new_time, rows, ignored_row=target)
|
|
1260
|
+
new_downbeat = (
|
|
1261
|
+
ccb._parse_bool(target.get("is_downbeat", "0"))
|
|
1262
|
+
if is_downbeat is None else is_downbeat
|
|
1263
|
+
)
|
|
1264
|
+
edits = _manual_edits(report)
|
|
1265
|
+
added = _find_edit(edits["added"], "time_seconds", old_time)
|
|
1266
|
+
adjusted = _find_edit(edits["adjusted"], "new_time_seconds", old_time)
|
|
1267
|
+
if added is not None:
|
|
1268
|
+
added.update(time_seconds=new_time, is_downbeat=new_downbeat)
|
|
1269
|
+
if note is not None:
|
|
1270
|
+
added["note"] = note.strip()
|
|
1271
|
+
reason = "manual_added"
|
|
1272
|
+
elif adjusted is not None:
|
|
1273
|
+
adjusted.update(new_time_seconds=new_time, is_downbeat=new_downbeat)
|
|
1274
|
+
if note is not None:
|
|
1275
|
+
adjusted["note"] = note.strip()
|
|
1276
|
+
reason = "manual_adjusted"
|
|
1277
|
+
else:
|
|
1278
|
+
edits["adjusted"].append(
|
|
1279
|
+
{
|
|
1280
|
+
"original_time_seconds": old_time,
|
|
1281
|
+
"new_time_seconds": new_time,
|
|
1282
|
+
"is_downbeat": new_downbeat,
|
|
1283
|
+
"note": "" if note is None else note.strip(),
|
|
1284
|
+
}
|
|
1285
|
+
)
|
|
1286
|
+
reason = "manual_adjusted"
|
|
1287
|
+
target["beat_time_seconds"] = f"{new_time:.9f}"
|
|
1288
|
+
target["is_downbeat"] = str(int(new_downbeat))
|
|
1289
|
+
target["reliability_class"] = "MANUAL_EDIT"
|
|
1290
|
+
target["reliability_score"] = "1.000000"
|
|
1291
|
+
target["reliability_reason"] = reason
|
|
1292
|
+
report["manual_edits"] = edits
|
|
1293
|
+
report.setdefault("result", {})["downbeat_count"] = sum(
|
|
1294
|
+
ccb._parse_bool(row.get("is_downbeat", "0")) for row in rows
|
|
1295
|
+
)
|
|
1296
|
+
_write_json_atomic(report_path, report)
|
|
1297
|
+
_write_beats_atomic(beats_path, rows)
|
|
1298
|
+
return next(
|
|
1299
|
+
item for item in list_beats(file_name, output_dir=output_dir)
|
|
1300
|
+
if abs(item.time_seconds - new_time) <= ccb.MANUAL_BEAT_COLLISION_TOLERANCE
|
|
1301
|
+
)
|
|
1302
|
+
|
|
1303
|
+
|
|
1304
|
+
def delete_beat(
|
|
1305
|
+
file_name: str | Path,
|
|
1306
|
+
beat_id: int,
|
|
1307
|
+
*,
|
|
1308
|
+
output_dir: str | Path = "results",
|
|
1309
|
+
) -> None:
|
|
1310
|
+
"""Delete one beat and remember the deletion across subsequent runs."""
|
|
1311
|
+
if not isinstance(beat_id, int) or isinstance(beat_id, bool):
|
|
1312
|
+
raise TypeError("beat_id must be an integer")
|
|
1313
|
+
beats_path, report_path, report, rows = _load_beat_state(file_name, output_dir)
|
|
1314
|
+
if not 1 <= beat_id <= len(rows):
|
|
1315
|
+
raise ItemNotFoundError(f"Beat does not exist: {beat_id}")
|
|
1316
|
+
target = rows.pop(beat_id - 1)
|
|
1317
|
+
old_time = float(target["beat_time_seconds"])
|
|
1318
|
+
edits = _manual_edits(report)
|
|
1319
|
+
added = _find_edit(edits["added"], "time_seconds", old_time)
|
|
1320
|
+
adjusted = _find_edit(edits["adjusted"], "new_time_seconds", old_time)
|
|
1321
|
+
if added is not None:
|
|
1322
|
+
edits["added"].remove(added)
|
|
1323
|
+
else:
|
|
1324
|
+
original_time = old_time
|
|
1325
|
+
if adjusted is not None:
|
|
1326
|
+
original_time = float(adjusted["original_time_seconds"])
|
|
1327
|
+
edits["adjusted"].remove(adjusted)
|
|
1328
|
+
if _find_edit(edits["deleted"], "original_time_seconds", original_time) is None:
|
|
1329
|
+
edits["deleted"].append({"original_time_seconds": original_time})
|
|
1330
|
+
report["manual_edits"] = edits
|
|
1331
|
+
report.setdefault("result", {})["beat_count"] = len(rows)
|
|
1332
|
+
report["result"]["downbeat_count"] = sum(
|
|
1333
|
+
ccb._parse_bool(row.get("is_downbeat", "0")) for row in rows
|
|
1334
|
+
)
|
|
1335
|
+
_write_json_atomic(report_path, report)
|
|
1336
|
+
_write_beats_atomic(beats_path, rows)
|
|
1337
|
+
|
|
1338
|
+
|
|
1339
|
+
def reset_beat_edits(
|
|
1340
|
+
file_name: str | Path,
|
|
1341
|
+
*,
|
|
1342
|
+
output_dir: str | Path = "results",
|
|
1343
|
+
) -> int:
|
|
1344
|
+
"""Forget all manual beat operations; call run() to restore the automatic grid."""
|
|
1345
|
+
_, report_path, report, _ = _load_beat_state(file_name, output_dir)
|
|
1346
|
+
edits = _manual_edits(report)
|
|
1347
|
+
count = sum(len(items) for items in edits.values())
|
|
1348
|
+
report["manual_edits"] = ccb.empty_manual_beat_edits()
|
|
1349
|
+
_write_json_atomic(report_path, report)
|
|
1350
|
+
return count
|
|
1351
|
+
|
|
1352
|
+
|
|
1353
|
+
def run(
|
|
1354
|
+
file_name: str | Path,
|
|
1355
|
+
output_dir: str | Path = "results",
|
|
1356
|
+
*,
|
|
1357
|
+
cache_dir: str | Path | None = None,
|
|
1358
|
+
no_click: bool = False,
|
|
1359
|
+
refresh_cache: bool = False,
|
|
1360
|
+
device: str = "cpu",
|
|
1361
|
+
checkpoint: str = "final0",
|
|
1362
|
+
hop_seconds: float = 10.0,
|
|
1363
|
+
sample_rate: int = 22050,
|
|
1364
|
+
min_bpm: float = 40.0,
|
|
1365
|
+
max_bpm: float = 240.0,
|
|
1366
|
+
change_ratio: float = 0.08,
|
|
1367
|
+
change_bpm: float = 8.0,
|
|
1368
|
+
min_change_beats: int = 4,
|
|
1369
|
+
min_change_seconds: float = 2.0,
|
|
1370
|
+
preserve_manual_edits: bool = True,
|
|
1371
|
+
) -> RunResult:
|
|
1372
|
+
"""Run CCB for one audio file and return its final artifacts.
|
|
1373
|
+
|
|
1374
|
+
Beat This! is initialized only when the hidden inference cache is absent or
|
|
1375
|
+
``refresh_cache`` is true. Unlike the CLI, this function does not rewrite
|
|
1376
|
+
the multi-file ``summary.csv`` because one API call represents one song.
|
|
1377
|
+
Manual beat operations are reapplied by default; pass
|
|
1378
|
+
``preserve_manual_edits=False`` to discard them and rebuild the automatic
|
|
1379
|
+
grid.
|
|
1380
|
+
"""
|
|
1381
|
+
audio_path = Path(file_name).expanduser()
|
|
1382
|
+
if not audio_path.is_file():
|
|
1383
|
+
raise ResourceNotFoundError(f"Audio file does not exist: {audio_path}")
|
|
1384
|
+
if not 0 < min_bpm < max_bpm:
|
|
1385
|
+
raise InvalidArgumentError("Require 0 < min_bpm < max_bpm")
|
|
1386
|
+
if not 0 < hop_seconds < 30:
|
|
1387
|
+
raise InvalidArgumentError("Require 0 < hop_seconds < 30")
|
|
1388
|
+
if sample_rate <= 0:
|
|
1389
|
+
raise InvalidArgumentError("sample_rate must be positive")
|
|
1390
|
+
if min_change_beats < 1:
|
|
1391
|
+
raise InvalidArgumentError("min_change_beats must be at least 1")
|
|
1392
|
+
if min_change_seconds < 0:
|
|
1393
|
+
raise InvalidArgumentError("min_change_seconds must be non-negative")
|
|
1394
|
+
if not isinstance(preserve_manual_edits, bool):
|
|
1395
|
+
raise TypeError("preserve_manual_edits must be a boolean")
|
|
1396
|
+
|
|
1397
|
+
output_root = Path(output_dir).expanduser()
|
|
1398
|
+
output_root.mkdir(parents=True, exist_ok=True)
|
|
1399
|
+
resolved_cache_dir = (
|
|
1400
|
+
None if cache_dir is None else Path(cache_dir).expanduser()
|
|
1401
|
+
)
|
|
1402
|
+
music_gain, click_gain = ccb.validate_mix_gains(_music_gain, _click_gain)
|
|
1403
|
+
args = argparse.Namespace(
|
|
1404
|
+
audio=[audio_path],
|
|
1405
|
+
output=output_root,
|
|
1406
|
+
sample_rate=sample_rate,
|
|
1407
|
+
min_bpm=min_bpm,
|
|
1408
|
+
max_bpm=max_bpm,
|
|
1409
|
+
change_ratio=change_ratio,
|
|
1410
|
+
change_bpm=change_bpm,
|
|
1411
|
+
min_change_beats=min_change_beats,
|
|
1412
|
+
min_change_seconds=min_change_seconds,
|
|
1413
|
+
beat_this_device=device,
|
|
1414
|
+
beat_this_checkpoint=checkpoint,
|
|
1415
|
+
beat_this_hop_seconds=hop_seconds,
|
|
1416
|
+
refresh_cache=refresh_cache,
|
|
1417
|
+
no_click=no_click,
|
|
1418
|
+
music_gain=music_gain,
|
|
1419
|
+
click_gain=click_gain,
|
|
1420
|
+
preserve_manual_edits=preserve_manual_edits,
|
|
1421
|
+
cache_dir=resolved_cache_dir,
|
|
1422
|
+
)
|
|
1423
|
+
|
|
1424
|
+
estimator: ccb.BeatThisEstimator | None = None
|
|
1425
|
+
if refresh_cache or not ccb.inference_cache_available(
|
|
1426
|
+
audio_path,
|
|
1427
|
+
output_root,
|
|
1428
|
+
resolved_cache_dir,
|
|
1429
|
+
sample_rate=sample_rate,
|
|
1430
|
+
checkpoint=checkpoint,
|
|
1431
|
+
hop_seconds=hop_seconds,
|
|
1432
|
+
):
|
|
1433
|
+
estimator = ccb.BeatThisEstimator(
|
|
1434
|
+
device,
|
|
1435
|
+
checkpoint,
|
|
1436
|
+
hop_seconds=hop_seconds,
|
|
1437
|
+
)
|
|
1438
|
+
ccb.analyse_file(audio_path, output_root, args, estimator)
|
|
1439
|
+
|
|
1440
|
+
result_dir = ccb.result_directory(audio_path, output_root)
|
|
1441
|
+
report_path = result_dir / "report.json"
|
|
1442
|
+
with report_path.open(encoding="utf-8") as handle:
|
|
1443
|
+
report = json.load(handle)
|
|
1444
|
+
return RunResult(
|
|
1445
|
+
audio=audio_path.resolve(),
|
|
1446
|
+
output_directory=result_dir.resolve(),
|
|
1447
|
+
beats_csv=(result_dir / "beats.csv").resolve(),
|
|
1448
|
+
click_wav=(None if no_click else (result_dir / "click.wav").resolve()),
|
|
1449
|
+
overview_png=(result_dir / "overview.png").resolve(),
|
|
1450
|
+
segments_csv=(result_dir / "segments.csv").resolve(),
|
|
1451
|
+
report_json=report_path.resolve(),
|
|
1452
|
+
report=report,
|
|
1453
|
+
)
|
|
1454
|
+
|
|
1455
|
+
|
|
1456
|
+
__all__ = [
|
|
1457
|
+
"API_VERSION",
|
|
1458
|
+
"__version__",
|
|
1459
|
+
"CCBError",
|
|
1460
|
+
"InvalidArgumentError",
|
|
1461
|
+
"ResourceNotFoundError",
|
|
1462
|
+
"ResultStateError",
|
|
1463
|
+
"EditConflictError",
|
|
1464
|
+
"ItemNotFoundError",
|
|
1465
|
+
"RunResult",
|
|
1466
|
+
"NoBeatRange",
|
|
1467
|
+
"Beat",
|
|
1468
|
+
"CacheEntry",
|
|
1469
|
+
"CacheCleanupResult",
|
|
1470
|
+
"ManualBeatAddition",
|
|
1471
|
+
"ManualBeatAdjustment",
|
|
1472
|
+
"ManualBeatDeletion",
|
|
1473
|
+
"ManualBeatEdits",
|
|
1474
|
+
"ReviewRange",
|
|
1475
|
+
"SongInfo",
|
|
1476
|
+
"ValidationResult",
|
|
1477
|
+
"run",
|
|
1478
|
+
"set_music_gain",
|
|
1479
|
+
"set_click_gain",
|
|
1480
|
+
"list_no_beat_ranges",
|
|
1481
|
+
"create_no_beat_range",
|
|
1482
|
+
"update_no_beat_range",
|
|
1483
|
+
"delete_no_beat_range",
|
|
1484
|
+
"clear_no_beat_ranges",
|
|
1485
|
+
"list_beats",
|
|
1486
|
+
"get_result",
|
|
1487
|
+
"inspect_song",
|
|
1488
|
+
"get_manual_beat_edits",
|
|
1489
|
+
"get_review_ranges",
|
|
1490
|
+
"validate_result",
|
|
1491
|
+
"create_beat",
|
|
1492
|
+
"update_beat",
|
|
1493
|
+
"delete_beat",
|
|
1494
|
+
"reset_beat_edits",
|
|
1495
|
+
"list_caches",
|
|
1496
|
+
"prune_caches",
|
|
1497
|
+
]
|