diff --git a/src/fairseq2/cli/_setup.py b/src/fairseq2/cli/_setup.py index f0dd766ad..6fd8a5ba1 100644 --- a/src/fairseq2/cli/_setup.py +++ b/src/fairseq2/cli/_setup.py @@ -62,15 +62,15 @@ CausalLMLossEvalConfig, CausalLMTrainConfig, InstructionFinetuneConfig, - POFinetuneConfig, OnlineFinetuneConfig, + POFinetuneConfig, TextGenerateConfig, load_clm_loss_evaluator, load_clm_trainer, load_instruction_finetuner, + load_online_finetuner, load_po_finetuner, load_text_generator, - load_online_finetuner, ) from fairseq2.recipes.mt import ( MTEvalConfig, diff --git a/src/fairseq2/data/text/tokenizers/huggingface_tokenizer.py b/src/fairseq2/data/text/tokenizers/huggingface_tokenizer.py index e83caee9b..482724d1e 100644 --- a/src/fairseq2/data/text/tokenizers/huggingface_tokenizer.py +++ b/src/fairseq2/data/text/tokenizers/huggingface_tokenizer.py @@ -12,6 +12,7 @@ import torch from torch import Tensor +from transformers import AutoTokenizer from typing_extensions import override from fairseq2.data import VocabularyInfo @@ -20,7 +21,6 @@ TextTokenEncoder, ) from fairseq2.typing import Device -from transformers import AutoTokenizer @final diff --git a/src/fairseq2/datasets/prompt.py b/src/fairseq2/datasets/prompt.py index 669fb294c..b12f7de51 100644 --- a/src/fairseq2/datasets/prompt.py +++ b/src/fairseq2/datasets/prompt.py @@ -30,9 +30,9 @@ DataPipelineReader, DataReadOptions, DatasetHubAccessor, + DatasetLoadError, LengthBatching, StaticBatching, - DatasetLoadError, UnknownSplitError, ) from fairseq2.datasets.utils._manifest import _load_files_and_weights @@ -74,11 +74,11 @@ class PromptBatch: def batch_size(self) -> int: """The size of the batch dimension.""" return len(self.prompts) - + @property def prompt_lengths(self) -> list[int]: return [len(p) for p in self.prompts] - + @override def to(self, device: Device, *, non_blocking: bool = False) -> None: # no device moving since we only carry tokens prompts here diff --git a/src/fairseq2/models/llama/_hg.py b/src/fairseq2/models/llama/_hg.py index ce70fc9bd..51ebc501e 100644 --- a/src/fairseq2/models/llama/_hg.py +++ b/src/fairseq2/models/llama/_hg.py @@ -9,10 +9,14 @@ from pathlib import Path from typing import cast -from torch import Tensor import torch +from torch import Tensor -from fairseq2.models.utils.checkpoint import convert_checkpoint, create_reverse_key_map, get_converted_key +from fairseq2.models.utils.checkpoint import ( + convert_checkpoint, + create_reverse_key_map, + get_converted_key, +) from fairseq2.models.utils.hg import save_hg_checkpoint # isort: split @@ -119,8 +123,8 @@ def _convert_to_hg_config(config: LLaMAConfig) -> dict[str, object]: } -def _convert_parameter(name: str, - parameter: torch.nn.Parameter, config: LLaMAConfig +def _convert_parameter( + name: str, parameter: torch.nn.Parameter, config: LLaMAConfig ) -> dict[str, object]: head_dim = config.model_dim // config.num_attn_heads @@ -159,4 +163,4 @@ def permute_rotary(w: Tensor, num_heads: int) -> Tensor: converted_name = get_converted_key(name, key_map) - return converted_name, parameter \ No newline at end of file + return converted_name, parameter diff --git a/src/fairseq2/models/qwen/_hg.py b/src/fairseq2/models/qwen/_hg.py index 7a70d15c8..a693c5b1b 100644 --- a/src/fairseq2/models/qwen/_hg.py +++ b/src/fairseq2/models/qwen/_hg.py @@ -10,7 +10,11 @@ import torch -from fairseq2.models.utils.checkpoint import convert_checkpoint, create_reverse_key_map, get_converted_key +from fairseq2.models.utils.checkpoint import ( + convert_checkpoint, + create_reverse_key_map, + get_converted_key, +) from fairseq2.models.utils.hg import save_hg_checkpoint # isort: split @@ -61,8 +65,8 @@ def _convert_to_hg_checkpoint( return hg_checkpoint -def _convert_parameter(name: str, - parameter: torch.nn.Parameter, config: QwenConfig +def _convert_parameter( + name: str, parameter: torch.nn.Parameter, config: QwenConfig ) -> dict[str, object]: key_map = { @@ -86,4 +90,4 @@ def _convert_parameter(name: str, converted_name = get_converted_key(name, key_map) - return converted_name, parameter \ No newline at end of file + return converted_name, parameter diff --git a/src/fairseq2/models/utils/checkpoint.py b/src/fairseq2/models/utils/checkpoint.py index e2982a2ae..c485a4c7e 100644 --- a/src/fairseq2/models/utils/checkpoint.py +++ b/src/fairseq2/models/utils/checkpoint.py @@ -94,6 +94,7 @@ def load_checkpoint( if errors: raise ValueError(" ".join(errors)) + def get_converted_key(key: str, key_map: Mapping[str, str]) -> str: for pattern, replacement in key_map.items(): if (converted_key := re.sub(pattern, replacement, key)) != key: @@ -101,6 +102,7 @@ def get_converted_key(key: str, key_map: Mapping[str, str]) -> str: return key + def convert_checkpoint( checkpoint: dict[str, object], key_map: Mapping[str, str] ) -> dict[str, object]: diff --git a/src/fairseq2/recipes/lm/__init__.py b/src/fairseq2/recipes/lm/__init__.py index dcdea9892..907c8a312 100644 --- a/src/fairseq2/recipes/lm/__init__.py +++ b/src/fairseq2/recipes/lm/__init__.py @@ -39,6 +39,82 @@ from fairseq2.recipes.lm._loss_eval import ( register_clm_loss_eval_configs as register_clm_loss_eval_configs, ) +from fairseq2.recipes.lm._online_finetune._generative_judge import ( + GeneralVerifierExtractorHandler as GeneralVerifierExtractorHandler, +) +from fairseq2.recipes.lm._online_finetune._generative_judge import ( + J1PairwiseScoreExtractor as J1PairwiseScoreExtractor, +) +from fairseq2.recipes.lm._online_finetune._generative_judge import ( + J1PairwiseScoreExtractorHandler as J1PairwiseScoreExtractorHandler, +) +from fairseq2.recipes.lm._online_finetune._generative_judge import ( + J1PointwiseExtractor as J1PointwiseExtractor, +) +from fairseq2.recipes.lm._online_finetune._generative_judge import ( + J1PointwiseExtractorHandler as J1PointwiseExtractorHandler, +) +from fairseq2.recipes.lm._online_finetune._generative_judge import ( + JudgmentExtractorHandler as JudgmentExtractorHandler, +) + +# from fairseq2.recipes.lm._online_finetune._group_dpo import ( +# GroupDpoFinetuneUnitHandler as GroupDpoFinetuneUnitHandler, +# ) +from fairseq2.recipes.lm._online_finetune._grpo import ( + GrpoFinetuneUnitHandler as GrpoFinetuneUnitHandler, +) +from fairseq2.recipes.lm._online_finetune._online_dpo import ( + OnlineDpoFinetuneUnitHandler as OnlineDpoFinetuneUnitHandler, +) +from fairseq2.recipes.lm._online_finetune._recipe import ( + OnlineFinetuneConfig, +) +from fairseq2.recipes.lm._online_finetune._recipe import ( + OnlineFinetuneDatasetSection as OnlineFinetuneDatasetSection, +) +from fairseq2.recipes.lm._online_finetune._recipe import ( + OnlineFinetuneUnitHandler, + load_online_finetuner, + register_online_finetune_configs, +) +from fairseq2.recipes.lm._online_finetune._remote_model import ( + NoEnvAtheneRewardPipeline as NoEnvAtheneRewardPipeline, +) +from fairseq2.recipes.lm._online_finetune._remote_model import ( + NoEnvGeneralVerifierPipeline as NoEnvGeneralVerifierPipeline, +) +from fairseq2.recipes.lm._online_finetune._remote_model import ( + RemoteModelHandler as RemoteModelHandler, +) +from fairseq2.recipes.lm._online_finetune._rewards import ( + AtheneVerifier as AtheneVerifier, +) +from fairseq2.recipes.lm._online_finetune._rewards import ( + AtheneVerifierHandler as AtheneVerifierHandler, +) +from fairseq2.recipes.lm._online_finetune._rewards import ( + GenerativePairwiseVerifier as GenerativePairwiseVerifier, +) +from fairseq2.recipes.lm._online_finetune._rewards import ( + GenerativePairwiseVerifierHandler as GenerativePairwiseVerifierHandler, +) +from fairseq2.recipes.lm._online_finetune._rewards import ( + GenerativePointwiseVerifier as GenerativePointwiseVerifier, +) +from fairseq2.recipes.lm._online_finetune._rewards import ( + GenerativePointwiseVerifierHandler as GenerativePointwiseVerifierHandler, +) +from fairseq2.recipes.lm._online_finetune._rewards import GSM8kVerifier as GSM8kVerifier +from fairseq2.recipes.lm._online_finetune._rewards import ( + GSM8kVerifierHandler as GSM8kVerifierHandler, +) +from fairseq2.recipes.lm._online_finetune._rewards import ( + MathVerifyHandler as MathVerifyHandler, +) +from fairseq2.recipes.lm._online_finetune._rewards import ( + VLLMOutputRewardHandler as VLLMOutputRewardHandler, +) from fairseq2.recipes.lm._preference_finetune._config import ( POCriterionSection as POCriterionSection, ) @@ -119,72 +195,6 @@ from fairseq2.recipes.lm._text_generate import ( register_text_generate_configs as register_text_generate_configs, ) -from fairseq2.recipes.lm._online_finetune._recipe import ( - register_online_finetune_configs, - OnlineFinetuneUnitHandler, - load_online_finetuner, - OnlineFinetuneConfig, -) - -from fairseq2.recipes.lm._online_finetune._recipe import ( - OnlineFinetuneDatasetSection as OnlineFinetuneDatasetSection, -) -from fairseq2.recipes.lm._online_finetune._online_dpo import ( - OnlineDpoFinetuneUnitHandler as OnlineDpoFinetuneUnitHandler, -) - -# from fairseq2.recipes.lm._online_finetune._group_dpo import ( -# GroupDpoFinetuneUnitHandler as GroupDpoFinetuneUnitHandler, -# ) -from fairseq2.recipes.lm._online_finetune._grpo import ( - GrpoFinetuneUnitHandler as GrpoFinetuneUnitHandler, -) -from fairseq2.recipes.lm._online_finetune._rewards import ( - VLLMOutputRewardHandler as VLLMOutputRewardHandler, -) -from fairseq2.recipes.lm._online_finetune._rewards import GSM8kVerifier as GSM8kVerifier -from fairseq2.recipes.lm._online_finetune._rewards import ( - GSM8kVerifierHandler as GSM8kVerifierHandler, -) - -from fairseq2.recipes.lm._online_finetune._rewards import ( - AtheneVerifier as AtheneVerifier, -) -from fairseq2.recipes.lm._online_finetune._rewards import ( - AtheneVerifierHandler as AtheneVerifierHandler, -) - -from fairseq2.recipes.lm._online_finetune._rewards import ( - MathVerifyHandler as MathVerifyHandler, -) - -from fairseq2.recipes.lm._online_finetune._rewards import ( - GenerativePointwiseVerifier as GenerativePointwiseVerifier, -) -from fairseq2.recipes.lm._online_finetune._rewards import ( - GenerativePointwiseVerifierHandler as GenerativePointwiseVerifierHandler, -) - -from fairseq2.recipes.lm._online_finetune._rewards import ( - GenerativePairwiseVerifier as GenerativePairwiseVerifier, -) -from fairseq2.recipes.lm._online_finetune._rewards import ( - GenerativePairwiseVerifierHandler as GenerativePairwiseVerifierHandler, -) - -from fairseq2.recipes.lm._online_finetune._remote_model import ( - RemoteModelHandler as RemoteModelHandler, -) - -from fairseq2.recipes.lm._online_finetune._remote_model import ( - NoEnvAtheneRewardPipeline as NoEnvAtheneRewardPipeline, -) - -from fairseq2.recipes.lm._online_finetune._remote_model import ( - NoEnvGeneralVerifierPipeline as NoEnvGeneralVerifierPipeline, -) - - from fairseq2.recipes.lm._train import CausalLMTrainConfig as CausalLMTrainConfig from fairseq2.recipes.lm._train import CausalLMTrainUnit as CausalLMTrainUnit from fairseq2.recipes.lm._train import TextDatasetSection as TextDatasetSection @@ -192,27 +202,3 @@ from fairseq2.recipes.lm._train import ( register_clm_train_configs as register_clm_train_configs, ) - -from fairseq2.recipes.lm._online_finetune._generative_judge import ( - JudgmentExtractorHandler as JudgmentExtractorHandler, -) - -from fairseq2.recipes.lm._online_finetune._generative_judge import ( - J1PointwiseExtractor as J1PointwiseExtractor, -) - -from fairseq2.recipes.lm._online_finetune._generative_judge import ( - J1PointwiseExtractorHandler as J1PointwiseExtractorHandler, -) - -from fairseq2.recipes.lm._online_finetune._generative_judge import ( - J1PairwiseScoreExtractor as J1PairwiseScoreExtractor, -) - -from fairseq2.recipes.lm._online_finetune._generative_judge import ( - J1PairwiseScoreExtractorHandler as J1PairwiseScoreExtractorHandler, -) - -from fairseq2.recipes.lm._online_finetune._generative_judge import ( - GeneralVerifierExtractorHandler as GeneralVerifierExtractorHandler, -) diff --git a/src/fairseq2/recipes/lm/_online_finetune/_common.py b/src/fairseq2/recipes/lm/_online_finetune/_common.py index fbefa8a1d..c5ed7aa70 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_common.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_common.py @@ -17,24 +17,22 @@ from torch import Tensor from vllm import RequestOutput - from fairseq2.data import ( CollateOptionsOverride, Collater, SequenceData, ) -from fairseq2.datasets.preference import PreferenceBatch -from fairseq2.datasets.prompt import PromptBatch -from fairseq2.gang import Gang, Gangs from fairseq2.datasets import ( SequenceBatch, ) +from fairseq2.datasets.preference import PreferenceBatch +from fairseq2.datasets.prompt import PromptBatch +from fairseq2.gang import Gang, Gangs +from fairseq2.logging import log +from fairseq2.metrics import Mean, MetricBag, Sum from fairseq2.nn._batch_layout import BatchLayout from fairseq2.nn.utils.padding import pad_seqs - -from fairseq2.logging import log from fairseq2.recipes.lm._online_finetune._remote_model import RemoteVllmModel -from fairseq2.metrics import Mean, Sum, MetricBag @dataclass(kw_only=True) @@ -595,3 +593,18 @@ def compute_reference_logps( ).seqs return ref_logps + + +def get_parameter_converter(model_config): + + from fairseq2.models.llama import LLaMAConfig + from fairseq2.models.qwen import QwenConfig + + if isinstance(model_config, QwenConfig): + from fairseq2.models.qwen._hg import _convert_parameter + elif isinstance(model_config, LLaMAConfig): + from fairseq2.models.llama._hg import _convert_parameter + else: + raise RuntimeError(f"{model_config} not supported in online recipe") + + return _convert_parameter diff --git a/src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py b/src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py index 492f1c70c..f261e7d4e 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py @@ -72,11 +72,13 @@ """ +import re from abc import ABC, abstractmethod +from typing import Any + from typing_extensions import override + from fairseq2.logging import log -from typing import Any -import re class JudgmentExtractorHandler(ABC): @@ -182,8 +184,8 @@ def __init__(self): try: from math_verify import parse from math_verify.parser import ( - LatexExtractionConfig, ExprExtractionConfig, + LatexExtractionConfig, NormalizationConfig, ) except ImportError: diff --git a/src/fairseq2/recipes/lm/_online_finetune/_grpo.py b/src/fairseq2/recipes/lm/_online_finetune/_grpo.py index 2b19e9e7c..58cd226a7 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_grpo.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_grpo.py @@ -8,50 +8,50 @@ from copy import copy from dataclasses import dataclass, field -from typing import Dict, Final, List, cast, final, Any, Union +from typing import Any, Dict, Final, List, Union, cast, final import torch from torch import Tensor from torch.nn import Module -from fairseq2.metrics import MetricBag from typing_extensions import override from vllm.distributed.device_communicators.pynccl import PyNcclCommunicator -from fairseq2.nn._batch_layout import BatchLayout from fairseq2.context import RuntimeContext +from fairseq2.datasets import ( + SequenceBatch, +) from fairseq2.datasets.preference import PreferenceBatch from fairseq2.datasets.prompt import PromptBatch from fairseq2.gang import Gang, Gangs from fairseq2.logging import log -from fairseq2.datasets import ( - SequenceBatch, -) +from fairseq2.metrics import MetricBag +from fairseq2.nn._batch_layout import BatchLayout from fairseq2.nn.data_parallel._fsdp import ( fsdp_summon_full_parameters as fsdp_summon_full_parameters, ) - +from fairseq2.recipes import Model, TrainUnit from fairseq2.recipes.lm._online_finetune._common import ( + StatefulRolloutBag, VllmSyncSection, + collate_with_target_mask, + compute_reference_logps, compute_token_level_entropy, - log_rollouts, - get_rollout_lengths, generate_rollouts, - StatefulRolloutBag, + get_rollout_lengths, + log_rollouts, update_avg_reward, update_avg_reward_len_norm, update_avg_rollout_length, update_batch_metrics, - update_logit_entropy, - update_grpo_loss, update_grpo_batch_metrics, - compute_reference_logps, - collate_with_target_mask, + update_grpo_loss, + update_logit_entropy, update_std_reward, ) from fairseq2.recipes.lm._online_finetune._handler import OnlineFinetuneUnitHandler from fairseq2.recipes.lm._online_finetune._remote_model import ( - RemoteVllmModel, RemoteHFModel, + RemoteVllmModel, maybe_sync_model, ) from fairseq2.recipes.lm._online_finetune._rewards import ( @@ -59,7 +59,6 @@ VLLMOutputReward, VLLMOutputRewardHandler, ) -from fairseq2.recipes import Model, TrainUnit from fairseq2.utils.structured import structure from fairseq2.utils.validation import validate @@ -525,18 +524,6 @@ def create( context=self._context, ) - # TODO: decide converter as part of the model handler - if "llama" in model.name: - from fairseq2.models.llama._hg import _convert_parameter - - model._convert_parameter = _convert_parameter - elif "qwen" in model.name: - from fairseq2.models.qwen._hg import _convert_parameter - - model._convert_parameter = _convert_parameter - else: - raise RuntimeError - # sync models here before we start training if config.vllm_sync.sync_model_every_n_steps > 0: maybe_sync_model(gangs, model, vllm_model, -1, -1, force_sync=True) diff --git a/src/fairseq2/recipes/lm/_online_finetune/_handler.py b/src/fairseq2/recipes/lm/_online_finetune/_handler.py index d959d0cf7..943528f51 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_handler.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_handler.py @@ -10,8 +10,8 @@ from torch.nn import Module -from fairseq2.gang import Gangs from fairseq2.datasets import SequenceBatch +from fairseq2.gang import Gangs from fairseq2.recipes import Model, TrainUnit diff --git a/src/fairseq2/recipes/lm/_online_finetune/_online_dpo.py b/src/fairseq2/recipes/lm/_online_finetune/_online_dpo.py index 6c48afa46..8f01c2ef7 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_online_dpo.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_online_dpo.py @@ -8,7 +8,7 @@ from copy import copy from dataclasses import dataclass, field -from typing import Dict, Final, List, cast, final, Any, Union +from typing import Any, Dict, Final, List, Union, cast, final import ray import torch @@ -23,51 +23,51 @@ from fairseq2.context import RuntimeContext from fairseq2.data import CollateOptionsOverride, Collater, SequenceData -from fairseq2.datasets.preference import PreferenceBatch from fairseq2.datasets import ( LengthBatching, SequenceBatch, StaticBatching, SyncMode, ) +from fairseq2.datasets.preference import PreferenceBatch from fairseq2.datasets.prompt import PromptBatch from fairseq2.gang import Gang, Gangs from fairseq2.logging import log +from fairseq2.metrics import Mean, MetricBag from fairseq2.models.clm import CausalLM - from fairseq2.nn.data_parallel._fsdp import ( fsdp_summon_full_parameters as fsdp_summon_full_parameters, ) from fairseq2.nn.utils.module import freeze_parameters + +# from fairseq2.recipes.model import Model +from fairseq2.recipes import Model, TrainUnit from fairseq2.recipes.common import setup_reference_model from fairseq2.recipes.common._distributed import broadcast_model from fairseq2.recipes.config import ( ReferenceModelSection, TrainerSection, ) +from fairseq2.recipes.lm._instruction_finetune import update_nll_loss from fairseq2.recipes.lm._online_finetune._common import ( VllmSyncSection, + compute_reference_logps, compute_token_level_entropy, - log_rollouts, - get_rollout_lengths, generate_rollouts, - StatefulRolloutBag, + get_rollout_lengths, + log_rollouts, + update_avg_loss_zeroer, update_avg_reward, update_avg_reward_len_norm, update_avg_rollout_length, update_batch_metrics, - update_logit_entropy, - update_grpo_loss, update_dpo_loss, - update_grpo_batch_metrics, - compute_reference_logps, - collate_with_target_mask, - update_avg_loss_zeroer, + update_logit_entropy, ) from fairseq2.recipes.lm._online_finetune._handler import OnlineFinetuneUnitHandler from fairseq2.recipes.lm._online_finetune._remote_model import ( - RemoteVllmModel, RemoteHFModel, + RemoteVllmModel, maybe_sync_model, ) from fairseq2.recipes.lm._online_finetune._rewards import ( @@ -75,17 +75,11 @@ VLLMOutputReward, VLLMOutputRewardHandler, ) -from fairseq2.recipes.lm._instruction_finetune import update_nll_loss -from fairseq2.metrics import Mean, MetricBag from fairseq2.recipes.lm._preference_finetune._common import ( _gather_lprobs_avg, update_logps_metrics, update_sequence_length_metrics, ) -from fairseq2.recipes.lm._online_finetune._common import compute_token_level_entropy - -# from fairseq2.recipes.model import Model -from fairseq2.recipes import Model, TrainUnit # from fairseq2.typing import DataType from fairseq2.utils.structured import structure @@ -465,18 +459,6 @@ def create( context=self._context, ) - # TODO: decide converter as part of the model handler - if "llama" in model.name: - from fairseq2.models.llama._hg import _convert_parameter - - model._convert_parameter = _convert_parameter - elif "qwen" in model.name: - from fairseq2.models.qwen._hg import _convert_parameter - - model._convert_parameter = _convert_parameter - else: - raise RuntimeError - return OnlineDpoFinetuneUnit( model, reference_model, vllm_model, vllm_actors, reward, gangs, config ) diff --git a/src/fairseq2/recipes/lm/_online_finetune/_recipe.py b/src/fairseq2/recipes/lm/_online_finetune/_recipe.py index 833bbf5ff..210900d09 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_recipe.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_recipe.py @@ -31,10 +31,12 @@ PromptDataset, PromptReadOptions, ) +from fairseq2.device import CPU from fairseq2.logging import log from fairseq2.models.clm import CausalLM from fairseq2.optim import ADAMW_OPTIMIZER, AdamWConfig from fairseq2.optim.lr_scheduler import COSINE_ANNEALING_LR, CosineAnnealingLRConfig +from fairseq2.recipes import Trainer from fairseq2.recipes.common import ( create_checkpoint_manager, create_lr_scheduler, @@ -43,11 +45,12 @@ load_dataset, load_text_tokenizer, register_extra_asset_paths, - setup_training_gangs, - setup_torch, setup_model, + setup_torch, + setup_training_gangs, ) from fairseq2.recipes.config import ( + ActivationCheckpointingSection, CommonSection, DatasetSection, FSDPSection, @@ -58,13 +61,14 @@ RegimeSection, TextTokenizerSection, TrainerSection, - ActivationCheckpointingSection, ) from fairseq2.recipes.lm._online_finetune._common import ( OnlineCriterionSection, - get_ray_actor, + get_parameter_converter, +) +from fairseq2.recipes.lm._online_finetune._grpo import ( + GrpoFinetuneConfig, ) -from fairseq2.recipes.lm._online_finetune._grpo import GrpoFinetuneConfig from fairseq2.recipes.lm._online_finetune._handler import ( OnlineFinetuneUnitHandler, UnknownOnlineFinetuneUnitError, @@ -72,17 +76,11 @@ from fairseq2.recipes.lm._online_finetune._online_dpo import ( # ONLINE_DPO_FINETUNE_UNIT, OnlineDpoFinetuneConfig, ) -from fairseq2.recipes.lm._online_finetune._grpo import ( - GrpoFinetuneConfig, -) - from fairseq2.recipes.lm._online_finetune._remote_model import ( + HFRayActorConfig, RemoteRayModelHandler, VllmRayActorConfig, - HFRayActorConfig, ) -from fairseq2.recipes import Trainer -from fairseq2.device import CPU from fairseq2.utils.rng import manual_seed from fairseq2.utils.structured import structure from fairseq2.utils.validation import validate @@ -273,6 +271,9 @@ def load_online_finetuner( checkpoint_manager, ) + # set parameter converter to sync with vllm + model._convert_parameter = get_parameter_converter(model.config) + optimizer = create_optimizer(context, config.optimizer, model) lr_scheduler = create_lr_scheduler( diff --git a/src/fairseq2/recipes/lm/_online_finetune/_remote_model.py b/src/fairseq2/recipes/lm/_online_finetune/_remote_model.py index 0ef2f5fe5..607a3cbac 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_remote_model.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_remote_model.py @@ -6,23 +6,14 @@ from __future__ import annotations +import os +import re from abc import ABC, abstractmethod from collections import Counter from dataclasses import dataclass, field -import os -from typing_extensions import override -from vllm.engine.arg_utils import PoolerConfig -from fairseq2.gang import Gangs -from fairseq2.nn._batch_layout import BatchLayout -from fairseq2.recipes.lm._online_finetune.third_party.athene import AtheneRewardPipeline -from fairseq2.recipes.lm._online_finetune.third_party.general_verifier import ( - GeneralVerifierPipeline, -) -from fairseq2.utils.structured import StructureError, structure from typing import Any, Dict, Union -from vllm.worker.worker import Worker + import ray -import re import torch from ray.util.placement_group import placement_group from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy @@ -30,9 +21,17 @@ from vllm import LLM, SamplingParams from vllm.engine.arg_utils import PoolerConfig from vllm.utils import get_ip, get_open_port +from vllm.worker.worker import Worker + +from fairseq2.context import RuntimeContext from fairseq2.gang import Gangs from fairseq2.logging import log -from fairseq2.context import RuntimeContext +from fairseq2.nn._batch_layout import BatchLayout +from fairseq2.recipes.lm._online_finetune.third_party.athene import AtheneRewardPipeline +from fairseq2.recipes.lm._online_finetune.third_party.general_verifier import ( + GeneralVerifierPipeline, +) +from fairseq2.utils.structured import StructureError, structure @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 a65a2a765..a4c0bae8c 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/_rewards.py +++ b/src/fairseq2/recipes/lm/_online_finetune/_rewards.py @@ -21,11 +21,11 @@ from fairseq2.datasets.prompt import PromptBatch from fairseq2.gang import Gangs from fairseq2.recipes.lm._online_finetune._common import ( + _mute_output, collate_with_target_mask, generate_rewards, generate_rewards_generative, prepare_preference_batch_random_pair, - _mute_output, ) from fairseq2.recipes.lm._online_finetune._generative_judge import ( JudgmentExtractorHandler, @@ -187,8 +187,8 @@ def __init__(self, answer_key, prompt_key, reward_name, gangs, context): try: from math_verify.metric import math_metric from math_verify.parser import ( - LatexExtractionConfig, ExprExtractionConfig, + LatexExtractionConfig, NormalizationConfig, ) except ImportError: diff --git a/src/fairseq2/recipes/lm/_online_finetune/third_party/athene.py b/src/fairseq2/recipes/lm/_online_finetune/third_party/athene.py index 3423fffcd..81205f241 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/third_party/athene.py +++ b/src/fairseq2/recipes/lm/_online_finetune/third_party/athene.py @@ -1,12 +1,13 @@ +from typing import Any, Dict, List, cast + import torch import torch.nn as nn -from typing import Any, List, cast, Dict from transformers import ( AutoModelForCausalLM, + AutoTokenizer, LlamaModel, LlamaPreTrainedModel, TextClassificationPipeline, - AutoTokenizer, ) diff --git a/src/fairseq2/recipes/lm/_online_finetune/third_party/general_verifier.py b/src/fairseq2/recipes/lm/_online_finetune/third_party/general_verifier.py index 617fcc242..43974943a 100644 --- a/src/fairseq2/recipes/lm/_online_finetune/third_party/general_verifier.py +++ b/src/fairseq2/recipes/lm/_online_finetune/third_party/general_verifier.py @@ -1,5 +1,5 @@ -from transformers import AutoTokenizer, AutoModelForCausalLM import torch +from transformers import AutoModelForCausalLM, AutoTokenizer def convert_decisions_to_binary(string_list): diff --git a/src/fairseq2/recipes/lm/_preference_finetune/_dpo.py b/src/fairseq2/recipes/lm/_preference_finetune/_dpo.py index f29491da8..7467bd47a 100644 --- a/src/fairseq2/recipes/lm/_preference_finetune/_dpo.py +++ b/src/fairseq2/recipes/lm/_preference_finetune/_dpo.py @@ -171,9 +171,7 @@ def __call__( / chosen_target_batch.num_target_elements() ) loss = ( - dpo_loss - + self._nll_scale - * nll_loss + dpo_loss + self._nll_scale * nll_loss ) # normalization applied locally per-rank return loss, chosen_target_batch.batch_size diff --git a/src/fairseq2/setup/_po_finetune_units.py b/src/fairseq2/setup/_po_finetune_units.py index f41850db8..4db4bd48e 100644 --- a/src/fairseq2/setup/_po_finetune_units.py +++ b/src/fairseq2/setup/_po_finetune_units.py @@ -7,30 +7,30 @@ from __future__ import annotations import ray + from fairseq2.context import RuntimeContext -from fairseq2.recipes.lm import ( +from fairseq2.recipes.lm import ( # GroupDpoFinetuneUnitHandler, + AtheneVerifierHandler, CpoFinetuneUnitHandler, DpoFinetuneUnitHandler, - OrpoFinetuneUnitHandler, - POFinetuneUnitHandler, - SimPOFinetuneUnitHandler, - OnlineDpoFinetuneUnitHandler, - # GroupDpoFinetuneUnitHandler, + GeneralVerifierExtractorHandler, + GenerativePairwiseVerifierHandler, + GenerativePointwiseVerifierHandler, GrpoFinetuneUnitHandler, - OnlineFinetuneUnitHandler, GSM8kVerifierHandler, + J1PairwiseScoreExtractorHandler, + J1PointwiseExtractorHandler, + JudgmentExtractorHandler, MathVerifyHandler, - AtheneVerifierHandler, - GenerativePointwiseVerifierHandler, - GenerativePairwiseVerifierHandler, - VLLMOutputRewardHandler, - RemoteModelHandler, NoEnvAtheneRewardPipeline, NoEnvGeneralVerifierPipeline, - JudgmentExtractorHandler, - GeneralVerifierExtractorHandler, - J1PointwiseExtractorHandler, - J1PairwiseScoreExtractorHandler, + OnlineDpoFinetuneUnitHandler, + OnlineFinetuneUnitHandler, + OrpoFinetuneUnitHandler, + POFinetuneUnitHandler, + RemoteModelHandler, + SimPOFinetuneUnitHandler, + VLLMOutputRewardHandler, ) diff --git a/src/fairseq2/setup/_recipes.py b/src/fairseq2/setup/_recipes.py index 971fde609..33bb2cf3b 100644 --- a/src/fairseq2/setup/_recipes.py +++ b/src/fairseq2/setup/_recipes.py @@ -12,8 +12,8 @@ register_clm_loss_eval_configs, register_clm_train_configs, register_instruction_finetune_configs, - register_po_finetune_configs, register_online_finetune_configs, + register_po_finetune_configs, register_text_generate_configs, ) from fairseq2.recipes.mt import ( diff --git a/src/fairseq2/setup/_root.py b/src/fairseq2/setup/_root.py index a44ead996..657e7d20e 100644 --- a/src/fairseq2/setup/_root.py +++ b/src/fairseq2/setup/_root.py @@ -40,8 +40,10 @@ from fairseq2.setup._metrics import _register_metric_descriptors from fairseq2.setup._models import _register_model_families from fairseq2.setup._optim import _register_optimizers -from fairseq2.setup._po_finetune_units import _register_po_finetune_units -from fairseq2.setup._po_finetune_units import _register_online_finetune_units +from fairseq2.setup._po_finetune_units import ( + _register_online_finetune_units, + _register_po_finetune_units, +) from fairseq2.setup._profilers import _register_profilers from fairseq2.setup._recipes import _register_recipes from fairseq2.setup._text_tokenizers import _register_text_tokenizer_families