mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-19 13:00:36 +08:00
* rename validate -> evaluate; implement test logic; allow multiple test_loaders * add test_step and test_end to LightningModule * add in_test_mode to pretraining to implement case 2 (test pretrained model) * fix code style issues * LightningTestModel: add optional second test set, implement test_step and test_end * implemented test for multiple test_dataloaders; fixed typo * add two test cases for #89 * add documentation for test_step, test_end; fix computation of loss in validation_step example * Update trainer.py * Update trainer.py * Update trainer.py * Update trainer.py * Update trainer.py * Update trainer.py * Added proper dp ddp routing calls for test mode * Update trainer.py * Update test_models.py * Update trainer.py * Update trainer.py * Update override_data_parallel.py * Update test_models.py * Update test_models.py * Update trainer.py * Update trainer.py * Update trainer.py * Update test_models.py * Update test_models.py * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * debug * Update trainer.py * Update override_data_parallel.py * Update debug.py * Update lm_test_module.py * Update test_models.py
54 lines
1.0 KiB
Python
54 lines
1.0 KiB
Python
import torch
|
|
|
|
|
|
class ModelHooks(torch.nn.Module):
|
|
|
|
def on_sanity_check_start(self):
|
|
"""
|
|
Called before starting evaluate
|
|
:return:
|
|
"""
|
|
pass
|
|
|
|
def on_batch_start(self, data_batch):
|
|
pass
|
|
|
|
def on_batch_end(self):
|
|
pass
|
|
|
|
def on_epoch_start(self):
|
|
pass
|
|
|
|
def on_epoch_end(self):
|
|
pass
|
|
|
|
def on_pre_performance_check(self):
|
|
pass
|
|
|
|
def on_post_performance_check(self):
|
|
pass
|
|
|
|
def on_tng_metrics(self, metrics):
|
|
pass
|
|
|
|
def on_before_zero_grad(self, optimizer):
|
|
"""
|
|
Called after optimizer.step() and before optimizer.zero_grad()
|
|
|
|
for optimizer in optimizers:
|
|
optimizer.step()
|
|
model.on_before_zero_grad(optimizer) # < ---- called here
|
|
optimizer.zero_grad
|
|
|
|
:param optimizer:
|
|
:return:
|
|
"""
|
|
pass
|
|
|
|
def on_after_backward(self):
|
|
"""
|
|
Called after loss.backward() and before optimizers do anything
|
|
:return:
|
|
"""
|
|
pass
|