[Auto Sync] Update backend.py, forward_batch_info.py, piece... (20251113) (#13221)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com> Co-authored-by: Stefan He <hebiaobuaa@gmail.com>
This commit is contained in:
co-authored by
github-actions[bot]
Stefan He
parent
b5904999c0
commit
7a2254b260
@@ -24,6 +24,16 @@ from sglang.srt.compilation.pass_manager import PostGradPassManager
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
SPLIT_OPS = [
|
||||||
|
"sglang.unified_attention_with_output",
|
||||||
|
"sglang.inplace_all_reduce",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def add_split_ops(ops):
|
||||||
|
SPLIT_OPS.extend(ops)
|
||||||
|
|
||||||
|
|
||||||
def make_compiler(config: CompilationConfig):
|
def make_compiler(config: CompilationConfig):
|
||||||
if config.compiler == "eager":
|
if config.compiler == "eager":
|
||||||
return EagerAdapter()
|
return EagerAdapter()
|
||||||
@@ -392,9 +402,9 @@ class SGLangBackend:
|
|||||||
self.configure_post_pass()
|
self.configure_post_pass()
|
||||||
|
|
||||||
self.split_gm, self.piecewise_graphs = split_graph(
|
self.split_gm, self.piecewise_graphs = split_graph(
|
||||||
graph, ["sglang.unified_attention_with_output", "sglang.inplace_all_reduce"]
|
graph,
|
||||||
|
SPLIT_OPS,
|
||||||
)
|
)
|
||||||
|
|
||||||
from torch._dynamo.utils import lazy_format_graph_code
|
from torch._dynamo.utils import lazy_format_graph_code
|
||||||
|
|
||||||
# depyf will hook lazy_format_graph_code and dump the graph
|
# depyf will hook lazy_format_graph_code and dump the graph
|
||||||
|
|||||||
@@ -201,6 +201,10 @@ class ForwardBatch:
|
|||||||
# The original sequence length without being chunked. Qwen-1M related.
|
# The original sequence length without being chunked. Qwen-1M related.
|
||||||
orig_seq_lens: Optional[torch.Tensor] = None
|
orig_seq_lens: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
# The indices of output tokens in the token_to_kv_pool_swa
|
||||||
|
# TODO(shiyang, biao): integrate out_cache_loc_swa into multiple attention backends
|
||||||
|
out_cache_loc_swa: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# Optional seq_lens on cpu
|
# Optional seq_lens on cpu
|
||||||
seq_lens_cpu: Optional[torch.Tensor] = None
|
seq_lens_cpu: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
|||||||
@@ -185,6 +185,9 @@ 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.max_num_tokens,), dtype=self._cache_loc_dtype()
|
||||||
|
)
|
||||||
self.positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64)
|
self.positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64)
|
||||||
self.tbo_plugin = TboCudaGraphRunnerPlugin()
|
self.tbo_plugin = TboCudaGraphRunnerPlugin()
|
||||||
|
|
||||||
@@ -238,9 +241,13 @@ class PiecewiseCudaGraphRunner:
|
|||||||
out_cache_loc=torch.zeros(
|
out_cache_loc=torch.zeros(
|
||||||
(num_tokens,), device=self.device, dtype=self._cache_loc_dtype()
|
(num_tokens,), device=self.device, dtype=self._cache_loc_dtype()
|
||||||
),
|
),
|
||||||
|
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,
|
||||||
|
extend_num_tokens=num_tokens,
|
||||||
extend_seq_lens=torch.tensor([num_tokens], device=self.device),
|
extend_seq_lens=torch.tensor([num_tokens], device=self.device),
|
||||||
extend_prefix_lens=torch.tensor([num_tokens], device=self.device),
|
extend_prefix_lens=torch.tensor([num_tokens], device=self.device),
|
||||||
extend_start_loc=torch.tensor([0], device=self.device),
|
extend_start_loc=torch.tensor([0], device=self.device),
|
||||||
@@ -334,6 +341,7 @@ class PiecewiseCudaGraphRunner:
|
|||||||
# Graph inputs
|
# Graph inputs
|
||||||
input_ids = self.input_ids[:num_tokens]
|
input_ids = self.input_ids[:num_tokens]
|
||||||
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]
|
||||||
positions = self.positions[:num_tokens]
|
positions = self.positions[:num_tokens]
|
||||||
|
|
||||||
# pipeline parallelism
|
# pipeline parallelism
|
||||||
@@ -365,9 +373,11 @@ class PiecewiseCudaGraphRunner:
|
|||||||
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=out_cache_loc,
|
out_cache_loc=out_cache_loc,
|
||||||
|
out_cache_loc_swa=out_cache_loc_swa,
|
||||||
seq_lens_sum=num_tokens,
|
seq_lens_sum=num_tokens,
|
||||||
encoder_lens=None,
|
encoder_lens=None,
|
||||||
return_logprob=False,
|
return_logprob=False,
|
||||||
|
extend_num_tokens=num_tokens,
|
||||||
extend_seq_lens=torch.tensor([num_tokens], device=self.device),
|
extend_seq_lens=torch.tensor([num_tokens], device=self.device),
|
||||||
extend_prefix_lens=torch.tensor([num_tokens], device=self.device),
|
extend_prefix_lens=torch.tensor([num_tokens], device=self.device),
|
||||||
extend_start_loc=torch.tensor([0], device=self.device),
|
extend_start_loc=torch.tensor([0], device=self.device),
|
||||||
@@ -437,16 +447,24 @@ 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_()
|
||||||
bs = forward_batch.batch_size
|
bs = forward_batch.batch_size
|
||||||
|
|
||||||
self.input_ids[:num_tokens].copy_(forward_batch.input_ids)
|
self.input_ids[:num_tokens].copy_(forward_batch.input_ids)
|
||||||
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:
|
||||||
|
self.out_cache_loc_swa[:num_tokens].copy_(forward_batch.out_cache_loc_swa)
|
||||||
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]
|
||||||
|
|
||||||
|
out_cache_loc_swa = (
|
||||||
|
self.out_cache_loc_swa[:static_num_tokens]
|
||||||
|
if forward_batch.out_cache_loc_swa is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
next_token_logits_buffer = None
|
next_token_logits_buffer = None
|
||||||
mrope_positions = None
|
mrope_positions = None
|
||||||
|
|
||||||
@@ -463,6 +481,7 @@ class PiecewiseCudaGraphRunner:
|
|||||||
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=out_cache_loc,
|
out_cache_loc=out_cache_loc,
|
||||||
|
out_cache_loc_swa=out_cache_loc_swa,
|
||||||
seq_lens_sum=forward_batch.seq_lens_sum,
|
seq_lens_sum=forward_batch.seq_lens_sum,
|
||||||
encoder_lens=forward_batch.encoder_lens,
|
encoder_lens=forward_batch.encoder_lens,
|
||||||
return_logprob=False,
|
return_logprob=False,
|
||||||
|
|||||||
Reference in New Issue
Block a user