Skip to content

clip

Minimal implementation of CLIPVisionModel intended to be only used within a vision language model.

Classes

fastvideo.models.encoders.clip.CLIPAttention

CLIPAttention(config: CLIPVisionConfig | CLIPTextConfig, quant_config: QuantizationConfig | None = None, prefix: str = '')

Bases: Module

Multi-headed attention from 'Attention Is All You Need' paper

Source code in fastvideo/models/encoders/clip.py
def __init__(
    self,
    config: CLIPVisionConfig | CLIPTextConfig,
    quant_config: QuantizationConfig | None = None,
    prefix: str = "",
):
    super().__init__()
    self.config = config
    self.embed_dim = config.hidden_size
    self.num_heads = config.num_attention_heads
    self.head_dim = self.embed_dim // self.num_heads
    if self.head_dim * self.num_heads != self.embed_dim:
        raise ValueError(
            "embed_dim must be divisible by num_heads "
            f"(got `embed_dim`: {self.embed_dim} and `num_heads`:"
            f" {self.num_heads}).")
    self.scale = self.head_dim**-0.5 if config.enable_scale else None
    self.dropout = config.attention_dropout

    self.qkv_proj = QKVParallelLinear(
        hidden_size=self.embed_dim,
        head_size=self.head_dim,
        total_num_heads=self.num_heads,
        quant_config=quant_config,
        prefix=f"{prefix}.qkv_proj",
    )

    self.out_proj = RowParallelLinear(
        input_size=self.embed_dim,
        output_size=self.embed_dim,
        quant_config=quant_config,
        prefix=f"{prefix}.out_proj",
    )

    self.tp_size = get_tp_world_size()
    self.num_heads_per_partition = divide(self.num_heads, self.tp_size)

    self.attn = LocalAttention(
        self.num_heads_per_partition,
        self.head_dim,
        self.num_heads_per_partition,
        softmax_scale=self.scale,
        causal=config.is_causal,
        supported_attention_backends=config._supported_attention_backends)

Methods:

fastvideo.models.encoders.clip.CLIPAttention.forward
forward(hidden_states: Tensor)

Input shape: Batch x Time x Channel

Source code in fastvideo/models/encoders/clip.py
def forward(
    self,
    hidden_states: torch.Tensor,
):
    """Input shape: Batch x Time x Channel"""

    qkv_states, _ = self.qkv_proj(hidden_states)
    query_states, key_states, value_states = qkv_states.chunk(3, dim=-1)
    # use flash_attn_func
    query_states = query_states.reshape(query_states.shape[0],
                                        query_states.shape[1],
                                        self.num_heads_per_partition,
                                        self.head_dim)
    key_states = key_states.reshape(key_states.shape[0],
                                    key_states.shape[1],
                                    self.num_heads_per_partition,
                                    self.head_dim)
    value_states = value_states.reshape(value_states.shape[0],
                                        value_states.shape[1],
                                        self.num_heads_per_partition,
                                        self.head_dim)
    attn_output = self.attn(query_states, key_states, value_states)
    attn_output = attn_output.reshape(
        attn_output.shape[0], attn_output.shape[1],
        self.num_heads_per_partition * self.head_dim)
    attn_output, _ = self.out_proj(attn_output)

    return attn_output, None

fastvideo.models.encoders.clip.CLIPEncoder

CLIPEncoder(config: CLIPVisionConfig | CLIPTextConfig, quant_config: QuantizationConfig | None = None, num_hidden_layers_override: int | None = None, prefix: str = '')

Bases: Module

Transformer encoder consisting of config.num_hidden_layers self attention layers. Each layer is a [CLIPEncoderLayer].

Parameters:

Name Type Description Default
config CLIPVisionConfig | CLIPTextConfig

CLIPConfig

required
Source code in fastvideo/models/encoders/clip.py
def __init__(
    self,
    config: CLIPVisionConfig | CLIPTextConfig,
    quant_config: QuantizationConfig | None = None,
    num_hidden_layers_override: int | None = None,
    prefix: str = "",
) -> None:
    super().__init__()

    self.config = config

    if num_hidden_layers_override is None:
        num_hidden_layers = config.num_hidden_layers
    else:
        num_hidden_layers = num_hidden_layers_override
    self.layers = nn.ModuleList([
        CLIPEncoderLayer(config=config,
                         quant_config=quant_config,
                         prefix=f"{prefix}.layers.{layer_idx}")
        for layer_idx in range(num_hidden_layers)
    ])

