{"cells":[{"metadata":{"papermill":{"duration":0.01433,"end_time":"2020-12-30T10:46:38.324027","exception":false,"start_time":"2020-12-30T10:46:38.309697","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# About this notebook\nI share an example how to use annotated images to improve score.  \nIf this notebook is helpful, feel free to upvote :)\n## Training strategy\n- [1st-stage training](https://www.kaggle.com/yasufuminakama/ranzcr-resnet200d-3-stage-training-step1)\n    - teacher model training for annotated image\n        - data: annotated data\n        - pretrained weight: imagenet weight\n        - `BCEWithLogitsLoss(y_preds, labels)`\n        - `y_preds: teacher model predictions for annotated image`\n- [2nd-stage training](https://www.kaggle.com/yasufuminakama/ranzcr-resnet200d-3-stage-training-step2)\n    - student model training with teacher model features\n        - data: annotated data\n        - student model pretrained weight: imagenet weight\n        - teacher model pretrained weight: 1st-stage weight\n        - `BCEWithLogitsLoss(y_preds, labels) + w * MSELoss(student_features, teacher_features)`\n        - `y_preds: student model predictions for normal image`\n        - `student_features: student model features for normal image`\n        - `teacher_features: teacher model features for annotated image`\n- [3rd-stage training](https://www.kaggle.com/yasufuminakama/ranzcr-resnet200d-3-stage-training-step3)\n    - model training\n        - data: all data\n        - pretrained weight: 2nd-stage weight\n        - `BCEWithLogitsLoss(y_preds, labels)`\n        - `y_preds: student model predictions for normal image`\n- [inference notebook](https://www.kaggle.com/yasufuminakama/ranzcr-resnet200d-3-stage-training-sub)\n    - I use private weight in this notebook"},{"metadata":{"papermill":{"duration":0.012911,"end_time":"2020-12-30T10:46:38.350228","exception":false,"start_time":"2020-12-30T10:46:38.337317","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Directory settings"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:38.383060Z","iopub.status.busy":"2020-12-30T10:46:38.382267Z","iopub.status.idle":"2020-12-30T10:46:38.384695Z","shell.execute_reply":"2020-12-30T10:46:38.385166Z"},"papermill":{"duration":0.022086,"end_time":"2020-12-30T10:46:38.385299","exception":false,"start_time":"2020-12-30T10:46:38.363213","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\nTEST_PATH = '../input/ranzcr-clip-catheter-line-classification/test'","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.012837,"end_time":"2020-12-30T10:46:38.411076","exception":false,"start_time":"2020-12-30T10:46:38.398239","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# CFG"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:38.443493Z","iopub.status.busy":"2020-12-30T10:46:38.442720Z","iopub.status.idle":"2020-12-30T10:46:38.445373Z","shell.execute_reply":"2020-12-30T10:46:38.444954Z"},"papermill":{"duration":0.021361,"end_time":"2020-12-30T10:46:38.445461","exception":false,"start_time":"2020-12-30T10:46:38.424100","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    debug=False\n    num_workers=4\n    model_name='resnet200d_320'\n    size=512\n    batch_size=128\n    seed=416\n    target_size=11\n    target_cols=['ETT - Abnormal', 'ETT - Borderline', 'ETT - Normal',\n                 'NGT - Abnormal', 'NGT - Borderline', 'NGT - Incompletely Imaged', 'NGT - Normal', \n                 'CVC - Abnormal', 'CVC - Borderline', 'CVC - Normal',\n                 'Swan Ganz Catheter Present']\n    n_fold=5\n    trn_fold=[0] # [0, 1, 2, 3, 4]","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.012783,"end_time":"2020-12-30T10:46:38.471440","exception":false,"start_time":"2020-12-30T10:46:38.458657","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Library"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:38.509521Z","iopub.status.busy":"2020-12-30T10:46:38.508973Z","iopub.status.idle":"2020-12-30T10:46:42.885519Z","shell.execute_reply":"2020-12-30T10:46:42.886514Z"},"papermill":{"duration":4.402244,"end_time":"2020-12-30T10:46:42.886702","exception":false,"start_time":"2020-12-30T10:46:38.484458","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport sys\nsys.path.append('../input/pytorch-image-models/pytorch-image-models-master')\n\nimport os\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\n\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport cv2\nfrom PIL import Image\n\nfrom matplotlib import pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\nfrom albumentations import (\n    Compose, OneOf, Normalize, Resize, RandomResizedCrop, RandomCrop, HorizontalFlip, VerticalFlip, \n    RandomBrightness, RandomContrast, RandomBrightnessContrast, Rotate, ShiftScaleRotate, Cutout, \n    IAAAdditiveGaussianNoise, Transpose\n    )\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\nimport timm\n\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport warnings \nwarnings.filterwarnings('ignore')\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.021856,"end_time":"2020-12-30T10:46:42.930516","exception":false,"start_time":"2020-12-30T10:46:42.908660","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Utils"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:42.982676Z","iopub.status.busy":"2020-12-30T10:46:42.981858Z","iopub.status.idle":"2020-12-30T10:46:42.998451Z","shell.execute_reply":"2020-12-30T10:46:42.999292Z"},"papermill":{"duration":0.047251,"end_time":"2020-12-30T10:46:42.999462","exception":false,"start_time":"2020-12-30T10:46:42.952211","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(y_true, y_pred):\n    scores = []\n    for i in range(y_true.shape[1]):\n        score = roc_auc_score(y_true[:,i], y_pred[:,i])\n        scores.append(score)\n    avg_score = np.mean(scores)\n    return avg_score, scores\n\n\ndef get_result(result_df):\n    preds = result_df[[f'pred_{c}' for c in CFG.target_cols]].values\n    labels = result_df[CFG.target_cols].values\n    score, scores = get_score(labels, preds)\n    LOGGER.info(f'Score: {score:<.4f}  Scores: {np.round(scores, decimals=4)}')\n\n\n@contextmanager\ndef timer(name):\n    t0 = time.time()\n    LOGGER.info(f'[{name}] start')\n    yield\n    LOGGER.info(f'[{name}] done in {time.time() - t0:.0f} s.')\n\n\ndef init_logger(log_file=OUTPUT_DIR+'inference.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\nLOGGER = init_logger()\n\n\ndef seed_torch(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_torch(seed=CFG.seed)","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.021159,"end_time":"2020-12-30T10:46:43.042180","exception":false,"start_time":"2020-12-30T10:46:43.021021","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Data Loading"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:43.092660Z","iopub.status.busy":"2020-12-30T10:46:43.091791Z","iopub.status.idle":"2020-12-30T10:46:43.260289Z","shell.execute_reply":"2020-12-30T10:46:43.261345Z"},"papermill":{"duration":0.197965,"end_time":"2020-12-30T10:46:43.261530","exception":false,"start_time":"2020-12-30T10:46:43.063565","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"oof_df = pd.read_csv('../input/ranzcr-exp12-step3-fold0/oof_df.csv')\nfor fold in CFG.trn_fold:\n    fold_oof_df = oof_df[oof_df['fold']==fold].reset_index(drop=True)\n    LOGGER.info(f\"========== fold: {fold} result ==========\")\n    get_result(fold_oof_df)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:43.336303Z","iopub.status.busy":"2020-12-30T10:46:43.335376Z","iopub.status.idle":"2020-12-30T10:46:43.366056Z","shell.execute_reply":"2020-12-30T10:46:43.366470Z"},"papermill":{"duration":0.073162,"end_time":"2020-12-30T10:46:43.366596","exception":false,"start_time":"2020-12-30T10:46:43.293434","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"if CFG.debug:\n    test = pd.read_csv('../input/ranzcr-clip-catheter-line-classification/sample_submission.csv', nrows=10)\nelse:\n    test = pd.read_csv('../input/ranzcr-clip-catheter-line-classification/sample_submission.csv')\n\nprint(test.shape)\ntest.head()","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.016099,"end_time":"2020-12-30T10:46:43.398877","exception":false,"start_time":"2020-12-30T10:46:43.382778","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Dataset"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:43.438622Z","iopub.status.busy":"2020-12-30T10:46:43.437967Z","iopub.status.idle":"2020-12-30T10:46:43.443053Z","shell.execute_reply":"2020-12-30T10:46:43.442609Z"},"papermill":{"duration":0.028259,"end_time":"2020-12-30T10:46:43.443148","exception":false,"start_time":"2020-12-30T10:46:43.414889","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\nclass TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['StudyInstanceUID'].values\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_name = self.file_names[idx]\n        file_path = f'{TEST_PATH}/{file_name}.jpg'\n        image = cv2.imread(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n        return image","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.015333,"end_time":"2020-12-30T10:46:43.473054","exception":false,"start_time":"2020-12-30T10:46:43.457721","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Transforms"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:43.515432Z","iopub.status.busy":"2020-12-30T10:46:43.514681Z","iopub.status.idle":"2020-12-30T10:46:43.517634Z","shell.execute_reply":"2020-12-30T10:46:43.517196Z"},"papermill":{"duration":0.028995,"end_time":"2020-12-30T10:46:43.517719","exception":false,"start_time":"2020-12-30T10:46:43.488724","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef get_transforms(*, data):\n    \n    if data == 'train':\n        return Compose([\n            #Resize(CFG.size, CFG.size),\n            RandomResizedCrop(CFG.size, CFG.size, scale=(0.85, 1.0)),\n            HorizontalFlip(p=0.5),\n            RandomBrightnessContrast(p=0.2, brightness_limit=(-0.2, 0.2), contrast_limit=(-0.2, 0.2)),\n            HueSaturationValue(p=0.2, hue_shift_limit=0.2, sat_shift_limit=0.2, val_shift_limit=0.2),\n            ShiftScaleRotate(p=0.2, shift_limit=0.0625, scale_limit=0.2, rotate_limit=20),\n            CoarseDropout(p=0.2),\n            Cutout(p=0.2, max_h_size=16, max_w_size=16, fill_value=(0., 0., 0.), num_holes=16),\n            Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n\n    elif data == 'valid':\n        return Compose([\n            Resize(CFG.size, CFG.size),\n            Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:43.556040Z","iopub.status.busy":"2020-12-30T10:46:43.555532Z","iopub.status.idle":"2020-12-30T10:46:44.061161Z","shell.execute_reply":"2020-12-30T10:46:44.061741Z"},"papermill":{"duration":0.528631,"end_time":"2020-12-30T10:46:44.061883","exception":false,"start_time":"2020-12-30T10:46:43.533252","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"train_dataset = TestDataset(test, transform=get_transforms(data='valid'))\n\nfor i in range(1):\n    image = train_dataset[i]\n    plt.imshow(image[0])\n    plt.show()\n    plt.imshow(image[0].flip(-1))\n    plt.show()","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.018854,"end_time":"2020-12-30T10:46:44.099473","exception":false,"start_time":"2020-12-30T10:46:44.080619","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# MODEL"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:44.148742Z","iopub.status.busy":"2020-12-30T10:46:44.146993Z","iopub.status.idle":"2020-12-30T10:46:44.149379Z","shell.execute_reply":"2020-12-30T10:46:44.149791Z"},"papermill":{"duration":0.031623,"end_time":"2020-12-30T10:46:44.149901","exception":false,"start_time":"2020-12-30T10:46:44.118278","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# MODEL\n# ====================================================\nclass CustomResNet200D(nn.Module):\n    def __init__(self, model_name='resnet200d_320', pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=False)\n        if pretrained:\n            pretrained_path = '../input/resnet200d-pretrained-weight/resnet200d_ra2-bdba9bf9.pth'\n            self.model.load_state_dict(torch.load(pretrained_path))\n            print(f'load {model_name} pretrained model')\n        n_features = self.model.fc.in_features\n        self.model.global_pool = nn.Identity()\n        self.model.fc = nn.Identity()\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(n_features, CFG.target_size)\n\n    def forward(self, x):\n        bs = x.size(0)\n        features = self.model(x)\n        pooled_features = self.pooling(features).view(bs, -1)\n        output = self.fc(pooled_features)\n        return output","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.017883,"end_time":"2020-12-30T10:46:44.185844","exception":false,"start_time":"2020-12-30T10:46:44.167961","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Helper functions"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:44.232971Z","iopub.status.busy":"2020-12-30T10:46:44.231189Z","iopub.status.idle":"2020-12-30T10:46:44.233599Z","shell.execute_reply":"2020-12-30T10:46:44.234065Z"},"papermill":{"duration":0.029995,"end_time":"2020-12-30T10:46:44.234173","exception":false,"start_time":"2020-12-30T10:46:44.204178","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\ndef inference(models, test_loader, device):\n    tk0 = tqdm(enumerate(test_loader), total=len(test_loader))\n    probs = []\n    for i, (images) in tk0:\n        images = images.to(device)\n        avg_preds = []\n        for model in models:\n            with torch.no_grad():\n                y_preds1 = model(images)\n                y_preds2 = model(images.flip(-1))\n            y_preds = (y_preds1.sigmoid().to('cpu').numpy() + y_preds2.sigmoid().to('cpu').numpy()) / 2\n            avg_preds.append(y_preds)\n        avg_preds = np.mean(avg_preds, axis=0)\n        probs.append(avg_preds)\n    probs = np.concatenate(probs)\n    return probs","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.018023,"end_time":"2020-12-30T10:46:44.270096","exception":false,"start_time":"2020-12-30T10:46:44.252073","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# inference"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:46:44.316293Z","iopub.status.busy":"2020-12-30T10:46:44.315652Z","iopub.status.idle":"2020-12-30T10:46:53.986207Z","shell.execute_reply":"2020-12-30T10:46:53.986600Z"},"papermill":{"duration":9.697736,"end_time":"2020-12-30T10:46:53.986726","exception":false,"start_time":"2020-12-30T10:46:44.288990","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"%%time\n\nmodel = CustomResNet200D(CFG.model_name, pretrained=False)\nmodel_path = '../input/ranzcr-exp12-step3-fold0/resnet200d_320_fold0_best_loss.pth'\nmodel.load_state_dict(torch.load(model_path)['model'])\nmodel.eval()\n\nmodels = [model.to(device)]","execution_count":null,"outputs":[]},{"metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","execution":{"iopub.execute_input":"2020-12-30T10:46:54.031937Z","iopub.status.busy":"2020-12-30T10:46:54.031263Z","iopub.status.idle":"2020-12-30T10:50:43.727882Z","shell.execute_reply":"2020-12-30T10:50:43.727137Z"},"papermill":{"duration":229.721991,"end_time":"2020-12-30T10:50:43.728007","exception":false,"start_time":"2020-12-30T10:46:54.006016","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"test_dataset = TestDataset(test, transform=get_transforms(data='valid'))\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, \n                         num_workers=CFG.num_workers, pin_memory=True)\npredictions = inference(models, test_loader, device)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-30T10:50:43.775557Z","iopub.status.busy":"2020-12-30T10:50:43.772419Z","iopub.status.idle":"2020-12-30T10:50:44.146048Z","shell.execute_reply":"2020-12-30T10:50:44.145570Z"},"papermill":{"duration":0.39865,"end_time":"2020-12-30T10:50:44.146154","exception":false,"start_time":"2020-12-30T10:50:43.747504","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# submission\ntest[CFG.target_cols] = predictions\ntest[['StudyInstanceUID'] + CFG.target_cols].to_csv(OUTPUT_DIR+'submission.csv', index=False)\ntest.head()","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}