{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":41880,"databundleVersionId":5677426,"sourceType":"competition"}],"dockerImageVersionId":30823,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\nimport random\nimport time\n\nimport json\nfrom tqdm import tqdm\nimport glob\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\n\nfrom sklearn.model_selection import train_test_split, StratifiedGroupKFold\nfrom sklearn.metrics import accuracy_score, average_precision_score\n\nimport warnings\nwarnings.filterwarnings(action='ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:03:35.374710Z","iopub.execute_input":"2024-12-28T06:03:35.375021Z","iopub.status.idle":"2024-12-28T06:03:40.055771Z","shell.execute_reply.started":"2024-12-28T06:03:35.374993Z","shell.execute_reply":"2024-12-28T06:03:40.054863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n    train_dir1 = \"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/defog\"\n    train_dir2 = \"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/tdcsfog\"\n\n    batch_size = 1024\n    window_size = 32\n    window_future = 8\n    window_past = window_size - window_future\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\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    \n    feature_list = ['AccV', 'AccML', 'AccAP']\n    label_list = ['StartHesitation', 'Turn', 'Walking']\n    \n    \ncfg = Config()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:03:43.483359Z","iopub.execute_input":"2024-12-28T06:03:43.483872Z","iopub.status.idle":"2024-12-28T06:03:43.537873Z","shell.execute_reply.started":"2024-12-28T06:03:43.483843Z","shell.execute_reply":"2024-12-28T06:03:43.536878Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Analysis of positive instances in each fold of our CV folds\n\nn1_sum = []\nn2_sum = []\nn3_sum = []\ncount = []\n\n# Here I am using the metadata file available during training. Since the code will run again during submission, if \n# I used the usual file from the competition folder, it would have been updated with the test files too.\nmetadata = pd.read_csv(\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/tdcsfog_metadata.csv\")\n\nfor f in tqdm(metadata['Id']):\n    fpath = f\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/tdcsfog/{f}.csv\"\n    df = pd.read_csv(fpath)\n    \n    n1_sum.append(np.sum(df['StartHesitation']))\n    n2_sum.append(np.sum(df['Turn']))\n    n3_sum.append(np.sum(df['Walking']))\n    count.append(len(df))\n    \nprint(f\"32 files have positive values in all 3 classes\")\n\nmetadata['n1_sum'] = n1_sum\nmetadata['n2_sum'] = n2_sum\nmetadata['n3_sum'] = n3_sum\nmetadata['count'] = count\n\nsgkf = StratifiedGroupKFold(n_splits=5, random_state=42, shuffle=True)\nfor i, (train_index, valid_index) in enumerate(sgkf.split(X=metadata['Id'], y=[1]*len(metadata), groups=metadata['Subject'])):\n    print(f\"Fold = {i}\")\n    train_ids = metadata.loc[train_index, 'Id']\n    valid_ids = metadata.loc[valid_index, 'Id']\n    \n    print(f\"Length of Train = {len(train_index)}, Length of Valid = {len(valid_index)}\")\n    n1_sum = metadata.loc[train_index, 'n1_sum'].sum()\n    n2_sum = metadata.loc[train_index, 'n2_sum'].sum()\n    n3_sum = metadata.loc[train_index, 'n3_sum'].sum()\n    print(f\"Train classes: {n1_sum:,}, {n2_sum:,}, {n3_sum:,}\")\n    \n    n1_sum = metadata.loc[valid_index, 'n1_sum'].sum()\n    n2_sum = metadata.loc[valid_index, 'n2_sum'].sum()\n    n3_sum = metadata.loc[valid_index, 'n3_sum'].sum()\n    print(f\"Valid classes: {n1_sum:,}, {n2_sum:,}, {n3_sum:,}\")\n    \n# FOLD 2 is the most well balanced\n# The actual train-test split (based on Fold 2)\n\nmetadata = pd.read_csv(\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/tdcsfog_metadata.csv\")\nsgkf = StratifiedGroupKFold(n_splits=5, random_state=42, shuffle=True)\nfor i, (train_index, valid_index) in enumerate(sgkf.split(X=metadata['Id'], y=[1]*len(metadata), groups=metadata['Subject'])):\n    if i != 2:\n        continue\n    print(f\"Fold = {i}\")\n    train_ids = metadata.loc[train_index, 'Id']\n    valid_ids = metadata.loc[valid_index, 'Id']\n    print(f\"Length of Train = {len(train_ids)}, Length of Valid = {len(valid_ids)}\")\n    \n    if i == 2:\n        break\n        \ntrain_fpaths_tdcs = [f\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/tdcsfog/{_id}.csv\" for _id in train_ids]\nvalid_fpaths_tdcs = [f\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/tdcsfog/{_id}.csv\" for _id in valid_ids]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:03:46.368743Z","iopub.execute_input":"2024-12-28T06:03:46.369030Z","iopub.status.idle":"2024-12-28T06:04:05.988970Z","shell.execute_reply.started":"2024-12-28T06:03:46.369007Z","shell.execute_reply":"2024-12-28T06:04:05.988138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Analysis of positive instances in each fold of our CV folds\n\nn1_sum = []\nn2_sum = []\nn3_sum = []\ncount = []\n\n# Here I am using the metadata file available during training. Since the code will run again during submission, if \n# I used the usual file from the competition folder, it would have been updated with the test files too.\nmetadata = pd.read_csv(\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/defog_metadata.csv\")\nmetadata['n1_sum'] = 0\nmetadata['n2_sum'] = 0\nmetadata['n3_sum'] = 0\nmetadata['count'] = 0\n\nfor f in tqdm(metadata['Id']):\n    fpath = f\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/defog/{f}.csv\"\n    if os.path.exists(fpath) == False:\n        continue\n        \n    df = pd.read_csv(fpath)\n    metadata.loc[metadata['Id'] == f, 'n1_sum'] = np.sum(df['StartHesitation'])\n    metadata.loc[metadata['Id'] == f, 'n2_sum'] = np.sum(df['Turn'])\n    metadata.loc[metadata['Id'] == f, 'n3_sum'] = np.sum(df['Walking'])\n    metadata.loc[metadata['Id'] == f, 'count'] = len(df)\n    \nmetadata = metadata[metadata['count'] > 0].reset_index()\n\nsgkf = StratifiedGroupKFold(n_splits=5, random_state=42, shuffle=True)\nfor i, (train_index, valid_index) in enumerate(sgkf.split(X=metadata['Id'], y=[1]*len(metadata), groups=metadata['Subject'])):\n    print(f\"Fold = {i}\")\n    train_ids = metadata.loc[train_index, 'Id']\n    valid_ids = metadata.loc[valid_index, 'Id']\n    \n    print(f\"Length of Train = {len(train_index)}, Length of Valid = {len(valid_index)}\")\n    n1_sum = metadata.loc[train_index, 'n1_sum'].sum()\n    n2_sum = metadata.loc[train_index, 'n2_sum'].sum()\n    n3_sum = metadata.loc[train_index, 'n3_sum'].sum()\n    print(f\"Train classes: {n1_sum:,}, {n2_sum:,}, {n3_sum:,}\")\n    \n    n1_sum = metadata.loc[valid_index, 'n1_sum'].sum()\n    n2_sum = metadata.loc[valid_index, 'n2_sum'].sum()\n    n3_sum = metadata.loc[valid_index, 'n3_sum'].sum()\n    print(f\"Valid classes: {n1_sum:,}, {n2_sum:,}, {n3_sum:,}\")\n    \n# FOLD 2 is the most well balanced\n# The actual train-test split (based on Fold 2)\n\nsgkf = StratifiedGroupKFold(n_splits=5, random_state=42, shuffle=True)\nfor i, (train_index, valid_index) in enumerate(sgkf.split(X=metadata['Id'], y=[1]*len(metadata), groups=metadata['Subject'])):\n    if i != 1:\n        continue\n    print(f\"Fold = {i}\")\n    train_ids = metadata.loc[train_index, 'Id']\n    valid_ids = metadata.loc[valid_index, 'Id']\n    print(f\"Length of Train = {len(train_ids)}, Length of Valid = {len(valid_ids)}\")\n    \n    if i == 2:\n        break\n        \ntrain_fpaths_de = [f\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/defog/{_id}.csv\" for _id in train_ids]\nvalid_fpaths_de = [f\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/defog/{_id}.csv\" for _id in valid_ids]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:04:13.288943Z","iopub.execute_input":"2024-12-28T06:04:13.289230Z","iopub.status.idle":"2024-12-28T06:04:35.570763Z","shell.execute_reply.started":"2024-12-28T06:04:13.289206Z","shell.execute_reply":"2024-12-28T06:04:35.569928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_fpaths = [(f, 'de') for f in train_fpaths_de] + [(f, 'tdcs') for f in train_fpaths_tdcs]\nvalid_fpaths = [(f, 'de') for f in valid_fpaths_de] + [(f, 'tdcs') for f in valid_fpaths_tdcs]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:04:38.325781Z","iopub.execute_input":"2024-12-28T06:04:38.326069Z","iopub.status.idle":"2024-12-28T06:04:38.330093Z","shell.execute_reply.started":"2024-12-28T06:04:38.326047Z","shell.execute_reply":"2024-12-28T06:04:38.329361Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FOGDataset(Dataset):\n    def __init__(self, fpaths, scale=9.806, split=\"train\"):\n        super(FOGDataset, self).__init__()\n        tm = time.time()\n        self.split = split\n        self.scale = scale\n        \n        self.fpaths = fpaths\n        self.dfs = [self.read(f[0], f[1]) for f in fpaths]\n        self.f_ids = [os.path.basename(f[0])[:-4] for f in self.fpaths]\n        \n        self.end_indices = []\n        self.shapes = []\n        _length = 0\n        for df in self.dfs:\n            self.shapes.append(df.shape[0])\n            _length += df.shape[0]\n            self.end_indices.append(_length)\n        \n        self.dfs = np.concatenate(self.dfs, axis=0).astype(np.float16)\n        self.length = self.dfs.shape[0]\n        \n        shape1 = self.dfs.shape[1]\n        \n        self.dfs = np.concatenate([np.zeros((cfg.wx*cfg.window_past, shape1)), self.dfs, np.zeros((cfg.wx*cfg.window_future, shape1))], axis=0)\n        print(f\"Dataset initialized in {time.time() - tm} secs!\")\n        gc.collect()\n        \n    def read(self, f, _type):\n        df = pd.read_csv(f)\n        if self.split == \"test\":\n            return np.array(df)\n        \n        if _type ==\"tdcs\":\n            df['Valid'] = 1\n            df['Task'] = 1\n            df['tdcs'] = 1\n        else:\n            df['tdcs'] = 0\n        \n        return np.array(df)\n            \n    def __getitem__(self, index):\n        if self.split == \"train\":\n            row_idx = random.randint(0, self.length-1) + cfg.wx*cfg.window_past\n        elif self.split == \"test\":\n            for i,e in enumerate(self.end_indices):\n                if index >= e:\n                    continue\n                df_idx = i\n                break\n\n            row_idx_true = self.shapes[df_idx] - (self.end_indices[df_idx] - index)\n            _id = self.f_ids[df_idx] + \"_\" + str(row_idx_true)\n            row_idx = index + cfg.wx*cfg.window_past\n        else:\n            row_idx = index + cfg.wx*cfg.window_past\n            \n        #scale = 9.806 if self.dfs[row_idx, -1] == 1 else 1.0\n        x = self.dfs[row_idx - cfg.wx*cfg.window_past : row_idx + cfg.wx*cfg.window_future, 1:4]\n        x = x[::cfg.wx, :][::-1, :]\n        x = torch.tensor(x.astype('float'))#/scale\n        \n        t = self.dfs[row_idx, -3]*self.dfs[row_idx, -2]\n        \n        if self.split == \"test\":\n            return _id, x, t\n        \n        y = self.dfs[row_idx, 4:7].astype('float')\n        y = torch.tensor(y)\n        \n        return x, y, t\n    \n    def __len__(self):\n        # return self.length\n        if self.split == \"train\":\n            return 5_000_000\n        return self.length","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:04:40.768946Z","iopub.execute_input":"2024-12-28T06:04:40.769228Z","iopub.status.idle":"2024-12-28T06:04:40.780375Z","shell.execute_reply.started":"2024-12-28T06:04:40.769207Z","shell.execute_reply":"2024-12-28T06:04:40.779586Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:04:44.449826Z","iopub.execute_input":"2024-12-28T06:04:44.450130Z","iopub.status.idle":"2024-12-28T06:04:44.583656Z","shell.execute_reply.started":"2024-12-28T06:04:44.450102Z","shell.execute_reply":"2024-12-28T06:04:44.582763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass AttentionBlock(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.attention = nn.Sequential(\n            nn.Linear(dim, dim),\n            nn.Tanh(),\n            nn.Linear(dim, 1),\n            nn.Softmax(dim=1)\n        )\n        \n    def forward(self, x):\n        # x shape: (batch_size, seq_len, dim)\n        attn_weights = self.attention(x)  # (batch_size, seq_len, 1)\n        attended = torch.sum(x * attn_weights, dim=1)  # (batch_size, dim)\n        return attended, attn_weights\n\nclass ConvBlock(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv = nn.Sequential(\n            nn.Conv1d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm1d(out_channels),\n            nn.ReLU(),\n            nn.Conv1d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm1d(out_channels),\n            nn.ReLU()\n        )\n        \n    def forward(self, x):\n        return self.conv(x)\n\nclass FOGModel(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n        self.window_size = cfg.window_size\n        self.n_channels = len(cfg.feature_list)\n        \n        # Convolutional feature extraction\n        self.conv_blocks = nn.ModuleList([\n            ConvBlock(self.n_channels, 32),\n            ConvBlock(32, 64),\n            ConvBlock(64, 128)\n        ])\n        \n        # Temporal attention\n        self.attention = AttentionBlock(128)\n        \n        # Feature processing\n        self.feature_net = nn.Sequential(\n            nn.Linear(128, cfg.model_hidden),\n            nn.LayerNorm(cfg.model_hidden),\n            nn.ReLU(),\n            nn.Dropout(cfg.model_dropout),\n            nn.Linear(cfg.model_hidden, cfg.model_hidden),\n            nn.LayerNorm(cfg.model_hidden),\n            nn.ReLU(),\n            nn.Dropout(cfg.model_dropout)\n        )\n        \n        # Class-specific heads\n        self.heads = nn.ModuleList([\n            nn.Sequential(\n                nn.Linear(cfg.model_hidden, cfg.model_hidden // 2),\n                nn.ReLU(),\n                nn.Dropout(cfg.model_dropout/2),\n                nn.Linear(cfg.model_hidden // 2, 1)\n            ) for _ in range(len(cfg.label_list))\n        ])\n        \n        # Initialize weights\n        self._init_weights()\n        \n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.xavier_uniform_(m.weight)\n                if m.bias is not None:\n                    nn.init.zeros_(m.bias)\n            elif isinstance(m, nn.Conv1d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.zeros_(m.bias)\n            elif isinstance(m, (nn.BatchNorm1d, nn.LayerNorm)):\n                nn.init.ones_(m.weight)\n                nn.init.zeros_(m.bias)\n    \n    def forward(self, x):\n        # x shape: (batch_size, window_size, channels)\n        batch_size = x.size(0)\n        \n        # Reshape for 1D convolution: (batch_size, channels, window_size)\n        x = x.transpose(1, 2)\n        \n        # Apply convolutional blocks\n        for conv in self.conv_blocks:\n            x = conv(x)\n            \n        # Reshape for attention: (batch_size, window_size, channels)\n        x = x.transpose(1, 2)\n        \n        # Apply attention\n        x, _ = self.attention(x)\n        \n        # Process features\n        x = self.feature_net(x)\n        \n        # Apply class-specific heads\n        outputs = []\n        for head in self.heads:\n            outputs.append(head(x))\n            \n        # Combine outputs\n        x = torch.cat(outputs, dim=1)\n        return x\n\ndef get_fog_model(cfg):\n    model = FOGModel(cfg)\n    \n    # Use custom weight initialization\n    def init_weights(m):\n        if isinstance(m, nn.Linear):\n            torch.nn.init.xavier_uniform_(m.weight)\n            if m.bias is not None:\n                torch.nn.init.zeros_(m.bias)\n    \n    model.apply(init_weights)\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:24:11.737057Z","iopub.execute_input":"2024-12-28T06:24:11.737354Z","iopub.status.idle":"2024-12-28T06:24:11.750938Z","shell.execute_reply.started":"2024-12-28T06:24:11.737328Z","shell.execute_reply":"2024-12-28T06:24:11.750144Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def count_parameters(model):\n    return sum(p.numel() for p in model.parameters() if p.requires_grad)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:11:35.768788Z","iopub.execute_input":"2024-12-28T06:11:35.769069Z","iopub.status.idle":"2024-12-28T06:11:35.773037Z","shell.execute_reply.started":"2024-12-28T06:11:35.769045Z","shell.execute_reply":"2024-12-28T06:11:35.772122Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, criterion, cfg, scheduler):\n    loss_sum = 0.\n    acc_sum = 0.\n    class_acc_sum = torch.zeros(len(cfg.label_list)).to(cfg.device)\n    scaler = GradScaler()\n    \n    model.train()\n    for x, y, t in tqdm(loader):\n        x = x.to(cfg.device).float()\n        y = y.to(cfg.device).float()\n        t = t.to(cfg.device).float()\n        \n        optimizer.zero_grad()\n        \n        with torch.cuda.amp.autocast():\n            y_pred = model(x)\n            loss = criterion(y_pred, y)\n            loss = torch.mean(loss * t.unsqueeze(-1), dim=1)\n            \n            t_sum = torch.sum(t)\n            if t_sum > 0:\n                loss = torch.sum(loss) / t_sum\n            else:\n                loss = torch.sum(loss) * 0.\n        \n        # Compute accuracy\n        with torch.no_grad():\n            acc = ((y_pred >= 0.5).float() == y).float().mean(dim=0)\n            acc_sum += acc.mean().item()\n            class_acc_sum += acc\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        scheduler.step()  # Move the scheduler step here\n        \n        loss_sum += loss.item()\n    \n    avg_loss = loss_sum / len(loader)\n    avg_acc = acc_sum / len(loader)\n    avg_class_acc = class_acc_sum / len(loader)\n    \n    return avg_loss, avg_acc, avg_class_acc\n\n\ndef validation_one_epoch(model, loader, criterion, cfg):\n    loss_sum = 0.\n    acc_sum = 0.\n    class_acc_sum = torch.zeros(len(cfg.label_list)).to(cfg.device)\n    y_true_epoch = []\n    y_pred_epoch = []\n    t_valid_epoch = []\n    \n    model.eval()\n    with torch.no_grad():\n        for x, y, t in tqdm(loader):\n            x = x.to(cfg.device).float()\n            y = y.to(cfg.device).float()\n            t = t.to(cfg.device).float()\n            \n            y_pred = model(x)\n            loss = criterion(y_pred, y)\n            loss = torch.mean(loss * t.unsqueeze(-1), dim=1)\n            \n            t_sum = torch.sum(t)\n            if t_sum > 0:\n                loss = torch.sum(loss) / t_sum\n            else:\n                loss = torch.sum(loss) * 0.\n            \n            acc = ((y_pred >= 0.5).float() == y).float().mean(dim=0)\n            acc_sum += acc.mean().item()\n            class_acc_sum += acc\n            \n            loss_sum += loss.item()\n            y_true_epoch.append(y.cpu().numpy())\n            y_pred_epoch.append(y_pred.cpu().numpy())\n            t_valid_epoch.append(t.cpu().numpy())\n    \n    avg_loss = loss_sum / len(loader)\n    avg_acc = acc_sum / len(loader)\n    avg_class_acc = class_acc_sum / len(loader)\n    \n    # Calculate metrics only for valid samples\n    y_true_epoch = np.concatenate(y_true_epoch, axis=0)\n    y_pred_epoch = np.concatenate(y_pred_epoch, axis=0)\n    t_valid_epoch = np.concatenate(t_valid_epoch, axis=0)\n    \n    mask = t_valid_epoch > 0\n    y_true_epoch = y_true_epoch[mask]\n    y_pred_epoch = y_pred_epoch[mask]\n    \n    # Calculate AP score for each class\n    scores = [\n        average_precision_score(y_true_epoch[:, i], y_pred_epoch[:, i])\n        for i in range(len(cfg.label_list))\n    ]\n    mean_score = np.mean(scores)\n    \n    return avg_loss, avg_acc, avg_class_acc, mean_score, scores","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:28:27.855803Z","iopub.execute_input":"2024-12-28T06:28:27.856121Z","iopub.status.idle":"2024-12-28T06:28:27.867896Z","shell.execute_reply.started":"2024-12-28T06:28:27.856093Z","shell.execute_reply":"2024-12-28T06:28:27.867071Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nprint(f\"Number of parameters in model - {count_parameters(model):,}\")\n\ntrain_dataset = FOGDataset(train_fpaths, split=\"train\")\nvalid_dataset = FOGDataset(valid_fpaths, split=\"valid\")\nprint(f\"lengths of datasets: train - {len(train_dataset)}, valid - {len(valid_dataset)}\")\n\ntrain_loader = DataLoader(train_dataset, batch_size=cfg.batch_size, num_workers=5, shuffle=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=cfg.batch_size, num_workers=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:24:45.342682Z","iopub.execute_input":"2024-12-28T06:24:45.343000Z","iopub.status.idle":"2024-12-28T06:25:30.349331Z","shell.execute_reply.started":"2024-12-28T06:24:45.342972Z","shell.execute_reply":"2024-12-28T06:25:30.348446Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:06:29.113183Z","iopub.execute_input":"2024-12-28T06:06:29.113456Z","iopub.status.idle":"2024-12-28T06:06:29.150616Z","shell.execute_reply.started":"2024-12-28T06:06:29.113435Z","shell.execute_reply":"2024-12-28T06:06:29.149764Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y, t = next(iter(train_loader))\nprint(\"Input shape:\", x.shape)\nprint(\"Label shape:\", y.shape)\nprint(\"Weight shape:\", t.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:13:09.517215Z","iopub.execute_input":"2024-12-28T06:13:09.517504Z","iopub.status.idle":"2024-12-28T06:13:10.376035Z","shell.execute_reply.started":"2024-12-28T06:13:09.517481Z","shell.execute_reply":"2024-12-28T06:13:10.375143Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize model\nmodel = get_fog_model(cfg).to(cfg.device)\n\n# Use AdamW with weight decay\noptimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=0.01)\n\n# Use OneCycleLR scheduler for better convergence\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=cfg.lr,\n    epochs=cfg.num_epochs,\n    steps_per_epoch=len(train_loader),\n    pct_start=0.1,\n    anneal_strategy='cos'\n)\n\n# Loss function with positive class weights to handle imbalance\npos_weight = torch.tensor([2.0, 2.0, 2.0]).to(cfg.device)  # Adjust these weights based on class distribution\ncriterion = nn.BCEWithLogitsLoss(reduction='none', pos_weight=pos_weight).to(cfg.device)\n\n# Training loop remains the same, but add scheduler step\nfor epoch in range(cfg.num_epochs):\n    print(f\"Epoch: {epoch}\")\n    train_loss, train_acc, train_class_acc = train_one_epoch(\n        model, train_loader, optimizer, criterion, cfg, scheduler\n    )\n    # ... rest of the training loop\n    print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}\")\n    print(f\"Class-wise Train Acc: {[f'{acc:.3f}' for acc in train_class_acc.tolist()]}\")\n    \n    # Validate\n    val_loss, val_acc, val_class_acc, mean_score, class_scores = validation_one_epoch(\n        model, valid_loader, criterion, cfg\n    )\n    print(f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}, Score: {mean_score:.3f}\")\n    \n    if mean_score > max_score:\n        max_score = mean_score\n        torch.save(model.state_dict(), \"best_model_state.h5\")\n        print(\"Saving Model ...\")\n    print(\"=\"*50)\n\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T06:28:31.669768Z","iopub.execute_input":"2024-12-28T06:28:31.670067Z","iopub.status.idle":"2024-12-28T07:22:03.462536Z","shell.execute_reply.started":"2024-12-28T06:28:31.670044Z","shell.execute_reply":"2024-12-28T07:22:03.461805Z"}},"outputs":[],"execution_count":null}]}