Spec-Zone.ru › PyTorch 2.14

torch.testing

Создано: 7 мая 2021 г. | Последнее обновление: 10 июня 2025 г.

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) [исходный код]

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

Если actual и expected имеют strided-разметку, не квантованы, содержат действительные и конечные значения, они считаются близкими, если

∣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), их элементы со strided-разметкой проверяются по отдельности. Индексы, а именно 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.

Параметры:
  • 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 преобразуются в тензоры со strided-разметкой перед сравнением.
  • check_stride (bool) – Если True и соответствующие тензоры имеют strided-разметку, проверяет, что их шаги совпадают.
  • 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.
  • AssertionError – Если check_layout — True, но соответствующие тензоры имеют разные layout.
  • AssertionError – Если только один из соответствующих тензоров квантован.
  • AssertionError – Если соответствующие тензоры квантованы, но их значения qscheme() различаются.
  • AssertionError – Если check_device — True, но соответствующие тензоры находятся на разных device.
  • AssertionError – Если check_dtype — True, но соответствующие тензоры имеют разные dtype.
  • AssertionError – Если check_stride — True, но соответствующие тензоры со strided-разметкой имеют разные шаги.
  • 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

другой

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!

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) [исходный код]

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

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

dtype

low

high

логический тип

0

2

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

0

10

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

-9

10

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

-9

9

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

-9

9

Параметры:
  • shape (Tuple[int, ...]) – Одно целое число или последовательность целых чисел, задающая форму выходного тензора.
  • dtype (torch.dtype) – Тип данных возвращаемого тензора.
  • device (Union[str, torch.device]) – Устройство возвращаемого тензора.
  • low (Optional[Number]) – Задаёт нижнюю границу указанного диапазона (включительно). Переданное число ограничивается наименьшим конечным значением, представимым указанным dtype. Если None (значение по умолчанию), это значение определяется на основе dtype (см. таблицу выше). По умолчанию: None.
  • high (Optional[Number]) –

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

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

  • requires_grad (Optional[bool]) – Следует ли autograd записывать операции с возвращаемым тензором. По умолчанию: False.
  • noncontiguous (Optional[bool]) – Если True, возвращаемый тензор будет неконтiguousным. Этот аргумент игнорируется, если созданный тензор содержит менее двух элементов. Несовместим с memory_format.
  • exclude_zero (Optional[bool]) – Если True, нули заменяются на малое положительное значение dtype в зависимости от dtype. Для логических и целочисленных типов ноль заменяется на единицу. Для типов с плавающей точкой он заменяется на наименьшее положительное нормальное число данного dtype (значение «tiny» объекта dtype finfo()), а для комплексных типов — на комплексное число, действительная и мнимая части которого равны наименьшему положительному нормальному числу, представимому комплексным типом. По умолчанию False.
  • memory_format (Optional[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='') [исходный код]

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

torch.testing.assert_allclose() устарела с версии 1.12 и будет удалена в одном из будущих выпусков. Вместо неё используйте torch.testing.assert_close(). Подробные инструкции по переходу можно найти здесь.

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

Spec-Zone.ru

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