{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":53379,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":44787}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport imageio.v3 as imageio\n\nfrom tqdm.notebook import tqdm\nimport librosa\nimport cv2\nimport pickle\nimport lzma\nimport os\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nimport typing\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\n\n\npd.set_option('display.max_rows', 1000)  # Set the maximum number of rows to display\n","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:23.830788Z","iopub.execute_input":"2024-05-22T05:00:23.831283Z","iopub.status.idle":"2024-05-22T05:00:23.838007Z","shell.execute_reply.started":"2024-05-22T05:00:23.831253Z","shell.execute_reply":"2024-05-22T05:00:23.836901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.__version__","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:24.91984Z","iopub.execute_input":"2024-05-22T05:00:24.920199Z","iopub.status.idle":"2024-05-22T05:00:24.927275Z","shell.execute_reply.started":"2024-05-22T05:00:24.92017Z","shell.execute_reply":"2024-05-22T05:00:24.926161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_metadata_df = pd.read_csv(\n        '/kaggle/input/birdclef-2024/train_metadata.csv',\n        dtype={\n            'secondary_labels': 'string',\n            'primary_label': 'category',\n        },\n    )\n\n# Convert secondary_labels to iterable tuple\ndef parse_secondary_labels(s):\n    s = s.strip(\"[']\")\n    s = s.split(\"', '\")\n    return tuple([e for e in s if len(e) > 0])\n\n# train_metadata_df['secondary_labels'] = train_metadata_df['secondary_labels'].apply(parse_secondary_labels)\n# train_metadata_df['type'] = train_metadata_df['type'].apply(parse_secondary_labels)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:25.805567Z","iopub.execute_input":"2024-05-22T05:00:25.805917Z","iopub.status.idle":"2024-05-22T05:00:26.011347Z","shell.execute_reply.started":"2024-05-22T05:00:25.805888Z","shell.execute_reply":"2024-05-22T05:00:26.010301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_metadata = train_metadata_df[[\"primary_label\",\"filename\"]]\ntrain_metadata","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:26.451069Z","iopub.execute_input":"2024-05-22T05:00:26.45142Z","iopub.status.idle":"2024-05-22T05:00:26.475494Z","shell.execute_reply.started":"2024-05-22T05:00:26.451391Z","shell.execute_reply":"2024-05-22T05:00:26.474533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_metadata['primary_label'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:26.768322Z","iopub.execute_input":"2024-05-22T05:00:26.769274Z","iopub.status.idle":"2024-05-22T05:00:26.773465Z","shell.execute_reply.started":"2024-05-22T05:00:26.769235Z","shell.execute_reply":"2024-05-22T05:00:26.772349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_metadata['primary_label'].unique()","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:27.130694Z","iopub.execute_input":"2024-05-22T05:00:27.131511Z","iopub.status.idle":"2024-05-22T05:00:27.146048Z","shell.execute_reply.started":"2024-05-22T05:00:27.131459Z","shell.execute_reply":"2024-05-22T05:00:27.145166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_metadata['filename']","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:27.512142Z","iopub.execute_input":"2024-05-22T05:00:27.51251Z","iopub.status.idle":"2024-05-22T05:00:27.522129Z","shell.execute_reply.started":"2024-05-22T05:00:27.512458Z","shell.execute_reply":"2024-05-22T05:00:27.521089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(100,50))\nplt.bar(train_metadata[\"primary_label\"].value_counts().reset_index()[\"primary_label\"],\n       train_metadata[\"primary_label\"].value_counts().reset_index()[\"count\"])\n","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:27.855868Z","iopub.execute_input":"2024-05-22T05:00:27.856564Z","iopub.status.idle":"2024-05-22T05:00:31.573028Z","shell.execute_reply.started":"2024-05-22T05:00:27.856527Z","shell.execute_reply":"2024-05-22T05:00:31.572057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Assuming your CSV has a column 'class' which represents the class labels\nclasses = train_metadata['primary_label'].unique()\n\n# Dictionary to hold the split data for each class\nsplit_data = {}\n\n# Loop through each class and split the data\nfor class_label in classes:\n    # Filter data for the current class\n    class_data = train_metadata[train_metadata['primary_label'] == class_label]\n    \n    # Split the data into 80% training and 20% testing\n    train_data, test_data = train_test_split(class_data, test_size=0.2, random_state=42)\n    \n    # Store the split data in the dictionary\n    split_data[class_label] = {'train': train_data, 'test': test_data}\n\n# Create separate DataFrames for 80% and 20% data for each class\ntrain_data_80 = pd.concat([split_data[class_label]['train'] for class_label in classes]).reset_index()\ntest_data_20 = pd.concat([split_data[class_label]['test'] for class_label in classes]).reset_index()\n\n# Accessing the combined split data\n# train_data_80 contains 80% of the data from all classes\n# test_data_20 contains 20% of the data from all classes\n","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:31.575124Z","iopub.execute_input":"2024-05-22T05:00:31.575567Z","iopub.status.idle":"2024-05-22T05:00:31.846904Z","shell.execute_reply.started":"2024-05-22T05:00:31.575535Z","shell.execute_reply":"2024-05-22T05:00:31.846175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_80","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:31.847843Z","iopub.execute_input":"2024-05-22T05:00:31.848096Z","iopub.status.idle":"2024-05-22T05:00:31.859266Z","shell.execute_reply.started":"2024-05-22T05:00:31.848075Z","shell.execute_reply":"2024-05-22T05:00:31.858298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_encoder = LabelEncoder()\n","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:31.861226Z","iopub.execute_input":"2024-05-22T05:00:31.861513Z","iopub.status.idle":"2024-05-22T05:00:31.868587Z","shell.execute_reply.started":"2024-05-22T05:00:31.861468Z","shell.execute_reply":"2024-05-22T05:00:31.867795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_80['primary_label'] = label_encoder.fit_transform(train_data_80['primary_label'])\ntrain_data_80","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:31.86954Z","iopub.execute_input":"2024-05-22T05:00:31.869817Z","iopub.status.idle":"2024-05-22T05:00:31.890693Z","shell.execute_reply.started":"2024-05-22T05:00:31.869794Z","shell.execute_reply":"2024-05-22T05:00:31.889698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data_20['primary_label'] = label_encoder.fit_transform(test_data_20['primary_label'])\ntest_data_20","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:31.891651Z","iopub.execute_input":"2024-05-22T05:00:31.891892Z","iopub.status.idle":"2024-05-22T05:00:31.904682Z","shell.execute_reply.started":"2024-05-22T05:00:31.891871Z","shell.execute_reply":"2024-05-22T05:00:31.903696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Mel","metadata":{}},{"cell_type":"code","source":"def plot_mel_spectrogram(audio_file):\n    # Load the audio file\n    y, sr = librosa.load(audio_file)\n\n    # Compute the Mel spectrogram\n    n_fft = 2048  # FFT window size\n    hop_length = 512  # Hop length\n    n_mels = 128  # Number of Mel bands\n    mel_spec = librosa.feature.melspectrogram(y=y, sr=sr, n_fft=n_fft, hop_length=hop_length, n_mels=n_mels)\n\n    # Convert to decibels (log scale)\n    mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max)\n    print('mel shape : ',mel_spec_db.shape)\n    # Visualize the spectrogram\n    plt.figure(figsize=(10, 4))\n    librosa.display.specshow(mel_spec_db, sr=sr, hop_length=hop_length, x_axis='time', y_axis='mel')\n    plt.colorbar(format='%+2.0f dB')\n    plt.title('Mel Spectrogram')\n    plt.show()\n    \n#     return mel_spec_db\n\n    # Optionally, save the spectrogram to an image file\n    # plt.savefig('mel_spectrogram.png', bbox_inches='tight', pad_inches=0)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:31.905948Z","iopub.execute_input":"2024-05-22T05:00:31.906305Z","iopub.status.idle":"2024-05-22T05:00:31.9167Z","shell.execute_reply.started":"2024-05-22T05:00:31.906269Z","shell.execute_reply":"2024-05-22T05:00:31.915442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_mel_spectrogram(\"/kaggle/input/birdclef-2024/train_audio/asbfly/XC164848.ogg\")","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:32.182845Z","iopub.execute_input":"2024-05-22T05:00:32.183527Z","iopub.status.idle":"2024-05-22T05:00:42.858588Z","shell.execute_reply.started":"2024-05-22T05:00:32.183473Z","shell.execute_reply":"2024-05-22T05:00:42.857642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_mel_spectrogram(\"/kaggle/input/birdclef-2024/train_audio/asbfly/XC134896.ogg\")","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:42.860648Z","iopub.execute_input":"2024-05-22T05:00:42.861583Z","iopub.status.idle":"2024-05-22T05:00:43.49345Z","shell.execute_reply.started":"2024-05-22T05:00:42.861551Z","shell.execute_reply":"2024-05-22T05:00:43.492538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_audio(audio_path, target_duration=30.0, n_mels=128, show=False, max_duration=None):\n    # Load the audio file\n    audio, sr = librosa.load(audio_path, sr=None)\n    \n    # Calculate the duration of the audio in seconds\n    audio_duration = len(audio) / sr\n    \n    # Calculate the target length in samples\n    target_length = int(target_duration * sr)\n    \n    if max_duration is None:\n        max_duration = target_duration\n    \n    # Pad or crop the audio to match the target duration\n    if audio_duration >= max_duration:\n        # If audio is longer than max duration, crop it\n        start_time = np.random.uniform(0, audio_duration - target_duration)\n        start_sample = int(start_time * sr)\n        end_sample = start_sample + target_length\n        processed_audio = audio[start_sample:end_sample]\n    else:\n        # If audio is shorter than max duration, pad it\n        pad_length = target_length - len(audio)\n        processed_audio = np.pad(audio, (0, pad_length), mode='constant')\n    \n    # Extract Mel spectrogram\n    n_fft = 2048  # FFT window size\n    hop_length = 512  # Hop length\n    mel_spec = librosa.feature.melspectrogram(y=processed_audio, sr=sr, n_fft=n_fft, hop_length=hop_length, n_mels=n_mels)\n#     print(\"mel spec shape : \",mel_spec.shape)\n    mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max)\n#     print(\"mel spec db : \",mel_spec_db.shape)\n    \n    if show:\n        # Plot Mel spectrogram if show is True\n        plt.figure(figsize=(10, 4))\n        librosa.display.specshow(mel_spec_db, sr=sr, hop_length=hop_length, x_axis='time', y_axis='mel')\n        plt.colorbar(format='%+2.0f dB')\n        plt.title('Mel Spectrogram')\n        plt.xlabel('Time')\n        plt.ylabel('Mel Frequencies')\n        plt.tight_layout()\n        plt.show()\n        \n#     # Transpose and convert to PyTorch tensor\n    if mel_spec_db.shape[1] < 938:\n        pad_width = ((0, 0), (0, 938 - mel_spec_db.shape[1]))\n        # Pad the array along axis 1\n        mel_spec_db = np.pad(mel_spec_db, pad_width, mode='constant')\n    \n    mel_spec_tensor = torch.FloatTensor(mel_spec_db).unsqueeze(0)\n    \n    return mel_spec_tensor","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:43.4948Z","iopub.execute_input":"2024-05-22T05:00:43.495456Z","iopub.status.idle":"2024-05-22T05:00:43.50801Z","shell.execute_reply.started":"2024-05-22T05:00:43.495424Z","shell.execute_reply":"2024-05-22T05:00:43.507033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"process_audio(\"/kaggle/input/birdclef-2024/train_audio/asbfly/XC134896.ogg\",target_duration=15.0,show=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:43.510455Z","iopub.execute_input":"2024-05-22T05:00:43.510815Z","iopub.status.idle":"2024-05-22T05:00:44.300885Z","shell.execute_reply.started":"2024-05-22T05:00:43.510785Z","shell.execute_reply":"2024-05-22T05:00:44.299914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"process_audio(\"/kaggle/input/birdclef-2024/train_audio/asbfly/XC164848.ogg\", target_duration=15.0,show=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:44.302462Z","iopub.execute_input":"2024-05-22T05:00:44.302813Z","iopub.status.idle":"2024-05-22T05:00:44.906022Z","shell.execute_reply.started":"2024-05-22T05:00:44.302785Z","shell.execute_reply":"2024-05-22T05:00:44.905048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataloader","metadata":{}},{"cell_type":"code","source":"class MelDataset(Dataset):\n    def __init__(self, df : pd.DataFrame , max_length_sec : int = 15, n_mels : int = 128):\n        self.df = df\n        self.max_length_sec = max_length_sec\n        self.n_mels = n_mels\n        self.root_path = \"/kaggle/input/birdclef-2024/train_audio\"\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        primary_label = self.df[\"primary_label\"][idx]\n        file_path = os.path.join(self.root_path,\n                                 self.df[\"filename\"][idx])\n        mel_features = process_audio(file_path,\n                                      target_duration =self.max_length_sec,\n                                      n_mels=self.n_mels)\n        return mel_features, primary_label","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:44.907074Z","iopub.execute_input":"2024-05-22T05:00:44.907328Z","iopub.status.idle":"2024-05-22T05:00:44.9142Z","shell.execute_reply.started":"2024-05-22T05:00:44.907306Z","shell.execute_reply":"2024-05-22T05:00:44.913305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = MelDataset(train_data_80)\ntrain_dataloader = DataLoader(train_data, batch_size=32, shuffle=True)\n\ntest_data = MelDataset(test_data_20)\ntest_dataloader = DataLoader(test_data, batch_size=32, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:44.915548Z","iopub.execute_input":"2024-05-22T05:00:44.91592Z","iopub.status.idle":"2024-05-22T05:00:44.930047Z","shell.execute_reply.started":"2024-05-22T05:00:44.915889Z","shell.execute_reply":"2024-05-22T05:00:44.929295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for batch in train_dataloader:\n    print(\"Batch label : \",batch[1])\n    print(\"Batch mfcc shape :\",batch[0].size())\n    break","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:44.930993Z","iopub.execute_input":"2024-05-22T05:00:44.931227Z","iopub.status.idle":"2024-05-22T05:00:48.457062Z","shell.execute_reply.started":"2024-05-22T05:00:44.931205Z","shell.execute_reply":"2024-05-22T05:00:48.456121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install torchviz\nfrom torchviz import make_dot","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### ViT","metadata":{}},{"cell_type":"code","source":"class PatchEmbedding(nn.Module):\n    def __init__(self, image_size, patch_size, in_channels, embed_dim):\n        super(PatchEmbedding, self).__init__()\n        self.projection = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)\n\n    def forward(self, x):\n        patches = self.projection(x)  # (B, C, H//P, W//P)\n        B, C, H, W = patches.shape\n        patches = patches.permute(0, 2, 3, 1).reshape(B, -1, C)  # (B, (H//P)*(W//P), C)\n        return patches\n\nclass VisionTransformer(nn.Module):\n    def __init__(self, image_size, patch_size, in_channels, embed_dim, num_classes, num_heads, num_layers, hidden_dim, dropout):\n        super(VisionTransformer, self).__init__()\n        self.patch_embed = PatchEmbedding(image_size, patch_size, in_channels, embed_dim)\n        self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim))\n        self.transformer_encoder = nn.TransformerEncoder(nn.TransformerEncoderLayer(d_model=embed_dim, nhead=num_heads, dim_feedforward=hidden_dim, dropout=dropout), num_layers=num_layers)\n        self.fc = nn.Linear(embed_dim, num_classes)\n\n    def forward(self, x):\n        B = x.size(0)\n        patches = self.patch_embed(x)\n        cls_tokens = self.cls_token.expand(B, -1, -1)  # (B, 1, C)\n        x = torch.cat([cls_tokens, patches], dim=1)  # (B, 1 + N, C)\n        x = self.transformer_encoder(x)\n        x = x[:, 0]  # take only the cls_token\n        x = self.fc(x)\n        return x\n\n# Example usage\nimage_size = (128, 938)\npatch_size = 16\nin_channels = 1  # Grayscale image\nembed_dim = 256\nnum_classes = 182\nnum_heads = 8\nnum_layers = 6\nhidden_dim = 512\ndropout = 0.1\n\nmodel = VisionTransformer(image_size=image_size,\n                          patch_size=patch_size,\n                          in_channels=in_channels,\n                          embed_dim=embed_dim,\n                          num_classes=num_classes,\n                          num_heads=num_heads,\n                          num_layers=num_layers,\n                          hidden_dim=hidden_dim,\n                          dropout=dropout)\n\n# Example input\nx = torch.randn(1, in_channels, image_size[0], image_size[1])\noutput = model(x)\nprint(output.shape)  # Output shape: (1, num_classes)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T08:51:04.951195Z","iopub.execute_input":"2024-05-10T08:51:04.952121Z","iopub.status.idle":"2024-05-10T08:51:05.060277Z","shell.execute_reply.started":"2024-05-10T08:51:04.952086Z","shell.execute_reply":"2024-05-10T08:51:05.059232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### CNN","metadata":{}},{"cell_type":"code","source":"class MelClassifier(nn.Module):\n    def __init__(self, num_classes):\n        super(MelClassifier, self).__init__()\n        self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1)\n        self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)\n        self.conv3 = nn.Conv2d(64, 128, kernel_size=3, padding=1)\n        self.conv4 = nn.Conv2d(128, 128, kernel_size=3, padding=1)\n        self.pool = nn.MaxPool2d(2, 2)\n        self.fc1 = nn.Linear(128 * 8 * 58, 512)\n        self.fc2 = nn.Linear(512, num_classes)\n\n    def forward(self, x):\n        x = self.pool(torch.relu(self.conv1(x)))\n        x = self.pool(torch.relu(self.conv2(x)))\n        x = self.pool(torch.relu(self.conv3(x)))\n        x = self.pool(torch.relu(self.conv4(x)))\n        x = torch.flatten(x, 1)\n        x = torch.relu(self.fc1(x))\n        x = self.fc2(x)\n        return x\n    \nmodel = MelClassifier(num_classes=182)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:48.761911Z","iopub.execute_input":"2024-05-22T05:00:48.76263Z","iopub.status.idle":"2024-05-22T05:00:49.067789Z","shell.execute_reply.started":"2024-05-22T05:00:48.7626Z","shell.execute_reply":"2024-05-22T05:00:49.06696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"  # Use GPU if available\ndevice","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:00:49.704632Z","iopub.execute_input":"2024-05-22T05:00:49.705337Z","iopub.status.idle":"2024-05-22T05:00:49.760259Z","shell.execute_reply.started":"2024-05-22T05:00:49.705308Z","shell.execute_reply":"2024-05-22T05:00:49.759157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torch.load(\"/kaggle/input/birdclef/pytorch/cnn/1/birdclef_mel_ep13_acc73.pth\")\nmodel","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:01:16.5825Z","iopub.execute_input":"2024-05-22T05:01:16.583345Z","iopub.status.idle":"2024-05-22T05:01:17.780661Z","shell.execute_reply.started":"2024-05-22T05:01:16.583312Z","shell.execute_reply":"2024-05-22T05:01:17.779682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create model, loss function, and optimizer\nmodel = model.to(device)\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=0.001)\n# Training the model\nnum_epochs = 13\ntotal_epoch = 20\nfor epoch in range(num_epochs, total_epoch):\n    model.train()  # Set the model to train mode\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    for inputs, labels in tqdm(train_dataloader):\n        inputs, labels = inputs.to(device), labels.to(device)\n        # Forward pass\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        \n        # Backward and optimize\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        # Update statistics\n        running_loss += loss.item()\n        _, predicted = torch.max(outputs.data, 1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n    \n    # Calculate epoch loss and accuracy\n    epoch_loss = running_loss / len(train_dataloader)\n    epoch_accuracy = 100 * correct / total\n    torch.save(model, f'birdclef_mel_ep{epoch+1}_acc{epoch_accuracy:.2f}.pth')\n    print(f'Epoch [{epoch+1}/{total_epoch}], Loss: {epoch_loss:.4f}, Accuracy: {epoch_accuracy:.2f}%')","metadata":{"execution":{"iopub.status.busy":"2024-05-22T05:03:45.846499Z","iopub.execute_input":"2024-05-22T05:03:45.847305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to('cpu')","metadata":{"execution":{"iopub.status.busy":"2024-05-16T16:44:27.296766Z","iopub.execute_input":"2024-05-16T16:44:27.297497Z","iopub.status.idle":"2024-05-16T16:44:27.523656Z","shell.execute_reply.started":"2024-05-16T16:44:27.29745Z","shell.execute_reply":"2024-05-16T16:44:27.522713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model,'birdclef_mel_ep10_acc69.pth')","metadata":{"execution":{"iopub.status.busy":"2024-05-16T16:45:24.983838Z","iopub.execute_input":"2024-05-16T16:45:24.984782Z","iopub.status.idle":"2024-05-16T16:45:25.140418Z","shell.execute_reply.started":"2024-05-16T16:45:24.984743Z","shell.execute_reply":"2024-05-16T16:45:25.139367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"probs = F.softmax(model(mfcc_features.unsqueeze(0)), dim=1)\nprobs","metadata":{"execution":{"iopub.status.busy":"2024-05-16T16:45:41.68962Z","iopub.execute_input":"2024-05-16T16:45:41.690214Z","iopub.status.idle":"2024-05-16T16:45:41.726859Z","shell.execute_reply.started":"2024-05-16T16:45:41.690181Z","shell.execute_reply":"2024-05-16T16:45:41.725667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_encoder.classes_[torch.argmax(probs, axis=1).item()]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Test the model\nmodel.eval()\nwith torch.no_grad():\n    correct = 0\n    total = 0\n    for inputs, labels in test_dataloader:\n        outputs = model(inputs)\n        _, predicted = torch.max(outputs.data, 1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n\n    print('Accuracy of the network on the test images: {} %'.format(100 * correct / total))","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}