{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"**Code for 1st place solution:**\n- Solution write-up: https://www.kaggle.com/competitions/rsna-breast-cancer-detection/discussion/392449\n- Training code: https://github.com/dangnh0611/kaggle_rsna_breast_cancer","metadata":{}},{"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!pip uninstall -y timm\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\n# try:\n#     import nvjpeg2k\n# except:\n#     # For NVJPEG2k\n#     !cp /kaggle/input/kaggle-rsna-pkgs/nvjpeg2k.so ./\n\nprint('Import done!')","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"scrolled":true,"execution":{"iopub.status.busy":"2023-06-08T15:34:36.350384Z","iopub.execute_input":"2023-06-08T15:34:36.350848Z","iopub.status.idle":"2023-06-08T15:39:31.927803Z","shell.execute_reply.started":"2023-06-08T15:34:36.35076Z","shell.execute_reply":"2023-06-08T15:39:31.926436Z"},"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":"2023-06-08T15:49:34.575313Z","iopub.execute_input":"2023-06-08T15:49:34.576472Z","iopub.status.idle":"2023-06-08T15:49:36.977364Z","shell.execute_reply.started":"2023-06-08T15:49:34.576388Z","shell.execute_reply":"2023-06-08T15:49:36.976258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metrics\nMetrics computation, for local validation only. This code is not well-refactored","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport sklearn\nfrom sklearn import metrics\nfrom sklearn.metrics import (auc, confusion_matrix,\n                             precision_recall_fscore_support, roc_curve)\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\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    }\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 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 = pfbeta_np(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)\n\n\ndef _compute_metrics(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, preds, 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, bin_preds, 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    }\n\n\ndef compute_metrics(df,\n                    plot_save_path='plot.png',\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']).mean()\n    reducer_gbmax = lambda df: df.groupby(['patient_id', 'laterality']).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 = _compute_metrics(gts, preds, None, thres_range, 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","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-08T15:51:24.467383Z","iopub.execute_input":"2023-06-08T15:51:24.467819Z","iopub.status.idle":"2023-06-08T15:51:24.551473Z","shell.execute_reply.started":"2023-06-08T15:51:24.467784Z","shell.execute_reply":"2023-06-08T15:51:24.550278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ROI extraction (YOLOX)\n\n- YOLOX-nano 416 x 416\n- Otsu thresholding + findContours() as fallback","metadata":{}},{"cell_type":"code","source":"%%writefile roi_extract.py\n\n# Separated in .py file instead of a notebook cell for easier multiprocessing (e.g spawn)\nimport os\nos.environ['CUDA_MODULE_LOADING'] = 'LAZY'\nimport sys\nimport cv2\nimport numpy as np\nimport torch\nimport torchvision\nsys.path.append('/kaggle/tmp/libs/')\nfrom torch2trt import TRTModule\nfrom torch.nn import functional as F\n\n_TORCH_VER = [int(x) for x in torch.__version__.split(\".\")[:2]]\n_TORCH11X = (_TORCH_VER >= [1, 10])\n\n\ndef meshgrid(*tensors):\n    if _TORCH11X:\n        return torch.meshgrid(*tensors, indexing=\"ij\")\n    else:\n        return torch.meshgrid(*tensors)\n\n\ndef extract_roi_otsu(img, gkernel=(5, 5)):\n    \"\"\"WARNING: this function modify input image inplace.\"\"\"\n    ori_h, ori_w = img.shape[:2]\n    # clip percentile: implant, white lines\n    upper = np.percentile(img, 95)\n    img[img > upper] = np.min(img)\n    # Gaussian filtering to reduce noise (optional)\n    if gkernel is not None:\n        img = cv2.GaussianBlur(img, gkernel, 0)\n    _, img_bin = cv2.threshold(img, 0, 255,\n                               cv2.THRESH_BINARY + cv2.THRESH_OTSU)\n    # dilation to improve contours connectivity\n    element = cv2.getStructuringElement(cv2.MORPH_RECT, (3, 3), (-1, -1))\n    img_bin = cv2.dilate(img_bin, element)\n    cnts, _ = cv2.findContours(img_bin, cv2.RETR_EXTERNAL,\n                               cv2.CHAIN_APPROX_SIMPLE)\n    if len(cnts) == 0:\n        return None, None, None\n    areas = np.array([cv2.contourArea(cnt) for cnt in cnts])\n    select_idx = np.argmax(areas)\n    cnt = cnts[select_idx]\n    area_pct = areas[select_idx] / (img.shape[0] * img.shape[1])\n    x0, y0, w, h = cv2.boundingRect(cnt)\n    # min-max for safety only\n    # x0, y0, x1, y1\n    x1 = min(max(int(x0 + w), 0), ori_w)\n    y1 = min(max(int(y0 + h), 0), ori_h)\n    x0 = min(max(int(x0), 0), ori_w)\n    y0 = min(max(int(y0), 0), ori_h)\n    return [x0, y0, x1, y1], area_pct, None\n\n\nclass RoiExtractor:\n\n    def __init__(self,\n                 engine_path,\n                 input_size,\n                 num_classes,\n                 conf_thres=0.5,\n                 nms_thres=0.9,\n                 class_agnostic=False,\n                 area_pct_thres=0.04,\n                 hw=None,\n                 strides=None,\n                 exp=None):\n        self.input_size = input_size\n        self.input_h, self.input_w = input_size\n        self.num_classes = num_classes\n        self.conf_thres = conf_thres\n        self.nms_thres = nms_thres\n        self.class_agnostic = class_agnostic\n        self.area_pct_thres = area_pct_thres\n\n        model = TRTModule()\n        model.load_state_dict(torch.load(engine_path))\n        self.model = model\n        if hw is None or strides is None:\n            assert exp is not None\n            self._set_meta(exp)\n        else:\n            self.hw = hw\n            self.strides = strides\n\n    def _set_meta(self, exp):\n        assert exp is not None\n        print(\"Start probing model metadata..\")\n        # dummy infer\n        torch_model = exp.get_model().cuda().eval()\n        _dummy = torch.ones(1, 3, exp.test_size[0], exp.test_size[1]).cuda()\n        torch_model(_dummy)\n        # set attributes\n        self.hw = torch_model.head.hw\n        self.strides = torch_model.head.strides\n        # cleanup\n        del torch_model, _dummy\n        import gc\n        gc.collect()\n        torch.cuda.empty_cache()\n        print('Done probbing model metadata..')\n\n    def decode_outputs(self, outputs):\n        dtype = outputs.type()\n        grids = []\n        strides = []\n        for (hsize, wsize), stride in zip(self.hw, self.strides):\n            yv, xv = meshgrid([torch.arange(hsize), torch.arange(wsize)])\n            grid = torch.stack((xv, yv), 2).view(1, -1, 2)\n            grids.append(grid)\n            shape = grid.shape[:2]\n            strides.append(torch.full((*shape, 1), stride))\n\n        grids = torch.cat(grids, dim=1).type(dtype)\n        strides = torch.cat(strides, dim=1).type(dtype)\n\n        outputs = torch.cat(\n            [(outputs[..., 0:2] + grids) * strides,\n             torch.exp(outputs[..., 2:4]) * strides, outputs[..., 4:]],\n            dim=-1)\n        return outputs\n\n    def post_process(self,\n                     pred,\n                     conf_thres=0.5,\n                     nms_thres=0.9,\n                     class_agnostic=False):\n        box_corner = pred.new(pred.shape)\n        box_corner[:, :, 0] = pred[:, :, 0] - pred[:, :, 2] / 2\n        box_corner[:, :, 1] = pred[:, :, 1] - pred[:, :, 3] / 2\n        box_corner[:, :, 2] = pred[:, :, 0] + pred[:, :, 2] / 2\n        box_corner[:, :, 3] = pred[:, :, 1] + pred[:, :, 3] / 2\n        pred[:, :, :4] = box_corner[:, :, :4]\n\n        output = [None for _ in range(len(pred))]\n        for i, image_pred in enumerate(pred):\n\n            # If none are remaining => process next image\n            if not image_pred.size(0):\n                continue\n            # Get score and class with highest confidence\n            class_conf, class_pred = torch.max(image_pred[:, 5:5 +\n                                                          self.num_classes],\n                                               1,\n                                               keepdim=True)\n\n            conf_mask = (image_pred[:, 4] * class_conf.squeeze() >=\n                         conf_thres).squeeze()\n            # Detections ordered as (x1, y1, x2, y2, obj_conf, class_conf, class_pred)\n            detections = torch.cat(\n                (image_pred[:, :5], class_conf, class_pred.float()), 1)\n            detections = detections[conf_mask]\n            if not detections.size(0):\n                continue\n\n            if class_agnostic:\n                nms_out_index = torchvision.ops.nms(\n                    detections[:, :4],\n                    detections[:, 4] * detections[:, 5],\n                    nms_thres,\n                )\n            else:\n                nms_out_index = torchvision.ops.batched_nms(\n                    detections[:, :4],\n                    detections[:, 4] * detections[:, 5],\n                    detections[:, 6],\n                    nms_thres,\n                )\n            detections = detections[nms_out_index]\n            if output[i] is None:\n                output[i] = detections\n            else:\n                output[i] = torch.cat((output[i], detections))\n        return output\n\n    def preprocess_single(self, img: torch.Tensor):\n        ori_h = img.size(0)\n        ori_w = img.size(1)\n        ratio = min(self.input_h / ori_h, self.input_w / ori_w)\n        # resize\n        resized_img = F.interpolate(img.view(1, 1, ori_h, ori_w),\n                                    mode=\"bilinear\",\n                                    scale_factor=ratio,\n                                    recompute_scale_factor=True)[0, 0]\n        # padding\n        padded_img = torch.full((self.input_h, self.input_w),\n                                114,\n                                dtype=resized_img.dtype,\n                                device='cuda')\n        padded_img[:resized_img.size(0), :resized_img.size(1)] = resized_img\n        # 1 channel --> 3 channels\n        padded_img = padded_img.unsqueeze(-1).expand(-1, -1, 3)\n        # HWC --> CHW\n        padded_img = padded_img.permute(2, 0, 1)\n        padded_img = padded_img.float()\n        return padded_img, resized_img, ratio, ori_h, ori_w\n\n    def detect_single(self, img):\n        padded_img, resized_img, ratio, ori_h, ori_w = self.preprocess_single(\n            img)\n        padded_img = padded_img.unsqueeze(0)\n        output = self.model(padded_img)\n        output = self.decode_outputs(output)\n        # x0, y0, x1, y1, box_conf, cls_conf, cls_id\n        output = self.post_process(output, self.conf_thres, self.nms_thres)[0]\n        if output is not None:\n            output[:, :4] = output[:, :4] / ratio\n            # re-compute: conf = box_conf * cls_conf\n            output[:, 4] = output[:, 4] * output[:, 5]\n            # select box with highest confident\n            output = output[output[:, 4].argmax()]\n            x0 = min(max(int(output[0]), 0), ori_w)\n            y0 = min(max(int(output[1]), 0), ori_h)\n            x1 = min(max(int(output[2]), 0), ori_w)\n            y1 = min(max(int(output[3]), 0), ori_h)\n            area_pct = (x1 - x0) * (y1 - y0) / (ori_h * ori_w)\n            if area_pct >= self.area_pct_thres:\n                # xyxy, area_pct, conf\n                return [x0, y0, x1, y1], area_pct, output[4]\n\n        # if YOLOX fail, try Otsu thresholding + find contours\n        xyxy, area_pct, _ = extract_roi_otsu(\n            resized_img.to(torch.uint8).cpu().numpy())\n        # if both fail, use full frame\n        if xyxy is not None:\n            if area_pct >= self.area_pct_thres:\n                print('ROI detection: using Otsu.')\n                x0, y0, x1, y1 = xyxy\n                x0 = min(max(int(x0 / ratio), 0), ori_w)\n                y0 = min(max(int(y0 / ratio), 0), ori_h)\n                x1 = min(max(int(x1 / ratio), 0), ori_w)\n                y1 = min(max(int(y1 / ratio), 0), ori_h)\n                return [x0, y0, x1, y1], area_pct, None\n        print('ROI detection: both fail.')\n        return None, area_pct, None","metadata":{"execution":{"iopub.status.busy":"2023-06-08T15:52:42.908633Z","iopub.execute_input":"2023-06-08T15:52:42.909036Z","iopub.status.idle":"2023-06-08T15:52:42.92333Z","shell.execute_reply.started":"2023-06-08T15:52:42.909003Z","shell.execute_reply":"2023-06-08T15:52:42.922135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Classification model (4 x ConvNext-small ensemble)","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport torch\nfrom timm.data import resolve_data_config\nfrom timm.models import create_model\nfrom torch import nn\n\n\nclass KFoldEnsembleModel(nn.Module):\n\n    def __init__(self, model_info, ckpt_paths):\n        super(KFoldEnsembleModel, self).__init__()\n        fmodels = []\n        for i, ckpt_path in enumerate(ckpt_paths):\n            print(f'Loading model from {ckpt_path}')\n            fmodel = create_model(\n                model_info['model_name'],\n                num_classes=model_info['num_classes'],\n                in_chans=model_info['in_chans'],\n                pretrained=False,\n                checkpoint_path=ckpt_path,\n                global_pool=model_info['global_pool'],\n            ).eval()\n            data_config = resolve_data_config({}, model=fmodel)\n            print('Data config:', data_config)\n            mean = np.array(data_config['mean']) * 255\n            std = np.array(data_config['std']) * 255\n            print(f'mean={mean}, std={std}')\n            fmodels.append(fmodel)\n        self.fmodels = nn.ModuleList(fmodels)\n\n        self.register_buffer('mean',\n                             torch.FloatTensor(mean).reshape(1, 3, 1, 1))\n        self.register_buffer('std', torch.FloatTensor(std).reshape(1, 3, 1, 1))\n\n    def forward(self, x):\n        #         x = x.sub(self.mean).div(self.std)\n        x = (x - self.mean) / self.std\n        probs = []\n        for fmodel in self.fmodels:\n            logits = fmodel(x)\n            #             prob = logits.softmax(dim=1)[:, 1]\n            prob = logits.sigmoid()[:, 0]\n            probs.append(prob)\n        probs = torch.stack(probs, dim=1)\n        return probs","metadata":{"execution":{"iopub.status.busy":"2023-06-08T15:56:08.790479Z","iopub.execute_input":"2023-06-08T15:56:08.791645Z","iopub.status.idle":"2023-06-08T15:56:08.80678Z","shell.execute_reply.started":"2023-06-08T15:56:08.791594Z","shell.execute_reply":"2023-06-08T15:56:08.805774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import roi_extract\nimportlib.reload(roi_extract)\nimport roi_extract\n\n# global vars\nJ2K_SUID = '1.2.840.10008.1.2.4.90'\nJ2K_HEADER = b\"\\x00\\x00\\x00\\x0C\"\nJLL_SUID = '1.2.840.10008.1.2.4.70'\nJLL_HEADER = b\"\\xff\\xd8\\xff\\xe0\"\nSUID2HEADER = {J2K_SUID: J2K_HEADER, JLL_SUID: JLL_HEADER}\nVOILUT_FUNCS_MAP = {'LINEAR': 0, 'LINEAR_EXACT': 1, 'SIGMOID': 2}\nVOILUT_FUNCS_INV_MAP = {v: k for k, v in VOILUT_FUNCS_MAP.items()}","metadata":{"execution":{"iopub.status.busy":"2023-06-08T15:56:10.055462Z","iopub.execute_input":"2023-06-08T15:56:10.055915Z","iopub.status.idle":"2023-06-08T15:56:10.079365Z","shell.execute_reply.started":"2023-06-08T15:56:10.055876Z","shell.execute_reply":"2023-06-08T15:56:10.07738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configs\n\nMost important configs such as binarization threshold, batch size, ..","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 2\n# binarization threshold for classification\nTHRES = 0.31\nAUTO_THRES = False\nAUTO_THRES_PERCENTILE = 0.97935\n\n# classification model\nUSE_TRT = True\n\n\n# roi detection\nROI_YOLOX_INPUT_SIZE = [416, 416]\nROI_YOLOX_CONF_THRES = 0.5\nROI_YOLOX_NMS_THRES = 0.9\nROI_YOLOX_HW = [(52, 52), (26, 26), (13, 13)]\nROI_YOLOX_STRIDES = [8, 16, 32]\nROI_AREA_PCT_THRES = 0.04\n\n# model\nMODEL_INPUT_SIZE = [2048, 1024]\n\nMODE = 'KAGGLE-TEST'\nassert MODE in ['LOCAL-VAL', 'KAGGLE-VAL', 'KAGGLE-TEST']\n\n# settings corresponding to each mode\nif MODE == 'KAGGLE-VAL':\n    TRT_MODEL_PATH = '/kaggle/input/rsna-breast-cancer-detection-best-ckpts/best_convnext_ensemble_batch2_fp32_torch2trt.engine'\n    TORCH_MODEL_CKPT_PATHS = [\n        f'/kaggle/input/rsna-breast-cancer-detection-best-ckpts/best_convnext_fold_{i}.pth.tar'\n        for i in range(4)\n    ]\n    ROI_YOLOX_ENGINE_PATH = '/kaggle/input/rsna-breast-cancer-detection-best-ckpts/yolox_nano_416_roi_trt_p100.pth'\n    CSV_PATH = '/kaggle/input/rsna-breast-cancer-detection-best-ckpts/_val_fold_0.csv'\n    DCM_ROOT_DIR = '/kaggle/input/rsna-breast-cancer-detection/train_images'\n    SAVE_IMG_ROOT_DIR = '/kaggle/tmp/pngs'\n    N_CHUNKS = 2\n    N_CPUS = 2\n    RM_DONE_CHUNK = False\nelif MODE == 'KAGGLE-TEST':\n    TRT_MODEL_PATH = '/kaggle/input/rsna-breast-cancer-detection-best-ckpts/best_convnext_ensemble_batch2_fp32_torch2trt.engine'\n    TORCH_MODEL_CKPT_PATHS = [\n        f'/kaggle/input/rsna-breast-cancer-detection-best-ckpts/best_convnext_fold_{i}.pth.tar'\n        for i in range(4)\n    ]\n    ROI_YOLOX_ENGINE_PATH = '/kaggle/input/rsna-breast-cancer-detection-best-ckpts/yolox_nano_416_roi_trt_p100.pth'\n    CSV_PATH = '/kaggle/input/rsna-breast-cancer-detection/test.csv'\n    DCM_ROOT_DIR = '/kaggle/input/rsna-breast-cancer-detection/test_images'\n    SAVE_IMG_ROOT_DIR = '/kaggle/tmp/pngs'\n    N_CHUNKS = 2\n    N_CPUS = 2\n    RM_DONE_CHUNK = True\nelif MODE == 'LOCAL-VAL':\n    TRT_MODEL_PATH = './assets/best_convnext_ensemble_batch2_fp32_torch2trt.engine'\n    TORCH_MODEL_CKPT_PATHS = [\n        f'./assets/best_convnext_fold_{i}.pth.tar'\n        for i in range(4)\n    ]\n    ROI_YOLOX_ENGINE_PATH = '../roi_det/YOLOX/YOLOX_outputs/yolox_nano_bre_416/model_trt.pth'\n    CSV_PATH = '../../datasets/cv/v1/val_fold_0.csv'\n    DCM_ROOT_DIR = '../../datasets/train_images/'\n    SAVE_IMG_ROOT_DIR = './temp_save'\n    N_CHUNKS = 2\n    N_CPUS = 2\n    RM_DONE_CHUNK = False","metadata":{"execution":{"iopub.status.busy":"2023-06-10T06:32:16.495867Z","iopub.execute_input":"2023-06-10T06:32:16.496306Z","iopub.status.idle":"2023-06-10T06:32:16.515376Z","shell.execute_reply.started":"2023-06-10T06:32:16.496273Z","shell.execute_reply":"2023-06-10T06:32:16.513557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helpers","metadata":{}},{"cell_type":"markdown","source":"## Dicom metadata","metadata":{}},{"cell_type":"code","source":"class PydicomMetadata:\n\n    def __init__(self, ds):\n        if \"WindowWidth\" not in ds or \"WindowCenter\" not in ds:\n            self.window_widths = []\n            self.window_centers = []\n        else:\n            ww = ds['WindowWidth']\n            wc = ds['WindowCenter']\n            self.window_widths = [float(e) for e in ww\n                                  ] if ww.VM > 1 else [float(ww.value)]\n\n            self.window_centers = [float(e) for e in wc\n                                   ] if wc.VM > 1 else [float(wc.value)]\n\n        # if nan --> LINEAR\n        self.voilut_func = str(ds.get('VOILUTFunction', 'LINEAR')).upper()\n        self.invert = (ds.PhotometricInterpretation == 'MONOCHROME1')\n        assert len(self.window_widths) == len(self.window_centers)\n\n\nclass DicomsdlMetadata:\n\n    def __init__(self, ds):\n        self.window_widths = ds.WindowWidth\n        self.window_centers = ds.WindowCenter\n        if self.window_widths is None or self.window_centers is None:\n            self.window_widths = []\n            self.window_centers = []\n        else:\n            try:\n                if not isinstance(self.window_widths, list):\n                    self.window_widths = [self.window_widths]\n                self.window_widths = [float(e) for e in self.window_widths]\n                if not isinstance(self.window_centers, list):\n                    self.window_centers = [self.window_centers]\n                self.window_centers = [float(e) for e in self.window_centers]\n            except:\n                self.window_widths = []\n                self.window_centers = []\n\n        # if nan --> LINEAR\n        self.voilut_func = ds.VOILUTFunction\n        if self.voilut_func is None:\n            self.voilut_func = 'LINEAR'\n        else:\n            self.voilut_func = str(self.voilut_func).upper()\n        self.invert = (ds.PhotometricInterpretation == 'MONOCHROME1')\n        assert len(self.window_widths) == len(self.window_centers)","metadata":{"execution":{"iopub.status.busy":"2023-06-10T06:32:17.671739Z","iopub.execute_input":"2023-06-10T06:32:17.67248Z","iopub.status.idle":"2023-06-10T06:32:17.688504Z","shell.execute_reply.started":"2023-06-10T06:32:17.672443Z","shell.execute_reply":"2023-06-10T06:32:17.687264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Windowing","metadata":{}},{"cell_type":"code","source":"# slow\n# from pydicom's source\ndef _apply_windowing_np_v1(arr,\n                           window_width=None,\n                           window_center=None,\n                           voi_func='LINEAR',\n                           y_min=0,\n                           y_max=255):\n    assert window_width > 0\n    y_range = y_max - y_min\n    # float64 needed (default) or just float32 ?\n    # arr = arr.astype(np.float64)\n    arr = arr.astype(np.float32)\n\n    if voi_func in ['LINEAR', 'LINEAR_EXACT']:\n        # PS3.3 C.11.2.1.2.1 and C.11.2.1.3.2\n        if voi_func == 'LINEAR':\n            if window_width < 1:\n                raise ValueError(\n                    \"The (0028,1051) Window Width must be greater than or \"\n                    \"equal to 1 for a 'LINEAR' windowing operation\")\n            window_center -= 0.5\n            window_width -= 1\n        below = arr <= (window_center - window_width / 2)\n        above = arr > (window_center + window_width / 2)\n        between = np.logical_and(~below, ~above)\n\n        arr[below] = y_min\n        arr[above] = y_max\n        if between.any():\n            arr[between] = ((\n                (arr[between] - window_center) / window_width + 0.5) * y_range\n                            + y_min)\n    elif voi_func == 'SIGMOID':\n        arr = y_range / (1 +\n                         np.exp(-4 *\n                                (arr - window_center) / window_width)) + y_min\n    else:\n        raise ValueError(\n            f\"Unsupported (0028,1056) VOI LUT Function value '{voi_func}'\")\n    return arr\n\n\ndef _apply_windowing_np_v2(arr,\n                           window_width=None,\n                           window_center=None,\n                           voi_func='LINEAR',\n                           y_min=0,\n                           y_max=255):\n    assert window_width > 0\n    y_range = y_max - y_min\n    # float64 needed (default) or just float32 ?\n    # arr = arr.astype(np.float64)\n    arr = arr.astype(np.float32)\n\n    if voi_func == 'LINEAR' or voi_func == 'LINEAR_EXACT':\n        # PS3.3 C.11.2.1.2.1 and C.11.2.1.3.2\n        if voi_func == 'LINEAR':\n            if window_width < 1:\n                raise ValueError(\n                    \"The (0028,1051) Window Width must be greater than or \"\n                    \"equal to 1 for a 'LINEAR' windowing operation\")\n            window_center -= 0.5\n            window_width -= 1\n\n        # simple trick to improve speed\n        s = y_range / window_width\n        b = (-window_center / window_width + 0.5) * y_range + y_min\n        arr = arr * s + b\n        arr = np.clip(arr, y_min, y_max)\n\n    elif voi_func == 'SIGMOID':\n        # simple trick to improve speed\n        s = -4 / window_width\n        arr = y_range / (1 + np.exp((arr - window_center) * s)) + y_min\n    else:\n        raise ValueError(\n            f\"Unsupported (0028,1056) VOI LUT Function value '{voi_func}'\")\n    return arr\n\n\ndef _apply_windowing_torch(arr,\n                           window_width=None,\n                           window_center=None,\n                           voi_func='LINEAR',\n                           y_min=0,\n                           y_max=255):\n    assert window_width > 0\n    y_range = y_max - y_min\n    # float64 needed (default) or just float32 ?\n    # arr = arr.double()\n    arr = arr.float()\n\n    if voi_func == 'LINEAR' or voi_func == 'LINEAR_EXACT':\n        # PS3.3 C.11.2.1.2.1 and C.11.2.1.3.2\n        if voi_func == 'LINEAR':\n            if window_width < 1:\n                raise ValueError(\n                    \"The (0028,1051) Window Width must be greater than or \"\n                    \"equal to 1 for a 'LINEAR' windowing operation\")\n            window_center -= 0.5\n            window_width -= 1\n\n        # simple trick to improve speed\n        s = y_range / window_width\n        b = (-window_center / window_width + 0.5) * y_range + y_min\n        arr = arr * s + b\n        arr = torch.clamp(arr, y_min, y_max)\n\n    elif voi_func == 'SIGMOID':\n        # simple trick to improve speed\n        s = -4 / window_width\n        arr = y_range / (1 + torch.exp((arr - window_center) * s)) + y_min\n    else:\n        raise ValueError(\n            f\"Unsupported (0028,1056) VOI LUT Function value '{voi_func}'\")\n    return arr\n\n\ndef apply_windowing(arr,\n                    window_width=None,\n                    window_center=None,\n                    voi_func='LINEAR',\n                    y_min=0,\n                    y_max=255,\n                    backend='np_v2'):\n    if backend == 'torch':\n        if isinstance(arr, torch.Tensor):\n            pass\n        elif isinstance(arr, np.ndarray):\n            if arr.dtype == np.uint16:\n                arr = torch.from_numpy(arr, torch.int16)\n            else:\n                arr = torch.from_numpy(arr)\n\n    if backend == 'np_v1':\n        windowing_func = _apply_windowing_np_v1\n    elif backend == 'np_v2':\n        windowing_func = _apply_windowing_np_v2\n    elif backend == 'torch':\n        windowing_func = _apply_windowing_torch\n    else:\n        raise ValueError(\n            f'Invalid backend {backend}, must be one of [\"np\", \"np_v2\", \"torch\"]'\n        )\n\n    arr = windowing_func(arr,\n                         window_width=window_width,\n                         window_center=window_center,\n                         voi_func=voi_func,\n                         y_min=y_min,\n                         y_max=y_max)\n    return arr","metadata":{"execution":{"iopub.status.busy":"2023-06-10T06:32:18.072781Z","iopub.execute_input":"2023-06-10T06:32:18.073161Z","iopub.status.idle":"2023-06-10T06:32:18.101Z","shell.execute_reply.started":"2023-06-10T06:32:18.07313Z","shell.execute_reply":"2023-06-10T06:32:18.099942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Others","metadata":{}},{"cell_type":"code","source":"def min_max_scale(img):\n    maxv = img.max()\n    minv = img.min()\n    if maxv > minv:\n        return (img - minv) / (maxv - minv)\n    else:\n        return img - minv  # ==0\n\n\n#@TODO: percentile on both min-max?\n# this version is not correctly implemented, but used in the winning submission\ndef percentile_min_max_scale(img, pct=99):\n    if isinstance(img, np.ndarray):\n        maxv = np.percentile(img, pct) - 1\n        minv = img.min()\n        assert maxv >= minv\n        if maxv > minv:\n            ret = (img - minv) / (maxv - minv)\n        else:\n            ret = img - minv  # ==0\n        ret = np.clip(ret, 0, 1)\n    elif isinstance(img, torch.Tensor):\n        maxv = torch.quantile(img, pct / 100) - 1\n        minv = img.min()\n        assert maxv >= minv\n        if maxv > minv:\n            ret = (img - minv) / (maxv - minv)\n        else:\n            ret = img - minv  # ==0\n        ret = torch.clamp(ret, 0, 1)\n    else:\n        raise ValueError(\n            'Invalid img type, should be numpy array or torch.Tensor')\n    return ret\n\n\ndef resize_and_pad(img, input_size=MODEL_INPUT_SIZE):\n    input_h, input_w = input_size\n    ori_h, ori_w = img.shape[:2]\n    ratio = min(input_h / ori_h, input_w / ori_w)\n    # resize\n    img = F.interpolate(img.view(1, 1, ori_h, ori_w),\n                        mode=\"bilinear\",\n                        scale_factor=ratio,\n                        recompute_scale_factor=True)[0, 0]\n    # padding\n    padded_img = torch.zeros((input_h, input_w),\n                             dtype=img.dtype,\n                             device='cuda')\n    cur_h, cur_w = img.shape\n    y_start = (input_h - cur_h) // 2\n    x_start = (input_w - cur_w) // 2\n    padded_img[y_start:y_start + cur_h, x_start:x_start + cur_w] = img\n    padded_img = padded_img.unsqueeze(-1).expand(-1, -1, 3)\n    return padded_img\n\n\ndef save_img_to_file(save_path, img, backend='cv2'):\n    file_ext = os.path.basename(save_path).split('.')[-1]\n    if backend == 'cv2':\n        if img.dtype == np.uint16:\n            # https://docs.opencv.org/3.4/d4/da8/group__imgcodecs.html#gabbc7ef1aa2edfaa87772f1202d67e0ce\n            assert file_ext in ['png', 'jp2', 'tiff', 'tif']\n            cv2.imwrite(save_path, img)\n        elif img.dtype == np.uint8:\n            cv2.imwrite(save_path, img)\n        else:\n            raise ValueError(\n                '`cv2` backend only support uint8 or uint16 images.')\n    elif backend == 'np':\n        assert file_ext == 'npy'\n        np.save(save_path, img)\n    else:\n        raise ValueError(f'Unsupported backend `{backend}`.')\n\n\ndef load_img_from_file(img_path, backend='cv2'):\n    if backend == 'cv2':\n        return cv2.imread(img_path, cv2.IMREAD_ANYDEPTH)\n    elif backend == 'np':\n        return np.load(img_path)\n    else:\n        raise ValueError()\n        \n\ndef make_uid_transfer_dict(df, dcm_root_dir):\n    machine_id_to_transfer = {}\n    machine_id = df.machine_id.unique()\n    for i in machine_id:\n        row = df[df.machine_id == i].iloc[0]\n        sample_dcm_path = os.path.join(dcm_root_dir, str(row.patient_id),\n                                       f'{row.image_id}.dcm')\n        dicom = pydicom.dcmread(sample_dcm_path)\n        machine_id_to_transfer[i] = dicom.file_meta.TransferSyntaxUID\n    return machine_id_to_transfer","metadata":{"execution":{"iopub.status.busy":"2023-06-10T06:32:18.530996Z","iopub.execute_input":"2023-06-10T06:32:18.531917Z","iopub.status.idle":"2023-06-10T06:32:18.55414Z","shell.execute_reply.started":"2023-06-10T06:32:18.531866Z","shell.execute_reply":"2023-06-10T06:32:18.553054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dicom decoding with DALI or dicomsdl\n\nHelpers/Utilizations for dicom decoding and further preprocessing","metadata":{}},{"cell_type":"code","source":"# DALI patch for INT16 support\n################################################################################\nDALI2TORCH_TYPES = {\n    types.DALIDataType.FLOAT: torch.float32,\n    types.DALIDataType.FLOAT64: torch.float64,\n    types.DALIDataType.FLOAT16: torch.float16,\n    types.DALIDataType.UINT8: torch.uint8,\n    types.DALIDataType.INT8: torch.int8,\n    types.DALIDataType.UINT16: torch.int16,\n    types.DALIDataType.INT16: torch.int16,\n    types.DALIDataType.INT32: torch.int32,\n    types.DALIDataType.INT64: torch.int64\n}\n\nTORCH_DTYPES = {\n    'uint8': torch.uint8,\n    'float16': torch.float16,\n    'float32': torch.float32,\n    'float64': torch.float64,\n}\n\n\n# @TODO: dangerous to copy from UINT16 to INT16 (memory layout?)\n# little/big endian ?\n# @TODO: faster reuse memory without copying: https://github.com/NVIDIA/DALI/issues/4126\ndef feed_ndarray(dali_tensor, arr, cuda_stream=None):\n    \"\"\"\n    Copy contents of DALI tensor to PyTorch's Tensor.\n\n    Parameters\n    ----------\n    `dali_tensor` : nvidia.dali.backend.TensorCPU or nvidia.dali.backend.TensorGPU\n                    Tensor from which to copy\n    `arr` : torch.Tensor\n            Destination of the copy\n    `cuda_stream` : torch.cuda.Stream, cudaStream_t or any value that can be cast to cudaStream_t.\n                    CUDA stream to be used for the copy\n                    (if not provided, an internal user stream will be selected)\n                    In most cases, using pytorch's current stream is expected (for example,\n                    if we are copying to a tensor allocated with torch.zeros(...))\n    \"\"\"\n    dali_type = DALI2TORCH_TYPES[dali_tensor.dtype]\n\n    assert dali_type == arr.dtype, (\n        \"The element type of DALI Tensor/TensorList\"\n        \" doesn't match the element type of the target PyTorch Tensor: \"\n        \"{} vs {}\".format(dali_type, arr.dtype))\n    assert dali_tensor.shape() == list(arr.size()), \\\n        (\"Shapes do not match: DALI tensor has size {0}, but PyTorch Tensor has size {1}\".\n            format(dali_tensor.shape(), list(arr.size())))\n    cuda_stream = types._raw_cuda_stream(cuda_stream)\n\n    # turn raw int to a c void pointer\n    c_type_pointer = ctypes.c_void_p(arr.data_ptr())\n    if isinstance(dali_tensor, (TensorGPU, TensorListGPU)):\n        stream = None if cuda_stream is None else ctypes.c_void_p(cuda_stream)\n        dali_tensor.copy_to_external(c_type_pointer, stream, non_blocking=True)\n    else:\n        dali_tensor.copy_to_external(c_type_pointer)\n    return arr\n\n\nclass _JStreamExternalSource:\n    \"\"\"DALI External Source for in-memory dicom decoding\"\"\"\n\n    def __init__(self, dcm_paths, batch_size=1):\n        self.dcm_paths = dcm_paths\n        self.len = len(dcm_paths)\n        self.batch_size = batch_size\n\n    def __call__(self, batch_info):\n        idx = batch_info.iteration\n        # print('IDX:', batch_info.iteration, batch_info.epoch_idx)\n        start = idx * self.batch_size\n        end = min(self.len, start + self.batch_size)\n        if end <= start:\n            raise StopIteration()\n\n        batch_dcm_paths = self.dcm_paths[start:end]\n        j_streams = []\n        inverts = []\n        windowing_params = []\n        voilut_funcs = []\n\n        for dcm_path in batch_dcm_paths:\n            ds = pydicom.dcmread(dcm_path)\n            pixel_data = ds.PixelData\n            offset = pixel_data.find(\n                SUID2HEADER[ds.file_meta.TransferSyntaxUID])\n            j_stream = np.array(bytearray(pixel_data[offset:]), np.uint8)\n            invert = (ds.PhotometricInterpretation == 'MONOCHROME1')\n            meta = PydicomMetadata(ds)\n            windowing_param = np.array(\n                [meta.window_centers, meta.window_widths], np.float16)\n            voilut_func = VOILUT_FUNCS_MAP[meta.voilut_func]\n            j_streams.append(j_stream)\n            inverts.append(invert)\n            windowing_params.append(windowing_param)\n            voilut_funcs.append(voilut_func)\n        return j_streams, np.array(inverts, dtype=np.bool_), \\\n            windowing_params, np.array(voilut_funcs, dtype=np.uint8)\n\n\n@dali.pipeline_def\ndef _dali_pipeline(eii):\n    jpeg, invert, windowing_param, voilut_func = dali.fn.external_source(\n        source=eii,\n        num_outputs=4,\n        dtype=[\n            dali.types.UINT8, dali.types.BOOL, dali.types.FLOAT16,\n            dali.types.UINT8\n        ],\n        batch=True,\n        batch_info=True,\n        parallel=True)\n    ori_img = dali.fn.experimental.decoders.image(\n        jpeg,\n        device='mixed',\n        output_type=dali.types.ANY_DATA,\n        dtype=dali.types.UINT16)\n    return ori_img, invert, windowing_param, voilut_func\n\n\ndef decode_crop_save_dali(roi_yolox_engine_path,\n                          dcm_paths,\n                          save_paths,\n                          save_backend='cv2',\n                          batch_size=1,\n                          num_threads=1,\n                          py_num_workers=1,\n                          py_start_method='fork',\n                          device_id=0):\n    \"\"\"DALI dicom decoding --> ROI cropping --> norm --> save as 8-bits PNG\"\"\"\n    \n    assert len(dcm_paths) == len(save_paths)\n    assert save_backend in ['cv2', 'np']\n    num_dcms = len(dcm_paths)\n\n    # dali to process with chunk in-memory\n    external_source = _JStreamExternalSource(dcm_paths, batch_size=batch_size)\n    pipe = _dali_pipeline(\n        external_source,\n        py_num_workers=py_num_workers,\n        py_start_method=py_start_method,\n        batch_size=batch_size,\n        num_threads=num_threads,\n        device_id=device_id,\n        debug=False,\n    )\n    pipe.build()\n\n    roi_extractor = roi_extract.RoiExtractor(engine_path=roi_yolox_engine_path,\n                                             input_size=ROI_YOLOX_INPUT_SIZE,\n                                             num_classes=1,\n                                             conf_thres=ROI_YOLOX_CONF_THRES,\n                                             nms_thres=ROI_YOLOX_NMS_THRES,\n                                             class_agnostic=False,\n                                             area_pct_thres=ROI_AREA_PCT_THRES,\n                                             hw=ROI_YOLOX_HW,\n                                             strides=ROI_YOLOX_STRIDES,\n                                             exp=None)\n    print('ROI extractor (YOLOX) loaded!')\n\n    num_batchs = num_dcms // batch_size\n    last_batch_size = batch_size\n    if num_dcms % batch_size > 0:\n        num_batchs += 1\n        last_batch_size = num_dcms % batch_size\n\n    cur_idx = -1\n    for _batch_idx in tqdm(range(num_batchs)):\n        try:\n            outs = pipe.run()\n        except Exception as e:\n            #             print('DALI exception occur:', e)\n            print(\n                f'Exception: One of {dcm_paths[_batch_idx * batch_size: (_batch_idx + 1) * batch_size]} can not be decoded.'\n            )\n            # ignore this batch and re-build pipeline\n            if _batch_idx < num_batchs - 1:\n                cur_idx += batch_size\n                del external_source, pipe\n                gc.collect()\n                torch.cuda.empty_cache()\n                external_source = _JStreamExternalSource(\n                    dcm_paths[(_batch_idx + 1) * batch_size:],\n                    batch_size=batch_size)\n                pipe = _dali_pipeline(\n                    external_source,\n                    py_num_workers=py_num_workers,\n                    py_start_method=py_start_method,\n                    batch_size=batch_size,\n                    num_threads=num_threads,\n                    device_id=device_id,\n                    debug=False,\n                )\n                pipe.build()\n            else:\n                cur_idx += last_batch_size\n            continue\n\n        imgs = outs[0]\n        inverts = outs[1]\n        windowing_params = outs[2]\n        voilut_funcs = outs[3]\n        for j in range(len(inverts)):\n            cur_idx += 1\n            save_path = save_paths[cur_idx]\n            img_dali = imgs[j]\n            img_torch = torch.empty(img_dali.shape(),\n                                    dtype=torch.int16,\n                                    device='cuda')\n            feed_ndarray(img_dali,\n                         img_torch,\n                         cuda_stream=torch.cuda.current_stream(device=0))\n            # @TODO: test whether copy uint16 to int16 pointer is safe in this case\n            if 0:\n                img_np = img_dali.as_cpu().squeeze(-1)  # uint16\n                print(type(img_np), img_np.shape)\n                img_np = torch.from_numpy(img_np, dtype=torch.int16)\n                diff = torch.max(torch.abs(img_np - img_torch))\n                assert diff == 0, f'{img_torch.shape}, {img_np.shape}, {diff}'\n\n            invert = inverts.at(j).item()\n            windowing_param = windowing_params.at(j)\n            voilut_func = voilut_funcs.at(j).item()\n            voilut_func = VOILUT_FUNCS_INV_MAP[voilut_func]\n\n            # YOLOX for ROI extraction\n            img_yolox = min_max_scale(img_torch)\n            img_yolox = (img_yolox * 255)  # float32\n            if invert:\n                img_yolox = 255 - img_yolox\n            # YOLOX infer\n            # who know if exception happen in hidden test ?\n            try:\n                xyxy, _area_pct, _conf = roi_extractor.detect_single(img_yolox)\n                if xyxy is not None:\n                    x0, y0, x1, y1 = xyxy\n                    crop = img_torch[y0:y1, x0:x1]\n                else:\n                    crop = img_torch\n            except:\n                print('ROI extract exception!')\n                crop = img_torch\n\n            # apply windowing\n            if windowing_param.shape[1] != 0:\n                default_window_center = windowing_param[0, 0]\n                default_window_width = windowing_param[1, 0]\n                crop = apply_windowing(crop,\n                                       window_width=default_window_width,\n                                       window_center=default_window_center,\n                                       voi_func=voilut_func,\n                                       y_min=0,\n                                       y_max=255,\n                                       backend='torch')\n            # if no window center/width in dcm file\n            # do simple min-max scaling\n            else:\n                print('No windowing param!')\n                crop = min_max_scale(crop)\n                crop = crop * 255\n            if invert:\n                crop = 255 - crop\n            crop = resize_and_pad(crop, MODEL_INPUT_SIZE)\n            crop = crop.to(torch.uint8)\n            crop = crop.cpu().numpy()\n            save_img_to_file(save_path, crop, backend=save_backend)\n\n\n#     assert cur_idx == len(\n#         save_paths) - 1, f'{cur_idx} != {len(save_paths) - 1}'\n    try:\n        del external_source, pipe, roi_extractor\n    except:\n        pass\n    gc.collect()\n    torch.cuda.empty_cache()\n    return\n\n\ndef decode_and_save_dali_parallel(\n        roi_yolox_engine_path,\n        dcm_paths,\n        save_paths,\n        save_backend='cv2',\n        batch_size=1,\n        num_threads=1,\n        py_num_workers=1,\n        py_start_method='fork',\n        device_id=0,\n        parallel_n_jobs=2,\n        parallel_n_chunks=4,\n        parallel_backend='joblib',  # joblib or multiprocessing\n        joblib_backend='loky'):\n    assert parallel_backend in ['joblib', 'multiprocessing']\n    assert joblib_backend in ['threading', 'multiprocessing', 'loky']\n    # py_num_workers > 0 means using multiprocessing worker\n    # 'fork' multiprocessing after CUDA init is not work (we must use 'spawn' instead)\n    # since our pipeline can be re-build (when a dicom can't be decoded on GPU),\n    # 2 options:\n    #       (py_num_workers = 0, py_start_method=?)\n    #       (py_num_workers > 0, py_start_method = 'spawn')\n    assert not (py_num_workers > 0 and py_start_method == 'fork')\n\n    if parallel_n_jobs == 1:\n        print('No parralel. Starting the tasks within current process.')\n        return decode_crop_save_dali(roi_yolox_engine_path,\n                                     dcm_paths,\n                                     save_paths,\n                                     save_backend=save_backend,\n                                     batch_size=batch_size,\n                                     num_threads=num_threads,\n                                     py_num_workers=py_num_workers,\n                                     py_start_method=py_start_method,\n                                     device_id=device_id)\n    else:\n        num_samples = len(dcm_paths)\n        num_samples_per_chunk = num_samples // parallel_n_chunks\n        if num_samples % parallel_n_chunks > 0:\n            num_samples_per_chunk += 1\n        starts = [num_samples_per_chunk * i for i in range(parallel_n_chunks)]\n        ends = [\n            min(start + num_samples_per_chunk, num_samples) for start in starts\n        ]\n        if isinstance(device_id, list):\n            assert len(device_id) == parallel_n_chunks\n        elif isinstance(device_id, int):\n            device_id = [device_id] * parallel_n_chunks\n\n        print(\n            f'Starting {parallel_n_jobs} jobs with backend `{parallel_backend}`, {parallel_n_chunks} chunks ...'\n        )\n        if parallel_backend == 'joblib':\n            _ = Parallel(n_jobs=parallel_n_jobs, backend=joblib_backend)(\n                delayed(decode_crop_save_dali)(\n                    roi_yolox_engine_path,\n                    dcm_paths[start:end],\n                    save_paths[start:end],\n                    save_backend=save_backend,\n                    batch_size=batch_size,\n                    num_threads=num_threads,\n                    py_num_workers=py_num_workers,  # ram_v3\n                    py_start_method=py_start_method,\n                    device_id=worker_device_id,\n                ) for start, end, worker_device_id in zip(\n                    starts, ends, device_id))\n        else:  # manually start multiprocessing's processes\n            workers = []\n            daemon = False if py_num_workers > 0 else True\n            for i in range(parallel_n_jobs):\n                start = starts[i]\n                end = ends[i]\n                worker_device_id = device_id[i]\n                worker = mp.Process(group=None,\n                                    target=decode_crop_save_dali,\n                                    args=(\n                                        roi_yolox_engine_path,\n                                        dcm_paths[start:end],\n                                        save_paths[start:end],\n                                    ),\n                                    kwargs={\n                                        'save_backend': save_backend,\n                                        'batch_size': batch_size,\n                                        'num_threads': num_threads,\n                                        'py_num_workers': py_num_workers,\n                                        'py_start_method': py_start_method,\n                                        'device_id': worker_device_id,\n                                    },\n                                    daemon=daemon)\n                workers.append(worker)\n            for worker in workers:\n                worker.start()\n            for worker in workers:\n                worker.join()\n    return\n\n\ndef _single_decode_crop_save_sdl(roi_extractor,\n                                 dcm_path,\n                                 save_path,\n                                 save_backend='cv2',\n                                 index=0):\n    dcm = dicomsdl.open(dcm_path)\n    meta = DicomsdlMetadata(dcm)\n    info = dcm.getPixelDataInfo()\n    if info['SamplesPerPixel'] != 1:\n        raise RuntimeError('SamplesPerPixel != 1')\n    else:\n        shape = [info['Rows'], info['Cols']]\n\n    ori_dtype = info['dtype']\n    img = np.empty(shape, dtype=ori_dtype)\n    dcm.copyFrameData(index, img)\n    img_torch = torch.from_numpy(img.astype(np.int16)).cuda()\n\n    # YOLOX for ROI extraction\n    img_yolox = min_max_scale(img_torch)\n    img_yolox = (img_yolox * 255)  # float32\n    # @TODO: subtract on large array --> should move after F.interpolate()\n    if meta.invert:\n        img_yolox = 255 - img_yolox\n    # YOLOX infer\n    try:\n        xyxy, _area_pct, _conf = roi_extractor.detect_single(img_yolox)\n        if xyxy is not None:\n            x0, y0, x1, y1 = xyxy\n            crop = img_torch[y0:y1, x0:x1]\n        else:\n            crop = img_torch\n    except:\n        print('ROI extract exception!')\n        crop = img_torch\n\n    # apply voi lut\n    if meta.window_widths:\n        crop = apply_windowing(crop,\n                               window_width=meta.window_widths[0],\n                               window_center=meta.window_centers[0],\n                               voi_func=meta.voilut_func,\n                               y_min=0,\n                               y_max=255,\n                               backend='torch')\n    else:\n        print('No windowing param!')\n        crop = min_max_scale(crop)\n        crop = crop * 255\n\n    if meta.invert:\n        crop = 255 - crop\n    crop = resize_and_pad(crop, MODEL_INPUT_SIZE)\n    crop = crop.to(torch.uint8)\n    crop = crop.cpu().numpy()\n    save_img_to_file(save_path, crop, backend=save_backend)\n\n\ndef decode_crop_save_sdl(roi_yolox_engine_path,\n                         dcm_paths,\n                         save_paths,\n                         save_backend='cv2'):\n    \"\"\"DicomSDL decoding --> ROI cropping --> norm --> save as 8-bits PNG\"\"\"\n    \n    assert len(dcm_paths) == len(save_paths)\n    roi_detector = roi_extract.RoiExtractor(engine_path=roi_yolox_engine_path,\n                                            input_size=ROI_YOLOX_INPUT_SIZE,\n                                            num_classes=1,\n                                            conf_thres=ROI_YOLOX_CONF_THRES,\n                                            nms_thres=ROI_YOLOX_NMS_THRES,\n                                            class_agnostic=False,\n                                            area_pct_thres=ROI_AREA_PCT_THRES,\n                                            hw=ROI_YOLOX_HW,\n                                            strides=ROI_YOLOX_STRIDES,\n                                            exp=None)\n    print('ROI extractor (YOLOX) loaded!')\n    for i in tqdm(range(len(dcm_paths))):\n        _single_decode_crop_save_sdl(roi_detector, dcm_paths[i], save_paths[i],\n                                     save_backend)\n\n    del roi_detector\n    gc.collect()\n    torch.cuda.empty_cache()\n    return\n\n\ndef decode_crop_save_sdl_parallel(roi_yolox_engine_path,\n                                  dcm_paths,\n                                  save_paths,\n                                  save_backend='cv2',\n                                  parallel_n_jobs=2,\n                                  parallel_n_chunks=4,\n                                  joblib_backend='loky'):\n    assert len(dcm_paths) == len(save_paths)\n    if parallel_n_jobs == 1:\n        print('No parralel. Starting the tasks within current process.')\n        return decode_crop_save_sdl(roi_yolox_engine_path, dcm_paths,\n                                    save_paths, save_backend)\n    else:\n        num_samples = len(dcm_paths)\n        num_samples_per_chunk = num_samples // parallel_n_chunks\n        if num_samples % parallel_n_chunks > 0:\n            num_samples_per_chunk += 1\n        starts = [num_samples_per_chunk * i for i in range(parallel_n_chunks)]\n        ends = [\n            min(start + num_samples_per_chunk, num_samples) for start in starts\n        ]\n\n        print(\n            f'Starting {parallel_n_jobs} jobs with backend `{joblib_backend}`, {parallel_n_chunks} chunks...'\n        )\n        _ = Parallel(n_jobs=parallel_n_jobs, backend=joblib_backend)(\n            delayed(decode_crop_save_sdl)(roi_yolox_engine_path,\n                                          dcm_paths[start:end],\n                                          save_paths[start:end], save_backend)\n            for start, end in zip(starts, ends))","metadata":{"execution":{"iopub.status.busy":"2023-06-10T06:32:19.52678Z","iopub.execute_input":"2023-06-10T06:32:19.527259Z","iopub.status.idle":"2023-06-10T06:32:19.608738Z","shell.execute_reply.started":"2023-06-10T06:32:19.52721Z","shell.execute_reply":"2023-06-10T06:32:19.607754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset & Dataloader","metadata":{}},{"cell_type":"code","source":"class ValTransform:\n\n    def __init__(self):\n        self.transform_fn = A.Compose([ToTensorV2(transpose_mask=True)])\n\n    def __call__(self, img):\n        return self.transform_fn(image=img)['image']\n\n\nclass RSNADataset(Dataset):\n\n    def __init__(self, df, img_root_dir, transform_fn=None):\n        self.img_paths = []\n        self.transform_fn = transform_fn\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(img_root_dir, img_name)\n            self.img_paths.append(img_path)\n        print(f'Done loading dataset with {len(self.img_paths)} samples.')\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        img = cv2.imread(img_path)\n        if img is None:\n            print('ERROR:', img_path)\n        if self.transform_fn:\n            img = self.transform_fn(img)\n        return img\n\n    def get_df(self):\n        return self.df","metadata":{"execution":{"iopub.status.busy":"2023-06-10T06:32:20.195085Z","iopub.execute_input":"2023-06-10T06:32:20.195478Z","iopub.status.idle":"2023-06-10T06:32:20.207713Z","shell.execute_reply.started":"2023-06-10T06:32:20.195445Z","shell.execute_reply":"2023-06-10T06:32:20.206721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Main\n- Preprocessing\n    + Decode dicom (jpeg)\n    + ROI cropping\n    + Normalization\n    + Save to disk as 8-bits PNG\n- Inference\n- Post-processing","metadata":{}},{"cell_type":"markdown","source":"## Preprocessing + inference in chunks","metadata":{}},{"cell_type":"code","source":"######################################################\n# MAIN CODE\n\nif MODE == 'KAGGLE-TEST':\n    global_df = pd.read_csv(CSV_PATH)\nelse:\n    global_df = pd.read_csv(CSV_PATH)[:500]\n    \nMACHINE_TO_SUID = make_uid_transfer_dict(global_df, DCM_ROOT_DIR)\nall_patients = list(global_df.patient_id.unique())\nnum_patients = len(all_patients)\n\n# Processing in chunk to prevent disk overflow while saving PNGs\nnum_patients_per_chunk = num_patients // N_CHUNKS + 1\nall_chunk_patients = [\n    all_patients[num_patients_per_chunk * i:num_patients_per_chunk * (i + 1)]\n    for i in range(N_CHUNKS)\n]\nprint(f'PATIENT CHUNKS: {[len(c) for c in all_chunk_patients]}')\n\npred_dfs = []\nfor chunk_idx, chunk_patients in enumerate(all_chunk_patients):\n    os.makedirs(SAVE_IMG_ROOT_DIR, exist_ok=True)\n    df = global_df[global_df.patient_id.isin(chunk_patients)].reset_index(\n        drop=True)\n    print(\n        f'Processing chunk {chunk_idx} with {len(chunk_patients)} patients, {len(df)} images'\n    )\n    if len(df) == 0:\n        continue\n    dcm_paths = []\n    save_paths = []\n    dali_dcm_paths = []\n    dali_save_paths = []\n    for i in range(len(df)):\n        patient_id = df.at[i, 'patient_id']\n        image_id = df.at[i, 'image_id']\n        suid = MACHINE_TO_SUID[df.at[i, 'machine_id']]\n        dcm_path = os.path.join(DCM_ROOT_DIR, str(patient_id),\n                                f'{image_id}.dcm')\n        save_path = os.path.join(SAVE_IMG_ROOT_DIR,\n                                 f'{patient_id}@{image_id}.png')\n        # if os.path.isfile(save_path):\n        #     continue\n        dcm_paths.append(dcm_path)\n        save_paths.append(save_path)\n        if suid == J2K_SUID or suid == JLL_SUID:\n            dali_dcm_paths.append(dcm_path)\n            dali_save_paths.append(save_path)\n            \n    # save images to disk as 8-bits PNG\n    if 1:\n        t0 = time.time()\n        # try to decode all .90 and .70 with DALI\n        decode_and_save_dali_parallel(\n            ROI_YOLOX_ENGINE_PATH,\n            dali_dcm_paths,\n            dali_save_paths,\n            save_backend='cv2',\n            batch_size=1,\n            num_threads=1,\n            py_num_workers=0,\n            py_start_method='fork',\n            device_id=0,\n            parallel_n_jobs=N_CPUS + 1,\n            parallel_n_chunks = N_CPUS + 1,\n            parallel_backend='joblib',  # joblib or multiprocessing\n            joblib_backend='loky')\n        gc.collect()\n        torch.cuda.empty_cache()\n        t1 = time.time()\n        print(f'DALI done in {t1 - t0} sec')\n\n\n        # CPU decode all others (exceptions) with dicomsdl\n        done_img_names = os.listdir(SAVE_IMG_ROOT_DIR)\n        save_img_names = [os.path.basename(p) for p in save_paths]\n        remain_img_names = list(set(save_img_names) - set(done_img_names))\n        remain_img_paths = [\n            os.path.join(SAVE_IMG_ROOT_DIR, name)\n            for name in remain_img_names\n        ]\n        remain_dcm_paths = []\n        for name in remain_img_names:\n            patient_id, image_id = os.path.basename(name).split(\n                '.')[0].split('@')\n            remain_dcm_paths.append(\n                os.path.join(DCM_ROOT_DIR, patient_id, f'{image_id}.dcm'))\n        num_remain = len(remain_dcm_paths)\n        print(f'Number of undecoded files: {num_remain}')\n        #         print(f'Remains: {remain_img_names}')\n        if num_remain > 0:\n            # 16 or just any > 0 number\n            if num_remain > 32 * N_CPUS:\n                sdl_n_jobs = N_CPUS\n                sdl_n_chunks = N_CPUS\n            else:\n                sdl_n_jobs = 1\n                sdl_n_chunks = 1\n            decode_crop_save_sdl_parallel(ROI_YOLOX_ENGINE_PATH,\n                                          remain_dcm_paths,\n                                          remain_img_paths,\n                                          save_backend='cv2',\n                                          parallel_n_jobs=sdl_n_jobs,\n                                          parallel_n_chunks=sdl_n_chunks,\n                                          joblib_backend='loky')\n            gc.collect()\n            torch.cuda.empty_cache()\n        else:\n            print('No remain files to decode.')\n        t2 = time.time()\n        print(f'SDL done in { t2 - t1} sec')\n        print(f'TOTAL DECODING TIME: {t2 - t0} sec')\n\n    # loading data\n    dataset = RSNADataset(df, SAVE_IMG_ROOT_DIR, transform_fn=ValTransform())\n    dataloader = DataLoader(\n        dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        num_workers=2,\n        pin_memory=True,\n        drop_last=False,\n    )\n    \n    # load model\n    if USE_TRT:\n        model = TRTModule()\n        assert os.path.isfile(TRT_MODEL_PATH)\n        model.load_state_dict(torch.load(TRT_MODEL_PATH))\n    else:\n        model_info = {\n            'model_name': 'convnext_small.fb_in22k_ft_in1k_384',\n            'num_classes': 1,\n            'in_chans': 3,\n            'global_pool': 'max',\n        }\n        model = KFoldEnsembleModel(model_info, TORCH_MODEL_CKPT_PATHS)\n        model.eval()\n        model.cuda()\n\n    # inference\n    all_probs = []\n    with torch.inference_mode():\n        for batch in tqdm(dataloader):\n            batch = batch.cuda().float()\n            probs = model(batch)\n            probs = probs.cpu().numpy()\n            all_probs.append(probs)\n    \n    # N * num_models\n    all_probs = np.concatenate(all_probs, axis=0)\n    all_probs = np.nan_to_num(all_probs, nan=0.0, posinf=None, neginf=None)\n    # simple avg for ensemble to get per-sample prediction\n    all_probs = all_probs.mean(axis=-1)\n    assert all_probs.shape[0] == len(df)\n\n    df['preds'] = all_probs\n    pred_dfs.append(df)\n    print(f'DONE CHUNK {chunk_idx} with {len(df)} samples')\n    del model\n    gc.collect()\n    torch.cuda.empty_cache()\n    if RM_DONE_CHUNK:\n        shutil.rmtree(SAVE_IMG_ROOT_DIR)\n        print(f'Removed save image directory {SAVE_IMG_ROOT_DIR}')\n    print('-----------------------------\\n\\n')","metadata":{"execution":{"iopub.status.busy":"2023-06-10T06:32:35.373697Z","iopub.execute_input":"2023-06-10T06:32:35.374081Z","iopub.status.idle":"2023-06-10T06:33:27.240127Z","shell.execute_reply.started":"2023-06-10T06:32:35.374048Z","shell.execute_reply":"2023-06-10T06:33:27.238825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Post-processing & Submit\n\nNote that ConvNext's raw classification probability is not calibrated, hence metrics based on absolute probability prediction (e.g pF1) will be affected by label smoothing, etc.","metadata":{}},{"cell_type":"code","source":"pred_df = pd.concat(pred_dfs).reset_index(drop=True)\nif 'prediction_id' not in pred_df.columns:\n    pred_df['prediction_id'] = pred_df.apply(lambda row: str(row.patient_id) + '_' + row.laterality, axis = 1)\nsubmit_df = pred_df[['prediction_id', 'preds']]\n\n# Simple avg for per-breast prediction\nsubmit_df = pred_df.groupby('prediction_id').mean()\n\n# # mean of top-3\n# submit_df = submit_df.groupby('prediction_id')['preds'].nlargest(3).mean(level = 0).to_frame()\n\n# every one hacked the metric to binary F1-score\nif AUTO_THRES:\n    thres = np.quantile(submit_df['preds'].values, AUTO_THRES_PERCENTILE)\nelse:\n    thres = THRES\nsubmit_df['cancer'] = (submit_df['preds'].values > thres).astype(int)\nsubmit_df = submit_df['cancer']\nsubmit_df.to_csv('submission.csv')\nsubmit_df","metadata":{"execution":{"iopub.status.busy":"2023-06-10T06:34:36.61817Z","iopub.execute_input":"2023-06-10T06:34:36.61856Z","iopub.status.idle":"2023-06-10T06:34:36.640549Z","shell.execute_reply.started":"2023-06-10T06:34:36.618525Z","shell.execute_reply":"2023-06-10T06:34:36.639543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation\nValidate on val data if needed","metadata":{}},{"cell_type":"code","source":"if MODE == 'KAGGLE-TEST':\n    pass\nelse:\n    pred_df['targets'] = pred_df['cancer']\n    pred_df.to_csv('prediction.csv')\n    metrics = compute_metrics(pred_df,\n                              plot_save_path='metric_plot.png',\n                              thres_range=(0, 1, 0.01),\n                              sort_by='pfbeta',\n                              additional_info=True)\n    print('METRICS:', metrics)","metadata":{"execution":{"iopub.status.busy":"2023-06-10T06:33:27.260833Z","iopub.execute_input":"2023-06-10T06:33:27.261472Z","iopub.status.idle":"2023-06-10T06:33:27.277603Z","shell.execute_reply.started":"2023-06-10T06:33:27.261434Z","shell.execute_reply":"2023-06-10T06:33:27.276539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}