Spec-Zone.ru › PyTorch 2.14

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

Тип возвращаемого значения:

tuple[Any, Any]

Ограничения:
  • В 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

Spec-Zone.ru

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