pytorch_adv_trainer.py 7.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195
  1. import os
  2. import time
  3. import json
  4. import torch
  5. import torch.optim as optim
  6. import pandas as pd
  7. import numpy as np
  8. from torch.utils.data import DataLoader
  9. from sklearn.model_selection import train_test_split
  10. from utils.training_interface import TrainerInterface
  11. from utils.file_utils import list_files_in_directory
  12. from utils.pytorch_utils import WakeWordDataset, CNNNetwork, preprocess_training_data, train_model, optimize_graph
  13. # 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 pytorch_adv_trainer(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.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  30. self.model = None
  31. self.optimizer = None
  32. self.criterion = torch.nn.BCEWithLogitsLoss().to(self.device)
  33. self.current_epoch = 0
  34. self.stop_training = False
  35. self.training_status = "Not started"
  36. self.train_loader = None
  37. self.val_loader = None
  38. def set_model_path(self, model_path):
  39. self.model_path = model_path
  40. self.model = self.load_model()
  41. self.model.to(self.device)
  42. self.optimizer = optim.Adam(self.model.parameters(), lr=self.learning_rate)
  43. def set_training_directories(self, directories):
  44. self.training_directories = directories
  45. def set_sample_directory(self, directory):
  46. self.sample_directory = directory
  47. def set_background_directory(self, directory):
  48. self.background_directory = directory
  49. def set_dropout_rate(self, rate):
  50. self.dropout_rate = rate
  51. def set_epochs(self, epochs):
  52. self.epochs = epochs
  53. def set_batch_size(self, batch_size):
  54. self.batch_size = batch_size
  55. def get_training_data(self):
  56. return self.training_directories
  57. def get_model_path(self):
  58. return self.model_path
  59. def get_epochs(self):
  60. return self.epochs
  61. def get_batch_size(self):
  62. return self.batch_size
  63. def get_dropout_rate(self):
  64. return self.dropout_rate
  65. def get_training_directories(self):
  66. return self.training_directories
  67. def get_sample_directory(self):
  68. return self.sample_directory
  69. def get_background_directory(self):
  70. return self.background_directory
  71. def set_training_data(self, features, labels):
  72. X_train, X_val, y_train, y_val = train_test_split(features, labels, test_size=0.2, random_state=42)
  73. self.train_loader = DataLoader(WakeWordDataset(X_train, y_train), batch_size=self.batch_size, shuffle=True)
  74. self.val_loader = DataLoader(WakeWordDataset(X_val, y_val), batch_size=self.batch_size, shuffle=False)
  75. def load_model(self):
  76. model = CNNNetwork()
  77. model_file_path = os.path.join(self.model_path, "hotword_model.pth")
  78. if os.path.exists(model_file_path):
  79. print(f"Loading model from {model_file_path}")
  80. checkpoint = torch.load(model_file_path)
  81. if 'model_state_dict' in checkpoint:
  82. state_dict = checkpoint['model_state_dict']
  83. new_state_dict = {}
  84. model_state_dict = model.state_dict()
  85. for key in model_state_dict.keys():
  86. if key in state_dict:
  87. if state_dict[key].shape == model_state_dict[key].shape:
  88. new_state_dict[key] = state_dict[key]
  89. else:
  90. print(f"Shape mismatch for {key}, skipping.")
  91. else:
  92. print(f"Missing key {key}, skipping.")
  93. model.load_state_dict(new_state_dict, strict=False)
  94. else:
  95. print(f"Key 'model_state_dict' not found in checkpoint. Initializing new model.")
  96. else:
  97. print(f"No model found at {model_file_path}. Initializing new model.")
  98. return model
  99. def start(self):
  100. self.training_status = "Training"
  101. self.stop_training = False
  102. self.run_training()
  103. def start_new(self):
  104. self.training_status = "Training"
  105. self.stop_training = False
  106. self.model = self.load_model()
  107. self.model.to(self.device)
  108. self.current_epoch = 0
  109. self.run_training()
  110. def pause(self):
  111. self.training_status = "Paused"
  112. self.stop_training = True
  113. def resume(self):
  114. self.training_status = "Training"
  115. self.stop_training = False
  116. self.run_training()
  117. def stop(self):
  118. self.training_status = "Stopped"
  119. self.stop_training = True
  120. def preprocess_training_data(self, update_progress_callback):
  121. preprocess_training_data(self.training_directories, self.background_directory, self.model_path, update_progress_callback)
  122. def run_training(self):
  123. global current_progress, current_loss, current_accuracy, total_epochs, total_batches
  124. self.preprocess_training_data(lambda progress: setattr(self, 'current_progress', progress))
  125. df = pd.read_pickle(self.model_path + "audio_data.csv")
  126. X = df["feature"].values
  127. X = np.concatenate(X, axis=0).reshape(len(X), 40)
  128. X = np.expand_dims(X, axis=1) # Add channel dimension
  129. y = np.array(df["class_label"].tolist())
  130. self.set_training_data(X, y)
  131. total_batches = int(np.ceil(len(self.train_loader.dataset) / self.batch_size))
  132. total_epochs = self.epochs
  133. for epoch in range(self.current_epoch, self.epochs):
  134. if self.stop_training:
  135. break
  136. self.current_epoch = epoch
  137. train_model(self.model, self.train_loader, self.optimizer, self.criterion, epoch, self.epochs, self.device)
  138. current_progress = (epoch + 1) / self.epochs * 100
  139. print(f"Epoch {epoch + 1}/{self.epochs}")
  140. if not self.stop_training:
  141. torch.save({'model_state_dict': self.model.state_dict()}, os.path.join(self.model_path, "hotword_model.pth"))
  142. self.evaluate_model(X, y)
  143. self.training_status = "Finished"
  144. def evaluate_model(self, X_test, y_test):
  145. self.model.eval()
  146. with torch.no_grad():
  147. X_test_tensor = torch.FloatTensor(X_test).to(self.device)
  148. y_test_tensor = torch.LongTensor(y_test).to(self.device)
  149. outputs = self.model(X_test_tensor)
  150. _, predicted = torch.max(outputs, 1)
  151. accuracy = (predicted == y_test_tensor).sum().item() / len(y_test_tensor)
  152. print(f"Test Accuracy: {accuracy * 100:.2f}%")
  153. def get_training_info(self):
  154. return {
  155. "status": self.training_status,
  156. "current_epoch": self.current_epoch,
  157. "total_epochs": self.epochs,
  158. "loss": current_loss,
  159. "accuracy": current_accuracy
  160. }