{"cells":[{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.layers import Input, Flatten, Dense, Dropout, GlobalAveragePooling2D\nfrom tensorflow.keras.applications.mobilenet import MobileNet, preprocess_input\nimport math\nimport pandas as pd","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"NUM_CLASSES = 5\nIMG_WIDTH, IMG_HEIGHT = 224,224\nBATCH_SIZE = 64","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"datagen = ImageDataGenerator(preprocessing_function=preprocess_input,\n                                     shear_range=0.3,\n                                   rotation_range=20,\n                                   width_shift_range=0.4,\n                                   height_shift_range=0.5,\n                                   zoom_range=0.3,\n                                  horizontal_flip=True,\n                                  vertical_flip=True,\n                                  validation_split=0.2)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data['label'] = data['label'].map({0: \"Cassava Bacterial Blight (CBB)\", 1: \"Cassava Brown Streak Disease (CBSD)\", 2: \"Cassava Green Mottle (CGM)\", 3: \"Cassava Mosaic Disease (CMD)\", 4: \"Healthy\"})","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_generator = datagen.flow_from_dataframe(dataframe=data, \n                    directory=\"../input/cassava-leaf-disease-classification/train_images\", \n                    x_col=\"image_id\", \n                    y_col=\"label\", \n                    class_mode=\"categorical\", \n                    target_size=(IMG_WIDTH, IMG_HEIGHT), \n                    batch_size=BATCH_SIZE,\n                    subset='training')\n\nvalidation_generator = datagen.flow_from_dataframe(dataframe=data, \n                    directory=\"../input/cassava-leaf-disease-classification/train_images\", \n                    x_col=\"image_id\", \n                    y_col=\"label\", \n                    class_mode=\"categorical\", \n                    target_size=(IMG_WIDTH, IMG_HEIGHT), \n                    batch_size=BATCH_SIZE,\n                    subset='validation')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"TRAIN_SAMPLES = 17118\nVALIDATION_SAMPLES = 4279","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def model_maker():\n    base_model = MobileNet(include_top=False,\n                           input_shape=(IMG_WIDTH, IMG_HEIGHT, 3))\n    for layer in base_model.layers[:-6]:\n        layer.trainable = False\n    input = Input(shape=(IMG_WIDTH, IMG_HEIGHT, 3))\n    custom_model = base_model(input)\n    custom_model = GlobalAveragePooling2D()(custom_model)\n    custom_model = Dense(80, activation='relu')(custom_model)\n    custom_model = Dropout(0.4)(custom_model)\n    predictions = Dense(NUM_CLASSES, activation='softmax')(custom_model)\n    return Model(inputs=input, outputs=predictions)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**As the dataset is Imbalanced We calculate class weights**"},{"metadata":{"trusted":true},"cell_type":"code","source":"from collections import Counter\n\ncounter = Counter(train_generator.classes)                          \nmax_val = float(max(counter.values()))       \nclass_weights = {class_id : max_val/num_images for class_id, num_images in counter.items()}   ","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Training the Model**"},{"metadata":{"trusted":true},"cell_type":"code","source":"model = model_maker()\nmodel.compile(loss='categorical_crossentropy', #FocalLoss(alpha=1)\n              optimizer=tf.keras.optimizers.Adam(0.001),\n              metrics=['acc'])\nmodel.fit_generator(\n    train_generator,\n    #class_weight=class_weights,\n    steps_per_epoch=math.ceil(float(TRAIN_SAMPLES) / BATCH_SIZE),\n    epochs=50,\n    validation_data=validation_generator,\n    validation_steps=math.ceil(float(VALIDATION_SAMPLES) / BATCH_SIZE))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Saving the Model**"},{"metadata":{"trusted":true},"cell_type":"code","source":"import time\nt = time.time()\n\nexport_path = \"saved_models/MobileNet/{}\".format(int(t))\nmodel.save(export_path)\n\nexport_path","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}