{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"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":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9114148,"sourceType":"datasetVersion","datasetId":5501147},{"sourceId":9423242,"sourceType":"datasetVersion","datasetId":5705584},{"sourceId":9482835,"sourceType":"datasetVersion","datasetId":5762105},{"sourceId":9484766,"sourceType":"datasetVersion","datasetId":5723776},{"sourceId":9516581,"sourceType":"datasetVersion","datasetId":5786322},{"sourceId":9578924,"sourceType":"datasetVersion","datasetId":5801157}],"dockerImageVersionId":30733,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Import Libralies","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nsys.path.append('/kaggle/input/timm-3d/')\nsys.path.append('/kaggle/input/lhwcv-rsna2024-code/')\nos.environ['OMP_NUM_THREADS'] = '1'\n\nimport pandas as pd\nimport numpy as np\nimport torch\nimport tqdm\nimport random\nimport gc\nimport pickle\n\nfrom concurrent.futures.thread import ThreadPoolExecutor\n\nfrom infer_scripts.v2.dataset import RSNA24DatasetTest_LHW_keypoint_3D_Axial, RSNA24DatasetTest_LHW_V2, \\\n    RSNA24DatasetTest_LHW_keypoint_3D_Saggital\nfrom infer_scripts.v2.models import RSNA24Model_Keypoint_3D_Sag_V2\nfrom infer_scripts.v2.dataset import RSNA24DatasetTest_LHW_V2, v2_collate_fn\nfrom infer_scripts.v2.models import HybridModel_V2\nfrom infer_scripts.v20.models import build_v20_sag_model\nfrom infer_scripts.v24.infer_utils import *\n\nimport warnings\nwarnings.filterwarnings('ignore', category=pd.errors.SettingWithCopyWarning)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:36:41.753428Z","iopub.execute_input":"2024-10-08T16:36:41.754065Z","iopub.status.idle":"2024-10-08T16:36:51.055596Z","shell.execute_reply.started":"2024-10-08T16:36:41.754029Z","shell.execute_reply":"2024-10-08T16:36:51.054796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def setup_seed(seed, deterministic=True):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = deterministic\n    torch.backends.cudnn.benchmark = False","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:36:51.057071Z","iopub.execute_input":"2024-10-08T16:36:51.057499Z","iopub.status.idle":"2024-10-08T16:36:51.062982Z","shell.execute_reply.started":"2024-10-08T16:36:51.057476Z","shell.execute_reply":"2024-10-08T16:36:51.062099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 8620\nsetup_seed(SEED, deterministic=True)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:36:51.064026Z","iopub.execute_input":"2024-10-08T16:36:51.064311Z","iopub.status.idle":"2024-10-08T16:36:51.086981Z","shell.execute_reply.started":"2024-10-08T16:36:51.064288Z","shell.execute_reply":"2024-10-08T16:36:51.086132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_root = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\ntest_series_descriptions_fn = 'test_series_descriptions.csv'\nimage_dir = 'test_images'\n\n# test_series_descriptions_fn = 'train_series_descriptions.csv'\n# image_dir = 'train_images'\n\nN_WORKERS = os.cpu_count()\nbatch_size = 4\n\nkeypoint_dir = './keypoints_pred/'\nCACHE_DIR = './cache/'\n\nmodel_dir = \"/kaggle/input/lhwcv-rsna2024-final-models/\"\nmodel_dir2 = \"/kaggle/input/lhwcv-rsna2024-final-models3/\"\nmodel_dir2_old = \"/kaggle/input/lhwcv-rsna2024-final-models2/\"","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:36:51.089251Z","iopub.execute_input":"2024-10-08T16:36:51.089575Z","iopub.status.idle":"2024-10-08T16:36:51.095051Z","shell.execute_reply.started":"2024-10-08T16:36:51.089546Z","shell.execute_reply":"2024-10-08T16:36:51.094207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.makedirs(keypoint_dir, exist_ok=True)\nos.makedirs(CACHE_DIR, exist_ok=True)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:36:51.095958Z","iopub.execute_input":"2024-10-08T16:36:51.096223Z","iopub.status.idle":"2024-10-08T16:36:51.102644Z","shell.execute_reply.started":"2024-10-08T16:36:51.096202Z","shell.execute_reply":"2024-10-08T16:36:51.101744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONDITIONS = [\n    'spinal_canal_stenosis', \n    'left_neural_foraminal_narrowing', \n    'right_neural_foraminal_narrowing',\n    'left_subarticular_stenosis',\n    'right_subarticular_stenosis'\n]\n\nLEVELS = [\n    'l1_l2',\n    'l2_l3',\n    'l3_l4',\n    'l4_l5',\n    'l5_s1',\n]\ndf = pd.read_csv(f'{data_root}/{test_series_descriptions_fn}')\nstudy_ids = list(df['study_id'].unique())\n\nrow_names = []\nfor si in study_ids:\n    for cond in CONDITIONS:\n        for level in LEVELS:\n            row_names.append(str(si) + '_' + cond + '_' + level)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:36:51.103710Z","iopub.execute_input":"2024-10-08T16:36:51.103992Z","iopub.status.idle":"2024-10-08T16:36:51.122571Z","shell.execute_reply.started":"2024-10-08T16:36:51.103971Z","shell.execute_reply":"2024-10-08T16:36:51.121914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_sub = pd.read_csv(f'{data_root}/sample_submission.csv')\nLABELS = list(sample_sub.columns[1:])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:36:51.123667Z","iopub.execute_input":"2024-10-08T16:36:51.123945Z","iopub.status.idle":"2024-10-08T16:36:51.142371Z","shell.execute_reply.started":"2024-10-08T16:36:51.123924Z","shell.execute_reply":"2024-10-08T16:36:51.132283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1. keypoint detection","metadata":{}},{"cell_type":"code","source":"\ndef infer_axial_3d_keypoints_v2(data_root,\n                                study_ids,\n                                model_dir,\n                                model_dir2,\n                                device,\n                                num_workers=8,\n                                is_parallel=False):\n    \"\"\"\n    this has moved to infer_axial_3d_keypoints_v24 for dataloader speed\n    \"\"\"\n    model_cfgs = [\n        {\n            'sub_dir': 'keypoint_3d_v2_axial/densenet161_lr_0.0006',\n            'backbone': 'densenet161',\n        }\n    ]\n    v2_folds = 3\n    models = []\n    for cfg in model_cfgs:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        for fold in range(v2_folds):\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_fold_{fold}_ema.pt'\n            model = RSNA24Model_Keypoint_3D(backbone,\n                                            num_classes=30,\n                                            pretrained=False)\n            model.load_state_dict(torch.load(fn))\n            model.to(device)\n            model.eval()\n            if is_parallel:\n                model = nn.DataParallel(model)\n            models.append(model)\n\n    dset = RSNA24DatasetTest_LHW_keypoint_3D_Axial(data_root,\n                                                   test_series_descriptions_fn,\n                                                   study_ids,\n                                                   image_dir=data_root + f'/{image_dir}/')\n\n    dloader = DataLoader(dset, batch_size=8, num_workers=num_workers)\n\n    study_id_to_pred_keypoints = {}\n    for volumns, study_ids_, sids, depths in tqdm.tqdm(dloader,desc=f'{device}'):\n        bs, _, _, _ = volumns.shape\n        volumns = volumns.unsqueeze(1).to(device)\n        keypoints = None\n        with torch.no_grad():\n            with autocast:\n                for i in range(len(models)):\n                    p = models[i](volumns).cpu().numpy().reshape(bs, 10, 3)\n                    if keypoints is None:\n                        keypoints = p\n                    else:\n                        keypoints += p\n        keypoints = keypoints / len(models)\n        for idx, study_id in enumerate(study_ids_):\n            study_id = int(study_id)\n            sid = int(sids[idx])\n            d = depths[idx]\n            # print(study_id, sid)\n            if study_id not in study_id_to_pred_keypoints.keys():\n                study_id_to_pred_keypoints[study_id] = {}\n            study_id_to_pred_keypoints[study_id][sid] = {\n                'points': keypoints[idx],\n                'd': int(d)\n            }\n    return study_id_to_pred_keypoints\n\n\ndef infer_axial_3d_keypoints_v24(data_root,\n                                 study_ids,\n                                 model_dir,\n                                 model_dir2,\n                                 device,\n                                 num_workers=8,\n                                 is_parallel=False,\n                                 ):\n    z_model_cfgs = [\n        {\n            'sub_dir': 'keypoint_3d_v24_axial/level_cls/convnext_small.in12k_ft_in1k_384',\n            'backbone': 'convnext_small.in12k_ft_in1k_384',\n        },\n    ]\n    z_folds = 1\n    z_models = []\n    for cfg in z_model_cfgs:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        for fold in range(z_folds):\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_fold_{fold}_ema.pt'\n            print('load: ', fn)\n            model = Axial_Level_Cls_Model_for_Test(backbone,\n                                                   pretrained=False).to(device)\n            model.load_state_dict(torch.load(fn))\n            model.eval()\n            z_models.append(model)\n\n    \n    #\n    xy_model_cfgs = [\n#         {\n#             'sub_dir': 'keypoint_3d_v24_axial/axial_2d_keypoints/densenet161_lr_0.0006',\n#             'backbone': 'densenet161',\n#         },\n        {\n            'sub_dir': 'keypoint_2d_v24_axial/xception65.tf_in1k_lr_0.0006/',\n            'backbone': 'xception65.tf_in1k',\n        },\n    ]\n    xy_folds = 5\n    xy_models = []\n    for cfg in xy_model_cfgs:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        for fold in range(xy_folds):\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_fold_{fold}_ema.pt'\n            if not os.path.exists(fn):\n                fn = os.path.join(model_dir2_old, sub_dir)\n                fn = fn + f'/best_fold_{fold}_ema.pt'\n            \n            print('load: ', fn)\n            model = RSNA24Model_Keypoint_2D(backbone,\n                                            pretrained=False,\n                                            num_classes=4,\n                                            ).to(device)\n            model.load_state_dict(torch.load(fn))\n            model.eval()\n            xy_models.append(model)\n\n    v24_axial_pred_keypoints_info = {}\n    v24_axial_pred_keypoints_info_2 = {}\n    v2_axial_pred_keypoints_info = {}\n\n    # v2 moved to here for dataloader speed\n    boost_v2_dataloader = False\n    if boost_v2_dataloader:\n        model_cfgs = [\n            {\n                'sub_dir': 'keypoint_3d_v2_axial/densenet161_lr_0.0006',\n                'backbone': 'densenet161',\n            }\n        ]\n        v2_folds = 3\n        v2_models = []\n        for cfg in model_cfgs:\n            sub_dir = cfg['sub_dir']\n            backbone = cfg['backbone']\n            for fold in range(v2_folds):\n                fn = os.path.join(model_dir, sub_dir)\n                fn = fn + f'/best_fold_{fold}_ema.pt'\n                model = RSNA24Model_Keypoint_3D(backbone,\n                                                num_classes=30,\n                                                pretrained=False).to(device)\n                model.load_state_dict(torch.load(fn))\n                model.eval()\n                v2_models.append(model)\n\n    dset = Axial_Level_Dataset_Multi_V24(data_root, study_ids,\n                                         test_series_descriptions_fn=test_series_descriptions_fn,\n                                         image_dir=image_dir)\n    dloader = DataLoader(dset, batch_size=1,\n                         num_workers=num_workers,\n                         collate_fn=axial_v24_collate_fn)\n\n\n    for data_dict in tqdm.tqdm(dloader,desc=f'{device}'):\n        study_id, series_id_list, pred_keypoints = axial_v24_infer_z(z_models, data_dict, device)\n        study_id = int(study_id)\n        pred_keypoints = axial_v24_infer_xy(xy_models, pred_keypoints, data_dict)\n\n        if boost_v2_dataloader:\n            v2_pred = axial_v2_infer_xyz(v2_models, data_dict)\n            v2_axial_pred_keypoints_info[study_id] = v2_pred\n\n        v24_axial_pred_keypoints_info[study_id] = {}\n        for i, sid in enumerate(series_id_list):\n            sid = int(sid)\n            v24_axial_pred_keypoints_info[study_id][sid] = {\n                'points': pred_keypoints[i]  # 2, 5, 4\n            }\n        \n    if not boost_v2_dataloader:\n        v2_axial_pred_keypoints_info = infer_axial_3d_keypoints_v2(data_root,\n                                                                   study_ids,\n                                                                   model_dir,\n                                                                   model_dir2,\n                                                                   device,\n                                                                   num_workers,\n                                                                   is_parallel)\n    \n    return (v24_axial_pred_keypoints_info, v2_axial_pred_keypoints_info)\n\n\ndef _infer_sag_xy(models, img):\n    # img: b, 1, h, w\n    pts = None\n    for i in range(len(models)):\n        with torch.no_grad():\n            with autocast:\n                p = models[i](img)\n                if pts is None:\n                    pts = p\n                else:\n                    pts += p\n    pts = pts / len(models)\n    return pts\n\n\ndef infer_sag_2d_keypoints_v2(data_root,\n                              study_ids,\n                              model_dir,\n                              model_dir2,\n                              device,\n                              num_workers=8,\n                              is_parallel=False):\n    model_cfgs_t1 = [\n#         {\n#             'sub_dir': 'keypoint_2d_v20_sag_t1/densenet161_lr_0.0006_sag_T1',\n#             'backbone': 'densenet161',\n#         },\n        {\n            'sub_dir': 'keypoint_2d_v20_sag_t1/fastvit_ma36.apple_dist_in1k_lr_0.0008_sag_T1_not_exclude_hard',\n            'backbone': 'fastvit_ma36.apple_dist_in1k',\n        },\n\n    ]\n    v2_folds = 5\n    models_t1 = []\n    for cfg in model_cfgs_t1:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        for fold in range(v2_folds):\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_fold_{fold}_ema.pt'\n            if not os.path.exists(fn):\n                fn = os.path.join(model_dir2, sub_dir)\n                fn = fn + f'/best_fold_{fold}_ema.pt'\n\n            model = RSNA24Model_Keypoint_2D(backbone,\n                                            num_classes=10,\n                                            pretrained=False)\n            print(f'{device} -->load: ',fn)\n            model.load_state_dict(torch.load(fn))\n            model.to(device)\n            model.eval()\n            if is_parallel:\n                model = nn.DataParallel(model)\n            models_t1.append(model)\n    #\n    model_cfgs_t2 = [\n#         {\n#             'sub_dir': 'keypoint_2d_v20_sag_t2/densenet161_lr_0.0006_sag_T2',\n#             'backbone': 'densenet161',\n#         },\n        {\n            'sub_dir': 'keypoint_2d_v20_sag_t2/convformer_s36.sail_in22k_ft_in1k_384_lr_0.0008_sag_T2_not_exclude_hard',\n            'backbone': 'convformer_s36.sail_in22k_ft_in1k_384',\n        },\n    ]\n    v2_folds = 5\n    models_t2 = []\n    for cfg in model_cfgs_t2:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        for fold in range(v2_folds):\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_fold_{fold}_ema.pt'\n            if not os.path.exists(fn):\n                fn = os.path.join(model_dir2, sub_dir)\n                fn = fn + f'/best_fold_{fold}_ema.pt'\n\n            model = RSNA24Model_Keypoint_2D(backbone,\n                                            num_classes=10,\n                                            pretrained=False)\n            model.load_state_dict(torch.load(fn))\n            model.to(device)\n            model.eval()\n            if is_parallel:\n                model = nn.DataParallel(model)\n            models_t2.append(model)\n\n    #\n    dset = RSNA24DatasetTest_LHW_V2(data_root, study_ids,\n                                    image_dir=data_root + f'/{image_dir}/',\n                                    test_series_descriptions_fn=test_series_descriptions_fn,\n                                    with_axial=False,\n                                    cache_dir=CACHE_DIR)\n    print(f'{device} build data done!')\n    dloader = DataLoader(dset, batch_size=8, num_workers=num_workers)\n\n    study_id_to_pred_keypoints = {}\n\n    for tensor_dict in tqdm.tqdm(dloader, desc=f'{device}'):\n        imgs = tensor_dict['img']\n        study_ids = tensor_dict['study_id']\n        bs, _, _, _ = imgs.shape\n        imgs = imgs.to(device)\n        s_t1 = imgs[:, :10]\n        s_t1 = s_t1[:, 5:6]  # take the center\n        s_t2 = imgs[:, 10: 20]\n        s_t2 = s_t2[:, 5: 6]\n\n        keypoints_t1 = _infer_sag_xy(models_t1, s_t1).reshape(bs, 1, 5, 2)\n        keypoints_t2 = _infer_sag_xy(models_t2, s_t2).reshape(bs, 1, 5, 2)\n        keypoints = torch.cat((keypoints_t1, keypoints_t2), dim=1).cpu().numpy()\n\n        for idx, study_id in enumerate(study_ids):\n            study_id = int(study_id)\n            study_id_to_pred_keypoints[study_id] = keypoints[idx] * 512\n    return study_id_to_pred_keypoints\n\n\ndef infer_sag_3d_keypoints_v2(data_root,\n                              study_ids,\n                              model_dir,\n                              model_dir2,\n                              device,\n                              num_workers=8,\n                              is_parallel=False,\n                              ):\n    model_cfgs = [\n        {\n            'sub_dir': 'keypoint_3d_v2_sag/densenet161_lr_0.0006/',\n            'backbone': 'densenet161',\n        }\n    ]\n    v2_folds = 3\n    models = []\n    for cfg in model_cfgs:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        for fold in range(v2_folds):\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_fold_{fold}_ema.pt'\n            print(f'{device} -->load: ', fn)\n            model = RSNA24Model_Keypoint_3D_Sag_V2(model_name=backbone,\n                                                   pretrained=False)\n            model.load_state_dict(torch.load(fn))\n            model.to(device)\n            model.eval()\n            if is_parallel:\n                model = nn.DataParallel(model)\n            models.append(model)\n    #\n    dset = RSNA24DatasetTest_LHW_keypoint_3D_Saggital(data_root,\n                                                      test_series_descriptions_fn,\n                                                      study_ids,\n                                                      image_dir=data_root + f'{image_dir}',\n                                                      cache_dir=CACHE_DIR)\n\n    dloader = DataLoader(dset, batch_size=8, num_workers=num_workers)\n\n    study_id_to_pred_keypoints_sag = {}\n    for volumns, study_ids_ in tqdm.tqdm(dloader, desc=f'{device}'):\n        with torch.no_grad():\n            volumns = volumns.to(device)\n            with autocast:\n                keypoints = None\n                for i in range(len(models)):\n                    y = models[i](volumns)\n                    bs = y.shape[0]\n                    if keypoints is None:\n                        keypoints = y.cpu().numpy().reshape(bs, 3, 5, 3)\n                    else:\n                        keypoints += y.cpu().numpy().reshape(bs, 3, 5, 3)\n                keypoints = keypoints / len(models)\n        for idx, study_id in enumerate(study_ids_):\n            study_id = int(study_id)\n            study_id_to_pred_keypoints_sag[study_id] = {\n                'points': keypoints[idx],\n            }\n    return study_id_to_pred_keypoints_sag\n\n\ndef infer_sag_3d_t1_keypoints_v20(data_root,\n                                  study_ids,\n                                  model_dir,\n                                  model_dir2,\n                                  device,\n                                  num_workers=8,\n                                  is_parallel=False):\n    with_cascade = False\n    if with_cascade:\n        model_cfgs_t1_2d = [\n            {\n                'sub_dir': 'keypoint_2d_v20_sag_t1/densenet161_lr_0.0006_sag_T1',\n                'backbone': 'densenet161',\n            }\n        ]\n        v2_folds = 5\n        models_t1_2d = []\n        for cfg in model_cfgs_t1_2d:\n            sub_dir = cfg['sub_dir']\n            backbone = cfg['backbone']\n            for fold in range(v2_folds):\n                fn = os.path.join(model_dir, sub_dir)\n                fn = fn + f'/best_fold_{fold}_ema.pt'\n                model = RSNA24Model_Keypoint_2D(backbone,\n                                                num_classes=10,\n                                                pretrained=False)\n                model.load_state_dict(torch.load(fn))\n                model.to(device)\n                model.eval()\n                if is_parallel:\n                    model = nn.DataParallel(model)\n                models_t1_2d.append(model)\n\n    #\n    pred_sag_keypoints_infos_3d_t1 = {}\n\n    model_cfgs = [\n#         {\n#             'sub_dir': 'keypoint_3d_v20_sag_t1/densenet161_lr_0.0006_sag_T1',\n#             'backbone': 'densenet161',\n#         },\n        {\n            'sub_dir': 'keypoint_3d_v20_sag_t1/gluon_resnet152_v1s_lr_0.0006_sag_T1_not_exclude_hard',\n            'backbone': 'gluon_resnet152_v1s',\n        }\n    ]\n    folds = 5\n    models = []\n    for cfg in model_cfgs:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        for fold in range(folds):\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_fold_{fold}_ema.pt'\n            if not os.path.exists(fn):\n                fn = os.path.join(model_dir2, sub_dir)\n                fn = fn + f'/best_fold_{fold}_ema.pt'\n            model = RSNA24Model_Keypoint_3D(backbone,\n                                            in_chans=1,\n                                            pretrained=False,\n                                            num_classes=30).to(device)\n            model.load_state_dict(torch.load(fn))\n            model.eval()\n            if is_parallel:\n                model = nn.DataParallel(model)\n            models.append(model)\n\n    dset = Sag_3D_Point_Dataset_V24(data_root,\n                                    study_ids,\n                                    test_series_descriptions_fn=test_series_descriptions_fn,\n                                    image_dir=image_dir,\n                                    series_description='Sagittal T1',\n                                    with_origin_arr=True if with_cascade else False,\n                                    cache_dir=CACHE_DIR\n                                    )\n\n    dloader = DataLoader(dset, batch_size=1 if with_cascade else 8,\n                         num_workers=num_workers)\n\n    for tensor_dict in tqdm.tqdm(dloader,desc=f'{device}'):\n        with torch.no_grad():\n            x = tensor_dict['imgs'].to(device)\n            study_id_list = tensor_dict['study_id']\n            series_id_list = tensor_dict['series_id']\n            with autocast:\n                y = None\n                for i in range(len(models)):\n                    p = models[i](x)\n                    if y is None:\n                        y = p\n                    else:\n                        y += p\n                y = y / len(models)\n            bs = y.shape[0]\n            for b in range(bs):\n                study_id = int(study_id_list[b])\n                series_id = int(series_id_list[b])\n                pts = y[b].reshape(-1, 5, 3).cpu().numpy()\n\n                if with_cascade:\n                    origin_depth = int(tensor_dict['origin_depth'][b])\n                    origin_imgs = tensor_dict['origin_imgs'][b].to(device)\n                    sub_imgs = []\n                    for n in range(pts.shape[0]):\n                        for level in range(5):\n                            z = pts[n, level, 2] * origin_depth\n                            z = int(np.round(z))\n                            if z < 0:\n                                z = 0\n                            if z > origin_depth - 1:\n                                z = origin_depth - 1\n\n                            img = origin_imgs[z].unsqueeze(0).unsqueeze(0)\n                            sub_imgs.append(img)\n                    sub_imgs = torch.cat(sub_imgs, dim=0)  # 2*5, 1, 512, 512\n                    xy_pred = _infer_sag_xy(models_t1_2d, sub_imgs).reshape(-1, 5, 5, 2).cpu().numpy()\n                    for n in range(pts.shape[0]):\n                        for level in range(5):\n                            pts[n, level, :2] = xy_pred[n, level, level]\n\n                if study_id not in pred_sag_keypoints_infos_3d_t1:\n                    pred_sag_keypoints_infos_3d_t1[study_id] = {}\n                pred_sag_keypoints_infos_3d_t1[study_id][series_id] = pts\n\n    return pred_sag_keypoints_infos_3d_t1\n\n\ndef infer_sag_3d_t2_keypoints_v20(data_root,\n                                  study_ids,\n                                  model_dir,\n                                  model_dir2,\n                                  device,\n                                  num_workers=8,\n                                  is_parallel=False):\n    with_cascade = False\n    if with_cascade:\n        model_cfgs_t1_2d = [\n            {\n                'sub_dir': 'keypoint_2d_v20_sag_t2/densenet161_lr_0.0006_sag_T2',\n                'backbone': 'densenet161',\n            }\n        ]\n        v2_folds = 5\n        models_t2_2d = []\n        for cfg in model_cfgs_t1_2d:\n            sub_dir = cfg['sub_dir']\n            backbone = cfg['backbone']\n            for fold in range(v2_folds):\n                fn = os.path.join(model_dir, sub_dir)\n                fn = fn + f'/best_fold_{fold}_ema.pt'\n                model = RSNA24Model_Keypoint_2D(backbone,\n                                                num_classes=10,\n                                                pretrained=False)\n                model.load_state_dict(torch.load(fn))\n                model.to(device)\n                model.eval()\n                if is_parallel:\n                    model = nn.DataParallel(model)\n                models_t2_2d.append(model)\n\n    #\n    pred_sag_keypoints_infos_3d_t2 = {}\n\n    model_cfgs = [\n#         {\n#             'sub_dir': 'keypoint_3d_v20_sag_t2/densenet161_lr_0.0006_sag_T2',\n#             'backbone': 'densenet161',\n#         },\n        {\n            'sub_dir': 'keypoint_3d_v20_sag_t2/gluon_resnet152_v1s_lr_0.0006_sag_T2_not_exclude_hard',\n            'backbone': 'gluon_resnet152_v1s',\n        }\n    ]\n    folds = 5\n    models = []\n    for cfg in model_cfgs:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        for fold in range(folds):\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_fold_{fold}_ema.pt'\n            if not os.path.exists(fn):\n                fn = os.path.join(model_dir2, sub_dir)\n                fn = fn + f'/best_fold_{fold}_ema.pt'\n            model = RSNA24Model_Keypoint_3D(backbone,\n                                            in_chans=1,\n                                            pretrained=False,\n                                            num_classes=15).to(device)\n            model.load_state_dict(torch.load(fn))\n            model.eval()\n            if is_parallel:\n                model = nn.DataParallel(model)\n            models.append(model)\n\n    dset = Sag_3D_Point_Dataset_V24(data_root,\n                                    study_ids,\n                                    test_series_descriptions_fn=test_series_descriptions_fn,\n                                    image_dir=image_dir,\n                                    series_description='Sagittal T2/STIR',\n                                    cache_dir=CACHE_DIR,\n                                    with_origin_arr=True if with_cascade else False\n                                    )\n\n    dloader = DataLoader(dset, batch_size=1 if with_cascade else 8,\n                         num_workers=num_workers)\n\n    for tensor_dict in tqdm.tqdm(dloader,desc=f'{device}'):\n        with torch.no_grad():\n            x = tensor_dict['imgs'].to(device)\n            study_id_list = tensor_dict['study_id']\n            series_id_list = tensor_dict['series_id']\n            with autocast:\n                y = None\n                for i in range(len(models)):\n                    p = models[i](x)\n                    if y is None:\n                        y = p\n                    else:\n                        y += p\n                y = y / len(models)\n            bs = y.shape[0]\n            for b in range(bs):\n                study_id = int(study_id_list[b])\n                series_id = int(series_id_list[b])\n                pts = y[b].reshape(-1, 5, 3).cpu().numpy()\n\n                if with_cascade:\n                    origin_depth = int(tensor_dict['origin_depth'][b])\n                    origin_imgs = tensor_dict['origin_imgs'][b].to(device)\n                    sub_imgs = []\n                    for n in range(pts.shape[0]):\n                        for level in range(5):\n                            z = pts[n, level, 2] * origin_depth\n                            z = int(np.round(z))\n                            if z < 0:\n                                z = 0\n                            if z > origin_depth - 1:\n                                z = origin_depth - 1\n\n                            img = origin_imgs[z].unsqueeze(0).unsqueeze(0)\n                            sub_imgs.append(img)\n                    sub_imgs = torch.cat(sub_imgs, dim=0)  # 2*5, 1, 512, 512\n                    xy_pred = _infer_sag_xy(models_t2_2d, sub_imgs).reshape(-1, 5, 5, 2).cpu().numpy()\n                    for n in range(pts.shape[0]):\n                        for level in range(5):\n                            pts[n, level, :2] = xy_pred[n, level, level]\n\n                if study_id not in pred_sag_keypoints_infos_3d_t2:\n                    pred_sag_keypoints_infos_3d_t2[study_id] = {}\n                pred_sag_keypoints_infos_3d_t2[study_id][series_id] = pts\n\n    return pred_sag_keypoints_infos_3d_t2","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:36:51.144116Z","iopub.execute_input":"2024-10-08T16:36:51.144460Z","iopub.status.idle":"2024-10-08T16:36:51.235864Z","shell.execute_reply.started":"2024-10-08T16:36:51.144428Z","shell.execute_reply":"2024-10-08T16:36:51.235128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef infer_all_keypoints(data_root,\n                        model_dir,\n                        model_dir2,\n                        with_v2_sag_center_slice_2d=True,\n                        with_v2_axial_3d=True,\n                        with_v2_sag_3d=True,\n                        with_v20_sag_t1_3d=True,\n                        with_v20_sag_t2_3d=True,\n                        with_v24_axial_3d=True\n                        ):\n    print('with_v2_sag_center_slice_2d: ', with_v2_sag_center_slice_2d)\n    print('with_v2_axial_3d: ', with_v2_axial_3d)\n    print('with_v2_sag_3d: ', with_v2_sag_3d)\n    print('with_v20_sag_t1_3d: ', with_v20_sag_t1_3d)\n    print('with_v20_sag_t2_3d: ', with_v20_sag_t2_3d)\n    print('with_v24_axial_3d: ', with_v24_axial_3d)\n\n    df = pd.read_csv(f'{data_root}/{test_series_descriptions_fn}')\n    study_ids = df['study_id'].unique().tolist()\n    # study_ids = study_ids[:16]\n\n    ret_dict = {}\n    device_0 = torch.device('cuda:0')\n\n\n    is_parallel = torch.cuda.device_count() > 1 and len(study_ids) >=2\n    print('is_parallel: ', is_parallel)\n    if with_v2_axial_3d or with_v24_axial_3d:\n        pred_dict_v24,  pred_dict_v2 = \\\n            infer_axial_3d_keypoints_v24(data_root, study_ids, model_dir,model_dir2,\n                                         device_0,N_WORKERS,is_parallel)\n\n        ret_dict['axial_3d_keypoints_v2'] = pred_dict_v2\n        ret_dict['axial_3d_keypoints_v24'] = pred_dict_v24\n\n\n    if with_v2_sag_3d:\n        print('infer_sag_3d_keypoints_v2...')\n        pred_dict = infer_sag_3d_keypoints_v2(data_root, study_ids,\n                                              model_dir, model_dir2, device_0, N_WORKERS,\n                                              is_parallel)\n        ret_dict['sag_3d_keypoints_v2'] = pred_dict\n\n\n    if with_v2_sag_center_slice_2d:\n        print('infer_sag_2d_keypoints_v2...')\n        pred_dict = infer_sag_2d_keypoints_v2(data_root, study_ids,\n                                              model_dir, model_dir2, device_0, N_WORKERS,\n                                              is_parallel)\n        ret_dict['sag_keypoints_v2'] = pred_dict\n\n\n    if with_v20_sag_t1_3d:\n        print('infer_sag_3d_t1_keypoints_v20...')\n        pred_dict = infer_sag_3d_t1_keypoints_v20(data_root, study_ids,\n                                                  model_dir, model_dir2, device_0, N_WORKERS,\n                                                  is_parallel)\n        ret_dict['sag_t1_3d_keypoints_v20'] = pred_dict\n\n    if with_v20_sag_t2_3d:\n        print('infer_sag_3d_t2_keypoints_v20..')\n        pred_dict = infer_sag_3d_t2_keypoints_v20(data_root, study_ids,\n                                                  model_dir, model_dir2, device_0, N_WORKERS,\n                                                  is_parallel)\n        ret_dict['sag_t2_3d_keypoints_v20'] = pred_dict\n\n    return ret_dict","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:36:51.237083Z","iopub.execute_input":"2024-10-08T16:36:51.237457Z","iopub.status.idle":"2024-10-08T16:36:51.249853Z","shell.execute_reply.started":"2024-10-08T16:36:51.237426Z","shell.execute_reply":"2024-10-08T16:36:51.248938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nif True:\n    \n    save_dir = keypoint_dir\n\n    sag_keypoints_v2 = True\n    axial_3d_keypoints_v2 = True\n    sag_3d_keypoints_v2 = True\n\n    sag_t1_3d_keypoints_v20 = True\n    sag_t2_3d_keypoints_v20 = True\n    axial_3d_keypoints_v24 = True\n\n    sag_keypoints_v2_save_fn = f'{save_dir}/sag_keypoints_v2.pkl'\n    axial_3d_keypoints_v2_save_fn = f'{save_dir}/axial_3d_keypoints_v2.pkl'\n    sag_3d_keypoints_v2_save_fn = f'{save_dir}/sag_3d_keypoints_v2.pkl'\n    sag_t1_3d_keypoints_v20_save_fn = f'{save_dir}/sag_t1_3d_keypoints_v20.pkl'\n    sag_t2_3d_keypoints_v20_save_fn = f'{save_dir}/sag_t2_3d_keypoints_v20.pkl'\n    axial_3d_keypoints_v24_fn = f'{save_dir}/axial_3d_keypoints_v24.pkl'\n\n    if os.path.exists(sag_keypoints_v2_save_fn):\n        sag_keypoints_v2 = False\n    if os.path.exists(axial_3d_keypoints_v2_save_fn):\n        axial_3d_keypoints_v2 = False\n    if os.path.exists(sag_3d_keypoints_v2_save_fn):\n        sag_3d_keypoints_v2 = False\n\n    if os.path.exists(sag_t1_3d_keypoints_v20_save_fn):\n        sag_t1_3d_keypoints_v20 = False\n    if os.path.exists(sag_t2_3d_keypoints_v20_save_fn):\n        sag_t2_3d_keypoints_v20 = False\n    if os.path.exists(axial_3d_keypoints_v24_fn):\n        axial_3d_keypoints_v24 = False\n\n    ret_dict = infer_all_keypoints(data_root,\n                                   model_dir,\n                                   model_dir2,\n                                   with_v2_sag_center_slice_2d=sag_keypoints_v2,\n                                   with_v2_axial_3d=axial_3d_keypoints_v2,\n                                   with_v2_sag_3d=sag_3d_keypoints_v2,\n                                   with_v20_sag_t1_3d=sag_t1_3d_keypoints_v20,\n                                   with_v20_sag_t2_3d=sag_t2_3d_keypoints_v20,\n                                   with_v24_axial_3d=axial_3d_keypoints_v24, )\n\n    for k, v in ret_dict.items():\n        with open(f'{save_dir}/{k}.pkl', 'wb') as file_handle:\n            pickle.dump(v, file_handle)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:36:51.253224Z","iopub.execute_input":"2024-10-08T16:36:51.253624Z","iopub.status.idle":"2024-10-08T16:39:28.032853Z","shell.execute_reply.started":"2024-10-08T16:36:51.253600Z","shell.execute_reply":"2024-10-08T16:39:28.031581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. condition classification","metadata":{}},{"cell_type":"code","source":"def infer_v2(data_root,\n             study_ids,\n             model_dir,\n             save_dir,\n             device,\n             num_workers=8,\n             folds=None,\n             is_parallel=False):\n    model_cfgs = [\n        {\n            'sub_dir': 'v2_cond/pvt_v2_b2.in1k_axial_pvt_v2_b2.in1k_axial_size_256_sag_size_128',\n            'backbone_sag': 'pvt_v2_b2.in1k',\n            'backbone_axial': 'pvt_v2_b2.in1k',\n            'weight': 0.3469085736266617,\n        },\n\n        {\n            'sub_dir': 'v2_cond/convnext_small.in12k_ft_in1k_384_axial_densenet161_axial_size_256_sag_size_128',\n            'backbone_sag': 'convnext_small.in12k_ft_in1k_384',\n            'backbone_axial': 'densenet161',\n            'weight': 0.20715647998779071,\n        },\n        {\n            'sub_dir': 'v2_cond/convnext_nano.in12k_ft_in1k_axial_convnext_nano.in12k_ft_in1k_axial_size_256_sag_size_128',\n            'backbone_sag': 'convnext_nano.in12k_ft_in1k',\n            'backbone_axial': 'convnext_nano.in12k_ft_in1k',\n            'weight': 0.1927349675794976,\n        },\n        {\n            'sub_dir': 'v2_cond/pvt_v2_b1.in1k_axial_pvt_v2_b1.in1k_axial_size_256_sag_size_128',\n            'backbone_sag': 'pvt_v2_b1.in1k',\n            'backbone_axial': 'pvt_v2_b1.in1k',\n            'weight': 0.13677624328967222,\n        },\n        {\n            'sub_dir': 'v2_cond/pvt_v2_b1.in1k_axial_densenet161_axial_size_256_sag_size_128',\n            'backbone_sag': 'pvt_v2_b1.in1k',\n            'backbone_axial': 'densenet161',\n            'weight': 0.11642373551637777,\n        },\n\n    ]\n    v2_folds = list(range(5))\n    if folds is not None:\n        v2_folds = folds\n\n    models = []\n    weights = []\n    for cfg in model_cfgs:\n        sub_dir = cfg['sub_dir']\n        backbone_sag = cfg['backbone_sag']\n        backbone_axial = cfg['backbone_axial']\n        w = cfg['weight']\n        for fold in v2_folds:\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_wll_model_fold-{fold}_score.pt'\n            print('load: ', fn)\n            model = HybridModel_V2(backbone_sag,\n                                   backbone_axial,\n                                   pretrained=False)\n            model.load_state_dict(torch.load(fn))\n            model.to(device)\n            model.eval()\n            if is_parallel:\n                model = nn.DataParallel(model)\n            models.append(model)\n            weights.append(w)\n\n    weights_sum = np.sum(weights)\n    #\n    with open(f'{keypoint_dir}/sag_keypoints_v2.pkl', 'rb') as f:\n        sag_keypoints_v2 = pickle.load(f)\n\n    with open(f'{keypoint_dir}/axial_3d_keypoints_v2.pkl', 'rb') as f:\n        axial_3d_keypoints_v2 = pickle.load(f)\n\n    valid_ds = RSNA24DatasetTest_LHW_V2(data_root,\n                                        study_ids,\n                                        test_series_descriptions_fn,\n                                        study_id_to_pred_keypoints_sag=sag_keypoints_v2,\n                                        study_id_to_pred_keypoints_axial=axial_3d_keypoints_v2,\n                                        image_dir=data_root + f'{image_dir}/',\n                                        cache_dir=CACHE_DIR)\n    valid_dl = DataLoader(\n        valid_ds,\n        batch_size=4,\n        shuffle=False,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=num_workers,\n        collate_fn=v2_collate_fn\n    )\n\n    all_preds = {}\n    for tensor_d in tqdm.tqdm(valid_dl):\n        with torch.no_grad():\n            x = tensor_d['img'].to(device)\n            axial_imgs = tensor_d['axial_imgs'].to(device)\n            n_list = tensor_d['n_list']\n            cond = tensor_d['cond'].to(device)\n            study_id_list = tensor_d['study_id']\n            preds = None\n            with autocast:\n                for i in range(len(models)):\n                    p = models[i](x, axial_imgs, cond, n_list)\n                    if preds is None:\n                        preds = (p.float() * weights[i])\n                    else:\n                        preds += (p.float() * weights[i])\n            preds = preds / weights_sum\n            bs, _ = preds.shape\n            # bs, 5_cond, 5_level, 3\n            preds = preds.reshape(bs, 5, 5, 3).cpu()\n            for b in range(bs):\n                study_id = int(study_id_list[b])\n                all_preds[study_id] = preds[b]\n\n    return all_preds\n\n\ndef infer_v24(  # df_valid,\n        data_root,\n        study_ids,\n        model_dir,\n        save_dir,\n        device,\n        num_workers=8,\n        folds=None,\n        is_parallel=False,\n):\n    def _forward_tta(model, x, cond):\n        # return model(x, cond)\n        x_flip = torch.flip(x, dims=[-1, ])  # A.HorizontalFlip(p=0.5),\n        x0 = model(x, cond)\n        x1 = model(x_flip, cond)\n        return (x0 + x1) / 2\n\n    # t1\n\n    model_cfgs = [\n        {\n            'sub_dir': 'v20_cond_t1/convnext_small.in12k_ft_in1k_384_z_imgs_3_seed_8620_h128_w128_legacy',\n            'backbone': 'convnext_small.in12k_ft_in1k_384',\n            'with_cond': True,\n            'z_imgs': 3,\n            'with_gru': False,\n            'with_level_lstm': False,\n            'data_key': '3_128_128',\n            'weight': 0.20538912558481712,\n        },\n        {\n            'sub_dir': 'v20_cond_t1/convnext_small.in12k_ft_in1k_384_z_imgs_3_seed_8620_h72_w128',\n            'backbone': 'convnext_small.in12k_ft_in1k_384',\n            'with_cond': False,\n            'z_imgs': 3,\n            'with_gru': False,\n            'with_level_lstm': False,\n            'data_key': '3_72_128',\n            'weight': 0.18552050200612183,\n        },\n        {\n            'sub_dir': 'v20_cond_t1/convnext_small.in12k_ft_in1k_384_z_imgs_3_seed_8620_h72_w128_with_gru',\n            'backbone': 'convnext_small.in12k_ft_in1k_384',\n            'with_cond': False,\n            'z_imgs': 3,\n            'with_gru': True,\n            'with_level_lstm': False,\n            'data_key': '3_72_128',\n            'weight': 0.306110664432234,\n        },\n        {\n            'sub_dir': 'v20_cond_t1/convnext_tiny.in12k_ft_in1k_384_z_imgs_3_seed_8620_h128_w128_with_gru',\n            'backbone': 'convnext_tiny.in12k_ft_in1k_384',\n            'with_cond': False,\n            'z_imgs': 3,\n            'with_gru': True,\n            'with_level_lstm': False,\n            'data_key': '3_128_128',\n            'weight': 0.30297970797682705,\n        }\n\n    ]\n    v24_folds = list(range(5))\n    if folds is not None:\n        v24_folds = folds\n\n    models = []\n    data_keys = []\n    weights = []\n    for cfg in model_cfgs:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        with_cond = cfg['with_cond']\n        with_gru = cfg['with_gru']\n        with_level_lstm = cfg['with_level_lstm']\n        z_imgs = cfg['z_imgs']\n        data_key = cfg['data_key']\n        w = cfg['weight']\n        for fold in v24_folds:\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_wll_model_fold-{fold}_ema.pt'\n            print('load: ', fn)\n            model = build_v20_sag_model(\n                backbone,\n                with_cond=with_cond,\n                with_gru=with_gru,\n                with_level_lstm=with_level_lstm,\n                z_imgs=z_imgs,\n                pretrained=False\n            )\n            model.load_state_dict(torch.load(fn))\n            model.to(device)\n            model.eval()\n            if is_parallel:\n                model = nn.DataParallel(model)\n\n            models.append(model)\n            weights.append(w)\n            data_keys.append(data_key)\n    weights_sum = np.sum(weights)\n    with open(f'{keypoint_dir}/sag_t1_3d_keypoints_v20.pkl', 'rb') as f:\n        pred_sag_keypoints_infos_3d_t1 = pickle.load(f)\n\n    other_crop_size_list = [\n        (3, 72, 128),\n    ]\n    dset = Sag_T1_Dataset_V24(data_root,\n                              study_ids,\n                              pred_sag_keypoints_infos_3d_t1,\n                              test_series_descriptions_fn=test_series_descriptions_fn,\n                              image_dir=image_dir,\n                              crop_size_h=128,\n                              crop_size_w=128,\n                              cache_dir=CACHE_DIR,\n                              other_crop_size_list=other_crop_size_list\n                              )\n    dloader = DataLoader(\n        dset,\n        batch_size=batch_size,\n        shuffle=False,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=num_workers,\n    )\n    all_preds_t1 = {}\n    sag_t1_preds = []\n    for tensor_d in tqdm.tqdm(dloader):\n        with torch.no_grad():\n            for k in tensor_d.keys():\n                if k not in ['study_id']:\n                    tensor_d[k] = tensor_d[k].to(device)\n            study_id_list = tensor_d['study_id']\n            cond = tensor_d['cond']\n            preds = None\n            with autocast:\n                for i in range(len(models)):\n                    p = _forward_tta(models[i], tensor_d[data_keys[i]], cond)\n                    if preds is None:\n                        preds = (p.float() * weights[i])\n                    else:\n                        preds += (p.float() * weights[i])\n            preds = preds / weights_sum\n            bs, _ = preds.shape\n            # bs, n_cond, 5_level, 3\n            preds = preds.reshape(bs, -1, 5, 3).cpu()\n            sag_t1_preds.append(preds)\n            for b in range(bs):\n                study_id = int(study_id_list[b])\n                all_preds_t1[study_id] = preds[b]\n    sag_t1_preds = torch.cat(sag_t1_preds, dim=0)\n\n    # t2\n    model_cfgs = [\n        {\n            'sub_dir': 'v20_cond_t2_part2/rexnetr_200.sw_in12k_ft_in1k_z_imgs_3_seed_8620_h128_w128_level_lstm',\n            'backbone': 'rexnetr_200.sw_in12k_ft_in1k',\n            'with_cond': False,\n            'z_imgs': 3,\n            'with_gru': False,\n            'with_level_lstm': True,\n            'data_key': '3_128_128',\n            'weight': 0.2523387488397007,\n        },\n\n        {\n            'sub_dir': 'v20_cond_t2/convnext_small.in12k_ft_in1k_384_z_imgs_3_seed_8620_h128_w128_with_gru',\n            'backbone': 'convnext_small.in12k_ft_in1k_384',\n            'with_cond': False,\n            'z_imgs': 3,\n            'with_gru': True,\n            'with_level_lstm': False,\n            'data_key': '3_128_128',\n            'weight': 0.24855049213914635,\n        },\n        {\n            'sub_dir': 'v20_cond_t2_part2/pvt_v2_b5.in1k_z_imgs_5_seed_8620_h128_w128_level_lstm',\n            'backbone': 'pvt_v2_b5.in1k',\n            'with_cond': False,\n            'z_imgs': 5,\n            'with_gru': False,\n            'with_level_lstm': True,\n            'data_key': '5_128_128',\n            'weight': 0.19532993129183054,\n        },\n        {\n            'sub_dir': 'v20_cond_t2/convnext_small.in12k_ft_in1k_384_z_imgs_3_seed_8620_h64_w96_level_lstm',\n            'backbone': 'convnext_small.in12k_ft_in1k_384',\n            'with_cond': False,\n            'z_imgs': 3,\n            'with_gru': False,\n            'with_level_lstm': True,\n            'data_key': '3_64_96',\n            'weight': 0.14061503567169062,\n        },\n        {\n            'sub_dir': 'v20_cond_t2_part2/convnext_tiny.in12k_ft_in1k_384_z_imgs_3_seed_8620_h64_w96_level_lstm',\n            'backbone': 'convnext_tiny.in12k_ft_in1k_384',\n            'with_cond': False,\n            'z_imgs': 3,\n            'with_gru': False,\n            'with_level_lstm': True,\n            'data_key': '3_64_96',\n            'weight': 0.090578570264412,\n        },\n\n        {\n            'sub_dir': 'v20_cond_t2_part2/convnext_tiny.in12k_ft_in1k_384_z_imgs_5_seed_8620_h128_w128_level_lstm',\n            'backbone': 'convnext_tiny.in12k_ft_in1k_384',\n            'with_cond': False,\n            'z_imgs': 5,\n            'with_gru': False,\n            'with_level_lstm': True,\n            'data_key': '5_128_128',\n            'weight': 0.0725872217932197,\n        },\n    ]\n    v24_folds = list(range(5))\n    if folds is not None:\n        v24_folds = folds\n\n    models = []\n    data_keys = []\n    weights = []\n    for cfg in model_cfgs:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        with_cond = cfg['with_cond']\n        with_gru = cfg['with_gru']\n        with_level_lstm = cfg['with_level_lstm']\n        z_imgs = cfg['z_imgs']\n        data_key = cfg['data_key']\n        w = cfg['weight']\n        for fold in v24_folds:\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_wll_model_fold-{fold}_ema.pt'\n            print('load: ', fn)\n            model = build_v20_sag_model(\n                backbone,\n                with_cond=with_cond,\n                with_gru=with_gru,\n                with_level_lstm=with_level_lstm,\n                z_imgs=z_imgs,\n                pretrained=False\n            )\n            model.load_state_dict(torch.load(fn))\n            model.to(device)\n            model.eval()\n            if is_parallel:\n                model = nn.DataParallel(model)\n\n            models.append(model)\n            weights.append(w)\n            data_keys.append(data_key)\n    weights_sum = np.sum(weights)\n\n    with open(f'{keypoint_dir}/sag_t2_3d_keypoints_v20.pkl', 'rb') as f:\n        pred_sag_keypoints_infos_3d_t2 = pickle.load(f)\n\n#     with open(f'{keypoint_dir}/sag_3d_keypoints_v2.pkl', 'rb') as f:\n#         pred_sag_keypoints_infos_3d = pickle.load(f)\n\n#     for study_id in pred_sag_keypoints_infos_3d_t2.keys():\n#         pred_keypoints = pred_sag_keypoints_infos_3d[study_id]['points'].reshape(15, 3)\n#         t2_pred_keypoints = pred_keypoints[:5]\n#         t2_pred_keypoints[:, :2] = t2_pred_keypoints[:, :2] / 4.0\n#         t2_pred_keypoints[:, 2] = t2_pred_keypoints[:, 2] / 16.0\n#         for sid in pred_sag_keypoints_infos_3d_t2[study_id].keys():\n#             pred_sag_keypoints_infos_3d_t2[study_id][sid] = t2_pred_keypoints\n\n    other_crop_size_list = [\n        (3, 64, 96),\n        (5, 128, 128),\n    ]\n    \n    dset = Sag_T2_Dataset_V24(data_root,\n                              study_ids,\n                              pred_sag_keypoints_infos_3d_t2,\n                              test_series_descriptions_fn=test_series_descriptions_fn,\n                              image_dir=image_dir,\n                              crop_size_h=128,\n                              crop_size_w=128,\n                              cache_dir=CACHE_DIR,\n                              other_crop_size_list=other_crop_size_list\n                              )\n    dloader = DataLoader(\n        dset,\n        batch_size=batch_size,\n        shuffle=False,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=num_workers,\n    )\n    all_preds_t2 = {}\n    sag_t2_preds = []\n    for tensor_d in tqdm.tqdm(dloader):\n        with torch.no_grad():\n            for k in tensor_d.keys():\n                if k not in ['study_id']:\n                    tensor_d[k] = tensor_d[k].to(device)\n            study_id_list = tensor_d['study_id']\n            cond = None  # tensor_d['cond']\n            preds = None\n            with autocast:\n                for i in range(len(models)):\n                    # p = _forward_tta(models[i], tensor_d['img'], cond)\n                    p = _forward_tta(models[i], tensor_d[data_keys[i]], cond)\n                    if preds is None:\n                        preds = (p.float() * weights[i])\n                    else:\n                        preds += (p.float() * weights[i])\n            preds = preds / weights_sum\n            bs, _ = preds.shape\n            # bs, n_cond, 5_level, 3\n            preds = preds.reshape(bs, -1, 5, 3).cpu()\n            sag_t2_preds.append(preds)\n            for b in range(bs):\n                study_id = int(study_id_list[b])\n                all_preds_t2[study_id] = preds[b]\n    sag_t2_preds = torch.cat(sag_t2_preds, dim=0)\n    # return all_preds_t2\n    # axial\n\n    model_cfgs = [\n        {\n            'sub_dir': 'v24_cond_axial/convnext_small.in12k_ft_in1k_384_z_imgs_3',\n            'backbone': 'convnext_small.in12k_ft_in1k_384',\n            'axial_in_channels': 1,\n            'weight': 0.19731044421522836,\n            'z_imgs': 3,\n            'data_key': 'axial_imgs_3_128_128'\n        },\n        {\n            'sub_dir': 'v24_cond_axial/convnext_tiny.in12k_ft_in1k_384_z_imgs_3',\n            'backbone': 'convnext_tiny.in12k_ft_in1k',\n            'axial_in_channels': 1,\n            'weight': 0.30013472174512595,\n            'z_imgs': 3,\n            'data_key': 'axial_imgs_3_128_128'\n        },\n        {\n            'sub_dir': 'v24_cond_axial/densenet161_z_imgs_3',\n            'backbone': 'densenet161',\n            'axial_in_channels': 1,\n            'weight': 0.1627137568736331,\n            'z_imgs': 3,\n            'data_key': 'axial_imgs_3_128_128'\n        },\n        {\n            'sub_dir': 'v24_cond_axial/pvt_v2_b1.in1k_z_imgs_3',\n            'backbone': 'pvt_v2_b1.in1k',\n            'axial_in_channels': 1,\n            'weight': 1.0,\n            'z_imgs': 0.0987451411551922,\n            'data_key': 'axial_imgs_3_128_128'\n        },\n        {\n            'sub_dir': 'v24_cond_axial/rexnetr_200.sw_in12k_ft_in1k_z_imgs_5',\n            'backbone': 'rexnetr_200.sw_in12k_ft_in1k',\n            'axial_in_channels': 3,\n            'weight': 0.2410959360108204,\n            'z_imgs': 5,\n            'data_key': 'axial_imgs_5_128_128'\n        },\n\n    ]\n    v24_folds = list(range(5))\n    if folds is not None:\n        v24_folds = folds\n\n    models = []\n    data_keys = []\n    weights = []\n    for cfg in model_cfgs:\n        sub_dir = cfg['sub_dir']\n        backbone = cfg['backbone']\n        data_key = cfg['data_key']\n        axial_in_channels = cfg['axial_in_channels']\n        w = cfg['weight']\n        for fold in v24_folds:\n            fn = os.path.join(model_dir, sub_dir)\n            fn = fn + f'/best_wll_model_fold-{fold}_ema.pt'\n            print('load: ', fn)\n            model = Axial_HybridModel_24(backbone,\n                                         backbone,\n                                         pretrained=False,\n                                         axial_in_channels=axial_in_channels\n                                         )\n            model.load_state_dict(torch.load(fn))\n            model.to(device)\n            model.eval()\n            if is_parallel:\n                model = nn.DataParallel(model)\n\n            models.append(model)\n            weights.append(w)\n            data_keys.append(data_key)\n    weights_sum = np.sum(weights)\n\n    with open(f'{keypoint_dir}/axial_3d_keypoints_v24.pkl', 'rb') as f:\n        axial_3d_keypoints_v24 = pickle.load(f)\n\n    other_crop_size_list = [\n        (5, 128, 128),\n    ]\n    dset = Axial_Cond_Dataset_Multi_V24(data_root,\n                                        study_ids,\n                                        axial_3d_keypoints_v24,\n                                        pred_sag_keypoints_infos_3d_t2,\n                                        test_series_descriptions_fn=test_series_descriptions_fn,\n                                        image_dir=image_dir,\n                                        z_imgs=3,\n                                        other_z_imgs_list=other_crop_size_list,\n                                        cache_dir=CACHE_DIR,\n                                        )\n    dloader = DataLoader(\n        dset,\n        batch_size=batch_size,\n        shuffle=False,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=num_workers,\n    )\n    all_preds_axial = {}\n    axial_preds = []\n    for tensor_d in tqdm.tqdm(dloader):\n        with torch.no_grad():\n            for k in tensor_d.keys():\n                if k not in ['study_id']:\n                    tensor_d[k] = tensor_d[k].to(device)\n            study_id_list = tensor_d['study_id']\n            preds = None\n            with autocast:\n                for i in range(len(models)):\n                    p = models[i](tensor_d['sag_t2'], tensor_d[data_keys[i]])\n                    if preds is None:\n                        preds = (p.float() * weights[i])\n                    else:\n                        preds += (p.float() * weights[i])\n            preds = preds / weights_sum\n            bs, _ = preds.shape\n            # bs, n_cond, 5_level, 3\n            preds = preds.reshape(bs, -1, 5, 3).cpu()\n            axial_preds.append(preds)\n            for b in range(bs):\n                study_id = int(study_id_list[b])\n                all_preds_axial[study_id] = preds[b]\n\n    axial_preds = torch.cat(axial_preds, dim=0)\n    all_preds = torch.cat((sag_t2_preds, sag_t1_preds, axial_preds), dim=1)\n    #all_preds = torch.cat((sag_t2_preds, sag_t1_preds), dim=1)\n    all_preds_dict = {}\n    for i, study_id in enumerate(study_ids):\n        all_preds_dict[study_id] = all_preds[i]\n    return all_preds_dict","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:39:28.034613Z","iopub.execute_input":"2024-10-08T16:39:28.034976Z","iopub.status.idle":"2024-10-08T16:39:28.104311Z","shell.execute_reply.started":"2024-10-08T16:39:28.034943Z","shell.execute_reply":"2024-10-08T16:39:28.103336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer_all(data_root,\n              model_dir,\n              ):\n    df = pd.read_csv(f'{data_root}/{test_series_descriptions_fn}')\n    study_ids = df['study_id'].unique().tolist()\n    \n    device_0 = torch.device('cuda:0')\n\n    is_parallel = torch.cuda.device_count() > 1 and len(study_ids) >=2\n    print('is_parallel: ', is_parallel)\n    folds = None\n    \n    pred_dict_v24 = infer_v24(data_root, study_ids, model_dir,\n                              save_dir, device_0, N_WORKERS, folds, is_parallel)\n\n\n    pred_dict_v2 = infer_v2(data_root, study_ids, model_dir,\n                            save_dir, device_0, N_WORKERS, folds, is_parallel)\n\n   \n    print(len(study_ids))\n    fold_preds = []\n    for study_id in study_ids:\n        pred = pred_dict_v2[study_id]\n        fold_preds.append(pred.unsqueeze(0))\n    fold_preds = torch.cat(fold_preds, dim=0)\n    \n    print('fold_preds shape: ', fold_preds.shape)\n\n    fold_preds_v24 = []\n    for study_id in study_ids:\n        pred = pred_dict_v24[study_id]\n        fold_preds_v24.append(pred.unsqueeze(0))\n    fold_preds_v24 = torch.cat(fold_preds_v24, dim=0)\n\n    fold_preds_en = fold_preds\n    fold_preds_en[:, 0:1] = 0.5896 * fold_preds[:, 0:1] + 0.4104 * fold_preds_v24[:, 0:1]\n    fold_preds_en[:, 1:3] = 0.09 * fold_preds[:, 1:3] + 0.91 * fold_preds_v24[:, 1:3]\n    fold_preds_en[:, 3:5] = 0.38 * fold_preds[:, 3:5] + 0.62 * fold_preds_v24[:, 3:5]\n\n    fold_preds_en = nn.Softmax(dim=-1)(fold_preds_en).cpu().numpy()\n    return fold_preds_en","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:39:28.105498Z","iopub.execute_input":"2024-10-08T16:39:28.106951Z","iopub.status.idle":"2024-10-08T16:39:28.119039Z","shell.execute_reply.started":"2024-10-08T16:39:28.106925Z","shell.execute_reply":"2024-10-08T16:39:28.118157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def infer_v2_seed666(data_root,\n#                      study_ids,\n#                      model_dir,\n#                      save_dir,\n#                      device,\n#                      num_workers=8,\n#                      folds=None,\n#                      is_parallel=False):\n#     model_cfgs = [\n#         {\n#             'sub_dir': 'v2_cond/pvt_v2_b1.in1k_axial_pvt_v2_b1.in1k_axial_size_256_sag_size_128',\n#             'backbone_sag': 'pvt_v2_b1.in1k',\n#             'backbone_axial': 'pvt_v2_b1.in1k',\n#             'weight': 0.34092462492421793,\n#         },\n#         {\n#             'sub_dir': 'v2_cond/pvt_v2_b1.in1k_axial_densenet161_axial_size_256_sag_size_128',\n#             'backbone_sag': 'pvt_v2_b1.in1k',\n#             'backbone_axial': 'densenet161',\n#             'weight': 0.2458979925774557,\n#         },\n#         {\n#             'sub_dir': 'v2_cond/pvt_v2_b2.in1k_axial_pvt_v2_b2.in1k_axial_size_256_sag_size_128',\n#             'backbone_sag': 'pvt_v2_b2.in1k',\n#             'backbone_axial': 'pvt_v2_b2.in1k',\n#             'weight': 0.21641664109925535,\n#         },\n#         {\n#             'sub_dir': 'v2_cond/convnext_nano.in12k_ft_in1k_axial_convnext_nano.in12k_ft_in1k_axial_size_256_sag_size_128',\n#             'backbone_sag': 'convnext_nano.in12k_ft_in1k',\n#             'backbone_axial': 'convnext_nano.in12k_ft_in1k',\n#             'weight': 0.19676074139907107,\n#         },\n\n#     ]\n#     v2_folds = list(range(5))\n#     if folds is not None:\n#         v2_folds = folds\n\n#     models = []\n#     weights = []\n#     for cfg in model_cfgs:\n#         sub_dir = cfg['sub_dir']\n#         backbone_sag = cfg['backbone_sag']\n#         backbone_axial = cfg['backbone_axial']\n#         w = cfg['weight']\n#         for fold in v2_folds:\n#             fn = os.path.join(model_dir, sub_dir)\n#             fn = fn + f'/best_wll_model_fold-{fold}_score.pt'\n#             print('load: ', fn)\n#             model = HybridModel_V2(backbone_sag,\n#                                    backbone_axial,\n#                                    pretrained=False,\n#                                    fea_dim=512)\n#             model.load_state_dict(torch.load(fn))\n#             model.to(device)\n#             model.eval()\n#             if is_parallel:\n#                 model = nn.DataParallel(model)\n#             models.append(model)\n#             weights.append(w)\n\n#     weights_sum = np.sum(weights)\n#     #\n#     with open(f'{keypoint_dir}/sag_keypoints_v2.pkl', 'rb') as f:\n#         sag_keypoints_v2 = pickle.load(f)\n\n#     with open(f'{keypoint_dir}/axial_3d_keypoints_v2.pkl', 'rb') as f:\n#         axial_3d_keypoints_v2 = pickle.load(f)\n#     #\n#     # with open(f'{keypoint_dir}/sag_t1_3d_keypoints_v20.pkl', 'rb') as f:\n#     #     pred_sag_keypoints_infos_3d_t1 = pickle.load(f)\n\n#     valid_ds = RSNA24DatasetTest_LHW_V2(data_root,\n#                                         study_ids,\n#                                         test_series_descriptions_fn,\n#                                         study_id_to_pred_keypoints_sag=sag_keypoints_v2,\n#                                         study_id_to_pred_keypoints_axial=axial_3d_keypoints_v2,\n#                                         # pred_sag_keypoints_infos_3d_t1=pred_sag_keypoints_infos_3d_t1,\n#                                         image_dir=data_root + f'{image_dir}/',\n#                                         cache_dir=CACHE_DIR)\n#     valid_dl = DataLoader(\n#         valid_ds,\n#         batch_size=4,\n#         shuffle=False,\n#         pin_memory=False,\n#         drop_last=False,\n#         num_workers=num_workers,\n#         collate_fn=v2_collate_fn\n#     )\n\n#     all_preds = {}\n#     for tensor_d in tqdm.tqdm(valid_dl):\n#         with torch.no_grad():\n#             x = tensor_d['img'].to(device)\n#             axial_imgs = tensor_d['axial_imgs'].to(device)\n#             n_list = tensor_d['n_list']\n#             cond = tensor_d['cond'].to(device)\n#             study_id_list = tensor_d['study_id']\n#             preds = None\n#             with autocast:\n#                 for i in range(len(models)):\n#                     p = models[i](x, axial_imgs, cond, n_list)\n#                     if preds is None:\n#                         preds = (p.float() * weights[i])\n#                     else:\n#                         preds += (p.float() * weights[i])\n#             preds = preds / weights_sum\n#             bs, _ = preds.shape\n#             # bs, 5_cond, 5_level, 3\n#             preds = preds.reshape(bs, 5, 5, 3).cpu()\n#             for b in range(bs):\n#                 study_id = int(study_id_list[b])\n#                 all_preds[study_id] = preds[b]\n\n#     return all_preds\n\n\n# def infer_v24_seed666(  # df_valid,\n#         data_root,\n#         study_ids,\n#         model_dir,\n#         save_dir,\n#         device,\n#         num_workers=8,\n#         folds=None,\n#         is_parallel=False,\n# ):\n#     def _forward_tta(model, x, cond):\n#         # return model(x, cond)\n#         x_flip = torch.flip(x, dims=[-1, ])  # A.HorizontalFlip(p=0.5),\n#         x0 = model(x, cond)\n#         x1 = model(x_flip, cond)\n#         return (x0 + x1) / 2\n\n#     # t1\n#     model_cfgs = [\n#         {\n#             'sub_dir': 'v20_cond_t1/convnext_small.in12k_ft_in1k_384_z_imgs_3_seed_666_h128_w128',\n#             'backbone': 'convnext_small.in12k_ft_in1k_384',\n#             'with_cond': False,\n#             'z_imgs': 3,\n#             'with_gru': False,\n#             'with_level_lstm': False,\n#             'data_key': '3_128_128',\n#             'weight': 0.27957959035255875,\n#         },\n#         {\n#             'sub_dir': 'v20_cond_t1/convnext_small.in12k_ft_in1k_384_z_imgs_3_seed_666_h72_w128_with_gru',\n#             'backbone': 'convnext_small.in12k_ft_in1k_384',\n#             'with_cond': False,\n#             'z_imgs': 3,\n#             'with_gru': True,\n#             'with_level_lstm': False,\n#             'data_key': '3_72_128',\n#             'weight': 0.2299256035844461,\n#         },\n#         {\n#             'sub_dir': 'v20_cond_t1/rexnetr_200.sw_in12k_ft_in1k_z_imgs_3_seed_666_h72_w128_with_gru',\n#             'backbone': 'rexnetr_200.sw_in12k_ft_in1k',\n#             'with_cond': False,\n#             'z_imgs': 3,\n#             'with_gru': True,\n#             'with_level_lstm': False,\n#             'data_key': '3_72_128',\n#             'weight': 0.22004214729630475,\n#         },\n#         {\n#             'sub_dir': 'v20_cond_t1/convnext_small.in12k_ft_in1k_384_z_imgs_3_seed_666_h72_w128',\n#             'backbone': 'convnext_small.in12k_ft_in1k_384',\n#             'with_cond': False,\n#             'z_imgs': 3,\n#             'with_gru': False,\n#             'with_level_lstm': False,\n#             'data_key': '3_72_128',\n#             'weight': 0.13829760915067793,\n#         },\n#         {\n#             'sub_dir': 'v20_cond_t1/convnext_tiny.in12k_ft_in1k_384_z_imgs_3_seed_666_h72_w128_with_gru',\n#             'backbone': 'convnext_tiny.in12k_ft_in1k_384',\n#             'with_cond': False,\n#             'z_imgs': 3,\n#             'with_gru': True,\n#             'with_level_lstm': False,\n#             'data_key': '3_72_128',\n#             'weight': 0.13215504961601252,\n#         }\n\n#     ]\n#     v24_folds = list(range(5))\n#     if folds is not None:\n#         v24_folds = folds\n\n#     models = []\n#     data_keys = []\n#     weights = []\n#     for cfg in model_cfgs:\n#         sub_dir = cfg['sub_dir']\n#         backbone = cfg['backbone']\n#         with_cond = cfg['with_cond']\n#         with_gru = cfg['with_gru']\n#         with_level_lstm = cfg['with_level_lstm']\n#         z_imgs = cfg['z_imgs']\n#         data_key = cfg['data_key']\n#         w = cfg['weight']\n#         for fold in v24_folds:\n#             fn = os.path.join(model_dir, sub_dir)\n#             fn = fn + f'/best_wll_model_fold-{fold}_ema.pt'\n#             print('load: ', fn)\n#             model = build_v20_sag_model(\n#                 backbone,\n#                 with_cond=with_cond,\n#                 with_gru=with_gru,\n#                 with_level_lstm=with_level_lstm,\n#                 z_imgs=z_imgs,\n#                 pretrained=False\n#             )\n#             model.load_state_dict(torch.load(fn))\n#             model.to(device)\n#             model.eval()\n#             if is_parallel:\n#                 model = nn.DataParallel(model)\n\n#             models.append(model)\n#             weights.append(w)\n#             data_keys.append(data_key)\n#     weights_sum = np.sum(weights)\n#     with open(f'{keypoint_dir}/sag_t1_3d_keypoints_v20.pkl', 'rb') as f:\n#         pred_sag_keypoints_infos_3d_t1 = pickle.load(f)\n\n#     other_crop_size_list = [\n#         (3, 72, 128),\n#     ]\n#     dset = Sag_T1_Dataset_V24(data_root,\n#                               study_ids,\n#                               pred_sag_keypoints_infos_3d_t1,\n#                               test_series_descriptions_fn=test_series_descriptions_fn,\n#                               image_dir=image_dir,\n#                               crop_size_h=128,\n#                               crop_size_w=128,\n#                               cache_dir=CACHE_DIR,\n#                               other_crop_size_list=other_crop_size_list\n#                               )\n#     dloader = DataLoader(\n#         dset,\n#         batch_size=batch_size,\n#         shuffle=False,\n#         pin_memory=False,\n#         drop_last=False,\n#         num_workers=num_workers,\n#     )\n#     all_preds_t1 = {}\n#     sag_t1_preds = []\n#     for tensor_d in tqdm.tqdm(dloader):\n#         with torch.no_grad():\n#             for k in tensor_d.keys():\n#                 if k not in ['study_id']:\n#                     tensor_d[k] = tensor_d[k].to(device)\n#             study_id_list = tensor_d['study_id']\n#             cond = tensor_d['cond']\n#             preds = None\n#             with autocast:\n#                 for i in range(len(models)):\n#                     p = _forward_tta(models[i], tensor_d[data_keys[i]], cond)\n#                     if preds is None:\n#                         preds = (p.float() * weights[i])\n#                     else:\n#                         preds += (p.float() * weights[i])\n#             preds = preds / weights_sum\n#             bs, _ = preds.shape\n#             # bs, n_cond, 5_level, 3\n#             preds = preds.reshape(bs, -1, 5, 3).cpu()\n#             sag_t1_preds.append(preds)\n#             for b in range(bs):\n#                 study_id = int(study_id_list[b])\n#                 all_preds_t1[study_id] = preds[b]\n#     sag_t1_preds = torch.cat(sag_t1_preds, dim=0)\n\n#     # t2\n#     model_cfgs = [\n#         {\n#             'sub_dir': 'v20_cond_t2/rexnetr_200.sw_in12k_ft_in1k_z_imgs_3_seed_666_h128_w128_level_lstm',\n#             'backbone': 'rexnetr_200.sw_in12k_ft_in1k',\n#             'with_cond': False,\n#             'z_imgs': 3,\n#             'with_gru': False,\n#             'with_level_lstm': True,\n#             'data_key': '3_128_128',\n#             'weight': 0.2676011522464479,\n#         },\n\n#         {\n#             'sub_dir': 'v20_cond_t2/convnext_small.in12k_ft_in1k_384_z_imgs_3_seed_666_h128_w128_with_gru',\n#             'backbone': 'convnext_small.in12k_ft_in1k_384',\n#             'with_cond': False,\n#             'z_imgs': 3,\n#             'with_gru': True,\n#             'with_level_lstm': False,\n#             'data_key': '3_128_128',\n#             'weight': 0.25537269026075604,\n#         },\n#         {\n#             'sub_dir': 'v20_cond_t2/pvt_v2_b1.in1k_z_imgs_3_seed_666_h64_w96_with_gru',\n#             'backbone': 'pvt_v2_b1.in1k',\n#             'with_cond': False,\n#             'z_imgs': 3,\n#             'with_gru': True,\n#             'with_level_lstm': True,\n#             'data_key': '3_64_96',\n#             'weight': 0.20531756213739447,\n#         },\n#         {\n#             'sub_dir': 'v20_cond_t2/convnext_tiny.in12k_ft_in1k_z_imgs_5_seed_666_h128_w128_level_lstm',\n#             'backbone': 'convnext_tiny.in12k_ft_in1k',\n#             'with_cond': False,\n#             'z_imgs': 5,\n#             'with_gru': False,\n#             'with_level_lstm': True,\n#             'data_key': '5_128_128',\n#             'weight': 0.16238787669016758,\n#         },\n#         {\n#             'sub_dir': 'v20_cond_t2/convnext_tiny.in12k_ft_in1k_z_imgs_3_seed_666_h64_w96_level_lstm',\n#             'backbone': 'convnext_tiny.in12k_ft_in1k',\n#             'with_cond': False,\n#             'z_imgs': 3,\n#             'with_gru': False,\n#             'with_level_lstm': True,\n#             'data_key': '3_64_96',\n#             'weight': 0.08318722948976684,\n#         },\n\n#         {\n#             'sub_dir': 'v20_cond_t2/pvt_v2_b2.in1k_z_imgs_3_seed_666_h128_w128_level_lstm',\n#             'backbone': 'pvt_v2_b2.in1k',\n#             'with_cond': False,\n#             'z_imgs': 3,\n#             'with_gru': False,\n#             'with_level_lstm': True,\n#             'data_key': '3_128_128',\n#             'weight': 0.02613348917546713,\n#         },\n#     ]\n#     v24_folds = list(range(5))\n#     if folds is not None:\n#         v24_folds = folds\n\n#     models = []\n#     data_keys = []\n#     weights = []\n#     for cfg in model_cfgs:\n#         sub_dir = cfg['sub_dir']\n#         backbone = cfg['backbone']\n#         with_cond = cfg['with_cond']\n#         with_gru = cfg['with_gru']\n#         with_level_lstm = cfg['with_level_lstm']\n#         z_imgs = cfg['z_imgs']\n#         data_key = cfg['data_key']\n#         w = cfg['weight']\n#         for fold in v24_folds:\n#             fn = os.path.join(model_dir, sub_dir)\n#             fn = fn + f'/best_wll_model_fold-{fold}_ema.pt'\n#             print('load: ', fn)\n#             model = build_v20_sag_model(\n#                 backbone,\n#                 with_cond=with_cond,\n#                 with_gru=with_gru,\n#                 with_level_lstm=with_level_lstm,\n#                 z_imgs=z_imgs,\n#                 pretrained=False\n#             )\n#             model.load_state_dict(torch.load(fn))\n#             model.to(device)\n#             model.eval()\n#             if is_parallel:\n#                 model = nn.DataParallel(model)\n\n#             models.append(model)\n#             weights.append(w)\n#             data_keys.append(data_key)\n#     weights_sum = np.sum(weights)\n\n#     with open(f'{keypoint_dir}/sag_t2_3d_keypoints_v20.pkl', 'rb') as f:\n#         pred_sag_keypoints_infos_3d_t2 = pickle.load(f)\n\n#     other_crop_size_list = [\n#         (3, 64, 96),\n#         (5, 128, 128),\n#     ]\n\n#     dset = Sag_T2_Dataset_V24(data_root,\n#                               study_ids,\n#                               pred_sag_keypoints_infos_3d_t2,\n#                               test_series_descriptions_fn=test_series_descriptions_fn,\n#                               image_dir=image_dir,\n#                               crop_size_h=128,\n#                               crop_size_w=128,\n#                               cache_dir=CACHE_DIR,\n#                               other_crop_size_list=other_crop_size_list\n#                               )\n#     dloader = DataLoader(\n#         dset,\n#         batch_size=batch_size,\n#         shuffle=False,\n#         pin_memory=False,\n#         drop_last=False,\n#         num_workers=num_workers,\n#     )\n#     all_preds_t2 = {}\n#     sag_t2_preds = []\n#     for tensor_d in tqdm.tqdm(dloader):\n#         with torch.no_grad():\n#             for k in tensor_d.keys():\n#                 if k not in ['study_id']:\n#                     tensor_d[k] = tensor_d[k].to(device)\n#             study_id_list = tensor_d['study_id']\n#             cond = None  # tensor_d['cond']\n#             preds = None\n#             with autocast:\n#                 for i in range(len(models)):\n#                     # p = _forward_tta(models[i], tensor_d['img'], cond)\n#                     p = _forward_tta(models[i], tensor_d[data_keys[i]], cond)\n#                     if preds is None:\n#                         preds = (p.float() * weights[i])\n#                     else:\n#                         preds += (p.float() * weights[i])\n#             preds = preds / weights_sum\n#             bs, _ = preds.shape\n#             # bs, n_cond, 5_level, 3\n#             preds = preds.reshape(bs, -1, 5, 3).cpu()\n#             sag_t2_preds.append(preds)\n#             for b in range(bs):\n#                 study_id = int(study_id_list[b])\n#                 all_preds_t2[study_id] = preds[b]\n#     sag_t2_preds = torch.cat(sag_t2_preds, dim=0)\n#     # return all_preds_t2\n\n#     # axial\n#     model_cfgs = [\n#         {\n#             'sub_dir': 'v24_cond_axial/rexnetr_200.sw_in12k_ft_in1k_z_imgs_5_666',\n#             'backbone': 'rexnetr_200.sw_in12k_ft_in1k',\n#             'axial_in_channels': 3,\n#             'weight': 0.44292661444289017,\n#             'z_imgs': 5,\n#             'data_key': 'axial_imgs_5_128_128'\n#         },\n#         {\n#             'sub_dir': 'v24_cond_axial/convnext_tiny.in12k_ft_in1k_384_z_imgs_3_666',\n#             'backbone': 'convnext_tiny.in12k_ft_in1k',\n#             'axial_in_channels': 1,\n#             'weight': 0.29758735999956626,\n#             'z_imgs': 3,\n#             'data_key': 'axial_imgs_3_128_128'\n#         },\n\n#         {\n#             'sub_dir': 'v24_cond_axial/convnext_small.in12k_ft_in1k_384_z_imgs_3_666',\n#             'backbone': 'convnext_small.in12k_ft_in1k_384',\n#             'axial_in_channels': 1,\n#             'weight': 0.2384201128638661,\n#             'z_imgs': 3,\n#             'data_key': 'axial_imgs_3_128_128'\n#         },\n#         {\n#             'sub_dir': 'v24_cond_axial/densenet161_z_imgs_3_666',\n#             'backbone': 'densenet161',\n#             'axial_in_channels': 1,\n#             'weight': 0.021065912693677462,\n#             'z_imgs': 3,\n#             'data_key': 'axial_imgs_3_128_128'\n#         },\n\n#     ]\n#     v24_folds = list(range(5))\n#     if folds is not None:\n#         v24_folds = folds\n\n#     models = []\n#     data_keys = []\n#     weights = []\n#     for cfg in model_cfgs:\n#         sub_dir = cfg['sub_dir']\n#         backbone = cfg['backbone']\n#         data_key = cfg['data_key']\n#         axial_in_channels = cfg['axial_in_channels']\n#         w = cfg['weight']\n#         for fold in v24_folds:\n#             fn = os.path.join(model_dir, sub_dir)\n#             fn = fn + f'/best_wll_model_fold-{fold}_ema.pt'\n#             print('load: ', fn)\n#             model = Axial_HybridModel_24(backbone,\n#                                          backbone,\n#                                          pretrained=False,\n#                                          axial_in_channels=axial_in_channels\n#                                          )\n#             model.load_state_dict(torch.load(fn))\n#             model.to(device)\n#             model.eval()\n#             if is_parallel:\n#                 model = nn.DataParallel(model)\n\n#             models.append(model)\n#             weights.append(w)\n#             data_keys.append(data_key)\n#     weights_sum = np.sum(weights)\n\n#     with open(f'{keypoint_dir}/axial_3d_keypoints_v24.pkl', 'rb') as f:\n#         axial_3d_keypoints_v24 = pickle.load(f)\n\n#     other_crop_size_list = [\n#         (5, 128, 128),\n#     ]\n\n#     dset = Axial_Cond_Dataset_Multi_V24(data_root,\n#                                         study_ids,\n#                                         axial_3d_keypoints_v24,\n#                                         pred_sag_keypoints_infos_3d_t2,\n#                                         test_series_descriptions_fn=test_series_descriptions_fn,\n#                                         image_dir=image_dir,\n#                                         z_imgs=3,\n#                                         other_z_imgs_list=other_crop_size_list,\n#                                         cache_dir=CACHE_DIR,\n#                                         )\n\n#     dloader = DataLoader(\n#         dset,\n#         batch_size=batch_size,\n#         shuffle=False,\n#         pin_memory=False,\n#         drop_last=False,\n#         num_workers=num_workers,\n#     )\n#     all_preds_axial = {}\n#     axial_preds = []\n#     for tensor_d in tqdm.tqdm(dloader):\n#         with torch.no_grad():\n#             for k in tensor_d.keys():\n#                 if k not in ['study_id']:\n#                     tensor_d[k] = tensor_d[k].to(device)\n#             study_id_list = tensor_d['study_id']\n#             preds = None\n#             with autocast:\n#                 for i in range(len(models)):\n#                     p = models[i](tensor_d['sag_t2'], tensor_d[data_keys[i]])\n#                     if preds is None:\n#                         preds = (p.float() * weights[i])\n#                     else:\n#                         preds += (p.float() * weights[i])\n#             preds = preds / weights_sum\n#             bs, _ = preds.shape\n#             # bs, n_cond, 5_level, 3\n#             preds = preds.reshape(bs, -1, 5, 3).cpu()\n#             axial_preds.append(preds)\n#             for b in range(bs):\n#                 study_id = int(study_id_list[b])\n#                 all_preds_axial[study_id] = preds[b]\n\n#     axial_preds = torch.cat(axial_preds, dim=0)\n#     all_preds = torch.cat((sag_t2_preds, sag_t1_preds, axial_preds), dim=1)\n#     all_preds_dict = {}\n#     for i, study_id in enumerate(study_ids):\n#         all_preds_dict[study_id] = all_preds[i]\n#     return all_preds_dict","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:39:28.120566Z","iopub.execute_input":"2024-10-08T16:39:28.120902Z","iopub.status.idle":"2024-10-08T16:39:28.145549Z","shell.execute_reply.started":"2024-10-08T16:39:28.120873Z","shell.execute_reply":"2024-10-08T16:39:28.144583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer_all_666(data_root,\n              model_dir,\n              ):\n    df = pd.read_csv(f'{data_root}/{test_series_descriptions_fn}')\n    study_ids = df['study_id'].unique().tolist()\n    \n    device_0 = torch.device('cuda:0')\n\n    is_parallel = torch.cuda.device_count() > 1 and len(study_ids) >=2\n    print('is_parallel: ', is_parallel)\n    folds = None\n    \n    pred_dict_v24 = infer_v2_seed666(data_root, study_ids, model_dir,\n                              save_dir, device_0, N_WORKERS, folds, is_parallel)\n\n\n    pred_dict_v2 = infer_v24_seed666(data_root, study_ids, model_dir,\n                            save_dir, device_0, N_WORKERS, folds, is_parallel)\n\n   \n    print(len(study_ids))\n    fold_preds = []\n    for study_id in study_ids:\n        pred = pred_dict_v2[study_id]\n        fold_preds.append(pred.unsqueeze(0))\n    fold_preds = torch.cat(fold_preds, dim=0)\n    \n    print('fold_preds shape: ', fold_preds.shape)\n\n    fold_preds_v24 = []\n    for study_id in study_ids:\n        pred = pred_dict_v24[study_id]\n        fold_preds_v24.append(pred.unsqueeze(0))\n    fold_preds_v24 = torch.cat(fold_preds_v24, dim=0)\n\n    fold_preds_en = fold_preds\n    fold_preds_en = fold_preds\n    fold_preds_en[:, 0:1] = 0.55 * fold_preds[:, 0:1] + 0.45 * fold_preds_v24[:, 0:1]\n    fold_preds_en[:, 1:3] = 0.1 * fold_preds[:, 1:3] + 0.9 * fold_preds_v24[:, 1:3]\n    fold_preds_en[:, 3:5] = 0.38 * fold_preds[:, 3:5] + 0.62 * fold_preds_v24[:, 3:5]\n\n\n    fold_preds_en = nn.Softmax(dim=-1)(fold_preds_en).cpu().numpy()\n    return fold_preds_en","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:39:28.146783Z","iopub.execute_input":"2024-10-08T16:39:28.147056Z","iopub.status.idle":"2024-10-08T16:39:28.158128Z","shell.execute_reply.started":"2024-10-08T16:39:28.147034Z","shell.execute_reply":"2024-10-08T16:39:28.157356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#lhwcv_preds_666 = infer_all_666(data_root, model_dir2).reshape(-1, 3)\nlhwcv_preds = infer_all(data_root, model_dir).reshape(-1, 3)\n#lhwcv_preds = lhwcv_preds * 0.55 + lhwcv_preds_666 * 0.45","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:39:28.159023Z","iopub.execute_input":"2024-10-08T16:39:28.159297Z","iopub.status.idle":"2024-10-08T16:46:15.051493Z","shell.execute_reply.started":"2024-10-08T16:39:28.159274Z","shell.execute_reply":"2024-10-08T16:46:15.050408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -r ./cache/\n!rm -r ./keypoints_pred/","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:46:15.053566Z","iopub.execute_input":"2024-10-08T16:46:15.053990Z","iopub.status.idle":"2024-10-08T16:46:17.206919Z","shell.execute_reply.started":"2024-10-08T16:46:15.053949Z","shell.execute_reply":"2024-10-08T16:46:17.205472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. make submission","metadata":{}},{"cell_type":"code","source":"sub = pd.DataFrame()\nsub['row_id'] = row_names\nsub[LABELS] = lhwcv_preds\n\nsub.to_csv('submission.csv', index=False)\nsub.head(25)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:46:17.208712Z","iopub.execute_input":"2024-10-08T16:46:17.209045Z","iopub.status.idle":"2024-10-08T16:46:17.241815Z","shell.execute_reply.started":"2024-10-08T16:46:17.209011Z","shell.execute_reply":"2024-10-08T16:46:17.240834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"previous_submit_df = sub\nprevious_submit_df = previous_submit_df.set_index('row_id')\nprint('previous_submit_df', previous_submit_df.shape)\nprint(previous_submit_df.dtypes)\nprint('')\n###########################################\n# heng submission\n\ntry: \n    import natsort  \nexcept:\n    #%pip install natsort\n    !pip install natsort --no-index --find-links=file://///kaggle/input/heng-rnas2024-final-01/\n    \nimport sys\nimport os  \nsys.path.append('/kaggle/input/heng-rnas2024-final-01')\n\nfrom data import *\nimport pandas as pd\nfrom natsort import natsorted\nimport matplotlib\nimport matplotlib.pyplot as plt\n\npd.set_option('display.width', 1000)\n\nprint('IMPORT OK !!!!!!!!!!!!!')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:46:17.243174Z","iopub.execute_input":"2024-10-08T16:46:17.243534Z","iopub.status.idle":"2024-10-08T16:46:30.567097Z","shell.execute_reply.started":"2024-10-08T16:46:17.243500Z","shell.execute_reply":"2024-10-08T16:46:30.566063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"KAGGLE_DATA_DIR =\\\n    '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\n   # '/home/hp/work/2024/kaggle/rsna2024-lumbar-spine/data/kaggle/rsna-2024-lumbar-spine-degenerative-classification'\n   \n\nMODE = 'submit'  # local #submit\nDEVICE = 'cuda'  # 'cpu' 'cuda'\n\n\n\nif MODE == 'local':\n    IMAGE_DIR = f'{KAGGLE_DATA_DIR}/train_images'\n    valid_df = pd.read_csv(f'{KAGGLE_DATA_DIR}/train_series_descriptions.csv')\n\nif MODE == 'submit':\n    IMAGE_DIR = f'{KAGGLE_DATA_DIR}/test_images'\n    valid_df = pd.read_csv(f'{KAGGLE_DATA_DIR}/test_series_descriptions.csv')\n\nWEIGHT_DIR =\\\n    '/kaggle/input/heng-rnas2024-final-01'\n    #'/home/hp/work/2024/kaggle/rsna2024-lumbar-spine/code/team-share-002/final-01/010/scs_model_weight'\n\nscs_cfg = dotdict(\n    image_size=320,\n    arch='pvt_v2_b4',\n    checkpoint=[\n        f'{WEIGHT_DIR}/scs_model_weight/fold0-00008857-00010420-swa.pth',\n        f'{WEIGHT_DIR}/scs_model_weight/fold1-00008416-00009994-swa.pth',\n        f'{WEIGHT_DIR}/scs_model_weight/fold2-00012259-00013858-swa.pth',\n        f'{WEIGHT_DIR}/scs_model_weight/fold3-00009486-00014229-swa.pth',\n        f'{WEIGHT_DIR}/scs_model_weight/fold4-00006276-00012029-swa.pth',\n    ],\n\n    IMAGE_DIR= IMAGE_DIR,\n    MODE = MODE,\n    DEVICE = DEVICE,\n)\nnfn_cfg = dotdict(\n    image_size=320,\n    arch='pvt_v2_b4',\n    checkpoint=[\n        f'{WEIGHT_DIR}/nfn_model_weight/fold0-00032592.fix-name.pth',\n        f'{WEIGHT_DIR}/nfn_model_weight/fold1-00032144.pth',\n        f'{WEIGHT_DIR}/nfn_model_weight/fold2-00036524.pth',\n        f'{WEIGHT_DIR}/nfn_model_weight/fold3-00032619.pth',\n        f'{WEIGHT_DIR}/nfn_model_weight/fold4-00029488.pth',\n        f'{WEIGHT_DIR}/fold2-fix-flip-aug-00015086.pth',\n        f'{WEIGHT_DIR}/fold3-fix-flip-aug-000031047.pth',\n    ],\n\n    IMAGE_DIR= IMAGE_DIR,\n    MODE = MODE,\n    DEVICE = DEVICE,\n)\nprint(valid_df)\nprint('SETTING OK !!!!!!!!!!!!!')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:46:30.568785Z","iopub.execute_input":"2024-10-08T16:46:30.569822Z","iopub.status.idle":"2024-10-08T16:46:30.589368Z","shell.execute_reply.started":"2024-10-08T16:46:30.569774Z","shell.execute_reply":"2024-10-08T16:46:30.588577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from run_scs_submit import *\nfrom run_nfn_submit import *\n\n#scs_df = run_scs(valid_df, scs_cfg)\n#scs_result = run_scs(valid_df, scs_cfg)\n#scs_df = make_scs_submit(scs_result)\n\nnfn_result = run_nfn(valid_df, nfn_cfg)\nnfn_df = make_nfn_submit(nfn_result)\n\n#----\n# example to merge scs_df to you your submit csv\n\n#submit_df = make_dummy_submit(valid_df) #your submission csv\nsubmit_df = previous_submit_df\n\n#merge\n#submit_df.loc[scs_df.index, grade_col] = \\\n#    0.8*submit_df.loc[scs_df.index, grade_col]+ 0.2*scs_df\nsubmit_df.loc[nfn_df.index, grade_col] = \\\n    0.75*submit_df.loc[nfn_df.index, grade_col]+ 0.25*nfn_df\n\n\nsubmit_df = submit_df.reset_index(drop=False)\nsubmit_df.to_csv('submission.csv', index=False)\n\nprint('** FINAL SUBMIT **')\nprint(submit_df.head(30))\nprint(submit_df.shape)\n\nprint('SUBMIT OK!!!!!')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T16:46:30.590349Z","iopub.execute_input":"2024-10-08T16:46:30.590674Z","iopub.status.idle":"2024-10-08T16:46:55.614354Z","shell.execute_reply.started":"2024-10-08T16:46:30.590620Z","shell.execute_reply":"2024-10-08T16:46:55.613427Z"},"trusted":true},"execution_count":null,"outputs":[]}]}