| 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 |
|
|
|
|
| |
| 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", |
| ) |
|
|
| |
| |
| 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) |