BrainTorrent: A Peer-to-Peer Environment for Decentralized Federated Learning

Abhijit Guha Roy, Shayan Siddiqui, Sebastian Pölsterl, Nassir Navab, Christian Wachinger

Introduction

Training deep neural networks (DNNs) effectively requires access to abundant annotated data. This is a common concern in medical applications, where annotations are both, expensive and time-consuming. For instance, manual labeling of a single 3D brain MRI scan can take up to a week by a trained neuroanatomist . It is therefore challenging to reach sample sizes with in-house curated datasets that enable an effective application of deep learning. Pooling data across medical centers could alleviate the limited data problem. However, data sharing is restricted due to ethical and legal regulations, preventing the aggregation of medical data from multiple medical centers. This leaves us with a scenario, where each center might have too limited data to effectively train DNNs but the combination of data across the centers that work on the same problem is not possible.

In such a decentralized environment, where data is distributed across centers, training a common DNN is challenging. One naïve approach would be to train in a sequentially incremental fashion . Such an approach trains the network at a given center and then passes the DNN weights to the next center, where it is fine-tuned on new data, and so on. A common problem encountered with such an approach is catastrophic forgetting , i.e., at every stage of fine-tuning, the new data overwrites the knowledge acquired from the previous training, thus deteriorating generalizability. Also, certain centers can have very limited data, which can risk overfitting the network severely.

Recently, a learning strategy was introduced to address challenges in training DNNs in such decentralized environments called Federated Learning (FL) . The environment consists of a central server, which is connected to all the centers coordinating the overall process. The main motivation of this framework is to provide assistance to mobile device users. In such a scenario, where users can easily scale into the millions, it is very convenient to have a central server body. Also, one of the main aims of FL is to minimize communication costs. In the scenario of collaborative learning within a community, like medical centers, the motivations are a bit different. Firstly, in contrast to millions of clients in FL, the number of medical centers forming a community is much lower (in the order of 10s). Secondly, each center can be expected to have a strong communication infrastructure, so that communication cycles are not a big bottleneck. Thirdly, it is difficult to have a central trusted server body in such a setting, rather every center may want to coordinate with the rest directly. Fourthly, if the whole community is dependent on the server and a fault occurs at the server, the whole system is non-operational, which is undesirable in a medical setting.

In this paper, we propose for the first time a server-less, peer-to-peer approach to federated learning where clients communicate directly among themselves. We term this decentralized environment The BrainTorrent. The design is motivated to fulfill the above mentioned requirements for a group of medical centers to collaborate. Absence of a central server body not only makes our environment resistant to failure but also precludes the need for a body everyone trusts. Further, any client at any point can initiate an update process very dynamically. As number of communication round is not a bottleneck for medical centers, they can interact more frequently. Due to this high frequency of interaction and update process per client, the models in BrainTorrent converges faster and reaches an accuracy similar to a model trained with pooling the data from all the clients.

The main contributions of the paper: (i) we introduce BrainTorrent a peer-to-peer, decentralized environment where multiple medical centers can collaborate and benefit from each other without sharing data among them, (ii) we propose a new training strategy of DNNs within this environment in a federated learning fashion that does not rely upon a central server-body to coordinate the process, and (iii) we demonstrate exemplar cases where different centers in the environment can have data with different age ranges and non-uniform data distribution, where BrainTorrent outperforms traditional, server-based federated learning.

Prior work: Federated learning (FL) was proposed by Mcmahan et. al. for training models from decentralized data distributed across different clients (mainly mobile users), preserving data privacy. Other works build on top of FL by improving communication efficiency , improving system scalability and improving encryption for better privacy . The traditional FL concept has recently been applied to medical image analysis for 2-class segmentation of brain tumors , demonstrating its feasibility. In contrast to that, we introduce a new peer-to-peer FL approach, and tackle the more challenging task of whole-brain segmentation with 2020 classes with severe class-imbalance. To the best of our knowledge, this is the first application of FL to whole-brain segmentation.

Method

