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) @spaces.GPU 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( '