Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 11 additions & 1 deletion python/freetoken/engine/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -425,7 +425,17 @@ def __init__(self, config: EngineConfig):
self._warmup_prefill()

def _init_communication(self, config: EngineConfig) -> torch.distributed.ProcessGroup:
if config.tp_info.size == 1 or config.use_pynccl:
use_pynccl = config.use_pynccl
if config.tp_info.size > 1 and use_pynccl:
from freetoken.kernel.backend import is_rocm

if is_rocm():
logger.warning_rank0(
"PyNCCL is NVIDIA-only; using PyTorch's ROCm/RCCL process group instead"
)
use_pynccl = False

if config.tp_info.size == 1 or use_pynccl:
torch.distributed.init_process_group(
backend="gloo",
rank=config.tp_info.rank,
Expand Down
74 changes: 74 additions & 0 deletions tests/engine/test_rocm_communication.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
from types import SimpleNamespace

import torch

import freetoken.engine.engine as engine_module
import freetoken.kernel.backend as kernel_backend
from freetoken.engine.engine import Engine


def _config(*, use_pynccl: bool = True):
return SimpleNamespace(
use_pynccl=use_pynccl,
tp_info=SimpleNamespace(size=2, rank=0),
distributed_timeout=10,
distributed_addr="tcp://127.0.0.1:29500",
max_forward_len=32,
model_config=SimpleNamespace(hidden_size=64),
)


def test_rocm_routes_tensor_parallel_communication_to_rccl(monkeypatch):
calls = []
cpu_group = object()

def reject_pynccl(*_args):
raise AssertionError("PyNCCL selected on ROCm")

monkeypatch.setattr(kernel_backend, "is_rocm", lambda: True)
monkeypatch.setattr(
torch.distributed,
"init_process_group",
lambda **kwargs: calls.append(("init", kwargs)),
)
monkeypatch.setattr(
torch.distributed,
"new_group",
lambda **kwargs: calls.append(("new", kwargs)) or cpu_group,
)
monkeypatch.setattr(engine_module, "enable_pynccl_distributed", reject_pynccl)

result = Engine._init_communication(SimpleNamespace(), _config())

assert result is cpu_group
assert calls[0][1]["backend"] == "nccl"
assert calls[1] == ("new", {"backend": "gloo"})


def test_cuda_keeps_custom_pynccl_path(monkeypatch):
calls = []
world_group = object()

def reject_new_group(**_kwargs):
raise AssertionError("unexpected RCCL path")

monkeypatch.setattr(kernel_backend, "is_rocm", lambda: False)
monkeypatch.setattr(
torch.distributed,
"init_process_group",
lambda **kwargs: calls.append(("init", kwargs)),
)
monkeypatch.setattr(torch.distributed, "group", SimpleNamespace(WORLD=world_group))
monkeypatch.setattr(torch.distributed, "new_group", reject_new_group)
monkeypatch.setattr(
engine_module,
"enable_pynccl_distributed",
lambda *args: calls.append(("pynccl", args)),
)

engine = SimpleNamespace(dtype=torch.float16)
result = Engine._init_communication(engine, _config())

assert result is world_group
assert calls[0][1]["backend"] == "gloo"
assert calls[1][0] == "pynccl"