[misc] fix ray folder lint (#22905)

This commit is contained in:
Qiaolin Yu
2026-04-15 15:08:18 -07:00
committed by GitHub
parent f9792166c3
commit 0b1b07db72
2 changed files with 17 additions and 33 deletions
@@ -93,9 +93,7 @@ class RayDataParallelController(DataParallelController):
# Create actors for each DP rank sequentially # Create actors for each DP rank sequentially
for dp_rank in range(server_args.dp_size): for dp_rank in range(server_args.dp_size):
self._launch_ray_tp_group( self._launch_ray_tp_group(server_args, dp_port_args_list[dp_rank], dp_rank)
server_args, dp_port_args_list[dp_rank], dp_rank
)
def launch_dp_attention_schedulers( def launch_dp_attention_schedulers(
self, server_args: ServerArgs, port_args: PortArgs self, server_args: ServerArgs, port_args: PortArgs
@@ -168,24 +166,20 @@ class RayDataParallelController(DataParallelController):
rank_port_args.detokenizer_ipc_name = ( rank_port_args.detokenizer_ipc_name = (
port_args.detokenizer_ipc_name port_args.detokenizer_ipc_name
) )
rank_port_args.tokenizer_ipc_name = ( rank_port_args.tokenizer_ipc_name = port_args.tokenizer_ipc_name
port_args.tokenizer_ipc_name
)
local_gpu_idx = (pp_rank % pp_per_node) * tp_per_node + ( local_gpu_idx = (pp_rank % pp_per_node) * tp_per_node + (
tp_rank % tp_per_node tp_rank % tp_per_node
) )
attn_cp_rank, moe_dp_rank, moe_ep_rank = ( attn_cp_rank, moe_dp_rank, moe_ep_rank = _compute_parallelism_ranks(
_compute_parallelism_ranks(server_args, tp_rank) server_args, tp_rank
) )
# Each DP group needs a unique dist_init_addr for its own # Each DP group needs a unique dist_init_addr for its own
# torch.distributed process group. Use nccl_port which is # torch.distributed process group. Use nccl_port which is
# unique per DP group (regular DP) or shared (DP attention). # unique per DP group (regular DP) or shared (DP attention).
dist_init_addr = ( dist_init_addr = f"{self.rank0_node_ip}:{rank_port_args.nccl_port}"
f"{self.rank0_node_ip}:{rank_port_args.nccl_port}"
)
actor = SchedulerActor.options( actor = SchedulerActor.options(
num_cpus=0, num_cpus=0,
+12 -22
View File
@@ -124,28 +124,24 @@ class RayEngine(Engine):
f"Use {gpus_per_node} GPUs/node, world_size={world_size}" f"Use {gpus_per_node} GPUs/node, world_size={world_size}"
) )
dist_init_addr = ( dist_init_addr = f"{rank0_node_ip}:{server_args.port + ZMQ_TCP_PORT_DELTA}"
f"{rank0_node_ip}:{server_args.port + ZMQ_TCP_PORT_DELTA}"
)
logger.info(f"dist_init_addr: {dist_init_addr}") logger.info(f"dist_init_addr: {dist_init_addr}")
scheduler_actors = [] scheduler_actors = []
for node_idx in range(nnodes): for node_idx in range(nnodes):
bundle_idx = bundle_for_node[node_idx] bundle_idx = bundle_for_node[node_idx]
pp_range, tp_range, pp_per_node, tp_per_node = ( pp_range, tp_range, pp_per_node, tp_per_node = _calculate_rank_ranges(
_calculate_rank_ranges( nnodes,
nnodes, server_args.pp_size,
server_args.pp_size, server_args.tp_size,
server_args.tp_size, node_rank=node_idx,
node_rank=node_idx,
)
) )
for pp_rank in pp_range: for pp_rank in pp_range:
for tp_rank in tp_range: for tp_rank in tp_range:
local_gpu_idx = ( local_gpu_idx = (pp_rank % pp_per_node) * tp_per_node + (
pp_rank % pp_per_node tp_rank % tp_per_node
) * tp_per_node + (tp_rank % tp_per_node) )
attn_cp_rank, moe_dp_rank, moe_ep_rank = ( attn_cp_rank, moe_dp_rank, moe_ep_rank = (
_compute_parallelism_ranks(server_args, tp_rank) _compute_parallelism_ranks(server_args, tp_rank)
@@ -182,12 +178,8 @@ class RayEngine(Engine):
try: try:
ray.kill(actor) ray.kill(actor)
except Exception: except Exception:
logger.error( logger.error(f"Failed to kill Ray scheduler actor: {actor}")
f"Failed to kill Ray scheduler actor: {actor}" raise RuntimeError(f"Scheduler actor failed to initialize: {e}")
)
raise RuntimeError(
f"Scheduler actor failed to initialize: {e}"
)
event_loop_refs = [ event_loop_refs = [
actor.run_event_loop.remote() for actor in scheduler_actors actor.run_event_loop.remote() for actor in scheduler_actors
@@ -197,9 +189,7 @@ class RayEngine(Engine):
try: try:
ray.get(event_loop_refs) ray.get(event_loop_refs)
except Exception as e: except Exception as e:
logger.error( logger.error(f"Ray scheduler actor terminated with error: {e}")
f"Ray scheduler actor terminated with error: {e}"
)
return ( return (
RaySchedulerInitResult( RaySchedulerInitResult(