import shutil
from contextlib import nullcontext
import datetime
import functools
import logging
import math
import os
import random
import time
import gc
import copy

import numpy as np
from tqdm import tqdm
from einops import rearrange
import torch
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.distributed.fsdp import (
    FullyShardedDataParallel as FSDP,
    ShardingStrategy,
)
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from torch.distributed import destroy_process_group, init_process_group
from torch.utils.data import DataLoader
from data_utils import PreprocessDataset
from audioloader import AudioLoaderDataset, AudioConfig
from modules.base import (
    apply_fsdp_checkpointing,
    configure_optimizers as base_configure_optimizers,
    estimate_mfu_no_model,
    LayerNorm,
    CausalSelfAttention,
    MLP,
    Block,
)
from utils.fsdp_policies import bfSixteen
from utils.helpers import (
    dist_barrier,
    hash_string_to_number,
    load_checkpoint,
    load_old_state_dict,
    load_old_optimizer_state_dict,
    print_with_time,
    print_with_time_master,
    save_checkpoint,
    save_old_checkpoint,
    save_dual_model_checkpoint,
    suppress_logging,
    verify_preload_model_args,
)
from utils.logging import build_gpu_memory_monitor, Color, NoColor
from utils.profiling import maybe_enable_memory_snapshot, maybe_enable_profiling
from models.model_selector import CodecConfig, DiscriminatorConfig, get_model
from modules.losses import LSGANLoss, MelSpectrogramLoss

color = Color if True else NoColor

# turn down some annoying fsdp logging
logging.getLogger("torch.distributed.fsdp._debug_utils").setLevel(logging.ERROR)
logging.getLogger("torch.distributed.fsdp._optim_utils").setLevel(logging.ERROR)
logging.getLogger("torch.distributed.checkpoint._dedup_tensors").setLevel(logging.ERROR)

os.umask(0o003)  # set umask to 0o003 to allow group write for created directories

data_dir = None
out_dir = None

enable_profiling = False
dump_folder = "/app/suno/gpt_profiling/"
save_traces_folder = "traces"
profile_freq = 50
enable_memory_snapshot = False
save_memory_snapshot_folder = "memory_snapshots"

master_addr = "localhost"
master_port = 12355

train_metas_filename = "metas_tr.jsonl"
val_metas_filename = "metas_val.jsonl"
debug_val_only = False
preload_checkpoint = None
preload_optimizer = False
local_cache_dir = None  # checkpoint will get copied here, 1 per node to allow for faster loading
preload_strict = True  # enforce keys in dict on load
suppress_compile_warnings = True
grad_checkpointing = False
checkpoint_save_old_format = True
# model params
codec_type = "dac_vae"
encoder_dim = 128
encoder_rates = [2, 3, 5, 8, 8]
latent_dim = 128
decoder_dim = 2048
decoder_rates = [8, 8, 5, 3, 2]
vae_dim = 128
is_frozen_encoder = False
discriminator_type = "dac_discriminator"
discriminator_rates = []
discriminator_periods = [2, 3, 5, 7, 11]
discriminator_fft_sizes = [2048, 1024, 512]
discriminator_bands = [(0.0, 0.1), (0.1, 0.25), (0.25, 0.5), (0.5, 0.75), (0.75, 1.0)]
# data
n_channels = 2
sample_rate = 48000
duration_s = 1.0
is_vae = False
is_mert = False
is_musicfm = False
gradient_accumulation_steps = 1
batch_size = 2  # if gradient_accumulation_steps > 1, this is the micro-batch size
batch_store_size = 4
# train params
weight_mel_loss = 15.0
weight_kl_loss = 0.0001
weight_feat_loss = 2.0
weight_adv_loss = 1.0
weight_disc_loss = 1.0

# eval items
custom_seed_offset = 0
eval_interval = 2000
log_interval = 25
eval_iters = 50
eval_only = False  # if True, script exits right after the first eval
debug_gradients = False
# wandb logging
wandb_log = False
wandb_project = "suno-test"
wandb_run_name = "test"
wandb_dir = None
# model
n_transformer_layers = 0
# adamw optimizer
learning_rate_codec = 1.5e-4  # codec/generator learning rate (matches impl 1)
learning_rate_disc = 3e-4  # discriminator learning rate (matches impl 1)
max_iters = 100_000  # total number of training iterations
step_save_iters = 20_000  # at this checkpoint we save the model
weight_decay = 0.0  # removed weight decay to match impl 3
beta1 = 0.8
beta2 = 0.99
grad_clip_gen = 10.0  # clip gradients at this value, or disable if == 0.0
grad_clip_disc = 1000.0  # clip gradients at this value, or disable if == 0.0
# LR scheduler selection: "inverse", "cosine", "exponential"
lr_scheduler_type = "inverse"  # "inverse", "cosine", or "exponential"
# InverseLR scheduler params
inverse_lr_gamma = 200000
inverse_lr_power = 0.5
inverse_lr_warmup = 0.999
# ExponentialLR scheduler params
exponential_lr_gamma = 0.999996

# system
device = "cuda"
dtype = "bfloat16"  # "float32", "bfloat16"
compile = False  # use PyTorch 2.0 to compile the model to be faster
fsdp = False  # fully sharded data parallel
sharding_strategy = "full_shard"
# -----------------------------------------------------------------------------
config_keys = [
    k
    for k, v in globals().items()
    if not k.startswith("_") and isinstance(v, (int, float, bool, str, type(None)))
]
exec(open("configurator.py").read())  # overrides from command line or config file
config = {k: globals()[k] for k in config_keys}  # will be useful for logging
# -----------------------------------------------------------------------------

# Validate scheduler type
if lr_scheduler_type not in ["inverse", "cosine", "exponential"]:
    raise ValueError(f"lr_scheduler_type must be 'inverse', 'cosine', or 'exponential', got: {lr_scheduler_type}")

use_raw_audio = True


assert dtype in ("bfloat16", "float32", "float16")


if debug_val_only or eval_only:
    train_metas_filename = val_metas_filename
    if not eval_only:
        wandb_log = False

