Support DCP for Kimi Linear model (#32612)
Co-authored-by: Julien Lin <jullin@nvidia.com> Co-authored-by: kpham-sgl <khoa.pham@radixark.ai>
This commit is contained in:
co-authored by
Julien Lin
kpham-sgl
parent
c4fc241fd3
commit
ef6c07008b
@@ -50,7 +50,11 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
||||
set_tc_piecewise_forward_context,
|
||||
)
|
||||
from sglang.srt.utils import is_hip
|
||||
from sglang.srt.utils.common import ceil_align, require_mlp_sync
|
||||
from sglang.srt.utils.common import (
|
||||
ceil_align,
|
||||
get_eager_max_batch_size,
|
||||
require_mlp_sync,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -101,14 +105,12 @@ class EagerRunner(BaseRunner):
|
||||
# (expand_for_topk_draft) before the eager fallback.
|
||||
max_bs *= sa.speculative_eagle_topk
|
||||
# Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies.
|
||||
if require_mlp_sync(sa):
|
||||
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
||||
|
||||
max_bs = ceil_align(max_bs, self.attn_tp_size)
|
||||
max_bs = ceil_align(max_bs, get_cp_padding_align_size())
|
||||
max_bs = get_eager_max_batch_size(sa, max_bs)
|
||||
prefill_ceiling = max(mr.max_total_num_tokens, sa.max_prefill_buffer_tokens())
|
||||
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req)
|
||||
if require_mlp_sync(sa):
|
||||
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
||||
|
||||
max_num_token = ceil_align(max_num_token, self.attn_tp_size)
|
||||
max_num_token = ceil_align(max_num_token, get_cp_padding_align_size())
|
||||
self._eager_max_bs = max_bs
|
||||
@@ -261,7 +263,9 @@ class EagerRunner(BaseRunner):
|
||||
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
||||
|
||||
if forward_batch.needs_forward_metadata_init():
|
||||
if hasattr(model_runner.model, "prepare_context_parallel_metadata_for_dcp"):
|
||||
if model_runner.dcp_size > 1 and hasattr(
|
||||
model_runner.model, "prepare_context_parallel_metadata_for_dcp"
|
||||
):
|
||||
# prepare kv cache buffer for dcp to gather kv cache
|
||||
forward_batch.attn_dcp_metadata = (
|
||||
model_runner.model.prepare_context_parallel_metadata_for_dcp(
|
||||
|
||||
Reference in New Issue
Block a user