Spaces:
Running on Zero
Running on Zero
Download app.py from teamup-tech/SongFormer: direct link, hf CLI and curl.
- Browser
- Download file 10.2 kB
-
https://huggingface.co/spaces/teamup-tech/SongFormer/resolve/main/app.py
- Command line
-
hf download hf://spaces/teamup-tech/SongFormer/app.py
-
curl -L -o app.py https://huggingface.co/spaces/teamup-tech/SongFormer/resolve/main/app.py
10.2 kB
| 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"], | |
| ) | |
| 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) | |