RobotsMali/RobotsMali_ASR_DEMO
0
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)