import datetime
import inspect
import logging
import os
import random
import time
import math
from contextlib import contextmanager, nullcontext

import numpy as np
import torch
import torch.nn.functional as F
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.distributed import init_process_group, destroy_process_group
import torchaudio.functional as aF
from torchaudio.transforms import MelSpectrogram

from discriminator import MultiScaleSTFTDiscriminator
from loss import stft
from model import CirceNet


@contextmanager
def suppress_logging(highest_level=logging.CRITICAL):
    previous_level = logging.root.manager.disable
    logging.disable(highest_level)
    try:
        yield
    finally:
        logging.disable(previous_level)


# data params
data_dir = None
out_dir = None
train_filename = None
val_filename = None
# model params
dimension = 512
n_filters = 64
sample_rate = 48_000
ratios = (8, 5, 4, 4)  # (8, 5, 4, 3, 2)
n_codebooks = 8
causal = False
# misc
n_steps_adv_start = 5000
cycle_sample_rate = True
disc_only = False
gen_only = False
skip_quantization = False
custom_seed_offset = 123
match_val_weights = True
randomize_stft = True
preload_checkpoint = None
preload_optimizer = False
preload_strict = True
preload_checkpoint_adv = None
preload_optimizer_adv = False
preload_strict_adv = True
suppress_compile_warnings = True
eval_interval = 2000
log_interval = 25
eval_iters = 500
eval_only = False # if True, script exits right after the first eval
always_save_checkpoint = True # if True, always save a checkpoint after each eval
# wandb logging
wandb_log = False # disabled by default
wandb_project = "suno"
wandb_run_name = "base"
# data
gradient_accumulation_steps = 1 # used to simulate larger batch sizes
batch_size = 16 # if gradient_accumulation_steps > 1, this is the micro-batch size
# adamw optimizer
learning_rate = 3e-4
max_iters = 250000 # total number of training iterations
beta1 = 0.5
beta2 = 0.9
grad_clip = 1.0 # clip gradients at this value, or disable if == 0.0
# learning rate decay settings
decay_lr = True # whether to decay the learning rate
warmup_iters = 1000 # how many steps to warm up for
lr_decay_iters = None # should be ~= max_iters per Chinchilla
min_lr = 0 # minimum learning rate, should be ~= learning_rate/10 per Chinchilla
# DDP settings
backend = "nccl" # "nccl", "gloo", etc.
# system
device = "cuda" # examples: "cpu", "cuda", "cuda:0", "cuda:1" etc., or try "mps" on macbooks
dtype = "float32" # "float32", "bfloat16", or "float16" (implements a GradScaler)
compile = False # use PyTorch 2.0 to compile the model to be faster

# TODO: technically missing the 0.1 l1_loss in time domain
loss_factor_map = {
    "rec": 25.0,  # 100
    "disc": 1.0,
    "gen_h": 4.0,
    "gen_f": 4.0,
    "comm": 1.0,  # 100
}

# -----------------------------------------------------------------------------
config_keys = [
    k for k,v in globals().items()
    if not k.startswith("_") and isinstance(v, (int, float, bool, str))
]
exec(open("custom_configurator.py").read()) # overrides from command line or config file
config = {k: globals()[k] for k in config_keys} # will be useful for logging
# -----------------------------------------------------------------------------

assert(dtype == "float32")
assert(not (gen_only and disc_only))
assert(gradient_accumulation_steps == 1)

eval_iters = int(eval_iters * gradient_accumulation_steps / 5)
if lr_decay_iters is None:
    lr_decay_iters = max_iters

# various inits, derived attributes, I/O setup
ddp = int(os.environ.get("RANK", -1)) != -1 # is this a ddp run?
if ddp:
    init_process_group(backend=backend)
    ddp_rank = int(os.environ["RANK"])
    ddp_local_rank = int(os.environ["LOCAL_RANK"])
    world_size = torch.distributed.get_world_size()
    device = f"cuda:{ddp_local_rank}"
    torch.cuda.set_device(device)
    master_process = ddp_rank == 0 # this process will do logging, checkpointing etc.
    seed_offset = ddp_rank # each process gets a different seed
else:
    # if not ddp, we are running on a single gpu, and one process
    master_process = True
    seed_offset = 0
    world_size = 1

