import os
import gc
import json
import time
import math
import wandb
import torch
import random
import datetime
import argparse
import numpy as np
import torch.nn as nn
import torch.distributed as dist

from tqdm import tqdm
from colorama import Fore, Style
from datetime import timedelta
from tokenizers import Tokenizer
from torch.nn import functional as F
from suno_utils.utils.text import read_jsonl
from torch.distributed import barrier, is_initialized, init_process_group

from prefix_model.model import DiffusionTransformer

RUN_START_TIME = time.strftime("%Y-%m-%d_%H-%M-%S")

CHECKPOINT_DIR = f"/app2/suno/checkpoints/{RUN_START_TIME}_s{random.randint(0, 9999)}"
os.makedirs(CHECKPOINT_DIR, exist_ok=True)


def is_ddp():
    return int(os.environ.get("RANK", -1)) != -1


def is_master():
    if is_ddp():
        return int(os.environ["RANK"]) == 0
    return True


def dist_barrier():
    if is_initialized():
        barrier()


def print_with_time(content):
    """Print the content with the current time."""
    print(f"[{datetime.datetime.now().strftime('%Y-%m-%d_%H:%M:%S')}]: {content}")


def print_with_time_master(content):
    if is_master():
        print_with_time(content)


FULL_PRECISION_KEY_FRAGMENTS = ("pos_emb", "inv_freq")


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 setup_distributed(master_addr, master_port):
    # Set environment variables based on parsed arguments
    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']}"
            )
        rank = int(os.environ["SLURM_PROCID"])
        local_rank = int(os.environ["SLURM_LOCALID"])
        world_size = int(os.environ["SLURM_JOB_NUM_NODES"]) * int(
            os.environ["SLURM_NTASKS_PER_NODE"]
        )
    else:  # Running locally
        rank = 0
        local_rank = 0
        world_size = 1
    os.environ["RANK"] = str(rank)
    os.environ["LOCAL_RANK"] = str(local_rank)
    print(f"Initializing distributed process group on rank {local_rank}")
    torch.cuda.set_device(local_rank)
    try:
        dist.init_process_group(
            backend="nccl",
            timeout=timedelta(hours=6),
            rank=rank,
            world_size=world_size,
            device_id=torch.device(f"cuda:{local_rank}"),
        )
    except Exception as e:
        print(f"Distributed error on rank {rank} with host {os.environ['HOSTNAME']}")
        raise e
    print(f"Done initializing on rank: {dist.get_rank()}")
    # barrier to check if nccl is working
    dist_barrier()
    print_with_time_master("distributed setup ready.")


def apply_specaugment_mask(
    latents: np.ndarray,
    time_mask_range=(20, 50),
    feature_mask_range=(8, 24),
    num_time_masks=1,
    num_feature_masks=1,
    time_mask_prob=0.5,
    feature_mask_prob=0.5,
    mask_value=0.0,
):
    """
    Apply SpecAugment-style masking to VAE latents.

    Args:
        latents (np.ndarray): Array of shape (T, D).
        time_mask_range (tuple): (min, max) width of time masks.
        feature_mask_range (tuple): (min, max) width of feature masks.
        num_time_masks (int): Number of time masks to attempt.
        num_feature_masks (int): Number of feature masks to attempt.
        time_mask_prob (float): Probability of applying each time mask.
        feature_mask_prob (float): Probability of applying each feature mask.
        mask_value (float): Value to use for masking (default: 0.0).
    Returns:
        np.ndarray: Masked latents.
    """
    T, D = latents.shape
    latents = latents.copy()  # avoid modifying original array

    # Time masking
    for _ in range(num_time_masks):
        if np.random.rand() < time_mask_prob:
            mask_width = np.random.randint(*time_mask_range)
            if T - mask_width > 0:
                t = np.random.randint(0, T - mask_width)
                latents[t : t + mask_width, :] = mask_value

    # Feature masking
    for _ in range(num_feature_masks):
        if np.random.rand() < feature_mask_prob:
            mask_width = np.random.randint(*feature_mask_range)
            if D - mask_width > 0:
                f = np.random.randint(0, D - mask_width)
                latents[:, f : f + mask_width] = mask_value

    return latents


