Skip to content
Draft
Show file tree
Hide file tree
Changes from 11 commits
Commits
Show all changes
35 commits
Select commit Hold shift + click to select a range
18dace7
change var name
jacklanchantin May 2, 2025
c6fc816
check if self._step_nr exists when syncing (#1160)
jacklanchantin May 2, 2025
7df6271
force_sync bug fix
jacklanchantin May 6, 2025
c27b822
2 verifiers
jacklanchantin May 9, 2025
e22fbed
brute force vllm reward synchronization
jacklanchantin May 9, 2025
fcb6755
working ckpt
jacklanchantin May 9, 2025
76c6775
reward mapper
jacklanchantin May 9, 2025
0c9d658
superclass
jacklanchantin May 9, 2025
cc0d6b1
change class
jacklanchantin May 9, 2025
da29581
Merge branch 'online_training' of github.com:facebookresearch/fairseq…
jacklanchantin May 9, 2025
a63bbd4
cleanup
jacklanchantin May 9, 2025
5db1f99
cleanup
jacklanchantin May 9, 2025
d135ee5
reorder
jacklanchantin May 9, 2025
dc72481
log
jacklanchantin May 9, 2025
52563e9
separate val logging
jacklanchantin May 9, 2025
4e425c9
Merge branch 'jacklanchantin/normalized_rewards' of github.com:facebo…
jacklanchantin May 9, 2025
31200d3
log tasks during training
jacklanchantin May 9, 2025
c7328a8
bug
jacklanchantin May 10, 2025
74a4f85
comm sleep
jacklanchantin May 10, 2025
fbf29f0
new ray.get method for reward model
jacklanchantin May 10, 2025
5aad554
mult rewards
jacklanchantin May 12, 2025
754fd82
comment
jacklanchantin May 12, 2025
dabea7a
remove math tokenizer/wrapper
jacklanchantin May 12, 2025
9fe1a89
log
jacklanchantin May 12, 2025
59cf67b
logs
jacklanchantin May 12, 2025
e6debd5
force athene verifier if batch mismatch
jacklanchantin May 12, 2025
1a9ea63
neurips checkpoint
jacklanchantin May 20, 2025
49c31c2
merge
jacklanchantin May 21, 2025
e6d2146
merge
jacklanchantin May 27, 2025
a790b85
merge
jacklanchantin Jun 4, 2025
ee71556
typo
jacklanchantin Jun 4, 2025
924a02e
merge fixes
jacklanchantin Jun 4, 2025
0a1d1d4
Merge branch 'online_training' of github.com:facebookresearch/fairseq…
jacklanchantin Jun 6, 2025
df7325f
Merge branch 'online_training' of github.com:facebookresearch/fairseq…
jacklanchantin Jun 10, 2025
692f79a
cant throw error for mixed tasks yet
jacklanchantin Jun 27, 2025
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 30 additions & 9 deletions src/fairseq2/recipes/lm/_online_finetune/_online_dpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,7 @@
RewardSection,
VLLMOutputReward,
VLLMOutputRewardHandler,
MultiVerifier,
)
from fairseq2.recipes.lm._preference_finetune._common import (
POCriterionSection,
Expand Down Expand Up @@ -89,7 +90,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__(
Expand All @@ -116,6 +117,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"
Expand Down Expand Up @@ -225,9 +227,17 @@ def __call__(self, prompt_batch: PromptBatch) -> tuple[Tensor, int]:
log_rollouts(prompt_batch, rollouts, "Train")

batch: PreferenceBatch
# Class MultiReward (superclass of current reward)
# subrewards AtheneReward, MathReward
# this super class should be working as is for 2 validation sets, since validation sets will also have task type in the jsonl, so hypothetically correct reward will also be used for logging
# if self._gangs.root.rank == 0:
# breakpoint()
# self._gangs.root.barrier()

batch, is_bad_batch, reward_output = self._reward.prepare_preference_batch(
prompt_batch, rollouts
) # loss_zeroer is used when entire batch has no valid prefrence pair

if is_bad_batch:
loss_zeroer = 0.0
else:
Expand Down Expand Up @@ -562,15 +572,26 @@ 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:
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_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,
Expand Down
88 changes: 84 additions & 4 deletions src/fairseq2/recipes/lm/_online_finetune/_rewards.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -34,6 +34,7 @@
)
from fairseq2.recipes.model import Model
from fairseq2.recipes.trainer import TrainUnit
from fairseq2.logging import log


@dataclass(kw_only=True)
Expand All @@ -45,7 +46,7 @@ class RewardModelConfig:

@dataclass(kw_only=True)
class RewardSection:
name: str = "dummy"
name: str | List[str] = "dummy"
config: RewardModelConfig = field(default_factory=lambda: RewardModelConfig())


Expand Down Expand Up @@ -179,6 +180,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
Expand All @@ -193,7 +195,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 (
Expand All @@ -209,6 +211,8 @@ 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
self.tokenizer = AutoTokenizer.from_pretrained("Nexusflow/Athene-RM-8B")

label_normalizer = NormalizationConfig(
basic_latex=True,
Expand All @@ -227,6 +231,16 @@ def __init__(self, answer_key, prompt_key, gangs):
precision=6,
)

def wrap_text(self, prompt_text, rollout_text):
wrapped_text = [
{"role": "user", "content": prompt_text},
{"role": "assistant", "content": rollout_text},
]
chat_str = self.tokenizer.apply_chat_template(wrapped_text, tokenize=False)
chat_str += "<|reserved_special_token_1|>"

return chat_str

def verify_answer(self, completion: str, answer: str):
# here we add extra $$ to label so that LatexExtractor works as expected
if not answer.startswith("$"):
Expand All @@ -246,12 +260,13 @@ 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

reference_answers = prompt_batch.meta_info.get(self.answer_key)

for i, i_batch_request_output in enumerate(vllm_outputs):
rollouts_text = []
rollouts_tokens = []
Expand All @@ -264,16 +279,29 @@ def process_rollouts(
rollout_output.text, i_reference_answer
)
rollouts_rewards.append(predicted_reward)
vllm_input = self.wrap_text("dummy prompt", "dummy rollout")
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:
log.info("MathVerifyVerifir generate_rewards()")
dummy_rewards = generate_rewards(

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

if i check for if self.vllm_math_reward_model is not None which it will be None for non combination runs, then i get an nccl error. not sure what's happening there

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
) -> PreferenceBatch:

log.info("MathVerifyVerifier prepare_preference_batch()")
reward_output = self.process_rollouts(rollouts, prompt_batch)

batch, is_bad_batch = prepare_preference_batch_random_pair(
Expand Down Expand Up @@ -601,6 +629,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 = []
Expand All @@ -625,6 +654,7 @@ def process_rollouts(
batch_text.append(rollouts_text)
batch_tokens.append(rollouts_tokens)

log.info("AtheneVerifier generate_rewards()")
batch_rewards = generate_rewards(
vllm_inputs, dp_gang=self._gangs.dp, vllm_model=self.reward_model
)
Expand All @@ -639,6 +669,7 @@ def prepare_preference_batch(
self, prompt_batch: PromptBatch, rollouts
) -> PreferenceBatch:

log.info("AtheneVerifier prepare_preference_batch()")
reward_output = self.process_rollouts(rollouts, prompt_batch)

chosen_batch = []
Expand Down Expand Up @@ -751,3 +782,52 @@ def prepare_grpo_batch(self, prompt_batch: PromptBatch, rollouts):
)

return grpo_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.
"""
assert (
len(set(prompt_batch.meta_info["reward_model"])) == 1
), "Not all batch prompts have the same 'reward_model' value."
reward_model_type = prompt_batch.meta_info["reward_model"][0]
if reward_model_type is None:
raise ValueError("Reward model type not found in meta_info.")

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

@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)