# set up distributed variables
os.environ["MASTER_ADDR"] = str(master_addr)
os.environ["MASTER_PORT"] = str(master_port)
if "SLURM_PROCID" in os.environ:  # Running on SLURM
    if int(os.environ["SLURM_NTASKS_PER_NODE"]) != torch.cuda.device_count():
        raise ValueError(
            f"SLURM_NTASKS_PER_NODE ({os.environ['SLURM_NTASKS_PER_NODE']}) does not match"
            f" the number of CUDA devices ({torch.cuda.device_count()}) on node {os.environ['HOSTNAME']}"
        )
    ddp_rank = int(os.environ["SLURM_PROCID"])
    ddp_local_rank = int(os.environ["SLURM_LOCALID"])
    world_size = int(os.environ["SLURM_JOB_NUM_NODES"]) * int(os.environ["SLURM_NTASKS_PER_NODE"])
elif "RANK" in os.environ and "LOCAL_RANK" in os.environ and "WORLD_SIZE" in os.environ:
    # Running with torchrun
    ddp_rank = int(os.environ["RANK"])
    ddp_local_rank = int(os.environ["LOCAL_RANK"])
    world_size = int(os.environ["WORLD_SIZE"])
    print(f"Detected torchrun: rank={ddp_rank}, local_rank={ddp_local_rank}, world_size={world_size}")
else:  # Running locally single GPU
    ddp_rank = 0
    ddp_local_rank = 0
    world_size = 1
os.environ["RANK"] = str(ddp_rank)
os.environ["LOCAL_RANK"] = str(ddp_local_rank)

# various inits, derived attributes, I/O setup
ddp = int(os.environ.get("RANK", -1)) != -1  # is this a ddp run?

if fsdp:
    assert ddp, "found fsdp = True but ddp is False"

if ddp:
    print(
        f"Attempting DDP initialization: rank={ddp_rank}, world_size={world_size}, local_rank={ddp_local_rank}"
    )
    try:
        print(f"Calling init_process_group with backend=nccl...")
        init_process_group(
            backend="nccl",
            timeout=datetime.timedelta(seconds=2 * 60 * 60),
            rank=ddp_rank,
            world_size=world_size,
            device_id=torch.device(f"cuda:{ddp_local_rank}"),
        )
        print(f"init_process_group completed successfully for rank {ddp_rank}")
    except Exception as e:
        print(f"Distributed error on rank {ddp_rank} with host {os.environ.get('HOSTNAME', 'Unknown')}")
        print(f"Exception details: {e}")
        raise e
    device = f"cuda:{ddp_local_rank}"
    torch.cuda.set_device(device)
    master_process = ddp_rank == 0  # this process will do logging, checkpointing etc.
    local_master_process = ddp_local_rank == 0
    seed_offset = ddp_rank + 1  # each process gets a different seed
    print_with_time(f"ddp init, rank {ddp_rank}, local_rank {ddp_local_rank}")
else:
    # if not ddp, we are running on a single gpu, and one process
    master_process = True
    local_master_process = True
    seed_offset = 1
n_gpus_per_node = torch.cuda.device_count()
dist_barrier()
print_with_time_master(f"ddp init: world size {world_size} ddp_rank {ddp_rank}.")

# make sure we offset seeds in a clever way
seed_offset += (
    custom_seed_offset * world_size + 0
    if preload_checkpoint is None
    else hash_string_to_number(preload_checkpoint)
)
torch.manual_seed(6006 + seed_offset)
random.seed(6006 + seed_offset)
np.random.seed(6006 + seed_offset)
torch.backends.cuda.matmul.allow_tf32 = True  # allow tf32 on matmul
torch.backends.cudnn.allow_tf32 = True  # allow tf32 on cudnn
device_type = "cuda" if "cuda" in device else "cpu"  # for later use in torch.autocast
ptdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16, "float16": torch.float16}[dtype]
# Use float32 for input audio when frozen encoder is enabled (for better quality)
input_dtype = torch.float32 if is_frozen_encoder else ptdtype
if is_frozen_encoder and dtype in ("bfloat16", "float16"):
    print_with_time_master(
        f"Mixed precision training enabled: Input+Encoder+Quantizer (float32), Decoder ({dtype})"
    )
ctx = (
    nullcontext()
    if device_type == "cpu" or fsdp
    else torch.amp.autocast(device_type=device_type, dtype=ptdtype)
)

loss_discount_map = {}
loss_discount_map["mel_loss"] = weight_mel_loss
loss_discount_map["kl_loss"] = weight_kl_loss
loss_discount_map["feat_loss"] = weight_feat_loss
loss_discount_map["adv_loss"] = weight_adv_loss
loss_discount_map["disc_loss"] = weight_disc_loss

for k, v in loss_discount_map.items():
    print_with_time_master(f"loss weight for {k}: {v:.3f}")

date_time_str = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
out_dir = os.path.join(out_dir, date_time_str)
if master_process and not debug_val_only:
    os.makedirs(out_dir, exist_ok=True)
    print_with_time_master(f"logging checkpoint here: {out_dir}")

# logging
if wandb_log and master_process:
    import wandb
    import importlib.metadata
    import sys

    wandb_log_cfg = {}
    wandb_log_cfg["run_config"] = {k: v for k, v in config.items()}
    wandb_log_cfg["world_size"] = world_size
    wandb_log_cfg["slurm_id"] = os.environ.get("SLURM_JOB_ID")
    wandb_log_cfg["slurm_name"] = os.environ.get("SLURM_JOB_NAME")
    wandb_log_cfg["slurm_script_path"] = os.environ.get("SLURM_SCRIPT_PATH")
    wandb_log_cfg["checkpoint_dir"] = out_dir
    wandb_log_cfg["pip_freeze"] = {
        dist.metadata["Name"]: dist.version for dist in importlib.metadata.distributions()
    }
    wandb_log_cfg["python_path"] = sys.executable

    wandb.init(project=wandb_project, name=wandb_run_name, config=wandb_log_cfg, dir=wandb_dir)
    wandb.run.log_code(".")
    print_with_time_master(f"Total world size {world_size}")

if not use_raw_audio:
    raise NotImplementedError("Raw audio not implemented")

dist_barrier()

