From 4eec16b3473ef7d8141bf81861b2ade6f78e2eb9 Mon Sep 17 00:00:00 2001 From: CJstate <1507965754@qq.com> Date: Wed, 9 Sep 2026 15:46:58 +0800 Subject: [PATCH] fix: validate ngram_size in get_repetition_penalty_reward Zero or negative ngram_size values were accepted at construction time but caused a ZeroDivisionError (or silently empty n-grams) during reward computation. Add construction-time validation consistent with the existing max_penalty check. Fixes #7015. --- tests/test_rewards.py | 5 +++++ trl/rewards/other_rewards.py | 2 ++ 2 files changed, 7 insertions(+) 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)