mmcgovern574's picture
Update app.py
2efcfde verified
Raw History Blame
7.36 kB
import json
import subprocess
from threading import Thread
import os
import torch
import spaces
import gradio as gr
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig, TextIteratorStreamer
from google.oauth2.service_account import Credentials
from googleapiclient.discovery import build
from optimum.bettertransformer import BetterTransformer
# Model Configuration
MODEL_ID = "mistralai/Mistral-Small-24B-Instruct-2501"
MODEL_NAME = MODEL_ID.split("/")[-1]
CONTEXT_LENGTH = 32768
EMOJI = "🌪️"
DESCRIPTION = f"Chat with my Digital Twin, powered by {MODEL_NAME}"
def load_system_message():
"""
Load the system prompt text from a private Google Doc
"""
doc_id = os.getenv("GOOGLE_DOC_ID", "")
if not doc_id:
print("Warning: No GOOGLE_DOC_ID found. Using default system message.")
return "You are a helpful assistant. First recognize user request and then reply carefully with thinking."
google_creds_json = os.getenv("GOOGLE_CREDS_JSON", "")
if not google_creds_json:
print("Warning: No GOOGLE_CREDS_JSON in environment. Using default message.")
return "You are a helpful assistant. First recognize user request and then reply carefully with thinking."
try:
creds_info = json.loads(google_creds_json)
creds = Credentials.from_service_account_info(
creds_info,
scopes=["https://www.googleapis.com/auth/documents.readonly"]
)
service = build("docs", "v1", credentials=creds)
doc = service.documents().get(documentId=doc_id).execute()
paragraphs = []
for element in doc.get("body", {}).get("content", []):
paragraph_elements = element.get("paragraph", {}).get("elements", [])
for run in paragraph_elements:
text_run = run.get("textRun", {})
if text_run.get("content"):
paragraphs.append(text_run["content"])
system_message = "".join(paragraphs).strip()
if not system_message:
print("Warning: Doc is empty. Using default system message.")
return "You are a helpful assistant. First recognize user request and then reply carefully with thinking."
return system_message
except Exception as e:
print(f"Error loading system message from Google Doc: {e}")
return "You are a helpful assistant. First recognize user request and then reply carefully with thinking."
SYSTEM_MESSAGE = load_system_message()
@spaces.GPU()
def predict(message, history):
# Format history using Mistral's chat template
messages = [{"role": "system", "content": SYSTEM_MESSAGE}]
for user, assistant in history:
messages.append({"role": "user", "content": user})
messages.append({"role": "assistant", "content": assistant})
messages.append({"role": "user", "content": message})
# Convert messages to Mistral format
prompt = tokenizer.apply_chat_template(messages, tokenize=False)
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
enc = tokenizer(prompt, return_tensors="pt", padding=True, truncation=True)
input_ids, attention_mask = enc.input_ids, enc.attention_mask
if input_ids.shape[1] > CONTEXT_LENGTH:
input_ids = input_ids[:, -CONTEXT_LENGTH:]
attention_mask = attention_mask[:, -CONTEXT_LENGTH:]
# Optimized generation parameters
generate_kwargs = dict(
input_ids=input_ids.to(device),
attention_mask=attention_mask.to(device),
streamer=streamer,
do_sample=True,
temperature=0.3,
max_new_tokens=400,
top_k=50,
repetition_penalty=1.1,
top_p=0.95,
use_cache=True,
num_beams=1,
pad_token_id=tokenizer.pad_token_id,
eos_token_id=tokenizer.eos_token_id
)
t = Thread(target=model.generate, kwargs=generate_kwargs)
t.start()
outputs = []
for new_token in streamer:
outputs.append(new_token)
yield "".join(outputs)
# Load model with optimized settings for Mistral-24B
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
quantization_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.bfloat16,
use_double_quant=True,
bnb_4bit_quant_type="nf4"
)
# Initialize tokenizer
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.pad_token_id = tokenizer.eos_token_id
# Load model with optimizations
try:
print("Loading model with optimizations...")
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
device_map="auto",
quantization_config=quantization_config,
torch_dtype=torch.bfloat16
)
# Apply Better Transformer optimization
print("Applying Better Transformer optimization...")
model = BetterTransformer.transform(model)
# Apply torch compile optimization
print("Applying torch compile optimization...")
model = torch.compile(model, mode="reduce-overhead")
print("Model loading and optimization complete!")
except Exception as e:
print(f"Warning: Could not apply all optimizations: {e}")
# Fallback to basic model loading if optimizations fail
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
device_map="auto",
quantization_config=quantization_config,
torch_dtype=torch.bfloat16
)
# Custom CSS
CSS = """
#title {
text-align: center !important;
}
#disclaimer-container {
display: flex !important;
justify-content: center !important;
width: 100% !important;
margin: 0 auto !important;
}
#disclaimer {
text-align: center !important;
color: rgba(153, 153, 153, 0.5) !important;
font-size: 0.7em !important;
margin: 20px auto !important;
padding-bottom: 20px !important;
max-width: 600px !important;
opacity: 0.4 !important;
line-height: 1.4 !important;
font-weight: 250 !important;
font-style: italic !important;
}
"""
# Create Gradio interface
with gr.Blocks(css=CSS) as demo:
gr.Markdown(
"""
# Chat with my Digital Twin!
""",
elem_id="title"
)
chat = gr.ChatInterface(
fn=predict,
chatbot=gr.Chatbot(height=400),
examples=[
["Tell me the story of your life, the choices you have made and why you made them."],
["What are some of your favorite books or ideas?"],
["What is a significant technical project you led during your career?"],
["What mental models have you developed and found useful?"],
["When have you applied your mental models?"],
["What are your thoughts on proprietary vs open source projects?"]
],
fill_height=True,
theme="Nymbo/Alyx_Theme",
title=None,
description=None
)
with gr.Row(elem_id="disclaimer-container"):
gr.Markdown(
f"""
*Powered by {MODEL_NAME}. Output may not always reflect my beliefs or be completely accurate. Additionally my viewpoints may change over time.*
""",
elem_id="disclaimer"
)
if __name__ == "__main__":
demo.queue().launch()