carlosduplar commited on
Commit
4ff45b3
·
1 Parent(s): b05542b

build-small-hackathon: switch to modal.asgi_app, fix endpoint body parsing, base64 audio

Browse files
Files changed (5) hide show
  1. README.md +15 -16
  2. llm_engine.py +2 -6
  3. modal_app.py +66 -134
  4. stt_engine.py +4 -5
  5. tts_engine.py +6 -23
README.md CHANGED
@@ -1,20 +1,19 @@
1
- <div align="center">
2
- <img width="1200" height="475" alt="GHBanner" src="https://github.com/user-attachments/assets/0aa67016-6eaf-458a-adb2-6e31a0763ed6" />
3
- </div>
 
 
 
 
 
 
 
 
4
 
5
- # Run and deploy your AI Studio app
6
 
7
- This contains everything you need to run your app locally.
8
 
9
- View your app in AI Studio: https://ai.studio/apps/5628e25f-e268-4a5d-9d6d-b6d96c02e6f7
10
 
11
- ## Run Locally
12
-
13
- **Prerequisites:** Node.js
14
-
15
-
16
- 1. Install dependencies:
17
- `npm install`
18
- 2. Set the `GEMINI_API_KEY` in [.env.local](.env.local) to your Gemini API key
19
- 3. Run the app:
20
- `npm run dev`
 
1
+ ---
2
+ title: Patient Virtuel · Hygiéniste Pro
3
+ emoji: 🦷
4
+ colorFrom: orange
5
+ colorTo: dark
6
+ sdk: gradio
7
+ sdk_version: 5.0
8
+ app_file: app.py
9
+ pinned: false
10
+ license: apache-2.0
11
+ ---
12
 
13
+ # Patient Virtuel · Hygiéniste Pro
14
 
