import tensorflow as tf
import numpy as np
import os.path
import matplotlib.pyplot as plt

# Callback, um die Fehlerwerte für jede Epoche zu sammeln und zu plotten
class PlotCallback(tf.keras.callbacks.Callback):
    def __init__(self):
        self.losses = []
        super().__init__()

    def on_epoch_end(self, epoch, logs=None):
        loss = logs['loss']
        self.losses.append(loss)

    def on_train_end(self, logs=None):
        plt.plot(self.losses)
        plt.xlabel('Epochen')
        plt.ylabel('Quadratischer Fehler')
        plt.show()

OVERWRITE_MODEL = False # True: Modell wird überschrieben
EPOCHS = 200 # Anzahl Iterationen (0: kein Training)
MODEL_NAME = "xor.model" # Für Speicherung

if os.path.exists(MODEL_NAME) and not OVERWRITE_MODEL:
    model = tf.keras.models.load_model(MODEL_NAME)
else:
    model = tf.keras.models.Sequential((
        tf.keras.layers.Input((2,)), # Input Layer mit Input-Dimension: (2,)
        tf.keras.layers.Dense(3, activation=tf.keras.activations.sigmoid), # Hidden Layer (3 Neuronen) mit Aktivierungsfunktion
        tf.keras.layers.Dense(1) # Output Layer (1 Neuron)
    ))

if EPOCHS > 0:
    x_train = np.array([[0, 0], [0, 1], [1, 0], [1, 1]]) # NumPy-Arrays verschnellern den Lernprozess wesentlich
    y_train = np.array([[0], [1], [1], [0]])

    model.compile(loss=tf.losses.mean_squared_error, # loss (loss-function): Funktion zur Berechnung des Fehlers
                  optimizer=tf.optimizers.Adam(learning_rate=0.2)) # optimizer: Funktion zur Anpassung der Gewichte (Adam = dyn. Anpassung der Lernrate)

    model.fit(x_train, # Inputs
              y_train, # Erwartete Outputs
              epochs=EPOCHS, # Anzahl Iterationen
              callbacks=[PlotCallback()]) # Callback für plotting
    if OVERWRITE_MODEL == False:
        if input('save? (y/n) ').lower() == 'y':
            model.save(MODEL_NAME)

pred = model.predict([[0, 1]]) # Vorhersage
print(pred)