from contextlib import contextmanager
import math
import os
import random
import tempfile
import torch

import numpy as np
import tqdm
from tokenizers import Tokenizer

from suno_utils.audio import Audio
from suno_utils.tasks.mert_25 import (
    preload_models as preload_semantic_models,
    encode as encode_semantic,
    EMBEDDING_RATE as SEMANTIC_HZ,
)
from suno_utils.utils.s3 import _download_s3_file, get_s3_checksum
from suno_utils.utils.text import get_filename

from dataset import unfold_tensor
from helpers import FULL_PRECISION_KEY_FRAGMENTS, load_checkpoint, simplify_whitespace
from sampling import (
    VDenoiser,
    get_sigmas_polyexponential,
    sample_dpmpp_3m_sde,
    sample_ddim,
    sample_flow_dpmpp,
    sample_discrete_euler,
    sample_rk4,
    sample_flow_pingpong,
    TimeShift,
)


# hold models in global scope to lazy load
global models
models = {}
ModelClass = None
g_model_type = None

codec_decode = None
codec_encode = None


TOKENIZER_FILEPATH = "s3://suno-data/georg/models/tokenizers/tokenizer_60k.json"
SEMANTIC_MODEL_FILEPATH = "s3://suno-data/georg/models/semantic/mert_25.pt"
SEMANTIC_CLUSTERS_FILEPATH = "s3://suno-data/georg/models/semantic/mert_25_2x4k.npy"
CODEC_FILEPATH = "s3://suno-data/christian/25hz_vae_peaq_kl_0.005.pth"
DIT_MODEL_FILEPATH = "s3://suno-data/georg/tmp/2b_prefix_1000k.pt"


os.environ["TOKENIZERS_PARALLELISM"] = "False"


def ratio_mask_semantic_codes(
    semantic_codes_chunk: torch.Tensor,
    semantic_mask_ratio: float = 0.0,
    cond_semantic_n_vocab: int = 4001,
) -> torch.Tensor:
    """
    Mask the semantic codes in the chunk by the given ratio.
    """
    if not 0 <= semantic_mask_ratio <= 1:
        raise ValueError("semantic_mask_ratio must be 0 and 1.")

    if not isinstance(semantic_codes_chunk, torch.Tensor):
        semantic_codes_chunk = torch.tensor(semantic_codes_chunk)

    if semantic_mask_ratio == 0:
        return semantic_codes_chunk.clone()

    if semantic_codes_chunk.shape[1] == 0:
        return semantic_codes_chunk

    step = 1 / semantic_mask_ratio
    indices = torch.arange(0, semantic_codes_chunk.shape[1], step).long()
    indices = indices[indices < semantic_codes_chunk.shape[1]]  # Ensure indices are within bounds
    indices = torch.unique(indices)
    masked = semantic_codes_chunk.clone()
    masked[:, indices] = cond_semantic_n_vocab - 1
    return masked


def get_model_if_needed(model_path, cache_dir):
    if not model_path.startswith("s3:"):
        return model_path
    # remote model, fetch if needed
    checksum = get_s3_checksum(model_path)

    # get temp path
    cache_path = os.path.join(cache_dir, f"{checksum}.pt")
    if not os.path.exists(cache_path):
        parent_dir = os.path.dirname(cache_path)
        os.makedirs(parent_dir, exist_ok=True)
        print(f"Downloading model `{get_filename(model_path, keep_ext=True)}`...")
        _download_s3_file(model_path, cache_path)
    return cache_path


@contextmanager
def _download_from_s3_if_needed(maybe_s3_filepath):
    tmp_filepath = maybe_s3_filepath
    if maybe_s3_filepath.startswith("s3://"):
        temp_dir = tempfile.TemporaryDirectory()
        filename = get_filename(maybe_s3_filepath, keep_ext=True)
        tmp_filepath = os.path.join(temp_dir.name, filename)
        _download_s3_file(maybe_s3_filepath, tmp_filepath)
    yield tmp_filepath


def convert_to_precision(module, weights_precision=torch.float16):
    for p_name, param in module.named_parameters():
        if not any(s in p_name for s in FULL_PRECISION_KEY_FRAGMENTS):
            param.data = param.data.to(weights_precision)


