# %% import os import glob import random from suno_utils.audio import Audio from suno_utils.audio.midi import Midi import shutil from tqdm import tqdm import torchcrepe import tempfile import torch import pesto import torchaudio import numpy as np import librosa import copy import matplotlib.pyplot as plt os.environ["CUDA_VISIBLE_DEVICES"] = "4,5,6,7" # %% trombone_champ_dir = "/app2/suno/data/victor/trombone_champ_vocals" # %% class MidiPair: def __init__(self, midi_file: str, audio_file: str): self.midi_file = midi_file self.audio_file = audio_file self.midi = None self.audio = None def load_midi(self): if self.midi is not None: return self.midi self.midi = Midi.from_path(self.midi_file) return self.midi def load_audio(self): if self.audio is not None: return self.audio self.audio = Audio.from_file(self.audio_file) return self.audio def load_all(self): self.load_midi() self.load_audio() def __str__(self): return f"MidiPair(midi_file={self.midi_file}, audio_file={self.audio_file})" def __repr__(self): return self.__str__() def play(self): self.load_all() stereo_audio = self.midi.make_stereo_comparison(self.audio) stereo_audio.play() def load_pairs(dir: str, audio_ext: str = "mp3", midi_ext: str = "mid"): midi_files = glob.glob(os.path.join(dir, "**", f"*.{midi_ext}"), recursive=True) audio_files = glob.glob(os.path.join(dir, "**", f"*.{audio_ext}"), recursive=True) # Create a mapping of base filenames to audio files audio_map = {} for audio_file in audio_files: base_name = os.path.splitext(os.path.basename(audio_file))[0] audio_map[base_name] = audio_file pairs = [] for midi_file in midi_files: base_name = os.path.splitext(os.path.basename(midi_file))[0] if base_name in audio_map: pairs.append(MidiPair(midi_file, audio_map[base_name])) return pairs # %% [markdown] # ## trombone champ # %% trombone_champ_pairs = load_pairs(trombone_champ_dir, audio_ext="opus") print(f"Loaded {len(trombone_champ_pairs)} trombone champ pairs") # %% [markdown] # ## fix octave shifts # %% def plot_f0_comparison(results, stem, midi): # Convert frequencies to MIDI note numbers (pitches) frequencies = torchcrepe.filter.median(results, 30)[0] # zero out sections that are silent hz = int(len(frequencies) / stem.duration_s) for i in range(int(stem.duration_s)): if stem.get_slice(i, i + 1).loudness < -35: frequencies[i * hz : (i + 1) * hz] = 20 pitches = librosa.hz_to_midi(frequencies) # Create time axis time_axis = [t / hz for t in range(len(frequencies))] plt.figure(figsize=(12, 6)) plt.plot(time_axis, pitches, label="Vocal F0") # filter out notes that are silent for i in reversed(range(len(midi.pmidi.instruments[0].notes))): audio_slice = stem.get_slice( midi.pmidi.instruments[0].notes[i].start, midi.pmidi.instruments[0].notes[i].end + 1 ) if np.mean(audio_slice.array_float**2) < 1e-6: midi.pmidi.instruments[0].notes.pop(i) # Overlay MIDI notes for note in midi.pmidi.instruments[0].notes: start_time = note.start end_time = note.end pitch = note.pitch plt.hlines( pitch, start_time, end_time, colors="red", linewidth=2, alpha=0.7, label="MIDI Notes" if note == midi.pmidi.instruments[0].notes[0] else "", ) plt.xlabel("Time (s)") plt.ylabel("MIDI Note Number") plt.legend() plt.show() # %% def get_f0_pesto(audio): with tempfile.NamedTemporaryFile(suffix=".wav") as f: audio.write_wav(f.name) x, sr = torchaudio.load(f.name) x = x.mean(dim=0) _, pitch, confidence, _ = pesto.predict(x, sr) return pitch.clone(), confidence.clone() def get_f0(audio): # Process in 60s chunks to avoid OOM chunk_duration = 60 total_duration = audio.duration_s #min(240, audio.duration_s) all_results = [] all_period_results = [] for start_time in range(0, int(total_duration), chunk_duration): end_time = min(start_time + chunk_duration, total_duration) with tempfile.NamedTemporaryFile(suffix=".wav") as f: audio.get_slice(start_time, end_time).write_wav(f.name) chunk_results, period_results = torchcrepe.predict_from_file( f.name, device="cuda:3", decoder=torchcrepe.decode.weighted_argmax, pad=False, return_periodicity=True ) all_results.append(chunk_results) all_period_results.append(period_results) # Concatenate results results = torch.cat([r[0] for r in all_results], dim=0).unsqueeze(0) period_final_results = torch.cat([r[0] for r in all_period_results], dim=0).unsqueeze(0) return results, period_final_results def octave_dist(note1, note2): diff = abs(note1 - note2) return min(diff, 12 - diff) def fix_midi_octave_errors(pair, verbose=False, pitch_algo="crepe", λ=500, confidence_threshold=0.05): midi = pair.load_midi() stem = pair.load_audio().resample(22050) if verbose: print(f"MIDI has {len(midi.pmidi.instruments)} instruments") if midi.pmidi.instruments: print(f"First instrument has {len(midi.pmidi.instruments[0].notes)} notes") else: print("No instruments found!") # get vocal f0 if pitch_algo == "pesto": frequencies, confidences = get_f0_pesto(stem) results = frequencies.unsqueeze(0) else: results, confidences = get_f0(stem) confidences = confidences[0] frequencies = torchcrepe.filter.median(results, 30)[0] # zero out sections that are silent OR low confidence hz = int(len(frequencies) / stem.duration_s) device = frequencies.device for i in range(int(stem.duration_s)): if stem.get_slice(i, i + 1).loudness < -35: frequencies[i * hz : (i + 1) * hz] = 20 if confidences is not None: confidences[i * hz : (i + 1) * hz] = 0 # Validate MIDI structure if not midi.pmidi.instruments: raise ValueError("MIDI file has no instruments") if not midi.pmidi.instruments[0].notes: raise ValueError("MIDI file has no notes in the first instrument") # viterbi octave correction Ks = torch.arange(-3, 4, device=device) # [-3, -2, -1, 0, 1, 2, 3] N = len(midi.pmidi.instruments[0].notes) dp = torch.full((N, len(Ks)), float('inf'), device=device) prev = torch.zeros((N, len(Ks)), dtype=torch.long, device=device) note_estimates = [] note_confidences = [] for note in midi.pmidi.instruments[0].notes: start_idx, end_idx = int(note.start * hz), int(note.end * hz) note_pitches = frequencies[start_idx:end_idx] if confidences is not None: note_confs = confidences[start_idx:end_idx] # filter by confidence and silence valid_mask = (note_pitches > 20) & (note_confs > confidence_threshold) valid_pitches = note_pitches[valid_mask] valid_confs = note_confs[valid_mask] if len(valid_pitches) > 0: # weighted median median_pitch = weighted_median_torch(valid_pitches, valid_confs) # average confidence for this note avg_confidence = torch.mean(valid_confs) else: median_pitch = torch.tensor(float('nan'), device=device) avg_confidence = torch.tensor(0.0, device=device) else: # fallback for CREPE valid_pitches = note_pitches[note_pitches > 20] if len(valid_pitches) > 0: median_pitch = torch.median(valid_pitches) else: median_pitch = torch.tensor(float('nan'), device=device) avg_confidence = torch.tensor(1.0, device=device) # Convert Hz to MIDI if torch.isnan(median_pitch): note_estimates.append(median_pitch) else: midi_pitch = 69 + 12 * torch.log2(median_pitch / 440) note_estimates.append(midi_pitch) note_confidences.append(avg_confidence) notes = midi.pmidi.instruments[0].notes # init with confidence weighting for ik, k in enumerate(Ks): if torch.isnan(note_estimates[0]): dp[0, ik] = 0 else: error = (note_estimates[0] - (notes[0].pitch + 12 * k)) ** 2 confidence_weight = 1.0 / torch.clamp(note_confidences[0], min=0.1) dp[0, ik] = error * confidence_weight # fill with confidence weighting for i in range(1, N): for ik, k in enumerate(Ks): if torch.isnan(note_estimates[i]): obs = torch.tensor(0.0, device=device) else: error = (note_estimates[i] - (notes[i].pitch + 12 * k)) ** 2 confidence_weight = 1.0 / torch.clamp(note_confidences[i], min=0.1) obs = error * confidence_weight # transition costs transition_costs = λ * torch.abs(k - Ks) costs = dp[i - 1, :] + transition_costs dp[i, ik] = obs + torch.min(costs) prev[i, ik] = torch.argmin(costs) # backtrack best_path = torch.zeros(N, dtype=torch.long, device=device) best_path[-1] = torch.argmin(dp[-1, :]) for i in range(N - 2, -1, -1): best_path[i] = prev[i + 1, best_path[i + 1]] if verbose: print(best_path.cpu().numpy()) cost = torch.tensor(0.0, device=device) for i in range(N): if torch.isnan(note_estimates[i]): cost += 0 else: diff = note_estimates[i] - (notes[i].pitch + 12 * Ks[best_path[i]]) cost += diff**2 normalized_cost = cost / N if verbose: print(f"Cost: {cost.item()}, Normalized cost: {normalized_cost.item()}") new_midi = copy.deepcopy(midi) # apply shifts - convert back to CPU for MIDI manipulation best_path_cpu = best_path.cpu() Ks_cpu = Ks.cpu() for i, note in enumerate(notes): note.pitch += 12 * Ks_cpu[best_path_cpu[i]].item() return results, new_midi def weighted_median_torch(values, weights): """Compute weighted median using torch operations""" if len(values) == 0: return torch.tensor(float('nan'), device=values.device) sorted_indices = torch.argsort(values) sorted_values = values[sorted_indices] sorted_weights = weights[sorted_indices] cumsum = torch.cumsum(sorted_weights, dim=0) total_weight = cumsum[-1] median_pos = total_weight / 2 median_idx = torch.searchsorted(cumsum, median_pos) # Clamp to valid range median_idx = torch.clamp(median_idx, 0, len(sorted_values) - 1) return sorted_values[median_idx] # %% #midi_pair = random.choice(trombone_champ_pairs) # %% #stem = midi_pair.load_audio().resample(22050) #midi = midi_pair.load_midi() #pitches, new_midi = fix_midi_octave_errors(midi_pair,λ=500) #plot_f0_comparison(pitches, midi_pair.load_audio(), new_midi) # %% #stem = midi_pair.load_audio().resample(22050) #midi = midi_pair.load_midi() #pitches, new_midi = fix_midi_octave_errors(midi_pair, pitch_algo="pesto",λ=350) #new_midi.make_stereo_comparison(stem).play() #plot_f0_comparison(pitches, midi_pair.load_audio(), new_midi) # %% pitch_algo = "crepe" λ = 500 output_dir = f"/app2/suno/data/sara/trombone_champ_vocals_octaved_{pitch_algo}_{λ}" os.makedirs(output_dir, exist_ok=True) print(f"Writing to {output_dir}") for midi_pair in tqdm(trombone_champ_pairs): filename = midi_pair.audio_file.split("/")[-1].split(".")[0] try: pitches, new_midi = fix_midi_octave_errors(midi_pair, pitch_algo=pitch_algo, λ = λ) except Exception as e: print(f"Skipping {filename}: {e}") continue midi_path = os.path.join(output_dir, filename + ".mid") audio_path = os.path.join(output_dir, filename + ".opus") new_midi.write(midi_path) shutil.copy(midi_pair.audio_file, audio_path) try: new_midi.write(midi_path) shutil.copy(midi_pair.audio_file, audio_path) except Exception as e: print(f"Failed to write {midi_path}: {e}")