import time from beat_processor import BeatProcessor import subprocess import tempfile import boto3 import concurrent.futures from lib.redis_client import RedisClient, getenv from audiocraft.models import MusicGen from audiocraft.data.audio import audio_write import uuid import torchaudio from lib.audio_utils import download_audio, upload_audio, resample_audio, CODEC_SPEEDS from lib.progress import ( OpProgressSet, OpProgress, OpProgressMinSet, OpProgressFfmpeg, ) class RetryableError(Exception): pass class MusicGenRedisClient(RedisClient): def __init__(self, model, s3_bucket, *args): super().__init__(*args) self.model = MusicGen.get_pretrained(model) self.s3 = boto3.client("s3") self.s3_bucket = s3_bucket self.beat_processor = BeatProcessor() @classmethod def from_env(cls): return super().from_env(getenv("MUSICGEN_MODEL"), getenv("MUSICGEN_S3_BUCKET")) def warmup(self): self.model.set_generation_params(duration=1) self.model.generate_unconditional(1) def _get_audio_from_s3(self, key): with tempfile.TemporaryDirectory() as temp_dir: return torchaudio.load( download_audio( self.s3, self.s3_bucket, key, temp_dir, # TODO resample? ) ) def _grid_audio(self, wav, sr, bpm, op_prog): with tempfile.TemporaryDirectory() as temp_dir: wav = wav.cpu() output_path = audio_write( f"{temp_dir}/output", wav, sr, strategy="loudness", loudness_compressor=True, format="wav", ) beats = self.beat_processor.detect_beats(output_path) op_prog(1) first_bar_time = next((b[0] for b in beats if b[1] == 1), None) if first_bar_time is None: raise RetryableError("Could not detect first bar start") adjusted_beats = [ b[0] - first_bar_time for b in beats if b[0] > first_bar_time + 1.0e-6 ] if len(adjusted_beats) < 2: raise RetryableError("Not enough beats detected") cut_wav = wav[:, int(first_bar_time * sr) :] average_dt = sum( b - a for a, b in zip([0] + adjusted_beats[:-1], adjusted_beats) ) / len(adjusted_beats) duration_at_avg_bpm = average_dt * len(adjusted_beats) duration_at_target_bpm = len(adjusted_beats) * 60 / bpm min_duration_change = 1.0e6 min_duration_factor = None for i in [0.25, 0.5, 1.0, 2.0, 4.0]: duration_change = abs( 1 - duration_at_target_bpm * i / (duration_at_avg_bpm + 1.0e-6) ) if duration_change < min_duration_change: min_duration_change = duration_change min_duration_factor = i bpm /= min_duration_factor timemap = [] resample_sr = 48000 score = 0 last_orig_time = 0 for i, orig_time in enumerate(adjusted_beats): orig_dt = orig_time - last_orig_time last_orig_time = orig_time percent_change_this_beat = abs(1 - orig_dt / (60 / bpm)) score = max(score, percent_change_this_beat) stretch_to_time = (i + 1) * 60 / bpm timemap.append( (int(orig_time * resample_sr), int(stretch_to_time * resample_sr)) ) with open(f"{temp_dir}/timemap.txt", "w") as f: for a, b in timemap: f.write(f"{a} {b}\n") ts_in_path = audio_write( f"{temp_dir}/ts_in", cut_wav, sr, strategy="loudness", loudness_compressor=True, format="wav", ) if sr != resample_sr: output_path = f"{temp_dir}/resampled.wav" resample_audio(ts_in_path, output_path, resample_sr) ts_in_path = output_path op_prog(2) subprocess.run( [ "rubberband", "--fine", "--centre-focus", "--duration", str(float(len(adjusted_beats) * 60 / bpm)), "--timemap", f"{temp_dir}/timemap.txt", ts_in_path, f"{temp_dir}/ts_out.wav", ], check=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, ) op_prog(3) return *torchaudio.load(f"{temp_dir}/ts_out.wav"), score def _put_audio_to_s3(self, wav, sr, use_compressor, op_prog=None): with tempfile.TemporaryDirectory() as temp_dir: output_path = audio_write( f"{temp_dir}/output", wav.cpu(), sr, strategy="loudness", loudness_compressor=use_compressor, format="wav", ) codec = "ogg" s3_key = f"{uuid.uuid4()}.{codec}" upload_audio(self.s3, self.s3_bucket, s3_key, output_path, codec, op_prog) return s3_key def _grid_audio_and_upload(self, wav, sr, bpm, op_prog): wav, sr, score = self._grid_audio(wav, sr, bpm, op_prog) s3_key = self._put_audio_to_s3(wav, sr, use_compressor=False) return s3_key, score def close(self): self.beat_processor.close() def handle_request(self, request_data, recorder, notify_progress, times_out_at): t0 = time.time() print("request: ", request_data) want_duration = request_data.get("duration", 10) op_set = OpProgressSet(notify_progress, debounce=0.5) op_generate = OpProgress(100, 50.0) op_set.add(op_generate) op_encode = OpProgressFfmpeg(want_duration, CODEC_SPEEDS.get("ogg", 22.0)) op_set.add(op_encode) audio_prompt_key = request_data.get("audioPrompt", None) if audio_prompt_key is not None: audio_prompt, audio_prompt_sr = self._get_audio_from_s3(audio_prompt_key) else: audio_prompt, audio_prompt_sr = None, None self.model.set_generation_params(duration=want_duration) def progress_callback(a, b): op_generate.total = b op_generate(a) self.model.set_custom_progress_callback(progress_callback) prompt = request_data.get("textPrompt", None) target_bpm = request_data.get("bpm", None) batch_size = 1 if target_bpm is None else 4 if target_bpm is not None: op_set_grid = OpProgressMinSet() op_set.add(op_set_grid) op_grids = [] for _ in range(batch_size): op_grids.append(OpProgress(3, 3.0 / want_duration)) op_set_grid.add(op_grids[-1]) if prompt is None and audio_prompt is None: wavs = self.model.generate_unconditional(batch_size, progress=True) elif audio_prompt is None: wavs = self.model.generate([prompt] * batch_size, progress=True) else: wavs = self.model.generate_with_chroma( [prompt if prompt is not None else ""] * batch_size, audio_prompt[None].expand(batch_size, -1, -1), audio_prompt_sr, progress=True, ) if target_bpm is not None: results = [] with concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor: for fut in [ executor.submit( self._grid_audio_and_upload, wavs[i], self.model.sample_rate, float(target_bpm), op, ) for i, op in enumerate(op_grids) ]: try: results.append(fut.result()) except RetryableError: continue if len(results) == 0: raise RuntimeError("Sample generation failed") results = list(sorted(results, key=lambda x: x[1])) print(f"({time.time() - t0:.2f}s) -> ", results) return { "status": "success", "s3Key": results[0][0], "results": [{"s3Key": r[0], "score": r[1]} for r in results], } else: s3_key = self._put_audio_to_s3( wavs[0], self.model.sample_rate, use_compressor=True, op_prog=op_encode ) print(f"({time.time() - t0:.2f}s) -> ", s3_key) return {"status": "success", "s3Key": s3_key} if __name__ == "__main__": client = MusicGenRedisClient.from_env() # client.handle_request( # { # "textPrompt": "A happy tune", # "duration": 5, # "bpm": 120, # }, # None, # print, # time.time() + 60, # ) client.run() client.close()