Управление потоком — Cond
Создано: 3 окт. 2023 г. | Последнее обновление: 14 февр. 2026 г.
torch.cond — это оператор структурированного управления потоком. Его можно использовать для задания управления потоком, аналогичного if-else, и логически представить как реализацию следующего вида.
def cond(
pred: Union[bool, torch.Tensor],
true_fn: Callable,
false_fn: Callable,
operands: Tuple[torch.Tensor]
):
if pred:
return true_fn(*operands)
else:
return false_fn(*operands)
Его уникальная особенность — возможность выражать управление потоком, зависящее от данных: он преобразуется в условный оператор (torch.ops.higher_order.cond), сохраняющий предикат, функцию для истинной ветви и функцию для ложной ветви. Это обеспечивает большую гибкость при написании и развертывании моделей, архитектура которых меняется в зависимости от значения или формы входных данных либо промежуточных результатов тензорных операций.
Предупреждение
torch.cond — экспериментальная функция PyTorch. Поддерживаются лишь некоторые типы входных и выходных данных. В будущих версиях PyTorch ожидается более стабильная реализация. Подробнее о классификации функций: https://pytorch.org/blog/pytorch-feature-classification-changes/#prototype
Примеры
Ниже приведен пример использования cond для выбора ветви в зависимости от формы входных данных:
import torch
def true_fn(x: torch.Tensor):
return x.cos()
def false_fn(x: torch.Tensor):
return x.sin()
class DynamicShapeCondPredicate(torch.nn.Module):
"""
A basic usage of cond based on dynamic shape predicate.
"""
def __init__(self):
super().__init__()
def forward(self, x: torch.Tensor) -> torch.Tensor:
return torch.cond(x.shape[0] > 4, true_fn, false_fn, (x,))
dyn_shape_mod = DynamicShapeCondPredicate()
Можно выполнить модель в обычном режиме и убедиться, что результаты меняются в зависимости от формы входных данных:
inp = torch.randn(3) inp2 = torch.randn(5) print(dyn_shape_mod(inp), false_fn(inp)) print(dyn_shape_mod(inp2), true_fn(inp2))
Модель можно экспортировать для дальнейших преобразований и развертывания. В результате получим экспортированную программу, показанную ниже:
inp = torch.randn(4, 3)
ep = torch.export.export(
DynamicShapeCondPredicate(),
(inp,),
dynamic_shapes={"x": {0: torch.export.Dim.DYNAMIC}}
)
print(ep)
Обратите внимание: torch.cond преобразуется в torch.ops.higher_order.cond, его предикат становится символьным выражением, зависящим от формы входных данных, а функции ветвей становятся двумя атрибутами подграфов модуля графа верхнего уровня.
Ниже приведен еще один пример, демонстрирующий, как выразить управление потоком, зависящее от данных:
def true_fn(x: torch.Tensor):
return x.cos() + x.sin()
def false_fn(x: torch.Tensor):
return x.sin()
class DataDependentCondPredicate(torch.nn.Module):
"""
A basic usage of cond based on data dependent predicate.
"""
def __init__(self):
super().__init__()
def forward(self, x: torch.Tensor) -> torch.Tensor:
return torch.cond(x.sum() > 4.0, true_fn, false_fn, (x,))
inp = torch.randn(4, 3)
ep = torch.export.export(DataDependentCondPredicate(), (inp,), dynamic_shapes={"x": {0: torch.export.Dim.DYNAMIC}})
print(ep)
Инварианты torch.ops.higher_order.cond
Для torch.ops.higher_order.cond важны несколько инвариантов:
-
Для предиката:
- Динамический характер предиката сохраняется (например,
gtв примере выше) - Если предикат в пользовательской программе является константой (например, константой типа bool в Python), то
predоператора будет константой.
- Динамический характер предиката сохраняется (например,
-
Для ветвей:
- Сигнатуры входных и выходных данных представляют собой расплющенный кортеж.
- Они являются
torch.fx.GraphModule. - Замыкания исходной функции становятся явными входными данными. Замыканий нет.
- Мутации входных данных или глобальных переменных не допускаются.
-
Для операндов:
- Они также представляют собой плоский кортеж.
- Вложенные вызовы
torch.condв пользовательской программе становятся вложенными модулями графа.
Справочник API
-
torch._higher_order_ops.cond.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), находящаяся в области видимости трассируемого кода. Истинная и ложная ветви должны иметь согласованные входные и выходные данные: входные данные должны совпадать, а тип и форма выходных данных должны быть одинаковыми. Также допускается выходное значение типа int. Оно преобразуется в symint, чтобы сделать выходное значение динамическим.
- operands (Tuple of possibly nested dict/list/tuple of torch.Tensor) – Кортеж входных данных для функций истинной и ложной ветвей. Он может быть пустым, если 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) должна удовлетворять всем следующим условиям:- Сигнатура функции должна соответствовать операндам.
- Функция должна возвращать тензор с теми же метаданными, например формой, типом данных и т. д.
- Функция не может выполнять мутации на месте глобальных переменных. (Примечание: тензорные операции на месте, такие как
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/higher_order_ops/cond.html