{"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":"gpu","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"}],"dockerImageVersionId":30664,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport os\nimport torch.nn.functional as F\n\nfrom tqdm import tqdm\nfrom torchvision.io import read_image, ImageReadMode\nfrom torchvision.datasets import ImageFolder\nfrom torchvision import transforms\nfrom torch.utils.data import DataLoader, Dataset, ConcatDataset","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-06-24T08:31:40.915651Z","iopub.execute_input":"2024-06-24T08:31:40.9162Z","iopub.status.idle":"2024-06-24T08:31:47.574204Z","shell.execute_reply.started":"2024-06-24T08:31:40.916172Z","shell.execute_reply":"2024-06-24T08:31:47.5734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')\n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nimage_dir = \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC\"","metadata":{"execution":{"iopub.status.busy":"2024-06-24T08:31:47.576183Z","iopub.execute_input":"2024-06-24T08:31:47.576907Z","iopub.status.idle":"2024-06-24T08:31:47.600979Z","shell.execute_reply.started":"2024-06-24T08:31:47.576873Z","shell.execute_reply":"2024-06-24T08:31:47.600059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision.models import regnet_y_128gf, RegNet_Y_128GF_Weights\n\nmodel_name = 'regnet_y_128gf_e2e'\nweights = RegNet_Y_128GF_Weights.IMAGENET1K_SWAG_E2E_V1\nmodel = regnet_y_128gf(weights=weights)\ntransform = weights.transforms()\nprint(transform)","metadata":{"execution":{"iopub.status.busy":"2024-06-24T08:31:47.602082Z","iopub.execute_input":"2024-06-24T08:31:47.602328Z","iopub.status.idle":"2024-06-24T08:34:26.895242Z","shell.execute_reply.started":"2024-06-24T08:31:47.602307Z","shell.execute_reply":"2024-06-24T08:34:26.89433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImagenetTrainClassDataset(Dataset):\n    def __init__(self, path:str, class_id:int, transform):\n        assert path.split('/')[-1] == 'train'\n        super().__init__()\n        class_names = sorted(os.listdir(path))\n        self.class_name = class_names[class_id]\n        self.class_path = path + '/' + self.class_name\n        \n        self.img_names = sorted(os.listdir(self.class_path))\n        self.transform = transform\n    \n    def __getitem__(self, idx):\n        img_path = self.class_path + '/' + self.img_names[idx]\n        image = read_image(img_path, ImageReadMode.RGB)\n        return self.transform(image)\n    \n    def __len__(self):\n        return len(self.img_names)","metadata":{"execution":{"iopub.status.busy":"2024-06-24T08:34:26.897524Z","iopub.execute_input":"2024-06-24T08:34:26.898216Z","iopub.status.idle":"2024-06-24T08:34:26.905269Z","shell.execute_reply.started":"2024-06-24T08:34:26.89818Z","shell.execute_reply":"2024-06-24T08:34:26.904195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path = image_dir + '/train'\n\nnum_classes = 1000\nchunk_size = 125 # number of classes in one chunk\n\nmodel = model.to(device)\nmodel.eval()\n\nfor chunk_id in range(5, 6):\n    subsets = []\n    name_list = []\n    for i in tqdm(range(chunk_id * chunk_size, (chunk_id+1) * chunk_size)):\n        class_subset = ImagenetTrainClassDataset(train_path, class_id=i, transform=transform)\n        subsets.append(class_subset)\n        name_list += class_subset.img_names\n\n    name_list = [name.split('.')[0] for name in name_list] # remove JPEG extension\n    subset = ConcatDataset(subsets)\n    train_dataloader = DataLoader(subset, batch_size=10, shuffle=False, num_workers=2)\n    train_probs = torch.empty((len(subset), 1000), dtype=torch.float16)\n\n    with torch.no_grad():\n        for i, images in tqdm(enumerate(train_dataloader)):\n            images = images.to(device)\n            logits = model(images)\n            probs = F.softmax(logits, dim=1)\n            train_probs[i*10: i*10 + probs.size(0)] = probs.detach().cpu().half()\n\n    output = {\n        'probs': train_probs,\n        'img_names': name_list\n    }\n\n    torch.save(output, f'{model_name}_train_{chunk_id}.pth')","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}