Поток управления — Map
Создано: 14 февр. 2026 | Последнее обновление: 14 февр. 2026
torch.map — это оператор структурированного потока управления, который применяет функцию к ведущему измерению входных тензоров. Логически его можно представить как реализованный следующим образом:
def map(
f: Callable[[PyTree, ...], PyTree],
xs: Union[PyTree, torch.Tensor],
*args,
):
out = []
for idx in range(xs.size(0)):
xs_sliced = xs.select(0, idx)
out.append(f(xs_sliced, *args))
return torch.stack(out)
Предупреждение
torch._higher_order_ops.map — экспериментальная функция PyTorch. Возможны ошибки компиляции. Подробнее о классификации функций: https://pytorch.org/blog/pytorch-feature-classification-changes/#prototype
Примеры
Ниже приведён пример использования map для применения функции к пакету данных:
import torch
from torch._higher_order_ops import map
def f(x):
return x.sin() + x.cos()
xs = torch.randn(3, 4, 5) # batch of 3 tensors, each 4x5
# Applies f to each of the 3 slices
result = map(f, xs) # returns tensor of shape [3, 4, 5]
print(result)
Мы можем экспортировать модель с помощью map для последующих преобразований и развёртывания. В этом примере используются динамические формы, чтобы допускать переменный размер пакета:
class MapModule(torch.nn.Module):
def __init__(self):
super().__init__()
def forward(self, xs: torch.Tensor) -> torch.Tensor:
def body_fn(x):
return x.sin() + x.cos()
return map(body_fn, xs)
mod = MapModule()
inp = torch.randn(3, 4)
ep = torch.export.export(mod, (inp,), dynamic_shapes={"xs": {0: torch.export.Dim.DYNAMIC}})
print(ep)
Обратите внимание, что torch.map преобразуется в torch.ops.higher_order.map_impl, а функция тела становится атрибутом подграфа модуля графа верхнего уровня.
Ограничения
- Отображаемые
xsмогут состоять только из тензоров. - Ведущие измерения всех тензоров в
xsдолжны быть согласованными и ненулевыми. - Функция тела не должна изменять входные данные.
Справочник API
-
torch._higher_order_ops.map.map(f, xs, *args)[исходный код] -
Выполняет отображение f по xs. Интуитивно семантику можно представить так:
out = [] for idx in len(xs.size(0)): xs_sliced = xs.select(0, idx) out.append(f(xs_sliced, *args)) torch.stack(out)Предупреждение
torch._higher_order_ops.map— экспериментальная функция PyTorch. В настоящее время она не поддерживает autograd; также возможны ошибки компиляции. Подробнее о классификации функций: https://pytorch.org/blog/pytorch-feature-classification-changes/#prototype- Параметры:
-
- f (Callable) – вызываемый объект, принимающий входное значение x, которое может быть отдельным тензором или вложенным словарём, списком тензоров и некоторыми дополнительными входными данными
- xs (Any | Tensor) – входные данные, к которым применяется отображение. Мы будем перебирать первое измерение каждого x и выполнять f для каждого среза.
- *args (TypeVarTuple) – дополнительные аргументы, передаваемые на каждом шаге f. Их также можно опустить: map сможет автоматически определить зависимость чтения.
- Возвращает:
-
объединённый в стек результат каждого шага f
Пример:
def f(xs): return xs[0] + xs[1] + const1 + const2 xs = [torch.randn(2, 3), torch.randn(2, 3)] const1 = torch.randn(2, 3) const2 = torch.randn(2, 3) # returns a tensor of shape [2, 2, 3] torch._higher_order_ops.map(f, 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/map.html