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
110 changes: 97 additions & 13 deletions tests/_unit_stubs.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,19 +43,110 @@ def real_module_available(name: str) -> bool:
return False


class _RayRemoteWrapper:
def __init__(self, target: Any):
self._target = target
if isinstance(target, type):
self.__ray_actor_class__ = target

def options(self, *args, **kwargs): # noqa: ARG002
return self

def remote(self, *args, **kwargs):
return self._target(*args, **kwargs)


def _ray_remote(*args, **kwargs): # noqa: ARG001
if len(args) == 1 and callable(args[0]) and not kwargs:
return _RayRemoteWrapper(args[0])

def decorator(target: Any):
return _RayRemoteWrapper(target)

return decorator


class _RaySchedulingStrategy:
def __init__(self, **kwargs):
self.kwargs = kwargs


def _install_ray_module_stub() -> None:
ray_mod = types.ModuleType("ray")
ray_mod.__path__ = []
ray_mod.get = lambda refs: refs
ray_mod.put = lambda value, **kwargs: value # noqa: ARG005
ray_mod.wait = lambda refs, timeout=None: (refs, []) # noqa: ARG005
ray_mod.kill = lambda actor, no_restart=False: None # noqa: ARG005
ray_mod.nodes = lambda: []
ray_mod.cluster_resources = lambda: {}
ray_mod.available_resources = lambda: {}
ray_mod.get_gpu_ids = lambda: [0]
ray_mod.remote = _ray_remote
ray_mod.ObjectRef = object

actor_mod = types.ModuleType("ray.actor")
actor_mod.ActorHandle = object

private_mod = types.ModuleType("ray._private")
private_mod.__path__ = []
services_mod = types.ModuleType("ray._private.services")
services_mod.get_node_ip_address = lambda: "127.0.0.1"
private_mod.services = services_mod

util_mod = types.ModuleType("ray.util")
util_mod.__path__ = []
util_mod.get_node_ip_address = services_mod.get_node_ip_address

scheduling_mod = types.ModuleType("ray.util.scheduling_strategies")
scheduling_mod.PlacementGroupSchedulingStrategy = _RaySchedulingStrategy
scheduling_mod.NodeAffinitySchedulingStrategy = _RaySchedulingStrategy

placement_group_mod = types.ModuleType("ray.util.placement_group")
placement_group_mod.PlacementGroup = object
placement_group_mod.placement_group = lambda *args, **kwargs: types.SimpleNamespace(ready=lambda: None) # noqa: ARG005

ray_mod.actor = actor_mod
ray_mod._private = private_mod
ray_mod.util = util_mod
util_mod.scheduling_strategies = scheduling_mod
util_mod.placement_group = placement_group_mod

sys.modules["ray"] = ray_mod
sys.modules["ray.actor"] = actor_mod
sys.modules["ray._private"] = private_mod
sys.modules["ray._private.services"] = services_mod
sys.modules["ray.util"] = util_mod
sys.modules["ray.util.scheduling_strategies"] = scheduling_mod
sys.modules["ray.util.placement_group"] = placement_group_mod


def ensure_ray_stub() -> None:
if real_module_available("ray"):
return
ray = MagicMock()
sys.modules["ray"] = ray
sys.modules["ray._private"] = MagicMock()
sys.modules["ray._private.services"] = MagicMock()
sys.modules["ray.actor"] = MagicMock()
_install_ray_module_stub()


def install_pyarrow_parquet_stub() -> None:
"""Avoid importing pyarrow's native parquet extension in CPU-only unit tests."""
pyarrow_mod = types.ModuleType("pyarrow")
pyarrow_mod.__path__ = []
parquet_mod = types.ModuleType("pyarrow.parquet")

class ParquetFile:
def __init__(self, *args, **kwargs): # noqa: ARG002
raise ImportError("pyarrow is required for parquet support")

parquet_mod.ParquetFile = ParquetFile
pyarrow_mod.parquet = parquet_mod
sys.modules["pyarrow"] = pyarrow_mod
sys.modules["pyarrow.parquet"] = parquet_mod


def install_rollout_optional_stubs() -> None:
"""Stub rollout-side optional imports when not installed."""
ensure_ray_stub()
install_pyarrow_parquet_stub()

install_vllm_router_stub()

Expand Down Expand Up @@ -216,14 +307,7 @@ def install_megatron_mpu_stub() -> MagicMock:


def install_ray_stub() -> None:
ray_mod = types.ModuleType("ray")
ray_mod.get = lambda refs: refs
ray_mod.ObjectRef = object
ray_mod.actor = types.ModuleType("ray.actor")
ray_mod.actor.ActorHandle = object
ray_mod._private = types.SimpleNamespace(services=types.SimpleNamespace(get_node_ip_address=lambda: "127.0.0.1"))
sys.modules.setdefault("ray", ray_mod)
sys.modules.setdefault("ray.actor", ray_mod.actor)
_install_ray_module_stub()


def install_vllm_cli_stubs() -> None:
Expand Down
6 changes: 6 additions & 0 deletions tests/test_sample.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,11 @@ def _make_sample(**overrides) -> Sample:
loss_mask=[1, 1, 0, 1, 1],
weight_versions=["v1"],
rollout_log_probs=[-0.1, -0.2],
consistency_metadata={
"schema_version": 1,
"sample": {"rollout_id": 7, "index": 42},
"fingerprint": "sha256:unit-test",
},
rollout_top_p_token_ids=[10, 11, 12, 20],
rollout_top_p_token_offsets=[0, 3, 4],
rollout_routed_experts=[[0, 1], [2, 3]],
Expand Down Expand Up @@ -133,6 +138,7 @@ def test_round_trip_preserves_every_field():
"loss_mask",
"weight_versions",
"rollout_log_probs",
"consistency_metadata",
"rollout_top_p_token_ids",
"rollout_top_p_token_offsets",
"rollout_routed_experts",
Expand Down
Loading