{"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 gc\nimport torch.nn.functional as F\n\nfrom tqdm import tqdm\nfrom torchvision.io import read_image, ImageReadMode\nfrom torchvision.models import *\nfrom torchvision.datasets import ImageFolder\nfrom torchvision import transforms\nfrom torch.utils.data import DataLoader, Dataset","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-15T05:55:39.348695Z","iopub.execute_input":"2024-05-15T05:55:39.349061Z","iopub.status.idle":"2024-05-15T05:55:45.56575Z","shell.execute_reply.started":"2024-05-15T05:55:39.349014Z","shell.execute_reply":"2024-05-15T05:55:45.564981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2024-05-15T05:55:45.567253Z","iopub.execute_input":"2024-05-15T05:55:45.567802Z","iopub.status.idle":"2024-05-15T05:55:45.571902Z","shell.execute_reply.started":"2024-05-15T05:55:45.567775Z","shell.execute_reply":"2024-05-15T05:55:45.571104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = '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-05-15T05:55:45.572886Z","iopub.execute_input":"2024-05-15T05:55:45.573162Z","iopub.status.idle":"2024-05-15T05:55:45.602494Z","shell.execute_reply.started":"2024-05-15T05:55:45.573138Z","shell.execute_reply":"2024-05-15T05:55:45.601556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def model_generator():\n    models = {\n        'alexnet_v1': (alexnet(AlexNet_Weights.IMAGENET1K_V1), AlexNet_Weights.IMAGENET1K_V1.transforms()),\n#         'resnet18_v1': (resnet18(ResNet18_Weights.IMAGENET1K_V1), ResNet18_Weights.IMAGENET1K_V1.transforms()),\n#         'resnet34_v1': (resnet34(ResNet34_Weights.IMAGENET1K_V1), ResNet34_Weights.IMAGENET1K_V1.transforms()),\n#         'resnet50_v1': (resnet50(ResNet50_Weights.IMAGENET1K_V1), ResNet50_Weights.IMAGENET1K_V1.transforms()),\n#         'resnet101_v1': (resnet101(ResNet101_Weights.IMAGENET1K_V1), ResNet101_Weights.IMAGENET1K_V1.transforms()),\n#         'resnet152_v1': (resnet152(ResNet152_Weights.IMAGENET1K_V1), ResNet152_Weights.IMAGENET1K_V1.transforms()),\n#         'resnet50_v2': (resnet50(ResNet50_Weights.IMAGENET1K_V2), ResNet50_Weights.IMAGENET1K_V2.transforms()),\n#         'resnet101_v2': (resnet101(ResNet101_Weights.IMAGENET1K_V2), ResNet101_Weights.IMAGENET1K_V2.transforms()),\n#         'resnet152_v2': (resnet152(ResNet152_Weights.IMAGENET1K_V2), ResNet152_Weights.IMAGENET1K_V2.transforms()),\n#         'vgg11': (vgg11(VGG11_Weights.IMAGENET1K_V1), VGG11_Weights.IMAGENET1K_V1.transforms()),\n#         'vgg11_bn': (vgg11_bn(VGG11_BN_Weights.IMAGENET1K_V1), VGG11_BN_Weights.IMAGENET1K_V1.transforms()),\n#         'vgg13': (vgg13(VGG13_Weights.IMAGENET1K_V1), VGG13_Weights.IMAGENET1K_V1.transforms()),\n#         'vgg13_bn': (vgg13_bn(VGG13_BN_Weights.IMAGENET1K_V1), VGG13_BN_Weights.IMAGENET1K_V1.transforms()),\n#         'vgg16': (vgg16(VGG16_Weights.IMAGENET1K_V1), VGG16_Weights.IMAGENET1K_V1.transforms()),\n#         'vgg16_bn': (vgg16_bn(VGG16_BN_Weights.IMAGENET1K_V1), VGG16_BN_Weights.IMAGENET1K_V1.transforms()),\n#         'vgg19': (vgg19(VGG19_Weights.IMAGENET1K_V1), VGG19_Weights.IMAGENET1K_V1.transforms()),\n#         'vgg19_bn': (vgg19_bn(VGG19_BN_Weights.IMAGENET1K_V1), VGG19_BN_Weights.IMAGENET1K_V1.transforms()),\n#         'densenet121': (densenet121(DenseNet121_Weights.IMAGENET1K_V1), DenseNet121_Weights.IMAGENET1K_V1.transforms()),\n#         'densenet161': (densenet161(DenseNet161_Weights.IMAGENET1K_V1), DenseNet161_Weights.IMAGENET1K_V1.transforms()),\n#         'densenet169': (densenet169(DenseNet169_Weights.IMAGENET1K_V1), DenseNet169_Weights.IMAGENET1K_V1.transforms()),\n#         'densenet201': (densenet201(DenseNet201_Weights.IMAGENET1K_V1), DenseNet201_Weights.IMAGENET1K_V1.transforms()),\n#         'swin_t': (swin_t(Swin_T_Weights.IMAGENET1K_V1), Swin_T_Weights.IMAGENET1K_V1.transforms()),\n#         'swin_b': (swin_b(Swin_B_Weights.IMAGENET1K_V1), Swin_B_Weights.IMAGENET1K_V1.transforms()),\n#         'swin_v2_t': (swin_v2_t(Swin_V2_T_Weights.IMAGENET1K_V1), Swin_V2_T_Weights.IMAGENET1K_V1.transforms()),\n#         'swin_v2_b': (swin_v2_b(Swin_V2_B_Weights.IMAGENET1K_V1), Swin_V2_B_Weights.IMAGENET1K_V1.transforms()),\n        'swin_s': (swin_s(Swin_S_Weights.IMAGENET1K_V1), Swin_S_Weights.IMAGENET1K_V1.transforms()),\n        'swin_v2_s': (swin_v2_s(Swin_V2_S_Weights.IMAGENET1K_V1), Swin_V2_S_Weights.IMAGENET1K_V1.transforms()),\n\n    }\n    for name, (model, transform) in models.items():\n        yield name, (model, transform)","metadata":{"execution":{"iopub.status.busy":"2024-05-15T05:55:51.252109Z","iopub.execute_input":"2024-05-15T05:55:51.252734Z","iopub.status.idle":"2024-05-15T05:55:51.259299Z","shell.execute_reply.started":"2024-05-15T05:55:51.252702Z","shell.execute_reply":"2024-05-15T05:55:51.258391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImagenetTestDataset(Dataset):\n    def __init__(self, path:str, transform):\n        super().__init__()\n        self.path = path\n        self.img_names = os.listdir(path)\n        self.img_names.sort()\n        self.transform = transform\n    \n    def __getitem__(self, idx):\n        img_path = self.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-05-15T05:55:57.827654Z","iopub.execute_input":"2024-05-15T05:55:57.828439Z","iopub.status.idle":"2024-05-15T05:55:57.834547Z","shell.execute_reply.started":"2024-05-15T05:55:57.828407Z","shell.execute_reply":"2024-05-15T05:55:57.833585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for name, (model, transform) in model_generator():\n    val_dataset = ImagenetTestDataset(image_dir + '/val', transform=transform)\n    val_dataloader = DataLoader(val_dataset, batch_size=100, shuffle=False, num_workers=2)    \n    val_probs = torch.empty((50000, 1000), dtype=torch.float32)\n\n    model = model.to(device)\n    model.eval()\n    with torch.no_grad():\n        for i, images in tqdm(enumerate(val_dataloader)):\n            images = images.to(device)\n            logits = model(images)\n            probs = F.softmax(logits, dim=1).detach().cpu()\n            val_probs[i*100: (i+1)*100] = probs\n    \n    val_probs = val_probs.half()\n    torch.save(val_probs, f'{name}_val.pth')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}