seed_offset += 1
seed_offset *= custom_seed_offset + 1  # multiply to not just shift
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
# note: float16 data type will automatically use a GradScaler
ptdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16, "float16": torch.float16}[dtype]
ctx = (
    nullcontext()
    if device_type == "cpu"
    else torch.amp.autocast(device_type=device_type, dtype=ptdtype)
)

# logging
if wandb_log and master_process:
    import wandb
    wandb.init(project=wandb_project, name=wandb_run_name, config=config)

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:
    os.makedirs(out_dir, exist_ok=True)
    print(f"logging checkpoint here: {out_dir}")

# load data
train_data = np.memmap(os.path.join(data_dir, train_filename), dtype=np.int16, mode="r")
val_data = np.memmap(os.path.join(data_dir, val_filename), dtype=np.int16, mode="r")

block_size = sample_rate
COMMON_SAMPLE_RATES = [8000, 16000, 24000, 32000, 44100, 48000]


def _cycle_sample_rate(waveform, from_sample_rate=48_000, to_sample_rate=8_000):
    assert(isinstance(waveform, torch.Tensor))
    assert(len(waveform.shape) == 2)
    resampled_waveform = aF.resample(waveform, from_sample_rate, to_sample_rate)
    cycled_waveform = aF.resample(resampled_waveform, to_sample_rate, from_sample_rate)
    assert(waveform.shape == cycled_waveform.shape)
    return cycled_waveform


def get_sample(split, is_training=False):
    data = train_data if split == "train" else val_data
    idx = random.randint(0, len(data)-block_size)
    arr = np.array(data[idx:idx+block_size])
    arr = torch.from_numpy(arr.astype(np.float32) / np.iinfo(np.int16).max)[None]
    if cycle_sample_rate and is_training and random.random() >= 0.75:
        arr = _cycle_sample_rate(
            arr, from_sample_rate=sample_rate, to_sample_rate=random.choice(COMMON_SAMPLE_RATES)
        )
    x = arr
    y = x.clone()
    return x, y


def get_batch(split, is_training=False):
    x_list = []
    y_list = []
    for _ in range(batch_size):
        x, y = get_sample(split, is_training=is_training)
        x_list.append(x)
        y_list.append(y)
    x = torch.stack(x_list)
    y = torch.stack(y_list)
    if device_type == "cuda":
        # pin arrays x,y, which allows us to move them to GPU asynchronously (non_blocking=True)
        x = x.pin_memory().to(device, non_blocking=True)
        y = y.pin_memory().to(device, non_blocking=True)
    else:
        x, y = x.to(device), y.to(device)
    del x_list, y_list
    return x, y


iter_num = 0
best_val_loss = 1e9

# init a new model from scratch
print("Initializing a new model from scratch")
model = CirceNet(
    dimension=dimension,
    n_filters=n_filters,
    ratios=ratios,
    causal=causal,
    skip_quantization=skip_quantization,
    n_codebooks=n_codebooks,
)
model.to(device)
model_adv = MultiScaleSTFTDiscriminator(32)
model_adv.to(device)

# optimizer
use_fused = (device_type == "cuda") and ("fused" in inspect.signature(torch.optim.AdamW).parameters)
print(f"using fused AdamW: {use_fused}")
extra_args = dict(fused=True) if use_fused else dict()
optimizer = torch.optim.AdamW(
    [{"params": model.parameters(), "lr": learning_rate}], betas=(beta1, beta2), **extra_args
)
optimizer_adv = torch.optim.AdamW(
    [{"params": model_adv.parameters(), "lr": learning_rate}], betas=(beta1, beta2), **extra_args
)

# load checkpoint
if preload_checkpoint is not None:
    print("preloading checkpoint")
    checkpoint = torch.load(preload_checkpoint, map_location="cpu")
    if "model" in checkpoint:
        state_dict = checkpoint["model"]
        # fix the keys of the state dictionary :(
        # honestly no idea how checkpoints sometimes get this prefix, have to debug more
        unwanted_prefix = "_orig_mod."
        for k, v in list(state_dict.items()):
            if k.startswith(unwanted_prefix):
                state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k)
    else:
        state_dict = checkpoint
    if not preload_strict:
        # clean up
        orig_state_dict = model.state_dict()
        n_dropped = 0
        n_groups = len(state_dict)
        for k in list(state_dict.keys()):
            if k not in orig_state_dict or state_dict[k].shape != orig_state_dict[k].shape:
                state_dict.pop(k)
                n_dropped += 1
        print(f"dropped {n_dropped}/{n_groups} state dict groups")
        del orig_state_dict
    model.load_state_dict(state_dict, strict=preload_strict)
    if preload_optimizer:
        print("preloading optimizer")
        optimizer.load_state_dict(checkpoint["optimizer"])
    del state_dict, checkpoint
