{"cells":[{"metadata":{},"cell_type":"markdown","source":"**From version 11** : I got a lot of unsuccessfull atempt in private repository. I will retry here with an efficientnet0 from a public kernel until I managed to got a good working implementation. NB : according to this [source](https://arxiv.org/pdf/1906.02629.pdf), a network trained with label smoothing is not really a got teacher, so I will import a network trained with cross-entropy as the loss."},{"metadata":{},"cell_type":"markdown","source":"**Version 14** : testing an alpha parameter of 0.9 as suggested by [this article](https://openaccess.thecvf.com/content_ICCV_2019/papers/Cho_On_the_Efficacy_of_Knowledge_Distillation_ICCV_2019_paper.pdf)."},{"metadata":{},"cell_type":"markdown","source":"**Version 15** : Update graphic to show distillation loss, student loss and total loss.\n\n**Version 16** : Reboot of the train generator.\n\n**Version 17** : Testing unfreezing.\n\n**Version 19** : Finetuning would be probably better if I reduce the learning rate."},{"metadata":{},"cell_type":"markdown","source":"# What is knowledge distillation"},{"metadata":{},"cell_type":"markdown","source":"As presented [in this discussion thread](https://www.kaggle.com/c/cassava-leaf-disease-classification/discussion/214959), knowledge distillation is defined as *simply trains another individual model to match the output of the ensemble.* [Source](https://www.microsoft.com/en-us/research/blog/three-mysteries-in-deep-learning-ensemble-knowledge-distillation-and-self-distillation/). It is in fact slightly more complicated : the second neural net (student) will made predictions on the images, but then, the losses will be a function of the its loss as well as a loss beased on the difference between his prediction and the one of its teacher or the ensemble. This approach allow to compress an ensemble into one model and by then reduce the inference time, or, if trained to match the output of one mode, increase the overall performance of the model. I discover this approach by looking at the top solution of the Plant Pathology 2020, such as [this one](https://www.kaggle.com/c/plant-pathology-2020-fgvc7/discussion/154056)."},{"metadata":{},"cell_type":"markdown","source":"I let you go to [to this source mention aboved to understand how it could potentially works](https://www.microsoft.com/en-us/research/blog/three-mysteries-in-deep-learning-ensemble-knowledge-distillation-and-self-distillation/). It does not seems sure, but it seems related to the learning of specific feature vs forcing the student to learn \"multiple view\", multiple type of feature to detect in the images. "},{"metadata":{},"cell_type":"markdown","source":"There is off course, no starting material to do it in R. Thanksfully there is a code example on the [website of keras](https://keras.io/examples/vision/knowledge_distillation/). In this example, they create a class of model, a distiller, to make the knowledge distillation. There is, however, one problem : **model are not inheritable in R**. There is example of inheritance with a R6 for callback, [like here](https://keras.rstudio.com/articles/training_callbacks.html), but the models are not a R6 class. To overcome this problems, I used the code example as a guide, and reproduced the steps by following the approach in this [guide for eager executation in keras with R](https://keras.rstudio.com/articles/eager_guide.html). I took other code from [the tensorflow website for R](https://tensorflow.rstudio.com/tutorials/advanced/).\n"},{"metadata":{},"cell_type":"markdown","source":"**The code is quite hard to understand at first glance**. The reason is that, everything is executed in a **single for loop**. Everything is done in eager mode. I did not seemed possible to do it differently. So there is a lot of variable around to collect metrics during training. If you want to understand the code just removed it or run it outside of the for loop, before reconstructing the loop around. I did not used tfdataset as shown on the guide for eager execution, so instead of make_iterator_one_shot() and iterator_get_next(), here we loop over the train_generator to produce the batches."},{"metadata":{"_uuid":"051d70d956493feee0c6d64651c6a088724dca2a","_execution_state":"idle","trusted":true},"cell_type":"code","source":"library(tidyverse)\nlibrary(tensorflow)\ntf$executing_eagerly()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"tensorflow::tf_version()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Here I flex with my own version of keras. Basically, it is a fork with application wrapper for the efficient net."},{"metadata":{},"cell_type":"markdown","source":"**Disclaimer : I did not writte the code for the really handy applications wrappers.** It came [from this commit](https://github.com/rstudio/keras/commit/c406ec55f7bb2864ac58a17f963448810a531c18) for which the PR is hold until the fully release of tf 2.3, as stated [in this PR](https://github.com/rstudio/keras/pull/1097). I am not sure why the PR is closed."},{"metadata":{"trusted":true},"cell_type":"code","source":"devtools::install_github(\"Cdk29/keras\", dependencies = FALSE)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"library(keras)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"labels<-read_csv('/kaggle/input/cassava-leaf-disease-classification//train.csv')\nhead(labels)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"levels(as.factor(labels$label))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"idx0<-which(labels$label==0)\nidx1<-which(labels$label==1)\nidx2<-which(labels$label==2)\nidx3<-which(labels$label==3)\nidx4<-which(labels$label==4)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"labels$CBB<-0\nlabels$CBSD<-0\nlabels$CGM<-0\nlabels$CMD<-0\nlabels$Healthy<-0","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"labels$CBB[idx0]<-1\nlabels$CBSD[idx1]<-1\nlabels$CGM[idx2]<-1\nlabels$CMD[idx3]<-1","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"\"Would it have been easier to create a function to convert the labelling ?\" You may ask."},{"metadata":{"trusted":true},"cell_type":"code","source":"labels$Healthy[idx4]<-1","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Probably."},{"metadata":{"trusted":true},"cell_type":"code","source":"#labels$label<-NULL","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"head(labels)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"val_labels<-read_csv('../input/efficientnetb0-with-r-and-tf2-cyclic-lr/validation_set.csv')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_labels<-labels[which(!labels$image_id %in% val_labels$image_id),]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"table(train_labels$image_id %in% val_labels$image_id)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_labels$label<-NULL\nval_labels$label<-NULL\n\nhead(train_labels)\nhead(val_labels)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"image_path<-'/kaggle/input/cassava-leaf-disease-classification//train_images'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#data augmentation\ndatagen <- image_data_generator(\n  rotation_range = 40,\n  width_shift_range = 0.2,\n  height_shift_range = 0.2,\n  shear_range = 0.2,\n  zoom_range = 0.5,\n  horizontal_flip = TRUE,\n  fill_mode = \"reflect\"\n)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"img_path<-\"/kaggle/input/cassava-leaf-disease-classification/train_images/1000015157.jpg\"\n\nimg <- image_load(img_path, target_size = c(448, 448))\nimg_array <- image_to_array(img)\nimg_array <- array_reshape(img_array, c(1, 448, 448, 3))\nimg_array<-img_array/255\n# Generated that will flow augmented images\naugmentation_generator <- flow_images_from_data(\n  img_array, \n  generator = datagen, \n  batch_size = 1 \n)\nop <- par(mfrow = c(2, 2), pty = \"s\", mar = c(1, 0, 1, 0))\nfor (i in 1:4) {\n  batch <- generator_next(augmentation_generator)\n  plot(as.raster(batch[1,,,]))\n}\npar(op)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Data generator"},{"metadata":{},"cell_type":"markdown","source":"Okay so here is an interresting thing, I will try to compress the code to call a train generator to make it easier to call it. \n\nWhy ? **Looking at version 7, 6 and 2**, you can see the for loop of the distiller stop working without obvious reason at epoch 7. The reason is that the validation generator can have an end !\n\nWhen we iterate over it, validation_generator yeld 8 images and 8 label, until the batch 267, than contains only 5 images (and create the bug when we try to add the loss of the batch to the loss of the epoch. Batch 268 does not exist. So solution seems to recreate on the fly the validation set and restart the iterations."},{"metadata":{"trusted":true},"cell_type":"code","source":"arg.list <- list(dataframe = val_labels, directory = image_path,\n                                              class_mode = \"other\",\n                                              x_col = \"image_id\",\n                                              y_col = c(\"CBB\",\"CBSD\", \"CGM\", \"CMD\", \"Healthy\"),\n                                              target_size = c(600, 600),\n                                              batch_size=8)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"validation_generator <- do.call(flow_images_from_dataframe, arg.list)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dim(validation_generator[266][[1]])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dim(validation_generator[267][[1]])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dim(val_labels)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"2141/8","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_generator <- flow_images_from_dataframe(dataframe = train_labels, \n                                              directory = image_path,\n                                              generator = datagen,\n                                              class_mode = \"other\",\n                                              x_col = \"image_id\",\n                                              y_col = c(\"CBB\",\"CBSD\", \"CGM\", \"CMD\", \"Healthy\"),\n                                              target_size = c(600, 600),\n                                              batch_size=8)\n\nvalidation_generator <- flow_images_from_dataframe(dataframe = val_labels, \n                                              directory = image_path,\n                                              class_mode = \"other\",\n                                              x_col = \"image_id\",\n                                              y_col = c(\"CBB\",\"CBSD\", \"CGM\", \"CMD\", \"Healthy\"),\n                                              target_size = c(600, 600),\n                                              batch_size=8)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_generator","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"conv_base<-keras::application_efficientnet_b0(weights = \"imagenet\", include_top = FALSE, input_shape = c(600, 600, 3))\n\nfreeze_weights(conv_base)\n\nmodel <- keras_model_sequential() %>%\n    conv_base %>% \n    layer_global_max_pooling_2d() %>% \n    layer_batch_normalization() %>% \n    layer_dropout(rate=0.5) %>%\n    layer_dense(units=5, activation=\"softmax\")\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#unfreeze_weights(model, from = 'block5a_expand_conv')\nunfreeze_weights(conv_base, from = 'block5a_expand_conv')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model %>% load_model_weights_hdf5(\"../input/efficientnetb0-with-r-and-tf2-cyclic-lr/checkpoints_fine_tuned/fine_tuned_eff_net_weights.15.hdf5\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"summary(model)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"conv_base_student<-keras::application_efficientnet_b0(weights = \"imagenet\", include_top = FALSE, input_shape = c(600, 600, 3))\n\nfreeze_weights(conv_base_student)\n\nstudent <- keras_model_sequential() %>%\n    conv_base_student %>% \n    layer_global_max_pooling_2d() %>% \n    layer_batch_normalization() %>% \n    layer_dropout(rate=0.5) %>%\n    layer_dense(units=5, activation=\"softmax\")\n\nstudent","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Source code and knowledge distillation"},{"metadata":{"trusted":true},"cell_type":"markdown","source":"Source code for knowledge distillation with Keras : https://keras.io/examples/vision/knowledge_distillation/  \nHelp for eager executation details in R and various usefull code : https://keras.rstudio.com/articles/eager_guide.html  \nOther source code in R : https://tensorflow.rstudio.com/tutorials/advanced/"},{"metadata":{"trusted":true},"cell_type":"code","source":"i=1\nalpha=0.9 #On_the_Efficacy_of_Knowledge_Distillation_ICCV_2019\ntemperature=3","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"optimizer <- optimizer_adam()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_loss <- tf$keras$metrics$Mean(name='student_loss')\ntrain_accuracy <-  tf$keras$metrics$CategoricalAccuracy(name='train_accuracy')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"nb_epoch<-12","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"nb_batch<-300\nval_step<-40","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_loss_plot<-c()\naccuracy_plot<-c()\ndistilation_loss_plot <- c()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"val_loss_plot <- c()\nval_accuracy_plot <- c()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"count_epoch<-0","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"collapsed":true},"cell_type":"code","source":"for (epoch in 1:nb_epoch) {\n    cat(\"Epoch: \", epoch, \" -----------\\n\")\n    # Init metrics\n    train_loss_epoch <- 0\n    accuracies_on_epoch <- c()\n    distilation_loss_epoch <- 0\n    val_loss_epoch <- 0\n    val_accuaries_on_epoch <- c()\n    \n    #Formula to not see the same batch over and over on each epoch\n    #Count epoch instead of epoch\n    count_epoch<-count_epoch+1\n    idx_batch <- (1+nb_batch*(count_epoch-1)):(nb_batch*count_epoch)\n    idx_val_set <- (1+val_step*(count_epoch-1)):(val_step*count_epoch)\n    \n    #Dirty solution to restart on a new validation batch generator before reaching the end of the other one \n    if (as.integer((dim(val_labels)[1]/8)-1) %in% idx_val_set) {\n        count_epoch<-1\n        idx_val_set <- (1+val_step*(count_epoch-1)):(val_step*count_epoch)\n        validation_generator <- do.call(flow_images_from_dataframe, arg.list)\n    }\n    #need the same if for train generator\n    if (as.integer((dim(train_labels)[1]/8)-1) %in% idx_batch) {\n        count_epoch<-1\n        idx_batch <- (1+nb_batch*(count_epoch-1)):(nb_batch*count_epoch)\n        train_generator <- do.call(flow_images_from_dataframe, arg.list)\n    }\n    \n    for (batch in idx_batch) {\n        x = train_generator[batch][[1]]\n        y = train_generator[batch][[2]]\n        # Forward pass of teacher\n        teacher_predictions = model(x)\n\n        with(tf$GradientTape() %as% tape, {\n            student_predictions = student(x)\n            student_loss = tf$losses$categorical_crossentropy(y, student_predictions)\n        \n            distillation_loss = tf$losses$categorical_crossentropy(tf$nn$softmax(teacher_predictions/temperature, axis=0L), \n                                                           tf$nn$softmax(student_predictions/temperature, axis=0L))\n        \n            loss = alpha * student_loss + (1 - alpha) * distillation_loss\n            })\n        \n        # Compute gradients\n        # Variating learning rate :\n        # optimizer <- optimizer_adam(lr = 0.0001)\n        gradients <- tape$gradient(loss, student$trainable_variables)\n        optimizer$apply_gradients(purrr::transpose(list(gradients, student$trainable_variables)))\n        \n        #Collect the metrics of the student\n        train_loss_epoch <- train_loss_epoch + student_loss\n        distilation_loss_epoch <- distilation_loss_epoch + distillation_loss\n        \n        accuracy_on_batch <- train_accuracy(y_true=y, y_pred=student_predictions)\n        accuracies_on_epoch <- c(accuracies_on_epoch, as.numeric(accuracy_on_batch))\n        \n    }\n\n    #Collect info on current epoch and for graphs and cat()\n    train_loss_epoch <- mean(as.vector(as.numeric(train_loss_epoch))/nb_batch)\n    train_loss_plot <- c(train_loss_plot, train_loss_epoch)\n    \n    distilation_loss_epoch <- mean(as.vector(as.numeric(distilation_loss_epoch))/nb_batch)\n    distilation_loss_plot <- c(distilation_loss_plot, distilation_loss_epoch)\n    \n    accuracies_on_epoch <- mean(accuracies_on_epoch)\n    accuracy_plot <- c(accuracy_plot, accuracies_on_epoch)\n    \n    \n    for (step in idx_val_set) {\n        # Unpack the data\n        x = validation_generator[step][[1]]\n        y = validation_generator[step][[2]]\n\n        # Compute predictions\n        student_predictions = student(x)\n\n        # Calculate the loss\n        student_loss = tf$losses$categorical_crossentropy(y, student_predictions)\n\n        #Collect the metrics of the student\n        #This line will create a bug of shape when val_loss end.\n        val_loss_epoch <- val_loss_epoch + student_loss\n        \n        accuracy_on_val_step <- train_accuracy(y_true=y, y_pred=student_predictions)\n        val_accuaries_on_epoch <- c(val_accuaries_on_epoch, as.numeric(accuracy_on_val_step))\n    }\n    \n    #Collect info on current epoch and for graphs and cat()\n    val_loss_epoch <- mean(as.vector(as.numeric(val_loss_epoch))/val_step)\n    val_loss_plot <- c(val_loss_plot, val_loss_epoch)\n    \n    val_accuaries_on_epoch <- mean(val_accuaries_on_epoch)\n    val_accuracy_plot <- c(val_accuracy_plot, val_accuaries_on_epoch)\n    \n    #Plotting\n    cat(\"Total loss (epoch): \", epoch, \": \", train_loss_epoch, \"\\n\")\n    cat(\"Distillater loss : \", epoch, \": \", distilation_loss_epoch, \"\\n\")\n    cat(\"Accuracy (epoch): \", epoch, \": \", accuracies_on_epoch, \"\\n\")\n    cat(\"Val loss : \", epoch, \": \", val_loss_epoch, \"\\n\")\n    cat(\"Val Accuracy (epoch): \", epoch, \": \", val_accuaries_on_epoch, \"\\n\")\n}","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"What about global_step = tf.train.get_or_create_global_step() describe [here](https://keras.rstudio.com/articles/eager_guide.html) ? It seems to only refers to the number of batches seen by the graph. [Source](https://stackoverflow.com/questions/41166681/what-does-global-step-mean-in-tensorflow)."},{"metadata":{},"cell_type":"markdown","source":"## Plotting"},{"metadata":{"trusted":true},"cell_type":"code","source":"total_loss_plot<-c()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#instead of collecting them during the training : \ntotal_loss_plot <- alpha * train_loss_plot + (1 - alpha) * distilation_loss_plot","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data <- data.frame(\"Student_loss\" = train_loss_plot, \n                    \"Distillation_loss\" = distilation_loss_plot,\n                   \"Total_loss\" = total_loss_plot,\n                    \"Epoch\" = 1:length(train_loss_plot),\n                    \"Val_loss\" = val_loss_plot,\n                    \"Train_accuracy\"= accuracy_plot,\n                    \"Val_accuracy\"= val_accuracy_plot)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"head(data)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Where total_loss is alpha * train_loss_plot * (1 - alpha) * distilation_loss_plot"},{"metadata":{"trusted":true},"cell_type":"code","source":"ggplot(data, aes(Epoch)) +\n  scale_colour_manual(values=c(Student_loss=\"#F8766D\",Val_loss=\"#00BFC4\", Distillation_loss=\"#DE8C00\", Total_loss=\"#1aff8c\")) +\n  geom_line(aes(y = Student_loss, colour = \"Student_loss\")) + \n  geom_line(aes(y = Val_loss, colour = \"Val_loss\")) + \n  geom_line(aes(y = Total_loss, colour = \"Total_loss\")) + \n  geom_line(aes(y = Distillation_loss, colour = \"Distillation_loss\"))\n#Validation set\nggplot(data, aes(Epoch)) + \n  geom_line(aes(y = Train_accuracy, colour = \"Train_accuracy\")) + \n  geom_line(aes(y = Val_accuracy, colour = \"Val_accuracy\"))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Fine tuning"},{"metadata":{"trusted":true},"cell_type":"code","source":"optimizer <- optimizer_adam(lr = 0.0001)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"unfreeze_weights(conv_base_student, from = 'block5a_expand_conv')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"nb_epoch<-25","execution_count":null,"outputs":[]},{"metadata":{"_kg_hide-input":true,"_kg_hide-output":false,"trusted":true},"cell_type":"code","source":"for (epoch in 1:nb_epoch) {\n    cat(\"Epoch: \", epoch, \" -----------\\n\")\n    # Init metrics\n    train_loss_epoch <- 0\n    accuracies_on_epoch <- c()\n    distilation_loss_epoch <- 0\n    val_loss_epoch <- 0\n    val_accuaries_on_epoch <- c()\n    \n    #Formula to not see the same batch over and over on each epoch\n    #Count epoch instead of epoch\n    count_epoch<-count_epoch+1\n    idx_batch <- (1+nb_batch*(count_epoch-1)):(nb_batch*count_epoch)\n    idx_val_set <- (1+val_step*(count_epoch-1)):(val_step*count_epoch)\n    \n    #Dirty solution to restart on a new validation batch generator before reaching the end of the other one \n    if (as.integer((dim(val_labels)[1]/8)-1) %in% idx_val_set) {\n        count_epoch<-1\n        idx_val_set <- (1+val_step*(count_epoch-1)):(val_step*count_epoch)\n        validation_generator <- do.call(flow_images_from_dataframe, arg.list)\n    }\n    #need the same if for train generator\n    if (as.integer((dim(train_labels)[1]/8)-1) %in% idx_batch) {\n        count_epoch<-1\n        idx_batch <- (1+nb_batch*(count_epoch-1)):(nb_batch*count_epoch)\n        train_generator <- do.call(flow_images_from_dataframe, arg.list)\n    }\n    \n    for (batch in idx_batch) {\n        x = train_generator[batch][[1]]\n        y = train_generator[batch][[2]]\n        # Forward pass of teacher\n        teacher_predictions = model(x)\n\n        with(tf$GradientTape() %as% tape, {\n            student_predictions = student(x)\n            student_loss = tf$losses$categorical_crossentropy(y, student_predictions)\n        \n            distillation_loss = tf$losses$categorical_crossentropy(tf$nn$softmax(teacher_predictions/temperature, axis=0L), \n                                                           tf$nn$softmax(student_predictions/temperature, axis=0L))\n        \n            loss = alpha * student_loss + (1 - alpha) * distillation_loss\n            })\n        \n        # Compute gradients\n        # Variating learning rate :\n        # optimizer <- optimizer_adam(lr = 0.0001)\n        gradients <- tape$gradient(loss, student$trainable_variables)\n        optimizer$apply_gradients(purrr::transpose(list(gradients, student$trainable_variables)))\n        \n        #Collect the metrics of the student\n        train_loss_epoch <- train_loss_epoch + student_loss\n        distilation_loss_epoch <- distilation_loss_epoch + distillation_loss\n        \n        accuracy_on_batch <- train_accuracy(y_true=y, y_pred=student_predictions)\n        accuracies_on_epoch <- c(accuracies_on_epoch, as.numeric(accuracy_on_batch))\n        \n    }\n\n    #Collect info on current epoch and for graphs and cat()\n    train_loss_epoch <- mean(as.vector(as.numeric(train_loss_epoch))/nb_batch)\n    train_loss_plot <- c(train_loss_plot, train_loss_epoch)\n    \n    distilation_loss_epoch <- mean(as.vector(as.numeric(distilation_loss_epoch))/nb_batch)\n    distilation_loss_plot <- c(distilation_loss_plot, distilation_loss_epoch)\n    \n    accuracies_on_epoch <- mean(accuracies_on_epoch)\n    accuracy_plot <- c(accuracy_plot, accuracies_on_epoch)\n    \n    \n    for (step in idx_val_set) {\n        # Unpack the data\n        x = validation_generator[step][[1]]\n        y = validation_generator[step][[2]]\n\n        # Compute predictions\n        student_predictions = student(x)\n\n        # Calculate the loss\n        student_loss = tf$losses$categorical_crossentropy(y, student_predictions)\n\n        #Collect the metrics of the student\n        #This line will create a bug of shape when val_loss end.\n        val_loss_epoch <- val_loss_epoch + student_loss\n        \n        accuracy_on_val_step <- train_accuracy(y_true=y, y_pred=student_predictions)\n        val_accuaries_on_epoch <- c(val_accuaries_on_epoch, as.numeric(accuracy_on_val_step))\n    }\n    \n    #Collect info on current epoch and for graphs and cat()\n    val_loss_epoch <- mean(as.vector(as.numeric(val_loss_epoch))/val_step)\n    val_loss_plot <- c(val_loss_plot, val_loss_epoch)\n    \n    val_accuaries_on_epoch <- mean(val_accuaries_on_epoch)\n    val_accuracy_plot <- c(val_accuracy_plot, val_accuaries_on_epoch)\n    \n    #Plotting\n}","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"total_loss_plot <- alpha * train_loss_plot + (1 - alpha) * distilation_loss_plot","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data <- data.frame(\"Student_loss\" = train_loss_plot, \n                    \"Distillation_loss\" = distilation_loss_plot,\n                   \"Total_loss\" = total_loss_plot,\n                    \"Epoch\" = 1:length(train_loss_plot),\n                    \"Val_loss\" = val_loss_plot,\n                    \"Train_accuracy\"= accuracy_plot,\n                    \"Val_accuracy\"= val_accuracy_plot)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"ggplot(data, aes(Epoch)) +\n  scale_colour_manual(values=c(Student_loss=\"#F8766D\",Val_loss=\"#00BFC4\", Distillation_loss=\"#DE8C00\", Total_loss=\"#1aff8c\")) +\n  geom_line(aes(y = Student_loss, colour = \"Student_loss\")) + \n  geom_line(aes(y = Val_loss, colour = \"Val_loss\")) + \n  geom_line(aes(y = Total_loss, colour = \"Total_loss\")) + \n  geom_line(aes(y = Distillation_loss, colour = \"Distillation_loss\"))\n#Validation set\nggplot(data, aes(Epoch)) + \n  geom_line(aes(y = Train_accuracy, colour = \"Train_accuracy\")) + \n  geom_line(aes(y = Val_accuracy, colour = \"Val_accuracy\"))","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"name":"ir","display_name":"R","language":"R"},"language_info":{"name":"R","codemirror_mode":"r","pygments_lexer":"r","mimetype":"text/x-r-source","file_extension":".r","version":"3.6.3"}},"nbformat":4,"nbformat_minor":4}