{"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":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":216567044,"sourceType":"kernelVersion"}],"dockerImageVersionId":30822,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Training notebook, coming after the data preprocessing: [here](https://www.kaggle.com/code/liimaxime/fork-of-czii-making-datasets)","metadata":{}},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session\n\nimport torch\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\ntorch.set_num_threads(4)\n\nDEVICE = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nprint(DEVICE)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-09T22:29:44.840914Z","iopub.execute_input":"2025-01-09T22:29:44.841265Z","iopub.status.idle":"2025-01-09T22:29:48.112433Z","shell.execute_reply.started":"2025-01-09T22:29:44.841237Z","shell.execute_reply":"2025-01-09T22:29:48.111470Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.device_count()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T22:29:48.113686Z","iopub.execute_input":"2025-01-09T22:29:48.114207Z","iopub.status.idle":"2025-01-09T22:29:48.138402Z","shell.execute_reply.started":"2025-01-09T22:29:48.114166Z","shell.execute_reply":"2025-01-09T22:29:48.137640Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PARTICLE_CONFS = [0.3, 0.0, 0.2, 0.5, 0.2, 0.5]  \nclasses_dict = {\n  0: 'apo-ferritin',\n  1: 'beta-amylase',\n  2: 'beta-galactosidase',\n  3: 'ribosome',\n  4: 'thyroglobulin',\n  5: 'virus-like-particle',\n}\nparticle_radius = {\n    'apo-ferritin': 60,\n    'beta-amylase': 65,\n    'beta-galactosidase': 90,\n    'ribosome': 150,\n    'thyroglobulin': 130,\n    'virus-like-particle': 135,\n}\nVOXEL_SPACING = 10.012444196428572\nSIZE = 800\nOSIZE = 630","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T22:29:48.139976Z","iopub.execute_input":"2025-01-09T22:29:48.140232Z","iopub.status.idle":"2025-01-09T22:29:48.144668Z","shell.execute_reply.started":"2025-01-09T22:29:48.140212Z","shell.execute_reply":"2025-01-09T22:29:48.143687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision\nfrom torchvision import tv_tensors\nfrom torchvision.transforms import v2 as T\n\nclass CZIIDataset(torch.utils.data.Dataset):\n    def __init__(self, train=True):\n        self.train = train\n        if train:\n            self.image_root = '/kaggle/input/fork-of-czii-making-datasets/datasets/my_czii_det2d/train/images'\n            self.label_root = '/kaggle/input/fork-of-czii-making-datasets/datasets/my_czii_det2d/train/labels'\n            self.labels = []\n            for path, subdirs, files in os.walk(self.label_root):\n                for name in files:\n                    self.labels.append(os.path.join(path, name))\n\n            self.labels = list(sorted(self.labels))\n        else:\n            self.image_root = '/kaggle/input/fork-of-czii-making-datasets/datasets/my_czii_det2d/test/images'\n            self.labels = None\n\n        # list of path of images\n        self.imgs = []\n        for path, subdirs, files in os.walk(self.image_root):\n            for name in files:\n                self.imgs.append(os.path.join(path, name))\n        self.imgs = list(sorted(self.imgs))\n\n        transforms = []\n        transforms.append(T.ToDtype(torch.float, scale=True))\n        transforms.append(T.ToPureTensor())\n        self.transforms =  T.Compose(transforms)\n                \n    def __getitem__(self, idx):\n        \"\"\"read a single image/label\"\"\"\n        img = torch.load(self.imgs[idx], weights_only=True)\n\n        if self.train:\n            target_file = np.loadtxt(self.labels[idx]).reshape(-1,5)\n            target_file = torch.from_numpy(target_file)\n\n            labels = target_file[:, 0].type(torch.int64)\n            boxes = target_file[:, 1:]\n\n        # convert img to tv_tensor\n        img = tv_tensors.Image(img)\n        target = {}\n        if self.transforms is not None:\n            img, target = self.transforms(img, target)\n\n        if not self.train:\n            return img\n        \n        if len(boxes) == 0:\n            target[\"boxes\"] = torch.zeros((0, 4), dtype=torch.float32)\n            target[\"labels\"] = torch.zeros((0,), dtype=torch.int64)\n            target[\"image_id\"] = idx\n        else:\n            target[\"boxes\"] = tv_tensors.BoundingBoxes(boxes, format=\"XYXY\", canvas_size=img.shape[1:])\n            target[\"labels\"] = labels\n            target[\"image_id\"] = idx\n            \n        return img, target\n        \n    def __len__(self):\n        return len(self.imgs)\n\ndef collate_fn(batch):\n    return tuple(zip(*batch))\n\ndataset = CZIIDataset()\nloader = torch.utils.data.DataLoader(\n    dataset,\n    batch_size=16,\n    shuffle=True,\n    collate_fn=collate_fn,\n    num_workers = 4\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T22:29:48.145810Z","iopub.execute_input":"2025-01-09T22:29:48.146089Z","iopub.status.idle":"2025-01-09T22:29:52.842542Z","shell.execute_reply.started":"2025-01-09T22:29:48.146054Z","shell.execute_reply":"2025-01-09T22:29:52.841850Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(dataset)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T22:29:52.843476Z","iopub.execute_input":"2025-01-09T22:29:52.843835Z","iopub.status.idle":"2025-01-09T22:29:52.848732Z","shell.execute_reply.started":"2025-01-09T22:29:52.843810Z","shell.execute_reply":"2025-01-09T22:29:52.848063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"classes_dict_ = {\n  0: 0,\n  1: 0,\n  2: 0,\n  3: 0,\n  4: 0,\n  5: 0,\n}\n\n\"\"\"\nfor elt, t in dataset:\n    for e in t['labels']:\n        classes_dict_[e.item()] += 1\nclasses_dict_\n\"\"\"\n#{0: 1250, 1: 0, 2: 593, 3: 3833, 4: 2158, 5: 717}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T22:29:52.849453Z","iopub.execute_input":"2025-01-09T22:29:52.849648Z","iopub.status.idle":"2025-01-09T22:29:52.868722Z","shell.execute_reply.started":"2025-01-09T22:29:52.849631Z","shell.execute_reply":"2025-01-09T22:29:52.868167Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F\n\nweight = torch.Tensor([62400, 4130, 3080, 1800, 10100, 8400]).to(DEVICE)\ndef fastrcnn_loss_(class_logits, box_regression, labels, regression_targets):\n    # type: (Tensor, Tensor, List[Tensor], List[Tensor]) -> Tuple[Tensor, Tensor]\n    \"\"\"\n    Computes the loss for Faster R-CNN.\n\n    Args:\n        class_logits (Tensor)\n        box_regression (Tensor)\n        labels (list[BoxList])\n        regression_targets (Tensor)\n\n    Returns:\n        classification_loss (Tensor)\n        box_loss (Tensor)\n    \"\"\"\n\n    labels = torch.cat(labels, dim=0)\n    regression_targets = torch.cat(regression_targets, dim=0)\n\n    classification_loss = F.cross_entropy(class_logits, labels, weight=weight)\n\n    # get indices that correspond to the regression targets for\n    # the corresponding ground truth labels, to be used with\n    # advanced indexing\n    sampled_pos_inds_subset = torch.where(labels > 0)[0]\n    labels_pos = labels[sampled_pos_inds_subset]\n    N, num_classes = class_logits.shape\n    box_regression = box_regression.reshape(N, box_regression.size(-1) // 4, 4)\n\n    box_loss = F.smooth_l1_loss(\n        box_regression[sampled_pos_inds_subset, labels_pos],\n        regression_targets[sampled_pos_inds_subset],\n        beta=1 / 9,\n        reduction=\"sum\",\n    )\n    box_loss = box_loss / labels.numel()\n\n    return classification_loss, box_loss\n\ntorchvision.models.detection.roi_heads.fastrcnn_loss = fastrcnn_loss_","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T22:29:52.869456Z","iopub.execute_input":"2025-01-09T22:29:52.869730Z","iopub.status.idle":"2025-01-09T22:29:53.059820Z","shell.execute_reply.started":"2025-01-09T22:29:52.869700Z","shell.execute_reply":"2025-01-09T22:29:53.058920Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\n\nimport torchvision\nfrom torchvision.models.detection import FasterRCNN\nfrom torchvision.models.detection.rpn import AnchorGenerator\n\n# load a pre-trained model for classification and return\n# only the features\nbackbone = torchvision.models.mobilenet_v3_small().features\n# ``FasterRCNN`` needs to know the number of\n# output channels in a backbone. For mobilenet_v2, it's 1280\n# so we need to add it here\nbackbone.out_channels = 576\n\n# let's make the RPN generate 5 x 3 anchors per spatial\n# location, with 5 different sizes and 3 different aspect\n# ratios. We have a Tuple[Tuple[int]] because each feature\n# map could potentially have different sizes and\n# aspect ratios\nsizes = tuple(2*int(value/VOXEL_SPACING*(SIZE/OSIZE)+1) for value in particle_radius.values()) + (32,)\nanchor_generator = AnchorGenerator(\n    sizes=sizes,\n    aspect_ratios=((0.5, 1.0, 2.0),)\n)\nprint(sizes)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T22:29:53.062741Z","iopub.execute_input":"2025-01-09T22:29:53.062967Z","iopub.status.idle":"2025-01-09T22:29:53.522474Z","shell.execute_reply.started":"2025-01-09T22:29:53.062949Z","shell.execute_reply":"2025-01-09T22:29:53.521569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# let's define what are the feature maps that we will\n# use to perform the region of interest cropping, as well as\n# the size of the crop after rescaling.\n# if your backbone returns a Tensor, featmap_names is expected to\n# be [0]. More generally, the backbone should return an\n# ``OrderedDict[Tensor]``, and in ``featmap_names`` you can choose which\n# feature maps to use.\nroi_pooler = torchvision.ops.MultiScaleRoIAlign(\n    featmap_names=['0'],\n    output_size=7,\n    sampling_ratio=2\n)\n\n# put the pieces together inside a Faster-RCNN model\nmodel = FasterRCNN(\n    backbone,\n    num_classes=len(classes_dict),\n    rpn_anchor_generator=anchor_generator,\n    box_roi_pool=roi_pooler\n)\n\n\"\"\"\n# load a model pre-trained on COCO\nmodel = torchvision.models.detection.fasterrcnn_resnet50_fpn(weights=\"DEFAULT\")\n\n# replace the classifier with a new one, that has\n# num_classes which is user-defined\nnum_classes = len(classes_dict)\n# get number of input features for the classifier\nin_features = model.roi_heads.box_predictor.cls_score.in_features\n# replace the pre-trained head with a new one\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)\n\"\"\"\n\nmodel.to(DEVICE)\n#model = torch.nn.DataParallel(model, device_ids=[0, 1])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"optimizer = torch.optim.AdamW(model.parameters())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T22:29:53.524025Z","iopub.execute_input":"2025-01-09T22:29:53.524339Z","iopub.status.idle":"2025-01-09T22:29:53.528588Z","shell.execute_reply.started":"2025-01-09T22:29:53.524316Z","shell.execute_reply":"2025-01-09T22:29:53.527893Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"epochs = 500\n\nmodel.train()\n\nfor epoch in range(epochs):\n\n    loss_total = 0.0\n    \n    for imgs, targs in loader:\n\n        # moove to GPU\n        imgs = list(image.to(DEVICE) for image in imgs)\n        targs = [{k: v.to(DEVICE) if isinstance(v, torch.Tensor) else v for k, v in t.items()} for t in targs]\n\n        # compute losses\n\n        loss_dict = model(imgs, targs)\n        losses = sum(loss for loss in loss_dict.values())\n\n        # update model\n        optimizer.zero_grad()\n        losses.backward()\n        optimizer.step()\n\n        loss_total += losses.item()\n    \n    # print state\n    state = f\"epoch = {epoch} | loss = {loss_total/len(dataset)}\\nloss_dict = {loss_dict}\"\n    print(state)\n    print('-'*5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-09T22:29:53.529367Z","iopub.execute_input":"2025-01-09T22:29:53.529577Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model prediction on the test dataset","metadata":{}},{"cell_type":"code","source":"import re\nfrom scipy.spatial import cKDTree\nimport json\n\ntransforms = []\ntransforms.append(T.ToDtype(torch.float, scale=True))\ntransforms.append(T.ToPureTensor())\ntransforms =  T.Compose(transforms)\n\ndef get_particle_boxes_and_types(model, test_root, run):\n\n    # list of images for the given run\n    files = os.listdir(os.path.join(test_root, run, 'denoised'))\n\n    labels = []\n    x = []\n    y = []\n    z = []\n    conf = []\n    for file in files:\n        if file.endswith('txt'):\n            continue\n        \n        # get original picture size\n        #vol = zarr.open(f'/kaggle/input/czii-cryo-et-object-identification/test/static/ExperimentRuns/{run}/VoxelSpacing10.000/denoised.zarr', mode='r')\n        #vol = vol[0]\n        #OSIZE1, OSIZE2 = vol.shape[1:]\n        #OSIZE1, OSIZE2 = np.loadtxt(os.path.join(test_root, run, 'denoised', 'dim.txt')).astype(int)\n        zarraypath = f'/kaggle/input/czii-cryo-et-object-identification/test/static/ExperimentRuns/{run}/VoxelSpacing10.000/denoised.zarr/0/.zarray'\n        with open(zarraypath, 'r') as zf:\n            config = json.load(zf)\n        OSIZE1, OSIZE2 = config['shape'][1:]\n        \n        # \n        path = os.path.join(test_root, run, 'denoised', file)\n        z_ = int(re.findall(r'\\d+', file)[0])\n        \n        img = torch.load(path).reshape(1,3,SIZE,SIZE)\n        res = model(transforms(img).to(DEVICE))[0]\n        \n        x_min = res['boxes'][:,0].cpu().detach().numpy() \n        y_min = res['boxes'][:,1].cpu().detach().numpy() \n        x_max = res['boxes'][:,2].cpu().detach().numpy() \n        y_max = res['boxes'][:,3].cpu().detach().numpy()\n\n        x_ = (x_min + x_max)/2.0 * OSIZE1 / SIZE * VOXEL_SPACING\n        y_ = (y_min + y_max)/2.0 * OSIZE2 / SIZE * VOXEL_SPACING\n        z_ = z_ * np.ones_like(x_) + 5\n\n        conf_ = res['scores'].cpu().detach().numpy()\n        \n        x.extend(x_.tolist())\n        y.extend(y_.tolist())\n        z.extend(z_.tolist())\n        conf.extend(conf_.tolist())\n        labels.extend(res['labels'].cpu().numpy().tolist())\n        \n    df = pd.DataFrame()\n    df['particle_type'] = labels\n    df['x'] = x\n    df['y'] = y\n    df['z'] = z    \n    df['conf'] = conf\n    return df\n\ndef get_particle_positions(df, run):\n    \n    df_out = []\n\n    for key,val in classes_dict.items():\n\n        if val == 'beta-amylase':\n            continue\n            \n        pdf = df[df['particle_type']==key]\n        n = len(pdf)\n\n        p_rad = particle_radius[val]\n            \n        # v2\n        # Define a function to filter connections based on conditions\n        def filter_connections(i, indices):\n            xi, yi, zi = points[i]\n            adjacency = []\n            for j in indices:\n                if j <= i:  # Avoid duplicate or invalid indices\n                    continue\n                xj, yj, zj = points[j]\n                dist_p2 = (xi - xj)**2 + (yi - yj)**2\n                xy_tol = p_rad / 8.0\n                xy_tol_p2 = xy_tol ** 2\n                if abs(zi - zj) <= 20 and dist_p2 < xy_tol_p2 and dist_p2 + (zi - zj)**2 < p_rad**2:\n                    adjacency.append(j)\n            return adjacency\n            \n        points = pdf[['x', 'y', 'z']].values\n        tree = cKDTree(points)\n        \n        # Query all points within the 3D radius\n        adjacency_list = [[] for _ in range(n)]\n        for i in range(n):\n            neighbors = tree.query_ball_point(points[i], r=p_rad)\n            adjacency_list[i] = filter_connections(i, neighbors)\n            for j in adjacency_list[i]:  # Add reciprocal connections\n                adjacency_list[j].append(i)\n        \n        # loop over graph to compute mean positions\n        df_ = pd.DataFrame()\n        passed = [False for _ in range(n)]\n        cxs = []\n        cys = []\n        czs = []\n        \n        def compute_sum(ind, nv, cx, cy, cz, conf_sum, passed):\n            passed[ind] = True\n            nv += 1\n            cx += pdf.iloc[ind].x\n            cy += pdf.iloc[ind].y\n            cz += pdf.iloc[ind].z\n            conf_sum += pdf.iloc[ind].conf\n\n            for next_v in adjacency_list[ind]:\n                if (passed[next_v]): continue\n                nv, cx, cy, cz, conf_sum = compute_sum(next_v, nv, cx, cy, cz, conf_sum, passed)\n                \n            return nv, cx, cy, cz, conf_sum\n        \n        for i in range(n):\n            nv = 0\n            cx = 0.0\n            cy = 0.0\n            cz = 0.0\n            conf_sum = 0.0\n\n            if not passed[i]:\n                nv, cx, cy, cz, conf_sum = compute_sum(i, nv, cx, cy, cz, conf_sum, passed)\n\n            if nv>=2 and conf_sum / (nv**0.5) > PARTICLE_CONFS[key]:\n                cxs.append(cx / nv)\n                cys.append(cy / nv)\n                czs.append(cz / nv)        \n\n        df_['experiment'] = [run] * len(cxs)\n        df_['particle_type'] = [val] * len(cys)\n        df_['x'] = cxs\n        df_['y'] = cys\n        df_['z'] = czs\n        df_out.append(df_)\n    return pd.concat(df_out, axis=0)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\n\ntest_root = '/kaggle/input/fork-of-czii-making-datasets/datasets/my_czii_det2d/test/images'\ntest_runs = os.listdir(test_root)\n\nsubmissions = []\nfor run in test_runs:\n    \n    # get boxes/class for each run\n    df = get_particle_boxes_and_types(model, test_root, run)\n    \n    # try to connect each box with a connectivity graph to compute average position for each particle\n    submissions.append(get_particle_positions(df, run))\n    \nsubmission = pd.concat(submissions, axis=0).reset_index(drop=True)\nsubmission.insert(0, 'id', range(len(submission)))\nsubmission.to_csv('submission.csv', index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_solutions(run):\n    train_root = f'/kaggle/input/czii-cryo-et-object-identification/train/overlay/ExperimentRuns/{run}/Picks'\n\n    solutions = []\n    for particle in particle_radius.keys():\n    \n        with open(os.path.join(train_root, f'{particle}.json'), 'r') as json_file:\n            content = json.load(json_file)\n\n        x_list = []\n        y_list = []\n        z_list = []\n        for elt in content['points']:\n            x_, y_, z_ = elt['location'].values()\n            x_list.append(x_)\n            y_list.append(y_)\n            z_list.append(z_)\n        solutions_ = pd.DataFrame()\n        solutions_['experiment'] =  [run] * len(x_list)\n        solutions_['particle_type'] = [particle] * len(x_list)\n        solutions_['x'] = x_list\n        solutions_['y'] = y_list\n        solutions_['z'] = z_list\n        solutions.append(solutions_)\n    return pd.concat(solutions, axis=0).reset_index(drop=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"solutions = []\n\nfor run in test_runs:\n    solutions.append(get_solutions(run))\n\nsolution = pd.concat(solutions, axis=0).reset_index(drop=True)\nsolution.insert(0, 'id', range(len(solution)))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Mock evaluation","metadata":{}},{"cell_type":"code","source":"\"\"\"\nDerived from:\nhttps://github.com/cellcanvas/album-catalog/blob/main/solutions/copick/compare-picks/solution.py\n\"\"\"\nfrom scipy.spatial import KDTree\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\ndef compute_metrics(reference_points, reference_radius, candidate_points):\n    num_reference_particles = len(reference_points)\n    num_candidate_particles = len(candidate_points)\n\n    if len(reference_points) == 0:\n        return 0, num_candidate_particles, 0\n\n    if len(candidate_points) == 0:\n        return 0, 0, num_reference_particles\n\n    ref_tree = KDTree(reference_points)\n    candidate_tree = KDTree(candidate_points)\n    raw_matches = candidate_tree.query_ball_tree(ref_tree, r=reference_radius)\n    matches_within_threshold = []\n    for match in raw_matches:\n        matches_within_threshold.extend(match)\n    # Prevent submitting multiple matches per particle.\n    # This won't be be strictly correct in the (extremely rare) case where true particles\n    # are very close to each other.\n    matches_within_threshold = set(matches_within_threshold)\n    tp = int(len(matches_within_threshold))\n    fp = int(num_candidate_particles - tp)\n    fn = int(num_reference_particles - tp)\n    return tp, fp, fn\n\n\ndef score(\n        solution: pd.DataFrame,\n        submission: pd.DataFrame,\n        row_id_column_name: str = None,\n        distance_multiplier: float = 1.0,\n        beta: int = 4) -> float:\n    '''\n    F_beta\n      - a true positive occurs when\n         - (a) the predicted location is within a threshold of the particle radius, and\n         - (b) the correct `particle_type` is specified\n      - raw results (TP, FP, FN) are aggregated across all experiments for each particle type\n      - f_beta is calculated for each particle type\n      - individual f_beta scores are weighted by particle type for final score\n    '''\n\n    particle_radius = {\n        'apo-ferritin': 60,\n        'beta-amylase': 65,\n        'beta-galactosidase': 90,\n        'ribosome': 150,\n        'thyroglobulin': 130,\n        'virus-like-particle': 135,\n    }\n\n    weights = {\n        'apo-ferritin': 1,\n        'beta-amylase': 0,\n        'beta-galactosidase': 2,\n        'ribosome': 1,\n        'thyroglobulin': 2,\n        'virus-like-particle': 1,\n    }\n\n    particle_radius = {k: v * distance_multiplier for k, v in particle_radius.items()}\n\n    # Filter submission to only contain experiments found in the solution split\n    split_experiments = set(solution['experiment'].unique())\n    submission = submission.loc[submission['experiment'].isin(split_experiments)]\n\n    # Only allow known particle types\n    if not set(submission['particle_type'].unique()).issubset(set(weights.keys())):\n        raise ParticipantVisibleError('Unrecognized `particle_type`.')\n\n    assert solution.duplicated(subset=['experiment', 'x', 'y', 'z']).sum() == 0\n    assert particle_radius.keys() == weights.keys()\n\n    results = {}\n    for particle_type in solution['particle_type'].unique():\n        results[particle_type] = {\n            'total_tp': 0,\n            'total_fp': 0,\n            'total_fn': 0,\n        }\n\n    for experiment in split_experiments:\n        for particle_type in solution['particle_type'].unique():\n            reference_radius = particle_radius[particle_type]\n            select = (solution['experiment'] == experiment) & (solution['particle_type'] == particle_type)\n            reference_points = solution.loc[select, ['x', 'y', 'z']].values\n\n            select = (submission['experiment'] == experiment) & (submission['particle_type'] == particle_type)\n            candidate_points = submission.loc[select, ['x', 'y', 'z']].values\n\n            if len(reference_points) == 0:\n                reference_points = np.array([])\n                reference_radius = 1\n\n            if len(candidate_points) == 0:\n                candidate_points = np.array([])\n\n            tp, fp, fn = compute_metrics(reference_points, reference_radius, candidate_points)\n\n            results[particle_type]['total_tp'] += tp\n            results[particle_type]['total_fp'] += fp\n            results[particle_type]['total_fn'] += fn\n\n    aggregate_fbeta = 0.0\n    for particle_type, totals in results.items():\n        tp = totals['total_tp']\n        fp = totals['total_fp']\n        fn = totals['total_fn']\n\n        precision = tp / (tp + fp) if tp + fp > 0 else 0\n        recall = tp / (tp + fn) if tp + fn > 0 else 0\n        fbeta = (1 + beta**2) * (precision * recall) / (beta**2 * precision + recall) if (precision + recall) > 0 else 0.0\n        aggregate_fbeta += fbeta * weights.get(particle_type, 1.0)\n\n    if weights:\n        aggregate_fbeta = aggregate_fbeta / sum(weights.values())\n    else:\n        aggregate_fbeta = aggregate_fbeta / len(results)\n    return aggregate_fbeta","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for coo in ['x', 'y', 'z']:\n    print(submission[coo].min(), submission[coo].max())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for coo in ['x', 'y', 'z']:\n    print(solution[coo].min(), solution[coo].max())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"score(submission = submission, solution=solution)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Some predictions","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom torchvision.utils import draw_bounding_boxes\n\nfig, axs = plt.subplots(2, 2, figsize=(30, 30), sharex = True, sharey=True)\n\ncounter = 0\n\nfor imgs, targets in loader:\n    for i in range(16):\n        \n        img = imgs[i]\n        target = targets[i]\n\n        if target['boxes'].shape[0]==0:\n            continue\n            \n        pred = model(torch.unsqueeze(imgs[i], 0).to(DEVICE))\n        pred_boxes = pred[0]['boxes'].detach().cpu()\n        pred_labels = [classes_dict[label] for label in pred[0]['labels'].detach().cpu().numpy()]\n        img_pred = draw_bounding_boxes(img, pred_boxes, pred_labels, colors=\"red\", font_size=10)\n        axs[counter, 0].imshow(img_pred.permute(1, 2, 0))\n\n        true_boxes = target['boxes']\n        true_labels = [classes_dict[label] for label in target['labels'].cpu().numpy()]\n        image_sol = draw_bounding_boxes(img, true_boxes, true_labels, colors=\"red\", font_size=10)\n        axs[counter, 1].imshow(image_sol.permute(1, 2, 0))\n\n        counter += 1\n        \n        if counter==2:\n            break\n    if counter == 2:\n        break","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"nboxes = len(true_boxes)\nfig, axs = plt.subplots(4, figsize=(30, 30), sharex = True, sharey=True)\n\nfor i in range(4):\n    xmin, ymin, xmax, ymax = true_boxes[i]\n    dx = xmax - xmin\n    dy = ymax - ymin\n    axs[i].imshow(image_sol.permute(1, 2, 0))\n    axs[i].set_xlim(xmin-2*dx, xmax+2*dx)\n    axs[i].set_ylim(ymin-2*dy, ymax+2*dy)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}