{"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_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nsys.path.insert(0, '/kaggle/input/l5kit-may31/l5kit/')","metadata":{"execution":{"iopub.status.busy":"2021-06-02T06:14:52.427186Z","iopub.execute_input":"2021-06-02T06:14:52.427739Z","iopub.status.idle":"2021-06-02T06:14:52.437117Z","shell.execute_reply.started":"2021-06-02T06:14:52.427654Z","shell.execute_reply":"2021-06-02T06:14:52.436008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from IPython.core.debugger import set_trace","metadata":{"execution":{"iopub.status.busy":"2021-06-02T06:15:03.760400Z","iopub.execute_input":"2021-06-02T06:15:03.760745Z","iopub.status.idle":"2021-06-02T06:15:03.765372Z","shell.execute_reply.started":"2021-06-02T06:15:03.760716Z","shell.execute_reply":"2021-06-02T06:15:03.764096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport os\nimport psutil\nimport torch\n\nfrom torch import nn, optim\nfrom torch.utils.data import DataLoader, random_split\nfrom torchvision.models.resnet import resnet50\nfrom tqdm.notebook import tqdm\nfrom typing import Dict\nfrom pprint import pprint\n\nfrom l5kit.configs import load_config_data\nfrom l5kit.data import LocalDataManager, ChunkedDataset\nfrom l5kit.dataset import AgentDataset, EgoDataset\nfrom l5kit.evaluation import write_pred_csv\nfrom l5kit.geometry import transform_points\nfrom l5kit.rasterization import build_rasterizer","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-06-02T06:15:05.494666Z","iopub.execute_input":"2021-06-02T06:15:05.495343Z","iopub.status.idle":"2021-06-02T06:15:09.926854Z","shell.execute_reply.started":"2021-06-02T06:15:05.495282Z","shell.execute_reply":"2021-06-02T06:15:09.925594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT_DIR = '/kaggle/input/lyft-motion-prediction-autonomous-vehicles'\nWEIGHTS_FILE = '/kaggle/input/cs535-resnet50-training/cs535_resnet50.pth'","metadata":{"execution":{"iopub.status.busy":"2021-06-02T06:15:17.122354Z","iopub.execute_input":"2021-06-02T06:15:17.122715Z","iopub.status.idle":"2021-06-02T06:15:17.126939Z","shell.execute_reply.started":"2021-06-02T06:15:17.122679Z","shell.execute_reply":"2021-06-02T06:15:17.125830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# set env variable for data\nos.environ[\"L5KIT_DATA_FOLDER\"] = INPUT_DIR\ndm = LocalDataManager(None)","metadata":{"execution":{"iopub.status.busy":"2021-06-02T06:15:18.405246Z","iopub.execute_input":"2021-06-02T06:15:18.405649Z","iopub.status.idle":"2021-06-02T06:15:18.409671Z","shell.execute_reply.started":"2021-06-02T06:15:18.405616Z","shell.execute_reply":"2021-06-02T06:15:18.408931Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = load_config_data(\"/kaggle/input/l5kit-may31/examples/agent_motion_prediction/agent_motion_config.yaml\")\n# pprint(cfg)","metadata":{"execution":{"iopub.status.busy":"2021-06-02T06:15:19.790823Z","iopub.execute_input":"2021-06-02T06:15:19.791437Z","iopub.status.idle":"2021-06-02T06:15:19.827374Z","shell.execute_reply.started":"2021-06-02T06:15:19.791403Z","shell.execute_reply":"2021-06-02T06:15:19.826317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg['model_params']['history_num_frames'] = 10\ncfg['val_data_loader']['batch_size'] = 12\ncfg['val_data_loader']['num_workers'] = 4\ncfg['val_data_loader']['key'] = 'scenes/test.zarr'","metadata":{"execution":{"iopub.status.busy":"2021-06-02T06:15:21.130980Z","iopub.execute_input":"2021-06-02T06:15:21.131329Z","iopub.status.idle":"2021-06-02T06:15:21.136340Z","shell.execute_reply.started":"2021-06-02T06:15:21.131299Z","shell.execute_reply":"2021-06-02T06:15:21.134957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Init test dataset","metadata":{}},{"cell_type":"code","source":"# ===== INIT DATASET\ntest_cfg = cfg[\"val_data_loader\"]\n\n# Rasterizer\nrasterizer = build_rasterizer(cfg, dm)\n\n# Test dataset/dataloader\ntest_zarr = ChunkedDataset(dm.require(test_cfg[\"key\"])).open()\ntest_mask = np.load(f\"{INPUT_DIR}/scenes/mask.npz\")[\"arr_0\"]\ntest_dataset = AgentDataset(cfg, test_zarr, rasterizer, agents_mask=test_mask)\n# test_dataset, _ = random_split(test_dataset, [100, 71122-100])\ntest_dataloader = DataLoader(test_dataset,\n                             shuffle=test_cfg[\"shuffle\"],\n                             batch_size=test_cfg[\"batch_size\"],\n                             num_workers=test_cfg[\"num_workers\"])\n\n\nprint(test_dataloader)\nprint(len(test_dataset))\nprint(len(test_dataloader))","metadata":{"execution":{"iopub.status.busy":"2021-06-02T06:15:25.172519Z","iopub.execute_input":"2021-06-02T06:15:25.173058Z","iopub.status.idle":"2021-06-02T06:15:42.388842Z","shell.execute_reply.started":"2021-06-02T06:15:25.173024Z","shell.execute_reply":"2021-06-02T06:15:42.387784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Build model","metadata":{}},{"cell_type":"code","source":"def build_model(cfg: Dict) -> torch.nn.Module:\n    # load pre-trained Conv2D model\n    model = resnet50(pretrained=False)\n\n    # change input channels number to match the rasterizer's output\n    num_history_channels = (cfg[\"model_params\"][\"history_num_frames\"] + 1) * 2\n    num_in_channels = 3 + num_history_channels\n    model.conv1 = nn.Conv2d(\n        num_in_channels,\n        model.conv1.out_channels,\n        kernel_size=model.conv1.kernel_size,\n        stride=model.conv1.stride,\n        padding=model.conv1.padding,\n        bias=False,\n    )\n    # change output size to (X, Y) * number of future states\n    num_targets = 2 * cfg[\"model_params\"][\"future_num_frames\"]\n    model.fc = nn.Linear(in_features=2048, out_features=num_targets)\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2021-06-02T06:15:42.390748Z","iopub.execute_input":"2021-06-02T06:15:42.391455Z","iopub.status.idle":"2021-06-02T06:15:42.399589Z","shell.execute_reply.started":"2021-06-02T06:15:42.391406Z","shell.execute_reply":"2021-06-02T06:15:42.398792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ==== INIT MODEL\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nmodel = build_model(cfg).to(device)\nmodel.load_state_dict(torch.load(WEIGHTS_FILE, map_location=device))\n# optimizer = optim.Adam(model.parameters(), lr=1e-3)\n# criterion = nn.MSELoss(reduction=\"none\")","metadata":{"execution":{"iopub.status.busy":"2021-06-02T06:15:42.401734Z","iopub.execute_input":"2021-06-02T06:15:42.402643Z","iopub.status.idle":"2021-06-02T06:15:45.459405Z","shell.execute_reply.started":"2021-06-02T06:15:42.402592Z","shell.execute_reply":"2021-06-02T06:15:45.458531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference loop","metadata":{}},{"cell_type":"code","source":"model.eval()\n\nfuture_coords_offsets_pd = []\ntimestamps = []\nagent_ids = []\n\nwith torch.no_grad():\n    dataiter = iter(test_dataloader)\n    \n    pbar = tqdm(dataiter)\n    for data in pbar:\n\n        inputs = data[\"image\"].to(device)\n        target_availabilities = data[\"target_availabilities\"].unsqueeze(-1).to(device)\n        targets = data[\"target_positions\"].to(device)\n        outputs = model(inputs).reshape(targets.shape)\n        \n        # convert agent coordinates into world offsets\n        agents_coords = outputs.cpu().numpy().copy()\n        world_from_agents = data[\"world_from_agent\"].numpy()\n        centroids = data[\"centroid\"].numpy()\n        coords_offset = transform_points(agents_coords, world_from_agents) - centroids[:, None, :2]\n        \n        future_coords_offsets_pd.append(coords_offset)\n        timestamps.append(data[\"timestamp\"].numpy().copy())\n        agent_ids.append(data[\"track_id\"].numpy().copy())\n        \n        pbar.set_description(f'RAM used: {psutil.virtual_memory().percent}%')","metadata":{"execution":{"iopub.status.busy":"2021-06-02T06:15:51.006195Z","iopub.execute_input":"2021-06-02T06:15:51.006599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Get submission file","metadata":{}},{"cell_type":"code","source":"write_pred_csv('submission.csv',\n               timestamps=np.concatenate(timestamps),\n               track_ids=np.concatenate(agent_ids),\n               coords=np.concatenate(future_coords_offsets_pd))","metadata":{},"execution_count":null,"outputs":[]}]}