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
import torchaudio.functional as AF
import torchaudio.transforms as AT

from tqdm import tqdm
from suno_utils.audio import Audio
from colorama import Fore, Style
from dataclasses import dataclass
from datetime import timedelta
from torch.nn import functional as F
from suno_utils.utils.text import read_jsonl
from typing import List, Tuple, Optional
from torch.distributed import barrier, is_initialized, init_process_group


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.")


# -----------------
# dataset stuff
# -----------------


def _fast_trim_mono(
    x: np.ndarray,  # shape: (samples,), float32/64 in [-1, 1]
    sr: int,  # sample rate (Hz)
    thresh_db_rel: float = -35,  # keep where RMS > max_RMS + thresh (dB)
    win_ms: float = 20.0,  # moving RMS window size (ms)
    pad_ms: float = 20.0,  # pad around kept regions (ms)
    min_keep_ms: float = 40.0,  # drop kept bits shorter than this (ms)
) -> Tuple[np.ndarray, List[Tuple[int, int]]]:
    """
    Ultra-fast silence trimmer for mono audio. No convolutions, all O(n).
    Returns (trimmed_audio, kept_spans) with kept_spans in original sample indices.
    """
    assert x.ndim == 1, "Expected mono waveform of shape (samples,)"
    n = x.size
    if n == 0:
        return x[:0], []

    # --- Moving RMS via cumulative sums (box filter), O(n) ---
    # Compute moving average of power over a window, then sqrt.
    win = max(1, int(round(sr * win_ms / 1000.0)))
    if win > n:
        win = n

    # power and cumulative sum (use float64 for numeric safety)
    sq = x.astype(np.float64) ** 2
    csum = np.empty(n + 1, dtype=np.float64)
    csum[0] = 0.0
    np.cumsum(sq, out=csum[1:])  # csum[k] = sum_{i<k} sq[i]

    # moving average (valid positions)
    # ma_valid[t] = mean of sq[t : t+win]
    ma_valid = (csum[win:] - csum[:-win]) / win  # length n - win + 1

    # Center-align to original length by padding equally on both sides
    left = win // 2
    right = n - (ma_valid.size + left)
    rms = np.sqrt(np.pad(ma_valid, (left, right), mode="edge"))

    # --- Threshold relative to max ---
    eps = 1e-12
    rel_db = 20.0 * np.log10(np.maximum(rms, eps) / (np.max(rms) + eps))
    mask = rel_db > thresh_db_rel  # True = keep

    # --- Turn mask into spans, expand by pad, merge, drop short ---
    pad = max(0, int(round(sr * pad_ms / 1000.0)))
    min_keep = max(1, int(round(sr * min_keep_ms / 1000.0)))

    # Find rising/falling edges
    m = mask.astype(np.int8)
    edges = np.flatnonzero(np.diff(m, prepend=0, append=0))
    # edges come in pairs [start0, end0, start1, end1, ...]
    starts = edges[::2]
    ends = edges[1::2]

    if starts.size == 0:
        return x[:0], []

    # Expand by pad and clamp
    starts = np.maximum(0, starts - pad)
    ends = np.minimum(n, ends + pad)

    # Merge overlaps and drop short spans
    spans: List[Tuple[int, int]] = []
    s_prev = int(starts[0])
    e_prev = int(ends[0])
    for s, e in zip(starts[1:], ends[1:]):
        s = int(s)
        e = int(e)
        if s <= e_prev:  # overlap/adjacent -> merge
            e_prev = max(e_prev, e)
        else:
            if (e_prev - s_prev) >= min_keep:
                spans.append((s_prev, e_prev))
            s_prev, e_prev = s, e
    # last span
    if (e_prev - s_prev) >= min_keep:
        spans.append((s_prev, e_prev))

    if not spans:
        return x[:0], []

    # --- Concatenate kept spans (one pass) ---
    parts = [x[a:b] for (a, b) in spans]
    y = np.concatenate(parts, axis=0).astype(x.dtype)
    return y, spans


def _get_segments(
    audio_np: np.ndarray,
    sample_rate: int,
    num_segments: int = 3,
    segment_duration_sec: float = 3.0,
):
    """
    audio_np: np.ndarray of shape (samples,)
    sample_rate: int
    num_segments: int, number of random segments to crop and return
    segment_duration_sec: float, length (seconds) of each segment to return
    Returns:
        segments: list of np.ndarray of shape (segment_samples,)
    """
    total_samples = len(audio_np)
    segment_samples = int(segment_duration_sec * sample_rate)
    if total_samples < segment_samples:
        # Pad if audio is too short
        padded = np.pad(audio_np, (0, segment_samples - total_samples), mode="constant")
        return [padded.copy() for _ in range(num_segments)]

    segments = []
    for _ in range(num_segments):
        start_idx = np.random.randint(0, total_samples - segment_samples + 1)
        seg = audio_np[start_idx : start_idx + segment_samples]
        segments.append(seg.copy())
    return segments