# model init
codec_args = dict(
    model_type=codec_type,
    encoder_dim=encoder_dim,
    encoder_rates=encoder_rates,
    latent_dim=latent_dim,
    decoder_dim=decoder_dim,
    decoder_rates=decoder_rates,
    vae_dim=vae_dim,
    sample_rate=sample_rate,
    is_frozen_encoder=is_frozen_encoder,
    n_transformer_layers=n_transformer_layers,
)
discriminator_args = dict(
    model_type=discriminator_type,
    rates=discriminator_rates,
    periods=discriminator_periods,
    fft_sizes=discriminator_fft_sizes,
    sample_rate=sample_rate,
    bands=discriminator_bands,
)

if preload_checkpoint is not None and not preload_checkpoint.endswith(".pt"):
    # verification for old style we do later on checkpoint load
    verify_preload_model_args(codec_args, preload_checkpoint, preload_strict=preload_strict)
gpu_memory_monitor = build_gpu_memory_monitor()
# init a new model from scratch
print_with_time_master("Initializing a new model from scratch")
codec_config = CodecConfig(**codec_args)
codec_model = get_model(codec_config)
discriminator_config = DiscriminatorConfig(**discriminator_args)
discriminator_model = get_model(discriminator_config)

if not fsdp:
    codec_model.to(device)
    discriminator_model.to(device)
    # Convert models to target dtype (bfloat16, float16, or float32)
    if dtype in ("bfloat16", "float16"):
        print_with_time_master(f"Converting models to {dtype} for DDP...")
        codec_model = codec_model.to(dtype=ptdtype)
        discriminator_model = discriminator_model.to(dtype=ptdtype)
        
        # If frozen encoder mode, keep encoder/quantizer in float32 for better quality
        if is_frozen_encoder:
            print_with_time_master("Keeping frozen encoder/quantizer in float32 for better quality...")
            codec_model.encoder = codec_model.encoder.float()
            codec_model.quantizer = codec_model.quantizer.float()

print(codec_args)
print(discriminator_args)
print(f"codec model params: {codec_model.get_num_params()}")
print(f"discriminator model params: {discriminator_model.get_num_params()}")

# this is needed to calculate MFU later
# it will get messed up by FSDP, so calculate now
raw_model_n_params = codec_model.get_num_params()

# compile the model
if compile:
    import torch._dynamo

    torch._dynamo.config.cache_size_limit = 512  # 64
    print_with_time_master("compiling the model... (takes a ~minute)")
    compile_ctx = suppress_logging if suppress_compile_warnings else nullcontext
    with compile_ctx():
        # codec_model = torch.compile(codec_model, fullgraph=True, mode="max-autotune")
        codec_model = torch.compile(codec_model)
        discriminator_model = torch.compile(discriminator_model)
    dist_barrier()
else:
    print_with_time_master("not compiling model.")

iter_num = 0
total_hours_processed = 0
rel_hours_processed = 0
best_val_loss = 1e9  # Tracks best validation mel loss (not overall loss)

if local_cache_dir is not None and local_master_process:
    shutil.rmtree(local_cache_dir, ignore_errors=True)
    os.makedirs(local_cache_dir, exist_ok=True)
    os.chmod(local_cache_dir, 0o774)

# load old single-file checkpoint
if preload_checkpoint is not None and preload_checkpoint.endswith(".pt"):
    # Load checkpoint file
    if local_cache_dir is not None:
        local_ckpt_fp = os.path.join(local_cache_dir, "ckpt.pt")
        if master_process:
            print_with_time_master("copying checkpoint file to local cachedir...")
            os.makedirs(local_cache_dir, exist_ok=True)
            shutil.copy2(preload_checkpoint, local_ckpt_fp)
        dist_barrier()
    else:
        local_ckpt_fp = preload_checkpoint

    # Load both models from checkpoint
    print_with_time_master("loading codec and discriminator state_dicts...")
    checkpoint = torch.load(local_ckpt_fp, map_location=device)

    # Load codec model
    if "codec_model" in checkpoint:
        codec_model.load_state_dict(checkpoint["codec_model"], strict=preload_strict)
        print_with_time_master("loaded codec model state_dict")
    elif "model" in checkpoint:
        # Fallback to old format where only codec was saved
        codec_model.load_state_dict(checkpoint["model"], strict=preload_strict)
        print_with_time_master("loaded codec model state_dict (legacy format)")

    # Load discriminator model if available
    if "discriminator_model" in checkpoint:
        discriminator_model.load_state_dict(checkpoint["discriminator_model"], strict=preload_strict)
        print_with_time_master("loaded discriminator model state_dict")
    else:
        print_with_time_master("discriminator state_dict not found in checkpoint - starting fresh")

    del checkpoint
    dist_barrier()

# order matters:
# FSDP: load model ckpt, wrap model, make optim (sharded), shard ckpt into optimizer
# DDP: load model ckpt, make optimizer, load model and optimizer checkpoints
if fsdp:
    print_with_time_master("wrapping models in FSDP...")
    codec_model = FSDP(
        codec_model,
        mixed_precision=bfSixteen,
        sharding_strategy=getattr(ShardingStrategy, sharding_strategy.upper()),
        device_id=torch.cuda.current_device(),
        sync_module_states=True,
        use_orig_params=True,
        # cpu_offload=torch.distributed.fsdp.CPUOffload(offload_params=True),
    )
    discriminator_model = FSDP(
        discriminator_model,
        mixed_precision=bfSixteen,
        sharding_strategy=getattr(ShardingStrategy, sharding_strategy.upper()),
        device_id=torch.cuda.current_device(),
        sync_module_states=True,
        use_orig_params=True,
        # cpu_offload=torch.distributed.fsdp.CPUOffload(offload_params=True),
    )
    gpu_mem_stats = gpu_memory_monitor.get_peak_stats()
    print_with_time_master(
        f"GPU memory usage for model: "
        f"{gpu_mem_stats.max_reserved_gib:.2f}GiB"
        f"({gpu_mem_stats.max_reserved_pct:.2f}%)"
    )
    if grad_checkpointing:
        apply_fsdp_checkpointing(codec_model)
        apply_fsdp_checkpointing(discriminator_model)
    generator_optimizer = base_configure_optimizers(
        codec_model,
        weight_decay,
        learning_rate_codec,  # Use separate LR for codec
        (beta1, beta2),
        device_type,
        use_fused=False,
        is_fsdp=True,
    )
    discriminator_optimizer = base_configure_optimizers(
        discriminator_model,
        weight_decay,
        learning_rate_disc,  # Use separate LR for discriminator
        (beta1, beta2),
        device_type,
        use_fused=False,
        is_fsdp=True,
    )
