diff --git a/src/fairseq2/datasets/prompt.py b/src/fairseq2/datasets/prompt.py index 3a59515a7..3f49c05fa 100644 --- a/src/fairseq2/datasets/prompt.py +++ b/src/fairseq2/datasets/prompt.py @@ -250,10 +250,13 @@ def process_examples(example: dict[str, Any]) -> dict[str, Any]: jsonl_content = {} for k in jsonl_keys: if k not in example: - raise KeyError( - f"Required key '{k}' not found in example dictionary." - ) - jsonl_content[k] = example[k] + print(f"Required key '{k}' not found in example dictionary.") + # raise KeyError( + # f"Required key '{k}' not found in example dictionary." + # ) + jsonl_content[k] = None + else: + jsonl_content[k] = example[k] else: jsonl_content = None diff --git a/src/fairseq2/recipes/lm/_online_finetune/_common.py b/src/fairseq2/recipes/lm/_online_finetune/_common.py index d1da34a31..4c957a093 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_common.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_common.py @@ -318,11 +318,9 @@ def find_first_value(lst, value): def generate_rollouts( - prompts: List[List[int]], - dp_gang, - vllm_model, - sampling_params=None, + prompts: List[List[int]], dp_gang, vllm_model, sampling_params=None ): + prompts_to_generate = [None] * dp_gang.size if dp_gang.rank == 0: dp_gang.gather_object(prompts, prompts_to_generate, 0) diff --git a/src/fairseq2/recipes/lm/_online_finetune/_online_dpo.py b/src/fairseq2/recipes/lm/_online_finetune/_online_dpo.py index 43afd10f4..f8d613fd1 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_online_dpo.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_online_dpo.py @@ -63,6 +63,7 @@ RewardSection, VLLMOutputReward, VLLMOutputRewardHandler, + MultiVerifier, ) from fairseq2.recipes.lm._preference_finetune._common import ( POCriterionSection, @@ -90,7 +91,7 @@ class OnlineDpoFinetuneUnit(TrainUnit[SequenceBatch]): _sync_vllm_model_every_n_steps: int _sync_ref_model_every_n_steps: int _display_name: str - _reward: VLLMOutputReward + _reward: dict | VLLMOutputReward _reference_offload: bool def __init__( @@ -117,6 +118,7 @@ def __init__( self._sync_vllm_model_every_n_steps = sync_vllm_model_every_n_steps self._sync_ref_model_every_n_steps = sync_ref_model_every_n_step self._reward = reward + # self._reward = reward self._metric_bag = OnlineDpoFinetuneMetricBag(gangs.dp) self._display_name = "online_dpo" @@ -164,6 +166,23 @@ def validate_reward(self, prompt_batch: PromptBatch) -> tuple[Tensor, int]: policy_sampling_params.__setattr__(k, v) else: policy_sampling_params = None + + if self._loss_config.mix_lengths: + try: + reward_model_type = prompt_batch.meta_info.get("reward_model", [None])[ + 0 + ] + if len(set(prompt_batch.meta_info["reward_model"])) != 1: + reward_model_type = "athene_verifier" + if reward_model_type == "math_verify": + max_tokens = 2048 + elif reward_model_type == "athene_verifier": + max_tokens = 1024 + except: + max_tokens = None + else: + max_tokens = None + rollouts = generate_rollouts( prompt_batch.prompts, dp_gang=self._gangs.dp, @@ -226,16 +245,35 @@ def __call__(self, prompt_batch: PromptBatch) -> tuple[Tensor, int]: self.maybe_sync_models() + if self._loss_config.mix_lengths: + try: + reward_model_type = prompt_batch.meta_info.get("reward_model", [None])[ + 0 + ] + if len(set(prompt_batch.meta_info["reward_model"])) != 1: + reward_model_type = "athene_verifier" + if reward_model_type == "math_verify": + max_tokens = 2048 + elif reward_model_type == "athene_verifier": + max_tokens = 1024 + except: + max_tokens = None + else: + max_tokens = None + rollouts = generate_rollouts( - prompt_batch.prompts, dp_gang=self._gangs.dp, vllm_model=self._vllm_model + prompt_batch.prompts, + dp_gang=self._gangs.dp, + vllm_model=self._vllm_model, ) if self._loss_config.log_rollouts: log_rollouts(prompt_batch, rollouts, "Train") batch: PreferenceBatch batch, is_bad_batch, reward_output = self._reward.prepare_preference_batch( - prompt_batch, rollouts + prompt_batch, rollouts, self._loss_config.reject_longest ) # loss_zeroer is used when entire batch has no valid prefrence pair + if is_bad_batch: loss_zeroer = 0.0 else: @@ -336,6 +374,9 @@ def __call__(self, prompt_batch: PromptBatch) -> tuple[Tensor, int]: avg_reward = torch.tensor(reward_output["rewards"]).float().mean() self._metric_bag.update_avg_reward(avg_reward) + if "reward_model" in prompt_batch.meta_info: + reward_model_type = prompt_batch.meta_info["reward_model"][0] + self._metric_bag.update_avg_task_reward(avg_reward, reward_model_type) if self._loss_config.nll_length_normalization: nll_loss = ( nll_loss @@ -410,6 +451,8 @@ class OnlineDpoFinetuneMetricBag(POFinetuneMetricBag): avg_reward: Mean avg_loss_zeroer: Mean logit_entropy: Mean + avg_math_reward: Mean + avg_wildchat_reward: Mean def __init__(self, gang: Gang) -> None: super().__init__(gang) @@ -431,6 +474,12 @@ def __init__(self, gang: Gang) -> None: self.register_metric( "logit_entropy", Mean(device=gang.device), persistent=False ) + self.register_metric( + "avg_math_reward", Mean(device=gang.device), persistent=False + ) + self.register_metric( + "avg_wildchat_reward", Mean(device=gang.device), persistent=False + ) @torch.inference_mode() def update_logit_entropy(self, logit_entropy: Tensor): @@ -461,6 +510,14 @@ def update_num_dummy_batches(self, batch: PreferenceBatch, num_dummy_batches: in def update_avg_reward(self, avg_reward): self.avg_reward.update(avg_reward, weight=1) + @torch.inference_mode() + def update_avg_task_reward(self, avg_reward, task): + # FIXME don't hardcode tasks + if task == "math_verify": + self.avg_math_reward.update(avg_reward, weight=1) + elif task == "athene_verifier": + self.avg_wildchat_reward.update(avg_reward, weight=1) + @torch.inference_mode() def update_avg_rollout_length(self, avg_rollout_length): self.avg_rollout_length.update(avg_rollout_length, weight=1) @@ -506,6 +563,12 @@ class DpoLossConfig: validation_vllm_sampling_params: Dict[str, Any] = field(default_factory=lambda: {}) """VLLM sampling params for validation. If not set, the same params as training will be used.""" + reject_longest: bool = False + """Longest sequence in the batch will be rejected. This is used to avoid the model to generate long sequences.""" + + mix_lengths: bool = False + """Mix lengths of the sequences in the batch. This is used to avoid the model to generate long sequences.""" + @dataclass(kw_only=True) class OnlineDpoFinetuneConfig: @@ -587,15 +650,27 @@ def create( gangs.root.barrier() vllm_model = vllm_actors[config.ray_policy_actor_name] - - vllm_reward_model = vllm_actors.get(config.vllm_reward_model_name, None) reward_registry = self._context.get_registry(VLLMOutputRewardHandler) - reward_handler = reward_registry.get(config.reward.name) - reward = reward_handler.create( - reward_model=vllm_reward_model, - reward_config=config.reward.config, - gangs=gangs, - ) + vllm_reward_model = vllm_actors.get(config.vllm_reward_model_name, None) + if type(config.reward.name) is list and len(config.reward.name) > 1: + reward_mapper = {} + for reward_name in config.reward.name: + reward_handler = reward_registry.get(reward_name) + reward_model = reward_handler.create( + reward_model=vllm_reward_model, + reward_config=config.reward.config, + gangs=gangs, + ) + reward_mapper[reward_name] = reward_model + reward = MultiVerifier(reward_mapper) + else: + reward_name = config.reward.name + reward_handler = reward_registry.get(reward_name) + reward = reward_handler.create( + reward_model=vllm_reward_model, + reward_config=config.reward.config, + gangs=gangs, + ) return OnlineDpoFinetuneUnit( model, diff --git a/src/fairseq2/recipes/lm/_online_finetune/_remote_model.py b/src/fairseq2/recipes/lm/_online_finetune/_remote_model.py index 4bb27ce13..6db6e47c4 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_remote_model.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_remote_model.py @@ -29,11 +29,10 @@ NoEnvAtheneRewardPipeline, stateless_init_process_group, ) +import time from fairseq2.logging import log from fairseq2.context import RuntimeContext -from fairseq2.recipes.lm._online_finetune._rewards import ( - VLLMOutputRewardHandler, -) +from fairseq2.recipes.lm._online_finetune._rewards import VLLMOutputRewardHandler @dataclass(kw_only=True) diff --git a/src/fairseq2/recipes/lm/_online_finetune/_rewards.py b/src/fairseq2/recipes/lm/_online_finetune/_rewards.py index d132b0179..a6887cc84 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_rewards.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_rewards.py @@ -9,7 +9,7 @@ import re from abc import ABC, abstractmethod from dataclasses import dataclass, field -from typing import Any, List +from typing import Any, List, Dict import torch from transformers import AutoTokenizer @@ -31,6 +31,7 @@ from fairseq2.recipes.lm._online_finetune._generative_prompts import POINTWISE_PROMPT from fairseq2.recipes.model import Model from fairseq2.recipes.trainer import TrainUnit +from fairseq2.logging import log @dataclass(kw_only=True) @@ -42,7 +43,7 @@ class RewardModelConfig: @dataclass(kw_only=True) class RewardSection: - name: str = "dummy" + name: str | List[str] = "dummy" config: RewardModelConfig = field(default_factory=lambda: RewardModelConfig()) @@ -66,7 +67,9 @@ class VLLMOutputReward(ABC): def process_rollouts(self, vllm_outputs: List[RequestOutput]): ... @abstractmethod - def prepare_preference_batch(self, prompt_batch: PromptBatch, rollouts): ... + def prepare_preference_batch( + self, prompt_batch: PromptBatch, rollouts, reject_longest: bool + ): ... class GSM8kVerifierHandler(VLLMOutputRewardHandler): @@ -141,9 +144,8 @@ def process_rollouts( return {"text": batch_text, "tokens": batch_tokens, "rewards": batch_rewards} def prepare_preference_batch( - self, prompt_batch: PromptBatch, rollouts + self, prompt_batch: PromptBatch, rollouts, reject_longest=False ) -> PreferenceBatch: - reward_output = self.process_rollouts(rollouts, prompt_batch) batch, is_bad_batch = prepare_preference_batch_random_pair( @@ -163,6 +165,7 @@ def create(self, reward_model, reward_config, gangs): answer_key=reward_config.answer_key, prompt_key=reward_config.prompt_key, gangs=gangs, + vllm_math_reward_model=reward_model, ) @property @@ -177,7 +180,7 @@ def config_kls(self): class MathVerifyVerifier(VLLMOutputReward): - def __init__(self, answer_key, prompt_key, gangs): + def __init__(self, answer_key, prompt_key, gangs, vllm_math_reward_model): try: from math_verify.metric import math_metric from math_verify.parser import ( @@ -193,6 +196,7 @@ def __init__(self, answer_key, prompt_key, gangs): self._gangs = gangs self.answer_key = answer_key self.prompt_key = prompt_key + self.vllm_math_reward_model = vllm_math_reward_model label_normalizer = NormalizationConfig( basic_latex=True, @@ -230,12 +234,23 @@ def process_rollouts( vllm_outputs: List[RequestOutput], prompt_batch: PromptBatch, ): + log.info("MathVerifyVerifier process_rollouts()") batch_text = [] batch_tokens = [] batch_rewards = [] + vllm_inputs = [] # dummy for vllm syncronization # FIXME + vllm_input = [ + { + "role": "user", + "content": "dummy_prompt", + }, + { + "role": "assistant", + "content": "dummy_rollout", + }, + ] reference_answers = prompt_batch.meta_info.get(self.answer_key) - for i, i_batch_request_output in enumerate(vllm_outputs): rollouts_text = [] rollouts_tokens = [] @@ -248,16 +263,26 @@ def process_rollouts( rollout_output.text, i_reference_answer ) rollouts_rewards.append(predicted_reward) + vllm_inputs.append(vllm_input) + batch_text.append(rollouts_text) batch_tokens.append(rollouts_tokens) batch_rewards.append(rollouts_rewards) + # dummy vllm call for syncronization # FIXME + # if self.vllm_math_reward_model is not None: + dummy_rewards = generate_rewards( + vllm_inputs, + dp_gang=self._gangs.dp, + vllm_model=self.vllm_math_reward_model, + ) + return {"text": batch_text, "tokens": batch_tokens, "rewards": batch_rewards} def prepare_preference_batch( - self, prompt_batch: PromptBatch, rollouts + self, prompt_batch: PromptBatch, rollouts, reject_longest=False ) -> PreferenceBatch: - + log.info("MathVerifyVerifier prepare_preference_batch()") reward_output = self.process_rollouts(rollouts, prompt_batch) batch, is_bad_batch = prepare_preference_batch_random_pair( @@ -331,6 +356,7 @@ def wrap_text(self, prompt_text, rollout_text): def process_rollouts( self, vllm_outputs: List[RequestOutput], prompt_batch: PromptBatch ): + log.info("AtheneVerifier process_rollouts()") vllm_inputs = [] batch_text = [] batch_tokens = [] @@ -366,9 +392,10 @@ def process_rollouts( return {"text": batch_text, "tokens": batch_tokens, "rewards": batch_rewards} def prepare_preference_batch( - self, prompt_batch: PromptBatch, rollouts + self, prompt_batch: PromptBatch, rollouts, reject_longest=False ) -> PreferenceBatch: + log.info("AtheneVerifier prepare_preference_batch()") reward_output = self.process_rollouts(rollouts, prompt_batch) chosen_batch = [] @@ -382,7 +409,22 @@ def prepare_preference_batch( ): chosen_rollout_position = i_batch_rewards.index(max(i_batch_rewards)) - rejected_rollout_position = i_batch_rewards.index(min(i_batch_rewards)) + if reject_longest: + max_length = -1 + rejected_rollout_position = i_batch_rewards.index( + min(i_batch_rewards) + ) # fallback + for idx, tokens in enumerate(i_batch_tokens): + # Only consider rollouts that are *not* the chosen one + if idx != chosen_rollout_position: + current_length = len(tokens) + if current_length > max_length: + max_length = current_length # Update the maximum length + rejected_rollout_position = ( + idx # Update the rejected position + ) + else: + rejected_rollout_position = i_batch_rewards.index(min(i_batch_rewards)) if chosen_rollout_position == rejected_rollout_position: # cant form preference pair when we dont have such rollouts @@ -444,6 +486,62 @@ def prepare_preference_batch( return batch, is_bad_batch, reward_output +class MultiVerifier(VLLMOutputReward): + """ + A reward model verifier that processes rollouts using multiple reward models. + This class evaluates rollouts generated by vLLM by wrapping the prompt and rollout text into a specific format and passing it through the appropriate reward model. + This class is designed to handle multiple reward models and their configurations. + Note: this relies on modified Athene-RM-8B code to ensure compatibility with vLLM. + """ + + def __init__(self, rewards_map: Dict[VLLMOutputReward]): + self.rewards_map = rewards_map + + def get_reward_model_type(self, prompt_batch: PromptBatch) -> str: + """ + Get the reward model type from the prompt batch meta info. + """ + reward_model_type = prompt_batch.meta_info.get("reward_model", [None])[0] + if reward_model_type is None: + raise ValueError("Reward model type not found in meta_info.") + + # assert ( + # len(set(prompt_batch.meta_info.get("reward_model", 0))) == 1 + # ), "Not all batch prompts have the same 'reward_model' value." + if len(set(prompt_batch.meta_info["reward_model"])) != 1: + # FIXME temp solution if the whole batch is not the same reward_model type + log.warning( + "Not all batch prompts have the same 'reward_model' value. Using athene_verifier." + ) + reward_model_type = "athene_verifier" + + log.info(f"Reward model type: {reward_model_type}") + return reward_model_type + + @override + def process_rollouts( + self, vllm_outputs: List[RequestOutput], prompt_batch: PromptBatch + ): + reward_model_type = self.get_reward_model_type(prompt_batch) + reward = self.rewards_map[reward_model_type] + return reward.process_rollouts(vllm_outputs, prompt_batch) + + @override + def prepare_preference_batch( + self, prompt_batch: PromptBatch, rollouts, reject_longest=False + ) -> PreferenceBatch: + + reward_model_type = self.get_reward_model_type(prompt_batch) + reward = self.rewards_map[reward_model_type] + return reward.prepare_preference_batch(prompt_batch, rollouts, reject_longest) + + @override + def prepare_grpo_batch(self, prompt_batch: PromptBatch, rollouts): + reward_model_type = self.get_reward_model_type(prompt_batch) + reward = self.rewards_map[reward_model_type] + return reward.prepare_grpo_batch(prompt_batch, rollouts) + + class GenerativePointwiseVerifierHandler(VLLMOutputRewardHandler): def __init__(self): pass