# Why is silu\_backward decomposed in fx graph trace?

**URL:** <https://dev-discuss.pytorch.org/t/why-is-silu-backward-decomposed-in-fx-graph-trace/1319>\
**Category:** FX\
**Created:** [June 12, 2023, 12:52pm UTC](https://dev-discuss.pytorch.org/t/why-is-silu-backward-decomposed-in-fx-graph-trace/1319 "2023-06-12T12:52:06Z")\
**Posts on this page:** 5\
**Page:** 1

<div class="post-metadata">

**Author:** ![liufengwei0103](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/liufengwei0103/32/656_2.png) [@liufengwei0103](https://dev-discuss.pytorch.org/u/liufengwei0103)\
**Post date:** [June 12, 2023, 12:52pm UTC](https://dev-discuss.pytorch.org/t/why-is-silu-backward-decomposed-in-fx-graph-trace/1319/1 "2023-06-12T12:52:06Z")

</div>

When I tried to trace silu and silu\_backward by below code

```auto
import torch
import torch.nn as nn
from torch.fx.experimental.proxy_tensor import ProxyTorchDispatchMode, PythonKeyTracer
from torch.fx._symbolic_trace import Tracer
from torch.utils.weak import WeakTensorKeyDictionary
import weakref

device = torch.device("cpu")

def inner_func():
  a = torch.tensor([2.0, 3.0], requires_grad=True, device=device)
  b = torch._C._nn.silu(a)
  c = torch.relu(b)
  c.backward(torch.tensor([1.0, 1.0], device=device))
python_fx_tracer = PythonKeyTracer()
with ProxyTorchDispatchMode(python_fx_tracer, 'real'):
  graph = python_fx_tracer.trace(inner_func)
print(graph)

```

I got a graph as below:

 ![Screenshot 2023-06-13 114048](https://canada1.discourse-cdn.com/flex036/uploads/pytorch1/original/2X/0/0315a464b6b6c3571db561573e0af65b1f7ef5ba.png)

I find there are two silu\_backward implementation in aten/src/ATen/native/Activation.cpp. **silu\_backward was decomposed and dispatched to the implementation as this link show.**

> <https://github.com/pytorch/pytorch/blob/37359c36fdb413df3b02996eb0ea2433c147db34/aten/src/ATen/native/Activation.cpp#LL547C8-L547C26>

Questions:

1. Why is silu\_backward dispatched to a implementation combined of other op in ProxyTorchDispatchMode? that is to say silu\_backward is decomposed in fx graph.
2. **If I want to get a single silu\_backward node without decomposition in fx graph, what should I do ?**

---

<div class="post-metadata">

**Author:** ![IvanYashchuk](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/ivanyashchuk/32/390_2.png) [@IvanYashchuk](https://dev-discuss.pytorch.org/u/IvanYashchuk)\
**Post date:** [June 16, 2023, 4:36pm UTC](https://dev-discuss.pytorch.org/t/why-is-silu-backward-decomposed-in-fx-graph-trace/1319/2 "2023-06-16T16:36:22Z")

</div>

The reason `silu_backward` gets decomposed unconditionally is because it has a “CompositeImplicitAutograd” dispatch path and by design “math\_silu\_backward” implementation is picked up by the AOTAutograd component of torch.compile.  
It would be great if there was a way to specify that certain operations of type “CompositeImplicitAutograd” shouldn’t be decomposed and traced through, I’m not aware of such a mechanism implemented.

There’s a related issue that unfortunately got no response [Nonoptimal trace of silu\_backward with AOT Autograd · Issue #86612 · pytorch/pytorch · GitHub](https://github.com/pytorch/pytorch/issues/86612)

---

<div class="post-metadata">

**Author:** ![liufengwei0103](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/liufengwei0103/32/656_2.png) [@liufengwei0103](https://dev-discuss.pytorch.org/u/liufengwei0103)\
**Post date:** [June 17, 2023, 3:30am UTC](https://dev-discuss.pytorch.org/t/why-is-silu-backward-decomposed-in-fx-graph-trace/1319/3 "2023-06-17T03:30:21Z")

</div>

I know the AOTAutograd component of torch.compile picks up ‘math\_silu\_backward’ implementation on purpose and ‘silu\_backward’ walks through different dispatch path in many dispatch mode contexts (for example, fake mode) compared to not being in such contexts.  
I try to understand why silu\_backward needs to walk through different dispatch path in these special context, but can’t figure out. Can you help to explain what the purpose is?  
Thanks.

---

<div class="post-metadata">

**Author:** ![IvanYashchuk](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/ivanyashchuk/32/390_2.png) [@IvanYashchuk](https://dev-discuss.pytorch.org/u/IvanYashchuk)\
**Post date:** [June 19, 2023, 1:29pm UTC](https://dev-discuss.pytorch.org/t/why-is-silu-backward-decomposed-in-fx-graph-trace/1319/4 "2023-06-19T13:29:29Z")

</div>

@Chillee might be able to answer this.

---

<div class="post-metadata">

**Author:** ![bdhirsh](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/bdhirsh/32/575_2.png) [@bdhirsh](https://dev-discuss.pytorch.org/u/bdhirsh)\
**Post date:** [June 21, 2023, 1:18am UTC](https://dev-discuss.pytorch.org/t/why-is-silu-backward-decomposed-in-fx-graph-trace/1319/5 "2023-06-21T01:18:50Z")

</div>

It looks like part of the reason is because `make_fx()` will unconditionally run `CompositeImplicitAutograd` decompositions, if there are any.

code that calls `decompose`: [pytorch/torch/fx/experimental/proxy\_tensor.py at ee83c646bb30d5f11b64013b54174768b733214b · pytorch/pytorch · GitHub](https://github.com/pytorch/pytorch/blob/ee83c646bb30d5f11b64013b54174768b733214b/torch/fx/experimental/proxy_tensor.py#L284)

decompose impl: [pytorch/torch/\_ops.py at ee83c646bb30d5f11b64013b54174768b733214b · pytorch/pytorch · GitHub](https://github.com/pytorch/pytorch/blob/ee83c646bb30d5f11b64013b54174768b733214b/torch/_ops.py#L461)

this isn’t exactly the nicest solution, but you can see that we rely on whether or not we return `NotImplemented` to determine if `decompose()` resulted in an actual decomposition. One way to avoid the decomposition would be to use the python dispatcher to register your own implementation of `silu_backward()`, that just returns `NotImplemented`:

```auto
  @torch.ops.aten.silu_backward.default.py_impl(torch._C.DispatchKey.CompositeImplicitAutograd)
  def my_silu_backward(grad_out, self):
      # use with some caution: this is only really valid to run in the context of proxy tensor tracing
      return NotImplemented

```
