{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "c56e1bfd-cd0e-46dd-803c-d6df8eafa16b",
   "metadata": {},
   "source": [
    "Introduction to PyTorch Lightning\n",
    "=================================\n",
    "\n",
    "Yet another Deep Learning framework?!\n",
    "-------------------------------------\n",
    "\n",
    "We learned how to train a neural network in PyTorch. PyTorch Lightning can be thought of as a \"wrapper\" around PyTorch. It is intended as a very flexible and versatile tool for AI researchers.\n",
    "\n",
    "\n",
    "\n",
    "**Benefits**\n",
    "\n",
    "- Clearly separate research and engineering code\n",
    "- Less \"boilerplate\" code - less sources of error\n",
    "- Easily scale from CPU to GPU to multiple GPUs on a high-performance computer\n",
    "- Production made easy\n",
    "\n",
    "**Drawbacks**\n",
    "\n",
    "- Learning curve\n",
    "- Requires some understanding of object oriented programming\n",
    "\n",
    "PyTorch vs Lightning\n",
    "--------------------\n",
    "\n",
    "Lightning builds on top of PyTorch. All of PyTorch is accessible and contained within Lightning.\n",
    "\n",
    "\n",
    "Installation\n",
    "------------\n",
    "\n",
    "`pip install lightning`\n",
    "\n",
    "Versioning follows PyTorch with short delay.\n",
    "\n",
    "\n",
    "Imports\n",
    "-------\n",
    "\n",
    "Standard import:\n",
    "\n",
    "`import lightning as L`"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9fbae6a4-0282-4cef-824d-0c0f30639239",
   "metadata": {},
   "outputs": [],
   "source": [
    "import lightning as L"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "ba0915ee-bcb8-48b7-ac34-559e5b994442",
   "metadata": {},
   "source": [
    "Lightning building blocks\n",
    "=========================\n",
    "\n",
    "To use Lightning, you need to re-organize your code by defining\n",
    "\n",
    "- `LightningModule`: defines the model architecture, optimizers, and behaviour during training / evaluation steps\n",
    "- `Trainer`: controls the training and evaluation loop\n",
    "\n",
    "Optionally, you can use the `LightningDataModule` as a wrapper for everything related to datasets and dataloaders."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "a01b4268-2bf5-4dd3-aede-9861223a3ad5",
   "metadata": {},
   "source": [
    "Model: `LightningModule`\n",
    "------------------------\n",
    "\n",
    "We will start by wrapping the deep learning model we developed in the previous exercise into a `LightningModule`. We first copy over the model from the previous exercise\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "bafe14cf-c59a-4cb8-9a22-3e8000ef3a5d",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import torch.nn as nn\n",
    "from torch.nn.functional import relu\n",
    "\n",
    "class Net(nn.Module):\n",
    "    def __init__(self):\n",
    "        super().__init__()\n",
    "        self.conv1 = nn.Conv2d(1, 6, 5)\n",
    "        self.pool = nn.MaxPool2d(2, 2)\n",
    "        self.conv2 = nn.Conv2d(6, 16, 5)\n",
    "        self.fc1 = nn.Linear(16 * 15 * 15, 120)\n",
    "        self.fc2 = nn.Linear(120, 84)\n",
    "        self.fc3 = nn.Linear(84, 12)\n",
    "\n",
    "    def forward(self, x):\n",
    "        x = self.pool(relu(self.conv1(x)))\n",
    "        x = self.pool(relu(self.conv2(x)))\n",
    "        x = torch.flatten(x, 1) # flatten all dimensions except batch\n",
    "        x = relu(self.fc1(x))\n",
    "        x = relu(self.fc2(x))\n",
    "        x = self.fc3(x)\n",
    "        return x"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "4417f804-2d91-49de-96b6-377df0774d35",
   "metadata": {},
   "source": [
    "Now we look at the structure of the `LightningModule`. It should define the following classes:\n",
    "\n",
    "```python\n",
    "class MyLightningModule(L.LightningModule):\n",
    "    def __init__(self, model, learning_rate):\n",
    "        super().__init__()\n",
    "        self.model = model\n",
    "        self.learning_rate = learning_rate\n",
    "        self.criterion = ... # define the loss criterion\n",
    "\n",
    "    def training_step(self, batch, batch_idx):\n",
    "        x, y = batch\n",
    "\n",
    "        # define what should happen in the training step\n",
    "        # MUST return the loss\n",
    "\n",
    "    def validation_step(self, batch, batch_idx):\n",
    "        x, y = batch\n",
    "\n",
    "        # define what should happen in the validation step\n",
    "\n",
    "    def configure_optimizers(self):\n",
    "        optimizer = torch.optim.SGD(self.model.parameters(), lr=self.learning_rate)\n",
    "        return optimizer\n",
    "```\n",
    "\n",
    "Let's use the following starter code to develop our own `LightningModule`.\n",
    "\n",
    "✏️ Complete the commented statements"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "202a3e2a-eea2-49e8-8d77-fbb6af99d625",
   "metadata": {},
   "outputs": [],
   "source": [
    "class MyLightningModule(L.LightningModule):\n",
    "    def __init__(self, model, learning_rate):\n",
    "        super().__init__()\n",
    "        self.model = model\n",
    "        self.learning_rate = learning_rate\n",
    "        self.loss_fn = ... # define the loss criterion\n",
    "\n",
    "    def training_step(self, batch, batch_idx):\n",
    "        x, y = batch\n",
    "\n",
    "        # define what should happen in the training step\n",
    "        # MUST return the loss\n",
    "\n",
    "    def validation_step(self, batch, batch_idx):\n",
    "        x, y = batch\n",
    "\n",
    "        # define what should happen in the validation step\n",
    "\n",
    "    def configure_optimizers(self):\n",
    "        optimizer = torch.optim.SGD(self.model.parameters(), lr=self.learning_rate)\n",
    "        return optimizer"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "14b64a2d-e01f-40d9-8694-601ca6c765b0",
   "metadata": {},
   "source": [
    "Data handling\n",
    "-------------\n",
    "\n",
    "We will reuse the data handling from the previous exercise. This cell is copied from the previous notebook, we just execute it here.\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "dfaa0b5f-a28d-4b71-a25c-ef2e3a082526",
   "metadata": {},
   "outputs": [],
   "source": [
    "import xarray as xr\n",
    "\n",
    "# We retrieve the absolute path of the file containing the temperatures\n",
    "path_20cr = %env path_20cr\n",
    "\n",
    "# We load the data into xarray datasets\n",
    "ds_20cr = xr.open_dataset(path_20cr)\n",
    "\n",
    "# We extract the months and create a tensor containing data and labels\n",
    "def create_dataset(ds):\n",
    "\n",
    "    n_time, n_lon, n_lat = ds[\"tas\"].shape\n",
    "    \n",
    "    targets = torch.empty(n_time, dtype=torch.long)\n",
    "    for i in range(n_time):\n",
    "        targets[i] = int(str(ds[\"time\"].values[i]).split(\"-\")[1]) - 1\n",
    "    \n",
    "    data = torch.as_tensor(ds[\"tas\"].values.reshape(n_time, 1, n_lon, n_lat))\n",
    "    data = (data - 273.15) / 20.5\n",
    "    dataset = torch.utils.data.TensorDataset(data, targets)\n",
    "    \n",
    "    return dataset, n_time\n",
    "\n",
    "data_20cr, n_time = create_dataset(ds_20cr)\n",
    "\n",
    "# We use a PyTorch function to split the dataset randomly\n",
    "n_valtest = int( n_time * 0.2 )\n",
    "n_training = n_time - 2 * n_valtest\n",
    "training_data, validation_data, test_data = torch.utils.data.random_split(data_20cr, [n_training, n_valtest, n_valtest])"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "357ad15b-1591-4a4e-92cb-b9827ed76020",
   "metadata": {},
   "source": [
    "Lightning also needs an iterator for the datasets. We use the PyTorch DataLoader again."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "425e8a54-562b-440d-be85-da3e560b7440",
   "metadata": {},
   "outputs": [],
   "source": [
    "# We use another PyTorch function that makes the iteration over the batches of the datasets easier\n",
    "from torch.utils.data import DataLoader\n",
    "\n",
    "batch_size = 16\n",
    "\n",
    "train_dataloader = DataLoader(training_data, batch_size=batch_size, shuffle=True)\n",
    "validation_dataloader = DataLoader(validation_data, batch_size=batch_size, shuffle=True)\n",
    "test_dataloader = DataLoader(test_data, batch_size=n_valtest, shuffle=False)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "9cee3089-ce2e-45be-a703-46c81baf9ea5",
   "metadata": {},
   "source": [
    "`Trainer`\n",
    "---------\n",
    "\n",
    "The Lightning `Trainer` is a wrapper for the full training loop. It controls for example the number of epochs we use for training.\n",
    "\n",
    "Looking back at the plain PyTorch code from the last exercise, this was the training loop:\n",
    "\n",
    "```python\n",
    "train_loss = np.zeros(num_epochs)\n",
    "val_loss = np.zeros(num_epochs)\n",
    "for epoch in tqdm(range(num_epochs)):\n",
    "    \n",
    "    for i, data in enumerate(train_dataloader, 0):\n",
    "        # We load the batch of samples\n",
    "        inputs, labels = data\n",
    "\n",
    "        # We reset the gradients to zero to avoid accumulation over iteration\n",
    "        optimizer.zero_grad()\n",
    "\n",
    "        # Forward pass\n",
    "        model.train() # If we use model.eval() later\n",
    "        outputs = model(inputs)\n",
    "        \n",
    "        # We calculate the loss value\n",
    "        loss = criterion(outputs, labels)\n",
    "        \n",
    "        # Backward pass\n",
    "        loss.backward()\n",
    "        \n",
    "        # We update the parameters in the NN\n",
    "        optimizer.step()\n",
    "\n",
    "        # We store the training loss value\n",
    "        train_loss[epoch] += loss.item()\n",
    "\n",
    "        # In practice, it is good to create a validation dataset\n",
    "        # and calculate the validation loss as well\n",
    "        \n",
    "        model.eval() # Should be used in the general case (turn off batch normalization and dropout layers)\n",
    "        \n",
    "        # We load a new batch of test samples\n",
    "        inputs, labels = next(dataiter)\n",
    "        \n",
    "        # We temporarily set all the requires_grad flag to false with torch.no_grad()\n",
    "        with torch.no_grad():\n",
    "            outputs = model(inputs)\n",
    "        \n",
    "        # We store the validation loss value\n",
    "        val_loss[epoch] += criterion(outputs, labels)\n",
    "            \n",
    "train_loss /= len(train_dataloader)\n",
    "val_loss /= len(train_dataloader)\n",
    "```\n",
    "\n",
    "Lightning can handle the following aspects of this training loop automatically:\n",
    "\n",
    "- No need for the nested `for` loop, `Trainer` iterates the provided dataloaders automatically\n",
    "- Calls to `optimizer` are handled automatically\n",
    "- `model.train()` and `model.eval()` is set automatically\n",
    "- Logging is supported internally (we will come to that later)\n",
    "- We do not need any `.to(device)` calls, the device is infered automatically\n",
    "\n",
    "Let's create our `Trainer`. It takes > 20 keyword arguments that control the training process. \n",
    "\n",
    "- `max_epochs`: Number of epochs that we will train\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "90d4abf0-208a-4e2e-96b5-432f1b8b0269",
   "metadata": {},
   "outputs": [],
   "source": [
    "trainer = L.Trainer(max_epochs=10)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "20160468-5521-4144-9d34-51176be9e696",
   "metadata": {},
   "source": [
    "Train a neural network with Pytorch Lightning\n",
    "---------------------------------------------\n",
    "\n",
    "We repeat the exercise from before with our new setup. Create a `Lightning` model and start training."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d4fb9976-2a6d-4c50-ad15-6fb7fb95e58c",
   "metadata": {},
   "outputs": [],
   "source": [
    "model = Net()\n",
    "module = MyLightningModule(model, learning_rate=0.001)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3e15f264-5866-40c3-b3d2-4c17f127f314",
   "metadata": {},
   "outputs": [],
   "source": [
    "trainer.fit(module, train_dataloader)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "eb7a996b-b118-4cd9-a590-7ffac70b0869",
   "metadata": {},
   "source": [
    "What happened?"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "323fe3f8-6ca8-4b06-b372-26625e512894",
   "metadata": {},
   "source": [
    "Add a validation loop\n",
    "---------------------\n",
    "\n",
    "Think back to the PyTorch exercise, where we added a validation loop. In Lightning, we can enable validation by passing a `valid_dataloader` to the `trainer.fit` call. Re-create the model and trainer and get started."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "10118bea-17e2-4992-b1a6-f485a40d03e8",
   "metadata": {},
   "outputs": [],
   "source": [
    "model = Net()\n",
    "module = MyLightningModule(model, learning_rate=0.001)\n",
    "\n",
    "# validation will also show up in the progress bar\n",
    "trainer = L.Trainer(max_epochs=10, enable_progress_bar=True)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "eb4c243d-426e-411f-b6ff-3c1e0ecee9f4",
   "metadata": {},
   "outputs": [],
   "source": [
    "trainer.fit(module, train_dataloader, validation_dataloader)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "2cf8186b-68fb-475f-b153-be1f698678e9",
   "metadata": {},
   "source": [
    "Add a Logger\n",
    "------------\n",
    "\n",
    "We would like to log certain metrics during training\n",
    "\n",
    "- training loss\n",
    "- validation loss\n",
    "\n",
    "We need to specify this in different places. In `training_step` and `validation_step`, we can use the following statements:\n",
    "\n",
    "```python\n",
    "self.log('train/loss', loss, on_step=True, on_epoch=False, prog_bar=True) # in train\n",
    "\n",
    "self.log('valid/loss', loss, on_step=False, on_epoch=True, prog_bar=True) # in validation\n",
    "```\n",
    "\n",
    "This would automatically output the metrics averaged across epochs to the progress bar. Useful for Jupyter notebooks, but if we want something more permanent, we should use a professional logger again.\n",
    "\n",
    "✏️ Add the log statements to the `LightningModule`"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ea6f2a95-05ea-45c7-a187-71f638d6f3da",
   "metadata": {},
   "outputs": [],
   "source": [
    "import getpass, os, mlflow\n",
    "\n",
    "from lightning.pytorch.loggers import MLFlowLogger\n",
    "# Here: available from workshop environment os.environ[\"MLFLOW_TRACKING_USERNAME\"] = getpass.getuser()\n",
    "# Here: available from workshop environment os.environ[\"MLFLOW_TRACKING_PASSWORD\"] = getpass.getpass(\"MLflow password: \")\n",
    "mlflow.set_workspace(\"bk1444\")\n",
    "logger = MLFlowLogger(tracking_uri=\"https://mlflow.dkrz.de\", experiment_name=\"DeepLearningLightning\", run_name=\"CHANGE-THIS\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3fe20ff4-459c-4ebd-b1e5-9cb1e5568504",
   "metadata": {},
   "outputs": [],
   "source": [
    "model = Net()\n",
    "module = MyLightningModule(model, learning_rate=0.001)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d87f7462-c200-4dcb-9019-a1947be72926",
   "metadata": {},
   "outputs": [],
   "source": [
    "trainer = L.Trainer(max_epochs=10, logger=logger, enable_progress_bar=True)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e440ac3d-7442-4475-ac76-3295f0fd5347",
   "metadata": {},
   "outputs": [],
   "source": [
    "trainer.fit(module, train_dataloader, validation_dataloader)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "0691135d-2295-4e6b-909f-71fbbe35185c",
   "metadata": {},
   "source": [
    "Look at the output of the model training logs together over at WandB ..."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "173631d5-b4b4-4132-952e-741205032556",
   "metadata": {},
   "source": [
    "Evaluate on a test set\n",
    "----------------------\n",
    "\n",
    "Lightning can also be used in inference. Update the model with a `test_step` that reproduces the code from the earlier exercise:\n",
    "\n",
    "```python\n",
    "# We create a function to compute the accuracy given a dataset\n",
    "def get_score(data_loader, data_type):\n",
    "    model.eval()\n",
    "    with torch.no_grad():\n",
    "        num_correct = 0\n",
    "        num_samples = 0\n",
    "        for i, data in enumerate(data_loader, 0):\n",
    "            inputs, labels = data\n",
    "            # We applied our trained model to each batch of samples\n",
    "            scores = model(inputs)\n",
    "            # The shape of scores is (16, 12) = (batch_size, number_of_months)\n",
    "            # From each sample of the batch (first axis), \n",
    "            # get the index of the output with highest value (most probable month)\n",
    "            _, predictions = scores.max(1)\n",
    "            num_correct += (predictions == labels).sum()\n",
    "            num_samples += predictions.size(0)\n",
    "        print(\"Accuracy {}: {:.2f}%\".format(data_type, 100 * num_correct / num_samples))\n",
    "```\n",
    "\n",
    "Use the same `Trainer` instance with the `test` function. It takes as arguments\n",
    "\n",
    "- the trained module\n",
    "- the test dataloader that iterates the test set"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "91891a33-738e-4b5f-98b9-d3b6f650144c",
   "metadata": {},
   "outputs": [],
   "source": [
    "trainer.test(module, test_dataloader)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "c86ffd16-42ec-4ab5-9106-c73a9271b0de",
   "metadata": {},
   "source": [
    "## Advanced: From your laptop to HPC\n",
    "\n",
    "With Lightning, you can run the same code on different platforms. Useful trainer flags (https://lightning.ai/docs/pytorch/stable/core-api/trainer) to know are\n",
    "\n",
    "- `fast_dev_run`: If `True`, run only one step each of training, validation, and test -> quickly spot failure before submitting to the queue\n",
    "- `accelerator`: Automatically chooses GPU (\"cuda\") device if it is available, otherwise cpu\n",
    "- `num_nodes`: Number of nodes, and `devices`: number of devices per node. For example, `num_nodes=2, devices=4` on Levante would use 2 nodes with 4 GPUs each"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b6d86c2e-1f2b-4e16-a34c-f84acc0aee06",
   "metadata": {},
   "outputs": [],
   "source": [
    "model = Net()\n",
    "module = MyLightningModule(model, learning_rate=0.001)\n",
    "trainer = L.Trainer(fast_dev_run=True, accelerator='auto', logger=logger)\n",
    "trainer.fit(module, train_dataloader, validation_dataloader)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "7c21cea4-ff11-46dc-a066-b178f841d021",
   "metadata": {},
   "source": [
    "Advanced: Callbacks\n",
    "-------------\n",
    "\n",
    "Callbacks are a great way to control training. They come from the standard Pytorch library and can be added to the `trainer` when it is created. Callbacks are executed at defined times in the training process automatically. We will use two callbacks here:\n",
    "\n",
    "- `EarlyStopping`: This is to avoid overfitting. It monitors a validation metric, and once that metric no longer decreases, we think we have reached an optimal set of model parameters. The training procedure stops automatically.\n",
    "- `ModelCheckpoint`: This saves the best model, again determined by monitoring a validation metric."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "aec6ec34-92d1-4d5f-9f23-73e931d61059",
   "metadata": {},
   "outputs": [],
   "source": [
    "from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7c500b60-f4c7-4f7e-8159-52ee7b5eedb9",
   "metadata": {},
   "outputs": [],
   "source": [
    "callbacks = [EarlyStopping(monitor='valid/loss', mode='min', patience=3), ModelCheckpoint(monitor='valid/loss')]\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "51b09bbb-8d1c-4c9e-af2f-901d0b82214c",
   "metadata": {},
   "outputs": [],
   "source": [
    "trainer = L.Trainer(max_epochs=50, enable_progress_bar=True, devices=1, accelerator='auto', logger=logger, callbacks=callbacks)\n"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "ML Class",
   "language": "python",
   "name": "mlclass"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.10.13"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
