{"cells":[{"metadata":{},"cell_type":"markdown","source":"**\n\nI'll be using TPU in this notebook. Being a beginner myself, I'll be referencing some notebooks to write my code.A good place to start will be this link : [https://www.kaggle.com/jessemostipak/getting-started-tpus-cassava-leaf-disease]\n**"},{"metadata":{},"cell_type":"markdown","source":"# Importing Libraries."},{"metadata":{"trusted":true},"cell_type":"code","source":"import os ,gc\nimport numpy as np \nimport pandas as pd \nimport matplotlib.pyplot as plt\nimport seaborn as sns \nfrom sklearn.model_selection import train_test_split\nimport tensorflow as tf \nimport albumentations\nfrom tensorflow import keras\nfrom tensorflow.keras import backend as k\nimport cv2\nfrom functools import partial\nimport datetime as dt \nimport json\nimport re\nfrom kaggle_datasets import KaggleDatasets\n%matplotlib inline ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#setting seed for reproducibility.\ndef set_seed(seed=7):\n    os.environ['PYTHONHASHSEED']=str(seed)\n    tf.random.set_seed(seed)\n    np.random.seed(seed)\nseed=7\nset_seed(seed=seed)    ","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Detecting TPU"},{"metadata":{"trusted":true},"cell_type":"code","source":"try:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\n    print('Device:', tpu.master())\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)\nexcept:\n    strategy = tf.distribute.get_strategy()\nprint('Number of replicas:', strategy.num_replicas_in_sync)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Setting Parameters"},{"metadata":{"trusted":true},"cell_type":"code","source":"AUTOTUNE = tf.data.experimental.AUTOTUNE\nGCS_PATH = KaggleDatasets().get_gcs_path()\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync\nIMAGE_SIZE = [512, 512]\nCLASSES = ['0', '1', '2', '3', '4']\nEPOCHS = 50\nmodel_filepath1='Cassava_effnet_b7_v1.h5'    #path to save model\nmodel_filepath2='EfficientNetB_V2.h5'","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Visualizing Images."},{"metadata":{"trusted":true},"cell_type":"code","source":"#reading the file with label description.\npath='../input/cassava-leaf-disease-classification/label_num_to_disease_map.json'\nwith open(path,'r') as f:\n    disease_labels=json.load(f)\ndisease_labels","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_df=pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\ntrain_df['label']=train_df['label'].astype(str)\ntrain_path='../input/cassava-leaf-disease-classification/train_images'\n\n#fuction to view randomly sampled images from dataset.\ndef show_sample_images(df):\n    dfs=df.sample(5)\n    plt.figure(figsize=(20,7))\n    for i,img in enumerate(dfs.image_id):\n        plt.subplot(1,5,i+1)\n        img=cv2.imread(os.path.join(train_path +'/' +img))\n        img=cv2.resize(img,(512,512))\n        plt.imshow(img)\n        plt.axis('off')\n    plt.tight_layout()\n    plt.show()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print('0: Cassava Bacterial Blight (CBB)')\nshow_sample_images(train_df[train_df.label=='0'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print('1:Cassava Brown Streak Disease (CBSD)')\nshow_sample_images(train_df[train_df.label=='1'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print('2:Cassava Green Mottle (CGM)')\nshow_sample_images(train_df[train_df.label=='2'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print('3:Cassava Mosaic Disease (CMD)')\nshow_sample_images(train_df[train_df.label=='3'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print('4:Healthy')\nshow_sample_images(train_df[train_df.label=='4'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# lets see the class balance in the dataset.\nplt.figure(figsize=(16,8))\nsns.countplot(train_df['label'])\nplt.title('Class Distribution')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"\n\n#helper fuctions :\n\n#decode image:\ndef decode_image(image):\n    image = tf.image.decode_jpeg(image, channels=3)\n    image = tf.cast(image, tf.float32) / 255.0\n    image = tf.reshape(image, [*IMAGE_SIZE, 3])\n    return image\n\n#Function to read TFrecords:\ndef read_tfrecord(example, labeled):\n    tfrecord_format = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"target\": tf.io.FixedLenFeature([], tf.int64)\n    } if labeled else {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"image_name\": tf.io.FixedLenFeature([], tf.string)\n    }\n    example = tf.io.parse_single_example(example, tfrecord_format)\n    image = decode_image(example['image'])\n    if labeled:\n        label = tf.cast(example['target'], tf.int32)\n        return image, label\n    idnum = example['image_name']\n    return image, idnum\n\ndef load_dataset(filenames, labeled=True, ordered=False):\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = False # disable order, increase speed\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTOTUNE) # automatically interleaves reads from multiple files\n    dataset = dataset.with_options(ignore_order) # uses data as soon as it streams in, rather than in its original order\n    dataset = dataset.map(partial(read_tfrecord, labeled=labeled), num_parallel_calls=AUTOTUNE)\n    return dataset\n\ndef count_data_items(filenames):\n    n = [int(re.compile(r\"-([0-9]*)\\.\").search(filename).group(1)) for filename in filenames]\n    return np.sum(n)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Train Test Split."},{"metadata":{"trusted":true},"cell_type":"code","source":"TRAINING_FILENAMES, VALID_FILENAMES = train_test_split(\n    tf.io.gfile.glob(GCS_PATH + '/train_tfrecords/ld_train*.tfrec'),\n    test_size=0.1,random_state=7)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Data Augmentation"},{"metadata":{"trusted":true},"cell_type":"code","source":"def data_augment(image,label):\n    image=tf.image.random_contrast(image,0.7,1.2)\n    image=tf.image.random_brightness(image,0.2)\n    image=tf.image.random_saturation(image,0.7,1.3)\n    image=tf.image.random_hue(image,0.1,0.3)\n    image = tf.image.random_flip_left_right(image)\n    image=tf.image.random_flip_up_down(image)\n    \n    return image,label","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Training and validation data**"},{"metadata":{"trusted":true},"cell_type":"code","source":"#getting training and validation datasets.\n\n#Training Dataset.\ndef get_training_dataset():\n    dataset = load_dataset(TRAINING_FILENAMES, labeled=True)  \n    dataset = dataset.map(data_augment, num_parallel_calls=AUTOTUNE)  \n    dataset = dataset.repeat()\n    dataset = dataset.shuffle(2048)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTOTUNE)\n    return dataset\n\n#validation dataset \ndef get_validation_dataset(ordered=False):\n    dataset = load_dataset(VALID_FILENAMES, labeled=True, ordered=ordered) \n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.cache()\n    dataset = dataset.prefetch(AUTOTUNE)\n    return dataset\n\n\n\n#loading training and validation data.\n#training_set\ntraining_dataset=get_training_dataset()\n\n#validation_set\nvalid_dataset=get_validation_dataset()\n\nNum_train = count_data_items(TRAINING_FILENAMES)\nNum_valid = count_data_items(VALID_FILENAMES)\n\nprint('Number of training Images {} \\n Number of Validation Images {}'.format(Num_train,Num_valid))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Building the model."},{"metadata":{},"cell_type":"markdown","source":"# Model 1 "},{"metadata":{"trusted":true},"cell_type":"code","source":"#Model 1 with InceptionResnet base:\n\n\nwith strategy.scope():\n    # loading Inception_ResnetV2:\n    incep_resnet=tf.keras.applications.InceptionResNetV2(include_top=False,weights='imagenet',input_shape=(512,512,3))\n    \n    model=keras.Sequential([\n    keras.Input(shape=(512,512,3)),\n    keras.layers.experimental.preprocessing.Normalization(),\n    incep_resnet,\n    \n        \n        \n    keras.layers.GlobalAveragePooling2D(),\n    keras.layers.Flatten(),\n    keras.layers.Dense(512,activation='relu'),\n    keras.layers.Dropout(.6,seed=7),\n    keras.layers.BatchNormalization(),\n        \n#     keras.layers.Dense(16,activation='relu'),\n#     keras.layers.Dropout(.3,seed=7),\n#     keras.layers.BatchNormalization(),\n    \n        \n    keras.layers.Dense(5,activation='softmax')])\n    \n    #compiling model.\n    model.compile(optimizer=keras.optimizers.Adam(lr=1e-3),                                \n                  loss='sparse_categorical_crossentropy',metrics=['sparse_categorical_accuracy'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#function to plot accuracy and loss\ndef plot_history(history):\n    history=pd.DataFrame(history.history)\n    plt.figure(figsize=(16,8))\n\n    #plotting accuracy:\n    plt.subplot(1,2,1)\n    plt.title('Accuracy')\n    plt.plot(range(len(history['loss'])),history['sparse_categorical_accuracy'],color='g',label='Training Accuracy')\n    plt.plot(range(len(history['loss'])),history['val_sparse_categorical_accuracy'],color='r',label='Validation Accuracy')\n    plt.legend()\n    #plotting loss \n    plt.subplot(1,2,2)\n    plt.title('Loss')\n    plt.plot(range(len(history['loss'])),history['loss'],color='g',label='Training_loss')\n    plt.plot(range(len(history['loss'])),history['val_loss'],color='r',label='Validation loss')\n\n    plt.legend()\n    plt.tight_layout()\n    plt.show()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"\nmodel_filepath='IncepResNetV2.h5'\n#callbacks:\n#to reduce learning rate by factor of .5 if val_loss does not improve after 2 epochs.\nreduce_lr=keras.callbacks.ReduceLROnPlateau(monitor='val_loss',factor=.5,patience=3,min_delta=0.001,verbosity=1)\n\n#stop training if validation loss does not decrease by atleast .001 in 5 epochs.\nearly_stopping=keras.callbacks.EarlyStopping(min_delta=.001,patience=10,monitor='val_sparse_categorical_accuracy',\n                                             restore_best_weights=True,mode='max')\n\n#save the best weights and the model.\nmodel_checkpoint=keras.callbacks.ModelCheckpoint(filepath=model_filepath,monitor='val_sparse_categorical_accuracy',\n                                                 save_best_only=True)\n\ncallbacks_v1=[reduce_lr,model_checkpoint,early_stopping]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"history1=model.fit(training_dataset,\n                 steps_per_epoch=Num_train//BATCH_SIZE,\n                 validation_data=valid_dataset,\n                 validation_steps=Num_valid//BATCH_SIZE,\n                 callbacks=callbacks_v1,\n                 epochs=EPOCHS)\nk.clear_session()\n\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Learning Curve.** "},{"metadata":{"trusted":true},"cell_type":"code","source":"plot_history(history=history1)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Model 2"},{"metadata":{"trusted":true},"cell_type":"code","source":"#model 2\nwith strategy.scope():\n    # loading Inception_ResnetV2:\n    effnet=tf.keras.applications.EfficientNetB7(include_top=False,weights='imagenet',input_shape=(512,512,3))\n    \n    model2=keras.Sequential([\n    keras.Input(shape=(512,512,3)),\n    keras.layers.experimental.preprocessing.Normalization(),\n    incep_resnet,\n    \n        \n        \n    keras.layers.GlobalAveragePooling2D(),\n    keras.layers.Flatten(),\n    keras.layers.Dense(256,activation='relu'),\n    keras.layers.Dropout(.5,seed=7),\n    keras.layers.BatchNormalization(),\n    \n    \n    keras.layers.Dense(16,activation='relu'),\n    keras.layers.Dropout(.5,seed=7),\n    keras.layers.BatchNormalization(),\n        \n    keras.layers.Dense(5,activation='softmax')])\n    \n    #compiling model.\n    model2.compile(optimizer=keras.optimizers.Adam(lr=1e-3),                                \n                  loss='sparse_categorical_crossentropy',metrics=['sparse_categorical_accuracy'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"\n#callbacks:\n#to reduce learning rate by factor of .5 if val_loss does not improve after 2 epochs.\nreduce_lr=keras.callbacks.ReduceLROnPlateau(monitor='val_loss',factor=.5,patience=2,min_delta=0.001)\n\n#stop training if validation loss does not decrease by atleast .001 in 5 epochs.\nearly_stopping=keras.callbacks.EarlyStopping(min_delta=.001,patience=10,monitor='val_sparse_categorical_accuracy',\n                                             restore_best_weights=True,mode='max')\n\n#save the best weights and the model.\nmodel_checkpoint=keras.callbacks.ModelCheckpoint(filepath=model_filepath2,monitor='val_sparse_categorical_accuracy',\n                                                 save_best_only=True)\n\ncallbacks_v1=[reduce_lr,model_checkpoint,early_stopping]\n\n\nhistory2=model2.fit(training_dataset,\n                 steps_per_epoch=Num_train//BATCH_SIZE,\n                 validation_data=valid_dataset,\n                 validation_steps=Num_valid//BATCH_SIZE,\n                 callbacks=callbacks_v1,\n                 epochs=EPOCHS)\n\nk.clear_session()\n\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Learning curve**"},{"metadata":{"trusted":true},"cell_type":"code","source":"plot_history(history=history2)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}