[Piecewise Cuda Graph] rename, refactor and add more logging (#13675)

Co-authored-by: Minglei Zhu <mingleizhu1122@gmail.com>
Co-authored-by: Ke Bao <ISPObaoke@163.com>
Co-authored-by: Oasis-Git <ayw.sirius19@gmail.com>
This commit is contained in:
Stefan He
2025-11-21 13:28:38 +08:00
committed by GitHub
co-authored by Minglei Zhu Ke Bao Oasis-Git
parent 475962a139
commit d754ce973e
4 changed files with 36 additions and 16 deletions
+4 -1
View File
@@ -20,6 +20,7 @@ from sglang.srt.compilation.compilation_counter import compilation_counter
from sglang.srt.compilation.compiler_interface import EagerAdapter, InductorAdaptor from sglang.srt.compilation.compiler_interface import EagerAdapter, InductorAdaptor
from sglang.srt.compilation.cuda_piecewise_backend import CUDAPiecewiseBackend from sglang.srt.compilation.cuda_piecewise_backend import CUDAPiecewiseBackend
from sglang.srt.compilation.pass_manager import PostGradPassManager from sglang.srt.compilation.pass_manager import PostGradPassManager
from sglang.srt.utils.common import rank0_log
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -357,6 +358,7 @@ class SGLangBackend:
config: CompilationConfig, config: CompilationConfig,
graph_pool: Any, graph_pool: Any,
): ):
rank0_log(f"Initializing SGLangBackend")
assert graph_pool is not None assert graph_pool is not None
self.graph_pool = graph_pool self.graph_pool = graph_pool
@@ -375,6 +377,7 @@ class SGLangBackend:
self.inductor_config["post_grad_custom_post_pass"] = self.post_grad_pass_manager self.inductor_config["post_grad_custom_post_pass"] = self.post_grad_pass_manager
def __call__(self, graph: fx.GraphModule, example_inputs) -> Callable: def __call__(self, graph: fx.GraphModule, example_inputs) -> Callable:
rank0_log(f"SGLangBackend __call__")
base_cache_dir = os.path.expanduser( base_cache_dir = os.path.expanduser(
os.getenv("SGLANG_CACHE_DIR", "~/.cache/sglang/") os.getenv("SGLANG_CACHE_DIR", "~/.cache/sglang/")
) )
@@ -441,7 +444,7 @@ class SGLangBackend:
with open(graph_path, "w") as f: with open(graph_path, "w") as f:
f.write(src) f.write(src)
logger.debug("Computation graph saved to %s", graph_path) rank0_log(f"Computation graph saved to {graph_path}")
self._called = True self._called = True
return self.split_gm return self.split_gm
+2
View File
@@ -11,6 +11,7 @@ from typing import Any, Callable, Optional, Union
import torch import torch
from sglang.srt.compilation.compilation_config import CompilationConfig from sglang.srt.compilation.compilation_config import CompilationConfig
from sglang.srt.utils.common import rank0_log
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -129,6 +130,7 @@ def install_torch_compiled(
fullgraph: bool = True, fullgraph: bool = True,
graph_pool: Any = None, graph_pool: Any = None,
): ):
rank0_log(f"install_torch_compiled")
unbound_fwd = module.__class__.forward unbound_fwd = module.__class__.forward
if not callable(unbound_fwd): if not callable(unbound_fwd):
raise TypeError("module.__class__.forward must be callable") raise TypeError("module.__class__.forward must be callable")
@@ -59,6 +59,27 @@ if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
@contextmanager
def disable_ca_comm(tp_group):
"""
Context manager to temporarily disable custom allreduce communication.
This is used during Piecewise CUDA graph capture to avoid custom allreduce operations
that may not be compatible with graph capture.
TODO(yuwei): Fix this
"""
old_disabled = None
try:
if tp_group.ca_comm is not None:
old_disabled = tp_group.ca_comm.disabled
tp_group.ca_comm.disabled = True
yield
finally:
if tp_group.ca_comm is not None and old_disabled is not None:
tp_group.ca_comm.disabled = old_disabled
@contextmanager @contextmanager
def freeze_gc(enable_cudagraph_gc: bool): def freeze_gc(enable_cudagraph_gc: bool):
""" """
@@ -207,7 +228,7 @@ class PiecewiseCudaGraphRunner:
) )
with set_compiled(True): with set_compiled(True):
self.warmup_and_capture() self.warmup_torch_compile()
# Capture # Capture
try: try:
@@ -219,7 +240,8 @@ class PiecewiseCudaGraphRunner:
self.raw_num_tokens = 0 self.raw_num_tokens = 0
def warmup_and_capture(self): def warmup_torch_compile(self):
"""Warmup the model with a simple forward pass before CUDA graph capture."""
num_tokens = 2 num_tokens = 2
with torch.device(self.device): with torch.device(self.device):
forward_batch = ForwardBatch( forward_batch = ForwardBatch(
@@ -283,7 +305,7 @@ class PiecewiseCudaGraphRunner:
with set_forward_context( with set_forward_context(
forward_batch, self.attention_layers, self.quant_config forward_batch, self.attention_layers, self.quant_config
): ), disable_ca_comm(self.model_runner.tp_group):
_ = self.model_runner.model.forward( _ = self.model_runner.model.forward(
forward_batch.input_ids, forward_batch.input_ids,
forward_batch.positions, forward_batch.positions,
@@ -311,10 +333,9 @@ class PiecewiseCudaGraphRunner:
# Trigger CUDA graph capture for specific shapes. # Trigger CUDA graph capture for specific shapes.
# Capture the large shapes first so that the smaller shapes # Capture the large shapes first so that the smaller shapes
# can reuse the memory pool allocated for the large shapes. # can reuse the memory pool allocated for the large shapes.
with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc): with freeze_gc(
if self.model_runner.tp_group.ca_comm is not None: self.model_runner.server_args.enable_cudagraph_gc
old_ca_disable = self.model_runner.tp_group.ca_comm.disabled ), disable_ca_comm(self.model_runner.tp_group):
self.model_runner.tp_group.ca_comm.disabled = True
avail_mem = get_available_gpu_memory( avail_mem = get_available_gpu_memory(
self.model_runner.device, self.model_runner.device,
self.model_runner.gpu_id, self.model_runner.gpu_id,
@@ -342,8 +363,6 @@ class PiecewiseCudaGraphRunner:
# Save gemlite cache after each capture # Save gemlite cache after each capture
save_gemlite_cache() save_gemlite_cache()
if self.model_runner.tp_group.ca_comm is not None:
self.model_runner.tp_group.ca_comm.disabled = old_ca_disable
def capture_one_batch_size(self, num_tokens: int): def capture_one_batch_size(self, num_tokens: int):
bs = 1 bs = 1
@@ -565,10 +584,7 @@ class PiecewiseCudaGraphRunner:
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
**kwargs, **kwargs,
) -> Union[LogitsProcessorOutput, PPProxyTensors]: ) -> Union[LogitsProcessorOutput, PPProxyTensors]:
with enable_piecewise_cuda_graph(): with enable_piecewise_cuda_graph(), disable_ca_comm(self.model_runner.tp_group):
if self.model_runner.tp_group.ca_comm is not None:
old_ca_disable = self.model_runner.tp_group.ca_comm.disabled
self.model_runner.tp_group.ca_comm.disabled = True
self.model_runner.attn_backend.init_forward_metadata(forward_batch) self.model_runner.attn_backend.init_forward_metadata(forward_batch)
static_forward_batch = self.replay_prepare(forward_batch, **kwargs) static_forward_batch = self.replay_prepare(forward_batch, **kwargs)
# Replay # Replay
@@ -599,8 +615,6 @@ class PiecewiseCudaGraphRunner:
raise NotImplementedError( raise NotImplementedError(
"PPProxyTensors is not supported in PiecewiseCudaGraphRunner yet." "PPProxyTensors is not supported in PiecewiseCudaGraphRunner yet."
) )
if self.model_runner.tp_group.ca_comm is not None:
self.model_runner.tp_group.ca_comm.disabled = old_ca_disable
def get_spec_info(self, num_tokens: int): def get_spec_info(self, num_tokens: int):
spec_info = None spec_info = None
@@ -127,6 +127,7 @@ class TiktokenTokenizer:
add_generation_prompt, add_generation_prompt,
tools=None, tools=None,
reasoning_effort=None, reasoning_effort=None,
**kwargs, # Accept additional parameters (e.g., return_dict) for compatibility
): ):
ret = self.chat_template_jinja.render( ret = self.chat_template_jinja.render(
messages=messages, add_generation_prompt=add_generation_prompt messages=messages, add_generation_prompt=add_generation_prompt