pytorch_trainer.py 9.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261
  1. import torch
  2. import torch.nn as nn
  3. import torch.optim as optim
  4. import os
  5. import pandas as pd
  6. import numpy as np
  7. from sklearn.model_selection import train_test_split
  8. from utils.training_interface import TrainerInterface
  9. from utils.file_utils import list_files_in_directory
  10. import time
  11. import json
  12. # Global variables for progress tracking
  13. current_progress = 0
  14. current_loss = 0.0
  15. current_accuracy = 0.0
  16. total_epochs = 0
  17. total_batches = 0
  18. class CustomDataset(torch.utils.data.Dataset):
  19. def __init__(self, features, labels):
  20. self.features = features
  21. self.labels = labels
  22. def __len__(self):
  23. return len(self.features)
  24. def __getitem__(self, idx):
  25. return self.features[idx], self.labels[idx]
  26. class PyTorchTrainer(TrainerInterface):
  27. def __init__(self, learning_rate=0.001):
  28. self.model_path = None
  29. self.training_directories = []
  30. self.sample_directory = None
  31. self.background_directory = None
  32. self.dropout_rate = 0.5
  33. self.epochs = 3000
  34. self.batch_size = 32
  35. self.learning_rate = learning_rate
  36. self.model = None
  37. self.optimizer = None
  38. self.criterion = nn.CrossEntropyLoss()
  39. self.current_epoch = 0
  40. self.stop_training = False
  41. self.training_status = "Not started"
  42. self.train_loader = None
  43. self.val_loader = None
  44. def set_model_path(self, model_path):
  45. self.model_path = model_path
  46. self.model = self.load_model()
  47. self.optimizer = optim.Adam(self.model.parameters(), lr=self.learning_rate)
  48. def set_training_directories(self, directories):
  49. self.training_directories = directories
  50. def set_sample_directory(self, directory):
  51. self.sample_directory = directory
  52. def set_background_directory(self, directory):
  53. self.background_directory = directory
  54. def set_dropout_rate(self, rate):
  55. self.dropout_rate = rate
  56. def set_epochs(self, epochs):
  57. self.epochs = epochs
  58. def set_batch_size(self, batch_size):
  59. self.batch_size = batch_size
  60. def get_training_data(self):
  61. return self.training_directories
  62. def get_model_path(self):
  63. return self.model_path
  64. def get_epochs(self):
  65. return self.epochs
  66. def get_batch_size(self):
  67. return self.batch_size
  68. def get_dropout_rate(self):
  69. return self.dropout_rate
  70. def get_training_directories(self):
  71. return self.training_directories
  72. def get_sample_directory(self):
  73. return self.sample_directory
  74. def get_background_directory(self):
  75. return self.background_directory
  76. def set_training_data(self, features, labels):
  77. X_train, X_val, y_train, y_val = train_test_split(features, labels, test_size=0.2, random_state=42)
  78. self.train_loader = torch.utils.data.DataLoader(CustomDataset(X_train, y_train), batch_size=self.batch_size, shuffle=True)
  79. self.val_loader = torch.utils.data.DataLoader(CustomDataset(X_val, y_val), batch_size=self.batch_size, shuffle=False)
  80. def preprocess_training_data(self):
  81. import librosa
  82. data_path_dict = {}
  83. data_path_dict[0] = list_files_in_directory(self.background_directory, extension='.wav')
  84. for idx, directory in enumerate(self.training_directories):
  85. data_path_dict[idx+1] = list_files_in_directory(directory, extension='.wav')
  86. all_data = []
  87. total_files = sum(len(files) for files in data_path_dict.values())
  88. processed_files = 0
  89. for class_label, list_of_files in data_path_dict.items():
  90. for single_file in list_of_files:
  91. try:
  92. audio, sample_rate = librosa.load(single_file)
  93. mfcc = librosa.feature.mfcc(y=audio, sr=sample_rate, n_mfcc=40)
  94. mfcc_processed = np.mean(mfcc.T, axis=0)
  95. all_data.append([mfcc_processed, class_label])
  96. processed_files += 1
  97. global current_progress
  98. current_progress = (processed_files / total_files) * 100
  99. except Exception as e:
  100. print(f"Exception: {str(e)}")
  101. df = pd.DataFrame(all_data, columns=["feature", "class_label"])
  102. df.to_pickle(self.model_path + "audio_data.csv")
  103. def load_model(self):
  104. num_classes = len(self.training_directories) + 1 # Assuming one extra for background or unknown
  105. model = nn.Sequential(
  106. nn.Linear(40, 128),
  107. nn.ReLU(),
  108. nn.Dropout(self.dropout_rate),
  109. nn.Linear(128, num_classes), # Update output layer to match num_classes
  110. nn.Softmax(dim=1)
  111. )
  112. model_path = self.model_path + "hotword_model.pth"
  113. if os.path.exists(model_path):
  114. state_dict = torch.load(model_path)
  115. try:
  116. model.load_state_dict(state_dict)
  117. except RuntimeError as e:
  118. print(f"Error loading model state dict: {e}")
  119. # Handle dimension mismatch by reinitializing the last layer
  120. model = nn.Sequential(
  121. nn.Linear(40, 128),
  122. nn.ReLU(),
  123. nn.Dropout(self.dropout_rate),
  124. nn.Linear(128, num_classes),
  125. nn.Softmax(dim=1)
  126. )
  127. print("Reinitialized the final layer to match the number of classes.")
  128. return model
  129. def start(self):
  130. self.training_status = "Training"
  131. self.stop_training = False
  132. self.run_training()
  133. def start_new(self):
  134. self.training_status = "Training"
  135. self.stop_training = False
  136. self.model = self.load_model()
  137. self.current_epoch = 0
  138. self.run_training()
  139. def pause(self):
  140. self.training_status = "Paused"
  141. self.stop_training = True
  142. def resume(self):
  143. self.training_status = "Training"
  144. self.stop_training = False
  145. self.run_training()
  146. def stop(self):
  147. self.training_status = "Stopped"
  148. self.stop_training = True
  149. def run_training(self):
  150. global current_progress, current_loss, current_accuracy, total_epochs, total_batches
  151. self.preprocess_training_data()
  152. start_time = time.time()
  153. df = pd.read_pickle(self.model_path + "audio_data.csv")
  154. X = df["feature"].values
  155. X = np.concatenate(X, axis=0).reshape(len(X), 40)
  156. y = np.array(df["class_label"].tolist())
  157. self.set_training_data(X, y)
  158. total_batches = len(self.train_loader)
  159. total_epochs = self.epochs
  160. training_info = {
  161. "trainer_script": "tensorflow_trainer",
  162. "training_time": 0,
  163. "model_size": 0,
  164. "epochs": self.epochs,
  165. "dropout_rate": self.dropout_rate,
  166. "batch_size": self.batch_size,
  167. "total_training_data": len(X),
  168. "training_directories": self.training_directories,
  169. "accuracy": [],
  170. "loss": [],
  171. "test_accuracy": 0
  172. }
  173. for epoch in range(self.current_epoch, self.epochs):
  174. if self.stop_training:
  175. break
  176. self.current_epoch = epoch
  177. self.model.train()
  178. for batch_idx, (batch_X, batch_y) in enumerate(self.train_loader):
  179. global current_loss, current_accuracy
  180. self.optimizer.zero_grad()
  181. outputs = self.model(batch_X)
  182. loss = self.criterion(outputs, batch_y.long()) # Convert labels to Long type
  183. loss.backward()
  184. self.optimizer.step()
  185. # Update global progress variables
  186. current_loss = loss.item()
  187. current_accuracy = (outputs.argmax(dim=1) == batch_y).float().mean().item()
  188. current_progress = (epoch * len(self.train_loader) + batch_idx + 1) / (self.epochs * len(self.train_loader)) * 100
  189. training_info["accuracy"].append(current_accuracy)
  190. training_info["loss"].append(current_loss)
  191. print(f"Epoch {epoch + 1}/{self.epochs}, Loss: {current_loss}, Accuracy: {current_accuracy}")
  192. epoch_time = time.time() - start_time
  193. training_info["training_time"] = epoch_time
  194. if not self.stop_training:
  195. torch.save(self.model.state_dict(), self.model_path + "hotword_model.pth")
  196. self.evaluate_model(X, y)
  197. self.save_training_info(training_info)
  198. self.training_status = "Finished"
  199. else:
  200. self.training_status = "Paused"
  201. def evaluate_model(self, X_test, y_test):
  202. self.model.eval()
  203. with torch.no_grad():
  204. X_test_tensor = torch.FloatTensor(X_test)
  205. y_test_tensor = torch.LongTensor(y_test)
  206. outputs = self.model(X_test_tensor)
  207. _, predicted = torch.max(outputs, 1)
  208. accuracy = (predicted == y_test_tensor).sum().item() / len(y_test_tensor)
  209. print(f"Test Accuracy: {accuracy * 100:.2f}%")
  210. def save_training_info(self, info):
  211. with open(self.model_path + "info_pytorch_trainer.json", "w") as f:
  212. json.dump(info, f)
  213. def get_training_info(self):
  214. return {
  215. "status": self.training_status,
  216. "current_epoch": self.current_epoch,
  217. "total_epochs": self.epochs,
  218. "loss": current_loss,
  219. "accuracy": current_accuracy
  220. }