mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
added cpu model test
This commit is contained in:
@@ -108,7 +108,7 @@ class LightningTemplateModel(LightningModule):
|
||||
|
||||
output = OrderedDict({
|
||||
'val_loss': loss_val,
|
||||
'val_acc': torch.tensor(val_acc).cuda(loss_val.device.index),
|
||||
'val_acc': torch.tensor(val_acc).type(loss_val.dtype),
|
||||
})
|
||||
|
||||
# can also return just a scalar instead of a dict (return loss_val)
|
||||
|
||||
@@ -436,6 +436,10 @@ class Trainer(TrainerIO):
|
||||
|
||||
self.__run_pretrain_routine(model)
|
||||
|
||||
# return 1 when finished
|
||||
# used for testing or when we need to know that training succeeded
|
||||
return 1
|
||||
|
||||
def dp_train(self, model):
|
||||
|
||||
# CHOOSE OPTIMIZER
|
||||
|
||||
Reference in New Issue
Block a user