{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":75346,"databundleVersionId":8282674,"sourceType":"competition"},{"sourceId":4421846,"sourceType":"datasetVersion","datasetId":2590074},{"sourceId":4804745,"sourceType":"datasetVersion","datasetId":2779893},{"sourceId":4998596,"sourceType":"datasetVersion","datasetId":2891303},{"sourceId":5877069,"sourceType":"datasetVersion","datasetId":2998419},{"sourceId":8174028,"sourceType":"datasetVersion","datasetId":4838166}],"dockerImageVersionId":30381,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Dependencies","metadata":{}},{"cell_type":"code","source":"# https://stackoverflow.com/questions/46288847/how-to-suppress-pip-upgrade-warning\n!pip config set global.disable-pip-version-check true\n!pip config set global.root-user-action ignore\n\n\nimport os\nos.environ['CUDA_MODULE_LOADING']='LAZY'\n\n!mkdir -p /kaggle/tmp/libs\n\n# # upgrade pytorch to 1.12 for torch_tensorrt\n!pip install /kaggle/input/pytorch112-cu113/{torch-1.12.1+cu113-cp37-cp37m-linux_x86_64.whl,torchvision-0.13.1+cu113-cp37-cp37m-linux_x86_64.whl}\n\n# install timm==0.8.11.dev0\n\n!cp -r /kaggle/input/kaggle-rsna-pkgs/timm /kaggle/tmp/libs\n%cd /kaggle/tmp/libs/timm\n!pip install -e .\n%cd /kaggle/working\n\n# install torch2trt\ntry: \n    import torch2trt\nexcept:\n    !pip install /kaggle/input/torch-tensorrt-pkg/nvidia_pyindex-1.0.9-py3-none-any.whl\n    !mkdir -p /tmp/pip/cache/\n    !cp /kaggle/input/torch-tensorrt-pkg/nvidia-cublas-cu11-2022.4.8.xyz /tmp/pip/cache/nvidia-cublas-cu11-2022.4.8.tar.gz\n    !cp /kaggle/input/torch-tensorrt-pkg/nvidia-cuda-runtime-cu11-2022.4.25.xyz /tmp/pip/cache/nvidia-cuda-runtime-cu11-2022.4.25.tar.gz\n    !cp /kaggle/input/torch-tensorrt-pkg/nvidia-cudnn-cu11-2022.5.19.xyz /tmp/pip/cache/nvidia-cudnn-cu11-2022.5.19.tar.gz\n    !cp /kaggle/input/torch-tensorrt-pkg/nvidia_cublas_cu117-11.10.1.25-py3-none-manylinux1_x86_64.whl /tmp/pip/cache/\n    !cp /kaggle/input/torch-tensorrt-pkg/nvidia_cuda_runtime_cu117-11.7.60-py3-none-manylinux1_x86_64.whl /tmp/pip/cache/\n    !cp /kaggle/input/torch-tensorrt-pkg/nvidia_cudnn_cu116-8.4.0.27-py3-none-manylinux1_x86_64.whl /tmp/pip/cache/\n    !cp /kaggle/input/torch-tensorrt-pkg/nvidia_tensorrt-8.4.3.1-cp37-none-linux_x86_64.whl /tmp/pip/cache/\n    !pip install --no-index --find-links /tmp/pip/cache/ nvidia_tensorrt\n    !pip install /kaggle/input/torch-tensorrt-pkg/torch_tensorrt-1.2.0-cp37-cp37m-linux_x86_64.whl\n    \n    # setup torch2trt\n    !cp -r /kaggle/input/kaggle-rsna-pkgs/torch2trt /kaggle/tmp/libs\n    %cd /kaggle/tmp/libs/torch2trt\n    !python setup.py install\n    !pip install -e .\n#     !cmake -B build . && cmake --build build --target install && ldconfig\n    %cd /kaggle/working/\n\ntry:\n    import dicomsdl\nexcept:\n    !pip install /kaggle/input/kaggle-rsna-pkgs/pylibjpeg-1.4.0-py3-none-any.whl\n    !pip install /kaggle/input/kaggle-rsna-pkgs/python_gdcm-3.0.21-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n    !pip install /kaggle/input/kaggle-rsna-pkgs/dicomsdl-0.109.1-cp37-cp37m-manylinux_2_12_x86_64.manylinux2010_x86_64.whl\ntry:\n    import dali\nexcept:\n    !pip install /kaggle/input/kaggle-rsna-pkgs/nvidia_dali_nightly_cuda110-1.23.0.dev20230210-7260679-py3-none-manylinux2014_x86_64.whl\n\nprint('Import done!')","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"scrolled":true,"execution":{"iopub.status.busy":"2024-04-20T14:03:23.497333Z","iopub.execute_input":"2024-04-20T14:03:23.497837Z","iopub.status.idle":"2024-04-20T14:04:09.074871Z","shell.execute_reply.started":"2024-04-20T14:03:23.497781Z","shell.execute_reply":"2024-04-20T14:04:09.073562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/tmp/libs/timm')\nsys.path.append('/opt/conda/lib/python3.7/site-packages/torch2trt-0.4.0-py3.7.egg')\nsys.path.append('/kaggle/tmp/libs/torch2trt')\nimport timm\nimport gc\nprint('Timm version:', timm.__version__)\n\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")\nimport os\n\nos.environ['CUDA_MODULE_LOADING'] = 'LAZY'\nimport ctypes\nimport gc\nimport importlib\nimport multiprocessing as mp\nimport shutil\n\nimport albumentations as A\nimport cv2\nimport dicomsdl\nimport numpy as np\nimport nvidia.dali as dali\nimport pandas as pd\nimport pydicom\n\nimport torch\nfrom albumentations.pytorch.transforms import ToTensorV2\nfrom joblib import Parallel, delayed\nfrom nvidia.dali import types\nfrom nvidia.dali.backend import TensorGPU, TensorListGPU\nfrom torch2trt import TRTModule\nfrom torch.nn import functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nimport time","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-04-20T14:04:09.077497Z","iopub.execute_input":"2024-04-20T14:04:09.077949Z","iopub.status.idle":"2024-04-20T14:04:09.09016Z","shell.execute_reply.started":"2024-04-20T14:04:09.077899Z","shell.execute_reply":"2024-04-20T14:04:09.08902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nfrom sklearn import metrics\n\n#################################################################################################\n\n\ndef compute_metrics_over_thresholds(preds,\n                                    gts,\n                                    thresholds=np.linspace(0, 1, 101),\n                                    eps=1e-3):\n    f1scores = []\n    precisions = []\n    recalls = []\n    for t in thresholds:\n        predict = (preds > t).astype(np.float32)\n\n        tp = ((predict >= 0.5) & (gts >= 0.5)).sum()\n        fp = ((predict >= 0.5) & (gts < 0.5)).sum()\n        fn = ((predict < 0.5) & (gts >= 0.5)).sum()\n\n        r = tp / (tp + fn + eps)\n        p = tp / (tp + fp + eps)\n        f1 = 2 * r * p / (r + p + eps)\n        f1scores.append(f1)\n        precisions.append(p)\n        recalls.append(r)\n    f1scores = np.array(f1scores)\n    precisions = np.array(precisions)\n    recalls = np.array(recalls)\n    return f1scores, precisions, recalls, thresholds\n\n\ndef compute_best_metrics(cancer_p, cancer_t):\n\n    fpr, tpr, thresholds = metrics.roc_curve(cancer_t, cancer_p)\n    auc = metrics.auc(fpr, tpr)\n\n    f1scores, precisions, recalls, thresholds = compute_metrics_over_thresholds(\n        cancer_p, cancer_t)\n    i = f1scores.argmax()\n    f1score, precision, recall, threshold = f1scores[i], precisions[\n        i], recalls[i], thresholds[i]\n\n    specificity = ((cancer_p < threshold) &\n                   ((cancer_t <= 0.5))).sum() / (cancer_t <= 0.5).sum()\n    sensitivity = ((cancer_p >= threshold) &\n                   ((cancer_t >= 0.5))).sum() / (cancer_t >= 0.5).sum()\n\n    return {\n        'auc': auc,\n        'threshold': threshold,\n        'f1score': f1score,\n        'precision': precision,\n        'recall': recall,\n        'sensitivity': sensitivity,\n        'specificity': specificity,\n    }\n\n\ndef compute_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            #cfp += 1 - prediction\n        else:\n            cfp += prediction\n\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp + 1e-8)\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) / (\n            beta_squared * c_precision + c_recall)\n        return result\n    else:\n        return 0\n\n\ndef print_all_metric(valid_df):\n\n    print(\n        f'{\"    \": <16}    \\tauc      @th     f1      | \tprec    recall  | \tsens    spec '\n    )\n    for site_id in [0, 1, 2]:\n        if site_id > 0:\n            site_df = valid_df[valid_df.site_id == site_id].reset_index(\n                drop=True)\n        else:\n            site_df = valid_df\n        # ---\n\n        gb = site_df\n        m = compute_best_metrics(gb.cancer_p, gb.cancer_t)\n        text = f'{\"single image\": <16} [{site_id}]'\n        text += f'\\t{m[\"auc\"]:0.5f}'\n        text += f'\\t{m[\"threshold\"]:0.5f}'\n        text += f'\\t{m[\"f1score\"]:0.5f} | '\n        text += f'\\t{m[\"precision\"]:0.5f}'\n        text += f'\\t{m[\"recall\"]:0.5f} | '\n        text += f'\\t{m[\"sensitivity\"]:0.5f}'\n        text += f'\\t{m[\"specificity\"]:0.5f}'\n        #text += '\\n'\n        print(text)\n\n        # ---\n\n        gb = site_df[['patient_id', 'laterality', 'cancer_t',\n                      'cancer_p']].groupby(['patient_id',\n                                            'laterality']).mean()\n        m = compute_best_metrics(gb.cancer_p, gb.cancer_t)\n        text = f'{\"grouby mean()\": <16} [{site_id}]'\n        text += f'\\t{m[\"auc\"]:0.5f}'\n        text += f'\\t{m[\"threshold\"]:0.5f}'\n        text += f'\\t{m[\"f1score\"]:0.5f} | '\n        text += f'\\t{m[\"precision\"]:0.5f}'\n        text += f'\\t{m[\"recall\"]:0.5f} | '\n        text += f'\\t{m[\"sensitivity\"]:0.5f}'\n        text += f'\\t{m[\"specificity\"]:0.5f}'\n        #text += '\\n'\n        print(text)\n\n        # ---\n        gb = site_df[['patient_id', 'laterality', 'cancer_t',\n                      'cancer_p']].groupby(['patient_id', 'laterality']).max()\n        m = compute_best_metrics(gb.cancer_p, gb.cancer_t)\n        text = f'{\"grouby max()\": <16} [{site_id}]'\n        text += f'\\t{m[\"auc\"]:0.5f}'\n        text += f'\\t{m[\"threshold\"]:0.5f}'\n        text += f'\\t{m[\"f1score\"]:0.5f} | '\n        text += f'\\t{m[\"precision\"]:0.5f}'\n        text += f'\\t{m[\"recall\"]:0.5f} | '\n        text += f'\\t{m[\"sensitivity\"]:0.5f}'\n        text += f'\\t{m[\"specificity\"]:0.5f}'\n        #text += '\\n'\n        print(text)\n        print(f'--------------\\n')\n\n\ndef compute_all(df, plot_save_path):\n    print(f'Saving plot to {plot_save_path}')\n    df['cancer_p'] = df['preds']\n    df['cancer_t'] = df['targets']\n    print_all_metric(df)\n\n    gb = df[['site_id', 'patient_id', 'laterality', 'cancer_t',\n             'cancer_p']].groupby(['patient_id', 'laterality']).mean()\n    gb.loc[:, 'cancer_t'] = gb.cancer_t.astype(int)\n    m = compute_best_metrics(gb.cancer_p, gb.cancer_t)\n    text = f'{\"grouby mean()\": <16}'\n    text += f'\\t{m[\"auc\"]:0.5f}'\n    text += f'\\t{m[\"threshold\"]:0.5f}'\n    text += f'\\t{m[\"f1score\"]:0.5f} | '\n    text += f'\\t{m[\"precision\"]:0.5f}'\n    text += f'\\t{m[\"recall\"]:0.5f} | '\n    text += f'\\t{m[\"sensitivity\"]:0.5f}'\n    text += f'\\t{m[\"specificity\"]:0.5f}'\n    text += '\\n'\n    print(text)\n\n    pfbeta = compute_pfbeta(gb.cancer_t.values, gb.cancer_p.values, beta=1)\n    print('PROBABILITY-FBETA:', pfbeta)\n\n    plot_pr_curve(gb, plot_save_path)\n\n\ndef plot_pr_curve(df, plot_save_path):\n    f1scores, precisions, recalls, thresholds = compute_metrics_over_thresholds(\n        df.cancer_p, df.cancer_t)\n    i = f1scores.argmax()\n    f1score_max, precision_max, recall_max, threshold_max = f1scores[\n        i], precisions[i], recalls[i], thresholds[i]\n    print(\n        f'f1score_max = {f1score_max}, precision_max = {precision_max}, recall_max = {recall_max}, threshold_max = {threshold_max}'\n    )\n\n    _, axs = plt.subplots(2, 2, figsize=(20, 15))\n\n    ############################################################################\n    ### PRECISION-RECALL CURVE\n    f_scores = [0.2, 0.3, 0.4, 0.5, 0.6, 0.7,\n                0.8]  #np.linspace(0.2, 0.8, num=8)\n    for f_score in f_scores:\n        x = np.linspace(0.01, 1)\n        y = f_score * x / (2 * x - f_score)\n        (l, ) = axs[0, 0].plot(x[y >= 0], y[y >= 0], color=\"gray\", alpha=0.2)\n        axs[0, 0].annotate(\"f1={0:0.1f}\".format(f_score),\n                           xy=(0.9, y[45] + 0.02))\n    axs[0, 0].plot([0, 1], [0, 1], color=\"gray\", alpha=0.2)\n\n    # overall\n    precision, recall, threshold = metrics.precision_recall_curve(\n        df.cancer_t, df.cancer_p)\n    auc = metrics.auc(recall, precision)\n    axs[0, 0].plot(recall, precision)\n    s = axs[0, 0].scatter(recall[:-1], precision[:-1], c=threshold, cmap='hsv')\n    axs[0, 0].scatter(recall_max, precision_max, s=30, c='k')\n\n    # for each site\n    precision, recall, threshold = metrics.precision_recall_curve(\n        df.cancer_t[df.site_id == 1], df.cancer_p[df.site_id == 1])\n    axs[0, 0].plot(recall, precision, '--', label='site_id=1')\n    precision, recall, threshold = metrics.precision_recall_curve(\n        df.cancer_t[df.site_id == 2], df.cancer_p[df.site_id == 2])\n    axs[0, 0].plot(recall, precision, '--', label='site_id=2')\n\n    axs[0, 0].set_xlim([0.0, 1.0])\n    axs[0, 0].set_ylim([0.0, 1.05])\n\n    text = ''\n    text += f'MAX f1score {f1score_max: 0.5f} @ th = {threshold_max: 0.5f}\\n'\n    text += f'prec {precision_max: 0.5f}, recall {recall_max: 0.5f}, pr-auc {auc: 0.5f}\\n'\n\n    axs[0, 0].legend()\n    axs[0, 0].set_title(text)\n    plt.colorbar(s, ax=axs[0, 0], label='threshold')\n    axs[0, 0].set_xlabel('recall')\n    axs[0, 0].set_ylabel('precision')\n\n    ############################################################################\n    # HISTOGRAM\n    spacing = 51\n\n    for site_type in [0, 1, 2]:\n        if site_type == 0:\n            ax = axs[0, 1]\n            sub_df = df\n            title = 'All site'\n        elif site_type == 1:\n            ax = axs[1, 0]\n            sub_df = df[df.site_id == site_type].reset_index(drop=True)\n            title = 'Site 1'\n        elif site_type == 2:\n            ax = axs[1, 1]\n            sub_df = df[df.site_id == site_type].reset_index(drop=True)\n            title = 'Site 2'\n\n        cancer_p = sub_df.cancer_p\n        cancer_t = sub_df.cancer_t\n        cancer_t = cancer_t.astype(int)\n        pos, bin = np.histogram(cancer_p[cancer_t == 1],\n                                np.linspace(0, 1, spacing))\n        neg, bin = np.histogram(cancer_p[cancer_t == 0],\n                                np.linspace(0, 1, spacing))\n        pos = pos / (cancer_t == 1).sum()\n        neg = neg / (cancer_t == 0).sum()\n        # plt.plot(bin[1:],neg, alpha=1)\n        # plt.plot(bin[1:],pos, alpha=1)\n        bin = (bin[1:] + bin[:-1]) / 2\n        ax.bar(bin, neg, width=1 / spacing, label='neg', alpha=0.5)\n        ax.bar(bin, pos, width=1 / spacing, label='pos', alpha=0.5)\n        ax.legend()\n        ax.set_title(title)\n\n    # plt.show()\n    plt.savefig(plot_save_path)","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:09.091864Z","iopub.execute_input":"2024-04-20T14:04:09.092281Z","iopub.status.idle":"2024-04-20T14:04:09.151861Z","shell.execute_reply.started":"2024-04-20T14:04:09.092241Z","shell.execute_reply":"2024-04-20T14:04:09.150649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\n# sys.path.append('/kaggle/tmp/libs/timm')\nimport timm\nprint('timm version:',timm.__version__)\n\nimport os\nimport torch\nimport numpy as np\nimport sklearn\nimport albumentations as A\n\nimport cv2\nfrom albumentations.augmentations.crops import functional as F\nfrom albumentations.pytorch.transforms import ToTensorV2","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:09.154863Z","iopub.execute_input":"2024-04-20T14:04:09.155244Z","iopub.status.idle":"2024-04-20T14:04:09.168671Z","shell.execute_reply.started":"2024-04-20T14:04:09.15521Z","shell.execute_reply":"2024-04-20T14:04:09.167522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from timm.models import create_model\nfrom timm.optim import create_optimizer_v2, optimizer_kwargs\nfrom timm.data import create_dataset, create_loader, resolve_data_config, Mixup, FastCollateMixup, AugMixDataset\nfrom timm import utils\nimport logging\nfrom timm.models import create_model, safe_model_name, resume_checkpoint, load_checkpoint, model_parameters\nfrom timm.loss import JsdCrossEntropy, SoftTargetCrossEntropy, BinaryCrossEntropy, LabelSmoothingCrossEntropy\n# from timm.scheduler import create_scheduler_v2, scheduler_kwargs\nfrom timm.scheduler import create_scheduler, scheduler\nimport torch.nn as nn\n\n_logger = logging.getLogger(__name__)\n\nclass ExpBuild:\n    def __init__(self, args):\n        # synchorize with change in args (alias of same obj)\n        self.args = args\n        self.data_config = None\n\n        # infer num channels\n        in_chans = 3\n        if self.args.in_chans is not None:\n            in_chans = args.in_chans\n        elif self.args.input_size is not None:\n            in_chans = args.input_size[0]\n        self.args.in_chans = in_chans\n\n\n    def build_model(self):\n        model = create_model(\n            self.args.model,\n            pretrained=self.args.pretrained,\n            in_chans=self.args.in_chans,\n            num_classes=self.args.num_classes,\n            drop_rate=self.args.drop,\n            drop_path_rate=self.args.drop_path,\n            drop_block_rate=self.args.drop_block,\n            global_pool=self.args.gp,\n            bn_momentum=self.args.bn_momentum,\n            bn_eps=self.args.bn_eps,\n            scriptable=self.args.torchscript,\n            checkpoint_path=self.args.initial_checkpoint,\n            **self.args.model_kwargs,\n        )\n\n        # if utils.is_primary(self.args):\n        #     _logger.info(\n        #     f'Model {safe_model_name(self.args.model)} created, param count:{sum([m.numel() for m in model.parameters()])}')\n        \n        assert self.data_config is None\n        self.data_config = resolve_data_config(vars(self.args), model=model, verbose=utils.is_primary(self.args))\n\n        return model\n\n\n    def build_train_dataset(self):\n        train_dataset = create_dataset(\n            self.args.dataset,\n            root=self.args.data_dir,\n            split=self.args.train_split,\n            is_training=True,\n            class_map=self.args.class_map,\n            download=self.args.dataset_download,\n            batch_size=self.args.batch_size,\n            seed=self.args.seed,\n            repeats=self.args.epoch_repeats,\n        )\n        return train_dataset\n\n\n    def build_val_dataset(self):\n        val_dataset = create_dataset(\n            self.args.dataset,\n            root=self.args.data_dir,\n            split=self.args.val_split,\n            is_training=False,\n            class_map=self.args.class_map,\n            download=self.args.dataset_download,\n            batch_size=self.args.batch_size,\n        )\n        return val_dataset\n\n\n    def build_train_loader(self, collate_fn = None):\n        train_dataset = self.build_train_dataset()\n\n        # wrap dataset in AugMix helper\n        if self.args.num_aug_splits > 1:\n            train_dataset = AugMixDataset(train_dataset, num_splits=self.args.num_aug_splits)\n\n        # create data loaders w/ augmentation pipeiine\n        train_interpolation = self.args.train_interpolation\n        if self.args.no_aug or not train_interpolation:\n            train_interpolation = self.data_config['interpolation']\n\n        train_loader = create_loader(\n            train_dataset,\n            input_size=self.data_config['input_size'],\n            batch_size=self.args.batch_size,\n            is_training=True,\n            use_prefetcher=self.args.prefetcher,\n            no_aug=self.args.no_aug,\n            re_prob=self.args.reprob,\n            re_mode=self.args.remode,\n            re_count=self.args.recount,\n            re_split=self.args.resplit,\n            scale=self.args.scale,\n            ratio=self.args.ratio,\n            hflip=self.args.hflip,\n            vflip=self.args.vflip,\n            color_jitter=self.args.color_jitter,\n            auto_augment=self.args.aa,\n            num_aug_repeats=self.args.aug_repeats,\n            num_aug_splits=self.args.num_aug_splits,\n            interpolation=train_interpolation,\n            mean=self.data_config['mean'],\n            std=self.data_config['std'],\n            num_workers=self.args.workers,\n            distributed=self.args.distributed,\n            collate_fn=collate_fn,\n            pin_memory=self.args.pin_mem,\n            device=self.args.device,\n            use_multi_epochs_loader=self.args.use_multi_epochs_loader,\n            worker_seeding=self.args.worker_seeding,\n        )\n        return train_loader\n\n\n    def build_val_loader(self):\n        val_dataset = self.build_val_dataset()\n\n        val_workers = self.args.workers\n        if self.args.distributed and ('tfds' in self.args.dataset or 'wds' in self.args.dataset):\n            # FIXME reduces validation padding issues when using TFDS, WDS w/ workers and distributed training\n            val_workers = min(2, self.args.workers)\n        val_loader = create_loader(\n            val_dataset,\n            input_size=self.data_config['input_size'],\n            batch_size=self.args.validation_batch_size or self.args.batch_size,\n            is_training=False,\n            use_prefetcher=self.args.prefetcher,\n            interpolation=self.data_config['interpolation'],\n            mean=self.data_config['mean'],\n            std=self.data_config['std'],\n            num_workers=val_workers,\n            distributed=self.args.distributed,\n            crop_pct=self.data_config['crop_pct'],\n            pin_memory=self.args.pin_mem,\n            device=self.args.device,\n        )\n        return val_loader\n\n\n    def build_train_loss_fn(self):\n        # setup loss function\n        if self.args.jsd_loss:\n            assert self.args.num_aug_splits > 1  # JSD only valid with aug splits set\n            train_loss_fn = JsdCrossEntropy(num_splits=self.args.num_aug_splits, smoothing=self.args.smoothing)\n        elif self.args.mixup_active:\n            # smoothing is handled with mixup target transform which outputs sparse, soft targets\n            if self.args.bce_loss:\n                train_loss_fn = BinaryCrossEntropy(target_threshold=self.args.bce_target_thresh)\n            else:\n                train_loss_fn = SoftTargetCrossEntropy()\n        elif self.args.smoothing:\n            if self.args.bce_loss:\n                train_loss_fn = BinaryCrossEntropy(smoothing=self.args.smoothing, target_threshold=self.args.bce_target_thresh)\n            else:\n                train_loss_fn = LabelSmoothingCrossEntropy(smoothing=self.args.smoothing)\n        else:\n            train_loss_fn = nn.CrossEntropyLoss()\n        train_loss_fn = train_loss_fn.to(device=self.args.device)\n        return train_loss_fn\n        \n\n    def build_val_loss_fn(self):\n        val_loss_fn = nn.CrossEntropyLoss().to(device=self.args.device)\n        return val_loss_fn\n\n\n    def build_optimizer(self, model):\n        optimizer =  create_optimizer_v2(\n        model,\n        **optimizer_kwargs(cfg=self.args),\n        **self.args.opt_kwargs,\n        )\n        return optimizer\n\n    \n    def build_lr_scheduler(self, optimizer):\n        lr_scheduler, num_epochs = create_scheduler_v2(\n        optimizer,\n        **scheduler_kwargs(self.args),\n        updates_per_epoch=self.args.updates_per_epoch,\n        )\n        return lr_scheduler, num_epochs\n","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:09.170601Z","iopub.execute_input":"2024-04-20T14:04:09.170915Z","iopub.status.idle":"2024-04-20T14:04:09.250143Z","shell.execute_reply.started":"2024-04-20T14:04:09.17088Z","shell.execute_reply":"2024-04-20T14:04:09.249122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ValAgumentBaseline:\n    def __init__(self):\n        pass\n    \n    def __call__(self, img):\n        return img\n\n\nTrainAugment = ValAgumentBaseline\n\n\nclass ValTransformBaseline:\n    \n    def __init__(self, input_size, interpolation=cv2.INTER_LINEAR):\n        self.input_size = input_size\n        self.interpolation = interpolation\n        self.max_h, self.max_w = input_size\n        \n        def _fit_resize(image, **kwargs):\n            img_h, img_w = image.shape[:2]\n            r = min(self.max_h / img_h, self.max_w / img_w)\n            new_h, new_w = int(img_h * r), int(img_w * r)\n            new_image = cv2.resize(image, (new_w, new_h),\n                                   interpolation=interpolation)\n            return new_image\n        \n        self.transform_fn = A.Compose([\n            A.Lambda(name=\"FitResize\",\n                     image=_fit_resize,\n                     always_apply=True,\n                     p=1.0),\n            A.PadIfNeeded(min_height=self.max_h,\n                          min_width=self.max_w,\n                          pad_height_divisor=None,\n                          pad_width_divisor=None,\n                          position=A.augmentations.geometric.transforms.\n                          PadIfNeeded.PositionType.CENTER,\n                          border_mode=cv2.BORDER_CONSTANT,\n                          value=0,\n                          mask_value=None,\n                          always_apply=True,\n                          p=1.0),\n            ToTensorV2(transpose_mask=True)\n        ])\n            \n    def __call__(self, img):\n        return self.transform_fn(image=img)['image']\n\nTrainTransform = ValTransformBaseline   \n\nclass ExpBaseline(ExpBuild):\n    \n    def __init__(self, args):\n        super(ExpBaseline, self).__init__(args)\n        \n    def compute_metrics(self,\n                        df,\n                        plot_save_path,\n                        thres_range=(0, 1, 0.01),\n                        sort_by='pfbeta',\n                        additional_info=False):\n        ori_df = df[[\n            'site_id', 'patient_id', 'laterality', 'cancer', 'preds', 'targets'\n        ]]\n        all_metrics = {}\n\n        reducer_single = lambda df: df\n        reducer_gbmean = lambda df: df.groupby(['patient_id', 'laterality']\n                                               ).mean()\n        reducer_gbmax = lambda df: df.groupby(['patient_id', 'laterality']\n                                              ).mean()\n        reducer_gbmean_site1 = lambda df: df[df.site_id == 1].reset_index(\n            drop=True).groupby(['patient_id', 'laterality']).mean()\n        reducer_gbmean_site2 = lambda df: df[df.site_id == 2].reset_index(\n            drop=True).groupby(['patient_id', 'laterality']).mean()\n\n        reducers = {\n            'single': reducer_single,\n            'gbmean': reducer_gbmean,\n            'gbmean_site1': reducer_gbmean_site1,\n            'gbmean_site2': reducer_gbmean_site2,\n            'gbmax': reducer_gbmax,\n        }\n\n        for reducer_name, reducer in reducers.items():\n            df = reducer(ori_df.copy())\n            preds = df['preds'].to_numpy()\n            gts = df['targets'].to_numpy()\n            # mean_sample_weights = mean_df['sample_weights']\n            _metrics = self._compute_metrics(gts, preds, None, thres_range,\n                                             sort_by)\n            all_metrics[f'{reducer_name}_best_thres'] = _metrics['best_thres']\n            all_metrics.update({\n                f'{reducer_name}_best_{k}': v\n                for k, v in _metrics['best_metric'].items()\n            })\n            all_metrics[f'{reducer_name}_pfbeta'] = _metrics['pfbeta']\n            all_metrics[f'{reducer_name}_auc'] = _metrics['auc']\n            all_metrics[f'{reducer_name}_prauc'] = _metrics['prauc']\n\n        # rank 0 only\n        if additional_info:\n            compute_all(ori_df, plot_save_path)\n        return all_metrics\n    \n    def _compute_metrics(self,\n                         gts,\n                         preds,\n                         sample_weights=None,\n                         thres_range=(0, 1, 0.01),\n                         sort_by='pfbeta'):\n        if isinstance(gts, torch.Tensor):\n            gts = gts.cpu().numpy()\n        if isinstance(preds, torch.Tensor):\n            preds = preds.cpu().numpy()\n        assert isinstance(gts, np.ndarray) and isinstance(preds, np.ndarray)\n        assert len(preds) == len(gts)\n\n        # Probabilistic-fbeta\n        pfbeta = pfbeta_np(gts, preds, beta=1.0)\n        # AUC\n        fpr, tpr, _thresholds = sklearn.metrics.roc_curve(gts,\n                                                          preds,\n                                                          pos_label=1)\n        auc = sklearn.metrics.auc(fpr, tpr)\n\n        # PR-AUC\n        precisions, recalls, _thresholds = sklearn.metrics.precision_recall_curve(\n            gts, preds)\n        pr_auc = sklearn.metrics.auc(recalls, precisions)\n\n        ##### METRICS FOR CATEGORICAL PREDICTION #####\n        # PER THRESHOLD METRIC\n        per_thres_metrics = []\n        for thres in np.arange(*thres_range):\n            bin_preds = (preds > thres).astype(np.uint8)\n            metric_at_thres = compute_usual_metrics(gts, bin_preds, beta=1.0)\n            pfbeta_at_thres = pfbeta_np(gts, bin_preds, beta=1.0)\n            metric_at_thres['pfbeta'] = pfbeta_at_thres\n\n            if sample_weights is not None:\n                w_metric_at_thres = compute_usual_metrics(gts,\n                                                          bin_preds,\n                                                          beta=1.0)\n                w_metric_at_thres = {\n                    f'w_{k}': v\n                    for k, v in w_metric_at_thres.items()\n                }\n                metric_at_thres.update(w_metric_at_thres)\n            per_thres_metrics.append((thres, metric_at_thres))\n\n        per_thres_metrics.sort(key=lambda x: x[1][sort_by], reverse=True)\n\n        # handle multiple thresholds with same scores\n        top_score = per_thres_metrics[0][1][sort_by]\n        same_scores = []\n        for j, (thres, metric_at_thres) in enumerate(per_thres_metrics):\n            if metric_at_thres[sort_by] == top_score:\n                same_scores.append(abs(thres - 0.5))\n            else:\n                assert metric_at_thres[sort_by] < top_score\n                break\n        if len(same_scores) == 1:\n            best_thres, best_metric = per_thres_metrics[0]\n        else:\n            # the nearer 0.5 threshold is --> better\n            best_idx = np.argmin(np.array(same_scores))\n            best_thres, best_metric = per_thres_metrics[best_idx]\n\n        # best thres, best results, all results\n        return {\n            'best_thres': best_thres,\n            'best_metric': best_metric,\n            'all_metrics': per_thres_metrics,\n            'pfbeta': pfbeta,\n            'auc': auc,\n            'prauc': pr_auc,\n            # 'pos_log_loss': pos_loss,\n            # 'neg_log_loss': neg_loss,\n            # 'log_loss': total_loss,\n        } ","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:09.251948Z","iopub.execute_input":"2024-04-20T14:04:09.252299Z","iopub.status.idle":"2024-04-20T14:04:09.288372Z","shell.execute_reply.started":"2024-04-20T14:04:09.25227Z","shell.execute_reply":"2024-04-20T14:04:09.287186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nimport random\nfrom typing import Any, Dict, List, Optional, Sequence, Tuple, Union\n\nimport cv2\nimport numpy as np\nfrom albumentations import random_utils\nfrom albumentations.augmentations.crops import functional as F\nfrom albumentations.augmentations.geometric import functional as FGeometric\nfrom albumentations.augmentations.utils import (_maybe_process_in_chunks,\n                                                preserve_shape)\nfrom albumentations.core.transforms_interface import (DualTransform,\n                                                      ImageOnlyTransform)\n\n\nclass _CustomBaseRandomSizedCropNoResize(DualTransform):\n    # Base class for RandomSizedCrop and RandomResizedCrop\n\n    def __init__(self, always_apply=False, p=1.0):\n        super(_CustomBaseRandomSizedCropNoResize,\n              self).__init__(always_apply, p)\n\n    def apply(self,\n              img,\n              crop_height=0,\n              crop_width=0,\n              h_start=0,\n              w_start=0,\n              interpolation=cv2.INTER_LINEAR,\n              **params):\n        return F.random_crop(img, crop_height, crop_width, h_start, w_start)\n\n    def apply_to_bbox(self,\n                      bbox,\n                      crop_height=0,\n                      crop_width=0,\n                      h_start=0,\n                      w_start=0,\n                      rows=0,\n                      cols=0,\n                      **params):\n        return F.bbox_random_crop(bbox, crop_height, crop_width, h_start,\n                                  w_start, rows, cols)\n\n    def apply_to_keypoint(self,\n                          keypoint,\n                          crop_height=0,\n                          crop_width=0,\n                          h_start=0,\n                          w_start=0,\n                          rows=0,\n                          cols=0,\n                          **params):\n        keypoint = F.keypoint_random_crop(keypoint, crop_height, crop_width,\n                                          h_start, w_start, rows, cols)\n        scale_x = self.width / crop_width\n        scale_y = self.height / crop_height\n        keypoint = FGeometric.keypoint_scale(keypoint, scale_x, scale_y)\n        return keypoint\n\n\nclass CustomRandomSizedCropNoResize(_CustomBaseRandomSizedCropNoResize):\n    \"\"\"Torchvision's variant of crop a random part of the input and rescale it to some size.\n\n    Args:\n        scale ((float, float)): range of size of the origin size cropped\n        ratio ((float, float)): range of aspect ratio of the origin aspect ratio cropped\n        interpolation (OpenCV flag): flag that is used to specify the interpolation algorithm. Should be one of:\n            cv2.INTER_NEAREST, cv2.INTER_LINEAR, cv2.INTER_CUBIC, cv2.INTER_AREA, cv2.INTER_LANCZOS4.\n            Default: cv2.INTER_LINEAR.\n        p (float): probability of applying the transform. Default: 1.\n\n    Targets:\n        image, mask, bboxes, keypoints\n\n    Image types:\n        uint8, float32\n    \"\"\"\n\n    def __init__(\n            self,\n            scale=(0.08, 1.0),\n            ratio=(0.75, 1.3333333333333333),\n            always_apply=False,\n            p=1.0,\n    ):\n\n        super(CustomRandomSizedCropNoResize,\n              self).__init__(always_apply=always_apply, p=p)\n        self.scale = scale\n        self.ratio = ratio\n\n    def get_params_dependent_on_targets(self, params):\n        img = params[\"image\"]\n        area = img.shape[0] * img.shape[1]\n\n        for _attempt in range(10):\n            target_area = random.uniform(*self.scale) * area\n            log_ratio = (math.log(self.ratio[0]), math.log(self.ratio[1]))\n            aspect_ratio = math.exp(random.uniform(*log_ratio))\n\n            w = int(round(math.sqrt(target_area *\n                                    aspect_ratio)))  # skipcq: PTC-W0028\n            h = int(round(math.sqrt(target_area /\n                                    aspect_ratio)))  # skipcq: PTC-W0028\n\n            if 0 < w <= img.shape[1] and 0 < h <= img.shape[0]:\n                i = random.randint(0, img.shape[0] - h)\n                j = random.randint(0, img.shape[1] - w)\n                return {\n                    \"crop_height\": h,\n                    \"crop_width\": w,\n                    \"h_start\": i * 1.0 / (img.shape[0] - h + 1e-10),\n                    \"w_start\": j * 1.0 / (img.shape[1] - w + 1e-10),\n                }\n\n        # Fallback to central crop\n        in_ratio = img.shape[1] / img.shape[0]\n        if in_ratio < min(self.ratio):\n            w = img.shape[1]\n            h = int(round(w / min(self.ratio)))\n        elif in_ratio > max(self.ratio):\n            h = img.shape[0]\n            w = int(round(h * max(self.ratio)))\n        else:  # whole image\n            w = img.shape[1]\n            h = img.shape[0]\n        i = (img.shape[0] - h) // 2\n        j = (img.shape[1] - w) // 2\n        return {\n            \"crop_height\": h,\n            \"crop_width\": w,\n            \"h_start\": i * 1.0 / (img.shape[0] - h + 1e-10),\n            \"w_start\": j * 1.0 / (img.shape[1] - w + 1e-10),\n        }\n\n    def get_params(self):\n        return {}\n\n    @property\n    def targets_as_params(self):\n        return [\"image\"]\n\n    def get_transform_init_args_names(self):\n        return \"scale\", \"ratio\"\n\n\n# @TODO: support other native dtype: float, uint16,.. or support higher bit-depth (current 8 bits, LUT size = 256)\n# cv2.LUT is tricky and other implementation with higher bit-depth can cause performance (speed) drop\n# https://stackoverflow.com/questions/27098831/lookup-table-for-16-bit-mat-efficient-way\n# https://stackoverflow.com/questions/71734861/opencv-python-lut-for-16bit-image\n# https://stackoverflow.com/questions/28423701/efficient-way-to-loop-through-pixels-of-16-bit-mat-in-opencv\n# https://answers.opencv.org/question/206755/lut-for-16bit-image/\n@preserve_shape\ndef move_tone_curve_allow_float_img(img, low_y, high_y):\n    \"\"\"Rescales the relationship between bright and dark areas of the image by manipulating its tone curve.\n    Args:\n        img (numpy.ndarray): RGB or grayscale image.\n        low_y (float): y-position of a Bezier control point used\n            to adjust the tone curve, must be in range [0, 1]\n        high_y (float): y-position of a Bezier control point used\n            to adjust image tone curve, must be in range [0, 1]\n    \"\"\"\n    input_dtype = img.dtype\n\n    if low_y < 0 or low_y > 1:\n        raise ValueError(\"low_shift must be in range [0, 1]\")\n    if high_y < 0 or high_y > 1:\n        raise ValueError(\"high_shift must be in range [0, 1]\")\n\n    if input_dtype != np.uint8:\n        # raise ValueError(\"Unsupported image type {}\".format(input_dtype))\n        assert input_dtype == np.float32\n        img = (img * 255).astype(np.uint8)\n\n    t = np.linspace(0.0, 1.0, 256)\n\n    # Defines responze of a four-point bezier curve\n    def evaluate_bez(t):\n        return 3 * (1 - t)**2 * t * low_y + 3 * (1 - t) * t**2 * high_y + t**3\n\n    evaluate_bez = np.vectorize(evaluate_bez)\n    remapping = np.rint(evaluate_bez(t) * 255).astype(np.uint8)\n\n    lut_fn = _maybe_process_in_chunks(cv2.LUT, lut=remapping)\n    img = lut_fn(img)\n    # convert back to float image in range [0, 1]\n    if input_dtype != np.uint8:\n        img = (img / 255).astype(np.float32)\n    return img\n\n\nclass RandomToneCurveAllowFloatImage(ImageOnlyTransform):\n    \"\"\"Randomly change the relationship between bright and dark areas of the image by manipulating its tone curve.\n    Args:\n        scale (float): standard deviation of the normal distribution.\n            Used to sample random distances to move two control points that modify the image's curve.\n            Values should be in range [0, 1]. Default: 0.1\n    Targets:\n        image\n    Image types:\n        uint8, float32\n    Note: if image in float32 dtype, first convert it to uint8, modify tone curve then convert back to float32.\n    \"\"\"\n\n    def __init__(\n        self,\n        scale=0.1,\n        always_apply=False,\n        p=0.5,\n    ):\n        super(RandomToneCurveAllowFloatImage, self).__init__(always_apply, p)\n        self.scale = scale\n\n    def apply(self, image, low_y, high_y, **params):\n        return move_tone_curve_allow_float_img(image, low_y, high_y)\n\n    def get_params(self):\n        return {\n            \"low_y\":\n            np.clip(random_utils.normal(loc=0.25, scale=self.scale), 0, 1),\n            \"high_y\":\n            np.clip(random_utils.normal(loc=0.75, scale=self.scale), 0, 1),\n        }\n\n    def get_transform_init_args_names(self):\n        return (\"scale\", )\n\n    ","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:09.28997Z","iopub.execute_input":"2024-04-20T14:04:09.290313Z","iopub.status.idle":"2024-04-20T14:04:09.333551Z","shell.execute_reply.started":"2024-04-20T14:04:09.290283Z","shell.execute_reply":"2024-04-20T14:04:09.332604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nimport torch\nimport numpy as np\nfrom torch.utils.data import Sampler\nimport torch.distributed as dist\n\n\nclass BalanceSampler(Sampler):\n\n    def __init__(self, dataset, ratio=8):\n        self.r = ratio-1\n        self.dataset = dataset\n        labels = dataset.get_labels()\n        self.pos_index = np.where(labels>0)[0]\n        self.neg_index = np.where(labels==0)[0]\n        print('Num pos:', len(self.pos_index))\n        print('Num neg:', len(self.neg_index))\n\n        self.neg_length = self.r*int(np.floor(len(self.neg_index)/self.r))\n        self.len = self.neg_length + self.neg_length // self.r\n\n    def __iter__(self):\n        pos_index = self.pos_index.copy()\n        neg_index = self.neg_index.copy()\n        np.random.shuffle(pos_index)\n        np.random.shuffle(neg_index)\n\n        neg_index = neg_index[:self.neg_length].reshape(-1,self.r)\n        pos_index = np.random.choice(pos_index, self.neg_length//self.r).reshape(-1,1)\n\n        index = np.concatenate([pos_index,neg_index],-1).reshape(-1)\n        return iter(index)\n\n    def __len__(self):\n        return self.len\n\n\n\nclass BalanceSamplerV2(Sampler):\n    def __init__(self, dataset, batch_size, num_sched_epochs, num_epochs, start_ratio = 1/4, end_ratio = 1/8, one_pos_mode = True, seed = 42):\n        np.random.seed(seed)\n        self.dataset = dataset\n        self.batch_size = batch_size\n        self.start_ratio = start_ratio\n        self.end_ratio = end_ratio\n        self.num_sched_epochs = num_sched_epochs\n        self.num_epochs = num_epochs\n        labels = dataset.get_labels()\n        \n        self.pos_idxs = np.where(labels>0)[0]\n        self.neg_idxs = np.where(labels==0)[0]\n        self.num_pos = len(self.pos_idxs)\n        self.num_neg = len(self.neg_idxs)\n        print(f'Num pos: {self.num_pos}, Num neg: {self.num_neg}')\n\n        # percentage of pos samples per epoch\n        if start_ratio != end_ratio:\n            self.ratios = np.linspace(start_ratio, end_ratio, num_sched_epochs)\n            self.ratios = np.concatenate([self.ratios, np.full((num_epochs - num_sched_epochs,), end_ratio)])\n        else:\n            self.ratios = np.full((num_epochs,), start_ratio)\n        print('RATIO PER EPOCHS:', self.ratios)\n        assert len(self.ratios) == num_epochs\n\n        # pre-compute sampler indexs per epoch\n        self.pre_compute_epoch_idxs = []\n        for i, ratio in enumerate(self.ratios):\n            epoch_idxs = self._pre_compute_epoch_idxs(ratio, one_pos_mode= one_pos_mode)\n            self.pre_compute_epoch_idxs.append(epoch_idxs)\n        self.cur_epoch = None\n\n        self.set_epoch(0)\n\n\n    def set_epoch(self, ep):\n        assert ep < self.num_epochs, f'WARNING: Invalid set_epoch() with ep={ep} while max_epoch={self.num_epochs}'\n        print(f'Set epoch to {ep} with sampler ratio = {self.ratios[ep]}')\n        self.cur_epoch = ep\n        self.len = len(self.pre_compute_epoch_idxs[ep])\n\n    def _pre_compute_epoch_idxs(self, ratio, one_pos_mode = True):\n        print(f'Pre-compute for ratio = {ratio}')\n        epoch_num_pos = int(ratio * self.num_neg)\n        epoch_num_iters = (epoch_num_pos + self.num_neg) // self.batch_size\n        epoch_num_total = epoch_num_iters * self.batch_size\n        # never downsampling neg samples if possible\n        epoch_num_pos = epoch_num_total - self.num_neg\n        \n        # sampling pos idxs\n        min_pos_per_iter = epoch_num_pos // epoch_num_iters\n        if min_pos_per_iter < 1:\n            if one_pos_mode:\n                min_pos_per_iter = 1\n                epoch_num_pos = epoch_num_iters\n                print(f'ONE POS MODE: Switch num_pos_samples to {epoch_num_pos}')\n            else:\n                print(f\"WARNING: At least one batch which has no positive sample: {epoch_num_pos}, {epoch_num_iters}.\\n\")\n            \n        ret_idxs = []\n        pool_pos_idxs = []\n        _count = 0\n        while _count < epoch_num_pos:\n            temp_pos_idxs = self.pos_idxs.copy()\n            np.random.shuffle(temp_pos_idxs)\n            pool_pos_idxs.append(temp_pos_idxs)\n            _count += len(temp_pos_idxs)\n        pool_pos_idxs = np.concatenate(pool_pos_idxs, axis = 0)\n        assert len(pool_pos_idxs) >= epoch_num_pos\n\n        _start = 0\n        _end = 0\n        for i in range(epoch_num_iters):\n            _start = i * min_pos_per_iter\n            _end = (i+1) * min_pos_per_iter\n            ret_idxs.append(pool_pos_idxs[_start:_end].tolist())\n        num_pos_remain = epoch_num_pos - _end\n        assert num_pos_remain == epoch_num_pos % epoch_num_iters\n        pool_remain_pos_idxs = pool_pos_idxs[_end:epoch_num_pos]\n        assert len(pool_remain_pos_idxs) == num_pos_remain\n        for i, j in enumerate(np.random.choice(np.arange(0, epoch_num_iters, 1), num_pos_remain, replace = False)):\n            ret_idxs[j].append(pool_remain_pos_idxs[i])\n            \n        # sampling neg idxs\n        pool_neg_idxs = self.neg_idxs.copy()\n        np.random.shuffle(pool_neg_idxs)\n\n        _cur = 0\n        for i in range(epoch_num_iters):\n            iter_idxs = ret_idxs[i]\n            assert len(iter_idxs) - min_pos_per_iter <= 1\n            _end = _cur + self.batch_size - len(iter_idxs)\n            iter_idxs.extend(pool_neg_idxs[_cur: _end].tolist())\n            _cur = _end\n        if not one_pos_mode:\n            assert _cur == len(pool_neg_idxs)\n\n        ret_idxs = np.array(ret_idxs)\n        assert ret_idxs.shape[0] == epoch_num_iters and ret_idxs.shape[1] == self.batch_size\n        ret_idxs = ret_idxs.reshape(-1)\n        return ret_idxs\n\n\n    def __iter__(self):\n        print(f'STARTING COMPUTE EPOCH {self.cur_epoch} SAMPLE INDEXS...')\n        print('Current epoch sampler ratio:', self.ratios[self.cur_epoch])\n        epoch_idxs = self.pre_compute_epoch_idxs[self.cur_epoch]\n        print(f'{len(epoch_idxs) // self.batch_size} iters with {len(epoch_idxs)} samples')\n        return iter(epoch_idxs)\n\n\n    def __len__(self):\n        return self.len\n","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:09.335153Z","iopub.execute_input":"2024-04-20T14:04:09.335584Z","iopub.status.idle":"2024-04-20T14:04:09.371139Z","shell.execute_reply.started":"2024-04-20T14:04:09.33554Z","shell.execute_reply":"2024-04-20T14:04:09.369942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\n\ncv2.setNumThreads(0)\ncv2.ocl.setUseOpenCL(False)\n\n\nclass RSNADataset(Dataset):\n\n    def __init__(self,\n                 datasets,\n                 augment_fn=None,\n                 transform_fn=None,\n                 n_channels=3,\n                 subset='train'):\n        assert subset in ['train', 'val']\n        self.subset = subset\n        self.img_paths = []\n        self.labels = []\n        self.augment_fn = augment_fn\n        self.transform_fn = transform_fn\n        self.n_channels = n_channels\n\n        print('----------------------------')\n        print(subset)\n        print(datasets)\n        print('----------------------------')\n\n        for data_name, data_info in datasets:\n            print('DATANAME:', data_name)\n            data_csv_path = data_info['csv_path']\n            data_img_dir = data_info['img_dir']\n            df = pd.read_csv(data_csv_path)\n            self.df = df\n            for i in tqdm(range(len(df))):\n                patient_id = df.at[i, 'patient_id']\n                image_id = df.at[i, 'image_id']\n                img_name = f'{patient_id}@{image_id}.png'\n                img_path = os.path.join(data_img_dir, img_name)\n                if i == 0:\n                    _tmp_img = cv2.imread(img_path)\n                    assert _tmp_img is not None\n                    del _tmp_img\n                label = df.at[i, 'cancer']\n                self.img_paths.append(img_path)\n                self.labels.append(label)\n            print(f'Done loading {data_name} with {len(df)} samples.')\n        print(\n            f'DATASET TOTAL LENGTH: {len(self.labels)} with positive percent = {sum(self.labels) / len(self.labels)}'\n        )\n\n    def __len__(self):\n        return len(self.img_paths)\n\n    def __getitem__(self, idx):\n        img_path = self.img_paths[idx]\n        label = self.labels[idx]\n        if self.n_channels == 3:\n            img = cv2.imread(img_path)\n        elif self.n_channels == 1:\n            img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n        else:\n            raise AssertionError()\n        if self.augment_fn:\n            img = self.augment_fn(img)\n        if self.transform_fn:\n            img = self.transform_fn(img)\n        return img, label\n\n    def get_sampler_weights(self, pos_neg_ratio):\n        assert pos_neg_ratio > 0\n        labels = np.array(self.labels)\n        num_pos = labels.sum()\n        num_neg = len(labels) - num_pos\n        ori_pos_neg_ratio = num_pos / num_neg\n        pos_weight = pos_neg_ratio / ori_pos_neg_ratio\n        print('Original pos/neg ratio:', ori_pos_neg_ratio)\n        print('Expect pos/neg ratio:', pos_neg_ratio)\n        print('Pos weight:', pos_weight)\n        weights = np.ones_like(labels, dtype=np.float32)\n        weights[labels == 1] = pos_weight\n        return weights\n\n    def get_labels(self):\n        return np.array(self.labels)\n\n    def get_df(self):\n        assert self.subset == 'val'\n        return self.df","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:09.372798Z","iopub.execute_input":"2024-04-20T14:04:09.373214Z","iopub.status.idle":"2024-04-20T14:04:09.397749Z","shell.execute_reply.started":"2024-04-20T14:04:09.373176Z","shell.execute_reply":"2024-04-20T14:04:09.396837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport sklearn\nfrom sklearn.metrics import (auc, confusion_matrix,\n                             precision_recall_fscore_support, roc_curve)\n\n\ndef _compute_fbeta(precision, recall, beta=1.0):\n    return (1 + beta**2) * precision * recall / (\n        (beta**2) * precision + recall)\n\n\ndef compute_usual_metrics(gts, preds, beta=1.0, sample_weights=None):\n    \"\"\"Binary prediction only.\"\"\"\n    cfm = confusion_matrix(gts,\n                           preds,\n                           labels=[0, 1],\n                           sample_weight=sample_weights)\n\n    tn, fp, fn, tp = cfm.ravel()\n    acc = (tp + tn) / (tn + fp + fn + tp)\n    precision = tp / (tp + fp)\n    recall = tp / (tp + fn)\n    fbeta = _compute_fbeta(precision, recall, beta=beta)\n    # frr = fp / (fp + tn)\n    # far = fn / (fn + tp)  # 1 - recall\n    # bacc_beta = _compute_fbeta(1 - frr, 1 - far, beta=beta)\n    return {\n        'acc': acc,\n        'precision': precision,\n        'recall': recall,\n        'fbeta': fbeta,\n        # 'bacc_beta': bacc_beta,\n        # 'frr': frr,\n        # 'far': far,\n    }","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:09.40207Z","iopub.execute_input":"2024-04-20T14:04:09.402455Z","iopub.status.idle":"2024-04-20T14:04:09.412307Z","shell.execute_reply.started":"2024-04-20T14:04:09.402425Z","shell.execute_reply":"2024-04-20T14:04:09.411372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\n\ntry:\n    import tensorflow as tf\nexcept:\n    pass\n\ntry:\n    from numba import jit, njit\nexcept:\n    pass\n\n\ndef pfbeta_py(gts, preds, beta=1):\n    y_true_count = 0\n    ctp = 0\n    cfp = 0\n\n    for idx in range(len(gts)):\n        prediction = min(max(preds[idx], 0), 1)\n        if (gts[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 / y_true_count\n    if (c_precision > 0 and c_recall > 0):\n        ret = (1 + beta_squared) * (c_precision * c_recall) / (\n            beta_squared * c_precision + c_recall)\n        return ret\n    else:\n        return 0\n\n\ndef pfbeta_np(gts, preds, beta=1):\n    preds = preds.clip(0, 1.)\n    y_true_count = gts.sum()\n    ctp = preds[gts == 1].sum()\n    cfp = preds[gts == 0].sum()\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        ret = (1 + beta_squared) * (c_precision * c_recall) / (\n            beta_squared * c_precision + c_recall)\n        return ret\n    else:\n        return 0.0\n\n\n# First run will compile --> subsequent calls will be faster\n# @jit(nopython=True)\n@njit\ndef pfbeta_numba(gts, preds, beta=1):\n    preds = preds.clip(0, 1.)\n    y_true_count = gts.sum()\n    ctp = preds[gts == 1].sum()\n    cfp = preds[gts == 0].sum()\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        ret = (1 + beta_squared) * (c_precision * c_recall) / (\n            beta_squared * c_precision + c_recall)\n        return ret\n    else:\n        return 0.0\n\n\ndef pfbeta_tf(gts, preds, beta=1):\n    preds = tf.clip_by_value(preds, 0, 1)\n    y_true_count = tf.reduce_sum(gts)\n    ctp = tf.reduce_sum(preds[gts == 1])\n    cfp = tf.reduce_sum(preds[gts == 0])\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        ret = (1 + beta_squared) * (c_precision * c_recall) / (\n            beta_squared * c_precision + c_recall)\n        return ret\n    else:\n        return 0.0\n\n\ndef pfbeta_torch(gts, preds, beta=1):\n    preds = preds.clip(0, 1)\n    y_true_count = gts.sum()\n    ctp = preds[gts == 1].sum()\n    cfp = preds[gts == 0].sum()\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        ret = (1 + beta_squared) * (c_precision * c_recall) / (\n            beta_squared * c_precision + c_recall)\n        return ret\n    else:\n        return 0.0","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:09.414236Z","iopub.execute_input":"2024-04-20T14:04:09.414564Z","iopub.status.idle":"2024-04-20T14:04:13.70932Z","shell.execute_reply.started":"2024-04-20T14:04:09.414535Z","shell.execute_reply":"2024-04-20T14:04:13.708404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Optional\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass BinaryCrossEntropyPosSmoothOnly(nn.Module):\n    \"\"\" BCE with optional one-hot from dense targets, label smoothing, thresholding\n    NOTE for experiments comparing CE to BCE /w label smoothing, may remove\n    \"\"\"\n    def __init__(\n            self, smoothing=0.1, target_threshold: Optional[float] = None, weight: Optional[torch.Tensor] = None,\n            reduction: str = 'mean', pos_weight: Optional[torch.Tensor] = None):\n        super(BinaryCrossEntropyPosSmoothOnly, self).__init__()\n        assert 0. <= smoothing < 1.0\n        self.smoothing = smoothing\n        self.target_threshold = target_threshold\n        self.reduction = reduction\n        self.register_buffer('weight', weight)\n        self.register_buffer('pos_weight', pos_weight)\n\n    def forward(self, x: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        assert x.shape[0] == target.shape[0]\n        if target.shape != x.shape:\n            # NOTE currently assume smoothing or other label softening is applied upstream if targets are already sparse\n            num_classes = x.shape[-1]\n            # FIXME should off/on be different for smoothing w/ BCE? Other impl out there differ\n            assert num_classes <= 2\n            if num_classes > 1:\n                on_value = 1. - self.smoothing\n                target = target.float().view(-1, 1).repeat(1, 2)\n                target[:, 1] = target[:, 1] * on_value\n                target[:, 0] = 1.0 - target[:, 1]\n            elif num_classes == 1:\n                off_value = 0\n                on_value = 1. - self.smoothing\n                target = torch.where(target.bool(), on_value, off_value).view(-1, 1)\n            else:\n                raise AssertionError()\n\n        if self.target_threshold is not None:\n            # Make target 0, or 1 if threshold set\n            target = target.gt(self.target_threshold).to(dtype=target.dtype)\n\n        return F.binary_cross_entropy_with_logits(\n            x, target,\n            self.weight,\n            pos_weight=self.pos_weight,\n            reduction=self.reduction)","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:13.710942Z","iopub.execute_input":"2024-04-20T14:04:13.711265Z","iopub.status.idle":"2024-04-20T14:04:13.726145Z","shell.execute_reply.started":"2024-04-20T14:04:13.711234Z","shell.execute_reply":"2024-04-20T14:04:13.725158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import logging\nimport math\nimport os\n\nimport albumentations as A\nimport cv2\nimport numpy as np\nimport sklearn\nimport torch\nimport torch.distributed as dist\nimport torch.nn as nn\nfrom albumentations.pytorch.transforms import ToTensorV2\nfrom ignite.distributed import DistributedProxySampler\n\nfrom timm.data import (AugMixDataset, FastCollateMixup, Mixup, create_dataset,\n                       create_loader, resolve_data_config)\nfrom timm.loss import (BinaryCrossEntropy,\n                       JsdCrossEntropy, LabelSmoothingCrossEntropy,\n                       SoftTargetCrossEntropy)\n\nfrom timm.models import (create_model, load_checkpoint, model_parameters,\n                         resume_checkpoint, safe_model_name)\nfrom timm.optim import create_optimizer_v2, optimizer_kwargs\nfrom timm.scheduler import create_scheduler_v2, scheduler_kwargs\nfrom torch.utils.data import Sampler, WeightedRandomSampler\n\n\nclass ValAugment:\n\n    def __init__(self):\n        pass\n\n    def __call__(self, img):\n        return img\n\n\nclass TrainAugment:\n\n    def __init__(self):\n        self.transform_fn = A.Compose(\n            [\n                # crop\n                CustomRandomSizedCropNoResize(scale=(0.5, 1.0),\n                                                   ratio=(0.5, 0.8),\n                                                   always_apply=False,\n                                                   p=0.4),\n\n                # flip\n                A.HorizontalFlip(p=0.5),\n                A.VerticalFlip(p=0.5),\n\n                # downscale\n                A.OneOf([\n                    A.Downscale(scale_min=0.75,\n                                scale_max=0.95,\n                                interpolation=dict(upscale=cv2.INTER_LINEAR,\n                                                   downscale=cv2.INTER_AREA),\n                                always_apply=False,\n                                p=0.1),\n                    A.Downscale(scale_min=0.75,\n                                scale_max=0.95,\n                                interpolation=dict(upscale=cv2.INTER_LANCZOS4,\n                                                   downscale=cv2.INTER_AREA),\n                                always_apply=False,\n                                p=0.1),\n                    A.Downscale(scale_min=0.75,\n                                scale_max=0.95,\n                                interpolation=dict(upscale=cv2.INTER_LINEAR,\n                                                   downscale=cv2.INTER_LINEAR),\n                                always_apply=False,\n                                p=0.8),\n                ],\n                        p=0.125),\n\n                # contrast\n                # relative dark/bright between region, like HDR\n                A.OneOf([\n                    A.RandomToneCurve(scale=0.3, always_apply=False, p=0.5),\n                    A.RandomBrightnessContrast(brightness_limit=(-0.1, 0.2),\n                                               contrast_limit=(-0.4, 0.5),\n                                               brightness_by_max=True,\n                                               always_apply=False,\n                                               p=0.5)\n                ],\n                        p=0.5),\n\n                # affine\n                A.OneOf(\n                    [\n                        A.ShiftScaleRotate(shift_limit=None,\n                                           scale_limit=[-0.15, 0.15],\n                                           rotate_limit=[-30, 30],\n                                           interpolation=cv2.INTER_LINEAR,\n                                           border_mode=cv2.BORDER_CONSTANT,\n                                           value=0,\n                                           mask_value=None,\n                                           shift_limit_x=[-0.1, 0.1],\n                                           shift_limit_y=[-0.2, 0.2],\n                                           rotate_method='largest_box',\n                                           always_apply=False,\n                                           p=0.6),\n\n                        # one of with other affine\n                        A.ElasticTransform(alpha=1,\n                                           sigma=20,\n                                           alpha_affine=10,\n                                           interpolation=cv2.INTER_LINEAR,\n                                           border_mode=cv2.BORDER_CONSTANT,\n                                           value=0,\n                                           mask_value=None,\n                                           approximate=False,\n                                           same_dxdy=False,\n                                           always_apply=False,\n                                           p=0.2),\n\n                        # distort\n                        A.GridDistortion(num_steps=5,\n                                         distort_limit=0.3,\n                                         interpolation=cv2.INTER_LINEAR,\n                                         border_mode=cv2.BORDER_CONSTANT,\n                                         value=0,\n                                         mask_value=None,\n                                         normalized=True,\n                                         always_apply=False,\n                                         p=0.2),\n                    ],\n                    p=0.5),\n\n                # random erase\n                A.CoarseDropout(max_holes=6,\n                                max_height=0.15,\n                                max_width=0.25,\n                                min_holes=1,\n                                min_height=0.05,\n                                min_width=0.1,\n                                fill_value=0,\n                                mask_fill_value=None,\n                                always_apply=False,\n                                p=0.25),\n            ],\n            p=0.9)\n\n        print('TRAIN AUG:\\n', self.transform_fn)\n\n    def __call__(self, img):\n        return self.transform_fn(image=img)['image']\n\n\nclass ValTransform:\n\n    def __init__(self, input_size, interpolation=cv2.INTER_LINEAR):\n        self.input_size = input_size\n        self.interpolation = interpolation\n        self.max_h, self.max_w = input_size\n\n        def _fit_resize(image, **kwargs):\n            img_h, img_w = image.shape[:2]\n            r = min(self.max_h / img_h, self.max_w / img_w)\n            new_h, new_w = int(img_h * r), int(img_w * r)\n            new_image = cv2.resize(image, (new_w, new_h),\n                                   interpolation=interpolation)\n            return new_image\n\n        self.transform_fn = A.Compose([\n            A.Lambda(name=\"FitResize\",\n                     image=_fit_resize,\n                     always_apply=True,\n                     p=1.0),\n            A.PadIfNeeded(min_height=self.max_h,\n                          min_width=self.max_w,\n                          pad_height_divisor=None,\n                          pad_width_divisor=None,\n                          position=A.augmentations.geometric.transforms.\n                          PadIfNeeded.PositionType.CENTER,\n                          border_mode=cv2.BORDER_CONSTANT,\n                          value=0,\n                          mask_value=None,\n                          always_apply=True,\n                          p=1.0),\n            ToTensorV2(transpose_mask=True)\n        ])\n\n    def __call__(self, img):\n        return self.transform_fn(image=img)['image']\n\n\nTrainTransform = ValTransform\n\n\nclass Exp(ExpBaseline):\n\n    def __init__(self, args):\n        super(Exp, self).__init__(args)\n        self.meta = {\n            'fold_idx': 0,\n            'num_sched_epochs': 6,\n            'num_epochs': 50,\n            'start_ratio': 1 / 3,\n            'end_ratio': 1 / 7,\n            'one_pos_mode': True,\n        }\n        old_meta_len = len(self.meta)\n        self.meta.update(self.args.exp_kwargs)\n        assert len(self.meta) == old_meta_len\n        print('\\n------\\nEXP METADATA:\\n', self.meta)\n        self.output_dir = os.path.join(\"/kaggle/working\",\n                                       'timm_classification')\n\n    def build_train_dataset(self):\n        assert self.data_config is not None\n        fold_idx = self.meta['fold_idx']\n        augment_fn = TrainAugment()\n        transform_fn = TrainTransform(self.data_config['input_size'][1:])\n\n        rsna_train_dataset_info = {\n            'csv_path':\n            os.path.join('/kaggle/input/split-4-folds/classification/rsna-breast-cancer-detection/cv/v2',\n                         f'train_fold_{fold_idx}.csv'),\n            'img_dir':\n            os.path.join('/kaggle/input/rsnaroiextracted/ROI_extracted_1024x2048','ROI_extracted_1024x2048')\n        }\n        train_datasets_info = [('rsna-breast-cancer-detection', rsna_train_dataset_info)]\n        \n        train_dataset = RSNADataset(\n            train_datasets_info,\n            augment_fn,\n            transform_fn,\n            n_channels=self.args.input_size[0],\n            subset='train')\n        return train_dataset\n\n    def build_train_loader(self, collate_fn=None):\n        train_dataset = self.build_train_dataset()\n\n        # wrap dataset in AugMix helper\n        if self.args.num_aug_splits > 1:\n            train_dataset = AugMixDataset(train_dataset,\n                                          num_splits=self.args.num_aug_splits)\n\n        # create data loaders w/ augmentation pipeiine\n        train_interpolation = self.args.train_interpolation\n        if self.args.no_aug or not train_interpolation:\n            train_interpolation = self.data_config['interpolation']\n\n        if self.args.distributed:\n            assert not isinstance(train_dataset,\n                                  torch.utils.data.IterableDataset)\n            sampler = DistributedProxySampler(\n                BalanceSamplerV2(\n                    train_dataset,\n                    batch_size=self.args.batch_size,\n                    num_sched_epochs=self.meta['num_sched_epochs'],\n                    num_epochs=self.meta['num_epochs'],\n                    start_ratio=self.meta['start_ratio'],\n                    end_ratio=self.meta['end_ratio'],\n                    one_pos_mode=self.meta['one_pos_mode'],\n                    seed=self.args.seed))\n        else:\n            sampler = BalanceSamplerV2(\n                train_dataset,\n                batch_size=self.args.batch_size,\n                num_sched_epochs=self.meta['num_sched_epochs'],\n                num_epochs=self.meta['num_epochs'],\n                start_ratio=self.meta['start_ratio'],\n                end_ratio=self.meta['end_ratio'],\n                one_pos_mode=self.meta['one_pos_mode'],\n                seed=self.args.seed)\n\n        train_loader = create_loader(\n            train_dataset,\n            input_size=self.data_config['input_size'],\n            batch_size=self.args.batch_size,\n            is_training=True,\n            use_prefetcher=self.args.prefetcher,\n            no_aug=self.args.no_aug,\n            re_prob=self.args.reprob,\n            re_mode=self.args.remode,\n            re_count=self.args.recount,\n            re_split=self.args.resplit,\n            scale=self.args.scale,\n            ratio=self.args.ratio,\n            hflip=self.args.hflip,\n            vflip=self.args.vflip,\n            color_jitter=self.args.color_jitter,\n            auto_augment=self.args.aa,\n            num_aug_repeats=self.args.aug_repeats,\n            num_aug_splits=self.args.num_aug_splits,\n            interpolation=train_interpolation,\n            mean=self.data_config['mean'],\n            std=self.data_config['std'],\n            num_workers=self.args.workers,\n            distributed=self.args.distributed,\n            collate_fn=collate_fn,\n            pin_memory=self.args.pin_mem,\n            device=self.args.device,\n            use_multi_epochs_loader=self.args.use_multi_epochs_loader,\n            worker_seeding=self.args.worker_seeding,\n            sampler=sampler,\n        )\n        return train_loader\n\n    def build_val_dataset(self):\n        assert self.data_config is not None\n\n        fold_idx = self.meta['fold_idx']\n\n        augment_fn = ValAugment()\n        transform_fn = ValTransform(self.data_config['input_size'][1:])\n\n        rsna_val_dataset_info = {\n            'csv_path':\n            os.path.join(\"/kaggle/input/split-4-folds/classification/rsna-breast-cancer-detection/cv/v2\",\n                         f'val_fold_{fold_idx}.csv'),\n            'img_dir':\n            os.path.join('/kaggle/input/rsnaroiextracted/ROI_extracted_1024x2048','ROI_extracted_1024x2048')\n        }\n        val_datasets_info = [('rsna-breast-cancer-detection', rsna_val_dataset_info)]\n\n        val_dataset = RSNADataset(\n            val_datasets_info,\n            augment_fn,\n            transform_fn,\n            n_channels=self.args.input_size[0],\n            subset='val')\n        return val_dataset\n\n    def build_train_loss_fn(self):\n        pos_weight = self.args.pos_weight\n        pos_weight = None if pos_weight <= 0 else pos_weight\n        # assert self.args.smoothing > 0 and self.args.bce_loss\n        # setup loss function\n        if self.args.jsd_loss:\n            assert self.args.num_aug_splits > 1  # JSD only valid with aug splits set\n            train_loss_fn = JsdCrossEntropy(\n                num_splits=self.args.num_aug_splits,\n                smoothing=self.args.smoothing)\n        elif self.args.mixup_active:\n            # smoothing is handled with mixup target transform which outputs sparse, soft targets\n            if self.args.bce_loss:\n                train_loss_fn = BinaryCrossEntropy(\n                    target_threshold=self.args.bce_target_thresh)\n            else:\n                train_loss_fn = SoftTargetCrossEntropy()\n        elif self.args.smoothing:\n            if self.args.bce_loss:\n                if pos_weight is not None:\n                    assert self.args.num_classes == 1\n                    pos_weight = torch.Tensor([pos_weight])\n                print('Using pos weight:', pos_weight)\n                train_loss_fn = BinaryCrossEntropyPosSmoothOnly(\n                    smoothing=self.args.smoothing,\n                    target_threshold=self.args.bce_target_thresh,\n                    pos_weight=pos_weight)\n            else:\n                train_loss_fn = LabelSmoothingCrossEntropy(\n                    smoothing=self.args.smoothing)\n        else:\n            train_loss_fn = nn.CrossEntropyLoss()\n        train_loss_fn = train_loss_fn.to(device=self.args.device)\n        return train_loss_fn\n","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:13.727961Z","iopub.execute_input":"2024-04-20T14:04:13.728399Z","iopub.status.idle":"2024-04-20T14:04:14.148798Z","shell.execute_reply.started":"2024-04-20T14:04:13.72836Z","shell.execute_reply":"2024-04-20T14:04:14.147744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import argparse\n\n\n\n\nargs_name_space = argparse.Namespace(exp='./exp7.py', \n                          exp_kwargs={'fold_idx': 2, 'num_sched_epochs': 10, \n                                      'num_epochs': 35, 'start_ratio': 0.1429, 'end_ratio': 0.1429, 'one_pos_mode': True},\n                          data=None, data_dir=None, dataset='', train_split='train', val_split='validation', \n                          dataset_download=False, class_map='', model='convnext_small.fb_in22k_ft_in1k_384',\n                          pretrained=True, initial_checkpoint='', resume='', no_resume_opt=False, num_classes=1,\n                          gp='max', img_size=None, in_chans=None, input_size=[3, 2048, 1024], crop_pct=1.0, \n                          mean=None, std=None, interpolation='', batch_size=2, validation_batch_size=2, \n                          channels_last=False, fuser='', grad_checkpointing=False, fast_norm=False, model_kwargs={},\n                          torchscript=False, torchcompile=None, aot_autograd=False, opt='sgd', opt_eps=None, \n                          opt_betas=None, momentum=0.9, weight_decay=2e-05, clip_grad=None, clip_mode='norm',\n                          layer_decay=None, opt_kwargs={}, sched='cosine', sched_on_updates=False, lr=0.003, \n                          lr_base=0.1, lr_base_size=256, lr_base_scale='', lr_noise=None, lr_noise_pct=0.67, \n                          lr_noise_std=1.0, lr_cycle_mul=1.0, lr_cycle_decay=0.5, lr_cycle_limit=1, lr_k_decay=1.0,\n                          warmup_lr=3e-05, min_lr=5e-05, epochs=30, epoch_repeats=0.0, start_epoch=None, decay_milestones=[90, 180, 270],\n                          decay_epochs=90, warmup_epochs=4, warmup_prefix=False, cooldown_epochs=1, patience_epochs=10, decay_rate=0.1,\n                          no_aug=True, scale=[0.08, 1.0], ratio=[0.75, 1.3333333333333333], hflip=0.5, vflip=0.0, color_jitter=0.4, aa=None,\n                          aug_repeats=0, aug_splits=0, jsd_loss=False, bce_loss=True, bce_target_thresh=None, reprob=0.0, remode='pixel',\n                          recount=1, resplit=False, mixup=0.0, cutmix=0.0, cutmix_minmax=None, mixup_prob=1.0, \n                          mixup_switch_prob=0.5, \n                          mixup_mode='batch', mixup_off_epoch=0, smoothing=0.1, train_interpolation='random', drop=0.5, drop_connect=None, \n                          drop_path=0.2, drop_block=None, bn_momentum=None, bn_eps=None, sync_bn=False, dist_bn='reduce', split_bn=False, \n                          model_ema=True, model_ema_force_cpu=False, model_ema_decay=0.9998, seed=42, worker_seeding='all', log_interval=500,\n                          recovery_interval=0, checkpoint_hist=100, workers=8, save_images=True, amp=True, amp_dtype='float16',\n                          amp_impl='native',\n                          no_ddp_bb=False, pin_mem=False, no_prefetcher=False, output='', experiment_name='reproduce_train_fold_2', \n                          eval_metric='gbmean_best_pfbeta', tta=0, local_rank=0, use_multi_epochs_loader=False, log_wandb=True, pos_weight=-1,\n                          dense_ckpt_epochs=[10, 18], dense_ckpt_bins=2)\n    ","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:14.150457Z","iopub.execute_input":"2024-04-20T14:04:14.150779Z","iopub.status.idle":"2024-04-20T14:04:14.169877Z","shell.execute_reply.started":"2024-04-20T14:04:14.150751Z","shell.execute_reply":"2024-04-20T14:04:14.168877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from setproctitle import setproctitle\n\nsetproctitle(\"python3 train.py\")\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")\n\n# # slow/deadlock with opencv in dataloader\n\nimport cv2\n\ncv2.setNumThreads(0)\ncv2.ocl.setUseOpenCL(False)\n\n\nimport gc\n\nimport torch\n\n\ndef force_cudnn_initialization():\n    s = 32\n    dev = torch.device('cuda')\n    torch.nn.functional.conv2d(torch.zeros(s, s, s, s, device=dev), torch.zeros(s, s, s, s, device=dev))\n\nimport argparse\nimport logging\nimport math\nimport os\nimport time\nfrom collections import OrderedDict\nfrom contextlib import suppress\nfrom datetime import datetime\nfrom functools import partial\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torchvision.utils\nimport yaml\nfrom timm import utils\nfrom timm.data import (AugMixDataset, FastCollateMixup, Mixup, create_dataset,\n                       create_loader, resolve_data_config)\n# from timm.exp.build import get_exp_fn\nfrom timm.layers import (convert_splitbn_model, convert_sync_batchnorm,\n                         set_fast_norm)\nfrom timm.loss import (BinaryCrossEntropy, JsdCrossEntropy,\n                       LabelSmoothingCrossEntropy, SoftTargetCrossEntropy)\nfrom timm.models import (create_model, load_checkpoint, model_parameters,\n                         resume_checkpoint, safe_model_name)\nfrom timm.optim import create_optimizer_v2, optimizer_kwargs\nfrom timm.scheduler import create_scheduler_v2, scheduler_kwargs\nfrom timm.utils import ApexScaler, NativeScaler\nfrom torch.nn.parallel import DistributedDataParallel as NativeDDP\nfrom tqdm import tqdm\n\ntry:\n    from apex import amp\n    from apex.parallel import DistributedDataParallel as ApexDDP\n    from apex.parallel import convert_syncbn_model\n    has_apex = True\nexcept ImportError:\n    has_apex = False\n\nhas_native_amp = False\ntry:\n    if getattr(torch.cuda.amp, 'autocast') is not None:\n        has_native_amp = True\nexcept AttributeError:\n    pass\n\ntry:\n    import wandb\n    has_wandb = True\nexcept ImportError:\n    has_wandb = False\n\ntry:\n    from functorch.compile import memory_efficient_fusion\n    has_functorch = True\nexcept ImportError as e:\n    has_functorch = False\n\nhas_compile = hasattr(torch, 'compile')","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:14.171376Z","iopub.execute_input":"2024-04-20T14:04:14.172297Z","iopub.status.idle":"2024-04-20T14:04:14.193165Z","shell.execute_reply.started":"2024-04-20T14:04:14.172255Z","shell.execute_reply":"2024-04-20T14:04:14.19234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validate(\n        exp,\n        model,\n        loader,\n        loss_fn,\n        args,\n        device=torch.device('cuda'),\n        amp_autocast=suppress,\n        log_suffix='',\n        plot_save_path = None,\n        pred_save_path = None,\n):\n    batch_time_m = utils.AverageMeter()\n    losses_m = utils.AverageMeter()\n    top1_m = utils.AverageMeter()\n    top5_m = utils.AverageMeter()\n\n    model.eval()\n\n    end = time.time()\n    last_idx = len(loader) - 1\n\n    # buffer\n    # we use modified OrderedDistributedSampler, a little change from timm's\n    val_len = len(loader.dataset)\n    num_samples_per_rank = int(math.ceil(val_len / args.world_size))\n    temp_total_size = num_samples_per_rank * args.world_size\n    if utils.is_primary(args):\n        _logger.info(\n            f'Val samples: {val_len}, buffer size: {temp_total_size}')\n\n    targets = torch.zeros((temp_total_size,), requires_grad=False)\n    preds = torch.zeros_like(targets)\n    sample_weights =  torch.zeros_like(targets)\n\n    with torch.no_grad():\n        cur_sample_idx = 0\n        for batch_idx, (input, target) in tqdm(enumerate(loader)):\n            last_batch = batch_idx == last_idx\n            if not args.prefetcher:\n                input = input.to(device)\n                target = target.to(device)\n            if args.channels_last:\n                input = input.contiguous(memory_format=torch.channels_last)\n\n            with amp_autocast():\n                output = model(input)\n                if isinstance(output, (tuple, list)):\n                    output = output[0]\n\n                # augmentation reduction\n                reduce_factor = args.tta\n                if reduce_factor > 1:\n                    output = output.unfold(0, reduce_factor, reduce_factor).mean(dim=2)\n                    target = target[0:target.size(0):reduce_factor]\n\n                num_classes = output.shape[-1]\n                if num_classes == 1:\n                    loss = loss_fn(output, target.float().view(-1, 1))\n                else:\n                    loss = loss_fn(output, target)\n            acc1, acc5 = utils.accuracy(output, target, topk=(1, 5))\n\n            if args.distributed:\n                reduced_loss = utils.reduce_tensor(loss.data, args.world_size)\n                acc1 = utils.reduce_tensor(acc1, args.world_size)\n                acc5 = utils.reduce_tensor(acc5, args.world_size)\n            else:\n                reduced_loss = loss.data\n\n            if device.type == 'cuda':\n                torch.cuda.synchronize()\n\n            losses_m.update(reduced_loss.item(), input.size(0))\n            top1_m.update(acc1.item(), output.size(0))\n            top5_m.update(acc5.item(), output.size(0))\n\n            # CUSTOM METRIC\n            if args.distributed:\n                # gather preds\n                # @TODO: only available in torch > 1.13 ?\n                # output = utils.all_gather_tensor_into_tensor(output, args.world_size)\n                output = utils.all_gather_tensor(output, args.world_size)\n                # is the list in ordered by rank?\n                # https://discuss.pytorch.org/t/order-of-the-list-returned-by-torch-distributed-all-gather/125273\n                # https://github.com/pytorch/pytorch/issues/23144\n                output = torch.cat(output, dim = 0)\n            \n                # gather targets\n                target = utils.all_gather_tensor(target, args.world_size)\n                target = torch.cat(target, dim = 0)\n\n            if output.shape[-1] == 2:\n                prob = output.softmax(dim = 1)[:, 1]\n            elif output.shape[-1] == 1:\n                prob = output.sigmoid().view(-1, )\n            else:\n                raise AssertionError()\n            \n            nan_mask = torch.isnan(prob)\n            if torch.sum(nan_mask.long()) > 0:\n                print('CONTAIN NAN:', prob[nan_mask], output[nan_mask])\n                prob[nan_mask] = 0.\n            # prob = torch.nan_to_num(prob, nan = 0.0)\n            \n            output = prob\n            _targets = target\n            _preds = output\n            _sample_weights = torch.ones_like(_preds)\n\n            cur_bs = output.size(0)\n            targets[cur_sample_idx:cur_sample_idx + cur_bs] = _targets\n            preds[cur_sample_idx:cur_sample_idx + cur_bs] = _preds\n            sample_weights[cur_sample_idx:cur_sample_idx + cur_bs] = _sample_weights\n            cur_sample_idx += cur_bs\n\n            batch_time_m.update(time.time() - end)\n            end = time.time()\n            if utils.is_primary(args) and (last_batch or batch_idx % args.log_interval == 0):\n                log_name = 'Test' + log_suffix\n                _logger.info(\n                    '{0}: [{1:>4d}/{2}]  '\n                    'Time: {batch_time.val:.3f} ({batch_time.avg:.3f})  '\n                    'Loss: {loss.val:>7.4f} ({loss.avg:>6.4f})  '\n                    'Acc@1: {top1.val:>7.4f} ({top1.avg:>7.4f})  '\n                    'Acc@5: {top5.val:>7.4f} ({top5.avg:>7.4f})'.format(\n                        log_name, batch_idx, last_idx,\n                        batch_time=batch_time_m,\n                        loss=losses_m,\n                        top1=top1_m,\n                        top5=top5_m)\n                )\n\n    # truncate duplicated samples (tail)\n    targets = targets[:val_len].cpu().numpy()\n    preds = preds[:val_len].cpu().numpy()\n    # why nan = 1 ?\n    preds = np.nan_to_num(preds, nan=1, posinf=1, neginf=0)\n    sample_weights = sample_weights[:val_len].cpu().numpy()\n    val_df = loader.dataset.get_df()\n    val_df['preds'] = preds\n    val_df['targets'] = targets\n    val_df['sample_weights'] = sample_weights\n    assert (val_df['targets'] == val_df['cancer']).all()\n\n    if utils.is_primary(args):\n        additional_info = True\n        val_df.to_csv(pred_save_path, index = False)\n    else:\n        additional_info = False\n    \n    metric_results = exp.compute_metrics(\n        val_df,\n        plot_save_path,\n        additional_info = additional_info\n    )\n\n    metrics =[('loss', losses_m.avg), ('top1', top1_m.avg)]         \n    metrics.extend([(k, v) for k, v in metric_results.items()])\n    metrics = OrderedDict(metrics)\n    # print('---------------')\n    # print(metrics)\n    # print('---------------')\n    return metrics","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:14.194879Z","iopub.execute_input":"2024-04-20T14:04:14.19524Z","iopub.status.idle":"2024-04-20T14:04:14.230224Z","shell.execute_reply.started":"2024-04-20T14:04:14.19521Z","shell.execute_reply":"2024-04-20T14:04:14.22909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_one_epoch(\n        exp,    # change\n        eval_loader,    # change\n        val_loss_fn,    # change\n        best_metric,    # change\n        best_epoch,     # change\n        epoch,\n        model,\n        loader,\n        optimizer,\n        loss_fn,\n        args,\n        device=torch.device('cuda'),\n        lr_scheduler=None,\n        saver=None,\n        output_dir=None,\n        amp_autocast=suppress,\n        loss_scaler=None,\n        model_ema=None,\n        mixup_fn=None,\n        num_updates = None,     # change\n):\n    if args.mixup_off_epoch and epoch >= args.mixup_off_epoch:\n        if args.prefetcher and loader.mixup_enabled:\n            loader.mixup_enabled = False\n        elif mixup_fn is not None:\n            mixup_fn.mixup_enabled = False\n\n    second_order = hasattr(optimizer, 'is_second_order') and optimizer.is_second_order\n    batch_time_m = utils.AverageMeter()\n    data_time_m = utils.AverageMeter()\n    losses_m = utils.AverageMeter()\n\n    model.train()\n\n    end = time.time()\n    num_batches_per_epoch = len(loader)\n    last_idx = num_batches_per_epoch - 1\n    if num_updates is None:\n        num_updates = epoch * num_batches_per_epoch\n    else:\n        print('Current number of updates:', num_updates)\n    epoch_num_updates = 0\n    dense_ckpt_start_epoch, dense_ckpt_end_epoch = args.dense_ckpt_epochs\n    dense_ckpt_interval = num_batches_per_epoch // args.dense_ckpt_bins + 1\n    print(f'DENSE CKPT START/END: {args.dense_ckpt_epochs}, interval = {dense_ckpt_interval}')\n    for batch_idx, (input, target) in tqdm(enumerate(loader)):\n        last_batch = batch_idx == last_idx\n        data_time_m.update(time.time() - end)\n        if not args.prefetcher:\n            input, target = input.to(device), target.to(device)\n            if mixup_fn is not None:\n                input, target = mixup_fn(input, target)\n        if args.channels_last:\n            input = input.contiguous(memory_format=torch.channels_last)\n\n        with amp_autocast():\n            output = model(input)\n            loss = loss_fn(output, target)\n\n        if not args.distributed:\n            losses_m.update(loss.item(), input.size(0))\n\n        optimizer.zero_grad()\n        if loss_scaler is not None:\n            loss_scaler(\n                loss, optimizer,\n                clip_grad=args.clip_grad,\n                clip_mode=args.clip_mode,\n                parameters=model_parameters(model, exclude_head='agc' in args.clip_mode),\n                create_graph=second_order\n            )\n        else:\n            loss.backward(create_graph=second_order)\n            if args.clip_grad is not None:\n                utils.dispatch_clip_grad(\n                    model_parameters(model, exclude_head='agc' in args.clip_mode),\n                    value=args.clip_grad,\n                    mode=args.clip_mode\n                )\n            optimizer.step()\n\n        if model_ema is not None:\n            model_ema.update(model)\n\n        torch.cuda.synchronize()\n\n        num_updates += 1\n        epoch_num_updates += 1\n        batch_time_m.update(time.time() - end)\n        if last_batch or batch_idx % args.log_interval == 0:\n            lrl = [param_group['lr'] for param_group in optimizer.param_groups]\n            lr = sum(lrl) / len(lrl)\n\n            if args.distributed:\n                reduced_loss = utils.reduce_tensor(loss.data, args.world_size)\n                losses_m.update(reduced_loss.item(), input.size(0))\n\n            if utils.is_primary(args):\n                _logger.info(\n                    'Train: {} [{:>4d}/{} ({:>3.0f}%)]  '\n                    'Loss: {loss.val:#.4g} ({loss.avg:#.3g})  '\n                    'Time: {batch_time.val:.3f}s, {rate:>7.2f}/s  '\n                    '({batch_time.avg:.3f}s, {rate_avg:>7.2f}/s)  '\n                    'LR: {lr:.3e}  '\n                    'Data: {data_time.val:.3f} ({data_time.avg:.3f})'.format(\n                        epoch,\n                        batch_idx, len(loader),\n                        100. * batch_idx / last_idx,\n                        loss=losses_m,\n                        batch_time=batch_time_m,\n                        rate=input.size(0) * args.world_size / batch_time_m.val,\n                        rate_avg=input.size(0) * args.world_size / batch_time_m.avg,\n                        lr=lr,\n                        data_time=data_time_m)\n                )\n\n                if args.save_images and output_dir:\n                    torchvision.utils.save_image(\n                        input,\n                        os.path.join(output_dir, 'train-batch-%d.jpg' % batch_idx),\n                        padding=0,\n                        normalize=True\n                    )\n\n        if saver is not None and args.recovery_interval and (\n                last_batch or (batch_idx + 1) % args.recovery_interval == 0):\n            saver.save_recovery(epoch, batch_idx=batch_idx)\n\n        if lr_scheduler is not None:\n            lr_scheduler.step_update(num_updates=num_updates, metric=losses_m.avg)\n\n        end = time.time()\n        # end for\n\n\n        ##########################################################\n        ##########################################################\n        ##########################################################\n        # perform dense validate and save model checkpointing\n        # START\n\n        if epoch >= dense_ckpt_start_epoch and epoch <= dense_ckpt_end_epoch and epoch_num_updates % dense_ckpt_interval == 0:\n            print('\\n-----DENSE CKPT VALIDATION-----\\n')\n            \n            if args.distributed and args.dist_bn in ('broadcast', 'reduce'):\n                if utils.is_primary(args):\n                    _logger.info(\"Distributing BatchNorm running means and vars\")\n                utils.distribute_bn(model, args.world_size, args.dist_bn == 'reduce')\n            \n            temp_epoch = round(epoch - 1 + epoch_num_updates / num_batches_per_epoch , 2) \n            if args.save_results_dir is not None:\n                plot_save_path = os.path.join(args.save_results_dir, f'plot_{temp_epoch}.jpg')\n                ema_plot_save_path = os.path.join(args.save_results_dir, f'ema_plot_{temp_epoch}.jpg')\n                pred_save_path = os.path.join(args.save_results_dir, f'pred_{temp_epoch}.csv')\n            else:\n                plot_save_path = None\n                ema_plot_save_path = None\n\n            eval_metrics = None\n            if model_ema is not None and not args.model_ema_force_cpu:\n                if args.distributed and args.dist_bn in ('broadcast', 'reduce'):\n                    utils.distribute_bn(model_ema, args.world_size, args.dist_bn == 'reduce')\n                eval_metrics = validate(\n                    exp,\n                    model_ema.module,\n                    eval_loader,\n                    val_loss_fn,\n                    args,\n                    amp_autocast=amp_autocast,\n                    log_suffix=' (EMA)',\n                    plot_save_path = ema_plot_save_path,\n                    pred_save_path = pred_save_path,\n                )\n            else:\n                eval_metrics = validate(\n                exp,\n                model,\n                eval_loader,\n                val_loss_fn,\n                args,\n                amp_autocast=amp_autocast,\n                plot_save_path = plot_save_path,\n                pred_save_path = pred_save_path,\n            )\n\n            # primary process/rank only\n            if output_dir is not None:\n                lrs = [param_group['lr'] for param_group in optimizer.param_groups]\n                utils.update_summary(\n                    temp_epoch,\n                    OrderedDict([('loss', losses_m.avg)]),\n                    eval_metrics,\n                    filename=os.path.join(output_dir, 'summary.csv'),\n                    lr=sum(lrs) / len(lrs),\n                    write_header=best_metric is None,\n                    log_wandb=args.log_wandb and has_wandb,\n                )\n\n            if saver is not None:\n                # save proper checkpoint with eval metric\n                save_metric = eval_metrics[args.eval_metric]\n                best_metric, best_epoch = saver.save_checkpoint(temp_epoch, metric=save_metric)\n\n            model.train()\n            print('\\n---------------------------------------\\n')\n        # END\n        ##########################################################\n        ##########################################################\n        ##########################################################\n\n    if hasattr(optimizer, 'sync_lookahead'):\n        optimizer.sync_lookahead()\n\n    return OrderedDict([('loss', losses_m.avg)]), num_updates, best_metric, best_epoch","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:14.231708Z","iopub.execute_input":"2024-04-20T14:04:14.232096Z","iopub.status.idle":"2024-04-20T14:04:14.27419Z","shell.execute_reply.started":"2024-04-20T14:04:14.232058Z","shell.execute_reply":"2024-04-20T14:04:14.273155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main():\n    utils.setup_default_logging()\n    args = args_name_space\n\n    if torch.cuda.is_available():\n        torch.backends.cuda.matmul.allow_tf32 = True\n        torch.backends.cudnn.benchmark = True\n    args.prefetcher = not args.no_prefetcher\n    device = utils.init_distributed_device(args)\n    # prevent OOM error while initializing CUDNN\n    force_cudnn_initialization()\n    torch.cuda.empty_cache()\n\n    args.device = device\n    if args.distributed:\n        _logger.info(\n            'Training in distributed mode with multiple processes, 1 device per process.'\n            f'Process {args.rank}, total {args.world_size}, device {args.device}.')\n    else:\n        _logger.info(f'Training with a single process on 1 device ({args.device}).')\n    assert args.rank >= 0\n\n    # resolve AMP arguments based on PyTorch / Apex availability\n    use_amp = None\n    amp_dtype = torch.float16\n    if args.amp:\n        if args.amp_impl == 'apex':\n            assert has_apex, 'AMP impl specified as APEX but APEX is not installed.'\n            use_amp = 'apex'\n            assert args.amp_dtype == 'float16'\n        else:\n            assert has_native_amp, 'Please update PyTorch to a version with native AMP (or use APEX).'\n            use_amp = 'native'\n            assert args.amp_dtype in ('float16', 'bfloat16')\n        if args.amp_dtype == 'bfloat16':\n            amp_dtype = torch.bfloat16\n\n    utils.random_seed(args.seed, args.rank)\n\n    if args.fuser:\n        utils.set_jit_fuser(args.fuser)\n    if args.fast_norm:\n        set_fast_norm()\n\n    exp_fn = Exp\n    exp = exp_fn(args)\n\n    ### BUILD MODEL\n    model = exp.build_model()\n    \n    if args.num_classes is None:\n        assert hasattr(model, 'num_classes'), 'Model must have `num_classes` attr if not set on cmd line/config.'\n        args.num_classes = model.num_classes  # FIXME handle model default vs config num_classes more elegantly\n\n    if args.grad_checkpointing:\n        model.set_grad_checkpointing(enable=True)\n\n    # setup augmentation batch splits for contrastive loss or split bn\n    num_aug_splits = 0\n    if args.aug_splits > 0:\n        assert args.aug_splits > 1, 'A split of 1 makes no sense'\n        num_aug_splits = args.aug_splits\n    args.num_aug_splits = num_aug_splits\n\n    # enable split bn (separate bn stats per batch-portion)\n    if args.split_bn:\n        assert num_aug_splits > 1 or args.resplit\n        model = convert_splitbn_model(model, max(num_aug_splits, 2))\n\n    # move model to GPU, enable channels last layout if set\n    model.to(device=device)\n    if args.channels_last:\n        model.to(memory_format=torch.channels_last)\n\n    # setup synchronized BatchNorm for distributed training\n    if args.distributed and args.sync_bn:\n        args.dist_bn = ''  # disable dist_bn when sync BN active\n        assert not args.split_bn\n        if has_apex and use_amp == 'apex':\n            # Apex SyncBN used with Apex AMP\n            # WARNING this won't currently work with models using BatchNormAct2d\n            model = convert_syncbn_model(model)\n        else:\n            model = convert_sync_batchnorm(model)\n        if utils.is_primary(args):\n            _logger.info(\n                'Converted model to use Synchronized BatchNorm. WARNING: You may have issues if using '\n                'zero initialized BN layers (enabled by default for ResNets) while sync-bn enabled.')\n\n    if args.torchscript:\n        assert not use_amp == 'apex', 'Cannot use APEX AMP with torchscripted model'\n        assert not args.sync_bn, 'Cannot use SyncBatchNorm with torchscripted model'\n        model = torch.jit.script(model)\n    elif args.torchcompile:\n        # FIXME dynamo might need move below DDP wrapping? TBD\n        assert has_compile, 'A version of torch w/ torch.compile() is required for --compile, possibly a nightly.'\n        torch._dynamo.reset()\n        model = torch.compile(model, backend=args.torchcompile)\n    elif args.aot_autograd:\n        assert has_functorch, \"functorch is needed for --aot-autograd\"\n        model = memory_efficient_fusion(model)\n\n    if not args.lr:\n        global_batch_size = args.batch_size * args.world_size\n        batch_ratio = global_batch_size / args.lr_base_size\n        if not args.lr_base_scale:\n            on = args.opt.lower()\n            args.lr_base_scale = 'sqrt' if any([o in on for o in ('ada', 'lamb')]) else 'linear'\n        if args.lr_base_scale == 'sqrt':\n            batch_ratio = batch_ratio ** 0.5\n        args.lr = args.lr_base * batch_ratio\n        if utils.is_primary(args):\n            _logger.info(\n                f'Learning rate ({args.lr}) calculated from base learning rate ({args.lr_base}) '\n                f'and global batch size ({global_batch_size}) with {args.lr_base_scale} scaling.')\n\n    ### BUILD OPTIMIZER\n    optimizer = exp.build_optimizer(model)\n\n    # setup automatic mixed-precision (AMP) loss scaling and op casting\n    amp_autocast = suppress  # do nothing\n    loss_scaler = None\n    if use_amp == 'apex':\n        assert device.type == 'cuda'\n        model, optimizer = amp.initialize(model, optimizer, opt_level='O1')\n        loss_scaler = ApexScaler()\n        if utils.is_primary(args):\n            _logger.info('Using NVIDIA APEX AMP. Training in mixed precision.')\n    elif use_amp == 'native':\n        amp_autocast = partial(torch.autocast, device_type=device.type, dtype=amp_dtype)\n        if device.type == 'cuda':\n            loss_scaler = NativeScaler()\n        if utils.is_primary(args):\n            _logger.info('Using native Torch AMP. Training in mixed precision.')\n    else:\n        if utils.is_primary(args):\n            _logger.info('AMP not enabled. Training in float32.')\n\n    # optionally resume from a checkpoint\n    resume_epoch = None\n    if args.resume:\n        resume_epoch = resume_checkpoint(\n            model,\n            args.resume,\n            optimizer=None if args.no_resume_opt else optimizer,\n            loss_scaler=None if args.no_resume_opt else loss_scaler,\n            log_info=utils.is_primary(args),\n        )\n\n    # setup exponential moving average of model weights, SWA could be used here too\n    model_ema = None\n    if args.model_ema:\n        # Important to create EMA model after cuda(), DP wrapper, and AMP but before DDP wrapper\n        model_ema = utils.ModelEmaV2(\n            model, decay=args.model_ema_decay, device='cpu' if args.model_ema_force_cpu else None)\n        if args.resume:\n            load_checkpoint(model_ema.module, args.resume, use_ema=True)\n\n    # setup distributed training\n    if args.distributed:\n        if has_apex and use_amp == 'apex':\n            # Apex DDP preferred unless native amp is activated\n            if utils.is_primary(args):\n                _logger.info(\"Using NVIDIA APEX DistributedDataParallel.\")\n            model = ApexDDP(model, delay_allreduce=True)\n        else:\n            if utils.is_primary(args):\n                _logger.info(\"Using native Torch DistributedDataParallel.\")\n            model = NativeDDP(model, device_ids=[device], broadcast_buffers=not args.no_ddp_bb)\n        # NOTE: EMA model does not need to be wrapped by DDP\n\n    # create the train and eval datasets\n    if args.data and not args.data_dir:\n        args.data_dir = args.data\n\n    ### BUILD LOADERS\n    # setup mixup / cutmix\n    collate_fn = None\n    mixup_fn = None\n    mixup_active = args.mixup > 0 or args.cutmix > 0. or args.cutmix_minmax is not None\n    if mixup_active:\n        mixup_args = dict(\n            mixup_alpha=args.mixup,\n            cutmix_alpha=args.cutmix,\n            cutmix_minmax=args.cutmix_minmax,\n            prob=args.mixup_prob,\n            switch_prob=args.mixup_switch_prob,\n            mode=args.mixup_mode,\n            label_smoothing=args.smoothing,\n            num_classes=args.num_classes\n        )\n        if args.prefetcher:\n            assert not num_aug_splits  # collate conflict (need to support deinterleaving in collate mixup)\n            collate_fn = FastCollateMixup(**mixup_args)\n        else:\n            mixup_fn = Mixup(**mixup_args)\n    args.mixup_active = mixup_active\n\n    train_loader = exp.build_train_loader(collate_fn)\n    eval_loader = exp.build_val_loader()\n\n    ### BUILD LOSSES\n    train_loss_fn = exp.build_train_loss_fn()\n    val_loss_fn = exp.build_val_loss_fn()\n    print('TRAIN LOSS FUNCTION:', train_loss_fn)\n    print('VAL LOSS FUNCTION:', val_loss_fn)\n\n    # setup checkpoint saver and eval metric tracking\n    eval_metric = args.eval_metric\n    best_metric = None\n    best_epoch = None\n    saver = None\n    output_dir = None\n    if utils.is_primary(args):\n        if args.experiment_name:\n            exp_name = args.experiment_name\n        else:\n            exp_name = '-'.join([\n                datetime.now().strftime(\"%Y%m%d-%H%M%S\"),\n                safe_model_name(args.model),\n                f'fold{args.exp_kwargs[\"fold_idx\"]}',\n                'x'.join([str(e) for e in exp.data_config['input_size']])\n            ])\n        if args.output:\n            output_dir = args.output\n        else:\n            try:\n                output_dir = exp.output_dir\n            except:\n                output_dir = './output/train'\n        output_dir = os.path.join(output_dir, exp_name)\n        assert not os.path.exists(output_dir), f'Output directory {output_dir} exist !'\n        os.makedirs(output_dir, exist_ok=False)\n\n        decreasing = True if eval_metric == 'loss' else False\n        saver = utils.CheckpointSaver(\n            model=model,\n            optimizer=optimizer,\n            args=args,\n            model_ema=model_ema,\n            amp_scaler=loss_scaler,\n            checkpoint_dir=output_dir,\n            recovery_dir=output_dir,\n            decreasing=decreasing,\n            max_history=args.checkpoint_hist\n        )\n        with open(os.path.join(output_dir, 'args.yaml'), 'w') as f:\n            f.write(args_text)\n\n    args.output_dir = output_dir\n    if utils.is_primary(args) and args.log_wandb:\n        if has_wandb:\n            wandb.init(project=args.experiment_name, config=args)\n        else:\n            _logger.warning(\n                \"You've requested to log metrics to wandb but package not found. \"\n                \"Metrics not being logged to wandb, try `pip install wandb`\")\n\n    # setup learning rate schedule and starting epoch\n    updates_per_epoch = len(train_loader)\n    args.updates_per_epoch = updates_per_epoch\n    lr_scheduler, num_epochs = exp.build_lr_scheduler(optimizer)\n    \n    start_epoch = 0\n    if args.start_epoch is not None:\n        # a specified start_epoch will always override the resume epoch\n        start_epoch = args.start_epoch\n    elif resume_epoch is not None:\n        start_epoch = resume_epoch\n    if lr_scheduler is not None and start_epoch > 0:\n        if args.sched_on_updates:\n            lr_scheduler.step_update(start_epoch * updates_per_epoch)\n        else:\n            lr_scheduler.step(start_epoch)\n\n    if utils.is_primary(args):\n        _logger.info(\n            f'Scheduled epochs: {num_epochs}. LR stepped per {\"epoch\" if lr_scheduler.t_in_epochs else \"update\"}.')\n\n    torch.cuda.empty_cache()\n    if output_dir is not None:\n        save_results_dir = os.path.join(output_dir, 'results')\n        os.makedirs(save_results_dir)\n    else:\n        save_results_dir = None\n    # print('MODEL:\\n', model)\n    args.save_results_dir = save_results_dir\n\n    assert train_loader.sampler.num_epochs > num_epochs\n\n    num_updates = 0\n    try:\n        for epoch in range(start_epoch, num_epochs):\n            print(f'START EPOCH {epoch}')\n            if hasattr(train_loader.dataset, 'set_epoch'):\n                train_loader.dataset.set_epoch(epoch)\n            elif hasattr(train_loader.sampler, 'set_epoch'):\n                train_loader.sampler.set_epoch(epoch)\n\n            train_metrics, num_updates, best_metric, best_epoch = train_one_epoch(\n                exp,\n                eval_loader,\n                val_loss_fn,\n                best_metric,\n                best_epoch,\n                epoch,\n                model,\n                train_loader,\n                optimizer,\n                train_loss_fn,\n                args,\n                lr_scheduler=lr_scheduler,\n                saver=saver,\n                output_dir=output_dir,\n                amp_autocast=amp_autocast,\n                loss_scaler=loss_scaler,\n                model_ema=model_ema,\n                mixup_fn=mixup_fn,\n                num_updates = num_updates,\n            )\n\n            if args.distributed and args.dist_bn in ('broadcast', 'reduce'):\n                if utils.is_primary(args):\n                    _logger.info(\"Distributing BatchNorm running means and vars\")\n                utils.distribute_bn(model, args.world_size, args.dist_bn == 'reduce')\n\n            if save_results_dir is not None:\n                plot_save_path = os.path.join(save_results_dir, f'plot_{epoch}.jpg')\n                ema_plot_save_path = os.path.join(save_results_dir, f'ema_plot_{epoch}.jpg')\n                pred_save_path = os.path.join(args.save_results_dir, f'pred_{epoch}.csv')\n            else:\n                plot_save_path = None\n                ema_plot_save_path = None\n\n            if model_ema is not None and not args.model_ema_force_cpu:\n                if args.distributed and args.dist_bn in ('broadcast', 'reduce'):\n                    utils.distribute_bn(model_ema, args.world_size, args.dist_bn == 'reduce')\n\n                eval_metrics = validate(\n                    exp,\n                    model_ema.module,\n                    eval_loader,\n                    val_loss_fn,\n                    args,\n                    amp_autocast=amp_autocast,\n                    log_suffix=' (EMA)',\n                    plot_save_path = ema_plot_save_path,\n                    pred_save_path= pred_save_path,\n                )\n            else:\n                eval_metrics = validate(\n                    exp,\n                    model,\n                    eval_loader,\n                    val_loss_fn,\n                    args,\n                    amp_autocast=amp_autocast,\n                    plot_save_path = plot_save_path,\n                    pred_save_path = pred_save_path,\n                )\n\n            # primary process/rank only\n            if output_dir is not None:\n                lrs = [param_group['lr'] for param_group in optimizer.param_groups]\n                utils.update_summary(\n                    epoch,\n                    train_metrics,\n                    eval_metrics,\n                    filename=os.path.join(output_dir, 'summary.csv'),\n                    lr=sum(lrs) / len(lrs),\n                    write_header=best_metric is None,\n                    log_wandb=args.log_wandb and has_wandb,\n                )\n\n            if saver is not None:\n                # save proper checkpoint with eval metric\n                save_metric = eval_metrics[eval_metric]\n                best_metric, best_epoch = saver.save_checkpoint(epoch, metric=save_metric)\n\n            if lr_scheduler is not None:\n                # step LR for next epoch\n                lr_scheduler.step(epoch + 1, eval_metrics[eval_metric])\n\n    except KeyboardInterrupt:\n        pass\n\n    if best_metric is not None:\n        _logger.info('*** Best metric: {0} (epoch {1})'.format(best_metric, best_epoch))","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:14.275892Z","iopub.execute_input":"2024-04-20T14:04:14.276728Z","iopub.status.idle":"2024-04-20T14:04:14.341387Z","shell.execute_reply.started":"2024-04-20T14:04:14.276684Z","shell.execute_reply":"2024-04-20T14:04:14.340395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"main()","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:04:14.342823Z","iopub.execute_input":"2024-04-20T14:04:14.34324Z","iopub.status.idle":"2024-04-20T14:04:41.16393Z","shell.execute_reply.started":"2024-04-20T14:04:14.343201Z","shell.execute_reply":"2024-04-20T14:04:41.16253Z"},"trusted":true},"execution_count":null,"outputs":[]}]}