Preface
Recently, I was running experiments and found that the framework I had built previously was a bit uncomfortable to use, so I decided to refactor it. However, since the new framework isn’t ready yet, I first built a simple federated learning simulator based on mpi4py to run some basic experiments.
The code is hosted in this GitHub repository: https://github.com/yyyanbj/mpi-fedsim
Code Structure
The code structure is quite simple, mainly divided into five parts:
model.py: Model definition, including the model’s definition and initializationutils.py: Some utility functions, as well as dataset downloading and splittingserver.pyandclient.py: Definitions for the server and clients, including receiving and sending models, as well as model updatesmain_sync.py: The main program for synchronous federated learningalgor: Some federated learning algorithms, includingfedavg.pyconfig: Some configuration and experimental parameters.
The logging module uses loguru and tensorboard, while model creation and training are done with PyTorch. Although it’s quite simple, it still has “all the necessary organs,” hahaha.
Execution Logic
The logic mainly resides in the main function 1. First, if rank=0, it acts as the server; otherwise, it is a training process. The server initializes the model first, then waits for each process to upload the number of samples and the assigned client ID, recording them. Other processes first assign themselves a client ID based on pre-defined rules, then load the corresponding dataset based on that ID. Upon receiving the model, they begin training. After each process finishes training, it uploads the model to the server. Once the server receives it, it updates the global model and sends the model to all processes. Upon receiving the model, the processes update their local models and start the next round of training.
Since my experiments are conducted in a single-machine multi-GPU environment, I can assign specific GPUs during thread allocation. However, since CNNs are quite small and won’t exhaust VRAM, I didn’t implement this part. If needed, you only need to modify the mapping relationship between main_sync.py and client_id with gpu_id.
Since each process contains multiple clients but communicates only once when receiving model parameters, data transmission might be slightly faster.
Actually, the number of threads doesn’t really matter. The limit on threads stems from VRAM constraints, so as long as there is enough VRAM, the number of threads can be set very high. This approach simply aims to utilize as much VRAM as possible for acceleration within limited VRAM.
Usage
| |
Benchmark
Here, a two-layer convolutional CNN is used for MNIST classification; I haven’t run too many tests. In practice, with 10 clients, each having 6,000 samples, a local epoch of 1, and 50 rounds, using 11 processes takes about 1 minute to complete.
To be added
Improvements
Actually, the code is still quite rough, and there are many parts that can be improved. First, regarding the model receiving and sending part, for synchronous learning, using a broadcast approach would be better, as it improves efficiency and eliminates the need for 1-to-1 checks.

