Spec-Zone.ru › PyTorch 2.14

Поток управления — 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

Spec-Zone.ru

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