|
|
@@ -1,11 +1,12 @@
|
|
|
-from dash import html, dcc
|
|
|
+from dash import html, dcc, callback_context, dash_table
|
|
|
import dash_bootstrap_components as dbc
|
|
|
-from utils.file_utils import list_directories_in_directory, list_files_in_directory, get_file_metadata, get_training_time, get_file_size, trainer_classes
|
|
|
-import json
|
|
|
+from dash.dependencies import Input, Output, State
|
|
|
+from utils.file_utils import list_directories_in_directory, list_files_in_directory, get_file_metadata, get_training_time, get_file_size, trainer_classes, get_Size_String
|
|
|
import os
|
|
|
+import json
|
|
|
+import plotly.graph_objs as go
|
|
|
|
|
|
TRAINING_HOTWORD_DIR = './training/hotword/'
|
|
|
-MODEL_DIR = './model/hotword/'
|
|
|
|
|
|
def HotwordsLayout():
|
|
|
hotword_models = list_directories_in_directory(TRAINING_HOTWORD_DIR)
|
|
|
@@ -36,27 +37,107 @@ def generate_model_table(models):
|
|
|
return dbc.Table(header + [html.Tbody(rows)], bordered=True, hover=True)
|
|
|
|
|
|
def ManageHotwordLayout(model_name):
|
|
|
+ page_size = 10
|
|
|
training_data = list_files_in_directory(f"{TRAINING_HOTWORD_DIR}/{model_name}/", ".wav")
|
|
|
- training_data_table = generate_training_data_table(model_name, training_data)
|
|
|
+ training_data_table, row_count = generate_training_data_table(model_name, training_data, 1, page_size)
|
|
|
+ trainer_options = [{'label': name, 'value': name} for name in trainer_classes.keys()]
|
|
|
|
|
|
- model_info = generate_model_info(model_name)
|
|
|
+ training_time = get_training_time(model_name)
|
|
|
+
|
|
|
+ tab_trainer = []
|
|
|
+ for trainer_info in trainer_options:
|
|
|
+ infoFile = f"./model/hotword/{model_name}/info_{trainer_info['value']}.json"
|
|
|
+ if os.path.exists(infoFile):
|
|
|
+ with open(infoFile, 'r') as file:
|
|
|
+ info = json.load(file)
|
|
|
+ epochs = '-'
|
|
|
+ batch_size = '-'
|
|
|
+ dropout_rate = '-'
|
|
|
+ training_method = '-'
|
|
|
+ duration = '-'
|
|
|
+ testAccuracy = '-'
|
|
|
+ accu = []
|
|
|
+ loss = []
|
|
|
+ size = '-'
|
|
|
+ if 'epochs' in info:
|
|
|
+ epochs = info['epochs']
|
|
|
+ if 'batch_size' in info:
|
|
|
+ batch_size = info['batch_size']
|
|
|
+ if 'dropout_rate' in info:
|
|
|
+ dropout_rate = info['dropout_rate']
|
|
|
+ if 'trainer_script' in info:
|
|
|
+ training_method = info['trainer_script']
|
|
|
+ if 'accuracy' in info:
|
|
|
+ accu = info['accuracy']
|
|
|
+ if 'loss' in info:
|
|
|
+ loss = info['loss']
|
|
|
+ if 'training_time' in info:
|
|
|
+ duration = info['training_time']
|
|
|
+ if 'test_accuracy' in info:
|
|
|
+ testAccuracy = info['test_accuracy']
|
|
|
+ if 'model_size' in info:
|
|
|
+ size = info['model_size']
|
|
|
+ chart_fig1 = go.Figure(data=[go.Scatter(y=accu)])
|
|
|
+ chart_fig2 = go.Figure(data=[go.Scatter(y=loss)])
|
|
|
+ entries = [
|
|
|
+ dcc.Graph(figure=chart_fig1),
|
|
|
+ dcc.Graph(figure=chart_fig2),
|
|
|
+ html.P("Epochen: "+str(epochs), className="card-text"),
|
|
|
+ html.P("Batch Size: "+str(batch_size), className="card-text"),
|
|
|
+ html.P("Dropout Rate: "+str(dropout_rate), className="card-text"),
|
|
|
+ html.P("Methode: "+str(training_method), className="card-text"),
|
|
|
+ html.P("Duration: "+str(duration), className="card-text"),
|
|
|
+ html.P("Test Accuracy: "+str(testAccuracy), className="card-text"),
|
|
|
+ html.P("Size: "+get_Size_String(size*1024), className="card-text"),
|
|
|
+ ]
|
|
|
+ tab_trainer.append(
|
|
|
+ dbc.Tab(entries
|
|
|
+ , label=trainer_info['label'], tab_id=trainer_info['value'])
|
|
|
+ )
|
|
|
+ else:
|
|
|
+ tab_trainer.append(
|
|
|
+ dbc.Tab("",label=trainer_info['label'], tab_id=trainer_info['value'], disabled=True)
|
|
|
+ )
|
|
|
+
|
|
|
+ # Pagination logic
|
|
|
+ num_pages = (row_count // page_size) + (1 if row_count % page_size > 0 else 0)
|
|
|
+ pagination = dbc.Pagination(id='data-table-pagination', max_value=num_pages, active_page=1)
|
|
|
|
|
|
return html.Div([
|
|
|
html.H2(f"Manage Hotword Model: {model_name}"),
|
|
|
- model_info,
|
|
|
+ dbc.Card([dbc.CardHeader("Model Infos"),dbc.CardBody(dbc.Tabs(tab_trainer))]),
|
|
|
dbc.Card([
|
|
|
dbc.CardHeader("Training Controls"),
|
|
|
dbc.CardBody([
|
|
|
- dbc.ButtonGroup([
|
|
|
- dbc.Button("Start Training", id="start-training-hotword", color="success", className="mb-3"),
|
|
|
- dbc.DropdownMenu(
|
|
|
- [dbc.DropdownMenuItem(name, id={'type': 'select-trainer', 'index': name}) for name in trainer_classes.keys()],
|
|
|
- label="Select Trainer", bs_size="lg", direction="down"
|
|
|
- ),
|
|
|
- dbc.Button("Optimize Training", id="optimize-training-hotword", color="primary", className="mb-3"),
|
|
|
- dbc.Button("Stop Training", id="stop-training-hotword", color="danger", className="mb-3")
|
|
|
- ]),
|
|
|
+ #dbc.Row([
|
|
|
+ #dbc.Col(dbc.Button("Start Training", id="start-training-hotword", color="success", className="mb-3")),
|
|
|
+ #dbc.Col(dbc.Button("Optimize Training", id="optimize-training-hotword", color="primary", className="mb-3 ml-3")),
|
|
|
+ #dbc.Col(dbc.Button("Stop Training", id="stop-training-hotword", color="danger", className="mb-3 ml-3")),
|
|
|
+ dbc.ButtonGroup([
|
|
|
+ dbc.DropdownMenu(
|
|
|
+ label="Start training",
|
|
|
+ children=[
|
|
|
+ dbc.DropdownMenuItem("Item 1"),
|
|
|
+ dbc.DropdownMenuItem("Item 2"),
|
|
|
+ dbc.DropdownMenuItem("Item 3"),
|
|
|
+ ],
|
|
|
+ ),
|
|
|
+ dbc.Button("Start Training", id="start-training-hotword", color="success", className="mb-3"),
|
|
|
+ dbc.Button("Optimize Training", id="optimize-training-hotword", color="primary", className="mb-3 ml-3"),
|
|
|
+ dbc.Button("Stop Training", id="stop-training-hotword", color="danger", className="mb-3 ml-3"),
|
|
|
+ ]),
|
|
|
+ #]),
|
|
|
dbc.Row([
|
|
|
+ dbc.Col(
|
|
|
+ dcc.Dropdown(
|
|
|
+ id="training-method-dropdown",
|
|
|
+ options=trainer_options,
|
|
|
+ value=trainer_options[0]["value"], # default value
|
|
|
+ clearable=False,
|
|
|
+ className="mb-3",
|
|
|
+ placeholder="Select Training Method",
|
|
|
+ )
|
|
|
+ ),
|
|
|
dbc.Col(
|
|
|
dcc.Input(
|
|
|
id="dropout-rate",
|
|
|
@@ -70,26 +151,35 @@ def ManageHotwordLayout(model_name):
|
|
|
)
|
|
|
),
|
|
|
dbc.Col(
|
|
|
- dcc.Input(
|
|
|
- id="num-epochs",
|
|
|
- type="number",
|
|
|
- value=10,
|
|
|
- step=1,
|
|
|
- min=1,
|
|
|
- className="mb-3",
|
|
|
- placeholder="Epochs"
|
|
|
- )
|
|
|
+ dcc.Slider(0, 10000, value=3000, id="num-epochs",tooltip={"placement": "bottom", "always_visible": False})
|
|
|
+ #dcc.Input(
|
|
|
+ # id="num-epochs",
|
|
|
+ # type="number",
|
|
|
+ # value=10,
|
|
|
+ # step=1,
|
|
|
+ # min=1,
|
|
|
+ # className="mb-3",
|
|
|
+ # placeholder="Epochs"
|
|
|
+ #)
|
|
|
),
|
|
|
dbc.Col(
|
|
|
- dcc.Input(
|
|
|
- id="batch-size",
|
|
|
- type="number",
|
|
|
- value=32,
|
|
|
- step=1,
|
|
|
- min=1,
|
|
|
- className="mb-3",
|
|
|
- placeholder="Batch Size"
|
|
|
- )
|
|
|
+ dcc.Slider(0, 256, marks={
|
|
|
+ 8: '8',
|
|
|
+ 16: '16',
|
|
|
+ 32: '32',
|
|
|
+ 64: '64',
|
|
|
+ 128: '128',
|
|
|
+ 256: '256',
|
|
|
+ }, value=32, id="batch-size", step=None)
|
|
|
+ #dcc.Input(
|
|
|
+ # id="batch-size",
|
|
|
+ # type="number",
|
|
|
+ # value=32,
|
|
|
+ # step=1,
|
|
|
+ # min=1,
|
|
|
+ # className="mb-3",
|
|
|
+ # placeholder="Batch Size"
|
|
|
+ #)
|
|
|
)
|
|
|
]),
|
|
|
html.Div(id="training-status", className='mt-3'),
|
|
|
@@ -97,7 +187,7 @@ def ManageHotwordLayout(model_name):
|
|
|
dbc.Card([
|
|
|
dbc.CardHeader("Training Progress"),
|
|
|
dbc.CardBody([
|
|
|
- dcc.Graph(id="training-chart"),
|
|
|
+ dcc.Graph(id="training-chart", figure=[]),
|
|
|
dbc.Progress(id="training-progress", striped=True, animated=True, style={"height": "20px"}),
|
|
|
])
|
|
|
]),
|
|
|
@@ -132,49 +222,12 @@ def ManageHotwordLayout(model_name):
|
|
|
dbc.Card([
|
|
|
dbc.CardHeader("Training Data"),
|
|
|
dbc.CardBody([
|
|
|
- training_data_table
|
|
|
+ training_data_table, pagination
|
|
|
])
|
|
|
])
|
|
|
])
|
|
|
|
|
|
-def generate_model_info(model_name):
|
|
|
- model_dir = f"./model/hotword/{model_name}/"
|
|
|
- model_info = []
|
|
|
- for extension in ["h5", "keras", "pth"]:
|
|
|
- model_path = model_dir + f"hotword_model.{extension}"
|
|
|
- if os.path.exists(model_path):
|
|
|
- json_path = model_dir + f"info_{extension}.json"
|
|
|
- try:
|
|
|
- with open(json_path, "r") as f:
|
|
|
- info = json.load(f)
|
|
|
- model_info.append(html.Div([
|
|
|
- html.H4(f"Model ({extension.upper()})"),
|
|
|
- html.P(f"Training Time: {info['training_time']} seconds"),
|
|
|
- html.P(f"Model Size: {info['model_size']} KB"),
|
|
|
- html.P(f"Trainer Script: {info['trainer_script']}"),
|
|
|
- html.P(f"Epochs: {info['epochs']}"),
|
|
|
- html.P(f"Dropout Rate: {info['dropout_rate']}"),
|
|
|
- html.P(f"Batch Size: {info['batch_size']}"),
|
|
|
- html.P(f"Total Training Data: {info['total_training_data']}"),
|
|
|
- html.P(f"Training Directories: {', '.join(info['training_directories'])}"),
|
|
|
- html.P(f"Accuracy: {info['accuracy'][-1] if info['accuracy'] else 'N/A'}"),
|
|
|
- html.P(f"Loss: {info['loss'][-1] if info['loss'] else 'N/A'}"),
|
|
|
- html.P(f"Test Accuracy: {info['test_accuracy']}")
|
|
|
- ], className='mb-3'))
|
|
|
- except FileNotFoundError:
|
|
|
- model_info.append(html.Div([
|
|
|
- html.H4(f"Model ({extension.upper()})"),
|
|
|
- html.P("No detailed information available")
|
|
|
- ], className='mb-3'))
|
|
|
-
|
|
|
- if not model_info:
|
|
|
- model_info = html.Div([
|
|
|
- html.P("No trained model available"),
|
|
|
- ], className='mb-3')
|
|
|
-
|
|
|
- return html.Div(model_info)
|
|
|
-
|
|
|
-def generate_training_data_table(model_name, files, page_size=10):
|
|
|
+def generate_training_data_table(model_name, files, page=1, page_size=10):
|
|
|
import os
|
|
|
header = [
|
|
|
html.Thead(html.Tr([html.Th("Audio"), html.Th("File Name"), html.Th("Display"), html.Th("Length (s)"), html.Th("Size (KB)"), html.Th("Actions")]))
|
|
|
@@ -219,5 +272,13 @@ def generate_training_data_table(model_name, files, page_size=10):
|
|
|
html.Td("-"),
|
|
|
html.Td("File not found")
|
|
|
]))
|
|
|
- pagination = dcc.Pagination(id='data-table-pagination', max_value=len(rows)//page_size + 1, page_size=page_size)
|
|
|
- return html.Div([dbc.Table(header + [html.Tbody(rows)], bordered=True, hover=True), pagination])
|
|
|
+
|
|
|
+ start = (page - 1) * 10
|
|
|
+ end = start + 10
|
|
|
+ row_count = len(rows)
|
|
|
+ table_body = html.Tbody(id='data-table-body', children=rows[start:end])
|
|
|
+
|
|
|
+ return html.Div([
|
|
|
+ dbc.Table(header + [table_body], bordered=True, hover=True)
|
|
|
+ ]), row_count
|
|
|
+
|