Spec-Zone.ru › PyTorch 2.14

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/#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), находящаяся в области видимости, которая трассируется. Ветви true и false должны иметь согласованные входные и выходные данные: входные данные должны быть одинаковыми, а выходные данные — одного типа и формы. Также допускается целочисленный выход. Выходные данные будут сделаны динамическими путём преобразования в symint.
  • operands (Tuple of possibly nested dict/list/tuple of torch.Tensor) – кортеж входных данных для функций true/false. Он может быть пустым, если 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) должна соответствовать всем следующим ограничениям:

    • Сигнатура функции должна соответствовать 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

Spec-Zone.ru

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