Spaces:
Runtime error
Runtime error
| """Pipeline de production : ASR Wav2Vec2 puis analyse de sentiment à trois classes.""" | |
| from __future__ import annotations | |
| from dataclasses import asdict, dataclass # Structure la réponse sans dictionnaire ambigu. | |
| from pathlib import Path # Accepte aussi bien une chaîne qu'un chemin système. | |
| from threading import Lock # Empêche deux requêtes simultanées de charger deux fois les poids. | |
| import numpy as np | |
| import torch | |
| # Wav2Vec2Processor est explicite : AutoProcessor sélectionne ici un décodeur avec | |
| # language model qui demanderait pyctcdecode, alors que notre code utilise argmax CTC. | |
| from transformers import AutoModelForCTC, AutoModelForSequenceClassification, AutoTokenizer, Wav2Vec2Processor | |
| from src.audio import load_and_preprocess_audio | |
| from src.errors import PipelineError | |
| # Modèle XLSR français recommandé dans l'énoncé ; il accepte les signaux 16 kHz. | |
| DEFAULT_ASR_MODEL = "jonatasgrosman/wav2vec2-large-xlsr-53-french" | |
| # XLM-RoBERTa est une variante BERT multilingue fine-tunée sur trois sentiments. | |
| DEFAULT_SENTIMENT_MODEL = "cardiffnlp/twitter-xlm-roberta-base-sentiment" | |
| # Les cartes de labels diffèrent légèrement selon la version/configuration du modèle. | |
| # Cette table convertit toujours le résultat final en français et dans les trois classes imposées. | |
| LABELS = { | |
| "0": "négatif", | |
| "1": "neutre", | |
| "2": "positif", | |
| "negative": "négatif", | |
| "neutral": "neutre", | |
| "positive": "positif", | |
| } | |
| class Prediction: | |
| """Format unique retourné par Gradio et par l'API REST.""" | |
| # Texte produit par l'étape ASR, affiché à l'utilisateur pour rendre le pipeline explicable. | |
| transcription: str | |
| # Une des trois classes finales : positif, négatif ou neutre. | |
| sentiment: str | |
| # Probabilité softmax de la classe retenue, arrondie seulement à la sortie. | |
| confidence: float | |
| def to_dict(self) -> dict[str, str | float]: | |
| # FastAPI sérialise naturellement un dictionnaire au format JSON. | |
| return asdict(self) | |
| class SentimentCallPipeline: | |
| """Charge les modèles une seule fois et expose une prédiction déterministe.""" | |
| def __init__( | |
| self, | |
| asr_model_name: str = DEFAULT_ASR_MODEL, | |
| sentiment_model_name: str = DEFAULT_SENTIMENT_MODEL, | |
| ) -> None: | |
| # Garder les identifiants en attribut rend possible de les remplacer dans une expérience. | |
| self.asr_model_name = asr_model_name | |
| self.sentiment_model_name = sentiment_model_name | |
| # CUDA est utilisé automatiquement lorsqu'il est disponible ; sinon CPU. | |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| # Aucun poids n'est chargé au démarrage : Gradio/API reste réactif jusqu'à la première prédiction. | |
| self._loaded = False | |
| # Le verrou est important quand FastAPI reçoit plusieurs requêtes simultanées. | |
| self._load_lock = Lock() | |
| def _ensure_models_loaded(self) -> None: | |
| """Télécharge/instancie paresseusement les poids pour accélérer le démarrage.""" | |
| # Cas habituel après la première requête : on réutilise les modèles déjà présents en mémoire. | |
| if self._loaded: | |
| return | |
| with self._load_lock: | |
| # Une seconde vérification est nécessaire : une autre requête a pu finir le chargement | |
| # pendant que celle-ci attendait le verrou. | |
| if self._loaded: | |
| return | |
| try: | |
| # Charge uniquement tokenizer + feature extractor. Le décodage CTC glouton | |
| # est exécuté plus bas avec argmax : pyctcdecode n'est donc pas nécessaire. | |
| self.asr_processor = Wav2Vec2Processor.from_pretrained(self.asr_model_name) | |
| # eval() désactive dropout ; la prédiction devient stable et adaptée à l'inférence. | |
| self.asr_model = AutoModelForCTC.from_pretrained(self.asr_model_name).to(self.device).eval() | |
| # XLM-RoBERTa repose sur SentencePiece. `use_fast=False` sélectionne le tokenizer | |
| # Python stable, ce qui évite une incompatibilité du tokenizer rapide avec certaines | |
| # versions récentes de Transformers/Python. Il est suffisamment rapide ici (un appel). | |
| self.sentiment_tokenizer = AutoTokenizer.from_pretrained( | |
| self.sentiment_model_name, use_fast=False | |
| ) | |
| self.sentiment_model = AutoModelForSequenceClassification.from_pretrained( | |
| self.sentiment_model_name | |
| ).to(self.device).eval() | |
| # Le drapeau n'est activé qu'après le chargement complet des deux étapes. | |
| self._loaded = True | |
| except Exception as exc: | |
| raise PipelineError("Chargement des modèles impossible. Vérifiez la connexion et les dépendances.") from exc | |
| def transcribe(self, audio: np.ndarray) -> str: | |
| """Décode les logits CTC de Wav2Vec2 en une transcription française.""" | |
| # Garantit que le processor et le modèle existent même lors du premier appel. | |
| self._ensure_models_loaded() | |
| # `return_tensors="pt"` crée un tenseur PyTorch ; padding facilite une future extension en lot. | |
| inputs = self.asr_processor(audio, sampling_rate=16_000, return_tensors="pt", padding=True) | |
| # Les logits sont les scores par caractère/token. Ils sont calculés sur GPU si disponible. | |
| logits = self.asr_model(inputs.input_values.to(self.device)).logits | |
| # Le décodage CTC simple retient le token le plus probable à chaque instant (argmax). | |
| token_ids = torch.argmax(logits, dim=-1) | |
| # batch_decode fusionne les répétitions et tokens blancs propres à CTC, puis rend le texte. | |
| text = self.asr_processor.batch_decode(token_ids)[0].strip() | |
| if not text: | |
| raise PipelineError("Aucune transcription n'a pu être produite pour cet audio.") | |
| return text | |
| def classify_sentiment(self, text: str) -> tuple[str, float]: | |
| """Retourne la classe à trois valeurs et la probabilité softmax associée.""" | |
| # Le modèle de sentiment peut être appelé seul lors de tests, d'où cette vérification. | |
| self._ensure_models_loaded() | |
| # Les entrées longues sont tronquées à 512 tokens : limite architecturale de BERT/XLM-R. | |
| encoded = self.sentiment_tokenizer(text, truncation=True, max_length=512, return_tensors="pt") | |
| # Déplacer tous les champs (input_ids, attention_mask, etc.) vers le même périphérique que le modèle. | |
| logits = self.sentiment_model(**{key: value.to(self.device) for key, value in encoded.items()}).logits | |
| # Softmax convertit les trois scores bruts en probabilités dont la somme vaut 1. | |
| probabilities = torch.softmax(logits, dim=-1)[0] | |
| # L'indice de la plus grande probabilité est la classe prédite. | |
| index = int(torch.argmax(probabilities).item()) | |
| raw_label = str(self.sentiment_model.config.id2label.get(index, index)).lower() | |
| # Les modèles Cardiff utilisent LABEL_0/1/2 ; la table garantit le contrat FR. | |
| sentiment = LABELS.get(raw_label.replace("label_", ""), raw_label) | |
| return sentiment, round(float(probabilities[index].item()), 4) | |
| def predict(self, audio_path: str | Path) -> Prediction: | |
| """Exécute la chaîne complète audio -> texte -> sentiment.""" | |
| # Chaque interface passe ici : elles bénéficient donc exactement des mêmes contrôles audio. | |
| audio = load_and_preprocess_audio(audio_path) | |
| # La transcription est d'abord calculée, car elle est l'entrée du classifieur de sentiment. | |
| transcription = self.transcribe(audio) | |
| # Le couple final contient la classe et la confiance associée à cette classe. | |
| sentiment, confidence = self.classify_sentiment(transcription) | |
| return Prediction(transcription, sentiment, confidence) | |