fastvideo.models.encoders.clip.CLIPTextTransformer

CLIPTextTransformer(config: CLIPTextConfig, quant_config: QuantizationConfig | None = None, num_hidden_layers_override: int | None = None, prefix: str = '')

Bases: Module

Source code in fastvideo/models/encoders/clip.py
def __init__(self,
             config: CLIPTextConfig,
             quant_config: QuantizationConfig | None = None,
             num_hidden_layers_override: int | None = None,
             prefix: str = ""):
    super().__init__()
    self.config = config
    embed_dim = config.hidden_size

    self.embeddings = CLIPTextEmbeddings(config)

    self.encoder = CLIPEncoder(
        config,
        quant_config=quant_config,
        num_hidden_layers_override=num_hidden_layers_override,
        prefix=prefix)

    self.final_layer_norm = nn.LayerNorm(embed_dim,
                                         eps=config.layer_norm_eps)

    # For `pooled_output` computation
    self.eos_token_id = config.eos_token_id

Methods:

fastvideo.models.encoders.clip.CLIPTextTransformer.forward
forward(input_ids: Tensor | None, position_ids: Tensor | None = None, attention_mask: Tensor | None = None, inputs_embeds: Tensor | None = None, output_hidden_states: bool | None = None) -> BaseEncoderOutput

Returns:

Source code in fastvideo/models/encoders/clip.py
def forward(
    self,
    input_ids: torch.Tensor | None,
    position_ids: torch.Tensor | None = None,
    attention_mask: torch.Tensor | None = None,
    inputs_embeds: torch.Tensor | None = None,
    output_hidden_states: bool | None = None,
) -> BaseEncoderOutput:
    r"""
    Returns:

    """
    output_hidden_states = (output_hidden_states
                            if output_hidden_states is not None else
                            self.config.output_hidden_states)

    if input_ids is None:
        raise ValueError("You have to specify input_ids")

    input_shape = input_ids.size()
    input_ids = input_ids.view(-1, input_shape[-1])

    hidden_states = self.embeddings(input_ids=input_ids,
                                    position_ids=position_ids)

    # CLIP's text model uses causal mask, prepare it here.
    # https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
    # causal_attention_mask = _create_4d_causal_attention_mask(
    #     input_shape, hidden_states.dtype, device=hidden_states.device
    # )

    # # expand attention_mask
    # if attention_mask is not None and not self._use_flash_attention_2:
    #     raise NotImplementedError("attention_mask is not supported for CLIPTextTransformer")
    #     # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
    #     attention_mask = _prepare_4d_attention_mask(attention_mask, hidden_states.dtype)

    encoder_outputs = self.encoder(
        inputs_embeds=hidden_states,
        # attention_mask=attention_mask,
        # causal_attention_mask=causal_attention_mask,
        # output_attentions=output_attentions,
        return_all_hidden_states=output_hidden_states,
        # return_dict=return_dict,
    )

    last_hidden_state = encoder_outputs[-1]
    last_hidden_state = self.final_layer_norm(last_hidden_state)

    if self.eos_token_id == 2:
        # The `eos_token_id` was incorrect before PR #24773: Let's keep what have been done here.
        # A CLIP model with such `eos_token_id` in the config can't work correctly with extra new tokens added
        # ------------------------------------------------------------
        # text_embeds.shape = [batch_size, sequence_length, transformer.width]
        # take features from the eot embedding (eot_token is the highest number in each sequence)
        # casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14
        pooled_output = last_hidden_state[
            torch.arange(last_hidden_state.shape[0],
                         device=last_hidden_state.device),
            input_ids.to(dtype=torch.int, device=last_hidden_state.device).
            argmax(dim=-1),
        ]
    else:
        # The config gets updated `eos_token_id` from PR #24773 (so the use of exta new tokens is possible)
        pooled_output = last_hidden_state[
            torch.arange(last_hidden_state.shape[0],
                         device=last_hidden_state.device),
            # We need to get the first position of `eos_token_id` value (`pad_token_ids` might equal to `eos_token_id`)
            # Note: we assume each sequence (along batch dim.) contains an  `eos_token_id` (e.g. prepared by the tokenizer)
            (input_ids.to(dtype=torch.int, device=last_hidden_state.device
                          ) == self.eos_token_id).int().argmax(dim=-1),
        ]

    return BaseEncoderOutput(
        last_hidden_state=last_hidden_state,
        pooler_output=pooled_output,
        hidden_states=encoder_outputs,
        # attentions=encoder_outputs.attentions,
    )

Functions: