import tensorflow as tf from tensorflow import keras from tensorflow.keras import layers class CNN_Encoder(keras.Model): def __init__(self, embedding_dim=256): super(CNN_Encoder, self).__init__() self.conv1 = layers.Conv2D( filters=32, kernel_size=3, strides=2, padding='same', name='conv_block1' ) self.bn1 = layers.BatchNormalization(name='bn1') self.leaky1 = layers.LeakyReLU(alpha=0.2, name='leaky1') self.pool1 = layers.MaxPool2D(pool_size=2, strides=2, padding='same', name='pool1') self.conv2 = layers.Conv2D( filters=64, kernel_size=3, strides=2, padding='same', name='conv_block2' ) self.bn2 = layers.BatchNormalization(name='bn2') self.leaky2 = layers.LeakyReLU(alpha=0.2, name='leaky2') self.pool2 = layers.MaxPool2D(pool_size=2, strides=2, padding='same', name='pool2') self.conv3 = layers.Conv2D( filters=128, kernel_size=3, strides=2, padding='same', name='conv_block3' ) self.bn3 = layers.BatchNormalization(name='bn3') self.leaky3 = layers.LeakyReLU(alpha=0.2, name='leaky3') self.pool3 = layers.MaxPool2D(pool_size=2, strides=2, padding='same', name='pool3') self.conv4 = layers.Conv2D( filters=256, kernel_size=3, strides=2, padding='same', name='conv_block4' ) self.bn4 = layers.BatchNormalization(name='bn4') self.leaky4 = layers.LeakyReLU(alpha=0.2, name='leaky4') self.reshape = layers.Reshape((-1, 256), name='reshape_features') self.fc = layers.Dense(embedding_dim, activation='relu', name='fc_projection') self.embedding_dim = embedding_dim def call(self, x, training=False): x = self.conv1(x) x = self.bn1(x, training=training) x = self.leaky1(x) x = self.pool1(x) x = self.conv2(x) x = self.bn2(x, training=training) x = self.leaky2(x) x = self.pool2(x) x = self.conv3(x) x = self.bn3(x, training=training) x = self.leaky3(x) x = self.pool3(x) x = self.conv4(x) x = self.bn4(x, training=training) x = self.leaky4(x) x = self.reshape(x) features = self.fc(x) return features class BahdanauAttention(keras.layers.Layer): def __init__(self, units): super(BahdanauAttention, self).__init__() self.W1 = layers.Dense(units, name='attention_W1') self.W2 = layers.Dense(units, name='attention_W2') self.V = layers.Dense(1, name='attention_V') self.units = units def call(self, features, hidden): hidden_with_time_axis = tf.expand_dims(hidden, 1) score = tf.nn.tanh(self.W1(features) + self.W2(hidden_with_time_axis)) attention_weights = self.V(score) attention_weights = tf.nn.softmax(attention_weights, axis=1) context_vector = attention_weights * features context_vector = tf.reduce_sum(context_vector, axis=1) attention_weights = tf.squeeze(attention_weights, axis=-1) return context_vector, attention_weights class RNN_Decoder(keras.Model): def __init__(self, embedding_dim=256, units=512, vocab_size=10000, rnn_type='lstm'): super(RNN_Decoder, self).__init__() self.units = units self.embedding_dim = embedding_dim self.vocab_size = vocab_size self.rnn_type = rnn_type.lower() self.embedding = layers.Embedding(vocab_size, embedding_dim, name='word_embedding') self.attention = BahdanauAttention(self.units) if self.rnn_type == 'lstm': self.rnn = layers.LSTM( self.units, return_sequences=True, return_state=True, recurrent_initializer='glorot_uniform', name='lstm_decoder' ) else: self.rnn = layers.GRU( self.units, return_sequences=True, return_state=True, recurrent_initializer='glorot_uniform', name='gru_decoder' ) self.fc1 = layers.Dense(self.units, activation='relu', name='fc1') self.fc2 = layers.Dense(vocab_size, name='fc2_output') def call(self, x, features, hidden, carry=None, training=False): context_vector, attention_weights = self.attention(features, hidden) x = self.embedding(x) context_vector_expanded = tf.expand_dims(context_vector, 1) x = tf.concat([context_vector_expanded, x], axis=-1) if self.rnn_type == 'lstm': if carry is None: carry = tf.zeros_like(hidden) output, state_h, state_c = self.rnn(x, initial_state=[hidden, carry], training=training) hidden = state_h carry = state_c else: output, state = self.rnn(x, initial_state=hidden, training=training) hidden = state carry = None x = self.fc1(output) x = tf.reshape(x, (-1, x.shape[2])) predictions = self.fc2(x) return predictions, hidden, carry, attention_weights def reset_state(self, batch_size): hidden = tf.zeros((batch_size, self.units)) if self.rnn_type == 'lstm': carry = tf.zeros((batch_size, self.units)) return [hidden, carry] else: return [hidden] if __name__ == "__main__": print("Testing Custom CNN Encoder-Decoder Model...") EMBEDDING_DIM = 256 UNITS = 512 VOCAB_SIZE = 10000 BATCH_SIZE = 4 IMG_SIZE = 299 encoder = CNN_Encoder(embedding_dim=EMBEDDING_DIM) print(f"\nEncoder created with embedding_dim={EMBEDDING_DIM}") decoder = RNN_Decoder( embedding_dim=EMBEDDING_DIM, units=UNITS, vocab_size=VOCAB_SIZE, rnn_type='lstm' ) print(f"Decoder created with units={UNITS}, vocab_size={VOCAB_SIZE}") print("\nTesting forward pass...") dummy_img = tf.random.normal((BATCH_SIZE, IMG_SIZE, IMG_SIZE, 3)) print(f"Input image shape: {dummy_img.shape}") features = encoder(dummy_img, training=False) print(f"Encoder output (features) shape: {features.shape}") decoder_states = decoder.reset_state(batch_size=BATCH_SIZE) hidden = decoder_states[0] carry = decoder_states[1] if len(decoder_states) > 1 else None print(f"Decoder hidden state shape: {hidden.shape}") dummy_token = tf.constant([[1], [2], [3], [4]]) predictions, new_hidden, new_carry, attn_weights = decoder( dummy_token, features, hidden, carry, training=False ) print(f"Decoder predictions shape: {predictions.shape}") print(f"Decoder attention weights shape: {attn_weights.shape}") print(f"New hidden state shape: {new_hidden.shape}") print("\n✅ Model test completed successfully!")