{"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":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","scrolled":true,"execution":{"iopub.status.busy":"2023-04-23T04:37:17.676023Z","iopub.execute_input":"2023-04-23T04:37:17.677297Z","iopub.status.idle":"2023-04-23T04:37:17.684118Z","shell.execute_reply.started":"2023-04-23T04:37:17.677243Z","shell.execute_reply":"2023-04-23T04:37:17.682830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install tflite_runtime==2.9.1","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:20.958753Z","iopub.execute_input":"2023-04-23T04:37:20.959140Z","iopub.status.idle":"2023-04-23T04:37:32.333723Z","shell.execute_reply.started":"2023-04-23T04:37:20.959105Z","shell.execute_reply":"2023-04-23T04:37:32.332300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install onnx\n!pip install onnx-tf","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:32.336511Z","iopub.execute_input":"2023-04-23T04:37:32.337261Z","iopub.status.idle":"2023-04-23T04:37:53.550914Z","shell.execute_reply.started":"2023-04-23T04:37:32.337210Z","shell.execute_reply":"2023-04-23T04:37:53.549560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Make necessary imports","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport pandas as pd\nimport numpy as np\n\ndevice = torch.device('cuda')\n\n\n!git clone https://github.com/codejaeger/plotly-wrapper\nimport sys\nsys.path.append('/kaggle/working/plotly-wrapper')\nfrom plotlywrapper import *","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:53.553840Z","iopub.execute_input":"2023-04-23T04:37:53.554618Z","iopub.status.idle":"2023-04-23T04:37:53.561408Z","shell.execute_reply.started":"2023-04-23T04:37:53.554564Z","shell.execute_reply":"2023-04-23T04:37:53.560140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Set the input paths to dataset","metadata":{}},{"cell_type":"code","source":"BASE_DIR = \"/kaggle/input/asl-signs\"\nLANDMARK_FILES_DIR = f\"{BASE_DIR}/train_landmark_files/\"\nPROCESSED_DATA_DIR = f\"{BASE_DIR}/processed_landmark_files/\"\nTRAIN_FILE = f\"{BASE_DIR}/train.csv\"","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:53.564437Z","iopub.execute_input":"2023-04-23T04:37:53.564839Z","iopub.status.idle":"2023-04-23T04:37:53.573357Z","shell.execute_reply.started":"2023-04-23T04:37:53.564808Z","shell.execute_reply":"2023-04-23T04:37:53.572301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Data reading function (provided by the competition evaluation page)","metadata":{}},{"cell_type":"code","source":"ROWS_PER_FRAME = 543\n\ndef load_relevant_data_subset(pq_path):\n    data_columns = ['x', 'y', 'z']\n    data = pd.read_parquet(pq_path, columns=data_columns)\n    n_frames = int(len(data) / ROWS_PER_FRAME)\n    data = data.values.reshape(n_frames, ROWS_PER_FRAME, len(data_columns))\n    return data.astype(np.float32)","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:53.575090Z","iopub.execute_input":"2023-04-23T04:37:53.576968Z","iopub.status.idle":"2023-04-23T04:37:53.585758Z","shell.execute_reply.started":"2023-04-23T04:37:53.576938Z","shell.execute_reply":"2023-04-23T04:37:53.584498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Contains right hand only frames.\nsample1 = load_relevant_data_subset(f\"{LANDMARK_FILES_DIR}/16069/1119409279.parquet\") ","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:53.587735Z","iopub.execute_input":"2023-04-23T04:37:53.588065Z","iopub.status.idle":"2023-04-23T04:37:53.617678Z","shell.execute_reply.started":"2023-04-23T04:37:53.588031Z","shell.execute_reply":"2023-04-23T04:37:53.616548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Add pytorch preprocessing layer\nclass PreprocessingLayerG(nn.Module):\n    \"\"\" Custom preprocessing layer \"\"\"\n    def __init__(self):\n        super(PreprocessingLayerG, self).__init__()\n        \n        # The output frame sequence size\n        self.SEQUENCE_SIZE = 32\n        \n        # The number of final landmarks to choose, this has to stay fixed.\n        # Change the curated landmarks to choose specific ones you need.\n        self.NUM_LANDMARKS = 32\n        \n        # How many frames to skip to detect difference in position as movement.\n        self.skip_frame_rate = 1\n        \n        # Weights used for landmarks to detect movement between frames.\n        kp_wgts = torch.cat([torch.tensor([0.1, 0.5, 0.9], device=device),\n                                torch.ones((15,), device=device),\n                                torch.zeros((11,), device=device),\n                                torch.tensor([0.05, 0.1, 0.2], device=device)], axis=-1)\n        self.kp_wgts_norm = kp_wgts / torch.sum(kp_wgts)\n        \n        # We only choose one of the hands (preference to left from dataset) since the \n        # authors specify signs should be detected with just one arm.\n        self.curated_landmarks = torch.tensor([\n            0, 1, 2, # The signing arm (mostly left)\n            3, 4, 5, 6, 7, 8, 9, 10, 11, 12, # (most essential hand keypoints required for signing wrist, thumb tip, other finger knuckles and tips)\n            13, 14, 15, 16, 17, # (extra hand keypoints if your model can handle larger inputs)\n            18, # nose\n            19, 20, # eyes\n            21, 22, 23, 24, 25, 26, 27, 28, # most important lip outline landmarks\n            29, 30, 31, # the other arm (mostly right)\n        ], dtype=torch.long, device=device)\n        \n        # A mapping from mediapipe body to MMpose body landmarks.\n        # We care only about the first value of each pair.\n        seleted_mapping_2 = torch.tensor([\n            (0, 0),\n            (11, 5),\n            (13, 7),\n            (15, 9), # LEFT ARM\n            (2, 1),\n            (5, 2), # HEAD - eyes \n            (78, 71), # RIGHT LIP END @ 6th\n            (308, 77), # LEFT LIP END\n            (41, 84), \n            (12, 85), \n            (271, 86), \n            (179, 90), \n            (15, 89), \n            (403, 88), # LIPS  @ 13th\n            (0, 91),\n            (4, 95),\n            (5, 96),\n            (8, 99),\n            (9, 100),\n            (12, 103),\n            (13, 104),\n            (16, 107),\n            (17, 108),\n            (20, 111),\n            (2, 93),\n            (6, 97),\n            (10, 101),\n            (14, 105),\n            (18, 109), # LEFT HAND\n            (12, 6),\n            (14, 8),\n            (16, 10), # OTHER ARM\n        ], dtype=torch.long, device=device)\n        \n        # Actual landmark indices in the given input.\n        NOSE_IDXS0 = 489 + seleted_mapping_2[:1][:, 0]\n        EYES_IDXS0 = 489 + seleted_mapping_2[4:6][:, 0]\n        LIPS_IDXS0 = seleted_mapping_2[6:14][:, 0]\n        LEFT_HAND_IDXS0  = 468 + seleted_mapping_2[14:24][:, 0]\n        LEFT_HAND_EXTRA_IDXS0  = 468 + seleted_mapping_2[24:29][:, 0]\n        RIGHT_HAND_IDXS0 = 522 + seleted_mapping_2[14:24][:, 0]\n        RIGHT_HAND_EXTRA_IDXS0 = 522 + seleted_mapping_2[24:29][:, 0]\n        LEFT_ARM_IDXS0       = 489 + seleted_mapping_2[1:4][:, 0]\n        RIGHT_ARM_IDXS0      = 489 + seleted_mapping_2[-3:][:, 0]\n        \n        # Left hand only landmarks\n        self.LEFT_HANDA_IDXS0 = torch.cat((LEFT_HAND_IDXS0,\n                                   LEFT_HAND_EXTRA_IDXS0))\n\n        # Right hand only landmarks\n        self.RIGHT_HANDA_IDXS0 = torch.cat((RIGHT_HAND_IDXS0,\n                                           RIGHT_HAND_EXTRA_IDXS0))\n\n        # Left hand + body only landmarks\n        self.LEFT_SET = torch.cat((LEFT_ARM_IDXS0, # 3\n                                           LEFT_HAND_IDXS0, # 10\n                                           LEFT_HAND_EXTRA_IDXS0, # 5\n                                   NOSE_IDXS0, EYES_IDXS0, LIPS_IDXS0, # 11\n                                   RIGHT_ARM_IDXS0 # 3\n                                  ))\n\n        # Left hand + body only landmarks\n        self.RIGHT_SET = torch.cat((RIGHT_ARM_IDXS0, \n                                           RIGHT_HAND_IDXS0,\n                                           RIGHT_HAND_EXTRA_IDXS0,\n                                   NOSE_IDXS0, EYES_IDXS0, LIPS_IDXS0,\n                                    LEFT_ARM_IDXS0\n                                   ))\n        \n        # Harcoding these since torch.arange is not supported in tflite conversion.\n        self.ar_18 = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, \\\n                                   8, 9, 10, 11, 12, 13, \\\n                                   14, 15, 16, 17], dtype=torch.long, device=device)\n        self.ar_18_29 = torch.tensor([18, 19, 20, 21, 22, \\\n                                      23, 24, 25, 26, 27, 28], dtype=torch.long, device=device)\n        self.ar_29_32 = torch.tensor([29, 30, 31], dtype=torch.long, device=device)\n        self.frame_sizes = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35, 36, 37, 38, 39, 40, 41, 42, 43, 44, 45, 46, 47, 48, 49, 50, 51, 52, 53, 54, 55, 56, 57, 58, 59, 60, 61, 62, 63, 64, 65, 66, 67, 68, 69, 70, 71, 72, 73, 74, 75, 76, 77, 78, 79, 80, 81, 82, 83, 84, 85, 86, 87, 88, 89, 90, 91, 92, 93, 94, 95, 96, 97, 98, 99, \n                                         100, 101, 102, 103, 104, 105, 106, 107, 108, 109, 110, 111, 112, 113, 114, 115, 116, 117, 118, 119, 120, 121, 122, 123, 124, 125, 126, 127, 128, 129, 130, 131, 132, 133, 134, 135, 136, 137, 138, 139, 140, 141, 142, 143, 144, 145, 146, 147, 148, 149, 150, 151, 152, 153, 154, 155, 156, 157, 158, 159, 160, 161, 162, 163, 164, 165, 166, 167, 168, 169, 170, 171, 172, 173, 174, 175, 176, 177, 178, 179, 180, 181, 182, 183, 184, 185, 186, 187, 188, 189, 190, 191, 192, 193, 194, 195, 196, 197, 198, 199, \n                                         200, 201, 202, 203, 204, 205, 206, 207, 208, 209, 210, 211, 212, 213, 214, 215, 216, 217, 218, 219, 220, 221, 222, 223, 224, 225, 226, 227, 228, 229, 230, 231, 232, 233, 234, 235, 236, 237, 238, 239, 240, 241, 242, 243, 244, 245, 246, 247, 248, 249, 250, 251, 252, 253, 254, 255, 256, 257, 258, 259, 260, 261, 262, 263, 264, 265, 266, 267, 268, 269, 270, 271, 272, 273, 274, 275, 276, 277, 278, 279, 280, 281, 282, 283, 284, 285, 286, 287, 288, 289, 290, 291, 292, 293, 294, 295, 296, 297, 298, 299, 300, 301, 302, 303, 304, 305, 306, 307, 308, 309, 310, 311, 312, 313, 314, 315, 316, 317, 318, 319, 320, 321, 322, 323, 324, 325, 326, 327, 328, 329, 330, 331, 332, 333, 334, 335, 336, 337, 338, 339, 340, 341, 342, 343, 344, 345, 346, 347, 348, 349, 350, 351, 352, 353, 354, 355, 356, 357, 358, 359, 360, 361, 362, 363, 364, 365, 366, 367, 368, 369, 370, 371, 372, 373, 374, 375, 376, 377, 378, 379, \n                                         380, 381, 382, 383, 384, 385, 386, 387, 388, 389, 390, 391, 392, 393, 394, 395, 396, 397, 398, 399, 400, 401, 402, 403, 404, 405, 406, 407, 408, 409, 410, 411, 412, 413, 414, 415, 416, 417, 418, 419, 420, 421, 422, 423, 424, 425, 426, 427, 428, 429, 430, 431, 432, 433, 434, 435, 436, 437, 438, 439, 440, 441, 442, 443, 444, 445, 446, 447, 448, 449, 450, 451, 452, 453, 454, 455, 456, 457, 458, 459, 460, 461, 462, 463, 464, 465, 466, 467, 468, 469, 470, 471, 472, 473, 474, 475, 476, 477, 478, 479, 480, 481, 482, 483, 484, 485, 486, 487, 488, 489, 490, 491, 492, 493, 494, 495, 496, 497, 498, 499, 500, 501, 502, 503, 504, 505, 506, 507, 508, 509, 510, 511, 512, 513, 514, 515, 516, 517, 518, 519, 520, 521, 522, 523, 524, 525, 526, 527, 528, 529, 530, 531, 532, 533, 534, 535, 536, 537, 538, 539, 540, 541, 542, 543, 544, 545, 546, 547, 548, 549, 550, 551, 552, 553, 554, 555, 556, 557, 558, 559, \n                                         560, 561, 562, 563, 564, 565, 566, 567, 568, 569, 570, 571, 572, 573, 574, 575, 576, 577, 578, 579, 580, 581, 582, 583, 584, 585, 586, 587, 588, 589, 590, 591, 592, 593, 594, 595, 596, 597, 598, 599, 600], dtype=torch.long, device=device)\n        \n    # Flip the hand about the nose as center. \n    def flip_hand(self, data, flip_mask):\n        lm_indices = torch.concatenate([self.ar_18, self.ar_29_32])\n        lm_indices_msk = torch.zeros((data.size()[1],), device=device)\n        lm_indices_msk = lm_indices_msk.scatter_(0, lm_indices, 1).repeat(data.size()[0], 1).unsqueeze(dim=-1)\n\n        lm_indices_msk = torch.concat([lm_indices_msk, torch.zeros_like(lm_indices_msk, device=device), torch.zeros_like(lm_indices_msk, device=device)], axis=-1)\n        \n        flipped_hand_update = (2 * data[:, 18, 0][:, None, None] - data)\n        hn_data = (flip_mask[:, None, None] * flipped_hand_update + (1-flip_mask.float())[:, None, None] * data) * \\\n                    lm_indices_msk + (1-lm_indices_msk.float()) * data\n        return hn_data\n        \n    # Find the frames with no nans and atleast one nans for the given landmark indices.\n    def find_nans(self, data, lm_indices):\n        f, lm, _ = data.size()\n        num_ind = lm_indices.size()[0]\n        _data = data[:, lm_indices, :]\n        nan_sum = torch.sum(torch.isnan(_data), dim=[1, 2])\n        no_nan_mask = (nan_sum == 0)\n        some_nan_mask = torch.multiply(nan_sum > 0, nan_sum < num_ind * 3)\n        return no_nan_mask, some_nan_mask\n\n    # Normalise left hand only.\n    def normalise_left_hand(self, data):\n        lm_indices = self.ar_18[3:]\n        lm_indices_msk = torch.zeros((data.size()[1],), device=device)\n        lm_indices_msk = lm_indices_msk.scatter_(0, lm_indices, 1).repeat(data.size()[0], 1).unsqueeze(dim=-1)\n\n        # no normalization for z coordinate\n        lm_indices_msk = torch.concat([lm_indices_msk, lm_indices_msk, torch.zeros_like(lm_indices_msk, device=device)], axis=-1)\n        \n        hand_min = torch.min(data[:, lm_indices], axis=1)[0]\n        hand_max = torch.max(data[:, lm_indices], axis=1)[0]\n        \n        _hand_min = hand_min - (hand_max - hand_min) * 0.1\n        _hand_max = hand_max + (hand_max - hand_min) * 0.1\n        \n        n_hand_data = (data - _hand_min[:, None, :]) / (_hand_max - _hand_min + torch.tensor(1e-6, device=device))[:, None, :]\n        \n        hn_data = n_hand_data * lm_indices_msk + (1-lm_indices_msk.float()) * data\n        \n        return hn_data\n    \n    # Normalise the eyes, arms and lips landmarks.\n    def normalise_body(self, data):\n        lm_indices = torch.concat([self.ar_18[:3], self.ar_18_29, self.ar_29_32], axis=-1)\n        lm_indices_msk = torch.zeros((data.size()[1],), device=device)\n        lm_indices_msk = lm_indices_msk.scatter_(0, lm_indices, 1).repeat(data.size()[0], 1).unsqueeze(dim=-1)\n\n        # no normalization for z coordinate\n        lm_indices_msk = torch.concat([lm_indices_msk, lm_indices_msk, torch.zeros_like(lm_indices_msk, device=device)], axis=-1)\n\n        center = data[:, 18]\n        \n        h_unit = torch.linalg.norm(torch.abs(data[:, 29, :2] - data[:, 0, :2]), axis=-1) / 2\n        displ_min = torch.concat([3 * h_unit[:, None], 6 * h_unit[:, None], torch.zeros_like(h_unit[:, None], device=device)], axis=-1)\n        displ_max = torch.concat([3 * h_unit[:, None], 0.5 * h_unit[:, None], torch.zeros_like(h_unit[:, None], device=device)], axis=-1)\n        \n        body_min = center - displ_min\n        body_max = center + displ_max\n        \n        n_body_data = (data - body_min[:, None, :]) / (body_max - body_min + torch.tensor(1e-6, device=device))[:, None, :]\n        b_data = n_body_data * lm_indices_msk + (1-lm_indices_msk.float()) * data\n        \n        return b_data\n    \n    # torch.nan_to_num fails pytorch conversion\n    def nan_to_num_(self, x):\n        return torch.where(torch.isnan(x), torch.zeros_like(x, device=device), x)\n    \n    # Interpolate missing landmarks for intermediate frames wherever possible.\n    def interpolate_lms(self, data, data_isnan, lm_indices, nonempty_mask, p_nonempty_mask):\n        empty_mask = (1-nonempty_mask.float())\n        _nonempty_mask = torch.where(torch.sum(nonempty_mask) == 0, empty_mask.bool(), nonempty_mask)\n        \n        nonempty_idx = torch.nonzero(_nonempty_mask).squeeze()\n        \n        lm_indices_msk = torch.zeros((data.size()[1],), device=device)\n        lm_indices_msk = lm_indices_msk.scatter_(0, lm_indices, 1).repeat(empty_mask.size()[0], 1).unsqueeze(dim=-1)\n\n        lm_indices_msk = torch.concat([lm_indices_msk, lm_indices_msk, lm_indices_msk], axis=-1)\n        \n        _ex = nonempty_idx.repeat(nonempty_mask.size()[0], 1)\n        \n        poss = self.frame_sizes[:nonempty_mask.size()[0]].unsqueeze(dim=-1)\n        high = _ex - poss\n        low = -1 * high\n        \n        high = (self.SEQUENCE_SIZE+1000) * (high < 0) + high * (high > 0)\n        low = (self.SEQUENCE_SIZE+1000) * (low < 0) + low * (low > 0)\n        \n        closest_high = torch.clamp(torch.min(high, dim=-1)[0], max=self.SEQUENCE_SIZE)\n        closest_high_msk = (closest_high==self.SEQUENCE_SIZE)\n        closest_high = (1-closest_high_msk.float()) * closest_high\n\n        closest_low = torch.clamp(torch.min(low, dim=-1)[0], max=self.SEQUENCE_SIZE)\n        closest_low_msk = (closest_low==self.SEQUENCE_SIZE)\n        closest_low = (1-closest_low_msk.float()) * closest_low\n        \n        high_pos = (closest_high + poss.squeeze()).long()\n        low_pos = (poss.squeeze() - closest_low).long()\n\n        intrpolable_mask = (closest_high >=1) * (closest_low >=1) * (high_pos > low_pos) * empty_mask\n\n        p_intrpolable_mask = intrpolable_mask * p_nonempty_mask\n        f_intrpolable_mask = torch.logical_xor(intrpolable_mask, p_intrpolable_mask)\n        \n        dif = data[high_pos] - data[low_pos]\n        val = torch.zeros_like(data, device=device)\n        val = data[low_pos] + dif * (closest_low / (closest_high + closest_low))[:, None, None]\n              \n\n        out = data.clone()\n        out = self.nan_to_num_(out)\n        val = self.nan_to_num_(val)\n\n        out = (nonempty_mask[:, None, None] * out + \\\n                (p_intrpolable_mask[:, None, None] * data_isnan * val + \\\n                p_intrpolable_mask[:, None, None] * (1-data_isnan.float()) * out) + \\\n                f_intrpolable_mask[:, None, None] * val) * lm_indices_msk + \\\n                self.nan_to_num_(data) * (1-lm_indices_msk.float())\n\n        final_non_empty_mask = torch.logical_or(nonempty_mask, intrpolable_mask)\n\n        return out, final_non_empty_mask\n    \n    # TFlite acos not supported\n    def acos(self, x):\n        return torch.sqrt(2 * torch.clamp((1 - x), min=0))\n        \n    # Identify key postures for large frame sequences and return most non-static frames.\n    def identify_key_postures(self, data, skip):\n        indexs = self.frame_sizes[:data.size()[0]:skip]\n        transit = data[indexs[1:]] - data[indexs[:-1]]\n                                \n        transit_ = torch.sum(transit * self.kp_wgts_norm[None, :, None], axis=1)\n        nrm = torch.norm(transit_[1:, :2], dim=-1) * torch.norm(transit_[:-1, :2], dim=-1) + torch.tensor(1e-6, device=device)\n        score = nrm * self.acos(torch.abs(torch.sum(transit_[1:, :2] * transit_[:-1, :2], axis=-1)) / nrm)\n        \n        sorted_frames = torch.argsort(score, dim=0, descending=True)\n        important_frames = torch.concat([torch.zeros((1,), device=device), indexs[sorted_frames+1]], axis=-1)[:self.SEQUENCE_SIZE]\n\n        frames, _ = torch.sort(important_frames)\n        \n        return frames.long()\n    \n    def forward(self, data0):\n        # Number of Frames in Video\n        N_FRAMES0 = data0.size()[0]\n        \n        non_empty_lhand_mask, p_non_empty_lhand_mask = self.find_nans(data0, self.LEFT_HANDA_IDXS0)\n        non_empty_rhand_mask, p_non_empty_rhand_mask = self.find_nans(data0, self.RIGHT_HANDA_IDXS0)\n        \n        out = torch.zeros((N_FRAMES0, self.NUM_LANDMARKS, 3), device=device)\n        \n        # interpolate\n        new_hand_idx = self.ar_18\n        \n        l_out, l_ind_mask = self.interpolate_lms(self.nan_to_num_(data0[:, self.LEFT_SET, :]), \\\n                                                 torch.isnan(data0[:, self.LEFT_SET, :]), \\\n                                                 new_hand_idx, non_empty_lhand_mask, p_non_empty_lhand_mask)\n\n        r_out, r_ind_mask = self.interpolate_lms(self.nan_to_num_(data0[:, self.RIGHT_SET, :]), \\\n                                                 torch.isnan(data0[:, self.RIGHT_SET, :]), \\\n                                                 new_hand_idx, non_empty_rhand_mask, p_non_empty_rhand_mask)\n        \n        final_r_ind_mask = torch.logical_xor(torch.logical_or(r_ind_mask, l_ind_mask), l_ind_mask)\n        \n        # Deal with left-right hand conundrum\n        out = out + l_ind_mask[:, None, None] * l_out + \\\n                    self.flip_hand(r_out, final_r_ind_mask) * final_r_ind_mask[:, None, None]\n        \n        final_non_empty_mask = torch.logical_or(l_ind_mask, final_r_ind_mask)\n\n        final_non_empty_indx = torch.nonzero(final_non_empty_mask).squeeze()\n        out = torch.reshape(self.nan_to_num_(out[final_non_empty_indx, :, :]), (-1, 32, 3))\n        \n        out = self.normalise_left_hand(out)\n        out = self.normalise_body(out)\n        \n        N_FRAMES = out.size()[0]\n        \n        out = out[self.identify_key_postures(out, skip=self.skip_frame_rate)]\n\n        K_N_FRAMES = out.size()[0]\n        diff_pad =  self.SEQUENCE_SIZE - K_N_FRAMES\n        diff_tnsr = torch.zeros((diff_pad, self.NUM_LANDMARKS, 3), device=device)\n        out = torch.concat([out, diff_tnsr], dim=0)\n\n        return out[:, self.curated_landmarks, :]","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:53.620191Z","iopub.execute_input":"2023-04-23T04:37:53.621115Z","iopub.status.idle":"2023-04-23T04:37:53.703749Z","shell.execute_reply.started":"2023-04-23T04:37:53.621075Z","shell.execute_reply":"2023-04-23T04:37:53.702478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pplg = PreprocessingLayerG()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:53.705850Z","iopub.execute_input":"2023-04-23T04:37:53.706359Z","iopub.status.idle":"2023-04-23T04:37:53.723340Z","shell.execute_reply.started":"2023-04-23T04:37:53.706316Z","shell.execute_reply":"2023-04-23T04:37:53.722192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"outg = pplg(torch.from_numpy(sample1).to(device))","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:53.725568Z","iopub.execute_input":"2023-04-23T04:37:53.726912Z","iopub.status.idle":"2023-04-23T04:37:53.741831Z","shell.execute_reply.started":"2023-04-23T04:37:53.726855Z","shell.execute_reply":"2023-04-23T04:37:53.740870Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Numerical checks\ndef print_checks(ot, mes):\n    print(mes, ot[0, 0], torch.sum(ot), torch.mean(ot))\n    \ndef tf_print_checks(ot, mes):\n    print(mes, ot[0, 0], tf.math.reduce_sum(ot), tf.reduce_mean(ot))\n\nprint_checks(outg, \"sm\")","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:53.746125Z","iopub.execute_input":"2023-04-23T04:37:53.746412Z","iopub.status.idle":"2023-04-23T04:37:53.759191Z","shell.execute_reply.started":"2023-04-23T04:37:53.746384Z","shell.execute_reply":"2023-04-23T04:37:53.757923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_eggds = np.array([(18, 0), (0, 1), (1, 2), \n                  (18, 29), (29, 30), (30, 31),\n                  (3, 13), (13, 4),\n                  (3, 5), (5, 14), (14, 6),\n                  (3, 7), (7, 15), (15, 8),\n                  (3, 9), (9, 16), (16, 10),\n                  (3, 11), (11, 17), (17, 12)])\n\nfor i in range(0, 10, 3):\n#     display(make_plot(Scatter2D(pts=_outs[3][i][18:, :2], text=3) + \\\n#             Scatter2D(pts=_outs[3][i][:3, :2])\n# #           Wireframe(v=_outs[3][i][3:18, :2], e=eggds)\n#          ).scale_axis_to_same().invert_y())\n    \n    display(make_plot(Scatter2D(pts=outg.cpu().numpy()[i][:, :2]) +\n          Wireframe(v=outg.cpu().numpy()[i][:, :2], e=all_eggds)\n         ).scale_axis_to_same().invert_y())\n    \n#     display(make_plot(Scatter2D(pts=_outs[3][i][:, :2]) +\n#           Wireframe(v=_outs[3][i][:, :2], e=_eggds)\n#          ).scale_axis_to_same().invert_y())","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-04-23T04:37:53.760989Z","iopub.execute_input":"2023-04-23T04:37:53.761277Z","iopub.status.idle":"2023-04-23T04:37:53.991141Z","shell.execute_reply.started":"2023-04-23T04:37:53.761249Z","shell.execute_reply":"2023-04-23T04:37:53.990044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Exporting to ONNX -> Tensorflow -> TFlite","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:03:32.377068Z","iopub.execute_input":"2023-04-23T04:03:32.378117Z","iopub.status.idle":"2023-04-23T04:03:32.396924Z","shell.execute_reply.started":"2023-04-23T04:03:32.378068Z","shell.execute_reply":"2023-04-23T04:03:32.395280Z"}}},{"cell_type":"code","source":"torch.onnx.export(pplg, torch.zeros((32, 543, 3)).to('cuda'), \"simplemodel2.onnx\", input_names=['input'], output_names=['output'], do_constant_folding=True, opset_version=12, dynamic_axes={\"input\": {0: \"batch\"}})","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:53.992711Z","iopub.execute_input":"2023-04-23T04:37:53.993590Z","iopub.status.idle":"2023-04-23T04:37:54.709866Z","shell.execute_reply.started":"2023-04-23T04:37:53.993542Z","shell.execute_reply":"2023-04-23T04:37:54.708630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from onnx_tf.backend import prepare\nimport onnx\n\nonnx_model = onnx.load(\"simplemodel2.onnx\")\n\ntf_rep = prepare(onnx_model, device='gpu') \ntf_rep.export_graph(\"model_tf2\") ","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:37:54.711772Z","iopub.execute_input":"2023-04-23T04:37:54.712182Z","iopub.status.idle":"2023-04-23T04:38:10.970400Z","shell.execute_reply.started":"2023-04-23T04:37:54.712139Z","shell.execute_reply":"2023-04-23T04:38:10.969164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\ntf_torch = tf.keras.models.load_model('model_tf2')","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:38:10.972139Z","iopub.execute_input":"2023-04-23T04:38:10.974413Z","iopub.status.idle":"2023-04-23T04:38:12.886542Z","shell.execute_reply.started":"2023-04-23T04:38:10.974368Z","shell.execute_reply":"2023-04-23T04:38:12.885386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m_inp = torch.from_numpy(sample1).cpu().numpy()\nprint(\"Raw input size\", m_inp.shape)\n\nott = tf_torch(input=m_inp)['output']\nprint(\"Output size\", ott.shape)\n\ntf_print_checks(ott, \"\\nSome check string\")","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:38:12.895481Z","iopub.execute_input":"2023-04-23T04:38:12.895823Z","iopub.status.idle":"2023-04-23T04:38:14.371810Z","shell.execute_reply.started":"2023-04-23T04:38:12.895791Z","shell.execute_reply":"2023-04-23T04:38:14.369532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Great it matches the torch version","metadata":{}},{"cell_type":"markdown","source":"At this point you can refer to https://www.kaggle.com/code/hengck23/lb-0-67-one-pytorch-transformer-solution?scriptVersionId=122239639&cellId=5\nto understand how to include this layer into your Tensorflow model. I am skipping it here to keep this notebook clean.","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nconverter = tf.lite.TFLiteConverter.from_saved_model(\"model_tf2\")\n\nconverter.target_spec.supported_ops = [\n#   tf.lite.OpsSet.TFLITE_BUILTINS, # enable TensorFlow Lite ops.\n#   tf.lite.OpsSet.SELECT_TF_OPS # enable TensorFlow ops.\n]\n\nconverter.optimizations = [tf.lite.Optimize.DEFAULT]\n# converter.target_spec.supported_types = [tf.float16]\n\ntflite_quant_model = converter.convert()\n\ntflite_model = converter.convert()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:38:14.373328Z","iopub.execute_input":"2023-04-23T04:38:14.373931Z","iopub.status.idle":"2023-04-23T04:38:21.189385Z","shell.execute_reply.started":"2023-04-23T04:38:14.373894Z","shell.execute_reply":"2023-04-23T04:38:21.188187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open('./model_tfl2.9_sing.tflite', 'wb') as f:\n    f.write(tflite_model)\n!du -sh ./model_tfl2.9_sing.tflite","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:38:21.191460Z","iopub.execute_input":"2023-04-23T04:38:21.192328Z","iopub.status.idle":"2023-04-23T04:38:22.328364Z","shell.execute_reply.started":"2023-04-23T04:38:21.192283Z","shell.execute_reply":"2023-04-23T04:38:22.326937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Its a very small addition to the model size!","metadata":{}},{"cell_type":"code","source":"import tflite_runtime.interpreter as tflite\ninterpreter = tflite.Interpreter('./model_tfl2.9_sing.tflite')\n\nfound_signatures = list(interpreter.get_signature_list().keys())\n\nif 'serving_default' not in found_signatures:\n    raise KernelEvalException('Required input signature not found.')\n\nprediction_fn = interpreter.get_signature_runner(\"serving_default\")\n\noutput = prediction_fn(input=m_inp)\ntf_print_checks(output['output'], \"message\")","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:38:22.330918Z","iopub.execute_input":"2023-04-23T04:38:22.331352Z","iopub.status.idle":"2023-04-23T04:38:22.356460Z","shell.execute_reply.started":"2023-04-23T04:38:22.331305Z","shell.execute_reply":"2023-04-23T04:38:22.355203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\nimport numpy as np\n\nn = 100\nt0 = time.time()\nfor i in range(n): output = prediction_fn(input=m_inp)\nt1 = time.time()\n(t1-t0)/n","metadata":{"execution":{"iopub.status.busy":"2023-04-23T04:38:22.358688Z","iopub.execute_input":"2023-04-23T04:38:22.359119Z","iopub.status.idle":"2023-04-23T04:38:22.718087Z","shell.execute_reply.started":"2023-04-23T04:38:22.359075Z","shell.execute_reply":"2023-04-23T04:38:22.716385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The preprocessing layer just adds 2 milliseconds to the total inference time per video, for 40k videos that is `2 * 40000 * 0.001 = 80 seconds` **<<<** `10 minutes` allowed for buffering and preprocessing overheads.","metadata":{}},{"cell_type":"markdown","source":"## Thanks for reading and please give an upvote if this notebook helped you in this competition!","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}}]}