CoolFace
Apppublic

mtg-upf/audio-difficulty

sourceHugging Faceupdated 1y agoView on Hugging Face
2likes
get_difficulty.py116 linesDownload Raw Back to root
1import os2import pdb3from statistics import mean4 5import torch6from torch import nn7import numpy as np8import librosa9from piano_transcription_inference import PianoTranscription, sample_rate, load_audio10import pretty_midi11from utils import prediction2label12from model import AudioModel13from scipy.signal import resample14 15 16def downsample_log_cqt(cqt_matrix, target_fs=5):17    original_fs = 44100 / 16018    ratio = original_fs / target_fs19    downsampled = resample(cqt_matrix, int(cqt_matrix.shape[0] / ratio), axis=0)20    return downsampled21 22def downsample_matrix(mat, original_fs, target_fs):23    ratio = original_fs / target_fs24    return resample(mat, int(mat.shape[0] / ratio), axis=0)25 26def get_cqt_from_mp3(mp3_path):27    sample_rate = 4410028    hop_length = 16029    y, sr = librosa.load(mp3_path, sr=sample_rate, mono=True)30    cqt = librosa.cqt(y, sr=sr, hop_length=hop_length, n_bins=88, bins_per_octave=12)31    log_cqt = librosa.amplitude_to_db(np.abs(cqt))32    log_cqt = log_cqt.T  # shape (T, 88)33    log_cqt = downsample_log_cqt(log_cqt, target_fs=5)34    cqt_tensor = torch.tensor(log_cqt, dtype=torch.float32).unsqueeze(0).unsqueeze(0).cpu()35    print(f"cqt shape: {log_cqt.shape}")36    return cqt_tensor37 38def get_pianoroll_from_mp3(mp3_path):39    audio, _ = load_audio(mp3_path, sr=sample_rate, mono=True)40    transcriptor = PianoTranscription(device="cuda" if torch.cuda.is_available() else "cpu")41    midi_path = "temp.mid"42    transcriptor.transcribe(audio, midi_path)43    midi_data = pretty_midi.PrettyMIDI(midi_path)44 45    fs = 5  # original frames per second46    piano_roll = midi_data.get_piano_roll(fs=fs)[21:109].T  # shape: (T, 88)47    piano_roll = piano_roll / 12748    time_steps = piano_roll.shape[0]49 50    onsets = np.zeros_like(piano_roll)51    for instrument in midi_data.instruments:52        for note in instrument.notes:53            pitch = note.pitch - 2154            onset_frame = int(note.start * fs)55            if 0 <= pitch < 88 and onset_frame < time_steps:56                onsets[onset_frame, pitch] = 1.057 58    pr_tensor = torch.tensor(piano_roll.T).unsqueeze(0).unsqueeze(1).cpu().float()59    on_tensor = torch.tensor(onsets.T).unsqueeze(0).unsqueeze(1).cpu().float()60    out_tensor = torch.cat([pr_tensor, on_tensor], dim=1)61    print(f"piano_roll shape: {out_tensor.shape}")62    return out_tensor.transpose(2, 3)63 64def predict_difficulty(mp3_path, model_name, rep):65    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")66 67    if "only_cqt" in rep:68        only_cqt, only_pr = True, False69        rep_clean = "multimodal5"70    elif "only_pr" in rep:71        only_cqt, only_pr = False, True72        rep_clean = "multimodal5"73    else:74        only_cqt = only_pr = False75        rep_clean = rep76 77    model = AudioModel(num_classes=11, rep=rep_clean, modality_dropout=False, only_cqt=only_cqt, only_pr=only_pr).to(device)78    checkpoint = [torch.load(f"models/{model_name}/checkpoint_{i}.pth", map_location=device, weights_only=False)79                  for i in range(5)]80 81    if rep == "cqt5":82        inp_data = get_cqt_from_mp3(mp3_path).to(device)83    elif rep == "pianoroll5":84        inp_data = get_pianoroll_from_mp3(mp3_path).to(device)85    elif rep_clean == "multimodal5":86        x1 = get_pianoroll_from_mp3(mp3_path).to(device)87        x2 = get_cqt_from_mp3(mp3_path).to(device)88        inp_data = [x1, x2]89    else:90        raise ValueError(f"Representation {rep} not supported")91 92    preds = []93    for cheks in checkpoint:94        model.load_state_dict(cheks["model_state_dict"])95        model.eval()96        with torch.inference_mode():97            logits = model(inp_data, None)98            pred = prediction2label(logits).item()99            preds.append(pred)100 101    return mean(preds)102 103if __name__ == "__main__":104    mp3_path = "yt_audio.mp3"105    model_name = "audio_midi_multi_ps_v5"106    pred_multi = predict_difficulty(mp3_path, model_name=model_name, rep="multimodal5")107    print(f"Multimodal: {pred_multi}")108 109    model_name = "audio_midi_pianoroll_ps_5_v4"110    pred_multi = predict_difficulty(mp3_path, model_name=model_name, rep="pianoroll5")111    print(f"Pianoroll: {pred_multi}")112 113    model_name = "audio_midi_multi_ps_v5"114    pred_multi = predict_difficulty(mp3_path, model_name=model_name, rep="pianoroll5")115    print(f"CQT: {pred_multi}")116