diff --git a/lightllm/server/core/objs/py_sampling_params.py b/lightllm/server/core/objs/py_sampling_params.py index 0a158d7f7..ba3368d57 100644 --- a/lightllm/server/core/objs/py_sampling_params.py +++ b/lightllm/server/core/objs/py_sampling_params.py @@ -6,7 +6,7 @@ from typing import List, Optional, Union, Tuple from transformers import GenerationConfig from lightllm.server.req_id_generator import MAX_BEST_OF -from .sampling_params import normalize_seed, validate_seed +from .sampling_params import MAX_SEED _SAMPLING_EPS = 1e-5 @@ -93,7 +93,7 @@ def __init__( self.invalid_token_ids = invalid_token_ids self.group_request_id = group_request_id self.suggested_dp_index = suggested_dp_index - self.seed = normalize_seed(seed) + self.seed = self._normalize_and_verify_seed(seed) if self.do_sample is False: self.temperature = 1.0 self.top_p = 1.0 @@ -155,8 +155,6 @@ def verify(self): raise ValueError( f"min_new_tokens must <= max_new_tokens, but got min {self.min_new_tokens}, max {self.max_new_tokens}." ) - validate_seed(self.seed) - if len(self.exponential_decay_length_penalty) != 2: raise ValueError( f"exponential_decay_length_penalty must be a tuple of (int, float), \ @@ -205,6 +203,13 @@ def verify(self): return + @staticmethod + def _normalize_and_verify_seed(seed: Optional[int]) -> int: + seed = -1 if seed is None else seed + if not -1 <= seed <= MAX_SEED: + raise ValueError(f"seed must be -1 (random), or an integer in [0, {MAX_SEED}], got {seed}") + return seed + def _verify_allowed_token_ids(self): if self.allowed_token_ids is not None: if (not isinstance(self.allowed_token_ids, list)) or ( diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index e23f13ec0..8e31c5062 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -23,15 +23,6 @@ MAX_SEED = (1 << 63) - 1 -def normalize_seed(seed: Optional[int]) -> int: - return -1 if seed is None else seed - - -def validate_seed(seed: int) -> None: - if seed < -1 or seed > MAX_SEED: - raise ValueError(f"seed must be -1 (random), or an integer in [0, {MAX_SEED}], got {seed}") - - class StopSequence(ctypes.Structure): _pack_ = 4 _fields_ = [ @@ -354,11 +345,8 @@ def init(self, tokenizer, **kwargs): self.add_special_tokens = kwargs.get("add_special_tokens", True) self.add_spaces_between_special_tokens = kwargs.get("add_spaces_between_special_tokens", True) self.print_eos_token = kwargs.get("print_eos_token", False) - seed = normalize_seed(kwargs.get("seed")) - # ctypes silently wraps out-of-range integers, so validate the Python - # value before assigning it to the c_int64 field. - validate_seed(seed) - self.seed = seed + # ctypes silently wraps overflowing integers assigned to c_int64. + self.seed = self._normalize_and_verify_seed(kwargs.get("seed")) prompt_logprobs = kwargs.get("prompt_logprobs", None) self.prompt_logprobs = -1 if prompt_logprobs is None else int(prompt_logprobs) @@ -462,12 +450,18 @@ def verify(self): raise ValueError(f"prompt_logprobs must be in [-1, {MAX_PROMPT_LOGPROBS}], got {self.prompt_logprobs}") if self.prompt_logprobs >= 0 and not get_env_start_args().enable_prompt_logprobs: raise ValueError("prompt_logprobs requires --enable_prompt_logprobs") - validate_seed(self.seed) self._verify_allowed_token_ids() self._verify_grammar_constraint() return + @staticmethod + def _normalize_and_verify_seed(seed: Optional[int]) -> int: + seed = -1 if seed is None else seed + if not -1 <= seed <= MAX_SEED: + raise ValueError(f"seed must be -1 (random), or an integer in [0, {MAX_SEED}], got {seed}") + return seed + def _verify_grammar_constraint(self): if self.guided_grammar.length != 0: if self.regular_constraint.length != 0: diff --git a/test/test_api/test_seed_validation.py b/test/test_api/test_seed_validation.py new file mode 100644 index 000000000..1168c5ac5 --- /dev/null +++ b/test/test_api/test_seed_validation.py @@ -0,0 +1,51 @@ +import pytest +from pydantic import ValidationError + +from lightllm.server.api_models import ChatCompletionRequest, CompletionRequest, MAX_SEED +from lightllm.server.core.objs.py_sampling_params import SamplingParams as PySamplingParams +from lightllm.server.core.objs.sampling_params import SamplingParams + + +@pytest.mark.parametrize( + ("request_type", "request_data"), + [ + (CompletionRequest, {"model": "test", "prompt": "hello"}), + (ChatCompletionRequest, {"messages": [{"role": "user", "content": "hello"}]}), + ], +) +def test_api_request_seed_range(request_type, request_data): + for seed in (None, -1, 0, MAX_SEED): + assert request_type(seed=seed, **request_data).seed == seed + + for seed in (-2, MAX_SEED + 1): + with pytest.raises(ValidationError): + request_type(seed=seed, **request_data) + + +@pytest.mark.parametrize( + ("seed", "expected_seed"), + [ + (None, -1), + (-1, -1), + (0, 0), + (MAX_SEED, MAX_SEED), + ], +) +def test_sampling_params_normalizes_and_accepts_seed(seed, expected_seed): + sampling_params = SamplingParams() + sampling_params.init(None, seed=seed) + assert sampling_params.seed == expected_seed + + py_sampling_params = PySamplingParams(seed=seed) + py_sampling_params.verify() + assert py_sampling_params.seed == expected_seed + + +@pytest.mark.parametrize("seed", [-2, MAX_SEED + 1]) +def test_sampling_params_rejects_out_of_range_seed(seed): + sampling_params = SamplingParams() + with pytest.raises(ValueError, match="seed must be -1"): + sampling_params.init(None, seed=seed) + + with pytest.raises(ValueError, match="seed must be -1"): + PySamplingParams(seed=seed)