[MLX] Add Metal profiling hooks to server profiler (#28122)
This commit is contained in:
@@ -0,0 +1,261 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import gzip
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Callable, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.managers.io_struct import ProfileReqOutput
|
||||||
|
from sglang.srt.utils.tensor_bridge import use_mlx
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class MetalCaptureProfiler:
|
||||||
|
label: str
|
||||||
|
trace_path: Path
|
||||||
|
stop_capture: Callable[[], None]
|
||||||
|
standalone: bool
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def start_mlx(cls, trace_path: Path):
|
||||||
|
trace_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
try:
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
mx.metal.start_capture(str(trace_path))
|
||||||
|
except RuntimeError as e:
|
||||||
|
return None, _capture_error("MLX", e)
|
||||||
|
|
||||||
|
return cls._started(
|
||||||
|
label="MLX",
|
||||||
|
trace_path=trace_path,
|
||||||
|
stop_capture=mx.metal.stop_capture,
|
||||||
|
standalone=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def start_mps(cls, trace_path: Path):
|
||||||
|
trace_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
try:
|
||||||
|
if not hasattr(torch, "mps") or not hasattr(torch.mps, "profiler"):
|
||||||
|
raise RuntimeError("torch.mps.profiler is not available")
|
||||||
|
context = torch.mps.profiler.metal_capture(str(trace_path))
|
||||||
|
context.__enter__()
|
||||||
|
except RuntimeError as e:
|
||||||
|
return None, _capture_error("MPS", e)
|
||||||
|
|
||||||
|
return cls._started(
|
||||||
|
label="MPS",
|
||||||
|
trace_path=trace_path,
|
||||||
|
stop_capture=lambda: context.__exit__(None, None, None),
|
||||||
|
standalone=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _started(
|
||||||
|
cls,
|
||||||
|
*,
|
||||||
|
label: str,
|
||||||
|
trace_path: Path,
|
||||||
|
stop_capture: Callable[[], None],
|
||||||
|
standalone: bool,
|
||||||
|
):
|
||||||
|
profiler = cls(
|
||||||
|
label=label,
|
||||||
|
trace_path=trace_path,
|
||||||
|
stop_capture=stop_capture,
|
||||||
|
standalone=standalone,
|
||||||
|
)
|
||||||
|
logger.info("%s Metal capture started, saving to %s", label, trace_path)
|
||||||
|
return profiler, ProfileReqOutput(success=True, message="Succeeded")
|
||||||
|
|
||||||
|
def stop(self) -> str:
|
||||||
|
self.stop_capture()
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"%s Metal capture stopped. Trace saved to: %s",
|
||||||
|
self.label,
|
||||||
|
self.trace_path,
|
||||||
|
)
|
||||||
|
return f" Metal trace: {self.trace_path}"
|
||||||
|
|
||||||
|
|
||||||
|
def _capture_error(label: str, error: RuntimeError) -> ProfileReqOutput:
|
||||||
|
return ProfileReqOutput(
|
||||||
|
success=False,
|
||||||
|
message=(
|
||||||
|
f"Failed to start {label} Metal capture: {error}. "
|
||||||
|
"Set MTL_CAPTURE_ENABLED=1 in the server's environment "
|
||||||
|
"before launching to enable GPU trace capture."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class MetalTorchProfiler:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
start_metal_capture: Callable[[Path], tuple[Any, ProfileReqOutput]],
|
||||||
|
torch_profiler: Optional[Any] = None,
|
||||||
|
):
|
||||||
|
self.start_metal_capture = start_metal_capture
|
||||||
|
self.torch_profiler = torch_profiler
|
||||||
|
self.metal_profiler = None
|
||||||
|
|
||||||
|
def start(self):
|
||||||
|
trace_path = _new_temp_gputrace_path()
|
||||||
|
self.metal_profiler, result = self.start_metal_capture(trace_path)
|
||||||
|
if not result.success:
|
||||||
|
raise RuntimeError(result.message)
|
||||||
|
if self.torch_profiler is not None:
|
||||||
|
try:
|
||||||
|
self.torch_profiler.start()
|
||||||
|
except Exception:
|
||||||
|
self.metal_profiler.stop()
|
||||||
|
raise
|
||||||
|
|
||||||
|
def stop(self):
|
||||||
|
try:
|
||||||
|
if self.torch_profiler is not None:
|
||||||
|
self.torch_profiler.stop()
|
||||||
|
finally:
|
||||||
|
if self.metal_profiler is not None:
|
||||||
|
self.metal_profiler.stop()
|
||||||
|
|
||||||
|
def export_chrome_trace(self, path: str):
|
||||||
|
if self.torch_profiler is not None:
|
||||||
|
self.torch_profiler.export_chrome_trace(path)
|
||||||
|
else:
|
||||||
|
_write_empty_chrome_trace(path)
|
||||||
|
|
||||||
|
if self.metal_profiler is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
final_path = _unique_gputrace_path_for_chrome_trace(path)
|
||||||
|
final_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
if self.metal_profiler.trace_path.exists():
|
||||||
|
shutil.move(str(self.metal_profiler.trace_path), str(final_path))
|
||||||
|
logger.info("Metal trace saved to: %s", final_path)
|
||||||
|
|
||||||
|
|
||||||
|
def apply_metal_profiler_patches() -> None:
|
||||||
|
if getattr(torch.profiler.profile, "_sglang_metal_patched", False):
|
||||||
|
return
|
||||||
|
|
||||||
|
original_profile = torch.profiler.profile
|
||||||
|
|
||||||
|
def profile(*args, **kwargs):
|
||||||
|
activities = _get_activities(args, kwargs)
|
||||||
|
if not _has_cuda_activity(activities):
|
||||||
|
return original_profile(*args, **kwargs)
|
||||||
|
|
||||||
|
if use_mlx():
|
||||||
|
return MetalTorchProfiler(
|
||||||
|
start_metal_capture=MetalCaptureProfiler.start_mlx
|
||||||
|
)
|
||||||
|
|
||||||
|
torch_activities = [
|
||||||
|
activity for activity in activities if not _is_cuda_activity(activity)
|
||||||
|
]
|
||||||
|
torch_profiler = None
|
||||||
|
if torch_activities:
|
||||||
|
patched_args, patched_kwargs = _replace_activities(
|
||||||
|
args, kwargs, torch_activities
|
||||||
|
)
|
||||||
|
torch_profiler = original_profile(*patched_args, **patched_kwargs)
|
||||||
|
|
||||||
|
return MetalTorchProfiler(
|
||||||
|
start_metal_capture=MetalCaptureProfiler.start_mps,
|
||||||
|
torch_profiler=torch_profiler,
|
||||||
|
)
|
||||||
|
|
||||||
|
profile._sglang_metal_patched = True
|
||||||
|
profile._sglang_original_profile = original_profile
|
||||||
|
torch.profiler.profile = profile
|
||||||
|
|
||||||
|
|
||||||
|
def _get_activities(args, kwargs):
|
||||||
|
if "activities" in kwargs:
|
||||||
|
return kwargs["activities"]
|
||||||
|
if args:
|
||||||
|
return args[0]
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _replace_activities(args, kwargs, activities):
|
||||||
|
kwargs = dict(kwargs)
|
||||||
|
if "activities" in kwargs:
|
||||||
|
kwargs["activities"] = activities
|
||||||
|
return args, kwargs
|
||||||
|
|
||||||
|
if args:
|
||||||
|
args = list(args)
|
||||||
|
args[0] = activities
|
||||||
|
return tuple(args), kwargs
|
||||||
|
|
||||||
|
kwargs["activities"] = activities
|
||||||
|
return args, kwargs
|
||||||
|
|
||||||
|
|
||||||
|
def _has_cuda_activity(activities) -> bool:
|
||||||
|
if activities is None:
|
||||||
|
return False
|
||||||
|
return any(_is_cuda_activity(activity) for activity in activities)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_cuda_activity(activity) -> bool:
|
||||||
|
return activity == torch.profiler.ProfilerActivity.CUDA
|
||||||
|
|
||||||
|
|
||||||
|
def _new_temp_gputrace_path() -> Path:
|
||||||
|
output_dir = Path(os.getenv("SGLANG_TORCH_PROFILER_DIR", "/tmp")).expanduser()
|
||||||
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
for i in range(100):
|
||||||
|
candidate = (
|
||||||
|
output_dir / f"sglang-metal-{os.getpid()}-{time.time_ns()}-{i}.gputrace"
|
||||||
|
)
|
||||||
|
if not candidate.exists():
|
||||||
|
return candidate
|
||||||
|
raise RuntimeError(f"Cannot find an unused Metal trace path in {output_dir}")
|
||||||
|
|
||||||
|
|
||||||
|
def _unique_gputrace_path_for_chrome_trace(path: str) -> Path:
|
||||||
|
chrome_path = Path(path).expanduser()
|
||||||
|
name = chrome_path.name
|
||||||
|
if name.endswith(".trace.json.gz"):
|
||||||
|
name = name[: -len(".trace.json.gz")] + ".gputrace"
|
||||||
|
else:
|
||||||
|
name = chrome_path.stem + ".gputrace"
|
||||||
|
|
||||||
|
base = chrome_path.with_name(name)
|
||||||
|
if not base.exists():
|
||||||
|
return base
|
||||||
|
|
||||||
|
stem = base.name[: -len(".gputrace")]
|
||||||
|
for i in range(100):
|
||||||
|
candidate = base.with_name(f"{stem}-{time.time_ns()}-{i}.gputrace")
|
||||||
|
if not candidate.exists():
|
||||||
|
return candidate
|
||||||
|
raise RuntimeError(f"Cannot find an unused Metal trace path for {base}")
|
||||||
|
|
||||||
|
|
||||||
|
def _write_empty_chrome_trace(path: str):
|
||||||
|
trace = {"traceEvents": []}
|
||||||
|
Path(path).expanduser().parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
if str(path).endswith(".gz"):
|
||||||
|
with gzip.open(path, "wt") as f:
|
||||||
|
json.dump(trace, f)
|
||||||
|
else:
|
||||||
|
with open(path, "w") as f:
|
||||||
|
json.dump(trace, f)
|
||||||
@@ -19,14 +19,16 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType
|
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import is_npu
|
from sglang.srt.utils import is_mps, is_npu
|
||||||
from sglang.srt.utils.profile_merger import ProfileMerger
|
from sglang.srt.utils.profile_merger import ProfileMerger
|
||||||
|
from sglang.srt.utils.profile_utils import ProfileManager
|
||||||
from sglang.srt.utils.torch_npu_patch_utils import apply_torch_npu_patches
|
from sglang.srt.utils.torch_npu_patch_utils import apply_torch_npu_patches
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
|
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
|
_is_mps = is_mps()
|
||||||
if _is_npu:
|
if _is_npu:
|
||||||
import torch_npu
|
import torch_npu
|
||||||
|
|
||||||
@@ -36,13 +38,14 @@ if _is_npu:
|
|||||||
["profiler.ProfilerActivity.CPU", torch_npu.profiler.ProfilerActivity.CPU],
|
["profiler.ProfilerActivity.CPU", torch_npu.profiler.ProfilerActivity.CPU],
|
||||||
]
|
]
|
||||||
apply_torch_npu_patches(torch_npu, patches)
|
apply_torch_npu_patches(torch_npu, patches)
|
||||||
|
elif _is_mps:
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import apply_metal_profiler_patches
|
||||||
|
|
||||||
|
apply_metal_profiler_patches()
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
from sglang.srt.utils.profile_utils import ProfileManager
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(kw_only=True)
|
@dataclass(kw_only=True)
|
||||||
class SchedulerProfilerManager:
|
class SchedulerProfilerManager:
|
||||||
ps: Any
|
ps: Any
|
||||||
@@ -230,7 +233,11 @@ class SchedulerProfilerManager:
|
|||||||
)
|
)
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
self.torch_profiler.start()
|
try:
|
||||||
|
self.torch_profiler.start()
|
||||||
|
except RuntimeError as e:
|
||||||
|
self.torch_profiler = None
|
||||||
|
return ProfileReqOutput(success=False, message=str(e))
|
||||||
self.profile_in_progress = True
|
self.profile_in_progress = True
|
||||||
|
|
||||||
if "MEM" in activities:
|
if "MEM" in activities:
|
||||||
|
|||||||
@@ -0,0 +1,251 @@
|
|||||||
|
"""Unit tests for MLX Metal profiling patch (hardware_backend/mlx/profiler.py).
|
||||||
|
|
||||||
|
Covers:
|
||||||
|
- apply_metal_profiler_patches() replaces torch.profiler.profile
|
||||||
|
- MLX path: MetalTorchProfiler.start/stop produces a .gputrace file
|
||||||
|
- MPS path: MetalTorchProfiler wraps torch.mps.profiler.metal_capture
|
||||||
|
- RuntimeError from start_capture is caught and returned as success=False
|
||||||
|
- SchedulerProfilerManager._start_profile returns success=False gracefully
|
||||||
|
when Metal capture fails (no MTL_CAPTURE_ENABLED)
|
||||||
|
|
||||||
|
Skips on non-Apple-Silicon platforms and when ``mlx`` is missing.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import importlib.util
|
||||||
|
import platform
|
||||||
|
import tempfile
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
_IS_APPLE_SILICON = platform.system() == "Darwin" and platform.machine() == "arm64"
|
||||||
|
_HAS_MLX = importlib.util.find_spec("mlx") is not None
|
||||||
|
_SKIP_REASON = "requires Apple Silicon and mlx"
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_IS_APPLE_SILICON and _HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestApplyMetalProfilerPatches(unittest.TestCase):
|
||||||
|
"""apply_metal_profiler_patches() replaces torch.profiler.profile."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
import torch
|
||||||
|
|
||||||
|
self._original_profile = getattr(
|
||||||
|
torch.profiler.profile, "_sglang_original_profile", None
|
||||||
|
)
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
import torch
|
||||||
|
|
||||||
|
if self._original_profile is not None:
|
||||||
|
torch.profiler.profile = self._original_profile
|
||||||
|
|
||||||
|
def test_patch_replaces_profile(self):
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import (
|
||||||
|
MetalTorchProfiler,
|
||||||
|
apply_metal_profiler_patches,
|
||||||
|
)
|
||||||
|
|
||||||
|
apply_metal_profiler_patches()
|
||||||
|
self.assertTrue(getattr(torch.profiler.profile, "_sglang_metal_patched", False))
|
||||||
|
p = torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CUDA])
|
||||||
|
self.assertIsInstance(p, MetalTorchProfiler)
|
||||||
|
|
||||||
|
def test_patch_is_idempotent(self):
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import (
|
||||||
|
apply_metal_profiler_patches,
|
||||||
|
)
|
||||||
|
|
||||||
|
apply_metal_profiler_patches()
|
||||||
|
first = torch.profiler.profile
|
||||||
|
apply_metal_profiler_patches()
|
||||||
|
self.assertIs(torch.profiler.profile, first)
|
||||||
|
|
||||||
|
def test_no_cuda_activity_uses_original(self):
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import (
|
||||||
|
MetalTorchProfiler,
|
||||||
|
apply_metal_profiler_patches,
|
||||||
|
)
|
||||||
|
|
||||||
|
apply_metal_profiler_patches()
|
||||||
|
p = torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU])
|
||||||
|
self.assertNotIsInstance(p, MetalTorchProfiler)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_IS_APPLE_SILICON and _HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestMetalCaptureProfilerMLX(unittest.TestCase):
|
||||||
|
"""MLX path: start_mlx produces a .gputrace and stop_capture is called."""
|
||||||
|
|
||||||
|
def test_start_mlx_success(self):
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import MetalCaptureProfiler
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
trace_path = Path(tmp) / "test.gputrace"
|
||||||
|
with patch.object(mx.metal, "start_capture"), patch.object(
|
||||||
|
mx.metal, "stop_capture"
|
||||||
|
):
|
||||||
|
profiler, result = MetalCaptureProfiler.start_mlx(trace_path)
|
||||||
|
|
||||||
|
self.assertTrue(result.success)
|
||||||
|
self.assertIsNotNone(profiler)
|
||||||
|
self.assertEqual(profiler.label, "MLX")
|
||||||
|
self.assertTrue(profiler.standalone)
|
||||||
|
|
||||||
|
def test_start_mlx_runtime_error_returns_failure(self):
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import MetalCaptureProfiler
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
trace_path = Path(tmp) / "test.gputrace"
|
||||||
|
with patch.object(
|
||||||
|
mx.metal,
|
||||||
|
"start_capture",
|
||||||
|
side_effect=RuntimeError("Capture layer is not inserted"),
|
||||||
|
):
|
||||||
|
profiler, result = MetalCaptureProfiler.start_mlx(trace_path)
|
||||||
|
|
||||||
|
self.assertIsNone(profiler)
|
||||||
|
self.assertFalse(result.success)
|
||||||
|
self.assertIn("MTL_CAPTURE_ENABLED", result.message)
|
||||||
|
|
||||||
|
def test_stop_calls_stop_capture(self):
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import MetalCaptureProfiler
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
trace_path = Path(tmp) / "test.gputrace"
|
||||||
|
with patch.object(mx.metal, "start_capture"), patch.object(
|
||||||
|
mx.metal, "stop_capture"
|
||||||
|
) as mock_stop:
|
||||||
|
profiler, _ = MetalCaptureProfiler.start_mlx(trace_path)
|
||||||
|
profiler.stop()
|
||||||
|
mock_stop.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_IS_APPLE_SILICON and _HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestMetalCaptureProfilerMPS(unittest.TestCase):
|
||||||
|
"""MPS path: start_mps wraps torch.mps.profiler.metal_capture."""
|
||||||
|
|
||||||
|
def test_start_mps_success(self):
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import MetalCaptureProfiler
|
||||||
|
|
||||||
|
mock_ctx = MagicMock()
|
||||||
|
mock_ctx.__enter__ = MagicMock(return_value=mock_ctx)
|
||||||
|
mock_ctx.__exit__ = MagicMock(return_value=False)
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
trace_path = Path(tmp) / "test.gputrace"
|
||||||
|
with patch.object(
|
||||||
|
torch.mps.profiler, "metal_capture", return_value=mock_ctx
|
||||||
|
):
|
||||||
|
profiler, result = MetalCaptureProfiler.start_mps(trace_path)
|
||||||
|
|
||||||
|
self.assertTrue(result.success)
|
||||||
|
self.assertIsNotNone(profiler)
|
||||||
|
self.assertEqual(profiler.label, "MPS")
|
||||||
|
self.assertFalse(profiler.standalone)
|
||||||
|
|
||||||
|
def test_start_mps_runtime_error_returns_failure(self):
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import MetalCaptureProfiler
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
trace_path = Path(tmp) / "test.gputrace"
|
||||||
|
with patch.object(
|
||||||
|
torch.mps.profiler,
|
||||||
|
"metal_capture",
|
||||||
|
side_effect=RuntimeError("MPS profiler unavailable"),
|
||||||
|
):
|
||||||
|
profiler, result = MetalCaptureProfiler.start_mps(trace_path)
|
||||||
|
|
||||||
|
self.assertIsNone(profiler)
|
||||||
|
self.assertFalse(result.success)
|
||||||
|
self.assertIn("MTL_CAPTURE_ENABLED", result.message)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_IS_APPLE_SILICON and _HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestSchedulerProfilerManagerMPS(unittest.TestCase):
|
||||||
|
"""SchedulerProfilerManager._start_profile handles Metal capture failures."""
|
||||||
|
|
||||||
|
def _make_manager(self, output_dir):
|
||||||
|
from sglang.srt.managers.scheduler_components.profiler_manager import (
|
||||||
|
SchedulerProfilerManager,
|
||||||
|
)
|
||||||
|
|
||||||
|
class FakePS:
|
||||||
|
tp_rank = dp_rank = pp_rank = moe_ep_rank = 0
|
||||||
|
dp_size = pp_size = moe_ep_size = 1
|
||||||
|
gpu_id = 0
|
||||||
|
|
||||||
|
mgr = SchedulerProfilerManager(
|
||||||
|
ps=FakePS(), dp_tp_cpu_group=None, get_forward_ct=lambda: 0
|
||||||
|
)
|
||||||
|
mgr._init_profile(output_dir, None, None, None, None, None, False, "test")
|
||||||
|
return mgr
|
||||||
|
|
||||||
|
def test_start_profile_failure_does_not_crash(self):
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import (
|
||||||
|
apply_metal_profiler_patches,
|
||||||
|
)
|
||||||
|
|
||||||
|
apply_metal_profiler_patches()
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
mgr = self._make_manager(tmp)
|
||||||
|
with patch.object(
|
||||||
|
mx.metal,
|
||||||
|
"start_capture",
|
||||||
|
side_effect=RuntimeError("Capture layer is not inserted"),
|
||||||
|
):
|
||||||
|
result = mgr._start_profile()
|
||||||
|
|
||||||
|
self.assertFalse(result.success)
|
||||||
|
self.assertFalse(mgr.profile_in_progress)
|
||||||
|
self.assertIsNone(mgr.torch_profiler)
|
||||||
|
|
||||||
|
def test_start_profile_success_with_mock_capture(self):
|
||||||
|
from unittest.mock import patch as mock_patch
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.profiler import (
|
||||||
|
apply_metal_profiler_patches,
|
||||||
|
)
|
||||||
|
|
||||||
|
apply_metal_profiler_patches()
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
mgr = self._make_manager(tmp)
|
||||||
|
with mock_patch.object(mx.metal, "start_capture"), mock_patch.object(
|
||||||
|
mx.metal, "stop_capture"
|
||||||
|
), mock_patch("torch.distributed.barrier"):
|
||||||
|
result = mgr._start_profile()
|
||||||
|
self.assertTrue(result.success)
|
||||||
|
self.assertTrue(mgr.profile_in_progress)
|
||||||
|
mgr._stop_profile()
|
||||||
|
self.assertFalse(mgr.profile_in_progress)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user