{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"},{"sourceId":10299670,"sourceType":"datasetVersion","datasetId":6370394},{"sourceId":10325472,"sourceType":"datasetVersion","datasetId":6393139}],"dockerImageVersionId":30823,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"### config\n\ntotal_epochs = 100\nbatch_size = 128\nnum_processes = 2\nimage_size = 224\ndrop_path = 0.05\n## Loss Function - CE (but try BCE)\n# Always choose \"SGD\" for CNNs and AdamW for ViTs - SGD is Difficult to Converge || We should use LAMB with Cosine LR \n## Multi-label --> Mixup and CutMix \nLR = 5e-3\nweight_decay = 0.02\nwarmup_epoch = 5\ndropout = 0\ndrop_path = 0.05","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:32.492244Z","iopub.execute_input":"2024-12-29T14:11:32.492614Z","iopub.status.idle":"2024-12-29T14:11:32.497284Z","shell.execute_reply.started":"2024-12-29T14:11:32.492583Z","shell.execute_reply":"2024-12-29T14:11:32.49628Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import wandb\n\n\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"WANDB\")\nwandb.login(key=secret_value_0)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:32.49864Z","iopub.execute_input":"2024-12-29T14:11:32.49893Z","iopub.status.idle":"2024-12-29T14:11:39.414025Z","shell.execute_reply.started":"2024-12-29T14:11:32.498901Z","shell.execute_reply":"2024-12-29T14:11:39.413377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pytorch-lightning -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:39.415302Z","iopub.execute_input":"2024-12-29T14:11:39.415636Z","iopub.status.idle":"2024-12-29T14:11:42.647247Z","shell.execute_reply.started":"2024-12-29T14:11:39.415615Z","shell.execute_reply":"2024-12-29T14:11:42.646053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install timm -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:42.649062Z","iopub.execute_input":"2024-12-29T14:11:42.649323Z","iopub.status.idle":"2024-12-29T14:11:45.803096Z","shell.execute_reply.started":"2024-12-29T14:11:42.649297Z","shell.execute_reply":"2024-12-29T14:11:45.802121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom PIL import Image\nimport numpy as np\nimport pandas as pd","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:45.804186Z","iopub.execute_input":"2024-12-29T14:11:45.804541Z","iopub.status.idle":"2024-12-29T14:11:47.731195Z","shell.execute_reply.started":"2024-12-29T14:11:45.804505Z","shell.execute_reply":"2024-12-29T14:11:47.730521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations\ntrain_aug = albumentations.Compose(           \n    [\n        albumentations.Resize(image_size,image_size, p=1),\n\n         albumentations.ShiftScaleRotate( \n                        shift_limit=0.0625,  \n                        scale_limit=0.1,  \n                        rotate_limit=10, p=0.8 \n                    ), \n                    albumentations.OneOf( \n                        [ \n                            albumentations.RandomGamma( \n                                gamma_limit=(90, 110) \n                            ), \n                            albumentations.RandomBrightnessContrast( \n                                brightness_limit=0.1,  \n                                contrast_limit=0.1 \n                            ), \n                        ], \n                        p=0.5, \n                    ), \n        albumentations.HorizontalFlip(),\n        \n        albumentations.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225],\n            max_pixel_value=255.0,\n            p=1.0,\n        )],\n    p=1.0,\n)\n\nvalid_aug = albumentations.Compose(\n    [\n        albumentations.Resize(image_size, image_size, p=1),\n        albumentations.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225],\n            max_pixel_value=255.0,\n            p=1.0,\n        ),\n    ],\n    p=1.0,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:47.731994Z","iopub.execute_input":"2024-12-29T14:11:47.732387Z","iopub.status.idle":"2024-12-29T14:11:48.546135Z","shell.execute_reply.started":"2024-12-29T14:11:47.732364Z","shell.execute_reply":"2024-12-29T14:11:48.545238Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ImageNetDataset(torch.utils.data.Dataset):\n    def __init__(self,image_path, augmentations=None,train=True):\n        self.image_path = image_path\n        self.augmentations = augmentations\n        self.df = pd.read_csv(\"/kaggle/input/imagenet-labels/imagenet_class_labels.csv\")\n        self.valid_df = pd.read_csv(\"/kaggle/input/imagenet-validation-classes/validation_classes.csv\")\n        self.train = train\n    def __len__(self):\n        return len(self.image_path)\n    \n    def __getitem__(self,item):\n        image_path = self.image_path[item]\n        image = Image.open(image_path)\n        image = image.convert(\"RGB\") \n        image = np.asarray(image)\n\n        ## center crop 95% area\n        H,W,C = image.shape\n        image = image[int(0.04*H):int(0.96*H), int(0.04*W): int(0.96*W),:]\n\n        if self.train:\n            class_id = str(self.image_path[item].split(\"/\")[-2])\n            targets = self.df[self.df[\"Index\"]==class_id][\"ID\"].values[0]-1\n        else:\n            class_id = str(self.image_path[item].split(\"/\")[-1][:-5])\n            targets = self.valid_df[self.valid_df[\"ImageId\"]==class_id][\"LabelId\"].values[0]-1\n            \n        if self.augmentations is not None:\n            augmented = self.augmentations(image=image)\n            image = augmented[\"image\"]\n            \n        image = np.transpose(image, (2, 0, 1)).astype(np.float32)\n        \n        return {\n            \"image\": torch.tensor(image, dtype=torch.float),\n            \"targets\": torch.tensor(targets, dtype=torch.long),\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:48.547Z","iopub.execute_input":"2024-12-29T14:11:48.54738Z","iopub.status.idle":"2024-12-29T14:11:48.554524Z","shell.execute_reply.started":"2024-12-29T14:11:48.547327Z","shell.execute_reply":"2024-12-29T14:11:48.553687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from timm.data.mixup import Mixup\nmixup_args = {\n    'mixup_alpha': 0.1,\n    'cutmix_alpha': 1.0,\n    'cutmix_minmax': None,\n    'prob': 0.7,\n    'switch_prob': 0,\n    'mode': 'batch',\n    'label_smoothing': 0.1,\n    'num_classes': 1000}\nmixup_fn = Mixup(**mixup_args)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:48.555201Z","iopub.execute_input":"2024-12-29T14:11:48.555452Z","iopub.status.idle":"2024-12-29T14:11:49.537876Z","shell.execute_reply.started":"2024-12-29T14:11:48.555433Z","shell.execute_reply":"2024-12-29T14:11:49.537204Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nimport random\n\ntrain_paths = glob.glob(\"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train/*/*.JPEG\")\nvalid_paths = glob.glob(\"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val/*.JPEG\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:49.540221Z","iopub.execute_input":"2024-12-29T14:11:49.540602Z","iopub.status.idle":"2024-12-29T14:11:52.993413Z","shell.execute_reply.started":"2024-12-29T14:11:49.540578Z","shell.execute_reply":"2024-12-29T14:11:52.992488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pytorch_lightning as pl\nfrom pytorch_lightning.loggers  import WandbLogger\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:52.994682Z","iopub.execute_input":"2024-12-29T14:11:52.994983Z","iopub.status.idle":"2024-12-29T14:11:53.706673Z","shell.execute_reply.started":"2024-12-29T14:11:52.994953Z","shell.execute_reply":"2024-12-29T14:11:53.705981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom timm import create_model\nfrom torchvision import transforms, datasets\nimport pytorch_lightning as L\n# from timm.scheduler.cosine_lr import CosineLRScheduler\n\n\nclass LitClassification(L.LightningModule):\n    def __init__(self):\n        super().__init__()\n        self.model = create_model(\"resnet50\",pretrained=False, drop_path_rate=drop_path)\n        # model = torch.nn.SyncBatchNorm.convert_sync_batchnorm(model)\n        \n        self.loss_fn = torch.nn.CrossEntropyLoss()\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch):\n        images, targets = batch[\"image\"], batch[\"targets\"]\n        outputs = self.model(images)\n        loss = self.loss_fn(outputs, targets)\n        acc1, acc5 = self.__accuracy(outputs, targets, topk=(1, 5))\n        self.log(\"train_loss\", loss)\n        self.log(\"train_acc1\", acc1, on_step=True, prog_bar=True, on_epoch=True, logger=True)\n        self.log(\"train_acc5\", acc5, on_step=True, on_epoch=True, logger=True)\n        return loss\n\n    def validation_step(self, batch):\n        images, targets = batch[\"image\"], batch[\"targets\"]\n        outputs = self(images)\n        loss = self.loss_fn(outputs, targets)\n        \n        acc1, acc5 = self.__accuracy(outputs, targets, topk=(1, 5))\n        self.log(\"valid_loss\", loss)\n        self.log(\"val_acc1\", acc1, on_step=True, prog_bar=True, on_epoch=True)\n        self.log(\"val_acc5\", acc5, on_step=True, on_epoch=True)\n\n    @staticmethod\n    def __accuracy(output, target, topk=(1,)):\n        \"\"\"Computes the accuracy over the k top predictions for the specified values of k.\"\"\"\n        with torch.no_grad():\n            maxk = max(topk)\n            batch_size = target.size(0)\n\n            _, pred = output.topk(maxk, 1, True, True)\n            pred = pred.t()\n            correct = pred.eq(target.view(1, -1).expand_as(pred))\n\n            res = []\n            for k in topk:\n                correct_k = correct[:k].reshape(-1).float().sum(0, keepdim=True)\n                res.append(correct_k.mul_(100.0 / batch_size))\n            return res\n            \n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(model.parameters(),lr=LR,weight_decay=weight_decay)\n        # scheduler = CosineLRScheduler(optimizer, t_initial=total_epochs, lr_min=2e-8,\n        #               cycle_mul=1.0, cycle_decay=1.0, cycle_limit=1,\n        #               warmup_t=warmup_epoch, warmup_lr_init=1e-6, warmup_prefix=False, t_in_epochs=True,\n        #               noise_range_t=None, noise_pct=0.67, noise_std=1.0,\n        #               noise_seed=42, k_decay=1.0, initialize=True)\n        scheduler  = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=LR, total_steps=total_epochs, epochs=warmup_epoch,\n                                                         steps_per_epoch=None, pct_start=0.3, anneal_strategy='cos', cycle_momentum=True, base_momentum=0.85, \n                                                         max_momentum=0.95, div_factor=25.0, final_div_factor=10000.0, three_phase=False, \n                                                         last_epoch=-1, verbose='deprecated')\n        return [optimizer], [scheduler]\n\n    def train_dataloader(self):\n        train_dataset = ImageNetDataset(train_paths[:1000], train_aug , train = True)\n        train_loader = torch.utils.data.DataLoader(train_dataset,batch_size=batch_size,shuffle=True,pin_memory=True) \n        return train_loader\n\n    def val_dataloader(self):\n        valid_dataset = ImageNetDataset(valid_paths[:1000], valid_aug , train = False)\n        valid_loader = torch.utils.data.DataLoader(valid_dataset,batch_size=batch_size,shuffle=False,pin_memory=True)  \n        return valid_loader\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:53.707517Z","iopub.execute_input":"2024-12-29T14:11:53.708042Z","iopub.status.idle":"2024-12-29T14:11:53.719416Z","shell.execute_reply.started":"2024-12-29T14:11:53.70801Z","shell.execute_reply":"2024-12-29T14:11:53.718501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"L.seed_everything(879246)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:53.720114Z","iopub.execute_input":"2024-12-29T14:11:53.720304Z","iopub.status.idle":"2024-12-29T14:11:53.741941Z","shell.execute_reply.started":"2024-12-29T14:11:53.720287Z","shell.execute_reply":"2024-12-29T14:11:53.74114Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"wandb_logger = WandbLogger(log_model=\"all\", project=\"ImageNet_Lightning\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:11:53.742644Z","iopub.execute_input":"2024-12-29T14:11:53.742903Z","iopub.status.idle":"2024-12-29T14:11:53.75766Z","shell.execute_reply.started":"2024-12-29T14:11:53.742883Z","shell.execute_reply":"2024-12-29T14:11:53.757035Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize a trainer\nbest_checkpoint_callback = L.callbacks.ModelCheckpoint(filename=\"bestmodel-{epoch}-monitor-{val_acc1}\", mode=\"max\")\nevery_epoch_checkpoint_callback = L.callbacks.ModelCheckpoint(filename=\"{epoch}_{val_acc1}\", every_n_epochs=10)\n\ntrainer = L.Trainer(max_epochs=total_epochs,\n                     devices=2, \n                     accelerator='gpu',\n                     logger=wandb_logger,\n                     # callbacks=[early_stop_callback],\n                     callbacks=[best_checkpoint_callback,every_epoch_checkpoint_callback],\n                   )\n\nmodel = LitClassification()\n\ntrainer.fit(model)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-29T14:13:29.703162Z","iopub.execute_input":"2024-12-29T14:13:29.703535Z"}},"outputs":[],"execution_count":null}]}