Initial commit

This commit is contained in:
2026-04-06 15:59:23 +02:00
commit 60e9dfa5ad
83 changed files with 2167 additions and 0 deletions
Binary file not shown.
Binary file not shown.
+18
View File
@@ -0,0 +1,18 @@
import tkinter as tk
from ui.app_state import AppState
from ui.front_page.front_page import FrontPage
from ui.icons import icons
class App(tk.Tk):
def __init__(self):
super().__init__()
self.app_state = AppState(auto_load=True)
icons.load_icons()
self.title("MNIST Training Center")
self.geometry("1024x720")
self.front_page = FrontPage(self, self.app_state)
self.front_page.pack(expand=1, fill="both")
+23
View File
@@ -0,0 +1,23 @@
import os.path
from data.mnist_loader import MNISTModelData
from neural_net.mnist import MNISTNeuralNet
from neural_net.neural_net import NeuralNet, ModelData
class AppState:
def __init__(self, auto_load=False):
self.trainers = []
if auto_load:
self.neural_net: NeuralNet = MNISTNeuralNet()
data_folder = "/projects/learning/datasets/minst"
self.model_data: ModelData = MNISTModelData(
os.path.join(data_folder, "train-images-idx3-ubyte"),
os.path.join(data_folder, "train-labels-idx1-ubyte"),
os.path.join(data_folder, "t10k-images-idx3-ubyte"),
os.path.join(data_folder, "t10k-labels-idx1-ubyte")
)
self.neural_net.recalculate_accuracy(self.model_data.test_inputs, self.model_data.test_labels)
self.neural_net.recalculate_loss(self.model_data.test_inputs, self.model_data.test_labels)
else:
self.neural_net: NeuralNet = None
self.model_data: ModelData = None
+61
View File
@@ -0,0 +1,61 @@
import tkinter as tk
import numpy as np
from PIL import ImageGrab, ImageTk
from PIL.Image import Resampling
class DigitDrawer(tk.Frame):
def __init__(self, parent, canvas_width, canvas_height):
super().__init__(parent)
self.canvas_width = canvas_width
self.canvas_height = canvas_height
self.brush_size = 3
self.update_ui()
def clear_ui(self):
for widget in self.winfo_children():
widget.destroy()
def update_ui(self):
self.clear_ui()
# Create a Canvas to draw on
self.canvas = tk.Canvas(self, width=self.canvas_width, height=self.canvas_height, bg='white')
self.canvas.pack(padx=10, pady=10)
self.canvas_demo = tk.Canvas(self, width=28, height=28, bg='white')
self.canvas_demo.pack(padx=10, pady=10)
# Clear Button
self.clear_button = tk.Button(self, text="Clear", command=self.clear_canvas)
self.clear_button.pack(expand=True, fill='both')
# Bind mouse events to draw on the canvas
self.canvas.bind("<B1-Motion>", self.paint)
def paint(self, event):
"""Draw on the canvas by creating ovals (circles) at mouse position."""
x1, y1 = (event.x - self.brush_size), (event.y - self.brush_size)
x2, y2 = (event.x + self.brush_size), (event.y + self.brush_size)
self.canvas.create_oval(x1, y1, x2, y2, fill='black', outline='black')
def clear_canvas(self):
"""Clear the canvas to allow the user to draw a new digit."""
self.canvas.delete("all")
def convert_to_array(self):
"""Convert the canvas drawing to a 28x28 grayscale array."""
# Get the canvas's pixel data and save it temporarily
x = self.winfo_rootx() + self.canvas.winfo_x()
y = self.winfo_rooty() + self.canvas.winfo_y()
x1 = x + self.canvas.winfo_width()
y1 = y + self.canvas.winfo_height()
# Capture the canvas area and convert it into a grayscale image using PIL
image = ImageGrab.grab((x, y, x1, y1)).convert("L").resize((28, 28), resample=Resampling.HAMMING)
self.demo_image = ImageTk.PhotoImage(image)
self.canvas_demo.create_image(0, 0, anchor=tk.NW, image=self.demo_image)
image_array = np.asarray(image) / 255.0
print(np.array(image_array).reshape((28, 28)))
flat_array = image_array.flatten()
return flat_array
+21
View File
@@ -0,0 +1,21 @@
import tkinter as tk
from ui.icons.icons import icons
class LabelWithRefresh(tk.Frame):
def __init__(self, parent, initial_text, callback, initial_state=tk.DISABLED):
super().__init__(parent)
self.callback = callback
self._create_ui(initial_text, initial_state)
def _create_ui(self, initial_text, initial_state):
self.refresh_button = tk.Button(self, image=icons["refresh"], state=initial_state, command=self.callback)
self.refresh_button.pack(side=tk.RIGHT, padx=5)
self.label = tk.Label(self, text=initial_text)
self.label.pack(side=tk.RIGHT, padx=5)
def set_state(self, state):
self.refresh_button.config(state=state)
def set_text(self, text):
self.label.config(text=text)
+14
View File
@@ -0,0 +1,14 @@
import tkinter as tk
class NumberSlider(tk.Frame):
def __init__(self, parent, value, from_, to, resolution):
super().__init__(parent)
self.value = value
self.update_ui(from_, to, resolution)
def update_ui(self, from_, to, resolution):
self.entry = tk.Entry(self, textvariable=self.value)
self.entry.pack(side=tk.RIGHT, padx=5)
self.scaler = tk.Scale(self, from_=from_, to=to, length=200, resolution=resolution, showvalue=False, orient=tk.HORIZONTAL, sliderrelief="flat", relief="flat", borderwidth=0, variable=self.value)
self.scaler.set(self.value.get())
self.scaler.pack(side=tk.RIGHT, padx=5)
+27
View File
@@ -0,0 +1,27 @@
import tkinter as tk
from matplotlib.backends.backend_tkagg import FigureCanvasTkAgg
from matplotlib.figure import Figure
from ui.plotters.plotter import Plotter
class PlotFrame(tk.Frame):
def __init__(self, parent, width=None, height=None):
super().__init__(parent, width=width, height=height)
if width is not None or height is not None:
self.pack_propagate(False)
self.figure = self.create_plot_figure()
self.plotter: Plotter = None
def create_plot_figure(self):
figure = Figure(layout="compressed", facecolor=(0,0,0))
# Create a matplotlib canvas to display the plot
canvas = FigureCanvasTkAgg(figure, self)
canvas.draw()
(canvas.get_tk_widget()
.pack(fill=tk.BOTH, expand=False, padx=0, pady=0, ipadx=0, ipady=0))
return figure
def update_data(self, data):
self.plotter.update_plot(data)
Binary file not shown.
+76
View File
@@ -0,0 +1,76 @@
import os
import tkinter as tk
from data.mnist_loader import MNISTModelData
from ui.app_state import AppState
from ui.front_page.sections.model_overview_section import NeuralNetInfo
from ui.front_page.sections.test_model_section import TestModelSection
from ui.front_page.sections.training_section import TrainingSection
class FrontPage(tk.Frame):
def __init__(self, parent, app_state: AppState):
super().__init__(parent)
self.parent = parent
self.app_state = app_state
self.main_frame = None
self.neural_net_info = None
self.model_actions_frame = None
self.start_training_section = None
self.test_model_section = None
self.training_section = None
self.test_model_section = None
self.create_ui()
def create_ui(self):
(tk.Label(self, text="Welcome to MNIST Learning Center", font=("Arial", 16))
.pack(side=tk.TOP, fill=tk.BOTH, expand=False, padx=5))
self.main_frame = tk.Frame(self)
self.main_frame.pack(fill=tk.BOTH, expand=True)
self.neural_net_info = NeuralNetInfo(self.main_frame, self.app_state, self.on_model_loaded, self.on_data_loaded)
self.neural_net_info.pack(side=tk.TOP, fill=tk.BOTH, expand=True, padx=5)
self.load_model_actions_frame()
def update(self):
if self.neural_net_info is not None:
self.neural_net_info.update()
self.load_model_actions_frame()
def load_model_actions_frame(self):
if self.model_actions_frame is None and self.app_state.neural_net is not None and self.app_state.model_data is not None:
self.model_actions_frame = tk.Frame(self.main_frame)
self.model_actions_frame.pack(side=tk.BOTTOM, fill=tk.BOTH, expand=True, padx=5)
self.training_section = TrainingSection(self.model_actions_frame, self.app_state, self.after_training)
self.training_section.pack(side=tk.LEFT, fill=tk.BOTH, expand=True, padx=5)
self.test_model_section = TestModelSection(self.model_actions_frame, self.app_state)
self.test_model_section.pack(side=tk.RIGHT, fill=tk.BOTH, expand=True, padx=5)
else:
if self.test_model_section is not None:
self.test_model_section.update()
if self.training_section is not None:
self.training_section.update()
def on_data_loaded(self):
print("Data loaded")
self.update()
def on_model_loaded(self):
print("Model loaded")
self.update()
def load_training_data(self):
data_folder = "/projects/learning/datasets/minst"
self.app_state.model_data = MNISTModelData(
os.path.join(data_folder, "train-images-idx3-ubyte"),
os.path.join(data_folder, "train-labels-idx1-ubyte"),
os.path.join(data_folder, "t10k-images-idx3-ubyte"),
os.path.join(data_folder, "t10k-labels-idx1-ubyte")
)
self.update()
def after_training(self):
self.update()
Binary file not shown.
+41
View File
@@ -0,0 +1,41 @@
from abc import ABC
from matplotlib.figure import Figure
from neural_net.epoch import Epoch
from neural_net.neural_net import NeuralNet
from ui.components.plot_figure import PlotFrame
from ui.plotters.plotter import Plotter
class GradientsPlot(PlotFrame):
def __init__(self, parent, neural_net: NeuralNet):
super().__init__(parent)
self.plotter = GradientsPlotter(self.figure, neural_net)
class GradientsPlotter(Plotter, ABC):
def __init__(self, figure: Figure, neural_net: NeuralNet):
super().__init__(figure)
self.neural_net = neural_net
self.axes = figure.subplots(1, 2)
def reset_plot(self):
self.axes[0].clear()
self.axes[0].set_xlabel('Neuron Index')
self.axes[0].set_ylabel('Input Index')
self.axes[1].clear()
self.axes[1].set_xlabel('Output Neuron Index')
self.axes[1].set_ylabel('Hidden Neuron Index')
def plot(self, data: Epoch):
gradients_layer1 = data.layer_dl_gradients[1][-1]
self.axes[0].imshow(gradients_layer1, cmap='coolwarm', aspect='auto')
gradients_layer2 = data.layer_dl_gradients[0][-1]
self.axes[1].imshow(gradients_layer2, cmap='coolwarm', aspect='auto')
def plot_gradients_histogram(self, current_epoch: Epoch):
gradients_layer1 = current_epoch.layer_dl_gradients[1][-1]
self.axes[0].hist(gradients_layer1.flatten(), bins=50, color='blue', alpha=0.7)
gradients_layer2 = current_epoch.layer_dl_gradients[0][-1]
self.axes[1].hist(gradients_layer2.flatten(), bins=50, color='green', alpha=0.7)
+39
View File
@@ -0,0 +1,39 @@
from ui.components.plot_figure import PlotFrame
import math
from abc import ABC
from matplotlib.figure import Figure
from neural_net.activation_layers.activation_layer import ActivationLayer
from neural_net.neural_net import NeuralNet
from ui.plotters.plotter import Plotter
from utils.matplotlib.utils import mpl_matshow
class LayerWeightsPlot(PlotFrame):
def __init__(self, parent, neural_net: NeuralNet, layer: ActivationLayer, rows, cols):
super().__init__(parent)
self.plotter = LayerWeightsPlotter(self.figure, neural_net, layer, rows, cols)
class LayerWeightsPlotter(Plotter, ABC):
def __init__(self, figure: Figure, neural_net: NeuralNet, layer: ActivationLayer, rows, columns):
super().__init__(figure)
self.neural_net = neural_net
self.layer = layer
self.axes = figure.subplots(nrows=rows, ncols=columns, squeeze=True,
gridspec_kw={'wspace': 0.05, 'hspace': 0.05})
def reset_plot(self):
for axes in self.axes:
for ax in axes:
ax.clear()
def plot(self, data):
weights = self.layer.weights.T
n_neurons = weights.shape[0]
n_pixels = weights.shape[1]
for i in range(n_neurons):
row = i // self.axes.shape[1]
col = i % self.axes.shape[1]
mpl_matshow(self.axes[row, col], weights[i], int(math.sqrt(n_pixels)))
+40
View File
@@ -0,0 +1,40 @@
from neural_net.trainer import NeuralNetTrainer
from ui.components.plot_figure import PlotFrame
from abc import ABC
from matplotlib.figure import Figure
from neural_net.neural_net import NeuralNet
from ui.plotters.plotter import Plotter
class LossPlot(PlotFrame):
def __init__(self, parent, neural_net: NeuralNet, trainer: NeuralNetTrainer):
super().__init__(parent)
self.plotter = LossPlotter(self.figure, neural_net, trainer)
class LossPlotter(Plotter, ABC):
def __init__(self, figure: Figure, neural_net: NeuralNet, trainer: NeuralNetTrainer):
super().__init__(figure)
self.neural_net = neural_net
self.trainer = trainer
self.axes = figure.add_subplot()
def reset_plot(self):
self.axes.clear()
self.axes.set_title('Loss')
self.axes.set_ylabel("Loss")
self.axes.set_xlabel("Epoch")
def plot(self, data):
losses = []
for epoch in self.trainer.epoch_history:
if epoch.finished:
losses.append(epoch.loss)
self.axes.plot(losses, marker='o', label=f"Loss")
for idx, loss in enumerate(losses):
self.axes.annotate(f"{loss:.4f}", xy=(idx, loss), rotation=45)
self.axes.legend()
self.axes.grid(True)
+37
View File
@@ -0,0 +1,37 @@
from ui.components.plot_figure import PlotFrame
from abc import ABC
from matplotlib.figure import Figure
from ui.plotters.plotter import Plotter
class PredictionsPlot(PlotFrame):
def __init__(self, parent):
super().__init__(parent, height=32)
self.plotter = PredictionsPlotter(self.figure)
class PredictionsPlotter(Plotter, ABC):
def __init__(self, figure: Figure):
super().__init__(figure)
self.axes = figure.add_subplot()
self.clean_axes()
def plot(self, data):
self.axes.imshow(data, cmap='coolwarm', aspect='auto')
for idx in range(10):
self.axes.annotate(f"{idx}", xy=(idx - 0.2, 0.2))
self.clean_axes()
def clean_axes(self):
# Remove axis ticks, labels, and spines
self.axes.set_xticks([]) # Remove x-ticks
self.axes.set_yticks([]) # Remove y-ticks
self.axes.spines['top'].set_visible(False)
self.axes.spines['bottom'].set_visible(False)
self.axes.spines['left'].set_visible(False)
self.axes.spines['right'].set_visible(False)
self.axes.set_facecolor((0, 0, 0))
def reset_plot(self):
self.axes.clear()
@@ -0,0 +1,75 @@
import os
import tkinter as tk
from data.mnist_loader import MNISTModelData
from neural_net.mnist import MNISTNeuralNet
from ui.app_state import AppState
from ui.front_page.sections.neural_net_info_widget import NeuralNetInfoWidget
class NeuralNetInfo(tk.LabelFrame):
def __init__(self, parent, app_state: AppState, on_load_model, on_load_data):
super().__init__(parent, text="Model overview")
self.app_state = app_state
self.cb_on_load_model = on_load_model
self.cb_on_load_data = on_load_data
self.create_ui()
def create_ui(self):
# Option to load model (could be a file dialog or dropdown in future)
self.load_model_button = tk.Button(self, text="Load model", command=self.on_load_model)
self.load_model_button.pack(padx=5, pady=5, side=tk.TOP)
if self.app_state.neural_net is None:
self.model_status = tk.Label(self, text="No model loaded")
self.model_status.pack(padx=5, pady=5)
else:
self.load_model_button.config(text="Reload model")
load_data_button = tk.Button(self, text="Load data", command=self.on_load_data)
load_data_button.pack(padx=5, pady=5, side=tk.TOP)
if self.app_state.model_data is None:
self.data_status = tk.Label(self, text="No data loaded")
self.data_status.pack(padx=5, pady=5)
else:
load_data_button.config(text="Reload data")
self.neural_net_info = NeuralNetInfoWidget(self, self.app_state)
self.neural_net_info.pack(padx=5, pady=5)
def update(self):
if self.app_state.neural_net is None and self.model_status is None:
self.model_status = tk.Label(self, text="No model loaded")
self.model_status.pack(padx=5, pady=5)
load_data_button = tk.Button(self, text="Load data", command=self.on_load_data)
load_data_button.pack(padx=5, pady=5, side=tk.TOP)
if self.app_state.model_data is None:
self.data_status = tk.Label(self, text="No data loaded")
self.data_status.pack(padx=5, pady=5)
else:
load_data_button.config(text="Reload data")
self.neural_net_info = NeuralNetInfoWidget(self, self.app_state)
self.neural_net_info.pack(padx=5, pady=5)
def on_load_data(self):
data_folder = "/projects/learning/datasets/minst"
self.app_state.model_data = MNISTModelData(
os.path.join(data_folder, "train-images-idx3-ubyte"),
os.path.join(data_folder, "train-labels-idx1-ubyte"),
os.path.join(data_folder, "t10k-images-idx3-ubyte"),
os.path.join(data_folder, "t10k-labels-idx1-ubyte")
)
if self.app_state.neural_net is not None:
self.app_state.neural_net.recalculate_loss(self.app_state.model_data.test_inputs, self.app_state.model_data.test_labels)
self.app_state.neural_net.recalculate_accuracy(self.app_state.model_data.test_inputs, self.app_state.model_data.test_labels)
if self.cb_on_load_data is not None:
self.cb_on_load_data()
def on_load_model(self):
self.app_state.neural_net = MNISTNeuralNet()
if self.app_state.model_data is not None:
self.app_state.neural_net.recalculate_loss(self.app_state.model_data.test_inputs, self.app_state.model_data.test_labels)
self.app_state.neural_net.recalculate_accuracy(self.app_state.model_data.test_inputs, self.app_state.model_data.test_labels)
if self.cb_on_load_model is not None:
self.cb_on_load_model()
@@ -0,0 +1,53 @@
import tkinter as tk
from ui.app_state import AppState
from ui.components.label_with_refresh import LabelWithRefresh
class NeuralNetInfoWidget(tk.Frame):
def __init__(self, parent, app_state: AppState):
super().__init__(parent)
self.app_state = app_state
self.update_ui()
def clear_ui(self):
for widget in self.winfo_children():
widget.destroy()
def update_ui(self):
self.clear_ui()
row = 0
if self.app_state.neural_net is not None:
for layer in self.app_state.neural_net.layers:
(tk.Label(self, text=f"{layer.type} {layer.index}")
.grid(column=0, row=row, padx=10, pady=5, sticky='w'))
tk.Label(self, text=f"{layer.input_dim} -> {layer.output_dim} neurons").grid(column=1, row=row, padx=10, pady=5, sticky='e')
row += 1
button_state = tk.DISABLED
if self.app_state.model_data is not None:
button_state = tk.NORMAL
tk.Label(self, text="Accuracy:").grid(column=0, row=row, padx=10, pady=5, sticky='w')
last_accuracy = "NA"
if self.app_state.neural_net.last_accuracy is not None:
last_accuracy = f"{self.app_state.neural_net.last_accuracy * 100:.2f}%"
self.accuracy_label = LabelWithRefresh(self, last_accuracy, callback=self.recalculate_accuracy, initial_state=button_state)
self.accuracy_label.grid(column=1, row=row, padx=10, pady=5, sticky='e')
row += 1
tk.Label(self, text="Current Loss:").grid(column=0, row=row, padx=10, pady=5, sticky='w')
last_loss = "NA"
if self.app_state.neural_net.last_loss is not None:
last_loss = f"{self.app_state.neural_net.last_loss:.4f}"
self.loss_label = LabelWithRefresh(self, last_loss, callback=self.recalculate_loss, initial_state=button_state)
self.loss_label.grid(column=1, row=row, padx=10, pady=5, sticky='e')
row += 1
def recalculate_accuracy(self):
self.app_state.neural_net.recalculate_accuracy(self.app_state.model_data.test_inputs, self.app_state.model_data.test_labels)
self.update_ui()
def recalculate_loss(self):
self.app_state.neural_net.recalculate_loss(self.app_state.model_data.test_inputs, self.app_state.model_data.test_labels)
self.update_ui()
@@ -0,0 +1,42 @@
import tkinter as tk
from ui.app_state import AppState
from ui.components.digit_drawer import DigitDrawer
from ui.front_page.plots.predictions import PredictionsPlot
class TestModelSection(tk.LabelFrame):
def __init__(self, parent, app_state: AppState):
super().__init__(parent, text="Model testing")
self.app_state = app_state
self.update_ui()
def clear_ui(self):
for widget in self.winfo_children():
widget.destroy()
def update_ui(self):
self.clear_ui()
self.digit_drawer = DigitDrawer(self, 100, 100)
self.digit_drawer.pack(fill=tk.BOTH, expand=True)
# Predict Button (converts drawing to 28x28 and shows the array)
self.predict_button = tk.Button(self, text="Predict", command=self.predict_number)
self.predict_button.pack(fill=tk.BOTH, expand=True)
frame_prediction = tk.Frame(self, height=200)
frame_prediction.pack(fill=tk.BOTH, expand=True)
(tk.Label(frame_prediction, text="Prediction: ")
.pack(side=tk.LEFT))
self.lbl_prediction = tk.Label(frame_prediction, text="/")
self.lbl_prediction.pack(side=tk.LEFT)
self.prediction_plot = PredictionsPlot(self)
self.prediction_plot.pack(side=tk.BOTTOM, anchor=tk.S, fill=tk.X, expand=True)
def predict_number(self):
inputs = self.digit_drawer.convert_to_array()
raw_predictions, predictions = self.app_state.neural_net.predict([inputs])
print(predictions)
self.lbl_prediction.config(text=f"{predictions[0]}")
self.prediction_plot.update_data(raw_predictions)
@@ -0,0 +1,43 @@
import tkinter as tk
from neural_net.epoch import Epoch
class EpochInformation(tk.LabelFrame):
def __init__(self, parent, epoch: Epoch):
super().__init__(parent, text="Last epoch info")
self.epoch = epoch
self.lbl_epoch_training_time = None
self.lbl_last_loss = None
self.create_ui()
def create_ui(self):
row = 0
tk.Label(self, text="Duration:", anchor=tk.W).grid(column=0, row=row,
sticky=tk.E,
padx=(10, 20), pady=5)
self.lbl_epoch_training_time = tk.Label(self, text=f"{self.epoch.duration:.2f}sec")
self.lbl_epoch_training_time.grid(column=1, row=row, sticky=tk.E, padx=10, pady=5)
row += 1
tk.Label(self, text="Loss value:", anchor=tk.W).grid(column=0, row=row,
sticky=tk.E, padx=(10, 20),
pady=5)
self.lbl_last_loss = tk.Label(self, text=f"{self.epoch.loss:.4f}")
self.lbl_last_loss.grid(column=1, row=row, sticky=tk.E, padx=10, pady=5)
row += 1
tk.Label(self, text="Learning rate:", anchor=tk.W).grid(column=0, row=row,
sticky=tk.E, padx=(10, 20),
pady=5)
self.lbl_learning_rate = tk.Label(self, text=f"{self.epoch.learning_rate:.4f}")
self.lbl_learning_rate.grid(column=1, row=row, sticky=tk.E, padx=10, pady=5)
def update(self):
print(f"Updating training data for epoch {self.epoch}")
self.lbl_epoch_training_time.config(text=f"{self.epoch.duration:.2f}sec")
self.lbl_last_loss.config(text=f"{self.epoch.loss:.4f}")
self.lbl_learning_rate.config(text=f"{self.epoch.learning_rate:.4f}")
def set_epoch(self, epoch: Epoch):
self.epoch = epoch
self.update()
@@ -0,0 +1,79 @@
import threading
import tkinter as tk
from neural_net.trainer import NeuralNetTrainer
from ui.app_state import AppState
from ui.components.number_slider import NumberSlider
from ui.training_page.training_page import EpochInformation
class TrainingSection(tk.LabelFrame):
def __init__(self, parent, app_state: AppState, on_update_neural_net_info):
super().__init__(parent, text="Model training")
self.app_state = app_state
self.on_update_neural_net_info = on_update_neural_net_info
self.batch_size = tk.IntVar()
self.batch_size.set(1000)
self.batch_size_slider = None
self.learning_rate = tk.DoubleVar()
self.learning_rate.set(0.0001)
self.learning_rate_slider = None
self.btn_start_stop = None
self.stop_button = None
self.training_information_container: EpochInformation = None
self.trainer: NeuralNetTrainer = NeuralNetTrainer(self.app_state.neural_net, self.app_state.model_data,
self.learning_rate.get(), self.batch_size.get())
self.create_ui()
def create_ui(self):
tk.Label(self, text="Batch size:").grid(column=0, row=0, padx=10, pady=5, sticky='w')
self.batch_size_slider = NumberSlider(self, self.batch_size, from_=100, to=10000, resolution=1)
self.batch_size_slider.grid(column=1, row=0, padx=10, pady=5, sticky='w')
tk.Label(self, text="Learning rate:").grid(column=0, row=1, padx=10, pady=5, sticky='w')
self.learning_rate_slider = NumberSlider(self, self.learning_rate, from_=0.0001, to=0.1, resolution=0.0001)
self.learning_rate_slider.grid(column=1, row=1, padx=10, pady=5, sticky='w')
self.btn_prev_epoch = tk.Button(self, text="<<", command=self.on_prev_epoch)
self.btn_prev_epoch.grid(column=0, row=2, padx=10, pady=10, sticky='w')
self.btn_start_stop = tk.Button(self, text="Start", command=self.toggle_state)
self.btn_start_stop.grid(column=1, row=2, padx=10, pady=10, sticky='w')
self.btn_next_epoch = tk.Button(self, text=">>", command=self.on_next_epoch)
self.btn_next_epoch.grid(column=2, row=2, padx=10, pady=10, sticky='w')
def update(self):
if self.trainer.is_running:
if self.training_information_container is None:
self.training_information_container = EpochInformation(self, self.trainer.epoch_history[-1])
self.training_information_container.grid(column=0, row=5, padx=10, pady=10, sticky='e')
self.btn_start_stop.config(text="Stop")
else:
print("Setting the epoch")
self.training_information_container.set_epoch(self.trainer.epoch_history[-1])
else:
self.btn_start_stop.config(text="Start")
def toggle_state(self):
if self.trainer.is_running:
self.trainer.stop()
else:
self.thread = threading.Thread(target=self.trainer.start, args=(self.on_epoch_finish, self.on_update_neural_net_info))
self.thread.start()
# self.trainer.start(self.on_epoch_finish, self.on_update_neural_net_info)
self.update()
def start(self):
self.thread = threading.Thread(target=self.trainer.start)
self.thread.start()
# self.trainer.start(on_epoch_finished=self.update_training_data)
def on_epoch_finish(self, epoch):
print("Updating the epoch")
self.update()
def on_prev_epoch(self):
pass
def on_next_epoch(self):
pass
Binary file not shown.
+14
View File
@@ -0,0 +1,14 @@
import tkinter as tk
from PIL import Image, ImageTk
from PIL.Image import Resampling
icons = {}
def _load_icon(path, size):
img = Image.open(path)
img = img.resize(size, resample=Resampling.HAMMING)
return ImageTk.PhotoImage(img)
def load_icons():
icons["refresh"] = _load_icon("ui/icons/refresh.png", (24, 24))
Binary file not shown.

