Spec-Zone.ru › PyTorch 2.14

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

Создано: 14 февр. 2026 г. | Последнее обновление: 14 февр. 2026 г.

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

def associative_scan(
    combine_fn: Callable[[pytree.PyTree, pytree.PyTree], pytree.PyTree],
    xs: pytree.PyTree,
    dim: int,
    reverse: bool = False,
) -> pytree.PyTree:
    result = []
    carry = xs.select(dim, 0)
    result.append(carry)
    for i in range(1, xs.size(dim)):
        carry = combine_fn(carry, xs.select(dim, i))
        result.append(carry)
    return torch.stack(result, dim=dim)

Поскольку combine_fn должна быть ассоциативной, вычисления можно распараллелить с помощью алгоритма древовидной редукции, а не выполнять последовательно. Это позволяет эффективно реализовывать на GPU такие операции, как накопительные суммы, произведения и другие ассоциативные накопления.

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

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

Примеры

Ниже приведён пример использования associative_scan для вычисления накопительной суммы:

import torch
from torch._higher_order_ops.associative_scan import associative_scan

def add(x: torch.Tensor, y: torch.Tensor):
    return x + y

xs = torch.arange(1, 5, dtype=torch.float32)  # [1, 2, 3, 4]
cumsum = associative_scan(add, xs, dim=0, combine_mode="generic")
print(cumsum)

Ниже приведён пример вычисления накопительного произведения:

def mul(x: torch.Tensor, y: torch.Tensor):
    return x * y

xs = torch.arange(1, 5, dtype=torch.float32)  # [1, 2, 3, 4]
cumprod = associative_scan(mul, xs, dim=0, combine_mode="generic")
print(cumprod)

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

class AssociativeScanModule(torch.nn.Module):
    def forward(self, xs: torch.Tensor) -> torch.Tensor:
        def combine_fn(x, y):
            return x + y

        return associative_scan(combine_fn, xs, dim=0, combine_mode="pointwise")

mod = AssociativeScanModule()
inp = torch.randn(5, 3, device="cuda")
dim_seq = torch.export.Dim("seq", min=2)
ep = torch.export.export(mod, (inp,), dynamic_shapes={"xs": {0: dim_seq}})
print(ep)
ExportedProgram:
    class GraphModule(torch.nn.Module):
        def forward(self, xs: "f32[s83, 3]"):
            # File: /data/users/angelayi/pytorch2/foo.py:25 in forward, code: return associative_scan(combine_fn, xs, dim=0, combine_mode="pointwise")
            movedim: "f32[s83, 3]" = torch.ops.aten.movedim.int(xs, 0, 0);  xs = None

            # File: <eval_with_key>.3:6 in forward, code: select_copy = torch.select_copy(l_leaves_xs_0_, 0, 0);  select_copy = None
            select_copy: "f32[3]" = torch.ops.aten.select_copy.int(movedim, 0, 0);  select_copy = None

            # File: <eval_with_key>.3:8 in forward, code: associative_scan = torch.ops.higher_order.associative_scan(associative_scan_combine_fn_0, [l_leaves_xs_0_], ());  associative_scan_combine_fn_0 = l_leaves_xs_0_ = None
            associative_scan_combine_graph_0 = self.associative_scan_combine_graph_0
            associative_scan = torch.ops.higher_order.associative_scan(associative_scan_combine_graph_0, [movedim], ());  associative_scan_combine_graph_0 = movedim = None
            getitem: "f32[s83, 3]" = associative_scan[0];  associative_scan = None

            # File: /data/users/angelayi/pytorch2/foo.py:25 in forward, code: return associative_scan(combine_fn, xs, dim=0, combine_mode="pointwise")
            movedim_1: "f32[s83, 3]" = torch.ops.aten.movedim.int(getitem, 0, 0);  getitem = None
            return (movedim_1,)

        class associative_scan_combine_graph_0(torch.nn.Module):
            def forward(self, arg0_1: "f32[3]", arg1_1: "f32[3]"):
                # File: <eval_with_key>.4:5 in forward, code: add = child + child_1;  child = child_1 = None
                add: "f32[3]" = torch.ops.aten.add.Tensor(arg0_1, arg1_1);  arg0_1 = arg1_1 = None
                return [add]

Graph signature:
    # inputs
    xs: USER_INPUT

    # outputs
    movedim_1: USER_OUTPUT

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

Ограничения

  • combine_fn должна быть ассоциативной: combine_fn(combine_fn(a, b), c) == combine_fn(a, combine_fn(b, c)).
  • combine_fn не должна изменять входные данные на месте.
  • combine_fn не должна обращаться к переменным из внешней области видимости (замыкания не поддерживаются).
  • Результат combine_fn не должен иметь общую память ни с одним из входных данных.

Справочник API

torch._higher_order_ops.associative_scan.associative_scan(combine_fn, xs, dim, reverse=False, combine_mode='pointwise') [исходный код]

Выполняет включительное сканирование с ассоциативной функцией объединения.

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

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

Для combine_mode="pointwise" эффективное выполнение требует генерации кода во время выполнения с помощью torch.compile, а генерация кода поддерживается только на бэкендах с поддержкой сканирования (в настоящее время CUDA и XPU). На других устройствах оператор по-прежнему выполняется в eager-режиме с использованием общей резервной реализации.

Параметры:
  • combine_fn (Callable) – Вызываемый объект с двумя аргументами и типом (Tensor, Tensor) -> Tensor либо, если входные данные являются pytree, (pytree, pytree) -> pytree. Эта функция должна быть чистой (на данный момент поддержка lifted-аргументов отсутствует), удовлетворять свойству ассоциативности и не иметь побочных эффектов.
  • xs (torch.Tensor) – Входной тензор или вложенный pytree тензоров.
  • dim (int) – измерение, по которому выполняется сканирование
  • reverse (bool) – Логическое значение, указывающее, следует ли выполнять сканирование в обратном направлении относительно dim; по умолчанию False.
  • combine_mode (str) – Строка, указывающая, является ли combine_fn pointwise или generic; по умолчанию pointwise. Если combine_mode=pointwise, combine_fn должна быть чистой и может содержать только поэлементные операции; при использовании torch.compile xs должна выполняться на бэкенде с поддержкой генерации кода для сканирования (CUDA или XPU), иначе используется общая резервная реализация. Во всех остальных случаях следует использовать combine_mode=generic. Примечание: combine_mode=pointwise эффективнее, чем combine_mode=generic.
Возвращает:

pytree той же структуры и формы, что и xs. Если размер измерения сканирования равен 0, выходные данные без изменений соответствуют (пустым) входным данным. Градиент относительно xs также пуст (размер 0 вдоль dim), поскольку нет элементов, по которым можно вычислить производную.

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

Tensor

Пример:

def add(x: torch.Tensor, y: torch.Tensor):
    return x + y


cumsum = associative_scan(add, x, dim)

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

Spec-Zone.ru

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