{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Preparing","metadata":{"id":"WWALhgZAH2wQ"}},{"cell_type":"markdown","source":"## Imports and constants","metadata":{"id":"2m8TfCquIIG5"}},{"cell_type":"code","source":"!pip install transformers -qqq\n!pip install opendatasets --upgrade -qqq","metadata":{"execution":{"iopub.status.busy":"2022-02-28T08:50:16.153079Z","iopub.execute_input":"2022-02-28T08:50:16.153425Z","iopub.status.idle":"2022-02-28T08:50:37.503244Z","shell.execute_reply.started":"2022-02-28T08:50:16.153342Z","shell.execute_reply":"2022-02-28T08:50:37.502112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport csv\nimport librosa\nimport numpy as np\nfrom tqdm import tqdm","metadata":{"id":"AAABVGakCI-X","outputId":"a4e4e61b-e00e-40ec-c4b0-81c60a2492f7","execution":{"iopub.status.busy":"2022-02-28T08:50:54.965610Z","iopub.execute_input":"2022-02-28T08:50:54.965898Z","iopub.status.idle":"2022-02-28T08:50:57.511392Z","shell.execute_reply.started":"2022-02-28T08:50:54.965869Z","shell.execute_reply":"2022-02-28T08:50:57.510241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Download dataset from kaggle","metadata":{"id":"Bd3q8tOrIMvD"}},{"cell_type":"code","source":"import opendatasets as od\ndataset_url = 'https://www.kaggle.com/c/classification-of-short-noisy-audio-speech/data'\nod.download(dataset_url)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Evaluating","metadata":{"id":"HE_eQDAaIROi"}},{"cell_type":"markdown","source":"## Model","metadata":{"id":"weyOqJC-KoRV"}},{"cell_type":"code","source":"USE_CUDA = True\nSAMPLING_RATE = 16000 # typical for speech recognition sampling rate (most of the models were trained on it)\n\nmodel_name_1 = \"hubert/checkpoint-X\"\nLOGITS_1_WEIGHT = 0.5 # for ansembling\nUSE_FP16_MODEL_1 = False # if we want to use in half less memory\n\nmodel_name_2 = \"wav2vec/checkpoint-X\"\nLOGITS_2_WEIGHT = 0.5 # for ansembling\nUSE_FP16_MODEL_2 = False # if we want to use in half less memory\n\n# models should be trained on the same labels, for example, on these\nlabels = ['down', 'go', 'left', 'no', 'off', 'on', 'right', 'stop', 'up', 'yes']","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoModelForAudioClassification, AutoFeatureExtractor\n\nfeature_extractor_1 = AutoFeatureExtractor.from_pretrained(model_name_1)\nmodel_1 = AutoModelForAudioClassification.from_pretrained(model_name_1)\nif USE_CUDA:\n    model_1 = model_1.cuda()\nif USE_FP16_MODEL_1:\n    model_1 = model_1.half()\n\nfeature_extractor_2 = AutoFeatureExtractor.from_pretrained(model_name_2)\nmodel_2 = AutoModelForAudioClassification.from_pretrained(model_name_2)\nif USE_CUDA:\n    model_2 = model_2.cuda()\nif USE_FP16_MODEL_2:\n    model_2 = model_2.half()","metadata":{"id":"FPkGuO6zKuG4"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference functions","metadata":{"id":"F0nFT8gBcL6X"}},{"cell_type":"code","source":"def get_pred(samples, print_scores):\n    inputs_1 = feature_extractor_1(samples, sampling_rate=feature_extractor_1.sampling_rate, return_tensors=\"pt\",\n                               max_length=int(feature_extractor_1.sampling_rate * 1.0), \n                               truncation=True)\n    inputs_2 = feature_extractor_2(samples, sampling_rate=feature_extractor_2.sampling_rate, return_tensors=\"pt\",\n                               max_length=int(feature_extractor_2.sampling_rate * 1.0), \n                               truncation=True)\n    if USE_CUDA:\n        inputs_1['input_values'] = inputs_1['input_values'].cuda()\n    if USE_FP16_MODEL_1:\n        inputs_1['input_values'] = inputs_1['input_values'].half()\n    if USE_CUDA:\n        inputs_2['input_values'] = inputs_2['input_values'].cuda()\n    if USE_FP16_MODEL_2:\n        inputs_2['input_values'] = inputs_2['input_values'].half()\n        \n    with torch.no_grad():\n        logits_1 = model_1(**inputs_1).logits\n        logits_2 = model_2(**inputs_2).logits\n        logits = logits_1 * LOGITS_1_WEIGHT + logits_2 * LOGITS_2_WEIGHT\n    logits = list(logits[0].cpu())\n    logits = np.array(logits)\n    predicted_class_ids = np.argmax(logits, axis=-1)\n    predicted_label = labels[predicted_class_ids]\n    return predicted_label","metadata":{"id":"yZ-xQOUDMryK"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_preds(filenames, data_path):\n    for filename in tqdm(filenames):\n        word, sr = librosa.load(os.path.join(data_path, filename), sr=SAMPLING_RATE)\n        res = get_pred(word)\n        yield res","metadata":{"id":"CzvWGkrIblXy"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_submit(data_path, csv_save_path):\n    filenames = os.listdir(data_path)\n    with open(csv_save_path, mode='w') as f:\n        writer = csv.writer(f)\n        writer.writerow(['id', 'category'])\n        for filename, pred in zip(filenames, get_preds(filenames, data_path=data_path)):\n            writer.writerow([filename.split('.')[0], pred])","metadata":{"id":"d0oraxUUGtUn"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submit","metadata":{"id":"FDS0WE-7cQ7t"}},{"cell_type":"code","source":"create_submit(data_path='classification-of-short-noisy-audio-speech/hackaton_ds/test', csv_save_path='submission.csv')","metadata":{"id":"YloVGxwqjt4M","outputId":"bbc97aa2-4e1e-4895-85d9-1e62ac2ede23"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!kaggle competitions submit -c classification-of-short-noisy-audio-speech -f submission.csv -m \"my_submission\"","metadata":{"id":"G66UFrs2ujBV","outputId":"82bfbce0-a551-4c53-cb3e-3ef23d45c98a"},"execution_count":null,"outputs":[]}]}