{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":23870,"databundleVersionId":1781260,"sourceType":"competition"},{"sourceId":2021283,"sourceType":"datasetVersion","datasetId":1209783},{"sourceId":2023916,"sourceType":"datasetVersion","datasetId":1211369}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Train Knowledge Distillation model\nThis notebook's idea is alomost same as 3 stage training's step2 in [RANZCR / ResNet200D / 3-stage training / step2](https://www.kaggle.com/yasufuminakama/ranzcr-resnet200d-3-stage-training-step2).\n\nKeras version implementation is borrowed from Keras official \"Knowledge Distillation\" example code.\n[Knowledge Distillation]\n(https://keras.io/examples/vision/knowledge_distillation/#construct-distiller-class)\n\n   \n\nTeacher model training notebook:  \n[[Keras TPU] RANZCR Train annotation](https://www.kaggle.com/enukuro/keras-tpu-ranzcr-train-annotation)\n\n\nCreatting annotation tfrecords notebook:   \n[Annotation RANZCR CLiP 900](https://www.kaggle.com/enukuro/annotation-ranzcr-clip-900)","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y tensorflow\n!pip install tensorflow==2.14.0","metadata":{"execution":{"iopub.status.busy":"2024-08-16T14:59:02.643047Z","iopub.execute_input":"2024-08-16T14:59:02.643794Z","iopub.status.idle":"2024-08-16T15:00:05.962603Z","shell.execute_reply.started":"2024-08-16T14:59:02.643761Z","shell.execute_reply":"2024-08-16T15:00:05.961626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install tensorflow-addons[tensorflow]","metadata":{"execution":{"iopub.status.busy":"2024-08-16T15:00:05.964393Z","iopub.execute_input":"2024-08-16T15:00:05.964700Z","iopub.status.idle":"2024-08-16T15:00:11.822501Z","shell.execute_reply.started":"2024-08-16T15:00:05.964667Z","shell.execute_reply":"2024-08-16T15:00:11.821657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\nimport numpy as np\nimport pandas as pd\nimport random\nimport tensorflow as tf\nimport tensorflow.keras.layers as L\nimport tensorflow_addons as tfa\nfrom keras import backend as K\nfrom kaggle_datasets import KaggleDatasets\nfrom keras.applications.xception import Xception as BaseModel\nfrom keras.applications.xception import preprocess_input\nimport itertools\nimport gc","metadata":{"id":"TB205EIUGzHu","outputId":"d88358d9-08cc-42c8-dfc2-4edf7fa1eeed","execution":{"iopub.status.busy":"2024-08-16T15:00:16.674796Z","iopub.execute_input":"2024-08-16T15:00:16.675161Z","iopub.status.idle":"2024-08-16T15:00:29.401674Z","shell.execute_reply.started":"2024-08-16T15:00:16.675127Z","shell.execute_reply":"2024-08-16T15:00:29.400963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\nprint(\"Tensorflow version \" + tf.__version__)\n\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\nexcept ValueError:\n    tpu = None\n\nif tpu:\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    strategy = tf.distribute.get_strategy()\n\nAUTOTUNE = tf.data.experimental.AUTOTUNE","metadata":{"id":"nUlL507YG1G0","outputId":"4ab30bb3-ec02-495f-817a-44a68ec4602d","execution":{"iopub.status.busy":"2024-08-16T15:00:45.605309Z","iopub.execute_input":"2024-08-16T15:00:45.605919Z","iopub.status.idle":"2024-08-16T15:00:45.613896Z","shell.execute_reply.started":"2024-08-16T15:00:45.605883Z","shell.execute_reply":"2024-08-16T15:00:45.613162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 42\n\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n    \nseed_everything(SEED)\n\nBATCH_SIZE = strategy.num_replicas_in_sync * 16\nIMG_SIZE = 900\nNUM_CLASSES = 11\n\ntarget_fold = 1","metadata":{"id":"1kOVhjvwG23G","execution":{"iopub.status.busy":"2024-08-16T15:00:47.758806Z","iopub.execute_input":"2024-08-16T15:00:47.759262Z","iopub.status.idle":"2024-08-16T15:00:47.764156Z","shell.execute_reply.started":"2024-08-16T15:00:47.759220Z","shell.execute_reply":"2024-08-16T15:00:47.763431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feature_map = {\n    'image': tf.io.FixedLenFeature([], tf.string),\n    'image_annotation': tf.io.FixedLenFeature([], tf.string),\n    'StudyInstanceUID': tf.io.FixedLenFeature([], tf.string),  \n    'ETT - Abnormal': tf.io.FixedLenFeature([], tf.int64),\n    'ETT - Borderline': tf.io.FixedLenFeature([], tf.int64),\n    'ETT - Normal': tf.io.FixedLenFeature([], tf.int64),\n    'NGT - Abnormal': tf.io.FixedLenFeature([], tf.int64),\n    'NGT - Borderline': tf.io.FixedLenFeature([], tf.int64),\n    'NGT - Incompletely Imaged': tf.io.FixedLenFeature([], tf.int64),\n    'NGT - Normal': tf.io.FixedLenFeature([], tf.int64),\n    'CVC - Abnormal': tf.io.FixedLenFeature([], tf.int64),\n    'CVC - Borderline': tf.io.FixedLenFeature([], tf.int64),\n    'CVC - Normal': tf.io.FixedLenFeature([], tf.int64),\n    'Swan Ganz Catheter Present': tf.io.FixedLenFeature([], tf.int64)}\n\ndef decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    image = tf.image.resize(image, [IMG_SIZE, IMG_SIZE])\n    image = tf.reshape(image, [IMG_SIZE, IMG_SIZE, 3])\n    return tf.cast(image, tf.float32)\n\ndef read_tfrecord(example):\n    example = tf.io.parse_single_example(example, feature_map)\n    image = decode_image(example['image'])\n    image_annotation = decode_image(example['image_annotation'])\n    target = [\n        example['ETT - Abnormal'],\n        example['ETT - Borderline'],\n        example['ETT - Normal'],\n        example['NGT - Abnormal'],\n        example['NGT - Borderline'],\n        example['NGT - Incompletely Imaged'],\n        example['NGT - Normal'],\n        example['CVC - Abnormal'],\n        example['CVC - Borderline'],\n        example['CVC - Normal'],\n        example['Swan Ganz Catheter Present']]\n    return [preprocess_input(image), preprocess_input(image_annotation)], tf.cast(target, tf.float32)\n\ndef data_augment(img, target):\n    img = tf.map_fn(lambda x: tf.image.random_flip_left_right(x), img)\n    return img, target\n\ndef get_dataset(filenames, shuffled=False, repeated=False, \n                cached=False, augmented=False):\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTOTUNE)\n    dataset = dataset.map(read_tfrecord, num_parallel_calls=AUTOTUNE)\n    if cached:\n        dataset = dataset.cache()\n    if shuffled:\n        dataset = dataset.shuffle(1024, seed=SEED)\n    if augmented:\n        dataset = dataset.map(data_augment, num_parallel_calls=AUTOTUNE)\n    if repeated:\n        dataset = dataset.repeat()\n    dataset = dataset.batch(BATCH_SIZE, drop_remainder=True)\n    dataset = dataset.prefetch(AUTOTUNE)\n    return dataset\n\n","metadata":{"id":"XJ3-_YipG53o","execution":{"iopub.status.busy":"2024-08-16T15:00:49.256053Z","iopub.execute_input":"2024-08-16T15:00:49.256430Z","iopub.status.idle":"2024-08-16T15:00:49.267783Z","shell.execute_reply.started":"2024-08-16T15:00:49.256398Z","shell.execute_reply":"2024-08-16T15:00:49.267036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://keras.io/examples/vision/knowledge_distillation/#construct-distiller-class\nclass Distiller(tf.keras.Model):\n    def __init__(self, student, teacher):\n        super(Distiller, self).__init__()\n        self.teacher = teacher\n        self.student = student\n\n    def compile(\n        self,\n        optimizer,\n        metrics,\n        student_loss_fn,\n        distillation_loss_fn,\n        alpha=0.1,\n    ):\n        \"\"\" Configure the distiller.\n\n        Args:\n            optimizer: Keras optimizer for the student weights\n            metrics: Keras metrics for evaluation\n            student_loss_fn: Loss function of difference between student\n                predictions and ground-truth\n            distillation_loss_fn: Loss function of difference between soft\n                student predictions and soft teacher predictions\n            alpha: weight to student_loss_fn and 1-alpha to distillation_loss_fn\n        \"\"\"\n        super(Distiller, self).compile(optimizer=optimizer, metrics=metrics)\n        self.student_loss_fn = student_loss_fn\n        self.distillation_loss_fn = distillation_loss_fn\n        self.alpha = alpha\n        \n    @tf.function\n    def train_step(self, data):\n        x, y = data\n        image, image_annotation = tf.split(x, 2, axis=1)\n        image = tf.squeeze(image)\n        image_annotation = tf.squeeze(image_annotation)\n        \n        teacher_predictions, teacher_features = self.teacher(image_annotation, training=False)\n        with tf.GradientTape() as tape:\n            student_predictions, student_features = self.student(image, training=True)\n         \n            student_loss = self.student_loss_fn(y, student_predictions)\n            student_loss = tf.reduce_sum(student_loss * (1. / BATCH_SIZE))\n            distillation_loss = self.distillation_loss_fn(tf.reshape(teacher_features, [BATCH_SIZE, -1]), tf.reshape(student_features, [BATCH_SIZE, -1]))\n            # distillation_loss = tf.reduce_sum(distillation_loss * (1. / BATCH_SIZE))\n            loss = self.alpha * student_loss + (1 - self.alpha) * distillation_loss\n            \n        # Compute gradients\n        trainable_vars = self.student.trainable_variables\n        gradients = tape.gradient(loss, trainable_vars)\n\n        # Update weights\n        self.optimizer.apply_gradients(zip(gradients, trainable_vars))\n\n        # Update the metrics configured in `compile()`.\n        self.compiled_metrics.update_state(y, student_predictions)\n\n        # Return a dict of performance\n        results = {m.name: m.result() for m in self.metrics}\n        results.update(\n            {\"student_loss\": student_loss, \"distillation_loss\": distillation_loss}\n        )\n        return results\n    \n    @tf.function\n    def valid_step(self, data):\n        x, y = data\n        image, image_annotation = tf.split(x, 2, axis=1)\n        image = tf.squeeze(image)\n        # Compute predictions\n        y_prediction, _ = self.student(image, training=False)\n\n        # Calculate the loss\n        student_loss = self.student_loss_fn(y, y_prediction)\n        student_loss = tf.reduce_sum(student_loss * (1. / BATCH_SIZE))\n        # Update the metrics.\n        self.compiled_metrics.update_state(y, y_prediction)\n\n        # Return a dict of performance\n        results = {m.name: m.result() for m in self.metrics}\n        results.update({\"student_loss\": student_loss})\n        return results\n    \n    @tf.function\n    def test_step(self, data):\n        x, y = data\n        image, image_annotation = tf.split(x, 2, axis=1)\n        image = tf.squeeze(image)\n        # Compute predictions\n        y_prediction, _ = self.student(image, training=False)\n\n        # Calculate the loss\n        student_loss = self.student_loss_fn(y, y_prediction)\n        student_loss = tf.reduce_sum(student_loss * (1. / BATCH_SIZE))\n        # Update the metrics.\n        self.compiled_metrics.update_state(y, y_prediction)\n\n        # Return a dict of performance\n        results = {m.name: m.result() for m in self.metrics}\n        results.update({\"student_loss\": student_loss})\n        return results","metadata":{"id":"SvcAybCiNcyD","execution":{"iopub.status.busy":"2024-08-16T15:00:50.911003Z","iopub.execute_input":"2024-08-16T15:00:50.911585Z","iopub.status.idle":"2024-08-16T15:00:50.925267Z","shell.execute_reply.started":"2024-08-16T15:00:50.911550Z","shell.execute_reply":"2024-08-16T15:00:50.924545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pwd","metadata":{"execution":{"iopub.status.busy":"2024-08-16T15:03:34.999257Z","iopub.execute_input":"2024-08-16T15:03:34.999759Z","iopub.status.idle":"2024-08-16T15:03:35.149254Z","shell.execute_reply.started":"2024-08-16T15:03:34.999717Z","shell.execute_reply":"2024-08-16T15:03:35.147890Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd ..","metadata":{"execution":{"iopub.status.busy":"2024-08-16T15:03:47.335823Z","iopub.execute_input":"2024-08-16T15:03:47.336267Z","iopub.status.idle":"2024-08-16T15:03:47.344331Z","shell.execute_reply.started":"2024-08-16T15:03:47.336231Z","shell.execute_reply":"2024-08-16T15:03:47.343466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%ls","metadata":{"execution":{"iopub.status.busy":"2024-08-16T15:03:52.589218Z","iopub.execute_input":"2024-08-16T15:03:52.589670Z","iopub.status.idle":"2024-08-16T15:03:52.739737Z","shell.execute_reply.started":"2024-08-16T15:03:52.589634Z","shell.execute_reply":"2024-08-16T15:03:52.738359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd input","metadata":{"execution":{"iopub.status.busy":"2024-08-16T15:04:01.926108Z","iopub.execute_input":"2024-08-16T15:04:01.927228Z","iopub.status.idle":"2024-08-16T15:04:01.933127Z","shell.execute_reply.started":"2024-08-16T15:04:01.927184Z","shell.execute_reply":"2024-08-16T15:04:01.932288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%ls","metadata":{"execution":{"iopub.status.busy":"2024-08-16T15:04:07.461373Z","iopub.execute_input":"2024-08-16T15:04:07.461816Z","iopub.status.idle":"2024-08-16T15:04:07.612939Z","shell.execute_reply.started":"2024-08-16T15:04:07.461784Z","shell.execute_reply":"2024-08-16T15:04:07.611738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd ranzcr-clip-catheter-line-classification/","metadata":{"execution":{"iopub.status.busy":"2024-08-16T15:04:19.037558Z","iopub.execute_input":"2024-08-16T15:04:19.037989Z","iopub.status.idle":"2024-08-16T15:04:19.044344Z","shell.execute_reply.started":"2024-08-16T15:04:19.037954Z","shell.execute_reply":"2024-08-16T15:04:19.043473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%ls","metadata":{"execution":{"iopub.status.busy":"2024-08-16T15:04:24.949084Z","iopub.execute_input":"2024-08-16T15:04:24.950091Z","iopub.status.idle":"2024-08-16T15:04:25.107317Z","shell.execute_reply.started":"2024-08-16T15:04:24.950053Z","shell.execute_reply":"2024-08-16T15:04:25.106020Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(name):\n    base_model = BaseModel(input_shape=(IMG_SIZE,IMG_SIZE,3), include_top=False, weights='imagenet', pooling=\"avg\")\n    base_model_output = base_model.output\n    x = L.Dropout(0.5)(base_model_output)\n    outputs = L.Dense(NUM_CLASSES, activation=\"sigmoid\")(x)\n\n    model = tf.keras.models.Model(inputs=base_model.input, outputs=[outputs, base_model_output], name=name)\n    return model\n\ndef get_distiller_model(fold=0):\n    with strategy.scope():\n        student = get_model('student')\n        teacher = get_model('teacher')\n        teacher.load_weights(f'../input/ranzcr-annotation-teacher/teacher_model_{fold}.h5')\n\n        distiller = Distiller(student=student, teacher=teacher)\n        distiller.compile(\n            optimizer=tf.keras.optimizers.Adam(lr=1e-3),\n            metrics=[tf.keras.metrics.AUC(multi_label=True)],\n            student_loss_fn=tfa.losses.SigmoidFocalCrossEntropy(alpha = 0.5, gamma = 2, reduction=tf.keras.losses.Reduction.NONE),\n            distillation_loss_fn=tf.keras.losses.MSE,\n            alpha=0.3\n        )\n    return distiller","metadata":{"id":"qoUr3sHuG7--","execution":{"iopub.status.busy":"2024-08-16T15:00:52.678871Z","iopub.execute_input":"2024-08-16T15:00:52.679502Z","iopub.status.idle":"2024-08-16T15:00:52.685815Z","shell.execute_reply.started":"2024-08-16T15:00:52.679467Z","shell.execute_reply":"2024-08-16T15:00:52.685140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TF_REC_DS_PATH = KaggleDatasets().get_gcs_path('ranzcr-annotation-900-tfrecords')\n\ntfrec_files = []\nfor fold in range(5):\n    training_files = [TF_REC_DS_PATH + f'/{fold}_{num}.tfrec' for num in range(0,5)]\n    random.shuffle(training_files)\n    tfrec_files.append(training_files)","metadata":{"id":"e5TH_x21G-57","outputId":"6a5b3eaf-5956-46da-f527-5989db994d19","execution":{"iopub.status.busy":"2024-08-16T15:00:54.232814Z","iopub.execute_input":"2024-08-16T15:00:54.233628Z","iopub.status.idle":"2024-08-16T15:00:54.238312Z","shell.execute_reply.started":"2024-08-16T15:00:54.233594Z","shell.execute_reply":"2024-08-16T15:00:54.237520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fold_item_counts = [1804, 1783, 1809, 1851, 1848]\n\nfor fold in range(5):\n    \n    if fold != target_fold:\n        continue\n        \n    print(f'fold_{fold} start')\n    train_filenames = list(itertools.chain.from_iterable([tfrec_files[i] for i in range(5) if i != fold]))\n    val_filenames = tfrec_files[fold]\n\n    random.shuffle(train_filenames)\n\n    train_dataset = get_dataset(train_filenames, shuffled=True, augmented=True, repeated=True)\n    val_dataset = get_dataset(val_filenames, shuffled=False, cached=True)\n\n    steps_per_epoch = (sum(fold_item_counts) - fold_item_counts[fold]) // BATCH_SIZE\n    validation_steps = fold_item_counts[fold] // BATCH_SIZE\n    \n    model = get_distiller_model(fold)\n    \n    sv = tf.keras.callbacks.ModelCheckpoint(f'distiller_model_{fold}.h5', monitor='val_student_loss', verbose=1, save_best_only=True, save_weights_only=True)\n    reduce_lr = tf.keras.callbacks.ReduceLROnPlateau(monitor='val_student_loss', verbose=1, factor=0.1, patience=3, min_delta=0.0001, min_lr=1e-6)\n\n    model.fit(\n        train_dataset,\n        steps_per_epoch=steps_per_epoch,\n        epochs=15,\n        callbacks=[reduce_lr, sv],\n        validation_data=val_dataset,\n    )\n    model.built = True\n    model.load_weights(f'./distiller_model_{fold}.h5')\n    model.get_layer('student').save_weights(f'student_model_{fold}.h5')\n    \n#     tf.keras.backend.clear_session()\n#     del model\n#     gc.collect()\n    ","metadata":{"id":"Zwj5TbIIHAr9","outputId":"e5b361e3-303f-405a-92ec-96d3657498d3","execution":{"iopub.status.busy":"2024-08-16T15:00:58.391724Z","iopub.execute_input":"2024-08-16T15:00:58.392079Z","iopub.status.idle":"2024-08-16T15:01:05.447756Z","shell.execute_reply.started":"2024-08-16T15:00:58.392052Z","shell.execute_reply":"2024-08-16T15:01:05.446484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}