{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Hello fellow Kagglers,\n\nThis notebook demonstrates the training process on a TPU in Tensorflow.\n\nThanks to the use of a [TPU (Tensor Processing Unit)](https://cloud.google.com/tpu) training takes about an hour.\n\nThe TFREcord dataset contains cropped images sized 1344x768, created in [this notebook](https://www.kaggle.com/code/markwijkhuizen/rsna-preprocessing-tfrecords-640x512-dataset).\n\n20% of the data is used for validation, which reaches ~0.20 pF1 with the best threshold.\n\n**Things that did not work for me:**\n\n* [SigmoidFocalCrossEntropy](https://www.tensorflow.org/addons/api_docs/python/tfa/losses/SigmoidFocalCrossEntropy)\n* Increasing model size to for example EfficientNetV2S\n\n**Things that did work for me:**\n\n* Class weights: give minority class weight of 10\n* Training on TPU instead of GPU: larger batch size (16x2->16x8) giving larger probability of having positive sample in batch\n* Cropping Images\n* Using Cropped Image Ratio\n\nI enjoy this competition and will update this notebook frequently, stay tuned!\n\n**V2**\n\n* Cropped images in 1344x768 resolution\n* EfficientNetV2T\n* Added augmentations\n* Single image modal instead of both CC and MLO views as input\n\n**V4**\n\n* Switch to modern ConvNextV2 models [paper: ConvNeXt V2: Co-designing and Scaling ConvNets with Masked Autoencoders](https://arxiv.org/pdf/2301.00808.pdf)\n* Reduced learning rate 1e-4xN_REPLICAS -> 5e-6xN_REPLICAS\n* Adam -> AdamW optimizer with 0.01 Weight Decay ratio\n* Removed Warmup Epochs\n\n**V5**\n\nThis will be the final update of my notebooks for this competitions, which should achieve a LB score in the low 0.50s. I will continue participating in this competition, however, I will not share my progress anymore.\n\n* Mainly updated preprocessing strategy, see [Preprocessing Notebook](https://www.kaggle.com/code/markwijkhuizen/rsna-cropped-tfrecords-768x1344-dataset)\n\n[RSNA Cropped TFRecords 768x1344 Dataset](https://www.kaggle.com/code/markwijkhuizen/rsna-cropped-tfrecords-768x1344-dataset)\n\n[Inference Notebook](https://www.kaggle.com/markwijkhuizen/rsna-efficientnetv2-inference-tensorflow)\n\nGood luck to all of you in the last month of this excisting competition!","metadata":{"id":"NU_zvhoTAcCe"}},{"cell_type":"markdown","source":"# Kaggle <->Colab settings\n","metadata":{"id":"z2iOzJ3o8UCL"}},{"cell_type":"code","source":"import os\nimport json\n# config\ndef set_kaggle_api_cfg():\n    '''\n    # API example of use\n    * submit\n    `!kaggle competitions submit -c open-problems-multimodal -f /content/drive/MyDrive/kaggle/single-cell-competition/data/output/submission.csv -m \"Message\"`\n\n    * check_submissions\n    `!kaggle competitions submissions -c open-problems-multimodal`\n\n    * show kernel list\n    `!kaggle kernels list`\n    '''\n    f = open(\"/content/drive/MyDrive/Takus-folder/api/kaggle.json\", 'r')\n    json_data = json.load(f)\n    os.environ['KAGGLE_USERNAME'] = json_data['username']\n    os.environ['KAGGLE_KEY'] = json_data['key']\n# Common\nIS_KAGGLE = False\ntry:\n    !pip install kaggle\n    from google.colab import drive\n    drive.mount(\"/content/drive\")\n    # Kaggle API\n    set_kaggle_api_cfg()\nexcept:\n    IS_KAGGLE = True\n   \n\nif not IS_KAGGLE:\n    print('Running in Colab...')\n    RSNA_2022_PATH = '/content/drive/MyDrive/Dataset/rsna-breast-cancer-detection'\n    TRAIN_IMAGES_PATH = ''\n    MODELS_PATH = ''\n","metadata":{"id":"adc2RSGTRI4Y","outputId":"8d3824e9-b792-466a-d41d-efc0f4ee1262","execution":{"iopub.status.busy":"2023-02-25T05:40:12.975673Z","iopub.execute_input":"2023-02-25T05:40:12.976317Z","iopub.status.idle":"2023-02-25T05:40:23.334968Z","shell.execute_reply.started":"2023-02-25T05:40:12.976166Z","shell.execute_reply":"2023-02-25T05:40:23.333821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Install ConvNextV2 Models From Keras CV Attention Models Pip Package\nif IS_KAGGLE:\n    !pip install -qq /kaggle/input/keras-cv-attention-models/keras_cv_attention_models-1.3.9-py3-none-any.whl\nelse:\n    !pip install -qq /content/drive/MyDrive/Dataset/keras-cv-attention-models/keras_cv_attention_models-1.3.9-py3-none-any.whl\n","metadata":{"id":"WxY-IcNwPox1","outputId":"3a8c1100-bb07-4447-e349-2e8e53847c8b","execution":{"iopub.status.busy":"2023-02-25T05:40:23.3384Z","iopub.execute_input":"2023-02-25T05:40:23.338835Z","iopub.status.idle":"2023-02-25T05:40:32.542825Z","shell.execute_reply.started":"2023-02-25T05:40:23.338778Z","shell.execute_reply":"2023-02-25T05:40:32.541619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nimport matplotlib.pyplot as plt\nimport matplotlib as mpl\n\nfrom tqdm.notebook import tqdm\nfrom multiprocessing import cpu_count\nfrom sklearn.model_selection import train_test_split\nfrom keras_cv_attention_models import convnext\n\nimport os\nimport time\nimport pickle\nimport math\nimport random\nimport sys\nimport cv2\nimport gc\nimport datetime\n\nprint(f'Tensorflow Version: {tf.__version__}')\nprint(f'Python Version: {sys.version}')","metadata":{"id":"u4TV-5q4Pox4","outputId":"9e7f8ea7-29de-42e9-d5ee-6cc05ff9923f","execution":{"iopub.status.busy":"2023-02-25T05:40:32.544412Z","iopub.execute_input":"2023-02-25T05:40:32.544746Z","iopub.status.idle":"2023-02-25T05:40:40.8241Z","shell.execute_reply.started":"2023-02-25T05:40:32.5447Z","shell.execute_reply":"2023-02-25T05:40:40.822406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# GCS config","metadata":{"id":"p1YbczzHHH_I"}},{"cell_type":"code","source":"# For TPU's the dataset needs to be stored in Google Cloud\n# Retrieve the Google Cloud location of the dataset\n\nif IS_KAGGLE:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    user_credential = user_secrets.get_gcloud_credential()\n    user_secrets.set_tensorflow_credential(user_credential)\n    from kaggle_datasets import KaggleDatasets\n    GCS_DS_PATH = KaggleDatasets().get_gcs_path('rsna-preprocessing-tfrecords-640x512-dataset-pub')\nelse:\n    # GCS settings for TPU\n    from google.colab import auth\n    auth.authenticate_user()\n    BUCKET_NAME = 'kds-800015fc04766f4d5026ebdcfa0f16de8f881ae5b9225eac4dfbd526'\n    \n    MOUNT_TO = 'gcs-dataset'\n    os.makedirs('gcs-dataset',exist_ok=True)\n\n    ! gcsfuse --implicit-dirs --limit-bytes-per-sec -1 --limit-ops-per-sec -1 {BUCKET_NAME} {MOUNT_TO}\n    %ls {MOUNT_TO}\n\n    !echo \"deb http://packages.cloud.google.com/apt gcsfuse-bionic main\" > /etc/apt/sources.list.d/gcsfuse.list\n    !curl https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add -\n    !apt update\n    !apt install gcsfuse\n    GCS_DS_PATH = 'gs://'+BUCKET_NAME\n    # GCS_DS_PATH = KaggleDatasets().get_gcs_path('rsna-preprocessing-tfrecords-garudakai-stkfold')\n\nprint(GCS_DS_PATH)","metadata":{"id":"07FCjETeAcCw","outputId":"25e39a44-9b11-45f7-d2c5-e646330515bc","execution":{"iopub.status.busy":"2023-02-25T05:40:40.825564Z","iopub.execute_input":"2023-02-25T05:40:40.825952Z","iopub.status.idle":"2023-02-25T05:40:43.360644Z","shell.execute_reply.started":"2023-02-25T05:40:40.825903Z","shell.execute_reply":"2023-02-25T05:40:43.35962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{"id":"7VBY7q5jAcCv"}},{"cell_type":"code","source":"now = datetime.datetime.now().strftime(\"%d-%b-%Y %H-%M-%S\")\nnp.save(now, np.array([now]))","metadata":{"id":"VQkB6nZ6AcCr","execution":{"iopub.status.busy":"2023-02-25T05:40:43.363823Z","iopub.execute_input":"2023-02-25T05:40:43.364232Z","iopub.status.idle":"2023-02-25T05:40:43.370495Z","shell.execute_reply.started":"2023-02-25T05:40:43.364154Z","shell.execute_reply":"2023-02-25T05:40:43.369325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# float32 or mixed_float16 (mixed precision: compute float16, variable float32)\n# TPU is fast enough and has enough memory to use float32\npolicy = tf.keras.mixed_precision.Policy('float32')\ntf.keras.mixed_precision.set_global_policy(policy)\n\nprint(f'Compute dtype: {tf.keras.mixed_precision.global_policy().compute_dtype}')\nprint(f'Variable dtype: {tf.keras.mixed_precision.global_policy().variable_dtype}')","metadata":{"id":"ScOLEc60AcCt","outputId":"2b56c570-9362-49d2-8c3a-eb2f54db0cc0","execution":{"iopub.status.busy":"2023-02-25T05:40:43.371982Z","iopub.execute_input":"2023-02-25T05:40:43.372352Z","iopub.status.idle":"2023-02-25T05:40:43.382935Z","shell.execute_reply.started":"2023-02-25T05:40:43.372308Z","shell.execute_reply":"2023-02-25T05:40:43.381957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# MatplotLib Global Settings\nmpl.rcParams.update(mpl.rcParamsDefault)\nmpl.rcParams['xtick.labelsize'] = 16\nmpl.rcParams['ytick.labelsize'] = 16\nmpl.rcParams['axes.labelsize'] = 18\nmpl.rcParams['axes.titlesize'] = 24","metadata":{"id":"9lbe0N0jAcCu","execution":{"iopub.status.busy":"2023-02-25T05:40:43.384283Z","iopub.execute_input":"2023-02-25T05:40:43.385116Z","iopub.status.idle":"2023-02-25T05:40:43.395824Z","shell.execute_reply.started":"2023-02-25T05:40:43.385075Z","shell.execute_reply":"2023-02-25T05:40:43.394919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Detect hardware, return appropriate distribution strategy\ntry:\n    TPU = tf.distribute.cluster_resolver.TPUClusterResolver()  # TPU detection. No parameters necessary if TPU_NAME environment variable is set. On Kaggle this is always the case.\n    print('Running on TPU ', TPU.master())\nexcept ValueError:\n    print('Running on GPU')\n    TPU = None\n\nif TPU:\n    IS_TPU = True\n    tf.config.experimental_connect_to_cluster(TPU)\n    tf.tpu.experimental.initialize_tpu_system(TPU)\n    STRATEGY = tf.distribute.experimental.TPUStrategy(TPU)\nelse:\n    IS_TPU = False\n    STRATEGY = tf.distribute.get_strategy() # default distribution strategy in Tensorflow. Works on CPU and single GPU.\n\nN_REPLICAS = STRATEGY.num_replicas_in_sync\nprint(f'N_REPLICAS: {N_REPLICAS}, IS_TPU: {IS_TPU}')","metadata":{"id":"peIPXUnnAcCv","outputId":"c0be7e73-a4f0-4de7-991c-360fd03cd920","execution":{"iopub.status.busy":"2023-02-25T05:40:43.397302Z","iopub.execute_input":"2023-02-25T05:40:43.39834Z","iopub.status.idle":"2023-02-25T05:40:43.41564Z","shell.execute_reply.started":"2023-02-25T05:40:43.398287Z","shell.execute_reply":"2023-02-25T05:40:43.414379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 42\nDEBUG = False\n\n# Is Interactive Flag and COrresponding Verbosity Method\nif IS_KAGGLE:\n    IS_INTERACTIVE = os.environ['KAGGLE_KERNEL_RUN_TYPE'] == 'Interactive'\nelse:\n    IS_INTERACTIVE = True\n    \n# Image dimensions\nIMG_HEIGHT = 1344\nIMG_WIDTH = 768\nN_CHANNELS = 1\nINPUT_SHAPE = (IMG_HEIGHT, IMG_WIDTH, 1)\nN_SAMPLES_TFRECORDS = 548\n\n# Peak Learning Rate\nLR_MAX = 5e-6 * N_REPLICAS\nWD_RATIO = 0.01\n\nN_WARMUP_EPOCHS = 0\nN_EPOCHS= 3 if not IS_INTERACTIVE else 1\n\n# Batch size\nBATCH_SIZE = 8 * N_REPLICAS\n\n\nVERBOSE = 1 if IS_INTERACTIVE else 2\n\n# Tensorflow AUTO flag\nAUTO = tf.data.experimental.AUTOTUNE\n\nprint(f'BATCH_SIZE: {BATCH_SIZE}')","metadata":{"id":"tJpc_VZwAcCw","outputId":"983da3da-0e74-4c48-bc03-c8ebaf5b65b6","execution":{"iopub.status.busy":"2023-02-25T05:40:43.417017Z","iopub.execute_input":"2023-02-25T05:40:43.418155Z","iopub.status.idle":"2023-02-25T05:40:43.427785Z","shell.execute_reply.started":"2023-02-25T05:40:43.41811Z","shell.execute_reply":"2023-02-25T05:40:43.426838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Seed","metadata":{"id":"dQLZnx-mAcCx"}},{"cell_type":"code","source":"# Seed all random number generators\ndef seed_everything(seed=SEED):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n\nseed_everything()","metadata":{"id":"ueWTA2pEAcCx","execution":{"iopub.status.busy":"2023-02-25T05:40:43.429262Z","iopub.execute_input":"2023-02-25T05:40:43.429807Z","iopub.status.idle":"2023-02-25T05:40:43.439028Z","shell.execute_reply.started":"2023-02-25T05:40:43.429769Z","shell.execute_reply":"2023-02-25T05:40:43.437936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utility Functions","metadata":{"id":"RkKCy5IkAcCx"}},{"cell_type":"code","source":"# short Tensorflow randin integer function\ndef tf_rand_int(minval, maxval, dtype=tf.int64):\n    minval = tf.cast(minval, dtype)\n    maxval = tf.cast(maxval, dtype)\n    return tf.random.uniform(shape=(), minval=minval, maxval=maxval, dtype=dtype)\n\n# chance of 1 in k\ndef one_in(k):\n    return 0 == tf_rand_int(0, k)","metadata":{"id":"TVc71WxpAcCy","execution":{"iopub.status.busy":"2023-02-25T05:40:43.44074Z","iopub.execute_input":"2023-02-25T05:40:43.441009Z","iopub.status.idle":"2023-02-25T05:40:43.455808Z","shell.execute_reply.started":"2023-02-25T05:40:43.440982Z","shell.execute_reply":"2023-02-25T05:40:43.454667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"id":"gYtszgQeAcCy"}},{"cell_type":"code","source":"# Function to benchmark the dataset\ndef benchmark_dataset(dataset, num_epochs=3, n_steps_per_epoch=10, bs=BATCH_SIZE):\n    start_time = time.perf_counter()\n    for epoch_num in range(num_epochs):\n        for idx, (inputs, labels) in enumerate(dataset.take(n_steps_per_epoch + 1)):\n            if idx == 0:\n                epoch_start = time.perf_counter()\n            elif idx == 1 and epoch_num == 0:\n                image = inputs['image']\n                print(f'image shape: {image.shape}, labels shape: {labels.shape}, image dtype: {image.dtype}, labels dtype: {labels.dtype}')\n            else:\n                pass\n        \n        epoch_t = time.perf_counter() - epoch_start\n        mean_step_t = round(epoch_t / n_steps_per_epoch * 1000, 1)\n        n_imgs_per_s = int(1 / (mean_step_t / 1000) * bs)\n        print(f'epoch {epoch_num} took: {round(epoch_t, 2)} sec, mean step duration: {mean_step_t}ms, images/s: {n_imgs_per_s}')","metadata":{"id":"HoUK8FVQAcCy","execution":{"iopub.status.busy":"2023-02-25T05:40:43.45763Z","iopub.execute_input":"2023-02-25T05:40:43.458164Z","iopub.status.idle":"2023-02-25T05:40:43.473565Z","shell.execute_reply.started":"2023-02-25T05:40:43.458117Z","shell.execute_reply":"2023-02-25T05:40:43.472387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plots a batch of images\ndef show_batch(dataset, n_rows=16, n_cols=4):\n    inputs, targets = next(iter(dataset))\n    images = inputs['image'].numpy().squeeze()\n    fig, axes = plt.subplots(nrows=n_rows, ncols=n_cols, figsize=(n_cols*4, n_rows*7))\n    for r in range(n_rows):\n        for c in range(n_cols):\n            idx = r * n_cols + c\n            # Image\n            img = images[idx]\n            axes[r, c].imshow(img)\n            # Target\n            target = targets[idx]\n            axes[r, c].set_title(f'target: {target}', fontsize=16, pad=5)\n        \n    plt.show()","metadata":{"id":"UUtFz-mqAcCz","execution":{"iopub.status.busy":"2023-02-25T05:40:43.474807Z","iopub.execute_input":"2023-02-25T05:40:43.475611Z","iopub.status.idle":"2023-02-25T05:40:43.487276Z","shell.execute_reply.started":"2023-02-25T05:40:43.475557Z","shell.execute_reply":"2023-02-25T05:40:43.485967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Decodes the TFRecords\ndef decode_image(record_bytes):\n    features = tf.io.parse_single_example(record_bytes, {\n        'image': tf.io.FixedLenFeature([], tf.string),\n        'target': tf.io.FixedLenFeature([], tf.int64),\n        'patient_id': tf.io.FixedLenFeature([], tf.int64),\n    })\n    \n    # Decode PNG Image\n    image = tf.io.decode_jpeg(features['image'], channels=N_CHANNELS)\n    # Explicit reshape needed for TPU\n    image = tf.reshape(image, [IMG_HEIGHT, IMG_WIDTH, N_CHANNELS])\n\n    target = features['target']\n    \n    return { 'image': image }, target","metadata":{"id":"kYavfatTAcCz","execution":{"iopub.status.busy":"2023-02-25T05:40:43.492291Z","iopub.execute_input":"2023-02-25T05:40:43.49267Z","iopub.status.idle":"2023-02-25T05:40:43.695287Z","shell.execute_reply.started":"2023-02-25T05:40:43.492623Z","shell.execute_reply":"2023-02-25T05:40:43.694104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def augment_image(X, y):\n    image = X['image']\n    \n    # Random Brightness\n    image = tf.image.random_brightness(image, 0.10)\n    \n    # Random Contrast\n    image = tf.image.random_contrast(image, 0.90, 1.10)\n    \n    # Random JPEG Quality\n    image = tf.image.random_jpeg_quality(image, 75, 100)\n    \n    # Random crop image with maximum of 10%\n    ratio = tf.random.uniform([], 0.75, 1.00)\n    img_height_crop = tf.cast(ratio * IMG_HEIGHT, tf.int32)\n    img_width_crop = tf.cast(ratio * IMG_WIDTH, tf.int32)\n    # Random offset for crop\n    img_height_offset = tf_rand_int(0, IMG_HEIGHT - img_height_crop)\n    img_width_offset = 0\n    # Crop And Resize\n    image = tf.slice(image, [img_height_offset, img_width_offset, 0], [img_height_crop, img_width_crop, N_CHANNELS])\n    image = tf.image.resize(image, [IMG_HEIGHT, IMG_WIDTH], method=tf.image.ResizeMethod.BILINEAR)\n    # Clip pixel values in range [0,255] to prevent underflow/overflow\n    image = tf.clip_by_value(image, 0, 255)\n    image = tf.cast(image, tf.uint8)\n    \n    return { 'image': image }, y","metadata":{"id":"j_z_N2agAcCz","execution":{"iopub.status.busy":"2023-02-25T05:40:43.696746Z","iopub.execute_input":"2023-02-25T05:40:43.697005Z","iopub.status.idle":"2023-02-25T05:40:43.71551Z","shell.execute_reply.started":"2023-02-25T05:40:43.696976Z","shell.execute_reply":"2023-02-25T05:40:43.713775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Undersample majority class (0/negative) by randomly dropping them\ndef undersample_majority(X, y):\n    # Filter 2/3 of negative samples to upsample positive samples by a factor 3\n    return y == 1 or tf.random.uniform([]) > 0.66","metadata":{"id":"IE1wC6rWAcC0","execution":{"iopub.status.busy":"2023-02-25T05:40:43.717859Z","iopub.execute_input":"2023-02-25T05:40:43.71834Z","iopub.status.idle":"2023-02-25T05:40:43.726898Z","shell.execute_reply.started":"2023-02-25T05:40:43.718289Z","shell.execute_reply":"2023-02-25T05:40:43.725808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TFRecord file paths\nTFRECORDS_FILE_PATHS = sorted(tf.io.gfile.glob(f'{GCS_DS_PATH}/*.tfrecords'))\nprint(f'Found {len(TFRECORDS_FILE_PATHS)} TFRecords')","metadata":{"id":"_PwBTln9AcC0","outputId":"f910675f-d1ac-44c3-a2bd-4f6ee2730ce6","execution":{"iopub.status.busy":"2023-02-25T05:40:43.728253Z","iopub.execute_input":"2023-02-25T05:40:43.728544Z","iopub.status.idle":"2023-02-25T05:40:44.132181Z","shell.execute_reply.started":"2023-02-25T05:40:43.728513Z","shell.execute_reply":"2023-02-25T05:40:44.130964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_dataset(tfrecords, bs=BATCH_SIZE, val=False, debug=True):\n    ignore_order = tf.data.Options()\n    ignore_order.experimental_deterministic = False\n    \n    # Initialize dataset with TFRecords\n    dataset = tf.data.TFRecordDataset(tfrecords, num_parallel_reads=AUTO, compression_type='GZIP')\n    \n    # Decode mapping\n    dataset = dataset.map(decode_image, num_parallel_calls=AUTO)\n\n    if not val:\n        dataset = dataset.filter(undersample_majority)\n        dataset = dataset.map(augment_image, num_parallel_calls=AUTO)\n        dataset = dataset.with_options(ignore_order)\n        if not debug:\n            dataset = dataset.shuffle(1024)\n        dataset = dataset.repeat()        \n\n    dataset = dataset.batch(bs, drop_remainder=not val)\n    dataset = dataset.prefetch(AUTO)\n    \n    return dataset","metadata":{"id":"SNuGCzhfAcC1","execution":{"iopub.status.busy":"2023-02-25T05:40:44.133945Z","iopub.execute_input":"2023-02-25T05:40:44.134332Z","iopub.status.idle":"2023-02-25T05:40:44.1435Z","shell.execute_reply.started":"2023-02-25T05:40:44.134286Z","shell.execute_reply":"2023-02-25T05:40:44.142243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# pF1 Metric\n\ninspiration: [RSNA-BCD: EfficientNet [TF][TPU-1VM][Train]](https://www.kaggle.com/code/awsaf49/rsna-bcd-efficientnet-tf-tpu-1vm-train#Metric)\n\nThe source implementation is however buggy, it is a moving average which does not reset each epoch. The implementation below does reset each epoch.","metadata":{"id":"aXcxGf2yAcC1"}},{"cell_type":"code","source":"# Tensorflow custom metric is just a conventional class object\nclass pF1(tf.keras.metrics.Metric):\n    # Initialize properties\n    def __init__(self, name='pF1', **kwargs):\n        super(pF1, self).__init__(name=name, **kwargs)\n        self.tc = self.add_weight(name='tc', initializer='zeros')\n        self.tp = self.add_weight(name='tp', initializer='zeros')\n        self.fp = self.add_weight(name='fp', initializer='zeros')\n\n    # Update state called on each batch with true and predicted labels\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        self.tc.assign_add(tf.cast(tf.reduce_sum(y_true), tf.float32))\n        self.tp.assign_add(tf.cast(tf.reduce_sum((y_pred[y_true == 1])), tf.float32))\n        self.fp.assign_add(tf.cast(tf.reduce_sum((y_pred[y_true == 0])), tf.float32))\n\n    # Result function is called to obtain result which is printed in progress bar\n    def result(self):\n        if self.tc == 0 or (self.tp + self.fp) == 0:\n            return 0.0\n        else:\n            precision = self.tp / (self.tp + self.fp)\n            recall = self.tp / (self.tc)\n            return 2 * (precision * recall) / (precision + recall)\n\n    # Reset state is called after each epoch to start fresh each epoch\n    def reset_state(self):\n        self.tc.assign(0)\n        self.tp.assign(0)\n        self.fp.assign(0)","metadata":{"id":"eMfvn3UnAcC1","execution":{"iopub.status.busy":"2023-02-25T05:40:44.145291Z","iopub.execute_input":"2023-02-25T05:40:44.145727Z","iopub.status.idle":"2023-02-25T05:40:44.161099Z","shell.execute_reply.started":"2023-02-25T05:40:44.145683Z","shell.execute_reply":"2023-02-25T05:40:44.159921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"id":"mJvW_KcpAcC2"}},{"cell_type":"code","source":"def normalize(image):\n    # Repeat channels to create 3 channel images required by pretrained ConvNextV2 models\n    image = tf.repeat(image, repeats=3, axis=3)\n    # Cast to float 32\n    image = tf.cast(image, tf.float32)\n    # Normalize with respect to ImageNet mean/std\n    image = tf.keras.applications.imagenet_utils.preprocess_input(image, mode='torch')\n\n    return image","metadata":{"id":"44liKMeDAcC2","execution":{"iopub.status.busy":"2023-02-25T05:40:44.163271Z","iopub.execute_input":"2023-02-25T05:40:44.163789Z","iopub.status.idle":"2023-02-25T05:40:44.176475Z","shell.execute_reply.started":"2023-02-25T05:40:44.163745Z","shell.execute_reply":"2023-02-25T05:40:44.17529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model():\n    # Verify Mixed Policy Settings\n    print(f'Compute dtype: {tf.keras.mixed_precision.global_policy().compute_dtype}')\n    print(f'Variable dtype: {tf.keras.mixed_precision.global_policy().variable_dtype}')\n    \n    with STRATEGY.scope():\n        # Set seed for deterministic weights initialization\n        seed_everything()\n        \n        # Inputs, note the names are equal to the dictionary keys in the dataset\n        image = tf.keras.layers.Input(INPUT_SHAPE, name='image', dtype=tf.uint8)\n        \n        # Normalize Input\n        image_norm = normalize(image)\n        \n        # Normalize Input\n        image_norm = normalize(image)\n\n        # CNN Prediction in range [0,1]\n        x = convnext.ConvNeXtV2Tiny(\n            input_shape=(IMG_HEIGHT, IMG_WIDTH, 3),\n            pretrained='imagenet21k-ft1k',\n            num_classes=0,\n        )(image_norm)\n        \n        # Average Pooling BxHxWxC -> BxC\n        x = tf.keras.layers.GlobalAveragePooling2D()(x)\n        # Dropout to prevent Overfitting\n        x = tf.keras.layers.Dropout(0.30)(x)\n        # Output value between [0, 1] using Sigmoid function\n        outputs = tf.keras.layers.Dense(1, activation='sigmoid')(x)\n\n        # We will use the famous AdamW optimizer for fast learning with weight decay\n        optimizer = tfa.optimizers.AdamW(learning_rate=LR_MAX, weight_decay=LR_MAX*WD_RATIO, epsilon=1e-6)\n\n        # Loss\n        loss = tf.keras.losses.BinaryCrossentropy(from_logits=False)\n        \n        # Metrics\n        metrics = [\n            pF1(),\n            tfa.metrics.F1Score(num_classes=1, threshold=0.50),\n            tf.keras.metrics.Precision(),\n            tf.keras.metrics.Recall(),\n            tf.keras.metrics.AUC(),\n            tf.keras.metrics.BinaryAccuracy(),\n        ]\n\n        model = tf.keras.models.Model(inputs=image, outputs=outputs)\n        \n        model.compile(optimizer=optimizer, loss=loss, metrics=metrics)\n\n        return model","metadata":{"id":"u03bGS5BAcC2","execution":{"iopub.status.busy":"2023-02-25T05:40:44.177736Z","iopub.execute_input":"2023-02-25T05:40:44.177974Z","iopub.status.idle":"2023-02-25T05:40:44.190652Z","shell.execute_reply.started":"2023-02-25T05:40:44.177949Z","shell.execute_reply":"2023-02-25T05:40:44.189673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Pretrained File Path: '/kaggle/input/sartorius-training-dataset/model.h5'\ntf.keras.backend.clear_session()\n# enable XLA optmizations\ntf.config.optimizer.set_jit(True)\n\nmodel = get_model()\n","metadata":{"id":"BB45_9-3AcC3","outputId":"0ad5fcbc-7607-4e99-ef15-f6aa621a6962","execution":{"iopub.status.busy":"2023-02-25T05:40:44.192359Z","iopub.execute_input":"2023-02-25T05:40:44.193085Z","iopub.status.idle":"2023-02-25T05:40:56.357253Z","shell.execute_reply.started":"2023-02-25T05:40:44.19304Z","shell.execute_reply":"2023-02-25T05:40:56.356001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot model summary\nmodel.summary()","metadata":{"id":"EAGPx2ysAcC3","outputId":"4f46e94b-7225-4bf6-f341-44e8976efacb","execution":{"iopub.status.busy":"2023-02-25T05:40:56.358692Z","iopub.execute_input":"2023-02-25T05:40:56.35895Z","iopub.status.idle":"2023-02-25T05:40:56.380955Z","shell.execute_reply.started":"2023-02-25T05:40:56.358923Z","shell.execute_reply":"2023-02-25T05:40:56.380007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Model architecture\ntf.keras.utils.plot_model(model, show_shapes=True, show_dtype=True, show_layer_names=True, expand_nested=False)","metadata":{"id":"rj8pvI6sAcC4","outputId":"2040c82b-1246-4fc6-baf5-8a7164dd4b92","execution":{"iopub.status.busy":"2023-02-25T05:40:56.382256Z","iopub.execute_input":"2023-02-25T05:40:56.382519Z","iopub.status.idle":"2023-02-25T05:40:57.829287Z","shell.execute_reply.started":"2023-02-25T05:40:56.382489Z","shell.execute_reply":"2023-02-25T05:40:57.828405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Validation metric on initialized model\n# _ = model.evaluate(\n#         get_dataset(TFRECORDS_VAL, val=True),\n#         verbose=VERBOSE,\n#         steps=VAL_STEPS_PER_EPOCH,\n#     )","metadata":{"id":"6fdsjMplAcC5","execution":{"iopub.status.busy":"2023-02-25T05:40:57.831166Z","iopub.execute_input":"2023-02-25T05:40:57.831567Z","iopub.status.idle":"2023-02-25T05:40:57.838462Z","shell.execute_reply.started":"2023-02-25T05:40:57.831525Z","shell.execute_reply":"2023-02-25T05:40:57.837387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Train Output Baseline\n# train_preds = model.predict(\n#         get_dataset(TFRECORDS_TRAIN, val=True),\n#         verbose=VERBOSE,\n#         steps=128,\n#     ).squeeze()","metadata":{"id":"kDmxN33RAcC5","execution":{"iopub.status.busy":"2023-02-25T05:40:57.840355Z","iopub.execute_input":"2023-02-25T05:40:57.840919Z","iopub.status.idle":"2023-02-25T05:40:57.853305Z","shell.execute_reply.started":"2023-02-25T05:40:57.840869Z","shell.execute_reply":"2023-02-25T05:40:57.85218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Initialized model train predictions: should not be saturated (all 0/1)\n# display(pd.Series(train_preds).describe().to_frame('Value').round(2))","metadata":{"id":"5Ae92qLhAcC6","execution":{"iopub.status.busy":"2023-02-25T05:40:57.855309Z","iopub.execute_input":"2023-02-25T05:40:57.855941Z","iopub.status.idle":"2023-02-25T05:40:57.865055Z","shell.execute_reply.started":"2023-02-25T05:40:57.855893Z","shell.execute_reply":"2023-02-25T05:40:57.864031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Predictions on initialized model, shoudl be random and not saturated (all 0 or  all 1)\n# plt.figure(figsize=(15,8))\n# plt.title(f'Training Predictions Initialized Model')\n# pd.Series(train_preds).plot(kind='hist', bins=32)\n# plt.xticks(np.arange(0, 1.1, 0.1))\n# plt.xlim(0, 1)\n# plt.grid()\n# plt.show()","metadata":{"id":"oN_ztY8IAcC6","execution":{"iopub.status.busy":"2023-02-25T05:40:57.866553Z","iopub.execute_input":"2023-02-25T05:40:57.867083Z","iopub.status.idle":"2023-02-25T05:40:57.876535Z","shell.execute_reply.started":"2023-02-25T05:40:57.867039Z","shell.execute_reply":"2023-02-25T05:40:57.87564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Learning Rate Scheduler","metadata":{"id":"RqPV_POCAcC6"}},{"cell_type":"code","source":"# Learning rate scheduler with logaritmic warmup and cosine decay\ndef lrfn(current_step, num_warmup_steps, lr_max, num_cycles=0.50, num_training_steps=N_EPOCHS):\n    \n    if current_step < num_warmup_steps:\n        return lr_max * 0.10 ** (num_warmup_steps - current_step)\n    else:\n        progress = float(current_step - num_warmup_steps) / float(max(1, num_training_steps - num_warmup_steps))\n\n        return max(0.0, 0.5 * (1.0 + math.cos(math.pi * float(num_cycles) * 2.0 * progress))) * lr_max","metadata":{"id":"YQBsKIsUAcC6","execution":{"iopub.status.busy":"2023-02-25T05:40:57.877667Z","iopub.execute_input":"2023-02-25T05:40:57.878494Z","iopub.status.idle":"2023-02-25T05:40:57.890613Z","shell.execute_reply.started":"2023-02-25T05:40:57.878456Z","shell.execute_reply":"2023-02-25T05:40:57.889356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot the learning rate scheduler\ndef plot_lr_schedule(lr_schedule, epochs):\n    fig = plt.figure(figsize=(20, 10))\n    plt.plot([None] + lr_schedule + [None])\n    # X Labels\n    x = np.arange(1, epochs + 1)\n    x_axis_labels = [i if epochs <= 40 or i % 5 == 0 or i == 1 else None for i in range(1, epochs + 1)]\n    plt.xlim([1, epochs])\n    plt.xticks(x, x_axis_labels) # set tick step to 1 and let x axis start at 1\n    \n    # Increase y-limit for better readability\n    plt.ylim([0, max(lr_schedule) * 1.1])\n    \n    # Title\n    schedule_info = f'start: {lr_schedule[0]:.1E}, max: {max(lr_schedule):.1E}, final: {lr_schedule[-1]:.1E}'\n    plt.title(f'Step Learning Rate Schedule, {schedule_info}', size=18, pad=12)\n    \n    # Plot Learning Rates\n    for x, val in enumerate(lr_schedule):\n        if epochs <= 40 or x % 5 == 0 or x is epochs - 1:\n            if x < len(lr_schedule) - 1:\n                if lr_schedule[x - 1] < val:\n                    ha = 'right'\n                else:\n                    ha = 'left'\n            elif x == 0:\n                ha = 'right'\n            else:\n                ha = 'left'\n            plt.plot(x + 1, val, 'o', color='black');\n            offset_y = (max(lr_schedule) - min(lr_schedule)) * 0.02\n            plt.annotate(f'{val:.1E}', xy=(x + 1, val + offset_y), size=12, ha=ha)\n    \n    plt.xlabel('Epoch', size=16, labelpad=5)\n    plt.ylabel('Learning Rate', size=16, labelpad=5)\n    plt.grid()\n    plt.show()\n\n# Learning rate for encoder\nLR_SCHEDULE = [lrfn(step, num_warmup_steps=N_WARMUP_EPOCHS, lr_max=LR_MAX, num_cycles=0.50) for step in range(N_EPOCHS)]\nplot_lr_schedule(LR_SCHEDULE, epochs=N_EPOCHS)","metadata":{"id":"eIkRmEQEAcC7","outputId":"c5746db3-b36f-47d3-e9e9-5965247f71e9","execution":{"iopub.status.busy":"2023-02-25T05:40:57.892299Z","iopub.execute_input":"2023-02-25T05:40:57.892828Z","iopub.status.idle":"2023-02-25T05:40:58.262667Z","shell.execute_reply.started":"2023-02-25T05:40:57.892792Z","shell.execute_reply":"2023-02-25T05:40:58.261753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Learning Rate Callback\nlr_callback = tf.keras.callbacks.LearningRateScheduler(lambda step: LR_SCHEDULE[step], verbose=0)","metadata":{"id":"atszUqpyAcC7","execution":{"iopub.status.busy":"2023-02-25T05:40:58.264138Z","iopub.execute_input":"2023-02-25T05:40:58.264516Z","iopub.status.idle":"2023-02-25T05:40:58.270901Z","shell.execute_reply.started":"2023-02-25T05:40:58.264473Z","shell.execute_reply":"2023-02-25T05:40:58.269859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Weight Decay Callback","metadata":{"id":"yYO62wZ6AcC8"}},{"cell_type":"code","source":"# Tensorflow Learning Rate Scheduler does not update weight decay, need to do it manually in a custom callback\nclass WeightDecayCallback(tf.keras.callbacks.Callback):\n    def __init__(self, wd_ratio=WD_RATIO):\n        self.step_counter = 0\n        self.wd_ratio = wd_ratio\n    \n    def on_epoch_begin(self, epoch, logs=None):\n        model.optimizer.weight_decay = model.optimizer.learning_rate * self.wd_ratio\n        print(f'learning rate: {model.optimizer.learning_rate.numpy():.2e}, weight decay: {model.optimizer.weight_decay.numpy():.2e}')","metadata":{"id":"KThGgbFbAcC8","execution":{"iopub.status.busy":"2023-02-25T05:40:58.272317Z","iopub.execute_input":"2023-02-25T05:40:58.272634Z","iopub.status.idle":"2023-02-25T05:40:58.284446Z","shell.execute_reply.started":"2023-02-25T05:40:58.2726Z","shell.execute_reply":"2023-02-25T05:40:58.28321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{"id":"-1J_7GXoAcC8"}},{"cell_type":"code","source":"# source: https://www.kaggle.com/code/sohier/probabilistic-f-score\n# Competition Leaderboard Metric\ndef pfbeta(labels, predictions, beta=1):\n    y_true_count = 0\n    ctp = 0\n    cfp = 0\n\n    for idx in range(len(labels)):\n        prediction = min(max(predictions[idx], 0), 1)\n        if (labels[idx]):\n            y_true_count += 1\n            ctp += prediction\n        else:\n            cfp += prediction\n\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / y_true_count\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall)\n        return result\n    else:\n        return 0","metadata":{"id":"TtmKLfY4AcC8","execution":{"iopub.status.busy":"2023-02-25T05:40:58.286091Z","iopub.execute_input":"2023-02-25T05:40:58.287087Z","iopub.status.idle":"2023-02-25T05:40:58.300352Z","shell.execute_reply.started":"2023-02-25T05:40:58.287037Z","shell.execute_reply":"2023-02-25T05:40:58.299337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"id":"o7DfV3DIAcC9","outputId":"9c8a4130-79d9-4395-97ec-8f91d1da3108","execution":{"iopub.status.busy":"2023-02-25T05:40:58.301948Z","iopub.execute_input":"2023-02-25T05:40:58.302354Z","iopub.status.idle":"2023-02-25T05:40:58.606083Z","shell.execute_reply.started":"2023-02-25T05:40:58.302285Z","shell.execute_reply":"2023-02-25T05:40:58.605361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold\nimport json\n\n\nos.makedirs('models',exist_ok=True)\nkf = KFold(n_splits=5,shuffle=True,random_state=SEED)\nconfig_dict = {}\n\n\n# Train Test Split\nfor fold, (train_index, val_index) in tqdm(enumerate(kf.split(TFRECORDS_FILE_PATHS))):\n    TFRECORDS_TRAIN, TFRECORDS_VAL = np.array(TFRECORDS_FILE_PATHS)[train_index],np.array(TFRECORDS_FILE_PATHS)[val_index],\n    print(f'# TFRECORDS_TRAIN: {len(TFRECORDS_TRAIN)}, # TFRECORDS_VAL: {len(TFRECORDS_VAL)}')\n    model = get_model()\n    \n    # Get Train/Validation datasets\n    train_dataset = get_dataset(TFRECORDS_TRAIN, val=False, debug=False)\n    val_dataset = get_dataset(TFRECORDS_VAL, val=True, debug=False)\n\n    TRAIN_STEPS_PER_EPOCH = len(TFRECORDS_TRAIN) * N_SAMPLES_TFRECORDS // BATCH_SIZE\n    VAL_STEPS_PER_EPOCH = len(TFRECORDS_VAL) * N_SAMPLES_TFRECORDS // BATCH_SIZE\n    print(f'TRAIN_STEPS_PER_EPOCH: {TRAIN_STEPS_PER_EPOCH}, VAL_STEPS_PER_EPOCH: {VAL_STEPS_PER_EPOCH}')\n    gc.collect()\n    # Train model on TPU!\n    history = model.fit(\n            train_dataset,\n            steps_per_epoch = TRAIN_STEPS_PER_EPOCH,\n            validation_data = val_dataset,\n            epochs = N_EPOCHS,\n            verbose = VERBOSE,\n            callbacks = [\n                lr_callback,\n                WeightDecayCallback(),\n            ],\n            class_weight = {\n                0: 1.0,\n                1: 5.0,\n            },\n        )\n    \n    # Save model weights for inference\n    model.save_weights(f'models/model-f{fold}.h5')\n    \n    \n    \n    \n    # Get true labels and predictions for validation set\n    y_true_val = []\n    y_pred_val = []\n    for X_batch, y_batch in tqdm(get_dataset(TFRECORDS_VAL, val=True), total=VAL_STEPS_PER_EPOCH):\n        y_true_val += y_batch.numpy().tolist()\n        y_pred_val += model.predict_on_batch(X_batch).squeeze().tolist()\n    del model\n    gc.collect()\n    # Plot pF1 by threshold plot to find best threshold\n    pf1_by_threshold = []\n    thresholds = np.arange(0, 1.01, 0.01)\n    for t in tqdm(thresholds):\n        # Compute pF1 for each threshold\n        pf1_by_threshold.append(\n            pfbeta(y_true_val, y_pred_val > t)\n        )\n        \n    # Best threshold and pF1 score\n    arg_max = np.argmax(pf1_by_threshold)\n    val_max = np.max(pf1_by_threshold)\n    threshold_best = thresholds[arg_max]\n    print(f'Best Threshold {threshold_best:.2f}, pF1 Score: {val_max:.2f}')\n    config_dict[f'model-f{fold}'] = {\n        'threshold':threshold_best,\n        'pf1_by_threshold':val_max,\n        'TFRECORDS_TRAIN': len(TFRECORDS_TRAIN),\n        'TFRECORDS_VAL': len(TFRECORDS_VAL),\n        'TRAIN_STEPS_PER_EPOCH': TRAIN_STEPS_PER_EPOCH,\n        'VAL_STEPS_PER_EPOCH': VAL_STEPS_PER_EPOCH,\n    }\n    \n    config_dict[f'model-f{fold}'] = dict(**config_dict[f'model-f{fold}'],**history.history)\n    del history\n    gc.collect()\n    # cast types to save as json\n#     for key in config_dict[f'model-f{fold}'].keys():\n#         if type(config_dict[f'model-f{fold}'][key]) == np.ndarray:\n#                  config_dict[f'model-f{fold}'][key] = config_dict[f'model-f{fold}'][key].tolist()\n#         try:\n#             config_dict[f'model-f{fold}'][key] = float(config_dict[f'model-f{fold}'][key])\n#             config_dict[f'model-f{fold}'][key] = list(map(float,config_dict[f'model-f{fold}'][key]))\n#         except: pass\n    gc.collect()\n    if IS_INTERACTIVE: break\n        \n\n# with open('models/models-config.json', 'w') as f:\n#     json.dump(config_dict,f)\n\nprint(config_dict)","metadata":{"id":"5OUkapeLAcC9","outputId":"e4ae0877-5afb-410b-9644-81c091b55727","execution":{"iopub.status.busy":"2023-02-25T05:40:58.607343Z","iopub.execute_input":"2023-02-25T05:40:58.607584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training History","metadata":{"id":"k_lE0HmhAcC9"}},{"cell_type":"code","source":"def plot_history_metric(metric, f_best=np.argmax, ylim=None, yscale=None, yticks=None):\n\n    fig,axs = plt.subplots(len(config_dict.keys()),1,figsize=(20, 10))\n    axs = axs.flatten()\n\n    for fold in range(len(config_dict.keys())):\n        hist_dict = config_dict[f'model-f{fold}']\n        values = hist_dict[metric]\n        N_EPOCHS = len(values)\n        val = 'val' in ''.join(hist_dict.keys())\n        # Epoch Ticks\n        if N_EPOCHS <= 20:\n            x = np.arange(1, N_EPOCHS + 1)\n        else:\n            x = [1, 5] + [10 + 5 * idx for idx in range((N_EPOCHS - 10) // 5 + 1)]\n\n        x_ticks = np.arange(1, N_EPOCHS+1)\n\n        # Validation\n        if val:\n            val_values = hist_dict[f'val_{metric}']\n            val_argmin = f_best(val_values)\n            plt.plot(x_ticks, val_values, label=f'val')\n\n        # summarize history for accuracy\n        axs[fold].plot(x_ticks, values, label=f'train')\n        argmin = f_best(values)\n        axs[fold].scatter(argmin + 1, values[argmin], color='red', s=75, marker='o', label=f'train_best')\n        if val:\n            axs[fold].scatter(val_argmin + 1, val_values[val_argmin], color='purple', s=75, marker='o', label=f'val_best')\n\n        axs[fold].set_title(f'fold{fold}-Model {metric}', fontsize=24, pad=2)\n        axs[fold].set_ylabel(metric, fontsize=20, labelpad=5)\n\n        if ylim:\n            axs[fold].set_ylim(ylim)\n\n        if yscale is not None:\n            axs[fold].set_yscale(yscale)\n\n        if yticks is not None:\n            axs[fold].set_yticklabels(yticks, fontsize=16)\n\n        axs[fold].set_xlabel('epoch', fontsize=20, labelpad=5)        \n        axs[fold].tick_params(axis='x', labelsize=8)\n        axs[fold].set_xticklabels(x, fontsize=16) # set tick step to 1 and let x axis start at 1\n        axs[fold].grid()\n        # axs[fold].set_yticks(fontsize=16)\n\n    plt.legend(prop={'size': 10})\n    plt.show()","metadata":{"id":"gbEr4g8lAcC9","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_history_metric('loss', f_best=np.argmin)","metadata":{"id":"CeNLRp3IAcC-","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_history_metric('pF1', ylim=[0,1], yticks=np.arange(0.0, 1.1, 0.1))","metadata":{"id":"zO9KpYQCAcC-","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_history_metric('f1_score', ylim=[0,1], yticks=np.arange(0.0, 1.1, 0.1))","metadata":{"id":"s0TAWQJYAcC-","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_history_metric('precision', ylim=[0,1], yticks=np.arange(0.0, 1.1, 0.1))","metadata":{"id":"v4pe5f6AAcC_","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_history_metric('recall', ylim=[0,1], yticks=np.arange(0.0, 1.1, 0.1))","metadata":{"id":"vFWEJEZ3AcC_","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_history_metric('auc', ylim=[0,1], yticks=np.arange(0.0, 1.1, 0.1))","metadata":{"id":"0oVeWMJMAcC_","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_history_metric('binary_accuracy', ylim=[0,1], yticks=np.arange(0.0, 1.1, 0.1))","metadata":{"id":"4QuA4Ny_AcC_","trusted":true},"execution_count":null,"outputs":[]}]}