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)