{"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":"none","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"}],"dockerImageVersionId":30664,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torchvision\nfrom torch import nn\nfrom torchvision.datasets import ImageFolder\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom torchvision.transforms import ToTensor, ToPILImage\nfrom torchvision.models import list_models, get_model, get_weight, get_model_weights\nfrom PIL import Image","metadata":{"execution":{"iopub.status.busy":"2024-03-12T19:22:03.906569Z","iopub.execute_input":"2024-03-12T19:22:03.906845Z","iopub.status.idle":"2024-03-12T19:22:09.385766Z","shell.execute_reply.started":"2024-03-12T19:22:03.906822Z","shell.execute_reply":"2024-03-12T19:22:09.385096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc, os, random, time\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom collections import defaultdict","metadata":{"execution":{"iopub.status.busy":"2024-03-12T19:22:09.387018Z","iopub.execute_input":"2024-03-12T19:22:09.387456Z","iopub.status.idle":"2024-03-12T19:22:09.390934Z","shell.execute_reply.started":"2024-03-12T19:22:09.387435Z","shell.execute_reply":"2024-03-12T19:22:09.390072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Data","metadata":{}},{"cell_type":"code","source":"mapping_path = '/kaggle/input/imagenet-object-localization-challenge/LOC_synset_mapping.txt'\n\nclass_id_to_name_dict = {}\nclass_num_to_name_dict = {}\nclass_id_to_num_dict = {}\nclass_num_to_id_dict = {}\n\ni = 0\nfor line in open(mapping_path):\n    class_id = line[:9].strip()\n    class_name = line[9:].strip()\n    \n    class_id_to_name_dict[class_id] = class_name\n    class_num_to_name_dict[i] = class_name\n    class_id_to_num_dict[class_id] = i\n    class_num_to_id_dict[i] = class_id\n    i += 1","metadata":{"execution":{"iopub.status.busy":"2024-03-12T19:22:09.391835Z","iopub.execute_input":"2024-03-12T19:22:09.392087Z","iopub.status.idle":"2024-03-12T19:22:09.407472Z","shell.execute_reply.started":"2024-03-12T19:22:09.392068Z","shell.execute_reply":"2024-03-12T19:22:09.4067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_PATH = \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train\"","metadata":{"execution":{"iopub.status.busy":"2024-03-12T19:22:09.409343Z","iopub.execute_input":"2024-03-12T19:22:09.409668Z","iopub.status.idle":"2024-03-12T19:22:09.412405Z","shell.execute_reply.started":"2024-03-12T19:22:09.409645Z","shell.execute_reply":"2024-03-12T19:22:09.411741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImageNet1K_TrainSubset(Dataset):\n    def __init__(self, train_path:str, chosen_class:int, preprocess):\n        super(ImageNet1K_TrainSubset, self).__init__()\n\n        self.images = []\n        self.labels = []\n        self.images_ids = []\n        self.preprocess = preprocess\n\n        i = 0\n        for train_class in tqdm(sorted(os.listdir(train_path))):\n            if i != chosen_class:\n                i += 1\n                continue\n            for image_path in sorted(os.listdir(train_path + '/' + train_class)):\n                path = train_path + '/' + train_class + '/' + image_path\n                image = ToTensor()(Image.open(path))\n                self.images.append(image)\n                label = class_id_to_num_dict[path.split('/')[-2]]\n                self.labels.append(label)\n                self.images_ids.append(image_path[:-5])\n            break\n        \n    def __getitem__(self, idx):\n        image_id = self.images_ids[idx]\n        image = self.images[idx]\n        label = self.labels[idx] \n        if image.shape[0] == 1:\n            image = image.repeat(3, 1, 1)\n        elif image.shape[0] == 4:\n            image = transforms.Lambda(lambda x: x[:3, :, :])(image)\n        return image_id, self.preprocess(image), label\n        \n    def __len__(self):\n        return len(self.images_ids)","metadata":{"execution":{"iopub.status.busy":"2024-03-12T19:22:09.413121Z","iopub.execute_input":"2024-03-12T19:22:09.413606Z","iopub.status.idle":"2024-03-12T19:22:09.422168Z","shell.execute_reply.started":"2024-03-12T19:22:09.413586Z","shell.execute_reply":"2024-03-12T19:22:09.421445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classification_models = list_models(module=torchvision.models)\nmodels_and_weights_per_preprocess = defaultdict(list)\n\nfor model in classification_models:\n    try:\n        weight = get_model_weights(model)\n        preprocess = weight.IMAGENET1K_V1.transforms(antialias=False)\n        models_and_weights_per_preprocess[str(preprocess)].append((model, str(weight.IMAGENET1K_V1)))\n    except:\n        pass\n    \n    try:\n        weight = get_model_weights(model)\n        preprocess = weight.IMAGENET1K_V2.transforms(antialias=False)\n        models_and_weights_per_preprocess[str(preprocess)].append((model, str(weight.IMAGENET1K_V2)))\n    except:\n        pass\n    \n    try:\n        weight = get_model_weights(model)\n        preprocess = weight.IMAGENET1K_SWAG_E2E_V1.transforms(antialias=False)\n        models_and_weights_per_preprocess[str(preprocess)].append((model, str(weight.IMAGENET1K_SWAG_E2E_V1)))\n    except:\n        pass\n    \n    try:\n        weight = get_model_weights(model)\n        preprocess = weight.IMAGENET1K_SWAG_LINEAR_V1.transforms(antialias=False)\n        models_and_weights_per_preprocess[str(preprocess)].append((model, str(weight.IMAGENET1K_SWAG_LINEAR_V1)))\n    except:\n        pass\n    \n    try:\n        weight = get_model_weights(model)\n        preprocess = weight.IMAGENET1K_FEATURES.transforms(antialias=False)\n        models_and_weights_per_preprocess[str(preprocess)].append((model, str(weight.IMAGENET1K_FEATURES)))\n    except:\n        pass","metadata":{"execution":{"iopub.status.busy":"2024-03-12T19:22:09.423036Z","iopub.execute_input":"2024-03-12T19:22:09.42328Z","iopub.status.idle":"2024-03-12T19:22:09.457611Z","shell.execute_reply.started":"2024-03-12T19:22:09.423261Z","shell.execute_reply":"2024-03-12T19:22:09.456743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"list(models_and_weights_per_preprocess.values())[0]","metadata":{"execution":{"iopub.status.busy":"2024-03-12T19:22:09.458899Z","iopub.execute_input":"2024-03-12T19:22:09.459205Z","iopub.status.idle":"2024-03-12T19:22:09.474493Z","shell.execute_reply.started":"2024-03-12T19:22:09.459185Z","shell.execute_reply":"2024-03-12T19:22:09.473613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nfor model_str, weight_str in list(models_and_weights_per_preprocess.values())[0]:\n    if model_str != 'vgg11':\n        continue\n    preprocess = get_weight(weight_str).transforms(antialias=False)\n    model = get_model(model_str, weights=weight_str)\n    # model.to(device)\n    for i in range(569, 1000):\n        output = {}\n        train_data = ImageNet1K_TrainSubset(train_path=TRAIN_PATH, chosen_class=i, preprocess=preprocess)\n        train_dataloader = DataLoader(train_data, batch_size=130, shuffle=False)\n        j = 0\n        for image_ids, images, labels in train_dataloader:\n            with torch.no_grad():\n                # image_ids, images, labels = next(iter(train_dataloader))\n                images = images.to(device)\n                probs = torch.nn.functional.softmax(model(images), dim=1)\n                output['image_ids'] = image_ids\n                output['probs'] = probs.cpu().detach().numpy()\n                torch.save(output, 'output__{0}__{1}__{2}__{3}.pth'.format(weight_str, i, class_num_to_id_dict[i], j))\n                # torch.save(output, 'output__{0}__{1}__{2}.pth'.format(weight_str, i, class_num_to_id_dict[i]))\n                j += 1\n    break","metadata":{"execution":{"iopub.status.busy":"2024-03-12T19:22:09.475771Z","iopub.execute_input":"2024-03-12T19:22:09.476018Z","iopub.status.idle":"2024-03-12T19:22:23.057013Z","shell.execute_reply.started":"2024-03-12T19:22:09.475997Z","shell.execute_reply":"2024-03-12T19:22:23.055738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(len(list(torch.argmax(probs, dim=1))))\n# list(torch.argmax(probs, dim=1))","metadata":{"execution":{"iopub.status.busy":"2024-03-12T19:22:23.057926Z","iopub.status.idle":"2024-03-12T19:22:23.05924Z","shell.execute_reply.started":"2024-03-12T19:22:23.059035Z","shell.execute_reply":"2024-03-12T19:22:23.059053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# loaded_data = torch.load('/kaggle/working/output_AlexNet_Weights.IMAGENET1K_V1.pth')\n# print(\"Loaded Data:\", loaded_data)\n# print(type(loaded_data))","metadata":{"execution":{"iopub.status.busy":"2024-03-12T19:22:23.060458Z","iopub.status.idle":"2024-03-12T19:22:23.061097Z","shell.execute_reply.started":"2024-03-12T19:22:23.060877Z","shell.execute_reply":"2024-03-12T19:22:23.060894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# END","metadata":{}}]}