MCPcopy Create free account
hub / github.com/Intelligent-Computing-Lab-Panda/NDA_SNN / train

Function train

main.py:41–73  ·  view source on GitHub ↗
(model, device, train_loader, criterion, optimizer, epoch, scaler, args)

Source from the content-addressed store, hash-verified

39
40
41def train(model, device, train_loader, criterion, optimizer, epoch, scaler, args):
42 running_loss = 0
43 model.train()
44 M = len(train_loader)
45 total = 0
46 correct = 0
47 s_time = time.time()
48 for i, (images, labels) in enumerate(train_loader):
49 optimizer.zero_grad()
50 labels = labels.to(device)
51 images = images.to(device)
52
53 if args.amp:
54 with autocast(device_type='cuda', dtype=torch.float16):
55 outputs = model(images)
56 mean_out = outputs.mean(1)
57 loss = criterion(mean_out, labels)
58 scaler.scale(loss.mean()).backward()
59 scaler.step(optimizer)
60 scaler.update()
61 else:
62 outputs = model(images)
63 mean_out = outputs.mean(1)
64 loss = criterion(mean_out, labels)
65 loss.mean().backward()
66 optimizer.step()
67
68 running_loss += loss.item()
69 total += float(labels.size(0))
70 _, predicted = mean_out.cpu().max(1)
71 correct += float(predicted.eq(labels.cpu()).sum().item())
72 e_time = time.time()
73 return running_loss / M, 100 * correct / total, (e_time-s_time)/60
74
75
76@torch.no_grad()

Callers 1

main.pyFile · 0.85

Calls 1

backwardMethod · 0.80

Tested by

no test coverage detected