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'
πŸ‘€ You: {msg["content"]}
', unsafe_allow_html=True) else: st.markdown(f'
πŸ€– AI: {msg["content"]}
', 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κΈ° (ꡰ볡무 쀑 ν•™μŠ΅) """)