hum / test_sampler.py
batoon's picture
Handle missing SheetSage beat grid in short inputs
bbcd177 verified
Raw History Blame Contribute Delete
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()