{"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":"- CAUSION: IN ORDER TO SHORTEN THE TRAINGING AND TESTING TIME WE ONLY PROVIDED THE CODE TRAINING ON THE SIZE OF 256*256","metadata":{}},{"cell_type":"markdown","source":"# install some packages","metadata":{}},{"cell_type":"code","source":"!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!pip install /kaggle/input/nvidia-dali-wheel/nvidia_dali_nightly_cuda110-1.22.0.dev20221213-6757685-py3-none-manylinux2014_x86_64.whl\n!pip install /kaggle/input/nvidia-dali-wheel/dicomsdl-0.109.1-cp37-cp37m-manylinux_2_12_x86_64.manylinux2010_x86_64.whl","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-04-09T04:50:52.501963Z","iopub.execute_input":"2023-04-09T04:50:52.502819Z","iopub.status.idle":"2023-04-09T04:51:47.463616Z","shell.execute_reply.started":"2023-04-09T04:50:52.502774Z","shell.execute_reply":"2023-04-09T04:51:47.461796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import packages","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/rsnacode/')\nimport os\n# os.environ[\"CUDA_VISIBLE_DEVICES\"] = '1'\nimport gc\nimport 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, StratifiedGroupKFold\nfrom sklearn.preprocessing import LabelEncoder\nimport sys\nimport timm\nfrom timm import create_model, list_models\nfrom timm.data import create_transform\nfrom torch.cuda.amp import GradScaler, autocast\nfrom tqdm import tqdm\nimport random\nimport wandb\nfrom wandb import AlertLevel\nimport torchvision\nfrom torch.utils.data import Dataset\nfrom torch.nn import functional as F\nfrom torch.nn import Module, Linear, Sequential, ModuleList, ReLU, Dropout, Flatten\nfrom torch import nn\nfrom torch.utils.data import DataLoader\nfrom torch.optim import Adam, SGD, AdamW, lr_scheduler\nimport warnings\nwarnings.filterwarnings('ignore')\n\n\nos.environ['WANDB_API_KEY'] = 'YOUR API KEY' # I set offline.\ngc.collect()\ntorch.cuda.empty_cache()\n\nfrom utils import seed_everything, init_logger, get_timediff, optimal_f1, gc_collect, add_weight_decay, get_parameter_number\nfrom model import GeM, BreastCancerModel\nfrom dataset import BreastCancerDataSet_16bit, BreastCancerDataSet_8bit, mixup_augmentation, get_transforms_8bit, get_transforms_16bit\n","metadata":{"execution":{"iopub.status.busy":"2023-04-09T04:52:18.663952Z","iopub.execute_input":"2023-04-09T04:52:18.664382Z","iopub.status.idle":"2023-04-09T04:52:18.864269Z","shell.execute_reply.started":"2023-04-09T04:52:18.664347Z","shell.execute_reply":"2023-04-09T04:52:18.862926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CFG","metadata":{}},{"cell_type":"code","source":"class CFG:\n    suff = \"0003\"\n    image_size =  (256, 256)\n    epochs = 5\n    model_arch = 'tf_efficientnetv2_s' # tf_efficientnetv2_s / convnextv2_tiny.fcmae_ft_in22k_in1k_384 / convnextv2_tiny\n    dropout = 0.0\n    fc_dropout=0.2\n    es_paitient = 3\n\n    onecycle = True\n    onecycle_pct_start = 0.1\n    max_lr = 1e-6\n    optim = \"AdamW\"\n    weight_decay = 0.01\n    accum_iter=1\n\n    positive_target_weight = 1\n    neg_downsample = 0.35\n    train_batch_size = 8\n    valid_batch_size = 16\n    mixup_rate = 0.5\n    mixup_alpha = 0.5\n\n    tta = True\n    \n    seed = 1788\n    num_workers = 5\n    n_folds = 5\n    folds = [0]\n    gpu_parallel = False\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    df_path = f'/kaggle/input/rsna-breast-cancer-detection/train.csv'\n\n    checkpoint_path = f'./input/checkpoint/0002_effv2s_PL_f4_ep6.pth'\n\n    normalize_mean= [0.485, 0.456, 0.406]  # [0.21596, 0.21596, 0.21596]\n    normalize_std = [0.229, 0.224, 0.225]  # [0.18558, 0.18558, 0.18558]\n\n    wandb_project = f'RSNA2023-V5' \n    wandb_run_name = f'{suff}_{model_arch}'\n\n    target = 'cancer'\n\n\ncomp_data_dir = '/kaggle/input/rsna-breast-cancer-detection'\nimages_dir = f'/kaggle/input/rsna-breast-cancer-256-pngs/'\noutput_dir = f'/output/{CFG.suff}'\nos.makedirs(output_dir, exist_ok=True)\n\nDEBUG = True\nWANDB_SWEEP = False\nTRAIN = True\nCV = True","metadata":{"execution":{"iopub.status.busy":"2023-04-09T05:06:03.947223Z","iopub.execute_input":"2023-04-09T05:06:03.947779Z","iopub.status.idle":"2023-04-09T05:06:03.961287Z","shell.execute_reply.started":"2023-04-09T05:06:03.947714Z","shell.execute_reply":"2023-04-09T05:06:03.95986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(CFG.seed)\n\nLOGGER = init_logger(f'{output_dir}/train_{CFG.suff}.log')","metadata":{"execution":{"iopub.status.busy":"2023-04-09T05:06:04.368753Z","iopub.execute_input":"2023-04-09T05:06:04.369214Z","iopub.status.idle":"2023-04-09T05:06:04.378514Z","shell.execute_reply.started":"2023-04-09T05:06:04.369174Z","shell.execute_reply":"2023-04-09T05:06:04.377002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LOGGER.info(f'run: {CFG.wandb_run_name}; folds:{CFG.folds}')\nLOGGER.info(f\"timm.version: {timm.__version__}\")\nLOGGER.info(f\"checkpoint: {CFG.checkpoint_path.split('/')[-1]}\")\n\n# try:\n#     df_train = pd.read_csv(CFG.df_path)\n# except:\nLOGGER.info(f\"Can't find {CFG.df_path}, creating one\")\ndf_train = pd.read_csv(f'{comp_data_dir}/train.csv')\nsplit = StratifiedGroupKFold(CFG.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[\"sample_rand\"] = np.random.rand(len(df_train)) \ndf_train.loc[df_train[\"cancer\"]==1, \"sample_rand\"] = 0.0\ndf_train.to_csv(f'/df_train_{CFG.suff}.csv', index=False)\n\ndf_train = df_train[df_train['image_id'].astype(str) != '1942326353'].reset_index(drop=True)\ndf_train = df_train[df_train['patient_id'].astype(str) != '27770'].reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T05:07:25.512919Z","iopub.execute_input":"2023-04-09T05:07:25.514175Z","iopub.status.idle":"2023-04-09T05:07:31.162266Z","shell.execute_reply.started":"2023-04-09T05:07:25.514123Z","shell.execute_reply":"2023-04-09T05:07:31.160877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"if DEBUG:\n    ds_train = BreastCancerDataSet_16bit(df_train, images_dir, CFG.target, get_transforms_16bit('train', CFG.image_size, CFG.normalize_mean, CFG.normalize_std))\n    X, y_cancer = ds_train[42]\n    print(f\"X.shape: {X.shape}, y_cancer.shape: {y_cancer.shape}\")\n\n\ndataset_show = BreastCancerDataSet_16bit(df_train, images_dir, CFG.target, get_transforms_16bit('train', CFG.image_size, CFG.normalize_mean, CFG.normalize_std))\nfor i in range(2):\n    f, axarr = plt.subplots(1,3, figsize=(10,8))\n    for p in range(0,3):\n        idx = np.random.randint(0, len(dataset_show))\n        img, cancer_target = dataset_show[idx]\n        img = ((img-img.min())/(img.max()-img.min())*255).to(torch.uint8).transpose(0, 1).transpose(1,2)\n        axarr[p].imshow(img)\n        axarr[p].set_title(str(cancer_target.item()))\n        axarr[p].axis('off')\n        plt.tight_layout()\n    plt.savefig(f'{output_dir}/show_{i}.jpg')\n    # plt.show()\n\n\n# 3. Model\nif DEBUG:\n    with torch.no_grad():\n        model = BreastCancerModel(model_arch=CFG.model_arch, dropout=0.0, fc_dropout=0.0)\n        pred = model(torch.randn(2, 3, 512, 512))\n        print('model output:', pred.shape)\n        LOGGER.info(get_parameter_number(model))\n    del model","metadata":{"execution":{"iopub.status.busy":"2023-04-09T05:10:40.833288Z","iopub.execute_input":"2023-04-09T05:10:40.833868Z","iopub.status.idle":"2023-04-09T05:10:46.024143Z","shell.execute_reply.started":"2023-04-09T05:10:40.833818Z","shell.execute_reply":"2023-04-09T05:10:46.022672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train and Test","metadata":{}},{"cell_type":"code","source":"def valid_one_epoch(model, dataloader):\n    model = model.to(CFG.device)\n    cancer_pred_list = []\n    with torch.no_grad():\n        model.eval()\n        losses = []; targets = []\n        with tqdm(dataloader, desc='Eval', mininterval=30) as progress:\n            for i, (X, y_c) in enumerate(progress):\n                with autocast(enabled=True):\n                    X = X.to(CFG.device)\n                    y_c = y_c.to(float).to(CFG.device)\n                    \n                    pred_c = model(X).view(-1)\n                    if CFG.tta:\n                        pred_c2 = model(torch.flip(X, dims=[-1])) # horizontal mirror\n                        pred_c = (pred_c + pred_c2) / 2\n\n                    loss = F.binary_cross_entropy_with_logits(pred_c, y_c, pos_weight=torch.tensor([CFG.positive_target_weight]).to(CFG.device)).item()\n                    loss = loss / CFG.accum_iter\n                    \n                    cancer_pred_list.append(torch.sigmoid(pred_c))\n                    losses.append(loss); targets.append(y_c.cpu().numpy())\n        \n        targets = np.concatenate(targets)\n        pred = torch.concat(cancer_pred_list).cpu().numpy()\n        pf1, thres = optimal_f1(targets, pred)\n        #      (best_pf1, best_thres)  pred_value  mean_all_losses\n        return (pf1,      thres),      pred,       np.mean(losses)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T05:11:02.096064Z","iopub.execute_input":"2023-04-09T05:11:02.096722Z","iopub.status.idle":"2023-04-09T05:11:02.113416Z","shell.execute_reply.started":"2023-04-09T05:11:02.096661Z","shell.execute_reply":"2023-04-09T05:11:02.111697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_one_epoch(model, dl, optim, scheduler, cancer_criterion, epoch, logger):\n    model.train()\n    scaler = GradScaler()\n    losses = []\n    with tqdm(dl, desc='Train', mininterval=10) as train_progress:\n        for batch_idx, (img, yc) in enumerate(train_progress):\n            img = img.to(CFG.device)\n            yc = yc.to(float).to(CFG.device)\n            \n            # Mixup- allowed\n            if torch.randn(1)[0] < CFG.mixup_rate and img.shape[0]>1:  \n                mixed_x, yc_j, yc_k, lam = mixup_augmentation(img, yc, alpha=CFG.mixup_alpha)\n                with autocast(enabled=True):\n                    pred_c = model(mixed_x).view(-1) \n                    # Mixup loss calculation\n                    loss_j = cancer_criterion(pred_c, yc_j, pos_weight=torch.tensor([CFG.positive_target_weight]).to(CFG.device)) \n                    loss_k = cancer_criterion(pred_c, yc_k, pos_weight=torch.tensor([CFG.positive_target_weight]).to(CFG.device))\n                    loss = lam * loss_j + (1 - lam) * loss_k\n        \n            # Mixup - not allowed\n            else:\n                if img.shape[0] <= 1:\n                    print('batch size is 1, skipping mixup') \n                with autocast(enabled=True):\n                    pred_c = model(img).view(-1) \n                    loss = cancer_criterion(pred_c, yc, pos_weight=torch.tensor([CFG.positive_target_weight]).to(CFG.device)) \n\n            if np.isinf(loss.item()) or np.isnan(loss.item()):\n                print(f'Bad loss, skipping the batch {batch_idx}')\n                del loss, pred_c\n                gc_collect()\n                continue\n            loss = loss / CFG.accum_iter\n            losses.append(loss.item())\n            scaler.scale(loss).backward() # scaler is needed to prevent \"gradient underflow\"\n            if (batch_idx + 1) % CFG.accum_iter == 0:\n                scaler.step(optim)\n                scaler.update()\n                optim.zero_grad()\n\n            if scheduler is not None:\n                scheduler.step()\n            \n            logger.log({'tr_loss': (loss.item()),\n                        'lr': scheduler.get_last_lr()[0] if scheduler else CFG.max_lr,\n                        'epoch': epoch})\n            \n    return model, np.mean(losses[-30:])","metadata":{"execution":{"iopub.status.busy":"2023-04-09T05:11:10.782931Z","iopub.execute_input":"2023-04-09T05:11:10.783487Z","iopub.status.idle":"2023-04-09T05:11:10.806217Z","shell.execute_reply.started":"2023-04-09T05:11:10.783427Z","shell.execute_reply":"2023-04-09T05:11:10.804697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_model(path, model=None):\n    state_dict = torch.load(path, map_location=CFG.device)\n    if model is None:\n        model = BreastCancerModel(state_dict['model_arch'], CFG.dropout, CFG.fc_dropout)\n    model.load_state_dict(state_dict['model'])\n    return model, state_dict['threshold'], state_dict['model_arch']\n\ndef train_loop(logger, fold, do_save_model=True):\n    # ====================================================\n    # data loader\n    # ====================================================\n    tr_df = df_train.query('split != @fold')\n    tr_df = tr_df[tr_df[\"sample_rand\"] <= CFG.neg_downsample].reset_index(drop=True)\n    va_df = df_train.query('split == @fold')\n    LOGGER.info(f\"train: {len(tr_df)}, train pos rate: {tr_df['cancer'].mean():.3f}\")\n    LOGGER.info(f\"valid: {len(va_df)}, valid pos rate: {va_df['cancer'].mean():.3f}\")\n\n    ds_train = BreastCancerDataSet_16bit(tr_df, images_dir, CFG.target, get_transforms_16bit('train', CFG.image_size, CFG.normalize_mean, CFG.normalize_std))\n    ds_valid = BreastCancerDataSet_16bit(va_df, images_dir, CFG.target, get_transforms_16bit('valid', CFG.image_size, CFG.normalize_mean, CFG.normalize_std))\n    dl_train = DataLoader(ds_train, batch_size=CFG.train_batch_size, shuffle=True,  num_workers=CFG.num_workers, pin_memory=True,  drop_last=True)\n    dl_valid = DataLoader(ds_valid, batch_size=CFG.valid_batch_size, shuffle=False, num_workers=CFG.num_workers, pin_memory=False, drop_last=False)\n    \n    # ====================================================\n    # model & optimizer & scheduler & loss\n    # ====================================================\n    # model\n    # model = BreastCancerModel(CFG.model_arch, CFG.dropout, CFG.fc_dropout).to(CFG.device)\n    # if CFG.gpu_parallel:    \n    #     from torch.nn import DataParallel\n    #     num_gpu = torch.cuda.device_count()\n    #     model = DataParallel(model, device_ids=range(num_gpu))\n    #     LOGGER.info(f\"enable gpu parallel, num_gpu: {num_gpu}\")\n    model = load_model(CFG.checkpoint_path)[0]\n    model = model.to(CFG.device)\n\n    # optimizer\n    if CFG.optim == \"AdamW\":\n        optim = AdamW(add_weight_decay(model, weight_decay=CFG.weight_decay, skip_list=['bias']), lr=CFG.max_lr, betas=(0.9, 0.999), weight_decay=CFG.weight_decay)\n    elif CFG.optim == \"Adam\":\n        optim = Adam(model.parameters())\n    \n    # scheduler\n    scheduler = None\n    if CFG.onecycle:\n        scheduler = lr_scheduler.OneCycleLR(optim, max_lr=CFG.max_lr, epochs=CFG.epochs, steps_per_epoch=len(dl_train), pct_start=CFG.onecycle_pct_start)\n    \n    # loss\n    cancer_criterion = F.binary_cross_entropy_with_logits\n\n    # ====================================================\n    # loop\n    # ====================================================    \n    best_valid_score = 0; best_valid_thres=0\n    n_es = 0\n    for epoch in range(CFG.epochs):\n        model, tr_loss = train_one_epoch(model, dl_train, optim, scheduler, cancer_criterion, epoch, logger)\n        (f1, thres), _, va_loss = valid_one_epoch(model, dl_valid)\n\n        n_es += 1\n        if f1 > best_valid_score:\n            n_es = 0\n            best_valid_score = f1\n            best_valid_thres = thres\n            if do_save_model and epoch > 0:\n                save_name = f'{output_dir}/{CFG.suff}_{CFG.model_arch}_FT_f{fold}_ep{epoch}.pth'\n                torch.save({'model': model.state_dict(), 'threshold': thres, 'model_arch': CFG.model_arch}, save_name)\n                best_dict[fold] = save_name\n\n        LOGGER.info(f'Epoch {epoch} - valid_f1: {f1:.4f} - valid_thres: {thres:.4f} - best_f1: {best_valid_score:.4f} - best_thres: {best_valid_thres:.4f}, lr: {scheduler.get_last_lr()[0] if scheduler else CFG.max_lr}')\n        LOGGER.info(f\"train_loss: {tr_loss:.4f} - valid_loss: {va_loss:.4f}\")\n        logger.log({\n            'fold':       fold,\n            'epoch':      epoch,\n            'va_loss':    va_loss,\n            'va_pf1':     f1,\n            'va_thres':   thres,\n            'best_pf1':   best_valid_score,\n            'best_thres': best_valid_thres,\n            })\n\n        if n_es > CFG.es_paitient:\n            LOGGER.info(f'Early Stopping - Epoch: {epoch}')\n            break","metadata":{"execution":{"iopub.status.busy":"2023-04-09T05:11:36.174165Z","iopub.execute_input":"2023-04-09T05:11:36.174683Z","iopub.status.idle":"2023-04-09T05:11:36.196383Z","shell.execute_reply.started":"2023-04-09T05:11:36.174638Z","shell.execute_reply":"2023-04-09T05:11:36.195166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN:\n    best_dict = {}\n    for fold in CFG.folds:\n        LOGGER.info(f'\\n========== Fold {fold} ==========')\n        with wandb.init(project=CFG.wandb_project, name=f'{CFG.wandb_run_name}-f{fold}', group=CFG.wandb_run_name,mode=\"offline\") as run:\n            gc_collect()\n            train_loop(run, fold)\n\ntry:\n    with open(f\"{output_dir}/finished\",\"w\") as f:\n        f.write(\"finish\")\nexcept:\n    LOGGER.info(\"write finished fail.\")","metadata":{"execution":{"iopub.status.busy":"2023-04-09T05:12:10.425305Z","iopub.execute_input":"2023-04-09T05:12:10.425893Z","iopub.status.idle":"2023-04-09T05:12:12.521456Z","shell.execute_reply.started":"2023-04-09T05:12:10.425839Z","shell.execute_reply":"2023-04-09T05:12:12.519713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}