diff --git a/python/sglang/bench_one_batch.py b/python/sglang/bench_one_batch.py index b462bb7d2..6c845af1d 100644 --- a/python/sglang/bench_one_batch.py +++ b/python/sglang/bench_one_batch.py @@ -92,11 +92,25 @@ from sglang.srt.utils.hf_transformers_utils import get_tokenizer from sglang.srt.utils.tensor_bridge import use_mlx -def start_profile(profile_activities, profile_record_shapes=False, rank_print=print): +def start_profile( + profile_activities, + profile_record_shapes=False, + rank_print=print, + trace_filename=None, +): """ Abstracted function to start profiling based on profile_activities. Returns profiler object (or None). """ + if use_mlx(): + import mlx.core as mx + + if trace_filename: + mlx_trace_filename = trace_filename.replace(".trace.json.gz", ".gputrace") + mx.metal.start_capture(mlx_trace_filename) + rank_print(f"MLX Metal capture started directly to {mlx_trace_filename}") + return "mlx" + if "CUDA_PROFILER" in profile_activities: try: torch.cuda.cudart().cudaProfilerStart() @@ -135,6 +149,19 @@ def stop_profile( Abstracted function to stop profiling based on profile_activities. Optionally saves trace results and prints completion messages. """ + if profiler == "mlx": + import mlx.core as mx + + mx.metal.stop_capture() + + if save_trace and trace_filename: + # Change SGLang's default torch extension to Apple's .gputrace extension + mlx_trace_filename = trace_filename.replace(".trace.json.gz", ".gputrace") + + stage_desc = f"for {stage}" if stage else "" + rank_print(f"MLX Metal gputrace {stage_desc} saved to {mlx_trace_filename}") + return + if "CUDA_PROFILER" in profile_activities: try: torch.cuda.cudart().cudaProfilerStop() @@ -659,11 +686,16 @@ def latency_test_run_once( profiler = None enable_profile_prefill = profile and profile_stage in ["all", "prefill"] + trace_filename_prefill = None if enable_profile_prefill: + trace_filename_prefill = _create_torch_profiler_filename( + profile_filename_prefix, batch_size, input_len, output_len, "prefill" + ) profiler = start_profile( profile_activities, profile_record_shapes=profile_record_shapes, rank_print=rank_print, + trace_filename=trace_filename_prefill, # pass it in here for the MLX path only ) model_runner.synchronize() @@ -673,15 +705,12 @@ def latency_test_run_once( prefill_latency = time.perf_counter() - tic if enable_profile_prefill: - trace_filename = _create_torch_profiler_filename( - profile_filename_prefix, batch_size, input_len, output_len, "prefill" - ) stop_profile( profiler, profile_activities, rank_print=rank_print, save_trace=True, - trace_filename=trace_filename, + trace_filename=trace_filename_prefill, stage="prefill", ) @@ -700,15 +729,20 @@ def latency_test_run_once( ) profile_end = profile_start + (profile_steps if profile_steps is not None else 1) enable_profile_decode = profile and profile_stage in ["all", "decode"] + trace_filename_decode = None profiler = None for i in range(output_len - 1): model_runner.synchronize() # Start profiler at the specified step if enable_profile_decode and i == profile_start: + trace_filename_decode = _create_torch_profiler_filename( + profile_filename_prefix, batch_size, input_len, output_len, "decode" + ) profiler = start_profile( profile_activities, profile_record_shapes=profile_record_shapes, rank_print=rank_print, + trace_filename=trace_filename_decode, ) tic = time.perf_counter() @@ -718,15 +752,12 @@ def latency_test_run_once( # Stop profiler after the specified number of steps if enable_profile_decode and profiler is not None and i >= profile_end - 1: - trace_filename = _create_torch_profiler_filename( - profile_filename_prefix, batch_size, input_len, output_len, "decode" - ) stop_profile( profiler, profile_activities, rank_print=rank_print, save_trace=True, - trace_filename=trace_filename, + trace_filename=trace_filename_decode, stage="decode", ) profiler = None