NVIDIA / NVIDIA/Megatron-LM

RFC: Async overlap of CPU and GPU compute during dynamic inference step

Open
#2,019 2 comments 0 reactions 3 assignees Claimed by @kanz-nv View on GitHub
enhancement
Dominant language
Python
Stars
17.9k
Forks
4.5k
Avg merge
4d 3h
Merged PRs (30d)
272

Description

This RFC the following plan for how to best optimize the dynamic inference step function.

There are several interconnected issues at play:

- Dynamic sampling code is currently very unoptimized. There is a PR draft that reimplements it.
- `async_generate_output_tokens_dynamic_batch` mixes CPU and GPU operations indiscriminately.
- `async_generate_output_tokens_dynamic_batch` may be declared async, but it has no good way of yielding the event loop. A lot of CPU time is wasted waiting for the GPU, and can be reclaimed.

The ideal solution appears to be:
- Fix dynamic sampling code.
- Clearly separate CPU and GPU operations.
- Provide a place to yield the event loop.

The PR series suggested by this RFC are:
1. Break `async_generate_output_tokens_dynamic_batch` apart into multiple sub-methods, which are clearly labeled as "CPU compute" vs "GPU compute".
- Achieved by #1992.
2. Implement barebones unoptimized dynamic sampling code.
- Achieved by #1927.
3. Tensorize the dynamic sampling bookkeeping.
- Achieved by #2105.
4. Reorganize the step function to allow for async step
- Achieved by #2192. No new functionality, just moving code around.
5. Reorder the sub-methods from point 1) so that CPU/GPU compute forms separate continuous blocks of code, and yield the event loop after the CPU compute via torch polling.
- Due to all the prep work, this will be a tiny, extremely readable, PR.
- In-progress at #2193.
6. Optimize dynamic sampling code via graphed FlashInfer sampling.
- In-progress at #2456.
7. Refactor dynamic logprobs computation to follow the same style as the new sampling code.
- In-progress at #2457.
~8. Reconcile with main's implementation of `top_n_logprobs`.~
9. Wait for a torch update, or brainstorm a way to yield the event loop without polling in the current version of pytorch.
- Maybe by sampling on a single rank, instead of the current sampling on every rank?
- Will discuss further in comments.

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.