{"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 timm","metadata":{"id":"Bjsd8qrfdGT0","outputId":"59829cc3-12dc-4cdf-cc28-4a0a8dd3b2e8","execution":{"iopub.status.busy":"2023-06-20T08:19:52.552753Z","iopub.execute_input":"2023-06-20T08:19:52.553103Z","iopub.status.idle":"2023-06-20T08:20:03.40623Z","shell.execute_reply.started":"2023-06-20T08:19:52.553078Z","shell.execute_reply":"2023-06-20T08:20:03.404924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nimport os\nimport importlib\nimport multiprocessing as mp\n\n# Pytorch\n# ------------------------------------------------------\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\n#------------------------------------------------------\nimport torchaudio\n\nimport timm\nfrom torch import nn\nimport torch\nimport torchaudio as ta\nfrom torch.cuda.amp import autocast\nimport random\nfrom torch.distributions import Beta\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import Dataset\nimport numpy as np\nimport librosa\nimport ast\nimport numpy as np\nimport pandas as pd\nimport importlib\nimport sys\nimport random\nfrom tqdm import tqdm\nimport gc\nimport argparse\nimport torch\nfrom torch import optim\nfrom torch.cuda.amp import GradScaler, autocast\nfrom collections import defaultdict\nimport cv2\nfrom copy import copy\nimport os\nfrom torch.utils.data import SequentialSampler, DataLoader\nfrom pathlib import Path\nfrom IPython.display import Audio\n\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import f1_score\nimport tensorflow as tf\n\nfrom sklearn.metrics import accuracy_score, precision_score, recall_score, confusion_matrix\n","metadata":{"id":"A9N935raKEwk","execution":{"iopub.status.busy":"2023-06-20T08:20:03.41009Z","iopub.execute_input":"2023-06-20T08:20:03.410455Z","iopub.status.idle":"2023-06-20T08:20:03.422973Z","shell.execute_reply.started":"2023-06-20T08:20:03.410423Z","shell.execute_reply":"2023-06-20T08:20:03.421985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2023-06-20T08:20:03.424485Z","iopub.execute_input":"2023-06-20T08:20:03.424803Z","iopub.status.idle":"2023-06-20T08:20:03.436891Z","shell.execute_reply.started":"2023-06-20T08:20:03.424774Z","shell.execute_reply":"2023-06-20T08:20:03.435596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = \"cuda\"","metadata":{"id":"HRxMaB8vePqJ","outputId":"03e635c5-21a3-4189-95c2-85ac17ed5550","execution":{"iopub.status.busy":"2023-06-20T08:20:03.440088Z","iopub.execute_input":"2023-06-20T08:20:03.44044Z","iopub.status.idle":"2023-06-20T08:20:03.446676Z","shell.execute_reply.started":"2023-06-20T08:20:03.440416Z","shell.execute_reply":"2023-06-20T08:20:03.445692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install tensorflow-extra","metadata":{"id":"vLpbI7Imp1Oa","outputId":"20d9a38a-fb14-4909-b044-a7dd6d482a84","execution":{"iopub.status.busy":"2023-06-20T08:20:03.448102Z","iopub.execute_input":"2023-06-20T08:20:03.448707Z","iopub.status.idle":"2023-06-20T08:20:14.682447Z","shell.execute_reply.started":"2023-06-20T08:20:03.448675Z","shell.execute_reply":"2023-06-20T08:20:14.681307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/birdclef-2023/train_metadata.csv')","metadata":{"id":"_moPYtw7M6Z_","execution":{"iopub.status.busy":"2023-06-20T08:20:14.684607Z","iopub.execute_input":"2023-06-20T08:20:14.685041Z","iopub.status.idle":"2023-06-20T08:20:14.776603Z","shell.execute_reply.started":"2023-06-20T08:20:14.685Z","shell.execute_reply":"2023-06-20T08:20:14.775604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = LabelEncoder()\ndf['primary_label_encoded'] = encoder.fit_transform(df['primary_label'])","metadata":{"id":"8hQimENCfAoJ","execution":{"iopub.status.busy":"2023-06-20T08:20:14.778237Z","iopub.execute_input":"2023-06-20T08:20:14.77861Z","iopub.status.idle":"2023-06-20T08:20:14.790181Z","shell.execute_reply.started":"2023-06-20T08:20:14.778575Z","shell.execute_reply":"2023-06-20T08:20:14.789144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=5)\nfor k, (_, val_ind) in enumerate(skf.split(X=df, y=df['primary_label_encoded'])):\n    df.loc[val_ind, 'kfold'] = k\n\n","metadata":{"id":"BEj1k4PXewXL","outputId":"28207a53-a79d-4289-fb36-c706c114b7d8","execution":{"iopub.status.busy":"2023-06-20T08:20:14.791608Z","iopub.execute_input":"2023-06-20T08:20:14.792188Z","iopub.status.idle":"2023-06-20T08:20:14.815388Z","shell.execute_reply.started":"2023-06-20T08:20:14.792156Z","shell.execute_reply":"2023-06-20T08:20:14.814463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def filter_data(df, thr=5):\n    # Count the number of samples for each class\n    class_counts = df.primary_label.value_counts()\n\n    # Condition that selects classes with fewer than `thr` samples\n    condition = df.primary_label.isin(class_counts[class_counts < thr].index.tolist())\n\n    # Add a new column to indicate samples for cross-validation\n    df['cv'] = True\n\n    # Set 'cv' as False for classes with fewer than 'thr' samples\n    df.loc[condition, 'cv'] = False\n\n    # Return the filtered dataframe\n    return df\n\ndef upsample_data(df, thr=20):\n    # Get the class distribution\n    class_distribution = df['primary_label'].value_counts()\n\n    # Identify the classes with fewer than 'thr' samples\n    undersampled_classes = class_distribution[class_distribution < thr].index.tolist()\n\n    # Create an empty list to store the upsampled dataframes\n    upsampled_dfs = []\n\n    # Loop through the undersampled classes and perform upsampling\n    for c in undersampled_classes:\n        # Get the dataframe for the current class\n        class_df = df.query(\"primary_label==@c\")\n        # Find the number of samples to add\n        num_samples_to_add = thr - class_df.shape[0]\n        # Upsample the dataframe\n        class_df = class_df.sample(n=num_samples_to_add, replace=True, random_state=42)\n        # Append the upsampled dataframe to the list\n        upsampled_dfs.append(class_df)\n\n    # Concatenate the upsampled dataframes with the original dataframe\n    upsampled_df = pd.concat([df] + upsampled_dfs, axis=0, ignore_index=True)\n\n    return upsampled_df\n\ndef downsample_data(df, thr=500):\n    # Get the class distribution\n    class_distribution = df['primary_label'].value_counts()\n\n    # Identify the classes with more than 'thr' samples\n    oversampled_classes = class_distribution[class_distribution > thr].index.tolist()\n\n    # Create an empty list to store the downsampled dataframes\n    downsampled_dfs = []\n\n    # Loop through the oversampled classes and perform downsampling\n    for c in oversampled_classes:\n        # Get the dataframe for the current class\n        class_df = df.query(\"primary_label==@c\")\n        # Remove the data for that class\n        df = df.query(\"primary_label!=@c\")\n        # Downsample the dataframe\n        class_df = class_df.sample(n=thr, replace=False, random_state=42)\n        # Append the downsampled dataframe to the list\n        downsampled_dfs.append(class_df)\n\n    # Concatenate the downsampled dataframes with the original dataframe\n    downsampled_df = pd.concat([df] + downsampled_dfs, axis=0, ignore_index=True)\n\n    return downsampled_df","metadata":{"id":"3UwRv-IPNOBe","execution":{"iopub.status.busy":"2023-06-20T08:20:14.817016Z","iopub.execute_input":"2023-06-20T08:20:14.817415Z","iopub.status.idle":"2023-06-20T08:20:14.832451Z","shell.execute_reply.started":"2023-06-20T08:20:14.817381Z","shell.execute_reply":"2023-06-20T08:20:14.831334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CUDA_LAUNCH_BLOCKING=1.\n","metadata":{"id":"OnQpAwDKq9vb","execution":{"iopub.status.busy":"2023-06-20T08:20:14.838019Z","iopub.execute_input":"2023-06-20T08:20:14.838341Z","iopub.status.idle":"2023-06-20T08:20:14.844606Z","shell.execute_reply.started":"2023-06-20T08:20:14.838311Z","shell.execute_reply":"2023-06-20T08:20:14.843714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = upsample_data(df)\ndf = downsample_data(df)","metadata":{"id":"utoavx_EnPzU","execution":{"iopub.status.busy":"2023-06-20T08:20:14.84648Z","iopub.execute_input":"2023-06-20T08:20:14.847168Z","iopub.status.idle":"2023-06-20T08:20:15.164612Z","shell.execute_reply.started":"2023-06-20T08:20:14.847137Z","shell.execute_reply":"2023-06-20T08:20:15.163694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"filter_data(df).head()","metadata":{"id":"MBrg6Z7LQuL5","outputId":"7b44b186-6af3-43cd-e5b3-6fa019be5e88","execution":{"iopub.status.busy":"2023-06-20T08:20:15.166518Z","iopub.execute_input":"2023-06-20T08:20:15.166883Z","iopub.status.idle":"2023-06-20T08:20:15.192111Z","shell.execute_reply.started":"2023-06-20T08:20:15.166841Z","shell.execute_reply":"2023-06-20T08:20:15.191035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Augmentation Functions","metadata":{"id":"95mZbimVipv8"}},{"cell_type":"code","source":"# Import required packages\nimport tensorflow_extra as tfe\n\n\n# Generates random integer\ndef random_int(shape=[], minval=0, maxval=1):\n    return tf.random.uniform(shape=shape, minval=minval, maxval=maxval, dtype=tf.int32)\n\n\n# Generats random float\ndef random_float(minval=0.0, maxval=1.0):\n    return torch.rand(1) * (maxval - minval) + minval\n\ndef time_shift(audio, prob=0.5):\n    # Randomly apply time shift with probability `prob`\n    if random_float() < prob:\n        # Calculate random shift value\n        shift = random_float(0, audio.shape[0])\n        # Randomly set the shift to be negative with 50% probability\n        if random_float() < 0.5:\n            shift = -shift\n        # Roll the audio signal by the shift value\n        audio = torch.roll(audio, shifts=shift, dims=0)\n    return audio\n\n# Apply random noise to audio data\n# @tf.function\ndef gaussian_noise(audio, std=[0.0025, 0.025], prob=0.5):\n    # Select a random value of standard deviation for Gaussian noise within the given range\n    std = random_float(std[0], std[1])\n    # Randomly apply Gaussian noise with probability `prob`\n    if random_float() < prob:\n        # Generate random Gaussian noise with the same shape as the audio signal\n        noise = torch.empty_like(audio).normal_(mean=0, std=float(std))\n        audio = audio + noise\n    return audio\n\n# Applies augmentation to Audio Signal\ndef audio_augmentation(audio):\n    # Apply time shift and Gaussian noise to the audio signal\n#     audio = time_shift(audio, prob=0.01)\n    audio = gaussian_noise(audio, prob= 0.35)\n    return audio\n\n# CutMix & MixUp\nmixup_layer = tfe.layers.MixUp(alpha=0.5, prob=0.7)\ncutmix_layer = tfe.layers.CutMix(alpha=2.5, prob=0.65)\n\ndef cutmix_up(audios, labels):\n    audios, labels = mixup_layer(audios, labels, training=True)\n    audios, labels = cutmix_layer(audios, labels, training=True)\n    return audios, labels","metadata":{"id":"R2jam0_9SVjK","execution":{"iopub.status.busy":"2023-06-20T08:20:15.193752Z","iopub.execute_input":"2023-06-20T08:20:15.194255Z","iopub.status.idle":"2023-06-20T08:20:15.209872Z","shell.execute_reply.started":"2023-06-20T08:20:15.194222Z","shell.execute_reply":"2023-06-20T08:20:15.208903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_audio(audio_path):\n\n    # Load an audio file\n    samples, sample_rate = librosa.load(audio_path)\n    samples = np.array(crop_or_pad(samples))\n\n    # Visualize the waveform\n    plt.figure(figsize=(14, 5))\n    librosa.display.waveshow(samples, sr=sample_rate)\n    plt.title('Waveform')\n","metadata":{"id":"bwQMR-8rmbkA","execution":{"iopub.status.busy":"2023-06-20T08:20:15.211067Z","iopub.execute_input":"2023-06-20T08:20:15.211401Z","iopub.status.idle":"2023-06-20T08:20:15.223603Z","shell.execute_reply.started":"2023-06-20T08:20:15.21137Z","shell.execute_reply":"2023-06-20T08:20:15.222595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_mfcc(audio_path):\n\n    # Load an audio file\n    samples, sample_rate = librosa.load(audio_path)\n    samples = np.array(crop_or_pad(samples))\n    # Compute the MFCCs\n    mfccs = librosa.feature.mfcc(y=samples, sr=sample_rate, n_mfcc=13)\n    # Visualize the MFCCs\n    plt.figure(figsize=(14, 5))\n    librosa.display.specshow(mfccs, sr=sample_rate, x_axis='time')\n    plt.colorbar()\n    plt.title('MFCCs')\n\n    # Show the plots\n    display(Audio(samples, rate=sample_rate))\n    plt.show()","metadata":{"id":"lssSplIUpqHA","execution":{"iopub.status.busy":"2023-06-20T08:20:15.225183Z","iopub.execute_input":"2023-06-20T08:20:15.225517Z","iopub.status.idle":"2023-06-20T08:20:15.237343Z","shell.execute_reply.started":"2023-06-20T08:20:15.225484Z","shell.execute_reply":"2023-06-20T08:20:15.23638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Decodes Audio\ndef audio_decoder(with_labels=True, dim=32000*10,\n                  take_first=False, num_classes=264):\n    def get_audio(filepath):\n        ftype = filepath[1]\n        filepath = filepath[0]\n        file_bytes = tf.io.read_file(filepath)\n        if ftype:\n            audio = tfio.audio.decode_vorbis(file_bytes) # decode .ogg file\n        else:\n            audio = tfio.audio.decode_mp3(file_bytes) # decode .mp3 file\n        audio = tf.cast(audio, tf.float32)\n        if tf.shape(audio)[1]>1: # stereo -> mono\n            audio = audio[...,0:1]\n        audio = tf.squeeze(audio, axis=-1)\n        return audio\n\n    def crop_or_pad(audio, target_len, pad_mode='constant', take_first=True):\n        audio_len = tf.shape(audio)[0]\n        diff_len = abs(target_len - audio_len)\n        if audio_len < target_len:\n            pad1 = tf.random.uniform([], maxval=diff_len, dtype=tf.int32)\n            pad2 = diff_len - pad1\n            audio = tf.pad(audio, paddings=[[pad1, pad2]], mode=pad_mode)\n        elif audio_len > target_len:\n            if take_first:\n                audio = audio[:target_len]\n            else:\n                idx = tf.random.uniform([], maxval=diff_len, dtype=tf.int32)\n                audio = audio[idx: (idx + target_len)]\n        return tf.reshape(audio, [target_len])\n\n    def get_target(target):\n        target = tf.reshape(target, [1])\n        target = tf.cast(tf.one_hot(target, num_classes), tf.float32)\n        target = tf.reshape(target, [num_classes])\n        return target\n\n    def decode(path):\n        audio = get_audio(path)\n        audio = crop_or_pad(audio, dim) # crop or pad audio to keep a fixed length\n        audio = tf.reshape(audio, [dim])\n        return audio\n\n    def decode_with_labels(path, label):\n        label = get_target(label)\n        return decode(path), label\n\n    return decode_with_labels if with_labels else decode","metadata":{"id":"C__V9UTRooDU","execution":{"iopub.status.busy":"2023-06-20T08:20:15.238749Z","iopub.execute_input":"2023-06-20T08:20:15.239108Z","iopub.status.idle":"2023-06-20T08:20:15.253683Z","shell.execute_reply.started":"2023-06-20T08:20:15.239076Z","shell.execute_reply":"2023-06-20T08:20:15.25268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def crop_or_pad(audio, target_len=10*32000, pad_mode='constant', take_first=True):\n    audio_len = audio.shape[0]\n    diff_len = abs(target_len - audio_len)\n    if audio_len < target_len:\n        pad1 = torch.randint(low=0, high=diff_len, size=())\n        pad2 = diff_len - pad1\n        audio = np.pad(audio, pad_width=((pad1, pad2)), mode=pad_mode)\n    elif audio_len > target_len:\n        if take_first:\n            audio = audio[:target_len]\n        else:\n            idx = torch.randint(low=0, high=diff_len, size=())\n            audio = audio[idx: (idx + target_len)]\n    return audio.reshape([target_len])","metadata":{"id":"vsJh7UFHbVku","execution":{"iopub.status.busy":"2023-06-20T08:20:15.255174Z","iopub.execute_input":"2023-06-20T08:20:15.255687Z","iopub.status.idle":"2023-06-20T08:20:15.26724Z","shell.execute_reply.started":"2023-06-20T08:20:15.255632Z","shell.execute_reply":"2023-06-20T08:20:15.266359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# samples, sample_rate = librosa.load(\"/content/train_audio/abythr1/XC115981.ogg\")\n# samples = np.array(crop_or_pad(samples))\n# # Compute the MFCCs\n# mfccs = librosa.feature.mfcc(y=samples, sr=sample_rate, n_mfcc=13)","metadata":{"id":"NRQ7kDeSb1eV","execution":{"iopub.status.busy":"2023-06-20T08:20:15.26863Z","iopub.execute_input":"2023-06-20T08:20:15.268975Z","iopub.status.idle":"2023-06-20T08:20:15.280844Z","shell.execute_reply.started":"2023-06-20T08:20:15.268944Z","shell.execute_reply":"2023-06-20T08:20:15.279809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BirdCLEFDataset(Dataset):\n    def __init__(self, df, target_sample_rate=32000, max_time=5, image_transforms=None):\n        self.file_paths = df['filename'].values\n        self.labels = df['primary_label_encoded'].values\n        self.target_sample_rate = target_sample_rate\n        num_samples = target_sample_rate * max_time\n        self.num_samples = num_samples\n        self.image_transforms = image_transforms\n\n    def __len__(self):\n        return len(self.file_paths)\n\n    def __getitem__(self, index):\n        filepath = f'/kaggle/input/birdclef-2023/train_audio/{self.file_paths[index]}'\n        audio, sample_rate = torchaudio.load(filepath)\n        audio = self.to_mono(audio)\n        \n        if sample_rate != self.target_sample_rate:\n            resample = Resample(sample_rate, self.target_sample_rate)\n            audio = resample(audio)\n        \n        if audio.shape[0] > self.num_samples:\n            audio = self.crop_audio(audio)\n            \n        if audio.shape[0] < self.num_samples:\n            audio = self.pad_audio(audio)\n            \n        \n        audio = audio_augmentation(audio)\n\n        mel_spectogram = torchaudio.transforms.MFCC(sample_rate=self.target_sample_rate, \n                                        n_mfcc=224, melkwargs={\"n_fft\": 1024,\"n_mels\": 1024, \"center\": False},\n                                        )\n        audio = gaussian_noise(audio)\n\n\n        mel = mel_spectogram(audio)\n        label = torch.tensor(self.labels[index])\n        \n        # Convert to Image\n        image = torch.stack([mel, mel, mel])        \n\n\n\n        # Normalize Image\n        max_val = torch.tensor(image).max()\n        min_val =  torch.tensor(image).min()\n        image = (image -  min_val) / (image - min_val)\n\n        return image, label\n\n    def pad_audio(self, audio):\n        pad_length = self.num_samples - audio.shape[0]\n        last_dim_padding = (0, pad_length)\n        audio = F.pad(audio, last_dim_padding)\n        return audio\n\n    def crop_audio(self, audio):\n        return audio[:self.num_samples]\n\n    def to_mono(self, audio):\n        return torch.mean(audio, axis=0)","metadata":{"id":"v9NFZ2Me6GFz","execution":{"iopub.status.busy":"2023-06-20T08:20:15.282343Z","iopub.execute_input":"2023-06-20T08:20:15.282708Z","iopub.status.idle":"2023-06-20T08:20:15.297614Z","shell.execute_reply.started":"2023-06-20T08:20:15.282676Z","shell.execute_reply":"2023-06-20T08:20:15.294939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        return self.gem(x, p=self.p, eps=self.eps)\n\n    def gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n\n    def __repr__(self):\n        return self.__class__.__name__ + \\\n                '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + \\\n                ', ' + 'eps=' + str(self.eps) + ')'\n\n\nclass BirdCLEFModel(nn.Module):\n    def __init__(self, model_name=\"tf_efficientnet_b4_ns\", embedding_size=768, pretrained=True):\n        super(BirdCLEFModel, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        in_features = self.model.classifier.in_features\n        self.model.classifier = nn.Identity()\n        self.model.global_pool = nn.Identity()\n        self.pooling = GeM()\n        self.embedding = nn.Linear(in_features, embedding_size)\n        self.fc = nn.Linear(embedding_size, 264)\n\n    def forward(self, images):\n        features = self.model(images)\n        pooled_features = self.pooling(features).flatten(1)\n        embedding = self.embedding(pooled_features)\n        output = self.fc(embedding)\n        return output","metadata":{"id":"ReyKqJmkateU","execution":{"iopub.status.busy":"2023-06-20T08:20:15.298858Z","iopub.execute_input":"2023-06-20T08:20:15.29924Z","iopub.status.idle":"2023-06-20T08:20:15.315366Z","shell.execute_reply.started":"2023-06-20T08:20:15.299211Z","shell.execute_reply":"2023-06-20T08:20:15.314469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def loss_fn(outputs, labels):\n    return nn.CrossEntropyLoss()(outputs, labels)\n\ndef train(model, data_loader, optimizer, scheduler, device, epoch):\n    model.train()\n\n    running_loss = 0\n    loop = tqdm(data_loader, position=0)\n    for i, (mels, labels) in enumerate(loop):\n#         print(type(mels))\n#         print(mels.shape)\n        mels = mels.to(device)\n        labels = labels.to(device)\n\n        outputs = model(mels)\n        _, preds = torch.max(outputs, 1)\n\n        loss = loss_fn(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n\n        if scheduler is not None:\n            scheduler.step()\n\n        running_loss += loss.item()\n\n        loop.set_description(f\"Epoch [{epoch+1}/{10}]\")\n        loop.set_postfix(loss=loss.item())\n\n    return running_loss/len(data_loader)","metadata":{"id":"W3p_4ou2cp2l","execution":{"iopub.status.busy":"2023-06-20T08:20:15.316857Z","iopub.execute_input":"2023-06-20T08:20:15.317208Z","iopub.status.idle":"2023-06-20T08:20:15.328468Z","shell.execute_reply.started":"2023-06-20T08:20:15.317156Z","shell.execute_reply":"2023-06-20T08:20:15.327518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def valid(model, data_loader, device, epoch):\n    model.eval()\n\n    running_loss = 0\n    pred = []\n    label = []\n\n    loop = tqdm(data_loader, position=0)\n    for mels, labels in loop:\n#         print(type(mels))\n#         print(mels)\n        mels = mels.to(device)\n        labels = labels.to(device)\n\n        outputs = model(mels)\n        _, preds = torch.max(outputs, 1)\n\n        loss = loss_fn(outputs, labels)\n\n        running_loss += loss.item()\n\n        pred.extend(preds.view(-1).cpu().detach().numpy())\n        label.extend(labels.view(-1).cpu().detach().numpy())\n\n        loop.set_description(f\"Epoch [{epoch+1}/{10}]\")\n        loop.set_postfix(loss=loss.item())\n\n    valid_f1 = f1_score(label, pred, average='macro')\n    valid_accuracy = accuracy_score(label, pred)\n    valid_precision = precision_score(label, pred, average='macro')\n    valid_recall = recall_score(label, pred, average='macro')\n\n    return running_loss/len(data_loader), valid_f1, valid_accuracy, valid_precision, valid_recall","metadata":{"id":"Yx930hJCcs6h","execution":{"iopub.status.busy":"2023-06-20T08:20:15.330072Z","iopub.execute_input":"2023-06-20T08:20:15.330386Z","iopub.status.idle":"2023-06-20T08:20:15.342199Z","shell.execute_reply.started":"2023-06-20T08:20:15.330358Z","shell.execute_reply":"2023-06-20T08:20:15.341603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_loaders(df, fold):\n    df_train = df[df.kfold != fold].reset_index(drop=True)\n    df_valid = df[df.kfold == fold].reset_index(drop=True)\n\n    train_dataset = BirdCLEFDataset(df_train, target_sample_rate=32000, max_time=5)\n    valid_dataset = BirdCLEFDataset(df_valid, target_sample_rate=32000, max_time=5)\n\n    train_loader = DataLoader(train_dataset, batch_size=8,\n                              num_workers=2, shuffle=True, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset, batch_size=8,\n                              num_workers=2, shuffle=False, pin_memory=True)\n\n    return train_loader, valid_loader","metadata":{"id":"0PxG7RN9cu-4","execution":{"iopub.status.busy":"2023-06-20T08:20:15.343291Z","iopub.execute_input":"2023-06-20T08:20:15.3443Z","iopub.status.idle":"2023-06-20T08:20:15.352469Z","shell.execute_reply.started":"2023-06-20T08:20:15.344277Z","shell.execute_reply":"2023-06-20T08:20:15.351827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = BirdCLEFModel().to(device)\noptimizer = Adam(model.parameters(), lr=1e-4)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, eta_min=1e-5, T_max=10)","metadata":{"id":"i4AVzK05c9CG","outputId":"5a15e2eb-a64d-4f81-8fa2-54aa483dfb81","execution":{"iopub.status.busy":"2023-06-20T08:20:15.35478Z","iopub.execute_input":"2023-06-20T08:20:15.355363Z","iopub.status.idle":"2023-06-20T08:20:16.001425Z","shell.execute_reply.started":"2023-06-20T08:20:15.355333Z","shell.execute_reply":"2023-06-20T08:20:16.000447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = df[df.kfold != 0].reset_index(drop=True)\ntrain_dataset = BirdCLEFDataset(df_train, target_sample_rate=32000, max_time=5)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T08:20:16.003357Z","iopub.execute_input":"2023-06-20T08:20:16.003732Z","iopub.status.idle":"2023-06-20T08:20:16.017604Z","shell.execute_reply.started":"2023-06-20T08:20:16.003699Z","shell.execute_reply":"2023-06-20T08:20:16.016485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"for i in range(5):\n    best_valid_f1 = 0\n    train_loader, valid_loader = prepare_loaders(df, i) # fold 0\n    total_valid_f1, total_valid_accuracy, total_valid_precision, total_valid_recall = 0, 0, 0, 0\n    for epoch in range(10):\n        train_loss = train(model, train_loader, optimizer, scheduler, device, epoch)\n        valid_loss, valid_f1, valid_accuracy, valid_precision, valid_recall = valid(model, valid_loader, device, epoch)\n        if valid_f1 > best_valid_f1:\n            torch.save(model.state_dict(), f'./mfcc_model_k{i}.bin')\n            \n    total_valid_f1 += valid_f1\n    total_valid_accuracy += valid_accuracy\n    total_valid_precision += valid_precision\n    total_valid_recall += valid_recall\n    \nprint(\"total_valid_f1: \", total_valid_f1/5)\nprint(\"total_valid_accuracy: \", total_valid_accuracy/5)\nprint(\"total_valid_precision: \", total_valid_precision/5)\nprint(\"total_valid_recall: \", total_valid_recall/5)","metadata":{"id":"iKLp8vd2c-8U","execution":{"iopub.status.busy":"2023-06-20T08:20:16.020765Z","iopub.execute_input":"2023-06-20T08:20:16.021024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df_train = df.reset_index(drop=True)\n# # df_valid = df[df.kfold == fold].reset_index(drop=True)\n\n# train_dataset = BirdCLEFDataset(df_train, target_sample_rate=32000, max_time=5)\n\n# train_loader = DataLoader(train_dataset, batch_size=8,\n#                           num_workers=2, shuffle=True, pin_memory=True, drop_last=True)\n\n# for epoch in range(10):\n#     print(epoch)\n#     train_loss = train(model, train_loader, optimizer, scheduler, device, epoch)\n#     torch.save(model.state_dict(), f'./mel_model.bin')\n#     print(f\"Saved model checkpoint at ./mel_model.bin\")\n#     best_valid_f1 = valid_f1","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}