{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import tensorflow as tf\nimport pandas as pd\nfrom kaggle_datasets import KaggleDatasets\nimport math, os, re, warnings, random\nimport numpy as np \nprint(\"Tensorflow version \" + tf.__version__)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"try:\n  tpu = tf.distribute.cluster_resolver.TPUClusterResolver()  # TPU detection\n  print('Running on TPU ', tpu.cluster_spec().as_dict()['worker'])\nexcept ValueError:\n  raise BaseException('ERROR: Not connected to a TPU runtime; please see the previous cell in this notebook for instructions!')\n\ntf.config.experimental_connect_to_cluster(tpu)\ntf.tpu.experimental.initialize_tpu_system(tpu)\ntpu_strategy = tf.distribute.experimental.TPUStrategy(tpu)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\nIMAGE_SIZE = [512, 512]\nbatch_size = 16 * tpu_strategy.num_replicas_in_sync\nLEARNING_RATE = 3e-5 * tpu_strategy.num_replicas_in_sync\nEPOCHS = 30\nHEIGHT = 512\nWIDTH = 512\nCHANNELS = 3\nN_CLASSES = 5\nES_PATIENCE = 10\nN_FOLDS = 5","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"GCS_PATH = KaggleDatasets().get_gcs_path(f'cassava-leaf-disease-tfrecords-center-{HEIGHT}x{WIDTH}')\n#TRAINING_FILENAMES = tf.io.gfile.glob(database_base_path + '/train_tfrecords/*.tfrec')  # Original TFRecords\nTRAINING_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/*.tfrec')\n#print(TRAINING_FILENAMES)\n#NUM_TRAINING_IMAGES = count_data_items(TRAINING_FILENAMES)\n#gcs_pattern = 'cassava-leaf-disease-tfrecords-center-512x512'\nvalidation_split = 0.19\nfilenames = tf.io.gfile.glob(TRAINING_FILENAMES) \nprint(filenames)\nsplit = int(0.9 * len(filenames))\ntrain_fns = filenames[:split]\nvalidation_fns = filenames[split:]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def decode_image(image_data):\n    \"\"\"\n        1. Decode a JPEG-encoded image to a uint8 tensor.\n        2. Cast tensor to float and normalizes (range between 0 and 1).\n        3. Resize and reshape images to the expected size.\n    \"\"\"\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    image = tf.cast(image, tf.float32) / 255.0\n                      \n    image = tf.image.resize(image, [HEIGHT, WIDTH])\n    image = tf.reshape(image, [HEIGHT, WIDTH, 3])\n    return image","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def parse_tfrecord(example, labeled=True):\n    \"\"\"\n        1. Parse data based on the 'TFREC_FORMAT' map.\n        2. Decode image.\n        3. If 'labeled' returns (image, label) if not (image, name).\n    \"\"\"\n    if labeled:\n        TFREC_FORMAT = {\n            'image': tf.io.FixedLenFeature([], tf.string), \n            'target': tf.io.FixedLenFeature([], tf.int64), \n        }\n    else:\n        TFREC_FORMAT = {\n            'image': tf.io.FixedLenFeature([], tf.string), \n            'image_name': tf.io.FixedLenFeature([], tf.string), \n        }\n    example = tf.io.parse_single_example(example, TFREC_FORMAT)\n    image = decode_image(example['image'])\n    if labeled:\n        label_or_name = tf.cast(example['target'], tf.int32)\n    else:\n        label_or_name =  example['image_name']\n    return image, label_or_name","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def load_dataset(filenames):\n      # Read from TFRecords. For optimal performance, we interleave reads from multiple files.\n    records = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTO)\n    ok=records.map(parse_tfrecord, num_parallel_calls=AUTO)\n    return ok","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_training_dataset():\n    dataset = load_dataset(train_fns)\n\n  # Create some additional training images by randomly flipping and\n  # increasing/decreasing the saturation of images in the training set. \n    def data_augment(image, one_hot_class):\n        modified = tf.image.random_flip_left_right(image)\n        modified = tf.image.random_saturation(modified, 0, 2)\n        return modified, one_hot_class\n    augmented = dataset.map(data_augment, num_parallel_calls=AUTO)\n\n  # Prefetch the next batch while training (autotune prefetch buffer size).\n    return augmented.repeat().shuffle(2048).batch(batch_size).prefetch(AUTO) ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"training_dataset = get_training_dataset()\nvalidation_dataset = load_dataset(validation_fns).batch(batch_size).prefetch(AUTO)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"CLASSES = ['Cassava Bacterial Blight', \n           'Cassava Brown Streak Disease', \n           'Cassava Green Mottle', \n           'Cassava Mosaic Disease', \n           'Healthy']","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def display_one_flower(image, title, subplot, color):\n  plt.subplot(subplot)\n  plt.axis('off')\n  plt.imshow(image)\n  plt.title(title, fontsize=16, color=color)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# If model is provided, use it to generate predictions.\ndef display_nine_flowers(images, titles, title_colors=None):\n  subplot = 512\n  plt.figure(figsize=(13,13))\n  for i in range(9):\n    color = 'black' if title_colors is None else title_colors[i]\n    display_one_flower(images[i], titles[i], 512+i, color)\n  plt.tight_layout()\n  plt.subplots_adjust(wspace=0.1, hspace=0.1)\n  plt.show()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_dataset_iterator(dataset, n_examples):\n  return dataset.unbatch().batch(n_examples).as_numpy_iterator()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"training_viz_iterator = get_dataset_iterator(training_dataset, 9)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from matplotlib import pyplot as plt\ndef batch_to_numpy_images_and_labels(data):\n    images, labels = data\n    numpy_images = images.numpy()\n    numpy_labels = labels.numpy()\n    if numpy_labels.dtype == object: # binary string in this case, these are image ID strings\n        numpy_labels = [None for _ in enumerate(numpy_images)]\n    # If no labels, only image IDs, return None for labels (this is the case for test data)\n    return numpy_images, numpy_labels","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def display_one_flower(image, title, subplot, red=False, titlesize=16):\n    plt.subplot(*subplot)\n    plt.axis('off')\n    plt.imshow(image)\n    if len(title) > 0:\n        plt.title(title, fontsize=int(titlesize) if not red else int(titlesize/1.2), color='red' if red else 'black', \n                  fontdict={'verticalalignment':'center'}, pad=int(titlesize/1.5))\n    return (subplot[0], subplot[1], subplot[2]+1)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def display_batch_of_images(databatch, predictions=None):\n    \"\"\"This will work with:\n    display_batch_of_images(images)\n    display_batch_of_images(images, predictions)\n    display_batch_of_images((images, labels))\n    display_batch_of_images((images, labels), predictions)\n    \"\"\"\n    # data\n    images, labels = batch_to_numpy_images_and_labels(databatch)\n    if labels is None:\n        labels = [None for _ in enumerate(images)]\n        \n    # auto-squaring: this will drop data that does not fit into square or square-ish rectangle\n    rows = int(math.sqrt(len(images)))\n    cols = len(images)//rows\n        \n    # size and spacing\n    FIGSIZE = 13.0\n    SPACING = 0.1\n    subplot=(rows,cols,1)\n    if rows < cols:\n        plt.figure(figsize=(FIGSIZE,FIGSIZE/cols*rows))\n    else:\n        plt.figure(figsize=(FIGSIZE/rows*cols,FIGSIZE))\n    \n    # display\n    for i, (image, label) in enumerate(zip(images[:rows*cols], labels[:rows*cols])):\n        title = '' if label is None else CLASSES[label]\n        correct = True\n        if predictions is not None:\n            title, correct = title_from_label_and_target(predictions[i], label)\n        dynamic_titlesize = FIGSIZE*SPACING/max(rows,cols)*40+3 # magic formula tested to work from 1x1 to 10x10 images\n        subplot = display_one_flower(image, title, subplot, not correct, titlesize=dynamic_titlesize)\n    \n    #layout\n    plt.tight_layout()\n    if label is None and predictions is None:\n        plt.subplots_adjust(wspace=0, hspace=0)\n    else:\n        plt.subplots_adjust(wspace=SPACING, hspace=SPACING)\n    plt.show()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_dataset = get_training_dataset()\ntrain_iter = iter(train_dataset.unbatch().batch(20))\n\ndisplay_batch_of_images(next(train_iter))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def create_model():\n  pretrained_model = tf.keras.applications.InceptionV3(input_shape=[*IMAGE_SIZE, 3], include_top=False)\n  pretrained_model.trainable = True\n  model = tf.keras.Sequential([\n    pretrained_model,\n    tf.keras.layers.GlobalAveragePooling2D(),\n    tf.keras.layers.Dense(5, activation='softmax')\n  ])\n  model.compile(\n    optimizer='adam',\n    loss = 'sparse_categorical_crossentropy',\n      #'sparse_categorical_crossentropy'\n    metrics=['accuracy']\n  )\n  return model","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"with tpu_strategy.scope(): # creating the model in the TPUStrategy scope means we will train the model on the TPU\n  model = create_model()\nmodel.summary()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def count_data_items(filenames):\n  # The number of data items is written in the name of the .tfrec files, i.e. flowers00-230.tfrec = 230 data items\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":{"trusted":true},"cell_type":"code","source":"n_train = count_data_items(train_fns)\nn_valid = count_data_items(validation_fns)\ntrain_steps = count_data_items(train_fns) // batch_size\nprint(\"TRAINING IMAGES: \", n_train, \", STEPS PER EPOCH: \", train_steps)\nprint(\"VALIDATION IMAGES: \", n_valid)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"EPOCHS = 12","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"start_lr = 0.00001\nmin_lr = 0.00001\nmax_lr = 0.00005 * tpu_strategy.num_replicas_in_sync\nrampup_epochs = 5\nsustain_epochs = 0\nexp_decay = .8","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def lrfn(epoch):\n  if epoch < rampup_epochs:\n    return (max_lr - start_lr)/rampup_epochs * epoch + start_lr\n  elif epoch < rampup_epochs + sustain_epochs:\n    return max_lr\n  else:\n    return (max_lr - min_lr) * exp_decay**(epoch-rampup_epochs-sustain_epochs) + min_lr","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"lr_callback = tf.keras.callbacks.LearningRateScheduler(lambda epoch: lrfn(epoch), verbose=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"rang = np.arange(EPOCHS)\ny = [lrfn(x) for x in rang]\nplt.plot(rang, y)\nprint('Learning rate per epoch:')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"my_model = model.fit(train_dataset, validation_data=validation_dataset, steps_per_epoch=train_steps, epochs=EPOCHS, callbacks=[lr_callback])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"final_accuracy = my_model.history[\"val_accuracy\"][-5:]\nprint(\"FINAL ACCURACY MEAN-5: \", np.mean(final_accuracy))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#reading training data\ntrain_df = pd.read_csv('../input/cassava-leaf-disease-tfrecords-center-512x512/train.csv')\n#maping the class labels mentioned in json file wiht its respective disease name\n#disease_names = open('../input/cassava-leaf-disease-classification/label_num_to_disease_map.json')\n#disease_names = json.load(disease_names)\ndisease_names = {0:'Cassava Bacterial Blight', \n           1:'Cassava Brown Streak Disease', \n           2:'Cassava Green Mottle', \n           3:'Cassava Mosaic Disease', \n           4:'Healthy'}\n#print(type(train_df['label'][1]))\n#print(train_df)\n#parse through every label value and identify the disease name based on label number from json file\ntrain_df['disease_name'] = train_df['label'].apply(lambda x: disease_names[x])#str()\nprint(train_df)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#visualizations\nimport plotly.graph_objs as go\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom plotly.subplots import make_subplots\nfrom sklearn.manifold import TSNE\nfrom tqdm.notebook import tqdm","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"fig = make_subplots(rows=1, cols=2,\n            specs=[[{\"type\": \"xy\"}, {\"type\": \"domain\"}]],)\n# value_counts: to count number of images in each class with respect to disease_name column\n# Bar plot \nt1 = go.Bar(x=train_df['disease_name'].value_counts().index, \n            y=train_df['disease_name'].value_counts().values,\n            text=train_df['disease_name'].value_counts().values,\n            textposition='auto',name='Count',\n           marker_color='indianred')\n#Pie chart with labels and counts\nt2 = go.Pie(labels=train_df['disease_name'].value_counts().index,\n           values=train_df['disease_name'].value_counts().values,\n           hole=0.3)\nfig.add_trace(t1,row=1, col=1)\nfig.add_trace(t2,row=1, col=2)\nfig.update_layout(title='Distribution of Class Labels')\nfig.show()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dataset = load_dataset(train_fns)\nnbr= {'Cassava Bacterial Blight':0, \n           'Cassava Brown Streak Disease':0, \n           'Cassava Green Mottle':0, \n           'Cassava Mosaic Disease':0, \n           'Healthy':0}\nfor et in dataset:\n    if et[1]==0:\n        nbr['Cassava Bacterial Blight']+=1\n    if et[1]==1:\n        nbr['Cassava Brown Streak Disease']+=1\n    if et[1]==2:\n        nbr['Cassava Green Mottle']+=1\n    if et[1]==3:\n        nbr['Cassava Mosaic Disease']+=1\n    if et[1]==4:\n        nbr['Healthy']+=1\n        \nprint(nbr)\n    \n\n    ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dataset_v = load_dataset(validation_fns)\nnbr_v= {'Cassava Bacterial Blight':0, \n           'Cassava Brown Streak Disease':0, \n           'Cassava Green Mottle':0, \n           'Cassava Mosaic Disease':0, \n           'Healthy':0}\nfor et in dataset_v:\n    if et[1]==0:\n        nbr_v['Cassava Bacterial Blight']+=1\n    if et[1]==1:\n        nbr_v['Cassava Brown Streak Disease']+=1\n    if et[1]==2:\n        nbr_v['Cassava Green Mottle']+=1\n    if et[1]==3:\n        nbr_v['Cassava Mosaic Disease']+=1\n    if et[1]==4:\n        nbr_v['Healthy']+=1\n        \nprint(nbr_v)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"x=['Cassava Bacterial Blight','Cassava Brown Streak Disease','Cassava Green Mottle','Cassava Mosaic Disease','Healthy']\ny=[942,1897,2068,11404,2232]\ny_v=[144,292,318,1754,344]\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#print(dataset)\nfig = go.Figure()\nt1 = go.Bar(name='Train',x=x,y=y,text=y,textposition='auto')\nt2 = go.Bar(name='Valid',x=x,y=y_v,text=y_v,textposition='auto')\nfig.add_trace(t1)\nfig.add_trace(t2)\n#x-axis and y axis title\nfig.update_xaxes(title_text=\"Class Labels\")\nfig.update_yaxes(title_text=\"Number of Images\")\nfig.update_layout(title='Train and Valid Split')\nfig.show()\n\n#Pie Chart\nfig = make_subplots(rows=1, cols=2,subplot_titles=['Train Data', 'Valid Data'],\n            specs=[[{\"type\": \"domain\"}, {\"type\": \"domain\"}]],)\n\n#Pie chart with labels and counts\nt1 = go.Pie(labels=x,\n           values=y,\n           hole=0.3)\nt2 = go.Pie(labels=x,\n           values=y_v,\n           hole=0.3)\nfig.add_trace(t1,row=1, col=1)\nfig.add_trace(t2,row=1, col=2)\nfig.update_layout(title='Distribution in Train and Valid Split')\nfig.show()","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}