{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.9","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"}],"dockerImageVersionId":30061,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Libraries and Configurations","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport tensorflow.keras.layers as L\nimport tensorflow_addons as tfa\nimport glob, random, os, warnings\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix, classification_report\nimport seaborn as sns\n\nprint('TensorFlow Version ' + tf.__version__)\n\ndef seed_everything(seed = 0):\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    os.environ['TF_DETERMINISTIC_OPS'] = '1'\n\nseed_everything()\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:05:06.198316Z","iopub.execute_input":"2026-06-10T01:05:06.198683Z","iopub.status.idle":"2026-06-10T01:05:11.764498Z","shell.execute_reply.started":"2026-06-10T01:05:06.198593Z","shell.execute_reply":"2026-06-10T01:05:11.763754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_size = 224\nbatch_size = 16\nn_classes = 5\n\ntrain_path = '/kaggle/input/cassava-leaf-disease-classification/train_images'\ntest_path = '/kaggle/input/cassava-leaf-disease-classification/test_images'\n\ndf_train = pd.read_csv('/kaggle/input/cassava-leaf-disease-classification/train.csv', dtype = 'str')\n\ntest_images = glob.glob(test_path + '/*.jpg')\ndf_test = pd.DataFrame(test_images, columns = ['image_path'])\n\nclasses = {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\"}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:05:14.001696Z","iopub.execute_input":"2026-06-10T01:05:14.002040Z","iopub.status.idle":"2026-06-10T01:05:14.042109Z","shell.execute_reply.started":"2026-06-10T01:05:14.002008Z","shell.execute_reply":"2026-06-10T01:05:14.041287Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Augmentations","metadata":{}},{"cell_type":"code","source":"def data_augment(image):\n    p_spatial = tf.random.uniform([], 0, 1.0, dtype = tf.float32)\n    p_rotate = tf.random.uniform([], 0, 1.0, dtype = tf.float32)\n \n    image = tf.image.random_flip_left_right(image)\n    image = tf.image.random_flip_up_down(image)\n    \n    if p_spatial > .75:\n        image = tf.image.transpose(image)\n        \n    # Rotates\n    if p_rotate > .75:\n        image = tf.image.rot90(image, k = 3) # rotate 270º\n    elif p_rotate > .5:\n        image = tf.image.rot90(image, k = 2) # rotate 180º\n    elif p_rotate > .25:\n        image = tf.image.rot90(image, k = 1) # rotate 90º\n        \n    return image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:05:16.393655Z","iopub.execute_input":"2026-06-10T01:05:16.393955Z","iopub.status.idle":"2026-06-10T01:05:16.399993Z","shell.execute_reply.started":"2026-06-10T01:05:16.393930Z","shell.execute_reply":"2026-06-10T01:05:16.399166Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Generator","metadata":{}},{"cell_type":"code","source":"datagen = tf.keras.preprocessing.image.ImageDataGenerator(samplewise_center = True,\n                                                          samplewise_std_normalization = True,\n                                                          validation_split = 0.2,\n                                                          preprocessing_function = data_augment)\n\ntrain_gen = datagen.flow_from_dataframe(dataframe = df_train,\n                                        directory = train_path,\n                                        x_col = 'image_id',\n                                        y_col = 'label',\n                                        subset = 'training',\n                                        batch_size = batch_size,\n                                        seed = 1,\n                                        color_mode = 'rgb',\n                                        shuffle = True,\n                                        class_mode = 'categorical',\n                                        target_size = (image_size, image_size))\n\nvalid_gen = datagen.flow_from_dataframe(dataframe = df_train,\n                                        directory = train_path,\n                                        x_col = 'image_id',\n                                        y_col = 'label',\n                                        subset = 'validation',\n                                        batch_size = batch_size,\n                                        seed = 1,\n                                        color_mode = 'rgb',\n                                        shuffle = False,\n                                        class_mode = 'categorical',\n                                        target_size = (image_size, image_size))\n\ntest_gen = datagen.flow_from_dataframe(dataframe = df_test,\n                                       x_col = 'image_path',\n                                       y_col = None,\n                                       batch_size = batch_size,\n                                       seed = 1,\n                                       color_mode = 'rgb',\n                                       shuffle = False,\n                                       class_mode = None,\n                                       target_size = (image_size, image_size))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:05:18.314988Z","iopub.execute_input":"2026-06-10T01:05:18.315324Z","iopub.status.idle":"2026-06-10T01:06:19.756059Z","shell.execute_reply.started":"2026-06-10T01:05:18.315293Z","shell.execute_reply":"2026-06-10T01:06:19.755248Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Sample Images Visualization","metadata":{}},{"cell_type":"code","source":"images = [train_gen[0][0][i] for i in range(16)]\nfig, axes = plt.subplots(3, 5, figsize = (10, 10))\n\naxes = axes.flatten()\n\nfor img, ax in zip(images, axes):\n    ax.imshow(img.reshape(image_size, image_size, 3))\n    ax.axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:06:34.576170Z","iopub.execute_input":"2026-06-10T01:06:34.576529Z","iopub.status.idle":"2026-06-10T01:06:41.368421Z","shell.execute_reply.started":"2026-06-10T01:06:34.576502Z","shell.execute_reply":"2026-06-10T01:06:41.367424Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Hyperparameters","metadata":{}},{"cell_type":"code","source":"learning_rate = 0.001\nweight_decay = 0.0001\nnum_epochs = 1\n\npatch_size = 7  # Size of the patches to be extract from the input images\nnum_patches = (image_size // patch_size) ** 2\nprojection_dim = 64\nnum_heads = 4\ntransformer_units = [\n    projection_dim * 2,\n    projection_dim,\n]  # Size of the transformer layers\ntransformer_layers = 8\nmlp_head_units = [56, 28]  # Size of the dense layers of the final classifier","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:06:46.297150Z","iopub.execute_input":"2026-06-10T01:06:46.297457Z","iopub.status.idle":"2026-06-10T01:06:46.301986Z","shell.execute_reply.started":"2026-06-10T01:06:46.297433Z","shell.execute_reply":"2026-06-10T01:06:46.301086Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Building the Model and it's Components","metadata":{}},{"cell_type":"markdown","source":"## 1. Multilayer Perceptron (MLP)","metadata":{}},{"cell_type":"code","source":"def mlp(x, hidden_units, dropout_rate):\n    for units in hidden_units:\n        x = L.Dense(units, activation = tf.nn.gelu)(x)\n        x = L.Dropout(dropout_rate)(x)\n    return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:06:49.306027Z","iopub.execute_input":"2026-06-10T01:06:49.306381Z","iopub.status.idle":"2026-06-10T01:06:49.311038Z","shell.execute_reply.started":"2026-06-10T01:06:49.306350Z","shell.execute_reply":"2026-06-10T01:06:49.310219Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Patch Creation Layer","metadata":{}},{"cell_type":"code","source":"class Patches(L.Layer):\n    def __init__(self, patch_size):\n        super(Patches, self).__init__()\n        self.patch_size = patch_size\n\n    def call(self, images):\n        batch_size = tf.shape(images)[0]\n        patches = tf.image.extract_patches(\n            images = images,\n            sizes = [1, self.patch_size, self.patch_size, 1],\n            strides = [1, self.patch_size, self.patch_size, 1],\n            rates = [1, 1, 1, 1],\n            padding = 'VALID',\n        )\n        patch_dims = patches.shape[-1]\n        patches = tf.reshape(patches, [batch_size, -1, patch_dims])\n        return patches","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:06:51.303834Z","iopub.execute_input":"2026-06-10T01:06:51.304180Z","iopub.status.idle":"2026-06-10T01:06:51.309872Z","shell.execute_reply.started":"2026-06-10T01:06:51.304151Z","shell.execute_reply":"2026-06-10T01:06:51.309075Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Sample Image Patches Visualization","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(4, 4))\n\nx = train_gen.next()\nimage = x[0][0]\n\nplt.imshow(image.astype('uint8'))\nplt.axis('off')\n\nresized_image = tf.image.resize(\n    tf.convert_to_tensor([image]), size = (image_size, image_size)\n)\n\npatches = Patches(patch_size)(resized_image)\nprint(f'Image size: {image_size} X {image_size}')\nprint(f'Patch size: {patch_size} X {patch_size}')\nprint(f'Patches per image: {patches.shape[1]}')\nprint(f'Elements per patch: {patches.shape[-1]}')\n\nn = int(np.sqrt(patches.shape[1]))\nplt.figure(figsize=(4, 4))\n\nfor i, patch in enumerate(patches[0]):\n    ax = plt.subplot(n, n, i + 1)\n    patch_img = tf.reshape(patch, (patch_size, patch_size, 3))\n    plt.imshow(patch_img.numpy().astype('uint8'))\n    plt.axis('off')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:06:53.229959Z","iopub.execute_input":"2026-06-10T01:06:53.230286Z","iopub.status.idle":"2026-06-10T01:07:28.535670Z","shell.execute_reply.started":"2026-06-10T01:06:53.230258Z","shell.execute_reply":"2026-06-10T01:07:28.534897Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Patch Encoding Layer\nThe `PatchEncoder` layer will linearly transform a patch by projecting it into a vector of size `projection_dim`. In addition, it adds a learnable position embedding to the projected vector.","metadata":{}},{"cell_type":"code","source":"class PatchEncoder(L.Layer):\n    def __init__(self, num_patches, projection_dim):\n        super(PatchEncoder, self).__init__()\n        self.num_patches = num_patches\n        self.projection = L.Dense(units = projection_dim)\n        self.position_embedding = L.Embedding(\n            input_dim = num_patches, output_dim = projection_dim\n        )\n\n    def call(self, patch):\n        positions = tf.range(start = 0, limit = self.num_patches, delta = 1)\n        encoded = self.projection(patch) + self.position_embedding(positions)\n        return encoded","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:07:28.554121Z","iopub.execute_input":"2026-06-10T01:07:28.554488Z","iopub.status.idle":"2026-06-10T01:07:28.561083Z","shell.execute_reply.started":"2026-06-10T01:07:28.554463Z","shell.execute_reply":"2026-06-10T01:07:28.560397Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Build the ViT model\nThe ViT model consists of multiple Transformer blocks, which use the `MultiHeadAttention` layer as a self-attention mechanism applied to the sequence of patches. The Transformer blocks produce a `[batch_size, num_patches, projection_dim]` tensor, which is processed via an classifier head with softmax to produce the final class probabilities output.\n\nUnlike the technique described in the paper, which prepends a learnable embedding to the sequence of encoded patches to serve as the image representation, all the outputs of the final Transformer block are reshaped with `Flatten()` and used as the image representation input to the classifier head. Note that the `GlobalAveragePooling1D` layer could also be used instead to aggregate the outputs of the Transformer block, especially when the number of patches and the projection dimensions are large.","metadata":{}},{"cell_type":"code","source":"def vision_transformer():\n    inputs = L.Input(shape = (image_size, image_size, 3))\n    \n    # Create patches.\n    patches = Patches(patch_size)(inputs)\n    \n    # Encode patches.\n    encoded_patches = PatchEncoder(num_patches, projection_dim)(patches)\n\n    # Create multiple layers of the Transformer block.\n    for _ in range(transformer_layers):\n        \n        # Layer normalization 1.\n        x1 = L.LayerNormalization(epsilon = 1e-6)(encoded_patches)\n        \n        # Create a multi-head attention layer.\n        attention_output = L.MultiHeadAttention(\n            num_heads = num_heads, key_dim = projection_dim, dropout = 0.1\n        )(x1, x1)\n        \n        # Skip connection 1.\n        x2 = L.Add()([attention_output, encoded_patches])\n        \n        # Layer normalization 2.\n        x3 = L.LayerNormalization(epsilon = 1e-6)(x2)\n        \n        # MLP.\n        x3 = mlp(x3, hidden_units = transformer_units, dropout_rate = 0.1)\n        \n        # Skip connection 2.\n        encoded_patches = L.Add()([x3, x2])\n\n    # Create a [batch_size, projection_dim] tensor.\n    representation = L.LayerNormalization(epsilon = 1e-6)(encoded_patches)\n    representation = L.Flatten()(representation)\n    representation = L.Dropout(0.5)(representation)\n    \n    # Add MLP.\n    features = mlp(representation, hidden_units = mlp_head_units, dropout_rate = 0.5)\n    \n    # Classify outputs.\n    logits = L.Dense(n_classes)(features)\n    \n    # Create the model.\n    model = tf.keras.Model(inputs = inputs, outputs = logits)\n    \n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:07:28.584858Z","iopub.execute_input":"2026-06-10T01:07:28.585156Z","iopub.status.idle":"2026-06-10T01:07:28.595162Z","shell.execute_reply.started":"2026-06-10T01:07:28.585132Z","shell.execute_reply":"2026-06-10T01:07:28.594510Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"decay_steps = train_gen.n // train_gen.batch_size\ninitial_learning_rate = learning_rate\n\nlr_decayed_fn = tf.keras.experimental.CosineDecay(initial_learning_rate, decay_steps)\n\nlr_scheduler = tf.keras.callbacks.LearningRateScheduler(lr_decayed_fn)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:07:28.596031Z","iopub.execute_input":"2026-06-10T01:07:28.596261Z","iopub.status.idle":"2026-06-10T01:07:28.608027Z","shell.execute_reply.started":"2026-06-10T01:07:28.596239Z","shell.execute_reply":"2026-06-10T01:07:28.607360Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = tf.keras.optimizers.Adam(learning_rate = learning_rate)\n\nmodel = vision_transformer()\n    \nmodel.compile(optimizer = optimizer, \n              loss = tf.keras.losses.CategoricalCrossentropy(label_smoothing = 0.1), \n              metrics = ['accuracy'])\n\n\nSTEP_SIZE_TRAIN = train_gen.n // train_gen.batch_size\nSTEP_SIZE_VALID = valid_gen.n // valid_gen.batch_size\n\nearlystopping = tf.keras.callbacks.EarlyStopping(monitor = 'val_accuracy',\n                                                 min_delta = 1e-4,\n                                                 patience = 5,\n                                                 mode = 'max',\n                                                 restore_best_weights = True,\n                                                 verbose = 1)\n\ncheckpointer = tf.keras.callbacks.ModelCheckpoint(filepath = './model.hdf5',\n                                                  monitor = 'val_accuracy', \n                                                  verbose = 1, \n                                                  save_best_only = True,\n                                                  save_weights_only = True,\n                                                  mode = 'max')\n\ncallbacks = [earlystopping, lr_scheduler, checkpointer]\n\nmodel.fit(x = train_gen,\n          steps_per_epoch = STEP_SIZE_TRAIN,\n          validation_data = valid_gen,\n          validation_steps = STEP_SIZE_VALID,\n          epochs = num_epochs,\n          callbacks = callbacks)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:07:28.609027Z","iopub.execute_input":"2026-06-10T01:07:28.609275Z","iopub.status.idle":"2026-06-10T01:18:47.169159Z","shell.execute_reply.started":"2026-06-10T01:07:28.609238Z","shell.execute_reply":"2026-06-10T01:18:47.168300Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Results","metadata":{}},{"cell_type":"code","source":"print('Training results')\nmodel.evaluate(train_gen)\n\nprint('Validation results')\nmodel.evaluate(valid_gen)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T01:21:35.372433Z","iopub.execute_input":"2026-06-10T01:21:35.372811Z","iopub.status.idle":"2026-06-10T01:27:38.037412Z","shell.execute_reply.started":"2026-06-10T01:21:35.372775Z","shell.execute_reply":"2026-06-10T01:27:38.036567Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Summary\n\nNote that the state of the art results reported in the paper are achieved by pre-training the ViT model using the JFT-300M dataset, then fine-tuning it on the target dataset. To improve the model quality without pre-training, you can try to train the model for more epochs, use a larger number of Transformer layers, resize the input images, change the patch size, or increase the projection dimensions. Besides, as mentioned in the paper, the quality of the model is affected not only by architecture choices, but also by parameters such as the learning rate schedule, optimizer, weight decay, etc. In practice, it's recommended to fine-tune a ViT model that was pre-trained using a large, high-resolution dataset. <br>\n\n**References:** <br>\nKeras Docs: https://keras.io/api/ <br>\nResearch Paper: https://arxiv.org/pdf/2010.11929.pdf","metadata":{}}]}