import streamlit as st import torch from peft import PeftConfig, PeftModel from transformers import AutoModelForCausalLM, AutoTokenizer import time import re from typing import Tuple from datasets import load_dataset import psutil st.set_page_config(page_title="LG AIMERS EXAONE 4.0 GSM8K", layout="wide") @st.cache_resource def load_model(): st.info("π μ΅μ LoRA λ‘λ© μ€...") REPO_ID = "paddacoco/exaone4_gsm8k2" SUBFOLDER = "exaone4_gsm8k/final_exaone4_gsm8k" try: # 1. LoRA μ€μ λ‘λ config = PeftConfig.from_pretrained(REPO_ID, subfolder=SUBFOLDER) st.write(f"β LoRA r={config.r}, target_modules={config.target_modules}") # 2. EXAONE 4.0 λ² μ΄μ€ base_model = AutoModelForCausalLM.from_pretrained( "LGAI-EXAONE/EXAONE-4.0-1.2B", torch_dtype=torch.float32, device_map="cpu", trust_remote_code=True ) # 3. LoRA μ΄λν° κ²°ν© model = PeftModel.from_pretrained(base_model, REPO_ID, subfolder=SUBFOLDER) # 4. ν ν¬λμ΄μ (LoRAμμ) tokenizer = AutoTokenizer.from_pretrained(REPO_ID, subfolder=SUBFOLDER, trust_remote_code=True) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token st.success("π LG AIMERS LoRA μμ λ‘λ!") st.metric("LoRA ν¬κΈ°", "43.7MB") st.metric("λ² μ΄μ€", "EXAONE 4.0 1.2B") return model, tokenizer except Exception as e: st.error(f"β {e}") return None, None def extract_math_answer(text: str) -> str: """μν λ΅λ³ μΆμΆ""" if "####" in text: answer_part = text.split("####")[-1].strip() numbers = re.findall(r"\d+", answer_part) return numbers[0] if numbers else answer_part[:50] numbers = re.findall(r"\d+", text) return numbers[-1] if numbers else "λ΅ μμ" def generate_response(model, tokenizer, prompt: str, max_length: int = 256, temperature: float = 0.7, top_p: float = 0.9) -> Tuple[str, float]: """λͺ¨λΈ μλ΅ μμ±""" if model is None or tokenizer is None: return "λͺ¨λΈ λ‘λ μ€ν¨", 0.0 start_time = time.time() try: # μν λͺ¨λ μ΅μ ν ν둬ννΈ if "μν" in st.session_state.get("mode", ""): system_prompt = """μν λ¬Έμ λ₯Ό λ¨κ³λ³λ‘ μ νν νμ΄μ£ΌμΈμ. λ¨μμ κ³μ° κ³Όμ μ λͺ νν μ€λͺ νμΈμ. μ΅μ’ λ΅μ #### λ€μ μ«μλ§ μ μΌμΈμ.""" full_prompt = f"{system_prompt}\n\nλ¬Έμ : {prompt}" else: full_prompt = prompt chat = [{'role': 'user', 'content': full_prompt}] formatted = tokenizer.apply_chat_template(chat, tokenize=False, add_generation_prompt=True) inputs = tokenizer(formatted, return_tensors="pt", truncation=True, max_length=1024) with torch.no_grad(): outputs = model.generate( **inputs, max_new_tokens=max_length, temperature=temperature, top_p=top_p, do_sample=True, pad_token_id=tokenizer.eos_token_id, repetition_penalty=1.1 ) response = tokenizer.decode(outputs[0][inputs['input_ids'].shape[1]:], skip_special_tokens=True) return response.strip(), time.time() - start_time except Exception as e: return f"μμ± μλ¬: {str(e)}", 0.0 @st.cache_data def evaluate_gsm8k(n_samples=100): """GSM8K λ²€μΉλ§ν¬ νκ° (ν΄μ»€ν€μ©)""" try: st.info(f"π GSM8K {n_samples}κ° μν νκ° μ€... (μκ° κ±Έλ¦Ό)") dataset = load_dataset("openai/gsm8k", "main")['test'].select(range(min(n_samples, 100))) correct = 0 total_time = 0 results = [] progress_bar = st.progress(0) status_text = st.empty() for i, example in enumerate(dataset): prompt = example['question'] start = time.time() response, _ = generate_response(st.session_state.model, st.session_state.tokenizer, prompt, max_length=128) elapsed = time.time() - start total_time += elapsed # κ°λ¨ μ«μ μΆμΆ pred = extract_math_answer(response) ground_truth = example['answer'].split('####')[-1].strip() is_correct = pred == ground_truth if is_correct: correct += 1 results.append({ 'question': prompt[:60], 'pred': pred, 'gt': ground_truth, 'correct': 'β ' if is_correct else 'β' }) progress_bar.progress((i + 1) / len(dataset)) status_text.text(f"νκ° μ§ν μ€... {i+1}/{len(dataset)} ({int(correct/(i+1)*100)}%)") progress_bar.empty() status_text.empty() accuracy = correct / len(dataset) * 100 avg_time = total_time / len(dataset) return { 'accuracy': accuracy, 'avg_time': avg_time, 'results': results } except Exception as e: st.error(f"νκ° μ€ν¨: {str(e)}") return None def get_memory_usage(): """νμ¬ λ©λͺ¨λ¦¬ μ¬μ©λ (MB)""" process = psutil.Process() return process.memory_info().rss / 1024**2 # μ± μμ if "model" not in st.session_state: st.session_state.model, st.session_state.tokenizer = load_model() model = st.session_state.model tokenizer = st.session_state.tokenizer if model: st.balloons() else: st.stop() # λ©μΈ UI st.markdown(""" # π€ EXAONE 1.2B LLM λμ보λ ### LoRA κ²½λν λͺ¨λΈ - **LG AIMERS 8th Cohort ν΄μ»€ν€** """) # μ¬μ΄λλ° μ€μ with st.sidebar: st.markdown("## βοΈ μ€μ ") st.session_state.mode = st.radio("π λͺ¨λ", ["μΌλ° μ±ν ", "μν νμ΄"], key="mode_selector") st.markdown("---") st.session_state.temperature = st.slider("π‘οΈ Temperature", 0.1, 2.0, 0.7, 0.1) st.session_state.top_p = st.slider("π― Top-P", 0.1, 1.0, 0.9, 0.05) st.session_state.max_length = st.slider("π κΈΈμ΄", 50, 512, 256, 50) st.markdown("---") col1, col2 = st.columns(2) with col1: st.metric("λͺ¨λΈ", "EXAONE 1.2B") with col2: st.metric("μ΅μ ν", "LoRA 43.7MB") st.markdown("---") eval_mode = st.checkbox("π GSM8K λ²€μΉλ§ν¬ (μκ° μμ)") if eval_mode: if st.button("π νκ° μμ", key="eval_button"): benchmark_results = evaluate_gsm8k(n_samples=50) if benchmark_results: st.success(f"β μ νλ: {benchmark_results['accuracy']:.1f}%") st.info(f"β±οΈ νκ· μκ°: {benchmark_results['avg_time']:.2f}s") # μν κ²°κ³Ό νμ st.markdown("#### νκ° μν (μμ 5κ°)") for result in benchmark_results['results'][:5]: st.write(f"{result['correct']} Q: {result['question']}...") st.write(f" μμΈ‘: {result['pred']} | μ λ΅: {result['gt']}") # ν UI tab1, tab2, tab3 = st.tabs(["π¬ μ±ν ", "π ν΅κ³", "π ν΄μ»€ν€ λΆμ"]) with tab1: st.markdown("### μ±λ΄κ³Ό λννμΈμ") if "messages" not in st.session_state: st.session_state.messages = [] # μ΄μ λ©μμ§ νμ for msg in st.session_state.messages: if msg["role"] == "user": st.markdown(f'
', unsafe_allow_html=True) else: st.markdown(f'', unsafe_allow_html=True) st.markdown("---") # μ λ ₯ UI col1, col2 = st.columns([4, 1]) with col1: user_input = st.text_input("λ©μμ§:", placeholder="μ§λ¬Έμ μ λ ₯νμΈμ", key="user_input") with col2: send_button = st.button("μ μ‘", use_container_width=True) # μ μ‘ μ²λ¦¬ if send_button and user_input: st.session_state.messages.append({"role": "user", "content": user_input}) with st.spinner("π€ μλ΅ μμ± μ€..."): response, processing_time = generate_response( model, tokenizer, user_input, st.session_state.max_length, st.session_state.temperature, st.session_state.top_p ) st.session_state.messages.append({"role": "assistant", "content": response}) # λ©νΈλ¦ νμ col1, col2, col3 = st.columns(3) with col1: st.metric("κΈΈμ΄", f"{len(response)} μ") with col2: st.metric("μκ°", f"{processing_time:.2f}μ΄") with col3: if "μν" in st.session_state.mode: answer = extract_math_answer(response) st.metric("λ΅", answer) st.rerun() if not st.session_state.messages: st.info("π EXAONE 1.2Bμ μ€μ κ²μ νμν©λλ€!\nμν λ¬Έμ λ μΌλ° μ§λ¬Έμ ν΄λ³΄μΈμ.") with tab2: st.markdown("### π λν ν΅κ³") if st.session_state.messages: assistant_msgs = [m for m in st.session_state.messages if m["role"] == "assistant"] col1, col2, col3 = st.columns(3) with col1: st.metric("λν μ", len(st.session_state.messages) // 2) with col2: avg_length = sum(len(m["content"]) for m in assistant_msgs) / len(assistant_msgs) if assistant_msgs else 0 st.metric("νκ· κΈΈμ΄", f"{int(avg_length)} μ") with col3: total_length = sum(len(m["content"]) for m in assistant_msgs) st.metric("μ΄ μμ±", f"{total_length:,} μ") else: st.info("μμ§ λνκ° μμ΅λλ€.") with tab3: st.markdown("### π LG AIMERS ν΄μ»€ν€ λΆμ") st.markdown("#### π κ²½λν ν¨κ³Ό") comparison_data = { "νλͺ©": ["μ 체 νλΌλ―Έν°", "νμ΅ νλΌλ―Έν°", "μΆλ‘ λ©λͺ¨λ¦¬", "νμ΅ VRAM"], "Base Model": ["1.2B", "1.2B", "8-10 GB", "20-30 GB β"], "LoRA (λΉμ )": ["1.2B", "11.4M (0.95%)", "9-11 GB", "T4 16GB β "] } st.table(comparison_data) st.markdown("#### πΎ λ©λͺ¨λ¦¬ ν¨μ¨μ±") col1, col2, col3 = st.columns(3) with col1: st.metric("νμ¬ λ©λͺ¨λ¦¬", f"{get_memory_usage():.1f} MB") with col2: st.metric("νλΌλ―Έν° μ μ½", "99.05% β") with col3: st.metric("νμ΅ μκ°", "~2μκ° (T4)") st.markdown("#### π― μμ ν¬μΈνΈ") st.success(""" β **99.1% νλΌλ―Έν° μ μ½** (1.2B β 11.4M νμ΅) β **T4 GPU νΈν** (16GB λ΄ νμ΅/μΆλ‘ μ±κ³΅) β **μ€μκ° Streamlit λ°λͺ¨** (μν + μ±ν λͺ¨λ) β **λ²€μΉλ§ν¬ μΈ‘μ ** (GSM8K μ νλ μΆμ ) β **λ©λͺ¨λ¦¬ vs Accuracy μ΅μ ν** μ¦λͺ """) st.markdown("#### π κ°μ λ‘λλ§΅") st.info(""" 1οΈβ£ **QLoRA** - 4bit μμνλ‘ μΆκ° 50% λ©λͺ¨λ¦¬ μ μ½ 2οΈβ£ **Layer-wise LoRA** - μ£Όμ μΈ΅λ§ fine-tuning (5-6M νλΌλ―Έν°) 3οΈβ£ **DPO + SFT** - μ νλ μΆκ° 3-5% κ°μ """) # νΈν° st.markdown("---") st.markdown(""" ### π LG AIMERS 8th Cohort ν΄μ»€ν€ νλ‘μ νΈ **EXAONE 4.0 1.2B + GSM8K LoRA μ΅μ ν** - λͺ¨λΈ: LGAI-EXAONE/EXAONE-4.0-1.2B - κΈ°λ²: LoRA (Low-Rank Adaptation) - λ°μ΄ν°: GSM8K (μν λ¬Έμ ) - νλ«νΌ: Streamlit + HuggingFace - κ°λ°μ: LG AIMERS 8κΈ° (ꡰ볡무 μ€ νμ΅) """)