if preload_checkpoint_adv is not None:
    print("preloading checkpoint adv")
    checkpoint = torch.load(preload_checkpoint_adv, map_location="cpu")
    # hack
    if "raw_model_adv" in checkpoint:
        checkpoint["model_adv"] = checkpoint["raw_model_adv"]
    if "model_adv" in checkpoint:
        state_dict = checkpoint["model_adv"]
        # fix the keys of the state dictionary :(
        # honestly no idea how checkpoints sometimes get this prefix, have to debug more
        unwanted_prefix = "_orig_mod."
        for k, v in list(state_dict.items()):
            if k.startswith(unwanted_prefix):
                state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k)
    else:
        state_dict = checkpoint
    if not preload_strict_adv:
        # clean up
        orig_state_dict = model_adv.state_dict()
        n_dropped = 0
        n_groups = len(state_dict)
        for k in list(state_dict.keys()):
            if k not in orig_state_dict or state_dict[k].shape != orig_state_dict[k].shape:
                state_dict.pop(k)
                n_dropped += 1
        print(f"dropped {n_dropped}/{n_groups} state dict groups")
        del orig_state_dict
    model_adv.load_state_dict(state_dict, strict=preload_strict_adv)
    if preload_optimizer:
        print("preloading optimizer")
        optimizer_adv.load_state_dict(checkpoint["optimizer_adv"])
    del state_dict, checkpoint

torch.cuda.empty_cache()
torch.cuda.synchronize()

# compile the model
if compile:
    print("compiling the model... (takes a ~minute)")
    compile_ctx = suppress_logging if suppress_compile_warnings else nullcontext
    with compile_ctx():
        model = torch.compile(model)  # requires PyTorch 2.0
        model_adv = torch.compile(model_adv)

# wrap model into DDP container
if ddp:
    model = DDP(model, device_ids=[ddp_local_rank])
    model_adv = DDP(model_adv, device_ids=[ddp_local_rank])


def reconstruction_loss(x, G_x, eps=1e-7):
    L = 100 * F.mse_loss(x, G_x)  # wav L1 loss
    for i in range(6, 11):
        s = 2**i
        melspec = MelSpectrogram(
            sample_rate=sample_rate,
            n_fft=s,
            hop_length=s // 4,
            n_mels=64,
            wkwargs={"device": x.device}).to(x.device)
        S_x = melspec(x)
        S_G_x = melspec(G_x)
        loss = ((S_x - S_G_x).abs().mean() + (
            ((torch.log(S_x.abs() + eps) - torch.log(S_G_x.abs() + eps))**2
             ).mean(dim=-2)**0.5).mean()) / (i)
        L += loss
    return L


def gen_loss(fmap_real, fmap_fake, logits_real, logits_fake):
    # hinge loss
    loss_h = torch.tensor(0.0, device=logits_fake[0].device, requires_grad=True)
    n_logits = len(logits_fake)
    for lf in logits_fake:
        loss_h = loss_h + F.relu(1 - lf).mean() / n_logits
    # f1 for features
    loss_f = torch.tensor(0.0, device=logits_fake[0].device, requires_grad=True)
    for er, ef in zip(fmap_real, fmap_fake):
        for eer, eef in zip(er, ef):
            # loss_f = loss_f + F.l1_loss(eer, eef)
            loss_f = loss_f + ((eer - eef).abs() / (eer.abs().mean())).mean()
    # missing factor of 100
    loss_f = loss_f / (len(logits_fake) * len(logits_fake[0]))
    # sim loss
    loss_s = torch.tensor(0.0, device=logits_fake[0].device, requires_grad=True)
    for lr, lf in zip(logits_real, logits_fake):
        loss_s = loss_s + F.mse_loss(lr, lf) / len(logits_fake)
    loss_s = loss_s / len(logits_fake)
    return loss_h, loss_f + loss_s


