{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Notes\n    This notebook was made to be trained locally on a 24 GB GPU. \n    The batch sizes were adjusted downward to account for lower memory P100 at Kaggle.\n    \n    Thanks to:\n    2.5D and \n    TTA post and model repos\n\n    Additional Info: \n    https://www.youtube.com/watch?v=YWfIx5ggeYI&t=3250s\n    https://www.youtube.com/watch?v=T0mWqsFrJpk&t=3320s\n    \n    Lowest mean = point of ink surface?\n    \n    https://paperswithcode.com/lib/timm\n    \n    Takeaways:\n    Widening tile size for val is useful as annoying padding gets reduced,\n    => Not sure if general but seems to improve up to a point only (wanna still use overlap) and increases optimal threshold\n    For some reason rotate90 sucks","metadata":{}},{"cell_type":"markdown","source":"# Fundamentals","metadata":{"tags":[]}},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nclass CFG: \n    broaden = 0\n    remove_background_only = True\n    # ============== comp exp name =============\n    comp_name = 'vesuvius'\n    comp_dir_path = '/kaggle/input/'\n    comp_folder_name = 'vesuvius-cv'\n    comp_dataset_path = f'{comp_dir_path}{comp_folder_name}/'  \n    exp_name = 'vesuvius_2d_tta'\n    # ============== pred target =============\n    target_size = 1\n    # ============== model cfg =============\n    backbone =  \"efficientnet-b7\" #efficientnet-b7 timm-efficientnet-b8 timm-efficientnet-l2\n    in_chans = 8 # 65\n    \n    # ============== training cfg =============\n    tile_size_train = 512\n    tile_size_val = 1536\n    stride_train = tile_size_train // 2\n    stride_val = tile_size_val // 2\n    train_batch_size = 6 # 16\n    \n    valid_batch_size = 1\n    use_amp = True\n    # ============== fold =============\n    metric_direction = 'maximize'  # maximize, 'minimize'\n    # ============== fixed =============\n    pretrained = True\n    inf_weight = 'best'  # 'best'\n    max_grad_norm = 1000\n    num_workers = 4\n    # ============== set dataset path =============\n    print('set dataset path')\n    outputs_path = f'outputs/{comp_name}/{exp_name}/'\n    submission_dir = outputs_path + 'submissions/'\n    submission_path = submission_dir + f'submission_{exp_name}.csv'\n    model_dir = outputs_path + f'{comp_name}-models/'\n    figures_dir = outputs_path + 'figures/'\n    log_dir = outputs_path + 'logs/'\n    log_path = log_dir + f'{exp_name}.txt'\n    # ============== augmentation =============\n    valid_aug_list = [\n        A.ToFloat(),\n        ToTensorV2(transpose_mask=True),   \n    ]","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-07-20T12:45:10.134916Z","iopub.execute_input":"2023-07-20T12:45:10.135290Z","iopub.status.idle":"2023-07-20T12:45:15.324626Z","shell.execute_reply.started":"2023-07-20T12:45:10.135259Z","shell.execute_reply":"2023-07-20T12:45:15.323634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Libraries","metadata":{}},{"cell_type":"code","source":"!pip install segmentation-models-pytorch\n!pip install warmup_scheduler","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:15.326573Z","iopub.execute_input":"2023-07-20T12:45:15.327489Z","iopub.status.idle":"2023-07-20T12:45:46.578554Z","shell.execute_reply.started":"2023-07-20T12:45:15.327452Z","shell.execute_reply":"2023-07-20T12:45:46.577350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score, accuracy_score, f1_score, log_loss\nimport os\nimport gc\nimport sys\nimport time\nimport random\nfrom pathlib import Path\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.cuda.amp import autocast, GradScaler\nimport segmentation_models_pytorch as smp\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\nfrom sklearn.metrics import fbeta_score\nfrom warmup_scheduler import GradualWarmupScheduler\nimport torchvision\nimport wandb","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-07-20T12:45:46.580449Z","iopub.execute_input":"2023-07-20T12:45:46.580836Z","iopub.status.idle":"2023-07-20T12:45:48.850365Z","shell.execute_reply.started":"2023-07-20T12:45:46.580789Z","shell.execute_reply":"2023-07-20T12:45:48.849345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Set Seed","metadata":{}},{"cell_type":"code","source":"def set_seed(seed):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n\ndef worker_init_fn(worker_id):\n    worker_seed = seed + worker_id\n    np.random.seed(worker_seed)\n    random.seed(worker_seed)\n    \nseed = 42\nset_seed(seed)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:48.852886Z","iopub.execute_input":"2023-07-20T12:45:48.853179Z","iopub.status.idle":"2023-07-20T12:45:48.865945Z","shell.execute_reply.started":"2023-07-20T12:45:48.853153Z","shell.execute_reply":"2023-07-20T12:45:48.864904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Login and Display","metadata":{}},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"WANDB\")","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:48.867454Z","iopub.execute_input":"2023-07-20T12:45:48.868387Z","iopub.status.idle":"2023-07-20T12:45:49.122462Z","shell.execute_reply.started":"2023-07-20T12:45:48.868353Z","shell.execute_reply":"2023-07-20T12:45:49.121454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def wandb_init(cv, name, group):\n    config = {k:v for k,v in dict(vars(CFG)).items() if '__' not in k}\n    config.update({\"cv\":int(cv)})\n    run = wandb.init(project=\"Vesuv\",\n                     entity='ralf-c-kinkel',\n                     name=name,\n                     config=config,\n                     group=group,\n                     save_code=True,)\n    return run\n\ndef log_wb(loss, val_loss, lr, score, combo_score, elapsed):\n    elapsed_time = time.time()  # Calculate elapsed time\n    \n    # Log the metrics\n    wandb.log({\n        'loss': loss,\n        'val_loss': val_loss,\n        'learning_rate': lr,\n        'score': score,\n        'combo_score': combo_score,\n        'elapsed_time': elapsed,\n    })\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\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\ndef init_logger(log_file):\n    from logging import getLogger, INFO, FileHandler, Formatter, StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\ndef make_dirs(cfg):\n    for dir in [cfg.model_dir, cfg.figures_dir, cfg.submission_dir, cfg.log_dir]:\n        os.makedirs(dir, exist_ok=True)   \ndef cfg_init(cfg, mode='train'):\n    if mode == 'train':\n        make_dirs(cfg)\n\nif wandb.run is not None:\n    wandb.finish()\nwandb.login(key=secret_value_0)\ntry: \n    Logger.info('-------- exp_info -----------------')\nexcept:\n    cfg_init(CFG)\n    Logger = init_logger(log_file=CFG.log_path)\n    Logger.info('-------- exp_info -----------------')\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:49.124020Z","iopub.execute_input":"2023-07-20T12:45:49.124719Z","iopub.status.idle":"2023-07-20T12:45:52.059352Z","shell.execute_reply.started":"2023-07-20T12:45:49.124686Z","shell.execute_reply":"2023-07-20T12:45:52.058266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Functions and Classes","metadata":{}},{"cell_type":"code","source":"def empty():\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.060854Z","iopub.execute_input":"2023-07-20T12:45:52.061460Z","iopub.status.idle":"2023-07-20T12:45:52.065825Z","shell.execute_reply.started":"2023-07-20T12:45:52.061402Z","shell.execute_reply":"2023-07-20T12:45:52.064798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_image_mask(fragment_id, offset=0):              # Modded to png and train_new\n    images = []\n    mid = 32\n    start = mid - CFG.in_chans // 2\n    end = mid + CFG.in_chans // 2\n    idxs = range(start+offset, end+offset)\n    for i in tqdm(idxs):\n        image = cv2.imread(CFG.comp_dataset_path + f\"train_new/{fragment_id}/surface_volume/{i:02}.png\", 0)\n        images.append(image)\n    images = np.stack(images, axis=2)\n    mask = cv2.imread(CFG.comp_dataset_path + f\"train_new/{fragment_id}/inklabels.png\", 0)\n    \n    mask = mask.astype('float32')\n    mask /= 255.0\n    \n    foreground = cv2.imread(CFG.comp_dataset_path + f\"train_new/{fragment_id}/mask.png\", 0)\n    foreground = foreground.astype('float32')\n    foreground /= 255.0\n    foreground = foreground.astype('uint8')\n    return images, mask, foreground\n\ndef pad_all(image, mask, foreground, ts):\n    pad0 = (ts - image.shape[0] % ts)\n    pad1 = (ts - image.shape[1] % ts)\n    image = np.pad(image, [(0, pad0), (0, pad1), (0, 0)], constant_values=0)\n    mask = np.pad(mask, [(0, pad0), (0, pad1)], constant_values=0)\n    foreground = np.pad(foreground, [(0, pad0), (0, pad1)], constant_values=0)\n    return image, mask, foreground\n\ndef get_train_valid_dataset(id, offset=0):\n    train_images = []\n    train_masks = []\n    valid_images = []\n    valid_masks = []\n    valid_xyxys = []\n    \n    id = str(id)\n    valid_id = id[0]\n    \n    for fragment_id in range(1, 4):\n        fragment_id = str(fragment_id)\n        \n        if fragment_id == valid_id:\n            # For validation set, we load the specific section of the image\n            image, mask, foreground = read_image_mask(f'{id}')\n            image, mask, foreground = pad_all(image, mask, foreground, CFG.tile_size_val)\n            \n            x1_list = list(range(0, image.shape[1]-CFG.tile_size_val+1, CFG.stride_val))\n            y1_list = list(range(0, image.shape[0]-CFG.tile_size_val+1, CFG.stride_val))\n            \n            for y1 in y1_list:\n                for x1 in x1_list:\n                    y2 = y1 + CFG.tile_size_val\n                    x2 = x1 + CFG.tile_size_val\n                    \n                    if np.sum(foreground[y1:y2, x1:x2]) > CFG.remove_background_only - 1:\n                        valid_images.append(image[y1:y2, x1:x2])\n                        valid_masks.append(mask[y1:y2, x1:x2, None])\n                        valid_xyxys.append([x1, y1, x2, y2])\n                        \n            valid_mask_gt = cv2.imread(CFG.comp_dataset_path + f\"train_new/{id}/inklabels.png\", 0)\n            valid_mask_gt = valid_mask_gt / 255\n            pad0 = (CFG.tile_size_val - valid_mask_gt.shape[0] % (CFG.tile_size_val))\n            pad1 = (CFG.tile_size_val - valid_mask_gt.shape[1] % (CFG.tile_size_val))\n            valid_mask_gt = np.pad(valid_mask_gt, [(0, pad0 ), (0, pad1)], constant_values=0)\n                        \n            # Load train images for rest of validation\n            if len(id) == 2:\n                stitch_dict = {'11':[['12','13']], '12':[['11'],['13']], '13':[['11','12']],\n                   '21':[['22','23','24','25']], '22':[['21'],['23','24','25']],  '23':[['21','22'],['24','25']], \n                       '24':[['21','22','23'],['25']], '25':[['21','22','23','24']], \n                   '31':[['32']], '32':[['31']]}\n                stitches = stitch_dict[id]\n                for stitch in stitches:\n                    # Concatenate all in stitch along 0 axis, image, mask and foreground\n                    images = []\n                    masks = []\n                    foregrounds = []\n                    for part in stitch:\n                        image, mask, foreground = read_image_mask(f'{part}')\n                        images.append(image)\n                        masks.append(mask)\n                        foregrounds.append(foreground)\n                     # Concatenate along y axis (axis=0)\n                    image = np.concatenate(images, axis=0)\n                    mask = np.concatenate(masks, axis=0)\n                    foreground = np.concatenate(foregrounds, axis=0)\n                    image, mask, foreground = pad_all(image, mask, foreground, CFG.tile_size_train)\n                    plt.figure()\n                    plt.imshow(mask)\n    \n                    x1_list = list(range(0, image.shape[1]-CFG.tile_size_train+1, CFG.stride_train))\n                    y1_list = list(range(0, image.shape[0]-CFG.tile_size_train+1, CFG.stride_train))\n                    for y1 in y1_list:\n                        for x1 in x1_list:\n                            y2 = y1 + CFG.tile_size_train\n                            x2 = x1 + CFG.tile_size_train\n                            if np.sum(foreground[y1:y2, x1:x2]) > CFG.remove_background_only - 1:\n                                train_images.append(image[y1:y2, x1:x2])\n                                train_masks.append(mask[y1:y2, x1:x2, None])\n                             \n        else:\n            # For training set, we load the entire image and specific sections from other images\n            for offset in range(-CFG.broaden,1+CFG.broaden):\n                image, mask, foreground = read_image_mask(fragment_id, offset)  \n                image, mask, foreground = pad_all(image, mask, foreground, CFG.tile_size_train)\n                x1_list = list(range(0, image.shape[1]-CFG.tile_size_train+1, CFG.stride_train))\n                y1_list = list(range(0, image.shape[0]-CFG.tile_size_train+1, CFG.stride_train))\n                for y1 in y1_list:\n                    for x1 in x1_list:\n                        y2 = y1 + CFG.tile_size_train\n                        x2 = x1 + CFG.tile_size_train\n                        if np.sum(foreground[y1:y2, x1:x2]) > CFG.remove_background_only - 1:\n                            train_images.append(image[y1:y2, x1:x2])\n                            train_masks.append(mask[y1:y2, x1:x2, None])\n                            \n    return train_images, train_masks, valid_images, valid_masks, np.stack(valid_xyxys), valid_mask_gt","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.067317Z","iopub.execute_input":"2023-07-20T12:45:52.067891Z","iopub.status.idle":"2023-07-20T12:45:52.101833Z","shell.execute_reply.started":"2023-07-20T12:45:52.067856Z","shell.execute_reply":"2023-07-20T12:45:52.100969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def determine_values(lr_mult, batches_mult, epochs=100, val_calcs=20, train_batch_size=CFG.train_batch_size, tile_size_train=CFG.tile_size_train):\n    lr_base = 512**2 / tile_size_train**2 * train_batch_size * 1e-5\n    val_every_epochs = epochs // val_calcs\n    batches_per_epoch_base = int(1000 * 512**2 / tile_size_train**2 / train_batch_size)\n    \n    lr = lr_base * lr_mult\n    batches_per_epoch = int(batches_per_epoch_base * batches_mult)\n    return lr, batches_per_epoch, val_every_epochs","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.103856Z","iopub.execute_input":"2023-07-20T12:45:52.105048Z","iopub.status.idle":"2023-07-20T12:45:52.118524Z","shell.execute_reply.started":"2023-07-20T12:45:52.105015Z","shell.execute_reply":"2023-07-20T12:45:52.117560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(data, cfg):\n    if data == 'train':\n        aug = A.Compose(train_aug_list)\n    elif data == 'valid':\n        aug = A.Compose(CFG.valid_aug_list)\n    return aug\nclass CustomDataset(Dataset):\n    def __init__(self, images, cfg, labels=None, transform=None):\n        self.images = images\n        self.cfg = cfg\n        self.labels = labels\n        self.transform = transform\n    def __len__(self):\n        return len(self.images)\n    def __getitem__(self, idx):\n        image = self.images[idx]\n        label = self.labels[idx]\n        if self.transform:\n            data = self.transform(image=image, mask=label)\n            image = data['image']\n            label = data['mask']\n        return image, label\n\nclass EvalDataset(Dataset):\n    def __init__(self, images, cfg, transform=None):\n        self.images = images\n        self.cfg = cfg\n        self.transform = transform\n    def __len__(self):\n        return len(self.images)\n    def __getitem__(self, idx):\n        image = self.images[idx]\n        if self.transform:\n            data = self.transform(image=image)\n            image = data['image']\n        return image","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.123164Z","iopub.execute_input":"2023-07-20T12:45:52.123465Z","iopub.status.idle":"2023-07-20T12:45:52.134163Z","shell.execute_reply.started":"2023-07-20T12:45:52.123437Z","shell.execute_reply":"2023-07-20T12:45:52.133186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    def __init__(self, cfg, weight=None):\n        super().__init__()\n        self.cfg = cfg\n        self.encoder = smp.UnetPlusPlus(\n            encoder_name=cfg.backbone, \n            encoder_weights=weight,\n            in_channels=cfg.in_chans,\n            classes=cfg.target_size,\n            activation=None)\n\n    def forward(self, image):\n        output = self.encoder(image)\n        return output\ndef build_model(cfg, weight='imagenet'): # noisy-student\n    print('backbone', cfg.backbone)\n    model = CustomModel(cfg, weight)\n    return model","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-07-20T12:45:52.135501Z","iopub.execute_input":"2023-07-20T12:45:52.136181Z","iopub.status.idle":"2023-07-20T12:45:52.145447Z","shell.execute_reply.started":"2023-07-20T12:45:52.136145Z","shell.execute_reply":"2023-07-20T12:45:52.144461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BCELoss = smp.losses.SoftBCEWithLogitsLoss()\ndef criterion(y_pred, y_true):\n    return BCELoss(y_pred, y_true)","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-07-20T12:45:52.148746Z","iopub.execute_input":"2023-07-20T12:45:52.149016Z","iopub.status.idle":"2023-07-20T12:45:52.159092Z","shell.execute_reply.started":"2023-07-20T12:45:52.148995Z","shell.execute_reply":"2023-07-20T12:45:52.158171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fn(train_loader, batches_per_epoch, model, criterion, optimizer, device):\n    model.train()\n    scaler = GradScaler(enabled=CFG.use_amp)\n    losses = AverageMeter()\n    with tqdm(total=batches_per_epoch) as pbar:\n        for step, (images, labels) in enumerate(train_loader):\n            if step + 1 == batches_per_epoch:\n                pbar.total = step + 1\n                pbar.update()\n                break\n            pbar.update()\n            \n            images = images.to(device)\n            labels = labels.to(device)\n            batch_size = labels.size(0)\n            with autocast(CFG.use_amp):\n                y_preds = model(images)\n                loss = criterion(y_preds, labels)  # CHANGE\n            losses.update(loss.item(), batch_size)\n            scaler.scale(loss).backward()\n            grad_norm = torch.nn.utils.clip_grad_norm_(\n                model.parameters(), CFG.max_grad_norm)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n    return losses.avg\n    \ndef valid_fn(valid_loader, model, criterion, device, valid_xyxys, valid_mask_gt):\n    mask_pred = np.zeros(valid_mask_gt.shape)\n    mask_count = np.zeros(valid_mask_gt.shape)\n    model.eval()\n    losses = AverageMeter()\n    for step, (images, labels) in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        with torch.no_grad():\n            y_preds = model(images)\n            loss = criterion(y_preds, labels)\n        losses.update(loss.item(), batch_size)\n        # make whole mask\n        y_preds = torch.sigmoid(y_preds).to('cpu').numpy()\n        start_idx = step*CFG.valid_batch_size\n        end_idx = start_idx + batch_size\n        for i, (x1, y1, x2, y2) in enumerate(valid_xyxys[start_idx:end_idx]):\n            mask_pred[y1:y2, x1:x2] += y_preds[i].squeeze(0)\n            mask_count[y1:y2, x1:x2] += np.ones((CFG.tile_size_val, CFG.tile_size_val))\n    print(f'mask_count_min: {mask_count.min()}')\n    mask_pred /= mask_count\n    return losses.avg, mask_pred","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-07-20T12:45:52.162300Z","iopub.execute_input":"2023-07-20T12:45:52.162678Z","iopub.status.idle":"2023-07-20T12:45:52.176995Z","shell.execute_reply.started":"2023-07-20T12:45:52.162655Z","shell.execute_reply":"2023-07-20T12:45:52.175677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fbeta_numpy(targets, preds, beta=0.5, smooth=1e-5):\n    y_true_count = targets.sum()\n    ctp = preds[targets==1].sum()\n    cfp = preds[targets==0].sum()\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp + smooth)\n    c_recall = ctp / (y_true_count + smooth)\n    dice = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall + smooth)\n    return dice\ndef calc_fbeta(mask, mask_pred, r=range(25, 81, 5)):\n    mask = mask.astype(int).flatten()\n    mask_pred = mask_pred.flatten()\n    best_th = 0\n    best_dice = 0\n    for th in np.array(r) / 100:   \n        dice = fbeta_numpy(mask, (mask_pred >= th).astype(int), beta=0.5)\n        print(f'th: {th}, fbeta: {dice}')\n        if dice > best_dice:\n            best_dice = dice\n            best_th = th\n    Logger.info(f'best_th: {best_th}, fbeta: {best_dice}')\n    return best_dice, best_th\ndef calc_cv(mask_gt, mask_pred):\n    best_dice, best_th = calc_fbeta(mask_gt, mask_pred)\n    return best_dice, best_th","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-07-20T12:45:52.178453Z","iopub.execute_input":"2023-07-20T12:45:52.179098Z","iopub.status.idle":"2023-07-20T12:45:52.191249Z","shell.execute_reply.started":"2023-07-20T12:45:52.179067Z","shell.execute_reply":"2023-07-20T12:45:52.190289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_eval_transforms():\n    valid_aug_list = [\n        A.ToFloat(),\n        ToTensorV2(transpose_mask=True)]\n    aug = A.Compose(valid_aug_list)\n    return aug","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.192662Z","iopub.execute_input":"2023-07-20T12:45:52.192986Z","iopub.status.idle":"2023-07-20T12:45:52.202515Z","shell.execute_reply.started":"2023-07-20T12:45:52.192955Z","shell.execute_reply":"2023-07-20T12:45:52.201584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EvalModel(nn.Module):\n    def __init__(self, cfg, weight=None):\n        super().__init__()\n        self.cfg = cfg\n        self.encoder = smp.UnetPlusPlus(\n            encoder_name=cfg.backbone, \n            encoder_weights=weight,\n            in_channels=cfg.in_chans,\n            classes=cfg.target_size,\n            activation=None)\n    def forward(self, image):\n        output = self.encoder(image)\n        return output\ndef build_eval_model(cfg, weight=\"imagenet\"):\n    print('backbone', cfg.backbone)\n    model = EvalModel(cfg, weight)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.204060Z","iopub.execute_input":"2023-07-20T12:45:52.204456Z","iopub.status.idle":"2023-07-20T12:45:52.216804Z","shell.execute_reply.started":"2023-07-20T12:45:52.204425Z","shell.execute_reply":"2023-07-20T12:45:52.215752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_eval_image_mask(fragment_id, broaden=0):\n    image_stack = []\n    mid = 65 // 2\n    start = mid - CFG.in_chans // 2\n    end = mid + CFG.in_chans // 2\n    idxs = range(start-broaden, end+broaden)\n    for i in tqdm(idxs):\n        image = cv2.imread(CFG.comp_dataset_path + ev_path + f\"/surface_volume/{i:02}.png\", 0)\n        image = np.pad(image, [(border, border), (border, border)], constant_values=0) #mode='mean') # Pad the border region\n        pad0 = min((tile_size_val - image.shape[0] % tile_size_val), (tile_size_val - image.shape[0] % tile_size_val)) # Pad for batch loading\n        pad1 = min((tile_size_val - image.shape[1] % tile_size_val), (tile_size_val - image.shape[1] % tile_size_val)) # Pad for batch loading\n        image = np.pad(image, [(0, pad0), (0, pad1)], constant_values=0) #mode='mean')\n        image_stack.append(image)\n    image = np.stack(image_stack, axis=2)\n    mask = cv2.imread(CFG.comp_dataset_path + ev_path + f\"/inklabels.png\", 0)\n    mask = mask.astype('float32')\n    mask /= 255.0\n    return image, mask, image.shape[:2]","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.219310Z","iopub.execute_input":"2023-07-20T12:45:52.219922Z","iopub.status.idle":"2023-07-20T12:45:52.229564Z","shell.execute_reply.started":"2023-07-20T12:45:52.219892Z","shell.execute_reply":"2023-07-20T12:45:52.228454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(data, cfg):\n    if data == 'train':\n        aug = A.Compose(train_aug_list)\n    elif data == 'valid':\n        aug = A.Compose(CFG.valid_aug_list)\n    return aug\nclass CustomDataset(Dataset):\n    def __init__(self, images, cfg, labels=None, transform=None):\n        self.images = images\n        self.cfg = cfg\n        self.labels = labels\n        self.transform = transform\n    def __len__(self):\n        return len(self.images)\n    def __getitem__(self, idx):\n        image = self.images[idx]\n        label = self.labels[idx]\n        if self.transform:\n            data = self.transform(image=image, mask=label)\n            image = data['image']\n            label = data['mask']\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.230832Z","iopub.execute_input":"2023-07-20T12:45:52.231375Z","iopub.status.idle":"2023-07-20T12:45:52.241906Z","shell.execute_reply.started":"2023-07-20T12:45:52.231344Z","shell.execute_reply":"2023-07-20T12:45:52.240653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from albumentations.core.transforms_interface import ImageOnlyTransform\nclass ReverseChannelOrder(ImageOnlyTransform):\n    \"\"\"Reverse the channel order of an image.\"\"\"\n    def apply(self, img, **params):\n        return img[..., ::-1]","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.243731Z","iopub.execute_input":"2023-07-20T12:45:52.244152Z","iopub.status.idle":"2023-07-20T12:45:52.255563Z","shell.execute_reply.started":"2023-07-20T12:45:52.244120Z","shell.execute_reply":"2023-07-20T12:45:52.254511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Plotting","metadata":{}},{"cell_type":"code","source":"plotting = False","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-07-20T12:45:52.257151Z","iopub.execute_input":"2023-07-20T12:45:52.258154Z","iopub.status.idle":"2023-07-20T12:45:52.265978Z","shell.execute_reply.started":"2023-07-20T12:45:52.258077Z","shell.execute_reply":"2023-07-20T12:45:52.264854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if plotting:\n    train_dataset[0][0].shape\n    plot_dataset = CustomDataset(\n        train_images, CFG, labels=train_masks)\n    transform = CFG.train_aug_list\n    transform = A.Compose(\n        [t for t in transform if not isinstance(t, (A.Normalize, ToTensorV2))])\n    plot_count = 0\n    for i in range(2000,3000):\n        image, mask = plot_dataset[i]\n        data = transform(image=image, mask=mask)\n        aug_image = data['image']\n        aug_mask = data['mask']\n        if mask.sum() == 0:\n            continue\n        fig, axes = plt.subplots(1, 4, figsize=(15, 8))\n        axes[0].imshow(image[..., 0], cmap=\"gray\")\n        axes[1].imshow(mask, cmap=\"gray\")\n        axes[2].imshow(aug_image[..., 0], cmap=\"gray\")\n        axes[3].imshow(aug_mask, cmap=\"gray\")\n        plt.savefig(CFG.figures_dir + f'aug_fold_{CFG.valid_id}_{plot_count}.png')\n        plot_count += 1\n        if plot_count == 5:\n            break\n    del plot_dataset\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.267271Z","iopub.execute_input":"2023-07-20T12:45:52.267787Z","iopub.status.idle":"2023-07-20T12:45:52.278085Z","shell.execute_reply.started":"2023-07-20T12:45:52.267756Z","shell.execute_reply":"2023-07-20T12:45:52.277290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Is randomrotate90 (1-3 times) with 2 times same as vertical flip?\nif plotting:\n    batch_index = 20\n    in_b_index = 5\n    train_aug_list = [#A.HorizontalFlip(p=1),\n                #A.VerticalFlip(p=1),\n                #A.RandomRotate90(),\n                #A.RandomGamma (gamma_limit=(80, 120), eps=None, always_apply=False, p=1),\n                #A.ShiftScaleRotate(scale_limit=0.3, rotate_limit=[30,40], p=1), #0.1, 45\n                A.ShiftScaleRotate(rotate_limit=[90,90], p=1), #0.1, 45\n                #A.GridDistortion(num_steps=4, distort_limit=0.3, p=0),\n                A.ToFloat(),\n                ToTensorV2(transpose_mask=True)]\n    train_dataset = CustomDataset(train_images, CFG, labels=train_masks, transform=get_transforms(data='train', cfg=CFG))\n    train_loader = DataLoader(train_dataset,batch_size=CFG.train_batch_size,shuffle=False,num_workers=CFG.num_workers, pin_memory=True, drop_last=True, worker_init_fn=worker_init_fn)\n    for i, data in enumerate(train_loader):\n        if i == batch_index:\n            images, masks = data[0], data[1]\n            break\n    images, masks = data[0], data[1]\n\n    # Choose a random image from the batch\n    image, mask = images[in_b_index], masks[in_b_index]\n\n    # Convert the tensor to a PIL image\n    image = torchvision.transforms.ToPILImage()(image[3:6])\n    print(image)\n\n    # Convert the mask to a numpy array\n    mask = mask.squeeze().numpy()\n\n    # Plot the image and mask\n    fig, ax = plt.subplots(1, 2, figsize=(10, 5))\n    ax[0].imshow(image)\n    ax[1].imshow(mask)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-20T12:45:52.279110Z","iopub.execute_input":"2023-07-20T12:45:52.280079Z","iopub.status.idle":"2023-07-20T12:45:52.290973Z","shell.execute_reply.started":"2023-07-20T12:45:52.280045Z","shell.execute_reply":"2023-07-20T12:45:52.290374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"cross_val = [1] \n\nepochs = 101\nweight_decays = [0.8]\nlr_mults = [4]\nbatch_mults = [0.3]        # Epoch multiplier atm\ngroup = \"efficientnet-b7-scratch-train\"   # pct_start back to 0.05\n\nbest = {}\nfor cv in cross_val:\n    total_best_score = -1\n    total_best_loss = 100\n    total_best_combo = -1\n    \n    # Reload train and validation images, labels etc for each cross validation\n    train_images, train_masks, valid_images, valid_masks, valid_xyxys, valid_mask_gt = get_train_valid_dataset(cv, offset=2) \n\n    # Train augmentations\n    train_aug_list = [\n                #8 Different orientation rotation combinations\n                A.HorizontalFlip(p=0.5),\n                A.VerticalFlip(p=0.5),\n                #A.RandomRotate90(p=0.75), \n        \n                A.RandomBrightnessContrast(p=0.75),\n                A.OneOf([A.GaussNoise(var_limit=[10, 50]), A.GaussianBlur(), A.MotionBlur()], p=0.4),\n                A.RandomGamma (gamma_limit=(150, 225), eps=None, always_apply=False, p=0.75),\n                A.ShiftScaleRotate(scale_limit=0.4, rotate_limit=90, p=0.90), #0.1, 45\n                A.GridDistortion(num_steps=4, distort_limit=0.3, p=0.75),\n                A.ToFloat(),\n                ToTensorV2(transpose_mask=True)]\n\n    # Valid stuff can be set up here, the train stuff needs to be setup for changes in training augmentation\n    valid_dataset = CustomDataset(valid_images, CFG, labels=valid_masks, transform=get_transforms(data='valid', cfg=CFG))\n    valid_loader = DataLoader(valid_dataset,batch_size=CFG.valid_batch_size,shuffle=False,num_workers=CFG.num_workers, pin_memory=True, drop_last=False, worker_init_fn=worker_init_fn)\n    train_dataset = CustomDataset(train_images, CFG, labels=train_masks, transform=get_transforms(data='train', cfg=CFG))\n    train_loader = DataLoader(train_dataset,batch_size=CFG.train_batch_size,shuffle=True,num_workers=CFG.num_workers, pin_memory=True, drop_last=True, worker_init_fn=worker_init_fn)               \n\n    # Iterate over learning rates\n    for lr_mult in lr_mults:\n        # Iterate over weight decays\n        for weight_decay in weight_decays:\n            # Iterate over other stuff\n            for batch_mult in batch_mults:\n                lr, batches_per_epoch, val_every_epochs = determine_values(lr_mult, batch_mult)\n                \n                name = str(cv) + \" \" +  str(lr_mult) + \" \" + str(weight_decay) + \" \" + str(batch_mult) \n                run = wandb_init(cv, name, group)        #WB\n                \n                #If data augmentation gets iterated over, put train_dataset and train_loader here            \n\n                model = build_model(CFG, weight='imagenet')    #'imagenet'\n                #check_point = torch.load('submit/efficientnet-b7_1_best_champ_65+.pth', map_location=torch.device('cpu'))\n                #weights = check_point['model']\n                #model.load_state_dict(weights)\n                model.to(device)\n                \n                optimizer = AdamW(model.parameters(), lr=lr,  weight_decay=weight_decay)\n                print(weight_decay)\n                scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=lr, steps_per_epoch=1, epochs=epochs, pct_start=0.05) \n    \n                best_score = -1 \n                best_loss = 100\n                best_combo = -1\n                for epoch in range(epochs):\n                    start_time = time.time()\n                    avg_loss = train_fn(train_loader, batches_per_epoch, model, criterion, optimizer, device)\n                    scheduler.step()\n                    \n                    if (epoch) % val_every_epochs == 0:\n                        avg_val_loss, mask_pred = valid_fn(valid_loader, model, criterion, device, valid_xyxys, valid_mask_gt)\n                        best_dice, best_th = calc_cv(valid_mask_gt, mask_pred)\n                        score = best_dice\n                        comboscore = score / avg_val_loss ** (1/3)\n                                     \n                        if score > best_score:\n                            best_score = score\n                            Logger.info(f'Epoch {epoch+1} - Save Best Score: {best_score:.4f}')\n                            if best_score > total_best_score:\n                                torch.save({'model': model.state_dict(),'preds': mask_pred}, CFG.model_dir + f'{CFG.backbone}_{cv}_best_score.pth')\n                                total_best_score = best_score\n                        if comboscore > best_combo:\n                            best_combo = comboscore\n                            Logger.info(f'Epoch {epoch+1} - Save Best Combo: {best_combo:.4f}')\n                            if best_combo > total_best_combo:\n                                torch.save({'model': model.state_dict(),'preds': mask_pred}, CFG.model_dir + f'{CFG.backbone}_{cv}_best_combo.pth')\n                                total_best_combo = best_combo\n                        if avg_val_loss < best_loss:\n                            best_loss = avg_val_loss\n                            Logger.info(f'Epoch {epoch+1} - Best Loss: {best_loss:.4f}')\n                            if best_loss < total_best_loss:\n                                total_best_loss = best_loss\n\n                    # Logging\n                    elapsed = time.time() - start_time\n                    lr_current = optimizer.param_groups[0]['lr']\n                    Logger.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n                    Logger.info(f'Epoch {epoch+1} - avgScore: {score:.4f}')\n                    Logger.info(f'Epoch {epoch+1} - ComboScore: {comboscore:.5f}')\n                    Logger.info(f'Epoch {epoch+1} - LR: {lr_current:.5f}')\n                    log_wb(avg_loss, avg_val_loss, lr_current, score, comboscore, elapsed)\n        \n                best[(lr, weight_decay, batch_mult, cv)] = (best_score, best_loss)\n                wandb.run.finish()    #WB         \nprint(best)","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-07-20T12:45:52.292430Z","iopub.execute_input":"2023-07-20T12:45:52.293047Z","iopub.status.idle":"2023-07-20T13:08:23.830328Z","shell.execute_reply.started":"2023-07-20T12:45:52.293016Z","shell.execute_reply":"2023-07-20T13:08:23.829097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Eval","metadata":{"tags":[]}},{"cell_type":"markdown","source":"    Optimizing for score: Slight overfit (?)\n    Optimizing for loss: Less good results, but less overfit on val set => Improvements from TTA","metadata":{}},{"cell_type":"code","source":"def most_certain_fraction(mask_pred, scroll_size, fraction_goal):\n    for th in range(0,501):  # It should never be close to reaching 100!\n        possible_pred = mask_pred > 0.01 * th\n        if np.sum(possible_pred) / scroll_size < fraction_goal:\n            print(th*0.01)\n            return possible_pred","metadata":{"execution":{"iopub.status.busy":"2023-07-20T13:08:23.835173Z","iopub.execute_input":"2023-07-20T13:08:23.836210Z","iopub.status.idle":"2023-07-20T13:08:23.843416Z","shell.execute_reply.started":"2023-07-20T13:08:23.836178Z","shell.execute_reply":"2023-07-20T13:08:23.842293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fragment_id = 1\ntile_size_val_while_train = 256 * 6\n\nbroaden = 0\ntile_size_val = 256 * 6\nborder = tile_size_val // 4\nvalid_batch_size = 1\nstride = tile_size_val // 2\nev_path = f\"train_new/{fragment_id}\"","metadata":{"execution":{"iopub.status.busy":"2023-07-20T13:08:23.845042Z","iopub.execute_input":"2023-07-20T13:08:23.845675Z","iopub.status.idle":"2023-07-20T13:08:23.866032Z","shell.execute_reply.started":"2023-07-20T13:08:23.845641Z","shell.execute_reply":"2023-07-20T13:08:23.865030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"load_id = fragment_id\n\n#check_point = torch.load('/kaggle/working/outputs/vesuvius/vesuvius_2d_tta/vesuvius-models/' + f'{CFG.backbone}_{fragment_id}_{CFG.inf_weight}_score.pth', map_location=torch.device('cpu'))\n#check_point = torch.load(CFG.model_dir + f'{CFG.backbone}_{fragment_id}_{CFG.inf_weight}_combo.pth', map_location=torch.device('cpu'))\ncheck_point = torch.load(CFG.comp_dir_path + 'models/efficientnet-b7_1_best_champ_65.pth', map_location=torch.device('cpu'))\n#check_point = torch.load(CFG.comp_dir_path + f'models/efficientnet-b7_{load_id}_best_combo.pth', map_location=torch.device('cpu'))\n\nmask_pred_orig = check_point['preds']\nvalid_mask_gt = cv2.imread(CFG.comp_dataset_path + f\"train_new/{fragment_id}/inklabels.png\", 0)\nvalid_mask_gt = valid_mask_gt / 255\nbinary_mask = cv2.imread(CFG.comp_dataset_path + f\"train_new/{fragment_id}/mask.png\", 0) // 255\npad0 = (tile_size_val_while_train - valid_mask_gt.shape[0] % (tile_size_val_while_train))\npad1 = (tile_size_val_while_train - valid_mask_gt.shape[1] % (tile_size_val_while_train))\nvalid_mask_gt = np.pad(valid_mask_gt, [(0, pad0 ), (0, pad1)], constant_values=0)\nbest_dice, best_th  = calc_fbeta(valid_mask_gt, mask_pred_orig, range(30,81,5))","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-07-20T13:21:10.305863Z","iopub.execute_input":"2023-07-20T13:21:10.306268Z","iopub.status.idle":"2023-07-20T13:26:04.716440Z","shell.execute_reply.started":"2023-07-20T13:21:10.306235Z","shell.execute_reply":"2023-07-20T13:26:04.715289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"weights = check_point['model']\neval_model = build_eval_model(CFG)\neval_model.load_state_dict(weights)\neval_model.eval()\neval_model.to(device) \nprint()","metadata":{"execution":{"iopub.status.busy":"2023-07-20T13:26:04.721041Z","iopub.execute_input":"2023-07-20T13:26:04.721341Z","iopub.status.idle":"2023-07-20T13:26:07.259393Z","shell.execute_reply.started":"2023-07-20T13:26:04.721314Z","shell.execute_reply":"2023-07-20T13:26:07.258333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(mask_pred_orig)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T13:26:07.260830Z","iopub.execute_input":"2023-07-20T13:26:07.261222Z","iopub.status.idle":"2023-07-20T13:26:16.007739Z","shell.execute_reply.started":"2023-07-20T13:26:07.261187Z","shell.execute_reply":"2023-07-20T13:26:16.006776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image, mask, im_shape = read_eval_image_mask(fragment_id, broaden=broaden)\nvalid_images = []\nvalid_masks = []\nvalid_xyxys = []\nx1_list = list(range(0, image.shape[1]-tile_size_val+1, stride))\ny1_list = list(range(0, image.shape[0]-tile_size_val+1, stride))\nfor y1 in y1_list:\n    for x1 in x1_list:\n        y2 = y1 + tile_size_val\n        x2 = x1 + tile_size_val\n        valid_images.append(image[y1:y2, x1:x2])\n        valid_masks.append(mask[y1:y2, x1:x2, None])\n        valid_xyxys.append([x1, y1, x2, y2])\n        \nvalid_xyxys = np.stack(valid_xyxys)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T13:26:16.010038Z","iopub.execute_input":"2023-07-20T13:26:16.011151Z","iopub.status.idle":"2023-07-20T13:26:45.863443Z","shell.execute_reply.started":"2023-07-20T13:26:16.011113Z","shell.execute_reply":"2023-07-20T13:26:45.862352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_dataset = EvalDataset(valid_images, CFG, transform=get_eval_transforms())\nvalid_loader = DataLoader(valid_dataset,batch_size=valid_batch_size,shuffle=False,num_workers=CFG.num_workers, pin_memory=True, drop_last=False, worker_init_fn=worker_init_fn)\nvalid_mask_gt = cv2.imread(CFG.comp_dataset_path + ev_path + \"/inklabels.png\", 0)\nvalid_mask_gt = valid_mask_gt // 255","metadata":{"execution":{"iopub.status.busy":"2023-07-20T13:26:45.865041Z","iopub.execute_input":"2023-07-20T13:26:45.865548Z","iopub.status.idle":"2023-07-20T13:26:46.107880Z","shell.execute_reply.started":"2023-07-20T13:26:45.865511Z","shell.execute_reply":"2023-07-20T13:26:46.106749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def TTA(x, model, rotations, flip):\n    shape = x.shape\n    y_preds = []\n    for i in rotations:\n        x_rot = torch.rot90(x, k=i, dims=(-2,-1))\n        x_rot = model(x_rot)\n        x_rot = torch.sigmoid(x_rot)\n        x_rot = x_rot.reshape(shape[0],*shape[2:])\n        x_rot = torch.rot90(x_rot,k=-i%4,dims=(-2,-1))\n        y_preds.append(x_rot)\n\n    if flip:\n        x_flip = x.flip(-1)\n        x_flip = model(x_flip)\n        x_flip = torch.sigmoid(x_flip)\n        x_flip = x_flip.reshape(shape[0],*shape[2:])\n        x_flip = x_flip.flip(-1)\n        y_preds.append(x_flip)\n\n        for i in rotations:\n            x_rot_flip = torch.rot90(x.flip(-1), k=i, dims=(-2,-1))\n            x_rot_flip = model(x_rot_flip)\n            x_rot_flip = torch.sigmoid(x_rot_flip)\n            x_rot_flip = x_rot_flip.reshape(shape[0],*shape[2:])\n            x_rot_flip = torch.rot90(x_rot_flip,k=-i%4,dims=(-2,-1)).flip(-1)\n            y_preds.append(x_rot_flip)\n\n    y_preds = torch.stack(y_preds,dim=0)\n    return y_preds.mean(0)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T13:26:46.109625Z","iopub.execute_input":"2023-07-20T13:26:46.110541Z","iopub.status.idle":"2023-07-20T13:26:46.123235Z","shell.execute_reply.started":"2023-07-20T13:26:46.110501Z","shell.execute_reply":"2023-07-20T13:26:46.122140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Implement second mean map to correct adaptively?\nrotations = [0,1,2,3]\napply_flip = False\n\ndef eval_fn(valid_loader, model, criterion, device, valid_xyxys, valid_mask_gt, binary_mask=1):\n    mask_pred = np.zeros(im_shape)\n    mask_count = np.zeros(im_shape)\n    model.eval()\n    for j in range(broaden*2+1):\n        for step, (images) in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n            images = images.to(device)\n            batch_size = images.size(0)\n            with torch.no_grad():\n                y_preds = TTA(images[:,j:j+8,:,:],model, rotations, apply_flip).cpu().numpy()\n            #y_preds /= np.mean(y_preds) * 1 + 0.55\n            start_idx = step*valid_batch_size\n            end_idx = start_idx + batch_size\n            for i, (x1, y1, x2, y2) in enumerate(valid_xyxys[start_idx:end_idx]):\n                mask_pred[y1+border:y2-border, x1+border:x2-border] += y_preds[i].reshape(mask_pred[y1:y2, x1:x2].shape)[border:-border, border:-border]\n                mask_count[y1+border:y2-border, x1+border:x2-border] += np.ones((tile_size_val // 2, tile_size_val // 2))\n        #plt.figure()\n        #plt.imshow(mask_count)\n        mask_count = mask_count[border:valid_mask_gt.shape[0]+border, border:valid_mask_gt.shape[1]+border]\n    return mask_pred[border:valid_mask_gt.shape[0]+border, border:valid_mask_gt.shape[1]+border] / (broaden*2 + 1) * binary_mask / mask_count","metadata":{"execution":{"iopub.status.busy":"2023-07-20T13:26:46.125625Z","iopub.execute_input":"2023-07-20T13:26:46.126463Z","iopub.status.idle":"2023-07-20T13:26:46.140513Z","shell.execute_reply.started":"2023-07-20T13:26:46.126391Z","shell.execute_reply":"2023-07-20T13:26:46.139498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask_pred = eval_fn(valid_loader, eval_model, criterion, device, valid_xyxys, valid_mask_gt, binary_mask)\n#best_dice, best_th  = calc_fbeta(valid_mask_gt, mask_pred, range(50,51,5))","metadata":{"execution":{"iopub.status.busy":"2023-07-20T13:26:46.142083Z","iopub.execute_input":"2023-07-20T13:26:46.142485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = most_certain_fraction(mask_pred, np.sum(binary_mask), fraction_goal=0.115)\nprint(fbeta_numpy(valid_mask_gt, preds))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 3, figsize=(15, 8))\naxes[0].imshow(valid_mask_gt)\naxes[1].imshow(mask_pred)\naxes[2].imshow((mask_pred>=best_th).astype(int))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}