{"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":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        os.path.join(dirname, filename)\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install numpy==1.18.5\n!pip install pandas==1.1.4\n!pip install matplotlib==3.3.2\n!pip install seaborn==0.11.0\n!pip install tqdm==4.51.0\n!pip install torch==1.9.0\n!pip install torchtext==0.9.1\n!pip install efficientnet_pytorch==0.7.0\n!pip install l5kit==1.0.6\n!pip install omegaconf","metadata":{"execution":{"iopub.status.busy":"2022-05-30T01:47:38.028199Z","iopub.execute_input":"2022-05-30T01:47:38.028432Z","iopub.status.idle":"2022-05-30T01:52:10.642336Z","shell.execute_reply.started":"2022-05-30T01:47:38.028407Z","shell.execute_reply":"2022-05-30T01:52:10.641547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install timm==0.4.12","metadata":{"execution":{"iopub.status.busy":"2022-05-30T01:52:10.645792Z","iopub.execute_input":"2022-05-30T01:52:10.646Z","iopub.status.idle":"2022-05-30T01:52:19.819095Z","shell.execute_reply.started":"2022-05-30T01:52:10.645972Z","shell.execute_reply":"2022-05-30T01:52:19.818232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nfrom IPython.display import Image\nfrom PIL import Image\nimport PIL\n\nimport matplotlib.pyplot as plt\nfrom matplotlib.image import imread\nimport seaborn as sns\n\n\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data.dataset import Dataset\n\nimport logging\nimport numpy as np\nimport pickle\nfrom collections import defaultdict\n\nfrom omegaconf import DictConfig\n# import hydra\n\nimport pandas as pd\nfrom tqdm import tqdm\n\nimport torch.optim as optim\nimport matplotlib.pyplot as plt\nfrom matplotlib.image import imread\nimport seaborn as sns\n\n#from hydra import initialize, initialize_config_module, initialize_config_dir, compose\nfrom omegaconf import DictConfig, OmegaConf\n\n\nimport omegaconf\nfrom torch.distributions.multivariate_normal import MultivariateNormal\nfrom torchvision.models import mobilenet_v2\nfrom efficientnet_pytorch import EfficientNet\n\nfrom torch.utils.data.dataset import Dataset\nfrom torch.utils.data.dataset import Subset\nfrom torch.utils.data import DataLoader\n\nfrom l5kit.geometry import transform_points\nfrom l5kit.data import ChunkedDataset, LocalDataManager\nfrom l5kit.dataset import AgentDataset\nfrom l5kit.rasterization import build_rasterizer","metadata":{"execution":{"iopub.status.busy":"2022-05-30T01:52:19.821668Z","iopub.execute_input":"2022-05-30T01:52:19.821894Z","iopub.status.idle":"2022-05-30T01:52:24.204234Z","shell.execute_reply.started":"2022-05-30T01:52:19.821865Z","shell.execute_reply":"2022-05-30T01:52:24.203465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TransformDataset(Dataset):\n    def __init__(self, dataset, cfg):\n        self.cfg = cfg\n        self.dataset = dataset\n        self.W = self.cfg['raster_params']['raster_size'][0]\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)\n\n    # Here batch is just 1 element\n    def transform(self, batch):\n\n        # (2,) agent coordinates in image frame of reference\n        agent_position_image = np.array([self.cfg['raster_params']['ego_center'][0]*self.W,\n                                         self.cfg['raster_params']['ego_center'][1]*self.W])\n\n        # (1,) meters per pixel\n        r = self.cfg['raster_params']['pixel_size'][0]\n\n        # initially positions are given in agent frame of reference but with wrong rotation\n        # transform into image frame of reference using rotation and affine transformation\n        # after that tranform into agent frame of reference back\n\n        batch[\"target_positions\"] = (transform_points(batch[\"target_positions\"] + batch[\"centroid\"],\n                                                      batch[\"world_to_image\"]) - agent_position_image)*r\n\n        batch['history_positions'] = (transform_points(batch['history_positions'] + batch[\"centroid\"],\n                                                       batch[\"world_to_image\"]) - agent_position_image)*r\n\n        return (batch[\"image\"].astype(np.float32),\n                batch[\"target_positions\"].astype(np.float32),\n                batch[\"target_availabilities\"].astype(np.float32),\n                batch['history_positions'].astype(np.float32),\n                batch['history_yaws'].astype(np.float32),\n                batch[\"centroid\"].astype(np.float32),\n                batch[\"world_to_image\"].astype(np.float32)\n                )","metadata":{"execution":{"iopub.status.busy":"2022-05-30T01:52:24.205895Z","iopub.execute_input":"2022-05-30T01:52:24.206133Z","iopub.status.idle":"2022-05-30T01:52:24.217087Z","shell.execute_reply.started":"2022-05-30T01:52:24.206099Z","shell.execute_reply":"2022-05-30T01:52:24.216353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FusionDCGAN(nn.Module):\n    \"\"\"\n    Classic DCGAN with fusion added for actor state.\n    This class handles Vanilla GAN and Wasserstein GAN.\n    \"\"\"\n\n    def __init__(self, nc, ndf,\n                 input_dim,\n                 gan_type,\n                 embedding_type,\n                 lstm_embedding_dim):\n\n        super().__init__()\n\n        self.nc = nc  # input channels number\n        self.ndf = ndf  # ndf - base for inner channels number\n        self.input_dim = input_dim\n        self.embedding_type = embedding_type\n        self.lstm_embedding_dim = lstm_embedding_dim\n\n        # in case of lstm encoding we first lstm-encode and later project with nn.linear to proper dimension\n        if self.embedding_type == 'lstm':\n            self.initial_embedding = LSTMEncoder(embedding_dim=self.lstm_embedding_dim,\n                                                 h_dim=self.lstm_embedding_dim,\n                                                 dropout=0.0)\n            self.input_dim = self.lstm_embedding_dim\n\n        # embedding_dim = 3*18*18 since we want to reshape it to (3, 18, 18)\n        self.state_encoder = nn.Linear(self.input_dim, 3*18*18)\n\n        # Conv with kernel 1x1 for state fusion\n        self.conv_state = nn.Conv2d(3, self.ndf * 8, kernel_size=1)\n\n        self.conv_block1 = nn.Sequential(\n            # input size is (nc) x 300 x 300\n            nn.Conv2d(self.nc, self.ndf, 4, 2, 1, bias=False),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block2 = nn.Sequential(\n            # input size is (ndf) x 150 x 150\n            nn.Conv2d(self.ndf, self.ndf * 2, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(self.ndf * 2),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block3 = nn.Sequential(\n            # input size is (ndf*2) x 75 x 75\n            nn.Conv2d(self.ndf * 2, self.ndf * 4, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(self.ndf * 4),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block4 = nn.Sequential(\n            # input size is (ndf*4) x 37 x 37\n            nn.Conv2d(self.ndf * 4, self.ndf * 8, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(self.ndf * 8),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block5 = nn.Sequential(\n            # input size is(ndf*8) x 18 x 18\n            nn.Conv2d(self.ndf * 8, self.ndf * 16, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(self.ndf * 16),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block6 = nn.Sequential(\n            # input size is (ndf*16) x 9 x 9\n            nn.Conv2d(self.ndf * 16, self.ndf * 32, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(self.ndf * 32),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        if gan_type == 'vanilla':\n            self.conv_block7 = nn.Sequential(\n                # input size is (ndf*32) x 4 x 4\n                nn.Conv2d(self.ndf * 32, 1, 4, 1, 0, bias=False),\n                # state size. 1 x 1 x 1\n                nn.Sigmoid()\n            )\n        else:  # for wasserstein GAN\n            self.conv_block7 = nn.Sequential(\n                # input size is (ndf*32) x 4 x 4\n                nn.Conv2d(self.ndf * 32, 1, 4, 1, 0, bias=False),\n                # state size. 1 x 1 x 1\n            )\n\n    def forward(self, image, actor_state):\n        print(\"fusion gan\")\n        # actor_state = [(batch_size, h_s, 2), (batch_size, h_s, 1)]\n\n        batch_size = image.shape[0]\n\n        actor_state = torch.cat(actor_state, dim=-1)  # (batch_size, h_s, 3)\n\n        if self.embedding_type == 'mlp':\n            actor_state = actor_state.reshape(batch_size, -1)  # (batch_size, 3*h_s)\n        elif self.embedding_type == 'lstm':\n            # add initial embedding in lstm-case\n            actor_state = self.initial_embedding(actor_state)  # (batch_size, lstm_embeding_dim)\n        else:\n            raise NotImplementedError\n\n        # actor_state fusion embedding\n        encoded_state = self.state_encoder(actor_state)  # (batch_size, 3*18*18)\n        encoded_state = encoded_state.reshape(batch_size, 3, 18, 18)  # (batch_size, 3, 18, 18)\n        encoded_state = self.conv_state(encoded_state)  # (batch_size, ndf * 8, 18, 18)\n\n        y = self.conv_block1(image)\n        y = self.conv_block2(y)\n        y = self.conv_block3(y)\n        y = self.conv_block4(y)\n\n        # fusion\n        y = y + encoded_state\n\n        y = self.conv_block5(y)\n        y = self.conv_block6(y)\n        y = self.conv_block7(y)\n        \n        print(y)\n\n        return y\n\n\nclass FusionDCGAN_gp(nn.Module):\n    \"\"\"\n    Classic DCGAN with fusion added for actor state.\n    This class handles Wasserstein GAN with Gradient Penalty.\n    There is no batchnorm in conv blocks.\n    \"\"\"\n\n    def __init__(self, nc, ndf,\n                 input_dim,\n                 embedding_type,\n                 lstm_embedding_dim):\n\n        super().__init__()\n\n        self.nc = nc  # input channels number\n        self.ndf = ndf  # ndf - base for inner channels number\n        self.input_dim = input_dim\n        self.embedding_type = embedding_type\n        self.lstm_embedding_dim = lstm_embedding_dim\n\n        # in case of lstm encoding we first lstm-encode and later project with nn.linear to proper dimension\n        if self.embedding_type == 'lstm':\n            self.initial_embedding = LSTMEncoder(embedding_dim=self.lstm_embedding_dim,\n                                                 h_dim=self.lstm_embedding_dim,\n                                                 dropout=0.0)\n            self.input_dim = self.lstm_embedding_dim\n\n        # embedding_dim = 3*18*18 since we want to reshape it to (3, 18, 18)\n        self.state_encoder = nn.Linear(self.input_dim, 3 * 18 * 18)  # torch.linear requires input dimensions\n\n        # Conv with kernel 1x1 for state fusion\n        self.conv_state = nn.Conv2d(3, self.ndf * 8, kernel_size=1)\n\n        self.conv_block1 = nn.Sequential(\n            # input size is (nc) x 300 x 300\n            nn.Conv2d(self.nc, self.ndf, 4, 2, 1, bias=False),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block2 = nn.Sequential(\n            # input size is (ndf) x 150 x 150\n            nn.Conv2d(self.ndf, self.ndf * 2, 4, 2, 1, bias=False),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block3 = nn.Sequential(\n            # input size is (ndf*2) x 75 x 75\n            nn.Conv2d(self.ndf * 2, self.ndf * 4, 4, 2, 1, bias=False),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block4 = nn.Sequential(\n            # input size is (ndf*4) x 37 x 37\n            nn.Conv2d(self.ndf * 4, self.ndf * 8, 4, 2, 1, bias=False),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block5 = nn.Sequential(\n            # input size is (ndf*8) x 18 x 18\n            nn.Conv2d(self.ndf * 8, self.ndf * 16, 4, 2, 1, bias=False),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block6 = nn.Sequential(\n            # input size is (ndf*16) x 9 x 9\n            nn.Conv2d(self.ndf * 16, self.ndf * 32, 4, 2, 1, bias=False),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n\n        self.conv_block7 = nn.Sequential(\n            # input size is (ndf*32) x 4 x 4\n            nn.Conv2d(self.ndf * 32, 1, 4, 1, 0, bias=False),\n            # state size. 1 x 1 x 1\n        )\n\n    def forward(self, image, actor_state):\n        # actor_state = [(batch_size, h_s, 2), (batch_size, h_s, 1)]\n\n        batch_size = image.shape[0]\n\n        actor_state = torch.cat(actor_state, dim=-1)  # (batch_size, h_s, 3)\n\n        if self.embedding_type == 'mlp':\n            actor_state = actor_state.reshape(batch_size, -1)  # (batch_size, 3*h_s)\n        elif self.embedding_type == 'lstm':\n            # add initial embedding in lstm-case\n            actor_state = self.initial_embedding(actor_state)  # (batch_size, lstm_embeding_dim)\n        else:\n            raise NotImplementedError\n\n        # actor_state fusion embedding\n        encoded_state = self.state_encoder(actor_state)  # (batch_size, 3*18*18)\n        encoded_state = encoded_state.reshape(batch_size, 3, 18, 18)  # (batch_size, 3, 18, 18)\n        encoded_state = self.conv_state(encoded_state)  # (batch_size, ndf * 8, 18, 18)\n\n        y = self.conv_block1(image)\n        y = self.conv_block2(y)\n        y = self.conv_block3(y)\n        y = self.conv_block4(y)\n\n        # fusion\n        y = y + encoded_state\n\n        y = self.conv_block5(y)\n        y = self.conv_block6(y)\n        y = self.conv_block7(y)\n\n        return y","metadata":{"execution":{"iopub.status.busy":"2022-05-30T01:52:24.218641Z","iopub.execute_input":"2022-05-30T01:52:24.219114Z","iopub.status.idle":"2022-05-30T01:52:24.257948Z","shell.execute_reply.started":"2022-05-30T01:52:24.219078Z","shell.execute_reply":"2022-05-30T01:52:24.257107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def f_get_raster_image(cfg,\n                       images,\n                       history_weight=0.9):\n\n    batch_size = images.shape[0]\n    image_size = images.shape[-1]\n\n    # get number of history steps\n    hnf = cfg['model_params']['history_num_frames']\n\n    ego_index = range(hnf+1, 2*hnf+2)\n\n    # iterate through ego-car's frames and sum them according to history_weight (history fading) in single channel.\n    ego_path_image = torch.zeros(size=(batch_size, image_size, image_size), device=cfg['device'])\n    for im_id in reversed(ego_index):\n        ego_path_image = (images[:, im_id, :, :] + ego_path_image * history_weight).clamp(0, 1)\n\n    # define agent's range\n    agents_index = range(0, hnf+1)\n\n    # iterate through agent-car's frames and sum them according to history_weight in single channel\n    agents_path_image = torch.zeros(size=(batch_size, image_size, image_size), device=cfg['device'])\n    for im_id in reversed(agents_index):\n        agents_path_image = (images[:, im_id, :, :] + agents_path_image*history_weight).clamp(0, 1)\n\n    #  RGB path for ego (red (255, 0, 0)); channels last\n    ego_path_image_rgb = torch.zeros((ego_path_image.shape[0],\n                                      ego_path_image.shape[1],\n                                      ego_path_image.shape[2],\n                                      3), device=cfg['device'])\n\n    ego_path_image_rgb[:, :, :, 0] = ego_path_image\n\n    # RGB paths for agents (yellow (255, 255, 0)); channels last\n    agents_path_image_rgb = torch.zeros((agents_path_image.shape[0],\n                                         agents_path_image.shape[1],\n                                         agents_path_image.shape[2],\n                                         3), device=cfg['device'])\n    # yellow\n    agents_path_image_rgb[:, :, :, 0] = agents_path_image\n    agents_path_image_rgb[:, :, :, 1] = agents_path_image\n\n    # generate full RGB image with all cars (ego + agents)\n    all_vehicles_image = ego_path_image_rgb + agents_path_image_rgb  # (batch_size, 3, H, H)\n\n    # get RGB image for scene from rasterizer (3 last images); channels last\n    scene_image_rgb = images[:, 2*hnf+2:, :, :].permute(0, 2, 3, 1)\n\n    scene_image_rgb[(all_vehicles_image > 0).any(dim=-1)] = 0.0\n\n    # generate final raster map\n    full_raster_image = (all_vehicles_image + scene_image_rgb).clamp(0, 1)\n\n    # channels as a second dimension\n    full_raster_image = full_raster_image.permute(0, 3, 1, 2)\n#     print('raster image executed')\n    return full_raster_image  # (batch_size, 3, W, W)\n","metadata":{"execution":{"iopub.status.busy":"2022-05-03T14:54:32.599308Z","iopub.execute_input":"2022-05-03T14:54:32.599628Z","iopub.status.idle":"2022-05-03T14:54:32.618342Z","shell.execute_reply.started":"2022-05-03T14:54:32.599595Z","shell.execute_reply":"2022-05-03T14:54:32.617316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef evaluate(cfg, predictor, data_loader):\n    \"\"\"\n    Evaluates model on the given dataset.\n    Returns a dictionary with metrics.\n    \"\"\"\n\n    predictor.eval()\n\n    metrics = {}\n    loss_l2_bo3_list = []\n    loss_l2_bo20_list = []\n\n    with torch.no_grad():\n\n        for batch in data_loader:\n\n            batch = [tensor.to(cfg['device']) for tensor in batch]\n            batch[0] = f_get_raster_image(cfg=cfg,\n                                          images=batch[0],\n                                          history_weight=cfg['model_params']['history_fading_weight'])\n\n            (image, target_positions, target_availabilities,\n             history_positions, history_yaws, centroid, world_to_image) = batch\n\n            actor_state = (history_positions, history_yaws)\n\n            # best of k average l2 loss (per trajectory/point)\n            loss_l2_bo3 = l2_loss_kmin(traj_real=target_positions,\n                                       generator_=predictor,\n                                       image=image,\n                                       actor_state=actor_state,\n                                       cfg=cfg,\n                                       kmin=3)  # (1,)\n\n            loss_l2_bo20 = l2_loss_kmin(traj_real=target_positions,\n                                        generator_=predictor,\n                                        image=image,\n                                        actor_state=actor_state,\n                                        cfg=cfg,\n                                        kmin=20)  # (1,)\n\n            loss_l2_bo3_list.append(loss_l2_bo3)\n            loss_l2_bo20_list.append(loss_l2_bo20)\n\n    loss_l2_bo3 = torch.stack(loss_l2_bo3_list, dim=0)  # (num_batches, )\n    loss_l2_bo20 = torch.stack(loss_l2_bo20_list, dim=0)\n\n    metrics['l2_best_of_3_loss'] = loss_l2_bo3.mean().item()\n    metrics['l2_best_of_20_loss'] = loss_l2_bo20.mean().item()\n\n    predictor.train()\n    return metrics","metadata":{"execution":{"iopub.status.busy":"2022-05-03T14:54:32.620226Z","iopub.execute_input":"2022-05-03T14:54:32.62063Z","iopub.status.idle":"2022-05-03T14:54:32.636565Z","shell.execute_reply.started":"2022-05-03T14:54:32.62057Z","shell.execute_reply":"2022-05-03T14:54:32.635497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.autograd import grad\n\n\ndef l2_loss_kmin(traj_real,\n                 generator_,\n                 image,\n                 actor_state,\n                 cfg,\n                 kmin,\n                 return_best_traj=False\n                 ):\n    \"\"\"\n    Apply current generator k times\n    and select minimal L2 distance to real trajectory among generated trajectories as loss.\n    \"\"\"\n\n    batch_size = image.shape[0]\n    noise_dim = cfg['gan_params']['noise_dim']\n    pred_len = cfg['model_params']['future_num_frames']\n\n    l2_losses = []\n    trajectories = []\n\n    # generate k trajectories\n    for i in range(kmin):\n        noise = torch.normal(size=(batch_size, noise_dim), mean=0.0, std=1.0,\n                             dtype=torch.float32, device=cfg['device'])\n\n        traj_fake = generator_(image, actor_state, noise)\n        trajectories.append(traj_fake)\n\n        current_l2_loss = l2_loss(traj_fake, traj_real)  # (batch_size, )\n        l2_losses.append(current_l2_loss)\n\n    # stack all k_min losses to select the best one\n    stacked_losses = torch.stack(l2_losses, dim=1)  # (batch_size, k_min)\n\n    # indices of best l2 loss for each element in batch\n    best_indices = stacked_losses.argmin(dim=-1)  # (batch_size,)\n\n    # select best loss for each element in batch\n    losses = torch.gather(stacked_losses, 1, best_indices[:, None])  # (batch_size, 1)\n\n    if cfg['losses']['variety_l2_mode'] == 'average':\n        # minimal average loss per trajectory point\n        l2_loss_ = torch.sum(losses) / (batch_size*pred_len)\n    elif cfg['losses']['variety_l2_mode'] == 'sum':\n        # minimal summary loss per whole trajectory\n        l2_loss_ = torch.sum(losses) / batch_size\n    else:\n        raise NotImplementedError\n\n    # return corresponding best trajectories\n    if return_best_traj:\n        stacked_trajectories = torch.stack(trajectories, dim=1)  # (batch_size, k_min, target_size, 2)\n        best_traj = stacked_trajectories[torch.arange(batch_size), best_indices]  # (batch_size, target_size, 2)\n        return l2_loss_, best_traj\n    else:\n        return l2_loss_\n\n\ndef l2_loss(traj_fake, traj_real):\n    \"\"\"\n    Returns summary losses for generated trajectories\n    traj_fake: Tensor of shape # (batch_size, target_size, 2). Predicted trajectory.\n    traj_real: Tensor of shape # (batch_size, target_size, 2). Ground truth predictions.\n    \"\"\"\n\n    loss = (traj_real - traj_fake)**2  # (batch_size, target_size, 2)\n\n    # batch of summary losses for each trajectory\n    loss = loss.sum(dim=2).sum(dim=1)  # (batch_size,)\n    return loss\n\n\ndef gradient_penalty(discrim, real_trajectory, fake_trajectory, in_image, in_actor_state, lambda_gp, device):\n    \"\"\"\n    Calculates the gradient penalty loss for Wasserstein GAN with GP (https://arxiv.org/abs/1704.00028)\n    in_image - scene context image\n    in_actor_state - history of actor states\n    \"\"\"\n\n    # Random weight term of shape (batch_size, target_size, 2) for interpolation between real and fake samples\n    batch_size = real_trajectory.shape[0]\n    epsilon = torch.rand(size=(batch_size, 1, 1), device=device)  # (batch_size, 1, 1)\n\n    # Get random interpolation between real and fake samples\n    interpolates = (epsilon * real_trajectory + ((1 - epsilon) * fake_trajectory)).requires_grad_(True)  # (batch_size, target_size, 2)\n    d_interpolates = discrim(interpolates, in_image, in_actor_state)  # (batch_size, 1)\n\n    fake = torch.ones(size=(batch_size, 1), requires_grad=False, device=device)\n\n    # Get gradient of d_interpolates w.r.t. interpolates:\n    # We use torch.autograd.grad and set grad_output==1.\n\n    gradients = grad(\n        outputs=d_interpolates,\n        inputs=interpolates,\n        grad_outputs=fake,\n        create_graph=True,\n        retain_graph=True,\n        only_inputs=True,\n    )[0]  # (batch_size, target_size, 2)\n\n    gradients = gradients.view(gradients.size(0), -1)  # (batch_size, 2*target_size)\n    gradient_penalty_ = ((gradients.norm(2, dim=1) - 1) ** 2).mean()\n\n    return lambda_gp * gradient_penalty_","metadata":{"execution":{"iopub.status.busy":"2022-05-03T14:54:32.640126Z","iopub.execute_input":"2022-05-03T14:54:32.640535Z","iopub.status.idle":"2022-05-03T14:54:32.663219Z","shell.execute_reply.started":"2022-05-03T14:54:32.640489Z","shell.execute_reply":"2022-05-03T14:54:32.661853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DifferentionalRasterizerLayer(nn.Module):\n    print(\"DifferentionalRasterizer\")\n    \"\"\"\n    Differential trajectory rasterizer from the article https://arxiv.org/abs/2004.06247\n    It provides an width x width image for each point of generated trajectory.\n    Each image is in separate channel (trajectory grids).\n    \"\"\"\n\n    def __init__(self,\n                 width,\n                 h_0,\n                 w_0,\n                 r,\n                 sigma,\n                 device):\n\n        super().__init__()\n\n        self.W = width\n        self.h_0 = h_0\n        self.w_0 = w_0\n        self.r = r\n        self.sigma = sigma\n        self.pi_ = np.pi\n        self.device = device\n        self.m = MultivariateNormal(torch.zeros(2, device=self.device),\n                                    torch.eye(2, device=self.device) * self.sigma**2)\n\n    def forward(self, input):\n        bs_ = input.shape[0]\n        hs_ = input.shape[1]\n\n        ones_ = torch.ones(size=(bs_, hs_, self.W, self.W), device=self.device)\n        ranges_ = torch.range(0, self.W-1, device=self.device).reshape((1, 1, 1, -1))\n\n        # element-wise product\n        delta_i = ones_*ranges_\n        delta_j = delta_i.permute((0, 1, 3, 2))\n\n        delta_i = (delta_i - self.h_0)*self.r\n        delta_j = (delta_j - self.w_0)*self.r\n\n        # delta (bs_, hs_, W, W, 2)\n        delta = torch.stack((delta_i, delta_j), dim=-1)\n        Delta = delta - input.reshape((bs_, hs_, 1, 1, 2))\n\n        G = torch.exp(self.m.log_prob(Delta))\n        print(G)\n\n        return G","metadata":{"execution":{"iopub.status.busy":"2022-05-03T14:54:32.66564Z","iopub.execute_input":"2022-05-03T14:54:32.666298Z","iopub.status.idle":"2022-05-03T14:54:32.681864Z","shell.execute_reply.started":"2022-05-03T14:54:32.666249Z","shell.execute_reply":"2022-05-03T14:54:32.680627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LSTMEncoder(nn.Module):\n    \"\"\"\n    Conv1D + LSTM layers to encode a sequence of ego-car's coordinates and yaws.\n    \"\"\"\n\n    def __init__(self,\n                 embedding_dim=128,\n                 h_dim=128,\n                 dropout=0.0):\n\n        super().__init__()\n\n        self.h_dim = h_dim\n        self.embedding_dim = embedding_dim\n        self.num_layers = 1\n\n        # (bs, history_size, 2) --> (bs, history_size, 128)\n        self.spatial_embedding = Conv1DEmbedder(embedding_dim=self.embedding_dim)\n\n        self.encoder = nn.LSTM(self.embedding_dim,\n                               self.h_dim,\n                               self.num_layers,\n                               dropout=dropout)\n\n    def forward(self, actor_state):\n        \"\"\"\n        Inputs: concatenated tuple of tensors\n        - history_positions: (batch, history_size, 2)\n        - history_yaws: (batch, history_size, 1)\n        Output:\n        - final_h: (batch, self.h_dim)\n        \"\"\"\n\n        # encode trajectory\n        history_data_embedding = self.spatial_embedding(actor_state)\n\n        output, state = self.encoder(history_data_embedding.permute(1, 0, 2))  # lstm input is (seq_len, batch, input_size)\n\n        final_h = state[0]\n#         print('lstm executed')\n\n        return final_h.permute(1, 0, 2).squeeze(1)\n\n\nclass Conv1DEmbedder(nn.Module):\n    def __init__(self, embedding_dim=128):\n\n        super().__init__()\n\n        self.embedding_dim = embedding_dim\n        self.conv1d = nn.Conv1d(3, self.embedding_dim, 3, padding=1)\n\n    def forward(self, history_data):\n        history_data = history_data.permute(0, 2, 1)  # (bs, 3, history_size)\n        history_data = self.conv1d(history_data)  # (bs, embedding_dim, history_size)\n#         print('1d-embedder executed')\n        return history_data.permute(0, 2, 1)  # (bs, history_size, embedding_dim)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T14:54:32.686272Z","iopub.execute_input":"2022-05-03T14:54:32.686743Z","iopub.status.idle":"2022-05-03T14:54:32.701758Z","shell.execute_reply.started":"2022-05-03T14:54:32.686693Z","shell.execute_reply":"2022-05-03T14:54:32.700654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Generator(nn.Module):\n    \"\"\"\n    Generator uses scene rgb image and actor state data (coordinates + yaws).\n    For image embedding we use pretrained MobileNet or EfficientNet as backbone.\n    For actor state embedding we use shallow mlp layer or conv1d + lstm.\n    \"\"\"\n\n    def __init__(self,\n                 input_dim,\n                 embedding_dim,\n                 decoder_dim,\n                 trajectory_dim,\n                 noise_dim,\n                 backbone_type,\n                 embedding_type):\n\n        super().__init__()\n\n        self.input_dim = input_dim\n        self.embedding_dim = embedding_dim\n        self.decoder_dim = decoder_dim\n        self.trajectory_dim = trajectory_dim\n        self.noise_dim = noise_dim\n        self.embedding_type = embedding_type\n\n        print('Backbone Type:', backbone_type)\n\n        if backbone_type == 'mobilenet':\n            self.backbone = mobilenet_v2(pretrained=True)\n            self.extracted_features = list(self.backbone.classifier.children())[-1].out_features  # 1000\n        elif 'efficientnet' in backbone_type:\n            self.backbone = EfficientNet.from_pretrained(backbone_type)\n            self.extracted_features = self.backbone._fc.out_features  # 1000\n        else:\n            raise NotImplementedError\n\n        if self.embedding_type == 'mlp':\n            self.state_encoder = nn.Linear(self.input_dim, self.embedding_dim)\n        elif self.embedding_type == 'lstm':\n            self.state_encoder = LSTMEncoder(embedding_dim=self.embedding_dim,\n                                             h_dim=self.embedding_dim,\n                                             dropout=0.0)\n        else:\n            raise NotImplementedError\n\n        # we add noise_dim for latent variable (noise)\n        self.decoder_1 = nn.Linear(self.extracted_features + self.embedding_dim + self.noise_dim,\n                                   self.decoder_dim)\n        # we predict flattened trajectory vector for both x and y\n        self.decoder_2 = nn.Linear(self.decoder_dim, 2*self.trajectory_dim)\n        print('End')\n\n    def forward(self, image, actor_state, noise):\n#         image = (batch_size, 3, W, W)\n#         actor_state = [(batch_size, h_s, 2), (batch_size, h_s, 1)] - pair: ego-car coordinates + yaws history\n#         noise = (batch_size, noise_dim)\n        #print('forward executed')\n        batch_size = image.shape[0]\n\n        # flattened scene context for image\n        scene_context = self.backbone(image)  # (batch_size, self.extracted_features)\n\n        actor_state = torch.cat(actor_state, dim=-1)  # (batch_size, h_s, 3)\n\n        # flatten input (x, y, angle) for shallow layer\n        if self.embedding_type == 'mlp':\n            actor_state = actor_state.reshape(batch_size, -1)  # (batch_size, 3*h_s)\n\n        # actor_state embedding\n        encoded_state = self.state_encoder(actor_state)  # (batch_size, self.embedding_dim)\n        trajectory_decoded=0\n\n\n        # concatenate all inputs\n        concat = torch.cat([scene_context, encoded_state, noise], dim=-1)  # (batch_size, self.extracted_features +\n                                                                           #  self.embedding_dim + self.noise_dim)\n\n        trajectory_decoded = self.decoder_1(concat)  # (batch_size, decoder_dim)\n        trajectory_decoded = self.decoder_2(trajectory_decoded)  # (batch_size, 2*trajectory_dim)\n\n        # reshape flattened vector to get (x,y) pairs\n        trajectory_decoded = trajectory_decoded.reshape(batch_size, self.trajectory_dim, 2)  # (batch_size, self.trajectory_dim, 2)\n        print(\"decoded trajectory\")\n        print(trajectory_decoded)\n        return trajectory_decoded","metadata":{"execution":{"iopub.status.busy":"2022-05-03T14:54:32.703656Z","iopub.execute_input":"2022-05-03T14:54:32.703998Z","iopub.status.idle":"2022-05-03T14:54:32.723117Z","shell.execute_reply.started":"2022-05-03T14:54:32.703953Z","shell.execute_reply":"2022-05-03T14:54:32.722278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Discriminator(nn.Module):\n    \"\"\"\n    Discriminator that uses differentiable  rasterization from article https://arxiv.org/abs/2004.06247\n    for generated trajectory along with scene context image\n    and history of actor states, that we add via so-called Fusion, as  proposed in https://arxiv.org/abs/1906.08469 .\n    Three types of GAN architectures supported: Vanilla DCGAN, Wasserstein GAN and Wasserstein GAN with Gradient Penalty.\n    \"\"\"\n\n    def __init__(self,\n                 width,\n                 h_0,\n                 w_0,\n                 r,\n                 sigma,\n                 channels_num,\n                 num_disc_feats,\n                 input_dim,\n                 device,\n                 gan_type,\n                 embedding_type,\n                 lstm_embedding_dim):\n        super().__init__()\n\n        self.width = width\n        self.h_0 = h_0\n        self.w_0 = w_0\n        self.r = r\n        self.sigma = sigma\n        self.channels_num = channels_num\n        self.num_disc_feats = num_disc_feats\n        self.input_dim = input_dim\n        self.device = device\n        self.gan_type = gan_type\n        self.embedding_type = embedding_type\n        self.lstm_embedding_dim = lstm_embedding_dim\n\n        self.diff_rasterizer = DifferentionalRasterizerLayer(self.width,\n                                                             self.h_0,\n                                                             self.w_0,\n                                                             self.r,\n                                                             self.sigma,\n                                                             self.device)\n\n        if gan_type == 'wasserstein_gp':\n            self.fusion_dcgan = FusionDCGAN_gp(nc=self.channels_num,\n                                               ndf=self.num_disc_feats,\n                                               input_dim=self.input_dim,\n                                               embedding_type=self.embedding_type,\n                                               lstm_embedding_dim=self.lstm_embedding_dim\n                                               )\n        else:  # vanilla/wasserstein\n            self.fusion_dcgan = FusionDCGAN(nc=self.channels_num,\n                                            ndf=self.num_disc_feats,\n                                            input_dim=self.input_dim,\n                                            gan_type=self.gan_type,\n                                            embedding_type=self.embedding_type,\n                                            lstm_embedding_dim=self.lstm_embedding_dim\n                                            )\n\n    def forward(self,\n                trajectory,\n                image,\n                actor_state):\n        # trajectory: (batch_size, target_size, 2) - predicted or ground truth trajectory\n        # image: (batch_size, 3, W, W) - rasterized scene context image\n        # actor_state = [(batch_size, h_s, 2), (batch_size, h_s, 1)] - coordinates + yaws\n\n        batch_size = image.shape[0]\n\n        # generate N = target_size images (grids)\n        # there is an image for each point in predicted/ground truth trajectory\n        trajectory_grids = self.diff_rasterizer(trajectory)  # (batch_size, target_size, W, W)\n\n        # add scene context image\n        trajectory_grids = torch.cat((trajectory_grids, image), dim=1) \n        print(\"dis_forward_executed\")\n        # (batch_size, target_size+3, W, W)\n        # add history of actor states via Fusion and process through DCGAN\n        out = self.fusion_dcgan(trajectory_grids, actor_state)  # (batch_size, 1, 1, 1)\n        out = out.reshape(batch_size, -1)\n        print(out)\n\n        return out","metadata":{"execution":{"iopub.status.busy":"2022-05-03T14:54:32.724991Z","iopub.execute_input":"2022-05-03T14:54:32.725456Z","iopub.status.idle":"2022-05-03T14:54:32.743312Z","shell.execute_reply.started":"2022-05-03T14:54:32.725367Z","shell.execute_reply":"2022-05-03T14:54:32.742229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def weights_init(m):\n    classname = m.__class__.__name__\n    if classname.find('Conv') != -1 and classname != 'Conv1DEmbedder':\n        nn.init.normal_(m.weight.data, 0.0, 0.02)\n    elif classname.find('BatchNorm') != -1:\n        nn.init.normal_(m.weight.data, 1.0, 0.02)\n        nn.init.constant_(m.bias.data, 0)\n\ndef tensor_to_image(tensor):\n    tensor = tensor*255\n    tensor=tensor.permute(0,2,3,1)\n    tensor = np.array(tensor.cpu(), dtype=np.uint8)\n    i=0\n    if np.ndim(tensor)>3:\n        for t in tensor:\n            img=PIL.Image.fromarray(t)\n            img.save(f\"./{i}.jpeg\")\n            i+=1","metadata":{"execution":{"iopub.status.busy":"2022-05-03T14:54:32.745322Z","iopub.execute_input":"2022-05-03T14:54:32.745918Z","iopub.status.idle":"2022-05-03T14:54:32.758816Z","shell.execute_reply.started":"2022-05-03T14:54:32.74587Z","shell.execute_reply":"2022-05-03T14:54:32.757633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install python-benedict\nfrom benedict import benedict\npath='../input/configuration/cfg.yaml'\ncfg = benedict.from_yaml(path)\ncfg[\"train_data_loader\"][\"key\"]='scenes/sample.zarr'\nprint(cfg)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T14:54:32.760705Z","iopub.execute_input":"2022-05-03T14:54:32.761082Z","iopub.status.idle":"2022-05-03T14:54:52.324324Z","shell.execute_reply.started":"2022-05-03T14:54:32.761002Z","shell.execute_reply":"2022-05-03T14:54:52.323111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"    os.environ[\"L5KIT_DATA_FOLDER\"] = cfg[\"l5kit_data_folder\"]\n    dm = LocalDataManager(None)\n\n    train_cfg = cfg[\"train_data_loader\"]\n    valid_cfg = cfg[\"valid_data_loader\"]\n\n    # rasterizer\n    rasterizer = build_rasterizer(cfg, dm)\n\n    train_path = train_cfg[\"key\"]\n    train_zarr = ChunkedDataset(dm.require(train_path)).open(cached=False)\n\n\n    train_agent_dataset = AgentDataset(cfg, train_zarr, rasterizer)\n\n    # transform dataset to the proper frame of reference\n    train_dataset = TransformDataset(train_agent_dataset, cfg)\n\n    if not train_cfg['subset'] == -1:\n        train_dataset = Subset(train_dataset, np.arange(train_cfg['subset']))\n\n    train_loader = DataLoader(train_dataset,\n                              shuffle=train_cfg[\"shuffle\"],\n                              batch_size=train_cfg[\"batch_size\"],\n                              num_workers=train_cfg[\"num_workers\"])\n\n\n    # loading custom mask for validation dataset\n#     logger.info(f\"Loading val mask in path {valid_cfg['mask_path']}\")\n#     val_custom_mask = np.load(valid_cfg['mask_path'])\n#     logger.info(f\"Length of validation mask is: {val_custom_mask.sum()}\")\n\n    valid_path = valid_cfg[\"key\"]\n    valid_zarr = ChunkedDataset(dm.require(valid_path)).open(cached=False)\n\n\n\n    valid_agent_dataset = AgentDataset(cfg, valid_zarr, rasterizer)\n\n    # transform validation dataset to the proper frame of reference\n    valid_dataset = TransformDataset(valid_agent_dataset, cfg)\n\n    if not valid_cfg['subset'] == -1:\n        valid_dataset = Subset(valid_dataset, valid_cfg['subset'])\n\n    valid_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    )\n\n\n    n_epochs = cfg['train_params']['num_epochs']\n\n    d_steps = cfg['train_params']['num_d_steps']\n    g_steps = cfg['train_params']['num_g_steps']\n\n    noise_dim = cfg['gan_params']['noise_dim']\n    g_learning_rate = cfg['train_params']['g_learning_rate']\n    d_learning_rate = cfg['train_params']['d_learning_rate']\n\n    if cfg['gan_params']['gan_type'] == 'vanilla':\n        cross_entropy = nn.BCELoss()\n ","metadata":{"execution":{"iopub.status.busy":"2022-05-03T14:54:52.329917Z","iopub.execute_input":"2022-05-03T14:54:52.332742Z","iopub.status.idle":"2022-05-03T15:00:19.105042Z","shell.execute_reply.started":"2022-05-03T14:54:52.332693Z","shell.execute_reply":"2022-05-03T15:00:19.103819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\n\ndef trainer():\n        generator = Generator(input_dim=cfg['gan_params']['input_dim'],\n                              embedding_dim=cfg['gan_params']['embedding_dim'],\n                              decoder_dim=cfg['gan_params']['decoder_dim'],\n                              trajectory_dim=cfg['model_params']['future_num_frames'],\n                              noise_dim=noise_dim,\n                              backbone_type=cfg['gan_params']['backbone_type'],\n                              embedding_type=cfg['gan_params']['embedding_type']\n                              )\n        generator.to(cfg['device'])\n        generator.train()\n        \n        \n        W = cfg['raster_params']['raster_size'][0]\n        \n        \n        discriminator = Discriminator(width=W,\n                                  h_0=cfg['raster_params']['ego_center'][0]*W,\n                                  w_0=cfg['raster_params']['ego_center'][1]*W,\n                                  r=cfg['raster_params']['pixel_size'][0],\n                                  sigma=cfg['gan_params']['sigma'],\n                                  channels_num=cfg['model_params']['future_num_frames']+3,\n                                  num_disc_feats=cfg['gan_params']['num_disc_feats'],\n                                  input_dim=cfg['gan_params']['input_dim'],\n                                  device=cfg['device'],\n                                  gan_type=cfg['gan_params']['gan_type'],\n                                  embedding_type=cfg['gan_params']['embedding_type'],\n                                  lstm_embedding_dim=cfg['gan_params']['embedding_dim']\n                                  )\n        discriminator.to(cfg['device'])\n        discriminator.apply(weights_init)\n        discriminator.train()\n        \n        if cfg['gan_params']['gan_type'] == 'wasserstein':\n            optimizer_g = optim.RMSprop(generator.parameters(), lr=g_learning_rate)\n            optimizer_d = optim.RMSprop(discriminator.parameters(), lr=d_learning_rate)\n        elif cfg['gan_params']['gan_type'] == 'wasserstein_gp':\n            betas = (0.0, 0.9)\n            optimizer_g = optim.Adam(generator.parameters(), lr=g_learning_rate, betas=betas)\n            optimizer_d = optim.Adam(discriminator.parameters(), lr=d_learning_rate, betas=betas)\n        else:\n            optimizer_g = optim.Adam(generator.parameters(), lr=g_learning_rate)\n            optimizer_d = optim.Adam(discriminator.parameters(), lr=d_learning_rate)\n\n        d_steps_left = d_steps\n        g_steps_left = g_steps\n\n        # variables for statistics\n        d_full_loss = []\n        g_full_loss = []\n        gp_values = []\n        l2_variety_values = []\n        metric_vals = []\n\n        # checkpoint dictionary\n        checkpoint = {\n            'G_losses': defaultdict(list),\n            'D_losses': defaultdict(list),\n            'counters': {\n                't': None,\n                'epoch': None,\n            },\n            'g_state': None,\n            'g_optim_state': None,\n            'd_state': None,\n            'd_optim_state': None\n        }\n\n        id_batch = 0\n\n        # total number of batches\n        len_of_epoch = len(train_loader)\n        \n        for epoch in range(3):\n            print(\"epoch \",epoch,\":\")\n            for batch in tqdm(train_loader,total=len(train_loader)):\n                batch = [tensor.to(cfg['device']) for tensor in batch]\n\n                # Creates single raster image from sequence of images from l5kit's AgentDataset\n                batch[0] = f_get_raster_image(cfg=cfg,\n                                              images=batch[0],\n                                              history_weight=cfg['model_params']['history_fading_weight'])\n                tensor_to_image(batch[0])\n                (image, target_positions, target_availabilities,\n                 history_positions, history_yaws, centroid, world_to_image) = batch\n\n                actor_state = (history_positions, history_yaws)\n\n                batch_size = image.shape[0]\n\n                # noise for generator\n                noise = torch.normal(size=(batch_size, noise_dim),\n                                     mean=0.0,\n                                     std=1.0,\n                                     dtype=torch.float32,\n                                     device=cfg['device'])\n                fake_trajectory = generator(image, actor_state, noise)\n\n\n                if d_steps_left > 0:\n                    d_steps_left -= 1\n\n                    for pd in discriminator.parameters():  # reset requires_grad\n                        pd.requires_grad = True  # they are set to False below in generator update\n\n                    # freeze generator while training discriminator\n                    for pg in generator.parameters():\n                        pg.requires_grad = False\n\n                    discriminator.zero_grad()\n\n                    # generate fake trajectories (batch_size, target_size, 2) for current batch\n                    #fake_trajectory = generator(image, actor_state, noise)\n\n                    # discriminator predictions (batch_size, 1) on real and fake trajectories\n                    d_real_pred = discriminator(target_positions, image, actor_state)\n                    d_g_pred = discriminator(fake_trajectory, image, actor_state)\n                    print(\"real_pred done\")\n                    print(\"g_pred done\")\n\n                    # loss\n                    if cfg['gan_params']['gan_type'] == 'vanilla':\n                        # tensor with true/fake labels of size (batch_size, 1)\n                        real_labels = torch.full((batch_size,), 1, dtype=torch.float, device=cfg['device'])\n                        fake_labels = torch.full((batch_size,), 0, dtype=torch.float, device=cfg['device'])\n\n                        real_loss = cross_entropy(d_real_pred, real_labels)\n                        fake_loss = cross_entropy(d_g_pred, fake_labels)\n\n                        total_loss = real_loss + fake_loss\n                    elif cfg['gan_params']['gan_type'] == 'wasserstein':  # D(fake) - D(real)\n                        total_loss = torch.mean(d_g_pred) - torch.mean(d_real_pred)\n                    elif cfg['gan_params']['gan_type'] == 'wasserstein_gp':\n                        gp_loss = gradient_penalty(discrim=discriminator,\n                                                   real_trajectory=target_positions,\n                                                   fake_trajectory=fake_trajectory,\n                                                   in_image=image,\n                                                   in_actor_state=actor_state,\n                                                   lambda_gp=cfg['losses']['lambda_gp'],\n                                                   device=cfg['device'])\n\n                        total_loss = torch.mean(d_g_pred) - torch.mean(d_real_pred) + gp_loss\n                    else:\n                        raise NotImplementedError\n\n                    # calculate gradients for this batch\n                    total_loss.backward()\n                    optimizer_d.step()\n\n                    # weight clipping for discriminator in pure Wasserstein GAN\n                    if cfg['gan_params']['gan_type'] == 'wasserstein':\n                        c = cfg['losses']['weight_clip']\n                        for p in discriminator.parameters():\n                            p.data.clamp_(-c, c)\n\n                    d_full_loss.append(total_loss.item())\n                    print(\"d_full_loss\",d_full_loss)\n\n                    if cfg['gan_params']['gan_type'] == 'wasserstein_gp':\n                        gp_values.append(gp_loss.item())\n\n                #######################################\n                #         TRAIN GENERATOR\n                #######################################\n\n                elif g_steps_left > 0:  # we either train generator or discriminator on current batch\n                    g_steps_left -= 1\n\n                    for pd in discriminator.parameters():\n                        pd.requires_grad = False  # avoid discriminator training\n\n                    # unfreeze generator\n                    for pg in generator.parameters():\n                        pg.requires_grad = True\n\n                    generator.zero_grad()\n\n                    if cfg['losses']['use_variety_l2']:\n                        l2_variety_loss, fake_trajectory = l2_loss_kmin(traj_real=target_positions,\n                                                                        generator_=generator,\n                                                                        image=image,\n                                                                        actor_state=actor_state,\n                                                                        cfg=cfg,\n                                                                        kmin=cfg['losses']['k_min'],\n                                                                        return_best_traj=True)\n                    else:\n                        fake_trajectory = generator(image, actor_state, noise)\n\n                    d_g_pred = discriminator(fake_trajectory, image, actor_state)\n\n                    if cfg['gan_params']['gan_type'] == 'vanilla':\n                        # while training generator we associate generated fake examples\n                        # with real labels in order to measure generator quality\n                        real_labels = torch.full((batch_size,), 1, dtype=torch.float, device=cfg['device'])\n                        fake_loss = cross_entropy(d_g_pred, real_labels)\n                    elif cfg['gan_params']['gan_type'] in ['wasserstein', 'wasserstein_gp']:  # -D(fake)\n                        fake_loss = -torch.mean(d_g_pred)\n                    else:\n                        raise NotImplementedError\n\n                    if cfg['losses']['use_variety_l2']:\n                        fake_loss += cfg['losses']['weight_variety_l2'] * l2_variety_loss\n\n                        l2_variety_values.append(l2_variety_loss.item())\n\n                    fake_loss.backward()\n                    optimizer_g.step()\n\n                    g_full_loss.append(fake_loss.item())\n                    print(\"g_loss: \",g_full_loss)\n                \n             if d_steps_left == 0 and g_steps_left == 0:\n                d_steps_left = d_steps\n                g_steps_left = g_steps\n\n            # print current model state on train dataset\n             if (id_batch > 0) and (id_batch % cfg['train_params']['print_every_n_steps'] == 0):\n\n                print_statistics(logger=logger,\n                                 cfg=cfg,\n                                 epoch=epoch,\n                                 len_of_epoch=len_of_epoch,\n                                 id_batch=id_batch,\n                                 d_full_loss=d_full_loss,\n                                 g_full_loss=g_full_loss,\n                                 gp_values=gp_values,\n                                 l2_variety_values=l2_variety_values,\n                                 print_over_n_last=1000)\n                \n            id_batch = id_batch + 1\n\n\n\n\n        ","metadata":{"execution":{"iopub.status.busy":"2022-05-03T15:00:19.108516Z","iopub.execute_input":"2022-05-03T15:00:19.109101Z","iopub.status.idle":"2022-05-03T15:00:19.151481Z","shell.execute_reply.started":"2022-05-03T15:00:19.109052Z","shell.execute_reply":"2022-05-03T15:00:19.150411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\n","metadata":{}}]}