{"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":"code","source":"import sys\nsys.path.append('/kaggle/input/pretrainedmodels/pretrainedmodels-0.7.4')\nsys.path.append('/kaggle/input/efficientnet-pytorch/EfficientNet-PyTorch-master')\nsys.path.append('/kaggle/input/timm-pytorch-image-models/pytorch-image-models-master')\nsys.path.append('/kaggle/input/segmentation-models-pytorch/segmentation_models.pytorch-master')\nsys.path.append('/kaggle/input/mmcv-pytorch')\nsys.path.append('/kaggle/input/addict')\nsys.path.append('/kaggle/input/ink-model')\nfrom model.ink_model import get_model\nfrom mmcv.cnn import ConvModule\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.utils.data as data\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.tensorboard import SummaryWriter\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport segmentation_models_pytorch as smp\nfrom segmentation_models_pytorch.decoders.unetplusplus.decoder import UnetPlusPlusDecoder\nfrom segmentation_models_pytorch.base import SegmentationHead\nfrom torch.cuda.amp import GradScaler\nfrom torchvision.utils import make_grid\nimport warnings\nimport os\nimport random\nimport pandas as pd\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport gc\n# 忽略所有警告\nwarnings.filterwarnings('ignore')\n\nclass CFG:\n    device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\n\n#     checpoint = '/kaggle/input/pretrain/Resnet3D-17.pkl'\n    # ============== comp exp name =============\n    comp_name = 'vesuvius'\n\n    # comp_dir_path = './'\n    comp_dir_path = '/kaggle/input/'\n    comp_folder_name = 'vesuvius-challenge-ink-detection'\n    # comp_dataset_path = f'{comp_dir_path}datasets/{comp_folder_name}/'\n    comp_dataset_path = f'{comp_dir_path}{comp_folder_name}/'\n#         comp_dir_path = './'\n#     comp_dir_path = ''\n#     comp_folder_name = 'data'\n#     # comp_dataset_path = f'{comp_dir_path}datasets/{comp_folder_name}/'\n#     comp_dataset_path = f'{comp_dir_path}{comp_folder_name}/'\n\n#     img_path = 'working/'\n    \n    encoder_name = 'resnet3d' # resnet3d、r2plus1d_18、r3d_18、mc3_18、convnext3d、unet3d_down\n    decoder_name = 'cnn' # cnn、uper_head、unetplusplus、unet3d_up\n    mix_up = False\n\n    # ============== pred target =============\n    target_size = 1\n\n    # ============== model cfg =============\n    model_name = encoder_name + '-' + decoder_name\n\n    # ============== model cfg =============\n\n    in_idx = [i for i in range(21, 43)]\n    \n    in_chans = len(in_idx)# 65\n    # ============== training cfg =============\n    size = [224]\n    tile_size = [224]\n    stride = [i // 6 for i in size]\n    \n    inplanes = [64, 128, 256, 512]\n\n    valid_batch_size = [16]\n    \n    threshhold = 0.45\n\n    num_workers = 4\n\n    seed = 42\n\n    all_best_dice = 0\n    all_best_loss = np.float('inf')\n\n    shape_list = []\n    test_shape_list = []\n\n    val_mask = None\n    val_label = None\n\n    valid_aug_list = [\n        A.Normalize(\n            mean= [0] * in_chans,\n            std= [1] * in_chans\n        ),\n        ToTensorV2(transpose_mask=True),\n    ]\nseed = CFG.seed\nos.environ['PYTHONHASHSEED'] = str(seed)\nnp.random.seed(seed)\nrandom.seed(seed)\ntorch.manual_seed(seed)\ntorch.cuda.manual_seed(seed)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:41:56.086613Z","iopub.execute_input":"2023-06-11T04:41:56.087268Z","iopub.status.idle":"2023-06-11T04:42:10.373610Z","shell.execute_reply.started":"2023-06-11T04:41:56.087228Z","shell.execute_reply":"2023-06-11T04:42:10.372333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Ink3DModel(nn.Module):\n    def __init__(self, encoder_name, decoder_name, mix_up=CFG.mix_up, **kwargs):\n        super().__init__()\n        self.encoder_name = encoder_name\n        self.decoder_name = decoder_name\n        self.mix_up = mix_up\n        self.encoder, self.decoder = get_model(encoder_name=encoder_name, decoder_name=decoder_name, **kwargs)\n\n    def forward(self, x):\n        feat_maps = self.encoder(x)\n        if self.mix_up:\n            pred_mask = []\n            for decoder in self.decoder:\n                mask = decoder(feat_maps)\n                pred_mask.append(mask)\n            pred_mask = torch.stack(pred_mask, dim=0)\n            pred_mask = torch.mean(pred_mask, dim=0)\n        else:\n            pred_mask = self.decoder(feat_maps)\n        return pred_mask","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:10.376295Z","iopub.execute_input":"2023-06-11T04:42:10.376719Z","iopub.status.idle":"2023-06-11T04:42:10.387265Z","shell.execute_reply.started":"2023-06-11T04:42:10.376670Z","shell.execute_reply":"2023-06-11T04:42:10.385959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_image_location(fragment_id, idx):\n    images = []\n    idxs = CFG.in_idx\n\n    for i in tqdm(idxs):\n        \n        image = cv2.imread(CFG.comp_dataset_path + f\"test/{fragment_id}/surface_volume/{i:02}.tif\", 0)\n\n        pad0 = (CFG.tile_size[idx] - image.shape[0] % CFG.tile_size[idx])\n        pad1 = (CFG.tile_size[idx] - image.shape[1] % CFG.tile_size[idx])\n\n        image = np.pad(image, [(0, pad0), (0, pad1)], constant_values=0)\n\n        images.append(image)\n    images = np.stack(images, axis=2)\n    \n    mask_location = cv2.imread(CFG.comp_dataset_path + f\"test/{fragment_id}/mask.png\", 0)\n    mask_location = np.pad(mask_location, [(0, pad0), (0, pad1)], constant_values=0)\n\n    mask_location = mask_location / 255\n    \n    return images, mask_location","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:10.389211Z","iopub.execute_input":"2023-06-11T04:42:10.390019Z","iopub.status.idle":"2023-06-11T04:42:10.405321Z","shell.execute_reply.started":"2023-06-11T04:42:10.389975Z","shell.execute_reply":"2023-06-11T04:42:10.404060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(data, cfg):\n    if data == 'train':\n        aug = A.Compose(cfg.train_aug_list)\n    elif data == 'valid':\n        aug = A.Compose(cfg.valid_aug_list)\n\n    # print(aug)\n    return aug\n\nclass Ink_Detection_Dataset(data.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\n    def __len__(self):\n        # return len(self.xyxys)\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        # x1, y1, x2, y2 = self.xyxys[idx]\n        image = self.images[idx]\n        data = self.transform(image=image)\n        image = data['image'].unsqueeze(0)\n        return image","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:10.409630Z","iopub.execute_input":"2023-06-11T04:42:10.410186Z","iopub.status.idle":"2023-06-11T04:42:10.421577Z","shell.execute_reply.started":"2023-06-11T04:42:10.410138Z","shell.execute_reply":"2023-06-11T04:42:10.420258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_test_dataset(fragment_id, idx):\n    test_images, mask_location = read_image_location(fragment_id, idx)\n    \n    x1_list = list(range(0, test_images.shape[1]-CFG.tile_size[idx]+1, CFG.stride[idx]))\n    y1_list = list(range(0, test_images.shape[0]-CFG.tile_size[idx]+1, CFG.stride[idx]))\n    \n    test_images_list = []\n    xyxys = []\n    for y1 in y1_list:\n        for x1 in x1_list:\n            y2 = y1 + CFG.tile_size[idx]\n            x2 = x1 + CFG.tile_size[idx]\n            if np.sum(mask_location[y1:y2, x1:x2]) == 0:\n                continue\n            test_images_list.append(test_images[y1:y2, x1:x2])\n            xyxys.append((x1, y1, x2, y2))\n    xyxys = np.stack(xyxys)\n            \n    test_dataset = Ink_Detection_Dataset(test_images_list, CFG, transform=get_transforms(data='valid', cfg=CFG))\n    \n    test_loader = data.DataLoader(test_dataset,\n                          batch_size=CFG.valid_batch_size[idx],\n                          shuffle=False,\n                          num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n    \n    return test_loader, xyxys","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:10.423256Z","iopub.execute_input":"2023-06-11T04:42:10.423903Z","iopub.status.idle":"2023-06-11T04:42:10.438528Z","shell.execute_reply.started":"2023-06-11T04:42:10.423857Z","shell.execute_reply":"2023-06-11T04:42:10.437253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    # pixels = (pixels >= thr).astype(int)\n    \n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:10.441073Z","iopub.execute_input":"2023-06-11T04:42:10.441745Z","iopub.status.idle":"2023-06-11T04:42:10.452957Z","shell.execute_reply.started":"2023-06-11T04:42:10.441698Z","shell.execute_reply":"2023-06-11T04:42:10.451714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    def __init__(self, cfg, weight=None, model=None):\n        super().__init__()\n        self.cfg = cfg\n\n        self.encoder = model\n\n    def forward(self, image):\n        output = self.encoder(image)\n#         output = output.squeeze(-1)\n        return output\n\ndef build_model(cfg, weight=\"imagenet\", model=None):\n    print('model_name', cfg.model_name)\n    print('backbone', cfg.backbone)\n\n    model = CustomModel(cfg, weight, model)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:10.455168Z","iopub.execute_input":"2023-06-11T04:42:10.455637Z","iopub.status.idle":"2023-06-11T04:42:10.467194Z","shell.execute_reply.started":"2023-06-11T04:42:10.455593Z","shell.execute_reply":"2023-06-11T04:42:10.465516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EnsembleModel:\n    def __init__(self, use_tta=False):\n        self.models = []\n        self.use_tta = use_tta\n\n    def __call__(self, x):\n        outputs = [torch.sigmoid(model(x)).to('cpu').numpy()\n                   for model in self.models]\n        avg_preds = np.mean(outputs, axis=0)\n        return avg_preds\n\n    def add_model(self, model):\n        self.models.append(model)\n\ndef build_ensemble_model(size):\n    model = EnsembleModel()\n    fold_dir_list = ['/kaggle/input/checpoint/' + str(size) + '/']\n    model_name_list = ['resnet']\n    for f in range(len(fold_dir_list)):\n        fold = os.listdir(fold_dir_list[f])\n        for i in range(len(fold)):\n            if model_name_list[f] == 'unet++':\n                _model = build_model(CFG, weight=None, model=smp.UnetPlusPlus(in_channels=CFG.in_chans, \n                                                                            classes=CFG.target_size, \n                                                                            encoder_name=CFG.backbone, \n                                                                            encoder_weights=None, \n                                                                            activation=None, \n                                                                            decoder_attention_type='scse'))\n            if model_name_list[f] == 'resnet':\n                print('resnet')\n                if CFG.decoder_name == 'unetplusplus':\n                    _model = Ink3DUnet(encoder_name=CFG.encoder_name, \n                                      decoder_name=CFG.decoder_name,\n                                      classes=CFG.target_size ,\n                                      decoder_attention_type='scse', \n                                      encoder_depth=4, \n                                      decoder_channels=CFG.inplanes[::-1],\n                                      encoder_weights=None)\n                else:\n                    _model = Ink3DModel(encoder_name=CFG.encoder_name, decoder_name=CFG.decoder_name, classes=CFG.target_size)\n                if size!=288:\n                    _model = nn.DataParallel(_model, device_ids=[0])\n                _model.to(CFG.device)\n            model_path = f'/kaggle/input/checpoint/{size}/{fold[i]}'\n            print(model_path)\n            state = torch.load(model_path, map_location=CFG.device)\n            _model.load_state_dict(state)\n            _model.eval()\n            model.add_model(_model)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:10.470007Z","iopub.execute_input":"2023-06-11T04:42:10.471049Z","iopub.status.idle":"2023-06-11T04:42:10.489195Z","shell.execute_reply.started":"2023-06-11T04:42:10.471003Z","shell.execute_reply":"2023-06-11T04:42:10.487788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fragment_ids = sorted(os.listdir(CFG.comp_dataset_path + 'test'))","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:10.492922Z","iopub.execute_input":"2023-06-11T04:42:10.493379Z","iopub.status.idle":"2023-06-11T04:42:10.506122Z","shell.execute_reply.started":"2023-06-11T04:42:10.493288Z","shell.execute_reply":"2023-06-11T04:42:10.504757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\nfor size in CFG.size:  \n    model = build_ensemble_model(size)\n    models.append(model)\nprint(models)","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:10.511662Z","iopub.execute_input":"2023-06-11T04:42:10.512011Z","iopub.status.idle":"2023-06-11T04:42:16.458879Z","shell.execute_reply.started":"2023-06-11T04:42:10.511982Z","shell.execute_reply":"2023-06-11T04:42:16.456852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def TTA(x:torch.Tensor,model:nn.Module):\n    #x.shape=(batch,c,h,w)\n    shape=x.shape\n    x=[x,*[torch.rot90(x,k=i,dims=(-2,-1)) for i in range(1,4)]]\n    x=torch.cat(x,dim=0)\n    x=torch.from_numpy(model(x))\n    x=x.reshape(4,shape[0],1,*shape[-2:])\n    x=[torch.rot90(x[i],k=-i,dims=(-2,-1)) for i in range(4)]\n    x=torch.stack(x,dim=0)\n    return x.mean(0)","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:16.460368Z","iopub.status.idle":"2023-06-11T04:42:16.461826Z","shell.execute_reply.started":"2023-06-11T04:42:16.461497Z","shell.execute_reply":"2023-06-11T04:42:16.461529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"results = []\nfor fragment_id in fragment_ids:\n    print(f'test{fragment_id}')\n    pred_masks = []\n    for idx, size in enumerate(CFG.size):\n        print(f'test:{size}')\n\n        test_loader, xyxys = make_test_dataset(fragment_id, idx)\n\n        binary_mask = cv2.imread(CFG.comp_dataset_path + f\"test/{fragment_id}/mask.png\", 0)\n        binary_mask = (binary_mask / 255).astype(int)\n\n        ori_h = binary_mask.shape[0]\n        ori_w = binary_mask.shape[1]\n        # mask = mask / 255\n\n        pad0 = (CFG.tile_size[idx] - binary_mask.shape[0] % CFG.tile_size[idx])\n        pad1 = (CFG.tile_size[idx] - binary_mask.shape[1] % CFG.tile_size[idx])\n\n        binary_mask = np.pad(binary_mask, [(0, pad0), (0, pad1)], constant_values=0)\n\n        mask_pred = np.zeros(binary_mask.shape)\n        mask_count = (1 - binary_mask).astype(np.float64)\n\n        for step, (images) in tqdm(enumerate(test_loader), total=len(test_loader)):\n            images = images.to(CFG.device)\n            batch_size = images.size(0)\n\n            with torch.no_grad():\n                y_preds = TTA(images,models[idx]).cpu().numpy()\n    #         y_preds = torch.sigmoid(y_preds).to('cpu').numpy()\n\n            start_idx = step * CFG.valid_batch_size[idx]\n            end_idx = start_idx + batch_size\n            for i, (x1, y1, x2, y2) in enumerate(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[idx], CFG.tile_size[idx]))\n\n        plt.imshow(mask_count)\n        plt.show()\n\n        print(f'mask_count_min: {mask_count.min()}')\n        mask_pred /= mask_count\n\n        mask_pred = mask_pred[:ori_h, :ori_w]\n        binary_mask = binary_mask[:ori_h, :ori_w]\n        has_nan = np.isnan(mask_pred).any()\n        print(has_nan)\n        pred_masks.append(mask_pred)\n    \n    mask_pred = np.mean(pred_masks, axis=0)\n    mask_pred = (mask_pred >= CFG.threshhold).astype(int)\n    mask_pred *= binary_mask\n    \n    plt.imshow(mask_pred)\n    plt.show()\n    \n    inklabels_rle = rle(mask_pred)\n    \n    results.append((fragment_id, inklabels_rle))\n    \n\n    del mask_pred, mask_count\n    del test_loader\n    \n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:16.463826Z","iopub.status.idle":"2023-06-11T04:42:16.464909Z","shell.execute_reply.started":"2023-06-11T04:42:16.464564Z","shell.execute_reply":"2023-06-11T04:42:16.464594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.DataFrame(results, columns=['Id', 'Predicted'])","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:16.466506Z","iopub.status.idle":"2023-06-11T04:42:16.467495Z","shell.execute_reply.started":"2023-06-11T04:42:16.467211Z","shell.execute_reply":"2023-06-11T04:42:16.467240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_sub = pd.read_csv(CFG.comp_dataset_path + 'sample_submission.csv')\nsample_sub = pd.merge(sample_sub[['Id']], sub, on='Id', how='left')","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:16.469327Z","iopub.status.idle":"2023-06-11T04:42:16.470679Z","shell.execute_reply.started":"2023-06-11T04:42:16.470371Z","shell.execute_reply":"2023-06-11T04:42:16.470403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_sub.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-06-11T04:42:16.472491Z","iopub.status.idle":"2023-06-11T04:42:16.473089Z","shell.execute_reply.started":"2023-06-11T04:42:16.472766Z","shell.execute_reply":"2023-06-11T04:42:16.472813Z"},"trusted":true},"execution_count":null,"outputs":[]}]}