Image-Text-to-Text
Transformers
Safetensors
English
qwen2_5vl_ca
feature-extraction
conversational
custom_code
Instructions to use kyutai/CASA-Qwen2_5-VL-3B with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use kyutai/CASA-Qwen2_5-VL-3B with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("image-text-to-text", model="kyutai/CASA-Qwen2_5-VL-3B", trust_remote_code=True) messages = [ { "role": "user", "content": [ {"type": "image", "url": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/p-blog/candy.JPG"}, {"type": "text", "text": "What animal is on the candy?"} ] }, ] pipe(text=messages)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("kyutai/CASA-Qwen2_5-VL-3B", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use kyutai/CASA-Qwen2_5-VL-3B with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "kyutai/CASA-Qwen2_5-VL-3B" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "kyutai/CASA-Qwen2_5-VL-3B", "messages": [ { "role": "user", "content": [ { "type": "text", "text": "Describe this image in one sentence." }, { "type": "image_url", "image_url": { "url": "https://cdn.britannica.com/61/93061-050-99147DCE/Statue-of-Liberty-Island-New-York-Bay.jpg" } } ] } ] }'Use Docker
docker model run hf.co/kyutai/CASA-Qwen2_5-VL-3B
- SGLang
How to use kyutai/CASA-Qwen2_5-VL-3B with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "kyutai/CASA-Qwen2_5-VL-3B" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "kyutai/CASA-Qwen2_5-VL-3B", "messages": [ { "role": "user", "content": [ { "type": "text", "text": "Describe this image in one sentence." }, { "type": "image_url", "image_url": { "url": "https://cdn.britannica.com/61/93061-050-99147DCE/Statue-of-Liberty-Island-New-York-Bay.jpg" } } ] } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "kyutai/CASA-Qwen2_5-VL-3B" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "kyutai/CASA-Qwen2_5-VL-3B", "messages": [ { "role": "user", "content": [ { "type": "text", "text": "Describe this image in one sentence." }, { "type": "image_url", "image_url": { "url": "https://cdn.britannica.com/61/93061-050-99147DCE/Statue-of-Liberty-Island-New-York-Bay.jpg" } } ] } ] }' - Docker Model Runner
How to use kyutai/CASA-Qwen2_5-VL-3B with Docker Model Runner:
docker model run hf.co/kyutai/CASA-Qwen2_5-VL-3B
Download cross_attention.py from kyutai/CASA-Qwen2_5-VL-3B: direct link, hf CLI and curl.
- Browser
- Download file 17.7 kB
-
https://huggingface.co/kyutai/CASA-Qwen2_5-VL-3B/resolve/5f7de59b9176f39e65882ab7f22fbebdcd5dc0a6/cross_attention.py
- Command line
-
hf download hf://kyutai/CASA-Qwen2_5-VL-3B@5f7de59b9176f39e65882ab7f22fbebdcd5dc0a6/cross_attention.py
-
curl -L -o cross_attention.py https://huggingface.co/kyutai/CASA-Qwen2_5-VL-3B/resolve/5f7de59b9176f39e65882ab7f22fbebdcd5dc0a6/cross_attention.py
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 | |