{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "377adf2b-d0df-4848-82a6-783caf5b9176",
   "metadata": {},
   "source": [
    "PyTorch Lightning summary\n",
    "=========================\n",
    "\n",
    "This notebook summarizes the `lightning_demo` notebook for later reference."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2124d961-9725-40ba-919a-d01003508c57",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "from torch.utils.data import DataLoader\n",
    "\n",
    "import xarray as xr\n",
    "\n",
    "import lightning as L\n",
    "from lightning.pytorch.loggers import MLFlowLogger\n",
    "import mlflow\n",
    "\n",
    "from lightning.pytorch.callbacks import EarlyStopping, ModelCheckpoint\n",
    "\n",
    "import getpass, os"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "55bf67ae-3541-42fa-8109-061b074eedfc",
   "metadata": {},
   "source": [
    "Data\n",
    "----"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4e7e2c5a-a97e-425a-b9ab-ecccf2d83b25",
   "metadata": {},
   "outputs": [],
   "source": [
    "class MyDataModule(L.LightningDataModule):\n",
    "    def __init__(self, path, batch_size):\n",
    "        super().__init__()\n",
    "\n",
    "        self.path = path\n",
    "        self.batch_size = batch_size        \n",
    "\n",
    "    def prepare_data(self):\n",
    "        ds = xr.open_dataset(self.path)\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",
    "        self.dataset = dataset\n",
    "        self.n_valtest = int( n_time * 0.2 )\n",
    "        self.n_training = n_time - 2 * self.n_valtest\n",
    "    \n",
    "    def setup(self):\n",
    "        training_data, validation_data, test_data = torch.utils.data.random_split(self.dataset, [self.n_training, self.n_valtest, self.n_valtest])\n",
    "\n",
    "        self.training_data = training_data\n",
    "        self.validation_data = validation_data\n",
    "        self.test_data = test_data\n",
    "\n",
    "    def train_dataloader(self):\n",
    "        return DataLoader(self.training_data, batch_size=self.batch_size, shuffle=True)\n",
    "\n",
    "    def val_dataloader(self):\n",
    "        return DataLoader(self.validation_data, batch_size=self.batch_size, shuffle=False)\n",
    "\n",
    "    def test_dataloader(self):\n",
    "        return DataLoader(self.test_data, batch_size=self.n_valtest, shuffle=False)\n",
    "\n",
    "    def predict_datalaoder(self):\n",
    "        return DataLoader(self.test_data, batch_size=self.n_valtest, shuffle=False)\n",
    "\n",
    "\n",
    "        "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7214da7e-b6a4-487b-9f9e-705d5388439e",
   "metadata": {},
   "outputs": [],
   "source": [
    "path_20cr = %env path_20cr \n",
    "print(path_20cr)\n",
    "# /work/bk1318/mlclass/data/20crv3-part.nc\n",
    "datamodule = MyDataModule(path=path_20cr, batch_size=16)\n",
    "datamodule.prepare_data()\n",
    "datamodule.setup()\n",
    "\n",
    "train_dataloader = datamodule.train_dataloader()\n",
    "validation_dataloader = datamodule.val_dataloader()\n",
    "test_dataloader = datamodule.test_dataloader()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "730f15c8-aab4-4cc6-b3f1-3124f3ec7a33",
   "metadata": {},
   "source": [
    "Model\n",
    "-----"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "370b10fc-55a5-4224-b60a-8054706fef29",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch.nn as nn\n",
    "from torch.nn.functional import relu\n",
    "\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": "code",
   "execution_count": null,
   "id": "6e9c0269-22a5-4674-85a1-3554dfe4e3a8",
   "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.criterion = nn.CrossEntropyLoss() # define the loss criterion\n",
    "\n",
    "    def training_step(self, batch, batch_idx):\n",
    "        inputs, labels = batch\n",
    "        outputs = self.model(inputs)\n",
    "        loss = self.criterion(outputs, labels)\n",
    "        self.log('train/loss', loss, on_step=True, on_epoch=False)\n",
    "        return loss\n",
    "\n",
    "    def validation_step(self, batch, batch_idx):\n",
    "        inputs, labels = batch\n",
    "        outputs = self.model(inputs)\n",
    "        loss = self.criterion(outputs, labels)\n",
    "        self.log('valid/loss', loss, on_step=False, on_epoch=True)\n",
    "\n",
    "    def test_step(self, batch, batch_idx):\n",
    "        inputs, labels = batch\n",
    "        outputs = self.model(inputs)\n",
    "\n",
    "        _, predictions = outputs.max(1)\n",
    "        num_correct = (predictions == labels).sum()\n",
    "        num_samples = predictions.size(0)\n",
    "        accuracy = num_correct / num_samples\n",
    "        \n",
    "        self.log('test/accuracy', accuracy)\n",
    "\n",
    "    def predict_step(self, batch, batch_idx):\n",
    "        inputs, labels = batch\n",
    "        outputs = self.model(inputs)\n",
    "\n",
    "        _, predictions = outputs.max(1)\n",
    "        return predictions\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": "115de6c0-50ea-402b-9e19-0c8f7830f0c7",
   "metadata": {},
   "source": [
    "Trainer\n",
    "-------"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "486c6008-3518-4a2f-80c9-0147e3be60e8",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Available from workshop environment: os.environ[\"MLFLOW_TRACKING_USERNAME\"] = getpass.getuser()\n",
    "# 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=\"test-run-lightning\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1210d093-2ab6-413a-8d06-c3f18ebef9bd",
   "metadata": {},
   "outputs": [],
   "source": [
    "callbacks = [EarlyStopping(monitor='valid/loss', mode='min', patience=3), ModelCheckpoint(monitor='valid/loss')]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "574f6420-01de-4122-a5d6-0a01235b1b02",
   "metadata": {},
   "outputs": [],
   "source": [
    "trainer = L.Trainer(max_epochs=50, enable_progress_bar=False, devices=1, accelerator='auto', logger=logger, callbacks=callbacks)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "23ca02de-f7c0-44fe-b9c3-19a0e87d0f07",
   "metadata": {},
   "source": [
    "Demo Training\n",
    "-------------"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "112eb5c2-81da-48df-9734-15d0a5ed9fd1",
   "metadata": {},
   "outputs": [],
   "source": [
    "model = Net()\n",
    "module = MyLightningModule(model, learning_rate=0.01)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c8c3885c-94c3-43d0-bb9f-fcb0ac2dedcd",
   "metadata": {},
   "outputs": [],
   "source": [
    "trainer.fit(module, train_dataloader, validation_dataloader)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "8877ddd6-94c3-4959-bb6f-1762bd3176ac",
   "metadata": {},
   "source": [
    "Demo Evaluation\n",
    "---------------"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "27a885ed-f1e8-4e2b-bd01-9ef2c6b81788",
   "metadata": {},
   "outputs": [],
   "source": [
    "trainer.test(module, test_dataloader);"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "28d0b5fd-2496-4fda-a4a5-bf1c568c8ff0",
   "metadata": {},
   "source": [
    "Demo Prediction\n",
    "---------------"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5119d502-d758-4435-8147-95896bf717d4",
   "metadata": {},
   "outputs": [],
   "source": [
    "predictions = trainer.predict(module, test_dataloader)[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2fe0a6fb-8403-45a7-bb22-33dbe44cb1db",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9e5b8952-427e-48d9-b6b7-603ba31a24e0",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "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
}
