Fix a bug

This commit is contained in:
oobabooga 2023-04-16 01:40:47 -03:00
parent 16a3a5b039
commit 9d3c6d2dc3

View File

@ -193,7 +193,8 @@ def generate_reply(question, state, eos_token=None, stopping_strings=[]):
# Handling the stopping strings
stopping_criteria_list = transformers.StoppingCriteriaList()
for st in (stopping_strings, ast.literal_eval(f"[{state['custom_stopping_strings']}]")]):
print(ast.literal_eval(f"[{state['custom_stopping_strings']}]"))
for st in (stopping_strings, ast.literal_eval(f"[{state['custom_stopping_strings']}]")):
if type(st) is list and len(st) > 0:
sentinel_token_ids = [encode(string, add_special_tokens=False) for string in st]
stopping_criteria_list.append(_SentinelTokenStoppingCriteria(sentinel_token_ids=sentinel_token_ids, starting_idx=len(input_ids[0])))