{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":52950,"databundleVersionId":5973250,"sourceType":"competition"}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Setup","metadata":{"_uuid":"11d1c4e4-668b-42f2-85cf-ea9be0b4e2e5","_cell_guid":"1aa5e70c-b5c2-4a43-a463-4bcdc30f8744","trusted":true}},{"cell_type":"code","source":"import torch \nfrom torch.nn import functional as F\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom torch import Tensor\n\nimport numpy as np\nimport pandas as pd\nimport pyarrow.parquet as pq\nimport json\nfrom tqdm.notebook import tqdm\nimport math\nfrom sklearn.metrics import accuracy_score\nimport os\nimport gc\nfrom typing import Optional,Tuple, Union\nimport matplotlib.pyplot as plt \n\nimport transformers\nfrom transformers.models.speech_to_text import Speech2TextConfig\nfrom transformers.models.speech_to_text.modeling_speech_to_text import shift_tokens_right, Speech2TextDecoder\nfrom timm.layers.norm_act import BatchNormAct2d\nimport torchaudio","metadata":{"_uuid":"14c8ef83-6676-49c6-bddb-173d9cbab870","_cell_guid":"4702c22e-3b54-48bb-896d-6dd69b6405ad","execution":{"iopub.status.busy":"2023-12-03T20:47:03.231375Z","iopub.execute_input":"2023-12-03T20:47:03.231767Z","iopub.status.idle":"2023-12-03T20:47:03.239662Z","shell.execute_reply.started":"2023-12-03T20:47:03.231741Z","shell.execute_reply":"2023-12-03T20:47:03.238436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"_uuid":"814e4ff3-d19d-422b-b384-fcae0ebfcbc3","_cell_guid":"9575265e-da30-4885-8e05-ac853cb3a3ec","trusted":true}},{"cell_type":"code","source":"# constants\nFEATURES = {\n    'hand': {\n        'left': range(21),\n        'right': range(21)\n    },\n    'pose': {\n        'left': [13, 15, 17, 19, 21],\n        'right': [14, 16, 18, 20, 22]\n    }\n}\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nVAL_SPLIT = 0.8\nTEST_SPLIT = 0.2\nBATCH_SIZE = 64\n\nFRAME_LEN = 128\nMAX_LEN = 400 # max number of frames\nMAX_PHRASE = 31 + 3 # max len from data + start, pad, end tokens\n\nSTART_TOKEN = 'S'\nPAD_TOKEN = 'P'\nEND_TOKEN = 'E'","metadata":{"_uuid":"3cd39c99-42c4-4cfb-ba08-0c8f7d450f91","_cell_guid":"0c072221-4506-458a-bbca-bdc3931f2bb5","execution":{"iopub.status.busy":"2023-12-03T20:47:03.24208Z","iopub.execute_input":"2023-12-03T20:47:03.242441Z","iopub.status.idle":"2023-12-03T20:47:03.250869Z","shell.execute_reply.started":"2023-12-03T20:47:03.242407Z","shell.execute_reply":"2023-12-03T20:47:03.249812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_columns(features: dict = FEATURES) -> list:\n    HAND_COLS = []\n    if 'hand' in features:\n        HAND_COLS = [\n            f'{d}_{o}_hand_{i}' \n            for o in ['right', 'left'] \n            for i in features['hand'][o] \n            for d in ['x', 'y', 'z']\n        ]\n        \n    POSE_COLS = []\n    if 'pose' in features:\n        POSE_COLS = [\n            f'{d}_pose_{i}' \n            for o in ['right', 'left'] \n            for i in features['pose'][o]\n            for d in ['x', 'y', 'z']\n        ]\n        \n    HEAD_COLS = []\n    if 'head' in features:\n        HEAD_COLS = [\n            f'{d}_head_{i}' \n            for i in features['head'] \n            for d in ['x', 'y', 'z']\n        ]\n    \n    return HAND_COLS + POSE_COLS + HEAD_COLS","metadata":{"_uuid":"f837e068-5230-4ff4-aaef-ef77e595448f","_cell_guid":"256f48d8-ff18-4634-8650-3e728be9f25c","execution":{"iopub.status.busy":"2023-12-03T20:47:03.25214Z","iopub.execute_input":"2023-12-03T20:47:03.252434Z","iopub.status.idle":"2023-12-03T20:47:03.260754Z","shell.execute_reply.started":"2023-12-03T20:47:03.252402Z","shell.execute_reply":"2023-12-03T20:47:03.259897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FingerspellingDataset(Dataset):\n    def __init__(self,\n                 features: dict = FEATURES,\n                 train: str = True,\n                 transform = None\n                ):\n        \n        if train:\n            dataset_path = '/kaggle/input/asl-fingerspelling/train_landmarks'\n            dataset_file = '/kaggle/input/asl-fingerspelling/train.csv'\n        else:\n            dataset_path = '/kaggle/input/asl-fingerspelling/supplemental_landmarks'\n            dataset_file = '/kaggle/input/asl-fingerspelling/supplemental_metadata.csv'\n        \n        self.dataset_path = dataset_path\n        self.dataset_df = pd.read_csv(dataset_file)\n        self.feature_columns = extract_columns(features)\n        self.transform = transform\n        \n        # fetch the data from the .parquet file\n        # filter out the non used columns\n        self.parquet_df = {\n            file.split('.')[0]: pq.read_table(\n                f\"{self.dataset_path}/{file.split('.')[0]}.parquet\",\n                columns=['sequence_id'] + self.feature_columns\n            ).to_pandas()\n            for file in os.listdir(dataset_path)\n        }\n        # convert parquet data to numpy\n        self.parquet_np = {\n            file_id: self.parquet_df[file_id].to_numpy() for file_id in self.parquet_df\n        }\n        \n        self.X_IDX = [i for i, col in enumerate(self.feature_columns)  if \"x_\" in col]\n        self.Y_IDX = [i for i, col in enumerate(self.feature_columns)  if \"y_\" in col]\n        self.Z_IDX = [i for i, col in enumerate(self.feature_columns)  if \"z_\" in col]\n    \n    def __len__(self):\n        return len(self.dataset_df)\n    \n    def __getitem__(self, index):\n        # convert to list of indices (if index is a tensor)\n        if torch.is_tensor(index):\n            index = index.tolist()\n    \n        # locate sample in dataset dataframe\n        sequence_id, file_id, phrase = self.dataset_df.iloc[index][['sequence_id', 'file_id', 'phrase']]\n        \n        # filter dataset and fetch entries for the relevant file_id\n        file_df = self.dataset_df.loc[self.dataset_df[\"file_id\"] == file_id]\n    \n        # filter the parquet data by the sequence_id of the sample\n        frames = self.parquet_np[str(file_id)][self.parquet_df[str(file_id)].index == sequence_id]\n        indices_lists = [self.X_IDX, self.Y_IDX, self.Z_IDX]\n        frames = np.stack([frames[:, indices] for indices in indices_lists], axis=-1)\n        frames = frames.reshape(frames.shape[0], -1, len(indices_lists))\n\n        sample = {\n            'data': frames, # numpy.ndarray\n            'phrase': phrase, # string\n        }\n        \n        # apply transformation(s)\n        if self.transform:\n            sample = self.transform(sample)\n    \n        return sample","metadata":{"_uuid":"1a2164e9-6842-4820-98d7-9e2d5df9d6c8","_cell_guid":"f31e0d10-d65e-4915-a940-38b9c3109e28","execution":{"iopub.status.busy":"2023-12-03T20:47:03.262354Z","iopub.execute_input":"2023-12-03T20:47:03.262735Z","iopub.status.idle":"2023-12-03T20:47:03.280254Z","shell.execute_reply.started":"2023-12-03T20:47:03.262702Z","shell.execute_reply":"2023-12-03T20:47:03.279191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_data_loaders(features: dict = FEATURES, val_split: float = 0.8, test_split: float = 0.2, batch_size: int = 16):\n    # load datasets\n    transform = transforms.Compose([\n                ToTensor(),\n                NormalizeAndFillNaNs(),\n                # Resample(),\n                # SpatialRandomAffine(),\n#                 SpatialMask(),\n#                 TemporalMask(),\n#                 TemporalCrop(),\n                # FlipLeftRight(),\n                InterpolateOrPad(), \n                SplitData(),\n                TokenizePhrase() \n    ])\n    testVal_transform = transforms.Compose([\n                ToTensor(),\n                NormalizeAndFillNaNs(),\n                InterpolateOrPad(), \n                SplitData(),\n                TokenizePhrase() \n    ])\n    train_data = FingerspellingDataset(features=features, train=True, transform=transform) \n    \n    test_val_data = FingerspellingDataset(features=features, train=False, transform=testVal_transform) \n    dataset_size = len(test_val_data)\n    val_data, test_data = torch.utils.data.random_split(test_val_data, [math.ceil(test_split * dataset_size), math.floor(val_split * dataset_size)])\n\n    # setup data loaders\n    train_loader = torch.utils.data.DataLoader(train_data, batch_size=BATCH_SIZE, shuffle=True)\n    val_loader = torch.utils.data.DataLoader(val_data, batch_size=BATCH_SIZE, shuffle=True)\n    test_loader = torch.utils.data.DataLoader(test_data, batch_size=BATCH_SIZE)\n\n    return train_loader, val_loader, test_loader","metadata":{"_uuid":"4f6a38c4-b8fd-4b67-8a48-dc2c6fe31a42","_cell_guid":"69a40e5b-4595-4769-ad87-771c15cae651","execution":{"iopub.status.busy":"2023-12-03T20:47:03.283784Z","iopub.execute_input":"2023-12-03T20:47:03.284421Z","iopub.status.idle":"2023-12-03T20:47:03.29309Z","shell.execute_reply.started":"2023-12-03T20:47:03.284395Z","shell.execute_reply":"2023-12-03T20:47:03.291986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Augmentations","metadata":{"_uuid":"8226a77b-697c-4f56-9429-2f50714518c2","_cell_guid":"660e6bfa-6581-4b95-b3e4-322147ef941a","trusted":true}},{"cell_type":"code","source":"class NormalizeAndFillNaNs(object):\n    def __init__(self):\n        super(NormalizeAndFillNaNs, self).__init__()\n    \n    def normalize(self,x):\n        nonan = x[~torch.isnan(x)].view(-1, x.shape[-1])\n        x = x - nonan.mean(0)[None, None, :]\n        x = x / nonan.std(0, unbiased=False)[None, None, :]\n        return x\n    \n    def fill_nans(self,x):\n        x[torch.isnan(x)] = 0\n        return x\n        \n    def __call__(self, sample):\n        x, phrase = sample['data'], sample['phrase']\n        #seq_len, 3* n_landmarks -> seq_len, n_landmarks, 3\n        \n        x = x.reshape(x.shape[0],3,-1).permute(0,2,1)\n        \n        # Normalize & fill nans\n        \n        x = self.normalize(x) # perhaps also remove this \n        x = self.fill_nans(x)\n        \n        return {\n            'data': x, \n            'phrase': phrase\n        }","metadata":{"_uuid":"63f3cc0c-e131-421b-bdfa-6e9a29dfc56f","_cell_guid":"13124b54-b3f0-4d0a-80e6-6c5ad2b210ab","execution":{"iopub.status.busy":"2023-12-03T20:47:03.294465Z","iopub.execute_input":"2023-12-03T20:47:03.294788Z","iopub.status.idle":"2023-12-03T20:47:03.30516Z","shell.execute_reply.started":"2023-12-03T20:47:03.294758Z","shell.execute_reply":"2023-12-03T20:47:03.304287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TokenizePhrase(object):\n    def __init__(self):\n        with open('/kaggle/input/asl-fingerspelling/character_to_prediction_index.json', \"r\") as f:\n            self.char_to_num = json.load(f)\n        n = len(self.char_to_num)\n        self.char_to_num[PAD_TOKEN] = n\n        self.char_to_num[START_TOKEN] = n + 1\n        self.char_to_num[END_TOKEN] = n + 2\n        \n        self.num_to_char = {j:i for i,j in self.char_to_num.items()}\n        \n    def __call__(self, sample):\n        left_hand, right_hand, left_pose, right_pose, data_mask, phrase = sample['left_hand'], sample['right_hand'], sample['left_pose'], sample['right_pose'], sample['data_mask'], sample['phrase']\n        \n        start_token_id = self.char_to_num[START_TOKEN]\n        end_token_id = self.char_to_num[END_TOKEN]\n        pad_token_id = self.char_to_num[PAD_TOKEN]\n        phrase_tokens = [self.char_to_num[char] for char in phrase]\n        if len(phrase_tokens) > MAX_PHRASE - 3:\n            phrase_tokens = phrase_tokens[:MAX_PHRASE - 3]\n        phrase_tokens = [start_token_id] + phrase_tokens + [end_token_id]\n        phrase_mask = [1] * len(phrase_tokens)\n        \n        to_pad = MAX_PHRASE - len(phrase_tokens)\n        phrase_tokens = torch.tensor(phrase_tokens + [pad_token_id] * to_pad)\n        phrase_mask = torch.tensor(phrase_mask + [0] * to_pad)\n        return {\n            'left_hand': left_hand, # tensor\n            'right_hand': right_hand, # tensor\n            'left_pose': left_pose, # tensor\n            'right_pose': right_pose, # tensor\n            'data_mask': data_mask, # tensor\n            \n            'phrase': phrase, # string\n            'phrase_tokens': phrase_tokens, # tensor (long)\n            'phrase_mask': phrase_mask # tensor (long)\n        }","metadata":{"_uuid":"03eee0a3-8cb9-473a-ad9d-f18fabe0faba","_cell_guid":"3f498f30-4564-4c1d-bbd4-82070cf73620","execution":{"iopub.status.busy":"2023-12-03T20:47:03.306212Z","iopub.execute_input":"2023-12-03T20:47:03.306497Z","iopub.status.idle":"2023-12-03T20:47:03.318153Z","shell.execute_reply.started":"2023-12-03T20:47:03.306472Z","shell.execute_reply":"2023-12-03T20:47:03.31729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class InterpolateOrPad(object):\n    def __init__(self, max_length: int = MAX_LEN):\n        self.max_length = max_length\n        \n    def __call__(self, sample):\n        data, phrase = sample['data'], sample['phrase']\n        diff = self.max_length - data.shape[0]\n        \n        # crop\n        if diff <= 0:\n            data = F.interpolate(data.permute(1,2,0),self.max_length).permute(2,0,1)\n            data_mask = torch.ones_like(data[:,0,0])\n            return {\n                'data': data,\n                'data_mask': data_mask,\n                'phrase': phrase\n            }\n        \n        # pad\n        coef = 0\n        padding = torch.ones((diff, data.shape[1], data.shape[2]))\n        data_mask = torch.ones_like(data[:,0,0])\n        data = torch.cat([data, padding * coef])\n        data_mask = torch.cat([data_mask, padding[:,0,0] * coef])\n        \n        return {\n            'data': data,\n            'data_mask': data_mask,\n            'phrase': phrase\n        }","metadata":{"_uuid":"00289c28-ec8c-42cf-8af6-9c15fc79dfe5","_cell_guid":"ce0dac4e-ca3e-4c18-adac-fcae67db0a4e","execution":{"iopub.status.busy":"2023-12-03T20:47:03.319331Z","iopub.execute_input":"2023-12-03T20:47:03.319619Z","iopub.status.idle":"2023-12-03T20:47:03.331417Z","shell.execute_reply.started":"2023-12-03T20:47:03.319595Z","shell.execute_reply":"2023-12-03T20:47:03.330546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SplitData(object):\n    def __init__(self, features: dict = FEATURES):\n        columns = extract_columns(features)\n        \n        self.X_IDX = [i for i, col in enumerate(columns)  if \"x_\" in col]\n        self.Y_IDX = [i for i, col in enumerate(columns)  if \"y_\" in col]\n        self.Z_IDX = [i for i, col in enumerate(columns)  if \"z_\" in col]\n        \n        self.RHAND_IDX = list(set([int(i/3) for i, col in enumerate(columns)  if \"right\" in col]))\n        self.LHAND_IDX = list(set([int(i/3) for i, col in enumerate(columns)  if  \"left\" in col]))\n        self.RPOSE_IDX = list(set([int(i/3) for i, col in enumerate(columns)  if  \"pose\" in col and int(col[-2:]) in FEATURES['pose']['right']]))\n        self.LPOSE_IDX = list(set([int(i/3) for i, col in enumerate(columns)  if  \"pose\" in col and int(col[-2:]) in FEATURES['pose']['left']]))\n        LEFT = self.LHAND_IDX + self.LPOSE_IDX\n        RIGHT = self.RHAND_IDX + self.RPOSE_IDX\n\n    def __call__(self, sample):\n        data, data_mask, phrase = sample['data'], sample['data_mask'], sample['phrase']\n        data, phrase = sample['data'], sample['phrase']\n        return {\n            'left_hand': data[:, self.LHAND_IDX],\n            'right_hand': data[:, self.RHAND_IDX],\n            'left_pose': data[:, self.LPOSE_IDX],\n            'right_pose': data[:, self.RPOSE_IDX],\n            'data_mask': data_mask,\n            \n            'phrase': phrase,\n        }","metadata":{"_uuid":"a4ab1751-d8e2-4a4f-9341-038ab95e554b","_cell_guid":"64a58238-a8a2-4ae2-9641-d2b5a07b4b71","execution":{"iopub.status.busy":"2023-12-03T20:47:03.332577Z","iopub.execute_input":"2023-12-03T20:47:03.332821Z","iopub.status.idle":"2023-12-03T20:47:03.345572Z","shell.execute_reply.started":"2023-12-03T20:47:03.332799Z","shell.execute_reply":"2023-12-03T20:47:03.344646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ToTensor(object):\n    def __call__(self, sample):\n        frames, phrase = sample['data'], sample['phrase']\n        return {\n            'data': torch.from_numpy(frames), \n            'phrase': phrase\n        }","metadata":{"_uuid":"982d25f8-f821-4056-9965-112390fd847e","_cell_guid":"cc0e3292-65d8-4fff-8b82-2dbf3e95e958","execution":{"iopub.status.busy":"2023-12-03T20:47:03.347059Z","iopub.execute_input":"2023-12-03T20:47:03.347327Z","iopub.status.idle":"2023-12-03T20:47:03.357897Z","shell.execute_reply.started":"2023-12-03T20:47:03.347304Z","shell.execute_reply":"2023-12-03T20:47:03.357054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Resample(object):\n    def __init__(self, rate=(0.8,1.2)):\n        self.rate = rate\n    \n    def interp1d_(self, x, new_size):\n        indices = torch.linspace(0, len(x) - 1, steps=new_size, dtype=torch.float32)\n        indices_floor = torch.floor(indices).to(torch.int64)\n        indices_frac = indices - indices_floor\n        indices_floor = torch.clamp(indices_floor, 0, len(x) - 2)\n \n        x0 = x[indices_floor]\n        x1 = x[indices_floor + 1]\n        indices_frac = indices_frac.view(-1, 1, 1)\n        new_x = x0 + (x1 - x0) * indices_frac\n        return new_x\n\n    def __call__(self, sample):\n        frames, phrase = sample['data'], sample['phrase']\n        if torch.rand(1)>0.8:\n            rate = torch.FloatTensor(1).uniform_(self.rate[0], self.rate[1])\n            length = frames.shape[0]\n            new_size = int(rate * length)\n            new_frames = self.interp1d_(frames, new_size)\n        else:\n            new_frames = frames\n        return {'data': new_frames, 'phrase': phrase}","metadata":{"_uuid":"48249082-27ef-42ee-ad25-9a5c8e850cd1","_cell_guid":"e3b1aa19-a5b7-4cdc-89fb-bff664200966","execution":{"iopub.status.busy":"2023-12-03T20:47:03.358908Z","iopub.execute_input":"2023-12-03T20:47:03.359191Z","iopub.status.idle":"2023-12-03T20:47:03.368803Z","shell.execute_reply.started":"2023-12-03T20:47:03.359168Z","shell.execute_reply":"2023-12-03T20:47:03.367902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SpatialRandomAffine(object):\n    def __init__(self, \n                 scale=(0.8, 1.2),\n                 shear=(-0.15, 0.15),\n                 shift=(-0.1, 0.1),\n                 degree=(-30, 30)\n                ):\n        self.scale = scale\n        self.shear = shear\n        self.shift = shift\n        self.degree = degree\n    \n    def __call__(self, sample):\n        data, phrase = sample['data'], sample['phrase']\n        if torch.rand(1)>0.75:\n            center = torch.tensor([0.5, 0.5])\n    \n            if self.scale is not None:\n                scale = torch.rand(1).item() * (self.scale[1] - self.scale[0]) + self.scale[0]\n                data = scale * data\n\n            if self.shear is not None:\n                xy = data[..., :2]\n                z = data[..., 2:]\n                shear_x = shear_y = torch.rand(1).item() * (self.shear[1] - self.shear[0]) + self.shear[0]\n                if torch.rand(1).item() < 0.5:\n                    shear_x = 0.0\n                else:\n                    shear_y = 0.0\n                shear_mat = torch.tensor([\n                    [1.0, shear_x],\n                    [shear_y, 1.0]\n                ])\n                xy = torch.matmul(xy, shear_mat)\n                center = center + torch.tensor([shear_y, shear_x])\n                data = torch.cat([xy, z], dim=-1)\n            \n            if self.degree is not None:\n                xy = data[..., :2]\n                z = data[..., 2:]\n                xy -= center\n                degree = torch.rand(1).item() * (self.degree[1] - self.degree[0]) + self.degree[0]\n                radian = degree / 180 * torch.tensor([3.14159265358979323846])\n                c = torch.cos(radian)\n                s = torch.sin(radian)\n                rotate_mat = torch.tensor([\n                    [c, s],\n                    [-s, c]\n                ])\n                xy = torch.matmul(xy, rotate_mat)\n                xy = xy + center\n                data = torch.cat([xy, z], dim=-1)\n\n            if self.shift is not None:\n                shift = torch.rand(1).item() * (self.shift[1] - self.shift[0]) + self.shift[0]\n                data = data + shift\n            \n        return {'data': data, 'phrase': phrase}","metadata":{"_uuid":"948a4fec-791b-4d4f-8cb1-b1fb99aff5f3","_cell_guid":"8966e4bd-8e81-43e7-987a-67c23b15ae24","execution":{"iopub.status.busy":"2023-12-03T20:47:03.370206Z","iopub.execute_input":"2023-12-03T20:47:03.37054Z","iopub.status.idle":"2023-12-03T20:47:03.386756Z","shell.execute_reply.started":"2023-12-03T20:47:03.370503Z","shell.execute_reply":"2023-12-03T20:47:03.38589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TemporalMask(object):\n    def __init__(self, size=(0.2,0.4), mask_value=float('nan')):\n        self.size = size\n        self.mask_value = mask_value\n        \n    def __call__(self, sample):\n        data, phrase = sample['data'], sample['phrase']\n        if torch.rand(1)>0.5:\n            l = data.shape[0]\n            mask_size = torch.rand(1).item() * (self.size[1] - self.size[0]) + self.size[0]\n            mask_size = int(l * mask_size)\n            mask_offset = torch.randint(0, l - mask_size + 1, (1,)).item()\n            mask_indices = torch.arange(mask_offset, mask_offset + mask_size).unsqueeze(1)\n            mask = torch.full((mask_size, 52, 3), self.mask_value, dtype=data.dtype)\n            data[mask_indices,...] = mask.unsqueeze(1)\n        \n        return {'data': data, 'phrase': phrase}","metadata":{"_uuid":"8d191d51-85a5-46e9-a738-ba6a4fda8c74","_cell_guid":"89e6f262-bf9c-4ab1-a8e3-ada589bc9fb2","execution":{"iopub.status.busy":"2023-12-03T20:47:03.388118Z","iopub.execute_input":"2023-12-03T20:47:03.388457Z","iopub.status.idle":"2023-12-03T20:47:03.399985Z","shell.execute_reply.started":"2023-12-03T20:47:03.38842Z","shell.execute_reply":"2023-12-03T20:47:03.399169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SpatialMask(object):\n    def __init__(self, size=(0.2,0.4), mask_value=float('nan')):\n        self.size = size\n        self.mask_value = mask_value\n        \n    def __call__(self, sample):\n        # TODO: determine if this works as intended (does x/y refer to xyz coordinates?)\n        xyz, phrase = sample['data'], sample['phrase']\n        if torch.rand(1)>0.5:\n            mask_offset_y = torch.rand(1).item()\n            mask_offset_x = torch.rand(1).item()\n            mask_size = torch.rand(1).item() * (self.size[1] - self.size[0]) + self.size[0]\n\n            mask_x = (mask_offset_x < xyz[..., 0]) & (xyz[..., 0] < mask_offset_x + mask_size)\n            mask_y = (mask_offset_y < xyz[..., 1]) & (xyz[..., 1] < mask_offset_y + mask_size)\n            mask = mask_x & mask_y\n\n            xyz = torch.where(mask.unsqueeze(-1), torch.tensor(self.mask_value), xyz)\n\n        \n        return {'data': xyz, 'phrase': phrase}","metadata":{"_uuid":"4507964a-8078-4893-96b0-64bf28d8952b","_cell_guid":"9e3a8f3d-6f25-4183-947a-734de1071955","_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-12-03T20:47:03.404776Z","iopub.execute_input":"2023-12-03T20:47:03.405165Z","iopub.status.idle":"2023-12-03T20:47:03.413839Z","shell.execute_reply.started":"2023-12-03T20:47:03.40514Z","shell.execute_reply":"2023-12-03T20:47:03.412684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FlipLeftRight(object):\n    def __init__(self, features: dict = FEATURES):\n        columns = extract_columns(features)\n\n        \n        self.RHAND_IDX = list(set([int(i/3) for i, col in enumerate(columns)  if \"right\" in col]))\n        self.LHAND_IDX = list(set([int(i/3) for i, col in enumerate(columns)  if  \"left\" in col]))\n        self.RPOSE_IDX = list(set([int(i/3) for i, col in enumerate(columns)  if  \"pose\" in col and int(col[-2:]) in FEATURES['pose']['right']]))\n        self.LPOSE_IDX = list(set([int(i/3) for i, col in enumerate(columns)  if  \"pose\" in col and int(col[-2:]) in FEATURES['pose']['left']]))\n        self.left = self.LHAND_IDX + self.LPOSE_IDX\n        self.right = self.RHAND_IDX + self.RPOSE_IDX\n\n        \n    # TODO: fix, not sure if non 3d data format will work (ie requires xyz to extend into 3rd dimension)\n    def __call__(self, sample):\n        xyz, phrase = sample['data'], sample['phrase']\n        if torch.rand(1)>0.5:\n            x, y, z = torch.unbind(xyz, dim=-1)\n            x = 1 - x\n            new_xyz = torch.stack([x, y, z], dim=-1)\n            new_xyz = new_xyz.transpose(0, 1)\n\n            l_x = new_xyz[self.left]\n            r_x = new_xyz[self.right]\n\n            for i in range(len(self.left)):\n                new_xyz[self.left[i]] = r_x[i]\n                new_xyz[self.right[i]] = l_x[i]\n\n            new_xyz = new_xyz.transpose(0, 1)\n        \n        else:\n            new_xyz = xyz\n\n        return {'data': new_xyz, 'phrase': phrase}","metadata":{"_uuid":"724b98ad-f556-4de4-8bea-d3efef774e56","_cell_guid":"04534cee-9a10-4b57-8967-d3dab2252ca9","execution":{"iopub.status.busy":"2023-12-03T20:47:03.415088Z","iopub.execute_input":"2023-12-03T20:47:03.415418Z","iopub.status.idle":"2023-12-03T20:47:03.428089Z","shell.execute_reply.started":"2023-12-03T20:47:03.415387Z","shell.execute_reply":"2023-12-03T20:47:03.427049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TemporalCrop(object):\n    def __init__(self, length=32):\n        self.length = length\n        \n    def __call__(self, sample):\n        data, phrase = sample['data'], sample['phrase']\n        \n        l = data.shape[0]\n        # TODO remove self.length<l-1 once fixed\n        if self.length is not None and self.length<l-1:\n            offset = torch.randint(0, l - self.length + 1, (1,)).item()\n            data = data[offset:offset + self.length]\n            \n        return {\n            'data': data, \n            'phrase': phrase\n        }","metadata":{"_uuid":"db335334-c4f0-403e-ab3a-8ba5fd7c7778","_cell_guid":"567e454c-717c-4333-b544-0ab2b4dd055d","execution":{"iopub.status.busy":"2023-12-03T20:47:03.429329Z","iopub.execute_input":"2023-12-03T20:47:03.429621Z","iopub.status.idle":"2023-12-03T20:47:03.440668Z","shell.execute_reply.started":"2023-12-03T20:47:03.429597Z","shell.execute_reply":"2023-12-03T20:47:03.439877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"_uuid":"a286ecca-590c-4e2f-8af6-50af90010dea","_cell_guid":"0194c7b2-0cd1-4e63-a9b2-33ad550da1b3","trusted":true}},{"cell_type":"code","source":"\nclass Swish(nn.Module):\n    def __init__(self) -> None:\n        super(Swish, self).__init__()\n\n    def forward(self, inputs: Tensor) -> Tensor:\n        return inputs * inputs.sigmoid()\n\n\nclass GLU(nn.Module):\n    def __init__(self, dim: int) -> None:\n        super(GLU, self).__init__()\n        self.dim = dim\n\n    def forward(self, inputs: Tensor) -> Tensor:\n        outputs, gate = inputs.chunk(2, dim=self.dim)\n        return outputs * gate.sigmoid()\n    \nclass RelativeMultiHeadAttention(nn.Module):\n    \"\"\"\n    Multi-head attention with relative positional encoding.\n    This concept was proposed in the \"Transformer-XL: Attentive Language Models Beyond a Fixed-Length Context\"\n    Args:\n        d_model (int): The dimension of model\n        num_heads (int): The number of attention heads.\n        dropout_p (float): probability of dropout\n    Inputs: query, key, value, pos_embedding, mask\n        - **query** (batch, time, dim): Tensor containing query vector\n        - **key** (batch, time, dim): Tensor containing key vector\n        - **value** (batch, time, dim): Tensor containing value vector\n        - **pos_embedding** (batch, time, dim): Positional embedding tensor\n        - **mask** (batch, 1, time2) or (batch, time1, time2): Tensor containing indices to be masked\n    Returns:\n        - **outputs**: Tensor produces by relative multi head attention module.\n    \"\"\"\n\n    def __init__(\n        self,\n        d_model: int = 512,\n        num_heads: int = 16,\n        dropout_p: float = 0.1,\n    ):\n        super(RelativeMultiHeadAttention, self).__init__()\n        assert d_model % num_heads == 0, \"d_model % num_heads should be zero.\"\n        self.d_model = d_model\n        self.d_head = int(d_model / num_heads)\n        self.num_heads = num_heads\n        self.sqrt_dim = math.sqrt(self.d_head)\n\n        self.query_proj = nn.Linear(d_model, d_model)\n        self.key_proj = nn.Linear(d_model, d_model)\n        self.value_proj = nn.Linear(d_model, d_model)\n        self.pos_proj = nn.Linear(d_model, d_model, bias=False)\n\n        self.dropout = nn.Dropout(p=dropout_p)\n        self.u_bias = nn.Parameter(torch.Tensor(self.num_heads, self.d_head))\n        self.v_bias = nn.Parameter(torch.Tensor(self.num_heads, self.d_head))\n        torch.nn.init.xavier_uniform_(self.u_bias)\n        torch.nn.init.xavier_uniform_(self.v_bias)\n\n        self.out_proj = nn.Linear(d_model, d_model)\n\n    def forward(\n        self,\n        query: Tensor,\n        key: Tensor,\n        value: Tensor,\n        pos_embedding: Tensor,\n        mask: Optional[Tensor] = None,\n    ) -> Tensor:\n        batch_size = value.size(0)\n\n        query = self.query_proj(query).view(batch_size, -1, self.num_heads, self.d_head)\n        key = self.key_proj(key).view(batch_size, -1, self.num_heads, self.d_head).permute(0, 2, 1, 3)\n        value = self.value_proj(value).view(batch_size, -1, self.num_heads, self.d_head).permute(0, 2, 1, 3)\n        pos_embedding = self.pos_proj(pos_embedding).view(batch_size, -1, self.num_heads, self.d_head)\n\n        content_score = torch.matmul((query + self.u_bias).transpose(1, 2), key.transpose(2, 3))\n        pos_score = torch.matmul((query + self.v_bias).transpose(1, 2), pos_embedding.permute(0, 2, 3, 1))\n        pos_score = self._relative_shift(pos_score)\n\n        score = (content_score + pos_score) / self.sqrt_dim\n\n        if mask is not None:\n            mask = mask.unsqueeze(1)\n            score.masked_fill_(mask, -1e9)\n\n        attn = F.softmax(score, -1)\n        attn = self.dropout(attn)\n\n        context = torch.matmul(attn, value).transpose(1, 2)\n        context = context.contiguous().view(batch_size, -1, self.d_model)\n\n        return self.out_proj(context)\n\n    def _relative_shift(self, pos_score: Tensor) -> Tensor:\n        batch_size, num_heads, seq_length1, seq_length2 = pos_score.size()\n        zeros = pos_score.new_zeros(batch_size, num_heads, seq_length1, 1)\n        padded_pos_score = torch.cat([zeros, pos_score], dim=-1)\n\n        padded_pos_score = padded_pos_score.view(batch_size, num_heads, seq_length2 + 1, seq_length1)\n        pos_score = padded_pos_score[:, :, 1:].view_as(pos_score)[:, :, :, : seq_length2 // 2 + 1]\n\n        return pos_score\n\n\nclass MultiHeadedSelfAttentionModule(nn.Module):\n    \"\"\"\n    Args:\n        d_model (int): The dimension of model\n        num_heads (int): The number of attention heads.\n        dropout_p (float): probability of dropout\n    Inputs: inputs, mask\n        - **inputs** (batch, time, dim): Tensor containing input vector\n        - **mask** (batch, 1, time2) or (batch, time1, time2): Tensor containing indices to be masked\n    Returns:\n        - **outputs** (batch, time, dim): Tensor produces by relative multi headed self attention module.\n    \"\"\"\n\n    def __init__(self, d_model: int, num_heads: int, dropout_p: float = 0.1):\n        super(MultiHeadedSelfAttentionModule, self).__init__()\n        self.positional_encoding = RelPositionalEncoding(d_model)\n        self.attention = RelativeMultiHeadAttention(d_model, num_heads, dropout_p)\n        self.dropout = nn.Dropout(p=dropout_p)\n\n    def forward(self, inputs: Tensor, mask: Optional[Tensor] = None):\n        batch_size = inputs.size(0)\n        pos_embedding = self.positional_encoding(inputs)\n        pos_embedding = pos_embedding.repeat(batch_size, 1, 1)\n\n        outputs = self.attention(inputs, inputs, inputs, pos_embedding=pos_embedding, mask=mask)\n\n        return self.dropout(outputs)\n    \nclass DepthwiseConv2dSubsampling(nn.Module):\n    \"\"\"\n    Depthwise Convolutional 2D subsampling (to 1/4 length)\n\n    Args:\n        in_channels (int): Number of channels in the input image\n        out_channels (int): Number of channels produced by the convolution\n    Inputs: inputs, input_lengths\n        - **inputs** (batch, time, dim): Tensor containing sequence of inputs\n        - **input_lengths** (batch): list of sequence input lengths\n    Returns: outputs, output_lengths\n        - **outputs** (batch, time, dim): Tensor produced by the convolution\n        - **output_lengths** (batch): list of sequence output lengths\n    \"\"\"\n\n    def __init__(self, in_channels: int, out_channels: int) -> None:\n        super(DepthwiseConv2dSubsampling, self).__init__()\n        self.sequential = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=2),\n            nn.ReLU(),\n            DepthwiseConv2d(out_channels, out_channels, kernel_size=3, stride=2),\n            nn.ReLU(),\n        )\n\n    def forward(self, inputs: Tensor, input_lengths: Tensor) -> Tuple[Tensor, Tensor]:\n        outputs = self.sequential(inputs.unsqueeze(1))\n        batch_size, channels, subsampled_lengths, subsampled_dim = outputs.size()\n\n        outputs = outputs.permute(0, 2, 1, 3)\n        outputs = outputs.contiguous().view(batch_size, subsampled_lengths, channels * subsampled_dim)\n\n        output_lengths = input_lengths >> 2\n        output_lengths -= 1\n\n        return outputs, output_lengths\n\n\nclass DepthwiseConv2d(nn.Module):\n    \"\"\"\n    When groups == in_channels and out_channels == K * in_channels, where K is a positive integer,\n    this operation is termed in literature as depthwise convolution.\n    ref : https://pytorch.org/docs/stable/generated/torch.nn.Conv2d.html\n\n    Args:\n        in_channels (int): Number of channels in the input\n        out_channels (int): Number of channels produced by the convolution\n        kernel_size (int or tuple): Size of the convolving kernel\n        stride (int, optional): Stride of the convolution. Default: 2\n        padding (int or tuple, optional): Zero-padding added to both sides of the input. Default: 0\n    Inputs: inputs\n        - **inputs** (batch, in_channels, time): Tensor containing input vector\n    Returns: outputs\n        - **outputs** (batch, out_channels, time): Tensor produces by depthwise 2-D convolution.\n    \"\"\"\n\n    def __init__(\n        self,\n        in_channels: int,\n        out_channels: int,\n        kernel_size: Union[int, Tuple],\n        stride: int = 2,\n        padding: int = 0,\n    ) -> None:\n        super(DepthwiseConv2d, self).__init__()\n        assert out_channels % in_channels == 0, \"out_channels should be constant multiple of in_channels\"\n        self.conv = nn.Conv2d(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            kernel_size=kernel_size,\n            stride=stride,\n            padding=padding,\n            groups=in_channels,\n        )\n\n    def forward(self, inputs: Tensor) -> Tensor:\n        return self.conv(inputs)\n\n\nclass DepthwiseConv1d(nn.Module):\n    \"\"\"\n    When groups == in_channels and out_channels == K * in_channels, where K is a positive integer,\n    this operation is termed in literature as depthwise convolution.\n    ref : https://pytorch.org/docs/stable/generated/torch.nn.Conv2d.html\n\n    Args:\n        in_channels (int): Number of channels in the input\n        out_channels (int): Number of channels produced by the convolution\n        stride (int, optional): Stride of the convolution. Default: 1\n        padding (int or tuple, optional): Zero-padding added to both sides of the input. Default: 0\n        bias (bool, optional): If True, adds a learnable bias to the output. Default: False\n    Inputs: inputs\n        - **inputs** (batch, in_channels, time): Tensor containing input vector\n    Returns: outputs\n        - **outputs** (batch, out_channels, time): Tensor produces by depthwise 1-D convolution.\n    \"\"\"\n\n    def __init__(\n        self,\n        in_channels: int,\n        out_channels: int,\n        kernel_size: int,\n        stride: int = 1,\n        padding: int = 0,\n        bias: bool = False,\n    ) -> None:\n        super(DepthwiseConv1d, self).__init__()\n        assert out_channels % in_channels == 0, \"out_channels should be constant multiple of in_channels\"\n        self.conv = nn.Conv1d(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            kernel_size=kernel_size,\n            groups=in_channels,\n            stride=stride,\n            padding=padding,\n            bias=bias,\n        )\n\n    def forward(self, inputs: Tensor) -> Tensor:\n        return self.conv(inputs)\n\n\nclass PointwiseConv1d(nn.Module):\n    \"\"\"\n    When kernel size == 1 conv1d, this operation is termed in literature as pointwise convolution.\n    This operation often used to match dimensions.\n\n    Args:\n        in_channels (int): Number of channels in the input\n        out_channels (int): Number of channels produced by the convolution\n        stride (int, optional): Stride of the convolution. Default: 1\n        padding (int or tuple, optional): Zero-padding added to both sides of the input. Default: 0\n        bias (bool, optional): If True, adds a learnable bias to the output. Default: True\n    Inputs: inputs\n        - **inputs** (batch, in_channels, time): Tensor containing input vector\n    Returns: outputs\n        - **outputs** (batch, out_channels, time): Tensor produces by pointwise 1-D convolution.\n    \"\"\"\n\n    def __init__(\n        self,\n        in_channels: int,\n        out_channels: int,\n        stride: int = 1,\n        padding: int = 0,\n        bias: bool = True,\n    ) -> None:\n        super(PointwiseConv1d, self).__init__()\n        self.conv = nn.Conv1d(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            kernel_size=1,\n            stride=stride,\n            padding=padding,\n            bias=bias,\n        )\n\n    def forward(self, inputs: Tensor) -> Tensor:\n        return self.conv(inputs)\n\n\nclass ConvModule(nn.Module):\n    \"\"\"\n    Convolution module starts with a pointwise convolution and a gated linear unit (GLU).\n    This is followed by a single 1-D depthwise convolution layer. Batchnorm is deployed just after the convolution\n    to aid training deep models.\n\n    Args:\n        in_channels (int): Number of channels in the input\n        kernel_size (int or tuple, optional): Size of the convolving kernel Default: 31\n        dropout_p (float, optional): probability of dropout\n    Inputs: inputs\n        inputs (batch, time, dim): Tensor contains input sequences\n    Outputs: outputs\n        outputs (batch, time, dim): Tensor produces by squeezeformer convolution module.\n    \"\"\"\n\n    def __init__(\n        self,\n        in_channels: int,\n        kernel_size: int = 31,\n        expansion_factor: int = 2,\n        dropout_p: float = 0.1,\n    ) -> None:\n        super(ConvModule, self).__init__()\n        assert (kernel_size - 1) % 2 == 0, \"kernel_size should be a odd number for 'SAME' padding\"\n        assert expansion_factor == 2, \"Currently, Only Supports expansion_factor 2\"\n\n        self.pw_conv_1 = PointwiseConv1d(in_channels, in_channels * expansion_factor, stride=1, padding=0, bias=True)\n        self.act1 = GLU(dim=1)\n        self.dw_conv = DepthwiseConv1d(in_channels, in_channels, kernel_size, stride=1, padding=(kernel_size - 1) // 2)\n        self.bn = nn.BatchNorm1d(in_channels)\n        self.act2 = Swish()\n        self.pw_conv_2 = PointwiseConv1d(in_channels, in_channels, stride=1, padding=0, bias=True)\n        self.do = nn.Dropout(p=dropout_p)\n\n    # mask_pad = mask.bool().unsqueeze(1)\n    def forward(self, x, mask_pad):\n        \"\"\"Compute convolution module.\n        Args:\n            x (torch.Tensor): Input tensor (#batch, time, channels).\n            mask_pad (torch.Tensor): used for batch padding (#batch, 1, time),\n                (0, 0, 0) means fake mask.\n        Returns:\n            torch.Tensor: Output tensor (#batch, time, channels).\n        Reference for masking : https://github.com/Ascend/ModelZoo-PyTorch/blob/master/PyTorch/built-in/audio/Wenet_Conformer_for_Pytorch/wenet/transformer/convolution.py#L26\n        \"\"\"\n        # mask batch padding\n        x = x.transpose(1, 2)\n        if mask_pad.size(2) > 0:  # time > 0\n            x = x.masked_fill(~mask_pad, 0.0)\n        x = self.pw_conv_1(x)\n        x = self.act1(x)\n        x = self.dw_conv(x)\n        # torch.Size([4, 128, 384])\n        x_bn = x.permute(0,2,1).reshape(-1, x.shape[1])\n        mask_bn = mask_pad.view(-1)\n        x_bn[mask_bn] = self.bn(x_bn[mask_bn])\n        x = x_bn.view(x.permute(0,2,1).shape).permute(0,2,1)\n        '''\n        x = self.bn(x)\n        '''\n        x = self.act2(x)\n        x = self.pw_conv_2(x)\n        x = self.do(x)\n        # mask batch padding\n        if mask_pad.size(2) > 0:  # time > 0\n            x = x.masked_fill(~mask_pad, 0.0)\n        x = x.transpose(1, 2)\n        return x\n\n\n\nclass TimeReductionLayer(nn.Module):\n    def __init__(\n        self,\n        in_channels: int = 1,\n        out_channels: int = 1,\n        kernel_size: int = 3,\n        stride: int = 2,\n    ) -> None:\n        super(TimeReductionLayer, self).__init__()\n        self.sequential = nn.Sequential(\n            DepthwiseConv2d(\n                in_channels=in_channels,\n                out_channels=out_channels,\n                kernel_size=kernel_size,\n                stride=stride,\n            ),\n            Swish(),\n        )\n\n    def forward(self, inputs: Tensor, input_lengths: Tensor) -> Tuple[Tensor, Tensor]:\n        outputs = self.sequential(inputs.unsqueeze(1))\n        batch_size, channels, subsampled_lengths, subsampled_dim = outputs.size()\n\n        outputs = outputs.permute(0, 2, 1, 3)\n        outputs = outputs.contiguous().view(batch_size, subsampled_lengths, channels * subsampled_dim)\n\n        output_lengths = input_lengths >> 1\n        output_lengths -= 1\n        return outputs, output_lengths\n    \nclass FeedForwardModule(nn.Module):\n    \"\"\"\n    Feed Forward Module follow pre-norm residual units and apply layer normalization within the residual unit\n    and on the input before the first linear layer. This module also apply Swish activation and dropout, which helps\n    regularizing the network.\n\n    Args:\n        encoder_dim (int): Dimension of squeezeformer encoder\n        expansion_factor (int): Expansion factor of feed forward module.\n        dropout_p (float): Ratio of dropout\n    Inputs: inputs\n        - **inputs** (batch, time, dim): Tensor contains input sequences\n    Outputs: outputs\n        - **outputs** (batch, time, dim): Tensor produces by feed forward module.\n    \"\"\"\n\n    def __init__(\n        self,\n        encoder_dim: int = 512,\n        expansion_factor: int = 4,\n        dropout_p: float = 0.1,\n    ) -> None:\n        super(FeedForwardModule, self).__init__()\n        self.sequential = nn.Sequential(\n            nn.Linear(encoder_dim, encoder_dim * expansion_factor, bias=True),\n            Swish(),\n            nn.Dropout(p=dropout_p),\n            nn.Linear(encoder_dim * expansion_factor, encoder_dim, bias=True),\n            nn.Dropout(p=dropout_p),\n        )\n\n    def forward(self, inputs: Tensor) -> Tensor:\n        return self.sequential(inputs)\n\n\nclass RelPositionalEncoding(nn.Module):\n    \"\"\"\n    Relative positional encoding module.\n    Args:\n        d_model: Embedding dimension.\n        max_len: Maximum input length.\n    \"\"\"\n\n    def __init__(self, d_model: int = 512, max_len: int = 5000) -> None:\n        super(RelPositionalEncoding, self).__init__()\n        self.d_model = d_model\n        self.pe = None\n        self.extend_pe(torch.tensor(0.0).expand(1, max_len))\n\n    def extend_pe(self, x):\n        if self.pe is not None:\n            if self.pe.size(1) >= x.size(1) * 2 - 1:\n                if self.pe.dtype != x.dtype or self.pe.device != x.device:\n                    self.pe = self.pe.to(dtype=x.dtype, device=x.device)\n                return\n\n        pe_positive = torch.zeros(x.size(1), self.d_model)\n        pe_negative = torch.zeros(x.size(1), self.d_model)\n        position = torch.arange(0, x.size(1), dtype=torch.float32).unsqueeze(1)\n        div_term = torch.exp(\n            torch.arange(0, self.d_model, 2, dtype=torch.float32) * -(math.log(10000.0) / self.d_model)\n        )\n        pe_positive[:, 0::2] = torch.sin(position * div_term)\n        pe_positive[:, 1::2] = torch.cos(position * div_term)\n        pe_negative[:, 0::2] = torch.sin(-1 * position * div_term)\n        pe_negative[:, 1::2] = torch.cos(-1 * position * div_term)\n\n        pe_positive = torch.flip(pe_positive, [0]).unsqueeze(0)\n        pe_negative = pe_negative[1:].unsqueeze(0)\n        pe = torch.cat([pe_positive, pe_negative], dim=1)\n        self.pe = pe.to(device=x.device, dtype=x.dtype)\n\n    def forward(self, x: torch.Tensor):\n        \"\"\"\n        Args:\n            x : Input tensor B X T X C\n        Returns:\n            torch.Tensor: Encoded tensor B X T X C\n        \"\"\"\n        self.extend_pe(x)\n        pos_emb = self.pe[\n            :,\n            self.pe.size(1) // 2 - x.size(1) + 1 : self.pe.size(1) // 2 + x.size(1),\n        ]\n        return pos_emb\n\n\nclass ResidualConnectionModule(nn.Module):\n    \"\"\"\n    Residual Connection Module.\n    outputs = (module(inputs) x module_factor + inputs x input_factor)\n    \"\"\"\n\n    def __init__(self, module: nn.Module, module_factor: float = 1.0) -> None:\n        super(ResidualConnectionModule, self).__init__()\n        self.module = module\n        self.module_factor = module_factor\n\n    def forward(self, inputs: Tensor) -> Tensor:\n        return (self.module(inputs) * self.module_factor) + inputs\n\n\nclass Transpose(nn.Module):\n    \"\"\"Wrapper class of torch.transpose() for Sequential module.\"\"\"\n\n    def __init__(self, shape: tuple) -> None:\n        super(Transpose, self).__init__()\n        self.shape = shape\n\n    def forward(self, x: Tensor) -> Tensor:\n        return x.transpose(*self.shape)\n\n\ndef recover_resolution(inputs: Tensor) -> Tensor:\n    outputs = list()\n\n    for idx in range(inputs.size(1) * 2):\n        outputs.append(inputs[:, idx // 2, :])\n    return torch.stack(outputs, dim=1)\n\nclass SqueezeformerEncoder(nn.Module):\n    \"\"\"\n    Squeezeformer encoder first processes the input with a convolution subsampling layer and then\n    with a number of squeezeformer blocks.\n\n    Args:\n        input_dim (int, optional): Dimension of input vector\n        encoder_dim (int, optional): Dimension of squeezeformer encoder\n        num_layers (int, optional): Number of squeezeformer blocks\n        reduce_layer_index (int, optional): The layer index to reduce sequence length\n        recover_layer_index (int, optional): The layer index to recover sequence length\n        num_attention_heads (int, optional): Number of attention heads\n        feed_forward_expansion_factor (int, optional): Expansion factor of feed forward module\n        conv_expansion_factor (int, optional): Expansion factor of squeezeformer convolution module\n        feed_forward_dropout_p (float, optional): Probability of feed forward module dropout\n        attention_dropout_p (float, optional): Probability of attention module dropout\n        conv_dropout_p (float, optional): Probability of squeezeformer convolution module dropout\n        conv_kernel_size (int or tuple, optional): Size of the convolving kernel\n        half_step_residual (bool): Flag indication whether to use half step residual or not\n    Inputs: inputs, input_lengths\n        - **inputs** (batch, time, dim): Tensor containing input vector\n        - **input_lengths** (batch): list of sequence input lengths\n    Returns: outputs, output_lengths\n        - **outputs** (batch, out_channels, time): Tensor produces by squeezeformer encoder.\n        - **output_lengths** (batch): list of sequence output lengths\n    \"\"\"\n\n    def __init__(\n        self,\n        input_dim: int = 80,\n        encoder_dim: int = 512,\n        num_layers: int = 16,\n        num_attention_heads: int = 8,\n        feed_forward_expansion_factor: int = 4,\n        conv_expansion_factor: int = 2,\n        input_dropout_p: float = 0.1,\n        feed_forward_dropout_p: float = 0.1,\n        attention_dropout_p: float = 0.1,\n        conv_dropout_p: float = 0.1,\n        conv_kernel_size: int = 31,\n    ):\n        super(SqueezeformerEncoder, self).__init__()\n        self.num_layers = num_layers\n        self.recover_tensor = None\n\n        self.blocks = nn.ModuleList()\n        for idx in range(num_layers):\n            self.blocks.append(\n                SqueezeformerBlock(\n                    encoder_dim=encoder_dim,\n                    num_attention_heads=num_attention_heads,\n                    feed_forward_expansion_factor=feed_forward_expansion_factor,\n                    conv_expansion_factor=conv_expansion_factor,\n                    feed_forward_dropout_p=feed_forward_dropout_p,\n                    attention_dropout_p=attention_dropout_p,\n                    conv_dropout_p=conv_dropout_p,\n                    conv_kernel_size=conv_kernel_size,\n                )\n            )\n\n    def count_parameters(self) -> int:\n        \"\"\"Count parameters of encoder\"\"\"\n        return sum([p.numel for p in self.parameters()])\n\n    def forward(self, x: Tensor, mask: Tensor):\n        \"\"\"\n        Forward propagate a `inputs` for  encoder training.\n        Args:\n            inputs (torch.FloatTensor): A input sequence passed to encoder. Typically for inputs this will be a padded\n                `FloatTensor` of size ``(batch, seq_length, dimension)``.\n            input_lengths (torch.LongTensor): The length of input tensor. ``(batch)``\n        Returns:\n            (Tensor, Tensor)\n            * outputs (torch.FloatTensor): A output sequence of encoder. `FloatTensor` of size\n                ``(batch, seq_length, dimension)``\n            * output_lengths (torch.LongTensor): The length of output tensor. ``(batch)``\n        \"\"\"\n\n        for idx, block in enumerate(self.blocks):\n            x = block(x, mask)\n\n        return x\n\n\ndef make_scale(encoder_dim):\n    scale = torch.nn.Parameter(torch.tensor([1.] * encoder_dim)[None,None,:])\n    bias = torch.nn.Parameter(torch.tensor([0.] * encoder_dim)[None,None,:])\n    return scale, bias\n\nclass SqueezeformerBlock(nn.Module):\n    \"\"\"\n    SqueezeformerBlock is a simpler block structure similar to the standard Transformer block,\n    where the MHA and convolution modules are each directly followed by a single feed forward module.\n\n    Args:\n        encoder_dim (int, optional): Dimension of squeezeformer encoder\n        num_attention_heads (int, optional): Number of attention heads\n        feed_forward_expansion_factor (int, optional): Expansion factor of feed forward module\n        conv_expansion_factor (int, optional): Expansion factor of squeezeformer convolution module\n        feed_forward_dropout_p (float, optional): Probability of feed forward module dropout\n        attention_dropout_p (float, optional): Probability of attention module dropout\n        conv_dropout_p (float, optional): Probability of squeezeformer convolution module dropout\n        conv_kernel_size (int or tuple, optional): Size of the convolving kernel\n        half_step_residual (bool): Flag indication whether to use half step residual or not\n    Inputs: inputs\n        - **inputs** (batch, time, dim): Tensor containing input vector\n    Returns: outputs\n        - **outputs** (batch, time, dim): Tensor produces by squeezeformer block.\n    \"\"\"\n\n    def __init__(\n        self,\n        encoder_dim: int = 512,\n        num_attention_heads: int = 8,\n        feed_forward_expansion_factor: int = 4,\n        conv_expansion_factor: int = 2,\n        feed_forward_dropout_p: float = 0.1,\n        attention_dropout_p: float = 0.1,\n        conv_dropout_p: float = 0.1,\n        conv_kernel_size: int = 31,\n    ):\n        super(SqueezeformerBlock, self).__init__()\n        \n        self.scale_mhsa, self.bias_mhsa = make_scale(encoder_dim)\n        self.scale_ff_mhsa, self.bias_ff_mhsa = make_scale(encoder_dim)\n        self.scale_conv, self.bias_conv = make_scale(encoder_dim)\n        self.scale_ff_conv, self.bias_ff_conv = make_scale(encoder_dim)\n        \n        self.mhsa = MultiHeadedSelfAttentionModule(\n                    d_model=encoder_dim,\n                    num_heads=num_attention_heads,\n                    dropout_p=attention_dropout_p,)\n        self.ln_mhsa = nn.LayerNorm(encoder_dim)\n        self.ff_mhsa = FeedForwardModule(\n                    encoder_dim=encoder_dim,\n                    expansion_factor=feed_forward_expansion_factor,\n                    dropout_p=feed_forward_dropout_p,\n                )\n        self.ln_ff_mhsa = nn.LayerNorm(encoder_dim)\n        self.conv = ConvModule(\n                    in_channels=encoder_dim,\n                    kernel_size=conv_kernel_size,\n                    expansion_factor=conv_expansion_factor,\n                    dropout_p=conv_dropout_p,\n                )\n        self.ln_conv = nn.LayerNorm(encoder_dim)\n        self.ff_conv = FeedForwardModule(\n                    encoder_dim=encoder_dim,\n                    expansion_factor=feed_forward_expansion_factor,\n                    dropout_p=feed_forward_dropout_p,\n                )\n        self.ln_ff_conv = nn.LayerNorm(encoder_dim)\n\n\n    def forward(self, x, mask):\n        mask_pad = ( mask).long().bool().unsqueeze(1)\n        mask_pad = ~( mask_pad.permute(0, 2,1) * mask_pad)\n        mask_flat = mask.view(-1).bool()\n        bs, slen, nfeats = x.shape\n        \n        residual = x\n        x = x * self.scale_mhsa + self.bias_mhsa\n        x = residual + self.mhsa(x, mask_pad )\n        # Skip pad #1\n        x_skip = x.view(-1, x.shape[-1])\n        x = x_skip[mask_flat].unsqueeze(0)\n        \n        x = self.ln_mhsa(x)\n\n        residual = x\n        x = x * self.scale_ff_mhsa + self.bias_ff_mhsa\n        x = residual + self.ff_mhsa(x)\n        x = self.ln_ff_mhsa(x)\n        \n        \n        # Unskip pad #1\n        x_skip[mask_flat] = x[0]\n        x = x_skip.view(bs, slen, nfeats)\n        residual = x\n        # torch.Size([16, 384, 128])\n        x = x * self.scale_conv + self.bias_conv\n        x = residual + self.conv(x, mask_pad = mask.bool().unsqueeze(1))\n        # Skip pad #2\n        x_skip = x.view(-1, x.shape[-1])\n        x = x_skip[mask_flat].unsqueeze(0)\n        \n        x = self.ln_conv(x)\n        \n        \n        residual = x\n        x = x * self.scale_ff_conv + self.bias_ff_conv\n        x = residual + self.ff_conv(x)\n        x = self.ln_ff_conv(x)\n        \n        # Unskip pad #2\n        x_skip[mask_flat] = x[0]\n        x = x_skip.view(bs, slen, nfeats)  \n        \n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-12-03T20:47:03.442302Z","iopub.execute_input":"2023-12-03T20:47:03.442689Z","iopub.status.idle":"2023-12-03T20:47:03.535874Z","shell.execute_reply.started":"2023-12-03T20:47:03.442652Z","shell.execute_reply":"2023-12-03T20:47:03.535039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 1D POSITION ENCODINGS","metadata":{}},{"cell_type":"code","source":"def get_emb(sin_inp):\n    \"\"\"\n    Gets a base embedding for one dimension with sin and cos intertwined\n    \"\"\"\n    emb = torch.stack((sin_inp.sin(), sin_inp.cos()), dim=-1)\n    return torch.flatten(emb, -2, -1)\n\nclass PositionalEncoding1D(nn.Module):\n    def __init__(self, channels):\n        \"\"\"\n        :param channels: The last dimension of the tensor you want to apply pos emb to.\n        \"\"\"\n        super(PositionalEncoding1D, self).__init__()\n        self.org_channels = channels\n        channels = int(np.ceil(channels / 2) * 2)\n        self.channels = channels\n        inv_freq = 1.0 / (10000 ** (torch.arange(0, channels, 2).float() / channels))\n        self.register_buffer(\"inv_freq\", inv_freq)\n        self.cached_penc = None\n\n    def forward(self, tensor):\n        \"\"\"\n        :param tensor: A 3d tensor of size (batch_size, x, ch)\n        :return: Positional Encoding Matrix of size (batch_size, x, ch)\n        \"\"\"\n        if len(tensor.shape) != 3:\n            raise RuntimeError(\"The input tensor has to be 3d!\")\n\n        if self.cached_penc is not None and self.cached_penc.shape == tensor.shape:\n            return self.cached_penc\n\n        self.cached_penc = None\n        batch_size, x, orig_ch = tensor.shape\n        pos_x = torch.arange(x, device=tensor.device).type(self.inv_freq.type())\n        sin_inp_x = torch.einsum(\"i,j->ij\", pos_x, self.inv_freq)\n        emb_x = get_emb(sin_inp_x)\n        emb = torch.zeros((x, self.channels), device=tensor.device).type(tensor.type())\n        emb[:, : self.channels] = emb_x\n\n        self.cached_penc = emb[None, :, :orig_ch].repeat(batch_size, 1, 1)\n        return self.cached_penc\n        ","metadata":{"execution":{"iopub.status.busy":"2023-12-03T20:47:03.536919Z","iopub.execute_input":"2023-12-03T20:47:03.537204Z","iopub.status.idle":"2023-12-03T20:47:03.552297Z","shell.execute_reply.started":"2023-12-03T20:47:03.537181Z","shell.execute_reply":"2023-12-03T20:47:03.551442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## SINUSOIDAL POSITIONAL ENCODING","metadata":{}},{"cell_type":"code","source":"class SinusoidalPositionalEmbedding(nn.Embedding): # for testing\n    \"\"\"This module produces sinusoidal positional embeddings of any length.\"\"\"\n    def __init__(self, num_positions: int, embedding_dim: int, padding_idx: Optional[int] = None) -> None:\n        super().__init__(num_positions, embedding_dim)\n        self.weight = self._init_weight(self.weight)\n\n    @staticmethod\n    def _init_weight(out: nn.Parameter) -> nn.Parameter:\n        \"\"\"\n        Interleaved sine and cosine position embeddings\n        \"\"\"\n        out.requires_grad = False\n        out.detach_()\n        \n        N, D = out.shape\n\n        ## TODO: Create a N x D//2 array of position encodings (argument to the sine/cosine)\n        inds = np.arange(0, D // 2)\n        k = np.arange(N)\n        denom = 1 / np.power(10_000, 2*inds / D)\n        position_enc = np.outer(k, denom)  # Efficiently make N x D//2 array for all positions/dimensions\n        #####\n\n        out[:, 0::2] = torch.FloatTensor(np.sin(position_enc))  # Even indices get sin\n        out[:, 1::2] = torch.FloatTensor(np.cos(position_enc))  # Odd indices get cos\n        return out\n\n    @torch.no_grad()\n    def forward(self, input_ids_shape: torch.Size, past_key_values_length: int = 0) -> torch.Tensor:\n        \"\"\"`input_ids_shape` is expected to be [bsz x seqlen].\"\"\"\n        bsz, seq_len = input_ids_shape[:2]\n        positions = torch.arange(\n            past_key_values_length, past_key_values_length + seq_len, dtype=torch.long, device=self.weight.device\n        )\n        return super().forward(positions)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T20:47:03.553379Z","iopub.execute_input":"2023-12-03T20:47:03.55364Z","iopub.status.idle":"2023-12-03T20:47:03.567321Z","shell.execute_reply.started":"2023-12-03T20:47:03.553617Z","shell.execute_reply":"2023-12-03T20:47:03.566434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## VANILLA TRANSFORMER\n","metadata":{}},{"cell_type":"code","source":"# VANILLA TRANSFORMER MODEL \n\nclass VanillaEncoder(nn.Module):\n    \"\"\"\n    PARAMS: \n        input_dim = input vector dimension \n        encoder_dim = dimension of encoder\n        num layers = num of encoder blocks (?? )\n        num_attention_heads = num of attention heads\n        \n    IN: inputs, input_lengths \n        -**inputs** (batch, time, dim): Tensor containing input vector\n        - **input_lengths** (batch): list of sequence input lengths\n    \n    OUT: outputs, output_lengths\n        - **outputs** (batch, out_channels, time): Tensor produces by encoder.\n        - **output_lengths** (batch): list of sequence output lengths\n    \n    \"\"\"\n    def __init__(\n        self, \n        d_model=512, \n        d_input = 208, \n        d_output = 62, \n        n_head=4, \n        n_layers=3,  \n        dropout_p: float = 0.1\n    ):\n        super().__init__()\n        self.d_model = d_model\n        self.n_head = n_head\n        self.n_layers = n_layers\n        self.dropout_p = dropout_p\n        self.input_transform = nn.Linear(d_input, d_model) # transform from 80 input features to 512 input features\n        self.d_input = d_input\n            \n        # Input embedding layers\n        # self.position_embed = RelPositionalEncoding(d_model)\n        self.position_embed = SinusoidalPositionalEmbedding(num_positions = 400, embedding_dim = self.d_model)\n        # self.position_embed = PositionalEncoding1D(d_model)\n\n            \n        # init the transformer encoder layer\n        self.encoder_layer = nn.TransformerEncoderLayer(\n            # normally use between 2-4x model for feedforward dim\n            d_model=self.d_model,\n            nhead=self.n_head,\n            dim_feedforward=2*self.d_model,\n            dropout=self.dropout_p,\n            activation=\"gelu\", \n            batch_first=True, \n            norm_first=True\n        )\n        \n        #Encoder\n        self.encoder = nn.TransformerEncoder(self.encoder_layer, num_layers = self.n_layers)      \n        \n        # class token for the encoder\n        self.cls = torch.nn.Parameter(torch.randn(1, 1, self.d_input))\n        \n        self.layer_norm = nn.LayerNorm(self.d_model)\n        self.classifier = nn.Linear(self.d_model, 1)\n        \n        \n    def forward(self, x, mask):\n        # replace first element of the sequence with the cls token\n        # print(\"before token replce: \",x.size())\n        \n        x[:, 0, :] = self.cls \n        \n        # print(\"after token replce: \", x.size())\n        \n        # pass x through input linear layer\n        x = self.input_transform(x)\n        \n       #  print(\"after input transform: \", x.size())\n        \n        #get position embedding from the input shape, add tp transformed inputs\n        # pe = self.position_embed(x)\n        pe = self.position_embed(x.shape)  # Sinusoidal \n        \n        # print(\"pe: \", pe.size()) \n        #add positional embeds to x\n        x = pe + x\n        # x = x + pe.to(x.device)\n        \n        # print(\"pe plus x: \", x.size()) \n        \n        # pass inputs through encoder\n        x = self.encoder(x)\n        \n        # print(\"x encoder result: \", x.size()) \n        \n        x = self.layer_norm(x)\n        # print(\"x layer norm result: \", x.size()) \n#         x = x[:, 0, :].reshape(-1, self.d_model)  # Predict only on cls token\n#         print(\"x reshape result: \", x.size()) \n#         x = self.classifier(x)\n#         print(\"x classifier: \", x.size()) \n#         x = torch.flatten(x)\n        \n        # print(\"x end encoder result: \", x.size()) \n        return x\n   ","metadata":{"execution":{"iopub.status.busy":"2023-12-03T20:47:03.568608Z","iopub.execute_input":"2023-12-03T20:47:03.568911Z","iopub.status.idle":"2023-12-03T20:47:03.582357Z","shell.execute_reply.started":"2023-12-03T20:47:03.568866Z","shell.execute_reply":"2023-12-03T20:47:03.581521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FeatureExtractor(nn.Module):\n    def __init__(self,\n                 n_landmarks,out_dim, conv_ch = 3):\n        super().__init__()   \n\n        self.in_channels = in_channels = 32 * math.ceil(n_landmarks / 2)\n        self.stem_linear = nn.Linear(in_channels,out_dim,bias=False)\n        self.stem_bn = nn.BatchNorm1d(out_dim, momentum=0.95)\n        self.conv_stem = nn.Conv2d(conv_ch, 32, kernel_size=(3, 3), stride=(1, 2), padding=(1, 1), bias=False)\n        self.bn_conv = BatchNormAct2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True,act_layer = nn.SiLU,drop_layer=None)\n        \n    def forward(self, data, mask):\n\n        xc = data.permute(0,3,1,2)\n        xc = self.conv_stem(xc)\n        xc = self.bn_conv(xc)\n        xc = xc.permute(0,2,3,1)\n        xc = xc.reshape(*data.shape[:2], -1)\n        \n        m = mask.to(torch.bool)  \n        x = self.stem_linear(xc)\n        \n        # Batchnorm without pads\n        bs,slen,nfeat = x.shape\n        x = x.view(-1, nfeat)\n        x_bn = x[mask.view(-1)==1].unsqueeze(0)\n        x_bn = self.stem_bn(x_bn.permute(0,2,1)).permute(0,2,1)\n        x[mask.view(-1)==1] = x_bn[0]\n        x = x.view(bs,slen,nfeat)\n        # Padding mask\n        x = x.masked_fill(~mask.bool().unsqueeze(-1), 0.0)\n        \n        return x\n\n    \nclass Decoder(nn.Module):\n    def __init__(self, decoder_config):\n        super(Decoder, self).__init__()\n        \n        self.config = decoder_config\n        self.decoder = Speech2TextDecoder(decoder_config) \n        self.lm_head = nn.Linear(decoder_config.d_model, decoder_config.vocab_size, bias=False)\n        \n        self.decoder_start_token_id = decoder_config.decoder_start_token_id\n        self.decoder_pad_token_id = decoder_config.pad_token_id #used for early stopping\n        self.decoder_end_token_id= decoder_config.eos_token_id\n        \n    def forward(self,x, labels=None, attention_mask = None, encoder_attention_mask = None):\n        \n        if labels is not None:\n            decoder_input_ids = shift_tokens_right(labels, self.config.pad_token_id, self.config.decoder_start_token_id)\n            \n        decoder_outputs = self.decoder(input_ids=decoder_input_ids,\n                                       encoder_hidden_states=x, \n                                       attention_mask = attention_mask,\n                                       encoder_attention_mask = encoder_attention_mask)\n        lm_logits = self.lm_head(decoder_outputs.last_hidden_state)\n        return lm_logits\n            \n    def generate(self, x, max_new_tokens=33, encoder_attention_mask=None):\n\n        decoder_input_ids = torch.ones((x.shape[0], 1), device=x.device, dtype=torch.long).fill_(self.decoder_start_token_id)\n        for i in range(max_new_tokens-1):  \n            decoder_outputs = self.decoder(input_ids=decoder_input_ids,encoder_hidden_states=x, encoder_attention_mask=encoder_attention_mask)\n            logits = self.lm_head(decoder_outputs.last_hidden_state)\n            decoder_input_ids = torch.cat([decoder_input_ids,logits.argmax(2)[:,-1:]],dim=1)\n\n            if torch.all((decoder_input_ids==self.decoder_end_token_id).sum(-1) > 0):\n                break\n                \n        return decoder_input_ids\n    \ndef count_parameters(model):\n    return sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nclass Net(nn.Module):\n\n    def __init__(self):\n        super(Net,self).__init__()\n\n        dim=208\n        n_handlandmarks = 21\n        n_poselandmarks = 5\n        n_landmarks = 2*n_handlandmarks+2*n_poselandmarks\n\n        d_cfg = Speech2TextConfig.from_pretrained(\"facebook/s2t-small-librispeech-asr\")\n        d_cfg.encoder_layers = 0\n        d_cfg.decoder_layers = 2\n        d_cfg.d_model = dim\n        d_cfg.max_target_positions = 1024 #?\n        d_cfg.num_hidden_layers = 1\n        d_cfg.vocab_size = 63\n        d_cfg.bos_token_id = 60\n        d_cfg.eos_token_id = 61\n        d_cfg.decoder_start_token_id = 60\n        d_cfg.pad_token_id = 59\n        d_cfg.num_conv_layers = 0\n        d_cfg.conv_kernel_sizes = []\n        d_cfg.max_length = dim\n        d_cfg.input_feat_per_channel = dim\n        d_cfg.num_beams = 1\n        d_cfg.attention_dropout = 0.2\n        d_cfg.decoder_ffn_dim = 512\n        d_cfg.init_std = 0.02\n        \n        self.feature_extractor = FeatureExtractor(n_landmarks=n_landmarks,out_dim=dim)\n        self.feature_extractor_lhand = FeatureExtractor(n_handlandmarks,out_dim=dim//4)\n        self.feature_extractor_rhand = FeatureExtractor(n_handlandmarks,out_dim=dim//4)\n        self.feature_extractor_lpose = FeatureExtractor(n_poselandmarks,out_dim=dim//4)\n        self.feature_extractor_rpose = FeatureExtractor(n_poselandmarks,out_dim=dim//4)\n        self.training = True\n#         self.encoder = SqueezeformerEncoder(\n#                       input_dim=dim,\n#                       encoder_dim=dim,\n#                       num_layers=5,\n#                       num_attention_heads= 4,\n#                       feed_forward_expansion_factor=1,\n#                       conv_expansion_factor= 2,\n#                       input_dropout_p=0.1,\n#                       feed_forward_dropout_p= 0.1,\n#                       attention_dropout_p= 0.1,\n#                       conv_dropout_p= 0.1,\n#                       conv_kernel_size= 51,)\n        self.encoder = VanillaEncoder(d_model = dim)\n        self.decoder = Decoder(d_cfg)\n        self.loss_fn = nn.CrossEntropyLoss() #done\n        self.max_phrase = MAX_PHRASE\n        print('n_params:',count_parameters(self))\n\n    def forward(self, batch):\n        #Concat,'rhand','lhand','lpose','rpose'\n        x = torch.cat([batch['left_hand'],batch['right_hand'],batch['left_pose'],batch['right_pose']],dim=-2)\n        labels = batch['phrase_tokens']\n        mask = batch['data_mask'].long()\n        label_mask = batch['phrase_mask']    \n\n        #maybe normalize\n        x_lhand = self.feature_extractor_lhand(batch['left_hand'].clone(), mask)\n        x_lpose = self.feature_extractor_lpose(batch['left_pose'].clone(), mask)\n        x_rhand = self.feature_extractor_rhand(batch['right_hand'].clone(), mask)\n        x_rpose = self.feature_extractor_rpose(batch['right_pose'].clone(), mask)\n        \n        x1 = torch.cat([x_lhand,x_rhand,x_lpose,x_rpose],dim=-1)\n        x = self.feature_extractor(x, mask)\n        x = x + x1\n        x = self.encoder(x, mask)\n        decoder_labels = labels.clone()        \n        \n        #??\n        # if self.training:\n        #     m = torch.rand(labels.shape) < self.decoder_mask_aug\n        #     decoder_labels[m] = 62\n        \n        logits = self.decoder(x,\n                            labels=decoder_labels, \n                            encoder_attention_mask=mask.long()\n                            )\n        \n        loss = self.loss_fn(logits.view(-1, self.decoder.config.vocab_size), labels.view(-1))\n        output = {'loss':loss}\n        \n        with torch.no_grad():\n            generated_ids_padded = torch.ones((x.shape[0],self.max_phrase), dtype=torch.long, device=x.device) * 59\n\n            generated_ids = self.decoder.generate(x,max_new_tokens=self.max_phrase + 1, encoder_attention_mask=mask.long())\n\n\n            cutoffs = (generated_ids==self.decoder.decoder_end_token_id).float().argmax(1).clamp(0,self.max_phrase)\n            for i, c in enumerate(cutoffs):\n                generated_ids_padded[i,:c] = generated_ids[i,:c]\n            output['generated_ids'] = generated_ids_padded\n            \n        return output","metadata":{"_uuid":"6bf46d90-ae3f-4340-a7a5-2f79002e0301","_cell_guid":"a7e3d045-817e-4899-a03f-2aaa5e4b48ac","execution":{"iopub.status.busy":"2023-12-03T20:47:03.583672Z","iopub.execute_input":"2023-12-03T20:47:03.58391Z","iopub.status.idle":"2023-12-03T20:47:03.61606Z","shell.execute_reply.started":"2023-12-03T20:47:03.583889Z","shell.execute_reply":"2023-12-03T20:47:03.615053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{"_uuid":"08c3c0ad-3b5c-4bf8-8254-b49188703874","_cell_guid":"804eabfa-0d2e-4e41-a541-dd47387bde31","trusted":true}},{"cell_type":"code","source":"def train(model, train_loader, optimizer,scheduler,epoch):\n    total_loss = 0\n    all_predictions = []\n    all_targets = []\n    loss_history = []\n    total_ld_score = 0\n    \n    model.training = True\n    # set model to training mode\n    model.train()  \n    \n    print(f\"Epoch {epoch + 1} started\")\n\n    for i, data in enumerate(tqdm(train_loader)):\n        optimizer.zero_grad()\n        batch = {key: data[key].to(DEVICE) for key in data if key != 'phrase'}\n        output = model(batch)\n        loss = output['loss']\n        preds = post_process_pipeline(output)\n        scores = calc_metrics(preds,data)\n        loss.backward()\n        \n        optimizer.step()\n        optimizer.zero_grad()\n        scheduler.step()\n\n        # track some values to compute statistics\n        total_loss += loss.item()\n        total_ld_score += scores['score']\n\n    final_loss = total_loss / len(train_loader)\n    final_ld_score = total_ld_score / len(train_loader)\n    # print average loss and accuracy\n    print(f\"learning Rate = {optimizer.param_groups[0]['lr']}. average train loss = {final_loss:.2f}. average lev distance = {final_ld_score:.2f}\")\n    return final_ld_score, final_loss\n\ndef post_process_pipeline(val_data):\n    \n    generated_ids = val_data['generated_ids'].cpu()\n    with open('/kaggle/input/asl-fingerspelling/character_to_prediction_index.json', \"r\") as f:\n            char_to_num = json.load(f)\n    \n    rev_character_map = {j:i for i,j in char_to_num.items()}\n    phrase_preds = [\"\".join([rev_character_map.get(s, \"\") for s in generated_id]) for generated_id in generated_ids.numpy()]\n    \n    \n    return {'phrase_preds':phrase_preds}\n\n\ndef validation(model, val_loader, epoch):\n    total_loss = 0\n    all_predictions = []\n    all_targets = []\n    total_ld_score = 0\n\n    model.training = False\n    # set model to evaluation mode\n    model.eval()  \n    for i, data in enumerate(tqdm(val_loader)):\n        with torch.no_grad():\n            batch = {key: data[key].to(DEVICE) for key in data if key != 'phrase'}\n            outputs = model(batch)\n            loss = outputs['loss']\n            preds = post_process_pipeline(outputs)\n            scores = calc_metrics(preds,data)\n            total_ld_score += scores['score']\n            # Track some values to compute statistics\n            total_loss += loss.item()\n\n    final_loss = total_loss / len(val_loader)\n    final_ld_score = total_ld_score / len(val_loader)\n    # Print average loss and accuracy\n    print(f\"Epoch {epoch + 1} done. average validation loss = {final_loss:.2f}, lev score = {final_ld_score}\")\n    return final_ld_score, final_loss\n\ndef test(model, test_loader, loss_fn):\n    total_loss = 0\n    all_predictions = []\n    all_targets = []\n\n    model.training = False\n    # set model to evaluation mode\n    model.eval()  \n    for i, data in enumerate(tqdm(test_loader)):\n        with torch.no_grad():\n            batch = {key: data[key].to(DEVICE) for key in data if key != 'phrase'}\n            outputs = model(batch)\n            loss = outputs['loss']\n            preds = post_process_pipeline(outputs)\n            scores = calc_metrics(preds,data)\n            total_ld_score += scores['score']\n            # Track some values to compute statistics\n            total_loss += loss.item()\n\n    final_loss = total_loss / len(test_loader)\n    final_ld_score = total_ld_score / len(test_loader)\n    # Print average loss and accuracy\n    print(f\"Epoch {epoch + 1} done. average test loss = {final_loss:.2f}, lev score = {final_ld_score}\")\n    return final_ld_score, final_loss","metadata":{"_uuid":"0f50b44b-799e-4547-8468-d57899f8ac9f","_cell_guid":"077096e3-77f1-4c75-a159-dc5a54a162a9","execution":{"iopub.status.busy":"2023-12-03T20:47:03.617523Z","iopub.execute_input":"2023-12-03T20:47:03.617763Z","iopub.status.idle":"2023-12-03T20:47:03.635194Z","shell.execute_reply.started":"2023-12-03T20:47:03.617741Z","shell.execute_reply":"2023-12-03T20:47:03.634375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(\n                data_loaders: tuple = None,\n                learning_rate: float = 4.5e-3, \n                weight_decay: float = 0.08, \n                features: dict = None,\n                warmup_steps=2,\n                training_steps = 14\n#                 warmup_steps=1,\n#                 training_steps = 1\n               ):\n    train_accs = []\n    train_losses = []\n    val_accs = []\n    val_losses = []\n    \n    if not features:\n        features = FEATURES\n        \n    if not data_loaders:\n        train_loader, val_loader, test_loader = get_data_loaders(features, VAL_SPLIT, TEST_SPLIT, BATCH_SIZE)\n    else:\n        train_loader, val_loader, test_loader = data_loaders\n\n    model = Net().to(DEVICE)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n    scheduler = transformers.get_cosine_schedule_with_warmup(\n        optimizer, num_warmup_steps=warmup_steps*len(train_loader),num_training_steps=training_steps*len(train_loader),num_cycles=0.5\n    )\n\n    max_epochs = warmup_steps + training_steps\n    best_ld_score = 0\n    dip_count = 0\n    num_epochs = 0\n    optimizer.zero_grad()\n    for e in range(max_epochs):\n        gc.collect()\n        train_ld_score, train_loss = train(model, train_loader, optimizer,scheduler,e)\n        val_ld_score, val_loss = validation(model, val_loader, e)\n        # early stopping based on train acc\n        if e%2 == 0:\n            \n            if val_ld_score >= best_ld_score:\n                best_ld_score = val_ld_score\n                dip_count = 0\n            else:\n                dip_count +=1\n\n        if dip_count >1:\n            break\n        train_accs.append(train_ld_score)\n        train_losses.append(train_loss)\n        val_accs.append(val_ld_score)\n        val_losses.append(val_loss)\n        \n        \n    # temporarily save model\n    torch.save(model.state_dict(), \"temp_model.pth\")\n\n    val_ld_score, val_loss = validation(model, val_loader, max_epochs-1)\n    print(f'Final Levenshtein Distance Score: {val_ld_score}')\n    torch.save(model.state_dict(), \"lev_model.pth\")\n    \n    return train_accs,train_losses,val_accs,val_losses\n","metadata":{"_uuid":"3a039f86-1103-4f01-a147-c1b949fe4dbb","_cell_guid":"078cccf9-b674-491f-933e-411c410cb123","execution":{"iopub.status.busy":"2023-12-03T20:47:03.636776Z","iopub.execute_input":"2023-12-03T20:47:03.63707Z","iopub.status.idle":"2023-12-03T20:47:03.649967Z","shell.execute_reply.started":"2023-12-03T20:47:03.637045Z","shell.execute_reply":"2023-12-03T20:47:03.649145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nfrom torchaudio.functional import edit_distance\n\ndef get_score(phrase_gt, phrase_preds):\n    \n#     score = edit_distance(phrase_gt, phrase_preds)\n    N = np.array([len(p) for p in phrase_gt])\n    D = np.array([edit_distance(p1,p2) for p1,p2 in zip(phrase_gt,phrase_preds)])\n    score = (N.sum() - D.sum()) / N.sum()\n    if score < 0:\n        score = -1 * score\n    return score\n\ndef calc_metrics(pp_out, val_df):\n    \n    \n    phrase_gt = val_df['phrase']\n    phrase_preds = pp_out['phrase_preds']\n    # print(phrase_gt,phrase_preds)\n    \n    score = get_score(phrase_gt, phrase_preds)\n    \n    \n    return {'score':score}","metadata":{"execution":{"iopub.status.busy":"2023-12-03T20:47:03.650925Z","iopub.execute_input":"2023-12-03T20:47:03.651186Z","iopub.status.idle":"2023-12-03T20:47:03.663692Z","shell.execute_reply.started":"2023-12-03T20:47:03.651165Z","shell.execute_reply":"2023-12-03T20:47:03.662779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(DEVICE)\ntrain_acc,train_loss,val_acc,val_loss = train_model()","metadata":{"_uuid":"82fef586-fc03-49b7-8b0f-f5877b8bdf9c","_cell_guid":"da44875e-bea6-4819-9266-d15f7703dcc5","execution":{"iopub.status.busy":"2023-12-03T20:47:03.66482Z","iopub.execute_input":"2023-12-03T20:47:03.66522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# x_axis = range(16)\n# x_axis_val = range(0,16,2)\n# plt.plot(x_axis,train_acc, color='r', label='Train', marker = 'o')\n# plt.plot(x_axis_val, val_acc, color='g', label='Val', marker = 'o')\n# plt.xlabel(\"Epochs\")\n# plt.ylabel('Levenshtein Accuracy')\n# plt.title(\"Train and Validation Accuracy over Training Epochs\") \n# plt.legend()\n# plt.show() \n\n# plt.plot(x_axis,train_loss, color='r', label='Train', marker = 'o')\n# plt.plot(x_axis_val, val_loss, color='g', label='Val', marker = 'o')\n# plt.xlabel(\"Epochs\")\n# plt.ylabel('Cross Entropy Loss')\n# plt.title(\"Train and Validation Loss over Training Epochs\") \n# plt.legend()\n# plt.show() \n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torch.save(model.state_dict(), \"lev_model.pth\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}