Download app.py from EarthSpeciesProject/NatureLM-Audio: direct link, hf CLI and curl.
- Browser
- Download file 21 kB
-
https://huggingface.co/spaces/EarthSpeciesProject/NatureLM-Audio/resolve/main/app.py
- Command line
-
hf download hf://spaces/EarthSpeciesProject/NatureLM-Audio/app.py
-
curl -L -o app.py https://huggingface.co/spaces/EarthSpeciesProject/NatureLM-Audio/resolve/main/app.py
21 kB
| import os | |
| import uuid | |
| from pathlib import Path | |
| import gradio as gr | |
| import librosa | |
| 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 log_interaction | |
| from naturelm_audio import GenerationConfig, NatureLM | |
| APP_DIR = Path(__file__).resolve().parent | |
| STATIC_DIR = APP_DIR / "static" | |
| ASSETS_DIR = APP_DIR / "assets" | |
| # TODO: Set these values carefully later. | |
| SAMPLE_RATE = 16000 | |
| MIN_AUDIO_DURATION: float = 0.5 # seconds | |
| MAX_HISTORY_TURNS = 3 | |
| MODEL_MAX_AUDIO_DURATION: float = 10.0 # seconds – model was trained on 10 s clips | |
| assert torch.cuda.is_available(), "CUDA is required to run this app" | |
| DEVICE = "cuda" | |
| # TODO: derive model version from model metadata or config instead of hardcoding | |
| MODEL_VERSION = "1.1" | |
| MODEL_REPO_ID = "EarthSpeciesProject/naturelm-audio-1.1.00-private" | |
| logger.info("Loading model from %s …", MODEL_REPO_ID) | |
| model = NatureLM.from_hf_hub(MODEL_REPO_ID) | |
| model = model.eval().to(DEVICE) | |
| logger.info("Model loaded successfully") | |
| def validate_audio(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 shorter than `MIN_AUDIO_DURATION`. | |
| """ | |
| info = sf.info(audio_path) | |
| duration = info.duration | |
| if duration < MIN_AUDIO_DURATION: | |
| raise gr.Error(f"Audio duration must be at least {MIN_AUDIO_DURATION} seconds.") | |
| def check_truncation_warning(audio_path: str | None) -> dict: | |
| """Return a visibility update for the truncation warning banner. | |
| Parameters | |
| ---------- | |
| audio_path : str | None | |
| Path to the uploaded audio file, or ``None`` when audio is cleared. | |
| Returns | |
| ------- | |
| dict | |
| A `gr.update` with ``visible=True`` when the audio exceeds | |
| `MODEL_MAX_AUDIO_DURATION`, otherwise ``visible=False``. | |
| """ | |
| if not audio_path: | |
| return gr.update(visible=False) | |
| try: | |
| duration = sf.info(audio_path).duration | |
| except Exception: | |
| return gr.update(visible=False) | |
| return gr.update(visible=duration > MODEL_MAX_AUDIO_DURATION) | |
| 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." | |
| ) | |
| # Load audio, mix to mono, and resample to model sample rate if needed | |
| audio_np, sr = sf.read(audio_input, dtype="float32") | |
| if audio_np.ndim > 1: | |
| audio_np = np.mean(audio_np, axis=1) | |
| if sr != SAMPLE_RATE: | |
| audio_np = librosa.resample( | |
| y=audio_np, orig_sr=sr, target_sr=SAMPLE_RATE, res_type="kaiser_best", scale=True | |
| ) | |
| max_samples = int(SAMPLE_RATE * MODEL_MAX_AUDIO_DURATION) | |
| if len(audio_np) > max_samples: | |
| audio_np = audio_np[:max_samples] | |
| audio_tensor = torch.from_numpy(audio_np).to(DEVICE) | |
| # Build chat-format messages for model.generate(). | |
| # Gradio may return content as a list of parts on subsequent turns, | |
| # so normalise to plain strings first. | |
| messages: list[dict[str, str]] = [] | |
| for msg in chatbot_history: | |
| text = msg["content"] if isinstance(msg["content"], str) else msg["content"][0]["text"] | |
| if msg["role"] in ("user", "assistant"): | |
| messages.append({"role": msg["role"], "content": text}) | |
| logger.debug("Messages: %s", messages) | |
| response = model.generate( | |
| audio=[audio_tensor], | |
| messages=[messages], | |
| generation_config=GenerationConfig(merging_alpha=0.7), | |
| )[0] | |
| logger.info("Model response: %s", response) | |
| except Exception as e: | |
| logger.exception("Error generating response: %s", e) | |
| response = "Error generating response. Please try again." | |
| chatbot_history.append({"role": "assistant", "content": response}) | |
| return chatbot_history | |
| def plot_spectrogram(audio: torch.Tensor, sample_rate: int) -> plt.Figure: | |
| """Generate a spectrogram from the audio tensor. | |
| Parameters | |
| ---------- | |
| audio : torch.Tensor | |
| Audio tensor. | |
| sample_rate : int | |
| Sample rate of the audio in Hz, used for time and frequency axis labels. | |
| 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) | |
| sample_rate = SAMPLE_RATE | |
| if audio_input: | |
| try: | |
| audio, sample_rate = torchaudio.load(audio_input) | |
| except Exception: | |
| logger.exception("Error loading audio file %s", audio_input) | |
| return plot_spectrogram(audio, sample_rate) | |
| 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. | |
| """ | |
| 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"] | |
| log_interaction(audio, user_text, model_response, session_id, model_version=MODEL_VERSION) | |
| def main() -> tuple[gr.Blocks, gr.themes.Base, str]: | |
| # 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" | |
| walrus_audio = ASSETS_DIR / "Walrus - Odobenus rosmarus.wav" | |
| examples = { | |
| "Species Identification (Lazuli Bunting)": [ | |
| str(laz_audio), | |
| "What is the common name for the focal species in the audio?", | |
| ], | |
| "Species Detection (Humpback Whale)": [ | |
| str(whale_audio), | |
| "What are the common names for the species in the audio, if any?", | |
| ], | |
| "Call Type (Green Tree Frog)": [ | |
| str(frog_audio), | |
| "What type of call is the frog making in this recording?", | |
| ], | |
| "Caption the audio (American Robin)": [ | |
| str(robin_audio), | |
| "Caption the audio, using the scientific name for any animal species.", | |
| ], | |
| "Multiple Species Identification (American Crow)": [ | |
| str(crow_audio), | |
| "List the common names of all species vocalizing in this audio clip.", | |
| ], | |
| "Taxonomy (Walrus)": [str(walrus_audio), "What is the taxonomic name of the focal species in the audio?"], | |
| } | |
| gr.set_static_paths(paths=[ASSETS_DIR]) | |
| theme = gr.themes.Base(primary_hue="blue", font=[gr.themes.GoogleFont("Noto Sans")]) | |
| css = (STATIC_DIR / "style.css").read_text() | |
| 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())) | |
| 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: | |
| truncation_warning = gr.HTML( | |
| '<div style="background:#FEFCE8; border:1px solid #F5E6A3;' | |
| " border-radius:8px; padding:10px 14px; color:#92820E;" | |
| ' font-size:14px;">' | |
| f"ⓘ Only the first {MODEL_MAX_AUDIO_DURATION:.0f}" | |
| " seconds will be analyzed. Trim to the most relevant" | |
| " section.</div>", | |
| visible=False, | |
| ) | |
| audio_input = gr.Audio( | |
| container=True, | |
| interactive=True, | |
| sources=["upload"], | |
| type="filepath", | |
| waveform_options=gr.WaveformOptions(waveform_progress_color="#3b82f6"), | |
| ) | |
| # Validate audio duration and sample rate on upload | |
| audio_input.change( | |
| fn=validate_audio, | |
| 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), 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?", | |
| "What species is vocalizing in this audio recording? Common name?", | |
| "Which of these is the focal species in the audio? Options: [add your options here]", | |
| "List the scientific names of all species vocalizing in this audio clip.", | |
| "What is the genus of the focal species in the audio?", | |
| "What is the common name of the species vocalizing in this audio recording?" | |
| " Provide your top 3 predictions in ranked order.", | |
| "What type of vocalization or call is this?", | |
| "Is the focal species an adult or juvenile?", | |
| "Caption the audio, using common names for any animal species.", | |
| "Is there a bird vocalizing in this recording? Answer: Yes or No.", | |
| "Based on the sounds, what habitat or environment do you think this was recorded in?", | |
| "How many individual vocalizations can you detect in this audio?", | |
| "First describe what you hear, then identify the species.", | |
| ], | |
| 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, truncation_warning], | |
| 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=check_truncation_warning, | |
| inputs=[audio_input], | |
| outputs=[truncation_warning], | |
| ).then( | |
| fn=make_spectrogram_figure, | |
| inputs=[audio_input], | |
| outputs=[plotter], | |
| ) | |
| 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, | |
| waveform_options=gr.WaveformOptions(waveform_progress_color="#3b82f6"), | |
| ) | |
| with gr.Tab("💡 Help"): | |
| gr.HTML((STATIC_DIR / "help.html").read_text()) | |
| return app, theme, css | |
| # Create and launch the app | |
| if __name__ == "__main__": | |
| app, theme, css = 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, | |
| css=css, | |
| root_path=root_path, | |
| allowed_paths=[str(ASSETS_DIR)], | |
| ) | |