diff --git a/lightllm/server/api_models.py b/lightllm/server/api_models.py index 1737d2774..d9fa86f5c 100644 --- a/lightllm/server/api_models.py +++ b/lightllm/server/api_models.py @@ -6,6 +6,8 @@ from typing import Any, Dict, List, Optional, Union, Literal, ClassVar from transformers import GenerationConfig +MAX_SEED = (1 << 63) - 1 + class ImageURL(BaseModel): url: str @@ -151,7 +153,7 @@ class CompletionRequest(BaseModel): top_k: Optional[int] = -1 repetition_penalty: Optional[float] = 1.0 ignore_eos: Optional[bool] = False - seed: Optional[int] = -1 + seed: Optional[int] = Field(default=None, ge=-1, le=MAX_SEED) # Class variables to store loaded default values _loaded_defaults: ClassVar[Dict[str, Any]] = {} @@ -231,7 +233,7 @@ class ChatCompletionRequest(BaseModel): top_k: Optional[int] = -1 repetition_penalty: Optional[float] = 1.0 ignore_eos: Optional[bool] = False - seed: Optional[int] = -1 + seed: Optional[int] = Field(default=None, ge=-1, le=MAX_SEED) role_settings: Optional[Dict[str, str]] = None character_settings: Optional[List[Dict[str, str]]] = None diff --git a/lightllm/server/core/objs/py_sampling_params.py b/lightllm/server/core/objs/py_sampling_params.py index 2514d9dac..0a158d7f7 100644 --- a/lightllm/server/core/objs/py_sampling_params.py +++ b/lightllm/server/core/objs/py_sampling_params.py @@ -6,6 +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 _SAMPLING_EPS = 1e-5 @@ -92,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 = seed + self.seed = normalize_seed(seed) if self.do_sample is False: self.temperature = 1.0 self.top_p = 1.0 @@ -154,6 +155,7 @@ 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( diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index c122884ea..e23f13ec0 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -20,6 +20,16 @@ JSON_SCHEMA_MAX_LENGTH = int(os.getenv("LIGHTLLM_JSON_SCHEMA_MAX_LENGTH", 2048)) INVALID_TOKEN_IDS_MAX_LENGTH = int(os.getenv("LIGHTLLM_INVALID_TOKEN_IDS_MAX_LENGTH", 10)) MAX_PROMPT_LOGPROBS = int(os.getenv("LIGHTLLM_MAX_PROMPT_LOGPROBS", 1024)) +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): @@ -344,7 +354,11 @@ 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) - self.seed = kwargs.get("seed", -1) + 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 prompt_logprobs = kwargs.get("prompt_logprobs", None) self.prompt_logprobs = -1 if prompt_logprobs is None else int(prompt_logprobs) @@ -448,6 +462,7 @@ 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()