Skip to content

feat(ep): add opt-in gfx1201 EP2/EP4 support - #712

Open
dongdongzhao1121 wants to merge 1 commit into
ROCm:mainfrom
dongdongzhao1121:feat/gfx1201-intranode-ep
Open

dongdongzhao1121 wants to merge 1 commit into
ROCm:mainfrom
dongdongzhao1121:feat/gfx1201-intranode-ep

Conversation

@dongdongzhao1121

@dongdongzhao1121 dongdongzhao1121 commented Sep 28, 2026 •

Copy link
Copy Markdown

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=1 selects specialized HIP kernels behind the existing EpDispatchCombineOp API. They communicate directly between GPUs over PCIe using IPC-mapped peer buffers managed by MORI SHMEM.

  • Dispatch: Map each token's selected experts to destination GPU ranks and reserve one receive row per token per destination rank. Multiple experts on the same rank share one hidden payload. EP2 pushes both the hidden payload and routing metadata to the destination; EP4 pushes the hidden payload while the destination pulls staged expert indices and optional FP32 weights. Dispatch also records the mapping back to the source token for Combine and routing-handle replay.
  • Combine: After the caller supplies the per-rank expert results, each destination uses that reverse mapping to push remote hidden contributions back to the source GPU. The source sums its local and received contributions in FP32, then converts the output to BF16/FP16. Optional FP32 weights are communicated and summed separately; expert computation and applying routing weights to expert outputs remain the caller's responsibility.
  • Wave32 and launch policy: Top-k=8 has dedicated kernel specializations. Other top-k values use runtime routing loops that cover all slots in wave32-sized batches, including slots 32–63. For up to 256 source tokens, Dispatch uses fewer blocks and Combine splits the hidden dimension into 512- or 1024-element tiles to expose more parallel work; larger inputs use whole-token copies and reductions.
  • Memory and synchronization: Hidden payloads use cached IPC allocations, while control buffers stay in the uncached static heap. System-scope publication, generation counters and completion acknowledgements coordinate peer visibility and buffer reuse across calls. Cached payload allocations have a 2 MiB minimum to avoid the stale peer visibility observed with smaller reused allocations on the tested gfx1201 stack.

The main implementation is in src/ops/dispatch_combine/intranode_ep2.hpp and intranode_ep4.hpp, with shared routing helpers in intranode_rdna4.hpp and launch selection in python/mori/ops/_rdna4_ep.py. docs/rdna4_ep.md documents 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_reducescatter backend. In vLLM's AgRsAll2AllManager, 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:

  • Hidden dimension: 2048, 3072, 4096, 5120, 7168, 8192.
  • Tokens per rank: 8, 32, 64, 128, 256, 512, 1024, 4096, 8192, 16384.

This gives 60 configurations per EP, 120 configurations in total. The summary, tables and heatmap all use this same matrix.

EP Correct configurations MORI faster Median latency reduction Geomean speedup
EP2 60/60 57/60 22.72% 1.279×
EP4 60/60 60/60 16.08% 1.244×

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.

gfx1201 BF16 top-k=8: EP2/EP4 Dispatch+Combine latency reduction versus RCCL AllGather+ReduceScatter

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

Tokens/rank EP2 MORI Dispatch+Combine (µs) EP2 RCCL AllGather+ReduceScatter (µs) EP2 reduction EP4 MORI Dispatch+Combine (µs) EP4 RCCL AllGather+ReduceScatter (µs) EP4 reduction
8 36.90 43.76 +15.68% 46.52 80.80 +42.43%
32 40.68 40.64 -0.10% 60.66 64.96 +6.62%
64 46.76 48.68 +3.94% 78.22 87.86 +10.97%
128 58.44 64.98 +10.06% 109.80 129.22 +15.03%
256 83.86 98.00 +14.43% 173.32 208.52 +16.88%
512 140.46 164.52 +14.62% 306.64 373.86 +17.98%
1,024 232.12 290.84 +20.19% 592.01 712.43 +16.90%
4,096 760.86 1,068.39 +28.78% 2,298.35 2,714.37 +15.33%
8,192 1,471.73 2,113.29 +30.36% 4,602.22 5,407.15 +14.89%
16,384 2,888.20 4,171.99 +30.77% 9,181.04 10,768.84 +14.74%

Hidden dimension = 3072

