MCPcopy Create free account
hub / github.com/Anoise/WTFlib / train

Function train

LDPS_Graph/layers/utils.py:193–226  ·  view source on GitHub ↗
(model, train_loader, optimizer, epoch, device, verbose = 0,
    lossFn = None, lr_schedule=None, 
    post_proc = lambda args: args)

Source from the content-addressed store, hash-verified

191
192
193def train(model, train_loader, optimizer, epoch, device, verbose = 0,
194 lossFn = None, lr_schedule=None,
195 post_proc = lambda args: args):
196
197 if lossFn is None:
198 lossFn = nn.MSELoss()
199
200 model.train()
201
202 total_loss = 0.
203
204 for batch_idx, (data, target) in enumerate(train_loader):
205
206 bs = len(data)
207 data, target = data.to(device), target.to(device)
208 optimizer.zero_grad()
209
210 output = model(data)
211
212 target = post_proc(target)
213 output = post_proc(output)
214 loss = lossFn(output.view(bs, -1), target.view(bs, -1))
215
216 loss.backward()
217 optimizer.step()
218 total_loss += loss.sum().item()
219 if lr_schedule is not None: lr_schedule.step()
220
221 if verbose>0:
222 print('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format(
223 epoch, batch_idx * len(data), len(train_loader.dataset),
224 100. * batch_idx / len(train_loader), loss.item()))
225
226 return total_loss/len(train_loader.dataset)
227
228
229def test(model, test_loader, device, verbose=0, lossFn=None,

Callers

nothing calls this directly

Calls 2

stepMethod · 0.80
trainMethod · 0.45

Tested by

no test coverage detected