mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
* New metric classes (#1326) * Create metrics package * Create metric.py * Create utils.py * Create __init__.py * add tests for metric utils * add docstrings for metrics utils * add function to recursively apply other function to collection * add tests for this function * update test * Update pytorch_lightning/metrics/metric.py Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * update metric name * remove example docs * fix tests * add metric tests * fix to tensor conversion * fix apply to collection * Update CHANGELOG.md * Update pytorch_lightning/metrics/metric.py Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * remove tests from init * add missing type annotations * rename utils to convertors * Create metrics.rst * Update index.rst * Update index.rst * Update pytorch_lightning/metrics/convertors.py Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * Update pytorch_lightning/metrics/convertors.py Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * Update pytorch_lightning/metrics/convertors.py Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * Update pytorch_lightning/metrics/metric.py Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * Update tests/utilities/test_apply_to_collection.py Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * Update tests/utilities/test_apply_to_collection.py Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * Update tests/metrics/convertors.py Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * Apply suggestions from code review Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * add doctest example * rename file and fix imports * added parametrized test * replace lambda with inlined function * rename apply_to_collection to apply_func * Separated class description from init args * Apply suggestions from code review Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * adjust random values * suppress output when seeding * remove gpu from doctest * Add requested changes and add ellipsis for doctest * forgot to push these files... * add explicit check for dtype to convert to * fix ddp tests * remove explicit ddp destruction Co-authored-by: Jirka Borovec <Borda@users.noreply.github.com> * move dtype device mixin to more general place * refactor to general device dtype mixin * add initial metric package description * change default to none for mac os * pep8 * fix import * Update index.rst * Update ci-testing.yml * Apply suggestions from code review Co-authored-by: Adrian Wälchli <aedu.waelchli@gmail.com> * Update CHANGELOG.md * Update pytorch_lightning/metrics/converters.py * readme * Update metric.py * Update pytorch_lightning/metrics/converters.py Co-authored-by: Jirka Borovec <Borda@users.noreply.github.com> Co-authored-by: William Falcon <waf2107@columbia.edu> Co-authored-by: Adrian Wälchli <aedu.waelchli@gmail.com> Co-authored-by: Jirka <jirka@pytorchlightning.ai>
86 lines
2.4 KiB
Python
86 lines
2.4 KiB
Python
import numpy as np
|
|
import torch
|
|
|
|
from pytorch_lightning.metrics.metric import Metric, TensorMetric, NumpyMetric
|
|
|
|
|
|
class DummyTensorMetric(TensorMetric):
|
|
def __init__(self):
|
|
super().__init__('dummy')
|
|
|
|
def forward(self, input1, input2):
|
|
assert isinstance(input1, torch.Tensor)
|
|
assert isinstance(input2, torch.Tensor)
|
|
return 1.
|
|
|
|
|
|
class DummyNumpyMetric(NumpyMetric):
|
|
def __init__(self):
|
|
super().__init__('dummy')
|
|
|
|
def forward(self, input1, input2):
|
|
assert isinstance(input1, np.ndarray)
|
|
assert isinstance(input2, np.ndarray)
|
|
return 1.
|
|
|
|
|
|
def _test_metric(metric: Metric):
|
|
input1, input2 = torch.tensor([1.]), torch.tensor([2.])
|
|
|
|
def change_and_check_device_dtype(device, dtype):
|
|
metric.to(device=device, dtype=dtype)
|
|
|
|
metric_val = metric(input1, input2)
|
|
assert isinstance(metric_val, torch.Tensor)
|
|
|
|
if device is not None:
|
|
assert metric.device in [device, torch.device(device)]
|
|
assert metric_val.device in [device, torch.device(device)]
|
|
|
|
if dtype is not None:
|
|
assert metric.dtype == dtype
|
|
assert metric_val.dtype == dtype
|
|
|
|
devices = [None, 'cpu']
|
|
if torch.cuda.is_available():
|
|
devices += ['cuda:0']
|
|
|
|
for device in devices:
|
|
for dtype in [None, torch.float32, torch.float64]:
|
|
change_and_check_device_dtype(device=device, dtype=dtype)
|
|
|
|
if torch.cuda.is_available():
|
|
metric.cuda(0)
|
|
assert metric.device == torch.device('cuda', index=0)
|
|
assert metric(input1, input2).device == torch.device('cuda', index=0)
|
|
|
|
metric.cpu()
|
|
assert metric.device == torch.device('cpu')
|
|
assert metric(input1, input2).device == torch.device('cpu')
|
|
|
|
metric.type(torch.int8)
|
|
assert metric.dtype == torch.int8
|
|
assert metric(input1, input2).dtype == torch.int8
|
|
|
|
metric.float()
|
|
assert metric.dtype == torch.float32
|
|
assert metric(input1, input2).dtype == torch.float32
|
|
|
|
metric.double()
|
|
assert metric.dtype == torch.float64
|
|
assert metric(input1, input2).dtype == torch.float64
|
|
|
|
if torch.cuda.is_available():
|
|
metric.cuda()
|
|
metric.half()
|
|
assert metric.dtype == torch.float16
|
|
assert metric(input1, input2).dtype == torch.float16
|
|
|
|
|
|
def test_tensor_metric():
|
|
_test_metric(DummyTensorMetric())
|
|
|
|
|
|
def test_numpy_metric():
|
|
_test_metric(DummyNumpyMetric())
|