Skip to content

Commit 1bca58d

Browse files
authored
Merge pull request NVIDIA#942 from NVIDIA/ipod/llm-sigterm
[LLM] Gracefully shutdown processes with scancel
2 parents 25cfb79 + 617b64d commit 1bca58d

13 files changed

Lines changed: 864 additions & 119 deletions

File tree

conf/experimental/vllm/test_scenario/vllm.toml

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,27 @@
1616

1717
name = "vllm"
1818

19+
[[Tests]]
20+
id = "vllm.disagg.4nodes"
21+
test_name = "vllm"
22+
num_nodes = 4
23+
time_limit = "00:30:00"
24+
25+
[Tests.cmd_args.prefill]
26+
num_nodes = 2
27+
enforce_eager = ""
28+
tensor_parallel_size = 8
29+
max_num_batched_tokens = 1024
30+
31+
[Tests.cmd_args.decode]
32+
num_nodes = 2
33+
enforce_eager = ""
34+
tensor_parallel_size = 8
35+
max_num_batched_tokens = 1024
36+
37+
[Tests.extra_env_vars]
38+
CUDA_VISIBLE_DEVICES = "0,1,2,3"
39+
1940
[[Tests]]
2041
id = "vllm.agg.1node"
2142
test_name = "vllm"

src/cloudai/workloads/common/llm_serving.py

