tensorflow_trainer.py 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231
  1. import os
  2. import logging
  3. import time
  4. import numpy as np
  5. import pandas as pd
  6. from sklearn.model_selection import train_test_split
  7. from tensorflow.keras.models import Sequential, load_model
  8. from tensorflow.keras.layers import Dense, Activation, Dropout
  9. from tensorflow.keras.utils import to_categorical
  10. from utils.training_interface import TrainerInterface
  11. from utils.file_utils import list_files_in_directory
  12. import json
  13. # Define global variables for tracking progress
  14. current_progress = 0
  15. current_loss = 0.0
  16. current_accuracy = 0.0
  17. total_epochs = 0
  18. total_batches = 0
  19. class TensorFlowTrainer(TrainerInterface):
  20. def __init__(self, learning_rate=0.001):
  21. self.model_path = None
  22. self.training_directories = []
  23. self.sample_directory = None
  24. self.background_directory = None
  25. self.dropout_rate = 0.5
  26. self.epochs = 10
  27. self.batch_size = 32
  28. self.learning_rate = learning_rate
  29. self.model = None
  30. self.current_epoch = 0
  31. self.stop_training = False
  32. self.training_status = "Not started"
  33. self.train_data = None
  34. self.val_data = None
  35. def set_model_path(self, model_path):
  36. self.model_path = model_path
  37. self.model = self.load_model()
  38. def set_training_directories(self, directories):
  39. self.training_directories = directories
  40. def set_sample_directory(self, directory):
  41. self.sample_directory = directory
  42. def set_background_directory(self, directory):
  43. self.background_directory = directory
  44. def set_dropout_rate(self, rate):
  45. self.dropout_rate = rate
  46. def set_epochs(self, epochs):
  47. self.epochs = epochs
  48. def set_batch_size(self, batch_size):
  49. self.batch_size = batch_size
  50. def get_training_data(self):
  51. return self.training_directories
  52. def get_model_path(self):
  53. return self.model_path
  54. def get_epochs(self):
  55. return self.epochs
  56. def get_batch_size(self):
  57. return self.batch_size
  58. def get_dropout_rate(self):
  59. return self.dropout_rate
  60. def get_training_directories(self):
  61. return self.training_directories
  62. def get_sample_directory(self):
  63. return self.sample_directory
  64. def get_background_directory(self):
  65. return self.background_directory
  66. def set_training_data(self, features, labels):
  67. X_train, X_val, y_train, y_val = train_test_split(features, labels, test_size=0.2, random_state=42)
  68. self.train_data = (X_train, y_train)
  69. self.val_data = (X_val, y_val)
  70. def load_model(self):
  71. if os.path.exists(self.model_path + "hotword_model.h5"):
  72. return load_model(self.model_path + "hotword_model.h5")
  73. else:
  74. model = Sequential([
  75. Dense(256, input_shape=(40,)),
  76. Activation('relu'),
  77. Dropout(self.dropout_rate),
  78. Dense(256),
  79. Activation('relu'),
  80. Dropout(self.dropout_rate),
  81. Dense(len(self.training_directories) + 1, activation='softmax')
  82. ])
  83. model.compile(
  84. loss="categorical_crossentropy",
  85. optimizer='adam',
  86. metrics=['accuracy']
  87. )
  88. return model
  89. def start(self):
  90. self.training_status = "Training"
  91. self.stop_training = False
  92. self.run_training()
  93. def start_new(self):
  94. self.training_status = "Training"
  95. self.stop_training = False
  96. self.model = self.load_model()
  97. self.current_epoch = 0
  98. self.run_training()
  99. def pause(self):
  100. self.training_status = "Paused"
  101. self.stop_training = True
  102. def resume(self):
  103. self.training_status = "Training"
  104. self.stop_training = False
  105. self.run_training()
  106. def stop(self):
  107. self.training_status = "Stopped"
  108. self.stop_training = True
  109. def preprocess_training_data(self):
  110. import librosa
  111. data_path_dict = {}
  112. data_path_dict[0] = list_files_in_directory(self.background_directory, extension='.wav')
  113. for idx, directory in enumerate(self.training_directories):
  114. data_path_dict[idx+1] = list_files_in_directory(directory, extension='.wav')
  115. all_data = []
  116. total_files = sum(len(files) for files in data_path_dict.values())
  117. processed_files = 0
  118. for class_label, list_of_files in data_path_dict.items():
  119. for single_file in list_of_files:
  120. try:
  121. audio, sample_rate = librosa.load(single_file)
  122. mfcc = librosa.feature.mfcc(y=audio, sr=sample_rate, n_mfcc=40)
  123. mfcc_processed = np.mean(mfcc.T, axis=0)
  124. all_data.append([mfcc_processed, class_label])
  125. processed_files += 1
  126. global current_progress
  127. current_progress = (processed_files / total_files) * 100
  128. except Exception as e:
  129. print(f"Exception: {str(e)}")
  130. df = pd.DataFrame(all_data, columns=["feature", "class_label"])
  131. df.to_pickle(self.model_path + "audio_data.csv")
  132. def run_training(self):
  133. global current_progress, current_loss, current_accuracy, total_epochs, total_batches
  134. self.preprocess_training_data()
  135. start_time = time.time()
  136. df = pd.read_pickle(self.model_path + "audio_data.csv")
  137. X = df["feature"].values
  138. X = np.concatenate(X, axis=0).reshape(len(X), 40)
  139. y = np.array(df["class_label"].tolist())
  140. y = to_categorical(y)
  141. self.set_training_data(X, y)
  142. X_train, y_train = self.train_data
  143. total_batches = int(np.ceil(len(X_train) / self.batch_size))
  144. total_epochs = self.epochs
  145. training_info = {
  146. "trainer_script": "tensorflow_trainer",
  147. "training_time": 0,
  148. "model_size": 0,
  149. "epochs": self.epochs,
  150. "dropout_rate": self.dropout_rate,
  151. "batch_size": self.batch_size,
  152. "total_training_data": len(X),
  153. "training_directories": self.training_directories,
  154. "accuracy": [],
  155. "loss": [],
  156. "test_accuracy": 0
  157. }
  158. for epoch in range(self.current_epoch, self.epochs):
  159. if self.stop_training:
  160. break
  161. self.current_epoch = epoch
  162. history = self.model.fit(X_train, y_train, epochs=1, batch_size=self.batch_size)
  163. current_loss = history.history['loss'][-1]
  164. current_accuracy = history.history['accuracy'][-1]
  165. current_progress = (epoch + 1) / self.epochs * 100
  166. training_info["accuracy"].append(current_accuracy)
  167. training_info["loss"].append(current_loss)
  168. print(f"Epoch {epoch + 1}/{self.epochs}, Loss: {current_loss}, Accuracy: {current_accuracy}")
  169. epoch_time = time.time() - start_time
  170. training_info["training_time"] = epoch_time
  171. if not self.stop_training:
  172. self.model.save(self.model_path + "hotword_model.h5")
  173. self.evaluate_model(X, y)
  174. training_info["test_accuracy"] = current_accuracy
  175. training_info["model_size"] = os.path.getsize(self.model_path + "hotword_model.h5") / 1024 # in KB
  176. self.save_training_info(training_info)
  177. self.training_status = "Finished"
  178. else:
  179. self.training_status = "Paused"
  180. def evaluate_model(self, X_test, y_test):
  181. loss, accuracy = self.model.evaluate(X_test, y_test)
  182. print(f"Test Accuracy: {accuracy * 100:.2f}%")
  183. def save_training_info(self, info):
  184. with open(self.model_path + "info_tensorflow_trainer.json", "w") as f:
  185. json.dump(info, f)
  186. def get_training_info(self):
  187. return {
  188. "status": self.training_status,
  189. "current_epoch": self.current_epoch,
  190. "total_epochs": self.epochs,
  191. "loss": current_loss,
  192. "accuracy": current_accuracy
  193. }