website / app.py
frogggggmonkeeeeey's picture
Update app.py
3cb983f verified
Raw
History Blame Contribute Delete
4.5 kB
import sys
import traceback
import streamlit as st
# ์šฐ๋ฆฌ๊ฐ€ ๋งŒ๋“  ๋ฐฑ์—”๋“œ ํŒŒ์ผ์—์„œ ์—”์ง„ ํ•จ์ˆ˜ ์ˆ˜์ž…ํ•ด์˜ค๊ธฐ
from backend import load_all_models, predict_ensemble
# 1. ์ŠคํŠธ๋ฆผ๋ฆฟ UI ๊ธฐ๋ณธ ์„ค์ •
st.set_page_config(
page_title="๋ฌธ์ฒด ๊ธฐ๋ฐ˜ ๋™์ผ์ธ ํŒ๋ณ„ ์‹œ์Šคํ…œ",
page_icon="๐Ÿง ",
layout="centered"
)
# 2. ๋ฆฌ์†Œ์Šค๋ฅผ ํ•œ ๋ฒˆ๋งŒ ๋กœ๋“œํ•˜๋„๋ก ์บ์‹ฑ ์„ธํŒ…
@st.cache_resource(show_spinner=False)
def get_cached_models():
try:
return load_all_models()
except ModuleNotFoundError as e:
raise ModuleNotFoundError(
"ํ•„์ˆ˜ ํŒจํ‚ค์ง€๊ฐ€ ์„ค์น˜๋˜์ง€ ์•Š์•˜์Šต๋‹ˆ๋‹ค. requirements.txt๋ฅผ ํ™•์ธํ•˜์„ธ์š”.\n\n"
f"์›๋ณธ ์˜ค๋ฅ˜: {e}"
)
# 3. ๋ฉ”์ธ ํ™”๋ฉด ๋ ˆ์ด์•„์›ƒ ๊ทธ๋ฆฌ๊ธฐ
st.title("๐Ÿง  ๋ฌธ์ฒด ๊ธฐ๋ฐ˜ ๋™์ผ์ธ ํŒ๋ณ„ ์‹œ์Šคํ…œ")
st.markdown(
"""
์ด ์‹œ์Šคํ…œ์€ **KcBERT, DeBERTa, Multi-Scale Char-CNN ์•™์ƒ๋ธ” ๋ชจ๋ธ**์„ ํ™œ์šฉํ•˜์—ฌ
๋‘ ํ…์ŠคํŠธ์˜ ๋™์ผ ์ž‘์„ฑ์ž ๊ฐ€๋Šฅ์„ฑ์„ ๋ถ„์„ํ•ฉ๋‹ˆ๋‹ค.
"""
)
st.info("๋‘ ํ…์ŠคํŠธ๋ฅผ ์ž…๋ ฅํ•œ ๋’ค [ํŒ๋ณ„ํ•˜๊ธฐ] ๋ฒ„ํŠผ์„ ๋ˆ„๋ฅด๋ฉด ๋™์ผ ์ž‘์„ฑ์ž ๊ฐ€๋Šฅ์„ฑ์„ ์ถœ๋ ฅํ•ฉ๋‹ˆ๋‹ค.")
# ์ž…๋ ฅ ํ…์ŠคํŠธ ์ƒ์ž
text_a = st.text_area("ํ…์ŠคํŠธ A", height=180, placeholder="์ฒซ ๋ฒˆ์งธ ํ…์ŠคํŠธ๋ฅผ ์ž…๋ ฅํ•˜์„ธ์š”.")
text_b = st.text_area("ํ…์ŠคํŠธ B", height=180, placeholder="๋‘ ๋ฒˆ์งธ ํ…์ŠคํŠธ๋ฅผ ์ž…๋ ฅํ•˜์„ธ์š”.")
col1, col2 = st.columns(2)
with col1:
analyze_button = st.button("ํŒ๋ณ„ํ•˜๊ธฐ", use_container_width=True)
with col2:
clear_button = st.button("์ž…๋ ฅ ์ดˆ๊ธฐํ™” ์•ˆ๋‚ด", use_container_width=True)
if clear_button:
st.warning("์ž…๋ ฅ์ฐฝ ๋‚ด์šฉ์€ ์ง์ ‘ ์ง€์šฐ๊ฑฐ๋‚˜ ๋ธŒ๋ผ์šฐ์ € ์ƒˆ๋กœ๊ณ ์นจ์„ ํ•˜๋ฉด ์ดˆ๊ธฐํ™”๋ฉ๋‹ˆ๋‹ค.")
# 4. ํŒ๋ณ„ ๋ฒ„ํŠผ ํด๋ฆญ ์‹œ ๋ฐฑ์—”๋“œ ์—”์ง„ ๊ตฌ๋™
if analyze_button:
if not text_a.strip() or not text_b.strip():
st.warning("ํ…์ŠคํŠธ A์™€ ํ…์ŠคํŠธ B๋ฅผ ๋ชจ๋‘ ์ž…๋ ฅํ•˜์„ธ์š”.")
else:
try:
with st.spinner("AI ๋ชจ๋ธ๋กœ ๋ถ„์„ ์ค‘์ž…๋‹ˆ๋‹ค. ์ฒ˜์Œ ์‹คํ–‰ ์‹œ ์‹œ๊ฐ„์ด ๋‹ค์†Œ ๊ฑธ๋ฆด ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค."):
# ์บ์‹œ๋œ ๋ชจ๋ธ ๊ฐ€์ ธ์˜ค๊ธฐ
model_bundle = get_cached_models()
# ๋ฐฑ์—”๋“œ ์—ฐ์‚ฐ ํ•จ์ˆ˜ ํ˜ธ์ถœ
result = predict_ensemble(text_a, text_b, model_bundle)
st.divider()
st.subheader("๋ถ„์„ ๊ฒฐ๊ณผ")
same_percent = result["same_prob"] * 100
diff_percent = result["diff_prob"] * 100
# ๊ฒฐ๊ณผ์— ๋”ฐ๋ฅธ UI ์—ฐ์ถœ
if result["pred_label"] == 0:
st.success("ํŒ๋ณ„ ๊ฒฐ๊ณผ: ๋™์ผ ์ž‘์„ฑ์ž ๊ฐ€๋Šฅ์„ฑ์ด ๋†’์Šต๋‹ˆ๋‹ค.")
else:
st.error("ํŒ๋ณ„ ๊ฒฐ๊ณผ: ๋™์ผ ์ž‘์„ฑ์ž ๊ฐ€๋Šฅ์„ฑ์ด ๋‚ฎ์Šต๋‹ˆ๋‹ค.")
metric_col1, metric_col2 = st.columns(2)
with metric_col1:
st.metric("๋™์ผ ์ž‘์„ฑ์ž ๊ฐ€๋Šฅ์„ฑ", f"{same_percent:.2f}%")
with metric_col2:
st.metric("๋น„๋™์ผ ์ž‘์„ฑ์ž ๊ฐ€๋Šฅ์„ฑ", f"{diff_percent:.2f}%")
st.progress(min(max(result["same_prob"], 0.0), 1.0))
# ๋ชจ๋ธ๋ณ„ ๋””ํ…Œ์ผ ์ˆ˜์น˜ ํ‘œ ์ถœ๋ ฅ
st.subheader("๋ชจ๋ธ๋ณ„ ์˜ˆ์ธก ํ™•๋ฅ ")
model_table = {
"๋ชจ๋ธ": ["KcBERT", "Multi-Scale Char-CNN", "DeBERTa"],
"Same(0)": [
f"{result['kc_prob'][0] * 100:.2f}%",
f"{result['cnn_prob'][0] * 100:.2f}%",
f"{result['deberta_prob'][0] * 100:.2f}%"
],
"Diff(1)": [
f"{result['kc_prob'][1] * 100:.2f}%",
f"{result['cnn_prob'][1] * 100:.2f}%",
f"{result['deberta_prob'][1] * 100:.2f}%"
]
}
st.table(model_table)
# ํ™•์žฅ ํƒญ ์ •๋ณด
with st.expander("๊ธฐ์ˆ  ์ •๋ณด"):
st.write("์‹คํ–‰ ์žฅ์น˜:", result["device"])
st.write("Python ์‹คํ–‰ ๊ฒฝ๋กœ:")
st.code(sys.executable)
st.write("๋ชจ๋ธ ๊ฒฝ๋กœ:")
st.code(result["model_dir"])
st.caption("์ฃผ์˜: ๋ณธ ๊ฒฐ๊ณผ๋Š” AI ๋ชจ๋ธ ๊ธฐ๋ฐ˜ ์ถ”์ •๊ฐ’์ด๋ฉฐ, ์‹ค์ œ ๋™์ผ์ธ์„ ๋‹จ์ •ํ•˜๋Š” ๋ฒ•์  ํŒ๋‹จ์ด ์•„๋‹™๋‹ˆ๋‹ค.")
except Exception as e:
st.error("๊ณ ๊ธ‰ ๋ชจ๋ธ ์‹คํ–‰ ์ค‘ ์˜ค๋ฅ˜๊ฐ€ ๋ฐœ์ƒํ–ˆ์Šต๋‹ˆ๋‹ค.")
st.code(str(e))
with st.expander("์ƒ์„ธ ์˜ค๋ฅ˜ ๋ณด๊ธฐ"):
st.code(traceback.format_exc())