Поток управления — ассоциативное сканирование
Создано: 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_fnpointwiseилиgeneric; по умолчаниюpointwise. Еслиcombine_mode=pointwise,combine_fnдолжна быть чистой и может содержать только поэлементные операции; при использованииtorch.compilexsдолжна выполняться на бэкенде с поддержкой генерации кода для сканирования (CUDA или XPU), иначе используется общая резервная реализация. Во всех остальных случаях следует использоватьcombine_mode=generic. Примечание:combine_mode=pointwiseэффективнее, чемcombine_mode=generic.
-
combine_fn (Callable) – Вызываемый объект с двумя аргументами и типом
- Возвращает:
-
pytree той же структуры и формы, что и
xs. Если размер измерения сканирования равен 0, выходные данные без изменений соответствуют (пустым) входным данным. Градиент относительноxsтакже пуст (размер 0 вдольdim), поскольку нет элементов, по которым можно вычислить производную. - Тип возвращаемого значения:
Пример:
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