Skip to content

fx — Symbolic Model Transformation

Overview

torch.fx traces your module into a graph of Nodes and lets you rewrite it — fuse ops, fold constants, quantize, inject hooks, or export. It's the engine under quantization tooling, torch.compile's front-end, and many research projects.

  • symbolic_trace(module)GraphModule with graph, code(), and full IR.
  • Graph is a DAG of call_function / call_module / call_method / get_attr / placeholder nodes.
  • Write transformations by walking nodes and mutating graph.
  • GraphModule re-compiles to a runnable Python module.

💡 fx tracing executes the module's forward with Proxys, so control flow must be representable: data-dependent branches break naive tracing (fx has Tracer guards).


Trace & Inspect

import torch, torch.nn as nn
from torch.fx import symbolic_trace

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(8, 16)
        self.fc2 = nn.Linear(16, 4)
    def forward(self, x):
        return self.fc2(torch.relu(self.fc1(x)))

gm = symbolic_trace(Net())
print(gm.graph)          # IR dump
# placeholder x -> call_module fc1 -> call_function relu -> call_module fc2 -> output
print(gm.code)           # readable python

Graph Structure (the IR)

for node in gm.graph.nodes:
    print(node.op, node.name, node.target, list(node.args))
# placeholder x x ()
# call_module fc1 fc1 (x,)
# call_function relu <built-in method relu> (fc1,)
# call_module fc2 fc2 (relu,)
# output output (fc2,)

Transformations — Replacing Ops

def replace_relu_with_hardtanh(gm):
    for node in gm.graph.nodes:
        if node.op == 'call_function' and node.target is torch.relu:
            node.target = lambda z: torch.clamp(z, 0, 1)   # new callable
    gm.recompile()
    return gm

gm2 = replace_relu_with_hardtanh(symbolic_trace(Net()))
print("code now uses torch.clamp:", "clamp" in gm2.code)

Fuse consecutive linears (demo pass)

def fuse_consecutive_linear(gm):
    linear_nodes = [n for n in gm.graph.nodes
                    if n.op == 'call_module' and 'fc' in n.target]
    # real passes inspect weight shapes; this is the shape of a pass:
    print(f"found {len(linear_nodes)} linear nodes to consider")
    return gm

fuse_consecutive_linear(symbolic_trace(Net()))

Writing a Pass That Inserts Nodes

from torch.fx import GraphModule
from torch.fx.graph import Graph

def add_relu_after_first_linear(gm):
    g = gm.graph
    with g.inserting_after(list(g.nodes)[1]):   # after placeholders/fc1? insert point
        pass
    return gm

✅ Real quantization passes (e.g., torch.ao.quantization.quantize_fx) are fx passes: they insert quantize/dequantize nodes around linear/conv.


Limitations & Where fx Falls Short

Limitation Workaround
Control flow (if data < 0) torch.export with guards / keep in Python
Python builtins opaque use torch.fx Tracer overrides, or torch.fx.wrap
In-place & aliasing avoid in traced code; use functional forms
Non-tensor args must be traced constants or torch.fx.wrap them

Wrap a python helper

import torch.fx as fx

@fx.wrap
def my_helper(a, b):
    return a.clamp(min=b)

class W(nn.Module):
    def forward(self, x):
        return my_helper(x, 0)
print(list(symbolic_trace(W()).graph.nodes))

fx ↔ Compilation Ecosystem

  • torch.compile(model) uses dynamo (not fx) today, but fx passes remain for research & tooling.
  • torch.ao.quantization.quantize_fx = official fx-based quantization.
  • torch.fx is also the front-end of torch.unittest-style graph tests.
model = torch.ao.quantization.quantize_fx.prepare_qat_fx(
    symbolic_trace(Net()), {'': torch.ao.quantization.default_qat_qconfig}, None)

Key Takeaways

  • fx = trace → IR (nodes) → rewrite → recompile (GraphModule).
  • Passes mutate node targets/args and call recompile().
  • Ideal for fusion, folding, and quantization, not for data-dependent control flow.
  • torch.compile/dynamo supersedes fx for everyday acceleration, but fx remains the tooling layer.