---
title: "Leaf Doctor: R EDA with Torch"
output:
  html_document:
    number_sections: true
    fig_caption: true
    toc: true
    fig_width: 5
    fig_height: 4
    theme: cosmo
    highlight: tango
    code_folding: show
---

```{r setup, include=FALSE}
knitr::opts_chunk$set(echo = TRUE)
remotes::install_local(path = "../input/torchvision/torchvision_0.1.0/torchvision")
```

<center>
![Illustration taken from https://in.pinterest.com/pin/21181060735895627/](https://i.imgur.com/ZRYZxLu.png)
</center>

# Introduction
This is a basic Exploratory Data Analysis for the Cassava Leaf Disease Classification competition with R and torch library. R is extensively and traditionally used in [biology](https://www.bioconductor.org/). Though there is a great deal of image processing libraries, R is not used in this field of study frequently. With a new [torch](https://torch.mlverse.org/) package, which is written in R using the **libtorch** library, R got new opportunities for deep learning and image analysis without the need to run third-party scripting languages under the hood.

The aim of this challenge is to implement a robust classifier of cassava images into four disease categories or a fifth category indicating a healthy leaf. This may help farmers to identify diseased plants.

This is a supervised machine learning problem which is evaluated on [categorization accuracy](https://developers.google.com/machine-learning/crash-course/classification/accuracy).

# Preparations {.tabset .tabset-fade}
## Libraries
Here we load the essential packages.
```{r load_lib, message=FALSE, warning=FALSE, results='hide'}
library(torch)
library(knitr)
library(magick)
library(rsample)
library(magrittr)
library(tidyverse)
library(torchvision)
```
## Constants
We set a seed, define a device on which torch tensors will be allocated, a path to data, etc...
```{r set_const, message=FALSE, warning=FALSE, results='hide'}
set.seed(0)
torch_manual_seed(0)
device <- if (cuda_is_available()) torch_device("cuda:0") else "cpu"
path <- "../input/cassava-leaf-disease-classification/"
epochs <- 50
img_size <- 224
batch_size <- 128
smoothing <- 0.02
print_every_n <- 3
early_stopping_steps <- 3
```
## Load Data
We load tabular data for this competition.
```{r load_tab, message=FALSE, warning=FALSE, results='hide'}
tr <- read_csv(str_c(path, "train.csv"))
sub <- read_csv(str_c(path, "sample_submission.csv"))
dmap <- jsonlite::read_json(str_c(path, "label_num_to_disease_map.json")) %>% 
  as_tibble() %>% 
  pivot_longer(matches("[[:digit:]]"), names_to = "label", values_to = "disease") %>% 
  mutate(label = as.integer(label),
         disease = str_remove(disease, "\\(.*\\)"))
```

# Overview
## Table Data
Let's check the content of the **train.csv** file:
```{r tr_csv, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
head(tr, 5) %>% kable()
```
<div></div>
It has a simple structure with two fields - the image file name and the ID code for the disease. The **sample_submission.csv** has the same columns:
```{r sub_csv, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
head(sub) %>% kable()
```
<div></div>
Also there is a *label_num_to_disease_map.json* file which maps each disease code into the real disease name:
```{r lab_json, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
kable(dmap) 
```
<div></div>
Let's check the balance between different classes in the train set:
```{r imbal, message=FALSE, warning=FALSE, echo=TRUE, results='show', fig.align='center'}
tr %>% 
  left_join(dmap, by = "label") %>% 
  ggplot(aes(x = disease)) +
  geom_bar(stat = "count", fill = "steelblue", width = 0.600) +
  geom_text(stat = "count", aes(label=..count..), vjust = 1.2, color = "white") +
  theme_minimal() +
  labs(y = "") +
  theme(axis.text.x = element_text(angle = 40, hjust = 1))
```

There is an imbalance between classes. Perhaps, we should try oversampling/undersampling to adjust the class distribution.

## Image Data
The dataset includes the large number of **jpg** images: 

* The training set has `r str_c(path, "train_images/") %>% list.files(pattern = "jpg$") %>% length()` files
* The test set has `r str_c(path, "test_images/") %>% list.files(pattern = "jpg$") %>% length()` files. In this competition, the test set is hidden. Thus, this one file is a sample. After submission the test set will be replaced by a bigger one to calculate the score.

We are provided with **tfrecords** files which contain the image files in tfrecord format, but I don't use them in this EDA.

Let's inspect the images for each label along with citations from Wikipedia. There will be 3 images in a row - original, after applying a sharpening kernel and fuzzy c-means segmentation.
```{r fun_view, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
show_img <- function(L) {
  image_id <- tr %>% 
    filter(label == L) %>% 
    sample_n(3) %$% 
    image_id
  
  images <- image_blank(0, 0)
  for (id in image_id) {
    im <- image_read(str_c(path, "train_images/", id))
    images <- image_append(c(images,
                             image_append(c(im, 
                                            image_convolve(im, 
                                                           kernel = '3x3: 0   -1.2  0
                                                                         -1.2  6   -1.2
                                                                          0   -1.2  0 ', 
                                                           bias = "0%"),
                                            image_fuzzycmeans(im, 25, 2.5)))),
                           stack = TRUE)
  }
  images
}
```

0. [Cassava Bacterial Blight](https://en.wikipedia.org/wiki/Bacterial_blight_of_cassava). "Originally discovered in Brazil in 1912, the disease has followed cultivation of cassava across the world. Among diseases which afflict cassava worldwide, bacterial blight causes the largest losses in terms of yield."

```{r view0, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
show_img(0)
```

1. [Cassava Brown Streak Disease](https://en.wikipedia.org/wiki/Cassava_brown_streak_virus_disease). "It was first identified in 1936 in Tanzania. This disease is considered to be the biggest threat to food security in coastal East Africa and around the eastern lakes."
```{r view1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
show_img(1)
```

2. [Cassava Green Mottle](https://en.wikipedia.org/wiki/Cassava_green_mottle_virus) "is a plant pathogenic virus of the family Secoviridae."
```{r view2, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
show_img(2)
```

3. [Cassava Mosaic Disease](https://en.wikipedia.org/wiki/Cassava_mosaic_virus). "The first report of cassava mosaic disease was from East Africa in 1894. Since then, epidemics have occurred throughout the African continent, resulting in great economic loss and devastating famine."
```{r view3, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
show_img(3)
```

4. Healthy Cassava Leaves. To me, some healthy leaves seem not very healthy. This means that some labels might be erroneous. 
```{r view4, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
show_img(4)
```

In some cases fuzzy c-means segmentation successfully highlights the damaged parts of the leaves.
Additionally, there are pictures of tubers among the images, which increases the difficulty of the problem.
```{r view_tubers, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
image_append(c(
  image_read(str_c(path, "train_images/", "1004389140.jpg")),
  image_read(str_c(path, "train_images/", "1014492188.jpg")),
  image_read(str_c(path, "train_images/", "1008244905.jpg"))
))
```

# Libtorch and AlexNet
This is the main part of the EDA. I'm going to describe how to load and finetune a pretrained [AlexNet](https://torchvision.mlverse.org/reference/model_alexnet.html) model for classification. 
But first we have to do some preparations.

## Dataset Class 
Here we define a dataset class, an abstract container that knows how to iterate over our data. We implement **.getitem()** method, which loads an image and performs basic transformations, and **.length()** method, which just returns the length of the dataset, in an [R6 class](https://cran.r-project.org/web/packages/R6/index.html). Also we define **do_augment()** helper function which performs different transformations on an image.

```{r ds_helper, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
n_classes <- tr$label %>% n_distinct()

to_device <- function(x, device) x$to(device = device)

do_augment <- . %>%
  transform_random_horizontal_flip(p = 0.5) %>%
  transform_random_vertical_flip(p = 0.5) 
```
```{r ds1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
cas_dataset <- dataset(
  name = "cassava_dataset",
  
  initialize = function(df, p, augment = FALSE, output_label = TRUE) {
    self$df <- df
    self$p <- p
    self$augment <- augment
    self$output_label <- output_label
  },
  
  .getitem = function(idx) {
    img <- str_c(self$p, self$df$image_id[idx]) %>% 
      base_loader() %>% 
      transform_to_tensor() %>% 
      to_device(device) %>% 
      transform_center_crop(img_size) %>% 
      {if (self$augment) do_augment(.) else .} %>% 
      transform_normalize(mean = c(0.485, 0.456, 0.406),
                          std = c(0.229, 0.224, 0.225))
    
    if (self$output_label) {
      list(img, 
           torch_tensor(self$df$label[idx] + 1, device = device, dtype = torch_long()))
    } else {
      list(img)
    }
  },
  
  .length = function() nrow(self$df),
)
```

## Step and Infer Functions
[R torch](https://cran.r-project.org/web/packages/torch/index.html) is a low-level framework, so we have to code the functions for training, validation and inferring by ourselves. The **step_nn()** function trains the libtorch model for 1 epoch or evaluates it. It returns both loss and metric:
```{r run1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
step_nn <- function(m_nn, dl, loss_fn, metric_fn, optimizer = NULL, scheduler = NULL, do_train = TRUE) {
  if (do_train) m_nn$train() else m_nn$eval()
  L <- c()
  M <- c()
  for (batch in enumerate(dl)) {
    Y_h <- m_nn(batch[[1]])
    loss <- loss_fn(Y_h, batch[[2]][, 1])
    L <- c(L, loss$item())
    metric <- metric_fn(Y_h, batch[[2]][, 1])
    M <- c(M, metric$item())
    if (do_train) {
      optimizer$zero_grad()
      loss$backward()
      optimizer$step()
      if (!is.null(scheduler)) scheduler$step()
    }
  }
  list(loss = mean(L), metric = mean(M))
}
```
The **infer()** function disables gradients and calculates the output of the model:
```{r infer1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
infer <- function(nn_file, dl) {
  preds <- c()
  m_nn <- torch_load(nn_file)
  m_nn$eval()
  for (batch in enumerate(dl)) {
    with_no_grad({
      batch_pred <- batch[[1]] %>% 
        m_nn() %>% 
        nnf_softmax(dim = 2) %>% 
        to_device(device = "cpu") %>% 
        as_array()
      preds <- rbind(preds, batch_pred)
    })
  }
  preds
}
```

## Dataloaders
Here I use a simple binary split into a training and validation sets:
```{r split1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
rs <- initial_split(tr, 0.85)
tri <- rs$in_id
```
Having defined a dataset class and a train/validation split we need to create data loaders for the train, validation and test sets on this basis. 
```{r dl1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
X <- cas_dataset(tr[tri, ], 
                 str_c(path, "train_images/"),
                 augment = TRUE, 
                 output_label = TRUE) %>% dataloader(batch_size)
X_val <- cas_dataset(tr[-tri, ], 
                     str_c(path, "train_images/"), 
                     augment = FALSE, 
                     output_label = TRUE) %>%  dataloader(batch_size)
X_te <- cas_dataset(sub, 
                    str_c(path, "test_images/"), 
                    augment = FALSE,
                    output_label = FALSE) %>% dataloader(batch_size)
```

## Pretrained Model
Now we are ready to implement the next task - building the model. One could obtain a pretrained model by using `m_nn <- model_alexnet(pretrained = TRUE)`, but this is a Research Code Competition and Internet access is disabled, so I just load the state dictionary of the AlexNet from my private dataset. In order to save the weights of the pretrained model we have to freeze them by disabling gradients calculation. Also, we have to modify the last block of the AlexNet:
```{r std1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
load_net <- function(n_classes, dpath) {
  state_dict <- load_state_dict(dpath)
  m_nn <- model_alexnet(pretrained = FALSE)
  m_nn$load_state_dict(state_dict)
  
  for (p in m_nn$parameters) p$requires_grad <- FALSE
  
  m_nn$classifier <- nn_sequential(nn_dropout(0.4),
                                   nn_linear(256 * 6 * 6, 4096),
                                   nn_relu(inplace = TRUE),
                                   nn_dropout(0.2),
                                   nn_linear(4096, 512),
                                   nn_relu(inplace = TRUE),
                                   nn_linear(512, n_classes))
  m_nn$to(device = device)
}
```

## Optimizer
We use a simple [SGD optimizer](https://torch.mlverse.org/docs/reference/optim_sgd.html) with Nesterov momentum, so there are a lot of capabilities for tuning. Also we define a [cross entropy loss function](https://torch.mlverse.org/docs/reference/nn_cross_entropy_loss.html) and accuracy metric:
```{r opt1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
def_optimizer <- function(m_nn, loss_fn, lr = 1e-04, momentum = 0.95, weight_decay = 1e-06, nesterov = TRUE) {
  list(loss_fn = loss_fn, 
       metric_fn = function(Y_h, y) (Y_h$argmax(dim = 2)$add(1) == y)$sum() / batch_size,
       optimizer = optim_sgd(m_nn$parameters, lr = lr, momentum = momentum, weight_decay = weight_decay, nesterov = nesterov))
}
```

## Training
In the training loop we call the `step_nn()` function to update network weights and evaluate our model: 
```{r train1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
fit_nn <- function(m_nn, X, X_val, loss_fn, metric_fn, optimizer, epochs, m_name = "m_nn.pt") {
  early_steps <- 0
  best_loss <- Inf
  best_metric <- -Inf
  for (epoch in seq_len(epochs)) {
    
    tr_loss_m <- step_nn(m_nn, X, loss_fn, metric_fn, optimizer)
    val_loss_m <- step_nn(m_nn, X_val, loss_fn, metric_fn, do_train = FALSE)
    
    if (best_loss > val_loss_m$loss) {
      best_metric <- val_loss_m$metric
      best_loss <- val_loss_m$loss
      torch_save(m_nn, m_name)
    } else {
      early_steps <- early_steps + 1
      if (early_steps >= early_stopping_steps) break
    }
    
    if (epoch %% print_every_n == 0 || epoch == 1) {
      sprintf(
        "Epoch %d:
          train_loss   % 8.4f | val_loss   % 8.4f | best val_loss   % 8.4f
          train_metric % 8.4f | val_metric % 8.4f | best val_metric % 8.4f\n", 
        epoch, 
        tr_loss_m$loss,   val_loss_m$loss,   best_loss,
        tr_loss_m$metric, val_loss_m$metric, best_metric) %>% 
        cat()
    }
  }
  sprintf("Stopped at %d epoch with best val_loss %.4f and val_metric %.4f\n", epoch, best_loss, best_metric) %>% cat()
}
```
```{r train2, message=FALSE, warning=FALSE, echo=TRUE, results='show', eval=TRUE}
m_nn <- load_net(n_classes, "../input/torchvision/alexnet.pth")
opt <- def_optimizer(m_nn, nn_cross_entropy_loss())
m_nn %>% 
  fit_nn(X, X_val, opt$loss_fn, opt$metric_fn, opt$optimizer, 1, "alexnet_t4_a73_cpu.pt")
```
Actually, to save the time I added the updated AlexNet (as **alexnet_t4_a73_cpu.pt**) to my private dataset, that's why the previous chunk have been running for 1 epoch only.

## Submission
Here we use updated model to predict the class of each image. This model achieves accuracy of about 0.75:
```{r sub1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
preds <- infer("../input/torchvision/alexnet_t4_a73_cpu.pt", X_te)
sub$label <- apply(preds, 1, which.max) - 1
write_csv(sub, "submission_0.753.csv")
```

## Model Evaluation
Let's build a confusion matrix for the validation set:
```{r val1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
preds_val <- infer("../input/torchvision/alexnet_t4_a73_cpu.pt", X_val)
preds_val <- apply(preds_val, 1, which.max) - 1
cm <- tibble(obs = as_factor(tr$label[-tri]), 
             pred = as_factor(preds_val)) %>% 
  yardstick::conf_mat(obs, pred)
summary(cm)
```
From all the metrics above we are interested in accuracy, which is not that bad by the way. Though sensitivity of the model is quite low. And this is the confusion matrix without normalization:
```{r val2, message=FALSE, warning=FALSE, echo=TRUE, results='show', fig.align='center'}
autoplot(cm, type = "heatmap") +
  scale_fill_gradient(low = "#D6EAF8", high = "#2E86C1")
```

Below we can see errors per class:
```{r val3, message=FALSE, warning=FALSE, echo=TRUE, results='show', fig.align='center'}
errors <- ((colSums(cm$table) - diag(cm$table)) / colSums(cm$table) * 100)
tibble(label = names(errors), error = round(errors, 2)) %>% 
  ggplot(aes(x = label, y = error)) +
  geom_col(fill = "steelblue", width = 0.600) +
  geom_text(aes(label = str_c(error, "%")), vjust = 1.2, color = "white") +
  theme_minimal() +
  labs(y = "") 
```

The less the number of elements in the class the higher the error - thus, Bacterial blight of cassava is the toughest class to predict for this model. The disbalance between classes distribution can cause great problems in future.

# Squeezing out more from the model
## More Augmentations
I did a few experiments with image data augmentation and modified the `do_augment()` function as follows:
```{r som1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
do_augment <- function(img, p = 0.5) {
  img %>% 
    { if (runif(1) > p) image_flip(.) else . } %>% 
    { if (runif(1) > p) image_flop(.) else . } %>% 
    { if (runif(1) > p) image_median(., 1.5) else . } %>% 
    { if (runif(1) > p) image_modulate(., brightness = 125) else . } %>%
    { if (runif(1) > p) image_modulate(., saturation = 125) else . } %>%
    { if (runif(1) > p) image_noise(., noisetype = "Impulse") else .}  %>%
    { if (runif(1) > p) image_blur(., 0, 1.5) else . } %>% 
    { if (runif(1) > p) image_rotate(., runif(1, -15, 15)) else . } %>% 
    { if (runif(1) > p) image_convolve(., kernel = '3x3: 0   -1.2  0
                                                        -1.2  6   -1.2
                                                         0   -1.2  0 ', bias = "0%") else . }
}
```
Here we use the R [magick package](https://cran.r-project.org/web/packages/magick/index.html), which wraps [the ImageMagick STL](https://www.imagemagick.org/Magick++/STL.html), to create transformed versions of the images. Below we can see some transformations of the image after applying of the new **do_augment()** function:
```{r som11, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
img <- str_c(path, "train_images/", "1000015157.jpg") %>% image_read()
image_append(c(img, map(1:4, ~ do_augment(img))))
```

I did a few experiments and they showed that the better score for my model can be achieved using only the `image_flop()` and `image_median()` transformations. Thus, the `do_augment()` function can be reduced to the following few lines:
```{r som111, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
do_augment <- function(img, p = 0.5) {
  img %>% 
    { if (runif(1) > p) image_flop(.) else . } %>% 
    { if (runif(1) > p) image_median(., 1.5) else . }
}
```
Also I modified the `cas_dataset` in order to make it work with the images transformed by the **magick package**:
```{r som2, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
cas_dataset <- dataset(
  name = "cassava_dataset",
  
  initialize = function(df, p, augment = FALSE, output_label = TRUE) {
    self$df <- df
    self$p <- p
    self$augment <- augment
    self$output_label <- output_label
  },
  
  .getitem = function(idx) {
    img <- str_c(self$p, self$df$image_id[idx]) %>% 
      image_read() %>% 
      { if (self$augment) do_augment(.) else . } %>% 
      image_crop(geometry_area(img_size, img_size), "Center") %>% 
      transform_to_tensor() %>% 
      to_device(device) %>% 
      transform_normalize(mean = c(0.485, 0.456, 0.406), 
                          std = c(0.229, 0.224, 0.225))
    
    if (self$output_label) {
      list(img, 
           torch_tensor(self$df$label[idx] + 1, device = device, dtype = torch_long()))
    } else {
      list(img)
    }
  },
  
  .length = function() nrow(self$df),
)
```
We need to recreate the dataloaders using the new definition of the `cas_dataset` class:
```{r som3_tr, message=FALSE, warning=FALSE, echo=TRUE, results='show', eval=TRUE}
X <- cas_dataset(tr[tri, ], 
                 str_c(path, "train_images/"),
                 augment = TRUE, 
                 output_label = TRUE) %>% dataloader(batch_size)
X_val <- cas_dataset(tr[-tri, ], 
                     str_c(path, "train_images/"), 
                     augment = FALSE, 
                     output_label = TRUE) %>% dataloader(batch_size)
X_te <- cas_dataset(sub, 
                    str_c(path, "test_images/"), 
                    augment = FALSE,
                    output_label = FALSE) %>% dataloader(batch_size)
```
## Label Smoothing
Label smoothing might help to create more stable predictions making a model less overconfident:
```{r som5_tr, message=FALSE, warning=FALSE, echo=TRUE, results='show', eval=TRUE}
nn_label_smoothing_loss <- nn_module(
  "nn_label_smoothing_loss",
  initialize = function(n_classes, smoothing = 0, dim = -1) {
    self$confidence <- 1 - smoothing
    self$smoothing <- smoothing
    self$n_classes <- n_classes
    self$dim <- dim
  },
  forward = function(pred, target) {
    pred <- pred$log_softmax(dim = self$dim)
    true_dist <- torch_full_like(pred, self$smoothing / (self$n_classes - 1), device = "cpu")
    true_dist$scatter_(2, target$detach()$cpu()$unsqueeze(2), self$confidence)
    torch_mean(torch_sum(-true_dist$to(device = device) * pred, dim = self$dim))
  }
) 
```
## Alexnet with Label Smoothing 
```{r som6_tr, message=FALSE, warning=FALSE, echo=TRUE, results='show', eval=TRUE}
m_nn <- load_net(n_classes, "../input/torchvision/alexnet.pth")
opt <- def_optimizer(m_nn, 
                     nn_label_smoothing_loss(n_classes, smoothing)$to(device = device))
m_nn %>% 
  fit_nn(X, X_val, opt$loss_fn, opt$metric_fn, opt$optimizer, 1, "alexnet_t7_a77_cpu.pt")
```
This new submission improves the previous score:
```{r som_sub1, message=FALSE, warning=FALSE, echo=TRUE, results='show'}
preds <- infer("../input/torchvision/alexnet_t7_a77_cpu.pt", X_te)
sub$label <- apply(preds, 1, which.max) - 1
write_csv(sub, "submission.csv")
```

To be continued...