Image-Text-to-Text
Transformers
Safetensors
English
qwen2_5vl_ca
feature-extraction
conversational
custom_code
CASA-Qwen2_5-VL-3B / utils.py
ameroyer's picture nielsr's picture
nielsr HF Staff
Super-squash branch 'main' using huggingface_hub
5f7de59
Raw History Blame
13.1 kB
# 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()