{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Import necessary modules"},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nfrom torch import nn\nfrom torch import Tensor\nfrom torchvision.transforms import Compose, Resize, ToTensor\nfrom torchvision import datasets, transforms, models\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.utils import make_grid\n\nimport os\nimport time\nimport pandas as pd\nimport numpy as np\nfrom PIL import Image, ImageEnhance\nimport matplotlib.pyplot as plt","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Install timm to access ViT PyTorch Models"},{"metadata":{"trusted":true},"cell_type":"code","source":"pip install '../input/timm034/timm-0.3.4-py3-none-any.whl'","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Import timm"},{"metadata":{"trusted":true},"cell_type":"code","source":"import timm","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Display Available Vision Transformer Models"},{"metadata":{"trusted":true},"cell_type":"code","source":"print(\"Available ViT Models: \")\ntimm.list_models(\"vit*\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data_path  = '../input/cassava-leaf-disease-classification/'\ntrain_path = '../input/cassava-leaf-disease-classification/train_images/'\ntest_path  = '../input/cassava-leaf-disease-classification/test_images/'\nmodel_path = '../input/vitbase16224/jx_vit_base_p16_224-80ecf9dd.pth'\nCassava_model = '../input/cassavanewaugtp95epochs3/CassavaViT_newaug_TP95_Epochs3_LR1-75e05.pt'","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Define ViTBase 16 class"},{"metadata":{"trusted":true},"cell_type":"code","source":"class ViTBase16(nn.Module):\n    def __init__(self, n_classes, pretrained=False):\n        \n        super(ViTBase16, self).__init__()\n        \n        self.model = timm.create_model(\"vit_base_patch16_224\", pretrained=False)\n        \n        if pretrained:\n            self.model.load_state_dict(torch.load(model_path))\n            \n        self.model.head = nn.Linear(self.model.head.in_features, n_classes)\n    \n    def forward(self, x):\n        x = self.model(x)\n        return x","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Create cassava pre-trained ViTBase16 model instance (5 classes)"},{"metadata":{"trusted":true},"cell_type":"code","source":"cassava_model = ViTBase16(n_classes=5, pretrained=True)\ncassava_model.load_state_dict(torch.load(Cassava_model))\ncassava_gpu_model = cassava_model.cuda()\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(cassava_gpu_model.parameters(), lr=1.5e-05)\ncassava_gpu_model","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Examine test image folder"},{"metadata":{"trusted":true},"cell_type":"code","source":"test_img_names = []\nfor folder, subfolders, filenames in os.walk(test_path):\n    print('************************************************FOLDER************************************************')\n    print(folder)\n    print('************************************************IMAGES************************************************')\n    for img in filenames:\n        print(img)\n        if img[-3:] == 'jpg':\n            test_img_names.append(img)        \nprint('Testing Images: ',len(test_img_names))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Create class to load test data (adds image names)"},{"metadata":{"trusted":true},"cell_type":"code","source":"class TestSet2(Dataset):\n    \"\"\"Cassava Disease Dataset\"\"\"\n\n    def __init__(self, root_dir, test_dir, transform=None):\n        \"\"\"\n        Args:\n            csv_file (string): Path to the csv file with test image names.\n            root_dir (string): Directory with all the test images.\n            transform (callable, optional): Optional transform to be applied\n            on a sample.\n        \"\"\"\n        super().__init__()\n        self.root_dir = root_dir\n        self.test_dir = test_dir\n        self.transform = transform\n        print(root_dir)\n        print(test_dir)\n        print(\"Cassava Disease Test Dataset Length = \", len(os.listdir(self.test_dir)))\n\n\n    def __len__(self):\n        return len(os.listdir(self.test_dir))\n\n\n    def __getitem__(self, idx):\n        img_path = self.test_dir + os.listdir(self.test_dir)[idx] \n        img = Image.open(img_path).convert(\"RGB\")\n        image_name = os.listdir(self.test_dir)[idx]\n\n        \n        if self.transform:\n            image = self.transform(img)\n    \n        return (image, image_name)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Create test transforms (needed for image resizing)"},{"metadata":{"trusted":true},"cell_type":"code","source":"test_transform = transforms.Compose([\n        transforms.Resize((224, 224)),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485, 0.456, 0.406],\n                             [0.229, 0.224, 0.225])\n    ])","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Create test dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"testset = TestSet2(root_dir = '' , test_dir=test_path, transform = test_transform)\nprint(testset)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Load the test dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"test_batch_size = 1\ntest_loader = DataLoader(dataset=testset, batch_size=test_batch_size, shuffle=True, pin_memory=True)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Check to see if test_loader is working"},{"metadata":{"trusted":true},"cell_type":"code","source":"print(test_loader)\n\n# Grab the first batch of 16 images\nfor images,names in test_loader: \n    break\n\nim = make_grid(images, nrow=4)  # the default nrow is 8\n\n# Inverse normalize the images\n#inv_normalize = transforms.Normalize(\n#    mean=[-0.485/0.229, -0.456/0.224, -0.406/0.225],\n#    std=[1/0.229, 1/0.224, 1/0.225]\n#)\n#im_inv = inv_normalize(im)\nprint(names)\n# Print the images\nplt.figure(figsize=(12,4))\nplt.imshow(np.transpose(im.numpy(), (1, 2, 0)));","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Used trained model to generate test data predictions\n# Generate submission file"},{"metadata":{"trusted":true},"cell_type":"code","source":"test_losses = []\ntest_correct = []\ncol_names =  ['image_id', 'label']\nsubmission_df = pd.DataFrame(columns = col_names)\ntic = time.time()\nwith torch.no_grad():\n    for b, (X_test, name) in enumerate(test_loader):\n        X_test = X_test.cuda()\n        # Apply the model\n        y_test_pred = cassava_gpu_model(X_test)\n        predicted = torch.max(y_test_pred.data, 1)[1]\n        label = int(predicted)\n        image_name = name[0]\n        new_row = {'image_id': image_name, 'label': label}\n        submission_df = submission_df.append(new_row, ignore_index=True)\n        \ntoc = time.time() - tic\nprint('Time for model val is ', toc)\n#display(submission_df)\nsubmission_df.to_csv('submission.csv', index=False)\ndf_sub_test = pd.read_csv('submission.csv')\ndisplay(df_sub_test)\n!ls","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}