Spaces:
Sleeping
Sleeping
| # app.py | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # μ±λ΄ μ€ν λ©μΈ μ± (SQLite μ κ±° λ²μ ). | |
| # νλ¦: intro β pre(μ¬μ μ€λ¬Έ) β chat(λν) β post(μ¬νμ€λ¬Έ) β done | |
| # μ μ₯: storage.pyμ μ΄μ€ μμ λ§(ꡬκΈμνΈ + λ©λͺ¨λ¦¬). DB νμΌμ΄ μλ€. | |
| # μ¬μ΄λλ°: (μ) μ°κ΅¬ μλ΄ μ€λͺ / (μλ) λ°μ΄ν° CSV λ€μ΄λ‘λ λ°±μ | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| import uuid | |
| import time | |
| import streamlit as st | |
| import storage | |
| from prompts import SYSTEM_PROMPTS, TASK_INSTRUCTION | |
| from llm import stream_response | |
| st.set_page_config(page_title="μ§λ‘ μ±λ΄ λν μ°κ΅¬", page_icon="π¬") | |
| # ===== μΈμ μν μ΄κΈ°ν ===== | |
| # Streamlitμ λ§€ μνΈμμ©λ§λ€ μ€ν¬λ¦½νΈλ₯Ό μβμλλ‘ λ€μ μ€ννλ€. | |
| # λ°λΌμ 'μ§κΈ μ΄λ λ¨κ³μΈκ°'λ₯Ό session_stateμ μ μ₯ν΄ λ¬μΌ νλ€. | |
| if "stage" not in st.session_state: | |
| st.session_state.stage = "intro" # νμ¬ λ¨κ³ | |
| st.session_state.participant_id = str(uuid.uuid4()) # μ΅λͺ μλ³μ(μ΄λ¦Β·μ΄λ©μΌ μ λ°μ) | |
| st.session_state.condition = None # 'A' λλ 'B' | |
| st.session_state.messages = [] # νλ©΄ νμμ© λν κΈ°λ‘ | |
| st.session_state.turn = 0 # λν ν΄ μ | |
| storage.init_state(st.session_state) # λμ 리μ€νΈ 3κ° μ€λΉ | |
| # ===== μ¬μ΄λλ°: μμͺ½ μλ΄ + μλμͺ½ λ€μ΄λ‘λ ===== | |
| def render_sidebar(): | |
| with st.sidebar: | |
| # ββ (μ) μ°κ΅¬ μλ΄ μ€λͺ ββ | |
| st.header("π μ°κ΅¬ μλ΄") | |
| st.markdown( | |
| "- **μμμκ°**: μ½ 10~15λΆ\n" | |
| "- **μ μ°¨**: μ¬μ μ€λ¬Έ β μ±λ΄ λν(5ν΄+) β μ¬ν μ€λ¬Έ\n" | |
| "- **μ΅λͺ μ±**: μ΄λ¦Β·μ΄λ©μΌμ μμ§νμ§ μμ΅λλ€.\n" | |
| "- **μ μ**: μ±λ΄μ AIμ΄λ©°, μ λ¬Έ μλ¬Έμ΄ νμνλ©΄ μ λ¬Έκ°μκ² λ³λ λ¬ΈμνμΈμ." | |
| ) | |
| st.divider() | |
| # ββ (μλ) λ°μ΄ν° λ€μ΄λ‘λ λ°±μ ββ | |
| # μμ λ§ 2: HF Dataset μ λ‘λκ° μ€ν¨ν΄λ μ΄ λ©λͺ¨λ¦¬ λμ λΆμΌλ‘ νμ κ°λ₯. | |
| # μ£Όμ: session_stateλ 'νμ¬ λΈλΌμ°μ μΈμ 'μ λ°μ΄ν°λ§ λ΄λλ€. | |
| # μ 체 μ°Έμ¬μ ν΅ν©λ³Έμ HF Datasetμμ λ°λλ€(μμ λ§ 1). | |
| st.header("β¬οΈ λ°μ΄ν° λ€μ΄λ‘λ") | |
| st.caption("μ΄ μΈμ μ μμΈ κΈ°λ‘(λ°±μ μ©). μ 체 ν΅ν©λ³Έμ HF Dataset μ°Έμ‘°.") | |
| n_p = len(st.session_state["rows_participants"]) | |
| n_m = len(st.session_state["rows_messages"]) | |
| n_s = len(st.session_state["rows_surveys"]) | |
| st.download_button("μ°Έμ¬μ CSV", storage.to_csv_bytes(st.session_state["rows_participants"]), | |
| "participants.csv", "text/csv", disabled=(n_p == 0)) | |
| st.download_button("λν CSV", storage.to_csv_bytes(st.session_state["rows_messages"]), | |
| "messages.csv", "text/csv", disabled=(n_m == 0)) | |
| st.download_button("μ€λ¬Έ CSV", storage.to_csv_bytes(st.session_state["rows_surveys"]), | |
| "surveys.csv", "text/csv", disabled=(n_s == 0)) | |
| st.caption(f"μ°Έμ¬μ {n_p} Β· λν {n_m}μ€ Β· μ€λ¬Έ {n_s}μ€") | |
| render_sidebar() | |
| # ===== Stage 1: μΈνΈλ‘ + λμ ===== | |
| if st.session_state.stage == "intro": | |
| st.title("AI μ§λ‘μλ΄ μ±λ΄ λν μ°κ΅¬") | |
| st.markdown( | |
| "μ‘Έμ ν μ§λ‘μ λν΄ μ±λ΄κ³Ό λννλ μ°κ΅¬μ λλ€. " | |
| "μ¬μ μ€λ¬Έ β λν β μ¬ν μ€λ¬Έ μμΌλ‘ μ§νλ©λλ€." | |
| ) | |
| if st.button("λμνκ³ μμνκΈ°", type="primary"): | |
| # μ°Έμ¬μλ³ λλ€ λ°°μ (μΈμ μ΄ λΆλ¦¬λΌ μμ΄ κ΅λ λ°°μ μ΄ λΆκ°νλ―λ‘) | |
| import random | |
| cond = random.choice(["A", "B"]) | |
| st.session_state.condition = cond | |
| storage.add_participant(st.session_state, st.session_state.participant_id, cond) | |
| st.session_state.stage = "pre" | |
| st.rerun() | |
| # ===== Stage 2: μ¬μ μ‘°μ¬ ===== | |
| elif st.session_state.stage == "pre": | |
| st.title("μ¬μ μ€λ¬Έ") | |
| with st.form("pre_survey"): | |
| year = st.selectbox("νλ ", ["1νλ ", "2νλ ", "3νλ ", "4νλ ", "μ‘Έμ μ μ/ν΄ν", "λνμμ"]) | |
| major = st.selectbox("μ 곡 κ³μ΄", ["μΈλ¬Έ", "μ¬ν", "μκ²½", "곡ν", "μμ°", "μμ½", "μ체λ₯", "κΈ°ν"]) | |
| d_stage = st.selectbox("νμ¬ μ§λ‘ κ²°μ λ¨κ³", | |
| ["μμ§ νμ μ€", "λ°©ν₯μ΄ μ΄λ μ λ μ‘ν", "ꡬ체μ κ³ν μ립 μ€", "κ±°μ νμ "]) | |
| anxiety = st.slider("μ΅κ·Ό μ§λ‘ λΆμ μ λ (1: μ ν μμ ~ 7: λ§€μ° νΌ)", 1, 7, 4) | |
| ai_freq = st.selectbox("AI μ±λ΄ μ¬μ© λΉλ", ["κ±°μ μ μ", "μ 1~2ν", "μ£Ό 1~2ν", "κ±°μ λ§€μΌ"]) | |
| if st.form_submit_button("λ€μ"): | |
| pid, cond = st.session_state.participant_id, st.session_state.condition | |
| for qid, val in [("year", year), ("major", major), ("decision_stage", d_stage), | |
| ("career_anxiety", anxiety), ("ai_usage", ai_freq)]: | |
| storage.add_survey(st.session_state, pid, cond, "pre", qid, val) | |
| st.session_state.stage = "chat" | |
| st.rerun() | |
| # ===== Stage 3: λν ===== | |
| elif st.session_state.stage == "chat": | |
| st.title("μ§λ‘ μ±λ΄κ³Ό λννκΈ°") | |
| st.info(TASK_INSTRUCTION) | |
| st.caption(f"μ΅μ 5ν΄ μ΄μ λν ν μλ λ²νΌμ λλ₯΄μΈμ. (νμ¬ {st.session_state.turn}ν΄)") | |
| # μ΄μ λ©μμ§ λ€μ 그리기(λ§€ μ€νλ§λ€ νλ©΄μ μλ‘ κ·Έλ¦¬λ―λ‘ νμ) | |
| for m in st.session_state.messages: | |
| with st.chat_message(m["role"]): | |
| st.markdown(m["content"]) | |
| if prompt := st.chat_input("λ©μμ§λ₯Ό μ λ ₯νμΈμ"): | |
| pid, cond = st.session_state.participant_id, st.session_state.condition | |
| st.session_state.turn += 1 | |
| # 1) μ¬μ©μ λ©μμ§: νλ©΄ νμ + μ μ₯ | |
| st.session_state.messages.append({"role": "user", "content": prompt}) | |
| storage.add_message(st.session_state, pid, cond, st.session_state.turn, "user", prompt) | |
| with st.chat_message("user"): | |
| st.markdown(prompt) | |
| # 2) μ±λ΄ μλ΅: μ€νΈλ¦¬λ°μΌλ‘ ν μ²ν¬μ© νμ | |
| with st.chat_message("assistant"): | |
| placeholder = st.empty() | |
| full, t0 = "", time.time() | |
| try: | |
| for delta in stream_response(SYSTEM_PROMPTS[cond], st.session_state.messages): | |
| full += delta | |
| placeholder.markdown(full + "β") # 컀μ ν¨κ³Ό | |
| placeholder.markdown(full) | |
| except Exception as e: | |
| full = f"[μ€λ₯κ° λ°μνμ΅λλ€: {e}]" | |
| placeholder.markdown(full) | |
| latency = int((time.time() - t0) * 1000) | |
| st.session_state.messages.append({"role": "assistant", "content": full}) | |
| storage.add_message(st.session_state, pid, cond, st.session_state.turn, | |
| "assistant", full, latency) | |
| # 5ν΄ μ΄μμ΄λ©΄ μ¬νμ€λ¬ΈμΌλ‘ λμ΄κ°λ λ²νΌ λ ΈμΆ | |
| if st.session_state.turn >= 5: | |
| if st.button("λν μ’ λ£νκ³ μ¬ν μ€λ¬ΈμΌλ‘", type="primary"): | |
| st.session_state.stage = "post" | |
| st.rerun() | |
| # ===== Stage 4: μ¬νμ‘°μ¬ ===== | |
| elif st.session_state.stage == "post": | |
| st.title("μ¬ν μ€λ¬Έ") | |
| with st.form("post_survey"): | |
| usefulness = st.slider("λ°μ μ λ³΄κ° μ μ©νλ€ (1~7)", 1, 7, 4) | |
| warmth = st.slider("μ±λ΄μ΄ λ°λ»νλ€ (1~7)", 1, 7, 4) | |
| competence = st.slider("μ±λ΄μ΄ μ λ₯νλ€ (1~7)", 1, 7, 4) | |
| trust = st.slider("μ±λ΄μ μ λ’°νλ€ (1~7)", 1, 7, 4) | |
| clarity = st.slider("λν ν μ§λ‘ κ³ λ―Όμ΄ λ μ 리λμλ€ (1~7)", 1, 7, 4) | |
| recommend = st.slider("μΉκ΅¬μκ² μΆμ²ν μν₯ (1~7)", 1, 7, 4) | |
| free_text = st.text_area("μμ μ견: κ°μ₯ λμμ΄ λ μ , μμ¬μ λ μ (μ ν)") | |
| if st.form_submit_button("μ μΆ"): | |
| pid, cond = st.session_state.participant_id, st.session_state.condition | |
| for qid, val in [("usefulness", usefulness), ("warmth", warmth), | |
| ("competence", competence), ("trust", trust), | |
| ("clarity", clarity), ("recommend", recommend), | |
| ("free_text", free_text)]: | |
| storage.add_survey(st.session_state, pid, cond, "post", qid, val) | |
| # μλ£ νμ(λ©λͺ¨λ¦¬): ν΄λΉ μ°Έμ¬μ νμ completedλ₯Ό 1λ‘ | |
| for r in st.session_state["rows_participants"]: | |
| if r["participant_id"] == pid: | |
| r["completed"] = 1 | |
| # μμ λ§ 1: μ΄λ² μΈμ λμ λΆμ HF Datasetμ μΌκ΄ μ λ‘λ | |
| # (μλ£ μμ μ ν λ²λ§. μ€ν¨ν΄λ λ©λͺ¨λ¦¬ λ°±μ μ΄ λ¨μΌλ―λ‘ νλ¦μ κ³μ) | |
| storage.push_to_dataset(st.session_state) | |
| st.session_state.stage = "done" | |
| st.rerun() | |
| # ===== Stage 5: μλ£ ===== | |
| elif st.session_state.stage == "done": | |
| st.title("μ°Έμ¬ν΄μ£Όμ μ κ°μ¬ν©λλ€ π") | |
| st.markdown(f"μ°Έμ¬μ ID: `{st.session_state.participant_id}`") | |
| st.caption("μλ΅μ΄ μμ νκ² κΈ°λ‘λμμ΅λλ€.") | |