from torch.nn.functional import scaled_dot_product_attention as sdpa class TFViTSelfAttention(keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) self.num_attention_heads = ATTN_HEADS self.attention_head_size = int(HIDDEN_SIZE / ATTN_HEADS) self.all_head_size = ATTN_HEADS * self.attention_head_size self.sqrt_att_head_size = math.sqrt(self.attention_head_size) self.query = keras.layers.Dense(self.all_head_size, name="query") self.key = keras.layers.Dense(self.all_head_size, name="key") self.value = keras.layers.Dense(self.all_head_size, name="value") def transpose_for_scores(self, tensor, batch_size: int): tensor = keras.ops.reshape(tensor, (batch_size, -1, ATTN_HEADS, self.attention_head_size)) return keras.ops.transpose(tensor, [0, 2, 1, 3]) def call(self, hidden_states, training=False): bs = hidden_states.shape[0] mixed_query_layer = self.query(inputs=hidden_states) mixed_key_layer = self.key(inputs=hidden_states) mixed_value_layer = self.value(inputs=hidden_states) query_layer = self.transpose_for_scores(mixed_query_layer, bs) key_layer = self.transpose_for_scores(mixed_key_layer, bs) value_layer = self.transpose_for_scores(mixed_value_layer, bs) sdpa_output = sdpa(query_layer, key_layer, value_layer) attention_output = keras.ops.transpose(sdpa_output,[0,2,1,3]) attention_output = keras.ops.reshape(attention_output, (bs, -1, self.all_head_size)) return (attention_output,)