class RewardModelMemmapDataset(torch.utils.data.Dataset):
    def __init__(
        self,
        dataset_dir,
        metas_filename,
        vae_memmap_filename,
        vae_scale_factor,
        mask_prob=0.5,
        vae_use_float16=True,
        vae_n_tokens=750,
        vae_dim=128,
    ):
        self.metas_filepath = os.path.join(dataset_dir, metas_filename)
        # self.metas = read_jsonl(self.metas_filepath)
        self.vae_scale_factor = vae_scale_factor
        self.mask_prob = mask_prob

        # print(f"Loaded {len(self.metas)} metas")

        # load the memmap files
        vae_data = np.memmap(
            os.path.join(dataset_dir, vae_memmap_filename),
            dtype=np.float16 if vae_use_float16 else np.float32,
            mode="r",
        )
        vae_data = vae_data.reshape(-1, vae_n_tokens, vae_dim)
        self.vae_data = vae_data
        print(self.vae_data.shape)

        if self.vae_data.shape[0] < 256:
            # repeat the data 256 times
            self.vae_data = np.concatenate([self.vae_data] * 2, axis=0)
            print(f"Repeating data x2 to {self.vae_data.shape}")

        # assert len(self.metas) == self.vae_data.shape[0]

    def __len__(self):
        return self.vae_data.shape[0]

    def __getitem__(self, idx):

        # only use even indices
        # 0, 2, 4, ...
        # so we have to convert idx to an even index using modulo
        # if idx is odd, we need to subtract 1
        if idx % 2 == 1:
            idx -= 1

        # meta = self.metas[idx]
        negative_latents = self.vae_data[idx] * self.vae_scale_factor
        positive_latents = self.vae_data[idx + 1] * self.vae_scale_factor

        if self.mask_prob > 0:
            positive_latents = apply_specaugment_mask(
                positive_latents,
                time_mask_prob=self.mask_prob,
                feature_mask_prob=self.mask_prob,
            )
            negative_latents = apply_specaugment_mask(
                negative_latents,
                time_mask_prob=self.mask_prob,
                feature_mask_prob=self.mask_prob,
            )

        negative_latents = torch.from_numpy(negative_latents).float()
        positive_latents = torch.from_numpy(positive_latents).float()

        return positive_latents, negative_latents


def load_tokenizer(
    tokenizer_filepath="s3://suno-data/georg/models/tokenizers/tokenizer_60k.json",
):
    tokenizer = Tokenizer.from_file(tokenizer_filepath)
    tokenizer.add_special_tokens(["\n"])
    tokenizer.pad_idx = tokenizer.token_to_id("[PAD]")
    return tokenizer


