Based on this link I am trying to write a new way of FL algorithm. I train all clients and send the model parameters of all clients to the server, and the server will weight average only the model parameters of 30% of all clients during the aggregation process. As a criterion for selecting model parameters of 30% of clients, I want to do a weighted average by using weights_delta of 30% of clients with less loss_sum of clients.
The code below is a modified code for this link.
@tf.function
def client_update(model, dataset, server_message, client_optimizer):
model_weights = model.weights
initial_weights = server_message.model_weights
tff.utils.assign(model_weights, initial_weights)
num_examples = tf.constant(0, dtype=tf.int32)
loss_sum = tf.constant(0, dtype=tf.float32)
for batch in iter(dataset):
with tf.GradientTape() as tape:
outputs = model.forward_pass(batch)
grads = tape.gradient(outputs.loss, model_weights.trainable)
grads_and_vars = zip(grads, model_weights.trainable)
client_optimizer.apply_gradients(grads_and_vars)
batch_size = tf.shape(batch['x'])[0]
num_examples += batch_size
loss_sum += outputs.loss * tf.cast(batch_size, tf.float32)
weights_delta = tf.nest.map_structure(lambda a, b: a - b,
model_weights.trainable,
initial_weights.trainable)
client_weight = tf.cast(num_examples, tf.float32)
client_loss = loss_sum #add
return ClientOutput(weights_delta, client_weight, loss_sum / client_weight,client_loss)
There are the following attributes in client_output
weights_delta = attr.ib()
client_weight = attr.ib()
model_output = attr.ib()
client_loss = attr.ib()
After that, I made the client_output in the form of a sequence through
collected_output = tff.federated_collect(client_output) and round_model_delta = tff.federated_map(selecting_fn,(collected_output,weight_denom))in here .
@tff.federated_computation(federated_server_state_type,
federated_dataset_type)
def run_one_round(server_state, federated_dataset):
server_message = tff.federated_map(server_message_fn, server_state)
server_message_at_client = tff.federated_broadcast(server_message)
client_outputs = tff.federated_map(
client_update_fn, (federated_dataset, server_message_at_client))
weight_denom = client_outputs.client_weight
collected_output = tff.federated_collect(client_outputs) # add
round_model_delta = tff.federated_map(selecting_fn,(collected_output,weight_denom)) #add
server_state = tff.federated_map(server_update_fn,(server_state, round_model_delta))
round_loss_metric = tff.federated_mean(client_outputs.model_output, weight=weight_denom)
return server_state, round_loss_metric
Also, the following code is added here to implement the selecting_fn function.
@tff.tf_computation() # append
def selecting_fn(collected_output,weight_denom):
#TODO
return round_model_delta
I am not sure if it is correct to write the code in the above way.
I tried in various ways, but mainly TypeError: The value to be mapped must be a FederatedType or implicitly convertible to a FederatedType (got a <<model_weights=<trainable=<float32[5,5,1,32],float32[32] ,float32[5,5,32,64],float32[64],float32[3136,512],float32[512],float32[512,10],float32[10]>,non_trainable=<>>,optimizer_state= <int64>,round_num=int32>@SERVER,{<float32[5,5,1,32],float32[32],float32[5,5,32,64],float32[64],float32[3136,512 ],float32[512],float32[512,10],float32[10]>}@CLIENTS>) I get this error.
I wonder how the sequence type collected_output accesses each client's client_loss(= loss_sum) and sorts them, and also wonders what method to use when calculating the weighted average with weight_denom applied.