Let us consider an environment with NN centers {C1,…,CN}\{C_{1},\dots,C_{N}\}, where each center CiC_{i} has training data Di={(x1,y1),…,(xai,yai)}\mathcal{D}_{i}=\{(x_{1},y_{1}),\dots,(x_{a_{i}},y_{a_{i}})\} with aia_{i} labeled samples. In a general setting, data from all centers are pooled together (a=∑iaia=\sum_{i}a_{i}) at a common server, where a common model is learned that is distributed across all the centers for usage. In a medical setting, data in each center is very sensitive containing patient specific information, which cannot be shared across centers or with a central server S\mathcal{S}.

The process of distributed training in traditional FL with server (FLS) is illustrated in Fig. 1 (a). Here, in a single round of training, all NN clients {Ci}i=1N\{C_{i}\}_{i=1}^{N} start with training their respective networks in parallel only for a few iterations (not till convergence). Let the weight parameter list for all clients be indicated by {W1,…,WN}\{\mathbf{W}^{1},\dots,\mathbf{W}^{N}\}. Next, all the clients send these partially trained parameters to the central server S\mathcal{S}, which aggregates them by weighted averaging WS=∑iaiaWi\mathbf{W}^{\mathcal{S}}=\sum_{i}\frac{a_{i}}{a}\mathbf{W}^{i}. The multiplicative factor is computed as the fraction of the total data belonging to a client. The rationale is to emphasize clients with more training data. Finally, the aggregated model WS\mathbf{W}^{\mathcal{S}} is distributed back to all clients for further training. We refer to for a more detailed description of the implementation.

Several rounds are executed until all client models converge. At the end of the training process, each client has its own personalized model Wi\mathbf{W}^{i}, fine-tuned to its local data, and the server model WS\mathbf{W}^{\mathcal{S}}, which is more generic to new unseen data. The central server S\mathcal{S} has the vital role of coordinating the aggregation and re-distributing the weights across clients. When a new client is added to the environment, it receives the server model WS\mathbf{W}^{\mathcal{S}} to start off.

2 BrainTorrent: Server-less Peer-to-Peer Federated Learning

Locally train each client in parallel for a few iterations using local dataset.

A random client CiC_{i} from the environment initiates the training process. It sends out a ‘ping_request’ to the rest of the clients to get their latest model versions to generate vnew\mathbf{v}_{\text{new}}. vold\mathbf{v}_{\text{old}} is initiated with client’s v\mathbf{v}.

All clients CjC_{j} with updates, i.e., voldj<vnewjv^{j}_{\text{old}}<v^{j}_{\text{new}}, send their weights Wj\mathbf{W}^{j} and the training sample size aja_{j} to CiC_{i}.

This subset of models is merged with CiC_{i}’s current model to a single model by weighted averaging. Then return to Step 1.

This comprises a single round of training. It must be noted that the definition of a round here is different from FLS. In a single round of FLS, all clients are updated once by fine-tuning, whereas in our framework only a single random client is updated. So, the number of updates per client in RR rounds of FLS is equivalent to R×NR\times N rounds of our framework. The steps are presented in Algorithm 1.

Experimental Settings

To demonstrate the effectiveness of BrainTorrent, we choose the challenging task of whole-brain segmentation of MRI T1 scans. We use the Multi-Atlas Labelling Challenge (MALC) dataset for our experiments. The dataset consists of 3030 annotated whole-brain MRI T1 scans from different patients, out of which we always use 2020 scans for training and the remaining 1010 for testing. Manual annotations were provided by Neuromorphometrics Inc. As a segmentation network, we decided to use the QuickNAT architecture , which demonstrated state-of-the-art performance for whole-brain segmentation. We combined the left and right brain structures in one class and all the cortical parcellations in a single cortex class, thus reducing the task to a 20-class segmentation problem. During fine-tuning at each client center, we fix the number of epochs to 22, without risking any client-specific overfitting. Initially, all clients had a learning rate of 0.0010.001, which was reduced by a factor of 0.50.5 after every 44 update rounds. We use Adam for optimization. We explore two experimental settings detailed below, where we compare our proposed BrainTorrent and FLS.

In this experiment, we randomly distributed the 2020 training scans among the clients in a uniform fashion, i.e., each client receives the same number of scans. Here, we also conduct experiments by varying the number of clients by {5,7,10,20}\{5,7,10,20\} in the environment.