else:  # both DDP and single-worker
    # Create separate optimizers for generator (codec) and discriminator
    generator_optimizer = base_configure_optimizers(
        codec_model,
        weight_decay,
        learning_rate_codec,  # Use separate LR for codec
        (beta1, beta2),
        device_type,
        use_fused=False,
        is_fsdp=False,
    )
    discriminator_optimizer = base_configure_optimizers(
        discriminator_model,
        weight_decay,
        learning_rate_disc,  # Use separate LR for discriminator
        (beta1, beta2),
        device_type,
        use_fused=False,
        is_fsdp=False,
    )

    if ddp:
        print_with_time_master("wrapping models in DDP")
        codec_model = DDP(codec_model, device_ids=[ddp_local_rank], find_unused_parameters=True)
        discriminator_model = DDP(discriminator_model, device_ids=[ddp_local_rank], find_unused_parameters=True)
        
    # After DDP wrapping, reconvert frozen encoder/quantizer to float32 for mixed precision
    if is_frozen_encoder and dtype in ("bfloat16", "float16"):
        print_with_time_master("Reconverting frozen encoder/quantizer to float32 for mixed precision...")
        # Access the actual model (unwrap DDP if needed)
        actual_codec_model = codec_model.module if ddp else codec_model
        actual_codec_model.encoder = actual_codec_model.encoder.float()
        actual_codec_model.quantizer = actual_codec_model.quantizer.float()
        
        # Verify dtype conversion
        if master_process:
            enc_dtype = next(actual_codec_model.encoder.parameters()).dtype
            quant_params = list(actual_codec_model.quantizer.parameters())
            quant_dtype = next(actual_codec_model.quantizer.parameters()).dtype if quant_params else "no params"
            dec_dtype = next(actual_codec_model.decoder.parameters()).dtype
            print_with_time_master(f"Dtype check - Encoder: {enc_dtype}, Quantizer: {quant_dtype}, Decoder: {dec_dtype}")
        
torch.cuda.empty_cache()
dist_barrier()

# load old single-file checkpoint optimizers
if preload_checkpoint is not None and preload_checkpoint.endswith(".pt") and preload_optimizer:
    local_ckpt_fp = (
        preload_checkpoint if local_cache_dir is None else os.path.join(local_cache_dir, "ckpt.pt")
    )
    print_with_time_master("loading optimizer state_dicts...")
    checkpoint = torch.load(local_ckpt_fp, map_location=device)

    # Load optimizer states
    if "generator_optimizer" in checkpoint:
        generator_optimizer.load_state_dict(checkpoint["generator_optimizer"])
        print_with_time_master("loaded generator optimizer state_dict")

    if "discriminator_optimizer" in checkpoint:
        discriminator_optimizer.load_state_dict(checkpoint["discriminator_optimizer"])
        print_with_time_master("loaded discriminator optimizer state_dict")

    # Load training state
    if "iter_num" in checkpoint:
        iter_num = checkpoint["iter_num"]
    if "best_val_loss" in checkpoint:
        best_val_loss = checkpoint["best_val_loss"]
    if "n_hours" in checkpoint:
        total_hours_processed = checkpoint["n_hours"]

    del checkpoint
    print_with_time_master(f"resumed from iteration {iter_num}, best_val_loss: {best_val_loss:.4f}")

# load new distributed checkpoint
if preload_checkpoint is not None and not preload_checkpoint.endswith(".pt"):
    # For now, distributed checkpoint loading for dual models not implemented
    # iter_num, total_tokens_processed, best_val_loss = load_checkpoint(
    #     preload_checkpoint,
    #     preload_optimizer,
    #     codec_model,
    #     generator_optimizer,
    # )
    print_with_time_master("distributed checkpoint loading for dual models not implemented yet")
    dist_barrier()

print_with_time_master("model setup done")

# set up loss functions
gan_loss = LSGANLoss(discriminator_model)
mel_loss = MelSpectrogramLoss().to(device)


# InverseLR scheduler
def get_inverse_lr(it, base_lr, inv_gamma=200000, power=0.5, warmup=0.999, final_lr=0.0):
    """Inverse decay learning rate schedule with exponential warmup.
    
    Args:
        it: Current iteration
        base_lr: Base learning rate
        inv_gamma: Inverse multiplicative factor of learning rate decay
        power: Exponential factor of learning rate decay
        warmup: Exponential warmup factor (0 <= warmup < 1, 0 to disable)
        final_lr: The final learning rate
    """
    warmup_factor = 1 - warmup ** (it + 1)
    lr_mult = (1 + it / inv_gamma) ** -power
    return warmup_factor * max(final_lr, base_lr * lr_mult)


# ExponentialLR scheduler
def get_exponential_lr(it, base_lr, gamma=0.999996):
    """Exponential decay learning rate schedule.
    
    Args:
        it: Current iteration
        base_lr: Base learning rate
        gamma: Multiplicative factor of learning rate decay per step
    
    Returns:
        Current learning rate
    """
    return base_lr * (gamma ** it)


def flatten_stereo(batch):
    return rearrange(batch, "b c t -> (b c) 1 t")


# audio config
audio_cfg = AudioConfig(
    sample_rate=sample_rate,
    n_channels=n_channels,
    duration_s=duration_s,
    is_vae=is_vae,
    is_mert=is_mert,
    is_musicfm=is_musicfm,
)


