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())