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перед сравнением перемещаются на 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. - и т.д. (остальные исключения переводятся аналогично)
-
ValueError – Если из входных данных нельзя создать
В следующей таблице показаны значения по умолчанию
rtolиatolдля различныхdtype. В случае несовпадающихdtypeиспользуется максимальное значение обеих погрешностей.dtypertolatolfloat161e-31e-5bfloat161.6e-21e-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.dtypelowhighТип boolean
02Тип целого беззнакового
010Типы целых со знаком
-910Плавающие типы
-99Комплексные типы
-99- Параметры
-
- форма (Кортеж[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не поддерживается этой функцией.
-
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='')[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