Spaces:
Running on A100
Running on A100
Download app.py from EarthSpeciesProject/NatureLM-Audio: direct link, hf CLI and curl.
- Browser
- Download file 21.1 kB
-
https://huggingface.co/spaces/EarthSpeciesProject/NatureLM-Audio/resolve/c3156f6a5929dc83a8fc447b048089dd465bbf39/app.py
- Command line
-
hf download hf://spaces/EarthSpeciesProject/NatureLM-Audio@c3156f6a5929dc83a8fc447b048089dd465bbf39/app.py
-
curl -L -o app.py https://huggingface.co/spaces/EarthSpeciesProject/NatureLM-Audio/resolve/c3156f6a5929dc83a8fc447b048089dd465bbf39/app.py
21.1 kB
| import os | |
| import uuid | |
| from pathlib import Path | |
| import gradio as gr | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import soundfile as sf | |
| import spaces | |
| import torch | |
| import torchaudio | |
| from esp_research.logging import logger | |
| from hub_logger import upload_data | |
| # from NatureLM.infer import Pipeline | |
| # from NatureLM.models.NatureLM import NatureLM | |
| from naturelm_audio import NatureLM # noqa: F401 | |
| APP_DIR = Path(__file__).resolve().parent | |
| STATIC_DIR = APP_DIR / "static" | |
| ASSETS_DIR = APP_DIR / "assets" | |
| SAMPLE_RATE = 16000 # Default sample rate for NatureLM-audio | |
| MIN_AUDIO_DURATION: float = 0.5 # seconds | |
| MAX_HISTORY_TURNS = 3 # Maximum number of conversation turns to include in context (user + assistant pairs) | |
| DEVICE: str = "cuda" if torch.cuda.is_available() else "cpu" | |
| # TODO: derive model version from model metadata or config instead of hardcoding | |
| MODEL_VERSION = "1.5" | |
| class _MockModel: | |
| """Placeholder model that returns dummy predictions.""" | |
| def __call__( | |
| self, | |
| audios: list[str], | |
| queries: list[str], | |
| **kwargs: object, | |
| ) -> list[list[dict]]: | |
| return [[{"prediction": "(mock) I don't know yet!"}] for _ in audios] | |
| # TODO: replace with real model loading | |
| # model = NatureLM.from_pretrained("EarthSpeciesProject/NatureLM-audio") | |
| # model = model.eval().to(DEVICE) | |
| # model = Pipeline(model) | |
| logger.info("Device: %s", DEVICE) | |
| model = _MockModel() | |
| def validate_audio_duration(audio_path: str) -> None: | |
| """Validate that the audio file meets the minimum duration requirement. | |
| Parameters | |
| ---------- | |
| audio_path : str | |
| Path to the audio file. | |
| Raises | |
| ------ | |
| Error | |
| If the audio duration is less than `MIN_AUDIO_DURATION`. | |
| """ | |
| info = sf.info(audio_path) | |
| duration = info.duration # info.num_frames / info.sample_rate | |
| if duration < MIN_AUDIO_DURATION: | |
| raise gr.Error(f"Audio duration must be at least {MIN_AUDIO_DURATION} seconds.") | |
| def prompt_lm( | |
| audios: list[str], | |
| queries: list[str] | str, | |
| window_length_seconds: float = 10.0, | |
| hop_length_seconds: float = 10.0, | |
| ) -> list[str]: | |
| """Generate response using the model. | |
| Parameters | |
| ---------- | |
| audios : list[str] | |
| List of audio file paths. | |
| queries : list[str] | str | |
| Query or list of queries to process. | |
| window_length_seconds : float | |
| Length of the window for processing audio. | |
| hop_length_seconds : float | |
| Hop length for processing audio. | |
| Returns | |
| ------- | |
| list[list[dict]] | |
| Nested list of prediction dictionaries for each audio-query pair. | |
| """ | |
| if model is None: | |
| return "❌ Model not loaded. Please check the model configuration." | |
| with torch.amp.autocast(device_type="cuda", dtype=torch.float16): | |
| results: list[list[dict]] = model( | |
| audios, | |
| queries, | |
| window_length_seconds=window_length_seconds, | |
| hop_length_seconds=hop_length_seconds, | |
| input_sample_rate=None, | |
| ) | |
| return results | |
| def get_response(chatbot_history: list[dict], audio_input: str) -> list[dict]: | |
| """Generate response from the model based on user input and audio file. | |
| Parameters | |
| ---------- | |
| chatbot_history : list[dict] | |
| Current chat history with conversation context. | |
| audio_input : str | |
| Path to the audio file. | |
| Returns | |
| ------- | |
| list[dict] | |
| Updated chat history with model response appended. | |
| """ | |
| try: | |
| # Warn if conversation is getting long | |
| num_turns = len(chatbot_history) | |
| if num_turns > MAX_HISTORY_TURNS * 2: # Each turn = user + assistant message | |
| gr.Warning( | |
| "⚠️ Long conversations may affect response quality." | |
| " Consider starting a new conversation with the Clear button." | |
| ) | |
| # Build conversation context from history | |
| conversation_context = [] | |
| for message in chatbot_history: | |
| if message["role"] == "user": | |
| conversation_context.append(f"User: {message['content']}") | |
| elif message["role"] == "assistant": | |
| conversation_context.append(f"Assistant: {message['content']}") | |
| # Get the last user message | |
| last_user_message = "" | |
| for message in reversed(chatbot_history): | |
| if message["role"] == "user": | |
| last_user_message = message["content"] | |
| break | |
| # Format the full prompt with conversation history | |
| if len(conversation_context) > 2: # More than just the current query | |
| # Include previous turns (limit to last MAX_HISTORY_TURNS exchanges) | |
| # recent_context = conversation_context[ | |
| # -(MAX_HISTORY_TURNS + 1) : -1 | |
| # ] # Exclude current message | |
| recent_context = conversation_context | |
| full_prompt = ( | |
| "Previous conversation:\n" + "\n".join(recent_context) + "\n\nCurrent question: " + last_user_message | |
| ) | |
| else: | |
| full_prompt = last_user_message | |
| logger.debug("Full prompt with history: %s", full_prompt) | |
| response = prompt_lm( | |
| audios=[audio_input], | |
| queries=[full_prompt.strip()], | |
| window_length_seconds=100_000, | |
| hop_length_seconds=100_000, | |
| ) | |
| # get first item | |
| if isinstance(response, list) and len(response) > 0: | |
| response = response[0][0]["prediction"] | |
| logger.info("Model response: %s", response) | |
| else: | |
| response = "No response generated." | |
| except Exception as e: | |
| logger.exception("Error generating response: %s", e) | |
| response = "Error generating response. Please try again." | |
| # Add model response to chat history | |
| chatbot_history.append({"role": "assistant", "content": response}) | |
| return chatbot_history | |
| def plot_spectrogram(audio: torch.Tensor) -> plt.Figure: | |
| """Generate a spectrogram from the audio tensor. | |
| Parameters | |
| ---------- | |
| audio : torch.Tensor | |
| Audio tensor. | |
| Returns | |
| ------- | |
| plt.Figure | |
| Matplotlib figure with the spectrogram. | |
| """ | |
| spectrogram = torchaudio.transforms.Spectrogram(n_fft=1024)(audio) | |
| spectrogram = spectrogram.numpy()[0].squeeze() | |
| fig, ax = plt.subplots(figsize=(13, 5)) | |
| ax.imshow(np.log(spectrogram + 1e-4), aspect="auto", origin="lower", cmap="viridis") | |
| ax.set_title("Spectrogram") | |
| # Set x ticks to reflect 0 to audio duration seconds | |
| if audio.dim() > 1: | |
| duration = audio.size(1) / SAMPLE_RATE | |
| else: | |
| duration = audio.size(0) / SAMPLE_RATE | |
| ax.set_xlabel("Time") | |
| ax.set_xticks([0, spectrogram.shape[1]]) | |
| ax.set_xticklabels(["0s", f"{duration:.2f}s"]) | |
| ax.set_ylabel("Frequency") | |
| ax.set_yticks( | |
| [ | |
| 0, | |
| spectrogram.shape[0] // 4, | |
| spectrogram.shape[0] // 2, | |
| 3 * spectrogram.shape[0] // 4, | |
| spectrogram.shape[0] - 1, | |
| ] | |
| ) | |
| # Set y ticks to reflect 0 to nyquist frequency (sample_rate/2) | |
| nyquist_freq = SAMPLE_RATE / 2 | |
| ax.set_yticklabels( | |
| [ | |
| "0 Hz", | |
| f"{nyquist_freq / 4:.0f} Hz", | |
| f"{nyquist_freq / 2:.0f} Hz", | |
| f"{3 * nyquist_freq / 4:.0f} Hz", | |
| f"{nyquist_freq:.0f} Hz", | |
| ] | |
| ) | |
| fig.tight_layout() | |
| return fig | |
| def make_spectrogram_figure(audio_input: str) -> plt.Figure: | |
| audio = torch.zeros(1, SAMPLE_RATE) | |
| if audio_input: | |
| try: | |
| audio, _ = torchaudio.load(audio_input) | |
| except Exception: | |
| logger.exception("Error loading audio file %s", audio_input) | |
| return plot_spectrogram(audio) | |
| def add_user_query(chatbot_history: list[dict], chat_input: str) -> list[dict]: | |
| """Add user message to chat history. | |
| Parameters | |
| ---------- | |
| chatbot_history : list[dict] | |
| Current chat history. | |
| chat_input : str | |
| User's input text. | |
| Returns | |
| ------- | |
| list[dict] | |
| Updated chat history with the user message appended. | |
| """ | |
| # Validate input | |
| if not chat_input.strip(): | |
| return chatbot_history | |
| chatbot_history.append({"role": "user", "content": chat_input.strip()}) | |
| return chatbot_history | |
| def log_to_hub(chatbot_history: list[dict], audio: str, session_id: str) -> None: | |
| """Upload data to hub.""" | |
| if not chatbot_history or len(chatbot_history) < 2: | |
| return | |
| user_text = chatbot_history[-2]["content"] | |
| model_response = chatbot_history[-1]["content"] | |
| upload_data(audio, user_text, model_response, session_id, model_version=MODEL_VERSION) | |
| def main() -> gr.Blocks: | |
| # Create placeholder audio files if they don't exist | |
| laz_audio = ASSETS_DIR / "Lazuli_Bunting_yell-YELLLAZB20160625SM303143.mp3" | |
| frog_audio = ASSETS_DIR / "nri-GreenTreeFrogEvergladesNP.mp3" | |
| robin_audio = ASSETS_DIR / "yell-YELLAMRO20160506SM3.mp3" | |
| whale_audio = ASSETS_DIR / "Humpback Whale - Megaptera novaeangliae.wav" | |
| crow_audio = ASSETS_DIR / "American Crow - Corvus brachyrhynchos.mp3" | |
| examples = { | |
| "Identifying Focal Species (Lazuli Bunting)": [ | |
| str(laz_audio), | |
| "What is the common name for the focal species in the audio?", | |
| ], | |
| "Caption the audio (Green Tree Frog)": [ | |
| str(frog_audio), | |
| "Caption the audio, using the common name for any animal species.", | |
| ], | |
| "Caption the audio (American Robin)": [ | |
| str(robin_audio), | |
| "Caption the audio, using the scientific name for any animal species.", | |
| ], | |
| "Identifying Focal Species (Megaptera novaeangliae)": [ | |
| str(whale_audio), | |
| "What is the scientific name for the focal species in the audio?", | |
| ], | |
| "Speaker Count (American Crow)": [ | |
| str(crow_audio), | |
| "How many individuals are vocalizing in this audio?", | |
| ], | |
| "Caption the audio (Humpback Whale)": [str(whale_audio), "Caption the audio."], | |
| } | |
| gr.set_static_paths(paths=[ASSETS_DIR]) | |
| theme = gr.themes.Base(primary_hue="blue", font=[gr.themes.GoogleFont("Noto Sans")]) | |
| with gr.Blocks( | |
| title="NatureLM-audio", | |
| ) as app: | |
| with gr.Row(): | |
| gr.HTML((STATIC_DIR / "header.html").read_text()) | |
| with gr.Tabs(): | |
| with gr.Tab("Analyze Audio"): | |
| session_id = gr.State(str(uuid.uuid4())) | |
| # uploaded_audio = gr.State() | |
| # Status indicator | |
| # status_text = gr.Textbox( | |
| # value=model_manager.get_status(), | |
| # label="Model Status", | |
| # interactive=False, | |
| # visible=True, | |
| # ) | |
| with gr.Column(visible=True) as onboarding_message: | |
| gr.HTML( | |
| (STATIC_DIR / "onboarding.html").read_text(), | |
| padding=False, | |
| ) | |
| with gr.Column(visible=True) as upload_section: | |
| audio_input = gr.Audio( | |
| container=True, | |
| interactive=True, | |
| sources=["upload"], | |
| ) | |
| # check that audio duration is greater than MIN_AUDIO_DURATION | |
| # raise | |
| audio_input.change( | |
| fn=validate_audio_duration, | |
| inputs=[audio_input], | |
| outputs=[], | |
| ) | |
| with gr.Accordion(label="Toggle Spectrogram", open=False, visible=False) as spectrogram: | |
| plotter = gr.Plot( | |
| plot_spectrogram(torch.zeros(1, SAMPLE_RATE)), | |
| label="Spectrogram", | |
| visible=False, | |
| elem_id="spectrogram-plot", | |
| ) | |
| with gr.Column(visible=False) as tasks: | |
| task_dropdown = gr.Dropdown( | |
| [ | |
| "What are the common names for the species in the audio, if any?", | |
| "Caption the audio, using the scientific name for any animal species.", | |
| "Caption the audio, using the common name for any animal species.", | |
| "What is the scientific name for the focal species in the audio?", | |
| "What is the common name for the focal species in the audio?", | |
| "What is the family of the focal species in the audio?", | |
| "What is the genus of the focal species in the audio?", | |
| "What is the taxonomic name of the focal species in the audio?", | |
| "What call types are heard from the focal species in the audio?", | |
| "What is the life stage of the focal species in the audio?", | |
| ], | |
| label="Pre-Loaded Tasks", | |
| info="Select a task, or write your own prompt below.", | |
| allow_custom_value=False, | |
| value=None, | |
| ) | |
| with gr.Group(visible=False) as chat: | |
| chatbot = gr.Chatbot( | |
| elem_id="chatbot", | |
| height=250, | |
| label="Chat", | |
| render_markdown=False, | |
| group_consecutive_messages=False, | |
| feedback_options=[ | |
| "like", | |
| "dislike", | |
| "wrong species", | |
| "incorrect response", | |
| "other", | |
| ], | |
| resizable=True, | |
| ) | |
| with gr.Column(): | |
| chat_input = gr.Textbox( | |
| placeholder="Type your message and press Enter to send", | |
| lines=1, | |
| show_label=False, | |
| submit_btn="Send", | |
| container=True, | |
| autofocus=False, | |
| elem_id="chat-input", | |
| ) | |
| with gr.Column(): | |
| gr.Examples( | |
| list(examples.values()), | |
| [audio_input, chat_input], | |
| [audio_input, chat_input], | |
| example_labels=list(examples.keys()), | |
| examples_per_page=20, | |
| ) | |
| def validate_and_submit(chatbot_history: list[dict], chat_input: str) -> tuple[list[dict], str]: | |
| if not chat_input or not chat_input.strip(): | |
| gr.Warning("Please enter a question or message before sending.") | |
| return chatbot_history, chat_input | |
| updated_history = add_user_query(chatbot_history, chat_input) | |
| return updated_history, "" | |
| clear_button = gr.ClearButton( | |
| components=[chatbot, chat_input, audio_input, plotter], | |
| visible=False, | |
| ) | |
| # if task_dropdown is selected, set chat_input to that value | |
| def set_query(task: str | None) -> dict: | |
| if task: | |
| return gr.update(value=task) | |
| return gr.update(value="") | |
| task_dropdown.select( | |
| fn=set_query, | |
| inputs=[task_dropdown], | |
| outputs=[chat_input], | |
| ) | |
| def start_chat_interface(audio_path: str) -> tuple: | |
| return ( | |
| gr.update(visible=False), # hide onboarding message | |
| gr.update(visible=True), # show upload section | |
| gr.update(visible=True), # show spectrogram | |
| gr.update(visible=True), # show tasks | |
| gr.update(visible=True), # show chat box | |
| gr.update(visible=True), # show plotter | |
| ) | |
| # When audio added, set spectrogram | |
| audio_input.change( | |
| fn=start_chat_interface, | |
| inputs=[audio_input], | |
| outputs=[ | |
| onboarding_message, | |
| upload_section, | |
| spectrogram, | |
| tasks, | |
| chat, | |
| plotter, | |
| ], | |
| ).then( | |
| fn=make_spectrogram_figure, | |
| inputs=[audio_input], | |
| outputs=[plotter], | |
| ) | |
| # When submit clicked first: | |
| # 1. Validate and add user query to chat history | |
| # 2. Get response from model | |
| # 3. Clear the chat input box | |
| # 4. Show clear button | |
| chat_input.submit( | |
| validate_and_submit, | |
| inputs=[chatbot, chat_input], | |
| outputs=[chatbot, chat_input], | |
| ).then( | |
| get_response, | |
| inputs=[chatbot, audio_input], | |
| outputs=[chatbot], | |
| ).then( | |
| lambda: gr.update(visible=True), # Show clear button | |
| None, | |
| [clear_button], | |
| ).then( | |
| log_to_hub, | |
| [chatbot, audio_input, session_id], | |
| None, | |
| ) | |
| clear_button.click(lambda: gr.ClearButton(visible=False), None, [clear_button]) | |
| with gr.Tab("Sample Library"): | |
| with gr.Row(): | |
| with gr.Column(): | |
| gr.Markdown("### Download Sample Audio") | |
| gr.Markdown( | |
| "Feel free to explore these sample audio files." | |
| " To download, click the button in the" | |
| " top-right corner of each audio file." | |
| " You can also find a large collection of" | |
| " publicly available animal sounds on" | |
| " [Xenocanto](https://xeno-canto.org/explore/taxonomy)" | |
| " and [Watkins Marine Mammal Sound Database]" | |
| "(https://whoicf2.whoi.edu/science/B/whalesounds/index.cfm)." | |
| ) | |
| samples = [ | |
| ( | |
| str(ASSETS_DIR / "Lazuli_Bunting_yell-YELLLAZB20160625SM303143.m4a"), | |
| "Lazuli Bunting", | |
| ), | |
| ( | |
| str(ASSETS_DIR / "nri-GreenTreeFrogEvergladesNP.mp3"), | |
| "Green Tree Frog", | |
| ), | |
| ( | |
| str(ASSETS_DIR / "American Crow - Corvus brachyrhynchos.mp3"), | |
| "American Crow", | |
| ), | |
| ( | |
| str(ASSETS_DIR / "Gray Wolf - Canis lupus italicus.m4a"), | |
| "Gray Wolf", | |
| ), | |
| ( | |
| str(ASSETS_DIR / "Humpback Whale - Megaptera novaeangliae.wav"), | |
| "Humpback Whale", | |
| ), | |
| (str(ASSETS_DIR / "Walrus - Odobenus rosmarus.wav"), "Walrus"), | |
| ] | |
| for row_i in range(0, len(samples), 3): | |
| with gr.Row(): | |
| for filepath, label in samples[row_i : row_i + 3]: | |
| with gr.Column(): | |
| gr.Audio( | |
| filepath, | |
| label=label, | |
| ) | |
| with gr.Tab("💡 Help"): | |
| gr.HTML((STATIC_DIR / "help.html").read_text()) | |
| app.css = (STATIC_DIR / "style.css").read_text() | |
| return app, theme | |
| # Create and launch the app | |
| if __name__ == "__main__": | |
| app, theme = main() | |
| # Docker-based HF Spaces require root_path so Gradio generates correct | |
| # URLs behind the reverse proxy (the Gradio SDK sets this automatically). | |
| root_path = os.environ.get("GRADIO_ROOT_PATH", "") | |
| app.launch( | |
| server_name="0.0.0.0", | |
| server_port=7860, | |
| theme=theme, | |
| root_path=root_path, | |
| allowed_paths=[str(ASSETS_DIR)], | |
| ) | |