# Bibliotheken laden
import board
import time
import terminalio
import displayio
from adafruit_display_text import label

# Display initialisieren
display = board.DISPLAY
display.root_group = None
display_group = displayio.Group()

# Bild darstellen
bitmap = displayio.OnDiskBitmap("/knn.bmp")
bitmap_grid = displayio.TileGrid(bitmap, pixel_shader=bitmap.pixel_shader)
bitmap_grid.x = 400  # Abstand vom linken Rand
bitmap_grid.y = 50   # Abstand vom oberen Rand
display_group.append(bitmap_grid)

# Einstellungen: Schrift
font       = terminalio.FONT
color      = 0xFFFFFF
ok_color   = 0x00FF00
fail_color = 0xFF0000
bg_color   = 0x000000

# Sigmoid-Funktion und Ableitung
def sigmoid(x): return 1 / (1 + pow(2.71828, -x))
def sigmoid_derivative(x): return x * (1 - x)

# Trainingsdatensätze (4)
# Eingänge [A, B] einer logischen Funktion
inputs = [ [0, 0], [0, 1], [1, 0], [1, 1] ] # nicht ändern

# Ausgang der logischen Funktion
targets = [ [0], [1], [1], [0] ] # XOR
#targets = [ [0], [0], [0], [1] ] # AND
#targets = [ [0], [1], [1], [1] ] # OR
#targets = [ [1], [1], [1], [0] ] # NAND

# Startgewichte mit festgelegten Werten
'''
w_input_hidden = [[0.5, -0.6], [-0.3, 0.8]] # 2x2
w_hidden_output = [0.7, -0.5]               # 1x2
'''

# Startgewichte mit zufälligen Werten
import random
w_input_hidden = [[random.uniform(-1, 1), random.uniform(-1, 1)], [random.uniform(-1, 1), random.uniform(-1, 1)]]
w_hidden_output = [random.uniform(-1, 1), random.uniform(-1, 1)]

# Biases (nicht ändern)
b_hidden = [0.0, 0.0]
b_output = 0.0

# Anzahl der Durchläufe (Epoche)
EPOCHS = 4200
SPEED = 1  # Sekunden

# Lernrate
lr = 0.5 # lernt gut (Standardwert)
#lr = 0 # lernt nichts
#lr = 100 # Netz instabil (explodiert)
#lr = 10 # instabil / osziliert
#lr = 1 # lernt schneller
#lr = 0.1 # lernt langsamer

# Start-Koordinaten für Liste
start_y = 40
next_y = 11

# Titelzeile
text_head = label.Label(font, text=f"Training eines neuronalen Netzes\tEpochen: max. {EPOCHS}\tLernrate: {lr}", color=color, scale=2)
text_head.x = 10
text_head.y = 20
display_group.append(text_head)
display.root_group = display_group

