{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-10-12T17:23:07.538888Z","iopub.execute_input":"2023-10-12T17:23:07.540091Z","iopub.status.idle":"2023-10-12T17:23:07.546015Z","shell.execute_reply.started":"2023-10-12T17:23:07.540048Z","shell.execute_reply":"2023-10-12T17:23:07.545073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def append_ext(fn):\n    return fn+\".png\"\n\ntrain_dataframe = pd.read_csv(\"/kaggle/input/aptos2019-blindness-detection/train.csv\",dtype=str)\ntrain_dataframe['id_code'] = train_dataframe['id_code'].apply(append_ext)\ntrain_dataframe.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:07.548309Z","iopub.execute_input":"2023-10-12T17:23:07.549018Z","iopub.status.idle":"2023-10-12T17:23:07.600814Z","shell.execute_reply.started":"2023-10-12T17:23:07.548985Z","shell.execute_reply":"2023-10-12T17:23:07.599851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"diagnosis_counts = train_dataframe.groupby('diagnosis')['id_code'].count()\n\n# Plot the counts\ndiagnosis_counts.plot(kind='bar')\n\n# Add labels and title to the plot\nplt.xlabel('Diagnosis')\nplt.ylabel('Count')\nplt.title('Count of id_code for each diagnosis')\n\n# Display the plot\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:07.602133Z","iopub.execute_input":"2023-10-12T17:23:07.603114Z","iopub.status.idle":"2023-10-12T17:23:07.889578Z","shell.execute_reply.started":"2023-10-12T17:23:07.603081Z","shell.execute_reply":"2023-10-12T17:23:07.888687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nX = train_dataframe['id_code']\ny = train_dataframe['diagnosis']\nX_train, X_valid, y_train, y_valid = train_test_split(X, y, test_size=0.1, random_state=42, shuffle=True)\nX_train, X_test, y_train, y_test = train_test_split(X_train, y_train, test_size=0.1, random_state=42, shuffle=True)\ntrain_df = pd.DataFrame({'id_code': X_train, 'diagnosis': y_train}).reset_index(drop=True)\nvalid_df = pd.DataFrame({'id_code': X_valid, 'diagnosis': y_valid}).reset_index(drop=True)\ntest_df = pd.DataFrame({'id_code': X_test, 'diagnosis': y_test}).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:07.893406Z","iopub.execute_input":"2023-10-12T17:23:07.894242Z","iopub.status.idle":"2023-10-12T17:23:07.914627Z","shell.execute_reply.started":"2023-10-12T17:23:07.894209Z","shell.execute_reply":"2023-10-12T17:23:07.913633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_datagen = tf.keras.preprocessing.image.ImageDataGenerator(rescale=1 / 255.0,\n                                                               rotation_range = 10, \n                                                               zoom_range = 0.30, \n                                                               shear_range = 0.30,\n                                                               fill_mode = \"nearest\"\n                                                               )\n\ntest_datagen = tf.keras.preprocessing.image.ImageDataGenerator(rescale=1 / 255.0)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:07.918187Z","iopub.execute_input":"2023-10-12T17:23:07.918805Z","iopub.status.idle":"2023-10-12T17:23:07.92687Z","shell.execute_reply.started":"2023-10-12T17:23:07.91877Z","shell.execute_reply":"2023-10-12T17:23:07.925985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images_directory = \"/kaggle/input/aptos2019-blindness-detection/train_images\"\ntrain_generator = train_datagen.flow_from_dataframe(dataframe = train_df,\n                                                    directory = images_directory,\n                                                    x_col = 'id_code',\n                                                    y_col = 'diagnosis',\n                                                    target_size=(224, 224),\n                                                    color_mode='rgb',\n                                                    class_mode='categorical',\n                                                    batch_size=16)\nvalid_generator = test_datagen.flow_from_dataframe(dataframe = valid_df,\n                                                    directory = images_directory,\n                                                    x_col = 'id_code',\n                                                    y_col = 'diagnosis',\n                                                    target_size=(224, 224),\n                                                    color_mode='rgb',\n                                                    class_mode='categorical',\n                                                    batch_size=16)\ntest_generator = test_datagen.flow_from_dataframe(dataframe = test_df,\n                                                    directory = images_directory,\n                                                    x_col = 'id_code',\n                                                    y_col = 'diagnosis',\n                                                    target_size=(224, 224),\n                                                    color_mode='rgb',\n                                                    class_mode='categorical',\n                                                    batch_size=16)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:07.928232Z","iopub.execute_input":"2023-10-12T17:23:07.928772Z","iopub.status.idle":"2023-10-12T17:23:12.436275Z","shell.execute_reply.started":"2023-10-12T17:23:07.928742Z","shell.execute_reply":"2023-10-12T17:23:12.435327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_teacher_model(model_name, num_classes=5):  # Specify the correct number of classes (5 in this case)\n    pre_trained_model = tf.keras.applications.__dict__[model_name](weights='imagenet', include_top=False)\n    \n    # Make the final layer compatible with the number of classes\n    x = tf.keras.layers.GlobalAveragePooling2D()(pre_trained_model.output)\n    x = tf.keras.layers.Dense(num_classes, activation='softmax')(x)\n\n    teacher_model = tf.keras.models.Model(inputs=pre_trained_model.input, outputs=x)\n    \n    return teacher_model\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:12.437565Z","iopub.execute_input":"2023-10-12T17:23:12.43848Z","iopub.status.idle":"2023-10-12T17:23:12.444438Z","shell.execute_reply.started":"2023-10-12T17:23:12.438447Z","shell.execute_reply":"2023-10-12T17:23:12.443294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense\n\n# Define the teacher models\n# Create the teacher models with num_classes=5\nteacher_models = {\n    'VGG16': create_teacher_model('VGG16', num_classes=5),\n    'VGG19': create_teacher_model('VGG19', num_classes=5),\n    'Xception': create_teacher_model('Xception', num_classes=5),\n    'InceptionV3': create_teacher_model('InceptionV3', num_classes=5),\n    'ResNet50': create_teacher_model('ResNet50', num_classes=5)\n}\n\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:12.445951Z","iopub.execute_input":"2023-10-12T17:23:12.446598Z","iopub.status.idle":"2023-10-12T17:23:30.669104Z","shell.execute_reply.started":"2023-10-12T17:23:12.446566Z","shell.execute_reply":"2023-10-12T17:23:30.668157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the student model (KNET)\ndef create_student_model(num_classes=5):  # Specify the correct number of classes (5 in this case)\n    model = keras.Sequential([\n        Conv2D(64, (3, 3), activation='relu', input_shape=(224, 224, 3)),\n        MaxPooling2D((2, 2)),\n        Conv2D(128, (3, 3), activation='relu'),\n        MaxPooling2D((2, 2)),\n        Conv2D(256, (3, 3), activation='relu'),\n        MaxPooling2D((2, 2)),\n        Flatten(),\n        Dense(256, activation='relu'),\n        Dense(num_classes, activation='softmax')\n    ])\n    return model\n\nstudent_model = create_student_model(num_classes=5)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:30.670642Z","iopub.execute_input":"2023-10-12T17:23:30.671201Z","iopub.status.idle":"2023-10-12T17:23:30.74077Z","shell.execute_reply.started":"2023-10-12T17:23:30.671167Z","shell.execute_reply":"2023-10-12T17:23:30.739844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Define the Distiller class\nclass Distiller(keras.Model):\n    def __init__(self, student, teachers):\n        super().__init__()\n        self.teachers = teachers\n        self.student = student\n\n    def compile(self, optimizer, metrics, student_loss_fn, distillation_loss_fn, alpha=0.1, temperature=3):\n        super().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        self.temperature = temperature\n\n    def train_step(self, data):\n        x, y = data\n        teacher_predictions = [teacher(x, training=False) for teacher in self.teachers]\n\n        with tf.GradientTape() as tape:\n            student_predictions = self.student(x, training=True)\n            student_loss = self.student_loss_fn(y, student_predictions)\n            distillation_loss = 0\n\n            for i, teacher_pred in enumerate(teacher_predictions):\n                distillation_loss += self.distillation_loss_fn(\n                    tf.nn.softmax(teacher_pred / self.temperature, axis=1),\n                    tf.nn.softmax(student_predictions / self.temperature, axis=1)\n                )\n\n            distillation_loss /= len(self.teachers)\n            loss = self.alpha * student_loss + (1 - self.alpha) * distillation_loss\n\n        trainable_vars = self.student.trainable_variables\n        gradients = tape.gradient(loss, trainable_vars)\n        self.optimizer.apply_gradients(zip(gradients, trainable_vars))\n\n        self.compiled_metrics.update_state(y, student_predictions)\n        results = {m.name: m.result() for m in self.metrics}\n        results.update({\"student_loss\": student_loss, \"distillation_loss\": distillation_loss})\n        return results\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:30.742055Z","iopub.execute_input":"2023-10-12T17:23:30.742373Z","iopub.status.idle":"2023-10-12T17:23:30.752041Z","shell.execute_reply.started":"2023-10-12T17:23:30.742339Z","shell.execute_reply":"2023-10-12T17:23:30.750904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n    def test_step(self, data):\n        x, y = data\n        student_predictions = self.student(x, training=False)\n        student_loss = self.student_loss_fn(y, student_predictions)\n        self.compiled_metrics.update_state(y, student_predictions)\n        results = {m.name: m.result() for m in self.metrics}\n        results.update({\"student_loss\": student_loss})\n        return results\n\n# Create instances of teacher and student models\n\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:30.753374Z","iopub.execute_input":"2023-10-12T17:23:30.753915Z","iopub.status.idle":"2023-10-12T17:23:30.76762Z","shell.execute_reply.started":"2023-10-12T17:23:30.753885Z","shell.execute_reply":"2023-10-12T17:23:30.766508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"teachers = [teacher_models[name] for name in ['VGG16', 'VGG19', 'Xception', 'InceptionV3', 'ResNet50']]\nstudent = create_student_model()\n\n# Create a distiller\ndistiller = Distiller(student=student, teachers=teachers)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:30.769041Z","iopub.execute_input":"2023-10-12T17:23:30.769575Z","iopub.status.idle":"2023-10-12T17:23:30.844436Z","shell.execute_reply.started":"2023-10-12T17:23:30.769542Z","shell.execute_reply":"2023-10-12T17:23:30.843566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compile the distiller\ndistiller.compile(\n    optimizer=keras.optimizers.Adam(),\n    metrics=[keras.metrics.CategoricalAccuracy()],\n    student_loss_fn=keras.losses.CategoricalCrossentropy(),\n    distillation_loss_fn=keras.losses.KLDivergence(),\n    alpha=0.1,\n    temperature=20\n)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:30.845696Z","iopub.execute_input":"2023-10-12T17:23:30.846051Z","iopub.status.idle":"2023-10-12T17:23:30.867335Z","shell.execute_reply.started":"2023-10-12T17:23:30.846022Z","shell.execute_reply":"2023-10-12T17:23:30.866396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DistillerModel(tf.keras.Model):\n    def __init__(self, distiller):\n        super().__init__()\n        self.distiller = distiller\n\n    def call(self, inputs, training=False):\n        # Forward pass through the distiller\n        return self.distiller.student(inputs, training=training)\n\n# Create an instance of the distiller model\ndistiller_model = DistillerModel(distiller)\n\n# Compile the distiller model\ndistiller_model.compile(\n    optimizer=keras.optimizers.Adam(),\n    metrics=[keras.metrics.CategoricalAccuracy()],\n    loss=keras.losses.CategoricalCrossentropy()\n)\n\n# Train the distiller model\ndistiller_model.fit(train_generator,\n                    validation_data=valid_generator,\n                    steps_per_epoch=train_generator.n // train_generator.batch_size,\n                    validation_steps=valid_generator.n // valid_generator.batch_size,\n                    epochs=50)\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T17:23:30.871468Z","iopub.execute_input":"2023-10-12T17:23:30.871732Z","iopub.status.idle":"2023-10-12T22:33:10.512109Z","shell.execute_reply.started":"2023-10-12T17:23:30.871709Z","shell.execute_reply":"2023-10-12T22:33:10.511081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# class Distiller(keras.Model):\n#     def __init__(self, student, teacher):\n#         super().__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#         temperature=3,\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#             temperature: Temperature for softening probability distributions.\n#                 Larger temperature gives softer distributions.\n#         \"\"\"\n#         super().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#         self.temperature = temperature\n\n#     def train_step(self, data):\n#         # Unpack data\n#         x, y = data\n\n#         # Forward pass of teacher\n#         teacher_predictions = self.teacher(x, training=False)\n\n#         with tf.GradientTape() as tape:\n#             # Forward pass of student\n#             student_predictions = self.student(x, training=True)\n\n#             # Compute losses\n#             student_loss = self.student_loss_fn(y, student_predictions)\n\n#             # Compute scaled distillation loss from https://arxiv.org/abs/1503.02531\n#             # The magnitudes of the gradients produced by the soft targets scale\n#             # as 1/T^2, multiply them by T^2 when using both hard and soft targets.\n#             distillation_loss = (\n#                 self.distillation_loss_fn(\n#                     tf.nn.softmax(teacher_predictions / self.temperature, axis=1),\n#                     tf.nn.softmax(student_predictions / self.temperature, axis=1),\n#                 )\n#                 * self.temperature**2\n#             )\n\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#     def test_step(self, data):\n#         # Unpack the data\n#         x, y = data\n\n#         # Compute predictions\n#         y_prediction = self.student(x, training=False)\n\n#         # Calculate the loss\n#         student_loss = self.student_loss_fn(y, y_prediction)\n\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","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:33:10.51338Z","iopub.execute_input":"2023-10-12T22:33:10.513702Z","iopub.status.idle":"2023-10-12T22:33:10.520982Z","shell.execute_reply.started":"2023-10-12T22:33:10.513656Z","shell.execute_reply":"2023-10-12T22:33:10.519994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def create_model(model_name):\n#     # Download the pre-trained model\n#     pre_trained_model = tf.keras.applications.__dict__[model_name](weights='imagenet', include_top=False, pooling = 'avg')\n    \n#     pre_trained_model.trainable = True\n\n#     model_output = pre_trained_model.output\n    \n#     # Add a dense layer with 10 units\n#     dense_layer = tf.keras.layers.Dense(5, activation='softmax')(model_output)\n    \n#     # Create a new model\n#     model = tf.keras.models.Model(inputs=pre_trained_model.input, outputs=dense_layer)\n    \n#     return model","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:33:10.522263Z","iopub.execute_input":"2023-10-12T22:33:10.522871Z","iopub.status.idle":"2023-10-12T22:33:10.552005Z","shell.execute_reply.started":"2023-10-12T22:33:10.52284Z","shell.execute_reply":"2023-10-12T22:33:10.551002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Create the teacher\n# teacher = create_model('DenseNet201')\n\n# # Create the student\n# student = create_model('ResNet50')","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:33:10.553283Z","iopub.execute_input":"2023-10-12T22:33:10.554106Z","iopub.status.idle":"2023-10-12T22:33:10.566819Z","shell.execute_reply.started":"2023-10-12T22:33:10.554077Z","shell.execute_reply":"2023-10-12T22:33:10.565808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# teacher.compile(\n#     optimizer=keras.optimizers.Adam(learning_rate=1e-3),\n#     loss=keras.losses.CategoricalCrossentropy(),\n#     metrics=[keras.metrics.CategoricalAccuracy()],\n# )\n\n# # Train and evaluate teacher on data.\n# teacher_history = teacher.fit(train_generator,\n#                     validation_data = valid_generator,\n#                     steps_per_epoch = train_generator.n//train_generator.batch_size,\n#                     validation_steps = valid_generator.n//valid_generator.batch_size,epochs=50)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:33:10.568083Z","iopub.execute_input":"2023-10-12T22:33:10.56888Z","iopub.status.idle":"2023-10-12T22:33:10.578634Z","shell.execute_reply.started":"2023-10-12T22:33:10.56885Z","shell.execute_reply":"2023-10-12T22:33:10.577724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:33:10.580001Z","iopub.execute_input":"2023-10-12T22:33:10.580691Z","iopub.status.idle":"2023-10-12T22:33:10.590812Z","shell.execute_reply.started":"2023-10-12T22:33:10.580661Z","shell.execute_reply":"2023-10-12T22:33:10.589983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import classification_report, accuracy_score, precision_score, recall_score, f1_score\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\n# Create a list of teacher models\n\n\nfor model_name, teacher_model in teacher_models.items():\n    print(f'Teacher Model: {model_name}')\n    \n    # Compile the teacher model\n    teacher_model.compile(\n        optimizer=keras.optimizers.Adam(),\n        loss=keras.losses.CategoricalCrossentropy(),\n        metrics=['accuracy']\n    )\n    \n    # Evaluate the teacher on the test dataset\n    teacher_score = teacher_model.evaluate(test_generator)\n    print('Categorical Accuracy:', teacher_score[1])\n\n    # Get teacher model predictions on the test dataset\n    teacher_predictions = np.argmax(teacher_model.predict(test_generator), axis=1)\n\n    # Print a classification report\n    print('Classification Report:\\n', classification_report(test_generator.classes, teacher_predictions, target_names=['No_DR', 'Mild', 'Moderate', 'Severe', 'Proliferative_DR']))\n\n    # Calculate and print other metrics\n    teacher_accuracy = accuracy_score(test_generator.classes, teacher_predictions)\n    teacher_precision = precision_score(test_generator.classes, teacher_predictions, average='weighted')\n    teacher_recall = recall_score(test_generator.classes, teacher_predictions, average='weighted')\n    teacher_f1 = f1_score(test_generator.classes, teacher_predictions, average='weighted')\n\n    print('Accuracy:', teacher_accuracy)\n    print('Precision:', teacher_precision)\n    print('Recall:', teacher_recall)\n    print('F1 Score:', teacher_f1)\n\n    # Visualize the confusion matrix for the teacher\n    cnf_matrix = confusion_matrix(test_generator.classes, teacher_predictions)\n    plt.figure(figsize=(8, 6))\n    sns.heatmap(cnf_matrix, annot=True, fmt='g', cmap='Blues', cbar=False)\n    plt.xlabel('Predicted labels')\n    plt.ylabel('True labels')\n    plt.title(f'Teacher Model {model_name} - Confusion Matrix')\n    plt.xticks(ticks=[0.5, 1.5, 2.5, 3.5, 4.5], labels=['No_DR', 'Mild', 'Moderate', 'Severe', 'Proliferative_DR'])\n    plt.yticks(ticks=[0.5, 1.5, 2.5, 3.5, 4.5], labels=['No_DR', 'Mild', 'Moderate', 'Severe', 'Proliferative_DR'])\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:33:10.592158Z","iopub.execute_input":"2023-10-12T22:33:10.592853Z","iopub.status.idle":"2023-10-12T22:39:20.077262Z","shell.execute_reply.started":"2023-10-12T22:33:10.592824Z","shell.execute_reply":"2023-10-12T22:39:20.076408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # summarize history for accuracy\n# plt.plot(teacher_history.history['categorical_accuracy'])\n# plt.plot(teacher_history.history['val_categorical_accuracy'])\n# plt.title('model accuracy')\n# plt.ylabel('accuracy')\n# plt.xlabel('epoch')\n# plt.legend(['train', 'validation'], loc='upper left')\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:39:20.078512Z","iopub.execute_input":"2023-10-12T22:39:20.079468Z","iopub.status.idle":"2023-10-12T22:39:20.08359Z","shell.execute_reply.started":"2023-10-12T22:39:20.079434Z","shell.execute_reply":"2023-10-12T22:39:20.08247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# score = teacher.evaluate(test_generator)\n# print('Test loss:', score[0])\n# print('Test accuracy:', score[1])","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:39:20.084801Z","iopub.execute_input":"2023-10-12T22:39:20.08567Z","iopub.status.idle":"2023-10-12T22:39:20.097855Z","shell.execute_reply.started":"2023-10-12T22:39:20.08564Z","shell.execute_reply":"2023-10-12T22:39:20.096872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from sklearn.metrics import confusion_matrix\n# teacher_predictions = np.argmax(teacher.predict(test_generator),axis=1)\n# cnf_matrix = confusion_matrix(test_generator.classes, teacher_predictions)\n# ax= plt.subplot()\n# sns.heatmap(cnf_matrix, annot=True, fmt='g', ax=ax);\n\n# # labels, title and ticks\n# ax.set_xlabel('Predicted labels');\n# ax.set_ylabel('True labels'); \n# ax.set_title('Confusion Matrix');\n# ax.xaxis.set_ticklabels(['No_DR', 'Mild','Moderate','Severe','Proliferative_DR']); \n# ax.yaxis.set_ticklabels(['No_DR', 'Mild','Moderate','Severe','Proliferative_DR']);","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:39:20.099156Z","iopub.execute_input":"2023-10-12T22:39:20.100058Z","iopub.status.idle":"2023-10-12T22:39:20.109597Z","shell.execute_reply.started":"2023-10-12T22:39:20.100025Z","shell.execute_reply":"2023-10-12T22:39:20.108499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# distiller = Distiller(student=student, teacher=teacher)\n# distiller.compile(\n#     optimizer=keras.optimizers.Adam(),\n#     metrics=[keras.metrics.CategoricalAccuracy()],\n#     student_loss_fn=keras.losses.CategoricalCrossentropy(),\n#     distillation_loss_fn=keras.losses.KLDivergence(),\n#     alpha=0.1,\n#     temperature=20,\n# )","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:39:20.110739Z","iopub.execute_input":"2023-10-12T22:39:20.111618Z","iopub.status.idle":"2023-10-12T22:39:20.124094Z","shell.execute_reply.started":"2023-10-12T22:39:20.111589Z","shell.execute_reply":"2023-10-12T22:39:20.123228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Distill teacher to student\n# distiller.fit(train_generator,\n#                     validation_data = valid_generator,\n#                     steps_per_epoch = train_generator.n//train_generator.batch_size,\n#                     validation_steps = valid_generator.n//valid_generator.batch_size,epochs=50)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:39:20.125153Z","iopub.execute_input":"2023-10-12T22:39:20.125947Z","iopub.status.idle":"2023-10-12T22:39:20.134527Z","shell.execute_reply.started":"2023-10-12T22:39:20.125883Z","shell.execute_reply":"2023-10-12T22:39:20.133716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Evaluate student on test dataset\n# score = distiller.evaluate(test_generator)\n# print('Categorical_accuracy:', score[0])\n# print('student_loss:', score[1])","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:39:20.135744Z","iopub.execute_input":"2023-10-12T22:39:20.136269Z","iopub.status.idle":"2023-10-12T22:39:20.144824Z","shell.execute_reply.started":"2023-10-12T22:39:20.136242Z","shell.execute_reply":"2023-10-12T22:39:20.143968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# student_predictions = np.argmax(distiller.student.predict(test_generator),axis=1)\n# cnf_matrix = confusion_matrix(test_generator.classes, student_predictions)\n# ax= plt.subplot()\n# sns.heatmap(cnf_matrix, annot=True, fmt='g', ax=ax);\n\n# # labels, title and ticks\n# ax.set_xlabel('Predicted labels');\n# ax.set_ylabel('True labels'); \n# ax.set_title('Confusion Matrix');\n# ax.xaxis.set_ticklabels(['No_DR', 'Mild','Moderate','Severe','Proliferative_DR']); \n# ax.yaxis.set_ticklabels(['No_DR', 'Mild','Moderate','Severe','Proliferative_DR']);","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:39:20.146205Z","iopub.execute_input":"2023-10-12T22:39:20.146852Z","iopub.status.idle":"2023-10-12T22:39:20.156111Z","shell.execute_reply.started":"2023-10-12T22:39:20.146787Z","shell.execute_reply":"2023-10-12T22:39:20.155187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the student model (KNET)\ndef create_student_model(num_classes=5):\n    model = keras.Sequential([\n        Conv2D(64, (3, 3), activation='relu', input_shape=(224, 224, 3)),\n        MaxPooling2D((2, 2)),\n        Conv2D(128, (3, 3), activation='relu'),\n        MaxPooling2D((2, 2)),\n        Conv2D(256, (3, 3), activation='relu'),\n        MaxPooling2D((2, 2)),\n        Flatten(),\n        Dense(256, activation='relu'),\n        Dense(num_classes, activation='softmax')\n    ])\n    return model\n\n# Create an instance of the student model\nstudent_model = create_student_model()\n\n# Create a distiller\ndistiller = Distiller(student=student_model, teachers=list(teacher_models.values()))  # Pass the list of teacher models\n\n# Compile the distiller\ndistiller.compile(\n    optimizer=keras.optimizers.Adam(),\n    metrics=[keras.metrics.CategoricalAccuracy()],\n    student_loss_fn=keras.losses.CategoricalCrossentropy(),\n    distillation_loss_fn=keras.losses.KLDivergence(),\n    alpha=0.1,\n    temperature=20\n)\n\n# Train the distiller\ndistiller.fit(train_generator,\n              validation_data=valid_generator,\n              steps_per_epoch=train_generator.n // train_generator.batch_size,\n              validation_steps=valid_generator.n // valid_generator.batch_size,\n              epochs=50)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:39:20.157183Z","iopub.execute_input":"2023-10-12T22:39:20.157944Z","iopub.status.idle":"2023-10-12T22:44:53.597467Z","shell.execute_reply.started":"2023-10-12T22:39:20.157898Z","shell.execute_reply":"2023-10-12T22:44:53.595972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define a function to calculate and print evaluation metrics\ndef evaluate_student(model, data_generator):\n    student_predictions = model.predict(data_generator)\n    student_predictions = np.argmax(student_predictions, axis=1)\n\n    # Calculate and print a classification report\n    print('Classification Report (Student):\\n', classification_report(data_generator.classes, student_predictions, target_names=['No_DR', 'Mild', 'Moderate', 'Severe', 'Proliferative_DR']))\n\n    # Calculate and print other metrics\n    student_accuracy = accuracy_score(data_generator.classes, student_predictions)\n    student_precision = precision_score(data_generator.classes, student_predictions, average='weighted')\n    student_recall = recall_score(data_generator.classes, student_predictions, average='weighted')\n    student_f1 = f1_score(data_generator.classes, student_predictions, average='weighted')\n\n    print('Accuracy (Student):', student_accuracy)\n    print('Precision (Student):', student_precision)\n    print('Recall (Student):', student_recall)\n    print('F1 Score (Student):', student_f1)\n\n# Evaluate the student on the test dataset\nevaluate_student(student_model, test_generator)\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:44:53.598514Z","iopub.status.idle":"2023-10-12T22:44:53.599398Z","shell.execute_reply.started":"2023-10-12T22:44:53.599152Z","shell.execute_reply":"2023-10-12T22:44:53.599176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"\n\n# Evaluate the student on the test dataset\nstudent_score = evaluate_student(student_model, test_generator)\nprint('Categorical Accuracy (Student):', student_score[1])\n\n# Get student model predictions on the test dataset\nstudent_predictions = np.argmax(student_model.predict(test_generator), axis=1)\n\n# Print a classification report for the student\nprint('Classification Report (Student):\\n', classification_report(test_generator.classes, student_predictions, target_names=['No_DR', 'Mild', 'Moderate', 'Severe', 'Proliferative_DR']))\n\n# Calculate and print other metrics for the student\nstudent_accuracy = accuracy_score(test_generator.classes, student_predictions)\nstudent_precision = precision_score(test_generator.classes, student_predictions, average='weighted')\nstudent_recall = recall_score(test_generator.classes, student_predictions, average='weighted')\nstudent_f1 = f1_score(test_generator.classes, student_predictions, average='weighted')\n\nprint('Accuracy (Student):', student_accuracy)\nprint('Precision (Student):', student_precision)\nprint('Recall (Student):', student_recall)\nprint('F1 Score (Student):', student_f1)\n\n# Visualize the confusion matrix for the student\ncnf_matrix_student = confusion_matrix(test_generator.classes, student_predictions)\nplt.figure(figsize=(8, 6))\nsns.heatmap(cnf_matrix_student, annot=True, fmt='g', cmap='Blues', cbar=False)\nplt.xlabel('Predicted labels')\nplt.ylabel('True labels')\nplt.title('Student Model - Confusion Matrix')\nplt.xticks(ticks=[0.5, 1.5, 2.5, 3.5, 4.5], labels=['No_DR', 'Mild', 'Moderate', 'Severe', 'Proliferative_DR'])\nplt.yticks(ticks=[0.5, 1.5, 2.5, 3.5, 4.5], labels=['No_DR', 'Mild', 'Moderate', 'Severe', 'Proliferative_DR'])\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-10-12T22:44:53.600983Z","iopub.status.idle":"2023-10-12T22:44:53.601743Z","shell.execute_reply.started":"2023-10-12T22:44:53.601477Z","shell.execute_reply":"2023-10-12T22:44:53.601499Z"},"trusted":true},"execution_count":null,"outputs":[]}]}