mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-21 13:20:08 +08:00
25 lines
682 B
Python
25 lines
682 B
Python
"""Models for testing."""
|
|
|
|
import torch
|
|
|
|
from .base import LightningTestModelBase
|
|
from .mixins import (
|
|
LightningValidationStepMixin,
|
|
LightningValidationMixin,
|
|
LightningValidationStepMultipleDataloadersMixin,
|
|
LightningValidationMultipleDataloadersMixin,
|
|
LightningTestStepMixin,
|
|
LightningTestMixin,
|
|
LightningTestStepMultipleDataloadersMixin,
|
|
LightningTestMultipleDataloadersMixin,
|
|
)
|
|
|
|
|
|
class LightningTestModel(LightningValidationMixin, LightningTestMixin, LightningTestModelBase):
|
|
"""
|
|
Most common test case. Validation and test dataloaders.
|
|
"""
|
|
|
|
def on_training_metrics(self, logs):
|
|
logs['some_tensor_to_test'] = torch.rand(1)
|