Spec-Zone.ru › PyTorch 2

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 перед сравнением перемещаются на CPU.
  • 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 – Если из входных данных нельзя создать torch.Tensor.
  • ValueError – Если указан только rtol или atol.
  • AssertionError – Если соответствующие входные данные не являются Python скалярами и не напрямую связаны.
  • AssertionError – Если allow_subclasses — False, но соответствующие входные данные не являются Python скалярами и имеют разные типы.
  • AssertionError – Если входные данные являются Sequence, но их длина не совпадает.
  • AssertionError – Если входные данные являются Mapping, но их наборы ключей не совпадают.
  • AssertionError – Если соответствующие тензоры не имеют одинакового shape.
  • и т.д. (остальные исключения переводятся аналогично)

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

dtype

rtol

atol

float16

1e-3

1e-5

bfloat16

1.6e-2

1e-5

... ... ...

Примечание

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!

Expected 1e-10 but got 1e-09.
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!

Expected nan but got nan.
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, memory_format=None) [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

Параметры
  • форма (Кортеж[int, ...]) – Целое число или последовательность целых чисел, определяющих форму выходного тензора.
  • тип данных (torch.dtype) – Тип данных возвращаемого тензора.
  • устройство (Объединение[str, torch.device]) – Устройство возвращаемого тензора.
  • низкий (Необязательный[Число]) – Устанавливает нижнюю границу (включительно) заданного диапазона. Если число указано, оно ограничивается наименьшим представимым конечным значением данного типа данных. Когда None (по умолчанию), это значение определяется на основе dtype (см. таблицу выше). По умолчанию: None.
  • высокий (Необязательный[Число]) –

    Устанавливает верхнюю границу (исключительно) заданного диапазона. Если число указано, оно ограничивается наибольшим представимым конечным значением данного типа данных. Когда None (по умолчанию), это значение определяется на основе dtype (см. таблицу выше). По умолчанию: None.

    Устарело начиная с версии 2.1: Передача low==high в make_tensor() для численных или комплексных типов устарело начиная с версии 2.1 и будет удалено в версии 2.3. Используйте torch.full() вместо этого.

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

Tensor

Примеры

>>> 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')
torch.testing.assert_allclose(actual, expected, rtol=None, atol=None, equal_nan=True, msg='') [source]

Предупреждение

torch.testing.assert_allclose() устарело начиная с 1.12 и будет удалено в будущей версии. Пожалуйста, используйте torch.testing.assert_close() вместо этого. Подробные инструкции по обновлению вы найдете здесь.

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

Spec-Zone.ru

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