[Piecewise] support disable decode cuda graph when enable piecewise cuda graph (#13965)
This commit is contained in:
@@ -433,19 +433,24 @@ class SGLangBackend:
|
|||||||
self,
|
self,
|
||||||
).run(*example_inputs)
|
).run(*example_inputs)
|
||||||
|
|
||||||
graph_path = os.path.join(local_cache_dir, "computation_graph.py")
|
rank = torch.distributed.get_rank()
|
||||||
if not os.path.exists(graph_path):
|
|
||||||
# code adapted from https://github.com/thuml/depyf/blob/dab831108a752d1facc00acdd6d4243891845c37/depyf/explain/patched_lazy_format_graph_code.py#L30 # noqa
|
|
||||||
# use `print_readable` because it can include submodules
|
|
||||||
src = (
|
|
||||||
"from __future__ import annotations\nimport torch\n"
|
|
||||||
+ self.split_gm.print_readable(print_output=False)
|
|
||||||
)
|
|
||||||
src = src.replace("<lambda>", "GraphModule")
|
|
||||||
with open(graph_path, "w") as f:
|
|
||||||
f.write(src)
|
|
||||||
|
|
||||||
rank0_log(f"Computation graph saved to {graph_path}")
|
if rank == 0:
|
||||||
|
graph_path = os.path.join(
|
||||||
|
local_cache_dir, f"computation_graph_{time.time()}.py"
|
||||||
|
)
|
||||||
|
if not os.path.exists(graph_path):
|
||||||
|
# code adapted from https://github.com/thuml/depyf/blob/dab831108a752d1facc00acdd6d4243891845c37/depyf/explain/patched_lazy_format_graph_code.py#L30 # noqa
|
||||||
|
# use `print_readable` because it can include submodules
|
||||||
|
src = (
|
||||||
|
"from __future__ import annotations\nimport torch\n"
|
||||||
|
+ self.split_gm.print_readable(print_output=False)
|
||||||
|
)
|
||||||
|
src = src.replace("<lambda>", "GraphModule")
|
||||||
|
with open(graph_path, "w") as f:
|
||||||
|
f.write(src)
|
||||||
|
|
||||||
|
rank0_log(f"Computation graph saved to {graph_path}")
|
||||||
|
|
||||||
self._called = True
|
self._called = True
|
||||||
return self.split_gm
|
return self.split_gm
|
||||||
|
|||||||
@@ -297,6 +297,7 @@ class ModelRunner:
|
|||||||
self.req_to_token_pool = req_to_token_pool
|
self.req_to_token_pool = req_to_token_pool
|
||||||
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
||||||
self.is_hybrid = model_config.is_hybrid
|
self.is_hybrid = model_config.is_hybrid
|
||||||
|
self.is_hybrid_swa = self.is_hybrid
|
||||||
self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA
|
self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA
|
||||||
self.attention_chunk_size = model_config.attention_chunk_size
|
self.attention_chunk_size = model_config.attention_chunk_size
|
||||||
self.forward_pass_id = 0
|
self.forward_pass_id = 0
|
||||||
@@ -1615,11 +1616,6 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def can_run_piecewise_cuda_graph(self):
|
def can_run_piecewise_cuda_graph(self):
|
||||||
if self.server_args.disable_cuda_graph:
|
|
||||||
log_info_on_rank0(
|
|
||||||
logger, "Disable piecewise CUDA graph because disable_cuda_graph is set"
|
|
||||||
)
|
|
||||||
return False
|
|
||||||
if self.server_args.enable_torch_compile:
|
if self.server_args.enable_torch_compile:
|
||||||
log_info_on_rank0(
|
log_info_on_rank0(
|
||||||
logger,
|
logger,
|
||||||
|
|||||||
@@ -201,8 +201,10 @@ class PiecewiseCudaGraphRunner:
|
|||||||
self.out_cache_loc = torch.zeros(
|
self.out_cache_loc = torch.zeros(
|
||||||
(self.max_num_tokens,), dtype=self._cache_loc_dtype()
|
(self.max_num_tokens,), dtype=self._cache_loc_dtype()
|
||||||
)
|
)
|
||||||
self.out_cache_loc_swa = torch.zeros(
|
self.out_cache_loc_swa = (
|
||||||
(self.max_num_tokens,), dtype=self._cache_loc_dtype()
|
torch.zeros((self.max_num_tokens,), dtype=torch.int64)
|
||||||
|
if model_runner.is_hybrid_swa
|
||||||
|
else None
|
||||||
)
|
)
|
||||||
self.positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64)
|
self.positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64)
|
||||||
self.mrope_positions = torch.zeros(
|
self.mrope_positions = torch.zeros(
|
||||||
@@ -245,6 +247,13 @@ class PiecewiseCudaGraphRunner:
|
|||||||
def warmup_torch_compile(self):
|
def warmup_torch_compile(self):
|
||||||
"""Warmup the model with a simple forward pass before CUDA graph capture."""
|
"""Warmup the model with a simple forward pass before CUDA graph capture."""
|
||||||
num_tokens = 2
|
num_tokens = 2
|
||||||
|
|
||||||
|
out_cache_loc = self.out_cache_loc[:num_tokens]
|
||||||
|
out_cache_loc_swa = (
|
||||||
|
self.out_cache_loc_swa[:num_tokens]
|
||||||
|
if self.out_cache_loc_swa is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
with torch.device(self.device):
|
with torch.device(self.device):
|
||||||
forward_batch = ForwardBatch(
|
forward_batch = ForwardBatch(
|
||||||
forward_mode=ForwardMode.EXTEND,
|
forward_mode=ForwardMode.EXTEND,
|
||||||
@@ -272,12 +281,8 @@ class PiecewiseCudaGraphRunner:
|
|||||||
req_to_token_pool=self.model_runner.req_to_token_pool,
|
req_to_token_pool=self.model_runner.req_to_token_pool,
|
||||||
token_to_kv_pool=self.model_runner.token_to_kv_pool,
|
token_to_kv_pool=self.model_runner.token_to_kv_pool,
|
||||||
attn_backend=self.model_runner.attn_backend,
|
attn_backend=self.model_runner.attn_backend,
|
||||||
out_cache_loc=torch.zeros(
|
out_cache_loc=out_cache_loc,
|
||||||
(num_tokens,), device=self.device, dtype=self._cache_loc_dtype()
|
out_cache_loc_swa=out_cache_loc_swa,
|
||||||
),
|
|
||||||
out_cache_loc_swa=torch.zeros(
|
|
||||||
(num_tokens,), device=self.device, dtype=self._cache_loc_dtype()
|
|
||||||
),
|
|
||||||
seq_lens_sum=num_tokens,
|
seq_lens_sum=num_tokens,
|
||||||
encoder_lens=None,
|
encoder_lens=None,
|
||||||
return_logprob=False,
|
return_logprob=False,
|
||||||
@@ -304,7 +309,6 @@ class PiecewiseCudaGraphRunner:
|
|||||||
|
|
||||||
# Attention backend
|
# Attention backend
|
||||||
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
|
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
|
||||||
|
|
||||||
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):
|
), disable_ca_comm(self.model_runner.tp_group):
|
||||||
@@ -378,7 +382,11 @@ class PiecewiseCudaGraphRunner:
|
|||||||
input_embeds = None
|
input_embeds = None
|
||||||
|
|
||||||
out_cache_loc = self.out_cache_loc[:num_tokens]
|
out_cache_loc = self.out_cache_loc[:num_tokens]
|
||||||
out_cache_loc_swa = self.out_cache_loc_swa[:num_tokens]
|
out_cache_loc_swa = (
|
||||||
|
self.out_cache_loc_swa[:num_tokens]
|
||||||
|
if self.out_cache_loc_swa is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
positions = self.positions[:num_tokens]
|
positions = self.positions[:num_tokens]
|
||||||
mrope_positions = self.mrope_positions[:, :num_tokens]
|
mrope_positions = self.mrope_positions[:, :num_tokens]
|
||||||
|
|
||||||
@@ -490,7 +498,8 @@ class PiecewiseCudaGraphRunner:
|
|||||||
self.raw_num_tokens = num_tokens
|
self.raw_num_tokens = num_tokens
|
||||||
if static_num_tokens != num_tokens:
|
if static_num_tokens != num_tokens:
|
||||||
self.out_cache_loc.zero_()
|
self.out_cache_loc.zero_()
|
||||||
self.out_cache_loc_swa.zero_()
|
if self.out_cache_loc_swa is not None:
|
||||||
|
self.out_cache_loc_swa.zero_()
|
||||||
bs = forward_batch.batch_size
|
bs = forward_batch.batch_size
|
||||||
|
|
||||||
if self.use_input_embeds:
|
if self.use_input_embeds:
|
||||||
@@ -500,8 +509,12 @@ class PiecewiseCudaGraphRunner:
|
|||||||
|
|
||||||
self.positions[:num_tokens].copy_(forward_batch.positions)
|
self.positions[:num_tokens].copy_(forward_batch.positions)
|
||||||
self.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
|
self.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
|
||||||
if forward_batch.out_cache_loc_swa is not None:
|
if self.out_cache_loc_swa is not None:
|
||||||
self.out_cache_loc_swa[:num_tokens].copy_(forward_batch.out_cache_loc_swa)
|
self.out_cache_loc_swa[: self.raw_num_tokens].copy_(
|
||||||
|
self.model_runner.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
|
||||||
|
forward_batch.out_cache_loc
|
||||||
|
)
|
||||||
|
)
|
||||||
input_ids = self.input_ids[:static_num_tokens]
|
input_ids = self.input_ids[:static_num_tokens]
|
||||||
positions = self.positions[:static_num_tokens]
|
positions = self.positions[:static_num_tokens]
|
||||||
out_cache_loc = self.out_cache_loc[:static_num_tokens]
|
out_cache_loc = self.out_cache_loc[:static_num_tokens]
|
||||||
@@ -522,9 +535,6 @@ class PiecewiseCudaGraphRunner:
|
|||||||
input_ids = self.input_ids[:static_num_tokens]
|
input_ids = self.input_ids[:static_num_tokens]
|
||||||
input_embeds = None
|
input_embeds = None
|
||||||
|
|
||||||
positions = self.positions[:static_num_tokens]
|
|
||||||
out_cache_loc = self.out_cache_loc[:static_num_tokens]
|
|
||||||
|
|
||||||
mrope_positions = (
|
mrope_positions = (
|
||||||
self.mrope_positions[:, :static_num_tokens]
|
self.mrope_positions[:, :static_num_tokens]
|
||||||
if forward_batch.mrope_positions is not None
|
if forward_batch.mrope_positions is not None
|
||||||
|
|||||||
Reference in New Issue
Block a user