{"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":"## About notebook\n\nThis notebook is pipeline to use transformer-based models (at least from [keras_cv_attention_models](https://github.com/leondgarse/keras_cv_attention_models) library)\n\nSome code and techniques are from [RSNA-BCD: EfficientNet [TF][TPU-1VM][Train]](https://www.kaggle.com/code/awsaf49/rsna-bcd-efficientnet-tf-tpu-1vm-train) and [DeBERTa LayerwiseLR LastLayerReinit TensorFlow](https://www.kaggle.com/code/electro/deberta-layerwiselr-lastlayerreinit-tensorflow) notebooks.\n\n* **Data**: *[ROI TFRecords 512x1024](https://www.kaggle.com/datasets/olegbaryshnikov/rsna-roi-tfrecords-512x1024)* (images and cancer column)\n* **Image augmentations**: Random brightness, contrast, crop and coarse dropout\n* **Model**: *[CoaTLiteSmall](https://github.com/leondgarse/keras_cv_attention_models/tree/main/keras_cv_attention_models/coat)* with Multi-Sample Dropout\n* **Losses**: *Focal* and *BCE*\n* **Metrics**: *pF1*, *thresholded pF1* (**maximized**) and *AUC*\n* **Optimizers and schedulers**: *AdamW* and *ExponentialDecay* with Layer-Wise Decay\n* **Logger**: *Wandb*\n* **Accelerator**: *TPU*\n\n\n<h4>\n    <span style=\"color:orange\">\n        Inference notebook: \n        <a href=\"https://www.kaggle.com/olegbaryshnikov/rsna-coat-tf-inference\">[RSNA] CoaT [TF][Inference]</a>\n    </span>\n</h4>","metadata":{}},{"cell_type":"markdown","source":"## Version History\n* **v3: CoaTTiny baseline**\n* **v5: CoaTSmall baseline, model quantization and wandb turn off options**","metadata":{}},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"!pip install -q --no-index --no-deps --find-links=/kaggle/input/keras-cv-attention-models keras-cv-attention-models","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:21:00.768904Z","iopub.execute_input":"2023-01-17T19:21:00.769422Z","iopub.status.idle":"2023-01-17T19:21:04.051133Z","shell.execute_reply.started":"2023-01-17T19:21:00.76932Z","shell.execute_reply":"2023-01-17T19:21:04.050112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -qU --no-deps tensorflow-addons==0.13.0","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:21:04.053515Z","iopub.execute_input":"2023-01-17T19:21:04.0539Z","iopub.status.idle":"2023-01-17T19:21:21.034266Z","shell.execute_reply.started":"2023-01-17T19:21:04.053863Z","shell.execute_reply":"2023-01-17T19:21:21.033275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -qU wandb","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:21:21.036064Z","iopub.execute_input":"2023-01-17T19:21:21.036417Z","iopub.status.idle":"2023-01-17T19:21:38.939959Z","shell.execute_reply.started":"2023-01-17T19:21:21.036382Z","shell.execute_reply":"2023-01-17T19:21:38.939001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install -qU tensorflow-model-optimization","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:21:38.944789Z","iopub.execute_input":"2023-01-17T19:21:38.945151Z","iopub.status.idle":"2023-01-17T19:21:51.91747Z","shell.execute_reply.started":"2023-01-17T19:21:38.945114Z","shell.execute_reply":"2023-01-17T19:21:51.915935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nimport math\nimport glob\nimport string\nimport random\nimport shutil\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nfrom IPython.display import display\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom tensorflow.keras import backend as K\n\nfrom tensorflow.keras.utils import plot_model\nimport tensorflow_model_optimization as tfmot\n\nfrom tqdm.notebook import tqdm\nimport gc\n\nfrom sklearn.model_selection import KFold\nfrom sklearn.utils.class_weight import compute_class_weight\n\nfrom kaggle_datasets import KaggleDatasets\nimport keras_cv_attention_models as kecam\nimport wandb","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-17T19:21:51.919679Z","iopub.execute_input":"2023-01-17T19:21:51.920165Z","iopub.status.idle":"2023-01-17T19:21:59.807673Z","shell.execute_reply.started":"2023-01-17T19:21:51.920104Z","shell.execute_reply":"2023-01-17T19:21:59.806217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Tensorflow version \" + tf.__version__)\nprint('tfa:', tfa.__version__)","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:21:59.809629Z","iopub.execute_input":"2023-01-17T19:21:59.810025Z","iopub.status.idle":"2023-01-17T19:21:59.822397Z","shell.execute_reply.started":"2023-01-17T19:21:59.809986Z","shell.execute_reply":"2023-01-17T19:21:59.820453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"Config={\n    \"seed\":1111,\n\n    \"dim_x\": 512,\n    \"dim_y\": 1024,\n    \"loss\":\"focal\",\n    \"model_type\": \"coat\",\n    \"model_name\":\"CoaTLiteSmall\",\n    \"model_weights\": \"imagenet\",\n    \"include_preprocessing\": True,\n    \n    \"exp_name\":\"rsna_CoaTLiteSmall_512x1024\",\n    \n    #Training config\n    \"epochs\": 7,\n    \"learning_rate\": 1e-4,\n    \"min_lr_mul\": 1e-2,\n    \"use_scheduler\": True,\n    \"use_layerwise_lr\": True,\n    \"last_layer_speedup\": False,\n    \"auguments\": True,\n    \n    #\"hidden_layers_num\": 8,\n    \n    \"optimizer\": \"AdamW\",\n    \"weight_decay\": 1e-6,\n    \n    \"quantize_model\": True,\n    \n    \"batch_size\": 4, \n    \"n_fold\": 4,\n    \"early_stopping_patience\" : 5,\n    \n    \"use_wandb\": True,\n}","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:21:59.824173Z","iopub.execute_input":"2023-01-17T19:21:59.824579Z","iopub.status.idle":"2023-01-17T19:21:59.846251Z","shell.execute_reply.started":"2023-01-17T19:21:59.824543Z","shell.execute_reply":"2023-01-17T19:21:59.84506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Detect TPU, return appropriate distribution strategy\nConfig[\"tpu\"]=None\ntry:\n    Config[\"tpu\"] = tf.distribute.cluster_resolver.TPUClusterResolver() \n    print('Running on TPU ', Config[\"tpu\"].master())\nexcept ValueError:\n    Config[\"tpu\"] = None\n    \nif Config[\"tpu\"]:\n    tf.config.experimental_connect_to_cluster(Config[\"tpu\"])\n    tf.tpu.experimental.initialize_tpu_system(Config[\"tpu\"])\n    Config[\"strategy\"] = tf.distribute.experimental.TPUStrategy(Config[\"tpu\"])\nelse:\n    Config[\"strategy\"] = tf.distribute.get_strategy() \n\nConfig[\"REPLICAS\"] = Config[\"strategy\"].num_replicas_in_sync\nConfig[\"AUTO\"] = tf.data.experimental.AUTOTUNE\nprint(\"REPLICAS: \", Config[\"REPLICAS\"])","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:21:59.848041Z","iopub.execute_input":"2023-01-17T19:21:59.849191Z","iopub.status.idle":"2023-01-17T19:21:59.866894Z","shell.execute_reply.started":"2023-01-17T19:21:59.849125Z","shell.execute_reply":"2023-01-17T19:21:59.865578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=42):\n    np.random.seed(seed)\n    random.seed(seed)\n    tf.random.set_seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(Config[\"seed\"])","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:21:59.868406Z","iopub.execute_input":"2023-01-17T19:21:59.868798Z","iopub.status.idle":"2023-01-17T19:21:59.879529Z","shell.execute_reply.started":"2023-01-17T19:21:59.868764Z","shell.execute_reply":"2023-01-17T19:21:59.878512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if(Config[\"use_wandb\"]):\n    try:\n        from kaggle_secrets import UserSecretsClient\n        user_secrets = UserSecretsClient()\n        api_key = user_secrets.get_secret(\"WANDB\")\n\n        wandb.login(key=api_key)\n        anonymous = None\n    except:\n        anonymous = \"must\"\n        print('To use your W&B account,\\nGo to Add-ons -> Secrets and provide your W&B access token. Use the Label name as WANDB. \\nGet your W&B access token from here: https://wandb.ai/authorize')","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:21:59.885828Z","iopub.execute_input":"2023-01-17T19:21:59.886209Z","iopub.status.idle":"2023-01-17T19:22:05.007848Z","shell.execute_reply.started":"2023-01-17T19:21:59.886177Z","shell.execute_reply":"2023-01-17T19:22:05.006566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def wandb_init(fold):\n    ignore_keys = [\"tpu\", \"strategy\", \"AUTO\"]\n    \n    config = {key:Config[key] for key in Config if key not in ignore_keys}\n    config[\"fold\"] = int(fold)\n    \n    run    = wandb.init(\n        project=\"rsna-bcd-public\",\n        name=f\"fold-{fold}|dim-{Config['dim_x']}x{Config['dim_y']}|model-{Config['model_name']}\",\n        config=config,\n        anonymous=anonymous,\n        group=Config[\"exp_name\"]\n    )\n    return run","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:05.009657Z","iopub.execute_input":"2023-01-17T19:22:05.010048Z","iopub.status.idle":"2023-01-17T19:22:05.017531Z","shell.execute_reply.started":"2023-01-17T19:22:05.010012Z","shell.execute_reply":"2023-01-17T19:22:05.016401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Datasets","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ntrain_df","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:05.019084Z","iopub.execute_input":"2023-01-17T19:22:05.01946Z","iopub.status.idle":"2023-01-17T19:22:05.194915Z","shell.execute_reply.started":"2023-01-17T19:22:05.019403Z","shell.execute_reply":"2023-01-17T19:22:05.193836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_weight = compute_class_weight(class_weight='balanced',\n                                        classes=train_df.cancer.unique(),\n                                        y=train_df.cancer.values)\nclass_weight = dict(zip(train_df.cancer.unique(), class_weight))\nclass_weight","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:05.196544Z","iopub.execute_input":"2023-01-17T19:22:05.197042Z","iopub.status.idle":"2023-01-17T19:22:05.222362Z","shell.execute_reply.started":"2023-01-17T19:22:05.19699Z","shell.execute_reply":"2023-01-17T19:22:05.221195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"GCS_PATH = KaggleDatasets().get_gcs_path(f'rsna-roi-tfrecords-{Config[\"dim_x\"]}x{Config[\"dim_y\"]}')\nprint(f'GCS_PATH: {GCS_PATH}')\n\ntfrecords_file_paths = tf.io.gfile.glob(f'{GCS_PATH}/train_{Config[\"dim_x\"]}x{Config[\"dim_y\"]}/*.tfrec')\ntfrecords_file_paths[0]","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:05.224032Z","iopub.execute_input":"2023-01-17T19:22:05.224403Z","iopub.status.idle":"2023-01-17T19:22:07.820229Z","shell.execute_reply.started":"2023-01-17T19:22:05.224369Z","shell.execute_reply":"2023-01-17T19:22:07.819298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def random_float(shape=[], minval=0.0, maxval=1.0):\n    rnd = tf.random.uniform(\n        shape=shape, minval=minval, maxval=maxval, dtype=tf.float32)\n    return rnd\n\ndef dropout(image,DIM=[Config[\"dim_y\"],Config[\"dim_x\"]], CT = 12, SZ = 0.08, mode='fill', fill_with = 0):\n    # input image - is one image of size [dim,dim,3] not a batch of [b,dim,dim,3]\n    # output - image with CT squares of side size SZ*DIM removed\n\n    # DO DROPOUT WITH PROBABILITY DEFINED ABOVE\n    if (CT==0)|(SZ==0): \n        return image\n    CT = tf.random.uniform([],1,CT,tf.int32)\n    SZ = random_float(minval=0.7, maxval=1.0)*SZ\n    for k in range(CT):\n        # CHOOSE RANDOM LOCATION\n        x = tf.cast( tf.random.uniform([],0,DIM[1]),tf.int32)\n        y = tf.cast( tf.random.uniform([],0,DIM[0]),tf.int32)\n        # COMPUTE SQUARE \n        WIDTH = tf.cast( SZ*min(DIM),tf.int32)\n        ya = tf.math.maximum(0,y-WIDTH//2)\n        yb = tf.math.minimum(DIM[0],y+WIDTH//2)\n        xa = tf.math.maximum(0,x-WIDTH//2)\n        xb = tf.math.minimum(DIM[1],x+WIDTH//2)\n        # DROPOUT IMAGE\n        one = image[ya:yb,0:xa,:]\n        \n        two = tf.ones([yb-ya,xb-xa,3], dtype = image.dtype)\n        if(mode=='fill'):\n             two = two * fill_with\n        elif(mode=='mean'):\n            two = two * tf.reduce_mean(image[ya:yb,xa:xb,:])\n        two = tf.cast(two, image.dtype)\n            \n        three = image[ya:yb,xb:DIM[1],:]\n        middle = tf.concat([one,two,three],axis=1)\n        image = tf.concat([image[0:ya,:,:],middle,image[yb:DIM[0],:,:]],axis=0)\n        image = tf.reshape(image,[*DIM,3])\n\n    return image","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:07.821765Z","iopub.execute_input":"2023-01-17T19:22:07.822376Z","iopub.status.idle":"2023-01-17T19:22:07.837133Z","shell.execute_reply.started":"2023-01-17T19:22:07.822342Z","shell.execute_reply":"2023-01-17T19:22:07.835595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def apply_augments(img):\n    if (random_float() < 0.3):\n        img = tf.image.random_brightness(img, 0.2)\n        img = tf.image.random_contrast(img, 0.85, 1.15)\n    if (random_float() < 0.2):\n        crop_size = random_float(minval=0.75, maxval=0.95)\n        crop_size_x = tf.cast(Config[\"dim_x\"]*crop_size,tf.int32)\n        crop_size_y = tf.cast(Config[\"dim_y\"]*crop_size,tf.int32)\n        img = tf.image.random_crop(img, [crop_size_y,crop_size_x,3])\n        \n        img = tf.image.resize(img, [Config[\"dim_y\"],Config[\"dim_x\"]])\n        img = tf.cast(img,np.dtype('uint8'))\n    if (random_float() < 0.1):\n        rnd = random_float()\n        if(rnd<0.7):\n            img = dropout(img,DIM=[Config[\"dim_y\"],Config[\"dim_x\"]], CT = 12, SZ = 0.08, mode='mean')\n        elif(rnd<0.85):\n            img = dropout(img,DIM=[Config[\"dim_y\"],Config[\"dim_x\"]], CT = 12, SZ = 0.08, mode='fill', fill_with = 0)\n        else:\n            img = dropout(img,DIM=[Config[\"dim_y\"],Config[\"dim_x\"]], CT = 12, SZ = 0.08, mode='fill', fill_with = 255)\n    \n    return img","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:07.838798Z","iopub.execute_input":"2023-01-17T19:22:07.839299Z","iopub.status.idle":"2023-01-17T19:22:07.853624Z","shell.execute_reply.started":"2023-01-17T19:22:07.839263Z","shell.execute_reply":"2023-01-17T19:22:07.852505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _parse_image_function(example_proto, augment, normalize):\n    image_feature_description = {\n        'image': tf.io.FixedLenFeature([], tf.string),\n        'cancer': tf.io.FixedLenFeature([], tf.int64),\n        'laterality': tf.io.FixedLenFeature([], tf.int64),\n        'view': tf.io.FixedLenFeature([], tf.int64),\n        'age': tf.io.FixedLenFeature([], tf.float32),\n        'implant': tf.io.FixedLenFeature([], tf.int64),\n        'machine_id': tf.io.FixedLenFeature([], tf.int64),\n        'site_id': tf.io.FixedLenFeature([], tf.int64),\n    }\n    \n    single_example = tf.io.parse_single_example(example_proto, image_feature_description)\n    \n    output_dict = {}\n    #for png images\n    #image = tf.reshape(tf.io.decode_png(single_example['image'],dtype=np.dtype('uint8')), (dim,dim,3))\n    #for raw images\n    image = tf.reshape(tf.io.decode_raw(single_example['image'],out_type=np.dtype('uint8')), (Config[\"dim_y\"],Config[\"dim_x\"],3))\n    cancer =  single_example['cancer']\n    #laterality =  single_example['laterality']\n    #view =  single_example['view']\n    #age =  single_example['age']\n    #implant =  single_example['implant']\n    #machine_id =  single_example['machine_id']\n    #site_id =  single_example['site_id']\n    \n    if(augment and Config[\"auguments\"]):\n        image = apply_augments(image)\n    \n    if(normalize and Config[\"include_preprocessing\"]):\n        image = tf.cast(image, tf.float32)\n        image = keras.applications.imagenet_utils.preprocess_input(image, mode='torch')\n        \n        if(Config[\"quantize_model\"]):\n            image = tf.cast(image, tf.float16)\n                         \n    \n    return (image, cancer)\n\n\ndef load_dataset(filenames, augment = True, normalize=True):\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=Config[\"AUTO\"], compression_type = 'GZIP')\n    dataset = dataset.map(lambda ex: _parse_image_function(ex,augment,normalize))\n    return dataset\n\n\ndef get_dataset(filenames, augment = True, normalize=True):\n    dataset = load_dataset(filenames, augment, normalize)\n    #dataset = dataset.shuffle(Config[\"REPLICAS\"]*Config[\"batch_size\"]*4, seed = Config[\"seed\"])\n    dataset = dataset.batch(Config[\"REPLICAS\"]*Config[\"batch_size\"],drop_remainder=True)\n    dataset = dataset.prefetch(Config[\"AUTO\"])\n    return dataset","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:07.855344Z","iopub.execute_input":"2023-01-17T19:22:07.855885Z","iopub.status.idle":"2023-01-17T19:22:07.872143Z","shell.execute_reply.started":"2023-01-17T19:22:07.855845Z","shell.execute_reply":"2023-01-17T19:22:07.871078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_dict=next(iter(get_dataset(tfrecords_file_paths, augment=False, normalize=False)))\n\nprint('image shape:',dataset_dict[0].shape)\ndisplay({'cancer':dataset_dict[1]})\n\nfig=plt.figure(figsize=(20, 10))\nfor i in range(0,Config[\"REPLICAS\"]*Config[\"batch_size\"]):\n    fig.add_subplot(Config[\"batch_size\"],Config[\"REPLICAS\"],i+1)\n    plt.imshow(dataset_dict[0][i], cmap='bone')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:07.873632Z","iopub.execute_input":"2023-01-17T19:22:07.873991Z","iopub.status.idle":"2023-01-17T19:22:14.525237Z","shell.execute_reply.started":"2023-01-17T19:22:07.873958Z","shell.execute_reply":"2023-01-17T19:22:14.523942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_dict=next(iter(get_dataset(tfrecords_file_paths, augment=True, normalize=False)))\n\nprint('image shape:',dataset_dict[0].shape)\ndisplay({'cancer':dataset_dict[1]})\n\nfig=plt.figure(figsize=(20, 10))\nfor i in range(0,Config[\"REPLICAS\"]*Config[\"batch_size\"]):\n    fig.add_subplot(Config[\"batch_size\"],Config[\"REPLICAS\"],i+1)\n    plt.imshow(dataset_dict[0][i], cmap='bone')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:14.526918Z","iopub.execute_input":"2023-01-17T19:22:14.527968Z","iopub.status.idle":"2023-01-17T19:22:21.936863Z","shell.execute_reply.started":"2023-01-17T19:22:14.527929Z","shell.execute_reply":"2023-01-17T19:22:21.935677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class StackMeanPool(keras.layers.Layer):\n    def __init__(self,layers_num):\n        super().__init__()\n        \n        self.layer_pooler = layers.GlobalAveragePooling2D()\n        \n        \n    def call(self, inputs, training=False):\n        pooled_layers = []\n        \n        for layer in inputs:\n            pooled_layer = self.layer_pooler(layer)\n            pooled_layers += [pooled_layer]\n        \n        output = tf.concat(pooled_layers,axis=1)\n        \n        return output","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:21.938634Z","iopub.execute_input":"2023-01-17T19:22:21.939Z","iopub.status.idle":"2023-01-17T19:22:21.946323Z","shell.execute_reply.started":"2023-01-17T19:22:21.938966Z","shell.execute_reply":"2023-01-17T19:22:21.945173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def quantize_model(model):\n    def apply_quantization_to_dense(layer):\n        if isinstance(layer, tf.keras.layers.Dense):\n            return tfmot.quantization.keras.quantize_annotate_layer(layer)\n        return layer\n\n    # Use `tf.keras.models.clone_model` to apply `apply_quantization_to_dense` \n    # to the layers of the model.\n    annotated_model = tf.keras.models.clone_model(\n        model,\n        clone_function=apply_quantization_to_dense,\n    )\n\n    # Now that the Dense layers are annotated,\n    # `quantize_apply` actually makes the model quantization aware.\n    quant_aware_model = tfmot.quantization.keras.quantize_apply(annotated_model)\n    return quant_aware_model","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:21.948055Z","iopub.execute_input":"2023-01-17T19:22:21.948403Z","iopub.status.idle":"2023-01-17T19:22:21.965871Z","shell.execute_reply.started":"2023-01-17T19:22:21.94837Z","shell.execute_reply":"2023-01-17T19:22:21.964432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TransformerModel(tf.keras.Model):\n    def __init__(self, pooling=False):\n        super().__init__()\n        \n        transformer_type = getattr(kecam, Config[\"model_type\"])\n        transformer_model = getattr(transformer_type,Config[\"model_name\"])(num_classes=0, input_shape=(Config['dim_y'],Config['dim_x'],3),\n                                                          pretrained=Config[\"model_weights\"])\n        self.base_model = quantize_model(transformer_model)\n        \n        self.pooling = pooling\n        if(pooling):\n            self.pooler = layers.GlobalAveragePooling2D()\n            \n        \n        self.output_dropout_list = [layers.Dropout(0.1*i) for i in range(1,6)]\n        self.dense = layers.Dense(64)\n        \n        self.leaky_relu = layers.LeakyReLU(alpha=0.05)\n        \n        self.classifier = layers.Dense(1, activation = \"sigmoid\")\n    \n    def call(self, inputs, training=False):\n        image = inputs\n        \n        x = self.base_model(image)\n        \n        if(self.pooling):\n            x = self.pooler(x)\n        \n        #multi-sample dropout\n        x_classified_list = []\n        for i in range(0,5):\n            x_to_dense = x\n            if(training):\n                x_to_dense = self.output_dropout_list[i](x_to_dense)\n            \n            x_densed = self.dense(x_to_dense)\n            x_densed_act = self.leaky_relu(x_densed)\n            \n            x_classified = self.classifier(x_densed_act)\n            x_classified_list += [x_classified]\n        \n        output_stacked = tf.stack(x_classified_list, axis = 0)\n        output = tf.reduce_mean(output_stacked, axis = 0)\n        \n        return output","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:21.968474Z","iopub.execute_input":"2023-01-17T19:22:21.969565Z","iopub.status.idle":"2023-01-17T19:22:21.98236Z","shell.execute_reply.started":"2023-01-17T19:22:21.969513Z","shell.execute_reply":"2023-01-17T19:22:21.981376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#check model\ninputs = layers.Input(shape = (Config[\"dim_y\"],Config[\"dim_x\"],3), dtype =\"float16\" if Config[\"include_preprocessing\"] else np.dtype('uint8'))\n\nwith Config[\"strategy\"].scope():\n    model = TransformerModel()\n    \n    model(inputs,training=True)\n\n    print(model.summary())\n\n#show layers\nmodel_plot = plot_model(\n    tf.keras.Model(\n        inputs=inputs,\n        outputs=model.call(inputs,training=True)\n    ),\n    to_file='model.png', \n    dpi=56,  \n    show_shapes=True, \n    show_layer_names=True,\n    expand_nested=False\n)\ndisplay(model_plot)\n\ndel model,model_plot\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:22:21.983837Z","iopub.execute_input":"2023-01-17T19:22:21.984169Z","iopub.status.idle":"2023-01-17T19:23:07.296688Z","shell.execute_reply.started":"2023-01-17T19:22:21.984137Z","shell.execute_reply":"2023-01-17T19:23:07.294783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def plot_history(history, fold):\n    history = pd.DataFrame(history.history)\n    plt.figure(figsize=[15, 5])\n    ax = sns.lineplot(data=history,linewidth=2,dashes=False)\n\n    ax.set(title=f\"Training ProtBert fold {fold}\",xlabel=\"epoch\")\n    sns.set_palette(\"colorblind\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:23:07.298298Z","iopub.execute_input":"2023-01-17T19:23:07.298654Z","iopub.status.idle":"2023-01-17T19:23:07.306479Z","shell.execute_reply.started":"2023-01-17T19:23:07.298623Z","shell.execute_reply":"2023-01-17T19:23:07.305718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# tensorflow\ndef pfbeta_tf(labels, preds, beta=1):\n    eps = 1e-5\n    preds = tf.clip_by_value(preds, 0, 1)\n    y_true_count = tf.reduce_sum(labels)\n    ctp = tf.reduce_sum(preds[labels==1])\n    cfp = tf.reduce_sum(preds[labels==0])\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp + eps)\n    c_recall = ctp / (y_true_count + eps)\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall + eps)\n        return result\n    else:\n        return tf.constant(0, dtype=tf.float32)\npfbeta_tf.__name__='pF1'\n\n\n# finds best pf1 using thresholds\ndef pfbeta_thr(labels, preds):\n    thrs = tf.range(0, 1, 0.05)\n    best_score = tf.constant(0, dtype=tf.float32)\n    for thr in thrs:\n        score = pfbeta_tf(labels, tf.cast(preds>thr, tf.float32))\n        best_score = tf.cond(score > best_score, lambda: score, lambda: best_score)\n    return best_score\n\npfbeta_thr.__name__='pF1_thr'\n\n# numpy\ndef pfbeta(labels, preds, beta=1):\n    eps = 1e-5\n    preds = preds.clip(0, 1)\n    y_true_count = labels.sum()\n    ctp = preds[labels==1].sum()\n    cfp = preds[labels==0].sum()\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp + eps)\n    c_recall = ctp / (y_true_count + eps)\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall + eps)\n        return result\n    else:\n        return 0.0","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:23:07.307685Z","iopub.execute_input":"2023-01-17T19:23:07.308372Z","iopub.status.idle":"2023-01-17T19:23:07.51367Z","shell.execute_reply.started":"2023-01-17T19:23:07.308342Z","shell.execute_reply":"2023-01-17T19:23:07.512692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def find_best_thr(labels, preds, fold):\n    thrs = np.arange(0,1,0.001)\n    scores = []\n    for thr in tqdm(thrs):\n        scores+=[pfbeta(labels.astype('float32'), preds.astype('float32')>thr)]\n        \n    best_score_idx = np.argmax(scores)\n    best_score = np.max(scores)\n    best_thr = thrs[best_score_idx]\n    print(f'\\nFold {fold} MAX pF1 = {best_score: 0.3f} @ {best_thr:0.3f}\\n')\n\n    plt.figure(figsize=[15, 5])\n    ax = sns.lineplot(x = thrs, y = scores,linewidth=2)\n    ax.axvline(x=best_thr, color='blue', ls='--')\n    ax.plot(best_thr, best_score, color='blue', marker='o', markersize=12)\n    ax.fill_between(thrs, scores, color='cyan', alpha=0.3, )\n    \n    ax.set(title=f\"Threshold Vs pF1 {fold}\", xlabel=\"Threshold\", ylabel=\"pF1\")\n    sns.set_palette(\"colorblind\")\n    plt.show()\n    \n    return best_thr","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:23:07.515115Z","iopub.execute_input":"2023-01-17T19:23:07.515742Z","iopub.status.idle":"2023-01-17T19:23:07.5277Z","shell.execute_reply.started":"2023-01-17T19:23:07.515707Z","shell.execute_reply":"2023-01-17T19:23:07.526776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.makedirs(\"./model_weights\",exist_ok =True)\n\ntfrecords_file_paths = np.array(tfrecords_file_paths)\n\nkf = KFold(n_splits=Config[\"n_fold\"],random_state=Config[\"seed\"], shuffle=True)\n\nbest_val_pF1_thrs = []\nbest_val_thrs = []\n\nfor fold,(train_index, val_index) in enumerate(kf.split(X=tfrecords_file_paths)):\n    print(f\"Fold {fold}:\")\n    \n    train_paths, val_paths = tfrecords_file_paths[train_index], tfrecords_file_paths[val_index]\n    train_ds, val_ds = get_dataset(train_paths, augment=True), get_dataset(val_paths, augment=False)\n    \n    \n    tf.keras.backend.clear_session()\n    \n    with Config[\"strategy\"].scope():\n        model = TransformerModel()\n        \n        optimizers=[]\n        \n        base_model_layers = model.base_model.layers\n        base_model_layers_reversed = list(reversed(base_model_layers))\n        #classifier layer (high lr)\n        dense_layer = model.dense\n        classifier_layer = model.classifier\n        \n        #layer-wise lr configs\n        llrdr = 0.995 if Config[\"use_layerwise_lr\"] else 1\n        \n        num_of_tfrecords = 54634\n        lr_sch_decay_steps = (num_of_tfrecords/Config[\"n_fold\"]*(Config[\"n_fold\"]-1)) // (Config[\"batch_size\"]*Config[\"REPLICAS\"])\n        \n        Config[\"min_learning_rate\"] = Config[\"learning_rate\"]*Config[\"min_lr_mul\"]\n        sch_decay = (Config[\"min_learning_rate\"]/Config[\"learning_rate\"])**(1/Config[\"epochs\"])\n        sch_decay = sch_decay if Config[\"use_scheduler\"] else 1\n        last_layer_mul = 10 if Config[\"last_layer_speedup\"] else 1\n        \n        lr_schedules_low_to_normal = [\n            keras.optimizers.schedules.ExponentialDecay(\n            initial_learning_rate=Config[\"learning_rate\"]*llrdr**i,\n            decay_steps=lr_sch_decay_steps, \n            decay_rate=sch_decay,\n            ) for i in range(len(base_model_layers_reversed))\n        ]\n        lr_schedules_fast = keras.optimizers.schedules.ExponentialDecay(\n            initial_learning_rate=Config[\"learning_rate\"]*1,\n            decay_steps=lr_sch_decay_steps, \n            decay_rate=sch_decay,\n        )\n        \n        optim = None\n        if(Config[\"optimizer\"]==\"Adam\"):\n            optim = lambda lr:keras.optimizers.Adam(learning_rate = lr)\n        elif(Config[\"optimizer\"]==\"AdamW\"):\n            optim = lambda lr:tfa.optimizers.AdamW(learning_rate = lr, weight_decay = Config[\"weight_decay\"])\n\n        optimizers += [(optim(lr_schedules_low_to_normal[i]),base_model_layers_reversed[i]) for i in range(len(base_model_layers_reversed))]\n        #high lr\n        optimizers += [(optim(lr_schedules_fast),[dense_layer,classifier_layer])]\n        #optimizers += [(keras.optimizers.Adam(learning_rate = Config[\"learning_rate\"]*10),[classifier_layer])]\n\n        #print(optimizers[0][0].lr(optimizers[0][0].iterations))\n        if(Config[\"loss\"]=='binary'):\n            loss = keras.losses.BinaryCrossentropy(label_smoothing=0.05)\n        elif(Config[\"loss\"]=='focal'):\n            loss = tfa.losses.SigmoidFocalCrossEntropy(alpha=0.80, gamma=2.0)\n           \n        if(Config[\"use_wandb\"]):\n            wandb_init(fold)\n        \n        model.compile(\n            optimizer=tfa.optimizers.MultiOptimizer(optimizers),\n            loss=loss,\n            metrics=[pfbeta_tf,pfbeta_thr,keras.metrics.AUC()]\n        )\n\n        early_stopping = tf.keras.callbacks.EarlyStopping(\n            monitor=\"val_pF1_thr\",\n            patience=Config[\"early_stopping_patience\"],\n            restore_best_weights=True,\n            min_delta=1e-6,\n            mode='max',\n            verbose=1\n        )\n        \n        chp_callback = keras.callbacks.ModelCheckpoint(\n            filepath = f\"./model_weights/best_model_fold_{fold}.h5\",\n            monitor=\"val_pF1_thr\",\n            mode=\"max\",\n            save_best_only=True,\n            verbose=1,\n            save_weights_only=True\n        )\n\n        callbacks = [early_stopping, chp_callback]\n        \n        if(Config[\"use_wandb\"]):\n            callbacks += [wandb.keras.WandbCallback(save_model=False, monitor=\"val_pF1_thr\", mode=\"max\")]\n    \n    \n    history = model.fit(\n        train_ds,\n        validation_data=val_ds,\n        callbacks = callbacks,\n        class_weight = class_weight,\n        epochs=Config[\"epochs\"]\n    )\n    \n    \n    #plot_history(history, fold)\n    \n    y_pred = model.predict(\n        val_ds\n    )\n    \n    y_val = []\n    for (_,label) in val_ds:\n        y_val += label.numpy().tolist()\n    y_val = np.array(y_val)\n    \n    display(y_pred)\n    \n    if(Config[\"use_wandb\"]):\n        wandb.run.finish()\n    \n    best_val_pF1_thrs += [np.max(history.history[\"val_pF1_thr\"])]\n    best_val_thrs += [find_best_thr(y_val ,y_pred, fold)]\n    \n    del model, train_ds, val_ds, history\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-01-17T19:23:07.529658Z","iopub.execute_input":"2023-01-17T19:23:07.530559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.save('best_val_pF1_thrs.npy', best_val_pF1_thrs)\nbest_val_pF1_thrs","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.save('best_val_thrs.npy', best_val_thrs)\nbest_val_thrs","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}