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",
)