def _load_dit_model(dit_model_filepath, use_ema_if_exists, weights_precision, compile=False):
    with _download_from_s3_if_needed(dit_model_filepath) as tmp_fp:
        state_dict, dit_config = load_checkpoint(tmp_fp, use_ema_if_exists=use_ema_if_exists)
    # assert (
    #    dit_config["block_size"] // dit_config["io_hz"] == dit_config["cond_semantic_len"] // SEMANTIC_HZ
    # )
    # assert (
    #    get_embedding_rate() == dit_config["io_hz"]
    # ), f"codec mismatch: {get_embedding_rate()}hz vs {dit_config['io_hz']}hz"
    # assert (
    #    models["codec_model"].latent_dim == dit_config["io_channels"]
    # ), f"codec mismatch: {models['codec_model'].latent_dim} vs {dit_config['io_channels']} channels
    # print(dit_config)
    dit_model = ModelClass(
        io_hz=dit_config["io_hz"],
        io_channels=dit_config["io_channels"],
        embed_dim=dit_config["embed_dim"],
        depth=dit_config["depth"],
        n_heads=dit_config["n_heads"],
        qk_norm=dit_config["qk_norm"],
        block_size=dit_config["block_size"],
        cond_semantic_n_vocab=dit_config["cond_semantic_n_vocab"],
        cond_semantic_len=dit_config["cond_semantic_len"],
        cond_text_n_vocab=dit_config["cond_text_n_vocab"],
        cond_text_len=dit_config["cond_text_len"],
        ctx_len=dit_config.get("ctx_len", 0),
        infill_ctx_len=dit_config.get("infill_ctx_len", 0),
        stem_ctx_len=dit_config.get("stem_ctx_len", 0),
        shared_ctx=dit_config.get("shared_ctx", False),
        use_rvq=dit_config.get("use_rvq", False),
        n_codebooks=dit_config.get("n_codebooks", 1),
    )
    dit_model.eval()
    # print("loading weights...")
    non_essential_keys = [
        "semantic_conditioner.pos_embedding.inv_freq",
        "text_conditioner.pos_embedding.inv_freq",
        "ctx_conditioner.pos_embedding.inv_freq",
    ]
    for k in non_essential_keys:
        state_dict.pop(k, None)

    if (dit_model.infill_ctx_len == 0 or dit_model.infill_ctx_len is None) and not dit_model.shared_ctx:
        print("warning: removing infill keys from state dict...")
        # remove infill_ctx_conditioner if it's not used
        infill_keys = [
            "infill_ctx_conditioner.proj_out.bias",
            "infill_ctx_conditioner.proj_out.weight",
            "vae_infill_pad_embed",
        ]
        for k in infill_keys:
            state_dict.pop(k, None)

    # TODO: gross hack, can be replaced with strict=True when we have new models
    model_sd_keys = set(dit_model.state_dict().keys())
    checkpoint_sd_keys = set(state_dict.keys())
    extra_keys = checkpoint_sd_keys - model_sd_keys
    missing_keys = model_sd_keys - checkpoint_sd_keys - set(non_essential_keys)
    assert len(extra_keys) == 0, f"extra keys in state dict: {extra_keys}"
    assert len(missing_keys) == 0, f"missing keys not in state dict: {missing_keys}"
    dit_model.load_state_dict(state_dict, strict=False)
    # print(f"converting model to precision {weights_precision}...")
    convert_to_precision(dit_model, weights_precision=weights_precision)

    # Only move to CUDA and compile if compile=True
    if compile:
        dit_model.to("cuda")
        dit_model = torch.compile(dit_model)

    if dit_model.infill_ctx_len is None:
        dit_model.infill_ctx_len = 0
    return dit_model, dit_config


def preload_models(
    tokenizer_filepath=TOKENIZER_FILEPATH,
    semantic_model_filepath=SEMANTIC_MODEL_FILEPATH,
    semantic_clusters_filepath=SEMANTIC_CLUSTERS_FILEPATH,
    codec_filepath=CODEC_FILEPATH,
    dit_model_filepath=DIT_MODEL_FILEPATH,
    dit_model_2_filepath=None,
    use_ema_if_exists=True,
    model_type="prefix",
    codec_scale_factor=2.5,
    patch_size=1,
    compile=True,
    weights_precision=torch.bfloat16,
):
    if model_type == "default":
        from model import DiffusionTransformer
    elif model_type == "prefix":
        from prefix_model.model import DiffusionTransformer
    else:
        raise NotImplementedError()
    global g_model_type
    g_model_type = model_type
    global ModelClass
    ModelClass = DiffusionTransformer
    global models
    # load semantic model
    print("loading semantic model...")
    _ = preload_semantic_models(
        checkpoint_filepath=semantic_model_filepath,
        centroids_filepath=semantic_clusters_filepath,
        device="cuda",
    )
    # load vae
    print("loading codec model...")

    if "dac_2c_25x12" in codec_filepath:
        from suno_utils.tasks.dac_2c import (
            preload_models as preload_codec_models,
            decode,
            encode,
            get_embedding_rate,
            load_model as load_codec_model,
        )
    elif "dac_vae_fixed_25hz" in codec_filepath or "dac_vae_tuned_25hz" in codec_filepath:
        from suno_utils.tasks.dac_vae_fixed_25hz import (
            preload_models as preload_codec_models,
            decode,
            encode,
            get_embedding_rate,
            load_model as load_codec_model,
        )
    else:
        from suno_utils.tasks.dac_vae_100hz_peaq import (  # NOTE: works for 25hz as well
            preload_models as preload_codec_models,
            decode,
            encode,
            get_embedding_rate,
            load_model as load_codec_model,
        )

    global codec_decode
    global codec_encode
    codec_decode = decode
    codec_encode = encode

    _ = preload_codec_models(checkpoint_filepath=codec_filepath, device="cuda")
    models["codec_model"] = load_codec_model()
    models["codec_scale_factor"] = codec_scale_factor
    models["patch_size"] = patch_size
    # load main diffusion model
    print("loading diffusion model...")
    dit_model, dit_config = _load_dit_model(
        dit_model_filepath, use_ema_if_exists, weights_precision, compile
    )
    models["dit_model"] = dit_model
    # load tokenizer
    print("loading tokenizer...")
    with _download_from_s3_if_needed(tokenizer_filepath) as tmp_fp:
        tokenizer = Tokenizer.from_file(tmp_fp)
    if dit_config["cond_text_n_vocab"] == 60005:
        # TODO: remove this legacy config
        tokenizer.add_special_tokens(["<tags>", "</tags>", "<lyrics>", "</lyrics>", "\n"])
    elif dit_config["cond_text_n_vocab"] == 60001:
        tokenizer.add_special_tokens(["\n"])
    else:
        raise NotImplementedError()
    tokenizer.pad_idx = tokenizer.token_to_id("[PAD]")
    models["tokenizer"] = tokenizer
    models["info"] = {
        "weights_precision": weights_precision,
    }
    print("done!")


