Markuspierre commited on
Commit
0c663e1
·
verified ·
1 Parent(s): fbe49da

Update asr-tts_service.py

Browse files
Files changed (1) hide show
  1. asr-tts_service.py +16 -24
asr-tts_service.py CHANGED
@@ -6,7 +6,7 @@ import logging
6
  import numpy as np
7
  import soundfile as sf
8
  import torch
9
- import google.generativeai as genai
10
  import requests
11
  import tempfile
12
  from flask import Flask, request, jsonify
@@ -14,17 +14,16 @@ from transformers import pipeline, AutoTokenizer
14
  from parler_tts import ParlerTTSForConditionalGeneration
15
  from pydub import AudioSegment
16
  from dotenv import load_dotenv
17
- from concurrent.futures import ThreadPoolExecutor
18
 
19
 
20
 
21
  # Charger les variables d'environnement
22
  load_dotenv()
23
 
24
- # Configuration Gemini
25
  google_api = os.getenv('GOOGLE_API_KEY')
26
- genai.configure(api_key=google_api)
27
- model_gemini = genai.GenerativeModel('gemini-2.5-flash') #gemini-3-flash-preview #gemini-flash-latest
28
 
29
  # Configuration générale
30
  device = "cpu"
@@ -77,19 +76,16 @@ def number_to_french(n: int) -> str:
77
  if n == 80: return base
78
  return base + "-" + number_to_french(n - 80)
79
 
80
- # GESTION DES CENTAINES
81
  if n < 1000:
82
  hundreds, rest = divmod(n, 100)
83
  base = "cent" if hundreds == 1 else UNITS.get(hundreds, str(hundreds)) + " cent"
84
  return base if rest == 0 else base + " " + number_to_french(rest)
85
 
86
- # GESTION DES MILLIERS ET PLUS (Sécurité pour éviter le KeyError)
87
  if n < 1000000:
88
  thousands, rest = divmod(n, 1000)
89
  base = "mille" if thousands == 1 else number_to_french(thousands) + " mille"
90
  return base if rest == 0 else base + " " + number_to_french(rest)
91
 
92
- # Par sécurité, si le nombre est trop grand (millions), on le renvoie tel quel
93
  return str(n)
94
 
95
  def convert_digits_in_text(text: str) -> str:
@@ -97,11 +93,8 @@ def convert_digits_in_text(text: str) -> str:
97
 
98
  def repl(match):
99
  s = match.group(0)
100
- # On épèle : "7 7 1 2 3..." au lieu de "Sept milliards..."
101
  if len(s) > 4:
102
  return " ".join([UNITS.get(int(digit), digit) for digit in s])
103
-
104
- # SINON (Petit nombre comme 15, 100, 2024)
105
  try:
106
  val = int(s)
107
  return number_to_french(val)
@@ -162,7 +155,7 @@ def generate_tts_optimized(text: str) -> str:
162
  prompt_input_ids=prompt_ids,
163
  max_new_tokens=2048,
164
  do_sample=True,
165
- temperature= 0.8,
166
  min_new_tokens=20
167
  )
168
  audio_np = audio.cpu().numpy().squeeze().astype(np.float32)
@@ -189,7 +182,10 @@ def french_to_wolof_with_gemini(text: str) -> str:
189
  Texte : {text}"""
190
 
191
  try:
192
- response = model_gemini.generate_content(prompt)
 
 
 
193
  return response.text.strip()
194
  except Exception as e:
195
  return f"Erreur de traduction : {str(e)}"
@@ -205,7 +201,10 @@ def wolof_to_french_gemini(text: str) -> str:
205
 
206
  Texte : {text}"""
207
  try:
208
- response = model_gemini.generate_content(prompt)
 
 
 
209
  return response.text.strip()
210
  except Exception as e:
211
  return 'Bonjour'
@@ -237,36 +236,30 @@ def transcribe_from_url():
237
  if not audio_url: return "Bonjour", 400
238
 
239
  try:
240
- # Téléchargement du fichier
241
  resp = requests.get(audio_url, stream=True)
242
  if resp.status_code != 200:
243
  logger.error(f"Erreur téléchargement audio: {resp.status_code}")
244
  return "Bonjour"
245
 
246
- # Utilisation d'un fichier temporaire sécurisé
247
  with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as tmp:
248
  tmp.write(resp.content)
249
- tmp.flush() # FORCE l'écriture des données sur le disque
250
- os.fsync(tmp.fileno()) # Assure la synchronisation physique
251
  tmp_path = tmp.name
252
 
253
  try:
254
- # Conversion via pydub (ffmpeg)
255
  audio = AudioSegment.from_file(tmp_path)
256
  wav_io = io.BytesIO()
257
  audio.set_frame_rate(16000).set_channels(1).export(wav_io, format="wav")
258
  wav_io.seek(0)
259
 
260
- # Lecture des données audio
261
  data, sr = sf.read(wav_io)
262
  data = np.asarray(data, dtype=np.float32)
263
  if data.ndim > 1: data = data.mean(axis=1)
264
 
