Spaces:
Running
Running
Download app.py from asriel14/article_classifier: direct link, hf CLI and curl.
- Browser
- Download file 6.73 kB
-
https://huggingface.co/spaces/asriel14/article_classifier/resolve/main/app.py
- Command line
-
hf download hf://spaces/asriel14/article_classifier/app.py
-
curl -L -o app.py https://huggingface.co/spaces/asriel14/article_classifier/resolve/main/app.py
6.73 kB
| 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": "", | |
| }, | |
| ] | |
| 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() | |