{"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":21154,"databundleVersionId":1243559,"sourceType":"competition"}],"dockerImageVersionId":30700,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\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":"56355547-281a-4169-91d1-e2228ce2290c","_cell_guid":"e8b50733-818a-4212-abeb-ab8d8f822813","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:26:51.749970Z","iopub.execute_input":"2024-05-07T13:26:51.750302Z","iopub.status.idle":"2024-05-07T13:26:51.770794Z","shell.execute_reply.started":"2024-05-07T13:26:51.750276Z","shell.execute_reply":"2024-05-07T13:26:51.769915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import io\nimport glob\nimport matplotlib.pyplot as plt\nimport seaborn as sns; sns.set()\nfrom PIL import Image\nfrom sklearn.metrics import f1_score, classification_report, confusion_matrix\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport cv2\nfrom torchvision import models\nfrom collections import Counter\n\nimport tensorflow as tf\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.optim.lr_scheduler import CosineAnnealingLR","metadata":{"_uuid":"bb320642-1b2c-47c7-87c5-212f7a7eb4db","_cell_guid":"58554dbb-64fe-4e40-8c56-f01cee40cacd","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:26:51.772426Z","iopub.execute_input":"2024-05-07T13:26:51.772710Z","iopub.status.idle":"2024-05-07T13:26:51.779841Z","shell.execute_reply.started":"2024-05-07T13:26:51.772687Z","shell.execute_reply":"2024-05-07T13:26:51.778746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_files = '/kaggle/input/tpu-getting-started/tfrecords-jpeg-224x224/train/*.tfrec'\nvalid_files = '/kaggle/input/tpu-getting-started/tfrecords-jpeg-224x224/val/*.tfrec'\ntest_files  = '/kaggle/input/tpu-getting-started/tfrecords-jpeg-224x224/test/*.tfrec'\ndevice      = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # hardware\nn_epochs    = 15                                                            # number of training epochs\nbatch_size  = 20                                                           # training batch size\nval_batch_size = 16                                                         # validation batch size\nnum_prints  = 10                                                            # number of losses to print per epoch\ntrain_size  = 12753                                                        # number of training data samples\nprint_freq  = train_size // (batch_size * num_prints) + 1                  # print if iteration is a multiple of this\ncheck_freq  = 1                                                            # save model if epoch is a multiple of this","metadata":{"_uuid":"867aea00-01bd-4c59-b1dd-d3b8eb60a191","_cell_guid":"d5e57760-0552-4e30-93fe-66423386cd3c","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:26:51.781026Z","iopub.execute_input":"2024-05-07T13:26:51.781332Z","iopub.status.idle":"2024-05-07T13:26:51.790823Z","shell.execute_reply.started":"2024-05-07T13:26:51.781310Z","shell.execute_reply":"2024-05-07T13:26:51.790066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def tfdf_to_pddf(file_path, test=False):\n    \n    \"\"\"\n    Parse TFRecord files into a Pandas DataFrame.\n    \n    Parameters:\n    file_path (str): Filepath pattern to read TFRecord files.\n    test (bool): If True, parse files as testing data, which lacks the 'class' label.\n    \n    Returns:\n    pd.DataFrame: Contains columns for 'id', 'image', and optionally 'label' if not test data.\n    \"\"\"\n    \n    def parse(pb, test=False):\n        # Define the features to be extracted from the TFRecord.\n        features = {\n            'id': tf.io.FixedLenFeature([], tf.string),\n            'image': tf.io.FixedLenFeature([], tf.string)\n        }\n        if not test:\n            features['class'] = tf.io.FixedLenFeature([], tf.int64)\n        \n        # Parse the example.\n        return tf.io.parse_single_example(pb, features)\n\n    # Initialize dictionary for DataFrame construction.\n    df = {'id': [], 'img': []}\n    if not test:\n        df['lab'] = []\n\n    # Create a dataset from the file pattern.\n    dataset = tf.data.TFRecordDataset(glob.glob(file_path))\n    \n    # Process each record into numpy arrays.\n    for sample in dataset.map(lambda pb: parse(pb, test)):\n        df['id'].append(sample['id'].numpy().decode('utf-8'))\n        df['img'].append(sample['image'].numpy())\n        if not test:\n            df['lab'].append(sample['class'].numpy())\n            \n    return pd.DataFrame(df)","metadata":{"_uuid":"d6f76b64-b8d9-4581-8366-a93a8af58c2e","_cell_guid":"3d966e15-f00e-4f6e-9b95-93695b28702a","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:26:51.792768Z","iopub.execute_input":"2024-05-07T13:26:51.793046Z","iopub.status.idle":"2024-05-07T13:26:51.802849Z","shell.execute_reply.started":"2024-05-07T13:26:51.793025Z","shell.execute_reply":"2024-05-07T13:26:51.801977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualise(dataset, n, cols):\n    \n    '''\n    Display a grid of labelled images of flowers.\n    \n    Parameters:\n    dataset : Dataset\n        Dataset containing the flower images and labels.\n        \n    n : int\n        Number of images to display.\n        \n    cols : int\n        Number of columns in the grid.\n    '''\n    \n    # Calculate number of rows needed\n    rows = n // cols if n % cols == 0 else n // cols + 1\n    \n    # Check if there are enough images\n    if n > len(dataset):\n        print(f\"Requested number of images ({n}) exceeds dataset size ({len(dataset)}). Displaying {len(dataset)} images instead.\")\n        n = len(dataset)\n    \n    plt.figure(figsize=(2 * cols, 2 * rows))\n    \n    for i in range(n):\n        plt.subplot(rows, cols, i + 1)\n        img, lab = dataset[i]\n        # Ensure the image tensor is in the correct format\n        if img.shape[0] == 3:  # Check if channels are first\n            img = img.permute(1, 2, 0)  # Change from CxHxW to HxWxC for Matplotlib\n        plt.imshow(img.numpy())\n        plt.title(f\"Label: {lab}\")\n        plt.axis('off')\n    plt.show()","metadata":{"_uuid":"06baff85-a8fe-478e-9cb7-3af36e205666","_cell_guid":"0d12037d-e080-412e-b341-04052a6108f0","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:26:51.803958Z","iopub.execute_input":"2024-05-07T13:26:51.804239Z","iopub.status.idle":"2024-05-07T13:26:51.814450Z","shell.execute_reply.started":"2024-05-07T13:26:51.804209Z","shell.execute_reply":"2024-05-07T13:26:51.813693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TransformDataset(Dataset):\n    '''\n    Unified representation of the dataset with advanced image transformations, caching, and optional label handling.\n    '''\n    def __init__(self, files, frac=1.0, test=False, cache_data=False):\n        '''\n        Initialize the dataset with optional image caching, transformations based on scaling, and label handling.\n\n        Parameters:\n        files (list): List of file paths to the dataset.\n        frac (float): Fraction of data samples to keep, between 0 and 1.\n        test (bool): If true, the dataset contains testing data (no labels available).\n        cache_data (bool): Whether to cache images after the first load.\n        '''\n        super().__init__()\n        if not (0 < frac <= 1):\n            raise ValueError(\"Fraction must be between 0 and 1.\")\n\n        self.df = tfdf_to_pddf(files, test).sample(frac=frac).reset_index(drop=True)\n        self.cache_data = cache_data\n        self.data_cache = {}\n        self.test = test\n        \n        if test:\n            self.transformations = A.Compose([\n                A.Resize(300, 300),\n                A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n                ToTensorV2()\n            ])\n        else:\n            self.transformations = A.Compose([\n                A.Resize(300, 300),\n                A.RandomCrop(width=224, height=224),\n                A.HorizontalFlip(p=0.5),\n                A.RandomBrightnessContrast(p=0.2),\n                A.Rotate(limit=15),\n                A.PadIfNeeded(min_height=300, min_width=300),\n                A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n                ToTensorV2()\n            ])\n            if 'lab' in self.df.columns:\n                self.labels = self.df['lab'].values\n            else:\n                raise AttributeError(\"Label column 'lab' is missing from the DataFrame.\")\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        if idx in self.data_cache:\n            return self.data_cache[idx]\n\n        row = self.df.iloc[idx]\n        image = cv2.imdecode(np.frombuffer(row['img'], np.uint8), cv2.IMREAD_COLOR)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)  # Convert BGR to RGB\n\n        transformed_image = self.transformations(image=image)['image']\n        result = (transformed_image, row.get('lab', row.get('id')))  # Retrieve 'lab' or 'id' based on availability\n\n        if self.cache_data:\n            self.data_cache[idx] = result\n\n        return result\n    \n    def get_weights(self):\n        if hasattr(self, 'labels'):\n            class_sample_count = np.array([len(np.where(self.labels == t)[0]) for t in np.unique(self.labels)])\n            weight = 1. / class_sample_count\n            samples_weight = np.array([weight[t] for t in self.labels])\n            return torch.from_numpy(samples_weight).double()\n        else:\n            raise RuntimeError(\"Attempting to calculate weights, but 'labels' are not set.\")\n\n    def weighted_loader(self, batch_size, num_workers=0):\n        weights = self.get_weights()\n        sampler = WeightedRandomSampler(weights, num_samples=len(weights), replacement=True)\n        \n        return DataLoader(self, batch_size=batch_size, sampler=sampler, num_workers=num_workers)\n","metadata":{"_uuid":"c41bbfc8-0c2d-44a4-ad76-89445e799083","_cell_guid":"d96772c6-9583-4559-9427-0f65f8ecbbaf","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:26:51.852369Z","iopub.execute_input":"2024-05-07T13:26:51.852977Z","iopub.status.idle":"2024-05-07T13:26:51.870797Z","shell.execute_reply.started":"2024-05-07T13:26:51.852952Z","shell.execute_reply":"2024-05-07T13:26:51.869642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class EfficientNetB0(nn.Module):\n#     def __init__(self, n_classes, learnable_modules=None, pretrained=True):\n#         super().__init__()\n#         self.efficientnet_b0 = models.efficientnet_b0(pretrained=pretrained)\n#         self.efficientnet_b0.classifier[1] = nn.Linear(self.efficientnet_b0.classifier[1].in_features, n_classes)\n\n#         if learnable_modules is None:\n#             learnable_modules = [\n#                 'features.5.2', 'features.6', 'features.7', 'features.8', 'classifier'\n#             ]\n\n#         # Set requires_grad=False initially\n#         for param in self.efficientnet_b0.parameters():\n#             param.requires_grad = False\n\n#         # Enable grad for specific modules\n#         for name, module in self.efficientnet_b0.named_modules():\n#             if any(subname in name for subname in learnable_modules):\n#                 for param in module.parameters():\n#                     param.requires_grad = True\n\n#     def forward(self, x):\n#         x = self.efficientnet_b0(x)\n#         return F.log_softmax(x, dim=1)","metadata":{"_uuid":"ef001ee4-abb9-4351-94ae-0e5ce047a73f","_cell_guid":"e9f2dcd1-b2d5-418a-bcfa-f32691225d02","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:26:51.872664Z","iopub.execute_input":"2024-05-07T13:26:51.872969Z","iopub.status.idle":"2024-05-07T13:26:51.882456Z","shell.execute_reply.started":"2024-05-07T13:26:51.872939Z","shell.execute_reply":"2024-05-07T13:26:51.881651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model ensembling\n\nDefine 3 models","metadata":{}},{"cell_type":"code","source":"class EfficientNetB0_Deep(nn.Module):\n    def __init__(self, n_classes, learnable_modules=None, pretrained=True):\n        super().__init__()\n        self.efficientnet_b0 = models.efficientnet_b0(pretrained=pretrained)\n        self.efficientnet_b0.classifier[1] = nn.Linear(self.efficientnet_b0.classifier[1].in_features, n_classes)\n\n        if learnable_modules is None:\n            learnable_modules = ['features.5.2', 'features.6', 'features.7', 'features.8', 'classifier']\n\n        for param in self.efficientnet_b0.parameters():\n            param.requires_grad = False\n\n        for name, module in self.efficientnet_b0.named_modules():\n            if any(sub in name for sub in learnable_modules):\n                for param in module.parameters():\n                    param.requires_grad = True\n\n    def forward(self, x):\n        return F.log_softmax(self.efficientnet_b0(x), dim=1)\n\nclass EfficientNetB0_Selective(nn.Module):\n    def __init__(self, n_classes, learnable_modules=None, pretrained=True):\n        super().__init__()\n        self.efficientnet_b0 = models.efficientnet_b0(pretrained=pretrained)\n        self.efficientnet_b0.classifier[1] = nn.Linear(self.efficientnet_b0.classifier[1].in_features, n_classes)\n\n        if learnable_modules is None:\n            learnable_modules = ['features.5.2', 'features.6', 'features.7', 'features.8', 'classifier']\n\n        for param in self.efficientnet_b0.parameters():\n            param.requires_grad = False\n\n        for name, module in self.efficientnet_b0.named_modules():\n            if any(sub in name for sub in learnable_modules):\n                for param in module.parameters():\n                    param.requires_grad = True\n\n    def forward(self, x):\n        return F.log_softmax(self.efficientnet_b0(x), dim=1)\n\nclass EfficientNetB0_Dropout(nn.Module):\n    def __init__(self, n_classes, learnable_modules=None, dropout_rate=0.5, pretrained=True):\n        super().__init__()\n        self.efficientnet_b0 = models.efficientnet_b0(pretrained=pretrained)\n        self.dropout = nn.Dropout(p=dropout_rate)\n        self.efficientnet_b0.classifier = nn.Sequential(\n            nn.Dropout(p=dropout_rate),\n            nn.Linear(self.efficientnet_b0.classifier[1].in_features, n_classes)\n        )\n\n        if learnable_modules is None:\n            learnable_modules = ['features.5.2', 'features.6', 'features.7', 'features.8', 'classifier']\n\n        for param in self.efficientnet_b0.parameters():\n            param.requires_grad = False\n\n        for name, module in self.efficientnet_b0.named_modules():\n            if any(sub in name for sub in learnable_modules):\n                for param in module.parameters():\n                    param.requires_grad = True\n\n    def forward(self, x):\n        x = self.efficientnet_b0.features(x)\n        x = self.efficientnet_b0.avgpool(x)\n        x = torch.flatten(x, 1)\n        x = self.efficientnet_b0.classifier(x)\n        return F.log_softmax(x, dim=1)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-07T13:26:51.883523Z","iopub.execute_input":"2024-05-07T13:26:51.883810Z","iopub.status.idle":"2024-05-07T13:26:51.902939Z","shell.execute_reply.started":"2024-05-07T13:26:51.883788Z","shell.execute_reply":"2024-05-07T13:26:51.902241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_workers = 2\n\n# Training data loader\ntrain_set = TransformDataset(files=train_files, test=False)\ntrain_loader = train_set.weighted_loader(batch_size=batch_size, num_workers=num_workers)\n\n# Validation data loader\nvalid_set = TransformDataset(files=valid_files, frac=0.20, test=False)  \nvalid_loader = DataLoader(valid_set, batch_size=val_batch_size, shuffle=False, num_workers=num_workers)\n\n# Test data loader\ntest_set = TransformDataset(files=test_files, test=True)  \ntest_loader = DataLoader(test_set, batch_size=val_batch_size, shuffle=False, num_workers=num_workers)","metadata":{"_uuid":"7120c50d-ea04-49c2-a4a4-c75e30d3afb7","_cell_guid":"dabeb9c7-0f24-4b7b-b3b4-8efe23501c17","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:26:51.905479Z","iopub.execute_input":"2024-05-07T13:26:51.905897Z","iopub.status.idle":"2024-05-07T13:26:58.129123Z","shell.execute_reply.started":"2024-05-07T13:26:51.905868Z","shell.execute_reply":"2024-05-07T13:26:58.128101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Display some training images and labels:\n# ----------------------------------------\nvisualise(train_set, n = 40, cols = 10)","metadata":{"_uuid":"0df127fe-9442-4a8e-b1d5-d5a125474a98","_cell_guid":"dbf60f62-f88c-42ea-ace1-e5d6437a5ba8","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:26:58.130784Z","iopub.execute_input":"2024-05-07T13:26:58.131061Z","iopub.status.idle":"2024-05-07T13:27:03.707882Z","shell.execute_reply.started":"2024-05-07T13:26:58.131038Z","shell.execute_reply":"2024-05-07T13:27:03.706668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_deep = nn.DataParallel(EfficientNetB0_Deep(n_classes=104, learnable_modules=('features.5.2', \n                                                                         'features.6', \n                                                                         'features.7', \n                                                                         'features.8', \n                                                                         'classifier'))).to(device)\nmodel_selective = nn.DataParallel(EfficientNetB0_Selective(n_classes=104, learnable_modules=('features.5.2', \n                                                                         'features.6', \n                                                                         'features.7', \n                                                                         'features.8', \n                                                                         'classifier'))).to(device)\nmodel_dropout = nn.DataParallel(EfficientNetB0_Dropout(n_classes=104, learnable_modules=('features.5.2', \n                                                                         'features.6', \n                                                                         'features.7', \n                                                                         'features.8', \n                                                                         'classifier'))).to(device)\n\n\noptimizer_deep = torch.optim.Adam([\n    {'params': model_deep.module.efficientnet_b0.features[5][2].parameters(), 'lr': 1e-4},\n    {'params': model_deep.module.efficientnet_b0.features[6].parameters(), 'lr': 1e-4},\n    {'params': model_deep.module.efficientnet_b0.features[7].parameters(), 'lr': 1e-4},\n    {'params': model_deep.module.efficientnet_b0.features[8].parameters(), 'lr': 1e-4},\n    {'params': model_deep.module.efficientnet_b0.classifier.parameters(), 'lr': 1e-3}\n], weight_decay=1e-4)\n\noptimizer_selective = torch.optim.Adam([\n    {'params': model_selective.module.efficientnet_b0.features[5][2].parameters(), 'lr': 1e-4},\n    {'params': model_selective.module.efficientnet_b0.features[6].parameters(), 'lr': 1e-4},\n    {'params': model_selective.module.efficientnet_b0.features[7].parameters(), 'lr': 1e-4},\n    {'params': model_selective.module.efficientnet_b0.features[8].parameters(), 'lr': 1e-4},\n    {'params': model_selective.module.efficientnet_b0.classifier.parameters(), 'lr': 1e-3}\n], weight_decay=1e-4)\n\noptimizer_dropout = torch.optim.Adam([\n    {'params': model_dropout.module.efficientnet_b0.features[5][2].parameters(), 'lr': 1e-4},\n    {'params': model_dropout.module.efficientnet_b0.features[6].parameters(), 'lr': 1e-4},\n    {'params': model_dropout.module.efficientnet_b0.features[7].parameters(), 'lr': 1e-4},\n    {'params': model_dropout.module.efficientnet_b0.features[8].parameters(), 'lr': 1e-4},\n    {'params': model_dropout.module.efficientnet_b0.classifier.parameters(), 'lr': 1e-3}\n], weight_decay=1e-4)\n\n# Scheduler for learning rate adjustment\nscheduler_deep = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer_deep, T_max=n_epochs)\nscheduler_selective = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer_selective, T_max=n_epochs)\nscheduler_dropout = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer_dropout, T_max=n_epochs)\n\n# Loss function for training\nloss_fn = torch.nn.functional.nll_loss","metadata":{"_uuid":"3d6eeaa4-672b-4001-b9dc-16950f34fd3d","_cell_guid":"acbe2097-a468-4f79-a347-d52223363bf5","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:27:03.709315Z","iopub.execute_input":"2024-05-07T13:27:03.709679Z","iopub.status.idle":"2024-05-07T13:27:04.289441Z","shell.execute_reply.started":"2024-05-07T13:27:03.709650Z","shell.execute_reply":"2024-05-07T13:27:04.288628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# losses = []\n# valid_f1s = []\n# for epoch in range(n_epochs):\n#     print(f\"\\nEpoch {epoch}:\")\n#     print('-' * 10)\n    \n#     model.train()\n#     for i, (x, y) in enumerate(train_loader):\n#         x, y = x.to(device), y.to(device)\n#         optimizer.zero_grad()\n#         output = model(x)\n#         loss = loss_fn(output, y)\n#         loss.backward()\n#         optimizer.step()\n\n#         if i % print_freq == 0:\n#             print(f'Loss {i}: {loss.item():.3f}')\n#             losses.append(loss.item())\n\n#     if epoch % check_freq == 0:\n#         model.eval()\n#         valid_true_labs, valid_pred_labs = [], []\n#         with torch.no_grad():\n#             for x, y in valid_loader:\n#                 x, y = x.to(device), y.to(device)\n#                 outputs = model(x)\n#                 valid_pred_labs.extend(outputs.argmax(dim=1).tolist())\n#                 valid_true_labs.extend(y.tolist())\n\n#         valid_f1 = f1_score(valid_true_labs, valid_pred_labs, average='weighted')\n#         valid_f1s.append(valid_f1)\n#         print(f'Validation F1: {valid_f1 * 100:.2f}%')\n        \n#         # Save the model\n#         torch.save(model.state_dict(), f'./epoch{epoch // check_freq}.pth')\n        \n\n#     scheduler.step()\n\n\ndef validate_and_checkpoint(model, loader, valid_f1s, epoch, model_name):\n    model.eval()\n    valid_true_labs, valid_pred_labs = [], []\n    with torch.no_grad():\n        for x, y in loader:\n            x, y = x.to(device), y.to(device)\n            outputs = model(x)\n            valid_pred_labs.extend(outputs.argmax(dim=1).tolist())\n            valid_true_labs.extend(y.tolist())\n\n    valid_f1 = f1_score(valid_true_labs, valid_pred_labs, average='weighted')\n    valid_f1s.append(valid_f1)\n    print(f'{model_name} Validation F1: {valid_f1 * 100:.2f}%')\n\n    # Save the model\n    torch.save(model.state_dict(), f'./{model_name}_epoch{epoch // check_freq}.pth')\n\nlosses_deep = []\nlosses_selective = []\nlosses_dropout = []\nvalid_f1s_deep = []\nvalid_f1s_selective = []\nvalid_f1s_dropout = []\n\nfor epoch in range(n_epochs):\n    print(f\"\\nEpoch {epoch}:\")\n    print('-' * 10)\n\n    model_deep.train()\n    model_selective.train()\n    model_dropout.train()\n\n    for i, (x, y) in enumerate(train_loader):\n        x, y = x.to(device), y.to(device)\n\n        # Train model_deep\n        optimizer_deep.zero_grad()\n        output_deep = model_deep(x)\n        loss_deep = loss_fn(output_deep, y)\n        loss_deep.backward()\n        optimizer_deep.step()\n\n        # Train model_selective\n        optimizer_selective.zero_grad()\n        output_selective = model_selective(x)\n        loss_selective = loss_fn(output_selective, y)\n        loss_selective.backward()\n        optimizer_selective.step()\n\n        # Train model_dropout\n        optimizer_dropout.zero_grad()\n        output_dropout = model_dropout(x)\n        loss_dropout = loss_fn(output_dropout, y)\n        loss_dropout.backward()\n        optimizer_dropout.step()\n\n        if i % print_freq == 0:\n            print(f'Model Deep Loss {i}: {loss_deep.item():.3f}')\n            print(f'Model Selective Loss {i}: {loss_selective.item():.3f}')\n            print(f'Model Dropout Loss {i}: {loss_dropout.item():.3f}')\n            losses_deep.append(loss_deep.item())\n            losses_selective.append(loss_selective.item())\n            losses_dropout.append(loss_dropout.item())\n\n    validate_and_checkpoint(model_deep, valid_loader, valid_f1s_deep, epoch, 'deep')\n    validate_and_checkpoint(model_selective, valid_loader, valid_f1s_selective, epoch, 'selective')\n    validate_and_checkpoint(model_dropout, valid_loader, valid_f1s_dropout, epoch, 'dropout')\n\n    scheduler_deep.step()\n    scheduler_selective.step()\n    scheduler_dropout.step()\n    ","metadata":{"_uuid":"3ff1a9e3-17ec-4977-b406-e5e8db7e9401","_cell_guid":"fbce4155-c37c-46cd-81cd-ad03abb85039","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-07T13:27:04.290758Z","iopub.execute_input":"2024-05-07T13:27:04.291067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# optimal_epoch = np.argmax(np.array(valid_f1s)) # highest validation F1 epoch\n\noptimal_epoch_deep = np.argmax(np.array(valid_f1s_deep))\noptimal_epoch_selective = np.argmax(np.array(valid_f1s_selective))\noptimal_epoch_dropout = np.argmax(np.array(valid_f1s_dropout))","metadata":{"_uuid":"27624bdf-76cd-4ef2-ac44-f57df7e7a64d","_cell_guid":"83e7dc73-8e97-4620-bba8-bf5807bbb7ec","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plt.figure(figsize = (10, 2.5))\n# plt.subplot(1, 2, 1)\n# plt.plot(np.arange(len(losses)) / n_epochs, losses, linewidth = 2)\n# plt.xlabel('Epoch')\n# plt.ylabel('Loss')\n# plt.title('Training')\n# plt.subplot(1, 2, 2)\n# plt.plot(np.arange(len(valid_f1s)) * check_freq, valid_f1s, linewidth = 2)\n# plt.vlines(optimal_epoch * check_freq, 0, valid_f1s[optimal_epoch], colors = 'black', linestyles = 'dashed', label = f'Optimal epoch ({optimal_epoch * check_freq})')\n# plt.xlabel('Epoch')\n# plt.ylabel('Weighted F1')\n# plt.ylim(0, 1)\n# plt.title('Validation')\n# plt.legend(loc = 'lower left')\n# plt.savefig('plot.png')\n# plt.show()\n\nplt.figure(figsize=(15, 6))\n\n# Plot training losses for all models\nplt.subplot(1, 2, 1)\nplt.plot(np.arange(len(losses_deep)) / len(losses_deep) * n_epochs, losses_deep, label='Model Deep', linewidth=2)\nplt.plot(np.arange(len(losses_selective)) / len(losses_selective) * n_epochs, losses_selective, label='Model Selective', linewidth=2)\nplt.plot(np.arange(len(losses_dropout)) / len(losses_dropout) * n_epochs, losses_dropout, label='Model Dropout', linewidth=2)\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.title('Training Loss')\nplt.legend()\n\n# Plot validation F1 scores for all models\nplt.subplot(1, 2, 2)\nplt.plot(np.arange(len(valid_f1s_deep)) * check_freq, valid_f1s_deep, label='Model Deep', linewidth=2)\nplt.plot(np.arange(len(valid_f1s_selective)) * check_freq, valid_f1s_selective, label='Model Selective', linewidth=2)\nplt.plot(np.arange(len(valid_f1s_dropout)) * check_freq, valid_f1s_dropout, label='Model Dropout', linewidth=2)\nplt.xlabel('Epoch')\nplt.ylabel('Weighted F1')\nplt.title('Validation F1 Score')\nplt.legend()\n\nplt.tight_layout()\nplt.savefig('training_validation_metrics.png')\nplt.show()\n","metadata":{"_uuid":"6716cd53-19b4-42c1-9f1d-ba6e81f56605","_cell_guid":"ae3b8268-fa6e-4247-9358-8b13d4d5758e","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = nn.DataParallel(EfficientNetB0(n_classes = 104, learnable_modules = ())).to(device)\n# model.load_state_dict(torch.load(f'./epoch{optimal_epoch}.pth'))\n\n# Define each model with DataParallel for potential use of multiple GPUs\nmodel_deep = nn.DataParallel(EfficientNetB0_Deep(n_classes=104)).to(device)\nmodel_selective = nn.DataParallel(EfficientNetB0_Selective(n_classes=104)).to(device)\nmodel_dropout = nn.DataParallel(EfficientNetB0_Dropout(n_classes=104)).to(device)\n\nmodel_deep.load_state_dict(torch.load(f'./deep_epoch{optimal_epoch_deep}.pth'))\nmodel_selective.load_state_dict(torch.load(f'./selective_epoch{optimal_epoch_selective}.pth'))\nmodel_dropout.load_state_dict(torch.load(f'./dropout_epoch{optimal_epoch_dropout}.pth'))","metadata":{"_uuid":"b1d28db9-9c32-49a6-a128-d7b10b75c4b8","_cell_guid":"4e1a0cd8-afca-4957-8fe9-a47e387291a5","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set all models to evaluation mode\nmodel_deep.eval()\nmodel_selective.eval()\nmodel_dropout.eval()\n\n# Define the ensemble predict function if not already defined\ndef ensemble_predict(models, data_loader):\n    ids = []\n    preds = []\n    with torch.no_grad():\n        for x, batch_ids in data_loader:\n            x = x.to(device)  # Move input to the appropriate device\n            \n            # Collect softmax outputs from each model\n            model_outputs = [F.softmax(model(x), dim=1) for model in models]\n            # Average the probabilities across models\n            avg_preds = torch.mean(torch.stack(model_outputs), dim=0)\n            predicted_labels = avg_preds.argmax(dim=1)  # Get the index of the max log-probability\n            \n            # Extend lists\n            ids.extend(batch_ids)  # Assuming 'batch_ids' are iterable\n            preds.extend(predicted_labels.cpu().numpy())\n    \n    return ids, preds\n\n# List of models for ensemble\nmodels = [model_deep, model_selective, model_dropout]\n\n# Assuming test_loader is your DataLoader for testing data\nids, predicted_labels = ensemble_predict(models, test_loader)\n\n# Creating a DataFrame for submission\nsubmission = pd.DataFrame({\n    'id': ids,\n    'label': predicted_labels\n})\n\n# Checking if the submission has the correct number of rows\nif len(submission) != 7382:\n    print(f\"Warning: Submission length is {len(submission)}, expected 7382.\")\n\n# Save to CSV\nsubmission.to_csv('submission.csv', index=False)\n\n# Display the first few rows to confirm\nprint(submission.head())","metadata":{"_uuid":"2044bffe-e2cc-47e4-9aa7-aa8587e9e846","_cell_guid":"8e83b687-0eb4-4e1c-8832-6263c03b5cbd","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]}]}