{"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":"!python -m pip install --no-index --find-links=/kaggle/input/omegaconf222py3 omegaconf","metadata":{"execution":{"iopub.status.busy":"2023-08-03T05:14:44.004626Z","iopub.execute_input":"2023-08-03T05:14:44.004997Z","iopub.status.idle":"2023-08-03T05:14:55.107829Z","shell.execute_reply.started":"2023-08-03T05:14:44.004966Z","shell.execute_reply":"2023-08-03T05:14:55.106577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python -m pip install --no-index --find-links=/kaggle/input/hydracore120py3 hydra_core","metadata":{"execution":{"iopub.status.busy":"2023-08-03T05:15:03.960375Z","iopub.execute_input":"2023-08-03T05:15:03.961167Z","iopub.status.idle":"2023-08-03T05:15:15.221215Z","shell.execute_reply.started":"2023-08-03T05:15:03.961128Z","shell.execute_reply":"2023-08-03T05:15:15.219995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp -r ../input/cassava-compe-code20230801ver2/cassava-competition ../working","metadata":{"execution":{"iopub.status.busy":"2023-08-03T05:15:15.223905Z","iopub.execute_input":"2023-08-03T05:15:15.224312Z","iopub.status.idle":"2023-08-03T05:15:16.191899Z","shell.execute_reply.started":"2023-08-03T05:15:15.224275Z","shell.execute_reply":"2023-08-03T05:15:16.190777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cd ../working/cassava-competition","metadata":{"execution":{"iopub.status.busy":"2023-08-03T05:14:58.101449Z","iopub.execute_input":"2023-08-03T05:14:58.102318Z","iopub.status.idle":"2023-08-03T05:14:58.109140Z","shell.execute_reply.started":"2023-08-03T05:14:58.102277Z","shell.execute_reply":"2023-08-03T05:14:58.108261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !python train.py","metadata":{"execution":{"iopub.status.busy":"2023-08-03T01:54:31.891339Z","iopub.execute_input":"2023-08-03T01:54:31.891696Z","iopub.status.idle":"2023-08-03T04:05:35.217659Z","shell.execute_reply.started":"2023-08-03T01:54:31.891667Z","shell.execute_reply":"2023-08-03T04:05:35.216500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from lightning_data_module import CassavaDataModule\nfrom lightning_module import CassvaImgClassifier\nfrom pytorch_lightning import seed_everything\nimport pandas as pd\nimport pytorch_lightning as pl\nimport numpy as np\nfrom sklearn.metrics import log_loss\nimport torch\n\nimport os\n\ndef main():\n    seed_everything(719)\n    train = pd.read_csv(\"/kaggle/input/cassava-leaf-disease-classification/train.csv\")\n    test = pd.DataFrame()\n    test['image_id'] = list(os.listdir(\"/kaggle/input/cassava-leaf-disease-classification/test_images/\"))\n\n    test_data_module = CassavaDataModule(\n        train_df=train,\n        test_df = test,\n        train_bs = 16,\n        valid_bs = 32,\n        num_workers = 2,\n        fold_num = 5,\n        seed = 719,\n        train_data_root = \"/kaggle/input/cassava-leaf-disease-classification/train_images/\",\n        test_data_root = \"/kaggle/input/cassava-leaf-disease-classification/test_images/\",\n    )\n\n    trainer = pl.Trainer(\n        precision=16,\n        accelerator=\"gpu\",\n    )\n\n    val_preds = []\n    tst_preds = []\n\n    for i, epoch in enumerate([6, 7, 8, 9]): \n        model = CassvaImgClassifier.load_from_checkpoint(\n            '/kaggle/input/checkpoints/tf_efficientnet_b4_ns_fold_0_epoch{}.ckpt'.format(epoch),\n            model_arch = \"tf_efficientnet_b4_ns\", \n            n_class = train.label.nunique(), \n            learning_rate = 1e-4,\n            T_0 = 10,\n            min_lr = 1e-6,\n            weights_path = \"weights/tf_efficientnet_b4_ns.pth\",\n        )\n        # 推論時にデータ拡張をfor文で行う（TTA）\n        for _ in range(3):\n\n            tst_output = trainer.predict(\n                model=model, \n                datamodule=test_data_module\n            )\n            tst_output = np.concatenate(tst_output, axis=0)\n            tst_preds += [[1, 1, 1, 1][i] / sum([1, 1, 1, 1]) / 3 * tst_output]\n\n    tst_preds = np.mean(tst_preds, axis=0)\n\n    test['label'] = np.argmax(tst_preds, axis=2).reshape(-1)\n    test.to_csv('submission.csv', index=False)\n\nif __name__ == '__main__':\n    main()","metadata":{"execution":{"iopub.status.busy":"2023-08-03T05:16:53.496662Z","iopub.execute_input":"2023-08-03T05:16:53.497041Z","iopub.status.idle":"2023-08-03T05:17:16.781883Z","shell.execute_reply.started":"2023-08-03T05:16:53.497010Z","shell.execute_reply":"2023-08-03T05:17:16.780829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp submission.csv ../","metadata":{"execution":{"iopub.status.busy":"2023-08-03T05:17:35.937646Z","iopub.execute_input":"2023-08-03T05:17:35.938051Z","iopub.status.idle":"2023-08-03T05:17:36.956347Z","shell.execute_reply.started":"2023-08-03T05:17:35.938017Z","shell.execute_reply":"2023-08-03T05:17:36.954969Z"},"trusted":true},"execution_count":null,"outputs":[]}]}