{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9550977,"sourceType":"datasetVersion","datasetId":5415880},{"sourceId":9551266,"sourceType":"datasetVersion","datasetId":5462370}],"dockerImageVersionId":30746,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\nIN_LOCAL= not os.path.exists(\"/kaggle/input\")\nif IN_LOCAL:\n    INPUT_DIR = '.'\nelse:\n    INPUT_DIR = \"/kaggle/input\"\nif IN_LOCAL:\n    sys.path.append(INPUT_DIR + '/src')\nelse:\n    sys.path.append(INPUT_DIR)\n    sys.path.append(INPUT_DIR + '/rsna2024src')\nimport tqdm\nimport torch.nn as nn\n\nfrom rsna.configs import *\nimport mylibs.training_kits_fn as mykits\nmykits.seed_everything(config.global_seed)\n\nimport mylibs.common_models  as mymds\nimport rsna.models as ms\nimport rsna.datasets as ds\nimport gc","metadata":{"execution":{"iopub.status.busy":"2024-09-12T01:33:28.019382Z","iopub.execute_input":"2024-09-12T01:33:28.019646Z","iopub.status.idle":"2024-09-12T01:33:35.312851Z","shell.execute_reply.started":"2024-09-12T01:33:28.019618Z","shell.execute_reply":"2024-09-12T01:33:35.311880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class data_cfg(Stage2Model0Cfg):\n    training = False\n\nclass model0_test_cfg(Stage2Model0Cfg):\n    training = False\n    \nclass model1_test_cfg(Stage2Model1Cfg):\n    training = False\n\nclass sagittal_loc_cfg(sagittal_loc_cfg):\n    training = False   \n\nclass axial_loc_cfg(axial_loc_cfg):\n    training = False   \n   ","metadata":{"execution":{"iopub.status.busy":"2024-09-12T01:33:35.314573Z","iopub.execute_input":"2024-09-12T01:33:35.314885Z","iopub.status.idle":"2024-09-12T01:33:35.319991Z","shell.execute_reply.started":"2024-09-12T01:33:35.314860Z","shell.execute_reply":"2024-09-12T01:33:35.318879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_LOCAL:\n    import  rsna.train_datasets as tds\n    import rsna.train_utils as tuts\n    # dataset = tds.RsnaDatasetT(cfg = test_cfg)\n    dataset = ds.RsnaDataset(cfg = data_cfg,test=False)\n    # (train_data,dataset),_ = tuts.split_dataset(test_cfg,0,no_imgs=False)\n    # dataset = tds.RsnaDatasetT2(cfg = test_cfg,fold = 0,training=False)\n    dataset = torch.utils.data.Subset(dataset,range(len(dataset))[0:10])\n    # dataset = torch.utils.data.Subset(dataset,range(len(dataset))[1131:1142])\n    score_fun = tuts.create_loss_fun(data_cfg)\n    # lw_fun = tuts.create_level_weighted_loss_fun(test_cfg)\nelse:\n    dataset = ds.RsnaDataset(cfg = data_cfg,test=not IN_LOCAL)\ndataloader = torch.utils.data.DataLoader(dataset,\n                                         batch_size=1, \n                                         shuffle=False, \n                                         collate_fn = ds.unbatch_collate_fn,\n                                         num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T01:33:35.391425Z","iopub.execute_input":"2024-09-12T01:33:35.392036Z","iopub.status.idle":"2024-09-12T01:33:35.415965Z","shell.execute_reply.started":"2024-09-12T01:33:35.392013Z","shell.execute_reply":"2024-09-12T01:33:35.415274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sloc_model = mymds.EnsembleModel([ms.SagittalLocModel(sagittal_loc_cfg),\n                                     ms.SagittalLocModel(sagittal_loc_cfg),\n                                     ms.SagittalLocModel(sagittal_loc_cfg),\n                                     ms.SagittalLocModel(sagittal_loc_cfg),\n                                     ms.SagittalLocModel(sagittal_loc_cfg),\n                                    ])\nsloc_model.load(config.PRETRAINED + '/' +sloc_model.get_name() + '.pt')\nsloc_model.to(data_cfg.device)\nsloc_model.eval()\n\naloc_model = mymds.EnsembleModel([ms.AxialLocModel(axial_loc_cfg),\n                                     ms.AxialLocModel(axial_loc_cfg),\n                                     ms.AxialLocModel(axial_loc_cfg),\n                                     ms.AxialLocModel(axial_loc_cfg),\n                                     ms.AxialLocModel(axial_loc_cfg),\n                                    ])\naloc_model.load(config.PRETRAINED + '/' +aloc_model.get_name() + '.pt')\naloc_model.to(data_cfg.device)\naloc_model.eval()\n\nsum(p.numel() for p in sloc_model.parameters() if p.requires_grad)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T01:33:35.416961Z","iopub.execute_input":"2024-09-12T01:33:35.417224Z","iopub.status.idle":"2024-09-12T01:33:41.904539Z","shell.execute_reply.started":"2024-09-12T01:33:35.417201Z","shell.execute_reply":"2024-09-12T01:33:41.903564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def ag_fun(x):\n    x = torch.stack(x)\n    x = x.softmax(-1)\n    x = torch.mean(x, dim=0)\n    return x\n\nmodel0 = mymds.EnsembleModel([\n    ms.RsnaModel0Plus(model0_test_cfg),\n    ms.RsnaModel0Plus(model0_test_cfg),\n    ms.RsnaModel0Plus(model0_test_cfg),\n    ms.RsnaModel0Plus(model0_test_cfg),\n    ms.RsnaModel0Plus(model0_test_cfg),\n    ],\n    ag_fun = ag_fun,\n    )\nmodel0.load(config.PRETRAINED + '/' +model0.get_name() + '.pt')\n\nmodel1 = mymds.EnsembleModel([\n    ms.RsnaModel1Plus(model1_test_cfg),\n    ms.RsnaModel1Plus(model1_test_cfg),\n    ms.RsnaModel1Plus(model1_test_cfg),\n    ms.RsnaModel1Plus(model1_test_cfg),\n    ms.RsnaModel1Plus(model1_test_cfg),\n    ],\n    ag_fun = ag_fun,\n    )\nmodel1.load(config.PRETRAINED + '/' +model1.get_name() + '.pt')\n\ntrained_model = mymds.EnsembleModel([\n#     * model0.models, \n    * model1.models\n],ag_fun = ag_fun)\n\ntrained_model.to(model1_test_cfg.device)\ntrained_model.eval()\nsum(p.numel() for p in trained_model.parameters() if p.requires_grad),len(trained_model.models)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T01:33:41.905650Z","iopub.execute_input":"2024-09-12T01:33:41.905929Z","iopub.status.idle":"2024-09-12T01:33:42.336796Z","shell.execute_reply.started":"2024-09-12T01:33:41.905906Z","shell.execute_reply":"2024-09-12T01:33:42.335864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nrow_ids = []\nprobs = []\nscores = []\nwith torch.autocast(device_type='cuda', dtype=torch.float16):\n    with torch.no_grad():\n        for row,d in enumerate(tqdm.tqdm(dataloader)):\n            ((simgs,simg_poses),(aimgs,aimgs_poses)) = ds.crop_imgs_by_loc(d,sloc_model,aloc_model,data_cfg)\n            d['s_imgs'] = torch.as_tensor(simgs)\n            d['a_imgs'] = torch.as_tensor(aimgs)\n            d['s_imgs_poses'] = simg_poses\n            d['a_imgs_poses'] = aimgs_poses\n            a_level_imgs_idx,_,t1_imgs_idxs,t2_imgs_idxs,_ = ds.choose_imgs(d)\n            d['a_imgs_idx'] = a_level_imgs_idx\n            d['t1_imgs_idx'] = t1_imgs_idxs\n            d['t2_imgs_idx'] = t2_imgs_idxs\n            torch.cuda.empty_cache()\n            gc.collect()\n            study_id = d['study_id']\n            prob = trained_model(d)\n            if IN_LOCAL:\n                import rsna.train_utils as tuts\n                if d.get('y_and_weight') is not None:\n                    # s1 = score_fun(prob,d['y_and_weight'].to(prob.device)).cpu().numpy()\n                    # s2 = tuts.score(d['y_and_weight'].unsqueeze(0).numpy(),prob.softmax(-1).cpu().numpy())\n                    # scores.append(score_fun(prob,d['y_and_weight'].to(prob.device)).cpu().numpy())\n                    scores.append(tuts.score(d['y_and_weight'].unsqueeze(0).numpy(),prob.cpu().numpy()))\n            prob = prob[0].cpu().numpy()\n            for i,c in enumerate(config.label_names):\n                row_ids.append(str(study_id)+'_'+c)\n                probs.append(prob[i])\n\n      \ndel dataloader\ndel sloc_model\ndel aloc_model\ndel trained_model\ntorch.cuda.empty_cache()\ngc.collect()      \n\nsub = pd.DataFrame()\nsub[config.sub_columns[0]] = row_ids\nsub[config.sub_columns[1:]] = probs\nsub.to_csv('submission.csv', index=False)\nif IN_LOCAL:\n    if len(scores) > 0:\n        print('score = ',np.mean(scores))\nsub.head(25)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T01:33:42.338231Z","iopub.execute_input":"2024-09-12T01:33:42.338671Z","iopub.status.idle":"2024-09-12T01:33:45.335011Z","shell.execute_reply.started":"2024-09-12T01:33:42.338635Z","shell.execute_reply":"2024-09-12T01:33:45.334043Z"},"trusted":true},"execution_count":null,"outputs":[]}]}