mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-12 12:40:20 +08:00
Tests: refactor cleanup (#1744)
* wip * cleaning * optim imports * - * default hparams * fix restore * fix imports
This commit is contained in:
+2
-59
@@ -1,61 +1,4 @@
|
||||
"""Models for testing."""
|
||||
|
||||
import torch
|
||||
|
||||
from tests.base.eval_model_template import EvalModelTemplate
|
||||
from tests.base.mixins import (
|
||||
LightEmptyTestStep,
|
||||
LightValidationStepMixin,
|
||||
LightValidationMixin,
|
||||
LightValidationStepMultipleDataloadersMixin,
|
||||
LightValidationMultipleDataloadersMixin,
|
||||
LightTestStepMixin,
|
||||
LightTestMixin,
|
||||
LightTestStepMultipleDataloadersMixin,
|
||||
LightTestMultipleDataloadersMixin,
|
||||
LightTestFitSingleTestDataloadersMixin,
|
||||
LightTestFitMultipleTestDataloadersMixin,
|
||||
LightValStepFitSingleDataloaderMixin,
|
||||
LightValStepFitMultipleDataloadersMixin,
|
||||
LightTrainDataloader,
|
||||
LightValidationDataloader,
|
||||
LightTestDataloader,
|
||||
LightInfTrainDataloader,
|
||||
LightInfValDataloader,
|
||||
LightInfTestDataloader,
|
||||
LightTestOptimizerWithSchedulingMixin,
|
||||
LightTestMultipleOptimizersWithSchedulingMixin,
|
||||
LightTestOptimizersWithMixedSchedulingMixin,
|
||||
LightTestReduceLROnPlateauMixin,
|
||||
LightTestNoneOptimizerMixin,
|
||||
LightZeroLenDataloader
|
||||
)
|
||||
from tests.base.models import TestModelBase, DictHparamsModel
|
||||
|
||||
|
||||
class LightningTestModel(LightTrainDataloader,
|
||||
LightValidationMixin,
|
||||
LightTestMixin,
|
||||
TestModelBase):
|
||||
"""Most common test case. Validation and test dataloaders."""
|
||||
|
||||
def on_training_metrics(self, logs):
|
||||
logs['some_tensor_to_test'] = torch.rand(1)
|
||||
|
||||
|
||||
class LightningTestModelWithoutHyperparametersArg(LightningTestModel):
|
||||
"""Without hparams argument in constructor """
|
||||
|
||||
def __init__(self):
|
||||
import tests.base.utils as tutils
|
||||
|
||||
# the user loads the hparams in some other way
|
||||
hparams = tutils.get_default_hparams()
|
||||
super().__init__(hparams)
|
||||
|
||||
|
||||
class LightningTestModelWithUnusedHyperparametersArg(LightningTestModelWithoutHyperparametersArg):
|
||||
"""It has hparams argument in constructor but is not used."""
|
||||
|
||||
def __init__(self, hparams):
|
||||
super().__init__()
|
||||
from tests.base.datasets import TrialMNIST
|
||||
from tests.base.model_template import EvalModelTemplate
|
||||
|
||||
Reference in New Issue
Block a user