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()