llama model and llama eagle3 model support dp-attn (#15268)
This commit is contained in:
@@ -13,6 +13,7 @@ from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner i
|
|||||||
EAGLEDraftNpuGraphRunner,
|
EAGLEDraftNpuGraphRunner,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.triton_backend import TritonMultiStepDraftBackend
|
from sglang.srt.layers.attention.triton_backend import TritonMultiStepDraftBackend
|
||||||
|
from sglang.srt.layers.dp_attention import get_attention_tp_group
|
||||||
from sglang.srt.layers.moe.utils import (
|
from sglang.srt.layers.moe.utils import (
|
||||||
speculative_moe_a2a_backend_context,
|
speculative_moe_a2a_backend_context,
|
||||||
speculative_moe_backend_context,
|
speculative_moe_backend_context,
|
||||||
@@ -118,7 +119,15 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
self.req_to_token_pool, self.token_to_kv_pool_allocator = (
|
self.req_to_token_pool, self.token_to_kv_pool_allocator = (
|
||||||
target_worker.get_memory_pool()
|
target_worker.get_memory_pool()
|
||||||
)
|
)
|
||||||
with empty_context(), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
|
||||||
|
# Init draft worker
|
||||||
|
if server_args.enable_dp_attention and self.speculative_algorithm.is_eagle3():
|
||||||
|
ctx = draft_tp_context(get_attention_tp_group())
|
||||||
|
else:
|
||||||
|
ctx = empty_context()
|
||||||
|
with (
|
||||||
|
ctx
|
||||||
|
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
||||||
# Init draft worker
|
# Init draft worker
|
||||||
self.draft_worker = TpModelWorker(
|
self.draft_worker = TpModelWorker(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
@@ -626,7 +635,9 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
# Draft prefill
|
# Draft prefill
|
||||||
model_worker_batch.capture_hidden_mode = CaptureHiddenMode.LAST
|
model_worker_batch.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||||
with speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
with self.draft_worker.draft_tp_context(
|
||||||
|
self.draft_worker.draft_runner.tp_group
|
||||||
|
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
||||||
batch_output.next_draft_input = (
|
batch_output.next_draft_input = (
|
||||||
self.draft_worker._draft_extend_for_prefill(
|
self.draft_worker._draft_extend_for_prefill(
|
||||||
model_worker_batch,
|
model_worker_batch,
|
||||||
@@ -644,14 +655,18 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
topk=self.topk,
|
topk=self.topk,
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
with speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
with self.draft_worker.draft_tp_context(
|
||||||
|
self.draft_worker.draft_runner.tp_group
|
||||||
|
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
||||||
verify_input: EagleVerifyInput = self.draft_worker.draft(
|
verify_input: EagleVerifyInput = self.draft_worker.draft(
|
||||||
model_worker_batch
|
model_worker_batch
|
||||||
)
|
)
|
||||||
assert verify_input.is_verify_input()
|
assert verify_input.is_verify_input()
|
||||||
model_worker_batch.spec_info = verify_input
|
model_worker_batch.spec_info = verify_input
|
||||||
batch_output = self.verify(model_worker_batch)
|
batch_output = self.verify(model_worker_batch)
|
||||||
with speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
with self.draft_worker.draft_tp_context(
|
||||||
|
self.draft_worker.draft_runner.tp_group
|
||||||
|
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
||||||
self.draft_worker._draft_extend_for_decode(
|
self.draft_worker._draft_extend_for_decode(
|
||||||
model_worker_batch, batch_output
|
model_worker_batch, batch_output
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user