{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"%reload_ext autoreload\n%autoreload 2\n%matplotlib inline","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"import math\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport os\nimport pandas as pd\nfrom PIL import Image\nimport random\nfrom sklearn.model_selection import train_test_split\nimport time\nfrom tqdm.notebook import tqdm\n\nimport torch\nfrom torch.utils.data.dataset import Dataset\nfrom torch.utils.data import DataLoader\nfrom torchvision import transforms, datasets\nfrom torch import nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.autograd import Variable\nimport torchvision.models as models\nfrom torchvision.utils import make_grid","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Seed"},{"metadata":{"trusted":true},"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"SEED = 17\nseed_everything(SEED)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Data Folder"},{"metadata":{"trusted":true},"cell_type":"code","source":"data_dir = '../input/cassava-leaf-disease-classification'\ntrain_dir = data_dir + '/train_images'\ntrain_csv = data_dir + '/train.csv'\ntest_dir = data_dir + '/test_images'\nname_json = data_dir + '/label_num_to_disease_map.json'\nsample_csv = data_dir + '/sample_submission.csv'","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Read CSV"},{"metadata":{"trusted":true},"cell_type":"code","source":"train_df = pd.read_csv(train_csv)\ntrain_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_df.label.value_counts()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"the dataset seems heavily unbalanced towards label 3."},{"metadata":{"trusted":true},"cell_type":"code","source":"sub_df = pd.read_csv(sample_csv)\nsub_df.head()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaDS(Dataset):\n    def __init__(self, df, data_dir, transforms=None):\n        super().__init__()\n        self.df_data = df.values\n        self.transforms = transforms\n        self.data_dir = data_dir\n\n    def __len__(self):\n        return len(self.df_data)\n\n    def __getitem__(self, index):\n        img_name, label = self.df_data[index]\n        img_path = os.path.join(self.data_dir, img_name)\n        img = Image.open(img_path).convert(\"RGB\")\n        if self.transforms is not None:\n            image = self.transforms(img)\n        return image, label","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"X_train, X_valid = train_test_split(train_df, test_size=0.1, \n                                                    random_state=SEED,\n                                                    stratify=train_df.label.values)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"X_train.shape, X_valid.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"normalize = transforms.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_tf = transforms.Compose([\n    transforms.Pad(4, padding_mode='reflect'),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomResizedCrop(224),\n    transforms.ToTensor(),\n    normalize\n])\n\nvalid_tf = transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(224),\n    transforms.ToTensor(),\n    normalize\n])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_ds = CassavaDS(X_train, train_dir, train_tf)\nvalid_ds = CassavaDS(X_valid, train_dir, valid_tf)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"bs = 64","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_loader = DataLoader(train_ds, batch_size=bs, shuffle=True)\nvalid_loader = DataLoader(valid_ds, batch_size=bs, shuffle=True)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Labels"},{"metadata":{"trusted":true},"cell_type":"code","source":"import json\n\nwith open(name_json, 'r') as f:\n    cat_to_name = json.load(f)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"cat_to_name","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Plot Images"},{"metadata":{"trusted":true},"cell_type":"code","source":"class UnNormalize(object):\n    def __init__(self, mean, std):\n        self.mean = mean\n        self.std = std\n\n    def __call__(self, tensor):\n        for t, m, s in zip(tensor, self.mean, self.std):\n            t.mul_(s).add_(m)\n        return tensor","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"unnorm = UnNormalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def display_img(img, label=None, unnorm_obj=None, invert=True, return_label=True):\n    if unnorm_obj != None:\n        img = unnorm_obj(img)\n\n    plt.imshow(img.permute(1, 2, 0))\n    \n    if label != None:\n        plt.title(cat_to_name[str(label)])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def display_batch(batch, unnorm_obj=None):    \n    imgs, labels = batch\n    \n    if unnorm_obj:\n        unnorm_imgs = []\n        for img in imgs:\n            unnorm_imgs.append(unnorm_obj(img))\n        imgs = unnorm_imgs\n    \n    ig, ax = plt.subplots(figsize=(16, 8))\n    ax.set_xticks([]); ax.set_yticks([])\n    ax.imshow(make_grid(imgs, nrow=16).permute(1, 2, 0))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"img, label = train_ds[0]\ndisplay_img(img, label)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"display_batch(next(iter(train_loader)))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Helper Functions"},{"metadata":{"trusted":true},"cell_type":"code","source":"class AvgStats(object):\n    def __init__(self):\n        self.reset()\n        \n    def reset(self):\n        self.losses =[]\n        self.precs =[]\n        self.its = []\n        \n    def append(self, loss, prec, it):\n        self.losses.append(loss)\n        self.precs.append(prec)\n        self.its.append(it)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def save_checkpoint(model, is_best, filename='./checkpoint.pth'):\n    \"\"\"Save checkpoint if a new best is achieved\"\"\"\n    if is_best:\n        torch.save(model.state_dict(), filename)  # save checkpoint\n    else:\n        print (\"=> Validation Accuracy did not improve\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def load_checkpoint(model, filename = './checkpoint.pth'):\n    sd = torch.load(filename, map_location=lambda storage, loc: storage)\n    names = set(model.state_dict().keys())\n    for n in list(sd.keys()): \n        if n not in names and n+'_raw' in names:\n            if n+'_raw' not in sd: sd[n+'_raw'] = sd[n]\n            del sd[n]\n    model.load_state_dict(sd)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Train and Test"},{"metadata":{"trusted":true},"cell_type":"code","source":"def train(loader, model, optimizer, device):\n    model.train()\n    correct, trn_loss, trn_time = 0., 0., 0\n    t = tqdm(loader, leave=False, total=len(loader))\n    bt_start = time.time()\n    for i, (ip, target) in enumerate(t):\n        ip, target = ip.to(device), target.to(device)                          \n        output = model(ip)\n        loss = criterion(output, target)\n        trn_loss += loss.item()\n        \n        # measure accuracy and record loss\n        _, pred = output.max(dim=1)\n        correct += torch.sum(pred == target.data)\n\n        # compute gradient and do SGD step\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n    \n    trn_time = time.time() - bt_start\n    trn_acc = correct * 100 / len(loader.dataset)\n    trn_loss /= len(loader)\n    return trn_acc, trn_loss, trn_time","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def valid(loader, model, optimizer, device):\n    model.eval()\n    with torch.no_grad():\n        correct, val_loss, val_time = 0., 0., 0\n        t = tqdm(loader, leave=False, total=len(loader))\n        bt_start = time.time()\n        for i, (ip, target) in enumerate(t):\n            ip, target = ip.to(device), target.to(device)                          \n            output = model(ip)\n            loss = criterion(output, target)\n            val_loss += loss.item()\n\n            # measure accuracy and record loss\n            _, pred = output.max(dim=1)\n            correct += torch.sum(pred == target.data)\n\n        val_time = time.time() - bt_start\n        val_acc = correct * 100 / len(loader.dataset)\n        val_loss /= len(loader)\n        return val_acc, val_loss, val_time","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def fit(model, sched, optimizer, device, epoch):\n    print(\"Epoch\\tTrn_loss\\tVal_loss\\tTrn_acc\\t\\tVal_acc\")\n    best_acc = 0.\n    for j in range(epoch):\n        trn_acc, trn_loss, trn_time = train(train_loader, model, optimizer, device)\n        trn_stat.append(trn_loss, trn_acc, trn_time)\n        val_acc, val_loss, val_time = valid(valid_loader, model, optimizer, device)\n        val_stat.append(val_acc, val_loss, val_time)\n        if sched:\n            sched.step()\n        if val_acc > best_acc:\n            best_acc = val_acc\n            save_checkpoint(model, True, './best_model.pth')\n        print(\"{}\\t{:06.8f}\\t{:06.8f}\\t{:06.8f}\\t{:06.8f}\"\n              .format(j+1, trn_loss, val_loss, trn_acc, val_acc))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Vision Transformer"},{"metadata":{},"cell_type":"markdown","source":"Using implementation from https://github.com/nachiket273/Vision_transformer_pytorch.git"},{"metadata":{},"cell_type":"markdown","source":"Using weights from pretrained-model on Imagenet-1k <br>\nFile is available at https://www.kaggle.com/nachiket273/visiontransformerpretrainedimagenet1kweights"},{"metadata":{},"cell_type":"markdown","source":"pytorch tpu kernel available @ https://www.kaggle.com/nachiket273/pytorch-tpu-vision-transformer"},{"metadata":{"trusted":true},"cell_type":"code","source":"!cp ../input/visiontransformerpretrainedimagenet1kweights/vit.py .\n!cp ../input/visiontransformerpretrainedimagenet1kweights/vit_16_224_imagenet1000.pth .","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from vit import ViT","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_model(out_features=5):\n    model = ViT(224, 16, drop_rate=0.1)\n    load_checkpoint(model, './vit_16_224_imagenet1000.pth')\n    model.out = nn.Linear(in_features=model.out.in_features, out_features=5)\n    for param in model.parameters():\n        param.require_grad = True\n    return model","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = get_model()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = model.to(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"trn_stat = AvgStats()\nval_stat = AvgStats()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"optimizer = torch.optim.SGD(model.parameters(), lr=1e-2, momentum=0.9, weight_decay=1e-4)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"epochs = 20","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sched = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, epochs, 1e-3)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"fit(model, sched, optimizer, device, epochs)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Predict"},{"metadata":{"trusted":true},"cell_type":"code","source":"test_tf = transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(224),\n    transforms.ToTensor(),\n    normalize\n])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"load_checkpoint(model, './best_model.pth')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def predict(test_dir, model, device):\n    img_names = []\n    preds = []\n    for name in os.listdir(test_dir):\n        img_path = os.path.join(test_dir, name)\n        img = Image.open(img_path).convert(\"RGB\")\n        img = test_tf(img)\n        img = img.unsqueeze(0)\n        img = img.to(device)\n        op = model(img)\n        _, pred = op.max(dim=1)\n        img_names.append(name)\n        preds.append(pred.item())\n    return img_names, preds","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"img_names, preds = predict(test_dir, model, device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"img_names, preds","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sub_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sub_df['image_id'] = img_names\nsub_df['label'] = preds","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sub_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sub_df.to_csv('submission.csv', index=False)","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}