From 8ca05416ef878b9b23220ad8507ed0d7a677f936 Mon Sep 17 00:00:00 2001 From: Sehlani042 <257166922+Sehlani042@users.noreply.github.com> Date: Sat, 15 Aug 2026 23:30:02 +0800 Subject: [PATCH] Fix repetition penalty forwarding --- scripts/data/generate_train_data.py | 6 +++-- tests/test_generate_train_data.py | 41 +++++++++++++++++++++++++++++ 2 files changed, 45 insertions(+), 2 deletions(-) create mode 100644 tests/test_generate_train_data.py diff --git a/scripts/data/generate_train_data.py b/scripts/data/generate_train_data.py index 2ab156d0..79c3a011 100644 --- a/scripts/data/generate_train_data.py +++ b/scripts/data/generate_train_data.py @@ -44,6 +44,8 @@ def validate_args(args): raise ValueError("top-k must be greater than 0") if args.min_p is not None and not 0.0 <= args.min_p <= 1.0: raise ValueError("min-p must be between 0.0 and 1.0") + if args.repetition_penalty is not None and args.repetition_penalty <= 0: + raise ValueError("repetition-penalty must be greater than 0") if args.max_tokens <= 0: raise ValueError("max-tokens must be greater than 0") if args.concurrency <= 0: @@ -77,14 +79,14 @@ def build_query_kwargs(args, messages, max_tokens=None): } if args.top_p is not None: query_kwargs["top_p"] = args.top_p - if args.repetition_penalty is not None: - query_kwargs["presence_penalty"] = args.repetition_penalty extra_body = {} if args.top_k is not None: extra_body["top_k"] = args.top_k if args.min_p is not None: extra_body["min_p"] = args.min_p + if args.repetition_penalty is not None: + extra_body["repetition_penalty"] = args.repetition_penalty if args.enable_thinking: extra_body.setdefault("chat_template_kwargs", {})["enable_thinking"] = True if args.disable_thinking: diff --git a/tests/test_generate_train_data.py b/tests/test_generate_train_data.py new file mode 100644 index 00000000..36072bf0 --- /dev/null +++ b/tests/test_generate_train_data.py @@ -0,0 +1,41 @@ +from types import SimpleNamespace + +import pytest + +from scripts.data.generate_train_data import build_query_kwargs, validate_args + + +def make_args(**overrides): + values = { + "model": "test-model", + "max_tokens": 128, + "temperature": 0.7, + "top_p": None, + "top_k": None, + "min_p": None, + "repetition_penalty": None, + "enable_thinking": False, + "disable_thinking": False, + "is_gpt_oss": False, + "concurrency": 1, + } + values.update(overrides) + return SimpleNamespace(**values) + + +def test_build_query_kwargs_forwards_repetition_penalty_to_sglang(): + kwargs = build_query_kwargs( + make_args(repetition_penalty=1.1), + [{"role": "user", "content": "hello"}], + ) + + assert kwargs["extra_body"]["repetition_penalty"] == 1.1 + assert "presence_penalty" not in kwargs + + +@pytest.mark.parametrize("repetition_penalty", [0.0, -1.0]) +def test_validate_args_rejects_non_positive_repetition_penalty( + repetition_penalty, +): + with pytest.raises(ValueError, match="repetition-penalty must be greater than 0"): + validate_args(make_args(repetition_penalty=repetition_penalty))