Kaynağa Gözat

Working PyTorch ADV Training

Patrick Baumgartner 2 yıl önce
ebeveyn
işleme
1beb6329f8

BIN
trainer/model/hotword/trixy/audio_data.csv


BIN
trainer/model/hotword/trixy/hotword_model.pth


+ 32 - 10
trainer/trainer_scripts/pytorch_adv_trainer.py

@@ -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)

+ 5 - 2
trainer/utils/file_utils.py

@@ -9,6 +9,7 @@ import matplotlib.pyplot as plt
 import numpy as np
 import pandas as pd
 import librosa.display
+import collections
 
 # Configure logging
 logging.basicConfig(level=logging.INFO)
@@ -22,13 +23,15 @@ def ensure_directory_exists(directory):
 def list_files_in_directory(directory, extension=None, deep=0):
     if not os.path.exists(directory):
         raise FileNotFoundError(f"Directory {directory} does not exist.")
-    isarray = isinstance(extension, list)
+    #isarray = type(extension).__name__ in ('list', 'tuple', 'dict')
+    isarray = isinstance(extension, (list, tuple, dict))
     files = []
     for file in os.listdir(directory):
         if os.path.isdir(directory+file+"/") and deep>0 and file != "processed":
-            dfs = list_files_in_directory(directory+file+"/", deep-1)
+            dfs = list_files_in_directory(directory+file+"/", extension=extension, deep=deep-1)
             for df in dfs:
                 files.append(df)
+            continue
         if extension:
             if isarray:
                 found=False

+ 27 - 4
trainer/utils/pytorch_utils.py

@@ -1,6 +1,14 @@
 import torch
 import torch.nn as nn
 from torch.utils.data import Dataset
+from utils.file_utils import list_files_in_directory
+import numpy as np
+import wave
+import logging
+from pydub import AudioSegment
+import matplotlib.pyplot as plt
+import pandas as pd
+import librosa.display
 
 class WakeWordDataset(Dataset):
     def __init__(self, features, labels):
@@ -37,9 +45,9 @@ def preprocess_training_data(training_directories, background_directory, model_p
     import librosa
     import pandas as pd
     data_path_dict = {}
-    data_path_dict[0] = list_files_in_directory(background_directory, extension='.wav')
+    data_path_dict[0] = list_files_in_directory(background_directory, extension=['.wav','mp3','.m4a'],deep=5)
     for idx, directory in enumerate(training_directories):
-        data_path_dict[idx+1] = list_files_in_directory(directory, extension='.wav')
+        data_path_dict[idx+1] = list_files_in_directory(directory, extension=['.wav','mp3','.m4a'],deep=5)
 
     all_data = []
     total_files = sum(len(files) for files in data_path_dict.values())
@@ -60,11 +68,16 @@ def preprocess_training_data(training_directories, background_directory, model_p
     df = pd.DataFrame(all_data, columns=["feature", "class_label"])
     df.to_pickle(model_path + "audio_data.csv")
 
-def train_model(model, train_loader, optimizer, criterion, epoch, total_epochs):
+def train_model(model, train_loader, optimizer, criterion, epoch, total_epochs, device):
     model.train()
+    running_loss = 0.0
+    correct = 0
+    total = 0
+
     for batch_X, batch_y in train_loader:
+        batch_X, batch_y = batch_X.to(device), batch_y.to(device)
         optimizer.zero_grad()
-        outputs = model(batch_X.unsqueeze(1).float())
+        outputs = model(batch_X.float())
         loss = criterion(outputs, batch_y.unsqueeze(1).float())
         loss.backward()
         optimizer.step()
@@ -72,5 +85,15 @@ def train_model(model, train_loader, optimizer, criterion, epoch, total_epochs):
         current_loss = loss.item()
         current_accuracy = (outputs.argmax(dim=1) == batch_y).float().mean().item()
 
+        running_loss += loss.item()
+        _, predicted = torch.max(outputs.data, 1)
+        total += batch_y.size(0)
+        correct += (predicted == batch_y).sum().item()
+
+    epoch_loss = running_loss / len(train_loader)
+    epoch_acc = 100 * correct / total
+
+    print(f'Epoch [{epoch + 1}/{total_epochs}], Loss: {epoch_loss:.4f}, Accuracy: {epoch_acc:.2f}%')
+
 def optimize_graph(model, device, optimizer, criterion, dataloader):
     pass