"""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", } @dataclass(frozen=True) 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 @torch.inference_mode() 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 @torch.inference_mode() 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)