{"cells":[{"metadata":{"_uuid":"476ee169-d65c-4163-ad02-04f1e532fd5c","_cell_guid":"4d3f86e7-2f60-447a-8549-1135f11e6fae","trusted":true},"cell_type":"code","source":"import os\n\nimport albumentations as A\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom tensorflow.keras.applications import EfficientNetB0\nfrom tensorflow.keras.layers import Dense, GlobalAveragePooling2D\nfrom tensorflow.keras.models import load_model, Sequential\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\n\n\nclass Transform():\n    def __init__(self):\n        self.aug = A.Compose([\n            A.RandomRotate90(),\n            A.Rotate(),\n            A.Flip(),\n        ])\n\n    def __call__(self, image):\n        return self.aug(image=image)['image']\n\n\nINPUT_DIR = '/kaggle/input/cassava-leaf-disease-classification'\nWORKING_DIR = '/kaggle/working'\n\nos.makedirs('checkpoints', exist_ok=True)\n\ndf = pd.read_csv(os.path.join(INPUT_DIR, 'train.csv'))\ndf['label'] = df['label'].astype(str)\n\ntrain_df, validation_df = train_test_split(\n    df,\n    test_size=0.1,\n    random_state=42,\n    stratify=df[['label']],\n)\ntest_df = pd.DataFrame({\n    'image_id': list(os.listdir(os.path.join(INPUT_DIR, 'test_images'))),\n    'label': '0',\n})\ntest_df.to_csv('submission.csv', index=False)\n\nbatch_size = 16\ntrain_datagen = ImageDataGenerator(preprocessing_function=Transform())\\\n    .flow_from_dataframe(\n        dataframe=train_df,\n        directory=os.path.join(INPUT_DIR, 'train_images'),\n        x_col='image_id',\n        y_col='label',\n        batch_size=batch_size,\n    )\nvalidation_datagen = ImageDataGenerator()\\\n    .flow_from_dataframe(\n        dataframe=validation_df,\n        directory=os.path.join(INPUT_DIR, 'train_images'),\n        x_col='image_id',\n        y_col='label',\n        batch_size=batch_size,\n    )\n\nmodel = Sequential([\n    EfficientNetB0(include_top=False),\n    GlobalAveragePooling2D(),\n    Dense(units=5, activation='softmax'),\n])\n\nmodel.compile('adam', 'categorical_crossentropy', ['accuracy'])\nmodel.fit(\n    train_datagen,\n    epochs=25,\n    validation_data=validation_datagen,\n    workers=4,\n    verbose=2,\n)\n\nmodel.save(os.path.join(WORKING_DIR, 'cassava.h5'), save_format='h5')","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}