Skip to content

fix(rocm): route tensor parallel communication through RCCL - #135

Draft
zihaomu wants to merge 9 commits into
FlashML-org:mainfrom
zihaomu:feat/rocm-rdna4-rccl
Draft

fix(rocm): route tensor parallel communication through RCCL#135
zihaomu wants to merge 9 commits into
FlashML-org:mainfrom
zihaomu:feat/rocm-rdna4-rccl

Conversation

@zihaomu

@zihaomu zihaomu commented Aug 24, 2026

Copy link
Copy Markdown

Draft follow-up to #132. This branch is based on the #132 head. Until #132 merges, GitHub Files changed also includes the foundation diff; the incremental review scope is only the engine routing change and its unit test.

Summary

  • disable the NVIDIA-only custom PyNCCL communicator for multi-GPU ROCm runs;
  • initialize PyTorch distributed with backend nccl, which the ROCm wheel implements through RCCL;
  • retain the existing PyNCCL path on CUDA;
  • add unit tests for both ROCm and CUDA routing decisions.

Incremental review scope

Relative to zihaomu:feat/rocm-rdna3-rdna4-foundation:

  • python/freetoken/engine/engine.py
  • tests/engine/test_rocm_communication.py

Diff size: 85 additions, 1 deletion.

Dependency

Validation

  • ROCm/CUDA communication-routing unit tests: 2 passed
  • two Radeon AI PRO R9700 (gfx1201) GPUs, two ranks, RCCL all-reduce: rank 0 and rank 1 both produced 3.0
  • final combined follow-up test set: 54 passed

Known limitation

This validates communication selection and an RCCL collective. Full multi-GPU model serving with real weights is not claimed yet. The PR remains Draft until #132 establishes the common ROCm foundation.

bouclem and others added 9 commits August 24, 2026 13:53
- Add hip_compat.h shim mapping CUDA runtime API to HIP equivalents
- Update pinned_tensor.cpp to compile under both nvcc and hipcc
- Add ROCm detection in arch.py (is_rocm, get_rocm_gfx_arch, is_gfx11xx_family)
- Guard NVIDIA arch checks to return None on ROCm
- Skip nvcc version check in _toolchain.py when on ROCm
- Add ROCm build path in setup.py (ROCM_HOME, amdhip64, --offload-arch)
- Add _hip_cflags() in kernel/utils.py for JIT compilation on ROCm
- Add is_rocm() and driver_hip_version() in backend.py
- Add rocm-smi fallback in __main__.py for clangd generation
- Add TODO(ROCm) for NCCL->RCCL, flashinfer/sgl_kernel ROCm builds,
  Triton autotune RDNA3 tuning, PDL equivalent, hiprtc JIT cache
- Add AMD ROCm classifier in pyproject.toml
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.

3 participants