# Understanding CUDAGraph Trees

**URL:** <https://dev-discuss.pytorch.org/t/understanding-cudagraph-trees/1967>\
**Category:** compiler\
**Created:** [March 26, 2024, 10:39am UTC](https://dev-discuss.pytorch.org/t/understanding-cudagraph-trees/1967 "2024-03-26T10:39:20Z")\
**Posts on this page:** 7\
**Page:** 1

<div class="post-metadata">

**Author:** ![Abhishekghosh1998](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/abhishekghosh1998/32/1813_2.png) [@Abhishekghosh1998](https://dev-discuss.pytorch.org/u/Abhishekghosh1998)\
**Post date:** [March 26, 2024, 10:39am UTC](https://dev-discuss.pytorch.org/t/understanding-cudagraph-trees/1967/1 "2024-03-26T10:39:20Z")

</div>

I was listening to the podcast by @ezyang about [CUDAGraph Trees](https://pytorch-dev-podcast.simplecast.com/episodes/cuda-graph-trees). There, he mentions the memory bloat when using CUDA Graphs — which led to the design of CUDA Graph Trees. I was thinking of some clarification regarding the approach.

In the eager CUDAGraph approach (not using `torch.compile`), we could deal with this memory bloat by [Sharing memory across captures](https://pytorch.org/docs/master/notes/cuda.html#sharing-memory-across-captures).

The example shown there is as follows:

```python
g1 = torch.cuda.CUDAGraph()
g2 = torch.cuda.CUDAGraph()

# (create static inputs for g1 and g2, run warmups of their workloads...)

# Captures g1
with torch.cuda.graph(g1):
    static_out_1 = g1_workload(static_in_1)

# Captures g2, hinting that g2 may share a memory pool with g1
with torch.cuda.graph(g2, pool=g1.pool()):
    static_out_2 = g2_workload(static_in_2)

static_in_1.copy_(real_data_1)
static_in_2.copy_(real_data_2)
g1.replay()
g2.replay()

```

After the first CUDAGraph capture in g1, we have a situation like the following:  
 ![image](https://canada1.discourse-cdn.com/flex036/uploads/pytorch1/original/2X/3/319f985ff54f7c0e314b3d659f6eb699db39eb63.png)

- We have a CUDAGraph g1, which uses the `static placeholder1` to get its inputs. (The memory address of this placeholder is baked in the CUDAGraph).
- For its internal operations, it uses the CUDAGraph Memory Pool.
- The `output1` we get after the graph replay is a pointer to a location in the CUDAGraph memory pool.

Next, we capture g2 and make g2 use the same memory pool as g1. And the situation which we have is something like the following:

 ![image](https://canada1.discourse-cdn.com/flex036/uploads/pytorch1/original/2X/c/c53b1e8b7a01f4fb92eed84682e7951d7d0c0ab4.png)

- g2 uses the same CUDAGraph Memory Pool as g1 for its internal operations.
- g2 gets its input from a different static placeholder 2.
- Output of g2 is again a pointer to the CUDAGraph memory pool.

From the [post](https://pytorch.org/docs/master/notes/cuda.html#sharing-memory-across-captures):

> With [`torch.cuda.make_graphed_callables()`](https://pytorch.org/docs/master/generated/torch.cuda.make_graphed_callables.html#torch.cuda.make_graphed_callables), if you want to graph several callables and you know they’ll always run in the same order (and never concurrently) pass them as a tuple in the same order they’ll run in the live workload, and [`make_graphed_callables()`](https://pytorch.org/docs/master/generated/torch.cuda.make_graphed_callables.html#torch.cuda.make_graphed_callables) will capture their graphs using a shared private pool.

This means that if the capture sequence is g1 → g2, we should follow the order of g1 → g2 during replay. (why?)

> [@CUDAGraphs in Pytorch 2.0](https://dev-discuss.pytorch.org/t/cudagraphs-in-pytorch-2-0/1428/1):
>
> Like Graph Callables, CUDA Graph Trees use a single memory pool across all graph captures.

I want to understand a few details about the approach:

1. I guess copying data to the static placeholder is not required in CUDAGraph trees (based on the diagrams and the [lightning talk](https://youtu.be/Lg8F4F_qZxk?si=15WtHyqUT_cPRoTE&t=278) by @eellison). The output address of g1 is baked into the graph g2 as its (static) input.
2. If we have a chain, say g1 → g2, I guess g1 and g2 share the same CUDA memory pool in the same manner as in the case of eager CUDA Graph capture.
3. Now, as in the post by @eellison following g1 → g2 if we have another graph g4 (a chain g1 → g2 → g4), I guess, g4 as well shares the same memory pool as that of g1, g2 right?
4. I am trying to understand why we need to branch out and checkpoint the memory state.
  - Let us say that we are in a situation where we have captured g1 followed by g2, and g2 shares the same memory pool as g1.
  - Now, suppose we reach a situation where we need to execute graph g3 after g1. And now, given the current situation of the memory pool, if we try capturing g3, it would lead to the dependency g1 → g2 → g3, right? Trying to replay in order g1 → g3 might result in an error because we are not following the capture order during replay.

> [@CUDAGraphs in Pytorch 2.0](https://dev-discuss.pytorch.org/t/cudagraphs-in-pytorch-2-0/1428/1):
>
> A CUDAGraphNode specializes the memory allocation patterns and tensor lifetimes from when it was recorded in order to check if it can re-execute. It checks that the same tensors are still alive from its parent node, and that any tensors that died after recording die again on re-execution, and that the path to the root is the same on execution as on recording.

1. Regarding the above quote, can a simple example be provided illustrating its importance?

> [@CUDAGraphs in Pytorch 2.0](https://dev-discuss.pytorch.org/t/cudagraphs-in-pytorch-2-0/1428/1):
>
> The same pattern of memory must be observed between recording and replay: if a tensor output of one graph dies subsequent to another graph during recording, it must also do so during replay.

1. Similarly, for the above quote, a simple example would help to understand it.

> [@CUDAGraphs in Pytorch 2.0](https://dev-discuss.pytorch.org/t/cudagraphs-in-pytorch-2-0/1428/1):
>
> Live memory in the cuda pool forces a dependence between two recordings

1. An example of the “live memory” discussed above and how it causes dependency between two recordings.

> [@CUDAGraphs in Pytorch 2.0](https://dev-discuss.pytorch.org/t/cudagraphs-in-pytorch-2-0/1428/1):
>
> Graph 1 gets replayed, and then we hit Graph 3 which we have not yet recorded. On graph replays the private memory pool is not updated, so y is not reflected in the allocator. Without care we would overwrite it. To support reusing the same memory pool after replaying other graphs, we checkpoint the memory pool back to its state at the end of graph 1. Checkpointing both updates the CUDACaching allocator to reflect the currently live tensors, and adds a deleter function to the live tensors so that when they die, the allocator will mark their memory as free. Now that our live tensors are reflected in the caching allocator, we are safe to run a new graph.

1. If I can understand the idea of “live tensors” (as asked in point 7), then the above quote will get clearer, I guess.

I am fascinated by the concept of `torch.compile`, and I find this section dealing with CUDAGraphs pretty interesting. So, couldn’t help but write down my doubts.

---

<div class="post-metadata">

**Author:** ![eellison](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/eellison/32/102_2.png) [@eellison](https://dev-discuss.pytorch.org/u/eellison)\
**Post date:** [March 28, 2024, 4:57pm UTC](https://dev-discuss.pytorch.org/t/understanding-cudagraph-trees/1967/2 "2024-03-28T16:57:57Z")

</div>

Hi! Seems like you have a great handle on things. Answering a few questions/clarifying a few things.

> 1. 
> - Now, suppose we reach a situation where we need to execute graph g3 after g1. And now, given the current situation of the memory pool, if we try capturing g3, it would lead to the dependency g1 → g2 → g3, right? Trying to replay in order g1 → g3 might result in an error because we are not following the capture order during replay.

We had previously captured g1 → g2, and now we are trying to execute g1 → g3. The reason we need to checkpoint the memory state is during the fast - path execution of cuda graphs we dont do any memory accounting in the cuda caching allocator. All of the allocations appear as deallocated. To capture a new graph and share memory pool, the cuda caching allocator needs to know which tensors are live so that new allocations made in that pool do not overwrite existing live tensors.

5/6 are pretty much the same concept. Run the below with TORCH\_LOGS=“cudagraphs”

```auto
import torch

@torch.compile(mode="reduce-overhead")
def foo(x):
    return x + 1, x + 2

@torch.compile(mode="reduce-overhead")
def fee(y):
    return y * 4

for _ in range(3):
    torch.compiler.cudagraph_mark_step_begin()
    inp = torch.rand([4], device="cuda")
    a, b = foo(inp)
    del a
    # a's memory can be reused here
    fee(b)

torch.compiler.cudagraph_mark_step_begin()
inp = torch.rand([4], device="cuda")
a, b = foo(inp)
print("Should checkpoint now")
# a no longer deleted, still live, cant reclaim meomry
fee(b)

```

7/8. Memory dependency is forced by live tensors. The caching allocator liveness state depends on prior outputs in the current tree.

Here are some examples.

```auto
import torch

@torch.compile(mode="reduce-overhead")
def foo(x):
    return x + 1, x + 2

@torch.compile(mode="reduce-overhead")
def fee(y):
    return y * 4

def get_curr_node(device):
    return torch._inductor.cudagraph_trees.get_container(device.index).tree_manager.current_node

for i in range(3):
    torch.compiler.cudagraph_mark_step_begin()
    inp = torch.rand([4], device="cuda")
    a, b = foo(inp)
    del a, b
    # no mem dependency
    fee(inp)
    assert get_curr_node(inp.device).parent is None

for i in range(3):
    torch.compiler.cudagraph_mark_step_begin()
    inp = torch.rand([6], device="cuda")
    a, b = foo(inp)
    # mem dependency
    fee(inp)
    assert get_curr_node(inp.device).parent is not None

```

---

<div class="post-metadata">

**Author:** ![Abhishekghosh1998](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/abhishekghosh1998/32/1813_2.png) [@Abhishekghosh1998](https://dev-discuss.pytorch.org/u/Abhishekghosh1998)\
**Post date:** [April 2, 2024, 7:10pm UTC](https://dev-discuss.pytorch.org/t/understanding-cudagraph-trees/1967/3 "2024-04-02T19:10:34Z")

</div>

Thanks @eellison for your insightful explanation.

> [@eellison](#):
>
> We had previously captured g1 → g2, and now we are trying to execute g1 → g3. The reason we need to checkpoint the memory state is during the fast - path execution of cuda graphs we dont do any memory accounting in the cuda caching allocator. All of the allocations appear as deallocated. To capture a new graph and share memory pool, the cuda caching allocator needs to know which tensors are live so that new allocations made in that pool do not overwrite existing live tensors.

Just to clarify on this:

 ![image](https://canada1.discourse-cdn.com/flex036/uploads/pytorch1/original/2X/c/c4372269f55c8f67135c8c03dab0daf63420ff81.png)

In the above example figure, we capture g1 → g2. Then we replay g1 → g2. As a result of this we have _live tensors_ `Output1` and `Output2` in the CUDAGraph Memory Pool. (Since you said that the allocations are not accounted for during graph-replay or fast-path, and the allocations appear as deallocated, I assume Output1 and Output2 as just pointers to memory location in the CUDAGraph Memory Pool, which can be assumed logically as a long stretch of memory region)

 ![image](https://canada1.discourse-cdn.com/flex036/uploads/pytorch1/original/2X/d/d47f008bb33b801601691746521bceb5ae3c4b33.png)

Next if we try to use the same CUDAGraph Memory Pool for capture of CUDAGraph g3, then it can just overwrite the memory location (for the allocations of g3) which are pointed to by `Output1` and/or `Output2`, since the CachingAllocator has no accounting information, the live tensors (`Output1` and `Output2`) are not actually allocated. We need to preserve the state of the live tensors - therefore we checkpoint. Right?

I guess, I got why do we need to checkpoint. I would like to get some clarifications, about how the checkpointing is done and how it solves the problem.

> [@CUDAGraphs in Pytorch 2.0](https://dev-discuss.pytorch.org/t/cudagraphs-in-pytorch-2-0/1428/1):
>
> To support reusing the same memory pool after replaying other graphs, we checkpoint the memory pool back to its state at the end of graph 1. Checkpointing both updates the CUDACaching allocator to reflect the currently live tensors, and adds a deleter function to the live tensors so that when they die, the allocator will mark their memory as free.

How checkpointing the memory pool back to the state it was at the end of graph g1, solves the issue of maintaining the live-tensor state? If we checkpoint as above, the state of live tensors shall be lost right? I am not getting this part.

> <https://github.com/pytorch/pytorch/blob/feabb645a7fbbd695d25aa94150e6b0e90fb07c6/c10/cuda/CUDACachingAllocator.cpp#L1584C3-L1587C60>

The comment in the CUDACachingAllocator.cpp seems to mention the same thing, but I am unable to see through it.

---

<div class="post-metadata">

**Author:** ![eellison](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/eellison/32/102_2.png) [@eellison](https://dev-discuss.pytorch.org/u/eellison)\
**Post date:** [April 11, 2024, 12:25am UTC](https://dev-discuss.pytorch.org/t/understanding-cudagraph-trees/1967/4 "2024-04-11T00:25:40Z")

</div>

> How checkpointing the memory pool back to the state it was at the end of graph g1, solves the issue of maintaining the live-tensor state? If we checkpoint as above, the state of live tensors shall be lost right? I am not getting this part.

The [test here](https://github.com/pytorch/pytorch/blob/main/test/test_cuda.py#L3954-L3976) might be helpful.

After `del outputs` if were to run `live_blocks(pool_id)` it would be equal 0. However, after we run checkpointing on the captured Cuda Caching Allocator state, it correctly accounts for 2 live allocated blocks of memory. The `_cuda_cudaCachingAllocator_raw_delete` is a stand-in for what happens in cudagraph trees where the tensors call raw\_delete in their deleter\_fn.

Note that we checkpoint before we record g3, and we apply any deltas in liveness that might have occurred between the end of g1 and start of g3.

---

<div class="post-metadata">

**Author:** ![Abhishekghosh1998](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/abhishekghosh1998/32/1813_2.png) [@Abhishekghosh1998](https://dev-discuss.pytorch.org/u/Abhishekghosh1998)\
**Post date:** [April 11, 2024, 2:00am UTC](https://dev-discuss.pytorch.org/t/understanding-cudagraph-trees/1967/5 "2024-04-11T02:00:20Z")

</div>

Thanks @eellison.

Do you have any CUDAGraph Trees design doc apart from [this](https://docs.google.com/document/d/1ZrxLGWz7T45MSX6gPsL6Ln4t0eZCSfWewtJ_qLd_D0E/)? Like some doc, upon which the source code of CUDAGraph Trees is based.

---

<div class="post-metadata">

**Author:** ![shink](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/shink/32/1970_2.png) [@shink](https://dev-discuss.pytorch.org/u/shink)\
**Post date:** [April 15, 2025, 6:55am UTC](https://dev-discuss.pytorch.org/t/understanding-cudagraph-trees/1967/6 "2025-04-15T06:55:42Z")

</div>

Hey thanks so much for this post. I just have a question here, what the difference between `backend='cudagraphs'` and `mode='reduce-overhead'` in `torch.compile()`? Seems they are both using cudagraph trees and cannot be used at the same time.

---

<div class="post-metadata">

**Author:** ![shink](https://yyz2.discourse-cdn.com/flex036/user_avatar/dev-discuss.pytorch.org/shink/32/1970_2.png) [@shink](https://dev-discuss.pytorch.org/u/shink)\
**Post date:** [April 15, 2025, 12:53pm UTC](https://dev-discuss.pytorch.org/t/understanding-cudagraph-trees/1967/7 "2025-04-15T12:53:08Z")

</div>

Oh seems cudagraph trees capture fx graphs when backend=cudagraphs, and capture triton code when mode=reduce-overhead. Right?
