{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"},{"sourceId":11243487,"sourceType":"datasetVersion","datasetId":7024989}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Visualisation Functions","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport pandas as pd\nfrom sklearn.metrics import ConfusionMatrixDisplay, accuracy_score, confusion_matrix\n\n\ndef origin_image_plot(train_dataset):\n    \"\"\"\n    Plot the original images from the dataset.\n\n    Args:\n        train_dataset: The dataset containing the images and labels.\n\n    Returns:\n        fig1: The figure containing different disease images.\n        fig2: The figure containing the healthy image.\n    \"\"\"\n    fig1, axes = plt.subplots(2, 2, figsize=(10, 10))\n    fig2, ax2 = plt.subplots(1, 1, figsize=(10, 10))\n\n    labels = [0, 1, 2, 3, 4]\n    labels_name = {\n        \"0\": \"Cassava Bacterial Blight (CBB)\",\n        \"1\": \"Cassava Brown Streak Disease (CBSD)\",\n        \"2\": \"Cassava Green Mottle (CGM)\",\n        \"3\": \"Cassava Mosaic Disease (CMD)\",\n        \"4\": \"Healthy\",\n    }\n    found_images = {}\n\n    for image, label in train_dataset:\n        label = label.numpy()\n        if label in labels and label not in found_images:\n            found_images[label] = image\n            labels.remove(label)\n        if len(labels) == 0:\n            break\n\n    # plot the healthy image\n    ax2.imshow(found_images[4])\n    ax2.set_xlabel(labels_name[\"4\"], fontsize=25)\n    # ax2.set_title(\"\")\n    ax2.set_xticks([])\n    ax2.set_yticks([])\n    ax2.spines[\"top\"].set_visible(False)\n    ax2.spines[\"bottom\"].set_visible(False)\n    ax2.spines[\"left\"].set_visible(False)\n    ax2.spines[\"right\"].set_visible(False)\n    found_images.pop(4)\n\n    # plot the rest of the images\n    for i, (label, image) in enumerate(found_images.items()):\n        row = i // 2\n        col = i % 2\n        axes[row, col].imshow(image)\n        axes[row, col].set_xlabel(labels_name[str(label)], fontsize=15)\n        # axes[row, col].set_title(\"\")\n        axes[row, col].set_xticks([])\n        axes[row, col].set_yticks([])\n        axes[row, col].spines[\"top\"].set_visible(False)\n        axes[row, col].spines[\"bottom\"].set_visible(False)\n        axes[row, col].spines[\"left\"].set_visible(False)\n        axes[row, col].spines[\"right\"].set_visible(False)\n\n    return fig1, fig2\n\n\ndef learning_curve(history):\n    \"\"\"\n    Plot the learning curve of the model.\n\n    Args:\n        history: The history object returned by the model's fit method.\n\n    Returns:\n        fig: The figure containing the learning curve.\n    \"\"\"\n    history_frame = pd.DataFrame(history.history)\n\n    fig, ax = plt.subplots(1, 1, figsize=(10, 8))\n    epochs = range(1, len(history_frame) + 1)\n\n    ax.plot(epochs, history_frame[\"loss\"], label=\"Training Loss\")\n    ax.plot(epochs, history_frame[\"val_loss\"], label=\"Validation Loss\")\n    ax.legend()\n    ax.set_xlabel(\"Epochs\")\n    ax.set_ylabel(\"Loss\")\n    ax.set_xticks(range(1, len(history_frame) + 1, 2))\n\n    return fig\n\n\ndef plot_cive_result(dataset):\n    \"\"\"\n    Plot the original image, mask image, and processed image.\n\n    Args:\n        dataset: The dataset containing the images and masks.\n\n    Returns:\n        fig: The figure containing the original image, mask image, and processed image.\n    \"\"\"\n    fig, axes = plt.subplots(1, 3, figsize=(12, 8))\n\n    for image, mask in dataset.skip(14).take(1):\n        # for image, mask in dataset.take(3):\n        axes[0].imshow(image)\n        axes[0].set_xlabel(\"Original Image\", fontsize=15)\n        axes[0].set_xticks([])\n        axes[0].set_yticks([])\n        axes[0].spines[\"top\"].set_visible(False)\n        axes[0].spines[\"bottom\"].set_visible(False)\n        axes[0].spines[\"left\"].set_visible(False)\n        axes[0].spines[\"right\"].set_visible(False)\n\n        axes[1].imshow(mask)\n        axes[1].set_xlabel(\"Mask Image\", fontsize=15)\n        axes[1].set_xticks([])\n        axes[1].set_yticks([])\n        axes[1].spines[\"top\"].set_visible(False)\n        axes[1].spines[\"bottom\"].set_visible(False)\n        axes[1].spines[\"left\"].set_visible(False)\n        axes[1].spines[\"right\"].set_visible(False)\n\n        image_pro = image * mask\n        axes[2].imshow(image_pro)\n        axes[2].set_xlabel(\"Processed Image\", fontsize=15)\n        axes[2].set_xticks([])\n        axes[2].set_yticks([])\n        axes[2].spines[\"top\"].set_visible(False)\n        axes[2].spines[\"bottom\"].set_visible(False)\n        axes[2].spines[\"left\"].set_visible(False)\n        axes[2].spines[\"right\"].set_visible(False)\n\n    return fig\n\n\ndef origin_Unet_result_plot(origin_image: list, mask_image: list, image_name: list):\n    \"\"\"\n    Plot the original image and the processed image including all the diseases.\n\n    Args:\n        origin_image: The original images.\n        mask_image: The mask images.\n        image_name: The names of the images.\n\n    Returns:\n        fig: The figure containing the original image and the processed image.\n    \"\"\"\n    fig, axes = plt.subplots(2, 4, figsize=(10, 6))\n    for i in range(4):\n        axes[0, i].imshow(origin_image[i])\n        axes[0, i].set_xlabel(image_name[i], fontsize=15)\n        axes[0, i].set_xticks([])\n        axes[0, i].set_yticks([])\n        axes[0, i].spines[\"top\"].set_visible(False)\n        axes[0, i].spines[\"bottom\"].set_visible(False)\n        axes[0, i].spines[\"left\"].set_visible(False)\n        axes[0, i].spines[\"right\"].set_visible(False)\n\n        axes[1, i].imshow(origin_image[i] * mask_image[i])\n        axes[1, i].set_xlabel(f\"Processed {image_name[i]}\", fontsize=15)\n        axes[1, i].set_xticks([])\n        axes[1, i].set_yticks([])\n        axes[1, i].spines[\"top\"].set_visible(False)\n        axes[1, i].spines[\"bottom\"].set_visible(False)\n        axes[1, i].spines[\"left\"].set_visible(False)\n        axes[1, i].spines[\"right\"].set_visible(False)\n        plt.tight_layout(pad=0.3)\n\n    return fig\n\n\ndef classfication_result(y_true, y_pred):\n    \"\"\"\n    Plot the confusion matrix and calculate the accuracy of the EfficientNetB1 model.\n\n    Args:\n        y_true: The true labels.\n        y_pred: The predicted labels.\n\n    Returns:\n        fig: The confusion matrix figure.\n        acc: The accuracy of the model.\n    \"\"\"\n    # calculate the accuracy\n    acc = accuracy_score(y_true, y_pred)\n\n    # calculate the confusion matrix\n    cm = confusion_matrix(y_true, y_pred)\n\n    # plot the confusion matrix\n    fig, ax = plt.subplots(figsize=(10, 10))\n    disp = ConfusionMatrixDisplay(confusion_matrix=cm)\n    disp.plot(ax=ax, cmap=\"Blues\", values_format=\".4g\")\n    ax.set_xlabel(\"Predicted Label\", fontsize=15)\n    ax.set_ylabel(\"True Label\", fontsize=15)\n\n    # set the x and y ticks\n    xticks = [\"CBB\", \"CBSD\", \"CGM\", \"CMD\", \"Healthy\"]\n    yticks = [\"CBB\", \"CBSD\", \"CGM\", \"CMD\", \"Healthy\"]\n    ax.set_xticks(range(len(xticks)))\n    ax.set_xticklabels(xticks, fontsize=12)\n    ax.set_yticks(range(len(yticks)))\n    ax.set_yticklabels(yticks, fontsize=12)\n\n    plt.tight_layout(pad=0.3)\n\n    return fig, acc\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:04.400414Z","iopub.execute_input":"2025-04-11T02:31:04.400626Z","iopub.status.idle":"2025-04-11T02:31:05.780749Z","shell.execute_reply.started":"2025-04-11T02:31:04.400606Z","shell.execute_reply":"2025-04-11T02:31:05.779813Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Image Preprocessing Functions","metadata":{}},{"cell_type":"code","source":"from functools import partial\n\nimport numpy as np\nimport tensorflow as tf\nfrom sklearn.model_selection import train_test_split\n\n\ndef UNet_preprocessing_pro(dataset, batch_size, gen):\n    \"\"\"\n    Preprocess the dataset including healthy and CBB images for UNet training.\n\n    Args:\n        dataset: tf.data.Dataset object.\n        batch_size: int, batch size for the dataset.\n        gen: tf.random.Generator object.\n\n    Returns:\n        unet_trainset: tf.data.Dataset object for training.\n        unet_valset: tf.data.Dataset object for validation.\n    \"\"\"\n    healthy_leaf_set = get_healthy_image(dataset)\n    unet_dataset = build_Unet_dataset(healthy_leaf_set)\n    cbb_leaf_set = get_CBB_image(dataset)\n    cbb_dataset = build_cbb_dataset(cbb_leaf_set)\n\n    # fig = vs.plot_cive_result(unet_dataset)\n    # fig.savefig(\"figures/Image after CIVE.png\")\n\n    image_list = []\n    mask_list = []\n\n    for img, msk in unet_dataset:\n        image_list.append(img)\n        mask_list.append(msk)\n\n    for img, msk in cbb_dataset:\n        image_list.append(img)\n        mask_list.append(msk)\n        if gen.uniform(()) > 0.8:\n            img_flip = tf.image.flip_left_right(img)\n            msk_flip = tf.image.flip_left_right(msk)\n            image_list.append(img_flip)\n            mask_list.append(msk_flip)\n        if gen.uniform(()) > 0.8:\n            img_flip = tf.image.flip_up_down(img)\n            msk_flip = tf.image.flip_up_down(msk)\n            image_list.append(img_flip)\n            mask_list.append(msk_flip)\n\n    image_list = np.array(image_list)\n    mask_list = np.array(mask_list)\n\n    print(image_list.shape)\n\n    x_train, x_val, y_train, y_val = train_test_split(\n        image_list, mask_list, test_size=0.2, random_state=711\n    )\n\n    unet_trainset = (\n        tf.data.Dataset.from_tensor_slices((x_train, y_train))\n        .shuffle(1000)\n        .batch(batch_size)\n    )\n    unet_valset = tf.data.Dataset.from_tensor_slices((x_val, y_val)).batch(batch_size)\n\n    return unet_trainset, unet_valset\n\n\ndef build_cbb_dataset(train_dataset):\n    \"\"\"\n    Build the dataset for CBB images.\n\n    Args:\n        train_dataset: tf.data.Dataset object.\n\n    Returns:\n        cbb_dataset: tf.data.Dataset object for CBB images.\n    \"\"\"\n    cbb_dataset = train_dataset.map(\n        create_mask_pair_cbb, num_parallel_calls=tf.data.experimental.AUTOTUNE\n    )\n    return cbb_dataset\n\n\ndef UNet_preprocessing(dataset, batch_size):\n    \"\"\"\n    Preprocess the dataset including healthy images for UNet training.\n\n    Args:\n        dataset: tf.data.Dataset object.\n        batch_size: int, batch size for the dataset.\n\n    Returns:\n        unet_trainset: tf.data.Dataset object for training.\n        unet_valset: tf.data.Dataset object for validation.\n    \"\"\"\n    healthy_leaf_set = get_healthy_image(dataset)\n    unet_dataset = build_Unet_dataset(healthy_leaf_set)\n\n    # fig = vs.plot_cive_result(unet_dataset)\n    # fig.savefig(\"figures/Image after CIVE.png\")\n\n    image = []\n    mask = []\n\n    for img, msk in unet_dataset:\n        image.append(img)\n        mask.append(msk)\n\n    image = np.array(image)\n    mask = np.array(mask)\n\n    print(image.shape)\n\n    x_train, x_val, y_train, y_val = train_test_split(\n        image, mask, test_size=0.2, random_state=711\n    )\n\n    unet_trainset = (\n        tf.data.Dataset.from_tensor_slices((x_train, y_train))\n        .shuffle(1000)\n        .batch(batch_size)\n    )\n    unet_valset = tf.data.Dataset.from_tensor_slices((x_val, y_val)).batch(batch_size)\n\n    return unet_trainset, unet_valset\n\n\ndef compute_cive(image):\n    \"\"\"\n    Compute the CIVE index for a given healthy cassava leaf image.\n\n    Args:\n        image: np.ndarray, input image.\n\n    Returns:\n        cive: np.ndarray, CIVE index.\n    \"\"\"\n    # CIVE = 0.441 * R - 0.811 * G + 0.385 * B + 18.78745\n    R = image[:, :, 0]\n    G = image[:, :, 1]\n    B = image[:, :, 2]\n    cive = 0.441 * R - 0.811 * G + 0.385 * B + 18.78745\n    return cive\n\n\ndef compute_cive_cbb(image):\n    \"\"\"\n    Compute the CIVE index for a given CBB cassava leaf image.\n\n    Args:\n        image: np.ndarray, input image.\n\n    Returns:\n        cive: np.ndarray, CIVE index.\n    \"\"\"\n    # CIVE = 10.441 * R + 0.611 * G - 1.885 * B - 48.787\n    R = image[:, :, 0]\n    G = image[:, :, 1]\n    B = image[:, :, 2]\n    cive = 10.441 * R + 0.611 * G - 1.885 * B - 48.787\n    return cive\n\n\ndef get_mask(cive, threshold=None):\n    \"\"\"\n    Create the mask based on the CIVE index.\n\n    Args:\n        cive: np.ndarray, CIVE index.\n        threshold: float, threshold value for mask creation.\n\n    Returns:\n        mask: np.ndarray, binary mask.\n    \"\"\"\n    if threshold is None:\n        # threshold = tf.reduce_mean(cive)\n        threshold = 18.7  # threshold is set to 18.7\n    mask = tf.cast(cive < threshold, tf.float32)\n    mask = tf.expand_dims(mask, axis=-1)\n    return mask\n\n\ndef get_mask_cbb(cive, threshold=None):\n    \"\"\"\n    Create the mask based on the CIVE index for CBB images.\n\n    Args:\n        cive: np.ndarray, CIVE index.\n        threshold: float, threshold value for mask creation.\n\n    Returns:\n        mask: np.ndarray, binary mask.\n    \"\"\"\n    if threshold is None:\n        threshold = tf.reduce_mean(cive)\n        # threshold = -48.63\n    mask = tf.cast(cive > threshold, tf.float32)\n    mask = tf.expand_dims(mask, axis=-1)\n    return mask\n\n\ndef create_mask_pair(image, label):\n    \"\"\"\n    Create a mask pair for the given image and label.\n\n    Args:\n        image: np.ndarray, input image.\n        label: np.ndarray, input label.\n\n    Returns:\n        image: np.ndarray, input image.\n        mask: np.ndarray, binary mask.\n    \"\"\"\n    cive = compute_cive(image)\n    mask = get_mask(cive, threshold=None)\n    return image, mask\n\n\ndef create_mask_pair_cbb(image, label):\n    \"\"\"\n    Create a mask pair for the given CBB image and label.\n\n    Args:\n        image: np.ndarray, input image.\n        label: np.ndarray, input label.\n\n    Returns:\n        image: np.ndarray, input image.\n        mask: np.ndarray, binary mask.\n    \"\"\"\n    cive = compute_cive_cbb(image)\n    mask = get_mask_cbb(cive, threshold=None)\n    return image, mask\n\n\ndef build_Unet_dataset(train_dataset):\n    \"\"\"\n    Build the dataset including images and masks for UNet training.\n\n    Args:\n        train_dataset: tf.data.Dataset object.\n\n    Returns:\n        Unet_dataset: tf.data.Dataset object for UNet training.\n    \"\"\"\n    unet_dataset = train_dataset.map(\n        create_mask_pair, num_parallel_calls=tf.data.experimental.AUTOTUNE\n    )\n    return unet_dataset\n\n\ndef which_classes(image, label, classes):\n    \"\"\"\n    Filter the dataset based on the given classes.\n\n    Args:\n        image: np.ndarray, input image.\n        label: np.ndarray, input label.\n        classes: int, class to filter.\n\n    Returns:\n        bool: True if the label matches the classes, False otherwise.\n    \"\"\"\n    return tf.equal(label, classes)\n\n\ndef get_healthy_image(train_dataset):\n    \"\"\"\n    Get the healthy images from the dataset.\n\n    Args:\n        train_dataset: tf.data.Dataset object.\n\n    Returns:\n        healthy_dataset: tf.data.Dataset object for healthy images.\n    \"\"\"\n    filter_fn = partial(which_classes, classes=4)\n    healthy_dataset = train_dataset.filter(filter_fn)\n    # healthy_dataset_num = count_data_items(healthy_dataset)\n    # print(f\"Number of healthy images: {healthy_dataset_num}\")\n    return healthy_dataset\n\n\ndef get_all_diseases_image(train_dataset):\n    \"\"\"\n    Get all the diseases images including CBB, CBSD, CGM, and CMD from the dataset.\n\n    Args:\n        train_dataset: tf.data.Dataset object.\n\n    Returns:\n        cbb_dataset: tf.data.Dataset object for CBB images.\n        cbsd_dataset: tf.data.Dataset object for CBSD images.\n        cgm_dataset: tf.data.Dataset object for CGM images.\n        cmd_dataset: tf.data.Dataset object for CMD images.\n    \"\"\"\n    cbb_dataset = get_CBB_image(train_dataset)\n    cbsd_dataset = get_CBSD_image(train_dataset)\n    cgm_dataset = get_CGM_image(train_dataset)\n    cmd_dataset = get_CMD_image(train_dataset)\n    return cbb_dataset, cbsd_dataset, cgm_dataset, cmd_dataset\n\n\ndef get_CBB_image(train_dataset):\n    \"\"\"\n    Get the CBB images from the dataset.\n\n    Args:\n        train_dataset: tf.data.Dataset object.\n\n    Returns:\n        cbb_dataset: tf.data.Dataset object for CBB images.\n    \"\"\"\n    filter_fn = partial(which_classes, classes=0)\n    cbb_dataset = train_dataset.filter(filter_fn)\n    # cbb_dataset_num = count_data_items(cbb_dataset)\n    # print(f\"Number of CBB images: {cbb_dataset_num}\")\n    return cbb_dataset\n\n\ndef get_CBSD_image(train_dataset):\n    \"\"\"\n    Get the CBSD images from the dataset.\n\n    Args:\n        train_dataset: tf.data.Dataset object.\n\n    Returns:\n        cbsd_dataset: tf.data.Dataset object for CBSD images.\n    \"\"\"\n    filter_fn = partial(which_classes, classes=1)\n    cbsd_dataset = train_dataset.filter(filter_fn)\n    # cbsd_dataset_num = count_data_items(cbsd_dataset)\n    # print(f\"Number of CBSD images: {cbsd_dataset_num}\")\n    return cbsd_dataset\n\n\ndef get_CGM_image(train_dataset):\n    \"\"\"\n    Get the CGM images from the dataset.\n\n    Args:\n        train_dataset: tf.data.Dataset object.\n\n    Returns:\n        cgm_dataset: tf.data.Dataset object for CGM images.\n    \"\"\"\n    filter_fn = partial(which_classes, classes=2)\n    cgm_dataset = train_dataset.filter(filter_fn)\n    # cgm_dataset_num = count_data_items(cgm_dataset)\n    # print(f\"Number of CGM images: {cgm_dataset_num}\")\n    return cgm_dataset\n\n\ndef get_CMD_image(train_dataset):\n    \"\"\"\n    Get the CMD images from the dataset.\n\n    Args:\n        train_dataset: tf.data.Dataset object.\n\n    Returns:\n        cmd_dataset: tf.data.Dataset object for CMD images.\n    \"\"\"\n    filter_fn = partial(which_classes, classes=3)\n    cmd_dataset = train_dataset.filter(filter_fn)\n    # cmd_dataset_num = count_data_items(cmd_dataset)\n    # print(f\"Number of CMD images: {cmd_dataset_num}\")\n    return cmd_dataset\n\n\ndef data_acquisition(gcs_path, train_file_path, image_size):\n    \"\"\"\n    Acquire the training and validation datasets for image segmentation.\n\n    Args:\n        gcs_path: str, path to the GCS bucket.\n        train_file_path: str, path to the training file.\n        image_size: list, size of the images.\n\n    Returns:\n        train_dataset: tf.data.Dataset object for training.\n        val_dataset: tf.data.Dataset object for validation.\n    \"\"\"\n    trainfile, valfile = split_data(gcs_path, train_file_path)\n    train_dataset = load_trainset(trainfile, image_size=image_size, labeled=True)\n    val_dataset = load_valset(valfile, image_size=image_size, labeled=True)\n    # # visualise the images\n    # fig1, fig2 = vs.origin_image_plot(train_dataset)\n    # fig2.savefig(\"figures/Healthy Leaf.png\")\n    # fig1.savefig(\"figures/Diseases Leaf.png\")\n    return train_dataset, val_dataset\n\n\ndef data_acquisition_classification(\n    gcs_path, train_file_path, image_size, train_ratio, val_ratio\n):\n    \"\"\"\n    Acquire the training, validation, and test datasets to build a new dataset for classification.\n\n    Args:\n        gcs_path: str, path to the GCS bucket.\n        train_file_path: str, path to the training file.\n        image_size: list, size of the images.\n        train_ratio: float, ratio of training data.\n        val_ratio: float, ratio of validation data.\n\n    Returns:\n        train_dataset: tf.data.Dataset object for training.\n        val_dataset: tf.data.Dataset object for validation.\n        test_dataset: tf.data.Dataset object for testing.\n    \"\"\"\n    trainfile, valfile, testfile = split_data_classification(\n        gcs_path, train_file_path, train_ratio, val_ratio\n    )\n    train_dataset = load_trainset(trainfile, image_size=image_size, labeled=True)\n    val_dataset = load_valset(valfile, image_size=image_size, labeled=True)\n    test_dataset = load_testset(testfile, image_size=image_size, labeled=True)\n    return train_dataset, val_dataset, test_dataset\n\n\ndef decode_raw_image(image_data, image_shape):\n    \"\"\"\n    Decode the raw image data and resize it to the specified shape.\n\n    Args:\n        image_data: bytes, raw image data.\n        image_shape: list, shape of the image.\n\n    Returns:\n        image: tf.Tensor, decoded and resized image.\n    \"\"\"\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    image = tf.image.resize(image, image_shape)\n    image = tf.cast(image, tf.float32) / 255.0\n    image = tf.reshape(image, [*image_shape, 3])\n    return image\n\n\ndef read_tfrecord(example, image_size, labeled):\n    \"\"\"\n    Read a single TFRecord example and decode the image data.\n\n    Args:\n        example: tf.train.Example, TFRecord example.\n        image_size: list, size of the images.\n        labeled: bool, whether the dataset is labeled.\n\n    Returns:\n        images: tf.Tensor, decoded and resized image.\n        labels: tf.Tensor, labels for the images (if labeled).\n    \"\"\"\n    tfrecord_format = (\n        {\n            \"image\": tf.io.FixedLenFeature([], tf.string),\n            \"target\": tf.io.FixedLenFeature([], tf.int64),\n        }\n        if labeled\n        else {\n            \"image\": tf.io.FixedLenFeature([], tf.string),\n            \"image_name\": tf.io.FixedLenFeature([], tf.string),\n        }\n    )\n    example = tf.io.parse_single_example(example, tfrecord_format)\n    images = decode_raw_image(example[\"image\"], image_shape=image_size)\n    if labeled:\n        labels = tf.cast(example[\"target\"], tf.int32)\n        return images, labels\n    image_name = example[\"image_name\"]\n    return images, image_name\n\n\ndef split_data(gcs_path, train_file_path):\n    \"\"\"\n    Split the tfrecords into training and validation use.\n\n    Args:\n        gcs_path: str, path to the GCS bucket.\n        train_file_path: str, path to the training file.\n\n    Returns:\n        train_filename: list, list of training filenames.\n        val_filename: list, list of validation filenames.\n    \"\"\"\n    train_filename, val_filename = train_test_split(\n        tf.io.gfile.glob(gcs_path + train_file_path), train_size=0.8, random_state=711\n    )\n    return train_filename, val_filename\n\n\ndef split_data_classification(gcs_path, train_file_path, train_ratio, val_ratio):\n    \"\"\"\n    Split the tfrecords into training, validation, and test use.\n\n    Args:\n        gcs_path: str, path to the GCS bucket.\n        train_file_path: str, path to the training file.\n        train_ratio: float, ratio of training data.\n        val_ratio: float, ratio of validation data.\n\n    Returns:\n        train_filename: list, list of training filenames.\n        val_filename: list, list of validation filenames.\n        test_filename: list, list of test filenames.\n    \"\"\"\n    train_filename, val_test_filename = train_test_split(\n        tf.io.gfile.glob(gcs_path + train_file_path),\n        train_size=train_ratio,\n        random_state=711,\n    )\n\n    val_test_ratio = val_ratio / (1 - train_ratio)\n    val_filename, test_filename = train_test_split(\n        val_test_filename, train_size=val_test_ratio, random_state=711\n    )\n\n    return train_filename, val_filename, test_filename\n\n\ndef load_trainset(filenames, image_size, labeled, ordered=False):\n    \"\"\"\n    Load the training dataset from TFRecord files.\n\n    Args:\n        filenames: list, list of TFRecord filenames.\n        image_size: list, size of the images.\n        labeled: bool, whether the dataset is labeled.\n        ordered: bool, whether to load the dataset in order.\n\n    Returns:\n        dataset: tf.data.Dataset object for training.\n    \"\"\"\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = True\n    dataset = tf.data.TFRecordDataset(\n        filenames,\n        num_parallel_reads=tf.data.experimental.AUTOTUNE,\n    )  # automatically interleaves reads from multiple files\n    dataset = dataset.with_options(\n        ignore_order\n    )  # uses data as soon as it streams in, rather than in its original order\n    dataset = dataset.map(\n        partial(read_tfrecord, image_size=image_size, labeled=labeled),\n        num_parallel_calls=tf.data.experimental.AUTOTUNE,\n    )\n    return dataset\n\n\ndef load_valset(filenames, image_size, labeled, ordered=False):\n    \"\"\"\n    Load the validation dataset from TFRecord files.\n\n    Args:\n        filenames: list, list of TFRecord filenames.\n        image_size: list, size of the images.\n        labeled: bool, whether the dataset is labeled.\n        ordered: bool, whether to load the dataset in order.\n\n    Returns:\n        dataset: tf.data.Dataset object for validation.\n    \"\"\"\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = True\n    dataset = tf.data.TFRecordDataset(\n        filenames,\n        num_parallel_reads=tf.data.experimental.AUTOTUNE,\n    )  # automatically interleaves reads from multiple files\n    dataset = dataset.with_options(\n        ignore_order\n    )  # uses data as soon as it streams in, rather than in its original order\n    dataset = dataset.map(\n        partial(read_tfrecord, image_size=image_size, labeled=labeled),\n        num_parallel_calls=tf.data.experimental.AUTOTUNE,\n    )\n    return dataset\n\n\ndef load_testset(filenames, image_size, labeled, ordered=False):\n    \"\"\"\n    Load the test dataset from TFRecord files.\n\n    Args:\n        filenames: list, list of TFRecord filenames.\n        image_size: list, size of the images.\n        labeled: bool, whether the dataset is labeled.\n        ordered: bool, whether to load the dataset in order.\n\n    Returns:\n        dataset: tf.data.Dataset object for testing.\n    \"\"\"\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = True\n    dataset = tf.data.TFRecordDataset(\n        filenames,\n        num_parallel_reads=tf.data.experimental.AUTOTUNE,\n    )  # automatically interleaves reads from multiple files\n    dataset = dataset.with_options(\n        ignore_order\n    )  # uses data as soon as it streams in, rather than in its original order\n    dataset = dataset.map(\n        partial(read_tfrecord, image_size=image_size, labeled=labeled),\n        num_parallel_calls=tf.data.experimental.AUTOTUNE,\n    )\n    return dataset\n\n\ndef count_data_items(trainset):\n    \"\"\"\n    Count the number of samples in the dataset.\n\n    Args:\n        trainset: tf.data.Dataset object.\n\n    Returns:\n        sample_count: int, number of samples in the dataset.\n    \"\"\"\n    sample_count = sum(1 for _ in trainset)\n    return sample_count\n\n\n# --------preprocessing for classification task--------\n\n\ndef create_classification_dataset_batch(model, dataset, batch_size=64, threshold=0.4):\n    \"\"\"\n    Create a new dataset after image segmentation.\n\n    Args:\n        model: tf.keras.Model object, trained UNet model for segmentation.\n        dataset: tf.data.Dataset object, input dataset.\n        batch_size: int, batch size for the dataset.\n        threshold: float, threshold value for mask creation.\n\n    Returns:\n        dataset: tf.data.Dataset object, new dataset after segmentation.\n    \"\"\"\n    images, labels = [], []\n\n    for image, label in dataset:\n        images.append(image)\n        labels.append(label)\n\n    images = tf.stack(images)\n    labels = tf.convert_to_tensor(labels)\n\n    # predict the mask\n    masks = model.predict(images, batch_size=batch_size)\n    masks = masks > threshold\n    masks = tf.cast(masks, tf.float32)\n\n    # segment the image\n    segmented_images = images * masks\n\n    return tf.data.Dataset.from_tensor_slices((segmented_images, labels))\n\n\ndef _bytes_feature(value):\n    return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))\n\n\ndef _int64_feature(value):\n    return tf.train.Feature(int64_list=tf.train.Int64List(value=[value]))\n\n\ndef create_example(image, label, image_type=tf.float32):\n    \"\"\"\n    Create a TFRecord example from the image and label.\n\n    Args:\n        image: tf.Tensor, input image.\n        label: tf.Tensor, input label.\n        image_type: tf.DType, type of the image.\n\n    Returns:\n        example: tf.train.Example, TFRecord example.\n    \"\"\"\n    # image_raw = tf.io.serialize_tensor(tf.cast(image, image_type)).numpy()\n    image_uint8 = tf.image.convert_image_dtype(image, tf.uint8)\n    image_raw = tf.image.encode_jpeg(image_uint8).numpy()\n    label = int(label.numpy())\n\n    feature = {\n        \"image\": _bytes_feature(image_raw),\n        \"target\": _int64_feature(label),\n    }\n    return tf.train.Example(features=tf.train.Features(feature=feature))\n\n\ndef save_tfrecord(dataset, filename):\n    \"\"\"\n    Save the dataset to a TFRecord file.\n\n    Args:\n        dataset: tf.data.Dataset object, input dataset.\n        filename: str, path to the output TFRecord file.\n    \"\"\"\n    with tf.io.TFRecordWriter(filename) as writer:\n        for image, label in dataset:\n            example = create_example(image, label)\n            writer.write(example.SerializeToString())\n    print(f\"{filename} saved\")\n\n\ndef count_data_items_from_tfrecord(filenames):\n    \"\"\"\n    Count the number of samples in the TFRecord files.\n\n    Args:\n        filenames: list, list of TFRecord filenames.\n\n    Returns:\n        count: int, number of samples in the TFRecord files.\n    \"\"\"\n    count = 0\n    for f in filenames:\n        for _ in tf.data.TFRecordDataset(f):\n            count += 1\n    return count\n\n\ndef preprocess_image_for_classification(\n    gcs_path,\n    train_file_path,\n    val_file_path,\n    test_file_path,\n    image_size,\n    batch_size,\n    autotune,\n):\n    \"\"\"\n    Preprocess the dataset for cassava leaf image classification based on new dataset.\n\n    Args:\n        gcs_path: str, path to the GCS bucket.\n        train_file_path: str, path to the training file.\n        val_file_path: str, path to the validation file.\n        test_file_path: str, path to the test file.\n        image_size: list, size of the images.\n        batch_size: int, batch size for the dataset.\n        autotune: tf.data.experimental.AUTOTUNE object.\n\n    Returns:\n        trainset: tf.data.Dataset object for training.\n        valset: tf.data.Dataset object for validation.\n        testset: tf.data.Dataset object for testing.\n        num_trainset: int, number of samples in the training dataset.\n        num_valset: int, number of samples in the validation dataset.\n    \"\"\"\n    trainset_path = gcs_path + train_file_path\n    trainset_files = tf.io.gfile.glob(trainset_path)\n    num_trainset = count_data_items_from_tfrecord(trainset_files)\n    valset_path = gcs_path + val_file_path\n    valset_files = tf.io.gfile.glob(valset_path)\n    num_valset = count_data_items_from_tfrecord(valset_files)\n    testset_path = gcs_path + test_file_path\n\n    trainset = load_trainset(trainset_files, image_size=image_size, labeled=True)\n    # trainset = trainset.repeat()\n    trainset = trainset.shuffle(2048, seed=711)\n    trainset = trainset.repeat()\n    trainset = trainset.batch(batch_size)\n    trainset = trainset.prefetch(autotune)\n    valset = load_valset(valset_path, image_size=image_size, labeled=True)\n    valset = valset.repeat().batch(batch_size).prefetch(autotune)\n    testset = load_testset(testset_path, image_size=image_size, labeled=True)\n    testset = testset.batch(batch_size).prefetch(autotune)\n\n    return trainset, valset, testset, num_trainset, num_valset\n\n\ndef preprocess_image_for_compare(\n    gcs_path,\n    train_file_path,\n    image_size,\n    train_ratio,\n    val_ratio,\n    batch_size,\n    autotune,\n):\n    \"\"\"\n    Preprocess the dataset for cassava leaf image classification based on original dataset.\n\n    Args:\n        gcs_path: str, path to the GCS bucket.\n        train_file_path: str, path to the training file.\n        image_size: list, size of the images.\n        train_ratio: float, ratio of training data.\n        val_ratio: float, ratio of validation data.\n        batch_size: int, batch size for the dataset.\n        autotune: tf.data.experimental.AUTOTUNE object.\n\n    Returns:\n        trainset: tf.data.Dataset object for training.\n        valset: tf.data.Dataset object for validation.\n        testset: tf.data.Dataset object for testing.\n        num_trainset: int, number of samples in the training dataset.\n        num_valset: int, number of samples in the validation dataset.\n    \"\"\"\n    trainfile, valfile, testfile = split_data_classification(\n        gcs_path, train_file_path, train_ratio, val_ratio\n    )\n    num_trainset = count_data_items_from_tfrecord(trainfile)\n    num_valset = count_data_items_from_tfrecord(valfile)\n\n    trainset = load_trainset(trainfile, image_size=image_size, labeled=True)\n    trainset = trainset.repeat()\n    trainset = trainset.shuffle(2048, seed=711)\n    trainset = trainset.batch(batch_size)\n    trainset = trainset.prefetch(autotune)\n    valset = load_valset(valfile, image_size=image_size, labeled=True)\n    valset = valset.repeat().batch(batch_size).prefetch(autotune)\n    testset = load_testset(testfile, image_size=image_size, labeled=True)\n    testset = testset.batch(batch_size).prefetch(autotune)\n\n    return trainset, valset, testset, num_trainset, num_valset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:05.782139Z","iopub.execute_input":"2025-04-11T02:31:05.782621Z","iopub.status.idle":"2025-04-11T02:31:17.245655Z","shell.execute_reply.started":"2025-04-11T02:31:05.782585Z","shell.execute_reply":"2025-04-11T02:31:17.244999Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Building and Training Functions","metadata":{}},{"cell_type":"code","source":"import os\nimport random\n\nimport numpy as np\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import initializers, layers, models\nfrom tensorflow.keras.applications import EfficientNetB1\n\n\ndef set_seed(seed=711):\n    \"\"\"\n    Set the seed for reproducibility.\n\n    Args:\n        seed (int): The seed value to set.\n\n    Returns:\n        tf.random.Generator: A TensorFlow random generator with the specified seed.\n    \"\"\"\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n    os.environ[\"TF_DETERMINISTIC_OPS\"] = \"1\"\n\n    return tf.random.Generator.from_seed(seed)\n\n\ndef Unet_Arch(dropout_num, input_shape=(256, 256, 3)):\n    \"\"\"\n    Define the U-Net architecture.\n\n    Args:\n        dropout_num (float): Dropout rate.\n        input_shape (tuple): Shape of the input image.\n\n    Returns:\n        tf.keras.Model: U-Net model.\n    \"\"\"\n    inputs = keras.Input(shape=input_shape)\n    conv_init = initializers.HeNormal(seed=711)\n    bias_init = initializers.Zeros()\n\n    # contracting path\n    c1 = layers.Conv2D(\n        16,\n        (3, 3),\n        activation=\"relu\",\n        kernel_initializer=conv_init,\n        bias_initializer=bias_init,\n        padding=\"same\",\n    )(\n        inputs\n    )  # 256 x 256 x 16\n    c2 = layers.Conv2D(\n        16,\n        (3, 3),\n        activation=\"relu\",\n        kernel_initializer=conv_init,\n        bias_initializer=bias_init,\n        padding=\"same\",\n    )(\n        c1\n    )  # 256 x 256 x 16\n    p1 = layers.MaxPooling2D((2, 2))(c2)  # 128 x 128 x 16\n    p1 = layers.Dropout(dropout_num)(p1)\n\n    c3 = layers.Conv2D(\n        32,\n        (3, 3),\n        activation=\"relu\",\n        kernel_initializer=conv_init,\n        bias_initializer=bias_init,\n        padding=\"same\",\n    )(\n        p1\n    )  # 128 x 128 x 32\n    c4 = layers.Conv2D(\n        32,\n        (3, 3),\n        activation=\"relu\",\n        kernel_initializer=conv_init,\n        bias_initializer=bias_init,\n        padding=\"same\",\n    )(\n        c3\n    )  # 128 x 128 x 32\n    p2 = layers.MaxPooling2D((2, 2))(c4)  # 64 x 64 x 32\n    p2 = layers.Dropout(dropout_num)(p2)\n\n    # bottleneck\n    c5 = layers.Conv2D(\n        64,\n        (3, 3),\n        activation=\"relu\",\n        kernel_initializer=conv_init,\n        bias_initializer=bias_init,\n        padding=\"same\",\n    )(\n        p2\n    )  # 64 x 64 x 64\n    c6 = layers.Conv2D(\n        64,\n        (3, 3),\n        activation=\"relu\",\n        kernel_initializer=conv_init,\n        bias_initializer=bias_init,\n        padding=\"same\",\n    )(\n        c5\n    )  # 64 x 64 x 64\n\n    # expansive path\n    u1 = layers.UpSampling2D((2, 2))(c6)  # 128 x 128 x 64\n    u1 = layers.Concatenate()([u1, c4])  # 128 x 128 x 96\n    u1 = layers.Dropout(dropout_num)(u1)\n    c7 = layers.Conv2D(\n        32,\n        (3, 3),\n        activation=\"relu\",\n        kernel_initializer=conv_init,\n        bias_initializer=bias_init,\n        padding=\"same\",\n    )(\n        u1\n    )  # 128 x 128 x 32\n    c8 = layers.Conv2D(\n        32,\n        (3, 3),\n        activation=\"relu\",\n        kernel_initializer=conv_init,\n        bias_initializer=bias_init,\n        padding=\"same\",\n    )(\n        c7\n    )  # 128 x 128 x 32\n\n    u2 = layers.UpSampling2D((2, 2))(c8)  # 256 x 256 x 32\n    u2 = layers.Concatenate()([u2, c2])  # 256 x 256 x 48\n    u2 = layers.Dropout(dropout_num)(u2)\n    c9 = layers.Conv2D(\n        16,\n        (3, 3),\n        activation=\"relu\",\n        kernel_initializer=conv_init,\n        bias_initializer=bias_init,\n        padding=\"same\",\n    )(\n        u2\n    )  # 256 x 256 x 16\n    c10 = layers.Conv2D(\n        16,\n        (3, 3),\n        activation=\"relu\",\n        kernel_initializer=conv_init,\n        bias_initializer=bias_init,\n        padding=\"same\",\n    )(\n        c9\n    )  # 256 x 256 x 16\n\n    outputs = layers.Conv2D(1, (1, 1), activation=\"sigmoid\")(c10)  # 256 x 256 x 1\n\n    model = models.Model(inputs=[inputs], outputs=[outputs])\n    return model\n\n\ndef Unet_train(model, train, val, epochs=10, learning_rate=0.001):\n    \"\"\"\n    Build the U-Net model.\n\n    Args:\n        model (tf.keras.Model): U-Net model.\n        train (tf.data.Dataset): Training dataset.\n        val (tf.data.Dataset): Validation dataset.\n        epochs (int): Number of epochs to train.\n        learning_rate (float): Learning rate.\n\n    Returns:\n        tf.keras.Model: Trained U-Net model.\n        matplotlib.figure.Figure: Learning curve figure.\n    \"\"\"\n    model.compile(\n        optimizer=tf.keras.optimizers.Adam(learning_rate=learning_rate),\n        loss=tf.keras.losses.BinaryCrossentropy(),\n        metrics=[tf.keras.metrics.BinaryAccuracy()],\n    )\n    history = model.fit(\n        train,\n        validation_data=val,\n        epochs=epochs,\n    )\n\n    fig = learning_curve(history)\n\n    return model, fig\n\n\ndef Unet_test(model, dataset):\n    \"\"\"\n    Test the U-Net model.\n\n    Args:\n        model (tf.keras.Model): U-Net model.\n        dataset (tf.data.Dataset): Dataset to test.\n\n    Returns:\n        matplotlib.figure.Figure: Figure showing the original image, mask, and processed image.\n    \"\"\"\n    cbb_dataset, cbsd_dataset, cgm_dataset, cmd_dataset = get_all_diseases_image(\n        dataset\n    )\n    # get the test sample\n    for image, _ in cbb_dataset.skip(56).take(1):  # skip(56).take(1)  11 71 116\n        cbb_image_test = image\n        cbb_image_test_batch = tf.expand_dims(cbb_image_test, axis=0)\n\n    for image, _ in cbsd_dataset.skip(55).take(1):\n        cbsd_image_test = image\n        cbsd_image_test_batch = tf.expand_dims(cbsd_image_test, axis=0)\n\n    for image, _ in cgm_dataset.skip(45).take(1):  # 45\n        cgm_image_test = image\n        cgm_image_test_batch = tf.expand_dims(cgm_image_test, axis=0)\n\n    for image, _ in cmd_dataset.take(1):\n        cmd_image_test = image\n        cmd_image_test_batch = tf.expand_dims(cmd_image_test, axis=0)\n\n    # predict the test sample\n    cbb_mask_test = model.predict(cbb_image_test_batch)\n    cbb_mask_test = tf.squeeze(cbb_mask_test, axis=0)\n    cbb_mask_test = (cbb_mask_test > 0.4).numpy().astype(\"uint8\")\n    cbsd_mask_test = model.predict(cbsd_image_test_batch)\n    cbsd_mask_test = tf.squeeze(cbsd_mask_test, axis=0)\n    cbsd_mask_test = (cbsd_mask_test > 0.4).numpy().astype(\"uint8\")\n    cgm_mask_test = model.predict(cgm_image_test_batch)\n    cgm_mask_test = tf.squeeze(cgm_mask_test, axis=0)\n    cgm_mask_test = (cgm_mask_test > 0.4).numpy().astype(\"uint8\")\n    cmd_mask_test = model.predict(cmd_image_test_batch)\n    cmd_mask_test = tf.squeeze(cmd_mask_test, axis=0)\n    cmd_mask_test = (cmd_mask_test > 0.4).numpy().astype(\"uint8\")\n\n    # plot the test sample\n    fig = origin_Unet_result_plot(\n        [cbb_image_test, cbsd_image_test, cgm_image_test, cmd_image_test],\n        [cbb_mask_test, cbsd_mask_test, cgm_mask_test, cmd_mask_test],\n        [\"CBB\", \"CBSD\", \"CGM\", \"CMD\"],\n    )\n\n    return fig\n\n\ndef EfficientNetB1_Arch(\n    dropout_ratio, layers_freezed, input_shape=(224, 224, 3), classes=5\n):\n    \"\"\"\n    Build a EfficientNetB1 model based on transfer learning.\n\n    Args:\n        dropout_ratio (float): Dropout rate.\n        layers_freezed (int): Number of layers to freeze.\n        input_shape (tuple): Shape of the input image.\n        classes (int): Number of classes.\n\n    Returns:\n        tf.keras.Model: EfficientNetB1 model.\n    \"\"\"\n    base_model = EfficientNetB1(\n        include_top=False, weights=\"imagenet\", input_shape=input_shape\n    )\n    base_model.trainable = True\n\n    # Freeze some layers\n    # cheak the number of layers in the base model\n    print(len(base_model.layers))\n    # Freeze the first 50 layers\n    for layer in base_model.layers[:layers_freezed]:\n        layer.trainable = False\n\n    model = models.Sequential(\n        [\n            base_model,\n            layers.GlobalAveragePooling2D(),\n            layers.Dense(128, activation=\"relu\"),\n            layers.Dropout(dropout_ratio),\n            layers.Dense(classes, activation=\"softmax\"),\n        ]\n    )\n\n    return model\n\n\ndef Effi_B1_train(\n    model,\n    trainset,\n    valset,\n    epochs,\n    learning_scheduler,\n    num_samples,\n    num_valset,\n    batch_size,\n):\n    \"\"\"\n    Train the EfficientNetB1 model.\n\n    Args:\n        model (tf.keras.Model): EfficientNetB1 model.\n        trainset (tf.data.Dataset): Training dataset.\n        valset (tf.data.Dataset): Validation dataset.\n        epochs (int): Number of epochs to train.\n        learning_scheduler: Learning rate schedule.\n        num_samples (int): Number of samples in the training set.\n        num_valset (int): Number of samples in the validation set.\n        batch_size (int): Batch size.\n\n    Returns:\n        tf.keras.Model: Trained EfficientNetB1 model.\n        matplotlib.figure.Figure: Learning curve figure.\n        tf.keras.callbacks.History: Training history.\n    \"\"\"\n    model.compile(\n        optimizer=tf.keras.optimizers.Adam(\n            learning_rate=learning_scheduler, epsilon=0.001\n        ),\n        loss=\"sparse_categorical_crossentropy\",\n        metrics=[\"sparse_categorical_accuracy\"],\n    )\n\n    history = model.fit(\n        trainset,\n        validation_data=valset,\n        epochs=epochs,\n        steps_per_epoch=num_samples // batch_size,\n        validation_steps=num_valset // batch_size,\n    )\n\n    fig = learning_curve(history)\n\n    return model, fig, history\n\n\ndef Effi_B1_test(model, testset):\n    \"\"\"\n    Test the EfficientNetB1 model.\n\n    Args:\n        model (tf.keras.Model): EfficientNetB1 model.\n        testset (tf.data.Dataset): Test dataset.\n\n    Returns:\n        matplotlib.figure.Figure: Figure of confusion matrix.\n        float: Classification accuracy.\n    \"\"\"\n    y_true = []\n    y_pred = []\n\n    # get the image and label from the testset\n    for images, labels in testset:\n        preds = model.predict(images)\n        preds = np.argmax(preds, axis=1)\n        y_true.extend(labels.numpy())\n        y_pred.extend(preds)\n\n    # convert to numpy array\n    y_true = np.array(y_true)\n    y_pred = np.array(y_pred)\n\n    # get the classification result\n    fig, acc = classfication_result(y_true, y_pred)\n\n    return fig, acc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:17.246858Z","iopub.execute_input":"2025-04-11T02:31:17.247394Z","iopub.status.idle":"2025-04-11T02:31:17.328014Z","shell.execute_reply.started":"2025-04-11T02:31:17.247369Z","shell.execute_reply":"2025-04-11T02:31:17.327144Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"set_seed()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:17.329590Z","iopub.execute_input":"2025-04-11T02:31:17.329931Z","iopub.status.idle":"2025-04-11T02:31:18.123678Z","shell.execute_reply.started":"2025-04-11T02:31:17.329901Z","shell.execute_reply":"2025-04-11T02:31:18.123008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nimport os, re\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport importlib\nfrom tensorflow.keras import optimizers","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:18.124433Z","iopub.execute_input":"2025-04-11T02:31:18.124703Z","iopub.status.idle":"2025-04-11T02:31:18.128135Z","shell.execute_reply.started":"2025-04-11T02:31:18.124668Z","shell.execute_reply":"2025-04-11T02:31:18.127490Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Set parameters","metadata":{}},{"cell_type":"code","source":"autotune = tf.data.experimental.AUTOTUNE\ngcs_path = \"/kaggle/input\"\ntrainset_path = \"/cassava-leaf-disease-classification/train_tfrecords/ld_train*.tfrec\"\n# parameters for image segmentation\nbatch_size = 256\nbatch_size_unet = 64\nimage_size = [256, 256]\nclasses = [\"0\", \"1\", \"2\", \"3\", \"4\"]\nepochs = 18\nepochs_pro = 26\nUnet_lr = 0.0005\nUnet_lr_pro = 0.0002\ndropout_num_pro = 0.12\ntrain_ratio_cnn = 0.8\nval_ratio_cnn = 0.1\nlearning_scheduler = optimizers.schedules.ExponentialDecay(\n    initial_learning_rate=0.001, decay_steps=10000, decay_rate=0.9\n)\n  \n# data path for classfication\ntrainset_path_enet = \"/new-dataset/New_Dataset/trainset_*.tfrec\"\nvalset_path_enet = \"/new-dataset/New_Dataset/validation_set.tfrec\"\ntestset_path_enet = \"/new-dataset/New_Dataset/test_set.tfrec\"\n\nimage_size_enet = [240, 240]\ndropout_ratio_ori = 0.7\ndropout_ratio_enet = 0.6\nlayer_freezed = 50\nepochs_ori = 12\nepochs_enet = 12\nbatch_size_enet = 64\nlearning_scheduler_ori = optimizers.schedules.ExponentialDecay(\n    initial_learning_rate=0.00008, decay_steps=500, decay_rate=0.4\n)\nlearning_scheduler_enet = optimizers.schedules.ExponentialDecay(\n    initial_learning_rate=0.0001, decay_steps=500, decay_rate=0.5\n) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:18.128919Z","iopub.execute_input":"2025-04-11T02:31:18.129217Z","iopub.status.idle":"2025-04-11T02:31:18.148637Z","shell.execute_reply.started":"2025-04-11T02:31:18.129186Z","shell.execute_reply":"2025-04-11T02:31:18.147791Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Preprocess the Image Data","metadata":{}},{"cell_type":"code","source":"trainset_ori, valset_ori, testset_ori, num_train_ori, num_val_ori = preprocess_image_for_compare(\n    gcs_path, \n    trainset_path, \n    image_size_enet, \n    train_ratio_cnn, \n    val_ratio_cnn, \n    batch_size_enet, \n    autotune,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:18.149365Z","iopub.execute_input":"2025-04-11T02:31:18.149619Z","iopub.status.idle":"2025-04-11T02:31:46.772375Z","shell.execute_reply.started":"2025-04-11T02:31:18.149592Z","shell.execute_reply":"2025-04-11T02:31:46.771718Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(num_train_ori)\n# print(num_val_ori)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:46.774489Z","iopub.execute_input":"2025-04-11T02:31:46.774701Z","iopub.status.idle":"2025-04-11T02:31:46.778037Z","shell.execute_reply.started":"2025-04-11T02:31:46.774683Z","shell.execute_reply":"2025-04-11T02:31:46.777242Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainset_enet, valset_enet, testset_enet, num_train_enet, num_val_enet = preprocess_image_for_classification(\n    gcs_path, \n    trainset_path_enet, \n    valset_path_enet, \n    testset_path_enet, \n    image_size_enet, \n    batch_size_enet, \n    autotune,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:46.779527Z","iopub.execute_input":"2025-04-11T02:31:46.780032Z","iopub.status.idle":"2025-04-11T02:31:55.095427Z","shell.execute_reply.started":"2025-04-11T02:31:46.780001Z","shell.execute_reply":"2025-04-11T02:31:55.094525Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(num_train_enet)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:55.096372Z","iopub.execute_input":"2025-04-11T02:31:55.096596Z","iopub.status.idle":"2025-04-11T02:31:55.100970Z","shell.execute_reply.started":"2025-04-11T02:31:55.096577Z","shell.execute_reply":"2025-04-11T02:31:55.100246Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(num_val_enet)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:55.101681Z","iopub.execute_input":"2025-04-11T02:31:55.101911Z","iopub.status.idle":"2025-04-11T02:31:55.122493Z","shell.execute_reply.started":"2025-04-11T02:31:55.101863Z","shell.execute_reply":"2025-04-11T02:31:55.121889Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Build the Transfer Learning Model","metadata":{}},{"cell_type":"code","source":"Enet_model_enet = EfficientNetB1_Arch(\n    dropout_ratio=dropout_ratio_enet, \n    layers_freezed=layer_freezed, \n    input_shape=(*image_size_enet, 3), \n    classes=5,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:55.123299Z","iopub.execute_input":"2025-04-11T02:31:55.123590Z","iopub.status.idle":"2025-04-11T02:31:58.398905Z","shell.execute_reply.started":"2025-04-11T02:31:55.123559Z","shell.execute_reply":"2025-04-11T02:31:58.398166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Enet_model_enet.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:58.399688Z","iopub.execute_input":"2025-04-11T02:31:58.399965Z","iopub.status.idle":"2025-04-11T02:31:58.423632Z","shell.execute_reply.started":"2025-04-11T02:31:58.399930Z","shell.execute_reply":"2025-04-11T02:31:58.422948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Enet_model_done, Enet_lr_fig, enet_history = Effi_B1_train(\n    Enet_model_enet, \n    trainset_enet, \n    valset_enet, \n    epochs=epochs_enet, \n    learning_scheduler=learning_scheduler_enet, \n    num_samples=num_train_enet, \n    num_valset=num_val_enet, \n    batch_size=batch_size_enet,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:31:58.424500Z","iopub.execute_input":"2025-04-11T02:31:58.424730Z","iopub.status.idle":"2025-04-11T02:47:01.280170Z","shell.execute_reply.started":"2025-04-11T02:31:58.424711Z","shell.execute_reply":"2025-04-11T02:47:01.278922Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Training (without Image Segmentation)","metadata":{}},{"cell_type":"code","source":"tf.keras.backend.clear_session()\nEnet_model_ori = EfficientNetB1_Arch(\n    dropout_ratio=dropout_ratio_ori, \n    layers_freezed=layer_freezed, \n    input_shape=(*image_size_enet, 3), \n    classes=5,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:47:01.281306Z","iopub.execute_input":"2025-04-11T02:47:01.281655Z","iopub.status.idle":"2025-04-11T02:47:03.816525Z","shell.execute_reply.started":"2025-04-11T02:47:01.281614Z","shell.execute_reply":"2025-04-11T02:47:03.815739Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Enet_model_ori.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:47:03.817281Z","iopub.execute_input":"2025-04-11T02:47:03.817550Z","iopub.status.idle":"2025-04-11T02:47:03.840817Z","shell.execute_reply.started":"2025-04-11T02:47:03.817516Z","shell.execute_reply":"2025-04-11T02:47:03.840200Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ori_model_done, ori_lr_fig, ori_history = Effi_B1_train(\n    Enet_model_ori, \n    trainset_ori, \n    valset_ori, \n    epochs=epochs_ori, \n    learning_scheduler=learning_scheduler_ori, \n    num_samples=num_train_ori, \n    num_valset=num_val_ori, \n    batch_size=batch_size_enet,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T02:47:03.841642Z","iopub.execute_input":"2025-04-11T02:47:03.841947Z","iopub.status.idle":"2025-04-11T03:02:30.509705Z","shell.execute_reply.started":"2025-04-11T02:47:03.841916Z","shell.execute_reply":"2025-04-11T03:02:30.508763Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Evaluation","metadata":{}},{"cell_type":"code","source":"fig_enet, acc_enet = Effi_B1_test(Enet_model_done, testset_enet)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T03:02:30.510741Z","iopub.execute_input":"2025-04-11T03:02:30.511086Z","iopub.status.idle":"2025-04-11T03:02:53.430585Z","shell.execute_reply.started":"2025-04-11T03:02:30.511053Z","shell.execute_reply":"2025-04-11T03:02:53.429738Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(acc_enet)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T03:02:53.431511Z","iopub.execute_input":"2025-04-11T03:02:53.431968Z","iopub.status.idle":"2025-04-11T03:02:53.436524Z","shell.execute_reply.started":"2025-04-11T03:02:53.431934Z","shell.execute_reply":"2025-04-11T03:02:53.435613Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Save the Models and Figures","metadata":{}},{"cell_type":"code","source":"figure_dir = \"/kaggle/working/figure\"\nmodel_dir = \"/kaggle/working/model\"\n\nif not os.path.exists(figure_dir):\n    os.makedirs(figure_dir)\n\nif not os.path.exists(model_dir):\n    os.makedirs(model_dir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T03:02:53.437392Z","iopub.execute_input":"2025-04-11T03:02:53.437661Z","iopub.status.idle":"2025-04-11T03:02:53.454583Z","shell.execute_reply.started":"2025-04-11T03:02:53.437624Z","shell.execute_reply":"2025-04-11T03:02:53.453955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # save the model\n# Enet_model_done.save(os.path.join(model_dir, \"EfficientNetB1_model.h5\"))\n# Enet_model_done.save(os.path.join(model_dir, \"EfficientNetB1_model.keras\"))\n# ori_model_done.save(os.path.join(model_dir, \"EfficientNetB1_model_ori.h5\"))\n# ori_model_done.save(os.path.join(model_dir, \"EfficientNetB1_model_ori.keras\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T03:02:53.455416Z","iopub.execute_input":"2025-04-11T03:02:53.455712Z","iopub.status.idle":"2025-04-11T03:02:53.468507Z","shell.execute_reply.started":"2025-04-11T03:02:53.455683Z","shell.execute_reply":"2025-04-11T03:02:53.467657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # save the learning curve\n# Enet_lr_fig.savefig(os.path.join(figure_dir, \"EfficientNetB1_model_learning_curve.png\"))\n# fig_enet.savefig(os.path.join(figure_dir, \"EfficientNetB1_model_test_result.png\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T03:02:53.469389Z","iopub.execute_input":"2025-04-11T03:02:53.469694Z","iopub.status.idle":"2025-04-11T03:02:53.482805Z","shell.execute_reply.started":"2025-04-11T03:02:53.469662Z","shell.execute_reply":"2025-04-11T03:02:53.482065Z"}},"outputs":[],"execution_count":null}]}