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