Encoder model support for the Transformers backend (#25174)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
This commit is contained in:
@@ -27,7 +27,7 @@ from transformers import (AutoModel, BatchFeature, PretrainedConfig,
|
||||
PreTrainedModel)
|
||||
from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
|
||||
|
||||
from vllm.attention import Attention
|
||||
from vllm.attention import Attention, AttentionType
|
||||
from vllm.compilation.decorators import support_torch_compile
|
||||
from vllm.config import (CacheConfig, DeviceConfig, ModelConfig,
|
||||
ParallelConfig, VllmConfig)
|
||||
@@ -452,8 +452,9 @@ class TransformersBase(nn.Module, SupportsQuant, SupportsLoRA, SupportsPP):
|
||||
self.pp_rank = self.pp_group.rank_in_group
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
|
||||
# To be updated in child classes for use in `load_weights`
|
||||
self.skip_prefixes: Optional[list[str]] = None
|
||||
# Weights to skip in `self.load_weights`
|
||||
self.skip_prefixes: list[str] = []
|
||||
self.skip_substrs: list[str] = []
|
||||
|
||||
# Set correct attn and init on "meta" to delay allocating GPU tensors
|
||||
# TODO: @raushan, use the public `model.set_attn_implementation()`
|
||||
@@ -596,7 +597,10 @@ class TransformersBase(nn.Module, SupportsQuant, SupportsLoRA, SupportsPP):
|
||||
|
||||
_tensor_parallel(self.model)
|
||||
|
||||
def create_attention_instances(self) -> dict[int, Attention]:
|
||||
def create_attention_instances(
|
||||
self,
|
||||
attn_type: AttentionType = AttentionType.DECODER
|
||||
) -> dict[int, Attention]:
|
||||
"""
|
||||
Create `Attention` instances to inform KV cache allocation.
|
||||
"""
|
||||
@@ -625,7 +629,8 @@ class TransformersBase(nn.Module, SupportsQuant, SupportsLoRA, SupportsPP):
|
||||
cache_config=self.cache_config,
|
||||
quant_config=self.quant_config,
|
||||
per_layer_sliding_window=per_layer_sliding_window,
|
||||
prefix=f"{i}.attn")
|
||||
prefix=f"{i}.attn",
|
||||
attn_type=attn_type)
|
||||
return attention_instances
|
||||
|
||||
def init_parameters(self, module: nn.Module):
|
||||
@@ -685,7 +690,11 @@ class TransformersBase(nn.Module, SupportsQuant, SupportsLoRA, SupportsPP):
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str,
|
||||
torch.Tensor]]) -> set[str]:
|
||||
loader = AutoWeightsLoader(self, skip_prefixes=self.skip_prefixes)
|
||||
loader = AutoWeightsLoader(
|
||||
self,
|
||||
skip_prefixes=self.skip_prefixes,
|
||||
skip_substrs=self.skip_substrs,
|
||||
)
|
||||
return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper)
|
||||
|
||||
|
||||
@@ -700,6 +709,37 @@ class TransformersModel(TransformersBase):
|
||||
"model.score": "score",
|
||||
})
|
||||
|
||||
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
|
||||
super().__init__(vllm_config=vllm_config, prefix=prefix)
|
||||
|
||||
# Some encoder models have the position_ids buffer in the checkpoint
|
||||
# vLLM will always pass position_ids as an argument, so we skip loading
|
||||
# the buffer if it exists
|
||||
self.skip_substrs.append("position_ids")
|
||||
|
||||
def create_attention_instances(
|
||||
self, attn_type: AttentionType = AttentionType.DECODER):
|
||||
# TODO(hmellor): Better way to detect encoder models
|
||||
# In encoder models, the attention layers will have `is_causal=False`
|
||||
is_encoder = lambda m: not getattr(m, "is_causal", True)
|
||||
# vLLM does not support encoder-decoder models, so if any encoder layer
|
||||
# is found, we assume the whole model is an encoder model
|
||||
if any(is_encoder(m) for m in self.model.modules()):
|
||||
attn_type = AttentionType.ENCODER_ONLY
|
||||
|
||||
# Check minimum transformers version for encoder models support
|
||||
if attn_type == AttentionType.ENCODER_ONLY:
|
||||
import transformers
|
||||
from packaging.version import Version
|
||||
installed = Version(transformers.__version__)
|
||||
required = Version("4.57.0.dev0")
|
||||
if installed < required:
|
||||
raise ValueError(
|
||||
"Encoder models with the Transformers backend require "
|
||||
f"transformers>={required}, but got {installed}")
|
||||
|
||||
return super().create_attention_instances(attn_type)
|
||||
|
||||
|
||||
@support_torch_compile(enable_if=can_enable_torch_compile)
|
||||
class TransformersForCausalLM(TransformersBase):
|
||||
@@ -710,7 +750,7 @@ class TransformersForCausalLM(TransformersBase):
|
||||
# Tell `TransformersBase.load_weights` to skip
|
||||
# `lm_head` if the model has tied word embeddings
|
||||
if self.text_config.tie_word_embeddings:
|
||||
self.skip_prefixes = ["lm_head."]
|
||||
self.skip_prefixes.append("lm_head.")
|
||||
|
||||
if get_pp_group().is_last_rank:
|
||||
self.unpadded_vocab_size = self.text_config.vocab_size
|
||||
|
||||
Reference in New Issue
Block a user