Spaces:
Sleeping
Sleeping
| 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) |