Repository navigation
feat(ep): add opt-in gfx1201 EP2/EP4 support - #712
dongdongzhao1121 wants to merge 1 commit into
Conversation
|
nice work, do you try to use mori-cco as the backend rather than shmem? @dongdongzhao1121 |
Thanks for the suggestion! I haven’t tried the CCO backend yet. The current implementation uses HIP IPC mappings for direct PCIe P2P access. If CCO also supports HIP IPC mappings and P2P access on RDNA4, using it as the backend would make sense. @kawhil-amd |
yeah, the mori-cco also support the ipc handle |
Using the CCO backend is also expected to have little impact on performance? |
Sorry for my mistake, CCO does not support IPC handles, but it does support VMM memory allocation. In principle, there should be no performance difference between the two, while CCO with VMM makes address resolution more convenient. |
yeah, cco would be more convenient. |
This PR adds opt-in single-node EP2/EP4 support for gfx1201 (RDNA4) through MORI's existing IntraNode Dispatch/Combine API, including BF16/FP16 and top-k 1–64.
How the gfx1201 kernels work
MORI_RDNA4_EP=1selects specialized HIP kernels behind the existingEpDispatchCombineOpAPI. They communicate directly between GPUs over PCIe using IPC-mapped peer buffers managed by MORI SHMEM.The main implementation is in
src/ops/dispatch_combine/intranode_ep2.hppandintranode_ep4.hpp, with shared routing helpers inintranode_rdna4.hppand launch selection inpython/mori/ops/_rdna4_ep.py.docs/rdna4_ep.mddocuments the supported configurations and protocol details.top-k=8 performance
MORI time is the combined latency of Dispatch + Combine. RCCL time is the combined latency of AllGather + ReduceScatter. Both are measured with GPU graph replay; all table entries are pair p50 latency in microseconds (µs).
We use RCCL AllGather + ReduceScatter as the baseline because the vLLM RDNA4 EP path used in our deployment implements dispatch/combine through the
allgather_reducescatterbackend. In vLLM'sAgRsAll2AllManager, all-gather gathers the inputs for dispatch, and reduce-scatter combines the rank contributions. This benchmark therefore compares MORI's specialized EP transport with the collective operations used by that vLLM EP path.The MORI measurements include routing and FP32 weight communication. The RCCL measurements time dense hidden-payload AllGather + ReduceScatter; vLLM's additional metadata communication, expert computation and framework overhead are outside the benchmark. These are communication measurements, not end-to-end vLLM throughput measurements.
Configuration: BF16, top-k=8, 256 experts, FP32 weights enabled, and capacity equal to tokens per rank. The reported matrix contains every combination of:
This gives 60 configurations per EP, 120 configurations in total. The summary, tables and heatmap all use this same matrix.
Latency reduction = 100 × (1 − MORI_p50 / RCCL_p50);speedup = RCCL_p50 / MORI_p50. Positive reduction means MORI is faster. Summary statistics weight each configuration equally.All 60 hidden-dimension × token-count combinations are listed below, with EP2 and EP4 results side by side. The MORI columns measure Dispatch + Combine; the RCCL columns measure AllGather + ReduceScatter.
Hidden dimension = 2048
Hidden dimension = 3072
Hidden dimension = 4096
Hidden dimension = 5120
Hidden dimension = 7168
Hidden dimension = 8192
Within this matrix, EP2 is slower in three configurations: H=2048/T=32 (−0.10%), H=7168/T=8 (−5.28%), and H=8192/T=8 (−6.76%). EP4 is faster in all 60 configurations. These observations apply to the listed shapes and routing distribution.
Measurement setup
2.12.0+rocm10.0.0, reported HIP runtime7.15.26333, RCCL2.30.4.MORI_RDNA4_EP=1, SHMEM static heap, external input buffers, no quantization.max round p50 / min round p50 > 1.20); this is a stability check, not a confidence interval.Reproduction: build gfx1201 support, then run the tests
1. Prepare the ROCm/PyTorch development environment
Use an Ubuntu ROCm development environment with a gfx1201-capable PyTorch build and HIP compiler. The versions used for the measurements are listed above. Run all subsequent commands in the same environment, using the same
pythoninterpreter. If using a container, expose the GPUs through/dev/kfdand/dev/dri; the host must have a working AMD GPU driver.Install the native build prerequisites inside that environment. These commands assume Ubuntu and
sudo; omitsudowhen already running as root in a container.Keep the environment's existing ROCm SDK path when it is already configured. The measured environment uses
/opt/python/lib/python3.14/site-packages/_rocm_sdk_devel; a conventional system installation uses/opt/rocm.Identify four idle gfx1201 GPUs by their HIP indices, which can differ from
rocm-smiindices. The commands below use HIP4,5for EP2 and4,5,6,7for EP4, matching the measured machine. Replace those lists with the gfx1201 indices on your machine.2. Check out the measured PR revision
Fetch this PR and check out the revision used for the measurements:
If that revision is already checked out locally, set
MORI_REPRO_ROOTto its directory and continue with the build.3. Compile and install MORI with gfx1201 support
MORI_GPU_ARCHS=gfx1201selects gfx1201 for the build and JIT compiler.MORI_RDNA4_EP=1, enabled in step 4, selects the specialized transport at runtime; both settings are needed for this reproduction.This builds the native libraries and Python bindings. The disabled options concern UMBP and C++ examples/tests/benchmarks; they do not disable the Python EP benchmark below.
requirements-build.txtincludes Cython for the required CCO extension.MORI_SKIP_PRECOMPILE=1defers device compilation to the explicit JIT step below.4. Enable the RDNA4 transport and compile its device kernels
Both native libraries and JIT sources now come from this PR checkout, with device kernels compiled for gfx1201. The 2 GiB static heap is for control/routing buffers; payload allocations require additional VRAM.
5. Run host and GPU correctness checks
The correctness entry point is
tests/python/ops/verify_rdna4_ep.py. It starts its own worker processes, so launch it withpython. The GPU commands below check top-k=8 with both BF16 and FP16, all six requested hidden dimensions, and capacity boundaries 512/256, including routing-handle interleaving and consecutive Combine calls.Each driver must complete and exit with status zero. To validate the PR's entire supported top-k range, replace
VERIFY_TOPKS=8withVERIFY_TOPKS=all.6. Reproduce all 60 performance configurations per EP
The performance entry point is
tests/python/ops/bench_rdna4_ep.py. Launch it throughtorch.distributed.runwith two or four workers. It checks both backends before timing and records raw samples and aggregate results in JSONL.Run EP2 and EP4 sequentially. Each output contains 60 configurations; use fresh output filenames on subsequent runs because the driver refuses to overwrite existing measurements. Change
--dtypes bf16to--dtypes bf16,fp16to measure both dtypes; the tables in this comment report BF16 only.The original measurements were taken as part of a larger sweep run in shards. The tables select exactly the requested hidden/token combinations, and the commands above reproduce their configuration and timing settings.