[JAX] Support Ring Attention (Context Parallelism) (#1059)
* Implement ring attention primative for Jax. Signed-off-by:Michael Goldfarb <mgoldfarb@nvidia.com> Signed-off-by:
Ming Huang <mingh@nvidia.com> --------- Signed-off-by:
Michael Goldfarb <mgoldfarb@nvidia.com> Signed-off-by:
Ming Huang <mingh@nvidia.com> Co-authored-by:
Michael Goldfarb <mgoldfarb@nvidia.com> Co-authored-by:
pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Showing
tests/jax/test_misc.py
0 → 100644
This diff is collapsed.
Please register or sign in to comment