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.

Previous Versions of Hello TensorFlow