{"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"}],"dockerImageVersionId":30587,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-25T20:45:56.543637Z","iopub.execute_input":"2023-12-25T20:45:56.543942Z","iopub.status.idle":"2023-12-25T20:45:56.884937Z","shell.execute_reply.started":"2023-12-25T20:45:56.543914Z","shell.execute_reply":"2023-12-25T20:45:56.884062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install xmltodict","metadata":{"execution":{"iopub.status.busy":"2023-12-25T20:45:58.531734Z","iopub.execute_input":"2023-12-25T20:45:58.532425Z","iopub.status.idle":"2023-12-25T20:46:11.781449Z","shell.execute_reply.started":"2023-12-25T20:45:58.532389Z","shell.execute_reply":"2023-12-25T20:46:11.780385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset\nimport os\nimport matplotlib.pyplot as plt\nfrom skimage import io, transform\nimport xmltodict\nfrom tqdm import tqdm\nfrom collections import defaultdict\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2023-12-25T21:00:35.242088Z","iopub.execute_input":"2023-12-25T21:00:35.242482Z","iopub.status.idle":"2023-12-25T21:00:35.24962Z","shell.execute_reply.started":"2023-12-25T21:00:35.242453Z","shell.execute_reply":"2023-12-25T21:00:35.248655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"names = [\"image_path\", \"image_cls\"]\ndummy_ds = pd.read_csv(\"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/ImageSets/CLS-LOC/train_cls.txt\", names=names, header=None, delim_whitespace=True)\ndummy_ds.iloc[0, 0].split(\"/\")[0]","metadata":{"execution":{"iopub.status.busy":"2023-12-25T20:48:50.566207Z","iopub.execute_input":"2023-12-25T20:48:50.566804Z","iopub.status.idle":"2023-12-25T20:48:51.902954Z","shell.execute_reply.started":"2023-12-25T20:48:50.566766Z","shell.execute_reply":"2023-12-25T20:48:51.902011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cls_dummy = pd.read_csv(\"/kaggle/input/imagenet-object-localization-challenge/LOC_synset_mapping.txt\", names=[\"class_frame\"], header=None, sep='/s+')\ncls_dummy[[\"class\", \"class_name\"]] = cls_dummy[\"class_frame\"].str.split(\" \", n=1, expand=True)\ncls_dummy = cls_dummy.drop('class_frame', axis=1)\ncls_dummy.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-25T20:48:56.378476Z","iopub.execute_input":"2023-12-25T20:48:56.378816Z","iopub.status.idle":"2023-12-25T20:48:56.416586Z","shell.execute_reply.started":"2023-12-25T20:48:56.378792Z","shell.execute_reply":"2023-12-25T20:48:56.415681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cls_dummy.index[cls_dummy['class'] == dummy_ds.iloc[0, 0].split(\"/\")[0]].tolist()[0]","metadata":{"execution":{"iopub.status.busy":"2023-12-25T20:48:58.052626Z","iopub.execute_input":"2023-12-25T20:48:58.053362Z","iopub.status.idle":"2023-12-25T20:48:58.060552Z","shell.execute_reply.started":"2023-12-25T20:48:58.053328Z","shell.execute_reply":"2023-12-25T20:48:58.059617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open('/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Annotations/CLS-LOC/val/ILSVRC2012_val_00000026.xml') as xml_f:\n    obj_list = xmltodict.parse(xml_f.read())['annotation']['object']\n    multiple_cls = False\n    if type(obj_list) == list:\n        curr_cls = obj_list[0]['name']\n        classes = [curr_cls]\n        for i in range(1, len(obj_list)):\n            if curr_cls != obj_list[i]['name']:\n                multiple_cls = True\n                classes.append(obj_list[i]['name'])\n        print(classes)\n        print(multiple_cls)\n    else:\n        print(obj_list['name'])","metadata":{"execution":{"iopub.status.busy":"2023-12-25T20:49:00.337936Z","iopub.execute_input":"2023-12-25T20:49:00.33828Z","iopub.status.idle":"2023-12-25T20:49:00.367201Z","shell.execute_reply.started":"2023-12-25T20:49:00.338254Z","shell.execute_reply":"2023-12-25T20:49:00.366325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GenerateImageNetDataset(Dataset):\n    \"\"\"\n    ImageNet Dataset\n    \"\"\"\n    def __init__(self, classes_txt_file: str, data_root_path: str,  split: str, transform=None) -> None:\n        \"\"\"\n        Initializes the ImageNet dataset\n        Args:\n            train_cls_txt (str): Path to the train_cls.txt file containing image paths and labels\n            root_path (str): Root path to the train images\n        \"\"\"\n        self.root_path = data_root_path\n        self.split = split\n        col_names = [\"image_path\", \"image_idx\"]\n        self.imagenet_frame = pd.read_csv(classes_txt_file, names=col_names, header=None, delim_whitespace=True)\n        self.classes_frame = self._get_classes(\"/kaggle/input/imagenet-object-localization-challenge/LOC_synset_mapping.txt\")\n        self.transform = transform\n                \n    def __len__(self) -> int:\n        return len(self.imagenet_frame)\n    \n    \n    def __getitem__(self, idx) -> dict:\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n        \n        img_path = os.path.join(self.root_path, f'{self.imagenet_frame.iloc[idx, 0]}.JPEG')\n        \n        img = io.imread(img_path)\n        if img.ndim == 2:\n            img = self._gray2rgb(img)\n        if img.ndim == 4:\n            img = self._rgba2rgb(img)\n        if self.split == 'train':\n            label = self._get_image_label(img_path, self.classes_frame, self.imagenet_frame)\n        elif self.split == 'val':\n            label = self._get_val_image_label(self.imagenet_frame.iloc[idx, 0], self.classes_frame)\n        sample = {'image': torch.tensor(img / 255.).permute(2, 0, 1), 'label': torch.tensor(label)}\n        \n        if self.transform:\n            sample = {'image': self.transform(sample['image']), 'label': torch.tensor(label)}\n            \n        return sample\n    \n    \n    def _get_classes(self, cls_file_path: str) -> pd.DataFrame:\n        cls_frame = pd.read_csv(cls_file_path, names=[\"class_frame\"], header=None, sep='/s+')\n        cls_frame[[\"class\", \"class_name\"]] = cls_frame[\"class_frame\"].str.split(\" \", n=1, expand=True)\n        return cls_frame.drop('class_frame', axis=1)\n    \n    \n    def _get_image_label(self, img_path: str, cls_frame: pd.DataFrame, target_frame: pd.DataFrame) -> int:\n        return cls_frame.index[cls_frame['class'] == target_frame.iloc[0, 0].split(\"/\")[0]].tolist()[0]\n    \n    \n    def _get_val_image_label(self, img_path: str, cls_frame: pd.DataFrame) -> int:\n        val_path = '/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Annotations/CLS-LOC/val'\n        xml_path = os.path.join(val_path, f'{img_path}.xml')\n        cls_name = ''\n        with open(xml_path) as xml_f:\n            obj_list = xmltodict.parse(xml_f.read())['annotation']['object']\n            if type(obj_list) == list:\n                cls_name = obj_list[0]['name']\n            elif type(obj_list) == dict:\n                cls_name = obj_list['name']\n            else:\n                raise TypeError(\"The object list couldn't be read properly\")\n        return cls_frame.index[cls_frame['class'] == cls_name].tolist()[0]\n    \n    \n    def _gray2rgb(self, img: np.ndarray) -> np.ndarray:\n        if img.ndim != 2:\n            raise ValueError('The image dimension should be 2')\n            \n        rgb_img = np.expand_dims(img, axis=2)\n        rgb_img = np.repeat(rgb_img, 3, axis=2)\n        \n        return rgb_img\n    \n    \n    def _rgba2rgb(self, img: np.ndarray) -> np.ndarray:\n\n        if img.ndim != 4:\n            raise ValueError('The image dimension should be 4!')\n        \n        row, column, channels = img.shape\n\n        rgb = np.zeros((row, column, 3), dtype=np.float32)\n        r, g, b, a = img[:, :, 0], img[:, :, 1], img[:, :, 2], img[:, :, 3]\n        a = np.asarray(a, dtype=np.float32) / 255.\n\n        R, G, B = 255, 255, 255\n\n        rgb[:, :, 0] = r * a + (1.0 - a) * R\n        rgb[:, :, 1] = g * a + (1.0 - a) * G\n        rgb [:, :, 2] = b * a + (1.0 - a) * B\n\n        return np.asarray(rgb, dtype=np.uint8)\n        \n    \n    \nclass AlexNet(torch.nn.Module):\n    \"\"\"\n    Creates AlexNet model\n    Args:\n        num_classes (int): Number of classes in the dataset\n        dropout (float): Dropout rate\n    \"\"\"\n    def __init__(self, num_classes: int=1000, dropout: float=0.5):\n        super(AlexNet, self).__init__()\n        self.num_classes = num_classes\n        self.dropout = dropout\n        self.layer1 = nn.Sequential(\n            nn.Conv2d(3, 96, kernel_size=11, stride=4, padding=0),\n            nn.BatchNorm2d(96), \n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=3, stride=2)\n        )\n        self.layer2 = nn.Sequential(\n            nn.Conv2d(96, 256, kernel_size=5, stride=1, padding=2),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=3, stride=2)\n        )\n        self.layer3 = nn.Sequential(\n            nn.Conv2d(256, 384, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(384),\n            nn.ReLU()\n        )\n        self.layer4 = nn.Sequential(\n            nn.Conv2d(384, 384, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(384),\n            nn.ReLU()\n        )\n        self.layer5 = nn.Sequential(\n            nn.Conv2d(384, 256, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=3, stride=2)\n        )\n        self.fc1 = nn.Sequential(\n            nn.Dropout(self.dropout),\n            nn.Linear(9216, 4096),\n            nn.ReLU()\n        )\n        self.fc2 = nn.Sequential(\n            nn.Dropout(self.dropout),\n            nn.Linear(4096, 4096),\n            nn.ReLU()\n        )\n        self.fc3 = nn.Sequential(\n            nn.Linear(4096, self.num_classes)\n        )\n        \n    def forward(self, x):\n        out = self.layer1(x)\n        out = self.layer2(out)\n        out = self.layer3(out)\n        out = self.layer4(out)\n        out = self.layer5(out)\n        out = out.reshape(out.size(0), -1)\n        out = self.fc1(out)\n        out = self.fc2(out)\n        out = self.fc3(out)\n        return out\n    \n    \nclass BalancedSampler(torch.utils.data.Sampler):\n    def __init__(self, dataset, num_samples_per_class, num_classes=1000):\n        self.dataset = dataset\n        self.num_samples_per_class = num_samples_per_class\n        self.num_classes = num_classes\n\n        # Use defaultdict for efficient class index creation\n        self.class_indices = defaultdict(list)\n        for i, data in tqdm(enumerate(self.dataset)):\n            self.class_indices[data['label']].append(i)\n\n        # Efficiently select a subset of classes\n        selected_classes = random.sample(self.class_indices.keys(), self.num_classes)\n        self.class_indices = [self.class_indices[class_name] for class_name in selected_classes]\n\n    def __iter__(self):\n        indices = []\n        for class_indices in self.class_indices:\n            # Pre-allocate memory for shuffling\n            shuffled_indices = class_indices[:self.num_samples_per_class]\n            random.shuffle(shuffled_indices)\n            indices.extend(shuffled_indices)\n\n        random.shuffle(indices)\n        return iter(indices)\n\n    def __len__(self):\n        return len(self.class_indices) * self.num_samples_per_class\n    \n    \ndef train_one_step(epoch_idx, sm_writer, training_loader, model, optimizer, loss_fn):\n    running_loss = 0\n    running_acc = 0.0\n    last_loss = 0\n    for i, data in enumerate(training_loader):\n        images, labels = data[\"image\"], data[\"label\"]\n        optimizer.zero_grad()\n        \n        outputs = model(samples)\n        \n        loss = loss_fn(outputs, labels.float())\n        loss.backward\n        \n        optimizer.step()\n        \n        running_loss += loss.item()\n        running_acc += (outputs.round() == labels).float().mean()\n        if i % 10 == 9:\n            last_loss = running_loss / 10\n            last_acc = running_acc / 10\n            sm_x = epoch_idx * len(training_loader) + i + 1\n            sm_writer.add_scalar('Loss/train', last_loss, sm_x)\n            sm_writer.add_scalar('Accuracy/train', last_acc, sm_x)\n            running_loss = 0.\n            running_acc = 0.\n            \n    return last_acc, last_loss\n\n\n","metadata":{"execution":{"iopub.status.busy":"2023-12-25T21:00:44.733782Z","iopub.execute_input":"2023-12-25T21:00:44.734159Z","iopub.status.idle":"2023-12-25T21:00:44.772857Z","shell.execute_reply.started":"2023-12-25T21:00:44.734114Z","shell.execute_reply":"2023-12-25T21:00:44.771912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = GenerateImageNetDataset(\"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/ImageSets/CLS-LOC/train_cls.txt\", \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train\", \"train\", transform=transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(227),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n]))\n\nfor i, data in enumerate(train_dataset):\n    plt.imshow(data['image'].permute(1, 2, 0)), plt.show()\n    print(data['image'].shape)\n    print(data['label'])\n    if i == 10:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-12-25T20:58:12.918822Z","iopub.execute_input":"2023-12-25T20:58:12.91955Z","iopub.status.idle":"2023-12-25T20:58:16.728482Z","shell.execute_reply.started":"2023-12-25T20:58:12.919516Z","shell.execute_reply":"2023-12-25T20:58:16.727644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_dataset = GenerateImageNetDataset(\"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/ImageSets/CLS-LOC/val.txt\", \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val\", \"val\", transform=transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(256),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n]))\nfor i, data in enumerate(val_dataset):\n    plt.imshow(data['image'].permute(1, 2, 0)), plt.show()\n    print(data['image'].shape)\n    print(data['label'])\n    if i == 10:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-12-25T20:58:18.943074Z","iopub.execute_input":"2023-12-25T20:58:18.943812Z","iopub.status.idle":"2023-12-25T20:58:21.672802Z","shell.execute_reply.started":"2023-12-25T20:58:18.943776Z","shell.execute_reply":"2023-12-25T20:58:21.671951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nnum_devices = torch.cuda.device_count()","metadata":{"execution":{"iopub.status.busy":"2023-12-25T20:50:06.780002Z","iopub.execute_input":"2023-12-25T20:50:06.780902Z","iopub.status.idle":"2023-12-25T20:50:06.844685Z","shell.execute_reply.started":"2023-12-25T20:50:06.780865Z","shell.execute_reply":"2023-12-25T20:50:06.843622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_sampler = BalancedSampler(train_dataset, 100, num_classes=1000)\nval_data_sampler = BalancedSampler(val_dataset, 10, num_classes=1000)\n\ntrain_dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=16, sampler=train_data_sampler, num_workers=12, pin_memory=True, device=device)\nval_dataloader = torch.utils.data.DataLoader(val_dataset, batch_size=4, sampler=val_data_sampler, num_workers=12, pin_memory=True, device=device)","metadata":{"execution":{"iopub.status.busy":"2023-12-25T21:00:54.266456Z","iopub.execute_input":"2023-12-25T21:00:54.267105Z","iopub.status.idle":"2023-12-25T22:13:05.08384Z","shell.execute_reply.started":"2023-12-25T21:00:54.267071Z","shell.execute_reply":"2023-12-25T22:13:05.082319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}