import uuid import demucs.api import tempfile import boto3 from concurrent.futures import ThreadPoolExecutor from lib.redis_client import RedisClient, getenv import torch import torchaudio as ta from lib.audio_utils import download_audio, upload_audio, CODEC_SPEEDS from lib.progress import ( OpProgress, OpProgressSet, OpProgressMinSet, OpProgressFfmpeg, ) SAMPLE_RATE = 44100 # only rate that demucs will ever support class DemucsRedisClient(RedisClient): def __init__(self, model, s3_bucket, *args): super().__init__(*args) self.separator = demucs.api.Separator(model=model) self.s3 = boto3.client("s3") self.s3_bucket = s3_bucket @classmethod def from_env(cls): return super().from_env(getenv("DEMUCS_MODEL"), getenv("DEMUCS_S3_BUCKET")) def warmup(self): self.separator.separate_tensor(torch.randn(2, SAMPLE_RATE, dtype=torch.float32)) def _separate_s3(self, s3_key, notify_progress, codec, duration): duration = max(0.1, float(duration)) if duration is not None else 10 op_set = OpProgressSet(notify_progress, debounce=0.5) decode_op = OpProgressFfmpeg(duration, 300.0) op_set.add(decode_op) demucs_op = OpProgress(duration, 22.0) op_set.add(demucs_op) encode_op_set = OpProgressMinSet() op_set.add(encode_op_set) encode_ops = {} for stem in self.separator._model.sources: encode_op = OpProgressFfmpeg(duration, CODEC_SPEEDS.get(codec, 22.0)) encode_op_set.add(encode_op) encode_ops[stem] = encode_op with tempfile.TemporaryDirectory() as temp_dir, ThreadPoolExecutor( max_workers=4 ) as executor: input_path = download_audio( self.s3, self.s3_bucket, s3_key, temp_dir, progress_handler=decode_op, sample_rate=SAMPLE_RATE, ) decode_op.finish() def sep_notify(d): if d["state"] != "start": return true_duration = max(0.1, d["audio_length"] / SAMPLE_RATE) demucs_op.total = true_duration encode_op_set.total = true_duration demucs_op(d["segment_offset"] / SAMPLE_RATE) self.separator.update_parameter(callback=sep_notify) audio_tensor, loaded_sample_rate = ta.load(input_path) assert loaded_sample_rate == SAMPLE_RATE was_mono = False # demucs crashes on mono audio - convert to stereo if audio_tensor.shape[0] == 1: was_mono = True audio_tensor = audio_tensor.repeat(2, 1).clone() _, separated = self.separator.separate_tensor( audio_tensor, self.separator.samplerate ) demucs_op.finish() futures = [] res = {} for stem, wav in separated.items(): if was_mono: wav = wav.mean(0, keepdim=True) output_path = f"{temp_dir}/{stem}.wav" demucs.api.save_audio( wav, output_path, samplerate=self.separator.samplerate ) s3_key = f"{uuid.uuid4()}.{codec}" res[stem] = s3_key futures.append( executor.submit( upload_audio, self.s3, self.s3_bucket, s3_key, output_path, codec, progress_handler=encode_ops[stem], ) ) for future in futures: future.result() notify_progress(100) return res def handle_request(self, request_data, recorder, notify_progress, times_out_at): print("request: ", request_data) stem_dict = self._separate_s3( request_data["s3Key"], notify_progress, request_data.get("codec", "ogg"), request_data.get("duration", None), ) return {"status": "success", "stems": stem_dict} if __name__ == "__main__": client = DemucsRedisClient.from_env() client.run()