Lines changed: 103 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -565,41 +565,103 @@ def generate_wait_for_health_function(self) -> str:
565565
return 1
566566
}}"""
567567

568-
def generate_cleanup_function(self, pid_vars: list[str], timeout: int = 15) -> str:
569-
if len(pid_vars) == 1:
570-
pid_var = pid_vars[0]
571-
return f"""\
572-
cleanup() {{
573-
echo "Cleaning up PIDs: {pid_var}=${pid_var}"
574-
kill -TERM "${pid_var}" 2>/dev/null
575-
i=0
576-
while kill -0 "${pid_var}" 2>/dev/null; do
577-
[ "$i" -ge {timeout} ] && echo "PID did not exit in time" && return 1
578-
sleep 1
579-
i=$((i+1))
580-
done
581-
}}
582-
trap cleanup EXIT"""
568+
def slurm_step_name(self, role: str) -> str:
569+
return f"cloudai-{self.workload_slug}-{role.replace('_', '-')}"
583570

584-
pid_values = " ".join(f"{pid_var}=${pid_var}" for pid_var in pid_vars)
585-
pid_array = " ".join(f'"${p}"' for p in pid_vars)
571+
@staticmethod
572+
def slurm_step_id_var(role: str) -> str:
573+
return f"{role.replace('-', '_').upper()}_STEP_IDS"
574+
575+
def cleanup_step_id_vars(self, pid_vars: list[str]) -> list[str]:
576+
return [f"{pid_var.removesuffix('_PID')}_STEP_IDS" for pid_var in pid_vars]
577+
578+
def cleanup_prelude(self) -> str:
579+
return ""
580+
581+
def cleanup_guard_var(self) -> str:
582+
run_id = re.sub(r"\W+", "_", f"{self.test_run.name}_{self.test_run.current_iteration}").strip("_")
583+
return f"CLOUDAI_{run_id.upper()}_CLEANUP_DONE"
584+
585+
def _with_slurm_step_name(self, srun_prefix: str, role: str) -> str:
586+
return f"{srun_prefix} --job-name={self.slurm_step_name(role)}"
587+
588+
def render_step_discovery(self, role: str, step_id_var: str | None = None, expected_count: int = 1) -> str:
589+
step_id_var = step_id_var or self.slurm_step_id_var(role)
590+
step_name = self.slurm_step_name(role)
591+
squeue_cmd = (
592+
'squeue --noheader --steps --job "$SLURM_JOB_ID" --format="%i %j" 2>/dev/null '
593+
f"| awk '$2 == \"{step_name}\" {{ print $1 }}'"
594+
)
586595
return f"""\
587-
cleanup() {{
588-
echo "Cleaning up PIDs: {pid_values}"
596+
{step_id_var}=
597+
for _ in {{1..10}}; do
598+
{step_id_var}=$({squeue_cmd})
599+
CLOUDAI_STEP_COUNT=$(printf '%s\\n' "${{{step_id_var}}}" | wc -w | tr -d ' ')
600+
if [ "$CLOUDAI_STEP_COUNT" -ge {expected_count} ]; then break; fi
601+
sleep 1
602+
done
603+
echo "Slurm step IDs for {role}: ${{{step_id_var}:-unknown}}"
604+
"""
589605

590-
for pid in {pid_array}; do
591-
[ -n "$pid" ] && kill -TERM "$pid" 2>/dev/null
592-
done
606+
def _render_cleanup_signal_block(self, pid_var: str, step_id_var: str, signal_name: str) -> str:
607+
pid_signal = "KILL" if signal_name == "KILL" else "TERM"
608+
return f"""\
609+
if [ -n "${{{step_id_var}:-}}" ]; then
610+
for step_id in ${{{step_id_var}}}; do
611+
scancel --signal={signal_name} "$step_id" 2>/dev/null || true
612+
done
613+
elif [ -n "${{{pid_var}:-}}" ]; then
614+
kill -{pid_signal} "${{{pid_var}}}" 2>/dev/null || true
615+
fi"""
593616

594-
for pid in {pid_array}; do
595-
[ -z "$pid" ] && continue
617+
def _render_cleanup_wait_block(self, pid_var: str, step_id_var: str, timeout: int) -> str:
618+
force_block = self._render_cleanup_signal_block(pid_var, step_id_var, "KILL")
619+
force_block = "\n".join(f" {line}" for line in force_block.splitlines())
620+
return f"""\
621+
if [ -n "${{{pid_var}:-}}" ]; then
596622
i=0
597-
while kill -0 "$pid" 2>/dev/null; do
598-
[ "$i" -ge {timeout} ] && echo "PID $pid did not exit in time" && return 1
623+
while kill -0 "${{{pid_var}}}" 2>/dev/null; do
624+
if [ "$i" -ge {timeout} ]; then
625+
echo "PID ${{{pid_var}}} did not exit in time"
626+
cleanup_status=1
627+
{force_block}
628+
break
629+
fi
599630
sleep 1
600631
i=$((i+1))
601632
done
602-
done
633+
fi"""
634+
635+
def generate_cleanup_function(self, pid_vars: list[str], timeout: int = 60) -> str:
636+
step_id_vars = self.cleanup_step_id_vars(pid_vars)
637+
guard_var = self.cleanup_guard_var()
638+
pid_values = " ".join(f"{pid_var}=${pid_var}" for pid_var in pid_vars)
639+
step_values = " ".join(f"{step_id_var}=${step_id_var}" for step_id_var in step_id_vars)
640+
prelude = self.cleanup_prelude()
641+
if prelude:
642+
prelude = f"{prelude}\n"
643+
term_blocks = "\n".join(
644+
self._render_cleanup_signal_block(pid_var, step_id_var, "TERM")
645+
for pid_var, step_id_var in zip(pid_vars, step_id_vars, strict=True)
646+
)
647+
wait_blocks = "\n".join(
648+
self._render_cleanup_wait_block(pid_var, step_id_var, timeout)
649+
for pid_var, step_id_var in zip(pid_vars, step_id_vars, strict=True)
650+
)
651+
return f"""\
652+
cleanup() {{
653+
if [ "${{{guard_var}:-0}}" = "1" ]; then
654+
return 0
655+
fi
656+
{guard_var}=1
657+
cleanup_status=0
658+
{prelude} echo "Cleaning up PIDs: {pid_values}"
659+
echo "Cleaning up Slurm step IDs: {step_values}"
660+
661+
{term_blocks}
662+
663+
{wait_blocks}
664+
return "$cleanup_status"
603665
}}
604666
trap cleanup EXIT"""
605667

@@ -728,12 +790,15 @@ def render_serve_launch(
728790
head_node_var: str,
729791
nodelist_var: str,
730792
) -> str:
731-
del role, node_count, nodelist_var
793+
del node_count, nodelist_var
794+
srun_prefix = self._with_slurm_step_name(self._single_role_srun_prefix(head_node_var), role)
795+
step_id_var = self.slurm_step_id_var(role)
732796
return f"""\
733-
{self._single_role_srun_prefix(head_node_var)} \\
797+
{srun_prefix} \\
734798
--output={self.test_run.output_path.absolute()}/{log_file} \\
735799
{self._with_custom_bash(command_tail)} &
736-
{pid_var}=$!"""
800+
{pid_var}=$!
801+
{self.render_step_discovery(role, step_id_var)}"""
737802

738803
def _expand_semantic_eval_args(self, args: str, *, host: str) -> str:
739804
replacements = {
@@ -792,11 +857,15 @@ def _gen_aggregated_script(self, serve_cmd: list[str], bench_cmd: str) -> str:
792857
node_setup = self.generate_aggregated_node_setup(serve_node_count)
793858
preamble = self.aggregated_script_preamble()
794859
if legacy_single_node:
860+
serve_srun_prefix = self._with_slurm_step_name(
861+
f"{srun_prefix} --overlap --ntasks-per-node=1 --ntasks=1", "serve"
862+
)
795863
serve_launch = f"""\
796-
{srun_prefix} --overlap --ntasks-per-node=1 --ntasks=1 \\
864+
{serve_srun_prefix} \\
797865
--output={(self.test_run.output_path / self.serve_log_file).absolute()} \\
798866
{self._with_custom_bash(serve_cmd_with_env)} &
799-
{self.serve_pid_var}=$!"""
867+
{self.serve_pid_var}=$!
868+
{self.render_step_discovery("serve", self.slurm_step_id_var("serve"))}"""
800869
else:
801870
serve_launch = self.render_serve_launch(
802871
"serve",
@@ -889,10 +958,11 @@ def _gen_disaggregated_script(self, serve_commands: list[list[str]], bench_cmd:
889958
{wait_block}
890959
891960
echo "Starting {self.proxy_router_name}..."
892-
{prefill_srun_prefix} \\
961+
{self._with_slurm_step_name(prefill_srun_prefix, "helper")} \\
893962
--output={self.test_run.output_path.absolute()}/{self.proxy_router_log_file} \\
894963
{self._with_custom_bash(" ".join(helper_cmd))} &
895964
{self.proxy_router_pid_var}=$!
965+
{self.render_step_discovery("helper", self.slurm_step_id_var("helper"))}
896966
897967
{wait_block_helper}
898968

src/cloudai/workloads/vllm/slurm_command_gen_strategy.py

Lines changed: 27 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -174,6 +174,22 @@ def disaggregated_cleanup_pid_vars(self) -> list[str]:
174174
pid_vars.insert(insert_at, "DECODE_RAY_PID")
175175
return pid_vars
176176

177+
def cleanup_step_id_vars(self, pid_vars: list[str]) -> list[str]:
178+
step_id_vars: list[str] = []
179+
for pid_var in pid_vars:
180+
base = pid_var.removesuffix("_PID")
181+
if base.endswith("_RAY"):
182+
step_id_vars.append(f"{base}_WORKER_STEP_IDS")
183+
continue
184+
185+
role = base.lower()
186+
if role in {"serve", "prefill", "decode"} and self._needs_ray(role):
187+
step_id_vars.append(f"{base}_RAY_HEAD_STEP_IDS")
188+
continue
189+
190+
step_id_vars.append(f"{base}_STEP_IDS")
191+
return step_id_vars
192+
177193
@property
178194
def proxy_router_healthcheck(self) -> str:
179195
fields_set = self.tdef.cmd_args.model_fields_set
@@ -216,12 +232,8 @@ def _ray_stop_cleanup_block(self) -> str:
216232
)
217233
return "\n".join(lines)
218234

219-
def generate_cleanup_function(self, pid_vars: list[str], timeout: int = 15) -> str:
220-
cleanup = super().generate_cleanup_function(pid_vars, timeout)
221-
ray_stop_block = self._ray_stop_cleanup_block()
222-
if not ray_stop_block:
223-
return cleanup
224-
return cleanup.replace("cleanup() {\n", f"cleanup() {{\n{ray_stop_block}\n", 1)
235+
def cleanup_prelude(self) -> str:
236+
return self._ray_stop_cleanup_block()
225237

226238
def render_serve_launch(
227239
self,
@@ -246,8 +258,12 @@ def render_serve_launch(
246258
ray_worker_log = f"{self.workload_slug}-{role}-ray-worker-%N.log"
247259
serve_log = f"{self.test_run.output_path.absolute()}/{log_file}"
248260
head_node_expr = f"${{{head_node_var}}}"
249-
worker_prefix = self._role_srun_prefix("$node")
250-
head_prefix = self._single_role_srun_prefix(head_node_var)
261+
worker_role = f"{role}-ray-worker"
262+
head_role = f"{role}-ray-head"
263+
worker_step_id_var = f"{role_prefix}_RAY_WORKER_STEP_IDS"
264+
head_step_id_var = f"{role_prefix}_RAY_HEAD_STEP_IDS"
265+
worker_prefix = self._with_slurm_step_name(self._role_srun_prefix("$node"), worker_role)
266+
head_prefix = self._with_slurm_step_name(self._single_role_srun_prefix(head_node_var), head_role)
251267
serve_cmd = self._with_custom_bash(f'env RAY_ADDRESS="{head_node_expr}:${{{ray_port_var}}}" {command_tail}')
252268
ray_head_args = self._ray_start_args(role, "head", {"head": True, "port": f'"${{{ray_port_var}}}"'})
253269
ray_worker_args = self._ray_start_args(
@@ -301,11 +317,13 @@ def render_serve_launch(
301317
wait
302318
) &
303319
{ray_pid_var}=$!
320+
{self.render_step_discovery(worker_role, worker_step_id_var, node_count - 1)}
304321
{head_prefix} \\
305322
--output={serve_log} \\
306323
--error={self.test_run.output_path.absolute()}/{ray_head_log} \\
307324
bash -lc {ray_head_command} &
308-
{pid_var}=$!"""
325+
{pid_var}=$!
326+
{self.render_step_discovery(head_role, head_step_id_var)}"""
309327

310328
def disaggregated_role_env(self, role: str, gpu_ids: list[int]) -> dict[str, str]:
311329
env = super().disaggregated_role_env(role, gpu_ids)

0 commit comments

Comments
 (0)