{"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":"markdown","source":"# Import","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch\nimport torchaudio\nimport math, random\nfrom IPython.display import Audio\nimport librosa\nfrom tqdm import tqdm\nimport warnings\nfrom torch.utils.data import Dataset, DataLoader\nimport os\nimport matplotlib.pyplot as plt\nfrom torchvision.transforms import Resize\nfrom sklearn.model_selection import train_test_split\nimport torchvision.models as models\nimport torch.nn as nn\nfrom timeit import default_timer as timer\nfrom sklearn import preprocessing\nfrom sklearn.metrics import f1_score, recall_score, confusion_matrix, classification_report\n\n# disable all warning messages\nwarnings.filterwarnings(\"ignore\")\n\ntrain_on_gpu= True","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:32:25.388955Z","iopub.execute_input":"2023-05-10T20:32:25.38939Z","iopub.status.idle":"2023-05-10T20:32:29.350955Z","shell.execute_reply.started":"2023-05-10T20:32:25.389354Z","shell.execute_reply":"2023-05-10T20:32:29.350021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed: int):\n    import random, os\n    import numpy as np\n    import torch\n    \n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    \nseed_everything(42)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:32:29.353005Z","iopub.execute_input":"2023-05-10T20:32:29.353577Z","iopub.status.idle":"2023-05-10T20:32:29.363855Z","shell.execute_reply.started":"2023-05-10T20:32:29.353544Z","shell.execute_reply":"2023-05-10T20:32:29.363089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config ={\n 'sample_rate' : 22050,\n 'clip_length' : 5,\n 'classes_number' : 264,\n 'learning_rate': 0.001,\n 'batch_size': 32,\n 'num_epochs': 10,\n \n}","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:32:29.367021Z","iopub.execute_input":"2023-05-10T20:32:29.367589Z","iopub.status.idle":"2023-05-10T20:32:29.373797Z","shell.execute_reply.started":"2023-05-10T20:32:29.367563Z","shell.execute_reply":"2023-05-10T20:32:29.372783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Train DataFrame","metadata":{}},{"cell_type":"markdown","source":"Because the audio dosn't have the same legnth we will create new dataframe contains information about the length of each Audio, </br>\nand then Create dataframe the contains the start and the end of each new audio which will be the input to the AI pipline","metadata":{}},{"cell_type":"code","source":"# Original_Df = pd.read_csv(\"/kaggle/input/birdclef-2023/train_metadata.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-05-08T16:19:41.677496Z","iopub.execute_input":"2023-05-08T16:19:41.67794Z","iopub.status.idle":"2023-05-08T16:19:41.842549Z","shell.execute_reply.started":"2023-05-08T16:19:41.677902Z","shell.execute_reply":"2023-05-08T16:19:41.841368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for i, audio_name in tqdm(enumerate(Original_Df.filename)):\n#     audio_data, Audio_sample_rate  = librosa.load(f\"/kaggle/input/birdclef-2023/train_audio/{audio_name}\")\n#     audio_duration = librosa.get_duration(y=audio_data, sr=Audio_sample_rate)\n#     Original_Df.loc[i, \"duration\"] = audio_duration","metadata":{"execution":{"iopub.status.busy":"2023-05-08T16:31:57.27897Z","iopub.execute_input":"2023-05-08T16:31:57.279388Z","iopub.status.idle":"2023-05-08T16:55:12.544075Z","shell.execute_reply.started":"2023-05-08T16:31:57.279351Z","shell.execute_reply":"2023-05-08T16:55:12.542862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Original_Df.to_csv(\"train_metadata_with_length.csv\", index =False)","metadata":{"execution":{"iopub.status.busy":"2023-05-08T16:55:12.931964Z","iopub.execute_input":"2023-05-08T16:55:12.932417Z","iopub.status.idle":"2023-05-08T16:55:13.139129Z","shell.execute_reply.started":"2023-05-08T16:55:12.93238Z","shell.execute_reply":"2023-05-08T16:55:13.137924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Split each audio file into equil lengths  DF","metadata":{}},{"cell_type":"code","source":"Df_length = pd.read_csv(\"/kaggle/working/train_metadata_with_length.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-05-10T12:09:38.967135Z","iopub.execute_input":"2023-05-10T12:09:38.96752Z","iopub.status.idle":"2023-05-10T12:09:39.09686Z","shell.execute_reply.started":"2023-05-10T12:09:38.967491Z","shell.execute_reply":"2023-05-10T12:09:39.095719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def Create_from_end(Df_length, clip_length=5, overlap_length=0):\n    \n    # Create New DataFrame\n    New_columns = np.append(Df_length.columns.values, [\"start\", \"end\"])\n    New_df = pd.DataFrame(columns=New_columns)\n    \n    # itterate over the audios\n    for i in tqdm(range(len(Df_length)), total=len(Df_length)):\n        row = Df_length.iloc[i]\n        audio_length = original_length = row.duration # Audio Duration\n        \n        start = 0\n        end =  min(clip_length , audio_length)\n\n        while audio_length > 1:\n            #Append Start and the End to the DataFrame\n            new_row = pd.concat([row, pd.Series([start, end], index = [\"start\", \"end\"])])\n            New_df = New_df.append(new_row, ignore_index=True)\n\n            # Update the start and the end\n            start = clip_length + (start - overlap_length)\n            end =  min(clip_length + start, original_length)\n            \n            # Decrease the length of the audio\n            audio_length = audio_length - (clip_length - overlap_length)\n    return New_df","metadata":{"execution":{"iopub.status.busy":"2023-05-10T14:23:21.557254Z","iopub.execute_input":"2023-05-10T14:23:21.557642Z","iopub.status.idle":"2023-05-10T14:23:21.567471Z","shell.execute_reply.started":"2023-05-10T14:23:21.557612Z","shell.execute_reply":"2023-05-10T14:23:21.566104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_df = Create_from_end(Df_length,clip_length=5, overlap_length =2)\n# train_df.to_csv(\"train_df_5s.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T14:26:19.732451Z","iopub.execute_input":"2023-05-10T14:26:19.732817Z","iopub.status.idle":"2023-05-10T15:27:00.806672Z","shell.execute_reply.started":"2023-05-10T14:26:19.732788Z","shell.execute_reply":"2023-05-10T15:27:00.805529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df","metadata":{"execution":{"iopub.status.busy":"2023-05-10T15:39:54.071909Z","iopub.execute_input":"2023-05-10T15:39:54.072334Z","iopub.status.idle":"2023-05-10T15:39:54.101969Z","shell.execute_reply.started":"2023-05-10T15:39:54.072299Z","shell.execute_reply":"2023-05-10T15:39:54.10075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DataSet =  pd.read_csv(\"/kaggle/working/train_df_5s.csv\")\n\nX_train, X_valid, _, _ = train_test_split(DataSet, DataSet[\"primary_label\"], shuffle=True, test_size=0.1,  random_state=42, stratify=DataSet[\"primary_label\"] )\n","metadata":{"execution":{"iopub.status.busy":"2023-05-10T19:46:42.869609Z","iopub.execute_input":"2023-05-10T19:46:42.869974Z","iopub.status.idle":"2023-05-10T19:46:44.043062Z","shell.execute_reply.started":"2023-05-10T19:46:42.869936Z","shell.execute_reply":"2023-05-10T19:46:44.042071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"le = preprocessing.LabelEncoder()\nle.fit(X_train[\"primary_label\"])\ntrain_label = le.transform(X_train[\"primary_label\"])\nvalid_label = le.transform(X_valid[\"primary_label\"])","metadata":{"execution":{"iopub.status.busy":"2023-05-10T19:47:00.368433Z","iopub.execute_input":"2023-05-10T19:47:00.368842Z","iopub.status.idle":"2023-05-10T19:47:00.46543Z","shell.execute_reply.started":"2023-05-10T19:47:00.368811Z","shell.execute_reply":"2023-05-10T19:47:00.464565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train[\"label\"] = train_label\nX_valid[\"label\"] = valid_label\n\nX_train.reset_index(inplace=True)\nX_valid.reset_index(inplace=True)\n\nX_train.to_csv(\"Splited_train.csv\", index=False)\nX_valid.to_csv(\"Splited_valid.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T19:47:48.661554Z","iopub.execute_input":"2023-05-10T19:47:48.661995Z","iopub.status.idle":"2023-05-10T19:47:52.826183Z","shell.execute_reply.started":"2023-05-10T19:47:48.661956Z","shell.execute_reply":"2023-05-10T19:47:52.824921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data loader","metadata":{}},{"cell_type":"code","source":"class Bird_Dataset(Dataset):\n    def __init__(self, csv_file, root_dir, mode=\"train\", duration_sec=5, transform=None, transform_Aug=None):\n        \"\"\"\n        Args:\n            csv_file (string): Path to the csv file with annotations.\n            root_dir (string): Directory with all the images.\n            transform (callable, optional): Optional transform to be applied\n                on a sample.\n        \"\"\"\n        self.Bird_audios = pd.read_csv(csv_file)\n        self.root_dir = root_dir\n        self.transform = transform\n        self.transform_Aug = transform_Aug\n        self.duration_sec=duration_sec\n        self.mode = mode\n        \n             \n    def __len__(self):\n        return len(self.Bird_audios)\n    \n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        audio_path = os.path.join(self.root_dir,\n                                str(self.Bird_audios.loc[idx, \"filename\"]))\n        \n        # Load the audio signal\n        waveform, sample_rate = torchaudio.load(audio_path)\n        \n        # Resample the Audio\n        waveform = torchaudio.functional.resample(waveform, orig_freq=sample_rate, new_freq=22050)\n\n        # Clip the audio\n        start_sample = int(self.Bird_audios.loc[idx, \"start\"] * config[\"sample_rate\"])\n        end_sample = int(self.Bird_audios.loc[idx, \"end\"] * config[\"sample_rate\"])\n        waveform = waveform[:, start_sample:end_sample]\n        \n#         # Padd if the shorter less than duration_sec\n#         target_frames = int(self.duration_sec * config[\"sample_rate\"])\n#         pad_transform = torchaudio.transforms.PadTrim(target_frames)\n#         waveform = pad_transform(waveform)\n        \n        # Compute the spectrogram\n        spec_transform = torchaudio.transforms.MelSpectrogram(\n                        n_fft=800,\n                        hop_length=320,\n                        n_mels=128\n                    )\n        specgram = spec_transform(waveform)\n        \n        specgram = torchaudio.transforms.AmplitudeToDB()(specgram)\n        resize_transform = Resize((128,224))\n        specgram = resize_transform(specgram)\n\n#         # Define the learnable parameter alpha\n#         alpha = torch.nn.Parameter(torch.tensor([1.0]))\n\n#         # Apply exponential transformation with alpha\n#         exp_specgram = torch.exp(alpha * specgram)\n\n#         # If alpha is a tensor, apply different values for each mel band\n#         if alpha.dim() == 1:\n#             exp_specgram = exp_specgram * alpha.view(1, -1, 1)\n\n        specgram = torch.cat([specgram, specgram, specgram], dim=0)\n      \n        label =  self.Bird_audios.loc[idx, \"label\"]\n        return (specgram, label ) if self.mode == \"train\" else specgram ","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:32:49.359879Z","iopub.execute_input":"2023-05-10T20:32:49.360673Z","iopub.status.idle":"2023-05-10T20:32:49.373828Z","shell.execute_reply.started":"2023-05-10T20:32:49.36064Z","shell.execute_reply":"2023-05-10T20:32:49.372825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_Dataset = Bird_Dataset(csv_file=\"/kaggle/working/Splited_train.csv\",\n                     root_dir=\"/kaggle/input/birdclef-2023/train_audio\")\n\nvalid_Dataset = Bird_Dataset(csv_file=\"/kaggle/working/Splited_train.csv\",\n                     root_dir=\"/kaggle/input/birdclef-2023/train_audio\")","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:32:51.938864Z","iopub.execute_input":"2023-05-10T20:32:51.939236Z","iopub.status.idle":"2023-05-10T20:32:53.633075Z","shell.execute_reply.started":"2023-05-10T20:32:51.939207Z","shell.execute_reply":"2023-05-10T20:32:53.632172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = torch.utils.data.DataLoader(train_Dataset, batch_size=128, shuffle=True, num_workers=8, pin_memory=True)\nvalid_loader = torch.utils.data.DataLoader(valid_Dataset, batch_size=128, shuffle=True, num_workers=8, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:33:17.809986Z","iopub.execute_input":"2023-05-10T20:33:17.810548Z","iopub.status.idle":"2023-05-10T20:33:17.81564Z","shell.execute_reply.started":"2023-05-10T20:33:17.810507Z","shell.execute_reply":"2023-05-10T20:33:17.814772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# show item from dataset\nx = None\nfor image, label in train_Dataset:\n#     specgram_db = torchaudio.transforms.AmplitudeToDB()(image)\n   \n\n    print(image.shape)\n    # plot the spectrogram\n    plt.imshow(image[0,:,:].numpy())\n    plt.xlabel('Time')\n    plt.ylabel('Frequency')\n    plt.show()\n\n    break","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:33:21.702794Z","iopub.execute_input":"2023-05-10T20:33:21.70345Z","iopub.status.idle":"2023-05-10T20:33:22.374584Z","shell.execute_reply.started":"2023-05-10T20:33:21.703416Z","shell.execute_reply":"2023-05-10T20:33:22.373799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_Type='resnet50'\nfeature= 'baseline_model'  \n\nrd = np.random.randint(100000)\nmodelName = f\"{model_Type}_{feature}_{rd}\"\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:33:26.263392Z","iopub.execute_input":"2023-05-10T20:33:26.26376Z","iopub.status.idle":"2023-05-10T20:33:26.289742Z","shell.execute_reply.started":"2023-05-10T20:33:26.263732Z","shell.execute_reply":"2023-05-10T20:33:26.288808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_pretrained_model(model_name):\n\n    if model_name==\"ViT\":\n        \n        model = torch.hub.load('facebookresearch/deit:main', 'deit_tiny_patch16_224', pretrained=True)\n        for param in model.parameters(): #freeze model\n            param.requires_grad = False\n        n_inputs = model.head.in_features\n        model.head = nn.Sequential(\n                    nn.Linear(n_inputs, 512),\n                    nn.ReLU(),\n                    nn.Dropout(0.3),\n                    nn.Linear(512, 7)\n        )\n        \n    if model_name==\"MaxViT\":\n        model = timm.create_model(\"maxvit_tiny_rw_224\", pretrained=False,img_size = 224)\n        model.head.fc = nn.Linear(512, 7, bias=True)\n        \n        \n        \n    if model_name == \"vgg16\":\n        model = models.vgg16(pretrained=True)\n\n        n_inputs = model.classifier[6].in_features\n        model.classifier[6] = nn.Sequential(\n            nn.Linear(n_inputs, 2048),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(2048, 4),\n            nn.LogSoftmax(dim=1)\n        )\n\n       \n    if model_name == \"resnet50\":\n        model = models.resnet50(pretrained=True)\n        n_inputs = model.fc.in_features\n        model.fc = nn.Sequential(\n            nn.Linear(n_inputs, config[\"classes_number\"]),\n            nn.LogSoftmax(dim=1),\n        )\n\n\n   \n    if model_name == \"alexnet\":\n        model = models.AlexNet()\n\n        model.classifier[-1]= nn.Sequential(\n            nn.Linear(4096, 7),\n            nn.LogSoftmax(dim=1)\n        )\n\n    # Move to gpu and parallelize\n    model = model.to(device)\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:33:29.945299Z","iopub.execute_input":"2023-05-10T20:33:29.945837Z","iopub.status.idle":"2023-05-10T20:33:29.963515Z","shell.execute_reply.started":"2023-05-10T20:33:29.9458Z","shell.execute_reply":"2023-05-10T20:33:29.962694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = get_pretrained_model(model_Type)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:33:30.309177Z","iopub.execute_input":"2023-05-10T20:33:30.309866Z","iopub.status.idle":"2023-05-10T20:33:34.794585Z","shell.execute_reply.started":"2023-05-10T20:33:30.309832Z","shell.execute_reply":"2023-05-10T20:33:34.793628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(\n    model,\n    criterion,\n    optimizer,\n    train_loader,\n    valid_loader,\n    save_file_name,\n    max_epochs_stop=3,\n    n_epochs=20,\n    print_every=1\n):\n    \"\"\"Train a PyTorch Model\n\n    Params\n    --------\n        model (PyTorch model): cnn to train\n        criterion (PyTorch loss): objective to minimize\n        optimizer (PyTorch optimizier): optimizer to compute gradients of model parameters\n        train_loader (PyTorch dataloader): training dataloader to iterate through\n        valid_loader (PyTorch dataloader): validation dataloader used for early stopping\n        save_file_name (str ending in '.pt'): file path to save the model state dict\n        max_epochs_stop (int): maximum number of epochs with no improvement in validation loss for early stopping\n        n_epochs (int): maximum number of training epochs\n        print_every (int): frequency of epochs to print training stats\n\n    Returns\n    --------\n        model (PyTorch model): trained cnn with best weights\n        history (DataFrame): history of train and validation loss and accuracy\n    \"\"\"\n\n    # Early stopping intialization\n    epochs_no_improve = 0\n    valid_loss_min = np.Inf\n    valid_acc_max = 0\n    valid_f1_max = 0 \n    valid_max_acc = 0\n    history = []\n\n    # Number of epochs already trained (if using loaded in model weights)\n    try:\n        print(f\"Model has been trained for: {model.epochs} epochs.\\n\")\n    except:\n        model.epochs = 0\n        print(f\"Starting Training from Scratch.\\n\")\n\n    overall_start = timer()\n\n    # Main loop\n    for epoch in range(n_epochs):\n\n        # keep track of training and validation loss each epoch\n        train_loss = 0.0\n        valid_loss = 0.0\n\n        train_acc = 0\n        valid_acc = 0\n        Train_f1_sc = 0\n        Valid_f1_sc = 0 \n        \n\n        # Set to training\n        model.train()\n        start = timer()\n  \n        # Training loop\n        for ii, (data, target) in enumerate(train_loader):\n            # Tensors to gpu\n            \n            if train_on_gpu:\n                data, target = data.cuda(), target.cuda()\n\n            # Clear gradients\n            optimizer.zero_grad()\n            # Predicted outputs are log probabilities\n            output = model(data)    #  .sigmoid()\n#             target = target.unsqueeze(1)\n        \n            # Loss and backpropagation of gradients\n            loss = criterion(output, target)\n            loss.backward()\n\n            # Update the parameters\n            optimizer.step()\n            \n            # Track train loss by multiplying average loss by number of examples in batch\n            train_loss += loss.item() * data.size(0)\n    \n            # Calculate accuracy by finding max log probability\n            _, pred = torch.max(output, dim=1)\n#             pred = torch.ge(output, 0.35)\n            \n            correct_tensor = pred.eq(target.data.view_as(pred))\n            # Need to convert correct tensor from int to float to average\n            accuracy = torch.mean(correct_tensor.type(torch.FloatTensor))\n            # Multiply average accuracy times the number of examples in batch\n            \n            Train_f1_sc = f1_score(target.cpu().data, pred.cpu(), average = \"macro\")\n            \n            \n            train_acc += accuracy.item() * data.size(0)\n\n            # Track training progress\n            print(\n                f\"Epoch: {epoch}\\t{100 * (ii + 1) / len(train_loader):.2f}% complete. {timer() - start:.2f} seconds elapsed in epoch.\",\n                end=\"\\r\",\n            )\n      \n\n        # After training loops ends, start validation\n        else:\n            model.epochs += 1\n\n            # Don't need to keep track of gradients\n            with torch.no_grad():\n                # Set to evaluation mode\n                model.eval()\n\n                # Validation loop\n                for data, target in valid_loader:\n                    # Tensors to gpu\n                    if train_on_gpu:\n                        data, target = data.cuda(), target.cuda()\n\n                    # Forward pass\n#                     output = model(data)\n\n                    # Validation loss\n                    output = model(data)    #.sigmoid()\n#                     target = target.unsqueeze(1)\n\n                    # Loss and backpropagation of gradients\n                    loss = criterion(output, target)\n                    # Multiply average loss times the number of examples in batch\n                    valid_loss += loss.item() * data.size(0)\n\n                    # Calculate validation accuracy\n                    _, pred = torch.max(output, dim=1)\n#                     pred = torch.ge(output, 0.35)\n\n                    correct_tensor = pred.eq(target.data.view_as(pred))\n                    accuracy = torch.mean(correct_tensor.type(torch.FloatTensor))\n                    # Multiply average accuracy times the number of examples\n                    valid_acc += accuracy.item() * data.size(0)\n                    \n                    Valid_f1_sc = f1_score(target.cpu().data, pred.cpu(), average = \"macro\")\n\n                    \n                    \n                    \n                    \n                # Calculate average losses\n                train_loss = train_loss / len(train_loader.dataset)\n                valid_loss = valid_loss / len(valid_loader.dataset)\n\n                # Calculate average accuracy\n                train_acc = train_acc / len(train_loader.dataset)\n                valid_acc = valid_acc / len(valid_loader.dataset)\n\n                history.append([train_loss, valid_loss, train_acc, valid_acc, Train_f1_sc, Valid_f1_sc])\n\n                # Print training and validation results\n                if (epoch + 1) % print_every == 0:\n                    print(\n                        f\"\\nEpoch: {epoch} \\tTraining Loss: {train_loss:.4f} \\tValidation Loss: {valid_loss:.4f}\"\n                    )\n                    print(\n                        f\"\\t\\tTraining Accuracy: {100 * train_acc:.2f}%\\t Train F1 score: {100 * Train_f1_sc:.2f}%\\t Validation Accuracy: {100 * valid_acc:.2f}%\\t Validation F1 score: {100 * Valid_f1_sc:.2f}%\"\n                    )\n\n                # Save the model if validation loss decreases\n                if valid_loss < valid_loss_min:\n#                 if Valid_f1_sc > valid_f1_max:\n                    # Save model\n                    torch.save(model.state_dict(), save_file_name)\n                    # Track improvement\n                    epochs_no_improve = 0\n                    valid_loss_min = valid_loss\n                    valid_acc_max = valid_acc\n                    valid_f1_max = Valid_f1_sc\n\n                    # valid_best_acc = valid_acc\n                    best_epoch = epoch\n\n                # Otherwise increment count of epochs with no improvement\n                elif   valid_loss >= valid_loss_min:\n                    epochs_no_improve += 1\n                    # Trigger early stopping\n                    if epochs_no_improve >= max_epochs_stop:\n                        print(\n                            f\"\\nEarly Stopping! Total epochs: {epoch}. Best epoch: {best_epoch} with loss: {valid_loss_min:.2f} and acc: {100 * valid_acc_max:.2f}%\"\n                        )\n                        total_time = timer() - overall_start\n                        print(\n                            f\"{total_time:.2f} total seconds elapsed. {total_time / (epoch+1):.2f} seconds per epoch.\"\n                        )\n\n                        # Load the best state dict\n                        model.load_state_dict(torch.load(save_file_name))\n                        # Attach the optimizer\n                        model.optimizer = optimizer\n\n                        # Format history\n                        history = pd.DataFrame(\n                            history,\n                            columns=[\n                                \"train_loss\",\n                                \"valid_loss\",\n                                \"train_acc\",\n                                \"valid_acc\",\n                                \"train_f1\",\n                                \"valid_f1\",\n                            ],\n                        )\n                        return model, history\n\n    # Attach the optimizer\n    model.optimizer = optimizer\n    # Record overall time and print out stats\n    total_time = timer() - overall_start\n    print(\n        f\"\\nBest epoch: {best_epoch} with loss: {valid_loss_min:.2f} and acc: {100 * valid_acc:.2f}%\"\n    )\n    print(\n        f\"{total_time:.2f} total seconds elapsed. {total_time / (epoch):.2f} seconds per epoch.\"\n    )\n    # Format history\n    history = pd.DataFrame(\n        history, columns=[\"train_loss\", \"valid_loss\", \"train_acc\", \"valid_acc\", \"train_f1\", \"valid_f1\"]\n    )\n    return model, history","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:33:34.796581Z","iopub.execute_input":"2023-05-10T20:33:34.796955Z","iopub.status.idle":"2023-05-10T20:33:34.824592Z","shell.execute_reply.started":"2023-05-10T20:33:34.796907Z","shell.execute_reply":"2023-05-10T20:33:34.823532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# criterion = nn.CrossEntropyLoss(weight =  torch.FloatTensor(class_w).to(device))\ncriterion = nn.CrossEntropyLoss()\n\n# criterion =torch.nn.BCEWithLogitsLoss(pos_weight =  torch.tensor(weight_for_1).to(device))\n\n# criterion =torch.nn.BCEWithLogitsLoss()\n\n# criterion = LabelSmoothingCrossEntropy(weight = class_w)\ncriterion = criterion.to(\"cuda\")\n\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-5)\n# optimizer = optim.Adam(model.parameters(), lr=0.003)\n\n# exp_lr_scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=4, gamma=0.97)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:33:34.825763Z","iopub.execute_input":"2023-05-10T20:33:34.826638Z","iopub.status.idle":"2023-05-10T20:33:34.841031Z","shell.execute_reply.started":"2023-05-10T20:33:34.826606Z","shell.execute_reply":"2023-05-10T20:33:34.840072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model, history = train(\n    model,\n    criterion,\n    optimizer,\n    train_loader,\n    valid_loader,\n    save_file_name=f\"{modelName}.pt\",\n    max_epochs_stop=4,\n    n_epochs=10,\n    print_every=1\n)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:33:37.026078Z","iopub.execute_input":"2023-05-10T20:33:37.027015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device","metadata":{"execution":{"iopub.status.busy":"2023-05-10T20:18:24.57727Z","iopub.execute_input":"2023-05-10T20:18:24.5784Z","iopub.status.idle":"2023-05-10T20:18:24.584949Z","shell.execute_reply.started":"2023-05-10T20:18:24.578353Z","shell.execute_reply":"2023-05-10T20:18:24.583919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}