{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":41880,"databundleVersionId":5677426,"sourceType":"competition"}],"dockerImageVersionId":30408,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np \nimport pandas as pd\nimport os\nfrom tqdm import tqdm\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torch.nn import CrossEntropyLoss\nimport torch.nn.functional as F\n\nMAIN_DIR = \"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/\"\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nFEATURES = [\"AccV\", \"AccML\", \"AccAP\"]\nTARGETS = [\"StartHesitation\", \"Turn\", \"Walking\"]\n\nN_EPOCHS = 1","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2025-08-05T17:21:55.707408Z","iopub.execute_input":"2025-08-05T17:21:55.708106Z","iopub.status.idle":"2025-08-05T17:21:58.875989Z","shell.execute_reply.started":"2025-08-05T17:21:55.708079Z","shell.execute_reply":"2025-08-05T17:21:58.875126Z"},"trusted":true},"outputs":[],"execution_count":1},{"cell_type":"markdown","source":"# 1. About the Data / Loading","metadata":{}},{"cell_type":"markdown","source":"CONTEXT : *People with FOG are filmed while performing certain tasks that are likely to increase its occurrence. Experts then review the video to score each frame, indicating when FOG occurred. While scoring in this manner is relatively reliable and sensitive, it is extremely time-consuming and requires specific expertise. Another method involves augmenting FOG-provoking testing with wearable devices. With more sensors, the detection of FOG becomes easier, however, compliance and usability may be reduced. Therefore, a combination of these two methods may be the best approach.*","metadata":{}},{"cell_type":"markdown","source":"The data series include three datasets, collected under distinct circumstances:\n\n- **tDCS FOG (tdcsfog)** dataset: data series collected in the lab, as subjects completed a FOG-provoking protocol.\n- **DeFOG (defog)** dataset: data series collected in the subject's home, as subjects completed a FOG-provoking protocol\n- **Daily Living (daily)** dataset: comprising one week of continuous 24/7 recordings from sixty-five subjects. Forty-five subjects exhibit FOG symptoms and also have series in the defog dataset, while the other twenty subjects do not exhibit FOG symptoms and do not have series elsewhere in the data.","metadata":{}},{"cell_type":"markdown","source":"**tDCS FOG** & **DeFOG** are annotated => Supervised Learning\n**Daily Living** is not annotated => Unsupervised/Semi-supervised Learning","metadata":{}},{"cell_type":"code","source":"# Reduce Memory Usage\n# reference : https://www.kaggle.com/code/arjanso/reducing-dataframe-memory-size-by-65 @ARJANGROEN\n\ndef reduce_memory_usage(df):\n    \n    start_mem = df.memory_usage().sum() / 1024**2\n    print('Memory usage of dataframe is {:.2f} MB'.format(start_mem))\n    \n    for col in df.columns:\n        col_type = df[col].dtype.name\n        if ((col_type != 'datetime64[ns]') & (col_type != 'category')):\n            if (col_type != 'object'):\n                c_min = df[col].min()\n                c_max = df[col].max()\n\n                if str(col_type)[:3] == 'int':\n                    if c_min > np.iinfo(np.int8).min and c_max < np.iinfo(np.int8).max:\n                        df[col] = df[col].astype(np.int8)\n                    elif c_min > np.iinfo(np.int16).min and c_max < np.iinfo(np.int16).max:\n                        df[col] = df[col].astype(np.int16)\n                    elif c_min > np.iinfo(np.int32).min and c_max < np.iinfo(np.int32).max:\n                        df[col] = df[col].astype(np.int32)\n                    elif c_min > np.iinfo(np.int64).min and c_max < np.iinfo(np.int64).max:\n                        df[col] = df[col].astype(np.int64)\n\n                else:\n                    if c_min > np.finfo(np.float16).min and c_max < np.finfo(np.float16).max:\n                        df[col] = df[col].astype(np.float16)\n                    elif c_min > np.finfo(np.float32).min and c_max < np.finfo(np.float32).max:\n                        df[col] = df[col].astype(np.float32)\n                    else:\n                        pass\n            else:\n                df[col] = df[col].astype('category')\n    mem_usg = df.memory_usage().sum() / 1024**2 \n    print(\"Memory usage became: \",mem_usg,\" MB\")\n    \n    return df","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-08-05T17:26:21.366138Z","iopub.execute_input":"2025-08-05T17:26:21.367044Z","iopub.status.idle":"2025-08-05T17:26:21.377519Z","shell.execute_reply.started":"2025-08-05T17:26:21.367007Z","shell.execute_reply":"2025-08-05T17:26:21.376554Z"},"trusted":true},"outputs":[],"execution_count":2},{"cell_type":"code","source":"def read_data(\n    dataset,\n    datatype,\n    subject_id = None):\n    \n    metadata = pd.read_csv(MAIN_DIR + dataset + \"_metadata.csv\")\n    \n    DATA_ROOT = MAIN_DIR + datatype + \"/\" + dataset\n    \n    if subject_id is not None:\n        files = [file for file in files if subject_id in file]\n    \n    df_res = pd.DataFrame()\n    for root, dirs, files in os.walk(DATA_ROOT):\n        for name in tqdm(files):\n            f = os.path.join(root, name)\n            query_datatype = pd.read_csv(f)\n            query_datatype[\"file\"] = name.replace(\".csv\", \"\")\n            df_res = pd.concat([df_res,query_datatype])\n    \n    df_res = metadata.merge(df_res,\n                          how = 'inner',\n                          left_on = 'Id',\n                          right_on = 'file')\n    df_res = df_res.drop([\"file\"], axis = 1)\n\n    df_res = reduce_memory_usage(df_res)\n        \n    return df_res\n    ","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:26:25.980638Z","iopub.execute_input":"2025-08-05T17:26:25.981361Z","iopub.status.idle":"2025-08-05T17:26:25.988288Z","shell.execute_reply.started":"2025-08-05T17:26:25.981324Z","shell.execute_reply":"2025-08-05T17:26:25.987433Z"},"trusted":true},"outputs":[],"execution_count":3},{"cell_type":"code","source":"#yehudit\ndef read_data_(\n    dataset=\"defog\",\n    datatype=\"train\",\n    subject_id = None):\n    \n    \n    metadata = pd.read_csv(MAIN_DIR + dataset + \"_metadata.csv\")\n    \n  \n        \n    DATA_ROOT = MAIN_DIR + datatype + \"/\" + dataset\n    \n    if subject_id is not None:\n        files = [file for file in files if subject_id in file]\n    \n    df_res = pd.DataFrame()\n    for root, dirs, files in os.walk(DATA_ROOT):\n        \n        for name in tqdm(files):\n            f = os.path.join(root, name)\n            query_datatype = pd.read_csv(f)\n            if dataset==\"defog\" and datatype==\"train\":\n                new_query_datatype= query_datatype.loc[query_datatype['Valid'] == True]\n            else:\n                new_query_datatype = query_datatype        \n            new_query_datatype[\"file\"] = name.replace(\".csv\", \"\")\n            df_res = pd.concat([df_res,new_query_datatype])\n        \n    if dataset == \"tdcsfog\" :\n        metadata = metadata.drop('Test', axis=1)\n    df_res = metadata.merge(df_res,\n                              how = 'inner',\n                              left_on = 'Id',\n                              right_on = 'file')\n        \n    df_res = df_res.drop([\"file\"], axis = 1)\n    \n    if (dataset == \"defog\") and (datatype==\"train\") : \n        df_res = df_res.drop('Valid', axis=1)\n        df_res = df_res.drop('Task', axis=1)\n        \n    \n    df_res = reduce_memory_usage(df_res)\n        \n    return df_res\n","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:26:29.589669Z","iopub.execute_input":"2025-08-05T17:26:29.59002Z","iopub.status.idle":"2025-08-05T17:26:29.599338Z","shell.execute_reply.started":"2025-08-05T17:26:29.589991Z","shell.execute_reply":"2025-08-05T17:26:29.598375Z"},"trusted":true},"outputs":[],"execution_count":4},{"cell_type":"code","source":"# #yehudit\n# def read_data_test_tdcs(\n#     dataset=\"defog\",\n#     datatype=\"train\",\n#     subject_id = None):\n    \n#     metadata = pd.read_csv(MAIN_DIR + dataset + \"_metadata.csv\")\n    \n  \n        \n#     DATA_ROOT = MAIN_DIR + datatype + \"/\" + dataset\n    \n#     if subject_id is not None:\n#         files = [file for file in files if subject_id in file]\n    \n#     df_res = pd.DataFrame()\n#     for root, dirs, files in os.walk(DATA_ROOT):\n        \n#         for name in tqdm(files):\n#             if name != \"0a89f859b5.csv\":\n#                 f = os.path.join(root, name)\n#                 query_datatype = pd.read_csv(f)\n#                 if dataset==\"defog\" and datatype==\"train\":\n#                     new_query_datatype= query_datatype.loc[query_datatype['Valid'] == True]\n#                 else:\n#                     new_query_datatype = query_datatype        \n#                 new_query_datatype[\"file\"] = name.replace(\".csv\", \"\")\n#                 df_res = pd.concat([df_res,new_query_datatype])\n        \n#     if dataset == \"tdcsfog\" :\n#         metadata = metadata.drop('Test', axis=1)\n#     df_res = metadata.merge(df_res,\n#                               how = 'inner',\n#                               left_on = 'Id',\n#                               right_on = 'file')\n        \n#     df_res = df_res.drop([\"file\"], axis = 1)\n    \n#     if (dataset == \"defog\") and (datatype==\"train\") : \n#         df_res = df_res.drop('Valid', axis=1)\n#         df_res = df_res.drop('Task', axis=1)\n        \n    \n#     df_res = reduce_memory_usage(df_res)\n        \n#     return df_res\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-05T17:26:34.534262Z","iopub.execute_input":"2025-08-05T17:26:34.534635Z","iopub.status.idle":"2025-08-05T17:26:34.540447Z","shell.execute_reply.started":"2025-08-05T17:26:34.534603Z","shell.execute_reply":"2025-08-05T17:26:34.539395Z"}},"outputs":[],"execution_count":5},{"cell_type":"markdown","source":"# ****create combined train dataset****","metadata":{}},{"cell_type":"code","source":"df_defog = read_data_( dataset=\"defog\",datatype=\"train\")\ndf_tdcsfog = read_data_( dataset=\"tdcsfog\",datatype=\"train\")\ncombined_ds_train = pd.concat([df_defog,df_tdcsfog])","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:26:38.864571Z","iopub.execute_input":"2025-08-05T17:26:38.865297Z","iopub.status.idle":"2025-08-05T17:29:14.804609Z","shell.execute_reply.started":"2025-08-05T17:26:38.865259Z","shell.execute_reply":"2025-08-05T17:29:14.803543Z"},"trusted":true},"outputs":[{"name":"stderr","text":"  0%|          | 0/91 [00:00<?, ?it/s]/opt/conda/lib/python3.7/site-packages/ipykernel_launcher.py:27: SettingWithCopyWarning: \nA value is trying to be set on a copy of a slice from a DataFrame.\nTry using .loc[row_indexer,col_indexer] = value instead\n\nSee the caveats in the documentation: https://pandas.pydata.org/pandas-docs/stable/user_guide/indexing.html#returning-a-view-versus-a-copy\n100%|██████████| 91/91 [00:23<00:00,  3.83it/s]\n","output_type":"stream"},{"name":"stdout","text":"Memory usage of dataframe is 374.50 MB\nMemory usage became:  97.52996635437012  MB\n","output_type":"stream"},{"name":"stderr","text":"100%|██████████| 833/833 [02:00<00:00,  6.89it/s]\n","output_type":"stream"},{"name":"stdout","text":"Memory usage of dataframe is 646.61 MB\nMemory usage became:  175.1631965637207  MB\n","output_type":"stream"}],"execution_count":6},{"cell_type":"markdown","source":"# ****create combined test dataset****","metadata":{}},{"cell_type":"code","source":"# df_defog_test = read_data_( dataset=\"defog\",datatype=\"test\")\ndf_tdcsfog_test = read_data_( dataset=\"tdcsfog\",datatype=\"test\")\n# combined_ds_test = pd.concat([df_defog_test,df_tdcsfog_test])","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:29:26.991482Z","iopub.execute_input":"2025-08-05T17:29:26.992173Z","iopub.status.idle":"2025-08-05T17:29:27.029932Z","shell.execute_reply.started":"2025-08-05T17:29:26.992136Z","shell.execute_reply":"2025-08-05T17:29:27.028728Z"},"trusted":true},"outputs":[{"name":"stderr","text":"100%|██████████| 1/1 [00:00<00:00, 92.24it/s]","output_type":"stream"},{"name":"stdout","text":"Memory usage of dataframe is 0.32 MB\nMemory usage became:  0.08963394165039062  MB\n","output_type":"stream"},{"name":"stderr","text":"\n","output_type":"stream"}],"execution_count":7},{"cell_type":"markdown","source":"# 2. DataLoader","metadata":{}},{"cell_type":"code","source":"class FOGDataset(Dataset):\n         \n    @staticmethod\n    def encode_target(data, targets_list):\n        conditions = []\n        for target in targets_list:\n            conditions.append((data[target] == 1))\n\n        event = np.select(conditions, targets_list, default='Normal')\n        le = LabelEncoder()\n        return le.fit_transform(event)\n\n    @staticmethod\n    def get_features_target(data, features_list, datatype):\n        if datatype == \"train\":\n            features, target = data[features_list], data[\"target\"]\n            return features, target\n        else:\n            features = data[features_list]\n            return features\n    \n    def __init__(self, dataset, datatype, features_list, targets_list, lookback):\n        self.datatype = datatype\n        #self.data = read_data(dataset = dataset, datatype = datatype)\n        self.data = dataset\n        self.features = features_list\n        self.targets = targets_list\n        self.data[\"Id_encoded\"], _ = pd.factorize(self.data[\"Id\"])\n        self.lookback = lookback\n        \n        if datatype == \"train\":\n            self.data = self.data[:1_000]\n            self.data[\"target\"] = FOGDataset.encode_target(self.data, self.targets)\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        if self.datatype == \"train\":            \n            features, targets = FOGDataset.get_features_target(self.data,\n                                               self.features,\n                                               self.datatype\n                                              )\n            \n            if idx < self.lookback :\n                features = features[0: self.lookback]\n                targets = targets[self.lookback]\n            \n            else:\n                features = features[idx - self.lookback: idx]\n                targets = targets[idx]\n                \n            features = torch.tensor(features.to_numpy(), dtype=torch.float32)\n            targets = torch.tensor(targets, dtype=torch.float32)\n            \n            return features, targets\n        else:\n            features = FOGDataset.get_features_target(self.data,\n                                               self.features,\n                                               self.datatype\n                                              )\n            \n            if idx < self.lookback :\n                features = features[0: self.lookback]\n            \n            else:\n                features = features[idx - self.lookback: idx]\n                \n            features = torch.tensor(features.to_numpy(), dtype=torch.float32)\n            \n            return features","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:29:32.802408Z","iopub.execute_input":"2025-08-05T17:29:32.803457Z","iopub.status.idle":"2025-08-05T17:29:32.815558Z","shell.execute_reply.started":"2025-08-05T17:29:32.803403Z","shell.execute_reply":"2025-08-05T17:29:32.814596Z"},"trusted":true},"outputs":[],"execution_count":8},{"cell_type":"code","source":"dataset_train_ = FOGDataset(\n    dataset = combined_ds_train,#\"tdcsfog\",#\n    datatype = \"train\",\n    features_list = FEATURES,\n    targets_list = TARGETS,\n    lookback = 2\n)\n","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:29:35.775414Z","iopub.execute_input":"2025-08-05T17:29:35.775796Z","iopub.status.idle":"2025-08-05T17:29:36.619752Z","shell.execute_reply.started":"2025-08-05T17:29:35.775765Z","shell.execute_reply":"2025-08-05T17:29:36.618763Z"},"trusted":true},"outputs":[{"name":"stderr","text":"/opt/conda/lib/python3.7/site-packages/ipykernel_launcher.py:33: SettingWithCopyWarning: \nA value is trying to be set on a copy of a slice from a DataFrame.\nTry using .loc[row_indexer,col_indexer] = value instead\n\nSee the caveats in the documentation: https://pandas.pydata.org/pandas-docs/stable/user_guide/indexing.html#returning-a-view-versus-a-copy\n","output_type":"stream"}],"execution_count":9},{"cell_type":"code","source":"dataset_test = FOGDataset(\n    dataset = df_tdcsfog_test,#\"tdcsfog\",#\n    datatype = \"test\",\n    features_list = FEATURES,\n    targets_list = TARGETS,\n    lookback = 2\n)","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:29:36.931578Z","iopub.execute_input":"2025-08-05T17:29:36.932283Z","iopub.status.idle":"2025-08-05T17:29:36.938034Z","shell.execute_reply.started":"2025-08-05T17:29:36.93225Z","shell.execute_reply":"2025-08-05T17:29:36.936987Z"},"trusted":true},"outputs":[],"execution_count":10},{"cell_type":"code","source":"dataloader_train = DataLoader(dataset_train_, batch_size = 8, shuffle = False)\ndataloader_test = DataLoader(dataset_test, batch_size = 1000, shuffle = False)\n\ncount = 0\n\nfor batch in dataloader_train:\n    features, target = batch\n    print(\"FEATURES EXAMPLES\")\n    print(features.shape)\n    print(features)\n    print(\"TARGET EXAMPLES\")\n    print(target)\n    print(\"\\n\")\n    if count > 1:  \n        break\n    count += 1","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:29:40.045863Z","iopub.execute_input":"2025-08-05T17:29:40.046204Z","iopub.status.idle":"2025-08-05T17:29:40.145458Z","shell.execute_reply.started":"2025-08-05T17:29:40.046174Z","shell.execute_reply":"2025-08-05T17:29:40.144504Z"},"trusted":true},"outputs":[{"name":"stdout","text":"FEATURES EXAMPLES\ntorch.Size([8, 2, 3])\ntensor([[[-9.7021e-01,  6.1615e-02, -2.6562e-01],\n         [-9.8438e-01,  4.4495e-02, -2.6562e-01]],\n\n        [[-9.7021e-01,  6.1615e-02, -2.6562e-01],\n         [-9.8438e-01,  4.4495e-02, -2.6562e-01]],\n\n        [[-9.7021e-01,  6.1615e-02, -2.6562e-01],\n         [-9.8438e-01,  4.4495e-02, -2.6562e-01]],\n\n        [[-9.8438e-01,  4.4495e-02, -2.6562e-01],\n         [-9.8438e-01,  2.9022e-02, -2.6562e-01]],\n\n        [[-9.8438e-01,  2.9022e-02, -2.6562e-01],\n         [-9.8438e-01,  1.5625e-02, -2.6562e-01]],\n\n        [[-9.8438e-01,  1.5625e-02, -2.6562e-01],\n         [-9.8486e-01,  1.5327e-02, -2.6562e-01]],\n\n        [[-9.8486e-01,  1.5327e-02, -2.6562e-01],\n         [-1.0000e+00,  1.5140e-04, -2.6562e-01]],\n\n        [[-1.0000e+00,  1.5140e-04, -2.6562e-01],\n         [-1.0000e+00,  1.5625e-02, -2.6562e-01]]])\nTARGET EXAMPLES\ntensor([0., 0., 0., 0., 0., 0., 0., 0.])\n\n\nFEATURES EXAMPLES\ntorch.Size([8, 2, 3])\ntensor([[[-1.0000e+00,  1.5625e-02, -2.6562e-01],\n         [-1.0000e+00,  1.5625e-02, -2.6562e-01]],\n\n        [[-1.0000e+00,  1.5625e-02, -2.6562e-01],\n         [-1.0000e+00,  1.5625e-02, -2.6562e-01]],\n\n        [[-1.0000e+00,  1.5625e-02, -2.6562e-01],\n         [-1.0000e+00,  1.5625e-02, -2.6562e-01]],\n\n        [[-1.0000e+00,  1.5625e-02, -2.6562e-01],\n         [-1.0000e+00,  1.6153e-04, -2.6562e-01]],\n\n        [[-1.0000e+00,  1.6153e-04, -2.6562e-01],\n         [-1.0000e+00,  0.0000e+00, -2.6562e-01]],\n\n        [[-1.0000e+00,  0.0000e+00, -2.6562e-01],\n         [-1.0000e+00,  0.0000e+00, -2.6562e-01]],\n\n        [[-1.0000e+00,  0.0000e+00, -2.6562e-01],\n         [-1.0000e+00,  0.0000e+00, -2.6562e-01]],\n\n        [[-1.0000e+00,  0.0000e+00, -2.6562e-01],\n         [-1.0000e+00,  0.0000e+00, -2.6562e-01]]])\nTARGET EXAMPLES\ntensor([0., 0., 0., 0., 0., 0., 0., 0.])\n\n\nFEATURES EXAMPLES\ntorch.Size([8, 2, 3])\ntensor([[[-1.0000,  0.0000, -0.2656],\n         [-1.0000,  0.0000, -0.2656]],\n\n        [[-1.0000,  0.0000, -0.2656],\n         [-1.0000,  0.0000, -0.2656]],\n\n        [[-1.0000,  0.0000, -0.2656],\n         [-1.0000,  0.0000, -0.2656]],\n\n        [[-1.0000,  0.0000, -0.2656],\n         [-1.0000,  0.0000, -0.2656]],\n\n        [[-1.0000,  0.0000, -0.2656],\n         [-0.9873,  0.0000, -0.2656]],\n\n        [[-0.9873,  0.0000, -0.2656],\n         [-0.9966,  0.0000, -0.2656]],\n\n        [[-0.9966,  0.0000, -0.2656],\n         [-0.9883,  0.0000, -0.2656]],\n\n        [[-0.9883,  0.0000, -0.2656],\n         [-0.9956,  0.0000, -0.2656]]])\nTARGET EXAMPLES\ntensor([0., 0., 0., 0., 0., 0., 0., 0.])\n\n\n","output_type":"stream"}],"execution_count":11},{"cell_type":"code","source":"class LSTMNet(nn.Module):\n    def __init__(self, input_size, hidden_size, num_layers, num_classes):\n        super().__init__()\n        self.hidden_size = hidden_size\n        self.num_layers = num_layers\n        self.lstm = nn.LSTM(input_size, hidden_size, num_layers, batch_first=True)\n        self.fc1 = nn.Linear(hidden_size, num_classes)\n\n    def forward(self, x):\n\n        hidden_state = torch.zeros((self.num_layers, x.size(0), self.hidden_size), dtype=torch.float32)\n        cell_state = torch.zeros((self.num_layers, x.size(0), self.hidden_size), dtype=torch.float32)\n\n        out, _ = self.lstm(x, (hidden_state, cell_state))\n        out = out[:, -1,:]\n        out = self.fc1(out)\n        return out\n","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:29:45.82927Z","iopub.execute_input":"2025-08-05T17:29:45.829919Z","iopub.status.idle":"2025-08-05T17:29:45.837201Z","shell.execute_reply.started":"2025-08-05T17:29:45.829888Z","shell.execute_reply":"2025-08-05T17:29:45.83627Z"},"trusted":true},"outputs":[],"execution_count":12},{"cell_type":"markdown","source":"# 4. Training","metadata":{}},{"cell_type":"code","source":"def train(model, dataloader, loss_fn, optimizer):\n    model.train()\n    total_loss = 0\n    for epoch in range(N_EPOCHS):\n        mean_precision = []\n        for (features, targets) in tqdm(dataloader):\n            optimizer.zero_grad()\n            preds = model(features)\n            loss = loss_fn(preds, targets.long())\n            mean_precision.append(loss.item())\n            loss.backward()\n            optimizer.step()\n        \n        print(\"Average Precision : \", np.mean(mean_precision))\n    \n    return model\n\ndef predict(model, dataloader): \n    model.eval()\n    predictions = np.empty(len(dataset_test))\n    count = 0\n    for features in tqdm(dataloader):\n        preds = model(features)\n        preds = torch.argmax(preds, dim = 1)\n        preds = preds.numpy()\n        predictions[count : count + len(preds)] = preds\n        count += len(preds)\n            \n    return predictions","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:29:49.573612Z","iopub.execute_input":"2025-08-05T17:29:49.574376Z","iopub.status.idle":"2025-08-05T17:29:49.582325Z","shell.execute_reply.started":"2025-08-05T17:29:49.574342Z","shell.execute_reply":"2025-08-05T17:29:49.581155Z"},"trusted":true},"outputs":[],"execution_count":13},{"cell_type":"code","source":"INPUT_SIZE = len(FEATURES)\nHIDDEN_SIZE = 10\nNUM_LAYERS = 1\nNUM_CLASSES = 4\nPARAMS = {\n    \"input_size\" : INPUT_SIZE,\n    \"hidden_size\" : HIDDEN_SIZE,\n    \"num_layers\" : NUM_LAYERS,\n    \"num_classes\" : NUM_CLASSES\n}\nmodel = LSTMNet(**PARAMS)\n\nloss_fn = CrossEntropyLoss()\n\noptimizer = optim.Adam(model.parameters(), lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:29:51.70908Z","iopub.execute_input":"2025-08-05T17:29:51.709959Z","iopub.status.idle":"2025-08-05T17:29:51.718704Z","shell.execute_reply.started":"2025-08-05T17:29:51.709921Z","shell.execute_reply":"2025-08-05T17:29:51.717736Z"},"trusted":true},"outputs":[],"execution_count":14},{"cell_type":"code","source":"model = train(\n    model, \n    dataloader_train,\n    loss_fn,\n    optimizer\n)","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:29:55.046956Z","iopub.execute_input":"2025-08-05T17:29:55.047639Z","iopub.status.idle":"2025-08-05T17:29:55.801899Z","shell.execute_reply.started":"2025-08-05T17:29:55.047604Z","shell.execute_reply":"2025-08-05T17:29:55.800929Z"},"trusted":true},"outputs":[{"name":"stderr","text":"100%|██████████| 125/125 [00:00<00:00, 167.65it/s]","output_type":"stream"},{"name":"stdout","text":"Average Precision :  1.1088211631774902\n","output_type":"stream"},{"name":"stderr","text":"\n","output_type":"stream"}],"execution_count":15},{"cell_type":"code","source":"preds_tdcsfog = predict(model, dataloader_test)","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:30:01.166015Z","iopub.execute_input":"2025-08-05T17:30:01.166949Z","iopub.status.idle":"2025-08-05T17:30:03.297849Z","shell.execute_reply.started":"2025-08-05T17:30:01.166912Z","shell.execute_reply":"2025-08-05T17:30:03.296779Z"},"trusted":true},"outputs":[{"name":"stderr","text":"100%|██████████| 5/5 [00:02<00:00,  2.35it/s]\n","output_type":"stream"}],"execution_count":16},{"cell_type":"markdown","source":"# 3. First Submission","metadata":{"execution":{"iopub.status.busy":"2023-03-28T09:27:08.999527Z","iopub.execute_input":"2023-03-28T09:27:09.000095Z","iopub.status.idle":"2023-03-28T09:27:09.013584Z","shell.execute_reply.started":"2023-03-28T09:27:09.000046Z","shell.execute_reply":"2023-03-28T09:27:09.011646Z"}}},{"cell_type":"code","source":"# dataset_train = read_data(dataset = \"tdcsfog\", datatype = \"test\")\n# dataloader_train = DataLoader(dataset_train, batch_size = 8, shuffle = False)\n# dataset_test = read_data(dataset = \"defog\", datatype = \"test\")\n# dataloader_test = DataLoader(dataset_test, batch_size = 1000, shuffle = False)\n\ntest_tdcsfog = read_data(dataset = \"tdcsfog\", datatype = \"test\")\ntest_defog = read_data(dataset = \"defog\", datatype = \"test\")\n\nlen_test_tdcsfog = len(test_tdcsfog)\nlen_test_defog = len(test_defog)\nprint(len(test_tdcsfog))\nprint(len(test_defog))","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:30:09.64835Z","iopub.execute_input":"2025-08-05T17:30:09.648718Z","iopub.status.idle":"2025-08-05T17:30:10.138396Z","shell.execute_reply.started":"2025-08-05T17:30:09.648685Z","shell.execute_reply":"2025-08-05T17:30:10.137388Z"},"trusted":true},"outputs":[{"name":"stderr","text":"100%|██████████| 1/1 [00:00<00:00, 145.92it/s]\n","output_type":"stream"},{"name":"stdout","text":"Memory usage of dataframe is 0.36 MB\nMemory usage became:  0.09409904479980469  MB\n","output_type":"stream"},{"name":"stderr","text":"100%|██████████| 1/1 [00:00<00:00,  3.38it/s]","output_type":"stream"},{"name":"stdout","text":"Memory usage of dataframe is 19.34 MB\nMemory usage became:  5.910381317138672  MB\n4682\n281688\n","output_type":"stream"},{"name":"stderr","text":"\n","output_type":"stream"}],"execution_count":17},{"cell_type":"code","source":"test_tdcsfog[\"y_pred\"] = preds_tdcsfog\ntest_defog[\"y_pred\"] = np.zeros(len_test_defog)","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:30:13.591287Z","iopub.execute_input":"2025-08-05T17:30:13.592128Z","iopub.status.idle":"2025-08-05T17:30:13.597829Z","shell.execute_reply.started":"2025-08-05T17:30:13.592093Z","shell.execute_reply":"2025-08-05T17:30:13.596847Z"},"trusted":true},"outputs":[],"execution_count":18},{"cell_type":"code","source":"sub_fmt = pd.read_csv(\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:30:15.509604Z","iopub.execute_input":"2025-08-05T17:30:15.510189Z","iopub.status.idle":"2025-08-05T17:30:15.695699Z","shell.execute_reply.started":"2025-08-05T17:30:15.510156Z","shell.execute_reply":"2025-08-05T17:30:15.694649Z"},"trusted":true},"outputs":[],"execution_count":19},{"cell_type":"code","source":"sub = pd.DataFrame()\nfor data in [test_tdcsfog, test_defog]:\n    temp = data.copy()\n    temp[\"Id\"] = temp.apply(lambda x : str(x.Id) + \"_\" + str(x.Time), axis = 1)\n    temp['StartHesitation'] = np.where(temp['y_pred']==1, 1, 0)\n    temp['Turn'] = np.where(temp['y_pred']==2, 1, 0)\n    temp['Walking'] = np.where(temp['y_pred']==3, 1, 0)\n    temp = temp[[\"Id\"] + TARGETS]\n    sub = pd.concat([sub, temp])","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:30:18.601077Z","iopub.execute_input":"2025-08-05T17:30:18.601912Z","iopub.status.idle":"2025-08-05T17:30:23.172047Z","shell.execute_reply.started":"2025-08-05T17:30:18.601879Z","shell.execute_reply":"2025-08-05T17:30:23.170914Z"},"trusted":true},"outputs":[],"execution_count":20},{"cell_type":"code","source":"test_tdcsfog.head()","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:30:26.259183Z","iopub.execute_input":"2025-08-05T17:30:26.259961Z","iopub.status.idle":"2025-08-05T17:30:26.277347Z","shell.execute_reply.started":"2025-08-05T17:30:26.25993Z","shell.execute_reply":"2025-08-05T17:30:26.276363Z"},"trusted":true},"outputs":[{"execution_count":21,"output_type":"execute_result","data":{"text/plain":"           Id Subject  Visit  Test Medication  Time      AccV     AccML  \\\n0  003f117e14  4dc2f8      3     2         on     0 -9.531250  0.566406   \n1  003f117e14  4dc2f8      3     2         on     1 -9.539062  0.563965   \n2  003f117e14  4dc2f8      3     2         on     2 -9.531250  0.561523   \n3  003f117e14  4dc2f8      3     2         on     3 -9.531250  0.564453   \n4  003f117e14  4dc2f8      3     2         on     4 -9.539062  0.562012   \n\n      AccAP  y_pred  \n0 -1.413086     0.0  \n1 -1.440430     0.0  \n2 -1.429688     0.0  \n3 -1.415039     0.0  \n4 -1.429688     0.0  ","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>Id</th>\n      <th>Subject</th>\n      <th>Visit</th>\n      <th>Test</th>\n      <th>Medication</th>\n      <th>Time</th>\n      <th>AccV</th>\n      <th>AccML</th>\n      <th>AccAP</th>\n      <th>y_pred</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>003f117e14</td>\n      <td>4dc2f8</td>\n      <td>3</td>\n      <td>2</td>\n      <td>on</td>\n      <td>0</td>\n      <td>-9.531250</td>\n      <td>0.566406</td>\n      <td>-1.413086</td>\n      <td>0.0</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>003f117e14</td>\n      <td>4dc2f8</td>\n      <td>3</td>\n      <td>2</td>\n      <td>on</td>\n      <td>1</td>\n      <td>-9.539062</td>\n      <td>0.563965</td>\n      <td>-1.440430</td>\n      <td>0.0</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>003f117e14</td>\n      <td>4dc2f8</td>\n      <td>3</td>\n      <td>2</td>\n      <td>on</td>\n      <td>2</td>\n      <td>-9.531250</td>\n      <td>0.561523</td>\n      <td>-1.429688</td>\n      <td>0.0</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>003f117e14</td>\n      <td>4dc2f8</td>\n      <td>3</td>\n      <td>2</td>\n      <td>on</td>\n      <td>3</td>\n      <td>-9.531250</td>\n      <td>0.564453</td>\n      <td>-1.415039</td>\n      <td>0.0</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>003f117e14</td>\n      <td>4dc2f8</td>\n      <td>3</td>\n      <td>2</td>\n      <td>on</td>\n      <td>4</td>\n      <td>-9.539062</td>\n      <td>0.562012</td>\n      <td>-1.429688</td>\n      <td>0.0</td>\n    </tr>\n  </tbody>\n</table>\n</div>"},"metadata":{}}],"execution_count":21},{"cell_type":"code","source":"sub.info()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-05T17:30:33.860697Z","iopub.execute_input":"2025-08-05T17:30:33.861084Z","iopub.status.idle":"2025-08-05T17:30:33.893507Z","shell.execute_reply.started":"2025-08-05T17:30:33.861051Z","shell.execute_reply":"2025-08-05T17:30:33.89231Z"}},"outputs":[{"name":"stdout","text":"<class 'pandas.core.frame.DataFrame'>\nInt64Index: 286370 entries, 0 to 281687\nData columns (total 4 columns):\n #   Column           Non-Null Count   Dtype \n---  ------           --------------   ----- \n 0   Id               286370 non-null  object\n 1   StartHesitation  286370 non-null  int64 \n 2   Turn             286370 non-null  int64 \n 3   Walking          286370 non-null  int64 \ndtypes: int64(3), object(1)\nmemory usage: 10.9+ MB\n","output_type":"stream"}],"execution_count":22},{"cell_type":"code","source":"sub_fmt.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-05T17:30:40.277799Z","iopub.execute_input":"2025-08-05T17:30:40.278142Z","iopub.status.idle":"2025-08-05T17:30:40.287818Z","shell.execute_reply.started":"2025-08-05T17:30:40.278113Z","shell.execute_reply":"2025-08-05T17:30:40.286873Z"}},"outputs":[{"execution_count":23,"output_type":"execute_result","data":{"text/plain":"             Id  StartHesitation  Turn  Walking\n0  003f117e14_0                0     0        0\n1  003f117e14_1                0     0        0\n2  003f117e14_2                0     0        0\n3  003f117e14_3                0     0        0\n4  003f117e14_4                0     0        0","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>Id</th>\n      <th>StartHesitation</th>\n      <th>Turn</th>\n      <th>Walking</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>003f117e14_0</td>\n      <td>0</td>\n      <td>0</td>\n      <td>0</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>003f117e14_1</td>\n      <td>0</td>\n      <td>0</td>\n      <td>0</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>003f117e14_2</td>\n      <td>0</td>\n      <td>0</td>\n      <td>0</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>003f117e14_3</td>\n      <td>0</td>\n      <td>0</td>\n      <td>0</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>003f117e14_4</td>\n      <td>0</td>\n      <td>0</td>\n      <td>0</td>\n    </tr>\n  </tbody>\n</table>\n</div>"},"metadata":{}}],"execution_count":23},{"cell_type":"code","source":"len(sub_fmt)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-05T17:30:43.824408Z","iopub.execute_input":"2025-08-05T17:30:43.824769Z","iopub.status.idle":"2025-08-05T17:30:43.831442Z","shell.execute_reply.started":"2025-08-05T17:30:43.824739Z","shell.execute_reply":"2025-08-05T17:30:43.829854Z"}},"outputs":[{"execution_count":24,"output_type":"execute_result","data":{"text/plain":"286370"},"metadata":{}}],"execution_count":24},{"cell_type":"code","source":"sub.to_csv(\"/kaggle/working/submission.csv\", index = False)","metadata":{"execution":{"iopub.status.busy":"2025-08-05T17:30:46.452932Z","iopub.execute_input":"2025-08-05T17:30:46.453287Z","iopub.status.idle":"2025-08-05T17:30:46.871549Z","shell.execute_reply.started":"2025-08-05T17:30:46.453256Z","shell.execute_reply":"2025-08-05T17:30:46.870195Z"},"trusted":true},"outputs":[],"execution_count":25},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}