I am trying to do transfer learning on my own dataset. I managed to do transfer learning with resnet and I got decent accuracy. But when I try VGG16 my accuracy stays the same all the time and loss is changing. I preprocessed my images with tf.keras.applications.vgg16.preprocess_input. Also, I tried to do transfer learning in a centralized manner and it works fine.
if base_model == "VGG16":
base_model = tf.keras.applications.vgg16.VGG16(
include_top=False,
weights="imagenet",
input_tensor=tf.keras.layers.Input(shape=(input_shape, input_shape, 3)),
pooling=None,
)
base_model.trainable = False
inputs = tf.keras.Input(shape=(input_shape, input_shape, 3))
x = base_model(inputs, training=False)
x = tf.keras.layers.GlobalAveragePooling2D()(x)
outputs = tf.keras.layers.Dense(num_classes, activation="softmax")(x)
model = tf.keras.Model(inputs, outputs)
return model
def create_FL_model():
"""create_FL_model_test _summary_
Returns:
tff.learning.Model: _description_
"""
keras_model = load_model(name, base_model)
return tff.learning.from_keras_model(
keras_model,
input_spec=input_spec.element_spec,
loss=tf.keras.losses.SparseCategoricalCrossentropy(),
metrics=[
tf.keras.metrics.SparseCategoricalAccuracy(),
],
)
# check again
if fed_alg == "FedAvg":
transfer_learning_iterative_process = (
tff.learning.build_federated_averaging_process(
create_FL_model,
client_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=0.02),
server_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=1.0),
)
)
keras_model = load_model(name, base_model)
state_transfer = transfer_learning_iterative_process.initialize()
state = tff.learning.state_with_new_model_weights(
state_transfer,
trainable_weights=[v.numpy() for v in keras_model.trainable_weights],
non_trainable_weights=[v.numpy() for v in keras_model.non_trainable_weights],
)
Training round:
for epoch in range(num_epochs):
client_data_train, client_data_valid = client_data.train_test_client_split(
client_data, num_test_clients=1, seed=12345
)
fed_valid_data = preprocess(
client_data_valid.create_tf_dataset_for_client(
client_data_valid.client_ids[0]
)
)
random_clients_ids = random.sample(client_data_train.client_ids, k=2)
federated_train_data = make_federated_data(
client_data_train, random_clients_ids
)
state, metrics = transfer_learning_iterative_process.next(
state, federated_train_data
)