David Robinson Milad Alizadeh commited on
Commit ·
7c474ca
1
Parent(s): 515204e
truncate audio to 10s and warn (#123)
Browse filesTruncate to 10s for model. Warn if going to truncate.
---------
Co-authored-by: Milad Alizadeh <git@mil.ad>
- Dockerfile +1 -3
- app.py +44 -7
Dockerfile
CHANGED
|
@@ -16,10 +16,8 @@ RUN apt-get update && apt-get install -y \
|
|
| 16 |
&& rm -rf /var/lib/apt/lists/* \
|
| 17 |
&& git lfs install
|
| 18 |
|
| 19 |
-
# TODO: Pin esp-research and esp-data revisions
|
| 20 |
-
# TODO: remove hf-app branch once merged
|
| 21 |
RUN --mount=type=secret,id=GH_TOKEN,mode=0444,required=true \
|
| 22 |
-
git clone -
|
| 23 |
git clone --single-branch --depth 1 https://$(cat /run/secrets/GH_TOKEN)@github.com/earthspecies/esp-data.git /app/esp-data
|
| 24 |
|
| 25 |
# esp-research installs esp-data from gcloud artifact registry, which is not
|
|
|
|
| 16 |
&& rm -rf /var/lib/apt/lists/* \
|
| 17 |
&& git lfs install
|
| 18 |
|
|
|
|
|
|
|
| 19 |
RUN --mount=type=secret,id=GH_TOKEN,mode=0444,required=true \
|
| 20 |
+
git clone --single-branch --depth 1 https://$(cat /run/secrets/GH_TOKEN)@github.com/earthspecies/esp-research.git /app/esp-research && \
|
| 21 |
git clone --single-branch --depth 1 https://$(cat /run/secrets/GH_TOKEN)@github.com/earthspecies/esp-data.git /app/esp-data
|
| 22 |
|
| 23 |
# esp-research installs esp-data from gcloud artifact registry, which is not
|
app.py
CHANGED
|
@@ -22,8 +22,8 @@ ASSETS_DIR = APP_DIR / "assets"
|
|
| 22 |
# TODO: Set these values carefully later.
|
| 23 |
SAMPLE_RATE = 16000
|
| 24 |
MIN_AUDIO_DURATION: float = 0.5 # seconds
|
| 25 |
-
MAX_AUDIO_DURATION: float = 60.0 # seconds
|
| 26 |
MAX_HISTORY_TURNS = 3
|
|
|
|
| 27 |
|
| 28 |
assert torch.cuda.is_available(), "CUDA is required to run this app"
|
| 29 |
DEVICE = "cuda"
|
|
@@ -39,7 +39,7 @@ logger.info("Model loaded successfully")
|
|
| 39 |
|
| 40 |
|
| 41 |
def validate_audio(audio_path: str) -> None:
|
| 42 |
-
"""Validate that the audio file meets duration
|
| 43 |
|
| 44 |
Parameters
|
| 45 |
----------
|
|
@@ -49,15 +49,35 @@ def validate_audio(audio_path: str) -> None:
|
|
| 49 |
Raises
|
| 50 |
------
|
| 51 |
Error
|
| 52 |
-
If the audio duration is
|
| 53 |
-
`MAX_AUDIO_DURATION`].
|
| 54 |
"""
|
| 55 |
info = sf.info(audio_path)
|
| 56 |
duration = info.duration
|
| 57 |
if duration < MIN_AUDIO_DURATION:
|
| 58 |
raise gr.Error(f"Audio duration must be at least {MIN_AUDIO_DURATION} seconds.")
|
| 59 |
-
|
| 60 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 61 |
|
| 62 |
|
| 63 |
@spaces.GPU
|
|
@@ -93,6 +113,10 @@ def get_response(chatbot_history: list[dict], audio_input: str) -> list[dict]:
|
|
| 93 |
audio_np = librosa.resample(
|
| 94 |
y=audio_np, orig_sr=sr, target_sr=SAMPLE_RATE, res_type="kaiser_best", scale=True
|
| 95 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
audio_tensor = torch.from_numpy(audio_np).to(DEVICE)
|
| 97 |
|
| 98 |
# Build chat-format messages for model.generate().
|
|
@@ -271,6 +295,15 @@ def main() -> gr.Blocks:
|
|
| 271 |
)
|
| 272 |
|
| 273 |
with gr.Column(visible=True) as upload_section:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 274 |
audio_input = gr.Audio(
|
| 275 |
container=True,
|
| 276 |
interactive=True,
|
|
@@ -355,7 +388,7 @@ def main() -> gr.Blocks:
|
|
| 355 |
return updated_history, ""
|
| 356 |
|
| 357 |
clear_button = gr.ClearButton(
|
| 358 |
-
components=[chatbot, chat_input, audio_input, plotter],
|
| 359 |
visible=False,
|
| 360 |
)
|
| 361 |
|
|
@@ -393,6 +426,10 @@ def main() -> gr.Blocks:
|
|
| 393 |
chat,
|
| 394 |
plotter,
|
| 395 |
],
|
|
|
|
|
|
|
|
|
|
|
|
|
| 396 |
).then(
|
| 397 |
fn=make_spectrogram_figure,
|
| 398 |
inputs=[audio_input],
|
|
|
|
| 22 |
# TODO: Set these values carefully later.
|
| 23 |
SAMPLE_RATE = 16000
|
| 24 |
MIN_AUDIO_DURATION: float = 0.5 # seconds
|
|
|
|
| 25 |
MAX_HISTORY_TURNS = 3
|
| 26 |
+
MODEL_MAX_AUDIO_DURATION: float = 10.0 # seconds – model was trained on 10 s clips
|
| 27 |
|
| 28 |
assert torch.cuda.is_available(), "CUDA is required to run this app"
|
| 29 |
DEVICE = "cuda"
|
|
|
|
| 39 |
|
| 40 |
|
| 41 |
def validate_audio(audio_path: str) -> None:
|
| 42 |
+
"""Validate that the audio file meets the minimum duration requirement.
|
| 43 |
|
| 44 |
Parameters
|
| 45 |
----------
|
|
|
|
| 49 |
Raises
|
| 50 |
------
|
| 51 |
Error
|
| 52 |
+
If the audio duration is shorter than `MIN_AUDIO_DURATION`.
|
|
|
|
| 53 |
"""
|
| 54 |
info = sf.info(audio_path)
|
| 55 |
duration = info.duration
|
| 56 |
if duration < MIN_AUDIO_DURATION:
|
| 57 |
raise gr.Error(f"Audio duration must be at least {MIN_AUDIO_DURATION} seconds.")
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
def check_truncation_warning(audio_path: str | None) -> dict:
|
| 61 |
+
"""Return a visibility update for the truncation warning banner.
|
| 62 |
+
|
| 63 |
+
Parameters
|
| 64 |
+
----------
|
| 65 |
+
audio_path : str | None
|
| 66 |
+
Path to the uploaded audio file, or ``None`` when audio is cleared.
|
| 67 |
+
|
| 68 |
+
Returns
|
| 69 |
+
-------
|
| 70 |
+
dict
|
| 71 |
+
A `gr.update` with ``visible=True`` when the audio exceeds
|
| 72 |
+
`MODEL_MAX_AUDIO_DURATION`, otherwise ``visible=False``.
|
| 73 |
+
"""
|
| 74 |
+
if not audio_path:
|
| 75 |
+
return gr.update(visible=False)
|
| 76 |
+
try:
|
| 77 |
+
duration = sf.info(audio_path).duration
|
| 78 |
+
except Exception:
|
| 79 |
+
return gr.update(visible=False)
|
| 80 |
+
return gr.update(visible=duration > MODEL_MAX_AUDIO_DURATION)
|
| 81 |
|
| 82 |
|
| 83 |
@spaces.GPU
|
|
|
|
| 113 |
audio_np = librosa.resample(
|
| 114 |
y=audio_np, orig_sr=sr, target_sr=SAMPLE_RATE, res_type="kaiser_best", scale=True
|
| 115 |
)
|
| 116 |
+
max_samples = int(SAMPLE_RATE * MODEL_MAX_AUDIO_DURATION)
|
| 117 |
+
if len(audio_np) > max_samples:
|
| 118 |
+
audio_np = audio_np[:max_samples]
|
| 119 |
+
|
| 120 |
audio_tensor = torch.from_numpy(audio_np).to(DEVICE)
|
| 121 |
|
| 122 |
# Build chat-format messages for model.generate().
|
|
|
|
| 295 |
)
|
| 296 |
|
| 297 |
with gr.Column(visible=True) as upload_section:
|
| 298 |
+
truncation_warning = gr.HTML(
|
| 299 |
+
'<div style="background:#FEFCE8; border:1px solid #F5E6A3;'
|
| 300 |
+
" border-radius:8px; padding:10px 14px; color:#92820E;"
|
| 301 |
+
' font-size:14px;">'
|
| 302 |
+
f"ⓘ Only the first {MODEL_MAX_AUDIO_DURATION:.0f}"
|
| 303 |
+
" seconds will be analyzed. Trim to the most relevant"
|
| 304 |
+
" section.</div>",
|
| 305 |
+
visible=False,
|
| 306 |
+
)
|
| 307 |
audio_input = gr.Audio(
|
| 308 |
container=True,
|
| 309 |
interactive=True,
|
|
|
|
| 388 |
return updated_history, ""
|
| 389 |
|
| 390 |
clear_button = gr.ClearButton(
|
| 391 |
+
components=[chatbot, chat_input, audio_input, plotter, truncation_warning],
|
| 392 |
visible=False,
|
| 393 |
)
|
| 394 |
|
|
|
|
| 426 |
chat,
|
| 427 |
plotter,
|
| 428 |
],
|
| 429 |
+
).then(
|
| 430 |
+
fn=check_truncation_warning,
|
| 431 |
+
inputs=[audio_input],
|
| 432 |
+
outputs=[truncation_warning],
|
| 433 |
).then(
|
| 434 |
fn=make_spectrogram_figure,
|
| 435 |
inputs=[audio_input],
|