import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer, StoppingCriteria, StoppingCriteriaList
from threading import Thread
checkpoint_path = "./Bormotuha-H0"
model = AutoModelForCausalLM.from_pretrained(
checkpoint_path,
device_map="cuda",
dtype="float16"
).half()
tokenizer = AutoTokenizer.from_pretrained(checkpoint_path, use_fast=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
class StopOnTokens(StoppingCriteria):
def __init__(self, stop_token_ids_list):
self.stop_token_ids_list = stop_token_ids_list
def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool:
for stop_ids in self.stop_token_ids_list:
if len(input_ids[0]) >= len(stop_ids):
if input_ids[0][-len(stop_ids):].tolist() == stop_ids:
return True
return False
stop_words = ["</answer>"]
stop_ids_list = [tokenizer.encode(word, add_special_tokens=False) for word in stop_words]
stopping_criteria = StoppingCriteriaList([StopOnTokens(stop_ids_list)])
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=False)
test_input=[]
test_input.append("<system>Ответить на вопрос кратко</system><newline><question>Чем знаменит Аристотель?</question><newline><answer>")
test_input.append("<system>Ответить на вопрос</system><newline><question>Что такое русский космизм и есть ли связь между ним и космической программой СССР?</question><newline><answer>")
for i, prompt in enumerate(test_input):
question=prompt.split("</system>")[1].split("</")[0]
question=question.replace("<question>","").replace("<newline>","\n")
print(f"\n\nЗадание {i}: {question}\n")
processed_input = "▁" + prompt.lstrip().replace(" ", "▁").replace("\n", "<newline>")
inputs = tokenizer(
processed_input,
return_tensors="pt",
add_special_tokens=False,
).to(model.device)
generation_kwargs = dict(
inputs,
streamer=streamer,
max_new_tokens=1024,
do_sample=True,
top_p=0.6,
top_k=40,
repetition_penalty=1.3,
pad_token_id=tokenizer.pad_token_id,
eos_token_id=tokenizer.eos_token_id,
stopping_criteria=stopping_criteria
)
thread = Thread(target=model.generate, kwargs=generation_kwargs)
thread.start()
full_response = ""
for new_text in streamer:
decoded_chunk = new_text.replace("▁", " ").replace("<newline>", "\n").replace("<space>", " ")
if decoded_chunk in stop_words:
break
print(decoded_chunk, end="", flush=True)