15
+ Hackathon submission for [build-small-hackathon](https://huggingface.co/build-small-hackathon).
16
 
17
+ **Track**: Backyard AI — a real tool for a real learner: a dental hygienist training professional French for a Swiss clinic.
18
 
19
+ See [`space_README.md`](space_README.md) for the full description.
 
 
 
 
 
 
 
 
 
llm_engine.py CHANGED
@@ -4,34 +4,30 @@ import httpx
4
  MODAL_ENDPOINT = os.environ.get("MODAL_ENDPOINT_QWEN", "")
5
  MODAL_AUTH_TOKEN = os.environ.get("MODAL_AUTH_TOKEN", "")
6
 
7
- TRUNCATION_LIMIT = 20_000 # tokens; oldest non-system turns trimmed when exceeded
8
 
9
 
10
  def chat(messages: list[dict]) -> str | None:
11
  if not MODAL_ENDPOINT:
12
  raise RuntimeError("MODAL_ENDPOINT_QWEN not set")
13
 
14
- # Trim history if too long: keep system prompt + last N turns
15
  _trim(messages)
16
 
17
  resp = httpx.post(
18
  MODAL_ENDPOINT,
19
  json={"messages": messages, "token": MODAL_AUTH_TOKEN},
20
- timeout=600, # long timeout for cold starts
21
  )
22
  resp.raise_for_status()
23
  return resp.json().get("text")
24
 
25
 
26
  def _trim(messages: list[dict]):
27
- """Drop oldest non-system turns if total tokens exceeds TRUNCATION_LIMIT."""
28
  if len(messages) < 4:
29
  return
30
- # rough estimate: 1 token ≈ 3.5 chars
31
  total_chars = sum(len(m.get("content", "")) for m in messages)
32
  if total_chars < TRUNCATION_LIMIT * 3.5:
33
  return
34
- # keep system prompt, drop oldest user/assistant pairs
35
  system = [m for m in messages if m.get("role") == "system"]
36
  rest = [m for m in messages if m.get("role") != "system"]
37
  while rest and total_chars >= TRUNCATION_LIMIT * 3.5:
 
4
  MODAL_ENDPOINT = os.environ.get("MODAL_ENDPOINT_QWEN", "")
5
  MODAL_AUTH_TOKEN = os.environ.get("MODAL_AUTH_TOKEN", "")
6
 
7
+ TRUNCATION_LIMIT = 20_000
8
 
9
 
10
  def chat(messages: list[dict]) -> str | None:
11
  if not MODAL_ENDPOINT:
12
  raise RuntimeError("MODAL_ENDPOINT_QWEN not set")
13
 
 
14
  _trim(messages)
15
 
16
  resp = httpx.post(
17
  MODAL_ENDPOINT,
18
  json={"messages": messages, "token": MODAL_AUTH_TOKEN},
19
+ timeout=600,
20
  )
21
  resp.raise_for_status()
22
  return resp.json().get("text")
23
 
24
 
25
  def _trim(messages: list[dict]):
 
26
  if len(messages) < 4:
27
  return
 
28
  total_chars = sum(len(m.get("content", "")) for m in messages)
29
  if total_chars < TRUNCATION_LIMIT * 3.5:
30
  return
 
31
  system = [m for m in messages if m.get("role") == "system"]
32
  rest = [m for m in messages if m.get("role") != "system"]
33
  while rest and total_chars >= TRUNCATION_LIMIT * 3.5:
modal_app.py CHANGED
@@ -1,181 +1,113 @@
 
1
  import os
2
  import tempfile
3
- import soundfile as sf
4
  import modal
5
 
6
  app = modal.App("patient-virtuel")
7
 
8
- # ---- Shared volumes for model cache ----
9
  qwen_vol = modal.Volume.from_name("qwen-cache", create_if_missing=True)
10
  whisper_vol = modal.Volume.from_name("whisper-cache", create_if_missing=True)
11
  CACHE_DIR = "/root/.cache/huggingface"
12
 
13
- # ---- Auth helper ----
14
- EXPECTED_TOKEN = os.environ.get("EXPECTED_TOKEN", "")
15
-
16
- def _check_token(token: str | None):
17
- if not token or token != EXPECTED_TOKEN:
18
- raise modal.exception.APIError("Unauthorized")
19
-
20
- default_qwen_timeout = 10 * 60 # 10m for cold start + inference
21
-
22
  # ---- 1. Qwen/Qwen3.6-27B LLM ----
23
  qwen_image = (
24
  modal.Image.debian_slim(python_version="3.12")
25
- .pip_install(
26
- "vllm>=0.8.5",
27
- "huggingface-hub>=0.25",
28
- "torch>=2.5",
29
- )
30
  )
31
 
32
- HF_QWEN = "Qwen/Qwen3.6-27B"
33
- GA = 15
34
- LP = 15
35
 
36
  @app.function(
37
  image=qwen_image,
38
- gpu=modal.gpu.A100(count=1, memory=40),
39
  volumes={CACHE_DIR: qwen_vol},
40
- secrets=[modal.Secret.from_name("hf-token")],
41
- timeout=default_qwen_timeout,
42
- scaledown_window=600, # keep warm 10 min after last call
43
- concurrency_limit=1,
44
- container_idle_timeout=600,
45
  )
46
- @modal.web_endpoint(method="POST", label="qwen-chat")
47
- def qwen_chat(messages: list[dict], token: str | None = None):
48
- _check_token(token)
 
 
49
 
50
- from vllm import LLM, SamplingParams, TokensPrompt
 
 
 
 
 
 
51
 
52
- llm = LLM(
53
- model=HF_QWEN,
54
- dtype="bfloat16",
55
- max_model_len=32768,
56
- trust_remote_code=True,
57
- )
58
 
59
- sp = SamplingParams(
60
- temperature=0.7,
61
- top_p=0.8,
62
- top_k=20,
63
- max_tokens=512,
64
- stop=["</s>", "<|im_end|>"],
65
- )
66
 
67
- outputs = llm.chat(
68
- messages=messages,
69
- sampling_params=sp,
70
- use_tqdm=False,
71
- )
 
 
 
72
 
73
- text = outputs[0].outputs[0].text.strip()
 
74
 
75
- # Strip reasoning tokens if present
76
- if "<think>" in text:
77
- text = text.split("</think>")[-1].strip()
 
78
 
79
- return {"text": text}
80
 
81
 
82
- # ---- 2. Whisper STT ----
83
  whisper_image = (
84
  modal.Image.debian_slim(python_version="3.12")
85
- .pip_install("faster-whisper", "numpy", "soundfile")
86
  )
87
 
 
88
  @app.function(
89
  image=whisper_image,
90
- gpu=modal.gpu.A10G(count=1),
91
  volumes={CACHE_DIR: whisper_vol},
92
  timeout=120,
93
  scaledown_window=300,
94
- concurrency_limit=1,
95
- container_idle_timeout=300,
96
  )
97
- @modal.web_endpoint(method="POST", label="whisper-stt")
98
- def whisper_stt(token: str | None = None, audio_bytes: bytes | None = None):
99
- _check_token(token)
100
 
101
- if not audio_bytes:
102
- return {"error": "No audio bytes provided"}, 400
103
 
104
- from faster_whisper import WhisperModel
 
 
 
 
 
 
105
 
106
- model = WhisperModel(
107
- "large-v3-turbo",
108
- device="cuda",
109
- compute_type="float16",
110
- download_root=CACHE_DIR,
111
- )
112
 
113
- with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
114
- f.write(audio_bytes if isinstance(audio_bytes, bytes) else audio_bytes.encode("latin1"))
115
- path = f.name
116
 
117
- segments, info = model.transcribe(
118
- path,
119
- language="fr",
120
- beam_size=5,
121
- vad_filter=True,
122
- condition_on_previous_text=False,
123
- )
124
- os.unlink(path)
125
 
126
- text = " ".join(seg.text for seg in segments)
127
- return {"text": text.strip()}
 
 
128
 
 
 
129
 
130
- # ---- 3. Mistral API TTS (wrapped for unified call from Space) ----
131
- @app.function(
132
- image=modal.Image.debian_slim(python_version="3.12").pip_install("httpx"),
133
- secrets=[modal.Secret.from_name("mistral-key")],
134
- timeout=30,
135
- scaledown_window=60,
136
- )
137
- @modal.web_endpoint(method="POST", label="mistral-tts")
138
- def mistral_tts(text: str, token: str | None = None):
139
- _check_token(token)
140
-
141
- import httpx as _httpx
142
-
143
- api_key = os.environ.get("MISTRAL_API_KEY", "")
144
- if not api_key:
145
- return {"error": "MISTRAL_API_KEY not set"}, 500
146
-
147
- resp = _httpx.post(
148
- "https://api.mistral.ai/v1/audio/speech",
149
- headers={"Authorization": f"Bearer {api_key}"},
150
- json={
151
- "model": "mistral-tts", # TODO: confirm exact model name on Mistral API
152
- "input": text,
153
- "voice": "french_female",
154
- "response_format": "wav",
155
- },
156
- timeout=25,
157
- )
158
- resp.raise_for_status()
159
- return resp.content # raw WAV bytes
160
-
161
-
162
- # ---- Local dev helper ----
163
- @app.local_entrypoint()
164
- def test():
165
- import httpx
166
-
167
- print("Testing Qwen endpoint...")
168
- r = httpx.post(
169
- "http://localhost:8000/qwen-chat",
170
- json={"messages": [{"role": "user", "content": "Dis bonjour en français"}], "token": EXPECTED_TOKEN},
171
- timeout=30,
172
- )
173
- print(f"Qwen: {r.json()}")
174
-
175
- print("Testing Mistral TTS endpoint...")
176
- r = httpx.post(
177
- "http://localhost:8000/mistral-tts",
178
- json={"text": "Bonjour, comment allez-vous?", "token": EXPECTED_TOKEN},
179
- timeout=30,
180
- )
181
- print(f"TTS: {len(r.content)} bytes")
 
1
+ import base64
2
  import os
3
  import tempfile
4
+
5
  import modal
6
 
7
  app = modal.App("patient-virtuel")
8
 
9
+ # ---- Shared volumes ----
10
  qwen_vol = modal.Volume.from_name("qwen-cache", create_if_missing=True)
11
  whisper_vol = modal.Volume.from_name("whisper-cache", create_if_missing=True)
12
  CACHE_DIR = "/root/.cache/huggingface"
13
 
 
 
 
 
 
 
 
 
 
14
  # ---- 1. Qwen/Qwen3.6-27B LLM ----
15
  qwen_image = (
16
  modal.Image.debian_slim(python_version="3.12")
17
+ .pip_install("vllm>=0.6.0", "bitsandbytes>=0.43", "huggingface-hub", "torch>=2.5", "fastapi[standard]")
 
 
 
 
18
  )
19
 
 
 
 
20
 
21
  @app.function(
22
  image=qwen_image,
23
+ gpu="A100:1",
24
  volumes={CACHE_DIR: qwen_vol},
25
+ secrets=[modal.Secret.from_name("hf-token"), modal.Secret.from_name("app-tokens")],
26
+ timeout=600,
27
+ scaledown_window=600,
 
 
28
  )
29
+ @modal.asgi_app()
30
+ def qwen_web():
31
+ from fastapi import FastAPI, HTTPException, Request
32
+
33
+ fastapi_app = FastAPI()
34
 
35
+ @fastapi_app.post("/")
36
+ async def handler(request: Request):
37
+ body = await request.json()
38
+ token = body.get("token", "")
39
+ expected = os.environ.get("EXPECTED_TOKEN", "")
40
+ if not token or token != expected:
41
+ raise HTTPException(401, "Unauthorized")
42
 
43
+ messages = body["messages"]
 
 
 
 
 
44
 
45
+ from vllm import LLM, SamplingParams
 
 
 
 
 
 
46
 
47
+ llm = LLM(
48
+ model="Qwen/Qwen3.6-27B",
49
+ dtype="bfloat16",
50
+ quantization="bitsandbytes",
51
+ load_format="bitsandbytes",
52
+ max_model_len=16384,
53
+ trust_remote_code=True,
54
+ )
55
 
56
+ sp = SamplingParams(temperature=0.7, top_p=0.8, top_k=20, max_tokens=512)
57
+ outputs = llm.chat(messages=messages, sampling_params=sp, use_tqdm=False)
58
 
59
+ text = outputs[0].outputs[0].text.strip()
60
+ if "<think>" in text:
61
+ text = text.split("</think>")[-1].strip()
62
+ return {"text": text}
63
 
64
+ return fastapi_app
65
 
66
 
67
+ # ---- 2. Whisper STT (JSON input — base64-encoded WAV) ----
68
  whisper_image = (
69
  modal.Image.debian_slim(python_version="3.12")
70
+ .pip_install("faster-whisper", "numpy", "fastapi[standard]")
71
  )
72
 
73
+
74
  @app.function(
75
  image=whisper_image,
76
+ gpu="A10G:1",
77
  volumes={CACHE_DIR: whisper_vol},
78
  timeout=120,
79
  scaledown_window=300,
 
 
80
  )
81
+ @modal.asgi_app()
82
+ def whisper_web():
83
+ from fastapi import FastAPI, HTTPException, Request
84
 
85
+ fastapi_app = FastAPI()
 
86
 
87
+ @fastapi_app.post("/")
88
+ async def handler(request: Request):
89
+ body = await request.json()
90
+ token = body.get("token", "")
91
+ expected = os.environ.get("EXPECTED_TOKEN", "")
92
+ if not token or token != expected:
93
+ raise HTTPException(401, "Unauthorized")
94
 
95
+ audio_base64 = body["audio_base64"]
 
 
 
 
 
96
 
97
+ from faster_whisper import WhisperModel
 
 
98
 
99
+ model = WhisperModel(
100
+ "large-v3-turbo", device="cuda", compute_type="float16", download_root=CACHE_DIR,
101
+ )
 
 
 
 
 
102
 
103
+ raw = base64.b64decode(audio_base64)
104
+ with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
105
+ f.write(raw)
106
+ path = f.name
107
 
108
+ segments, _ = model.transcribe(path, language="fr", beam_size=5, vad_filter=True)
109
+ os.unlink(path)
110
 
111
+ return {"text": " ".join(s.text for s in segments)}
112
+
113
+ return fastapi_app
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
stt_engine.py CHANGED
@@ -1,3 +1,4 @@
 
1
  import os
2
  import httpx
3
 
@@ -10,14 +11,12 @@ def transcribe(audio_path: str) -> str | None:
10
  raise RuntimeError("MODAL_ENDPOINT_WHISPER not set")
11
 
12
  with open(audio_path, "rb") as f:
13
- audio_bytes = f.read()
14
 
15
  resp = httpx.post(
16
  MODAL_ENDPOINT,
17
- data={"token": MODAL_AUTH_TOKEN},
18
- files={"audio_bytes": audio_bytes},
19
  timeout=30,
20
  )
21
  resp.raise_for_status()
22
- data = resp.json()
23
- return data.get("text")
 
1
+ import base64
2
  import os
3
  import httpx
4
 
 
11
  raise RuntimeError("MODAL_ENDPOINT_WHISPER not set")
12
 
13
  with open(audio_path, "rb") as f:
14
+ b64 = base64.b64encode(f.read()).decode()
15
 
16
  resp = httpx.post(
17
  MODAL_ENDPOINT,
18
+ json={"audio_base64": b64, "token": MODAL_AUTH_TOKEN},
 
19
  timeout=30,
20
  )
21
  resp.raise_for_status()
22
+ return resp.json().get("text")
 
tts_engine.py CHANGED
@@ -1,21 +1,18 @@
1
  import os
2
  import httpx
3
 
4
- MODAL_ENDPOINT = os.environ.get("MODAL_ENDPOINT_MISTRAL_TTS", "")
5
  MISTRAL_API_KEY = os.environ.get("MISTRAL_API_KEY", "")
6
- MODAL_AUTH_TOKEN = os.environ.get("MODAL_AUTH_TOKEN", "")
7
-
8
- VOICE = "french_female"
9
- MODEL = "mistral-tts"
10
 
11
 
12
  def synthesize(text: str) -> bytes | None:
13
- """Tier-1: call Mistral API directly from the Space (simplest, free tier)."""
14
  if not MISTRAL_API_KEY:
15
- return _fallback_modal(text)
16
 
17
  resp = httpx.post(
18
- "https://api.mistral.ai/v1/audio/speech",
19
  headers={"Authorization": f"Bearer {MISTRAL_API_KEY}"},
20
  json={
21
  "model": MODEL,
@@ -26,19 +23,5 @@ def synthesize(text: str) -> bytes | None:
26
  timeout=25,
27
  )
28
  if resp.is_error:
29
- return _fallback_modal(text)
30
- return resp.content
31
-
32
-
33
- def _fallback_modal(text: str) -> bytes | None:
34
- """Tier-2: fall back to Modal-hosted Mistral TTS."""
35
- if not MODAL_ENDPOINT:
36
- raise RuntimeError("No TTS endpoint available")
37
-
38
- resp = httpx.post(
39
- MODAL_ENDPOINT,
40
- json={"text": text, "token": MODAL_AUTH_TOKEN},
41
- timeout=30,
42
- )
43
- resp.raise_for_status()
44
  return resp.content
 
1
  import os
2
  import httpx
3
 
 
4
  MISTRAL_API_KEY = os.environ.get("MISTRAL_API_KEY", "")
5
+ VOICE = os.environ.get("TTS_VOICE", "french_female")
6
+ MODEL = os.environ.get("TTS_MODEL", "mistral-tts")
7
+ API_URL = "https://api.mistral.ai/v1/audio/speech"
 
8
 
9
 
10
  def synthesize(text: str) -> bytes | None:
 
11
  if not MISTRAL_API_KEY:
12
+ raise RuntimeError("MISTRAL_API_KEY not set")
13
 
14
  resp = httpx.post(
15
+ API_URL,
16
  headers={"Authorization": f"Bearer {MISTRAL_API_KEY}"},
17
  json={
18
  "model": MODEL,
 
23
  timeout=25,
24
  )
25
  if resp.is_error:
26
+ raise RuntimeError(f"Mistral TTS error: {resp.status_code} {resp.text}")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
27
  return resp.content