Skip to content
Merged
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
13 changes: 9 additions & 4 deletions lightllm/server/core/objs/py_sampling_params.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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), \
Expand Down Expand Up @@ -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 (
Expand Down
24 changes: 9 additions & 15 deletions lightllm/server/core/objs/sampling_params.py
Original file line number Diff line number Diff line change
Expand Up @@ -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_ = [
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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:
Expand Down
51 changes: 51 additions & 0 deletions test/test_api/test_seed_validation.py
Original file line number Diff line number Diff line change
@@ -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)
Loading