MCPcopy Create free account
hub / github.com/Sirwenhao/Deep-Learning-Notes / train_one_epoch

Function train_one_epoch

CV/Pytorch_classification/RegNet/utils.py:118–144  ·  view source on GitHub ↗
(model, optimizer, data_loader, device, epoch)

Source from the content-addressed store, hash-verified

116
117
118def train_one_epoch(model, optimizer, data_loader, device, epoch):
119 model.train()
120 loss_function = torch.nn.CrossEntropyLoss()
121 mean_loss = torch.zeros(1).to(device)
122 optimizer.zero_grad()
123
124 data_loader = tqdm(data_loader, file=sys.stdout)
125
126 for step, data in enumerate(data_loader):
127 images, labels = data
128
129 pred = model(images.to(device))
130
131 loss = loss_function(pred, labels.to(device))
132 loss.backward()
133 mean_loss = (mean_loss * step + loss.detach()) / (step + 1) # update mean losses
134
135 data_loader.desc = "[epoch {}] mean loss {}".format(epoch, round(mean_loss.item(), 3))
136
137 if not torch.isfinite(loss):
138 print('WARNING: non-finite loss, ending training ', loss)
139 sys.exit(1)
140
141 optimizer.step()
142 optimizer.zero_grad()
143
144 return mean_loss.item()
145
146
147@torch.no_grad()

Callers 1

mainFunction · 0.90

Calls 1

backwardMethod · 0.80

Tested by

no test coverage detected