Spec-Zone.ru › PyTorch 1

torch.testing

torch.testing.assert_close(actual, expected, *, allow_subclasses=True, rtol=None, atol=None, equal_nan=False, check_device=True, check_dtype=True, check_layout=True, check_stride=False, msg=None) [source]

Проверяет, что actual и expected близки.

Если actual и expected являются с шагом, без квантования, вещественными и конечными, они считаются близкими, если

∣actual−expected∣≤atol+rtol⋅∣expected∣\lvert \text{actual} - \text{expected} \rvert \le \texttt{atol} + \texttt{rtol} \cdot \lvert \text{expected} \rvert

Неконечные значения (-inf и inf) считаются близкими только тогда и только тогда, когда они равны. NaN считаются равными друг другу только в том случае, если equal_nan является True.

Кроме того, они считаются близкими, если у них одинаковы: - device (если check_device является True), - dtype (если check_dtype является True), - layout (если check_layout является True), и - шаг (если check_stride является True). Если хотя бы actual или expected является мета-тензором, будут проверены только атрибуты.

Если actual и expected являются разреженными (имеют структуру COO, CSR, CSC, BSR или BSC), их с шагом члены проверяются по отдельности. Индексы, а именно indices для COO, crow_indices и col_indices для CSR и BSR, или ccol_indices и row_indices для CSC и BSC соответственно, всегда проверяются на равенство, а значения проверяются на близость в соответствии с определением выше.

Если actual и expected квантованы, они считаются близкими, если у них одинаковая qscheme() и результат dequantize() близок в соответствии с определением выше.

actual и expected могут быть Tensor или любыми тензорно-или-скалярными значениями, из которых torch.Tensor могут быть построены с помощью torch.as_tensor(). За исключением Python-скаляров, типы входных данных должны быть напрямую связаны. Кроме того, actual и expected могут быть Sequence или Mapping, в этом случае они считаются близкими, если их структура совпадает, и все их элементы считаются близкими в соответствии с вышеуказанным определением.

Примечание

Python-скаляры являются исключением из требования связи типов, так как их type(), т.е. int, float и complex, эквивалентна dtype тензорного объекта. Таким образом, можно проверять Python-скаляры разных типов, но требуется check_dtype=False.

END_OF_DOCUMENT_MARKER
Параметры:
  • actual (Any) – Фактический ввод.
  • expected (Any) – Ожидаемый ввод.
  • allow_subclasses (bool) – Если True, (по умолчанию) и за исключением Python-скаляров, допускаются входные данные непосредственно связанных типов. В противном случае требуется равенство типов.
  • rtol (Optional[float]) – Относительная погрешность. Если указано, atol, также должно быть указано. Если опущено, значения по умолчанию выбираются на основе dtype, в соответствии с таблицей ниже.
  • atol (Optional[float]) – Абсолютная погрешность. Если указано, rtol, также должно быть указано. Если опущено, значения по умолчанию выбираются на основе dtype, в соответствии с таблицей ниже.
  • equal_nan (Union[bool, str]) – Если True, два NaN значения будут считаться равными.
  • check_device (bool) – Если True (по умолчанию), проверяется, что соответствующие тензоры находятся на одном device. Если эта проверка отключена, тензоры на разных device перемещаются на ЦП перед сравнением.
  • check_dtype (bool) – Если True (по умолчанию), проверяется, что соответствующие тензоры имеют одинаковый dtype. Если эта проверка отключена, тензоры с разными dtype преобразуются к общему dtype (в соответствии с torch.promote_types()) перед сравнением.
  • check_layout (bool) – Если True (по умолчанию), проверяется, что соответствующие тензоры имеют одинаковый layout. Если эта проверка отключена, тензоры с разными layout преобразуются в тензоры с шагом перед сравнением.
  • check_stride (bool) – Если True и соответствующие тензоры имеют шаг, проверяется, что у них одинаковый шаг.
  • msg (Optional[Union[str, Callable[[str], str]]]) – Необязательное сообщение об ошибке для использования в случае сбоя при сравнении. Также может быть передано как вызываемая функция, в этом случае она будет вызвана с сгенерированным сообщением и должна вернуть новое сообщение.