We investigate the FL performance as the number of clients increases in the environment. Since the number of training scans are fixed to 2020, scans per client reduces with increasing clients. We vary the number of clients between {5,7,10,20}\{5,7,10,20\}, which results in {4,3,2,1}\{4,3,2,1\} training scans per client, respectively. For the setting with 77 clients, one client had 22 scans whereas the rest had 33 scans. Under this setting, we compare FL and BrainTorrent and present the results in Tab. 1. We compare the average Dice score across all clients for FLS and BrainTorrent for all the configurations on the 1010 test scans. Also, we compare their aggregated model, i.e., the model which would be provided to a new client when it first joins the environment. For FL, this is the server model, whereas for BrainTorrent, we create a model by averaging the model weights of all the clients in the environment. As an upper bound model, we trained a model by pooling all the data across the clients termed as ‘pooled model’. We observe that irrespective of the number of clients, BrainTorrent outperforms FLS for both, average Dice score over clients and Dice score for aggregated model. Also, as the number of clients increases (and therefore number of scans per client decreases), the performance degrades. This drop in performance is only marginal at the beginning up to 1010 clients, and drops by a huge margin when each center has only 11 annotated scan, simulating an extremely limited data scenario. For the aggregated model, we observe that BrainTorrent outperforms FL by 1−2%1-2\% Dice points. Also, we observe that BrainTorrent achieves the same level of segmentation accuracy that would be reached by the ‘pooled model’, which is striking given the constraints. This performance is sustained for number of clients 5 to 7 with only 4-3 training scans per center.

Analysis of client-wise performance

We take a closer look at the segmentation performance of the client-specific models. We select the configuration with 1010 clients, where every client has 22 annotated training scans. We report the performance of each client model for both BrainTorrent and FLS in Tab. 2. Also, as lower bound analysis, we train client-specific models with only 22 scans, referred to as ‘only client’ models and report their performance on the same validation set. First, we observe that both FLS and BrainTorrent outperform the ‘only client’ model by an average of 24%24\% and 26%26\% Dice points, substantiating the immense effectiveness of the federated learning approach. Further, BrainTorrent achieves an average 2%2\% higher Dice score than FLS, where 77 out of 1010 client models performed better in BrainTorrent than in FLS. This reaffirms our claim that BrainTorrent does not only result in a stronger aggregated model but also in more robust client-level personalized models.

2 Experiment 2

In this experiment, we distribute the 2020 training scans across 55 clients, where each client has scans for a specific, non-overlapping age range, see Tab. 3. This experiment simulates the scenario that each client has data with unique characteristics. In addition, it also provides a scenario for non-uniform data distribution, where the number of training scans differs among clients, yielding a realistic clinical use-case.

Tab. 4 reports the results for FLS and BrainTorrent. We observe that under such an uneven distribution, the aggregated model of BrainTorrent achieves the performance of ‘pooled model’, whereas FLS had a performance 3%3\% Dice points below that. Comparing average Dice scores across clients, BrainTorrent outperforms FLS by a margin of 7%7\% Dice points. This demonstrates that in a scenario of non-uniform data distribution, performance of BrainTorrent is unaffected, whereas FLS performance degrades. Also, it must be noted that the performance of C3C_{3} and C4C_{4}, which have only 2 and 1 annotated scans, respectively, is comparatively low for FLS. One possible cause can be slight overfitting. In contrast, these clients perform very well in the BrainTorrent framework.

Conclusion

In this paper, we introduced BrainTorrent, a server-less peer-to-peer federated learning environment for decentralized training. In contrast to traditional FL with server, our framework does not rely on a central server body for orchestrating the training process. We presented a proof-of-concept study tackling the challenging task of whole-brain segmentation, training a complex fully convolutional neural network in a decentralized fashion. We demonstrated in our experiments that BrainTorrent achieves a better performance than FLS under different experimental settings. The margin extends up to 7%7\% Dice points in scenarios where clients have unequal numbers of training scans. Overall, BrainTorrent does not only resolve the issue relying on a central server but also enables more robust training of clients through highly dynamic updates, reaching performance similar to a model trained on data pooled across clients. Although we focused on image segmentation, our proposed method is generic and can be used for training any machine learning model.

References