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-разметку, не квантованы, содержат действительные и конечные значения, они считаются близкими, еслиНефинитные значения (
-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 – Если значения соответствующих тензоров не являются близкими согласно приведённому выше определению.
-
ValueError – Если из входного аргумента невозможно создать
В следующей таблице приведены значения
rtolиatolпо умолчанию для различныхdtype. Если значенияdtypeне совпадают, используется большее из двух значений погрешности.dtypertolatolfloat161e-31e-5bfloat161.6e-21e-5float321.3e-61e-5float641e-71e-7complex321e-31e-5complex641.3e-61e-5complex1281e-71e-7quint81.3e-61e-5quint2x41.3e-61e-5quint4x21.3e-61e-5qint81.3e-61e-5qint321.3e-61e-5другой
0.00.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.dtypelowhighлогический тип
02беззнаковый целочисленный тип
010знаковые целочисленные типы
-910типы с плавающей точкой
-99комплексные типы
-99- Параметры:
-
- 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» объектаdtypefinfo()), а для комплексных типов — на комплексное число, действительная и мнимая части которого равны наименьшему положительному нормальному числу, представимому комплексным типом. По умолчанию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не поддерживается этой функцией.
-
ValueError – Если
- Тип возвращаемого значения:
Примеры
>>> 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