kirimaru's picture
Add more examples & Fix height issue of Prediction
96e7759
Raw History Blame Contribute Delete
5.71 kB
"""
Step 1: Run Your Script
Save the above code in a file, for example, app.py. Then, run the script using Python:
python app.py
Step 2: Access the Gradio Interface
After running the script, you should see output in the terminal indicating that the Gradio interface is running. It will provide a local URL (usually http://localhost:7860) that you can open in your web browser to interact with your Gradio app.
Additional Features
Gradio supports various input and output types, including images, audio, and more. You can customize the interface further by adding more inputs and outputs, changing the layout, and adding examples.
For more advanced usage and features, you can refer to the Gradio documentation.
"""
import gradio as gr
import matplotlib.pyplot as plt
import pandas as pd
import numpy as np
import bat_detect.utils.detector_utils as du
import bat_detect.utils.audio_utils as au
import bat_detect.utils.plot_utils as viz
# setup the arguments
args = {}
args = du.get_default_bd_args()
args['detection_threshold'] = 0.3
args['time_expansion_factor'] = 1
args['model_path'] = 'models/Net2DFast_UK_same.pth.tar'
max_duration = 10.0
# load the model
model, params = du.load_model(args['model_path'])
prediction_df = gr.Dataframe(
headers=["species", "time", "detection_prob", "species_prob"],
datatype=["str", "str", "str", "str"],
row_count=1,
col_count=(4, "fixed"),
# max_height=300,
elem_classes="qa-pairs",
label='Predictions'
)
prediction_css = """
.qa-pairs .table-wrap {
min-height: 300px;
max-height: 300px;
}
#visualisation-container {
overflow-x: auto;
width: 100%;
height: 200px; /* Fixed height */
}
#visualisation-container img {
height: 100%; /* fixed height */
width: auto; /* flexible width */
object-fit: contain;
}
"""
examples = [['example_data/audio/20170701_213954-MYOMYS-LR_0_0.5.wav', 0.3],
['example_data/audio/20180530_213516-EPTSER-LR_0_0.5.wav', 0.3],
['example_data/audio/20180627_215323-RHIFER-LR_0_0.5.wav', 0.3],
['example_data/audio/Myotis daubentonii_A004034_PVTRQZEYAT.flac', 0.3],
# ['example_data/audio/Eptesicus serotinus_A003974_RLNMCAZDEJ.flac', 0.3],
# ['example_data/audio/Nyctalus noctula_A004093_DOEYCSPEZD.flac', 0.3],
]
def make_prediction(file_name=None, detection_threshold=0.3):
if file_name is not None:
audio_file = file_name
else:
return "You must provide an input audio file."
if detection_threshold is not None and detection_threshold != '':
args['detection_threshold'] = float(detection_threshold)
# process the file to generate predictions
results = du.process_file(audio_file, model, params, args, max_duration=max_duration)
# results = du.process_file(audio_file, model, params, args, max_duration=False)
anns = [ann for ann in results['pred_dict']['annotation']]
clss = [aa['class'] for aa in anns]
st_time = [aa['start_time'] for aa in anns]
cls_prob = [aa['class_prob'] for aa in anns]
det_prob = [aa['det_prob'] for aa in anns]
data = {'species': clss, 'time': st_time, 'detection_prob': det_prob, 'species_prob': cls_prob}
prediction_df = pd.DataFrame(data=data)
im = generate_results_image(audio_file, anns)
return [prediction_df, im]
def generate_results_image(audio_file, anns):
# load audio
sampling_rate, audio = au.load_audio_file(audio_file, args['time_expansion_factor'],
params['target_samp_rate'], params['scale_raw_audio'], max_duration=max_duration)
duration = audio.shape[0] / sampling_rate
# generate spec
spec, spec_viz = au.generate_spectrogram(audio, sampling_rate, params, True, False)
# create fig
plt.close('all')
# Adjust figsize calculation to control the width of the spectrogram
fig_width = max(6, spec.shape[1] / 100) # Minimum width of 6 inches, scales with spectrogram width
fig_height = spec.shape[0] / 100 # Adjust the divisor for height scaling (higher = smaller)
fig = plt.figure(1, figsize=(fig_width, fig_height), dpi=100, frameon=False)
spec_duration = au.x_coords_to_time(spec.shape[1], sampling_rate, params['fft_win_length'], params['fft_overlap'])
viz.create_box_image(spec, fig, anns, 0, spec_duration, spec_duration, params, spec.max()*1.1, False, True)
plt.ylabel('Freq - kHz')
plt.xlabel('Time - secs')
plt.tight_layout()
# convert fig to image
fig.canvas.draw()
data = np.frombuffer(fig.canvas.buffer_rgba(), dtype=np.uint8)
w, h = fig.canvas.get_width_height()
im = data.reshape((int(h), int(w), -1))
return im
descr_txt = "Demo of Bat identification tools. " \
"<br>It is based on two state-of-the-art models BatDetect2 and Bat-cli. " \
"For the demo purposes, the input file is longer than 10 seconds, only the first 10 seconds will be processed." \
# "<br>Check out the two papers for more details [here](https://www.biorxiv.org/content/10.1101/2022.12.14.520490v1) and [here](https://ar5iv.labs.arxiv.org/html/2309.11218)."
Gradio_interface = gr.Interface(
fn = make_prediction,
inputs = [gr.Audio(sources=["upload"], type="filepath"),
gr.Dropdown([0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9])],
outputs = [prediction_df, gr.Image(interactive=False, label="Visualisation", elem_id="visualisation-container")],
# theme = "huggingface",
title = "Bat Identification Demo",
description = descr_txt,
examples = examples,
allow_flagging = 'never',
css=prediction_css
)
Gradio_interface.launch()