Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
4 changes: 2 additions & 2 deletions src/fairseq2/cli/_setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7,20 +7,20 @@
from __future__ import annotations

from collections.abc import Sequence
from pathlib import Path

Check failure on line 10 in src/fairseq2/data/text/tokenizers/huggingface_tokenizer.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

'pathlib.Path' imported but unused
from typing import final

import torch
from torch import Tensor
from transformers import AutoTokenizer
from typing_extensions import override

from fairseq2.data import VocabularyInfo

Check failure on line 18 in src/fairseq2/data/text/tokenizers/huggingface_tokenizer.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

'fairseq2.data.VocabularyInfo' imported but unused
from fairseq2.data.text.tokenizers import (
TextTokenDecoder,
TextTokenEncoder,
)
from fairseq2.typing import Device
from transformers import AutoTokenizer


@final
Expand Down
6 changes: 3 additions & 3 deletions src/fairseq2/datasets/prompt.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,12 +11,12 @@
from collections.abc import Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Final, List, cast, final

Check failure on line 14 in src/fairseq2/datasets/prompt.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

'typing.cast' imported but unused

import torch

Check failure on line 16 in src/fairseq2/datasets/prompt.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

'torch' imported but unused
from typing_extensions import override

from fairseq2.data import (

Check failure on line 19 in src/fairseq2/datasets/prompt.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

'fairseq2.data.create_bucket_sizes' imported but unused

Check failure on line 19 in src/fairseq2/datasets/prompt.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

'fairseq2.data.SequenceData' imported but unused

Check failure on line 19 in src/fairseq2/datasets/prompt.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

'fairseq2.data.Collater' imported but unused

Check failure on line 19 in src/fairseq2/datasets/prompt.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

'fairseq2.data.CollateOptionsOverride' imported but unused
CollateOptionsOverride,
Collater,
DataPipeline,
Expand All @@ -26,13 +26,13 @@
read_sequence,
)
from fairseq2.data.text.tokenizers import TextTokenizer
from fairseq2.datasets import (

Check failure on line 29 in src/fairseq2/datasets/prompt.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

'fairseq2.datasets.LengthBatching' imported but unused
DataPipelineReader,
DataReadOptions,
DatasetHubAccessor,
DatasetLoadError,
LengthBatching,
StaticBatching,
DatasetLoadError,
UnknownSplitError,
)
from fairseq2.datasets.utils._manifest import _load_files_and_weights
Expand Down Expand Up @@ -74,11 +74,11 @@
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
Expand Down
14 changes: 9 additions & 5 deletions src/fairseq2/models/llama/_hg.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
return converted_name, parameter
12 changes: 8 additions & 4 deletions src/fairseq2/models/qwen/_hg.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 = {
Expand All @@ -86,4 +90,4 @@ def _convert_parameter(name: str,

converted_name = get_converted_key(name, key_map)

return converted_name, parameter
return converted_name, parameter
2 changes: 2 additions & 0 deletions src/fairseq2/models/utils/checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,13 +94,15 @@ 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:
return converted_key

return key


def convert_checkpoint(
checkpoint: dict[str, object], key_map: Mapping[str, str]
) -> dict[str, object]:
Expand Down
166 changes: 76 additions & 90 deletions src/fairseq2/recipes/lm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -119,100 +195,10 @@
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
from fairseq2.recipes.lm._train import load_clm_trainer as load_clm_trainer
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,
)
27 changes: 20 additions & 7 deletions src/fairseq2/recipes/lm/_online_finetune/_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -72,16 +72,18 @@
"""


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):
@abstractmethod
def create(self): ...

Check failure on line 86 in src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Function is missing a return type annotation

@property
@abstractmethod
Expand Down Expand Up @@ -111,7 +113,7 @@
def prompt(self) -> str: ...

@abstractmethod
def format_prompt(self, prompt_text, **kwargs: Any) -> str: ...

Check failure on line 116 in src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Function is missing a type annotation for one or more arguments

"""
Format the prompt text and additional arguments into a string suitable for input to the reward model.
Expand All @@ -124,7 +126,7 @@
"""

@abstractmethod
def extract(self, generation) -> float | str: ...

Check failure on line 129 in src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Function is missing a type annotation for one or more arguments

"""
Extract the final scalar reward score from the model's response.
Expand All @@ -143,7 +145,7 @@
"""

@abstractmethod
def aggregate(self, judgments) -> float | str: ...

Check failure on line 148 in src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Function is missing a type annotation for one or more arguments

"""
Aggregate multiple responses (judgments) from the reward model into a single value.
Expand All @@ -159,31 +161,31 @@


class GeneralVerifierExtractorHandler(JudgmentExtractorHandler):
def __init__(self):

Check failure on line 164 in src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Function is missing a return type annotation
pass

@override
def create(self):

Check failure on line 168 in src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Function is missing a return type annotation
return GeneralVerifierExtractor()

@property
@override
def name(self):

Check failure on line 173 in src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Function is missing a return type annotation
return "general_verifier_extractor"

@property
@override
def config_kls(self):

Check failure on line 178 in src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Function is missing a return type annotation
return None


class GeneralVerifierExtractor(JudgmentExtractor):
def __init__(self):

Check failure on line 183 in src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Function is missing a return type annotation
try:
from math_verify import parse

Check failure on line 185 in src/fairseq2/recipes/lm/_online_finetune/_generative_judge.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Cannot find implementation or library stub for module named "math_verify"
from math_verify.parser import (
LatexExtractionConfig,
ExprExtractionConfig,
LatexExtractionConfig,
NormalizationConfig,
)
except ImportError:
Expand Down
Loading
Loading