"""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