Управление потоком — Scan
Создано: 14 февр. 2026 г. | Последнее обновление: 14 февр. 2026 г.
torch.scan — это оператор структурированного управления потоком, выполняющий инклюзивное сканирование с помощью комбинирующей функции. Обычно он используется для кумулятивных операций, таких как cumsum, cumprod, или более общих рекуррентных вычислений. Логически его можно представить в следующем виде:
def scan(
combine_fn: Callable[[PyTree, PyTree], tuple[PyTree, PyTree]],
init: PyTree,
xs: PyTree,
*,
dim: int = 0,
reverse: bool = False,
) -> tuple[PyTree, PyTree]:
carry = init
ys = []
for i in range(xs.size(dim)):
x_slice = xs.select(dim, i)
carry, y = combine_fn(carry, x_slice)
ys.append(y)
return carry, torch.stack(ys)
Предупреждение
torch.scan — экспериментальная функция в PyTorch. Возможны ошибки компиляции. Подробнее о классификации функций см. по ссылке: https://pytorch.org/blog/pytorch-feature-classification-changes/#prototype
Примеры
Ниже приведён пример использования scan для вычисления кумулятивной суммы:
import torch
from torch._higher_order_ops import scan
def add(carry: torch.Tensor, x: torch.Tensor):
next_carry = carry + x
y = next_carry.clone() # clone to avoid output-output aliasing
return next_carry, y
init = torch.zeros(1)
xs = torch.arange(5, dtype=torch.float32)
final_carry, cumsum = scan(add, init=init, xs=xs)
print(final_carry)
print(cumsum)
Мы можем экспортировать модель с scan для последующих преобразований и развёртывания. В этом примере используются динамические формы, позволяющие задавать переменную длину последовательности:
class ScanModule(torch.nn.Module):
def forward(self, xs: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
def combine_fn(carry, x):
next_carry = carry + x
return next_carry, next_carry.clone()
init = torch.zeros_like(xs[0])
return scan(combine_fn, init=init, xs=xs)
mod = ScanModule()
inp = torch.randn(5, 3)
ep = torch.export.export(mod, (inp,), dynamic_shapes={"xs": {0: torch.export.Dim.DYNAMIC}})
print(ep)
Обратите внимание, что комбинирующая функция становится атрибутом подграфа модуля графа верхнего уровня.
Ограничения
-
combine_fnдолжна возвращать тензоры с одинаковыми метаданными (формой, типом данных) дляnext_carryиinit. -
combine_fnне должна изменять входные данные на месте. Перед изменением необходимо создать копию. -
combine_fnне должна изменять переменные Python (например, list/dict), созданные вне функции. -
Выходные данные
combine_fnне должны ссылаться на те же данные, что и какие-либо входные данные. Необходимо создать копию.
Справочник API
-
torch._higher_order_ops.scan.scan(combine_fn, init, xs, *, dim=0, reverse=False, length=None)[исходный код] -
Выполняет инклюзивное сканирование с помощью комбинирующей функции.
Предупреждение
torch.scan— экспериментальная функция в PyTorch. Возможны ошибки компиляции. Подробнее о классификации функций см. по ссылке: https://pytorch.org/blog/pytorch-feature-classification-changes/#prototype- Параметры:
-
-
combine_fn (Callable) – Вызываемый объект с двумя аргументами типа
(Tensor, Tensor) -> (Tensor, Tensor)или, если xs является pytree,(pytree, pytree) -> (pytree, pytree). Первый аргументcombine_fn— это предыдущее или начальное значение переноса сканирования, а второй входной элементcombine_fn— срез входных данных вдоль dim. Первый выходной элементcombine_fn— это следующее значение переноса сканирования, а второй выходной элементcombine_fnпредставляет собой срез выходных данных. Эта функция должна быть чистой: на данный момент не поддерживаются поднятые аргументы, также у неё не должно быть побочных эффектов. -
init (torch.Tensor или pytree с тензорными листьями) – Начальное значение переноса сканирования: тензор или вложенная pytree-структура тензоров. Ожидается, что
initбудет иметь ту же структуру pytree, что и первый выходной элемент (то есть значение переноса)combine_fn. -
xs (torch.Tensor или pytree с тензорными листьями или None) – Входной тензор или вложенная pytree-структура тензоров. Может быть равен
None, если заданlength; в этом случаеcombine_fnполучаетNoneкакxна каждой итерации (режим цикла со счётчиком).
-
combine_fn (Callable) – Вызываемый объект с двумя аргументами типа
- Именованные аргументы:
-
- dim (int) – измерение, вдоль которого выполняется сканирование; по умолчанию 0.
-
reverse (bool) – Логическое значение, указывающее, следует ли выполнять сканирование в обратном направлении относительно
dim; по умолчаниюFalse. -
length (int или None) – Необязательное количество итераций сканирования; по умолчанию
None. Еслиxsсодержит тензорные листья,lengthуказывать необязательно; если оно задано, его значение должно совпадать сxs.shape[dim], и оно используется только для проверки согласованности (ограничение не действует, когдаlengthравноNone). Еслиxsне содержит листьев (Noneили пустая pytree-структура),lengthопределяет число итераций, аcombine_fnполучаетx=Noneна каждой итерации.length=0без тензоров xs поддерживается только в eager-режиме; он не поддерживается вtorch.compile.
- Возвращает:
-
- final_carry (torch.Tensor или pytree с тензорными листьями),
-
итоговое значение переноса операции сканирования с той же структурой pytree, что и у init.
- out (torch.Tensor или pytree с тензорными листьями),
-
каждый тензорный лист представляет собой выходные данные, объединённые вдоль первого измерения; каждый срез соответствует результату одной итерации сканирования. Если размер измерения сканирования равен 0,
final_carryне изменяется и равенinit, а размер каждого выходного листа вдольdimравен 0. Градиентfinal_carryпо отношению кinitравен тождественному преобразованию (а не нулю), поскольку тело функции не вызывается и значение переноса проходит без изменений.
- Тип возвращаемого значения:
- Ограничения:
-
-
- В combine_fn не должно быть алиасинга между входными данными, между входными и выходными данными, а также между выходными данными. Например, возврат представления
-
или того же тензора, что и на входе, не поддерживается. В качестве обходного решения можно создать копию выходных данных, чтобы избежать алиасинга.
-
- combine_fn не должна изменять входные данные. Вскоре мы снимем ограничение на изменение данных для вывода. Пожалуйста, создайте запрос,
-
если для обучения необходима поддержка изменения входных данных.
- Начальное значение переноса combine_fn должно совпадать со значением next_carry по структуре pytree и метаданным тензора.
-
Пример:
def add(x: torch.Tensor, y: torch.Tensor): next_carry = y = x + y # clone the output to avoid output-output aliasing return next_carry, y.clone() i0 = torch.zeros(1) xs = torch.arange(5) # returns torch.tensor([10.]), torch.tensor([[0], [1.], [3.], [6.], [10.]]) last_carry, cumsum = scan(add, init=i0, xs=xs)
© 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/scan.html