{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Selecting device","metadata":{}},{"cell_type":"code","source":"import torch \nimport torch.nn as nn\n\ndevice = \"cpu\"\nif torch.cuda.is_available():\n    device = \"cuda:1\"","metadata":{"execution":{"iopub.status.busy":"2023-07-04T18:18:23.060996Z","iopub.execute_input":"2023-07-04T18:18:23.061327Z","iopub.status.idle":"2023-07-04T18:18:23.066446Z","shell.execute_reply.started":"2023-07-04T18:18:23.061299Z","shell.execute_reply":"2023-07-04T18:18:23.065643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# General function to test a model","metadata":{}},{"cell_type":"code","source":"import numpy as np\n\ndef test_model(model):\n    model.eval()\n#     starter, ender = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)\n#     timings = []\n#     #GPU-WARM-UP\n#     i=0\n#     for data in val_loader:\n#         if(i>1000):\n#             break\n#         images, labels = data\n#         images = images.to(device)\n#         _ = model(images)\n#         i += 1\n    \n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for data in val_loader:\n            images, labels = data\n            images = images.to(device)\n            labels = labels.to(device)\n            \n#             starter.record()\n            outputs = model(images)\n#             ender.record()\n            \n#             # WAIT FOR GPU SYNC\n#             torch.cuda.synchronize()\n#             curr_time = starter.elapsed_time(ender)\n#             timings.append(curr_time)\n            \n            _, predicted = torch.max(outputs.data, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n\n    print('Accuracy of the network on the 10000 test images: '+str(100 * correct / total))\n    \n#     tot = np.sum(timings)\n#     mean_syn_per_batch = np.sum(timings) / len(timings)\n#     std_syn_per_batch = np.std(timings)\n#     print(\"Total inference time for test data: \"+str(tot))\n#     print(\"Mean inference time per test batch: \"+str(mean_syn_per_batch))\n#     print(\"Standard deviation of inference times per batch: \"+str(std_syn_per_batch))\n    model.train()","metadata":{"execution":{"iopub.status.busy":"2023-07-04T18:18:23.067942Z","iopub.execute_input":"2023-07-04T18:18:23.068196Z","iopub.status.idle":"2023-07-04T18:18:23.08103Z","shell.execute_reply.started":"2023-07-04T18:18:23.068166Z","shell.execute_reply":"2023-07-04T18:18:23.07989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading and normalizing images using TorchVision\n","metadata":{"papermill":{"duration":0.022376,"end_time":"2020-11-06T00:15:34.434234","exception":false,"start_time":"2020-11-06T00:15:34.411858","status":"completed"},"tags":[]}},{"cell_type":"code","source":"! wget https://raw.githubusercontent.com/raghakot/keras-vis/master/resources/imagenet_class_index.json\n! wget https://gist.githubusercontent.com/paulgavrikov/3af1efe6f3dff63f47d48b91bb1bca6b/raw/00bad6903b5e4f84c7796b982b72e2e617e5fde1/ILSVRC2012_val_labels.json","metadata":{"execution":{"iopub.status.busy":"2023-07-04T18:18:23.082882Z","iopub.execute_input":"2023-07-04T18:18:23.083424Z","iopub.status.idle":"2023-07-04T18:18:24.069585Z","shell.execute_reply.started":"2023-07-04T18:18:23.083348Z","shell.execute_reply":"2023-07-04T18:18:24.068495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision\nimport torchvision.transforms as transforms\nimport torchvision.datasets as datasets\nimport os\nfrom torch.utils.data import Dataset\nfrom PIL import Image\nimport json\nfrom torch.utils.data import DataLoader\nimport torch\nfrom tqdm import tqdm","metadata":{"papermill":{"duration":1.135699,"end_time":"2020-11-06T00:15:35.592273","exception":false,"start_time":"2020-11-06T00:15:34.456574","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-04T18:18:24.072431Z","iopub.execute_input":"2023-07-04T18:18:24.073407Z","iopub.status.idle":"2023-07-04T18:18:24.180166Z","shell.execute_reply.started":"2023-07-04T18:18:24.07335Z","shell.execute_reply":"2023-07-04T18:18:24.179303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImageNetKaggle(Dataset):\n    def __init__(self, root, split, transform=None):\n        self.samples = []\n        self.targets = []\n        self.transform = transform\n        self.syn_to_class = {}\n        with open(\"./imagenet_class_index.json\", \"rb\") as f:\n                    json_file = json.load(f)\n                    for class_id, v in json_file.items():\n                        self.syn_to_class[v[0]] = int(class_id)\n        with open(\"./ILSVRC2012_val_labels.json\", \"rb\") as f:\n                    self.val_to_syn = json.load(f)\n        samples_dir = os.path.join(root, \"ILSVRC/Data/CLS-LOC\", split)\n        for entry in os.listdir(samples_dir):\n            if split == \"train\":\n                syn_id = entry\n                target = self.syn_to_class[syn_id]\n                syn_folder = os.path.join(samples_dir, syn_id)\n                for sample in os.listdir(syn_folder):\n                    sample_path = os.path.join(syn_folder, sample)\n                    self.samples.append(sample_path)\n                    self.targets.append(target)\n            elif split == \"val\":\n                syn_id = self.val_to_syn[entry]\n                target = self.syn_to_class[syn_id]\n                sample_path = os.path.join(samples_dir, entry)\n                self.samples.append(sample_path)\n                self.targets.append(target)\n            elif split == \"test\":\n                target = 0 #Dummy class. Not to be used for calculating test accuracy\n                sample_path = os.path.join(samples_dir, entry)\n                self.samples.append(sample_path)\n                self.targets.append(target)\n    def __len__(self):\n            return len(self.samples)\n    def __getitem__(self, idx):\n            x = Image.open(self.samples[idx]).convert(\"RGB\")\n            if self.transform:\n                x = self.transform(x)\n            return x, self.targets[idx]\n        \n\nPATH = \"/kaggle/input/imagenet-object-localization-challenge\"\n\nmean = (0.485, 0.456, 0.406)\nstd = (0.229, 0.224, 0.225)\nval_test_transform = transforms.Compose(\n            [\n                transforms.Resize(256),\n                transforms.CenterCrop(224),\n                transforms.ToTensor(),\n                transforms.Normalize(mean, std),\n            ]\n        )\ntrain_transform = transforms.Compose([\n        transforms.RandomResizedCrop(224),\n        transforms.RandomHorizontalFlip(),\n        transforms.ToTensor(),\n        transforms.Normalize(mean, std),\n    ])\n\nval_dataset = ImageNetKaggle(PATH, \"val\", val_test_transform)\nval_loader = DataLoader(\n            val_dataset,\n            batch_size=256, # may need to reduce this depending on your GPU \n            num_workers=2, # may need to reduce this depending on your num of CPUs and RAM\n            shuffle=True,\n            drop_last=False,\n            pin_memory=True\n        )\n\ntest_dataset = ImageNetKaggle(PATH, \"test\", val_test_transform)\ntest_loader = DataLoader(\n            test_dataset,\n            batch_size=256, # may need to reduce this depending on your GPU \n            num_workers=2, # may need to reduce this depending on your num of CPUs and RAM\n            shuffle=True,\n            drop_last=False,\n            pin_memory=True\n        )\n\ntrain_dataset = ImageNetKaggle(PATH, \"train\", train_transform)\ntrain_loader = DataLoader(\n            train_dataset,\n            batch_size=256, # may need to reduce this depending on your GPU \n            num_workers=2, # may need to reduce this depending on your num of CPUs and RAM\n            shuffle=True,\n            drop_last=False,\n            pin_memory=True\n        )","metadata":{"execution":{"iopub.status.busy":"2023-07-04T18:18:24.181295Z","iopub.execute_input":"2023-07-04T18:18:24.18153Z","iopub.status.idle":"2023-07-04T18:18:27.261025Z","shell.execute_reply.started":"2023-07-04T18:18:24.181509Z","shell.execute_reply":"2023-07-04T18:18:27.26031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nimport math\nimport torch.utils.model_zoo as model_zoo\n\n__all__ = ['ResNet', 'resnet18', 'resnet34', 'resnet50', 'resnet101',\n           'resnet152']\n\nmodel_urls = {\n    'resnet18': 'https://download.pytorch.org/models/resnet18-5c106cde.pth',\n    'resnet34': 'https://download.pytorch.org/models/resnet34-333f7ec4.pth',\n    'resnet50': 'https://download.pytorch.org/models/resnet50-19c8e357.pth',\n    'resnet101': 'https://download.pytorch.org/models/resnet101-5d3b4d8f.pth',\n    'resnet152': 'https://download.pytorch.org/models/resnet152-b121ed2d.pth',\n}\n\n\ndef conv3x3(in_planes, out_planes, stride=1):\n    \"3x3 convolution with padding\"\n    return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,\n                     padding=1, bias=False)\n\n\nclass BasicBlock(nn.Module):\n    expansion = 1\n\n    def __init__(self, inplanes, planes, stride=1, downsample=None):\n        super(BasicBlock, self).__init__()\n        self.conv1 = conv3x3(inplanes, planes, stride)\n        self.bn1 = nn.BatchNorm2d(planes)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = conv3x3(planes, planes)\n        self.bn2 = nn.BatchNorm2d(planes)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        residual = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n\n        if self.downsample is not None:\n            residual = self.downsample(x)\n\n        out += residual\n        out = self.relu(out)\n\n        return out\n\n\nclass Bottleneck(nn.Module):\n    expansion = 4\n\n    def __init__(self, inplanes, planes, stride=1, downsample=None):\n        super(Bottleneck, self).__init__()\n        self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(planes)\n        self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=stride,\n                               padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(planes)\n        self.conv3 = nn.Conv2d(planes, planes * 4, kernel_size=1, bias=False)\n        self.bn3 = nn.BatchNorm2d(planes * 4)\n        self.relu = nn.ReLU(inplace=True)\n        # self.relu = nn.ReLU(inplace=False)\n\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        residual = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n\n        out = self.conv3(out)\n        out = self.bn3(out)\n\n        if self.downsample is not None:\n            residual = self.downsample(x)\n\n        out += residual\n        out = self.relu(out)\n        return out\n\n\nclass ResNet(nn.Module):\n\n    def __init__(self, block, layers, num_classes=1000):\n        self.inplanes = 64\n        super(ResNet, self).__init__()\n        self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3,\n                               bias=False)\n        self.bn1 = nn.BatchNorm2d(64)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n        self.layer1 = self._make_layer(block, 64, layers[0])\n        self.layer2 = self._make_layer(block, 128, layers[1], stride=2)\n        self.layer3 = self._make_layer(block, 256, layers[2], stride=2)\n        self.layer4 = self._make_layer(block, 512, layers[3], stride=2)\n        self.avgpool = nn.AvgPool2d(7)\n        self.fc = nn.Linear(512 * block.expansion, num_classes)\n\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels\n                m.weight.data.normal_(0, math.sqrt(2. / n))\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\n    def _make_layer(self, block, planes, blocks, stride=1):\n        downsample = None\n        if stride != 1 or self.inplanes != planes * block.expansion:\n            downsample = nn.Sequential(\n                nn.Conv2d(self.inplanes, planes * block.expansion,\n                          kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(planes * block.expansion),\n            )\n\n        layers = []\n        layers.append(block(self.inplanes, planes, stride, downsample))\n        self.inplanes = planes * block.expansion\n        for i in range(1, blocks):\n            layers.append(block(self.inplanes, planes))\n\n        return nn.Sequential(*layers)\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.maxpool(x)\n\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n\n        x = self.avgpool(x)\n        x = x.view(x.size(0), -1)\n        x = self.fc(x)\n\n        return x\n\n\ndef resnet18(pretrained=False, **kwargs):\n    \"\"\"Constructs a ResNet-18 model.\n    Args:\n        pretrained (bool): If True, returns a model pre-trained on ImageNet\n    \"\"\"\n    model = ResNet(BasicBlock, [2, 2, 2, 2], **kwargs)\n    if pretrained:\n        model.load_state_dict(model_zoo.load_url(model_urls['resnet18']))\n        print('ResNet-18 Use pretrained model for initalization')\n    return model\n\n\ndef resnet34(pretrained=False, **kwargs):\n    \"\"\"Constructs a ResNet-34 model.\n    Args:\n        pretrained (bool): If True, returns a model pre-trained on ImageNet\n    \"\"\"\n    model = ResNet(BasicBlock, [3, 4, 6, 3], **kwargs)\n    if pretrained:\n        model.load_state_dict(model_zoo.load_url(model_urls['resnet34']))\n        print('ResNet-34 Use pretrained model for initalization')\n    return model\n\n\ndef resnet50(pretrained=False, **kwargs):\n    \"\"\"Constructs a ResNet-50 model.\n    Args:\n        pretrained (bool): If True, returns a model pre-trained on ImageNet\n    \"\"\"\n    model = ResNet(Bottleneck, [3, 4, 6, 3], **kwargs)\n    if pretrained:\n        model.load_state_dict(model_zoo.load_url(model_urls['resnet50']))\n        print('ResNet-50 Use pretrained model for initalization')\n    return model\n\n\ndef resnet101(pretrained=False, **kwargs):\n    \"\"\"Constructs a ResNet-101 model.\n    Args:\n        pretrained (bool): If True, returns a model pre-trained on ImageNet\n    \"\"\"\n    model = ResNet(Bottleneck, [3, 4, 23, 3], **kwargs)\n    if pretrained:\n        model.load_state_dict(model_zoo.load_url(model_urls['resnet101']))\n        print('ResNet-101 Use pretrained model for initalization')\n    return model\n\n\ndef resnet152(pretrained=False, **kwargs):\n    \"\"\"Constructs a ResNet-152 model.\n    Args:\n        pretrained (bool): If True, returns a model pre-trained on ImageNet\n    \"\"\"\n    model = ResNet(Bottleneck, [3, 8, 36, 3], **kwargs)\n    if pretrained:\n        model.load_state_dict(model_zoo.load_url(model_urls['resnet152']))\n        print('ResNet-152 Use pretrained model for initalization')\n    return model\n","metadata":{"execution":{"iopub.status.busy":"2023-07-04T18:18:27.262202Z","iopub.execute_input":"2023-07-04T18:18:27.262611Z","iopub.status.idle":"2023-07-04T18:18:27.312419Z","shell.execute_reply.started":"2023-07-04T18:18:27.262588Z","shell.execute_reply":"2023-07-04T18:18:27.311423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"arch = \"resnet152\"\nusePretrain = True","metadata":{"execution":{"iopub.status.busy":"2023-07-04T18:18:27.314196Z","iopub.execute_input":"2023-07-04T18:18:27.314506Z","iopub.status.idle":"2023-07-04T18:18:27.33086Z","shell.execute_reply.started":"2023-07-04T18:18:27.314418Z","shell.execute_reply":"2023-07-04T18:18:27.329452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# The unpruned model","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/resnet-imagenet-fpgm/testing\")\nimport models\n\nunpruned_model = models.__dict__[arch](pretrained=usePretrain)\nunpruned_model.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-07-04T18:18:27.332545Z","iopub.execute_input":"2023-07-04T18:18:27.332914Z","iopub.status.idle":"2023-07-04T18:18:28.640114Z","shell.execute_reply.started":"2023-07-04T18:18:27.33288Z","shell.execute_reply":"2023-07-04T18:18:28.639096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Testing the accuracy of the unpruned model","metadata":{}},{"cell_type":"code","source":"test_model(unpruned_model)","metadata":{"execution":{"iopub.status.busy":"2023-07-04T18:18:28.644217Z","iopub.execute_input":"2023-07-04T18:18:28.644518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pruning using FPGM","metadata":{}},{"cell_type":"code","source":"! python3 /kaggle/input/resnet-imagenet-fpgm/testing/pruning_imagenet.py -a resnet152 --use_pretrain --lr 0.01 --save_dir ./snapshots/resnet152-rate-0.7 --rate_norm 1 --rate_dist 0.4 --layer_begin 0 --layer_end 462 --layer_inter 3 /kaggle/input ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading the pruned (only zeroed out) model","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/resnet-imagenet-fpgm/testing\")\nimport models\nfrom utils import convert_secs2time, time_string, time_file_str\n\npruned_model = models.__dict__[arch](pretrained=usePretrain)\npruned_model.to(device)\n\nsave_dir = \"./snapshots/resnet152-rate-0.7\"\nprefix = time_file_str()\nfilename = os.path.join(save_dir, 'checkpoint.{:}.{:}.pth.tar'.format(arch, prefix))\npruned_model.load_state_dict(torch.load(filename)['state_dict'])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Saving the pruned (only zeroed out) model","metadata":{}},{"cell_type":"code","source":"torch.save(pruned_model, './resnet_imagenet_pruned_net.pth') # without .state_dict","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Let's test the accuracy of the pruned (only zeroed out) model","metadata":{}},{"cell_type":"code","source":"test_model(pruned_model)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Changing the architecture","metadata":{}},{"cell_type":"code","source":"!pip install torch-pruning\nimport torch_pruning as tp\n    \nfor name, module in pruned_model.named_modules():\n    if isinstance(module, torch.nn.Conv2d): #Iterating over all the conv2d layers of the model\n        channel_indices = [] #Stores indices of the channels to prune within this conv layer\n        t = module.weight.clone().detach()\n        t = t.reshape(t.shape[0], -1)\n        z = torch.all(t == 0, dim=1)\n        z = z.tolist()\n        \n        for i, flag in enumerate(z):\n            if(flag):\n                channel_indices.append(i)\n\n        if(channel_indices == []):\n            continue\n        \n        # 1. build dependency graph for vgg\n        DG = tp.DependencyGraph().build_dependency(pruned_model, example_inputs=torch.randn(1,3,32,32).to(device))\n\n        # 2. Specify the to-be-pruned channels. Here we prune those channels indexed by idxs.\n        group = DG.get_pruning_group(module, tp.prune_conv_out_channels, idxs=channel_indices)\n        #print(group)\n\n        # 3. prune all grouped layers that are coupled with the conv layer (included).\n        if DG.check_pruning_group(group): # avoid full pruning, i.e., channels=0.\n            group.prune()\n    \n# 4. Save & Load\npruned_model.zero_grad() # We don't want to store gradient information\ntorch.save(pruned_model, './vgg_cifar100_arch_pruned_net.pth') # without .state_dict","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Let's test the accuracy of the pruned model after the architecture modifications","metadata":{}},{"cell_type":"code","source":"test_model(pruned_model)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Arch pruned model reload check","metadata":{}},{"cell_type":"code","source":"reloaded_model = torch.load('./vgg_cifar100_arch_pruned_net.pth')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_model(reloaded_model)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}