{"cells":[{"metadata":{"_uuid":"88804b28-2e9c-4bd9-a505-b040233c8dad","_cell_guid":"d49780c9-ad99-4592-a8b1-2fccb6be2673","trusted":true},"cell_type":"markdown","source":"# Flowers image classification","execution_count":null},{"metadata":{"_uuid":"6d95ab39-0f1c-4b18-9e40-cadaa82db9e8","_cell_guid":"40bdadb2-1a67-4d92-9613-5e20f88c00be","trusted":true},"cell_type":"markdown","source":"## Importing required modules","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"import os \nimport torch\nimport torchvision\nimport torch.nn as nn\nfrom tqdm.notebook import tqdm\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.utils.data import DataLoader\nfrom torchvision.utils import make_grid\nfrom torchvision.datasets import ImageFolder","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Preparing the data","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"TRAIN_DIR = \"../input/104-flowers-garden-of-eden/jpeg-224x224/train\"\nVAL_DIR = \"../input/104-flowers-garden-of-eden/jpeg-224x224/val\"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"transform_train = T.Compose([\n    T.RandomCrop(128, padding_mode=\"reflect\"),\n    T.RandomHorizontalFlip(),\n    T.ToTensor()\n])\ntrain_ds = ImageFolder(\n    root=TRAIN_DIR,\n    transform=transform_train\n)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"transform_val = T.Compose([\n    T.ToTensor()\n])\n\nval_ds = ImageFolder(\n    root=VAL_DIR,\n    transform=transform_val\n\n)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"batch_size=128","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"The data loader will allow to access the data in batches.","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"train_dl = DataLoader(train_ds, batch_size, shuffle=True, num_workers=3, pin_memory=True)\nval_dl = DataLoader(val_ds, batch_size, num_workers=3, pin_memory=True)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"This is a helper method to show a batch of images and make sure that everything is working.","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"def show_batch(train_dl):\n    for images,_ in train_dl:\n        fig, ax = plt.subplots(figsize=(8,8))\n        ax.set_xticks([]); ax.set_yticks([])\n        ax.imshow(make_grid(images[:32], nrow=8).permute(1,2,0))\n        break\n        \nshow_batch(train_dl)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Moving to the GPU","execution_count":null},{"metadata":{},"cell_type":"markdown","source":"    The following code will be used to make sure a GPU is being used","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_device():\n    if torch.cuda.is_available():\n        return torch.device(\"cuda\") #if the GPU is availble this method will return cuda.\n    else:\n        return torch.device(\"cpu\")\n    \ndef to_device(data, device): #in here we move the data to device of our choice, the GPU\n    if isinstance(data, (list,tuple)):\n        return [to_device(x, device) for x in data]\n    return data.to(device, non_blocking=True)\n\nclass DeviceDataLoader():\n    def __init__(self, dl, device):\n        self.dl = dl\n        self.device = device\n        \n    def __iter__(self):\n        for x in self.dl:\n            yield to_device(x, self.device)\n            \n    def __len__(self):\n        return len(self.dl)\n    \ndevice = get_device()\ndevice","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## The model","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"def accuracy(out, labels):\n    _, preds = torch.max(out, dim=1)\n    return torch.tensor(torch.sum(preds == labels).item() / len(preds))\n\nclass ImageClassificationBase(nn.Module):\n    def training_step(self, batch):\n        images, labels = batch\n        out =self(images)\n        loss = F.cross_entropy(out, labels)\n        return loss\n    \n    def validation_step(self, batch):\n        images, labels = batch\n        out = self(images)\n        loss = F.cross_entropy(out, labels)\n        acc = accuracy(out, labels)\n        return {\"val_loss\": loss.detach(), \"val_acc\": acc}\n    \n    def validation_epoch_end(self, outputs):\n        batch_loss = [x[\"val_loss\"] for x in outputs]\n        epoch_loss = torch.stack(batch_loss).mean()\n        batch_acc = [x[\"val_acc\"] for x in outputs]\n        epoch_acc = torch.stack(batch_acc).mean()\n        return {\"val_loss\": epoch_loss.item(), \"val_acc\": epoch_acc.item()}\n    \n    def epoch_end(self, epoch, epochs, result):\n        print(\"Epoch: [{}/{}], last_lr: {:.4f}, train_loss: {:.4f}, val_loss: {:.4f}, val_acc: {:.4f}\".format(\n        epoch+1, epochs, result[\"lrs\"][-1], result[\"train_loss\"], result[\"val_loss\"], result[\"val_acc\"]))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class ResNet(ImageClassificationBase):\n    def __init__(self):\n        super().__init__()\n        self.network = models.resnet34(pretrained=True)\n        number_of_features = self.network.fc.in_features\n        self.network.fc = nn.Linear(number_of_features, 104)\n        \n    def forward(self, xb):\n        return self.network(xb)\n    \n    def freeze(self): #by freezing all the layers but the last one we allow it to warm up (the others are already good at training)\n        for param in self.network.parameters():\n            param.require_grad=False\n        for param in self.network.fc.parameters():\n            param.require_grad=True\n            \n    def unfreeze(self):\n        for param in self.network.parameters():\n            param.require_grad=True","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = ResNet()\nmodel","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = to_device(model, device) #let's move the model to the GPU","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_dl = DeviceDataLoader(train_dl, device)\nval_dl = DeviceDataLoader(val_dl, device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"@torch.no_grad()\ndef evaluate(model, val_dl):\n    model.eval()\n    outputs = [model.validation_step(batch) for batch in val_dl]\n    return model.validation_epoch_end(outputs)\n\ndef get_lr(optimizer):\n    for param_group in optimizer.param_groups:\n        return param_group[\"lr\"]\n    \ndef fit_one_cycle(epochs, max_lr, model, train_dl, val_dl, weight_decay=0,\n                 grad_clip=None, opt_func=torch.optim.Adam):\n    torch.cuda.empty_cache()\n    \n    history = []\n    opt = opt_func(model.parameters(), max_lr, weight_decay=weight_decay)\n    sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr, epochs=epochs,\n                                               steps_per_epoch=len(train_dl))\n    \n    for epoch in range(epochs):\n        model.train()\n        train_loss = []\n        lrs = []\n        for batch in tqdm(train_dl):\n            loss = model.training_step(batch)\n            train_loss.append(loss)\n            loss.backward()\n            \n            if grad_clip:\n                nn.utils.clip_grad_value_(model.parameters(), grad_clip)\n                \n            opt.step()\n            opt.zero_grad()\n            \n            lrs.append(get_lr(opt))\n            sched.step()\n            \n        result = evaluate(model, val_dl)\n        result[\"train_loss\"] = torch.stack(train_loss).mean().item()\n        result[\"lrs\"] = lrs\n        model.epoch_end(epoch, epochs, result)\n        history.append(result)\n    return history","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"result = evaluate(model, val_dl) #let's check the model performance before training it\nresult","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model.freeze()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"epochs = 10\nmax_lr = 10e-4\ngrad_clip = 0.1\nweight_decay = 1e-4\nopt_func = torch.optim.Adam","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"%%time\n\nhistory = fit_one_cycle(epochs, max_lr, model, train_dl, val_dl,\n                       weight_decay=weight_decay, grad_clip=grad_clip,\n                       opt_func=opt_func)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model.unfreeze()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"epochs = 10\nmax_lr = 0.0005\ngrad_clip = 0.1\nweight_decay = 1e-4\nopt_func = torch.optim.Adam","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"%%time\n\nhistory = fit_one_cycle(epochs, max_lr, model, train_dl, val_dl,\n                       weight_decay=weight_decay, grad_clip=grad_clip,\n                       opt_func=opt_func)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Model performance","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"val_loss = [x[\"val_loss\"] for x in history]\ntrain_loss = [x.get(\"train_loss\") for x in history]\nplt.plot(val_loss, \"-rx\")\nplt.plot(train_loss, \"-gx\")\nplt.title(\"Loss vs number of epochs\")\nplt.legend([\"Validation loss\", \"Train loss\"])\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Loss\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"accuracy = [x[\"val_acc\"] for x in history]\nplt.plot(accuracy, \"-bx\")\nplt.title(\"Acccuracy vs number of epochs\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Accuracy\")","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Predictions","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"import pandas as pd\nfrom torch.utils.data import Dataset\nfrom torchvision.datasets.folder import default_loader","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df = pd.read_csv(\"../input/flowers/flowers\")\ndf.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class TestData(Dataset):\n    def __init__(self, root_dir, csv_file, transform=None):\n        self.root_dir = root_dir\n        self.label = pd.read_csv(csv_file)\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.label)\n    \n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.item()\n        \n        label = self.label.iloc[idx,1]\n        image_path = os.path.join(self.root_dir, f\"{self.label.iloc[idx,0]}.jpeg\")\n        \n        image = default_loader(image_path)\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        return (image, label)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"TEST_DIR = \"../input/104-flowers-garden-of-eden/jpeg-224x224/test\"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"transform_test = T.Compose([\n    T.ToTensor()\n])\n\ntest_ds = TestData(\n    root_dir=TEST_DIR,\n    csv_file=\"../input/flowers/flowers\",\n    transform=transform_test\n\n)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def predict_image(image):\n    xb = to_device(image.unsqueeze(0), device)\n    out = model(xb)\n    _, preds = torch.max(out, dim=1)\n    prediction = preds[0].item()\n    return prediction","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"image, label = test_ds[20]\nprint(\"Label:\", label)\nprint(\"Prediction:\", predict_image(image))\nplt.imshow(image.permute(1,2,0))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","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}