import os import sys import json import math import tempfile import uuid from pathlib import Path import spaces ROOT = Path(__file__).resolve().parent SONGFORMER = ROOT / "src" / "SongFormer" os.chdir(SONGFORMER) sys.path.insert(0, str(ROOT / "src" / "third_party")) sys.path.insert(0, str(SONGFORMER)) import numpy as np import gradio as gr import torch import librosa from omegaconf import OmegaConf from ema_pytorch import EMA from huggingface_hub import hf_hub_download from muq import MuQ from musicfm.model.musicfm_25hz import MusicFM25Hz from postprocessing.functional import postprocess_functional_structure from dataset.label2id import DATASET_ID_ALLOWED_LABEL_IDS, DATASET_LABEL_TO_DATASET_ID from pyharp import ModelCard, build_endpoint AFTER_DOWNSAMPLING_FRAME_RATES = 8.333 DATASET_LABEL = "SongForm-HX-8Class" DATASET_IDS = [5] TIME_DUR = 420 INPUT_SAMPLING_RATE = 24000 OUTPUT_ROOT = Path(tempfile.gettempdir()) / "songformer_outputs" def initialize_models(): from models.SongFormer import Model from safetensors.torch import load_file device = torch.device("cuda") hp = OmegaConf.load("configs/SongFormer.yaml") muq_model = MuQ.from_pretrained("OpenMuQ/MuQ-large-msd-iter") muq_model = muq_model.to(device).eval() musicfm_model = MusicFM25Hz( is_flash=False, stat_path=hf_hub_download("minzwon/MusicFM", "msd_stats.json"), model_path=hf_hub_download("minzwon/MusicFM", "pretrained_msd.pt"), ) musicfm_model = musicfm_model.to(device).eval() msa_model = Model(hp) checkpoint = hf_hub_download("ASLP-lab/SongFormer", "SongFormer.safetensors") model_ema = EMA(msa_model, include_online_model=False) model_ema.load_state_dict(load_file(checkpoint, device="cpu")) msa_model.load_state_dict(model_ema.ema_model.state_dict()) msa_model.to(device).eval() return muq_model, musicfm_model, msa_model, hp, device MODELS = initialize_models() def process_audio(audio_path, muq_model, musicfm_model, msa_model, hp, device, win_size=420, hop_size=420, num_classes=128): wav, _ = librosa.load(audio_path, sr=INPUT_SAMPLING_RATE) if len(wav) < 5 * INPUT_SAMPLING_RATE: raise gr.Error("Audio must be at least 5 seconds long.") audio = torch.tensor(wav).to(device) # Prepare output total_len = ( (audio.shape[0] // INPUT_SAMPLING_RATE) // TIME_DUR * TIME_DUR ) + TIME_DUR total_frames = math.ceil(total_len * AFTER_DOWNSAMPLING_FRAME_RATES) logits = { "function_logits": np.zeros([total_frames, num_classes]), "boundary_logits": np.zeros([total_frames]), } logits_num = { "function_logits": np.zeros([total_frames, num_classes]), "boundary_logits": np.zeros([total_frames]), } # Prepare label masks dataset_id2label_mask = {} for key, allowed_ids in DATASET_ID_ALLOWED_LABEL_IDS.items(): dataset_id2label_mask[key] = np.ones(num_classes, dtype=bool) dataset_id2label_mask[key][allowed_ids] = False lens = 0 i = 0 with torch.no_grad(): while True: start_idx = i * INPUT_SAMPLING_RATE end_idx = min((i + win_size) * INPUT_SAMPLING_RATE, audio.shape[-1]) if start_idx >= audio.shape[-1]: break if end_idx - start_idx <= 1024: break audio_seg = audio[start_idx:end_idx] # Get embeddings muq_output = muq_model(audio_seg.unsqueeze(0), output_hidden_states=True) muq_embd_420s = muq_output["hidden_states"][10] del muq_output torch.cuda.empty_cache() _, musicfm_hidden_states = musicfm_model.get_predictions( audio_seg.unsqueeze(0) ) musicfm_embd_420s = musicfm_hidden_states[10] del musicfm_hidden_states torch.cuda.empty_cache() # Process 30-second segments wraped_muq_embd_30s = [] wraped_musicfm_embd_30s = [] for idx_30s in range(i, i + hop_size, 30): start_idx_30s = idx_30s * INPUT_SAMPLING_RATE end_idx_30s = min( (idx_30s + 30) * INPUT_SAMPLING_RATE, audio.shape[-1], (i + hop_size) * INPUT_SAMPLING_RATE, ) if start_idx_30s >= audio.shape[-1]: break if end_idx_30s - start_idx_30s <= 1024: continue wraped_muq_embd_30s.append( muq_model( audio[start_idx_30s:end_idx_30s].unsqueeze(0), output_hidden_states=True, )["hidden_states"][10] ) torch.cuda.empty_cache() wraped_musicfm_embd_30s.append( musicfm_model.get_predictions( audio[start_idx_30s:end_idx_30s].unsqueeze(0) )[1][10] ) torch.cuda.empty_cache() wraped_muq_embd_30s = torch.concatenate(wraped_muq_embd_30s, dim=1) wraped_musicfm_embd_30s = torch.concatenate( wraped_musicfm_embd_30s, dim=1 ) all_embds = [ wraped_musicfm_embd_30s, wraped_muq_embd_30s, musicfm_embd_420s, muq_embd_420s, ] # Align embedding lengths min_embd_len = min(x.shape[1] for x in all_embds) all_embds = [x[:, :min_embd_len, :] for x in all_embds] embd = torch.concatenate(all_embds, axis=-1) # Inference dataset_ids = torch.Tensor(DATASET_IDS).to(device, dtype=torch.long) msa_info, chunk_logits = msa_model.infer( input_embeddings=embd, dataset_ids=dataset_ids, label_id_masks=torch.Tensor( dataset_id2label_mask[ DATASET_LABEL_TO_DATASET_ID[DATASET_LABEL] ] ) .to(device, dtype=bool) .unsqueeze(0) .unsqueeze(0), with_logits=True, ) # Accumulate logits start_frame = int(i * AFTER_DOWNSAMPLING_FRAME_RATES) end_frame = start_frame + min( math.ceil(hop_size * AFTER_DOWNSAMPLING_FRAME_RATES), chunk_logits["boundary_logits"][0].shape[0], ) logits["function_logits"][start_frame:end_frame, :] += ( chunk_logits["function_logits"][0].detach().cpu().numpy() ) logits["boundary_logits"][start_frame:end_frame] = ( chunk_logits["boundary_logits"][0].detach().cpu().numpy() ) logits_num["function_logits"][start_frame:end_frame, :] += 1 logits_num["boundary_logits"][start_frame:end_frame] += 1 lens += end_frame - start_frame i += hop_size # Average logits logits["function_logits"] /= np.maximum(logits_num["function_logits"], 1) logits["boundary_logits"] /= np.maximum(logits_num["boundary_logits"], 1) logits["function_logits"] = torch.from_numpy( logits["function_logits"][:lens] ).unsqueeze(0) logits["boundary_logits"] = torch.from_numpy( logits["boundary_logits"][:lens] ).unsqueeze(0) # Post-process msa_infer_output = postprocess_functional_structure(logits, hp) return logits, msa_infer_output def rule_post_processing(msa_list): if len(msa_list) <= 2: return msa_list result = msa_list.copy() while len(result) > 2: first_duration = result[1][0] - result[0][0] if first_duration < 1.0: result[0] = (result[0][0], result[1][1]) result = [result[0]] + result[2:] else: break while len(result) > 2: last_label_duration = result[-1][0] - result[-2][0] if last_label_duration < 1.0: result = result[:-2] + [result[-1]] else: break while len(result) > 2: if result[0][1] == result[1][1] and result[1][0] <= 10.0: result = [(result[0][0], result[0][1])] + result[2:] else: break while len(result) > 2: last_duration = result[-1][0] - result[-2][0] if result[-2][1] == result[-3][1] and last_duration <= 10.0: result = result[:-2] + [result[-1]] else: break return result model_card = ModelCard( name="SongFormer", description="Analyze the section boundaries and labels of a music recording.", author="ASLP Lab", tags=["music-structure-analysis", "music-information-retrieval"], ) @spaces.GPU(duration=240) def process_fn(audio_path: str | None) -> str: if not audio_path: raise gr.Error("Upload a music recording.") _, msa_output = process_audio(audio_path, *MODELS) msa_output = rule_post_processing(msa_output) segments = [ { "start": round(float(msa_output[index][0]), 2), "end": round(float(msa_output[index + 1][0]), 2), "label": msa_output[index][1], } for index in range(len(msa_output) - 1) ] output_dir = OUTPUT_ROOT / uuid.uuid4().hex output_dir.mkdir(parents=True) output_path = output_dir / "structure.json" output_path.write_text(json.dumps({"segments": segments}, indent=2) + "\n") return str(output_path) with gr.Blocks(title="SongFormer Music Structure Analysis") as demo: build_endpoint( model_card=model_card, input_components=[ gr.Audio(type="filepath", label="Music Recording") .harp_required(True) .set_info("Upload a music recording for section analysis."), ], output_components=[ gr.File(type="filepath", file_types=[".json"], label="Structure") .set_info("Section start, end, and label in seconds."), ], process_fn=process_fn, ) if __name__ == "__main__": demo.queue(default_concurrency_limit=1).launch(show_error=True)