def _retrieve_models():
    global models
    return models


def generate(
    audio,
    lyrics="",
    aligned_lyrics=None,
    tags="",
    text_cfg_coef=2.5,
    tag_cfg_coef=1.0,
    sem_cfg_coef=1.0,
    ctx_cfg_coef=1.0,
    model_cfg_coef=1.0,
    infill_ctx_cfg_coef=1.0,
    stem_ctx_cfg_coef=1.0,
    steps=16,
    sampling_method="dpmp",  #  dpmp|ddim|rk4|euler|pingpong
    sampler_type="dpmp",  # dpmp|rk4|euler|pingpong
    objective="v",  # v|rf_denoiser|rectified_flow
    oracle_first_chunk=False,
    init_latents=None,
    sigma_min=0.5,
    sigma_max=50,
    rho=1.0,
    time_shift=0.0,  # Time shift for rectified flow models (0=no shift, 1-3=more detail)
    seed=None,
    lyrics_idx_pad=2,
    semantic_skip_factor=1,
    semantic_noise=None,
    use_repeated_sem_pad=False,
    truncate_at_last_chunk=False,
    downscale_ctx_vector=True,
    use_ctx_vector=True,
    ctx_vector_scale=1.0,
    codec_cycle_ctx_vector=False,
    noise_ctx_vector=0.0,
    noise_ctx_vector_pad_size=0,
    history_ctx_vae=None,
    infill_ctx_vae=None,
    infill_ctx_mask=None,
    infill_variation=False,
    stem_ctx_vae=None,
    stem_ctx_mask=None,
    chunk_size=None,
    sampler_eta=1.0,
    sampler_s_noise=1.0,
    ctx_len_used=None,
    ctx_steps_used=None,
    return_latents=False,
    use_codec=True,
    distilled=False,
    normalize_volume=True,
):
    """
    Generate audio from a given audio or semantic codes.
    Args:
        audio (Audio or torch.Tensor): Audio to generate from, if a tensor, it is assumed to be semantic codes.
        lyrics (str): Lyrics to use for text conditioning.
        aligned_lyrics (list[dict]): Aligned lyrics to use for text conditioning.
        tags (str): Tags to use for text conditioning.
        text_cfg_coef (float): Text conditioning cfg scale.
        tag_cfg_coef (float): Tag conditioning cfg scale.
        sem_cfg_coef (float): Semantic conditioning cfg scale.
        ctx_cfg_coef (float): Context conditioning cfg scale.
        steps (int): Number of sampling steps.
        sampling_method (str): Sampling method to use.
        init_latents (torch.Tensor): Initial latents to use for sampling.
        sigma_max (float): Maximum sigma value for sampling.
        seed (int): Random seed for reproducibility.
        lyrics_idx_pad (int): Number of lyrics indices to pad on left and right.
        downscale_ctx_vector (bool): Whether to downscale the context vector by the codec scale factor.
        noise_ctx_vector (float): Amount of noise to add to the context vector.
        vae_ctx_vector (torch.Tensor): Context vector to use for VAE conditioning.
        vae_ctx_mask (torch.Tensor): Mask to use for VAE conditioning.
        return_latents (bool): Whether to return the latents.
    Returns:
        Audio: Generated audio.
    """
    global models, g_model_type
    if seed is not None:
        torch.manual_seed(seed)

    if ctx_steps_used is None:
        ctx_steps_used = steps

    # simplify inputs
    assert sampling_method in ["dpmp", "ddim"]
    if steps == 1:
        sampling_method = "ddim"
    if aligned_lyrics is not None:
        assert len(lyrics) == 0
        assert isinstance(aligned_lyrics, list)
        if len(aligned_lyrics) > 0:
            assert isinstance(aligned_lyrics[0], dict)
        aligned_lyrics = [m for m in aligned_lyrics if "word" in m]
    tags = simplify_whitespace(tags.replace("[", " ").replace("]", " "), retain_newlines=False)
    lyrics = simplify_whitespace(lyrics, retain_newlines=True)
    # semantic encoding
    codec_codes = None
    if isinstance(audio, Audio):
        original_audio_len = audio.duration_s
        if normalize_volume:
            audio = audio.normalize_volume()
        if models["dit_model"].ctx_len is None:
            audio = audio.get_segment(
                to_s=int(round(models["dit_model"].block_size / models["dit_model"].io_hz)) + 0.01
            )
        # get semantic conditioning
        semantic_codes = encode_semantic(
            audio.convert(sample_rate=24_000, byte_width=2, n_channels=1)
        ).astype(np.int64)[:, 0]
        semantic_codes = torch.from_numpy(semantic_codes).long()
        print(f"semantic_codes: {semantic_codes.shape}")
        if oracle_first_chunk:
            codec_codes = codec_encode(audio.convert(sample_rate=48_000, byte_width=2, n_channels=2))
        if models["dit_model"].stem_ctx_len > 0:
            stem_ctx_vae = codec_encode(audio.convert(sample_rate=48_000, byte_width=2, n_channels=2))
            stem_ctx_vae = torch.from_numpy(stem_ctx_vae).float()[None]
            semantic_codes[:] = models["dit_model"].cond_semantic_n_vocab - 1
            # assert semantic_codes.shape[0] == stem_ctx_vae.shape[1], (
            #     f"{semantic_codes.shape[0]} != {stem_ctx_vae.shape[1]}"
            # )
    elif isinstance(audio, torch.Tensor):
        semantic_codes = audio
    else:
        raise NotImplementedError()
    if len(semantic_codes.shape) == 1:
        semantic_codes = semantic_codes[None]
    assert semantic_codes.shape[0] == 1  # single channel

    if chunk_size is None:
        chunk_window_size = models["dit_model"].block_size
    else:
        chunk_window_size = chunk_size

    if ctx_len_used is None:
        ctx_len_used = models["dit_model"].ctx_len
    else:
        ctx_len_used = int(ctx_len_used)

    if models["dit_model"].ctx_len is None:
        n_chunks = 1
    else:
        n_chunks = int(
            math.ceil(
                semantic_codes.shape[1] / SEMANTIC_HZ * models["dit_model"].io_hz / chunk_window_size
            )
        )

    out_pred_z = torch.zeros(
        (semantic_codes.shape[0], n_chunks * chunk_window_size, models["dit_model"].io_channels),
        dtype=torch.float32,
    )
    for n_chunk in tqdm.tqdm(range(n_chunks), disable=n_chunks == 1):
        if oracle_first_chunk and n_chunk == 0:
            out_pred_z[:, :chunk_window_size] = (
                torch.from_numpy(codec_codes[:chunk_window_size])[None]
                .repeat(out_pred_z.shape[0], 1, 1)
                .float()
                * models["codec_scale_factor"]
            )
            continue
        # text conditioning
        lyrics_chunk = lyrics
        if models["dit_model"].ctx_len is not None and aligned_lyrics is not None:
            # use roughly aligned text if available
            chunk_start_s = n_chunk * chunk_window_size / models["dit_model"].io_hz
            chunk_end_s = (n_chunk + 1) * chunk_window_size / models["dit_model"].io_hz

            chunk_start_indices = [
                idx
                for idx, m in enumerate(aligned_lyrics)
                if "start_s" in m and m["start_s"] >= chunk_start_s
            ]
            chunk_end_indices = [
                idx for idx, m in enumerate(aligned_lyrics) if "end_s" in m and m["end_s"] <= chunk_end_s
            ]

            if len(chunk_start_indices) > 0 and len(chunk_end_indices) > 0:
                chunk_start_idx = max(0, chunk_start_indices[0] - lyrics_idx_pad)
                chunk_end_idx = min(len(aligned_lyrics), chunk_end_indices[-1] + lyrics_idx_pad)
                lyrics_chunk = "".join(
                    [m["word"] for m in aligned_lyrics[chunk_start_idx:chunk_end_idx]]
                )
                lyrics_chunk = simplify_whitespace(lyrics_chunk, retain_newlines=True)
            else:
                lyrics_chunk = ""

        use_legacy_format = models["dit_model"].cond_text_n_vocab == 60005
        if use_legacy_format:
            text = f"<tags>{tags}</tags><lyrics>{lyrics_chunk}</lyrics>"
            empty_text = "<tags></tags><lyrics></lyrics>"
            empty_tag = "<tags></tags><lyrics>{lyrics_chunk}</lyrics>"
        else:
            text_pieces = []
            if len(tags) > 0:
                text_pieces.append(f"[{tags}]")
            if len(lyrics_chunk) > 0:
                text_pieces.append(lyrics_chunk)
            text = "\n\n".join(text_pieces)
            empty_text = ""
            empty_tag = lyrics_chunk
        # text conditioning
        text_codes = torch.full(
            (1, models["dit_model"].cond_text_len), models["tokenizer"].pad_idx, dtype=torch.long
        )
        for n, codes_row in enumerate(models["tokenizer"].encode_batch([text])):
            codes_row = torch.tensor(
                codes_row.ids[: models["dit_model"].cond_text_len], dtype=torch.long
            )
            text_codes[n, : len(codes_row)] = codes_row
        empty_text_codes = torch.full(
            (1, models["dit_model"].cond_text_len), models["tokenizer"].pad_idx, dtype=torch.long
        )
        for n, codes_row in enumerate(models["tokenizer"].encode_batch([empty_text])):
            codes_row = torch.tensor(
                codes_row.ids[: models["dit_model"].cond_text_len], dtype=torch.long
            )
            empty_text_codes[n, : len(codes_row)] = codes_row
        empty_tag_codes = torch.full(
            (1, models["dit_model"].cond_text_len), models["tokenizer"].pad_idx, dtype=torch.long
        )
        for n, codes_row in enumerate(models["tokenizer"].encode_batch([empty_tag])):
            codes_row = torch.tensor(
                codes_row.ids[: models["dit_model"].cond_text_len], dtype=torch.long
            )
            empty_tag_codes[n, : len(codes_row)] = codes_row
        text_codes = text_codes.cuda()
        empty_text_codes = empty_text_codes.cuda()
        empty_tag_codes = empty_tag_codes.cuda()
        with torch.no_grad(), torch.autocast("cuda", dtype=models["info"]["weights_precision"]):
            text_condition = models["dit_model"].text_conditioner(text_codes)
            empty_text_condition = models["dit_model"].text_conditioner(empty_text_codes)
            empty_tag_condition = models["dit_model"].text_conditioner(empty_tag_codes)
        # semantic conditioning
        start_idx = n_chunk * chunk_window_size
        end_idx = (n_chunk + 1) * chunk_window_size
        semantic_codes_chunk = semantic_codes[:, start_idx:end_idx]
        if models["dit_model"].stem_ctx_len > 0:
            stem_ctx_chunk = stem_ctx_vae[:, start_idx:end_idx]
            # pad if necessary (use repeated segment)
            if stem_ctx_chunk.shape[1] < models["dit_model"].stem_ctx_len:
                stem_ctx_chunk = torch.concat(
                    [stem_ctx_chunk]
                    * int(math.ceil(models["dit_model"].stem_ctx_len / stem_ctx_chunk.shape[1])),
                    dim=1,
                )[:, : models["dit_model"].stem_ctx_len]
        no_semantic_codes_chunk = torch.full(
            (1, models["dit_model"].cond_semantic_len),
            models["dit_model"].cond_semantic_n_vocab - 1,
            dtype=torch.long,
        )

        if semantic_codes_chunk.shape[1] < models["dit_model"].cond_semantic_len:
            if use_repeated_sem_pad:
                # pad if necessary (use repeated segment)
                semantic_codes_chunk = torch.concat(
                    [semantic_codes_chunk]
                    * int(
                        math.ceil(models["dit_model"].cond_semantic_len / semantic_codes_chunk.shape[1])
                    ),
                    dim=1,
                )[:, : models["dit_model"].cond_semantic_len]
            else:
                # pad if necessary (use repeated pad token)
                semantic_codes_chunk = torch.concat(
                    [
                        semantic_codes_chunk,
                        torch.full(
                            (1, models["dit_model"].cond_semantic_len - semantic_codes_chunk.shape[1]),
                            models["dit_model"].cond_semantic_n_vocab - 1,
                            dtype=torch.long,
                        ),
                    ],
                    dim=1,
                )
        semantic_codes_chunk = semantic_codes_chunk.cuda()
        no_semantic_codes_chunk = no_semantic_codes_chunk.cuda()

        # semantic skip if available
        semantic_skip_factor = int(semantic_skip_factor)
        if semantic_skip_factor > 1:
            n_phase = random.randint(0, semantic_skip_factor - 1)
            for nn in range(semantic_skip_factor - 1):
                shifted_idx = (nn + n_phase) % semantic_skip_factor
                semantic_codes_chunk[:, shifted_idx::semantic_skip_factor] = (
                    models["dit_model"].cond_semantic_n_vocab - 1
                )

        with torch.no_grad(), torch.autocast("cuda", dtype=models["info"]["weights_precision"]):
            if models["dit_model"].cond_semantic_len > 0:
                sem_condition = models["dit_model"].semantic_conditioner(
                    semantic_codes_chunk, noise=semantic_noise
                )
                no_sem_condition = models["dit_model"].semantic_conditioner(
                    no_semantic_codes_chunk, noise=semantic_noise
                )
                # combine and add context vector if needed
                cross_attn_cond = torch.concat([text_condition, sem_condition], dim=1)
                no_text_cross_attn_cond = torch.concat([empty_text_condition, sem_condition], dim=1)
                no_tag_cross_attn_cond = torch.concat([empty_tag_condition, sem_condition], dim=1)
                no_sem_cross_attn_cond = torch.concat([text_condition, no_sem_condition], dim=1)
                no_ctx_cross_attn_cond = None
            else:
                cross_attn_cond = text_condition
                no_text_cross_attn_cond = empty_text_condition
                no_tag_cross_attn_cond = empty_tag_condition
                no_sem_cross_attn_cond = None
                no_ctx_cross_attn_cond = None

            if models["dit_model"].ctx_len > 0:
                ctx_len = models["dit_model"].ctx_len
                if models["dit_model"].shared_ctx:
                    if infill_ctx_vae is not None:
                        print("doing infill...")
                        # apply the mask to the infill_ctx_vector
                        expanded_vae_pad = (
                            models["dit_model"]
                            .vae_infill_pad_embed.reshape(1, 1, -1)
                            .repeat(1, ctx_len, 1)
                        )
                        # add the infill_ctx_vector to the vae_pad_embed
                        infill_ctx_vector = infill_ctx_vae + expanded_vae_pad

                        if noise_ctx_vector > 0.0:
                            # if n_chunk != n_chunks - 1 and :
                            # don't add noise to the last chunk
                            # add noise to the centre of the chunk
                            ps = noise_ctx_vector_pad_size  # latents on the right are un-noised
                            if ps > 0:
                                infill_ctx_vector[:, :-ps, :] = infill_ctx_vector[:, :-ps, :] + (
                                    torch.randn_like(infill_ctx_vector[:, :-ps, :]) * noise_ctx_vector
                                )
                            else:
                                ctx_vector = infill_ctx_vector + (
                                    torch.randn_like(infill_ctx_vector) * noise_ctx_vector
                                )

                        # pass through ctx_conditioner
                        ctx_condition = models["dit_model"].infill_ctx_conditioner(infill_ctx_vector)

                    elif n_chunk == 0 or not use_ctx_vector:
                        expanded_default_embed = (
                            models["dit_model"].vae_default_embed.reshape(1, 1, -1).repeat(1, ctx_len, 1)
                        )
                        # no need to pass through ctx_conditioner
                        ctx_condition = expanded_default_embed.cuda()
                    else:
                        expanded_vae_pad = (
                            models["dit_model"].vae_pad_embed.reshape(1, 1, -1).repeat(1, ctx_len, 1)
                        )
                        ctx_vector = torch.zeros(1, ctx_len, models["dit_model"].io_channels).cuda()
                        # copy context from out_pred_z from the end
                        end_read_idx = n_chunk * chunk_window_size
                        start_read_idx = end_read_idx - ctx_len_used

                        if start_read_idx < 0:
                            start_read_idx = 0

                        chunk_to_write = out_pred_z[:, start_read_idx:end_read_idx, :]
                        chunk_to_write_size = chunk_to_write.shape[1]
                        # write the chunk to the ctx_vector
                        ctx_vector[:, ctx_len - chunk_to_write_size :, :] = chunk_to_write

                        if noise_ctx_vector > 0.0:
                            # if n_chunk != n_chunks - 1 and :
                            # don't add noise to the last chunk
                            # add noise to the centre of the chunk
                            ps = noise_ctx_vector_pad_size  # latents on the right are un-noised
                            if ps > 0:
                                ctx_vector[:, :-ps, :] = ctx_vector[:, :-ps, :] + (
                                    torch.randn_like(ctx_vector[:, :-ps, :]) * noise_ctx_vector
                                )
                            else:
                                ctx_vector = ctx_vector + (
                                    torch.randn_like(ctx_vector) * noise_ctx_vector
                                )

                        # add the ctx_vector to the vae_pad_embed
                        ctx_vector = ctx_vector + expanded_vae_pad

                        # pass through ctx_conditioner
                        ctx_vector = ctx_vector.cuda()
                        ctx_condition = models["dit_model"].ctx_conditioner(ctx_vector)

                    cross_attn_cond = torch.concat([cross_attn_cond, ctx_condition], dim=1)
                    no_text_cross_attn_cond = torch.concat(
                        [no_text_cross_attn_cond, ctx_condition], dim=1
                    )
                else:
                    if (n_chunk == 0 or not use_ctx_vector) and ctx_len > 0 and history_ctx_vae is None:
                        ctx_vector = (
                            models["dit_model"].vae_pad_embed.reshape(1, 1, -1).repeat(1, ctx_len, 1)
                        )
                    elif ctx_len > 0:
                        if history_ctx_vae is not None:
                            ctx_vector = history_ctx_vae
                        else:
                            # start with a fully pad ctx_vae
                            ctx_vector = (
                                models["dit_model"].vae_pad_embed.reshape(1, 1, -1).repeat(1, ctx_len, 1)
                            )
                            # copy context from out_pred_z from the end
                            end_read_idx = n_chunk * chunk_window_size
                            start_read_idx = end_read_idx - ctx_len_used

                            if start_read_idx < 0:
                                start_read_idx = 0

                            chunk_to_write = out_pred_z[:, start_read_idx:end_read_idx, :]
                            chunk_to_write_size = chunk_to_write.shape[1]
                            # write the chunk to the ctx_vector
                            ctx_vector[:, ctx_len - chunk_to_write_size :, :] = chunk_to_write

                        if ctx_vector_scale != 1.0:
                            print("scaling ctx_vector by", ctx_vector_scale)
                            ctx_vector *= ctx_vector_scale

                        if downscale_ctx_vector:
                            ctx_vector = ctx_vector / models["codec_scale_factor"]

                        if noise_ctx_vector > 0.0:
                            # if n_chunk != n_chunks - 1 and :
                            # don't add noise to the last chunk
                            # add noise to the centre of the chunk
                            ps = noise_ctx_vector_pad_size  # latents on the right are un-noised
                            if ps > 0:
                                ctx_vector[:, :-ps, :] = ctx_vector[:, :-ps, :] + (
                                    torch.randn_like(ctx_vector[:, :-ps, :]) * noise_ctx_vector
                                )
                            else:
                                ctx_vector = ctx_vector + (
                                    torch.randn_like(ctx_vector) * noise_ctx_vector
                                )

                    ctx_vector = ctx_vector.cuda()

                    ctx_condition = models["dit_model"].ctx_conditioner(ctx_vector)
                    empty_ctx_condition = models["dit_model"].ctx_conditioner(
                        models["dit_model"]
                        .vae_pad_embed.reshape(1, 1, -1)
                        .repeat(1, models["dit_model"].ctx_len, 1)
                    )
                    cross_attn_cond = torch.concat([cross_attn_cond, ctx_condition], dim=1)
                    no_text_cross_attn_cond = torch.concat(
                        [no_text_cross_attn_cond, ctx_condition], dim=1
                    )

                    if models["dit_model"].cond_semantic_len > 0:
                        no_sem_cross_attn_cond = torch.concat(
                            [no_sem_cross_attn_cond, ctx_condition], dim=1
                        )  # not really used atm

            if models["dit_model"].infill_ctx_len > 0:
                infill_ctx_len = models["dit_model"].infill_ctx_len
                if infill_ctx_vae is None:
                    infill_ctx_vector = torch.zeros(
                        1, infill_ctx_len, models["dit_model"].io_channels
                    ).cuda()
                    infill_ctx_mask = torch.zeros(1, infill_ctx_len).bool().cuda()
                else:
                    infill_ctx_vector = infill_ctx_vae
                    infill_ctx_mask = infill_ctx_mask

                # apply mask to infill_ctx_vector
                expanded_vae_pad = (
                    models["dit_model"]
                    .vae_infill_pad_embed.reshape(1, 1, -1)
                    .repeat(1, infill_ctx_len, 1)
                )
                inverted_mask = (~infill_ctx_mask).float().unsqueeze(-1)
                # Use the mask to blend the original input with the pad embedding
                vae_input = (
                    infill_ctx_vector * infill_ctx_mask.float().unsqueeze(-1)
                    + expanded_vae_pad * inverted_mask
                )
                # move to cuda
                vae_input = vae_input.cuda()
                infill_ctx_condition = models["dit_model"].infill_ctx_conditioner(vae_input)
                empty_infill_ctx_condition = models["dit_model"].infill_ctx_conditioner(
                    models["dit_model"]
                    .vae_infill_pad_embed.reshape(1, 1, -1)
                    .repeat(1, infill_ctx_len, 1)
                )
                # add infill_ctx_condition to cross_attn_cond
                cross_attn_cond = torch.concat([cross_attn_cond, infill_ctx_condition], dim=1)
                no_text_cross_attn_cond = torch.concat(
                    [no_text_cross_attn_cond, infill_ctx_condition], dim=1
                )

            if models["dit_model"].stem_ctx_len > 0:
                stem_ctx_mask = torch.zeros(1, models["dit_model"].stem_ctx_len).bool().cuda()
                if stem_ctx_vae is None:
                    stem_ctx_vector = torch.zeros(
                        1, models["dit_model"].stem_ctx_len, models["dit_model"].io_channels
                    ).cuda()
                else:
                    stem_ctx_vector = stem_ctx_chunk * models["codec_scale_factor"]

                stem_ctx_vector = stem_ctx_vector.cuda()
                stem_ctx_mask = stem_ctx_mask.cuda()

                stem_ctx_condition = models["dit_model"].stem_ctx_conditioner(stem_ctx_vector)
                # add stem_ctx_condition to cross_attn_cond
                cross_attn_cond = torch.concat([cross_attn_cond, stem_ctx_condition], dim=1)
                no_text_cross_attn_cond = torch.concat(
                    [no_text_cross_attn_cond, stem_ctx_condition], dim=1
                )

        if g_model_type == "prefix":
            n_cfg = (
                int(text_cfg_coef != 1.0)
                + int(tag_cfg_coef != 1.0)
                + int(sem_cfg_coef != 1.0)
                + int(ctx_cfg_coef != 1.0)
                + int(infill_ctx_cfg_coef != 1.0)
                + int(stem_ctx_cfg_coef != 1.0)
            )
            models["dit_model"].transformer.setup_caches(
                1 + n_cfg,
                models["dit_model"].cond_text_len
                + models["dit_model"].cond_semantic_len
                + models["dit_model"].ctx_len
                + models["dit_model"].infill_ctx_len
                + models["dit_model"].stem_ctx_len,
            )
            if models.get("dit_model_2"):
                models["dit_model_2"].transformer.setup_caches(
                    1 + n_cfg,
                    models["dit_model_2"].cond_text_len
                    + models["dit_model_2"].cond_semantic_len
                    + models["dit_model_2"].ctx_len
                    + models["dit_model_2"].infill_ctx_len
                    + models["dit_model_2"].stem_ctx_len,
                )
        # do inference sampling
        x = torch.randn(
            [1, models["dit_model"].io_channels, models["dit_model"].block_size], device="cuda"
        )  # noise input
        if init_latents is not None:
            assert x.shape == init_latents.shape
            x += init_latents.cuda()

        extra_args = {
            "cross_attn_cond": cross_attn_cond,
            "no_text_cross_attn_cond": no_text_cross_attn_cond,
            # "no_tag_cross_attn_cond": no_tag_cross_attn_cond,
            # "no_sem_cross_attn_cond": no_sem_cross_attn_cond,
            # "no_ctx_cross_attn_cond": no_ctx_cross_attn_cond,
            # "no_infill_ctx_cross_attn_cond": no_infill_ctx_cross_attn_cond,
            "text_cfg_scale": text_cfg_coef,
            "tag_cfg_scale": tag_cfg_coef,
            "sem_cfg_scale": sem_cfg_coef,
            "ctx_cfg_scale": ctx_cfg_coef,
            "infill_ctx_cfg_scale": infill_ctx_cfg_coef,
            "stem_ctx_cfg_scale": stem_ctx_cfg_coef,
        }

        if objective == "v":
            denoiser = VDenoiser(models["dit_model"])
            # sigma_min = 0.5
            # rho = 1.0
            sigmas = get_sigmas_polyexponential(steps, sigma_min, sigma_max, rho, device="cuda")
            x = x * sigmas[0]
            with torch.no_grad(), torch.autocast("cuda", dtype=models["info"]["weights_precision"]):
                pred_z = sample_dpmpp_3m_sde(
                    denoiser,
                    x,
                    sigmas,
                    eta=sampler_eta,
                    s_noise=sampler_s_noise,
                    disable=n_chunks > 1,
                    extra_args=extra_args,
                )
        elif objective in ["rf_denoiser", "rectified_flow"]:
            if sigma_max > 1:
                sigma_max = 1

            model_fn = models["dit_model"].forward_inference

            logsnr_max = math.log(((1 - sigma_max) / sigma_max) + 1e-6) if sigma_max < 1 else -6
            logsnr = torch.linspace(logsnr_max, 2, steps + 1)
            t = torch.sigmoid(-logsnr)

            t[0] = sigma_max
            t[-1] = 0

            # Apply time shift if specified
            dist_shift_obj = None
            if time_shift != 0.0:
                dist_shift_obj = TimeShift(shift=time_shift)

            sampler_fn = {
                "euler": sample_discrete_euler,
                "rk4": sample_rk4,
                "dpmpp": sample_flow_dpmpp,
                "pingpong": sample_flow_pingpong,
            }[sampler_type]

            with torch.no_grad(), torch.autocast("cuda", dtype=models["info"]["weights_precision"]):
                pred_z = sampler_fn(
                    model_fn,
                    x,
                    sigmas=t,
                    sigma_max=sigma_max,
                    dist_shift=dist_shift_obj,
                    **extra_args,
                )

        # crop pred_z to chunk_window_size
        pred_z_chunk = pred_z[:, :, :chunk_window_size]
        pred_z_chunk = pred_z_chunk.swapaxes(2, 1)

        # if doing an infill, replace pred_z_chunk with the infill_ctx_vae
        # with the proper mask
        if infill_ctx_vae is not None:
            if not infill_variation:
                # replace the output with the infill_ctx_vae
                masked_pred_z_chunk = pred_z_chunk * (~infill_ctx_mask.unsqueeze(-1))
                pred_z_chunk = masked_pred_z_chunk + infill_ctx_vae * infill_ctx_mask.unsqueeze(-1)

        start_idx = n_chunk * chunk_window_size
        end_idx = start_idx + chunk_window_size
        out_pred_z[:, start_idx:end_idx, :] = pred_z_chunk.cpu()  # non-scaled

    # decode latents
    input_audio_len = int(round(semantic_codes.shape[1] / SEMANTIC_HZ * models["dit_model"].io_hz))
    if truncate_at_last_chunk and input_audio_len >= 25:
        # truncate last 0.2s to avoid artifacts of incomplete chunks
        input_audio_len -= 5
    scaled_pred_z = (out_pred_z[:, :input_audio_len] / models["codec_scale_factor"])[0]

    if models["patch_size"] > 1:
        scaled_pred_z = unfold_tensor(
            scaled_pred_z, models["dit_model"].io_channels // models["patch_size"]
        )

    if use_codec:
        pred_audio = codec_decode(scaled_pred_z)
    else:
        pred_audio = scaled_pred_z

    if return_latents:
        return pred_audio, scaled_pred_z
    else:
        return pred_audio


def simple_generate(model, semantic_codes, text_codes, ctx_vae=None, ctx_mask=None, n_steps=8):
    """Used for computing metrics during training"""
    x = torch.randn(semantic_codes.shape[0], model.io_channels, model.block_size, device=model.device)
    sigmas = get_sigmas_polyexponential(50, 0.1, n_steps, device=model.device)
    x = x * sigmas[0]
    x = x.to(torch.bfloat16)
    denoiser = VDenoiser(model)
    out_pred_z = sample_dpmpp_3m_sde(
        denoiser,
        x,
        sigmas,
        extra_args={
            "text_codes": text_codes,
            "semantic_codes": semantic_codes,
            "ctx_vae": ctx_vae,
            "ctx_mask": ctx_mask,
        },
        disable=True,
    )
    return out_pred_z
