{"cells":[{"metadata":{},"cell_type":"markdown","source":"lyft package installation"},{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install --no-index -q --use-feature=2020-resolver -f ../input/kaggle-l5kit-110 l5kit","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"importing modules"},{"metadata":{"trusted":true},"cell_type":"code","source":"import gc\nimport os\nfrom pathlib import Path\nimport random\nimport sys\nfrom l5kit.data import ChunkedDataset, LocalDataManager\nfrom l5kit.dataset import EgoDataset, AgentDataset\nfrom tqdm.notebook import tqdm\nimport numpy as np\nimport pandas as pd\nimport scipy as sp\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfrom IPython.core.display import display, HTML\n\n# --- plotly ---\nfrom plotly import tools, subplots\nimport plotly.offline as py\npy.init_notebook_mode(connected=True)\nimport plotly.graph_objs as go\nimport plotly.express as px\nimport plotly.figure_factory as ff\nimport plotly.io as pio\npio.templates.default = \"plotly_dark\"\n\n# --- models ---\nfrom sklearn import preprocessing\nfrom sklearn.model_selection import KFold\nimport lightgbm as lgb\nimport xgboost as xgb\nimport catboost as cb\n\n# --- setup ---\npd.set_option('max_columns', 50)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"ba9a0d10-e340-474d-bcec-f50497076d68","_cell_guid":"cffd50a0-66e6-4845-89a5-2daa624e4010","trusted":true},"cell_type":"code","source":"\nfrom l5kit.data import ChunkedDataset, LocalDataManager\nfrom l5kit.dataset import EgoDataset, AgentDataset\nfrom l5kit.evaluation import write_pred_csv\nfrom l5kit.rasterization import build_rasterizer\nfrom l5kit.configs import load_config_data\nfrom l5kit.visualization import draw_trajectory, TARGET_POINTS_COLOR\nfrom l5kit.geometry import transform_points","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nfrom pathlib import Path\n# !pip install pytorch_pfn_extras\n# import pytorch_pfn_extras as ppe\nfrom math import ceil\n# from pytorch_pfn_extras.training import IgniteExtensionsManager\n# from pytorch_pfn_extras.training.triggers import MinValueTrigger\nfrom torch import nn, optim\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data.dataset import Subset\n# import pytorch_pfn_extras.training.extensions as E","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"f40537c2-cbfd-4372-876b-8049f634544b","_cell_guid":"c0f7e074-168e-4663-a6d4-2d1748b564f3","trusted":true},"cell_type":"markdown","source":"Data Loader"},{"metadata":{"_uuid":"bdf655ca-25d3-467e-a3fd-3541c7c7ec19","_cell_guid":"94e38fee-ea00-4a79-9a26-817f94e6422b","trusted":true},"cell_type":"code","source":"# --- Dataset utils ---\nfrom typing import Callable\n\nfrom torch.utils.data.dataset import Dataset\n\nclass TransformDataset(Dataset):\n    def __init__(self, dataset: Dataset, transform: Callable):\n        self.dataset = dataset\n        self.transform = transform\n\n    def __getitem__(self, index):\n        batch = self.dataset[index]\n        return self.transform(batch)\n\n    def __len__(self):\n        return len(self.dataset)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"f5e2a4e5-c40c-486f-8733-215c24d14091","_cell_guid":"f8ea4fa1-c025-4fd6-82c5-e1a349134cb2","trusted":true},"cell_type":"markdown","source":"Loss function"},{"metadata":{"_uuid":"2c4d902a-dfe6-44e9-b57e-3435ef74b378","_cell_guid":"348222c9-c0aa-4e06-90be-12e389ab494e","trusted":true},"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)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"2db30941-42dd-46df-9f7e-eeeadd17a418","_cell_guid":"cdb6e401-03c3-4810-8328-d38e09299eec","trusted":true},"cell_type":"markdown","source":"Model, where I use resnet-18 followed by attention model. If you want to use LSTM also then un-comment lines after x = self.head(x)"},{"metadata":{"_uuid":"aed24705-95c7-4a5d-aa95-ca10c4620e33","_cell_guid":"3155b072-eb66-442c-8f64-5b124b8e6f79","trusted":true},"cell_type":"code","source":"# --- Model utils ---\nimport torch\nfrom torchvision.models import resnet18, resnet50\nfrom torch import nn\nfrom typing import Dict\nimport torch.nn.functional as F\n\nclass LyftMultiModelAttn(nn.Module):\n\n    def __init__(self, cfg: Dict, num_modes=3):\n        super().__init__()\n        \n        backbone = resnet18(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            1,\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        # This is 512 for resnet18 and resnet34;\n        # And it is 2048 for the other resnets\n        backbone_out_features = 512        \n        self.backbone.layer5 = nn.Conv2d(\n            backbone_out_features,\n            128,\n            kernel_size=2,\n            stride=2,\n            # padding=self.backbone.conv1.padding,\n            bias=False,\n        )        \n\n#         self.backbone.layer6 = nn.Conv2d(\n#             1024,\n#             2048,\n#             kernel_size=2,\n#             stride=2,\n#             # padding=self.backbone.conv1.padding,\n#             bias=False,\n#         )        \n\n#         self.backbone.layer7 = nn.Conv2d(\n#             2048,\n#             2048,\n#             kernel_size=2,\n#             stride=2,\n#             # padding=self.backbone.conv1.padding,\n#             bias=False,\n#         )        \n        \n\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=128*7*7, out_features=4096),\n        )\n\n        self.num_preds = num_targets * num_modes\n        self.num_modes = num_modes\n        self.attn_layer = nn.Sequential(\n            nn.Linear(128, 512, False),\n            nn.BatchNorm1d(512),\n            nn.Tanh(),          \n            # nn.Dropout(0.5),\n            nn.Linear(512, 1, False)\n        )\n        self.logit = nn.Linear(1024, out_features=self.num_preds + num_modes)\n\n        num_layers = 1\n        self.lstm = nn.LSTM(input_size=128*7*7,\n                            hidden_size=1024,\n                            num_layers=num_layers)\n\n        self.bs = cfg['train_data_loader']['batch_size']\n\n    def forward(self, x):\n\n        batch_size, time_steps, height, width = x.size()\n        x = x.view(batch_size * time_steps, 1, height, width)\n\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        x = self.backbone.layer5(x)\n        # x = self.backbone.layer6(x)\n        # x = self.backbone.layer7(x)\n        _, _, height, width = x.size()\n        x = x.view(batch_size * height * width * time_steps, 128)\n        alpha = self.attn_layer(x)\n        alpha = alpha.view(batch_size * time_steps, height * width)\n        alpha = F.softmax(alpha, dim=1)\n        alpha = alpha.view(batch_size * height * width * time_steps, 1).clone().repeat(1, 128)\n        x = x * alpha\n        # x = x.view(batch_size, time_steps, height * width *128)\n\n\n        # x = self.backbone.avgpool(x)\n        # x = torch.flatten(x, 1)\n\n        # x = self.head(x)\n        \n        x = x.view(batch_size, time_steps, -1)\n\n        x = x.permute(1, 0, 2)\n\n        _, (x, _) = self.lstm(x)\n        x = x.squeeze()\n        del _\n        torch.cuda.empty_cache()\n        # x = torch.flatten(x, 1)\n        x = self.logit(x)\n\n        # pred (bs)x(modes)x(time)x(2D coords)\n        # confidences (bs)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","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"e0e0bf5e-5aae-4cf2-85d0-15fc4ba0c143","_cell_guid":"89fd1c9e-0b2c-4731-b065-d4010c1d004e","trusted":true},"cell_type":"code","source":"class LyftMultiRegressor(nn.Module):\n    \"\"\"Single mode prediction\"\"\"\n\n    def __init__(self, predictor, lossfun=pytorch_neg_multi_log_likelihood_batch):\n        super().__init__()\n        self.predictor = predictor\n        self.lossfun = lossfun\n\n    def forward(self, image, targets, target_availabilities):\n        pred, confidences = self.predictor(image)\n        loss = self.lossfun(targets, pred, confidences, target_availabilities)\n        metrics = {\n            \"loss\": loss.item(),\n            \"nll\": pytorch_neg_multi_log_likelihood_batch(targets, pred, confidences, target_availabilities).item()\n        }\n        # ppe.reporting.report(metrics, self)\n        return loss, metrics","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Inference"},{"metadata":{"trusted":true},"cell_type":"code","source":"def run_prediction(predictor, data_loader):\n    predictor.eval()\n\n    pred_coords_list = []\n    confidences_list = []\n    timestamps_list = []\n    track_id_list = []\n\n    with torch.no_grad():\n        dataiter = tqdm(data_loader)\n        for data in dataiter:\n            image = data[\"image\"].to(device)\n            # target_availabilities = data[\"target_availabilities\"].to(device)\n            # targets = data[\"target_positions\"].to(device)\n            pred, confidences = predictor(image)\n\n            pred_coords_list.append(pred.cpu().numpy().copy())\n            confidences_list.append(confidences.cpu().numpy().copy())\n            timestamps_list.append(data[\"timestamp\"].numpy().copy())\n            track_id_list.append(data[\"track_id\"].numpy().copy())\n    timestamps = np.concatenate(timestamps_list)\n    track_ids = np.concatenate(track_id_list)\n    coords = np.concatenate(pred_coords_list)\n    confs = np.concatenate(confidences_list)\n    return timestamps, track_ids, coords, confs","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"6cf61668-dfff-45f6-9bd1-991073f33b8b","_cell_guid":"921eb45a-8673-492e-8b0e-21ef5ae58adc","trusted":true},"cell_type":"raw","source":"Trainer"},{"metadata":{"_uuid":"976264ff-3393-4929-88b7-0c7de3ccf1ed","_cell_guid":"1f231204-74b0-4af5-aeaa-61a159b5bbc9","trusted":true},"cell_type":"code","source":"# --- Training utils ---\nfrom ignite.engine import Engine\ndef create_trainer(model, optimizer, device) -> Engine:\n    model.to(device)\n    def update_fn(engine, batch):\n        model.train()\n        optimizer.zero_grad()\n        loss, metrics = model(*[elem.to(device) for elem in batch])\n        loss.backward()\n        optimizer.step()\n        return metrics\n    trainer = Engine(update_fn)\n    return trainer","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"588c75c0-042a-4b55-9c2a-6da4a12bdf9d","_cell_guid":"c7705a68-42b4-4947-a74b-cd219bcb242a","trusted":true},"cell_type":"code","source":"# --- Utils ---\nimport yaml\n\n\ndef save_yaml(filepath, content, width=120):\n    with open(filepath, 'w') as f:\n        yaml.dump(content, f, width=width)\n\n\ndef load_yaml(filepath):\n    with open(filepath, 'r') as f:\n        content = yaml.safe_load(f)\n    return content\n\nclass DotDict(dict):\n    \"\"\"dot.notation access to dictionary attributes\n\n    Refer: https://stackoverflow.com/questions/2352181/how-to-use-a-dot-to-access-members-of-dictionary/23689767#23689767\n    \"\"\"  # NOQA\n\n    __getattr__ = dict.get\n    __setattr__ = dict.__setitem__\n    __delattr__ = dict.__delitem__","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"config file"},{"metadata":{"trusted":true},"cell_type":"code","source":"# --- Lyft configs ---\ncfg = {\n    'format_version': 4,\n    'model_params': {\n        'model_architecture': 'resnet50',\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    },\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\n    'train_data_loader': {\n        'key': 'scenes/train.zarr',\n        'batch_size':5,\n        'shuffle': True,\n        'num_workers': 4 \n    },\n\n    'valid_data_loader': {\n        'key': 'scenes/validate.zarr',\n        'batch_size': 5,\n        'shuffle': False,\n        'num_workers': 4\n    },\n    'test_data_loader': {\n        'key': 'scenes/test.zarr',\n        'batch_size': 5,\n        'shuffle': False,\n        'num_workers': 4\n    },\n    'train_params': {\n        'max_num_steps': 10000,\n        'checkpoint_every_n_steps': 5000,\n\n        # 'eval_every_n_steps': -1\n    }\n}","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"b930764b-4aca-4713-9332-3a803a4bf088","_cell_guid":"98a53d40-24ab-4a9d-879b-f7cfefefb8fe","trusted":true},"cell_type":"code","source":"device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\n\nflags_dict = {\n    \"debug\": True,\n    # --- Data configs ---\n    \"l5kit_data_folder\": \"/kaggle/input/lyft-motion-prediction-autonomous-vehicles\",\n    # --- Model configs ---\n    \"pred_mode\": \"multi\",\n    # --- Training configs ---\n    \"device\": device,\n    \"out_dir\": \"results/multi_train\",\n    \"epoch\": 20,\n    \"snapshot_freq\": 50,\n}\nprint(device)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"9dfe6858-9a66-4885-bc91-45bf3bcc4639","_cell_guid":"0717586b-6d96-4e9d-b628-915135aed753","trusted":true},"cell_type":"code","source":"flags = DotDict(flags_dict)\nout_dir = Path(flags.out_dir)\nos.makedirs(str(out_dir), exist_ok=True)\nprint(f\"flags: {flags_dict}\")\nsave_yaml(out_dir / 'flags.yaml', flags_dict)\nsave_yaml(out_dir / 'cfg.yaml', cfg)\ndebug = flags.debug","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Data loading "},{"metadata":{"trusted":true},"cell_type":"code","source":"# set env variable for data\nos.environ[\"L5KIT_DATA_FOLDER\"] = flags.l5kit_data_folder\ndm = LocalDataManager(None)\n\nprint(\"Load dataset...\")\ntrain_cfg = cfg[\"train_data_loader\"]\nvalid_cfg = cfg[\"valid_data_loader\"]\n\n# Rasterizer\nrasterizer = build_rasterizer(cfg, dm)\n\n# Train dataset/dataloader\ndef transform(batch):\n    return batch[\"image\"], batch[\"target_positions\"], batch[\"target_availabilities\"]\n\ntrain_path = \"scenes/sample.zarr\" if debug else train_cfg[\"key\"]\ntrain_zarr = ChunkedDataset(dm.require(train_path)).open()\nprint(\"train_zarr\", type(train_zarr))\ntrain_agent_dataset = AgentDataset(cfg, train_zarr, rasterizer)\ntrain_dataset = TransformDataset(train_agent_dataset, transform)\nif debug:\n    # Only use 1000 dataset for fast check...\n    train_dataset = Subset(train_dataset, np.arange(1400))\nelse:\n    train_dataset = Subset(train_dataset, np.arange(100000))\n\nprint(train_agent_dataset)\n\nvalid_path = \"scenes/sample.zarr\" if debug else valid_cfg[\"key\"]\nvalid_zarr = ChunkedDataset(dm.require(valid_path)).open()\nprint(\"valid_zarr\", type(train_zarr))\nvalid_agent_dataset = AgentDataset(cfg, valid_zarr, rasterizer)\nvalid_dataset = TransformDataset(valid_agent_dataset, transform)\nif debug:\n    # Only use 100 dataset for fast check...\n    valid_dataset = Subset(valid_dataset, np.arange(100))\nelse:\n    # Only use 1000 dataset for fast check...\n    valid_dataset = Subset(valid_dataset, np.arange(1000))\n\nprint(valid_agent_dataset)\nprint(\"# AgentDataset train:\", len(train_agent_dataset), \"#valid\", len(valid_agent_dataset))\n\n# AgentDataset train: 22496709 #valid 21624612\n# ActualDataset train: 100 #valid 100","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"PyTorch Dataloader declaration "},{"metadata":{"trusted":true},"cell_type":"code","source":"train_cfg = cfg[\"train_data_loader\"]\nvalid_cfg = cfg[\"valid_data_loader\"]\ntrain_loader = DataLoader(train_dataset,\n                          shuffle=train_cfg[\"shuffle\"],\n                          batch_size=train_cfg[\"batch_size\"],\n                          num_workers=train_cfg[\"num_workers\"])\nvalid_loader = DataLoader(\n    valid_dataset,\n    shuffle=valid_cfg[\"shuffle\"],\n    batch_size=valid_cfg[\"batch_size\"],\n    num_workers=valid_cfg[\"num_workers\"]\n)\nprint(\"# ActualDataset train:\", len(train_dataset), \"#valid\", len(valid_dataset))\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Test Data loading "},{"metadata":{"_uuid":"a62d89dc-5f87-4904-897c-f492cf9c901a","_cell_guid":"fc457058-5529-4736-bfb6-95dfdf1000ea","trusted":true},"cell_type":"code","source":"# set env variable for data\nl5kit_data_folder = \"/kaggle/input/lyft-motion-prediction-autonomous-vehicles\"\nos.environ[\"L5KIT_DATA_FOLDER\"] = l5kit_data_folder\ndm = LocalDataManager(None)\n\nprint(\"Load dataset...\")\ndefault_test_cfg = {\n    'key': 'scenes/test.zarr',\n    'batch_size': 32,\n    'shuffle': False,\n    'num_workers': 4\n}\ntest_cfg = cfg.get(\"test_data_loader\", default_test_cfg)\n\n# Rasterizer\nrasterizer = build_rasterizer(cfg, dm)\n\ntest_path = test_cfg[\"key\"]\nprint(f\"Loading from {test_path}\")\ntest_zarr = ChunkedDataset(dm.require(test_path)).open()\nprint(\"test_zarr\", type(test_zarr))\ntest_mask = np.load(f\"{l5kit_data_folder}/scenes/mask.npz\")[\"arr_0\"]\ntest_agent_dataset = AgentDataset(cfg, test_zarr, rasterizer, agents_mask=test_mask)\ntest_dataset = test_agent_dataset\n# if debug:\n    # Only use 100 dataset for fast check...\n    # test_dataset = Subset(test_dataset, np.arange(100))\ntest_loader = DataLoader(\n    test_dataset,\n    shuffle=test_cfg[\"shuffle\"],\n    batch_size=test_cfg[\"batch_size\"],\n    num_workers=test_cfg[\"num_workers\"],\n    pin_memory=True,\n)\n\nprint(test_agent_dataset)\nprint(\"# AgentDataset test:\", len(test_agent_dataset))\nprint(\"# ActualDataset test:\", len(test_dataset))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Model declaration"},{"metadata":{"_uuid":"443acba6-5698-4b82-8a41-039956f7245b","_cell_guid":"bea2b730-0727-4908-9bc9-8da550585e56","trusted":true},"cell_type":"code","source":"device = torch.device(flags.device)\n\nif flags.pred_mode == \"multi\":\n\n    predictor = LyftMultiModelAttn(cfg)\n  \n    model = LyftMultiRegressor(predictor)\n\nelse:\n    raise ValueError(f\"[ERROR] Unexpected value flags.pred_mode={flags.pred_mode}\")\n\nmodel.to(device)\noptimizer = optim.Adam(model.parameters(), lr=1e-3)\nscheduler = torch.optim.lr_scheduler.ExponentialLR(\n    optimizer, gamma=0.99999)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Main function with inference"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Train setup\ntrainer = create_trainer(model, optimizer, device)\nMAX = 1e16 + 7\ndef eval_func(*batch):\n    loss, metrics = model(*[elem.to(device) for elem in batch])\n    return loss, metrics\n\ndef eval(model, loader, eval_func):\n    \n    model.eval()\n    error = 0\n    count = 0\n    for batch_i, batch in enumerate(loader):\n        count += 1\n        with torch.no_grad():\n            loss, metrics = model(*[elem.to(device) for elem in batch])\n        error += loss.item()\n\n        del metrics\n        torch.cuda.empty_cache()\n    print(\"Validation loss per batch {}\".format(error/count))\n    return loss\n\ndef train(model, loader, eval_func, optimizer):\n    model.train()\n    error = 0\n    count = 0\n    lastcheckpoint = flags.out_dir+'/intermediate_model.pth'\n    if os.path.isfile(lastcheckpoint):\n        print(\"loading ...\")\n        t = torch.load(lastcheckpoint, map_location=lambda storage, loc: storage)   \n        model.predictor.load_state_dict(t['state_dict'])\n        print(\"done\")\n    else:\n        print(\"file not found\")\n        \n    for batch in tqdm(loader):\n\n        count += 1\n        optimizer.zero_grad()        \n        loss, metrics = model(*[elem.to(device) for elem in batch])\n        loss.backward()\n        optimizer.step()        \n        scheduler.step()\n        del metrics\n        torch.cuda.empty_cache()        \n        if count%10 == 0:\n            print(\"saving at \", flags.out_dir+'/intermediate_model.pth')\n            torch.save({'count': count, 'state_dict': model.predictor.state_dict()},\n                       flags.out_dir+'/intermediate_model.pth')\n            print(\"Epoch no. {} TR loss {} lr {}\".format(count, error/count, optimizer.param_groups[0]['lr']))       \n        error += loss.item()\n        del loss\n        torch.cuda.empty_cache()    \n    print(\"training loss per batch {}\".format(error/count))\n    return loss\n\n\nepoch = flags.epoch\nfor epoch_n in range(epoch):\n    print(\"epoch no.\", epoch_n)\n    tl = train(model, train_loader, eval_func, optimizer)\n    vl = eval(model, valid_loader, eval_func)\n    if vl < MAX:\n        timestamps, track_ids, coords, confs = run_prediction(model.predictor, test_loader)\n        def saving_csv():\n            csv_path = \"submission.csv\"\n            write_pred_csv(\n                csv_path,\n                timestamps=timestamps,\n                track_ids=track_ids,\n                coords=coords,\n                confs=confs)\n            print(f\"Saved to {csv_path}\")\n        saving_csv()\n    basename = \"epoch {} train loss {} val loss{}\".format(epoch_n, tl, vl)\n    torch.save({'epoch': epoch_n, 'state_dict': model.predictor.state_dict()},\n               flags.out_dir+'/model.pth')","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat":4,"nbformat_minor":4}