# -*- coding: utf-8 -*-
from model import *
from data import *
import matplotlib.pyplot as plt

print("start")
print("Num GPUs Available: ", len(tf.config.list_physical_devices('GPU')))
print("training")
data_gen_args = dict(rotation_range=0.2,
                    width_shift_range=0.05,
                    height_shift_range=0.05,
                    shear_range=0.05,
                    zoom_range=0.05,
                    horizontal_flip=True,
                    fill_mode='nearest')
myLeads = trainGenerator(2,'data/train','images','masks',data_gen_args,save_to_dir = None)
sat = "VIIRS"
#sat = "MODIS"
model = unet()
model_checkpoint = ModelCheckpoint('unet_lead.'+sat+'.h5', monitor='loss',verbose=1, save_best_only=True)
EPOCHS = 200
Estep = 30
model.fit_generator(myLeads,steps_per_epoch=Estep ,epochs=EPOCHS,callbacks=[model_checkpoint])
model_history = model.fit_generator(myLeads,steps_per_epoch=Estep ,epochs=EPOCHS,callbacks=[model_checkpoint])
loss = model_history.history['loss']
accuracy = model_history.history['accuracy']
epochs = range(EPOCHS)
plt.figure()
plt.plot(epochs, loss, 'r', label='Loss')
plt.plot(epochs, accuracy, 'b', label='Accuracy')
plt.title('Accuracy and Loss')
plt.xlabel('Epoch')
plt.ylabel('Value')
plt.ylim([0, 1])
plt.legend()
plt.savefig(sat+'.loss_plot.png', dpi=200)

testModel = testGenerator("data/test/images")
results = model.predict_generator(testModel,30,verbose=1)
saveResult("data/test/result",results)