{
  "nbformat": 4,
  "nbformat_minor": 0,
  "metadata": {
    "colab": {
      "provenance": []
    },
    "kernelspec": {
      "name": "python3",
      "display_name": "Python 3"
    },
    "language_info": {
      "name": "python"
    }
  },
  "cells": [
    {
      "cell_type": "markdown",
      "source": [
        "# Réseaux de neurones et classification d'images avec MLP\n",
        "\n",
        "Le but de ce TP est de e familiariser avec la classification d'images sur un corpus d'images simples, avant d'approfondir la semaine prochaines sur des images plus réalistes.\n",
        "\n",
        "Le but du TP est de remplir les zones de code vides pour implémenter un programme qui peut classer correctement les images (avec un taux d'erreur inférieur à 10% )"
      ],
      "metadata": {
        "id": "fmcLY7w4uoqO"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "import torch\n",
        "from torch.utils.data import Dataset, DataLoader\n",
        "\n",
        "\n",
        "from torchvision import datasets\n",
        "from torchvision.transforms import ToTensor, Compose, Resize\n",
        "import matplotlib.pyplot as plt\n"
      ],
      "metadata": {
        "id": "Ghnz6EbLvdb6"
      },
      "execution_count": 2,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "source": [
        "## Utilisation GPU\n",
        "le code suivant vérifie que le GPU est disponible, et l'utilise si c'est le cas.\n",
        "Si la valeur de `use_cuda` est `False`, vous devez changer le type de machine pour avoir accès à un GPU."
      ],
      "metadata": {
        "id": "GGio5tdXvgs0"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "use_cuda = torch.cuda.is_available()\n",
        "print(use_cuda)\n",
        "device = torch.device(\"cuda:0\" if use_cuda else \"cpu\")"
      ],
      "metadata": {
        "colab": {
          "base_uri": "https://localhost:8080/"
        },
        "id": "TX666b8yvhHm",
        "outputId": "63f23f07-dde4-4fac-f192-7a1115749e0f"
      },
      "execution_count": 3,
      "outputs": [
        {
          "output_type": "stream",
          "name": "stdout",
          "text": [
            "False\n"
          ]
        }
      ]
    },
    {
      "cell_type": "markdown",
      "source": [
        "## Préparation des données\n",
        "\n",
        "Torchvision contient des jeux de données pour entraîner et évaluer les réseaux.\n",
        "Nous utilisons [MNIST](https://en.wikipedia.org/wiki/MNIST_database).\n",
        "\n",
        "On récupère les données d'entraînement et de test, puis on réserve une partie des données d'entraînement, qu'on appelle données de validation, pour sélectionner le meilleur modèle.\n",
        "Ensuite, on nomme les différentes classes d'images."
      ],
      "metadata": {
        "id": "AUngjgx_wL9S"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "training_data = datasets.MNIST(\n",
        "    root=\"data\",\n",
        "    train=True,\n",
        "    download=True,\n",
        "    transform=ToTensor()\n",
        ")\n",
        "\n",
        "test_data = datasets.MNIST(\n",
        "    root=\"data\",\n",
        "    train=False,\n",
        "    download=True,\n",
        "    transform=ToTensor()\n",
        ")\n",
        "\n",
        "training_data, val_data = torch.utils.data.random_split(training_data, [50000, 10000])\n",
        "\n",
        "NB_CLASSES=10\n",
        "\n",
        "labels_map = {\n",
        "    0: \"Zero\",\n",
        "    1: \"One\",\n",
        "    2: \"Two\",\n",
        "    3: \"Three\",\n",
        "    4: \"Four\",\n",
        "    5: \"Five\",\n",
        "    6: \"Six\",\n",
        "    7: \"Seven\",\n",
        "    8: \"Eight\",\n",
        "    9: \"Nine\",\n",
        "}\n"
      ],
      "metadata": {
        "id": "cB-xilm-wLBI"
      },
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "source": [
        "Ce corpus est constitué d'images en noir et blanc de taille 28×28. Ces images représentent des chiffres réparties en 10 classes différentes.\n",
        "\n",
        "Chacune des cases (28$\\times$ 28=784) contient une valeur entre 0 et 1, qui indique un niveau de gris, avec 0 pour les points noirs et 1 pour les points blancs.\n",
        "\n",
        "Notez bien que même si ces images peuvent être représentées par une matrice 28×28, elles sont en faites stockées comme des tenseurs 1×28×28 de façon à expliciter qu'au lieu des habituels 3 canaux RGB indiquant les couleurs, qu'on verra en détails la semaine prochaine, on n'a ici qu'un seul canal qui indique le niveau de gris.\n",
        "\n",
        "La fonction suivante affiche une sélection aléatoire de chiffres:\n",
        "\n"
      ],
      "metadata": {
        "id": "YUof2QEVxIpw"
      }
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "id": "G25YG_mAuiRr"
      },
      "outputs": [],
      "source": [
        "def display_examples(cols=5, rows=2, net=None, data=training_data):\n",
        "  figure = plt.figure(figsize=(16, 8))\n",
        "  for i in range(1, cols * rows + 1):\n",
        "    sample_idx = torch.randint(len(data), size=(1,)).item()\n",
        "    img, label = data[sample_idx]\n",
        "    figure.add_subplot(rows, cols, i)\n",
        "\n",
        "    if net is None:\n",
        "      label = \"Correct Label: \" + labels_map[label]\n",
        "    else:\n",
        "      imgdev = img.to(device)\n",
        "      scores = net(imgdev)\n",
        "      pred = torch.argmax(scores)\n",
        "      label = \"Predicted Label: \" + labels_map[pred.item()]\n",
        "\n",
        "    plt.title(label)\n",
        "    plt.axis(\"off\")\n",
        "    plt.imshow(img.squeeze(), cmap=\"gray\")\n",
        "  plt.show()\n",
        "\n",
        "display_examples()"
      ]
    },