# Reward Model Training # Train a reward model using Bradley-Terry preference loss # Uses the same paired preference data as DPO training # Negative samples have even indices, positive samples have odd indices from collections import defaultdict from contextlib import nullcontext import datetime import funcy import functools import json import logging import math import os import random import shutil import time import gc from typing import Tuple import numpy as np import torch from torch.nn.parallel import DistributedDataParallel as DDP from torch.distributed.fsdp import ( FullyShardedDataParallel as FSDP, ShardingStrategy, ) from torch.nn import functional as F from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy from torch.distributed import destroy_process_group, init_process_group from data_utils_mmap import read_jsonl, write_jsonl from utils.dpo_data_utils import get_batch from modules.base import ( apply_fsdp_checkpointing, configure_optimizers as base_configure_optimizers, estimate_mfu_no_model, ) from modules.gpt import GPTConfig, GPTTrainConfig, GPT, Block from utils.fsdp_policies import bfSixteen from utils.helpers import ( dist_barrier, load_checkpoint, load_old_state_dict, load_old_optimizer_state_dict, print_with_time, print_with_time_master, save_checkpoint, save_old_checkpoint, suppress_logging, verify_preload_model_args, ) from utils.logging import build_gpu_memory_monitor, Color, NoColor color = Color if True else NoColor # turn down some annoying fsdp logging logging.getLogger("torch.distributed.fsdp._debug_utils").setLevel(logging.ERROR) logging.getLogger("torch.distributed.fsdp._optim_utils").setLevel(logging.ERROR) logging.getLogger("torch.distributed.checkpoint._dedup_tensors").setLevel(logging.ERROR) master_addr = "localhost" master_port = 12355 # Reward model specific parameters freeze_base_model = False # whether to freeze GPT base model layers (only train reward head) use_reward_head = True # enable reward head in model label_smoothing = 0.1 # label smoothing for noisy preference data (~70% accuracy) # Token-level reward training (RAD-inspired) # NOTE: Disabled by default - labels are at sequence level, not token level # Token-level outputs are kept for RAD inference (reward-guided sampling) use_token_level_loss = False # enable per-token Bradley-Terry loss lambda_token = 0.0 # weight for token-level loss (0.0 = disabled) token_level_beta = 1.0 # temperature for token-level Bradley-Terry # Data parameters data_dir = None local_data_shard_dir = None allow_data_shard_reuse = False out_dir = None train_filename = "data_tr.bin" train_metas_filename = "meta_tr.jsonl" train_info_filename = "info_tr.json" val_filename = "data_val.bin" val_metas_filename = "meta_val.jsonl" val_info_filename = "info_val.json" tokenizer_filename = "tokenizer_60k.json" debug_val_only = False dummy_data = False preload_checkpoint = None preload_optimizer = False local_cache_dir = None preload_strict = True suppress_compile_warnings = True grad_checkpointing = False weights_multiplier = None is_finetune = False suppress_text = False checkpoint_save_old_format = True # vocab/time constants text_vocab_size = 60_032 text_codebook_size = 60_001 text_pad_token = text_codebook_size semantic_n_codebooks = 1 semantic_vocab_size = 4032 semantic_codebook_size = 4000 semantic_rate_hz = 25 semantic_shift_factor = 50 coarse_vocab_size = 2112 coarse_codebook_size = 2048 coarse_n_codebooks = 12 data_coarse_n_codebooks = 12 coarse_rate_hz = 25 coarse_shift_factor = 5 t_text = 1152 t_audio = 3136 t_memmap = 6016 t_data_memmap = 6016 block_size = 4288 use_rotary_pos_emb = True rope_theta = 500_000 use_qk_norm = True activation_f = "silu" embed_scale_factor = 1.0 # train params mask_padding = True pack = False layer_init = False infill_augment = False allow_artist = False allow_cover = False use_text_loss = False use_mmbert = False use_vae_input = False output_paradigm = "gpt" output_distribution = "semantic" use_hoot = False use_ditto = False # eval items custom_seed_offset = 0 eval_interval = 2000 log_interval = 25 eval_iters = 250 eval_only = False model_as_bfloat16 = False debug_gradients = False # wandb logging wandb_log = False wandb_project = "suno-reward-model" wandb_run_name = "reward-model-test" wandb_dir = None # data gradient_accumulation_steps = 1 batch_size = 8 eval_loss_batch_size = 8 # model n_layer = 24 n_head = 16 n_kv_head = 4 d_head = 64 dropout = 0.0 bias = False # adamw optimizer - conservative settings for noisy preference data learning_rate = 1e-6 # very low LR for stability with noisy labels min_lr = 1e-7 max_iters = 100_000 warmup_iters = 5_000 # longer warmup for stability lr_decay_iters = None step_save_iters = 10_000 weight_decay = 0.1 # higher weight decay for regularization beta1 = 0.9 beta2 = 0.999 grad_clip = 0.5 # stronger gradient clipping attention_type = "tao" attention_sliding_window_size = 1024 global_every_n_layers = 1 shuffle_data = False local_shuffle_data = False # system device = "cuda" dtype = "bfloat16" compile = False fsdp = False sharding_strategy = "no_shard" # ----------------------------------------------------------------------------- config_keys = [ k for k, v in globals().items() if not k.startswith("_") and isinstance(v, (int, float, bool, str)) ] exec(open("configurator.py").read()) # overrides from command line or config file config = {k: globals()[k] for k in config_keys} # will be useful for logging # ----------------------------------------------------------------------------- # auto set a few params batch_size_tokens = block_size * batch_size eval_loss_batch_size_tokens = block_size * eval_loss_batch_size if pack: batch_size = 1 assert t_text + t_audio <= block_size assert dtype in ("bfloat16", "float32") if debug_val_only or eval_only: train_filename = val_filename train_metas_filename = val_metas_filename train_info_filename = val_info_filename if not eval_only: wandb_log = False eval_iters = int(eval_iters * gradient_accumulation_steps) if lr_decay_iters is None: lr_decay_iters = max_iters # set up distributed variables os.environ["MASTER_ADDR"] = str(master_addr) os.environ["MASTER_PORT"] = str(master_port) if "SLURM_PROCID" in os.environ: if int(os.environ["SLURM_NTASKS_PER_NODE"]) != torch.cuda.device_count(): raise ValueError( f"SLURM_NTASKS_PER_NODE ({os.environ['SLURM_NTASKS_PER_NODE']}) does not match" f" the number of CUDA devices ({torch.cuda.device_count()}) on node {os.environ['HOSTNAME']}" ) ddp_rank = int(os.environ["SLURM_PROCID"]) ddp_local_rank = int(os.environ["SLURM_LOCALID"]) world_size = int(os.environ["SLURM_JOB_NUM_NODES"]) * int(os.environ["SLURM_NTASKS_PER_NODE"]) else: ddp_rank = 0 ddp_local_rank = 0 world_size = 1 os.environ["RANK"] = str(ddp_rank) os.environ["LOCAL_RANK"] = str(ddp_local_rank) ddp = int(os.environ.get("RANK", -1)) != -1 if fsdp: assert ddp, "found fsdp = True but ddp is False" if ddp: try: init_process_group( backend="nccl", timeout=datetime.timedelta(seconds=24 * 60 * 60), rank=ddp_rank, world_size=world_size, device_id=torch.device(f"cuda:{ddp_local_rank}"), ) except Exception as e: print(f"Distributed error on rank {ddp_rank} with host {os.environ['HOSTNAME']}") raise e device = f"cuda:{ddp_local_rank}" torch.cuda.set_device(device) master_process = ddp_rank == 0 seed_offset = ddp_rank + 1 print_with_time(f"ddp init, rank {ddp_rank}, local_rank {ddp_local_rank}") else: master_process = True seed_offset = 1 n_gpus_per_node = torch.cuda.device_count() dist_barrier() print_with_time_master(f"ddp init: world size {world_size} ddp_rank {ddp_rank}.") gc.disable() seed_offset *= custom_seed_offset + 1 torch.manual_seed(6006 + seed_offset) random.seed(6006 + seed_offset) np.random.seed(6006 + seed_offset) torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True device_type = "cuda" if "cuda" in device else "cpu" ptdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16}[dtype] ctx = ( nullcontext() if device_type == "cpu" or fsdp 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, dir=wandb_dir) wandb.run.log_code(".") print_with_time_master(f"Total world size {world_size}") date_time_str = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S") out_dir = os.path.join(out_dir, date_time_str) if master_process and not debug_val_only: os.makedirs(out_dir, exist_ok=True) print_with_time_master(f"logging checkpoint here: {out_dir}") # shard data if necessary (same as DPO) dist_barrier() if allow_data_shard_reuse: assert local_data_shard_dir is not None for fn in [ tokenizer_filename, val_filename, val_info_filename, val_metas_filename, train_filename, train_info_filename, train_metas_filename, ]: assert os.path.isfile(os.path.join(local_data_shard_dir, fn)), os.path.join( local_data_shard_dir, fn ) if local_data_shard_dir is not None and not allow_data_shard_reuse and ddp_local_rank == 0: print_with_time_master("sharding data...") shutil.rmtree(local_data_shard_dir, ignore_errors=True) os.makedirs(local_data_shard_dir) # copy over tokenizer and val for fn in [tokenizer_filename, val_filename, val_metas_filename, val_info_filename]: shutil.copyfile( os.path.join(data_dir, fn), os.path.join(local_data_shard_dir, fn), ) # load from data_dir and shard based on fraction that node should receive from_frac = ddp_rank / world_size to_frac = (ddp_rank + n_gpus_per_node) / world_size assert 0 <= from_frac <= 1 assert 0 <= to_frac <= 1 with open(os.path.join(data_dir, train_info_filename)) as f: train_info = json.load(f) for dset_name in train_info.keys(): if "idx_map" in train_info[dset_name]: train_info[dset_name]["idx_map"] = { int(k): v for k, v in train_info[dset_name]["idx_map"].items() } new_train_info = {} new_idx_offset = 0 orig_idx_seq = [] for dset_name, info in train_info.items(): if info.get("task", "default") == "default": assert "idx_list" in info idx_list = info["idx_list"][:] from_n_sample = int(round(from_frac * len(idx_list))) to_n_sample = int(round(to_frac * len(idx_list))) keep_idx_list = idx_list[from_n_sample:to_n_sample] new_train_info[dset_name] = { "idx_list": list(range(new_idx_offset, new_idx_offset + len(keep_idx_list))), "task": "default", } orig_idx_seq.extend(keep_idx_list) new_idx_offset += len(keep_idx_list) elif info["task"] == "covers": idx_map_list = [(k, v) for k, v in info["idx_map"].items()] from_n_sample = int(round(from_frac * len(idx_map_list))) to_n_sample = int(round(to_frac * len(idx_map_list))) keep_idx_list = [] new_idx_map = defaultdict(list) for k, v in idx_map_list[from_n_sample:to_n_sample]: keep_idx_list.append(k) keep_idx_list.extend(v) new_idx_map[new_idx_offset] = list( range(new_idx_offset + 1, new_idx_offset + 1 + len(v)) ) new_idx_offset += 1 + len(v) new_train_info[dset_name] = {"idx_map": dict(new_idx_map), "task": "covers"} orig_idx_seq.extend(keep_idx_list) else: raise ValueError(f"unknown task for {dset_name} in info file") print_with_time_master(f"shard size: {len(orig_idx_seq):,}") with open(os.path.join(local_data_shard_dir, train_info_filename), "w") as f: json.dump(new_train_info, f) print_with_time_master("done with info shard") train_data = np.memmap(os.path.join(data_dir, train_filename), dtype=np.uint16, mode="r") train_data = train_data.reshape(-1, t_data_memmap, semantic_n_codebooks + data_coarse_n_codebooks) train_metas = read_jsonl( os.path.join(data_dir, train_metas_filename), parse_idx_set=set(orig_idx_seq) ) assert len(train_data) == len(train_metas) new_train_metas = [train_metas[idx] for idx in orig_idx_seq] assert not any(m is None for m in new_train_metas) write_jsonl(new_train_metas, os.path.join(local_data_shard_dir, train_metas_filename)) print_with_time_master("done with metas shard") new_train_data = np.memmap( os.path.join(local_data_shard_dir, train_filename), dtype=np.uint16, mode="w+", shape=(1, t_data_memmap, semantic_n_codebooks + data_coarse_n_codebooks), ) new_idx = 0 orig_idx_seq_chunks = list(funcy.chunks(100_000, orig_idx_seq)) for n_chunk, orig_idx_seq_chunk in enumerate(orig_idx_seq_chunks): new_train_data = np.memmap( os.path.join(local_data_shard_dir, train_filename), dtype=np.uint16, mode="r+", shape=( new_idx + len(orig_idx_seq_chunk), t_data_memmap, semantic_n_codebooks + data_coarse_n_codebooks, ), ) for n1, n2 in sorted( list( zip( range(new_idx, new_idx + len(orig_idx_seq_chunk)), orig_idx_seq_chunk, ) ), key=lambda x: x[-1], ): new_train_data[n1] = train_data[n2] new_idx += len(orig_idx_seq_chunk) new_train_data.flush() print_with_time_master(f"processed memmap chunk {n_chunk + 1}/{len(orig_idx_seq_chunks)}") del new_train_data, new_train_metas, new_train_info del train_metas, train_data, train_info gc.collect() print_with_time_master("done sharding.") if local_data_shard_dir is not None: data_dir = local_data_shard_dir dist_barrier() # load data print_with_time_master("loading data...") if weights_multiplier is None: weights_multiplier_map = {} elif weights_multiplier == "base": weights_multiplier_map = { "youtube_music_lyrics": 2, "youtube_music_lyrics_foreign": 3, "genius_hq_lyrics": 2, "genius_hq_lyrics_foreign": 3, "deezer_lyrics": 2, "deezer_lyrics_foreign": 3, } elif weights_multiplier == "finetune": weights_multiplier_map = { "youtube_music_lyrics": 2, "youtube_music_lyrics_foreign": 3, "genius_hq_lyrics": 6, "genius_hq_lyrics_foreign": 8, "deezer_lyrics": 2, "deezer_lyrics_foreign": 3, } else: weights_multiplier_map = { k.split(":")[0].strip(): float(k.split(":")[1].strip()) for k in weights_multiplier.strip(";").split(";") if ":" in k } def load_dataset( data_dir: str, filename: str, info_filename: str, metas_filename: str, weights_multiplier_map: dict, is_finetune: bool, ) -> Tuple: """Load dataset for reward model training.""" dataset_names = [] data_idx_lists = [] data_weights = [] data = np.memmap(os.path.join(data_dir, filename), dtype=np.uint16, mode="r") data = data.reshape(-1, t_data_memmap, semantic_n_codebooks + data_coarse_n_codebooks) assert data[:100, :, :semantic_n_codebooks].max() <= semantic_vocab_size if data_coarse_n_codebooks > 0: assert data[:100, :, semantic_n_codebooks:].max() <= coarse_vocab_size with open(os.path.join(data_dir, info_filename)) as f: infos = json.load(f) metas = read_jsonl(os.path.join(data_dir, metas_filename)) assert len(data) == len(metas), (len(data), len(metas)) artist_to_songs = defaultdict(list) for i, m in enumerate(metas): if "artist" in m: artist_to_songs[f"{m['dataset']}__{m['artist']}"].append(i) artist_to_songs = {k: v for k, v in artist_to_songs.items() if len(v) > 1} if allow_artist: assert len(artist_to_songs) > 0, "no artist data found" print_with_time_master(f"found {len(artist_to_songs):,} samples with artists on main process") idx_set = set() has_cover = False for dset_name in sorted(infos.keys()): if "idx_map" in infos[dset_name]: infos[dset_name]["idx_map"] = {int(k): v for k, v in infos[dset_name]["idx_map"].items()} for dset_name in sorted(infos.keys()): info = infos[dset_name] dataset_names.append(dset_name) if info.get("task", "default") == "default": assert "idx_list" in info idx_list = info["idx_list"][:] idx_set |= set(idx_list) idx_list = sorted([idx for idx in info["idx_list"]]) elif info["task"] == "covers": assert pack, "for now pack needs to be active to do covers" assert batch_size_tokens >= t_memmap * 2 + t_text, "for covers we need double the blocksize" has_cover = True idx_list = [] n_covers = 0 for idx, child_idx_l in info["idx_map"].items(): idx_list.append(int(idx)) idx_set.add(int(idx)) idx_set |= set(child_idx_l) n_covers += len(child_idx_l) print_with_time_master( f"found {len(idx_list):,} samples with {n_covers:,} total covers on main process" ) else: raise ValueError(f"unknown task for {dset_name} in info file") data_idx_lists.append(idx_list) data_weights.append(len(idx_list) * weights_multiplier_map.get(dset_name, 1.0)) if allow_cover: assert has_cover, "no cover data found" weights_norm = np.sum(data_weights) data_weights = [v / weights_norm for v in data_weights] if not is_finetune: print_with_time_master(f"indexed {len(idx_set) / len(data) * 100:.1f}% of data") for k in weights_multiplier_map.keys(): assert k in dataset_names del idx_set shard_info = "" if local_data_shard_dir is None else " (sharded)" print_with_time_master(f"{len(data):,} lines of {filename} loaded.{shard_info}") assert len(data) == len(metas) assert len(infos) == len(dataset_names) == len(data_weights) == len(data_idx_lists) return ( dataset_names, data_idx_lists, data_weights, data, metas, infos, artist_to_songs, ) ( val_dataset_names, val_data_idx_lists, val_data_weights, val_data, val_metas, val_info, val_artist_to_songs, ) = load_dataset( data_dir, val_filename, val_info_filename, val_metas_filename, weights_multiplier_map, is_finetune, ) ( train_dataset_names, train_data_idx_lists, train_data_weights, train_data, train_metas, train_info, train_artist_to_songs, ) = load_dataset( data_dir, train_filename, train_info_filename, train_metas_filename, weights_multiplier_map, is_finetune, ) if master_process: weights_str = "train data weights:" for k, v in zip(train_dataset_names, train_data_weights): weights_str += f"\n {round(v * 100, 1)}% {k}" print_with_time_master(weights_str) print_with_time_master("done loading data") dist_barrier() def compute_loss_end_indices(Y: torch.Tensor, seq_len: int, semantic_pad_token: int = 4000) -> list[int]: """Compute where valid data ends for each sample (before padding). Y[j]=-1 OR Y[j]=pad_token means X[j+1] is padding, so end_index = last_valid_j + 2 Args: Y: Targets (batch, n_codebooks, seq_len-1) with -1 or pad_token for padding seq_len: Sequence length of X/rewards semantic_pad_token: Padding token value for semantic codebook (default 4000) Returns: loss_end_index_list: End index for each sample """ batch_size = Y.shape[0] loss_end_index_list = [] for i in range(batch_size): y_sample = Y[i, 0, :] # First codebook (semantic) # Check for BOTH -1 (marked by mask_middle_padding) AND actual pad token # Non-padding means: not -1 AND not semantic_pad_token non_pad_mask = (y_sample != -1) & (y_sample != semantic_pad_token) if non_pad_mask.any(): last_valid = torch.where(non_pad_mask)[0][-1].item() end_idx = last_valid + 2 # Y-shift: Y[j]=pad → X[j+1] padding else: end_idx = seq_len loss_end_index_list.append(end_idx) return loss_end_index_list def extract_scalar_rewards( reward_logits: torch.Tensor, loss_start_index_list: list[int], loss_end_index_list: list[int] = None, ) -> torch.Tensor: """Extract scalar rewards by averaging between start and end indices. Args: reward_logits: Token-level rewards (batch, seq_len) loss_start_index_list: Where generation starts for each sample loss_end_index_list: Where valid data ends (None = seq_len, i.e., no padding) Returns: scalar_rewards: (batch,) """ batch_size, seq_len = reward_logits.shape device = reward_logits.device # Position indices positions = torch.arange(seq_len, device=device).unsqueeze(0).expand(batch_size, -1) # Start indices start_indices = ( torch.tensor(loss_start_index_list, device=device).unsqueeze(1) if loss_start_index_list else torch.zeros(batch_size, 1, dtype=torch.long, device=device) ) # End indices (default to seq_len if not provided) end_indices = ( torch.tensor(loss_end_index_list, device=device).unsqueeze(1) if loss_end_index_list else torch.full((batch_size, 1), seq_len, dtype=torch.long, device=device) ) # Simple mask: start <= position < end mask = (positions >= start_indices) & (positions < end_indices) # Average masked_rewards = reward_logits * mask valid_counts = mask.sum(dim=1, keepdim=True).clamp(min=1) scalar_rewards = masked_rewards.sum(dim=1) / valid_counts.squeeze(1) return scalar_rewards def token_level_reward_loss( reward_logits_chosen: torch.Tensor, reward_logits_rejected: torch.Tensor, loss_start_index_list_chosen: list, loss_start_index_list_rejected: list, loss_end_index_list_chosen: list, loss_end_index_list_rejected: list, beta: float = 1.0, ) -> Tuple[torch.Tensor, torch.Tensor]: """Token-level Bradley-Terry preference loss (RAD-inspired). Trains the reward model to score chosen tokens higher than rejected tokens at EACH position in the sequence. This provides much richer training signal than sequence-level loss alone (10-100x more gradients per batch). For each valid token position t: loss[t] = -log(sigmoid(beta * (reward_chosen[t] - reward_rejected[t]))) Final loss is the average over all valid positions. Args: reward_logits_chosen: Token-level rewards for chosen samples, shape (batch, seq_len) reward_logits_rejected: Token-level rewards for rejected samples, shape (batch, seq_len) loss_start_index_list_chosen: Generation start positions for chosen loss_start_index_list_rejected: Generation start positions for rejected loss_end_index_list_chosen: Generation end positions for chosen loss_end_index_list_rejected: Generation end positions for rejected beta: Temperature parameter for Bradley-Terry (default 1.0) Returns: loss: Token-level preference loss (scalar) accuracy: Fraction of tokens where chosen > rejected (scalar) """ batch_size, seq_len = reward_logits_chosen.shape device = reward_logits_chosen.device # Create position indices positions = torch.arange(seq_len, device=device).unsqueeze(0).expand(batch_size, -1) # Create masks inline (no helper needed) start_c = torch.tensor(loss_start_index_list_chosen, device=device).unsqueeze(1) end_c = torch.tensor(loss_end_index_list_chosen, device=device).unsqueeze(1) mask_chosen = (positions >= start_c) & (positions < end_c) start_r = torch.tensor(loss_start_index_list_rejected, device=device).unsqueeze(1) end_r = torch.tensor(loss_end_index_list_rejected, device=device).unsqueeze(1) mask_rejected = (positions >= start_r) & (positions < end_r) # Only compute loss where BOTH chosen and rejected are valid (aligned comparison) # NOTE: For sequences of different lengths, this excludes non-overlapping portions. # This is correct for pairwise Bradley-Terry - we can only compare positions where # both sequences exist. The sequence-level loss (always enabled) already captures # overall preference using each sequence's full independent length. mask = mask_chosen & mask_rejected # (batch, seq_len) # Token-level Bradley-Terry loss # For each token: P(chosen[t] > rejected[t]) = sigmoid(beta * (r_chosen[t] - r_rejected[t])) token_diff = reward_logits_chosen - reward_logits_rejected # (batch, seq_len) token_losses = -F.logsigmoid(beta * token_diff) # (batch, seq_len) # Apply mask and average over all valid token positions masked_losses = token_losses * mask loss = masked_losses.sum() / mask.sum().clamp(min=1) # Token-level accuracy: How many positions have chosen > rejected? token_correct = (reward_logits_chosen > reward_logits_rejected) & mask accuracy = token_correct.sum().float() / mask.sum().float().clamp(min=1) return loss, accuracy # Bradley-Terry preference loss for reward modeling def reward_model_loss( reward_chosen: torch.Tensor, reward_rejected: torch.Tensor, label_smoothing: float = 0.0, ) -> Tuple[torch.Tensor, torch.Tensor]: """Compute Bradley-Terry preference loss for reward model with label smoothing. Bradley-Terry model: P(chosen > rejected) = sigmoid(reward_chosen - reward_rejected) Loss with label smoothing: Interpolate between perfect labels and uniform distribution -log(sigmoid(r_c - r_r)) * (1-eps) - log(sigmoid(r_r - r_c)) * eps Label smoothing helps with noisy preference data (~70% accuracy) by preventing the model from being overconfident on potentially mislabeled pairs. Args: reward_chosen: Scalar rewards for chosen/positive samples, shape (batch_size,) reward_rejected: Scalar rewards for rejected/negative samples, shape (batch_size,) label_smoothing: Smoothing factor (0.0 = no smoothing, typical: 0.1 for noisy data) Returns: loss: Bradley-Terry loss with label smoothing (scalar) accuracy: Reward ranking accuracy (scalar) """ # Bradley-Terry loss with label smoothing # Interpolate between correct preference and reversed preference losses = ( -F.logsigmoid(reward_chosen - reward_rejected) * (1 - label_smoothing) - F.logsigmoid(reward_rejected - reward_chosen) * label_smoothing ) loss = losses.mean() # Accuracy: how often is chosen reward > rejected reward? accuracy = (reward_chosen > reward_rejected).float().mean() return loss, accuracy # model init model_args = dict( n_layer=n_layer, n_head=n_head, n_kv_head=n_kv_head, d_head=d_head, block_size=block_size, bias=bias, text_vocab_size=text_vocab_size, text_codebook_size=text_codebook_size, text_pad_token=text_pad_token, semantic_vocab_size=semantic_vocab_size, semantic_codebook_size=semantic_codebook_size, semantic_n_codebooks=semantic_n_codebooks, semantic_rate_hz=semantic_rate_hz, semantic_shift_factor=semantic_shift_factor, coarse_vocab_size=coarse_vocab_size, coarse_codebook_size=coarse_codebook_size, coarse_n_codebooks=coarse_n_codebooks, coarse_rate_hz=coarse_rate_hz, coarse_shift_factor=coarse_shift_factor, t_text=t_text, t_audio=t_audio, use_rotary_pos_emb=use_rotary_pos_emb, rope_theta=rope_theta, use_qk_norm=use_qk_norm, activation_f=activation_f, embed_scale_factor=embed_scale_factor, attention_sliding_window_size=attention_sliding_window_size, global_every_n_layers=global_every_n_layers, use_text_loss=use_text_loss, use_mmbert=use_mmbert, use_vae_input=use_vae_input, output_paradigm=output_paradigm, output_distribution=output_distribution, use_hoot=use_hoot, use_ditto=use_ditto, use_reward_head=use_reward_head, # Enable reward head ) train_model_args = dict(dropout=dropout, attention_type=attention_type, layer_init=layer_init) if preload_checkpoint is not None and not preload_checkpoint.endswith(".pt"): verify_preload_model_args(model_args, preload_checkpoint, preload_strict=preload_strict) elif preload_checkpoint is not None: ckpt = torch.load(preload_checkpoint, mmap=True, weights_only=False) if "model_args" in ckpt: loaded_model_args = ckpt["model_args"] print_with_time_master(f"loaded model args from checkpoint: {loaded_model_args}") # Filter out any unknown parameters (e.g., use_mt5 from old checkpoints) # Get valid GPTConfig parameter names import inspect valid_params = set(inspect.signature(GPTConfig.__init__).parameters.keys()) - {"self"} # Only keep valid parameters filtered_model_args = {k: v for k, v in loaded_model_args.items() if k in valid_params} unknown_params = set(loaded_model_args.keys()) - valid_params if unknown_params: print_with_time_master(f"Filtering out unknown parameters from checkpoint: {unknown_params}") model_args = filtered_model_args # Override with reward model specific args model_args["block_size"] = block_size model_args["t_text"] = t_text model_args["t_audio"] = t_audio model_args["use_reward_head"] = use_reward_head print_with_time_master( f"Overriding model args for reward model: use_reward_head={use_reward_head}" ) gpu_memory_monitor = build_gpu_memory_monitor() if preload_checkpoint is None: raise ValueError("Reward model training requires a pre-trained checkpoint") # init model print_with_time_master( f"Initializing reward model from checkpoint with use_reward_head={model_args.get('use_reward_head', False)}" ) gptconf = GPTConfig(**model_args) print_with_time_master(f"GPTConfig created with use_reward_head={gptconf.use_reward_head}") gpttrainconf = GPTTrainConfig(**train_model_args) model = GPT(gptconf, gpttrainconf) # Verify reward head was created if not model.config.use_reward_head: raise RuntimeError( f"Model was created with use_reward_head={model.config.use_reward_head}, expected True" ) if "reward_head" not in model.output_modules: raise RuntimeError( f"Model does not have reward_head. Output modules: {list(model.output_modules.keys())}" ) print_with_time_master(f"✓ Reward head verified in model: {list(model.output_modules.keys())}") if model_as_bfloat16: model.to(torch.bfloat16) if not fsdp: model.to(device) cfg = model.config train_cfg = model.train_config print_with_time_master("finish init reward model") # calculate params raw_model_n_params = model.get_num_params() all_param = sum(p.numel() for p in model.parameters()) trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) print_with_time_master( f"trainable params: {trainable_params:,d} || " f"all params: {all_param:,d} || " f"trainable%: {100 * trainable_params / all_param:.4f}" ) # Freeze base model if requested (only train reward head) if freeze_base_model: print_with_time_master("Freezing base model, only training reward head") for name, param in model.named_parameters(): if "reward_head" not in name: param.requires_grad = False trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) print_with_time_master( f"After freezing - trainable params: {trainable_params:,d} || " f"trainable%: {100 * trainable_params / all_param:.4f}" ) # compile the model if compile: print_with_time_master("compiling the model... (takes a ~minute)") compile_ctx = suppress_logging if suppress_compile_warnings else nullcontext with compile_ctx(): model = torch.compile(model) dist_barrier() else: print_with_time_master("not compiling model.") iter_num = 0 total_tokens_processed = 0 rel_tokens_processed = 0 best_val_loss = 1e9 # load old single-file checkpoint if preload_checkpoint is not None and preload_checkpoint.endswith(".pt"): print_with_time_master("start loading state dict") load_old_state_dict( model_args, preload_checkpoint, model, local_cache_dir, preload_strict=False, # Allow missing reward head use_mmap=True, ) print_with_time_master("finish loading state dict") dist_barrier() # order matters for FSDP/DDP if fsdp: print_with_time_master("wrapping model in FSDP ....") auto_wrap_policy = functools.partial( transformer_auto_wrap_policy, transformer_layer_cls={Block}, ) model = FSDP( model, auto_wrap_policy=auto_wrap_policy, mixed_precision=bfSixteen, sharding_strategy=getattr(ShardingStrategy, sharding_strategy.upper()), device_id=torch.cuda.current_device(), sync_module_states=True, use_orig_params=True, ) gpu_mem_stats = gpu_memory_monitor.get_peak_stats() print_with_time_master( f"GPU memory usage for model: " f"{gpu_mem_stats.max_reserved_gib:.2f}GiB" f"({gpu_mem_stats.max_reserved_pct:.2f}%)" ) if grad_checkpointing: apply_fsdp_checkpointing(model) optimizer = base_configure_optimizers( model, weight_decay, learning_rate, (beta1, beta2), device_type, use_fused=False, is_fsdp=True, ) else: optimizer = model.configure_optimizers(weight_decay, learning_rate, (beta1, beta2), device_type) if ddp: print_with_time_master("wrapping model in DDP") model = DDP(model, device_ids=[ddp_local_rank]) torch.cuda.empty_cache() dist_barrier() # load optimizer state if requested if preload_checkpoint is not None and preload_checkpoint.endswith(".pt") and preload_optimizer: iter_num, total_tokens_processed, best_val_loss = load_old_optimizer_state_dict( model, optimizer, local_cache_dir, ) # load new distributed checkpoint if preload_checkpoint is not None and not preload_checkpoint.endswith(".pt"): iter_num, total_tokens_processed, best_val_loss = load_checkpoint( preload_checkpoint, preload_optimizer, model, optimizer, ) dist_barrier() print_with_time_master("model setup done") # learning rate decay scheduler (cosine with warmup) def get_lr(it: int) -> float: """Get learning rate for iteration it.""" if it < warmup_iters: return min_lr + (learning_rate - min_lr) * it / warmup_iters if it > lr_decay_iters: return min_lr 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)) return min_lr + coeff * (learning_rate - min_lr) data_sampling_info = { "cfg": cfg, "train_cfg": train_cfg, "batch_size": batch_size, "batch_size_tokens": batch_size_tokens, "tokenizer_fp": os.path.join(data_dir, tokenizer_filename), "device": device, "device_type": device_type, "train": { "data": train_data, "metas": train_metas, "infos": train_info, "artist_to_songs": train_artist_to_songs, "names": train_dataset_names, "weights": train_data_weights, "idx_lists": train_data_idx_lists, "all_idx_lists": sorted([idx for sublist in train_data_idx_lists for idx in sublist]), }, "val": { "data": val_data, "metas": val_metas, "infos": val_info, "artist_to_songs": val_artist_to_songs, "names": val_dataset_names, "weights": val_data_weights, "idx_lists": val_data_idx_lists, "all_idx_lists": sorted([idx for sublist in val_data_idx_lists for idx in sublist]), }, } eval_loss_data_sampling_info = data_sampling_info.copy() eval_loss_data_sampling_info["batch_size"] = eval_loss_batch_size eval_loss_data_sampling_info["batch_size_tokens"] = eval_loss_batch_size_tokens @torch.no_grad() def estimate_loss(): """Evaluate reward model on train and val sets.""" model.eval() # Verify model has reward head enabled (unwrap DDP/FSDP if needed) raw_model = model.module if hasattr(model, "module") else model if not raw_model.config.use_reward_head: raise RuntimeError(f"Model use_reward_head is {raw_model.config.use_reward_head}, expected True") if "reward_head" not in raw_model.output_modules: raise RuntimeError("Model does not have reward_head in output_modules") if master_process: print_with_time_master(f"Eval: use_reward_head={raw_model.config.use_reward_head}") out = {} for split in ["train", "val"]: losses = [] accuracies = [] chosen_rewards = [] rejected_rewards = [] margins = [] losses_seq = [] losses_tok = [] accs_tok = [] for _ in range(eval_iters): # Load paired preference data idxs, X, Y, loss_start_index_list = get_batch( data_sampling_info, split, inference=True, dummy_data=dummy_data, suppress_text=suppress_text, load_dpo_pair=True, return_idx=True, ) with ctx: # Get output from model (dict with rewards and generation logits) output = model(X, return_logits=True) reward_logits = output["reward_logits"] # (batch, seq_len) # Compute end indices from Y loss_end_index_list = compute_loss_end_indices(Y, reward_logits.shape[1]) # Sequence-level loss rewards = extract_scalar_rewards( reward_logits, loss_start_index_list, loss_end_index_list ) # (batch,) reward_rejected_seq = rewards[::2] reward_chosen_seq = rewards[1::2] loss_seq, acc_seq = reward_model_loss( reward_chosen_seq, reward_rejected_seq, label_smoothing=label_smoothing ) # Token-level loss if use_token_level_loss: reward_logits_chosen = reward_logits[1::2] reward_logits_rejected = reward_logits[::2] loss_tok, acc_tok = token_level_reward_loss( reward_logits_chosen, reward_logits_rejected, loss_start_index_list[1::2], # chosen start loss_start_index_list[::2], # rejected start loss_end_index_list[1::2], # chosen end loss_end_index_list[::2], # rejected end beta=token_level_beta, ) loss = loss_seq + lambda_token * loss_tok losses_seq.append(loss_seq.item()) losses_tok.append(loss_tok.item()) accs_tok.append(acc_tok.item()) else: loss = loss_seq losses.append(loss.item()) accuracies.append(acc_seq.item()) chosen_rewards.append(reward_chosen_seq.mean().item()) rejected_rewards.append(reward_rejected_seq.mean().item()) margins.append((reward_chosen_seq - reward_rejected_seq).mean().item()) out[f"{split}/loss"] = float(np.mean(losses)) out[f"{split}/accuracy"] = float(np.mean(accuracies)) out[f"{split}/chosen_reward"] = float(np.mean(chosen_rewards)) out[f"{split}/rejected_reward"] = float(np.mean(rejected_rewards)) out[f"{split}/margin"] = float(np.mean(margins)) if use_token_level_loss: out[f"{split}/loss_seq"] = float(np.mean(losses_seq)) out[f"{split}/loss_tok"] = float(np.mean(losses_tok)) out[f"{split}/acc_tok"] = float(np.mean(accs_tok)) model.train() return out @torch.no_grad() def debug_first_batches(): """Print detailed info about first batches from each dataset to debug initial metrics.""" model.eval() print_with_time_master("\n" + "=" * 80) print_with_time_master("DEBUG: First batch analysis") print_with_time_master("=" * 80) # Debug val set print_with_time_master("\n--- VALIDATION SET (first batch) ---") idxs, X, Y, loss_start_index_list = get_batch( data_sampling_info, "val", inference=True, dummy_data=dummy_data, suppress_text=suppress_text, load_dpo_pair=True, return_idx=True, ) with ctx: output = model(X, return_logits=True) reward_logits = output["reward_logits"] loss_end_index_list = compute_loss_end_indices(Y, reward_logits.shape[1]) rewards = extract_scalar_rewards(reward_logits, loss_start_index_list, loss_end_index_list) # Rewards are interleaved: [rejected_0, chosen_0, rejected_1, chosen_1, ...] reward_rejected = rewards[::2] reward_chosen = rewards[1::2] print_with_time_master(f"Batch size: {len(idxs)}, Pairs: {len(reward_chosen)}") print_with_time_master(f"Y shape: {Y.shape}") print_with_time_master(f"loss_start_indices: {loss_start_index_list[:4]}...") print_with_time_master(f"loss_end_indices: {loss_end_index_list[:4]}...") # Check if Y actually has padding y_first_codebook = Y[:, 0, :] # (batch, seq_len-1) has_padding = (y_first_codebook == -1).any(dim=1) num_with_padding = has_padding.sum().item() print_with_time_master(f"Samples with padding (-1 in Y): {num_with_padding}/{Y.shape[0]}") # For each pair, show actual sequence lengths AND where padding starts print_with_time_master(f"\nPair-wise sequence lengths (DEBUGGING PADDING POSITIONS):") for i in range(min(3, len(reward_chosen))): rej_idx = i * 2 cho_idx = i * 2 + 1 rej_end = loss_end_index_list[rej_idx] cho_end = loss_end_index_list[cho_idx] rej_len = rej_end - loss_start_index_list[rej_idx] cho_len = cho_end - loss_start_index_list[cho_idx] same_len = "✓ SAME" if rej_len == cho_len else "✗ DIFFERENT" # Find where -1 FIRST appears in Y for each sequence y_rej = Y[rej_idx, 0, :] y_cho = Y[cho_idx, 0, :] rej_pad_locs = torch.where(y_rej == -1)[0] cho_pad_locs = torch.where(y_cho == -1)[0] rej_first_pad = rej_pad_locs[0].item() if len(rej_pad_locs) > 0 else "NONE" cho_first_pad = cho_pad_locs[0].item() if len(cho_pad_locs) > 0 else "NONE" rej_last_pad = rej_pad_locs[-1].item() if len(rej_pad_locs) > 0 else "NONE" cho_last_pad = cho_pad_locs[-1].item() if len(cho_pad_locs) > 0 else "NONE" print_with_time_master( f" Pair {i}: REJ end={rej_end} pad[{rej_first_pad}:{rej_last_pad}], " f"CHO end={cho_end} pad[{cho_first_pad}:{cho_last_pad}] [{same_len}]" ) print_with_time_master(f"\nFirst 5 pairs:") for i in range(min(5, len(reward_chosen))): margin = (reward_chosen[i] - reward_rejected[i]).item() correct = reward_chosen[i] > reward_rejected[i] print_with_time_master( f" Pair {i}: chosen={reward_chosen[i].item():.4f}, " f"rejected={reward_rejected[i].item():.4f}, " f"margin={margin:.4f}, correct={correct}" ) # Overall stats loss, acc = reward_model_loss(reward_chosen, reward_rejected, label_smoothing=label_smoothing) avg_margin = (reward_chosen - reward_rejected).mean().item() print_with_time_master( f"\nVal batch stats: loss={loss.item():.4f}, acc={acc.item():.3f}, " f"avg_margin={avg_margin:.4f}" ) print_with_time_master( f"Chosen rewards: mean={reward_chosen.mean().item():.4f}, " f"std={reward_chosen.std().item():.4f}" ) print_with_time_master( f"Rejected rewards: mean={reward_rejected.mean().item():.4f}, " f"std={reward_rejected.std().item():.4f}" ) # Check raw reward distribution (all tokens, not just scalar) reward_logits_flat = reward_logits.flatten() print_with_time_master( f"Raw token rewards (all): mean={reward_logits_flat.mean().item():.4f}, " f"std={reward_logits_flat.std().item():.4f}, " f"min={reward_logits_flat.min().item():.4f}, " f"max={reward_logits_flat.max().item():.4f}" ) # Debug train set (mixed from all datasets, like actual training) print_with_time_master(f"\n--- TRAIN SET (first batch - mixed datasets) ---") idxs, X, Y, loss_start_index_list = get_batch( data_sampling_info, "train", inference=True, dummy_data=dummy_data, suppress_text=suppress_text, load_dpo_pair=True, return_idx=True, ) with ctx: output = model(X, return_logits=True) reward_logits = output["reward_logits"] loss_end_index_list = compute_loss_end_indices(Y, reward_logits.shape[1]) rewards = extract_scalar_rewards(reward_logits, loss_start_index_list, loss_end_index_list) reward_rejected = rewards[::2] reward_chosen = rewards[1::2] print_with_time_master(f"Batch size: {len(idxs)}, Pairs: {len(reward_chosen)}") print_with_time_master(f"Y shape: {Y.shape}") print_with_time_master(f"loss_start_indices: {loss_start_index_list[:4]}...") print_with_time_master(f"loss_end_indices: {loss_end_index_list[:4]}...") # Check if Y actually has padding y_first_codebook = Y[:, 0, :] # (batch, seq_len-1) has_padding = (y_first_codebook == -1).any(dim=1) num_with_padding = has_padding.sum().item() print_with_time_master(f"Samples with padding (-1 in Y): {num_with_padding}/{Y.shape[0]}") # For each pair, show actual sequence lengths print_with_time_master(f"\nPair-wise sequence lengths:") for i in range(min(3, len(reward_chosen))): rej_idx = i * 2 cho_idx = i * 2 + 1 rej_end = loss_end_index_list[rej_idx] cho_end = loss_end_index_list[cho_idx] rej_len = rej_end - loss_start_index_list[rej_idx] cho_len = cho_end - loss_start_index_list[cho_idx] same_len = "✓ SAME" if rej_len == cho_len else "✗ DIFFERENT" print_with_time_master( f" Pair {i}: rejected len={rej_len}, chosen len={cho_len} [{same_len}]" ) print_with_time_master(f"\nFirst 5 pairs:") for i in range(min(5, len(reward_chosen))): margin = (reward_chosen[i] - reward_rejected[i]).item() correct = reward_chosen[i] > reward_rejected[i] print_with_time_master( f" Pair {i}: chosen={reward_chosen[i].item():.4f}, " f"rejected={reward_rejected[i].item():.4f}, " f"margin={margin:.4f}, correct={correct}" ) # Overall stats loss, acc = reward_model_loss(reward_chosen, reward_rejected, label_smoothing=label_smoothing) avg_margin = (reward_chosen - reward_rejected).mean().item() print_with_time_master( f"\nTrain batch stats: loss={loss.item():.4f}, acc={acc.item():.3f}, " f"avg_margin={avg_margin:.4f}" ) print_with_time_master( f"Chosen rewards: mean={reward_chosen.mean().item():.4f}, " f"std={reward_chosen.std().item():.4f}" ) print_with_time_master( f"Rejected rewards: mean={reward_rejected.mean().item():.4f}, " f"std={reward_rejected.std().item():.4f}" ) # Check raw reward distribution reward_logits_flat = reward_logits.flatten() print_with_time_master( f"Raw token rewards (all): mean={reward_logits_flat.mean().item():.4f}, " f"std={reward_logits_flat.std().item():.4f}, " f"min={reward_logits_flat.min().item():.4f}, " f"max={reward_logits_flat.max().item():.4f}" ) print_with_time_master("=" * 80 + "\n") model.train() # training loop print_with_time_master("training reward model...") t0 = time.time() t00 = time.time() t_start = time.time() local_iter_num = 0 mfu = 0 tokens_per_s = 0 effective_tokens_per_s_per_node = 0 running_loss = [] running_accuracy = [] running_margin = [] running_loss_seq = [] running_loss_tok = [] running_acc_tok = [] # get initial batch local_seen_idxs = set() data_idx_lists_flat = sorted([idx for sublist in train_data_idx_lists for idx in sublist]) local_idx_size = int(math.ceil(len(data_idx_lists_flat) / world_size / 2)) local_idxs_fixed = [ (local_idx_size * ddp_rank + i) % len(data_idx_lists_flat) for i in range(local_idx_size) ] if local_shuffle_data: random.seed(custom_seed_offset) random.shuffle(local_idxs_fixed) input_ids = local_idxs_fixed[0 : batch_size // 2] idxs, X, Y, loss_start_index_list = get_batch( data_sampling_info, "train", dummy_data=dummy_data, row_idx=None if shuffle_data else input_ids, suppress_text=suppress_text, load_dpo_pair=True, return_idx=True, ) print_with_time_master(f"First input_ids: {input_ids}") for idx in idxs: local_seen_idxs.add(idx) gpu_memory_monitor.reset_peak_stats() # Debug first batches to understand initial metrics dist_barrier() debug_first_batches() dist_barrier() while True: # determine and set the learning rate for this iteration lr = get_lr(iter_num) for param_group in optimizer.param_groups: param_group["lr"] = lr # evaluate the loss on train/val sets and write checkpoints if iter_num % eval_interval == 0 or iter_num == max_iters - 1: dist_barrier() time_since_last_loss = time.time() - t00 t00 = time.time() losses = estimate_loss() estimation_time = time.time() - t00 eval_time_pct = np.clip(estimation_time / time_since_last_loss * 100, 0, 100) if master_process: print_with_time_master( f"loss estimation took {estimation_time:.1f} seconds. ({eval_time_pct:.1f}% of loop)" ) if use_token_level_loss: print_with_time_master( f"step {iter_num}: " f"train loss {losses['train/loss']:.4f} (seq:{losses['train/loss_seq']:.4f} tok:{losses['train/loss_tok']:.4f}), " f"val loss {losses['val/loss']:.4f} (seq:{losses['val/loss_seq']:.4f} tok:{losses['val/loss_tok']:.4f}), " f"train acc {losses['train/accuracy']:.3f} (tok:{losses['train/acc_tok']:.3f}), " f"val acc {losses['val/accuracy']:.3f} (tok:{losses['val/acc_tok']:.3f}), " f"margin {losses['train/margin']:.3f}" ) else: print_with_time_master( f"step {iter_num}: train loss {losses['train/loss']:.4f}, " f"val loss {losses['val/loss']:.4f}, " f"train acc {losses['train/accuracy']:.3f}, " f"val acc {losses['val/accuracy']:.3f}, " f"train margin {losses['train/margin']:.3f}, " f"val margin {losses['val/margin']:.3f}" ) if wandb_log: log_dict = { "iter": iter_num, "n_tokens": total_tokens_processed, "lr": lr, } for k, v in losses.items(): log_dict[k] = v wandb.log(log_dict) dist_barrier() if iter_num > 0: if checkpoint_save_old_format: best_val_loss = save_old_checkpoint( out_dir, model, optimizer, best_val_loss, losses["val/loss"], step_save_iters, time_since_last_loss, model_args=model_args, iter_num=iter_num, n_tokens=total_tokens_processed, debug_val_only=debug_val_only, save_best_ckpt=False, # Don't save best checkpoint save_last_ckpt=True, # Only save last_ckpt_infer.pt ) else: best_val_loss = save_checkpoint( out_dir, model, optimizer, best_val_loss, losses["val/loss"], step_save_iters, time_since_last_loss, model_args=model_args, iter_num=iter_num, n_tokens=total_tokens_processed, debug_val_only=debug_val_only, ) dist_barrier() if eval_only: print_with_time_master("eval test done.") break # forward backward update for micro_step in range(gradient_accumulation_steps): if ddp and micro_step < gradient_accumulation_steps - 1: grad_sync_context = model.no_sync else: grad_sync_context = nullcontext with grad_sync_context(): with ctx: # Get output from model (dict with rewards and generation logits) output = model(X, return_logits=True) reward_logits = output["reward_logits"] # (batch, seq_len) # Compute end indices from Y loss_end_index_list = compute_loss_end_indices(Y, reward_logits.shape[1]) # Sequence-level loss (existing) rewards = extract_scalar_rewards( reward_logits, loss_start_index_list, loss_end_index_list ) # (batch,) reward_rejected_seq = rewards[::2] # even indices reward_chosen_seq = rewards[1::2] # odd indices loss_seq, acc_seq = reward_model_loss( reward_chosen_seq, reward_rejected_seq, label_smoothing=label_smoothing ) # Token-level loss (new - RAD-inspired) if use_token_level_loss: reward_logits_chosen = reward_logits[1::2] # odd indices reward_logits_rejected = reward_logits[::2] # even indices loss_tok, acc_tok = token_level_reward_loss( reward_logits_chosen, reward_logits_rejected, loss_start_index_list[1::2], # chosen start loss_start_index_list[::2], # rejected start loss_end_index_list[1::2], # chosen end loss_end_index_list[::2], # rejected end beta=token_level_beta, ) # Combined loss loss = loss_seq + lambda_token * loss_tok # Track both metrics loss_seq_val = loss_seq.item() loss_tok_val = loss_tok.item() acc_tok_val = acc_tok.item() else: loss = loss_seq loss_seq_val = loss_seq.item() loss_tok_val = 0.0 acc_tok_val = 0.0 loss = loss / gradient_accumulation_steps loss_val = loss.item() * gradient_accumulation_steps accuracy_val = acc_seq.item() margin_val = (reward_chosen_seq - reward_rejected_seq).mean().item() total_tokens_processed += X.shape[0] * X.shape[-1] * world_size rel_tokens_processed += X.shape[0] * X.shape[-1] * world_size # prefetch next batch iter_retrieval_start_idx = (local_iter_num + 1) * batch_size // 2 iter_retrieval_end_idx = (local_iter_num + 2) * batch_size // 2 input_ids = [ local_idxs_fixed[i % len(local_idxs_fixed)] for i in range(iter_retrieval_start_idx, iter_retrieval_end_idx) ] idxs, X, Y, loss_start_index_list = get_batch( data_sampling_info, "train", dummy_data=dummy_data, row_idx=None if shuffle_data else input_ids, suppress_text=suppress_text, load_dpo_pair=True, return_idx=True, ) for idx in idxs: local_seen_idxs.add(idx) if iter_num == 0 and micro_step == 0: torch.cuda.empty_cache() loss.backward() running_loss.append(loss_val) running_accuracy.append(accuracy_val) running_margin.append(margin_val) running_loss_seq.append(loss_seq_val) running_loss_tok.append(loss_tok_val) running_acc_tok.append(acc_tok_val) # clip gradient if grad_clip != 0.0: if fsdp: grad_norm = model.clip_grad_norm_(grad_clip) if torch.isnan(grad_norm): raise RuntimeError("Found NaN grad") else: grad_norm = torch.nn.utils.clip_grad_norm_( model.parameters(), grad_clip, error_if_nonfinite=True ) grad_norm = grad_norm.item() optimizer.step() optimizer.zero_grad(set_to_none=True) if master_process and wandb_log and grad_norm is not None: wandb.log( { "iter": iter_num, "n_tokens": total_tokens_processed, "misc/grad_norm": grad_norm, } ) if iter_num % 300 == 1: assert not gc.isenabled() gc.collect() # timing and logging t1 = time.time() dt = t1 - t0 t0 = t1 if iter_num % log_interval == 0 or iter_num == max_iters - 1: if master_process: tok_per_batch = block_size if not pack else batch_size_tokens if local_iter_num >= 5: mfu = estimate_mfu_no_model( raw_model_n_params, n_layer, n_head, n_head * d_head, block_size, batch_size * gradient_accumulation_steps, dt, ) mfu *= tok_per_batch / block_size tokens_per_s = world_size * batch_size * tok_per_batch * gradient_accumulation_steps / dt effective_tokens_per_s_per_node = ( rel_tokens_processed / (time.time() - t_start) / int(round(world_size / n_gpus_per_node)) ) avg_loss = np.mean(running_loss) avg_accuracy = np.mean(running_accuracy) avg_margin = np.mean(running_margin) avg_loss_seq = np.mean(running_loss_seq) if len(running_loss_seq) > 0 else 0.0 avg_loss_tok = np.mean(running_loss_tok) if len(running_loss_tok) > 0 else 0.0 avg_acc_tok = np.mean(running_acc_tok) if len(running_acc_tok) > 0 else 0.0 running_loss = [] running_accuracy = [] running_margin = [] running_loss_seq = [] running_loss_tok = [] running_acc_tok = [] gpu_mem_stats = gpu_memory_monitor.get_peak_stats() if use_token_level_loss: print_with_time_master( f"{color.cyan}iter {iter_num}:" f"{color.green} loss {avg_loss:.3f} (seq:{avg_loss_seq:.3f} tok:{avg_loss_tok:.3f})," f"{color.blue} acc {avg_accuracy:.3f} (tok:{avg_acc_tok:.3f})," f"{color.magenta} margin {avg_margin:.3f}," f"{color.yellow} step {dt * 1000:.1f}ms," f"{color.red} {tokens_per_s / 1e3:,.0f}k tok/s" f"{color.reset}" ) else: print_with_time_master( f"{color.cyan}iter {iter_num}:" f"{color.green} loss {avg_loss:.3f}," f"{color.blue} acc {avg_accuracy:.3f}," f"{color.magenta} margin {avg_margin:.3f}," f"{color.yellow} step_time {dt * 1000:.1f}ms," f"{color.red} throughput {tokens_per_s / 1e3:,.0f}k tok/s," f"{color.reset}" ) if wandb_log: log_dict = { "iter": iter_num, "n_tokens": total_tokens_processed, "train/running_loss": avg_loss, "train/running_accuracy": avg_accuracy, "train/running_margin": avg_margin, "lr": lr, "mfu": mfu * 100, "eff_tok/s/node": effective_tokens_per_s_per_node, "tok/s": tokens_per_s, "memory/max_reserved(GiB)": gpu_mem_stats.max_reserved_gib, "memory/max_reserved(%)": gpu_mem_stats.max_reserved_pct, } if use_token_level_loss: log_dict.update( { "train/loss_seq": avg_loss_seq, "train/loss_tok": avg_loss_tok, "train/acc_tok": avg_acc_tok, } ) wandb.log(log_dict) gpu_memory_monitor.reset_peak_stats() iter_num += 1 local_iter_num += 1 # termination if iter_num >= max_iters: print_with_time_master("done.") break dist_barrier() if ddp: destroy_process_group()