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)