A-tags / app.py
Avisl's picture
Update app.py
538a0c5 verified
Raw
History Blame Contribute Delete
6.57 kB
import os
import traceback
from typing import Optional, List
import gradio as gr
from inference import (
clean_text,
load_tags_model,
infer_tags_for_texts,
)
TAGS_REPO = "DanielNRU/Avito_tags-ruroberta"
USE_CUDA_ENV = os.getenv("USE_CUDA", "1")
USE_CUDA = USE_CUDA_ENV == "1"
class TagsService:
def __init__(self, tags_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.all_tags,
self.threshold,
self.head_tokens,
self.tail_tokens,
self.strategy,
) = load_tags_model(tags_model_dir, device)
self.device = device
def analyze_text(self, text: str, threshold: Optional[float] = None) -> List[str]:
thr = threshold if threshold is not None else self.threshold
text_clean = clean_text(text)
tags_list = infer_tags_for_texts(
[text_clean],
self.model,
self.tokenizer,
self.max_length,
self.all_tags,
thr,
self.device,
return_probs=False,
head_tokens=self.head_tokens,
tail_tokens=self.tail_tokens,
strategy=self.strategy,
)
if not tags_list:
return []
return tags_list[0] or []
def analyze_text_with_probs(self, text: str, threshold: Optional[float] = None):
thr = threshold if threshold is not None else self.threshold
text_clean = clean_text(text)
tags_list, probs_list = infer_tags_for_texts(
[text_clean],
self.model,
self.tokenizer,
self.max_length,
self.all_tags,
thr,
self.device,
return_probs=True,
head_tokens=self.head_tokens,
tail_tokens=self.tail_tokens,
strategy=self.strategy,
)
if not tags_list:
return [], []
thr_val = thr
result_pairs = [
(tag, float(prob))
for tag, prob in zip(self.all_tags, probs_list[0])
if prob >= thr_val
]
result_tags = [t for t, _ in result_pairs]
return result_tags, result_pairs
_service: Optional[TagsService] = None
def get_service() -> TagsService:
global _service
if _service is None:
_service = TagsService(TAGS_REPO, use_cuda=USE_CUDA)
return _service
def analyze_single_text(text: str, tags_thr: float):
"""
Единая точка входа для клиента и UI.
При пустом тексте возвращаем '—', чтобы не падать.
"""
if not text or not str(text).strip():
return "—", "—"
try:
svc = get_service()
tags_list, probs_list = svc.analyze_text_with_probs(text, threshold=float(tags_thr))
except Exception as e:
print(f"[ERROR] analyze_single_text / A-tags: {e}")
traceback.print_exc()
return "", ""
if not tags_list:
return "", ""
tags_str = ", ".join(tags_list)
probs_str = ", ".join(f"{tag}: {prob*100:.1f}%" for tag, prob in probs_list)
return tags_str, probs_str
# ── Issue #222 / Шаг 8: батч-endpoint ────────────────────────────────────────
def analyze_batch(
texts: List[str],
tags_thr: float = 0.3,
) -> List[List[str]]:
"""Батч-анализ тегов (A).
Принимает список текстов, возвращает [[tags_str, probs_str], ...].
Пустые строки получают ['', ''] без вызова модели.
HFBatchSender вызывает:
client.predict(texts, tags_thr, api_name='/analyze_batch')
Returns:
[[tags_str, probs_str], ...] той же длины что и texts.
"""
svc = get_service()
results: List[List[str]] = []
for text in texts:
if not text or not str(text).strip():
results.append(["", ""])
continue
try:
tags_list, probs_list = svc.analyze_text_with_probs(
str(text), threshold=float(tags_thr)
)
tags_str = ", ".join(tags_list) if tags_list else ""
probs_str = (
", ".join(f"{tag}: {prob*100:.1f}%" for tag, prob in probs_list)
if probs_list else ""
)
results.append([tags_str, probs_str])
except Exception as e:
print(f"[ERROR] analyze_batch / A-tags (item): {e}")
results.append(["ошибка", ""])
return results
# ─────────────────────────────────────────────────────────────────────────────
# Дефолтный порог для слайдера
_default_thr = None
try:
_tmp_svc = TagsService(TAGS_REPO, use_cuda=False)
_default_thr = float(_tmp_svc.threshold)
except Exception as e:
print(f"[WARN] Не удалось инициализировать TagsService для _default_thr: {e}")
_default_thr = 0.3
with gr.Blocks(title="Теги сообщения") as demo:
gr.Markdown("# Предсказание тематических тегов")
inp_text = gr.Textbox(
label="Текст сообщения",
placeholder="Вставьте сообщение...",
lines=8,
)
tags_thr_slider = gr.Slider(
minimum=0.0,
maximum=1.0,
value=_default_thr,
step=0.01,
label="Порог тегов",
)
btn = gr.Button("Анализировать")
out_tags = gr.Textbox(label="Теги", interactive=False)
out_probs = gr.Textbox(label="Вероятности тегов", interactive=False)
btn.click(
fn=analyze_single_text,
inputs=[inp_text, tags_thr_slider],
outputs=[out_tags, out_probs],
api_name="analyze_single_text",
)
# Issue #222 / Шаг 8: батч-endpoint
# HFBatchSender: client.predict(texts, tags_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)