[3/N][CP] Implement zigzag CP strategy (#28421)
This commit is contained in:
@@ -123,6 +123,13 @@ from sglang.srt.layers.attention.attention_registry import (
|
||||
)
|
||||
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
|
||||
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
||||
from sglang.srt.layers.cp.utils import (
|
||||
cp_gather_after_forward,
|
||||
cp_split_before_forward,
|
||||
get_cp_strategy,
|
||||
is_cp_v2_active,
|
||||
prepare_cp_forward,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
DpPaddingMode,
|
||||
get_attention_tp_group,
|
||||
@@ -3411,6 +3418,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
self.prefill_cuda_graph_runner is not None
|
||||
and self.prefill_cuda_graph_runner.can_run(forward_batch)
|
||||
)
|
||||
if get_cp_strategy() is not None:
|
||||
can_run_graph = False
|
||||
if can_run_graph:
|
||||
# TODO: device_timer.wrap is too broad here — it also includes
|
||||
# replay_prepare time. Move timing into the prefill cuda graph
|
||||
@@ -3434,6 +3443,21 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
# e.g. Moss-VL's prefill cross-attention custom mask.
|
||||
self.model.prepare_forward_batch(forward_batch)
|
||||
self.attn_backend.init_forward_metadata(forward_batch)
|
||||
cp_v2_active = is_cp_v2_active(forward_batch)
|
||||
forward_positions = forward_batch.positions
|
||||
if cp_v2_active:
|
||||
prepare_cp_forward(forward_batch)
|
||||
complete_hidden_states = kwargs.get("input_embeds")
|
||||
if complete_hidden_states is None:
|
||||
embed_layer = self.model.get_input_embeddings()
|
||||
complete_hidden_states = embed_layer(forward_batch.input_ids)
|
||||
sharded_hidden_states, sharded_positions = cp_split_before_forward(
|
||||
complete_hidden_states,
|
||||
forward_batch.positions,
|
||||
forward_batch,
|
||||
)
|
||||
kwargs["input_embeds"] = sharded_hidden_states
|
||||
forward_positions = sharded_positions
|
||||
|
||||
ctx = (
|
||||
self.device_timer.wrap(metadata={"category": "extend"})
|
||||
@@ -3441,7 +3465,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
else contextlib.nullcontext()
|
||||
)
|
||||
with ctx:
|
||||
if _is_hip and self.prefill_cuda_graph_runner is not None:
|
||||
if (
|
||||
_is_hip
|
||||
and self.prefill_cuda_graph_runner is not None
|
||||
and not cp_v2_active
|
||||
):
|
||||
# AMD/HIP: when PCG is enabled but the batch exceeds max captured
|
||||
# size, run eagerly under enable_tc_piecewise_cuda_graph() and
|
||||
# set_tc_piecewise_forward_context() so that (a) Dynamo guards on
|
||||
@@ -3461,14 +3489,47 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
):
|
||||
ret = self.model.forward(
|
||||
forward_batch.input_ids,
|
||||
forward_batch.positions,
|
||||
forward_positions,
|
||||
forward_batch,
|
||||
**kwargs,
|
||||
)
|
||||
elif cp_v2_active:
|
||||
hidden_states = self.model.model(
|
||||
forward_batch.input_ids,
|
||||
forward_positions,
|
||||
forward_batch,
|
||||
input_embeds=kwargs.get("input_embeds"),
|
||||
pp_proxy_tensors=kwargs.get("pp_proxy_tensors"),
|
||||
)
|
||||
|
||||
aux_hidden_states = None
|
||||
capture_aux_hidden_states = getattr(
|
||||
self.model, "capture_aux_hidden_states", False
|
||||
)
|
||||
if capture_aux_hidden_states:
|
||||
hidden_states, aux_hidden_states = hidden_states
|
||||
|
||||
if self.model.pp_group.is_last_rank:
|
||||
hidden_states = cp_gather_after_forward(
|
||||
hidden_states,
|
||||
forward_batch,
|
||||
torch.cuda.current_stream(),
|
||||
)
|
||||
ret = self.model.logits_processor(
|
||||
forward_batch.input_ids,
|
||||
hidden_states,
|
||||
self.model.lm_head,
|
||||
forward_batch,
|
||||
aux_hidden_states,
|
||||
)
|
||||
elif capture_aux_hidden_states:
|
||||
ret = hidden_states, aux_hidden_states
|
||||
else:
|
||||
ret = hidden_states
|
||||
else:
|
||||
ret = self.model.forward(
|
||||
forward_batch.input_ids,
|
||||
forward_batch.positions,
|
||||
forward_positions,
|
||||
forward_batch,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user