
    {
      "cell_type": "markdown",
      "source": [
        "## On charge les données\n",
        "\n",
        "pour pouvoir être chargées facilement on crée des `DataLoader` qui vont nous permettre d'itérer facilement sur les données\n",
        "\n"
      ],
      "metadata": {
        "id": "hM28_r4ayyHH"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "BATCH_SIZE=512 #nombre d'images à traiter en parallèle\n",
        "\n",
        "train_dataloader = DataLoader(training_data, batch_size=BATCH_SIZE, shuffle=True, pin_memory=True)\n",
        "val_dataloader = DataLoader(val_data, batch_size=BATCH_SIZE, shuffle=False, pin_memory=True)\n",
        "test_dataloader = DataLoader(test_data, batch_size=BATCH_SIZE, shuffle=False, pin_memory=True)"
      ],
      "metadata": {
        "id": "t9NZHrq2yZbR"
      },
      "execution_count": 7,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "source": [
        "On dispose maintenant d'itérateurs sur les données `next`. On peut écrire par exemple :\n",
        "\n"
      ],
      "metadata": {
        "id": "2BVSFmOlzFN5"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "# Display image and label.\n",
        "train_iter = iter(train_dataloader) #get an iterator over data\n",
        "train_features, train_labels = next(train_iter) # get a batch of data\n",
        "print(f\"Feature batch shape: {train_features.size()}\")\n",
        "print(f\"Labels batch shape: {train_labels.size()}\")\n",
        "img = train_features[0].squeeze() # reformat the first pic as 28x28\n",
        "label = train_labels[0]\n",
        "plt.imshow(img, cmap=\"gray\") #it's B&W\n",
        "plt.show()\n",
        "print(f\"Label: {label}\")"
      ],
      "metadata": {
        "id": "j6pwuyZGzL0V"
      },
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "source": [
        "Que contient notre image ?"
      ],
      "metadata": {
        "id": "Hrj7534kzZaN"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "print(img)"
      ],
      "metadata": {
        "id": "JBJXPAFpzffk"
      },
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "source": [
        "## Évaluation\n",
        "La fonction suivante compare, pour un jeu de données passé en paramètre, la classe `argmax` calculée par un réseau et le label qui devait être prédit.\n",
        "Si ces valeurs sont égales alors elles comptent pour de bonnes réponses.\n",
        "On retourne le ratio de bonnes réponses."
      ],
      "metadata": {
        "id": "31hlPXqdzoKB"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "def eval(net, dataloader, device):\n",
        "\n",
        "  total_test = 0.0\n",
        "  correct_test = 0.0\n",
        "\n",
        "  net.eval()\n",
        "\n",
        "  #iterate over data\n",
        "  for local_batch, local_labels in dataloader:\n",
        "    # Transfer to GPU\n",
        "    local_batch, local_labels = local_batch.to(device), local_labels.to(device)\n",
        "\n",
        "    preds = net(local_batch)\n",
        "    total_test += preds.size(0)\n",
        "    correct_test += (torch.argmax(preds, dim=1) == local_labels).sum()\n",
        "\n",
        "  return correct_test/total_test"
      ],
      "metadata": {
        "id": "WG1tcWSdznzs"
      },
      "execution_count": 11,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "source": [
        "# Implémentation d'un Perceptron multicouche\n",
        "\n",
        "Implementer la class MLP pour les perceptrons multicouches.\n",
        "\n",
        "Le constructeur prend en entrée:\n",
        "\n",
        "    input_size la taille du vecteur d'entrée\n",
        "    hidden_sizes la liste des tailles des couches internes\n",
        "    output_size la taille de la couche de sortie\n",
        "    activation la classe qui implémente la fonction à appliquer après chaque transformation linéaire\n",
        "\n",
        "Aide: vous pouvez utiliser la classe Sequential pour modéliser la liste des transformations de votre MLP\n",
        "\n",
        "La fonction forward prend entrée une image. Il faut d'abord transformer l'image (le tenseur à 3 dimensions) en un vecteur, puis on applique successivement chaque transformation linéaire, suivie d'une activation si on n'est pas sur la dernière couche."
      ],
      "metadata": {
        "id": "5k9wZOAw0WnP"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "class MLP(torch.nn.Module):\n",
        "\n",
        "  def __init__(self, input_size, hidden_sizes, output_size, activation):\n",
        "    super().__init__()\n",
        "\n",
        "    #your code begins here\n",
        "    pass\n",
        "    #your code ends here\n",
        "\n",
        "  def forward(self,x):\n",
        "    #your code begins here\n",
        "    pass\n",
        "    #your code ends here\n",
        ""
      ],
      "metadata": {
        "id": "HBnWXDW-0e_O"
      },
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "source": [
        "## Boucle d'entraînement\n",
        "\n",
        "Écrire la fonction `train` qui implément la boucle d'apprentissage"
      ],
      "metadata": {
        "id": "9qKGNQm10ljU"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "def train(net, optimizer, epochs, train_data, val_data, test_data, batch_size, device, eval_frq=10):\n",
        "\n",
        "  loss_fn = torch.nn.CrossEntropyLoss()\n",
        "\n",
        "  total_loss = 0.0\n",
        "\n",
        "  n_batches = 0\n",
        "  for epoch in range(epochs):\n",
        "\n",
        "    net.train()\n",
        "    # iterate through train data\n",
        "    #   -> get new batch\n",
        "    #   -> send it to device\n",
        "    #   -> compute class scores with net\n",
        "    #   -> compute loss and compute the gradient of the loss\n",
        "    #   -> call step() on optimizer to perform SGD, reset gradients to zero\n",
        "\n",
        "    #your code begins here\n",
        "    pass\n",
        "    #your code ends here\n",
        "\n",
        "    # eval if it's time...\n",
        "    if epoch % eval_frq == 0:\n",
        "      print(f\"Epoch: {epoch}, Mean loss: {total_loss/n_batches}\")\n",
        "      rv=eval(net, val_loader, device)\n",
        "      #  --> Display eval scores rt and rv\n",
        "      display_examples(cols=5, rows=2, net=net, data=val_data):\n",
        "      rt=eval(net, test_loader, device)\n",
        "      display_examples(cols=5, rows=2, net=net, data=test_data):\n",
        "\n"
      ],
      "metadata": {
        "id": "P7Sw2-IE0WRG"
      },
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "source": [
        "## Création du réseau\n",
        "\n",
        "Donner les bon paramètres pour la création de `net`.\n",
        "\n",
        "\n",
        "Enuiste, modifier la fonction `display_examples` pour que si le paramètre `net` n'est pas `None` alors au lieu d'afficher le label correct des images, on affiche le label prédit par `net`, c'est-à-dire la classes `argmax`."
      ],
      "metadata": {
        "id": "3LSbUqLD2awa"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "net = MLP(None, None, None, None)\n",
        "net = net.to(device)\n",
        "\n",
        "display_examples(net=net)"
      ],
      "metadata": {
        "id": "KjEsUmYe2Q9W"
      },
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "source": [
        "## TEST"
      ],
      "metadata": {
        "id": "s1-Xvh_k2prc"
      }
    },
    {
      "cell_type": "code",
      "source": [
        "\n",
        "opt = torch.optim.SGD(net.parameters(),lr=0.1)\n",
        "\n",
        "train(net, opt, epochs=50, train_data=training_data, val_data=val_data, test_data=test_data, batch_size=BATCH_SIZE, model_name=\"model\", device=device, eval_frq=5)\n",
        "\n"
      ],
      "metadata": {
        "id": "a4xaAm052src"
      },
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "source": [
        "## Autres choses\n",
        "\n",
        "1. Changer la fonction d'activation (`tanh`, `sigmoid`, `ReLU` etc) et comparer la convergence.\n",
        "1. Changer le `learning rate` ($\\alpha$) de SGD, et comparer la vitesse de convergence\n",
        "2. Au lieu de SGD, utiliser Adam et comparer les vitesses de convergence\n",
        "\n"
      ],
      "metadata": {
        "id": "fSYIoXXu3P3O"
      }
    },
    {
      "cell_type": "code",
      "source": [],
      "metadata": {
        "id": "5YqYEV9N3ZAb"
      },
      "execution_count": null,
      "outputs": []
    }
  ]
}