{"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":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":53431,"sourceType":"modelInstanceVersion","modelInstanceId":44833}],"dockerImageVersionId":30699,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\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-05-22T09:38:04.305968Z","iopub.execute_input":"2024-05-22T09:38:04.306814Z","iopub.status.idle":"2024-05-22T09:38:11.035275Z","shell.execute_reply.started":"2024-05-22T09:38:04.306776Z","shell.execute_reply":"2024-05-22T09:38:11.034017Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# imports\nimport os\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport torch\nimport torchaudio\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader, Subset, random_split\nfrom torch.utils.data import Dataset\nimport torch.nn as nn\nimport torchvision.models as models\nfrom torchaudio.transforms import MelSpectrogram, Resample\nfrom IPython.display import Audio","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:38:39.834920Z","iopub.execute_input":"2024-05-22T09:38:39.835441Z","iopub.status.idle":"2024-05-22T09:38:46.569566Z","shell.execute_reply.started":"2024-05-22T09:38:39.835409Z","shell.execute_reply":"2024-05-22T09:38:46.568672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"audio_filepath = '/kaggle/input/birdclef-2024/train_audio/'\nBASE_path = '/kaggle/input/birdclef-2024'","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:38:49.190027Z","iopub.execute_input":"2024-05-22T09:38:49.190844Z","iopub.status.idle":"2024-05-22T09:38:49.195992Z","shell.execute_reply.started":"2024-05-22T09:38:49.190811Z","shell.execute_reply":"2024-05-22T09:38:49.194858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/birdclef-2024/train_metadata.csv')\ndf.head(10)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:38:55.663457Z","iopub.execute_input":"2024-05-22T09:38:55.663865Z","iopub.status.idle":"2024-05-22T09:38:55.875396Z","shell.execute_reply.started":"2024-05-22T09:38:55.663834Z","shell.execute_reply":"2024-05-22T09:38:55.874353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.shape","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:39:07.158866Z","iopub.execute_input":"2024-05-22T09:39:07.159753Z","iopub.status.idle":"2024-05-22T09:39:07.168068Z","shell.execute_reply.started":"2024-05-22T09:39:07.159681Z","shell.execute_reply":"2024-05-22T09:39:07.166732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_audio(file_path):\n    waveform, sample_rate = torchaudio.load(file_path)\n    if waveform.shape[0] > 1:\n        waveform = torch.mean(waveform, dim=0, keepdim=True)\n    return waveform.squeeze(0), sample_rate","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:39:15.582280Z","iopub.execute_input":"2024-05-22T09:39:15.582643Z","iopub.status.idle":"2024-05-22T09:39:15.588393Z","shell.execute_reply.started":"2024-05-22T09:39:15.582616Z","shell.execute_reply":"2024-05-22T09:39:15.587223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"audio_path = os.path.join(audio_filepath, df['filename'][0])\naudio, rate = load_audio(audio_path)\n\nprint(audio)\nAudio(audio.numpy(), rate=rate)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:39:22.254650Z","iopub.execute_input":"2024-05-22T09:39:22.255058Z","iopub.status.idle":"2024-05-22T09:39:22.483575Z","shell.execute_reply.started":"2024-05-22T09:39:22.255027Z","shell.execute_reply":"2024-05-22T09:39:22.482439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def crop_or_pad_audio(audio, target_length, sample_rate):\n    # Calculate the target length in samples\n    target_samples = int(target_length * sample_rate)\n\n    # Get the current length of the audio\n    current_samples = audio.shape[0]\n\n    # If the current length is greater than the target length, crop the audio\n    if current_samples > target_samples:\n        audio = audio[:target_samples]\n    # If the current length is less than the target length, pad the audio\n    elif current_samples < target_samples:\n        # Calculate the number of samples to pad\n        num_padding = target_samples - current_samples\n        # Pad the audio with zeros at the end\n        audio = torch.nn.functional.pad(audio, (0, num_padding))\n    \n    return audio","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:39:32.327553Z","iopub.execute_input":"2024-05-22T09:39:32.327973Z","iopub.status.idle":"2024-05-22T09:39:32.334811Z","shell.execute_reply.started":"2024-05-22T09:39:32.327941Z","shell.execute_reply":"2024-05-22T09:39:32.333746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"audio = crop_or_pad_audio(audio, target_length=10, sample_rate=32000)\nprint(audio)\nAudio(audio.numpy(), rate=32000)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:39:39.527496Z","iopub.execute_input":"2024-05-22T09:39:39.527914Z","iopub.status.idle":"2024-05-22T09:39:39.552311Z","shell.execute_reply.started":"2024-05-22T09:39:39.527883Z","shell.execute_reply":"2024-05-22T09:39:39.551143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_melspectrogram(audio):\n    nfft = 2048\n    window = 2048\n    hop_length = 1024\n    sample_rate = 32000\n    n_mels = 256\n    f_min = 0\n    f_max = 16000\n\n    # Ensure the audio tensor is in float32 format\n    audio = audio.float()\n\n    # Create a spectrogram\n    spectrogram_transform = torchaudio.transforms.Spectrogram(\n        n_fft=nfft,\n        win_length=window,\n        hop_length=hop_length,\n        power=2\n    )\n    spect = spectrogram_transform(audio)\n\n    # Convert the spectrogram to mel scale\n    mel_spectrogram_transform = torchaudio.transforms.MelScale(\n        n_mels=n_mels,\n        sample_rate=sample_rate,\n        f_min=f_min,\n        f_max=f_max,\n        n_stft=spect.size(0)  # the number of frequency bins in the spectrogram\n    )\n    mel_spectrogram = mel_spectrogram_transform(spect)\n    \n    mel_spectrogram = mel_spectrogram.transpose(0, 1)\n\n    return mel_spectrogram","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:39:47.334881Z","iopub.execute_input":"2024-05-22T09:39:47.335643Z","iopub.status.idle":"2024-05-22T09:39:47.342805Z","shell.execute_reply.started":"2024-05-22T09:39:47.335613Z","shell.execute_reply":"2024-05-22T09:39:47.341669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_melspectrogram(mel_spectrogram):\n    # Convert PyTorch tensor to NumPy array and transpose for plotting\n    mel_spectrogram_np = mel_spectrogram.numpy().T\n\n    plt.figure()\n    # Use imshow to plot the mel spectrogram\n    plt.imshow(mel_spectrogram_np, cmap='viridis')\n\n    plt.colorbar(format='%+2.0f dB')\n    plt.title('Mel-Spectrogram')\n    plt.xlabel('Time')\n    plt.ylabel('Frequency')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:39:54.519761Z","iopub.execute_input":"2024-05-22T09:39:54.520140Z","iopub.status.idle":"2024-05-22T09:39:54.526426Z","shell.execute_reply.started":"2024-05-22T09:39:54.520114Z","shell.execute_reply":"2024-05-22T09:39:54.525274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mel_spectrogram = get_melspectrogram(audio)\nprint('Audio shape:', audio.shape)\nprint('Melspectrogram shape:', mel_spectrogram.shape)\nplot_melspectrogram(mel_spectrogram)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:40:02.407605Z","iopub.execute_input":"2024-05-22T09:40:02.408007Z","iopub.status.idle":"2024-05-22T09:40:02.935625Z","shell.execute_reply.started":"2024-05-22T09:40:02.407976Z","shell.execute_reply":"2024-05-22T09:40:02.934425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_dbscale_melspectrogram(mel_spectrogram, top_db=80):\n    # Convert power/amplitude to decibels\n    transform = torchaudio.transforms.AmplitudeToDB(stype=\"power\", top_db=80)\n    \n    dbscale_mel_spectrogram = transform(mel_spectrogram)\n    \n    return dbscale_mel_spectrogram","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:40:09.190860Z","iopub.execute_input":"2024-05-22T09:40:09.191281Z","iopub.status.idle":"2024-05-22T09:40:09.197966Z","shell.execute_reply.started":"2024-05-22T09:40:09.191251Z","shell.execute_reply":"2024-05-22T09:40:09.196441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_dbscale_melspectrogram(dbscale_mel_spectrogram):\n    # Convert PyTorch tensor to NumPy array and transpose for plotting\n    dbscale_mel_spectrogram_np = dbscale_mel_spectrogram.numpy().T\n\n    plt.figure()\n    # Use imshow to plot the mel spectrogram\n    plt.imshow(dbscale_mel_spectrogram_np, cmap='viridis')\n\n    plt.colorbar(format='%+2.0f dB')\n    plt.title('DB-Scaled Mel-Spectrogram')\n    plt.xlabel('Time')\n    plt.ylabel('Frequency')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:40:15.455648Z","iopub.execute_input":"2024-05-22T09:40:15.456661Z","iopub.status.idle":"2024-05-22T09:40:15.462897Z","shell.execute_reply.started":"2024-05-22T09:40:15.456624Z","shell.execute_reply":"2024-05-22T09:40:15.461774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dbscale_mel_spectrogram = get_dbscale_melspectrogram(mel_spectrogram)\nprint('Audio shape:', audio.shape)\nprint('DB Scale Mel Spectrogram shape:', dbscale_mel_spectrogram.shape)\nplot_dbscale_melspectrogram(dbscale_mel_spectrogram)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:40:36.638011Z","iopub.execute_input":"2024-05-22T09:40:36.638383Z","iopub.status.idle":"2024-05-22T09:40:37.026979Z","shell.execute_reply.started":"2024-05-22T09:40:36.638354Z","shell.execute_reply":"2024-05-22T09:40:37.025899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def frequency_mask(spect, param=10):\n    # Initialize the FrequencyMasking transform\n    frequency_masking = torchaudio.transforms.FrequencyMasking(freq_mask_param=param)\n    \n    # Apply the frequency mask\n    masked_spect = frequency_masking(spect)\n    \n    return masked_spect","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:40:50.899579Z","iopub.execute_input":"2024-05-22T09:40:50.900597Z","iopub.status.idle":"2024-05-22T09:40:50.905863Z","shell.execute_reply.started":"2024-05-22T09:40:50.900558Z","shell.execute_reply":"2024-05-22T09:40:50.904745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"freq_masked_spectrogram = frequency_mask(dbscale_mel_spectrogram)\nprint('Audio shape:', audio.shape)\nprint('DB Scaled Mel Spectrogram shape:', freq_masked_spectrogram.shape)\nplot_dbscale_melspectrogram(freq_masked_spectrogram)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:40:57.528911Z","iopub.execute_input":"2024-05-22T09:40:57.529276Z","iopub.status.idle":"2024-05-22T09:40:57.919008Z","shell.execute_reply.started":"2024-05-22T09:40:57.529250Z","shell.execute_reply":"2024-05-22T09:40:57.918057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def time_mask(spect, param=10):\n    # Initialize the TimeMasking transform\n    time_masking = torchaudio.transforms.TimeMasking(time_mask_param=param)\n    \n    # Apply the time mask\n    masked_spect = time_masking(spect)\n    \n    return masked_spect","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:41:07.494651Z","iopub.execute_input":"2024-05-22T09:41:07.495477Z","iopub.status.idle":"2024-05-22T09:41:07.500834Z","shell.execute_reply.started":"2024-05-22T09:41:07.495446Z","shell.execute_reply":"2024-05-22T09:41:07.499710Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"time_masked_spectrogram = time_mask(freq_masked_spectrogram)\nprint('Audio shape:', audio.shape)\nprint('DB Scaled Mel Spectrogram shape:', time_masked_spectrogram.shape)\nplot_dbscale_melspectrogram(time_masked_spectrogram)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:41:16.559828Z","iopub.execute_input":"2024-05-22T09:41:16.560759Z","iopub.status.idle":"2024-05-22T09:41:17.207975Z","shell.execute_reply.started":"2024-05-22T09:41:16.560723Z","shell.execute_reply":"2024-05-22T09:41:17.206974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_spectrogram(spectrogram):\n    # Ensure the input is a PyTorch tensor\n    if not isinstance(spectrogram, torch.Tensor):\n        spectrogram = torch.tensor(spectrogram)\n\n    # Expand dimensions to create a single-channel image\n    spectrogram_image = spectrogram.unsqueeze(0)  # Adds a channel dimension\n\n    # Replicate the single channel to create a three-channel (RGB) image\n    spectrogram_image = spectrogram_image.repeat(3, 1, 1)  # Repeat across the channel dimension\n\n    # Define the resize transformation\n    resize = transforms.Resize((224, 224), antialias = None)  # Expected size for pretrained ResNet and ViT architectures\n\n    # Apply resize transformation\n    spectrogram_image = resize(spectrogram_image)\n\n    return spectrogram_image\n\nspect_image = preprocess_spectrogram(dbscale_mel_spectrogram)\nprint('Spectrogram shape:', spect_image.shape)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:41:29.585452Z","iopub.execute_input":"2024-05-22T09:41:29.586304Z","iopub.status.idle":"2024-05-22T09:41:29.604880Z","shell.execute_reply.started":"2024-05-22T09:41:29.586260Z","shell.execute_reply":"2024-05-22T09:41:29.603749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_and_process_data(file_path, label):\n    audio, sample_rate = load_audio(file_path)\n    audio = crop_or_pad_audio(audio,10,sample_rate)\n    mel_spect = get_melspectrogram(audio)\n    freq_masked_spect = frequency_mask(mel_spect)\n    time_masked_spect = time_mask(freq_masked_spect)\n    spect_image = preprocess_spectrogram(time_masked_spect)\n    return spect_image, label","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:41:35.902986Z","iopub.execute_input":"2024-05-22T09:41:35.903722Z","iopub.status.idle":"2024-05-22T09:41:35.910602Z","shell.execute_reply.started":"2024-05-22T09:41:35.903662Z","shell.execute_reply":"2024-05-22T09:41:35.909380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def conver_to_one_hot(labels, num_classes):\n    labels_tensor = torch.tensor(labels, dtype=torch.long)\n    \n    one_hot_labels = torch.nn.functional.one_hot(labels_tensor, num_classes = num_classes)\n    \n    return one_hot_labels","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:41:41.959221Z","iopub.execute_input":"2024-05-22T09:41:41.959595Z","iopub.status.idle":"2024-05-22T09:41:41.964953Z","shell.execute_reply.started":"2024-05-22T09:41:41.959558Z","shell.execute_reply":"2024-05-22T09:41:41.963871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 8\nnum_classes = len(df['primary_label'].unique())\n\nclass Audio_preprocessing_dataset(Dataset):\n    def __init__(self, file_paths, labels):\n        self.file_paths = file_paths\n        self.labels = labels\n\n    def __len__(self):\n        return len(self.file_paths)\n\n    def __getitem__(self, index):\n        # Here, load_and_process_data should handle loading the file and all preprocessing steps\n        data, label = load_and_process_data(self.file_paths[index], self.labels[index])\n        return data, label\n\ndef create_dataset(df):\n    file_path = [os.path.join(audio_filepath, filename) for filename in df['filename']]\n    labels = df['primary_label']\n    \n    label_to_index = {label: i for i, label in enumerate(np.unique(labels))}\n    labels = [label_to_index[label] for label in labels]\n    labels = conver_to_one_hot(labels, len(label_to_index))\n    dataset = Audio_preprocessing_dataset(file_path, labels)\n    print(len(dataset))\n#     data_loader = DataLoader(dataset, batch_size=batch_size, shuffle=True)\n#     print(len(data_loader))\n    \n    return dataset","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:41:50.864508Z","iopub.execute_input":"2024-05-22T09:41:50.865419Z","iopub.status.idle":"2024-05-22T09:41:50.878630Z","shell.execute_reply.started":"2024-05-22T09:41:50.865384Z","shell.execute_reply":"2024-05-22T09:41:50.877714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def split_dataset(dataset, train_split=0.8):\n    # Calculate the number of samples for training and validation\n    num_samples = len(dataset)\n    num_train_samples = int(train_split * num_samples)\n    num_val_samples = num_samples - num_train_samples\n\n    # Split the dataset into training and validation sets\n    train_dataset, val_dataset = random_split(dataset, [num_train_samples, num_val_samples])\n\n    return train_dataset, val_dataset","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:41:56.999576Z","iopub.execute_input":"2024-05-22T09:41:56.999949Z","iopub.status.idle":"2024-05-22T09:41:57.006076Z","shell.execute_reply.started":"2024-05-22T09:41:56.999924Z","shell.execute_reply":"2024-05-22T09:41:57.004740Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = create_dataset(df)\ndataset","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:42:03.182856Z","iopub.execute_input":"2024-05-22T09:42:03.183248Z","iopub.status.idle":"2024-05-22T09:42:03.294815Z","shell.execute_reply.started":"2024-05-22T09:42:03.183218Z","shell.execute_reply":"2024-05-22T09:42:03.293727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset, val_dataset = split_dataset(dataset)\ntrain_dataset, val_dataset","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:42:10.703443Z","iopub.execute_input":"2024-05-22T09:42:10.704246Z","iopub.status.idle":"2024-05-22T09:42:10.712857Z","shell.execute_reply.started":"2024-05-22T09:42:10.704214Z","shell.execute_reply":"2024-05-22T09:42:10.711765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:42:16.007989Z","iopub.execute_input":"2024-05-22T09:42:16.008845Z","iopub.status.idle":"2024-05-22T09:42:16.014711Z","shell.execute_reply.started":"2024-05-22T09:42:16.008801Z","shell.execute_reply":"2024-05-22T09:42:16.013494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_loader)","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:42:21.927633Z","iopub.execute_input":"2024-05-22T09:42:21.928541Z","iopub.status.idle":"2024-05-22T09:42:21.934582Z","shell.execute_reply.started":"2024-05-22T09:42:21.928507Z","shell.execute_reply":"2024-05-22T09:42:21.933415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for data, target in train_loader:\n    print(data.shape, target.shape)\n    img = data[6]\n    img_np = img.cpu().numpy()\n    img_np = np.transpose(img_np, (1, 2, 0))\n    plt.imshow(img_np, cmap='viridis')\n    plt.axis('off')\n    plt.show()\n    break","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:42:27.640553Z","iopub.execute_input":"2024-05-22T09:42:27.641538Z","iopub.status.idle":"2024-05-22T09:42:28.206098Z","shell.execute_reply.started":"2024-05-22T09:42:27.641495Z","shell.execute_reply":"2024-05-22T09:42:28.204614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\ntorch.manual_seed(42)\nclass MultiClassTrainer:\n    def __init__(self, train_loader, val_loader, num_classes=182, lr=0.001, device='cuda'):\n        # Initialize parameters\n        self.train_loader = train_loader\n        self.val_loader = val_loader\n        self.device = device if device else torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        self.num_classes = num_classes\n        \n        # Initialize model\n        self.model = models.resnet18(weights=None)\n        self.model.fc = nn.Linear(self.model.fc.in_features, self.num_classes)\n\n        self.checkpoint = torch.load('/kaggle/input/birdclef_resnet18/pytorch/checkpoint_4/1/resnet18_checkpoint_4.pt', map_location=device)\n#         self.checkpoint.pop('fc.weight')\n#         self.checkpoint.pop('fc.bias')\n\n        # Load the model checkpoint\n        self.model.load_state_dict(self.checkpoint, strict=False)\n        # Define the final fully connected layer to match the number of classes\n#         self.model.fc = nn.Linear(self.model.fc.in_features, self.num_classes)\n        self.model = self.model.to(self.device)\n        \n        # Initialize loss function and optimizer\n        self.criterion = nn.CrossEntropyLoss()\n        self.optimizer = torch.optim.Adam(self.model.parameters(), lr=lr)\n        \n        # Initialize history tracking\n        self.train_loss_history = []\n        self.val_loss_history = []\n\n    def train(self, num_epochs=100, model_checkpoint_prefix='model_checkpoint', save_interval=10):\n        for epoch in range(num_epochs):\n            epoch_train_loss = 0.0\n            epoch_train_corrects = 0\n            self.model.train()\n            for inputs, targets in tqdm(self.train_loader):\n                inputs = inputs.to(self.device)\n                targets = targets.to(self.device)\n\n                self.optimizer.zero_grad()\n                outputs = self.model(inputs)\n                _, preds = torch.max(outputs, 1)\n        \n                loss = self.criterion(outputs, targets.argmax(dim=1))\n                loss.backward()\n                self.optimizer.step()\n\n                epoch_train_loss += loss.item() * inputs.size(0)\n                epoch_train_corrects += torch.sum(preds == targets.argmax(dim=1))\n\n            average_train_loss = epoch_train_loss / len(self.train_loader.dataset)\n            train_accuracy = epoch_train_corrects.double() / len(self.train_loader.dataset)\n            self.train_loss_history.append(average_train_loss)\n\n            # Validation\n            epoch_val_loss = 0.0\n            epoch_val_corrects = 0\n            self.model.eval()\n            with torch.no_grad():\n                for inputs, targets in tqdm(self.val_loader):\n                    inputs = inputs.to(self.device)\n                    targets = targets.to(self.device)\n\n                    outputs = self.model(inputs)\n                    _, preds = torch.max(outputs, 1)\n                    val_loss = self.criterion(outputs, targets.argmax(dim=1))\n\n                    epoch_val_loss += val_loss.item() * inputs.size(0)\n                    epoch_val_corrects += torch.sum(preds == targets.argmax(dim=1))\n\n            average_val_loss = epoch_val_loss / len(self.val_loader.dataset)\n            val_accuracy = epoch_val_corrects.double() / len(self.val_loader.dataset)\n            self.val_loss_history.append(average_val_loss)\n\n            print(f\"Epoch [{epoch+1}/{num_epochs}], Train Loss: {average_train_loss:.4f}, Train Acc: {train_accuracy:.4f}, Val Loss: {average_val_loss:.4f}, Val Acc: {val_accuracy:.4f}\")\n\n            # Save the model checkpoint after every save_interval epochs\n            if (epoch + 1) % save_interval == 0:\n                model_checkpoint_path = f'{model_checkpoint_prefix}_{epoch+1}.pt'\n                torch.save(self.model.state_dict(), model_checkpoint_path)\n                print(f\"Model saved at epoch {epoch+1}\")\n\n    def plot_loss(self):\n        plt.plot(range(1, len(self.train_loss_history) + 1), self.train_loss_history, label='Training Loss')\n        plt.plot(range(1, len(self.val_loss_history) + 1), self.val_loss_history, label='Validation Loss')\n        plt.title('Training and Validation Loss vs Epoch')\n        plt.xlabel('Epoch')\n        plt.ylabel('Loss')\n        plt.legend()\n        plt.show()\n\n\ntrainer = MultiClassTrainer(train_loader, val_loader, num_classes=182)\ntrainer.train(num_epochs=25, model_checkpoint_prefix='resnet18_checkpoint', save_interval=5)\ntrainer.plot_loss()\n","metadata":{"execution":{"iopub.status.busy":"2024-05-22T09:43:45.370042Z","iopub.execute_input":"2024-05-22T09:43:45.370874Z","iopub.status.idle":"2024-05-22T09:44:05.262439Z","shell.execute_reply.started":"2024-05-22T09:43:45.370837Z","shell.execute_reply":"2024-05-22T09:44:05.260725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}