@@ -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}}
604666trap 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
891960echo "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
0 commit comments