David Robinson Milad Alizadeh commited on
Commit
7c474ca
·
1 Parent(s): 515204e

truncate audio to 10s and warn (#123)

Browse files

Truncate to 10s for model. Warn if going to truncate.

---------

Co-authored-by: Milad Alizadeh <git@mil.ad>

Files changed (2) hide show
  1. Dockerfile +1 -3
  2. 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 -b update-hf-app --single-branch --depth 1 https://$(cat /run/secrets/GH_TOKEN)@github.com/earthspecies/esp-research.git /app/esp-research && \
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 requirements.
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 outside [`MIN_AUDIO_DURATION`,
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
- if duration > MAX_AUDIO_DURATION:
60
- raise gr.Error(f"Audio duration must be at most {MAX_AUDIO_DURATION} seconds.")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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"&#9432; 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],