Spaces:
Runtime error
Runtime error
Download streamlit_app.py from paddacoco/exaone-llm: direct link, hf CLI and curl.
- Browser
- Download file 11.9 kB
-
https://huggingface.co/spaces/paddacoco/exaone-llm/resolve/main/streamlit_app.py
- Command line
-
hf download hf://spaces/paddacoco/exaone-llm/streamlit_app.py
-
curl -L -o streamlit_app.py https://huggingface.co/spaces/paddacoco/exaone-llm/resolve/main/streamlit_app.py
11.9 kB
| 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") | |
| 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 | |
| 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'<div class="message-user"><strong>π€ You:</strong> {msg["content"]}</div>', unsafe_allow_html=True) | |
| else: | |
| st.markdown(f'<div class="message-assistant"><strong>π€ AI:</strong> {msg["content"]}</div>', 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κΈ° (ꡰ볡무 μ€ νμ΅) | |
| """) | |