{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":10794954,"sourceType":"datasetVersion","datasetId":6699410},{"sourceId":232113,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":198006,"modelId":219826},{"sourceId":241743,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":206501,"modelId":228247},{"sourceId":246068,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":210269,"modelId":231966}],"dockerImageVersionId":30840,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\nimport time\nimport random\nfrom math import ceil\nfrom pathlib import Path\nfrom collections import OrderedDict\n\nimport cv2\n#import zarr\nimport yaml\nimport timm\nimport torch\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nfrom torch.nn import DataParallel\n\n\nsys.path.append('/kaggle/input/czii-code-9972862/cryoet-main/src')\nfrom component_factory import create_optimizer, create_scheduler, create_criterion, create_metric\nfrom data_processing.dataset import SlidingWindowDataset #,build_infer_loader\nfrom utils import sigmoid\nfrom models import build_model_from_config\nfrom evaluation import pred_dicts_to_df\nimport constants\n\n\nfrom losses import from_cfg\n\nfrom logger import Logger, EmaCalculator\n\n#from models.postprocessing import PostProcessor\nfrom utils.loop import to_device\nfrom models.postprocessing import PostProcessor\nfrom data_processing.dataset import dataset #,build_infer_loader","metadata":{"_uuid":"eced986d-7767-4464-9f4c-84889cd9ee6d","_cell_guid":"bd7a6bf4-7121-4ed9-88fe-2e1f8f5d754c","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.029254Z","iopub.execute_input":"2025-02-19T14:35:38.029598Z","iopub.status.idle":"2025-02-19T14:35:38.035303Z","shell.execute_reply.started":"2025-02-19T14:35:38.029565Z","shell.execute_reply":"2025-02-19T14:35:38.034467Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ls /kaggle/input/czii-code-9972862/cryoet-main/src","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-19T14:35:37.886195Z","iopub.execute_input":"2025-02-19T14:35:37.886469Z","iopub.status.idle":"2025-02-19T14:35:38.028107Z","shell.execute_reply.started":"2025-02-19T14:35:37.886424Z","shell.execute_reply":"2025-02-19T14:35:38.027334Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Prepare DataFrame","metadata":{"_uuid":"b838341b-69e2-44a2-bf94-66f749b9718a","_cell_guid":"3af2e024-1bb6-48a1-8b0e-1bf1d806f98f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def create_infer_df(\n        save_path = '../data/processed/train_df.csv',\n        base_path = Path('../data/original'),\n        crop_size = (128, 128, 128),\n        crop_stride = (56, 64, 64),\n        resolution_hierarchy = 0,\n    ):\n    \"\"\"\n    This function will create a train_df.csv file that will be used to generate the dataset.\n    This may create a new fold_df.csv file if it does not exist.\n    \"\"\"\n    if not isinstance(base_path, Path):\n        base_path = Path(base_path)\n    if not isinstance(save_path, Path):\n        save_path = Path(save_path)\n    \n    #fold_df = create_folds()\n    #fold_df.set_index('experiment_id', inplace=True)\n\n    train_df = {\n        \"experiment_id\": [],\n        \"zarr_path\": [],\n        \"crop_origin_d0\": [],\n        \"crop_origin_d1\": [],\n        \"crop_origin_d2\": [],\n        \"source\": [],\n        #\"fold\": [],\n    }\n    arr_shape = (184, 630, 630)\n    n_crops_0 = ceil((arr_shape[0] - crop_size[0]) / crop_stride[0]) + 1\n    n_crops_1 = ceil((arr_shape[1] - crop_size[1]) / crop_stride[1]) + 1\n    n_crops_2 = ceil((arr_shape[2] - crop_size[2]) / crop_stride[2]) + 1\n    crop_origins_d0 = [ii * crop_stride[0] for ii in range(n_crops_0-1)] + [arr_shape[0] - crop_size[0]]\n    crop_origins_d1 = [ii * crop_stride[1] for ii in range(n_crops_1-1)] + [arr_shape[1] - crop_size[1]]\n    crop_origins_d2 = [ii * crop_stride[2] for ii in range(n_crops_2-1)] + [arr_shape[2] - crop_size[2]]\n    for experiment_id in os.listdir(base_path / 'test' / 'static' / 'ExperimentRuns'):\n        zarr_path = base_path/'test'/'static'/'ExperimentRuns'/experiment_id/'VoxelSpacing10.000'/'denoised.zarr'# only used 'denoised.zarr' as only this is available on test data\n        \n        for i in range(n_crops_0):\n            crop_start_0 = crop_origins_d0[i]\n            for j in range(n_crops_1):\n                crop_start_1 = crop_origins_d1[j]\n                for k in range(n_crops_2):\n                    crop_start_2 = crop_origins_d2[k]\n                    train_df['experiment_id'].append(experiment_id)\n                    train_df['zarr_path'].append(str(zarr_path))\n                    train_df['crop_origin_d0'].append(crop_start_0)\n                    train_df['crop_origin_d1'].append(crop_start_1)\n                    train_df['crop_origin_d2'].append(crop_start_2)\n                    train_df['source'].append('original')\n                    #train_df['fold'].append(fold_df.loc[experiment_id, 'fold'])\n    train_df = pd.DataFrame(train_df)\n    train_df.to_csv(save_path, index=False)","metadata":{"_uuid":"026edef3-184b-42ba-983a-1f9792902bc4","_cell_guid":"ec73ad36-33df-408b-ad7e-a1b090404f8f","trusted":true,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.036259Z","iopub.execute_input":"2025-02-19T14:35:38.036565Z","iopub.status.idle":"2025-02-19T14:35:38.047721Z","shell.execute_reply.started":"2025-02-19T14:35:38.036537Z","shell.execute_reply":"2025-02-19T14:35:38.046990Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"create_infer_df(save_path = '/kaggle/working/infer_df.csv', base_path = '/kaggle/input/czii-cryo-et-object-identification')\ninfer_df = pd.read_csv('/kaggle/working/infer_df.csv')","metadata":{"_uuid":"acb13673-004c-4a4d-b292-0d57f240b4f3","_cell_guid":"abcda028-3b63-4603-bb19-18de384da9c6","trusted":true,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.048395Z","iopub.execute_input":"2025-02-19T14:35:38.048662Z","iopub.status.idle":"2025-02-19T14:35:38.071768Z","shell.execute_reply.started":"2025-02-19T14:35:38.048642Z","shell.execute_reply":"2025-02-19T14:35:38.070819Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collate_fn(batch):\n    out = tuple(zip(*batch))\n    images, targets = out\n    images = torch.stack(images)\n    return images, targets","metadata":{"_uuid":"d8927d72-4237-4ffd-8b12-b72fd0d89232","_cell_guid":"7dd18268-f30b-49d5-aa22-65471d73aec4","trusted":true,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.072627Z","iopub.execute_input":"2025-02-19T14:35:38.072864Z","iopub.status.idle":"2025-02-19T14:35:38.076626Z","shell.execute_reply.started":"2025-02-19T14:35:38.072840Z","shell.execute_reply":"2025-02-19T14:35:38.075702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Dataset\nfrom typing import Union, Tuple\nclass SlidingWindowDataset(Dataset):\n    def __init__(self,\n            index,df,image_path: Union[Path, str] = Path('../data/original'),return_targets=False,\n            do_augmentation: bool = False, augmentation_args: dict = {}, aurgementation_probabilities: dict = {}, path_modifier = None,\n            # dataset specific arguments\n            external_path: Union[Path, str] = Path('../data/external'),\n            image_size: Tuple[int, int, int] = (128, 128, 128), resolution_hierarchy: int = 0,\n            # input preprocessor\n            rescale_factor: float = 1e5, clip_percentile = (0.1, 99.9), standardize = False,\n            # target generator\n            kernel_size_multiplier: int = 7, kernel_sigma_multiplier: float = 0.5,\n            detr: bool = False, return_num_points: bool = False,\n            epoch=None, domain_blend_params=None\n\n        ):\n        if detr or return_num_points:\n            raise ValueError('SlidingWindowDataset does not support DETR or return_num_points')\n        image_sizes = {\n            0: (184, 630, 630),\n            1: (92, 315, 315),\n            2: (46, 157, 157),\n        }\n        self.inner_df = df[\n            (df['crop_origin_d0'] == 0) & (df['crop_origin_d1'] == 0) & (df['crop_origin_d2'] == 0)]\n        #self.inner_df.set_index('experiment_id', inplace=True, drop=False)\n        #print(self.inner_df)\n        self.inner_df_index = self.inner_df.index.values\n        self.experiment_ids = list(self.inner_df['experiment_id'])\n        self.ds = dataset(\n            self.inner_df_index, self.inner_df, image_path, return_targets, False, {}, {}, path_modifier,\n            external_path, image_sizes[resolution_hierarchy], resolution_hierarchy,\n            rescale_factor, clip_percentile, standardize,\n            kernel_size_multiplier, kernel_sigma_multiplier, False, False,\n            epoch=epoch, domain_blend_params=domain_blend_params\n        )\n        self.image_size = np.asarray(image_size)\n        self.index = index\n\n        self.cached_images = {}\n        self.num_access = {}\n        self.df = df\n        self.tiles_per_experiment = len(self.df[self.df['experiment_id'] == self.df['experiment_id'].iloc[0]])\n    \n    def __getitem__(self, idx):\n        idx = self.index[idx]\n        entry = self.df.loc[idx]\n        experiment_id = entry['experiment_id'] \n        if experiment_id not in self.cached_images:\n            #print(self.ds[idx])\n            image, target = self.ds[self.experiment_ids.index(experiment_id)]\n            self.cached_images[experiment_id] = image\n            if self.image_size[0] > 184:\n                self.cached_images[experiment_id] = torch.nn.functional.pad(self.cached_images[experiment_id], (\n                    0, 0,\n                    0, 0,\n                    0, self.image_size[0] - 184,\n                ))\n            self.num_access[experiment_id] = 0\n\n        crop_start = entry['crop_origin_d0'], entry['crop_origin_d1'], entry['crop_origin_d2']\n        image = self.cached_images[experiment_id][\n            :,\n            crop_start[0]:crop_start[0]+self.image_size[0],\n            crop_start[1]:crop_start[1]+self.image_size[1],\n            crop_start[2]:crop_start[2]+self.image_size[2]\n        ]\n        self.num_access[experiment_id] += 1\n        if self.num_access[experiment_id] >= self.tiles_per_experiment:\n            del self.cached_images[experiment_id]\n            del self.num_access[experiment_id]\n        return image, None\n    def __len__(self):\n        return len(self.index)\n    @staticmethod\n    def collate_fn(batch):\n        out = tuple(zip(*batch))\n        images, targets = out\n        images = torch.stack(images)\n        if targets[0] is None:\n            return images, None\n        targets = {k: torch.stack([t[k] for t in targets]) for k in targets[0].keys()}\n        return images, targets","metadata":{"_uuid":"ec44e4ca-a1b8-40cf-a6f8-cda8f47c78df","_cell_guid":"aac83c7c-3b9c-407c-8c73-e9b64397e5c4","trusted":true,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.077323Z","iopub.execute_input":"2025-02-19T14:35:38.077572Z","iopub.status.idle":"2025-02-19T14:35:38.090233Z","shell.execute_reply.started":"2025-02-19T14:35:38.077545Z","shell.execute_reply":"2025-02-19T14:35:38.089536Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_infer_loader(infer_df, cfg):\n    infer_ds = SlidingWindowDataset(\n        index = infer_df.index.values,\n        df = infer_df,\n        image_path = cfg['data_path'] + '/test',\n        return_targets=False,\n        **cfg['feature_extractor_args']\n    )\n    infer_dl = torch.utils.data.DataLoader(\n        infer_ds,\n        batch_size = cfg['val_batch_size'],\n        shuffle = False,\n        num_workers = 0,\n        #prefetch_factor = 2,\n        collate_fn = collate_fn,\n        pin_memory = True\n    )\n    return infer_dl","metadata":{"_uuid":"2924952b-2613-4d96-9c50-50516f229dd3","_cell_guid":"fac216d5-e287-4bae-b452-054850791ffa","trusted":true,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.090977Z","iopub.execute_input":"2025-02-19T14:35:38.091182Z","iopub.status.idle":"2025-02-19T14:35:38.104011Z","shell.execute_reply.started":"2025-02-19T14:35:38.091165Z","shell.execute_reply":"2025-02-19T14:35:38.103327Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Configuration","metadata":{"_uuid":"90002b32-bd8b-445f-8090-d9203e5d2372","_cell_guid":"0d4512c7-8aca-494b-9b9a-53b9ad28909e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"experiment_names = [\n    #'46_resnet50d.ra2_in1k_20241215',\n    #'58_maxvit_tiny_tf_224.in1k_20241218',\n    #'153_resnet50d.ra2_in1k_20250106',\n    #'130_resnet50d.ra2_in1k_20250102',\n    #'86_resnet50d.ra2_in1k_20241222',\n    #'184_resnet50d.ra2_in1k_20250111',\n    #'192_resnet50d.ra2_in1k_20250112',\n    #'196_resnet50d.ra2_in1k_20250113',\n    '209_resnet50d.ra2_in1k_20250116',\n    #'242_resnet50d.ra2_in1k_20250123',\n    #'219_resnet50d.ra2_in1k_20250118',\n    #\"237_tf_efficientnetv2_s.in21k_ft_in1k_20250122\",\n    #\"242_resnet50d.ra2_in1k_20250123\",\n    #\"247_resnet50d.ra2_in1k_20250124\",\n    #\"246_resnet50d.ra2_in1k_20250124\",\n    \"249_tf_efficientnetv2_m.in21k_ft_in1k_20250125\",\n    #\"251_resnet50d.ra2_in1k_20250126\",\n    #\"255_resnet50d.ra2_in1k_20250128\",\n    #\"256_resnet50d.ra2_in1k_20250128\",\n    #\"264_tf_efficientnetv2_m.in21k_ft_in1k_20250130\",\n    \"265_resnet50d.ra2_in1k_20250130\",\n    #\"266_maxvit_tiny_tf_224.in1k_20250130\",\n    #\"277_tf_efficientnetv2_m.in21k_ft_in1k_20250202\",\n    #\"283_resnet50d.ra2_in1k_20250203\",\n    #\"284_resnet50d.ra2_in1k_20250203\",\n    #\"290_tf_efficientnetv2_m.in21k_ft_in1k_20250204\",\n    #\"291_resnet50d.ra2_in1k_20250205\",\n    #\"293_resnet50d.ra2_in1k_20250205\",\n]\nsoup=False\nif not soup:\n    weight_paths = [\n        [\n        f'/kaggle/input/{e}/pytorch/default/1/{e}/weights/fold_{i}/best.pth' for i in range(4)\n        ] for e in experiment_names\n    ]\n    model_config_paths = [\n        [\n        f'/kaggle/input/{e}/pytorch/default/1/{e}/config.yml' for i in range(4)\n        ] for e in experiment_names\n    ]\nelse:\n    weight_paths = [\n        [\n        f'/kaggle/input/{e}/pytorch/default/1/{e}/weights/fold_{i}/best.pth' for i in range(4)\n        for e in experiment_names]\n    ]\n    model_config_paths = [\n        [\n        f'/kaggle/input/{e}/pytorch/default/1/{e}/config.yml' for i in range(4)\n        for e in experiment_names]\n    ]","metadata":{"_uuid":"72fb60dd-23fb-4c4b-8cd2-d39400e2b484","_cell_guid":"faef2bef-cfe7-415c-9666-15541c493c1f","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.104784Z","iopub.execute_input":"2025-02-19T14:35:38.105074Z","iopub.status.idle":"2025-02-19T14:35:38.121506Z","shell.execute_reply.started":"2025-02-19T14:35:38.105046Z","shell.execute_reply":"2025-02-19T14:35:38.120716Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Transformations=[\n    'e', # no transformation\n    'b', #flip1\n    'ba2', # flip2\n    'a2', # flip12,\n    #'a', # rot12\n    #'a3', # rot12*3\n    #'ba', # flip1, rot12\n    #'ba3', # flip1, rot12*3\n    #'c', # flip0\n    #'cb', # flip0, flip1\n    #'cba2', # flip0, flip2\n    #'ca2', # flip0, flip12,\n    #'ca', # flip0, rot12\n    #'ca3', # flip0, rot12*3\n    #'cba', # flip0, flip1, rot12\n    #'cba3', # flip0, flip1, rot12*3\n]","metadata":{"_uuid":"2275dc9a-6251-4d22-ba21-fd12ef220189","_cell_guid":"602489c2-250f-4ed7-bea7-fa43ac20bdad","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.124440Z","iopub.execute_input":"2025-02-19T14:35:38.124678Z","iopub.status.idle":"2025-02-19T14:35:38.139066Z","shell.execute_reply.started":"2025-02-19T14:35:38.124660Z","shell.execute_reply":"2025-02-19T14:35:38.138339Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"_uuid":"fdd22932-a2e5-4021-837a-55ed7df78867","_cell_guid":"eaf753bd-de28-4f11-91b0-cbc41ae41303","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.140092Z","iopub.execute_input":"2025-02-19T14:35:38.140361Z","iopub.status.idle":"2025-02-19T14:35:38.153739Z","shell.execute_reply.started":"2025-02-19T14:35:38.140340Z","shell.execute_reply":"2025-02-19T14:35:38.152934Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"eed1a431-da80-43b4-b995-cb717ddc5ac8","_cell_guid":"15d613ab-6d78-4ef5-afad-63bef431d895","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!nvidia-smi","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.154575Z","iopub.execute_input":"2025-02-19T14:35:38.154789Z","iopub.status.idle":"2025-02-19T14:35:38.406180Z","shell.execute_reply.started":"2025-02-19T14:35:38.154763Z","shell.execute_reply":"2025-02-19T14:35:38.405125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# set random seed for everything\ndef seed_everything(seed):\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\nconfig_file = '/kaggle/input/265_resnet50d.ra2_in1k_20250130/pytorch/default/1/265_resnet50d.ra2_in1k_20250130/config.yml'\nwith open(config_file, encoding='utf-8')as f:\n    cfg = yaml.safe_load(f)\n    \ncfg['model']['pretrained'] = False\ncfg['base_path'] = \"/kaggle/input/czii-cryo-et-object-identification\"\ncfg['train_df_path'] = '/kaggle/working/infer_df.csv'\ncfg['data_path'] = \"/kaggle/input/czii-cryo-et-object-identification\"\n\nTEST = False\nADD_FP = False\nDOUBLE_PRED = False\n#PARTICLES_TO_TEST = 'beta-galactosidase'\n#PARTICLES_TO_TEST = 'thyroglobulin'\n#PARTICLES_TO_TEST = 'ribosome'\n#PARTICLES_TO_TEST = 'apo-ferritin'\nPARTICLES_TO_TEST = 'virus-like-particle'\n\nEXPORT_TO_ONNX = False\n\nUSE_TENSORRT = False\n\nOFFSET=False\n\n# general postprocessing settings\ncfg['postprocessing']['weight_center'] = True\n#cfg['postprocessing']['edge_begin'] = 0.20\n#cfg['postprocessing']['min_weight'] = 0.1\ncfg['postprocessing']['erosion_after_threshold'] = 0\n#cfg[\"postprocessing\"][\"threshold\"] = [0.5, 0.0, 0.0, 0.4, 0.5]\ncfg[\"postprocessing\"][\"threshold\"] = [0.3, 0.2, -0.2, -0.2, -0.2]\n#cfg[\"postprocessing\"][\"threshold\"] = 0.0#[0.3, 0.2, -0.2, 0.0, -0.2] #[0.7, 0.7, 0.5, 0.7, 0.5]\ncfg['postprocessing']['gaussian_blur_size'] = 5\ncfg['postprocessing']['gaussian_blur_sigma'] = 1.0#[0.5, 1.5, 2.0, 1.5, 1.5]\ncfg['postprocessing']['method'] = 'pooling'\ncfg['postprocessing']['method_args'] = {'ksize': 3}\n\n# window size\ncfg[\"feature_extractor_args\"][\"image_size\"] = (192, 128, 128)\ninfer_df = infer_df[infer_df[\"crop_origin_d0\"] == 0]\n\n# GPU\ncfg['amp_dtype'] = 'fp16'\ncfg['val_batch_size'] = 4\n\n# ensemble settings\ncfg['ensemble'] = {}\ncfg['ensemble']['method'] = 'avg'\ncfg[\"postprocessing\"]['wbf_radius_multiplier'] = 0.05\ncfg['ensemble']['radius_multiplier'] = 0.5\ncfg['ensemble']['min_votes'] = 0\n\npseudo_label_epochs = 0\n\n#torch.backends.cudnn.benchmark = True\n\nif 'pretrain_weight_path' in cfg['model']:\n    del cfg['model']['pretrain_weight_path']\n    \nif cfg['amp_dtype'] == 'fp16':\n    AMP_DTYPE = torch.float16\nelif cfg['amp_dtype'] == 'bf16':\n    AMP_DTYPE = torch.bfloat16\nelif cfg['amp_dtype'] == 'fp32':\n    AMP_DTYPE = torch.float32\n    \nseed_everything(cfg['seed'])\nBASE_PATH = cfg['base_path']\n\n\n# The outer loop is ensemble with separate models\n# inner loop is for ensemble with model weight averaging\n\n\nmodels = []\n\nif not USE_TENSORRT:\n    for exp, ps, mcfgs in zip(experiment_names, weight_paths, model_config_paths):\n        models_i = []\n        for p, mcfg in zip(ps, mcfgs):\n            with open(mcfg, encoding='utf-8') as f:\n                model_cfg = yaml.safe_load(f)[\"model\"]\n            model_cfg['pretrained'] = False\n            if 'ema' in model_cfg:\n                del model_cfg['ema']\n            if 'pretrain_weight_path' in model_cfg:\n                del model_cfg['pretrain_weight_path']\n            model = build_model_from_config(model_cfg)\n            state_dict = torch.load(p, map_location=cfg['device'])\n            if '_orig_mod.segmentation_head.weight' in state_dict:\n                state_dict = {k.replace('_orig_mod.', ''): v for k, v in state_dict.items()}\n            if 'module.segmentation_head.weight' in state_dict:\n                state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}\n            model.load_state_dict(state_dict)\n            del state_dict\n            model.eval()\n            models_i.append(model)\n        model = models_i[0]\n        for k in model.state_dict():\n            model.state_dict()[k] = sum([m.state_dict()[k] for m in models_i])/len(models_i)\n        #model = torch.compile(model)\n    \n        if EXPORT_TO_ONNX:\n            dummy_input = torch.randn(4, 1, 192, 128, 128, device='cuda', dtype=torch.float16)\n            with torch.no_grad():\n                with torch.autocast(device_type=cfg['device'], dtype=AMP_DTYPE):\n                    torch.onnx.export(\n                        model,\n                        dummy_input,\n                        f'{exp}.onnx',\n                        input_names=['input'],\n                        output_names=['heatmap'],\n                        opset_version=11\n                    )\n        model = DataParallel(model)\n        models.append(model)\n        torch.cuda.empty_cache()\nelse:\n    for exp in experiment_names:\n        engine = TRTInference(f'/kaggle/input/cryoet-trts/{exp.split(\"_\")[0]}.exp')\n\nif EXPORT_TO_ONNX:\n    exit()\n\n# data preparation\ninfer_dl = build_infer_loader(infer_df, cfg)","metadata":{"_uuid":"c9c8f7c9-d214-481b-9f4e-c61efba041d8","_cell_guid":"fa074e2a-918a-43f1-85be-283b6c8c937b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-02-19T14:35:38.407510Z","iopub.execute_input":"2025-02-19T14:35:38.407873Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Pseudo labelling","metadata":{"_uuid":"411c5e98-d1d0-4b3f-9bdc-f63f50c05e6d","_cell_guid":"cee247ad-31a1-44c8-9820-a9a4980de91d","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"_uuid":"30af63d5-ac48-4b8e-a2da-a93e19b56fcf","_cell_guid":"3498801d-9f36-42a3-944a-560acf2e698e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!nvidia-smi","metadata":{"_uuid":"b59013f4-ce2e-4201-8740-a7d8ca44fc59","_cell_guid":"073b3e0d-62f3-40df-b3d5-d4e9a7bdee09","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copy","metadata":{"_uuid":"10b7bdf8-0b08-4e29-b2c6-ce1fbec3001b","_cell_guid":"54bfadbb-a6e6-43bd-a4ce-adc5525a2e3d","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Run inference","metadata":{"_uuid":"eb09cfc8-77a7-4797-9229-c3508310c11a","_cell_guid":"6762f40d-641c-466a-9882-3b7d2859fb5c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"transformations = {\n    'e': lambda x: x,\n    'b': lambda x: x.flip(dims=(-2,)),\n    'ba2': lambda x: x.flip(dims=(-1,)),\n    'a2': lambda x: x.flip(dims=(-1, -2)),\n    'a': lambda x: x.rot90(dims=(-1, -2), k=1),\n    'a3': lambda x: x.rot90(dims=(-1, -2), k=3),\n    'ba': lambda x: x.flip(dims=(-2,)).rot90(dims=(-1, -2), k=1),\n    'ba3': lambda x: x.flip(dims=(-2,)).rot90(dims=(-1, -2), k=3),\n    'c': lambda x: x.flip(dims=(-3,)),\n    'cb': lambda x: x.flip(dims=(-3, -2)),\n    'cba2': lambda x: x.flip(dims=(-3, -1)),\n    'ca2': lambda x: x.flip(dims=(-3, -2, -1)),\n    'ca': lambda x: x.flip(dims=(-3,)).rot90(dims=(-2, -1), k=1),\n    'ca3': lambda x: x.flip(dims=(-3,)).rot90(dims=(-2, -1), k=3),\n    'cba': lambda x: x.flip(dims=(-3, -2)).rot90(dims=(-1, -2), k=1),\n    'cba3': lambda x: x.flip(dims=(-3, -2)).rot90(dims=(-1, -2), k=3),\n}\n\nreverse_transformations = {\n    'e': 'e',\n    'b': 'b',\n    'ba2': 'ba2',\n    'a2': 'a2',\n    'a': 'a3',\n    'a3': 'a',\n    'ba': 'ba',\n    'ba3': 'ba3',\n    'c': 'c',\n    'cb': 'cb',\n    'cba2': 'cba2',\n    'ca2': 'ca2',\n    'ca': 'ca3',\n    'ca3': 'ca',\n    'cba': 'cba',\n    'cba3': 'cba3',\n}\n\ndef infer_transformerd(model, images, transformation):\n    out = model(transformations[transformation](images))\n    if isinstance(out, tuple):\n        out = (\n            transformations[reverse_transformations[transformation]](out[0]),\n            transformations[reverse_transformations[transformation]](out[1])\n        )\n        if out[0].shape[2] > 184:\n            out = (\n                out[0][:, :, :184, :, :],\n                out[1][:, :, :184, :, :]\n            )\n    else:\n        out = transformations[reverse_transformations[transformation]](out)\n        if out.shape[2] > 184:\n            out = out[:, :, :184, :, :]\n    \n    return out\n\n\n\ndef infer(models, images, TTA):\n    out = 0\n    with torch.inference_mode():\n    #with torch.no_grad():\n        for model in models:\n            with torch.autocast(device_type=cfg['device'], dtype=AMP_DTYPE):\n                for t in TTA:\n                    out += infer_transformerd(model, images, t)\n        out /= len(models)\n        out /= len(TTA)\n    return out\n\ndef infer_offset(models, images, TTA):\n    out = [0, 0]\n    with torch.inference_mode():\n        for model in models:\n            with torch.autocast(device_type=cfg['device'], dtype=AMP_DTYPE):\n                for t in TTA:\n                    out_tup = infer_transformerd(model, images, t)\n                    out[0] += out_tup[0]\n                    out[1] += out_tup[1]\n        out[0] /= len(models)\n        out[1] /= len(models)\n        out[0] /= len(TTA)\n        out[1] /= len(TTA)\n    return tuple(out)","metadata":{"_uuid":"efa49c82-4746-41e0-aa47-03e564d351b3","_cell_guid":"e23c7110-bb50-43eb-b49e-7a279edd04a5","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def TTA_avg(models, pbar, TTA, cfg, AMP_DTYPE):\n    postprocess_size = list(cfg['feature_extractor_args']['image_size'])\n    postprocess_size[0] = min(184, postprocess_size[0])\n    classes=[\n            'apo-ferritin',\n            # 'beta-amylase', # ignore beta-amylase for now as it is deemed impossible to detect\n            'beta-galactosidase',\n            'ribosome',\n            'thyroglobulin',\n            'virus-like-particle'\n        ]\n    post_processor = PostProcessor(\n        classes = classes,\n        #tiles_per_experiment = 162,\n        tiles_per_experiment = (infer_df[\"experiment_id\"] == infer_df[\"experiment_id\"][0]).sum(),\n        window_size=postprocess_size,\n        **cfg['postprocessing'],\n        ignore_uncovered=True,\n        #keep_heatmaps=True\n    )\n    time_postrocess = 0\n    time_inference = 0\n    start = time.time()\n    for i, (images, targets) in enumerate(pbar):\n        start_inference = time.time()\n        if OFFSET:\n            out = infer_offset(models, images, TTA)\n        else:\n            out = infer(models, images, TTA)\n        time_inference += time.time() - start_inference\n\n        start_postprocess = time.time()\n        crop_origins = infer_df.iloc[i * cfg['val_batch_size']: (i + 1) * cfg['val_batch_size'], [2, 3, 4]].values\n        experiment_ids = infer_df.iloc[i * cfg['val_batch_size']: (i + 1) * cfg['val_batch_size'], 0].values\n        post_processor.accumulate(out, crop_origins, experiment_ids)\n        time_postrocess += time.time() - start_postprocess\n    assert post_processor.accumulated_data == {}\n    end = time.time()\n    return post_processor.predictions, end-start, time_inference, time_postrocess\n\nfrom models.postprocessing import weighted_box_fusion\ndef TTA_wbf(models, pbar, TTA, cfg, AMP_DTYPE):\n    postprocess_size = list(cfg['feature_extractor_args']['image_size'])\n    postprocess_size[0] = min(184, postprocess_size[0])\n    classes=[\n            'apo-ferritin',\n            # 'beta-amylase', # ignore beta-amylase for now as it is deemed impossible to detect\n            'beta-galactosidase',\n            'ribosome',\n            'thyroglobulin',\n            'virus-like-particle'\n        ]\n    post_processors = [\n        PostProcessor(\n            classes = classes,\n            #tiles_per_experiment = 162,\n            tiles_per_experiment = (infer_df[\"experiment_id\"] == infer_df[\"experiment_id\"][0]).sum(),\n            window_size=postprocess_size,\n            **cfg['postprocessing'],\n            ignore_uncovered=True,\n            #keep_heatmaps=True\n        ) for _ in range(len(models))\n    ]\n    time_postprocess = 0\n    time_inference = 0\n    start = time.time()\n    for i, (images, targets) in enumerate(pbar):\n        #print(images.shape)\n        #images, targets = to_device(images, targets, cfg['device'])\n        images = images.to(cfg['device'])\n\n        start_inference = time.time()\n        outs = [0 for _ in range(len(models))]\n        for mi, model in enumerate(models):\n            if OFFSET:\n                out = infer_offset([model], images, TTA)\n            else:\n                out = infer([model], images, TTA)\n            outs[mi] = out\n        outs[0][:, 1] = 0\n        time_inference += time.time() - start_inference\n\n        start_postprocess = time.time()\n        crop_origins = infer_df.iloc[i * cfg['val_batch_size']: (i + 1) * cfg['val_batch_size'], [2, 3, 4]].values\n        experiment_ids = infer_df.iloc[i * cfg['val_batch_size']: (i + 1) * cfg['val_batch_size'], 0].values\n        for out, post_processor in zip(outs, post_processors, strict=True):\n            post_processor.accumulate(out, crop_origins, experiment_ids)\n        time_postprocess += time.time() - start_postprocess\n    \n    preds = {}\n    # exp_id -> class -> points\n    #       |-> heatmap\n    all_predictions = [pp.predictions for pp in post_processors]\n    for exp_id in all_predictions[0]:\n        preds[exp_id] = {}\n        for c in classes:\n        #for c in preds[exp_id]:\n            if c == 'beta-amylase':\n                continue\n            preds[exp_id][c] = {\n                'points': np.vstack([pp[exp_id][c]['points'] for pp in all_predictions]),\n                'confidence': np.concatenate([pp[exp_id][c]['confidence'] for pp in all_predictions])\n            } \n    for exp_id in preds:\n        for c in classes:\n            if c == 'beta-amylase':\n                continue\n            if len(preds[exp_id][c]['points']) == 0:\n                continue\n            preds[exp_id][c] = weighted_box_fusion(\n                preds[exp_id][c],\n                constants.particle_radius[c] * cfg['ensemble']['radius_multiplier'],\n                cfg['ensemble']['min_votes'],\n            )\n    end = time.time()\n        \n    return preds, end-start, time_inference, time_postprocess","metadata":{"_uuid":"28d020b7-79bb-4bff-a28e-f74ecd2f411c","_cell_guid":"ea665b0e-fff8-40bb-a2bb-16dfca0d18b6","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pbar = infer_dl\nif cfg['ensemble']['method'] == 'avg':\n    preds, total_time, total_time_inference, total_time_postprocess = TTA_avg(models, pbar, Transformations, cfg, AMP_DTYPE)\nelif cfg['ensemble']['method'] == 'wbf':\n    preds, total_time, total_time_inference, total_time_postprocess = TTA_wbf(models, pbar, Transformations, cfg, AMP_DTYPE)\n\nprint(f'time taken to NN inference: {total_time_inference} seconds')\nprint(f'time taken to post-processing: {total_time_postprocess} seconds')\nprint(f'(time taken to NN infer) / (time taken to post-process) = {total_time_inference / total_time_postprocess}')\n\nprint(f'time taken to process {len(infer_dl) * cfg[\"val_batch_size\"]} tiles: {total_time} seconds.')\nprint(f'this roughly translates to {((81 * 500) / (len(infer_dl) * cfg[\"val_batch_size\"])) * (total_time)} seconds on the actual test data.')\n\ndummy_FP = np.array([\n    [    0,    0,    0],\n]*40)\nif ADD_FP:\n    for exp_id in preds:\n        preds[exp_id][PARTICLES_TO_TEST]['points'] = np.vstack([\n            preds[exp_id][PARTICLES_TO_TEST]['points'],\n            dummy_FP\n        ]) if len(preds[exp_id][PARTICLES_TO_TEST]['points']) > 0 else dummy_FP\n\nif DOUBLE_PRED:\n    for exp_id in preds:\n        preds[exp_id][PARTICLES_TO_TEST]['points'] = np.vstack([\n            preds[exp_id][PARTICLES_TO_TEST]['points'],\n            preds[exp_id][PARTICLES_TO_TEST]['points']\n        ]) if len(preds[exp_id][PARTICLES_TO_TEST]['points']) > 0 else preds[exp_id][PARTICLES_TO_TEST]['points']\npred_sub_df = pred_dicts_to_df(preds)\nif TEST:\n    pred_sub_df = pred_sub_df[pred_sub_df['particle_type'] == PARTICLES_TO_TEST]\npred_sub_df.to_csv('submission.csv', index=False)","metadata":{"_uuid":"4b4ea5b3-64fa-477b-b49b-2dfd999201ea","_cell_guid":"423ddaa4-14b9-4fe5-ba27-3e2408be6c05","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(pred_sub_df)","metadata":{"_uuid":"1eefbd04-fcc9-4e10-aa3a-f96369aeac1b","_cell_guid":"16153865-97fe-483f-b109-e55c058e3ac2","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"74de2801-d0ad-44d6-b86c-dc4977ff4d9a","_cell_guid":"33a00ef0-c7df-4d20-b3f2-19ffcfea8142","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"af9f1c03-84dc-4949-b1c2-ae4b606128b8","_cell_guid":"2c0a29ce-4e73-4574-b6a6-cb72f484f965","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"32740d0d-ffbe-45f2-8d91-9e6f09d90ca7","_cell_guid":"a0dc6abe-c102-4b95-bed2-2efa77988c70","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"309a78ff-36c1-45fc-96ef-aae301f6eb05","_cell_guid":"f62d8c64-3553-47a8-859c-872d3da748ae","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"4562e8be-c252-4d35-bf95-33d48660a1b3","_cell_guid":"4eeca4ff-5425-4f4c-856b-75f1538c9792","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"6846d1a8-2a7f-4f87-87f8-4d7fe71c2e1c","_cell_guid":"a4379c26-f493-45df-bd68-76ddbe1d994a","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"1301d6d0-daf8-4cb9-98ce-30eb28ed7e08","_cell_guid":"a914df3e-610c-47d8-a851-fb26bb3e35cf","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"cd83b3de-40bc-4617-a9b9-95e48324e1df","_cell_guid":"8e448eaa-5e6a-4e22-8095-9d1e772dfc91","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"95193dc7-83de-4983-a93f-2a472821c400","_cell_guid":"8741cfb8-6a59-4c0f-9c82-c0afffe6ac6d","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"3588ccbc-b83e-46f4-a539-95ab9974122b","_cell_guid":"efd958ef-cc46-424e-8f09-e44965af2e99","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"acf35a26-c1f0-4c19-a48d-84d674eabdf7","_cell_guid":"f7dea544-7097-4400-90c0-4b6a35e6d009","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"c929f21a-b4ec-46b9-a3aa-d384d0c5d5b4","_cell_guid":"914a7b82-ab80-4fdb-9cb7-97b2d222af68","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"99811ca8-0d67-4b44-9f52-726c8adccbc6","_cell_guid":"a7b665a8-192f-4e40-a632-7fa91f9c2c37","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}