Skip to content

new collective development for JAX/XLA - #657

Open
i-chaochen wants to merge 1 commit into
mainfrom
pavel/xla-mori-kernels-backup
Open

i-chaochen wants to merge 1 commit into
mainfrom
pavel/xla-mori-kernels-backup

Conversation

@i-chaochen

@i-chaochen i-chaochen commented Sep 10, 2026 •

Copy link
Copy Markdown
Contributor

the initial work: https://github.com/ROCm/mori/tree/pemeliya/collectives-develop

Summary

Add a header-only, CCO-backed collective implementation intended for
XLA integration. The new path uses CCO symmetric windows and SDMA
copy-plus-completion packets for intra-node GPU communication without
depending on MORI's SHMEM collective path.

This PR provides the MORI-side facade and kernels; it does not yet add
JAX/XLA registration or Bazel integration.

Changes

  • Add CollectivesFacade for symmetric allocation, lifecycle management,
    kernel dispatch, and runtime tuning.
  • Add SDMA push implementations for:
    • reduce-scatter
    • all-reduce
    • all-gather
    • all-to-all
    • collective-permute
  • Add a direct-XGMI pull implementation for reduce-scatter.
  • Add sliced/pipelined reduce-scatter and all-reduce execution.
  • Add SDMA ADD32 completion packets and supporting runtime plumbing.
  • Add a unified multi-GPU benchmark with device-side correctness checks
    and runtime selection of collective, size, algorithm, and slice count.

Current scope

  • Intra-node LSA/XGMI communication only.
  • Real kernels are compiled for gfx942, gfx950, and gfx1250.
  • Reduction dispatch defaults to f32 with sum.
  • Collective-permute currently supports one destination per rank.
  • Back-to-back launches require stream completion and a host-side
    ccoBarrierAll.
  • Send, receive, device barrier, quiet, and fence APIs are placeholders.

Validation

The benchmark validates each rank's output after timing.

Before submission, record:

  • GPU and ROCm versions:
  • PE counts and message sizes:
  • Collectives and modes tested:
  • Build command:
  • Benchmark commands and results:

Performance

Any RCCL comparison should be rerun on the same hardware and software
configuration using the same bandwidth definition. No performance claim
is made from the preliminary checked-in timing log.

@i-chaochen
i-chaochen requested a review from jhchouuu September 10, 2026 13:41
Introduce a header-only CollectivesFacade backed by a symmetric CCO
heap and SDMA device communication.

Add single-node implementations of reduce-scatter, all-reduce,
all-gather, all-to-all, and collective-permute for gfx942, gfx950, and
gfx1250. Reduce-scatter supports runtime-selectable SDMA push and
direct-load pull modes.

Add a runtime-configurable multi-GPU benchmark with correctness checks,
along with the SDMA ADD32 completion support required by the new path.

The default reduction dispatch instantiates f32/sum only.

Co-authored-by: Pavel Emeliyanenko <pavel.emeliyanenko@amd.com>
@i-chaochen
i-chaochen force-pushed the pavel/xla-mori-kernels-backup branch from 03c43bc to 1f3c1ac Compare September 10, 2026 14:08

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants