Skip to content
Open
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
13 changes: 11 additions & 2 deletions roll/utils/send_recv_utils.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import os
from typing import Dict

import torch
Expand Down Expand Up @@ -280,6 +281,8 @@ def _bucket_named_tensors(named_tensors: list[tuple[str, torch.Tensor]]) -> tupl


def named_tensors_from_bucket(bucket: "torch.Tensor", tensors_meta: list[dict]) -> list[tuple[str, torch.Tensor]]:
if not isinstance(bucket, torch.Tensor):
bucket = torch.from_numpy(bucket)
reconstructed = []
for i, meta in enumerate(tensors_meta):
tensor = bucket[meta["start_idx"] : meta["end_idx"]].view(meta["dtype"]).reshape(torch.Size(meta["shape"]))
Expand Down Expand Up @@ -311,12 +314,18 @@ def serialize_named_weights(named_weights: list[tuple[str, torch.Tensor]], infer

bucket, tensors_meta = _bucket_named_tensors(named_weights)

# Managed AutoDL containers block pidfd_getfd, which makes CUDA IPC
# deserialization fail even for colocated workers. Preserve the default
# zero-copy path unless this explicit portable transport is requested.
if os.getenv("ROLL_WEIGHT_SYNC_USE_CPU", "0") == "1":
bucket = bucket.detach().to("cpu").contiguous().numpy()
# PumpkinComment:
# FSDP2 will fail if using CPUOffload Policy without this check
if not getattr(bucket, "is_cuda", False):
elif not getattr(bucket, "is_cuda", False):
bucket = bucket.to(current_platform.device_type).contiguous()

monkey_patch_torch_reductions()
if getattr(bucket, "is_cuda", False):
monkey_patch_torch_reductions()

serialized_tensors = MultiprocessingSerializer.serialize({"bucket": bucket, "tensors_meta": tensors_meta})
return serialized_tensors
52 changes: 52 additions & 0 deletions tests/utils/test_send_recv_cpu_staging.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
import numpy as np
import torch

from roll.platforms import current_platform
from roll.utils.cuda_ipc_utils import MultiprocessingSerializer
from roll.utils.send_recv_utils import (
named_tensors_from_bucket,
serialize_named_weights,
)


def _make_named_weights():
return [
("weight_a", torch.arange(0, 6, dtype=torch.float32).reshape(2, 3)),
("weight_b", torch.arange(6, 10, dtype=torch.float32)),
]


def test_serialize_named_weights_cpu_staging(monkeypatch):
"""The opt-in CPU-staging transport serializes a numpy payload without CUDA IPC."""
monkeypatch.setenv("ROLL_WEIGHT_SYNC_USE_CPU", "1")
named_weights = _make_named_weights()

serialized = serialize_named_weights(named_weights, infer_strategy="vllm")

assert isinstance(serialized, bytes)
payload = MultiprocessingSerializer.deserialize(serialized)
assert isinstance(payload["bucket"], np.ndarray)

reconstructed = dict(named_tensors_from_bucket(payload["bucket"], payload["tensors_meta"]))
assert set(reconstructed) == {"weight_a", "weight_b"}
assert torch.equal(reconstructed["weight_a"], named_weights[0][1])
assert torch.equal(reconstructed["weight_b"], named_weights[1][1])


def test_serialize_named_weights_default_path_cpu_only(monkeypatch):
"""Without the opt-in env, CPU tensors take the existing `.to(current_platform.device_type)` path."""
monkeypatch.delenv("ROLL_WEIGHT_SYNC_USE_CPU", raising=False)
# Emulate a CPU-only instance so the default path never touches an accelerator.
monkeypatch.setattr(type(current_platform), "device_type", "cpu")
named_weights = _make_named_weights()

serialized = serialize_named_weights(named_weights, infer_strategy="vllm")

assert isinstance(serialized, bytes)
payload = MultiprocessingSerializer.deserialize(serialized)
assert isinstance(payload["bucket"], torch.Tensor)
assert payload["bucket"].device.type == "cpu"

reconstructed = dict(named_tensors_from_bucket(payload["bucket"], payload["tensors_meta"]))
assert torch.equal(reconstructed["weight_a"], named_weights[0][1])
assert torch.equal(reconstructed["weight_b"], named_weights[1][1])