{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\n\nimport cassava_utils as utils\n\nimport torch\nfrom torch.utils.data import DataLoader\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"DEVICE = \"cuda\"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"folder_path = \"../input/cassava-leaf-disease-classification\"\nmodels_path = \"../input/resnet50-fold\"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# get the path of all models\nmodels_list = []\nfor model in os.listdir(models_path):\n    model = os.path.join(models_path, model)\n    models_list.append(model)\nmodels_list.sort()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# define albumentations for test data\ndata_albums = {\n    'test': A.Compose([\n        A.Resize(height=400, width=400),\n        A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n        ToTensorV2()])}","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def create_test_dataloader(root_dir, data_albums):\n    test_files = os.listdir(root_dir)\n    test_dataset = utils.CassavaTestDataset(test_files, root_dir, albums=data_albums[\"test\"])\n    test_dataloader = {\"test\": DataLoader(test_dataset, batch_size=16, num_workers=8)}\n    return test_dataloader","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_dataloader = create_test_dataloader(os.path.join(folder_path, \"test_images\"), data_albums)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# load all models from model_path\nmodels = {}\nfor fold, model_path in enumerate(models_list):\n    models[fold] = torch.load(model_path)\n    models[fold].eval()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_inference_df(dataloader, models):\n    test_eng = utils.Engine(model=None, optimizer=None, device=DEVICE)\n    inference_df = test_eng.ensemble_predict(dataloader[\"test\"], models)\n    return inference_df ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"inference_df = get_inference_df(test_dataloader, models)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"inference_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}