pytorch_utils.py 3.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103
  1. import torch
  2. import torch.nn as nn
  3. import numpy as np
  4. import pandas as pd
  5. from torch.utils.data import Dataset
  6. class WakeWordDataset(Dataset):
  7. def __init__(self, features, labels):
  8. self.features = features
  9. self.labels = labels
  10. def __len__(self):
  11. return len(self.features)
  12. def __getitem__(self, idx):
  13. feature = torch.tensor(self.features[idx], dtype=torch.float32)
  14. label = torch.tensor(self.labels[idx], dtype=torch.long)
  15. return feature, label
  16. class CNNNetwork(nn.Module):
  17. def __init__(self):
  18. super(CNNNetwork, self).__init__()
  19. self.network = nn.Sequential(
  20. nn.Conv2d(1, 32, kernel_size=3, stride=1, padding=1),
  21. nn.ReLU(),
  22. nn.MaxPool2d(kernel_size=2, stride=2),
  23. nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1),
  24. nn.ReLU(),
  25. nn.MaxPool2d(kernel_size=2, stride=2),
  26. nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),
  27. nn.ReLU(),
  28. nn.MaxPool2d(kernel_size=2, stride=2),
  29. nn.Flatten(),
  30. nn.Linear(128 * 4 * 4, 128),
  31. nn.ReLU(),
  32. nn.Dropout(0.5),
  33. nn.Linear(128, 1),
  34. nn.Sigmoid()
  35. )
  36. def forward(self, x):
  37. return self.network(x)
  38. def preprocess_training_data(training_directories, background_directory, model_path, update_progress_callback):
  39. import librosa
  40. data_path_dict = {}
  41. data_path_dict[0] = list_files_in_directory(background_directory, extension='.wav')
  42. for idx, directory in enumerate(training_directories):
  43. data_path_dict[idx+1] = list_files_in_directory(directory, extension='.wav')
  44. all_data = []
  45. total_files = sum(len(files) for files in data_path_dict.values())
  46. processed_files = 0
  47. for class_label, list_of_files in data_path_dict.items():
  48. for single_file in list_of_files:
  49. try:
  50. audio, sample_rate = librosa.load(single_file)
  51. mfcc = librosa.feature.mfcc(y=audio, sr=sample_rate, n_mfcc=40)
  52. mfcc_processed = np.mean(mfcc.T, axis=0)
  53. all_data.append([mfcc_processed, class_label])
  54. processed_files += 1
  55. update_progress_callback(processed_files / total_files * 100)
  56. except Exception as e:
  57. print(f"Exception: {str(e)}")
  58. df = pd.DataFrame(all_data, columns=["feature", "class_label"])
  59. df.to_pickle(model_path + "audio_data.csv")
  60. def train_model(model, train_loader, optimizer, criterion, epoch, total_epochs):
  61. model.train()
  62. device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
  63. model.to(device)
  64. for batch_idx, (data, target) in enumerate(train_loader):
  65. data, target = data.to(device), target.to(device)
  66. optimizer.zero_grad()
  67. output = model(data)
  68. loss = criterion(output, target.float().unsqueeze(1))
  69. loss.backward()
  70. optimizer.step()
  71. if batch_idx % 10 == 0:
  72. 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}')
  73. def optimize_graph(model, device, optimizer, criterion, dataloader):
  74. model.train()
  75. total_loss = 0
  76. correct = 0
  77. for data, target in dataloader:
  78. data, target = data.to(device), target.to(device)
  79. optimizer.zero_grad()
  80. output = model(data)
  81. loss = criterion(output, target.float().unsqueeze(1))
  82. loss.backward()
  83. optimizer.step()
  84. total_loss += loss.item()
  85. pred = torch.round(output)
  86. correct += pred.eq(target.float().unsqueeze(1)).sum().item()
  87. accuracy = correct / len(dataloader.dataset)
  88. return total_loss / len(dataloader), accuracy