Tokens/rank EP2 MORI Dispatch+Combine (µs) EP2 RCCL AllGather+ReduceScatter (µs) EP2 reduction EP4 MORI Dispatch+Combine (µs) EP4 RCCL AllGather+ReduceScatter (µs) EP4 reduction
8 37.92 53.80 +29.52% 50.30 113.04 +55.50%
32 44.28 45.24 +2.12% 72.38 78.76 +8.10%
64 53.12 56.60 +6.15% 94.94 107.72 +11.86%
128 71.08 81.78 +13.08% 137.22 171.58 +20.03%
256 104.54 129.68 +19.39% 232.66 290.30 +19.86%
512 187.38 222.36 +15.73% 448.51 534.73 +16.12%
1,024 322.88 417.93 +22.74% 867.41 1,032.60 +16.00%
4,096 1,108.93 1,593.98 +30.43% 3,420.08 4,048.50 +15.52%
8,192 2,161.18 3,163.24 +31.68% 6,823.81 8,079.94 +15.55%
16,384 4,275.62 6,271.21 +31.82% 13,653.35 16,113.63 +15.27%

Hidden dimension = 4096

Tokens/rank EP2 MORI Dispatch+Combine (µs) EP2 RCCL AllGather+ReduceScatter (µs) EP2 reduction EP4 MORI Dispatch+Combine (µs) EP4 RCCL AllGather+ReduceScatter (µs) EP4 reduction
8 39.28 62.92 +37.57% 53.92 140.14 +61.52%
32 46.80 48.82 +4.14% 80.06 86.68 +7.64%
64 57.80 65.40 +11.62% 105.38 128.32 +17.88%
128 82.40 98.24 +16.12% 169.20 208.04 +18.67%
256 125.38 164.98 +24.00% 299.44 371.73 +19.44%
512 234.06 289.84 +19.25% 588.73 713.41 +17.48%
1,024 415.78 548.53 +24.20% 1,147.74 1,366.62 +16.02%
4,096 1,465.02 2,122.38 +30.97% 4,549.13 5,394.35 +15.67%
8,192 2,857.39 4,213.73 +32.19% 9,091.08 10,740.64 +15.36%
16,384 5,656.54 8,383.50 +32.53% 18,199.88 21,426.17 +15.06%

Hidden dimension = 5120

Tokens/rank EP2 MORI Dispatch+Combine (µs) EP2 RCCL AllGather+ReduceScatter (µs) EP2 reduction EP4 MORI Dispatch+Combine (µs) EP4 RCCL AllGather+ReduceScatter (µs) EP4 reduction
8 41.34 74.82 +44.75% 54.94 164.26 +66.55%
32 51.26 53.44 +4.08% 87.82 95.78 +8.31%
64 65.38 73.52 +11.07% 119.08 149.14 +20.16%
128 94.64 113.40 +16.54% 198.66 248.96 +20.20%
256 145.76 189.62 +23.13% 367.77 452.97 +18.81%
512 280.05 353.39 +20.75% 735.66 870.66 +15.51%
1,024 511.01 683.40 +25.22% 1,424.69 1,708.95 +16.63%
4,096 1,816.04 2,632.34 +31.01% 5,641.29 6,753.63 +16.47%
8,192 3,560.28 5,227.21 +31.89% 11,302.84 13,460.13 +16.03%
16,384 7,040.19 10,373.21 +32.13% 22,668.04 26,879.69 +15.67%

Hidden dimension = 7168

Tokens/rank EP2 MORI Dispatch+Combine (µs) EP2 RCCL AllGather+ReduceScatter (µs) EP2 reduction EP4 MORI Dispatch+Combine (µs) EP4 RCCL AllGather+ReduceScatter (µs) EP4 reduction
8 42.24 40.12 -5.28% 59.74 62.68 +4.69%
32 55.80 60.88 +8.34% 99.20 116.64 +14.95%
64 76.16 89.80 +15.19% 148.52 186.52 +20.37%
128 114.00 147.48 +22.70% 257.74 329.72 +21.83%
256 192.62 255.62 +24.65% 501.34 616.78 +18.72%
512 374.64 484.32 +22.65% 1,000.04 1,197.98 +16.52%
1,024 694.74 941.68 +26.22% 1,984.87 2,377.71 +16.52%
4,096 2,527.24 3,679.29 +31.31% 7,879.15 9,388.77 +16.08%
8,192 4,962.81 7,295.06 +31.97% 15,760.29 18,768.19 +16.03%
16,384 9,822.25 14,518.29 +32.35% 31,595.09 37,536.02 +15.83%

Hidden dimension = 8192

