Spaces:
Sleeping
Sleeping
File size: 4,503 Bytes
cf8541f 783a1e9 6e47160 783a1e9 6e47160 783a1e9 6e47160 783a1e9 cf8541f 783a1e9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 |
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())
|