class RewardModelDataset(torch.utils.data.Dataset):
    def __init__(
        self,
        metas_filepath,
        vae_scale_factor,
        chunk_size=750,
        mask_prob=0.5,
        use_preference_labels=True,
        cond_text_len=1536,
    ):
        self.metas_filepath = metas_filepath
        self.metas = read_jsonl(metas_filepath)
        self.vae_scale_factor = vae_scale_factor
        self.chunk_size = chunk_size
        self.mask_prob = mask_prob
        self.use_preference_labels = use_preference_labels
        self.cond_text_len = cond_text_len

        print(f"Loaded {len(self.metas)} metas")

    def __len__(self):
        return len(self.metas)

    def __getitem__(self, idx):
        meta = self.metas[idx]
        positive_latents = np.load(meta["pos_vae_latents_filepath"])[
            "vae_latents"
        ].astype(np.float32)
        negative_latents = np.load(meta["neg_vae_latents_filepath"])[
            "vae_latents"
        ].astype(np.float32)

        # semantic_codes = np.load(meta["semantic_codes_filepath"])
        # semantic_codes = semantic_codes[:, 0]
        # semantic_codes = torch.from_numpy(semantic_codes).long()

        if self.use_preference_labels:
            # If latents are smaller than chunk size, pad with zeros instead of random crop
            if positive_latents.shape[0] < self.chunk_size:
                pad_width = self.chunk_size - positive_latents.shape[0]
                positive_latents = np.pad(
                    positive_latents,
                    ((0, pad_width), (0, 0)),
                    mode="constant",
                    constant_values=0,
                )
                negative_latents = np.pad(
                    negative_latents,
                    ((0, pad_width), (0, 0)),
                    mode="constant",
                    constant_values=0,
                )
            else:
                # select a random chunk from the positive latents
                start_idx = np.random.randint(
                    0, positive_latents.shape[0] - self.chunk_size + 1
                )
                end_idx = start_idx + self.chunk_size
                positive_latents = positive_latents[start_idx:end_idx]
                negative_latents = negative_latents[start_idx:end_idx]
                # semantic_codes = semantic_codes[start_idx:end_idx]

        else:
            # sometimes use the positive, sometimes use the negative
            latents = positive_latents if np.random.random() < 0.5 else negative_latents
            # now construct pairs
            # the positive latent will be an earlier chunk
            # the negative latent will be a later chunk
            max_start_for_positive = latents.shape[0] - (self.chunk_size * 2)
            if max_start_for_positive < 0:
                # fallback: just take the first and last chunk
                positive_latents = latents[: self.chunk_size]
                negative_latents = latents[-self.chunk_size :]
            else:
                pos_start_idx = np.random.randint(0, max_start_for_positive + 1)
                pos_end_idx = pos_start_idx + self.chunk_size
                positive_latents = latents[pos_start_idx:pos_end_idx]

                neg_start_min = pos_end_idx
                neg_start_max = latents.shape[0] - self.chunk_size
                if neg_start_min >= neg_start_max:
                    neg_start_idx = neg_start_min
                else:
                    neg_start_idx = np.random.randint(neg_start_min, neg_start_max + 1)
                neg_end_idx = neg_start_idx + self.chunk_size
                negative_latents = latents[neg_start_idx:neg_end_idx]

        positive_latents = apply_specaugment_mask(
            positive_latents,
            time_mask_prob=self.mask_prob,
            feature_mask_prob=self.mask_prob,
        )
        negative_latents = apply_specaugment_mask(
            negative_latents,
            time_mask_prob=self.mask_prob,
            feature_mask_prob=self.mask_prob,
        )

        # apply the vae scale factor to move to std ~1
        positive_latents = torch.from_numpy(positive_latents) * self.vae_scale_factor
        negative_latents = torch.from_numpy(negative_latents) * self.vae_scale_factor
        positive_latents = positive_latents.float()
        negative_latents = negative_latents.float()

        return (
            positive_latents,  # .permute(1, 0),
            negative_latents,  # .permute(1, 0),
            # semantic_codes,
        )


class SinusoidalPositionalEncoding(nn.Module):
    def __init__(self, dim, max_len=2048):
        super().__init__()
        pe = torch.zeros(max_len, dim)
        position = torch.arange(0, max_len).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, dim, 2) * -(math.log(10000.0) / dim))
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)
        self.register_buffer("pe", pe)

    def forward(self, x):
        # x: (B, T, D)
        seq_len = x.size(1)
        return x + self.pe[:seq_len].unsqueeze(0).to(x.dtype)  # (1, T, D)


@torch.no_grad()
def mae_random_masking(x, mask_ratio=0.75):
    """
    x: (B, N, C) tokens BEFORE the encoder
    Returns:
      x_vis: (B, N_keep, C)
      mask: (B, N) bool, True where masked
      ids_restore: (B, N) to restore original order
      ids_keep: (B, N_keep) indices of visible tokens
    """
    B, N, C = x.shape
    N_keep = int(N * (1.0 - mask_ratio))

    noise = torch.rand(B, N, device=x.device)
    ids_shuffle = torch.argsort(noise, dim=1)  # ascending
    ids_restore = torch.argsort(ids_shuffle, dim=1)

    ids_keep = ids_shuffle[:, :N_keep]
    x_vis = torch.gather(x, 1, ids_keep.unsqueeze(-1).expand(-1, -1, C))

    mask = torch.ones(B, N, device=x.device, dtype=torch.bool)
    mask.scatter_(1, ids_keep, False)
    mask = torch.gather(mask, 1, ids_restore)  # unshuffle to original order
    return x_vis, mask, ids_restore, ids_keep


