Spec-Zone.ru › PyTorch 2.14

Управление потоком — 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/#prototype

cond — оператор структурированного управления потоком. Иными словами, он похож на инструкцию 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 не требуют входных данных. Значение по умолчанию — ().
Тип возвращаемого значения:

Any

Пример:

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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API