ML2-hw4 / app.py
Alisa
arxiv urls added
390519d
Raw
History Blame Contribute Delete
10.6 kB
import json
import re
import numpy as np
import torch
import gradio as gr
import httpx
from xml.etree import ElementTree as ET
from transformers import AutoTokenizer, AutoModelForSequenceClassification
# ── Константы ──────────────────────────────────────────────────────────────
MODEL_DIR = "model_best"
LABEL_MAP = "data/label_mapping.json"
MAX_LEN = 256
THRESHOLD = 0.95
CATEGORY_NAMES = {
"cs": "Computer Science",
"math": "Mathematics",
"physics": "Physics",
"q-bio": "Quantitative Biology",
"stat": "Statistics",
"econ": "Economics",
"eess": "Electrical Engineering",
"q-fin": "Quantitative Finance",
}
# ── Загрузка модели ────────────────────────────────────────────────────────
def load_model():
with open(LABEL_MAP) as f:
label_info = json.load(f)
id2cat = {int(k): v for k, v in label_info["id2cat"].items()}
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
tokenizer = AutoTokenizer.from_pretrained(MODEL_DIR)
model = AutoModelForSequenceClassification.from_pretrained(MODEL_DIR)
model.to(device).eval()
return tokenizer, model, id2cat, device
tokenizer, model, id2cat, device = load_model()
# ── Утилиты ────────────────────────────────────────────────────────────────
def clean_category(cat: str) -> str:
match = re.search(r"term':\s*'(\w[\w-]*)", cat)
if match:
return match.group(1)
return cat.strip()
def extract_arxiv_id(url: str) -> str | None:
"""Вытаскивает arxiv id из ссылки любого формата."""
# Поддерживаем форматы:
# https://arxiv.org/abs/2301.00001
# https://arxiv.org/pdf/2301.00001
# 2301.00001
match = re.search(r"(\d{4}\.\d{4,5}(v\d+)?)", url)
return match.group(1) if match else None
def fetch_arxiv(url: str):
arxiv_id = extract_arxiv_id(url.strip())
if not arxiv_id:
return None, None, "❌ Не удалось распознать arxiv ID. Пример: https://arxiv.org/abs/1706.03762"
try:
api_url = f"https://export.arxiv.org/api/query?id_list={arxiv_id}"
response = httpx.get(api_url, timeout=8.0) # жёсткий таймаут 8 сек
response.raise_for_status()
ns = {"atom": "http://www.w3.org/2005/Atom"}
root = ET.fromstring(response.text)
entry = root.find("atom:entry", ns)
if entry is None:
return None, None, f"❌ Статья {arxiv_id} не найдена."
title = entry.find("atom:title", ns).text.strip().replace("\n", " ")
abstract = entry.find("atom:summary", ns).text.strip().replace("\n", " ")
return title, abstract, None
except httpx.TimeoutException:
return None, None, "❌ arxiv не отвечает — попробуй ещё раз."
except Exception as e:
return None, None, f"❌ Ошибка: {e}"
# ── Предсказание ───────────────────────────────────────────────────────────
def predict(title: str, abstract: str):
if not title or not title.strip():
return "⚠️ Введите название статьи.", None
enc = tokenizer(
title.strip(),
abstract.strip() if abstract else "",
max_length=MAX_LEN,
truncation=True,
padding="max_length",
return_tensors="pt",
)
with torch.no_grad():
logits = model(
input_ids=enc["input_ids"].to(device),
attention_mask=enc["attention_mask"].to(device),
).logits
probs = torch.softmax(logits, dim=-1).squeeze().cpu().numpy()
sorted_ids = np.argsort(probs)[::-1]
results, cumsum = [], 0.0
for idx in sorted_ids:
cat = clean_category(id2cat[idx])
results.append((cat, float(probs[idx])))
cumsum += probs[idx]
if cumsum >= THRESHOLD:
break
lines = []
if not abstract or not abstract.strip():
lines.append("ℹ️ Abstract не введён — классификация только по названию.\n")
for i, (cat, prob) in enumerate(results):
name = CATEGORY_NAMES.get(cat, cat)
prefix = "🥇" if i == 0 else ("🥈" if i == 1 else "▪️")
lines.append(f"{prefix} **{name}** (`{cat}`) — {prob:.1%}")
total = sum(p for _, p in results)
lines.append(f"\n_Суммарная вероятность: {total:.1%}_")
all_sorted = [(clean_category(id2cat[i]), float(probs[i])) for i in np.argsort(probs)[::-1]]
plot_data = {
"Category": [CATEGORY_NAMES.get(c, c) for c, _ in all_sorted],
"Probability": [round(p, 4) for _, p in all_sorted],
}
return "\n".join(lines), plot_data
def predict_from_url(url: str):
"""Парсит arxiv и классифицирует."""
if not url or not url.strip():
return "", "", "⚠️ Введите ссылку на статью.", None
title, abstract, error = fetch_arxiv(url)
if error:
return "", "", error, None
result_text, plot_data = predict(title, abstract)
return title, abstract, result_text, plot_data
# ── Интерфейс ──────────────────────────────────────────────────────────────
with gr.Blocks(title="arXiv Classifier") as demo:
gr.Markdown(
"""
# 📄 arXiv Article Classifier
Определяет тематику научной статьи по названию и аннотации.
Показывает категории с суммарной вероятностью **≥ 95%**.
"""
)
with gr.Tabs():
# ── Вкладка 1: ввод вручную ────────────────────────────────────────
with gr.Tab("✏️ Ввод вручную"):
with gr.Row():
with gr.Column(scale=2):
title_input = gr.Textbox(
label="Название статьи",
placeholder="Например: Attention Is All You Need",
)
abstract_input = gr.Textbox(
label="Abstract (необязательно)",
placeholder="Вставьте аннотацию статьи сюда…",
lines=7,
)
btn_manual = gr.Button("Классифицировать", variant="primary")
with gr.Column(scale=3):
result_text_manual = gr.Markdown(label="Результат")
result_plot_manual = gr.BarPlot(
x="Category",
y="Probability",
title="Вероятности по категориям",
y_lim=[0, 1],
tooltip=["Category", "Probability"],
)
btn_manual.click(
fn=predict,
inputs=[title_input, abstract_input],
outputs=[result_text_manual, result_plot_manual],
)
gr.Examples(
examples=[
["Attention Is All You Need",
"We propose a new simple network architecture, the Transformer, based solely on attention mechanisms."],
["A proof of the Riemann Hypothesis using spectral methods", ""],
["CRISPR-Cas9 genome editing in mammalian cells",
"We demonstrate efficient genome editing in human cells using the CRISPR-Cas9 system."],
],
inputs=[title_input, abstract_input],
label="Примеры",
)
# ── Вкладка 2: ссылка на arxiv ─────────────────────────────────────
with gr.Tab("🔗 Ссылка на arxiv"):
with gr.Row():
with gr.Column(scale=2):
url_input = gr.Textbox(
label="Ссылка на статью",
placeholder="https://arxiv.org/abs/1706.03762",
)
btn_url = gr.Button("Загрузить и классифицировать", variant="primary")
parsed_title = gr.Textbox(label="Название (спарсено)", interactive=False)
parsed_abstract = gr.Textbox(label="Abstract (спарсен)", lines=5, interactive=False)
with gr.Column(scale=3):
result_text_url = gr.Markdown(label="Результат")
result_plot_url = gr.BarPlot(
x="Category",
y="Probability",
title="Вероятности по категориям",
y_lim=[0, 1],
tooltip=["Category", "Probability"],
)
btn_url.click(
fn=predict_from_url,
inputs=[url_input],
outputs=[parsed_title, parsed_abstract, result_text_url, result_plot_url],
)
gr.Examples(
examples=[
["https://arxiv.org/abs/1706.03762"], # Attention is all you need
["https://arxiv.org/abs/1810.04805"], # BERT
["https://arxiv.org/abs/2303.08774"], # GPT-4
],
inputs=[url_input],
label="Примеры ссылок",
)
with gr.Accordion("Поддерживаемые категории", open=False):
gr.Markdown("\n".join(f"- `{k}` — {v}" for k, v in CATEGORY_NAMES.items()))
if __name__ == "__main__":
demo.launch()