{"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":"## Parkinson FoG Prediction TF Transformer Model\nModel adapted from https://keras.io/examples/timeseries/timeseries_transformer_classification/","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport numpy as np\nfrom numpy.random import default_rng\nimport pandas as pd\nfrom tqdm.auto import tqdm\nfrom glob import glob\nfrom os.path import basename, dirname, join, exists\nfrom time import perf_counter\nfrom collections import defaultdict as dd\nfrom functools import partial\n\nfrom sklearn.model_selection import train_test_split, StratifiedKFold, StratifiedGroupKFold\nfrom sklearn.metrics import average_precision_score\nfrom sklearn.preprocessing import StandardScaler as Scaler\nfrom scipy.special import expit\n\nimport tensorflow as tf\nfrom tensorflow.keras import layers\nprint(f\"TF version: {tf.__version__}\")\nAUTO = tf.data.experimental.AUTOTUNE","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-21T13:31:40.057382Z","iopub.execute_input":"2023-04-21T13:31:40.057726Z","iopub.status.idle":"2023-04-21T13:31:49.386909Z","shell.execute_reply.started":"2023-04-21T13:31:40.057693Z","shell.execute_reply":"2023-04-21T13:31:49.385819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Constants\n\nBASE_DIR = \"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction\"\nTRAIN_DIR = join(BASE_DIR, \"train\")\nTEST_DIR = join(BASE_DIR, \"test\")\n\nIS_PUBLIC = len(glob(join(TEST_DIR, \"*/*.csv\")))==2","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:31:49.391855Z","iopub.execute_input":"2023-04-21T13:31:49.394586Z","iopub.status.idle":"2023-04-21T13:31:49.412573Z","shell.execute_reply.started":"2023-04-21T13:31:49.394542Z","shell.execute_reply":"2023-04-21T13:31:49.411598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    train_sub_dirs = [\n        join(TRAIN_DIR, \"defog\"),\n        join(TRAIN_DIR, \"tdcsfog\")\n    ]\n    \n    metadata_paths = [\n        join(BASE_DIR, \"defog_metadata.csv\"),\n        join(BASE_DIR, \"tdcsfog_metadata.csv\")\n    ]\n    \n    splits = 5 # 10\n    fold = 0\n    \n    train_limit = 100 # 5_000_000\n\n    batch_size = 512 # 1024\n    window_size = 64*4 # 32 # \n    window_future = 16*4 # 8 # \n    window_past = window_size - window_future # Includes current value\n    \n    wx = 8\n    \n    model_dropout = 0.2\n    model_hidden = 512\n    model_nblocks = 3\n    \n    lr = 0.00015\n    num_epochs = 8 # 10 # 9 # 1 if IS_PUBLIC else \n    \n    feature_list = ['Time_frac', 'AccV', 'AccML', 'AccAP']\n    label_list = ['StartHesitation', 'Turn', 'Walking']\n    \n    n_features = len(feature_list)\n    n_labels = len(label_list)\n    \n    norm = True\n    norm_list = ['AccV', 'AccML', 'AccAP']    \n    \ncfg = Config()","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:35:55.562546Z","iopub.execute_input":"2023-04-21T13:35:55.563534Z","iopub.status.idle":"2023-04-21T13:35:55.573776Z","shell.execute_reply.started":"2023-04-21T13:35:55.563465Z","shell.execute_reply":"2023-04-21T13:35:55.572571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def format_time(seconds):\n    if seconds > 3600:\n        return f\"{seconds/3600:.2f} hrs\"\n    if seconds > 60:\n        return f\"{seconds/60:.2f} mins\"\n    return f\"{seconds:.2f} secs\"","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:31:49.435992Z","iopub.execute_input":"2023-04-21T13:31:49.436613Z","iopub.status.idle":"2023-04-21T13:31:49.447129Z","shell.execute_reply.started":"2023-04-21T13:31:49.436580Z","shell.execute_reply":"2023-04-21T13:31:49.445535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Stratified Group K Fold","metadata":{}},{"cell_type":"code","source":"def reduce_categories(categories, cutoff):\n    cum_sum = 0\n    to_change = []\n    last = -1\n    categories = np.array(categories)\n    for idx, count in reversed(list(pd.Series(categories).value_counts().items())):\n        if last > -1:\n            to_change.append(last)\n        last = idx\n        cum_sum += count\n        if cum_sum >= cutoff:\n            break\n    for idx in to_change:\n        categories[categories==idx] = last\n    return categories","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:31:49.454523Z","iopub.execute_input":"2023-04-21T13:31:49.455212Z","iopub.status.idle":"2023-04-21T13:31:49.463438Z","shell.execute_reply.started":"2023-04-21T13:31:49.455171Z","shell.execute_reply":"2023-04-21T13:31:49.461999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create Mapping between Id and Subject\nid2sub_df = pd.concat([\n    pd.read_csv(f, usecols=['Id', 'Subject']).assign(Module=basename(f).split('_')[0]) for f in cfg.metadata_paths\n]).astype(\"category\").set_index(\"Id\")\nprint(f\"id2sub_df length: {len(id2sub_df)}, unique Ids: {id2sub_df.index.nunique()}, unique Subjects: {id2sub_df.Subject.nunique()}\")","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:31:49.465586Z","iopub.execute_input":"2023-04-21T13:31:49.466523Z","iopub.status.idle":"2023-04-21T13:31:49.525936Z","shell.execute_reply.started":"2023-04-21T13:31:49.466485Z","shell.execute_reply":"2023-04-21T13:31:49.525082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Read csv files and add metadata (Id, Subject, Event)\ndef reader(filepath, usecols, getid=False, getsub=False, getevent=False, dtype=None, exclude=['notype']):\n    fog_type = basename(dirname(filepath))\n    if fog_type in exclude:\n        return None\n    df = pd.read_csv(filepath, index_col=\"Time\", usecols=usecols, dtype=dtype)\n    if getid:\n        df['Id'] = basename(filepath).split('.')[0] + '_' + df.index.astype(str)\n    if getsub:\n        df['Subject'] = id2sub_df.loc[basename(filepath).split('.')[0], 'Subject']\n    if getevent:\n        df['Event'] = np.select(\n            [df[col].astype(bool) for col in cfg.label_list], \n            np.arange(1,cfg.n_labels+1), default=0\n        ).astype('int8')\n    return df","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:31:49.529994Z","iopub.execute_input":"2023-04-21T13:31:49.532224Z","iopub.status.idle":"2023-04-21T13:31:49.542665Z","shell.execute_reply.started":"2023-04-21T13:31:49.532172Z","shell.execute_reply":"2023-04-21T13:31:49.541711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create common train Dataframe\ntrain_paths = glob(join(TRAIN_DIR, '*/*.csv'))\ndtype = {col:'int8' for col in cfg.label_list}\ndtype['Time'] = 'int32'\nusecols = ['Time', *cfg.label_list]\n\ntrain_reader = partial(reader, usecols=usecols, dtype=dtype, getsub=True, getevent=True)\ntrain_df = pd.concat([train_reader(f) for f in tqdm(train_paths)]).reset_index(drop=True)\ntrain_df.Subject = train_df.Subject.astype('category')\ndisplay(train_df.Event.value_counts().to_frame().style.background_gradient())","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:31:49.547294Z","iopub.execute_input":"2023-04-21T13:31:49.549970Z","iopub.status.idle":"2023-04-21T13:32:29.670731Z","shell.execute_reply.started":"2023-04-21T13:31:49.549932Z","shell.execute_reply":"2023-04-21T13:32:29.669607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save paths for each Stratified Group Fold for defog and tdcsfog separately\nsgkf = StratifiedGroupKFold(n_splits=cfg.splits, random_state=42, shuffle=True)\nfold_train_fpaths, fold_valid_fpaths = {'defog': [], 'tdcsfog':[]}, {'defog': [], 'tdcsfog':[]}\ndf_paths = {'defog':glob(join(cfg.train_sub_dirs[0],'*.csv')), 'tdcsfog': glob(join(cfg.train_sub_dirs[1],'*.csv'))}\nfor module, paths in df_paths.items():\n    print(f\"{module}:\")\n    sub_train_df = train_df[train_df.Subject.isin(id2sub_df.loc[id2sub_df.Module==module, 'Subject'])].reset_index(drop=True)\n    for i, (train_index, test_index) in enumerate(sgkf.split(sub_train_df.index, sub_train_df.Event, groups=sub_train_df.Subject)):\n        print(f\"\\tFold {i}:\", end=\" \")\n        train_subs = sub_train_df.loc[train_index, 'Subject'].unique()\n        test_subs = sub_train_df.loc[test_index, 'Subject'].unique()\n        print(f\"Subjects->train:{len(train_subs)}|test:{len(test_subs)}\")\n        train_ids = set(id2sub_df[id2sub_df.Subject.isin(train_subs)].index)\n        test_ids = set(id2sub_df[id2sub_df.Subject.isin(test_subs)].index)\n        fold_train_fpaths[module].append([f for f in paths if basename(f).split('.')[0] in train_ids])\n        fold_valid_fpaths[module].append([f for f in paths if basename(f).split('.')[0] in test_ids])\n        del train_subs, test_subs, train_ids, test_ids\n        gc.collect()\n    del sub_train_df\ndel train_df\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:32:29.672601Z","iopub.execute_input":"2023-04-21T13:32:29.672966Z","iopub.status.idle":"2023-04-21T13:33:20.590339Z","shell.execute_reply.started":"2023-04-21T13:32:29.672927Z","shell.execute_reply":"2023-04-21T13:33:20.587555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fold_train_fpaths, fold_valid_fpaths = dd(list), dd(list)\n\n# for train_sub_dir, metadata_path in zip(cfg.train_sub_dirs, cfg.metadata_paths):\n#     fog_type = basename(train_sub_dir).replace('fog', '')\n#     print(f\"Evaluating {fog_type}fog folder:\")\n#     # Analysis of positive instances in each fold of our CV folds\n\n#     n_sum = [list() for _ in range(len(cfg.label_list))]\n#     count = []\n#     rejected_ids = []\n#     categories = []\n    \n#     metadata = pd.read_csv(metadata_path)\n\n#     for fid in tqdm(metadata['Id']):\n#         fpath = join(train_sub_dir, f\"{fid}.csv\")\n#         if not exists(fpath):\n#             rejected_ids.append(fid)\n#             continue\n#         df = pd.read_csv(fpath)\n#         category = 0\n#         for i, col in enumerate(cfg.label_list):\n#             n_sum[i].append(np.sum(df[col]))\n#             category = (category<<1) + int(n_sum[i][-1]>0)\n#         count.append(len(df))\n#         categories.append(category)\n        \n#     categories = reduce_categories(categories, cfg.splits)\n#     # display(pd.Series(categories).value_counts())\n\n#     if rejected_ids:\n#         print(f\"\\tRemoving {len(rejected_ids)} Ids from metadata\")\n#         metadata = metadata[~metadata['Id'].isin(rejected_ids)].reset_index().copy()\n\n#     all_pos_files = (np.array(n_sum) > 0).all(axis=0).sum()\n#     print(f\"\\t{all_pos_files} files have positive values in all 3 classes\")\n\n#     sum_cols = []\n#     for i, sum_lst in enumerate(n_sum, 1):\n#         sum_cols.append(f'n{i}_sum')\n#         metadata[sum_cols[-1]] = sum_lst # TODO: Change naming\n#     metadata['count'] = count\n\n#     sgkf = StratifiedGroupKFold(n_splits=cfg.splits, random_state=42, shuffle=True)\n#     for i, (train_index, valid_index) in enumerate(sgkf.split(X=metadata['Id'], y=categories, groups=metadata['Subject'])):\n#         print(f\"\\tFold = {i}\")\n\n#         print(f\"\\t Length of Train = {len(train_index)}, Length of Valid = {len(valid_index)}\")\n\n#         col_sums = metadata.loc[train_index, sum_cols].sum()\n#         total = col_sums.sum()\n#         col_sums = [f\"{s:,}({s/total:.2f})\" for s in col_sums]\n#         print(f\"\\t Train classes:\", \", \".join(col_sums))\n\n#         col_sums = metadata.loc[valid_index, sum_cols].sum()\n#         total = col_sums.sum()\n#         col_sums = [f\"{s:,}({s/total:.2f})\" for s in col_sums]\n#         print(f\"\\t Valid classes:\", \", \".join(col_sums))\n        \n#         train_ids = metadata.loc[train_index, 'Id']\n#         valid_ids = metadata.loc[valid_index, 'Id']\n\n#         fold_train_fpaths[i].extend([join(train_sub_dir, f\"{fid}.csv\") for fid in train_ids])\n#         fold_valid_fpaths[i].extend([join(train_sub_dir, f\"{fid}.csv\") for fid in valid_ids])\n\n# for i in range(cfg.splits):\n#     print(f\"Fold {i}: {len(fold_train_fpaths[i])}, {len(fold_valid_fpaths[i])}\")","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:33:20.594969Z","iopub.execute_input":"2023-04-21T13:33:20.595503Z","iopub.status.idle":"2023-04-21T13:33:20.601386Z","shell.execute_reply.started":"2023-04-21T13:33:20.595469Z","shell.execute_reply":"2023-04-21T13:33:20.600235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Calibrate StandardScaler based on core feature columns split by 'AccV'","metadata":{}},{"cell_type":"code","source":"if cfg.norm:\n    # TODO: Try also folderwise\n    tdcsfog_fit_values, other_fit_values = [], []\n    for df_path in tqdm(glob(join(TRAIN_DIR, \"*/*.csv\"))):\n        df = pd.read_csv(df_path, index_col=\"Time\", usecols=['Time', *cfg.norm_list])\n        if basename(dirname(df_path))==\"tdcsfog\":\n            tdcsfog_fit_values.append(df[cfg.norm_list])\n        else:\n            other_fit_values.append(df[cfg.norm_list])\n\n    print(f\"Fitting {len(tdcsfog_fit_values)} tdcsfog values and {len(other_fit_values)} other values.\")\n    tdcsfog_scaler = Scaler().fit(pd.concat(tdcsfog_fit_values))\n    other_scaler = Scaler().fit(pd.concat(other_fit_values))","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:33:20.602910Z","iopub.execute_input":"2023-04-21T13:33:20.603539Z","iopub.status.idle":"2023-04-21T13:33:54.272766Z","shell.execute_reply.started":"2023-04-21T13:33:20.603499Z","shell.execute_reply":"2023-04-21T13:33:54.271592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class FOGSequence(tf.keras.utils.Sequence):\n\n    def __init__(self, df_paths, cfg=cfg, split=\"train\", trim=False, verbose=True):\n        _time = perf_counter()\n        \n        self.rng = default_rng(42)\n        self.cfg = cfg\n        self.split = split\n        self.trim = trim\n        \n        self.past_pad = self.cfg.wx*(self.cfg.window_past-1)\n        self.future_pad = self.cfg.wx*self.cfg.window_future\n        \n        if self.split == \"test\":\n            self.Ids = []\n        _values = [self._read(f) for f in df_paths]\n        \n        self.end_indices = []\n        self.shapes = []\n        self.mapping = []\n        _length = 0\n        for _value in _values:\n            _shape = _value.shape[0]\n            self.mapping.extend(range(_length+self.past_pad, _length+_shape-self.future_pad))\n            self.shapes.append(_shape)\n            _length += _shape\n            self.end_indices.append(_length)\n            \n        self.values = np.concatenate(_values, axis=0)\n        self.mapping = np.array(self.mapping)\n        if self.split != \"test\":\n            _valid_pos = self.values[self.mapping,self.valid_position] > 0\n            _task_pos = self.values[self.mapping,self.task_position] > 0\n            self.mapping = self.mapping[_valid_pos&_task_pos]\n        self.length = self.mapping.shape[0]\n        \n        if verbose:\n            print(f\"Valid Dataset of size {self.length:,} initialized in {perf_counter() - _time:.3f} secs!\")\n        gc.collect()\n        \n    def _read(self, path):\n        _is_tdcs = basename(dirname(path)).startswith('tdcs')\n        df = pd.read_csv(path)\n        if self.cfg.norm:\n            if basename(dirname(path))==\"tdcsfog\":\n                df[cfg.norm_list] = tdcsfog_scaler.transform(df[cfg.norm_list])\n            else:\n                df[cfg.norm_list] = other_scaler.transform(df[cfg.norm_list])\n        df['Time_frac'] = (df.Time/df.Time.max()).astype('float16')\n        \n        if self.split == \"test\":\n            _ids = basename(path).split('.')[0] + '_' + df.Time.astype(str)\n            self.Ids.extend(_ids.tolist())\n            return self._df_to_array(df, self.cfg.feature_list)\n        \n        _cols = [*self.cfg.feature_list, *self.cfg.label_list, 'Valid', 'Task']\n        self.valid_position = self.cfg.n_features + self.cfg.n_labels\n        self.task_position = self.valid_position + 1\n        \n        if _is_tdcs:\n            # Fill Valid and Task columns for tdcsfog\n            df['Valid'] = 1\n            df['Task'] = 1\n            \n        return self._df_to_array(df, _cols)\n        \n    def _df_to_array(self, df, cols):\n        _values = df[cols].values.astype(np.float16)\n        return np.pad(_values, ((self.past_pad, self.future_pad),(0,0)), 'edge')\n        \n\n    def __len__(self):\n        if self.trim and self.split != \"test\":\n            return self.cfg.train_limit\n        return int(np.ceil(self.length / self.cfg.batch_size))\n\n    def __getitem__(self, idx):\n        \n        if self.split == \"train\":\n            _idxs = self.rng.choice(self.mapping, size=self.cfg.batch_size, replace=False)\n        else:\n            _idxs = self._get_indices(idx)\n            \n        # For test return only features\n        if self.split == \"test\":\n            return self._get_X(_idxs)\n        # For train and val splits return y also\n        return self._get_X(_idxs), self._get_y(_idxs)\n    \n    def _get_indices(self, idx):\n        _low = idx * self.cfg.batch_size\n        # Cap high at self.length so overflow does not occur\n        _high = min(_low + self.cfg.batch_size, self.length)\n        return self.mapping[_low:_high]\n    \n    def _get_X(self, indices):\n        _X = np.empty((len(indices), self.cfg.window_size, self.cfg.n_features), dtype=np.float16)\n        for i, idx in enumerate(indices):\n            _X[i] = self.values[idx-self.past_pad:idx+self.future_pad+1:self.cfg.wx, :self.cfg.n_features]\n        return _X\n    \n    def _get_y(self, indices):\n        return self.values[indices, self.cfg.n_features:self.cfg.n_features + self.cfg.n_labels]","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:36:14.544044Z","iopub.execute_input":"2023-04-21T13:36:14.544501Z","iopub.status.idle":"2023-04-21T13:36:14.569110Z","shell.execute_reply.started":"2023-04-21T13:36:14.544463Z","shell.execute_reply":"2023-04-21T13:36:14.568152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Transformer Model","metadata":{}},{"cell_type":"code","source":"# average_precision_score with positive sample added if no true positive cases are present\ndef calculate_precision(y_true, y_pred):\n    pad_width = ((0,0),(0,0)) if y_true.any(axis=0).all() else ((1,0),(0,0))\n    y_true, y_pred = np.pad(y_true, pad_width, constant_values=1), np.pad(y_pred, pad_width, constant_values=1)\n    return average_precision_score(y_true, y_pred)","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:33:54.301472Z","iopub.execute_input":"2023-04-21T13:33:54.301921Z","iopub.status.idle":"2023-04-21T13:33:54.315894Z","shell.execute_reply.started":"2023-04-21T13:33:54.301881Z","shell.execute_reply":"2023-04-21T13:33:54.314932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AveragePrecision(tf.keras.metrics.Metric):\n\n    def __init__(self, num_classes, thresholds=None, name='avg_precision', **kwargs):\n        super(AveragePrecision, self).__init__(name=name, **kwargs)\n        self.class_precision = [tf.keras.metrics.Precision(thresholds) for _ in range(num_classes)]\n\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        for i, precision in enumerate(self.class_precision):\n            precision.update_state(y_true[:,i], y_pred[:, i])\n\n    def result(self):\n        return tf.math.reduce_mean([precision.result() for precision in self.class_precision])\n    \n    def reset_state(self):\n        for precision in self.class_precision:\n            precision.reset_state()","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:33:54.317171Z","iopub.execute_input":"2023-04-21T13:33:54.317631Z","iopub.status.idle":"2023-04-21T13:33:54.330035Z","shell.execute_reply.started":"2023-04-21T13:33:54.317593Z","shell.execute_reply":"2023-04-21T13:33:54.329013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def transformer_encoder(inputs, head_size, num_heads, ff_dim, dropout=0):\n    # Normalization and Attention\n    x = layers.LayerNormalization(epsilon=1e-6)(inputs)\n    x = layers.MultiHeadAttention(\n        key_dim=head_size, num_heads=num_heads, dropout=dropout\n    )(x, x)\n    x = layers.Dropout(dropout)(x)\n    res = x + inputs\n\n    # Feed Forward Part\n    x = layers.LayerNormalization(epsilon=1e-6)(res)\n    x = layers.Conv1D(filters=ff_dim, kernel_size=1, activation=\"relu\")(x)\n    x = layers.Dropout(dropout)(x)\n    x = layers.Conv1D(filters=inputs.shape[-1], kernel_size=1)(x)\n    return x + res","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:33:54.331459Z","iopub.execute_input":"2023-04-21T13:33:54.331867Z","iopub.status.idle":"2023-04-21T13:33:54.341698Z","shell.execute_reply.started":"2023-04-21T13:33:54.331830Z","shell.execute_reply":"2023-04-21T13:33:54.340703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_model(\n    input_shape, n_classes, head_size, num_heads, ff_dim, num_transformer_blocks, mlp_units, dropout=0, mlp_dropout=0\n):\n    inputs = tf.keras.Input(shape=input_shape)\n    x = inputs\n    for _ in range(num_transformer_blocks):\n        x = transformer_encoder(x, head_size, num_heads, ff_dim, dropout)\n\n    x = layers.GlobalAveragePooling1D(data_format=\"channels_last\")(x) # channels_last->(batch, steps, features)\n    for dim in mlp_units:\n        x = layers.Dense(dim, activation=\"relu\")(x)\n        x = layers.Dropout(mlp_dropout)(x)\n    outputs = layers.Dense(n_classes)(x) # No activation\n    return tf.keras.Model(inputs, outputs)","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:33:54.343167Z","iopub.execute_input":"2023-04-21T13:33:54.343581Z","iopub.status.idle":"2023-04-21T13:33:54.353247Z","shell.execute_reply.started":"2023-04-21T13:33:54.343545Z","shell.execute_reply":"2023-04-21T13:33:54.352242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(checkpoint_path = None):\n    # Model adapted from https://keras.io/examples/timeseries/timeseries_transformer_classification/\n    model = build_model(\n        (cfg.window_size, cfg.n_features),\n        cfg.n_labels,\n        head_size=256,\n        num_heads=4,\n        ff_dim=4,\n        num_transformer_blocks=4,\n        mlp_units=[128],\n        mlp_dropout=cfg.model_dropout*2,\n        dropout=cfg.model_dropout,\n    )\n\n    if checkpoint_path is not None:\n        model.load_weights(checkpoint_path)\n    model.compile(\n        tf.keras.optimizers.Adam(learning_rate=cfg.lr), \n        loss = tf.keras.losses.BinaryCrossentropy(from_logits=True),\n        # metrics=[AveragePrecision(cfg.n_labels, thresholds=0.0)]\n    )\n    return model\n\nget_model().summary()","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:36:24.617692Z","iopub.execute_input":"2023-04-21T13:36:24.618449Z","iopub.status.idle":"2023-04-21T13:36:25.161142Z","shell.execute_reply.started":"2023-04-21T13:36:24.618406Z","shell.execute_reply":"2023-04-21T13:36:25.160307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"# scores = []\n# monitor = 'val_avg_precision'\n# mode = 'max'\n# for i in range(cfg.splits):\n#     print(f\"***Fold {i}{'*'*75}\")\n#     model = get_model()\n    \n#     train_ds = FOGSequence(fold_train_fpaths[i], trim=False) # IS_PUBLIC\n#     val_ds = FOGSequence(fold_valid_fpaths[i], split=\"val\")\n\n#     prec_ckpt = tf.keras.callbacks.ModelCheckpoint(\n#         f\"fold{i}_best_model_state.h5\", monitor=monitor, save_best_only=True, save_weights_only=True, mode=mode)\n#     es_cb = tf.keras.callbacks.EarlyStopping(monitor=monitor, patience=(cfg.num_epochs+1)//2)\n\n#     history = model.fit(\n#         train_ds, epochs=cfg.num_epochs, verbose=1, validation_data=val_ds, workers=5, use_multiprocessing=True, callbacks=[prec_ckpt, es_cb])\n    \n#     scores.append(np.max(history.history[monitor]))\n    \n#     if IS_PUBLIC:\n#         break","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:33:57.633890Z","iopub.execute_input":"2023-04-21T13:33:57.634252Z","iopub.status.idle":"2023-04-21T13:33:57.644842Z","shell.execute_reply.started":"2023-04-21T13:33:57.634206Z","shell.execute_reply":"2023-04-21T13:33:57.643901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_select_model(fold, ds, model_save_dir=''):\n    best_path, best_score = None, -1\n    print(f\"Validation for fold{fold}:\")\n    for model_path in sorted(glob(join(model_save_dir, f\"fold{fold}_*.h5\"))):\n        pred_time = perf_counter()\n        gc.collect()\n        score = calculate_precision(\n            ds.values[ds.mapping, cfg.n_features:cfg.n_features + cfg.n_labels], \n            expit(get_model(model_path).predict(ds, verbose=0)) # expit converts to sigmoid output\n        )\n        if best_score < score:\n            best_score = score\n            best_path = model_path\n        gc.collect()\n        print(\"\\t\", basename(model_path), f\": score-{score:.4f} in {format_time(perf_counter()-pred_time)}\")\n    print(basename(best_path), \"selected with score\", best_score)\n    return best_path","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:33:57.645820Z","iopub.execute_input":"2023-04-21T13:33:57.646165Z","iopub.status.idle":"2023-04-21T13:33:57.660539Z","shell.execute_reply.started":"2023-04-21T13:33:57.646127Z","shell.execute_reply":"2023-04-21T13:33:57.659661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_loop(train_paths, valid_paths, fold, model_save_dir=''):\n    gc.collect()\n    \n    train_ds = FOGSequence(train_paths, trim=False) # IS_PUBLIC\n    val_ds = FOGSequence(valid_paths, split=\"val\")\n    \n    model = get_model()\n    ckpt = tf.keras.callbacks.ModelCheckpoint(join(model_save_dir, f\"fold{fold}_model_\"+\"{epoch:02d}.h5\"), save_weights_only=True)\n    history = model.fit(train_ds, epochs=cfg.num_epochs, verbose=2, workers=5, use_multiprocessing=True, callbacks=[ckpt]) # validation_data=val_ds, \n    \n    best_model_path = predict_select_model(fold, val_ds, model_save_dir)\n    \n    del train_ds, val_ds, model, ckpt, history\n    gc.collect()\n    return best_model_path","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:33:57.661451Z","iopub.execute_input":"2023-04-21T13:33:57.661763Z","iopub.status.idle":"2023-04-21T13:33:57.672953Z","shell.execute_reply.started":"2023-04-21T13:33:57.661728Z","shell.execute_reply":"2023-04-21T13:33:57.672151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Main training loop\nmodel_paths = {'defog': [], 'tdcsfog':[]} # (17.83, 31.59)\n# for module in model_paths:\nmodule = 'defog'\nmodule_start = perf_counter()\nprint(f\"***Training {module}{'*'*75}\")\nif not exists(module): \n    os.mkdir(module)\ntrain_fpaths, valid_fpaths = fold_train_fpaths[module][cfg.fold], fold_valid_fpaths[module][cfg.fold]\n# (1:17, 9), (2:30, 11.35)\nprint(f\"Fold {cfg.fold}{'-'*25}\")\nmodel_paths[module].append(train_loop(train_fpaths, valid_fpaths, cfg.fold, model_save_dir=module))\nprint(f\"***{module} done in {format_time(perf_counter()-module_start)}{'*'*50}\\n\")","metadata":{"execution":{"iopub.status.busy":"2023-04-21T14:15:23.437425Z","iopub.execute_input":"2023-04-21T14:15:23.438439Z","iopub.status.idle":"2023-04-21T14:16:57.376104Z","shell.execute_reply.started":"2023-04-21T14:15:23.438396Z","shell.execute_reply":"2023-04-21T14:16:57.371479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # val_ds = FOGSequence(fold_valid_fpaths[i], split=\"val\")\n# y_true = val_ds.values[val_ds.mapping, cfg.n_features:cfg.n_features + cfg.n_labels].astype(np.int8)\n# y_pred = expit(model.predict(val_ds, verbose=1))\n# average_precision_score(y_true, y_pred)","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:34:39.347757Z","iopub.status.idle":"2023-04-21T13:34:39.348255Z","shell.execute_reply.started":"2023-04-21T13:34:39.348026Z","shell.execute_reply":"2023-04-21T13:34:39.348055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def valid_length(path):\n#     df = pd.read_csv(path)\n#     if 'Valid' not in df.columns:\n#         return len(df)\n#     return len(df[(df.Valid>0)&(df.Task>0)])\n\n# y_pred.shape, y_true.shape, val_ds.length, sum(map(valid_length, fold_valid_fpaths[i]))","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:34:39.349940Z","iopub.status.idle":"2023-04-21T13:34:39.350754Z","shell.execute_reply.started":"2023-04-21T13:34:39.350474Z","shell.execute_reply":"2023-04-21T13:34:39.350501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False: # IS_PUBLIC\n    all_ds = FOGSequence(glob(join(TRAIN_DIR, \"defog/*.csv\")) + glob(join(TRAIN_DIR, \"tdcsfog/*.csv\")), split=\"val\")\n    y_true = all_ds.values[all_ds.mapping, cfg.n_features:cfg.n_features + cfg.n_labels].astype(np.int8)\n    \n    y_pred_list = []\n    for model_path in glob(\"*.h5\"):\n        model = get_model(model_path)\n        print(f\"For model from {model_path}:\")\n        y_pred_list.append(expit(model.predict(all_ds, verbose=2)))\n    y_pred_list = np.array(y_pred_list)\n\n    y_pred = y_pred_list.mean(axis=0)\n    y_pred_median = y_pred_list[scores>=np.median(scores)].mean(axis=0)\n    print(\"All:\", average_precision_score(y_true, y_pred), \"Median:\", average_precision_score(y_true, y_pred_median))","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:34:39.352582Z","iopub.status.idle":"2023-04-21T13:34:39.353154Z","shell.execute_reply.started":"2023-04-21T13:34:39.352847Z","shell.execute_reply":"2023-04-21T13:34:39.352872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"# test_defog_paths = glob(join(TEST_DIR, \"defog/*.csv\"))\n# test_tdcsfog_paths = glob(join(TEST_DIR, \"tdcsfog/*.csv\"))\n\n# test_ds = FOGSequence(test_defog_paths+test_tdcsfog_paths, split=\"test\")\n\n# y_pred_list = []\n# for model_path in glob(\"*.h5\"):\n#     model = get_model(model_path)\n#     y_pred_list.append(expit(model.predict(test_ds, verbose=0)))\n        \n# y_pred = np.mean(y_pred_list, axis=0)\n# assert y_pred.shape[0]==len(test_ds.Ids)\n\n# # Predictions to DataFrame\n# preds = expit(preds).mean(axis=0).clip(0.0,1.0) # mean to avg all predictions & clip func just for safety\n# assert preds.shape == (len(ids),3) # Assert that preds and ids match\n# submission = pd.DataFrame({'Id': test_ds.Ids, 'StartHesitation': y_pred[:,0], 'Turn': y_pred[:,1], 'Walking': y_pred[:,2]})","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:34:39.355142Z","iopub.status.idle":"2023-04-21T13:34:39.355640Z","shell.execute_reply.started":"2023-04-21T13:34:39.355381Z","shell.execute_reply":"2023-04-21T13:34:39.355406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_defog_paths = glob(join(TEST_DIR, \"defog/*.csv\"))\n# test_tdcsfog_paths = glob(join(TEST_DIR, \"tdcsfog/*.csv\"))\n\n# test_ds_dict = {\n#     'defog':FOGSequence(test_defog_paths, split=\"test\"), \n#     'tdcsfog':FOGSequence(test_tdcsfog_paths, split=\"test\")\n# }\n\n# # Get test predictions\n# df_list = []\n# for module, test_ds in test_ds_dict.items():\n#     y_pred_list = []\n#     for model_path in model_paths[module]:\n#         model = get_model(model_path)\n#         y_pred_list.append(expit(model.predict(test_ds, verbose=0)))  # expit converts to sigmoid output\n#     y_pred = np.mean(y_pred_list, axis=0)\n#     df_list.append(pd.DataFrame(\n#         {'Id': test_ds.Ids, 'StartHesitation': y_pred[:,0], 'Turn': y_pred[:,1], 'Walking': y_pred[:,2]}))\n    \n# # Concatenate Prediction to DataFrames\n# submission = pd.concat(df_list)","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:34:39.357417Z","iopub.status.idle":"2023-04-21T13:34:39.357920Z","shell.execute_reply.started":"2023-04-21T13:34:39.357664Z","shell.execute_reply":"2023-04-21T13:34:39.357689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sample_submission = pd.read_csv(join(BASE_DIR, \"sample_submission.csv\"))\n\n# # Only keep Ids in sample_submission\n# submission = pd.merge(sample_submission[['Id']], submission, how='left', on='Id').fillna(0.0)\n# submission.to_csv(\"submission.csv\", index=False, float_format='%.5f') # round to 5 decimal places while keeping point notation","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:34:39.360768Z","iopub.status.idle":"2023-04-21T13:34:39.361620Z","shell.execute_reply.started":"2023-04-21T13:34:39.361351Z","shell.execute_reply":"2023-04-21T13:34:39.361379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !head -5 submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-04-21T13:34:39.362974Z","iopub.status.idle":"2023-04-21T13:34:39.363888Z","shell.execute_reply.started":"2023-04-21T13:34:39.363626Z","shell.execute_reply":"2023-04-21T13:34:39.363653Z"},"trusted":true},"execution_count":null,"outputs":[]}]}