{"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":"code","source":"!pip install pytorch_metric_learning\n!pip install -U torchaudio","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:50:45.497343Z","iopub.execute_input":"2022-02-27T14:50:45.497653Z","iopub.status.idle":"2022-02-27T14:52:07.577820Z","shell.execute_reply.started":"2022-02-27T14:50:45.497573Z","shell.execute_reply":"2022-02-27T14:52:07.576873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport pandas as pd\nimport numpy as np\nfrom tqdm import tqdm\n\nfrom sklearn.metrics import accuracy_score\nfrom sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:52:07.580344Z","iopub.execute_input":"2022-02-27T14:52:07.580657Z","iopub.status.idle":"2022-02-27T14:52:08.460290Z","shell.execute_reply.started":"2022-02-27T14:52:07.580615Z","shell.execute_reply":"2022-02-27T14:52:08.459574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.tensorboard import SummaryWriter\n\nimport torchaudio\nfrom pytorch_metric_learning import losses","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:52:08.461654Z","iopub.execute_input":"2022-02-27T14:52:08.461920Z","iopub.status.idle":"2022-02-27T14:52:09.385603Z","shell.execute_reply.started":"2022-02-27T14:52:08.461886Z","shell.execute_reply":"2022-02-27T14:52:09.384766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bundle = torchaudio.pipelines.WAV2VEC2_BASE","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:52:09.388257Z","iopub.execute_input":"2022-02-27T14:52:09.388691Z","iopub.status.idle":"2022-02-27T14:52:09.393641Z","shell.execute_reply.started":"2022-02-27T14:52:09.388642Z","shell.execute_reply":"2022-02-27T14:52:09.392968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wav2vec2 = bundle.get_model()","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:52:09.396953Z","iopub.execute_input":"2022-02-27T14:52:09.397137Z","iopub.status.idle":"2022-02-27T14:52:30.702485Z","shell.execute_reply.started":"2022-02-27T14:52:09.397114Z","shell.execute_reply":"2022-02-27T14:52:30.701726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wav2vec2.encoder.transformer.layers = wav2vec2.encoder.transformer.layers[:-4]","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:52:30.703967Z","iopub.execute_input":"2022-02-27T14:52:30.704290Z","iopub.status.idle":"2022-02-27T14:52:30.712540Z","shell.execute_reply.started":"2022-02-27T14:52:30.704248Z","shell.execute_reply":"2022-02-27T14:52:30.711897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"root_dir = '/kaggle/input/classification-of-short-noisy-audio-speech/hackaton_ds/train/'\nbatch_size = 64\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\ndevice","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:52:30.713873Z","iopub.execute_input":"2022-02-27T14:52:30.714315Z","iopub.status.idle":"2022-02-27T14:52:30.766769Z","shell.execute_reply.started":"2022-02-27T14:52:30.714280Z","shell.execute_reply":"2022-02-27T14:52:30.766149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CommandDataset(Dataset):\n\n    def __init__(self, meta, root_dir, sample_rate, labelmap):\n        self.meta = meta\n        self.root_dir = root_dir\n        self.sample_rate = sample_rate\n        self.labelmap = labelmap\n\n    def __len__(self):\n        return len(self.meta)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        file_name = self.meta['path'].iloc[idx]\n        waveform, sample_rate = torchaudio.load(file_name)\n        \n        if random.randint(0, 1):\n        \n            effects = [\n              [\"speed\", str(np.random.random() + 0.5)],  # reduce the speed\n                                 # This only changes sample rate, so it is necessary to\n                                 # add `rate` effect with original sample rate after this.\n              [\"rate\", f\"{sample_rate}\"],\n            ]\n\n            # Apply effects\n            waveform, sample_rate = torchaudio.sox_effects.apply_effects_tensor(\n                waveform, sample_rate, effects)\n            \n        if random.randint(0, 1):\n        \n            effects = [\n              [\"rate\", f\"{sample_rate}\"],\n              [\"reverb\", \"-w\"],  # Reverbration gives some dramatic feeling\n            ]\n\n            # Apply effects\n            waveform, sample_rate = torchaudio.sox_effects.apply_effects_tensor(\n                waveform, sample_rate, effects)\n        \n        waveform = torchaudio.functional.resample(waveform, sample_rate, bundle.sample_rate)#[:, :10**5]\n        waveform = torch.nn.functional.pad(waveform, (16000-waveform.shape[1], 0))[0]\n            \n        label = self.meta['label'].iloc[idx]\n\n        return waveform, self.labelmap[label]","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:52:30.768834Z","iopub.execute_input":"2022-02-27T14:52:30.769228Z","iopub.status.idle":"2022-02-27T14:52:30.780463Z","shell.execute_reply.started":"2022-02-27T14:52:30.769191Z","shell.execute_reply":"2022-02-27T14:52:30.779672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = {\n    'yes': 0, \n    'no': 1, \n    'up': 2, \n    'down': 3, \n    'left': 4, \n    'right': 5, \n    'on': 6, \n    'off': 7, \n    'stop': 8, \n    'go': 9, \n}","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:52:30.781715Z","iopub.execute_input":"2022-02-27T14:52:30.782370Z","iopub.status.idle":"2022-02-27T14:52:30.790784Z","shell.execute_reply.started":"2022-02-27T14:52:30.782332Z","shell.execute_reply":"2022-02-27T14:52:30.790141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = pd.DataFrame([\n    {'label': i[0].split('/')[-1], 'path': i[0] + '/' + j}\n    for i in os.walk(root_dir)\n    for j in i[2]\n])","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:52:30.792129Z","iopub.execute_input":"2022-02-27T14:52:30.792571Z","iopub.status.idle":"2022-02-27T14:53:40.797620Z","shell.execute_reply.started":"2022-02-27T14:52:30.792531Z","shell.execute_reply":"2022-02-27T14:53:40.796866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data.label.value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:53:40.799105Z","iopub.execute_input":"2022-02-27T14:53:40.799347Z","iopub.status.idle":"2022-02-27T14:53:40.835341Z","shell.execute_reply.started":"2022-02-27T14:53:40.799314Z","shell.execute_reply":"2022-02-27T14:53:40.834680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train, val, _, _ = train_test_split(data, data['label'], test_size=0.1)","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:53:40.836866Z","iopub.execute_input":"2022-02-27T14:53:40.837121Z","iopub.status.idle":"2022-02-27T14:53:40.860162Z","shell.execute_reply.started":"2022-02-27T14:53:40.837082Z","shell.execute_reply":"2022-02-27T14:53:40.859582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CommandDataset(meta=train, root_dir=root_dir, sample_rate=bundle.sample_rate, labelmap=labels)\ntrain_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=12)\n\nval_dataset = CommandDataset(meta=val, root_dir=root_dir, sample_rate=bundle.sample_rate, labelmap=labels)\nval_dataloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=12)","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:53:40.861425Z","iopub.execute_input":"2022-02-27T14:53:40.861885Z","iopub.status.idle":"2022-02-27T14:53:40.870711Z","shell.execute_reply.started":"2022-02-27T14:53:40.861847Z","shell.execute_reply":"2022-02-27T14:53:40.869977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CommandClassifier(nn.Module):\n    def __init__(self, feature_extractor):\n        super(CommandClassifier, self).__init__()\n        self.feature_extractor = feature_extractor\n        self.linear = nn.Linear(768, len(labels))\n        \n    def forward(self, X):\n        features = self.get_embeddings(X)\n        logits = self.linear(features)\n        return logits\n    \n    def get_embeddings(self, X):\n        embeddings = self.feature_extractor(X)[0].mean(axis=1)\n        return nn.functional.normalize(embeddings)","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:53:40.875395Z","iopub.execute_input":"2022-02-27T14:53:40.875916Z","iopub.status.idle":"2022-02-27T14:53:40.882611Z","shell.execute_reply.started":"2022-02-27T14:53:40.875880Z","shell.execute_reply":"2022-02-27T14:53:40.881731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CommandClassifier(wav2vec2)\n#model.load_state_dict(torch.load('model.pth'))\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:53:40.884020Z","iopub.execute_input":"2022-02-27T14:53:40.884547Z","iopub.status.idle":"2022-02-27T14:53:43.667307Z","shell.execute_reply.started":"2022-02-27T14:53:40.884509Z","shell.execute_reply":"2022-02-27T14:53:43.666608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCHS = 10\nlr = 0.00001\n\noptimizer = optim.AdamW(model.parameters(), lr)\n\ncriterion = losses.ArcFaceLoss(len(labels), 768).to(device)","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:53:43.668508Z","iopub.execute_input":"2022-02-27T14:53:43.669043Z","iopub.status.idle":"2022-02-27T14:53:43.676368Z","shell.execute_reply.started":"2022-02-27T14:53:43.669000Z","shell.execute_reply":"2022-02-27T14:53:43.675413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(EPOCHS):\n    \n    model.train()     \n        \n    train_loss = []\n    for batch, targets in tqdm(train_dataloader, desc=f\"Epoch: {epoch}\"):\n        optimizer.zero_grad()\n        batch = batch.to(device)\n        targets = targets.to(device)\n        \n        predictions = model.get_embeddings(batch)\n\n        loss = criterion(predictions, targets) \n        loss.backward()\n        \n        optimizer.step()\n\n        train_loss.append(loss.item())\n        \n    print('Training loss:', np.mean(train_loss))\n    \n    model.eval()\n        \n    val_loss = []\n    for batch, targets in tqdm(val_dataloader, desc=f\"Epoch: {epoch}\"):\n        \n        with torch.no_grad():\n        \n            batch = batch.to(device)\n            targets = targets.to(device)\n            \n            predictions = model.get_embeddings(batch)\n\n            loss = criterion(predictions, targets) \n\n            val_loss.append(loss.item())\n        \n    print('Val loss:', np.mean(val_loss))","metadata":{"execution":{"iopub.status.busy":"2022-02-27T14:53:43.677501Z","iopub.execute_input":"2022-02-27T14:53:43.677842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCHS = 20\nlr = 0.00001\n\noptimizer = optim.AdamW(model.parameters(), lr)\n\ncriterion = nn.CrossEntropyLoss()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), 'model.pth')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"writer = SummaryWriter()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(EPOCHS):\n    \n    model.train()\n        \n    train_loss = []\n    train_predictions = []\n    train_targets = []\n    for batch, targets in tqdm(train_dataloader, desc=f\"Epoch: {epoch}\"):\n        optimizer.zero_grad()\n        \n        batch = batch.to(device)\n        targets = targets.to(device)\n        \n        predictions = model(batch)\n        \n        loss = criterion(predictions, targets) \n        loss.backward()\n        optimizer.step()\n\n        train_loss.append(loss.item())\n        \n        train_predictions.extend(predictions.cpu().detach().numpy().argmax(axis=1))\n        train_targets.extend(targets.cpu().detach().numpy())\n        \n    \n    train_loss = np.mean(train_loss)\n    train_accuracy = accuracy_score(train_targets, train_predictions)\n    \n    print('Training loss:', train_loss, end=' ')\n    print('Train accuracy:', train_accuracy)\n    \n    model.eval()\n        \n    val_predictions = []\n    val_targets = []\n    val_loss = []\n    for batch, targets in tqdm(val_dataloader, desc=f\"Epoch: {epoch}\"):\n        \n        with torch.no_grad():\n        \n            batch = batch.to(device)\n            targets = targets.to(device)\n            predictions = model(batch)\n            loss = criterion(predictions, targets) \n            \n            val_loss.append(loss.item())\n\n            val_predictions.extend(predictions.cpu().numpy().argmax(axis=1))\n            val_targets.extend(targets.cpu().numpy())\n        \n    val_loss = np.mean(val_loss)\n    val_accuracy = accuracy_score(val_targets, val_predictions)\n    \n    print('Val loss:', val_loss, end=' ')\n    print('Val accuracy:', val_accuracy, end=' ')\n    \n    torch.save(model.state_dict(), 'model.pth')\n    \n    writer.add_scalars(\n        'Accuracy', \n        {'train': train_accuracy, 'val': val_accuracy,}, \n        epoch\n    )\n    writer.add_scalars(\n        'Loss', \n        {'train': train_loss, 'val': val_loss,}, \n        epoch\n    )","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dir = '/kaggle/input/classification-of-short-noisy-audio-speech/hackaton_ds/test/'","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\n\npred = []\nfor i in tqdm(os.listdir(test_dir)):\n    \n    waveform, sample_rate = torchaudio.load(f'{test_dir}/{i}')\n    waveform = torchaudio.functional.resample(waveform, sample_rate, bundle.sample_rate)#[:, :10**5]\n    waveform = torch.nn.functional.pad(waveform, (16000-waveform.shape[1], 0))[0]\n    \n    with torch.no_grad():\n        predictions = model(waveform.unsqueeze(0).to(device))[0].cpu()\n    \n    text_lab = list(labels.keys())[predictions.argmax()]\n    \n    pred.append({'id': i.replace('.wav', ''), 'category': text_lab})\npred = pd.DataFrame(pred)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred.to_csv('submission.csv', index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}