TensorVizion's picture
Upload 3 files
6f6aff6 verified
Raw History Blame Contribute Delete
2.34 kB
import gradio as gr
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
MODEL_ID = "TensorVizion/mistral-nemo-alpaca-finetune"
# Configure 4-bit quantization to fit the 12B model in a T4 GPU (16GB VRAM)
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
)
print("Loading tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
print("Loading model in 4-bit...")
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
quantization_config=bnb_config,
device_map="auto",
torch_dtype=torch.bfloat16,
)
def generate_response(message, history):
# Format chat history for the model
chat_history = []
for user_msg, assistant_msg in history:
chat_history.append({"role": "user", "content": user_msg})
chat_history.append({"role": "assistant", "content": assistant_msg})
chat_history.append({"role": "user", "content": message})
# Apply the model's built-in chat template (handles Mistral Nemo formatting)
input_ids = tokenizer.apply_chat_template(
chat_history,
add_generation_prompt=True,
return_tensors="pt"
).to(model.device)
# Generate response
with torch.no_grad():
output_ids = model.generate(
input_ids,
max_new_tokens=512,
temperature=0.7,
top_p=0.9,
do_sample=True,
pad_token_id=tokenizer.eos_token_id
)
# Decode only the newly generated tokens
response = tokenizer.decode(output_ids[0][input_ids.shape[1]:], skip_special_tokens=True)
return response
# Create Gradio Chat Interface
demo = gr.ChatInterface(
fn=generate_response,
title="Mistral Nemo 12B Alpaca Finetune",
description="Chat with the TensorVizion Mistral Nemo 12B model. Runs efficiently in 4-bit quantization.",
examples=[
"Explain the concept of quantum entanglement in simple terms.",
"Write a short Python script to scrape a website's title.",
"What are the main differences between supervised and unsupervised learning?"
],
retry_btn=None,
undo_btn=None,
clear_btn="Clear Chat"
)
if __name__ == "__main__":
demo.launch()