|
|
@@ -28,9 +28,10 @@ class pytorch_adv_trainer(TrainerInterface):
|
|
|
self.epochs = 10
|
|
|
self.batch_size = 32
|
|
|
self.learning_rate = learning_rate
|
|
|
+ self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
|
self.model = None
|
|
|
self.optimizer = None
|
|
|
- self.criterion = torch.nn.BCEWithLogitsLoss()
|
|
|
+ self.criterion = torch.nn.BCEWithLogitsLoss().to(self.device)
|
|
|
self.current_epoch = 0
|
|
|
self.stop_training = False
|
|
|
self.training_status = "Not started"
|
|
|
@@ -40,6 +41,7 @@ class pytorch_adv_trainer(TrainerInterface):
|
|
|
def set_model_path(self, model_path):
|
|
|
self.model_path = model_path
|
|
|
self.model = self.load_model()
|
|
|
+ self.model.to(self.device)
|
|
|
self.optimizer = optim.Adam(self.model.parameters(), lr=self.learning_rate)
|
|
|
|
|
|
def set_training_directories(self, directories):
|
|
|
@@ -91,11 +93,29 @@ class pytorch_adv_trainer(TrainerInterface):
|
|
|
|
|
|
def load_model(self):
|
|
|
model = CNNNetwork()
|
|
|
- if os.path.exists(self.model_path + "hotword_model.pth"):
|
|
|
- state_dict = torch.load(self.model_path + "hotword_model.pth")
|
|
|
- # Update the state_dict keys to match the model keys
|
|
|
- new_state_dict = {f'network.{k}': v for k, v in state_dict.items()}
|
|
|
- model.load_state_dict(new_state_dict)
|
|
|
+ model_file_path = os.path.join(self.model_path, "hotword_model.pth")
|
|
|
+ if os.path.exists(model_file_path):
|
|
|
+ print(f"Loading model from {model_file_path}")
|
|
|
+ checkpoint = torch.load(model_file_path)
|
|
|
+ if 'model_state_dict' in checkpoint:
|
|
|
+ state_dict = checkpoint['model_state_dict']
|
|
|
+ new_state_dict = {}
|
|
|
+
|
|
|
+ model_state_dict = model.state_dict()
|
|
|
+ for key in model_state_dict.keys():
|
|
|
+ if key in state_dict:
|
|
|
+ if state_dict[key].shape == model_state_dict[key].shape:
|
|
|
+ new_state_dict[key] = state_dict[key]
|
|
|
+ else:
|
|
|
+ print(f"Shape mismatch for {key}, skipping.")
|
|
|
+ else:
|
|
|
+ print(f"Missing key {key}, skipping.")
|
|
|
+
|
|
|
+ model.load_state_dict(new_state_dict, strict=False)
|
|
|
+ else:
|
|
|
+ print(f"Key 'model_state_dict' not found in checkpoint. Initializing new model.")
|
|
|
+ else:
|
|
|
+ print(f"No model found at {model_file_path}. Initializing new model.")
|
|
|
return model
|
|
|
|
|
|
def start(self):
|
|
|
@@ -107,6 +127,7 @@ class pytorch_adv_trainer(TrainerInterface):
|
|
|
self.training_status = "Training"
|
|
|
self.stop_training = False
|
|
|
self.model = self.load_model()
|
|
|
+ self.model.to(self.device)
|
|
|
self.current_epoch = 0
|
|
|
self.run_training()
|
|
|
|
|
|
@@ -133,6 +154,7 @@ class pytorch_adv_trainer(TrainerInterface):
|
|
|
df = pd.read_pickle(self.model_path + "audio_data.csv")
|
|
|
X = df["feature"].values
|
|
|
X = np.concatenate(X, axis=0).reshape(len(X), 40)
|
|
|
+ X = np.expand_dims(X, axis=1) # Add channel dimension
|
|
|
y = np.array(df["class_label"].tolist())
|
|
|
self.set_training_data(X, y)
|
|
|
|
|
|
@@ -143,21 +165,21 @@ class pytorch_adv_trainer(TrainerInterface):
|
|
|
if self.stop_training:
|
|
|
break
|
|
|
self.current_epoch = epoch
|
|
|
- train_model(self.model, self.train_loader, self.optimizer, self.criterion, epoch, self.epochs)
|
|
|
+ train_model(self.model, self.train_loader, self.optimizer, self.criterion, epoch, self.epochs, self.device)
|
|
|
|
|
|
current_progress = (epoch + 1) / self.epochs * 100
|
|
|
print(f"Epoch {epoch + 1}/{self.epochs}")
|
|
|
|
|
|
if not self.stop_training:
|
|
|
- torch.save(self.model.state_dict(), self.model_path + "hotword_model.pth")
|
|
|
+ torch.save({'model_state_dict': self.model.state_dict()}, os.path.join(self.model_path, "hotword_model.pth"))
|
|
|
self.evaluate_model(X, y)
|
|
|
self.training_status = "Finished"
|
|
|
|
|
|
def evaluate_model(self, X_test, y_test):
|
|
|
self.model.eval()
|
|
|
with torch.no_grad():
|
|
|
- X_test_tensor = torch.FloatTensor(X_test)
|
|
|
- y_test_tensor = torch.LongTensor(y_test)
|
|
|
+ X_test_tensor = torch.FloatTensor(X_test).to(self.device)
|
|
|
+ y_test_tensor = torch.LongTensor(y_test).to(self.device)
|
|
|
outputs = self.model(X_test_tensor)
|
|
|
_, predicted = torch.max(outputs, 1)
|
|
|
accuracy = (predicted == y_test_tensor).sum().item() / len(y_test_tensor)
|