# Durchlauf der Epochen
for epoch in range(EPOCHS+1):
    # Start der Epoche mit Fehler = 0
    total_error = 0
    
    # Eingangszustände
    for i in range(4):
        x = inputs[i]
        y = targets[i][0]
        
        #if epoch % 100 == 0: print(f"Beispiel {i+1}: Eingabe: {x}, Ziel: {y}")
        
        # Forward Pass: Eingabe -> Hidden
        h_input = [0, 0]
        h_output = [0, 0]
        # Diese Schleife verarbeitet jedes Neuron im Hidden Layer
        for j in range(2):
            # Berechnet für jedes Neuron den gewichteten Eingang
            h_input[j] = x[0] * w_input_hidden[0][j] + x[1] * w_input_hidden[1][j] + b_hidden[j]
            # Berechnet für jedes Neuron die Aktivierung durch die Sigmoid-Funktion
            h_output[j] = sigmoid(h_input[j]) # Ausgabe des Hidden-Neurons
        #if epoch % 100 == 0: print(f" Hidden-Layer-Aktivierungen: {['%.3f' % h for h in h_output]}")

        # Forward Pass: Hidden -> Output
        # Ausgabe der Hidden-Neuronen wird gewichtet aufs Ausgabeneuron summiert
        o_input = h_output[0] * w_hidden_output[0] + h_output[1] * w_hidden_output[1] + b_output
        # Sigmoid-Aktivierung für den finalen Output
        o_output = sigmoid(o_input)
        #if epoch % 100 == 0: print(f" Netzwerk-Ausgabe: {o_output:.3f}")
        
        # Fehler der Vorhersage wird berechnet (Differenz zur Zielausgabe)
        error = y - o_output
        # Der quadratische Fehler fließt in die Statistik ein
        total_error += error ** 2
        #if epoch % 100 == 0: print(f" Fehler: {error:.3f}")
        
        # Backpropagation: Output-Layer
        # Delta für den Output-Layer wird berechnet
        d_output = error * sigmoid_derivative(o_output) # unter Berücksichtigung der Ableitung der Sigmoid-Funktion
        
        # Gewichte Hidden: Output anpassen
        # Aktualisiert die Gewichte vom Hidden-Layer zum Output-Layer
        for j in range(2):
            delta_w = lr * d_output * h_output[j]
            #if epoch % 100 == 0: print(f" w_hidden_output[{j}]: {delta_w:.4f}")
            w_hidden_output[j] += delta_w
        # Bias des Output-Neurons wird angepasst
        b_output += lr * d_output
        #if epoch % 100 == 0: print(f" bias_output: {lr * d_output:.4f}")
        
        # Backpropagation: Hidden Layer (Fehler und Updates)
        for j in range(2):
            # Fehlersignal (Delta) für das Hidden-Neuron j
            d_hidden = d_output * w_hidden_output[j] * sigmoid_derivative(h_output[j])
            # Gewichtsänderungen in Abhängigkeit von Lernrate (rt), Fehlersignal (d_hidden) und der Signale an den Eingängen (x)
            delta_w0 = lr * d_hidden * x[0]
            delta_w1 = lr * d_hidden * x[1]
            #if epoch % 100 == 0: print(f" w_input_hidden[0][{j}]: {delta_w0:.4f}, w_input_hidden[1][{j}]: {delta_w1:.4f}")
            # Gewichtsänderungen zu den aktuellen Gewichten addieren
            w_input_hidden[0][j] += delta_w0
            w_input_hidden[1][j] += delta_w1
            # Bias des Hidden-Neurons j wird angepasst
            b_hidden[j] += lr * d_hidden
            if epoch % 100 == 0:
                #print(f" bias_hidden[{j}]: {lr * d_hidden:.4f}")
                text_output = label.Label(font, text=str(round(w_input_hidden[0][j],4)) + " ", color=color, background_color=bg_color)
                text_output.x = 522
                text_output.y = 155 + (j * 120)
                display_group.append(text_output)
                text_output = label.Label(font, text=str(round(w_input_hidden[1][j],4)) + " ", color=color, background_color=bg_color)
                text_output.x = 522
                text_output.y = 185 + (j * 120)
                display_group.append(text_output)
    
    if epoch % 100 == 0:
        start_y += next_y
        text_error = round(total_error, 4)
        str_epoch = f"{epoch:>{len(str(EPOCHS))}}"
        text_total_error = label.Label(font, text=f"Epoche: {str_epoch} / Gesamtfehler: {text_error:.4f}", color=color, background_color=bg_color)
        text_total_error.x = 20
        text_total_error.y = start_y
        display_group.append(text_total_error)
        display.root_group = display_group
        
        # Nach dem Training
        txt_output = "Test nach dem Training\n\n"
        
        for x in inputs:
            h_output = [0, 0] 
            for j in range(2):
                h_output[j] = sigmoid(x[0] * w_input_hidden[0][j] + x[1] * w_input_hidden[1][j] + b_hidden[j])
                text_output = label.Label(font, text=str(round(h_output[j],4)), color=color, background_color=bg_color)
                text_output.x = 740
                text_output.y = 145 + (j * 175)
                display_group.append(text_output)
            o_output = sigmoid(h_output[0] * w_hidden_output[0] + h_output[1] * w_hidden_output[1] + b_output)
            txt_output += (f"input: {x} -> output: {o_output:.4f} -> {'1' if o_output > 0.5 else '0'} (erwartet: {targets[inputs.index(x)][0]})\n")
        
        text_output = label.Label(font, text=txt_output, color=color, background_color=bg_color)
        text_output.x = 400
        text_output.y = 450
        display_group.append(text_output)
        display.root_group = display_group
        # Warten
        time.sleep(SPEED)
    
    # Training abbrechen/beenden
    if total_error < 0.01 or epoch >= EPOCHS:
        break

# Warten: 10 Minuten
time.sleep(600)

# Reset/Neustart
import microcontroller
microcontroller.reset()