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являются с шагом, без квантования, вещественными и конечными, они считаются близкими, еслиНеконечные значения (
-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, в этом случае они считаются близкими, если их структура совпадает, и все их элементы считаются близкими в соответствии с вышеуказанным определением.
- Параметры:
-
- 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_subclassesFalse, но соответствующие входные данные не являются скалярами Python и имеют разные типы. - AssertionError – Если входные данные являются последовательностями, но их длина не совпадает.
- AssertionError – Если входные данные являются отображениями, но их наборы ключей не совпадают.
-
AssertionError – Если соответствующие тензоры не имеют одинакового
shape. -
AssertionError – Если
check_layoutTrue, но соответствующие тензоры не имеют одинаковогоlayout. - AssertionError – Если только один из соответствующих тензоров квантован.
-
AssertionError – Если соответствующие тензоры квантованы, но имеют разные
qscheme(). -
AssertionError – Если
check_deviceTrue, но соответствующие тензоры не находятся на одномdevice. -
AssertionError – Если
check_dtypeTrue, но соответствующие тензоры не имеют одинаковогоdtype. -
AssertionError – Если
check_strideTrue, но соответствующие тензоры со смещением не имеют одинакового шага. - AssertionError – Если значения соответствующих тензоров не близки в соответствии с вышеизложенным определением.
В следующей таблице показаны значения по умолчанию
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-5other
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! 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.dtypelowhighтип boolean
02тип целочисленный беззнаковый
010целочисленные типы со знаком
-910типы с плавающей точкой
-99типы комплексных чисел
-99- Параметры:
-
- 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’sfinfo()), а для типов комплексных чисел — на комплексное число, действительная и мнимая части которого равны наименьшему положительному нормализованному числу, представимому типом комплексных чисел. По умолчаниюFalse.
- Исключения:
-
-
ValueError – если
requires_grad=Trueпередано для целочисленногоdtype -
ValueError – Если
low > high. -
ValueError – Если либо
lowилиhighравноnan. -
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')
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/testing.html