---
title: "Work with fast.ai from R"
author: "by Turgut Abdullayev"
date: '`r Sys.Date()`'
output:
  html_document:
    number_sections: true
    toc: true
---


```{r setup, include=FALSE}
knitr::opts_chunk$set(echo = TRUE, eval = TRUE)
```

## R interface to fastai

The fastai package provides R wrappers to [fastai](https://github.com/fastai/fastai).

The fastai library simplifies training fast and accurate neural nets using modern best practices. See the [fastai website](https://henry090.github.io/fastai/) to get started. The library is based on research into deep learning best practices undertaken at ```fast.ai```, and includes "out of the box" support for ```vision```, ```text```, ```tabular```, [audio](https://github.com/fastaudio/fastaudio), [time-series](https://github.com/tcapelle/timeseries_fastai) and ```collab``` (collaborative filtering) models. 

                                                                                             
```{r eval=T}
devtools::install_github('rstudio/reticulate',dependencies=FALSE)
devtools::install_github("henry090/fastai",dependencies=FALSE)
library(fastai)
install_fastai(gpu = TRUE)
detach("package:fastai", unload = TRUE)
```

## Dataset

Read train dataset:

```{r}
library(fastai)
library(magrittr)

path = "../input/cassava-leaf-disease-classification"

df = data.table::fread(paste(path,"train.csv",sep="/"))
df[['label']] = as.character(df[['label']])

```


## Datablock

Create datablock:

```{r}
batch_tfms = list(RandomResizedCrop(256),
                  #DeterministicDihedral(),
                  #Warp(), Hue(), Saturation(),
                  Hue(0.2, 0.5),
                  aug_transforms(size = 300, max_rotate = 180.,
                                 max_lighting = 0.2, max_warp = 0.4,
                                 flip_vert = TRUE, max_zoom = 2.),
                  Normalize_from_stats( imagenet_stats() ))

plant = DataBlock(blocks = list(ImageBlock(), CategoryBlock()),
                         get_x = function(x) {paste(paste(path,'train_images',sep = '/'), x[[0]], sep = '/')},
                         get_y = ColReader('label'),
                         item_tfms = Resize(300),
                         splitter = RandomSplitter(),
                         batch_tfms = batch_tfms) 

dls = plant %>% dataloaders(df, bs = 50)

dls %>% show_batch(max_n = 16)
```

Efficientnet archs:
    
```{r}
grep('^efficient',timm_list_models(),value = TRUE)
```

All models:
    
```{r}
timm_list_models()
```


## Learner

Create ```Learner```:

```{r}
#learn = cnn_learner(audio_dbunch, 
#                    resnet101(), 
#                    opt_func = Adam,
#                    loss_func = BCEWithLogitsLossFlat(),
#                    metrics = accuracy_multi())

# alternatively work with models from timm package
learn = timm_learner(dls, 
                     'efficientnet_lite0',
                     opt_func=Adam,
                     loss_func=CrossEntropyLossFlat(),
                     metrics=list(accuracy,error_rate))

learn %>% summary()
```


## Training

Train your model:

```{r}
# uncomment and train
#tune = learn %>% fine_tune(1,freeze_epochs = 1)
```

## Conclusion

And make predictions for test dataset:

```{r}
test_files = as.character(list.files(paste(path,'test_images',sep = '/'),full.names = T))
test_dl = learn$dls$test_dl(test_files)
predictions = learn$get_preds(dl = test_dl, with_decoded = TRUE)

#names without full path
test_files = as.character(list.files(paste(path,'test_images',sep = '/'),full.names = F))

submission = data.frame(image_id=test_files,label=as_array(predictions[[3]]))

data.table::fwrite(submission,'submission.csv')
```