{"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":"import sys\nsys.path.append('../input/transmil-update187')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-29T05:05:27.056851Z","iopub.execute_input":"2022-09-29T05:05:27.05761Z","iopub.status.idle":"2022-09-29T05:05:27.081339Z","shell.execute_reply.started":"2022-09-29T05:05:27.057494Z","shell.execute_reply":"2022-09-29T05:05:27.080346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip install pytorch_toolbelt","metadata":{"execution":{"iopub.status.busy":"2022-09-29T05:05:27.083887Z","iopub.execute_input":"2022-09-29T05:05:27.084157Z","iopub.status.idle":"2022-09-29T05:05:40.171907Z","shell.execute_reply.started":"2022-09-29T05:05:27.084133Z","shell.execute_reply":"2022-09-29T05:05:40.170741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip install addict","metadata":{"execution":{"iopub.status.busy":"2022-09-29T05:05:40.174636Z","iopub.execute_input":"2022-09-29T05:05:40.175018Z","iopub.status.idle":"2022-09-29T05:05:50.244381Z","shell.execute_reply.started":"2022-09-29T05:05:40.174981Z","shell.execute_reply":"2022-09-29T05:05:50.242968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip install nystrom_attention","metadata":{"execution":{"iopub.status.busy":"2022-09-29T05:05:50.247039Z","iopub.execute_input":"2022-09-29T05:05:50.247709Z","iopub.status.idle":"2022-09-29T05:06:00.800871Z","shell.execute_reply.started":"2022-09-29T05:05:50.247668Z","shell.execute_reply":"2022-09-29T05:06:00.799669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import argparse\nfrom pathlib import Path\nimport numpy as np\nimport glob\n\nfrom transMIL_datasets import DataInterface\nfrom transMIL_models import ModelInterface\nfrom utils.utils import *\n\n# pytorch_lightning\nimport pytorch_lightning as pl\nfrom pytorch_lightning import Trainer\n\n#--->Setting parameters\ndef make_parse():\n    parser = argparse.ArgumentParser()\n    parser.add_argument('--stage', default='train', type=str)\n    parser.add_argument('--config', default='Camelyon/TransMIL.yaml',type=str)\n    parser.add_argument('--gpus', default = [2])\n    parser.add_argument('--fold', default = 0)\n    args = parser.parse_args()\n    return args\n\nclass parse_():\n    stage = 'train'\n    config = '../input/transmil-update187/Camelyon/TransMIL.yaml'\n    gpus = [2]\n    fold = 0\n\n#---->main\ndef main(cfg):\n\n    #---->Initialize seed\n    pl.seed_everything(cfg.General.seed)\n\n    #---->load loggers\n    cfg.load_loggers = load_loggers(cfg)\n\n    #---->load callbacks\n    cfg.callbacks = load_callbacks(cfg)\n\n    #---->Define Data \n    DataInterface_dict = {'train_batch_size': cfg.Data.train_dataloader.batch_size,\n                'train_num_workers': cfg.Data.train_dataloader.num_workers,\n                'test_batch_size': cfg.Data.test_dataloader.batch_size,\n                'test_num_workers': cfg.Data.test_dataloader.num_workers,\n                'dataset_name': cfg.Data.dataset_name,\n                'dataset_cfg': cfg.Data,}\n    dm = DataInterface(**DataInterface_dict)\n\n    #---->Define Model\n    ModelInterface_dict = {'model': cfg.Model,\n                            'loss': cfg.Loss,\n                            'optimizer': cfg.Optimizer,\n                            'data': cfg.Data,\n                            'log': cfg.log_path\n                            }\n    model = ModelInterface(**ModelInterface_dict)\n    \n    #---->Instantiate Trainer\n    trainer = Trainer(\n        num_sanity_val_steps=0, \n        logger=cfg.load_loggers,\n        callbacks=cfg.callbacks,\n        max_epochs= cfg.General.epochs,\n        gpus=0,\n#         amp_level=cfg.General.amp_level,\n        precision=cfg.General.precision,  \n        accumulate_grad_batches=cfg.General.grad_acc,\n        deterministic=True,\n        check_val_every_n_epoch=1,\n        # amp_backend='apex', \n        accelerator='gpu'\n    )\n\n    #---->train or test\n    if cfg.General.server == 'train':\n        trainer.fit(model = model, datamodule = dm)\n    else:\n        model_paths = list(cfg.log_path.glob('*.ckpt'))\n        model_paths = [str(model_path) for model_path in model_paths if 'epoch' in str(model_path)]\n        for path in model_paths:\n            print(path)\n            new_model = model.load_from_checkpoint(checkpoint_path=path, cfg=cfg)\n            trainer.test(model=new_model, datamodule=dm)\n\nif __name__ == '__main__':\n\n    args = parse_\n    cfg = read_yaml(args.config)\n\n    #---->update\n    cfg.config = args.config\n    cfg.General.gpus = args.gpus\n    cfg.General.server = args.stage\n    cfg.Data.fold = args.fold\n\n    #---->main\n    main(cfg)\n ","metadata":{"execution":{"iopub.status.busy":"2022-09-29T05:06:00.803807Z","iopub.execute_input":"2022-09-29T05:06:00.804239Z"},"trusted":true},"execution_count":null,"outputs":[]}]}