diff --git a/tests/test_rewards.py b/tests/test_rewards.py index f5fd582f68d..25beec5b573 100644 --- a/tests/test_rewards.py +++ b/tests/test_rewards.py @@ -161,6 +161,11 @@ def test_positive_max_penalty_raises(self): with pytest.raises(ValueError): get_repetition_penalty_reward(ngram_size=2, max_penalty=0.5) + @pytest.mark.parametrize("ngram_size", [0, -1]) + def test_non_positive_ngram_size_raises(self, ngram_size): + with pytest.raises(ValueError): + get_repetition_penalty_reward(ngram_size=ngram_size, max_penalty=-1.0) + def test_extra_kwargs_are_ignored(self): """Trainers pass prompts/completions/etc. as kwargs; the reward must accept and ignore them.""" reward_fn = get_repetition_penalty_reward(ngram_size=2, max_penalty=-1.0) diff --git a/trl/rewards/other_rewards.py b/trl/rewards/other_rewards.py index 4848d10d183..dd30114593f 100644 --- a/trl/rewards/other_rewards.py +++ b/trl/rewards/other_rewards.py @@ -56,6 +56,8 @@ def get_repetition_penalty_reward(ngram_size: int = 3, max_penalty: float = -1.0 """ if max_penalty > 0: raise ValueError(f"max_penalty {max_penalty} should not be positive") + if ngram_size <= 0: + raise ValueError(f"ngram_size {ngram_size} should be greater than 0") return _RepetitionPenalty(ngram_size, max_penalty)