def make_data_iter(split, is_eval=False):
    metas_filename = train_metas_filename if split == "train" else val_metas_filename
    audio_dataset = AudioLoaderDataset(
        audio_cfg,
        os.path.join(data_dir, metas_filename),
        split=split,
    )
    audio_dataloader = DataLoader(
        audio_dataset,
        shuffle=False,
        num_workers=3 if is_eval else 6,
        prefetch_factor=batch_store_size * (2 if is_eval else 4),
        batch_size=None,
        worker_init_fn=lambda x: random.seed(x + seed_offset * 1000),
    )
    audio_dataloader_iter = iter(audio_dataloader)

    dataset = PreprocessDataset(
        audio_dataloader_iter,
        batch_size,
        audio_cfg,
    )
    dataloader = DataLoader(
        dataset,
        shuffle=False,
        num_workers=0,  # no multiprocessing so its in sync across workers
        batch_size=None,
        worker_init_fn=lambda x: random.seed(x + seed_offset * 1000),
    )

    def cached_dataloader_iter_fn(dataloader):
        # preload batch_store_size batches in memory.
        # #this way we make new batches every 10 steps, smoothing out the load
        dataloader_iter = iter(dataloader)
        batch_store = []
        while True:
            if not batch_store:
                for _ in tqdm(range(batch_store_size), desc="preloading batches", disable=True):
                    try:
                        batch = next(dataloader_iter)
                    except StopIteration:
                        print("dataloader_iter exhausted, resetting")
                        dataloader_iter = iter(dataloader)
                        batch = next(dataloader_iter)
                    batch_store.append(batch)

            batch = batch_store.pop(0)
            yield batch

    return cached_dataloader_iter_fn(dataloader)


tr_dataloader_iter = make_data_iter("train", is_eval=False)
val_dataloaders = {}
for k in ["train", "val"]:
    val_dataloaders[k] = {}
    val_dataloaders[k][0] = make_data_iter(k, is_eval=True)
train_dataset_names = ["main"]
val_dataset_names = ["main"]


@torch.no_grad()
def estimate_loss():
    n_loss_modalities = len(train_dataset_names) + len(val_dataset_names)
    effective_eval_iters = int(round(eval_iters / n_loss_modalities))
    if fsdp:
        modules_for_eval = (
            torch.nn.Linear,
            torch.nn.Dropout,
            torch.nn.Embedding,
            torch.nn.SiLU,
            LayerNorm,
            MLP,
            CausalSelfAttention,
        )
        for name, module in codec_model.named_modules():
            module.train(False)
        for name, module in discriminator_model.named_modules():
            module.train(False)
    else:
        codec_model.eval()
        discriminator_model.eval()
    loss_prefix = "loss"
    n_loss_entries = n_loss_modalities * len(loss_discount_map)
    loss_tensor = torch.zeros(n_loss_entries, device=device)
    loss_tensor_keys = []
    n_loss_entry = 0
    for split in ["train", "val"]:
        n_datasets = len(train_dataset_names) if split == "train" else len(val_dataset_names)
        dataset_names = train_dataset_names if split == "train" else val_dataset_names
        for dataset_idx in range(n_datasets):
            losses = []
            for _ in range(effective_eval_iters):
                data_list = next(val_dataloaders[split][dataset_idx])

                with ctx:
                    # For evaluation, just compute codec reconstruction loss
                    raw_audio_batch = [line.data_wav for line in data_list]
                    if is_mert:
                        input_batch = [line.data_mert for line in data_list]

                    elif is_musicfm:
                        input_batch = [line.data_musicfm for line in data_list]
                    elif is_vae:
                        input_batch = [line.data_vae for line in data_list]
                    else:
                        input_batch = raw_audio_batch
                    input_batch = torch.from_numpy(np.stack(input_batch)).to(device, dtype=input_dtype)
                    raw_audio_batch = torch.from_numpy(np.stack(raw_audio_batch)).to(
                        device, dtype=input_dtype
                    )

                    outp = codec_model(input_batch)
                    gen_audio = outp["audio"]

                    # Compute all losses for evaluation
                    disc_loss_val = gan_loss.discriminator_loss(gen_audio, raw_audio_batch)
                    mel_loss_val = mel_loss(flatten_stereo(gen_audio), flatten_stereo(raw_audio_batch))
                    mel_loss_val += mel_loss(gen_audio.mean(dim=1), raw_audio_batch.mean(dim=1))
                    mel_loss_val /= 2
                    kl_loss_val = outp.get("kl", torch.tensor(0.0, device=device, dtype=ptdtype))

                    adv_loss_val, feat_loss_val = gan_loss.generator_loss(gen_audio, raw_audio_batch)

                    # Create loss dict matching loss_discount_map keys
                    loss_dict = {
                        "mel_loss": mel_loss_val,
                        "kl_loss": kl_loss_val,
                        "feat_loss": feat_loss_val,
                        "adv_loss": adv_loss_val,
                        "disc_loss": disc_loss_val,
                    }

                losses.append(
                    [
                        loss_dict[k].item() if isinstance(loss_dict[k], torch.Tensor) else loss_dict[k]
                        for k in loss_discount_map.keys()
                    ]
                )
            for n, loss_name in enumerate(loss_discount_map.keys()):
                loss_tensor[n_loss_entry] = float(np.mean([e[n] for e in losses]))
                loss_tensor_keys.append(
                    f"{split}/{loss_prefix}_{dataset_names[dataset_idx]}_{loss_name}"
                )
                n_loss_entry += 1
    if ddp:
        torch.distributed.all_reduce(loss_tensor, op=torch.distributed.ReduceOp.AVG)
    tmp_out = {k: loss_tensor[n].item() for n, k in enumerate(loss_tensor_keys)}
    # add extra loss items
    out = {k: v for k, v in tmp_out.items()}
    out[f"train/{loss_prefix}"] = float(
        np.mean([v for k, v in tmp_out.items() if k.startswith("train/")])
    )
    out[f"val/{loss_prefix}"] = float(np.mean([v for k, v in tmp_out.items() if k.startswith("val/")]))
    codec_model.train()
    discriminator_model.train()
    for name, module in codec_model.named_modules():
        assert module.training  # make sure we can undo everything
    for name, module in discriminator_model.named_modules():
        assert module.training  # make sure we can undo everything
    return out


gc.disable()  # manually gc to avoid slowdowns https://imbue.com/research/70b-infrastructure/

