janasumit2911 commited on
Commit
548819b
·
verified ·
1 Parent(s): fe9c6e5

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +125 -0
app.py ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ try:
2
+ import os
3
+ import tensorflow as tf
4
+ import tensorflow_hub as hub
5
+ import numpy as np
6
+ import csv
7
+ import requests
8
+ import json
9
+ import logging
10
+ import scipy
11
+ from scipy.io import wavfile
12
+ from pydub import AudioSegment
13
+ import io
14
+ from io import BytesIO
15
+
16
+ model = hub.load('Audio_Multiple_v1')
17
+
18
+ def class_names_from_csv(class_map_csv_text):
19
+ """Returns list of class names corresponding to score vector."""
20
+ class_names = []
21
+ with tf.io.gfile.GFile(class_map_csv_text) as csvfile:
22
+ reader = csv.DictReader(csvfile)
23
+ for row in reader:
24
+ class_names.append(row['display_name'])
25
+ return class_names
26
+
27
+ class_map_path = model.class_map_path().numpy()
28
+ class_names = class_names_from_csv(class_map_path)
29
+
30
+
31
+ def ensure_sample_rate(original_sample_rate, waveform, desired_sample_rate=16000):
32
+ if original_sample_rate != desired_sample_rate:
33
+ desired_length = int(round(float(len(waveform)) / original_sample_rate * desired_sample_rate))
34
+ waveform = scipy.signal.resample(waveform, desired_length)
35
+ return desired_sample_rate, waveform
36
+
37
+
38
+ def convert_mp3_to_wav(mp3_data):
39
+ audio = AudioSegment.from_file(io.BytesIO(mp3_data), format="mp3")
40
+ wav_buffer = io.BytesIO()
41
+ audio.export(wav_buffer, format='wav')
42
+ wav_buffer.seek(0)
43
+ return wav_buffer
44
+
45
+
46
+ def process_audio_file(file_data, url):
47
+ try:
48
+ sample_rate, wav_data = wavfile.read(BytesIO(file_data))
49
+ if wav_data.ndim > 1:
50
+ wav_data = np.mean(wav_data, axis=1)
51
+
52
+ sample_rate, wav_data = ensure_sample_rate(sample_rate, wav_data)
53
+ waveform = wav_data / tf.int16.max
54
+
55
+ scores, embeddings, spectrogram = model(waveform)
56
+
57
+ scores_np = scores.numpy()
58
+ spectrogram_np = spectrogram.numpy()
59
+ mean_scores = np.mean(scores, axis=0)
60
+
61
+ top_two_indices = np.argsort(mean_scores)[-2:][::-1]
62
+ inferred_class = class_names[top_two_indices[0]]
63
+
64
+ if inferred_class == "Silence" and len(top_two_indices) > 1:
65
+ inferred_class = class_names[top_two_indices[1]]
66
+
67
+ answer_dict = {'file_name': url, 'class_name': inferred_class}
68
+ return answer_dict
69
+ except Exception as e:
70
+ logging.error(f"Error processing {url}: {e}")
71
+ return None
72
+
73
+
74
+ def get_audio_data(url):
75
+ response = requests.get(url)
76
+ response.raise_for_status()
77
+ return response.content
78
+
79
+ def send_results_to_api(data, result_url):
80
+ # Example function to send results to an API
81
+ headers = {"Content-Type": "application/json"}
82
+ response = requests.post(result_url, json=data, headers=headers)
83
+ if response.status_code == 200:
84
+ return response.json() # Return any response from the API if needed
85
+ else:
86
+ return {"error": f"Failed to send results to API: {response.status_code}"}
87
+
88
+ def process_audio(params):
89
+ try:
90
+ params = json.loads(params)
91
+ except json.JSONDecodeError as e:
92
+ return {"error": f"Invalid JSON input: {e.msg} at line {e.lineno} column {e.colno}"}
93
+
94
+ audio_files = params.get("audio_files", [])
95
+ api = params.get("api", "")
96
+ job_id = params.get("job_id", "")
97
+
98
+ solutions = []
99
+ for audio_url in audio_files:
100
+ audio_data = get_audio_data(audio_url)
101
+
102
+ if audio_url.endswith(".mp3"):
103
+ wav_data = convert_mp3_to_wav(audio_data)
104
+ result = process_audio_file(wav_data, audio_url)
105
+
106
+ elif audio_url.endswith(".wav"):
107
+ result = process_audio_file(audio_data, audio_url)
108
+
109
+ if result:
110
+ solutions.append(result)
111
+
112
+ result_url = f"{api}/{job_id}"
113
+ send_results_to_api(solutions, result_url)
114
+
115
+ return json.dumps({"solutions": solutions}, indent=4)
116
+
117
+ import gradio as gr
118
+ inputt = gr.Textbox(label="Parameters (JSON format) Eg. {'audio_files':['file1.mp3','file2.wav'], 'api':'https://api.example.com', 'job_id':'1001'}")
119
+ outputs = gr.JSON()
120
+
121
+ application = gr.Interface(fn=process_audio, inputs=inputt, outputs=outputs, title="Audio Classification with API Integration")
122
+ application.launch()
123
+
124
+ except Exception as e:
125
+ print(e)