Hello TensorFlow with Job API
This example uses TensorFlow/Keras and NVIDIA FLARE’s Job and Client APIs to train an MNIST classifier with federated averaging. The runnable source is in the hello-tf directory.
The setup has one server and two clients. In each round, every client evaluates
the received global model, trains it on its local half of MNIST, and sends its
updated layer weights and accuracy to the server. The server aggregates the
updates with FedAvg.
Running the Example
From the repository root, create an environment, install the example requirements, and run the job recipe:
$ cd examples/hello-world/hello-tf
$ python3 -m pip install -r requirements.txt
$ python3 job.py
Use --n_clients and --num_rounds to change the defaults:
$ python3 job.py --n_clients 3 --num_rounds 5
The recipe executes with SimEnv. Results
and logs are written under /tmp/nvflare/simulation by default.
Job Recipe
The example’s job.py constructs a TensorFlow
FedAvgRecipe with the
actual model object and client training script:
1
2import argparse
3
4from model import Net
5
6from nvflare.app_opt.tf.recipes.fedavg import FedAvgRecipe
7from nvflare.recipe import SimEnv, add_experiment_tracking
8
9
10def define_parser():
11 parser = argparse.ArgumentParser()
12 parser.add_argument("--n_clients", type=int, default=2)
13 parser.add_argument("--num_rounds", type=int, default=3)
14 return parser.parse_args()
15
16
17if __name__ == "__main__":
18 args = define_parser()
19 n_clients = args.n_clients
20 num_rounds = args.num_rounds
21 train_script = "client.py"
22
23 recipe = FedAvgRecipe(
24 name="hello-tf_fedavg",
25 num_rounds=num_rounds,
26 # Model can be specified as class instance or dict config:
27 model=Net(),
28 # Alternative: model={"class_path": "model.Net", "args": {}},
29 # For pre-trained weights: initial_ckpt="/server/path/to/model.h5",
30 min_clients=n_clients,
31 train_script=train_script,
32 )
33 add_experiment_tracking(recipe, tracking_type="tensorboard")
34
35 env = SimEnv(num_clients=n_clients)
36 run = recipe.execute(env=env)
37 print()
38 print("Result can be found in :", run.get_result())
39 print("Job Status is:", run.get_status())
40 print()
The important inputs are model=Net() and train_script="client.py".
add_experiment_tracking also adds TensorBoard event handling to the job.
Model
model.py defines the Keras Net class used to initialize and persist the
global model:
1
2from tensorflow.keras import layers, models
3
4
5class Net(models.Sequential):
6 def __init__(self, input_shape=(None, 28, 28)):
7 super().__init__()
8 self._input_shape = input_shape
9 self.add(layers.Flatten())
10 self.add(layers.Dense(128, activation="relu"))
11 self.add(layers.Dropout(0.2))
12 self.add(layers.Dense(10))
Client Training
client.py is ordinary TensorFlow training code with a small Client API loop:
1
2import tensorflow as tf
3from model import Net
4
5import nvflare.client as flare
6from nvflare.client.tracking import SummaryWriter
7
8WEIGHTS_PATH = "./tf_model.weights.h5"
9
10
11def main():
12 flare.init()
13 writer = SummaryWriter()
14
15 sys_info = flare.system_info()
16 print(f"system info is: {sys_info}", flush=True)
17
18 model = Net()
19 model.build(input_shape=(None, 28, 28))
20 model.compile(
21 optimizer="adam", loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=["accuracy"]
22 )
23 model.summary()
24
25 (train_images, train_labels), (
26 test_images,
27 test_labels,
28 ) = tf.keras.datasets.mnist.load_data()
29 train_images, test_images = (
30 train_images / 255.0,
31 test_images / 255.0,
32 )
33
34 # simulate separate datasets for each client by dividing MNIST dataset in half
35 client_name = sys_info["site_name"]
36 if client_name == "site-1":
37 train_images = train_images[: len(train_images) // 2]
38 train_labels = train_labels[: len(train_labels) // 2]
39 test_images = test_images[: len(test_images) // 2]
40 test_labels = test_labels[: len(test_labels) // 2]
41 elif client_name == "site-2":
42 train_images = train_images[len(train_images) // 2 :]
43 train_labels = train_labels[len(train_labels) // 2 :]
44 test_images = test_images[len(test_images) // 2 :]
45 test_labels = test_labels[len(test_labels) // 2 :]
46
47 while flare.is_running():
48 input_model = flare.receive()
49 print(f"current_round={input_model.current_round}")
50
51 sys_info = flare.system_info()
52 print(f"system info is: {sys_info}")
53
54 for k, v in input_model.params.items():
55 model.get_layer(k).set_weights(v)
56
57 _, test_global_acc = model.evaluate(test_images, test_labels, verbose=2)
58 print(
59 f"Accuracy of the received model on round {input_model.current_round} on the test images: {test_global_acc * 100} %"
60 )
61 writer.add_scalar(tag="local_acc", scalar=test_global_acc, global_step=input_model.current_round)
62
63 # training
64 model.fit(train_images, train_labels, epochs=1, validation_data=(test_images, test_labels))
65
66 print("Finished Training")
67
68 model.save_weights(WEIGHTS_PATH)
69
70 sys_info = flare.system_info()
71 print(f"system info is: {sys_info}", flush=True)
72 print(f"finished round: {input_model.current_round}", flush=True)
73
74 output_model = flare.FLModel(
75 params={layer.name: layer.get_weights() for layer in model.layers},
76 params_type="FULL",
77 metrics={"accuracy": test_global_acc},
78 current_round=input_model.current_round,
79 )
80
81 flare.send(output_model)
82
83
84if __name__ == "__main__":
85 main()
The essential Client API calls are:
flare.init()to initialize the trainer in the Client Job process.flare.receive()to receive the current global model.flare.send()to return updated weights and metrics.
Generated Application Configuration
When the recipe is exported, the server persistor refers to the model that the
example actually supplies in model.py:
{
"id": "persistor",
"path": "nvflare.app_opt.tf.model_persistor.TFModelPersistor",
"args": {
"model": {
"path": "model.Net",
"args": {}
}
}
}
The generated client configuration uses the unified
ClientAPIExecutor in in_process mode and points to the actual bundled
training file, client.py:
{
"tasks": ["*"],
"executor": {
"path": "nvflare.app_common.executors.client_api_executor.ClientAPIExecutor",
"args": {
"execution_mode": "in_process",
"task_script_path": "client.py",
"params_exchange_format": "keras_layer_weights",
"server_expected_format": "numpy"
}
}
}
The trainer runs inside the Client Job process in this example. Use Client API Attach Mode only when the trainer process is started and owned independently of NVFLARE.
Notes on Running with GPUs
TensorFlow may allocate most available GPU memory at startup. When simulating multiple clients on one host, enable memory growth and asynchronous allocation:
$ TF_FORCE_GPU_ALLOW_GROWTH=true TF_GPU_ALLOCATOR=cuda_malloc_async python3 job.py
For GPU environments, the NVIDIA TensorFlow container is recommended.