mirror of
https://github.com/wassname/vllm.git
synced 2026-08-20 12:50:59 +08:00
[Model] Composite weight loading for multimodal Qwen2 (#10944)
Signed-off-by: DarkLight1337 <tlleungac@connect.ust.hk>
This commit is contained in:
@@ -101,12 +101,10 @@ def _initialize_model(
|
||||
vllm_config: VllmConfig,
|
||||
*,
|
||||
prefix: str = "",
|
||||
architectures: Optional[list[str]] = None,
|
||||
) -> nn.Module:
|
||||
"""Initialize a model with the given configurations."""
|
||||
model_config = vllm_config.model_config
|
||||
model_class, _ = get_model_architecture(model_config,
|
||||
architectures=architectures)
|
||||
model_class, _ = get_model_architecture(model_config)
|
||||
|
||||
signatures = inspect.signature(model_class.__init__)
|
||||
all_params = [param.name for param in signatures.parameters.values()]
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
"""Utilities for selecting and loading models."""
|
||||
import contextlib
|
||||
from typing import Optional, Tuple, Type
|
||||
from typing import Tuple, Type
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -20,12 +20,8 @@ def set_default_torch_dtype(dtype: torch.dtype):
|
||||
|
||||
|
||||
def get_model_architecture(
|
||||
model_config: ModelConfig,
|
||||
*,
|
||||
architectures: Optional[list[str]] = None,
|
||||
) -> Tuple[Type[nn.Module], str]:
|
||||
if architectures is None:
|
||||
architectures = getattr(model_config.hf_config, "architectures", [])
|
||||
model_config: ModelConfig) -> Tuple[Type[nn.Module], str]:
|
||||
architectures = getattr(model_config.hf_config, "architectures", [])
|
||||
|
||||
# Special handling for quantized Mixtral.
|
||||
# FIXME(woosuk): This is a temporary hack.
|
||||
|
||||
Reference in New Issue
Block a user