SongFormer / app.py
LING
Tighten SongFormer documentation and inference flow
884eb42
Raw History Blame Contribute Delete
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"],
)
@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)