# pylint: disable=protected-access """Utils to handle CA layers construction""" from contextlib import contextmanager from dataclasses import dataclass, fields from typing import Any, Callable, Generic, Literal, Sequence, TypeVar, overload from typing import cast as type_cast import torch def __split_n_merge__( x: torch.Tensor, sample_lengths: list[int], padding_side: Literal["left", "right"] = "right", pad_value: int | float | bool = 0, ) -> torch.Tensor: max_sample_length = max(sample_lengths) pad_tuple = tuple(0 for _ in range((x.ndim - 1) * 2)) return torch.stack( [ torch.nn.functional.pad( _x, pad_tuple + (0, max_sample_length - _x.shape[0]) if padding_side == "right" else pad_tuple + (max_sample_length - _x.shape[0], 0), value=pad_value, ) for _x in torch.split(x, sample_lengths, dim=0) ], dim=0, ) @overload def insert_image_tokens( inputs_embeds: torch.Tensor, image_embeds: torch.Tensor | Sequence[torch.Tensor], image_embeds_insertion_points: list[torch.Tensor], recover_batch_dim: Literal[True], attention_mask: torch.Tensor | None = None, padding_side: Literal["left", "right"] = "right", keep_only_attended: bool = False, pad_output: int | float | bool = 0.0, ) -> tuple[ torch.Tensor, None, torch.Tensor | None, torch.Tensor, ]: ... @overload def insert_image_tokens( inputs_embeds: torch.Tensor, image_embeds: torch.Tensor | Sequence[torch.Tensor], image_embeds_insertion_points: list[torch.Tensor], recover_batch_dim: Literal[False], attention_mask: torch.Tensor | None = None, padding_side: Literal["left", "right"] = "right", keep_only_attended: bool = False, pad_output: int | float | bool = 0.0, ) -> tuple[ torch.Tensor, list[int], torch.Tensor | None, torch.Tensor, ]: ... def insert_image_tokens( inputs_embeds: torch.Tensor, image_embeds: torch.Tensor | Sequence[torch.Tensor], image_embeds_insertion_points: list[torch.Tensor], recover_batch_dim: bool = True, attention_mask: torch.Tensor | None = None, padding_side: Literal["left", "right"] = "right", keep_only_attended: bool = False, pad_output: int | float | bool = 0.0, ) -> tuple[ torch.Tensor | torch.Tensor, list[int] | None, torch.Tensor | torch.Tensor | None, torch.Tensor | torch.Tensor, ]: """ Insert image embeddings into text embeddings Args: inputs_embeds (torch.Tensor): (B, S, D) input token embeddings. image_embeds (torch.Tensor | list[torch.Tensor]): (N_images, Nt, D) | List[(Nt, D)] image token embeddings. image_embeds_insertion_points (list[torch.Tensor]): Insertion indices. attention_mask (torch.Tensor, optional): (B, S) attention mask. padding_side (Literal["left", "right"]): Padding scheme. Controls behavior for padded images. return_indices (bool): Whether to return gather indices or the fused sequence directly. keep_only_attended: This is only applicable when recover_batch_dim is False; whether to remove any non-attended tokens in the whole array. In this case, the attention mask returned is **still the original one**, so we can remember which indices have been removed Returns: output (torch.Tensor): (B, S + Ni * Nt) gather indices or (B, S + Ni * Nt, D) fused sequence image_embeds (torch.Tensor): (B, Ni * Nt) image embeds, padded and batch if input was a list attention_mask (torch.Tensor): Same shape, 1 for real tokens, 0 for image and text padding. image_tokens_mask (torch.Tensor): (B, S + Ni * Nt, 1), marks image token positions. """ if isinstance(image_embeds, list) and len(image_embeds) == 0: batch_size, text_seq_length, token_dim = inputs_embeds.shape if recover_batch_dim: return ( inputs_embeds, None, attention_mask, torch.zeros((batch_size, text_seq_length, 1), dtype=torch.bool), ) else: flattened_seq_length = inputs_embeds.shape[0] * inputs_embeds.shape[1] return ( torch.reshape(inputs_embeds, (flattened_seq_length, inputs_embeds.shape[2])), [text_seq_length] * inputs_embeds.shape[0], attention_mask.flatten() if attention_mask is not None else None, torch.zeros((flattened_seq_length, 1), dtype=torch.bool), ) # Sanity checks if isinstance(image_embeds, torch.Tensor): assert inputs_embeds.shape[-1] == image_embeds.shape[-1] else: assert all(inputs_embeds.shape[-1] == _x.shape[-1] for _x in image_embeds) batch_size, text_seq_length, token_dim = inputs_embeds.shape image_seq_length = [x.shape[0] for x in image_embeds] # Flatten insertion points insertion_offset = [] counter, offset_from_text, offset_from_image = 0, 0, 0 for sample in image_embeds_insertion_points: for pt in sample: insertion_offset.append(pt + offset_from_image + offset_from_text) offset_from_image += image_seq_length[counter] counter += 1 offset_from_text += text_seq_length image_insert_positions = [ x for idx, pt in enumerate(insertion_offset) for x in range(pt, pt + image_seq_length[idx]) ] # Flatten image embeds if isinstance(image_embeds, list): image_embeds = torch.cat(image_embeds, dim=0) else: image_embeds = type_cast(torch.Tensor, image_embeds) image_embeds = torch.reshape(image_embeds, (-1, token_dim)) # Flatten text embeds across batch dim (B x S, D) inputs_embeds = torch.reshape(inputs_embeds, (-1, token_dim)) flattened_seq_length = inputs_embeds.shape[0] + sum(image_seq_length) text_insert_positions = sorted( set(range(flattened_seq_length)).difference(set(image_insert_positions)) ) # Scatter image embeds in the flattened dict # scatter text related stuff output = torch.empty( (flattened_seq_length, token_dim), device=inputs_embeds.device, dtype=inputs_embeds.dtype, ) txt_positions_tensor = torch.Tensor(text_insert_positions).to( dtype=torch.long, device=inputs_embeds.device ) output.scatter_(0, txt_positions_tensor[:, None].expand(-1, token_dim), inputs_embeds) attention_mask_new: torch.Tensor | None = None if attention_mask is not None: attention_mask_new = torch.ones( (flattened_seq_length,), dtype=torch.bool, device=inputs_embeds.device ) attention_mask_new.scatter_( 0, txt_positions_tensor, attention_mask.flatten().to(torch.bool) ) # scatter image related stuff image_tokens_mask = torch.zeros( (flattened_seq_length,), dtype=torch.bool, device=inputs_embeds.device ) img_positions_tensor = torch.Tensor(image_insert_positions).to( device=inputs_embeds.device, dtype=torch.long ) output.scatter_(0, img_positions_tensor[:, None].expand(-1, token_dim), image_embeds) image_tokens_mask.scatter_(0, img_positions_tensor, True) # Compute expected sample length, taking into account the real batch # i.e. recover the batch dimension of image embeddings sample_lengths = [] counter = 0 for sample_idx, pts in enumerate(image_embeds_insertion_points): num_image_tokens = 0 for _ in pts: num_image_tokens += image_seq_length[counter] counter += 1 if keep_only_attended and attention_mask is not None: attended_seq_length = torch.sum(attention_mask[sample_idx]).cpu().item() sample_lengths.append(attended_seq_length + num_image_tokens) else: sample_lengths.append(text_seq_length + num_image_tokens) # For CA attention, we can keep stuff flatten and return # the sample_lengths for the blockwise attention if not recover_batch_dim: if keep_only_attended and attention_mask_new is not None: output = output[attention_mask_new] image_tokens_mask = image_tokens_mask[attention_mask_new] return output, sample_lengths, attention_mask_new, image_tokens_mask[..., None] # Otherwise, time to (pad) and reshape # Easy case: everything has the same length if all(x == sample_lengths[0] for x in sample_lengths): output = torch.reshape(output, (batch_size, sample_lengths[0], token_dim)) image_tokens_mask = torch.reshape(image_tokens_mask, (batch_size, sample_lengths[0], 1)) if attention_mask_new is not None: attention_mask_new = torch.reshape(attention_mask_new, (batch_size, sample_lengths[0])) # if there is any size mismatch we break into a # list and pad again else: # split and merge output = __split_n_merge__(output, sample_lengths, padding_side, pad_value=pad_output) # note that the extra padding tokens are also marked as image tokens to be removed later image_tokens_mask = __split_n_merge__( image_tokens_mask, sample_lengths, padding_side, True )[:, :, None] if attention_mask_new is not None: attention_mask_new = __split_n_merge__( attention_mask_new, sample_lengths, padding_side, 0 ) # Return return output, sample_lengths, attention_mask_new, image_tokens_mask class SharedModuleType(type): """Wrapper to build shared Pytorch modules. This can be used as a metaclass to build shared modules; see an example in attention.py""" _instances = {} def __call__(cls, *args: Any, **kwargs: Any) -> Any: if cls not in cls._instances: cls._instances[cls] = super(SharedModuleType, cls).__call__(*args, **kwargs) return cls._instances[cls] @dataclass class StreamingState: """Streaming State used by CA layers at inference to save e.g. the offset and other persistent states""" offset: int = 0 def _is_valid_field(self, key: str) -> bool: return key in {x.name for x in fields(self)} def _init_field(self, key: str) -> None: """Init function for non-argument dependent defaults""" assert self._is_valid_field(key) if key == "offset": self.offset = 0 else: # for fields which should be set explicitly and cannot be auto-initialized setattr(self, key, None) def init(self) -> None: for key in [x.name for x in fields(self)]: self._init_field(key) def _reset_field(self, name: str) -> None: """Resets the given field""" self._init_field(name) def reset(self) -> None: for f in fields(self): self._reset_field(f.name) def _get_field(self, f: str) -> Any: """Get field and init if not""" assert self._is_valid_field(f) if getattr(self, f) is None: self._init_field(f) return getattr(self, f) def _set_field(self, f: str, value: Any) -> None: assert self._is_valid_field(f) setattr(self, f, value) StreamingStateT = TypeVar("StreamingStateT", bound=StreamingState) class StreamingModule(torch.nn.Module, Generic[StreamingStateT]): # pylint: disable=abstract-method """Streaming-aware module base class""" def __init__(self, state_class: type) -> None: torch.nn.Module.__init__(self) self.is_streaming: bool = False self.enable_viz: tuple[str, ...] = () self._streaming_state: StreamingStateT = state_class() @property def streaming_state(self) -> StreamingStateT: return self._streaming_state def _apply_named_streaming(self, fn: Callable): """Apply function to all streaming modules""" for name, module in self.named_modules(): if isinstance(module, StreamingModule): fn(name, module) def reset_streaming(self): """Reset the streaming state.""" def _reset(_: str, module: StreamingModule): module._streaming_state.reset() self._apply_named_streaming(_reset) def _set_streaming(self, streaming: bool, viz: tuple[str, ...] = ()): """Set all streaming modules in streaming mode""" def _set_streaming(_, module: StreamingModule) -> None: module.is_streaming = streaming module.enable_viz = viz if streaming: module.streaming_state.init() self._apply_named_streaming(_set_streaming) @contextmanager def streaming(self, stream: bool = True, viz: tuple[str, ...] = ()): """Context manager to enter streaming mode. Reset streaming state on exit.""" self._set_streaming(stream, viz) try: yield finally: self._set_streaming(False, ()) self.reset_streaming()