{"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":"\ntry:\n    import pylibjpeg\nexcept:\n    !pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.whl,pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}\n","metadata":{"papermill":{"duration":112.187954,"end_time":"2022-12-05T00:23:08.74552","exception":false,"start_time":"2022-12-05T00:21:16.557566","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-15T08:11:07.9245Z","iopub.execute_input":"2022-12-15T08:11:07.924864Z","iopub.status.idle":"2022-12-15T08:11:07.930582Z","shell.execute_reply.started":"2022-12-15T08:11:07.924835Z","shell.execute_reply":"2022-12-15T08:11:07.929405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nimport os\nimport pydicom as dicom\n\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport torch\nimport sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\nfrom timm import create_model\nfrom tqdm.notebook import tqdm\n\npd.set_option('display.max_rows', 1000)\npd.set_option('display.max_columns', 1000)\nplt.rcParams['figure.figsize'] = (20, 5)\n\n\n\nDEBUG = True\n\nRSNA_2022_PATH = '/kaggle/input/rsna-breast-cancer-detection'\nPNG_TEST_IMAGES_PATH = f'test'\nMODELS_PATH = '/kaggle/input/breast-cancer-seresnext50-32x4d'\nTRAIN_IMAGES_PATH = f'{RSNA_2022_PATH}/train_images'\nTEST_IMAGES_PATH = f'{RSNA_2022_PATH}/test_images'\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nif DEVICE == 'cuda':\n    BATCH_SIZE = 16\nelse:\n    BATCH_SIZE = 2\nAUX_TARGET_NCLASSES = [2, 2, 6, 2, 2, 2, 4, 5, 2, 10, 10]\n\nclass Config:\n    ONE_CYCLE = True\n    ONE_CYCLE_PCT_START = 0.1\n    ADAMW = False\n    ADAMW_DECAY = 0.024\n    ONE_CYCLE_MAX_LR = 0.0008\n    EPOCHS = 3\n    MODEL_TYPE = 'seresnext50_32x4d'\n    DROPOUT = 0.2\n    AUG = False\n    AUX_LOSS_WEIGHT = 94\n    POSITIVE_TARGET_WEIGHT=20\n    BATCH_SIZE = 32\n    AUTO_AUG_M = 10\n    AUTO_AUG_N = 2\n    TTA = False","metadata":{"papermill":{"duration":3.570942,"end_time":"2022-12-05T00:23:12.324909","exception":false,"start_time":"2022-12-05T00:23:08.753967","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-15T08:11:38.076847Z","iopub.execute_input":"2022-12-15T08:11:38.077206Z","iopub.status.idle":"2022-12-15T08:11:38.111031Z","shell.execute_reply.started":"2022-12-15T08:11:38.077177Z","shell.execute_reply":"2022-12-15T08:11:38.109818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(f'{RSNA_2022_PATH}/train.csv')\ndf_train.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-15T08:11:38.448593Z","iopub.execute_input":"2022-12-15T08:11:38.449313Z","iopub.status.idle":"2022-12-15T08:11:38.529941Z","shell.execute_reply.started":"2022-12-15T08:11:38.449276Z","shell.execute_reply":"2022-12-15T08:11:38.528882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. Loading test dataframe\n\n ","metadata":{"papermill":{"duration":0.005183,"end_time":"2022-12-05T00:23:12.335995","exception":false,"start_time":"2022-12-05T00:23:12.330812","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def load_df_test():\n    df_test = pd.read_csv(f'{RSNA_2022_PATH}/test.csv')\n    return df_test\n\ndf_test = load_df_test()\ndf_test","metadata":{"papermill":{"duration":0.047641,"end_time":"2022-12-05T00:23:12.38885","exception":false,"start_time":"2022-12-05T00:23:12.341209","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-15T08:11:41.719877Z","iopub.execute_input":"2022-12-15T08:11:41.720908Z","iopub.status.idle":"2022-12-15T08:11:41.738085Z","shell.execute_reply.started":"2022-12-15T08:11:41.72087Z","shell.execute_reply":"2022-12-15T08:11:41.736992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. Dataset class\n","metadata":{"papermill":{"duration":0.005225,"end_time":"2022-12-05T00:23:12.399738","exception":false,"start_time":"2022-12-05T00:23:12.394513","status":"completed"},"tags":[]}},{"cell_type":"code","source":"test_images = glob.glob(f\"{TEST_IMAGES_PATH}/*/*.dcm\")\nlen(test_images)  ","metadata":{"papermill":{"duration":0.020585,"end_time":"2022-12-05T00:23:12.425702","exception":false,"start_time":"2022-12-05T00:23:12.405117","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-15T08:11:44.269922Z","iopub.execute_input":"2022-12-15T08:11:44.270282Z","iopub.status.idle":"2022-12-15T08:11:44.279184Z","shell.execute_reply.started":"2022-12-15T08:11:44.270251Z","shell.execute_reply":"2022-12-15T08:11:44.277927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom(path):\n    \"\"\"\n    This supports loading both regular and compressed JPEG images. \n    See the first sell with `pip install` commands for the necessary dependencies\n    \"\"\"\n    img=dicom.dcmread(path)\n    img.PhotometricInterpretation = 'YBR_FULL'\n    data = img.pixel_array    \n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n        data=(data * 255).astype(np.uint8)\n    return cv2.cvtColor(data, cv2.COLOR_GRAY2RGB), img\n\n\nim, meta = load_dicom(f'{TRAIN_IMAGES_PATH}/10006/1459541791.dcm')\nplt.figure()\nplt.imshow(im)\nplt.title('regular image')\n\nim, meta = load_dicom(f'{TRAIN_IMAGES_PATH}/10006/1864590858.dcm')\nplt.figure()\nplt.imshow(im)\nplt.title('regular image')","metadata":{"execution":{"iopub.status.busy":"2022-12-15T08:11:45.429536Z","iopub.execute_input":"2022-12-15T08:11:45.430204Z","iopub.status.idle":"2022-12-15T08:11:56.609551Z","shell.execute_reply.started":"2022-12-15T08:11:45.430161Z","shell.execute_reply":"2022-12-15T08:11:56.60856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor\nimport re\nimport pydicom\n\ndef fit_image(fname, size=1024):\n    # 1. Read, resize\n    patient = fname.split('/')[-2]\n    image = fname.split('/')[-1][:-4]\n    dicom = pydicom.dcmread(fname)\n    img = dicom.pixel_array\n    img = (img - img.min()) / (img.max() - img.min())\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n    img = cv2.resize(img, (size, size))\n    \n    # 2. Crop\n    X = img\n    # Some images have narrow exterior \"frames\" that complicate selection of the main data. Cutting off the frame\n    X = X[5:-5, 5:-5]\n    \n    # regions of non-empty pixels\n    output= cv2.connectedComponentsWithStats((X > 0.05).astype(np.uint8)[:, :], 8, cv2.CV_32S)\n\n    # stats.shape == (N, 5), where N is the number of regions, 5 dimensions correspond to:\n    # left, top, width, height, area_size\n    stats = output[2]\n    \n    # finding max area which always corresponds to the breast data. \n    idx = stats[1:, 4].argmax() + 1\n    x1, y1, w, h = stats[idx][:4]\n    x2 = x1 + w\n    y2 = y1 + h\n    \n    # cutting out the breast data\n    X_fit = X[y1: y2, x1: x2]\n    patient_id, im_id = os.path.basename(os.path.dirname(fname)), os.path.basename(fname)[:-4]\n    os.makedirs(f'{PNG_TEST_IMAGES_PATH}/test_images/{patient_id}', exist_ok=True)\n    cv2.imwrite(f'{PNG_TEST_IMAGES_PATH}/test_images/{patient_id}/{im_id}.png', (X_fit[:, :] * 255).astype(np.uint8))\n\ndef fit_all_images(all_images):\n    with ThreadPoolExecutor(2) as p:\n        for i in tqdm(p.map(fit_image, all_images), total=len(all_images)):\n            pass\n\nall_images = glob.glob('/kaggle/input/rsna-breast-cancer-detection/test_images/*/*') \n# all_images = glob.glob('/kaggle/input/rsna-breast-cancer-detection/train_images/10006/*')\nfit_all_images(all_images)","metadata":{"papermill":{"duration":6.642302,"end_time":"2022-12-05T00:23:19.073521","exception":false,"start_time":"2022-12-05T00:23:12.431219","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-15T08:11:57.820197Z","iopub.execute_input":"2022-12-15T08:11:57.820605Z","iopub.status.idle":"2022-12-15T08:12:00.167584Z","shell.execute_reply.started":"2022-12-15T08:11:57.820569Z","shell.execute_reply":"2022-12-15T08:12:00.166423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!find test | head","metadata":{"execution":{"iopub.status.busy":"2022-12-15T08:12:00.171438Z","iopub.execute_input":"2022-12-15T08:12:00.172197Z","iopub.status.idle":"2022-12-15T08:12:01.141848Z","shell.execute_reply.started":"2022-12-15T08:12:00.172158Z","shell.execute_reply":"2022-12-15T08:12:01.140659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision\nfrom PIL import Image\n\ndef get_transforms(aug=False):\n\n    def transforms(img):\n        img = img.convert('RGB')#.resize((512, 512))\n        if aug:\n            tfm = [\n                torchvision.transforms.RandomHorizontalFlip(0.5),\n                torchvision.transforms.RandomRotation(degrees=(-5, 5)), \n                torchvision.transforms.RandomResizedCrop((1024, 512), scale=(0.8, 1), ratio=(0.45, 0.55)) \n                  ]\n        else:\n            tfm = [\n                torchvision.transforms.RandomHorizontalFlip(0.5),\n                torchvision.transforms.Resize((1024, 512))\n            ]\n        img = torchvision.transforms.Compose(tfm + [            \n            torchvision.transforms.ToTensor(),\n            torchvision.transforms.Normalize(mean=0.2179, std=0.0529),\n            \n        ])(img)\n        return img\n    return lambda img: transforms(img)\n\nif DEBUG:\n    tfm = get_transforms(aug=False)\n    img = Image.open(f\"{PNG_TEST_IMAGES_PATH}/test_images/10008/68070693.png\")\n    plt.imshow(np.array(img), cmap='gray')\n    plt.show()\n\n    plt.figure(figsize=(20, 20))\n    for i in range(8):\n        v = tfm(img).permute(1, 2, 0)\n        v -= v.min()\n        v /= v.max()\n        # plt.imshow(v)\n        # break\n        plt.subplot(2, 4, i + 1).imshow(v)\n    plt.tight_layout()","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-15T08:12:04.027481Z","iopub.execute_input":"2022-12-15T08:12:04.028246Z","iopub.status.idle":"2022-12-15T08:12:07.246447Z","shell.execute_reply.started":"2022-12-15T08:12:04.028205Z","shell.execute_reply":"2022-12-15T08:12:07.245023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\n\n\nclass BreastCancerDataSet(torch.utils.data.Dataset):\n    def __init__(self, df, path, transforms=None):\n        super().__init__()\n        self.df = df\n        self.path = path\n        self.transforms = transforms\n\n    def __getitem__(self, i):\n        path = f'{self.path}/test_images/{self.df.iloc[i].patient_id}/{self.df.iloc[i].image_id}.png'\n        try:\n            img = Image.open(path).convert('RGB')\n        except Exception as ex:\n            print(path, ex)\n            return None\n\n        if self.transforms is not None:\n            img = self.transforms(img)\n\n\n        return img\ndef __len__(self):\n        return len(self.df)\n\nds_test = BreastCancerDataSet(df_test, PNG_TEST_IMAGES_PATH, get_transforms(False))\nif DEBUG:\n    X, y_cancer, y_aux = ds_test[2]\n    print(X.shape, y_cancer.shape, y_aux.shape)","metadata":{"papermill":{"duration":0.091858,"end_time":"2022-12-05T00:23:19.183339","exception":false,"start_time":"2022-12-05T00:23:19.091481","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-15T08:12:07.248094Z","iopub.execute_input":"2022-12-15T08:12:07.249044Z","iopub.status.idle":"2022-12-15T08:12:07.291741Z","shell.execute_reply.started":"2022-12-15T08:12:07.249006Z","shell.execute_reply":"2022-12-15T08:12:07.290674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4. Model\n","metadata":{"papermill":{"duration":0.006036,"end_time":"2022-12-05T00:23:19.195636","exception":false,"start_time":"2022-12-05T00:23:19.1896","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class BreastCancerModel(torch.nn.Module):\n    def __init__(self, aux_classes, model_type='seresnext50_32x4d', dropout=0.):\n        super().__init__()\n        self.model = create_model(model_type, pretrained=False, num_classes=0, drop_rate=dropout)\n\n        self.backbone_dim = self.model(torch.randn(1, 3, 512, 512)).shape[-1]\n\n        self.nn_cancer = torch.nn.Sequential(\n            torch.nn.Linear(self.backbone_dim, 1),\n        )\n        self.nn_aux = torch.nn.ModuleList([\n            torch.nn.Linear(self.backbone_dim, n) for n in aux_classes\n        ])\n    def forward(self, x):\n        # returns logits\n        x = self.model(x)\n\n        cancer = self.nn_cancer(x).squeeze()\n        aux = []\n        for nn in self.nn_aux:\n            aux.append(nn(x).squeeze())\n        return cancer, aux\n    \n    def predict(self, x):\n        cancer, aux = self.forward(x)\n        sigaux = []\n        for a in aux:\n            sigaux.append(torch.softmax(a, dim=-1))\n        return torch.sigmoid(cancer), sigaux\n\nif DEBUG:\n    with torch.no_grad():\n        model = BreastCancerModel(AUX_TARGET_NCLASSES, model_type='seresnext50_32x4d')\n        pred, aux = model.predict(torch.randn(2, 3, 512, 512))\n        print('seresnext', pred.shape, len(aux))\n\n        model = BreastCancerModel(AUX_TARGET_NCLASSES, model_type='efficientnet_b4')\n        pred, aux = model.predict(torch.randn(2, 3, 512, 512))\n        print('efficientnet_b4', pred.shape, len(aux))\n        \n    del model\n","metadata":{"papermill":{"duration":2.459131,"end_time":"2022-12-05T00:23:21.660809","exception":false,"start_time":"2022-12-05T00:23:19.201678","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-15T08:12:14.396648Z","iopub.execute_input":"2022-12-15T08:12:14.397002Z","iopub.status.idle":"2022-12-15T08:12:21.210217Z","shell.execute_reply.started":"2022-12-15T08:12:14.396972Z","shell.execute_reply":"2022-12-15T08:12:21.209023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_model(name, dir='.', model=None):\n    data = torch.load(os.path.join(dir, f'{name}'), map_location=DEVICE)\n    if model is None:\n        model = BreastCancerModel(AUX_TARGET_NCLASSES, data['model_type'])\n    model.load_state_dict(data['model'])\n    # print(data['threshold'], data['model_type'])\n    return model, data['threshold'], data['model_type']\n","metadata":{"papermill":{"duration":0.014644,"end_time":"2022-12-05T00:23:21.681888","exception":false,"start_time":"2022-12-05T00:23:21.667244","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-15T08:12:24.107742Z","iopub.execute_input":"2022-12-15T08:12:24.108295Z","iopub.status.idle":"2022-12-15T08:12:24.120709Z","shell.execute_reply.started":"2022-12-15T08:12:24.108253Z","shell.execute_reply":"2022-12-15T08:12:24.119784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\nfor fname in tqdm(sorted(os.listdir(MODELS_PATH))):\n    model, thres, model_type = load_model(fname, MODELS_PATH)\n    model = model.to(DEVICE)\n    models.append((model, thres))\n    print(f'fname:{fname}, model_type:{model_type}, thres:{thres}')\n","metadata":{"papermill":{"duration":12.829022,"end_time":"2022-12-05T00:23:34.516853","exception":false,"start_time":"2022-12-05T00:23:21.687831","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-15T08:12:26.232055Z","iopub.execute_input":"2022-12-15T08:12:26.232437Z","iopub.status.idle":"2022-12-15T08:12:35.751109Z","shell.execute_reply.started":"2022-12-15T08:12:26.232403Z","shell.execute_reply":"2022-12-15T08:12:35.750009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"THRES = 0.80 \nTHRES = 0.51","metadata":{"execution":{"iopub.status.busy":"2022-12-15T08:12:35.753205Z","iopub.execute_input":"2022-12-15T08:12:35.753879Z","iopub.status.idle":"2022-12-15T08:12:35.758484Z","shell.execute_reply.started":"2022-12-15T08:12:35.753839Z","shell.execute_reply":"2022-12-15T08:12:35.757516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 5. Submission","metadata":{"papermill":{"duration":0.005913,"end_time":"2022-12-05T00:23:34.529375","exception":false,"start_time":"2022-12-05T00:23:34.523462","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def models_predict(models, ds, max_batches=1e9):\n    dl_test = torch.utils.data.DataLoader(ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=os.cpu_count())\n    for m, thres in models:\n        m.eval()\n\n    with torch.no_grad():\n        predictions = []\n        for idx, X in enumerate(tqdm(dl_test, mininterval=30)):\n            pred = torch.zeros(len(X), len(models))\n            for idx, (m, thres) in enumerate(models):\n                preds = m.predict(X.to(DEVICE))[0].squeeze()\n                pred[:, idx] = preds.cpu()\n            predictions.append(pred.mean(dim=-1))\n            \n            if idx >= max_batches:\n                break\n        return torch.concat(predictions).numpy()\n\n# Quick test\nif DEBUG:\n    print(models_predict([(BreastCancerModel(AUX_TARGET_NCLASSES, 'seresnext50_32x4d').to(DEVICE), 0.5),\n                          (BreastCancerModel(AUX_TARGET_NCLASSES, 'seresnext50_32x4d').to(DEVICE), 0.1)], ds_test, max_batches=2))","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"papermill":{"duration":2.339608,"end_time":"2022-12-05T00:23:36.875001","exception":false,"start_time":"2022-12-05T00:23:34.535393","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-15T08:12:38.002635Z","iopub.execute_input":"2022-12-15T08:12:38.003006Z","iopub.status.idle":"2022-12-15T08:12:40.693944Z","shell.execute_reply.started":"2022-12-15T08:12:38.002976Z","shell.execute_reply":"2022-12-15T08:12:40.692309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models_pred = models_predict(models, ds_test)","metadata":{"papermill":{"duration":0.544972,"end_time":"2022-12-05T00:23:37.427741","exception":false,"start_time":"2022-12-05T00:23:36.882769","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test['cancer'] = models_pred\n\ndf_sub = df_test.groupby('prediction_id')[['cancer']].mean()\ndf_sub","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub['cancer'] = (df_sub.cancer > THRES).astype(float)\ndf_sub","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub.to_csv('submission.csv', index=True)\n!head submission.csv\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}