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()