Spec-Zone.ru › PyTorch 2.14

Управление потоком — Switch

Создано: 24 июня 2026 г. | Последнее обновление: 24 июня 2026 г.

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

def switch(
    index: Union[int, torch.Tensor],
    branches: Tuple[Callable, ...],
    operands: Tuple[torch.Tensor]
):
    return branches[index](*operands)

Его уникальная возможность заключается в том, что он позволяет выражать управление потоком с несколькими ветвями, зависящее от данных: он преобразуется в оператор switch (torch.ops.higher_order.switch), который сохраняет индекс, все функции ветвей и операнды. Это обеспечивает эффективную компиляцию и развёртывание моделей с N-ветвлением, основанным на значении или форме входных данных либо промежуточных результатов.

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

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

Примеры

Ниже приведён пример, в котором switch используется для выбора одной из нескольких операций на основе входного индекса:

import torch
from torch._higher_order_ops.switch import switch

def branch0(x: torch.Tensor):
    return x.cos()

def branch1(x: torch.Tensor):
    return x.sin()

def branch2(x: torch.Tensor):
    return x.tan()

class BasicSwitch(torch.nn.Module):
    """
    A basic usage of switch with multiple branches.
    """

    def __init__(self):
        super().__init__()

    def forward(self, index: torch.Tensor, x: torch.Tensor) -> torch.Tensor:
        return switch(index, [branch0, branch1, branch2], (x,))

switch_mod = BasicSwitch()

Мы можем выполнить модель в eager-режиме и ожидать, что результаты будут различаться в зависимости от индекса:

x = torch.randn(3)
idx0 = torch.tensor(0)
idx1 = torch.tensor(1)
idx2 = torch.tensor(2)
print(switch_mod(idx0, x), branch0(x))
print(switch_mod(idx1, x), branch1(x))
print(switch_mod(idx2, x), branch2(x))

Мы можем экспортировать модель для дальнейших преобразований и развёртывания:

x = torch.randn(4, 3)
idx = torch.tensor(1)
ep = torch.export.export(
    BasicSwitch(),
    (idx, x),
    dynamic_shapes={"index": None, "x": {0: torch.export.Dim.DYNAMIC}}
)
print(ep)

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

Вот ещё один пример, демонстрирующий switch с индексом, зависящим от данных:

def branch0(x: torch.Tensor):
    return x * 2

def branch1(x: torch.Tensor):
    return x + 10

def branch2(x: torch.Tensor):
    return x ** 2

class DataDependentSwitch(torch.nn.Module):
    """
    A usage of switch with data-dependent index.
    """
    def __init__(self):
        super().__init__()

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # Select branch based on the sign of the sum
        index = torch.clamp((x.sum() > 0).long() + (x.sum() > 5).long(), 0, 2)
        return switch(index, [branch0, branch1, branch2], (x,))

x = torch.randn(4, 3)
ep = torch.export.export(
    DataDependentSwitch(),
    (x,),
    dynamic_shapes={"x": {0: torch.export.Dim.DYNAMIC}}
)
print(ep)

Инварианты torch.ops.higher_order.switch

Для torch.ops.higher_order.switch определено несколько полезных инвариантов:

  • Для индекса:

    • Если индекс является константой (например, целым числом Python), оператор может специализироваться на одной ветви
    • Если индекс — тензор, он должен содержать ровно один элемент
    • Индексы за пределами допустимого диапазона ограничиваются диапазоном [0, len(branches)-1]
  • Для ветвей:

    • Все ветви должны иметь одинаковые сигнатуры входных и выходных данных
    • Сигнатура входных и выходных данных будет представлять собой уплощённый кортеж
    • Они должны быть torch.fx.GraphModule
    • Замыкания в исходных функциях становятся явными входными данными. Замыкания не допускаются.
    • Изменение входных данных или глобальных переменных не допускается
    • Выходные данные ветвей должны быть тензорами или, возможно, вложенными кортежами, списками или словарями тензоров. Листья, не являющиеся тензорами, должны быть int или None. Различающиеся между ветвями значения int объединяются в SymInt для динамических форм; None должны занимать одинаковые позиции во всех ветвях.
  • Для операндов:

    • Это будет плоский кортеж тензоров
  • Вложенные вызовы torch.switch в пользовательской программе становятся вложенными модулями графа

Справочник API

torch._higher_order_ops.switch.switch(index, branches, operands=()) [исходный код]

Выбирает и запускает одну из N функций ветвей по индексу.

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

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

Эквивалентно: branches[index](*operands) с индексом из [0, len(branches)).

Параметры:
  • index (Union[int, torch.Tensor]) – Целое число или целочисленный тензор с одним элементом, указывающий, какую ветвь запустить. Значения за пределами допустимого диапазона ограничиваются диапазоном [0, len(branches)).
  • branches (Union[tuple[Callable, ...], list[Callable]]) – Непустая последовательность вызываемых объектов. Каждый из них должен принимать операнды и возвращать выходные данные одинаковой структуры.
  • operands (Tuple of possibly nested dict/list/tuple of torch.Tensor) – Входные данные для функций ветвей. По умолчанию — ().
Тип возвращаемого значения:

Any

Ограничения:
  • Каждая ветвь должна иметь ту же сигнатуру, что и операнды, и возвращать выходные данные той же структуры (форма, dtype и т. д.). В выходных данных ветвей также допускаются константные листья int и None; они объединяются между ветвями (если листья int различаются между ветвями, вводится неподкреплённый SymInt).
  • В ветвях не допускаются операции изменения на месте входных данных или глобальных переменных.

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

Spec-Zone.ru

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