{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71916,"databundleVersionId":7876483,"sourceType":"competition"}],"dockerImageVersionId":30665,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nimport torchaudio\nimport torchmetrics\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader, Dataset\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-06T13:07:30.646817Z","iopub.execute_input":"2024-03-06T13:07:30.647247Z","iopub.status.idle":"2024-03-06T13:07:30.654298Z","shell.execute_reply.started":"2024-03-06T13:07:30.647211Z","shell.execute_reply":"2024-03-06T13:07:30.653400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_DIR_PATH = '/kaggle/input/msu-robust-speech-commands-classification/train/train'\nTEST_DIR_PATH = '/kaggle/input/msu-robust-speech-commands-classification/adv_test/adv_test'\n","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:30.655932Z","iopub.execute_input":"2024-03-06T13:07:30.656230Z","iopub.status.idle":"2024-03-06T13:07:30.662932Z","shell.execute_reply.started":"2024-03-06T13:07:30.656206Z","shell.execute_reply":"2024-03-06T13:07:30.662003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 512\nN_WORKERS = 8\nN_CLASSES = 35\nEPOCHS = 50\nLR = 0.01\n","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:30.663984Z","iopub.execute_input":"2024-03-06T13:07:30.664245Z","iopub.status.idle":"2024-03-06T13:07:30.672355Z","shell.execute_reply.started":"2024-03-06T13:07:30.664223Z","shell.execute_reply":"2024-03-06T13:07:30.671576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEVICE = torch.device('cpu')\nif torch.cuda.is_available():\n    DEVICE = torch.device('cuda:0')\nelif torch.backends.mps.is_available():\n    DEVICE = torch.device('mps')\n\nDEVICE\n","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:30.674217Z","iopub.execute_input":"2024-03-06T13:07:30.674494Z","iopub.status.idle":"2024-03-06T13:07:30.683280Z","shell.execute_reply.started":"2024-03-06T13:07:30.674471Z","shell.execute_reply":"2024-03-06T13:07:30.682471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SpeechCommandDataset(Dataset):\n    def __init__(self, dir_path, data, labels=None, dict_label_to_index=None, transform=None):\n        self.dir_path = dir_path\n        self.data = data\n        self.labels = labels\n        self.dict_label_to_index = dict_label_to_index\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        file_name = self.data[idx]\n        waveform = np.load(os.path.join(self.dir_path, file_name))\n        if waveform.shape[1] < 16000:\n            waveform = np.pad(\n                waveform, pad_width=((0, 0), (0, 16000 - waveform.shape[1])),\n                mode='constant',\n                constant_values=0\n            )\n\n        waveform = torch.from_numpy(waveform)\n\n        if self.transform != None:\n            waveform = self.transform(waveform)\n        \n        out_labels = []\n        if self.labels is not None:\n            if self.labels[idx] in self.dict_label_to_index:\n                out_labels = self.dict_label_to_index[self.labels[idx]]\n\n        return waveform, out_labels, int(file_name.split('.')[0])","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:35:29.633836Z","iopub.execute_input":"2024-03-06T13:35:29.634686Z","iopub.status.idle":"2024-03-06T13:35:29.645808Z","shell.execute_reply.started":"2024-03-06T13:35:29.634646Z","shell.execute_reply":"2024-03-06T13:35:29.644672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(\n    os.path.join(TRAIN_DIR_PATH, 'metadata.csv')\n)\ndict_label_to_index = {}\ndict_index_to_label = {}\nfor index, key in enumerate(df_train['label'].unique()):\n    dict_label_to_index[key] = index\n    dict_index_to_label[index] = key\n\ndict_label_to_index","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:30.696059Z","iopub.execute_input":"2024-03-06T13:07:30.696323Z","iopub.status.idle":"2024-03-06T13:07:30.837816Z","shell.execute_reply.started":"2024-03-06T13:07:30.696300Z","shell.execute_reply":"2024-03-06T13:07:30.836890Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_data, df_val_data = train_test_split(\n    df_train,\n    test_size=0.2,\n    random_state=42,\n    shuffle=True\n)\n\ntrain_data = df_train_data.file_name.values\ntrain_labels = df_train_data.label.values\n\nval_data = df_val_data.file_name.values\nval_labels = df_val_data.label.values\n\n","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:30.839082Z","iopub.execute_input":"2024-03-06T13:07:30.839366Z","iopub.status.idle":"2024-03-06T13:07:30.859585Z","shell.execute_reply.started":"2024-03-06T13:07:30.839341Z","shell.execute_reply":"2024-03-06T13:07:30.858889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = torch.nn.Sequential(\n    torchaudio.transforms.MelSpectrogram()\n    # torchaudio.transforms.MFCC()\n)\n\nval_transform = torch.nn.Sequential(\n    torchaudio.transforms.MelSpectrogram(),\n    # torchaudio.transforms.MFCC()\n)\n\ntrain_dataloader = DataLoader(\n    SpeechCommandDataset(\n        dir_path=TRAIN_DIR_PATH,\n        data=train_data,\n        labels=train_labels,\n        dict_label_to_index=dict_label_to_index,\n        # transform=train_transforms\n    ),\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=N_WORKERS\n)\n\nvalid_dataloader = DataLoader(\n    SpeechCommandDataset(\n        dir_path=TRAIN_DIR_PATH,\n        data=val_data,\n        labels=val_labels,\n        dict_label_to_index=dict_label_to_index,\n        # transform=train_transforms\n    ),\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=N_WORKERS\n)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:30.860656Z","iopub.execute_input":"2024-03-06T13:07:30.860947Z","iopub.status.idle":"2024-03-06T13:07:30.877942Z","shell.execute_reply.started":"2024-03-06T13:07:30.860923Z","shell.execute_reply":"2024-03-06T13:07:30.877004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for item in train_dataloader:\n    print(item[0].shape)\n    break","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:30.903354Z","iopub.execute_input":"2024-03-06T13:07:30.903670Z","iopub.status.idle":"2024-03-06T13:07:32.863304Z","shell.execute_reply.started":"2024-03-06T13:07:30.903644Z","shell.execute_reply":"2024-03-06T13:07:32.862290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class M5(nn.Module):\n    def __init__(self, n_input=1, n_output=35, stride=16, n_channel=32):\n        super().__init__()\n        self.conv1 = nn.Conv1d(n_input, n_channel, kernel_size=80, stride=stride)\n        self.bn1 = nn.BatchNorm1d(n_channel)\n        self.pool1 = nn.MaxPool1d(4)\n        self.conv2 = nn.Conv1d(n_channel, n_channel, kernel_size=3)\n        self.bn2 = nn.BatchNorm1d(n_channel)\n        self.pool2 = nn.MaxPool1d(4)\n        self.conv3 = nn.Conv1d(n_channel, 2 * n_channel, kernel_size=3)\n        self.bn3 = nn.BatchNorm1d(2 * n_channel)\n        self.pool3 = nn.MaxPool1d(4)\n        self.conv4 = nn.Conv1d(2 * n_channel, 2 * n_channel, kernel_size=3)\n        self.bn4 = nn.BatchNorm1d(2 * n_channel)\n        self.pool4 = nn.MaxPool1d(4)\n        self.fc1 = nn.Linear(2 * n_channel, n_output)\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = F.relu(self.bn1(x))\n        x = self.pool1(x)\n        x = self.conv2(x)\n        x = F.relu(self.bn2(x))\n        x = self.pool2(x)\n        x = self.conv3(x)\n        x = F.relu(self.bn3(x))\n        x = self.pool3(x)\n        x = self.conv4(x)\n        x = F.relu(self.bn4(x))\n        x = self.pool4(x)\n        x = F.avg_pool1d(x, x.shape[-1])\n        x = x.permute(0, 2, 1)\n        x = self.fc1(x)\n        \n        return F.log_softmax(x, dim=2)","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:32.865525Z","iopub.execute_input":"2024-03-06T13:07:32.865834Z","iopub.status.idle":"2024-03-06T13:07:32.878273Z","shell.execute_reply.started":"2024-03-06T13:07:32.865806Z","shell.execute_reply":"2024-03-06T13:07:32.877397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = M5()\nmodel = model.to(DEVICE)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:32.879395Z","iopub.execute_input":"2024-03-06T13:07:32.879633Z","iopub.status.idle":"2024-03-06T13:07:32.895313Z","shell.execute_reply.started":"2024-03-06T13:07:32.879612Z","shell.execute_reply":"2024-03-06T13:07:32.894616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_image = torch.rand(4, 1, 16000).to(DEVICE)\nmodel = model.to(DEVICE)\nresult = model(input_image)\n\nprint(result.size())\n","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:32.896461Z","iopub.execute_input":"2024-03-06T13:07:32.896737Z","iopub.status.idle":"2024-03-06T13:07:32.906172Z","shell.execute_reply.started":"2024-03-06T13:07:32.896714Z","shell.execute_reply":"2024-03-06T13:07:32.905273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model: nn.Module, train_data: DataLoader, valid_data: DataLoader):\n    optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=0.0001)\n    scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=20, gamma=0.1)\n    criterion = nn.NLLLoss()\n\n    accuracy_train = torchmetrics.classification.Accuracy(task=\"multiclass\", num_classes=N_CLASSES).to(DEVICE)\n    accuracy_val = torchmetrics.classification.Accuracy(task=\"multiclass\", num_classes=N_CLASSES).to(DEVICE)\n\n    for epoch in range(EPOCHS):\n        train_loss = 0.0\n        val_loss = 0.0\n\n        model.train()\n        for x, y, _ in train_data:\n            x = x.to(DEVICE)\n            y = y.to(DEVICE)\n\n            optimizer.zero_grad()\n\n            y_hat = model(x).squeeze()\n            loss = criterion(y_hat, y)\n\n            loss.backward()\n            optimizer.step()\n\n            train_loss += loss.item() * x.size(0)\n            _, preds = torch.max(y_hat, 1)\n\n            accuracy_train(\n                y_hat,\n                y\n            )\n\n        model.eval()\n        for x, y, _ in valid_data:\n            x = x.to(DEVICE)\n            y = y.to(DEVICE)\n\n            y_hat = model(x).squeeze()\n            loss = criterion(y_hat, y)\n\n            val_loss += loss.item() * x.size(0)\n            _, preds = torch.max(y_hat, 1)\n\n            accuracy_val(\n                y_hat,\n                y\n            )\n\n        train_loss = train_loss / len(train_dataloader.dataset)\n        val_loss = val_loss / len(valid_dataloader.dataset)\n\n        scheduler.step()\n\n        print(f\"Epoch {epoch + 1}/{EPOCHS}\")\n        print(f\"Train Loss: {train_loss:.4f}, Train Acc: {accuracy_train.compute():.4f}\")\n        print(f\"Val Loss: {val_loss:.4f}, Val Acc: {accuracy_val.compute():.4f}\")\n        ","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:32.908323Z","iopub.execute_input":"2024-03-06T13:07:32.908596Z","iopub.status.idle":"2024-03-06T13:07:32.920299Z","shell.execute_reply.started":"2024-03-06T13:07:32.908573Z","shell.execute_reply":"2024-03-06T13:07:32.919382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_model(\n    model=model,\n    train_data=train_dataloader,\n    valid_data=valid_dataloader\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:07:32.921550Z","iopub.execute_input":"2024-03-06T13:07:32.922073Z","iopub.status.idle":"2024-03-06T13:25:36.786212Z","shell.execute_reply.started":"2024-03-06T13:07:32.922041Z","shell.execute_reply":"2024-03-06T13:25:36.785052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test = pd.read_csv(\n    os.path.join(TEST_DIR_PATH, 'metadata.csv')\n)\ntest_dataloader = DataLoader(\n    SpeechCommandDataset(\n        dir_path=TEST_DIR_PATH,\n        data=df_test.file_name.values,\n        labels=None,\n        dict_label_to_index=dict_label_to_index,\n        # transform=train_transforms\n    ),\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=N_WORKERS\n)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:35:48.761673Z","iopub.execute_input":"2024-03-06T13:35:48.762681Z","iopub.status.idle":"2024-03-06T13:35:48.783601Z","shell.execute_reply.started":"2024-03-06T13:35:48.762636Z","shell.execute_reply":"2024-03-06T13:35:48.782751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ENSEMBLE PREDICTIONS AND SUBMIT\nresults = {\n    'id': [],\n    'label': []\n}\n\nmodel.eval()\nfor x, y, ids in test_dataloader:\n    x = x.float().to(DEVICE)\n    with torch.no_grad():\n        y_hat = model(x).squeeze()\n        _, preds = torch.max(y_hat, 1)\n        for i in range(len(preds)):\n            results[\"id\"].append(ids[i].item())\n            results[\"label\"].append(dict_index_to_label[int(preds[i].item())])\n        \n\npd.DataFrame(results).to_csv(\n    'submission.csv',\n    columns=['id', 'label'],\n    index=False\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:35:51.720300Z","iopub.execute_input":"2024-03-06T13:35:51.720887Z","iopub.status.idle":"2024-03-06T13:35:58.589256Z","shell.execute_reply.started":"2024-03-06T13:35:51.720829Z","shell.execute_reply":"2024-03-06T13:35:58.588156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import FileLink\n\nFileLink(r'submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-03-06T13:35:58.591338Z","iopub.execute_input":"2024-03-06T13:35:58.591733Z","iopub.status.idle":"2024-03-06T13:35:58.598737Z","shell.execute_reply.started":"2024-03-06T13:35:58.591702Z","shell.execute_reply":"2024-03-06T13:35:58.597754Z"},"trusted":true},"execution_count":null,"outputs":[]}]}