Spaces:
Running
Running
File size: 6,729 Bytes
2f97561 | 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 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 | import os
import time
from pathlib import Path
import pandas as pd
import streamlit as st
from src.predictor import ArticleTopicPredictor
st.set_page_config(
page_title="Article Topic Classifier",
page_icon="🧠",
layout="centered",
)
MODEL_DIR = os.getenv("MODEL_DIR", "artifacts/article_topic_model")
EXAMPLES = [
{
"title": "Attention-based models for scientific document understanding",
"abstract": "We propose a transformer-based architecture for classification and retrieval of scientific papers.",
},
{
"title": "A new benchmark for graph representation learning",
"abstract": "This paper introduces a benchmark suite for graph neural networks and evaluates generalization.",
},
{
"title": "Quantum error correction with surface codes",
"abstract": "",
},
]
@st.cache_resource(show_spinner=False)
def load_predictor(model_dir: str) -> ArticleTopicPredictor:
return ArticleTopicPredictor(model_dir=model_dir)
def render_intro() -> None:
st.title("🧠 Классификатор тематик научных статей")
st.markdown(
"""
Введите заголовок статьи и, при наличии, abstract.
Сервис покажет наиболее вероятные тематики и набор **top-95%** классов:
классы по убыванию вероятности, пока суммарная вероятность не превысит 95%.
"""
)
st.caption("Если abstract пустой, классификация выполняется только по title.")
def render_sidebar() -> None:
with st.sidebar:
st.header("О приложении")
st.write("Модель загружается один раз и кэшируется между перезапусками интерфейса.")
st.write("Поддерживается режим работы только по title.")
st.write(f"Папка модели: `{MODEL_DIR}`")
st.divider()
st.subheader("Быстрые примеры")
for i, example in enumerate(EXAMPLES):
if st.button(f"Подставить пример {i + 1}", use_container_width=True):
st.session_state["title_input"] = example["title"]
st.session_state["abstract_input"] = example["abstract"]
def validate_inputs(title: str, abstract: str) -> str | None:
if not title.strip() and not abstract.strip():
return "Введите хотя бы заголовок статьи или abstract."
if len(title.strip()) > 600:
return "Заголовок слишком длинный. Пожалуйста, сократите title до 600 символов."
if len(abstract.strip()) > 8000:
return "Abstract слишком длинный. Пожалуйста, сократите текст до 8000 символов."
return None
def main() -> None:
render_intro()
render_sidebar()
if "title_input" not in st.session_state:
st.session_state["title_input"] = ""
if "abstract_input" not in st.session_state:
st.session_state["abstract_input"] = ""
model_path = Path(MODEL_DIR)
if not model_path.exists():
st.error(
"Папка с моделью не найдена. Сначала обучите модель и сохраните её в "
f"`{MODEL_DIR}`, либо задайте переменную окружения MODEL_DIR."
)
st.stop()
title = st.text_input(
"Название статьи",
key="title_input",
placeholder="Например: Attention-based methods for scientific text classification",
)
abstract = st.text_area(
"Abstract (необязательно)",
key="abstract_input",
placeholder="Вставьте аннотацию статьи. Поле можно оставить пустым.",
height=220,
)
col1, col2 = st.columns([1, 1])
with col1:
run = st.button("Классифицировать", type="primary", use_container_width=True)
with col2:
clear = st.button("Очистить", use_container_width=True)
if clear:
st.session_state["title_input"] = ""
st.session_state["abstract_input"] = ""
st.rerun()
if run:
error_message = validate_inputs(title, abstract)
if error_message:
st.warning(error_message)
st.stop()
try:
predictor = load_predictor(MODEL_DIR)
with st.spinner("Считаю вероятности классов..."):
started = time.perf_counter()
result = predictor.predict(title=title, abstract=abstract, top95_threshold=0.95)
elapsed = time.perf_counter() - started
st.success(f"Готово. Время инференса: {elapsed:.2f} сек.")
st.subheader("Top-95% тематики")
top95_df = pd.DataFrame(result["top95"])
top95_df["probability"] = top95_df["probability"].map(lambda x: round(float(x), 4))
top95_df["cumulative_probability"] = top95_df["cumulative_probability"].map(lambda x: round(float(x), 4))
st.dataframe(
top95_df.rename(
columns={
"label": "Тема",
"probability": "Вероятность",
"cumulative_probability": "Накопленная вероятность",
}
),
use_container_width=True,
hide_index=True,
)
st.subheader("Все вероятности")
full_df = pd.DataFrame(result["all_probs"])
full_df["probability"] = full_df["probability"].map(float)
st.bar_chart(full_df.set_index("label")["probability"])
st.dataframe(
full_df.rename(columns={"label": "Тема", "probability": "Вероятность"}),
use_container_width=True,
hide_index=True,
)
st.caption(
"Top-95% — это минимальный набор классов, чей суммарный вес достигает не менее 95%."
)
except Exception as exc: # noqa: BLE001
st.error("Во время инференса произошла ошибка, но приложение не упало.")
st.exception(exc)
if __name__ == "__main__":
main()
|