CoolFace
Apppublic

RobotsMali/RobotsMali_ASR_DEMO

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
app.py123 linesDownload Raw Back to root
1import os, shlex, subprocess, tempfile, traceback, glob, gc, shutil2import time3import psutil4import humanize5import gradio as gr6 7# Les imports lourds (torch, nemo) sont déplacés dans les fonctions8# pour assurer un démarrage léger de l'Espace.9 10# Configuration11DEVICE = "cpu" # Forcé CPU pour stabilité au démarrage, sera vérifié dynamiquement12MODEL_CACHE = {}13 14def get_model(name, repo, arch_type):15    """Téléchargement et chargement différé d'un modèle NeMo"""16    try:17        # Patch pour éviter l'erreur de compatibilité ConfidenceConfig avec tdt_include_duration18        try:19            import nemo.collections.asr.parts.utils.asr_confidence_utils as asr_confidence_utils20            import inspect21            original_init = asr_confidence_utils.ConfidenceConfig.__init__22            def patched_init(self, *args, **kwargs):23                try:24                    valid_keys = set(inspect.signature(asr_confidence_utils.ConfidenceConfig).parameters.keys())25                    filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_keys}26                except Exception:27                    filtered_kwargs = kwargs28                original_init(self, *args, **filtered_kwargs)29            asr_confidence_utils.ConfidenceConfig.__init__ = patched_init30            print("🩹 Successfully patched NeMo ConfidenceConfig.__init__")31        except Exception as e:32            print(f"⚠️ Could not patch NeMo ConfidenceConfig: {e}")33 34        import torch35        from huggingface_hub import snapshot_download36        from nemo.collections.asr.models import EncDecCTCModel, EncDecRNNTModel, EncDecHybridRNNTCTCBPEModel37        38        device = "cuda" if torch.cuda.is_available() else "cpu"39        token = os.environ.get("HF_TOKEN")40        41        print(f"📥 Folder download: {repo}...")42        folder = snapshot_download(repo, local_dir_use_symlinks=False, token=token)43        nemo_file = glob.glob(os.path.join(folder, "*.nemo"))[0]44        45        print(f"📥 Loading model {name} on {device}...")46        if arch_type == "ctc":47            model = EncDecCTCModel.restore_from(nemo_file, map_location=torch.device(device))48        elif "soloni" in repo.lower():49            model = EncDecHybridRNNTCTCBPEModel.restore_from(nemo_file, map_location=torch.device(device))50        else:51            model = EncDecRNNTModel.restore_from(nemo_file, map_location=torch.device(device))52            53        model.eval()54        return model55    except Exception as e:56        print(f"❌ Erreur lors du chargement du modèle {name}:")57        print(traceback.format_exc())58        return None59 60MODELS = {61    "Soloba V3 (CTC)":           ("RobotsMali/soloba-ctc-0.6b-v3", "ctc"),62    "Soloni V3 (TDT-CTC)":       ("RobotsMali/soloni-114m-tdt-ctc-v3", "hybrid"),63}64 65def pipeline(audio_in, model_name, progress=gr.Progress()):66    if not audio_in:67        return "⚠️ Erreur", None, "Veuillez fournir un fichier audio."68    69    try:70        import torch71        import psutil72        73        # Nettoyage mémoire préventif74        gc.collect()75        if torch.cuda.is_available():76            torch.cuda.empty_cache()77            78        repo, arch_type = MODELS[model_name]79        80        # Chargement à la demande81        progress(0.2, desc="Chargement du modèle...")82        model = get_model(model_name, repo, arch_type)83        84        if model is None:85            return "❌ Erreur", None, f"Impossible de charger le modèle {model_name}. Vérifiez les logs."86 87        progress(0.5, desc="Transcription en cours...")88        transcription = model.transcribe([audio_in])[0]89        90        # Nettoyage post-exécution91        del model92        gc.collect()93        if torch.cuda.is_available():94            torch.cuda.empty_cache()95            96        return "✅ Succès", None, transcription97    except Exception as e:98        return "❌ Erreur", None, str(e)99 100# UI101with gr.Blocks(title="RobotsMali ASR", theme=gr.themes.Soft()) as demo:102    gr.Markdown("# 🎙️ RobotsMali ASR Demo")103    gr.Markdown("Transcription automatique de la parole pour le Bambara et les langues du Mali.")104    105    with gr.Row():106        with gr.Column():107            audio_input = gr.Audio(sources=["upload", "microphone"], type="filepath", label="Audio")108            model_dropdown = gr.Dropdown(choices=list(MODELS.keys()), value="Soloni V3 (TDT-CTC)", label="Modèle")109            submit_btn = gr.Button("Transcrire", variant="primary")110            111        with gr.Column():112            status_out = gr.Textbox(label="Statut")113            text_output = gr.Textbox(label="Transcription", lines=10)114 115    submit_btn.click(116        fn=pipeline,117        inputs=[audio_input, model_dropdown],118        outputs=[status_out, gr.State(), text_output]119    )120 121if __name__ == "__main__":122    print("🚀 Démarrage de l'interface Gradio...")123    demo.queue().launch(server_name="0.0.0.0", server_port=7860)