Team Ai
Apppublic

GMI-AI/Speech-Processing-Lab-Project-App

sourceHugging Facemitupdated 3mo agoView on Hugging Face
1likes
app.py396 linesDownload Raw Back to root
1import torch2import torch.nn as nn3import numpy as np4import librosa5import gradio as gr6import os7import tempfile8from datetime import datetime9from transformers import WhisperProcessor, WhisperForConditionalGeneration10 11# Essayez d'ajouter les safe globals, mais gérez l'erreur si cette fonctionnalité n'existe pas12try:13    torch.serialization.add_safe_globals(["numpy._core.multiarray._reconstruct"])14except (AttributeError, ImportError):15    print("La version de PyTorch ne prend pas en charge add_safe_globals, on continue sans.")16 17# Paramètres18SR = 16000  # Fréquence d'échantillonnage19MAX_LEN = 424  # Longueur maximale des features20 21# Définition du modèle de reconnaissance de locuteur22class SpeakerRecognitionCNN(nn.Module):23    def __init__(self, num_classes):24        super(SpeakerRecognitionCNN, self).__init__()25 26        # Premier bloc convolutionnel27        self.conv1 = nn.Sequential(28            nn.Conv2d(1, 16, kernel_size=(3, 3), stride=1, padding=1),29            nn.BatchNorm2d(16),30            nn.ReLU(),31            nn.MaxPool2d(kernel_size=(2, 2))32        )33 34        # Deuxième bloc convolutionnel35        self.conv2 = nn.Sequential(36            nn.Conv2d(16, 32, kernel_size=(3, 3), stride=1, padding=1),37            nn.BatchNorm2d(32),38            nn.ReLU(),39            nn.MaxPool2d(kernel_size=(2, 2))40        )41 42        # Troisième bloc convolutionnel43        self.conv3 = nn.Sequential(44            nn.Conv2d(32, 64, kernel_size=(3, 3), stride=1, padding=1),45            nn.BatchNorm2d(64),46            nn.ReLU(),47            nn.MaxPool2d(kernel_size=(2, 2))48        )49 50        # Quatrième bloc convolutionnel51        self.conv4 = nn.Sequential(52            nn.Conv2d(64, 128, kernel_size=(3, 3), stride=1, padding=1),53            nn.BatchNorm2d(128),54            nn.ReLU(),55            nn.MaxPool2d(kernel_size=(2, 2))56        )57 58        # Calcul dynamique de la taille après les convolutions59        self.fc_input_size = 128 * 2 * 2660 61        # Couches fully connected62        self.fc = nn.Sequential(63            nn.Linear(self.fc_input_size, 256),64            nn.ReLU(),65            nn.Dropout(0.5),66            nn.Linear(256, 128),67            nn.ReLU(),68            nn.Dropout(0.3),69            nn.Linear(128, num_classes)70        )71 72    def forward(self, x):73        # Appliquer les blocs convolutionnels74        x = self.conv1(x)75        x = self.conv2(x)76        x = self.conv3(x)77        x = self.conv4(x)78 79        # Aplatir pour les couches fully connected80        x = x.view(x.size(0), -1)81 82        # Appliquer les couches fully connected83        x = self.fc(x)84        return x85 86# Fonctions d'extraction de caractéristiques87def extract_mfcc(signal, sample_rate, n_mfcc=13, n_fft=2048, hop_length=512, include_deltas=True):88    """Extrait les MFCC d'un signal audio"""89    # Extraction des MFCC90    mfccs = librosa.feature.mfcc(y=signal,91                                 sr=sample_rate,92                                 n_mfcc=n_mfcc,93                                 n_fft=n_fft,94                                 hop_length=hop_length95                                 )96 97    # Normalisation des MFCC98    mfccs = librosa.util.normalize(mfccs, axis=1)99 100    if include_deltas:101        # Calcul des deltas (première dérivée)102        delta_mfccs = librosa.feature.delta(mfccs)103 104        # Calcul des delta-delta (seconde dérivée)105        delta2_mfccs = librosa.feature.delta(mfccs, order=2)106 107        # Concaténation des caractéristiques108        features = np.concatenate([mfccs, delta_mfccs, delta2_mfccs])109        return features110    else:111        return mfccs112 113def pad_or_truncate(features, max_len):114    """Ajuste la longueur des features à max_len"""115    if features.shape[1] < max_len:116        # ajouter des zéros à droite117        pad_width = max_len - features.shape[1]118        features_padded = np.pad(features, ((0, 0), (0, pad_width)), mode='constant')119    else:120        # couper pour correspondre à max_len121        features_padded = features[:, :max_len]122    return features_padded123 124def extract_features(audio_path):125    """Extrait les caractéristiques d'un fichier audio"""126    # Chargement du fichier audio127    y, sr = librosa.load(audio_path, sr=SR)128    129    # Utiliser la fonction extract_mfcc130    features = extract_mfcc(y, sr, n_mfcc=13, include_deltas=True)131    132    # Normalisation133    features = (features - np.mean(features)) / (np.std(features) + 1e-8)134    135    # Ajuster la longueur avec pad_or_truncate136    features = pad_or_truncate(features, MAX_LEN)137    138    # Préparation pour le modèle139    features = np.expand_dims(features, axis=0)  # Ajouter dimension batch140    features = np.expand_dims(features, axis=0)  # Ajouter dimension canal141 142    return torch.FloatTensor(features)143 144# Fonction pour l'inférence de la reconnaissance de locuteur145def predict_speaker(audio_data, model, class_names):146    """Prédit le locuteur à partir d'un fichier audio"""147    if audio_data is None:148        return "Aucun audio détecté. Veuillez enregistrer un échantillon."149    150    try:151        # Dans les nouvelles versions de Gradio, audio_data peut être:152        # - Un tuple (sr, data) si pas d'enregistrement de fichier153        # - Un chemin de fichier si enregistré154        155        if isinstance(audio_data, tuple):156            # C'est un tuple (sr, data) - sauvegardons-le temporairement157            sr, data = audio_data158            temp_dir = tempfile.gettempdir()159            temp_path = os.path.join(temp_dir, f"temp_audio_{datetime.now().strftime('%Y%m%d_%H%M%S')}.wav")160            161            # Convertir en format wav et sauvegarder162            import scipy.io.wavfile as wavfile163            wavfile.write(temp_path, sr, data)164            165            # Utiliser ce chemin temporaire166            audio_path = temp_path167        else:168            # C'est déjà un chemin de fichier169            audio_path = audio_data170        171        # Extraire les caractéristiques172        features = extract_features(audio_path)173        174        # Prédiction175        model.eval()176        with torch.no_grad():177            outputs = model(features)178            probabilities = torch.nn.functional.softmax(outputs, dim=1)179            confidence, predicted = torch.max(probabilities, 1)180            predicted_class = class_names[predicted.item()]181            confidence_value = confidence.item()182        183        # Déterminer si c'est un locuteur connu ou inconnu184        if predicted_class == "GMI" and confidence_value > 0.7:185            return f"Locuteur reconnu: GOURINE Mohammed Islem (Confiance: {confidence_value:.2f})"186        else:187            return f"Locuteur inconnu (Confiance: {confidence_value:.2f})"188            189    except Exception as e:190        return f"Erreur lors de la prédiction: {str(e)}"191    finally:192        # Nettoyer le fichier temporaire si créé193        if isinstance(audio_data, tuple) and 'temp_path' in locals() and os.path.exists(temp_path):194            os.remove(temp_path)195 196# Fonction pour la transcription audio avec Whisper197def transcribe_audio(audio_data, whisper_processor, whisper_model, language=None):198    """Transcrit l'audio en texte en utilisant le modèle Whisper"""199    if audio_data is None:200        return "Aucun audio détecté. Veuillez enregistrer un échantillon."201    202    try:203        # Gérer le format d'entrée audio204        if isinstance(audio_data, tuple):205            # C'est un tuple (sr, data)206            sr, data = audio_data207            temp_dir = tempfile.gettempdir()208            temp_path = os.path.join(temp_dir, f"temp_audio_{datetime.now().strftime('%Y%m%d_%H%M%S')}.wav")209            210            # Convertir en format wav et sauvegarder211            import scipy.io.wavfile as wavfile212            wavfile.write(temp_path, sr, data)213            214            # Utiliser ce chemin temporaire215            audio_path = temp_path216        else:217            # C'est déjà un chemin de fichier218            audio_path = audio_data219        220        # Charger l'audio avec librosa pour préparer l'input pour Whisper221        audio_array, sampling_rate = librosa.load(audio_path, sr=16000)222        223        # Préparer les inputs pour le modèle224        input_features = whisper_processor(225            audio_array, 226            sampling_rate=sampling_rate, 227            return_tensors="pt"228        ).input_features229        230        # Générer les tokens231        forced_decoder_ids = None232        if language:233            # Si une langue spécifique est demandée, la forcer234            forced_decoder_ids = whisper_processor.get_decoder_prompt_ids(235                language=language, task="transcribe"236            )237        238        # Transcription avec le modèle239        with torch.no_grad():240            predicted_ids = whisper_model.generate(241                input_features,242                forced_decoder_ids=forced_decoder_ids,243                max_length=448244            )245        246        # Décoder les tokens pour obtenir la transcription247        transcription = whisper_processor.batch_decode(248            predicted_ids, skip_special_tokens=True249        )[0]250        251        return transcription252        253    except Exception as e:254        return f"Erreur lors de la transcription: {str(e)}"255    finally:256        # Nettoyer le fichier temporaire si créé257        if isinstance(audio_data, tuple) and 'temp_path' in locals() and os.path.exists(temp_path):258            try:259                os.remove(temp_path)260            except:261                pass262 263# Charger le modèle Whisper264def load_whisper_model():265    """Charge le modèle Whisper tiny"""266    try:267        # Nom du modèle pour Whisper Large V3 Turbo268        model_id = "openai/whisper-tiny"269        270        # Charger le processeur et le modèle271        processor = WhisperProcessor.from_pretrained(model_id)272        model = WhisperForConditionalGeneration.from_pretrained(model_id)273        274        print("Modèle Whisper chargé avec succès!")275        return processor, model276    except Exception as e:277        print(f"Erreur lors du chargement du modèle Whisper: {e}")278        return None, None279 280# Fonction principale pour l'application Gradio281def create_gradio_app():282    # Charger le modèle de reconnaissance de locuteur283    try:284        # Essayer d'abord sans weights_only (sécurité PyTorch 2.6)285        try:286            model_path = "GMI-Speech-recognition-CNN-1.0.pth"287            checkpoint = torch.load(model_path, map_location=torch.device('cpu'), weights_only=False)288        except TypeError:289            # Si weights_only n'est pas un paramètre valide (versions antérieures de PyTorch)290            checkpoint = torch.load(model_path, map_location=torch.device('cpu'))291        292        # Récupérer les noms de classes293        if isinstance(checkpoint, dict) and 'class_names' in checkpoint:294            class_names = checkpoint['class_names']295        else:296            class_names = ["GMI", "Unknown"]297            print("Noms de classes non trouvés, utilisation de valeurs par défaut.")298        299        # Initialiser le modèle300        speaker_model = SpeakerRecognitionCNN(len(class_names))301        302        # Charger les poids du modèle303        if 'model_state_dict' in checkpoint:304            speaker_model.load_state_dict(checkpoint['model_state_dict'])305        else:306            speaker_model.load_state_dict(checkpoint)307        308        speaker_model.eval()309        print("Modèle de reconnaissance de locuteur chargé avec succès!")310        311    except Exception as e:312        print(f"Erreur lors du chargement du modèle de reconnaissance: {e}")313        # Créer un modèle factice pour la démonstration314        speaker_model = SpeakerRecognitionCNN(2)315        class_names = ["GMI", "Unknown"]316        print("Utilisation d'un modèle de démonstration pour la reconnaissance de locuteur.")317    318    # Charger le modèle Whisper319    whisper_processor, whisper_model = load_whisper_model()320    321    # Créer l'interface Gradio322    with gr.Blocks(title="Laboratoire de Traitement de la Parole") as app:323        gr.Markdown("# Laboratoire de Traitement de la Parole")324        325        with gr.Tab("Reconnaissance de Locuteur"):326            gr.Markdown("## Reconnaissance de Locuteur")327            gr.Markdown("Enregistrez votre voix pour vérifier si vous êtes GOURINE Mohammed Islem ou un locuteur inconnu.")328            329            with gr.Row():330                with gr.Column(scale=1):331                    speaker_audio_input = gr.Audio(332                        label="Enregistrez votre voix"333                    )334                    335                with gr.Column(scale=1):336                    speaker_result_output = gr.Textbox(label="Résultat")337                    recognize_btn = gr.Button("Reconnaître le locuteur")338            339            # Connecter le bouton à la fonction de prédiction340            recognize_btn.click(341                fn=lambda x: predict_speaker(x, speaker_model, class_names),342                inputs=speaker_audio_input,343                outputs=speaker_result_output344            )345               # Ajouter des exemples pour la reconnaissance de locuteur346            gr.Markdown("### Exemples de la voix de GOURINE Mohammed Islem")347            gr.Examples(348                examples=[349                    ["exemples/gmi_voix1.wav"],350                    ["exemples/gmi_voix2.wav"],351                  352                ],353                inputs=speaker_audio_input,354                outputs=speaker_result_output,355                fn=lambda x: predict_speaker(x, speaker_model, class_names),356                cache_examples=True,357            )358        359        with gr.Tab("Reconnaissance Vocale (Whisper)"):360            gr.Markdown("## Reconnaissance Vocale avec Whisper")361            gr.Markdown("Enregistrez ou téléchargez un fichier audio pour le transcrire en texte.")362            363            with gr.Row():364                with gr.Column(scale=1):365                    whisper_audio_input = gr.Audio(366                        label="Audio à transcrire"367                    )368                    369                    language_dropdown = gr.Dropdown(370                        choices=["auto", "français", "english", "arabic", "español", "deutsch", "italiano"],371                        value="auto",372                        label="Langue de l'audio"373                    )374                    375                with gr.Column(scale=1):376                    transcript_output = gr.Textbox(label="Transcription", lines=10)377                    transcribe_btn = gr.Button("Transcrire l'audio")378            379            # Connecter le bouton à la fonction de transcription380            transcribe_btn.click(381                fn=lambda x, lang: transcribe_audio(382                    x, 383                    whisper_processor, 384                    whisper_model,385                    None if lang == "auto" else lang386                ),387                inputs=[whisper_audio_input, language_dropdown],388                outputs=transcript_output389            )390    391    return app392 393# Lancement de l'application394if __name__ == "__main__":395    app = create_gradio_app()396    app.launch()