[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__)
|
||||
|
||||
|
||||
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):
|
||||
"""
|
||||
The sampling parameters.
|
||||
@@ -321,3 +291,33 @@ def _max_length_from_subpattern(subpattern: sre_parse.SubPattern):
|
||||
total += MAX_LEN
|
||||
|
||||
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