Image-Text-to-Text
Transformers
Safetensors
English
qwen2_5vl_ca
feature-extraction
conversational
custom_code
CASA-Qwen2_5-VL-3B / cross_attention.py
ameroyer's picture nielsr's picture
nielsr HF Staff
Super-squash branch 'main' using huggingface_hub
5f7de59
Raw History Blame
17.7 kB
"""Modular cross-attention module"""
import bisect
import copy
from itertools import accumulate
from typing import TYPE_CHECKING, Literal, TypedDict, TypeVar
from typing import cast as type_cast
import torch
from .utils import StreamingModule, StreamingState
if TYPE_CHECKING:
from transformers.configuration_utils import PretrainedConfig
try:
from flash_attn import flash_attn_varlen_func
except ImportError:
flash_attn_varlen_func = None # type: ignore
AttentionBaseT = TypeVar("AttentionBaseT", bound=StreamingModule)
WindowsComputeKwargs = TypedDict(
"WindowsComputeKwargs",
{
"num_post_image_tokens": int,
"num_pre_image_tokens": int,
},
total=False,
)
def get_sample_lengths_for_xa(
image_embeds_insertion_points: list[torch.Tensor],
image_embeds: torch.Tensor | list[torch.Tensor] | None,
total_seq_len: int,
attention_mask: torch.Tensor | None = None,
**kwargs: WindowsComputeKwargs,
) -> tuple[list[tuple[int, bool]], list[int], torch.Tensor | None]:
"""Sample lengths for cross-attention. Compared to other functions in this file,
it also returns a mask on the text tokens to mark tokens which do not relate
to any images (e.g. squashed text-only-samples, or BoS in prefix_after_bos)
"""
num_post_image_tokens = type_cast(int, kwargs.get("num_post_image_tokens", 0))
num_pre_image_tokens = type_cast(int, kwargs.get("num_pre_image_tokens", 0))
squashed_samples_lengths = type_cast(
list[list[int]] | None, kwargs.get("squashed_samples_lengths", None)
)
if squashed_samples_lengths is not None:
assert len(squashed_samples_lengths) == len(image_embeds_insertion_points)
def __insert_next_sample__(
batch_idx: int, insrt_pt: int, last_insrt_pt: int, end_of_batch_sample: bool = False
) -> None:
nonlocal attention_mask, active_tokens
nonlocal text_sample_lengths, full_sample_lengths
nonlocal cum_samples_lengths, current_image_offset
# Add the sample between [last_insrt_pt, insrt_pt] with breaks in
# between any squashed samples we find on the way
nonlocal last_image_idx, current_image_idx, current_length
added_sample = False
start_pt = bisect.bisect_left(cum_samples_lengths, last_insrt_pt)
for end_of_sample in cum_samples_lengths[start_pt:]:
# we will break the loop at the end when end_of_sample = insrt_pt
end_of_sample = min(end_of_sample, insrt_pt)
# Add between [last_insrt_pt, end_of_sample]
current_length = end_of_sample - last_insrt_pt
num_padding_tokens = 0
if attention_mask is not None:
num_padding_tokens = int(
torch.sum(~attention_mask[batch_idx, last_insrt_pt:end_of_sample]).item()
)
current_length -= num_padding_tokens
num_image_tokens = 0
if current_length > 0:
# add image tokens to current_length
added_sample = True
if current_image_idx > 0 and image_embeds is not None:
images_in_sample = [
img_idx
for img_idx in range(last_image_idx, current_image_idx)
if img_idx < len(image_embeds_insertion_points[batch_idx])
and last_insrt_pt
<= image_embeds_insertion_points[batch_idx][img_idx]
< end_of_sample
]
if len(images_in_sample) > 0:
num_image_tokens = sum(
_x.shape[0]
for _x in image_embeds[
current_image_offset + images_in_sample[0] : current_image_offset
+ images_in_sample[-1]
+ 1
]
)
# If no image, we should not insert and instead make it as inactive
if num_image_tokens > 0:
text_sample_lengths.append(
(current_length, end_of_batch_sample and insrt_pt == end_of_sample)
)
full_sample_lengths.append(current_length + num_image_tokens)
# Active tokens
active_tokens += [int(num_image_tokens > 0)] * (current_length + num_padding_tokens)
# prepare for next loop
last_insrt_pt = end_of_sample
if end_of_sample == insrt_pt:
break
# End of loop: catching edge case where we end up on a span full of padding
if end_of_batch_sample:
assert added_sample, "Weird edge case. Don't do that, thank you"
text_sample_lengths[-1] = (text_sample_lengths[-1][0], True)
current_image_offset = 0
text_sample_lengths, full_sample_lengths = [], []
cum_samples_lengths: list[int] = []
active_tokens = []
current_length, last_insrt_pt, last_image_idx, current_image_idx = 0, 0, 0, 0
for batch_idx, pts in enumerate(image_embeds_insertion_points):
if squashed_samples_lengths is not None:
cum_samples_lengths = list(accumulate(squashed_samples_lengths[batch_idx]))
else:
cum_samples_lengths = [total_seq_len]
for current_image_idx, insrt_pt in enumerate(pts.cpu().tolist()):
# check if the images are consecutive in which way we want
# them to belong to the same window
if current_image_idx >= 1 and insrt_pt == (
image_embeds_insertion_points[batch_idx][current_image_idx - 1]
+ num_pre_image_tokens
+ num_post_image_tokens
):
continue
# Otherwise, we found a new sample
# not very important but for completeness: the insertion points come *after*
# the pre-image tokens per design but for the document-id mask it is more consistent to
# have them correspond to the same image
insrt_pt -= num_pre_image_tokens
# Compute length between the two insertion points
current_length = insrt_pt - last_insrt_pt
if attention_mask is not None:
current_length -= int(
torch.sum(~attention_mask[batch_idx, last_insrt_pt:insrt_pt]).item()
)
# Update text and full sample lengths
if insrt_pt > last_insrt_pt:
__insert_next_sample__(
batch_idx, insrt_pt, last_insrt_pt, end_of_batch_sample=False
)
last_image_idx = current_image_idx
last_insrt_pt = insrt_pt
# End of batch: add sample in progress and reset
current_image_idx += 1
if cum_samples_lengths[-1] > last_insrt_pt:
__insert_next_sample__(
batch_idx, cum_samples_lengths[-1], last_insrt_pt, end_of_batch_sample=True
)
current_length, last_insrt_pt, last_image_idx, current_image_idx = 0, 0, 0, 0
current_image_offset += len(pts)
# Sample lengths
if image_embeds is None:
return text_sample_lengths, full_sample_lengths, None
return (
text_sample_lengths,
full_sample_lengths,
torch.tensor(active_tokens, dtype=torch.bool, device=image_embeds[0].device),
)
class CrossAttentionHandler:
def __init__(
self,
inputs_embeds: torch.Tensor,
image_embeds: torch.Tensor | list[torch.Tensor],
image_embeds_insertion_points: list[torch.Tensor] | None,
# info for building text->image windows link
ca_windows_info: None | WindowsComputeKwargs = None,
training: bool = True,
):
if image_embeds_insertion_points is None:
image_embeds_insertion_points = [
torch.tensor(
[0] * len(image_embeds), # type: ignore[arg-type]
dtype=torch.long,
device=image_embeds[0].device, # type: ignore[index]
)
]
# Create cu_seq_lens for queries (text tokens)
# Compute sample lengths based on image insertion points to get cu_seq_lens
text_sample_lengths, full_sample_lengths, self.active_tokens_mask = (
get_sample_lengths_for_xa(
image_embeds_insertion_points=image_embeds_insertion_points,
image_embeds=image_embeds,
total_seq_len=inputs_embeds.shape[1],
**(ca_windows_info or {}), # pyright: ignore[reportArgumentType]
)
)
if self.active_tokens_mask is None:
self.active_tokens_mask = torch.zeros(
(inputs_embeds.shape[0] * inputs_embeds.shape[1],),
dtype=torch.bool,
device=inputs_embeds.device,
)
assert sum(_x[0] for _x in text_sample_lengths) == int(
torch.sum(self.active_tokens_mask).item()
), "Sanity check"
self.cu_seqlens_q = torch.Tensor(
list(accumulate([_x[0] for _x in text_sample_lengths], initial=0))
).to(dtype=torch.int32, device=inputs_embeds.device)
self.max_seqlen_q = max(_x[0] for _x in text_sample_lengths)
# Create cu_seq_lens for keys values (the image tokens) while grouping the
# images which are consecutive
image_lens = [(l2 - l1) for (l1, _), l2 in zip(text_sample_lengths, full_sample_lengths)]
self.cu_seqlens_kv = torch.Tensor(list(accumulate(image_lens, initial=0))).to(
dtype=torch.int32, device=inputs_embeds.device
)
self.max_seqlen_kv = max(image_lens)
self.image_embeds = torch.cat([_x for _x in image_embeds], dim=0)[None, :, :]
def get_active_tokens(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Return tokens to be used as queries while ignoring
tokens who have nothing to do with an image"""
channels = hidden_states.shape[-1]
src = hidden_states.flatten(0, 1)
if self.active_tokens_mask is None:
return src.reshape((1, -1, channels))
return torch.masked_select(src, self.active_tokens_mask[:, None]).reshape((1, -1, channels))
def replace_active_tokens(
self, token_updates: torch.Tensor, hidden_states_in: torch.Tensor
) -> torch.Tensor:
if self.active_tokens_mask is None:
return token_updates
updates = torch.zeros_like(hidden_states_in.flatten(0, 1))
updates.masked_scatter_(source=token_updates, mask=self.active_tokens_mask[:, None])
return updates
def tie_qkvo_projections(self_attn: torch.nn.Module, cross_attn: "CrossAttention") -> None:
"""Alias the q/k/v/o projections of `cross_attn` onto those of `self_attn`.
Used in the `xa_share_qkvo` setting so the cross-attention reuses the host
self-attention's projection weights (a single set of parameters). The two attention
calls and their independent softmaxes are otherwise unchanged.
:param self_attn: the host self-attention module owning the q/k/v/o projections
:param cross_attn: the cross-attention module whose projections are aliased
"""
for proj in ("q_proj", "k_proj", "v_proj", "o_proj"):
setattr(cross_attn, proj, getattr(self_attn, proj))
class CrossAttention(StreamingModule[StreamingState]):
"""Attention module between images and text tokens"""
def __init__(
self,
config: "PretrainedConfig",
layer_idx: int | None,
input_layernorm: torch.nn.Module | None = None,
):
super().__init__(StreamingState)
self.head_dim = config.head_dim
self.config = config
self.is_first_ca_layer = layer_idx == (min(config.xa_layers) if config.xa_layers else 0)
# When weights are shared with the host self-attention, the q/k/v/o projections
# are not owned by this module; they are aliased onto the self-attn projections
# by the host model (see tie_qkvo_projections).
self.xa_share_qkvo: bool = getattr(config, "xa_share_qkvo", False)
if not self.xa_share_qkvo:
self.q_proj = self.init_from_config_proj("q", config)
self.k_proj = self.init_from_config_proj("k", config)
self.v_proj = self.init_from_config_proj("v", config)
self.o_proj = self.init_from_config_proj("o", config)
self.norm: torch.nn.Module | None = copy.deepcopy(input_layernorm)
self.cross_attention_handler = None
# (source image embeddings, projected keys, projected values); the projections
# are reused across decoding steps as long as the source tensor is unchanged
self._cached_image_kv: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None
def init_from_config_proj(
self, key: Literal["q", "o", "k", "v"], config: "PretrainedConfig"
) -> torch.nn.Linear:
"""Initialize the Linear proj in this module"""
num_heads = config.num_key_value_heads if key in {"k", "v"} else config.num_attention_heads
return torch.nn.Linear(
config.hidden_size,
num_heads * config.head_dim,
bias=config.attention_bias if key != "o" else False,
)
def reset_streaming(self):
super().reset_streaming()
self._cached_image_kv = None
def forward( # pyright: ignore[reportIncompatibleMethodOverride]
self, hidden_states: torch.Tensor, cross_attention_handler: CrossAttentionHandler | None
) -> torch.Tensor | None:
if self.is_streaming:
if self.cross_attention_handler is None:
self.cross_attention_handler = cross_attention_handler
else:
# extend the shared handler
cross_attention_handler = self.cross_attention_handler
if self.is_first_ca_layer:
cross_attention_handler.active_tokens_mask = None
cross_attention_handler.cu_seqlens_q = torch.tensor(
range(0, hidden_states.shape[0] + 1),
dtype=cross_attention_handler.cu_seqlens_q.dtype,
device=cross_attention_handler.cu_seqlens_q.device,
)
# Case of text-only samples (or inference when no handler was cached)
# in this case we just want to skip cross-attention
if cross_attention_handler is None:
return None
# Currently, we assume that images never change/get added on the fly at inference
if self.is_streaming and self.streaming_state.offset > 0:
assert hidden_states.shape[0] == (len(cross_attention_handler.cu_seqlens_kv) - 1)
og_dtype = hidden_states.dtype
og_shape = hidden_states.shape
# kv inputs: (1, num_total_image_tokens, dim)
q_inputs = cross_attention_handler.get_active_tokens(hidden_states)
kv_inputs = cross_attention_handler.image_embeds
if self.norm is not None:
q_inputs = self.norm(q_inputs)
assert q_inputs.shape[0] == kv_inputs.shape[0] == 1
# Compute QKV for the blockwise attention
bs = 1
hidden_shape_q = (bs, q_inputs.shape[1], -1, self.head_dim)
query_states = self.q_proj(q_inputs).view(*hidden_shape_q)
# The image keys/values are identical at every decoding step, so at inference
# they are computed on the first call and reused until the images change
if self._cached_image_kv is not None and self._cached_image_kv[0] is kv_inputs:
_, key_states, value_states = self._cached_image_kv
else:
normed_kv = self.norm(kv_inputs) if self.norm is not None else kv_inputs
hidden_shape_kv = (bs, kv_inputs.shape[1], -1, self.head_dim)
key_states = self.k_proj(normed_kv).view(*hidden_shape_kv)
value_states = self.v_proj(normed_kv).view(*hidden_shape_kv)
if self.is_streaming:
self._cached_image_kv = (kv_inputs, key_states, value_states)
assert flash_attn_varlen_func is not None, (
"flash_attention is not installed but required for block-wise attention"
)
assert cross_attention_handler.cu_seqlens_q[-1] == query_states.shape[1], (
f"{cross_attention_handler.cu_seqlens_q[-1]} != {query_states.shape[1]}"
)
attn_output: torch.Tensor = flash_attn_varlen_func(
query_states[0].to(torch.bfloat16),
key_states[0].to(torch.bfloat16),
value_states[0].to(torch.bfloat16),
cu_seqlens_q=cross_attention_handler.cu_seqlens_q,
cu_seqlens_k=cross_attention_handler.cu_seqlens_kv,
max_seqlen_q=cross_attention_handler.max_seqlen_q,
max_seqlen_k=cross_attention_handler.max_seqlen_kv,
dropout_p=0.0,
# No need for causality when cross-attending to image tokens since
# image tokens are never padded
causal=False,
).to(og_dtype)
attn_output = attn_output.reshape(hidden_shape_q[1], -1).contiguous()
attn_output = self.o_proj(attn_output)
# Reshape from flattened to non-flattened
attn_output = cross_attention_handler.replace_active_tokens(attn_output, hidden_states)
attn_output = attn_output.reshape(og_shape)
if self.is_streaming:
self.streaming_state.offset += attn_output.shape[1]
return attn_output