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