# How to trace torch.autograd.backward or torch.autograd.grad?

**URL:** <https://dev-discuss.pytorch.org/t/how-to-trace-torch-autograd-backward-or-torch-autograd-grad/1684>\
**Category:** FX\
**Created:** [November 27, 2023, 7:08am UTC](https://dev-discuss.pytorch.org/t/how-to-trace-torch-autograd-backward-or-torch-autograd-grad/1684 "2023-11-27T07:08:12Z")\
**Posts on this page:** 4\
**Page:** 1

<div class="post-metadata">

**Author:** ![botbw](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/botbw/32/1625_2.png) [@botbw](https://dev-discuss.pytorch.org/u/botbw)\
**Post date:** [November 27, 2023, 7:08am UTC](https://dev-discuss.pytorch.org/t/how-to-trace-torch-autograd-backward-or-torch-autograd-grad/1684/1 "2023-11-27T07:08:12Z")

</div>

Hi there, I would like to trace the backward graph, which has multiple outputs (`stage_output`) and inputs (`input_values`), and some of outputs require grad (specified by `outputs_with_grads_idxs`, which looks like below and contains `backward()` operation.

```python
def stage_backward(stage_output, output_grads, input_values, outputs_with_grads_idxs: List[int]):
    # some preprocessing code
    torch.autograd.backward(
        stage_output_tensors, # outputs that need backward, i.e. stage_output_tensors = stage_output[outputs_with_grad_idxs]
        grad_tensors=output_grad_tensors
    )

```

(check full code [here](https://github.com/pytorch/PiPPy/blob/15dfcd8ea1f445a30627693334ac0b9be160d01b/pippy/backward.py#L9)).

I tried to trace into functions like

```python
def stateless_backward(params, buffers, activations, kwargs_for_stage_backward):
        func_out= stage_backward(**kwargs_for_stage_backward)
        grads = {k: v.grad for k, v in params.items()}
        return func_out, grads

gm = make_fx(stateless_backward, 'fake')(*args)

```

where params and buffers are `FakeTensor`s saved from forward tracing (something like `make_fx(stateless_forward, 'fake')`), and `activations` are `FakeTensor`s collected by iterating `_saved_xxx` of `grad_fn` for all `stage_output`.

The problem is that the resulting graph always contains `_tensor_constant`, and I suspect that I missed some tensors in my `stateless_backward` function arguments so they got cloned and saved in traced code.

I’m a beginner with the stateless graph and `FakeTensor` so I’m even not sure if it is a not-even-wrong intention. Please share if you have any idea :), any suggestion will be appreciated!

---

<div class="post-metadata">

**Author:** ![shuokay](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/shuokay/32/594_2.png) [@shuokay](https://dev-discuss.pytorch.org/u/shuokay)\
**Post date:** [December 5, 2023, 8:55am UTC](https://dev-discuss.pytorch.org/t/how-to-trace-torch-autograd-backward-or-torch-autograd-grad/1684/2 "2023-12-05T08:55:14Z")

</div>

`aot_export_module` should provide you some help, otherwise you may need to use ` __torch_dispatch__ ` to achieve your requirements.

---

<div class="post-metadata">

**Author:** ![shuokay](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/shuokay/32/594_2.png) [@shuokay](https://dev-discuss.pytorch.org/u/shuokay)\
**Post date:** [December 5, 2023, 8:58am UTC](https://dev-discuss.pytorch.org/t/how-to-trace-torch-autograd-backward-or-torch-autograd-grad/1684/3 "2023-12-05T08:58:02Z")

</div>

These discussions might provide some help: [How does torch.compile work with autograd? - #4 by Chillee](https://dev-discuss.pytorch.org/t/how-does-torch-compile-work-with-autograd/1621/4)

---

<div class="post-metadata">

**Author:** ![botbw](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/botbw/32/1625_2.png) [@botbw](https://dev-discuss.pytorch.org/u/botbw)\
**Post date:** [December 5, 2023, 9:40am UTC](https://dev-discuss.pytorch.org/t/how-to-trace-torch-autograd-backward-or-torch-autograd-grad/1684/4 "2023-12-05T09:40:53Z")

</div>

Thanks shuokay! The `default_partition` func is just what I need.

Btw do you have any idea if I only want to trace the backward graph (instead of tracing a joint one and partitioning it)?
