{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","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":9867543,"sourceType":"datasetVersion","datasetId":6040935},{"sourceId":210966406,"sourceType":"kernelVersion"}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Setup","metadata":{}},{"cell_type":"code","source":"!cp -r '/kaggle/input/hengck-czii-cryo-et-01/wheel_file' '/kaggle/working/'\n!pip install /kaggle/working/wheel_file/asciitree-0.3.3/asciitree-0.3.3\n!pip install --no-index --find-links=/kaggle/working/wheel_file zarr\n!pip install --no-index --find-links=/kaggle/working/wheel_file connected-components-3d","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:50:51.245095Z","iopub.execute_input":"2024-12-09T07:50:51.24548Z","iopub.status.idle":"2024-12-09T07:51:54.812569Z","shell.execute_reply.started":"2024-12-09T07:50:51.245448Z","shell.execute_reply":"2024-12-09T07:51:54.811502Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 42\nimport os\nimport sys\nimport pickle\nimport random as rnd\nimport numpy as np\nfrom numpy import random as np_rnd\nimport pandas as pd\nfrom tqdm import tqdm\nimport zarr\nimport glob\nimport json\nimport time\nimport gc\n\nimport torch\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.optim import AdamW\nfrom transformers import get_polynomial_decay_schedule_with_warmup\nfrom datasets import Dataset\nfrom torch.utils.data import DataLoader\nfrom torchvision import transforms\nfrom scipy.optimize import linear_sum_assignment\nfrom scipy.spatial import KDTree\nimport albumentations as A\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\ndevice = torch.device(\"cuda\") if  torch.cuda.is_available() else torch.device(\"cpu\")\ndevice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:51:54.814631Z","iopub.execute_input":"2024-12-09T07:51:54.814937Z","iopub.status.idle":"2024-12-09T07:52:33.575896Z","shell.execute_reply.started":"2024-12-09T07:51:54.814907Z","shell.execute_reply":"2024-12-09T07:52:33.575066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    debug = False\n    dp_version = \"a1\"\n    max_coord_values = {\"x\": 6400, \"y\": 6400, \"z\": 1840}\n    # 0: high, 1: medium, 2: low\n    resolution_multiplier = {0: 10, 1: 20, 2: 40}\n    resolution = 0\n    channel_sampling_interval = 20\n    img_size = (224, 224)\n    n_channels = 10\n    label2id = {\n        \"no-object\": 0,\n        \"apo-ferritin\": 1,\n        \"beta-galactosidase\": 2,\n        \"ribosome\": 3,\n        \"thyroglobulin\": 4,\n        \"virus-like-particle\": 5,\n    }\n    id2label = {\n        0: \"no-object\",\n        1: \"apo-ferritin\",\n        2: \"beta-galactosidase\",\n        3: \"ribosome\",\n        4: \"thyroglobulin\",\n        5: \"virus-like-particle\",\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    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    beta = 4\n    n_folds = 7\n    epochs = 5 if debug else 40\n    early_stopping_rounds = 10\n    batch_size = 3\n    n_aug = 2 if debug else 16\n    eta = 5e-5\n    weight_decay = 1e-2\n    n_quries = 512\n    out_channels = 64\n    embed_dim = 256\n    n_heads = 4\n    ffn_dim = 1024\n    n_encoders = 3\n    n_decoders = 3\n    unmatched_loss_weight = 10.0\n    n_coords = 3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.577275Z","iopub.execute_input":"2024-12-09T07:52:33.578102Z","iopub.status.idle":"2024-12-09T07:52:33.586139Z","shell.execute_reply.started":"2024-12-09T07:52:33.57804Z","shell.execute_reply":"2024-12-09T07:52:33.585203Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed=42):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    # python random\n    rnd.seed(seed)\n    # numpy random\n    np_rnd.seed(seed)\n    # RAPIDS random\n    try:\n        cupy.random.seed(seed)\n    except:\n        pass\n    # tf random\n    try:\n        tf_rnd.set_seed(seed)\n    except:\n        pass\n    # pytorch random\n    try:\n        torch.backends.cudnn.benchmark = False\n        torch.backends.cudnn.deterministic = True\n        torch.manual_seed(seed)\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n    except:\n        pass\n\ndef pickleIO(obj, src, op=\"r\"):\n    if op==\"w\":\n        with open(src, op + \"b\") as f:\n            pickle.dump(obj, f)\n    elif op==\"r\":\n        with open(src, op + \"b\") as f:\n            tmp = pickle.load(f)\n        return tmp\n    else:\n        print(\"unknown operation\")\n        return obj\n    \ndef createFolder(directory):\n    try:\n        if not os.path.exists(directory):\n            os.makedirs(directory)\n    except OSError:\n        print('Error: Creating directory. ' + directory)\n\ndef findIdx(data_x, col_names):\n    return [int(i) for i, j in enumerate(data_x) if j in col_names]\n\ndef diff(first, second):\n    second = set(second)\n    return [item for item in first if item not in second]\n\ndef minmax_scaler(data, feature_range=(0, 255)):\n    min_val, max_val = feature_range\n    scaled_data = (data - np.min(data)) / (np.max(data) - np.min(data))\n    scaled_data = scaled_data * (max_val - min_val) + min_val\n    return scaled_data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.588748Z","iopub.execute_input":"2024-12-09T07:52:33.589399Z","iopub.status.idle":"2024-12-09T07:52:33.603238Z","shell.execute_reply.started":"2024-12-09T07:52:33.589371Z","shell.execute_reply":"2024-12-09T07:52:33.602525Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seed_everything(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.604444Z","iopub.execute_input":"2024-12-09T07:52:33.60476Z","iopub.status.idle":"2024-12-09T07:52:33.625204Z","shell.execute_reply.started":"2024-12-09T07:52:33.604729Z","shell.execute_reply":"2024-12-09T07:52:33.624373Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loading preprocessed data","metadata":{}},{"cell_type":"code","source":"df_annot = pickleIO(None, f\"/kaggle/input/czii-data-pipeline-{CFG.dp_version}/df_annot.pkl\", \"r\")\ndf_img = np.load(f\"/kaggle/input/czii-data-pipeline-{CFG.dp_version}/df_img.npz\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.62632Z","iopub.execute_input":"2024-12-09T07:52:33.626578Z","iopub.status.idle":"2024-12-09T07:52:33.754749Z","shell.execute_reply.started":"2024-12-09T07:52:33.626554Z","shell.execute_reply":"2024-12-09T07:52:33.753726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# drop beta-amylase & target encoding\ndf_annot = df_annot[df_annot[\"obj\"] != \"beta-amylase\"]\ndf_annot[\"obj\"] = df_annot[\"obj\"].apply(lambda x: CFG.label2id[x])\ndf_annot[\"obj\"].value_counts(normalize=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.756085Z","iopub.execute_input":"2024-12-09T07:52:33.756459Z","iopub.status.idle":"2024-12-09T07:52:33.78436Z","shell.execute_reply.started":"2024-12-09T07:52:33.756418Z","shell.execute_reply":"2024-12-09T07:52:33.783455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_annot","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.785692Z","iopub.execute_input":"2024-12-09T07:52:33.785992Z","iopub.status.idle":"2024-12-09T07:52:33.805287Z","shell.execute_reply.started":"2024-12-09T07:52:33.785963Z","shell.execute_reply":"2024-12-09T07:52:33.804405Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_annot.info()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.806433Z","iopub.execute_input":"2024-12-09T07:52:33.806711Z","iopub.status.idle":"2024-12-09T07:52:33.824511Z","shell.execute_reply.started":"2024-12-09T07:52:33.806685Z","shell.execute_reply":"2024-12-09T07:52:33.823602Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_annot.groupby(\"run\").size()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.82708Z","iopub.execute_input":"2024-12-09T07:52:33.827378Z","iopub.status.idle":"2024-12-09T07:52:33.837922Z","shell.execute_reply.started":"2024-12-09T07:52:33.82735Z","shell.execute_reply":"2024-12-09T07:52:33.837095Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define helper functions","metadata":{}},{"cell_type":"code","source":"def compute_metrics(reference_points, reference_radius, candidate_points):\n    num_reference_particles = len(reference_points)\n    num_candidate_particles = len(candidate_points)\n    # exception\n    if len(reference_points) == 0:\n        return 0, num_candidate_particles, 0\n    if len(candidate_points) == 0:\n        return 0, 0, num_reference_particles\n    # for prediction\n    candidate_tree = KDTree(candidate_points)\n    # for ground truth\n    ref_tree = KDTree(reference_points)\n    # matchin\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\ndef get_fbeta(df_gt, df_pred, particle_radius, weights, beta):\n    results = {}\n    for particle_type in df_gt['obj'].unique():\n        results[particle_type] = {\n            'total_tp': 0,\n            'total_fp': 0,\n            'total_fn': 0,\n        }\n    \n    for particle_type in df_gt['obj'].unique():\n        reference_points = df_gt.loc[(df_gt['obj'] == particle_type).values, ['x', 'y', 'z']].values\n        candidate_points = df_pred.loc[(df_pred['obj'] == particle_type).values, ['x', 'y', 'z']].values\n        # exception\n        if len(reference_points) == 0:\n            reference_points = np.array([])\n            reference_radius = 1\n        if len(candidate_points) == 0:\n            candidate_points = np.array([])\n        # compute\n        tp, fp, fn = compute_metrics(reference_points, particle_radius[particle_type], candidate_points)\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    \n    return aggregate_fbeta","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.839127Z","iopub.execute_input":"2024-12-09T07:52:33.839428Z","iopub.status.idle":"2024-12-09T07:52:33.852469Z","shell.execute_reply.started":"2024-12-09T07:52:33.839401Z","shell.execute_reply":"2024-12-09T07:52:33.851552Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_optimizer_params(model, eta, weight_decay):\n    no_decay = [\"bias\", \"LayerNorm.bias\", \"LayerNorm.weight\"]\n    optimizer_parameters = [\n        # apply weight decay\n        {'params': [p for n, p in model.named_parameters() if not any(nd in n for nd in no_decay)],\n         'lr': eta, 'weight_decay': weight_decay},\n        # don't apply weight decay for LayerNormalization layer\n        {'params': [p for n, p in model.named_parameters() if any(nd in n for nd in no_decay)],\n         'lr': eta, 'weight_decay': 0.0},\n    ]\n    return optimizer_parameters\n\ndef get_scheduler(optimizer, num_warmup_steps, num_training_steps):\n    scheduler = get_polynomial_decay_schedule_with_warmup(\n        optimizer, num_warmup_steps=num_warmup_steps, num_training_steps=num_training_steps, power=1.0, lr_end=1e-7\n    )\n    return scheduler\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self, name, fmt=':f'):\n        self.name = name\n        self.fmt = fmt\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n    def __str__(self):\n        fmtstr = '{name} {val' + self.fmt + '} ({avg' + self.fmt + '})'\n        return fmtstr.format(**self.__dict__)\n\ndef calculate_loss(output, batch, criterion):\n    loss = 0.0\n    for pred_coord, pred_label, true_coord, true_label in zip(output[\"reg\"], output[\"cls\"], batch[\"coord\"], batch[\"label\"]):\n        # get cost matrix\n        cost_coord = -1.0 * (1 / (torch.cdist(pred_coord, true_coord.to(device)) + 1.0))\n        cost_label = -1.0 * torch.stack([torch.stack([y_prob[y_true] for y_true in true_label], dim=0) for y_prob in pred_label.softmax(dim=-1)], dim=0)\n        # range of cost is -1 to 0\n        cost = (cost_coord + cost_label) / 2.0\n        # matching\n        matched = linear_sum_assignment(cost.detach().cpu().numpy())\n        # matched loss\n        matched_loss = torch.sqrt(criterion[\"reg\"](pred_coord[matched[0]], true_coord[matched[1]].to(device)).sum(dim=-1)) + criterion[\"cls\"](pred_label[matched[0]], true_label[matched[1]].to(device))\n        # unmatched loss\n        unmatched_mask = ~np.isin(np.arange(len(pred_coord)), matched[0])\n        unmatched_loss = criterion[\"cls\"](pred_label[unmatched_mask], torch.zeros(unmatched_mask.sum(), dtype=torch.int64).to(device))\n        # final loss\n        matched_loss = matched_loss.sum() / len(pred_coord)\n        unmatched_loss = unmatched_loss.sum() / len(pred_coord)\n        loss += (matched_loss + (unmatched_loss.sum() / CFG.unmatched_loss_weight)) / len(output[\"reg\"])\n    return loss\n    \ndef train_fn(model, dl, criterion, optimizer, scheduler, grad_scaler):\n    model.train()\n    metrics = {\n        \"loss\": AverageMeter(\"loss\", fmt=\":.5f\"),\n    }\n    \n    for idx, batch in enumerate(tqdm(dl)):\n        with torch.cuda.amp.autocast():\n            output = model(batch[\"pixel_values\"].to(device))\n            loss = calculate_loss(output, batch, criterion)\n        # initialization gradients to zero\n        optimizer.zero_grad()\n        # get scaled gradients by float16 (default)\n        grad_scaler.scale(loss).backward()\n        # apply original gradients (unscaling) to parameters\n        # if these gradients do not contain infs or NaNs, optimizer.step() is then called.\n        # otherwise, optimizer.step() is skipped.\n        grad_scaler.step(optimizer)\n        grad_scaler.update()\n        # update scheduler\n        scheduler.step()\n        # calcuate metrics\n        metrics[\"loss\"].update(loss.item())\n        del batch, output, loss\n        gc.collect()\n        torch.cuda.empty_cache()\n\n        if CFG.debug:\n            if idx >= 10:\n                break\n    \n    return metrics\n\n@torch.no_grad()\ndef valid_fn(model, dl, criterion):\n    model.eval()\n    metrics = {\n        \"loss\": AverageMeter(\"loss\", fmt=\":.5f\"),\n    }\n\n    y_pred = {\n        \"reg\": [],\n        \"cls\": [],\n    }\n    for idx, batch in enumerate(tqdm(dl)):\n        output = model(batch[\"pixel_values\"].to(device))\n        loss = calculate_loss(output, batch, criterion)\n        # calcuate metrics\n        metrics[\"loss\"].update(loss.item())\n        # get output\n        y_pred[\"reg\"].append(output[\"reg\"].detach().cpu().numpy())\n        y_pred[\"cls\"].append(output[\"cls\"].detach().cpu().numpy())\n        del batch, output, loss\n        gc.collect()\n        torch.cuda.empty_cache()\n\n        if CFG.debug:\n            if idx >= 10:\n                break\n    \n    return metrics, y_pred\n\n@torch.no_grad()\ndef infer_fn(model, dl):\n    model.eval()\n    y_pred = {\n        \"reg\": [],\n        \"cls\": [],\n    }\n    for idx, batch in enumerate(tqdm(dl)):\n        output = model(batch[\"pixel_values\"].to(device))\n        # get output\n        y_pred[\"reg\"].append(output[\"reg\"].detach().cpu().numpy())\n        y_pred[\"cls\"].append(output[\"cls\"].detach().cpu().numpy())\n        del batch, output\n        gc.collect()\n        torch.cuda.empty_cache()\n    return y_pred\n\ndef do_training(fold, model, model_params, df_gt, train_dl, valid_dl):\n    # set loss & optimizer\n    optimizer_parameters = get_optimizer_params(\n        model,\n        eta=CFG.eta,\n        weight_decay=CFG.weight_decay\n    )\n    optimizer = AdamW(optimizer_parameters, lr=CFG.eta, weight_decay=CFG.weight_decay)\n    scheduler = get_scheduler(\n        optimizer,\n        num_warmup_steps=0,\n        num_training_steps=len(train_dl) * CFG.epochs\n    )\n    criterion = {\n        \"reg\": nn.MSELoss(reduction=\"none\"),\n        \"cls\": nn.CrossEntropyLoss(reduction=\"none\"),\n    }\n    grad_scaler = torch.cuda.amp.GradScaler()\n    \n    best_score = np.inf\n    early_stopping_cnt = 0\n    for epoch in range(CFG.epochs):\n        seed_everything(epoch)\n        epoch_start_time = time.time()\n        \n        # training\n        train_metrics = train_fn(model, train_dl, criterion, optimizer, scheduler, grad_scaler)\n        # validation\n        valid_metrics, valid_pred = valid_fn(model, valid_dl, criterion)\n        \n        df_pred = []\n        for run, coord, label in zip(df_gt.index.unique(), np.concatenate(valid_pred[\"reg\"]), np.concatenate(valid_pred[\"cls\"])):\n            df = pd.DataFrame(index=[run] * len(coord))\n            df[\"obj\"] = label.argmax(axis=-1)\n            df[[\"x\", \"y\", \"z\"]] = np.clip(coord, a_min=0.01, a_max=0.99)\n            df_pred.append(df)\n            print(pd.DataFrame( F.softmax(torch.tensor(label), dim=-1)) )\n        df_pred = pd.concat(df_pred)\n        for axis, val in CFG.max_coord_values.items():\n            df_pred[axis] *= val\n        df_pred[\"obj\"] = df_pred[\"obj\"].map(CFG.id2label)\n        print(df_pred)\n        print(df_pred[\"obj\"].value_counts(normalize=True))\n        print(df_gt)\n        fbeta = get_fbeta(df_gt, df_pred, particle_radius=CFG.particle_radius, weights=CFG.weights, beta=CFG.beta)\n\n        score = {\n            **{f\"train_{k}\": v.avg for k, v in train_metrics.items()},\n            **{f\"valid_{k}\": v.avg for k, v in valid_metrics.items()},\n            \"valid_fbeta\": fbeta,\n        }\n        msg = [f\"Epoch[{epoch+1}/{CFG.epochs}]\"]\n        msg += [f\"{k}: {round(v, 5)}\" for k, v in score.items()]\n        msg += [f\"eta: {round(optimizer.param_groups[-1]['lr'], 5)}\", f\"elapsed: {round(time.time() - epoch_start_time, 3)}\"]\n        print(\"\\n\".join(msg))\n                \n        if score[\"valid_loss\"] < best_score:\n            return_score_dic = score\n            best_score = score[\"valid_loss\"]\n            print(\"INFO: Found best weight\\n\\n\")\n            torch.save(\n                {'model': model.state_dict(), \"model_params\": model_params},\n                f\"fold{fold}_model_best.pth\",\n            )\n            early_stopping_cnt = 0\n        else:\n            early_stopping_cnt += 1\n        \n        if early_stopping_cnt == CFG.early_stopping_rounds:\n            break   \n    torch.save(\n        {'model': model.state_dict(), \"model_params\": model_params},\n        f\"fold{fold}_model_last.pth\",\n    )\n    model.load_state_dict(torch.load(f\"fold{fold}_model_best.pth\")[\"model\"])\n    return return_score_dic","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.853905Z","iopub.execute_input":"2024-12-09T07:52:33.854297Z","iopub.status.idle":"2024-12-09T07:52:33.88329Z","shell.execute_reply.started":"2024-12-09T07:52:33.854257Z","shell.execute_reply":"2024-12-09T07:52:33.882402Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define model","metadata":{}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, image_container, annot_container, n_aug):\n        self.image_container = image_container\n        self.annot_container = annot_container\n        self.runs = list(image_container.keys())\n        self.transform = transform = A.Compose([\n            A.Rotate(limit=180, p=1.0),\n        ], keypoint_params=A.KeypointParams(format='xy', remove_invisible=False))\n        self.n_aug = n_aug\n\n    def __len__(self):\n        return len(self.runs)\n\n    def __getitem__(self, idx):\n        if self.n_aug == 0:\n            batch = {\n                \"pixel_values\": [torch.tensor(self.image_container[self.runs[i]], dtype=torch.float32) for i in idx],\n                \"coord\": [torch.tensor(self.annot_container.loc[self.runs[i]][:, :-1] / np.array([[CFG.img_size[0], CFG.img_size[1], CFG.n_channels]]), dtype=torch.float32) for i in idx],\n                \"label\": [torch.tensor(self.annot_container.loc[self.runs[i]][:, -1], dtype=torch.int64) for i in idx],\n            }\n        else:\n            batch = {\n                \"pixel_values\": [],\n                \"coord\": [],\n                \"label\": [],\n            }\n            for i in idx:\n                pixel_values = torch.tensor(self.image_container[self.runs[i]], dtype=torch.float32).permute(1, 2, 0).detach().cpu().numpy()\n                coord = self.annot_container.loc[self.runs[i]][:, :-1]\n                label = self.annot_container.loc[self.runs[i]][:, -1]\n                for _ in range(self.n_aug):\n                    aug = self.transform(image=pixel_values, keypoints=[tuple(i[:2]) for i in coord])\n                    batch[\"pixel_values\"].append(torch.tensor(aug[\"image\"], dtype=torch.float32).permute(2, 0, 1))\n                    aug[\"keypoints\"] = np.concatenate([np.array(aug[\"keypoints\"]), coord[:, [-1]]], axis=1) / np.array([[CFG.img_size[0], CFG.img_size[1], CFG.n_channels]])\n                    mask = (aug[\"keypoints\"] >= 0).all(axis=1)\n                    if mask.sum() == 0:\n                        continue\n                    batch[\"coord\"].append(torch.tensor(aug[\"keypoints\"][mask], dtype=torch.float32))\n                    batch[\"label\"].append(torch.tensor(label[mask], dtype=torch.int64))\n        return batch\n\ndef collate_fn(samples):\n    batch = {\n        \"pixel_values\": torch.stack([sample[\"pixel_values\"] for sample in samples]),\n        \"coord\": [sample[\"coord\"] for sample in samples],\n        \"label\": [sample[\"label\"] for sample in samples],\n    }\n    return batch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T07:52:33.884521Z","iopub.execute_input":"2024-12-09T07:52:33.884809Z","iopub.status.idle":"2024-12-09T07:52:33.898533Z","shell.execute_reply.started":"2024-12-09T07:52:33.884784Z","shell.execute_reply":"2024-12-09T07:52:33.89764Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ResidualBlock(nn.Module):\n    def __init__(self, in_c, out_c, activation):\n        super(ResidualBlock, self).__init__()\n        self.activation = activation\n        self.conv1 = nn.Sequential(\n            nn.MaxPool2d(kernel_size=2),\n            nn.Conv2d(in_c, out_c, kernel_size=3, groups=in_c, padding=\"same\"),\n            self.activation,\n            nn.Conv2d(out_c, out_c, kernel_size=1, padding=\"same\"),\n            self.activation,\n        )\n        self.conv2 = nn.Sequential(\n            nn.Conv2d(out_c, out_c, kernel_size=3, groups=in_c, padding=\"same\"),\n            self.activation,\n            nn.Conv2d(out_c, out_c, kernel_size=1, padding=\"same\"),\n            self.activation,\n        )\n        self.conv3 = nn.Sequential(\n            nn.Conv2d(out_c, out_c, kernel_size=3, groups=in_c, padding=\"same\"),\n            self.activation,\n            nn.Conv2d(out_c, out_c, kernel_size=1, padding=\"same\"),\n        )\n\n    def forward(self, x):\n        h0 = self.conv1(x)\n        h1 = self.conv2(h0)\n        h2 = self.activation(self.conv3(h1) + h0)\n        return h2\n\nclass CNNEmbedding(nn.Module):\n    def __init__(self, img_size, in_c, out_c, activation):\n        super(ResNetEmbedding, self).__init__()\n        self.activation = activation\n        self.input_conv = nn.Sequential(\n            nn.Conv2d(in_c, out_c, kernel_size=7, padding=\"same\"),\n            self.activation,\n        )\n        self.conv_blocks = nn.Sequential(\n            ResidualBlock(out_c * 2**0, out_c * 2**1, activation),\n            ResidualBlock(out_c * 2**1, out_c * 2**2, activation),\n            ResidualBlock(out_c * 2**2, out_c * 2**3, activation),\n            ResidualBlock(out_c * 2**3, out_c * 2**4, activation),\n        )\n\n    def forward(self, x):\n        x = self.input_conv(x)\n        x = self.conv_blocks(x)\n        return x\n\nclass PatchEmbedding(nn.Module):\n    def __init__(self, img_size, n_channels, embed_dim):\n        super(PatchEmbedding, self).__init__()\n        self.img_size = img_size\n        self.n_channels = n_channels\n        self.embed_dim = embed_dim\n        self.proj = nn.Conv2d(n_channels, embed_dim, kernel_size=1)\n\n    def forward(self, x):\n        # (B, in_c, P, P) > (B, D, P, P)\n        x = self.proj(x)\n        # (B, D, P, P) > (B, D, S)\n        x = x.flatten(2)\n        # (B, D, S) > (B, S, D)\n        x = x.transpose(1, 2)\n        return x\n\n# class PositionalEncoding(nn.Module):\n#     def __init__(self, vocab_size, embed_dim):\n#         super(PositionalEncoding, self).__init__()\n#         self.register_buffer('buf', torch.arange(vocab_size, dtype=torch.int64))\n#         self.pos_embedding = nn.Embedding(vocab_size, embed_dim)\n\n#     def forward(self, x):\n#         return x + self.pos_embedding(self.get_buffer('buf')[:x.shape[1]].unsqueeze(0))\n\nclass PositionalEncoding(nn.Module):\n    def __init__(self, seq_len, embed_dim):\n        super(PositionalEncoding, self).__init__()\n        pos_embedding = torch.zeros((seq_len, embed_dim), dtype=torch.float32)\n        pos_embedding[:, ::2] += torch.sin(torch.arange(seq_len)[:, None] / (10000 ** (torch.arange(0, embed_dim, 2) / seq_len)))\n        pos_embedding[:, 1::2] += torch.cos(torch.arange(seq_len)[:, None] / (10000 ** (torch.arange(1, embed_dim, 2) / seq_len)))\n        self.register_buffer('buf', pos_embedding)\n\n    def forward(self, x):\n        return x + self.get_buffer('buf')[:x.shape[1]].unsqueeze(0)\n\nclass DETR_CustomModel(nn.Module):\n    def __init__(self, n_quries, img_size, n_channels, out_channels, embed_dim, n_heads, ffn_dim, activation, n_encoders, n_decoders, n_coords, n_classes):\n        super(DETR_CustomModel, self).__init__()\n        self.n_quries = n_quries\n        self.img_size = img_size\n        self.n_channels = n_channels\n        self.embed_dim = embed_dim\n        self.n_heads = n_heads\n        self.ffn_dim = ffn_dim\n        self.activation = activation\n        self.n_encoders = n_encoders\n        self.n_decoders = n_decoders\n        self.n_classes = n_classes\n        self.cnn_embedding = CNNEmbedding(img_size, n_channels, out_channels, activation)\n        self.patch_embedding = PatchEmbedding(img_size, out_channels * 2**len(self.cnn_embedding.conv_blocks), embed_dim)\n        # self.patch_pos_embedding = PositionalEncoding(out_channels * 2**len(self.cnn_embedding.conv_blocks), embed_dim)\n        self.patch_pos_embedding = PositionalEncoding(out_channels * 2**len(self.cnn_embedding.conv_blocks), embed_dim)\n        self.query = nn.Parameter(torch.zeros(n_quries, embed_dim))\n        # self.query_pos_embedding = PositionalEncoding(n_quries, embed_dim)\n        self.query_pos_embedding = PositionalEncoding(n_quries, embed_dim)\n        self.encoder = nn.ModuleList(\n            nn.TransformerEncoderLayer(d_model=embed_dim, nhead=n_heads, dim_feedforward=ffn_dim, activation=activation, batch_first=True)\n            for _ in range(n_encoders)\n        )\n        self.decoder = nn.ModuleList(\n            nn.TransformerDecoderLayer(d_model=embed_dim, nhead=n_heads, dim_feedforward=ffn_dim, activation=activation, batch_first=True)\n            for _ in range(n_decoders)\n        )\n        self.head_reg = nn.Linear(embed_dim, n_coords)\n        self.head_cls = nn.Linear(embed_dim, n_classes)\n\n    def forward(self, x):\n        # embedding\n        x = self.cnn_embedding(x)\n        x = self.patch_embedding(x)\n        x_pos = self.patch_pos_embedding(x)\n        q = self.query.unsqueeze(0).repeat(x.shape[0], 1, 1)\n        q_pos = self.query_pos_embedding(q)\n        # encoding\n        for m in self.encoder:\n            x = m(x + x_pos)\n        # decoding\n        for m in self.decoder:\n            q = m(q + q_pos, x + x_pos)\n        # header\n        return {\"encode_values\": x, \"decode_values\": q, \"reg\": self.head_reg(q), \"cls\": self.head_cls(q)}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T08:24:59.92376Z","iopub.execute_input":"2024-12-09T08:24:59.924172Z","iopub.status.idle":"2024-12-09T08:24:59.944301Z","shell.execute_reply.started":"2024-12-09T08:24:59.924138Z","shell.execute_reply":"2024-12-09T08:24:59.943328Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"fold_pred = []\nfold_score = []\n\nfor fold in range(CFG.n_folds):\n    print(f\"\\n=== FOLD {fold} ===\")\n    # split train & valid\n    train_runs = df_annot.loc[df_annot[\"fold\"] != fold].index.unique()\n    valid_runs = df_annot.loc[df_annot[\"fold\"] == fold].index.unique()\n    # create ground truth df\n    df_train = df_annot.loc[df_annot.index.isin(train_runs), [\"obj\", \"x\", \"y\", \"z\", \"w\", \"h\", \"c\"]]\n    df_train[[\"x\", \"y\", \"z\"]] = (df_train[[\"x\", \"y\", \"z\"]].values / df_train[[\"w\", \"h\", \"c\"]]).values\n    for axis, val in CFG.max_coord_values.items():\n        df_train[axis] *= val\n    df_train[\"obj\"] = df_train[\"obj\"].map(CFG.id2label)\n    df_train = df_train[[\"obj\", \"x\", \"y\", \"z\"]]\n    df_valid = df_annot.loc[df_annot.index.isin(valid_runs), [\"obj\", \"x\", \"y\", \"z\", \"w\", \"h\", \"c\"]]\n    df_valid[[\"x\", \"y\", \"z\"]] = (df_valid[[\"x\", \"y\", \"z\"]].values / df_valid[[\"w\", \"h\", \"c\"]]).values\n    for axis, val in CFG.max_coord_values.items():\n        df_valid[axis] *= val\n    df_valid[\"obj\"] = df_valid[\"obj\"].map(CFG.id2label)\n    df_valid = df_valid[[\"obj\", \"x\", \"y\", \"z\"]]\n    # create dataset\n    train_ds = CustomDataset(\n        {run: df_img[run] for run in train_runs},\n        df_annot[df_annot.index.isin(train_runs)].groupby(\"run\").agg({\"x\": list, \"y\": list, \"z\": list, \"obj\": list}).apply(lambda x: np.stack(x, axis=-1), axis=1),\n        n_aug=CFG.n_aug,\n    )\n    valid_ds = CustomDataset(\n        {run: df_img[run] for run in valid_runs},\n        df_annot[df_annot.index.isin(valid_runs)].groupby(\"run\").agg({\"x\": list, \"y\": list, \"z\": list, \"obj\": list}).apply(lambda x: np.stack(x, axis=-1), axis=1),\n        n_aug=0,\n    )\n    print(\"iteration ->\", len(train_ds))\n    # create model\n    model_params = {\n        \"n_quries\": CFG.n_quries,\n        \"img_size\": CFG.img_size[0],\n        \"n_channels\": CFG.n_channels,\n        \"out_channels\": CFG.out_channels,\n        \"embed_dim\": CFG.embed_dim,\n        \"n_heads\": CFG.n_heads,\n        \"ffn_dim\": CFG.ffn_dim,\n        \"activation\": nn.LeakyReLU(),\n        \"n_encoders\": CFG.n_encoders,\n        \"n_decoders\": CFG.n_decoders,\n        \"n_coords\": CFG.n_coords,\n        \"n_classes\": len(CFG.label2id),\n    }\n    model = DETR_CustomModel(**model_params)\n    if torch.cuda.device_count() > 1:\n        model = nn.DataParallel(model)\n    model.to(device)\n    # training\n    return_score_dic = do_training(\n        fold, model, model_params, df_valid,\n        DataLoader(train_ds, collate_fn=collate_fn, batch_size=CFG.batch_size, shuffle=True),\n        DataLoader(valid_ds, collate_fn=collate_fn, batch_size=1, shuffle=False)\n    )\n    # validation\n    y_pred = infer_fn(model, DataLoader(train_ds, collate_fn=collate_fn, batch_size=1, shuffle=False))\n    df_pred = []\n    for run, coord, label in zip(valid_ds.runs, np.concatenate(y_pred[\"reg\"]), np.concatenate(y_pred[\"cls\"])):\n        df = pd.DataFrame(index=[run] * len(coord))\n        df[\"obj\"] = label.argmax(axis=-1)\n        df[[\"x\", \"y\", \"z\"]] = np.clip(coord, a_min=0.01, a_max=0.99)\n        df_pred.append(df)\n    df_pred = pd.concat(df_pred)\n    for axis, val in CFG.max_coord_values.items():\n        df_pred[axis] *= val\n    df_pred[\"obj\"] = df_pred[\"obj\"].map(CFG.id2label)\n    train_fbeta = get_fbeta(df_train, df_pred, particle_radius=CFG.particle_radius, weights=CFG.weights, beta=CFG.beta)\n    y_pred = infer_fn(model, DataLoader(valid_ds, collate_fn=collate_fn, batch_size=1, shuffle=False))\n    print(df_pred[\"obj\"].value_counts(normalize=True))\n    df_pred = []\n    for run, coord, label in zip(valid_ds.runs, np.concatenate(y_pred[\"reg\"]), np.concatenate(y_pred[\"cls\"])):\n        df = pd.DataFrame(index=[run] * len(coord))\n        df[\"obj\"] = label.argmax(axis=-1)\n        df[[\"x\", \"y\", \"z\"]] = np.clip(coord, a_min=0.01, a_max=0.99)\n        df_pred.append(df)\n    df_pred = pd.concat(df_pred)\n    for axis, val in CFG.max_coord_values.items():\n        df_pred[axis] *= val\n    df_pred[\"obj\"] = df_pred[\"obj\"].map(CFG.id2label)\n    valid_fbeta = get_fbeta(df_valid, df_pred, particle_radius=CFG.particle_radius, weights=CFG.weights, beta=CFG.beta)\n    print(df_pred[\"obj\"].value_counts(normalize=True))\n    fold_pred.append(df_pred)\n    # evaluation\n    return_score_dic[\"train_fbeta\"] = train_fbeta\n    return_score_dic[\"valid_fbeta\"] = valid_fbeta\n    fold_score.append(return_score_dic)\n    print(\"[SCORE]\")\n    print(pd.Series(fold_score[-1]))\n    del train_ds, valid_ds, model\n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T08:25:00.348994Z","iopub.execute_input":"2024-12-09T08:25:00.349711Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pickleIO(fold_pred, f\"fold_pred.pkl\", \"w\")\npickleIO(fold_score, f\"fold_score.pkl\", \"w\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Validation Summary","metadata":{}},{"cell_type":"code","source":"df_score = pd.DataFrame(fold_score)\ndf_score.loc[\"average\"] = df_score.mean()\ndf_score.round(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-09T08:02:47.124476Z","iopub.status.idle":"2024-12-09T08:02:47.124767Z","shell.execute_reply.started":"2024-12-09T08:02:47.124631Z","shell.execute_reply":"2024-12-09T08:02:47.124645Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}