CoolFace
Apppublic

ItCodinTime/aoe-demo

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
streamlit_app.py242 linesDownload Raw Back to root
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()