{"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":"# Context\n\nThis notebook demonstrates the use of **an algorithm to encrypt any AI model's parameters. The aim is to degrade the model's performance for unauthorised users** without a secret key. The algorithm was proposed in the paper [AdvParams: An Active DNN Intellectual Property Protection Technique via Adversarial Perturbation Based Parameter Encryption](https://arxiv.org/abs/2105.13697) by Xue et al.\n\nFor more details and documentation on reusing this algorithm, please see [this Github repository](https://github.com/Madhav-Malhotra/ML-parameter-encryption).\n\n------------","metadata":{}},{"cell_type":"markdown","source":"# Setup\n\nLoad libraries, the encryption dataset, and the model.","metadata":{}},{"cell_type":"code","source":"import gzip\nimport numpy\nimport torch \nimport random\nfrom torchvision import transforms\nfrom torch.utils.data import TensorDataset, DataLoader\n\ndev = None\nif torch.cuda.is_available():\n    dev = torch.device('cuda')\nelse: \n    dev = torch.device('cpu')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-29T19:36:12.873464Z","iopub.execute_input":"2023-04-29T19:36:12.874045Z","iopub.status.idle":"2023-04-29T19:36:16.306538Z","shell.execute_reply.started":"2023-04-29T19:36:12.874009Z","shell.execute_reply":"2023-04-29T19:36:16.305277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Helper functions to convert images to processed tensors\nconvert_img = transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(224),\n    transforms.ToTensor()\n])\nnormalise_img = transforms.Compose([\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\n\ndef load_imgs(root_dir : str, max_per_class : int):\n    '''\n    Loads images into a matrix of preprocessed tensors.\n    \n    Parameters\n    -------------------\n    root_dir (type: string)\n    - A filepath to a folder with subfolders with images. \n    - Each subfolder corresponds to one class.\n    max_per_class (type: int)\n    - The maximum number of images to save from each class.\n    \n    Returns\n    -------------------\n    torch.Tensor\n    - Shape is num_images x num_channels x width x height\n    '''\n    \n    images = []\n    labels = []\n    y = 0\n    \n    # Go through sub-class folders\n    for subclass in os.listdir(root_dir):\n        folder_name = os.path.join(root_dir, subclass)\n        files = os.listdir(folder_name)\n        print('.', end=\"\")\n        \n        # Go through up to max_per_class files in each sub-class\n        i = 0\n        max_files = min(len(files), max_per_class)\n        while (i < max_files):\n            \n            # Save valid files\n            file_name = os.path.join(folder_name, files[i])\n            if os.path.isfile(file_name):\n                converted = convert_img(Image.open(file_name))\n                if (converted.shape == (3, 224, 224)):\n                    images.append(normalise_img(converted))\n                    labels.append(y)\n                    \n            i += 1\n        y += 1\n        \n    # return output\n    encrypt_imgs = torch.stack(tuple(images), 0).to(dev)\n    encrypt_labels = torch.tensor(labels, dtype=torch.int16, device=dev)\n    return encrypt_imgs, encrypt_labels ","metadata":{"execution":{"iopub.status.busy":"2023-04-29T19:36:21.365545Z","iopub.execute_input":"2023-04-29T19:36:21.366056Z","iopub.status.idle":"2023-04-29T19:36:21.378828Z","shell.execute_reply.started":"2023-04-29T19:36:21.366023Z","shell.execute_reply":"2023-04-29T19:36:21.377699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_data(load_mode : bool=True, imgs_per_class : int=1) -> list:\n    '''\n    Returns image set for encryption (images, labels)\n    \n    Parameters\n    ------------------\n    load_mode (type: bool)\n    - Whether to load saved data or process new data\n    imgs_per_class (type: int)\n    - Number of images to load for each class\n    \n    Returns\n    encrypt_imgs (type: torch.Tensor, dim: N x C x H x W)\n    - Preprocessed images to use for encryption\n    encrypt_labels (type: torch.Tensor, dim: N)\n    - Labels for above images\n    '''\n    \n    # Load saved tensors\n    if (load_mode):\n        encrypt_imgs = torch.from_numpy(\n            numpy.load('/kaggle/input/tempsampleilsvrc/encrypt_imgs.npy')\n        ).to(dev)\n        encrypt_labels = torch.from_numpy(\n            numpy.load('/kaggle/input/tempsampleilsvrc/encrypt_labels.npy')\n        ).to(dev)\n    \n    # Process new tensors and then save\n    else:\n        encrypt_imgs, encrypt_labels = load_imgs(\n            '/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train', \n            imgs_per_class)\n        \n        f = gzip.GzipFile(\"encrypt_imgs.npy.gz\", \"w\")\n        numpy.save(file=f, arr=encrypt_imgs.detach().cpu().numpy())\n        f.close()\n\n        f = gzip.GzipFile(\"encrypt_labels.npy.gz\", \"w\")\n        numpy.save(file=f, arr=encrypt_labels.detach().cpu().numpy())\n        f.close()\n    \n    return encrypt_imgs, encrypt_labels","metadata":{"execution":{"iopub.status.busy":"2023-04-29T19:36:23.98366Z","iopub.execute_input":"2023-04-29T19:36:23.984027Z","iopub.status.idle":"2023-04-29T19:36:23.99263Z","shell.execute_reply.started":"2023-04-29T19:36:23.983993Z","shell.execute_reply":"2023-04-29T19:36:23.991424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load data\nencrypt_imgs, encrypt_labels = get_data()\n\n# Load MobileNetV2\nmodel = torch.hub.load('pytorch/vision:v0.10.0', 'mobilenet_v2', pretrained=True).to(dev)\nmodel.train(False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Encryption Functions","metadata":{}},{"cell_type":"code","source":"def get_layer_set(max_layers : int, model : torch.nn.Module) -> list:\n    '''\n    Randomly selects model layers to choose parameters for encryption\n    \n    Parameters\n    --------------------\n    max_layers (type: int)\n    - The maximum number of layers to be selected in the encryption set\n    model (type: torch.nn.Module)\n    - The model from which to draw layers from\n    \n    Returns\n    --------------------\n    list\n    - Has names of the selected layers\n    '''\n    \n    # Arrange layers into list\n    selected_layers = [name for name,_ in model.named_parameters()]\n\n    # Prune layers randomly if too many selected\n    while len(selected_layers) > max_layers:\n        rand_i = random.randint(0, len(selected_layers) - 1)\n        del selected_layers[rand_i]\n    \n    return selected_layers","metadata":{"execution":{"iopub.status.busy":"2023-04-29T19:36:56.004904Z","iopub.execute_input":"2023-04-29T19:36:56.005911Z","iopub.status.idle":"2023-04-29T19:36:56.01367Z","shell.execute_reply.started":"2023-04-29T19:36:56.005868Z","shell.execute_reply":"2023-04-29T19:36:56.012153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class encryption_utils:\n    def __init__(\n        self, \n        model : torch.nn.Module, \n        encrypt_data : torch.utils.data.DataLoader, \n        encrypt_layers : list, \n        max_per_layer : int, \n        loss_threshold : float, \n        step_size : float, \n        boundary_distance : float):\n        \n        self.model = model\n        self.encrypt_data = encrypt_data\n        self.encrypt_layers = encrypt_layers\n        self.max_per_layer = max_per_layer\n        self.loss_threshold = loss_threshold\n        self.step_size = step_size\n        self.boundary_distance = boundary_distance\n    \n    def compute_bounds(self, layer_params):\n        '''\n        Computes boundaries for parameter values in selected layer\n        \n        Parameters\n        --------------------\n        layer_params (type: torch.Tensor, dim: variable)\n        - Weights (multidimensional) or biases (1D) of current layer.\n        \n        Returns\n        --------------------\n        bound_low (type: float)\n        - lower limit on parameter values\n        bound_high (type: float)\n        - upper limit on parameter values\n        '''\n        \n        with torch.no_grad():\n            l_max = torch.max(layer_params).detach()\n            l_min = torch.min(layer_params).detach()\n            l_range = self.boundary_distance * (l_max - l_min)\n        \n        return l_min+l_range, l_max-l_range\n    \n    def compute_avg_loss(self):\n        '''\n        Computes average loss and gradient across minibatches\n        '''\n        \n        avg_loss = 0\n        batches = 0\n\n        # Compute gradient and average loss for batches\n        for x,y in self.encrypt_data:\n            print('.', end=\"\")\n            batches += 1\n            \n            pred = self.model(x)\n            loss = torch.nn.functional.cross_entropy(pred, y.long())\n            \n            avg_loss += loss.item()\n            loss.backward()\n        \n        return avg_loss / batches, batches\n\n    def get_max_grad(self, batches, mask, param):\n        '''\n        Finds the maximum gradient vector component\n        \n        Parameters\n        --------------------\n        batches (type: int)\n        - Number of batches to average gradient over. \n        mask (type: torch.Tensor, dim: variable)\n        - A tensor of 1s/0s for each parameter in current layer \n        - 1 indicates a parameter is updatable. \n        - Has same dim as current layer\n        param (type: torch.nn.Parameter)\n        - Layer parameters to update. \n        \n        Returns\n        --------------------\n        val (type: float)\n        - Maximum gradient component for current layer's params\n        i (type: int)\n        - Max gradient component's index for current layer's params\n        - Layers with multi-dimensional parameters are flattened\n          before index is computed.\n        '''\n        \n        with torch.no_grad():\n            # Return the maximum gradient val/index\n            avg_grad = (param.grad / batches) * mask\n            val = torch.max(avg_grad).item()\n            # returns flattened index for multidimensional tensors\n            i = torch.argmax(avg_grad)\n\n            return val, int(i)\n    \n    def unroll_index(self, i, new_dim):\n        '''\n        Turns a one-dimensional index into a new dimensional shape\n        \n        Parameters\n        --------------------\n        i (type: int)\n        - index of flattened tensor\n        new_dim (type: tuple)\n        - dimensions to adjust the flattened index into\n        \n        Returns\n        --------------------\n        list\n        - 1+ elements representing indices along different dimensions\n        '''\n        \n        # Handle exception for low-dimensional index\n        if len(new_dim) == 1:\n            return [i]\n        elif not len(new_dim):\n            return [0]\n        \n        # Create useful variables\n        out_i = []\n        new_dim = list(new_dim)\n        prod = int(new_dim[0])\n        for j in range(1, len(new_dim)):\n            prod *= int(new_dim[j])\n        \n        while (len(new_dim)):\n            # Get current index\n            prod /= int(new_dim[0])\n            out_i.append(int(i // prod))\n            \n            # Prepare for next round\n            i = int(i % prod)\n            del new_dim[0]\n        \n        return out_i\n            \n                \n    def compute_update(self, param, grad, param_index):\n        '''\n        Computes update step for selected parameter\n        \n        Parameters\n        --------------------\n        param (type: torch.nn.Parameter)\n        - Layer housing selected parameter\n        grad (type: float)\n        - Gradient component for selected parameter\n        param_index (type: int)\n        - Index of selected parameter within the layer\n        \n        Returns\n        --------------------\n        update (type: float)\n        -  new value for selected parameter\n        unrolled_i (type: list)\n        - the index of the element to update\n        '''\n        \n        # compute update step\n        param_range = torch.max(param).detach() - torch.min(param).detach()\n        step = self.step_size * int(grad > 0) * param_range\n\n        # get new param val\n        unrolled_i = self.unroll_index(param_index, param.shape)\n        new_val = param[tuple(unrolled_i)].detach() + step\n        return new_val, unrolled_i\n                \n\n    def encrypt_parameters(self):\n        '''\n        Makes selectively-targeted adversarial modifications to parameters.\n\n        Parameters\n        --------------------\n        encrypt_data (type: torch.utils.data.DataLoader)\n        - generates batched images and labels to evaluate loss\n        encrypt_layers (type: dict)\n        - keys are layer names and values are model parameters to adjust\n        max_per_layer (type: int)\n        - the maximum number of parameters which can be modified per layer\n        loss_threshold (type: float)\n        - loss at which the encryption process stops\n\n        Returns: \n        ---------------------\n        dict\n        - has keys for each encrypted layer and values indicating their \n        adjusted weights\n        '''\n\n        modifications = {}\n\n        # Enable grad for layers to encrypt\n        for name, param in self.model.named_parameters():\n            if name not in self.encrypt_layers:\n                continue\n            else:\n                param.requires_grad = True\n            print(name)\n            \n            # Generate mask to track which params in each layer are unusable\n            modifications[name] = []\n            mask = torch.ones(param.shape, dtype=bool, device=dev).detach()\n            bound_low, bound_high = self.compute_bounds(param)\n\n            # Only run updates on current layer up to max iters\n            for i in range(self.max_per_layer):\n                print(i, end=\"\")\n                model.zero_grad()\n                avg_loss, batches = self.compute_avg_loss()\n\n                # Abort if loss raised sufficiently\n                if avg_loss > self.loss_threshold:\n                    return modifications\n\n                # Get adversarial update\n                val, i = self.get_max_grad(batches, mask, param)\n                new_val, unrolled_i = self.compute_update(param, val, i)\n                \n                # Update mask and parameter\n                if new_val < bound_low or new_val > bound_high:\n                    mask[tuple(unrolled_i)] = False\n                \n                with torch.no_grad():\n                    clipped = torch.clamp(new_val, min=bound_low, max=bound_high)\n                    j = tuple(unrolled_i)\n                    modifications[name].append((j, clipped - param[j]))\n                    param[j] = clipped\n            \n            # Disable grad after layer operations finished\n            param.requires_grad = False\n            \n        return modifications","metadata":{"execution":{"iopub.status.busy":"2023-04-29T19:36:58.485196Z","iopub.execute_input":"2023-04-29T19:36:58.485933Z","iopub.status.idle":"2023-04-29T19:36:58.510312Z","shell.execute_reply.started":"2023-04-29T19:36:58.485896Z","shell.execute_reply":"2023-04-29T19:36:58.508532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Encrypt Model Parameters","metadata":{}},{"cell_type":"code","source":"# Init hyperparameters\ndataset = TensorDataset(encrypt_imgs, encrypt_labels)\nencrypt_data = DataLoader(dataset, batch_size=32)\nencrypt_layers = get_layer_set(25, model) # max 25 layers adjusted\n\nmax_per_layer = 25         # max weights adjusted per layer\nstep_size = 0.1            # for gradient updates\n# lower bound: 0 = encrypted weight values can be anywhere in range of current weight values\n# upper bound: 0.5 = encrypted weight values go at the midpoint of current weight values\nboundary_distance = 0.1    \n\n# Test current loss before encryption\nx,y = next(iter(encrypt_data))\nwith torch.no_grad():\n    loss = torch.nn.functional.cross_entropy(model(x), y.long())\n    print(\"Loss before\", loss)\nloss_threshold = loss.item() * 5","metadata":{"execution":{"iopub.status.busy":"2023-04-29T19:37:21.414757Z","iopub.execute_input":"2023-04-29T19:37:21.415205Z","iopub.status.idle":"2023-04-29T19:37:25.504915Z","shell.execute_reply.started":"2023-04-29T19:37:21.415167Z","shell.execute_reply":"2023-04-29T19:37:25.502916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"instance = encryption_utils(model, encrypt_data, encrypt_layers, max_per_layer, loss_threshold, step_size, boundary_distance)","metadata":{"execution":{"iopub.status.busy":"2023-04-29T19:37:30.649054Z","iopub.execute_input":"2023-04-29T19:37:30.650047Z","iopub.status.idle":"2023-04-29T19:37:30.655171Z","shell.execute_reply.started":"2023-04-29T19:37:30.650006Z","shell.execute_reply":"2023-04-29T19:37:30.653368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"modifications = instance.encrypt_parameters()","metadata":{"execution":{"iopub.status.busy":"2023-04-29T19:37:33.872677Z","iopub.execute_input":"2023-04-29T19:37:33.873518Z","iopub.status.idle":"2023-04-29T20:00:41.104273Z","shell.execute_reply.started":"2023-04-29T19:37:33.873476Z","shell.execute_reply":"2023-04-29T20:00:41.101515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Show loss after encryption\nx,y = next(iter(encrypt_data))\nwith torch.no_grad():\n    loss = torch.nn.functional.cross_entropy(model(x), y.long())\n    print(\"Encrypted: \", loss)","metadata":{"execution":{"iopub.status.busy":"2023-04-29T20:01:15.229183Z","iopub.execute_input":"2023-04-29T20:01:15.230224Z","iopub.status.idle":"2023-04-29T20:01:15.258218Z","shell.execute_reply.started":"2023-04-29T20:01:15.230169Z","shell.execute_reply":"2023-04-29T20:01:15.256956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Decrypt Parameters","metadata":{}},{"cell_type":"code","source":"def decrypt_parameters(model : torch.nn.Module, modifications : dict) -> None:\n    '''\n    Decrypts parameters using secret key\n    \n    Parameters\n    --------------------\n    model (type: torch.nn.Module)\n    - The model containing parameters to decrypt\n    modifications (type: dict)\n    - An object recording the modifications made to parameters by layer.10\n    \n    Returns\n    --------------------\n    None\n    '''\n    \n    encrypted_layers = modifications.keys()\n    \n    # Select model layer\n    for name, param in model.named_parameters():\n        if name in encrypted_layers:\n            print(name)\n            \n            # Replace modified parameter\n            for index, value in modifications[name]:\n                with torch.no_grad():\n                    param[index] = param[index] - value","metadata":{"execution":{"iopub.status.busy":"2023-04-29T20:01:23.077604Z","iopub.execute_input":"2023-04-29T20:01:23.078895Z","iopub.status.idle":"2023-04-29T20:01:23.085992Z","shell.execute_reply.started":"2023-04-29T20:01:23.078844Z","shell.execute_reply":"2023-04-29T20:01:23.084704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"decrypt_parameters(model, modifications)","metadata":{"execution":{"iopub.status.busy":"2023-04-29T20:01:26.536554Z","iopub.execute_input":"2023-04-29T20:01:26.537216Z","iopub.status.idle":"2023-04-29T20:01:26.578868Z","shell.execute_reply.started":"2023-04-29T20:01:26.537164Z","shell.execute_reply":"2023-04-29T20:01:26.577693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Test loss after decryption\nx,y = next(iter(encrypt_data))\nwith torch.no_grad():\n    loss = torch.nn.functional.cross_entropy(model(x), y.long())\n    print(\"Decrypted: \", loss)","metadata":{"execution":{"iopub.status.busy":"2023-04-29T20:01:29.101046Z","iopub.execute_input":"2023-04-29T20:01:29.101831Z","iopub.status.idle":"2023-04-29T20:01:29.128434Z","shell.execute_reply.started":"2023-04-29T20:01:29.101792Z","shell.execute_reply":"2023-04-29T20:01:29.127158Z"},"trusted":true},"execution_count":null,"outputs":[]}]}