def mae_prepare_decoder_input(z_enc, ids_restore, mask_token):
    """
    z_enc: (B, N_keep, C_d) encoder features already projected to decoder dim
    ids_restore: (B, N) from mae_random_masking
    mask_token: (1, 1, C_d)
    returns z_dec_in: (B, N, C_d) in original order, masked slots filled with mask_token
    """
    B, N = ids_restore.shape
    C_d = z_enc.size(-1)
    N_keep = z_enc.size(1)
    N_mask = N - N_keep

    mask_tokens = mask_token.expand(B, N_mask, C_d)
    z_ = torch.cat([z_enc, mask_tokens], dim=1)  # concat then unshuffle
    z_dec_in = torch.gather(z_, 1, ids_restore.unsqueeze(-1).expand(-1, -1, C_d))
    return z_dec_in


class LearnablePositionalEncoding(nn.Module):
    def __init__(self, dim, max_len=8192):
        super().__init__()
        self.pe = nn.Parameter(torch.zeros(1, max_len, dim))
        nn.init.trunc_normal_(self.pe, std=0.02)

    def forward(self, x):  # x: (B, T, D)
        return x + self.pe[:, : x.size(1)]


class MAEEncoder(nn.Module):
    def __init__(
        self,
        in_dim,
        embed_dim=768,
        num_layers=12,
        num_heads=12,
        ff_dim=3072,
        dropout=0.0,
        max_len=8192,
        pos_type="learned",
    ):
        super().__init__()
        self.in_proj = nn.Linear(in_dim, embed_dim)
        self.pos = (
            LearnablePositionalEncoding(embed_dim, max_len)
            if pos_type == "learned"
            else SinusoidalPositionalEncoding(embed_dim, max_len)
        )
        enc_layer = nn.TransformerEncoderLayer(
            d_model=embed_dim,
            nhead=num_heads,
            dim_feedforward=ff_dim,
            dropout=dropout,
            batch_first=True,
        )
        self.encoder = nn.TransformerEncoder(enc_layer, num_layers=num_layers)
        self.norm = nn.LayerNorm(embed_dim)

    def forward(self, x_vis, key_padding_mask=None):
        """
        x_vis: (B, N_keep, in_dim) tokens for encoder (visible only)
        key_padding_mask: (B, N_keep) True where PAD (optional)
        """
        h = self.in_proj(x_vis)
        h = self.pos(h)
        h = self.encoder(h, src_key_padding_mask=key_padding_mask)
        return self.norm(h)  # (B, N_keep, E)


class MAEDecoder(nn.Module):
    def __init__(
        self,
        out_dim,
        dec_dim=512,
        num_layers=8,
        num_heads=16,
        ff_dim=2048,
        dropout=0.0,
        max_len=8192,
        pos_type="learned",
    ):
        super().__init__()
        self.pos = (
            LearnablePositionalEncoding(dec_dim, max_len)
            if pos_type == "learned"
            else SinusoidalPositionalEncoding(dec_dim, max_len)
        )
        dec_layer = nn.TransformerEncoderLayer(
            d_model=dec_dim,
            nhead=num_heads,
            dim_feedforward=ff_dim,
            dropout=dropout,
            batch_first=True,
        )
        self.decoder = nn.TransformerEncoder(dec_layer, num_layers=num_layers)
        self.norm = nn.LayerNorm(dec_dim)
        self.pred = nn.Linear(dec_dim, out_dim)  # reconstruct original token

        # learned mask token in decoder space
        self.mask_token = nn.Parameter(torch.zeros(1, 1, dec_dim))
        nn.init.trunc_normal_(self.mask_token, std=0.02)

    def forward(self, z_enc_proj, ids_restore):
        """
        z_enc_proj: (B, N_keep, dec_dim) encoder feats already projected to dec_dim
        ids_restore: (B, N)
        returns y_pred: (B, N, out_dim)
        """
        z_in = mae_prepare_decoder_input(z_enc_proj, ids_restore, self.mask_token)
        z_in = self.pos(z_in)
        z = self.decoder(z_in)
        z = self.norm(z)
        return self.pred(z)  # (B, N, out_dim)


