Skip to content

t5

PyTorch T5 & UMT5 model.

Classes

fastvideo.models.encoders.t5.AttentionType

Attention type. Use string to be compatible with torch.compile.

fastvideo.models.encoders.t5.T5Attention

T5Attention(config: T5Config, attn_type: str, has_relative_attention_bias=False, quant_config: QuantizationConfig | None = None, prefix: str = '')

Bases: Module

Source code in fastvideo/models/encoders/t5.py
def __init__(
    self,
    config: T5Config,
    attn_type: str,
    has_relative_attention_bias=False,
    quant_config: QuantizationConfig | None = None,
    prefix: str = "",
):
    super().__init__()
    self.attn_type = attn_type
    # Cross-attention has no relative pos encoding anyway
    self.is_decoder = attn_type == AttentionType.DECODER
    self.has_relative_attention_bias = has_relative_attention_bias
    self.relative_attention_num_buckets = (
        config.relative_attention_num_buckets
    )
    self.relative_attention_max_distance = (
        config.relative_attention_max_distance
    )
    self.d_model = config.d_model
    self.key_value_proj_dim = config.d_kv
    self.total_num_heads = self.total_num_kv_heads = config.num_heads

    # Partition heads across multiple tensor parallel GPUs.
    tp_world_size = get_tp_world_size()
    assert config.num_heads % tp_world_size == 0
    self.n_heads = config.num_heads // tp_world_size

    self.inner_dim = self.n_heads * self.key_value_proj_dim
    # No GQA in t5.
    # self.n_kv_heads = self.n_heads

    self.qkv_proj = QKVParallelLinear(
        self.d_model,
        self.key_value_proj_dim,
        self.total_num_heads,
        self.total_num_kv_heads,
        bias=False,
        quant_config=quant_config,
        prefix=f"{prefix}.qkv_proj",
    )

    self.attn = T5MultiHeadAttention()

    if self.has_relative_attention_bias:
        self.relative_attention_bias = VocabParallelEmbedding(
            self.relative_attention_num_buckets,
            self.total_num_heads,
            org_num_embeddings=self.relative_attention_num_buckets,
            padding_size=self.relative_attention_num_buckets,
            quant_config=quant_config,
        )
    self.o = RowParallelLinear(
        self.total_num_heads * self.key_value_proj_dim,
        self.d_model,
        bias=False,
        quant_config=quant_config,
        prefix=f"{prefix}.o_proj",
    )

Methods:

fastvideo.models.encoders.t5.T5Attention.compute_bias
compute_bias(query_length, key_length, device=None) -> Tensor

Compute binned relative position bias

Source code in fastvideo/models/encoders/t5.py
def compute_bias(
    self, query_length, key_length, device=None
) -> torch.Tensor:
    """Compute binned relative position bias"""
    if device is None:
        device = self.relative_attention_bias.weight.device
    context_position = torch.arange(
        query_length, dtype=torch.long, device=device
    )[:, None]
    memory_position = torch.arange(
        key_length, dtype=torch.long, device=device
    )[None, :]
    # max_seq_len, nh
    relative_position = memory_position - context_position
    relative_position_bucket = self._relative_position_bucket(
        relative_position,  # shape (query_length, key_length)
        bidirectional=(not self.is_decoder),
        num_buckets=self.relative_attention_num_buckets,
        max_distance=self.relative_attention_max_distance,
    )
    values = self.relative_attention_bias(
        relative_position_bucket
    )  # shape (query_length, key_length, num_heads)
    x = values.permute([2, 0, 1]).unsqueeze(
        0
    )  # shape (1, num_heads, query_length, key_length)
    return x

Functions: