Disk FSM cache and adjust code. (#63)

This commit is contained in:
Liangsheng Yin
2024-01-20 21:26:11 -08:00
committed by GitHub
parent 0b2efc2adc
commit ca13f3b8c5
8 changed files with 207 additions and 299 deletions
@@ -45,7 +45,7 @@ class Req:
# for constrained decoding
self.regex_fsm = None
self.regex_fsm_state = None
self.regex_fsm_state = 0
def max_new_tokens(self):
return self.sampling_params.max_new_tokens
+12 -3
View File
@@ -111,7 +111,13 @@ class ModelRpcServer(rpyc.Service):
self.stream_interval = server_args.stream_interval
# Init the FSM cache for constrained generation
self.regex_fsm_cache = FSMCache(self.tokenizer)
self.regex_fsm_cache = FSMCache(
server_args.tokenizer_path,
{
"tokenizer_mode": server_args.tokenizer_mode,
"trust_remote_code": server_args.trust_remote_code,
},
)
# Init new token estimation
self.new_token_ratio = min(0.4 * server_args.schedule_conservativeness, 1.0)
@@ -213,6 +219,10 @@ class ModelRpcServer(rpyc.Service):
req.stream = recv_req.stream
req.tokenizer = self.tokenizer
# Init regex fsm
if req.sampling_params.regex is not None:
req.regex_fsm = self.regex_fsm_cache.init_fsm(req.sampling_params.regex)
# Truncate long prompts
req.input_ids = req.input_ids[: self.model_config.context_len - 1]
req.sampling_params.max_new_tokens = min(
@@ -322,11 +332,10 @@ class ModelRpcServer(rpyc.Service):
self.model_config.vocab_size, self.int_token_logit_bias
)
# init the regex fsm before first sampling
# Reset regex fsm state before first sampling due to retractions
for req in batch.reqs:
if req.sampling_params.regex is not None:
req.regex_fsm_state = 0
req.regex_fsm = self.regex_fsm_cache.get_fsm(req.sampling_params.regex)
if batch.extend_num_tokens != 0:
# Forward