class MAEModel(nn.Module):
    """
    Input tokens are VAE latents per patch: x: (B, N, vae_dim)
    We reconstruct those latents with MSE.
    """

    def __init__(
        self,
        vae_dim,
        enc_embed_dim=768,
        enc_layers=12,
        enc_heads=12,
        enc_ff=3072,
        dec_embed_dim=512,
        dec_layers=8,
        dec_heads=16,
        dec_ff=2048,
        dropout=0.0,
        max_len=750,
        pos_type="learned",
        proj_to_dec=True,
    ):
        super().__init__()
        self.vae_dim = vae_dim
        self.encoder = MAEEncoder(
            vae_dim,
            enc_embed_dim,
            enc_layers,
            enc_heads,
            enc_ff,
            dropout,
            max_len,
            pos_type,
        )
        self.proj_to_dec = (
            nn.Linear(enc_embed_dim, dec_embed_dim) if proj_to_dec else nn.Identity()
        )
        self.decoder = MAEDecoder(
            vae_dim,
            dec_embed_dim,
            dec_layers,
            dec_heads,
            dec_ff,
            dropout,
            max_len,
            pos_type,
        )

    def forward(self, x, mask_ratio=0.75, pad_mask=None):
        """
        x: (B, N, vae_dim)
        pad_mask: (B, N) True where PAD (optional, rare for images)
        Returns:
          loss, dict
        """
        # 1) random masking (no mask tokens yet)
        x_vis, mask, ids_restore, ids_keep = mae_random_masking(x, mask_ratio)

        # 2) encoder on visible tokens only
        kp_vis = pad_mask.gather(1, ids_keep) if pad_mask is not None else None
        z = self.encoder(x_vis, key_padding_mask=kp_vis)  # (B, N_keep, E)

        # 3) project to decoder and reinsert masked positions with a learned token
        z_dec_in = self.proj_to_dec(z)  # (B, N_keep, D)
        y_pred = self.decoder(z_dec_in, ids_restore)  # (B, N, vae_dim)

        # 4) compute reconstruction loss only on masked positions (MAE default)
        if pad_mask is None:
            valid = torch.ones_like(mask, dtype=torch.bool)
        else:
            valid = ~pad_mask  # True where real tokens

        recon_mask = mask & valid
        # (optional) normalize targets (channel-wise or per-token); here: none
        loss = ((y_pred - x) ** 2).mean(dim=-1)  # (B, N)
        # avoid empty sets
        denom = recon_mask.sum().clamp_min(1)
        loss = (loss * recon_mask.float()).sum() / denom

        return loss, {
            "y_pred": y_pred,
            "mask": mask,
            "ids_restore": ids_restore,
        }

    @torch.no_grad()
    def encode(self, x, mask_ratio=0.0, pad_mask=None):
        """
        Feature extraction: by default don’t mask (mask_ratio=0).
        If you want stochastic features, pass a small mask_ratio.
        """
        if mask_ratio > 0:
            x_vis, _, _, ids_keep = mae_random_masking(x, mask_ratio)
            kp_vis = pad_mask.gather(1, ids_keep) if pad_mask is not None else None
            z = self.encoder(x_vis, kp_vis)
        else:
            z = self.encoder(x, pad_mask)
        return z  # (B, N, enc_dim)


class MAEReward(nn.Module):
    def __init__(self, mae: MAEModel, pool="mean"):
        super().__init__()
        self.backbone = mae.encoder
        for p in self.backbone.parameters():
            p.requires_grad = False

        E = self.backbone.norm.normalized_shape[0]  # encoder embed dim
        self.pool = pool
        if pool == "attn":
            self.pool_query = nn.Parameter(torch.randn(1, 1, E))
            self.pool_proj = nn.Linear(E, E, bias=False)
        self.head = nn.Sequential(nn.Linear(E, E), nn.Tanh(), nn.Linear(E, 1))

    @torch.no_grad()
    def encode(self, x, pad_mask=None):
        return self.backbone(x, key_padding_mask=pad_mask)  # (B, N, E)

    def _pool(self, h, pad_mask=None):
        if self.pool == "mean":
            if pad_mask is None:
                return h.mean(dim=1)
            keep = (~pad_mask).float().unsqueeze(-1)
            return (h * keep).sum(dim=1) / keep.sum(dim=1).clamp_min(1.0)
        # attn pooling
        q = self.pool_query.expand(h.size(0), -1, -1)
        k = self.pool_proj(h)
        attn = (q @ k.transpose(1, 2)) / (h.size(-1) ** 0.5)
        if pad_mask is not None:
            attn = attn.masked_fill(pad_mask.unsqueeze(1), float("-inf"))
        attn = torch.softmax(attn, dim=-1)
        return (attn @ h).squeeze(1)

    def forward(self, x, pad_mask=None):
        with torch.no_grad():
            h = self.encode(x, pad_mask)  # (B, N, E)
        g = self._pool(h, pad_mask)
        return self.head(g).squeeze(-1)


