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
6 changes: 4 additions & 2 deletions lightllm/server/api_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]] = {}
Expand Down Expand Up @@ -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

Expand Down
4 changes: 3 additions & 1 deletion lightllm/server/core/objs/py_sampling_params.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
17 changes: 16 additions & 1 deletion lightllm/server/core/objs/sampling_params.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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()

Expand Down
Loading