diff --git a/nodes/audio/FL_Audio_Beat_Prompt_Schedule.py b/nodes/audio/FL_Audio_Beat_Prompt_Schedule.py index 846f2ca..c8e5f2b 100644 --- a/nodes/audio/FL_Audio_Beat_Prompt_Schedule.py +++ b/nodes/audio/FL_Audio_Beat_Prompt_Schedule.py @@ -857,6 +857,9 @@ class FL_Audio_Beat_Prompt_Schedule(io.ComfyNode): if internal_analysis is not None: ui_payload["audio_file"] = internal_analysis.get("audio_file", audio_file) ui_payload["cache_key"] = internal_analysis.get("cache_key", "") + source_analysis = internal_analysis.get("source_analysis") + if isinstance(source_analysis, dict) and "beat_times" in source_analysis: + ui_payload["source_analysis"] = source_analysis waveform = ( internal_analysis["waveform_preview"] if internal_analysis is not None diff --git a/nodes/audio/audio_timeline.py b/nodes/audio/audio_timeline.py index aa80f7f..2ecb8ed 100644 --- a/nodes/audio/audio_timeline.py +++ b/nodes/audio/audio_timeline.py @@ -14,8 +14,9 @@ from .audio_separation import load_cached_stem from .beat_this_detector import MODEL_FPS, MODEL_SHA256, analyze_beats as analyze_beat_this -ANALYSIS_VERSION = 6 +ANALYSIS_VERSION = 7 DETECTOR_VERSION = f"beat-this-final0-{MODEL_SHA256[:12]}" +SOURCE_ANALYSIS_VERSION = 1 _WAVEFORM_BUCKETS_PER_SECOND = 60 _MAX_WAVEFORM_BUCKETS = 8192 _WAVEFORM_SCALE = 32767 @@ -114,6 +115,27 @@ def waveform_preview(waveform, sample_rate): } +def crop_waveform_preview(preview, start, duration): + if not isinstance(preview, dict) or preview.get("version") != 1 or duration <= 0: + return None + peaks = preview.get("peaks") + source_duration = preview.get("duration") + if not isinstance(peaks, list) or len(peaks) < 2 or len(peaks) % 2 or not source_duration: + return None + bucket_count = len(peaks) // 2 + start_bucket = min(bucket_count - 1, max(0, math.floor(start / source_duration * bucket_count))) + end_bucket = min( + bucket_count, + max(start_bucket + 1, math.ceil((start + duration) / source_duration * bucket_count)), + ) + return { + "version": 1, + "duration": float(duration), + "scale": preview["scale"], + "peaks": peaks[start_bucket * 2:end_bucket * 2], + } + + def _detect_drums(waveform, sample_rate, onset_frames, onset_times): stft_mag = np.abs(librosa.stft(waveform)) frequencies = librosa.fft_frequencies(sr=sample_rate) @@ -275,15 +297,7 @@ def apply_beat_offset(analysis, fps, beat_offset_ms=0, beat_grid_density="every_ return result -def analyze_audio( - audio, - fps, - half_time=False, - beat_offset_ms=0, - beat_grid_density="every_beat", - detect_beats=True, - beat_audio=None, -): +def analyze_audio(audio, detect_beats=True, beat_audio=None): waveform = mono_numpy(audio) sample_rate = int(audio["sample_rate"]) duration = len(waveform) / sample_rate @@ -302,7 +316,6 @@ def analyze_audio( "version": ANALYSIS_VERSION, "detector_version": DETECTOR_VERSION, "onset_times": onsets, - "onset_frames": [round(value * fps) for value in onsets], "sample_rate": sample_rate, "audio_duration": float(duration), "drum_times": drum_times, @@ -320,20 +333,8 @@ def analyze_audio( downbeat_confidences = np.asarray( detected["downbeat_confidences"], dtype=np.float64 ) - if half_time: - retained = np.arange(0, len(beat_times), 2) - retained_indices = set(retained.tolist()) - downbeat_retained = [] - for index, downbeat in enumerate(downbeat_times): - nearest = int(np.argmin(np.abs(beat_times - downbeat))) - if nearest in retained_indices: - downbeat_retained.append(index) - beat_times = beat_times[retained] - beat_confidences = beat_confidences[retained] - downbeat_times = downbeat_times[downbeat_retained] - downbeat_confidences = downbeat_confidences[downbeat_retained] if len(beat_times) < 2: - raise ValueError("Beat This must detect at least two beats in the selected audio range.") + raise ValueError("Beat This must detect at least two beats in the source audio.") interval = float(np.median(np.diff(beat_times))) analysis.update({ @@ -348,32 +349,141 @@ def analyze_audio( "detected_beat_confidences": beat_confidences.tolist(), "base_detected_downbeat_confidences": downbeat_confidences.tolist(), "detected_downbeat_confidences": downbeat_confidences.tolist(), - "beat_frames": [round(value * fps) for value in beat_times], - "downbeat_frames": [round(value * fps) for value in downbeat_times], - "detected_beat_frames": [round(value * fps) for value in beat_times], "num_beats": len(beat_times), "num_downbeats": len(downbeat_times), "detector": detected["detector"], }) - return apply_beat_offset(analysis, fps, beat_offset_ms, beat_grid_density) + return analysis + + +def apply_half_time(analysis, half_time): + result = dict(analysis) + if not half_time or "beat_times" not in result: + return result + + beat_times = list(result["beat_times"]) + beat_confidences = list(result.get("detected_beat_confidences", [])) + retained = list(range(0, len(beat_times), 2)) + if len(retained) < 2: + raise ValueError("Beat This must detect at least two beats after half-time filtering.") + retained_indices = set(retained) + downbeat_times = list(result.get("downbeat_times", [])) + downbeat_confidences = list(result.get("detected_downbeat_confidences", [])) + retained_downbeats = [] + retained_downbeat_confidences = [] + for index, downbeat in enumerate(downbeat_times): + nearest = min(range(len(beat_times)), key=lambda position: abs(beat_times[position] - downbeat)) + if nearest in retained_indices: + retained_downbeats.append(downbeat) + if index < len(downbeat_confidences): + retained_downbeat_confidences.append(downbeat_confidences[index]) + + retained_beats = [beat_times[index] for index in retained] + retained_beat_confidences = [ + beat_confidences[index] for index in retained if index < len(beat_confidences) + ] + interval = float(np.median(np.diff(retained_beats))) + result.update({ + "bpm": 60.0 / interval, + "base_grid_interval_seconds": interval, + "beat_times": retained_beats, + "downbeat_times": retained_downbeats, + "detected_beat_times": retained_beats, + "detected_downbeat_times": retained_downbeats, + "base_detected_beat_confidences": retained_beat_confidences, + "detected_beat_confidences": retained_beat_confidences, + "base_detected_downbeat_confidences": retained_downbeat_confidences, + "detected_downbeat_confidences": retained_downbeat_confidences, + "num_beats": len(retained_beats), + "num_downbeats": len(retained_downbeats), + }) + return result + + +def _crop_times(values, start, end): + return [float(value - start) for value in values if start <= value < end] + + +def _crop_times_with_values(times, values, start, end): + cropped_times = [] + cropped_values = [] + for index, value in enumerate(times): + if start <= value < end: + cropped_times.append(float(value - start)) + if index < len(values): + cropped_values.append(values[index]) + return cropped_times, cropped_values + + +def project_analysis(analysis, crop, fps): + start = float(crop["source_start"]) + duration = float(crop["audio_duration"]) + end = start + duration + result = dict(analysis) + + for time_key, confidence_key in ( + ("base_detected_beat_times", "base_detected_beat_confidences"), + ("detected_beat_times", "detected_beat_confidences"), + ("base_detected_downbeat_times", "base_detected_downbeat_confidences"), + ("detected_downbeat_times", "detected_downbeat_confidences"), + ): + times, confidences = _crop_times_with_values( + analysis.get(time_key, []), + analysis.get(confidence_key, []), + start, + end, + ) + result[time_key] = times + result[confidence_key] = confidences + + for key in ( + "base_beat_times", + "beat_times", + "base_downbeat_times", + "downbeat_times", + "onset_times", + ): + if key in analysis: + result[key] = _crop_times(analysis[key], start, end) + + drums = dict(analysis.get("drum_times", {})) + for key in ("kick_times", "snare_times", "hihat_times"): + drums[key] = _crop_times(drums.get(key, []), start, end) + drums.update({ + "duration": duration, + "total_kicks": len(drums["kick_times"]), + "total_snares": len(drums["snare_times"]), + "total_hihats": len(drums["hihat_times"]), + }) + + result.update(crop) + result.update({ + "beat_frames": [round(value * fps) for value in result.get("beat_times", [])], + "downbeat_frames": [round(value * fps) for value in result.get("downbeat_times", [])], + "detected_beat_frames": [ + round(value * fps) for value in result.get("detected_beat_times", []) + ], + "onset_frames": [round(value * fps) for value in result.get("onset_times", [])], + "drum_times": drums, + "num_beats": len(result.get("beat_times", [])), + "num_downbeats": len(result.get("downbeat_times", [])), + "waveform_preview": crop_waveform_preview( + analysis.get("waveform_preview"), + start, + duration, + ), + }) + return result def analysis_cache_key( path, - fps, - trim_start_frame, - length_frames, - half_time, analysis_source, detect_beats=True, ): values = { "analysis_version": ANALYSIS_VERSION, "audio_sha256": audio_file_hash(path), - "fps": float(fps), - "trim_start_frame": int(trim_start_frame), - "length_frames": int(length_frames), - "half_time": bool(half_time), "analysis_source": analysis_source, "detect_beats": bool(detect_beats), "detector_version": DETECTOR_VERSION, @@ -426,50 +536,57 @@ def analyze_audio_file( path = resolve_audio_path(filename) _, master_audio = load_audio_file(filename) cropped_audio, crop = crop_audio(master_audio, fps, trim_start_frame, length_frames) - if analysis_source == "mix": - analysis_audio = cropped_audio - else: - stem_audio = load_cached_stem(filename, analysis_source) - analysis_audio, _ = crop_audio(stem_audio, fps, trim_start_frame, length_frames) cache_key = analysis_cache_key( path, - fps, - trim_start_frame, - length_frames, - half_time, analysis_source, detect_beats, ) cache_path = _cache_path(cache_key) cache_hit = cache_path.is_file() if cache_hit: - analysis = json.loads(cache_path.read_text(encoding="utf-8")) + source_analysis = json.loads(cache_path.read_text(encoding="utf-8")) else: - analysis = analyze_audio( - analysis_audio, - fps, - half_time, - detect_beats=detect_beats, - beat_audio=cropped_audio if analysis_source != "mix" and detect_beats else None, + analysis_audio = ( + master_audio + if analysis_source == "mix" + else load_cached_stem(filename, analysis_source) ) - analysis.update(crop) - analysis.update({ + source_analysis = analyze_audio( + analysis_audio, + detect_beats=detect_beats, + beat_audio=master_audio if analysis_source != "mix" and detect_beats else None, + ) + source_analysis.update({ "audio_file": filename, "analysis_source": analysis_source, "beat_analysis_source": "mix" if detect_beats else None, "cache_key": cache_key, + "audio_duration": crop["source_duration"], + "source_duration": crop["source_duration"], + "source_start": 0.0, }) temporary_path = cache_path.with_suffix(".tmp") - temporary_path.write_text(json.dumps(analysis, separators=(",", ":")), encoding="utf-8") + temporary_path.write_text( + json.dumps(source_analysis, separators=(",", ":")), + encoding="utf-8", + ) temporary_path.replace(cache_path) if detect_beats: analysis = apply_beat_offset( - analysis, + apply_half_time(source_analysis, half_time), fps, beat_offset_ms, beat_grid_density, ) else: - analysis = dict(analysis) + analysis = dict(source_analysis) + analysis = project_analysis(analysis, crop, fps) + analysis["source_analysis"] = { + **source_analysis, + "type": "fl_audio_source_analysis", + "version": SOURCE_ANALYSIS_VERSION, + "analysis_version": ANALYSIS_VERSION, + "analysis_cache_hit": cache_hit, + } analysis["analysis_cache_hit"] = cache_hit return analysis, cropped_audio diff --git a/tests/test_audio_beat_prompt_schedule.py b/tests/test_audio_beat_prompt_schedule.py index e58aa2e..2dd83c8 100644 --- a/tests/test_audio_beat_prompt_schedule.py +++ b/tests/test_audio_beat_prompt_schedule.py @@ -498,6 +498,11 @@ class BeatPromptScheduleTests(unittest.TestCase): "hihat_times": [], "duration": 1.0, }, + "source_analysis": { + "type": "fl_audio_source_analysis", + "version": 1, + "beat_times": [0.5, 1.0, 1.5], + }, } with mock.patch.object(schedule, "analyze_audio_file", return_value=(analysis, audio)): output = schedule.FL_Audio_Beat_Prompt_Schedule.execute( @@ -524,6 +529,8 @@ class BeatPromptScheduleTests(unittest.TestCase): self.assertEqual(payload["beat_analysis_source"], "mix") self.assertEqual(payload["drum_times"]["kick_times"], [0.05]) self.assertEqual(payload["source_start"], 0.5) + self.assertIs(payload["source_analysis"], analysis["source_analysis"]) + self.assertNotIn("source_analysis", output.result[0]) def test_cached_analysis_restores_a_missing_audio_widget_value(self): audio = {"waveform": torch.zeros(1, 1, 24000), "sample_rate": 24000} diff --git a/tests/test_audio_timeline.py b/tests/test_audio_timeline.py index 883ca69..9632e72 100644 --- a/tests/test_audio_timeline.py +++ b/tests/test_audio_timeline.py @@ -42,7 +42,7 @@ class AudioTimelineTests(unittest.TestCase): mock.patch.object(timeline, "_detect_drums", return_value={}), mock.patch.object(timeline, "waveform_preview", return_value={}), ): - analysis = timeline.analyze_audio(audio, fps=24.0) + analysis = timeline.analyze_audio(audio) self.assertAlmostEqual(analysis["bpm"], 120.0) self.assertEqual(analysis["bpm_source"], "beat_this_intervals_median") @@ -64,7 +64,7 @@ class AudioTimelineTests(unittest.TestCase): mock.patch.object(timeline, "_detect_drums", return_value={}), mock.patch.object(timeline, "waveform_preview", return_value={}), ): - analysis = timeline.analyze_audio(audio, fps=24.0, detect_beats=False) + analysis = timeline.analyze_audio(audio, detect_beats=False) detector.assert_not_called() self.assertNotIn("beat_times", analysis) @@ -88,7 +88,7 @@ class AudioTimelineTests(unittest.TestCase): mock.patch.object(timeline, "_detect_drums", return_value={}), mock.patch.object(timeline, "waveform_preview", return_value={}), ): - timeline.analyze_audio(stem, fps=24.0, beat_audio=mix) + timeline.analyze_audio(stem, beat_audio=mix) self.assertEqual(beat_this.call_args.args[1], 200) timeline.np.testing.assert_array_equal(beat_this.call_args.args[0], timeline.np.ones(400)) @@ -284,7 +284,7 @@ class AudioTimelineTests(unittest.TestCase): with ( mock.patch.object(timeline, "resolve_audio_path", return_value=pathlib.Path("song.wav")), mock.patch.object(timeline, "load_audio_file", return_value=(pathlib.Path("song.wav"), master)), - mock.patch.object(timeline, "load_cached_stem", return_value=stem), + mock.patch.object(timeline, "load_cached_stem", return_value=stem) as load_stem, mock.patch.object(timeline, "analysis_cache_key", return_value="key"), mock.patch.object(timeline, "_cache_path", return_value=cache_path), mock.patch.object(timeline, "analyze_audio", return_value=analysis) as analyze, @@ -295,10 +295,20 @@ class AudioTimelineTests(unittest.TestCase): length_frames=24, analysis_source="drums", ) + _, second_crop = timeline.analyze_audio_file( + "song.wav", + fps=24.0, + trim_start_frame=12, + length_frames=12, + analysis_source="drums", + ) self.assertEqual(cropped["waveform"].mean(), 1.0) + self.assertEqual(second_crop["waveform"].mean(), 1.0) self.assertEqual(analyze.call_args.args[0]["waveform"].mean(), 2.0) self.assertEqual(analyze.call_args.kwargs["beat_audio"]["waveform"].mean(), 1.0) + self.assertEqual(analyze.call_count, 1) + self.assertEqual(load_stem.call_count, 1) def test_offset_changes_reuse_the_base_analysis_cache(self): master = {"waveform": torch.ones(1, 1, 48000), "sample_rate": 48000} @@ -366,14 +376,151 @@ class AudioTimelineTests(unittest.TestCase): self.assertEqual(every_two["beat_times"], [0.0, 0.5]) self.assertEqual(every_beat["cache_key"], every_two["cache_key"]) + def test_different_crops_reuse_full_source_analysis_and_rebase_events(self): + master = {"waveform": torch.ones(1, 1, 96000), "sample_rate": 24000} + source_analysis = { + "bpm": 60.0, + "base_grid_interval_seconds": 1.0, + "beat_times": [0.25, 1.25, 2.25, 3.25], + "downbeat_times": [0.25, 2.25], + "detected_beat_times": [0.25, 1.25, 2.25, 3.25], + "detected_downbeat_times": [0.25, 2.25], + "detected_beat_confidences": [0.9, 0.8, 0.7, 0.6], + "detected_downbeat_confidences": [0.95, 0.75], + "onset_times": [0.5, 1.5, 2.5, 3.5], + "audio_duration": 4.0, + "waveform_preview": { + "version": 1, + "duration": 4.0, + "scale": 32767, + "peaks": list(range(16)), + }, + "drum_times": { + "kick_times": [1.25, 3.25], + "snare_times": [2.25], + "hihat_times": [], + }, + } + with tempfile.TemporaryDirectory() as directory: + cache_path = pathlib.Path(directory) / "analysis.json" + with ( + mock.patch.object(timeline, "resolve_audio_path", return_value=pathlib.Path("song.wav")), + mock.patch.object(timeline, "load_audio_file", return_value=(pathlib.Path("song.wav"), master)), + mock.patch.object(timeline, "analysis_cache_key", return_value="key"), + mock.patch.object(timeline, "_cache_path", return_value=cache_path), + mock.patch.object(timeline, "analyze_audio", return_value=source_analysis) as analyze, + ): + first, first_audio = timeline.analyze_audio_file( + "song.wav", fps=24.0, trim_start_frame=24, length_frames=48 + ) + second, second_audio = timeline.analyze_audio_file( + "song.wav", fps=24.0, trim_start_frame=48, length_frames=48 + ) + + self.assertEqual(analyze.call_count, 1) + self.assertEqual(analyze.call_args.args[0]["waveform"].shape[-1], 96000) + self.assertEqual(first_audio["waveform"].shape[-1], 48000) + self.assertEqual(second_audio["waveform"].shape[-1], 48000) + self.assertEqual(first["beat_times"], [0.25, 1.25]) + self.assertEqual(second["beat_times"], [0.25, 1.25]) + self.assertEqual(first["detected_beat_confidences"], [0.8, 0.7]) + self.assertEqual(first["onset_times"], [0.5, 1.5]) + self.assertEqual(first["drum_times"]["kick_times"], [0.25]) + self.assertEqual(first["source_start"], 1.0) + self.assertEqual(second["source_start"], 2.0) + self.assertEqual(first["source_analysis"]["beat_times"], [0.25, 1.25, 2.25, 3.25]) + self.assertFalse(first["analysis_cache_hit"]) + self.assertTrue(second["analysis_cache_hit"]) + + def test_fps_half_time_and_density_project_from_one_source_cache(self): + master = {"waveform": torch.ones(1, 1, 96000), "sample_rate": 24000} + source_analysis = { + "bpm": 120.0, + "base_grid_interval_seconds": 0.5, + "beat_times": [0.0, 0.5, 1.0, 1.5, 2.0, 2.5, 3.0, 3.5], + "downbeat_times": [0.0, 2.0], + "detected_beat_times": [0.0, 0.5, 1.0, 1.5, 2.0, 2.5, 3.0, 3.5], + "detected_downbeat_times": [0.0, 2.0], + "detected_beat_confidences": [0.9] * 8, + "detected_downbeat_confidences": [0.95, 0.85], + "onset_times": [], + "audio_duration": 4.0, + "waveform_preview": None, + "drum_times": {}, + } + with tempfile.TemporaryDirectory() as directory: + cache_path = pathlib.Path(directory) / "analysis.json" + with ( + mock.patch.object(timeline, "resolve_audio_path", return_value=pathlib.Path("song.wav")), + mock.patch.object(timeline, "load_audio_file", return_value=(pathlib.Path("song.wav"), master)), + mock.patch.object(timeline, "analysis_cache_key", return_value="key"), + mock.patch.object(timeline, "_cache_path", return_value=cache_path), + mock.patch.object(timeline, "analyze_audio", return_value=source_analysis) as analyze, + ): + regular, _ = timeline.analyze_audio_file( + "song.wav", fps=24.0, trim_start_frame=12, length_frames=48 + ) + half_time, _ = timeline.analyze_audio_file( + "song.wav", + fps=48.0, + trim_start_frame=24, + length_frames=96, + half_time=True, + ) + every_two, _ = timeline.analyze_audio_file( + "song.wav", + fps=24.0, + trim_start_frame=12, + length_frames=48, + beat_grid_density="every_2_beats", + ) + + self.assertEqual(analyze.call_count, 1) + self.assertEqual(regular["beat_times"], [0.0, 0.5, 1.0, 1.5]) + self.assertEqual(half_time["beat_times"], [0.5, 1.5]) + self.assertEqual(every_two["beat_times"], [0.5, 1.5]) + self.assertEqual(regular["cache_key"], half_time["cache_key"]) + self.assertEqual(regular["cache_key"], every_two["cache_key"]) + + def test_projection_uses_half_open_boundaries_and_keeps_confidence_pairs(self): + analysis = { + "beat_times": [1.0, 1.5, 2.0], + "downbeat_times": [1.0, 2.0], + "base_beat_times": [1.0, 1.5, 2.0], + "base_downbeat_times": [1.0, 2.0], + "detected_beat_times": [1.0, 1.5, 2.0], + "base_detected_beat_times": [1.0, 1.5, 2.0], + "detected_beat_confidences": [0.9, 0.8, 0.7], + "base_detected_beat_confidences": [0.9, 0.8, 0.7], + "detected_downbeat_times": [1.0, 2.0], + "base_detected_downbeat_times": [1.0, 2.0], + "detected_downbeat_confidences": [0.95, 0.75], + "base_detected_downbeat_confidences": [0.95, 0.75], + "onset_times": [1.0, 2.0], + "drum_times": {"kick_times": [1.0, 2.0]}, + "waveform_preview": None, + } + crop = { + "source_duration": 3.0, + "source_start": 1.0, + "trim_start_frame": 24, + "length_frames": 24, + "audio_duration": 1.0, + "sample_rate": 48000, + } + + projected = timeline.project_analysis(analysis, crop, fps=24.0) + + self.assertEqual(projected["beat_times"], [0.0, 0.5]) + self.assertEqual(projected["detected_beat_confidences"], [0.9, 0.8]) + self.assertEqual(projected["downbeat_times"], [0.0]) + self.assertEqual(projected["onset_times"], [0.0]) + self.assertEqual(projected["drum_times"]["kick_times"], [0.0]) + def test_beat_detection_mode_has_a_separate_versioned_cache_key(self): with mock.patch.object(timeline, "audio_file_hash", return_value="audio"): - beats = timeline.analysis_cache_key( - pathlib.Path("song.wav"), 24.0, 0, 0, False, "mix", True - ) - auxiliary = timeline.analysis_cache_key( - pathlib.Path("song.wav"), 24.0, 0, 0, False, "mix", False - ) + beats = timeline.analysis_cache_key(pathlib.Path("song.wav"), "mix", True) + auxiliary = timeline.analysis_cache_key(pathlib.Path("song.wav"), "mix", False) self.assertNotEqual(beats, auxiliary) diff --git a/tests/test_audio_timeline_coordinates.mjs b/tests/test_audio_timeline_coordinates.mjs new file mode 100644 index 0000000..e645186 --- /dev/null +++ b/tests/test_audio_timeline_coordinates.mjs @@ -0,0 +1,61 @@ +import assert from "node:assert/strict"; +import { readFile } from "node:fs/promises"; +import test from "node:test"; + +const moduleSource = await readFile( + new URL("../web/nodes/audio/audio_timeline_coordinates.js", import.meta.url), + "utf8", +); +const { + cropTimes, + cropTimesWithValues, + sourceTimes, + sourceTimeAtLocalFrame, + sourceTimeToLocalFrame, + waveformBinRange, +} = await import(`data:text/javascript;base64,${Buffer.from(moduleSource).toString("base64")}`); + +test("crop projection uses half-open bounds and preserves paired values", () => { + assert.deepEqual(cropTimes([1, 1.5, 2], 1, 2), [0, 0.5]); + assert.deepEqual( + cropTimesWithValues([1, 1.5, 2], [0.9, 0.8, 0.7], 1, 2), + [[0, 0.5], [0.9, 0.8]], + ); +}); + +test("waveform and analyzed beats share the same crop origin", () => { + const fps = 24; + const sourceTime = 10.25; + const firstCrop = 2; + const secondCrop = 3.5; + const firstFrame = sourceTimeToLocalFrame(sourceTime, firstCrop, fps); + const secondFrame = sourceTimeToLocalFrame(sourceTime, secondCrop, fps); + + assert.equal(firstFrame - secondFrame, (secondCrop - firstCrop) * fps); + assert.equal(sourceTimeAtLocalFrame(firstFrame, firstCrop, fps), sourceTime); + assert.equal(sourceTimeAtLocalFrame(secondFrame, secondCrop, fps), sourceTime); + + const firstBins = waveformBinRange(firstFrame, firstFrame + 1, firstCrop, fps, 0, 20, 2000); + const secondBins = waveformBinRange(secondFrame, secondFrame + 1, secondCrop, fps, 0, 20, 2000); + assert.deepEqual(firstBins, secondBins); +}); + +test("beat offset moves only the working grid", () => { + const fps = 24; + const cropStart = 2; + const detectedSourceTime = 10.25; + const detectedFrame = sourceTimeToLocalFrame(detectedSourceTime, cropStart, fps); + const waveformSourceTime = sourceTimeAtLocalFrame(detectedFrame, cropStart, fps); + const shiftedGridFrame = sourceTimeToLocalFrame(detectedSourceTime + 0.125, cropStart, fps); + + assert.equal(waveformSourceTime, detectedSourceTime); + assert.equal(shiftedGridFrame - detectedFrame, 3); +}); + +test("legacy crop-local markers are promoted without destructive reprojection", () => { + const source = sourceTimes([0.25, 1.25, 2.25], 4); + + assert.deepEqual(cropTimes(source, 4, 6), [0.25, 1.25]); + assert.deepEqual(cropTimes(source, 5, 7), [0.25, 1.25]); + assert.deepEqual(cropTimes(source, 4, 6), [0.25, 1.25]); +}); diff --git a/web/nodes/audio/FL_Audio_Beat_Prompt_Schedule.js b/web/nodes/audio/FL_Audio_Beat_Prompt_Schedule.js index b683c90..d105a8a 100644 --- a/web/nodes/audio/FL_Audio_Beat_Prompt_Schedule.js +++ b/web/nodes/audio/FL_Audio_Beat_Prompt_Schedule.js @@ -1,13 +1,21 @@ import { app } from "../../../../scripts/app.js"; import { api } from "../../../../scripts/api.js"; +import { + cropTimes, + cropTimesWithValues, + sourceTimes, + sourceTimeToLocalFrame, + waveformBinRange, +} from "./audio_timeline_coordinates.js"; const STYLE_ID = "fl-beat-prompt-sequencer-styles"; const INSTANCES = new Map(); const HEADER_RE = /^\s*\[\s*([0-9]+(?:\.[0-9]+)?)\s*-\s*([0-9]+(?:\.[0-9]+)?)(?:\s*\|\s*(.*?))?\s*\]\s*$/; const HEADER_START_RE = /^\s*\[\s*[0-9]+(?:\.[0-9]+)?\s*-/; const EPSILON = 1e-6; -const FORMAT_VERSION = 13; -const COMPATIBLE_FORMAT_VERSIONS = new Set([6, 7, 8, 9, 10, 11, 12, FORMAT_VERSION]); +const FORMAT_VERSION = 15; +const COMPATIBLE_FORMAT_VERSIONS = new Set([6, 7, 8, 9, 10, 11, 12, 13, 14, FORMAT_VERSION]); +const SOURCE_ANALYSIS_VERSION = 1; const LEGACY_BPM_METHODS = new Set(["beat_intervals", "onset_strength"]); const COMPACT_NODE_WIDTH = 380; const MEDIA_FILE_RE = /\.(?:aac|aiff?|flac|m4a|mka|mkv|mov|mp3|mp4|oga|ogg|opus|wav|webm|wma)$/i; @@ -597,7 +605,7 @@ function setWidgetValue(widget, value) { } function restoreCachedAudioWidgets(widgets, saved) { - const beatData = saved?.beatData; + const beatData = saved?.sourceAnalysis || saved?.beatData; if (!beatData) return; const audioFile = String(beatData.audioFile || ""); const cacheKey = String(beatData.cacheKey || ""); @@ -894,6 +902,165 @@ function cropWaveformPreview(preview, startSeconds, duration) { }; } +function analysisArray(value, snake, camel = snake) { + const values = value?.[snake] ?? value?.[camel]; + return Array.isArray(values) ? values.map((entry) => finiteNumber(entry)) : []; +} + +function sourceAnalysisValue(value) { + if (!value || value.type !== "fl_audio_source_analysis" || + finiteNumber(value.version) !== SOURCE_ANALYSIS_VERSION) { + return null; + } + const duration = finiteNumber( + value.source_duration ?? value.sourceDuration ?? value.audio_duration ?? value.audioDuration, + ); + if (!(duration > 0)) return null; + return { + type: "fl_audio_source_analysis", + version: SOURCE_ANALYSIS_VERSION, + analysisVersion: finiteNumber(value.analysis_version ?? value.analysisVersion), + bpm: finiteNumber(value.bpm), + baseGridIntervalSeconds: finiteNumber( + value.base_grid_interval_seconds ?? value.baseGridIntervalSeconds, + ), + beatTimes: analysisArray(value, "beat_times", "beatTimes"), + downbeatTimes: analysisArray(value, "downbeat_times", "downbeatTimes"), + detectedBeatTimes: analysisArray(value, "detected_beat_times", "detectedBeatTimes"), + detectedDownbeatTimes: analysisArray(value, "detected_downbeat_times", "detectedDownbeatTimes"), + detectedBeatConfidences: analysisArray(value, "detected_beat_confidences", "detectedBeatConfidences"), + detectedDownbeatConfidences: analysisArray( + value, + "detected_downbeat_confidences", + "detectedDownbeatConfidences", + ), + onsetTimes: analysisArray(value, "onset_times", "onsetTimes"), + drumTimes: value.drum_times || value.drumTimes || {}, + duration, + supportsHalfTime: value.supports_half_time == null && value.supportsHalfTime == null + ? true + : Boolean(value.supports_half_time ?? value.supportsHalfTime), + waveformPreview: normalizeWaveformPreview(value.waveform_preview || value.waveformPreview), + waveformPreviewStart: finiteNumber( + value.waveform_preview_start ?? value.waveformPreviewStart, + ), + cacheKey: String(value.cache_key || value.cacheKey || ""), + audioFile: String(value.audio_file || value.audioFile || ""), + detector: value.detector || null, + detectorVersion: String(value.detector_version || value.detectorVersion || ""), + bpmSource: String(value.bpm_source || value.bpmSource || ""), + analysisSource: String(value.analysis_source || value.analysisSource || "mix"), + beatAnalysisSource: String(value.beat_analysis_source || value.beatAnalysisSource || "mix"), + analysisCacheHit: Boolean(value.analysis_cache_hit ?? value.analysisCacheHit), + }; +} + +function sourceAnalysisFromCropPayload(value) { + if (!value) return null; + const sourceStart = Math.max(0, finiteNumber(value.source_start ?? value.sourceStart)); + const cropDuration = Math.max(0, finiteNumber(value.audio_duration ?? value.audioDuration)); + const sourceDuration = Math.max( + sourceStart + cropDuration, + finiteNumber(value.source_duration ?? value.sourceDuration), + ); + if (!(cropDuration > 0) || !(sourceDuration > 0)) return null; + + const offset = finiteNumber(value.beat_offset_ms ?? value.beatOffsetMs) / 1000; + const payloadBeats = analysisArray(value, "beat_times", "beatTimes"); + const payloadDownbeats = analysisArray(value, "downbeat_times", "downbeatTimes"); + const baseBeats = analysisArray(value, "base_beat_times", "baseBeatTimes"); + const baseDownbeats = analysisArray(value, "base_downbeat_times", "baseDownbeatTimes"); + const baseDetectedBeats = analysisArray( + value, + "base_detected_beat_times", + "baseDetectedBeatTimes", + ); + const baseDetectedDownbeats = analysisArray( + value, + "base_detected_downbeat_times", + "baseDetectedDownbeatTimes", + ); + const baseBeatConfidences = analysisArray( + value, + "base_detected_beat_confidences", + "baseDetectedBeatConfidences", + ); + const baseDownbeatConfidences = analysisArray( + value, + "base_detected_downbeat_confidences", + "baseDetectedDownbeatConfidences", + ); + const drums = { ...(value.drum_times || value.drumTimes || {}) }; + for (const [snake, camel] of [ + ["kick_times", "kickTimes"], + ["snare_times", "snareTimes"], + ["hihat_times", "hihatTimes"], + ]) { + drums[snake] = sourceTimes(analysisArray(drums, snake, camel), sourceStart); + } + + return sourceAnalysisValue({ + type: "fl_audio_source_analysis", + version: SOURCE_ANALYSIS_VERSION, + analysis_version: value.analysis_version ?? value.analysisVersion, + bpm: value.bpm, + base_grid_interval_seconds: value.base_grid_interval_seconds ?? value.baseGridIntervalSeconds, + beat_times: sourceTimes( + baseBeats.length ? baseBeats : payloadBeats.map((entry) => entry - offset), + sourceStart, + ), + downbeat_times: sourceTimes( + baseDownbeats.length ? baseDownbeats : payloadDownbeats.map((entry) => entry - offset), + sourceStart, + ), + detected_beat_times: sourceTimes( + baseDetectedBeats.length + ? baseDetectedBeats + : analysisArray(value, "detected_beat_times", "detectedBeatTimes"), + sourceStart, + ), + detected_downbeat_times: sourceTimes( + baseDetectedDownbeats.length + ? baseDetectedDownbeats + : analysisArray(value, "detected_downbeat_times", "detectedDownbeatTimes"), + sourceStart, + ), + detected_beat_confidences: baseBeatConfidences.length + ? baseBeatConfidences + : analysisArray(value, "detected_beat_confidences", "detectedBeatConfidences"), + detected_downbeat_confidences: baseDownbeatConfidences.length + ? baseDownbeatConfidences + : analysisArray(value, "detected_downbeat_confidences", "detectedDownbeatConfidences"), + onset_times: sourceTimes(analysisArray(value, "onset_times", "onsetTimes"), sourceStart), + drum_times: drums, + source_duration: sourceDuration, + supports_half_time: false, + waveform_preview: value.waveform_preview || value.waveformPreview, + waveform_preview_start: sourceStart, + cache_key: value.cache_key || value.cacheKey, + audio_file: value.audio_file || value.audioFile, + detector: value.detector, + detector_version: value.detector_version || value.detectorVersion, + bpm_source: value.bpm_source || value.bpmSource, + analysis_source: value.analysis_source || value.analysisSource, + beat_analysis_source: value.beat_analysis_source || value.beatAnalysisSource, + analysis_cache_hit: value.analysis_cache_hit ?? value.analysisCacheHit, + }); +} + +function medianInterval(values) { + const intervals = values + .slice(1) + .map((value, index) => value - values[index]) + .filter((value) => value > EPSILON) + .sort((left, right) => left - right); + if (!intervals.length) return 0; + const middle = Math.floor(intervals.length / 2); + return intervals.length % 2 + ? intervals[middle] + : (intervals[middle - 1] + intervals[middle]) / 2; +} + function audioViewURL(value) { const match = String(value || "").match(/^(.*?)(?:\s+\[(input|output|temp)\])?$/); const relative = (match?.[1] || "").replace(/\\/g, "/"); @@ -933,6 +1100,7 @@ class BeatPromptSequencer { this.hover = null; this.sourceWaveformPreview = null; this.sourceAudioDuration = 0; + this.sourceAnalysis = null; this.audioElement = null; this.audioURL = ""; this.playbackFrameRequest = null; @@ -952,9 +1120,12 @@ class BeatPromptSequencer { const saved = node.properties?.flBeatPromptSequencer || {}; const savedCompatible = COMPATIBLE_FORMAT_VERSIONS.has(finiteNumber(saved.formatVersion)); - this.beatData = savedCompatible ? saved.beatData || null : null; + this.sourceAnalysis = savedCompatible ? sourceAnalysisValue(saved.sourceAnalysis) : null; + this.beatData = savedCompatible && !this.sourceAnalysis ? saved.beatData || null : null; if (this.beatData) { this.beatData.waveformPreview = normalizeWaveformPreview(this.beatData.waveformPreview); + this.sourceAnalysis = sourceAnalysisFromCropPayload(this.beatData); + if (this.sourceAnalysis) this.beatData = null; } restoreCachedAudioWidgets(this.widgets, saved); this.dataFresh = false; @@ -1015,7 +1186,7 @@ class BeatPromptSequencer { No audio selected Choose audio to load the timeline -