{"cells":[{"metadata":{},"cell_type":"markdown","source":"Based on [Notebook - Cassava classification - EDA & fastai starter](https://www.kaggle.com/tanlikesmath/cassava-classification-eda-fastai-starter)"},{"metadata":{},"cell_type":"markdown","source":"# Setup"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Install PyTorch Image Models package (TIMM)\n!pip install ../input/timm031/timm-0.3.1-py3-none-any.whl","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import numpy as np\nimport os\nimport pandas as pd\nimport time\n\nfrom fastai.vision.all import *","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"nb_start = time.time()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Notebook variables\ndata_dir = Path('../input/cassava-leaf-disease-classification')\nsample_fraction = 1\nseed = 999","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"set_seed(seed)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Preprocess"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Read\ntrain_df = pd.read_csv(data_dir/'train.csv')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Process\ntrain_df = (train_df\n    .assign(path=train_df['image_id'].map(lambda x:data_dir/'train_images'/x))\n    .drop(columns=['image_id'])\n    .sample(frac=sample_fraction)\n    .reset_index(drop=True))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Showcase\nprint(train_df.shape[0])\ntrain_df.head()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# EDA"},{"metadata":{"trusted":true},"cell_type":"code","source":"from PIL import Image\n\nim = Image.open(train_df['path'][0])\nwidth, height = im.size\nprint(width,height)\nim","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Data Loader"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Create data loader\nitem_tfms = RandomResizedCrop(460, min_scale=0.75, ratio=(1.,1.))\nbatch_tfms = [*aug_transforms(size=224, max_warp=0), Normalize.from_stats(*imagenet_stats)]\nbs=32\n\ndls = ImageDataLoaders.from_df(\n    df=train_df,\n    valid_pct=0.2,\n    seed=seed,\n    label_col=0, # label is in the first column of the DataFrame\n    fn_col=1, # filename/path is in the second column of the DataFrame\n    bs=bs, # pass in batch size\n    item_tfms=item_tfms, # pass in item_tfms\n    batch_tfms=batch_tfms) # pass in batch_tfms","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Showcase data loader\ndls.show_batch()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Model Training"},{"metadata":{},"cell_type":"markdown","source":"## Setup"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Needed for making pretrained weights work without needing to find the default filename\n# EfficientNet-B3 model\nif not os.path.exists('/root/.cache/torch/hub/checkpoints/'):\n        os.makedirs('/root/.cache/torch/hub/checkpoints/')\n!cp '../input/timmefficientnet/tf_efficientnet_b3_ns-9d44bf68.pth' '/root/.cache/torch/hub/checkpoints/tf_efficientnet_b3_ns-9d44bf68.pth'","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Functions from: [walkwithfastai - Utilizing the timm Library Inside of fastai](https://walkwithfastai.com/vision.external.timm)"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Utilities\nfrom timm import create_model\nfrom fastai.vision.learner import _update_first_layer\n\n\ndef create_timm_body(arch:str, pretrained=True, cut=None, n_in=3):\n    \"Creates a body from any model in the `timm` library.\"\n    model = create_model(arch, pretrained=pretrained, num_classes=0, global_pool='')\n    _update_first_layer(model, n_in, pretrained)\n    if cut is None:\n        ll = list(enumerate(model.children()))\n        cut = next(i for i,o in reversed(ll) if has_pool_type(o))\n    if isinstance(cut, int): return nn.Sequential(*list(model.children())[:cut])\n    elif callable(cut): return cut(model)\n    else: raise NamedError(\"cut must be either integer or function\")\n\n        \ndef create_timm_model(arch:str, n_out, cut=None, pretrained=True, n_in=3, init=nn.init.kaiming_normal_, custom_head=None,\n                     concat_pool=True, **kwargs):\n    \"Create custom architecture using `arch`, `n_in` and `n_out` from the `timm` library\"\n    body = create_timm_body(arch, pretrained, None, n_in)\n    if custom_head is None:\n        nf = num_features_model(nn.Sequential(*body.children())) * (2 if concat_pool else 1)\n        head = create_head(nf, n_out, concat_pool=concat_pool, **kwargs)\n    else: head = custom_head\n    model = nn.Sequential(body, head)\n    if init is not None: apply_init(model[1], init)\n    return model\n\n\ndef timm_learner(dls, arch:str, loss_func=None, pretrained=True, cut=None, splitter=None,\n                y_range=None, config=None, n_out=None, normalize=True, **kwargs):\n    \"Build a convnet style learner from `dls` and `arch` using the `timm` library\"\n    if config is None: config = {}\n    if n_out is None: n_out = get_c(dls)\n    assert n_out, \"`n_out` is not defined, and could not be inferred from data, set `dls.c` or pass `n_out`\"\n    if y_range is None and 'y_range' in config: y_range = config.pop('y_range')\n    model = create_timm_model(arch, n_out, default_split, pretrained, y_range=y_range, **config)\n    learn = Learner(dls, model, loss_func=loss_func, splitter=default_split, **kwargs)\n    if pretrained: learn.freeze()\n    return learn","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Train (Stage 1)\nTrain frozen pre-trained model for single epoch "},{"metadata":{"trusted":true},"cell_type":"code","source":"# Define learner\nlearn = timm_learner(\n    dls=dls, \n    arch='tf_efficientnet_b3_ns',\n    loss_func=LabelSmoothingCrossEntropy(),\n    opt_func=ranger,\n    metrics=[accuracy]\n).to_native_fp16()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# # Find optimal learning rate for pre-trained model\n# start = time.time()\n# learn.lr_find()\n# print(\"{:.2f}min\".format(int(time.time() - start) / 60))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Train frozen pretrained model for single epoch\nstart = time.time()\nlearn.freeze()\nlearn.fit_flat_cos(1, 10e-2, wd=0.5, cbs=[MixUp()])\nprint(\"{:.2f}min\".format(int(time.time() - start) / 60))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Save stage-1 model\nlearn.save('stage-1')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Read stage-1 model\nlearn = learn.load('stage-1')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Validation loss \nlearn.recorder.plot_loss()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Train (Stage 2)\nTrain entire model for several epochs"},{"metadata":{"trusted":true},"cell_type":"code","source":"# # Find optimal learning rate for model\n# start = time.time()\n# learn.unfreeze()\n# learn.lr_find()\n# print(\"{:.2f}min\".format(int(time.time() - start) / 60))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"start = time.time()\nlearn.unfreeze()\nlearn.fit_flat_cos(5, 2e-3,pct_start=0, cbs=[MixUp()])\nprint(\"{:.2f}min\".format(int(time.time() - start) / 60))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learn.recorder.plot_loss()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learn.save('stage-2')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Analyze Model"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Plot confusion matrix\nlearn_32 = learn.to_native_fp32()\ninterp = ClassificationInterpretation.from_learner(learn_32)\ninterp.plot_confusion_matrix()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Inference"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Read sample dataset\nsample_df = pd.read_csv(data_dir/'sample_submission.csv')\nsample_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Create submission dataset\n_sample_df = sample_df.copy()\n_sample_df['path'] = _sample_df['image_id'].map(lambda x:data_dir/'test_images'/x)\n_sample_df = _sample_df.drop(columns=['image_id'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Create test set data loader\ntest_dl = dls.test_dl(_sample_df)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Showcase test set data loader\ntest_dl.show_batch()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Create predictions\npreds, _ = learn.tta(dl=test_dl, n=8, beta=0)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Submission"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Create and save submission file\nsample_df['label'] = preds.argmax(dim=-1).numpy()\nsample_df.to_csv('submission.csv',index=False)\nsample_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print(\"{:.2f}min\".format(int(time.time() - nb_start) / 60))","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat":4,"nbformat_minor":4}