Исключения:
  • ValueError – Если из ввода нельзя построить тензор.
  • ValueError – Если указан только rtol или atol.
  • NotImplementedError – Если тензор является мета-тензором. Это временное ограничение и будет снято в будущем.
  • AssertionError – Если соответствующие входные данные не являются скалярами Python и не связаны напрямую.
  • AssertionError – Если allow_subclasses False, но соответствующие входные данные не являются скалярами Python и имеют разные типы.
  • AssertionError – Если входные данные являются последовательностями, но их длина не совпадает.
  • AssertionError – Если входные данные являются отображениями, но их наборы ключей не совпадают.
  • AssertionError – Если соответствующие тензоры не имеют одинакового shape.
  • AssertionError – Если check_layout True, но соответствующие тензоры не имеют одинакового layout.
  • AssertionError – Если только один из соответствующих тензоров квантован.
  • AssertionError – Если соответствующие тензоры квантованы, но имеют разные qscheme().
  • AssertionError – Если check_device True, но соответствующие тензоры не находятся на одном device.
  • AssertionError – Если check_dtype True, но соответствующие тензоры не имеют одинакового dtype.
  • AssertionError – Если check_stride True, но соответствующие тензоры со смещением не имеют одинакового шага.
  • AssertionError – Если значения соответствующих тензоров не близки в соответствии с вышеизложенным определением.

В следующей таблице показаны значения по умолчанию rtol и atol для разных dtype. В случае несовпадения dtype используется максимальное значение обеих погрешностей.

dtype

rtol

atol

float16

1e-3

1e-5

bfloat16

1.6e-2

1e-5

float32

1.3e-6

1e-5

float64

1e-7

1e-7

complex32

1e-3

1e-5

complex64

1.3e-6

1e-5

complex128

1e-7

1e-7

quint8

1.3e-6

1e-5

quint2x4

1.3e-6

1e-5

quint4x2

1.3e-6

1e-5

qint8

1.3e-6

1e-5

qint32

1.3e-6

1e-5

other

0.0

0.0

Примечание

assert_close() имеет высоко настраиваемые значения по умолчанию. Пользователям рекомендуется partial() его для соответствия их случаю использования. Например, если требуется проверка на равенство, можно определить assert_equal, который использует нулевые значения допуска для всех dtype по умолчанию:

>>> import functools
>>> assert_equal = functools.partial(torch.testing.assert_close, rtol=0, atol=0)
>>> assert_equal(1e-9, 1e-10)
Traceback (most recent call last):
...
AssertionError: Scalars are not equal!

Absolute difference: 9.000000000000001e-10
Relative difference: 9.0

Примеры

>>> # tensor to tensor comparison
>>> expected = torch.tensor([1e0, 1e-1, 1e-2])
>>> actual = torch.acos(torch.cos(expected))
>>> torch.testing.assert_close(actual, expected)
>>> # scalar to scalar comparison
>>> import math
>>> expected = math.sqrt(2.0)
>>> actual = 2.0 / math.sqrt(2.0)
>>> torch.testing.assert_close(actual, expected)
>>> # numpy array to numpy array comparison
>>> import numpy as np
>>> expected = np.array([1e0, 1e-1, 1e-2])
>>> actual = np.arccos(np.cos(expected))
>>> torch.testing.assert_close(actual, expected)
>>> # sequence to sequence comparison
>>> import numpy as np
>>> # The types of the sequences do not have to match. They only have to have the same
>>> # length and their elements have to match.
>>> expected = [torch.tensor([1.0]), 2.0, np.array(3.0)]
>>> actual = tuple(expected)
>>> torch.testing.assert_close(actual, expected)
>>> # mapping to mapping comparison
>>> from collections import OrderedDict
>>> import numpy as np
>>> foo = torch.tensor(1.0)
>>> bar = 2.0
>>> baz = np.array(3.0)
>>> # The types and a possible ordering of mappings do not have to match. They only
>>> # have to have the same set of keys and their elements have to match.
>>> expected = OrderedDict([("foo", foo), ("bar", bar), ("baz", baz)])
>>> actual = {"baz": baz, "bar": bar, "foo": foo}
>>> torch.testing.assert_close(actual, expected)
>>> expected = torch.tensor([1.0, 2.0, 3.0])
>>> actual = expected.clone()
>>> # By default, directly related instances can be compared
>>> torch.testing.assert_close(torch.nn.Parameter(actual), expected)
>>> # This check can be made more strict with allow_subclasses=False
>>> torch.testing.assert_close(
...     torch.nn.Parameter(actual), expected, allow_subclasses=False
... )
Traceback (most recent call last):
...
TypeError: No comparison pair was able to handle inputs of type
<class 'torch.nn.parameter.Parameter'> and <class 'torch.Tensor'>.
>>> # If the inputs are not directly related, they are never considered close
>>> torch.testing.assert_close(actual.numpy(), expected)
Traceback (most recent call last):
...
TypeError: No comparison pair was able to handle inputs of type <class 'numpy.ndarray'>
and <class 'torch.Tensor'>.
>>> # Exceptions to these rules are Python scalars. They can be checked regardless of
>>> # their type if check_dtype=False.
>>> torch.testing.assert_close(1.0, 1, check_dtype=False)
>>> # NaN != NaN by default.
>>> expected = torch.tensor(float("Nan"))
>>> actual = expected.clone()
>>> torch.testing.assert_close(actual, expected)
Traceback (most recent call last):
...
AssertionError: Scalars are not close!

