File size: 2,339 Bytes
ab52edb
6f6aff6
 
ab52edb
6f6aff6
ab52edb
6f6aff6
 
 
 
 
 
 
ab52edb
6f6aff6
 
ab52edb
6f6aff6
 
 
 
 
 
 
ab52edb
6f6aff6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ab52edb
6f6aff6
 
 
 
 
 
 
 
 
ab52edb
6f6aff6
 
 
ab52edb
 
 
6f6aff6
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
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()