{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","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":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nimport torch\nfrom torch import nn\nfrom torch.optim import Adam\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import classification_report, confusion_matrix, f1_score, precision_score, recall_score, roc_auc_score\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfrom tqdm.notebook import tqdm, trange\nimport glob\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-29T23:25:14.698243Z","iopub.execute_input":"2024-05-29T23:25:14.699014Z","iopub.status.idle":"2024-05-29T23:25:14.706451Z","shell.execute_reply.started":"2024-05-29T23:25:14.698977Z","shell.execute_reply":"2024-05-29T23:25:14.705343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# defog_metadata = pd.read_csv('/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/defog_metadata.csv')\n# tdcsfog_metadata = pd.read_csv('/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/tdcsfog_metadata.csv')\n\ndefog_path = '/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/defog/'\ntdcsfog_path = '/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/tdcsfog/'\n\n# defog_metadata['filepath'] = defog_metadata['Id'].apply(lambda x: defog_path + x + '.csv')\n# tdcsfog_metadata['filepath'] = tdcsfog_metadata['Id'].apply(lambda x: tdcsfog_path + x + '.csv')\n\n# defog_metadata = defog_metadata.drop(0)\n# tdcsfog_metadata = tdcsfog_metadata.drop(0)\n\n# train = pd.concat([defog_metadata, tdcsfog_metadata.drop('Test', axis = 1)], axis = 0).reset_index(drop=True)\n\nfeatures = ['AccV', 'AccML', 'AccAP']\ntargets = ['StartHesitation', 'Turn', 'Walking']","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:25:15.732319Z","iopub.execute_input":"2024-05-29T23:25:15.733009Z","iopub.status.idle":"2024-05-29T23:25:15.737893Z","shell.execute_reply.started":"2024-05-29T23:25:15.732978Z","shell.execute_reply":"2024-05-29T23:25:15.736946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"defog_files = glob.glob(defog_path + '*')\ntdcsfog_files = glob.glob(tdcsfog_path + '*')","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:25:18.112483Z","iopub.execute_input":"2024-05-29T23:25:18.113337Z","iopub.status.idle":"2024-05-29T23:25:18.136000Z","shell.execute_reply.started":"2024-05-29T23:25:18.113280Z","shell.execute_reply":"2024-05-29T23:25:18.135302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.DataFrame({'path' : defog_files + tdcsfog_files})","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:25:19.072402Z","iopub.execute_input":"2024-05-29T23:25:19.072997Z","iopub.status.idle":"2024-05-29T23:25:19.078082Z","shell.execute_reply.started":"2024-05-29T23:25:19.072964Z","shell.execute_reply":"2024-05-29T23:25:19.077119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ndef get_pos_frac(path):\n    df = pd.read_csv(path)\n    \n    return [len(df), df[targets[0]].mean(), df[targets[1]].mean(), df[targets[2]].mean()]\n\n# get_pos_frac('/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/tdcsfog/0506d9a39f.csv')\ndf['metadata'] = df['path'].apply(lambda x: get_pos_frac(x))","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:25:21.492455Z","iopub.execute_input":"2024-05-29T23:25:21.493137Z","iopub.status.idle":"2024-05-29T23:25:43.158748Z","shell.execute_reply.started":"2024-05-29T23:25:21.493105Z","shell.execute_reply":"2024-05-29T23:25:43.157747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metadata = pd.DataFrame(np.array(np.concatenate([np.array(item).reshape(-1, 1) for item in df['metadata'].values], axis=1)).T, columns=['length', f'{targets[0]}_frac', f'{targets[1]}_frac', f'{targets[2]}_frac'])\ndf = pd.concat([df, metadata], axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:25:43.160329Z","iopub.execute_input":"2024-05-29T23:25:43.160638Z","iopub.status.idle":"2024-05-29T23:25:43.169851Z","shell.execute_reply.started":"2024-05-29T23:25:43.160613Z","shell.execute_reply":"2024-05-29T23:25:43.169093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[df.columns[-3:]].plot.hist(bins=100, subplots=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:25:43.170797Z","iopub.execute_input":"2024-05-29T23:25:43.171050Z","iopub.status.idle":"2024-05-29T23:25:44.518813Z","shell.execute_reply.started":"2024-05-29T23:25:43.171027Z","shell.execute_reply":"2024-05-29T23:25:44.517967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Preprocessing Pipeline","metadata":{}},{"cell_type":"code","source":"defog_path = '/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/defog/'\ntdcsfog_path = '/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/tdcsfog/'\n\nfeatures = ['AccV', 'AccML', 'AccAP']\ntargets = ['StartHesitation', 'Turn', 'Walking']\n\ndefog_files = glob.glob(defog_path + '*')\ntdcsfog_files = glob.glob(tdcsfog_path + '*')\n\ndf = pd.DataFrame({'path' : defog_files + tdcsfog_files})\n\ndef get_pos_frac(path):\n    df = pd.read_csv(path)\n    \n    return [len(df), df[targets[0]].mean(), df[targets[1]].mean(), df[targets[2]].mean()]\n\n# get_pos_frac('/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/tdcsfog/0506d9a39f.csv')\ndf['metadata'] = df['path'].apply(lambda x: get_pos_frac(x))\n\nmetadata = pd.DataFrame(np.array(np.concatenate([np.array(item).reshape(-1, 1) for item in df['metadata'].values], axis=1)).T, columns=['length', f'{targets[0]}_frac', f'{targets[1]}_frac', f'{targets[2]}_frac'])\ndf = pd.concat([df, metadata], axis=1)\n\nlabel_map = {\n    'triple_pos': 3,\n    'dubble_pos': 2,\n    'single_pos': 1,\n    'no_pos': 0\n}\n\nsingle_pos = ((df[df.columns[-3]] + df[df.columns[-2]] + df[df.columns[-1]]) > 0).astype(int)*label_map['single_pos']\ndubble_pos = ((df[df.columns[-3]]*df[df.columns[-2]] + df[df.columns[-3]]*df[df.columns[-1]] + df[df.columns[-2]]*df[df.columns[-1]]) > 0).astype(int)*label_map['dubble_pos']\ntriple_pos = ((df[df.columns[-3]]*df[df.columns[-2]]*df[df.columns[-1]]) > 0).astype(int)*label_map['triple_pos']\n\nlabel = single_pos.to_frame().rename(columns={0:'label'})\nlabel.loc[dubble_pos[dubble_pos != 0].index,:] = label_map['dubble_pos']\nlabel.loc[triple_pos[triple_pos != 0].index,:] = label_map['triple_pos']\n# label\n\ndf['label'] = label\n\ndef stratified_train_test_split(df):\n    no_pos = df[df['label'] == 0]\n    pos = df[df['label'] != 0]\n    \n#     groups = [pos[pos['label'] == _class] for _class in pos['label'].unique()]\n#     splited_groups = train_test_split(*groups, test_size=0.2)\n    \n    train_groups = []\n    test_groups = []\n    \n    for _class in pos['label'].unique():\n        selected_df = df[df['label'] == _class]\n        train, test = train_test_split(selected_df, test_size=0.2)\n        \n        train_groups.append(train)\n        test_groups.append(test)\n        \n    train = pd.concat(train_groups, axis=0)\n    test = pd.concat(test_groups, axis=0)\n    \n    return train, test\n\n\ntrain, test = stratified_train_test_split(df)\n\ntrain['label'].value_counts().plot.bar()","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:27:26.259459Z","iopub.execute_input":"2024-05-29T23:27:26.260337Z","iopub.status.idle":"2024-05-29T23:27:48.118758Z","shell.execute_reply.started":"2024-05-29T23:27:26.260305Z","shell.execute_reply":"2024-05-29T23:27:48.117861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_map = {\n    'triple_pos': 3,\n    'dubble_pos': 2,\n    'single_pos': 1,\n    'no_pos': 0\n}\n\nsingle_pos = ((df[df.columns[-3]] + df[df.columns[-2]] + df[df.columns[-1]]) > 0).astype(int)*label_map['single_pos']\ndubble_pos = ((df[df.columns[-3]]*df[df.columns[-2]] + df[df.columns[-3]]*df[df.columns[-1]] + df[df.columns[-2]]*df[df.columns[-1]]) > 0).astype(int)*label_map['dubble_pos']\ntriple_pos = ((df[df.columns[-3]]*df[df.columns[-2]]*df[df.columns[-1]]) > 0).astype(int)*label_map['triple_pos']","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:27:48.120741Z","iopub.execute_input":"2024-05-29T23:27:48.121090Z","iopub.status.idle":"2024-05-29T23:27:48.131593Z","shell.execute_reply.started":"2024-05-29T23:27:48.121057Z","shell.execute_reply":"2024-05-29T23:27:48.130570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label = single_pos.to_frame().rename(columns={0:'label'})\nlabel.loc[dubble_pos[dubble_pos != 0].index,:] = label_map['dubble_pos']\nlabel.loc[triple_pos[triple_pos != 0].index,:] = label_map['triple_pos']\n# label","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:27:48.132739Z","iopub.execute_input":"2024-05-29T23:27:48.133056Z","iopub.status.idle":"2024-05-29T23:27:48.146671Z","shell.execute_reply.started":"2024-05-29T23:27:48.133026Z","shell.execute_reply":"2024-05-29T23:27:48.145840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label.value_counts().plot.bar()","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:27:48.148576Z","iopub.execute_input":"2024-05-29T23:27:48.148868Z","iopub.status.idle":"2024-05-29T23:27:48.326174Z","shell.execute_reply.started":"2024-05-29T23:27:48.148845Z","shell.execute_reply":"2024-05-29T23:27:48.325347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"positive_frac = ((metadata[f\"{targets[0]}_frac\"] + metadata[f\"{targets[1]}_frac\"] + metadata[f\"{targets[2]}_frac\"]) / 3).mean()\nnegative_frac_in_no_pos_class = metadata.loc[label[label['label'] == 0].index, 'length'].sum() / metadata['length'].sum() ","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:27:48.327282Z","iopub.execute_input":"2024-05-29T23:27:48.327619Z","iopub.status.idle":"2024-05-29T23:27:48.335312Z","shell.execute_reply.started":"2024-05-29T23:27:48.327587Z","shell.execute_reply":"2024-05-29T23:27:48.334305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['label'] = label\n# df","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:27:48.336560Z","iopub.execute_input":"2024-05-29T23:27:48.336919Z","iopub.status.idle":"2024-05-29T23:27:48.345507Z","shell.execute_reply.started":"2024-05-29T23:27:48.336848Z","shell.execute_reply":"2024-05-29T23:27:48.344594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def stratified_train_test_split(df):\n    no_pos = df[df['label'] == 0]\n    pos = df[df['label'] != 0]\n    \n#     groups = [pos[pos['label'] == _class] for _class in pos['label'].unique()]\n#     splited_groups = train_test_split(*groups, test_size=0.2)\n    \n    train_groups = []\n    test_groups = []\n    \n    for _class in pos['label'].unique():\n        selected_df = df[df['label'] == _class]\n        train, test = train_test_split(selected_df, test_size=0.2)\n        \n        train_groups.append(train)\n        test_groups.append(test)\n        \n    train = pd.concat(train_groups, axis=0)\n    test = pd.concat(test_groups, axis=0)\n    \n    return train, test\n\n\ntrain, test = stratified_train_test_split(df)\n        \n#     for _class in pos['label'].unique():\n#         selected_df = pos[pos['label'] == _class]\n#         train, test = train_test_split(selected_df)","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:27:48.346515Z","iopub.execute_input":"2024-05-29T23:27:48.346788Z","iopub.status.idle":"2024-05-29T23:27:48.361210Z","shell.execute_reply.started":"2024-05-29T23:27:48.346766Z","shell.execute_reply":"2024-05-29T23:27:48.360439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train['label'].value_counts().plot.bar()","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:27:48.362126Z","iopub.execute_input":"2024-05-29T23:27:48.362434Z","iopub.status.idle":"2024-05-29T23:27:48.579090Z","shell.execute_reply.started":"2024-05-29T23:27:48.362411Z","shell.execute_reply":"2024-05-29T23:27:48.578124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:27:48.580371Z","iopub.execute_input":"2024-05-29T23:27:48.580996Z","iopub.status.idle":"2024-05-29T23:27:48.600318Z","shell.execute_reply.started":"2024-05-29T23:27:48.580961Z","shell.execute_reply":"2024-05-29T23:27:48.599338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FOGDataset(Dataset):\n    def __init__(self, df, treshold=5000):\n        super().__init__()\n        self.df = df.reset_index()\n        self.treshold = treshold\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        filepath = self.df.loc[idx, 'path']\n        data_file = pd.read_csv(filepath)\n        \n        X = data_file[features]\n        y = data_file[targets]\n        \n        X = torch.tensor(X.values)\n        y = torch.tensor(y.values)\n        \n        length = len(data_file)\n        if (length == self.treshold):\n            pass\n        elif (length > self.treshold):\n            start_idx = np.random.choice(data_file.index[:-self.treshold])\n            X = X[start_idx:start_idx + self.treshold]\n            y = y[start_idx:start_idx + self.treshold]\n        else:\n            X_zeros = torch.zeros(self.treshold, X.shape[-1])\n            X_zeros[:X.shape[0], :] = X\n            X = X_zeros\n            \n            y_zeros = torch.zeros(self.treshold, y.shape[-1])\n            y_zeros[:y.shape[0], :] = y\n            y = y_zeros\n            \n            del X_zeros, y_zeros\n        \n        return X.T, y.T\n    \ndataset = FOGDataset(train)","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:28:38.513068Z","iopub.execute_input":"2024-05-29T23:28:38.513751Z","iopub.status.idle":"2024-05-29T23:28:38.524573Z","shell.execute_reply.started":"2024-05-29T23:28:38.513720Z","shell.execute_reply":"2024-05-29T23:28:38.523521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X, y = dataset[120]\nplt.subplots(2, 1)\nplt.subplot(2, 1, 1)\nplt.plot(X.T.cpu().detach().numpy())\nplt.subplot(2, 1, 2)\nplt.plot(y.T.cpu().detach().numpy())\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:28:39.357287Z","iopub.execute_input":"2024-05-29T23:28:39.358158Z","iopub.status.idle":"2024-05-29T23:28:39.639594Z","shell.execute_reply.started":"2024-05-29T23:28:39.358123Z","shell.execute_reply":"2024-05-29T23:28:39.638692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_gen = FOGDataset(train)\ntrain_gen = DataLoader(train_gen, batch_size=32)\ntest_gen = FOGDataset(test)\ntest_gen = DataLoader(test_gen, batch_size=32)","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:28:41.598012Z","iopub.execute_input":"2024-05-29T23:28:41.598362Z","iopub.status.idle":"2024-05-29T23:28:41.607334Z","shell.execute_reply.started":"2024-05-29T23:28:41.598335Z","shell.execute_reply":"2024-05-29T23:28:41.606501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"next(iter(train_gen))[1].shape","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:28:43.052687Z","iopub.execute_input":"2024-05-29T23:28:43.053468Z","iopub.status.idle":"2024-05-29T23:28:45.674731Z","shell.execute_reply.started":"2024-05-29T23:28:43.053435Z","shell.execute_reply":"2024-05-29T23:28:45.673692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\nclass WaveBlock(nn.Module):\n    def __init__(self, input_features, num_dilation_conv, kernel_size, num_features):\n        super().__init__()\n        self.casual_conv1 = nn.Conv1d(input_features, num_features, kernel_size, padding='same')\n        \n        dilation_rates = [2**i for i in range(num_dilation_conv)]\n        \n        self.tanh_conv = nn.Sequential(\n            *[nn.Conv1d(num_features, num_features, kernel_size, dilation = rate, padding='same') for rate in dilation_rates]\n        )\n        \n        self.sig_conv = nn.Sequential(\n            *[nn.Conv1d(num_features, num_features, kernel_size, dilation = rate, padding='same') for rate in dilation_rates]\n        )\n        self.batch_norm1 = nn.BatchNorm1d(num_features)\n        \n        self.casual_conv2 = nn.Conv1d(num_features, num_features, kernel_size, padding='same')\n        self.batch_norm2 = nn.BatchNorm1d(num_features)\n        self.sigmoid = nn.Sigmoid()\n        self.tanh = nn.Tanh()\n        \n    def forward(self, x):\n        x = self.casual_conv1(x)\n        x_res = x\n        \n        x = self.tanh(self.tanh_conv(x))*self.sigmoid(self.sig_conv(x))\n        x = self.batch_norm1(x)\n        x = self.casual_conv2(x)\n        x += x_res\n        x = self.batch_norm2(x)\n        del x_res\n        \n        return x\n\n\nclass FogModel(nn.Module):\n    def __init__(\n        self, \n        input_features, \n        feature_learner_kernel_size, \n        wave_block_kernel_size,\n        hidden_size,\n        num_rnn_layers\n    ):\n        \n        super().__init__()\n        self.feature_learner1 = nn.Conv1d(3, input_features, feature_learner_kernel_size, padding='same')\n        self.wavenet = nn.Sequential(*[\n            WaveBlock(input_features, 16, wave_block_kernel_size, feature_learner_kernel_size),\n            WaveBlock(feature_learner_kernel_size, 8, wave_block_kernel_size, feature_learner_kernel_size*2),\n            WaveBlock(feature_learner_kernel_size*2, 4, wave_block_kernel_size, feature_learner_kernel_size*4),\n            WaveBlock(feature_learner_kernel_size*4, 2, wave_block_kernel_size, feature_learner_kernel_size*8)\n        ])\n        self.rnn = nn.LSTM(feature_learner_kernel_size*8, hidden_size, num_rnn_layers, batch_first=False)\n        self.linear = nn.Linear(hidden_size, 3)\n        self.relu = nn.ReLU()\n        self.sigmoid = nn.Sigmoid()\n        \n    def forward(self, x):\n        batch_shape = x.shape\n        \n        x = self.feature_learner1(x)\n#         print(x.shape)\n        x = self.wavenet(x)\n#         print(x.shape)\n        x = x.permute(0, -1, 1)\n        x, _ = self.rnn(x)\n#         print(x.shape)\n        x = x.flatten(end_dim=-2)\n        x = self.linear(x)\n#         print(x.shape)\n        x = x.reshape(batch_shape)\n        \n        x = self.sigmoid(x)\n#         print(x.shape)\n        return x\n    \nmodel = FogModel(8, 8, 3, 128, 4)\nx = torch.rand((32, 3, 5000))\npred = model(x)\npred.shape","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:29:24.019178Z","iopub.execute_input":"2024-05-29T23:29:24.019740Z","iopub.status.idle":"2024-05-29T23:29:26.546830Z","shell.execute_reply.started":"2024-05-29T23:29:24.019710Z","shell.execute_reply":"2024-05-29T23:29:26.545916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nepochs = 15\nmodel = FogModel(8, 64, 3, 256, 1).to(device)\nlr = 3e-5\noptimizer = Adam(model.parameters(), lr=lr)\ncriterion = nn.BCELoss()\n","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:29:26.548245Z","iopub.execute_input":"2024-05-29T23:29:26.548540Z","iopub.status.idle":"2024-05-29T23:29:26.651863Z","shell.execute_reply.started":"2024-05-29T23:29:26.548503Z","shell.execute_reply":"2024-05-29T23:29:26.651144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_score = 0\n\nfor epoch in trange(epochs):\n    model.train()\n    y_true = []\n    y_pred = []\n    for i, (X, y) in enumerate(tqdm(train_gen)):\n        X = X.to(device).to(torch.float)\n        y = y.to(device).to(torch.float)\n        \n        optimizer.zero_grad()\n        pred = model(X)\n        loss = criterion(pred, y)\n        loss.backward()\n        optimizer.step()\n        \n        y_true.append(y.reshape(-1, 1).cpu().detach().numpy())\n        y_pred.append(pred.reshape(-1, 1).cpu().detach().numpy())\n        \n    y_true = np.concatenate(y_true)\n    y_pred = np.concatenate(y_pred)\n    train_eval_score = f1_score(y_true, (y_pred > 0.5).astype(int))\n    \n    model.eval()\n    y_true = []\n    y_pred = []\n    with torch.no_grad():\n        for j, (X, y) in enumerate(tqdm(test_gen)):\n            X = X.to(device).to(torch.float)\n            y = y.to(device).to(torch.float)\n        \n#             optimizer.zero_grad()\n            pred = model(X)\n            loss = criterion(pred, y)\n            loss = loss.cpu().detach().numpy()\n        \n            y_true.append(y.reshape(-1, 1).cpu().detach().numpy())\n            y_pred.append(pred.reshape(-1, 1).cpu().detach().numpy())\n        \n        y_true = np.concatenate(y_true)\n        y_pred = np.concatenate(y_pred)\n        test_eval_score = f1_score(y_true, (y_pred > 0.5).astype(int))\n        \n        if (test_eval_score > best_score):\n            torch.save(model.state_dict(), '/kaggle/working/best_model.h5')\n            best_score = test_eval_score\n            print(\"Model state_dict saved!\")\n            \n        print(f\"Epoch: {epoch}, BCEScore: {loss}, Test_F1_Score: {test_eval_score}, Train_F1_Score: {train_eval_score}\")\n#             loss.backward()\n#             optimizer.step()","metadata":{"execution":{"iopub.status.busy":"2024-05-29T23:34:26.653675Z","iopub.execute_input":"2024-05-29T23:34:26.654592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}