{"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":"## Imports","metadata":{}},{"cell_type":"code","source":"!pip install onnxruntime\n!pip install onnx-tf\n!pip install tflite-runtime","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:03:51.080950Z","iopub.execute_input":"2023-04-16T02:03:51.081672Z","iopub.status.idle":"2023-04-16T02:04:27.262957Z","shell.execute_reply.started":"2023-04-16T02:03:51.081617Z","shell.execute_reply":"2023-04-16T02:04:27.261272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We first import the libraries used in developing and training the model. PyTorch is used to make the model.","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nimport json\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchinfo\n\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-16T02:04:27.270141Z","iopub.execute_input":"2023-04-16T02:04:27.270460Z","iopub.status.idle":"2023-04-16T02:04:28.655290Z","shell.execute_reply.started":"2023-04-16T02:04:27.270427Z","shell.execute_reply":"2023-04-16T02:04:28.654251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"These libraries are used to convert the PyTorch model to TensorFlow Lite. The model is converted from PyTorch to ONNX to TensorFlow to TFLite.","metadata":{}},{"cell_type":"code","source":"import onnx\nimport onnxruntime\nimport onnx_tf\nimport tensorflow as tf\nimport tflite_runtime.interpreter as tflite","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:04:28.656736Z","iopub.execute_input":"2023-04-16T02:04:28.657358Z","iopub.status.idle":"2023-04-16T02:04:36.451937Z","shell.execute_reply.started":"2023-04-16T02:04:28.657327Z","shell.execute_reply":"2023-04-16T02:04:36.450862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Setup path constants for Kaggle Notebook.","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT_DIR = Path('/kaggle/input/')\nASL_DIR = INPUT_DIR / 'asl-signs'\nDATASET_DIR = INPUT_DIR / 'asl-dataset'","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:04:36.454818Z","iopub.execute_input":"2023-04-16T02:04:36.456206Z","iopub.status.idle":"2023-04-16T02:04:36.462266Z","shell.execute_reply.started":"2023-04-16T02:04:36.456162Z","shell.execute_reply":"2023-04-16T02:04:36.461181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'mps' if torch.backends.mps.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:04:36.466250Z","iopub.execute_input":"2023-04-16T02:04:36.466541Z","iopub.status.idle":"2023-04-16T02:04:36.528076Z","shell.execute_reply.started":"2023-04-16T02:04:36.466510Z","shell.execute_reply":"2023-04-16T02:04:36.526979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loading Data","metadata":{}},{"cell_type":"markdown","source":"This is the competition provided function for loading data.","metadata":{}},{"cell_type":"code","source":"ROWS_PER_FRAME = 543  # number of landmarks per frame\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-16T02:04:37.985416Z","iopub.execute_input":"2023-04-16T02:04:37.986517Z","iopub.status.idle":"2023-04-16T02:04:37.994728Z","shell.execute_reply.started":"2023-04-16T02:04:37.986476Z","shell.execute_reply":"2023-04-16T02:04:37.993681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"To save time, the data used for training is preprocessed in a [separate notebook](https://www.kaggle.com/code/lameuler/asl-dataset) before being used here.","metadata":{}},{"cell_type":"code","source":"with (INPUT_DIR / 'asl-dataset' / 'landmarks.json').open() as f:\n    landmarks = json.load(f)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:04:37.654845Z","iopub.execute_input":"2023-04-16T02:04:37.655221Z","iopub.status.idle":"2023-04-16T02:04:37.662595Z","shell.execute_reply.started":"2023-04-16T02:04:37.655183Z","shell.execute_reply":"2023-04-16T02:04:37.661471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"signs_df = pd.read_csv(DATASET_DIR / 'train.csv')\nsigns_df","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:04:37.744796Z","iopub.execute_input":"2023-04-16T02:04:37.745742Z","iopub.status.idle":"2023-04-16T02:04:37.933676Z","shell.execute_reply.started":"2023-04-16T02:04:37.745701Z","shell.execute_reply":"2023-04-16T02:04:37.932583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Split the data for training and validation, with a 80:20 split.","metadata":{}},{"cell_type":"code","source":"train_df = signs_df.sample(frac=0.8)\ntest_df = signs_df.drop(train_df.index).sample(frac=1)\ntrain_df.shape, test_df.shape","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:04:37.935010Z","iopub.execute_input":"2023-04-16T02:04:37.937308Z","iopub.status.idle":"2023-04-16T02:04:37.973584Z","shell.execute_reply.started":"2023-04-16T02:04:37.937274Z","shell.execute_reply":"2023-04-16T02:04:37.972680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The PyTorch dataset is used to load the data and preprocess it.","metadata":{}},{"cell_type":"code","source":"class ASLDataset(Dataset):\n    def __init__(self, dataset_df, agg):\n        files = np.load(DATASET_DIR / 'data.npz')\n        self.items = [torch.Tensor(files[str(i)]).to(device) for i in tqdm(dataset_df.sequence_id, desc='Loading data', total=len(dataset_df))]\n        self.labels = torch.Tensor(dataset_df.label.values).long().to(device)\n        self.agg = agg\n    \n    def __len__(self):\n        return len(self.items)\n        \n    def __getitem__(self, index):\n        return self.agg(self.items[index]).float(), self.labels[index]","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:04:37.974839Z","iopub.execute_input":"2023-04-16T02:04:37.975284Z","iopub.status.idle":"2023-04-16T02:04:37.983819Z","shell.execute_reply.started":"2023-04-16T02:04:37.975245Z","shell.execute_reply":"2023-04-16T02:04:37.982580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"POINTS = torch.cat([torch.tensor(value).unfold(0,3,1) for value in landmarks.values()])","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:04:37.996137Z","iopub.execute_input":"2023-04-16T02:04:37.996836Z","iopub.status.idle":"2023-04-16T02:04:38.004535Z","shell.execute_reply.started":"2023-04-16T02:04:37.996794Z","shell.execute_reply":"2023-04-16T02:04:38.003611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INDICES = np.load(DATASET_DIR / 'indices.npy')","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:04:38.005981Z","iopub.execute_input":"2023-04-16T02:04:38.006495Z","iopub.status.idle":"2023-04-16T02:04:38.016914Z","shell.execute_reply.started":"2023-04-16T02:04:38.006456Z","shell.execute_reply":"2023-04-16T02:04:38.015817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare(x):\n    # Find \"angles\"\n    view = x[:,POINTS]\n    vectors = torch.stack(\n        (\n            view[...,1,:] - view[...,0,:],\n            view[...,2,:] - view[...,1,:]\n        ),\n        dim=-2\n    ).float()\n    angles = torch.div(\n        vectors.prod(dim=-2).sum(dim=-1),\n        vectors.square().sum(dim=-1).sqrt().prod(dim=-1)\n    )#.acos()\n\n    # Coordinate normalisation\n    coord_counts = (~x.isnan()).sum(dim=(0,1))\n    coord_no_nan = x.clone()\n    coord_no_nan[coord_no_nan.isnan()] = 0\n    coord_mean = coord_no_nan.sum(dim=(0,1)) / coord_counts\n    normed = x - coord_mean\n    #normed[normed.isnan()] = 0\n    #normed = nn.functional.normalize(normed, dim=-1)\n\n    # Coords + Angles\n    tensor = torch.cat((\n        normed.flatten(-2),\n        angles),1)\n\n    # Mean\n    counts = (~tensor.isnan()).sum(dim=0)\n    no_nan = tensor.clone()\n    no_nan[no_nan.isnan()] = 0\n    mean = no_nan.sum(dim=0) / counts\n\n    # Standard Deviation\n    diff = tensor - mean\n    diff[diff.isnan()] = 0\n    correction = 1\n    std = (diff.square().sum(dim=0) / (counts-correction)).float().sqrt()\n\n    out = torch.cat((mean,std))\n    out[out.isnan()] = 0\n    return out","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = ASLDataset(train_df, prepare)\ntrain_preload = train_dataset.items\ntest_dataset = ASLDataset(test_df, prepare)\ntest_preload = test_dataset.items\nlen(train_dataset), len(test_dataset)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:04:38.175024Z","iopub.execute_input":"2023-04-16T02:04:38.176022Z","iopub.status.idle":"2023-04-16T02:09:48.667975Z","shell.execute_reply.started":"2023-04-16T02:04:38.175984Z","shell.execute_reply":"2023-04-16T02:09:48.666892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Separate PyTorch dataloaders for training and validation.","metadata":{}},{"cell_type":"code","source":"train_dataloader = DataLoader(train_dataset, batch_size=128, shuffle=True)\ntest_dataloader = DataLoader(test_dataset, batch_size=128, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:09:48.669572Z","iopub.execute_input":"2023-04-16T02:09:48.670246Z","iopub.status.idle":"2023-04-16T02:09:48.676322Z","shell.execute_reply.started":"2023-04-16T02:09:48.670204Z","shell.execute_reply":"2023-04-16T02:09:48.675099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class MLPModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.sequential = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(534*2, 1024),\n            nn.BatchNorm1d(1024),\n            nn.Dropout(0.3),\n            nn.ReLU(),\n            nn.Linear(1024, 2048),\n            nn.BatchNorm1d(2048),\n            nn.Dropout(0.4),\n            nn.ReLU(),\n            nn.Linear(2048, 1024),\n            nn.BatchNorm1d(1024),\n            nn.Dropout(0.3),\n            nn.ReLU(),\n            nn.Linear(1024, 250)\n        )\n        \n    def forward(self, x):\n        # x = self.agg(x)\n        x = self.sequential(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:52:57.374425Z","iopub.execute_input":"2023-04-16T02:52:57.374824Z","iopub.status.idle":"2023-04-16T02:52:57.388462Z","shell.execute_reply.started":"2023-04-16T02:52:57.374785Z","shell.execute_reply":"2023-04-16T02:52:57.387206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"markdown","source":"`CrossEntropyLoss` is used with the `AdamW` optimizer and the `CosineAnnelingWarmRestarts` learning rate scheduler. A weight decay is used to reduce overfitting.","metadata":{}},{"cell_type":"code","source":"model = MLPModel().to(device)\nlearning_rate = 5e-4\nweight_decay = 0.1\ncycle = 8\n\nloss_fn = nn.CrossEntropyLoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=cycle, eta_min=learning_rate / 10)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T03:20:29.182744Z","iopub.execute_input":"2023-04-16T03:20:29.183664Z","iopub.status.idle":"2023-04-16T03:20:29.240947Z","shell.execute_reply.started":"2023-04-16T03:20:29.183610Z","shell.execute_reply":"2023-04-16T03:20:29.239956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This function loops through all the training batches and performs backpropagation, then it loops through all the validation batches to calculate metrics. A `tqdm` progress bar is used to track the training progress and display the metrics.","metadata":{}},{"cell_type":"code","source":"def train_val_loop(epoch, train_dataloader, val_dataloader, model, loss_fn, optimizer, scheduler, n_offset=1):\n    total_batches = len(train_dataloader)\n    train_size, train_batches = 0, 0\n    train_loss, train_correct = 0, 0\n    val_size, val_batches = 0, 0\n    val_loss, val_correct = 0, 0\n    \n    with tqdm(desc=f'Epoch {epoch+n_offset}', total=total_batches) as bar:\n        for batch, (X, y) in enumerate(train_dataloader):\n            # Compute prediction and loss\n            pred = model(X)\n            loss = loss_fn(pred, y)\n            train_loss += loss.item()\n            train_correct += (pred.argmax(1) == y).type(torch.float).sum().item()\n            train_size += len(y)\n            train_batches += 1\n\n            # Backpropagation\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n            scheduler.step(epoch + batch / total_batches)\n            \n            bar.update()\n\n            if batch % 10 == 0:\n                bar.set_postfix(accuracy = train_correct / train_size, loss = train_loss / train_batches, lr=scheduler.get_last_lr())\n            \n        bar.set_postfix(accuracy = train_correct / train_size, loss = train_loss / train_batches)\n            \n        with torch.no_grad():\n\n            for batch, (X, y) in enumerate(val_dataloader):\n                pred = model(X)\n                val_loss += loss_fn(pred, y).item()\n                val_correct += (pred.argmax(1) == y).type(torch.float).sum().item()\n                val_size += len(y)\n                val_batches += 1\n\n                if batch % 10 == 0 or batch+1 == len(val_dataloader):\n                    bar.set_postfix(\n                        accuracy = train_correct / train_size, loss = train_loss / train_batches,\n                        val_accuracy = val_correct / val_size, val_loss = val_loss / val_batches\n                    )\n                    \n    if scheduler.T_0 - scheduler.T_cur < 0.1:\n        print()\n            \n    return train_correct / train_size, train_loss / train_batches, val_correct / val_size, val_loss / val_batches","metadata":{"execution":{"iopub.status.busy":"2023-04-16T03:20:29.987736Z","iopub.execute_input":"2023-04-16T03:20:29.988097Z","iopub.status.idle":"2023-04-16T03:20:30.002230Z","shell.execute_reply.started":"2023-04-16T03:20:29.988063Z","shell.execute_reply":"2023-04-16T03:20:30.000990Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"For each epoch in the training loop, if the val_loss is the best so far, the model's parameter are saved.","metadata":{}},{"cell_type":"code","source":"best_loss = float('inf')\nsaved_state = model.state_dict()\nsaved_epoch = 0\n\nepochs = 32\n\nfor epoch in range(epochs):\n    acc, loss, v_acc, v_loss = train_val_loop(epoch, train_dataloader, test_dataloader, model, loss_fn, optimizer, scheduler)\n    if v_loss<best_loss:\n        best_loss = v_loss\n        saved_state = model.state_dict()\n        saved_epoch = epoch + 1","metadata":{"execution":{"iopub.status.busy":"2023-04-16T03:20:31.201583Z","iopub.execute_input":"2023-04-16T03:20:31.202722Z","iopub.status.idle":"2023-04-16T03:21:30.772590Z","shell.execute_reply.started":"2023-04-16T03:20:31.202673Z","shell.execute_reply":"2023-04-16T03:21:30.771448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"saved_epoch","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(saved_state, 'model_weights.pth')","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:17:33.636375Z","iopub.execute_input":"2023-04-16T02:17:33.636789Z","iopub.status.idle":"2023-04-16T02:17:33.690499Z","shell.execute_reply.started":"2023-04-16T02:17:33.636749Z","shell.execute_reply":"2023-04-16T02:17:33.689315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torchinfo.summary(model, (1,534*2))","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:17:33.692046Z","iopub.execute_input":"2023-04-16T02:17:33.692857Z","iopub.status.idle":"2023-04-16T02:17:33.718842Z","shell.execute_reply.started":"2023-04-16T02:17:33.692805Z","shell.execute_reply":"2023-04-16T02:17:33.717694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### PyTorch → ONNX","metadata":{}},{"cell_type":"markdown","source":"Preprocessing module used to handle the raw input data.","metadata":{}},{"cell_type":"code","source":"class Preprocess(nn.Module):\n    def __init__(self, agg):\n        super().__init__()\n        self.agg = agg\n        \n    def forward(self, x):\n        x = x[:,INDICES]\n        x = self.agg(x)\n        return x","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preprocess = Preprocess(prepare)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eval_model = MLPModel().to(device)\neval_model.load_state_dict(saved_state)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:17:55.608208Z","iopub.execute_input":"2023-04-16T02:17:55.608754Z","iopub.status.idle":"2023-04-16T02:17:55.736771Z","shell.execute_reply.started":"2023-04-16T02:17:55.608704Z","shell.execute_reply":"2023-04-16T02:17:55.735703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preprocess.eval()\npreprocess_sample = torch.rand((23, 543, 3)).to(device) # 23 is an arbitrary number of frames, 543 is the number of rows/landmarks, 3 is the x, y, z columns\nonnx_preprocess_path = 'preprocess.onnx'\ntorch.onnx.export(preprocess,\n                  preprocess_sample,\n                  onnx_preprocess_path,\n                  opset_version=12,\n                  input_names = ['inputs'],\n                  output_names = ['outputs'],\n                  dynamic_axes={'inputs': {0: 'frames'}})\n\neval_model.eval()\nmodel_sample = torch.rand((1, 534*2)).to(device) # 1 is the batch size\nonnx_model_path = 'model.onnx'\ntorch.onnx.export(eval_model,\n                  model_sample,\n                  onnx_model_path,\n                  opset_version=12,\n                  input_names = ['inputs'],\n                  output_names = ['outputs'],\n                  dynamic_axes={'inputs': {0: 'batch_size'}})","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:17:56.476808Z","iopub.execute_input":"2023-04-16T02:17:56.477290Z","iopub.status.idle":"2023-04-16T02:17:56.840963Z","shell.execute_reply.started":"2023-04-16T02:17:56.477245Z","shell.execute_reply":"2023-04-16T02:17:56.839746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Will raise an exception if checks fail\nonnx_preprocess = onnx.load(onnx_preprocess_path)\nonnx.checker.check_model(onnx_preprocess)\nonnx_model = onnx.load(onnx_model_path)\nonnx.checker.check_model(onnx_model)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:17:57.747297Z","iopub.execute_input":"2023-04-16T02:17:57.747747Z","iopub.status.idle":"2023-04-16T02:17:57.824260Z","shell.execute_reply.started":"2023-04-16T02:17:57.747702Z","shell.execute_reply":"2023-04-16T02:17:57.823119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### ONNX → Tensorflow","metadata":{}},{"cell_type":"code","source":"tf_preprocess_path = 'tf_preprocess'\ntf_preprocess = onnx_tf.backend.prepare(onnx_preprocess)\ntf_preprocess.export_graph(tf_preprocess_path)\n\ntf_model_path = 'tf_model'\ntf_model = onnx_tf.backend.prepare(onnx_model)\ntf_model.export_graph(tf_model_path)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:18:01.847135Z","iopub.execute_input":"2023-04-16T02:18:01.847494Z","iopub.status.idle":"2023-04-16T02:18:09.644852Z","shell.execute_reply.started":"2023-04-16T02:18:01.847460Z","shell.execute_reply":"2023-04-16T02:18:09.643518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class InferenceModel(tf.Module):\n    def __init__(self):\n        super().__init__()\n        self.preprocess = tf.saved_model.load(tf_preprocess_path)\n        self.model = tf.saved_model.load(tf_model_path)\n        self.preprocess.trainable = False\n        self.model.trainable = False\n    \n    @tf.function(input_signature=[\n      tf.TensorSpec(shape=[None, 543, 3], dtype=tf.float32, name='inputs')\n    ])\n    def call(self, x):\n        outputs = {}\n        preprocessed = self.preprocess(**{'inputs':x})['outputs']\n        pred = self.model(**{'inputs':tf.expand_dims(preprocessed, 0)})['outputs'][0,:]\n        #pred = tf.nn.softmax(pred)\n        return {\n            'outputs': pred\n        }","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:18:09.646928Z","iopub.execute_input":"2023-04-16T02:18:09.647349Z","iopub.status.idle":"2023-04-16T02:18:09.656482Z","shell.execute_reply.started":"2023-04-16T02:18:09.647305Z","shell.execute_reply":"2023-04-16T02:18:09.655060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tf_inference = InferenceModel()\ntf_inference_path = 'tf_inference'\ntf.saved_model.save(tf_inference, tf_inference_path, signatures={'serving_default': tf_inference.call})","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:18:09.658713Z","iopub.execute_input":"2023-04-16T02:18:09.659203Z","iopub.status.idle":"2023-04-16T02:18:10.792409Z","shell.execute_reply.started":"2023-04-16T02:18:09.659159Z","shell.execute_reply":"2023-04-16T02:18:10.791153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Tensorflow → TFLite","metadata":{}},{"cell_type":"code","source":"model_converter = tf.lite.TFLiteConverter.from_saved_model(tf_inference_path) # path to the SavedModel directory\ntflite_model = model_converter.convert()\n\n# Save the model.\nwith open('model.tflite', 'wb') as f:\n    f.write(tflite_model)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:18:10.795764Z","iopub.execute_input":"2023-04-16T02:18:10.797950Z","iopub.status.idle":"2023-04-16T02:18:12.734606Z","shell.execute_reply.started":"2023-04-16T02:18:10.797902Z","shell.execute_reply":"2023-04-16T02:18:12.733460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The submission file (submission.zip) is created by compressing the TFLite model.","metadata":{}},{"cell_type":"code","source":"!zip submission.zip model.tflite","metadata":{"execution":{"iopub.status.busy":"2023-04-16T02:18:12.736909Z","iopub.execute_input":"2023-04-16T02:18:12.737291Z","iopub.status.idle":"2023-04-16T02:18:15.520412Z","shell.execute_reply.started":"2023-04-16T02:18:12.737253Z","shell.execute_reply":"2023-04-16T02:18:15.519244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluating the TFLite model","metadata":{}},{"cell_type":"code","source":"import tflite_runtime.interpreter as tflite\ninterpreter = tflite.Interpreter('model.tflite')\n\nfound_signatures = list(interpreter.get_signature_list().keys())\n\n# if REQUIRED_SIGNATURE not in found_signatures:\n#     raise KernelEvalException('Required input signature not found.')\n\nframes = load_relevant_data_subset('/kaggle/input/asl-signs/train_landmark_files/16069/100015657.parquet')\n\nprediction_fn = interpreter.get_signature_runner(\"serving_default\")\noutput = prediction_fn(inputs=frames)\nsign = np.argmax(output[\"outputs\"])\n\nprint(sign, output[\"outputs\"].shape)","metadata":{"execution":{"iopub.status.busy":"2023-03-30T07:44:02.025954Z","iopub.status.idle":"2023-03-30T07:44:02.026694Z","shell.execute_reply.started":"2023-03-30T07:44:02.026464Z","shell.execute_reply":"2023-03-30T07:44:02.026491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntests = []\nfor index, row in signs_df.iloc[:50].iterrows():\n    frames = load_relevant_data_subset(ASL_DIR / row.path)\n    interpreter = tflite.Interpreter('model.tflite')\n    prediction_fn = interpreter.get_signature_runner(\"serving_default\")\n    output = prediction_fn(inputs=frames)\n    #output = tf_inference.call(frames)\n    #output = model(preprocess(torch.Tensor(frames)).unsqueeze(0))\n    sign = np.argmax(output[\"outputs\"])\n    #sign = torch.argmax(output)\n    \n    tests.append((sign, row.label))","metadata":{"execution":{"iopub.status.busy":"2023-03-30T07:44:02.034554Z","iopub.status.idle":"2023-03-30T07:44:02.035336Z","shell.execute_reply.started":"2023-03-30T07:44:02.035077Z","shell.execute_reply":"2023-03-30T07:44:02.035109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tests","metadata":{"execution":{"iopub.status.busy":"2023-03-30T07:44:02.036674Z","iopub.status.idle":"2023-03-30T07:44:02.037410Z","shell.execute_reply.started":"2023-03-30T07:44:02.037146Z","shell.execute_reply":"2023-03-30T07:44:02.037205Z"},"trusted":true},"execution_count":null,"outputs":[]}]}