{"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":"code","source":"#https://www.kaggle.com/code/tivfrvqhs5/torch-tensorrt-infer-fp16-and-fp32-benchmarks/notebook\n\nimport os\nos.environ['CUDA_MODULE_LOADING']='LAZY'\n\ntry: \n    import torch_tensorrt\n    \nexcept:\n    #upgrade pytorch to 1.12\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    !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    #install torch_tensorrt\n    !pip install /kaggle/input/torch-tensorrt-pkg/torch_tensorrt-1.2.0-cp37-cp37m-linux_x86_64.whl\n\ntry: \n    import dicomsdl\n    \nexcept:\n    !pip install /kaggle/input/rsna-2022-whl/pylibjpeg-1.4.0-py3-none-any.whl\n    !pip install /kaggle/input/rsna-2022-whl/python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n    !pip install /kaggle/input/rsna-bcd-whl-ds/dicomsdl-0.109.1-cp37-cp37m-manylinux_2_12_x86_64.manylinux2010_x86_64.whl\n    !cp /kaggle/input/easy-load-the-image-with-nvjpeg2000/nvjpeg2k.so ./\n    \n\n\nimport torch_tensorrt\nimport tensorrt\nimport torch\nprint(torch.__version__)\n\nimport sys\n\nprint('install ok')","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:29:18.280723Z","iopub.execute_input":"2023-01-30T04:29:18.281561Z","iopub.status.idle":"2023-01-30T04:34:48.017487Z","shell.execute_reply.started":"2023-01-30T04:29:18.281476Z","shell.execute_reply":"2023-01-30T04:34:48.016117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport cv2\n\nfrom timeit import default_timer as timer\nfrom joblib import Parallel, delayed\nfrom glob import glob\n##from tqdm import tqdm\nfrom tqdm.notebook import tqdm\n\nimport os\nimport sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\nimport timm\nimport pydicom\n\nimport dicomsdl\nimport nvjpeg2k\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import SequentialSampler","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:36:07.472301Z","iopub.execute_input":"2023-01-30T04:36:07.473024Z","iopub.status.idle":"2023-01-30T04:36:09.63164Z","shell.execute_reply.started":"2023-01-30T04:36:07.472986Z","shell.execute_reply":"2023-01-30T04:36:09.630548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = True\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nBATCH_SIZE = 4\n\n# MODE = 'local' #测试用\n\nMODE = 'submit'","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:07.957473Z","iopub.execute_input":"2023-01-30T04:56:07.957876Z","iopub.status.idle":"2023-01-30T04:56:07.962817Z","shell.execute_reply.started":"2023-01-30T04:56:07.957822Z","shell.execute_reply":"2023-01-30T04:56:07.961701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from albumentations import (\n    HorizontalFlip, VerticalFlip, IAAPerspective, ShiftScaleRotate, CLAHE, RandomRotate90,\n    Transpose, ShiftScaleRotate, Blur, OpticalDistortion, GridDistortion, HueSaturationValue,\n    IAAAdditiveGaussianNoise, GaussNoise, MotionBlur, MedianBlur, IAAPiecewiseAffine, RandomResizedCrop,\n    IAASharpen, IAAEmboss, RandomBrightnessContrast, Flip, OneOf, Compose, Normalize, Cutout, CoarseDropout,\n    ShiftScaleRotate, CenterCrop, Resize, Rotate, PadIfNeeded\n)\nfrom albumentations.pytorch import ToTensorV2\n\nclass CFG:\n    # dicom to png size \n    resize_dim = 2048    # 保存的png分辨率，2048x2048\n    aspect_ratio = True\n    \n    # size of training image\n    img_size = [2048, 1024]    # 裁剪后resize的分辨率\n    \n    drop_rate = 0.5\n    drop_path_rate = 0.5\n    num_classes = 1\n#     backbone = 'tf_efficientnet_b0'\n    val_transforms = Compose([\n        CLAHE(clip_limit=4.0,tile_grid_size=(4,4),p=1),\n        PadIfNeeded(min_height=img_size[0], min_width=img_size[1], border_mode=cv2.BORDER_CONSTANT,\n                    value=[0, 0, 0], p=1),\n        Resize(1536, 960),    # 输入模型的分辨率\n        Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n        ToTensorV2(p=1.0),\n    ], p=1.)","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:40.716543Z","iopub.execute_input":"2023-01-30T04:56:40.716946Z","iopub.status.idle":"2023-01-30T04:56:40.727734Z","shell.execute_reply.started":"2023-01-30T04:56:40.716913Z","shell.execute_reply":"2023-01-30T04:56:40.726718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMG_DIR = '/tmp/Dataset/rsna-bcd'\n# IMG_DIR = '/kaggle/working/Dataset/rsna-bcd'\nos.makedirs(f'{IMG_DIR}', exist_ok = True)\n\nif MODE == 'local':\n    csv_file = '/kaggle/input/rsna-breast-cancer-detection/train.csv'\n    dcm_dir  = '/kaggle/input/rsna-breast-cancer-detection/train_images'\n\nif MODE == 'submit':\n    csv_file = '/kaggle/input/rsna-breast-cancer-detection/test.csv'\n    dcm_dir  = '/kaggle/input/rsna-breast-cancer-detection/test_images'\n    \ntest_df = pd.read_csv(csv_file)\nif 0:\n    test_df = test_df.head(10000)\nif MODE == 'local':\n    test_df = test_df.head(99)\n    test_df['prediction_id'] = test_df['patient_id'].astype(\"str\") + '_' + test_df['laterality'].astype(\"str\")\ntest_df['dicom_path'] = f'{dcm_dir}'\\\n                    + '/' + test_df.patient_id.astype(str)\\\n                    + '/' + test_df.image_id.astype(str)\\\n                    + '.dcm'\ntest_df['image_path'] = test_df.dicom_path.str.replace('.dcm','.png').str.replace(dcm_dir, IMG_DIR)\nprint('\\nTest:')\ntest_df\n","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:41.449012Z","iopub.execute_input":"2023-01-30T04:56:41.450117Z","iopub.status.idle":"2023-01-30T04:56:41.587743Z","shell.execute_reply.started":"2023-01-30T04:56:41.45007Z","shell.execute_reply":"2023-01-30T04:56:41.586896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\n\n# roi裁剪图片有效区域\ndef img2roi(img):\n    # Binarize the image\n    bin_img = cv2.threshold(img, 20, 255, cv2.THRESH_BINARY)[1]\n\n    # Make contours around the binarized image, keep only the largest contour\n    contours, _ = cv2.findContours(bin_img, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)\n    contour = max(contours, key=cv2.contourArea)\n\n    # Find ROI from largest contour\n    ys = contour.squeeze()[:, 0]\n    xs = contour.squeeze()[:, 1]\n    roi =  img[np.min(xs):np.max(xs), np.min(ys):np.max(ys)]\n    \n    return roi\n\n# def read_xray(path, fix_monochrome = True):\n#     dicom = dicomsdl.open(path)\n#     data = dicom.pixelData(storedvalue=False)  # storedvalue = True for int16 return otherwise float32\n#     data = data - np.min(data)\n#     data = data / np.max(data)\n#     if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n#         data = 1.0 - data\n#     return data\n\n# def resize_and_save(file_path):\n#     img = read_xray(file_path)\n#     h, w = img.shape[:2]  # orig hw\n#     if CFG.aspect_ratio:\n#         r = CFG.resize_dim / max(h, w)  # resize image to img_size\n#         interp = cv2.INTER_LINEAR\n#         if r != 1:  # always resize down, only resize up if training with augmentation\n#             img = cv2.resize(img, (int(w * r), int(h * r)), interpolation=interp)\n#     else:\n#         img = cv2.resize(img, (CFG.resize_dim, CFG.resize_dim), cv2.INTER_LINEAR)\n    \n#     img = (img * 255).astype(np.uint8)\n#     img = img2roi(img)\n#     img = cv2.resize(img, CFG.img_size[::-1], cv2.INTER_LINEAR)\n    \n#     sub_path = file_path.split(\"/\",5)[-1].split('.dcm')[0] + '.png'\n#     infos = sub_path.split('/')\n#     pid = infos[-2]\n#     iid = infos[-1]; iid = iid.replace('.png','')\n#     new_path = os.path.join(IMG_DIR, sub_path)\n#     os.makedirs(new_path.rsplit('/',1)[0], exist_ok=True)\n#     cv2.imwrite(new_path, img)\n#     return pid,iid,w,h","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:42.699783Z","iopub.execute_input":"2023-01-30T04:56:42.7002Z","iopub.status.idle":"2023-01-30T04:56:42.709087Z","shell.execute_reply.started":"2023-01-30T04:56:42.700167Z","shell.execute_reply":"2023-01-30T04:56:42.707982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%time\n# from joblib import Parallel, delayed\n# file_paths = test_df.dicom_path.tolist()\n# # print(len(file_paths))\n# imgsize = Parallel(n_jobs=2,backend='threading')(delayed(resize_and_save)(file_path)\\\n#                                                   for file_path in tqdm(file_paths, leave=True, position=0))","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:43.618009Z","iopub.execute_input":"2023-01-30T04:56:43.618368Z","iopub.status.idle":"2023-01-30T04:56:43.622925Z","shell.execute_reply.started":"2023-01-30T04:56:43.618337Z","shell.execute_reply":"2023-01-30T04:56:43.621793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"11min44s","metadata":{}},{"cell_type":"code","source":"# 这部分不用看，直接用\ndef read_image(df, image_dir):\n    image = []\n    for t,d in df.iterrows():\n        image_file = f'{image_dir}/{d.patient_id}/{d.image_id}.png'\n        m = cv2.imread(image_file,cv2.IMREAD_ANYDEPTH)\n        image.append(m)\n    return image\n\ndef make_transfer_syntax_uid(df, dcm_dir):\n    machine_id_to_transfer = {}\n    machine_id = df.machine_id.unique()\n    for i in machine_id:\n        d = df[df.machine_id == i].iloc[0]\n        f = f'{dcm_dir}/{d.patient_id}/{d.image_id}.dcm'\n        dicom = pydicom.dcmread(f)\n        machine_id_to_transfer[i] = dicom.file_meta.TransferSyntaxUID\n    return machine_id_to_transfer\n\ndef normalised_to_8bit(image, photometric_interpretation):\n    xmin = image.min()\n    xmax = image.max()\n\n    norm = np.empty_like(image, dtype=np.uint8)\n    dicomsdl.util.convert_to_uint8(image, norm, xmin, xmax)\n    if photometric_interpretation == 'MONOCHROME1':\n        norm = 255 - norm\n    return norm\n\ndef resize_image_to_height(image, image_height):\n    h, w = image.shape[:2]\n    s = image_height/h\n    if image_height!=h:\n        image = cv2.resize(image, dsize=None, fx=s, fy=s, interpolation=cv2.INTER_LINEAR)\n    return image","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:44.660778Z","iopub.execute_input":"2023-01-30T04:56:44.661836Z","iopub.status.idle":"2023-01-30T04:56:44.672748Z","shell.execute_reply.started":"2023-01-30T04:56:44.6618Z","shell.execute_reply":"2023-01-30T04:56:44.671789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 图片有两种类型，一种可以用jpeg2000,一种只能用dicomsdl\n# 但是过程都一样，都是读图，resize，roi裁剪，resize到1596x960\n#----------------------------------------------------------------\n# dicomsdl reader\ndef dicomsdl_to_numpy_image(ds, index=0):\n    # https://stackoverflow.com/questions/44659924/returning-numpy-arrays-via-pybind11\n    info = ds.getPixelDataInfo()\n    if info['SamplesPerPixel'] != 1:\n        raise RuntimeError('SamplesPerPixel != 1')\n\n    shape = [info['Rows'], info['Cols']]\n    dtype = info['dtype']\n    outarr = np.empty(shape, dtype=dtype)\n    ds.copyFrameData(index, outarr)\n    return outarr\n\ndef dicomsdl_parallel_process(d, dcm_dir, image_dir):\n    # 读图\n    dcm_file = f'{dcm_dir}/{d.patient_id}/{d.image_id}.dcm'\n    ds = dicomsdl.open(dcm_file)\n    dc = pydicom.dcmread(dcm_file)\n    img = dicomsdl_to_numpy_image(ds)\n#     image = resize_image_to_height(image, image_height)\n    # resize\n    h, w = img.shape[:2]  # orig hw\n    if CFG.aspect_ratio:\n        r = CFG.resize_dim / max(h, w)  # resize image to img_size\n        interp = cv2.INTER_LINEAR\n        if r != 1:  # always resize down, only resize up if training with augmentation\n            img = cv2.resize(img, (int(w * r), int(h * r)), interpolation=interp)\n    else:\n        img = cv2.resize(img, (CFG.resize_dim, CFG.resize_dim), cv2.INTER_LINEAR)\n\n    img = img.astype(np.float32)\n    img = normalised_to_8bit(img, dc.PhotometricInterpretation)\n    # roi\n    img = img2roi(img)\n    # resize again\n    img = cv2.resize(img, CFG.img_size[::-1], cv2.INTER_LINEAR)\n\n    # save as png\n    os.makedirs(f'{image_dir}/{d.patient_id}', exist_ok=True)\n    cv2.imwrite(f'{image_dir}/{d.patient_id}/{d.image_id}.png', img)\n\ndef process_non_j2k(df, dcm_dir, image_dir, n_jobs):\n    #https://stackoverflow.com/questions/56659294/does-joblib-parallel-keep-the-original-order-of-data-passed\n    #Parallel(n_jobs=2, backend='multiprocessing')(\n    Parallel(n_jobs=n_jobs)(\n        delayed(dicomsdl_parallel_process)(d, dcm_dir, image_dir)\n        for t,d in tqdm(df.iterrows())\n    )\n\n#----------------------------------------------------------------\n# nvjpeg2k reader\n\n'''\nTransferSyntaxUID\n1.2.840.10008.1.2.4.70 = JPEG Lossless, Nonhierarchical, First- Order Prediction (Processes 14)\n1.2.840.10008.1.2.4.90 = JPEG 2000 Image Compression (Lossless Only)\n'''\nj2k_decoder = nvjpeg2k.Decoder()\n\ndef process_j2k(df, dcm_dir, image_dir):\n    for t, d in tqdm(df.iterrows()):\n        dcm_file = f'{dcm_dir}/{d.patient_id}/{d.image_id}.dcm'\n        dc = pydicom.dcmread(dcm_file)\n        offset = dc.PixelData.find(b'\\x00\\x00\\x00\\x0C')\n        jpeg_stream = bytearray(dc.PixelData[offset:])\n        img = j2k_decoder.decode(jpeg_stream)\n        \n#         image = resize_image_to_height(image, resize_dim)\n        h, w = img.shape[:2]  # orig hw\n        if CFG.aspect_ratio:\n            r = CFG.resize_dim / max(h, w)  # resize image to img_size\n            interp = cv2.INTER_LINEAR\n            if r != 1:  # always resize down, only resize up if training with augmentation\n                img = cv2.resize(img, (int(w * r), int(h * r)), interpolation=interp)\n        else:\n            img = cv2.resize(img, (CFG.resize_dim, CFG.resize_dim), cv2.INTER_LINEAR)\n\n        img = img.astype(np.float32)\n        img = normalised_to_8bit(img, dc.PhotometricInterpretation)\n        img = img2roi(img)\n        img = cv2.resize(img, CFG.img_size[::-1], cv2.INTER_LINEAR)\n        \n\n        # save as png\n        os.makedirs(f'{image_dir}/{d.patient_id}', exist_ok=True)\n        cv2.imwrite(f'{image_dir}/{d.patient_id}/{d.image_id}.png', img)","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:45.237747Z","iopub.execute_input":"2023-01-30T04:56:45.238151Z","iopub.status.idle":"2023-01-30T04:56:45.256411Z","shell.execute_reply.started":"2023-01-30T04:56:45.238118Z","shell.execute_reply":"2023-01-30T04:56:45.255084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"machine_id_to_transfer = make_transfer_syntax_uid(test_df, dcm_dir)\ntest_df.loc[:, 'i'] = np.arange(len(test_df))\ntest_df.loc[:, 'TransferSyntaxUID'] = test_df.machine_id.map(machine_id_to_transfer)","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:46.146095Z","iopub.execute_input":"2023-01-30T04:56:46.146477Z","iopub.status.idle":"2023-01-30T04:56:46.186338Z","shell.execute_reply.started":"2023-01-30T04:56:46.146445Z","shell.execute_reply":"2023-01-30T04:56:46.185375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 如上所述两种类型图片\nj2k_df = test_df[test_df.TransferSyntaxUID == '1.2.840.10008.1.2.4.90'].reset_index(drop=True)\nnon_j2k_df = test_df[test_df.TransferSyntaxUID != '1.2.840.10008.1.2.4.90'].reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:47.071947Z","iopub.execute_input":"2023-01-30T04:56:47.072624Z","iopub.status.idle":"2023-01-30T04:56:47.0811Z","shell.execute_reply.started":"2023-01-30T04:56:47.072583Z","shell.execute_reply":"2023-01-30T04:56:47.079977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 分别处理\n%%time\nprint(f'process_j2k(): {len(j2k_df)}')\nprocess_j2k(j2k_df, dcm_dir, IMG_DIR)","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:48.158386Z","iopub.execute_input":"2023-01-30T04:56:48.158774Z","iopub.status.idle":"2023-01-30T04:56:55.910726Z","shell.execute_reply.started":"2023-01-30T04:56:48.15874Z","shell.execute_reply":"2023-01-30T04:56:55.9097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nprint(f'process_non_j2k(): {len(non_j2k_df)}')\nprocess_non_j2k(non_j2k_df, dcm_dir, IMG_DIR, n_jobs=2)  ","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:56:57.675989Z","iopub.execute_input":"2023-01-30T04:56:57.676357Z","iopub.status.idle":"2023-01-30T04:57:09.497159Z","shell.execute_reply.started":"2023-01-30T04:56:57.676325Z","shell.execute_reply":"2023-01-30T04:57:09.495942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n# m = cv2.imread('/kaggle/working/test_images/10102/618254763.png',cv2.IMREAD_GRAYSCALE)\n# print(m.shape)\n# print(m)\n# plt.imshow(m,cmap='bone')\n# plt.show() ","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:57:10.781429Z","iopub.execute_input":"2023-01-30T04:57:10.781933Z","iopub.status.idle":"2023-01-30T04:57:10.791191Z","shell.execute_reply.started":"2023-01-30T04:57:10.781886Z","shell.execute_reply":"2023-01-30T04:57:10.789832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dataset 设置\nif 1:\n    from torch.utils.data import Dataset\n    # from utils import *\n    from PIL import Image\n    import cv2\n\n    def resize_image(image, size, letterbox_image):\n        image = Image.fromarray(image)\n        iw, ih = image.size\n        # print(image.size)\n        h, w = size\n\n        if letterbox_image:\n            scale  = min(w/iw, h/ih)\n            nw = int(iw*scale)\n            nh = int(ih*scale)\n            image  = image.resize((nw,nh), Image.Resampling.BICUBIC)\n            new_image = Image.new('RGB', (w,h), (0,0,0))\n            new_image.paste(image, ((w-nw)//2, (h-nh)//2))\n        else:\n            new_image = image.resize((w, h), Image.Resampling.BICUBIC)\n        return np.array(new_image)\n\n    def get_img(path):\n        im_bgr = cv2.imread(path)\n        im_rgb = im_bgr[:, :, ::-1]\n        # print(im_rgb)\n        return im_rgb\n\n    class CustomDataset(Dataset):\n        def __init__(self, df, cfg, data_root,\n                     transforms=None,\n                     output_label=True,\n                     one_hot_label=False,\n                ):\n\n            super().__init__()\n            self.df = df.reset_index(drop=True).copy()\n            self.transforms = transforms\n            self.data_root = data_root\n            self.cfg = cfg\n\n            self.output_label = output_label\n            self.one_hot_label = one_hot_label\n\n            self.predict_id = self.df['prediction_id'].values\n\n            if output_label == True:\n                self.labels = self.df['cancer'].values\n\n                if one_hot_label is True:\n                    self.labels = np.eye(self.df['cancer'].max() + 1)[self.labels]\n\n        def __len__(self):\n            return self.df.shape[0]\n\n        def __getitem__(self, index: int):\n\n            prediction_id = self.predict_id[index]\n\n            # get labels\n            if self.output_label:\n                target = self.labels[index]\n\n            img_id = self.df.loc[index]['image_id']\n            patient_id = self.df.loc[index]['patient_id']\n    #         print(\"{}/{}/{}.png \\n\".format(self.data_root, patient_id, img_id))\n            # get images\n            try:\n                img = get_img(\"{}/{}/{}.png\".format(self.data_root, patient_id, img_id))\n            except:\n                print(\"{}/{}/{}.png \\n\".format(self.data_root, patient_id, img_id))\n\n            # breast pic turn right\n            if img[:,-10:,:].mean() > img[:,10:,:].mean():\n                img = cv2.flip(img, 1)\n\n            img = resize_image(img, self.cfg.img_size, letterbox_image=True)\n\n            # get transforms\n            if self.transforms:\n                img = self.transforms(image=img)['image']\n\n            if self.output_label == True:\n                return img, target, prediction_id\n            else:\n                return img, prediction_id","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:57:11.462243Z","iopub.execute_input":"2023-01-30T04:57:11.462641Z","iopub.status.idle":"2023-01-30T04:57:11.481032Z","shell.execute_reply.started":"2023-01-30T04:57:11.462609Z","shell.execute_reply":"2023-01-30T04:57:11.479841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch.nn as nn\n# class RSNAClassifier(nn.Module):\n#     def __init__(self, pretrained=False):\n#         super().__init__()\n#         self.model = timm.create_model('efficientnet_b0', pretrained=pretrained, num_classes=1, drop_rate=CFG.drop_rate, drop_path_rate=CFG.drop_path_rate)\n\n#     def forward(self, x):\n#         x = self.model(x)\n#         return x","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:57:13.072301Z","iopub.execute_input":"2023-01-30T04:57:13.072685Z","iopub.status.idle":"2023-01-30T04:57:13.077263Z","shell.execute_reply.started":"2023-01-30T04:57:13.072651Z","shell.execute_reply":"2023-01-30T04:57:13.076212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# if 1:\n#     model = RSNAClassifier()\n#     model.to(DEVICE)\n\n#     # load weights\n\n\n#     model.eval()\n#     trt_model_fp16 = torch_tensorrt.compile(\n#         model,\n#         inputs=[\n#             torch_tensorrt.Input(min_shape=[1, 3, 2048, 1024],\n#                                  opt_shape=[BATCH_SIZE, 3, 2048, 1024],\n#                                  max_shape=[BATCH_SIZE, 3, 2048, 1024],dtype=torch.half)],\n#         enabled_precisions={torch.half},  # Run with FP16\n#         workspace_size=1 << 32,\n# #         require_full_compilation=True,\n#     ) \n#     torch.jit.save(trt_model_fp16, 'kaggle-nextvit-b-1536-gpu-aug0-01-swa.trt_fp16.ts')\n#     #compy this file to your own dataset to use again in submission\n#     print('trt_fp16 ok')","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:57:13.437694Z","iopub.execute_input":"2023-01-30T04:57:13.438645Z","iopub.status.idle":"2023-01-30T04:57:13.443977Z","shell.execute_reply.started":"2023-01-30T04:57:13.438587Z","shell.execute_reply":"2023-01-30T04:57:13.442966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 提交新模型的时候，需要将0改为1，同时把tensorrt_compile.py 中的模型改成新模型\n# 目前tensorrt_compile.py 是nextvit模型，提交convnextv2的话直接用timm\nif 0:\n    !python /kaggle/input/tensorrt-compile/tensorrt_compile.py\n    \n# 运行完毕会生成一个 .ts文件，这个文件保存下来，在下面调用\n","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:57:14.979479Z","iopub.execute_input":"2023-01-30T04:57:14.980214Z","iopub.status.idle":"2023-01-30T04:58:50.920276Z","shell.execute_reply.started":"2023-01-30T04:57:14.980169Z","shell.execute_reply":"2023-01-30T04:58:50.919043Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if 1:  \n    \n    # 这里是.ts文件集合，若干个模型\n    model = [\n    #             '/kaggle/input/rsna-breast-mammography-weight-10/kaggle-nextvit-b-1536-gpu-aug0-01-swa.trt_fp16.ts',\n                '/kaggle/input/nextvit-tensorrt/kaggle-nextvit-b-1536-gpu-aug0-01-swa.trt_fp16.ts'\n            ]\n    num_net = len(model)\n\n    net = []\n    for i in range(num_net):\n        n = torch.jit.load(model[i])\n        net.append(n)\n\n\n    test_dataset = CustomDataset(test_df, CFG, IMG_DIR, transforms=CFG.val_transforms, output_label=False)\n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=BATCH_SIZE,\n        sampler=SequentialSampler(test_dataset),\n        drop_last=False,\n        num_workers=2,\n    )","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:36:56.418149Z","iopub.execute_input":"2023-01-30T04:36:56.418528Z","iopub.status.idle":"2023-01-30T04:36:59.220629Z","shell.execute_reply.started":"2023-01-30T04:36:56.418495Z","shell.execute_reply":"2023-01-30T04:36:59.219706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 这个是为了处理最后不足一个完整batch的情况\nimport torch.nn.functional as F\ndef pad_to_batch_size(image, batch_size):\n    B = len(image)\n    if B == batch_size:\n        return image, False\n    pad = F.pad(input=image, pad=(0, 0, 0, 0, 0, 0, 0, batch_size - B), mode='constant', value=0)\n    return pad, True","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:39:49.21309Z","iopub.execute_input":"2023-01-30T04:39:49.213476Z","iopub.status.idle":"2023-01-30T04:39:49.220063Z","shell.execute_reply.started":"2023-01-30T04:39:49.213445Z","shell.execute_reply":"2023-01-30T04:39:49.218801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 模型推理，融合\n# 按照prediction_id取mean\n\nfrom torch.cuda.amp import autocast\ndef models_predict(models, dl_test, max_batches=1e9):\n    for m in models:\n        m.eval()\n        \n    result = {\n        'probability': [[] for i in range(num_net)],\n    }\n    test_num = 0\n    \n    for t, (X,_) in enumerate(tqdm(test_loader)):\n        B = len(X)\n        image, is_pad = pad_to_batch_size(X.cuda().half(), BATCH_SIZE)\n        image0 = image\n#         image1 = torch.flip(image, dims=[3, ])  # TTA\n        #print(image0.shape)\n\n        p = 0\n        count = 0 \n        with torch.no_grad():\n#             with autocast(enabled=True):\n            for i in range(num_net):\n                p += models[i](image0).squeeze()\n                count += 1\n\n#                     p += models[i](image1).squeeze()\n#                     count += 1\n\n        p = p / count\n        if is_pad:\n            p = p[:B]\n\n        result['probability'].append(p.float().data.cpu().numpy())\n        test_num += B\n        torch.cuda.empty_cache()\n\n    # ---\n    probability = np.concatenate(result['probability'])\n    probability = np.nan_to_num(probability, nan=0, posinf=1, neginf=0)\n    np.save('probability.npy', probability)\n\n\n        \n#         predictions = []\n#         with torch.no_grad():\n#             with autocast(enabled=True):\n#                 for idx, (X,_) in enumerate(tqdm(dl_test, mininterval=30)):\n#                     B = len(X)\n#                     X, is_pad = pad_to_batch_size(X.to(DEVICE).half(), BATCH_SIZE)\n                    \n#                     p = 0\n#                     count = 0 \n#                     for i in range(num_net):\n#                         p += net[i](image0)\n#                         count += 1\n\n#                         p += net[i](image1)\n#                         count += 1\n                    \n#                     pred = torch.zeros(len(X), len(models))\n#                     for idx, m in enumerate(models):\n#                         preds = torch.sigmoid(m(X.to(DEVICE).half()).squeeze())\n#         #                 print(preds)\n#                         pred[:, idx] = preds.cpu()\n#                     predictions.append(pred.mean(dim=-1))\n\n#                     if idx >= max_batches:\n#                         break\n#         return torch.concat(predictions).numpy()","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:42:21.403732Z","iopub.execute_input":"2023-01-30T04:42:21.404454Z","iopub.status.idle":"2023-01-30T04:42:21.418062Z","shell.execute_reply.started":"2023-01-30T04:42:21.404419Z","shell.execute_reply":"2023-01-30T04:42:21.417093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nmodels_pred = models_predict(net, test_loader)\nprint(models_pred)","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:42:22.112524Z","iopub.execute_input":"2023-01-30T04:42:22.113247Z","iopub.status.idle":"2023-01-30T04:42:40.880018Z","shell.execute_reply.started":"2023-01-30T04:42:22.113203Z","shell.execute_reply":"2023-01-30T04:42:40.878797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"probability = np.load('probability.npy')\nprint('probability', probability.shape)\nprint('')\n\n# threshold = 0.46939\nthreshold = 0.20\n\nsubmit_df = pd.DataFrame({\n    'prediction_id': test_df.prediction_id,\n    'cancer': probability,\n})\n\nsubmit_df = submit_df.groupby('prediction_id').mean()\nsubmit_df = submit_df.sort_index()\npredict = submit_df.cancer.values\n\nsubmit_df.loc[:, 'cancer'] = (submit_df.cancer.values > threshold).astype(np.float32)\npredict_threshold = submit_df.cancer.values\n\nsubmit_df.to_csv('submission.csv', index=True)\n# submit_df","metadata":{"execution":{"iopub.status.busy":"2023-01-30T04:43:24.517752Z","iopub.execute_input":"2023-01-30T04:43:24.518527Z","iopub.status.idle":"2023-01-30T04:43:24.540441Z","shell.execute_reply.started":"2023-01-30T04:43:24.518488Z","shell.execute_reply":"2023-01-30T04:43:24.53943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# THRES = 0.46939","metadata":{"execution":{"iopub.status.busy":"2023-01-29T22:23:37.531167Z","iopub.execute_input":"2023-01-29T22:23:37.531558Z","iopub.status.idle":"2023-01-29T22:23:37.536724Z","shell.execute_reply.started":"2023-01-29T22:23:37.531524Z","shell.execute_reply":"2023-01-29T22:23:37.53545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_df['cancer'] = models_pred\n\n\n# df_sub = test_df.groupby('prediction_id')[['cancer']].mean()\n# df_sub","metadata":{"execution":{"iopub.status.busy":"2023-01-29T22:23:38.83958Z","iopub.execute_input":"2023-01-29T22:23:38.839979Z","iopub.status.idle":"2023-01-29T22:23:38.866182Z","shell.execute_reply.started":"2023-01-29T22:23:38.839938Z","shell.execute_reply":"2023-01-29T22:23:38.865027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df_sub['cancer'] = (df_sub.cancer > THRES).astype(float)\n# df_sub","metadata":{"execution":{"iopub.status.busy":"2023-01-29T22:23:45.191758Z","iopub.execute_input":"2023-01-29T22:23:45.192135Z","iopub.status.idle":"2023-01-29T22:23:45.204838Z","shell.execute_reply.started":"2023-01-29T22:23:45.192103Z","shell.execute_reply":"2023-01-29T22:23:45.203727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df_sub.to_csv('submission.csv', index=True)\n# !head submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-01-29T22:23:48.363227Z","iopub.execute_input":"2023-01-29T22:23:48.363747Z","iopub.status.idle":"2023-01-29T22:23:49.612022Z","shell.execute_reply.started":"2023-01-29T22:23:48.363705Z","shell.execute_reply":"2023-01-29T22:23:49.610837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}