forked from Karylab-cklius/vllm
[CI Bugfix] Lazily import Qwen warmup dependencies (#47539)
Signed-off-by: LopezCastroRoberto <rocastro@redhat.com> Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
parent
978de83353
commit
d7192cfccf
@@ -8,19 +8,6 @@ from typing import TYPE_CHECKING
|
||||
import torch
|
||||
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.fla.ops.fused_gdn_prefill_post_conv import (
|
||||
fused_post_conv_prep,
|
||||
)
|
||||
from vllm.model_executor.layers.fla.ops.fused_sigmoid_gating import (
|
||||
fused_sigmoid_gating_delta_rule_update,
|
||||
)
|
||||
from vllm.model_executor.layers.mamba.mamba_utils import is_conv_state_dim_first
|
||||
from vllm.model_executor.layers.mamba.ops.causal_conv1d import (
|
||||
causal_conv1d_fn,
|
||||
)
|
||||
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID, PAD_SLOT_ID
|
||||
from vllm.v1.worker.block_table import BlockTable
|
||||
from vllm.v1.worker.utils import _zero_kv_blocks_kernel
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.v1.worker.gpu_model_runner import GPUModelRunner
|
||||
@@ -128,6 +115,10 @@ def _qwen_gdn_warmup_config(
|
||||
continue
|
||||
|
||||
conv_cache, ssm_state = cache_tensors
|
||||
from vllm.model_executor.layers.mamba.mamba_utils import (
|
||||
is_conv_state_dim_first,
|
||||
)
|
||||
|
||||
conv_state = (
|
||||
conv_cache if is_conv_state_dim_first() else conv_cache.transpose(-1, -2)
|
||||
)
|
||||
@@ -191,6 +182,8 @@ def _warm_zero_kv_blocks_with_runner_zeroer(runner: object) -> bool:
|
||||
def _warm_zero_kv_blocks_kernel(
|
||||
device: torch.device, config: _ZeroKvWarmupConfig
|
||||
) -> None:
|
||||
from vllm.v1.worker.utils import _zero_kv_blocks_kernel
|
||||
|
||||
max_n_blocks = max(_ZERO_KV_N_BLOCKS)
|
||||
scratch = torch.empty(
|
||||
max_n_blocks * config.page_size_el,
|
||||
@@ -219,6 +212,8 @@ def _warm_zero_kv_blocks_kernel(
|
||||
|
||||
|
||||
def _warm_compute_slot_mapping_kernel(device: torch.device) -> None:
|
||||
from vllm.v1.worker.block_table import BlockTable
|
||||
|
||||
# num_tokens/max_num_tokens are do_not_specialize; keep the launch tiny.
|
||||
num_tokens = 1
|
||||
query_start_loc = torch.tensor([0, num_tokens], dtype=torch.int32, device=device)
|
||||
@@ -244,6 +239,11 @@ def _warm_compute_slot_mapping_kernel(device: torch.device) -> None:
|
||||
def _warm_causal_conv1d_fwd_kernel(
|
||||
device: torch.device, config: _QwenGDNWarmupConfig
|
||||
) -> None:
|
||||
from vllm.model_executor.layers.mamba.ops.causal_conv1d import (
|
||||
causal_conv1d_fn,
|
||||
)
|
||||
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID, PAD_SLOT_ID
|
||||
|
||||
x_storage = torch.empty(
|
||||
(1, config.conv_dim), dtype=config.conv_dtype, device=device
|
||||
)
|
||||
@@ -276,6 +276,10 @@ def _warm_causal_conv1d_fwd_kernel(
|
||||
def _warm_fused_post_conv_kernel(
|
||||
device: torch.device, config: _QwenGDNWarmupConfig
|
||||
) -> None:
|
||||
from vllm.model_executor.layers.fla.ops.fused_gdn_prefill_post_conv import (
|
||||
fused_post_conv_prep,
|
||||
)
|
||||
|
||||
qkv_dim = 2 * config.h * config.k + config.hv * config.v
|
||||
for length in _FLA_POST_CONV_WARMUP_LENGTHS:
|
||||
conv_output = torch.empty(
|
||||
@@ -302,6 +306,10 @@ def _warm_fused_sigmoid_gating_delta_rule_update_kernel(
|
||||
device: torch.device,
|
||||
config: _QwenGDNWarmupConfig,
|
||||
) -> None:
|
||||
from vllm.model_executor.layers.fla.ops.fused_sigmoid_gating import (
|
||||
fused_sigmoid_gating_delta_rule_update,
|
||||
)
|
||||
|
||||
q = torch.empty((1, 1, config.h, config.k), dtype=config.conv_dtype, device=device)
|
||||
k = torch.empty_like(q)
|
||||
v = torch.empty((1, 1, config.hv, config.v), dtype=config.conv_dtype, device=device)
|
||||
|
||||
Reference in New Issue
Block a user