MCPcopy Create free account
hub / github.com/CHENGY12/PLOT / forward_backward

Method forward_backward

plot-pp/trainers/plotpp.py:333–355  ·  view source on GitHub ↗
(self, batch)

Source from the content-addressed store, hash-verified

331 self.scaler = GradScaler() if cfg.TRAINER.PLOTPP.PREC == "amp" else None
332
333 def forward_backward(self, batch):
334 image, label = self.parse_batch_train(batch)
335
336 prec = self.cfg.TRAINER.PLOTPP.PREC
337 if prec == "amp":
338 with autocast():
339 output = self.model(image, label)
340 self.optim.zero_grad()
341 self.scaler.scale(output).backward()
342 self.scaler.step(self.optim)
343 self.scaler.update()
344 else:
345 output = self.model(image)
346 loss = F.cross_entropy(output, label)
347 self.model_backward_and_update(loss)
348
349 loss_summary = {"loss": loss.item(),
350 "acc": compute_accuracy(output, label)[0].item()}
351
352 if (self.batch_idx + 1) == self.num_batches:
353 self.update_lr()
354
355 return loss_summary
356
357 def parse_batch_train(self, batch):
358 input = batch["img"]

Callers

nothing calls this directly

Calls 1

parse_batch_trainMethod · 0.95

Tested by

no test coverage detected