summaryrefslogtreecommitdiff
path: root/neurobench_testing/examples/gsc/train_ANN.py
diff options
context:
space:
mode:
Diffstat (limited to 'neurobench_testing/examples/gsc/train_ANN.py')
-rw-r--r--neurobench_testing/examples/gsc/train_ANN.py112
1 files changed, 112 insertions, 0 deletions
diff --git a/neurobench_testing/examples/gsc/train_ANN.py b/neurobench_testing/examples/gsc/train_ANN.py
new file mode 100644
index 0000000..55c5ba5
--- /dev/null
+++ b/neurobench_testing/examples/gsc/train_ANN.py
@@ -0,0 +1,112 @@
+import torch
+import numpy as np
+import torch.optim as optim
+from tqdm import tqdm
+from torch.utils.data import DataLoader
+import torchaudio
+
+import torch.nn.functional as F
+
+
+from neurobench.datasets import SpeechCommands
+
+from ANN import M5
+
+BATCH_SIZE = 256
+NUM_WORKERS = 8
+EPOCHS = 50
+
+# Check if GPU is available
+device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
+print(device)
+
+# Load the dataset
+data_dir = "../../../data/speech_commands/"
+train_set = SpeechCommands(path=data_dir, subset="training")
+val_set = SpeechCommands(path=data_dir, subset="validation")
+test_set = SpeechCommands(path=data_dir, subset="testing")
+
+# Create the dataloaders
+train_loader = DataLoader(train_set, batch_size=BATCH_SIZE, shuffle=True)
+val_loader = DataLoader(val_set, batch_size=BATCH_SIZE, shuffle=True)
+test_loader = DataLoader(test_set, batch_size=BATCH_SIZE, shuffle=True)
+
+# Model
+model = M5()
+model.to(device)
+
+transform = torchaudio.transforms.Resample(orig_freq=16000, new_freq=8000)
+optimizer = optim.Adam(model.parameters(), lr=0.01, weight_decay=0.0001)
+scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=20, gamma=0.1)
+
+def get_likely_index(tensor):
+ # find most likely label index for each element in the batch
+ return tensor.argmax(dim=-1)
+
+def number_of_correct(pred, target):
+ # count number of correct predictions
+ return pred.squeeze().eq(target).sum().item()
+
+def train(model, epoch, log_interval):
+ model.train()
+ for batch_idx, (data, target) in enumerate(train_loader):
+
+ data = data.permute(0, 2, 1).to(device)
+ target = target.to(device)
+
+ # apply transform and model on whole batch directly on device
+ data = transform(data)
+ output = model(data)
+
+ # negative log-likelihood for a tensor of size (batch x 1 x n_output)
+ loss = F.nll_loss(output.squeeze(), target)
+
+ optimizer.zero_grad()
+ loss.backward()
+ optimizer.step()
+
+ # print training stats
+ if batch_idx % log_interval == 0:
+ print(f"Train Epoch: {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)} ({100. * batch_idx / len(train_loader):.0f}%)]\tLoss: {loss.item():.6f}")
+
+
+def validate(model, epoch):
+ model.eval()
+ correct = 0
+ for batch_idx, (data, target) in enumerate(val_loader):
+ data = data.permute(0, 2, 1).to(device)
+ target = target.to(device)
+
+ # apply transform and model on whole batch directly on device
+ data = transform(data)
+ output = model(data)
+
+ pred = get_likely_index(output)
+ correct += number_of_correct(pred, target)
+
+ return correct / len(val_loader.dataset)
+
+# Training Start
+best_acc = 0
+for epoch in range(EPOCHS):
+ print(f"Epoch {epoch}:")
+ model.train()
+ train(model, epoch, log_interval=20)
+
+ val_acc = []
+ # validate
+ val_acc.append(validate(model, epoch))
+
+ print(f"Validation Accuracy: {np.mean(val_acc) * 100:.2f}%")
+
+ if np.mean(val_acc) > best_acc:
+ print("New Best Validation Accuracy. Saving...")
+ best_acc = np.mean(val_acc)
+ torch.save(model.state_dict(), "model_data/m5_ann")
+
+ scheduler.step()
+
+ print(f"---------------------\n")
+
+# Load the weights into the network for inference
+model.load_state_dict(torch.load("model_data/m5_ann")) \ No newline at end of file