DL_FINAL_EXAMEN / src /pipeline.py
Your NameBOLLO22
Ajout du projet complet de BOLLO
c6b0fdb
Raw
History Blame Contribute Delete
8.05 kB
"""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)