LT_QA_demo / streaming.py
tetragrama's picture
Upload streaming.py
580b318 verified
Raw
History Blame Contribute Delete
2.57 kB
import torch
from unsloth import FastModel
from transformers import TextStreamer
from unsloth.chat_templates import get_chat_template
LORA_MODEL_DIR = "gemma-3-12b-it-lt-lora"
MAX_SEQ_LENGTH = 2048
LOAD_IN_4BIT = False
MAX_NEW_TOKENS = 256
def load_model():
model, tokenizer = FastModel.from_pretrained(
model_name=LORA_MODEL_DIR,
max_seq_length=MAX_SEQ_LENGTH,
load_in_4bit=LOAD_IN_4BIT,
)
tokenizer = get_chat_template(
tokenizer,
chat_template="gemma-3",
)
FastModel.for_inference(model)
return model, tokenizer
def make_text_message(role: str, text: str):
return {
"role": role,
"content": [
{"type": "text", "text": text}
],
}
def stream_reply(model, tokenizer, messages, max_new_tokens=MAX_NEW_TOKENS):
inputs = tokenizer.apply_chat_template(
messages,
add_generation_prompt=True,
tokenize=True,
return_tensors="pt",
return_dict=True,
).to(model.device)
streamer = TextStreamer(
tokenizer,
skip_prompt=True,
skip_special_tokens=True,
)
with torch.inference_mode():
outputs = model.generate(
**inputs,
streamer=streamer,
max_new_tokens=max_new_tokens,
do_sample=True,
temperature=0.7,
top_p=0.9,
top_k=64,
pad_token_id=tokenizer.eos_token_id,
)
# Kad galėtume atsakymą įdėti atgal į history
prompt_len = inputs["input_ids"].shape[1]
new_tokens = outputs[0][prompt_len:]
reply = tokenizer.decode(new_tokens, skip_special_tokens=True).strip()
return reply
def main():
model, tokenizer = load_model()
messages = []
print("Model loaded. Type 'exit' to quit.\n")
while True:
try:
user_text = input("You: ").strip()
except (EOFError, KeyboardInterrupt):
print("\nBye.")
break
if not user_text:
continue
if user_text.lower() in {"exit", "quit"}:
print("Bye.")
break
messages.append(make_text_message("user", user_text))
try:
print("Assistant: ", end="", flush=True)
reply = stream_reply(model, tokenizer, messages)
print() # nauja eilutė po streaminimo
except Exception as e:
print(f"\nError during generation: {e}")
break
messages.append(make_text_message("assistant", reply))
if __name__ == "__main__":
main()