[Refactor] Move sampling tokenizer validation helper (#32694)
This commit is contained in:
@@ -42,36 +42,6 @@ TOP_K_ALL = 1 << 30
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def raise_if_tokenizer_required(
|
|
||||||
tokenizer, stop_strs, stop_regex_strs, min_new_tokens=0
|
|
||||||
):
|
|
||||||
"""Raise ValueError if tokenizer-dependent features are used without a tokenizer.
|
|
||||||
|
|
||||||
String-based stop conditions (stop_strs, stop_regex_strs) require tokenizer.decode()
|
|
||||||
to convert output token IDs to text for matching. min_new_tokens requires the
|
|
||||||
tokenizer's eos_token_id to penalize. When skip_tokenizer_init=True, these cannot
|
|
||||||
be used.
|
|
||||||
"""
|
|
||||||
if tokenizer is not None:
|
|
||||||
return
|
|
||||||
|
|
||||||
if stop_strs:
|
|
||||||
raise ValueError(
|
|
||||||
f"stop={stop_strs!r} is unavailable when skip_tokenizer_init=True "
|
|
||||||
"(requires tokenizer to decode tokens to text for matching)."
|
|
||||||
)
|
|
||||||
if stop_regex_strs:
|
|
||||||
raise ValueError(
|
|
||||||
f"stop_regex={stop_regex_strs!r} is unavailable when skip_tokenizer_init=True "
|
|
||||||
"(requires tokenizer to decode tokens to text for matching)."
|
|
||||||
)
|
|
||||||
if min_new_tokens > 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"min_new_tokens={min_new_tokens} is unavailable when skip_tokenizer_init=True "
|
|
||||||
"(requires tokenizer for eos_token_id)."
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class SamplingParams(msgspec.Struct, kw_only=True, array_like=True):
|
class SamplingParams(msgspec.Struct, kw_only=True, array_like=True):
|
||||||
"""
|
"""
|
||||||
The sampling parameters.
|
The sampling parameters.
|
||||||
@@ -321,3 +291,33 @@ def _max_length_from_subpattern(subpattern: sre_parse.SubPattern):
|
|||||||
total += MAX_LEN
|
total += MAX_LEN
|
||||||
|
|
||||||
return total
|
return total
|
||||||
|
|
||||||
|
|
||||||
|
def raise_if_tokenizer_required(
|
||||||
|
tokenizer, stop_strs, stop_regex_strs, min_new_tokens=0
|
||||||
|
):
|
||||||
|
"""Raise ValueError if tokenizer-dependent features are used without a tokenizer.
|
||||||
|
|
||||||
|
String-based stop conditions (stop_strs, stop_regex_strs) require tokenizer.decode()
|
||||||
|
to convert output token IDs to text for matching. min_new_tokens requires the
|
||||||
|
tokenizer's eos_token_id to penalize. When skip_tokenizer_init=True, these cannot
|
||||||
|
be used.
|
||||||
|
"""
|
||||||
|
if tokenizer is not None:
|
||||||
|
return
|
||||||
|
|
||||||
|
if stop_strs:
|
||||||
|
raise ValueError(
|
||||||
|
f"stop={stop_strs!r} is unavailable when skip_tokenizer_init=True "
|
||||||
|
"(requires tokenizer to decode tokens to text for matching)."
|
||||||
|
)
|
||||||
|
if stop_regex_strs:
|
||||||
|
raise ValueError(
|
||||||
|
f"stop_regex={stop_regex_strs!r} is unavailable when skip_tokenizer_init=True "
|
||||||
|
"(requires tokenizer to decode tokens to text for matching)."
|
||||||
|
)
|
||||||
|
if min_new_tokens > 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"min_new_tokens={min_new_tokens} is unavailable when skip_tokenizer_init=True "
|
||||||
|
"(requires tokenizer for eos_token_id)."
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user