def bradley_terry_loss(
    r_i: torch.Tensor, r_j: torch.Tensor, labels: torch.Tensor
) -> torch.Tensor:
    """
    Compute Bradley-Terry loss for paired comparisons.

    Args:
        r_i: Logits/scores for first options in pairs, shape (batch_size,)
        r_j: Logits/scores for second options in pairs, shape (batch_size,)
        labels: Binary tensor indicating whether first option (0) or second option (1)
               was preferred, shape (batch_size,)

    Returns:
        Mean loss value as a torch.Tensor
    """
    # Compute negative log likelihood using logsigmoid for numerical stability
    loss = -(
        (labels) * F.logsigmoid(r_j - r_i) + (1 - labels) * F.logsigmoid(r_i - r_j)
    )
    return loss.mean()


def train_step(model, batch, pretrain=False):
    model.train()
    pos_vae, neg_vae = batch
    pos_vae = pos_vae.cuda()
    neg_vae = neg_vae.cuda()

    if pretrain:

        acc = torch.tensor(0.0)

        batch_z = pos_vae if np.random.random() < 0.5 else neg_vae

        loss, _ = model(batch_z, mask_ratio=0.7)

        return loss, acc

    else:
        pos_logits = model(pos_vae)
        neg_logits = model(neg_vae)

        labels = torch.zeros_like(pos_logits)
        loss = bradley_terry_loss(pos_logits, neg_logits, labels)

        # compute accuracy
        preds = (pos_logits > neg_logits).float()
        acc = (preds == torch.ones_like(preds)).float().mean()

    return loss, acc


def validate(model, dataloader, pretrain=False):
    model.eval()
    total_loss = 0
    total_acc = 0
    for idx, batch in enumerate(dataloader):
        with torch.no_grad():
            loss, acc = train_step(model, batch, pretrain=pretrain)
            total_loss += loss.item()
            total_acc += acc.item()
    return total_loss / len(dataloader), total_acc / len(dataloader)


def log_metrics(master_process, loss, grad_norm, acc, world_size, global_step, epoch):
    metrics = torch.tensor([loss.item(), grad_norm, acc]).to(loss.device)
    dist.all_reduce(metrics, op=dist.ReduceOp.SUM)
    metrics = metrics / world_size

    avg_loss, avg_grad_norm, avg_acc = metrics[:3]

    if master_process:
        log_data = {
            "train/loss": avg_loss.item(),
            "train/grad_norm": avg_grad_norm.item(),
            "train/acc": avg_acc.item(),
            "trainer/global_step": global_step,
            "trainer/epoch": epoch,
        }
        wandb.log(log_data)
        print_with_time_master(
            f"{Fore.GREEN}Training metrics:{Style.RESET_ALL} "
            + ", ".join(
                [
                    (
                        f"{Fore.YELLOW}{k.replace('train/', '')}:{Style.RESET_ALL} {v:.6f}"
                        if isinstance(v, float)
                        else f"{Fore.YELLOW}{k.replace('train/', '')}:{Style.RESET_ALL} {v}"
                    )
                    for k, v in log_data.items()
                ]
            )
        )


def log_val_metrics(master_process, loss, acc, global_step, epoch):
    if master_process:
        log_data = {
            "val/loss": loss,
            "val/acc": acc,
            "trainer/global_step": global_step,
            "trainer/epoch": epoch,
        }
        wandb.log(log_data)
        print_with_time_master(
            f"{Fore.RED}Validation metrics:{Style.RESET_ALL} "
            + ", ".join(
                [
                    f"{Fore.YELLOW}{k.replace('val/', '')}:{Style.RESET_ALL} {v:.6f}"
                    for k, v in log_data.items()
                ]
            )
        )


