Управление потоком — 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) – Входные данные для функций ветвей. По умолчанию — ().
-
index (Union[int, torch.Tensor]) – Целое число или целочисленный тензор с одним элементом, указывающий, какую ветвь запустить. Значения за пределами допустимого диапазона ограничиваются диапазоном
- Тип возвращаемого значения:
- Ограничения:
-
- Каждая ветвь должна иметь ту же сигнатуру, что и операнды, и возвращать выходные данные той же структуры (форма, dtype и т. д.). В выходных данных ветвей также допускаются константные листья
intиNone; они объединяются между ветвями (если листьяintразличаются между ветвями, вводится неподкреплённый SymInt). - В ветвях не допускаются операции изменения на месте входных данных или глобальных переменных.
- Каждая ветвь должна иметь ту же сигнатуру, что и операнды, и возвращать выходные данные той же структуры (форма, dtype и т. д.). В выходных данных ветвей также допускаются константные листья
© 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