{"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, the encryption dataset, and the model.","metadata":{}},{"cell_type":"code","source":"import torch \nimport struct\nimport random\nimport numpy as np\nfrom datetime import datetime\nfrom torchvision import transforms\nfrom torch.utils.data import TensorDataset, DataLoader\n\n# init randint generator\nrandom.seed(datetime.now().timestamp())\n\n# Only CPU currently supported\ndev = None\nif torch.cuda.is_available():\n    dev = torch.device('cpu')\nelse: \n    dev = torch.device('cpu')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-11T01:19:40.295754Z","iopub.execute_input":"2023-05-11T01:19:40.296166Z","iopub.status.idle":"2023-05-11T01:19:40.303585Z","shell.execute_reply.started":"2023-05-11T01:19:40.296097Z","shell.execute_reply":"2023-05-11T01:19:40.302518Z"},"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-05-11T01:19:48.907133Z","iopub.execute_input":"2023-05-11T01:19:48.907614Z","iopub.status.idle":"2023-05-11T01:19:48.921781Z","shell.execute_reply.started":"2023-05-11T01:19:48.907562Z","shell.execute_reply":"2023-05-11T01:19:48.920507Z"},"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            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)\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-11T01:20:06.211176Z","iopub.execute_input":"2023-05-11T01:20:06.212241Z","iopub.status.idle":"2023-05-11T01:20:06.222483Z","shell.execute_reply.started":"2023-05-11T01:20:06.212197Z","shell.execute_reply":"2023-05-11T01:20:06.22135Z"},"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)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T01:23:21.17183Z","iopub.execute_input":"2023-05-11T01:23:21.172664Z","iopub.status.idle":"2023-05-11T01:23:27.323336Z","shell.execute_reply.started":"2023-05-11T01:23:21.172623Z","shell.execute_reply":"2023-05-11T01:23:27.321422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Byte Helper Functions","metadata":{}},{"cell_type":"code","source":"def float_to_bytes(num : float) -> bytes:\n    return struct.pack('f', num)\n\n\ndef bytes_to_float(b : bytes) -> float:\n    # Returns the float at the first element of the tuple\n    return struct.unpack('f', b)[0]\n\n\ndef xor_bytes(bytes1 : bytes, bytes2 : bytes) -> bytes:\n    ''' Returns new bytes object after XORing each byte '''\n    return bytes([b1 ^ b2 for b1, b2 in zip(bytes1, bytes2)])\n\n\ndef rotate_bytes(byte : bytes) -> bytes:\n    ''' Returns new bytes object after rotating each byte '''\n    return byte[1:] + byte[:1]\n\n\ndef print_bytes(byte : bytes) -> None:\n    ''' Prints a bytes object as a byte string '''\n\n    out = ''\n    for b in byte:\n        # Get rid of 0b prefix and pad start with zeros\n        string = bin(b)[2:]\n        string = '0' * (8 - len(string)) + string\n        # Add a space every 8 bits\n        out += string + ' '\n    \n    print(out)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T01:23:31.772888Z","iopub.execute_input":"2023-05-11T01:23:31.773639Z","iopub.status.idle":"2023-05-11T01:23:31.784842Z","shell.execute_reply.started":"2023-05-11T01:23:31.773597Z","shell.execute_reply":"2023-05-11T01:23:31.781695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Encryption Functions","metadata":{}},{"cell_type":"code","source":"def update_round_const(prev_const : int) -> int:\n    ''' Returns new round constant (integer) based on last round '''\n    \n    # Keep doubling \n    update = prev_const << 1\n    if update < 256:\n        return update\n    # unless constant > 1 byte.\n    else: \n        return update ^ 0x11b\n\n\ndef sub_bytes(byte : bytearray, inverse : bool = False) -> bytes:\n    ''' \n    Performs AES sub-bytes on each byte \n    \n    Parameters\n    ----------------\n    byte (type: bytearray)\n    - Bytes object to perform AES sub-bytes on\n    inverse (type: boolean)\n    - When true, runs byte through inverse rijndael s-box\n\n    Returns\n    ----------------\n    bytes\n    - Bytes object after AES sub-bytes\n    '''\n\n    # Create dict of s-box values and inverse s-box values\n    sbox = {\n        '00': '63', '01': '7c', '02': '77', '03': '7b', '04': 'f2', '05': '6b', '06': '6f', '07': 'c5', '08': '30', '09': '01', '0a': '67', '0b': '2b', '0c': 'fe', '0d': 'd7', '0e': 'ab', '0f': '76', \n        '10': 'ca', '11': '82', '12': 'c9', '13': '7d', '14': 'fa', '15': '59', '16': '47', '17': 'f0', '18': 'ad', '19': 'd4', '1a': 'a2', '1b': 'af', '1c': '9c', '1d': 'a4', '1e': '72', '1f': 'c0', \n        '20': 'b7', '21': 'fd', '22': '93', '23': '26', '24': '36', '25': '3f', '26': 'f7', '27': 'cc', '28': '34', '29': 'a5', '2a': 'e5', '2b': 'f1', '2c': '71', '2d': 'd8', '2e': '31', '2f': '15', \n        '30': '04', '31': 'c7', '32': '23', '33': 'c3', '34': '18', '35': '96', '36': '05', '37': '9a', '38': '07', '39': '12', '3a': '80', '3b': 'e2', '3c': 'eb', '3d': '27', '3e': 'b2', '3f': '75', \n        '40': '09', '41': '83', '42': '2c', '43': '1a', '44': '1b', '45': '6e', '46': '5a', '47': 'a0', '48': '52', '49': '3b', '4a': 'd6', '4b': 'b3', '4c': '29', '4d': 'e3', '4e': '2f', '4f': '84', \n        '50': '53', '51': 'd1', '52': '00', '53': 'ed', '54': '20', '55': 'fc', '56': 'b1', '57': '5b', '58': '6a', '59': 'cb', '5a': 'be', '5b': '39', '5c': '4a', '5d': '4c', '5e': '58', '5f': 'cf', \n        '60': 'd0', '61': 'ef', '62': 'aa', '63': 'fb', '64': '43', '65': '4d', '66': '33', '67': '85', '68': '45', '69': 'f9', '6a': '02', '6b': '7f', '6c': '50', '6d': '3c', '6e': '9f', '6f': 'a8', \n        '70': '51', '71': 'a3', '72': '40', '73': '8f', '74': '92', '75': '9d', '76': '38', '77': 'f5', '78': 'bc', '79': 'b6', '7a': 'da', '7b': '21', '7c': '10', '7d': 'ff', '7e': 'f3', '7f': 'd2', \n        '80': 'cd', '81': '0c', '82': '13', '83': 'ec', '84': '5f', '85': '97', '86': '44', '87': '17', '88': 'c4', '89': 'a7', '8a': '7e', '8b': '3d', '8c': '64', '8d': '5d', '8e': '19', '8f': '73', \n        '90': '60', '91': '81', '92': '4f', '93': 'dc', '94': '22', '95': '2a', '96': '90', '97': '88', '98': '46', '99': 'ee', '9a': 'b8', '9b': '14', '9c': 'de', '9d': '5e', '9e': '0b', '9f': 'db', \n        'a0': 'e0', 'a1': '32', 'a2': '3a', 'a3': '0a', 'a4': '49', 'a5': '06', 'a6': '24', 'a7': '5c', 'a8': 'c2', 'a9': 'd3', 'aa': 'ac', 'ab': '62', 'ac': '91', 'ad': '95', 'ae': 'e4', 'af': '79', \n        'b0': 'e7', 'b1': 'c8', 'b2': '37', 'b3': '6d', 'b4': '8d', 'b5': 'd5', 'b6': '4e', 'b7': 'a9', 'b8': '6c', 'b9': '56', 'ba': 'f4', 'bb': 'ea', 'bc': '65', 'bd': '7a', 'be': 'ae', 'bf': '08', \n        'c0': 'ba', 'c1': '78', 'c2': '25', 'c3': '2e', 'c4': '1c', 'c5': 'a6', 'c6': 'b4', 'c7': 'c6', 'c8': 'e8', 'c9': 'dd', 'ca': '74', 'cb': '1f', 'cc': '4b', 'cd': 'bd', 'ce': '8b', 'cf': '8a', \n        'd0': '70', 'd1': '3e', 'd2': 'b5', 'd3': '66', 'd4': '48', 'd5': '03', 'd6': 'f6', 'd7': '0e', 'd8': '61', 'd9': '35', 'da': '57', 'db': 'b9', 'dc': '86', 'dd': 'c1', 'de': '1d', 'df': '9e', \n        'e0': 'e1', 'e1': 'f8', 'e2': '98', 'e3': '11', 'e4': '69', 'e5': 'd9', 'e6': '8e', 'e7': '94', 'e8': '9b', 'e9': '1e', 'ea': '87', 'eb': 'e9', 'ec': 'ce', 'ed': '55', 'ee': '28', 'ef': 'df', \n        'f0': '8c', 'f1': 'a1', 'f2': '89', 'f3': '0d', 'f4': 'bf', 'f5': 'e6', 'f6': '42', 'f7': '68', 'f8': '41', 'f9': '99', 'fa': '2d', 'fb': '0f', 'fc': 'b0', 'fd': '54', 'fe': 'bb', 'ff': '16'\n    }\n\n    sbox_inverse = {\n        '00': '52', '01': '09', '02': '6a', '03': 'd5', '04': '30', '05': '36', '06': 'a5', '07': '38', '08': 'bf', '09': '40', '0a': 'a3', '0b': '9e', '0c': '81', '0d': 'f3', '0e': 'd7', '0f': 'fb', \n        '10': '7c', '11': 'e3', '12': '39', '13': '82', '14': '9b', '15': '2f', '16': 'ff', '17': '87', '18': '34', '19': '8e', '1a': '43', '1b': '44', '1c': 'c4', '1d': 'de', '1e': 'e9', '1f': 'cb', \n        '20': '54', '21': '7b', '22': '94', '23': '32', '24': 'a6', '25': 'c2', '26': '23', '27': '3d', '28': 'ee', '29': '4c', '2a': '95', '2b': '0b', '2c': '42', '2d': 'fa', '2e': 'c3', '2f': '4e', \n        '30': '08', '31': '2e', '32': 'a1', '33': '66', '34': '28', '35': 'd9', '36': '24', '37': 'b2', '38': '76', '39': '5b', '3a': 'a2', '3b': '49', '3c': '6d', '3d': '8b', '3e': 'd1', '3f': '25', \n        '40': '72', '41': 'f8', '42': 'f6', '43': '64', '44': '86', '45': '68', '46': '98', '47': '16', '48': 'd4', '49': 'a4', '4a': '5c', '4b': 'cc', '4c': '5d', '4d': '65', '4e': 'b6', '4f': '92', \n        '50': '6c', '51': '70', '52': '48', '53': '50', '54': 'fd', '55': 'ed', '56': 'b9', '57': 'da', '58': '5e', '59': '15', '5a': '46', '5b': '57', '5c': 'a7', '5d': '8d', '5e': '9d', '5f': '84', \n        '60': '90', '61': 'd8', '62': 'ab', '63': '00', '64': '8c', '65': 'bc', '66': 'd3', '67': '0a', '68': 'f7', '69': 'e4', '6a': '58', '6b': '05', '6c': 'b8', '6d': 'b3', '6e': '45', '6f': '06', \n        '70': 'd0', '71': '2c', '72': '1e', '73': '8f', '74': 'ca', '75': '3f', '76': '0f', '77': '02', '78': 'c1', '79': 'af', '7a': 'bd', '7b': '03', '7c': '01', '7d': '13', '7e': '8a', '7f': '6b', \n        '80': '3a', '81': '91', '82': '11', '83': '41', '84': '4f', '85': '67', '86': 'dc', '87': 'ea', '88': '97', '89': 'f2', '8a': 'cf', '8b': 'ce', '8c': 'f0', '8d': 'b4', '8e': 'e6', '8f': '73', \n        '90': '96', '91': 'ac', '92': '74', '93': '22', '94': 'e7', '95': 'ad', '96': '35', '97': '85', '98': 'e2', '99': 'f9', '9a': '37', '9b': 'e8', '9c': '1c', '9d': '75', '9e': 'df', '9f': '6e', \n        'a0': '47', 'a1': 'f1', 'a2': '1a', 'a3': '71', 'a4': '1d', 'a5': '29', 'a6': 'c5', 'a7': '89', 'a8': '6f', 'a9': 'b7', 'aa': '62', 'ab': '0e', 'ac': 'aa', 'ad': '18', 'ae': 'be', 'af': '1b', \n        'b0': 'fc', 'b1': '56', 'b2': '3e', 'b3': '4b', 'b4': 'c6', 'b5': 'd2', 'b6': '79', 'b7': '20', 'b8': '9a', 'b9': 'db', 'ba': 'c0', 'bb': 'fe', 'bc': '78', 'bd': 'cd', 'be': '5a', 'bf': 'f4', \n        'c0': '1f', 'c1': 'dd', 'c2': 'a8', 'c3': '33', 'c4': '88', 'c5': '07', 'c6': 'c7', 'c7': '31', 'c8': 'b1', 'c9': '12', 'ca': '10', 'cb': '59', 'cc': '27', 'cd': '80', 'ce': 'ec', 'cf': '5f', \n        'd0': '60', 'd1': '51', 'd2': '7f', 'd3': 'a9', 'd4': '19', 'd5': 'b5', 'd6': '4a', 'd7': '0d', 'd8': '2d', 'd9': 'e5', 'da': '7a', 'db': '9f', 'dc': '93', 'dd': 'c9', 'de': '9c', 'df': 'ef', \n        'e0': 'a0', 'e1': 'e0', 'e2': '3b', 'e3': '4d', 'e4': 'ae', 'e5': '2a', 'e6': 'f5', 'e7': 'b0', 'e8': 'c8', 'e9': 'eb', 'ea': 'bb', 'eb': '3c', 'ec': '83', 'ed': '53', 'ee': '99', 'ef': '61', \n        'f0': '17', 'f1': '2b', 'f2': '04', 'f3': '7e', 'f4': 'ba', 'f5': '77', 'f6': 'd6', 'f7': '26', 'f8': 'e1', 'f9': '69', 'fa': '14', 'fb': '63', 'fc': '55', 'fd': '21', 'fe': '0c', 'ff': '7d'\n    }\n    \n    # select sbox and make bytes mutable\n    sb = sbox_inverse if inverse else sbox\n\n    for i in range(len(byte)):\n        # extract and pad hex code\n        code = hex(byte[i])[2:]\n        new_val = sb[ '0' * (2 - len(code)) + code ]\n        # replace original byte with new \n        byte[i] = int(new_val, 16)\n    \n    return byte\n\n\ndef get_key(current_state : list = None, round_const : int = 1) -> tuple:\n    ''' \n    Generates a 128 or 256 bit AES key for current round\n\n    Parameters\n    ----------------\n    current_state (type: list of bytearrays)\n    - The current key grouped by 32 bit vectors\n    round_const (type: integer)\n    - The round constant for the last round\n\n    Returns\n    ----------------\n    current_state (type: list of bytearrays)\n    - The next key grouped by 32 bit vectors\n    round_const (type: integer)\n    - The round constant for the current round\n    '''\n\n    # rotate, sub, and add constant to last vector\n    transformed = rotate_bytes(current_state[-1])    \n    transformed = sub_bytes(transformed)\n    transformed[0] ^= round_const\n\n    # xor vectors until end of round\n    for i in range(len(current_state) - 1):\n        current_state[i] = xor_bytes(current_state[i], transformed)\n        transformed = current_state[i]\n\n    # update round constant\n    round_const = update_round_const(round_const)\n\n    return current_state, round_const\n\ndef encrypt_model(model : torch.nn.Module, key_256 : bool) -> None:\n    '''\n    Generates a key and uses it to encrypt model parameters.\n\n    Parameters\n    ----------------\n    model (type: torch.nn.Module)\n    - The model to encrypt\n    key_256 (type: boolean)\n    - Whether to use a 128 or 256 bit key\n\n    Returns\n    ----------------\n    master_key (type: bytes)\n    - The key you'll need to decrypt the model\n    '''\n    \n    # Each vec has 32 bits to get 128 or 256 bit key\n    num_vecs = 8 if key_256 else 4\n    master_key = bytearray()\n    key = []\n\n    # Initialise master key\n    for i in range(num_vecs):\n        vector = bytearray()\n        for j in range(4):\n            vector.append(random.randint(0, 255)) \n\n        # Save key bytes       \n        key.append(vector)\n        master_key += vector\n\n    # Get first round key to start encryption\n    key, round_const = get_key(key)\n\n    # Go through model weights\n    for param in model.parameters():\n        print(\".\", end=\"\")\n        # Convert parameter tensor to bytes\n        param_bytes = bytearray(param.data.numpy().tobytes())\n\n        # Encrypt 4 floats or doubles at a time\n        for i in range(0, len(param_bytes), num_vecs * 4):\n            param_bytes[i:i + num_vecs * 4] = sub_bytes(bytearray(\n                xor_bytes(param_bytes[i:i + num_vecs * 4], b''.join(key))\n            ))\n\n            # Update key\n            key, round_const = get_key(key, round_const)\n\n        # Update parameter\n        dtype = np.float32 if param.data.dtype == torch.float32 else np.float64\n        param.data = torch.from_numpy(\n            np.frombuffer(param_bytes, dtype=dtype).reshape(param.data.shape)\n        )\n    \n    return bytes(master_key)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T01:23:42.89211Z","iopub.execute_input":"2023-05-11T01:23:42.892509Z","iopub.status.idle":"2023-05-11T01:23:42.955663Z","shell.execute_reply.started":"2023-05-11T01:23:42.892475Z","shell.execute_reply":"2023-05-11T01:23:42.954473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Encrypt Model Parameters","metadata":{}},{"cell_type":"code","source":"# Init dataset\ndataset = TensorDataset(encrypt_imgs, encrypt_labels)\nencrypt_data = DataLoader(dataset, batch_size=32)\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)\n\n# Encrypt model\nmaster_key = encrypt_model(model, False)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T01:25:41.114036Z","iopub.execute_input":"2023-05-11T01:25:41.114933Z","iopub.status.idle":"2023-05-11T01:26:46.072379Z","shell.execute_reply.started":"2023-05-11T01:25:41.114894Z","shell.execute_reply":"2023-05-11T01:26:46.071183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Show master key\nprint_bytes(bytes(master_key))\n\n# Show loss after encryption (NaN since values keep exploding due to very large floats)\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-05-11T01:27:59.833816Z","iopub.execute_input":"2023-05-11T01:27:59.834578Z","iopub.status.idle":"2023-05-11T01:28:02.76934Z","shell.execute_reply.started":"2023-05-11T01:27:59.834534Z","shell.execute_reply":"2023-05-11T01:28:02.768166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Decrypt Parameters","metadata":{}},{"cell_type":"code","source":"def decrypt_model(model : torch.nn.Module, key : bytes) -> None:\n    '''\n    Decrypts model using provided key\n\n    Parameters\n    ----------------\n    model (type: torch.nn.Module)\n    - The model to decrypt\n    key (type: bytes)\n    - The key to decrypt the model\n\n    Returns\n    ----------------\n    None\n    '''\n    \n    # Init round key and constant\n    key = bytearray(key)\n    key = [key[i:i + 4] for i in range(0, len(key), 4)]\n    key, round_const = get_key(key)\n\n    # Go through model weights\n    for param in model.parameters():\n        print(\".\", end=\"\")\n        # Convert parameter tensor to bytes\n        param_bytes = bytearray(param.data.numpy().tobytes())\n\n        # Decrypt 4 floats or doubles at a time\n        for i in range(0, len(param_bytes), len(key) * 4):\n            param_bytes[i:i + len(key) * 4] = xor_bytes(\n                 b''.join(key), \n                sub_bytes(param_bytes[i:i + len(key) * 4], inverse = True) \n            )\n\n            # Update key\n            key, round_const = get_key(key, round_const)\n\n        # Update parameter\n        dtype = np.float32 if param.data.dtype == torch.float32 else np.float64\n        param.data = torch.from_numpy(\n            np.frombuffer(param_bytes, dtype=dtype).reshape(param.data.shape)\n        )","metadata":{"execution":{"iopub.status.busy":"2023-05-11T01:28:59.657667Z","iopub.execute_input":"2023-05-11T01:28:59.658195Z","iopub.status.idle":"2023-05-11T01:28:59.67203Z","shell.execute_reply.started":"2023-05-11T01:28:59.658103Z","shell.execute_reply":"2023-05-11T01:28:59.670772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"decrypt_model(model, bytes(master_key))","metadata":{"execution":{"iopub.status.busy":"2023-05-11T01:29:01.339573Z","iopub.execute_input":"2023-05-11T01:29:01.340556Z","iopub.status.idle":"2023-05-11T01:30:03.475495Z","shell.execute_reply.started":"2023-05-11T01:29:01.340502Z","shell.execute_reply":"2023-05-11T01:30:03.471769Z"},"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-05-11T01:30:10.829494Z","iopub.execute_input":"2023-05-11T01:30:10.830106Z","iopub.status.idle":"2023-05-11T01:30:13.35113Z","shell.execute_reply.started":"2023-05-11T01:30:10.830068Z","shell.execute_reply":"2023-05-11T01:30:13.349938Z"},"trusted":true},"execution_count":null,"outputs":[]}]}