Download app.py from mmcgovern574/DigitalTwin-Mistral-Small-24B: direct link, hf CLI and curl.
- Browser
- Download file 7.36 kB
-
https://huggingface.co/spaces/mmcgovern574/DigitalTwin-Mistral-Small-24B/resolve/2efcfde83309988d13879e0dc707bcb08f4573d5/app.py
- Command line
-
hf download hf://spaces/mmcgovern574/DigitalTwin-Mistral-Small-24B@2efcfde83309988d13879e0dc707bcb08f4573d5/app.py
-
curl -L -o app.py https://huggingface.co/spaces/mmcgovern574/DigitalTwin-Mistral-Small-24B/resolve/2efcfde83309988d13879e0dc707bcb08f4573d5/app.py
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() | |
| 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() |