{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":4804745,"sourceType":"datasetVersion","datasetId":2779893},{"sourceId":4998596,"sourceType":"datasetVersion","datasetId":2891303},{"sourceId":5877069,"sourceType":"datasetVersion","datasetId":2998419},{"sourceId":9201122,"sourceType":"datasetVersion","datasetId":5562864},{"sourceId":9202106,"sourceType":"datasetVersion","datasetId":5562911},{"sourceId":9207197,"sourceType":"datasetVersion","datasetId":5567013},{"sourceId":159351763,"sourceType":"kernelVersion"},{"sourceId":159353975,"sourceType":"kernelVersion"}],"dockerImageVersionId":30381,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Dependencies","metadata":{}},{"cell_type":"code","source":"!pip config set global.disable-pip-version-check true\n!pip config set global.root-user-action ignore\n!pip install --no-index /kaggle/input/package-rsna/torch-2.3.1cpu.cxx11.abi-cp310-cp310-linux_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2024-08-20T05:33:18.606822Z","iopub.execute_input":"2024-08-20T05:33:18.607787Z","iopub.status.idle":"2024-08-20T05:33:23.369033Z","shell.execute_reply.started":"2024-08-20T05:33:18.607746Z","shell.execute_reply":"2024-08-20T05:33:23.367918Z"},"trusted":true},"execution_count":7,"outputs":[{"name":"stdout","text":"Writing to /root/.config/pip/pip.conf\nWriting to /root/.config/pip/pip.conf\n\u001b[31mERROR: torch-2.3.1cpu.cxx11.abi-cp310-cp310-linux_x86_64.whl is not a supported wheel on this platform.\u001b[0m\u001b[31m\n\u001b[0m","output_type":"stream"}]},{"cell_type":"code","source":"\nimport os\nos.environ['CUDA_MODULE_LOADING']='LAZY'\n\n!mkdir -p /kaggle/tmp/libs\n# # upgrade pytorch to 1.12 for torch_tensorrt\n\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":"2024-08-20T05:23:16.372343Z","iopub.execute_input":"2024-08-20T05:23:16.372772Z","iopub.status.idle":"2024-08-20T05:25:03.802031Z","shell.execute_reply.started":"2024-08-20T05:23:16.372728Z","shell.execute_reply":"2024-08-20T05:25:03.80067Z"},"trusted":true},"execution_count":4,"outputs":[{"name":"stdout","text":"Writing to /root/.config/pip/pip.conf\nWriting to /root/.config/pip/pip.conf\nFound existing installation: timm 0.8.11.dev0\nUninstalling timm-0.8.11.dev0:\n  Successfully uninstalled timm-0.8.11.dev0\n/kaggle/tmp/libs/timm\nObtaining file:///kaggle/tmp/libs/timm\n  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hRequirement already satisfied: torch>=1.7 in /opt/conda/lib/python3.7/site-packages (from timm==0.8.11.dev0) (1.12.1)\nRequirement already satisfied: torchvision in /opt/conda/lib/python3.7/site-packages (from timm==0.8.11.dev0) (0.12.0)\nRequirement already satisfied: pyyaml in /opt/conda/lib/python3.7/site-packages (from timm==0.8.11.dev0) (6.0)\nRequirement already satisfied: huggingface_hub in /opt/conda/lib/python3.7/site-packages (from timm==0.8.11.dev0) (0.10.1)\nRequirement already satisfied: typing-extensions in /opt/conda/lib/python3.7/site-packages (from torch>=1.7->timm==0.8.11.dev0) (4.1.1)\nRequirement already satisfied: packaging>=20.9 in /opt/conda/lib/python3.7/site-packages (from huggingface_hub->timm==0.8.11.dev0) (23.0)\nRequirement already satisfied: tqdm in /opt/conda/lib/python3.7/site-packages (from huggingface_hub->timm==0.8.11.dev0) (4.64.0)\nRequirement already satisfied: filelock in /opt/conda/lib/python3.7/site-packages (from huggingface_hub->timm==0.8.11.dev0) (3.7.1)\nRequirement already satisfied: importlib-metadata in /opt/conda/lib/python3.7/site-packages (from huggingface_hub->timm==0.8.11.dev0) (6.0.0)\nRequirement already satisfied: requests in /opt/conda/lib/python3.7/site-packages (from huggingface_hub->timm==0.8.11.dev0) (2.28.1)\nRequirement already satisfied: numpy in /opt/conda/lib/python3.7/site-packages (from torchvision->timm==0.8.11.dev0) (1.21.6)\nRequirement already satisfied: pillow!=8.3.*,>=5.3.0 in /opt/conda/lib/python3.7/site-packages (from torchvision->timm==0.8.11.dev0) (9.1.1)\nRequirement already satisfied: zipp>=0.5 in /opt/conda/lib/python3.7/site-packages (from importlib-metadata->huggingface_hub->timm==0.8.11.dev0) (3.8.0)\nRequirement already satisfied: charset-normalizer<3,>=2 in /opt/conda/lib/python3.7/site-packages (from requests->huggingface_hub->timm==0.8.11.dev0) (2.1.0)\nRequirement already satisfied: idna<4,>=2.5 in /opt/conda/lib/python3.7/site-packages (from requests->huggingface_hub->timm==0.8.11.dev0) (3.3)\nRequirement already satisfied: certifi>=2017.4.17 in /opt/conda/lib/python3.7/site-packages (from requests->huggingface_hub->timm==0.8.11.dev0) (2022.12.7)\nRequirement already satisfied: urllib3<1.27,>=1.21.1 in /opt/conda/lib/python3.7/site-packages (from requests->huggingface_hub->timm==0.8.11.dev0) (1.26.14)\nInstalling collected packages: timm\n  Running setup.py develop for timm\nSuccessfully installed timm-0.8.11.dev0\n/kaggle/working\nProcessing /kaggle/input/torch-tensorrt-pkg/nvidia_pyindex-1.0.9-py3-none-any.whl\nnvidia-pyindex is already installed with the same version as the provided wheel. Use --force-reinstall to force an installation of the wheel.\nLooking in links: /tmp/pip/cache/\nRequirement already satisfied: nvidia_tensorrt in /opt/conda/lib/python3.7/site-packages (8.4.3.1)\nRequirement already satisfied: nvidia-cublas-cu11 in /opt/conda/lib/python3.7/site-packages (from nvidia_tensorrt) (2022.4.8)\nRequirement already satisfied: nvidia-cudnn-cu11 in /opt/conda/lib/python3.7/site-packages (from nvidia_tensorrt) (2022.5.19)\nRequirement already satisfied: nvidia-cuda-runtime-cu11 in /opt/conda/lib/python3.7/site-packages (from nvidia_tensorrt) (2022.4.25)\nRequirement already satisfied: nvidia-cublas-cu117 in /opt/conda/lib/python3.7/site-packages (from nvidia-cublas-cu11->nvidia_tensorrt) (11.10.1.25)\nRequirement already satisfied: nvidia-cuda-runtime-cu117 in /opt/conda/lib/python3.7/site-packages (from nvidia-cuda-runtime-cu11->nvidia_tensorrt) (11.7.60)\nRequirement already satisfied: nvidia-cudnn-cu116 in /opt/conda/lib/python3.7/site-packages (from nvidia-cudnn-cu11->nvidia_tensorrt) (8.4.0.27)\nRequirement already satisfied: wheel in /opt/conda/lib/python3.7/site-packages (from nvidia-cublas-cu117->nvidia-cublas-cu11->nvidia_tensorrt) (0.37.1)\nRequirement already satisfied: setuptools in /opt/conda/lib/python3.7/site-packages (from nvidia-cublas-cu117->nvidia-cublas-cu11->nvidia_tensorrt) (59.8.0)\nProcessing /kaggle/input/torch-tensorrt-pkg/torch_tensorrt-1.2.0-cp37-cp37m-linux_x86_64.whl\nRequirement already satisfied: torch<1.13.0,>=1.12.0+cu116 in /opt/conda/lib/python3.7/site-packages (from torch-tensorrt==1.2.0) (1.12.1)\nRequirement already satisfied: typing-extensions in /opt/conda/lib/python3.7/site-packages (from torch<1.13.0,>=1.12.0+cu116->torch-tensorrt==1.2.0) (4.1.1)\ntorch-tensorrt is already installed with the same version as the provided wheel. Use --force-reinstall to force an installation of the wheel.\n/kaggle/tmp/libs/torch2trt\nrunning install\n/opt/conda/lib/python3.7/site-packages/setuptools/command/install.py:37: SetuptoolsDeprecationWarning: setup.py install is deprecated. Use build and pip and other standards-based tools.\n  setuptools.SetuptoolsDeprecationWarning,\n/opt/conda/lib/python3.7/site-packages/setuptools/command/easy_install.py:159: EasyInstallDeprecationWarning: easy_install command is deprecated. Use build and pip and other standards-based tools.\n  EasyInstallDeprecationWarning,\nrunning bdist_egg\nrunning egg_info\nwriting torch2trt.egg-info/PKG-INFO\nwriting dependency_links to torch2trt.egg-info/dependency_links.txt\nwriting top-level names to torch2trt.egg-info/top_level.txt\nreading manifest file 'torch2trt.egg-info/SOURCES.txt'\nadding license file 'LICENSE.md'\nwriting manifest file 'torch2trt.egg-info/SOURCES.txt'\ninstalling library code to build/bdist.linux-x86_64/egg\nrunning install_lib\nrunning build_py\ncopying torch2trt/test.py -> build/lib/torch2trt\ncopying torch2trt/torch2trt.py -> build/lib/torch2trt\ncopying torch2trt/__init__.py -> build/lib/torch2trt\ncopying torch2trt/flatten_module_test.py -> build/lib/torch2trt\ncopying torch2trt/module_test.py -> build/lib/torch2trt\ncopying torch2trt/dataset_calibrator_test.py -> build/lib/torch2trt\ncopying torch2trt/dataset_calibrator.py -> build/lib/torch2trt\ncopying torch2trt/dataset.py -> build/lib/torch2trt\ncopying torch2trt/dynamic_shape_test.py -> build/lib/torch2trt\ncopying torch2trt/flattener_test.py -> build/lib/torch2trt\ncopying torch2trt/flatten_module.py -> build/lib/torch2trt\ncopying torch2trt/dataset_test.py -> build/lib/torch2trt\ncopying torch2trt/utils.py -> build/lib/torch2trt\ncopying torch2trt/flattener.py -> build/lib/torch2trt\ncopying torch2trt/tests/test_tensor_ne.py -> build/lib/torch2trt/tests\ncopying torch2trt/tests/__init__.py -> build/lib/torch2trt/tests\ncopying torch2trt/tests/test_contiguous.py -> build/lib/torch2trt/tests\ncopying torch2trt/tests/test_legacy_max_batch_size.py -> build/lib/torch2trt/tests\ncopying torch2trt/tests/test_flatten_dynamic.py -> build/lib/torch2trt/tests\ncopying torch2trt/tests/test_tensor_shape_div_batch.py -> build/lib/torch2trt/tests\ncopying torch2trt/tests/test_interpolate_dynamic.py -> build/lib/torch2trt/tests\ncopying torch2trt/tests/test_tensor_shape.py -> build/lib/torch2trt/tests\ncopying torch2trt/converters/BatchNorm2d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/prod.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/Conv1d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/sub.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/layer_norm.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/div.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/permute.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/batch_norm.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/getitem.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/identity.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/__init__.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/AdaptiveAvgPool2d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/view.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/adaptive_max_pool3d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/chunk.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/ne.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/avg_pool.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/mul.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/pow.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/sum.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/normalize.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/max.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/BatchNorm3d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/matmul.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/conv_functional.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/clamp.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/example_plugin.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/gelu.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/interpolate.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/relu6.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/unsqueeze.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/mod.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/sigmoid.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/ConvTranspose2d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/prelu.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/relu.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/compare.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/split.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/tensor.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/pad.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/activation.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/stack.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/silu.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/squeeze.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/expand.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/max_pool1d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/einsum.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/mean.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/Linear.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/AdaptiveAvgPool3d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/min.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/adaptive_max_pool2d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/adaptive_avg_pool2d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/Conv2d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/group_norm.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/BatchNorm1d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/reflection_pad_2d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/dummy_converters.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/softmax.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/adaptive_avg_pool3d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/cat.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/Conv.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/add.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/instance_norm.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/roll.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/unary.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/tanh.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/max_pool3d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/getitem_test.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/clone.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/narrow.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/ConvTranspose.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/LogSoftmax.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/flatten.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/max_pool2d.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/transpose.py -> build/lib/torch2trt/converters\ncopying torch2trt/converters/floordiv.py -> build/lib/torch2trt/converters\ncopying torch2trt/contrib/__init__.py -> build/lib/torch2trt/contrib\ncopying torch2trt/tests/timm/__init__.py -> build/lib/torch2trt/tests/timm\ncopying torch2trt/tests/timm/test_maxvit.py -> build/lib/torch2trt/tests/timm\ncopying torch2trt/tests/torchvision/save_load.py -> build/lib/torch2trt/tests/torchvision\ncopying torch2trt/tests/torchvision/__init__.py -> build/lib/torch2trt/tests/torchvision\ncopying torch2trt/tests/torchvision/segmentation.py -> build/lib/torch2trt/tests/torchvision\ncopying torch2trt/tests/torchvision/classification.py -> build/lib/torch2trt/tests/torchvision\ncopying torch2trt/contrib/qat/__init__.py -> build/lib/torch2trt/contrib/qat\ncopying torch2trt/contrib/qat/converters/__init__.py -> build/lib/torch2trt/contrib/qat/converters\ncopying torch2trt/contrib/qat/converters/QuantConvBN.py -> build/lib/torch2trt/contrib/qat/converters\ncopying torch2trt/contrib/qat/converters/QuantConv.py -> build/lib/torch2trt/contrib/qat/converters\ncopying torch2trt/contrib/qat/converters/QuantRelu.py -> build/lib/torch2trt/contrib/qat/converters\ncopying torch2trt/contrib/qat/layers/quant_activation.py -> build/lib/torch2trt/contrib/qat/layers\ncopying torch2trt/contrib/qat/layers/__init__.py -> build/lib/torch2trt/contrib/qat/layers\ncopying torch2trt/contrib/qat/layers/quant_conv.py -> build/lib/torch2trt/contrib/qat/layers\ncopying torch2trt/contrib/qat/layers/_utils.py -> build/lib/torch2trt/contrib/qat/layers\ncreating build/bdist.linux-x86_64/egg\ncreating build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/test.py -> build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/torch2trt.py -> build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/__init__.py -> build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/flatten_module_test.py -> build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/module_test.py -> build/bdist.linux-x86_64/egg/torch2trt\ncreating build/bdist.linux-x86_64/egg/torch2trt/tests\ncopying build/lib/torch2trt/tests/test_tensor_ne.py -> build/bdist.linux-x86_64/egg/torch2trt/tests\ncopying build/lib/torch2trt/tests/__init__.py -> build/bdist.linux-x86_64/egg/torch2trt/tests\ncopying build/lib/torch2trt/tests/test_contiguous.py -> build/bdist.linux-x86_64/egg/torch2trt/tests\ncopying build/lib/torch2trt/tests/test_legacy_max_batch_size.py -> build/bdist.linux-x86_64/egg/torch2trt/tests\ncopying build/lib/torch2trt/tests/test_flatten_dynamic.py -> build/bdist.linux-x86_64/egg/torch2trt/tests\ncopying build/lib/torch2trt/tests/test_tensor_shape_div_batch.py -> build/bdist.linux-x86_64/egg/torch2trt/tests\ncreating build/bdist.linux-x86_64/egg/torch2trt/tests/timm\ncopying build/lib/torch2trt/tests/timm/__init__.py -> build/bdist.linux-x86_64/egg/torch2trt/tests/timm\ncopying build/lib/torch2trt/tests/timm/test_maxvit.py -> build/bdist.linux-x86_64/egg/torch2trt/tests/timm\ncopying build/lib/torch2trt/tests/test_interpolate_dynamic.py -> build/bdist.linux-x86_64/egg/torch2trt/tests\ncreating build/bdist.linux-x86_64/egg/torch2trt/tests/torchvision\ncopying build/lib/torch2trt/tests/torchvision/save_load.py -> build/bdist.linux-x86_64/egg/torch2trt/tests/torchvision\ncopying build/lib/torch2trt/tests/torchvision/__init__.py -> build/bdist.linux-x86_64/egg/torch2trt/tests/torchvision\ncopying build/lib/torch2trt/tests/torchvision/segmentation.py -> build/bdist.linux-x86_64/egg/torch2trt/tests/torchvision\ncopying build/lib/torch2trt/tests/torchvision/classification.py -> build/bdist.linux-x86_64/egg/torch2trt/tests/torchvision\ncopying build/lib/torch2trt/tests/test_tensor_shape.py -> build/bdist.linux-x86_64/egg/torch2trt/tests\ncopying build/lib/torch2trt/dataset_calibrator_test.py -> build/bdist.linux-x86_64/egg/torch2trt\ncreating build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/BatchNorm2d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/prod.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/Conv1d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/sub.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/layer_norm.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/div.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/permute.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/batch_norm.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/getitem.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/identity.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/__init__.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/AdaptiveAvgPool2d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/view.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/adaptive_max_pool3d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/chunk.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/ne.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/avg_pool.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/mul.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/pow.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/sum.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/normalize.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/max.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/BatchNorm3d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/matmul.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/conv_functional.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/clamp.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/example_plugin.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/gelu.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/interpolate.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/relu6.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/unsqueeze.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/mod.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/sigmoid.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/ConvTranspose2d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/prelu.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/relu.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/compare.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/split.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/tensor.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/pad.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/activation.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/stack.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/silu.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/squeeze.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/expand.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/max_pool1d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/einsum.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/mean.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/Linear.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/AdaptiveAvgPool3d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/min.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/adaptive_max_pool2d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/adaptive_avg_pool2d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/Conv2d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/group_norm.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/BatchNorm1d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/reflection_pad_2d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/dummy_converters.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/softmax.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/adaptive_avg_pool3d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/cat.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/Conv.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/add.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/instance_norm.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/roll.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/unary.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/tanh.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/max_pool3d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/getitem_test.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/clone.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/narrow.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/ConvTranspose.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/LogSoftmax.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/flatten.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/max_pool2d.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/transpose.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/converters/floordiv.py -> build/bdist.linux-x86_64/egg/torch2trt/converters\ncopying build/lib/torch2trt/dataset_calibrator.py -> build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/dataset.py -> build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/dynamic_shape_test.py -> build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/flattener_test.py -> build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/flatten_module.py -> build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/dataset_test.py -> build/bdist.linux-x86_64/egg/torch2trt\ncreating build/bdist.linux-x86_64/egg/torch2trt/contrib\ncopying build/lib/torch2trt/contrib/__init__.py -> build/bdist.linux-x86_64/egg/torch2trt/contrib\ncreating build/bdist.linux-x86_64/egg/torch2trt/contrib/qat\ncopying build/lib/torch2trt/contrib/qat/__init__.py -> build/bdist.linux-x86_64/egg/torch2trt/contrib/qat\ncreating build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/converters\ncopying build/lib/torch2trt/contrib/qat/converters/__init__.py -> build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/converters\ncopying build/lib/torch2trt/contrib/qat/converters/QuantConvBN.py -> build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/converters\ncopying build/lib/torch2trt/contrib/qat/converters/QuantConv.py -> build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/converters\ncopying build/lib/torch2trt/contrib/qat/converters/QuantRelu.py -> build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/converters\ncreating build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/layers\ncopying build/lib/torch2trt/contrib/qat/layers/quant_activation.py -> build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/layers\ncopying build/lib/torch2trt/contrib/qat/layers/__init__.py -> build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/layers\ncopying build/lib/torch2trt/contrib/qat/layers/quant_conv.py -> build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/layers\ncopying build/lib/torch2trt/contrib/qat/layers/_utils.py -> build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/layers\ncopying build/lib/torch2trt/utils.py -> build/bdist.linux-x86_64/egg/torch2trt\ncopying build/lib/torch2trt/flattener.py -> build/bdist.linux-x86_64/egg/torch2trt\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/test.py to test.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/torch2trt.py to torch2trt.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/__init__.py to __init__.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/flatten_module_test.py to flatten_module_test.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/module_test.py to module_test.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/test_tensor_ne.py to test_tensor_ne.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/__init__.py to __init__.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/test_contiguous.py to test_contiguous.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/test_legacy_max_batch_size.py to test_legacy_max_batch_size.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/test_flatten_dynamic.py to test_flatten_dynamic.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/test_tensor_shape_div_batch.py to test_tensor_shape_div_batch.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/timm/__init__.py to __init__.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/timm/test_maxvit.py to test_maxvit.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/test_interpolate_dynamic.py to test_interpolate_dynamic.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/torchvision/save_load.py to save_load.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/torchvision/__init__.py to __init__.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/torchvision/segmentation.py to segmentation.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/torchvision/classification.py to classification.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/tests/test_tensor_shape.py to test_tensor_shape.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/dataset_calibrator_test.py to dataset_calibrator_test.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/BatchNorm2d.py to BatchNorm2d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/prod.py to prod.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/Conv1d.py to Conv1d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/sub.py to sub.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/layer_norm.py to layer_norm.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/div.py to div.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/permute.py to permute.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/batch_norm.py to batch_norm.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/getitem.py to getitem.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/identity.py to identity.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/__init__.py to __init__.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/AdaptiveAvgPool2d.py to AdaptiveAvgPool2d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/view.py to view.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/adaptive_max_pool3d.py to adaptive_max_pool3d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/chunk.py to chunk.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/ne.py to ne.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/avg_pool.py to avg_pool.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/mul.py to mul.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/pow.py to pow.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/sum.py to sum.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/normalize.py to normalize.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/max.py to max.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/BatchNorm3d.py to BatchNorm3d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/matmul.py to matmul.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/conv_functional.py to conv_functional.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/clamp.py to clamp.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/example_plugin.py to example_plugin.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/gelu.py to gelu.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/interpolate.py to interpolate.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/relu6.py to relu6.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/unsqueeze.py to unsqueeze.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/mod.py to mod.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/sigmoid.py to sigmoid.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/ConvTranspose2d.py to ConvTranspose2d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/prelu.py to prelu.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/relu.py to relu.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/compare.py to compare.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/split.py to split.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/tensor.py to tensor.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/pad.py to pad.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/activation.py to activation.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/stack.py to stack.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/silu.py to silu.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/squeeze.py to squeeze.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/expand.py to expand.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/max_pool1d.py to max_pool1d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/einsum.py to einsum.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/mean.py to mean.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/Linear.py to Linear.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/AdaptiveAvgPool3d.py to AdaptiveAvgPool3d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/min.py to min.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/adaptive_max_pool2d.py to adaptive_max_pool2d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/adaptive_avg_pool2d.py to adaptive_avg_pool2d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/Conv2d.py to Conv2d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/group_norm.py to group_norm.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/BatchNorm1d.py to BatchNorm1d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/reflection_pad_2d.py to reflection_pad_2d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/dummy_converters.py to dummy_converters.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/softmax.py to softmax.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/adaptive_avg_pool3d.py to adaptive_avg_pool3d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/cat.py to cat.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/Conv.py to Conv.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/add.py to add.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/instance_norm.py to instance_norm.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/roll.py to roll.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/unary.py to unary.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/tanh.py to tanh.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/max_pool3d.py to max_pool3d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/getitem_test.py to getitem_test.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/clone.py to clone.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/narrow.py to narrow.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/ConvTranspose.py to ConvTranspose.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/LogSoftmax.py to LogSoftmax.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/flatten.py to flatten.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/max_pool2d.py to max_pool2d.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/transpose.py to transpose.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/converters/floordiv.py to floordiv.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/dataset_calibrator.py to dataset_calibrator.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/dataset.py to dataset.cpython-37.pyc\nbuild/bdist.linux-x86_64/egg/torch2trt/dataset.py:61: SyntaxWarning: assertion is always true, perhaps remove parentheses?\n  assert(len(self) > 0, 'Cannot create default flattener without input data.')\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/dynamic_shape_test.py to dynamic_shape_test.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/flattener_test.py to flattener_test.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/flatten_module.py to flatten_module.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/dataset_test.py to dataset_test.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/contrib/__init__.py to __init__.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/__init__.py to __init__.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/converters/__init__.py to __init__.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/converters/QuantConvBN.py to QuantConvBN.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/converters/QuantConv.py to QuantConv.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/converters/QuantRelu.py to QuantRelu.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/layers/quant_activation.py to quant_activation.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/layers/__init__.py to __init__.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/layers/quant_conv.py to quant_conv.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/contrib/qat/layers/_utils.py to _utils.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/utils.py to utils.cpython-37.pyc\nbyte-compiling build/bdist.linux-x86_64/egg/torch2trt/flattener.py to flattener.cpython-37.pyc\ncreating build/bdist.linux-x86_64/egg/EGG-INFO\ncopying torch2trt.egg-info/PKG-INFO -> build/bdist.linux-x86_64/egg/EGG-INFO\ncopying torch2trt.egg-info/SOURCES.txt -> build/bdist.linux-x86_64/egg/EGG-INFO\ncopying torch2trt.egg-info/dependency_links.txt -> build/bdist.linux-x86_64/egg/EGG-INFO\ncopying torch2trt.egg-info/top_level.txt -> build/bdist.linux-x86_64/egg/EGG-INFO\nzip_safe flag not set; analyzing archive contents...\ntorch2trt.contrib.qat.layers.__pycache__._utils.cpython-37: module MAY be using inspect.stack\ncreating 'dist/torch2trt-0.4.0-py3.7.egg' and adding 'build/bdist.linux-x86_64/egg' to it\nremoving 'build/bdist.linux-x86_64/egg' (and everything under it)\nProcessing torch2trt-0.4.0-py3.7.egg\ncreating /opt/conda/lib/python3.7/site-packages/torch2trt-0.4.0-py3.7.egg\nExtracting torch2trt-0.4.0-py3.7.egg to /opt/conda/lib/python3.7/site-packages\n/opt/conda/lib/python3.7/site-packages/torch2trt-0.4.0-py3.7.egg/torch2trt/dataset.py:61: SyntaxWarning: assertion is always true, perhaps remove parentheses?\n  assert(len(self) > 0, 'Cannot create default flattener without input data.')\nRemoving torch2trt 0.4.0 from easy-install.pth file\nAdding torch2trt 0.4.0 to easy-install.pth file\n\nInstalled /opt/conda/lib/python3.7/site-packages/torch2trt-0.4.0-py3.7.egg\nProcessing dependencies for torch2trt==0.4.0\nFinished processing dependencies for torch2trt==0.4.0\nObtaining file:///kaggle/tmp/libs/torch2trt\n  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hInstalling collected packages: torch2trt\n  Attempting uninstall: torch2trt\n    Found existing installation: torch2trt 0.4.0\n    Uninstalling torch2trt-0.4.0:\n      Successfully uninstalled torch2trt-0.4.0\n  Running setup.py develop for torch2trt\nSuccessfully installed torch2trt-0.4.0\n/kaggle/working\nProcessing /kaggle/input/kaggle-rsna-pkgs/nvidia_dali_nightly_cuda110-1.23.0.dev20230210-7260679-py3-none-manylinux2014_x86_64.whl\nRequirement already satisfied: gast<=0.4.0,>=0.2.1 in /opt/conda/lib/python3.7/site-packages (from nvidia-dali-nightly-cuda110==1.23.0.dev20230210) (0.4.0)\nRequirement already satisfied: astunparse>=1.6.0 in /opt/conda/lib/python3.7/site-packages (from nvidia-dali-nightly-cuda110==1.23.0.dev20230210) (1.6.3)\nRequirement already satisfied: six<2.0,>=1.6.1 in /opt/conda/lib/python3.7/site-packages (from astunparse>=1.6.0->nvidia-dali-nightly-cuda110==1.23.0.dev20230210) (1.15.0)\nRequirement already satisfied: wheel<1.0,>=0.23.0 in /opt/conda/lib/python3.7/site-packages (from astunparse>=1.6.0->nvidia-dali-nightly-cuda110==1.23.0.dev20230210) (0.37.1)\nnvidia-dali-nightly-cuda110 is already installed with the same version as the provided wheel. Use --force-reinstall to force an installation of the wheel.\nImport done!\n","output_type":"stream"}]},{"cell_type":"code","source":"!pip install --no-index /kaggle/input/package-rsna/torch-2.3.1cpu.cxx11.abi-cp310-cp310-linux_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2024-08-20T05:25:03.805221Z","iopub.execute_input":"2024-08-20T05:25:03.805843Z","iopub.status.idle":"2024-08-20T05:25:05.735256Z","shell.execute_reply.started":"2024-08-20T05:25:03.805806Z","shell.execute_reply":"2024-08-20T05:25:05.733996Z"},"trusted":true},"execution_count":5,"outputs":[{"name":"stdout","text":"\u001b[31mERROR: torch-2.3.1cpu.cxx11.abi-cp310-cp310-linux_x86_64.whl is not a supported wheel on this platform.\u001b[0m\u001b[31m\n\u001b[0m","output_type":"stream"}]},{"cell_type":"code","source":"import torch\nprint(torch.__version__)\nimport torchvision\nprint(torchvision.__version__)\n","metadata":{"execution":{"iopub.status.busy":"2024-08-20T05:25:05.736766Z","iopub.execute_input":"2024-08-20T05:25:05.737116Z","iopub.status.idle":"2024-08-20T05:25:06.728722Z","shell.execute_reply.started":"2024-08-20T05:25:05.737084Z","shell.execute_reply":"2024-08-20T05:25:06.727701Z"},"trusted":true},"execution_count":6,"outputs":[{"name":"stdout","text":"1.12.1+cu102\n0.12.0\n","output_type":"stream"},{"name":"stderr","text":"/opt/conda/lib/python3.7/site-packages/torchvision/io/image.py:13: UserWarning: Failed to load image Python extension: /opt/conda/lib/python3.7/site-packages/torchvision/image.so: undefined symbol: _ZN5torch3jit17parseSchemaOrNameERKNSt7__cxx1112basic_stringIcSt11char_traitsIcESaIcEEE\n  warn(f\"Failed to load image Python extension: {e}\")\n","output_type":"stream"}]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/tmp/libs/timm')\nsys.path.append('/opt/conda/lib/python3.7/site-packages/torch2trt-0.4.0-py3.7.egg')\nsys.path.append('/kaggle/tmp/libs/torch2trt')\nimport timm\nimport gc\nprint('Timm version:', timm.__version__)\n\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")\nimport os\n\nos.environ['CUDA_MODULE_LOADING'] = 'LAZY'\nimport ctypes\nimport gc\nimport importlib\nimport multiprocessing as mp\nimport shutil\n\nimport albumentations as A\nimport cv2\nimport dicomsdl\nimport numpy as np\nimport nvidia.dali as dali\nimport pandas as pd\nimport pydicom\n\nimport torch\nfrom albumentations.pytorch.transforms import ToTensorV2\nfrom joblib import Parallel, delayed\nfrom nvidia.dali import types\nfrom nvidia.dali.backend import TensorGPU, TensorListGPU\nfrom torch2trt import TRTModule\nfrom torch.nn import functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nimport time","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-08-19T08:36:29.485522Z","iopub.execute_input":"2024-08-19T08:36:29.486522Z","iopub.status.idle":"2024-08-19T08:36:32.211168Z","shell.execute_reply.started":"2024-08-19T08:36:29.486483Z","shell.execute_reply":"2024-08-19T08:36:32.21023Z"},"trusted":true},"execution_count":4,"outputs":[{"name":"stdout","text":"Timm version: 0.8.11dev0\n","output_type":"stream"}]},{"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,"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# KAN Model","metadata":{}},{"cell_type":"code","source":"!pip install timm","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom functools import partial\n\nfrom timm.models.vision_transformer import VisionTransformer, _cfg, Block, Attention\nfrom timm.models._registry import register_model\nfrom timm.models.layers import trunc_normal_","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport math\nfrom typing import *\n\n\nclass SplineLinear(nn.Linear):\n    def __init__(self, in_features: int, out_features: int, init_scale: float = 0.1, **kw) -> None:\n        self.init_scale = init_scale\n        super().__init__(in_features, out_features, bias=False, **kw)\n\n    def reset_parameters(self) -> None:\n        nn.init.xavier_uniform_(self.weight)  # Using Xavier Uniform initialization\n\n\nclass ReflectionalSwitchFunction(nn.Module):\n    def __init__(\n            self,\n            grid_min: float = -2.,\n            grid_max: float = 2.,\n            num_grids: int = 8,\n            exponent: int = 2,\n            denominator: float = 0.33,  # larger denominators lead to smoother basis\n    ):\n        super().__init__()\n        grid = torch.linspace(grid_min, grid_max, num_grids)\n        self.grid = torch.nn.Parameter(grid, requires_grad=False)\n        self.denominator = denominator  # or (grid_max - grid_min) / (num_grids - 1)\n        # self.exponent = exponent\n        self.inv_denominator = 1 / self.denominator  # Cache the inverse of the denominator\n\n    def forward(self, x):\n        diff = (x[..., None] - self.grid)\n        diff_mul = diff.mul(self.inv_denominator)\n        diff_tanh = torch.tanh(diff_mul)\n        diff_pow = -diff_tanh.mul(diff_tanh)\n        diff_pow += 1\n        # diff_pow *= 0.667\n        return diff_pow  # Replace pow with multiplication for squaring\n\n\nclass FasterKANLayer(nn.Module):\n    def __init__(\n            self,\n            input_dim: int,\n            output_dim: int,\n            grid_min: float = -2.,\n            grid_max: float = 2.,\n            num_grids: int = 8,\n            exponent: int = 2,\n            denominator: float = 0.33,\n            use_base_update: bool = True,\n            base_activation=F.silu,\n            spline_weight_init_scale: float = 0.1,\n    ) -> None:\n        super().__init__()\n        self.layernorm = nn.LayerNorm(input_dim)\n        self.rbf = ReflectionalSwitchFunction(grid_min, grid_max, num_grids, exponent, denominator)\n        self.spline_linear = SplineLinear(input_dim * num_grids, output_dim, spline_weight_init_scale)\n        # self.use_base_update = use_base_update\n        # if use_base_update:\n        #    self.base_activation = base_activation\n        #    self.base_linear = nn.Linear(input_dim, output_dim)\n\n    def forward(self, x, time_benchmark=False):\n        if not time_benchmark:\n            spline_basis = self.rbf(self.layernorm(x)).view(x.shape[0], -1)\n            # print(\"spline_basis:\", spline_basis.shape)\n        else:\n            spline_basis = self.rbf(x).view(x.shape[0], -1)\n            # print(\"spline_basis:\", spline_basis.shape)\n        # print(\"-------------------------\")\n        # ret = 0\n        ret = self.spline_linear(spline_basis)\n        # print(\"spline_basis.shape[:-2]:\", spline_basis.shape[:-2])\n        # print(\"*spline_basis.shape[:-2]:\", *spline_basis.shape[:-2])\n        # print(\"spline_basis.view(*spline_basis.shape[:-2], -1):\", spline_basis.view(*spline_basis.shape[:-2], -1).shape)\n        # print(\"ret:\", ret.shape)\n        # print(\"-------------------------\")\n        # if self.use_base_update:\n        # base = self.base_linear(self.base_activation(x))\n        # print(\"self.base_activation(x):\", self.base_activation(x).shape)\n        # print(\"base:\", base.shape)\n        # print(\"@@@@@@@@@\")\n        # ret += base\n        return ret\n\n        # spline_basis = spline_basis.reshape(x.shape[0], -1)  # Reshape to [batch_size, input_dim * num_grids]\n        # print(\"spline_basis:\", spline_basis.shape)\n\n        # spline_weight = self.spline_weight.view(-1, self.spline_weight.shape[0])  # Reshape to [input_dim * num_grids, output_dim]\n        # print(\"spline_weight:\", spline_weight.shape)\n\n        # spline = torch.matmul(spline_basis, spline_weight)  # Resulting shape: [batch_size, output_dim]\n\n        # print(\"-------------------------\")\n        # print(\"Base shape:\", base.shape)\n        # print(\"Spline shape:\", spline.shape)\n        # print(\"@@@@@@@@@\")\n\n\nclass FasterKAN(nn.Module):\n    def __init__(\n            self,\n            layers_hidden: List[int],\n            grid_min: float = -2.,\n            grid_max: float = 2.,\n            num_grids: int = 8,\n            exponent: int = 2,\n            denominator: float = 0.33,\n            use_base_update: bool = True,\n            base_activation=F.silu,\n            spline_weight_init_scale: float = 0.667,\n    ) -> None:\n        super().__init__()\n        self.layers = nn.ModuleList([\n            FasterKANLayer(\n                in_dim, out_dim,\n                grid_min=grid_min,\n                grid_max=grid_max,\n                num_grids=num_grids,\n                exponent=exponent,\n                denominator=denominator,\n                use_base_update=use_base_update,\n                base_activation=base_activation,\n                spline_weight_init_scale=spline_weight_init_scale,\n            ) for in_dim, out_dim in zip(layers_hidden[:-1], layers_hidden[1:])\n        ])\n\n    def forward(self, x):\n        for layer in self.layers:\n            x = layer(x)\n        return x","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class kanBlock(Block):\n\n    def __init__(self, dim, num_heads=8, hdim_kan=192, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,\n                 drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm):\n        super().__init__(dim, num_heads)\n        self.norm1 = norm_layer(dim)\n        self.attn = Attention(\n            dim, num_heads=num_heads, qkv_bias=qkv_bias, attn_drop=attn_drop, proj_drop=drop)\n        # NOTE: drop path for stochastic depth, we shall see if this is better than dropout here\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n        self.norm2 = norm_layer(dim)\n        mlp_hidden_dim = int(dim * mlp_ratio)\n        # self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)\n        self.kan = FasterKAN([dim, hdim_kan, dim])\n\n    def forward(self, x):\n        b, t, d = x.shape\n        x = x + self.drop_path(self.attn(self.norm1(x)))\n        x = x + self.drop_path(self.kan(self.norm2(x).reshape(-1, x.shape[-1])).reshape(b, t, d))\n        return x","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VisionKAN(nn.Module):\n    def __init__(self, model_name=\"deit_base_patch16_384_KAN\", pretrained=True, out_dim=1, hdim_kan=192):\n        super().__init__()\n        self.backbone = VisionTransformer(\n            img_size=(1536,1024), patch_size=16, in_chans=3, num_classes=1, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4., qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6)\n        )\n        \n        # Modify the blocks to use kanBlock if required\n        self.backbone.blocks = nn.ModuleList([\n            kanBlock(dim=self.backbone.embed_dim, num_heads=12, hdim_kan=hdim_kan)\n            for _ in range(12)\n        ])\n\n        # Load pretrained weights if required\n        if pretrained:\n            checkpoint = torch.hub.load_state_dict_from_url(\n                url=\"https://dl.fbaipublicfiles.com/deit/deit_base_patch16_384-8de9b5d1.pth\",\n                map_location=\"cpu\", check_hash=True\n            )\n            self.backbone.load_state_dict(checkpoint, strict=False)\n\n        # Adjust the classifier to match the output dimensions\n        self.backbone.head = nn.Linear(self.backbone.head.in_features, out_dim)\n\n    def forward(self, x):\n        x = self.backbone.patch_embed(x)\n        cls_token = self.backbone.cls_token.expand(x.shape[0], -1, -1)\n        x = torch.cat((cls_token, x), dim=1)\n        x = x + self.backbone.pos_embed\n        x = self.backbone.pos_drop(x)\n\n        for block in self.backbone.blocks:\n            x = block(x)\n\n        x = self.backbone.norm(x)\n        x = self.backbone.head(x[:, 0])\n        return x","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Classification model (4 x ConvNext-small ensemble)","metadata":{}},{"cell_type":"markdown","source":"","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":{"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":{"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 = [1536, 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":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Windowing","metadata":{}},{"cell_type":"code","source":"from nvidia.dali import types\n\nDALI2TORCH_TYPES = {\n    types.FLOAT: torch.float32,\n    types.FLOAT64: torch.float64,\n    types.FLOAT16: torch.float16,\n    # Thêm các loại khác nếu cần\n}\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"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":{"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":{"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\nfrom nvidia.dali.types import DALIDataType\nfrom nvidia.dali import types as dali_types\n################################################################################\n# DALI2TORCH_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# }\nDALI2TORCH_TYPES = {\n    dali_types.FLOAT: torch.float32,\n    dali_types.FLOAT64: torch.float64,\n    dali_types.FLOAT16: torch.float16,\n    dali_types.UINT8: torch.uint8,\n    dali_types.INT8: torch.int8,\n    dali_types.UINT16: torch.int16,\n    dali_types.INT16: torch.int16,\n    dali_types.INT32: torch.int32,\n    dali_types.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\n# def 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#     cuda_steam=torch.cuda.current_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\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\n    cuda_stream = torch.cuda.current_stream() if cuda_stream is None else 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 = ctypes.c_void_p(cuda_stream.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\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=1,\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":{"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":{"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":"# Config","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch\nimport torch.nn as nn\nfrom functools import partial\nimport logging\n_logger = logging.getLogger(__name__)\n\n\nclass Config:\n    def __init__(self):\n        # General configuration\n        self.fold_idx = 1\n        self.num_sched_epochs = 10\n        self.num_epochs = 35 #35\n        self.start_ratio = 0.1429\n        self.end_ratio = 0.1429\n        self.one_pos_mode = True\n        self.seed = 42\n        self.device = torch.device('cuda:1' if torch.cuda.is_available() else 'cpu')\n        self.output_dir = './mnt/primary/RSNA/checkpoint'  # Output directory for model checkpoints and logs\n        self.save_results_dir='./mnt/primary/RSNA/results'\n\n        # Data configuration\n        self.data = None\n        self.dataset = ''  # Dataset type and name\n        self.data_dir = 'datasets'  # Path to dataset directory\n        self.train_split = 'train'  # Dataset train split\n        self.val_split = 'validation'  # Dataset validation split\n        self.class_map = ''  # Path to class to idx mapping file\n        self.dataset_download = False  # Allow download of dataset if available\n        self.input_size = [3, 1536, 1024]  # [C, H, W] - Number of channels, height, width\n        self.img_size = (1536, 1024)  # Image size (height and width)\n        self.batch_size = 1  # Batch size for training and validation\n        self.validation_batch_size = 2  # Validation batch size override\n        self.num_aug_splits = 0  # Number of augmentation splits\n        self.train_interpolation = 'random'  # Interpolation method for training\n        self.prefetcher = False  # Use fast prefetcher for data loading\n        self.no_aug = False  # Disable all training augmentations\n        self.reprob = 0.0  # Probability for random erasing\n        self.remode = 'pixel'  # Mode for random erasing\n        self.recount = 1  # Number of areas to erase\n        self.resplit = False  # Use different erasing in each augmentation split\n        self.scale = [0.08, 1.0]  # Scale range for random resized crop\n        self.ratio = [3./4., 4./3.]  # Aspect ratio range for random resized crop\n        self.hflip = 0.5  # Probability of horizontal flip\n        self.vflip = 0.0  # Probability of vertical flip\n        self.color_jitter = 0.4  # Factor for color jittering\n        self.aa = None  # AutoAugment policy\n        self.aug_repeats = 0  # Number of augmentation repetitions\n        self.workers = 8  # Number of data loader workers\n        self.distributed = False # Whether to use distributed training\n        self.pin_mem = False  # Pin CPU memory in DataLoader\n        self.use_multi_epochs_loader = False  # Use multi-epochs data loader\n        self.worker_seeding = 'all'  # Worker seeding mode\n\n        # Model configuration\n        self.model = \"VisionKAN\"  #model name to use\n        self.pretrained = False  # Use pretrained model\n        self.initial_checkpoint = ''  # Path to initialize model from checkpoint\n        self.resume = ''  # Path to resume model and optimizer state from checkpoint\n        self.no_resume_opt = False  # Prevent resume of optimizer state\n        self.num_classes = 1  # Number of label classes\n        self.gp = 'max'  # Global pooling type\n        self.in_chans = 3  # Number of image input channels\n        self.crop_pct = 1  # Input image center crop percent\n        self.mean = [0.485, 0.456, 0.406]  # Mean pixel value of dataset\n        self.std = [0.229, 0.224, 0.225]  # Std deviation of dataset\n        self.interpolation = 'bilinear'  # Image resize interpolation type\n        self.channels_last = False  # Use channels_last memory layout\n        self.fuser = 'nvfuser'  # JIT fuser\n        self.grad_checkpointing = False  # Enable gradient checkpointing\n        self.fast_norm = False  # Enable experimental fast norm\n        self.model_ema = True  # Enable exponential moving average of model weights\n        self.model_ema_force_cpu = False  # Force EMA to be tracked on CPU\n        self.model_ema_decay = 0.9998  # Decay factor for EMA\n        self.torchscript = False\n        self.torchcompile = True\n\n        # Optimizer configuration\n        self.opt = 'sgd'  # Optimizer type\n        self.opt_eps = None  # Optimizer epsilon\n        self.opt_betas = None  # Optimizer betas\n        self.momentum = 0.9  # Optimizer momentum\n        self.weight_decay = 2e-5  # Weight decay\n        self.clip_grad = None  # Clip gradient norm\n        self.clip_mode = 'norm'  # Gradient clipping mode\n        self.layer_decay = None  # Layer-wise learning rate decay\n\n        # Learning rate schedule\n        self.sched = 'cosine'  # Learning rate scheduler\n        self.sched_on_updates = False  # Apply LR scheduler step on update instead of epoch end\n        self.lr = 3e-3  # Learning rate\n        self.lr_base = 0.1  # Base learning rate\n        self.lr_base_size = 256  # Base learning rate batch size divisor\n        self.lr_base_scale = ''  # Base learning rate vs batch_size scaling\n        self.lr_noise = None  # Learning rate noise on/off epoch percentages\n        self.lr_noise_pct = 0.67  # Learning rate noise limit percent\n        self.lr_noise_std = 1.0  # Learning rate noise std-dev\n        self.lr_cycle_mul = 1.0  # Learning rate cycle len multiplier\n        self.lr_cycle_decay = 0.5  # Amount to decay each learning rate cycle\n        self.lr_cycle_limit = 1  # Learning rate cycle limit\n        self.lr_k_decay = 1.0  # Learning rate k-decay for cosine/poly\n        self.warmup_lr = 3e-5  # Warmup learning rate\n        self.min_lr = 5e-5  # Lower LR bound for cyclic schedulers\n        self.epochs = 30 #30  # Number of epochs to train\n        self.epoch_repeats = 0.  # Epoch repeat multiplier\n        self.start_epoch = None  # Manual epoch number (useful on restarts)\n        self.decay_milestones = [90, 180, 270]  # Decay epoch indices\n        self.decay_epochs = 90  # Epoch interval to decay LR\n        self.warmup_epochs = 4  # Epochs to warmup LR\n        self.warmup_prefix = False  # Exclude warmup period from decay schedule\n        self.cooldown_epochs = 1  # Epochs to cooldown LR at min_lr\n        self.patience_epochs = 10  # Patience epochs for Plateau LR scheduler\n        self.decay_rate = 0.1  # LR decay rate\n\n        # Augmentation & regularization\n        self.no_aug = False  # Disable all training augmentations\n        self.scale = [0.08, 1.0]  # Random resize scale\n        self.ratio = [3./4., 4./3.]  # Random resize aspect ratio\n        self.hflip = 0.5  # Horizontal flip probability\n        self.vflip = 0.0  # Vertical flip probability\n        self.color_jitter = 0.4  # Color jitter factor\n        self.aa = None  # AutoAugment policy\n        self.aug_repeats = 0  # Number of augmentation repetitions\n        self.aug_splits = 0  # Number of augmentation splits\n        self.jsd_loss = False  # Enable Jensen-Shannon Divergence + CE loss\n        self.mixup_active = False  # Enable Mixup/CutMix use\n        self.bce_loss = True  # Enable BCE loss w/ Mixup/CutMix use\n        self.bce_target_thresh = None  # Threshold for binarizing softened BCE targets\n        self.reprob = 0.  # Random erase probability\n        self.remode = 'pixel'  # Random erase mode\n        self.recount = 1  # Random erase count\n        self.resplit = False  # Do not random erase first augmentation split\n        self.mixup = 0.0  # Mixup alpha, mixup enabled if > 0.\n        self.cutmix = 0.0  # Cutmix alpha, cutmix enabled if > 0.\n        self.cutmix_minmax = None  # Cutmix min/max ratio\n        self.mixup_prob = 1.0  # Probability of performing mixup or cutmix\n        self.mixup_switch_prob = 0.5  # Probability of switching to cutmix\n        self.mixup_mode = 'batch'  # How to apply mixup/cutmix paramseval-metric\n        self.mixup_off_epoch = 0  # Turn off mixup after this epoch\n        self.smoothing = 0.1  # Label smoothing\n        self.train_interpolation = 'random'  # Training interpolation\n        self.drop = 0.5  # Dropout rate\n        self.drop_connect = None  # Drop connect rate\n        self.drop_path = 0.2  # Drop path rate\n        self.drop_block = None  # Drop block rate\n\n        # Batch norm parameters\n        self.bn_momentum = None  # BatchNorm momentum override\n        self.bn_eps = None  # BatchNorm epsilon override\n        self.sync_bn = False  # Enable synchronized BatchNorm\n        self.dist_bn = 'reduce'  # Distribute BatchNorm stats between nodes\n        self.split_bn = False  # Enable separate BN layers per augmentation split\n\n        # Exponential Moving Average (EMA)\n        self.model_ema = True  # Enable tracking moving average of model weights\n        self.model_ema_force_cpu = False  # Force EMA to be tracked on CPU\n        self.model_ema_decay = 0.9998  # Decay factor for model weights moving average\n\n        # Miscellaneous\n        self.seed = 42  # Random seed\n        self.worker_seeding = 'all'  # Worker seed mode\n        self.log_interval = 500  # How many batches to wait before logging training status\n        self.recovery_interval = 0  # How many batches to wait before writing recovery checkpoint\n        self.checkpoint_hist = 100  # Number of checkpoints to keep\n        self.save_images = True  # Save images of input batches every log interval for debugging\n        self.amp = True  # Use NVIDIA Apex AMP or Native AMP for mixed precision training\n        self.amp_dtype = 'float16'  # Lower precision AMP dtype\n        self.amp_impl = 'native'  # AMP implementation to use\n        self.no_ddp_bb = False  # Force broadcast buffers for native DDP to off\n        self.pin_mem = True  # Pin CPU memory in DataLoader\n        self.no_prefetcher = False  # Disable fast prefetcher\n        self.output = ''  # Path to output folder\n        self.experiment_name = ''  # Name of train experiment\n        self.eval_metric = 'gbmean_best_pfbeta'  # Best metric to evaluate\n        self.tta = 0  # Test/inference time augmentation factor\n        self.local_rank = 0  # Local rank for distributed training\n        self.use_multi_epochs_loader = False  # Use the multi-epochs-loader\n        self.log_wandb = True  # Log training and validation metrics to wandb\n\n        # Custom additions\n        self.pos_weight = 0.9  # Positive weight used for loss computation\n        self.dense_ckpt_epochs = [10, 18]  # Dense checkpointing epochs\n        self.dense_ckpt_bins = 2  # Bins for dense checkpointing\n        self.exp_kwargs = {\n            \n        }  # Extra keyword arguments for the experiment\n\n        # New additions based on missing parameters\n        self.updates_per_epoch = 0  # Number of updates per epoch, to be calculated later\n        self.fold_idx=1","metadata":{"execution":{"iopub.status.busy":"2024-08-19T08:37:29.245661Z","iopub.execute_input":"2024-08-19T08:37:29.24612Z","iopub.status.idle":"2024-08-19T08:37:29.284797Z","shell.execute_reply.started":"2024-08-19T08:37:29.246085Z","shell.execute_reply":"2024-08-19T08:37:29.283764Z"},"trusted":true},"execution_count":7,"outputs":[]},{"cell_type":"code","source":"import torch\nprint(torch.__version__)","metadata":{"execution":{"iopub.status.busy":"2024-08-19T08:36:43.840218Z","iopub.execute_input":"2024-08-19T08:36:43.840875Z","iopub.status.idle":"2024-08-19T08:36:43.846784Z","shell.execute_reply.started":"2024-08-19T08:36:43.840829Z","shell.execute_reply":"2024-08-19T08:36:43.845638Z"},"trusted":true},"execution_count":5,"outputs":[{"name":"stdout","text":"1.12.1+cu102\n","output_type":"stream"}]},{"cell_type":"code","source":"import pickle\npath = '/kaggle/working/tmp.pth/data.pkl'\nwith open(path, 'rb') as file:\n    model = pickle.load(file)","metadata":{"execution":{"iopub.status.busy":"2024-08-19T08:44:39.889297Z","iopub.execute_input":"2024-08-19T08:44:39.889767Z","iopub.status.idle":"2024-08-19T08:44:39.919116Z","shell.execute_reply.started":"2024-08-19T08:44:39.88973Z","shell.execute_reply":"2024-08-19T08:44:39.917617Z"},"trusted":true},"execution_count":11,"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mUnpicklingError\u001b[0m                           Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_3834/3100238960.py\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m      2\u001b[0m \u001b[0mpath\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;34m'/kaggle/working/tmp.pth/data.pkl'\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m      3\u001b[0m \u001b[0;32mwith\u001b[0m \u001b[0mopen\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mpath\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m'rb'\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;32mas\u001b[0m \u001b[0mfile\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 4\u001b[0;31m     \u001b[0mmodel\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpickle\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mload\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mfile\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m","\u001b[0;31mUnpicklingError\u001b[0m: A load persistent id instruction was encountered,\nbut no persistent_load function was specified."],"ename":"UnpicklingError","evalue":"A load persistent id instruction was encountered,\nbut no persistent_load function was specified.","output_type":"error"}]},{"cell_type":"code","source":"path = '/kaggle/working/tmp.pth/data.pkl'\n#model = VisionKAN(pretrained=True, out_dim=1)\nimport types\nimport sys\nsys.modules['config'] = types.ModuleType('config')\nsys.modules['config'].Config = Config()\ncheckpoint = torch.load(path)\nprint(checkpoint.keys())","metadata":{"execution":{"iopub.status.busy":"2024-08-19T08:43:48.598915Z","iopub.execute_input":"2024-08-19T08:43:48.60003Z","iopub.status.idle":"2024-08-19T08:43:48.645757Z","shell.execute_reply.started":"2024-08-19T08:43:48.599986Z","shell.execute_reply":"2024-08-19T08:43:48.644421Z"},"trusted":true},"execution_count":10,"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mUnpicklingError\u001b[0m                           Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_3834/3964636808.py\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m      5\u001b[0m \u001b[0msys\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmodules\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m'config'\u001b[0m\u001b[0;34m]\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtypes\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mModuleType\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m'config'\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m      6\u001b[0m \u001b[0msys\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmodules\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m'config'\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mConfig\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mConfig\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 7\u001b[0;31m \u001b[0mcheckpoint\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mload\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mpath\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m      8\u001b[0m \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mcheckpoint\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mkeys\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/opt/conda/lib/python3.7/site-packages/torch/serialization.py\u001b[0m in \u001b[0;36mload\u001b[0;34m(f, map_location, pickle_module, **pickle_load_args)\u001b[0m\n\u001b[1;32m    711\u001b[0m                     \u001b[0;32mreturn\u001b[0m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mjit\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mload\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mopened_file\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    712\u001b[0m                 \u001b[0;32mreturn\u001b[0m \u001b[0m_load\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mopened_zipfile\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mmap_location\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mpickle_module\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mpickle_load_args\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 713\u001b[0;31m         \u001b[0;32mreturn\u001b[0m \u001b[0m_legacy_load\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mopened_file\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mmap_location\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mpickle_module\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mpickle_load_args\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    714\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    715\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/opt/conda/lib/python3.7/site-packages/torch/serialization.py\u001b[0m in \u001b[0;36m_legacy_load\u001b[0;34m(f, map_location, pickle_module, **pickle_load_args)\u001b[0m\n\u001b[1;32m    918\u001b[0m             \"functionality.\")\n\u001b[1;32m    919\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 920\u001b[0;31m     \u001b[0mmagic_number\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpickle_module\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mload\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mf\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mpickle_load_args\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    921\u001b[0m     \u001b[0;32mif\u001b[0m \u001b[0mmagic_number\u001b[0m \u001b[0;34m!=\u001b[0m \u001b[0mMAGIC_NUMBER\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    922\u001b[0m         \u001b[0;32mraise\u001b[0m \u001b[0mRuntimeError\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"Invalid magic number; corrupt file?\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;31mUnpicklingError\u001b[0m: A load persistent id instruction was encountered,\nbut no persistent_load function was specified."],"ename":"UnpicklingError","evalue":"A load persistent id instruction was encountered,\nbut no persistent_load function was specified.","output_type":"error"}]},{"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    import types\n    sys.modules['config'] = types.ModuleType('config')\n    sys.modules['config'].Config = Config()\n\n    # Giả lập module config\n    #sys.modules['config'] = types.ModuleType('config')\n\n#     path='/kaggle/input/checkpoint-visionkan-v1/model_best.pth.tar'\n#     model= VisionKAN(pretrained=False, out_dim=1)\n#     checkpoint = torch.load(path,map_location=torch.device('cuda'))\n#     model.load_state_dict(checkpoint['model_state_dict'],strict=False)\n#     model.cuda()\n#     model.eval()\n    path = '/kaggle/input/checkpoint-visionkan-v1/model_best.pth.tar'\n    model = VisionKAN(pretrained=True, out_dim=1)\n    checkpoint = torch.load(path)\n    #print(checkpoint.keys())\n    model.load_state_dict(checkpoint['model_state_dict'], strict=False)\n    model = model.cuda()\n    model.eval()\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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"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":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}