{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":19990,"databundleVersionId":1472735,"sourceType":"competition"},{"sourceId":1551779,"sourceType":"datasetVersion","datasetId":884260},{"sourceId":43970252,"sourceType":"kernelVersion"},{"sourceId":238011,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":203275,"modelId":225008}],"dockerImageVersionId":30009,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Lyft: Complete train and prediction pipeline\n\nThis notebook is modified slightly from https://www.kaggle.com/huanvo/lyft-complete-train-and-prediction-pipeline\n\nChange log:\n- v13 (2020-11-19): simplify prediction and faster inference (perhaps) \n- v12 (2020-10-29): faster prediction - reduce batch size due to machine limit on Kaggle\n- v10 (2020-10-29): faster prediction\n- v9 (2020-10-28): get some real training\n- v8 (2020-10-25): skip computing loss to see if prediction will be faster.","metadata":{}},{"cell_type":"markdown","source":"# Environment setup\n\n - Please add [pestipeti/lyft-l5kit-unofficial-fix](https://www.kaggle.com/pestipeti/lyft-l5kit-unofficial-fix) as utility script.\n    - Official utility script \"[philculliton/kaggle-l5kit](https://www.kaggle.com/mathurinache/kaggle-l5kit)\" does not work with pytorch GPU.\n\nClick \"File\" botton on top-left, and choose \"Add utility script\". For the pop-up search window, you need to remove \"Your Work\" filter, and search [pestipeti/lyft-l5kit-unofficial-fix](https://www.kaggle.com/pestipeti/lyft-l5kit-unofficial-fix) on top-right of the search window. Then you can add the kaggle-l5kit utility script. It is much faster to do this rather than !pip install l5kit every time you run the notebook. \n\nIf successful, you can see \"usr/lib/lyft-l5kit-unofficial-fix\" is added to the \"Data\" section of this kernel page on right side of the kernel.\n\n- Also please add [pretrained baseline model](https://www.kaggle.com/huanvo/lyft-pretrained-model-hv)\n\nClick on the button \"Add data\" in the \"Data\" section and search for lyft-pretrained-model-hv. If you find the model useful, please upvote it as well.  ","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","collapsed":true,"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":false,"jupyter":{"outputs_hidden":true}}},{"cell_type":"code","source":"!pip install -U torch==1.7.0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:07:41.430262Z","iopub.execute_input":"2025-01-22T04:07:41.430482Z","iopub.status.idle":"2025-01-22T04:08:45.815645Z","shell.execute_reply.started":"2025-01-22T04:07:41.430445Z","shell.execute_reply":"2025-01-22T04:08:45.814539Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import Dict\n\nfrom tempfile import gettempdir\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch import nn, optim\nfrom torch.utils.data import DataLoader\nimport torchvision\nfrom torchvision.models.resnet import resnet50, resnet18, resnet34, resnet101\nfrom tqdm import tqdm\n\nimport l5kit\nfrom l5kit.configs import load_config_data\nfrom l5kit.data import LocalDataManager, ChunkedDataset\nfrom l5kit.dataset import AgentDataset, EgoDataset\nfrom l5kit.rasterization import build_rasterizer\nfrom l5kit.evaluation import write_pred_csv, compute_metrics_csv, read_gt_csv, create_chopped_dataset\nfrom l5kit.evaluation.chop_dataset import MIN_FUTURE_STEPS\nfrom l5kit.evaluation.metrics import neg_multi_log_likelihood, time_displace\nfrom l5kit.geometry import transform_points\nfrom l5kit.visualization import PREDICTED_POINTS_COLOR, TARGET_POINTS_COLOR, draw_trajectory\nfrom prettytable import PrettyTable\nfrom pathlib import Path\n\nimport matplotlib.pyplot as plt\n\nimport os\nimport random\nimport time\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nfrom IPython.display import display\nfrom tqdm import tqdm_notebook\nimport gc, psutil\n\nprint(l5kit.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:08:45.817468Z","iopub.execute_input":"2025-01-22T04:08:45.817756Z","iopub.status.idle":"2025-01-22T04:08:49.372358Z","shell.execute_reply.started":"2025-01-22T04:08:45.817720Z","shell.execute_reply":"2025-01-22T04:08:49.371504Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Memory measurement\ndef memory(verbose=True):\n    mem = psutil.virtual_memory()\n    gb = 1024*1024*1024\n    if verbose:\n        print('Physical memory:',\n              '%.2f GB (used),'%((mem.total - mem.available) / gb),\n              '%.2f GB (available)'%((mem.available) / gb), '/',\n              '%.2f GB'%(mem.total / gb))\n    return (mem.total - mem.available) / gb\n\ndef gc_memory(verbose=True):\n    m = gc.collect()\n    if verbose:\n        print('GC:', m, end=' | ')\n        memory()\n\nmemory();","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:09:57.401986Z","iopub.execute_input":"2025-01-22T04:09:57.402271Z","iopub.status.idle":"2025-01-22T04:09:57.412278Z","shell.execute_reply.started":"2025-01-22T04:09:57.402248Z","shell.execute_reply":"2025-01-22T04:09:57.411501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\nset_seed(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:10:01.265177Z","iopub.execute_input":"2025-01-22T04:10:01.265456Z","iopub.status.idle":"2025-01-22T04:10:01.271648Z","shell.execute_reply.started":"2025-01-22T04:10:01.265432Z","shell.execute_reply":"2025-01-22T04:10:01.270806Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Configs","metadata":{}},{"cell_type":"code","source":"# # --- Lyft configs ---\n# cfg = {\n#     'format_version': 4,\n#     'data_path': '/kaggle/input/lyft-motion-prediction-autonomous-vehicles',\n#     'model_params': {\n#         'model_architecture': 'resnet34',\n#         'history_num_frames': 10,\n#         'history_step_size': 1,\n#         'history_delta_time': 0.1,\n#         'future_num_frames': 50,\n#         'future_step_size': 1,\n#         'future_delta_time': 0.1,\n#         'model_name': \"model_resnet34_output\",\n#         'lr': 1e-3,\n#         'weight_path': '/kaggle/input/lyft-pretrained-model-hv/model_multi_update_lyft_public.pth',\n#         'train': False,\n#         'predict': True,\n#     },\n#     'raster_params': {\n#         'raster_size': [224, 224],\n#         'pixel_size': [0.5, 0.5],\n#         'ego_center': [0.25, 0.5],\n#         'map_type': 'py_semantic',\n#         'satellite_map_key': 'aerial_map/aerial_map.png',\n#         'semantic_map_key': 'semantic_map/semantic_map.pb',\n#         'dataset_meta_key': 'meta.json',\n#         'filter_agents_threshold': 0.5,\n#     },\n#     'train_data_loader': {\n#         'key': 'scenes/train.zarr',\n#         'batch_size': 16,\n#         'shuffle': True,\n#         'num_workers': 4,\n#     },    \n#     'test_data_loader': {\n#         'key': 'scenes/test.zarr',\n#         'batch_size': 128,\n#         'shuffle': False,\n#         'num_workers': 4,\n#     },\n#     'train_params': {\n# #         'steps': 100,\n# #         'update_steps': 10,\n# #         'checkpoint_steps': 50,\n#         'steps': 12000,\n#         'update_steps': 100,\n#         'checkpoint_steps': 3000,\n#     }\n# }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:47:40.968163Z","iopub.execute_input":"2025-01-22T04:47:40.968467Z","iopub.status.idle":"2025-01-22T04:47:40.972552Z","shell.execute_reply.started":"2025-01-22T04:47:40.968444Z","shell.execute_reply":"2025-01-22T04:47:40.971792Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Lyft configs ---\ncfg = {\n    'format_version': 4,\n    'data_path': '/kaggle/input/lyft-motion-prediction-autonomous-vehicles',\n    'model_params': {\n        'model_architecture': 'resnet34',\n        'history_num_frames': 10,\n        'history_step_size': 1,\n        'history_delta_time': 0.1,\n        'future_num_frames': 50,\n        'future_step_size': 1,\n        'future_delta_time': 0.1,\n        'model_name': \"model_resnet34_output\",\n        'lr': 1e-3,\n        'train': False,\n        'predict': True,\n    },\n    'raster_params': {\n        'raster_size': [224, 224],\n        'pixel_size': [0.5, 0.5],\n        'ego_center': [0.25, 0.5],\n        'map_type': 'py_semantic',\n        'satellite_map_key': 'aerial_map/aerial_map.png',\n        'semantic_map_key': 'semantic_map/semantic_map.pb',\n        'dataset_meta_key': 'meta.json',\n        'filter_agents_threshold': 0.5,\n    },\n    'train_data_loader': {\n        'key': 'scenes/train.zarr',\n        'batch_size': 16,\n        'shuffle': True,\n        'num_workers': 4,\n    },    \n    'test_data_loader': {\n        'key': 'scenes/test.zarr',\n        'batch_size': 128,\n        'shuffle': False,\n        'num_workers': 4,\n    },\n    'train_params': {\n        'steps': 12000,\n        'update_steps': 100,\n        'checkpoint_steps': 3000,\n    }\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:47:47.982130Z","iopub.execute_input":"2025-01-22T04:47:47.982412Z","iopub.status.idle":"2025-01-22T04:47:47.989198Z","shell.execute_reply.started":"2025-01-22T04:47:47.982388Z","shell.execute_reply":"2025-01-22T04:47:47.988455Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Couple of things to note:\n\n - **model_architecture:** you can put 'resnet18', 'resnet34' or 'resnet50'. For the pretrained model we use resnet18 so we need to use 'resnet18' in the config.\n - **weight_path:** path to the pretrained model. If you don't have a pretrained model and want to train from scratch, put **weight_path** = False. \n - **model_name:** the name of the model that will be saved as output, this is only when **train**= True.\n - **train:** True if you want to train the model.\n - **predict:** True if you want to predict and submit to Kaggle.\n - **lr:** learning rate of the model.\n - **raster_size:** specify the size of the image, the default is [224,224]. Increase **raster_size** can improve the score. However the training time will be significantly longer. \n - **batch_size:** number of samples for one forward pass\n - **steps:** number of batches of data that the model will be trained on. (note this is not epoch)\n - **checkpoint_every_n_steps:** the model will be saved at every n steps, again change this number as to how you want to keep track of the model.\n \n \n Note (Louis): The original pretrained model doesn't save the state of optimizer, so continute training doesn't work too well.","metadata":{}},{"cell_type":"markdown","source":"# Load the train and test datasets","metadata":{}},{"cell_type":"code","source":"# set env variable for data\nDIR_INPUT = cfg[\"data_path\"]\nos.environ[\"L5KIT_DATA_FOLDER\"] = DIR_INPUT\ndm = LocalDataManager()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:10:10.442296Z","iopub.execute_input":"2025-01-22T04:10:10.442577Z","iopub.status.idle":"2025-01-22T04:10:10.446377Z","shell.execute_reply.started":"2025-01-22T04:10:10.442552Z","shell.execute_reply":"2025-01-22T04:10:10.445457Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n# Build rasterizer\nrasterizer = build_rasterizer(cfg, dm)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:10:13.770215Z","iopub.execute_input":"2025-01-22T04:10:13.770492Z","iopub.status.idle":"2025-01-22T04:10:17.721793Z","shell.execute_reply.started":"2025-01-22T04:10:13.770469Z","shell.execute_reply":"2025-01-22T04:10:17.721041Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n# Train dataset\ntrain_cfg = cfg[\"train_data_loader\"]\ntrain_zarr = ChunkedDataset(dm.require(train_cfg[\"key\"])).open(cached=False)  # to prevent run out of memory\ntrain_dataset = AgentDataset(cfg, train_zarr, rasterizer)\ntrain_dataloader = DataLoader(train_dataset, shuffle=train_cfg[\"shuffle\"],\n                              batch_size=train_cfg[\"batch_size\"], num_workers=train_cfg[\"num_workers\"],\n                              pin_memory=True, prefetch_factor=8,\n                             )\nprint(train_dataset)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:10:19.857895Z","iopub.execute_input":"2025-01-22T04:10:19.858173Z","iopub.status.idle":"2025-01-22T04:12:14.263448Z","shell.execute_reply.started":"2025-01-22T04:10:19.858150Z","shell.execute_reply":"2025-01-22T04:12:14.262635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n# Test dataset\ntest_cfg = cfg[\"test_data_loader\"]\ntest_zarr = ChunkedDataset(dm.require(test_cfg[\"key\"])).open(cached=False)  # to prevent run out of memory\ntest_mask = np.load(f\"{DIR_INPUT}/scenes/mask.npz\")[\"arr_0\"]\ntest_dataset = AgentDataset(cfg, test_zarr, rasterizer, agents_mask=test_mask)\ntest_dataloader = DataLoader(test_dataset, shuffle=test_cfg[\"shuffle\"],\n                             batch_size=test_cfg[\"batch_size\"], num_workers=test_cfg[\"num_workers\"],\n                             pin_memory=False, prefetch_factor=4,\n                            )\nprint(test_dataset)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:12:14.265411Z","iopub.execute_input":"2025-01-22T04:12:14.265637Z","iopub.status.idle":"2025-01-22T04:12:15.411537Z","shell.execute_reply.started":"2025-01-22T04:12:14.265614Z","shell.execute_reply":"2025-01-22T04:12:15.410794Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(train_dataset), len(test_dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:12:15.412930Z","iopub.execute_input":"2025-01-22T04:12:15.413277Z","iopub.status.idle":"2025-01-22T04:12:15.418178Z","shell.execute_reply.started":"2025-01-22T04:12:15.413238Z","shell.execute_reply":"2025-01-22T04:12:15.417186Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Note that the train set size is much bigger than our steps * batch_size. So we will not even finish 1 epoch of training here.","metadata":{}},{"cell_type":"markdown","source":"# Simple visualization\n\nLet us visualize how an input to the model looks like.","metadata":{}},{"cell_type":"code","source":"def visualize_trajectory(dataset, index, title=\"target_positions movement with draw_trajectory\"):\n    data = dataset[index]\n    im = data[\"image\"].transpose(1, 2, 0)\n    im = dataset.rasterizer.to_rgb(im)\n    target_positions_pixels = transform_points(data[\"target_positions\"] + data[\"centroid\"][:2], data[\"world_to_image\"])\n    draw_trajectory(im, target_positions_pixels, TARGET_POINTS_COLOR, radius=1, yaws=data[\"target_yaws\"])\n\n    plt.title(title)\n    plt.imshow(im, origin='lower')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:12:15.419608Z","iopub.execute_input":"2025-01-22T04:12:15.419960Z","iopub.status.idle":"2025-01-22T04:12:15.428542Z","shell.execute_reply.started":"2025-01-22T04:12:15.419926Z","shell.execute_reply":"2025-01-22T04:12:15.427949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"i_plot = 4\n\nplt.figure(figsize=(8, 6))\nvisualize_trajectory(train_dataset, index=i_plot)\n\nplt.figure(figsize=(15, 15))\nfor i in range(25):\n    plt.subplot(5, 5, i+1).set_title(f'{i}')\n    plt.imshow(train_dataset[i_plot]['image'][i])\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:12:15.430259Z","iopub.execute_input":"2025-01-22T04:12:15.430528Z","iopub.status.idle":"2025-01-22T04:12:19.635190Z","shell.execute_reply.started":"2025-01-22T04:12:15.430472Z","shell.execute_reply":"2025-01-22T04:12:19.634475Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loss function\n\nFor this competition it is important to use the correct loss function when train the model. Our goal is to predict three possible paths together with the confidence score, so we need to use the loss function that takes that into account, simply using RMSE will not lead to an accurate model. More information about the loss function can be found here [negative log likelihood](https://github.com/lyft/l5kit/blob/master/competition.md).","metadata":{}},{"cell_type":"code","source":"# --- Function utils ---\n# Original code from https://github.com/lyft/l5kit/blob/20ab033c01610d711c3d36e1963ecec86e8b85b6/l5kit/l5kit/evaluation/metrics.py\nimport numpy as np\n\nimport torch\nfrom torch import Tensor\n\n\ndef pytorch_neg_multi_log_likelihood_batch(\n    gt: Tensor, pred: Tensor, confidences: Tensor, avails: Tensor\n) -> Tensor:\n    \"\"\"\n    Compute a negative log-likelihood for the multi-modal scenario.\n    log-sum-exp trick is used here to avoid underflow and overflow, For more information about it see:\n    https://en.wikipedia.org/wiki/LogSumExp#log-sum-exp_trick_for_log-domain_calculations\n    https://timvieira.github.io/blog/post/2014/02/11/exp-normalize-trick/\n    https://leimao.github.io/blog/LogSumExp/\n    Args:\n        gt (Tensor): array of shape (bs)x(time)x(2D coords)\n        pred (Tensor): array of shape (bs)x(modes)x(time)x(2D coords)\n        confidences (Tensor): array of shape (bs)x(modes) with a confidence for each mode in each sample\n        avails (Tensor): array of shape (bs)x(time) with the availability for each gt timestep\n    Returns:\n        Tensor: negative log-likelihood for this example, a single float number\n    \"\"\"\n    assert len(pred.shape) == 4, f\"expected 3D (MxTxC) array for pred, got {pred.shape}\"\n    batch_size, num_modes, future_len, num_coords = pred.shape\n\n    assert gt.shape == (batch_size, future_len, num_coords), f\"expected 2D (Time x Coords) array for gt, got {gt.shape}\"\n    assert confidences.shape == (batch_size, num_modes), f\"expected 1D (Modes) array for gt, got {confidences.shape}\"\n    assert torch.allclose(torch.sum(confidences, dim=1), confidences.new_ones((batch_size,))), \"confidences should sum to 1\"\n    assert avails.shape == (batch_size, future_len), f\"expected 1D (Time) array for gt, got {avails.shape}\"\n    # assert all data are valid\n    assert torch.isfinite(pred).all(), \"invalid value found in pred\"\n    assert torch.isfinite(gt).all(), \"invalid value found in gt\"\n    assert torch.isfinite(confidences).all(), \"invalid value found in confidences\"\n    assert torch.isfinite(avails).all(), \"invalid value found in avails\"\n\n    # convert to (batch_size, num_modes, future_len, num_coords)\n    gt = torch.unsqueeze(gt, 1)  # add modes\n    avails = avails[:, None, :, None]  # add modes and cords\n\n    # error (batch_size, num_modes, future_len)\n    error = torch.sum(((gt - pred) * avails) ** 2, dim=-1)  # reduce coords and use availability\n\n    with np.errstate(divide=\"ignore\"):  # when confidence is 0 log goes to -inf, but we're fine with it\n        # error (batch_size, num_modes)\n        error = torch.log(confidences) - 0.5 * torch.sum(error, dim=-1)  # reduce time\n\n    # use max aggregator on modes for numerical stability\n    # error (batch_size, num_modes)\n    max_value, _ = error.max(dim=1, keepdim=True)  # error are negative at this point, so max() gives the minimum one\n    error = -torch.log(torch.sum(torch.exp(error - max_value), dim=-1, keepdim=True)) - max_value  # reduce modes\n    # print(\"error\", error)\n    return torch.mean(error)\n\n\ndef pytorch_neg_multi_log_likelihood_single(\n    gt: Tensor, pred: Tensor, avails: Tensor\n) -> Tensor:\n    \"\"\"\n\n    Args:\n        gt (Tensor): array of shape (bs)x(time)x(2D coords)\n        pred (Tensor): array of shape (bs)x(time)x(2D coords)\n        avails (Tensor): array of shape (bs)x(time) with the availability for each gt timestep\n    Returns:\n        Tensor: negative log-likelihood for this example, a single float number\n    \"\"\"\n    # pred (bs)x(time)x(2D coords) --> (bs)x(mode=1)x(time)x(2D coords)\n    # create confidence (bs)x(mode=1)\n    batch_size, future_len, num_coords = pred.shape\n    confidences = pred.new_ones((batch_size, 1))\n    return pytorch_neg_multi_log_likelihood_batch(gt, pred.unsqueeze(1), confidences, avails)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:12:36.742779Z","iopub.execute_input":"2025-01-22T04:12:36.743079Z","iopub.status.idle":"2025-01-22T04:12:36.755876Z","shell.execute_reply.started":"2025-01-22T04:12:36.743054Z","shell.execute_reply":"2025-01-22T04:12:36.754942Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model\nNext we define the baseline model. Note that this model will return three possible trajectories together with confidence score for each trajectory.","metadata":{}},{"cell_type":"code","source":"# class LyftMultiModel(nn.Module):\n#     def __init__(self, cfg: Dict, num_modes=3):\n#         super().__init__()\n\n#         architecture = cfg[\"model_params\"][\"model_architecture\"]\n#         backbone = eval(architecture)(pretrained=True, progress=True)\n#         self.backbone = backbone\n\n#         num_history_channels = (cfg[\"model_params\"][\"history_num_frames\"] + 1) * 2\n#         num_in_channels = 3 + num_history_channels\n\n#         self.backbone.conv1 = nn.Conv2d(\n#             num_in_channels,\n#             self.backbone.conv1.out_channels,\n#             kernel_size=self.backbone.conv1.kernel_size,\n#             stride=self.backbone.conv1.stride,\n#             padding=self.backbone.conv1.padding,\n#             bias=False,\n#         )\n\n#         # This is 512 for resnet18 and resnet34\n#         # And it is 2048 for the other resnets        \n#         if architecture == \"resnet50\":\n#             backbone_out_features = 2048\n#         else:\n#             backbone_out_features = 512\n\n#         # X, Y coords for the future positions (output shape: batch_sizex50x2)\n#         self.future_len = cfg[\"model_params\"][\"future_num_frames\"]\n#         num_targets = 2 * self.future_len\n\n#         # You can add more layers here.\n#         self.head = nn.Sequential(\n#             # nn.Dropout(0.2),\n#             nn.Linear(in_features=backbone_out_features, out_features=4096),\n#         )\n\n#         self.num_preds = num_targets * num_modes\n#         self.num_modes = num_modes\n\n#         self.logit = nn.Linear(4096, out_features=self.num_preds + num_modes)\n\n#     def forward(self, x):\n#         x = self.backbone.conv1(x)\n#         x = self.backbone.bn1(x)\n#         x = self.backbone.relu(x)\n#         x = self.backbone.maxpool(x)\n\n#         x = self.backbone.layer1(x)\n#         x = self.backbone.layer2(x)\n#         x = self.backbone.layer3(x)\n#         x = self.backbone.layer4(x)\n\n#         x = self.backbone.avgpool(x)\n#         x = torch.flatten(x, 1)\n\n#         x = self.head(x)\n#         x = self.logit(x)\n\n#         # pred (batch_size)x(modes)x(time)x(2D coords)\n#         # confidences (batch_size)x(modes)\n#         bs, _ = x.shape\n#         pred, confidences = torch.split(x, self.num_preds, dim=1)\n#         pred = pred.view(bs, self.num_modes, self.future_len, 2)\n#         assert confidences.shape == (bs, self.num_modes)\n#         confidences = torch.softmax(confidences, dim=1)\n#         return pred, confidences","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:38:04.381233Z","iopub.execute_input":"2025-01-22T04:38:04.381629Z","iopub.status.idle":"2025-01-22T04:38:04.386159Z","shell.execute_reply.started":"2025-01-22T04:38:04.381597Z","shell.execute_reply":"2025-01-22T04:38:04.385290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install transformers","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:42:43.534246Z","iopub.execute_input":"2025-01-22T04:42:43.534580Z","iopub.status.idle":"2025-01-22T04:42:49.058708Z","shell.execute_reply.started":"2025-01-22T04:42:43.534538Z","shell.execute_reply":"2025-01-22T04:42:49.057889Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass PatchEmbedding(nn.Module):\n    \"\"\"\n    Split image into patches and embed them.\n    \"\"\"\n    def __init__(self, img_size=224, patch_size=16, in_channels=3, embed_dim=768):\n        super().__init__()\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.n_patches = (img_size // patch_size) ** 2\n\n        self.proj = nn.Conv2d(\n            in_channels=in_channels,\n            out_channels=embed_dim,\n            kernel_size=patch_size,\n            stride=patch_size,\n        )\n\n    def forward(self, x):\n        \"\"\"\n        Input shape: (batch_size, in_channels, img_size, img_size)\n        Output shape: (batch_size, n_patches, embed_dim)\n        \"\"\"\n        x = self.proj(x)  # (batch_size, embed_dim, n_patches_h, n_patches_w)\n        x = x.flatten(2)  # (batch_size, embed_dim, n_patches)\n        x = x.transpose(1, 2)  # (batch_size, n_patches, embed_dim)\n        return x\n\n\nclass MultiHeadAttention(nn.Module):\n    \"\"\"\n    Multi-head self-attention mechanism.\n    \"\"\"\n    def __init__(self, embed_dim=768, num_heads=12, dropout=0.1):\n        super().__init__()\n        self.embed_dim = embed_dim\n        self.num_heads = num_heads\n        self.head_dim = embed_dim // num_heads\n\n        assert self.head_dim * num_heads == embed_dim, \"embed_dim must be divisible by num_heads\"\n\n        self.qkv = nn.Linear(embed_dim, embed_dim * 3)\n        self.dropout = nn.Dropout(dropout)\n        self.proj = nn.Linear(embed_dim, embed_dim)\n\n    def forward(self, x):\n        \"\"\"\n        Input shape: (batch_size, n_patches, embed_dim)\n        Output shape: (batch_size, n_patches, embed_dim)\n        \"\"\"\n        batch_size, n_patches, embed_dim = x.shape\n\n        # Linear projection for Q, K, V\n        qkv = self.qkv(x).reshape(batch_size, n_patches, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv[0], qkv[1], qkv[2]\n\n        # Scaled dot-product attention\n        attn = (q @ k.transpose(-2, -1)) * (self.head_dim ** -0.5)\n        attn = F.softmax(attn, dim=-1)\n        attn = self.dropout(attn)\n\n        # Combine heads\n        x = (attn @ v).transpose(1, 2).reshape(batch_size, n_patches, embed_dim)\n        x = self.proj(x)\n        return x\n\n\nclass TransformerBlock(nn.Module):\n    \"\"\"\n    A single transformer block.\n    \"\"\"\n    def __init__(self, embed_dim=768, num_heads=12, dropout=0.1):\n        super().__init__()\n        self.norm1 = nn.LayerNorm(embed_dim)\n        self.attn = MultiHeadAttention(embed_dim, num_heads, dropout)\n        self.norm2 = nn.LayerNorm(embed_dim)\n        self.mlp = nn.Sequential(\n            nn.Linear(embed_dim, embed_dim * 4),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(embed_dim * 4, embed_dim),\n            nn.Dropout(dropout),\n        )\n\n    def forward(self, x):\n        \"\"\"\n        Input shape: (batch_size, n_patches, embed_dim)\n        Output shape: (batch_size, n_patches, embed_dim)\n        \"\"\"\n        x = x + self.attn(self.norm1(x))\n        x = x + self.mlp(self.norm2(x))\n        return x\n\n\nclass VisionTransformer(nn.Module):\n    \"\"\"\n    Simplified Vision Transformer (ViT) backbone.\n    \"\"\"\n    def __init__(self, img_size=224, patch_size=16, in_channels=3, embed_dim=768, depth=12, num_heads=12, dropout=0.1):\n        super().__init__()\n        self.patch_embed = PatchEmbedding(img_size, patch_size, in_channels, embed_dim)\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))\n        self.pos_embed = nn.Parameter(torch.zeros(1, self.patch_embed.n_patches + 1, embed_dim))\n        self.dropout = nn.Dropout(dropout)\n        self.blocks = nn.ModuleList([TransformerBlock(embed_dim, num_heads, dropout) for _ in range(depth)])\n        self.norm = nn.LayerNorm(embed_dim)\n\n    def forward(self, x):\n        \"\"\"\n        Input shape: (batch_size, in_channels, img_size, img_size)\n        Output shape: (batch_size, embed_dim)\n        \"\"\"\n        batch_size = x.shape[0]\n\n        # Patch embedding\n        x = self.patch_embed(x)  # (batch_size, n_patches, embed_dim)\n\n        # Add class token and position embedding\n        cls_token = self.cls_token.expand(batch_size, -1, -1)  # (batch_size, 1, embed_dim)\n        x = torch.cat((cls_token, x), dim=1)  # (batch_size, n_patches + 1, embed_dim)\n        x = x + self.pos_embed\n        x = self.dropout(x)\n\n        # Transformer blocks\n        for block in self.blocks:\n            x = block(x)\n\n        # Extract class token\n        x = self.norm(x)\n        return x[:, 0]  # (batch_size, embed_dim)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:44:15.228647Z","iopub.execute_input":"2025-01-22T04:44:15.229101Z","iopub.status.idle":"2025-01-22T04:44:15.256798Z","shell.execute_reply.started":"2025-01-22T04:44:15.229061Z","shell.execute_reply":"2025-01-22T04:44:15.255863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LyftMultiModel(nn.Module):\n    def __init__(self, cfg: Dict, num_modes=3):\n        super().__init__()\n\n        # Calculate input channels\n        num_history_channels = (cfg[\"model_params\"][\"history_num_frames\"] + 1) * 2\n        num_in_channels = 3 + num_history_channels\n\n        # Initialize Vision Transformer backbone\n        self.backbone = VisionTransformer(\n            img_size=224,  # Adjust based on your input size\n            patch_size=16,  # Adjust based on your needs\n            in_channels=num_in_channels,\n            embed_dim=768,  # ViT base model size\n            depth=12,  # Number of transformer blocks\n            num_heads=12,  # Number of attention heads\n            dropout=0.1,  # Dropout rate\n        )\n\n        # Future trajectory length\n        self.future_len = cfg[\"model_params\"][\"future_num_frames\"]\n        num_targets = 2 * self.future_len\n\n        # Head layers\n        self.head = nn.Sequential(\n            nn.Linear(in_features=768, out_features=4096),\n            nn.ReLU(),  # Add activation if needed\n            # nn.Dropout(0.2),  # Optional dropout\n        )\n\n        # Number of predictions and modes\n        self.num_preds = num_targets * num_modes\n        self.num_modes = num_modes\n\n        # Final logit layer\n        self.logit = nn.Linear(4096, out_features=self.num_preds + num_modes)\n\n    def forward(self, x):\n        # Forward pass through Vision Transformer\n        x = self.backbone(x)\n\n        # Forward pass through head\n        x = self.head(x)\n\n        # Final logits\n        x = self.logit(x)\n\n        # Split into predictions and confidences\n        bs, _ = x.shape\n        pred, confidences = torch.split(x, self.num_preds, dim=1)\n        pred = pred.view(bs, self.num_modes, self.future_len, 2)\n        assert confidences.shape == (bs, self.num_modes)\n        confidences = torch.softmax(confidences, dim=1)\n\n        return pred, confidences","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:44:31.947764Z","iopub.execute_input":"2025-01-22T04:44:31.948080Z","iopub.status.idle":"2025-01-22T04:44:31.957735Z","shell.execute_reply.started":"2025-01-22T04:44:31.948051Z","shell.execute_reply":"2025-01-22T04:44:31.956964Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # from transformers import ViTModel\n\n# class LyftMultiModel(nn.Module):\n#     def __init__(self, cfg: Dict, num_modes=3):\n#         super().__init__()\n\n#         # Load pre-trained ViT model\n#         vit_path = \"/kaggle/input/pretrained_vit/pytorch/default/1\"\n#         self.backbone = ViTModel.from_pretrained(vit_path)\n\n#         # Freeze the backbone if needed (optional)\n#         for param in self.backbone.parameters():\n#             param.requires_grad = False  # Set to True if you want to fine-tune\n\n#         # Calculate input channels\n#         num_history_channels = (cfg[\"model_params\"][\"history_num_frames\"] + 1) * 2\n#         num_in_channels = 3 + num_history_channels\n\n#         # Modify the first layer of ViT to accept the correct number of input channels\n#         original_embedding = self.backbone.embeddings.patch_embeddings\n#         self.backbone.embeddings.patch_embeddings = nn.Conv2d(\n#             in_channels=num_in_channels,\n#             out_channels=original_embedding.proj.out_channels,\n#             kernel_size=original_embedding.proj.kernel_size,\n#             stride=original_embedding.proj.stride,\n#             padding=original_embedding.proj.padding,\n#             bias=False,\n#         )\n\n#         # Get the output feature size of ViT\n#         backbone_out_features = self.backbone.config.hidden_size  # Typically 768 for base ViT\n\n#         # Future trajectory length\n#         self.future_len = cfg[\"model_params\"][\"future_num_frames\"]\n#         num_targets = 2 * self.future_len\n\n#         # Head layers\n#         self.head = nn.Sequential(\n#             nn.Linear(in_features=backbone_out_features, out_features=4096),\n#             nn.ReLU(),  # Add activation if needed\n#             # nn.Dropout(0.2),  # Optional dropout\n#         )\n\n#         # Number of predictions and modes\n#         self.num_preds = num_targets * num_modes\n#         self.num_modes = num_modes\n\n#         # Final logit layer\n#         self.logit = nn.Linear(4096, out_features=self.num_preds + num_modes)\n\n#     def forward(self, x):\n#         # Forward pass through ViT\n#         x = self.backbone.embeddings.patch_embeddings(x)\n#         x = x.flatten(2).transpose(1, 2)  # Reshape for ViT\n#         x = self.backbone.encoder(x)[0]  # Extract features from ViT\n\n#         # Global average pooling\n#         x = x.mean(dim=1)  # Average over sequence length\n\n#         # Forward pass through head\n#         x = self.head(x)\n\n#         # Final logits\n#         x = self.logit(x)\n\n#         # Split into predictions and confidences\n#         bs, _ = x.shape\n#         pred, confidences = torch.split(x, self.num_preds, dim=1)\n#         pred = pred.view(bs, self.num_modes, self.future_len, 2)\n#         assert confidences.shape == (bs, self.num_modes)\n#         confidences = torch.softmax(confidences, dim=1)\n\n#         return pred, confidences","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:44:20.604053Z","iopub.execute_input":"2025-01-22T04:44:20.604341Z","iopub.status.idle":"2025-01-22T04:44:20.608765Z","shell.execute_reply.started":"2025-01-22T04:44:20.604316Z","shell.execute_reply":"2025-01-22T04:44:20.607805Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Add `compute_loss` flag to skip computing loss during prediction.","metadata":{}},{"cell_type":"code","source":"def forward(data, model, device, criterion=pytorch_neg_multi_log_likelihood_batch, compute_loss=True):\n    inputs = data[\"image\"].to(device)\n    target_availabilities = data[\"target_availabilities\"].to(device)\n    targets = data[\"target_positions\"].to(device)\n    # Forward pass\n    preds, confidences = model(inputs)\n    # skip compute loss if we are doing prediction\n    loss = criterion(targets, preds, confidences, target_availabilities) if compute_loss else 0\n    return loss, preds, confidences","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:44:35.547713Z","iopub.execute_input":"2025-01-22T04:44:35.548051Z","iopub.status.idle":"2025-01-22T04:44:35.553407Z","shell.execute_reply.started":"2025-01-22T04:44:35.548023Z","shell.execute_reply":"2025-01-22T04:44:35.552488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==== INIT MODEL =================\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n# 初始化模型\nmodel = LyftMultiModel(cfg)\n\n# 将模型移动到设备\nmodel.to(device)\n\n# 初始化优化器\noptimizer = optim.Adam(model.parameters(), lr=cfg[\"model_params\"][\"lr\"])\n\nprint(f'device {device}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:48:06.122175Z","iopub.execute_input":"2025-01-22T04:48:06.122467Z","iopub.status.idle":"2025-01-22T04:48:06.776838Z","shell.execute_reply.started":"2025-01-22T04:48:06.122443Z","shell.execute_reply":"2025-01-22T04:48:06.775792Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Now let us initialize the model and load the pretrained weights. Note that since the pretrained model was trained on GPU, you also need to enable GPU when running this notebook.","metadata":{}},{"cell_type":"code","source":"def init_weights(m):\n    if isinstance(m, nn.Conv2d) or isinstance(m, nn.Linear):\n        nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n        if m.bias is not None:\n            nn.init.constant_(m.bias, 0)\n    elif isinstance(m, nn.BatchNorm2d):\n        nn.init.constant_(m.weight, 1)\n        nn.init.constant_(m.bias, 0)\n\n%%time\n# ==== INIT MODEL =================\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n# 初始化模型\nmodel = LyftMultiModel(cfg)\n\n# 手动初始化权重\nmodel.apply(init_weights)\n\n# 将模型移动到设备\nmodel.to(device)\n\n# 初始化优化器\noptimizer = optim.Adam(model.parameters(), lr=cfg[\"model_params\"][\"lr\"])\n\nprint(f'device {device}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:49:19.754828Z","iopub.execute_input":"2025-01-22T04:49:19.755157Z","iopub.status.idle":"2025-01-22T04:49:19.765838Z","shell.execute_reply.started":"2025-01-22T04:49:19.755129Z","shell.execute_reply":"2025-01-22T04:49:19.764829Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:49:27.305826Z","iopub.execute_input":"2025-01-22T04:49:27.306151Z","iopub.status.idle":"2025-01-22T04:49:27.311673Z","shell.execute_reply.started":"2025-01-22T04:49:27.306121Z","shell.execute_reply":"2025-01-22T04:49:27.310776Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training loop\nNext let us implement the training loop, when the **train** parameter is set to True. ","metadata":{}},{"cell_type":"code","source":"%%time\nif cfg[\"model_params\"][\"train\"]:\n    tr_it = iter(train_dataloader)\n    n_steps = cfg[\"train_params\"][\"steps\"]\n    progress_bar = tqdm_notebook(range(1, 1 + n_steps), mininterval=5.)\n    losses = []\n    iterations = []\n    metrics = []\n    memorys = []\n    times = []\n    model_name = cfg[\"model_params\"][\"model_name\"]\n    update_steps = cfg['train_params']['update_steps']\n    checkpoint_steps = cfg['train_params']['checkpoint_steps']\n    t_start = time.time()\n    torch.set_grad_enabled(True)\n        \n    for i in progress_bar:\n        try:\n            data = next(tr_it)\n        except StopIteration:\n            tr_it = iter(train_dataloader)\n            data = next(tr_it)\n        model.train()   # somehow we need this is ever batch or it perform very bad (not sure why)\n        loss, _, _ = forward(data, model, device)\n\n        # Backward pass\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        loss_v = loss.item()\n        losses.append(loss_v)\n        \n        if i % update_steps == 0:\n            mean_losses = np.mean(losses)\n            timespent = (time.time() - t_start) / 60\n            print('i: %5d'%i,\n                  'loss: %10.5f'%loss_v, 'loss(avg): %10.5f'%mean_losses, \n                  '%.2fmins'%timespent, end=' | ')\n            mem = memory()\n            if i % checkpoint_steps == 0:\n                torch.save(model.state_dict(), f'{model_name}_{i}.pth')\n                torch.save(optimizer.state_dict(), f'{model_name}_optimizer_{i}.pth')\n            iterations.append(i)\n            metrics.append(mean_losses)\n            memorys.append(mem)\n            times.append(timespent)\n\n    torch.save(model.state_dict(), f'{model_name}_final.pth')\n    torch.save(optimizer.state_dict(), f'{model_name}_optimizer_final.pth')\n    results = pd.DataFrame({\n        'iterations': iterations, \n        'metrics (avg)': metrics,\n        'elapsed_time (mins)': times,\n        'memory (GB)': memorys,\n    })\n    results.to_csv(f'train_metrics_{model_name}_{n_steps}.csv', index=False)\n    print(f'Total training time is {(time.time() - t_start) / 60} mins')\n    memory()\n    display(results)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:49:35.721307Z","iopub.execute_input":"2025-01-22T04:49:35.721615Z","iopub.status.idle":"2025-01-22T04:49:35.729808Z","shell.execute_reply.started":"2025-01-22T04:49:35.721585Z","shell.execute_reply":"2025-01-22T04:49:35.728877Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cfg[\"model_params\"][\"train\"]:\n    plt.figure(figsize=(12, 4))\n    plt.plot(results['iterations'], results['metrics (avg)'])\n    plt.xlabel('steps'); plt.ylabel('metrics (avg)')\n    plt.grid(); plt.show()\n\n    plt.figure(figsize=(12, 4))\n    plt.plot(results['iterations'], results['memory (GB)'])\n    plt.xlabel('steps'); plt.ylabel('memory (GB)')\n    plt.grid(); plt.show()\n\n    plt.figure(figsize=(12, 4))\n    plt.plot(results['iterations'], results['elapsed_time (mins)'])\n    plt.xlabel('steps'); plt.ylabel('elapsed_time (mins)')\n    plt.grid(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:49:39.590480Z","iopub.execute_input":"2025-01-22T04:49:39.590807Z","iopub.status.idle":"2025-01-22T04:49:39.597349Z","shell.execute_reply.started":"2025-01-22T04:49:39.590779Z","shell.execute_reply":"2025-01-22T04:49:39.596477Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prediction\n\nFinally we implement the inference to submit to Kaggle when **predict** param is set to True.","metadata":{}},{"cell_type":"code","source":"print('Number of batches for predictoin:', int(np.ceil(len(test_dataset) / cfg['test_data_loader']['batch_size'])))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:49:42.288021Z","iopub.execute_input":"2025-01-22T04:49:42.288304Z","iopub.status.idle":"2025-01-22T04:49:42.293233Z","shell.execute_reply.started":"2025-01-22T04:49:42.288281Z","shell.execute_reply":"2025-01-22T04:49:42.292494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def init_weights(m):\n    if isinstance(m, nn.Conv2d) or isinstance(m, nn.Linear):\n        nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n        if m.bias is not None:\n            nn.init.constant_(m.bias, 0)\n    elif isinstance(m, nn.BatchNorm2d):\n        nn.init.constant_(m.weight, 1)\n        nn.init.constant_(m.bias, 0)\n\n# 初始化模型\nmodel = LyftMultiModel(cfg)\n\n# 手动初始化权重\nmodel.apply(init_weights)\n\n# 将模型移动到设备\nmodel.to(device)\n\n# 初始化优化器\noptimizer = optim.Adam(model.parameters(), lr=cfg[\"model_params\"][\"lr\"])\n\nprint(f'device {device}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:53:38.840159Z","iopub.execute_input":"2025-01-22T04:53:38.840489Z","iopub.status.idle":"2025-01-22T04:53:40.559208Z","shell.execute_reply.started":"2025-01-22T04:53:38.840460Z","shell.execute_reply":"2025-01-22T04:53:40.558357Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\nif cfg[\"model_params\"][\"predict\"]:\n    \n    model.eval()\n    torch.set_grad_enabled(False)\n\n    # store information for evaluation\n    future_coords_offsets_pd = []\n    timestamps = []\n    confidences_list = []\n    agent_ids = []\n    memorys_pred = []\n    t0 = time.time()\n    times_pred = []\n    iterations_pred = []\n\n    for i, data in enumerate(tqdm_notebook(test_dataloader, mininterval=5.)):\n        \n        _, preds, confidences = forward(data, model, device, compute_loss=False)\n        \n        # rotation (batch) x (2d world coord) x (2d agent coord)\n        # preds (batch) x (mode) x (time) x (2d agent coord)\n        rotation = data[\"world_from_agent\"][:, :2, :2].float().to(device)\n        preds = torch.sum(preds[:, :, :, None, :] * rotation[:, None, None, :, :], dim=-1).cpu().numpy()\n        # same as: preds = torch.einsum('bmti,bji->bmtj', preds, rotation).cpu().numpy()\n    \n        future_coords_offsets_pd.append(preds.copy())\n        confidences_list.append(confidences.cpu().numpy().copy())\n        timestamps.append(data[\"timestamp\"].numpy().copy())\n        agent_ids.append(data[\"track_id\"].numpy().copy()) \n        \n        if i%50 == 0:\n            t = ((time.time() - t0) / 60)\n            print('%4d'%i, '%6.2fmins'%t, end=' | ')\n            mem = memory()\n            iterations_pred.append(i)\n            memorys_pred.append(mem)\n            times_pred.append(t)\n    print('Total timespent: %6.2fmins'%((time.time() - t0) / 60))\n    memory()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T04:53:42.673174Z","iopub.execute_input":"2025-01-22T04:53:42.673463Z","iopub.status.idle":"2025-01-22T06:02:21.876337Z","shell.execute_reply.started":"2025-01-22T04:53:42.673439Z","shell.execute_reply":"2025-01-22T06:02:21.875458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(12, 4))\nplt.plot(iterations_pred, memorys_pred)\nplt.xlabel('steps'); plt.ylabel('memory (GB)')\nplt.grid(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T06:02:21.878869Z","iopub.execute_input":"2025-01-22T06:02:21.879101Z","iopub.status.idle":"2025-01-22T06:02:22.023968Z","shell.execute_reply.started":"2025-01-22T06:02:21.879071Z","shell.execute_reply":"2025-01-22T06:02:22.023203Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(12, 4))\nplt.plot(iterations_pred, times_pred)\nplt.xlabel('steps'); plt.ylabel('elapsed_time (mins)')\nplt.grid(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T06:02:22.025127Z","iopub.execute_input":"2025-01-22T06:02:22.025458Z","iopub.status.idle":"2025-01-22T06:02:22.158438Z","shell.execute_reply.started":"2025-01-22T06:02:22.025422Z","shell.execute_reply":"2025-01-22T06:02:22.157680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n# create submission to submit to Kaggle\npred_path = 'submission.csv'\nwrite_pred_csv(\n    pred_path,\n    timestamps=np.concatenate(timestamps),\n    track_ids=np.concatenate(agent_ids),\n    coords=np.concatenate(future_coords_offsets_pd),\n    confs=np.concatenate(confidences_list),\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T06:02:22.159305Z","iopub.execute_input":"2025-01-22T06:02:22.159515Z","iopub.status.idle":"2025-01-22T06:02:44.174659Z","shell.execute_reply.started":"2025-01-22T06:02:22.159493Z","shell.execute_reply":"2025-01-22T06:02:44.173945Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Examine submission","metadata":{}},{"cell_type":"code","source":"df_sub = pd.read_csv(pred_path)\ndf_sub = df_sub.set_index(['timestamp', 'track_id'])\ndisplay(df_sub)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T06:02:44.176932Z","iopub.execute_input":"2025-01-22T06:02:44.177153Z","iopub.status.idle":"2025-01-22T06:02:47.151518Z","shell.execute_reply.started":"2025-01-22T06:02:44.177132Z","shell.execute_reply":"2025-01-22T06:02:47.150842Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plot functions\nimport matplotlib.patches as mpatches\n\ndef row_to_confs(row):\n    return [row[f'conf_{i}'] for i in range(3)]\ndef row_to_coords(row):\n    return row[3:].values.reshape(3, 50, 2)\n\n# here I use matplotlib default colors\ncmap = plt.get_cmap(\"tab10\")\nmatplotlib_colors_in_rgb_int = [\n    [int(255 * x) for x in cmap(i)[:3]] for i in range(10)\n]\n\ndef generate_image_predicted_trajectory(dataset, df_sub, index):\n    data = dataset[index]\n    im = data['image'].transpose(1, 2, 0)\n    im = dataset.rasterizer.to_rgb(im)\n    row = df_sub.loc[(data['timestamp'], data['track_id'])]\n    # note submission coordinate system = world - centroid\n    predicted_target_positions_in_sub = row_to_coords(row)\n    predicted_target_positions_in_world = predicted_target_positions_in_sub + data['centroid']\n    for i, coords in enumerate(predicted_target_positions_in_world):\n        target_positions_pixels = transform_points(coords, data['raster_from_world'])\n        draw_trajectory(im, target_positions_pixels, rgb_color=matplotlib_colors_in_rgb_int[i])\n    return im, row_to_confs(row)\n\ndef plot_predicted_trajectory(dataset, df_sub, indices, width=12, height=4, n_cols=3, title=''):\n    if not isinstance(indices, (list, np.ndarray)):\n        indices = [indices]\n    n_rows = len(indices) // n_cols + len(indices) % n_cols\n    plt.figure(figsize=(width, height*n_rows))\n    for k, index in enumerate(indices):\n        plt.subplot(n_rows, n_cols, 1+k).set_title(str(index))\n        im, confs = generate_image_predicted_trajectory(dataset, df_sub, index)\n        patches = [mpatches.Patch(color=cmap(m), label='%.3f'%conf) for m, conf in enumerate(confs)]\n        plt.imshow(im, origin='lower')\n        plt.legend(handles=patches)\n    if title:\n        plt.suptitle(title)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T06:02:47.152907Z","iopub.execute_input":"2025-01-22T06:02:47.153126Z","iopub.status.idle":"2025-01-22T06:02:47.166392Z","shell.execute_reply.started":"2025-01-22T06:02:47.153104Z","shell.execute_reply":"2025-01-22T06:02:47.165521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_predicted_trajectory(test_dataset, df_sub, [18431], width=6, height=6, n_cols=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T06:02:47.167640Z","iopub.execute_input":"2025-01-22T06:02:47.167984Z","iopub.status.idle":"2025-01-22T06:02:47.533768Z","shell.execute_reply.started":"2025-01-22T06:02:47.167949Z","shell.execute_reply":"2025-01-22T06:02:47.532930Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"i_plots = np.random.randint(len(test_dataset), size=9)\nplot_predicted_trajectory(test_dataset, df_sub, i_plots)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T06:02:47.534960Z","iopub.execute_input":"2025-01-22T06:02:47.535274Z","iopub.status.idle":"2025-01-22T06:02:49.923417Z","shell.execute_reply.started":"2025-01-22T06:02:47.535244Z","shell.execute_reply":"2025-01-22T06:02:49.922768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom tqdm import tqdm\n\ndef compute_negative_log_likelihood(model, test_dataloader, device):\n    \"\"\"\n    计算模型在测试数据集上的负对数似然（NLL）。\n\n    参数:\n        model: 训练好的模型。\n        test_dataloader: 测试数据加载器。\n        device: 设备（如 'cuda:0' 或 'cpu'）。\n\n    返回:\n        float: 测试数据集上的平均负对数似然。\n    \"\"\"\n    model.eval()  # 将模型设置为评估模式\n    torch.set_grad_enabled(False)  # 禁用梯度计算\n\n    total_nll = 0.0\n    num_samples = 0\n\n    # 遍历测试数据集\n    for data in tqdm(test_dataloader, desc=\"Computing NLL\"):\n        # 将数据移动到设备\n        inputs = data[\"image\"].to(device)\n        targets = data[\"target_positions\"].to(device)\n        target_availabilities = data[\"target_availabilities\"].to(device)\n\n        # 前向传播\n        preds, confidences = model(inputs)\n\n        # 计算负对数似然\n        nll = pytorch_neg_multi_log_likelihood_batch(targets, preds, confidences, target_availabilities)\n        total_nll += nll.item() * inputs.size(0)  # 乘以批次大小\n        num_samples += inputs.size(0)\n\n    # 返回平均负对数似然\n    return total_nll / num_samples\n\n# 调用函数计算 NLL\naverage_nll = compute_negative_log_likelihood(model, test_dataloader, device)\nprint(f\"Average Negative Log-Likelihood (NLL) on Test Dataset: {average_nll:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-22T06:09:38.091575Z","iopub.execute_input":"2025-01-22T06:09:38.091926Z","iopub.status.idle":"2025-01-22T07:17:44.339086Z","shell.execute_reply.started":"2025-01-22T06:09:38.091895Z","shell.execute_reply":"2025-01-22T07:17:44.337938Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}