torch.fx.subgraph_rewriter.replace_pattern
-
torch.fx.subgraph_rewriter.replace_pattern(gm, pattern, replacement)[источник] -
Находит все возможные непересекающиеся наборы операторов и их зависимостей по данным (
pattern) в графе GraphModule (gm), а затем заменяет каждый найденный подграф другим подграфом (replacement).- Параметры:
-
- gm (GraphModule) – GraphModule, содержащий граф, над которым выполняется операция
-
pattern (Callable[[...], Any] | GraphModule) – Подграф, который нужно найти в
gmдля замены -
replacement (Callable[[...], Any] | GraphModule) – Подграф, которым нужно заменить
pattern
- Возвращает:
-
Список объектов
Match, представляющих места в исходном графе, которым соответствуетpattern. Если совпадений нет, список пуст.Matchопределяется следующим образом:class Match(NamedTuple): # Node from which the match was found anchor: Node # Maps nodes in the pattern subgraph to nodes in the larger graph nodes_map: Dict[Node, Node] - Тип возвращаемого значения:
-
List[Match]
Примеры:
import torch from torch.fx import symbolic_trace, subgraph_rewriter class M(torch.nn.Module): def __init__(self) -> None: super().__init__() def forward(self, x, w1, w2): m1 = torch.cat([w1, w2]).sum() m2 = torch.cat([w1, w2]).sum() return x + torch.max(m1) + torch.max(m2) def pattern(w1, w2): return torch.cat([w1, w2]) def replacement(w1, w2): return torch.stack([w1, w2]) traced_module = symbolic_trace(M()) subgraph_rewriter.replace_pattern(traced_module, pattern, replacement)Приведённый выше код сначала найдёт
patternв методеforwardобъектаtraced_module. Сопоставление с образцом выполняется на основе отношений использования и определения, а не имён узлов. Например, если бы вpatternу вас былp = torch.cat([a, b]), вы могли бы найтиm = torch.cat([a, b])в исходной функцииforward, несмотря на различие имён переменных (pиm).Оператор
returnвpatternсопоставляется только по значению; он может совпасть или не совпасть с операторомreturnв большем графе. Иными словами, образец не обязан доходить до конца большего графа.После сопоставления образец удаляется из большей функции и заменяется на
replacement. Если в большей функции есть несколько совпадений дляpattern, будут заменены все непересекающиеся совпадения. Если совпадения пересекаются, будет заменено первое найденное совпадение из набора пересекающихся совпадений. («Первое» здесь определяется как первое в топологическом порядке отношений использования и определения узлов. В большинстве случаев первым узлом является параметр, непосредственно следующий заself, а последним — узел, возвращаемый функцией.)Важно отметить, что параметры Callable
patternдолжны использоваться в самом Callable, а параметры Callablereplacementдолжны соответствовать образцу. Первое правило объясняет, почему в приведённом выше блоке кода функцияforwardимеет параметрыx, w1, w2, а функцияpattern— только параметрыw1, w2.patternне используетx, поэтому не следует указыватьxв качестве параметра. В качестве примера второго правила рассмотрим заменуdef pattern(x, y): return torch.neg(x) + torch.relu(y)на
def replacement(x, y): return torch.relu(x)В этом случае
replacementдолжна иметь столько же параметров, сколько иpattern(оба имеютxиy), даже если параметрyне используется вreplacement.После вызова
subgraph_rewriter.replace_patternсгенерированный код Python выглядит следующим образом:def forward(self, x, w1, w2): stack_1 = torch.stack([w1, w2]) sum_1 = stack_1.sum() stack_2 = torch.stack([w1, w2]) sum_2 = stack_2.sum() max_1 = torch.max(sum_1) add_1 = x + max_1 max_2 = torch.max(sum_2) add_2 = add_1 + max_2 return add_2Примечание
Для этого API гарантируется обратная совместимость.
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.fx.subgraph_rewriter.replace_pattern.html