From 5d06c096821540a3012c117c9fdf69345e4bae0e Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 6 Aug 2026 20:10:58 +0800 Subject: [PATCH 1/6] fix(api): validate and normalize request seeds --- lightllm/server/api_models.py | 4 +- .../server/core/objs/py_sampling_params.py | 4 +- lightllm/server/core/objs/sampling_params.py | 5 +- test/test_api/test_seed_params.py | 51 +++++++++++++++++++ 4 files changed, 60 insertions(+), 4 deletions(-) create mode 100644 test/test_api/test_seed_params.py diff --git a/lightllm/server/api_models.py b/lightllm/server/api_models.py index 1737d2774d..eda4349911 100644 --- a/lightllm/server/api_models.py +++ b/lightllm/server/api_models.py @@ -151,7 +151,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) # Class variables to store loaded default values _loaded_defaults: ClassVar[Dict[str, Any]] = {} @@ -231,7 +231,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) 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 2514d9dacb..1300ccc17a 100644 --- a/lightllm/server/core/objs/py_sampling_params.py +++ b/lightllm/server/core/objs/py_sampling_params.py @@ -92,7 +92,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 = -1 if seed is None else seed if self.do_sample is False: self.temperature = 1.0 self.top_p = 1.0 @@ -154,6 +154,8 @@ 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}." ) + if self.seed < -1: + raise ValueError(f"seed must be -1 (random), or a non-negative integer, got {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 c122884ea1..9179c4c3b8 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -344,7 +344,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) - self.seed = kwargs.get("seed", -1) + seed = kwargs.get("seed") + self.seed = -1 if seed is None else seed prompt_logprobs = kwargs.get("prompt_logprobs", None) self.prompt_logprobs = -1 if prompt_logprobs is None else int(prompt_logprobs) @@ -448,6 +449,8 @@ 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") + if self.seed < -1: + raise ValueError(f"seed must be -1 (random), or a non-negative integer, got {self.seed}") self._verify_allowed_token_ids() self._verify_grammar_constraint() diff --git a/test/test_api/test_seed_params.py b/test/test_api/test_seed_params.py new file mode 100644 index 0000000000..7569f50a6e --- /dev/null +++ b/test/test_api/test_seed_params.py @@ -0,0 +1,51 @@ +import pytest +from pydantic import ValidationError + +from lightllm.server.api_models import ChatCompletionRequest, CompletionRequest +from lightllm.server.core.objs.py_sampling_params import SamplingParams as PySamplingParams +from lightllm.server.core.objs.sampling_params import SamplingParams + + +def _chat_request(**kwargs): + return ChatCompletionRequest(messages=[{"role": "user", "content": "hi"}], **kwargs) + + +def _completion_request(**kwargs): + return CompletionRequest(model="test-model", prompt="hi", **kwargs) + + +@pytest.mark.parametrize("request_factory", [_chat_request, _completion_request]) +def test_request_seed_defaults_to_none(request_factory): + assert request_factory().seed is None + assert request_factory(seed=None).seed is None + + +@pytest.mark.parametrize("request_factory", [_chat_request, _completion_request]) +@pytest.mark.parametrize("seed", [-1, 0, 42]) +def test_request_seed_accepts_supported_values(request_factory, seed): + assert request_factory(seed=seed).seed == seed + + +@pytest.mark.parametrize("request_factory", [_chat_request, _completion_request]) +def test_request_seed_rejects_values_below_random_sentinel(request_factory): + with pytest.raises(ValidationError, match="greater than or equal to -1"): + request_factory(seed=-2) + + +def test_sampling_params_normalize_none_seed_to_random_sentinel(): + params = SamplingParams() + params.init(tokenizer=None, seed=None) + assert params.seed == -1 + + py_params = PySamplingParams(seed=None) + assert py_params.seed == -1 + + +def test_sampling_params_reject_values_below_random_sentinel(): + with pytest.raises(ValueError, match="seed must be -1"): + params = SamplingParams() + params.init(tokenizer=None, seed=-2) + + py_params = PySamplingParams(seed=-2) + with pytest.raises(ValueError, match="seed must be -1"): + py_params.verify() From 64d4c393b390191a264efc8dc747e25d926025bc Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 6 Aug 2026 20:36:15 +0800 Subject: [PATCH 2/6] fix(api): cap seeds at int64 max --- lightllm/server/api_models.py | 6 ++-- .../server/core/objs/py_sampling_params.py | 5 +-- lightllm/server/core/objs/sampling_params.py | 10 ++++-- test/test_api/test_seed_params.py | 31 ++++++++++++++++++- 4 files changed, 44 insertions(+), 8 deletions(-) diff --git a/lightllm/server/api_models.py b/lightllm/server/api_models.py index eda4349911..d9fa86f5c5 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] = Field(default=None, ge=-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] = Field(default=None, ge=-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 1300ccc17a..19c9f2a52f 100644 --- a/lightllm/server/core/objs/py_sampling_params.py +++ b/lightllm/server/core/objs/py_sampling_params.py @@ -11,6 +11,7 @@ _SAMPLING_EPS = 1e-5 # 用环境变量控制是否进行输入惩罚的默认值 DEFAULT_INPUT_PENALTY = os.getenv("INPUT_PENALTY", "False").upper() in ["ON", "TRUE", "1"] +MAX_SEED = (1 << 63) - 1 class SamplingParams: @@ -154,8 +155,8 @@ 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}." ) - if self.seed < -1: - raise ValueError(f"seed must be -1 (random), or a non-negative integer, got {self.seed}") + if self.seed < -1 or self.seed > MAX_SEED: + raise ValueError(f"seed must be -1 (random), or an integer in [0, {MAX_SEED}], got {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 9179c4c3b8..ea46127557 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -20,6 +20,7 @@ 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 class StopSequence(ctypes.Structure): @@ -345,7 +346,10 @@ def init(self, tokenizer, **kwargs): 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 = kwargs.get("seed") - self.seed = -1 if seed is None else seed + seed = -1 if seed is None else seed + if seed < -1 or seed > MAX_SEED: + raise ValueError(f"seed must be -1 (random), or an integer in [0, {MAX_SEED}], got {seed}") + self.seed = seed prompt_logprobs = kwargs.get("prompt_logprobs", None) self.prompt_logprobs = -1 if prompt_logprobs is None else int(prompt_logprobs) @@ -449,8 +453,8 @@ 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") - if self.seed < -1: - raise ValueError(f"seed must be -1 (random), or a non-negative integer, got {self.seed}") + if self.seed < -1 or self.seed > MAX_SEED: + raise ValueError(f"seed must be -1 (random), or an integer in [0, {MAX_SEED}], got {self.seed}") self._verify_allowed_token_ids() self._verify_grammar_constraint() diff --git a/test/test_api/test_seed_params.py b/test/test_api/test_seed_params.py index 7569f50a6e..bbf5f06c0a 100644 --- a/test/test_api/test_seed_params.py +++ b/test/test_api/test_seed_params.py @@ -5,6 +5,8 @@ from lightllm.server.core.objs.py_sampling_params import SamplingParams as PySamplingParams from lightllm.server.core.objs.sampling_params import SamplingParams +MAX_SEED = (1 << 63) - 1 + def _chat_request(**kwargs): return ChatCompletionRequest(messages=[{"role": "user", "content": "hi"}], **kwargs) @@ -21,7 +23,7 @@ def test_request_seed_defaults_to_none(request_factory): @pytest.mark.parametrize("request_factory", [_chat_request, _completion_request]) -@pytest.mark.parametrize("seed", [-1, 0, 42]) +@pytest.mark.parametrize("seed", [-1, 0, 42, MAX_SEED]) def test_request_seed_accepts_supported_values(request_factory, seed): assert request_factory(seed=seed).seed == seed @@ -32,6 +34,13 @@ def test_request_seed_rejects_values_below_random_sentinel(request_factory): request_factory(seed=-2) +@pytest.mark.parametrize("request_factory", [_chat_request, _completion_request]) +@pytest.mark.parametrize("seed", [MAX_SEED + 1, (1 << 64) - 1, 10**100]) +def test_request_seed_rejects_values_above_int64_max(request_factory, seed): + with pytest.raises(ValidationError, match=f"less than or equal to {MAX_SEED}"): + request_factory(seed=seed) + + def test_sampling_params_normalize_none_seed_to_random_sentinel(): params = SamplingParams() params.init(tokenizer=None, seed=None) @@ -41,6 +50,15 @@ def test_sampling_params_normalize_none_seed_to_random_sentinel(): assert py_params.seed == -1 +def test_sampling_params_accept_int64_max_seed(): + params = SamplingParams() + params.init(tokenizer=None, seed=MAX_SEED) + assert params.seed == MAX_SEED + + py_params = PySamplingParams(seed=MAX_SEED) + py_params.verify() + + def test_sampling_params_reject_values_below_random_sentinel(): with pytest.raises(ValueError, match="seed must be -1"): params = SamplingParams() @@ -49,3 +67,14 @@ def test_sampling_params_reject_values_below_random_sentinel(): py_params = PySamplingParams(seed=-2) with pytest.raises(ValueError, match="seed must be -1"): py_params.verify() + + +@pytest.mark.parametrize("seed", [MAX_SEED + 1, (1 << 64) - 1, 10**100]) +def test_sampling_params_reject_seed_before_int64_wraparound(seed): + with pytest.raises(ValueError, match=f"integer in \\[0, {MAX_SEED}\\]"): + params = SamplingParams() + params.init(tokenizer=None, seed=seed) + + py_params = PySamplingParams(seed=seed) + with pytest.raises(ValueError, match=f"integer in \\[0, {MAX_SEED}\\]"): + py_params.verify() From 659b21801c51bfd53e0fd920c553b122d5e5441d Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 6 Aug 2026 20:38:46 +0800 Subject: [PATCH 3/6] style: format seed boundary tests --- test/test_api/test_seed_params.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/test/test_api/test_seed_params.py b/test/test_api/test_seed_params.py index bbf5f06c0a..5767b9a5ce 100644 --- a/test/test_api/test_seed_params.py +++ b/test/test_api/test_seed_params.py @@ -35,7 +35,7 @@ def test_request_seed_rejects_values_below_random_sentinel(request_factory): @pytest.mark.parametrize("request_factory", [_chat_request, _completion_request]) -@pytest.mark.parametrize("seed", [MAX_SEED + 1, (1 << 64) - 1, 10**100]) +@pytest.mark.parametrize("seed", [MAX_SEED + 1, (1 << 64) - 1, 10 ** 100]) def test_request_seed_rejects_values_above_int64_max(request_factory, seed): with pytest.raises(ValidationError, match=f"less than or equal to {MAX_SEED}"): request_factory(seed=seed) @@ -69,7 +69,7 @@ def test_sampling_params_reject_values_below_random_sentinel(): py_params.verify() -@pytest.mark.parametrize("seed", [MAX_SEED + 1, (1 << 64) - 1, 10**100]) +@pytest.mark.parametrize("seed", [MAX_SEED + 1, (1 << 64) - 1, 10 ** 100]) def test_sampling_params_reject_seed_before_int64_wraparound(seed): with pytest.raises(ValueError, match=f"integer in \\[0, {MAX_SEED}\\]"): params = SamplingParams() From fab23ca72e6e81ef9264d28541f93acf47a67385 Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 6 Aug 2026 20:48:10 +0800 Subject: [PATCH 4/6] refactor: centralize seed validation --- lightllm/server/api_models.py | 3 +-- lightllm/server/core/objs/py_sampling_params.py | 7 +++---- lightllm/server/core/objs/sampling_params.py | 11 +++-------- lightllm/utils/seed_utils.py | 16 ++++++++++++++++ test/test_api/test_seed_params.py | 9 +++------ 5 files changed, 26 insertions(+), 20 deletions(-) create mode 100644 lightllm/utils/seed_utils.py diff --git a/lightllm/server/api_models.py b/lightllm/server/api_models.py index d9fa86f5c5..b992dfca03 100644 --- a/lightllm/server/api_models.py +++ b/lightllm/server/api_models.py @@ -5,8 +5,7 @@ from pydantic import BaseModel, Field, field_validator, model_validator from typing import Any, Dict, List, Optional, Union, Literal, ClassVar from transformers import GenerationConfig - -MAX_SEED = (1 << 63) - 1 +from lightllm.utils.seed_utils import MAX_SEED class ImageURL(BaseModel): diff --git a/lightllm/server/core/objs/py_sampling_params.py b/lightllm/server/core/objs/py_sampling_params.py index 19c9f2a52f..8db42db0fa 100644 --- a/lightllm/server/core/objs/py_sampling_params.py +++ b/lightllm/server/core/objs/py_sampling_params.py @@ -6,12 +6,12 @@ from typing import List, Optional, Union, Tuple from transformers import GenerationConfig from lightllm.server.req_id_generator import MAX_BEST_OF +from lightllm.utils.seed_utils import normalize_seed, validate_seed _SAMPLING_EPS = 1e-5 # 用环境变量控制是否进行输入惩罚的默认值 DEFAULT_INPUT_PENALTY = os.getenv("INPUT_PENALTY", "False").upper() in ["ON", "TRUE", "1"] -MAX_SEED = (1 << 63) - 1 class SamplingParams: @@ -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 = -1 if seed is None else seed + self.seed = normalize_seed(seed) if self.do_sample is False: self.temperature = 1.0 self.top_p = 1.0 @@ -155,8 +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}." ) - if self.seed < -1 or self.seed > MAX_SEED: - raise ValueError(f"seed must be -1 (random), or an integer in [0, {MAX_SEED}], got {self.seed}") + 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 ea46127557..ee9fa1cdb3 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -4,6 +4,7 @@ from transformers import GenerationConfig from lightllm.server.req_id_generator import MAX_BEST_OF from lightllm.utils.envs_utils import get_env_start_args +from lightllm.utils.seed_utils import normalize_seed, validate_seed from .pd_kv_trans_params import PDKVTransParamObj _SAMPLING_EPS = 1e-5 @@ -20,7 +21,6 @@ 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 class StopSequence(ctypes.Structure): @@ -345,11 +345,7 @@ 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 = kwargs.get("seed") - seed = -1 if seed is None else seed - if seed < -1 or seed > MAX_SEED: - raise ValueError(f"seed must be -1 (random), or an integer in [0, {MAX_SEED}], got {seed}") - self.seed = seed + self.seed = normalize_seed(kwargs.get("seed")) prompt_logprobs = kwargs.get("prompt_logprobs", None) self.prompt_logprobs = -1 if prompt_logprobs is None else int(prompt_logprobs) @@ -453,8 +449,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") - if self.seed < -1 or self.seed > MAX_SEED: - raise ValueError(f"seed must be -1 (random), or an integer in [0, {MAX_SEED}], got {self.seed}") + validate_seed(self.seed) self._verify_allowed_token_ids() self._verify_grammar_constraint() diff --git a/lightllm/utils/seed_utils.py b/lightllm/utils/seed_utils.py new file mode 100644 index 0000000000..649670e029 --- /dev/null +++ b/lightllm/utils/seed_utils.py @@ -0,0 +1,16 @@ +from typing import Optional + + +RANDOM_SEED = -1 +MAX_SEED = (1 << 63) - 1 + + +def validate_seed(seed: int) -> None: + if seed < RANDOM_SEED or seed > MAX_SEED: + raise ValueError(f"seed must be {RANDOM_SEED} (random), or an integer in [0, {MAX_SEED}], got {seed}") + + +def normalize_seed(seed: Optional[int]) -> int: + seed = RANDOM_SEED if seed is None else seed + validate_seed(seed) + return seed diff --git a/test/test_api/test_seed_params.py b/test/test_api/test_seed_params.py index 5767b9a5ce..be5f556200 100644 --- a/test/test_api/test_seed_params.py +++ b/test/test_api/test_seed_params.py @@ -4,8 +4,7 @@ from lightllm.server.api_models import ChatCompletionRequest, CompletionRequest from lightllm.server.core.objs.py_sampling_params import SamplingParams as PySamplingParams from lightllm.server.core.objs.sampling_params import SamplingParams - -MAX_SEED = (1 << 63) - 1 +from lightllm.utils.seed_utils import MAX_SEED def _chat_request(**kwargs): @@ -64,9 +63,8 @@ def test_sampling_params_reject_values_below_random_sentinel(): params = SamplingParams() params.init(tokenizer=None, seed=-2) - py_params = PySamplingParams(seed=-2) with pytest.raises(ValueError, match="seed must be -1"): - py_params.verify() + PySamplingParams(seed=-2) @pytest.mark.parametrize("seed", [MAX_SEED + 1, (1 << 64) - 1, 10 ** 100]) @@ -75,6 +73,5 @@ def test_sampling_params_reject_seed_before_int64_wraparound(seed): params = SamplingParams() params.init(tokenizer=None, seed=seed) - py_params = PySamplingParams(seed=seed) with pytest.raises(ValueError, match=f"integer in \\[0, {MAX_SEED}\\]"): - py_params.verify() + PySamplingParams(seed=seed) From 2cb7d9648ad30bb0857acde7857cd55ec34717ec Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 6 Aug 2026 20:51:03 +0800 Subject: [PATCH 5/6] test: remove seed parameter tests --- test/test_api/test_seed_params.py | 77 ------------------------------- 1 file changed, 77 deletions(-) delete mode 100644 test/test_api/test_seed_params.py diff --git a/test/test_api/test_seed_params.py b/test/test_api/test_seed_params.py deleted file mode 100644 index be5f556200..0000000000 --- a/test/test_api/test_seed_params.py +++ /dev/null @@ -1,77 +0,0 @@ -import pytest -from pydantic import ValidationError - -from lightllm.server.api_models import ChatCompletionRequest, CompletionRequest -from lightllm.server.core.objs.py_sampling_params import SamplingParams as PySamplingParams -from lightllm.server.core.objs.sampling_params import SamplingParams -from lightllm.utils.seed_utils import MAX_SEED - - -def _chat_request(**kwargs): - return ChatCompletionRequest(messages=[{"role": "user", "content": "hi"}], **kwargs) - - -def _completion_request(**kwargs): - return CompletionRequest(model="test-model", prompt="hi", **kwargs) - - -@pytest.mark.parametrize("request_factory", [_chat_request, _completion_request]) -def test_request_seed_defaults_to_none(request_factory): - assert request_factory().seed is None - assert request_factory(seed=None).seed is None - - -@pytest.mark.parametrize("request_factory", [_chat_request, _completion_request]) -@pytest.mark.parametrize("seed", [-1, 0, 42, MAX_SEED]) -def test_request_seed_accepts_supported_values(request_factory, seed): - assert request_factory(seed=seed).seed == seed - - -@pytest.mark.parametrize("request_factory", [_chat_request, _completion_request]) -def test_request_seed_rejects_values_below_random_sentinel(request_factory): - with pytest.raises(ValidationError, match="greater than or equal to -1"): - request_factory(seed=-2) - - -@pytest.mark.parametrize("request_factory", [_chat_request, _completion_request]) -@pytest.mark.parametrize("seed", [MAX_SEED + 1, (1 << 64) - 1, 10 ** 100]) -def test_request_seed_rejects_values_above_int64_max(request_factory, seed): - with pytest.raises(ValidationError, match=f"less than or equal to {MAX_SEED}"): - request_factory(seed=seed) - - -def test_sampling_params_normalize_none_seed_to_random_sentinel(): - params = SamplingParams() - params.init(tokenizer=None, seed=None) - assert params.seed == -1 - - py_params = PySamplingParams(seed=None) - assert py_params.seed == -1 - - -def test_sampling_params_accept_int64_max_seed(): - params = SamplingParams() - params.init(tokenizer=None, seed=MAX_SEED) - assert params.seed == MAX_SEED - - py_params = PySamplingParams(seed=MAX_SEED) - py_params.verify() - - -def test_sampling_params_reject_values_below_random_sentinel(): - with pytest.raises(ValueError, match="seed must be -1"): - params = SamplingParams() - params.init(tokenizer=None, seed=-2) - - with pytest.raises(ValueError, match="seed must be -1"): - PySamplingParams(seed=-2) - - -@pytest.mark.parametrize("seed", [MAX_SEED + 1, (1 << 64) - 1, 10 ** 100]) -def test_sampling_params_reject_seed_before_int64_wraparound(seed): - with pytest.raises(ValueError, match=f"integer in \\[0, {MAX_SEED}\\]"): - params = SamplingParams() - params.init(tokenizer=None, seed=seed) - - with pytest.raises(ValueError, match=f"integer in \\[0, {MAX_SEED}\\]"): - PySamplingParams(seed=seed) From 36e6afaf33d02573ad2d40544956c61e57008f3a Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 6 Aug 2026 20:56:24 +0800 Subject: [PATCH 6/6] refactor: keep seed helpers with sampling params --- lightllm/server/api_models.py | 3 ++- lightllm/server/core/objs/py_sampling_params.py | 2 +- lightllm/server/core/objs/sampling_params.py | 17 +++++++++++++++-- lightllm/utils/seed_utils.py | 16 ---------------- 4 files changed, 18 insertions(+), 20 deletions(-) delete mode 100644 lightllm/utils/seed_utils.py diff --git a/lightllm/server/api_models.py b/lightllm/server/api_models.py index b992dfca03..d9fa86f5c5 100644 --- a/lightllm/server/api_models.py +++ b/lightllm/server/api_models.py @@ -5,7 +5,8 @@ from pydantic import BaseModel, Field, field_validator, model_validator from typing import Any, Dict, List, Optional, Union, Literal, ClassVar from transformers import GenerationConfig -from lightllm.utils.seed_utils import MAX_SEED + +MAX_SEED = (1 << 63) - 1 class ImageURL(BaseModel): diff --git a/lightllm/server/core/objs/py_sampling_params.py b/lightllm/server/core/objs/py_sampling_params.py index 8db42db0fa..0a158d7f7d 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 lightllm.utils.seed_utils import normalize_seed, validate_seed +from .sampling_params import normalize_seed, validate_seed _SAMPLING_EPS = 1e-5 diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index ee9fa1cdb3..e23f13ec09 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -4,7 +4,6 @@ from transformers import GenerationConfig from lightllm.server.req_id_generator import MAX_BEST_OF from lightllm.utils.envs_utils import get_env_start_args -from lightllm.utils.seed_utils import normalize_seed, validate_seed from .pd_kv_trans_params import PDKVTransParamObj _SAMPLING_EPS = 1e-5 @@ -21,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): @@ -345,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 = normalize_seed(kwargs.get("seed")) + 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) diff --git a/lightllm/utils/seed_utils.py b/lightllm/utils/seed_utils.py deleted file mode 100644 index 649670e029..0000000000 --- a/lightllm/utils/seed_utils.py +++ /dev/null @@ -1,16 +0,0 @@ -from typing import Optional - - -RANDOM_SEED = -1 -MAX_SEED = (1 << 63) - 1 - - -def validate_seed(seed: int) -> None: - if seed < RANDOM_SEED or seed > MAX_SEED: - raise ValueError(f"seed must be {RANDOM_SEED} (random), or an integer in [0, {MAX_SEED}], got {seed}") - - -def normalize_seed(seed: Optional[int]) -> int: - seed = RANDOM_SEED if seed is None else seed - validate_seed(seed) - return seed