github-actions[bot] commited on
Commit
bded519
·
1 Parent(s): 5f74901

Deploy from Achraf-cyber/hackton-locallang@76aa2de3965e28e4be5261c706fe0cd7a30d1778

Browse files
.gitattributes DELETED
@@ -1,35 +0,0 @@
1
- *.7z filter=lfs diff=lfs merge=lfs -text
2
- *.arrow filter=lfs diff=lfs merge=lfs -text
3
- *.bin filter=lfs diff=lfs merge=lfs -text
4
- *.bz2 filter=lfs diff=lfs merge=lfs -text
5
- *.ckpt filter=lfs diff=lfs merge=lfs -text
6
- *.ftz filter=lfs diff=lfs merge=lfs -text
7
- *.gz filter=lfs diff=lfs merge=lfs -text
8
- *.h5 filter=lfs diff=lfs merge=lfs -text
9
- *.joblib filter=lfs diff=lfs merge=lfs -text
10
- *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
- *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
- *.model filter=lfs diff=lfs merge=lfs -text
13
- *.msgpack filter=lfs diff=lfs merge=lfs -text
14
- *.npy filter=lfs diff=lfs merge=lfs -text
15
- *.npz filter=lfs diff=lfs merge=lfs -text
16
- *.onnx filter=lfs diff=lfs merge=lfs -text
17
- *.ot filter=lfs diff=lfs merge=lfs -text
18
- *.parquet filter=lfs diff=lfs merge=lfs -text
19
- *.pb filter=lfs diff=lfs merge=lfs -text
20
- *.pickle filter=lfs diff=lfs merge=lfs -text
21
- *.pkl filter=lfs diff=lfs merge=lfs -text
22
- *.pt filter=lfs diff=lfs merge=lfs -text
23
- *.pth filter=lfs diff=lfs merge=lfs -text
24
- *.rar filter=lfs diff=lfs merge=lfs -text
25
- *.safetensors filter=lfs diff=lfs merge=lfs -text
26
- saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
- *.tar.* filter=lfs diff=lfs merge=lfs -text
28
- *.tar filter=lfs diff=lfs merge=lfs -text
29
- *.tflite filter=lfs diff=lfs merge=lfs -text
30
- *.tgz filter=lfs diff=lfs merge=lfs -text
31
- *.wasm filter=lfs diff=lfs merge=lfs -text
32
- *.xz filter=lfs diff=lfs merge=lfs -text
33
- *.zip filter=lfs diff=lfs merge=lfs -text
34
- *.zst filter=lfs diff=lfs merge=lfs -text
35
- *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
Dockerfile ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM python:3.11-slim
2
+
3
+ RUN apt-get update && apt-get install -y --no-install-recommends ffmpeg \
4
+ && rm -rf /var/lib/apt/lists/*
5
+
6
+ WORKDIR /app
7
+
8
+ COPY requirements.txt .
9
+ RUN pip install --no-cache-dir -r requirements.txt
10
+
11
+ COPY . .
12
+
13
+ EXPOSE 7860
14
+
15
+ CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "7860"]
README.md CHANGED
@@ -1,10 +1,44 @@
1
  ---
2
- title: Model Service
3
- emoji: 👀
4
- colorFrom: pink
5
- colorTo: purple
6
  sdk: docker
7
- pinned: false
8
  ---
9
 
10
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ title: model-service
3
+ emoji: 🗣️
4
+ colorFrom: blue
5
+ colorTo: green
6
  sdk: docker
7
+ app_port: 7860
8
  ---
9
 
10
+ # model-service
11
+
12
+ Service Python FastAPI dont le seul role est d'exposer l'acces aux langues
13
+ Dioula et Moore : reconnaissance vocale (ASR), traduction et synthese vocale
14
+ (TTS). Aucune logique de LLM, de base de donnees ou de simplification n'est
15
+ geree ici.
16
+
17
+ ## Structure
18
+
19
+ ```
20
+ model-service/
21
+ ├── app/
22
+ │ ├── main.py # point d'entree FastAPI
23
+ │ ├── deps.py # configuration (Settings)
24
+ │ └── services/
25
+ │ ├── asr.py # reconnaissance vocale
26
+ │ ├── translation.py # traduction
27
+ │ └── tts.py # synthese vocale
28
+ └── tests/
29
+ ```
30
+
31
+ ## Developpement local
32
+
33
+ ```bash
34
+ pip install -r requirements.txt
35
+ uvicorn app.main:app --reload --port 8000
36
+ ```
37
+
38
+ ## Deploiement sur Hugging Face Spaces
39
+
40
+ Ce dossier est deploye comme un Space Docker (header YAML ci-dessus :
41
+ `sdk: docker`, `app_port: 7860`), automatiquement a chaque push sur `main`
42
+ via le workflow `.github/workflows/deploy-model-service.yml` (voir la racine
43
+ du monorepo). Le Space lit `ALLOWED_ORIGINS` et `HF_TOKEN` depuis ses propres
44
+ "Secrets" (Settings > Variables and secrets du Space), pas depuis ce repo.
app/__init__.py ADDED
File without changes
app/deps.py ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from functools import lru_cache
2
+ from typing import Literal
3
+
4
+ from pydantic_settings import BaseSettings, SettingsConfigDict
5
+
6
+
7
+ class Settings(BaseSettings):
8
+ model_config = SettingsConfigDict(env_file=".env", env_file_encoding="utf-8")
9
+
10
+ ALLOWED_ORIGINS: list[str] = ["*"]
11
+
12
+ # ASR temporaire via l'API d'inference Hugging Face pendant que
13
+ # facebook/mms-1b-all finit de telecharger en local (voir asr.py).
14
+ ASR_BACKEND: Literal["local", "hf_api"] = "local"
15
+ HF_TOKEN: str | None = None
16
+
17
+
18
+ @lru_cache
19
+ def get_settings() -> Settings:
20
+ return Settings()
app/main.py ADDED
@@ -0,0 +1,127 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """model-service - FastAPI app.
2
+
3
+ Lancement local :
4
+ uvicorn app.main:app --reload --port 8000
5
+ """
6
+
7
+ import logging
8
+ import time
9
+ import uuid
10
+ from pathlib import Path
11
+ from typing import Literal
12
+
13
+ from fastapi import FastAPI, File, Form, UploadFile
14
+ from fastapi.middleware.cors import CORSMiddleware
15
+ from fastapi.staticfiles import StaticFiles
16
+ from pydantic import BaseModel
17
+
18
+ from app.deps import get_settings
19
+ from app.services.asr import ASR
20
+ from app.services.translator import Translator
21
+ from app.services.tts import TTS
22
+
23
+ logging.basicConfig(level=logging.INFO)
24
+ logger = logging.getLogger("model-service")
25
+
26
+ settings = get_settings()
27
+
28
+ MEDIA_DIR = Path(__file__).resolve().parent.parent / "media"
29
+ MEDIA_DIR.mkdir(parents=True, exist_ok=True)
30
+
31
+ app = FastAPI(
32
+ title="model-service",
33
+ description="Expose ASR, traduction et TTS pour le Dioula et le Mooré.",
34
+ )
35
+
36
+ app.add_middleware(
37
+ CORSMiddleware,
38
+ allow_origins=settings.ALLOWED_ORIGINS,
39
+ allow_credentials=True,
40
+ allow_methods=["*"],
41
+ allow_headers=["*"],
42
+ )
43
+
44
+ app.mount("/media", StaticFiles(directory=str(MEDIA_DIR)), name="media")
45
+
46
+
47
+ class TranscribeResponse(BaseModel):
48
+ text: str
49
+
50
+
51
+ class LocalizeRequest(BaseModel):
52
+ text_fr: str
53
+ lang: Literal["dyu", "mos"]
54
+
55
+
56
+ class LocalizeResponse(BaseModel):
57
+ translated: str
58
+ audio_url: str
59
+
60
+
61
+ class ToFrenchRequest(BaseModel):
62
+ text: str
63
+ lang: Literal["dyu", "mos"]
64
+
65
+
66
+ class ToFrenchResponse(BaseModel):
67
+ text_fr: str
68
+
69
+
70
+ @app.get("/health")
71
+ def health() -> dict[str, str]:
72
+ return {"status": "ok"}
73
+
74
+
75
+ @app.post("/transcribe", response_model=TranscribeResponse)
76
+ async def transcribe(
77
+ file: UploadFile = File(...),
78
+ lang: Literal["dyu", "mos", "fra"] = Form(...),
79
+ ) -> TranscribeResponse:
80
+ start = time.perf_counter()
81
+
82
+ suffix = Path(file.filename or "audio").suffix or ".wav"
83
+ tmp_path = MEDIA_DIR / f"{uuid.uuid4()}{suffix}"
84
+ contents = await file.read()
85
+ tmp_path.write_bytes(contents)
86
+
87
+ try:
88
+ asr = ASR.get_instance()
89
+ text = asr.transcribe(str(tmp_path), lang)
90
+ finally:
91
+ tmp_path.unlink(missing_ok=True)
92
+
93
+ elapsed = time.perf_counter() - start
94
+ logger.info("POST /transcribe lang=%s duration=%.3fs", lang, elapsed)
95
+
96
+ return TranscribeResponse(text=text)
97
+
98
+
99
+ @app.post("/localize", response_model=LocalizeResponse)
100
+ def localize(payload: LocalizeRequest) -> LocalizeResponse:
101
+ start = time.perf_counter()
102
+
103
+ translator = Translator.get_instance()
104
+ translated = translator.translate(payload.text_fr, src="fr", tgt=payload.lang)
105
+
106
+ tts = TTS.get_instance()
107
+ filename = f"{uuid.uuid4()}.wav"
108
+ output_path = MEDIA_DIR / filename
109
+ tts.speak(translated, lang=payload.lang, output_path=str(output_path))
110
+
111
+ elapsed = time.perf_counter() - start
112
+ logger.info("POST /localize lang=%s duration=%.3fs", payload.lang, elapsed)
113
+
114
+ return LocalizeResponse(translated=translated, audio_url=f"/media/{filename}")
115
+
116
+
117
+ @app.post("/to-french", response_model=ToFrenchResponse)
118
+ def to_french(payload: ToFrenchRequest) -> ToFrenchResponse:
119
+ start = time.perf_counter()
120
+
121
+ translator = Translator.get_instance()
122
+ text_fr = translator.translate(payload.text, src=payload.lang, tgt="fr")
123
+
124
+ elapsed = time.perf_counter() - start
125
+ logger.info("POST /to-french lang=%s duration=%.3fs", payload.lang, elapsed)
126
+
127
+ return ToFrenchResponse(text_fr=text_fr)
app/services/__init__.py ADDED
File without changes
app/services/asr.py ADDED
@@ -0,0 +1,137 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Reconnaissance vocale (speech-to-text) pour le Dioula, le Moore et le francais
2
+ via facebook/mms-1b-all.
3
+
4
+ Deux backends, choisis par Settings.ASR_BACKEND :
5
+ - "local" (defaut) : Wav2Vec2ForCTC + AutoProcessor charges en local.
6
+ - "hf_api" : pont temporaire vers l'API d'inference Hugging Face, utile tant
7
+ que le modele local (~3.86 Go) n'est pas entierement telecharge.
8
+ ATTENTION : facebook/mms-1b-all n'est deploye sur AUCUN provider
9
+ d'inference HF (verifie : liste de providers vide). Le backend "hf_api"
10
+ utilise donc openai/whisper-large-v3 a la place, qui NE supporte PAS
11
+ officiellement le Dioula ni le Moore (~99 langues entrainees, dyu/mos
12
+ absentes) : fiable seulement pour lang="fra", best-effort pour dyu/mos.
13
+ Le contrat de transcribe(audio_path, lang) est identique dans les deux cas.
14
+ """
15
+
16
+ import logging
17
+
18
+ import numpy as np
19
+ import torch
20
+ from huggingface_hub import InferenceClient
21
+ from pydub import AudioSegment
22
+ from transformers import AutoProcessor, Wav2Vec2ForCTC
23
+
24
+ from app.deps import get_settings
25
+
26
+ logger = logging.getLogger("model-service.asr")
27
+
28
+ MODEL_NAME = "facebook/mms-1b-all"
29
+ HF_API_MODEL_NAME = "openai/whisper-large-v3"
30
+
31
+ MMS_LANG_CODES = {
32
+ "dyu": "dyu",
33
+ "mos": "mos",
34
+ "fra": "fra",
35
+ }
36
+
37
+ TARGET_SAMPLE_RATE = 16_000
38
+ WINDOW_SECONDS = 30
39
+ OVERLAP_SECONDS = 2
40
+
41
+
42
+ class ASR:
43
+ _instance = None
44
+
45
+ def __init__(self) -> None:
46
+ settings = get_settings()
47
+ self.backend = settings.ASR_BACKEND
48
+
49
+ if self.backend == "hf_api":
50
+ self._client = InferenceClient(model=HF_API_MODEL_NAME, token=settings.HF_TOKEN)
51
+ return
52
+
53
+ self.device = "cuda" if torch.cuda.is_available() else "cpu"
54
+ self.processor = AutoProcessor.from_pretrained(MODEL_NAME)
55
+ self.model = Wav2Vec2ForCTC.from_pretrained(MODEL_NAME).to(self.device)
56
+ self.model.eval()
57
+ self._current_lang: str | None = None
58
+
59
+ @classmethod
60
+ def get_instance(cls) -> "ASR":
61
+ if cls._instance is None:
62
+ cls._instance = cls()
63
+ return cls._instance
64
+
65
+ def _set_lang(self, lang: str) -> None:
66
+ if lang not in MMS_LANG_CODES:
67
+ raise ValueError(f"Langue non supportee: {lang}")
68
+ target_lang = MMS_LANG_CODES[lang]
69
+ if self._current_lang == target_lang:
70
+ return
71
+ self.processor.tokenizer.set_target_lang(target_lang)
72
+ self.model.load_adapter(target_lang)
73
+ self._current_lang = target_lang
74
+
75
+ def _load_audio(self, audio_path: str) -> np.ndarray:
76
+ audio = AudioSegment.from_file(audio_path)
77
+ audio = audio.set_channels(1).set_frame_rate(TARGET_SAMPLE_RATE)
78
+ samples = np.array(audio.get_array_of_samples()).astype(np.float32)
79
+ max_val = float(1 << (8 * audio.sample_width - 1))
80
+ samples /= max_val
81
+ return samples
82
+
83
+ def _transcribe_chunk(self, chunk: np.ndarray) -> str:
84
+ inputs = self.processor(
85
+ chunk, sampling_rate=TARGET_SAMPLE_RATE, return_tensors="pt"
86
+ ).to(self.device)
87
+ with torch.no_grad():
88
+ logits = self.model(**inputs).logits
89
+ ids = torch.argmax(logits, dim=-1)
90
+ return self.processor.batch_decode(ids)[0]
91
+
92
+ def _transcribe_hf_api(self, audio_path: str, lang: str) -> str:
93
+ if lang not in MMS_LANG_CODES:
94
+ raise ValueError(f"Langue non supportee: {lang}")
95
+ if lang != "fra":
96
+ logger.warning(
97
+ "ASR hf_api (whisper-large-v3) ne supporte pas officiellement '%s' "
98
+ "(dyu/mos absents des langues entrainees) : resultat best-effort.",
99
+ lang,
100
+ )
101
+ output = self._client.automatic_speech_recognition(audio_path)
102
+ return output.text.strip()
103
+
104
+ def transcribe(self, audio_path: str, lang: str) -> str:
105
+ if self.backend == "hf_api":
106
+ return self._transcribe_hf_api(audio_path, lang)
107
+
108
+ self._set_lang(lang)
109
+ samples = self._load_audio(audio_path)
110
+
111
+ window_size = WINDOW_SECONDS * TARGET_SAMPLE_RATE
112
+ overlap_size = OVERLAP_SECONDS * TARGET_SAMPLE_RATE
113
+ step = window_size - overlap_size
114
+
115
+ if len(samples) <= window_size:
116
+ return self._transcribe_chunk(samples).strip()
117
+
118
+ transcripts = []
119
+ start = 0
120
+ while start < len(samples):
121
+ chunk = samples[start : start + window_size]
122
+ if len(chunk) == 0:
123
+ break
124
+ transcripts.append(self._transcribe_chunk(chunk))
125
+ start += step
126
+
127
+ return " ".join(t.strip() for t in transcripts if t.strip())
128
+
129
+
130
+ if __name__ == "__main__":
131
+ import sys
132
+
133
+ asr = ASR.get_instance()
134
+ path = sys.argv[1] if len(sys.argv) > 1 else "sample.wav"
135
+ lang_arg = sys.argv[2] if len(sys.argv) > 2 else "dyu"
136
+ text = asr.transcribe(path, lang_arg)
137
+ print(f"Transcription ({lang_arg}): {text}")
app/services/translation.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Traduction entre le francais et le Dioula / Moore."""
app/services/translator.py ADDED
@@ -0,0 +1,64 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Traduction entre le francais et le Dioula / Moore via facebook/nllb-200-distilled-600M."""
2
+
3
+ import re
4
+
5
+ import torch
6
+ from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
7
+
8
+ MODEL_NAME = "facebook/nllb-200-distilled-600M"
9
+
10
+ NLLB_LANG_CODES = {
11
+ "fr": "fra_Latn",
12
+ "dyu": "dyu_Latn",
13
+ "mos": "mos_Latn",
14
+ }
15
+
16
+ _SENTENCE_SPLIT_RE = re.compile(r"(?<=[.!?])\s+")
17
+
18
+
19
+ class Translator:
20
+ _instance = None
21
+
22
+ def __init__(self) -> None:
23
+ self.device = "cuda" if torch.cuda.is_available() else "cpu"
24
+ self.tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
25
+ self.model = AutoModelForSeq2SeqLM.from_pretrained(MODEL_NAME).to(self.device)
26
+ self.model.eval()
27
+
28
+ @classmethod
29
+ def get_instance(cls) -> "Translator":
30
+ if cls._instance is None:
31
+ cls._instance = cls()
32
+ return cls._instance
33
+
34
+ def _split_sentences(self, text: str) -> list[str]:
35
+ sentences = [s.strip() for s in _SENTENCE_SPLIT_RE.split(text.strip()) if s.strip()]
36
+ return sentences or [text.strip()]
37
+
38
+ def _translate_sentence(self, sentence: str, src: str, tgt: str) -> str:
39
+ self.tokenizer.src_lang = NLLB_LANG_CODES[src]
40
+ inputs = self.tokenizer(sentence, return_tensors="pt").to(self.device)
41
+ forced_bos_token_id = self.tokenizer.convert_tokens_to_ids(NLLB_LANG_CODES[tgt])
42
+ with torch.no_grad():
43
+ generated = self.model.generate(
44
+ **inputs,
45
+ forced_bos_token_id=forced_bos_token_id,
46
+ num_beams=4,
47
+ max_length=256,
48
+ )
49
+ return self.tokenizer.batch_decode(generated, skip_special_tokens=True)[0]
50
+
51
+ def translate(self, text: str, src: str, tgt: str) -> str:
52
+ if src not in NLLB_LANG_CODES or tgt not in NLLB_LANG_CODES:
53
+ raise ValueError(f"Langue non supportee: src={src}, tgt={tgt}")
54
+ sentences = self._split_sentences(text)
55
+ translated = [self._translate_sentence(s, src, tgt) for s in sentences]
56
+ return " ".join(translated)
57
+
58
+
59
+ if __name__ == "__main__":
60
+ translator = Translator.get_instance()
61
+ example = "Bonjour. Comment allez-vous aujourd'hui ? J'espere que tout va bien."
62
+ result = translator.translate(example, src="fr", tgt="dyu")
63
+ print(f"FR: {example}")
64
+ print(f"DYU: {result}")
app/services/tts.py ADDED
@@ -0,0 +1,102 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Synthese vocale (text-to-speech) pour le Dioula et le Moore via les modeles
2
+ VITS facebook/mms-tts-dyu et facebook/mms-tts-mos."""
3
+
4
+ import re
5
+
6
+ import numpy as np
7
+ import soundfile as sf
8
+ import torch
9
+ from transformers import VitsModel, VitsTokenizer
10
+
11
+ MMS_TTS_MODEL_NAMES = {
12
+ "dyu": "facebook/mms-tts-dyu",
13
+ "mos": "facebook/mms-tts-mos",
14
+ }
15
+
16
+ MAX_CHARS_BEFORE_SPLIT = 500
17
+ SILENCE_SECONDS = 0.3
18
+ MIN_SEGMENT_LETTERS = 4
19
+
20
+ _SENTENCE_SPLIT_RE = re.compile(r"(?<=[.!?])\s+")
21
+ _LETTERS_RE = re.compile(r"[^a-zA-ZÀ-ÖØ-öø-ÿ]")
22
+
23
+
24
+ class TTS:
25
+ _instance = None
26
+
27
+ def __init__(self) -> None:
28
+ self.device = "cuda" if torch.cuda.is_available() else "cpu"
29
+ self._models: dict[str, VitsModel] = {}
30
+ self._tokenizers: dict[str, VitsTokenizer] = {}
31
+
32
+ @classmethod
33
+ def get_instance(cls) -> "TTS":
34
+ if cls._instance is None:
35
+ cls._instance = cls()
36
+ return cls._instance
37
+
38
+ def _get_model(self, lang: str) -> tuple[VitsModel, VitsTokenizer]:
39
+ if lang not in MMS_TTS_MODEL_NAMES:
40
+ raise ValueError(f"Langue non supportee: {lang}")
41
+ if lang not in self._models:
42
+ model_name = MMS_TTS_MODEL_NAMES[lang]
43
+ self._tokenizers[lang] = VitsTokenizer.from_pretrained(model_name)
44
+ model = VitsModel.from_pretrained(model_name).to(self.device)
45
+ model.eval()
46
+ self._models[lang] = model
47
+ return self._models[lang], self._tokenizers[lang]
48
+
49
+ def _split_text(self, text: str) -> list[str]:
50
+ text = text.strip()
51
+ if len(text) <= MAX_CHARS_BEFORE_SPLIT:
52
+ return [text]
53
+ raw_segments = [s.strip() for s in _SENTENCE_SPLIT_RE.split(text) if s.strip()]
54
+ return self._merge_short_segments(raw_segments)
55
+
56
+ def _merge_short_segments(self, segments: list[str]) -> list[str]:
57
+ """Fusionne les fragments trop courts (ex. '1.' d'une liste numerotee)
58
+ avec le fragment suivant : VITS plante (narrow(): length must be
59
+ non-negative) sur une sequence phonemisee trop courte."""
60
+ merged: list[str] = []
61
+ buffer = ""
62
+ for seg in segments:
63
+ buffer = f"{buffer} {seg}".strip() if buffer else seg
64
+ if len(_LETTERS_RE.sub("", buffer)) >= MIN_SEGMENT_LETTERS:
65
+ merged.append(buffer)
66
+ buffer = ""
67
+ if buffer:
68
+ if merged:
69
+ merged[-1] = f"{merged[-1]} {buffer}".strip()
70
+ else:
71
+ merged.append(buffer)
72
+ return merged
73
+
74
+ def _synthesize_segment(self, text: str, lang: str) -> np.ndarray:
75
+ model, tokenizer = self._get_model(lang)
76
+ inputs = tokenizer(text, return_tensors="pt").to(self.device)
77
+ with torch.no_grad():
78
+ output = model(**inputs).waveform
79
+ return output.squeeze().cpu().numpy()
80
+
81
+ def speak(self, text: str, lang: str, output_path: str) -> str:
82
+ model, _ = self._get_model(lang)
83
+ sample_rate = model.config.sampling_rate
84
+
85
+ segments = self._split_text(text)
86
+ silence = np.zeros(int(SILENCE_SECONDS * sample_rate), dtype=np.float32)
87
+
88
+ waveforms = []
89
+ for i, segment in enumerate(segments):
90
+ waveforms.append(self._synthesize_segment(segment, lang))
91
+ if i < len(segments) - 1:
92
+ waveforms.append(silence)
93
+
94
+ audio = np.concatenate(waveforms) if len(waveforms) > 1 else waveforms[0]
95
+ sf.write(output_path, audio, sample_rate)
96
+ return output_path
97
+
98
+
99
+ if __name__ == "__main__":
100
+ tts = TTS.get_instance()
101
+ out = tts.speak("I ni ce. An be here?", lang="dyu", output_path="demo_dyu.wav")
102
+ print(f"Audio genere: {out}")
pytest.ini ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ [pytest]
2
+ markers =
3
+ slow: tests lents qui chargent de vrais modeles (exclus par defaut avec -m "not slow")
requirements.txt ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ fastapi
2
+ uvicorn[standard]
3
+ python-multipart
4
+ transformers
5
+ torch
6
+ accelerate
7
+ sentencepiece
8
+ pydantic-settings
9
+ python-dotenv
10
+ soundfile
11
+ scipy
12
+ pydub
13
+ audioop-lts; python_version >= "3.13"
14
+ pillow
15
+ requests
16
+ pytest
17
+ httpx
tests/__init__.py ADDED
File without changes
tests/conftest.py ADDED
@@ -0,0 +1,51 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import pytest
3
+ import soundfile as sf
4
+ from fastapi.testclient import TestClient
5
+
6
+ from app.main import app
7
+ from app.services.asr import ASR
8
+ from app.services.translator import Translator
9
+ from app.services.tts import TTS
10
+
11
+ FIXED_TRANSCRIPT = "ceci est un texte fixe"
12
+ FIXED_TRANSLATION = "traduction fixe"
13
+
14
+
15
+ class FakeASR:
16
+ def transcribe(self, audio_path: str, lang: str) -> str:
17
+ return FIXED_TRANSCRIPT
18
+
19
+
20
+ class FakeTranslator:
21
+ def translate(self, text: str, src: str, tgt: str) -> str:
22
+ return FIXED_TRANSLATION
23
+
24
+
25
+ class FakeTTS:
26
+ def speak(self, text: str, lang: str, output_path: str) -> str:
27
+ silence = np.zeros(1, dtype=np.float32)
28
+ sf.write(output_path, silence, 16_000)
29
+ return output_path
30
+
31
+
32
+ @pytest.fixture(autouse=True)
33
+ def mock_heavy_services(request, monkeypatch):
34
+ """Evite le chargement de vrais modeles pendant les tests rapides.
35
+
36
+ Les tests marques @pytest.mark.slow veulent les vrais services : on ne
37
+ patche rien pour eux.
38
+ """
39
+ if request.node.get_closest_marker("slow"):
40
+ yield
41
+ return
42
+
43
+ monkeypatch.setattr(ASR, "get_instance", classmethod(lambda cls: FakeASR()))
44
+ monkeypatch.setattr(Translator, "get_instance", classmethod(lambda cls: FakeTranslator()))
45
+ monkeypatch.setattr(TTS, "get_instance", classmethod(lambda cls: FakeTTS()))
46
+ yield
47
+
48
+
49
+ @pytest.fixture
50
+ def client():
51
+ return TestClient(app)
tests/test_api.py ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ def test_health(client):
2
+ response = client.get("/health")
3
+ assert response.status_code == 200
4
+ assert response.json() == {"status": "ok"}
5
+
6
+
7
+ def test_localize_mocked(client):
8
+ response = client.post(
9
+ "/localize",
10
+ json={"text_fr": "Bonjour tout le monde", "lang": "dyu"},
11
+ )
12
+ assert response.status_code == 200
13
+ data = response.json()
14
+ assert "translated" in data
15
+ assert "audio_url" in data
16
+
17
+
18
+ def test_localize_missing_lang(client):
19
+ response = client.post(
20
+ "/localize",
21
+ json={"text_fr": "Bonjour tout le monde"},
22
+ )
23
+ assert response.status_code == 422
24
+
25
+
26
+ def test_to_french_mocked(client):
27
+ response = client.post(
28
+ "/to-french",
29
+ json={"text": "i ni ce", "lang": "dyu"},
30
+ )
31
+ assert response.status_code == 200
32
+ assert "text_fr" in response.json()
33
+
34
+
35
+ def test_to_french_missing_lang(client):
36
+ response = client.post("/to-french", json={"text": "i ni ce"})
37
+ assert response.status_code == 422
38
+
39
+
40
+ def test_transcribe_missing_file(client):
41
+ response = client.post(
42
+ "/transcribe",
43
+ data={"lang": "dyu"},
44
+ )
45
+ assert response.status_code == 422
tests/test_services.py ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import pytest
2
+
3
+ from app.services.translator import Translator
4
+ from app.services.tts import TTS
5
+
6
+
7
+ @pytest.mark.slow
8
+ def test_translate_fr_to_dyu_real():
9
+ translator = Translator.get_instance()
10
+ result = translator.translate("Bonjour, comment allez-vous ?", src="fr", tgt="dyu")
11
+ assert isinstance(result, str)
12
+ assert result.strip() != ""
13
+
14
+
15
+ def test_tts_merges_bare_numbered_list_markers():
16
+ """Regression : une liste numerotee ('1. Xxx. 2. Yyy.') donne des segments
17
+ '1.'/'2.' de 2 caracteres apres decoupe par phrase ; VITS plante
18
+ (narrow(): length must be non-negative) si on les synthetise seuls."""
19
+ # Instanciation directe (pas get_instance()) : construire un TTS ne charge
20
+ # aucun modele (lazy par langue), donc pas besoin du mock get_instance()
21
+ # utilise pour les autres tests rapides.
22
+ tts = TTS()
23
+ text = (
24
+ "1. " + ("Allez a la mairie avec vos papiers d identite. " * 6)
25
+ + "2. " + ("Presentez votre demande au guichet. " * 6)
26
+ )
27
+ assert len(text) > 500 # declenche la decoupe en segments
28
+
29
+ segments = tts._split_text(text)
30
+
31
+ assert all(len(s) > 5 for s in segments), segments
32
+
33
+
34
+ @pytest.mark.slow
35
+ def test_tts_speak_numbered_list_real():
36
+ """Reproduit le crash original avec une vraie synthese VITS."""
37
+ tts = TTS.get_instance()
38
+ text = (
39
+ "1. Allez a la mairie avec vos papiers d identite et un justificatif de "
40
+ "domicile recent, puis attendez votre tour dans la file d attente prevue "
41
+ "a cet effet pour les demandes administratives courantes. "
42
+ "2. Presentez votre demande au guichet approprie et attendez votre tour "
43
+ "patiemment en respectant les horaires d ouverture affiches devant le "
44
+ "batiment principal de la mairie. "
45
+ "3. Payez les frais administratifs requis pour le traitement de votre "
46
+ "dossier officiel aupres du caissier designe et conservez precieusement "
47
+ "votre recu de paiement."
48
+ )
49
+ out = tts.speak(text, lang="dyu", output_path="test_numbered_list.wav")
50
+ assert out == "test_numbered_list.wav"