feat(servers): wrappers Kyutai/Cohere locaux + fix chargement Cohere#44
feat(servers): wrappers Kyutai/Cohere locaux + fix chargement Cohere#44benoitvx wants to merge 3 commits into
Conversation
…oading Add the two local OpenAI-compatible wrapper servers (referenced by the main README's Kyutai/Cohere sections) under servers/, so the benchmark's local models are versioned and reproducible: - kyutai_server.py : Kyutai stt-1b-en_fr, 30s silence-aligned chunking - cohere_server.py : Cohere Transcribe 03-2026 Cohere fix: load via the native CohereAsrForConditionalGeneration class instead of the remote-code path (AutoModelForSpeechSeq2Seq + trust_remote_code). Under transformers 5.9 the remote-code path misapplies generation_config (decoder_start_token_id) → multilingual garbage that ignores the audio. The native class applies the config correctly. float32 kept (the -1e9 attention mask overflows float16). Co-Authored-By: Claude Opus 4.8 (1M context) <[email protected]>
There was a problem hiding this comment.
Code Review
This pull request introduces two minimal, OpenAI-compatible HTTP servers wrapping local ASR models (Kyutai STT and Cohere Transcribe) for benchmarking, along with documentation and dependency configuration. Feedback on the implementation highlights a critical type mismatch in kyutai_server.py when using float16 on MPS, which requires explicitly casting input tensors to the model's dtype. Additionally, both servers define their transcription endpoints as asynchronous despite performing blocking CPU-bound operations; these should be refactored to synchronous functions with robust temporary file cleanup to prevent event loop blocking and resource leaks. Finally, it is recommended to move generated tokens to the CPU before decoding to avoid potential device compatibility issues.
Important
The consumer version of Gemini Code Assist on GitHub is being sunset. Starting June 18, 2026, new organization installations will be blocked, and all code review activity will officially cease on July 17, 2026.
For more details on the timeline and next steps, please review the Help Documentation.
| if len(chunk) < min_seg: | ||
| continue | ||
| inputs = processor(audio=chunk, sampling_rate=SR, return_tensors="pt") | ||
| inputs = inputs.to(DEVICE) |
There was a problem hiding this comment.
Lorsque le modèle est chargé en float16 (sur MPS), les tenseurs d'entrée renvoyés par le processeur (qui sont en float32 par défaut) doivent être convertis dans le même type que le modèle pour éviter une erreur de type (RuntimeError: Input type and weight type should be the same). Il est nécessaire de passer explicitement dtype=model.dtype lors de l'appel à .to().
| inputs = inputs.to(DEVICE) | |
| inputs = inputs.to(DEVICE, dtype=model.dtype) |
| @app.post("/v1/audio/transcriptions") | ||
| async def transcribe( | ||
| file: UploadFile = File(...), | ||
| model: str | None = Form(None), | ||
| language: str | None = Form(None), | ||
| response_format: str | None = Form(None), | ||
| ): | ||
| suffix = Path(file.filename or "audio.wav").suffix or ".wav" | ||
| with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp: | ||
| tmp.write(await file.read()) | ||
| tmp_path = tmp.name | ||
| try: | ||
| text = transcribe_path(tmp_path) | ||
| finally: | ||
| os.unlink(tmp_path) | ||
| return {"text": text} |
There was a problem hiding this comment.
La fonction transcribe est définie comme async def, mais elle appelle transcribe_path qui effectue des calculs lourds et bloquants (chargement audio, inférence). Cela bloque complètement la boucle d'événements de FastAPI, empêchant le serveur de traiter d'autres requêtes en parallèle. De plus, si une exception survient pendant await file.read(), le fichier temporaire créé avec delete=False ne sera jamais supprimé, ce qui entraîne une fuite de ressources. Définir la route comme une fonction synchrone (def) permet à FastAPI de l'exécuter automatiquement dans un thread séparé, et l'utilisation d'un bloc try...finally robuste garantit la suppression du fichier temporaire.
| @app.post("/v1/audio/transcriptions") | |
| async def transcribe( | |
| file: UploadFile = File(...), | |
| model: str | None = Form(None), | |
| language: str | None = Form(None), | |
| response_format: str | None = Form(None), | |
| ): | |
| suffix = Path(file.filename or "audio.wav").suffix or ".wav" | |
| with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp: | |
| tmp.write(await file.read()) | |
| tmp_path = tmp.name | |
| try: | |
| text = transcribe_path(tmp_path) | |
| finally: | |
| os.unlink(tmp_path) | |
| return {"text": text} | |
| @app.post("/v1/audio/transcriptions") | |
| def transcribe( | |
| file: UploadFile = File(...), | |
| model: str | None = Form(None), | |
| language: str | None = Form(None), | |
| response_format: str | None = Form(None), | |
| ): | |
| suffix = Path(file.filename or "audio.wav").suffix or ".wav" | |
| tmp = tempfile.NamedTemporaryFile(delete=False, suffix=suffix) | |
| tmp_path = tmp.name | |
| try: | |
| with tmp: | |
| tmp.write(file.file.read()) | |
| text = transcribe_path(tmp_path) | |
| finally: | |
| try: | |
| os.unlink(tmp_path) | |
| except OSError: | |
| pass | |
| return {"text": text} |
| @app.post("/v1/audio/transcriptions") | ||
| async def transcribe( | ||
| file: UploadFile = File(...), | ||
| model: str | None = Form(None), | ||
| language: str | None = Form(None), | ||
| response_format: str | None = Form(None), | ||
| ): | ||
| suffix = Path(file.filename or "audio.wav").suffix or ".wav" | ||
| with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp: | ||
| tmp.write(await file.read()) | ||
| tmp_path = tmp.name | ||
| try: | ||
| text = transcribe_path(tmp_path, language or DEFAULT_LANGUAGE) | ||
| finally: | ||
| os.unlink(tmp_path) | ||
| return {"text": text} |
There was a problem hiding this comment.
La fonction transcribe est définie comme async def, mais elle appelle transcribe_path qui effectue des calculs lourds et bloquants. Cela bloque la boucle d'événements de FastAPI. De plus, si une exception survient pendant await file.read(), le fichier temporaire créé avec delete=False ne sera pas supprimé. Définir la route comme synchrone (def) et utiliser un bloc try...finally robuste résout ces deux problèmes.
| @app.post("/v1/audio/transcriptions") | |
| async def transcribe( | |
| file: UploadFile = File(...), | |
| model: str | None = Form(None), | |
| language: str | None = Form(None), | |
| response_format: str | None = Form(None), | |
| ): | |
| suffix = Path(file.filename or "audio.wav").suffix or ".wav" | |
| with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp: | |
| tmp.write(await file.read()) | |
| tmp_path = tmp.name | |
| try: | |
| text = transcribe_path(tmp_path, language or DEFAULT_LANGUAGE) | |
| finally: | |
| os.unlink(tmp_path) | |
| return {"text": text} | |
| @app.post("/v1/audio/transcriptions") | |
| def transcribe( | |
| file: UploadFile = File(...), | |
| model: str | None = Form(None), | |
| language: str | None = Form(None), | |
| response_format: str | None = Form(None), | |
| ): | |
| suffix = Path(file.filename or "audio.wav").suffix or ".wav" | |
| tmp = tempfile.NamedTemporaryFile(delete=False, suffix=suffix) | |
| tmp_path = tmp.name | |
| try: | |
| with tmp: | |
| tmp.write(file.file.read()) | |
| text = transcribe_path(tmp_path, language or DEFAULT_LANGUAGE) | |
| finally: | |
| try: | |
| os.unlink(tmp_path) | |
| except OSError: | |
| pass | |
| return {"text": text} |
| inputs = processor(audio=chunk, sampling_rate=SR, return_tensors="pt") | ||
| inputs = inputs.to(DEVICE) | ||
| with torch.no_grad(): | ||
| output_tokens = model.generate(**inputs) |
There was a problem hiding this comment.
Il est recommandé de déplacer explicitement les jetons générés (output_tokens) vers le CPU avant de les passer à processor.batch_decode. Certains décodeurs de Hugging Face peuvent lever des erreurs ou être inefficaces s'ils reçoivent des tenseurs situés sur un périphérique non-CPU (comme MPS).
| output_tokens = model.generate(**inputs) | |
| output_tokens = model.generate(**inputs).to("cpu") |
| audio_chunk_index = inputs.get("audio_chunk_index") | ||
| inputs = inputs.to(model.device, dtype=model.dtype) | ||
| with torch.no_grad(): | ||
| outputs = model.generate(**inputs, max_new_tokens=MAX_NEW_TOKENS) |
There was a problem hiding this comment.
Il est recommandé de déplacer explicitement les tenseurs générés (outputs) vers le CPU avant de les passer à processor.decode pour éviter des erreurs de compatibilité de périphérique avec certains décodeurs de Hugging Face.
| outputs = model.generate(**inputs, max_new_tokens=MAX_NEW_TOKENS) | |
| outputs = model.generate(**inputs, max_new_tokens=MAX_NEW_TOKENS).to("cpu") |
…load lock - pyproject: transformers>=5.3 (la classe native CohereAsr requise par le fix n'existe pas en 4.53) ; renomme le projet en eval-transcript-servers (héberge désormais les 2 wrappers, plus seulement Kyutai). - cohere_server: docstring corrigée (classe native, plus de trust_remote_code). - les 2 serveurs: verrou (threading.Lock) sur le lazy-load pour éviter un double chargement concurrent (pic mémoire/OOM) ; commentaire sur les params OpenAI acceptés mais ignorés (model/response_format). - README: transformers ≥ 5.3. Co-Authored-By: Claude Opus 4.8 (1M context) <[email protected]>
🔎 Revue de code (Claude Code)PR additive pure (+309/−0, 4 fichiers neufs, aucune ligne du package principal touchée) → risque d'intégration nul. Versionne les 2 wrappers serveurs locaux (Kyutai/Cohere) + le fix de chargement Cohere. 🔴 À corriger (correctness / repro)
🟡 À considérer
🟢 Points positifs
Tests & sécurité
✅ Corrections appliquées (commit
|
- HIGH: endpoints async def → def (FastAPI les exécute en threadpool, plus de
blocage de l'event loop) + cleanup temp file robuste en try/finally (plus de
fuite si la lecture du body échoue). Les 2 serveurs.
- MEDIUM: rapatrier les tokens générés sur CPU avant decode (compat décodeurs
HF). Les 2 serveurs. Vérifié sur clip FR.
- CRITICAL (cast dtype Kyutai) NON appliqué : faux positif. Kyutai charge en
float16 mais garde des biais conv en float32 ; forcer les inputs en float16
casse l'encodeur ("Input type (c10::Half) and bias type (float)..."). Les
inputs float32 d'origine sont corrects (confirmé par smoke-test). Commenté.
Co-Authored-By: Claude Opus 4.8 (1M context) <[email protected]>
Réponse aux findings Gemini (commit
|
Contexte
Les modèles locaux du bench (Kyutai, Cohere) tournent derrière des petits serveurs OpenAI-compatibles, jusqu'ici non versionnés (
~/Dev/kyutai-server). LeREADME.mdprincipal les documente sans les fournir. Cette PR les intègre sousservers/pour rendre le bench des modèles locaux reproductible.Contenu
servers/kyutai_server.py— wrapper Kyutaistt-1b-en_fr(chunking ~30 s aligné silences,generate()neuf par segment).servers/cohere_server.py— wrapper Cohere Transcribe 03-2026.servers/pyproject.toml,servers/README.md.Fix Cohere (le cœur de la PR)
Chargement via la classe native
CohereAsrForConditionalGenerationau lieu du chemin remote-code (AutoModelForSpeechSeq2Seq+trust_remote_code).Symptôme : sous
transformers 5.9, le chemin remote-code applique mal lageneration_config(decoder_start_token_id) → le modèle sort du texte multilingue aberrant en ignorant l'audio (reproduit sur CPU comme MPS).Correctif : la classe native (intégrée à transformers ≥ 5.x) applique correctement la config → transcription attendue. Validé sur extrait FR.
float32conservé (le masque d'attention à-1e9débordefloat16).Test
Bench complet rejoué sur les 8 samples du corpus public : Cohere produit des transcriptions FR correctes (8/8), poussées sur
AgentPublic/eval-stt-results.🤖 Generated with Claude Code