ItCodinTime/aoe-demo
0
1import streamlit as st2import torch3import os4from transformers import AutoTokenizer, AutoModelForCausalLM5import traceback6from typing import Optional7 8# Configure the page9st.set_page_config(10 page_title="LLM Comparison: GPT-4 vs Gemini vs AOE",11 page_icon="⚔️",12 layout="wide"13)14 15def load_aoe_model():16 """Load the AoE model and tokenizer from outputs/student/ directory"""17 model_path = "outputs/student/"18 19 try:20 if not os.path.exists(model_path):21 st.error(f"Model directory '{model_path}' not found. Please ensure the model files are present.")22 return None, None23 24 # Check if required files exist25 required_files = ["config.json", "pytorch_model.bin", "tokenizer.json"]26 missing_files = [f for f in required_files if not os.path.exists(os.path.join(model_path, f))]27 28 if missing_files:29 st.warning(f"Some model files may be missing: {missing_files}. Attempting to load anyway...")30 31 # Load tokenizer and model32 tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)33 model = AutoModelForCausalLM.from_pretrained(34 model_path, 35 torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,36 device_map="auto" if torch.cuda.is_available() else None,37 trust_remote_code=True38 )39 40 return model, tokenizer41 42 except Exception as e:43 st.error(f"Error loading AoE model: {str(e)}")44 st.text(f"Traceback: {traceback.format_exc()}")45 return None, None46 47def generate_aoe_response(model, tokenizer, prompt, max_length=512):48 """Generate response from the AoE model"""49 try:50 # Tokenize input51 inputs = tokenizer.encode(prompt, return_tensors="pt")52 53 # Move to same device as model if CUDA is available54 if torch.cuda.is_available() and next(model.parameters()).is_cuda:55 inputs = inputs.cuda()56 57 # Generate response58 with torch.no_grad():59 outputs = model.generate(60 inputs,61 max_length=len(inputs[0]) + max_length,62 num_return_sequences=1,63 temperature=0.7,64 do_sample=True,65 pad_token_id=tokenizer.eos_token_id66 )67 68 # Decode response69 response = tokenizer.decode(outputs[0], skip_special_tokens=True)70 71 # Remove the input prompt from the response72 if response.startswith(prompt):73 response = response[len(prompt):].strip()74 75 return response76 77 except Exception as e:78 return f"Error generating AoE response: {str(e)}"79 80def query_gpt4_api(prompt: str) -> str:81 """Query GPT-4 API using environment variable for API key"""82 api_key = os.getenv('OPENAI_API_KEY')83 84 if not api_key:85 return "❌ GPT-4 API key not found in environment variables. Please set OPENAI_API_KEY environment variable to use GPT-4."86 87 try:88 # This is a placeholder implementation - would need actual OpenAI API integration89 return "🤖 GPT-4 response would appear here with proper API configuration."90 except Exception as e:91 return f"Error querying GPT-4: {str(e)}"92 93def query_gemini_api(prompt: str) -> str:94 """Query Gemini API using environment variable for API key"""95 api_key = os.getenv('GOOGLE_API_KEY')96 97 if not api_key:98 return "❌ Gemini API key not found in environment variables. Please set GOOGLE_API_KEY environment variable to use Gemini."99 100 try:101 # This is a placeholder implementation - would need actual Google Gemini API integration102 return "🤖 Gemini response would appear here with proper API configuration."103 except Exception as e:104 return f"Error querying Gemini: {str(e)}"105 106def main():107 st.title("⚔️ LLM Comparison: GPT-4 vs Gemini vs AOE")108 st.markdown("Compare responses from three different language models side by side.")109 110 # Initialize session state for model caching111 if 'aoe_model' not in st.session_state:112 st.session_state.aoe_model = None113 st.session_state.aoe_tokenizer = None114 st.session_state.aoe_loaded = False115 116 # Load AOE model on first run117 if not st.session_state.aoe_loaded:118 with st.spinner("Loading AOE model from outputs/student/..."):119 model, tokenizer = load_aoe_model()120 if model is not None and tokenizer is not None:121 st.session_state.aoe_model = model122 st.session_state.aoe_tokenizer = tokenizer123 st.session_state.aoe_loaded = True124 st.success("✅ AOE model loaded successfully!")125 else:126 st.error("❌ Failed to load AOE model. Check error messages above.")127 128 # Configuration section129 st.markdown("---")130 st.subheader("🔧 Configuration")131 132 col1, col2 = st.columns(2)133 134 with col1:135 max_length = st.slider(136 "Max Response Length",137 min_value=100,138 max_value=1000,139 value=512,140 step=50,141 help="Maximum length for generated responses"142 )143 144 with col2:145 # Display API key status146 openai_key_status = "✅ Found" if os.getenv('OPENAI_API_KEY') else "❌ Missing"147 google_key_status = "✅ Found" if os.getenv('GOOGLE_API_KEY') else "❌ Missing"148 149 st.info(f"**API Key Status:**\n\nOpenAI API Key: {openai_key_status}\n\nGoogle API Key: {google_key_status}")150 151 # Main comparison interface152 st.markdown("---")153 st.subheader("💬 Compare LLM Responses")154 155 # User input156 user_prompt = st.text_area(157 "Enter your prompt:",158 placeholder="Type your prompt here to compare responses from all three models...",159 height=120,160 help="Enter a prompt to see how different LLMs respond"161 )162 163 # Generate responses button164 if st.button("🚀 Generate All Responses", type="primary"):165 if not user_prompt.strip():166 st.warning("Please enter a prompt first.")167 else:168 # Create three columns for side-by-side comparison169 col1, col2, col3 = st.columns(3)170 171 with col1:172 st.markdown("### 🤖 GPT-4")173 with st.spinner("Generating GPT-4 response..."):174 gpt4_response = query_gpt4_api(user_prompt)175 st.markdown("**Response:**")176 st.write(gpt4_response)177 178 with col2:179 st.markdown("### 🌟 Gemini")180 with st.spinner("Generating Gemini response..."):181 gemini_response = query_gemini_api(user_prompt)182 st.markdown("**Response:**")183 st.write(gemini_response)184 185 with col3:186 st.markdown("### 🏰 AOE (Local)")187 if st.session_state.aoe_loaded:188 with st.spinner("Generating AOE response..."):189 aoe_response = generate_aoe_response(190 st.session_state.aoe_model,191 st.session_state.aoe_tokenizer,192 user_prompt,193 max_length194 )195 st.markdown("**Response:**")196 st.write(aoe_response)197 else:198 st.error("AOE model not loaded. Please reload the page.")199 200 # Model information sidebar201 with st.sidebar:202 st.header("ℹ️ Model Information")203 204 st.markdown("**🤖 GPT-4**")205 openai_status = "✅ Configured" if os.getenv('OPENAI_API_KEY') else "❌ Environment variable OPENAI_API_KEY not set"206 st.write(f"Status: {openai_status}")207 st.write("Provider: OpenAI")208 209 st.markdown("**🌟 Gemini**")210 google_status = "✅ Configured" if os.getenv('GOOGLE_API_KEY') else "❌ Environment variable GOOGLE_API_KEY not set"211 st.write(f"Status: {google_status}")212 st.write("Provider: Google")213 214 st.markdown("**🏰 AOE (Local)**")215 st.write(f"Status: {'✅ Loaded' if st.session_state.aoe_loaded else '❌ Not loaded'}")216 st.write("Path: outputs/student/")217 if st.session_state.aoe_loaded:218 try:219 device_info = f"Device: {next(st.session_state.aoe_model.parameters()).device}"220 st.write(device_info)221 except:222 pass223 224 if st.button("🔄 Reload AOE Model"):225 st.session_state.aoe_loaded = False226 st.experimental_rerun()227 228 st.markdown("---")229 st.markdown("**📋 Instructions:**")230 st.markdown("1. Set OPENAI_API_KEY and GOOGLE_API_KEY environment variables")231 st.markdown("2. Enter your prompt in the text area")232 st.markdown("3. Click 'Generate All Responses'")233 st.markdown("4. Compare responses side by side")234 235 st.markdown("---")236 st.markdown("**⚠️ Notes:**")237 st.markdown("- GPT-4 and Gemini require valid API keys in environment variables")238 st.markdown("- AOE model runs locally from outputs/student/")239 st.markdown("- Responses are generated independently")240 241if __name__ == "__main__":242 main()