config: the draft runner carries its own attention backend
`build_draft_tp_worker` built a `ServerArgs` variant whose only job was to make four config reads answer with the draft's backend instead of the target's, and published it for the duration of the build so the bags agreed. The backend is a per-runner fact — target and draft coexist in one process — so it moves onto the runner, and the variant and the construction-time publish both go away. `ModelRunner` takes `draft_attention_backend` and resolves the runner's effective value once (`resolve_draft_attention_backend`: the algorithm's resolved backend, else `--speculative-draft-attention-backend`, else None for a target runner); `TpModelWorker` threads it to both runner constructions. `resolve_attention_backend_strs` reads it off the runner, and `ModelRunner` stamps the resolved pair *before* building backends so a backend can read it while it constructs — which is what the FlashInfer KV-access check needs now that it no longer asks the config. `configure_kv_cache_dtype` and the draft backend factory read the runner too. One latent bug falls out: the non-hybrid branch of the backend build ignored the resolved pair and re-read `server_args.attention_backend`, which is why the variant had to set that field as well as the split pair. It now uses the value that was resolved for the runner. `draft_server_args_overrides` and the `preserve_config()` publish switch are deleted; with them goes the last production `ServerArgs.derive` outside pre-publish config building, and the last construction-time publish. The chunked-prefix gate the target resolved simply stays in the bags, since nothing re-projects them.
This commit is contained in:
@@ -398,23 +398,16 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
|
||||
allowed = {"triton", "trtllm_mha", "flashinfer"}
|
||||
else:
|
||||
allowed = {"triton", "trtllm_mha", "fa4"}
|
||||
attn_be = runner.server_args.attention_backend
|
||||
prefill_be = runner.server_args.prefill_attention_backend
|
||||
decode_be = runner.server_args.decode_attention_backend
|
||||
# When using split prefill/decode backends, check each individually
|
||||
if prefill_be and decode_be:
|
||||
assert prefill_be in allowed and decode_be in allowed, (
|
||||
f"Only {allowed} backends are supported on Blackwell GPUs for hybrid GDN models. "
|
||||
f"Got prefill={prefill_be}, decode={decode_be}."
|
||||
)
|
||||
else:
|
||||
assert attn_be in allowed, (
|
||||
f"Only {allowed} backends are supported on Blackwell GPUs for hybrid GDN models. "
|
||||
f"Got attention_backend={attn_be}."
|
||||
)
|
||||
prefill_be = runner.prefill_attention_backend_str
|
||||
decode_be = runner.decode_attention_backend_str
|
||||
assert prefill_be in allowed and decode_be in allowed, (
|
||||
f"Only {allowed} backends are supported on Blackwell GPUs for hybrid GDN models. "
|
||||
f"Got prefill={prefill_be}, decode={decode_be}."
|
||||
)
|
||||
elif is_npu():
|
||||
assert (
|
||||
runner.server_args.attention_backend == "ascend"
|
||||
runner.prefill_attention_backend_str == "ascend"
|
||||
and runner.decode_attention_backend_str == "ascend"
|
||||
), "ascend backend is the only supported backend on NPU for hybrid GDN models, use --attention-backend ascend to specify the backend."
|
||||
logger.info(f"Using hybrid linear attention backend for hybrid GDN models.")
|
||||
linear_attn_backend = GDNAttnBackend(runner)
|
||||
|
||||
@@ -319,9 +319,8 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
self.decode_kv_access = self.kv_cache_quant_method.resolve_attention_access(
|
||||
"decode", "flashinfer"
|
||||
)
|
||||
prefill_backend, decode_backend = (
|
||||
model_runner.server_args.get_attention_backends()
|
||||
)
|
||||
prefill_backend = model_runner.prefill_attention_backend_str
|
||||
decode_backend = model_runner.decode_attention_backend_str
|
||||
if self.__class__ is FlashInferAttnBackend:
|
||||
if prefill_backend == "flashinfer":
|
||||
self._check_kv_attention_access("prefill", self.prefill_kv_access)
|
||||
|
||||
@@ -310,6 +310,7 @@ class TpModelWorker(BaseTpWorker):
|
||||
memory_pool_config: Optional[MemoryPoolConfig] = None,
|
||||
is_multi_layer_eagle: bool = False,
|
||||
context_length: Optional[int] = None,
|
||||
draft_attention_backend: Optional[str] = None,
|
||||
):
|
||||
# Parse args
|
||||
self.server_args = server_args
|
||||
@@ -325,6 +326,8 @@ class TpModelWorker(BaseTpWorker):
|
||||
# Draft worker: target's effective context length; the draft runs at
|
||||
# absolute target positions. None keeps server_args.context_length.
|
||||
self.context_length = context_length
|
||||
# Draft worker: the attention backend the algorithm resolved for it.
|
||||
self.draft_attention_backend = draft_attention_backend
|
||||
|
||||
# MTP model runners
|
||||
self.model_runner_list: List[ModelRunner] = []
|
||||
@@ -459,6 +462,7 @@ class TpModelWorker(BaseTpWorker):
|
||||
req_to_token_pool=self.req_to_token_pool,
|
||||
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||
memory_pool_config=self.memory_pool_config,
|
||||
draft_attention_backend=self.draft_attention_backend,
|
||||
draft_model_idx=0 if self.is_multi_layer_eagle else None,
|
||||
)
|
||||
|
||||
@@ -479,6 +483,7 @@ class TpModelWorker(BaseTpWorker):
|
||||
req_to_token_pool=self.req_to_token_pool,
|
||||
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||
memory_pool_config=self.memory_pool_config,
|
||||
draft_attention_backend=self.draft_attention_backend,
|
||||
draft_model_idx=i,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -924,7 +924,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
ret.extend_prefix_lens = extend_prefix_lens
|
||||
ret.extend_num_tokens = batch.extend_num_tokens
|
||||
positions, ret.extend_start_loc = compute_position(
|
||||
model_runner.server_args.attention_backend,
|
||||
model_runner.prefill_attention_backend_str,
|
||||
ret.extend_prefix_lens,
|
||||
ret.extend_seq_lens,
|
||||
ret.extend_num_tokens,
|
||||
|
||||
@@ -112,6 +112,7 @@ from sglang.srt.model_executor.model_runner_components.attention_backend_setup i
|
||||
build_attention_backends,
|
||||
configure_aux_hidden_state_capture,
|
||||
get_attention_backend,
|
||||
resolve_attention_backend_strs,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import (
|
||||
capture_cuda_graphs,
|
||||
@@ -261,6 +262,24 @@ class ModelRunnerOutput:
|
||||
indexer_topk_output: Optional[TopkCaptureOutput] = None
|
||||
|
||||
|
||||
def resolve_draft_attention_backend(
|
||||
*,
|
||||
draft_attention_backend: Optional[str],
|
||||
server_args: ServerArgs,
|
||||
is_draft_worker: bool,
|
||||
) -> Optional[str]:
|
||||
"""The attention backend a runner uses because it is a draft runner.
|
||||
|
||||
``None`` for a target runner. For a draft: the backend the algorithm that
|
||||
built it resolved (the supported-backend fallback in
|
||||
``build_draft_tp_worker``), else ``--speculative-draft-attention-backend``.
|
||||
It belongs to the runner, not the process: target and draft coexist.
|
||||
"""
|
||||
if not is_draft_worker:
|
||||
return None
|
||||
return draft_attention_backend or server_args.speculative_draft_attention_backend
|
||||
|
||||
|
||||
class ModelRunner:
|
||||
"""ModelRunner runs the forward passes of the models."""
|
||||
|
||||
@@ -277,6 +296,7 @@ class ModelRunner:
|
||||
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
|
||||
memory_pool_config: Optional[MemoryPoolConfig] = None,
|
||||
draft_model_idx: Optional[int] = None,
|
||||
draft_attention_backend: Optional[str] = None,
|
||||
):
|
||||
# Parse args
|
||||
self.mem_fraction_static = mem_fraction_static
|
||||
@@ -293,6 +313,11 @@ class ModelRunner:
|
||||
self.dist_port = nccl_port
|
||||
self.server_args = server_args
|
||||
self.is_draft_worker = is_draft_worker
|
||||
self.draft_attention_backend = resolve_draft_attention_backend(
|
||||
draft_attention_backend=draft_attention_backend,
|
||||
server_args=server_args,
|
||||
is_draft_worker=is_draft_worker,
|
||||
)
|
||||
# This runner's own load format, resolved before anything keys off it:
|
||||
# the remote-instance transfer engine is initialized at the top of
|
||||
# initialize(), long before the weights are loaded.
|
||||
@@ -902,12 +927,15 @@ class ModelRunner:
|
||||
dflash_target_layer_ids=self.spec_aux_config.dflash_target_layer_ids,
|
||||
is_dspark=self.spec_algorithm.is_dspark(),
|
||||
)
|
||||
# Resolve before building: backends read the pair off the runner while
|
||||
# they construct (the FlashInfer KV-access check).
|
||||
resolved = resolve_attention_backend_strs(model_runner=self)
|
||||
self.prefill_attention_backend_str = resolved.prefill
|
||||
self.decode_attention_backend_str = resolved.decode
|
||||
backends = build_attention_backends(model_runner=self)
|
||||
self.attn_backend = backends.attn_backend
|
||||
self.decode_attn_backend = backends.decode_attn_backend
|
||||
self.decode_attn_backend_group = backends.decode_attn_backend_group
|
||||
self.prefill_attention_backend_str = backends.prefill_attention_backend_str
|
||||
self.decode_attention_backend_str = backends.decode_attention_backend_str
|
||||
|
||||
if self.server_args.dcp_size > 1 and get_parallel().dcp_replicate_q_proj:
|
||||
self._prepare_replicated_q_proj()
|
||||
@@ -1061,7 +1089,9 @@ class ModelRunner:
|
||||
)
|
||||
|
||||
maybe_trigger_remote_instance_nccl_send_group(
|
||||
server_args=self.server_args, tp_rank=self.ps.tp_rank
|
||||
server_args=self.server_args,
|
||||
tp_rank=self.ps.tp_rank,
|
||||
load_format=draft_load_format,
|
||||
)
|
||||
|
||||
with self._load_format_scope(draft_load_format):
|
||||
@@ -1248,9 +1278,7 @@ class ModelRunner:
|
||||
if spec_algorithm is not None
|
||||
else False
|
||||
),
|
||||
speculative_draft_attention_backend=getattr(
|
||||
self.server_args, "speculative_draft_attention_backend", None
|
||||
),
|
||||
speculative_draft_attention_backend=self.draft_attention_backend,
|
||||
)
|
||||
)
|
||||
# This runner's OWN resolved dtype string (target or draft). Attention
|
||||
|
||||
+20
-11
@@ -17,7 +17,6 @@ from sglang.srt.utils import init_cublas
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -73,8 +72,13 @@ def build_attention_backends(*, model_runner: ModelRunner) -> AttentionBackends:
|
||||
if model_runner.device in ("cuda", "musa"):
|
||||
init_cublas()
|
||||
|
||||
resolved = _resolve_attention_backend_strs(
|
||||
server_args=server_args, is_draft_worker=model_runner.is_draft_worker
|
||||
# Already resolved and stamped on the runner before this call.
|
||||
resolved = ResolvedAttentionBackendStr(
|
||||
prefill=model_runner.prefill_attention_backend_str,
|
||||
decode=model_runner.decode_attention_backend_str,
|
||||
is_draft_override=bool(
|
||||
model_runner.is_draft_worker and model_runner.draft_attention_backend
|
||||
),
|
||||
)
|
||||
|
||||
if server_args.enable_pdmux:
|
||||
@@ -140,10 +144,7 @@ def get_attention_backend(
|
||||
*, model_runner: ModelRunner, init_new_workspace: bool = False
|
||||
) -> AttentionBackend:
|
||||
"""Init attention kernel backend."""
|
||||
resolved = _resolve_attention_backend_strs(
|
||||
server_args=model_runner.server_args,
|
||||
is_draft_worker=model_runner.is_draft_worker,
|
||||
)
|
||||
resolved = resolve_attention_backend_strs(model_runner=model_runner)
|
||||
return _build_resolved_backend(
|
||||
model_runner=model_runner,
|
||||
resolved=resolved,
|
||||
@@ -151,10 +152,18 @@ def get_attention_backend(
|
||||
)
|
||||
|
||||
|
||||
def _resolve_attention_backend_strs(
|
||||
*, server_args: ServerArgs, is_draft_worker: bool
|
||||
def resolve_attention_backend_strs(
|
||||
*, model_runner: ModelRunner
|
||||
) -> ResolvedAttentionBackendStr:
|
||||
draft_attn_backend = server_args.speculative_draft_attention_backend
|
||||
"""The (prefill, decode) backends this runner runs.
|
||||
|
||||
A draft runner's backend is its own (``ModelRunner.draft_attention_backend``):
|
||||
target and draft coexist in one process, so it cannot come from the
|
||||
process-wide config.
|
||||
"""
|
||||
server_args = model_runner.server_args
|
||||
is_draft_worker = model_runner.is_draft_worker
|
||||
draft_attn_backend = model_runner.draft_attention_backend
|
||||
if is_draft_worker and draft_attn_backend:
|
||||
logger.warning(f"Overriding draft attention backend to {draft_attn_backend}.")
|
||||
# Single backend for all draft modes (no prefill/decode split).
|
||||
@@ -218,7 +227,7 @@ def _build_resolved_backend(
|
||||
else:
|
||||
attn_backend = _build_backend_from_str(
|
||||
model_runner=model_runner,
|
||||
backend_str=model_runner.server_args.attention_backend,
|
||||
backend_str=resolved.prefill,
|
||||
init_new_workspace=init_new_workspace,
|
||||
)
|
||||
return attn_backend
|
||||
|
||||
@@ -1425,7 +1425,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
"DFLASH prefill expected out_cache_loc, but got None."
|
||||
)
|
||||
positions, _ = compute_position(
|
||||
self.model_runner.server_args.attention_backend,
|
||||
self.model_runner.prefill_attention_backend_str,
|
||||
draft_seq_lens,
|
||||
ctx_lens,
|
||||
int(sum(batch.extend_lens)),
|
||||
|
||||
@@ -39,7 +39,8 @@ class DraftBackendFactory:
|
||||
self.topk = topk
|
||||
self.speculative_num_steps = speculative_num_steps
|
||||
self.seed_dsa_topk_from_draft_extend = seed_dsa_topk_from_draft_extend
|
||||
self.draft_attn_backend = server_args.speculative_draft_attention_backend
|
||||
# The draft runner's own backend, not the process-wide config.
|
||||
self.draft_attn_backend = draft_model_runner.draft_attention_backend
|
||||
|
||||
def _create_backend(
|
||||
self, backend_name: str, backend_map: dict, error_template: str
|
||||
|
||||
@@ -9,7 +9,6 @@ import torch
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
|
||||
from sglang.srt.runtime_context import get_context, get_schedule
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
|
||||
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
||||
@@ -63,26 +62,6 @@ def _resolve_draft_attention_backend_fallback(
|
||||
return draft_backend
|
||||
|
||||
|
||||
def draft_server_args_overrides(draft_backend) -> dict:
|
||||
"""The fields a draft variant must carry: its attention backend.
|
||||
|
||||
Backend selection reads them off the config object the draft runner holds --
|
||||
``speculative_draft_attention_backend`` in ``_resolve_attention_backend_strs``
|
||||
and ``configure_kv_cache_dtype``, ``attention_backend`` in the non-hybrid
|
||||
branch of the backend build, and the split pair must not shadow either with
|
||||
the target's. ``disable_chunked_prefix_cache`` is the target's resolved gate,
|
||||
which lives in the bags only: publishing the variant re-projects the bags
|
||||
from it, so the value has to travel on the variant.
|
||||
"""
|
||||
return dict(
|
||||
speculative_draft_attention_backend=draft_backend,
|
||||
prefill_attention_backend=None,
|
||||
decode_attention_backend=None,
|
||||
attention_backend=draft_backend,
|
||||
disable_chunked_prefix_cache=get_schedule().disable_chunked_prefix_cache,
|
||||
)
|
||||
|
||||
|
||||
def build_draft_tp_worker(
|
||||
*,
|
||||
server_args: ServerArgs,
|
||||
@@ -101,23 +80,17 @@ def build_draft_tp_worker(
|
||||
server_args=server_args, algo_label=algo_label
|
||||
)
|
||||
)
|
||||
draft_server_args = server_args.derive(
|
||||
"draft_worker.build", **draft_server_args_overrides(draft_backend)
|
||||
draft_worker = TpModelWorker(
|
||||
server_args=server_args,
|
||||
gpu_id=gpu_id,
|
||||
ps=ps,
|
||||
nccl_port=nccl_port,
|
||||
is_draft_worker=True,
|
||||
# The draft runs at absolute target positions.
|
||||
context_length=target_model_config.context_len,
|
||||
draft_attention_backend=draft_backend,
|
||||
)
|
||||
|
||||
# The draft's layers must resolve config from the draft's own bags.
|
||||
with get_context().preserve_config():
|
||||
get_context().set_server_args(draft_server_args)
|
||||
draft_worker = TpModelWorker(
|
||||
server_args=draft_server_args,
|
||||
gpu_id=gpu_id,
|
||||
ps=ps,
|
||||
nccl_port=nccl_port,
|
||||
is_draft_worker=True,
|
||||
# The draft runs at absolute target positions.
|
||||
context_length=target_model_config.context_len,
|
||||
)
|
||||
|
||||
draft_model_runner = draft_worker.model_runner
|
||||
draft_worker.draft_runner = draft_model_runner
|
||||
return DraftWorkerBundle(
|
||||
|
||||
@@ -444,7 +444,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
batch.prefix_lens, dtype=torch.int32, device=device
|
||||
)
|
||||
positions, _ = compute_position(
|
||||
self.model_runner.server_args.attention_backend,
|
||||
self.model_runner.prefill_attention_backend_str,
|
||||
draft_seq_lens,
|
||||
ctx_lens,
|
||||
int(sum(batch.extend_lens)),
|
||||
|
||||
@@ -322,6 +322,11 @@ class MockModelRunner(ModelRunner):
|
||||
self.dtype = dtype
|
||||
self.kv_cache_dtype = dtype
|
||||
self.kv_cache_dtype_str = "auto"
|
||||
# This runner's own resolved backends (production stamps these in
|
||||
# ModelRunner.initialize); a draft runner would carry its own.
|
||||
self.prefill_attention_backend_str = case.backend
|
||||
self.decode_attention_backend_str = case.backend
|
||||
self.draft_attention_backend = None
|
||||
self.gpu_id = 0
|
||||
self.canary_manager = None
|
||||
self.page_size = case.page_size
|
||||
|
||||
@@ -293,6 +293,11 @@ class DSAMockModelRunner(ModelRunner):
|
||||
# `set_mla_kv_buffer` does the quantize on the way in.
|
||||
self.kv_cache_dtype = torch.float8_e4m3fn if fp8_kv_cache else dtype
|
||||
self.kv_cache_dtype_str = "auto"
|
||||
# This runner's own resolved backends (production stamps these in
|
||||
# ModelRunner.initialize); a draft runner would carry its own.
|
||||
self.prefill_attention_backend_str = case.backend
|
||||
self.decode_attention_backend_str = case.backend
|
||||
self.draft_attention_backend = None
|
||||
# For TARGET_VERIFY / DRAFT_EXTEND, the DSA backend uses
|
||||
# `self.speculative_num_draft_tokens` to size `seqlens_expanded`
|
||||
# (`dsa_backend.py:482-486,510-515`). When zero, deep_gemm's
|
||||
|
||||
@@ -332,6 +332,11 @@ class MockDSV4ModelRunner:
|
||||
self.dtype = dtype
|
||||
self.kv_cache_dtype = dtype
|
||||
self.kv_cache_dtype_str = "auto"
|
||||
# This runner's own resolved backends (production stamps these in
|
||||
# ModelRunner.initialize); a draft runner would carry its own.
|
||||
self.prefill_attention_backend_str = case.backend
|
||||
self.decode_attention_backend_str = case.backend
|
||||
self.draft_attention_backend = None
|
||||
self.gpu_id = 0
|
||||
self.canary_manager = None
|
||||
self.page_size = case.page_size
|
||||
|
||||
@@ -322,6 +322,11 @@ class DualChunkMockModelRunner(ModelRunner):
|
||||
self.dtype = dtype
|
||||
self.kv_cache_dtype = dtype
|
||||
self.kv_cache_dtype_str = "auto"
|
||||
# This runner's own resolved backends (production stamps these in
|
||||
# ModelRunner.initialize); a draft runner would carry its own.
|
||||
self.prefill_attention_backend_str = case.backend
|
||||
self.decode_attention_backend_str = case.backend
|
||||
self.draft_attention_backend = None
|
||||
self.gpu_id = 0
|
||||
self.canary_manager = None
|
||||
self.page_size = case.page_size
|
||||
|
||||
@@ -217,6 +217,11 @@ class MockGDNModelRunner(ModelRunner):
|
||||
self.dtype = dtype
|
||||
self.kv_cache_dtype = dtype
|
||||
self.kv_cache_dtype_str = "auto"
|
||||
# This runner's own resolved backends (production stamps these in
|
||||
# ModelRunner.initialize); a draft runner would carry its own.
|
||||
self.prefill_attention_backend_str = case.backend
|
||||
self.decode_attention_backend_str = case.backend
|
||||
self.draft_attention_backend = None
|
||||
self.gpu_id = 0
|
||||
self.ps = ParallelState.trivial()
|
||||
self.canary_manager = None
|
||||
|
||||
@@ -222,6 +222,11 @@ class MockKDAModelRunner(ModelRunner):
|
||||
self.dtype = dtype
|
||||
self.kv_cache_dtype = dtype
|
||||
self.kv_cache_dtype_str = "auto"
|
||||
# This runner's own resolved backends (production stamps these in
|
||||
# ModelRunner.initialize); a draft runner would carry its own.
|
||||
self.prefill_attention_backend_str = case.backend
|
||||
self.decode_attention_backend_str = case.backend
|
||||
self.draft_attention_backend = None
|
||||
self.gpu_id = 0
|
||||
self.ps = ParallelState.trivial()
|
||||
self.canary_manager = None
|
||||
|
||||
@@ -230,6 +230,11 @@ class MockLightningModelRunner(ModelRunner):
|
||||
self.dtype = dtype
|
||||
self.kv_cache_dtype = dtype
|
||||
self.kv_cache_dtype_str = "auto"
|
||||
# This runner's own resolved backends (production stamps these in
|
||||
# ModelRunner.initialize); a draft runner would carry its own.
|
||||
self.prefill_attention_backend_str = case.backend
|
||||
self.decode_attention_backend_str = case.backend
|
||||
self.draft_attention_backend = None
|
||||
self.gpu_id = 0
|
||||
self.ps = ParallelState.trivial()
|
||||
self.canary_manager = None
|
||||
|
||||
@@ -315,6 +315,11 @@ class MockMamba2ModelRunner(ModelRunner):
|
||||
self.dtype = dtype
|
||||
self.kv_cache_dtype = dtype
|
||||
self.kv_cache_dtype_str = "auto"
|
||||
# This runner's own resolved backends (production stamps these in
|
||||
# ModelRunner.initialize); a draft runner would carry its own.
|
||||
self.prefill_attention_backend_str = case.backend
|
||||
self.decode_attention_backend_str = case.backend
|
||||
self.draft_attention_backend = None
|
||||
self.gpu_id = 0
|
||||
self.ps = ParallelState.trivial()
|
||||
self.canary_manager = None
|
||||
|
||||
@@ -233,6 +233,11 @@ class MockMLAModelRunner(ModelRunner):
|
||||
# does the BF16->FP8 cast on the way in.
|
||||
self.kv_cache_dtype = torch.float8_e4m3fn if fp8_kv_cache else dtype
|
||||
self.kv_cache_dtype_str = "fp8_e4m3" if fp8_kv_cache else "auto"
|
||||
# This runner's own resolved backends (production stamps these in
|
||||
# ModelRunner.initialize); a draft runner would carry its own.
|
||||
self.prefill_attention_backend_str = case.backend
|
||||
self.decode_attention_backend_str = case.backend
|
||||
self.draft_attention_backend = None
|
||||
self.gpu_id = 0
|
||||
self.canary_manager = None
|
||||
self.page_size = case.page_size
|
||||
|
||||
Reference in New Issue
Block a user