{"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_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! ls -l /kaggle/input/\n! ls -l /kaggle/input/plant-pathology-2021-fgvc8\n! ls -l /kaggle/input/plant-pathology-model-zoo","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-05-20T21:03:38.631055Z","iopub.execute_input":"2021-05-20T21:03:38.631421Z","iopub.status.idle":"2021-05-20T21:03:40.527099Z","shell.execute_reply.started":"2021-05-20T21:03:38.631342Z","shell.execute_reply":"2021-05-20T21:03:40.526215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip install -q --no-index /kaggle/input/plant-pathology-model-zoo/*.whl\n! pip install -q --no-index /kaggle/input/plant-pathology-model-zoo/kaggle_plant_pathology-0.4.1-*\n! pip list | grep torch\n! pip list | grep kornia\n! pip list | grep kaggle","metadata":{"execution":{"iopub.status.busy":"2021-05-20T21:03:43.437118Z","iopub.execute_input":"2021-05-20T21:03:43.43745Z","iopub.status.idle":"2021-05-20T21:04:56.331163Z","shell.execute_reply.started":"2021-05-20T21:03:43.437416Z","shell.execute_reply":"2021-05-20T21:04:56.330089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom kaggle_plantpatho.data import PlantPathologyDM\nfrom kaggle_plantpatho.models import LitResnet, MultiPlantPathology\n\ndm = PlantPathologyDM(\n    base_path='/kaggle/input/plant-pathology-2021-fgvc8',\n    path_csv='/kaggle/input/plant-pathology-2021-fgvc8/train.csv',\n    simple=False,\n    batch_size=32,\n)\ndm.setup()\n\nnet = torch.load('/kaggle/input/plant-pathology-model-zoo/fgvc8_resnet50.pt')\n# net = torch.load('/kaggle/input/plant-pathology-model-zoo/fgvc8_resnext101_32x8d.pt')\nmodel = MultiPlantPathology(model=net)\n# ckpt = torch.load('/kaggle/input/plant-pathology-model-zoo/resnet50-epoch38-valid_acc0.9678-valid_f10.9028.ckpt', map_location=torch.device('cpu'))\n# model.load_state_dict(ckpt['state_dict'])","metadata":{"execution":{"iopub.status.busy":"2021-05-20T21:11:09.208302Z","iopub.execute_input":"2021-05-20T21:11:09.208637Z","iopub.status.idle":"2021-05-20T21:11:10.173627Z","shell.execute_reply.started":"2021-05-20T21:11:09.208606Z","shell.execute_reply":"2021-05-20T21:11:10.172721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\npreds = []\nmodel.cuda().eval()\n\nfor imgs, names in dm.test_dataloader():\n    with torch.no_grad():\n        onehots = model(imgs.cuda()).cpu()\n    print(np.round(onehots.detach().numpy(), decimals=2))\n    for oh, name in zip(onehots, names):\n        print(dm.labels_unique[torch.argmax(oh)])\n        lbs = dm.onehot_to_labels(oh)\n        preds.append(dict(image=name, labels=\" \".join(lbs[::-1])))\n\ndf_preds = pd.DataFrame(preds)\nprint(df_preds)","metadata":{"execution":{"iopub.status.busy":"2021-05-20T21:11:11.521676Z","iopub.execute_input":"2021-05-20T21:11:11.522003Z","iopub.status.idle":"2021-05-20T21:11:12.799993Z","shell.execute_reply.started":"2021-05-20T21:11:11.521973Z","shell.execute_reply":"2021-05-20T21:11:12.798287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_preds['labels'] = [lbs or \"healthy\" for lbs in df_preds['labels']]\nprint(df_preds)\ndf_preds.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2021-05-20T21:11:17.842306Z","iopub.execute_input":"2021-05-20T21:11:17.842645Z","iopub.status.idle":"2021-05-20T21:11:17.853785Z","shell.execute_reply.started":"2021-05-20T21:11:17.842613Z","shell.execute_reply":"2021-05-20T21:11:17.85265Z"},"trusted":true},"execution_count":null,"outputs":[]}]}