Skip to content
Draft
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
6 changes: 4 additions & 2 deletions scripts/data/generate_train_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
41 changes: 41 additions & 0 deletions tests/test_generate_train_data.py
Original file line number Diff line number Diff line change
@@ -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))