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)
