# How to get the backward graph while using torch.export?

**URL:** <https://dev-discuss.pytorch.org/t/how-to-get-the-backward-graph-while-using-torch-export/2318>\
**Category:** compiler\
**Created:** [July 26, 2024, 12:42am UTC](https://dev-discuss.pytorch.org/t/how-to-get-the-backward-graph-while-using-torch-export/2318 "2024-07-26T00:42:18Z")\
**Posts on this page:** 5\
**Page:** 1

<div class="post-metadata">

**Author:** ![fishmingyu](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/fishmingyu/32/1597_2.png) [@fishmingyu](https://dev-discuss.pytorch.org/u/fishmingyu)\
**Post date:** [July 26, 2024, 12:42am UTC](https://dev-discuss.pytorch.org/t/how-to-get-the-backward-graph-while-using-torch-export/2318/1 "2024-07-26T00:42:18Z")

</div>

I have read the references such as [topic-1621](https://dev-discuss.pytorch.org/t/how-does-torch-compile-work-with-autograd/1621/2) and [aot\_autograd](https://pytorch.org/functorch/stable/notebooks/aot_autograd_optimizations.html). However, I am unable to find a solution that enables exporting the backward graph while using torch.export. Is there any trick to do this? I provide a simple code script below:

```python
import torch
from torch._functorch.aot_autograd import aot_module

class M(torch.nn.Module):
    def __init__ (self):
        super(). __init__ ()
        self.linear = torch.nn.Linear(3, 3)

    def forward(self, x):
        x = self.linear(x)
        return x

inp = torch.randn(2, 3, device="cuda", requires_grad=True)
m = M().to(device="cuda")
ep = torch.export.export(m, (inp,))

# check graph module
print(ep.module().code)

# Run it with torch.compile
compile_module = torch.compile(ep.module(), backend="inductor")
res = compile_module(inp)

res.sum().backward()
# how print the backward graph?

```

Thank you for your support!

---

<div class="post-metadata">

**Author:** ![fishmingyu](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/fishmingyu/32/1597_2.png) [@fishmingyu](https://dev-discuss.pytorch.org/u/fishmingyu)\
**Post date:** [July 26, 2024, 4:11am UTC](https://dev-discuss.pytorch.org/t/how-to-get-the-backward-graph-while-using-torch-export/2318/2 "2024-07-26T04:11:47Z")

</div>

I can see the backward graph in the debug logs after calling `TORCH_COMPILE_DEBUG=1`. However, can I extract this graph directly from APIs?

---

<div class="post-metadata">

**Author:** ![wmhst7](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/wmhst7/32/2013_2.png) [@wmhst7](https://dev-discuss.pytorch.org/u/wmhst7)\
**Post date:** [July 26, 2024, 8:01pm UTC](https://dev-discuss.pytorch.org/t/how-to-get-the-backward-graph-while-using-torch-export/2318/3 "2024-07-26T20:01:16Z")

</div>

This post may be helpful: [How to set wrap function using TorchDynamo graph capture?](https://dev-discuss.pytorch.org/t/how-to-set-wrap-function-using-torchdynamo-graph-capture/2185)

Another way to do this as I know is like:

```auto
from functorch.compile import aot_module

captured_graphs = []

def custom_compiler(m: torch.fx.GraphModule, _):
    captured_graphs.append(m)
    return make_boxed_func(m.forward)

aot_model = aot_module(model, fw_compiler=custom_compiler)
y = aot_model(inputs)
y.sum().backward()

```

---

<div class="post-metadata">

**Author:** ![fishmingyu](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/fishmingyu/32/1597_2.png) [@fishmingyu](https://dev-discuss.pytorch.org/u/fishmingyu)\
**Post date:** [July 27, 2024, 12:13am UTC](https://dev-discuss.pytorch.org/t/how-to-get-the-backward-graph-while-using-torch-export/2318/4 "2024-07-27T00:13:12Z")

</div>

Thank for your reply. I have tried this before but failed with some models. I think aot module is less expressive than using torch.export directly.

---

<div class="post-metadata">

**Author:** ![wmhst7](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/wmhst7/32/2013_2.png) [@wmhst7](https://dev-discuss.pytorch.org/u/wmhst7)\
**Post date:** [July 29, 2024, 10:42pm UTC](https://dev-discuss.pytorch.org/t/how-to-get-the-backward-graph-while-using-torch-export/2318/5 "2024-07-29T22:42:22Z")

</div>

It makes sense. Also wonder how to use export to do backward capture. 😃
