Image-Text-to-Text
Transformers
Safetensors
English
qwen2_5vl_ca
feature-extraction
conversational
custom_code
CASA-Qwen2_5-VL-3B / modeling_qwen2_5vl_ca.py
ameroyer's picture nielsr's picture
nielsr HF Staff
Super-squash branch 'main' using huggingface_hub
5f7de59
Raw History Blame
12.8 kB
from functools import partial
from typing import Any, Sequence
from typing import cast as type_cast
import torch
from transformers.cache_utils import DynamicCache
from transformers.generation.utils import GenerateOutput
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
Qwen2_5_VLCausalLMOutputWithPast,
Qwen2_5_VLForConditionalGeneration,
)
from .cross_attention import CrossAttentionHandler, tie_qkvo_projections
from .image_encoder import Qwen25VLEncoder
from .configuration_qwen2_5vl_ca import Qwen2_5_VLCAConfig
from .language_qwen2_5vl_ca import (
Qwen2_5_VLAttention_CrossAttention,
QwenCrossAttention,
maybe_replace_with_cross_attention_layers,
)
class V2Qwen2_5VL(Qwen2_5_VLForConditionalGeneration): # pyright: ignore[reportIncompatibleMethodOverride]
config_class = Qwen2_5_VLCAConfig
def __init__(self, config: Qwen2_5_VLCAConfig, **kwargs: Any) -> None:
del kwargs
super().__init__(config)
# Wrap the Qwen visual encoder so its output matches our CA interface
self.image_prefix = Qwen25VLEncoder(self.visual) # type: ignore[assignment]
self.visual = None
self.model.apply(
partial(maybe_replace_with_cross_attention_layers, xa_layers=self.config.xa_layers)
)
# The cross-attention layers are swapped in after the base post_init, so register
# the shared-weight alias keys and (re-)tie now that the CA modules exist.
if config.xa_share_qkvo:
shared_keys: list[str] = []
for i, layer in enumerate(self.model.layers):
if isinstance(layer.self_attn, Qwen2_5_VLAttention_CrossAttention):
for proj, biased in (
("q_proj", True),
("k_proj", True),
("v_proj", True),
("o_proj", False),
):
prefix = f"model.layers.{i}.self_attn.cross_attn.{proj}"
shared_keys.append(f"{prefix}.weight")
if biased:
shared_keys.append(f"{prefix}.bias")
self._tied_weights_keys = list(self._tied_weights_keys or []) + shared_keys
self.tie_weights()
def _tie_weights(self) -> None:
if not getattr(self.config, "xa_share_qkvo", False):
return
for layer in self.model.layers:
if isinstance(layer.self_attn, Qwen2_5_VLAttention_CrossAttention):
tie_qkvo_projections(layer.self_attn, layer.self_attn.cross_attn)
def get_device(self) -> str:
"""Return the device type of the model"""
return next(self.parameters()).device.type
@property
def token_dim(self) -> int:
"""Returns the number of dimensions for the token representation"""
return self.config.hidden_size
def _update_model_kwargs_for_generation(
self,
outputs: Any,
model_kwargs: dict[str, Any],
is_encoder_decoder: bool = False,
num_new_tokens: int = 1,
):
"""Override to handle multi-turn generation and propagate updated attention masks"""
if (am := outputs.get("updated_attention_mask", None)) is not None:
model_kwargs["attention_mask"] = am
if "updated_cache_position" in outputs:
model_kwargs["cache_position"] = outputs.get("updated_cache_position")
else:
start = 0
if (kv := model_kwargs.get("past_key_values", None)) is not None:
start = kv._seen_tokens - am.shape[1]
model_kwargs["cache_position"] = torch.arange(
start,
start + am.shape[1],
dtype=model_kwargs["cache_position"].dtype,
device=model_kwargs["cache_position"].device,
)
# Call parent to get default updates
model_kwargs = super()._update_model_kwargs_for_generation(
outputs, model_kwargs, is_encoder_decoder, num_new_tokens
)
# Used by prepare_inputs_for_generation
model_kwargs["__is_first_gen_call__"] = False
return model_kwargs
def prepare_inputs_for_generation( # pyright: ignore[reportIncompatibleMethodOverride]
self,
input_ids: torch.Tensor,
past_key_values: DynamicCache | None = None,
**kwargs: Any,
):
"""Override to avoid Qwen erasing pixel_values on subsequent generation calls"""
backup = None
__is_first_gen_call__ = kwargs.get("__is_first_gen_call__", True)
if __is_first_gen_call__:
backup = kwargs.get("pixel_values", None)
if past_key_values is not None and (
kwargs.get("cache_position") is None
or type_cast(torch.Tensor, kwargs.get("cache_position")).shape[0] == 0
):
# We're continuing from a cached state
past_length = past_key_values._seen_tokens
kwargs["cache_position"] = torch.arange(
past_length,
past_length + (input_ids.shape[1] if __is_first_gen_call__ else 1),
dtype=torch.long,
device=input_ids.device,
)
out = super().prepare_inputs_for_generation(
input_ids,
past_key_values=past_key_values,
**kwargs,
)
if backup is not None:
out["pixel_values"] = backup
return out
def prepare_multimodal_inputs(
self,
input_ids: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
image_embeds_insertion_points: list[torch.Tensor] | None = None,
labels: torch.Tensor | None = None,
pixel_values: torch.Tensor | list[torch.Tensor] | None = None,
pre_image_tokens: list[int] | None = None,
post_image_tokens: list[int] | None = None,
**_kwargs: Any,
) -> dict:
"""Get a batch data mixing text and image data"""
del _kwargs
processed_inputs: dict = {
"input_ids": input_ids,
"inputs_embeds": inputs_embeds,
"labels": labels,
"attention_mask": attention_mask,
"image_embeds_insertion_points": image_embeds_insertion_points,
}
if pixel_values is not None:
processed_inputs.update(self.image_prefix(pixel_values))
image_embeds = processed_inputs.get("image_embeds")
assert image_embeds is not None
assert (isinstance(image_embeds, torch.Tensor) and image_embeds.ndim == 3) or (
isinstance(image_embeds, list) and all(_x.ndim == 2 for _x in image_embeds)
)
# Add kwargs necessary to compute cu_seqlens windows for CA
processed_inputs["ca_windows_info"] = {
"num_post_image_tokens": 0 if post_image_tokens is None else len(post_image_tokens),
"num_pre_image_tokens": 0 if pre_image_tokens is None else len(pre_image_tokens),
}
return processed_inputs
def forward( # type: ignore[override] # pylint: disable=W0221
self,
input_ids: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
pixel_values: torch.Tensor | list[torch.Tensor] | None = None,
return_loss: bool = True,
labels: torch.Tensor | None = None,
image_embeds_insertion_points: list[torch.Tensor] | None = None,
pre_image_tokens: list[int] | None = None,
post_image_tokens: list[int] | None = None,
**kwargs: Any,
) -> tuple | Qwen2_5_VLCausalLMOutputWithPast:
"""Multi-modal forward pass"""
if self.training:
assert return_loss is True, (
"Qwen2.5VL always computes its own labels/losses in train mode"
)
if inputs_embeds is None:
assert input_ids is not None
inputs_embeds = type_cast(torch.Tensor, self.model.embed_tokens(input_ids))
# Case 1: First generation call — compute image embeddings and set up CA handler
if kwargs.pop("__is_first_gen_call__", True):
processed_inputs = self.prepare_multimodal_inputs(
input_ids=input_ids,
inputs_embeds=inputs_embeds,
attention_mask=attention_mask,
image_embeds_insertion_points=image_embeds_insertion_points,
pixel_values=pixel_values,
labels=labels,
pre_image_tokens=pre_image_tokens,
post_image_tokens=post_image_tokens,
)
image_embeds = processed_inputs.get("image_embeds", None)
inst_points = processed_inputs.get("image_embeds_insertion_points", None)
# Only build a handler when images are actually present
cross_attention_handler: CrossAttentionHandler | None = None
if image_embeds is not None and len(image_embeds) > 0:
cross_attention_handler = CrossAttentionHandler(
inputs_embeds=torch.zeros_like(inputs_embeds),
image_embeds=image_embeds,
image_embeds_insertion_points=inst_points,
ca_windows_info=processed_inputs.pop("ca_windows_info", None),
training=self.training,
)
self.update_cross_attention_states(cross_attention_handler)
# Run Qwen with the attention layers replaced to use cross-attention
assert inputs_embeds is not None, "Could not compute input embeddings!"
out = super().forward(
inputs_embeds=inputs_embeds, # type: ignore[arg-type]
attention_mask=attention_mask,
pixel_values=None,
**kwargs,
)
return out
@property
def default_generation_eos_token_id(self) -> int | Sequence[int] | None:
return self.generation_config.eos_token_id if self.generation_config is not None else None
@torch.no_grad()
def generate_from_image( # pyright: ignore[reportInconsistentOverload]
self,
reset_streaming: bool = True,
temperature: float | None = 0.0,
eos_token_id: int | Sequence[int] | None = None,
**kwargs: Any,
) -> GenerateOutput | torch.LongTensor:
"""Custom generate function"""
# init self-attention KVCache
if kwargs.get("past_key_values", None) is None:
kwargs["past_key_values"] = DynamicCache()
if eos_token_id is None:
eos_token_id = self.default_generation_eos_token_id
# To avoid generate warning
if kwargs.get("pad_token_id", None) is None:
kwargs["pad_token_id"] = kwargs.get("eos_token_id", None)
if isinstance(kwargs["pad_token_id"], (list, tuple)):
kwargs["pad_token_id"] = kwargs["pad_token_id"][0]
if "pre_image_tokens" not in kwargs:
kwargs["pre_image_tokens"] = list(self.config.pre_image_tokens)
if "post_image_tokens" not in kwargs:
kwargs["post_image_tokens"] = list(self.config.post_image_tokens)
if not kwargs.get("do_sample", False):
temperature = None
kwargs.pop("top_p", None)
kwargs.pop("top_k", None)
# Generate
self.start_ca_streaming_states()
outputs = self.generate(
use_cache=True,
eos_token_id=eos_token_id,
temperature=temperature,
**kwargs,
)
if reset_streaming:
self.reset_ca_streaming_states()
return outputs
def update_cross_attention_states(self, handler: CrossAttentionHandler | None):
"""Push the new handler into all CA attention layers"""
def __update__(m: torch.nn.Module):
nonlocal handler
if isinstance(m, Qwen2_5_VLAttention_CrossAttention):
m.cross_attention_handler = handler
self.apply(__update__)
def reset_ca_streaming_states(self) -> None:
def __reset__(m: torch.nn.Module):
if isinstance(m, QwenCrossAttention):
m._set_streaming(False, ())
m.reset_streaming()
if hasattr(m, "cross_attention_handler"):
del m.cross_attention_handler
m.cross_attention_handler = None
self.apply(__reset__)
def start_ca_streaming_states(self) -> None:
def __start__(m: torch.nn.Module):
if isinstance(m, QwenCrossAttention):
m._set_streaming(True, ())
self.apply(__start__)