Download test_sampler.py from batoon/hum: direct link, hf CLI and curl.
- Browser
- Download file 12.9 kB
-
https://huggingface.co/spaces/batoon/hum/resolve/main/test_sampler.py
- Command line
-
hf download hf://spaces/batoon/hum/test_sampler.py
-
curl -L -o test_sampler.py https://huggingface.co/spaces/batoon/hum/resolve/main/test_sampler.py
12.9 kB
| """Sampler note timing and audio effect contracts.""" | |
| from __future__ import annotations | |
| from pathlib import Path | |
| import json | |
| import tempfile | |
| import unittest | |
| from unittest.mock import patch | |
| class SamplerTest(unittest.TestCase): | |
| def test_empty_sheet_sage_melody_uses_f0_for_sampler(self): | |
| import numpy as np | |
| import pretty_midi | |
| import soundfile as sf | |
| from pipeline import generate_all | |
| with tempfile.TemporaryDirectory() as directory: | |
| folder = Path(directory) | |
| wave = folder / "hum.wav" | |
| t = np.arange(24000) / 24000 | |
| sf.write(wave, .12 * np.sin(2 * np.pi * 261.626 * t), 24000) | |
| empty, abc, beat = (folder / name for name in | |
| ("empty.mid", "empty.abc", "beat.lab")) | |
| pretty_midi.PrettyMIDI().write(str(empty)) | |
| abc.write_text("unusable empty melody") | |
| def fake_mp3(_, target, __, **___): | |
| target.write_bytes(b"fake mp3") | |
| with (patch("pipeline.decode", return_value=(wave, wave, 1.)), | |
| patch("pipeline.transcribe", return_value=(empty, abc, abc, beat)), | |
| patch("pipeline.choose_bass_route", return_value=(-24, folder / "bass.sf2", 0, "test")), | |
| patch("pipeline.render_instrument", return_value=(wave, empty)), | |
| patch("pipeline.render_drums", return_value=(wave, empty, 3)), | |
| patch("pipeline.pair_with_drums", return_value=(wave, wave)), | |
| patch("pipeline.render_contour", return_value=wave), | |
| patch("pipeline.render_preset", return_value=(wave, empty)), | |
| patch("pipeline.apply_effect", return_value=wave), | |
| patch("pipeline.encode_mp3", side_effect=fake_mp3)): | |
| result = list(generate_all(str(wave), "Флейта", {}, 4096, False))[-1] | |
| self.assertTrue(result[-1]["complete"]) | |
| self.assertEqual(result[-1]["transcription_source"], "pYIN fallback") | |
| self.assertIn("SheetSage2 не нашёл нот", result[-1]["errors"]["sheet_sage"]) | |
| self.assertEqual(len(pretty_midi.PrettyMIDI(str(empty)).instruments[0].notes), 1) | |
| self.assertNotIn("Vocal", abc.read_text()) | |
| self.assertIsNone(result[12]) | |
| self.assertEqual(result[-1]["errors"]["instrument_drums"], | |
| "SheetSage2 не дал сетку долей; барабаны пропущены") | |
| self.assertNotIn("instrument_drums", result[-1]["source_timing_preserved"]) | |
| def test_archive_keeps_original_bytes_and_uses_opaque_name(self): | |
| from pipeline import archive_recording, source_sha256 | |
| with tempfile.TemporaryDirectory() as directory: | |
| root = Path(directory) | |
| source = root / "personal-name.m4a" | |
| source.write_bytes(b"original encoded audio bytes") | |
| archive_root = root / "private-bucket" | |
| archive_root.mkdir() | |
| digest = source_sha256(source) | |
| with patch("pipeline.RECORDINGS_DIR", archive_root): | |
| result = archive_recording(source, digest, 2.25) | |
| self.assertTrue(result["saved"]) | |
| self.assertNotIn("personal-name", result["key"]) | |
| saved_folder = archive_root / result["key"] | |
| self.assertEqual((saved_folder / "source.m4a").read_bytes(), source.read_bytes()) | |
| metadata = json.loads((saved_folder / "metadata.json").read_text()) | |
| self.assertEqual(metadata["source_sha256"], digest) | |
| self.assertEqual(metadata["source_seconds"], 2.25) | |
| def test_bass_font_sounds_at_both_ends_of_development_note_range(self): | |
| import numpy as np | |
| import pretty_midi | |
| import soundfile as sf | |
| from pipeline import SOUNDFONT_DIR, choose_bass_route, render_preset | |
| font = SOUNDFONT_DIR / "finger-bass.sf2" | |
| if not font.is_file(): | |
| self.skipTest("Pinned FreePats bass SoundFont not installed") | |
| # The seen development recording spans MIDI 54..68. Keep this tiny | |
| # synthetic fixture separate from the source audio and its title. | |
| with tempfile.TemporaryDirectory() as directory: | |
| folder = Path(directory) | |
| score = pretty_midi.PrettyMIDI() | |
| voice = pretty_midi.Instrument(program=0) | |
| voice.notes = [pretty_midi.Note(90, 54, .1, .35), | |
| pretty_midi.Note(90, 68, .6, .85)] | |
| score.instruments.append(voice) | |
| midi = folder / "notes.mid" | |
| score.write(str(midi)) | |
| shift, chosen_font, program, source = choose_bass_route(midi) | |
| self.assertEqual((shift, chosen_font, program), (-24, font, 0)) | |
| self.assertEqual(source, "FreePats Finger Bass YR") | |
| wav, _ = render_preset(midi, folder, "bass", 0, 1., folder / "log", | |
| soundfont=chosen_font, transpose=shift) | |
| audio, rate = sf.read(wav, dtype="float32", always_2d=True) | |
| for start in (.12, .62): | |
| span = audio[round(start * rate):round((start + .1) * rate)] | |
| self.assertGreater(float(np.sqrt(np.mean(span**2))), .001) | |
| def test_optional_yue_does_not_reclassify_sampler_outputs_as_missing(self): | |
| import numpy as np | |
| import pretty_midi | |
| import soundfile as sf | |
| from pipeline import generate_all | |
| abc_text = ('X:1\nM:4/4\nL:1/4\n' | |
| 'V: Vocal clef=treble name="Vocal Melody" snm="Vocal"\n' | |
| 'V: Ins clef=treble name="Ins Melody" snm="Inst."\n' | |
| 'K:C\nV: Vocal\nC|\nV: Ins\nZ|\n') | |
| with tempfile.TemporaryDirectory() as directory: | |
| folder = Path(directory) | |
| wave = folder / "input.wav" | |
| sf.write(wave, np.full(48000, .1, dtype="float32"), 48000) | |
| midi, abc, beat = (folder / name for name in | |
| ("notes.mid", "melody.abc", "beat.lab")) | |
| score = pretty_midi.PrettyMIDI() | |
| track = pretty_midi.Instrument(0) | |
| track.notes.append(pretty_midi.Note(90, 60, .1, .8)) | |
| score.instruments.append(track) | |
| score.write(str(midi)) | |
| abc.write_text(abc_text) | |
| beat.write_text("0 1 4 4\n.5 2 4 4\n1 3 4 4\n") | |
| def fake_yue(_, __, output, ___, ____, _____): | |
| for name in ("yue_solo", "yue_drums", "yue_arrangement"): | |
| target = output / name | |
| target.mkdir() | |
| sf.write(target / "raw.wav", np.full((48000, 2), .1, | |
| dtype="float32"), 48000) | |
| (target / "yue.json").write_text(json.dumps({"timing": {}})) | |
| yield name | |
| def fake_mp3(_, target, __, **___): | |
| target.write_bytes(b"fake mp3") | |
| with (patch("pipeline.decode", return_value=(wave, wave, 1.)), | |
| patch("pipeline.transcribe", return_value=(midi, abc, abc, beat)), | |
| patch("pipeline.choose_bass_route", return_value=(-24, folder / "bass.sf2", 0, "test")), | |
| patch("pipeline.render_instrument", return_value=(wave, midi)), | |
| patch("pipeline.render_drums", return_value=(wave, midi, 3)), | |
| patch("pipeline.pair_with_drums", return_value=(wave, wave)), | |
| patch("pipeline.render_contour", return_value=wave), | |
| patch("pipeline.render_preset", return_value=(wave, midi)), | |
| patch("pipeline.apply_effect", return_value=wave), | |
| patch("pipeline.encode_mp3", side_effect=fake_mp3), | |
| patch("pipeline.render_yue_batch", side_effect=fake_yue)): | |
| final = list(generate_all(str(wave), "Флейта", {}, 4096, True))[-1] | |
| self.assertEqual(final[-1]["errors"], {}) | |
| self.assertFalse(final[-1]["recording_archive"]["requested"]) | |
| with (patch("pipeline.decode", return_value=(wave, wave, 1.)), | |
| patch("pipeline.archive_recording", side_effect=OSError("bucket unavailable")), | |
| patch("pipeline.transcribe", return_value=(midi, abc, abc, beat)), | |
| patch("pipeline.choose_bass_route", return_value=(-24, folder / "bass.sf2", 0, "test")), | |
| patch("pipeline.render_instrument", return_value=(wave, midi)), | |
| patch("pipeline.render_drums", return_value=(wave, midi, 3)), | |
| patch("pipeline.pair_with_drums", return_value=(wave, wave)), | |
| patch("pipeline.render_contour", return_value=wave), | |
| patch("pipeline.render_preset", return_value=(wave, midi)), | |
| patch("pipeline.apply_effect", return_value=wave), | |
| patch("pipeline.encode_mp3", side_effect=fake_mp3), | |
| patch("pipeline.render_yue_batch", side_effect=fake_yue)): | |
| failed = list(generate_all(str(wave), "Флейта", {}, 4096, True, | |
| save_recording=True))[-1] | |
| self.assertEqual(failed[-1]["errors"]["recording_archive"], "bucket unavailable") | |
| self.assertFalse(failed[-1]["recording_archive"]["saved"]) | |
| self.assertTrue(failed[-1]["complete"]) | |
| def test_bass_register_preserves_note_times(self): | |
| import pretty_midi | |
| from pipeline import choose_bass_route, render_preset | |
| with tempfile.TemporaryDirectory() as directory: | |
| folder = Path(directory) | |
| original = pretty_midi.PrettyMIDI() | |
| voice = pretty_midi.Instrument(program=0) | |
| voice.notes = [ | |
| pretty_midi.Note(90, 60, .125, .375), | |
| pretty_midi.Note(90, 64, .5, .875), | |
| ] | |
| original.instruments.append(voice) | |
| source = folder / "source.mid" | |
| original.write(str(source)) | |
| shift, _, _, _ = choose_bass_route(source) | |
| self.assertEqual(shift, -24) | |
| with patch("pipeline.render_midi", return_value=folder / "bass.wav"): | |
| _, path = render_preset(source, folder, "bass", 0, 1., | |
| folder / "log", transpose=shift) | |
| notes = pretty_midi.PrettyMIDI(str(path)).instruments[0].notes | |
| self.assertEqual([n.pitch for n in notes], [36, 40]) | |
| for before, after in zip(voice.notes, notes): | |
| self.assertAlmostEqual(before.start, after.start, places=4) | |
| self.assertAlmostEqual(before.end, after.end, places=4) | |
| def test_bass_route_automatically_shifts_or_uses_full_bank(self): | |
| import pretty_midi | |
| from pipeline import choose_bass_route, GM_SOUNDFONT | |
| with tempfile.TemporaryDirectory() as directory: | |
| midi = Path(directory) / "notes.mid" | |
| for pitches, expected in (((42, 54), (-12, "FreePats Finger Bass YR")), | |
| ((50, 80), (-24, "FluidR3 Fingered Bass fallback"))): | |
| score = pretty_midi.PrettyMIDI() | |
| track = pretty_midi.Instrument(0) | |
| track.notes = [pretty_midi.Note(90, pitch, index * .3, | |
| index * .3 + .2) | |
| for index, pitch in enumerate(pitches)] | |
| score.instruments.append(track) | |
| score.write(str(midi)) | |
| shift, font, program, source = choose_bass_route(midi) | |
| self.assertEqual((shift, source), expected) | |
| if source.endswith("fallback"): | |
| self.assertEqual((font, program), (GM_SOUNDFONT, 33)) | |
| def test_effects_preserve_duration_and_change_waveform(self): | |
| import numpy as np | |
| import soundfile as sf | |
| from pipeline import apply_effect | |
| rate = 48000 | |
| t = np.arange(rate, dtype=np.float32) / rate | |
| source_wave = (.12 * np.sin(2 * np.pi * 220 * t)).astype("float32") | |
| with tempfile.TemporaryDirectory() as directory: | |
| folder = Path(directory) | |
| source = folder / "source.wav" | |
| sf.write(source, source_wave, rate, subtype="FLOAT") | |
| for kind in ("guitar", "bass"): | |
| result = apply_effect(source, folder / f"{kind}.wav", kind, | |
| 1., folder / "log") | |
| output, output_rate = sf.read(result, dtype="float32", | |
| always_2d=True) | |
| self.assertEqual(output_rate, rate) | |
| self.assertEqual(len(output), rate) | |
| self.assertTrue(np.isfinite(output).all()) | |
| self.assertGreater(np.mean(np.abs(output[:, 0] - source_wave)), .001) | |
| if __name__ == "__main__": | |
| unittest.main() | |