Spec-Zone.ru › PyTorch 2.14

Поток управления — цикл while

Создано: Feb 14, 2026 | Последнее обновление: Feb 14, 2026

torch.while_loop — это оператор структурированного потока управления, который выполняет функцию тела цикла, пока условие истинно. Его логически можно представить в следующем виде:

def while_loop(
    cond_fn: Callable[..., bool],
    body_fn: Callable[..., tuple],
    carried_inputs: tuple,
):
    val = carried_inputs
    while cond_fn(*val):
        val = body_fn(*val)
    return val

Предупреждение

torch.while_loop — экспериментальная функция PyTorch. Поддерживаются лишь некоторые типы входных и выходных данных. Ожидайте более стабильную реализацию в будущей версии PyTorch. Подробнее о классификации функций: https://pytorch.org/blog/pytorch-feature-classification-changes/#prototype

Примеры

Ниже приведен простой пример использования while_loop для выполнения итераций до тех пор, пока не будет выполнено условие:

import torch
from torch._higher_order_ops import while_loop

class M(torch.nn.Module):

    def cond_fn(self, iter_count, x):
        return iter_count.sum() > 0

    def body_fn(self, iter_count, x):
        return iter_count - 1, x * 2

    def forward(self, init_iter, init_x):
        final_iter, final_x = while_loop(self.cond_fn, self.body_fn, (init_iter, init_x))
        return final_iter, final_x

m = M()

Мы можем выполнить модель в eager-режиме и ожидать, что результаты будут различаться в зависимости от формы входных данных:

_, final_x = m(torch.tensor([3]), torch.ones(3))
assert torch.equal(final_x, torch.ones(3) * 2**3)

_, final_x = m(torch.tensor([10]), torch.ones(3))
assert torch.equal(final_x, torch.ones(3) * 2**10)

Мы можем экспортировать модель для дальнейших преобразований и развертывания. В результате получаем экспортированную программу, сохраняющую структуру while_loop:

ep = torch.export.export(M(), (torch.tensor([10]), torch.ones(3)))
print(ep)

Обратите внимание, что функции условия и тела цикла становятся атрибутами подграфов модуля графа верхнего уровня.

Ограничения

  • body_fn должна возвращать тензоры или целые числа с теми же метаданными (формой, dtype), что и входные данные.
  • body_fn и cond_fn не должны изменять carried_inputs на месте. Перед изменением необходимо создать копию.
  • body_fn и cond_fn не должны изменять переменные Python (например, list/dict), созданные вне функции.
  • Выходные данные body_fn и cond_fn не могут ссылаться на те же объекты, что и входные данные. Необходимо создать копию.

Справочник API

torch._higher_order_ops.while_loop.while_loop(cond_fn, body_fn, carried_inputs) [исходный код]

Выполняет body_fn(*carried_inputs), пока cond_fn(*carried_inputs) возвращает скалярный тензор со значением True. Возвращает результат body_fn или исходные carried_inputs.

Предупреждение

torch.while_loop — экспериментальная функция PyTorch. Поддерживаются лишь некоторые типы входных и выходных данных; обучение пока не поддерживается. Ожидайте более стабильную реализацию в будущей версии PyTorch. Подробнее о классификации функций: https://pytorch.org/blog/pytorch-feature-classification-changes/#prototype

while_loop — это оператор структурированного потока управления. Он сохраняет семантику цикла при использовании torch.compile и torch.export.

while_loop эквивалентен следующему:

def while_loop(cond_fn, body_fn, carried_inputs):
    val = carried_inputs
    while cond_fn(*val):
        val = body_fn(*val)
    return val
Параметры:
  • cond_fn (Callable) – Вызываемая функция, возвращающая скалярный булев тензор или логическое значение Python.
  • body_fn (Callable) – Вызываемая функция, принимающая те же входные данные, что и cond_fn, и возвращающая кортеж тензоров или целых чисел.
  • carried_inputs (Tuple of possibly nested dict/list/tuple of tensors or ints) – Кортеж входных данных для cond_fn и body_fn. Это также начальные значения состояний, сохраняемых между итерациями. Обратите внимание: если в качестве carry передано целое число, соответствующее возвращаемое значение while_loop будет другим целым числом с неизвестным значением, поскольку количество итераций while_loop заранее неизвестно.

Пример 1:

def cond_fn(iter, x):
    return iter.sum() < 10


def body_fn(iter, x):
    return iter + 1, x.sin()


while_loop(cond_fn, body_fn, (torch.zeros(1), torch.randn(3, 4)))

Пример 2:

def cond_fn(int_iter, x):
    return 2 * int_iter < x.shape[0]


def body_fn(int_iter, x):
    return int_iter + 1, x + int_iter


while_loop(cond_fn, body_fn, (0, torch.randn(3, 4)))

Ограничения:

  • body_fn должна возвращать тензоры или целые числа с теми же метаданными (например, формой и dtype), что и входные данные.
  • body_fn и cond_fn не должны изменять carried_inputs на месте. Перед изменением необходимо создать копию.
  • body_fn и cond_fn не должны изменять переменные Python (например, list/dict), созданные вне body_fn.
  • Выходные данные body_fn и cond_fn не могут ссылаться на те же объекты, что и входные данные. Необходимо создать копию.
  • Во время инференса body_fn и cond_fn могут изменять на месте тензоры, не входящие в carried_inputs, например буферы модуля и захваченные тензоры из внешней области видимости.

© 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/while_loop.html

Spec-Zone.ru

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