Milad Alizadeh commited on
Commit
515204e
·
1 Parent(s): 2cec591

tweaks to beats config/init (#114)

Browse files

<!-- av pr metadata
This information is embedded by the av CLI when creating PRs to track
the status of stacks when using Aviator. Please do not delete or edit
this section of the PR.
```
{"parent":"main","parentHead":"","trunk":"main"}
```
-->

Files changed (5) hide show
  1. Dockerfile +2 -1
  2. app.py +52 -133
  3. hub_logger.py +40 -25
  4. infer.py +0 -347
  5. pyproject.toml +1 -0
Dockerfile CHANGED
@@ -11,6 +11,7 @@ COPY --from=ghcr.io/astral-sh/uv:latest /uv /uvx /bin/
11
  RUN apt-get update && apt-get install -y \
12
  git \
13
  git-lfs \
 
14
  && apt-get clean \
15
  && rm -rf /var/lib/apt/lists/* \
16
  && git lfs install
@@ -18,7 +19,7 @@ RUN apt-get update && apt-get install -y \
18
  # TODO: Pin esp-research and esp-data revisions
19
  # TODO: remove hf-app branch once merged
20
  RUN --mount=type=secret,id=GH_TOKEN,mode=0444,required=true \
21
- git clone -b hf-app --single-branch --depth 1 https://$(cat /run/secrets/GH_TOKEN)@github.com/earthspecies/esp-research.git /app/esp-research && \
22
  git clone --single-branch --depth 1 https://$(cat /run/secrets/GH_TOKEN)@github.com/earthspecies/esp-data.git /app/esp-data
23
 
24
  # esp-research installs esp-data from gcloud artifact registry, which is not
 
11
  RUN apt-get update && apt-get install -y \
12
  git \
13
  git-lfs \
14
+ ffmpeg \
15
  && apt-get clean \
16
  && rm -rf /var/lib/apt/lists/* \
17
  && git lfs install
 
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
app.py CHANGED
@@ -3,6 +3,7 @@ import uuid
3
  from pathlib import Path
4
 
5
  import gradio as gr
 
6
  import matplotlib.pyplot as plt
7
  import numpy as np
8
  import soundfile as sf
@@ -11,49 +12,34 @@ import torch
11
  import torchaudio
12
 
13
  from esp_research.logging import logger
14
- from hub_logger import upload_data
15
-
16
- # from NatureLM.infer import Pipeline
17
- # from NatureLM.models.NatureLM import NatureLM
18
- from naturelm_audio import NatureLM # noqa: F401
19
 
20
  APP_DIR = Path(__file__).resolve().parent
21
  STATIC_DIR = APP_DIR / "static"
22
  ASSETS_DIR = APP_DIR / "assets"
23
 
24
- SAMPLE_RATE = 16000 # Default sample rate for NatureLM-audio
 
25
  MIN_AUDIO_DURATION: float = 0.5 # seconds
26
- MAX_HISTORY_TURNS = 3 # Maximum number of conversation turns to include in context (user + assistant pairs)
 
27
 
28
- DEVICE: str = "cuda" if torch.cuda.is_available() else "cpu"
 
29
 
30
  # TODO: derive model version from model metadata or config instead of hardcoding
31
- MODEL_VERSION = "1.5"
32
-
33
-
34
- class _MockModel:
35
- """Placeholder model that returns dummy predictions."""
36
-
37
- def __call__(
38
- self,
39
- audios: list[str],
40
- queries: list[str],
41
- **kwargs: object,
42
- ) -> list[list[dict]]:
43
- return [[{"prediction": "(mock) I don't know yet!"}] for _ in audios]
44
-
45
 
46
- # TODO: replace with real model loading
 
 
 
47
 
48
- # model = NatureLM.from_pretrained("EarthSpeciesProject/NatureLM-audio")
49
- # model = model.eval().to(DEVICE)
50
- # model = Pipeline(model)
51
- logger.info("Device: %s", DEVICE)
52
- model = _MockModel()
53
 
54
-
55
- def validate_audio_duration(audio_path: str) -> None:
56
- """Validate that the audio file meets the minimum duration requirement.
57
 
58
  Parameters
59
  ----------
@@ -63,53 +49,18 @@ def validate_audio_duration(audio_path: str) -> None:
63
  Raises
64
  ------
65
  Error
66
- If the audio duration is less than `MIN_AUDIO_DURATION`.
 
67
  """
68
  info = sf.info(audio_path)
69
- duration = info.duration # info.num_frames / info.sample_rate
70
  if duration < MIN_AUDIO_DURATION:
71
  raise gr.Error(f"Audio duration must be at least {MIN_AUDIO_DURATION} seconds.")
 
 
72
 
73
 
74
  @spaces.GPU
75
- def prompt_lm(
76
- audios: list[str],
77
- queries: list[str] | str,
78
- window_length_seconds: float = 10.0,
79
- hop_length_seconds: float = 10.0,
80
- ) -> list[str]:
81
- """Generate response using the model.
82
-
83
- Parameters
84
- ----------
85
- audios : list[str]
86
- List of audio file paths.
87
- queries : list[str] | str
88
- Query or list of queries to process.
89
- window_length_seconds : float
90
- Length of the window for processing audio.
91
- hop_length_seconds : float
92
- Hop length for processing audio.
93
-
94
- Returns
95
- -------
96
- list[list[dict]]
97
- Nested list of prediction dictionaries for each audio-query pair.
98
- """
99
- if model is None:
100
- return "❌ Model not loaded. Please check the model configuration."
101
-
102
- with torch.amp.autocast(device_type="cuda", dtype=torch.float16):
103
- results: list[list[dict]] = model(
104
- audios,
105
- queries,
106
- window_length_seconds=window_length_seconds,
107
- hop_length_seconds=hop_length_seconds,
108
- input_sample_rate=None,
109
- )
110
- return results
111
-
112
-
113
  def get_response(chatbot_history: list[dict], audio_input: str) -> list[dict]:
114
  """Generate response from the model based on user input and audio file.
115
 
@@ -134,56 +85,38 @@ def get_response(chatbot_history: list[dict], audio_input: str) -> list[dict]:
134
  " Consider starting a new conversation with the Clear button."
135
  )
136
 
137
- # Build conversation context from history
138
- conversation_context = []
139
- for message in chatbot_history:
140
- if message["role"] == "user":
141
- conversation_context.append(f"User: {message['content']}")
142
- elif message["role"] == "assistant":
143
- conversation_context.append(f"Assistant: {message['content']}")
144
-
145
- # Get the last user message
146
- last_user_message = ""
147
- for message in reversed(chatbot_history):
148
- if message["role"] == "user":
149
- last_user_message = message["content"]
150
- break
151
-
152
- # Format the full prompt with conversation history
153
- if len(conversation_context) > 2: # More than just the current query
154
- # Include previous turns (limit to last MAX_HISTORY_TURNS exchanges)
155
- # recent_context = conversation_context[
156
- # -(MAX_HISTORY_TURNS + 1) : -1
157
- # ] # Exclude current message
158
- recent_context = conversation_context
159
-
160
- full_prompt = (
161
- "Previous conversation:\n" + "\n".join(recent_context) + "\n\nCurrent question: " + last_user_message
162
  )
163
- else:
164
- full_prompt = last_user_message
165
-
166
- logger.debug("Full prompt with history: %s", full_prompt)
167
-
168
- response = prompt_lm(
169
- audios=[audio_input],
170
- queries=[full_prompt.strip()],
171
- window_length_seconds=100_000,
172
- hop_length_seconds=100_000,
173
- )
174
- # get first item
175
- if isinstance(response, list) and len(response) > 0:
176
- response = response[0][0]["prediction"]
177
- logger.info("Model response: %s", response)
178
- else:
179
- response = "No response generated."
 
 
180
  except Exception as e:
181
  logger.exception("Error generating response: %s", e)
182
  response = "Error generating response. Please try again."
183
 
184
- # Add model response to chat history
185
  chatbot_history.append({"role": "assistant", "content": response})
186
-
187
  return chatbot_history
188
 
189
 
@@ -269,7 +202,6 @@ def add_user_query(chatbot_history: list[dict], chat_input: str) -> list[dict]:
269
  list[dict]
270
  Updated chat history with the user message appended.
271
  """
272
- # Validate input
273
  if not chat_input.strip():
274
  return chatbot_history
275
 
@@ -283,7 +215,7 @@ def log_to_hub(chatbot_history: list[dict], audio: str, session_id: str) -> None
283
  return
284
  user_text = chatbot_history[-2]["content"]
285
  model_response = chatbot_history[-1]["content"]
286
- upload_data(audio, user_text, model_response, session_id, model_version=MODEL_VERSION)
287
 
288
 
289
  def main() -> gr.Blocks:
@@ -331,14 +263,6 @@ def main() -> gr.Blocks:
331
  with gr.Tabs():
332
  with gr.Tab("Analyze Audio"):
333
  session_id = gr.State(str(uuid.uuid4()))
334
- # uploaded_audio = gr.State()
335
- # Status indicator
336
- # status_text = gr.Textbox(
337
- # value=model_manager.get_status(),
338
- # label="Model Status",
339
- # interactive=False,
340
- # visible=True,
341
- # )
342
 
343
  with gr.Column(visible=True) as onboarding_message:
344
  gr.HTML(
@@ -351,11 +275,11 @@ def main() -> gr.Blocks:
351
  container=True,
352
  interactive=True,
353
  sources=["upload"],
 
354
  )
355
- # check that audio duration is greater than MIN_AUDIO_DURATION
356
- # raise
357
  audio_input.change(
358
- fn=validate_audio_duration,
359
  inputs=[audio_input],
360
  outputs=[],
361
  )
@@ -475,11 +399,6 @@ def main() -> gr.Blocks:
475
  outputs=[plotter],
476
  )
477
 
478
- # When submit clicked first:
479
- # 1. Validate and add user query to chat history
480
- # 2. Get response from model
481
- # 3. Clear the chat input box
482
- # 4. Show clear button
483
  chat_input.submit(
484
  validate_and_submit,
485
  inputs=[chatbot, chat_input],
 
3
  from pathlib import Path
4
 
5
  import gradio as gr
6
+ import librosa
7
  import matplotlib.pyplot as plt
8
  import numpy as np
9
  import soundfile as sf
 
12
  import torchaudio
13
 
14
  from esp_research.logging import logger
15
+ from hub_logger import log_interaction
16
+ from naturelm_audio import GenerationConfig, NatureLM
 
 
 
17
 
18
  APP_DIR = Path(__file__).resolve().parent
19
  STATIC_DIR = APP_DIR / "static"
20
  ASSETS_DIR = APP_DIR / "assets"
21
 
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"
30
 
31
  # TODO: derive model version from model metadata or config instead of hardcoding
32
+ MODEL_VERSION = "1.1"
33
+ MODEL_REPO_ID = "EarthSpeciesProject/naturelm-audio-1.1.00-private"
 
 
 
 
 
 
 
 
 
 
 
 
34
 
35
+ logger.info("Loading model from %s …", MODEL_REPO_ID)
36
+ model = NatureLM.from_hf_hub(MODEL_REPO_ID)
37
+ model = model.eval().to(DEVICE)
38
+ 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
  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
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
64
  def get_response(chatbot_history: list[dict], audio_input: str) -> list[dict]:
65
  """Generate response from the model based on user input and audio file.
66
 
 
85
  " Consider starting a new conversation with the Clear button."
86
  )
87
 
88
+ # Load audio, mix to mono, and resample to model sample rate if needed
89
+ audio_np, sr = sf.read(audio_input, dtype="float32")
90
+ if audio_np.ndim > 1:
91
+ audio_np = np.mean(audio_np, axis=1)
92
+ if sr != SAMPLE_RATE:
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().
99
+ # Gradio may return content as a list of parts on subsequent turns,
100
+ # so normalise to plain strings first.
101
+ messages: list[dict[str, str]] = []
102
+ for msg in chatbot_history:
103
+ text = msg["content"] if isinstance(msg["content"], str) else msg["content"][0]["text"]
104
+ if msg["role"] in ("user", "assistant"):
105
+ messages.append({"role": msg["role"], "content": text})
106
+
107
+ logger.debug("Messages: %s", messages)
108
+
109
+ response = model.generate(
110
+ audio=[audio_tensor],
111
+ messages=[messages],
112
+ generation_config=GenerationConfig(merging_alpha=0.7),
113
+ )[0]
114
+ logger.info("Model response: %s", response)
115
  except Exception as e:
116
  logger.exception("Error generating response: %s", e)
117
  response = "Error generating response. Please try again."
118
 
 
119
  chatbot_history.append({"role": "assistant", "content": response})
 
120
  return chatbot_history
121
 
122
 
 
202
  list[dict]
203
  Updated chat history with the user message appended.
204
  """
 
205
  if not chat_input.strip():
206
  return chatbot_history
207
 
 
215
  return
216
  user_text = chatbot_history[-2]["content"]
217
  model_response = chatbot_history[-1]["content"]
218
+ log_interaction(audio, user_text, model_response, session_id, model_version=MODEL_VERSION)
219
 
220
 
221
  def main() -> gr.Blocks:
 
263
  with gr.Tabs():
264
  with gr.Tab("Analyze Audio"):
265
  session_id = gr.State(str(uuid.uuid4()))
 
 
 
 
 
 
 
 
266
 
267
  with gr.Column(visible=True) as onboarding_message:
268
  gr.HTML(
 
275
  container=True,
276
  interactive=True,
277
  sources=["upload"],
278
+ type="filepath",
279
  )
280
+ # Validate audio duration and sample rate on upload
 
281
  audio_input.change(
282
+ fn=validate_audio,
283
  inputs=[audio_input],
284
  outputs=[],
285
  )
 
399
  outputs=[plotter],
400
  )
401
 
 
 
 
 
 
402
  chat_input.submit(
403
  validate_and_submit,
404
  inputs=[chatbot, chat_input],
hub_logger.py CHANGED
@@ -1,3 +1,5 @@
 
 
1
  import json
2
  import os
3
  import uuid
@@ -8,59 +10,72 @@ from huggingface_hub import HfApi, HfFileSystem
8
  DATASET_REPO = "EarthSpeciesProject/naturelm-audio-space-logs"
9
  SPLIT = "test"
10
  TESTING = os.getenv("TESTING", "0") == "1"
11
- api = HfApi(token=os.getenv("HF_TOKEN", None))
12
- # Upload audio
13
- # check if file exists
14
- hf_fs = HfFileSystem(token=os.getenv("HF_TOKEN", None))
 
 
15
 
16
 
17
- def upload_data(
18
  audio: str | Path,
19
  user_text: str,
20
  model_response: str,
21
  session_id: str = "",
22
  model_version: str = "",
23
  ) -> None:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24
  data_id = str(uuid.uuid4())
25
 
26
  if TESTING:
27
  data_id = "test-" + data_id
28
  session_id = "test-" + session_id
29
 
30
- # Audio path in repo
31
  suffix = Path(audio).suffix
32
- audio_p = f"{SPLIT}/audio/" + session_id + suffix
33
 
34
- if not hf_fs.exists(f"datasets/{DATASET_REPO}/{audio_p}"):
35
- api.upload_file(
36
  path_or_fileobj=str(audio),
37
- path_in_repo=audio_p,
38
  repo_id=DATASET_REPO,
39
  repo_type="dataset",
40
  )
41
 
42
- text = {
43
  "user_message": user_text,
44
  "model_response": model_response,
45
- "file_name": "audio/" + session_id + suffix, # has to be relative to metadata.jsonl
46
- "original_fn": os.path.basename(audio),
47
  "id": data_id,
48
  "session_id": session_id,
49
  "model_version": model_version,
50
  }
51
 
52
- # Append to a jsonl file in the repo
53
- # APPEND DOESN'T WORK, have to open first
54
- if hf_fs.exists(f"datasets/{DATASET_REPO}/{SPLIT}/metadata.jsonl"):
55
- with hf_fs.open(f"datasets/{DATASET_REPO}/{SPLIT}/metadata.jsonl", "r") as f:
 
56
  lines = f.readlines()
57
- lines.append(json.dumps(text) + "\n")
58
- with hf_fs.open(f"datasets/{DATASET_REPO}/{SPLIT}/metadata.jsonl", "w") as f:
59
  f.writelines(lines)
60
  else:
61
- with hf_fs.open(f"datasets/{DATASET_REPO}/{SPLIT}/metadata.jsonl", "w") as f:
62
- f.write(json.dumps(text) + "\n")
63
-
64
- # Write a separate file instead
65
- # with hf_fs.open(f"datasets/{DATASET_REPO}/{data_id}.json", "w") as f:
66
- # json.dump(text, f)
 
1
+ """Log user interactions to a HuggingFace Hub dataset."""
2
+
3
  import json
4
  import os
5
  import uuid
 
10
  DATASET_REPO = "EarthSpeciesProject/naturelm-audio-space-logs"
11
  SPLIT = "test"
12
  TESTING = os.getenv("TESTING", "0") == "1"
13
+
14
+ _hf_token = os.getenv("HF_TOKEN", None)
15
+ _api = HfApi(token=_hf_token)
16
+ _fs = HfFileSystem(token=_hf_token)
17
+
18
+ _METADATA_PATH = f"datasets/{DATASET_REPO}/{SPLIT}/metadata.jsonl"
19
 
20
 
21
+ def log_interaction(
22
  audio: str | Path,
23
  user_text: str,
24
  model_response: str,
25
  session_id: str = "",
26
  model_version: str = "",
27
  ) -> None:
28
+ """Log a single user/model exchange (audio + text) to the Hub dataset.
29
+
30
+ Parameters
31
+ ----------
32
+ audio : str | Path
33
+ Local path to the audio file uploaded by the user.
34
+ user_text : str
35
+ The user's query text.
36
+ model_response : str
37
+ The model's generated response.
38
+ session_id : str
39
+ Unique identifier for the user session.
40
+ model_version : str
41
+ Version string of the model that produced the response.
42
+ """
43
  data_id = str(uuid.uuid4())
44
 
45
  if TESTING:
46
  data_id = "test-" + data_id
47
  session_id = "test-" + session_id
48
 
 
49
  suffix = Path(audio).suffix
50
+ audio_repo_path = f"{SPLIT}/audio/{session_id}{suffix}"
51
 
52
+ if not _fs.exists(f"datasets/{DATASET_REPO}/{audio_repo_path}"):
53
+ _api.upload_file(
54
  path_or_fileobj=str(audio),
55
+ path_in_repo=audio_repo_path,
56
  repo_id=DATASET_REPO,
57
  repo_type="dataset",
58
  )
59
 
60
+ record = {
61
  "user_message": user_text,
62
  "model_response": model_response,
63
+ "file_name": f"audio/{session_id}{suffix}",
64
+ "original_fn": Path(audio).name,
65
  "id": data_id,
66
  "session_id": session_id,
67
  "model_version": model_version,
68
  }
69
 
70
+ line = json.dumps(record) + "\n"
71
+
72
+ # HfFileSystem doesn't support append, so read-then-write.
73
+ if _fs.exists(_METADATA_PATH):
74
+ with _fs.open(_METADATA_PATH, "r") as f:
75
  lines = f.readlines()
76
+ lines.append(line)
77
+ with _fs.open(_METADATA_PATH, "w") as f:
78
  f.writelines(lines)
79
  else:
80
+ with _fs.open(_METADATA_PATH, "w") as f:
81
+ f.write(line)
 
 
 
 
infer.py DELETED
@@ -1,347 +0,0 @@
1
- # """Run NatureLM-audio over a set of audio files paths or a directory with audio files."""
2
-
3
- # import argparse
4
- # from pathlib import Path
5
-
6
- # import librosa
7
- # import numpy as np
8
- # import pandas as pd
9
- # import torch
10
-
11
- # from NatureLM.config import Config
12
- # from NatureLM.models import NatureLM
13
- # from NatureLM.processors import NatureLMAudioProcessor
14
- # from NatureLM.utils import move_to_device
15
-
16
- # _MAX_LENGTH_SECONDS = 10
17
- # _MIN_CHUNK_LENGTH_SECONDS = 0.5
18
- # _SAMPLE_RATE = 16000 # Assuming the model uses a sample rate of 16kHz
19
- # _AUDIO_FILE_EXTENSIONS = [".wav", ".mp3", ".flac", ".ogg", ".mp4"] # Add other audio file formats as needed
20
- # _DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
21
- # __root_dir = Path(__file__).parent.parent
22
- # _DEFAULT_CONFIG_PATH = __root_dir / "configs" / "inference.yml"
23
-
24
-
25
- # def load_model_and_config(
26
- # cfg_path: str | Path = _DEFAULT_CONFIG_PATH, device: str = _DEVICE
27
- # ) -> tuple[NatureLM, Config]:
28
- # """Load the NatureLM model and configuration.
29
- # Returns:
30
- # tuple: The loaded model and configuration.
31
- # """
32
- # model = NatureLM.from_pretrained("EarthSpeciesProject/NatureLM-audio")
33
- # model = model.to(device).eval()
34
- # model.llama_tokenizer.pad_token_id = model.llama_tokenizer.eos_token_id
35
- # model.llama_model.generation_config.pad_token_id = model.llama_tokenizer.pad_token_id
36
-
37
- # cfg = Config.from_sources(cfg_path)
38
- # return model, cfg
39
-
40
-
41
- # def output_template(model_output: str, start_time: float, end_time: float) -> str:
42
- # """Format the output of the model.
43
-
44
- # Returns
45
- # -------
46
- # str
47
- # Formatted string with timestamps and model output.
48
- # """
49
- # return f"#{start_time:.2f}s - {end_time:.2f}s#: {model_output}\n"
50
-
51
-
52
- # def sliding_window_inference(
53
- # audio: str | Path | np.ndarray,
54
- # query: str,
55
- # processor: NatureLMAudioProcessor,
56
- # model: NatureLM,
57
- # cfg: Config,
58
- # window_length_seconds: float = 10.0,
59
- # hop_length_seconds: float = 10.0,
60
- # input_sr: int = _SAMPLE_RATE,
61
- # device: str = _DEVICE,
62
- # ) -> list[dict[str, any]]:
63
- # """Run inference on a long audio file using sliding window approach.
64
-
65
- # Args:
66
- # audio (str | Path | np.ndarray): Path to the audio file.
67
- # query (str): Query for the model.
68
- # processor (NatureLMAudioProcessor): Audio processor.
69
- # model (NatureLM): NatureLM model.
70
- # cfg (Config): Model configuration.
71
- # window_length_seconds (float): Length of the sliding window in seconds.
72
- # hop_length_seconds (float): Hop length for the sliding window in seconds.
73
- # input_sr (int): Sample rate of the audio file.
74
-
75
- # Returns:
76
- # str: The output of the model.
77
-
78
- # Raises:
79
- # ValueError: If the audio file is too short or if the audio file path is invalid.
80
- # """
81
- # if isinstance(audio, str) or isinstance(audio, Path):
82
- # audio_array, input_sr = librosa.load(str(audio), sr=None, mono=False)
83
- # elif isinstance(audio, np.ndarray):
84
- # audio_array = audio
85
- # print(f"Using provided sample rate: {input_sr}")
86
-
87
- # audio_array = audio_array.squeeze()
88
- # if audio_array.ndim > 1:
89
- # axis_to_average = int(np.argmin(audio_array.shape))
90
- # audio_array = audio_array.mean(axis=axis_to_average)
91
- # audio_array = audio_array.squeeze()
92
-
93
- # # Do initial check that the audio is long enough
94
- # if audio_array.shape[-1] < int(_MIN_CHUNK_LENGTH_SECONDS * input_sr):
95
- # raise ValueError(f"Audio is too short. Minimum length is {_MIN_CHUNK_LENGTH_SECONDS} seconds.")
96
-
97
- # start = 0
98
- # stride = int(hop_length_seconds * input_sr)
99
- # window_length = int(window_length_seconds * input_sr)
100
- # window_id = 0
101
-
102
- # output = [] # Initialize output list
103
- # while True:
104
- # chunk = audio_array[start : start + window_length]
105
- # if chunk.shape[-1] < int(_MIN_CHUNK_LENGTH_SECONDS * input_sr):
106
- # break
107
-
108
- # # Resamples, pads, truncates and creates torch Tensor
109
- # audio_tensor, prompt_list = processor([chunk], [query], [input_sr])
110
-
111
- # input_to_model = {
112
- # "raw_wav": audio_tensor,
113
- # "prompt": prompt_list[0],
114
- # "audio_chunk_sizes": 1,
115
- # "padding_mask": torch.zeros_like(audio_tensor).to(torch.bool),
116
- # }
117
- # input_to_model = move_to_device(input_to_model, device)
118
-
119
- # # generate
120
- # prediction: str = model.generate(input_to_model, cfg.generate, prompt_list)[0]
121
-
122
- # # Post-process the prediction
123
- # # prediction = output_template(prediction, start / input_sr, (start + window_length) / input_sr)
124
- # # output += prediction
125
- # output.append(
126
- # {
127
- # "start_time": start / input_sr,
128
- # "end_time": (start + window_length) / input_sr,
129
- # "prediction": prediction,
130
- # "window_number": window_id,
131
- # }
132
- # )
133
-
134
- # # Move the window
135
- # start += stride
136
-
137
- # if start + window_length > audio_array.shape[-1]:
138
- # break
139
-
140
- # return output
141
-
142
-
143
- # class Pipeline:
144
- # """Pipeline for running NatureLM-audio inference on a list of audio files or audio arrays"""
145
-
146
- # def __init__(self, model: NatureLM = None, cfg_path: str | Path = _DEFAULT_CONFIG_PATH) -> None:
147
- # self.cfg_path = cfg_path
148
-
149
- # # Load model and config
150
- # if model is not None:
151
- # self.cfg = Config.from_sources(cfg_path)
152
- # self.model = model
153
- # else:
154
- # # Download model from hub
155
- # self.model, self.cfg = load_model_and_config(cfg_path)
156
-
157
- # self.processor = NatureLMAudioProcessor(sample_rate=_SAMPLE_RATE, max_length_seconds=_MAX_LENGTH_SECONDS)
158
-
159
- # def __call__(
160
- # self,
161
- # audios: list[str | Path | np.ndarray],
162
- # queries: str | list[str],
163
- # window_length_seconds: float = 10.0,
164
- # hop_length_seconds: float = 10.0,
165
- # input_sample_rate: int = _SAMPLE_RATE,
166
- # verbose: bool = False,
167
- # ) -> list[str]:
168
- # """Run inference on a list of audio file paths or a single audio file with a
169
- # single query or a list of queries. If multiple queries are provided,
170
- # we assume that they are in the same order as the audio files. If a single query
171
- # is provided, it will be used for all audio files.
172
-
173
- # Args:
174
- # audios (list[str | Path | np.ndarray]): List of audio file paths or a single audio
175
- # file path or audio array(s)
176
- # queries (str | list[str]): Queries for the model.
177
- # window_length_seconds (float): Length of the sliding window in seconds. Defaults to 10.0.
178
- # hop_length_seconds (float): Hop length for the sliding window in seconds. Defaults to 10.0.
179
- # input_sample_rate (int): Sample rate of the audio. Defaults to 16000, which is the model's sample rate.
180
- # verbose (bool): If True, print the output of the model for each audio file.
181
- # Defaults to False.
182
-
183
- # Returns:
184
- # list[list[dict]]: List of model outputs for each audio file. Each output is a list of dictionaries
185
- # containing the start time, end time, and prediction for each chunk of audio.
186
-
187
- # Raises:
188
- # ValueError: If the number of audio files and queries do not match.
189
- # """
190
- # if isinstance(audios, str) or isinstance(audios, Path):
191
- # audios = [audios]
192
-
193
- # if isinstance(queries, str):
194
- # queries = [queries] * len(audios)
195
-
196
- # if len(audios) != len(queries):
197
- # raise ValueError("Number of audio files and queries must match.")
198
-
199
- # # Run inference
200
- # results = []
201
- # for audio, query in zip(audios, queries, strict=False):
202
- # output = sliding_window_inference(
203
- # audio,
204
- # query,
205
- # self.processor,
206
- # self.model,
207
- # self.cfg,
208
- # window_length_seconds,
209
- # hop_length_seconds,
210
- # input_sr=input_sample_rate,
211
- # )
212
- # results.append(output)
213
- # if verbose:
214
- # print(f"Processed {audio}, model output:\n=======\n{output}\n=======")
215
- # return results
216
-
217
-
218
- # def parse_args() -> argparse.Namespace:
219
- # parser = argparse.ArgumentParser("Run NatureLM-audio inference")
220
- # parser.add_argument(
221
- # "-a",
222
- # "--audio",
223
- # type=str,
224
- # required=True,
225
- # help="Path to an audio file or a directory containing audio files",
226
- # )
227
- # parser.add_argument("-q", "--query", type=str, required=True, help="Query for the model")
228
- # parser.add_argument(
229
- # "--cfg-path",
230
- # type=str,
231
- # default="configs/inference.yml",
232
- # help="Path to the configuration file for the model",
233
- # )
234
- # parser.add_argument(
235
- # "--output_path",
236
- # type=str,
237
- # default="inference_output.jsonl",
238
- # help="Output path for the results",
239
- # )
240
- # parser.add_argument(
241
- # "--window_length_seconds",
242
- # type=float,
243
- # default=10.0,
244
- # help="Length of the sliding window in seconds",
245
- # )
246
- # parser.add_argument(
247
- # "--hop_length_seconds",
248
- # type=float,
249
- # default=10.0,
250
- # help="Hop length for the sliding window in seconds",
251
- # )
252
- # args = parser.parse_args()
253
-
254
- # return args
255
-
256
-
257
- # def main(
258
- # cfg_path: str | Path,
259
- # audio_path: str | Path,
260
- # query: str,
261
- # output_path: str,
262
- # window_length_seconds: float,
263
- # hop_length_seconds: float,
264
- # ) -> None:
265
- # """Main function to run the NatureLM-audio inference script.
266
- # It takes command line arguments for audio file path, query, output path,
267
- # window length, and hop length. It processes the audio files and saves the
268
- # results to a CSV file.
269
-
270
- # Args:
271
- # cfg_path (str | Path): Path to the configuration file.
272
- # audio_path (str | Path): Path to the audio file or directory.
273
- # query (str): Query for the model.
274
- # output_path (str): Path to save the output results.
275
- # window_length_seconds (float): Length of the sliding window in seconds.
276
- # hop_length_seconds (float): Hop length for the sliding window in seconds.
277
-
278
- # Raises:
279
- # ValueError: If the audio file path is invalid or if the query is empty.
280
- # ValueError: If no audio files are found.
281
- # ValueError: If the audio file extension is not supported.
282
- # """
283
-
284
- # # Prepare sample
285
- # audio_path = Path(audio_path)
286
- # if audio_path.is_dir():
287
- # audio_paths = []
288
- # print(f"Searching for audio files in {str(audio_path)} with extensions {', '.join(_AUDIO_FILE_EXTENSIONS)}")
289
- # for ext in _AUDIO_FILE_EXTENSIONS:
290
- # audio_paths.extend(list(audio_path.rglob(f"*{ext}")))
291
-
292
- # print(f"Found {len(audio_paths)} audio files in {str(audio_path)}")
293
- # else:
294
- # # check that the extension is valid
295
- # if not any(audio_path.suffix == ext for ext in _AUDIO_FILE_EXTENSIONS):
296
- # raise ValueError(
297
- # f"Invalid audio file extension. Supported extensions are: {', '.join(_AUDIO_FILE_EXTENSIONS)}"
298
- # )
299
- # audio_paths = [audio_path]
300
-
301
- # # check that query is not empty
302
- # if not query:
303
- # raise ValueError("Query cannot be empty")
304
- # if not audio_paths:
305
- # raise ValueError("No audio files found. Please check the path or file extensions.")
306
-
307
- # # Load model and config
308
- # model, cfg = load_model_and_config(cfg_path)
309
-
310
- # # Load audio processor
311
- # processor = NatureLMAudioProcessor(sample_rate=_SAMPLE_RATE, max_length_seconds=_MAX_LENGTH_SECONDS)
312
-
313
- # # Run inference
314
- # results = {"audio_path": [], "output": []}
315
- # for path in audio_paths:
316
- # output = sliding_window_inference(
317
- # path,
318
- # query,
319
- # processor,
320
- # model,
321
- # cfg,
322
- # window_length_seconds,
323
- # hop_length_seconds,
324
- # )
325
- # results["audio_path"].append(str(path))
326
- # results["output"].append(output)
327
- # print(f"Processed {path}, model output:\n=======\n{output}\n=======\n")
328
-
329
- # # Save results as a csv
330
- # output_path = Path(output_path)
331
- # output_path.parent.mkdir(parents=True, exist_ok=True)
332
-
333
- # df = pd.DataFrame(results)
334
- # df.to_json(output_path, orient="records", lines=True)
335
- # print(f"Results saved to {output_path}")
336
-
337
-
338
- # if __name__ == "__main__":
339
- # args = parse_args()
340
- # main(
341
- # cfg_path=args.cfg_path,
342
- # audio_path=args.audio,
343
- # query=args.query,
344
- # output_path=args.output_path,
345
- # window_length_seconds=args.window_length_seconds,
346
- # hop_length_seconds=args.hop_length_seconds,
347
- # )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
pyproject.toml CHANGED
@@ -15,6 +15,7 @@ dependencies = [
15
  "torchaudio>=2.7.1",
16
  "matplotlib>=3.10.8",
17
  "numpy>=2.3.5",
 
18
  ]
19
 
20
  [tool.uv.sources]
 
15
  "torchaudio>=2.7.1",
16
  "matplotlib>=3.10.8",
17
  "numpy>=2.3.5",
18
+ "librosa>=0.9.2",
19
  ]
20
 
21
  [tool.uv.sources]