{"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":"<div class=\"alert alert-block alert-success\" style=\"font-size:30px\">\n[train] Pytorch:aux targets+weighted loss+thresholds\n</div>\n\n\n<div class=\"alert alert-block alert-info\">\n    <ul>\n        <li>\n            📌 This is a training part. For <b>inference</b> refer to: <a href=\"https://www.kaggle.com/code/vslaykovsky/infer-rsna-breast-cancer-effnetv2\">[infer] RSNA Breast Cancer EffNetV2</a>\n        </li>\n        <li>\n            📌 To run the notebook in non-interactive mode, <b>set Add-ons->Secrets->WANDB_API_KEY</b> secret key to the value of your <a href=\"https://wandb.ai/authorize\">Wandb API key. </a>\n        </li>\n</div>\n\n\n\nSome notes on the implementation:\n* \"cancer\" targets are imbalanced, so we use weighted loss to counteract. See `CANCER_LOSS_WEIGHT` for details. Additionally, best values of thresholds are selected based on performance on evaluation set.\n* [**Preprocessed dataset of 1024x512 png images**](https://www.kaggle.com/code/vslaykovsky/rsna-cut-off-empty-space-from-images) is used to train the model. This speeds up dataloader by ~10x with no performance degradation. Additionally empty space is removed from images\n* Cross-entropy auxilliary targets are made of the following `train.csv` columns: `['site_id', 'laterality', 'view', 'implant', 'biopsy', 'invasive', 'BIRADS', 'density', 'difficult_negative_case', 'machine_id', 'age']`. Auxilliary targets help to learn combined *.CSV + *.PNG data distribution which helps to improve performance of the main classifier.\n\n\n**UPDATE1**\n* Migrated to `timm`\n* Random augmentations with `timm.data.create_transform`\n* Hyperparameter optimization with Wandb sweeps.\n\n**UPDATE2**\n* Larger images (only region of interest 1024x512)\n\n<div class=\"alert alert-block alert-danger\" style=\"text-align:center; font-size:16px;\">\n    Thanks for ▲upvoting▲ Kaggle notebooks when you click \"Copy&Edit\". This really motivates authors to produce more quality work.\n</div>\n","metadata":{"papermill":{"duration":0.009487,"end_time":"2022-12-01T15:09:08.527773","exception":false,"start_time":"2022-12-01T15:09:08.518286","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n1. Imports, constants, dependencies\n</div>","metadata":{"papermill":{"duration":0.007035,"end_time":"2022-12-01T15:09:08.543072","exception":false,"start_time":"2022-12-01T15:09:08.536037","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\nimport gc\nimport os\n\n# import cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport sklearn\nimport torch\nfrom PIL import Image\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.preprocessing import LabelEncoder\nimport sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\nfrom timm import create_model, list_models\nfrom timm.data import create_transform\nfrom torch.cuda.amp import GradScaler, autocast\nfrom tqdm import tqdm\n\nimport wandb\n\nplt.rcParams['figure.figsize'] = (20, 5)\npd.set_option('display.max_rows', 100)\npd.set_option('display.max_columns', 1000)\n\n# Common\ntry:\n    from kaggle_secrets import UserSecretsClient\n    IS_KAGGLE = True\nexcept:\n    IS_KAGGLE = False\n\nos.environ[\"WANDB_MODE\"] = \"online\"\nif os.environ[\"WANDB_MODE\"] == \"online\":\n    if IS_KAGGLE:\n        os.environ['WANDB_API_KEY'] = UserSecretsClient().get_secret(\"WANDB_API_KEY\")\n\n        \nRSNA_2022_PATH = '../input/rsna-breast-cancer-detection'\nTRAIN_IMAGES_PATH = f'/kaggle/input/rsna-cut-off-empty-space-from-images'\nMAX_TRAIN_BATCHES = 40000\nMAX_EVAL_BATCHES = 400\nMODELS_PATH = '/kaggle/input/wandb-models/models'\nNUM_WORKERS = 8\nPREDICT_MAX_BATCHES = 1e9\nN_FOLDS = 5\nFOLDS = np.array(os.environ.get('FOLDS', '0,1,2,3,4').split(',')).astype(int)\nWANDB_SWEEP_PROJECT = 'rsna-breast-cancer-sweeps'\nWANDB_PROJECT = 'RSNA-breast-cancer-v4'\n\nCATEGORY_AUX_TARGETS = ['site_id', 'laterality', 'view', 'implant', 'biopsy', 'invasive', 'BIRADS', 'density', 'difficult_negative_case', 'machine_id', 'age']\nTARGET = 'cancer'\nALL_FEAT = [TARGET] + CATEGORY_AUX_TARGETS\n\nif not IS_KAGGLE:\n    print('Running locally')\n    RSNA_2022_PATH = 'data'\n    TRAIN_IMAGES_PATH = 'data/roi1024/train_images'\n    MODELS_PATH = 'models_roi_1024_v2'\n    os.environ['WANDB_API_KEY'] = 'adc8abc0714ba20c3a534b907b9d6beec640f847'\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"lines_to_next_cell":2,"papermill":{"duration":2.953187,"end_time":"2022-12-01T15:10:25.264744","exception":false,"start_time":"2022-12-01T15:10:22.311557","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:39:52.905088Z","iopub.execute_input":"2022-12-13T22:39:52.905729Z","iopub.status.idle":"2022-12-13T22:39:58.458737Z","shell.execute_reply.started":"2022-12-13T22:39:52.905649Z","shell.execute_reply":"2022-12-13T22:39:58.457489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Switchers\nDEBUG = os.environ.get('DEBUG', 'true').lower() == 'true'\nWANDB_SWEEP = False\nTRAIN = False\nCV = True","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-13T22:39:58.461427Z","iopub.execute_input":"2022-12-13T22:39:58.462341Z","iopub.status.idle":"2022-12-13T22:39:58.469529Z","shell.execute_reply.started":"2022-12-13T22:39:58.462287Z","shell.execute_reply":"2022-12-13T22:39:58.467609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Config\n\nclass Config:\n    # These are optimal parameters collected from https://wandb.ai/vslaykovsky/rsna-breast-cancer-sweeps/sweeps/k281hlr9?workspace=user-vslaykovsky\n    ONE_CYCLE = True\n    ONE_CYCLE_PCT_START = 0.1\n    ADAMW = False\n    ADAMW_DECAY = 0.024\n    # ONE_CYCLE_MAX_LR = float(os.environ.get('LR', '0.0008'))\n    ONE_CYCLE_MAX_LR = float(os.environ.get('LR', '0.0004'))\n    EPOCHS = int(os.environ.get('EPOCHS', 3))\n    MODEL_TYPE = os.environ.get('MODEL', 'seresnext50_32x4d')\n    DROPOUT = float(os.environ.get('DROPOUT', 0.0))\n    AUG = os.environ.get('AUG', 'true').lower() == 'true'\n    AUX_LOSS_WEIGHT = 94\n    POSITIVE_TARGET_WEIGHT=20\n    # BATCH_SIZE = 32\n    BATCH_SIZE = 16\n    AUTO_AUG_M = 10\n    AUTO_AUG_N = 2\n    TTA = False\n\n\nWANDB_RUN_NAME = f'{Config.MODEL_TYPE}_lr{Config.ONE_CYCLE_MAX_LR}_ep{Config.EPOCHS}_bs{Config.BATCH_SIZE}_pw{Config.POSITIVE_TARGET_WEIGHT}_' +\\\nf'aux{Config.AUX_LOSS_WEIGHT}_{\"adamw\" if Config.ADAMW else \"adam\"}_{\"aug\" if Config.AUG else \"noaug\"}_drop{Config.DROPOUT}'\nprint('run', WANDB_RUN_NAME, 'folds', FOLDS)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-13T22:39:58.471686Z","iopub.execute_input":"2022-12-13T22:39:58.472298Z","iopub.status.idle":"2022-12-13T22:39:58.493808Z","shell.execute_reply.started":"2022-12-13T22:39:58.472264Z","shell.execute_reply":"2022-12-13T22:39:58.492534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n2. Loading train/eval/test dataframes\n</div>","metadata":{"papermill":{"duration":0.007596,"end_time":"2022-12-01T15:10:25.2807","exception":false,"start_time":"2022-12-01T15:10:25.273104","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"1. Loading data from competition dataset folder `/kaggle/input/rsna-breast-cancer-detection/train.csv`\n2. Adding `Splits` column to facilitate train/eval splits.","metadata":{"papermill":{"duration":0.007673,"end_time":"2022-12-01T15:10:25.296109","exception":false,"start_time":"2022-12-01T15:10:25.288436","status":"completed"},"tags":[]}},{"cell_type":"code","source":"df_train = pd.read_csv(f'{RSNA_2022_PATH}/train.csv')\ndf_train","metadata":{"papermill":{"duration":0.144598,"end_time":"2022-12-01T15:10:25.448614","exception":false,"start_time":"2022-12-01T15:10:25.304016","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:39:58.496665Z","iopub.execute_input":"2022-12-13T22:39:58.497205Z","iopub.status.idle":"2022-12-13T22:39:58.683428Z","shell.execute_reply.started":"2022-12-13T22:39:58.497047Z","shell.execute_reply":"2022-12-13T22:39:58.680258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train","metadata":{"execution":{"iopub.status.busy":"2022-12-13T22:39:58.685956Z","iopub.execute_input":"2022-12-13T22:39:58.686372Z","iopub.status.idle":"2022-12-13T22:39:58.741324Z","shell.execute_reply.started":"2022-12-13T22:39:58.686333Z","shell.execute_reply":"2022-12-13T22:39:58.736008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedGroupKFold\n\nsplit = StratifiedGroupKFold(N_FOLDS)\nfor k, (_, test_idx) in enumerate(split.split(df_train, df_train.cancer, groups=df_train.patient_id)):\n    df_train.loc[test_idx, 'split'] = k\ndf_train.split = df_train.split.astype(int)\ndf_train.groupby('split').cancer.mean()","metadata":{"lines_to_next_cell":2,"papermill":{"duration":0.0785,"end_time":"2022-12-01T15:10:25.535694","exception":false,"start_time":"2022-12-01T15:10:25.457194","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:39:58.744998Z","iopub.execute_input":"2022-12-13T22:39:58.745368Z","iopub.status.idle":"2022-12-13T22:40:03.338572Z","shell.execute_reply.started":"2022-12-13T22:39:58.745329Z","shell.execute_reply":"2022-12-13T22:40:03.337316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.age.fillna(df_train.age.mean(), inplace=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-13T22:40:03.340505Z","iopub.execute_input":"2022-12-13T22:40:03.340924Z","iopub.status.idle":"2022-12-13T22:40:03.347974Z","shell.execute_reply.started":"2022-12-13T22:40:03.340884Z","shell.execute_reply":"2022-12-13T22:40:03.346757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train['age'] = pd.qcut(df_train.age, 10, labels=range(10), retbins=False).astype(int)\ndf_train","metadata":{"execution":{"iopub.status.busy":"2022-12-13T22:40:03.349931Z","iopub.execute_input":"2022-12-13T22:40:03.3503Z","iopub.status.idle":"2022-12-13T22:40:03.385322Z","shell.execute_reply.started":"2022-12-13T22:40:03.350264Z","shell.execute_reply":"2022-12-13T22:40:03.384253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train[CATEGORY_AUX_TARGETS] = df_train[CATEGORY_AUX_TARGETS].apply(LabelEncoder().fit_transform)","metadata":{"lines_to_next_cell":2,"execution":{"iopub.status.busy":"2022-12-13T22:40:03.386482Z","iopub.execute_input":"2022-12-13T22:40:03.387008Z","iopub.status.idle":"2022-12-13T22:40:03.451703Z","shell.execute_reply.started":"2022-12-13T22:40:03.386978Z","shell.execute_reply":"2022-12-13T22:40:03.450744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train[ALL_FEAT]","metadata":{"execution":{"iopub.status.busy":"2022-12-13T22:40:03.455633Z","iopub.execute_input":"2022-12-13T22:40:03.455918Z","iopub.status.idle":"2022-12-13T22:40:03.477507Z","shell.execute_reply.started":"2022-12-13T22:40:03.455891Z","shell.execute_reply":"2022-12-13T22:40:03.47638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n 3. Dataset class\n</div>\n\n`BreastCancerDataSet` class returns individual images. It uses a dataframe parameter `df` as a source of metadata to locate and load images from `path` folder. It accepts transforms parameter to apply transforms on images.\n\n`get_transforms` generates a default set of transforms for an ImageNet-based model. Additionally with is_training=True it applies a set of random augmentations","metadata":{"papermill":{"duration":0.008451,"end_time":"2022-12-01T15:10:25.608024","exception":false,"start_time":"2022-12-01T15:10:25.599573","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import torchvision\n\ndef get_transforms(aug=False):\n    \"\"\"\n    # old transforms\n    create_transform(\n        (1024, 512), \n        mean=0.53, #(0.53, 0.53, 0.53),\n        std=0.23, #(0.23, 0.23, 0.23),\n        is_training=is_training, \n        auto_augment=f'rand-m{config.AUTO_AUG_M}-n{config.AUTO_AUG_N}'\n    )\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\n    return lambda img: transforms(img)\n\nif DEBUG:\n    tfm = get_transforms(aug=True)\n    img = Image.open(f\"{TRAIN_IMAGES_PATH}/10006/1459541791.png\")\n    print(img.size)\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":{"execution":{"iopub.status.busy":"2022-12-13T22:40:03.480323Z","iopub.execute_input":"2022-12-13T22:40:03.480934Z","iopub.status.idle":"2022-12-13T22:40:06.649254Z","shell.execute_reply.started":"2022-12-13T22:40:03.480897Z","shell.execute_reply":"2022-12-13T22:40:06.648279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\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\n        path = f'{self.path}/{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        if TARGET in self.df.columns:\n            cancer_target = torch.as_tensor(self.df.iloc[i].cancer)\n            cat_aux_targets = torch.as_tensor(self.df.iloc[i][CATEGORY_AUX_TARGETS])\n            return img, cancer_target, cat_aux_targets\n\n        return img\n\n    def __len__(self):\n        return len(self.df)\n\nds_train = BreastCancerDataSet(df_train, TRAIN_IMAGES_PATH, get_transforms(aug=True))\nif DEBUG:\n    X, y_cancer, y_aux = ds_train[42]\n    print(X.shape, y_cancer.shape, y_aux.shape)","metadata":{"papermill":{"duration":0.01949,"end_time":"2022-12-01T15:10:25.690462","exception":false,"start_time":"2022-12-01T15:10:25.670972","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:40:06.650247Z","iopub.execute_input":"2022-12-13T22:40:06.650562Z","iopub.status.idle":"2022-12-13T22:40:06.696705Z","shell.execute_reply.started":"2022-12-13T22:40:06.650531Z","shell.execute_reply":"2022-12-13T22:40:06.695662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n    3. Model\n</div>\n\nWe create a backbone with `timm.create_model`. The backbone is strip of classification layer (num_classes=0). We create multiple linear heads (see self.nn_cancer, self.nn_aux). nn_aux is used to predict other columns of train.csv\n","metadata":{"papermill":{"duration":0.008737,"end_time":"2022-12-01T15:10:27.974999","exception":false,"start_time":"2022-12-01T15:10:27.966262","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class BreastCancerModel(torch.nn.Module):\n    def __init__(self, aux_classes, model_type=Config.MODEL_TYPE, dropout=0.):\n        super().__init__()\n        self.model = create_model(model_type, pretrained=True, 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\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\nAUX_TARGET_NCLASSES = df_train[CATEGORY_AUX_TARGETS].max() + 1\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","metadata":{"execution":{"iopub.status.busy":"2022-12-13T22:40:06.698474Z","iopub.execute_input":"2022-12-13T22:40:06.698879Z","iopub.status.idle":"2022-12-13T22:40:24.016746Z","shell.execute_reply.started":"2022-12-13T22:40:06.698844Z","shell.execute_reply":"2022-12-13T22:40:24.015668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n    4. Train: training/evaluation loop\n</div>\n\n* We use low precision to speed up training (see `autocast`)\n* We use pF1 target metric (see `pfbeta`). After every epoch we calculate the best value of the metric by search for an optimal threshold.","metadata":{"papermill":{"duration":0.00872,"end_time":"2022-12-01T15:10:30.059819","exception":false,"start_time":"2022-12-01T15:10:30.051099","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def save_model(name, model, thres, model_type):\n    torch.save({'model': model.state_dict(), 'threshold': thres, 'model_type': model_type}, f'{name}')","metadata":{"papermill":{"duration":0.018495,"end_time":"2022-12-01T15:10:30.087362","exception":false,"start_time":"2022-12-01T15:10:30.068867","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:40:24.018268Z","iopub.execute_input":"2022-12-13T22:40:24.018972Z","iopub.status.idle":"2022-12-13T22:40:24.025603Z","shell.execute_reply.started":"2022-12-13T22:40:24.018931Z","shell.execute_reply":"2022-12-13T22:40:24.024564Z"},"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    return model, data['threshold'], data['model_type']\n\n\nif DEBUG:\n    # quick test\n    model = torch.nn.Linear(2, 1)\n    save_model('testmodel', model, thres=0.123, model_type='abc')\n\n    model1, thres, model_type = load_model('testmodel', model=torch.nn.Linear(2, 1))\n    assert torch.all(\n        next(iter(model1.parameters())) == next(iter(model.parameters()))\n    ).item(), \"Loading/saving is inconsistent!\"\n    print(thres, model_type)","metadata":{"papermill":{"duration":0.022738,"end_time":"2022-12-01T15:10:30.119085","exception":false,"start_time":"2022-12-01T15:10:30.096347","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:40:24.026803Z","iopub.execute_input":"2022-12-13T22:40:24.02745Z","iopub.status.idle":"2022-12-13T22:40:26.893679Z","shell.execute_reply.started":"2022-12-13T22:40:24.027415Z","shell.execute_reply":"2022-12-13T22:40:26.891918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pfbeta(labels, predictions, beta=1.):\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        else:\n            cfp += prediction\n\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / max(y_true_count, 1)  # avoid / 0\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\ndef optimal_f1(labels, predictions):\n    thres = np.linspace(0, 1, 101)\n    f1s = [pfbeta(labels, predictions > thr) for thr in thres]\n    idx = np.argmax(f1s)\n    return f1s[idx], thres[idx]\n\ndef evaluate_model(model: BreastCancerModel, ds, max_batches=PREDICT_MAX_BATCHES, shuffle=False, config=Config):\n    torch.manual_seed(42)\n    model = model.to(DEVICE)\n    dl_test = torch.utils.data.DataLoader(ds, batch_size=config.BATCH_SIZE, shuffle=shuffle, num_workers=NUM_WORKERS, pin_memory=False)\n    pred_cancer = []\n    with torch.no_grad():\n        \n        model.eval()\n        cancer_losses = []\n        aux_losses = []\n        losses = []\n        targets = []\n        with tqdm(dl_test, desc='Eval', mininterval=30) as progress:\n            for i, (X, y_cancer, y_aux) in enumerate(progress):\n                with autocast(enabled=True):\n                    y_aux = y_aux.to(DEVICE)\n                    X = X.to(DEVICE)\n                    y_cancer_pred, aux_pred = model.forward(X)\n                    if config.TTA:\n                        y_cancer_pred2, aux_pred2 = model.forward(torch.flip(X, dims=[-1])) # horizontal mirror\n                        y_cancer_pred = (y_cancer_pred + y_cancer_pred2) / 2\n                        aux_pred = [(v1 + v2) / 2 for v1, v2 in zip(aux_pred, aux_pred2)]\n\n                    cancer_loss = torch.nn.functional.binary_cross_entropy_with_logits(\n                        y_cancer_pred, \n                        y_cancer.to(float).to(DEVICE),\n                        pos_weight=torch.tensor([config.POSITIVE_TARGET_WEIGHT]).to(DEVICE)\n                    ).item()\n                    aux_loss = torch.mean(torch.stack([torch.nn.functional.cross_entropy(aux_pred[i], y_aux[:, i]) for i in range(y_aux.shape[-1])])).item()\n                    pred_cancer.append(torch.sigmoid(y_cancer_pred))\n                    cancer_losses.append(cancer_loss)\n                    aux_losses.append(aux_loss)\n                    losses.append(cancer_loss + config.AUX_LOSS_WEIGHT * aux_loss)\n                    targets.append(y_cancer.cpu().numpy())\n                if i >= max_batches:\n                    break\n        targets = np.concatenate(targets)\n        pred = torch.concat(pred_cancer).cpu().numpy()\n        pf1, thres = optimal_f1(targets, pred)\n        return np.mean(cancer_losses), (pf1, thres), pred, np.mean(losses), np.mean(aux_losses)\n\n\n# quick test\nif DEBUG:\n\n    m = BreastCancerModel(AUX_TARGET_NCLASSES)\n    closs, f1, pred, loss, aloss = evaluate_model(m, ds_train, max_batches=2)\n    del m\n    closs, f1, pred.shape, loss, aloss","metadata":{"papermill":{"duration":3.102896,"end_time":"2022-12-01T15:10:33.26317","exception":false,"start_time":"2022-12-01T15:10:30.160274","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:40:26.895118Z","iopub.execute_input":"2022-12-13T22:40:26.89569Z","iopub.status.idle":"2022-12-13T22:40:42.955995Z","shell.execute_reply.started":"2022-12-13T22:40:26.895653Z","shell.execute_reply":"2022-12-13T22:40:42.954858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def gc_collect():\n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"papermill":{"duration":0.016662,"end_time":"2022-12-01T15:10:33.289627","exception":false,"start_time":"2022-12-01T15:10:33.272965","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:40:42.957998Z","iopub.execute_input":"2022-12-13T22:40:42.95843Z","iopub.status.idle":"2022-12-13T22:40:42.967553Z","shell.execute_reply.started":"2022-12-13T22:40:42.958384Z","shell.execute_reply":"2022-12-13T22:40:42.966151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def add_weight_decay(model, weight_decay=1e-5, skip_list=()):\n    decay = []\n    no_decay = []\n    for name, param in model.named_parameters():\n        if not param.requires_grad:\n            continue\n        if len(param.shape) == 1 or np.any([v in name.lower()  for v in skip_list]):\n            # print(name, 'no decay')\n            no_decay.append(param)\n        else:\n            # print(name, 'decay')\n            decay.append(param)\n    return [\n        {'params': no_decay, 'weight_decay': 0.},\n        {'params': decay, 'weight_decay': weight_decay}]","metadata":{"execution":{"iopub.status.busy":"2022-12-13T22:40:42.969421Z","iopub.execute_input":"2022-12-13T22:40:42.969937Z","iopub.status.idle":"2022-12-13T22:40:42.981899Z","shell.execute_reply.started":"2022-12-13T22:40:42.969895Z","shell.execute_reply":"2022-12-13T22:40:42.980685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(ds_train, ds_eval, logger, name, config=Config, do_save_model=True):\n    torch.manual_seed(42)\n    dl_train = torch.utils.data.DataLoader(ds_train, batch_size=config.BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True)\n\n    model = BreastCancerModel(AUX_TARGET_NCLASSES, config.MODEL_TYPE, config.DROPOUT).to(DEVICE)\n\n    if config.ADAMW:\n        optim = torch.optim.AdamW(add_weight_decay(model, weight_decay=config.ADAMW_DECAY, skip_list=['bias']), lr=config.ONE_CYCLE_MAX_LR, betas=(0.9, 0.999), weight_decay=config.ADAMW_DECAY)\n    else:\n        optim = torch.optim.Adam(model.parameters())\n\n\n    scheduler = None\n    if config.ONE_CYCLE:\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(optim, max_lr=config.ONE_CYCLE_MAX_LR, epochs=config.EPOCHS,\n                                                        steps_per_epoch=len(dl_train),\n                                                        pct_start=config.ONE_CYCLE_PCT_START)\n        \n    \n\n    scaler = GradScaler()\n    best_eval_score = 0\n    for epoch in tqdm(range(config.EPOCHS), desc='Epoch'):\n\n        model.train()\n        with tqdm(dl_train, desc='Train', mininterval=30) as train_progress:\n            for batch_idx, (X, y_cancer, y_aux) in enumerate(train_progress):\n                y_aux = y_aux.to(DEVICE)\n\n                optim.zero_grad()\n                # Using mixed precision training\n                with autocast():\n                    y_cancer_pred, aux_pred = model.forward(X.to(DEVICE))\n                    cancer_loss = torch.nn.functional.binary_cross_entropy_with_logits(\n                        y_cancer_pred,\n                        y_cancer.to(float).to(DEVICE),\n                        pos_weight=torch.tensor([config.POSITIVE_TARGET_WEIGHT]).to(DEVICE)\n                    )\n                    aux_loss = torch.mean(torch.stack([torch.nn.functional.cross_entropy(aux_pred[i], y_aux[:, i]) for i in range(y_aux.shape[-1])]))\n                    loss = cancer_loss + config.AUX_LOSS_WEIGHT * aux_loss\n                    if np.isinf(loss.item()) or np.isnan(loss.item()):\n                        print(f'Bad loss, skipping the batch {batch_idx}')\n                        del loss, cancer_loss, y_cancer_pred\n                        gc_collect()\n                        continue\n\n                # scaler is needed to prevent \"gradient underflow\"\n                scaler.scale(loss).backward()\n                scaler.step(optim)\n                if scheduler is not None:\n                    scheduler.step()\n                    \n                scaler.update()\n\n                lr = scheduler.get_last_lr()[0] if scheduler else config.ONE_CYCLE_MAX_LR\n                logger.log({'loss': (loss.item()),\n                            'cancer_loss': cancer_loss.item(),\n                            'aux_loss': aux_loss.item(),\n                            'lr': lr,\n                            'epoch': epoch})\n\n\n        if ds_eval is not None and MAX_EVAL_BATCHES > 0:\n            cancer_loss, (f1, thres), _, loss, aux_loss = evaluate_model(\n                model, ds_eval, max_batches=MAX_EVAL_BATCHES, shuffle=False, config=config)\n\n            if f1 > best_eval_score:\n                best_eval_score = f1\n                if do_save_model:\n                    save_model(name, model, thres, config.MODEL_TYPE)\n                    art = wandb.Artifact(\"rsna-breast-cancer\", type=\"model\")\n                    art.add_file(f'{name}')\n                    logger.log_artifact(art)\n\n            logger.log(\n                {\n                    'eval_cancer_loss': cancer_loss,\n                    'eval_f1': f1,\n                    'max_eval_f1': best_eval_score,\n                    'eval_f1_thres': thres,\n                    'eval_loss': loss,\n                    'eval_aux_loss': aux_loss,\n                    'epoch': epoch\n                }\n            )\n\n    return model\n\n\n# N-fold models. Can be used to estimate accurate CV score and in ensembled submissions.\nif TRAIN:\n    for fold in FOLDS:\n        name = f'{WANDB_RUN_NAME}-f{fold}'\n        with wandb.init(project=WANDB_PROJECT, name=name, group=WANDB_RUN_NAME) as run:\n            gc_collect()\n            ds_train = BreastCancerDataSet(df_train.query('split != @fold'), TRAIN_IMAGES_PATH, get_transforms(aug=Config.AUG))\n            ds_eval = BreastCancerDataSet(df_train.query('split == @fold'), TRAIN_IMAGES_PATH, get_transforms(aug=False))\n            train_model(ds_train, ds_eval, run, f'model-f{fold}')","metadata":{"papermill":{"duration":0.026306,"end_time":"2022-12-01T15:10:33.325102","exception":false,"start_time":"2022-12-01T15:10:33.298796","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:40:42.984082Z","iopub.execute_input":"2022-12-13T22:40:42.98478Z","iopub.status.idle":"2022-12-13T22:40:43.014391Z","shell.execute_reply.started":"2022-12-13T22:40:42.984746Z","shell.execute_reply":"2022-12-13T22:40:43.013332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%wandb vslaykovsky/RSNA-breast-cancer-v4 -h 1000","metadata":{"papermill":{"duration":0.021622,"end_time":"2022-12-01T15:10:33.35569","exception":false,"start_time":"2022-12-01T15:10:33.334068","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-18T11:46:08.533321Z","iopub.execute_input":"2022-12-18T11:46:08.533685Z","iopub.status.idle":"2022-12-18T11:46:19.315116Z","shell.execute_reply.started":"2022-12-18T11:46:08.533653Z","shell.execute_reply":"2022-12-18T11:46:19.313315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n    5. Wandb Sweeps\n</div>\n\n\nEnable `WANDB_SEEP` to use the feature.\n\n1. Run the cell once to generate SWEEP_ID\n2. Set the SWEEP_ID environment variable to run the agent.\n3. Run one or more agents to start hyperparameter optimization.","metadata":{}},{"cell_type":"code","source":"#  %env SWEEP_ID=tfi5ayrd","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-13T22:40:43.045544Z","iopub.execute_input":"2022-12-13T22:40:43.047706Z","iopub.status.idle":"2022-12-13T22:40:43.053163Z","shell.execute_reply.started":"2022-12-13T22:40:43.047672Z","shell.execute_reply":"2022-12-13T22:40:43.052283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nif WANDB_SWEEP:\n    sweep_id = os.environ.get('SWEEP_ID')\n    print('wandb sweep ', sweep_id)\n\n    if sweep_id is None:\n        \"\"\"\n        First run. Generate sweep_id.\n        \"\"\"\n        sweep_id = wandb.sweep(sweep={\n            'method': 'bayes',\n            'name': 'rsna-sweep',\n            'metric': {'goal': 'maximize', 'name': 'max_eval_f1'},\n            'parameters':\n                {\n                    'ONE_CYCLE': {'values': [True, False]},\n                    'ONE_CYCLE_PCT_START': {'values': [0.1]},\n                    'ADAMW': {'values': [True, False]},\n                    'ADAMW_DECAY': {'min': 0.001, 'max': 0.1, 'distribution': 'log_uniform_values'},\n                    'ONE_CYCLE_MAX_LR': {'min': 1e-5, 'max': 1e-3, 'distribution': 'log_uniform_values'},\n                    'EPOCHS': {'min': 1, 'max': 12, 'distribution': 'q_log_uniform_values'},\n                    'MODEL_TYPE': {'values': ['resnext50_32x4d', 'efficientnetv2_rw_s', 'seresnext50_32x4d', 'inception_v4', 'efficientnet_b4']},\n                    'DROPOUT': {'values': [0., 0.2]},\n                    'AUG': {'values': [True, False]},\n                    'AUX_LOSS_WEIGHT': {'min': 0.01, 'max': 100., 'distribution': 'log_uniform_values'},\n                    'POSITIVE_TARGET_WEIGHT': {'min': 1., 'max': 60., 'distribution': 'uniform'},\n                    'BATCH_SIZE': {'values': [32]},\n                    'AUTO_AUG_M': {'min': 1, 'max': 20, 'distribution': 'q_log_uniform_values'},\n                    'AUTO_AUG_N': {'min': 1, 'max': 6, 'distribution': 'q_uniform'},\n                    'TTA': {'values': [False]},\n                }\n        }, project=WANDB_SWEEP_PROJECT)\n        print('Generated sweep id', sweep_id)\n    else:\n        \"\"\"\n        Agent run. Use sweep_id generated above to produce (semi)-random hyperparameters run.config\n        \"\"\"\n        def wandb_callback():\n            with wandb.init() as run:\n                print('params', run.config)\n                fold = 0\n                ds_train = BreastCancerDataSet(df_train.query('split != @fold'), TRAIN_IMAGES_PATH, get_transforms(aug=run.config.AUG))\n                ds_eval = BreastCancerDataSet(df_train.query('split == @fold'), TRAIN_IMAGES_PATH, get_transforms(aug=False))\n                train_model(ds_train, ds_eval, run, f'model-f{fold}', config=run.config, do_save_model=False)\n\n\n        # Start sweep job.\n        wandb.agent(sweep_id, project=WANDB_SWEEP_PROJECT, function=wandb_callback, count=100000)","metadata":{"execution":{"iopub.status.busy":"2022-12-13T22:40:43.057756Z","iopub.execute_input":"2022-12-13T22:40:43.060355Z","iopub.status.idle":"2022-12-13T22:40:43.077285Z","shell.execute_reply.started":"2022-12-13T22:40:43.05825Z","shell.execute_reply":"2022-12-13T22:40:43.076267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%wandb -h 1800 vslaykovsky/rsna-breast-cancer-sweeps/sweeps/k281hlr9","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-13T22:40:43.082086Z","iopub.execute_input":"2022-12-13T22:40:43.084719Z","iopub.status.idle":"2022-12-13T22:40:43.364088Z","shell.execute_reply.started":"2022-12-13T22:40:43.084685Z","shell.execute_reply":"2022-12-13T22:40:43.363176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n    6. Cross-validation\n</div>","metadata":{}},{"cell_type":"code","source":"\ndef gen_predictions(models, df_train):\n    df_train_predictions = []\n    with tqdm(enumerate(models), total=len(models), desc='Folds') as progress:\n        for fold, model in progress:\n            if model is not None:\n                ds_eval = BreastCancerDataSet(df_train.query('split == @fold'), TRAIN_IMAGES_PATH, get_transforms(aug=False))\n\n                cancer_loss, (f1, thres), pred_cancer = evaluate_model(model, ds_eval, PREDICT_MAX_BATCHES)[:3]\n                progress.set_description(f'Eval fold:{fold} pF1:{f1:.02f}')\n                df_pred = pd.DataFrame(data=pred_cancer,\n                                              columns=['cancer_pred_proba'])\n                df_pred['cancer_pred'] = df_pred.cancer_pred_proba > thres\n\n                df = pd.concat(\n                    [df_train.query('split == @fold').reset_index(drop=True), df_pred],\n                    axis=1\n                ).sort_values(['patient_id', 'image_id'])\n                df_train_predictions.append(df)\n    df_train_predictions = pd.concat(df_train_predictions)\n    return df_train_predictions\n\nif CV:\n    models = [load_model(model, MODELS_PATH, BreastCancerModel(AUX_TARGET_NCLASSES))[0] for model in sorted(os.listdir(MODELS_PATH))]\n    df_pred = gen_predictions(models, df_train)\n    df_pred.to_csv('train_predictions.csv', index=False)\n    !head train_predictions.csv","metadata":{"papermill":{"duration":10.196675,"end_time":"2022-12-01T15:10:43.579891","exception":false,"start_time":"2022-12-01T15:10:33.383216","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:40:43.366275Z","iopub.execute_input":"2022-12-13T22:40:43.367189Z","iopub.status.idle":"2022-12-13T22:53:58.644755Z","shell.execute_reply.started":"2022-12-13T22:40:43.367152Z","shell.execute_reply":"2022-12-13T22:53:58.643592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CV:\n    df_pred = pd.read_csv('train_predictions.csv')\n    print('F1 CV score (multiple thresholds):', sklearn.metrics.f1_score(df_pred.cancer, df_pred.cancer_pred))    \n    df_pred = df_pred.groupby(['patient_id', 'laterality']).agg(\n        cancer_max=('cancer_pred_proba', 'max'), cancer_mean=('cancer_pred_proba', 'mean'), cancer=('cancer', 'max')\n    )\n    print('pF1 CV score. Mean aggregation, single threshold:', optimal_f1(df_pred.cancer.values, df_pred.cancer_mean.values))\n    print('pF1 CV score. Max aggregation, single threshold:', optimal_f1(df_pred.cancer.values, df_pred.cancer_max.values))","metadata":{"execution":{"iopub.status.busy":"2022-12-13T22:53:58.646742Z","iopub.execute_input":"2022-12-13T22:53:58.647891Z","iopub.status.idle":"2022-12-13T22:54:09.177839Z","shell.execute_reply.started":"2022-12-13T22:53:58.647842Z","shell.execute_reply":"2022-12-13T22:54:09.175887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IS_KAGGLE:\n    !rm -rf wandb\n    pass","metadata":{"papermill":{"duration":1.058896,"end_time":"2022-12-01T15:31:38.170913","exception":false,"start_time":"2022-12-01T15:31:37.112017","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-13T22:54:09.179188Z","iopub.execute_input":"2022-12-13T22:54:09.179584Z","iopub.status.idle":"2022-12-13T22:54:10.187991Z","shell.execute_reply.started":"2022-12-13T22:54:09.179549Z","shell.execute_reply":"2022-12-13T22:54:10.186589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-danger\" style=\"text-align:center; font-size:20px;\">\n    ❤️ Dont forget to ▲upvote▲ if you find this notebook usefull!  ❤️\n</div>","metadata":{"execution":{"iopub.execute_input":"2022-11-30T19:50:57.290957Z","iopub.status.busy":"2022-11-30T19:50:57.290345Z","iopub.status.idle":"2022-11-30T19:50:57.305056Z","shell.execute_reply":"2022-11-30T19:50:57.302113Z","shell.execute_reply.started":"2022-11-30T19:50:57.290918Z"},"papermill":{"duration":0.011509,"end_time":"2022-12-01T15:31:38.193438","exception":false,"start_time":"2022-12-01T15:31:38.181929","status":"completed"},"tags":[]}}]}