Absolute difference: nan (up to 1e-05 allowed)
Relative difference: nan (up to 1.3e-06 allowed)
>>> torch.testing.assert_close(actual, expected, equal_nan=True)
>>> expected = torch.tensor([1.0, 2.0, 3.0])
>>> actual = torch.tensor([1.0, 4.0, 5.0])
>>> # The default error message can be overwritten.
>>> torch.testing.assert_close(actual, expected, msg="Argh, the tensors are not close!")
Traceback (most recent call last):
...
AssertionError: Argh, the tensors are not close!
>>> # If msg is a callable, it can be used to augment the generated message with
>>> # extra information
>>> torch.testing.assert_close(
...     actual, expected, msg=lambda msg: f"Header\n\n{msg}\n\nFooter"
... )
Traceback (most recent call last):
...
AssertionError: Header

Tensor-likes are not close!

Mismatched elements: 2 / 3 (66.7%)
Greatest absolute difference: 2.0 at index (1,) (up to 1e-05 allowed)
Greatest relative difference: 1.0 at index (1,) (up to 1.3e-06 allowed)

Footer
torch.testing.make_tensor(*shape, dtype, device, low=None, high=None, requires_grad=False, noncontiguous=False, exclude_zero=False) [source]

Создаёт тензор с заданными shape, device, и dtype, заполненный значениями, равномерно взятыми из [low, high).

Если low или high указаны и находятся вне диапазона представленных конечных значений dtype, они обрезаются до наименьшего или наибольшего конечного представленного значения соответственно. Если None, тогда следующая таблица описывает значения по умолчанию для low и high, которые зависят от dtype.

dtype

low

high

тип boolean

0

2

тип целочисленный беззнаковый

0

10

целочисленные типы со знаком

-9

10

типы с плавающей точкой

-9

9

типы комплексных чисел

-9

9

Параметры:
  • shape (Кортеж[int, ...]) – Целое число или последовательность целых чисел, определяющая форму выходного тензора.
  • dtype (torch.dtype) – Тип данных возвращаемого тензора.
  • device (Объединение[строка, torch.device]) – Устройство возвращаемого тензора.
  • low (Необязательно[Число]) – Устанавливает нижнюю границу (включительно) заданного диапазона. Если задано число, оно обрезается до наименьшего представимого конечного значения для данного типа данных. При None (по умолчанию), это значение определяется на основе dtype (см. таблицу выше). По умолчанию: None.
  • high (Необязательно[Число]) – Устанавливает верхнюю границу (исключительно) заданного диапазона. Если задано число, оно обрезается до наибольшего представимого конечного значения для данного типа данных. При None (по умолчанию), это значение определяется на основе dtype (см. таблицу выше). По умолчанию: None.
  • requires_grad (Необязательно[bool]) – Если автограду следует записывать операции над возвращаемым тензором. По умолчанию: False.
  • noncontiguous (Необязательно[bool]) – Если True, возвращаемый тензор будет несмежным. Этот аргумент игнорируется, если создаваемый тензор содержит меньше двух элементов.
  • exclude_zero (Необязательно[bool]) – Если True , то нули заменяются на малое положительное значение типа данных в зависимости от dtype. Для булевых и целочисленных типов ноль заменяется единицей. Для типов с плавающей точкой он заменяется на наименьшее положительное нормализованное число типа данных (значение «tiny» объекта dtype’s finfo() ), а для типов комплексных чисел — на комплексное число, действительная и мнимая части которого равны наименьшему положительному нормализованному числу, представимому типом комплексных чисел. По умолчанию False.
Исключения:
  • ValueError – если requires_grad=True передано для целочисленного dtype
  • ValueError – Если low > high.
  • ValueError – Если либо low или high равно nan.
  • TypeError – Если dtype не поддерживается этой функцией.
Тип возвращаемого значения:

Тензор

Примеры

>>> from torch.testing import make_tensor
>>> # Creates a float tensor with values in [-1, 1)
>>> make_tensor((3,), device='cpu', dtype=torch.float32, low=-1, high=1)
tensor([ 0.1205, 0.2282, -0.6380])
>>> # Creates a bool tensor on CUDA
>>> make_tensor((2, 2), device='cuda', dtype=torch.bool)
tensor([[False, False],
        [False, True]], device='cuda:0')

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/testing.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API