def save_checkpoint(
    model,
    optimizer,
    lr_scheduler,
    global_step,
    checkpoint_dir,
    run_config,
    epoch,
    ckpt_name="last_ckpt.pt",
):
    checkpoint = {
        "model": model.state_dict(),
        "optimizer": optimizer.state_dict(),
        "lr_scheduler": lr_scheduler.state_dict(),
        "global_step": global_step,
        "run_config": run_config,
        "epoch": epoch,
    }
    checkpoint_path = os.path.join(checkpoint_dir, ckpt_name)
    torch.save(checkpoint, checkpoint_path)
    print_with_time_master(f"Saved checkpoint to {checkpoint_path}")


def parse_args():
    parser = argparse.ArgumentParser(description="Train the diffusion model")

    parser.add_argument(
        "--master_addr", type=str, default="localhost", help="Master node address"
    )
    parser.add_argument(
        "--master_port", type=str, default="12355", help="Master node port"
    )
    parser.add_argument(
        "--config_path", type=str, default=None, help="Path to config file to override"
    )
    return parser.parse_args()


if __name__ == "__main__":

    args = parse_args()

    setup_distributed(args.master_addr, args.master_port)
    master_process = int(os.environ["RANK"]) == 0

    # load run config
    with open(args.config_path, "r") as f:
        run_config = json.load(f)

    # Initialize datasets
    print_with_time_master("loading datasets...")
    if run_config["data"]["dataset_type"] == "memmap":
        train_dataset = RewardModelMemmapDataset(
            dataset_dir=run_config["data"]["dataset_dir"],
            metas_filename=run_config["data"]["train_metas_filename"],
            vae_memmap_filename=run_config["data"]["train_vae_memmap_filename"],
            vae_scale_factor=run_config["data"]["vae_scale_factor"],
            mask_prob=run_config["data"]["mask_prob"],
        )
        val_dataset = RewardModelMemmapDataset(
            dataset_dir=run_config["data"]["dataset_dir"],
            metas_filename=run_config["data"]["val_metas_filename"],
            vae_memmap_filename=run_config["data"]["val_vae_memmap_filename"],
            vae_scale_factor=run_config["data"]["vae_scale_factor"],
            mask_prob=0.0,  # no augmentation for validation
        )
    elif run_config["data"]["dataset_type"] == "jsonl":
        train_dataset = RewardModelDataset(
            metas_filepath=os.path.join(
                run_config["data"]["dataset_dir"],
                run_config["data"]["train_metas_filename"],
            ),
            vae_scale_factor=run_config["data"]["vae_scale_factor"],
            mask_prob=run_config["data"]["mask_prob"],
        )
        val_dataset = RewardModelDataset(
            metas_filepath=os.path.join(
                run_config["data"]["dataset_dir"],
                run_config["data"]["val_metas_filename"],
            ),
            vae_scale_factor=run_config["data"]["vae_scale_factor"],
            mask_prob=0.0,  # no augmentation for validation
        )

    # Setup
    ddp_rank = int(os.environ["RANK"])
    ddp_local_rank = int(os.environ["LOCAL_RANK"])
    world_size = dist.get_world_size()
    group_size = min(world_size, 8)
    device = f"cuda:{ddp_local_rank}"
    torch.cuda.set_device(device)
    master_process = ddp_rank == 0

    # Initialize wandb for logging (it will be disabled if in debug mode)
    if master_process:
        wandb.init(
            project=run_config["training"]["wandb_project"],
            name=run_config["training"]["wandb_name"],
            config={
                "run_config": run_config,
                "world_size": world_size,
                "slurm_id": os.environ.get("SLURM_JOB_ID"),
                "slurm_name": os.environ.get("SLURM_JOB_NAME"),
                "slurm_script_path": os.environ.get("SLURM_SCRIPT_PATH"),
                "checkpoint_dir": CHECKPOINT_DIR,
            },
        )
        wandb.run.log_code(".")

    train_sampler = torch.utils.data.DistributedSampler(
        train_dataset,
        shuffle=True,
        rank=ddp_rank,
        num_replicas=world_size,
        drop_last=True,
        seed=run_config["training"]["seed_offset"],
    )
    val_sampler = torch.utils.data.DistributedSampler(
        val_dataset,
        shuffle=False,
        rank=ddp_rank,
        num_replicas=world_size,
        drop_last=True,
        seed=run_config["training"]["seed_offset"],
    )

    # Prepare data loaders
    train_dataloader = torch.utils.data.DataLoader(
        train_dataset,
        batch_size=run_config["training"]["batch_size"],
        sampler=train_sampler,
        drop_last=True,
        num_workers=run_config["training"]["num_workers"],
    )
    val_dataloader = torch.utils.data.DataLoader(
        val_dataset,
        batch_size=run_config["training"]["batch_size"],
        sampler=val_sampler,
        drop_last=True,
        num_workers=0,
    )

    # create the model
    print_with_time_master("setting up model...")

    model = MAEModel(**run_config["model"])
    model.to(device)

    if not run_config["training"]["pretrain"]:
        model = MAEReward(model, pool="mean").to(device)

    num_params = sum(p.numel() for p in model.parameters())
    num_trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    print_with_time_master(f"Number of model parameters: {num_params:,}")
    print_with_time_master(f"Number of trainable parameters: {num_trainable_params:,}")

    if run_config["training"]["compile"]:
        print_with_time_master("compiling model...")
        model_copy_if_compiled = model
        model = torch.compile(model, dynamic=False)

    # create the optimizer
    optimizer = torch.optim.AdamW(model.parameters(), lr=run_config["training"]["lr"])

    # create the lr scheduler
    lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
        optimizer, T_max=run_config["training"]["max_steps"]
    )

    torch.cuda.empty_cache()
    gc.collect()
    dist_barrier()
    print_with_time_master("finished model initialization.")

    torch.manual_seed(ddp_rank)
    random.seed(ddp_rank)
    np.random.seed(ddp_rank)
    global_step = 0
    epoch = 0
    best_val_loss = np.inf

    # train
    print_with_time_master("starting training...")
    log_every = run_config["training"]["log_every"]
    val_every = run_config["training"]["val_every"]

    while global_step < run_config["training"]["max_steps"]:
        train_sampler.set_epoch(epoch)
        for batch in tqdm(
            train_dataloader,
            desc=f"{Fore.CYAN}Epoch {epoch + 1} - Training{Style.RESET_ALL}",
            disable=True,
        ):
            if val_every > 0:
                if global_step % val_every == 0:
                    val_loss, val_acc = validate(
                        model,
                        val_dataloader,
                        pretrain=run_config["training"]["pretrain"],
                    )
                    log_val_metrics(
                        master_process, val_loss, val_acc, global_step, epoch
                    )

            # Save checkpoint every ckpt_every steps
            if (
                run_config["training"]["ckpt_every"] > 0
                and global_step % run_config["training"]["ckpt_every"] == 0
            ):
                save_checkpoint(
                    model,
                    optimizer,
                    lr_scheduler,
                    global_step,
                    CHECKPOINT_DIR,
                    run_config,
                    epoch,
                    ckpt_name="last_ckpt.pt",
                )

            # Save checkpoint if val_acc improves
            if "val_loss" in locals() and val_loss < best_val_loss:
                best_val_loss = val_loss
                save_checkpoint(
                    model,
                    optimizer,
                    lr_scheduler,
                    global_step,
                    CHECKPOINT_DIR,
                    run_config,
                    epoch,
                    ckpt_name="best_ckpt.pt",
                )

            loss, acc = train_step(
                model, batch, pretrain=run_config["training"]["pretrain"]
            )
            optimizer.zero_grad()
            loss.backward()
            grad_norm = torch.nn.utils.clip_grad_norm_(
                model.parameters(), run_config["training"]["grad_clip"]
            )
            optimizer.step()
            lr_scheduler.step()
            global_step += 1

            if global_step % log_every == 0:
                log_metrics(
                    master_process,
                    loss,
                    grad_norm,
                    acc,
                    world_size,
                    global_step,
                    epoch,
                )
        epoch += 1

    # Cleanup
    dist_barrier()
    dist.destroy_process_group()
    wandb.finish()