Tokens/rank EP2 MORI Dispatch+Combine (µs) EP2 RCCL AllGather+ReduceScatter (µs) EP2 reduction EP4 MORI Dispatch+Combine (µs) EP4 RCCL AllGather+ReduceScatter (µs) EP4 reduction
8 45.18 42.32 -6.76% 63.28 66.50 +4.84%
32 60.58 65.72 +7.82% 104.10 128.24 +18.82%
64 84.08 98.12 +14.31% 164.78 207.58 +20.62%
128 126.64 163.52 +22.55% 298.08 371.84 +19.84%
256 214.92 293.64 +26.81% 576.75 721.33 +20.04%
512 424.69 549.13 +22.66% 1,164.77 1,374.70 +15.27%
1,024 781.69 1,071.81 +27.07% 2,289.43 2,718.37 +15.78%
4,096 2,872.88 4,194.15 +31.50% 9,044.66 10,771.60 +16.03%
8,192 5,656.22 8,338.39 +32.17% 18,045.16 21,505.41 +16.09%
16,384 11,201.75 16,552.10 +32.32% 35,938.97 42,980.04 +16.38%

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

  • Hardware: 2 or 4 AMD Radeon AI PRO R9700 32 GB GPUs, gfx1201, connected over PCIe.
  • Software: PyTorch 2.12.0+rocm10.0.0, reported HIP runtime 7.15.26333, RCCL 2.30.4.
  • MORI configuration: MORI_RDNA4_EP=1, SHMEM static heap, external input buffers, no quantization.
  • Routing: each token selects eight distinct experts uniformly across 256 experts. MORI sends one hidden payload per destination rank.
  • Timing: 10 ordinary-call warmups and 10 graph warmups per backend; 4 rounds × 30 samples, alternating backend order between rounds. Each sample uses the maximum latency across ranks; p50 is taken over those 120 maxima. Pair p50 is measured directly, rather than adding separately aggregated phase p50 values.
  • Both backends pass correctness checks for every listed configuration. None of these 120 points exceeds the recorded round-spread threshold (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 python interpreter. If using a container, expose the GPUs through /dev/kfd and /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; omit sudo when already running as root in a container.

sudo apt-get update
sudo apt-get install -y \
  git build-essential python3-dev python3-pip cmake ninja-build \
  libpci-dev libibverbs-dev ibverbs-utils libnuma-dev libdrm-dev

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.

export ROCM_PATH="${ROCM_PATH:-/opt/rocm}"
export PATH="$ROCM_PATH/bin:$ROCM_PATH/llvm/bin:$PATH"
test -x "$ROCM_PATH/bin/hipcc"
"$ROCM_PATH/bin/hipcc" --version

python - <<'PY'
import torch
assert torch.version.hip is not None, "Use a ROCm build of PyTorch"
assert torch.cuda.is_available(), "No usable ROCm GPUs are visible"
print("PyTorch:", torch.__version__, "HIP:", torch.version.hip)
for device in range(torch.cuda.device_count()):
    p = torch.cuda.get_device_properties(device)
    print(device, p.name, p.gcnArchName, p.total_memory)
PY

Identify four idle gfx1201 GPUs by their HIP indices, which can differ from rocm-smi indices. The commands below use HIP 4,5 for EP2 and 4,5,6,7 for 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:

export MORI_PR_NUMBER=712
git clone --branch main https://github.com/ROCm/mori.git mori-gfx1201
cd mori-gfx1201
git fetch origin "refs/pull/${MORI_PR_NUMBER}/head"
git checkout --detach 283fa687e047380568b1367de7d0729483f90d33
test "$(git rev-parse HEAD^{tree})" = 72cd4fcfa1436fc328af795cf22a4760dfcdd63c
git submodule update --init --recursive
export MORI_REPRO_ROOT="$PWD"

If that revision is already checked out locally, set MORI_REPRO_ROOT to its directory and continue with the build.

3. Compile and install MORI with gfx1201 support

MORI_GPU_ARCHS=gfx1201 selects 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.

python -m pip install -r requirements-build.txt pytest

export MORI_GPU_ARCHS=gfx1201
export BUILD_UMBP=OFF
export BUILD_EXAMPLES=OFF
export BUILD_BENCHMARK=OFF
export BUILD_TESTS=OFF
export MORI_WITH_MPI=OFF
export MORI_SKIP_PRECOMPILE=1

python -m pip install --no-build-isolation -v .
python tools/verify_install.py

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.txt includes Cython for the required CCO extension. MORI_SKIP_PRECOMPILE=1 defers device compilation to the explicit JIT step below.

4. Enable the RDNA4 transport and compile its device kernels

export PYTHONPATH="$MORI_REPRO_ROOT/python:$MORI_REPRO_ROOT"
export MORI_GPU_ARCHS=gfx1201
export MORI_RDNA4_EP=1
export MORI_EP_COMM=shmem
export MORI_SHMEM_MODE=static_heap
export MORI_SHMEM_HEAP_SIZE=2G
export MORI_JIT_CACHE_DIR="$MORI_REPRO_ROOT/build/topk8_repro_jit"
export OMP_NUM_THREADS=1
export MORI_TEST_WORKER_TIMEOUT=2400
export HSA_NO_SCRATCH_RECLAIM=1
export HIP_FORCE_DEV_KERNARG=1

python - <<'PY'
from mori.jit import compile_genco
for name in ("ep_intranode", "shmem_kernels"):
    print(name, compile_genco(name))
PY

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 with python. 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.

MORI_RDNA4_EP=0 python -m pytest -q \
  tests/python/test_rdna4_ep.py \
  tests/python/test_rdna4_ep_integration.py \
  tests/python/test_tuning_config_quant_types.py

HIP_VISIBLE_DEVICES=4,5 VERIFY_WORLD_SIZE=2 VERIFY_TOPKS=8 \
  VERIFY_HIDDEN_SIZES=2048,3072,4096,5120,7168,8192 \
  VERIFY_CAPACITIES=512,256 VERIFY_HANDLE_INTERLEAVE=1 VERIFY_DOUBLE_COMBINE=1 \
  python tests/python/ops/verify_rdna4_ep.py

HIP_VISIBLE_DEVICES=4,5,6,7 VERIFY_WORLD_SIZE=4 VERIFY_TOPKS=8 \
  VERIFY_HIDDEN_SIZES=2048,3072,4096,5120,7168,8192 \
  VERIFY_CAPACITIES=512,256 VERIFY_HANDLE_INTERLEAVE=1 VERIFY_DOUBLE_COMBINE=1 \
  python tests/python/ops/verify_rdna4_ep.py

Each driver must complete and exit with status zero. To validate the PR's entire supported top-k range, replace VERIFY_TOPKS=8 with VERIFY_TOPKS=all.

6. Reproduce all 60 performance configurations per EP

The performance entry point is tests/python/ops/bench_rdna4_ep.py. Launch it through torch.distributed.run with two or four workers. It checks both backends before timing and records raw samples and aggregate results in JSONL.

HIP_VISIBLE_DEVICES=4,5 python -m torch.distributed.run \
  --standalone --nproc_per_node=2 tests/python/ops/bench_rdna4_ep.py \
  --topks 8 --dtypes bf16 \
  --hidden-sizes 2048,3072,4096,5120,7168,8192 \
  --tokens 8,32,64,128,256,512,1024,4096,8192,16384 \
  --experts 256 --weights on --warmup 10 --rounds 4 --samples 30 \
  --seed 20261225 --output logs/topk8-ep2-new.jsonl

HIP_VISIBLE_DEVICES=4,5,6,7 python -m torch.distributed.run \
  --standalone --nproc_per_node=4 tests/python/ops/bench_rdna4_ep.py \
  --topks 8 --dtypes bf16 \
  --hidden-sizes 2048,3072,4096,5120,7168,8192 \
  --tokens 8,32,64,128,256,512,1024,4096,8192,16384 \
  --experts 256 --weights on --warmup 10 --rounds 4 --samples 30 \
  --seed 20261225 --output logs/topk8-ep4-new.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 bf16 to --dtypes bf16,fp16 to 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.

@kawhil-amd
kawhil-amd requested review from jhchouuu and zhangfei829 and removed request for zhangfei829 September 28, 2026 03:13
@kawhil-amd

kawhil-amd commented Sep 28, 2026 •

Copy link
Copy Markdown
Contributor

nice work, do you try to use mori-cco as the backend rather than shmem? @dongdongzhao1121

@dongdongzhao1121

dongdongzhao1121 commented Sep 28, 2026 •

Copy link
Copy Markdown
Author

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

@kawhil-amd

Copy link
Copy Markdown
Contributor

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

@dongdongzhao1121

Copy link
Copy Markdown
Author

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?

@kawhil-amd

Copy link
Copy Markdown
Contributor

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.

@dongdongzhao1121

Copy link
Copy Markdown
Author

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 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