{"cells":[{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"import numpy as np\nimport cv2\nimport gc\nimport random\nimport torch\nimport os\nimport pandas as pd\nfrom torch import optim\nfrom torch import nn\nimport albumentations as A\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torchvision.models import ResNet, Bottleneck\nimport time\nfrom tqdm import tqdm","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = ResNet(Bottleneck, [3, 4, 6, 3], num_classes=5, groups=32, width_per_group=4)\nmodel.load_state_dict(?)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"image_size = 380","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"config = dict(\n    # random seed\n    seed=22,\n    # name of the current approach\n    experiment_name='modified',\n    # paths to various things\n    test_location='../input/cassava-leaf-disease-classification/test_images',\n    checkpoint_path='../input/saved-models',\n    checkpoint='baselineepoch20.pt',\n    # model config\n    model='efficientnet-b4',\n    # number of epochs\n    epochs=10,\n    batch_size=16,\n    # number of worker processes\n    workers=8,\n    # transform config\n    inference_augmentations=[\n        dict(\n            name='RandomResizedCrop',\n            params=dict(\n                height=image_size,\n                width=image_size\n            )\n        ),\n        dict(\n            name='HorizontalFlip',\n            params=dict(\n                always_apply=False,\n                p=0.5,\n            )\n        ),\n        dict(\n            name='Transpose',\n            params=dict(\n                p=0.5\n            )\n        ),\n        dict(\n            name='VerticalFlip',\n            params=dict(\n                always_apply=False,\n                p=0.5\n                            )\n        ),\n        dict(\n            name='HueSaturationValue',\n            params=dict(\n                hue_shift_limit=0.2,\n                sat_shift_limit=0.2,\n                val_shift_limit=0.2,\n                p=0.5\n            )\n        ),\n        dict(\n            name='RandomBrightnessContrast',\n            params=dict(\n                brightness_limit=(-0.1, 0.1),\n                contrast_limit=(-0.1, 0.1),\n                p=0.5\n            )\n        ),\n        dict(\n            name='Normalize',\n            params=dict(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n                max_pixel_value=255.0,\n                p=1.0\n            )\n        ),\n        dict(\n            name='CoarseDropout',\n            params=dict(\n                p=0.5\n            )\n        ),\n        dict(\n            name='Cutout',\n            params=dict(\n                p=0.5\n            )\n        )\n    ]\n)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test = pd.DataFrame()\ntest['image_id'] = list(os.listdir(config['test_location']))\ntest.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# seed everything!\ndef seed(seed=22):\n    random.seed(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 = True\n    os.environ['PYTHONHASHSEED'] = str(seed)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"seed(config['seed'])\nassert torch.cuda.is_available(), 'cuda not available'\ndevice = torch.device('cuda')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# load model from checkpoint\ndef load_model():\n    model = EfficientNet.from_name(config['model'], num_classes=5)\n    checkpoint = torch.load(os.path.join(config['checkpoint_path'], config['checkpoint']))\n    if 'model' in checkpoint:\n        model.load_state_dict(checkpoint['model'])\n    else:\n        model.load_state_dict(checkpoint)\n    if 'epoch' in checkpoint:\n        epoch = int(checkpoint['epoch'])\n    if 'train_loss' in checkpoint:\n        train_loss = checkpoint['train_loss']\n    if 'val_loss' in checkpoint:\n        val_loss = checkpoint['val_loss']\n    if 'metrics' in checkpoint:\n        metrics = checkpoint['metrics']\n    if 'lr' in checkpoint:\n        lr = checkpoint['lr']\n\n    print('Loading model from checkpoint...')\n    print('Epoch', epoch)\n    print('Train loss', train_loss)\n    print('Validation loss', val_loss)\n    print('Accuracy', metrics)\n    print('Learning rate', lr)\n\n    return model.to(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Get train and test augmentations from the albumentations module\ndef get_transforms():\n    transforms = [getattr(A, item['name'])(**item['params']) for item in config['inference_augmentations']]\n    comp = A.Compose(transforms)\n    return comp","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaDataset(Dataset):\n    def __init__(self, images, transforms):\n        self.images = images\n        self.transforms = transforms\n\n    def __getitem__(self, n):\n        image = cv2.imread(os.path.join(config['test_location'], self.images[n]))\n        image = self.transforms(image=image)['image']\n        image = np.moveaxis(image, -1, 0)\n        image = torch.FloatTensor(image)\n        return image\n\n    def __len__(self):\n        return len(self.images)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Create a Torch DataLoader with our data\ndef get_dataloader():\n    transforms = get_transforms()\n    test_data = np.array(test['image_id'])\n\n    data = CassavaDataset(test_data, transforms)\n    dataloader = DataLoader(data, shuffle=False, batch_size=config['batch_size'], pin_memory=False, num_workers=config['workers'])\n\n    return dataloader","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# test loop\ndef infer(model, dataloader):\n    print('Running inference...')\n    model.eval()\n    predictions = []\n\n    with torch.no_grad():\n        for batch in tqdm(dataloader):\n            batch = batch.squeeze(1).to(device)\n            batch_hat = model(batch)\n            predictions.append(batch_hat)\n\n    return torch.cat(predictions, dim=0)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"if __name__ == '__main__':\n    torch.cuda.empty_cache()\n\n    dataloader = get_dataloader()\n    model = load_model()\n    predictions = None\n    print('Inferring experiment', config['experiment_name'])\n\n    # main training loop\n    for epoch in range(config['epochs']):\n        print('Epoch:', epoch)\n        start_time = time.time()\n\n        if epoch == 0:\n            predictions = infer(model, dataloader)\n        else:\n            predictions += infer(model, dataloader)\n\n        print('Time:', time.time() - start_time)\n        torch.cuda.empty_cache()\n        gc.collect()\n\n    # average across epochs\n    predictions /= config['epochs']\n    results = predictions.cpu().numpy()\n\n    # create submission file\n    test['label'] = np.argmax(results, axis=-1)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test.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}