Поток управления — цикл 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/#prototypewhile_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