frogggggmonkeeeeey's picture
Update app.py
74ca577 verified
Raw
History Blame Contribute Delete
7.92 kB
import os
import pickle
import sys
import traceback
import numpy as np
import streamlit as st
import torch
import torch.nn.functional as F
from transformers import AutoModelForSequenceClassification, AutoTokenizer
# --- [1] ํŒŒ์ผ ๋ฐ ํด๋” ๊ฒฝ๋กœ ์„ค์ • ---
BASE_DIR = os.path.dirname(__file__)
KC_DIR = os.path.join(BASE_DIR, "kcbert_web")
DEBERTA_DIR = os.path.join(BASE_DIR, "deberta_web")
CNN_PATH = os.path.join(BASE_DIR, "char_cnn_web.pt")
VOCAB_PATH = os.path.join(BASE_DIR, "vocab.pkl")
# --- [2] ๋ชจ๋ธ ๋กœ๋“œ ํ•จ์ˆ˜ (์„œ๋ฒ„ ๊ธฐ๋™ ์‹œ ๋”ฑ ํ•œ ๋ฒˆ ์‹คํ–‰) ---
@st.cache_resource(show_spinner=False)
def load_verification_models():
try:
# 1. KcBERT ๋กœ๋“œ
kc_tokenizer = AutoTokenizer.from_pretrained(KC_DIR)
kc_model = AutoModelForSequenceClassification.from_pretrained(KC_DIR)
kc_model.eval()
# 2. DeBERTa ๋กœ๋“œ
deberta_tokenizer = AutoTokenizer.from_pretrained(DEBERTA_DIR)
deberta_model = AutoModelForSequenceClassification.from_pretrained(DEBERTA_DIR)
deberta_model.eval()
# 3. Char-CNN ๋ฐ ๋ณด์นด ๋กœ๋“œ
vocab = None
if os.path.exists(VOCAB_PATH):
with open(VOCAB_PATH, "rb") as f:
vocab = pickle.load(f)
cnn_model = None # ์ถ”ํ›„ ๊ณ ์œ  CNN ๊ตฌ์กฐ ํ•„์š” ์‹œ ์—ฐ๊ฒฐ
return kc_tokenizer, kc_model, deberta_tokenizer, deberta_model, cnn_model
except Exception as e:
raise RuntimeError(
"๋ชจ๋ธ ๋กœ๋“œ ์ค‘ ์˜ค๋ฅ˜๊ฐ€ ๋ฐœ์ƒํ–ˆ์Šต๋‹ˆ๋‹ค. ๋ชจ๋ธ ํด๋” ๊ฒฝ๋กœ ๋ฐ ํŒŒ์ผ๋“ค์„ ํ™•์ธํ•ด์ฃผ์„ธ์š”.\n"
f"์ƒ์„ธ ์˜ค๋ฅ˜: {e}"
)
# ๋ชจ๋ธ ๋กœ๋“œ ๊ตฌ๋™
try:
kc_tok, kc_mod, deb_tok, deb_mod, cnn_mod = load_verification_models()
device = "cuda" if torch.cuda.is_available() else "cpu"
# ๋ชจ๋ธ๋“ค์„ ์ ์ ˆํ•œ ๋””๋ฐ”์ด์Šค๋กœ ์ด๋™ (CPU/GPU)
kc_mod.to(device)
deb_mod.to(device)
except Exception as e:
st.error("โš ๏ธ ์‹œ์Šคํ…œ ์ดˆ๊ธฐํ™” ์‹คํŒจ (๋ชจ๋ธ ๋กœ๋“œ ์—๋Ÿฌ)")
st.code(str(e))
st.stop()
# --- [3] ๋‘ ๋ฌธ์žฅ์˜ ์Šคํƒ€์ผ ์œ ์‚ฌ๋„๋ฅผ ๊ณ„์‚ฐํ•˜๋Š” ํ•จ์ˆ˜ ---
def calculate_similarity(text_a, text_b):
with torch.no_grad():
# --- 1) KcBERT ์Šคํƒ€์ผ ๋ฒกํ„ฐ ๋ถ„์„ ---
inputs_a = kc_tok(text_a, return_tensors="pt", truncation=True, max_length=128).to(device)
outputs_a = kc_mod(**inputs_a)
prob_a_kc = F.softmax(outputs_a.logits, dim=-1).flatten()
inputs_b = kc_tok(text_b, return_tensors="pt", truncation=True, max_length=128).to(device)
outputs_b = kc_mod(**inputs_b)
prob_b_kc = F.softmax(outputs_b.logits, dim=-1).flatten()
kc_sim = F.cosine_similarity(prob_a_kc.unsqueeze(0), prob_b_kc.unsqueeze(0)).item() * 100
# --- 2) DeBERTa ์Šคํƒ€์ผ ๋ฒกํ„ฐ ๋ถ„์„ ---
deb_inputs_a = deb_tok(text_a, return_tensors="pt", truncation=True, max_length=128).to(device)
deb_outputs_a = deb_mod(**deb_inputs_a)
prob_a_deb = F.softmax(deb_outputs_a.logits, dim=-1).flatten()
deb_inputs_b = deb_tok(text_b, return_tensors="pt", truncation=True, max_length=128).to(device)
deb_outputs_b = deb_mod(**deb_inputs_b)
prob_b_deb = F.softmax(deb_outputs_b.logits, dim=-1).flatten()
deb_sim = F.cosine_similarity(prob_a_deb.unsqueeze(0), prob_b_deb.unsqueeze(0)).item() * 100
# --- 3) ์ตœ์ข… ์•™์ƒ๋ธ” ์œ ์‚ฌ๋„ ์‚ฐ์ถœ ---
final_similarity = (kc_sim + deb_sim) / 2.0
return kc_sim, deb_sim, final_similarity
# --- [4] Streamlit UI ๋””์ž์ธ (๋™์ผ์ธ ์‹๋ณ„ ์ „์šฉ) ---
st.set_page_config(page_title="์ €์ž ์‹๋ณ„ ์‹œ์Šคํ…œ", layout="wide", page_icon="๐Ÿ•ต๏ธโ€โ™‚๏ธ")
st.title("๐Ÿ•ต๏ธโ€โ™‚๏ธ ์•ŒํŒŒ ํ”„๋กœ์ ํŠธ: ๋ฌธ์ฒดํ•™ ๊ธฐ๋ฐ˜ ์ €์ž ๋™์ผ์„ฑ ๊ฒ€์ฆ ์‹œ์Šคํ…œ")
st.subheader("๋‘ ๊ฐœ์˜ ๊ธ€์„ ๋น„๊ตํ•˜์—ฌ ๋™์ผ ์ธ๋ฌผ์ด ์ž‘์„ฑํ–ˆ๋Š”์ง€ ์‹ค์‹œ๊ฐ„์œผ๋กœ ๋ถ„์„ํ•ฉ๋‹ˆ๋‹ค.")
st.write("---")
# ํ™”๋ฉด์„ ์™ผ์ชฝ, ์˜ค๋ฅธ์ชฝ ๋‘ ์นธ์œผ๋กœ ๋ถ„ํ• 
col1, col2 = st.columns(2)
with col1:
st.markdown("### ๐Ÿ“ ๋ถ„์„ ๋Œ€์ƒ ๊ธ€ A")
text_a = st.text_area(
"์ฒซ ๋ฒˆ์งธ ๊ธ€์„ ์ž…๋ ฅํ•˜์„ธ์š”:",
placeholder="๋น„๊ตํ•  ์ฒซ ๋ฒˆ์งธ ๋ณธ๋ฌธ์„ ์ž…๋ ฅํ•˜์„ธ์š”.",
height=250,
key="text_a",
)
with col2:
st.markdown("### ๐Ÿ“ ๋ถ„์„ ๋Œ€์ƒ ๊ธ€ B")
text_b = st.text_area(
"๋‘ ๋ฒˆ์งธ ๊ธ€์„ ์ž…๋ ฅํ•˜์„ธ์š”:",
placeholder="๋น„๊ตํ•  ๋‘ ๋ฒˆ์งธ ๋ณธ๋ฌธ์„ ์ž…๋ ฅํ•˜์„ธ์š”.",
height=250,
key="text_b",
)
st.write("---")
# ์‹คํ–‰ ๋ฒ„ํŠผ
if st.button("๐Ÿ” ๋™์ผ์ธ ์—ฌ๋ถ€ ์ •๋ฐ€ ๊ฒ€์ฆ ์‹œ์ž‘", use_container_width=True):
if not text_a.strip() or not text_b.strip():
st.error("โš ๏ธ ๊ธ€ A์™€ ๊ธ€ B ๋ชจ๋‘ ํ…์ŠคํŠธ๋ฅผ ์ž…๋ ฅํ•ด์•ผ ๋ถ„์„์ด ๊ฐ€๋Šฅํ•ฉ๋‹ˆ๋‹ค!")
else:
try:
with st.spinner("3๋Œ€์žฅ ์•Œ๊ณ ๋ฆฌ์ฆ˜์ด ๋‘ ๋ฌธ์žฅ์˜ ๋ฌธ์ฒด ํŒจํ„ด์„ ๋Œ€์กฐํ•˜๋Š” ์ค‘..."):
# ์œ ์‚ฌ๋„ ์—ฐ์‚ฐ ์‹คํ–‰
kc_sim, deb_sim, final_sim = calculate_similarity(text_a, text_b)
# --- [5] ํŒ์ • ๊ฒฐ๊ณผ ์‹œ๊ฐํ™” ---
st.success("๐ŸŽ‰ ๋ฌธ์ฒด ๋Œ€์กฐ ๋ถ„์„ ์™„๋ฃŒ!")
# ๊ธฐ์ค€์ (Threshold) ์„ค์ •
THRESHOLD = 85.0
is_same_author = final_sim >= THRESHOLD
# ๋Œ€ํ˜• ํŒจ๋„๋กœ ๊ฒฐ๊ณผ ๋…ธ์ถœ
rect_col1, rect_col2 = st.columns(2)
with rect_col1:
if is_same_author:
st.metric(label="์ตœ์ข… ํŒ์ • ๊ฒฐ๊ณผ", value="๐ŸŸข ๋™์ผ์ธ ๊ฐ€๋Šฅ์„ฑ ๋งค์šฐ ๋†’์Œ")
else:
st.metric(label="์ตœ์ข… ํŒ์ • ๊ฒฐ๊ณผ", value="๐Ÿ”ด ๋‹ค๋ฅธ ์ธ๋ฌผ์ผ ๊ฐ€๋Šฅ์„ฑ ๋†’์Œ")
with rect_col2:
st.metric(label="์ตœ์ข… ๋ฌธ์ฒด ์œ ์‚ฌ๋„ ์ ์ˆ˜", value=f"{final_sim:.2f} / 100์ ")
st.write("---")
# ์ง„ํ–‰ ๋ฐ” ์‹œ๊ฐํ™” ์ถ”๊ฐ€
st.progress(min(max(final_sim / 100.0, 0.0), 1.0))
# ๋ ˆ์ด์•„์›ƒ ๋ถ„ํ• : ์™ผ์ชฝ์€ ์ฐจํŠธ, ์˜ค๋ฅธ์ชฝ์€ ์ƒ์„ธ ํ‘œ
result_col1, result_col2 = st.columns([4, 3])
with result_col1:
st.markdown("### ๐Ÿ“Š ๋ชจ๋ธ๋ณ„ ๋ฌธ์ฒด ๋Œ€์กฐ ์Šค์ฝ”์–ด")
scores = {
"KcBERT ๋Œ€์กฐ ์ ์ˆ˜": kc_sim,
"DeBERTa ๋Œ€์กฐ ์ ์ˆ˜": deb_sim,
"์ตœ์ข… ์•™์ƒ๋ธ” ๊ฒฐ๋ก ": final_sim,
}
st.bar_chart(scores)
with result_col2:
st.markdown("### ๐Ÿ“‹ ๋ชจ๋ธ๋ณ„ ๊ฒฐ๊ณผ ์ˆ˜์น˜")
model_table = {
"ํ‰๊ฐ€ ํ•ญ๋ชฉ": ["KcBERT", "DeBERTa", "์ตœ์ข… ์•™์ƒ๋ธ” ๊ฒฐ๊ณผ"],
"์œ ์‚ฌ๋„ ์Šค์ฝ”์–ด": [f"{kc_sim:.2f}์ ", f"{deb_sim:.2f}์ ", f"{final_sim:.2f}์ "]
}
st.table(model_table)
# ์—ฐ๊ตฌ์‹ค ๋ณด๊ณ ์šฉ ์ƒ์„ธ ํ…์ŠคํŠธ ํ”ผ๋“œ๋ฐฑ
st.info(
f"๐Ÿ’ก **๋ถ„์„ ๊ฒฐ๊ณผ ์š”์•ฝ**: ๋‘ ๋ฌธ์žฅ์˜ ๋ฌธ์žฅ ๊ตฌ์กฐ, ๋‹จ์–ด ์„ ํƒ ํŒจํ„ด, ์–ด์กฐ๋ฅผ ์ข…ํ•ฉํ•œ ๊ฒฐ๊ณผ "
f"์ตœ์ข… **{final_sim:.1f}%**์˜ ์ผ์น˜์œจ์„ ๋ณด์˜€์Šต๋‹ˆ๋‹ค. (ํŒ์ • ๊ธฐ์ค€์„ : {THRESHOLD}%)"
)
# ๊ณ ๊ธ‰ ํ™•์žฅ ํƒญ ์ •๋ณด (๊ธฐ์กด ๋””๋ฒ„๊น…์šฉ UI ๊ณ„์Šน)
with st.expander("๐Ÿ› ๏ธ ์‹œ์Šคํ…œ ๊ธฐ์ˆ  ์ •๋ณด"):
st.write("์‹คํ–‰ ์žฅ์น˜ (Device):", device)
st.write("Python ์‹คํ–‰ ๊ฒฝ๋กœ:")
st.code(sys.executable)
st.write("๋ชจ๋ธ ๊ธฐ๋ณธ ๋””๋ ‰ํ† ๋ฆฌ:")
st.code(BASE_DIR)
st.caption("์ฃผ์˜: ๋ณธ ๊ฒฐ๊ณผ๋Š” AI ๋ชจ๋ธ ๊ธฐ๋ฐ˜ ์ถ”์ •๊ฐ’์ด๋ฉฐ, ์‹ค์ œ ๋™์ผ์ธ์„ ๋‹จ์ •ํ•˜๋Š” ๋ฒ•์  ํŒ๋‹จ์ด ์•„๋‹™๋‹ˆ๋‹ค.")
except Exception as e:
st.error("๊ณ ๊ธ‰ ๋ชจ๋ธ ์‹คํ–‰ ์ค‘ ์˜ค๋ฅ˜๊ฐ€ ๋ฐœ์ƒํ–ˆ์Šต๋‹ˆ๋‹ค.")
st.code(str(e))
with st.expander("์ƒ์„ธ ์˜ค๋ฅ˜ ๋ณด๊ธฐ (Traceback)"):
st.code(traceback.format_exc())