From 7c38e30bf63c9581e1f237cfe959f0f58842b9f8 Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Fri, 14 Mar 2025 08:40:50 +0100 Subject: [PATCH 01/18] udoate datasets.py --- mlx_lm/tuner/datasets.py | 97 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 97 insertions(+) diff --git a/mlx_lm/tuner/datasets.py b/mlx_lm/tuner/datasets.py index a6f3bd295..8d8e1e974 100644 --- a/mlx_lm/tuner/datasets.py +++ b/mlx_lm/tuner/datasets.py @@ -7,6 +7,103 @@ 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 __getitem__(self, idx: int): + return { + "chosen": self._chosen_data[idx], + "rejected": self._rejected_data[idx], + "preference_score": self._scores[idx] + } + + class Dataset: """ Light-weight wrapper to hold a dataset. From a6a5b2b67e1bc9b1106b87453433cd9d43f2272f Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Fri, 14 Mar 2025 08:49:31 +0100 Subject: [PATCH 02/18] update lora.py --- mlx_lm/lora.py | 121 +++++++++++++++++++++++++++++++++++++++++-------- 1 file changed, 101 insertions(+), 20 deletions(-) diff --git a/mlx_lm/lora.py b/mlx_lm/lora.py index 042b40e2b..7de27d0cb 100644 --- a/mlx_lm/lora.py +++ b/mlx_lm/lora.py @@ -15,6 +15,7 @@ from .tokenizer_utils import TokenizerWrapper from .tuner.datasets import load_dataset from .tuner.trainer import TrainingArgs, TrainingCallback, evaluate, train +from .tuner.orpo_trainer import ORPOTrainingArgs, evaluate_orpo, train_orpo from .tuner.utils import ( build_schedule, linear_to_lora_layers, @@ -43,6 +44,7 @@ "model": "mlx_model", "train": False, "fine_tune_type": "lora", + "training_mode": "normal", "optimizer": "adam", "optimizer_config": { "adam": {}, @@ -68,6 +70,11 @@ "lr_schedule": None, "lora_parameters": {"rank": 8, "alpha": 16, "dropout": 0.0, "scale": 10.0}, "mask_prompt": False, + + # ORPO args + "beta": 0.1, + "reference_model_path": None, + "reward_scaling": 1.0, } @@ -100,6 +107,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, @@ -180,6 +193,20 @@ def build_parser(): default=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 @@ -255,33 +282,87 @@ def train_model( opt = opt_class(learning_rate=lr, **optimizer_config) - # Train model - train( - model=model, - tokenizer=tokenizer, - args=training_args, - optimizer=opt, - train_dataset=train_set, - val_dataset=valid_set, - training_callback=training_callback, - ) + # Train model based on training mode + 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, + tokenizer=tokenizer, + optimizer=opt, + train_dataset=train_set, + val_dataset=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 using SFT + train( + model=model, + tokenizer=tokenizer, + args=training_args, + optimizer=opt, + train_dataset=train_set, + val_dataset=valid_set, + training_callback=training_callback, + ) def evaluate_model(args, model: nn.Module, tokenizer: TokenizerWrapper, test_set): model.eval() - test_loss = evaluate( - model=model, - dataset=test_set, - tokenizer=tokenizer, - batch_size=args.batch_size, - num_batches=args.test_batches, - max_seq_length=args.max_seq_length, - ) + 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, + tokenizer=tokenizer, + 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): From 61e1504139469a485c42c48bc4a8f35479501d9c Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Fri, 14 Mar 2025 09:01:32 +0100 Subject: [PATCH 03/18] adding orpo_trainer.py + datastes.py fix --- mlx_lm/tuner/datasets.py | 44 +++-- mlx_lm/tuner/orpo_trainer.py | 357 +++++++++++++++++++++++++++++++++++ 2 files changed, 384 insertions(+), 17 deletions(-) create mode 100644 mlx_lm/tuner/orpo_trainer.py diff --git a/mlx_lm/tuner/datasets.py b/mlx_lm/tuner/datasets.py index 8d8e1e974..f74d5440e 100644 --- a/mlx_lm/tuner/datasets.py +++ b/mlx_lm/tuner/datasets.py @@ -2,7 +2,7 @@ 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 @@ -220,24 +220,34 @@ def create_dataset( 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 Dataset(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 Dataset(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" in sample and "rejected" in sample: + return ORPODataset(data, tokenizer) + 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." + ) 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..36dd48b94 --- /dev/null +++ b/mlx_lm/tuner/orpo_trainer.py @@ -0,0 +1,357 @@ +import time +from pathlib import Path +from dataclasses import dataclass, field + +import mlx.nn as nn +import mlx.core as mx +import numpy as np +from mlx.utils import tree_flatten +from mlx.nn.utils import average_gradients +from .trainer import TrainingArgs, grad_checkpoint, TrainingCallback + + +@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 orpo_loss(model, chosen, rejected, chosen_masks, rejected_masks, preference_scores, beta=0.1): + def get_logps(model, x, mask): + inputs = x[:, :-1] + targets = x[:, 1:] + logits = model(inputs) + logp = -nn.losses.cross_entropy(logits, targets, reduction='none') + seq_lengths = mask[:, :-1].sum(-1) + logp_sum = (logp * mask[:, :-1]).sum(-1) / seq_lengths + logits_mean = (logits * mask[:, :-1, None]).sum() / mask[:, :-1].sum() + return logp_sum, logits_mean + + policy_chosen_logps, chosen_logits_mean = get_logps(model, chosen, chosen_masks) + policy_rejected_logps, rejected_logits_mean = get_logps(model, rejected, rejected_masks) + + policy_chosen_logps = policy_chosen_logps * preference_scores + + log_odds = (policy_chosen_logps - policy_rejected_logps) - ( + mx.log1p(-mx.exp(policy_chosen_logps)) - mx.log1p(-mx.exp(policy_rejected_logps)) + ) + + ratio = nn.log_sigmoid(log_odds) + loss = -beta * ratio + + chosen_reward = beta * policy_chosen_logps + rejected_reward = beta * policy_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_rejected_logps': mx.mean(policy_rejected_logps), + 'policy_chosen_logps': mx.mean(policy_chosen_logps), + 'rejected_logits_mean': mx.mean(rejected_logits_mean), + 'chosen_logits_mean': mx.mean(chosen_logits_mean) + } + + 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 + lvalue, reward, toks, metrics = orpo_loss( + model=model, + chosen=chosen, + rejected=rejected, + 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, + tokenizer, + 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 + + (lvalue, reward, toks, metrics), grad = loss_value_and_grad( + model, + chosen, + rejected, + chosen_masks, + rejected_masks, + preference_scores=preference_scores, + ) + + grad = average_gradients(grad) + optimizer.update(model, grad) + + return lvalue, reward, toks, metrics + + def loss_wrapper(model, chosen, rejected, chosen_masks, rejected_masks, preference_scores): + return loss( + model=model, + chosen=chosen, + rejected=rejected, + 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) + + # Training loop with progress tracking + 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.metal.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}.") \ No newline at end of file From 2b0b8deef5269444d14cc8507135f80e56008fa2 Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Fri, 14 Mar 2025 09:10:57 +0100 Subject: [PATCH 04/18] udpate LORA.md + lora_config.yaml --- mlx_lm/LORA.md | 54 +++++++++++++++++++++++++++++++- mlx_lm/examples/lora_config.yaml | 3 ++ 2 files changed, 56 insertions(+), 1 deletion(-) diff --git a/mlx_lm/LORA.md b/mlx_lm/LORA.md index e863abc46..04fa0b5f5 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) @@ -82,7 +83,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 36bc1dff8..bafe2ccce 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: From 84a9535447a0cff50872b25d964a167fdd96778d Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Fri, 14 Mar 2025 09:52:15 +0100 Subject: [PATCH 05/18] nits + update Acnowledgements.md --- ACKNOWLEDGMENTS.md | 2 +- mlx_lm/lora.py | 13 ------------- 2 files changed, 1 insertion(+), 14 deletions(-) diff --git a/ACKNOWLEDGMENTS.md b/ACKNOWLEDGMENTS.md index 7c8021509..c9dc2c1ab 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 `MiniCPM`, `Helium`, `Mamba version 1`, `OLMoE` archtectures and support for `full-fine-tuning`. +- Gökdeniz Gülmez: Added support for the following architectures: OpenBMB's `MiniCPM`, Kyutai's `Helium`, State-Space's`Mamba v1`, and Allenai's `OLMoE`; added support for the following training algorithms: `full-fine-tuning` and `Odds Ratio Preference Optimization (ORPO)`. \ No newline at end of file diff --git a/mlx_lm/lora.py b/mlx_lm/lora.py index 7de27d0cb..4d1eb5dd8 100644 --- a/mlx_lm/lora.py +++ b/mlx_lm/lora.py @@ -252,19 +252,6 @@ 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 From 6cc3567954b44b1c091ad49c2764bae1cc0720d3 Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Tue, 18 Mar 2025 23:05:44 +0100 Subject: [PATCH 06/18] formatting --- mlx_lm/lora.py | 21 ++-- mlx_lm/tuner/datasets.py | 74 +++++++----- mlx_lm/tuner/orpo_trainer.py | 215 ++++++++++++++++++++--------------- 3 files changed, 185 insertions(+), 125 deletions(-) diff --git a/mlx_lm/lora.py b/mlx_lm/lora.py index 4d1eb5dd8..cb5b2f161 100644 --- a/mlx_lm/lora.py +++ b/mlx_lm/lora.py @@ -14,8 +14,8 @@ from .tokenizer_utils import TokenizerWrapper from .tuner.datasets import load_dataset -from .tuner.trainer import TrainingArgs, TrainingCallback, evaluate, train from .tuner.orpo_trainer import ORPOTrainingArgs, evaluate_orpo, train_orpo +from .tuner.trainer import TrainingArgs, TrainingCallback, evaluate, train from .tuner.utils import ( build_schedule, linear_to_lora_layers, @@ -70,7 +70,6 @@ "lr_schedule": None, "lora_parameters": {"rank": 8, "alpha": 16, "dropout": 0.0, "scale": 10.0}, "mask_prompt": False, - # ORPO args "beta": 0.1, "reference_model_path": None, @@ -199,13 +198,13 @@ def build_parser(): "--beta", type=float, help="Temperature parameter for ORPO training.", - default=0.1 + default=0.1, ) parser.add_argument( "--reward-scaling", type=float, help="Reward scaling factor for ORPO training, not implemented.", - default=1.0 + default=1.0, ) return parser @@ -282,9 +281,9 @@ def train_model( max_seq_length=args.max_seq_length, grad_checkpoint=args.grad_checkpoint, beta=args.beta, - reward_scaling=args.reward_scaling + reward_scaling=args.reward_scaling, ) - + train_orpo( model=model, tokenizer=tokenizer, @@ -292,7 +291,7 @@ def train_model( train_dataset=train_set, val_dataset=valid_set, args=training_args, - training_callback=training_callback + training_callback=training_callback, ) else: training_args = TrainingArgs( @@ -304,7 +303,7 @@ def train_model( steps_per_save=args.save_every, adapter_file=adapter_file, max_seq_length=args.max_seq_length, - grad_checkpoint=args.grad_checkpoint + grad_checkpoint=args.grad_checkpoint, ) # Train model using SFT @@ -329,10 +328,12 @@ def evaluate_model(args, model: nn.Module, tokenizer: TokenizerWrapper, test_set batch_size=args.batch_size, num_batches=args.test_batches, max_seq_length=args.max_seq_length, - beta=args.beta + 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( + 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(): diff --git a/mlx_lm/tuner/datasets.py b/mlx_lm/tuner/datasets.py index f74d5440e..1023c1087 100644 --- a/mlx_lm/tuner/datasets.py +++ b/mlx_lm/tuner/datasets.py @@ -16,64 +16,86 @@ def __init__( chosen_key: str = "chosen", rejected_key: str = "rejected", preference_score_key: str = "preference_score", - system_key: str = None + 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}] - + 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]}) + 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", "")}) + 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]}) + 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", "")}) + 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}, - ]) - + + 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): @@ -100,7 +122,7 @@ def __getitem__(self, idx: int): return { "chosen": self._chosen_data[idx], "rejected": self._rejected_data[idx], - "preference_score": self._scores[idx] + "preference_score": self._scores[idx], } diff --git a/mlx_lm/tuner/orpo_trainer.py b/mlx_lm/tuner/orpo_trainer.py index 36dd48b94..28d463e80 100644 --- a/mlx_lm/tuner/orpo_trainer.py +++ b/mlx_lm/tuner/orpo_trainer.py @@ -1,45 +1,50 @@ import time -from pathlib import Path from dataclasses import dataclass, field +from pathlib import Path -import mlx.nn as nn import mlx.core as mx +import mlx.nn as nn import numpy as np -from mlx.utils import tree_flatten from mlx.nn.utils import average_gradients -from .trainer import TrainingArgs, grad_checkpoint, TrainingCallback +from mlx.utils import tree_flatten + +from .trainer import TrainingArgs, TrainingCallback, grad_checkpoint @dataclass class ORPOTrainingArgs(TrainingArgs): beta: float = field( - default=0.1, - metadata={"help": "Temperature parameter for ORPO training."} + 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."} + metadata={"help": "Reward scaling factor for ORPO training, not implemented."}, ) -def orpo_loss(model, chosen, rejected, chosen_masks, rejected_masks, preference_scores, beta=0.1): +def orpo_loss( + model, chosen, rejected, chosen_masks, rejected_masks, preference_scores, beta=0.1 +): def get_logps(model, x, mask): inputs = x[:, :-1] targets = x[:, 1:] logits = model(inputs) - logp = -nn.losses.cross_entropy(logits, targets, reduction='none') + logp = -nn.losses.cross_entropy(logits, targets, reduction="none") seq_lengths = mask[:, :-1].sum(-1) logp_sum = (logp * mask[:, :-1]).sum(-1) / seq_lengths logits_mean = (logits * mask[:, :-1, None]).sum() / mask[:, :-1].sum() return logp_sum, logits_mean policy_chosen_logps, chosen_logits_mean = get_logps(model, chosen, chosen_masks) - policy_rejected_logps, rejected_logits_mean = get_logps(model, rejected, rejected_masks) + policy_rejected_logps, rejected_logits_mean = get_logps( + model, rejected, rejected_masks + ) policy_chosen_logps = policy_chosen_logps * preference_scores log_odds = (policy_chosen_logps - policy_rejected_logps) - ( - mx.log1p(-mx.exp(policy_chosen_logps)) - mx.log1p(-mx.exp(policy_rejected_logps)) + mx.log1p(-mx.exp(policy_chosen_logps)) + - mx.log1p(-mx.exp(policy_rejected_logps)) ) ratio = nn.log_sigmoid(log_odds) @@ -52,12 +57,12 @@ def get_logps(model, x, mask): 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_rejected_logps': mx.mean(policy_rejected_logps), - 'policy_chosen_logps': mx.mean(policy_chosen_logps), - 'rejected_logits_mean': mx.mean(rejected_logits_mean), - 'chosen_logits_mean': mx.mean(chosen_logits_mean) + "accuracies": mx.mean((chosen_reward > rejected_reward).astype(mx.float32)), + "margins": mx.mean(chosen_reward - rejected_reward), + "policy_rejected_logps": mx.mean(policy_rejected_logps), + "policy_chosen_logps": mx.mean(policy_chosen_logps), + "rejected_logits_mean": mx.mean(rejected_logits_mean), + "chosen_logits_mean": mx.mean(chosen_logits_mean), } return mx.mean(loss), reward, num_tokens, metrics @@ -65,66 +70,87 @@ def get_logps(model, x, mask): 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'])) - + 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)] - + + 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)) + 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) + + 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) - + 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_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_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) + mx.array(preference_scores), ) - + if not train: break -def evaluate_orpo(model, dataset, batch_size, num_batches, beta: float, max_seq_length=2048): +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, @@ -142,12 +168,12 @@ def evaluate_orpo(model, dataset, batch_size, num_batches, beta: float, max_seq_ chosen_masks=chosen_masks, rejected_masks=rejected_masks, preference_scores=preference_scores, - beta=beta + 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: @@ -159,11 +185,11 @@ def evaluate_orpo(model, dataset, batch_size, num_batches, beta: float, max_seq_ 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 @@ -181,7 +207,7 @@ def train_orpo( world = mx.distributed.init() world_size = world.size() rank = world.rank() - + if world_size > 1: print(f"Node {rank} of {world_size}") @@ -192,12 +218,12 @@ def train_orpo( def step(batch): chosen, rejected, chosen_masks, rejected_masks, preference_scores = batch - + (lvalue, reward, toks, metrics), grad = loss_value_and_grad( - model, - chosen, - rejected, - chosen_masks, + model, + chosen, + rejected, + chosen_masks, rejected_masks, preference_scores=preference_scores, ) @@ -207,7 +233,9 @@ def step(batch): return lvalue, reward, toks, metrics - def loss_wrapper(model, chosen, rejected, chosen_masks, rejected_masks, preference_scores): + def loss_wrapper( + model, chosen, rejected, chosen_masks, rejected_masks, preference_scores + ): return loss( model=model, chosen=chosen, @@ -215,9 +243,9 @@ def loss_wrapper(model, chosen, rejected, chosen_masks, rejected_masks, preferen chosen_masks=chosen_masks, rejected_masks=rejected_masks, preference_scores=preference_scores, - beta=args.beta + beta=args.beta, ) - + loss_value_and_grad = nn.value_and_grad(model, loss_wrapper) # Training loop with progress tracking @@ -227,14 +255,14 @@ def loss_wrapper(model, chosen, rejected, chosen_masks, rejected_masks, preferen 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 + "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), @@ -253,7 +281,7 @@ def loss_wrapper(model, chosen, rejected, chosen_masks, rejected_masks, preferen batch_size=args.batch_size, num_batches=args.val_batches, max_seq_length=args.max_seq_length, - beta=args.beta + beta=args.beta, ) val_time = time.perf_counter() - stop if rank == 0: @@ -269,14 +297,16 @@ def loss_wrapper(model, chosen, rejected, chosen_masks, rejected_masks, preferen ) 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, - }) + 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() @@ -296,15 +326,20 @@ def loss_wrapper(model, chosen, rejected, chosen_masks, rejected_masks, preferen 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()} + 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.metal.get_peak_memory() / 1e9 - + if rank == 0: print( f"Iter {it}: Train loss {train_loss:.3f}, " @@ -320,18 +355,20 @@ def loss_wrapper(model, chosen, rejected, chosen_masks, rejected_masks, preferen ) 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, - }) + 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,)) @@ -354,4 +391,4 @@ def loss_wrapper(model, chosen, rejected, chosen_masks, rejected_masks, preferen 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}.") \ No newline at end of file + print(f"Saved final weights to {args.adapter_file}.") From 74d55e9efcde3a94222bb10dc4c8fb01620eaf2c Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Wed, 19 Mar 2025 13:45:38 +0100 Subject: [PATCH 07/18] making key names customizable --- mlx_lm/examples/lora_config.yaml | 7 ++++++- mlx_lm/tuner/datasets.py | 16 ++++++++++++++-- 2 files changed, 20 insertions(+), 3 deletions(-) diff --git a/mlx_lm/examples/lora_config.yaml b/mlx_lm/examples/lora_config.yaml index bafe2ccce..3ac753468 100644 --- a/mlx_lm/examples/lora_config.yaml +++ b/mlx_lm/examples/lora_config.yaml @@ -89,4 +89,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/tuner/datasets.py b/mlx_lm/tuner/datasets.py index 1023c1087..973dc8d19 100644 --- a/mlx_lm/tuner/datasets.py +++ b/mlx_lm/tuner/datasets.py @@ -239,6 +239,10 @@ 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") @@ -263,8 +267,16 @@ def create_dataset( "https://github.com/ml-explore/mlx-examples/blob/main/llms/mlx_lm/LORA.md#data." ) else: - if "chosen" in sample and "rejected" in sample: - return ORPODataset(data, tokenizer) + 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" From c8922781e0ff505422896af0440505337b20d694 Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Thu, 20 Mar 2025 21:49:13 +0100 Subject: [PATCH 08/18] remove reference model arg --- mlx_lm/lora.py | 1 - 1 file changed, 1 deletion(-) diff --git a/mlx_lm/lora.py b/mlx_lm/lora.py index cb5b2f161..2c3145bb7 100644 --- a/mlx_lm/lora.py +++ b/mlx_lm/lora.py @@ -72,7 +72,6 @@ "mask_prompt": False, # ORPO args "beta": 0.1, - "reference_model_path": None, "reward_scaling": 1.0, } From 866a99f30ae1a49ac7a671c0d4032f9fe4399878 Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Tue, 25 Mar 2025 09:12:05 +0100 Subject: [PATCH 09/18] remove matal in clear cache + type definition --- mlx_lm/tuner/orpo_trainer.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/mlx_lm/tuner/orpo_trainer.py b/mlx_lm/tuner/orpo_trainer.py index 28d463e80..cc4198842 100644 --- a/mlx_lm/tuner/orpo_trainer.py +++ b/mlx_lm/tuner/orpo_trainer.py @@ -23,7 +23,13 @@ class ORPOTrainingArgs(TrainingArgs): def orpo_loss( - model, chosen, rejected, chosen_masks, rejected_masks, preference_scores, beta=0.1 + model: nn.Module, + chosen, + rejected, + chosen_masks, + rejected_masks, + preference_scores, + beta: float = 0.1 ): def get_logps(model, x, mask): inputs = x[:, :-1] @@ -65,6 +71,7 @@ def get_logps(model, x, mask): "chosen_logits_mean": mx.mean(chosen_logits_mean), } + mx.clear_cache() return mx.mean(loss), reward, num_tokens, metrics From 289cff09cb310075ad64e6bd70974da56f7fd964 Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Tue, 25 Mar 2025 09:16:55 +0100 Subject: [PATCH 10/18] nits --- mlx_lm/tuner/orpo_trainer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mlx_lm/tuner/orpo_trainer.py b/mlx_lm/tuner/orpo_trainer.py index cc4198842..19da414ad 100644 --- a/mlx_lm/tuner/orpo_trainer.py +++ b/mlx_lm/tuner/orpo_trainer.py @@ -71,7 +71,7 @@ def get_logps(model, x, mask): "chosen_logits_mean": mx.mean(chosen_logits_mean), } - mx.clear_cache() + mx.metal.clear_cache() return mx.mean(loss), reward, num_tokens, metrics From f04ae963c45c2067afbf46ff210c235ed339f9ad Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Tue, 25 Mar 2025 16:16:30 +0100 Subject: [PATCH 11/18] nits --- mlx_lm/tuner/orpo_trainer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mlx_lm/tuner/orpo_trainer.py b/mlx_lm/tuner/orpo_trainer.py index 19da414ad..cc4198842 100644 --- a/mlx_lm/tuner/orpo_trainer.py +++ b/mlx_lm/tuner/orpo_trainer.py @@ -71,7 +71,7 @@ def get_logps(model, x, mask): "chosen_logits_mean": mx.mean(chosen_logits_mean), } - mx.metal.clear_cache() + mx.clear_cache() return mx.mean(loss), reward, num_tokens, metrics From 037796f046188b7265c80b604d1abca5bc818301 Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Thu, 27 Mar 2025 20:17:22 +0100 Subject: [PATCH 12/18] fix --- mlx_lm/lora.py | 2 ++ mlx_lm/tuner/datasets.py | 6 +++--- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/mlx_lm/lora.py b/mlx_lm/lora.py index f2b25e34f..0e6735276 100644 --- a/mlx_lm/lora.py +++ b/mlx_lm/lora.py @@ -252,6 +252,8 @@ def train_model( adapter_file = adapter_path / "adapters.safetensors" save_config(vars(args), adapter_path / "adapter_config.json") + model.train() + # Initialize the selected optimizer lr = build_schedule(args.lr_schedule) if args.lr_schedule else args.learning_rate diff --git a/mlx_lm/tuner/datasets.py b/mlx_lm/tuner/datasets.py index 4d7924d76..29938d4c4 100644 --- a/mlx_lm/tuner/datasets.py +++ b/mlx_lm/tuner/datasets.py @@ -177,12 +177,12 @@ def process(self, d): tokens = self.tokenizer.apply_chat_template(messages, tools=tools) if self.mask_prompt: messages = messages[:-1] - offset = len(tokenizer.apply_chat_template(messages, tools=tools)) + offset = len(self.tokenizer.apply_chat_template(messages, tools=tools)) return (tokens, offset) else: return tokens - def itemlen(idx: int): + def itemlen(self, idx: int): return len(self._data[idx]) def __getitem__(self, idx: int): @@ -299,7 +299,7 @@ def create_dataset( elif text_feature in sample: if mask_prompt: raise ValueError("Prompt masking not supported for text dataset.") - return Dataset(data, tokenizer, text_key=text_feature) + return TextDataset(data, tokenizer, text_key=text_feature) else: raise ValueError( "Unsupported data format, check the supported formats here:\n" From 3c124e1c32cd8b90c4d79a8439ff86eb5e6d6fbb Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Thu, 8 May 2025 21:33:08 +0200 Subject: [PATCH 13/18] Refactor orpo_loss function to improve log probability calculations and reward estimation; remove unused tokenizer parameter from training functions. --- mlx_lm/lora.py | 2 -- mlx_lm/tuner/orpo_trainer.py | 52 +++++++++++++++++------------------- 2 files changed, 24 insertions(+), 30 deletions(-) diff --git a/mlx_lm/lora.py b/mlx_lm/lora.py index 69bb0a0db..9bc148a7b 100644 --- a/mlx_lm/lora.py +++ b/mlx_lm/lora.py @@ -285,7 +285,6 @@ def train_model( train_orpo( model=model, - tokenizer=tokenizer, optimizer=opt, train_dataset=train_set, val_dataset=valid_set, @@ -307,7 +306,6 @@ def train_model( train( model=model, - tokenizer=tokenizer, args=training_args, optimizer=opt, train_dataset=train_set, diff --git a/mlx_lm/tuner/orpo_trainer.py b/mlx_lm/tuner/orpo_trainer.py index cc4198842..40703ac22 100644 --- a/mlx_lm/tuner/orpo_trainer.py +++ b/mlx_lm/tuner/orpo_trainer.py @@ -29,35 +29,33 @@ def orpo_loss( chosen_masks, rejected_masks, preference_scores, - beta: float = 0.1 + beta: float = 0.1, ): - def get_logps(model, x, mask): - inputs = x[:, :-1] - targets = x[:, 1:] + def get_logps(model, tokens, mask): + inputs = tokens[:, :-1] + targets = tokens[:, 1:] logits = model(inputs) - logp = -nn.losses.cross_entropy(logits, targets, reduction="none") - seq_lengths = mask[:, :-1].sum(-1) - logp_sum = (logp * mask[:, :-1]).sum(-1) / seq_lengths - logits_mean = (logits * mask[:, :-1, None]).sum() / mask[:, :-1].sum() - return logp_sum, logits_mean - - policy_chosen_logps, chosen_logits_mean = get_logps(model, chosen, chosen_masks) - policy_rejected_logps, rejected_logits_mean = get_logps( - model, rejected, rejected_masks - ) + 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 - policy_chosen_logps = policy_chosen_logps * preference_scores + chosen_logps, chosen_logits_mean = get_logps(model, chosen, chosen_masks) + rejected_logps, rejected_logits_mean = get_logps(model, rejected, rejected_masks) - log_odds = (policy_chosen_logps - policy_rejected_logps) - ( - mx.log1p(-mx.exp(policy_chosen_logps)) - - mx.log1p(-mx.exp(policy_rejected_logps)) - ) + # Apply preference weighting + 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 - chosen_reward = beta * policy_chosen_logps - rejected_reward = beta * policy_rejected_logps + # 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() @@ -65,10 +63,10 @@ def get_logps(model, x, mask): metrics = { "accuracies": mx.mean((chosen_reward > rejected_reward).astype(mx.float32)), "margins": mx.mean(chosen_reward - rejected_reward), - "policy_rejected_logps": mx.mean(policy_rejected_logps), - "policy_chosen_logps": mx.mean(policy_chosen_logps), - "rejected_logits_mean": mx.mean(rejected_logits_mean), - "chosen_logits_mean": mx.mean(chosen_logits_mean), + "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() @@ -202,7 +200,6 @@ def evaluate_orpo( def train_orpo( model, - tokenizer, optimizer, train_dataset, val_dataset, @@ -255,7 +252,6 @@ def loss_wrapper( loss_value_and_grad = nn.value_and_grad(model, loss_wrapper) - # Training loop with progress tracking losses = 0 rewards = mx.zeros((2,)) n_tokens = 0 @@ -345,7 +341,7 @@ def loss_wrapper( it_sec = args.steps_per_report / (stop - start) tokens_sec = float(n_tokens) / (stop - start) trained_tokens += n_tokens - peak_mem = mx.metal.get_peak_memory() / 1e9 + peak_mem = mx.get_peak_memory() / 1e9 if rank == 0: print( From dbc7cd4d321d66055a38fd4f842629e2f0dc6c6c Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Thu, 8 May 2025 21:39:50 +0200 Subject: [PATCH 14/18] fix evaluation --- mlx_lm/lora.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/mlx_lm/lora.py b/mlx_lm/lora.py index 9bc148a7b..083187975 100644 --- a/mlx_lm/lora.py +++ b/mlx_lm/lora.py @@ -314,7 +314,7 @@ def train_model( ) -def evaluate_model(args, model: nn.Module, tokenizer: TokenizerWrapper, test_set): +def evaluate_model(args, model: nn.Module, test_set): model.eval() if args.training_mode == "orpo": @@ -338,7 +338,6 @@ def evaluate_model(args, model: nn.Module, tokenizer: TokenizerWrapper, test_set test_loss = evaluate( model=model, dataset=test_set, - tokenizer=tokenizer, batch_size=args.batch_size, num_batches=args.test_batches, max_seq_length=args.max_seq_length, From 5713919983522f07f623a7b8f6c4160ffafc855c Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Sun, 11 May 2025 00:20:44 +0200 Subject: [PATCH 15/18] use CacheDataset() --- mlx_lm/lora.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/mlx_lm/lora.py b/mlx_lm/lora.py index 2c86413b0..01206642c 100644 --- a/mlx_lm/lora.py +++ b/mlx_lm/lora.py @@ -292,8 +292,8 @@ def train_model( train_orpo( model=model, optimizer=opt, - train_dataset=train_set, - val_dataset=valid_set, + train_dataset=CacheDataset(train_set), + val_dataset=CacheDataset(valid_set), args=training_args, training_callback=training_callback, ) @@ -314,8 +314,8 @@ def train_model( model=model, args=training_args, optimizer=opt, - train_dataset=train_set, - val_dataset=valid_set, + train_dataset=CacheDataset(train_set), + val_dataset=CacheDataset(valid_set), training_callback=training_callback, ) From c143a64f2acd75c3a52efabcb3ac6c13533f03f7 Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Sun, 11 May 2025 00:45:47 +0200 Subject: [PATCH 16/18] nits --- mlx_lm/tuner/datasets.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mlx_lm/tuner/datasets.py b/mlx_lm/tuner/datasets.py index 5c137ea37..0cf25e939 100644 --- a/mlx_lm/tuner/datasets.py +++ b/mlx_lm/tuner/datasets.py @@ -321,7 +321,7 @@ def create_dataset( 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." + "https://github.com/ml-explore/mlx-examples/blob/main/llms/mlx_lm/LORA.md#GRPO-Training." ) From a8e264f534ba55c7feaa53a69676f01654036897 Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Thu, 15 May 2025 12:06:54 +0200 Subject: [PATCH 17/18] fix --- mlx_lm/tuner/datasets.py | 3 ++ mlx_lm/tuner/orpo_trainer.py | 74 ++++++++++++++++++++---------------- 2 files changed, 45 insertions(+), 32 deletions(-) diff --git a/mlx_lm/tuner/datasets.py b/mlx_lm/tuner/datasets.py index 0cf25e939..0cf28dad4 100644 --- a/mlx_lm/tuner/datasets.py +++ b/mlx_lm/tuner/datasets.py @@ -117,6 +117,9 @@ def _extract_content(self, data): 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], diff --git a/mlx_lm/tuner/orpo_trainer.py b/mlx_lm/tuner/orpo_trainer.py index 40703ac22..4d84faab2 100644 --- a/mlx_lm/tuner/orpo_trainer.py +++ b/mlx_lm/tuner/orpo_trainer.py @@ -1,14 +1,16 @@ -import time from dataclasses import dataclass, field from pathlib import Path +import time +from mlx.nn.utils import average_gradients +from mlx.utils import tree_flatten 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 .trainer import TrainingArgs, TrainingCallback, grad_checkpoint +from mlx_lm.tuner.callbacks import TrainingCallback + +from .trainer import TrainingArgs, grad_checkpoint @dataclass @@ -22,30 +24,28 @@ class ORPOTrainingArgs(TrainingArgs): ) +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( - model: nn.Module, - chosen, - rejected, + chosen_logps, + chosen_logits_mean, + rejected_logps, + rejected_logits_mean, chosen_masks, rejected_masks, preference_scores, beta: float = 0.1, ): - 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 - - chosen_logps, chosen_logits_mean = get_logps(model, chosen, chosen_masks) - rejected_logps, rejected_logits_mean = get_logps(model, rejected, rejected_masks) - - # Apply preference weighting chosen_logps = chosen_logps * preference_scores # Stable log-odds computation @@ -166,10 +166,15 @@ def evaluate_orpo( ), ): 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( - model=model, - chosen=chosen, - rejected=rejected, + chosen_logps, + chosen_logits_mean, + rejected_logps, + rejected_logits_mean, chosen_masks=chosen_masks, rejected_masks=rejected_masks, preference_scores=preference_scores, @@ -223,10 +228,14 @@ def train_orpo( 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( - model, - chosen, - rejected, + chosen_logps, + chosen_logits_mean, + rejected_logps, + rejected_logits_mean, chosen_masks, rejected_masks, preference_scores=preference_scores, @@ -238,12 +247,13 @@ def step(batch): return lvalue, reward, toks, metrics def loss_wrapper( - model, chosen, rejected, chosen_masks, rejected_masks, preference_scores + chosen_logps, chosen_logits_mean, rejected_logps, rejected_logits_mean, chosen_masks, rejected_masks, preference_scores ): return loss( - model=model, - chosen=chosen, - rejected=rejected, + 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, From bea9bd8003e2d738a1d0980ebc2ed4b0b7d9c264 Mon Sep 17 00:00:00 2001 From: Goekdeniz-Guelmez Date: Thu, 15 May 2025 12:08:24 +0200 Subject: [PATCH 18/18] format --- mlx_lm/lora.py | 3 +-- mlx_lm/tuner/datasets.py | 6 ++++-- mlx_lm/tuner/orpo_trainer.py | 22 ++++++++++++++++------ 3 files changed, 21 insertions(+), 10 deletions(-) diff --git a/mlx_lm/lora.py b/mlx_lm/lora.py index 01206642c..0b4126d20 100644 --- a/mlx_lm/lora.py +++ b/mlx_lm/lora.py @@ -11,9 +11,9 @@ import numpy as np import yaml -from .tuner.orpo_trainer import ORPOTrainingArgs, evaluate_orpo, train_orpo 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, @@ -70,7 +70,6 @@ "lora_parameters": {"rank": 8, "dropout": 0.0, "scale": 10.0}, "mask_prompt": False, "wandb": None, - # ORPO args "beta": 0.1, "reward_scaling": 1.0, diff --git a/mlx_lm/tuner/datasets.py b/mlx_lm/tuner/datasets.py index 0cf28dad4..a37a3673b 100644 --- a/mlx_lm/tuner/datasets.py +++ b/mlx_lm/tuner/datasets.py @@ -286,7 +286,9 @@ def create_dataset( 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") + 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") @@ -319,7 +321,7 @@ def create_dataset( prompt_key=prompt_feature, chosen_key=chosen_feature, rejected_key=rejected_feature, - preference_score_key=preference_score_feature + preference_score_key=preference_score_feature, ) else: raise ValueError( diff --git a/mlx_lm/tuner/orpo_trainer.py b/mlx_lm/tuner/orpo_trainer.py index 4d84faab2..378bdd1ad 100644 --- a/mlx_lm/tuner/orpo_trainer.py +++ b/mlx_lm/tuner/orpo_trainer.py @@ -1,12 +1,12 @@ +import time from dataclasses import dataclass, field from pathlib import Path -import time -from mlx.nn.utils import average_gradients -from mlx.utils import tree_flatten 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 @@ -168,7 +168,9 @@ def evaluate_orpo( 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) + rejected_logps, rejected_logits_mean = get_logps( + model, rejected, rejected_masks + ) lvalue, reward, toks, metrics = orpo_loss( chosen_logps, @@ -229,7 +231,9 @@ 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) + rejected_logps, rejected_logits_mean = get_logps( + model, rejected, rejected_masks + ) (lvalue, reward, toks, metrics), grad = loss_value_and_grad( chosen_logps, @@ -247,7 +251,13 @@ def step(batch): return lvalue, reward, toks, metrics def loss_wrapper( - chosen_logps, chosen_logits_mean, rejected_logps, rejected_logits_mean, chosen_masks, rejected_masks, preference_scores + chosen_logps, + chosen_logits_mean, + rejected_logps, + rejected_logits_mean, + chosen_masks, + rejected_masks, + preference_scores, ): return loss( chosen_logps=chosen_logps,