def disc_loss(logits_real, logits_fake):
    loss = torch.tensor(0.0, device=logits_fake[0].device, requires_grad=True)
    n_logits = len(logits_fake)
    for lr, lf in zip(logits_real, logits_fake):
        loss = loss + (F.relu(1-lr) + F.relu(1+lf)).mean() / n_logits
    return loss


# helps estimate an arbitrarily accurate loss over either split using many batches
@torch.no_grad()
def estimate_loss():
    loss_types = ["rec_loss", "disc_loss", "gen_h_loss", "gen_f_loss", "comm_loss"]
    n_loss_modalities = 1 + 1  # train + val
    effective_eval_iters = int(round(eval_iters / n_loss_modalities))
    model.eval()
    n_loss_entries = n_loss_modalities * len(loss_types)
    loss_tensor = torch.zeros(n_loss_entries, device=device)
    tensor_keys = []
    n_loss_entry = 0
    for split in ["train", "val"]:
        losses = [[] for _ in loss_types]
        for k in range(effective_eval_iters):
            X, Y = get_batch(split)
            with ctx:
                y_pred, loss_dict = model(X, Y)
                y_adv_real, fmap_real = model_adv(Y)
                y_adv_fake, fmap_fake = model_adv(y_pred)
                loss_adv = disc_loss(y_adv_real, y_adv_fake)
                loss_gen_h, loss_gen_f = gen_loss(fmap_real, fmap_fake, y_adv_real, y_adv_fake)
            # losses[0].append(reconstruction_loss(Y, y_pred).item())
            losses[0].append((loss_dict["sc_loss"] + loss_dict["mag_loss"]).item())
            losses[1].append(loss_adv.item())
            losses[2].append(loss_gen_h.item())
            losses[3].append(loss_gen_f.item())
            losses[4].append(0 if loss_dict["comm_loss"] is None else loss_dict["comm_loss"].item())
        for loss_type_idx, loss_type in enumerate(loss_types):
            loss_tensor[n_loss_entry] = np.mean(losses[loss_type_idx])
            tensor_keys.append(f"{split}/{loss_type}")
            n_loss_entry += 1
    if ddp:
        torch.distributed.all_reduce(loss_tensor, op=torch.distributed.ReduceOp.AVG)
    out = {k: loss_tensor[n].item() for n, k in enumerate(tensor_keys)}
    # add extra loss items
    for split in ["train", "val"]:
        out[f"{split}/loss"] = (
            out[f"{split}/rec_loss"] + out[f"{split}/gen_h_loss"] + out[f"{split}/gen_f_loss"]
        )
    model.train()
    for name, module in model.named_modules():
        assert module.training  # make sure we can undo everything
    return out


# save spec images
def _get_im(x, fft_size=2048, win_length=1024, hop_size=64*2):
    assert(isinstance(x, torch.Tensor))
    assert(len(x.shape) == 2)
    assert(x.shape[0] == 1)
    x = x.detach().cpu()
    window = torch.hann_window(win_length)
    x_mag = stft(x, fft_size, hop_size, win_length, window).numpy()[0]
    im_data = np.fliplr(np.log(x_mag)).T
    del window
    return im_data


@torch.no_grad()
def _get_images(n_samples=10):
    model.eval()
    im_data_real = []
    im_data_fake = []
    for n in range(n_samples):
        offs = int(len(val_data) * n / n_samples)
        x = torch.from_numpy(
            np.array(val_data[offs:offs+block_size])[None].astype(np.float32) /
            np.iinfo(np.int16).max
        )
        x_cycle, _ = model(x[None])
        x_cycle = x_cycle[0]
        im_data_real.append(_get_im(x))
        im_data_fake.append(_get_im(x_cycle))
    model.train()
    for name, module in model.named_modules():
        assert module.training  # make sure we can undo everything
    return im_data_real, im_data_fake


# learning rate decay scheduler (cosine with warmup)
def get_lr(it):
    # 1) linear warmup for warmup_iters steps
    if it < warmup_iters:
        return learning_rate * it / warmup_iters
    # 2) if it > lr_decay_iters, return min learning rate
    if it > lr_decay_iters:
        return min_lr
    # 3) in between, use cosine decay down to min learning rate
    decay_ratio = (it - warmup_iters) / (lr_decay_iters - warmup_iters)
    assert 0 <= decay_ratio <= 1
    coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio)) # coeff ranges 0..1
    return min_lr + coeff * (learning_rate - min_lr)


