{"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":{"papermill":{"duration":0.008323,"end_time":"2023-07-10T18:02:29.286402","exception":false,"start_time":"2023-07-10T18:02:29.278079","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import torch \nimport torch.nn as nn\n\ndevice = \"cpu\"\nif torch.cuda.is_available():\n    device = \"cuda\"","metadata":{"papermill":{"duration":4.085723,"end_time":"2023-07-10T18:02:33.379681","exception":false,"start_time":"2023-07-10T18:02:29.293958","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# General function to test a model","metadata":{"papermill":{"duration":0.007024,"end_time":"2023-07-10T18:02:33.394362","exception":false,"start_time":"2023-07-10T18:02:33.387338","status":"completed"},"tags":[]}},{"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 test_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 test_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 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 test batch: \"+str(std_syn_per_batch))\n    model.train()","metadata":{"papermill":{"duration":0.021344,"end_time":"2023-07-10T18:02:33.422958","exception":false,"start_time":"2023-07-10T18:02:33.401614","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading and normalizing images using TorchVision\n","metadata":{"papermill":{"duration":0.007297,"end_time":"2023-07-10T18:02:33.439227","exception":false,"start_time":"2023-07-10T18:02:33.43193","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import numpy as np\nimport os\nimport shutil\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torchvision import datasets, transforms\nfrom torch.autograd import Variable\nimport os, sys, shutil, time, random\nfrom scipy.spatial import distance","metadata":{"papermill":{"duration":0.713415,"end_time":"2023-07-10T18:02:34.160257","exception":false,"start_time":"2023-07-10T18:02:33.446842","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = datasets.Flowers102('./', split = \"train\", download=True,\n                             transform=transforms.Compose([\n                                 transforms.Pad(4),\n                                 transforms.RandomHorizontalFlip(),\n                                 transforms.ToTensor(),\n                                 transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n                             ]))\n        \nval_data = datasets.Flowers102('./', split = \"val\", download=True,\n                             transform=transforms.Compose([\n                                 transforms.Pad(4),\n                                 transforms.RandomHorizontalFlip(),\n                                 transforms.ToTensor(),\n                                 transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n                             ]))\n        \ncombined_train_data = torch.utils.data.ConcatDataset([train_data, val_data])\n        \nmin_height = combined_train_data[0][0].shape[1]\nmin_width = combined_train_data[0][0].shape[2]\n\nmax_height = combined_train_data[0][0].shape[1]\nmax_width = combined_train_data[0][0].shape[2]\n\nfor i in range(len(combined_train_data)):\n    if combined_train_data[i][0].shape[1] < min_height :\n        min_height = combined_train_data[i][0].shape[1]\n    if combined_train_data[i][0].shape[2] < min_width:\n        min_width = combined_train_data[i][0].shape[2]\n                \n    if combined_train_data[i][0].shape[1] > max_height:\n        max_height = combined_train_data[i][0].shape[1]\n    if combined_train_data[i][0].shape[2] > max_width:\n        max_width = combined_train_data[i][0].shape[2]\n        \nnew_size = ((min_height+max_height)//2, (min_width+max_width)//2)\n        \ntrain_data = datasets.Flowers102('./', split = \"train\", download=True,\n                             transform=transforms.Compose([\n                                 transforms.Pad(4),\n                                 transforms.Resize(new_size),\n                                 transforms.RandomHorizontalFlip(),\n                                 transforms.ToTensor(),\n                                 transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n                             ]))\n        \nval_data = datasets.Flowers102('./', split = \"val\", download=True,\n                             transform=transforms.Compose([\n                                 transforms.Pad(4),\n                                 transforms.Resize(new_size),\n                                 transforms.RandomHorizontalFlip(),\n                                 transforms.ToTensor(),\n                                 transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n                             ]))\n        \ncombined_train_data = torch.utils.data.ConcatDataset([train_data, val_data])\n        \ntrain_loader = torch.utils.data.DataLoader(\n            combined_train_data,\n            batch_size=16, shuffle=True)\n        \ntest_loader = torch.utils.data.DataLoader(\n            datasets.Flowers102('./', split = \"test\", download=True,\n                             transform=transforms.Compose([\n                                 transforms.Pad(4),\n                                 transforms.Resize(new_size),\n                                 transforms.ToTensor(),\n                                 transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n                             ])),\n            batch_size=16, shuffle=False)","metadata":{"papermill":{"duration":83.431447,"end_time":"2023-07-10T18:03:57.599459","exception":false,"start_time":"2023-07-10T18:02:34.168012","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\n\nimport torch\nimport torch.nn as nn\nfrom torch.autograd import Variable\n\nclass alexnet(nn.Module):\n    def __init__(self, dataset='flowers102', init_weights=False):\n        super(alexnet, self).__init__()\n\n\n        self.imgnet_an = torch.hub.load('pytorch/vision:v0.10.0', 'alexnet', pretrained=True)\n        self.an_features = getattr(self.imgnet_an, 'features')\n        self.an_avgpool = getattr(self.imgnet_an, 'avgpool')\n        self.an_classifier = getattr(self.imgnet_an, 'classifier')\n\n        if dataset == 'cifar10':\n            num_classes = 10\n        elif dataset == 'cifar100':\n            num_classes = 100\n        elif dataset == 'flowers102':\n            num_classes = 102\n            \n        self.classifier = nn.Sequential(\n              nn.Linear(1000, 512),\n              nn.BatchNorm1d(512),\n              nn.ReLU(inplace=True),\n              nn.Linear(512, num_classes)\n            )\n        \n        if init_weights:\n            self._initialize_weights()\n\n    def forward(self, x):\n        \n        x = self.an_features(x)\n        x = self.an_avgpool(x)\n        x = torch.flatten(x, 1)\n        x = self.an_classifier(x)\n        \n        y = self.classifier(x)\n        \n        return y\n\n    def _initialize_weights(self):\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                if m.bias is not None:\n                    m.bias.data.zero_()\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(0.5)\n                m.bias.data.zero_()\n            elif isinstance(m, nn.Linear):\n                m.weight.data.normal_(0, 0.01)\n                m.bias.data.zero_()\n","metadata":{"papermill":{"duration":0.033561,"end_time":"2023-07-10T18:03:57.650744","exception":false,"start_time":"2023-07-10T18:03:57.617183","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"arch = \"alexnet\"","metadata":{"papermill":{"duration":0.025986,"end_time":"2023-07-10T18:03:57.694868","exception":false,"start_time":"2023-07-10T18:03:57.668882","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# The unpruned model","metadata":{"papermill":{"duration":0.017214,"end_time":"2023-07-10T18:03:57.729711","exception":false,"start_time":"2023-07-10T18:03:57.712497","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/alexnet-fpgm/testing\")\nimport models\n\nunpruned_model = models.__dict__[arch](dataset='flowers102')\nunpruned_model.to(device)\n\ntotal = 0\nprint('\\nTrainable parameters:')\n\nfor n, module in unpruned_model.named_modules():\n    for name, param in module.named_parameters():\n        if param.requires_grad:\n            print(n+\".\"+name, '\\t', param.numel())\n            total += param.numel()\nprint()\nprint('Total', '\\t', total)","metadata":{"papermill":{"duration":6.658328,"end_time":"2023-07-10T18:04:04.405138","exception":false,"start_time":"2023-07-10T18:03:57.74681","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.autograd import Variable\n\n# def train(train_loader, model, optimizer, epoch, m=0):\n#     model.train()\n#     avg_loss = 0. \n#     train_acc = 0.\n#     for batch_idx, (data, target) in enumerate(train_loader):\n#         if torch.cuda.is_available():\n#             data, target = data.cuda(), target.cuda()\n#         data, target = Variable(data), Variable(target)\n#         optimizer.zero_grad()\n#         output = model(data)\n#         loss = F.cross_entropy(output, target)\n#         avg_loss += loss.item()\n#         pred = output.data.max(1, keepdim=True)[1]\n#         train_acc += pred.eq(target.data.view_as(pred)).cpu().sum()\n#         loss.backward()\n#         optimizer.step()\n#         if batch_idx % 10 == 0:\n#             print('Train Epoch: {} [{}/{} ({:.1f}%)]\\tLoss: {:.6f}'.format(\n#                 epoch, batch_idx * len(data), len(train_loader.dataset),\n#                        100. * batch_idx / len(train_loader), loss.item()))","metadata":{"papermill":{"duration":0.027622,"end_time":"2023-07-10T18:04:04.451325","exception":false,"start_time":"2023-07-10T18:04:04.423703","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn.functional as F\nimport torch.optim as optim\n\n# optimizer = optim.SGD(unpruned_model.parameters(), lr=0.001, momentum=0.9, weight_decay=1e-4)\n\n# best_prec1 = \"NULL\"\n# for epoch in range(0, 160):\n#     if epoch in [160 * 0.5, 160 * 0.75]:\n#         for param_group in optimizer.param_groups:\n#             param_group['lr'] *= 0.1\n#     train(train_loader, unpruned_model, optimizer, epoch)","metadata":{"papermill":{"duration":0.028633,"end_time":"2023-07-10T18:04:04.498124","exception":false,"start_time":"2023-07-10T18:04:04.469491","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torch.save(unpruned_model, './alexnet_fl102_unpruned_net.pth') # without .state_dict","metadata":{"papermill":{"duration":0.027361,"end_time":"2023-07-10T18:04:04.544204","exception":false,"start_time":"2023-07-10T18:04:04.516843","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Testing the accuracy of the unpruned model","metadata":{"papermill":{"duration":0.020795,"end_time":"2023-07-10T18:04:04.58364","exception":false,"start_time":"2023-07-10T18:04:04.562845","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# test_model(unpruned_model)","metadata":{"papermill":{"duration":0.025711,"end_time":"2023-07-10T18:04:04.627269","exception":false,"start_time":"2023-07-10T18:04:04.601558","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pruning using FPGM","metadata":{"papermill":{"duration":0.018302,"end_time":"2023-07-10T18:04:04.663645","exception":false,"start_time":"2023-07-10T18:04:04.645343","status":"completed"},"tags":[]}},{"cell_type":"code","source":"cfg = [64, 170, 256, 180, 150]\ntot_conv_filters = [64, 192, 384, 256, 256]","metadata":{"papermill":{"duration":0.02574,"end_time":"2023-07-10T18:04:04.707772","exception":false,"start_time":"2023-07-10T18:04:04.682032","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! python3  /kaggle/input/alexnet-fpgm/testing/pruning_fl102_alexnet.py . --batch-size 16 --test-batch-size 16 --dataset flowers102 --arch alexnet --save_path ./logs/alexnet_pretrain/prune_precfg_epoch160 --rate_norm 1 --rate_dist 0.5 --cfg 64,192,192,180,150 --use_state_dict --lr 0.001 --epochs 135 --epoch_prune 5 --use_precfg","metadata":{"papermill":{"duration":36345.994775,"end_time":"2023-07-11T04:09:50.720819","exception":false,"start_time":"2023-07-10T18:04:04.726044","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading the pruned (only zeroed out) model","metadata":{"papermill":{"duration":0.285824,"end_time":"2023-07-11T04:09:51.292429","exception":false,"start_time":"2023-07-11T04:09:51.006605","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/alexnet-fpgm/testing\")\nimport models\n\npruned_model = models.__dict__[arch](dataset='flowers102')\npruned_model.to(device)\n\nfilepath = './logs/alexnet_pretrain/prune_precfg_epoch160'\nfilename = os.path.join(filepath, 'checkpoint.pth.tar')\npruned_model.load_state_dict(torch.load(filename)['state_dict'])","metadata":{"papermill":{"duration":1.51277,"end_time":"2023-07-11T04:09:53.087603","exception":false,"start_time":"2023-07-11T04:09:51.574833","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Saving the pruned (only zeroed out) model","metadata":{"papermill":{"duration":0.988042,"end_time":"2023-07-11T04:09:54.441687","exception":false,"start_time":"2023-07-11T04:09:53.453645","status":"completed"},"tags":[]}},{"cell_type":"code","source":"torch.save(pruned_model, './alexnet_fl102_pruned_net.pth') # without .state_dict","metadata":{"papermill":{"duration":0.771849,"end_time":"2023-07-11T04:09:55.505484","exception":false,"start_time":"2023-07-11T04:09:54.733635","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Let's test the accuracy of the pruned (only zeroed out) model","metadata":{"papermill":{"duration":0.286492,"end_time":"2023-07-11T04:09:56.081839","exception":false,"start_time":"2023-07-11T04:09:55.795347","status":"completed"},"tags":[]}},{"cell_type":"code","source":"test_model(pruned_model)","metadata":{"papermill":{"duration":292.38411,"end_time":"2023-07-11T04:14:48.754917","exception":false,"start_time":"2023-07-11T04:09:56.370807","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Changing the architecture","metadata":{"papermill":{"duration":0.294625,"end_time":"2023-07-11T04:14:49.342781","exception":false,"start_time":"2023-07-11T04:14:49.048156","status":"completed"},"tags":[]}},{"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        shp = combined_train_data[0][0].shape\n        DG = tp.DependencyGraph().build_dependency(pruned_model, example_inputs=torch.randn(1,shp[0],shp[1],shp[2]).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, './alexnet_fl102_arch_pruned_net.pth') # without .state_dict","metadata":{"papermill":{"duration":15.598867,"end_time":"2023-07-11T04:15:05.365267","exception":true,"start_time":"2023-07-11T04:14:49.7664","status":"failed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Let's test the accuracy of the pruned model after the architecture modifications","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"cell_type":"code","source":"test_model(pruned_model)","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Arch pruned model reload check","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"cell_type":"code","source":"reloaded_model = torch.load('./alexnet_fl102_arch_pruned_net.pth')","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_model(reloaded_model)","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"execution_count":null,"outputs":[]}]}