{"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":"The dataset I am using in this notebook comes from here: https://www.kaggle.com/datasets/theoviel/rsna-breast-cancer-256-pngs\n\nThe inference notebook is [here](https://www.kaggle.com/snnclsr/rsna-pytorch-baseline-inference).\n\n**Please upvote if you find this notebook useful! It's too much appreciated.**","metadata":{}},{"cell_type":"code","source":"!pip install timm","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import math\nimport time\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import StratifiedGroupKFold\n\nimport cv2\nimport timm\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\nimport albumentations\nfrom albumentations.pytorch import ToTensorV2","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:01.219018Z","iopub.execute_input":"2022-11-30T16:22:01.219412Z","iopub.status.idle":"2022-11-30T16:22:02.967409Z","shell.execute_reply.started":"2022-11-30T16:22:01.219332Z","shell.execute_reply":"2022-11-30T16:22:02.966231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls ../input/rsna-breast-cancer-detection/","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:02.969718Z","iopub.execute_input":"2022-11-30T16:22:02.970269Z","iopub.status.idle":"2022-11-30T16:22:03.97815Z","shell.execute_reply.started":"2022-11-30T16:22:02.970233Z","shell.execute_reply":"2022-11-30T16:22:03.976752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"BASE_IMG_DIR = \"../input/rsna-breast-cancer-256-pngs/\"\n\nclass CFG:\n    \n    model_name = \"efficientnet_b1\"\n    n_folds = 5\n    n_classes = 1\n    n_epochs = 5\n    train_batch_size = 64\n    valid_batch_size = 64\n    lr = 1e-4\n    wd = 1e-6\n    gradient_accumulation_steps = 1\n    max_grad_norm = 1000\n    print_every = 100","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:03.98048Z","iopub.execute_input":"2022-11-30T16:22:03.980993Z","iopub.status.idle":"2022-11-30T16:22:03.988151Z","shell.execute_reply.started":"2022-11-30T16:22:03.980945Z","shell.execute_reply":"2022-11-30T16:22:03.986547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV Split","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:03.991014Z","iopub.execute_input":"2022-11-30T16:22:03.991406Z","iopub.status.idle":"2022-11-30T16:22:04.059074Z","shell.execute_reply.started":"2022-11-30T16:22:03.991372Z","shell.execute_reply":"2022-11-30T16:22:04.05811Z"}}},{"cell_type":"code","source":"df_all = pd.read_csv(\"../input/rsna-breast-cancer-detection/train.csv\")\ndf_all[\"fold\"] = -1\n\ngkfold = StratifiedGroupKFold(n_splits=CFG.n_folds)\nfor fold_idx, (train_idx, val_idx) in enumerate(gkfold.split(df_all, y=df_all.cancer, groups=df_all.patient_id)):\n    df_all.loc[val_idx, \"fold\"] = fold_idx","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:04.060849Z","iopub.execute_input":"2022-11-30T16:22:04.061319Z","iopub.status.idle":"2022-11-30T16:22:08.934201Z","shell.execute_reply.started":"2022-11-30T16:22:04.061282Z","shell.execute_reply":"2022-11-30T16:22:08.933167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_all.groupby(\"fold\")[\"cancer\"].value_counts().plot(kind=\"bar\")","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:08.936034Z","iopub.execute_input":"2022-11-30T16:22:08.93642Z","iopub.status.idle":"2022-11-30T16:22:09.181414Z","shell.execute_reply.started":"2022-11-30T16:22:08.936382Z","shell.execute_reply":"2022-11-30T16:22:09.180414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"def read_img_and_cvt_format(img_path, clr_format=cv2.COLOR_BGR2RGB):\n    return cv2.cvtColor(cv2.imread(img_path), clr_format)\n\nclass RSNADataset(Dataset):\n    \n    def __init__(self, df, is_test=False, transforms=None):\n        super(RSNADataset, self).__init__()\n        self.df = df\n        self.is_test = is_test\n        self.transforms = transforms\n            \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = BASE_IMG_DIR + f\"{row.patient_id}_{row.image_id}.png\"\n        img = read_img_and_cvt_format(img_path)\n        if self.transforms:\n            img = self.transforms(image=img)[\"image\"]\n            \n        label = -1\n        if not self.is_test:\n            label = torch.tensor(row.cancer, dtype=torch.float32).float()\n\n        return img, label\n            ","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:09.184259Z","iopub.execute_input":"2022-11-30T16:22:09.185792Z","iopub.status.idle":"2022-11-30T16:22:09.19413Z","shell.execute_reply.started":"2022-11-30T16:22:09.185755Z","shell.execute_reply":"2022-11-30T16:22:09.193124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Augmentations","metadata":{}},{"cell_type":"code","source":"def _get_train_transforms_without_aug():\n    return albumentations.Compose([\n        albumentations.Normalize(\n            mean=[0.485, 0.456, 0.406], \n            std=[0.229, 0.224, 0.225]\n        ),\n        ToTensorV2()\n    ])","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:09.196676Z","iopub.execute_input":"2022-11-30T16:22:09.197414Z","iopub.status.idle":"2022-11-30T16:22:09.204623Z","shell.execute_reply.started":"2022-11-30T16:22:09.197376Z","shell.execute_reply":"2022-11-30T16:22:09.203718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = _get_train_transforms_without_aug()\ndataset = RSNADataset(df_all, is_test=False, transforms=train_transforms)\ndata_loader = DataLoader(dataset, batch_size=2)\n# plt.imshow(dataset[0][0], cmap=\"bone\")","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:09.206014Z","iopub.execute_input":"2022-11-30T16:22:09.20647Z","iopub.status.idle":"2022-11-30T16:22:09.214849Z","shell.execute_reply.started":"2022-11-30T16:22:09.206436Z","shell.execute_reply":"2022-11-30T16:22:09.214107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class RSNAModel(nn.Module):\n    \n    def __init__(self, model_name, pretrained=True):\n        super(RSNAModel, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        if \"efficientnet\" in CFG.model_name:\n            in_features = self.model.classifier.in_features\n            self.model.classifier = nn.Linear(in_features, CFG.n_classes)\n        elif \"resnet\" in CFG.model_name:\n            in_features = self.model.fc.in_features\n            self.model.fc = nn.Linear(in_features, CFG.n_classes)\n\n    def forward(self, img):\n        return self.model(img)","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:09.215888Z","iopub.execute_input":"2022-11-30T16:22:09.216827Z","iopub.status.idle":"2022-11-30T16:22:09.225034Z","shell.execute_reply.started":"2022-11-30T16:22:09.216787Z","shell.execute_reply":"2022-11-30T16:22:09.224415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utilities","metadata":{}},{"cell_type":"code","source":"LOGS_PATH = Path(\"logs\")\nLOGS_PATH.mkdir(exist_ok=True)\n\nclass AverageMeter:\n    \n    def __init__(self):\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        \ndef as_minutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return f\"{m}m {s}s\"\n\n\ndef time_since(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / percent\n    rs = es - s\n    return f\"{as_minutes(s)} (remain {as_minutes(rs)})\"\n\n\ndef init_logger(log_file=LOGS_PATH / 'train.log'):\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\n\n\nLOGGER = init_logger()","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:09.226143Z","iopub.execute_input":"2022-11-30T16:22:09.227106Z","iopub.status.idle":"2022-11-30T16:22:09.238907Z","shell.execute_reply.started":"2022-11-30T16:22:09.227069Z","shell.execute_reply":"2022-11-30T16:22:09.238217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train and Evaluation Steps","metadata":{}},{"cell_type":"code","source":"def train_step(model, data_loader, criterion, optimizer, epoch, scheduler, device):\n    \"\"\"\n    There is no scheduler update currently.\n    \"\"\"\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    # scores = AverageMeter()\n    \n    model.train()\n    start = end = time.time()\n    # global_step = 0\n    total_len = len(data_loader)\n    \n    for step, (images, labels) in enumerate(data_loader):\n        \n        data_time.update(time.time() - end)\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        preds = model(images).squeeze()\n        loss = criterion(preds, labels)\n        losses.update(loss.item(), batch_size)\n        \n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        \n        loss.backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), \n                                                   CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            optimizer.step()\n            optimizer.zero_grad()\n            # global_step += 1\n        \n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_every == 0 or step == (total_len - 1):\n            print(f\"Epoch: [{epoch+1}][{step}/{total_len}] \"\n                  f\"Data: {data_time.val:.3f} ({data_time.avg:.3f}) \"\n                  f\"Batch: {batch_time.val:.3f} ({batch_time.avg:.3f}) \"\n                  f\"Elapsed: {time_since(start, float(step + 1) / (total_len))} \"\n                  f\"Loss: {losses.val:.5f}({losses.avg:.5f}) \"\n                  f\"Grad: {grad_norm:.4f}\" # LR: {lr:.6f}\n                 )\n    \n    return losses.avg\n            \n\ndef valid_step(model, data_loader, criterion, device):\n    \n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    \n    model.eval()\n    start = end = time.time()\n    total_len = len(data_loader)\n    predictions = []\n    \n    for step, (images, labels) in enumerate(data_loader):\n        data_time.update(time.time() - end)\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        \n        with torch.no_grad():\n            preds = model(images).squeeze()\n        \n        loss = criterion(preds, labels)\n        losses.update(loss.item(), batch_size)\n        predictions.append(preds.sigmoid().cpu().numpy())\n        \n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n            \n        batch_time.update(time.time() - end)\n        end = time.time()\n        \n        if step % CFG.print_every == 0 or step == (total_len - 1):\n            print(f\"Eval: [{step}/{total_len}] \"\n                  f\"Data: {data_time.val:.3f} ({data_time.avg:.3f}) \"\n                  f\"Batch: {batch_time.val:.3f} ({batch_time.avg:.3f}) \"\n                  f\"Elapsed: {time_since(start, float(step + 1) / total_len)} \"\n                  f\"Loss: {losses.val:.5f} ({losses.avg:.5f})\"\n                 )\n    \n    predictions = np.concatenate(predictions)\n    return losses.avg, predictions","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:09.240371Z","iopub.execute_input":"2022-11-30T16:22:09.2412Z","iopub.status.idle":"2022-11-30T16:22:09.25997Z","shell.execute_reply.started":"2022-11-30T16:22:09.241158Z","shell.execute_reply":"2022-11-30T16:22:09.25907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pfbeta(labels, predictions, beta):\n    \"\"\"\n    from here: https://www.kaggle.com/code/sohier/probabilistic-f-score\n    \"\"\"\n    y_true_count = 0\n    ctp = 0\n    cfp = 0\n\n    for idx in range(len(labels)):\n        prediction = min(max(predictions[idx], 0), 1)\n        if (labels[idx]):\n            y_true_count += 1\n            ctp += prediction\n            cfp += 1 - prediction\n        else:\n            cfp += prediction\n\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / y_true_count\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall)\n        return result\n    else:\n        return 0\n","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:09.2632Z","iopub.execute_input":"2022-11-30T16:22:09.263456Z","iopub.status.idle":"2022-11-30T16:22:09.273337Z","shell.execute_reply.started":"2022-11-30T16:22:09.263433Z","shell.execute_reply":"2022-11-30T16:22:09.272339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODELS_DIR = Path(\"models\")\nMODELS_DIR.mkdir(exist_ok=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:09.276707Z","iopub.execute_input":"2022-11-30T16:22:09.277163Z","iopub.status.idle":"2022-11-30T16:22:09.285913Z","shell.execute_reply.started":"2022-11-30T16:22:09.277137Z","shell.execute_reply":"2022-11-30T16:22:09.285006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Full Training","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\noof_df = pd.DataFrame()\n\nfor fold in range(CFG.n_folds):\n    LOGGER.info(f\"===== STARTING FOLD {fold}: ======\")\n    df_train = df_all[df_all.fold != fold].reset_index(drop=True)\n    df_valid = df_all[df_all.fold == fold].reset_index(drop=True)\n    \n    train_dataset = RSNADataset(df_train, transforms=train_transforms)\n    valid_dataset = RSNADataset(df_valid, transforms=train_transforms)\n    train_data_loader = DataLoader(train_dataset, batch_size=CFG.train_batch_size, shuffle=True)\n    valid_data_loader = DataLoader(valid_dataset, batch_size=CFG.valid_batch_size, shuffle=False)\n    \n    model = RSNAModel(CFG.model_name)\n    model.to(device)\n    optimizer = optim.Adam(model.parameters(), \n                           lr=CFG.lr, \n                           weight_decay=CFG.wd)\n    criterion = nn.BCEWithLogitsLoss()\n    best_score = 0.0\n    best_loss = np.inf\n\n    for epoch in range(CFG.n_epochs):\n        \n        start_time = time.time()\n        avg_epoch_loss = train_step(model, \n                                    train_data_loader, \n                                    criterion, \n                                    optimizer, \n                                    epoch, \n                                    scheduler=None, \n                                    device=device)\n\n        avg_valid_loss, valid_preds = valid_step(model, \n                                                 valid_data_loader, \n                                                 criterion, \n                                                 device)\n        score = pfbeta(df_valid.cancer, valid_preds, beta=1)\n        elapsed = time.time() - start_time\n        LOGGER.info(f\"Epoch: {epoch+1} - avg_epoch_loss: {avg_epoch_loss:.5f} - avg_val_loss: {avg_valid_loss:.5f} - time: {elapsed:.0f}s\")\n        \n        if score > best_score:\n            best_score = score\n            LOGGER.info(f\"Epoch: {epoch+1} - Save best score: {best_score:.4f}\")\n            torch.save({\n                \"model\": model.state_dict(),\n                \"preds\": valid_preds\n            }, str(MODELS_DIR / f\"{CFG.model_name}_fold_{fold}_best.pth\"))\n\n    check_point = torch.load(str(MODELS_DIR / f\"{CFG.model_name}_fold_{fold}_best.pth\"))\n    df_tmp = pd.DataFrame()\n    df_tmp[\"labels\"] = df_valid.cancer\n    df_tmp[\"preds\"] = check_point[\"preds\"]\n    df_tmp[\"fold\"] = fold\n    oof_df = pd.concat([oof_df, df_tmp])","metadata":{"execution":{"iopub.status.busy":"2022-11-30T16:22:12.797222Z","iopub.execute_input":"2022-11-30T16:22:12.79769Z","iopub.status.idle":"2022-11-30T17:01:08.20802Z","shell.execute_reply.started":"2022-11-30T16:22:12.797652Z","shell.execute_reply":"2022-11-30T17:01:08.207032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV Score","metadata":{}},{"cell_type":"code","source":"cv_score = pfbeta(oof_df.labels.values, oof_df.preds.values, beta=1)\nprint(f\"CV Score: {cv_score:.5f}\")","metadata":{"execution":{"iopub.status.busy":"2022-11-30T17:06:56.394384Z","iopub.execute_input":"2022-11-30T17:06:56.394765Z","iopub.status.idle":"2022-11-30T17:06:56.494096Z","shell.execute_reply.started":"2022-11-30T17:06:56.394732Z","shell.execute_reply":"2022-11-30T17:06:56.492818Z"},"trusted":true},"execution_count":null,"outputs":[]}]}