{"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 onnxsim\n!pip install tflite-runtime","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:05:29.279244Z","iopub.execute_input":"2023-04-16T09:05:29.280296Z","iopub.status.idle":"2023-04-16T09:06:14.248169Z","shell.execute_reply.started":"2023-04-16T09:05:29.280232Z","shell.execute_reply":"2023-04-16T09:06:14.246935Z"},"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\nimport math\nimport time\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchinfo\n\nfrom tqdm import tqdm # Progress bar","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-16T09:06:14.250769Z","iopub.execute_input":"2023-04-16T09:06:14.251147Z","iopub.status.idle":"2023-04-16T09:06:15.251866Z","shell.execute_reply.started":"2023-04-16T09:06:14.251102Z","shell.execute_reply":"2023-04-16T09:06:15.250755Z"},"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 onnxsim\nimport onnx_tf\nimport tensorflow as tf\nimport tflite_runtime.interpreter as tflite","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:06:15.253544Z","iopub.execute_input":"2023-04-16T09:06:15.254183Z","iopub.status.idle":"2023-04-16T09:06:24.105533Z","shell.execute_reply.started":"2023-04-16T09:06:15.254144Z","shell.execute_reply":"2023-04-16T09:06:24.104451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Setup path constants for Kaggle Notebook.","metadata":{}},{"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-16T09:06:24.108603Z","iopub.execute_input":"2023-04-16T09:06:24.109398Z","iopub.status.idle":"2023-04-16T09:06:24.118030Z","shell.execute_reply.started":"2023-04-16T09:06:24.109356Z","shell.execute_reply":"2023-04-16T09:06:24.117069Z"},"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-16T09:06:24.119533Z","iopub.execute_input":"2023-04-16T09:06:24.119892Z","iopub.status.idle":"2023-04-16T09:06:24.195603Z","shell.execute_reply.started":"2023-04-16T09:06:24.119855Z","shell.execute_reply":"2023-04-16T09:06:24.194562Z"},"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-16T09:06:24.198775Z","iopub.execute_input":"2023-04-16T09:06:24.199623Z","iopub.status.idle":"2023-04-16T09:06:24.206544Z","shell.execute_reply.started":"2023-04-16T09:06:24.199582Z","shell.execute_reply":"2023-04-16T09:06:24.205445Z"},"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 (DATASET_DIR / 'landmarks.json').open() as f:\n    landmarks = json.load(f)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:06:24.208128Z","iopub.execute_input":"2023-04-16T09:06:24.208549Z","iopub.status.idle":"2023-04-16T09:06:24.227812Z","shell.execute_reply.started":"2023-04-16T09:06:24.208513Z","shell.execute_reply":"2023-04-16T09:06:24.226888Z"},"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-16T09:06:24.230337Z","iopub.execute_input":"2023-04-16T09:06:24.230933Z","iopub.status.idle":"2023-04-16T09:06:24.483679Z","shell.execute_reply.started":"2023-04-16T09:06:24.230895Z","shell.execute_reply":"2023-04-16T09:06:24.482590Z"},"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-16T09:06:24.485405Z","iopub.execute_input":"2023-04-16T09:06:24.485780Z","iopub.status.idle":"2023-04-16T09:06:24.529276Z","shell.execute_reply.started":"2023-04-16T09:06:24.485742Z","shell.execute_reply":"2023-04-16T09:06:24.528321Z"},"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, prepare):\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.prepare = prepare\n    \n    def __len__(self):\n        return len(self.items)\n        \n    def __getitem__(self, index):\n        return self.prepare(self.items[index]).float(), self.labels[index]","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:06:24.533519Z","iopub.execute_input":"2023-04-16T09:06:24.533810Z","iopub.status.idle":"2023-04-16T09:06:24.540976Z","shell.execute_reply.started":"2023-04-16T09:06:24.533782Z","shell.execute_reply":"2023-04-16T09:06:24.539820Z"},"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()])\nINDICES = np.load(DATASET_DIR / 'indices.npy')","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:06:24.542791Z","iopub.execute_input":"2023-04-16T09:06:24.543153Z","iopub.status.idle":"2023-04-16T09:06:24.596442Z","shell.execute_reply.started":"2023-04-16T09:06:24.543112Z","shell.execute_reply":"2023-04-16T09:06:24.595504Z"},"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\n    # Coords + Angles\n    tensor = torch.cat((\n        normed.flatten(-2),\n        angles),1)\n\n    tensor[tensor.isnan()] = 0\n\n    return tensor","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:06:24.598216Z","iopub.execute_input":"2023-04-16T09:06:24.598954Z","iopub.status.idle":"2023-04-16T09:06:24.608079Z","shell.execute_reply.started":"2023-04-16T09:06:24.598917Z","shell.execute_reply":"2023-04-16T09:06:24.606994Z"},"trusted":true},"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-16T09:06:24.609496Z","iopub.execute_input":"2023-04-16T09:06:24.609999Z","iopub.status.idle":"2023-04-16T09:11:46.832125Z","shell.execute_reply.started":"2023-04-16T09:06:24.609952Z","shell.execute_reply":"2023-04-16T09:11:46.830882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Creates a single tensor from a list of inputs of varying length using padding and also returning the mask.","metadata":{}},{"cell_type":"code","source":"def pad_batch(batch):\n    max_frames = max([len(entry) for entry in batch])\n    size = (max_frames, len(batch), len(batch[0][0]))\n    padded = torch.zeros(size).to(device)\n    mask = torch.full((len(batch), max_frames), True).to(device)\n    for index, entry in enumerate(batch):\n        frames = len(entry)\n        padded[:frames, index] = entry\n        mask[index, :frames] = False\n        \n    return padded, mask","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:11:46.833887Z","iopub.execute_input":"2023-04-16T09:11:46.834334Z","iopub.status.idle":"2023-04-16T09:11:46.842377Z","shell.execute_reply.started":"2023-04-16T09:11:46.834291Z","shell.execute_reply":"2023-04-16T09:11:46.840849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Pad the batches in the dataloader.","metadata":{}},{"cell_type":"code","source":"def collate(batch):\n    transposed = list(zip(*batch))\n    sequence = list(transposed[0])\n    X = [torch.Tensor(x).nan_to_num(nan=0).flatten(1).float().to(device) for x in sequence]\n    src, mask = pad_batch(X)\n    y = torch.Tensor(transposed[1]).long().to(device)\n    return src, mask, y","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:11:46.844190Z","iopub.execute_input":"2023-04-16T09:11:46.844592Z","iopub.status.idle":"2023-04-16T09:11:46.855230Z","shell.execute_reply.started":"2023-04-16T09:11:46.844555Z","shell.execute_reply":"2023-04-16T09:11:46.854234Z"},"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, collate_fn=collate)\ntest_dataloader = DataLoader(test_dataset, batch_size=128, shuffle=True, collate_fn=collate)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:11:46.856608Z","iopub.execute_input":"2023-04-16T09:11:46.857017Z","iopub.status.idle":"2023-04-16T09:11:46.867318Z","shell.execute_reply.started":"2023-04-16T09:11:46.856980Z","shell.execute_reply":"2023-04-16T09:11:46.866318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"markdown","source":"Positional encoding layer of the transformer","metadata":{}},{"cell_type":"code","source":"class PositionalEncoding(nn.Module):\n\n    def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 5000):\n        super().__init__()\n        self.dropout = nn.Dropout(p=dropout)\n\n        position = torch.arange(max_len).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))\n        pe = torch.zeros(max_len, 1, d_model)\n        pe[:, 0, 0::2] = torch.sin(position * div_term)\n        pe[:, 0, 1::2] = torch.cos(position * div_term)\n        self.register_buffer('pe', pe)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x = x + self.pe[:x.size(0)]\n        return self.dropout(x).to(device)\n        #return self.pe","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:11:46.868752Z","iopub.execute_input":"2023-04-16T09:11:46.869220Z","iopub.status.idle":"2023-04-16T09:11:46.879442Z","shell.execute_reply.started":"2023-04-16T09:11:46.869178Z","shell.execute_reply":"2023-04-16T09:11:46.878315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The transformer model used in training uses the PyTorch `TransformerEncoderLayer` and the positonal encoding implemented above.","metadata":{}},{"cell_type":"code","source":"class TransformerModel(nn.Module):\n\n    def __init__(self):\n        super().__init__()\n        self.model_type = 'Transformer'\n        self.in_tokens = 534\n        self.d_model = 800\n        self.d_embed_ff = 1200\n        self.out_tokens = 250\n        self.nhead = 8\n        self.d_ff = 800\n        self.nlayers = 1\n        self.dropout = 0.4\n        \n        self.embed = nn.Sequential(\n            nn.Linear(self.in_tokens, self.d_embed_ff),\n            nn.LayerNorm(self.d_embed_ff),\n            nn.ReLU(),\n            nn.Linear(self.d_embed_ff, self.d_model)\n        )\n        self.pos_encoder = PositionalEncoding(self.d_model, self.dropout)\n        \n        encoder_layers = nn.TransformerEncoderLayer(self.d_model, self.nhead, self.d_ff, self.dropout)\n        self.transformer_encoder = nn.TransformerEncoder(encoder_layers, self.nlayers)\n        \n        self.decoder = nn.Linear(self.d_model, self.out_tokens)\n\n    def forward(self, src, mask) -> torch.Tensor:\n        src = self.embed(src)\n        src = self.pos_encoder(src)\n        src = torch.cat([torch.zeros((1,src.size(1),self.d_model)).to(src), src],0)\n        mask = torch.cat([torch.full((src.size(1), 1), False).to(mask), mask],1)\n        \n        output = self.transformer_encoder(src, src_key_padding_mask=mask)\n        output = output[0]\n        output = self.decoder(output)\n        \n        return output","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:11:46.881026Z","iopub.execute_input":"2023-04-16T09:11:46.881411Z","iopub.status.idle":"2023-04-16T09:11:46.894060Z","shell.execute_reply.started":"2023-04-16T09:11:46.881356Z","shell.execute_reply":"2023-04-16T09:11:46.893298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = TransformerModel().to(device)","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:11:46.895468Z","iopub.execute_input":"2023-04-16T09:11:46.896182Z","iopub.status.idle":"2023-04-16T09:11:47.045019Z","shell.execute_reply.started":"2023-04-16T09:11:46.896138Z","shell.execute_reply":"2023-04-16T09:11:47.043962Z"},"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":"learning_rate = 5e-4\nweight_decay = 1e-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-16T09:11:47.046660Z","iopub.execute_input":"2023-04-16T09:11:47.047062Z","iopub.status.idle":"2023-04-16T09:11:47.054639Z","shell.execute_reply.started":"2023-04-16T09:11:47.047022Z","shell.execute_reply":"2023-04-16T09:11:47.053512Z"},"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        \n        # Training\n        for batch, (src, mask, y) in enumerate(train_dataloader):\n            \n            # Compute prediction and loss\n            pred = model(src, mask)\n            loss = loss_fn(pred, y)\n            \n            # Compute metrics\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            torch.nn.utils.clip_grad_norm_(model.parameters(), 0.5)\n            optimizer.step()\n                \n            scheduler.step(epoch + batch / total_batches)\n            \n            # Update progress bar\n            bar.update()\n            bar.set_postfix(accuracy = train_correct / train_size, loss = train_loss / train_batches,\n                           lr=scheduler.get_last_lr())\n            \n        bar.set_postfix(accuracy = train_correct / train_size, loss = train_loss / train_batches)\n        #scheduler.step()\n           \n        # Validation\n        with torch.no_grad():\n\n            for batch, (src, mask, y) in enumerate(val_dataloader):\n                \n                # Compute prediction and loss\n                pred = model(src, mask)\n                val_loss += loss_fn(pred, y).item()\n                \n                # Compute metrics\n                val_correct += (pred.argmax(1) == y).type(torch.float).sum().item()\n                val_size += len(y)\n                val_batches += 1\n\n                # Update progress bar\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-16T09:11:47.056564Z","iopub.execute_input":"2023-04-16T09:11:47.057538Z","iopub.status.idle":"2023-04-16T09:11:47.073356Z","shell.execute_reply.started":"2023-04-16T09:11:47.057489Z","shell.execute_reply":"2023-04-16T09:11:47.072290Z"},"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()\n\nepochs = 64\n\n# Iterate through epochs for training\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()","metadata":{"execution":{"iopub.status.busy":"2023-04-16T09:11:47.075027Z","iopub.execute_input":"2023-04-16T09:11:47.075426Z","iopub.status.idle":"2023-04-16T09:15:15.534310Z","shell.execute_reply.started":"2023-04-16T09:11:47.075386Z","shell.execute_reply":"2023-04-16T09:15:15.532583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The model with the best val_loss is saved and used for conversion to TFLite.","metadata":{}},{"cell_type":"code","source":"torch.save(saved_state, 'model_weights.pth')","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:19:28.116687Z","iopub.execute_input":"2023-04-15T16:19:28.117041Z","iopub.status.idle":"2023-04-15T16:19:28.190624Z","shell.execute_reply.started":"2023-04-15T16:19:28.117008Z","shell.execute_reply":"2023-04-15T16:19:28.189601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preparing Model for Conversion to TFLite","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":"markdown","source":"In order to convert to TFLite, I had to create my own implementation of the transformer. Most of the code is adapted from the actual PyTorch source code but stripped down to only the code I need.","metadata":{}},{"cell_type":"code","source":"# https://pytorch.org/docs/stable/_modules/torch/nn/modules/activation.html#MultiheadAttention\n# https://github.com/pytorch/pytorch/blob/5ca3afd1bfeb16c7b873009d2b8044fa745d31b4/torch/nn/functional.py#L5044\n# https://github.com/pytorch/pytorch/blob/5ca3afd1bfeb16c7b873009d2b8044fa745d31b4/torch/nn/functional.py#L4751\nclass InferenceMultiheadAttention(nn.MultiheadAttention):\n    def forward(self,key,query,value,key_padding_mask = None):\n        # set up shape vars\n        tgt_len, bsz, _ = query.shape\n        src_len, _, _ = key.shape\n        \n        w_q, w_k, w_v = self.in_proj_weight.chunk(3)\n        if self.in_proj_bias is None:\n            b_q = b_k = b_v = None\n        else:\n            b_q, b_k, b_v = self.in_proj_bias.chunk(3)\n        q = F.linear(query, w_q, b_q)\n        k = F.linear(key, w_k, b_k)\n        v = F.linear(value, w_v, b_v)\n\n        # reshape q, k, v for multihead attention and make em batch first\n        q = q.view(tgt_len, bsz * self.num_heads, self.head_dim).transpose(0, 1)\n        k = k.view(k.shape[0], bsz * self.num_heads, self.head_dim).transpose(0, 1)\n        v = v.view(v.shape[0], bsz * self.num_heads, self.head_dim).transpose(0, 1)\n\n        # update source sequence length after adjustments\n        src_len = k.size(1)\n\n        # merge key padding and attention masks\n        attn_mask = None\n        if key_padding_mask is not None:\n            key_padding_mask = key_padding_mask.view(bsz, 1, 1, src_len).   \\\n                expand(-1, num_heads, -1, -1).reshape(bsz * self.num_heads, 1, src_len)\n            attn_mask = key_padding_mask\n            \n        # calculate attention and out projection\n        if attn_mask is not None:\n            if attn_mask.size(0) == 1 and attn_mask.dim() == 3:\n                attn_mask = attn_mask.unsqueeze(0)\n            else:\n                attn_mask = attn_mask.view(bsz, num_heads, -1, src_len)\n\n        q = q.view(bsz, self.num_heads, tgt_len, self.head_dim)\n        k = k.view(bsz, self.num_heads, src_len, self.head_dim)\n        v = v.view(bsz, self.num_heads, src_len, self.head_dim)\n\n        dropout = nn.Dropout(self.dropout)\n        dropout.train(self.training)\n        attn_output = attention(q, k, v, self.head_dim, mask=attn_mask, dropout=dropout)\n        attn_output = attn_output.permute(2, 0, 1, 3).contiguous().view(bsz * tgt_len, self.embed_dim)\n\n        attn_output = F.linear(attn_output, self.out_proj.weight, self.out_proj.bias)\n        attn_output = attn_output.view(tgt_len, bsz, attn_output.size(1))\n        \n        return attn_output, None","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:19:31.495593Z","iopub.execute_input":"2023-04-15T16:19:31.496505Z","iopub.status.idle":"2023-04-15T16:19:31.510984Z","shell.execute_reply.started":"2023-04-15T16:19:31.496459Z","shell.execute_reply":"2023-04-15T16:19:31.509793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The `attention` function is the only thing not adapted from PyTorch source code, but rather from a tutorial I found.","metadata":{}},{"cell_type":"code","source":"# https://towardsdatascience.com/how-to-code-the-transformer-in-pytorch-24db27c8f9ec#3fa3\ndef attention(q, k, v, d_k, mask=None, dropout=None):\n    \n    scores = torch.matmul(q, k.transpose(-2, -1)) /  math.sqrt(d_k)\n    if mask is not None:\n        mask = mask.unsqueeze(1)\n        scores = scores.masked_fill(mask == 0, -1e9)\n    scores = F.softmax(scores, dim=-1)\n    \n    if dropout is not None:\n        scores = dropout(scores)\n        \n    output = torch.matmul(scores, v)\n    return output","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:19:32.020213Z","iopub.execute_input":"2023-04-15T16:19:32.020585Z","iopub.status.idle":"2023-04-15T16:19:32.028920Z","shell.execute_reply.started":"2023-04-15T16:19:32.020551Z","shell.execute_reply":"2023-04-15T16:19:32.027717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://pytorch.org/docs/stable/_modules/torch/nn/modules/transformer.html#TransformerEncoderLayer\nclass InferenceTransformerEncoderLayer(nn.Module):\n    def __init__(self, d_model: int, nhead: int, d_ff: int = 2048, dropout: float = 0.1,\n                 layer_norm_eps: float = 1e-5,\n                 device=None, dtype=None) -> None:\n        factory_kwargs = {'device': device, 'dtype': dtype}\n        super().__init__()\n        # Multihead Attention\n        self.self_attn = InferenceMultiheadAttention(d_model, nhead, dropout=dropout)\n        \n        # Implementation of Feedforward model\n        self.linear1 = nn.Linear(d_model, d_ff, **factory_kwargs)\n        self.dropout = nn.Dropout(dropout)\n        self.linear2 = nn.Linear(d_ff, d_model, **factory_kwargs)\n\n        self.norm1 = nn.LayerNorm(d_model, eps=layer_norm_eps, **factory_kwargs)\n        self.norm2 = nn.LayerNorm(d_model, eps=layer_norm_eps, **factory_kwargs)\n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n        \n        self.activation = nn.ReLU()\n        \n    def forward(self, src: torch.Tensor, src_mask = None, src_key_padding_mask: torch.Tensor = None) -> torch.Tensor:\n        x = src\n        x = self.norm1(x + self._sa_block(x, src_key_padding_mask))\n        x = self.norm2(x + self._ff_block(x))\n        return x\n    \n    # self-attention block\n    def _sa_block(self, x, key_padding_mask):\n        x = self.self_attn(x, x, x,key_padding_mask=key_padding_mask)[0]\n        return self.dropout1(x)\n\n    # feed forward block\n    def _ff_block(self, x):\n        x = self.linear2(self.dropout(self.activation(self.linear1(x))))\n        return self.dropout2(x)","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:19:32.474336Z","iopub.execute_input":"2023-04-15T16:19:32.474715Z","iopub.status.idle":"2023-04-15T16:19:32.486617Z","shell.execute_reply.started":"2023-04-15T16:19:32.474681Z","shell.execute_reply":"2023-04-15T16:19:32.485563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The EvalModel is mostly the same as the model used in training, except it used the custom transformer implemented above and is only made to accept one input instead of a batch.","metadata":{}},{"cell_type":"code","source":"class EvalModel(TransformerModel):\n    \n    def __init__(self):\n        super().__init__()\n        encoder_layers = InferenceTransformerEncoderLayer(self.d_model, self.nhead, self.d_ff, self.dropout)\n        self.transformer_encoder = nn.TransformerEncoder(encoder_layers, self.nlayers)\n    \n    # src should be a single sequence of frames of size [frames, 534]\n    def forward(self, src: torch.Tensor):\n        src = src.unsqueeze(1)\n        src = self.embed(src)\n        src = self.pos_encoder(src)\n        src = torch.cat([torch.zeros((1,1,self.d_model)).to(device), src],0)\n        \n        output = self.transformer_encoder(src) # [frames+1,1,d_model]\n        output = output[0] # [1,d_model]\n        output = self.decoder(output) # [1,250]\n        output = output[0] # [250]\n        \n        return output","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:19:33.065753Z","iopub.execute_input":"2023-04-15T16:19:33.066452Z","iopub.status.idle":"2023-04-15T16:19:33.074753Z","shell.execute_reply.started":"2023-04-15T16:19:33.066416Z","shell.execute_reply":"2023-04-15T16:19:33.073476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eval_model = EvalModel()","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:19:33.801603Z","iopub.execute_input":"2023-04-15T16:19:33.801968Z","iopub.status.idle":"2023-04-15T16:19:33.940048Z","shell.execute_reply.started":"2023-04-15T16:19:33.801938Z","shell.execute_reply":"2023-04-15T16:19:33.938987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eval_model.load_state_dict(saved_state)\neval_model.eval()","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:19:34.656049Z","iopub.execute_input":"2023-04-15T16:19:34.656418Z","iopub.status.idle":"2023-04-15T16:19:34.675352Z","shell.execute_reply.started":"2023-04-15T16:19:34.656386Z","shell.execute_reply":"2023-04-15T16:19:34.674321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torchinfo.summary(eval_model, input_size=(105,534))","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:19:35.619077Z","iopub.execute_input":"2023-04-15T16:19:35.620048Z","iopub.status.idle":"2023-04-15T16:19:35.660815Z","shell.execute_reply.started":"2023-04-15T16:19:35.619994Z","shell.execute_reply":"2023-04-15T16:19:35.659703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### PyTorch → ONNX","metadata":{}},{"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={\n                      'inputs': {0: 'frames'},\n                      'outputs': {0: 'frames'}\n                  })\n\neval_model.eval()\nmodel_sample = torch.rand((23, 534)).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: 'frames'}})","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:19:45.351347Z","iopub.execute_input":"2023-04-15T16:19:45.351729Z","iopub.status.idle":"2023-04-15T16:19:45.941937Z","shell.execute_reply.started":"2023-04-15T16:19:45.351696Z","shell.execute_reply":"2023-04-15T16:19:45.940766Z"},"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-15T16:19:47.751082Z","iopub.execute_input":"2023-04-15T16:19:47.751704Z","iopub.status.idle":"2023-04-15T16:19:47.877447Z","shell.execute_reply.started":"2023-04-15T16:19:47.751661Z","shell.execute_reply":"2023-04-15T16:19:47.876387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"onnx_preprocess, _ = onnxsim.simplify(onnx_preprocess)\nonnx.checker.check_model(onnx_preprocess)\nonnx_model, _ = onnxsim.simplify(onnx_model)\nonnx.checker.check_model(onnx_model)","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:19:49.661657Z","iopub.execute_input":"2023-04-15T16:19:49.662672Z","iopub.status.idle":"2023-04-15T16:19:52.049124Z","shell.execute_reply.started":"2023-04-15T16:19:49.662603Z","shell.execute_reply":"2023-04-15T16:19:52.048025Z"},"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-15T16:19:55.139522Z","iopub.execute_input":"2023-04-15T16:19:55.140225Z","iopub.status.idle":"2023-04-15T16:20:03.554199Z","shell.execute_reply.started":"2023-04-15T16:19:55.140189Z","shell.execute_reply":"2023-04-15T16:20:03.553123Z"},"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=preprocessed)['outputs']\n        #pred = tf.nn.softmax(pred)\n        return {\n            'outputs': pred\n        }","metadata":{"execution":{"iopub.status.busy":"2023-04-15T16:20:12.651679Z","iopub.execute_input":"2023-04-15T16:20:12.652275Z","iopub.status.idle":"2023-04-15T16:20:12.660430Z","shell.execute_reply.started":"2023-04-15T16:20:12.652239Z","shell.execute_reply":"2023-04-15T16:20:12.659368Z"},"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-15T16:20:14.875257Z","iopub.execute_input":"2023-04-15T16:20:14.875627Z","iopub.status.idle":"2023-04-15T16:20:17.888816Z","shell.execute_reply.started":"2023-04-15T16:20:14.875594Z","shell.execute_reply":"2023-04-15T16:20:17.887760Z"},"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\nmodel_converter.optimizations = [tf.lite.Optimize.DEFAULT]\nmodel_converter.target_spec.supported_types = [tf.float16]\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-15T16:20:22.542480Z","iopub.execute_input":"2023-04-15T16:20:22.543194Z","iopub.status.idle":"2023-04-15T16:20:26.092610Z","shell.execute_reply.started":"2023-04-15T16:20:22.543155Z","shell.execute_reply":"2023-04-15T16:20:26.091535Z"},"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-15T16:20:27.686987Z","iopub.execute_input":"2023-04-15T16:20:27.687359Z","iopub.status.idle":"2023-04-15T16:20:30.089248Z","shell.execute_reply.started":"2023-04-15T16:20:27.687327Z","shell.execute_reply":"2023-04-15T16:20:30.087972Z"},"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-04-15T16:20:31.235209Z","iopub.execute_input":"2023-04-15T16:20:31.235609Z","iopub.status.idle":"2023-04-15T16:20:31.444910Z","shell.execute_reply.started":"2023-04-15T16:20:31.235572Z","shell.execute_reply":"2023-04-15T16:20:31.443720Z"},"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tests","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}