{"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":84493,"databundleVersionId":9871156,"sourceType":"competition"},{"sourceId":216145856,"sourceType":"kernelVersion"}],"dockerImageVersionId":30822,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import polars as pl\nimport random\nfrom tqdm import tqdm\nimport os\nimport pandas as pd\nimport polars as pl\nimport numpy as np\nimport gc\nfrom matplotlib import pyplot as plt\n\nfrom torch.utils.data import DataLoader, random_split, Dataset, TensorDataset\nimport torch\nimport time\nfrom torch import nn\nimport torch.optim as optim\n\nfrom torch.utils.tensorboard import SummaryWriter\n\nimport optuna","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:16:25.126952Z","iopub.execute_input":"2025-01-06T00:16:25.127222Z","iopub.status.idle":"2025-01-06T00:16:36.462328Z","shell.execute_reply.started":"2025-01-06T00:16:25.127196Z","shell.execute_reply":"2025-01-06T00:16:36.461699Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'using device: {device}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:16:36.463036Z","iopub.execute_input":"2025-01-06T00:16:36.463478Z","iopub.status.idle":"2025-01-06T00:16:36.517745Z","shell.execute_reply.started":"2025-01-06T00:16:36.463457Z","shell.execute_reply":"2025-01-06T00:16:36.516760Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# start=49\n# inputs_tensor = torch.load(f'/kaggle/input/seq-data/dataset_fea_{start}.pt', weights_only=True)\n# outputs_tensor = torch.load(f'/kaggle/input/seq-data/dataset_rsp_{start}.pt', weights_only=True)\n# weight_tensor = torch.load(f'/kaggle/input/seq-data/dataset_wgh_{start}.pt', weights_only=True)\n# times_tensor = torch.load(f'/kaggle/input/seq-data/dataset_tim_{start}.pt', weights_only=True)\n\n# num_samples = outputs_tensor.shape[0]\n# pad_time = times_tensor.shape[1]\n# pad_symbols = weight_tensor.shape[2]\n# pad_feature = inputs_tensor.shape[2] // pad_symbols\n# pad_responder = outputs_tensor.shape[2] // pad_symbols\n# pad_time = inputs_tensor.shape[1]\n# input_dim = inputs_tensor.shape[2]   # Example input feature dimension\n# output_dim = outputs_tensor.shape[2]\n# print(f'num samples per: {num_samples}, \\ninput dim: {input_dim}, \\noutput dim: {output_dim}, \\npad_symbols: {pad_symbols}, \\npad_feature: {pad_feature}, \\npad_responder: {pad_responder}, \\npad_time: {pad_time}')\n# print(f'feature: {inputs_tensor.shape}, \\nresponder: {outputs_tensor.shape}, \\ntimes: {times_tensor.shape}, \\nweight: {weight_tensor.shape}')\n\n# del inputs_tensor, outputs_tensor, weight_tensor, times_tensor\n# num_files = 19\n# inputs_tensor = torch.zeros((num_samples * num_files, pad_time, pad_symbols * pad_feature), dtype=torch.float16)\n# outputs_tensor = torch.zeros((num_samples * num_files, pad_time, pad_symbols * pad_responder), dtype=torch.float16)\n# weight_tensor = torch.zeros((num_samples * num_files, pad_time, pad_symbols), dtype=torch.float16)\n# times_tensor = torch.zeros((num_samples * num_files, pad_time), dtype=torch.int16)\n\n# for i in tqdm(range(num_files)):\n#     part = start + i * 50\n#     inputs_tensor[num_samples*i:num_samples*(i+1)] = torch.load(f'/kaggle/input/seq-data/dataset_fea_{part}.pt', weights_only=True)\n#     outputs_tensor[num_samples*i:num_samples*(i+1)] = torch.load(f'/kaggle/input/seq-data/dataset_rsp_{part}.pt', weights_only=True)\n#     weight_tensor[num_samples*i:num_samples*(i+1)] = torch.load(f'/kaggle/input/seq-data/dataset_wgh_{part}.pt', weights_only=True)\n#     times_tensor[num_samples*i:num_samples*(i+1)] = torch.load(f'/kaggle/input/seq-data/dataset_tim_{part}.pt', weights_only=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:16:36.518748Z","iopub.execute_input":"2025-01-06T00:16:36.518978Z","iopub.status.idle":"2025-01-06T00:16:36.534446Z","shell.execute_reply.started":"2025-01-06T00:16:36.518960Z","shell.execute_reply":"2025-01-06T00:16:36.533590Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"inputs_tensor = torch.load(f'/kaggle/input/seq-data/dataset_fea.pt', weights_only=True).to(device).to(torch.float32)\noutputs_tensor = torch.load(f'/kaggle/input/seq-data/dataset_rsp.pt', weights_only=True).to(device).to(torch.float32)\nweight_tensor = torch.load(f'/kaggle/input/seq-data/dataset_wgh.pt', weights_only=True).to(device).to(torch.float32)\n\nnum_samples, input_t, pad_symbols, pad_feature = inputs_tensor.shape\npad_responder = outputs_tensor.shape[-1]\npad_time = 968\nprint(f'num samples: {num_samples}, \\npad_symbols: {pad_symbols}, \\npad_feature: {pad_feature}, \\npad_responder: {pad_responder}, \\npad_time: {pad_time}')\nprint(f'feature: {inputs_tensor.shape}, \\nresponder: {outputs_tensor.shape}, \\nweight: {weight_tensor.shape}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:16:36.535363Z","iopub.execute_input":"2025-01-06T00:16:36.535693Z","iopub.status.idle":"2025-01-06T00:17:41.866225Z","shell.execute_reply.started":"2025-01-06T00:16:36.535670Z","shell.execute_reply":"2025-01-06T00:17:41.865495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mean_features = torch.mean(inputs_tensor, dim=(0,1)).to(device).to(torch.float32)\nstd_features = torch.std(inputs_tensor, dim=(0,1)).to(device).to(torch.float32)\nmean_responders = torch.mean(outputs_tensor, dim=(0,1)).to(device).to(torch.float32)\nstd_responders = torch.std(outputs_tensor, dim=(0,1)).to(device).to(torch.float32)\ntorch.cuda.empty_cache()\nprint(mean_features.shape, std_features.shape, mean_responders.shape, std_responders.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:17:41.866999Z","iopub.execute_input":"2025-01-06T00:17:41.867296Z","iopub.status.idle":"2025-01-06T00:17:41.965305Z","shell.execute_reply.started":"2025-01-06T00:17:41.867267Z","shell.execute_reply":"2025-01-06T00:17:41.964489Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n# fig, axes = plt.subplots(2, 2, figsize=(10, 10))  # 1 row, 2 columns\n\n# print(type(axes))\n# # Heatmap 1\n# im1 = axes[0,0].imshow(mean_features.numpy().reshape(pad_symbols, pad_feature), cmap='viridis')\n# axes[0,0].set_title(\"mean feature by symbol\")\n# axes[0,0].set_xlabel(\"feature\")\n# axes[0,0].set_ylabel(\"symbol\")\n# fig.colorbar(im1, ax=axes[0,0])\n\n# # Heatmap 1\n# im2 = axes[0,1].imshow(std_features.numpy().reshape(pad_symbols, pad_feature), cmap='viridis')\n# axes[0,1].set_title(\"std feature by symbol\")\n# axes[0,1].set_xlabel(\"feature\")\n# axes[0,1].set_ylabel(\"symbol\")\n# fig.colorbar(im2, ax=axes[0,1])\n\n# # Heatmap 1\n# im3 = axes[1,0].imshow(mean_responders.numpy().reshape(pad_symbols, pad_responder), cmap='viridis')\n# axes[1,0].set_title(\"mean responder by symbol\")\n# axes[1,0].set_xlabel(\"responder\")\n# axes[1,0].set_ylabel(\"symbol\")\n# fig.colorbar(im3, ax=axes[1,0])\n\n# # Heatmap 1\n# im4 = axes[1,1].imshow(std_responders.numpy().reshape(pad_symbols, pad_responder), cmap='viridis')\n# axes[1,1].set_title(\"std responder by symbol\")\n# axes[1,1].set_xlabel(\"responder\")\n# axes[1,1].set_ylabel(\"symbol\")\n# fig.colorbar(im4, ax=axes[1,1])\n\n# # Adjust layout\n# plt.tight_layout()\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:17:41.967596Z","iopub.execute_input":"2025-01-06T00:17:41.967814Z","iopub.status.idle":"2025-01-06T00:17:41.970987Z","shell.execute_reply.started":"2025-01-06T00:17:41.967796Z","shell.execute_reply":"2025-01-06T00:17:41.970339Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\n\nclass TwoDaySeqDataset(Dataset):\n    def __init__(self, inputs_tensor, outputs_tensor, weight_tensor):\n        self.inputs_tensor = inputs_tensor\n        self.outputs_tensor = outputs_tensor\n        self.weight_tensor = weight_tensor\n        \n        # Verify that all tensors have the same number of days\n        self.num_days = self.inputs_tensor.shape[0]\n        \n    def __len__(self):\n        return self.num_days - 1\n    \n    def __getitem__(self, idx):\n        # Retrieve two consecutive days for inputs\n        input_day1 = self.inputs_tensor[idx]     # Shape: [968, 42, 100]\n        input_day2 = self.inputs_tensor[idx + 1] # Shape: [968, 42, 100]\n        input_concat = torch.cat([input_day1, input_day2], dim=0)      # Shape: [968*2, 42, 100]\n        input_flat = input_concat.view(-1, pad_symbols * pad_feature)                   # Shape: [968*2, 4200]\n        \n        # Retrieve two consecutive days for outputs\n        output_day1 = self.outputs_tensor[idx]    # Shape: [968, 42, 9]\n        output_day2 = self.outputs_tensor[idx + 1]# Shape: [968, 42, 9]\n        output_concat = torch.cat([output_day1, output_day2], dim=0)  # Shape: [968*2, 42, 9]\n        output_flat = output_concat.view(-1, pad_symbols * pad_responder)                  # Shape: [968*2, 378]\n        \n        # Retrieve two consecutive days for weights\n        weight_day1 = self.weight_tensor[idx]     # Shape: [968, 42]\n        weight_day2 = self.weight_tensor[idx + 1] # Shape: [968, 42]\n        weight_concat = torch.cat([weight_day1, weight_day2], dim=0)  # Shape: [968*2, 42]\n        \n        return input_flat, output_flat, weight_concat\n\ntrain_size = int(0.95 * len(inputs_tensor))\ndataset_train = TwoDaySeqDataset(inputs_tensor[:train_size], outputs_tensor[:train_size], weight_tensor[:train_size])\ndataset_val = TwoDaySeqDataset(inputs_tensor[train_size:], outputs_tensor[train_size:], weight_tensor[train_size:])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:17:41.971961Z","iopub.execute_input":"2025-01-06T00:17:41.972250Z","iopub.status.idle":"2025-01-06T00:17:41.991817Z","shell.execute_reply.started":"2025-01-06T00:17:41.972230Z","shell.execute_reply":"2025-01-06T00:17:41.990942Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class NaiveDecoder(nn.Module):\n    def __init__(self, d_in, d_out, d_model, nhead, dim_feedforward, num_layers, x_mean, x_std, y_mean, y_std, dropout_rate):\n        super(NaiveDecoder, self).__init__()\n        self.embedding = nn.Linear(d_in, d_model)\n        self.target_embedding = nn.Linear(d_out, d_model)\n        self.embedding_dropout = nn.Dropout(p=dropout_rate)\n        self.target_embedding_dropout = nn.Dropout(p=dropout_rate)\n        self.decoder_layer = nn.TransformerDecoderLayer(\n            d_model=d_model, nhead=nhead, dim_feedforward=dim_feedforward, dropout=dropout_rate, batch_first=True\n        )\n        self.decoder = nn.TransformerDecoder(\n            self.decoder_layer, num_layers=num_layers\n        )\n        self.fc_out = nn.Linear(d_model, d_out)\n\n        self.d_in = d_in\n        self.d_out = d_out\n        self.x_mean = x_mean.flatten()\n        self.x_std = x_std.flatten()\n        self.y_mean = y_mean.flatten()\n        self.y_std = y_std.flatten()\n        t = torch.arange(0, pad_time*2).to(device).to(torch.float32)-pad_time\n        self.t_emb_x = self.time_embedding(t, self.d_in)\n        self.t_emb_y = self.time_embedding(t, self.d_out)\n\n    # def time_embedding(self, t, d_model):\n    #     t_emb = torch.zeros((t.shape[0], t.shape[1], d_model)).to(device)\n    #     for i in range(d_model):\n    #         if i % 2 == 0:\n    #             t_emb[..., i] = torch.sin(t / (10000 ** (i / d_model)))\n    #         else:\n    #             t_emb[..., i] = torch.cos(t / (10000 ** (i / d_model)))\n    #     return t_emb\n        \n    def time_embedding(self, pos, d_model=4):\n        dim_indices = torch.arange(d_model, device=device).float()  # Shape: (d_model,)\n        angle_rates = 1 / (10000 ** (dim_indices / d_model))  # Shape: (d_model,)\n        pos_expanded = pos.unsqueeze(-1)  # Shape: pos.shape + (1,)\n        angles = pos_expanded * angle_rates  # Shape: pos.shape + (d_model,)\n        PE = torch.zeros_like(angles, device=device)\n    \n        # Apply sin to even indices (0, 2, 4, ...) and cos to odd indices (1, 3, 5, ...)\n        PE[..., 0::2] = torch.sin(angles[..., 0::2])\n        PE[..., 1::2] = torch.cos(angles[..., 1::2])\n    \n        return PE\n\n    def preprocess(self, x, y):\n        x = (x-self.x_mean) / self.x_std\n        y = (y-self.y_mean) / self.y_std\n        x = torch.nan_to_num(x, nan=0.0)\n        y = torch.nan_to_num(y, nan=0.0)\n        return x+self.t_emb_x, y+self.t_emb_y\n        \n    def postprocess(self, y):\n        y = y - self.t_emb_y\n        y = y * self.y_std + self.y_mean\n        return y\n        \n    def forward(self, x, y):\n        # Preprocess\n        x, y = self.preprocess(x, y)\n        # Embed inputs\n        x_emb = self.embedding_dropout(self.embedding(x))\n        y_emb = self.target_embedding_dropout(self.target_embedding(y))\n\n        # Create mask for causal decoding\n        y_mask = torch.triu(torch.ones(y.size(1), y.size(1)), diagonal=1).bool()\n        y_mask = y_mask.to(device)\n\n        # Pass through the decoder\n        output = self.decoder(\n            tgt=y_emb, memory=x_emb, tgt_mask=y_mask\n        )\n        output = self.fc_out(output)\n        # Postprocess\n        output = self.postprocess(output)\n        return output\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:17:41.992678Z","iopub.execute_input":"2025-01-06T00:17:41.992952Z","iopub.status.idle":"2025-01-06T00:17:42.007987Z","shell.execute_reply.started":"2025-01-06T00:17:41.992927Z","shell.execute_reply":"2025-01-06T00:17:42.007127Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MSELossWeighted(nn.Module):\n    def __init__(self, weights):\n        super(MSELossWeighted, self).__init__()\n        self.weights = (weights / weights.mean()).to(device)\n    def forward(self, y_pred, y_true):\n        weighted_loss = self.weights * (y_pred - y_true) ** 2\n        return weighted_loss.mean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:17:42.008695Z","iopub.execute_input":"2025-01-06T00:17:42.008958Z","iopub.status.idle":"2025-01-06T00:17:42.022696Z","shell.execute_reply.started":"2025-01-06T00:17:42.008938Z","shell.execute_reply":"2025-01-06T00:17:42.021955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def inference(model, x, y_start, t):\n#     with torch.no_grad():\n#         b_x, s_x, d_x = x.shape\n#         b_y, s_y, d_y = y_start.shape\n        \n#         y_cum = torch.zeros((b_x, s_x, d_y), device=device)\n#         y_cum[:, :s_y] = y_start\n    \n#         steps = s_x - s_y\n#         for i in range(steps):\n#             l = s_y + i + 1\n#             # Predict next step\n#             y_cum[:, :l] = model(x[:, :l], y_cum[:, :l], t[:, :l])\n#     return y_cum\n\n# def inference(model, x, y_start, t, mean_responders):\n#     with torch.no_grad():\n#         b_x, s_x, d_x = x.shape\n#         b_y, s_y, d_y = y_start.shape\n        \n#         y_cum = torch.zeros((b_x, s_x, d_y), device=device)\n#         y_cum[:, :] = mean_responders\n#         y_cum[:, :s_y] = y_start\n    \n#         y_cum = model(x, y_cum)\n#     return y_cum\n\ndef evaluate(model, criterion, data_loader, pad_symbols, pad_responder, pad_time, mean_responders):\n    with torch.no_grad():\n        total_sum = torch.zeros(pad_symbols * pad_responder, device=device)\n        total_sqr = 0.0\n        mean_responders = mean_responders.flatten()\n        \n        total_loss = 0.0\n        total_samples = 0.0\n\n        R_sq_num = torch.zeros(pad_time, device=device)\n        R_sq_den = torch.zeros(pad_time, device=device)\n    \n        for i, (fea_gt, rsp_gt, w) in tqdm(enumerate(data_loader), total=len(data_loader)):\n            bs = fea_gt.shape[0]\n            fea_gt = fea_gt.to(torch.float32).to(device)\n            rsp_gt = rsp_gt.to(torch.float32).to(device)\n            rsp_st = rsp_gt.clone()\n            rsp_st[:, -pad_time:, :] = mean_responders\n            w = w[:, -pad_time:, :].to(torch.float32).to(device)\n            \n            rsp_pred = model(fea_gt, rsp_st)\n            \n            # First, get the stats\n            total_sum += torch.sum(rsp_pred.view(-1, pad_symbols * pad_responder), dim=0)\n            total_sqr += torch.sum((rsp_pred.view(-1, pad_symbols * pad_responder)-mean_responders)**2)\n            \n            # Second, get the loss\n            loss = criterion(rsp_pred, rsp_gt)\n            total_loss += loss.item() * bs\n            total_samples += bs\n    \n            # Third, do step wise metric\n            # rsp_pred = inference(model, fea_gt, rsp_gt[:, :-pad_time], t)\n            rsp_pred_6 = rsp_pred.view(bs, -1, pad_symbols, pad_responder)[:, -pad_time:, :, 6]\n            rsp_gt_6 = rsp_gt.view(bs, -1, pad_symbols, pad_responder)[:, -pad_time:, :, 6] # bs, 8, 42\n\n            w_inf = w[:, -pad_time:, :]\n            R_sq_num += torch.sum(w_inf * (rsp_gt_6-rsp_pred_6) * (rsp_gt_6-rsp_pred_6), dim=(0,2))\n            R_sq_den += torch.sum(w_inf * rsp_gt_6 * rsp_gt_6, dim=(0,2))\n    \n        \n        # Compute average loss and root mean squared error (RMSE)\n        mean_preds = total_sum / (total_samples*pad_time)\n        mean_offset = torch.norm(mean_preds-mean_responders).to('cpu').numpy() / np.sqrt(pad_symbols * pad_responder)\n        std_val = torch.sqrt(total_sqr / (total_samples*pad_time))\n        loss_val = total_loss / total_samples\n        R_sq = 1 - (R_sq_num/R_sq_den).to('cpu').numpy()\n\n    return loss_val, R_sq, mean_offset, std_val\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:17:42.023404Z","iopub.execute_input":"2025-01-06T00:17:42.023701Z","iopub.status.idle":"2025-01-06T00:17:42.035332Z","shell.execute_reply.started":"2025-01-06T00:17:42.023670Z","shell.execute_reply":"2025-01-06T00:17:42.034652Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_best_fit_slope(y_values):\n    x_values = np.arange(len(y_values))\n    \n    # Compute means\n    x_mean = x_values.mean()\n    y_mean = y_values.mean()\n    \n    # Compute numerator and denominator\n    numerator = np.sum((x_values - x_mean) * (y_values - y_mean))\n    denominator = np.sum((x_values - x_mean) ** 2)\n    \n    # Return the slope with a small epsilon to prevent division by zero\n    slope = (numerator + 1e-6) / (denominator + 1e-6)\n    return slope\n    \ndef save_model(model, optimizer, scheduler, global_step, acc_avg, path):\n    checkpoint = {\n        \"model_state_dict\": model.state_dict(),\n        \"optimizer_state_dict\": optimizer.state_dict(),\n        \"scheduler_state_dict\": scheduler.state_dict(),\n        \"global_step\": global_step,\n        \"acc_avg\": acc_avg,\n    }\n    os.makedirs(os.path.dirname(path), exist_ok=True)\n    torch.save(checkpoint, path)\n    print(f\"Checkpoint saved at {path}\")\n\ndef load_model(model, optimizer, scheduler, path):\n    checkpoint = torch.load(path)\n\n    # Restore model, optimizer, and scheduler states\n    model.load_state_dict(checkpoint[\"model_state_dict\"])\n    optimizer.load_state_dict(checkpoint[\"optimizer_state_dict\"])\n    scheduler.load_state_dict(checkpoint[\"scheduler_state_dict\"])\n\n    # Restore global step and best Acc\n    global_step = checkpoint[\"global_step\"]\n    acc_avg = checkpoint[\"acc_avg\"]\n\n    # Adjust scheduler's internal state to resume correctly\n    scheduler.last_epoch = global_step - 1\n\n    print(f\"Checkpoint loaded from {path}. Resuming from step {global_step}. Metric predicted {acc_avg}\")\n    return global_step, acc_avg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:17:42.036098Z","iopub.execute_input":"2025-01-06T00:17:42.036281Z","iopub.status.idle":"2025-01-06T00:17:42.051096Z","shell.execute_reply.started":"2025-01-06T00:17:42.036265Z","shell.execute_reply":"2025-01-06T00:17:42.050407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def pipeline(model, train_loader, test_loader, lr, total_steps, coeff, time_budget, writer=None, optuna_trial=None, load_ckpt=None, save_ckpt=None, dy=0.0, record_len=10, eval_steps=32, log=True):\n    torch.cuda.empty_cache()\n    # Define optimizer, scheduler, and criterion\n    optimizer = torch.optim.Adam(model.parameters(), lr=lr)\n\n    # Linear learning rate decay\n    def linear_decay(step):\n        return 1 - (step / total_steps)\n\n    scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=linear_decay)\n\n    weights_rsp = (1-coeff)/(pad_responder-1) * torch.ones(pad_responder, dtype=torch.float32)\n    weights_rsp[6] = coeff\n    weights_rsp = weights_rsp.repeat(pad_symbols)\n    criterion = MSELossWeighted(weights_rsp)\n\n    global_step = 0\n    track_loss = np.ones(record_len)\n    track_metric = np.zeros(record_len)\n    metric_slope = 0.0\n    loss_slope = 0.0\n    metric_avg = 0.0\n    flag = False\n    \n    if load_ckpt is not None:\n        global_step, metric_avg = load_model(model, optimizer, scheduler, load_ckpt)\n        track_metric[:] = metric_avg\n        \n    model.train()\n    start_time = time.time()  # Start the timer\n\n    epochs = (total_steps-global_step)//len(train_loader)\n    print(f'epochs: {epochs}')\n    for epoch in range(epochs):  # Train for a maximum of 16 epochs\n        for step, (x, y_gt, w) in enumerate(train_loader):\n            x = x.to(torch.float32).to(device)\n            y_gt = y_gt.to(torch.float32).to(device)\n            y_st = torch.clone(y_gt)\n            y_st[:, -pad_time:, :] = mean_responders.flatten()\n            \n            w = w.to(torch.float32).to(device)\n            \n            # Forward pass\n            y_hat = model(x, y_st)\n            y_hat = torch.nan_to_num(y_hat, nan=0.0, posinf=1.0, neginf=-1.0)\n            loss = criterion(y_hat, y_gt)\n            \n            optimizer.zero_grad()\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n            scheduler.step()\n\n            # Log training loss to TensorBoard\n            if writer:\n                writer.add_scalar(\"Loss/Train\", loss.item(), global_step)\n\n            global_step += 1\n\n            if global_step > 0 and global_step % eval_steps == 0:\n                model.eval()\n                loss_val, metric_vec, mean_offset, std_val = evaluate(\n                    model,\n                    criterion,\n                    test_loader,\n                    pad_symbols, \n                    pad_responder, \n                    pad_time, \n                    mean_responders\n                    )\n                metric_val = np.mean(metric_vec)\n                model.train()\n                \n                # Log evaluation metrics to TensorBoard\n                track_loss[:-1] = track_loss[1:]\n                track_loss[-1] = loss_val\n                track_metric[:-1] = track_metric[1:]\n                track_metric[-1] = metric_val\n                metric_slope = calculate_best_fit_slope(track_metric)\n                loss_slope = calculate_best_fit_slope(track_loss)\n\n                metric_avg = np.mean(track_metric)\n                if log:\n                    print()\n                    print(global_step)\n                    print(f'lr : {optimizer.param_groups[0][\"lr\"]}')\n                    print(f'Mean offset : {mean_offset}')\n                    print(f'Std offset : {std_val}')\n                    print(f'Val Loss : {loss_val}')\n                    print(f'Val Metric : {metric_vec}')\n                    print(f'Avg Metric : {metric_avg}')\n                    print(f'delta Loss : {loss_slope}')\n                    print(f'delta Metric : {metric_slope}')\n                if writer:\n                    writer.add_scalar(\"Loss/Eval\", loss_val, global_step)\n                    writer.add_scalar(\"Metric-1/Avg\", metric_avg, global_step)\n                    for i in range(len(metric_vec)):\n                        writer.add_scalar(f\"Metric{i}\", metric_vec[i], global_step)\n\n                trial_step = epoch*len(train_loader)+step\n                if trial_step // eval_steps >= record_len and time.time() - start_time > time_budget:\n                    print(\"Time budget exceeded!\")\n                    flag = True\n                    break\n                \n                if trial_step // eval_steps >= record_len and metric_slope<dy:\n                    print(f\"Turning point found : Slope {metric_slope} < {dy}\")\n                    flag = True\n                    break\n                \n                if trial_step // eval_steps >= record_len and loss_slope>(-1*dy):\n                    print(f\"Overfit : Slope {loss_slope} > -{dy}\")\n                    flag = True\n                    break\n                    \n                if trial_step // eval_steps >= record_len and optuna_trial:\n                    optuna_trial.report(metric_val, step=global_step)\n                    if optuna_trial.should_prune():\n                        print(\"Optuna pruned!\")\n                        flag = True\n                        break\n        if flag:\n            break\n    if save_ckpt is not None:\n        save_model(model, optimizer, scheduler, global_step, metric_avg, save_ckpt)\n\n    return metric_avg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:30:42.660930Z","iopub.execute_input":"2025-01-06T00:30:42.661217Z","iopub.status.idle":"2025-01-06T00:30:42.675320Z","shell.execute_reply.started":"2025-01-06T00:30:42.661194Z","shell.execute_reply":"2025-01-06T00:30:42.674517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def objective(trial):\n    # Suggest hyperparameters\n    lr = trial.suggest_float(\"lr\", 1e-5, 5e-3, log=True)\n    batch_size = trial.suggest_categorical(\"batch_size\", [1,2])\n    total_steps = trial.suggest_categorical(\"total_steps\", [4096])\n    \n    embed_dim = trial.suggest_categorical(\"embed_dim\", [512])\n    num_heads = trial.suggest_categorical(\"num_heads\", [2,4,8])\n    num_encoder_layers = trial.suggest_categorical(\"num_encoder_layers\", [4,8])\n    ff_dim = trial.suggest_categorical(\"ff_dim\", [2048])\n    dropout = trial.suggest_float(\"dropout\", 0.01, 0.1, log=True)\n\n    coeff = trial.suggest_float(\"coeff\", 0.25, 0.99, log=True)\n\n    run_name = f\"trial_{trial.number}-st{total_steps}-bs{batch_size}-lr{lr:.4f}-embed_dim{embed_dim}-num_heads{num_heads}-num_encoder_layers{num_encoder_layers}-ff_dim{ff_dim}-dropout{dropout:.4f}-coeff{coeff:.4f}\"\n    trial.set_user_attr(\"name\", run_name)\n    \n    # TensorBoard writer\n    writer = SummaryWriter(log_dir=f\"runs/{run_name}\")\n    \n    print()\n    print('='*50)\n    print(f'Trial: {run_name}')\n\n    train_loader = DataLoader(dataset_train, batch_size=batch_size, shuffle=False)\n    val_loader = DataLoader(dataset_val, batch_size=batch_size, shuffle=False)\n    \n    # Instantiate the model\n    model = NaiveDecoder(pad_feature*pad_symbols, pad_responder*pad_symbols, embed_dim, num_heads, ff_dim, num_encoder_layers, mean_features, std_features, mean_responders, std_responders, dropout)\n    model = model.to(device)\n    metric = pipeline(\n        model,\n        train_loader,\n        val_loader,\n        lr,\n        total_steps,\n        coeff,\n        900,\n        record_len=4,\n        eval_steps=64,\n        writer=writer,\n        optuna_trial=trial,\n        load_ckpt=None,\n        save_ckpt=None #f'./ckpt/{run_name}'\n        )\n\n    # Close the writer\n    writer.close()\n    # Return the best Acc for Optuna to maximize\n    return metric\n\nfrom optuna.importance import get_param_importances\n\n# Create an Optuna study\nstudy = optuna.create_study(direction=\"maximize\")  # Minimizing eval_loss\nstudy.optimize(objective, n_trials=20)  # Run 20 trials for demonstration\n\nprint(\"Best trial:\")\ntrial_opt = study.best_trial\nprint(f\"  Value: {trial_opt.value}\")\nprint(f\"  Name: {trial_opt.user_attrs['name']}\")\nprint(\"  Params: \")\nfor key, value in trial_opt.params.items():\n    print(f\"    {key}: {value}\")\n# Assuming you already have a completed study\n# Calculate parameter importances\nimportances = get_param_importances(study)\n\n# Print the parameter importances\nprint(\"Parameter importances:\")\nfor param, importance in importances.items():\n    print(f\"{param}: {importance:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:29:46.859927Z","iopub.execute_input":"2025-01-06T00:29:46.860245Z","iopub.status.idle":"2025-01-06T00:30:01.938289Z","shell.execute_reply.started":"2025-01-06T00:29:46.860216Z","shell.execute_reply":"2025-01-06T00:30:01.937198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# lr = 0.00001\n# batch_size = 64\n\n# total_steps = 20000\n\n# embed_dim = 2048\n# num_heads = 4\n# num_encoder_layers = 8\n# ff_dim = 32\n# dropout = 0.05\n\n# coeff = 0.7247692683890931\n\n# run_name = 'opt'\n\n# # lr = trial_opt.params['lr']\n# # batch_size = trial_opt.params['batch_size']\n# # embed_dim = trial_opt.params['embed_dim']\n# # num_heads = trial_opt.params['num_heads']\n# # num_encoder_layers = trial_opt.params['num_encoder_layers']\n# # ff_dim = trial_opt.params['ff_dim']\n# # dropout = trial_opt.params['dropout']\n# # ckpt = trial_opt.user_attrs['name']\n\n\n# ckpt_path = '/kaggle/input/opttest/pytorch/default/1/trial_4-st8192-bs128-lr0.0003-embed_dim512-num_heads256-num_encoder_layers4-ff_dim32-dropout0.0265-coeff0.6653'\n\n# writer = SummaryWriter(log_dir=f\"runs/{run_name}\")\n\n# train_loader = DataLoader(dataset_train, batch_size=batch_size, shuffle=True, drop_last=True)\n# val_loader = DataLoader(dataset_val, batch_size=batch_size, shuffle=False, drop_last=True)\n# # Instantiate the model\n# model = NaiveDecoder(input_dim, output_dim, embed_dim, num_heads, ff_dim, num_encoder_layers, mean_features, std_features, mean_responders, std_responders, dropout)\n# model = model.to(device)\n\n# pipeline(\n#     model,\n#     train_loader,\n#     val_loader,\n#     lr,\n#     total_steps,\n#     coeff,\n#     3600,\n#     record_len=4, # 32\n#     eval_steps=4, # 64\n#     writer=writer,\n#     load_ckpt=None,\n#     save_ckpt=None)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:17:44.395244Z","iopub.status.idle":"2025-01-06T00:17:44.395555Z","shell.execute_reply":"2025-01-06T00:17:44.395402Z"}},"outputs":[],"execution_count":null}]}