diff --git a/roll/utils/send_recv_utils.py b/roll/utils/send_recv_utils.py index 348d525ae..8e21bd2b7 100644 --- a/roll/utils/send_recv_utils.py +++ b/roll/utils/send_recv_utils.py @@ -1,3 +1,4 @@ +import os from typing import Dict import torch @@ -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"])) @@ -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 diff --git a/tests/utils/test_send_recv_cpu_staging.py b/tests/utils/test_send_recv_cpu_staging.py new file mode 100644 index 000000000..a288418d3 --- /dev/null +++ b/tests/utils/test_send_recv_cpu_staging.py @@ -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])