class ResNet(tf.keras.Model): def __init__(self, num_classes, block_counts,blocktype, initial_filters=64): super(ResNet, self).__init__() self.conv1 = tf.keras.layers.Conv2D(initial_filters, 7, strides=2, padding="same",kernel_regularizer=tf.keras.regularizers.L2(0.001)) self.bn1 = tf.keras.layers.BatchNormalization() self.relu = tf.keras.layers.ReLU() self.pool = tf.keras.layers.MaxPooling2D(pool_size=3, strides=2, padding="same") # Replace list with Sequential for residual blocks self.residual_blocks = tf.keras.Sequential(name="residual_blocks") filters = initial_filters for i, count in enumerate(block_counts): for j in range(count): strides = 2 if j == 0 and i > 0 else 1 # Downsample at the start of a new stage self.residual_blocks.add( blocktype(filters, strides=strides, downsample=(strides == 2)) ) filters *= 2 self.global_pool = tf.keras.layers.GlobalAveragePooling2D() self.dropout = tf.keras.layers.Dropout(0.3) self.fc = tf.keras.layers.Dense(num_classes, activation="softmax") def call(self, inputs, training=False): x = self.conv1(inputs) x = self.bn1(x, training=training) x = self.relu(x) x = self.pool(x) x = self.residual_blocks(x, training=training) # Pass through all residual blocks x = self.global_pool(x) x = self.dropout(x,training=training) return self.fc(x) __ __