diff --git a/ACKNOWLEDGMENTS.md b/ACKNOWLEDGMENTS.md index b9686e733..53d75fb9f 100644 --- a/ACKNOWLEDGMENTS.md +++ b/ACKNOWLEDGMENTS.md @@ -9,4 +9,4 @@ MLX LM was developed with contributions from the following individuals: - Shunta Saito: Added support for PLaMo models. - Prince Canuma: Helped add support for `Starcoder2` models. -- Gökdeniz Gülmez: Added support for the following architectures: OpenBMB's `MiniCPM` and `MiniCPM3`, Kyutai's `Helium`, State-Space's`Mamba v1`, Z.ai & THUKEG's `GLM4`, and Allenai's `OLMoE`; Added support for the following training algorithms: `full-fine-tuning`; Added support for the following other features: `Multiple Optimizers to choose for training`, and `reporting training metrics to WandB (Weights & Biases)`. +- Gökdeniz Gülmez: Added support for the following architectures: OpenBMB's `MiniCPM` and `MiniCPM3`, Kyutai's `Helium`, State-Space's`Mamba v1`, Z.ai & THUKEG's `GLM4`, and Allenai's `OLMoE`; Added support for the following training algorithms: `full-fine-tuning`, and `Odds Ratio Preference Optimization (ORPO)`; Added support for the following other features: `Multiple Optimizers to choose for training`, and `reporting training metrics to WandB (Weights & Biases)`. diff --git a/mlx_lm/LORA.md b/mlx_lm/LORA.md index c0756c1ef..c697a31fb 100644 --- a/mlx_lm/LORA.md +++ b/mlx_lm/LORA.md @@ -18,6 +18,7 @@ LoRA (QLoRA).[^qlora] LoRA fine-tuning works with the following model families: - [Run](#Run) - [Fine-tune](#Fine-tune) + - [ORPO-Training](#ORPO-Training) - [Evaluate](#Evaluate) - [Generate](#Generate) - [Fuse](#Fuse) @@ -87,7 +88,58 @@ The default training computes a loss for every token in the sample. You can ignore the prompt and compute loss for just the completion by passing `--mask-prompt`. Note this is only supported for `chat` and `completion` datasets. For `chat` datasets the final message in the message list is -considered the completion. See the [dataset section](#Data) for more details. +considered the completion. See the [dataset section](#Data) for more details. + +### ORPO-Training + +Odds Ratio Preference Optimization (ORPO) training fine-tunes models using human preference data. Usage: + +```shell +mlx_lm.lora \ + --model \ + --train \ + --training-mode orpo \ + --data \ + --beta 0.1 +``` + +Parameters: + +- `--beta`: Temperature for logistic function (default: 0.1) + +Data format (JSONL): + +```jsonl +# Basic format with string responses +{"prompt": "User prompt", "chosen": "Preferred response", "rejected": "Less preferred response"} + +# With custom preference score +{"prompt": "User prompt", "chosen": "Preferred response", "rejected": "Less preferred response", "preference_score": 8.0} + +# With system message +{"prompt": "User prompt", "chosen": "Preferred response", "rejected": "Less preferred response", "system": "System instruction"} + +# With full conversation objects +{ + "prompt": "User prompt", + "chosen": { + "messages": [ + {"role": "system", "content": "System instruction"}, + {"role": "user", "content": "User message"}, + {"role": "assistant", "content": "Assistant response"} + ] + }, + "rejected": { + "messages": [ + {"role": "system", "content": "System instruction"}, + {"role": "user", "content": "User message"}, + {"role": "assistant", "content": "Assistant response"} + ] + } +} +``` + +The trainer assigns binary rewards (1.0 chosen, 0.0 rejected) if no explicit rewards provided via `preference_score`. ### Evaluate diff --git a/mlx_lm/examples/lora_config.yaml b/mlx_lm/examples/lora_config.yaml index a4fe069ba..49153df56 100644 --- a/mlx_lm/examples/lora_config.yaml +++ b/mlx_lm/examples/lora_config.yaml @@ -7,6 +7,9 @@ train: true # The fine-tuning method: "lora", "dora", or "full". fine_tune_type: lora +# The training-mode: "normal", or "dpo" +training_mode: normal + # The Optimizer with its possible inputs optimizer: adamw # optimizer_config: @@ -89,4 +92,9 @@ lora_parameters: # valid_split: "train[-100:]" # prompt_feature: "text" # completion_feature: "summary" - +# For ORPO training +# prompt_feature: "prompt" +# system_feature: "system" +# chosen_feature: "chosen" +# rejected_feature: "rejected" +# preference_score_feature: "preference_score" diff --git a/mlx_lm/lora.py b/mlx_lm/lora.py index 32a09a1c0..0b4126d20 100644 --- a/mlx_lm/lora.py +++ b/mlx_lm/lora.py @@ -13,6 +13,7 @@ from .tuner.callbacks import WandBCallback from .tuner.datasets import CacheDataset, load_dataset +from .tuner.orpo_trainer import ORPOTrainingArgs, evaluate_orpo, train_orpo from .tuner.trainer import TrainingArgs, TrainingCallback, evaluate, train from .tuner.utils import ( build_schedule, @@ -42,6 +43,7 @@ "model": "mlx_model", "train": False, "fine_tune_type": "lora", + "training_mode": "normal", "optimizer": "adam", "optimizer_config": { "adam": {}, @@ -68,6 +70,9 @@ "lora_parameters": {"rank": 8, "dropout": 0.0, "scale": 10.0}, "mask_prompt": False, "wandb": None, + # ORPO args + "beta": 0.1, + "reward_scaling": 1.0, } @@ -100,6 +105,12 @@ def build_parser(): choices=["lora", "dora", "full"], help="Type of fine-tuning to perform: lora, dora, or full.", ) + parser.add_argument( + "--training-mode", + type=str, + choices=["normal", "dpo", "orpo"], + help="Training mode: normal, DPO or ORPO.", + ) parser.add_argument( "--optimizer", type=str, @@ -186,6 +197,20 @@ def build_parser(): help="WandB project name to report training metrics. Disabled if None.", ) parser.add_argument("--seed", type=int, help="The PRNG seed") + + # ORPO args + parser.add_argument( + "--beta", + type=float, + help="Temperature parameter for ORPO training.", + default=0.1, + ) + parser.add_argument( + "--reward-scaling", + type=float, + help="Reward scaling factor for ORPO training, not implemented.", + default=1.0, + ) return parser @@ -231,18 +256,7 @@ def train_model( adapter_file = adapter_path / "adapters.safetensors" save_config(vars(args), adapter_path / "adapter_config.json") - # init training args - training_args = TrainingArgs( - batch_size=args.batch_size, - iters=args.iters, - val_batches=args.val_batches, - steps_per_report=args.steps_per_report, - steps_per_eval=args.steps_per_eval, - steps_per_save=args.save_every, - adapter_file=adapter_file, - max_seq_length=args.max_seq_length, - grad_checkpoint=args.grad_checkpoint, - ) + model.train() # Initialize the selected optimizer lr = build_schedule(args.lr_schedule) if args.lr_schedule else args.learning_rate @@ -259,29 +273,84 @@ def train_model( opt = opt_class(learning_rate=lr, **optimizer_config) - # Train model - train( - model=model, - args=training_args, - optimizer=opt, - train_dataset=CacheDataset(train_set), - val_dataset=CacheDataset(valid_set), - training_callback=training_callback, - ) + if args.training_mode == "orpo": + training_args = ORPOTrainingArgs( + batch_size=args.batch_size, + iters=args.iters, + val_batches=args.val_batches, + steps_per_report=args.steps_per_report, + steps_per_eval=args.steps_per_eval, + steps_per_save=args.save_every, + adapter_file=adapter_file, + max_seq_length=args.max_seq_length, + grad_checkpoint=args.grad_checkpoint, + beta=args.beta, + reward_scaling=args.reward_scaling, + ) + + train_orpo( + model=model, + optimizer=opt, + train_dataset=CacheDataset(train_set), + val_dataset=CacheDataset(valid_set), + args=training_args, + training_callback=training_callback, + ) + else: + training_args = TrainingArgs( + batch_size=args.batch_size, + iters=args.iters, + val_batches=args.val_batches, + steps_per_report=args.steps_per_report, + steps_per_eval=args.steps_per_eval, + steps_per_save=args.save_every, + adapter_file=adapter_file, + max_seq_length=args.max_seq_length, + grad_checkpoint=args.grad_checkpoint, + ) + + train( + model=model, + args=training_args, + optimizer=opt, + train_dataset=CacheDataset(train_set), + val_dataset=CacheDataset(valid_set), + training_callback=training_callback, + ) def evaluate_model(args, model: nn.Module, test_set): - test_loss = evaluate( - model=model, - dataset=CacheDataset(test_set), - batch_size=args.batch_size, - num_batches=args.test_batches, - max_seq_length=args.max_seq_length, - ) + model.eval() + + if args.training_mode == "orpo": + test_loss, test_rewards, _, test_metrics = evaluate_orpo( + model=model, + dataset=test_set, + batch_size=args.batch_size, + num_batches=args.test_batches, + max_seq_length=args.max_seq_length, + beta=args.beta, + ) + test_ppl = math.exp(test_loss) + print( + f"Test loss {test_loss:.3f}, Test ppl {test_ppl:.3f}, Rewards: {test_rewards[0]:.3f}, {test_rewards[1]:.3f}" + ) + + print("ORPO Test Metrics:") + for metric_name, metric_value in test_metrics.items(): + print(f" {metric_name}: {float(metric_value):.3f}") + else: + test_loss = evaluate( + model=model, + dataset=test_set, + batch_size=args.batch_size, + num_batches=args.test_batches, + max_seq_length=args.max_seq_length, + ) - test_ppl = math.exp(test_loss) + test_ppl = math.exp(test_loss) - print(f"Test loss {test_loss:.3f}, Test ppl {test_ppl:.3f}.") + print(f"Test loss {test_loss:.3f}, Test ppl {test_ppl:.3f}.") def run(args, training_callback: TrainingCallback = None): diff --git a/mlx_lm/tuner/datasets.py b/mlx_lm/tuner/datasets.py index c98263a4b..a37a3673b 100644 --- a/mlx_lm/tuner/datasets.py +++ b/mlx_lm/tuner/datasets.py @@ -1,11 +1,133 @@ import json import types from pathlib import Path -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Union from transformers import PreTrainedTokenizer +class ORPODataset: + def __init__( + self, + data: List[Dict[str, Union[str, Dict, List]]], + tokenizer: PreTrainedTokenizer, + prompt_key: str = "prompt", + chosen_key: str = "chosen", + rejected_key: str = "rejected", + preference_score_key: str = "preference_score", + system_key: str = None, + ): + self._chosen_data = [] + self._rejected_data = [] + self._scores = [] + + for d in data: + prompt_content = d.get(prompt_key, d.get("question", "")) + + if system_key and system_key in d: + base_messages = [{"role": "system", "content": d[system_key]}] + chosen_messages = base_messages + [ + {"role": "user", "content": prompt_content} + ] + rejected_messages = base_messages + [ + {"role": "user", "content": prompt_content} + ] + + if isinstance(d[chosen_key], str): + chosen_messages.append( + {"role": "assistant", "content": d[chosen_key]} + ) + elif isinstance(d[chosen_key], dict): + if "messages" in d[chosen_key]: + chosen_messages.extend(d[chosen_key]["messages"]) + else: + chosen_messages.append( + { + "role": "assistant", + "content": d[chosen_key].get("content", ""), + } + ) + elif isinstance(d[chosen_key], list): + chosen_messages.extend(d[chosen_key]) + + if isinstance(d[rejected_key], str): + rejected_messages.append( + {"role": "assistant", "content": d[rejected_key]} + ) + elif isinstance(d[rejected_key], dict): + if "messages" in d[rejected_key]: + rejected_messages.extend(d[rejected_key]["messages"]) + else: + rejected_messages.append( + { + "role": "assistant", + "content": d[rejected_key].get("content", ""), + } + ) + elif isinstance(d[rejected_key], list): + rejected_messages.extend(d[rejected_key]) + + chosen_text = tokenizer.apply_chat_template(chosen_messages) + rejected_text = tokenizer.apply_chat_template(rejected_messages) + + else: + chosen_content = self._extract_content(d[chosen_key]) + rejected_content = self._extract_content(d[rejected_key]) + + chosen_text = tokenizer.apply_chat_template( + [ + {"role": "user", "content": prompt_content}, + {"role": "assistant", "content": chosen_content}, + ] + ) + rejected_text = tokenizer.apply_chat_template( + [ + {"role": "user", "content": prompt_content}, + {"role": "assistant", "content": rejected_content}, + ] + ) + + self._chosen_data.append(chosen_text) + self._rejected_data.append(rejected_text) + + if preference_score_key in d: + self._scores.append(float(d[preference_score_key])) + else: + self._scores.append(1.0) + + def _extract_content(self, data): + """Helper method to extract content from various data formats.""" + if isinstance(data, str): + return data + elif isinstance(data, dict): + if "messages" in data: + last_message = data["messages"][-1] + return last_message.get("content", last_message.get("messages", "")) + return data.get("content", "") + elif isinstance(data, list): + last_message = data[-1] + if isinstance(last_message, dict): + if "content" in last_message: + return last_message["content"] + elif "messages" in last_message: + return last_message["messages"] + return last_message if isinstance(last_message, str) else "" + return "" + + def __len__(self): + return len(self._chosen_data) + + def process(self, d): + return d + + def __getitem__(self, idx: int): + return { + "chosen": self._chosen_data[idx], + "rejected": self._rejected_data[idx], + "preference_score": self._scores[idx], + } + + class TextDataset: """ Light-weight wrapper to hold a dataset. @@ -161,27 +283,51 @@ def create_dataset( ): mask_prompt = getattr(config, "mask_prompt", False) prompt_feature = getattr(config, "prompt_feature", "prompt") + system_feature = getattr(config, "system_feature", "system") + chosen_feature = getattr(config, "chosen_feature", "chosen") + rejected_feature = getattr(config, "rejected_feature", "rejected") + preference_score_feature = getattr( + config, "preference_score_feature", "preference_score" + ) text_feature = getattr(config, "text_feature", "text") completion_feature = getattr(config, "completion_feature", "completion") chat_feature = getattr(config, "chat_feature", "messages") + training_mode = getattr(config, "training_mode", "normal") sample = data[0] - if prompt_feature in sample and completion_feature in sample: - return CompletionsDataset( - data, tokenizer, prompt_feature, completion_feature, mask_prompt - ) - elif chat_feature in sample: - return ChatDataset( - data, tokenizer, chat_key=chat_feature, mask_prompt=mask_prompt - ) - elif text_feature in sample: - if mask_prompt: - raise ValueError("Prompt masking not supported for text dataset.") - return TextDataset(data, tokenizer, text_key=text_feature) + if training_mode == "normal": + if prompt_feature in sample and completion_feature in sample: + return CompletionsDataset( + data, tokenizer, prompt_feature, completion_feature, mask_prompt + ) + elif chat_feature in sample: + return ChatDataset( + data, tokenizer, chat_key=chat_feature, mask_prompt=mask_prompt + ) + elif text_feature in sample: + if mask_prompt: + raise ValueError("Prompt masking not supported for text dataset.") + return TextDataset(data, tokenizer, text_key=text_feature) + else: + raise ValueError( + "Unsupported data format, check the supported formats here:\n" + "https://github.com/ml-explore/mlx-examples/blob/main/llms/mlx_lm/LORA.md#data." + ) else: - raise ValueError( - "Unsupported data format, check the supported formats here:\n" - "https://github.com/ml-explore/mlx-examples/blob/main/llms/mlx_lm/LORA.md#data." - ) + if chosen_feature in sample and rejected_feature in sample: + return ORPODataset( + data=data, + tokenizer=tokenizer, + system_key=system_feature, + prompt_key=prompt_feature, + chosen_key=chosen_feature, + rejected_key=rejected_feature, + preference_score_key=preference_score_feature, + ) + else: + raise ValueError( + "Unsupported data format, check the supported formats here:\n" + "https://github.com/ml-explore/mlx-examples/blob/main/llms/mlx_lm/LORA.md#GRPO-Training." + ) def load_local_dataset( diff --git a/mlx_lm/tuner/orpo_trainer.py b/mlx_lm/tuner/orpo_trainer.py new file mode 100644 index 000000000..378bdd1ad --- /dev/null +++ b/mlx_lm/tuner/orpo_trainer.py @@ -0,0 +1,417 @@ +import time +from dataclasses import dataclass, field +from pathlib import Path + +import mlx.core as mx +import mlx.nn as nn +import numpy as np +from mlx.nn.utils import average_gradients +from mlx.utils import tree_flatten + +from mlx_lm.tuner.callbacks import TrainingCallback + +from .trainer import TrainingArgs, grad_checkpoint + + +@dataclass +class ORPOTrainingArgs(TrainingArgs): + beta: float = field( + default=0.1, metadata={"help": "Temperature parameter for ORPO training."} + ) + reward_scaling: float = field( + default=1.0, + metadata={"help": "Reward scaling factor for ORPO training, not implemented."}, + ) + + +def get_logps(model, tokens, mask): + inputs = tokens[:, :-1] + targets = tokens[:, 1:] + logits = model(inputs) + log_probs = -nn.losses.cross_entropy(logits, targets, reduction="none") + mask = mask[:, :-1] + seq_lengths = mask.sum(-1) + logp_seq_avg = (log_probs * mask).sum(-1) / seq_lengths + logits_mean = logits.sum() / mask.sum() + return logp_seq_avg, logits_mean + + +def orpo_loss( + chosen_logps, + chosen_logits_mean, + rejected_logps, + rejected_logits_mean, + chosen_masks, + rejected_masks, + preference_scores, + beta: float = 0.1, +): + chosen_logps = chosen_logps * preference_scores + + # Stable log-odds computation + log_odds = chosen_logps - rejected_logps + ratio = nn.log_sigmoid(log_odds) + loss = -beta * ratio + + # Reward estimation + chosen_reward = beta * chosen_logps + rejected_reward = beta * rejected_logps + reward = mx.stack([mx.mean(chosen_reward), mx.mean(rejected_reward)]) + + num_tokens = chosen_masks.sum() + rejected_masks.sum() + + metrics = { + "accuracies": mx.mean((chosen_reward > rejected_reward).astype(mx.float32)), + "margins": mx.mean(chosen_reward - rejected_reward), + "policy_chosen_logps": mx.mean(chosen_logps), + "policy_rejected_logps": mx.mean(rejected_logps), + "chosen_logits_mean": chosen_logits_mean, + "rejected_logits_mean": rejected_logits_mean, + } + + mx.clear_cache() + return mx.mean(loss), reward, num_tokens, metrics + + +def iterate_orpo_batches(dataset, batch_size, max_seq_length, train=False): + """Batch iterator for ORPO with preference scores""" + idx = sorted(range(len(dataset)), key=lambda idx: len(dataset[idx]["chosen"])) + + if len(dataset) < batch_size: + raise ValueError( + f"Dataset must have at least batch_size={batch_size}" + f" examples but only has {len(dataset)}." + ) + + step = mx.distributed.init().size() + if batch_size % step != 0: + raise ValueError("Batch size must be divisible by number of workers") + + batch_idx = [ + idx[i : i + batch_size : step] + for i in range(0, len(idx) - batch_size + 1, batch_size) + ] + + while True: + indices = ( + np.random.permutation(len(batch_idx)) if train else range(len(batch_idx)) + ) + for i in indices: + batch = [dataset[j] for j in batch_idx[i]] + + chosen_lengths = [len(x["chosen"]) for x in batch] + rejected_lengths = [len(x["rejected"]) for x in batch] + max_length = min( + max(max(chosen_lengths), max(rejected_lengths)), max_seq_length + ) + pad_to = 8 + max_length_in_batch = pad_to * ((max_length + pad_to - 1) // pad_to) + + batch_size_per_device = batch_size // step + chosen_arr = np.zeros( + (batch_size_per_device, max_length_in_batch), np.int32 + ) + rejected_arr = np.zeros( + (batch_size_per_device, max_length_in_batch), np.int32 + ) + chosen_masks = np.zeros( + (batch_size_per_device, max_length_in_batch), np.float32 + ) + rejected_masks = np.zeros( + (batch_size_per_device, max_length_in_batch), np.float32 + ) + + preference_scores = np.array( + [x.get("preference_score", 1.0) for x in batch], np.float32 + ) + + for j in range(batch_size_per_device): + chosen_length = min(chosen_lengths[j], max_length_in_batch) + rejected_length = min(rejected_lengths[j], max_length_in_batch) + + chosen_arr[j, :chosen_length] = batch[j]["chosen"][:chosen_length] + chosen_masks[j, :chosen_length] = 1.0 + rejected_arr[j, :rejected_length] = batch[j]["rejected"][ + :rejected_length + ] + rejected_masks[j, :rejected_length] = 1.0 + + yield ( + mx.array(chosen_arr), + mx.array(rejected_arr), + mx.array(chosen_masks), + mx.array(rejected_masks), + mx.array(preference_scores), + ) + + if not train: + break + + +def evaluate_orpo( + model, dataset, batch_size, num_batches, beta: float, max_seq_length=2048 +): + all_losses = 0 + all_rewards = mx.zeros((2,)) + all_metrics = None + ntokens = 0 + + index_iterator = iter(range(num_batches)) if num_batches != -1 else iter(int, 1) + for _, batch in zip( + index_iterator, + iterate_orpo_batches( + dataset=dataset, + batch_size=batch_size, + max_seq_length=max_seq_length, + ), + ): + chosen, rejected, chosen_masks, rejected_masks, preference_scores = batch + + chosen_logps, chosen_logits_mean = get_logps(model, chosen, chosen_masks) + rejected_logps, rejected_logits_mean = get_logps( + model, rejected, rejected_masks + ) + + lvalue, reward, toks, metrics = orpo_loss( + chosen_logps, + chosen_logits_mean, + rejected_logps, + rejected_logits_mean, + chosen_masks=chosen_masks, + rejected_masks=rejected_masks, + preference_scores=preference_scores, + beta=beta, + ) + all_losses += lvalue * toks + all_rewards += reward * toks + ntokens += toks + + if all_metrics is None: + all_metrics = {k: v * toks for k, v in metrics.items()} + else: + for k, v in metrics.items(): + all_metrics[k] += v * toks + + mx.eval(all_losses, all_rewards, ntokens) + all_losses = mx.distributed.all_sum(all_losses) + all_rewards = mx.distributed.all_sum(all_rewards) + ntokens = mx.distributed.all_sum(ntokens) + all_metrics = {k: mx.distributed.all_sum(v) for k, v in all_metrics.items()} + + avg_metrics = {k: (v / ntokens).item() for k, v in all_metrics.items()} + avg_rewards = (all_rewards / ntokens).tolist() + avg_loss = (all_losses / ntokens).item() + + return avg_loss, avg_rewards, ntokens, avg_metrics + + +def train_orpo( + model, + optimizer, + train_dataset, + val_dataset, + loss: callable = orpo_loss, + args: ORPOTrainingArgs = ORPOTrainingArgs(), + training_callback: TrainingCallback = None, +): + print(f"Starting ORPO training..., iters: {args.iters}") + world = mx.distributed.init() + world_size = world.size() + rank = world.rank() + + if world_size > 1: + print(f"Node {rank} of {world_size}") + + if args.grad_checkpoint: + grad_checkpoint(model.layers[0]) + + state = [model.state, optimizer.state] + + def step(batch): + chosen, rejected, chosen_masks, rejected_masks, preference_scores = batch + + chosen_logps, chosen_logits_mean = get_logps(model, chosen, chosen_masks) + rejected_logps, rejected_logits_mean = get_logps( + model, rejected, rejected_masks + ) + + (lvalue, reward, toks, metrics), grad = loss_value_and_grad( + chosen_logps, + chosen_logits_mean, + rejected_logps, + rejected_logits_mean, + chosen_masks, + rejected_masks, + preference_scores=preference_scores, + ) + + grad = average_gradients(grad) + optimizer.update(model, grad) + + return lvalue, reward, toks, metrics + + def loss_wrapper( + chosen_logps, + chosen_logits_mean, + rejected_logps, + rejected_logits_mean, + chosen_masks, + rejected_masks, + preference_scores, + ): + return loss( + chosen_logps=chosen_logps, + chosen_logits_mean=chosen_logits_mean, + rejected_logps=rejected_logps, + rejected_logits_mean=rejected_logits_mean, + chosen_masks=chosen_masks, + rejected_masks=rejected_masks, + preference_scores=preference_scores, + beta=args.beta, + ) + + loss_value_and_grad = nn.value_and_grad(model, loss_wrapper) + + losses = 0 + rewards = mx.zeros((2,)) + n_tokens = 0 + steps = 0 + trained_tokens = 0 + accumulated_metrics = { + "accuracies": 0, + "margins": 0, + "policy_rejected_logps": 0, + "policy_chosen_logps": 0, + "rejected_logits_mean": 0, + "chosen_logits_mean": 0, + } + + start = time.perf_counter() + for it, batch in zip( + range(1, args.iters + 1), + iterate_orpo_batches( + dataset=train_dataset, + batch_size=args.batch_size, + max_seq_length=args.max_seq_length, + train=True, + ), + ): + if it == 1 or it % args.steps_per_eval == 0 or it == args.iters: + stop = time.perf_counter() + val_loss, val_rewards, val_ntokens, val_metrics = evaluate_orpo( + model=model, + dataset=val_dataset, + batch_size=args.batch_size, + num_batches=args.val_batches, + max_seq_length=args.max_seq_length, + beta=args.beta, + ) + val_time = time.perf_counter() - stop + if rank == 0: + print( + f"Iter {it}: " + f"Val loss {val_loss:.3f}, " + f"Val chosen reward {val_rewards[0]:.3f}, " + f"Val rejected reward {val_rewards[1]:.3f}, " + f"Val accuracy {val_metrics['accuracies']:.3f}, " + f"Val margin {val_metrics['margins']:.3f}, " + f"Val took {val_time:.3f}s", + flush=True, + ) + + if training_callback is not None: + training_callback.on_val_loss_report( + { + "iteration": it, + "val_loss": val_loss, + "val_chosen_reward": val_rewards[0], + "val_rejected_reward": val_rewards[1], + **{f"val_{k}": v for k, v in val_metrics.items()}, + "val_time": val_time, + } + ) + + start = time.perf_counter() + + # Training step + lvalue, reward, toks, metrics = step(batch) + losses += lvalue + rewards += reward + n_tokens += toks + steps += 1 + + for k, v in metrics.items(): + accumulated_metrics[k] += v + + mx.eval(state, losses, rewards, n_tokens) + + if it % args.steps_per_report == 0 or it == args.iters: + stop = time.perf_counter() + + train_loss = mx.distributed.all_sum(losses).item() / (steps * world_size) + train_rewards = [ + r / (steps * world_size) + for r in mx.distributed.all_sum(rewards).tolist() + ] + avg_metrics = { + k: v / (steps * world_size) for k, v in accumulated_metrics.items() + } + n_tokens = mx.distributed.all_sum(n_tokens).item() + learning_rate = optimizer.learning_rate.item() + it_sec = args.steps_per_report / (stop - start) + tokens_sec = float(n_tokens) / (stop - start) + trained_tokens += n_tokens + peak_mem = mx.get_peak_memory() / 1e9 + + if rank == 0: + print( + f"Iter {it}: Train loss {train_loss:.3f}, " + f"Chosen reward {train_rewards[0]:.3f}, " + f"Rejected reward {train_rewards[1]:.3f}, " + f"Accuracy {avg_metrics['accuracies']:.3f}, " + f"Margin {avg_metrics['margins']:.3f}, " + f"Learning Rate {learning_rate:.3e}, " + f"It/sec {it_sec:.3f}, " + f"Tokens/sec {tokens_sec:.3f}, " + f"Peak mem {peak_mem:.3f} GB", + flush=True, + ) + + if training_callback is not None: + training_callback.on_train_loss_report( + { + "iteration": it, + "train_loss": train_loss, + "train_chosen_reward": train_rewards[0], + "train_rejected_reward": train_rewards[1], + **{f"train_{k}": v for k, v in avg_metrics.items()}, + "learning_rate": learning_rate, + "iterations_per_second": it_sec, + "tokens_per_second": tokens_sec, + "trained_tokens": trained_tokens, + "peak_memory": peak_mem, + } + ) + + losses = 0 + rewards = mx.zeros((2,)) + n_tokens = 0 + steps = 0 + accumulated_metrics = {k: 0 for k in accumulated_metrics} + start = time.perf_counter() + + if it % args.steps_per_save == 0: + adapter_weights = dict(tree_flatten(model.trainable_parameters())) + mx.save_safetensors(str(args.adapter_file), adapter_weights) + checkpoint = ( + Path(args.adapter_file).parent / f"{it:07d}_adapters.safetensors" + ) + mx.save_safetensors(str(checkpoint), adapter_weights) + print( + f"Iter {it}: Saved adapter weights to " + f"{args.adapter_file} and {checkpoint}." + ) + + adapter_weights = dict(tree_flatten(model.trainable_parameters())) + mx.save_safetensors(str(args.adapter_file), adapter_weights) + print(f"Saved final weights to {args.adapter_file}.")