torch.cond
-
torch.cond(pred, true_fn, false_fn, operands=())[исходный код] -
Условно применяет
true_fnилиfalse_fn.Предупреждение
torch.cond— экспериментальная функция в PyTorch. Она имеет ограниченную поддержку типов входных и выходных данных. Ожидайте более стабильную реализацию в будущей версии PyTorch. Подробнее о классификации функций: https://pytorch.org/blog/pytorch-feature-classification-changes/#prototypecond— оператор структурированного потока управления. То есть он похож на оператор if в Python, но имеет ограничения наtrue_fn,false_fnиoperands, благодаря которым его можно захватывать с помощью torch.compile и torch.export.Если ограничения на аргументы
condсоблюдены,condэквивалентен следующему:def cond(pred, true_branch, false_branch, operands): if pred: return true_branch(*operands) else: return false_branch(*operands)- Параметры:
-
- pred (Union[bool, torch.Tensor]) – логическое выражение или тензор с одним элементом, указывающий, какую ветвь применить.
- true_fn (Callable) – вызываемая функция (a -> b), находящаяся в области видимости, которая трассируется.
- false_fn (Callable) – вызываемая функция (a -> b), находящаяся в области видимости, которая трассируется. Ветви true и false должны иметь согласованные входные и выходные данные: входные данные должны быть одинаковыми, а выходные данные — одного типа и формы. Также допускается целочисленный выход. Выходные данные будут сделаны динамическими путём преобразования в symint.
- operands (Tuple of possibly nested dict/list/tuple of torch.Tensor) – кортеж входных данных для функций true/false. Он может быть пустым, если true_fn/false_fn не требует входных данных. По умолчанию — ().
- Тип возвращаемого значения:
Пример:
def true_fn(x: torch.Tensor): return x.cos() def false_fn(x: torch.Tensor): return x.sin() return cond(x.shape[0] > 4, true_fn, false_fn, (x,))- Ограничения:
-
-
Условное выражение (также называемое
pred) должно соответствовать одному из следующих ограничений:- Это
torch.Tensor, содержащий только один элемент, с типом данных torch.bool - Это логическое выражение, например
x.shape[0] > 10илиx.dim() > 1 and x.shape[1] > 10
- Это
-
Функция ветви (также называемая
true_fn/false_fn) должна соответствовать всем следующим ограничениям:- Сигнатура функции должна соответствовать operands.
- Функция должна возвращать тензор с теми же метаданными, например формой, типом данных и т. д.
- Функция не может выполнять мутации на месте глобальных переменных. (Примечание: операции с тензорами на месте, такие как
add_для промежуточных результатов, разрешены в ветви.) -
Во время инференса функция может выполнять мутации на месте входных тензоров (то есть когда
torch.is_grad_enabled()имеет значение False). Примечание: при использованииtorch.compile()с непостоянным предикатом выходные данные всегда будут новыми тензорами, не имеющими идентичности объектов с исходными входными данными.Пример:
def true_fn(x): return x.sin_() def false_fn(x): return x + 1 def f(x): return cond(x.sum() > 0, true_fn, false_fn, (x,)) x = torch.ones(4) with torch.no_grad(): result = torch.compile(f)(x) assert result is not x # result is a new tensor, not the original x
-
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.cond.html