import numpy as np import torch import torchaudio from transformers import AutoProcessor, MusicgenForConditionalGeneration # transformers-native MusicGen. the audiocraft package is abandoned and # hard-pins torch==2.1.0 / xformers<0.0.23, which breaks the HF Space build # against gradio 6.x. transformers supports audio-prompted continuation # directly, so we don't need audiocraft at all. MODEL_ID = "facebook/musicgen-large" MUSICGEN_SR = 32000 FRAME_RATE = 50 # musicgen decoder tokens per second _model = None _processor = None def _load_model(): global _model, _processor if _model is None: _processor = AutoProcessor.from_pretrained(MODEL_ID) device = "cuda" if torch.cuda.is_available() else "cpu" dtype = torch.float16 if device == "cuda" else torch.float32 _model = MusicgenForConditionalGeneration.from_pretrained( MODEL_ID, torch_dtype=dtype ).to(device) _model.eval() return _model, _processor def _load_tail(path, prompt_duration): """load audio, return mono 32kHz numpy tail of `prompt_duration` seconds.""" track, sr = torchaudio.load(path) if sr != MUSICGEN_SR: track = torchaudio.transforms.Resample(sr, MUSICGEN_SR)(track) if track.shape[0] > 1: track = track.mean(dim=0, keepdim=True) tail_samples = int(prompt_duration * MUSICGEN_SR) if track.shape[1] > tail_samples: track = track[:, -tail_samples:] return track.squeeze(0).numpy() def continue_track(path, prompt_duration=10, gen_duration=15, key=None, bpm=None): """ takes the last `prompt_duration` seconds of the input track and generates `gen_duration` seconds of continuation. returns (continuation_only as 1-D float32 numpy, sample_rate). """ model, processor = _load_model() tail = _load_tail(path, prompt_duration) desc = "continue this song" if key and bpm: desc = f"continue this song in {key} at {round(bpm)} bpm" elif key: desc = f"continue this song in {key}" inputs = processor( audio=tail, sampling_rate=MUSICGEN_SR, text=[desc], padding=True, return_tensors="pt", ).to(model.device) # cast audio prompt to model dtype (fp16 on gpu) if "input_values" in inputs: inputs["input_values"] = inputs["input_values"].to(model.dtype) with torch.no_grad(): output = model.generate( **inputs, do_sample=True, guidance_scale=3.0, max_new_tokens=int(gen_duration * FRAME_RATE), ) audio = output[0, 0].float().cpu().numpy() # generate_continuation-style output contains the prompt audio at the # start; trim it so we return only the new material. if audio.shape[0] > tail.shape[0]: audio = audio[tail.shape[0]:] return audio, MUSICGEN_SR def stitch_with_crossfade(original, continuation, sr, fade_seconds=0.5): """ join original track and continuation with an equal-power crossfade so the seam doesn't click. both inputs 1-D numpy at the same sr. """ fade = int(fade_seconds * sr) fade = min(fade, len(original), len(continuation)) if fade <= 0: return np.concatenate([original, continuation]) t = np.linspace(0.0, np.pi / 2, fade, dtype=np.float32) fade_out = np.cos(t) fade_in = np.sin(t) head = original[:-fade] seam = original[-fade:] * fade_out + continuation[:fade] * fade_in rest = continuation[fade:] out = np.concatenate([head, seam, rest]).astype(np.float32) peak = np.abs(out).max() if peak > 1.0: out = out / peak return out