| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195 |
- import os
- 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
- # Global variables for tracking progress
- current_progress = 0
- current_loss = 0.0
- current_accuracy = 0.0
- total_epochs = 0
- total_batches = 0
- class pytorch_adv_trainer(TrainerInterface):
- def __init__(self, learning_rate=0.001):
- self.model_path = None
- self.training_directories = []
- self.sample_directory = None
- self.background_directory = None
- self.dropout_rate = 0.5
- 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().to(self.device)
- self.current_epoch = 0
- self.stop_training = False
- self.training_status = "Not started"
- self.train_loader = None
- self.val_loader = None
- 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):
- self.training_directories = directories
- def set_sample_directory(self, directory):
- self.sample_directory = directory
- def set_background_directory(self, directory):
- self.background_directory = directory
- def set_dropout_rate(self, rate):
- self.dropout_rate = rate
- def set_epochs(self, epochs):
- self.epochs = epochs
- def set_batch_size(self, batch_size):
- self.batch_size = batch_size
- def get_training_data(self):
- return self.training_directories
- def get_model_path(self):
- return self.model_path
- def get_epochs(self):
- return self.epochs
- def get_batch_size(self):
- return self.batch_size
- def get_dropout_rate(self):
- return self.dropout_rate
- def get_training_directories(self):
- return self.training_directories
- def get_sample_directory(self):
- return self.sample_directory
- def get_background_directory(self):
- return self.background_directory
- def set_training_data(self, features, labels):
- X_train, X_val, y_train, y_val = train_test_split(features, labels, test_size=0.2, random_state=42)
- self.train_loader = DataLoader(WakeWordDataset(X_train, y_train), batch_size=self.batch_size, shuffle=True)
- self.val_loader = DataLoader(WakeWordDataset(X_val, y_val), batch_size=self.batch_size, shuffle=False)
- def load_model(self):
- model = CNNNetwork()
- 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):
- self.training_status = "Training"
- self.stop_training = False
- self.run_training()
- def start_new(self):
- 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()
- def pause(self):
- self.training_status = "Paused"
- self.stop_training = True
- def resume(self):
- self.training_status = "Training"
- self.stop_training = False
- self.run_training()
- def stop(self):
- self.training_status = "Stopped"
- self.stop_training = True
- def preprocess_training_data(self, update_progress_callback):
- preprocess_training_data(self.training_directories, self.background_directory, self.model_path, update_progress_callback)
- def run_training(self):
- global current_progress, current_loss, current_accuracy, total_epochs, total_batches
- self.preprocess_training_data(lambda progress: setattr(self, 'current_progress', progress))
- 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)
-
- total_batches = int(np.ceil(len(self.train_loader.dataset) / self.batch_size))
- total_epochs = self.epochs
- for epoch in range(self.current_epoch, self.epochs):
- if self.stop_training:
- break
- self.current_epoch = epoch
- 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({'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).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)
- print(f"Test Accuracy: {accuracy * 100:.2f}%")
- def get_training_info(self):
- return {
- "status": self.training_status,
- "current_epoch": self.current_epoch,
- "total_epochs": self.epochs,
- "loss": current_loss,
- "accuracy": current_accuracy
- }
|