{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":21154,"databundleVersionId":1243559,"sourceType":"competition"}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Transfer Learning - EfficientNet\n\nThe goal of this competition is to develop a model capable of accurately predicting the flower depicted in an image.\n\n","metadata":{}},{"cell_type":"markdown","source":"\n   ## Contents:\n* <h1 style=\"padding: 1rem;\n          color:black;\n          text-align:left;\n          margin:0 auto;\n          font-size:1.5rem;\"><a href=\"#library\">Libraries</a></h1>\n* <h1 style=\"padding: 1rem;\n          color:black;\n          text-align:left;\n          margin:0 auto;\n          font-size:1.5rem;\"><a href=\"#transform\">Dataset Transform and Dataloader</a></h1>\n* <h1 style=\"padding: 1rem;\n          color:black;\n          text-align:left;\n          margin:0 auto;\n          font-size:1.5rem;\"><a href=\"#efficient\">Transfer Learning On EfficientNet Model</a></h1>\n               \n* <h1 style=\"padding: 1rem;\n          color:black;\n          text-align:left;\n          margin:0 auto;\n          font-size:1.5rem;\"><a href=\"#training\">Configs and Training Loop</a></h1>\n\n* <h1 style=\"padding: 1rem;\n          color:black;\n          text-align:left;\n          margin:0 auto;\n          font-size:1.5rem;\"><a href=\"#submission\">Inference and Submission</a></h1>","metadata":{}},{"cell_type":"markdown","source":"<div class=\"alert alert-warning\" role=\"alert\">\n<a id=\"library\"><h1 style=\"padding: 2rem;\n          color:black;\n          text-align:center;\n          margin:0 auto;\n          font-size:2rem;\">\n Libraries\n</h1></a>\n</div>","metadata":{}},{"cell_type":"code","source":"!pip install tfrecord\n!pip install --upgrade scipy\n\nimport time\nimport io\nimport pandas as pd\nimport tensorflow as tf\nimport timm\nimport torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm.auto import tqdm\nimport pandas as pd\nimport tfrecord\nimport os\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom torch import nn\nimport torchvision\nfrom torchinfo import summary\nfrom PIL import Image\nimport warnings\nwarnings.simplefilter('ignore')\n%matplotlib inline\nprint(f\"Make sure that the PyTorch version is the same as yours.\")\nprint(f\"PyTorch version: {torch.__version__}\\ntorchvision version: {torchvision.__version__}\")","metadata":{"execution":{"iopub.status.busy":"2023-11-27T16:14:01.530788Z","iopub.execute_input":"2023-11-27T16:14:01.531562Z","iopub.status.idle":"2023-11-27T16:14:52.263466Z","shell.execute_reply.started":"2023-11-27T16:14:01.531527Z","shell.execute_reply":"2023-11-27T16:14:52.262497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FLOWER_NAMES = ['pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',         # 00 - 09\n           'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily', # 10 - 19\n           'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',         # 20 - 29\n           'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',           # 30 - 39\n           'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',      # 40 - 49\n           'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',    # 50 - 59\n           'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',            # 60 - 69\n           'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',             # 70 - 79\n           'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',            # 80 - 89\n           'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',        # 90 - 99\n           'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose']                                                                                                                                               # 100 - 102\n","metadata":{"execution":{"iopub.status.busy":"2023-11-27T16:16:09.29233Z","iopub.execute_input":"2023-11-27T16:16:09.293187Z","iopub.status.idle":"2023-11-27T16:16:09.302405Z","shell.execute_reply.started":"2023-11-27T16:16:09.293158Z","shell.execute_reply":"2023-11-27T16:16:09.301459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nprint(f\"Using device: {device}\")","metadata":{"execution":{"iopub.status.busy":"2023-11-27T16:16:12.208726Z","iopub.execute_input":"2023-11-27T16:16:12.209485Z","iopub.status.idle":"2023-11-27T16:16:12.238069Z","shell.execute_reply.started":"2023-11-27T16:16:12.209453Z","shell.execute_reply":"2023-11-27T16:16:12.237155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-warning\" role=\"alert\">\n<a id=\"transform\"><h1 style=\"padding: 2rem;\n          color:black;\n          text-align:center;\n          margin:0 auto;\n          font-size:2rem;\">\n Dataset Transform\n</h1></a>\n</div>","metadata":{}},{"cell_type":"code","source":"DATASET_PATH = \"/kaggle/input/tpu-getting-started/tfrecords-jpeg-224x224\"\n\ndef transform_tf_to_df(subset_data):\n    df = pd.DataFrame({\"id\": pd.Series(dtype=\"str\"), \n                       \"class\": pd.Series(dtype=\"int\"), \n                       \"img\": pd.Series(dtype=\"object\")})    \n    tf_files = []\n    \n    for subdir, dirs, files in os.walk(DATASET_PATH):\n        if subdir.split(\"/\")[-1] == subset_data:\n            for file in files:\n                filepath = subdir + os.sep + file\n                tf_files.append(filepath)\n\n    for tf_file in tf_files:\n        if subset_data == \"test\":\n            loader = tfrecord.tfrecord_loader(tf_file, None, {\"id\": \"byte\", \"image\": \"byte\"})\n        else:\n            loader = tfrecord.tfrecord_loader(tf_file, None, {\"id\": \"byte\",\"image\": \"byte\", \"class\": \"int\"})\n\n        for record in loader:\n            id_label = record[\"id\"].decode('utf-8')\n            label = record[\"class\"][0].item() if subset_data != \"test\" else None\n            img_bytes = np.frombuffer(record[\"image\"], dtype=np.uint8)\n            img = cv2.imdecode(img_bytes, cv2.IMREAD_COLOR)\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n            df.loc[len(df.index)] = [id_label, label, img]\n    return df\n\ndf_validation = transform_tf_to_df('val')\ndf_train = transform_tf_to_df('train')\ndf_test = transform_tf_to_df('test')\n\nprint(df_train.dtypes, df_train.shape)\nplt.imshow(df_train.iloc[2]['img'])","metadata":{"execution":{"iopub.status.busy":"2023-11-27T16:16:14.717408Z","iopub.execute_input":"2023-11-27T16:16:14.71812Z","iopub.status.idle":"2023-11-27T16:18:16.873773Z","shell.execute_reply.started":"2023-11-27T16:16:14.71809Z","shell.execute_reply":"2023-11-27T16:18:16.872791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CLASSES_NUM = len(FLOWER_NAMES)\nCLASSES_NUM","metadata":{"execution":{"iopub.status.busy":"2023-11-27T16:18:16.875458Z","iopub.execute_input":"2023-11-27T16:18:16.875752Z","iopub.status.idle":"2023-11-27T16:18:16.881834Z","shell.execute_reply.started":"2023-11-27T16:18:16.875726Z","shell.execute_reply":"2023-11-27T16:18:16.880933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FlowersDataset(Dataset):\n    def __init__(self, data, transform=None) -> None:\n        self.data = data\n        self.transform = transform\n        \n    def __getitem__(self, idx):\n        \"Iterable function which applies to each row\"\n        img_id = self.data.iloc[idx, 0]\n        label = self.data.iloc[idx, 1]\n        image = self.data.iloc[idx, 2]\n        image = Image.fromarray(image)\n        if self.transform:\n            image = self.transform(image)\n        y = np.zeros(CLASSES_NUM, dtype=np.float32)\n        y[label] = int(1)\n        return img_id, y, image\n    \n    def __len__(self) -> int:\n        \"Returns the total number of samples.\"\n        return len(self.data)","metadata":{"execution":{"iopub.status.busy":"2023-11-27T16:18:36.628607Z","iopub.execute_input":"2023-11-27T16:18:36.628959Z","iopub.status.idle":"2023-11-27T16:18:36.636481Z","shell.execute_reply.started":"2023-11-27T16:18:36.628934Z","shell.execute_reply":"2023-11-27T16:18:36.635481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# weights = torchvision.models.EfficientNet_B0_Weights.DEFAULT\n# auto_transforms = weights.transforms()\n# print(auto_transforms)\n\ndef rgb2hsl_torch(rgb: torch.Tensor) -> torch.Tensor:\n    cmax, cmax_idx = torch.max(rgb, dim=0, keepdim=True)\n    cmin = torch.min(rgb, dim=0, keepdim=True)[0]\n    delta = cmax - cmin\n    hsl_h = torch.empty_like(rgb[0:1, :, :])\n    cmax_idx[delta == 0] = 3\n    hsl_h[cmax_idx == 0] = (((rgb[1:2] - rgb[2:3]) / delta) % 6)[cmax_idx == 0]\n    hsl_h[cmax_idx == 1] = (((rgb[2:3] - rgb[0:1]) / delta) + 2)[cmax_idx == 1]\n    hsl_h[cmax_idx == 2] = (((rgb[0:1] - rgb[1:2]) / delta) + 4)[cmax_idx == 2]\n    hsl_h[cmax_idx == 3] = 0.\n    hsl_h /= 6.\n\n    hsl_l = (cmax + cmin) / 2.\n    hsl_s = torch.empty_like(hsl_h)\n    hsl_s[hsl_l == 0] = 0\n    hsl_s[hsl_l == 1] = 0\n    hsl_l_ma = torch.bitwise_and(hsl_l > 0, hsl_l < 1)\n    hsl_l_s0_5 = torch.bitwise_and(hsl_l_ma, hsl_l <= 0.5)\n    hsl_l_l0_5 = torch.bitwise_and(hsl_l_ma, hsl_l > 0.5)\n    hsl_s[hsl_l_s0_5] = ((cmax - cmin) / (hsl_l * 2.))[hsl_l_s0_5]\n    hsl_s[hsl_l_l0_5] = ((cmax - cmin) / (- hsl_l * 2. + 2.))[hsl_l_l0_5]\n    return torch.cat([rgb, hsl_h, hsl_s, hsl_l], dim=0)\n\nstats = ((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))\ndata_transform = torchvision.transforms.Compose(\n                        [\n                        torchvision.transforms.RandomHorizontalFlip(),\n                        torchvision.transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1, hue=0.1),\n                        torchvision.transforms.ToTensor(), \n                        torchvision.transforms.Normalize(*stats, inplace=True),\n                        torchvision.transforms.Lambda(rgb2hsl_torch)\n                        ])\nauto_transforms = data_transform\n\ntrain_data = FlowersDataset(df_train, transform=auto_transforms)\nvalidation_data = FlowersDataset(df_validation, transform=auto_transforms)\ntest_data = FlowersDataset(df_test, transform=auto_transforms)","metadata":{"execution":{"iopub.status.busy":"2023-11-27T16:19:47.298829Z","iopub.execute_input":"2023-11-27T16:19:47.29919Z","iopub.status.idle":"2023-11-27T16:19:47.312588Z","shell.execute_reply.started":"2023-11-27T16:19:47.299162Z","shell.execute_reply":"2023-11-27T16:19:47.311629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 256\n\ntrain_dataloader = DataLoader(dataset=train_data,\n                              batch_size=batch_size,\n                              num_workers=2)\n\nvalidation_dataloader = DataLoader(dataset=validation_data, \n                              batch_size=batch_size,\n                              num_workers=1)\n\ntest_dataloader = DataLoader(dataset=test_data, \n                             batch_size=batch_size, \n                             num_workers=1, \n                             shuffle=False)\n\nprint(f\"Length of training dataloader: {len(train_dataloader)} batches of {64}\")\nprint(f\"Length of validation dataloader: {len(validation_dataloader)} batches of {64}\")\nprint(f\"Length of test_dataloader dataloader: {len(test_dataloader)} batches of {64}\")\nimage_id, label, image = next(iter(validation_dataloader))\nprint(f\"Data shape: {image.shape}, labels shape: {label.shape}\")\ndata_shape = image.shape","metadata":{"execution":{"iopub.status.busy":"2023-11-27T16:19:50.159097Z","iopub.execute_input":"2023-11-27T16:19:50.159499Z","iopub.status.idle":"2023-11-27T16:19:53.983409Z","shell.execute_reply.started":"2023-11-27T16:19:50.159448Z","shell.execute_reply":"2023-11-27T16:19:53.9821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-warning\" role=\"alert\">\n<a id=\"efficient\"><h1 style=\"padding: 2rem;\n          color:black;\n          text-align:center;\n          margin:0 auto;\n          font-size:2rem;\">\n Transfer Learning On EfficientNet Model\n</h1></a>\n</div>","metadata":{}},{"cell_type":"code","source":"model_cfg = {\n    \"input_channels\": data_shape[1],\n    \"input_size\": data_shape[2],\n    \"output_size\": CLASSES_NUM,\n}\n\nclass BaselineCNN(nn.Module):\n    def __init__(self, model_cfg):\n        super().__init__()\n\n        self.cnn = nn.Sequential(\n            nn.Conv2d(model_cfg[\"input_channels\"], 16, kernel_size=3, stride=1, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=2, stride=2),\n            nn.Conv2d(16, 32, kernel_size=3, stride=1, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=2, stride=2),\n            nn.Conv2d(32, 32, kernel_size=3, stride=1, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=2, stride=2),\n            nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=2, stride=2),\n            nn.Conv2d(64, 64, kernel_size=3, stride=1, padding=1),\n            nn.ReLU(),\n        )\n\n        self.output_figure_size = int(model_cfg[\"input_size\"] / 2**4)\n        self.dense = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(64 * self.output_figure_size**2, 128),\n            nn.ReLU(),\n            nn.Linear(128, 128),\n            nn.ReLU(),\n            nn.Linear(128, model_cfg[\"output_size\"]),\n        )\n\n    def forward(self, x: torch.Tensor):\n        cnn_output = self.cnn(x)\n\n        logits = self.dense(cnn_output)\n\n        return logits\n    \nmodel = BaselineCNN(model_cfg).to(device)\n\nsummary(model, \n        input_size=(batch_size, data_shape[1], 224, 224),\n        verbose=0,\n        col_names=[\"input_size\", \"output_size\", \"num_params\", \"trainable\"],\n        col_width=9,\n        row_settings=[\"var_names\"]\n)","metadata":{"execution":{"iopub.status.busy":"2023-11-27T16:19:58.74491Z","iopub.execute_input":"2023-11-27T16:19:58.745284Z","iopub.status.idle":"2023-11-27T16:19:59.145809Z","shell.execute_reply.started":"2023-11-27T16:19:58.745256Z","shell.execute_reply":"2023-11-27T16:19:59.144925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-warning\" role=\"alert\">\n<a id=\"training\"><h1 style=\"padding: 2rem;\n          color:black;\n          text-align:center;\n          margin:0 auto;\n          font-size:2rem;\">\n Configs and Training Loop\n</h1></a>\n</div>","metadata":{}},{"cell_type":"code","source":"training_loss_count, validation_loss_count  = [], []\ntraining_acc_count, validation_acc_count  = [], []\n\ndef accuracy_step(y_true, y_pred):\n    correct = torch.eq(y_true, y_pred).sum().item()\n    acc = (correct / len(y_pred)) * 100\n    return acc\n\ndef train_step(model: torch.nn.Module,\n               data_loader: torch.utils.data.DataLoader,\n               loss_fn: torch.nn.Module,\n               optimizer: torch.optim.Optimizer,\n               accuracy_step,\n               device: torch.device = device):\n    \n    train_loss, train_acc = 0, 0\n    model.to(device)\n    \n    for img_ids, y, X in data_loader:\n        X, y = X.to(device), y.to(device)\n        #forward\n        y_pred = model(X)\n        #calculate loss\n        loss = loss_fn(y_pred, y)\n        train_loss += loss\n        train_acc += accuracy_step(y_true=y.argmax(dim=1), y_pred=y_pred.argmax(dim=1))\n        #zerograd\n        optimizer.zero_grad()\n        #backward\n        loss.backward()\n        #optimizer\n        optimizer.step()\n    train_loss /= len(data_loader)\n    training_loss_count.append(train_loss.cpu().item())\n    train_acc /= len(data_loader)\n    training_acc_count.append(train_acc)\n    print(f\"Train loss: {train_loss:.5f} | Train accuracy: {train_acc:.2f}%\")\n    \n\ndef validation_step(data_loader: torch.utils.data.DataLoader,\n              model: torch.nn.Module,\n              loss_fn: torch.nn.Module,\n              accuracy_step,\n              device: torch.device = device):\n    valid_loss, valid_acc = 0, 0\n    model.to(device)\n    model.eval()\n    \n    with torch.inference_mode(): \n        for img_ids, y, X in data_loader:\n            X, y = X.to(device), y.to(device)\n            valid_pred = model(X)\n            valid_loss += loss_fn(valid_pred, y)\n            valid_acc += accuracy_step(y_true=y.argmax(dim=1), y_pred=valid_pred.argmax(dim=1))\n        valid_loss /= len(data_loader)\n        validation_loss_count.append(valid_loss.cpu().item())\n        valid_acc /= len(data_loader)\n        validation_acc_count.append(valid_acc)\n        print(f\"Validation loss: {valid_loss:.5f} | Validation accuracy: {valid_acc:.2f}%\\n\")","metadata":{"execution":{"iopub.status.busy":"2023-11-27T02:12:38.344002Z","iopub.execute_input":"2023-11-27T02:12:38.34465Z","iopub.status.idle":"2023-11-27T02:12:38.358224Z","shell.execute_reply.started":"2023-11-27T02:12:38.344599Z","shell.execute_reply":"2023-11-27T02:12:38.357277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCHS = 30\n\nloss_fn = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.003)\n# optimizer = torch.optim.Adam(model.parameters(), lr=0.01)\n\nfor epoch in tqdm(range(EPOCHS)):\n    start_time = time.perf_counter()\n    print(f\"Epoch: {epoch}\\n---------\")\n    train_step(data_loader=train_dataloader,\n        model=model,\n        loss_fn=loss_fn,\n        optimizer=optimizer,\n        accuracy_step=accuracy_step,\n        device=device)\n    \n    validation_step(data_loader=validation_dataloader,\n        model=model,\n        loss_fn=loss_fn,\n        accuracy_step=accuracy_step,\n        device=device)\n    \n    print(f\"Epoch time: {time.perf_counter() - start_time}\")\n    ","metadata":{"execution":{"iopub.status.busy":"2023-11-27T02:12:44.883141Z","iopub.execute_input":"2023-11-27T02:12:44.88392Z","iopub.status.idle":"2023-11-27T03:29:59.159996Z","shell.execute_reply.started":"2023-11-27T02:12:44.883888Z","shell.execute_reply":"2023-11-27T03:29:59.15878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"traCategoricalining_loss_count = training_loss_count[:30]\nvalidation_loss_count = validation_loss_count[:30]\ntraining_acc_count = training_acc_count[:30]\nvalidation_acc_count = validation_acc_count[:30]","metadata":{"execution":{"iopub.status.busy":"2023-11-27T02:09:10.88458Z","iopub.execute_input":"2023-11-27T02:09:10.885493Z","iopub.status.idle":"2023-11-27T02:09:10.890711Z","shell.execute_reply.started":"2023-11-27T02:09:10.885458Z","shell.execute_reply":"2023-11-27T02:09:10.889637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logs = pd.DataFrame(\n{\n    \"tr_loss\": training_loss_count,\n    \"tr_acc\": training_acc_count,\n    \"val_loss\": validation_loss_count,\n    \"val_acc\": validation_acc_count,\n})","metadata":{"execution":{"iopub.status.busy":"2023-11-27T03:36:44.764929Z","iopub.execute_input":"2023-11-27T03:36:44.765405Z","iopub.status.idle":"2023-11-27T03:36:44.77284Z","shell.execute_reply.started":"2023-11-27T03:36:44.76537Z","shell.execute_reply":"2023-11-27T03:36:44.771677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logs.to_csv(\"rgb+hsv.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-11-27T03:36:54.487497Z","iopub.execute_input":"2023-11-27T03:36:54.487864Z","iopub.status.idle":"2023-11-27T03:36:54.495731Z","shell.execute_reply.started":"2023-11-27T03:36:54.487838Z","shell.execute_reply":"2023-11-27T03:36:54.494491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(training_loss_count, label='Training Loss')\nplt.plot(validation_loss_count, label='Validation Loss')\nplt.legend(frameon=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-27T01:00:23.452132Z","iopub.execute_input":"2023-11-27T01:00:23.452901Z","iopub.status.idle":"2023-11-27T01:00:23.756872Z","shell.execute_reply.started":"2023-11-27T01:00:23.452846Z","shell.execute_reply":"2023-11-27T01:00:23.755963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-warning\" role=\"alert\">\n<a id=\"submission\"><h1 style=\"padding: 2rem;\n          color:black;\n          text-align:center;\n          margin:0 auto;\n          font-size:2rem;\">\n Inference and Submission\n</h1></a>\n</div>","metadata":{}},{"cell_type":"code","source":"submission_data = []\n\nwith torch.inference_mode():\n    for img_ids, y, X in test_dataloader:\n        X = X.to(device)\n        y_preds = model_effnet(X)\n        y_preds = y_preds.argmax(dim=1)\n        for img_id, y_pred in zip(img_ids, y_preds.cpu()):\n            submission_data.append({'id': img_id, 'label': y_pred.item()})\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = pd.DataFrame(submission_data)\nsubmission_df.to_csv('submission.csv', index=False)        \n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the number of columns for the subplot grid\nnum_cols = len(submission_df.head(7))\nfig, axes = plt.subplots(1, num_cols, figsize=(15, 5))\n\nfor i, (index, row) in enumerate(submission_df.head(7).iterrows()):\n    name = FLOWER_NAMES[row['label']]\n    selected_img = df_test.loc[df_test['id'] == row['id'], 'img'].squeeze()\n    ax = axes[i]\n    ax.imshow(selected_img)\n    ax.set_title(name)\n    ax.axis('off')\n    \nplt.subplots_adjust(wspace=0.1)\nplt.show()","metadata":{},"execution_count":null,"outputs":[]}]}