{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Objec Detection\n\n![](http://www.l5kit.org/_images/av.jpg)\n"},{"metadata":{"trusted":true,"_kg_hide-input":true,"_kg_hide-output":true},"cell_type":"code","source":"!pip install pytorch-pfn-extras==0.2.1","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-input":true},"cell_type":"code","source":"import gc\nimport os\nfrom pathlib import Path\nimport random\nimport sys\n\nfrom tqdm.notebook import tqdm\nimport numpy as np\nimport pandas as pd\nimport scipy as sp\n\n\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":{"trusted":true,"_kg_hide-input":true,"_kg_hide-output":true},"cell_type":"code","source":"import zarr\n\nimport l5kit\nfrom l5kit.data import ChunkedDataset, LocalDataManager\nfrom l5kit.dataset import EgoDataset, AgentDataset\n\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\nfrom tqdm import tqdm\nfrom collections import Counter\nfrom l5kit.data import PERCEPTION_LABELS\nfrom prettytable import PrettyTable\n\nfrom matplotlib import animation, rc\nfrom IPython.display import HTML\n\nrc('animation', html='jshtml')\nprint(\"l5kit version:\", l5kit.__version__)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-input":true},"cell_type":"code","source":"import torch\nfrom pathlib import Path\n\nimport pytorch_pfn_extras as ppe\nfrom math import ceil\nfrom pytorch_pfn_extras.training import IgniteExtensionsManager\nfrom pytorch_pfn_extras.training.triggers import MinValueTrigger\n\nfrom torch import nn, optim\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data.dataset import Subset\nimport pytorch_pfn_extras.training.extensions as E","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-input":true},"cell_type":"code","source":"# --- Dataset utils ---\nfrom typing import Callable\n\nfrom torch.utils.data.dataset import Dataset\n\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)\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Function\n\nTo define loss function to calculate competition evaluation metric **in batch**.<br/>\nIt works with **pytorch tensor, so it is differentiable** and can be used for training Neural Network."},{"metadata":{"trusted":true,"_kg_hide-input":true},"cell_type":"code","source":"# --- Function utils ---\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    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":{},"cell_type":"markdown","source":"## Model\n\npytorch model definition. Here model outputs both **multi-mode trajectory prediction & confidence of each trajectory**."},{"metadata":{"trusted":true,"_kg_hide-input":false},"cell_type":"code","source":"# --- Model utils ---\nimport torch\nfrom torchvision.models import resnet18\nfrom torch import nn\nfrom typing import Dict\n\n\nclass LyftMultiModel(nn.Module):\n\n    def __init__(self, cfg: Dict, num_modes=3):\n        super().__init__()\n\n        # TODO: support other than resnet18?\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            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        backbone_out_features = 512\n\n        # X, Y coords for the future positions (output shape: Bx50x2)\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 (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\n\n    ","execution_count":null,"outputs":[]},{"metadata":{"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\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-input":true},"cell_type":"code","source":"# --- Training utils ---\nfrom ignite.engine import Engine\n\n\ndef create_trainer(model, optimizer, device) -> Engine:\n    model.to(device)\n\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":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"# Modified to work with pytorch_pfn_extras\n\nimport os\nimport sys\nfrom copy import deepcopy\n\nfrom IPython.core.display import display\nfrom ipywidgets import HTML\n\nfrom pytorch_pfn_extras.training.extensions.print_report import PrintReport\n\nfrom pytorch_pfn_extras.training import extension\nfrom pytorch_pfn_extras.training.extensions import log_report \\\n    as log_report_module\nfrom pytorch_pfn_extras.training.extensions import util\n\n\nclass PrintReportNotebook(PrintReport):\n\n    \"\"\"An extension to print the accumulated results.\n\n    This extension uses the log accumulated by a :class:`LogReport` extension\n    to print specified entries of the log in a human-readable format.\n\n    Args:\n        entries (list of str ot None): List of keys of observations to print.\n            If `None` is passed, automatically infer keys from reported dict.\n        log_report (str or LogReport): Log report to accumulate the\n            observations. This is either the name of a LogReport extensions\n            registered to the manager, or a LogReport instance to use\n            internally.\n        out: Stream to print the bar. Standard output is used by default.\n\n    \"\"\"\n\n    def __init__(self, entries=None, log_report='LogReport', out=sys.stdout):\n        super(PrintReportNotebook, self).__init__(entries=entries, log_report=log_report, out=out)\n        self._widget = HTML()\n\n    def initialize(self, trainer):\n        display(self._widget)\n\n    @property\n    def widget(self):\n        return self._widget\n\n    def __call__(self, manager):\n        log_report = self.get_log_report(manager)\n        df = log_report.to_dataframe()\n        if self._infer_entries:\n            # --- update entries ---\n            self._update_entries(log_report)\n        self._widget.value = df[self._entries].to_html(index=False, na_rep='')\n","execution_count":null,"outputs":[]},{"metadata":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"# Code referenced from https://github.com/grafi-tt/chaineripy/blob/master/chaineripy/extensions/progress_bar.py by @grafi-tt\n# Modified to work with pytorch_pfn_extras\n\nfrom pytorch_pfn_extras.training import extension, trigger\nimport datetime\nimport time\n\nfrom IPython.core.display import display\nfrom ipywidgets import FloatProgress, HBox, HTML, VBox\n\n\nclass ProgressBarNotebook(extension.Extension):\n\n    \"\"\"Trainer extension to print a progress bar and recent training status.\n    This extension prints a progress bar at every call. It watches the current\n    iteration and epoch to print the bar.\n    Args:\n        training_length (tuple): Length of whole training. It consists of an\n            integer and either ``'epoch'`` or ``'iteration'``. If this value is\n            omitted and the stop trigger of the trainer is\n            :class:`IntervalTrigger`, this extension uses its attributes to\n            determine the length of the training.\n        update_interval (int): Number of iterations to skip printing the\n            progress bar.\n        bar_length (int): Length of the progress bar in characters.\n        out: Stream to print the bar. Standard output is used by default.\n    \"\"\"\n\n    def __init__(self, training_length=None, update_interval=100,\n                 bar_length=50):\n        self._training_length = training_length\n        if training_length is not None:\n            self._init_status_template()\n        self._update_interval = update_interval\n        self._recent_timing = []\n\n        self._total_bar = FloatProgress(description='total',\n                                        min=0, max=1, value=0,\n                                        bar_style='info')\n        self._total_html = HTML()\n        self._epoch_bar = FloatProgress(description='this epoch',\n                                        min=0, max=1, value=0,\n                                        bar_style='info')\n        self._epoch_html = HTML()\n        self._status_html = HTML()\n\n        self._widget = VBox([HBox([self._total_bar, self._total_html]),\n                             HBox([self._epoch_bar, self._epoch_html]),\n                             self._status_html])\n\n    def initialize(self, manager):\n        if self._training_length is None:\n            t = manager._stop_trigger\n            if not isinstance(t, trigger.IntervalTrigger):\n                raise TypeError(\n                    'cannot retrieve the training length from %s' % type(t))\n            self._training_length = t.period, t.unit\n            self._init_status_template()\n\n        updater = manager.updater\n        self.update(updater.iteration, updater.epoch_detail)\n        display(self._widget)\n\n    def __call__(self, manager):\n        length, unit = self._training_length\n\n        updater = manager.updater\n        iteration, epoch_detail = updater.iteration, updater.epoch_detail\n\n        if unit == 'iteration':\n            is_finished = iteration == length\n        else:\n            is_finished = epoch_detail == length\n\n        if iteration % self._update_interval == 0 or is_finished:\n            self.update(iteration, epoch_detail)\n\n    def finalize(self):\n        if self._total_bar.value != 1:\n            self._total_bar.bar_style = 'warning'\n            self._epoch_bar.bar_style = 'warning'\n\n    @property\n    def widget(self):\n        return self._widget\n\n    def update(self, iteration, epoch_detail):\n        length, unit = self._training_length\n\n        recent_timing = self._recent_timing\n        now = time.time()\n\n        recent_timing.append((iteration, epoch_detail, now))\n\n        if unit == 'iteration':\n            rate = iteration / length\n        else:\n            rate = epoch_detail / length\n        self._total_bar.value = rate\n        self._total_html.value = \"{:6.2%}\".format(rate)\n\n        epoch_rate = epoch_detail - int(epoch_detail)\n        self._epoch_bar.value = epoch_rate\n        self._epoch_html.value = \"{:6.2%}\".format(epoch_rate)\n\n        status = self._status_template.format(iteration=iteration,\n                                              epoch=int(epoch_detail))\n\n        if rate == 1:\n            self._total_bar.bar_style = 'success'\n            self._epoch_bar.bar_style = 'success'\n\n        old_t, old_e, old_sec = recent_timing[0]\n        span = now - old_sec\n        if span != 0:\n            speed_t = (iteration - old_t) / span\n            speed_e = (epoch_detail - old_e) / span\n        else:\n            speed_t = float('inf')\n            speed_e = float('inf')\n\n        if unit == 'iteration':\n            estimated_time = (length - iteration) / speed_t\n        else:\n            estimated_time = (length - epoch_detail) / speed_e\n        estimate = ('{:10.5g} iters/sec. Estimated time to finish: {}.'\n                    .format(speed_t,\n                            datetime.timedelta(seconds=estimated_time)))\n\n        self._status_html.value = status + estimate\n\n        if len(recent_timing) > 100:\n            del recent_timing[0]\n\n    def _init_status_template(self):\n        self._status_template = (\n            '{iteration:10} iter, {epoch} epoch / %s %ss<br />' %\n            self._training_length)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-input":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\n\nclass DotDict(dict):\n\n\n    __getattr__ = dict.get\n    __setattr__ = dict.__setitem__\n    __delattr__ = dict.__delitem__\n\n    ","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Configs"},{"metadata":{"trusted":true,"_kg_hide-input":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': 12,\n        'shuffle': True,\n        'num_workers': 4\n    },\n\n    'valid_data_loader': {\n        'key': 'scenes/validate.zarr',\n        'batch_size': 32,\n        'shuffle': False,\n        'num_workers': 4\n    },\n\n    'train_params': {\n        'max_num_steps': 10000,\n        'checkpoint_every_n_steps': 5000,\n\n        # 'eval_every_n_steps': -1\n    }\n}\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-input":true},"cell_type":"code","source":"flags_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\": \"cuda:0\",\n    \"out_dir\": \"results/multi_train\",\n    \"epoch\": 2,\n    \"snapshot_freq\": 50,\n}","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Main script\n\nNow finished defining all the util codes. Let's start writing main script to train the model!"},{"metadata":{"trusted":true,"_kg_hide-input":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\n\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-input":true,"_kg_hide-output":false},"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\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(1000))\ntrain_loader = DataLoader(train_dataset,\n                          shuffle=train_cfg[\"shuffle\"],\n                          batch_size=train_cfg[\"batch_size\"],\n                          num_workers=train_cfg[\"num_workers\"])\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))\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)\n\nprint(valid_agent_dataset)\nprint(\"# AgentDataset train:\", len(train_agent_dataset), \"#valid\", len(valid_agent_dataset))\nprint(\"# ActualDataset train:\", len(train_dataset), \"#valid\", len(valid_dataset))\n# AgentDataset train: 22496709 #valid 21624612\n# ActualDataset train: 100 #valid 100","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Prepare model & optimizer"},{"metadata":{"trusted":true},"cell_type":"code","source":"device = torch.device(flags.device)\n\nif flags.pred_mode == \"multi\":\n    predictor = LyftMultiModel(cfg)\n    model = LyftMultiRegressor(predictor)\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)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-input":true},"cell_type":"code","source":"# Train setup\ntrainer = create_trainer(model, optimizer, device)\n\n\ndef eval_func(*batch):\n    loss, metrics = model(*[elem.to(device) for elem in batch])\n\n\nvalid_evaluator = E.Evaluator(\n    valid_loader,\n    model,\n    progress_bar=False,\n    eval_func=eval_func,\n)\n\nlog_trigger = (10 if debug else 1000, \"iteration\")\nlog_report = E.LogReport(trigger=log_trigger)\n\n\nextensions = [\n    log_report,  # Save `log` to file\n    valid_evaluator,  # Run evaluation for valid dataset in each epoch.\n    # E.FailOnNonNumber()  # Stop training when nan is detected.\n]\n\nis_notebook = True  # Make it False when you run code in local machine using console.\nif is_notebook:\n    extensions.extend([\n        ProgressBarNotebook(update_interval=10 if debug else 100),  # Show progress bar during training\n        PrintReportNotebook(),  # Show \"log\" on jupyter notebook  \n    ])\nelse:\n    extensions.extend([\n        E.ProgressBar(update_interval=10 if debug else 100),  # Show progress bar during training\n        E.PrintReport(),  # Print \"log\" to terminal\n    ])\n\n\nepoch = flags.epoch\n\nmodels = {\"main\": model}\noptimizers = {\"main\": optimizer}\nmanager = IgniteExtensionsManager(\n    trainer,\n    models,\n    optimizers,\n    epoch,\n    extensions=extensions,\n    out_dir=str(out_dir),\n)\n# Save predictor.pt every epoch\nmanager.extend(E.snapshot_object(predictor, \"predictor.pt\"),\n               trigger=(flags.snapshot_freq, \"iteration\"))\n# Check & Save best validation predictor.pt every epoch\n# manager.extend(E.snapshot_object(predictor, \"best_predictor.pt\"),\n#                trigger=MinValueTrigger(\"validation/main/nll\", trigger=(flags.snapshot_freq, \"iteration\")))\n# --- lr scheduler ---\n# scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n#     optimizer, mode='min', factor=0.7, patience=5, min_lr=1e-10)\nscheduler = torch.optim.lr_scheduler.ExponentialLR(\n    optimizer, gamma=0.99999)\nmanager.extend(lambda manager: scheduler.step(), trigger=(1, \"iteration\"))\n# Show \"lr\" column in log\nmanager.extend(E.observe_lr(optimizer=optimizer), trigger=log_trigger)\n\ntrainer.run(train_loader, max_epochs=epoch)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"You can obtrain training history results really easily by just accessing `LogReport` class, which is useful for managing a lot of experiments during kaggle competitions."},{"metadata":{"trusted":true},"cell_type":"code","source":"df = log_report.to_dataframe()\ndf.to_csv(out_dir/\"log.csv\", index=False)\ndf[[\"epoch\", \"iteration\", \"main/loss\", \"main/nll\", \"validation/main/loss\", \"validation/main/nll\", \"lr\", \"elapsed_time\"]]","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}