Skip to content
Open
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
688 changes: 688 additions & 0 deletions cookbook/rl/mopd/cross_token_trainer_on_policy_two_teacher.py

Large diffs are not rendered by default.

348 changes: 348 additions & 0 deletions cookbook/rl/mopd/gold_trainer.py

Large diffs are not rendered by default.

14 changes: 14 additions & 0 deletions src/twinkle/data_format/sampling.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
@dataclass
class SamplingParams:
max_tokens: Optional[int] = None
min_tokens: Optional[int] = None
seed: Optional[int] = None
stop: Union[str, Sequence[str], Sequence[int], None] = None
temperature: float = 1.0
Expand Down Expand Up @@ -59,6 +60,16 @@ def __post_init__(self):
if self.max_tokens < 0:
raise ValueError(f'max_tokens must be >= 1, got {self.max_tokens}')

if self.min_tokens is not None:
if not isinstance(self.min_tokens, int):
raise ValueError(f'min_tokens must be an int or None, got {type(self.min_tokens)}')
if self.min_tokens < 0:
raise ValueError(f'min_tokens must be >= 0, got {self.min_tokens}')
if self.max_tokens is not None and self.min_tokens > self.max_tokens:
raise ValueError(
f'min_tokens ({self.min_tokens}) must be <= max_tokens ({self.max_tokens})'
)

if not isinstance(self.repetition_penalty, (int, float)):
raise ValueError(f'repetition_penalty must be a number, got {type(self.repetition_penalty)}')
if self.repetition_penalty <= 0:
Expand All @@ -79,6 +90,9 @@ def to_vllm(self, **kwargs):
if self.max_tokens is not None:
kwargs['max_tokens'] = self.max_tokens

if self.min_tokens is not None:
kwargs['min_tokens'] = self.min_tokens

if self.seed is not None:
kwargs['seed'] = self.seed

Expand Down
4 changes: 4 additions & 0 deletions src/twinkle/loss/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@
from .infonce import InfonceLoss
from .liger_fused_linear_cross_entropy import LigerFusedLinearCrossEntropyLoss
from .mse import MSELoss
from .gold import GOLDLoss
from .cross_token import CrossTokenLoss
from .value import PPOValueLoss

torch_loss_mapping = {
Expand All @@ -33,4 +35,6 @@
'orpo': ORPOLoss,
# Embedding / contrastive losses
'infonce': InfonceLoss,
'gold': GOLDLoss,
'cross_token': CrossTokenLoss
}
Loading