abtExp/source_separation
1
1import gradio as gr2import torch3import torchaudio4from timeit import default_timer as timer5from data_setups import audio_preprocess, resample6import gdown7 8url = 'https://drive.google.com/uc?id=1X5CR18u0I-ZOi_8P0cNptCe5JGk9Ro0C'9output = 'piano.wav'10gdown.download(url, output, quiet=False)11url = 'https://drive.google.com/uc?id=1W-8HwmGR5SiyDbUcGAZYYDKdCIst07__'12output= 'torch_efficientnet_fold2_CNN.pth'13gdown.download(url, output, quiet=False)14device = "cuda" if torch.cuda.is_available() else "cpu"15SAMPLE_RATE = 4410016AUDIO_LEN = 2.9017model = torch.load("torch_efficientnet_fold2_CNN.pth", map_location=torch.device('cpu'))18LABELS = [19 "Cello", "Clarinet", "Flute", "Acoustic Guitar", "Electric Guitar", "Organ", "Piano", "Saxophone", "Trumpet", "Violin", "Voice"20]21example_list = [22 ["piano.wav"]23]24 25 26def predict(audio_path):27 start_time = timer()28 wavform, sample_rate = torchaudio.load(audio_path)29 wav = resample(wavform, sample_rate, SAMPLE_RATE)30 if len(wav) > int(AUDIO_LEN * SAMPLE_RATE):31 wav = wav[:int(AUDIO_LEN * SAMPLE_RATE)]32 else:33 print(f"input length {len(wav)} too small!, need over {int(AUDIO_LEN * SAMPLE_RATE)}")34 return35 img = audio_preprocess(wav, SAMPLE_RATE).unsqueeze(0)36 model.eval()37 with torch.inference_mode():38 pred_probs = torch.softmax(model(img), dim=1)39 pred_labels_and_probs = {LABELS[i]: float(pred_probs[0][i]) for i in range(len(LABELS))}40 pred_time = round(timer() - start_time, 5)41 return pred_labels_and_probs, pred_time42 43demo = gr.Interface(fn=predict,44 inputs=gr.Audio(type="filepath"),45 outputs=[gr.Label(num_top_classes=11, label="Predictions"), 46 gr.Number(label="Prediction time (s)")],47 examples=example_list,48 cache_examples=False49 )50 51demo.launch(debug=False)