A-relevance / app.py
Avisl's picture
Update app.py
568b371 verified
Raw
History Blame Contribute Delete
5.49 kB
import os
import traceback
from typing import Optional, List
import gradio as gr
from inference import (
clean_text,
load_relevance_model,
infer_relevance,
)
RELEVANCE_REPO = "DanielNRU/Avito-relevance-rubert-20260516"
USE_CUDA_ENV = os.getenv("USE_CUDA", "1")
USE_CUDA = USE_CUDA_ENV == "1"
class RelevanceService:
def __init__(self, relevance_model_dir: str, use_cuda: bool = True):
import torch
device = torch.device(
"cuda" if use_cuda and torch.cuda.is_available() else "cpu"
)
(
self.model,
self.tokenizer,
self.max_length,
self.threshold,
) = load_relevance_model(relevance_model_dir, device)
self.device = device
def analyze_text(self, text: str, threshold: Optional[float] = None):
thr = threshold if threshold is not None else self.threshold
text_clean = clean_text(text)
preds, probs = infer_relevance(
[text_clean],
self.model,
self.tokenizer,
self.max_length,
thr,
self.device,
batch_size=1,
)
if len(preds) == 0:
return 0, 0.0
label = int(preds[0])
prob = float(probs[0])
return label, prob
_service: Optional[RelevanceService] = None
def get_service() -> RelevanceService:
global _service
if _service is None:
_service = RelevanceService(RELEVANCE_REPO, use_cuda=USE_CUDA)
return _service
def _format_pct(p: float) -> str:
return f"{p * 100:.1f}%"
def analyze_single_text(text: str, rel_thr: float):
if not text or not text.strip():
return "пустой текст", "0.0%"
try:
svc = get_service()
label, prob = svc.analyze_text(text, threshold=rel_thr)
label_str = "релевантно" if label == 1 else "нерелевантно"
prob_str = _format_pct(prob)
return label_str, prob_str
except Exception as e:
print(f"[ERROR] analyze_single_text / A-relevance: {e}")
traceback.print_exc()
return "ошибка", "0.0%"
# ── Issue #222 / Шаг 8: батч-endpoint ────────────────────────────────────────
def analyze_batch(
texts: List[str],
thr: float = 0.5,
) -> List[List[str]]:
"""Батч-анализ релевантности (A).
Принимает список текстов, возвращает [[label_str, prob_pct], ...].
Пустые строки получают ['нерелевантно', '0.0%'] без вызова модели.
HFBatchSender вызывает:
client.predict(texts, thr, api_name='/analyze_batch')
Returns:
[[label_str, prob_pct], ...] той же длины что и texts.
"""
svc = get_service()
results: List[List[str]] = []
for text in texts:
if not text or not str(text).strip():
results.append(["нерелевантно", "0.0%"])
continue
try:
label, prob = svc.analyze_text(str(text), threshold=float(thr))
label_str = "релевантно" if label == 1 else "нерелевантно"
prob_str = _format_pct(prob)
results.append([label_str, prob_str])
except Exception as e:
print(f"[ERROR] analyze_batch / A-relevance (item): {e}")
results.append(["ошибка", "0.0%"])
return results
# ─────────────────────────────────────────────────────────────────────────────
# Чтобы взять дефолтный порог из модели для слайдера
_default_thr = None
try:
_tmp_svc = RelevanceService(RELEVANCE_REPO, use_cuda=False)
_default_thr = float(_tmp_svc.threshold)
except Exception as e:
print(f"[WARN] Не удалось инициализировать RelevanceService для _default_thr: {e}")
_default_thr = 0.5
with gr.Blocks(title="Релевантность сообщения") as demo:
gr.Markdown("# Определение релевантности сообщения")
inp_text = gr.Textbox(
label="Текст сообщения",
placeholder="Вставьте сообщение...",
lines=8,
)
rel_thr_slider = gr.Slider(
minimum=0.0,
maximum=1.0,
value=_default_thr,
step=0.01,
label="Порог релевантности",
)
btn = gr.Button("Анализировать")
out_label = gr.Textbox(
label="Метка (релевантно / нерелевантно)",
interactive=False,
)
out_prob = gr.Textbox(
label="Вероятность релевантности",
interactive=False,
)
btn.click(
fn=analyze_single_text,
inputs=[inp_text, rel_thr_slider],
outputs=[out_label, out_prob],
api_name="analyze_single_text",
)
# Issue #222 / Шаг 8: батч-endpoint
# HFBatchSender: client.predict(texts, thr, api_name='/analyze_batch')
gr.api(
fn=analyze_batch,
api_name="analyze_batch",
)
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0", server_port=7860, show_error=True)