{"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 [Deep-Lock: Secure Authorization for Deep Neural Networks](https://arxiv.org/abs/2008.05966) by Alam 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 and prepare data loading functions","metadata":{}},{"cell_type":"code","source":"import os\nimport math\nimport gzip\nimport torch\nimport hashlib\nimport torchvision\nimport numpy as np\nfrom PIL import Image\nfrom torchvision import transforms\nfrom torch.utils.data import TensorDataset, DataLoader\nfrom cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes\n\ndev = None\nif torch.cuda.is_available():\n    dev = torch.device('cuda')\nelse: \n    dev = torch.device('cpu')","metadata":{"execution":{"iopub.status.busy":"2023-05-25T14:45:15.673214Z","iopub.execute_input":"2023-05-25T14:45:15.674046Z","iopub.status.idle":"2023-05-25T14:45:18.624072Z","shell.execute_reply.started":"2023-05-25T14:45:15.673988Z","shell.execute_reply":"2023-05-25T14:45:18.622933Z"},"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])\n\ndef load_imgs(root_dir : str, max_per_class : int, process : torchvision.transforms):\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    process (type: torchvision.transforms or derivative class)\n    - Transforms images converted to tensors (ex: normalisation/shuffling)\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(process(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-05-25T14:45:46.660897Z","iopub.execute_input":"2023-05-25T14:45:46.661554Z","iopub.status.idle":"2023-05-25T14:45:46.677619Z","shell.execute_reply.started":"2023-05-25T14:45:46.661509Z","shell.execute_reply":"2023-05-25T14:45:46.676488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_data(process : torchvision.transforms=None, 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    process (type: torchvision.transforms or derivative class)\n    - Transforms images converted to tensors (ex: normalisation/shuffling)\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            np.load('/kaggle/input/tempsampleilsvrc/encrypt_imgs.npy')\n        ).to(dev)\n        encrypt_labels = torch.from_numpy(\n            np.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, process)\n        \n        f = gzip.GzipFile(\"encrypt_imgs.npy.gz\", \"w\")\n        np.save(file=f, arr=encrypt_imgs.detach().cpu().numpy())\n        f.close()\n\n        f = gzip.GzipFile(\"encrypt_labels.npy.gz\", \"w\")\n        np.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-05-25T14:46:09.889911Z","iopub.execute_input":"2023-05-25T14:46:09.890402Z","iopub.status.idle":"2023-05-25T14:46:09.902728Z","shell.execute_reply.started":"2023-05-25T14:46:09.890362Z","shell.execute_reply":"2023-05-25T14:46:09.901304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Shuffle Transforms Class\n\nCreate a class to shuflle tensor input","metadata":{}},{"cell_type":"code","source":"# Create a custom transform class that shuffles the data\nclass ShuffleTransform():\n    \"\"\"\n    Divides input tensor image into blocks. Shuffles each block using key.\n\n    Parameters:\n    -----------------\n    key (type: int) \n        key to use for shuffling\n    block_size (type: int)\n        size of each block\n    \"\"\"\n\n    def __init__(self, block_size, key = None):\n        if (key):\n            self.key = key\n        else: \n            self.key = self.gen_key()\n        self.block_size = block_size\n    \n    def __call__(self, img: torch.Tensor) -> torch.Tensor:\n        \n        # Setup seed and batch dim\n        self.gen_seed(img)\n        batched = (len(img.shape) == 4)\n\n        # Shuffle tensor using seeded random number generator\n        block = self.tensor_to_blocks(img, batched)\n        shuffled = self.shuffle_block(block, batched)\n\n        return torch.reshape(shuffled, img.shape)\n    \n    def gen_key(self, key_256 : bool=True):\n        '''\n        Initialises secret key to process data. \n        \n        Parameters\n        -----------------\n        key_256 (type: bool)\n        - Whether to generate a 256-bit key (not 128 bit key)\n        - Use 256 bit keys for SHA-256/SHA3-256 hashes\n        \n        Returns\n        -----------------\n        master_key (type: bytearray)\n        - randomly initialised key \n        '''\n        \n        num_bytes = 32 if key_256 else 16\n        master_key = bytearray(os.urandom(num_bytes))\n        \n        return master_key\n        \n\n    def gen_seed(self, tensor : torch.Tensor) -> int:\n        '''\n        Generates seed from tensor using key\n\n        Parameters:\n        -----------------\n        tensor (type: torch.Tensor)\n            tensor to generate seed from\n\n        Returns:\n        -----------------\n        None, but saves seed (type: int) as a class attribute\n            seed generated from tensor and key\n        '''\n\n        # Convert tensor to bytes\n        hashed = hashlib.sha3_256(tensor.numpy().tobytes())\n        \n        # Setup AES encryption\n        init_vector = os.urandom(16)\n        cipher = Cipher(algorithms.AES(self.key), modes.CBC(init_vector))\n        encryptor = cipher.encryptor()\n        \n        # Generate seed using AES encryption\n        cyphertext = encryptor.update(hashed.digest()) + encryptor.finalize()\n        seed = int.from_bytes(cyphertext[::2][:4], byteorder=\"big\")\n\n        self.seed = seed\n\n    def tensor_to_blocks(self, raw: torch.Tensor, batched : bool) -> torch.Tensor:\n        '''\n        Divides tensor into blocks of size block_size\n        \n        Parameters:\n        -----------------\n        raw (type: torch.Tensor, dim: batch_size x num_channels x height x width) \n            tensor to divide into blocks\n        batched (type: bool)\n        - Whether the input has a batch dimension at the start. \n\n        Returns:\n        -----------------\n        blocks (type: torch.Tensor, dim: batch_size x num_blocks x (block_size^2 x num_channels))\n            Contains raw values, divided into blacks\n        '''\n\n        assert(len(raw.shape) == 4 or len(raw.shape) == 3)\n        # Adjust for batched versus non batched input\n        batch_dim = raw.shape[0] if batched else 1\n        \n        # For every block_size x block_size square in height x width\n        # select and flatten values in all channels in that block\n        num_blocks = math.ceil(raw.shape[-1] / self.block_size) * \\\n            math.ceil(raw.shape[-2] / self.block_size)\n        num_channels = raw.shape[1] if batched else raw.shape[0]\n        blocks = torch.zeros(batch_dim, num_blocks, self.block_size**2 * num_channels)\n\n        # Number batches and blocks\n        for batch in range(batch_dim):\n            for block in range(num_blocks):\n                \n                # Step through each block in height and width dimensions \n                for h_step in range(0, raw.shape[-2], self.block_size):\n                    for w_step in range(0, raw.shape[-1], self.block_size):\n                        \n                        # Save flattened block\n                        if batched:\n                            highdim = raw[batch, :, h_step : h_step + self.block_size, \n                                    w_step : w_step + self.block_size]\n                        else: \n                            highdim = raw[:, h_step : h_step + self.block_size, \n                                    w_step : w_step + self.block_size]\n                            \n                        blocks[batch, block] = torch.flatten(highdim)\n        \n        return blocks\n\n    def shuffle_block(self, raw : torch.Tensor, batched : bool) -> torch.Tensor:\n        '''\n        Randomly shuffles each block in the input tensor. \n\n        Parameters:\n        -----------------\n        raw (type: torch.Tensor, dim: batch_size x num_blocks x (block_size^2 x num_channels))\n            Tensor split into blocks.\n        batched (type: bool)\n        - Whether the input has a batch dimension at the start. \n\n        Returns:\n        -----------------\n        shuffled (type: torch.Tensor, dim: batch_size x num_blocks x (block_size^2 x num_channels))\n            Shuffled tensor. \n        '''\n\n        # Setup random number generator\n        rng = torch.Generator()\n        rng.manual_seed(self.seed)\n\n        for batch in range(raw.size(0)):\n            for block in range(raw.size(1)):\n                # Shuffle each block\n                idx = torch.randperm(raw.size(-1), generator=rng)\n                raw[batch, block] = raw[batch, block][idx]\n                \n        if raw.size(0) == 1:\n            raw = torch.squeeze(raw, 0)\n                \n        return raw","metadata":{"execution":{"iopub.status.busy":"2023-05-25T14:52:12.715106Z","iopub.execute_input":"2023-05-25T14:52:12.716432Z","iopub.status.idle":"2023-05-25T14:52:12.74534Z","shell.execute_reply.started":"2023-05-25T14:52:12.71637Z","shell.execute_reply":"2023-05-25T14:52:12.744048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load data\nblock_size = 8\nmaster_key = bytearray(os.urandom(32))\n\nprocess = transforms.Compose([\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ShuffleTransform(block_size, key = master_key)\n])\n\nencrypt_imgs, encrypt_labels = get_data(process=process, load_mode=False)\n\n# Load MobileNetV2\nmodel = torch.hub.load('pytorch/vision:v0.10.0', 'mobilenet_v2', pretrained=True).to(dev)","metadata":{"execution":{"iopub.status.busy":"2023-05-25T14:52:18.111569Z","iopub.execute_input":"2023-05-25T14:52:18.111977Z","iopub.status.idle":"2023-05-25T14:57:25.579432Z","shell.execute_reply.started":"2023-05-25T14:52:18.111941Z","shell.execute_reply":"2023-05-25T14:57:25.577704Z"},"trusted":true},"execution_count":null,"outputs":[]}]}