After

Width:  |  Height:  |  Size: 35 KiB

Binary file not shown.
+7
View File
@@ -0,0 +1,7 @@
from abc import ABC
from matplotlib.figure import Figure
from neural_net.epoch import Epoch
from neural_net.neural_net import NeuralNet
from ui.plotters.plotter import Plotter
+1
View File
@@ -0,0 +1 @@
+30
View File
@@ -0,0 +1,30 @@
from abc import abstractmethod
from matplotlib.figure import Figure
from neural_net.epoch import Epoch
class Plotter:
def __init__(self, figure: Figure):
self.figure = figure
def initialize_plots(self):
self.figure.show()
@abstractmethod
def update_plot(self, data):
self.reset_plot()
self.plot(data)
self.figure.canvas.draw()
self.figure.canvas.flush_events()
@abstractmethod
def reset_plot(self):
pass
@abstractmethod
def plot(self, current_epoch: Epoch):
pass
View File
View File
+103
View File
@@ -0,0 +1,103 @@
import threading
import tkinter as tk
from matplotlib.backends.backend_tkagg import FigureCanvasTkAgg
from matplotlib.figure import Figure
from neural_net.epoch import Epoch
from neural_net.trainer import NeuralNetTrainer
from ui.app_state import AppState
from ui.front_page.plots.gradients import GradientsPlot
from ui.front_page.plots.layer_weights import LayerWeightsPlot
from ui.front_page.plots.loss import LossPlot
from ui.front_page.sections.training_information import EpochInformation
class TrainingPage(tk.Frame):
def __init__(self, parent, app_state: AppState, on_training_finished=None):
super().__init__(parent)
self.app_state = app_state
self.on_training_finished = on_training_finished
self.trainer: NeuralNetTrainer = None
# trainer = NeuralNetTrainer(self.app_state.neural_net, self.app_state.model_data, learning_rate, nr_epochs)
# self.app_state.trainers.append(trainer)
# self.trainer = trainer
self.create_ui()
def start(self, learning_rate, nr_epochs, batch_size, callback=None):
if self.trainer is not None:
self.trainer.stop()
self.trainer = NeuralNetTrainer(self.app_state.neural_net, self.app_state.model_data,
learning_rate=learning_rate, nr_epochs=nr_epochs, batch_size=batch_size,
on_epoch_callback=self.update_training_data,
on_finished_callback=self.on_training_finished)
self.trainer.on_epoch_callback = self.update_training_data
self.thread = threading.Thread(target=self.trainer.start)
self.thread.start()
# self.trainer.start(on_epoch_finished=self.update_training_data)
if callback is not None:
callback()
def update_training_data(self, training_run, data: Epoch):
print(f"Updating training data {data.epoch}")
if self.trainer.is_running:
self.training_information_container.update_training_data(training_run, data)
self.loss_plot.update_training_data(training_run, data)
if data.epoch % 5 == 0:
self.gradients_plot.update_training_data(training_run, data)
self.layer0_weights_plot.update_training_data(training_run, data)
self.layer1_weights_plot.update_training_data(training_run, data)
def create_ui(self):
# Training center
self.training_information_container = EpochInformation(self, self.app_state.neural_net, self.trainer)
self.training_information_container.pack(side=tk.TOP, fill=tk.X, expand=False, pady=10, padx=10, ipady=10,
ipadx=10)
actions_frame = tk.Frame(self)
actions_frame.pack(side=tk.TOP, fill=tk.X, expand=False, pady=10, padx=10, ipady=10, ipadx=10)
btn_text = "Pause"
self.btn_toggle_pause = tk.Button(actions_frame, text=btn_text, command=self.toggle_state)
self.btn_toggle_pause.pack(side=tk.LEFT)
btn_stop = tk.Button(actions_frame, text="Stop", command=self.trainer.stop)
btn_stop.pack(side=tk.LEFT)
# Plot tabs
plot_tab_control = tk.Notebook(self)
plot_tab_control.pack(side=tk.BOTTOM, fill=tk.BOTH, expand=True, pady=0, padx=0, ipady=0, ipadx=0)
self.loss_plot = LossPlot(plot_tab_control, self.app_state.neural_net)
plot_tab_control.add(self.loss_plot, text="Loss Function")
self.gradients_plot = GradientsPlot(plot_tab_control, self.app_state.neural_net)
plot_tab_control.add(self.gradients_plot, text="Gradients")
self.layer0_weights_plot = LayerWeightsPlot(plot_tab_control, self.app_state.neural_net,
self.app_state.neural_net.layers[0],
11, 11)
plot_tab_control.add(self.layer0_weights_plot, text="Weights layer 0")
self.layer1_weights_plot = LayerWeightsPlot(plot_tab_control, self.app_state.neural_net,
self.app_state.neural_net.layers[1],
2, 5)
plot_tab_control.add(self.layer1_weights_plot, text="Weights layer 1")
def toggle_state(self):
self.trainer.toggle_state()
if self.trainer.training_paused:
self.btn_toggle_pause.config(text="Resume")
else:
self.btn_toggle_pause.config(text="Pause")
@staticmethod
def create_plot_figure(tab):
figure = Figure()
# Create a matplotlib canvas to display the plot
canvas = FigureCanvasTkAgg(figure, tab)
canvas.draw()
canvas.get_tk_widget().pack(fill=tk.BOTH, expand=True)
return figure
View File