# collate function
def artist_segment_collate_fn(batch):
    """
    Collate function for batching artist segment samples.

    Each element in `batch` is a tuple:
        (artist_id: str, artist_index: int, segments: List[np.ndarray])

    Returns:
        flat_artist_ids: List[str]                  # len = batch_size * num_segments
        flat_artist_indices: torch.LongTensor       # shape (batch_size * num_segments,)
        flat_segments: torch.FloatTensor            # shape (batch_size * num_segments, segment_samples)
    """
    flat_artist_ids = []
    flat_artist_indices = []
    flat_segments = []

    for artist_id, artist_index, segments in batch:
        # segments: list of np.ndarray (num_segments, segment_samples)
        for seg in segments:
            flat_artist_ids.append(artist_id)
            flat_artist_indices.append(artist_index)
            flat_segments.append(torch.tensor(seg, dtype=torch.float32))

    flat_artist_indices = torch.tensor(flat_artist_indices, dtype=torch.long)
    flat_segments = torch.stack(
        flat_segments, dim=0
    )  # (batch_size * num_segments, segment_samples)

    return flat_artist_ids, flat_artist_indices, flat_segments


class BasicIterableDataset(torch.utils.data.IterableDataset):
    def __init__(self, metas, num_segments: int = 3, segment_duration_s: float = 3.0):
        super(BasicIterableDataset, self).__init__()
        self.metas = metas
        self.num_segments = num_segments
        self.segment_duration_s = segment_duration_s

        # Create artist_id to index mapping for classifier
        self.artist_ids = list(metas.keys())
        self.artist_id_to_index = {
            artist_id: idx for idx, artist_id in enumerate(self.artist_ids)
        }
        self.num_artists = len(self.artist_ids)

        print(f"Dataset initialized with {self.num_artists} artists")

    def __iter__(self):
        while True:
            # Randomly sample an artist
            artist_id = random.choice(self.artist_ids)
            artist_index = self.artist_id_to_index[artist_id]

            # Randomly sample a stem from this artist
            stems_dicts = self.metas[artist_id]
            stem_dict = random.choice(stems_dicts)

            # load this audio file
            audio = Audio.from_file(stem_dict["path"], n_channels=1)

            # trim silence
            audio_trim, _ = _fast_trim_mono(audio.array_float, audio.sample_rate)

            # get segments
            segments = _get_segments(
                audio_trim,
                audio.sample_rate,
                self.num_segments,
                self.segment_duration_s,
            )

            # don't yield the segments separately
            # we will return list with the artist ids and indices
            # and then use a special collate function to merge them
            yield (artist_id, artist_index, segments)


# -----------------
# model stuff
# -----------------


