{"cells":[{"metadata":{},"cell_type":"markdown","source":"GitHub: https://github.com/lucidrains/vit-pytorch"},{"metadata":{"trusted":true},"cell_type":"code","source":"import pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import StepLR\nimport torchvision\nfrom torchvision import datasets, transforms\nfrom PIL import Image\nfrom torch.utils.data import DataLoader, Dataset\nfrom sklearn.model_selection import train_test_split","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install vit-pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\nsub = pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Hyperparams"},{"metadata":{"trusted":true},"cell_type":"code","source":"NUM_CLASSES = train[\"label\"].nunique()\nTARGET_SIZE = (320, 320)\nIMG_SIZE = 320\nEPOCHS = 25\nVALIDATION_SPLIT = 0.15\nBATCH_SIZE = 32\nLR = 2e-5\nGAMMA = 0.7","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_transforms = transforms.Compose(\n    [\n        transforms.Resize(TARGET_SIZE),\n        transforms.RandomResizedCrop(IMG_SIZE),\n        transforms.RandomHorizontalFlip(),\n        transforms.ToTensor(),\n    ]\n)\n\nval_transforms = transforms.Compose(\n    [\n        transforms.Resize(TARGET_SIZE),\n        transforms.ToTensor(),\n    ]\n)\n\ntest_transforms = transforms.Compose(\n    [\n        transforms.Resize(TARGET_SIZE),\n        transforms.ToTensor(),\n    ]\n)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_list, valid_list = train_test_split(train, \n                                          test_size=0.15,\n                                          stratify=train[\"label\"],\n                                          random_state=2987346)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"\nclass MyDataset(torch.utils.data.Dataset):\n    def __init__(self, dataframe, transform = None):\n        self.dataframe = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n        row = self.dataframe.iloc[index]\n        img = (Image.open(\"../input/cassava-leaf-disease-classification/train_images/\" + row[\"image_id\"]))\n        label = row[\"label\"]\n        if not self.transform:\n            return torchvision.transforms.functional.to_tensor(img), label\n        img_transformed = self.transform(img)\n        return img_transformed, label\n\n\ntrain_dataset = MyDataset(train_list, train_transforms)\nval_dataset = MyDataset(valid_list, val_transforms)\n\n\nclass TestDataset(torch.utils.data.Dataset):\n    def __init__(self, dataframe, transform = None):\n        self.dataframe = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n        row = self.dataframe.iloc[index]\n        img = (Image.open(\"../input/cassava-leaf-disease-classification/test_images/\" + row[\"image_id\"]))\n        if not self.transform:\n            return torchvision.transforms.functional.to_tensor(img), label\n        img_transformed = self.transform(img)\n        return img_transformed\n    \ntest_dataset = TestDataset(sub, test_transforms)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_loader = DataLoader(dataset = train_dataset, batch_size=BATCH_SIZE, shuffle=True)\nval_loader = DataLoader(dataset = val_dataset, batch_size=BATCH_SIZE, shuffle=True)\ntest_loader = DataLoader(dataset = test_dataset, batch_size=BATCH_SIZE, shuffle=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print(len(train_dataset), len(train_loader))\nprint(len(val_dataset), len(val_loader))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device = ('cuda' if torch.cuda.is_available() else 'cpu')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from vit_pytorch import ViT\n\n\nmodel = ViT(\n    image_size = IMG_SIZE,\n    patch_size = 16,\n    num_classes = NUM_CLASSES,\n    dim = 1024,\n    depth = 6,\n    heads = 8,\n    mlp_dim = 2048,\n    dropout = 0.1,\n    emb_dropout = 0.1\n).to(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# loss function\ncriterion = nn.CrossEntropyLoss()\n# optimizer\noptimizer = optim.Adam(model.parameters(), lr=LR)\n# scheduler\nscheduler = StepLR(optimizer, step_size=1, gamma=GAMMA)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from tqdm.notebook import tqdm","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Training"},{"metadata":{"trusted":true},"cell_type":"code","source":"for epoch in range(EPOCHS):\n    epoch_loss = 0\n    epoch_accuracy = 0\n\n    for data, label in tqdm(train_loader):\n        data = data.to(device)\n        label = label.to(device)\n\n        output = model(data)\n        loss = criterion(output, label)\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        acc = (output.argmax(dim=1) == label).float().mean()\n        epoch_accuracy += acc / len(train_loader)\n        epoch_loss += loss / len(train_loader)\n\n    with torch.no_grad():\n        epoch_val_accuracy = 0\n        epoch_val_loss = 0\n        for data, label in val_loader:\n            data = data.to(device)\n            label = label.to(device)\n\n            val_output = model(data)\n            val_loss = criterion(val_output, label)\n\n            acc = (val_output.argmax(dim=1) == label).float().mean()\n            epoch_val_accuracy += acc / len(val_loader)\n            epoch_val_loss += val_loss / len(val_loader)\n\n    print(\n        f\"Epoch : {epoch+1} - loss : {epoch_loss:.4f} - acc: {epoch_accuracy:.4f} - val_loss : {epoch_val_loss:.4f} - val_acc: {epoch_val_accuracy:.4f}\\n\"\n    )","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Make predictions"},{"metadata":{"trusted":true},"cell_type":"code","source":"import numpy as np\n\nmodel.eval()\npreds = []\n\nfor data in test_loader:\n    inputs = data.to(device)\n\n    with torch.no_grad():\n        outputs = model(inputs)\n\n    preds.append(outputs.sigmoid().detach().cpu().numpy())\n\npreds = np.concatenate(preds)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"predictions = np.argmax(preds, axis=1)\npredictions","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sub[\"label\"] = predictions\nsub.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}