# How does torch.compile work with autograd?

**URL:** <https://dev-discuss.pytorch.org/t/how-does-torch-compile-work-with-autograd/1621>\
**Category:** Uncategorized\
**Created:** [October 30, 2023, 1:43pm UTC](https://dev-discuss.pytorch.org/t/how-does-torch-compile-work-with-autograd/1621 "2023-10-30T13:43:13Z")\
**Posts on this page:** 1\
**Showing post:** 4

<div class="post-metadata">

**Author:** ![Chillee](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/chillee/32/10_2.png) [@Chillee](https://dev-discuss.pytorch.org/u/Chillee)\
**Post date:** [October 30, 2023, 8:18pm UTC](https://dev-discuss.pytorch.org/t/how-does-torch-compile-work-with-autograd/1621/4 "2023-10-30T20:18:22Z")

</div>

Haha, AOTAutograd used to be a lot simpler. And conceptually, it is not so complicated.

Basically, the pseudocode for how it works is:

Say we have our forwards

```python
def f(*inputs):
    return outputs

```

And say that we can call `trace(f)(*inputs)` to get our forwards graph.

In order to get the backwards pass, we trace something like

```python
def joint_fw_bw(fw_inputs, grad_outs):
     fw_out = f(*fw_inputs)
     grad_inps = torch.autograd.grad(fw_out, leaves=fw_inputs, gradOuts = grad_outs)
     return fw_out, grad_inps

```

Then, we simply partition this graph into two to give us the forwards pass and the backwards pass.

---

_[View the full topic](https://dev-discuss.pytorch.org/t/how-does-torch-compile-work-with-autograd/1621)._
