{"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":"import torch\nimport torchaudio\nimport pandas as pd\nimport numpy as np\nimport torch.nn as nn\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import classification_report, roc_auc_score\nfrom sklearn.manifold import TSNE\nfrom IPython.display import Audio","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:51:05.082967Z","iopub.execute_input":"2022-02-24T14:51:05.083446Z","iopub.status.idle":"2022-02-24T14:51:05.088763Z","shell.execute_reply.started":"2022-02-24T14:51:05.083409Z","shell.execute_reply":"2022-02-24T14:51:05.088094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_dict = {\n    0 : 'speech',\n    1 : 'music',\n    2 : 'noise'    \n}","metadata":{"execution":{"iopub.status.busy":"2022-02-24T15:55:49.408594Z","iopub.execute_input":"2022-02-24T15:55:49.409301Z","iopub.status.idle":"2022-02-24T15:55:49.412509Z","shell.execute_reply.started":"2022-02-24T15:55:49.409263Z","shell.execute_reply":"2022-02-24T15:55:49.411857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Посмотрим на датасет и послушаем примеры","metadata":{}},{"cell_type":"code","source":"data = pd.read_csv('/kaggle/input/silero-audio-classifier/train.csv', index_col=0)\ndata.shape","metadata":{"execution":{"iopub.status.busy":"2022-02-24T16:01:07.886199Z","iopub.execute_input":"2022-02-24T16:01:07.886462Z","iopub.status.idle":"2022-02-24T16:01:08.124921Z","shell.execute_reply.started":"2022-02-24T16:01:07.886433Z","shell.execute_reply":"2022-02-24T16:01:08.124205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"music = data[data['label'] == 'music']\nspeech = data[data['label'] == 'speech']\nnoise = data[data['label'] == 'noise']\n\nmusic = music.sample(1000)\nspeech = speech.sample(1000)\nnoise = noise.sample(1000)\n\ndf = pd.concat([music, speech, noise])","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:02:43.262754Z","iopub.execute_input":"2022-02-24T14:02:43.263192Z","iopub.status.idle":"2022-02-24T14:02:43.409392Z","shell.execute_reply.started":"2022-02-24T14:02:43.263154Z","shell.execute_reply":"2022-02-24T14:02:43.408685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('music')\nAudio(data_folder+music.iloc[10]['wav_path'])","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:53:11.378631Z","iopub.execute_input":"2022-02-24T14:53:11.378959Z","iopub.status.idle":"2022-02-24T14:53:11.397425Z","shell.execute_reply.started":"2022-02-24T14:53:11.378920Z","shell.execute_reply":"2022-02-24T14:53:11.396712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('speech')\nAudio(data_folder+speech.iloc[10]['wav_path'])","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:53:05.031872Z","iopub.execute_input":"2022-02-24T14:53:05.032481Z","iopub.status.idle":"2022-02-24T14:53:05.046402Z","shell.execute_reply.started":"2022-02-24T14:53:05.032447Z","shell.execute_reply":"2022-02-24T14:53:05.045612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('noise')\nAudio(data_folder+noise.iloc[10]['wav_path'])","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:52:51.307952Z","iopub.execute_input":"2022-02-24T14:52:51.308534Z","iopub.status.idle":"2022-02-24T14:52:51.322138Z","shell.execute_reply.started":"2022-02-24T14:52:51.308498Z","shell.execute_reply":"2022-02-24T14:52:51.321422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.sample(10)","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:02:49.240809Z","iopub.execute_input":"2022-02-24T14:02:49.241296Z","iopub.status.idle":"2022-02-24T14:02:49.254734Z","shell.execute_reply.started":"2022-02-24T14:02:49.241258Z","shell.execute_reply":"2022-02-24T14:02:49.254046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_folder = '/kaggle/input/silero-audio-classifier/train/'","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:07:24.676523Z","iopub.execute_input":"2022-02-24T14:07:24.677230Z","iopub.status.idle":"2022-02-24T14:07:24.680998Z","shell.execute_reply.started":"2022-02-24T14:07:24.677191Z","shell.execute_reply":"2022-02-24T14:07:24.679920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AudioDataset(Dataset):\n    def __init__(self, data, target_sr=8000):\n        super().__init__()\n        wavs = data['wav_path'].to_list()\n        self.audios = []\n        for wav in tqdm(wavs):\n            audio, sr = torchaudio.load(data_folder+wav)\n            audio = torchaudio.transforms.Resample(sr, target_sr).forward(audio)\n            self.audios.append(audio.squeeze(0))\n        self.targets = data['target'].to_list()\n    \n    def __getitem__(self, idx):\n        return {\n            \"audio\": self.audios[idx],\n            \"target\": self.targets[idx]\n        }\n    \n    def __len__(self):\n        return len(self.audios)","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:08:46.697125Z","iopub.execute_input":"2022-02-24T14:08:46.697979Z","iopub.status.idle":"2022-02-24T14:08:46.704812Z","shell.execute_reply.started":"2022-02-24T14:08:46.697936Z","shell.execute_reply":"2022-02-24T14:08:46.704010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, val_df = train_test_split(df, test_size=0.3, shuffle=True, stratify=df['target'])","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:10:55.956119Z","iopub.execute_input":"2022-02-24T14:10:55.956846Z","iopub.status.idle":"2022-02-24T14:10:55.966934Z","shell.execute_reply.started":"2022-02-24T14:10:55.956805Z","shell.execute_reply":"2022-02-24T14:10:55.965984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = AudioDataset(train_df)\nval_dataset = AudioDataset(val_df)","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:11:16.909108Z","iopub.execute_input":"2022-02-24T14:11:16.909362Z","iopub.status.idle":"2022-02-24T14:11:23.087361Z","shell.execute_reply.started":"2022-02-24T14:11:16.909333Z","shell.execute_reply":"2022-02-24T14:11:23.086588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dl = DataLoader(train_dataset, batch_size=8, shuffle=True)\nval_dl = DataLoader(val_dataset, batch_size=8)","metadata":{"execution":{"iopub.status.busy":"2022-02-24T14:11:35.328051Z","iopub.execute_input":"2022-02-24T14:11:35.328305Z","iopub.status.idle":"2022-02-24T14:11:35.332255Z","shell.execute_reply.started":"2022-02-24T14:11:35.328276Z","shell.execute_reply":"2022-02-24T14:11:35.331607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Обучить Fully-Connected модель","metadata":{}},{"cell_type":"code","source":"class FullyConnectedClassifier(nn.Module):\n    def __init__(self, hidden_size=256):\n        super().__init__()\n        self.features = torchaudio.transforms.MelSpectrogram(8000)\n        self.fc1 = nn.Linear(15488, hidden_size)\n        self.relu1 = nn.ReLU()\n        self.dropout1 = nn.Dropout(0.2)\n        self.fc2 = nn.Linear(hidden_size, hidden_size)\n        self.relu2 = nn.ReLU()\n        self.dropout2 = nn.Dropout(0.2)\n        self.fc3 = nn.Linear(hidden_size, 3)\n        \n    def forward(self, batch):\n        x = torch.flatten(self.features(batch), start_dim=1)\n        x = self.dropout1(self.relu1(self.fc1(x)))\n        x = self.dropout2(self.relu2(self.fc2(x)))\n        x = self.fc3(x)\n        return torch.softmax(x, dim=1)","metadata":{"execution":{"iopub.status.busy":"2022-02-24T16:04:48.791999Z","iopub.execute_input":"2022-02-24T16:04:48.792540Z","iopub.status.idle":"2022-02-24T16:04:48.801008Z","shell.execute_reply.started":"2022-02-24T16:04:48.792501Z","shell.execute_reply":"2022-02-24T16:04:48.800210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = FullyConnectedClassifier()\nmodel.to('cuda')\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.SGD(model.parameters(), lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2022-02-24T16:04:49.014958Z","iopub.execute_input":"2022-02-24T16:04:49.015200Z","iopub.status.idle":"2022-02-24T16:04:49.054196Z","shell.execute_reply.started":"2022-02-24T16:04:49.015174Z","shell.execute_reply":"2022-02-24T16:04:49.053580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(25):  \n\n    running_loss = 0.0\n    for i, data in tqdm(enumerate(train_dl, 0), total=len(train_dl)):\n        audios = data['audio'].to('cuda')\n        labels = data['target']\n\n        # zero the parameter gradients\n        optimizer.zero_grad()\n\n        # forward + backward + optimize\n        outputs = model(audios).cpu()\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        # print statistics\n        running_loss += loss.item()\n        if i % 100 == 99:   \n            tqdm.write(f'[{epoch + 1}, {i + 1:5d}] loss: {running_loss / 100:.3f}')\n            running_loss = 0.0\n            \nprint('Finished Training')","metadata":{"execution":{"iopub.status.busy":"2022-02-24T16:04:49.831378Z","iopub.execute_input":"2022-02-24T16:04:49.832031Z","iopub.status.idle":"2022-02-24T16:05:03.550886Z","shell.execute_reply.started":"2022-02-24T16:04:49.831991Z","shell.execute_reply":"2022-02-24T16:05:03.549998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\n\nfor i, data in tqdm(enumerate(val_dl, 0), total=len(val_dl)):\n    audios = data['audio'].to('cuda')\n    labels = data['target']\n\n    with torch.no_grad():\n        outputs = model(audios).cpu()\n\n    predictions += torch.argmax(outputs, axis=-1).tolist()\n    \nprint(classification_report(predictions, val_dataset.targets, target_names=label_dict.values()))","metadata":{"execution":{"iopub.status.busy":"2022-02-24T16:05:08.940192Z","iopub.execute_input":"2022-02-24T16:05:08.940660Z","iopub.status.idle":"2022-02-24T16:05:09.113528Z","shell.execute_reply.started":"2022-02-24T16:05:08.940609Z","shell.execute_reply":"2022-02-24T16:05:09.112832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Обучим CNN модель","metadata":{}},{"cell_type":"code","source":"class CNNClassifier(nn.Module):\n    def __init__(self, hidden_size=256):\n        super().__init__()\n        self.features = torchaudio.transforms.MelSpectrogram(8000)\n        self.conv1 = nn.Conv1d(128, hidden_size, 3)\n        self.pool1 = nn.MaxPool1d(3)\n        self.conv2 = nn.Conv1d(hidden_size, hidden_size*2, 3)\n        self.pool2 = nn.MaxPool1d(3)\n        self.conv3 = nn.Conv1d(hidden_size*2, hidden_size//4, 5)\n        self.pool3 = nn.MaxPool1d(2)\n        self.fc = nn.Linear(256, 3)\n        \n    def forward(self, batch):\n        x = self.pool1(self.conv1(self.features(batch)))\n        x = self.pool2(self.conv2(x))\n        x = self.pool3(self.conv3(x))\n        emb = torch.flatten(x, start_dim=1)\n        x = self.fc(emb)\n        return torch.softmax(x, dim=1), emb","metadata":{"execution":{"iopub.status.busy":"2022-02-24T16:05:54.115969Z","iopub.execute_input":"2022-02-24T16:05:54.116228Z","iopub.status.idle":"2022-02-24T16:05:54.124752Z","shell.execute_reply.started":"2022-02-24T16:05:54.116201Z","shell.execute_reply":"2022-02-24T16:05:54.123993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CNNClassifier()\nmodel.to('cuda')\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.SGD(model.parameters(), lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2022-02-24T16:05:54.531421Z","iopub.execute_input":"2022-02-24T16:05:54.531856Z","iopub.status.idle":"2022-02-24T16:05:54.546801Z","shell.execute_reply.started":"2022-02-24T16:05:54.531822Z","shell.execute_reply":"2022-02-24T16:05:54.546140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(25):  \n\n    running_loss = 0.0\n    for i, data in tqdm(enumerate(train_dl, 0), total=len(train_dl)):\n        audios = data['audio'].to('cuda')\n        labels = data['target']\n\n        # zero the parameter gradients\n        optimizer.zero_grad()\n\n        # forward + backward + optimize\n        outputs, _ = model(audios)\n        outputs = outputs.cpu()\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        # print statistics\n        running_loss += loss.item()\n        if i % 100 == 99:   \n            tqdm.write(f'[{epoch + 1}, {i + 1:5d}] loss: {running_loss / 100:.3f}')\n            running_loss = 0.0\n            \nprint('Finished Training')","metadata":{"execution":{"iopub.status.busy":"2022-02-24T16:06:00.398041Z","iopub.execute_input":"2022-02-24T16:06:00.398310Z","iopub.status.idle":"2022-02-24T16:06:19.897485Z","shell.execute_reply.started":"2022-02-24T16:06:00.398280Z","shell.execute_reply":"2022-02-24T16:06:19.896795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\nembs = []\n\nfor i, data in tqdm(enumerate(val_dl, 0), total=len(val_dl)):\n    audios = data['audio'].to('cuda')\n    labels = data['target']\n\n    with torch.no_grad():\n        outputs, emb = model(audios)\n        outputs = outputs.cpu()\n    \n    embs += emb.tolist()\n    predictions += torch.argmax(outputs, axis=-1).tolist()\n    \nprint(classification_report(predictions, val_dataset.targets, target_names=label_dict.values()))","metadata":{"execution":{"iopub.status.busy":"2022-02-24T16:06:22.723933Z","iopub.execute_input":"2022-02-24T16:06:22.724200Z","iopub.status.idle":"2022-02-24T16:06:22.955626Z","shell.execute_reply.started":"2022-02-24T16:06:22.724171Z","shell.execute_reply":"2022-02-24T16:06:22.954851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Посмотрим на пространство эмбеддингов модели","metadata":{}},{"cell_type":"code","source":"tsne = TSNE(n_components=2)\nspace = tsne.fit_transform(np.array(embs))\nx, y = zip(*space)\nscatter = plt.scatter(x, y, c=val_dataset.targets)\nplt.legend(handles=scatter.legend_elements()[0], \n           title=\"Classes\", labels=label_dict.values())\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-02-24T16:06:55.502617Z","iopub.execute_input":"2022-02-24T16:06:55.503173Z","iopub.status.idle":"2022-02-24T16:07:02.246075Z","shell.execute_reply.started":"2022-02-24T16:06:55.503133Z","shell.execute_reply":"2022-02-24T16:07:02.245434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}