import sys import torch import torchaudio from stable_audio_tools.models.conditioners import ( CodecConditioner, SemanticConditioner, MERTConditioner, DACConditioner, ) from suno_utils.tasks.mert_25 import ( preload_models as preload_semantic_models, encode as semantic_encode, ) _ = preload_semantic_models( checkpoint_filepath="s3://suno-data/georg/models/semantic/mert_25.pt", centroids_filepath="s3://suno-data/georg/models/semantic/mert_25_2x4k.npy", device="cuda", ) N = 262144 codec_cond = CodecConditioner(768).cuda() codec_codes = [ torch.randint(0, 2048, size=(250, 12), device="cuda"), torch.randint(0, 2048, size=(250, 12), device="cuda"), ] codec_embeds, _ = codec_cond(codec_codes) print(codec_codes[0].shape, codec_embeds.shape) sys.exit() mert_cond = MERTConditioner(768) # test by passing audio and internally generating codes cond_dicts = [ {"audio": torch.randn(2, N), "codes": None}, {"audio": torch.randn(2, N), "codes": None}, ] latents, _ = mert_cond(cond_dicts) print(latents.shape) # now test by manually passing codes audios = [torch.randn(2, N), torch.randn(2, N)] audios = [audio.mean(dim=0, keepdim=True) for audio in audios] audios = [torchaudio.functional.resample(audio, 48000, 24000) for audio in audios] codes = semantic_encode(audios) codes = [torch.from_numpy(code[:, 0]).long() for code in codes] print(codes) cond_dicts = [ {"audio": torch.randn(2, N), "codes": codes[0]}, {"audio": torch.randn(2, N), "codes": codes[1]}, ] latents, _ = mert_cond(cond_dicts) print(latents.shape) # codec print() print("dac") codec_ckpt = "/home/christian/christian/stable-audio-tools/checkpoints/dac_2c_25x12.pt" dac_cond = DACConditioner(768, 48000, codec_ckpt, codebook_dropout=True) cond_dicts = [ {"audio": torch.randn(2, N), "codes": None}, {"audio": torch.randn(2, N), "codes": None}, ] latents, _ = dac_cond(cond_dicts) print(latents.shape) # now test manually extracting codes and passing to conditioner audios = [torch.randn(2, N), torch.randn(2, N)] audios = torch.stack(audios) z, codes, latents, commitment_loss, codebook_loss = dac_cond.model.encode(audios, 12) print(codes.shape) # z = dac_cond.model.quantizer.from_codes(codes) # b, n, t # print(z.shape) # z = z.permute(0, 2, 1) cond_dicts = [ {"audio": torch.randn(2, N), "codes": codes[0]}, {"audio": torch.randn(2, N), "codes": codes[1]}, ] latents = dac_cond(cond_dicts)