Initial commit
This commit is contained in:
Binary file not shown.
@@ -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.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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)
|
||||
@@ -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)))
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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
|
||||
Reference in New Issue
Block a user