# training loop
print_with_time_master("training...")
t0 = time.time()
t00 = time.time()
t_start = time.time()  # absolute time since starting to train
t_data = 0  # time spent loading data
t_wait = 0  # time spent waiting in loop
t_model = 0  # time spent in master node model fw+bw
local_iter_num = 0  # number of iterations in the lifetime of this process
# raw_model = model.module if ddp else model  # unwrap DDP container if needed
mfu = 0
batch_per_s = 0
effective_batch_per_s_per_node = 0
running_gen_loss = []
running_disc_loss = []
running_mel_loss = []
running_kl_loss = []
running_feat_loss = []
running_adv_loss = []
# with torch.autograd.set_detect_anomaly(True):
gpu_memory_monitor.reset_peak_stats()
with (
    maybe_enable_profiling(
        enable_profiling, dump_folder, save_traces_folder, profile_freq, global_step=iter_num
    ) as torch_profiler,
    maybe_enable_memory_snapshot(
        enable_memory_snapshot,
        dump_folder,
        save_memory_snapshot_folder,
        profile_freq,
        global_step=iter_num,
    ) as memory_profiler,
):
    while True:
        # determine and set the learning rate for this iteration
        if lr_scheduler_type == "inverse":
            lr_gen = get_inverse_lr(iter_num, learning_rate_codec, 
                                   inv_gamma=inverse_lr_gamma, 
                                   power=inverse_lr_power, 
                                   warmup=inverse_lr_warmup)
            lr_disc = get_inverse_lr(iter_num, learning_rate_disc,
                                    inv_gamma=inverse_lr_gamma,
                                    power=inverse_lr_power,
                                    warmup=inverse_lr_warmup)
        elif lr_scheduler_type == "exponential":
            lr_gen = get_exponential_lr(iter_num, learning_rate_codec, 
                                       gamma=exponential_lr_gamma)
            lr_disc = get_exponential_lr(iter_num, learning_rate_disc,
                                        gamma=exponential_lr_gamma)
        else:
            raise ValueError(f"Unknown lr_scheduler_type: {lr_scheduler_type}")
        
        for param_group in generator_optimizer.param_groups:
            param_group["lr"] = lr_gen
        for param_group in discriminator_optimizer.param_groups:
            param_group["lr"] = lr_disc

        # evaluate the loss on train/val sets and write checkpoints
        if iter_num % eval_interval == 0 or iter_num == max_iters - 1:
            dist_barrier()
            time_since_last_loss = time.time() - t00
            t00 = time.time()
            losses = estimate_loss()
            estimation_time = time.time() - t00
            eval_time_pct = np.clip(estimation_time / time_since_last_loss * 100, 0, 100)

            # Only master process runs evaluation to avoid conflicts
            if master_process:
                print_with_time_master(
                    f"loss estimation took {estimation_time:.1f} seconds. ({eval_time_pct:.1f}% of loop)"
                )
                print_with_time_master(f"step {iter_num}: val loss {losses['val/loss']:.4f}, val mel loss {losses['val/loss_main_mel_loss']:.4f}")
                if wandb_log:
                    log_dict = {
                        "iter": iter_num,
                        "n_hours": total_hours_processed,
                    }
                    for k, v in losses.items():
                        log_dict[k] = v
                    wandb.log(log_dict)
            else:
                # Other processes create dummy loss dict for checkpoint logic
                losses = {"val/loss": best_val_loss, "val/loss_main_mel_loss": best_val_loss}

            dist_barrier()
            if iter_num > 0:
                # Save checkpoints using proper FSDP/DDP handling
                if checkpoint_save_old_format:
                    # Use the dual model checkpoint saving function (best determined by mel loss)
                    best_val_loss = save_dual_model_checkpoint(
                        out_dir=out_dir,
                        codec_model=codec_model,
                        discriminator_model=discriminator_model,
                        generator_optimizer=generator_optimizer,
                        discriminator_optimizer=discriminator_optimizer,
                        best_val_loss=best_val_loss,
                        current_val_loss=losses["val/loss_main_mel_loss"],
                        step_save_iters=step_save_iters,
                        time_since_last_loss=time_since_last_loss,
                        codec_args=codec_args,
                        discriminator_args=discriminator_args,
                        iter_num=iter_num,
                        n_hours=total_hours_processed,
                        debug_val_only=False,
                        save_best_ckpt=True,
                        save_periodic_ckpt=True,
                    )
                else:
                    # Use distributed checkpoint format (implement later if needed)
                    print_with_time_master("distributed checkpoint saving not implemented yet")
                    pass
            dist_barrier()

        # end if eval test only
        if eval_only:
            print_with_time_master("eval test done.")
            break

        # GAN training - alternating discriminator and generator updates
        latest_loss_dict = {}  # Store latest loss values for logging
        grad_norm_disc = None
        grad_norm_gen = None
        grad_norm_gen_raw = None

        accum_gen_loss = 0.0
        accum_disc_loss = 0.0
        accum_mel_loss = 0.0
        accum_kl_loss = 0.0
        accum_feat_loss = 0.0
        accum_adv_loss = 0.0

        # ===============================
        # Load batches for this training iteration
        # ===============================
        batches = []
        for micro_step in range(gradient_accumulation_steps):
            t0_tmp = time.time()
            data_list = next(tr_dataloader_iter)
            t_data += time.time() - t0_tmp

            # Prepare audio data
            raw_audio_batch = [line.data_wav for line in data_list]
            if is_mert:
                input_batch = [line.data_mert for line in data_list]
            elif is_musicfm:
                input_batch = [line.data_musicfm for line in data_list]
            elif is_vae:
                input_batch = [line.data_vae for line in data_list]
            else:
                input_batch = raw_audio_batch

            input_batch = torch.from_numpy(np.stack(input_batch)).to(device, dtype=input_dtype)
            raw_audio_batch = torch.from_numpy(np.stack(raw_audio_batch)).to(device, dtype=input_dtype)
            
            batches.append((input_batch, raw_audio_batch))

        # ===============================
        # Phase 1: Accumulate and Update Discriminator
        # ===============================
        discriminator_optimizer.zero_grad(set_to_none=True)
        
        for micro_step in range(gradient_accumulation_steps):
            if ddp and micro_step < gradient_accumulation_steps - 1:
                disc_grad_sync_context = discriminator_model.no_sync
            else:
                disc_grad_sync_context = nullcontext

            t0_tmp = time.time()
            input_batch, raw_audio_batch = batches[micro_step]

            # Generate fake audio for discriminator training
            with torch.no_grad():
                codec_output = codec_model(input_batch)
                fake_audio = codec_output["audio"]

            # Train Discriminator
            with disc_grad_sync_context():
                with ctx:
                    # Compute discriminator loss
                    # Note: discriminator_loss internally detaches fake_audio
                    disc_loss = gan_loss.discriminator_loss(fake_audio, raw_audio_batch)
                    disc_loss = disc_loss / gradient_accumulation_steps

                # Backward pass for discriminator
                disc_loss.backward()

            accum_disc_loss += float(disc_loss.item())
            t_model += time.time() - t0_tmp

        # Clip discriminator gradients after accumulation
        if grad_clip_disc != 0.0:
            if fsdp:
                grad_norm_disc = discriminator_model.clip_grad_norm_(grad_clip_disc)
                if isinstance(grad_norm_disc, torch.Tensor):
                    if torch.isnan(grad_norm_disc).any():
                        raise RuntimeError("Found NaN in discriminator grad")
                elif math.isnan(float(grad_norm_disc)):
                    raise RuntimeError("Found NaN in discriminator grad")
            else:
                grad_norm_disc = torch.nn.utils.clip_grad_norm_(
                    discriminator_model.parameters(), grad_clip_disc, error_if_nonfinite=True
                )
            grad_norm_disc = (
                grad_norm_disc.item() if isinstance(grad_norm_disc, torch.Tensor) else grad_norm_disc
            )
        else:
            if isinstance(grad_norm_disc, torch.Tensor):
                grad_norm_disc = grad_norm_disc.item()

        # Step discriminator
        discriminator_optimizer.step()
        discriminator_optimizer.zero_grad(set_to_none=True)

        # ===============================
        # Phase 2: Accumulate and Update Generator (with updated discriminator)
        # ===============================
        generator_optimizer.zero_grad(set_to_none=True)
        
        for micro_step in range(gradient_accumulation_steps):
            if ddp and micro_step < gradient_accumulation_steps - 1:
                codec_grad_sync_context = codec_model.no_sync
            else:
                codec_grad_sync_context = nullcontext

            t0_tmp = time.time()
            # Use the SAME batches as discriminator training
            input_batch, raw_audio_batch = batches[micro_step]

            # Generate fake audio for generator training (with gradients)
            codec_output = codec_model(input_batch)
            fake_audio = codec_output["audio"]

            # Train Generator
            with codec_grad_sync_context():
                with ctx:
                    # Compute all generator losses (discriminator now uses updated weights)
                    mel_loss_val = mel_loss(flatten_stereo(fake_audio), flatten_stereo(raw_audio_batch))
                    mel_loss_val += mel_loss(fake_audio.mean(dim=1), raw_audio_batch.mean(dim=1))
                    mel_loss_val /= 2
                    kl_loss_val = codec_output.get("kl", torch.tensor(0.0, device=device, dtype=ptdtype))
                    adv_loss_val, feat_loss_val = gan_loss.generator_loss(fake_audio, raw_audio_batch)

                    # Combine losses with weights
                    total_gen_loss = (
                        weight_mel_loss * mel_loss_val
                        + weight_kl_loss * kl_loss_val
                        + weight_feat_loss * feat_loss_val
                        + weight_adv_loss * adv_loss_val
                    )
                    total_gen_loss = total_gen_loss / gradient_accumulation_steps

                    # Store loss values for logging
                    loss_dict = {
                        "mel_loss": mel_loss_val,
                        "kl_loss": kl_loss_val,
                        "feat_loss": feat_loss_val,
                        "adv_loss": adv_loss_val,
                        "disc_loss": accum_disc_loss / gradient_accumulation_steps,  # Use accumulated disc loss
                    }
                    latest_loss_dict = loss_dict  # Store for later logging

                # Backward pass for generator
                total_gen_loss.backward()

            accum_gen_loss += float(total_gen_loss.item())
            # Component losses need to be scaled for correct logging
            accum_mel_loss += float((mel_loss_val / gradient_accumulation_steps).item() if isinstance(mel_loss_val, torch.Tensor) else mel_loss_val / gradient_accumulation_steps)
            accum_kl_loss += float((kl_loss_val / gradient_accumulation_steps).item() if isinstance(kl_loss_val, torch.Tensor) else kl_loss_val / gradient_accumulation_steps)
            accum_feat_loss += float((feat_loss_val / gradient_accumulation_steps).item() if isinstance(feat_loss_val, torch.Tensor) else feat_loss_val / gradient_accumulation_steps)
            accum_adv_loss += float((adv_loss_val / gradient_accumulation_steps).item() if isinstance(adv_loss_val, torch.Tensor) else adv_loss_val / gradient_accumulation_steps)

            # Update processed hours (assuming each sample is duration_s seconds)
            batch_hours = (batch_size * duration_s / 3600.0) * world_size
            total_hours_processed += batch_hours
            rel_hours_processed += batch_hours

            if debug_gradients and wandb_log and master_process:
                d = {
                    "iter": iter_num,
                    "n_hours": total_hours_processed,
                    "total_gen_loss": total_gen_loss.item(),
                    "disc_loss": accum_disc_loss / gradient_accumulation_steps,
                }
                # Log individual loss components
                for k, v in loss_dict.items():
                    d[f"loss/{k}"] = v.item() if isinstance(v, torch.Tensor) else v
                wandb.log(d)

            t_model += time.time() - t0_tmp

        # Clip generator gradients after accumulation
        if grad_clip_gen != 0.0:
            if fsdp:
                grad_norm_gen = codec_model.clip_grad_norm_(grad_clip_gen)
                if isinstance(grad_norm_gen, torch.Tensor):
                    if torch.isnan(grad_norm_gen).any():
                        raise RuntimeError("Found NaN in generator grad")
                elif math.isnan(float(grad_norm_gen)):
                    raise RuntimeError("Found NaN in generator grad")
            else:
                grad_norm_gen_raw = torch.nn.utils.clip_grad_norm_(
                    codec_model.parameters(), float("inf")
                )
                grad_norm_gen_raw = (
                    grad_norm_gen_raw.item()
                    if isinstance(grad_norm_gen_raw, torch.Tensor)
                    else float(grad_norm_gen_raw)
                )
                if grad_norm_gen_raw > 100.0:
                    print(f"⚠️ WARNING: Large generator gradient: {grad_norm_gen_raw:.2f}")
                grad_norm_gen = torch.nn.utils.clip_grad_norm_(
                    codec_model.parameters(), grad_clip_gen, error_if_nonfinite=True
                )
            grad_norm_gen = (
                grad_norm_gen.item() if isinstance(grad_norm_gen, torch.Tensor) else grad_norm_gen
            )
            if grad_norm_gen_raw is not None and isinstance(grad_norm_gen_raw, torch.Tensor):
                grad_norm_gen_raw = grad_norm_gen_raw.item()
        else:
            if isinstance(grad_norm_gen, torch.Tensor):
                grad_norm_gen = grad_norm_gen.item()

        # Step generator
        generator_optimizer.step()
        generator_optimizer.zero_grad(set_to_none=True)

        running_gen_loss.append(accum_gen_loss)
        running_disc_loss.append(accum_disc_loss)
        running_mel_loss.append(accum_mel_loss)
        running_kl_loss.append(accum_kl_loss)
        running_feat_loss.append(accum_feat_loss)
        running_adv_loss.append(accum_adv_loss)

        if master_process and wandb_log:
            wandb.log(
                {
                    "iter": iter_num,
                    "n_hours": total_hours_processed,
                    "misc/grad_norm_gen": grad_norm_gen if grad_norm_gen is not None else 0.0,
                    "misc/grad_norm_disc": grad_norm_disc if grad_norm_disc is not None else 0.0,
                }
            )

        if iter_num % 300 == 1:
            # manually collect garbage in sync to avoid slowdowns over time
            # this slows down the iteration by ~30%, dont gc too often
            # do it on mod 1 to avoid it showing up in the logs
            assert not gc.isenabled()
            gc.collect()

        # timing and logging
        t1 = time.time()
        dt = t1 - t0
        t0 = t1
        if iter_num % log_interval == 0 or iter_num == max_iters - 1:
            if master_process:
                if local_iter_num >= 5:  # let the training loop settle a bit
                    # Note: we process gradient_accumulation_steps batches (same batches for both D and G)
                    batch_per_s = world_size * batch_size * gradient_accumulation_steps / dt
                    effective_batch_per_s_per_node = (
                        rel_hours_processed
                        * 3600
                        / duration_s  # convert hours back to batches
                        / (time.time() - t_start)
                        / max(1, int(round(world_size / n_gpus_per_node)))
                    )
                pct_wait = t_wait / (time.time() - t_start) * 100
                pct_data = t_data / (time.time() - t_start) * 100
                pct_model = t_model / (time.time() - t_start) * 100
                pct_overhead = 100 - pct_wait - pct_data - pct_model
                avg_gen_loss = np.mean(running_gen_loss) if running_gen_loss else 0.0
                avg_disc_loss = np.mean(running_disc_loss) if running_disc_loss else 0.0
                avg_mel_loss = np.mean(running_mel_loss) if running_mel_loss else 0.0
                avg_kl_loss = np.mean(running_kl_loss) if running_kl_loss else 0.0
                avg_feat_loss = np.mean(running_feat_loss) if running_feat_loss else 0.0
                avg_adv_loss = np.mean(running_adv_loss) if running_adv_loss else 0.0
                running_gen_loss = []
                running_disc_loss = []
                running_mel_loss = []
                running_kl_loss = []
                running_feat_loss = []
                running_adv_loss = []
                gpu_mem_stats = gpu_memory_monitor.get_peak_stats()

                print_with_time_master(
                    f"{color.cyan}iter {iter_num}:"
                    f"{color.green} gen_loss {avg_gen_loss:.3f},"
                    f"{color.red} disc_loss {avg_disc_loss:.3f},"
                    f"{color.magenta} mel_loss {avg_mel_loss:.3f},"
                    f"{color.cyan} kl_loss {avg_kl_loss:.6f},"
                    f"{color.yellow} feat_loss {avg_feat_loss:.3f},"
                    f"{color.blue} adv_loss {avg_adv_loss:.3f},"
                    f"{color.cyan} step_time {dt * 1000:.1f}ms,"
                    f"{color.green} throughput {batch_per_s:.1f} batch/s,"
                    f"{color.yellow} memory {gpu_mem_stats.max_reserved_gib:5.2f}GiB"
                    f"({gpu_mem_stats.max_reserved_pct:.2f}%)"
                    f"{color.reset}"
                )
                if wandb_log:
                    log_dict = {
                        "iter": iter_num,
                        "n_hours": total_hours_processed,
                        "avg_gen_loss": avg_gen_loss,
                        "avg_disc_loss": avg_disc_loss,
                        "avg_mel_loss": avg_mel_loss,
                        "avg_kl_loss": avg_kl_loss,
                        "avg_feat_loss": avg_feat_loss,
                        "avg_adv_loss": avg_adv_loss,
                        "lr_gen": lr_gen,
                        "lr_disc": lr_disc,
                        "batch/s": batch_per_s,
                        "eff_batch/s/node": effective_batch_per_s_per_node,
                        "perf/pct_wait": pct_wait,
                        "perf/pct_data": pct_data,
                        "perf/pct_model": pct_model,
                        "perf/pct_overhead": pct_overhead,
                        "perf/t_data": t_data,
                        "perf/t_model": t_model,
                        "perf/t_wait": t_wait,
                        "memory/max_active(GiB)": gpu_mem_stats.max_active_gib,
                        "memory/max_active(%)": gpu_mem_stats.max_active_pct,
                        "memory/max_reserved(GiB)": gpu_mem_stats.max_reserved_gib,
                        "memory/max_reserved(%)": gpu_mem_stats.max_reserved_pct,
                        "memory/num_alloc_retries": gpu_mem_stats.num_alloc_retries,
                        "memory/num_ooms": gpu_mem_stats.num_ooms,
                    }
                    wandb.log(log_dict)
                gpu_memory_monitor.reset_peak_stats()
        iter_num += 1
        local_iter_num += 1

        # signals the profiler that the next profiling step has started
        if torch_profiler:
            torch_profiler.step()

        if memory_profiler:
            memory_profiler.step()

        # termination conditions
        if iter_num >= max_iters:
            print_with_time_master("done.")
            break

dist_barrier()
print_with_time_master("removing cache dir.")
if local_cache_dir is not None and local_master_process:
    shutil.rmtree(local_cache_dir, ignore_errors=True)

dist_barrier()
print_with_time_master("done.")
dist_barrier()
if ddp:
    destroy_process_group()

