Browse Source

2nd Commit2

Patrick Baumgartner 2 years ago
parent
commit
a50c6a83d2
2 changed files with 23 additions and 44 deletions
  1. 7 1
      trainer/trainer_scripts/pytorch_adv_trainer.py
  2. 16 43
      trainer/utils/pytorch_utils.py

+ 7 - 1
trainer/trainer_scripts/pytorch_adv_trainer.py

@@ -3,7 +3,10 @@ import time
 import json
 import torch
 import torch.optim as optim
+import pandas as pd
+import numpy as np
 from torch.utils.data import DataLoader
+from sklearn.model_selection import train_test_split
 from utils.training_interface import TrainerInterface
 from utils.file_utils import list_files_in_directory
 from utils.pytorch_utils import WakeWordDataset, CNNNetwork, preprocess_training_data, train_model, optimize_graph
@@ -89,7 +92,10 @@ class pytorch_adv_trainer(TrainerInterface):
     def load_model(self):
         model = CNNNetwork()
         if os.path.exists(self.model_path + "hotword_model.pth"):
-            model.load_state_dict(torch.load(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)
         return model
 
     def start(self):

+ 16 - 43
trainer/utils/pytorch_utils.py

@@ -1,7 +1,5 @@
 import torch
 import torch.nn as nn
-import numpy as np
-import pandas as pd
 from torch.utils.data import Dataset
 
 class WakeWordDataset(Dataset):
@@ -13,29 +11,23 @@ class WakeWordDataset(Dataset):
         return len(self.features)
 
     def __getitem__(self, idx):
-        feature = torch.tensor(self.features[idx], dtype=torch.float32)
-        label = torch.tensor(self.labels[idx], dtype=torch.long)
-        return feature, label
+        return self.features[idx], self.labels[idx]
 
 class CNNNetwork(nn.Module):
     def __init__(self):
         super(CNNNetwork, self).__init__()
         self.network = nn.Sequential(
-            nn.Conv2d(1, 32, kernel_size=3, stride=1, padding=1),
+            nn.Conv1d(in_channels=1, out_channels=32, kernel_size=5, stride=1, padding=2),
             nn.ReLU(),
-            nn.MaxPool2d(kernel_size=2, stride=2),
-            nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1),
+            nn.MaxPool1d(kernel_size=2, stride=2),
+            nn.Conv1d(in_channels=32, out_channels=64, kernel_size=5, stride=1, padding=2),
             nn.ReLU(),
-            nn.MaxPool2d(kernel_size=2, stride=2),
-            nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),
-            nn.ReLU(),
-            nn.MaxPool2d(kernel_size=2, stride=2),
+            nn.MaxPool1d(kernel_size=2, stride=2),
             nn.Flatten(),
-            nn.Linear(128 * 4 * 4, 128),
+            nn.Linear(64 * 10, 128),
             nn.ReLU(),
             nn.Dropout(0.5),
-            nn.Linear(128, 1),
-            nn.Sigmoid()
+            nn.Linear(128, 1)
         )
 
     def forward(self, x):
@@ -43,6 +35,7 @@ class CNNNetwork(nn.Module):
 
 def preprocess_training_data(training_directories, background_directory, model_path, update_progress_callback):
     import librosa
+    import pandas as pd
     data_path_dict = {}
     data_path_dict[0] = list_files_in_directory(background_directory, extension='.wav')
     for idx, directory in enumerate(training_directories):
@@ -60,7 +53,7 @@ def preprocess_training_data(training_directories, background_directory, model_p
                 mfcc_processed = np.mean(mfcc.T, axis=0)
                 all_data.append([mfcc_processed, class_label])
                 processed_files += 1
-                update_progress_callback(processed_files / total_files * 100)
+                update_progress_callback((processed_files / total_files) * 100)
             except Exception as e:
                 print(f"Exception: {str(e)}")
 
@@ -69,35 +62,15 @@ def preprocess_training_data(training_directories, background_directory, model_p
 
 def train_model(model, train_loader, optimizer, criterion, epoch, total_epochs):
     model.train()
-    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
-    model.to(device)
-    
-    for batch_idx, (data, target) in enumerate(train_loader):
-        data, target = data.to(device), target.to(device)
+    for batch_X, batch_y in train_loader:
         optimizer.zero_grad()
-        output = model(data)
-        loss = criterion(output, target.float().unsqueeze(1))
+        outputs = model(batch_X.unsqueeze(1).float())
+        loss = criterion(outputs, batch_y.unsqueeze(1).float())
         loss.backward()
         optimizer.step()
-
-        if batch_idx % 10 == 0:
-            print(f'Train Epoch: {epoch}/{total_epochs} [{batch_idx * len(data)}/{len(train_loader.dataset)} ({100. * batch_idx / len(train_loader):.0f}%)]\tLoss: {loss.item():.6f}')
+        global current_loss, current_accuracy, total_batches
+        current_loss = loss.item()
+        current_accuracy = (outputs.argmax(dim=1) == batch_y).float().mean().item()
 
 def optimize_graph(model, device, optimizer, criterion, dataloader):
-    model.train()
-    total_loss = 0
-    correct = 0
-
-    for data, target in dataloader:
-        data, target = data.to(device), target.to(device)
-        optimizer.zero_grad()
-        output = model(data)
-        loss = criterion(output, target.float().unsqueeze(1))
-        loss.backward()
-        optimizer.step()
-        total_loss += loss.item()
-        pred = torch.round(output)
-        correct += pred.eq(target.float().unsqueeze(1)).sum().item()
-
-    accuracy = correct / len(dataloader.dataset)
-    return total_loss / len(dataloader), accuracy
+    pass