# ----------------------------
# TDNN building blocks
# ----------------------------
class TDNNBlock(nn.Module):
    """
    1D time-dilated conv (x-vector style) with ReLU+BN.
    Input:  (B, C_in, T)
    Output: (B, C_out, T)
    """

    def __init__(self, c_in: int, c_out: int, kernel: int = 5, dilation: int = 1):
        super().__init__()
        pad = dilation * (kernel // 2)
        self.conv = nn.Conv1d(
            c_in, c_out, kernel_size=kernel, dilation=dilation, padding=pad, bias=False
        )
        self.bn = nn.BatchNorm1d(c_out)
        self.act = nn.ReLU(inplace=True)

    def forward(self, x):
        x = self.conv(x)
        x = self.bn(x)
        x = self.act(x)
        return x


class StatsPooling(nn.Module):
    """
    Mean+Std pooling over time.
    Input:  (B, C, T)
    Output: (B, 2*C)
    """

    def forward(self, x):
        # x: (B, C, T)
        mean = x.mean(dim=-1)
        std = x.std(dim=-1, unbiased=False)
        return torch.cat([mean, std], dim=1)


# ----------------------------
# Speaker Encoder -> Embedding
# ----------------------------
@dataclass
class EncoderConfig:
    n_mels: int = 80
    tdnn_channels: tuple = (256, 512, 1024, 2048, 2048)  # 5 TDNN layers
    tdnn_kernels: tuple = (5, 5, 7, 1, 1)
    tdnn_dilations: tuple = (1, 2, 3, 1, 1)
    emb_hidden: int = 256  # penultimate projection before final embedding
    emb_dim: int = 128  # final embedding dimension
    dropout: float = 0.1


class SpeakerEncoder(nn.Module):
    """
    Input:  log-mel features (B, F, T); F = n_mels
    Output: L2-normalized embedding (B, emb_dim)
    """

    def __init__(self, cfg: EncoderConfig):
        super().__init__()
        self.cfg = cfg

        C = [cfg.n_mels] + list(cfg.tdnn_channels)
        self.tdnn = nn.Sequential(
            *[
                TDNNBlock(
                    C[i],
                    C[i + 1],
                    kernel=cfg.tdnn_kernels[i],
                    dilation=cfg.tdnn_dilations[i],
                )
                for i in range(len(cfg.tdnn_channels))
            ]
        )

        self.pool = StatsPooling()  # (B, 2*C_last)
        pooled_dim = 2 * cfg.tdnn_channels[-1]

        self.fc1 = nn.Linear(pooled_dim, cfg.emb_hidden, bias=False)
        self.bn1 = nn.BatchNorm1d(cfg.emb_hidden)
        self.drop = nn.Dropout(p=cfg.dropout)

        self.fc2 = nn.Linear(cfg.emb_hidden, cfg.emb_dim, bias=True)

        # Kaiming for convs; Xavier for linears is fine
        self.apply(self._init_weights)

    @staticmethod
    def _init_weights(m):
        if isinstance(m, nn.Conv1d):
            nn.init.kaiming_normal_(m.weight, nonlinearity="relu")
        elif isinstance(m, nn.Linear):
            nn.init.xavier_uniform_(m.weight)
            if m.bias is not None:
                nn.init.zeros_(m.bias)

    def forward(self, mels: torch.Tensor) -> torch.Tensor:
        """
        mels: (B, F, T)
        returns: L2-normalized embeddings (B, emb_dim)
        """
        x = mels  # (B, F, T)
        x = self.tdnn(x)  # (B, C, T)
        x = self.pool(x)  # (B, 2C)
        x = self.fc1(x)  # (B, H)
        x = self.bn1(x)
        x = F.relu(x, inplace=True)
        x = self.drop(x)
        x = self.fc2(x)  # (B, D)
        # L2 normalize to put on the unit hypersphere
        x = F.normalize(x, p=2, dim=-1)
        return x

    @torch.no_grad()
    def embed(self, mels: torch.Tensor) -> torch.Tensor:
        self.eval()
        return self.forward(mels)


# ----------------------------
# ArcFace / AAM-Softmax Head
# ----------------------------
class ArcMarginProduct(nn.Module):
    """
    Implements AAM-Softmax (ArcFace) logits on-the-fly.
    - Weight matrix W is L2-normalized per row.
    - Inputs are expected already L2-normalized.
    logits = s * cos(theta + m) for the target class, s * cos(theta) otherwise.

    Args:
      in_features:  embedding dim D
      num_classes:  number of speakers C
      s:            scale (30 is common)
      m:            angular margin (0.2~0.3 common)
      easy_margin:  if True, use the easy-margin variant
      ls_eps:       label smoothing epsilon (optional)
    """

    def __init__(
        self,
        in_features: int,
        num_classes: int,
        s: float = 30.0,
        m: float = 0.2,
        easy_margin: bool = False,
        ls_eps: float = 0.0,
    ):
        super().__init__()
        self.in_features = in_features
        self.num_classes = num_classes
        self.s = s
        self.m = m
        self.easy_margin = easy_margin
        self.ls_eps = ls_eps

        self.weight = nn.Parameter(torch.empty(num_classes, in_features))
        nn.init.xavier_uniform_(self.weight)

        # Precompute margin constants
        self.cos_m = math.cos(m)
        self.sin_m = math.sin(m)
        self.th = math.cos(math.pi - m)  # cos(pi - m)
        self.mm = math.sin(math.pi - m) * m

    def forward(self, emb: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
        """
        emb:    (B, D) L2-normalized
        labels: (B,)   int64 in [0, C-1]
        returns: scaled logits for CE loss, shape (B, C)
        """
        # Normalize class weights
        W = F.normalize(self.weight, p=2, dim=1)  # (C, D)

        # Cosine similarity between emb and each class weight
        # cos_theta: (B, C)
        cos_theta = torch.matmul(emb, W.t()).clamp(-1.0, 1.0)

        # Gather target cosine
        idx = torch.arange(emb.size(0), device=emb.device)
        cos_theta_y = cos_theta[idx, labels]  # (B,)

        # Compute cos(theta + m) via trig identity
        sin_theta_y = torch.sqrt(torch.clamp(1.0 - cos_theta_y * cos_theta_y, min=0.0))
        cos_theta_m = cos_theta_y * self.cos_m - sin_theta_y * self.sin_m  # (B,)

        if self.easy_margin:
            # If easy margin: if cos(theta_y) > 0 use margin, else keep original cos
            cond = (cos_theta_y > 0).to(cos_theta.dtype)
            cos_theta_y_m = cond * cos_theta_m + (1 - cond) * cos_theta_y
        else:
            # Classic ArcFace margin decision
            cond = (cos_theta_y > self.th).to(cos_theta.dtype)
            cos_theta_y_m = cond * cos_theta_m + (1 - cond) * (cos_theta_y - self.mm)

        # Replace target logit
        logits = cos_theta.clone()
        logits[idx, labels] = cos_theta_y_m

        # Scale
        logits = logits * self.s

        # Optional label smoothing (applied in CE). We return logits; apply CE outside,
        # but provide a helper to build smoothed targets if needed.
        return logits


# ----------------------------
# Full Model = Encoder + ArcFace head
# ----------------------------
class SpeakerEmbeddingModel(nn.Module):
    """
    Training:
        logits = model(mels, labels) -> (B, C)  # pass to nn.CrossEntropyLoss
    Inference:
        emb = model.embed(mels) -> (B, D) L2-normalized
    """

    def __init__(
        self,
        n_classes: int,
        cfg: Optional[EncoderConfig] = None,
        s: float = 30.0,
        m: float = 0.2,
        easy_margin: bool = False,
        ls_eps: float = 0.0,
    ):
        super().__init__()
        self.cfg = cfg or EncoderConfig()
        self.encoder = SpeakerEncoder(self.cfg)
        self.arcface = ArcMarginProduct(
            in_features=self.cfg.emb_dim,
            num_classes=n_classes,
            s=s,
            m=m,
            easy_margin=easy_margin,
            ls_eps=ls_eps,
        )

    def forward(self, mels: torch.Tensor, labels: Optional[torch.Tensor] = None):
        """
        mels:   (B, F, T) float
        labels: (B,) long, required for training with ArcFace
        """
        emb = self.encoder(mels)  # (B, D), L2-normalized
        if labels is None:
            return emb
        logits = self.arcface(emb, labels)  # (B, C)
        return logits

    @torch.no_grad()
    def embed(self, mels: torch.Tensor) -> torch.Tensor:
        return self.encoder.embed(mels)


def wav_to_logmels(
    wav: torch.Tensor,
    sr: int,
    target_sr: int = 16_000,
    n_mels: int = 80,
    win_ms: float = 25,
    hop_ms: float = 10,
    fmin: float = 20.0,
    fmax: float | None = None,
    top_db: float = 80.0,
    cmvn: bool = True,
) -> torch.Tensor:
    # mix to mono if (B, C, T)
    if wav.dim() == 3:
        wav = wav.mean(1)
    # resample if needed
    if sr != target_sr:
        wav = AF.resample(wav, sr, target_sr)
    n_fft = int(target_sr * win_ms / 1000)
    win_length = n_fft
    hop_length = int(target_sr * hop_ms / 1000)

    # Ensure MelSpectrogram and its buffers are on the same device as input
    device = wav.device

    # Build MelSpectrogram on correct device & move module to input's device
    mel_spect = AT.MelSpectrogram(
        sample_rate=target_sr,
        n_fft=n_fft,
        win_length=win_length,
        hop_length=hop_length,
        f_min=fmin,
        f_max=fmax or target_sr / 2,
        n_mels=n_mels,
        window_fn=lambda window_length: torch.hann_window(window_length, device=device),
        power=2.0,
        mel_scale="slaney",
        norm="slaney",
    )
    mel_spect = mel_spect.to(device)
    mel = mel_spect(wav)  # (B, n_mels, T)

    # AmplitudeToDB -- ensure on same device as well
    db_xfm = AT.AmplitudeToDB(stype="power", top_db=top_db).to(device)
    logmel = db_xfm(mel)

    if cmvn:
        mu = logmel.mean(dim=(-1, -2), keepdim=True)
        sd = logmel.std(dim=(-1, -2), keepdim=True).clamp_min(1e-5)
        logmel = (logmel - mu) / sd

    return logmel


# -----------------
# Training loop
# -----------------


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...")
    train_dataset = BasicIterableDataset(
        run_config["dataset"]["train_metas"],
        num_segments=run_config["dataset"]["num_segments"],
        segment_duration_s=run_config["dataset"]["segment_duration_s"],
    )
    val_dataset = BasicIterableDataset(
        run_config["dataset"]["val_metas"],
        num_segments=run_config["dataset"]["num_segments"],
        segment_duration_s=run_config["dataset"]["segment_duration_s"],
    )

    # 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(".")

    # Prepare data loaders
    train_dataloader = torch.utils.data.DataLoader(
        train_dataset,
        batch_size=run_config["training"]["batch_size"],
        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"],
        drop_last=True,
        num_workers=0,
    )

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

    model = SpeakerEmbeddingModel(
        n_classes=run_config["model"]["n_classes"], cfg=EncoderConfig()
    )
    model.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"]:
        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()
