import math import keras HIDDEN_SIZE = 768 IMG_SIZE = 224 PATCH_SIZE = 16 ATTN_HEADS = 12 NUM_LAYERS = 12 INTER_SZ = 4*HIDDEN_SIZE N_LABELS = 2 class TFViTEmbeddings(keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) self.patch_embeddings = TFViTPatchEmbeddings() num_patches = self.patch_embeddings.num_patches self.cls_token = self.add_weight((1, 1, HIDDEN_SIZE)) self.position_embeddings = self.add_weight((1, num_patches+1, HIDDEN_SIZE)) def call(self, pixel_values, training=False): bs, num_channels, height, width = pixel_values.shape embeddings = self.patch_embeddings(pixel_values, training=training) cls_tokens = keras.ops.repeat(self.cls_token, repeats=bs, axis=0) embeddings = keras.ops.concatenate((cls_tokens, embeddings), axis=1) embeddings = embeddings + self.position_embeddings return embeddings class TFViTPatchEmbeddings(keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) patch_size = (PATCH_SIZE, PATCH_SIZE) image_size = (IMG_SIZE, IMG_SIZE) num_patches = (image_size[1]//patch_size[1]) * \ (image_size[0]//patch_size[0]) self.patch_size = patch_size self.num_patches = num_patches self.projection = keras.layers.Conv2D( filters=HIDDEN_SIZE, kernel_size=patch_size, strides=patch_size, padding="valid", data_format="channels_last" ) def call(self, pixel_values, training=False): bs, num_channels, height, width = pixel_values.shape pixel_values = keras.ops.transpose(pixel_values, (0, 2, 3, 1)) projection = self.projection(pixel_values) num_patches = (width // self.patch_size[1]) * \ (height // self.patch_size[0]) embeddings = keras.ops.reshape(projection, (bs, num_patches, -1)) return embeddings 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) key_layer_T = keras.ops.transpose(key_layer, [0,1,3,2]) attention_scores = keras.ops.matmul(query_layer, key_layer_T) dk = keras.ops.cast(self.sqrt_att_head_size, dtype=attention_scores.dtype) attention_scores = keras.ops.divide(attention_scores, dk) attention_probs = keras.ops.softmax(attention_scores+1e-9, axis=-1) attention_output = keras.ops.matmul(attention_probs, value_layer) attention_output = keras.ops.transpose(attention_output,[0,2,1,3]) attention_output = keras.ops.reshape(attention_output, (bs, -1, self.all_head_size)) return (attention_output,) class TFViTSelfOutput(keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) self.dense = keras.layers.Dense(HIDDEN_SIZE) def call(self, hidden_states, input_tensor, training = False): return self.dense(inputs=hidden_states) class TFViTAttention(keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) self.self_attention = TFViTSelfAttention() self.dense_output = TFViTSelfOutput() def call(self, input_tensor, training = False): self_outputs = self.self_attention( hidden_states=input_tensor, training=training ) attention_output = self.dense_output( hidden_states=self_outputs[0], input_tensor=input_tensor, training=training ) return (attention_output,) class TFViTIntermediate(keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) self.dense = keras.layers.Dense(INTER_SZ) self.intermediate_act_fn = keras.activations.gelu def call(self, hidden_states): hidden_states = self.dense(hidden_states) hidden_states = self.intermediate_act_fn(hidden_states) return hidden_states class TFViTOutput(keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) self.dense = keras.layers.Dense(HIDDEN_SIZE) def call(self, hidden_states, input_tensor, training: bool = False): hidden_states = self.dense(inputs=hidden_states) hidden_states = hidden_states + input_tensor return hidden_states class TFViTLayer(keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) self.attention = TFViTAttention() self.intermediate = TFViTIntermediate() self.vit_output = TFViTOutput() self.layernorm_before = keras.layers.LayerNormalization( epsilon=1e-12 ) self.layernorm_after = keras.layers.LayerNormalization( epsilon=1e-12 ) def call(self, hidden_states, training=False): attention_outputs = self.attention( input_tensor=self.layernorm_before(inputs=hidden_states), training=training, ) attention_output = attention_outputs[0] hidden_states = attention_output + hidden_states layer_output = self.layernorm_after(hidden_states) intermediate_output = self.intermediate(layer_output) layer_output = self.vit_output( hidden_states=intermediate_output, input_tensor=hidden_states, training=training ) outputs = (layer_output,) return outputs class TFViTEncoder(keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) self.layer = [TFViTLayer(name=f"layer_{i}") for i in range(NUM_LAYERS)] def call(self, hidden_states, training=False): for i, layer_module in enumerate(self.layer): layer_outputs = layer_module( hidden_states=hidden_states, training=training, ) hidden_states = layer_outputs[0] return tuple([hidden_states]) class TFViTMainLayer(keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) self.embeddings = TFViTEmbeddings() self.encoder = TFViTEncoder() self.layernorm = keras.layers.LayerNormalization(epsilon=1e-12) def call(self, pixel_values, training=False): embedding_output = self.embeddings( pixel_values=pixel_values, training=training, ) encoder_outputs = self.encoder( hidden_states=embedding_output, training=training, ) sequence_output = encoder_outputs[0] sequence_output = self.layernorm(inputs=sequence_output) return (sequence_output,) class TFViTForImageClassification(keras.Model): def __init__(self, *inputs, **kwargs): super().__init__(*inputs, **kwargs) self.vit = TFViTMainLayer() self.classifier = keras.layers.Dense(N_LABELS) def call(self, pixel_values, training=False): outputs = self.vit(pixel_values, training=training) sequence_output = outputs[0] logits = self.classifier(inputs=sequence_output[:, 0, :]) return (logits,)