import numpy as np
import pandas as pd #数据分析支持库
import matplotlib.pyplot as plt
import tensorflow as tf
from tensorflow.keras import models, layers
from tensorflow.keras.preprocessing.image import ImageDataGenerator
from keras.callbacks import ModelCheckpoint, EarlyStopping, ReduceLROnPlateau
from tensorflow.keras.applications import VGG16, ResNet50, DenseNet121, EfficientNetB0
from keras.optimizers import Adam
import os, cv2, json

# ignoring warnings
import warnings
warnings.simplefilter("ignore")

# For easy access to files
WORK_DIR = "../input/cassava-leaf-disease-classification/";
print(os.listdir(WORK_DIR))

with open('../input/cassava-leaf-disease-classification/label_num_to_disease_map.json', 'r') as file:
     labels = json.load(file)
print(labels)
     #with open('','r') as file:
     #print(file.read()) 以只读文件的方式打开某文件
data = pd.read_csv(WORK_DIR + "train.csv")#csv 其文件以纯文本的形式存储表格数据。该文件是一个字符序列，可以由任意数目的记录组成，记录间以某种换行符分割
print(data.head(n=5)) #return the first n rows
print(data.shape[0]) #number of training set images
data.label = data.label.astype("str") #astype可改变data.label里元素的数据类型，由原来的object改成string
print(data.label.value_counts())


##image visualization
IMG_SIZE = 300

plt.figure(figsize=(15,12))
data_sample = data.sample(9).reset_index(drop = True) #随机抽取九张图并且把它们的序号标为0-8

print(data_sample)

for i in range(8):
    plt.subplot(2,4,i+1)
    img = cv2.imread(WORK_DIR + "train_images/" + data_sample.image_id[i]) #cv2.imread读取图像，读出来图片颜色是BGR格式
    img = cv2.resize(img,(IMG_SIZE, IMG_SIZE))
    img = cv2.cvtColor(img,cv2.COLOR_BGR2RGB) #转换图片的颜色空间，从BGR格式到RGB格式
    
    plt.axis("off")
    plt.imshow(img)
    plt.title(labels.get(data_sample.label[i]))
    
    plt.tight_layout()
    plt.show()
    
train_generator = ImageDataGenerator(
                                    #featurewise_center=False,                                    
                                    #samplewise_center=False,
                                    #featurewise_std_normalization=False,
                                    #samplewise_std_normalization=False, 
                                    #zca_whitening=False,
                                    #zca_epsilon=1e-06,
                                    #rotation_range=90,
                                    width_shift_range=0.2,
                                    height_shift_range=0.2,
                                    #brightness_range=None,
                                    shear_range=25,
                                    zoom_range=0.3,
                                    #channel_shift_range=0.0,
                                    #fill_mode="nearest",
                                    #cval=0.0,
                                    horizontal_flip=True,
                                    vertical_flip=True,
                                    #rescale=None,
                                    #preprocessing_function=None,
                                    #data_format=None,
                                    validation_split=0.2,
                                    #dtype=None,
) \
        .flow_from_dataframe(
                            data,
                            directory = WORK_DIR + "train_images",
                            x_col = "image_id",
                            y_col = "label",
                            #weight_col = None,
                            target_size = (IMG_SIZE, IMG_SIZE),
                            #color_mode = "rgb",
                            #classes = None,
                            class_mode = "categorical",
                            batch_size = 32,
                            shuffle = True,
                            #seed = 34,
                            #save_to_dir = None,
                            #save_prefix = "",
                            #save_format = "png",
                            subset = "training",
                            #interpolation = "nearest",
                            #validate_filenames = True   
)   
        
valid_generator = ImageDataGenerator(
                                    validation_split = 0.2
) \
        .flow_from_dataframe(
                            data,
                            directory = WORK_DIR + "train_images",
                            x_col = "image_id",
                            y_col = "label",
                            target_size = (IMG_SIZE, IMG_SIZE),
                            class_mode = "categorical",
                            batch_size = 32,
                            shuffle = True,
                            #seed = 34,
                            subset = "validation")

print(valid_generator.class_indices)

###################################MODEL###############################################



def modelEfficientNetB0():
    
    model = models.Sequential()
    model.add(EfficientNetB0(include_top = False, weights = "imagenet",
                            input_shape=(IMG_SIZE,IMG_SIZE, 3)))
    model.add(layers.GlobalAveragePooling2D())
    model.add(layers.Dense(256, activation = 'relu'))
    model.add(layers.Dropout(0.5))
    model.add(layers.Dense(5, activation = "softmax"))
    
    return model 

model = modelEfficientNetB0()
model.summary()



from tensorflow.keras import utils
from keras.utils import plot_model

plot_model(model, show_shapes=True, to_file = 'model.png')

#ModelCheckpoint: callback to save the keras model or model weights at some frequency
model_check = ModelCheckpoint(
                            "./firstTry.h5",
                            monitor = "val_loss",
                            verbose = 1,
                            save_best_only = True,
                            save_weights_only = False,
                            mode = "min")

#EarlyStopping: stop training when a monitored metric has stopped improving
early_stop= EarlyStopping(
                                monitor = "val_loss",
                                min_delta=0.001,
                                patience=3,
                                verbose=1,
                                mode="min",
                                #baseline=None,
                                restore_best_weights=False)

#ReduceLROnPlateau: reduce learning rate when a metric has stopped improving
reduce_lr = ReduceLROnPlateau(
                                monitor="val_loss",
                                factor=0.1,
                                patience=3,
                                verbose=1,
                                mode="min",
                                min_delta=0.0001,
                                #cooldown=0,
                                #min_lr=0
)

model.compile(optimizer = "adam",
            loss = "categorical_crossentropy",
            metrics = ["accuracy"])

history = model.fit_generator(train_generator,
                            epochs = 10,
                            validation_data = valid_generator,
                             callbacks = [model_check,early_stop,reduce_lr])

# save model and architecture to single file
model.save("model.h5")
print("Saved model to disk")








    