# training loop
X, Y = get_batch("train", is_training=True) # fetch the very first batch
t0 = time.time()
t00 = time.time()
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
raw_model_adv = model_adv.module if ddp else model_adv
mfu = 0
tokens_per_s = 0
running_loss = []
running_loss_adv = []
# with torch.autograd.set_detect_anomaly(True):
while True:

    # determine and set the learning rate for this iteration
    lr = get_lr(iter_num) if decay_lr else learning_rate
    for param_group in optimizer.param_groups:
        param_group["lr"] = lr
    for param_group in optimizer_adv.param_groups:
        param_group["lr"] = lr

    # evaluate the loss on train/val sets and write checkpoints
    if iter_num % eval_interval == 0:
        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)
        if master_process:
            print(
                f"loss estimation took {estimation_time:.1f} seconds."
                f" ({eval_time_pct:.1f}% of loop)"
            )
            print(
                f"step {iter_num}: train loss {losses['train/loss']:.4f},"
                f" val loss {losses['val/loss']:.4f}"
            )
            if wandb_log:
                log_dict = {
                    "iter": iter_num,
                    "lr": lr,
                    "mfu": mfu,  # convert to percentage
                    "tok/s": tokens_per_s,
                }
                for k, v in losses.items():
                    log_dict[k] = v
                im_data_real, im_data_fake = _get_images()
                for n, im_data in enumerate(im_data_real):
                    log_dict[f"image/real_{n}"] = wandb.Image(im_data)
                for n, im_data in enumerate(im_data_fake):
                    log_dict[f"image/fake_{n}"] = wandb.Image(im_data)
                wandb.log(log_dict)
            if losses["val/loss"] < best_val_loss or always_save_checkpoint:
                if iter_num > 0:
                    checkpoint = {
                        "model": raw_model.state_dict(),
                        "model_adv": raw_model_adv.state_dict(),
                        "optimizer": optimizer.state_dict(),
                        "optimizer_adv": optimizer_adv.state_dict(),
                        "iter_num": iter_num,
                        "best_val_loss": losses["val/loss"],
                        "config": config,
                    }
                    print(f"saving checkpoint to {out_dir}")
                    if losses["val/loss"] < best_val_loss:
                        torch.save(checkpoint, os.path.join(out_dir, "best_ckpt.pt"))
                    if always_save_checkpoint:
                        torch.save(checkpoint, os.path.join(out_dir, "last_ckpt.pt"))
                        reduced_checkpoint = {
                            k: v
                            for k, v in checkpoint.items()
                            if k in ["model", "config", "best_val_loss"]
                        }
                        torch.save(
                            reduced_checkpoint, os.path.join(out_dir, "last_ckpt_infer.pt")
                        )
                if losses["val/loss"] < best_val_loss:
                    best_val_loss = losses["val/loss"]

    if iter_num == 0 and eval_only:
        break

    # forward backward update, with optional gradient accumulation to simulate larger batch size
    # and using the GradScaler if data type is float16
    if ddp:
        # in DDP training we only need to sync gradients at the last micro step.
        # the official way to do this is with model.no_sync() context manager, but
        # I really dislike that this bloats the code and forces us to repeat code
        # looking at the source of that context manager, it just toggles this variable
        model.require_backward_grad_sync = True
        model_adv.require_backward_grad_sync = True

    with ctx:
        for update_type in ["generator", "discriminator"]:
            model.zero_grad()
            model_adv.zero_grad()
            y_pred, loss_dict = model(X, Y, randomize_stft=randomize_stft)
            if update_type == "generator":
                # update generator

                # reconstruction loss
                # loss_rec = reconstruction_loss(Y, y_pred)
                loss_rec = loss_dict["sc_loss"] + loss_dict["mag_loss"]
                if skip_quantization:
                    loss_comm = loss_dict["comm_loss"]
                    assert(loss_comm is None)
                else:
                    loss_comm = loss_dict["comm_loss"]

                # generative loss
                y_adv_real, fmap_real = model_adv(Y)
                y_adv_fake, fmap_fake = model_adv(y_pred)
                loss_gen_h, loss_gen_f = gen_loss(fmap_real, fmap_fake, y_adv_real, y_adv_fake)

                if not disc_only:
                    loss = loss_rec * loss_factor_map["rec"]
                    if loss_comm is not None:
                        loss = loss + loss_comm * loss_factor_map["comm"]
                    if not gen_only and iter_num >= n_steps_adv_start:
                        loss = (
                            loss +
                            loss_gen_h * loss_factor_map["gen_h"] +
                            loss_gen_f * loss_factor_map["gen_f"]
                        )
                    loss.backward()
                    for _, param in model.named_parameters():
                        if torch.isnan(param.grad).any() or torch.isinf(param.grad).any():
                            print("nan found, setting to 0")
                        param.grad = torch.nan_to_num(param.grad, nan=0.0, posinf=0.0, neginf=0.0)
                    if grad_clip != 0.0:
                        grad_norm = torch.nn.utils.clip_grad_norm_(
                            model.parameters(), grad_clip, error_if_nonfinite=True
                        )
                    else:
                        grad_norm = 0
                    optimizer.step()
            else:
                # update discriminator
                if iter_num < n_steps_adv_start or iter_num % 2 == 0:
                    # only update every other step
                    loss_adv = None
                    grad_norm_adv = None
                    continue
                y_adv_real, _ = model_adv(Y)
                y_adv_fake, _ = model_adv(y_pred.detach())
                loss_adv = disc_loss(y_adv_real, y_adv_fake)
                if not gen_only:
                    (loss_adv * loss_factor_map["disc"]).backward()
                    for _, param in model_adv.named_parameters():
                        if torch.isnan(param.grad).any() or torch.isinf(param.grad).any():
                            print("nan found, setting to 0")
                            param.grad = torch.nan_to_num(
                                param.grad, nan=0.0, posinf=0.0, neginf=0.0
                            )
                    if grad_clip != 0.0:
                        grad_norm_adv = torch.nn.utils.clip_grad_norm_(
                            model_adv.parameters(), grad_clip, error_if_nonfinite=True
                        )
                    else:
                        grad_norm_adv = 0
                    optimizer_adv.step()
            optimizer.zero_grad(set_to_none=True)
            optimizer_adv.zero_grad(set_to_none=True)

    running_loss.append(loss_rec.item())
    if loss_adv is not None:
        running_loss_adv.append(loss_adv.item())

    # immediately async prefetch next batch while model is doing the forward pass on the GPU
    X, Y = get_batch("train", is_training=True)

    if master_process and wandb_log:
        log_dict = {
            "misc/loss_recon": loss_rec.item(),
            "misc/loss_gen_h": loss_gen_h.item(),
            "misc/loss_gen_f": loss_gen_f.item(),
            "misc/grad_norm": grad_norm.item(),
        }
        if grad_norm_adv is not None:
            log_dict["misc/grad_norm_adv"] = grad_norm_adv.item()
        if loss_adv is not None:
            log_dict["misc/loss_adv"] = loss_adv.item()
        if loss_comm is not None:
            log_dict["misc/loss_comm"] = loss_comm.item()
        wandb.log(log_dict)

    # timing and logging
    t1 = time.time()
    dt = t1 - t0
    t0 = t1
    if iter_num % log_interval == 0 and master_process:
        if local_iter_num >= 5:  # let the training loop settle a bit
            mfu = raw_model.estimate_mfu(batch_size * gradient_accumulation_steps, dt)
            tokens_per_s = (
                world_size * batch_size * block_size * gradient_accumulation_steps / dt
            )
        avg_loss = np.mean(running_loss)
        avg_loss_adv = np.mean(running_loss_adv)
        running_loss = []
        running_loss_adv = []
        print(
            f"iter {iter_num}: loss_recon {avg_loss:.3f}, loss_adv {avg_loss_adv:.3f},"
            f" step_time {dt*1000:.1f}ms, mfu {mfu:.1f},"
            f" throughput {tokens_per_s/1e3:,.0f}k tok/s"
        )
    iter_num += 1
    local_iter_num += 1

    # termination conditions
    if iter_num > max_iters:
        break

if ddp:
    destroy_process_group()