265
- # Traitement ASR
266
  wolof_text = asr(normalize_audio(data))["text"]
267
 
268
  finally:
269
- # Nettoyage du fichier temporaire même en cas d'erreur de décodage
270
  if os.path.exists(tmp_path):
271
  os.remove(tmp_path)
272
 
@@ -277,7 +270,6 @@ def transcribe_from_url():
277
 
278
  except Exception as e:
279
  logger.error(f"Erreur WhatsApp ASR: {e}")
280
- # Log supplémentaire pour débugger ffmpeg si l'erreur persiste
281
  return "Bonjour"
282
 
283
  @app.route("/tts", methods=["POST"])
@@ -291,4 +283,4 @@ def tts():
291
  return jsonify({"wolof_text": wolof_text, "audio": audio_base64})
292
 
293
  if __name__ == "__main__":
294
- app.run(host="0.0.0.0", port=7860)
 
6
  import numpy as np
7
  import soundfile as sf
8
  import torch
9
+ from google import genai as google_genai
10
  import requests
11
  import tempfile
12
  from flask import Flask, request, jsonify
 
14
  from parler_tts import ParlerTTSForConditionalGeneration
15
  from pydub import AudioSegment
16
  from dotenv import load_dotenv
 
17
 
18
 
19
 
20
  # Charger les variables d'environnement
21
  load_dotenv()
22
 
23
+ # Configuration Gemini (nouveau SDK google-genai)
24
  google_api = os.getenv('GOOGLE_API_KEY')
25
+ gemini_client = google_genai.Client(api_key=google_api)
26
+ GEMINI_MODEL = "gemini-2.5-flash"
27
 
28
  # Configuration générale
29
  device = "cpu"
 
76
  if n == 80: return base
77
  return base + "-" + number_to_french(n - 80)
78
 
 
79
  if n < 1000:
80
  hundreds, rest = divmod(n, 100)
81
  base = "cent" if hundreds == 1 else UNITS.get(hundreds, str(hundreds)) + " cent"
82
  return base if rest == 0 else base + " " + number_to_french(rest)
83
 
 
84
  if n < 1000000:
85
  thousands, rest = divmod(n, 1000)
86
  base = "mille" if thousands == 1 else number_to_french(thousands) + " mille"
87
  return base if rest == 0 else base + " " + number_to_french(rest)
88
 
 
89
  return str(n)
90
 
91
  def convert_digits_in_text(text: str) -> str:
 
93
 
94
  def repl(match):
95
  s = match.group(0)
 
96
  if len(s) > 4:
97
  return " ".join([UNITS.get(int(digit), digit) for digit in s])
 
 
98
  try:
99
  val = int(s)
100
  return number_to_french(val)
 
155
  prompt_input_ids=prompt_ids,
156
  max_new_tokens=2048,
157
  do_sample=True,
158
+ temperature=0.8,
159
  min_new_tokens=20
160
  )
161
  audio_np = audio.cpu().numpy().squeeze().astype(np.float32)
 
182
  Texte : {text}"""
183
 
184
  try:
185
+ response = gemini_client.models.generate_content(
186
+ model=GEMINI_MODEL,
187
+ contents=prompt
188
+ )
189
  return response.text.strip()
190
  except Exception as e:
191
  return f"Erreur de traduction : {str(e)}"
 
201
 
202
  Texte : {text}"""
203
  try:
204
+ response = gemini_client.models.generate_content(
205
+ model=GEMINI_MODEL,
206
+ contents=prompt
207
+ )
208
  return response.text.strip()
209
  except Exception as e:
210
  return 'Bonjour'
 
236
  if not audio_url: return "Bonjour", 400
237
 
238
  try:
 
239
  resp = requests.get(audio_url, stream=True)
240
  if resp.status_code != 200:
241
  logger.error(f"Erreur téléchargement audio: {resp.status_code}")
242
  return "Bonjour"
243
 
 
244
  with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as tmp:
245
  tmp.write(resp.content)
246
+ tmp.flush()
247
+ os.fsync(tmp.fileno())
248
  tmp_path = tmp.name
249
 
250
  try:
 
251
  audio = AudioSegment.from_file(tmp_path)
252
  wav_io = io.BytesIO()
253
  audio.set_frame_rate(16000).set_channels(1).export(wav_io, format="wav")
254
  wav_io.seek(0)
255
 
 
256
  data, sr = sf.read(wav_io)
257
  data = np.asarray(data, dtype=np.float32)
258
  if data.ndim > 1: data = data.mean(axis=1)
259
 
 
260
  wolof_text = asr(normalize_audio(data))["text"]
261
 
262
  finally:
 
263
  if os.path.exists(tmp_path):
264
  os.remove(tmp_path)
265
 
 
270
 
271
  except Exception as e:
272
  logger.error(f"Erreur WhatsApp ASR: {e}")
 
273
  return "Bonjour"
274
 
275
  @app.route("/tts", methods=["POST"])
 
283
  return jsonify({"wolof_text": wolof_text, "audio": audio_base64})
284
 
285
  if __name__ == "__main__":
286
+ app.run(host="0.0.0.0", port=7860)