import torch import torchaudio from audiocraft.models import MusicGen _model = None def _load_model(): global _model if _model is None: _model = MusicGen.get_pretrained('facebook/musicgen-large') return _model 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. key and bpm are hints for the text prompt. """ model = _load_model() model.set_generation_params(duration=gen_duration) track, sr = torchaudio.load(path) # grab the tail end as context tail_samples = int(prompt_duration * sr) if track.shape[1] > tail_samples: tail = track[:, -tail_samples:] else: tail = track # resample to 32kHz if needed (musicgen expects this) if sr != 32000: resampler = torchaudio.transforms.Resample(sr, 32000) tail = resampler(tail) # mono if tail.shape[0] > 1: tail = tail.mean(dim=0, keepdim=True) tail = tail.unsqueeze(0) # batch dim # build a natural description desc = "continue this song" if key and bpm: desc = f"continue this song in {key} at {bpm} bpm" elif key: desc = f"continue this song in {key}" with torch.no_grad(): output = model.generate_continuation(tail, 32000, [desc]) result = output[0].cpu() return result, 32000