Spaces:
Running
Running
| import os | |
| import traceback | |
| from typing import Optional, Dict, Tuple, Any, List | |
| import gradio as gr | |
| from inference import ( | |
| clean_text, | |
| load_sentiment_model, | |
| infer_sentiment_single, | |
| _SENTIMENT_NEGATIVE_THRESHOLD, | |
| _SENTIMENT_POSITIVE_THRESHOLD, | |
| ) | |
| SENTIMENT_REPO = "DanielNRU/aiforever-rubert-large-tbsa-avito-sentiment" | |
| USE_CUDA_ENV = os.getenv("USE_CUDA", "1") | |
| USE_CUDA = USE_CUDA_ENV == "1" | |
| # ── Нормализация ────────────────────────────────────────────────────────────── | |
| def normalize_tone_label(value: Any) -> str: | |
| if value is None: | |
| return "—" | |
| text = str(value).strip().lower() | |
| if not text or text in {"—", "-", "none", "null"}: | |
| return "—" | |
| if any(x in text for x in ["негатив", "negative", "neg"]): | |
| return "негатив" | |
| if any(x in text for x in ["нейтрал", "neutral", "neu"]): | |
| return "нейтрально" | |
| if any(x in text for x in ["позитив", "positive", "pos"]): | |
| return "позитив" | |
| if text in {"-1", "-1.0"}: | |
| return "негатив" | |
| if text in {"0", "0.0"}: | |
| return "нейтрально" | |
| if text in {"1", "1.0"}: | |
| return "позитив" | |
| return "—" | |
| def normalize_prob_value(v: Any) -> Optional[float]: | |
| try: | |
| x = float(v) | |
| except Exception: | |
| return None | |
| if x > 1.0: | |
| x = x / 100.0 | |
| if 0.0 <= x <= 1.0: | |
| return x | |
| return None | |
| def normalize_prob_dict(probs: Any) -> Dict[str, float]: | |
| result: Dict[str, float] = {} | |
| if not isinstance(probs, dict): | |
| return result | |
| for k, v in probs.items(): | |
| key = normalize_tone_label(k) | |
| if key == "—": | |
| continue | |
| val = normalize_prob_value(v) | |
| if val is None: | |
| continue | |
| result[key] = val | |
| return result | |
| def choose_tone_from_probs(tone_probs: Dict[str, float]) -> str: | |
| if not tone_probs: | |
| return "—" | |
| best_label = None | |
| best_prob = -1.0 | |
| for key in ["негатив", "нейтрально", "позитив"]: | |
| val = tone_probs.get(key) | |
| if val is not None and val > best_prob: | |
| best_prob = val | |
| best_label = key | |
| return best_label if best_label is not None else "—" | |
| def build_probs_block(tone_probs: Dict[str, float]) -> str: | |
| if not tone_probs: | |
| return "—" | |
| lines = [] | |
| for key in ["негатив", "нейтрально", "позитив"]: | |
| if key in tone_probs: | |
| try: | |
| lines.append(f"{key}: {tone_probs[key] * 100:.1f}%") | |
| except Exception: | |
| continue | |
| return "\n".join(lines) if lines else "—" | |
| # ── Логика порогов neg_thr / pos_thr ───────────────────────────────────────── | |
| def _apply_thresholds( | |
| tone_probs: Dict[str, float], | |
| neg_thr: Optional[float] = None, | |
| pos_thr: Optional[float] = None, | |
| ) -> str: | |
| """Определяет тональность по раздельным порогам neg_thr / pos_thr. | |
| Логика (A-модель всегда работает в режиме neg/pos, без tone_thr): | |
| - P(neg) >= neg_thr → 'негатив' | |
| - P(pos) >= pos_thr → 'позитив' | |
| - иначе → 'нейтрально' | |
| Fallback при отсутствии порогов — argmax. | |
| """ | |
| if not tone_probs: | |
| return "—" | |
| p_neg = tone_probs.get("негатив", 0.0) | |
| p_pos = tone_probs.get("позитив", 0.0) | |
| if neg_thr is not None or pos_thr is not None: | |
| _neg_thr = neg_thr if neg_thr is not None else 1.0 | |
| _pos_thr = pos_thr if pos_thr is not None else 1.0 | |
| if p_neg >= _neg_thr: | |
| return "негатив" | |
| if p_pos >= _pos_thr: | |
| return "позитив" | |
| return "нейтрально" | |
| # fallback argmax | |
| return choose_tone_from_probs(tone_probs) | |
| # ── SentimentService ────────────────────────────────────────────────────────── | |
| class SentimentService: | |
| def __init__(self, sentiment_model_dir: str, use_cuda: bool = True): | |
| import torch | |
| device = torch.device("cuda" if use_cuda and torch.cuda.is_available() else "cpu") | |
| print(f"[INFO] SentimentService device: {device}") | |
| print(f"[INFO] Загружаем модель: {sentiment_model_dir}") | |
| self.model, self.tokenizer, self.max_length = load_sentiment_model( | |
| sentiment_model_dir, device | |
| ) | |
| self.device = device | |
| def analyze_text( | |
| self, | |
| text: str, | |
| negative_threshold: Optional[float] = None, | |
| positive_threshold: Optional[float] = None, | |
| ) -> Tuple[str, Dict[str, float]]: | |
| """Анализирует тональность одного текста. | |
| Использует neg_thr/pos_thr нативно через infer_sentiment_single. | |
| Дополнительно применяет _apply_thresholds() для согласованности | |
| с SL-sentiment и HFBatchSender. | |
| """ | |
| text_clean = clean_text(text) | |
| if not text_clean: | |
| return "—", {} | |
| _neg_thr = negative_threshold if negative_threshold is not None else _SENTIMENT_NEGATIVE_THRESHOLD | |
| _pos_thr = positive_threshold if positive_threshold is not None else _SENTIMENT_POSITIVE_THRESHOLD | |
| try: | |
| res: Any = infer_sentiment_single( | |
| text_clean, | |
| self.model, | |
| self.tokenizer, | |
| self.max_length, | |
| self.device, | |
| negative_threshold=_neg_thr, | |
| positive_threshold=_pos_thr, | |
| ) | |
| print(f"[DEBUG] infer_sentiment_single raw result: {res!r}") | |
| except Exception as e: | |
| print(f"[ERROR] infer_sentiment_single exception: {e}") | |
| traceback.print_exc() | |
| return "—", {} | |
| if not isinstance(res, dict): | |
| print(f"[WARN] unexpected type from infer_sentiment_single: {type(res)}") | |
| return "—", {} | |
| tone_probs = normalize_prob_dict(res.get("tone_probs")) | |
| # Применяем _apply_thresholds для явного контроля (консистентно с SL-sentiment) | |
| tone_str = _apply_thresholds(tone_probs, neg_thr=_neg_thr, pos_thr=_pos_thr) | |
| if tone_str == "—": | |
| tone_str = normalize_tone_label(res.get("tone_str")) | |
| print(f"[DEBUG] final tone_str: {tone_str!r}, probs: {tone_probs}") | |
| return tone_str, tone_probs | |
| _service: Optional[SentimentService] = None | |
| def get_service() -> SentimentService: | |
| global _service | |
| if _service is None: | |
| print("[INFO] Инициализация SentimentService...") | |
| _service = SentimentService(SENTIMENT_REPO, use_cuda=USE_CUDA) | |
| return _service | |
| # ── UI functions ────────────────────────────────────────────────────────────── | |
| def analyze_single_text(text: str, neg_threshold: float, pos_threshold: float): | |
| text = text or "" | |
| if not text.strip(): | |
| return "—", "—" | |
| try: | |
| svc = get_service() | |
| tone_str, tone_probs = svc.analyze_text( | |
| text, | |
| negative_threshold=float(neg_threshold), | |
| positive_threshold=float(pos_threshold), | |
| ) | |
| except Exception as e: | |
| print(f"[ERROR] analyze_single_text / A-sentiment: {e}") | |
| traceback.print_exc() | |
| return "—", "—" | |
| tone_short = normalize_tone_label(tone_str) | |
| if tone_short == "—" and tone_probs: | |
| tone_short = choose_tone_from_probs(tone_probs) | |
| probs_block = build_probs_block(tone_probs) | |
| print(f"[DEBUG] tone_short: {tone_short!r}") | |
| print(f"[DEBUG] tone_probs: {tone_probs}") | |
| print(f"[DEBUG] probs_block:\n{probs_block}") | |
| return tone_short, probs_block | |
| def analyze_single_text_alias(text: str, neg_threshold: float, pos_threshold: float): | |
| return analyze_single_text(text, neg_threshold, pos_threshold) | |
| # ── Batch endpoint (Issue #222 / Шаг 8) ────────────────────────────────────── | |
| def analyze_batch( | |
| texts: List[str], | |
| neg_thr: float = _SENTIMENT_NEGATIVE_THRESHOLD, | |
| pos_thr: float = _SENTIMENT_POSITIVE_THRESHOLD, | |
| ) -> List[List[str]]: | |
| """Батч-анализ тональности (A). | |
| A-модель нативно работает с neg_thr/pos_thr — tone_thr не используется. | |
| HFBatchSender вызывает: | |
| client.predict(texts, neg_thr, pos_thr, api_name='/analyze_batch') | |
| Returns: | |
| [[tone_label, probs_block], ...] той же длины что и texts. | |
| """ | |
| svc = get_service() | |
| results: List[List[str]] = [] | |
| for text in texts: | |
| if not text or not str(text).strip(): | |
| results.append(["нейтрально", "—"]) | |
| continue | |
| try: | |
| tone_str, tone_probs = svc.analyze_text( | |
| str(text), | |
| negative_threshold=float(neg_thr), | |
| positive_threshold=float(pos_thr), | |
| ) | |
| tone_short = normalize_tone_label(tone_str) | |
| if tone_short == "—" and tone_probs: | |
| tone_short = choose_tone_from_probs(tone_probs) | |
| probs_block = build_probs_block(tone_probs) | |
| results.append([tone_short, probs_block]) | |
| except Exception as e: | |
| print(f"[ERROR] analyze_batch / A-sentiment (item): {e}") | |
| results.append(["ошибка", "—"]) | |
| return results | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| with gr.Blocks(title="Тональность сообщения") as demo: | |
| gr.Markdown("# Определение тональности сообщения") | |
| gr.Markdown( | |
| "Модель: `aiforever/ruroberta-large` (fine-tuned на датасете Авито, tbsa-формат) \n" | |
| f"Дефолтный порог негатива: `{_SENTIMENT_NEGATIVE_THRESHOLD}` · " | |
| f"Дефолтный порог позитива: `{_SENTIMENT_POSITIVE_THRESHOLD}`" | |
| ) | |
| inp_text = gr.Textbox( | |
| label="Текст сообщения", | |
| placeholder="Вставьте сообщение...", | |
| lines=8, | |
| ) | |
| with gr.Row(): | |
| neg_threshold_slider = gr.Slider( | |
| minimum=0.0, | |
| maximum=1.0, | |
| value=_SENTIMENT_NEGATIVE_THRESHOLD, | |
| step=0.01, | |
| label=f"Порог негатива (neg_thr, default={_SENTIMENT_NEGATIVE_THRESHOLD})", | |
| info="P(neg) >= neg_thr → негатив", | |
| ) | |
| pos_threshold_slider = gr.Slider( | |
| minimum=0.0, | |
| maximum=1.0, | |
| value=_SENTIMENT_POSITIVE_THRESHOLD, | |
| step=0.01, | |
| label=f"Порог позитива (pos_thr, default={_SENTIMENT_POSITIVE_THRESHOLD})", | |
| info="P(pos) >= pos_thr → позитив", | |
| ) | |
| btn = gr.Button("Анализировать") | |
| out_tone = gr.Textbox(label="Тональность", interactive=False) | |
| out_details = gr.Textbox( | |
| label="Подробные вероятности", | |
| interactive=False, | |
| lines=6, | |
| ) | |
| btn.click( | |
| fn=analyze_single_text, | |
| inputs=[inp_text, neg_threshold_slider, pos_threshold_slider], | |
| outputs=[out_tone, out_details], | |
| api_name="/analyze_single_text", | |
| ) | |
| gr.Button(visible=False).click( | |
| fn=analyze_single_text_alias, | |
| inputs=[inp_text, neg_threshold_slider, pos_threshold_slider], | |
| outputs=[out_tone, out_details], | |
| api_name="analyze_single_text", | |
| ) | |
| # Issue #222 / Шаг 8: батч-endpoint | |
| # HFBatchSender: client.predict(texts, neg_thr, pos_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) |