diff --git a/.devcontainer/Dockerfile b/.devcontainer/Dockerfile
index 2cf82414df56..016c12af2426 100644
--- a/.devcontainer/Dockerfile
+++ b/.devcontainer/Dockerfile
@@ -29,8 +29,8 @@ RUN apt install -y curl wget gnupg python3 python-is-python3 python3-pip git \
build-essential tmux vim
RUN python -m pip install \
- pip==23.1.2 \
- setuptools==68.0.0 \
+ pip==23.3.1 \
+ setuptools==68.2.2 \
poetry==1.5.1
USER $USERNAME
diff --git a/.github/actions/bootstrap/action.yml b/.github/actions/bootstrap/action.yml
index 3865cad1def6..584ae2634d9e 100644
--- a/.github/actions/bootstrap/action.yml
+++ b/.github/actions/bootstrap/action.yml
@@ -6,10 +6,10 @@ inputs:
default: 3.8
pip-version:
description: "Version of pip to be installed using pip"
- default: 23.1.2
+ default: 23.3.1
setuptools-version:
description: "Version of setuptools to be installed using pip"
- default: 68.0.0
+ default: 68.2.2
poetry-version:
description: "Version of poetry to be installed using pip"
default: 1.5.1
diff --git a/.github/workflows/cpp.yml b/.github/workflows/cpp.yml
index 16cd672ef034..35fe9813329e 100644
--- a/.github/workflows/cpp.yml
+++ b/.github/workflows/cpp.yml
@@ -35,9 +35,14 @@ jobs:
sudo apt-get update
sudo apt-get install -y clang-format cmake g++ clang-tidy cppcheck
- - name: Check Formatting
+ - name: Check source Formatting
run: |
- find src/cc/flwr -name '*.cc' -or -name '*.h' | xargs clang-format -i
+ find src/cc/flwr/src -name '*.cc' | xargs clang-format -i
+ git diff --exit-code
+
+ - name: Check header Formatting
+ run: |
+ find src/cc/flwr/include -name '*.h' -not -path "src/cc/flwr/include/flwr/*" | xargs clang-format -i
git diff --exit-code
- name: Build
diff --git a/README.md b/README.md
index efed9b0e477e..002d16066e78 100644
--- a/README.md
+++ b/README.md
@@ -23,22 +23,21 @@
Flower (`flwr`) is a framework for building federated learning systems. The
design of Flower is based on a few guiding principles:
-* **Customizable**: Federated learning systems vary wildly from one use case to
+- **Customizable**: Federated learning systems vary wildly from one use case to
another. Flower allows for a wide range of different configurations depending
on the needs of each individual use case.
-* **Extendable**: Flower originated from a research project at the University of
+- **Extendable**: Flower originated from a research project at the University of
Oxford, so it was built with AI research in mind. Many components can be
extended and overridden to build new state-of-the-art systems.
-* **Framework-agnostic**: Different machine learning frameworks have different
+- **Framework-agnostic**: Different machine learning frameworks have different
strengths. Flower can be used with any machine learning framework, for
example, [PyTorch](https://pytorch.org),
- [TensorFlow](https://tensorflow.org), [Hugging Face Transformers](https://huggingface.co/), [PyTorch Lightning](https://pytorchlightning.ai/), [MXNet](https://mxnet.apache.org/), [scikit-learn](https://scikit-learn.org/), [JAX](https://jax.readthedocs.io/), [TFLite](https://tensorflow.org/lite/), [fastai](https://www.fast.ai/), [Pandas](https://pandas.pydata.org/
-) for federated analytics, or even raw [NumPy](https://numpy.org/)
+ [TensorFlow](https://tensorflow.org), [Hugging Face Transformers](https://huggingface.co/), [PyTorch Lightning](https://pytorchlightning.ai/), [MXNet](https://mxnet.apache.org/), [scikit-learn](https://scikit-learn.org/), [JAX](https://jax.readthedocs.io/), [TFLite](https://tensorflow.org/lite/), [fastai](https://www.fast.ai/), [Pandas](https://pandas.pydata.org/) for federated analytics, or even raw [NumPy](https://numpy.org/)
for users who enjoy computing gradients by hand.
-* **Understandable**: Flower is written with maintainability in mind. The
+- **Understandable**: Flower is written with maintainability in mind. The
community is encouraged to both read and contribute to the codebase.
Meet the Flower community on [flower.dev](https://flower.dev)!
@@ -58,11 +57,11 @@ Flower's goal is to make federated learning accessible to everyone. This series
2. **Using Strategies in Federated Learning**
[![Open in Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/adap/flower/blob/main/doc/source/tutorial-use-a-federated-learning-strategy-pytorch.ipynb) (or open the [Jupyter Notebook](https://github.com/adap/flower/blob/main/doc/source/tutorial-use-a-federated-learning-strategy-pytorch.ipynb))
-
+
3. **Building Strategies for Federated Learning**
[![Open in Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/adap/flower/blob/main/doc/source/tutorial-series-use-a-federated-learning-strategy-pytorch.ipynb) (or open the [Jupyter Notebook](https://github.com/adap/flower/blob/main/doc/source/tutorial-series-use-a-federated-learning-strategy-pytorch.ipynb))
-
+
4. **Custom Clients for Federated Learning**
[![Open in Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/adap/flower/blob/main/doc/source/tutorial-series-customize-the-client-pytorch.ipynb) (or open the [Jupyter Notebook](https://github.com/adap/flower/blob/main/doc/source/tutorial-series-customize-the-client-pytorch.ipynb))
@@ -73,39 +72,39 @@ Stay tuned, more tutorials are coming soon. Topics include **Privacy and Securit
[![Open in Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/adap/flower/blob/main/examples/flower-in-30-minutes/tutorial.ipynb) (or open the [Jupyter Notebook](https://github.com/adap/flower/blob/main/examples/flower-in-30-minutes/tutorial.ipynb))
-
## Documentation
[Flower Docs](https://flower.dev/docs):
-* [Installation](https://flower.dev/docs/framework/how-to-install-flower.html)
-* [Quickstart (TensorFlow)](https://flower.dev/docs/framework/tutorial-quickstart-tensorflow.html)
-* [Quickstart (PyTorch)](https://flower.dev/docs/framework/tutorial-quickstart-pytorch.html)
-* [Quickstart (Hugging Face)](https://flower.dev/docs/framework/tutorial-quickstart-huggingface.html)
-* [Quickstart (PyTorch Lightning [code example])](https://flower.dev/docs/framework/tutorial-quickstart-pytorch-lightning.html)
-* [Quickstart (MXNet)](https://flower.dev/docs/framework/example-mxnet-walk-through.html)
-* [Quickstart (Pandas)](https://flower.dev/docs/framework/tutorial-quickstart-pandas.html)
-* [Quickstart (fastai)](https://flower.dev/docs/framework/tutorial-quickstart-fastai.html)
-* [Quickstart (JAX)](https://flower.dev/docs/framework/tutorial-quickstart-jax.html)
-* [Quickstart (scikit-learn)](https://flower.dev/docs/framework/tutorial-quickstart-scikitlearn.html)
-* [Quickstart (Android [TFLite])](https://flower.dev/docs/framework/tutorial-quickstart-android.html)
-* [Quickstart (iOS [CoreML])](https://flower.dev/docs/framework/tutorial-quickstart-ios.html)
+
+- [Installation](https://flower.dev/docs/framework/how-to-install-flower.html)
+- [Quickstart (TensorFlow)](https://flower.dev/docs/framework/tutorial-quickstart-tensorflow.html)
+- [Quickstart (PyTorch)](https://flower.dev/docs/framework/tutorial-quickstart-pytorch.html)
+- [Quickstart (Hugging Face)](https://flower.dev/docs/framework/tutorial-quickstart-huggingface.html)
+- [Quickstart (PyTorch Lightning [code example])](https://flower.dev/docs/framework/tutorial-quickstart-pytorch-lightning.html)
+- [Quickstart (MXNet)](https://flower.dev/docs/framework/example-mxnet-walk-through.html)
+- [Quickstart (Pandas)](https://flower.dev/docs/framework/tutorial-quickstart-pandas.html)
+- [Quickstart (fastai)](https://flower.dev/docs/framework/tutorial-quickstart-fastai.html)
+- [Quickstart (JAX)](https://flower.dev/docs/framework/tutorial-quickstart-jax.html)
+- [Quickstart (scikit-learn)](https://flower.dev/docs/framework/tutorial-quickstart-scikitlearn.html)
+- [Quickstart (Android [TFLite])](https://flower.dev/docs/framework/tutorial-quickstart-android.html)
+- [Quickstart (iOS [CoreML])](https://flower.dev/docs/framework/tutorial-quickstart-ios.html)
## Flower Baselines
Flower Baselines is a collection of community-contributed experiments that reproduce the experiments performed in popular federated learning publications. Researchers can build on Flower Baselines to quickly evaluate new ideas:
-* [FedAvg](https://arxiv.org/abs/1602.05629):
- * [MNIST](https://github.com/adap/flower/tree/main/baselines/flwr_baselines/flwr_baselines/publications/fedavg_mnist)
-* [FedProx](https://arxiv.org/abs/1812.06127):
- * [MNIST](https://github.com/adap/flower/tree/main/baselines/fedprox/)
-* [FedBN: Federated Learning on non-IID Features via Local Batch Normalization](https://arxiv.org/abs/2102.07623):
- * [Convergence Rate](https://github.com/adap/flower/tree/main/baselines/flwr_baselines/flwr_baselines/publications/fedbn/convergence_rate)
-* [Adaptive Federated Optimization](https://arxiv.org/abs/2003.00295):
- * [CIFAR-10/100](https://github.com/adap/flower/tree/main/baselines/flwr_baselines/flwr_baselines/publications/adaptive_federated_optimization)
+- [FedAvg](https://arxiv.org/abs/1602.05629):
+ - [MNIST](https://github.com/adap/flower/tree/main/baselines/flwr_baselines/flwr_baselines/publications/fedavg_mnist)
+- [FedProx](https://arxiv.org/abs/1812.06127):
+ - [MNIST](https://github.com/adap/flower/tree/main/baselines/fedprox/)
+- [FedBN: Federated Learning on non-IID Features via Local Batch Normalization](https://arxiv.org/abs/2102.07623):
+ - [Convergence Rate](https://github.com/adap/flower/tree/main/baselines/flwr_baselines/flwr_baselines/publications/fedbn/convergence_rate)
+- [Adaptive Federated Optimization](https://arxiv.org/abs/2003.00295):
+ - [CIFAR-10/100](https://github.com/adap/flower/tree/main/baselines/flwr_baselines/flwr_baselines/publications/adaptive_federated_optimization)
-Check the Flower documentation to learn more: [Using Baselines](https://flower.dev/docs/baselines/using-baselines.html)
+Check the Flower documentation to learn more: [Using Baselines](https://flower.dev/docs/baselines/how-to-use-baselines.html)
-The Flower community loves contributions! Make your work more visible and enable others to build on it by contributing it as a baseline: [Contributing Baselines](https://flower.dev/docs/baselines/contributing-baselines.html)
+The Flower community loves contributions! Make your work more visible and enable others to build on it by contributing it as a baseline: [Contributing Baselines](https://flower.dev/docs/baselines/how-to-contribute-baselines.html)
## Flower Usage Examples
@@ -113,26 +112,26 @@ Several code examples show different usage scenarios of Flower (in combination w
Quickstart examples:
-* [Quickstart (TensorFlow)](https://github.com/adap/flower/tree/main/examples/quickstart-tensorflow)
-* [Quickstart (PyTorch)](https://github.com/adap/flower/tree/main/examples/quickstart-pytorch)
-* [Quickstart (Hugging Face)](https://github.com/adap/flower/tree/main/examples/quickstart-huggingface)
-* [Quickstart (PyTorch Lightning)](https://github.com/adap/flower/tree/main/examples/quickstart-pytorch-lightning)
-* [Quickstart (fastai)](https://github.com/adap/flower/tree/main/examples/quickstart-fastai)
-* [Quickstart (Pandas)](https://github.com/adap/flower/tree/main/examples/quickstart-pandas)
-* [Quickstart (MXNet)](https://github.com/adap/flower/tree/main/examples/quickstart-mxnet)
-* [Quickstart (JAX)](https://github.com/adap/flower/tree/main/examples/quickstart-jax)
-* [Quickstart (scikit-learn)](https://github.com/adap/flower/tree/main/examples/sklearn-logreg-mnist)
-* [Quickstart (Android [TFLite])](https://github.com/adap/flower/tree/main/examples/android)
-* [Quickstart (iOS [CoreML])](https://github.com/adap/flower/tree/main/examples/ios)
+- [Quickstart (TensorFlow)](https://github.com/adap/flower/tree/main/examples/quickstart-tensorflow)
+- [Quickstart (PyTorch)](https://github.com/adap/flower/tree/main/examples/quickstart-pytorch)
+- [Quickstart (Hugging Face)](https://github.com/adap/flower/tree/main/examples/quickstart-huggingface)
+- [Quickstart (PyTorch Lightning)](https://github.com/adap/flower/tree/main/examples/quickstart-pytorch-lightning)
+- [Quickstart (fastai)](https://github.com/adap/flower/tree/main/examples/quickstart-fastai)
+- [Quickstart (Pandas)](https://github.com/adap/flower/tree/main/examples/quickstart-pandas)
+- [Quickstart (MXNet)](https://github.com/adap/flower/tree/main/examples/quickstart-mxnet)
+- [Quickstart (JAX)](https://github.com/adap/flower/tree/main/examples/quickstart-jax)
+- [Quickstart (scikit-learn)](https://github.com/adap/flower/tree/main/examples/sklearn-logreg-mnist)
+- [Quickstart (Android [TFLite])](https://github.com/adap/flower/tree/main/examples/android)
+- [Quickstart (iOS [CoreML])](https://github.com/adap/flower/tree/main/examples/ios)
Other [examples](https://github.com/adap/flower/tree/main/examples):
-* [Raspberry Pi & Nvidia Jetson Tutorial](https://github.com/adap/flower/tree/main/examples/embedded-devices)
-* [PyTorch: From Centralized to Federated](https://github.com/adap/flower/tree/main/examples/pytorch-from-centralized-to-federated)
-* [MXNet: From Centralized to Federated](https://github.com/adap/flower/tree/main/examples/mxnet-from-centralized-to-federated)
-* [Advanced Flower with TensorFlow/Keras](https://github.com/adap/flower/tree/main/examples/advanced-tensorflow)
-* [Advanced Flower with PyTorch](https://github.com/adap/flower/tree/main/examples/advanced-pytorch)
-* Single-Machine Simulation of Federated Learning Systems ([PyTorch](https://github.com/adap/flower/tree/main/examples/simulation_pytorch)) ([Tensorflow](https://github.com/adap/flower/tree/main/examples/simulation_tensorflow))
+- [Raspberry Pi & Nvidia Jetson Tutorial](https://github.com/adap/flower/tree/main/examples/embedded-devices)
+- [PyTorch: From Centralized to Federated](https://github.com/adap/flower/tree/main/examples/pytorch-from-centralized-to-federated)
+- [MXNet: From Centralized to Federated](https://github.com/adap/flower/tree/main/examples/mxnet-from-centralized-to-federated)
+- [Advanced Flower with TensorFlow/Keras](https://github.com/adap/flower/tree/main/examples/advanced-tensorflow)
+- [Advanced Flower with PyTorch](https://github.com/adap/flower/tree/main/examples/advanced-pytorch)
+- Single-Machine Simulation of Federated Learning Systems ([PyTorch](https://github.com/adap/flower/tree/main/examples/simulation_pytorch)) ([Tensorflow](https://github.com/adap/flower/tree/main/examples/simulation_tensorflow))
## Community
@@ -144,12 +143,12 @@ Flower is built by a wonderful community of researchers and engineers. [Join Sla
## Citation
-If you publish work that uses Flower, please cite Flower as follows:
+If you publish work that uses Flower, please cite Flower as follows:
```bibtex
@article{beutel2020flower,
title={Flower: A Friendly Federated Learning Research Framework},
- author={Beutel, Daniel J and Topal, Taner and Mathur, Akhil and Qiu, Xinchi and Fernandez-Marques, Javier and Gao, Yan and Sani, Lorenzo and Kwing, Hei Li and Parcollet, Titouan and Gusmão, Pedro PB de and Lane, Nicholas D},
+ author={Beutel, Daniel J and Topal, Taner and Mathur, Akhil and Qiu, Xinchi and Fernandez-Marques, Javier and Gao, Yan and Sani, Lorenzo and Kwing, Hei Li and Parcollet, Titouan and Gusmão, Pedro PB de and Lane, Nicholas D},
journal={arXiv preprint arXiv:2007.14390},
year={2020}
}
diff --git a/baselines/depthfl/.gitignore b/baselines/depthfl/.gitignore
new file mode 100644
index 000000000000..fb7448bbcb01
--- /dev/null
+++ b/baselines/depthfl/.gitignore
@@ -0,0 +1,4 @@
+dataset/
+outputs/
+prev_grads/
+multirun/
\ No newline at end of file
diff --git a/baselines/depthfl/LICENSE b/baselines/depthfl/LICENSE
new file mode 100644
index 000000000000..d64569567334
--- /dev/null
+++ b/baselines/depthfl/LICENSE
@@ -0,0 +1,202 @@
+
+ Apache License
+ Version 2.0, January 2004
+ http://www.apache.org/licenses/
+
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
+
+ 1. Definitions.
+
+ "License" shall mean the terms and conditions for use, reproduction,
+ and distribution as defined by Sections 1 through 9 of this document.
+
+ "Licensor" shall mean the copyright owner or entity authorized by
+ the copyright owner that is granting the License.
+
+ "Legal Entity" shall mean the union of the acting entity and all
+ other entities that control, are controlled by, or are under common
+ control with that entity. For the purposes of this definition,
+ "control" means (i) the power, direct or indirect, to cause the
+ direction or management of such entity, whether by contract or
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
+ outstanding shares, or (iii) beneficial ownership of such entity.
+
+ "You" (or "Your") shall mean an individual or Legal Entity
+ exercising permissions granted by this License.
+
+ "Source" form shall mean the preferred form for making modifications,
+ including but not limited to software source code, documentation
+ source, and configuration files.
+
+ "Object" form shall mean any form resulting from mechanical
+ transformation or translation of a Source form, including but
+ not limited to compiled object code, generated documentation,
+ and conversions to other media types.
+
+ "Work" shall mean the work of authorship, whether in Source or
+ Object form, made available under the License, as indicated by a
+ copyright notice that is included in or attached to the work
+ (an example is provided in the Appendix below).
+
+ "Derivative Works" shall mean any work, whether in Source or Object
+ form, that is based on (or derived from) the Work and for which the
+ editorial revisions, annotations, elaborations, or other modifications
+ represent, as a whole, an original work of authorship. For the purposes
+ of this License, Derivative Works shall not include works that remain
+ separable from, or merely link (or bind by name) to the interfaces of,
+ the Work and Derivative Works thereof.
+
+ "Contribution" shall mean any work of authorship, including
+ the original version of the Work and any modifications or additions
+ to that Work or Derivative Works thereof, that is intentionally
+ submitted to Licensor for inclusion in the Work by the copyright owner
+ or by an individual or Legal Entity authorized to submit on behalf of
+ the copyright owner. For the purposes of this definition, "submitted"
+ means any form of electronic, verbal, or written communication sent
+ to the Licensor or its representatives, including but not limited to
+ communication on electronic mailing lists, source code control systems,
+ and issue tracking systems that are managed by, or on behalf of, the
+ Licensor for the purpose of discussing and improving the Work, but
+ excluding communication that is conspicuously marked or otherwise
+ designated in writing by the copyright owner as "Not a Contribution."
+
+ "Contributor" shall mean Licensor and any individual or Legal Entity
+ on behalf of whom a Contribution has been received by Licensor and
+ subsequently incorporated within the Work.
+
+ 2. Grant of Copyright License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ copyright license to reproduce, prepare Derivative Works of,
+ publicly display, publicly perform, sublicense, and distribute the
+ Work and such Derivative Works in Source or Object form.
+
+ 3. Grant of Patent License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ (except as stated in this section) patent license to make, have made,
+ use, offer to sell, sell, import, and otherwise transfer the Work,
+ where such license applies only to those patent claims licensable
+ by such Contributor that are necessarily infringed by their
+ Contribution(s) alone or by combination of their Contribution(s)
+ with the Work to which such Contribution(s) was submitted. If You
+ institute patent litigation against any entity (including a
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
+ or a Contribution incorporated within the Work constitutes direct
+ or contributory patent infringement, then any patent licenses
+ granted to You under this License for that Work shall terminate
+ as of the date such litigation is filed.
+
+ 4. Redistribution. You may reproduce and distribute copies of the
+ Work or Derivative Works thereof in any medium, with or without
+ modifications, and in Source or Object form, provided that You
+ meet the following conditions:
+
+ (a) You must give any other recipients of the Work or
+ Derivative Works a copy of this License; and
+
+ (b) You must cause any modified files to carry prominent notices
+ stating that You changed the files; and
+
+ (c) You must retain, in the Source form of any Derivative Works
+ that You distribute, all copyright, patent, trademark, and
+ attribution notices from the Source form of the Work,
+ excluding those notices that do not pertain to any part of
+ the Derivative Works; and
+
+ (d) If the Work includes a "NOTICE" text file as part of its
+ distribution, then any Derivative Works that You distribute must
+ include a readable copy of the attribution notices contained
+ within such NOTICE file, excluding those notices that do not
+ pertain to any part of the Derivative Works, in at least one
+ of the following places: within a NOTICE text file distributed
+ as part of the Derivative Works; within the Source form or
+ documentation, if provided along with the Derivative Works; or,
+ within a display generated by the Derivative Works, if and
+ wherever such third-party notices normally appear. The contents
+ of the NOTICE file are for informational purposes only and
+ do not modify the License. You may add Your own attribution
+ notices within Derivative Works that You distribute, alongside
+ or as an addendum to the NOTICE text from the Work, provided
+ that such additional attribution notices cannot be construed
+ as modifying the License.
+
+ You may add Your own copyright statement to Your modifications and
+ may provide additional or different license terms and conditions
+ for use, reproduction, or distribution of Your modifications, or
+ for any such Derivative Works as a whole, provided Your use,
+ reproduction, and distribution of the Work otherwise complies with
+ the conditions stated in this License.
+
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
+ any Contribution intentionally submitted for inclusion in the Work
+ by You to the Licensor shall be under the terms and conditions of
+ this License, without any additional terms or conditions.
+ Notwithstanding the above, nothing herein shall supersede or modify
+ the terms of any separate license agreement you may have executed
+ with Licensor regarding such Contributions.
+
+ 6. Trademarks. This License does not grant permission to use the trade
+ names, trademarks, service marks, or product names of the Licensor,
+ except as required for reasonable and customary use in describing the
+ origin of the Work and reproducing the content of the NOTICE file.
+
+ 7. Disclaimer of Warranty. Unless required by applicable law or
+ agreed to in writing, Licensor provides the Work (and each
+ Contributor provides its Contributions) on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
+ implied, including, without limitation, any warranties or conditions
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
+ PARTICULAR PURPOSE. You are solely responsible for determining the
+ appropriateness of using or redistributing the Work and assume any
+ risks associated with Your exercise of permissions under this License.
+
+ 8. Limitation of Liability. In no event and under no legal theory,
+ whether in tort (including negligence), contract, or otherwise,
+ unless required by applicable law (such as deliberate and grossly
+ negligent acts) or agreed to in writing, shall any Contributor be
+ liable to You for damages, including any direct, indirect, special,
+ incidental, or consequential damages of any character arising as a
+ result of this License or out of the use or inability to use the
+ Work (including but not limited to damages for loss of goodwill,
+ work stoppage, computer failure or malfunction, or any and all
+ other commercial damages or losses), even if such Contributor
+ has been advised of the possibility of such damages.
+
+ 9. Accepting Warranty or Additional Liability. While redistributing
+ the Work or Derivative Works thereof, You may choose to offer,
+ and charge a fee for, acceptance of support, warranty, indemnity,
+ or other liability obligations and/or rights consistent with this
+ License. However, in accepting such obligations, You may act only
+ on Your own behalf and on Your sole responsibility, not on behalf
+ of any other Contributor, and only if You agree to indemnify,
+ defend, and hold each Contributor harmless for any liability
+ incurred by, or claims asserted against, such Contributor by reason
+ of your accepting any such warranty or additional liability.
+
+ END OF TERMS AND CONDITIONS
+
+ APPENDIX: How to apply the Apache License to your work.
+
+ To apply the Apache License to your work, attach the following
+ boilerplate notice, with the fields enclosed by brackets "[]"
+ replaced with your own identifying information. (Don't include
+ the brackets!) The text should be enclosed in the appropriate
+ comment syntax for the file format. We also recommend that a
+ file or class name and description of purpose be included on the
+ same "printed page" as the copyright notice for easier
+ identification within third-party archives.
+
+ Copyright [yyyy] [name of copyright owner]
+
+ Licensed under the Apache License, Version 2.0 (the "License");
+ you may not use this file except in compliance with the License.
+ You may obtain a copy of the License at
+
+ http://www.apache.org/licenses/LICENSE-2.0
+
+ Unless required by applicable law or agreed to in writing, software
+ distributed under the License is distributed on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ See the License for the specific language governing permissions and
+ limitations under the License.
diff --git a/baselines/depthfl/README.md b/baselines/depthfl/README.md
new file mode 100644
index 000000000000..b8ab7ed18571
--- /dev/null
+++ b/baselines/depthfl/README.md
@@ -0,0 +1,171 @@
+---
+title: DepthFL:Depthwise Federated Learning for Heterogeneous Clients
+url: https://openreview.net/forum?id=pf8RIZTMU58
+labels: [image classification, system heterogeneity, cross-device, knowledge distillation]
+dataset: [CIFAR-100]
+---
+
+# DepthFL: Depthwise Federated Learning for Heterogeneous Clients
+
+> Note: If you use this baseline in your work, please remember to cite the original authors of the paper as well as the Flower paper.
+
+**Paper:** [openreview.net/forum?id=pf8RIZTMU58](https://openreview.net/forum?id=pf8RIZTMU58)
+
+**Authors:** Minjae Kim, Sangyoon Yu, Suhyun Kim, Soo-Mook Moon
+
+**Abstract:** Federated learning is for training a global model without collecting private local data from clients. As they repeatedly need to upload locally-updated weights or gradients instead, clients require both computation and communication resources enough to participate in learning, but in reality their resources are heterogeneous. To enable resource-constrained clients to train smaller local models, width scaling techniques have been used, which reduces the channels of a global model. Unfortunately, width scaling suffers from heterogeneity of local models when averaging them, leading to a lower accuracy than when simply excluding resource-constrained clients from training. This paper proposes a new approach based on depth scaling called DepthFL. DepthFL defines local models of different depths by pruning the deepest layers off the global model, and allocates them to clients depending on their available resources. Since many clients do not have enough resources to train deep local models, this would make deep layers partially-trained with insufficient data, unlike shallow layers that are fully trained. DepthFL alleviates this problem by mutual self-distillation of knowledge among the classifiers of various depths within a local model. Our experiments show that depth-scaled local models build a global model better than width-scaled ones, and that self-distillation is highly effective in training data-insufficient deep layers.
+
+
+## About this baseline
+
+**What’s implemented:** The code in this directory replicates the experiments in DepthFL: Depthwise Federated Learning for Heterogeneous Clients (Kim et al., 2023) for CIFAR100, which proposed the DepthFL algorithm. Concretely, it replicates the results for CIFAR100 dataset in Table 2, 3 and 4.
+
+**Datasets:** CIFAR100 from PyTorch's Torchvision
+
+**Hardware Setup:** These experiments were run on a server with Nvidia 3090 GPUs. Any machine with 1x 8GB GPU or more would be able to run it in a reasonable amount of time. With the default settings, clients make use of 1.3GB of VRAM. Lower `num_gpus` in `client_resources` to train more clients in parallel on your GPU(s).
+
+**Contributors:** Minjae Kim
+
+
+## Experimental Setup
+
+**Task:** Image Classification
+
+**Model:** ResNet18
+
+**Dataset:** This baseline only includes the CIFAR100 dataset. By default it will be partitioned into 100 clients following IID distribution. The settings are as follow:
+
+| Dataset | #classes | #partitions | partitioning method |
+| :------ | :---: | :---: | :---: |
+| CIFAR100 | 100 | 100 | IID or Non-IID |
+
+**Training Hyperparameters:**
+The following table shows the main hyperparameters for this baseline with their default value (i.e. the value used if you run `python -m depthfl.main` directly)
+
+| Description | Default Value |
+| ----------- | ----- |
+| total clients | 100 |
+| local epoch | 5 |
+| batch size | 50 |
+| number of rounds | 1000 |
+| participation ratio | 10% |
+| learning rate | 0.1 |
+| learning rate decay | 0.998 |
+| client resources | {'num_cpus': 1.0, 'num_gpus': 0.5 }|
+| data partition | IID |
+| optimizer | SGD with dynamic regularization |
+| alpha | 0.1 |
+
+
+## Environment Setup
+
+To construct the Python environment follow these steps:
+
+```bash
+# Set python version
+pyenv install 3.10.6
+pyenv local 3.10.6
+
+# Tell poetry to use python 3.10
+poetry env use 3.10.6
+
+# Install the base Poetry environment
+poetry install
+
+# Activate the environment
+poetry shell
+```
+
+
+## Running the Experiments
+
+To run this DepthFL, first ensure you have activated your Poetry environment (execute `poetry shell` from this directory), then:
+
+```bash
+# this will run using the default settings in the `conf/config.yaml`
+python -m depthfl.main # 'accuracy' : accuracy of the ensemble model, 'accuracy_single' : accuracy of each classifier.
+
+# you can override settings directly from the command line
+python -m depthfl.main exclusive_learning=true model_size=1 # exclusive learning - 100% (a)
+python -m depthfl.main exclusive_learning=true model_size=4 # exclusive learning - 25% (d)
+python -m depthfl.main fit_config.feddyn=false fit_config.kd=false # DepthFL (FedAvg)
+python -m depthfl.main fit_config.feddyn=false fit_config.kd=false fit_config.extended=false # InclusiveFL
+```
+
+To run using HeteroFL:
+```bash
+# since sbn takes too long, we test global model every 50 rounds.
+python -m depthfl.main --config-name="heterofl" # HeteroFL
+python -m depthfl.main --config-name="heterofl" exclusive_learning=true model_size=1 # exclusive learning - 100% (a)
+```
+
+### Stateful clients comment
+
+To implement `feddyn`, stateful clients that store prev_grads information are needed. Since flwr does not yet officially support stateful clients, it was implemented as a temporary measure by loading `prev_grads` from disk when creating a client, and then storing it again on disk after learning. Specifically, there are files that store the state of each client in the `prev_grads` folder. When the strategy is instantiated (for both `FedDyn` and `HeteroFL`) the content of `prev_grads` is reset.
+
+
+## Expected Results
+
+With the following command we run DepthFL (FedDyn / FedAvg), InclusiveFL, and HeteroFL to replicate the results of table 2,3,4 in DepthFL paper. Tables 2, 3, and 4 may contain results from the same experiment in multiple tables.
+
+```bash
+# table 2 (HeteroFL row)
+python -m depthfl.main --config-name="heterofl"
+python -m depthfl.main --config-name="heterofl" --multirun exclusive_learning=true model.scale=false model_size=1,2,3,4
+
+# table 2 (DepthFL(FedAvg) row)
+python -m depthfl.main fit_config.feddyn=false fit_config.kd=false
+python -m depthfl.main --multirun fit_config.feddyn=false fit_config.kd=false exclusive_learning=true model_size=1,2,3,4
+
+# table 2 (DepthFL row)
+python -m depthfl.main
+python -m depthfl.main --multirun exclusive_learning=true model_size=1,2,3,4
+```
+
+**Table 2**
+
+100% (a), 75%(b), 50%(c), 25% (d) cases are exclusive learning scenario. 100% (a) exclusive learning means, the global model and every local model are equal to the smallest local model, and 100% clients participate in learning. Likewise, 25% (d) exclusive learning means, the global model and every local model are equal to the larget local model, and only 25% clients participate in learning.
+
+| Scaling Method | Dataset | Global Model | 100% (a) | 75% (b) | 50% (c) | 25% (d) |
+| :---: | :---: | :---: | :---: | :---: | :---: | :---: |
+| HeteroFL DepthFL (FedAvg) DepthFL | CIFAR100 | 57.61 72.67 76.06 | 64.39 67.08 69.68 | 66.08 70.78 73.21 | 62.03 68.41 70.29 | 51.99 59.17 60.32 |
+
+```bash
+# table 3 (Width Scaling - Duplicate results from table 2)
+python -m depthfl.main --config-name="heterofl"
+python -m depthfl.main --config-name="heterofl" --multirun exclusive_learning=true model.scale=false model_size=1,2,3,4
+
+# table 3 (Depth Scaling : Exclusive Learning, DepthFL(FedAvg) rows - Duplicate results from table 2)
+python -m depthfl.main fit_config.feddyn=false fit_config.kd=false
+python -m depthfl.main --multirun fit_config.feddyn=false fit_config.kd=false exclusive_learning=true model_size=1,2,3,4
+
+## table 3 (Depth Scaling - InclusiveFL row)
+python -m depthfl.main fit_config.feddyn=false fit_config.kd=false fit_config.extended=false
+```
+
+**Table 3**
+
+Accuracy of global sub-models compared to exclusive learning on CIFAR-100.
+
+| Method | Algorithm | Classifier 1/4 | Classifier 2/4 | Classifier 3/4 | Classifier 4/4 |
+| :---: | :---: | :---: | :---: | :---: | :---: |
+| Width Scaling | Exclusive Learning HeteroFL| 64.39 51.08 | 66.08 55.89 | 62.03 58.29 | 51.99 57.61 |
+
+| Method | Algorithm | Classifier 1/4 | Classifier 2/4 | Classifier 3/4 | Classifier 4/4 |
+| :---: | :---: | :---: | :---: | :---: | :---: |
+| Depth Scaling | Exclusive Learning InclusiveFL DepthFL (FedAvg) | 67.08 47.61 66.18 | 68.00 53.88 67.56 | 66.19 59.48 67.97 | 56.78 60.46 68.01 |
+
+```bash
+# table 4
+python -m depthfl.main --multirun fit_config.kd=true,false dataset_config.iid=true,false
+```
+
+**Table 4**
+
+Accuracy of the global model with/without self distillation on CIFAR-100.
+
+| Distribution | Dataset | KD | Classifier 1/4 | Classifier 2/4 | Classifier 3/4 | Classifier 4/4 | Ensemble |
+| :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: |
+| IID | CIFAR100 | ✗ ✓ | 70.13 71.74 | 69.63 73.35 | 68.92 73.57 | 68.92 73.55 | 74.48 76.06 |
+| non-IID | CIFAR100 | ✗ ✓ | 67.94 70.33 | 68.68 71.88 | 68.46 72.43 | 67.78 72.34 | 73.18 74.92 |
+
diff --git a/baselines/depthfl/depthfl/__init__.py b/baselines/depthfl/depthfl/__init__.py
new file mode 100644
index 000000000000..3343905e1879
--- /dev/null
+++ b/baselines/depthfl/depthfl/__init__.py
@@ -0,0 +1 @@
+"""Flower summer of reproducibility : DepthFL (ICLR' 23)."""
diff --git a/baselines/depthfl/depthfl/client.py b/baselines/depthfl/depthfl/client.py
new file mode 100644
index 000000000000..481ac90f1c79
--- /dev/null
+++ b/baselines/depthfl/depthfl/client.py
@@ -0,0 +1,181 @@
+"""Defines the DepthFL Flower Client and a function to instantiate it."""
+
+import copy
+import pickle
+from collections import OrderedDict
+from typing import Callable, Dict, List, Tuple
+
+import flwr as fl
+import numpy as np
+import torch
+from flwr.common.typing import NDArrays, Scalar
+from hydra.utils import instantiate
+from omegaconf import DictConfig
+from torch.utils.data import DataLoader
+
+from depthfl.models import test, train
+
+
+def prune(state_dict, param_idx):
+ """Prune width of DNN (for HeteroFL)."""
+ ret_dict = {}
+ for k in state_dict.keys():
+ if "num" not in k:
+ ret_dict[k] = state_dict[k][torch.meshgrid(param_idx[k])]
+ else:
+ ret_dict[k] = state_dict[k]
+ return copy.deepcopy(ret_dict)
+
+
+class FlowerClient(
+ fl.client.NumPyClient
+): # pylint: disable=too-many-instance-attributes
+ """Standard Flower client for CNN training."""
+
+ def __init__(
+ self,
+ net: torch.nn.Module,
+ trainloader: DataLoader,
+ valloader: DataLoader,
+ device: torch.device,
+ num_epochs: int,
+ learning_rate: float,
+ learning_rate_decay: float,
+ prev_grads: Dict,
+ cid: int,
+ ): # pylint: disable=too-many-arguments
+ self.net = net
+ self.trainloader = trainloader
+ self.valloader = valloader
+ self.device = device
+ self.num_epochs = num_epochs
+ self.learning_rate = learning_rate
+ self.learning_rate_decay = learning_rate_decay
+ self.prev_grads = prev_grads
+ self.cid = cid
+ self.param_idx = {}
+ state_dict = net.state_dict()
+
+ # for HeteroFL
+ for k in state_dict.keys():
+ self.param_idx[k] = [
+ torch.arange(size) for size in state_dict[k].shape
+ ] # store client's weights' shape (for HeteroFL)
+
+ def get_parameters(self, config: Dict[str, Scalar]) -> NDArrays:
+ """Return the parameters of the current net."""
+ return [val.cpu().numpy() for _, val in self.net.state_dict().items()]
+
+ def set_parameters(self, parameters: NDArrays) -> None:
+ """Change the parameters of the model using the given ones."""
+ params_dict = zip(self.net.state_dict().keys(), parameters)
+ state_dict = OrderedDict({k: torch.tensor(v) for k, v in params_dict})
+ self.net.load_state_dict(prune(state_dict, self.param_idx), strict=True)
+
+ def fit(
+ self, parameters: NDArrays, config: Dict[str, Scalar]
+ ) -> Tuple[NDArrays, int, Dict]:
+ """Implement distributed fit function for a given client."""
+ self.set_parameters(parameters)
+ num_epochs = self.num_epochs
+
+ curr_round = int(config["curr_round"]) - 1
+
+ # consistency weight for self distillation in DepthFL
+ consistency_weight_constant = 300
+ current = np.clip(curr_round, 0.0, consistency_weight_constant)
+ phase = 1.0 - current / consistency_weight_constant
+ consistency_weight = float(np.exp(-5.0 * phase * phase))
+
+ train(
+ self.net,
+ self.trainloader,
+ self.device,
+ epochs=num_epochs,
+ learning_rate=self.learning_rate * self.learning_rate_decay**curr_round,
+ config=config,
+ consistency_weight=consistency_weight,
+ prev_grads=self.prev_grads,
+ )
+
+ with open(f"prev_grads/client_{self.cid}", "wb") as prev_grads_file:
+ pickle.dump(self.prev_grads, prev_grads_file)
+
+ return self.get_parameters({}), len(self.trainloader), {"cid": self.cid}
+
+ def evaluate(
+ self, parameters: NDArrays, config: Dict[str, Scalar]
+ ) -> Tuple[float, int, Dict]:
+ """Implement distributed evaluation for a given client."""
+ self.set_parameters(parameters)
+ loss, accuracy, accuracy_single = test(self.net, self.valloader, self.device)
+ return (
+ float(loss),
+ len(self.valloader),
+ {"accuracy": float(accuracy), "accuracy_single": accuracy_single},
+ )
+
+
+def gen_client_fn( # pylint: disable=too-many-arguments
+ num_epochs: int,
+ trainloaders: List[DataLoader],
+ valloaders: List[DataLoader],
+ learning_rate: float,
+ learning_rate_decay: float,
+ models: List[DictConfig],
+) -> Callable[[str], FlowerClient]:
+ """Generate the client function that creates the Flower Clients.
+
+ Parameters
+ ----------
+ num_epochs : int
+ The number of local epochs each client should run the training for before
+ sending it to the server.
+ trainloaders: List[DataLoader]
+ A list of DataLoaders, each pointing to the dataset training partition
+ belonging to a particular client.
+ valloaders: List[DataLoader]
+ A list of DataLoaders, each pointing to the dataset validation partition
+ belonging to a particular client.
+ learning_rate : float
+ The learning rate for the SGD optimizer of clients.
+ learning_rate_decay : float
+ The learning rate decay ratio per round for the SGD optimizer of clients.
+ models : List[DictConfig]
+ A list of DictConfigs, each pointing to the model config of client's local model
+
+ Returns
+ -------
+ Callable[[str], FlowerClient]
+ client function that creates Flower Clients
+ """
+
+ def client_fn(cid: str) -> FlowerClient:
+ """Create a Flower client representing a single organization."""
+ # Load model
+ device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
+
+ # each client gets a different model config (different width / depth)
+ net = instantiate(models[int(cid)]).to(device)
+
+ # Note: each client gets a different trainloader/valloader, so each client
+ # will train and evaluate on their own unique data
+ trainloader = trainloaders[int(cid)]
+ valloader = valloaders[int(cid)]
+
+ with open(f"prev_grads/client_{int(cid)}", "rb") as prev_grads_file:
+ prev_grads = pickle.load(prev_grads_file)
+
+ return FlowerClient(
+ net,
+ trainloader,
+ valloader,
+ device,
+ num_epochs,
+ learning_rate,
+ learning_rate_decay,
+ prev_grads,
+ int(cid),
+ )
+
+ return client_fn
diff --git a/baselines/depthfl/depthfl/conf/config.yaml b/baselines/depthfl/depthfl/conf/config.yaml
new file mode 100644
index 000000000000..5a126229956e
--- /dev/null
+++ b/baselines/depthfl/depthfl/conf/config.yaml
@@ -0,0 +1,42 @@
+---
+
+num_clients: 100 # total number of clients
+num_epochs: 5 # number of local epochs
+batch_size: 50
+num_rounds: 1000
+fraction: 0.1 # participation ratio
+learning_rate: 0.1
+learning_rate_decay : 0.998 # per round
+static_bn: false # static batch normalization (HeteroFL)
+exclusive_learning: false # exclusive learning baseline in DepthFL paper
+model_size: 1 # model size for exclusive learning
+
+client_resources:
+ num_cpus: 1
+ num_gpus: 0.5
+
+server_device: cuda
+
+dataset_config:
+ iid: true
+ beta: 0.5
+
+fit_config:
+ feddyn: true
+ kd: true
+ alpha: 0.1 # alpha for FedDyn
+ extended: true # if not extended : InclusiveFL
+ drop_client: false # with FedProx, clients shouldn't be dropped even if they are stragglers
+
+model:
+ _target_: depthfl.resnet.multi_resnet18
+ n_blocks: 4 # depth (1 ~ 4)
+ num_classes: 100
+
+strategy:
+ _target_: depthfl.strategy.FedDyn
+ fraction_fit: 0.00001 # because we want the number of clients to sample on each round to be solely defined by min_fit_clients
+ fraction_evaluate: 0.0
+ # min_fit_clients: ${clients_per_round}
+ min_evaluate_clients: 0
+ # min_available_clients: ${clients_per_round}
\ No newline at end of file
diff --git a/baselines/depthfl/depthfl/conf/heterofl.yaml b/baselines/depthfl/depthfl/conf/heterofl.yaml
new file mode 100644
index 000000000000..ad0bb8c8f8b8
--- /dev/null
+++ b/baselines/depthfl/depthfl/conf/heterofl.yaml
@@ -0,0 +1,43 @@
+---
+
+num_clients: 100 # total number of clients
+num_epochs: 5 # number of local epochs
+batch_size: 50
+num_rounds: 1000
+fraction: 0.1 # participation ratio
+learning_rate: 0.1
+learning_rate_decay : 0.998 # per round
+static_bn: true # static batch normalization (HeteroFL)
+exclusive_learning: false # exclusive learning baseline in DepthFL paper
+model_size: 1 # model size for exclusive learning
+
+client_resources:
+ num_cpus: 1
+ num_gpus: 0.5
+
+server_device: cuda
+
+dataset_config:
+ iid: true
+ beta: 0.5
+
+fit_config:
+ feddyn: false
+ kd: false
+ alpha: 0.1 # unused
+ extended: false # unused
+ drop_client: false # with FedProx, clients shouldn't be dropped even if they are stragglers
+
+model:
+ _target_: depthfl.resnet_hetero.resnet18
+ n_blocks: 4 # width (1 ~ 4)
+ num_classes: 100
+ scale: true # scaler module in HeteroFL
+
+strategy:
+ _target_: depthfl.strategy_hetero.HeteroFL
+ fraction_fit: 0.00001 # because we want the number of clients to sample on each round to be solely defined by min_fit_clients
+ fraction_evaluate: 0.0
+ # min_fit_clients: ${clients_per_round}
+ min_evaluate_clients: 0
+ # min_available_clients: ${clients_per_round}
\ No newline at end of file
diff --git a/baselines/depthfl/depthfl/dataset.py b/baselines/depthfl/depthfl/dataset.py
new file mode 100644
index 000000000000..c2024fe068a0
--- /dev/null
+++ b/baselines/depthfl/depthfl/dataset.py
@@ -0,0 +1,60 @@
+"""CIFAR100 dataset utilities for federated learning."""
+
+from typing import Optional, Tuple
+
+import torch
+from omegaconf import DictConfig
+from torch.utils.data import DataLoader, random_split
+
+from depthfl.dataset_preparation import _partition_data
+
+
+def load_datasets( # pylint: disable=too-many-arguments
+ config: DictConfig,
+ num_clients: int,
+ val_ratio: float = 0.0,
+ batch_size: Optional[int] = 32,
+ seed: Optional[int] = 41,
+) -> Tuple[DataLoader, DataLoader, DataLoader]:
+ """Create the dataloaders to be fed into the model.
+
+ Parameters
+ ----------
+ config: DictConfig
+ Parameterises the dataset partitioning process
+ num_clients : int
+ The number of clients that hold a part of the data
+ val_ratio : float, optional
+ The ratio of training data that will be used for validation (between 0 and 1),
+ by default 0.1
+ batch_size : int, optional
+ The size of the batches to be fed into the model, by default 32
+ seed : int, optional
+ Used to set a fix seed to replicate experiments, by default 42
+
+ Returns
+ -------
+ Tuple[DataLoader, DataLoader, DataLoader]
+ The DataLoader for training, validation, and testing.
+ """
+ print(f"Dataset partitioning config: {config}")
+ datasets, testset = _partition_data(
+ num_clients,
+ iid=config.iid,
+ beta=config.beta,
+ seed=seed,
+ )
+ # Split each partition into train/val and create DataLoader
+ trainloaders = []
+ valloaders = []
+ for dataset in datasets:
+ len_val = 0
+ if val_ratio > 0:
+ len_val = int(len(dataset) / (1 / val_ratio))
+ lengths = [len(dataset) - len_val, len_val]
+ ds_train, ds_val = random_split(
+ dataset, lengths, torch.Generator().manual_seed(seed)
+ )
+ trainloaders.append(DataLoader(ds_train, batch_size=batch_size, shuffle=True))
+ valloaders.append(DataLoader(ds_val, batch_size=batch_size))
+ return trainloaders, valloaders, DataLoader(testset, batch_size=batch_size)
diff --git a/baselines/depthfl/depthfl/dataset_preparation.py b/baselines/depthfl/depthfl/dataset_preparation.py
new file mode 100644
index 000000000000..006491c7679e
--- /dev/null
+++ b/baselines/depthfl/depthfl/dataset_preparation.py
@@ -0,0 +1,125 @@
+"""Dataset(CIFAR100) preparation for DepthFL."""
+
+from typing import List, Optional, Tuple
+
+import numpy as np
+import torchvision.transforms as transforms
+from torch.utils.data import Dataset, Subset
+from torchvision.datasets import CIFAR100
+
+
+def _download_data() -> Tuple[Dataset, Dataset]:
+ """Download (if necessary) and returns the CIFAR-100 dataset.
+
+ Returns
+ -------
+ Tuple[CIFAR100, CIFAR100]
+ The dataset for training and the dataset for testing CIFAR100.
+ """
+ transform_train = transforms.Compose(
+ [
+ transforms.ToTensor(),
+ transforms.RandomCrop(32, padding=4),
+ transforms.RandomHorizontalFlip(),
+ transforms.Normalize((0.5071, 0.4867, 0.4408), (0.2675, 0.2565, 0.2761)),
+ ]
+ )
+
+ transform_test = transforms.Compose(
+ [
+ transforms.ToTensor(),
+ transforms.Normalize((0.5071, 0.4867, 0.4408), (0.2675, 0.2565, 0.2761)),
+ ]
+ )
+
+ trainset = CIFAR100(
+ "./dataset", train=True, download=True, transform=transform_train
+ )
+ testset = CIFAR100(
+ "./dataset", train=False, download=True, transform=transform_test
+ )
+ return trainset, testset
+
+
+def _partition_data(
+ num_clients,
+ iid: Optional[bool] = True,
+ beta=0.5,
+ seed=41,
+) -> Tuple[List[Dataset], Dataset]:
+ """Split training set to simulate the federated setting.
+
+ Parameters
+ ----------
+ num_clients : int
+ The number of clients that hold a part of the data
+ iid : bool, optional
+ Whether the data should be independent and identically distributed
+ or if the data should first be sorted by labels and distributed by
+ noniid manner to each client, by default true
+ beta : hyperparameter for dirichlet distribution
+ seed : int, optional
+ Used to set a fix seed to replicate experiments, by default 42
+
+ Returns
+ -------
+ Tuple[List[Dataset], Dataset]
+ A list of dataset for each client and a
+ single dataset to be use for testing the model.
+ """
+ trainset, testset = _download_data()
+
+ datasets: List[Subset] = []
+
+ if iid:
+ distribute_iid(num_clients, seed, trainset, datasets)
+
+ else:
+ distribute_noniid(num_clients, beta, seed, trainset, datasets)
+
+ return datasets, testset
+
+
+def distribute_iid(num_clients, seed, trainset, datasets):
+ """Distribute dataset in iid manner."""
+ np.random.seed(seed)
+ num_sample = int(len(trainset) / (num_clients))
+ index = list(range(len(trainset)))
+ for _ in range(num_clients):
+ sample_idx = np.random.choice(index, num_sample, replace=False)
+ index = list(set(index) - set(sample_idx))
+ datasets.append(Subset(trainset, sample_idx))
+
+
+def distribute_noniid(num_clients, beta, seed, trainset, datasets):
+ """Distribute dataset in non-iid manner."""
+ labels = np.array([label for _, label in trainset])
+ min_size = 0
+ np.random.seed(seed)
+
+ while min_size < 10:
+ idx_batch = [[] for _ in range(num_clients)]
+ # for each class in the dataset
+ for k in range(np.max(labels) + 1):
+ idx_k = np.where(labels == k)[0]
+ np.random.shuffle(idx_k)
+ proportions = np.random.dirichlet(np.repeat(beta, num_clients))
+ # Balance
+ proportions = np.array(
+ [
+ p * (len(idx_j) < labels.shape[0] / num_clients)
+ for p, idx_j in zip(proportions, idx_batch)
+ ]
+ )
+ proportions = proportions / proportions.sum()
+ proportions = (np.cumsum(proportions) * len(idx_k)).astype(int)[:-1]
+ idx_batch = [
+ idx_j + idx.tolist()
+ for idx_j, idx in zip(idx_batch, np.split(idx_k, proportions))
+ ]
+ min_size = min([len(idx_j) for idx_j in idx_batch])
+
+ for j in range(num_clients):
+ np.random.shuffle(idx_batch[j])
+ # net_dataidx_map[j] = np.array(idx_batch[j])
+ datasets.append(Subset(trainset, np.array(idx_batch[j])))
diff --git a/baselines/depthfl/depthfl/main.py b/baselines/depthfl/depthfl/main.py
new file mode 100644
index 000000000000..7bf1d9563eae
--- /dev/null
+++ b/baselines/depthfl/depthfl/main.py
@@ -0,0 +1,135 @@
+"""DepthFL main."""
+
+import copy
+
+import flwr as fl
+import hydra
+from flwr.common import ndarrays_to_parameters
+from flwr.server.client_manager import SimpleClientManager
+from hydra.core.hydra_config import HydraConfig
+from hydra.utils import instantiate
+from omegaconf import DictConfig, OmegaConf
+
+from depthfl import client, server
+from depthfl.dataset import load_datasets
+from depthfl.utils import save_results_as_pickle
+
+
+@hydra.main(config_path="conf", config_name="config", version_base=None)
+def main(cfg: DictConfig) -> None:
+ """Run the baseline.
+
+ Parameters
+ ----------
+ cfg : DictConfig
+ An omegaconf object that stores the hydra config.
+ """
+ print(OmegaConf.to_yaml(cfg))
+
+ # partition dataset and get dataloaders
+ trainloaders, valloaders, testloader = load_datasets(
+ config=cfg.dataset_config,
+ num_clients=cfg.num_clients,
+ batch_size=cfg.batch_size,
+ )
+
+ # exclusive learning baseline in DepthFL paper
+ # (model_size, % of clients) = (a,100), (b,75), (c,50), (d,25)
+ if cfg.exclusive_learning:
+ cfg.num_clients = int(
+ cfg.num_clients - (cfg.model_size - 1) * (cfg.num_clients // 4)
+ )
+
+ models = []
+ for i in range(cfg.num_clients):
+ model = copy.deepcopy(cfg.model)
+
+ # each client gets different model depth / width
+ model.n_blocks = i // (cfg.num_clients // 4) + 1
+
+ # In exclusive learning, every client has same model depth / width
+ if cfg.exclusive_learning:
+ model.n_blocks = cfg.model_size
+
+ models.append(model)
+
+ # prepare function that will be used to spawn each client
+ client_fn = client.gen_client_fn(
+ num_epochs=cfg.num_epochs,
+ trainloaders=trainloaders,
+ valloaders=valloaders,
+ learning_rate=cfg.learning_rate,
+ learning_rate_decay=cfg.learning_rate_decay,
+ models=models,
+ )
+
+ # get function that will executed by the strategy's evaluate() method
+ # Set server's device
+ device = cfg.server_device
+
+ # Static Batch Normalization for HeteroFL
+ if cfg.static_bn:
+ evaluate_fn = server.gen_evaluate_fn_hetero(
+ trainloaders, testloader, device=device, model_cfg=model
+ )
+ else:
+ evaluate_fn = server.gen_evaluate_fn(testloader, device=device, model=model)
+
+ # get a function that will be used to construct the config that the client's
+ # fit() method will received
+ def get_on_fit_config():
+ def fit_config_fn(server_round):
+ # resolve and convert to python dict
+ fit_config = OmegaConf.to_container(cfg.fit_config, resolve=True)
+ fit_config["curr_round"] = server_round # add round info
+ return fit_config
+
+ return fit_config_fn
+
+ net = instantiate(cfg.model)
+ # instantiate strategy according to config. Here we pass other arguments
+ # that are only defined at run time.
+ strategy = instantiate(
+ cfg.strategy,
+ cfg,
+ net,
+ evaluate_fn=evaluate_fn,
+ on_fit_config_fn=get_on_fit_config(),
+ initial_parameters=ndarrays_to_parameters(
+ [val.cpu().numpy() for _, val in net.state_dict().items()]
+ ),
+ min_fit_clients=int(cfg.num_clients * cfg.fraction),
+ min_available_clients=int(cfg.num_clients * cfg.fraction),
+ )
+
+ # Start simulation
+ history = fl.simulation.start_simulation(
+ client_fn=client_fn,
+ num_clients=cfg.num_clients,
+ config=fl.server.ServerConfig(num_rounds=cfg.num_rounds),
+ client_resources={
+ "num_cpus": cfg.client_resources.num_cpus,
+ "num_gpus": cfg.client_resources.num_gpus,
+ },
+ strategy=strategy,
+ server=server.ServerFedDyn(
+ client_manager=SimpleClientManager(), strategy=strategy
+ ),
+ )
+
+ # Experiment completed. Now we save the results and
+ # generate plots using the `history`
+ print("................")
+ print(history)
+
+ # Hydra automatically creates an output directory
+ # Let's retrieve it and save some results there
+ save_path = HydraConfig.get().runtime.output_dir
+
+ # save results as a Python pickle using a file_path
+ # the directory created by Hydra for each run
+ save_results_as_pickle(history, file_path=save_path, extra_results={})
+
+
+if __name__ == "__main__":
+ main()
diff --git a/baselines/depthfl/depthfl/models.py b/baselines/depthfl/depthfl/models.py
new file mode 100644
index 000000000000..df3eebf9f9ce
--- /dev/null
+++ b/baselines/depthfl/depthfl/models.py
@@ -0,0 +1,301 @@
+"""ResNet18 model architecutre, training, and testing functions for CIFAR100."""
+
+
+from typing import List, Tuple
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from omegaconf import DictConfig
+from torch.utils.data import DataLoader
+
+
+class KLLoss(nn.Module):
+ """KL divergence loss for self distillation."""
+
+ def __init__(self):
+ super().__init__()
+ self.temperature = 1
+
+ def forward(self, pred, label):
+ """KL loss forward."""
+ predict = F.log_softmax(pred / self.temperature, dim=1)
+ target_data = F.softmax(label / self.temperature, dim=1)
+ target_data = target_data + 10 ** (-7)
+ with torch.no_grad():
+ target = target_data.detach().clone()
+
+ loss = (
+ self.temperature
+ * self.temperature
+ * ((target * (target.log() - predict)).sum(1).sum() / target.size()[0])
+ )
+ return loss
+
+
+def train( # pylint: disable=too-many-arguments
+ net: nn.Module,
+ trainloader: DataLoader,
+ device: torch.device,
+ epochs: int,
+ learning_rate: float,
+ config: dict,
+ consistency_weight: float,
+ prev_grads: dict,
+) -> None:
+ """Train the network on the training set.
+
+ Parameters
+ ----------
+ net : nn.Module
+ The neural network to train.
+ trainloader : DataLoader
+ The DataLoader containing the data to train the network on.
+ device : torch.device
+ The device on which the model should be trained, either 'cpu' or 'cuda'.
+ epochs : int
+ The number of epochs the model should be trained for.
+ learning_rate : float
+ The learning rate for the SGD optimizer.
+ config : dict
+ training configuration
+ consistency_weight : float
+ hyperparameter for self distillation
+ prev_grads : dict
+ control variate for feddyn
+ """
+ criterion = torch.nn.CrossEntropyLoss()
+ optimizer = torch.optim.SGD(net.parameters(), lr=learning_rate, weight_decay=1e-3)
+ global_params = {
+ k: val.detach().clone().flatten() for (k, val) in net.named_parameters()
+ }
+
+ for k, _ in net.named_parameters():
+ prev_grads[k] = prev_grads[k].to(device)
+
+ net.train()
+ for _ in range(epochs):
+ _train_one_epoch(
+ net,
+ global_params,
+ trainloader,
+ device,
+ criterion,
+ optimizer,
+ config,
+ consistency_weight,
+ prev_grads,
+ )
+
+ # update prev_grads for FedDyn
+ if config["feddyn"]:
+ update_prev_grads(config, net, prev_grads, global_params)
+
+
+def update_prev_grads(config, net, prev_grads, global_params):
+ """Update prev_grads for FedDyn."""
+ for k, param in net.named_parameters():
+ curr_param = param.detach().clone().flatten()
+ prev_grads[k] = prev_grads[k] - config["alpha"] * (
+ curr_param - global_params[k]
+ )
+ prev_grads[k] = prev_grads[k].to(torch.device(torch.device("cpu")))
+
+
+def _train_one_epoch( # pylint: disable=too-many-locals, too-many-arguments
+ net: nn.Module,
+ global_params: dict,
+ trainloader: DataLoader,
+ device: torch.device,
+ criterion: torch.nn.CrossEntropyLoss,
+ optimizer: torch.optim.SGD,
+ config: dict,
+ consistency_weight: float,
+ prev_grads: dict,
+):
+ """Train for one epoch.
+
+ Parameters
+ ----------
+ net : nn.Module
+ The neural network to train.
+ global_params : List[Parameter]
+ The parameters of the global model (from the server).
+ trainloader : DataLoader
+ The DataLoader containing the data to train the network on.
+ device : torch.device
+ The device on which the model should be trained, either 'cpu' or 'cuda'.
+ criterion : torch.nn.CrossEntropyLoss
+ The loss function to use for training
+ optimizer : torch.optim.Adam
+ The optimizer to use for training
+ config : dict
+ training configuration
+ consistency_weight : float
+ hyperparameter for self distillation
+ prev_grads : dict
+ control variate for feddyn
+ """
+ criterion_kl = KLLoss().cuda()
+
+ for images, labels in trainloader:
+ images, labels = images.to(device), labels.to(device)
+ loss = torch.zeros(1).to(device)
+ optimizer.zero_grad()
+ output_lst = net(images)
+
+ for i, branch_output in enumerate(output_lst):
+ # only trains last classifier in InclusiveFL
+ if not config["extended"] and i != len(output_lst) - 1:
+ continue
+
+ loss += criterion(branch_output, labels)
+
+ # self distillation term
+ if config["kd"] and len(output_lst) > 1:
+ for j, output in enumerate(output_lst):
+ if j == i:
+ continue
+
+ loss += (
+ consistency_weight
+ * criterion_kl(branch_output, output.detach())
+ / (len(output_lst) - 1)
+ )
+
+ # Dynamic regularization in FedDyn
+ if config["feddyn"]:
+ for k, param in net.named_parameters():
+ curr_param = param.flatten()
+
+ lin_penalty = torch.dot(curr_param, prev_grads[k])
+ loss -= lin_penalty
+
+ quad_penalty = (
+ config["alpha"]
+ / 2.0
+ * torch.sum(torch.square(curr_param - global_params[k]))
+ )
+ loss += quad_penalty
+
+ loss.backward()
+ optimizer.step()
+
+
+def test( # pylint: disable=too-many-locals
+ net: nn.Module, testloader: DataLoader, device: torch.device
+) -> Tuple[float, float, List[float]]:
+ """Evaluate the network on the entire test set.
+
+ Parameters
+ ----------
+ net : nn.Module
+ The neural network to test.
+ testloader : DataLoader
+ The DataLoader containing the data to test the network on.
+ device : torch.device
+ The device on which the model should be tested, either 'cpu' or 'cuda'.
+
+ Returns
+ -------
+ Tuple[float, float, List[float]]
+ The loss and the accuracy of the global model
+ and the list of accuracy for each classifier on the given data.
+ """
+ criterion = torch.nn.CrossEntropyLoss()
+ correct, total, loss = 0, 0, 0.0
+ correct_single = [0] * 4 # accuracy of each classifier within model
+ net.eval()
+ with torch.no_grad():
+ for images, labels in testloader:
+ images, labels = images.to(device), labels.to(device)
+ output_lst = net(images)
+
+ # ensemble classfiers' output
+ ensemble_output = torch.stack(output_lst, dim=2)
+ ensemble_output = torch.sum(ensemble_output, dim=2) / len(output_lst)
+
+ loss += criterion(ensemble_output, labels).item()
+ _, predicted = torch.max(ensemble_output, 1)
+ total += labels.size(0)
+ correct += (predicted == labels).sum().item()
+
+ for i, single in enumerate(output_lst):
+ _, predicted = torch.max(single, 1)
+ correct_single[i] += (predicted == labels).sum().item()
+
+ if len(testloader.dataset) == 0:
+ raise ValueError("Testloader can't be 0, exiting...")
+ loss /= len(testloader.dataset)
+ accuracy = correct / total
+ accuracy_single = [correct / total for correct in correct_single]
+ return loss, accuracy, accuracy_single
+
+
+def test_sbn( # pylint: disable=too-many-locals
+ nets: List[nn.Module],
+ trainloaders: List[DictConfig],
+ testloader: DataLoader,
+ device: torch.device,
+) -> Tuple[float, float, List[float]]:
+ """Evaluate the networks on the entire test set.
+
+ Parameters
+ ----------
+ nets : List[nn.Module]
+ The neural networks to test. Each neural network has different width
+ trainloaders : List[DataLoader]
+ The List of dataloaders containing the data to train the network on
+ testloader : DataLoader
+ The DataLoader containing the data to test the network on.
+ device : torch.device
+ The device on which the model should be tested, either 'cpu' or 'cuda'.
+
+ Returns
+ -------
+ Tuple[float, float, List[float]]
+ The loss and the accuracy of the global model
+ and the list of accuracy for each classifier on the given data.
+ """
+ # static batch normalization
+ for trainloader in trainloaders:
+ with torch.no_grad():
+ for model in nets:
+ model.train()
+ for _batch_idx, (images, labels) in enumerate(trainloader):
+ images, labels = images.to(device), labels.to(device)
+ output = model(images)
+
+ model.eval()
+
+ criterion = torch.nn.CrossEntropyLoss()
+ correct, total, loss = 0, 0, 0.0
+ correct_single = [0] * 4
+
+ # test each network of different width
+ with torch.no_grad():
+ for images, labels in testloader:
+ images, labels = images.to(device), labels.to(device)
+
+ output_lst = []
+
+ for model in nets:
+ output_lst.append(model(images)[0])
+
+ output = output_lst[-1]
+
+ loss += criterion(output, labels).item()
+ _, predicted = torch.max(output, 1)
+ total += labels.size(0)
+ correct += (predicted == labels).sum().item()
+
+ for i, single in enumerate(output_lst):
+ _, predicted = torch.max(single, 1)
+ correct_single[i] += (predicted == labels).sum().item()
+
+ if len(testloader.dataset) == 0:
+ raise ValueError("Testloader can't be 0, exiting...")
+ loss /= len(testloader.dataset)
+ accuracy = correct / total
+ accuracy_single = [correct / total for correct in correct_single]
+ return loss, accuracy, accuracy_single
diff --git a/baselines/depthfl/depthfl/resnet.py b/baselines/depthfl/depthfl/resnet.py
new file mode 100644
index 000000000000..04348ae17441
--- /dev/null
+++ b/baselines/depthfl/depthfl/resnet.py
@@ -0,0 +1,386 @@
+"""ResNet18 for DepthFL."""
+
+import torch.nn as nn
+
+
+class MyGroupNorm(nn.Module):
+ """Group Normalization layer."""
+
+ def __init__(self, num_channels):
+ super().__init__()
+ # change num_groups to 32
+ self.norm = nn.GroupNorm(
+ num_groups=16, num_channels=num_channels, eps=1e-5, affine=True
+ )
+
+ def forward(self, x):
+ """GN forward."""
+ x = self.norm(x)
+ return x
+
+
+class MyBatchNorm(nn.Module):
+ """Batch Normalization layer."""
+
+ def __init__(self, num_channels):
+ super().__init__()
+ self.norm = nn.BatchNorm2d(num_channels, track_running_stats=True)
+
+ def forward(self, x):
+ """BN forward."""
+ x = self.norm(x)
+ return x
+
+
+def conv3x3(in_planes, out_planes, stride=1):
+ """Convolution layer 3x3."""
+ return nn.Conv2d(
+ in_planes, out_planes, kernel_size=3, stride=stride, padding=1, bias=False
+ )
+
+
+def conv1x1(in_planes, planes, stride=1):
+ """Convolution layer 1x1."""
+ return nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride, bias=False)
+
+
+class SepConv(nn.Module):
+ """Bottleneck layer module."""
+
+ def __init__( # pylint: disable=too-many-arguments
+ self,
+ channel_in,
+ channel_out,
+ kernel_size=3,
+ stride=2,
+ padding=1,
+ norm_layer=MyGroupNorm,
+ ):
+ super().__init__()
+ self.operations = nn.Sequential(
+ nn.Conv2d(
+ channel_in,
+ channel_in,
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=padding,
+ groups=channel_in,
+ bias=False,
+ ),
+ nn.Conv2d(channel_in, channel_in, kernel_size=1, padding=0, bias=False),
+ norm_layer(channel_in),
+ nn.ReLU(inplace=False),
+ nn.Conv2d(
+ channel_in,
+ channel_in,
+ kernel_size=kernel_size,
+ stride=1,
+ padding=padding,
+ groups=channel_in,
+ bias=False,
+ ),
+ nn.Conv2d(channel_in, channel_out, kernel_size=1, padding=0, bias=False),
+ norm_layer(channel_out),
+ nn.ReLU(inplace=False),
+ )
+
+ def forward(self, x):
+ """SepConv forward."""
+ return self.operations(x)
+
+
+class BasicBlock(nn.Module):
+ """Basic Block for ResNet18."""
+
+ expansion = 1
+
+ def __init__(
+ self, inplanes, planes, stride=1, downsample=None, norm_layer=None
+ ): # pylint: disable=too-many-arguments
+ super().__init__()
+ self.conv1 = conv3x3(inplanes, planes, stride)
+ self.bn1 = norm_layer(planes)
+ self.relu = nn.ReLU(inplace=True)
+ self.conv2 = conv3x3(planes, planes)
+ self.bn2 = norm_layer(planes)
+ self.downsample = downsample
+ self.stride = stride
+
+ def forward(self, x):
+ """BasicBlock forward."""
+ residual = x
+
+ output = self.conv1(x)
+ output = self.bn1(output)
+ output = self.relu(output)
+
+ output = self.conv2(output)
+ output = self.bn2(output)
+
+ if self.downsample is not None:
+ residual = self.downsample(x)
+
+ output += residual
+ output = self.relu(output)
+ return output
+
+
+class MultiResnet(nn.Module): # pylint: disable=too-many-instance-attributes
+ """Resnet model.
+
+ Args:
+ block (class): block type, BasicBlock or BottleneckBlock
+ layers (int list): layer num in each block
+ n_blocks (int) : Depth of network
+ num_classes (int): class num.
+ norm_layer (class): type of normalization layer.
+ """
+
+ def __init__( # pylint: disable=too-many-arguments
+ self,
+ block,
+ layers,
+ n_blocks,
+ num_classes=1000,
+ norm_layer=MyBatchNorm,
+ ):
+ super().__init__()
+ self.n_blocks = n_blocks
+ self.inplanes = 64
+ self.norm_layer = norm_layer
+ self.conv1 = nn.Conv2d(
+ 3, self.inplanes, kernel_size=3, stride=1, padding=1, bias=False
+ )
+ self.bn1 = norm_layer(self.inplanes)
+
+ self.relu = nn.ReLU(inplace=True)
+ # self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
+
+ self.layer1 = self._make_layer(block, 64, layers[0])
+
+ self.middle_fc1 = nn.Linear(512 * block.expansion, num_classes)
+ # self.feature_fc1 = nn.Linear(512 * block.expansion, 512 * block.expansion)
+ self.scala1 = nn.Sequential(
+ SepConv(
+ channel_in=64 * block.expansion,
+ channel_out=128 * block.expansion,
+ norm_layer=norm_layer,
+ ),
+ SepConv(
+ channel_in=128 * block.expansion,
+ channel_out=256 * block.expansion,
+ norm_layer=norm_layer,
+ ),
+ SepConv(
+ channel_in=256 * block.expansion,
+ channel_out=512 * block.expansion,
+ norm_layer=norm_layer,
+ ),
+ nn.AdaptiveAvgPool2d(1),
+ )
+
+ self.attention1 = nn.Sequential(
+ SepConv(
+ channel_in=64 * block.expansion,
+ channel_out=64 * block.expansion,
+ norm_layer=norm_layer,
+ ),
+ norm_layer(64 * block.expansion),
+ nn.ReLU(),
+ nn.Upsample(scale_factor=2, mode="bilinear", align_corners=False),
+ nn.Sigmoid(),
+ )
+
+ if n_blocks > 1:
+ self.layer2 = self._make_layer(block, 128, layers[1], stride=2)
+ self.middle_fc2 = nn.Linear(512 * block.expansion, num_classes)
+ # self.feature_fc2 = nn.Linear(512 * block.expansion, 512 * block.expansion)
+ self.scala2 = nn.Sequential(
+ SepConv(
+ channel_in=128 * block.expansion,
+ channel_out=256 * block.expansion,
+ norm_layer=norm_layer,
+ ),
+ SepConv(
+ channel_in=256 * block.expansion,
+ channel_out=512 * block.expansion,
+ norm_layer=norm_layer,
+ ),
+ nn.AdaptiveAvgPool2d(1),
+ )
+ self.attention2 = nn.Sequential(
+ SepConv(
+ channel_in=128 * block.expansion,
+ channel_out=128 * block.expansion,
+ norm_layer=norm_layer,
+ ),
+ norm_layer(128 * block.expansion),
+ nn.ReLU(),
+ nn.Upsample(scale_factor=2, mode="bilinear", align_corners=False),
+ nn.Sigmoid(),
+ )
+
+ if n_blocks > 2:
+ self.layer3 = self._make_layer(block, 256, layers[2], stride=2)
+ self.middle_fc3 = nn.Linear(512 * block.expansion, num_classes)
+ # self.feature_fc3 = nn.Linear(512 * block.expansion, 512 * block.expansion)
+ self.scala3 = nn.Sequential(
+ SepConv(
+ channel_in=256 * block.expansion,
+ channel_out=512 * block.expansion,
+ norm_layer=norm_layer,
+ ),
+ nn.AdaptiveAvgPool2d(1),
+ )
+ self.attention3 = nn.Sequential(
+ SepConv(
+ channel_in=256 * block.expansion,
+ channel_out=256 * block.expansion,
+ norm_layer=norm_layer,
+ ),
+ norm_layer(256 * block.expansion),
+ nn.ReLU(),
+ nn.Upsample(scale_factor=2, mode="bilinear", align_corners=False),
+ nn.Sigmoid(),
+ )
+
+ if n_blocks > 3:
+ self.layer4 = self._make_layer(block, 512, layers[3], stride=2)
+ self.fc_layer = nn.Linear(512 * block.expansion, num_classes)
+ self.scala4 = nn.AdaptiveAvgPool2d(1)
+
+ for module in self.modules():
+ if isinstance(module, nn.Conv2d):
+ nn.init.kaiming_normal_(
+ module.weight, mode="fan_out", nonlinearity="relu"
+ )
+ elif isinstance(module, (nn.BatchNorm2d, nn.GroupNorm)):
+ nn.init.constant_(module.weight, 1)
+ nn.init.constant_(module.bias, 0)
+
+ def _make_layer(
+ self, block, planes, layers, stride=1, norm_layer=None
+ ): # pylint: disable=too-many-arguments
+ """Create a block with layers.
+
+ Args:
+ block (class): block type
+ planes (int): output channels = planes * expansion
+ layers (int): layer num in the block
+ stride (int): the first layer stride in the block.
+ norm_layer (class): type of normalization layer.
+ """
+ norm_layer = self.norm_layer
+ downsample = None
+ if stride != 1 or self.inplanes != planes * block.expansion:
+ downsample = nn.Sequential(
+ conv1x1(self.inplanes, planes * block.expansion, stride),
+ norm_layer(planes * block.expansion),
+ )
+ layer = []
+ layer.append(
+ block(
+ self.inplanes,
+ planes,
+ stride=stride,
+ downsample=downsample,
+ norm_layer=norm_layer,
+ )
+ )
+ self.inplanes = planes * block.expansion
+ for _i in range(1, layers):
+ layer.append(block(self.inplanes, planes, norm_layer=norm_layer))
+
+ return nn.Sequential(*layer)
+
+ def forward(self, x):
+ """Resnet forward."""
+ x = self.conv1(x)
+ x = self.bn1(x)
+ x = self.relu(x)
+ # x = self.maxpool(x)
+
+ x = self.layer1(x)
+ fea1 = self.attention1(x)
+ fea1 = fea1 * x
+ out1_feature = self.scala1(fea1).view(x.size(0), -1)
+ middle_output1 = self.middle_fc1(out1_feature)
+ # out1_feature = self.feature_fc1(out1_feature)
+
+ if self.n_blocks == 1:
+ return [middle_output1]
+
+ x = self.layer2(x)
+ fea2 = self.attention2(x)
+ fea2 = fea2 * x
+ out2_feature = self.scala2(fea2).view(x.size(0), -1)
+ middle_output2 = self.middle_fc2(out2_feature)
+ # out2_feature = self.feature_fc2(out2_feature)
+ if self.n_blocks == 2:
+ return [middle_output1, middle_output2]
+
+ x = self.layer3(x)
+ fea3 = self.attention3(x)
+ fea3 = fea3 * x
+ out3_feature = self.scala3(fea3).view(x.size(0), -1)
+ middle_output3 = self.middle_fc3(out3_feature)
+ # out3_feature = self.feature_fc3(out3_feature)
+
+ if self.n_blocks == 3:
+ return [middle_output1, middle_output2, middle_output3]
+
+ x = self.layer4(x)
+ out4_feature = self.scala4(x).view(x.size(0), -1)
+ output4 = self.fc_layer(out4_feature)
+
+ return [middle_output1, middle_output2, middle_output3, output4]
+
+
+def multi_resnet18(n_blocks=1, norm="bn", num_classes=100):
+ """Create resnet18 for HeteroFL.
+
+ Parameters
+ ----------
+ n_blocks: int
+ depth of network
+ norm: str
+ normalization layer type
+ num_classes: int
+ # of labels
+
+ Returns
+ -------
+ Callable [ [nn.Module,List[int],int,int,nn.Module], nn.Module]
+ """
+ if norm == "gn":
+ norm_layer = MyGroupNorm
+
+ elif norm == "bn":
+ norm_layer = MyBatchNorm
+
+ return MultiResnet(
+ BasicBlock,
+ [2, 2, 2, 2],
+ n_blocks,
+ num_classes=num_classes,
+ norm_layer=norm_layer,
+ )
+
+
+# if __name__ == "__main__":
+# from ptflops import get_model_complexity_info
+
+# model = MultiResnet18(n_blocks=4, num_classes=100)
+
+# with torch.cuda.device(0):
+# macs, params = get_model_complexity_info(
+# model,
+# (3, 32, 32),
+# as_strings=True,
+# print_per_layer_stat=False,
+# verbose=True,
+# units="MMac",
+# )
+
+# print("{:<30} {:<8}".format("Computational complexity: ", macs))
+# print("{:<30} {:<8}".format("Number of parameters: ", params))
diff --git a/baselines/depthfl/depthfl/resnet_hetero.py b/baselines/depthfl/depthfl/resnet_hetero.py
new file mode 100644
index 000000000000..a84c07b881b2
--- /dev/null
+++ b/baselines/depthfl/depthfl/resnet_hetero.py
@@ -0,0 +1,280 @@
+"""ResNet18 for HeteroFL."""
+
+import numpy as np
+import torch.nn as nn
+
+
+class Scaler(nn.Module):
+ """Scaler module for HeteroFL."""
+
+ def __init__(self, rate, scale):
+ super().__init__()
+ if scale:
+ self.rate = rate
+ else:
+ self.rate = 1
+
+ def forward(self, x):
+ """Scaler forward."""
+ output = x / self.rate if self.training else x
+ return output
+
+
+class MyBatchNorm(nn.Module):
+ """Static Batch Normalization for HeteroFL."""
+
+ def __init__(self, num_channels, track=True):
+ super().__init__()
+ # change num_groups to 32
+ self.norm = nn.BatchNorm2d(num_channels, track_running_stats=track)
+
+ def forward(self, x):
+ """BatchNorm forward."""
+ x = self.norm(x)
+ return x
+
+
+def conv3x3(in_planes, out_planes, stride=1):
+ """Convolution layer 3x3."""
+ return nn.Conv2d(
+ in_planes, out_planes, kernel_size=3, stride=stride, padding=1, bias=False
+ )
+
+
+def conv1x1(in_planes, planes, stride=1):
+ """Convolution layer 1x1."""
+ return nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride, bias=False)
+
+
+class BasicBlock(nn.Module): # pylint: disable=too-many-instance-attributes
+ """Basic Block for ResNet18."""
+
+ expansion = 1
+
+ def __init__( # pylint: disable=too-many-arguments
+ self,
+ inplanes,
+ planes,
+ stride=1,
+ scaler_rate=1,
+ downsample=None,
+ track=True,
+ scale=True,
+ ):
+ super().__init__()
+ self.conv1 = conv3x3(inplanes, planes, stride)
+ self.scaler = Scaler(scaler_rate, scale)
+ self.bn1 = MyBatchNorm(planes, track)
+ self.relu = nn.ReLU(inplace=True)
+ self.conv2 = conv3x3(planes, planes)
+ self.bn2 = MyBatchNorm(planes, track)
+ self.downsample = downsample
+ self.stride = stride
+
+ def forward(self, x):
+ """BasicBlock forward."""
+ residual = x
+
+ output = self.conv1(x)
+ output = self.scaler(output)
+ output = self.bn1(output)
+ output = self.relu(output)
+
+ output = self.conv2(output)
+ output = self.scaler(output)
+ output = self.bn2(output)
+
+ if self.downsample is not None:
+ residual = self.downsample(x)
+
+ output += residual
+ output = self.relu(output)
+ return output
+
+
+class Resnet(nn.Module): # pylint: disable=too-many-instance-attributes
+ """Resnet model."""
+
+ def __init__( # pylint: disable=too-many-arguments
+ self, hidden_size, block, layers, num_classes, scaler_rate, track, scale
+ ):
+ super().__init__()
+
+ self.inplanes = hidden_size[0]
+ self.norm_layer = MyBatchNorm
+ self.conv1 = nn.Conv2d(
+ 3, self.inplanes, kernel_size=3, stride=1, padding=1, bias=False
+ )
+ self.scaler = Scaler(scaler_rate, scale)
+ self.bn1 = self.norm_layer(self.inplanes, track)
+
+ self.relu = nn.ReLU(inplace=True)
+ # self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
+
+ self.layer1 = self._make_layer(
+ block,
+ hidden_size[0],
+ layers[0],
+ scaler_rate=scaler_rate,
+ track=track,
+ scale=scale,
+ )
+ self.layer2 = self._make_layer(
+ block,
+ hidden_size[1],
+ layers[1],
+ stride=2,
+ scaler_rate=scaler_rate,
+ track=track,
+ scale=scale,
+ )
+ self.layer3 = self._make_layer(
+ block,
+ hidden_size[2],
+ layers[2],
+ stride=2,
+ scaler_rate=scaler_rate,
+ track=track,
+ scale=scale,
+ )
+ self.layer4 = self._make_layer(
+ block,
+ hidden_size[3],
+ layers[3],
+ stride=2,
+ scaler_rate=scaler_rate,
+ track=track,
+ scale=scale,
+ )
+ self.fc_layer = nn.Linear(hidden_size[3] * block.expansion, num_classes)
+ self.scala = nn.AdaptiveAvgPool2d(1)
+
+ for module in self.modules():
+ if isinstance(module, nn.Conv2d):
+ nn.init.kaiming_normal_(
+ module.weight, mode="fan_out", nonlinearity="relu"
+ )
+ elif isinstance(module, (nn.BatchNorm2d, nn.GroupNorm)):
+ nn.init.constant_(module.weight, 1)
+ nn.init.constant_(module.bias, 0)
+
+ def _make_layer( # pylint: disable=too-many-arguments
+ self, block, planes, layers, stride=1, scaler_rate=1, track=True, scale=True
+ ):
+ """Create a block with layers.
+
+ Args:
+ block (class): block type
+ planes (int): output channels = planes * expansion
+ layers (int): layer num in the block
+ stride (int): the first layer stride in the block.
+ scaler_rate (float): for scaler module
+ track (bool): static batch normalization
+ scale (bool): for scaler module.
+ """
+ norm_layer = self.norm_layer
+ downsample = None
+ if stride != 1 or self.inplanes != planes * block.expansion:
+ downsample = nn.Sequential(
+ conv1x1(self.inplanes, planes * block.expansion, stride),
+ norm_layer(planes * block.expansion, track),
+ )
+ layer = []
+ layer.append(
+ block(
+ self.inplanes,
+ planes,
+ stride=stride,
+ scaler_rate=scaler_rate,
+ downsample=downsample,
+ track=track,
+ scale=scale,
+ )
+ )
+ self.inplanes = planes * block.expansion
+ for _i in range(1, layers):
+ layer.append(
+ block(
+ self.inplanes,
+ planes,
+ scaler_rate=scaler_rate,
+ track=track,
+ scale=scale,
+ )
+ )
+
+ return nn.Sequential(*layer)
+
+ def forward(self, x):
+ """Resnet forward."""
+ x = self.conv1(x)
+ x = self.scaler(x)
+ x = self.bn1(x)
+ x = self.relu(x)
+ # x = self.maxpool(x)
+
+ x = self.layer1(x)
+ x = self.layer2(x)
+ x = self.layer3(x)
+ x = self.layer4(x)
+ out = self.scala(x).view(x.size(0), -1)
+ out = self.fc_layer(out)
+
+ return [out]
+
+
+def resnet18(n_blocks=4, track=False, scale=True, num_classes=100):
+ """Create resnet18 for HeteroFL.
+
+ Parameters
+ ----------
+ n_blocks: int
+ corresponds to width (divided by 4)
+ track: bool
+ static batch normalization
+ scale: bool
+ scaler module
+ num_classes: int
+ # of labels
+
+ Returns
+ -------
+ Callable [ [List[int],nn.Module,List[int],int,float,bool,bool], nn.Module]
+ """
+ # width pruning ratio : (0.25, 0.50, 0.75, 0.10)
+ model_rate = n_blocks / 4
+ classes_size = num_classes
+
+ hidden_size = [64, 128, 256, 512]
+ hidden_size = [int(np.ceil(model_rate * x)) for x in hidden_size]
+
+ scaler_rate = model_rate
+
+ return Resnet(
+ hidden_size,
+ BasicBlock,
+ [2, 2, 2, 2],
+ num_classes=classes_size,
+ scaler_rate=scaler_rate,
+ track=track,
+ scale=scale,
+ )
+
+
+# if __name__ == "__main__":
+# from ptflops import get_model_complexity_info
+
+# model = resnet18(100, 1.0)
+
+# with torch.cuda.device(0):
+# macs, params = get_model_complexity_info(
+# model,
+# (3, 32, 32),
+# as_strings=True,
+# print_per_layer_stat=False,
+# verbose=True,
+# units="MMac",
+# )
+
+# print("{:<30} {:<8}".format("Computational complexity: ", macs))
+# print("{:<30} {:<8}".format("Number of parameters: ", params))
diff --git a/baselines/depthfl/depthfl/server.py b/baselines/depthfl/depthfl/server.py
new file mode 100644
index 000000000000..dc99ae2fc5de
--- /dev/null
+++ b/baselines/depthfl/depthfl/server.py
@@ -0,0 +1,209 @@
+"""Server for DepthFL baseline."""
+
+import copy
+from collections import OrderedDict
+from logging import DEBUG, INFO
+from typing import Callable, Dict, List, Optional, Tuple, Union
+
+import torch
+from flwr.common import FitRes, Parameters, Scalar, parameters_to_ndarrays
+from flwr.common.logger import log
+from flwr.common.typing import NDArrays
+from flwr.server.client_proxy import ClientProxy
+from flwr.server.server import Server, fit_clients
+from hydra.utils import instantiate
+from omegaconf import DictConfig
+from torch.utils.data import DataLoader
+
+from depthfl.client import prune
+from depthfl.models import test, test_sbn
+from depthfl.strategy import aggregate_fit_depthfl
+from depthfl.strategy_hetero import aggregate_fit_hetero
+
+FitResultsAndFailures = Tuple[
+ List[Tuple[ClientProxy, FitRes]],
+ List[Union[Tuple[ClientProxy, FitRes], BaseException]],
+]
+
+
+def gen_evaluate_fn(
+ testloader: DataLoader,
+ device: torch.device,
+ model: DictConfig,
+) -> Callable[
+ [int, NDArrays, Dict[str, Scalar]],
+ Tuple[float, Dict[str, Union[Scalar, List[float]]]],
+]:
+ """Generate the function for centralized evaluation.
+
+ Parameters
+ ----------
+ testloader : DataLoader
+ The dataloader to test the model with.
+ device : torch.device
+ The device to test the model on.
+ model : DictConfig
+ model configuration for instantiating
+
+ Returns
+ -------
+ Callable[ [int, NDArrays, Dict[str, Scalar]],
+ Optional[Tuple[float, Dict[str, Scalar]]] ]
+ The centralized evaluation function.
+ """
+
+ def evaluate(
+ server_round: int, parameters_ndarrays: NDArrays, config: Dict[str, Scalar]
+ ) -> Tuple[float, Dict[str, Union[Scalar, List[float]]]]:
+ # pylint: disable=unused-argument
+ """Use the entire CIFAR-100 test set for evaluation."""
+ net = instantiate(model)
+ params_dict = zip(net.state_dict().keys(), parameters_ndarrays)
+ state_dict = OrderedDict({k: torch.tensor(v) for k, v in params_dict})
+ net.load_state_dict(state_dict, strict=True)
+ net.to(device)
+
+ loss, accuracy, accuracy_single = test(net, testloader, device=device)
+ # return statistics
+ return loss, {"accuracy": accuracy, "accuracy_single": accuracy_single}
+
+ return evaluate
+
+
+def gen_evaluate_fn_hetero(
+ trainloaders: List[DataLoader],
+ testloader: DataLoader,
+ device: torch.device,
+ model_cfg: DictConfig,
+) -> Callable[
+ [int, NDArrays, Dict[str, Scalar]],
+ Tuple[float, Dict[str, Union[Scalar, List[float]]]],
+]:
+ """Generate the function for centralized evaluation.
+
+ Parameters
+ ----------
+ trainloaders : List[DataLoader]
+ The list of dataloaders to calculate statistics for BN
+ testloader : DataLoader
+ The dataloader to test the model with.
+ device : torch.device
+ The device to test the model on.
+ model_cfg : DictConfig
+ model configuration for instantiating
+
+ Returns
+ -------
+ Callable[ [int, NDArrays, Dict[str, Scalar]],
+ Optional[Tuple[float, Dict[str, Scalar]]] ]
+ The centralized evaluation function.
+ """
+
+ def evaluate( # pylint: disable=too-many-locals
+ server_round: int, parameters_ndarrays: NDArrays, config: Dict[str, Scalar]
+ ) -> Tuple[float, Dict[str, Union[Scalar, List[float]]]]:
+ # pylint: disable=unused-argument
+ """Use the entire CIFAR-100 test set for evaluation."""
+ # test per 50 rounds (sbn takes a long time)
+ if server_round % 50 != 0:
+ return 0.0, {"accuracy": 0.0, "accuracy_single": [0] * 4}
+
+ # models with different width
+ models = []
+ for i in range(4):
+ model_tmp = copy.deepcopy(model_cfg)
+ model_tmp.n_blocks = i + 1
+ models.append(model_tmp)
+
+ # load global parameters
+ param_idx_lst = []
+ nets = []
+ net_tmp = instantiate(models[-1], track=False)
+ for model in models:
+ net = instantiate(model, track=True, scale=False)
+ nets.append(net)
+ param_idx = {}
+ for k in net_tmp.state_dict().keys():
+ param_idx[k] = [
+ torch.arange(size) for size in net.state_dict()[k].shape
+ ]
+ param_idx_lst.append(param_idx)
+
+ params_dict = zip(net_tmp.state_dict().keys(), parameters_ndarrays)
+ state_dict = OrderedDict({k: torch.tensor(v) for k, v in params_dict})
+
+ for net, param_idx in zip(nets, param_idx_lst):
+ net.load_state_dict(prune(state_dict, param_idx), strict=False)
+ net.to(device)
+ net.train()
+
+ loss, accuracy, accuracy_single = test_sbn(
+ nets, trainloaders, testloader, device=device
+ )
+ # return statistics
+ return loss, {"accuracy": accuracy, "accuracy_single": accuracy_single}
+
+ return evaluate
+
+
+class ServerFedDyn(Server):
+ """Sever for FedDyn."""
+
+ def fit_round(
+ self,
+ server_round: int,
+ timeout: Optional[float],
+ ) -> Optional[
+ Tuple[Optional[Parameters], Dict[str, Scalar], FitResultsAndFailures]
+ ]:
+ """Perform a single round."""
+ # Get clients and their respective instructions from strategy
+ client_instructions = self.strategy.configure_fit(
+ server_round=server_round,
+ parameters=self.parameters,
+ client_manager=self._client_manager,
+ )
+
+ if not client_instructions:
+ log(INFO, "fit_round %s: no clients selected, cancel", server_round)
+ return None
+ log(
+ DEBUG,
+ "fit_round %s: strategy sampled %s clients (out of %s)",
+ server_round,
+ len(client_instructions),
+ self._client_manager.num_available(),
+ )
+
+ # Collect `fit` results from all clients participating in this round
+ results, failures = fit_clients(
+ client_instructions=client_instructions,
+ max_workers=self.max_workers,
+ timeout=timeout,
+ )
+ log(
+ DEBUG,
+ "fit_round %s received %s results and %s failures",
+ server_round,
+ len(results),
+ len(failures),
+ )
+
+ if "HeteroFL" in str(type(self.strategy)):
+ aggregate_fit = aggregate_fit_hetero
+ else:
+ aggregate_fit = aggregate_fit_depthfl
+
+ aggregated_result: Tuple[
+ Optional[Parameters],
+ Dict[str, Scalar],
+ ] = aggregate_fit(
+ self.strategy,
+ server_round,
+ results,
+ failures,
+ parameters_to_ndarrays(self.parameters),
+ )
+
+ parameters_aggregated, metrics_aggregated = aggregated_result
+ return parameters_aggregated, metrics_aggregated, (results, failures)
diff --git a/baselines/depthfl/depthfl/strategy.py b/baselines/depthfl/depthfl/strategy.py
new file mode 100644
index 000000000000..3414c28c4518
--- /dev/null
+++ b/baselines/depthfl/depthfl/strategy.py
@@ -0,0 +1,136 @@
+"""Strategy for DepthFL."""
+
+import os
+import pickle
+from logging import WARNING
+from typing import Dict, List, Optional, Tuple, Union
+
+import numpy as np
+import torch
+import torch.nn as nn
+from flwr.common import (
+ NDArrays,
+ Parameters,
+ Scalar,
+ ndarrays_to_parameters,
+ parameters_to_ndarrays,
+)
+from flwr.common.logger import log
+from flwr.common.typing import FitRes
+from flwr.server.client_proxy import ClientProxy
+from flwr.server.strategy import FedAvg
+from omegaconf import DictConfig
+
+
+class FedDyn(FedAvg):
+ """Applying dynamic regularization in FedDyn paper."""
+
+ def __init__(self, cfg: DictConfig, net: nn.Module, *args, **kwargs):
+ self.cfg = cfg
+ self.h_variate = [np.zeros(v.shape) for (k, v) in net.state_dict().items()]
+
+ # tagging real weights / biases
+ self.is_weight = []
+ for k in net.state_dict().keys():
+ if "weight" not in k and "bias" not in k:
+ self.is_weight.append(False)
+ else:
+ self.is_weight.append(True)
+
+ # prev_grads file for each client
+ prev_grads = [
+ {k: torch.zeros(v.numel()) for (k, v) in net.named_parameters()}
+ ] * cfg.num_clients
+
+ if not os.path.exists("prev_grads"):
+ os.makedirs("prev_grads")
+
+ for idx in range(cfg.num_clients):
+ with open(f"prev_grads/client_{idx}", "wb") as prev_grads_file:
+ pickle.dump(prev_grads[idx], prev_grads_file)
+
+ super().__init__(*args, **kwargs)
+
+
+def aggregate_fit_depthfl(
+ strategy,
+ server_round: int,
+ results: List[Tuple[ClientProxy, FitRes]],
+ failures: List[Union[Tuple[ClientProxy, FitRes], BaseException]],
+ origin: NDArrays,
+) -> Tuple[Optional[Parameters], Dict[str, Scalar]]:
+ """Aggregate fit results using weighted average."""
+ if not results:
+ return None, {}
+ # Do not aggregate if there are failures and failures are not accepted
+ if not strategy.accept_failures and failures:
+ return None, {}
+
+ # Convert results
+ weights_results = [
+ (parameters_to_ndarrays(fit_res.parameters), fit_res.num_examples)
+ for _, fit_res in results
+ ]
+ parameters_aggregated = ndarrays_to_parameters(
+ aggregate(
+ weights_results,
+ origin,
+ strategy.h_variate,
+ strategy.is_weight,
+ strategy.cfg,
+ )
+ )
+
+ # Aggregate custom metrics if aggregation fn was provided
+ metrics_aggregated = {}
+ if strategy.fit_metrics_aggregation_fn:
+ fit_metrics = [(res.num_examples, res.metrics) for _, res in results]
+ metrics_aggregated = strategy.fit_metrics_aggregation_fn(fit_metrics)
+ elif server_round == 1: # Only log this warning once
+ log(WARNING, "No fit_metrics_aggregation_fn provided")
+
+ return parameters_aggregated, metrics_aggregated
+
+
+def aggregate(
+ results: List[Tuple[NDArrays, int]],
+ origin: NDArrays,
+ h_list: List,
+ is_weight: List,
+ cfg: DictConfig,
+) -> NDArrays:
+ """Aggregate model parameters with different depths."""
+ param_count = [0] * len(origin)
+ weights_sum = [np.zeros(v.shape) for v in origin]
+
+ # summation & counting of parameters
+ for parameters, _ in results:
+ for i, layer in enumerate(parameters):
+ weights_sum[i] += layer
+ param_count[i] += 1
+
+ # update parameters
+ for i, weight in enumerate(weights_sum):
+ if param_count[i] > 0:
+ weight = weight / param_count[i]
+ # print(np.isscalar(weight))
+
+ # update h variable for FedDyn
+ h_list[i] = (
+ h_list[i]
+ - cfg.fit_config.alpha
+ * param_count[i]
+ * (weight - origin[i])
+ / cfg.num_clients
+ )
+
+ # applying h only for weights / biases
+ if is_weight[i] and cfg.fit_config.feddyn:
+ weights_sum[i] = weight - h_list[i] / cfg.fit_config.alpha
+ else:
+ weights_sum[i] = weight
+
+ else:
+ weights_sum[i] = origin[i]
+
+ return weights_sum
diff --git a/baselines/depthfl/depthfl/strategy_hetero.py b/baselines/depthfl/depthfl/strategy_hetero.py
new file mode 100644
index 000000000000..7544204cde2f
--- /dev/null
+++ b/baselines/depthfl/depthfl/strategy_hetero.py
@@ -0,0 +1,136 @@
+"""Strategy for HeteroFL."""
+
+import os
+import pickle
+from logging import WARNING
+from typing import Dict, List, Optional, Tuple, Union
+
+import numpy as np
+import torch
+import torch.nn as nn
+from flwr.common import (
+ NDArrays,
+ Parameters,
+ Scalar,
+ ndarrays_to_parameters,
+ parameters_to_ndarrays,
+)
+from flwr.common.logger import log
+from flwr.common.typing import FitRes
+from flwr.server.client_proxy import ClientProxy
+from flwr.server.strategy import FedAvg
+from hydra.utils import instantiate
+from omegaconf import DictConfig
+
+
+class HeteroFL(FedAvg):
+ """Custom FedAvg for HeteroFL."""
+
+ def __init__(self, cfg: DictConfig, net: nn.Module, *args, **kwargs):
+ self.cfg = cfg
+ self.parameters = [np.zeros(v.shape) for (k, v) in net.state_dict().items()]
+ self.param_idx_lst = []
+
+ model = cfg.model
+ # store parameter shapes of different width
+ for i in range(4):
+ model.n_blocks = i + 1
+ net_tmp = instantiate(model)
+ param_idx = []
+ for k in net_tmp.state_dict().keys():
+ param_idx.append(
+ [torch.arange(size) for size in net_tmp.state_dict()[k].shape]
+ )
+
+ # print(net_tmp.state_dict()['conv1.weight'].shape[0])
+ self.param_idx_lst.append(param_idx)
+
+ self.is_weight = []
+
+ # tagging real weights / biases
+ for k in net.state_dict().keys():
+ if "num" in k:
+ self.is_weight.append(False)
+ else:
+ self.is_weight.append(True)
+
+ # prev_grads file for each client
+ prev_grads = [
+ {k: torch.zeros(v.numel()) for (k, v) in net.named_parameters()}
+ ] * cfg.num_clients
+
+ if not os.path.exists("prev_grads"):
+ os.makedirs("prev_grads")
+
+ for idx in range(cfg.num_clients):
+ with open(f"prev_grads/client_{idx}", "wb") as prev_grads_file:
+ pickle.dump(prev_grads[idx], prev_grads_file)
+
+ super().__init__(*args, **kwargs)
+
+ def aggregate_hetero(
+ self, results: List[Tuple[NDArrays, Union[bool, bytes, float, int, str]]]
+ ):
+ """Aggregate function for HeteroFL."""
+ for i, params in enumerate(self.parameters):
+ count = np.zeros(params.shape)
+ tmp_v = np.zeros(params.shape)
+ if self.is_weight[i]:
+ for weights, cid in results:
+ if self.cfg.exclusive_learning:
+ cid = self.cfg.model_size * (self.cfg.num_clients // 4) - 1
+
+ tmp_v[
+ torch.meshgrid(
+ self.param_idx_lst[cid // (self.cfg.num_clients // 4)][i]
+ )
+ ] += weights[i]
+ count[
+ torch.meshgrid(
+ self.param_idx_lst[cid // (self.cfg.num_clients // 4)][i]
+ )
+ ] += 1
+ tmp_v[count > 0] = np.divide(tmp_v[count > 0], count[count > 0])
+ params[count > 0] = tmp_v[count > 0]
+
+ else:
+ for weights, _ in results:
+ tmp_v += weights[i]
+ count += 1
+ tmp_v = np.divide(tmp_v, count)
+ params = tmp_v
+
+
+def aggregate_fit_hetero(
+ strategy,
+ server_round: int,
+ results: List[Tuple[ClientProxy, FitRes]],
+ failures: List[Union[Tuple[ClientProxy, FitRes], BaseException]],
+ origin: NDArrays,
+) -> Tuple[Optional[Parameters], Dict[str, Scalar]]:
+ """Aggregate fit results using weighted average."""
+ if not results:
+ return None, {}
+ # Do not aggregate if there are failures and failures are not accepted
+ if not strategy.accept_failures and failures:
+ return None, {}
+
+ # Convert results
+ weights_results = [
+ (parameters_to_ndarrays(fit_res.parameters), fit_res.metrics["cid"])
+ for _, fit_res in results
+ ]
+
+ strategy.parameters = origin
+ strategy.aggregate_hetero(weights_results)
+ parameters_aggregated = ndarrays_to_parameters(strategy.parameters)
+
+ # Aggregate custom metrics if aggregation fn was provided
+ metrics_aggregated = {}
+ if strategy.fit_metrics_aggregation_fn:
+ fit_metrics = [(res.num_examples, res.metrics) for _, res in results]
+ metrics_aggregated = strategy.fit_metrics_aggregation_fn(fit_metrics)
+ elif server_round == 1: # Only log this warning once
+ log(WARNING, "No fit_metrics_aggregation_fn provided")
+
+ return parameters_aggregated, metrics_aggregated
diff --git a/baselines/depthfl/depthfl/utils.py b/baselines/depthfl/depthfl/utils.py
new file mode 100644
index 000000000000..fad2afcad4be
--- /dev/null
+++ b/baselines/depthfl/depthfl/utils.py
@@ -0,0 +1,66 @@
+"""Contains utility functions for CNN FL on MNIST."""
+
+import pickle
+from pathlib import Path
+from secrets import token_hex
+from typing import Dict, Union
+
+from flwr.server.history import History
+
+
+def save_results_as_pickle(
+ history: History,
+ file_path: Union[str, Path],
+ extra_results: Dict,
+ default_filename: str = "results.pkl",
+) -> None:
+ """Save results from simulation to pickle.
+
+ Parameters
+ ----------
+ history: History
+ History returned by start_simulation.
+ file_path: Union[str, Path]
+ Path to file to create and store both history and extra_results.
+ If path is a directory, the default_filename will be used.
+ path doesn't exist, it will be created. If file exists, a
+ randomly generated suffix will be added to the file name. This
+ is done to avoid overwritting results.
+ extra_results : Dict
+ A dictionary containing additional results you would like
+ to be saved to disk. Default: {} (an empty dictionary)
+ default_filename: Optional[str]
+ File used by default if file_path points to a directory instead
+ to a file. Default: "results.pkl"
+ """
+ path = Path(file_path)
+
+ # ensure path exists
+ path.mkdir(exist_ok=True, parents=True)
+
+ def _add_random_suffix(path_: Path):
+ """Add a randomly generated suffix to the file name."""
+ print(f"File `{path_}` exists! ")
+ suffix = token_hex(4)
+ print(f"New results to be saved with suffix: {suffix}")
+ return path_.parent / (path_.stem + "_" + suffix + ".pkl")
+
+ def _complete_path_with_default_name(path_: Path):
+ """Append the default file name to the path."""
+ print("Using default filename")
+ return path_ / default_filename
+
+ if path.is_dir():
+ path = _complete_path_with_default_name(path)
+
+ if path.is_file():
+ # file exists already
+ path = _add_random_suffix(path)
+
+ print(f"Results will be saved into: {path}")
+
+ data = {"history": history, **extra_results}
+
+ # save results to pickle
+ with open(str(path), "wb") as handle:
+ pickle.dump(data, handle, protocol=pickle.HIGHEST_PROTOCOL)
diff --git a/baselines/depthfl/pyproject.toml b/baselines/depthfl/pyproject.toml
new file mode 100644
index 000000000000..2f928c2d3553
--- /dev/null
+++ b/baselines/depthfl/pyproject.toml
@@ -0,0 +1,141 @@
+[build-system]
+requires = ["poetry-core>=1.4.0"]
+build-backend = "poetry.masonry.api"
+
+[tool.poetry]
+name = "depthfl" # <----- Ensure it matches the name of your baseline directory containing all the source code
+version = "1.0.0"
+description = "DepthFL: Depthwise Federated Learning for Heterogeneous Clients"
+license = "Apache-2.0"
+authors = ["Minjae Kim "]
+readme = "README.md"
+homepage = "https://flower.dev"
+repository = "https://github.com/adap/flower"
+documentation = "https://flower.dev"
+classifiers = [
+ "Development Status :: 3 - Alpha",
+ "Intended Audience :: Developers",
+ "Intended Audience :: Science/Research",
+ "License :: OSI Approved :: Apache Software License",
+ "Operating System :: MacOS :: MacOS X",
+ "Operating System :: POSIX :: Linux",
+ "Programming Language :: Python",
+ "Programming Language :: Python :: 3",
+ "Programming Language :: Python :: 3 :: Only",
+ "Programming Language :: Python :: 3.8",
+ "Programming Language :: Python :: 3.9",
+ "Programming Language :: Python :: 3.10",
+ "Programming Language :: Python :: 3.11",
+ "Programming Language :: Python :: Implementation :: CPython",
+ "Topic :: Scientific/Engineering",
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
+ "Topic :: Scientific/Engineering :: Mathematics",
+ "Topic :: Software Development",
+ "Topic :: Software Development :: Libraries",
+ "Topic :: Software Development :: Libraries :: Python Modules",
+ "Typing :: Typed",
+]
+
+[tool.poetry.dependencies]
+python = ">=3.10.0, <3.11.0"
+flwr = { extras = ["simulation"], version = "1.5.0" }
+hydra-core = "1.3.2" # don't change this
+matplotlib = "3.7.1"
+torch = { url = "https://download.pytorch.org/whl/cu116/torch-1.13.1%2Bcu116-cp310-cp310-linux_x86_64.whl"}
+torchvision = { url = "https://download.pytorch.org/whl/cu116/torchvision-0.14.1%2Bcu116-cp310-cp310-linux_x86_64.whl"}
+
+
+[tool.poetry.dev-dependencies]
+isort = "==5.11.5"
+black = "==23.1.0"
+docformatter = "==1.5.1"
+mypy = "==1.4.1"
+pylint = "==2.8.2"
+flake8 = "==3.9.2"
+pytest = "==6.2.4"
+pytest-watch = "==4.2.0"
+ruff = "==0.0.272"
+types-requests = "==2.27.7"
+
+[tool.isort]
+line_length = 88
+indent = " "
+multi_line_output = 3
+include_trailing_comma = true
+force_grid_wrap = 0
+use_parentheses = true
+
+[tool.black]
+line-length = 88
+target-version = ["py38", "py39", "py310", "py311"]
+
+[tool.pytest.ini_options]
+minversion = "6.2"
+addopts = "-qq"
+testpaths = [
+ "flwr_baselines",
+]
+
+[tool.mypy]
+ignore_missing_imports = true
+strict = false
+plugins = "numpy.typing.mypy_plugin"
+
+[tool.pylint."MESSAGES CONTROL"]
+disable = "bad-continuation,duplicate-code,too-few-public-methods,useless-import-alias"
+good-names = "i,j,k,_,x,y,X,Y"
+signature-mutators="hydra.main.main"
+
+[tool.pylint.typecheck]
+generated-members="numpy.*, torch.*, tensorflow.*"
+
+[[tool.mypy.overrides]]
+module = [
+ "importlib.metadata.*",
+ "importlib_metadata.*",
+]
+follow_imports = "skip"
+follow_imports_for_stubs = true
+disallow_untyped_calls = false
+
+[[tool.mypy.overrides]]
+module = "torch.*"
+follow_imports = "skip"
+follow_imports_for_stubs = true
+
+[tool.docformatter]
+wrap-summaries = 88
+wrap-descriptions = 88
+
+[tool.ruff]
+target-version = "py38"
+line-length = 88
+select = ["D", "E", "F", "W", "B", "ISC", "C4"]
+fixable = ["D", "E", "F", "W", "B", "ISC", "C4"]
+ignore = ["B024", "B027"]
+exclude = [
+ ".bzr",
+ ".direnv",
+ ".eggs",
+ ".git",
+ ".hg",
+ ".mypy_cache",
+ ".nox",
+ ".pants.d",
+ ".pytype",
+ ".ruff_cache",
+ ".svn",
+ ".tox",
+ ".venv",
+ "__pypackages__",
+ "_build",
+ "buck-out",
+ "build",
+ "dist",
+ "node_modules",
+ "venv",
+ "proto",
+]
+
+[tool.ruff.pydocstyle]
+convention = "numpy"
diff --git a/baselines/fedper/LICENSE b/baselines/fedper/LICENSE
new file mode 100644
index 000000000000..d64569567334
--- /dev/null
+++ b/baselines/fedper/LICENSE
@@ -0,0 +1,202 @@
+
+ Apache License
+ Version 2.0, January 2004
+ http://www.apache.org/licenses/
+
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
+
+ 1. Definitions.
+
+ "License" shall mean the terms and conditions for use, reproduction,
+ and distribution as defined by Sections 1 through 9 of this document.
+
+ "Licensor" shall mean the copyright owner or entity authorized by
+ the copyright owner that is granting the License.
+
+ "Legal Entity" shall mean the union of the acting entity and all
+ other entities that control, are controlled by, or are under common
+ control with that entity. For the purposes of this definition,
+ "control" means (i) the power, direct or indirect, to cause the
+ direction or management of such entity, whether by contract or
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
+ outstanding shares, or (iii) beneficial ownership of such entity.
+
+ "You" (or "Your") shall mean an individual or Legal Entity
+ exercising permissions granted by this License.
+
+ "Source" form shall mean the preferred form for making modifications,
+ including but not limited to software source code, documentation
+ source, and configuration files.
+
+ "Object" form shall mean any form resulting from mechanical
+ transformation or translation of a Source form, including but
+ not limited to compiled object code, generated documentation,
+ and conversions to other media types.
+
+ "Work" shall mean the work of authorship, whether in Source or
+ Object form, made available under the License, as indicated by a
+ copyright notice that is included in or attached to the work
+ (an example is provided in the Appendix below).
+
+ "Derivative Works" shall mean any work, whether in Source or Object
+ form, that is based on (or derived from) the Work and for which the
+ editorial revisions, annotations, elaborations, or other modifications
+ represent, as a whole, an original work of authorship. For the purposes
+ of this License, Derivative Works shall not include works that remain
+ separable from, or merely link (or bind by name) to the interfaces of,
+ the Work and Derivative Works thereof.
+
+ "Contribution" shall mean any work of authorship, including
+ the original version of the Work and any modifications or additions
+ to that Work or Derivative Works thereof, that is intentionally
+ submitted to Licensor for inclusion in the Work by the copyright owner
+ or by an individual or Legal Entity authorized to submit on behalf of
+ the copyright owner. For the purposes of this definition, "submitted"
+ means any form of electronic, verbal, or written communication sent
+ to the Licensor or its representatives, including but not limited to
+ communication on electronic mailing lists, source code control systems,
+ and issue tracking systems that are managed by, or on behalf of, the
+ Licensor for the purpose of discussing and improving the Work, but
+ excluding communication that is conspicuously marked or otherwise
+ designated in writing by the copyright owner as "Not a Contribution."
+
+ "Contributor" shall mean Licensor and any individual or Legal Entity
+ on behalf of whom a Contribution has been received by Licensor and
+ subsequently incorporated within the Work.
+
+ 2. Grant of Copyright License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ copyright license to reproduce, prepare Derivative Works of,
+ publicly display, publicly perform, sublicense, and distribute the
+ Work and such Derivative Works in Source or Object form.
+
+ 3. Grant of Patent License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ (except as stated in this section) patent license to make, have made,
+ use, offer to sell, sell, import, and otherwise transfer the Work,
+ where such license applies only to those patent claims licensable
+ by such Contributor that are necessarily infringed by their
+ Contribution(s) alone or by combination of their Contribution(s)
+ with the Work to which such Contribution(s) was submitted. If You
+ institute patent litigation against any entity (including a
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
+ or a Contribution incorporated within the Work constitutes direct
+ or contributory patent infringement, then any patent licenses
+ granted to You under this License for that Work shall terminate
+ as of the date such litigation is filed.
+
+ 4. Redistribution. You may reproduce and distribute copies of the
+ Work or Derivative Works thereof in any medium, with or without
+ modifications, and in Source or Object form, provided that You
+ meet the following conditions:
+
+ (a) You must give any other recipients of the Work or
+ Derivative Works a copy of this License; and
+
+ (b) You must cause any modified files to carry prominent notices
+ stating that You changed the files; and
+
+ (c) You must retain, in the Source form of any Derivative Works
+ that You distribute, all copyright, patent, trademark, and
+ attribution notices from the Source form of the Work,
+ excluding those notices that do not pertain to any part of
+ the Derivative Works; and
+
+ (d) If the Work includes a "NOTICE" text file as part of its
+ distribution, then any Derivative Works that You distribute must
+ include a readable copy of the attribution notices contained
+ within such NOTICE file, excluding those notices that do not
+ pertain to any part of the Derivative Works, in at least one
+ of the following places: within a NOTICE text file distributed
+ as part of the Derivative Works; within the Source form or
+ documentation, if provided along with the Derivative Works; or,
+ within a display generated by the Derivative Works, if and
+ wherever such third-party notices normally appear. The contents
+ of the NOTICE file are for informational purposes only and
+ do not modify the License. You may add Your own attribution
+ notices within Derivative Works that You distribute, alongside
+ or as an addendum to the NOTICE text from the Work, provided
+ that such additional attribution notices cannot be construed
+ as modifying the License.
+
+ You may add Your own copyright statement to Your modifications and
+ may provide additional or different license terms and conditions
+ for use, reproduction, or distribution of Your modifications, or
+ for any such Derivative Works as a whole, provided Your use,
+ reproduction, and distribution of the Work otherwise complies with
+ the conditions stated in this License.
+
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
+ any Contribution intentionally submitted for inclusion in the Work
+ by You to the Licensor shall be under the terms and conditions of
+ this License, without any additional terms or conditions.
+ Notwithstanding the above, nothing herein shall supersede or modify
+ the terms of any separate license agreement you may have executed
+ with Licensor regarding such Contributions.
+
+ 6. Trademarks. This License does not grant permission to use the trade
+ names, trademarks, service marks, or product names of the Licensor,
+ except as required for reasonable and customary use in describing the
+ origin of the Work and reproducing the content of the NOTICE file.
+
+ 7. Disclaimer of Warranty. Unless required by applicable law or
+ agreed to in writing, Licensor provides the Work (and each
+ Contributor provides its Contributions) on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
+ implied, including, without limitation, any warranties or conditions
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
+ PARTICULAR PURPOSE. You are solely responsible for determining the
+ appropriateness of using or redistributing the Work and assume any
+ risks associated with Your exercise of permissions under this License.
+
+ 8. Limitation of Liability. In no event and under no legal theory,
+ whether in tort (including negligence), contract, or otherwise,
+ unless required by applicable law (such as deliberate and grossly
+ negligent acts) or agreed to in writing, shall any Contributor be
+ liable to You for damages, including any direct, indirect, special,
+ incidental, or consequential damages of any character arising as a
+ result of this License or out of the use or inability to use the
+ Work (including but not limited to damages for loss of goodwill,
+ work stoppage, computer failure or malfunction, or any and all
+ other commercial damages or losses), even if such Contributor
+ has been advised of the possibility of such damages.
+
+ 9. Accepting Warranty or Additional Liability. While redistributing
+ the Work or Derivative Works thereof, You may choose to offer,
+ and charge a fee for, acceptance of support, warranty, indemnity,
+ or other liability obligations and/or rights consistent with this
+ License. However, in accepting such obligations, You may act only
+ on Your own behalf and on Your sole responsibility, not on behalf
+ of any other Contributor, and only if You agree to indemnify,
+ defend, and hold each Contributor harmless for any liability
+ incurred by, or claims asserted against, such Contributor by reason
+ of your accepting any such warranty or additional liability.
+
+ END OF TERMS AND CONDITIONS
+
+ APPENDIX: How to apply the Apache License to your work.
+
+ To apply the Apache License to your work, attach the following
+ boilerplate notice, with the fields enclosed by brackets "[]"
+ replaced with your own identifying information. (Don't include
+ the brackets!) The text should be enclosed in the appropriate
+ comment syntax for the file format. We also recommend that a
+ file or class name and description of purpose be included on the
+ same "printed page" as the copyright notice for easier
+ identification within third-party archives.
+
+ Copyright [yyyy] [name of copyright owner]
+
+ Licensed under the Apache License, Version 2.0 (the "License");
+ you may not use this file except in compliance with the License.
+ You may obtain a copy of the License at
+
+ http://www.apache.org/licenses/LICENSE-2.0
+
+ Unless required by applicable law or agreed to in writing, software
+ distributed under the License is distributed on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ See the License for the specific language governing permissions and
+ limitations under the License.
diff --git a/baselines/fedper/README.md b/baselines/fedper/README.md
new file mode 100644
index 000000000000..157bc22d2da5
--- /dev/null
+++ b/baselines/fedper/README.md
@@ -0,0 +1,152 @@
+---
+title: Federated Learning with Personalization Layers
+url: https://arxiv.org/abs/1912.00818
+labels: [system heterogeneity, image classification, personalization, horizontal data partition]
+dataset: [CIFAR-10, FLICKR-AES]
+---
+
+# Federated Learning with Personalization Layers
+
+> Note: If you use this baseline in your work, please remember to cite the original authors of the paper as well as the Flower paper.
+
+**Paper:** [arxiv.org/abs/1912.00818](https://arxiv.org/abs/1912.00818)
+
+**Authors:** Manoj Ghuhan Arivazhagan, Vinay Aggarwal, Aaditya Kumar Singh, and Sunav Choudhary
+
+**Abstract:** The emerging paradigm of federated learning strives to enable collaborative training of machine learning models on the network edge without centrally aggregating raw data and hence, improving data privacy. This sharply deviates from traditional machine learning and necessitates design of algorithms robust to various sources of heterogeneity. Specifically, statistical heterogeneity of data across user devices can severely degrade performance of standard federated averaging for traditional machine learning applications like personalization with deep learning. This paper proposes `FedPer`, a base + personalization layer approach for federated training of deep feed forward neural networks, which can combat the ill-effects of statistical heterogeneity. We demonstrate effectiveness of `FedPer` for non-identical data partitions of CIFAR datasets and on a personalized image aesthetics dataset from Flickr.
+
+## About this baseline
+
+**What’s implemented:** The code in this directory replicates the experiments in _Federated Learning with Personalization Layers_ (Arivazhagan et al., 2019) for CIFAR10 and FLICKR-AES datasets, which proposed the `FedPer` model. Specifically, it replicates the results found in figures 2, 4, 7, and 8 in their paper. __Note__ that there is typo in the caption of Figure 4 in the article, it should be CIFAR10 and __not__ CIFAR100.
+
+**Datasets:** CIFAR10 from PyTorch's Torchvision and FLICKR-AES. FLICKR-AES was proposed as dataset in _Personalized Image Aesthetics_ (Ren et al., 2017) and can be downloaded using a link provided on thier [GitHub](https://github.com/alanspike/personalizedImageAesthetics). One must first download FLICKR-AES-001.zip (5.76GB), extract all inside and place in baseline/FedPer/datasets. To this location, also download the other 2 related files: (1) FLICKR-AES_image_labeled_by_each_worker.csv, and (2) FLICKR-AES_image_score.txt. Images are also scaled to 224x224 for both datasets. This is not explicitly stated in the paper but seems to be boosting performance. Also, for FLICKR dataset, it is stated in the paper that they use data from clients with more than 60 and less than 290 rated images. This amounts to circa 60 clients and we randomly select 30 out of these (as in paper). Therefore, the results might differ somewhat but only slighly. Since the pre-processing steps in the paper are somewhat obscure, the metric values in the plots below may differ slightly, but not the overall results and findings.
+
+```bash
+# These steps are not needed if you are only interested in CIFAR-10
+
+# Create the `datasets` directory if it doesn't exist already
+mkdir datasets
+
+# move/copy the downloaded FLICKR-AES-001.zip file to `datasets/`
+
+# unzip dataset to a directory named `flickr`
+cd datasets
+unzip FLICKR-AES-001.zip -d flickr
+
+# then move the .csv files inside flickr
+mv FLICKR-AES_image_labeled_by_each_worker.csv flickr
+mv FLICKR-AES_image_score.txt flickr
+```
+
+**Hardware Setup:** Experiments have been carried out on GPU. 2 different computers managed to run experiments:
+
+- GeForce RTX 3080 16GB
+- GeForce RTX 4090 24GB
+
+It's worth mentioning that GPU memory for each client is ~7.5GB. When training on powerful GPUs, one can reduce ratio of GPU needed for each client in the configuration setting to e.g. `num_gpus` to 0.33.
+
+> NOTE: One experiment carried out using 1 GPU (RTX 4090) takes somehwere between 1-3h depending on dataset and model. Running ResNet34 compared to MobileNet-v1 takes approximately 10-15% longer.
+
+**Contributors:** [William Lindskog](https://github.com/WilliamLindskog)
+
+
+## Experimental Setup
+
+**Task:** Image Classification
+
+**Model:** This directory implements 2 models:
+
+- ResNet34 which can be imported directly (after having installed the packages) from PyTorch, using `from torchvision.models import resnet34
+- MobileNet-v1
+
+Please see how models are implemented using a so called model_manager and model_split class since FedPer uses head and base layers in a neural network. These classes are defined in the models.py file and thereafter called when building new models in the directory /implemented_models. Please, extend and add new models as you wish.
+
+**Dataset:** CIFAR10, FLICKR-AES. CIFAR10 will be partitioned based on number of classes for data that each client shall recieve e.g. 4 allocated classes could be [1, 3, 5, 9]. FLICKR-AES is an unbalanced dataset, so there we only apply random sampling.
+
+**Training Hyperparameters:** The hyperparameters can be found in conf/base.yaml file which is the configuration file for the main script.
+
+| Description | Default Value |
+| ----------- | ----- |
+| num_clients | 10 |
+| clients per round | 10 |
+| number of rounds | 50 |
+| client resources | {'num_cpus': 4, 'num_gpus': 1 }|
+| learning_rate | 0.01 |
+| batch_size | 128 |
+| optimizer | SGD |
+| algorithm | fedavg|
+
+**Stateful Clients:**
+In this Baseline (FedPer), we must store the state of the local client head while aggregation of body parameters happen at the server. Flower is currently making this possible but for the time being, we reside to storing client _head_ state in a folder called client_states. We store the values after each fit and evaluate function carried out on each client, and call for the state before executing these funcitons. Moreover, the state of a unique client is accessed using the client ID.
+
+> NOTE: This is a work-around so that the local head parameters are not reset before each fit and evaluate. Nevertheless, it can come to change with future releases.
+
+
+## Environment Setup
+
+To construct the Python environment follow these steps:
+
+```bash
+# Set Python 3.10
+pyenv local 3.10.6
+# Tell poetry to use python 3.10
+poetry env use 3.10.6
+
+# Install the base Poetry environment
+poetry install
+
+# Activate the environment
+poetry shell
+```
+
+## Running the Experiments
+```bash
+python -m fedper.main # this will run using the default settings in the `conf/base.yaml`
+
+# When running models for flickr dataset, it is important to keep batch size at 4 or lower since some clients (for reproducing experiment) will have very few examples of one class
+```
+
+While the config files contain a large number of settings, the ones below are the main ones you'd likely want to modify to .
+```bash
+algorithm: fedavg, fedper # these are currently supported
+server_device: 'cuda:0', 'cpu'
+dataset.name: 'cifar10', 'flickr'
+num_classes: 10, 5 # respectively
+dataset.num_classes: 4, 8, 10 # for non-iid split assigning n num_classes to each client (these numbers for CIFAR10 experiments)
+model_name: mobile, resnet
+```
+
+To run multiple runs, one can also reside to `HYDRA`'s multirun option.
+```bash
+# for CIFAR10
+python -m fedper.main --multirun --config_name cifar10 dataset.num_classes=4,8,10 model_name=resnet,mobile algorithm=fedper,fedavg model.num_head_layers=2,3
+
+# to repeat each run 5 times, one can also add
+python -m fedper.main --multirun --config_name cifar10 dataset.num_classes=4,8,10 model_name=resnet,mobile algorithm=fedper,fedavg model.num_head_layers=2,3 '+repeat_num=range(5)'
+```
+
+
+## Expected Results
+
+To reproduce figures make `fedper/run_figures.sh` executable and run it. By default all experiments will be run:
+
+```bash
+# Make fedper/run_figures.sh executable
+chmod u+x fedper/run_figures.sh
+# Run the script
+bash fedper/run_figures.sh
+```
+
+Having run the `run_figures.sh`, the expected results should look something like this:
+
+**MobileNet-v1 and ResNet-34 on CIFAR10**
+
+
+
+**MobileNet-v1 and ResNet-34 on CIFAR10 using varying size of head**
+
+
+
+**MobileNet-v1 and ResNet-34 on FLICKR-AES**
+
+
\ No newline at end of file
diff --git a/baselines/fedper/_static/mobile_plot_figure_2.png b/baselines/fedper/_static/mobile_plot_figure_2.png
new file mode 100644
index 000000000000..b485b850fb39
Binary files /dev/null and b/baselines/fedper/_static/mobile_plot_figure_2.png differ
diff --git a/baselines/fedper/_static/mobile_plot_figure_flickr.png b/baselines/fedper/_static/mobile_plot_figure_flickr.png
new file mode 100644
index 000000000000..76e99927df36
Binary files /dev/null and b/baselines/fedper/_static/mobile_plot_figure_flickr.png differ
diff --git a/baselines/fedper/_static/mobile_plot_figure_num_head.png b/baselines/fedper/_static/mobile_plot_figure_num_head.png
new file mode 100644
index 000000000000..9dcb9f0a3f33
Binary files /dev/null and b/baselines/fedper/_static/mobile_plot_figure_num_head.png differ
diff --git a/baselines/fedper/_static/resnet_plot_figure_2.png b/baselines/fedper/_static/resnet_plot_figure_2.png
new file mode 100644
index 000000000000..14e3a7145a23
Binary files /dev/null and b/baselines/fedper/_static/resnet_plot_figure_2.png differ
diff --git a/baselines/fedper/_static/resnet_plot_figure_flickr.png b/baselines/fedper/_static/resnet_plot_figure_flickr.png
new file mode 100644
index 000000000000..4e6ba71489b7
Binary files /dev/null and b/baselines/fedper/_static/resnet_plot_figure_flickr.png differ
diff --git a/baselines/fedper/_static/resnet_plot_figure_num_head.png b/baselines/fedper/_static/resnet_plot_figure_num_head.png
new file mode 100644
index 000000000000..03c6ac88b84a
Binary files /dev/null and b/baselines/fedper/_static/resnet_plot_figure_num_head.png differ
diff --git a/baselines/fedper/fedper/__init__.py b/baselines/fedper/fedper/__init__.py
new file mode 100644
index 000000000000..a5e567b59135
--- /dev/null
+++ b/baselines/fedper/fedper/__init__.py
@@ -0,0 +1 @@
+"""Template baseline package."""
diff --git a/baselines/fedper/fedper/client.py b/baselines/fedper/fedper/client.py
new file mode 100644
index 000000000000..83babbd9613f
--- /dev/null
+++ b/baselines/fedper/fedper/client.py
@@ -0,0 +1,353 @@
+"""Client implementation - can call FedPer and FedAvg clients."""
+import pickle
+from collections import OrderedDict, defaultdict
+from pathlib import Path
+from typing import Any, Callable, Dict, List, Tuple, Type, Union
+
+import numpy as np
+import torch
+from flwr.client import NumPyClient
+from flwr.common import NDArrays, Scalar
+from omegaconf import DictConfig
+from torch.utils.data import DataLoader, Subset, random_split
+from torchvision import transforms
+from torchvision.datasets import ImageFolder
+
+from fedper.constants import MEAN, STD
+from fedper.dataset_preparation import call_dataset
+from fedper.implemented_models.mobile_model import MobileNetModelManager
+from fedper.implemented_models.resnet_model import ResNetModelManager
+
+PROJECT_DIR = Path(__file__).parent.parent.absolute()
+
+
+class ClientDataloaders:
+ """Client dataloaders."""
+
+ def __init__(
+ self,
+ trainloader: DataLoader,
+ testloader: DataLoader,
+ ) -> None:
+ """Initialize the client dataloaders."""
+ self.trainloader = trainloader
+ self.testloader = testloader
+
+
+class ClientEssentials:
+ """Client essentials."""
+
+ def __init__(
+ self,
+ client_id: str,
+ client_state_save_path: str = "",
+ ) -> None:
+ """Set client state save path and client ID."""
+ self.client_id = int(client_id)
+ self.client_state_save_path = (
+ (client_state_save_path + f"/client_{self.client_id}")
+ if client_state_save_path != ""
+ else None
+ )
+
+
+class BaseClient(NumPyClient):
+ """Implementation of Federated Averaging (FedAvg) Client."""
+
+ def __init__(
+ self,
+ data_loaders: ClientDataloaders,
+ config: DictConfig,
+ client_essentials: ClientEssentials,
+ model_manager_class: Union[
+ Type[MobileNetModelManager], Type[ResNetModelManager]
+ ],
+ ):
+ """Initialize client attributes.
+
+ Args:
+ config: dictionary containing the client configurations.
+ client_id: id of the client.
+ model_manager_class: class to be used as the model manager.
+ """
+ super().__init__()
+
+ self.train_id = 1
+ self.test_id = 1
+ self.client_id = int(client_essentials.client_id)
+ self.client_state_save_path = client_essentials.client_state_save_path
+ self.hist: Dict[str, Dict[str, Any]] = defaultdict(dict)
+ self.num_epochs: int = config["num_epochs"]
+ self.model_manager = model_manager_class(
+ client_id=self.client_id,
+ config=config,
+ trainloader=data_loaders.trainloader,
+ testloader=data_loaders.testloader,
+ client_save_path=self.client_state_save_path,
+ learning_rate=config["learning_rate"],
+ )
+
+ def get_parameters(self, config: Dict[str, Scalar]) -> NDArrays:
+ """Return the current local model parameters."""
+ return self.model_manager.model.get_parameters()
+
+ def set_parameters(
+ self, parameters: List[np.ndarray], evaluate: bool = False
+ ) -> None:
+ """Set the local model parameters to the received parameters.
+
+ Args:
+ parameters: parameters to set the model to.
+ """
+ _ = evaluate
+ model_keys = [
+ k
+ for k in self.model_manager.model.state_dict().keys()
+ if k.startswith("_body") or k.startswith("_head")
+ ]
+ params_dict = zip(model_keys, parameters)
+
+ state_dict = OrderedDict({k: torch.tensor(v) for k, v in params_dict})
+
+ self.model_manager.model.set_parameters(state_dict)
+
+ def perform_train(
+ self,
+ ) -> Dict[str, Union[List[Dict[str, float]], int, float]]:
+ """Perform local training to the whole model.
+
+ Returns
+ -------
+ Dict with the train metrics.
+ """
+ epochs = self.num_epochs
+
+ self.model_manager.model.enable_body()
+ self.model_manager.model.enable_head()
+
+ return self.model_manager.train(
+ epochs=epochs,
+ )
+
+ def fit(
+ self, parameters: NDArrays, config: Dict[str, Scalar]
+ ) -> Tuple[NDArrays, int, Dict[str, Union[bool, bytes, float, int, str]]]:
+ """Train the provided parameters using the locally held dataset.
+
+ Args:
+ parameters: The current (global) model parameters.
+ config: configuration parameters for training sent by the server.
+
+ Returns
+ -------
+ Tuple containing the locally updated model parameters, \
+ the number of examples used for training and \
+ the training metrics.
+ """
+ self.set_parameters(parameters)
+
+ train_results = self.perform_train()
+
+ # Update train history
+ self.hist[str(self.train_id)] = {
+ **self.hist[str(self.train_id)],
+ "trn": train_results,
+ }
+ print("<------- TRAIN RESULTS -------> :", train_results)
+
+ self.train_id += 1
+
+ return self.get_parameters(config), self.model_manager.train_dataset_size(), {}
+
+ def evaluate(
+ self, parameters: NDArrays, config: Dict[str, Scalar]
+ ) -> Tuple[float, int, Dict[str, Union[bool, bytes, float, int, str]]]:
+ """Evaluate the provided global parameters using the locally held dataset.
+
+ Args:
+ parameters: The current (global) model parameters.
+ config: configuration parameters for training sent by the server.
+
+ Returns
+ -------
+ Tuple containing the test loss, \
+ the number of examples used for evaluation and \
+ the evaluation metrics.
+ """
+ self.set_parameters(parameters, evaluate=True)
+
+ # Test the model
+ tst_results = self.model_manager.test()
+ print("<------- TEST RESULTS -------> :", tst_results)
+
+ # Update test history
+ self.hist[str(self.test_id)] = {
+ **self.hist[str(self.test_id)],
+ "tst": tst_results,
+ }
+ self.test_id += 1
+
+ return (
+ tst_results.get("loss", 0.0),
+ self.model_manager.test_dataset_size(),
+ {k: v for k, v in tst_results.items() if not isinstance(v, (dict, list))},
+ )
+
+
+class FedPerClient(BaseClient):
+ """Implementation of Federated Personalization (FedPer) Client."""
+
+ def get_parameters(self, config: Dict[str, Scalar]) -> NDArrays:
+ """Return the current local body parameters."""
+ return [
+ val.cpu().numpy()
+ for _, val in self.model_manager.model.body.state_dict().items()
+ ]
+
+ def set_parameters(self, parameters: List[np.ndarray], evaluate=False) -> None:
+ """Set the local body parameters to the received parameters.
+
+ Args:
+ parameters: parameters to set the body to.
+ evaluate: whether the client is evaluating or not.
+ """
+ model_keys = [
+ k
+ for k in self.model_manager.model.state_dict().keys()
+ if k.startswith("_body")
+ ]
+
+ if not evaluate:
+ # Only update client's local head if it hasn't trained yet
+ print("Setting head parameters to global head parameters.")
+ model_keys.extend(
+ [
+ k
+ for k in self.model_manager.model.state_dict().keys()
+ if k.startswith("_head")
+ ]
+ )
+
+ params_dict = zip(model_keys, parameters)
+
+ state_dict = OrderedDict({k: torch.tensor(v) for k, v in params_dict})
+
+ self.model_manager.model.set_parameters(state_dict)
+
+
+def get_client_fn_simulation(
+ config: DictConfig,
+ client_state_save_path: str = "",
+) -> Callable[[str], Union[FedPerClient, BaseClient]]:
+ """Generate the client function that creates the Flower Clients.
+
+ Parameters
+ ----------
+ model : DictConfig
+ The model configuration.
+ cleint_state_save_path : str
+ The path to save the client state.
+
+ Returns
+ -------
+ Tuple[Callable[[str], FlowerClient], DataLoader]
+ A tuple containing the client function that creates Flower Clients and
+ the DataLoader that will be used for testing
+ """
+ assert config.model_name.lower() in [
+ "mobile",
+ "resnet",
+ ], f"Model {config.model.name} not implemented"
+
+ # load dataset and clients' data indices
+ if config.dataset.name.lower() == "cifar10":
+ try:
+ partition_path = (
+ PROJECT_DIR / "datasets" / config.dataset.name / "partition.pkl"
+ )
+ print(f"Loading partition from {partition_path}")
+ with open(partition_path, "rb") as pickle_file:
+ partition = pickle.load(pickle_file)
+ data_indices: Dict[int, Dict[str, List[int]]] = partition["data_indices"]
+ except FileNotFoundError as error:
+ print(f"Partition not found at {partition_path}")
+ raise error
+
+ # - you can define your own data transformation strategy here -
+ general_data_transform = transforms.Compose(
+ [
+ transforms.Resize((224, 224)),
+ transforms.RandomCrop(224, padding=4),
+ # transforms.RandomHorizontalFlip(),
+ # transforms.ToTensor(),
+ transforms.Normalize(
+ MEAN[config.dataset.name], STD[config.dataset.name]
+ ),
+ ]
+ )
+ # ------------------------------------------------------------
+
+ def client_fn(cid: str) -> BaseClient:
+ """Create a Flower client representing a single organization."""
+ cid_use = int(cid)
+ if config.dataset.name.lower() == "flickr":
+ transform = transforms.Compose(
+ [
+ transforms.Resize((224, 224)),
+ transforms.ToTensor(),
+ ]
+ )
+ data_path = (
+ PROJECT_DIR / "datasets" / config.dataset.name / "tmp" / f"client_{cid}"
+ )
+ dataset = ImageFolder(root=data_path, transform=transform)
+ trainset, testset = random_split(
+ dataset,
+ [int(len(dataset) * 0.8), len(dataset) - int(len(dataset) * 0.8)],
+ )
+ else:
+ dataset = call_dataset(
+ dataset_name=config.dataset.name,
+ root=PROJECT_DIR / "datasets" / config.dataset.name,
+ general_data_transform=general_data_transform,
+ )
+
+ trainset = Subset(dataset, indices=[])
+ testset = Subset(dataset, indices=[])
+ trainset.indices = data_indices[cid_use]["train"]
+ testset.indices = data_indices[cid_use]["test"]
+
+ # Create the train loader
+ trainloader = DataLoader(trainset, config.batch_size, shuffle=False)
+ # Create the test loader
+ testloader = DataLoader(testset, config.batch_size)
+
+ manager: Union[
+ Type[MobileNetModelManager], Type[ResNetModelManager]
+ ] = MobileNetModelManager
+ if config.model_name.lower() == "resnet":
+ manager = ResNetModelManager
+ elif config.model_name.lower() == "mobile":
+ manager = MobileNetModelManager
+ else:
+ raise NotImplementedError("Model not implemented, check name.")
+ client_data_loaders = ClientDataloaders(trainloader, testloader)
+ client_essentials = ClientEssentials(
+ client_id=cid,
+ client_state_save_path=client_state_save_path,
+ )
+ if client_state_save_path != "":
+ return FedPerClient(
+ data_loaders=client_data_loaders,
+ client_essentials=client_essentials,
+ config=config,
+ model_manager_class=manager,
+ )
+ return BaseClient(
+ data_loaders=client_data_loaders,
+ client_essentials=client_essentials,
+ config=config,
+ model_manager_class=manager,
+ )
+
+ return client_fn
diff --git a/baselines/fedper/fedper/conf/base.yaml b/baselines/fedper/fedper/conf/base.yaml
new file mode 100644
index 000000000000..b0b9778d4682
--- /dev/null
+++ b/baselines/fedper/fedper/conf/base.yaml
@@ -0,0 +1,44 @@
+---
+num_clients: 10 # total number of clients
+num_epochs: 4 # number of local epochs
+batch_size: 128
+num_rounds: 100
+clients_per_round: 10
+learning_rate: 0.01
+algorithm: fedper
+model_name: resnet
+
+client_resources:
+ num_cpus: 4
+ num_gpus: 1
+
+server_device: cuda:0
+
+dataset:
+ name : "cifar10"
+ split: sample
+ num_classes: 10
+ seed: 42
+ num_clients: ${num_clients}
+ fraction: 0.83
+
+model:
+ _target_: null
+ num_head_layers: 2
+ num_classes: 10
+
+fit_config:
+ drop_client: false
+ epochs : ${num_epochs}
+ batch_size: ${batch_size}
+
+strategy:
+ _target_: fedPer.server.DefaultStrategyPipeline
+ fraction_fit: 0.00001 # because we want the number of clients to sample on each roudn to be solely defined by min_fit_clients
+ min_fit_clients: ${clients_per_round}
+ fraction_evaluate: 0.0
+ min_evaluate_clients: ${clients_per_round}
+ min_available_clients: ${num_clients}
+ algorithm: ${algorithm}
+ evaluate_fn: None
+ on_evaluate_config_fn: None
\ No newline at end of file
diff --git a/baselines/fedper/fedper/conf/cifar10.yaml b/baselines/fedper/fedper/conf/cifar10.yaml
new file mode 100644
index 000000000000..66a06d481507
--- /dev/null
+++ b/baselines/fedper/fedper/conf/cifar10.yaml
@@ -0,0 +1,44 @@
+---
+num_clients: 10 # total number of clients
+num_epochs: 4 # number of local epochs
+batch_size: 128
+num_rounds: 50
+clients_per_round: 10
+learning_rate: 0.01
+algorithm: fedavg
+model_name: resnet
+
+client_resources:
+ num_cpus: 4
+ num_gpus: 1
+
+server_device: cuda:0
+
+dataset:
+ name : "cifar10"
+ split: sample
+ num_classes: 10
+ seed: 42
+ num_clients: ${num_clients}
+ fraction: 0.83
+
+model:
+ _target_: null
+ num_head_layers: 2
+ num_classes: 10
+
+fit_config:
+ drop_client: false
+ epochs : ${num_epochs}
+ batch_size: ${batch_size}
+
+strategy:
+ _target_: fedPer.server.DefaultStrategyPipeline
+ fraction_fit: 0.00001 # because we want the number of clients to sample on each roudn to be solely defined by min_fit_clients
+ min_fit_clients: ${clients_per_round}
+ fraction_evaluate: 0.0
+ min_evaluate_clients: ${clients_per_round}
+ min_available_clients: ${num_clients}
+ algorithm: ${algorithm}
+ evaluate_fn: None
+ on_evaluate_config_fn: None
\ No newline at end of file
diff --git a/baselines/fedper/fedper/conf/flickr.yaml b/baselines/fedper/fedper/conf/flickr.yaml
new file mode 100644
index 000000000000..341b1c0ac6c2
--- /dev/null
+++ b/baselines/fedper/fedper/conf/flickr.yaml
@@ -0,0 +1,44 @@
+---
+num_clients: 30 # total number of clients
+num_epochs: 4 # number of local epochs
+batch_size: 4
+num_rounds: 35
+clients_per_round: 30
+learning_rate: 0.01
+algorithm: fedper
+model_name: resnet
+
+client_resources:
+ num_cpus: 4
+ num_gpus: 1
+
+server_device: cuda:0
+
+dataset:
+ name : "flickr"
+ split: sample
+ num_classes: 5
+ seed: 42
+ num_clients: ${num_clients}
+ fraction: 0.80
+
+model:
+ _target_: null
+ num_head_layers: 2
+ num_classes: 5
+
+fit_config:
+ drop_client: false
+ epochs : ${num_epochs}
+ batch_size: ${batch_size}
+
+strategy:
+ _target_: fedPer.server.DefaultStrategyPipeline
+ fraction_fit: 0.00001 # because we want the number of clients to sample on each roudn to be solely defined by min_fit_clients
+ min_fit_clients: ${clients_per_round}
+ fraction_evaluate: 0.0
+ min_evaluate_clients: ${clients_per_round}
+ min_available_clients: ${num_clients}
+ algorithm: ${algorithm}
+ evaluate_fn: None
+ on_evaluate_config_fn: None
\ No newline at end of file
diff --git a/baselines/fedper/fedper/constants.py b/baselines/fedper/fedper/constants.py
new file mode 100644
index 000000000000..3eda77c5134e
--- /dev/null
+++ b/baselines/fedper/fedper/constants.py
@@ -0,0 +1,23 @@
+"""Constants used in machine learning pipeline."""
+from enum import Enum
+
+
+# FL Algorithms
+class Algorithms(Enum):
+ """Enum for FL algorithms."""
+
+ FEDAVG = "FedAvg"
+ FEDPER = "FedPer"
+
+
+# FL Default Train and Fine-Tuning Epochs
+DEFAULT_TRAIN_EP = 5
+DEFAULT_FT_EP = 5
+
+MEAN = {
+ "cifar10": [0.4915, 0.4823, 0.4468],
+}
+
+STD = {
+ "cifar10": [0.2470, 0.2435, 0.2616],
+}
diff --git a/baselines/fedper/fedper/dataset.py b/baselines/fedper/fedper/dataset.py
new file mode 100644
index 000000000000..81a95286b1b8
--- /dev/null
+++ b/baselines/fedper/fedper/dataset.py
@@ -0,0 +1,85 @@
+"""Handle basic dataset creation.
+
+In case of PyTorch it should return dataloaders for your dataset (for both the clients
+and the server). If you are using a custom dataset class, this module is the place to
+define it. If your dataset requires to be downloaded (and this is not done
+automatically -- e.g. as it is the case for many dataset in TorchVision) and
+partitioned, please include all those functions and logic in the
+`dataset_preparation.py` module. You can use all those functions from functions/methods
+defined here of course.
+"""
+import os
+import pickle
+import sys
+from pathlib import Path
+
+import numpy as np
+
+from fedper.dataset_preparation import (
+ call_dataset,
+ flickr_preprocess,
+ randomly_assign_classes,
+)
+
+# working dir is two up
+WORKING_DIR = Path(__file__).resolve().parent.parent
+FL_BENCH_ROOT = WORKING_DIR.parent
+
+sys.path.append(FL_BENCH_ROOT.as_posix())
+
+
+def dataset_main(config: dict) -> None:
+ """Prepare the dataset."""
+ dataset_name = config["name"].lower()
+ dataset_folder = Path(WORKING_DIR, "datasets")
+ dataset_root = Path(dataset_folder, dataset_name)
+
+ if not os.path.isdir(dataset_root):
+ os.makedirs(dataset_root)
+
+ if dataset_name == "cifar10":
+ dataset = call_dataset(dataset_name=dataset_name, root=dataset_root)
+
+ # randomly assign classes
+ assert config["num_classes"] > 0, "Number of classes must be positive"
+ config["num_classes"] = max(1, min(config["num_classes"], len(dataset.classes)))
+ # partition, stats = randomly_assign_classes(
+ partition = randomly_assign_classes(
+ dataset=dataset,
+ client_num=config["num_clients"],
+ class_num=config["num_classes"],
+ )
+
+ clients_4_train = list(range(config["num_clients"]))
+ clients_4_test = list(range(config["num_clients"]))
+
+ partition["separation"] = {
+ "train": clients_4_train,
+ "test": clients_4_test,
+ "total": config["num_clients"],
+ }
+ for client_id, idx in enumerate(partition["data_indices"]):
+ if config["split"] == "sample":
+ num_train_samples = int(len(idx) * config["fraction"])
+
+ np.random.shuffle(idx)
+ idx_train, idx_test = idx[:num_train_samples], idx[num_train_samples:]
+ partition["data_indices"][client_id] = {
+ "train": idx_train,
+ "test": idx_test,
+ }
+ else:
+ if client_id in clients_4_train:
+ partition["data_indices"][client_id] = {"train": idx, "test": []}
+ else:
+ partition["data_indices"][client_id] = {"train": [], "test": idx}
+ with open(dataset_root / "partition.pkl", "wb") as pickle_file:
+ pickle.dump(partition, pickle_file)
+
+ # with open(dataset_root / "all_stats.json", "w") as f:
+ # json.dump(stats, f)
+
+ elif dataset_name.lower() == "flickr":
+ flickr_preprocess(dataset_root, config)
+ else:
+ raise RuntimeError("Please implement the dataset preparation for your dataset.")
diff --git a/baselines/fedper/fedper/dataset_preparation.py b/baselines/fedper/fedper/dataset_preparation.py
new file mode 100644
index 000000000000..0b8b53782aac
--- /dev/null
+++ b/baselines/fedper/fedper/dataset_preparation.py
@@ -0,0 +1,209 @@
+"""Dataset preparation."""
+import os
+import random
+from collections import Counter
+from pathlib import Path
+from typing import Any, Dict, List, Union
+
+import numpy as np
+import pandas as pd
+import torch
+import torchvision
+from torch.utils.data import Dataset
+from torchvision import transforms
+
+
+class BaseDataset(Dataset):
+ """Base class for all datasets."""
+
+ def __init__(
+ self,
+ root: Path = Path("datasets/cifar10"),
+ general_data_transform: transforms.transforms.Compose = None,
+ ) -> None:
+ """Initialize the dataset."""
+ self.root = root
+ self.classes = None
+ self.data: torch.tensor = None
+ self.targets: torch.tensor = None
+ self.general_data_transform = general_data_transform
+
+ def __getitem__(self, index):
+ """Get the item at the given index."""
+ data, targets = self.data[index], self.targets[index]
+ if self.general_data_transform is not None:
+ data = self.general_data_transform(data)
+ return data, targets
+
+ def __len__(self):
+ """Return the length of the dataset."""
+ return len(self.targets)
+
+
+class CIFAR10(BaseDataset):
+ """CIFAR10 dataset."""
+
+ def __init__(
+ self,
+ root: Path = Path("datasets/cifar10"),
+ general_data_transform=None,
+ ):
+ super().__init__()
+ train_part = torchvision.datasets.CIFAR10(root, True, download=True)
+ test_part = torchvision.datasets.CIFAR10(root, False, download=True)
+ train_data = torch.tensor(train_part.data).permute([0, -1, 1, 2]).float()
+ test_data = torch.tensor(test_part.data).permute([0, -1, 1, 2]).float()
+ train_targets = torch.tensor(train_part.targets).long().squeeze()
+ test_targets = torch.tensor(test_part.targets).long().squeeze()
+ self.data = torch.cat([train_data, test_data])
+ self.targets = torch.cat([train_targets, test_targets])
+ self.classes = train_part.classes
+ self.general_data_transform = general_data_transform
+
+
+def flickr_preprocess(root, config):
+ """Preprocess the FLICKR dataset."""
+ print("Preprocessing FLICKR dataset...")
+ # create a tmp folder to store the preprocessed data
+ tmp_folder = Path(root, "tmp")
+ if not os.path.isdir(tmp_folder):
+ os.makedirs(tmp_folder)
+
+ # remove any folder or file in tmp folder, even if it is not empty
+ os.system(f"rm -rf {tmp_folder.as_posix()}/*")
+
+ # get number of clients
+ num_clients = config["num_clients"]
+ # get flickr image labels per clients
+ df_labelled_igms = pd.read_csv(
+ Path(root, "FLICKR-AES_image_labeled_by_each_worker.csv")
+ )
+ # take num_clients random workers from df
+ # #where workers have minimum 60 images and maximum 290
+ df_labelled_igms = df_labelled_igms.groupby("worker").filter(
+ lambda x: len(x) >= 60 and len(x) <= 290
+ )
+ # only take workers that have at least 1 image for each score (1-5)
+ df_labelled_igms = df_labelled_igms.groupby("worker").filter(
+ lambda x: len(x[" score"].unique()) == 5
+ )
+ df_labelled_igms = df_labelled_igms.groupby("worker").filter(
+ lambda x: x[" score"].value_counts().min() >= 4
+ )
+ # only take workers that have at least 4 images for each score (1-5)
+
+ # get num_clients random workers
+ clients = np.random.choice(
+ df_labelled_igms["worker"].unique(), num_clients, replace=False
+ )
+ for i, client in enumerate(clients):
+ print(f"Processing client {i}...")
+ df_client = df_labelled_igms[df_labelled_igms["worker"] == client]
+ client_path = Path(tmp_folder, f"client_{i}")
+ if not os.path.isdir(client_path):
+ os.makedirs(client_path)
+ # create score folder in client folder, scores go from 1-5
+ for score in range(1, 6):
+ score_path = Path(client_path, str(score))
+ if not os.path.isdir(score_path):
+ os.makedirs(score_path)
+ # copy images to score folder
+ for _, row in df_client.iterrows():
+ img_path = Path(root, "40K", row[" imagePair"])
+ score_path = Path(client_path, str(row[" score"]))
+ if os.path.isfile(img_path):
+ os.system(f"cp {img_path} {score_path}")
+
+
+def call_dataset(dataset_name, root, **kwargs):
+ """Call the dataset."""
+ if dataset_name == "cifar10":
+ return CIFAR10(root, **kwargs)
+ raise ValueError(f"Dataset {dataset_name} not supported.")
+
+
+def randomly_assign_classes(
+ dataset: Dataset, client_num: int, class_num: int
+) -> Dict[str, Union[Dict[Any, Any], List[Any]]]:
+ # ) -> Dict[str, Any]:
+ """Randomly assign number classes to clients."""
+ partition: Dict[str, Union[Dict, List]] = {"separation": {}, "data_indices": []}
+ data_indices: List[List[int]] = [[] for _ in range(client_num)]
+ targets_numpy = np.array(dataset.targets, dtype=np.int32)
+ label_list = list(range(len(dataset.classes)))
+
+ data_idx_for_each_label = [
+ np.where(targets_numpy == i)[0].tolist() for i in label_list
+ ]
+
+ assigned_labels = []
+ selected_times = [0 for _ in label_list]
+ for _ in range(client_num):
+ sampled_labels = random.sample(label_list, class_num)
+ assigned_labels.append(sampled_labels)
+ for j in sampled_labels:
+ selected_times[j] += 1
+
+ batch_sizes = _get_batch_sizes(
+ targets_numpy=targets_numpy,
+ label_list=label_list,
+ selected_times=selected_times,
+ )
+
+ data_indices = _get_data_indices(
+ batch_sizes=batch_sizes,
+ data_indices=data_indices,
+ data_idx_for_each_label=data_idx_for_each_label,
+ assigned_labels=assigned_labels,
+ client_num=client_num,
+ )
+
+ partition["data_indices"] = data_indices
+
+ return partition # , stats
+
+
+def _get_batch_sizes(
+ targets_numpy: np.ndarray,
+ label_list: List[int],
+ selected_times: List[int],
+) -> np.ndarray:
+ """Get batch sizes for each label."""
+ labels_count = Counter(targets_numpy)
+ batch_sizes = np.zeros_like(label_list)
+ for i in label_list:
+ print(f"label: {i}, count: {labels_count[i]}")
+ print(f"selected times: {selected_times[i]}")
+ batch_sizes[i] = int(labels_count[i] / selected_times[i])
+
+ return batch_sizes
+
+
+def _get_data_indices(
+ batch_sizes: np.ndarray,
+ data_indices: List[List[int]],
+ data_idx_for_each_label: List[List[int]],
+ assigned_labels: List[List[int]],
+ client_num: int,
+) -> List[List[int]]:
+ for i in range(client_num):
+ for cls in assigned_labels[i]:
+ if len(data_idx_for_each_label[cls]) < 2 * batch_sizes[cls]:
+ batch_size = len(data_idx_for_each_label[cls])
+ else:
+ batch_size = batch_sizes[cls]
+ selected_idx = random.sample(data_idx_for_each_label[cls], batch_size)
+ data_indices_use: np.ndarray = np.concatenate(
+ [data_indices[i], selected_idx], axis=0
+ ).astype(np.int64)
+ data_indices[i] = data_indices_use.tolist()
+ # data_indices[i]: np.ndarray = np.concatenate(
+ # [data_indices[i], selected_idx], axis=0
+ # ).astype(np.int64)
+ data_idx_for_each_label[cls] = list(
+ set(data_idx_for_each_label[cls]) - set(selected_idx)
+ )
+
+ data_indices[i] = data_indices[i]
+
+ return data_indices
diff --git a/baselines/fedper/fedper/implemented_models/mobile_model.py b/baselines/fedper/fedper/implemented_models/mobile_model.py
new file mode 100644
index 000000000000..57d3210c9511
--- /dev/null
+++ b/baselines/fedper/fedper/implemented_models/mobile_model.py
@@ -0,0 +1,258 @@
+"""MobileNet-v1 model, model manager and model split."""
+from typing import Dict, List, Optional, Tuple, Union
+
+import torch
+import torch.nn as nn
+from omegaconf import DictConfig
+from torch.utils.data import DataLoader
+
+from fedper.models import ModelManager, ModelSplit
+
+# Set model architecture
+ARCHITECTURE = {
+ "layer_1": {"conv_dw": [32, 64, 1]},
+ "layer_2": {"conv_dw": [64, 128, 2]},
+ "layer_3": {"conv_dw": [128, 128, 1]},
+ "layer_4": {"conv_dw": [128, 256, 2]},
+ "layer_5": {"conv_dw": [256, 256, 1]},
+ "layer_6": {"conv_dw": [256, 512, 2]},
+ "layer_7": {"conv_dw": [512, 512, 1]},
+ "layer_8": {"conv_dw": [512, 512, 1]},
+ "layer_9": {"conv_dw": [512, 512, 1]},
+ "layer_10": {"conv_dw": [512, 512, 1]},
+ "layer_11": {"conv_dw": [512, 512, 1]},
+ "layer_12": {"conv_dw": [512, 1024, 2]},
+ "layer_13": {"conv_dw": [1024, 1024, 1]},
+}
+
+
+class MobileNet(nn.Module):
+ """Model from MobileNet-v1 (https://github.com/wjc852456/pytorch-mobilenet-v1)."""
+
+ def __init__(
+ self,
+ num_head_layers: int = 1,
+ num_classes: int = 10,
+ ) -> None:
+ super(MobileNet, self).__init__()
+
+ self.architecture = ARCHITECTURE
+
+ def conv_bn(inp, oup, stride):
+ return nn.Sequential(
+ nn.Conv2d(inp, oup, 3, stride, 1, bias=False),
+ nn.BatchNorm2d(oup),
+ nn.ReLU(inplace=True),
+ )
+
+ def conv_dw(inp, oup, stride):
+ return nn.Sequential(
+ nn.Conv2d(inp, inp, 3, stride, 1, groups=inp, bias=False),
+ nn.BatchNorm2d(inp),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(inp, oup, 1, 1, 0, bias=False),
+ nn.BatchNorm2d(oup),
+ nn.ReLU(inplace=True),
+ )
+
+ self.body = nn.Sequential()
+ self.body.add_module("initial_batch_norm", conv_bn(3, 32, 2))
+ for i in range(1, 13):
+ for _, value in self.architecture[f"layer_{i}"].items():
+ self.body.add_module(f"conv_dw_{i}", conv_dw(*value))
+
+ self.body.add_module("avg_pool", nn.AvgPool2d([7]))
+ self.body.add_module("fc", nn.Linear(1024, num_classes))
+
+ if num_head_layers == 1:
+ self.head = nn.Sequential(
+ nn.AvgPool2d([7]), nn.Flatten(), nn.Linear(1024, num_classes)
+ )
+ self.body.avg_pool = nn.Identity()
+ self.body.fc = nn.Identity()
+ elif num_head_layers == 2:
+ self.head = nn.Sequential(
+ conv_dw(1024, 1024, 1),
+ nn.AvgPool2d([7]),
+ nn.Flatten(),
+ nn.Linear(1024, num_classes),
+ )
+ self.body.conv_dw_13 = nn.Identity()
+ self.body.avg_pool = nn.Identity()
+ self.body.fc = nn.Identity()
+ elif num_head_layers == 3:
+ self.head = nn.Sequential(
+ conv_dw(512, 1024, 2),
+ conv_dw(1024, 1024, 1),
+ nn.AvgPool2d([7]),
+ nn.Flatten(),
+ nn.Linear(1024, num_classes),
+ )
+ self.body.conv_dw_12 = nn.Identity()
+ self.body.conv_dw_13 = nn.Identity()
+ self.body.avg_pool = nn.Identity()
+ self.body.fc = nn.Identity()
+ elif num_head_layers == 4:
+ self.head = nn.Sequential(
+ conv_dw(512, 512, 1),
+ conv_dw(512, 1024, 2),
+ conv_dw(1024, 1024, 1),
+ nn.AvgPool2d([7]),
+ nn.Flatten(),
+ nn.Linear(1024, num_classes),
+ )
+ self.body.conv_dw_11 = nn.Identity()
+ self.body.conv_dw_12 = nn.Identity()
+ self.body.conv_dw_13 = nn.Identity()
+ self.body.avg_pool = nn.Identity()
+ self.body.fc = nn.Identity()
+ else:
+ raise NotImplementedError("Number of head layers not implemented.")
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ """Forward pass of the model."""
+ x = self.body(x)
+ return self.head(x)
+
+
+class MobileNetModelSplit(ModelSplit):
+ """Split MobileNet model into body and head."""
+
+ def _get_model_parts(self, model: MobileNet) -> Tuple[nn.Module, nn.Module]:
+ return model.body, model.head
+
+
+class MobileNetModelManager(ModelManager):
+ """Manager for models with Body/Head split."""
+
+ def __init__(
+ self,
+ client_id: int,
+ config: DictConfig,
+ trainloader: DataLoader,
+ testloader: DataLoader,
+ client_save_path: Optional[str] = "",
+ learning_rate: float = 0.01,
+ ):
+ """Initialize the attributes of the model manager.
+
+ Args:
+ client_id: The id of the client.
+ config: Dict containing the configurations to be used by the manager.
+ """
+ super().__init__(
+ model_split_class=MobileNetModelSplit,
+ client_id=client_id,
+ config=config,
+ )
+ self.trainloader, self.testloader = trainloader, testloader
+ self.device = self.config["server_device"]
+ self.client_save_path = client_save_path if client_save_path != "" else None
+ self.learning_rate = learning_rate
+
+ def _create_model(self) -> nn.Module:
+ """Return MobileNet-v1 model to be splitted into head and body."""
+ try:
+ return MobileNet(
+ num_head_layers=self.config["model"]["num_head_layers"],
+ num_classes=self.config["model"]["num_classes"],
+ ).to(self.device)
+ except AttributeError:
+ self.device = self.config["server_device"]
+ return MobileNet(
+ num_head_layers=self.config["model"]["num_head_layers"],
+ num_classes=self.config["model"]["num_classes"],
+ ).to(self.device)
+
+ def train(
+ self,
+ epochs: int = 1,
+ ) -> Dict[str, Union[List[Dict[str, float]], int, float]]:
+ """Train the model maintained in self.model.
+
+ Method adapted from simple MobileNet-v1 (PyTorch) \
+ https://github.com/wjc852456/pytorch-mobilenet-v1.
+
+ Args:
+ epochs: number of training epochs.
+
+ Returns
+ -------
+ Dict containing the train metrics.
+ """
+ # Load client state (head) if client_save_path is not None and it is not empty
+ if self.client_save_path is not None:
+ try:
+ self.model.head.load_state_dict(torch.load(self.client_save_path))
+ except FileNotFoundError:
+ print("No client state found, training from scratch.")
+ pass
+
+ criterion = torch.nn.CrossEntropyLoss()
+ optimizer = torch.optim.SGD(
+ self.model.parameters(), lr=self.learning_rate, momentum=0.9
+ )
+ correct, total = 0, 0
+ loss: torch.Tensor = 0.0
+ # self.model.train()
+ for _ in range(epochs):
+ for images, labels in self.trainloader:
+ optimizer.zero_grad()
+ outputs = self.model(images.to(self.device))
+ labels = labels.to(self.device)
+ loss = criterion(outputs, labels)
+ loss.backward()
+ optimizer.step()
+ total += labels.size(0)
+ correct += (torch.max(outputs.data, 1)[1] == labels).sum().item()
+
+ # Save client state (head)
+ if self.client_save_path is not None:
+ torch.save(self.model.head.state_dict(), self.client_save_path)
+
+ return {"loss": loss.item(), "accuracy": correct / total}
+
+ def test(
+ self,
+ ) -> Dict[str, float]:
+ """Test the model maintained in self.model.
+
+ Returns
+ -------
+ Dict containing the test metrics.
+ """
+ # Load client state (head)
+ if self.client_save_path is not None:
+ self.model.head.load_state_dict(torch.load(self.client_save_path))
+
+ criterion = torch.nn.CrossEntropyLoss()
+ correct, total, loss = 0, 0, 0.0
+ # self.model.eval()
+ with torch.no_grad():
+ for images, labels in self.testloader:
+ outputs = self.model(images.to(self.device))
+ labels = labels.to(self.device)
+ loss += criterion(outputs, labels).item()
+ total += labels.size(0)
+ correct += (torch.max(outputs.data, 1)[1] == labels).sum().item()
+ print("Test Accuracy: {:.4f}".format(correct / total))
+
+ if self.client_save_path is not None:
+ torch.save(self.model.head.state_dict(), self.client_save_path)
+
+ return {
+ "loss": loss / len(self.testloader.dataset),
+ "accuracy": correct / total,
+ }
+
+ def train_dataset_size(self) -> int:
+ """Return train data set size."""
+ return len(self.trainloader)
+
+ def test_dataset_size(self) -> int:
+ """Return test data set size."""
+ return len(self.testloader)
+
+ def total_dataset_size(self) -> int:
+ """Return total data set size."""
+ return len(self.trainloader) + len(self.testloader)
diff --git a/baselines/fedper/fedper/implemented_models/resnet_model.py b/baselines/fedper/fedper/implemented_models/resnet_model.py
new file mode 100644
index 000000000000..0d9837b118a3
--- /dev/null
+++ b/baselines/fedper/fedper/implemented_models/resnet_model.py
@@ -0,0 +1,272 @@
+"""ResNet model, model manager and split."""
+from typing import Dict, List, Optional, Tuple, Union
+
+import torch
+import torch.nn as nn
+from omegaconf import DictConfig
+from torch.utils.data import DataLoader
+from torchvision.models.resnet import resnet34
+
+from fedper.models import ModelManager, ModelSplit
+
+
+def conv3x3(
+ in_planes: int, out_planes: int, stride: int = 1, groups: int = 1, dilation: int = 1
+) -> nn.Conv2d:
+ """3x3 convolution with padding."""
+ return nn.Conv2d(
+ in_planes,
+ out_planes,
+ kernel_size=3,
+ stride=stride,
+ padding=dilation,
+ groups=groups,
+ bias=False,
+ dilation=dilation,
+ )
+
+
+def conv1x1(in_planes: int, out_planes: int, stride: int = 1) -> nn.Conv2d:
+ """1x1 convolution."""
+ return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)
+
+
+class BasicBlock(nn.Module):
+ """Basic block for ResNet."""
+
+ expansion: int = 1
+
+ def __init__(
+ self,
+ inplanes: int,
+ planes: int,
+ stride: int = 1,
+ downsample: Optional[nn.Module] = None,
+ ) -> None:
+ super().__init__()
+ norm_layer = nn.BatchNorm2d
+ # Both self.conv1 and self.downsample layers downsample input when stride != 1
+ self.conv1 = conv3x3(inplanes, planes, stride)
+ self.bn1 = norm_layer(planes)
+ self.relu = nn.ReLU(inplace=True)
+ self.conv2 = conv3x3(planes, planes)
+ self.bn2 = norm_layer(planes)
+ self.downsample = downsample
+ self.stride = stride
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ """Forward inputs through the block."""
+ identity = x
+
+ out = self.conv1(x)
+ out = self.bn1(out)
+ out = self.relu(out)
+
+ out = self.conv2(out)
+ out = self.bn2(out)
+
+ if self.downsample is not None:
+ identity = self.downsample(x)
+
+ out += identity
+ out = self.relu(out)
+
+ return out
+
+
+class ResNet(nn.Module):
+ """ResNet model."""
+
+ def __init__(
+ self,
+ num_head_layers: int = 1,
+ num_classes: int = 10,
+ ) -> None:
+ super(ResNet, self).__init__()
+ assert (
+ num_head_layers > 0 and num_head_layers <= 17
+ ), "num_head_layers must be greater than 0 and less than 16"
+
+ self.num_head_layers = num_head_layers
+ self.body = resnet34()
+
+ # if only one head layer
+ if self.num_head_layers == 1:
+ self.head = self.body.fc
+ self.body.fc = nn.Identity()
+ elif self.num_head_layers == 2:
+ self.head = nn.Sequential(
+ BasicBlock(512, 512),
+ nn.AdaptiveAvgPool2d((1, 1)),
+ nn.Flatten(),
+ nn.Linear(512, num_classes),
+ )
+ # remove head layers from body
+ self.body = nn.Sequential(*list(self.body.children())[:-2])
+ body_layer4 = list(self.body.children())[-1]
+ self.body = nn.Sequential(*list(self.body.children())[:-1])
+ self.body.layer4 = nn.Sequential(*list(body_layer4.children())[:-1])
+ elif self.num_head_layers == 3:
+ self.head = nn.Sequential(
+ BasicBlock(512, 512),
+ BasicBlock(512, 512),
+ nn.AdaptiveAvgPool2d((1, 1)),
+ nn.Flatten(),
+ nn.Linear(512, num_classes),
+ )
+ # remove head layers from body
+ self.body = nn.Sequential(*list(self.body.children())[:-2])
+ body_layer4 = list(self.body.children())[-1]
+ self.body = nn.Sequential(*list(self.body.children())[:-1])
+ self.body.layer4 = nn.Sequential(*list(body_layer4.children())[:-2])
+ else:
+ raise NotImplementedError("Only 1 or 2 head layers supported")
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ """Forward inputs through the model."""
+ print("Forwarding through ResNet model")
+ x = self.body(x)
+ return self.head(x)
+
+
+class ResNetModelSplit(ModelSplit):
+ """Split ResNet model into body and head."""
+
+ def _get_model_parts(self, model: ResNet) -> Tuple[nn.Module, nn.Module]:
+ return model.body, model.head
+
+
+class ResNetModelManager(ModelManager):
+ """Manager for models with Body/Head split."""
+
+ def __init__(
+ self,
+ client_save_path: Optional[str],
+ client_id: int,
+ config: DictConfig,
+ trainloader: DataLoader,
+ testloader: DataLoader,
+ learning_rate: float = 0.01,
+ ):
+ """Initialize the attributes of the model manager.
+
+ Args:
+ client_save_path: Path to save the client state.
+ client_id: The id of the client.
+ config: Dict containing the configurations to be used by the manager.
+ trainloader: DataLoader containing the train data.
+ testloader: DataLoader containing the test data.
+ learning_rate: Learning rate for the optimizer.
+ """
+ super().__init__(
+ model_split_class=ResNetModelSplit,
+ client_id=client_id,
+ config=config,
+ )
+ self.client_save_path = client_save_path
+ self.trainloader, self.testloader = trainloader, testloader
+ self.device = self.config["server_device"]
+ self.learning_rate = learning_rate
+
+ def _create_model(self) -> nn.Module:
+ """Return MobileNet-v1 model to be splitted into head and body."""
+ try:
+ return ResNet(
+ num_head_layers=self.config["model"]["num_head_layers"],
+ num_classes=self.config["model"]["num_classes"],
+ ).to(self.device)
+ except AttributeError:
+ self.device = self.config["server_device"]
+ return ResNet(
+ num_head_layers=self.config["model"]["num_head_layers"],
+ num_classes=self.config["model"]["num_classes"],
+ ).to(self.device)
+
+ def train(
+ self,
+ epochs: int = 1,
+ ) -> Dict[str, Union[List[Dict[str, float]], int, float]]:
+ """Train the model maintained in self.model.
+
+ Method adapted from simple MobileNet-v1 (PyTorch) \
+ https://github.com/wjc852456/pytorch-mobilenet-v1.
+
+ Args:
+ epochs: number of training epochs.
+
+ Returns
+ -------
+ Dict containing the train metrics.
+ """
+ # Load client state (head) if client_save_path is not None and it is not empty
+ if self.client_save_path is not None:
+ try:
+ self.model.head.load_state_dict(torch.load(self.client_save_path))
+ except FileNotFoundError:
+ print("No client state found, training from scratch.")
+ pass
+
+ criterion = torch.nn.CrossEntropyLoss()
+ optimizer = torch.optim.SGD(
+ self.model.parameters(), lr=self.learning_rate, momentum=0.9
+ )
+ correct, total = 0, 0
+ loss: torch.Tensor = 0.0
+ # self.model.train()
+ for _ in range(epochs):
+ for images, labels in self.trainloader:
+ optimizer.zero_grad()
+ outputs = self.model(images.to(self.device))
+ labels = labels.to(self.device)
+ loss = criterion(outputs, labels)
+ loss.backward()
+
+ optimizer.step()
+ total += labels.size(0)
+ correct += (torch.max(outputs.data, 1)[1] == labels).sum().item()
+
+ # Save client state (head)
+ if self.client_save_path is not None:
+ torch.save(self.model.head.state_dict(), self.client_save_path)
+
+ return {"loss": loss.item(), "accuracy": correct / total}
+
+ def test(
+ self,
+ ) -> Dict[str, float]:
+ """Test the model maintained in self.model."""
+ # Load client state (head)
+ if self.client_save_path is not None:
+ self.model.head.load_state_dict(torch.load(self.client_save_path))
+
+ criterion = torch.nn.CrossEntropyLoss()
+ correct, total, loss = 0, 0, 0.0
+ # self.model.eval()
+ with torch.no_grad():
+ for images, labels in self.testloader:
+ outputs = self.model(images.to(self.device))
+ labels = labels.to(self.device)
+ loss += criterion(outputs, labels).item()
+ total += labels.size(0)
+ correct += (torch.max(outputs.data, 1)[1] == labels).sum().item()
+ print("Test Accuracy: {:.4f}".format(correct / total))
+
+ if self.client_save_path is not None:
+ torch.save(self.model.head.state_dict(), self.client_save_path)
+
+ return {
+ "loss": loss / len(self.testloader.dataset),
+ "accuracy": correct / total,
+ }
+
+ def train_dataset_size(self) -> int:
+ """Return train data set size."""
+ return len(self.trainloader)
+
+ def test_dataset_size(self) -> int:
+ """Return test data set size."""
+ return len(self.testloader)
+
+ def total_dataset_size(self) -> int:
+ """Return total data set size."""
+ return len(self.trainloader) + len(self.testloader)
diff --git a/baselines/fedper/fedper/main.py b/baselines/fedper/fedper/main.py
new file mode 100644
index 000000000000..b421b2e0442c
--- /dev/null
+++ b/baselines/fedper/fedper/main.py
@@ -0,0 +1,126 @@
+"""Create and connect the building blocks for your experiments; start the simulation.
+
+It includes processioning the dataset, instantiate strategy, specify how the global
+model is going to be evaluated, etc. At the end, this script saves the results.
+"""
+
+from pathlib import Path
+
+import flwr as fl
+import hydra
+from hydra.core.hydra_config import HydraConfig
+from hydra.utils import instantiate
+from omegaconf import DictConfig, OmegaConf
+
+from fedper.dataset import dataset_main
+from fedper.utils import (
+ get_client_fn,
+ get_create_model_fn,
+ plot_metric_from_history,
+ save_results_as_pickle,
+ set_client_state_save_path,
+ set_model_class,
+ set_num_classes,
+ set_server_target,
+)
+
+
+@hydra.main(config_path="conf", config_name="base", version_base=None)
+def main(cfg: DictConfig) -> None:
+ """Run the baseline.
+
+ Parameters
+ ----------
+ cfg : DictConfig
+ An omegaconf object that stores the hydra config.
+ """
+ # 1. Print parsed config
+ # Set the model class, server target, and number of classes
+ cfg = set_model_class(cfg)
+ cfg = set_server_target(cfg)
+ cfg = set_num_classes(cfg)
+
+ print(OmegaConf.to_yaml(cfg))
+
+ # Create directory to store client states if it does not exist
+ # Client state has subdirectories with the name of current time
+ client_state_save_path = set_client_state_save_path()
+
+ # 2. Prepare your dataset
+ dataset_main(cfg.dataset)
+
+ # 3. Define your clients
+ # Get client function
+ client_fn = get_client_fn(
+ config=cfg,
+ client_state_save_path=client_state_save_path,
+ )
+
+ # get a function that will be used to construct the config that the client's
+ # fit() method will received
+ def get_on_fit_config():
+ def fit_config_fn(server_round: int):
+ # resolve and convert to python dict
+ fit_config = OmegaConf.to_container(cfg.fit_config, resolve=True)
+ _ = server_round
+ return fit_config
+
+ return fit_config_fn
+
+ # get a function that will be used to construct the model
+ create_model, split = get_create_model_fn(cfg)
+
+ # 4. Define your strategy
+ strategy = instantiate(
+ cfg.strategy,
+ create_model=create_model,
+ on_fit_config_fn=get_on_fit_config(),
+ model_split_class=split,
+ )
+
+ # 5. Start Simulation
+ history = fl.simulation.start_simulation(
+ client_fn=client_fn,
+ num_clients=cfg.num_clients,
+ config=fl.server.ServerConfig(num_rounds=cfg.num_rounds),
+ client_resources={
+ "num_cpus": cfg.client_resources.num_cpus,
+ "num_gpus": cfg.client_resources.num_gpus,
+ },
+ strategy=strategy,
+ )
+
+ # Experiment completed. Now we save the results and
+ # generate plots using the `history`
+ print("................")
+ print(history)
+
+ # 6. Save your results
+ save_path = Path(HydraConfig.get().runtime.output_dir)
+
+ # save results as a Python pickle using a file_path
+ # the directory created by Hydra for each run
+ save_results_as_pickle(
+ history,
+ file_path=save_path,
+ )
+ # plot results and include them in the readme
+ strategy_name = strategy.__class__.__name__
+ file_suffix: str = (
+ f"_{strategy_name}"
+ f"_C={cfg.num_clients}"
+ f"_B={cfg.batch_size}"
+ f"_E={cfg.num_epochs}"
+ f"_R={cfg.num_rounds}"
+ f"_lr={cfg.learning_rate}"
+ )
+
+ plot_metric_from_history(
+ history,
+ save_path,
+ (file_suffix),
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/baselines/fedper/fedper/models.py b/baselines/fedper/fedper/models.py
new file mode 100644
index 000000000000..2a2ebde158f8
--- /dev/null
+++ b/baselines/fedper/fedper/models.py
@@ -0,0 +1,189 @@
+"""Abstract class for splitting a model into body and head."""
+from abc import ABC, abstractmethod
+from collections import OrderedDict
+from typing import Any, Dict, List, Tuple, Type, Union
+
+import numpy as np
+from omegaconf import DictConfig
+from torch import Tensor
+from torch import nn as nn
+
+
+class ModelSplit(ABC, nn.Module):
+ """Abstract class for splitting a model into body and head."""
+
+ def __init__(
+ self,
+ model: nn.Module,
+ ):
+ """Initialize the attributes of the model split.
+
+ Args:
+ model: dict containing the vocab sizes of the input attributes.
+ """
+ super().__init__()
+
+ self._body, self._head = self._get_model_parts(model)
+
+ @abstractmethod
+ def _get_model_parts(self, model: nn.Module) -> Tuple[nn.Module, nn.Module]:
+ """Return the body and head of the model.
+
+ Args:
+ model: model to be split into head and body
+
+ Returns
+ -------
+ Tuple where the first element is the body of the model
+ and the second is the head.
+ """
+
+ @property
+ def body(self) -> nn.Module:
+ """Return model body."""
+ return self._body
+
+ @body.setter
+ def body(self, state_dict: "OrderedDict[str, Tensor]") -> None:
+ """Set model body.
+
+ Args:
+ state_dict: dictionary of the state to set the model body to.
+ """
+ self.body.load_state_dict(state_dict, strict=True)
+
+ @property
+ def head(self) -> nn.Module:
+ """Return model head."""
+ return self._head
+
+ @head.setter
+ def head(self, state_dict: "OrderedDict[str, Tensor]") -> None:
+ """Set model head.
+
+ Args:
+ state_dict: dictionary of the state to set the model head to.
+ """
+ self.head.load_state_dict(state_dict, strict=True)
+
+ def get_parameters(self) -> List[np.ndarray]:
+ """Get model parameters (without fixed head).
+
+ Returns
+ -------
+ Body and head parameters
+ """
+ return [
+ val.cpu().numpy()
+ for val in [
+ *self.body.state_dict().values(),
+ *self.head.state_dict().values(),
+ ]
+ ]
+
+ def set_parameters(self, state_dict: Dict[str, Tensor]) -> None:
+ """Set model parameters.
+
+ Args:
+ state_dict: dictionary of the state to set the model to.
+ """
+ ordered_state_dict = OrderedDict(self.state_dict().copy())
+ # Update with the values of the state_dict
+ ordered_state_dict.update(dict(state_dict.items()))
+ self.load_state_dict(ordered_state_dict, strict=False)
+
+ def enable_head(self) -> None:
+ """Enable gradient tracking for the head parameters."""
+ for param in self.head.parameters():
+ param.requires_grad = True
+
+ def enable_body(self) -> None:
+ """Enable gradient tracking for the body parameters."""
+ for param in self.body.parameters():
+ param.requires_grad = True
+
+ def disable_head(self) -> None:
+ """Disable gradient tracking for the head parameters."""
+ for param in self.head.parameters():
+ param.requires_grad = False
+
+ def disable_body(self) -> None:
+ """Disable gradient tracking for the body parameters."""
+ for param in self.body.parameters():
+ param.requires_grad = False
+
+ def forward(self, inputs: Any) -> Any:
+ """Forward inputs through the body and the head."""
+ x = self.body(inputs)
+ return self.head(x)
+
+
+class ModelManager(ABC):
+ """Manager for models with Body/Head split."""
+
+ def __init__(
+ self,
+ client_id: int,
+ config: DictConfig,
+ model_split_class: Type[Any], # ModelSplit
+ ):
+ """Initialize the attributes of the model manager.
+
+ Args:
+ client_id: The id of the client.
+ config: Dict containing the configurations to be used by the manager.
+ model_split_class: Class to be used to split the model into body and head\
+ (concrete implementation of ModelSplit).
+ """
+ super().__init__()
+
+ self.client_id = client_id
+ self.config = config
+ self._model = model_split_class(self._create_model())
+
+ @abstractmethod
+ def _create_model(self) -> nn.Module:
+ """Return model to be splitted into head and body."""
+
+ @abstractmethod
+ def train(
+ self,
+ epochs: int = 1,
+ ) -> Dict[str, Union[List[Dict[str, float]], int, float]]:
+ """Train the model maintained in self.model.
+
+ Args:
+ epochs: number of training epochs.
+
+ Returns
+ -------
+ Dict containing the train metrics.
+ """
+
+ @abstractmethod
+ def test(
+ self,
+ ) -> Dict[str, float]:
+ """Test the model maintained in self.model.
+
+ Returns
+ -------
+ Dict containing the test metrics.
+ """
+
+ @abstractmethod
+ def train_dataset_size(self) -> int:
+ """Return train data set size."""
+
+ @abstractmethod
+ def test_dataset_size(self) -> int:
+ """Return test data set size."""
+
+ @abstractmethod
+ def total_dataset_size(self) -> int:
+ """Return total data set size."""
+
+ @property
+ def model(self) -> nn.Module:
+ """Return model."""
+ return self._model
diff --git a/baselines/fedper/fedper/run_figures.sh b/baselines/fedper/fedper/run_figures.sh
new file mode 100755
index 000000000000..9f7382412465
--- /dev/null
+++ b/baselines/fedper/fedper/run_figures.sh
@@ -0,0 +1,36 @@
+#!/bin/bash
+
+# CIFAR10 Mobile and Resnet (non-iid n classes (FIGURE 2a&b))
+for model in mobile resnet
+do
+ for num_classes in 4 8 10
+ do
+ for algorithm in fedper fedavg
+ do
+ python -m fedper.main --config-path conf --config-name cifar10 dataset.num_classes=${num_classes} model_name=${model} algorithm=${algorithm}
+ done
+ done
+done
+
+
+# CIFAR10 Mobile (n head layers (FIGURE 4a))
+for num_head_layers in 2 3 4
+do
+ python -m fedper.main --config-path conf --config-name cifar10 dataset.num_classes=4 model.num_head_layers=${num_head_layers} num_rounds=25 model_name=mobile algorithm=fedper
+done
+python -m fedper.main --config-path conf --config-name cifar10 num_rounds=25 model_name=mobile dataset.num_classes=4
+
+# CIFAR10 Resnet (n head layers (FIGURE 4b))
+for num_head_layers in 1 2 3
+do
+ python -m fedper.main --config-path conf --config-name cifar10 dataset.num_classes=4 model.num_head_layers=${num_head_layers} num_rounds=25 model_name=resnet algorithm=fedper
+done
+python -m fedper.main --config-path conf --config-name cifar10 num_rounds=25 model_name=resnet dataset.num_classes=4
+
+# FLICKR
+for model in mobile resnet
+do
+ python -m fedper.main --config-path conf --config-name flickr model.num_head_layers=2 model_name=${model} algorithm=fedper num_rounds=35
+ python -m fedper.main --config-path conf --config-name flickr model_name=${model} algorithm=fedavg num_rounds=35
+done
+
diff --git a/baselines/fedper/fedper/server.py b/baselines/fedper/fedper/server.py
new file mode 100644
index 000000000000..93616f50f45a
--- /dev/null
+++ b/baselines/fedper/fedper/server.py
@@ -0,0 +1,24 @@
+"""Server strategies pipelines for FedPer."""
+from flwr.server.strategy.fedavg import FedAvg
+
+from fedper.strategy import (
+ AggregateBodyStrategy,
+ AggregateFullStrategy,
+ ServerInitializationStrategy,
+)
+
+
+class InitializationStrategyPipeline(ServerInitializationStrategy):
+ """Initialization strategy pipeline."""
+
+
+class AggregateBodyStrategyPipeline(
+ InitializationStrategyPipeline, AggregateBodyStrategy, FedAvg
+):
+ """Aggregate body strategy pipeline."""
+
+
+class DefaultStrategyPipeline(
+ InitializationStrategyPipeline, AggregateFullStrategy, FedAvg
+):
+ """Default strategy pipeline."""
diff --git a/baselines/fedper/fedper/strategy.py b/baselines/fedper/fedper/strategy.py
new file mode 100644
index 000000000000..5ae55086db2f
--- /dev/null
+++ b/baselines/fedper/fedper/strategy.py
@@ -0,0 +1,437 @@
+"""FL server strategies."""
+from collections import OrderedDict
+from pathlib import Path
+from typing import Any, Callable, Dict, List, Optional, Tuple, Type, Union
+
+import torch
+from flwr.common import (
+ EvaluateIns,
+ EvaluateRes,
+ FitIns,
+ FitRes,
+ NDArrays,
+ Parameters,
+ Scalar,
+ ndarrays_to_parameters,
+ parameters_to_ndarrays,
+)
+from flwr.server.client_manager import ClientManager
+from flwr.server.client_proxy import ClientProxy
+from flwr.server.strategy.fedavg import FedAvg
+from torch import nn as nn
+
+from fedper.constants import Algorithms
+from fedper.implemented_models.mobile_model import MobileNetModelSplit
+from fedper.implemented_models.resnet_model import ResNetModelSplit
+from fedper.models import ModelSplit
+
+
+class ServerInitializationStrategy(FedAvg):
+ """Server FL Parameter Initialization strategy implementation."""
+
+ def __init__(
+ self,
+ *args: Any,
+ model_split_class: Union[
+ Type[MobileNetModelSplit], Type[ModelSplit], Type[ResNetModelSplit]
+ ],
+ create_model: Callable[[], nn.Module],
+ initial_parameters: Optional[Parameters] = None,
+ on_fit_config_fn: Optional[Callable[[int], Dict[str, Any]]] = None,
+ evaluate_fn: Optional[
+ Callable[
+ [int, NDArrays, Dict[str, Scalar]],
+ Optional[Tuple[float, Dict[str, Scalar]]],
+ ]
+ ] = None,
+ min_available_clients: int = 1,
+ min_evaluate_clients: int = 1,
+ min_fit_clients: int = 1,
+ algorithm: str = Algorithms.FEDPER.value,
+ **kwargs: Any,
+ ) -> None:
+ super().__init__(*args, **kwargs)
+ _ = evaluate_fn
+ self.on_fit_config_fn = on_fit_config_fn
+ self.initial_parameters = initial_parameters
+ self.min_available_clients = min_available_clients
+ self.min_evaluate_clients = min_evaluate_clients
+ self.min_fit_clients = min_fit_clients
+ self.algorithm = algorithm
+ self.model = model_split_class(model=create_model())
+
+ def initialize_parameters(
+ self, client_manager: ClientManager
+ ) -> Optional[Parameters]:
+ """Initialize the (global) model parameters.
+
+ Args:
+ client_manager: ClientManager. The client manager which holds all currently
+ connected clients.
+
+ Returns
+ -------
+ If parameters are returned, then the server will treat these as the
+ initial global model parameters.
+ """
+ initial_parameters: Optional[Parameters] = self.initial_parameters
+ self.initial_parameters = None # Don't keep initial parameters in memory
+ if initial_parameters is None and self.model is not None:
+ if self.algorithm == Algorithms.FEDPER.value:
+ initial_parameters_use = [
+ val.cpu().numpy() for _, val in self.model.body.state_dict().items()
+ ]
+ else: # FedAvg
+ initial_parameters_use = [
+ val.cpu().numpy() for _, val in self.model.state_dict().items()
+ ]
+
+ if isinstance(initial_parameters_use, list):
+ initial_parameters = ndarrays_to_parameters(initial_parameters_use)
+ return initial_parameters
+
+
+class AggregateFullStrategy(ServerInitializationStrategy):
+ """Full model aggregation strategy implementation."""
+
+ def __init__(self, *args, save_path: Path = Path(""), **kwargs) -> None:
+ super().__init__(*args, **kwargs)
+ self.save_path = save_path if save_path != "" else None
+ if save_path is not None:
+ self.save_path = save_path / "models"
+ self.save_path.mkdir(parents=True, exist_ok=True)
+
+ def configure_evaluate(
+ self, server_round: int, parameters: Parameters, client_manager: ClientManager
+ ) -> List[Tuple[ClientProxy, EvaluateIns]]:
+ """Configure the next round of evaluation.
+
+ Args:
+ server_round: The current round of federated learning.
+ parameters: The current (global) model parameters.
+ client_manager: The client manager which holds all currently
+ connected clients.
+
+ Returns
+ -------
+ A list of tuples. Each tuple in the list identifies a `ClientProxy` and the
+ `EvaluateIns` for this particular `ClientProxy`. If a particular
+ `ClientProxy` is not included in this list, it means that this
+ `ClientProxy` will not participate in the next round of federated
+ evaluation.
+ """
+ # Same as superclass method but adds the head
+
+ # Parameters and config
+ config: Dict[Any, Any] = {}
+
+ weights = parameters_to_ndarrays(parameters)
+
+ parameters = ndarrays_to_parameters(weights)
+
+ evaluate_ins = EvaluateIns(parameters, config)
+
+ # Sample clients
+ if server_round >= 0:
+ # Sample clients
+ sample_size, min_num_clients = self.num_evaluation_clients(
+ client_manager.num_available()
+ )
+ clients = client_manager.sample(
+ num_clients=sample_size,
+ min_num_clients=min_num_clients,
+ )
+ else:
+ clients = list(client_manager.all().values())
+
+ # Return client/config pairs
+ return [(client, evaluate_ins) for client in clients]
+
+ def aggregate_fit(
+ self,
+ server_round: int,
+ results: List[Tuple[ClientProxy, FitRes]],
+ failures: List[Union[Tuple[ClientProxy, FitRes], BaseException]],
+ ) -> Tuple[Optional[Parameters], Dict[str, Scalar]]:
+ """Aggregate received local parameters, set global model parameters and save.
+
+ Args:
+ server_round: The current round of federated learning.
+ results: Successful updates from the previously selected and configured
+ clients. Each pair of `(ClientProxy, FitRes)` constitutes a
+ successful update from one of the previously selected clients. Not
+ that not all previously selected clients are necessarily included in
+ this list: a client might drop out and not submit a result. For each
+ client that did not submit an update, there should be an `Exception`
+ in `failures`.
+ failures: Exceptions that occurred while the server was waiting for client
+ updates.
+
+ Returns
+ -------
+ If parameters are returned, then the server will treat these as the
+ new global model parameters (i.e., it will replace the previous
+ parameters with the ones returned from this method). If `None` is
+ returned (e.g., because there were only failures and no viable
+ results) then the server will no update the previous model
+ parameters, the updates received in this round are discarded, and
+ the global model parameters remain the same.
+ """
+ agg_params, agg_metrics = super().aggregate_fit(
+ server_round=server_round, results=results, failures=failures
+ )
+ if agg_params is not None:
+ # Update Server Model
+ parameters = parameters_to_ndarrays(agg_params)
+ model_keys = [
+ k
+ for k in self.model.state_dict().keys()
+ if k.startswith("_body") or k.startswith("_head")
+ ]
+ params_dict = zip(model_keys, parameters)
+ state_dict = OrderedDict({k: torch.tensor(v) for k, v in params_dict})
+ self.model.set_parameters(state_dict)
+
+ if self.save_path is not None:
+ # Save Model
+ torch.save(self.model, self.save_path / f"model-ep_{server_round}.pt")
+
+ return agg_params, agg_metrics
+
+ def aggregate_evaluate(
+ self,
+ server_round: int,
+ results: List[Tuple[ClientProxy, EvaluateRes]],
+ failures: List[Union[Tuple[ClientProxy, EvaluateRes], BaseException]],
+ ) -> Tuple[Optional[float], Dict[str, Scalar]]:
+ """Aggregate the received local parameters and store the test aggregated.
+
+ Args:
+ server_round: The current round of federated learning.
+ results: Successful updates from the
+ previously selected and configured clients. Each pair of
+ `(ClientProxy, FitRes` constitutes a successful update from one of the
+ previously selected clients. Not that not all previously selected
+ clients are necessarily included in this list: a client might drop out
+ and not submit a result. For each client that did not submit an update,
+ there should be an `Exception` in `failures`.
+ failures: Exceptions that occurred while the server
+ was waiting for client updates.
+
+ Returns
+ -------
+ Optional `float` representing the aggregated evaluation result. Aggregation
+ typically uses some variant of a weighted average.
+ """
+ aggregated_loss, aggregated_metrics = super().aggregate_evaluate(
+ server_round=server_round, results=results, failures=failures
+ )
+ _ = aggregated_metrics # Avoid unused variable warning
+
+ # Weigh accuracy of each client by number of examples used
+ accuracies: List[float] = []
+ for _, res in results:
+ accuracy: float = float(res.metrics["accuracy"])
+ accuracies.append(accuracy)
+ print(f"Round {server_round} accuracies: {accuracies}")
+
+ # Aggregate and print custom metric
+ averaged_accuracy = sum(accuracies) / len(accuracies)
+ print(f"Round {server_round} accuracy averaged: {averaged_accuracy}")
+ return aggregated_loss, {"accuracy": averaged_accuracy}
+
+
+class AggregateBodyStrategy(ServerInitializationStrategy):
+ """Body Aggregation strategy implementation."""
+
+ def __init__(self, *args, save_path: Path = Path(""), **kwargs) -> None:
+ super().__init__(*args, **kwargs)
+ self.save_path = save_path if save_path != "" else None
+ if save_path is not None:
+ self.save_path = save_path / "models"
+ self.save_path.mkdir(parents=True, exist_ok=True)
+
+ def configure_fit(
+ self, server_round: int, parameters: Parameters, client_manager: ClientManager
+ ) -> List[Tuple[ClientProxy, FitIns]]:
+ """Configure the next round of training.
+
+ Args:
+ server_round: The current round of federated learning.
+ parameters: The current (global) model parameters.
+ client_manager: The client manager which holds all
+ currently connected clients.
+
+ Returns
+ -------
+ A list of tuples. Each tuple in the list identifies a `ClientProxy` and the
+ `FitIns` for this particular `ClientProxy`. If a particular `ClientProxy`
+ is not included in this list, it means that this `ClientProxy`
+ will not participate in the next round of federated learning.
+ """
+ # Same as superclass method but adds the head
+
+ config = {}
+ if self.on_fit_config_fn is not None:
+ # Custom fit config function provided
+ config = self.on_fit_config_fn(server_round)
+
+ weights = parameters_to_ndarrays(parameters)
+
+ # Add head parameters to received body parameters
+ weights.extend(
+ [val.cpu().numpy() for _, val in self.model.head.state_dict().items()]
+ )
+
+ parameters = ndarrays_to_parameters(weights)
+
+ fit_ins = FitIns(parameters, config)
+
+ # Sample clients
+ clients = client_manager.sample(
+ num_clients=self.min_available_clients, min_num_clients=self.min_fit_clients
+ )
+
+ # Return client/config pairs
+ return [(client, fit_ins) for client in clients]
+
+ def configure_evaluate(
+ self, server_round: int, parameters: Parameters, client_manager: ClientManager
+ ) -> List[Tuple[ClientProxy, EvaluateIns]]:
+ """Configure the next round of evaluation.
+
+ Args:
+ server_round: The current round of federated learning.
+ parameters: The current (global) model parameters.
+ client_manager: The client manager which holds all currently
+ connected clients.
+
+ Returns
+ -------
+ A list of tuples. Each tuple in the list identifies a `ClientProxy` and the
+ `EvaluateIns` for this particular `ClientProxy`. If a particular
+ `ClientProxy` is not included in this list, it means that this
+ `ClientProxy` will not participate in the next round of federated
+ evaluation.
+ """
+ # Same as superclass method but adds the head
+
+ # Parameters and config
+ config: Dict[Any, Any] = {}
+
+ weights = parameters_to_ndarrays(parameters)
+
+ # Add head parameters to received body parameters
+ weights.extend(
+ [val.cpu().numpy() for _, val in self.model.head.state_dict().items()]
+ )
+
+ parameters = ndarrays_to_parameters(weights)
+
+ evaluate_ins = EvaluateIns(parameters, config)
+
+ # Sample clients
+ if server_round >= 0:
+ # Sample clients
+ sample_size, min_num_clients = self.num_evaluation_clients(
+ client_manager.num_available()
+ )
+ clients = client_manager.sample(
+ num_clients=sample_size,
+ min_num_clients=min_num_clients,
+ )
+ else:
+ clients = list(client_manager.all().values())
+
+ # Return client/config pairs
+ return [(client, evaluate_ins) for client in clients]
+
+ def aggregate_fit(
+ self,
+ server_round: int,
+ results: List[Tuple[ClientProxy, FitRes]],
+ failures: List[Union[Tuple[ClientProxy, FitRes], BaseException]],
+ ) -> Tuple[Optional[Parameters], Dict[str, Union[bool, bytes, float, int, str]]]:
+ """Aggregate received local parameters, set global model parameters and save.
+
+ Args:
+ server_round: The current round of federated learning.
+ results: Successful updates from the previously selected and configured
+ clients. Each pair of `(ClientProxy, FitRes)` constitutes a
+ successful update from one of the previously selected clients. Not
+ that not all previously selected clients are necessarily included in
+ this list: a client might drop out and not submit a result. For each
+ client that did not submit an update, there should be an `Exception`
+ in `failures`.
+ failures: Exceptions that occurred while the server was waiting for client
+ updates.
+
+ Returns
+ -------
+ If parameters are returned, then the server will treat these as the
+ new global model parameters (i.e., it will replace the previous
+ parameters with the ones returned from this method). If `None` is
+ returned (e.g., because there were only failures and no viable
+ results) then the server will no update the previous model
+ parameters, the updates received in this round are discarded, and
+ the global model parameters remain the same.
+ """
+ agg_params, agg_metrics = super().aggregate_fit(
+ server_round=server_round, results=results, failures=failures
+ )
+ if agg_params is not None:
+ parameters = parameters_to_ndarrays(agg_params)
+ model_keys = [
+ k for k in self.model.state_dict().keys() if k.startswith("_body")
+ ]
+ params_dict = zip(model_keys, parameters)
+ state_dict = OrderedDict({k: torch.tensor(v) for k, v in params_dict})
+ self.model.set_parameters(state_dict)
+
+ if self.save_path is not None:
+ # Save Model
+ torch.save(self.model, self.save_path / f"model-ep_{server_round}.pt")
+
+ return agg_params, agg_metrics
+
+ def aggregate_evaluate(
+ self,
+ server_round: int,
+ results: List[Tuple[ClientProxy, EvaluateRes]],
+ failures: List[Union[Tuple[ClientProxy, EvaluateRes], BaseException]],
+ ) -> Tuple[Optional[float], Dict[str, Scalar]]:
+ """Aggregate the received local parameters and store the test aggregated.
+
+ Args:
+ server_round: The current round of federated learning.
+ results: Successful updates from the
+ previously selected and configured clients. Each pair of
+ `(ClientProxy, FitRes` constitutes a successful update from one of the
+ previously selected clients. Not that not all previously selected
+ clients are necessarily included in this list: a client might drop out
+ and not submit a result. For each client that did not submit an update,
+ there should be an `Exception` in `failures`.
+ failures: Exceptions that occurred while the server
+ was waiting for client updates.
+
+ Returns
+ -------
+ Optional `float` representing the aggregated evaluation result. Aggregation
+ typically uses some variant of a weighted average.
+ """
+ aggregated_loss, aggregated_metrics = super().aggregate_evaluate(
+ server_round=server_round, results=results, failures=failures
+ )
+ _ = aggregated_metrics # Avoid unused variable warning
+
+ # Weigh accuracy of each client by number of examples used
+ accuracies: List[float] = []
+ for _, res in results:
+ accuracy: float = float(res.metrics["accuracy"])
+ accuracies.append(accuracy)
+ print(f"Round {server_round} accuracies: {accuracies}")
+
+ # Aggregate and print custom metric
+ averaged_accuracy = sum(accuracies) / len(accuracies)
+ print(f"Round {server_round} accuracy averaged: {averaged_accuracy}")
+ return aggregated_loss, {"accuracy": averaged_accuracy}
diff --git a/baselines/fedper/fedper/utils.py b/baselines/fedper/fedper/utils.py
new file mode 100644
index 000000000000..00b4c5318729
--- /dev/null
+++ b/baselines/fedper/fedper/utils.py
@@ -0,0 +1,225 @@
+"""Utility functions for FedPer."""
+import os
+import pickle
+import time
+from pathlib import Path
+from secrets import token_hex
+from typing import Callable, Optional, Type, Union
+
+import matplotlib.pyplot as plt
+import numpy as np
+from flwr.server.history import History
+from omegaconf import DictConfig
+
+from fedper.client import BaseClient, FedPerClient, get_client_fn_simulation
+from fedper.implemented_models.mobile_model import MobileNet, MobileNetModelSplit
+from fedper.implemented_models.resnet_model import ResNet, ResNetModelSplit
+
+
+def set_model_class(config: DictConfig) -> DictConfig:
+ """Set model class based on the model name in the config file."""
+ # Set the model class
+ if config.model_name.lower() == "resnet":
+ config.model["_target_"] = "fedper.implemented_models.resnet_model.ResNet"
+ elif config.model_name.lower() == "mobile":
+ config.model["_target_"] = "fedper.implemented_models.mobile_model.MobileNet"
+ else:
+ raise NotImplementedError(f"Model {config.model.name} not implemented")
+ return config
+
+
+def set_num_classes(config: DictConfig) -> DictConfig:
+ """Set the number of classes based on the dataset name in the config file."""
+ # Set the number of classes
+ if config.dataset.name.lower() == "cifar10":
+ config.model.num_classes = 10
+ elif config.dataset.name.lower() == "flickr":
+ config.model.num_classes = 5
+ # additionally for flickr
+ config.batch_size = 4
+ config.num_clients = 30
+ config.clients_per_round = 30
+ else:
+ raise NotImplementedError(f"Dataset {config.dataset.name} not implemented")
+ return config
+
+
+def set_server_target(config: DictConfig) -> DictConfig:
+ """Set the server target based on the algorithm in the config file."""
+ # Set the server target
+ if config.algorithm.lower() == "fedper":
+ config.strategy["_target_"] = "fedper.server.AggregateBodyStrategyPipeline"
+ elif config.algorithm.lower() == "fedavg":
+ config.strategy["_target_"] = "fedper.server.DefaultStrategyPipeline"
+ else:
+ raise NotImplementedError(f"Algorithm {config.algorithm} not implemented")
+ return config
+
+
+def set_client_state_save_path() -> str:
+ """Set the client state save path."""
+ client_state_save_path = time.strftime("%Y-%m-%d")
+ client_state_sub_path = time.strftime("%H-%M-%S")
+ client_state_save_path = (
+ f"./client_states/{client_state_save_path}/{client_state_sub_path}"
+ )
+ if not os.path.exists(client_state_save_path):
+ os.makedirs(client_state_save_path)
+ return client_state_save_path
+
+
+def get_client_fn(
+ config: DictConfig, client_state_save_path: str = ""
+) -> Callable[[str], Union[FedPerClient, BaseClient]]:
+ """Get client function."""
+ # Get algorithm
+ algorithm = config.algorithm.lower()
+ # Get client fn
+ if algorithm == "fedper":
+ client_fn = get_client_fn_simulation(
+ config=config,
+ client_state_save_path=client_state_save_path,
+ )
+ elif algorithm == "fedavg":
+ client_fn = get_client_fn_simulation(
+ config=config,
+ )
+ else:
+ raise NotImplementedError
+ return client_fn
+
+
+def get_create_model_fn(
+ config: DictConfig,
+) -> tuple[
+ Callable[[], Union[type[MobileNet], type[ResNet]]],
+ Union[type[MobileNetModelSplit], type[ResNetModelSplit]],
+]:
+ """Get create model function."""
+ device = config.server_device
+ split: Union[
+ Type[MobileNetModelSplit], Type[ResNetModelSplit]
+ ] = MobileNetModelSplit
+ if config.model_name.lower() == "mobile":
+
+ def create_model() -> Union[Type[MobileNet], Type[ResNet]]:
+ """Create initial MobileNet-v1 model."""
+ return MobileNet(
+ num_head_layers=config.model.num_head_layers,
+ num_classes=config.model.num_classes,
+ ).to(device)
+
+ elif config.model_name.lower() == "resnet":
+ split = ResNetModelSplit
+
+ def create_model() -> Union[Type[MobileNet], Type[ResNet]]:
+ """Create initial ResNet model."""
+ return ResNet(
+ num_head_layers=config.model.num_head_layers,
+ num_classes=config.model.num_classes,
+ ).to(device)
+
+ else:
+ raise NotImplementedError("Model not implemented, check name. ")
+ return create_model, split
+
+
+def plot_metric_from_history(
+ hist: History,
+ save_plot_path: Path,
+ suffix: Optional[str] = "",
+) -> None:
+ """Plot from Flower server History.
+
+ Parameters
+ ----------
+ hist : History
+ Object containing evaluation for all rounds.
+ save_plot_path : Path
+ Folder to save the plot to.
+ suffix: Optional[str]
+ Optional string to add at the end of the filename for the plot.
+ """
+ metric_type = "distributed"
+ metric_dict = (
+ hist.metrics_centralized
+ if metric_type == "centralized"
+ else hist.metrics_distributed
+ )
+ _, values = zip(*metric_dict["accuracy"])
+
+ # let's extract decentralized loss (main metric reported in FedProx paper)
+ rounds_loss, values_loss = zip(*hist.losses_distributed)
+
+ _, axs = plt.subplots(nrows=2, ncols=1, sharex="row")
+ axs[0].plot(np.asarray(rounds_loss), np.asarray(values_loss))
+ axs[1].plot(np.asarray(rounds_loss), np.asarray(values))
+
+ axs[0].set_ylabel("Loss")
+ axs[1].set_ylabel("Accuracy")
+
+ axs[0].grid()
+ axs[1].grid()
+ # plt.title(f"{metric_type.capitalize()} Validation - MNIST")
+ plt.xlabel("Rounds")
+ # plt.legend(loc="lower right")
+
+ plt.savefig(Path(save_plot_path) / Path(f"{metric_type}_metrics{suffix}.png"))
+ plt.close()
+
+
+def save_results_as_pickle(
+ history: History,
+ file_path: Union[str, Path],
+ default_filename: Optional[str] = "results.pkl",
+) -> None:
+ """Save results from simulation to pickle.
+
+ Parameters
+ ----------
+ history: History
+ History returned by start_simulation.
+ file_path: Union[str, Path]
+ Path to file to create and store both history and extra_results.
+ If path is a directory, the default_filename will be used.
+ path doesn't exist, it will be created. If file exists, a
+ randomly generated suffix will be added to the file name. This
+ is done to avoid overwritting results.
+ extra_results : Optional[Dict]
+ A dictionary containing additional results you would like
+ to be saved to disk. Default: {} (an empty dictionary)
+ default_filename: Optional[str]
+ File used by default if file_path points to a directory instead
+ to a file. Default: "results.pkl"
+ """
+ path = Path(file_path)
+
+ # ensure path exists
+ path.mkdir(exist_ok=True, parents=True)
+
+ def _add_random_suffix(path_: Path):
+ """Add a random suffix to the file name."""
+ print(f"File `{path_}` exists! ")
+ suffix = token_hex(4)
+ print(f"New results to be saved with suffix: {suffix}")
+ return path_.parent / (path_.stem + "_" + suffix + ".pkl")
+
+ def _complete_path_with_default_name(path_: Path):
+ """Append the default file name to the path."""
+ print("Using default filename")
+ if default_filename is None:
+ return path_
+ return path_ / default_filename
+
+ if path.is_dir():
+ path = _complete_path_with_default_name(path)
+
+ if path.is_file():
+ path = _add_random_suffix(path)
+
+ print(f"Results will be saved into: {path}")
+ # data = {"history": history, **extra_results}
+ data = {"history": history}
+ # save results to pickle
+ with open(str(path), "wb") as handle:
+ pickle.dump(data, handle, protocol=pickle.HIGHEST_PROTOCOL)
diff --git a/baselines/fedper/pyproject.toml b/baselines/fedper/pyproject.toml
new file mode 100644
index 000000000000..efcdf25eface
--- /dev/null
+++ b/baselines/fedper/pyproject.toml
@@ -0,0 +1,143 @@
+[build-system]
+requires = ["poetry-core>=1.4.0"]
+build-backend = "poetry.masonry.api"
+
+[tool.poetry]
+name = "fedper" # <----- Ensure it matches the name of your baseline directory containing all the source code
+version = "1.0.0"
+description = "Federated Learning with Personalization Layers"
+license = "Apache-2.0"
+authors = ["The Flower Authors ", "William Lindskog "]
+readme = "README.md"
+homepage = "https://flower.dev"
+repository = "https://github.com/adap/flower"
+documentation = "https://flower.dev"
+classifiers = [
+ "Development Status :: 3 - Alpha",
+ "Intended Audience :: Developers",
+ "Intended Audience :: Science/Research",
+ "License :: OSI Approved :: Apache Software License",
+ "Operating System :: MacOS :: MacOS X",
+ "Operating System :: POSIX :: Linux",
+ "Programming Language :: Python",
+ "Programming Language :: Python :: 3",
+ "Programming Language :: Python :: 3 :: Only",
+ "Programming Language :: Python :: 3.8",
+ "Programming Language :: Python :: 3.9",
+ "Programming Language :: Python :: 3.10",
+ "Programming Language :: Python :: 3.11",
+ "Programming Language :: Python :: Implementation :: CPython",
+ "Topic :: Scientific/Engineering",
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
+ "Topic :: Scientific/Engineering :: Mathematics",
+ "Topic :: Software Development",
+ "Topic :: Software Development :: Libraries",
+ "Topic :: Software Development :: Libraries :: Python Modules",
+ "Typing :: Typed",
+]
+
+[tool.poetry.dependencies]
+python = ">=3.10.0, <3.11.0" # don't change this
+flwr = {extras = ["simulation"], version = "1.5.0" }
+hydra-core = "1.3.2" # don't change this
+pandas = "^2.0.3"
+matplotlib = "^3.7.2"
+tqdm = "^4.66.1"
+torch = { url = "https://download.pytorch.org/whl/cu117/torch-2.0.1%2Bcu117-cp310-cp310-linux_x86_64.whl"}
+torchvision = { url = "https://download.pytorch.org/whl/cu117/torchvision-0.15.2%2Bcu117-cp310-cp310-linux_x86_64.whl"}
+
+
+[tool.poetry.dev-dependencies]
+isort = "==5.11.5"
+black = "==23.1.0"
+docformatter = "==1.5.1"
+mypy = "==1.4.1"
+pylint = "==2.8.2"
+flake8 = "==3.9.2"
+pytest = "==6.2.4"
+pytest-watch = "==4.2.0"
+ruff = "==0.0.272"
+types-requests = "==2.27.7"
+
+[tool.isort]
+line_length = 88
+indent = " "
+multi_line_output = 3
+include_trailing_comma = true
+force_grid_wrap = 0
+use_parentheses = true
+
+[tool.black]
+line-length = 88
+target-version = ["py38", "py39", "py310", "py311"]
+
+[tool.pytest.ini_options]
+minversion = "6.2"
+addopts = "-qq"
+testpaths = [
+ "flwr_baselines",
+]
+
+[tool.mypy]
+ignore_missing_imports = true
+strict = false
+plugins = "numpy.typing.mypy_plugin"
+
+[tool.pylint."MESSAGES CONTROL"]
+disable = "bad-continuation,duplicate-code,too-few-public-methods,useless-import-alias"
+good-names = "i,j,k,_,x,y,X,Y"
+signature-mutators="hydra.main.main"
+
+[tool.pylint."TYPECHECK"]
+generated-members="numpy.*, torch.*, tensorflow.*"
+
+[[tool.mypy.overrides]]
+module = [
+ "importlib.metadata.*",
+ "importlib_metadata.*",
+]
+follow_imports = "skip"
+follow_imports_for_stubs = true
+disallow_untyped_calls = false
+
+[[tool.mypy.overrides]]
+module = "torch.*"
+follow_imports = "skip"
+follow_imports_for_stubs = true
+
+[tool.docformatter]
+wrap-summaries = 88
+wrap-descriptions = 88
+
+[tool.ruff]
+target-version = "py38"
+line-length = 88
+select = ["D", "E", "F", "W", "B", "ISC", "C4"]
+fixable = ["D", "E", "F", "W", "B", "ISC", "C4"]
+ignore = ["B024", "B027"]
+exclude = [
+ ".bzr",
+ ".direnv",
+ ".eggs",
+ ".git",
+ ".hg",
+ ".mypy_cache",
+ ".nox",
+ ".pants.d",
+ ".pytype",
+ ".ruff_cache",
+ ".svn",
+ ".tox",
+ ".venv",
+ "__pypackages__",
+ "_build",
+ "buck-out",
+ "build",
+ "dist",
+ "node_modules",
+ "venv",
+ "proto",
+]
+
+[tool.ruff.pydocstyle]
+convention = "numpy"
\ No newline at end of file
diff --git a/baselines/fedwav2vec2/.gitignore b/baselines/fedwav2vec2/.gitignore
new file mode 100644
index 000000000000..df43bf9803df
--- /dev/null
+++ b/baselines/fedwav2vec2/.gitignore
@@ -0,0 +1,2 @@
+outputs/
+data/
\ No newline at end of file
diff --git a/baselines/fedwav2vec2/LICENSE b/baselines/fedwav2vec2/LICENSE
new file mode 100644
index 000000000000..d64569567334
--- /dev/null
+++ b/baselines/fedwav2vec2/LICENSE
@@ -0,0 +1,202 @@
+
+ Apache License
+ Version 2.0, January 2004
+ http://www.apache.org/licenses/
+
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
+
+ 1. Definitions.
+
+ "License" shall mean the terms and conditions for use, reproduction,
+ and distribution as defined by Sections 1 through 9 of this document.
+
+ "Licensor" shall mean the copyright owner or entity authorized by
+ the copyright owner that is granting the License.
+
+ "Legal Entity" shall mean the union of the acting entity and all
+ other entities that control, are controlled by, or are under common
+ control with that entity. For the purposes of this definition,
+ "control" means (i) the power, direct or indirect, to cause the
+ direction or management of such entity, whether by contract or
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
+ outstanding shares, or (iii) beneficial ownership of such entity.
+
+ "You" (or "Your") shall mean an individual or Legal Entity
+ exercising permissions granted by this License.
+
+ "Source" form shall mean the preferred form for making modifications,
+ including but not limited to software source code, documentation
+ source, and configuration files.
+
+ "Object" form shall mean any form resulting from mechanical
+ transformation or translation of a Source form, including but
+ not limited to compiled object code, generated documentation,
+ and conversions to other media types.
+
+ "Work" shall mean the work of authorship, whether in Source or
+ Object form, made available under the License, as indicated by a
+ copyright notice that is included in or attached to the work
+ (an example is provided in the Appendix below).
+
+ "Derivative Works" shall mean any work, whether in Source or Object
+ form, that is based on (or derived from) the Work and for which the
+ editorial revisions, annotations, elaborations, or other modifications
+ represent, as a whole, an original work of authorship. For the purposes
+ of this License, Derivative Works shall not include works that remain
+ separable from, or merely link (or bind by name) to the interfaces of,
+ the Work and Derivative Works thereof.
+
+ "Contribution" shall mean any work of authorship, including
+ the original version of the Work and any modifications or additions
+ to that Work or Derivative Works thereof, that is intentionally
+ submitted to Licensor for inclusion in the Work by the copyright owner
+ or by an individual or Legal Entity authorized to submit on behalf of
+ the copyright owner. For the purposes of this definition, "submitted"
+ means any form of electronic, verbal, or written communication sent
+ to the Licensor or its representatives, including but not limited to
+ communication on electronic mailing lists, source code control systems,
+ and issue tracking systems that are managed by, or on behalf of, the
+ Licensor for the purpose of discussing and improving the Work, but
+ excluding communication that is conspicuously marked or otherwise
+ designated in writing by the copyright owner as "Not a Contribution."
+
+ "Contributor" shall mean Licensor and any individual or Legal Entity
+ on behalf of whom a Contribution has been received by Licensor and
+ subsequently incorporated within the Work.
+
+ 2. Grant of Copyright License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ copyright license to reproduce, prepare Derivative Works of,
+ publicly display, publicly perform, sublicense, and distribute the
+ Work and such Derivative Works in Source or Object form.
+
+ 3. Grant of Patent License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ (except as stated in this section) patent license to make, have made,
+ use, offer to sell, sell, import, and otherwise transfer the Work,
+ where such license applies only to those patent claims licensable
+ by such Contributor that are necessarily infringed by their
+ Contribution(s) alone or by combination of their Contribution(s)
+ with the Work to which such Contribution(s) was submitted. If You
+ institute patent litigation against any entity (including a
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
+ or a Contribution incorporated within the Work constitutes direct
+ or contributory patent infringement, then any patent licenses
+ granted to You under this License for that Work shall terminate
+ as of the date such litigation is filed.
+
+ 4. Redistribution. You may reproduce and distribute copies of the
+ Work or Derivative Works thereof in any medium, with or without
+ modifications, and in Source or Object form, provided that You
+ meet the following conditions:
+
+ (a) You must give any other recipients of the Work or
+ Derivative Works a copy of this License; and
+
+ (b) You must cause any modified files to carry prominent notices
+ stating that You changed the files; and
+
+ (c) You must retain, in the Source form of any Derivative Works
+ that You distribute, all copyright, patent, trademark, and
+ attribution notices from the Source form of the Work,
+ excluding those notices that do not pertain to any part of
+ the Derivative Works; and
+
+ (d) If the Work includes a "NOTICE" text file as part of its
+ distribution, then any Derivative Works that You distribute must
+ include a readable copy of the attribution notices contained
+ within such NOTICE file, excluding those notices that do not
+ pertain to any part of the Derivative Works, in at least one
+ of the following places: within a NOTICE text file distributed
+ as part of the Derivative Works; within the Source form or
+ documentation, if provided along with the Derivative Works; or,
+ within a display generated by the Derivative Works, if and
+ wherever such third-party notices normally appear. The contents
+ of the NOTICE file are for informational purposes only and
+ do not modify the License. You may add Your own attribution
+ notices within Derivative Works that You distribute, alongside
+ or as an addendum to the NOTICE text from the Work, provided
+ that such additional attribution notices cannot be construed
+ as modifying the License.
+
+ You may add Your own copyright statement to Your modifications and
+ may provide additional or different license terms and conditions
+ for use, reproduction, or distribution of Your modifications, or
+ for any such Derivative Works as a whole, provided Your use,
+ reproduction, and distribution of the Work otherwise complies with
+ the conditions stated in this License.
+
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
+ any Contribution intentionally submitted for inclusion in the Work
+ by You to the Licensor shall be under the terms and conditions of
+ this License, without any additional terms or conditions.
+ Notwithstanding the above, nothing herein shall supersede or modify
+ the terms of any separate license agreement you may have executed
+ with Licensor regarding such Contributions.
+
+ 6. Trademarks. This License does not grant permission to use the trade
+ names, trademarks, service marks, or product names of the Licensor,
+ except as required for reasonable and customary use in describing the
+ origin of the Work and reproducing the content of the NOTICE file.
+
+ 7. Disclaimer of Warranty. Unless required by applicable law or
+ agreed to in writing, Licensor provides the Work (and each
+ Contributor provides its Contributions) on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
+ implied, including, without limitation, any warranties or conditions
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
+ PARTICULAR PURPOSE. You are solely responsible for determining the
+ appropriateness of using or redistributing the Work and assume any
+ risks associated with Your exercise of permissions under this License.
+
+ 8. Limitation of Liability. In no event and under no legal theory,
+ whether in tort (including negligence), contract, or otherwise,
+ unless required by applicable law (such as deliberate and grossly
+ negligent acts) or agreed to in writing, shall any Contributor be
+ liable to You for damages, including any direct, indirect, special,
+ incidental, or consequential damages of any character arising as a
+ result of this License or out of the use or inability to use the
+ Work (including but not limited to damages for loss of goodwill,
+ work stoppage, computer failure or malfunction, or any and all
+ other commercial damages or losses), even if such Contributor
+ has been advised of the possibility of such damages.
+
+ 9. Accepting Warranty or Additional Liability. While redistributing
+ the Work or Derivative Works thereof, You may choose to offer,
+ and charge a fee for, acceptance of support, warranty, indemnity,
+ or other liability obligations and/or rights consistent with this
+ License. However, in accepting such obligations, You may act only
+ on Your own behalf and on Your sole responsibility, not on behalf
+ of any other Contributor, and only if You agree to indemnify,
+ defend, and hold each Contributor harmless for any liability
+ incurred by, or claims asserted against, such Contributor by reason
+ of your accepting any such warranty or additional liability.
+
+ END OF TERMS AND CONDITIONS
+
+ APPENDIX: How to apply the Apache License to your work.
+
+ To apply the Apache License to your work, attach the following
+ boilerplate notice, with the fields enclosed by brackets "[]"
+ replaced with your own identifying information. (Don't include
+ the brackets!) The text should be enclosed in the appropriate
+ comment syntax for the file format. We also recommend that a
+ file or class name and description of purpose be included on the
+ same "printed page" as the copyright notice for easier
+ identification within third-party archives.
+
+ Copyright [yyyy] [name of copyright owner]
+
+ Licensed under the Apache License, Version 2.0 (the "License");
+ you may not use this file except in compliance with the License.
+ You may obtain a copy of the License at
+
+ http://www.apache.org/licenses/LICENSE-2.0
+
+ Unless required by applicable law or agreed to in writing, software
+ distributed under the License is distributed on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ See the License for the specific language governing permissions and
+ limitations under the License.
diff --git a/baselines/fedwav2vec2/README.md b/baselines/fedwav2vec2/README.md
new file mode 100644
index 000000000000..0b41c6172976
--- /dev/null
+++ b/baselines/fedwav2vec2/README.md
@@ -0,0 +1,131 @@
+---
+title: Federated Learning for ASR based on Wav2vec2.0
+url: https://ieeexplore.ieee.org/document/10096426
+labels: [speech, asr, cross-device]
+dataset: [TED-LIUM 3]
+---
+
+# Federated Learning for ASR Based on wav2vec 2.0
+
+> Note: If you use this baseline in your work, please remember to cite the original authors of the paper as well as the Flower paper.
+
+**Paper:** [ieeexplore.ieee.org/document/10096426](https://ieeexplore.ieee.org/document/10096426)
+
+**Authors:** Tuan Nguyen, Salima Mdhaffar, Natalia Tomashenko, Jean-François Bonastre, Yannick Estève
+
+**Abstract:** This paper presents a study on the use of federated learning to train an ASR model based on a wav2vec 2.0 model pre-trained by self supervision. Carried out on the well-known TED-LIUM 3 dataset, our experiments show that such a model can obtain, with no use of a language model, a word error rate of 10.92% on the official TEDLIUM 3 test set, without sharing any data from the different users. We also analyse the ASR performance for speakers depending to their participation to the federated learning. Since federated learning was first introduced for privacy purposes, we also measure its ability to protect speaker identity. To do that, we exploit an approach to analyze information contained in exchanged models based on a neural network footprint on an indicator dataset. This analysis is made layer-wise and shows which layers in an exchanged wav2vec 2.0 based model bring the speaker identity information.
+
+
+## About this baseline
+
+**What’s implemented:** Figure 1 in the paper. However, this baseline only provide the SSL from figure 1. However, this baseline exclusively offers the self-supervised learning (SSL) approach as depicted in Figure 1 due to it superior performance. If you wish to implement non-SSL methods yourself, you can use the provided recipe and pre-trained model by Speechbrain, available at this link: [Speechbrain Recipe for Non-SSL](https://github.com/speechbrain/speechbrain/tree/develop/recipes/CommonVoice/ASR/seq2seq).
+
+**Datasets:** TED-LIUM 3 dataset. It requires a 54GB download. Once extracted it is ~60 GB. You can read more about this dataset in the [TED-LIUM 3](https://arxiv.org/abs/1805.04699) paper. A more concise description of this dataset can be found in the [OpenSLR](https://www.openslr.org/51/) site.
+
+**Hardware Setup:** Training `wav2vec2.0` is a bit memory intensive so you'd need at least a 24GB GPU. With the current settings, each client requires ~15GB of VRAM. This suggest you could run the experiment fine on a 16GB GPU but not if you also need to pack the global model evaluation stage on the same GPU. On a single RTX 3090Ti (24GB VRAM) each round takes between 20 and 40 minutes (depending on which clients are sampled, some clients have more data than others).
+
+**Contributors:** [Tuan Nguyen](https://www.linkedin.com/in/manh-tuan-nguyen-595898203)
+
+## Experimental Setup
+
+**Task:** Automatic Speech Recognition (ASR)
+
+**Model:** Wav2vec2.0-large [from Huggingface](https://huggingface.co/facebook/wav2vec2-large-lv60) totalling 317M parameters. Read more in the [wav2vec2.0 paper](https://arxiv.org/abs/2006.11477).
+
+
+**Dataset:** In this paper, we divided the training dataset of TED-LIUM 3 into 1943 clients, where each of them is represented by a speaker from TED-LIUM 3. The clients are ordered by CID, with `client_0` having the largest amount of speech hours and `client_1943` having the smallest. Each client's data will be divided into training, development, and test sets with an 80-10-10 ratio. For client who has more than 10 minutes, we extract 5 minutes from their training set for analysis purposes. This portion will not be used during training or in any part of this baseline. For clients with duration less than 10 minutes, all the speaker data will represent the local dataset for the client. The full structure breakdown is below:
+```bash
+├── data
+│ ├── client_{cid}
+│ │ ├── ted_train.csv
+│ │ ├── ted_dev.csv
+│ │ ├── ted_test.csv
+│ │ ├── ted_train_full5.csv {Analysis dataset contains only 5m from ted_train.csv}
+│ │ ├── ted_train_wo5.csv {the training file for client who has more than 10m}
+│ ├── server
+│ │ ├── ted_train.csv {all TED-LIUM 3 train set}
+│ │ ├── ted_dev.csv {all TED-LIUM 3 valid set}
+│ │ ├── ted_test.csv {all TED-LIUM 3 test set}
+
+```
+For more details, please refer to the relevant section in the paper.
+
+**Training Hyperparameters:**
+| Hyperparameter | Default Value | Description |
+| ------- | ----- | ------- |
+| `pre_train_model_path` | `null` | Path to pre-trained model or checkpoint. The best checkpoint could be found [here](https://github.com/tuanct1997/Federated-Learning-ASR-based-on-wav2vec-2.0/tree/main/material/pre-trained) |
+| `save_checkpoint` | `null` | Path to folder where server model will be saved at each round |
+| `label_path` | `docs/pretrained_wav2vec2` | Label each character for every client to ensure consistency during training phase|
+| `sb_config` | `fedwav2vec2/conf/sb_config/w2v2.yaml` | Speechbrain config file for architecture model. Please refer to [SpeechBrain](https://github.com/speechbrain/speechbrain) for more information |
+| `rounds` | `100` | Indicate the number of Federated Learning (FL) rounds|
+| `local_epochs` | `20` | Specify the number of training epochs at the client side |
+| `total_clients` | `1943` | Size of client pool, with a maxium set at 1943 clients|
+| `server_cid` | `19999` | ID of the server to distinguish from the client's ID |
+| `server_device` | `cuda` | You can choose between `cpu` or `cuda` for centralised evaluation, but it is recommended to use `cuda`|
+| `parallel_backend` | `false` | Multi-gpus training. Only active if you have more than 1 gpu per client |
+| `strategy.min_fit_client` | `20` | Number of clients involve per round. Default is 20 as indicated in the paper |
+| `strategy.fraction_fit` | `0.01` | Ratio of client pool to involve during training |
+| `strategy.weight_strategy` | `num`| Different way to average clients weight. Could be chose between `num`,`loss`,`wer` |
+| `client_resources.num_cpus` | `8`| Number of cpus per client. Recommended to have more than 8 |
+| `client_resources.num_gpus` | `1`| Number of gpus per client. Recommended to have at least 1 with VRAM > 24GB |
+
+
+By default, long audio sequences (>10s) are excluded from training. This is done so to keep the VRAM usage low enough to train a client on a 16GB GPU. This hyperparameter is defined in the `sb_config` under the `avoid_if_longer_than` tag.
+
+## Environment Setup
+
+Once you have installed `pyenv` and `poetry`, run the commands below to setup your python environment:
+
+```bash
+# Set a recent version of Python for your environment
+pyenv local 3.10.6
+poetry env use 3.10.6
+
+# Install your environment
+poetry install
+
+# Activate your environment
+poetry shell
+```
+
+When you run this baseline for the first time, you need first to download the data-to-client mapping files as well as the `TED-LIUM-3`` dataset.
+
+```bash
+# Then create a directory using the same name as you'll use for `dada_dir` in your config (see conf/base.yaml)
+mkdir data
+
+# Clone client mapping (note content will be moved to your data dir)
+git clone https://github.com/tuanct1997/Federated-Learning-ASR-based-on-wav2vec-2.0.git _temp && mv _temp/data/* data/ && rm -rf _temp
+
+# Download dataset, extract and prepare dataset partitions
+# This might take a while depending on your internet connection
+python -m fedwav2vec2.dataset_preparation
+```
+
+
+## Running the Experiments
+
+```bash
+# Run with default arguments (one client per GPU)
+python -m fedwav2vec2.main
+
+# if you have a large GPU (32GB+) you migth want to fit two per GPU
+python -m fedwav2vec2.main client_resources.num_gpus=0.5
+
+# the global model can be saved at the end of each round if you specify a checkpoint path
+python -m fedwav2vec2.main save_checkpoint= # if directory doesn't exist, it will be created
+
+# then you can use it as the starting point for your global model like so:
+python -m fedwav2vec2.main pre_train_model_path=/last_checkpoint.pt
+```
+
+When running the experiment, a structure of directories `/` will be created by Hydra. Inside you'll find a directory for each client (where their log is recorded). Another directory at the same level `/server` is created where the server log is recorded. For this baseline the metric of interes it the Word Error Rate (`WER`) which is logged in `train_log.txt` at the end of each round.
+
+
+## Expected Results
+
+Running the command above will generate the `SSL` results as shown on the plot below. The results should closely follow those in Figure 1 in the paper.
+
+
+
+
diff --git a/baselines/fedwav2vec2/_static/fedwav2vec.png b/baselines/fedwav2vec2/_static/fedwav2vec.png
new file mode 100644
index 000000000000..27a2a7c5d7c1
Binary files /dev/null and b/baselines/fedwav2vec2/_static/fedwav2vec.png differ
diff --git a/baselines/fedwav2vec2/docs/label_encoder.txt b/baselines/fedwav2vec2/docs/label_encoder.txt
new file mode 100644
index 000000000000..654e01e1065d
--- /dev/null
+++ b/baselines/fedwav2vec2/docs/label_encoder.txt
@@ -0,0 +1,54 @@
+'t' => 50
+'h' => 1
+'e' => 2
+'a' => 3
+'_' => 4
+'o' => 5
+'n' => 6
+'l' => 7
+'i' => 8
+'r' => 9
+'s' => 10
+'p' => 11
+'d' => 12
+'w' => 13
+'u' => 14
+'k' => 15
+'c' => 16
+'m' => 17
+'y' => 18
+'v' => 19
+'z' => 20
+'f' => 21
+'b' => 22
+'g' => 23
+'j' => 24
+"'" => 25
+'x' => 26
+'q' => 27
+'4' => 28
+'2' => 29
+'7' => 30
+'[' => 31
+']' => 32
+'1' => 33
+'9' => 34
+'0' => 35
+'5' => 36
+'3' => 37
+'6' => 38
+'=' => 39
+'%' => 40
+'$' => 41
+'8' => 42
+'#' => 43
+'ā' => 44
+'&' => 45
+'+' => 46
+'@' => 47
+'^' => 48
+'\\' => 49
+'' => 0
+================
+'starting_index' => 0
+'blank_label' => ''
diff --git a/baselines/fedwav2vec2/fedwav2vec2/__init__.py b/baselines/fedwav2vec2/fedwav2vec2/__init__.py
new file mode 100644
index 000000000000..a5e567b59135
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/__init__.py
@@ -0,0 +1 @@
+"""Template baseline package."""
diff --git a/baselines/fedwav2vec2/fedwav2vec2/client.py b/baselines/fedwav2vec2/fedwav2vec2/client.py
new file mode 100644
index 000000000000..319580a83845
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/client.py
@@ -0,0 +1,172 @@
+"""Define your client class and a function to construct such clients.
+
+Please overwrite `flwr.client.NumPyClient` or `flwr.client.Client` and create a function
+to instantiate your client.
+"""
+
+
+import gc
+import logging
+from math import exp
+
+import flwr as fl
+import speechbrain as sb
+import torch
+from flwr.common import (
+ Code,
+ EvaluateIns,
+ EvaluateRes,
+ FitIns,
+ FitRes,
+ GetParametersIns,
+ GetParametersRes,
+ NDArrays,
+ Status,
+ ndarrays_to_parameters,
+ parameters_to_ndarrays,
+)
+from omegaconf import DictConfig
+
+from fedwav2vec2.models import int_model
+from fedwav2vec2.sb_recipe import get_weights, set_weights
+
+
+class SpeechBrainClient(fl.client.Client):
+ """Flower client for SpeechBrain."""
+
+ def __init__(self, cid: str, asr_brain, dataset):
+ self.cid = cid
+ self.params = asr_brain.hparams
+ self.modules = asr_brain.modules
+ self.asr_brain = asr_brain
+ self.dataset = dataset
+
+ fl.common.logger.log(logging.DEBUG, "Starting client %s", cid)
+
+ def get_parameters(self, _: GetParametersIns) -> GetParametersRes:
+ """Return the parameters of the current net."""
+ weights: NDArrays = get_weights(self.modules)
+ parameters = ndarrays_to_parameters(weights)
+ gc.collect()
+ status = Status(code=Code.OK, message="Success")
+ return GetParametersRes(status=status, parameters=parameters)
+
+ def fit(self, ins: FitIns) -> FitRes:
+ """Implement distributed fit function for a given client."""
+ weights: NDArrays = fl.common.parameters_to_ndarrays(ins.parameters)
+ config = ins.config
+
+ # Read training configuration
+ epochs = int(config["epochs"])
+
+ (_, num_examples, avg_loss, avg_wer) = self._train_speech_recogniser(
+ weights, epochs
+ )
+ metrics = {"train_loss": avg_loss, "wer": avg_wer}
+
+ parameters = self.get_parameters(GetParametersIns(config={})).parameters
+ del self.asr_brain.modules
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ gc.collect()
+
+ status = Status(code=Code.OK, message="Success")
+
+ return FitRes(
+ status=status,
+ parameters=parameters,
+ num_examples=num_examples,
+ metrics=metrics,
+ )
+
+ def evaluate(self, ins: EvaluateIns) -> EvaluateRes:
+ """Implement distributed evaluation for a given client."""
+ weights = parameters_to_ndarrays(ins.parameters)
+
+ num_examples, loss, wer = self.evaluate_train_speech_recogniser(
+ server_params=weights,
+ epochs=1,
+ )
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ gc.collect()
+
+ status = Status(code=Code.OK, message="Success")
+ # Return the number of evaluation examples and the evaluation result (loss)
+ return EvaluateRes(
+ status=status,
+ num_examples=num_examples,
+ loss=float(loss),
+ metrics={"Error rate": float(wer)},
+ )
+
+ def evaluate_train_speech_recogniser(self, server_params, epochs):
+ """Evaluate aggerate/server model."""
+ _, _, test_data = self._setup_task(server_params, epochs)
+ self.params.wer_file = self.params.output_folder + "/wer_test.txt"
+
+ batch_count, loss, wer = self.asr_brain.evaluate(
+ test_data,
+ test_loader_kwargs=self.params.test_dataloader_options,
+ )
+
+ return batch_count, float(loss), float(wer)
+
+ def _setup_task(
+ self,
+ server_params,
+ epochs,
+ ):
+ self.params.epoch_counter.limit = epochs
+ self.params.epoch_counter.current = 0
+
+ train_data, valid_data, test_data = self.dataset
+ # Set the parameters to the ones given by the server
+ if server_params is not None:
+ set_weights(server_params, self.modules, self.params.device)
+ return train_data, valid_data, test_data
+
+ def _train_speech_recogniser(self, server_params, epochs):
+ train_data, valid_data, _ = self._setup_task(server_params, epochs)
+
+ # Training
+ count_sample, avg_loss, avg_wer = self.asr_brain.fit(
+ self.params.epoch_counter,
+ train_data,
+ valid_data,
+ train_loader_kwargs=self.params.dataloader_options,
+ valid_loader_kwargs=self.params.test_dataloader_options,
+ )
+ # exp operation to avg_loss and avg_wer
+ avg_wer = 100 if avg_wer > 100 else avg_wer
+ avg_loss = exp(-avg_loss)
+ avg_wer = exp(100 - avg_wer)
+
+ # retrieve the parameters to return
+ params_list = get_weights(self.modules)
+
+ # Manage when last batch isn't full w.r.t batch size
+ train_set = sb.dataio.dataloader.make_dataloader(
+ train_data, **self.params.dataloader_options
+ )
+ if count_sample > len(train_set) * self.params.batch_size * epochs:
+ count_sample = len(train_set) * self.params.batch_size * epochs
+
+ del train_data, valid_data
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ gc.collect()
+ return (params_list, count_sample, avg_loss, avg_wer)
+
+
+def get_client_fn(config: DictConfig, save_path: str):
+ """Return a function that creates a Flower client."""
+
+ def client_fn(cid: str) -> fl.client.Client:
+ """Generate the simulated clients."""
+ device = "cuda" if torch.cuda.is_available() else "cpu"
+
+ asr_brain, dataset = int_model(cid, config, device=device, save_path=save_path)
+ return SpeechBrainClient(cid, asr_brain, dataset)
+
+ return client_fn
diff --git a/baselines/fedwav2vec2/fedwav2vec2/conf/base.yaml b/baselines/fedwav2vec2/fedwav2vec2/conf/base.yaml
new file mode 100644
index 000000000000..0df942aa1c12
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/conf/base.yaml
@@ -0,0 +1,43 @@
+---
+# this is the config that will be loaded as default by main.py
+# Please follow the provided structure (this will ensuring all baseline follow
+# a similar configuration structure and hence be easy to customise)
+
+pre_train_model_path: null # Path to checkpoint exp: docs/checkpoint/last_checkpoint.pt
+save_checkpoint: null # Path to folder for checkpoint
+
+# Path for label encoder file if want to ensure the same encode for every client
+label_path: docs/label_encoder.txt
+
+huggingface_model_save_path: docs/pretrained_wav2vec2
+sb_config: fedwav2vec2/conf/sb_config/w2v2.yaml # config with SpeechBrain recipe for Wav2Vec 2.0
+data_path: data # if you change this, ensure you `git cloned` the author's own repo to a directory with the new name
+rounds: 100 # global FL rounds
+local_epochs: 20 # local epochs for each client
+total_clients: 1943
+server_cid: 19999
+
+# Device setup
+server_device: cuda
+parallel_backend: false # If using multi-gpus per client (disable it if using server_device=cpu)
+
+client_resources:
+ num_cpus: 8
+ num_gpus: 1
+
+dataset:
+ download_filename: TEDLIUM_release-3.tgz
+
+ extract_subdirectory: audio
+
+
+strategy:
+ _target_: fedwav2vec2.strategy.CustomFedAvg
+ min_fit_clients: 20
+ fraction_fit: 0.01
+ fraction_evaluate: 0.00
+ min_available_clients: ${total_clients}
+ weight_strategy: num # strategy of weighting clients in: [num, loss, wer]
+ on_fit_config_fn:
+ _target_: fedwav2vec2.server.get_on_fit_config_fn
+ local_epochs: ${local_epochs}
diff --git a/baselines/fedwav2vec2/fedwav2vec2/conf/sb_config/w2v2.yaml b/baselines/fedwav2vec2/fedwav2vec2/conf/sb_config/w2v2.yaml
new file mode 100644
index 000000000000..87bf21080b0e
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/conf/sb_config/w2v2.yaml
@@ -0,0 +1,178 @@
+# ################################
+# Model: seq2seq ASR on TIMIT with wav2vec2 + CTC/Attention
+# Authors:
+# * Titouan Parcollet 2021
+# ################################
+
+# Seed needs to be set at top of yaml, before objects with parameters are made
+seed: 1234
+__set_seed: !!python/object/apply:torch.manual_seed [!ref ]
+output_folder: !ref docs/results/fl_wav2vec2/
+wer_file: !ref /wer.txt
+save_folder: !ref /save
+train_log: !ref /train_log.txt
+label_encode: docs/label_encoder.txt
+
+wav2vec_output: docs/pretrained_wav2vec2
+
+# URL for the biggest huggingface english wav2vec2 model.
+
+wav2vec2_hub: "facebook/wav2vec2-large-lv60"
+# Data files
+data_folder: data/audio/TEDLIUM_release-3/data
+
+accented_letters: False
+language: en # use 'it' for Italian, 'rw' for Kinyarwanda, 'en' for english
+train_csv: data/ted_train.csv
+valid_csv: data/ted_dev.csv
+test_csv: data/ted_test.csv
+skip_prep: False # Skip data preparation
+device: cpu # will be overwritten by `server_device` in `conf/base.yaml`
+# Training parameters
+
+avoid_if_longer_than: 10.0
+avoid_if_smaller_than: 0.1
+
+
+# Decoding parameters
+blank_index: 0
+bos_index: 1
+eos_index: 2
+
+# Training parameters
+number_of_epochs: 30
+lr: 0.001
+lr_wav2vec: 0.0001
+
+ctc_weight: 0.2
+sorting: descending
+
+# With data_parallel batch_size is split into N jobs
+# With DDP batch_size is multiplied by N jobs
+# Must be 6 per GPU to fit 16GB of VRAM
+batch_size: 4
+test_batch_size: 1
+
+dataloader_options:
+ batch_size: !ref
+ num_workers: 4
+test_dataloader_options:
+ batch_size: !ref
+ num_workers: 4
+
+# BPE parameters
+token_type: unigram # ["unigram", "bpe", "char"]
+character_coverage: 1.0
+
+# Feature parameters (FBANKS etc)
+sample_rate: 16000
+
+
+# Model parameters
+activation: !name:torch.nn.LeakyReLU
+dropout: 0.5
+dnn_neurons: 500
+dec_neurons: 500
+
+#encoder with w2v
+enc_dnn_layers: 1
+enc_dnn_neurons: 1024
+
+# Outputs
+output_neurons: 52
+
+# Decoding parameters
+# Be sure that the bos and eos index match with the BPEs ones
+# blank_index: 0
+beam_size: 20
+temperature: 1.50
+
+#
+# Functions and classes
+#
+epoch_counter: !new:speechbrain.utils.epoch_loop.EpochCounter
+ limit: !ref
+
+
+
+wav2vec2: !new:speechbrain.lobes.models.huggingface_wav2vec.HuggingFaceWav2Vec2
+ source: !ref
+ output_norm: True
+ freeze: False
+
+ save_path: !ref /wav2vec2_checkpoint
+
+ # A simple DNN that receive as inputs the output of the wav2vec2 model
+ # Here the output dimensionality of the LARGE wav2vec2 is 1024.
+enc: !new:speechbrain.lobes.models.VanillaNN.VanillaNN
+ input_shape: [null, null, 1024]
+ activation: !ref
+ dnn_blocks: !ref
+ dnn_neurons: !ref
+
+
+ctc_lin: !new:speechbrain.nnet.linear.Linear
+ input_size: !ref
+ n_neurons: !ref
+ bias: True
+
+log_softmax: !new:speechbrain.nnet.activations.Softmax
+ apply_log: True
+
+
+modules:
+ wav2vec2: !ref
+ enc: !ref
+ ctc_lin: !ref
+
+model: !new:torch.nn.ModuleList
+ - [!ref , !ref ]
+
+adam_opt_class: !name:torch.optim.Adam
+ lr: !ref
+
+wav2vec_opt_class: !name:torch.optim.Adam
+ lr: !ref
+
+ctc_cost: !name:speechbrain.nnet.losses.ctc_loss
+ blank_index: !ref
+
+lr_annealing_adam: !new:speechbrain.nnet.schedulers.NewBobScheduler
+ initial_value: !ref
+ improvement_threshold: 0.0025
+ annealing_factor: 0.8
+ patient: 0
+
+lr_annealing_wav2vec: !new:speechbrain.nnet.schedulers.NewBobScheduler
+ initial_value: !ref
+ improvement_threshold: 0.0025
+ annealing_factor: 0.9
+
+checkpointer: !new:speechbrain.utils.checkpoints.Checkpointer
+ checkpoints_dir: !ref
+ recoverables:
+ model: !ref
+ wav2vec2: !ref
+ lr_annealing_adam: !ref
+ lr_annealing_wav2vec: !ref
+ counter: !ref
+
+train_logger: !new:speechbrain.utils.train_logger.FileTrainLogger
+ save_file: !ref
+
+
+ctc_computer: !name:speechbrain.utils.metric_stats.MetricStats
+ metric: !name:speechbrain.nnet.losses.ctc_loss
+ blank_index: !ref
+ reduction: batch
+
+
+error_rate_computer: !name:speechbrain.utils.metric_stats.ErrorRateStats
+
+cer_computer: !name:speechbrain.utils.metric_stats.ErrorRateStats
+ merge_tokens: True
+
+coer_computer: !name:speechbrain.utils.metric_stats.ErrorRateStats
+
+cver_computer: !name:speechbrain.utils.metric_stats.ErrorRateStats
+
diff --git a/baselines/fedwav2vec2/fedwav2vec2/dataset.py b/baselines/fedwav2vec2/fedwav2vec2/dataset.py
new file mode 100644
index 000000000000..65fee90faf38
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/dataset.py
@@ -0,0 +1,125 @@
+"""Handle basic dataset creation.
+
+In case of PyTorch it should return dataloaders for your dataset (for both the clients
+and the server). If you are using a custom dataset class, this module is the place to
+define it. If your dataset requires to be downloaded (and this is not done
+automatically -- e.g. as it is the case for many dataset in TorchVision) and
+partitioned, please include all those functions and logic in the
+`dataset_preparation.py` module. You can use all those functions from functions/methods
+defined here of course.
+"""
+
+
+import speechbrain as sb
+import torchaudio
+
+
+# Define custom data procedure
+def dataio_prepare(hparams):
+ """Create Dataset objects from the CSV files."""
+ # 1. Define datasets
+ data_folder = hparams["data_folder"]
+
+ train_data = sb.dataio.dataset.DynamicItemDataset.from_csv(
+ csv_path=hparams["train_csv"],
+ replacements={"data_root": data_folder},
+ )
+
+ if hparams["sorting"] == "ascending":
+ # we sort training data to speed up training and get better results.
+ train_data = train_data.filtered_sorted(
+ sort_key="duration",
+ key_max_value={"duration": hparams["avoid_if_longer_than"]},
+ key_min_value={"duration": hparams["avoid_if_smaller_than"]},
+ )
+ # when sorting do not shuffle in dataloader ! otherwise is pointless
+ hparams["dataloader_options"]["shuffle"] = False
+
+ elif hparams["sorting"] == "descending":
+ train_data = train_data.filtered_sorted(
+ sort_key="duration",
+ reverse=True,
+ key_max_value={"duration": hparams["avoid_if_longer_than"]},
+ key_min_value={"duration": hparams["avoid_if_smaller_than"]},
+ )
+ # when sorting do not shuffle in dataloader ! otherwise is pointless
+ hparams["dataloader_options"]["shuffle"] = False
+
+ elif hparams["sorting"] == "random":
+ pass
+
+ else:
+ raise NotImplementedError("sorting must be random, ascending or descending")
+
+ valid_data = sb.dataio.dataset.DynamicItemDataset.from_csv(
+ csv_path=hparams["valid_csv"],
+ replacements={"data_root": data_folder},
+ )
+ # We also sort the validation data so it is faster to validate
+ valid_data = valid_data.filtered_sorted(
+ sort_key="duration",
+ reverse=True,
+ key_max_value={"duration": hparams["avoid_if_longer_than"]},
+ key_min_value={"duration": hparams["avoid_if_smaller_than"]},
+ )
+
+ test_data = sb.dataio.dataset.DynamicItemDataset.from_csv(
+ csv_path=hparams["test_csv"],
+ replacements={"data_root": data_folder},
+ )
+ # We also sort the test data so it is faster to validate
+ test_data = test_data.filtered_sorted(
+ sort_key="duration",
+ reverse=True,
+ key_max_value={"duration": hparams["avoid_if_longer_than"]},
+ key_min_value={"duration": hparams["avoid_if_smaller_than"]},
+ )
+
+ datasets = [train_data, valid_data, test_data]
+
+ label_encoder = sb.dataio.encoder.CTCTextEncoder()
+
+ # 2. Define audio pipeline:
+ @sb.utils.data_pipeline.takes("wav", "start_seg", "end_seg")
+ @sb.utils.data_pipeline.provides("sig")
+ def audio_pipeline(wav, start_seg, end_seg):
+ info = torchaudio.info(wav)
+ start = int(float(start_seg) * hparams["sample_rate"])
+ stop = int(float(end_seg) * hparams["sample_rate"])
+ speech_segment = {"file": wav, "start": start, "stop": stop}
+ sig = sb.dataio.dataio.read_audio(speech_segment)
+ # resample to correct 16Hz if different or else remain the same
+ resampled = torchaudio.transforms.Resample(
+ info.sample_rate,
+ hparams["sample_rate"],
+ )(sig)
+ return resampled
+
+ sb.dataio.dataset.add_dynamic_item(datasets, audio_pipeline)
+
+ # 3. Define text pipeline:
+ @sb.utils.data_pipeline.takes("char")
+ @sb.utils.data_pipeline.provides("char_list", "char_encoded")
+ def text_pipeline(char):
+ char_list = char.strip().split()
+ yield char_list
+ char_encoded = label_encoder.encode_sequence_torch(char_list)
+ yield char_encoded
+
+ sb.dataio.dataset.add_dynamic_item(datasets, text_pipeline)
+
+ lab_enc_file = hparams["label_encoder"]
+ label_encoder.load_or_create(
+ path=lab_enc_file,
+ from_didatasets=[train_data],
+ output_key="char_list",
+ special_labels={"blank_label": hparams["blank_index"]},
+ sequence_input=True,
+ )
+
+ # 4. Set output:
+ sb.dataio.dataset.set_output_keys(
+ datasets,
+ ["id", "sig", "char_encoded"],
+ )
+ return train_data, valid_data, test_data, label_encoder
diff --git a/baselines/fedwav2vec2/fedwav2vec2/dataset_preparation.py b/baselines/fedwav2vec2/fedwav2vec2/dataset_preparation.py
new file mode 100644
index 000000000000..255664d3d2c3
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/dataset_preparation.py
@@ -0,0 +1,130 @@
+"""Handle the dataset partitioning and (optionally) complex downloads.
+
+Please add here all the necessary logic to either download, uncompress, pre/post-process
+your dataset (or all of the above). If the desired way of running your baseline is to
+first download the dataset and partition it and then run the experiments, please
+uncomment the lines below and tell us in the README.md (see the "Running the Experiment"
+block) that this file should be executed first.
+"""
+
+
+import os
+import ssl
+import tarfile
+import urllib.request
+from shutil import rmtree
+
+import hydra
+import pandas as pd
+from hydra.core.hydra_config import HydraConfig
+from omegaconf import DictConfig, OmegaConf
+
+
+def _download_file(url, filename):
+ """Download the file and show a progress bar."""
+ print(f"Downloading {url}...")
+ retries = 3
+ while retries > 0:
+ try:
+ with urllib.request.urlopen(
+ url,
+ # pylint: disable=protected-access
+ context=ssl._create_unverified_context(),
+ ) as response, open(filename, "wb") as out_file:
+ total_size = int(response.getheader("Content-Length"))
+ block_size = 1024 * 8
+ count = 0
+ while True:
+ data = response.read(block_size)
+ if not data:
+ break
+ count += 1
+ out_file.write(data)
+ percent = int(count * block_size * 100 / total_size)
+ print(
+ f"\rDownload: {percent}% [{count * block_size}/{total_size}]",
+ end="",
+ )
+ print("\nDownload complete.")
+ break
+ except Exception as error: # pylint: disable=broad-except
+ print(f"\nError occurred during download: {error}")
+ retries -= 1
+ if retries > 0:
+ print(f"Retrying ({retries} retries left)...")
+ else:
+ print("Download failed.")
+ raise error
+
+
+def _extract_file(filename, extract_path):
+ """Extract the contents and show a progress bar."""
+ print(f"Extracting {filename}...")
+ with tarfile.open(filename, "r:gz") as tar:
+ members = tar.getmembers()
+ total_files = len(members)
+ current_file = 0
+ for member in members:
+ current_file += 1
+ tar.extract(member, path=extract_path)
+ percent = int(current_file * 100 / total_files)
+ print(f"\rExtracting: {percent}% [{current_file}/{total_files}]", end="")
+ print("\nExtraction complete.")
+
+
+def _delete_file(filename):
+ """Delete the downloaded file."""
+ os.remove(filename)
+ print(f"Deleted {filename}.")
+
+
+def _csv_path_audio(extract_path: str):
+ """Change the path corespond to your actual path."""
+ for subdir, _dirs, files in os.walk("./data"):
+ for file in files:
+ if file.endswith(".csv"):
+ if "client" in subdir:
+ path = path = os.path.join(extract_path, "legacy/train/sph")
+ else:
+ if "train" in file:
+ path = os.path.join(extract_path, "legacy/train/sph")
+ elif "dev" in file:
+ path = os.path.join(extract_path, "legacy/dev/sph")
+ else:
+ path = os.path.join(extract_path, "legacy/test/sph")
+ d_f = pd.read_csv(os.path.join(subdir, file))
+ d_f["wav"] = d_f["wav"].str.replace("path", path)
+ d_f.to_csv(os.path.join(subdir, file), index=False)
+
+
+@hydra.main(config_path="./conf", config_name="base", version_base=None)
+def download_and_extract(cfg: DictConfig) -> None:
+ """Download and extract TEDIUM-3 dataset."""
+ print(OmegaConf.to_yaml(cfg))
+ url = (
+ "https://projets-lium.univ-lemans.fr"
+ "/wp-content/uploads/corpus/TED-LIUM/TEDLIUM_release-3.tgz"
+ )
+ # URL = "https://www.openslr.org/resources/51/TEDLIUM_release-3.tgz"
+ filename = f"{cfg.data_path}/{cfg.dataset.download_filename}"
+ extract_path = f"{cfg.data_path}/{cfg.dataset.extract_subdirectory}"
+
+ print(f"{extract_path = }")
+ print(f"{filename = }")
+
+ if not os.path.exists(extract_path):
+ try:
+ _download_file(url, filename)
+ _extract_file(filename, extract_path)
+ finally:
+ _delete_file(filename)
+
+ _csv_path_audio(f"{extract_path}/TEDLIUM_release-3")
+
+ # remove output dir. No need to keep it around
+ save_path = HydraConfig.get().runtime.output_dir
+ rmtree(save_path)
+
+
+if __name__ == "__main__":
+ download_and_extract()
diff --git a/baselines/fedwav2vec2/fedwav2vec2/main.py b/baselines/fedwav2vec2/fedwav2vec2/main.py
new file mode 100644
index 000000000000..5011b2ac15e2
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/main.py
@@ -0,0 +1,60 @@
+"""Create and connect the building blocks for your experiments; start the simulation.
+
+It includes processioning the dataset, instantiate strategy, specify how the global
+model is going to be evaluated, etc. At the end, this script saves the results.
+"""
+
+import flwr as fl
+import hydra
+from hydra.core.hydra_config import HydraConfig
+from hydra.utils import instantiate
+from omegaconf import DictConfig, OmegaConf
+
+from fedwav2vec2.client import get_client_fn
+from fedwav2vec2.models import pre_trained_point
+from fedwav2vec2.server import get_evaluate_fn
+
+
+@hydra.main(config_path="conf", config_name="base", version_base=None)
+def main(cfg: DictConfig) -> None:
+ """Run the baseline.
+
+ Parameters
+ ----------
+ cfg : DictConfig
+ An omegaconf object that stores the hydra config.
+ """
+ # 1. Print parsed config
+ print(OmegaConf.to_yaml(cfg))
+
+ # Hydra automatically creates an output directory
+ # Let's retrieve it and save some results there
+ save_path = HydraConfig.get().runtime.output_dir
+
+ if cfg.pre_train_model_path is not None:
+ print("PRETRAINED INITIALIZE")
+
+ pretrained = pre_trained_point(save_path, cfg, cfg.server_device)
+ else:
+ pretrained = None
+
+ strategy = instantiate(
+ cfg.strategy,
+ initial_parameters=pretrained,
+ evaluate_fn=get_evaluate_fn(
+ cfg, server_device=cfg.server_device, save_path=save_path
+ ),
+ )
+
+ fl.simulation.start_simulation(
+ client_fn=get_client_fn(cfg, save_path),
+ num_clients=cfg.total_clients,
+ client_resources=cfg.client_resources,
+ config=fl.server.ServerConfig(num_rounds=cfg.rounds),
+ strategy=strategy,
+ ray_init_args={"include_dashboard": False},
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/baselines/fedwav2vec2/fedwav2vec2/models.py b/baselines/fedwav2vec2/fedwav2vec2/models.py
new file mode 100644
index 000000000000..38916aaac555
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/models.py
@@ -0,0 +1,148 @@
+"""Define our models, and training and eval functions.
+
+If your model is 100% off-the-shelf (e.g. directly from torchvision without requiring
+modifications) you might be better off instantiating your model directly from the Hydra
+config. In this way, swapping your model for another one can be done without changing
+the python code at all
+"""
+
+
+import gc
+import os
+
+import speechbrain as sb
+import torch
+from flwr.common import ndarrays_to_parameters
+from hyperpyyaml import load_hyperpyyaml
+from omegaconf import DictConfig
+
+from fedwav2vec2.dataset import dataio_prepare
+from fedwav2vec2.sb_recipe import ASR, get_weights
+
+
+def int_model( # pylint: disable=too-many-arguments,too-many-locals
+ cid,
+ config: DictConfig,
+ device: str,
+ save_path,
+ evaluate=False,
+):
+ """Set up the experiment.
+
+ Loading the hyperparameters from config files and command-line overrides, setting
+ the correct path for the corresponding clients, and creating the model.
+ """
+ # Load hyperparameters file with command-line overrides
+
+ if cid == 19999:
+ save_path = save_path + "server"
+ else:
+ save_path = save_path + "/client_" + str(cid)
+
+ # Override with FLOWER PARAMS
+ if evaluate:
+ overrides = {
+ "output_folder": save_path,
+ "number_of_epochs": 1,
+ "test_batch_size": 4,
+ "device": device,
+ "wav2vec_output": config.huggingface_model_save_path,
+ }
+
+ else:
+ overrides = {
+ "output_folder": save_path,
+ "wav2vec_output": config.huggingface_model_save_path,
+ }
+
+ label_path_ = config.label_path
+ if label_path_ is None:
+ label_path_ = os.path.join(save_path, "label_encoder.txt")
+
+ _, run_opts, _ = sb.parse_arguments(config.sb_config)
+ run_opts["device"] = device
+ run_opts["data_parallel_backend"] = config.parallel_backend
+ run_opts["noprogressbar"] = True # disable tqdm progress bar
+
+ with open(config.sb_config) as fin:
+ params = load_hyperpyyaml(fin, overrides)
+
+ # This logic follow the data_path is a path to csv folder file
+ # All train/dev/test csv files are in the same name format for server and client
+ # Example:
+ # server: /users/server/train.csv
+ # client: /users/client_1/train.csv
+ # Modify (if needed) the if else logic to fit with path format
+
+ if int(cid) != config.server_cid:
+ params["data_folder"] = os.path.join(config.data_path, "client_" + str(cid))
+ else:
+ params["data_folder"] = os.path.join(config.data_path, "server")
+
+ print(f'{params["data_folder"] = }')
+ params["train_csv"] = params["data_folder"] + "/ted_train.csv"
+ params["valid_csv"] = params["data_folder"] + "/ted_dev.csv"
+ params["test_csv"] = params["data_folder"] + "/ted_test.csv"
+
+ if int(cid) < 1341:
+ params["train_csv"] = params["data_folder"] + "/ted_train_wo5.csv"
+ params["label_encoder"] = label_path_
+
+ # Create experiment directory
+ sb.create_experiment_directory(
+ experiment_directory=params["output_folder"],
+ hyperparams_to_save=config.sb_config,
+ overrides=overrides,
+ )
+
+ # Create the datasets objects as well as tokenization and encoding :-D
+ train_data, valid_data, test_data, label_encoder = dataio_prepare(params)
+ # Trainer initialization
+
+ asr_brain = ASR(
+ modules=params["modules"],
+ hparams=params,
+ run_opts=run_opts,
+ checkpointer=params["checkpointer"],
+ )
+ asr_brain.label_encoder = label_encoder
+ asr_brain.label_encoder.add_unk()
+
+ # Adding objects to trainer.
+ gc.collect()
+ return asr_brain, [train_data, valid_data, test_data]
+
+
+def pre_trained_point(save, config: DictConfig, server_device: str):
+ """Return a pre-trained model from a path and hyperparameters."""
+ state_dict = torch.load(config.pre_train_model_path)
+
+ overrides = {"output_folder": save}
+
+ hparams = config.sb_config
+ _, run_opts, _ = sb.parse_arguments(hparams)
+ with open(hparams) as fin:
+ params = load_hyperpyyaml(fin, overrides)
+
+ run_opts["device"] = server_device
+ run_opts["data_parallel_backend"] = config.parallel_backend
+ run_opts["noprogressbar"] = True # disable tqdm progress bar
+
+ asr_brain = ASR(
+ modules=params["modules"],
+ hparams=params,
+ run_opts=run_opts,
+ checkpointer=params["checkpointer"],
+ )
+
+ asr_brain.modules.load_state_dict(state_dict)
+ weights = get_weights(asr_brain.modules)
+ pre_trained = ndarrays_to_parameters(weights)
+
+ # Free up space after initialized
+ del asr_brain, weights
+ gc.collect()
+
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ return pre_trained
diff --git a/baselines/fedwav2vec2/fedwav2vec2/sb_recipe.py b/baselines/fedwav2vec2/fedwav2vec2/sb_recipe.py
new file mode 100644
index 000000000000..390edfe246d3
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/sb_recipe.py
@@ -0,0 +1,473 @@
+"""Main SpeechBrain training and testing logic."""
+
+
+import gc
+import time
+from collections import OrderedDict
+from enum import Enum, auto
+from typing import Dict, Optional
+
+import flwr as fl
+import numpy as np
+import speechbrain as sb
+import torch
+from speechbrain.dataio.dataloader import LoopedLoader
+from torch.utils.data import DataLoader
+from tqdm.contrib import tqdm
+
+# Recipe for training a sequence-to-sequence ASR system with CommonVoice.
+# The system employs a wav2vec2 encoder and a CTC decoder.
+# Decoding is performed with greedy decoding (will be extended to beam search).
+
+# To run this recipe, do the following:
+# > python train_with_wav2vec2.py hparams/train_with_wav2vec2.yaml
+
+# With the default hyperparameters, the system employs a pretrained wav2vec2 encoder.
+# The wav2vec2 model is pretrained following the model given in the hprams file.
+# It may be dependent on the language.
+
+# The neural network is trained with CTC on sub-word units estimated with
+# Byte Pairwise Encoding (BPE).
+
+# The experiment file is flexible enough to support a large variety of
+# different systems. By properly changing the parameter files, you can try
+# different encoders, decoders, tokens (e.g, characters instead of BPE),
+# training languages (all CommonVoice languages), and many
+# other possible variations.
+
+# Authors
+# * Titouan Parcollet 2021
+
+
+class Stage(Enum):
+ """Simple enum to track stage of experiments."""
+
+ TRAIN = auto()
+ VALID = auto()
+ TEST = auto()
+
+
+def set_weights(weights: fl.common.NDArrays, modules, device) -> None:
+ """Set model weights from a list of NumPy ndarrays."""
+ state_dict = OrderedDict()
+ valid_keys = modules.state_dict().keys()
+ for key, value in zip(valid_keys, weights):
+ weight = torch.Tensor(np.array(value))
+ weight = weight.to(device)
+ state_dict[key] = weight
+
+ modules.load_state_dict(state_dict, strict=True)
+
+
+def get_weights(modules) -> fl.common.NDArrays:
+ """Get model weights as a list of NumPy ndarrays."""
+ weights = []
+ for _, value in modules.state_dict().items():
+ weights.append(value.cpu().numpy())
+ return weights
+
+
+# pylint: disable=E1101,W0201,R0902
+class ASR(sb.core.Brain):
+ """Override of SpeechBrain default Brain class."""
+
+ def compute_forward(self, batch, _):
+ """Forward computations from the waveform batches to the output.
+
+ probabilities.
+ """
+ batch = batch.to(self.device)
+ wavs, wav_lens = batch.sig
+ # Forward pass
+ self.feats = self.modules.wav2vec2(wavs)
+
+ encoded_features = self.modules.enc(self.feats)
+ logits = self.modules.ctc_lin(encoded_features)
+ p_ctc = self.hparams.log_softmax(logits)
+
+ return p_ctc, wav_lens
+
+ def compute_objectives(self, predictions, batch, stage):
+ """Compute the CTC loss given predictions and targets."""
+ ids = batch.id
+ p_ctc, wav_lens = predictions
+ chars, char_lens = batch.char_encoded
+
+ loss = self.hparams.ctc_cost(p_ctc, chars, wav_lens, char_lens)
+ sequence = sb.decoders.ctc_greedy_decode(
+ p_ctc, wav_lens, self.hparams.blank_index
+ )
+ # ==============================Add by Salima=======================
+ # ==================================================================
+
+ if stage != sb.Stage.TRAIN:
+ self.cer_metric.append(
+ ids=ids,
+ predict=sequence,
+ target=chars,
+ target_len=char_lens,
+ ind2lab=self.label_encoder.decode_ndim,
+ )
+ self.coer_metric.append(
+ ids=ids,
+ predict=sequence,
+ target=chars,
+ target_len=char_lens,
+ ind2lab=self.label_encoder.decode_ndim,
+ )
+ self.cver_metric.append(
+ ids=ids,
+ predict=sequence,
+ target=chars,
+ target_len=char_lens,
+ ind2lab=self.label_encoder.decode_ndim,
+ )
+ self.ctc_metric.append(ids, p_ctc, chars, wav_lens, char_lens)
+
+ return loss
+
+ def init_optimizers(self):
+ """Initialize the wav2vec2 optimizer and model optimizer."""
+ self.wav2vec_optimizer = self.hparams.wav2vec_opt_class(
+ self.modules.wav2vec2.parameters()
+ )
+ self.adam_optimizer = self.hparams.adam_opt_class(
+ self.hparams.model.parameters()
+ )
+
+ def fit_batch(self, batch):
+ """Train the parameters given a single batch in input."""
+ batch = batch.to(self.device)
+ wavs, wav_lens = batch.sig
+
+ wavs, wav_lens = wavs.to(self.device), wav_lens.to(self.device)
+
+ stage = sb.Stage.TRAIN
+
+ predictions = self.compute_forward(batch, stage)
+ loss = self.compute_objectives(predictions, batch, stage)
+ loss.backward()
+ if self.check_gradients(loss):
+ self.wav2vec_optimizer.step()
+ self.adam_optimizer.step()
+
+ self.wav2vec_optimizer.zero_grad()
+ self.adam_optimizer.zero_grad()
+
+ return loss.detach().cpu()
+
+ def evaluate_batch(self, batch, stage):
+ """Compute validation/test batches."""
+ # Get data.
+ batch = batch.to(self.device)
+
+ predictions = self.compute_forward(batch, stage)
+ with torch.no_grad():
+ loss = self.compute_objectives(predictions, batch, stage=stage)
+ return loss.detach()
+
+ def on_stage_start(self, stage, epoch=None):
+ """Call when a stage (either training, validation, test) starts."""
+ _ = epoch
+ # self.ctc_metrics = self.hparams.ctc_stats()
+ if stage != sb.Stage.TRAIN:
+ self.cer_metric = self.hparams.cer_computer()
+ self.ctc_metric = self.hparams.ctc_computer()
+ self.coer_metric = self.hparams.coer_computer()
+ self.cver_metric = self.hparams.cver_computer()
+ # self.wer_metric = self.hparams.error_rate_computer()
+
+ def on_stage_end(self, stage, stage_loss, epoch=None):
+ """Call at the end of a stage."""
+ # Compute/store important stats
+ stage_stats = {"loss": stage_loss}
+
+ # if stage == sb.Stage.TRAIN:
+ # self.train_loss = stage_loss
+ if stage == sb.Stage.TRAIN:
+ self.train_loss = stage_loss
+ else:
+ # cer = self.cer_metrics.summarize("error_rate")
+ stage_stats["WER"] = self.cer_metric.summarize("error_rate")
+ stage_stats["COER"] = self.coer_metric.summarize("error_rate")
+ stage_stats["CVER"] = self.cver_metric.summarize("error_rate")
+
+ # Perform end-of-iteration things, like annealing, logging, etc.
+ if stage == sb.Stage.VALID:
+ old_lr_adam, new_lr_adam = self.hparams.lr_annealing_adam(
+ stage_stats["loss"]
+ )
+ old_lr_wav2vec, new_lr_wav2vec = self.hparams.lr_annealing_wav2vec(
+ stage_stats["loss"]
+ )
+ sb.nnet.schedulers.update_learning_rate(self.adam_optimizer, new_lr_adam)
+ sb.nnet.schedulers.update_learning_rate(
+ self.wav2vec_optimizer, new_lr_wav2vec
+ )
+
+ self.hparams.train_logger.log_stats(
+ stats_meta={
+ "epoch": epoch,
+ "lr_adam": old_lr_adam,
+ "lr_wav2vec": old_lr_wav2vec,
+ },
+ train_stats={"loss": self.train_loss},
+ valid_stats=stage_stats,
+ )
+
+ self.stage_wer = stage_stats["WER"]
+
+ elif stage == sb.Stage.TEST:
+ self.hparams.train_logger.log_stats(
+ stats_meta={"Epoch loaded": self.hparams.epoch_counter.current},
+ test_stats=stage_stats,
+ )
+ with open(self.hparams.wer_file, "w") as wer_file:
+ wer_file.write("CTC loss stats:\n")
+ self.ctc_metric.write_stats(wer_file)
+ wer_file.write("\nCER stats:\n")
+ self.cer_metric.write_stats(wer_file)
+ print("CTC and WER stats written to ", self.hparams.wer_file)
+
+ self.stage_wer = stage_stats["WER"]
+
+ def fit( # pylint: disable=W0102,R0912,R0913,R0914,R0915
+ self,
+ epoch_counter,
+ train_set,
+ valid_set=None,
+ progressbar=None,
+ train_loader_kwargs=Optional[Dict],
+ valid_loader_kwargs=Optional[Dict],
+ ):
+ """Iterate epochs and datasets to improve objective.
+
+ Relies on the existence of multiple functions that can (or should) be
+ overridden. The following methods are used and expected to have a
+ certain behavior:
+
+ * ``fit_batch()``
+ * ``evaluate_batch()``
+ * ``update_average()``
+
+ If the initialization was done with distributed_count > 0 and the
+ distributed_backend is ddp, this will generally handle multiprocess
+ logic, like splitting the training data into subsets for each device and
+ only saving a checkpoint on the main process.
+
+ Arguments
+ ---------
+ epoch_counter : iterable
+ Each call should return an integer indicating the epoch count.
+ train_set : Dataset, DataLoader
+ A set of data to use for training. If a Dataset is given, a
+ DataLoader is automatically created. If a DataLoader is given, it is
+ used directly.
+ valid_set : Dataset, DataLoader
+ A set of data to use for validation. If a Dataset is given, a
+ DataLoader is automatically created. If a DataLoader is given, it is
+ used directly.
+ train_loader_kwargs : Optional[Dict]
+ Kwargs passed to `make_dataloader()` for making the train_loader
+ (if train_set is a Dataset, not DataLoader).
+ E.G. batch_size, num_workers.
+ DataLoader kwargs are all valid.
+ valid_loader_kwargs : Optional[Dict]
+ Kwargs passed to `make_dataloader()` for making the valid_loader
+ (if valid_set is a Dataset, not DataLoader).
+ E.g., batch_size, num_workers.
+ DataLoader kwargs are all valid.
+ progressbar : bool
+ Whether to display the progress of each epoch in a progressbar.
+ """
+ if not isinstance(train_set, (DataLoader, LoopedLoader)):
+ train_set = self.make_dataloader(
+ train_set, stage=sb.Stage.TRAIN, **train_loader_kwargs
+ )
+ if valid_set is not None and not isinstance(
+ valid_set, (DataLoader, LoopedLoader)
+ ):
+ valid_set = self.make_dataloader(
+ valid_set,
+ stage=sb.Stage.VALID,
+ ckpt_prefix=None,
+ **valid_loader_kwargs,
+ )
+
+ self.on_fit_start()
+
+ if progressbar is None:
+ progressbar = not self.noprogressbar
+ self.modules = self.modules.to(self.device)
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ gc.collect()
+ # Iterate epochs
+ batch_count = 0
+ for epoch in epoch_counter:
+ # Training stage
+ self.on_stage_start(sb.Stage.TRAIN, epoch)
+ self.modules.train()
+
+ # Reset nonfinite count to 0 each epoch
+ self.nonfinite_count = 0
+
+ if self.train_sampler is not None and hasattr(
+ self.train_sampler, "set_epoch"
+ ):
+ self.train_sampler.set_epoch(epoch)
+
+ # Time since last intra-epoch checkpoint
+ last_ckpt_time = time.time()
+
+ # Only show progressbar if requested and main_process
+ enable = progressbar and sb.utils.distributed.if_main_process()
+ with tqdm(
+ train_set,
+ initial=self.step,
+ dynamic_ncols=True,
+ disable=not enable,
+ ) as progress_bar:
+ for batch in progress_bar:
+ self.step += 1
+ loss = self.fit_batch(batch)
+ _, wav_lens = batch.sig
+ batch_count += wav_lens.shape[0]
+ self.avg_train_loss = self.update_average(loss, self.avg_train_loss)
+ progress_bar.set_postfix(train_loss=self.avg_train_loss)
+
+ # Debug mode only runs a few batches
+ if self.debug and self.step == self.debug_batches:
+ break
+
+ if (
+ self.checkpointer is not None
+ and self.ckpt_interval_minutes > 0
+ and time.time() - last_ckpt_time
+ >= self.ckpt_interval_minutes * 60.0
+ ):
+ # This should not use run_on_main, because that
+ # includes a DDP barrier. That eventually leads to a
+ # crash when the processes'
+ # time.time() - last_ckpt_time differ and some
+ # processes enter this block while others don't,
+ # missing the barrier.
+ if sb.utils.distributed.if_main_process():
+ self._save_intra_epoch_ckpt()
+ last_ckpt_time = time.time()
+
+ if epoch == epoch_counter.limit:
+ avg_loss = self.avg_train_loss
+ # Run train "on_stage_end" on all processes
+ self.on_stage_end(sb.Stage.TRAIN, self.avg_train_loss, epoch)
+ self.avg_train_loss = 0.0
+ self.step = 0
+
+ # Validation stage
+ if valid_set is not None:
+ self.on_stage_start(sb.Stage.VALID, epoch)
+ self.modules.eval()
+ avg_valid_loss = 0.0
+ with torch.no_grad():
+ for batch in tqdm(
+ valid_set, dynamic_ncols=True, disable=not enable
+ ):
+ self.step += 1
+ loss = self.evaluate_batch(batch, stage=sb.Stage.VALID)
+ avg_valid_loss = self.update_average(loss, avg_valid_loss)
+
+ # Debug mode only runs a few batches
+ if self.debug and self.step == self.debug_batches:
+ break
+
+ # Only run validation "on_stage_end" on main process
+ self.step = 0
+ self.on_stage_end(sb.Stage.VALID, avg_valid_loss, epoch)
+ valid_wer = self.stage_wer
+ if epoch == epoch_counter.limit:
+ valid_wer_last = valid_wer
+
+ # Debug mode only runs a few epochs
+ if self.debug and epoch == self.debug_epochs:
+ break
+ if self.device == "cpu":
+ self.modules = self.modules.to("cpu")
+ gc.collect()
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ return batch_count, avg_loss, valid_wer_last
+
+ def evaluate( # pylint: disable=W0102,R0913
+ self,
+ test_set,
+ max_key=None,
+ min_key=None,
+ progressbar=None,
+ test_loader_kwargs=Optional[Dict],
+ ):
+ """Iterate test_set and evaluate brain performance. By default, loads the best-.
+
+ performing checkpoint (as recorded using the checkpointer).
+
+ Arguments
+ ---------
+ test_set : Dataset, DataLoader
+ If a DataLoader is given, it is iterated directly. Otherwise passed
+ to ``self.make_dataloader()``.
+ max_key : str
+ Key to use for finding best checkpoint, passed to
+ ``on_evaluate_start()``.
+ min_key : str
+ Key to use for finding best checkpoint, passed to
+ ``on_evaluate_start()``.
+ progressbar : bool
+ Whether to display the progress in a progressbar.
+ test_loader_kwargs : Optional[Dict]
+ Kwargs passed to ``make_dataloader()`` if ``test_set`` is not a
+ DataLoader. NOTE: ``loader_kwargs["ckpt_prefix"]`` gets
+ automatically overwritten to ``None`` (so that the test DataLoader
+ is not added to the checkpointer).
+
+ Returns
+ -------
+ average test loss
+ """
+ if progressbar is None:
+ progressbar = not self.noprogressbar
+
+ if not isinstance(test_set, (DataLoader, LoopedLoader)):
+ test_loader_kwargs["ckpt_prefix"] = None
+ test_set = self.make_dataloader(
+ test_set, sb.Stage.TEST, **test_loader_kwargs
+ )
+ self.modules = self.modules.to(self.device)
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ gc.collect()
+ self.on_evaluate_start(max_key=max_key, min_key=min_key)
+ self.on_stage_start(sb.Stage.TEST, None)
+ self.modules.eval()
+ avg_test_loss = 0.0
+ batch_count = 0
+ with torch.no_grad():
+ for batch in tqdm(test_set, dynamic_ncols=True, disable=not progressbar):
+ self.step += 1
+ _, wav_lens = batch.sig
+ batch_count += wav_lens.shape[0]
+ loss = self.evaluate_batch(batch, stage=sb.Stage.TEST)
+ avg_test_loss = self.update_average(loss, avg_test_loss)
+
+ # Debug mode only runs a few batches
+ if self.debug and self.step == self.debug_batches:
+ break
+
+ self.on_stage_end(sb.Stage.TEST, avg_test_loss, None)
+ cer = self.stage_wer
+ self.step = 0
+ if self.device == "cpu":
+ self.modules = self.modules.to("cpu")
+ gc.collect()
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+
+ return batch_count, avg_test_loss, cer
diff --git a/baselines/fedwav2vec2/fedwav2vec2/server.py b/baselines/fedwav2vec2/fedwav2vec2/server.py
new file mode 100644
index 000000000000..db42c0f069a5
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/server.py
@@ -0,0 +1,69 @@
+"""Create global evaluation function.
+
+Optionally, also define a new Server class (please note this is not needed in most
+settings).
+"""
+
+import gc
+import os
+from typing import Callable, Dict
+
+import flwr as fl
+import torch
+from flwr.common import Scalar
+from omegaconf import DictConfig
+
+from fedwav2vec2.client import SpeechBrainClient
+from fedwav2vec2.models import int_model
+
+
+def get_on_fit_config_fn(local_epochs: int) -> Callable[[int], Dict[str, str]]:
+ """Return a function which returns training configurations."""
+
+ def fit_config(rnd: int) -> Dict[str, str]:
+ """Return a configuration with static batch size and (local) epochs."""
+ config = {"epoch_global": str(rnd), "epochs": str(local_epochs)}
+ return config
+
+ return fit_config
+
+
+def get_evaluate_fn(config: DictConfig, server_device: str, save_path: str):
+ """Return function to execute during global evaluation."""
+ config_ = config
+
+ def evaluate_fn(
+ server_round: int, weights: fl.common.NDArrays, config: Dict[str, Scalar]
+ ):
+ """Run centralized evaluation."""
+ _ = (server_round, config)
+ # int model
+ asr_brain, dataset = int_model(
+ config_.server_cid,
+ config_,
+ server_device,
+ save_path,
+ evaluate=True,
+ )
+
+ client = SpeechBrainClient(config_.server_cid, asr_brain, dataset)
+
+ _, lss, err = client.evaluate_train_speech_recogniser(
+ server_params=weights,
+ epochs=1,
+ )
+ # Save model if indicated
+ if config_.save_checkpoint is not None:
+ if not os.path.exists(config_.save_checkpoint):
+ os.mkdir(config_.save_checkpoint)
+ checkpoint = os.path.join(config_.save_checkpoint, "last_checkpoint.pt")
+ torch.save(asr_brain.modules.state_dict(), checkpoint)
+ print(f"Checkpoint saved for round {server_round}")
+
+ del client, asr_brain, dataset
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ gc.collect()
+ return lss, {"Error rate": err}
+
+ return evaluate_fn
diff --git a/baselines/fedwav2vec2/fedwav2vec2/strategy.py b/baselines/fedwav2vec2/fedwav2vec2/strategy.py
new file mode 100644
index 000000000000..6aa6007e4ccd
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/strategy.py
@@ -0,0 +1,69 @@
+"""Optionally define a custom strategy.
+
+Needed only when the strategy is not yet implemented in Flower or because you want to
+extend or modify the functionality of an existing strategy.
+"""
+
+
+import gc
+from typing import Dict, List, Optional, Tuple, Union
+
+import flwr as fl
+import torch
+from flwr.common import (
+ FitRes,
+ Parameters,
+ ndarrays_to_parameters,
+ parameters_to_ndarrays,
+)
+from flwr.server.client_proxy import ClientProxy
+from flwr.server.strategy.aggregate import aggregate
+
+
+class CustomFedAvg(fl.server.strategy.FedAvg):
+ """Custom strategy to aggregate using metrics instead of number of samples."""
+
+ def __init__(self, *args, weight_strategy, **kwargs) -> None:
+ super().__init__(*args, **kwargs)
+ self.weight_strategy = weight_strategy
+
+ def aggregate_fit(
+ self,
+ _: int,
+ results: List[Tuple[ClientProxy, FitRes]],
+ failures: List[Union[Tuple[ClientProxy, FitRes], BaseException]],
+ ) -> Tuple[Optional[Parameters], Dict[str, Union[bool, bytes, float, int, str]]]:
+ """Aggregate results using different weighing metrics (train_loss or WER)."""
+ if not results:
+ return None, {}
+ # Do not aggregate if there are failures and failures are not accepted
+ if not self.accept_failures and failures:
+ return None, {}
+
+ # Convert results
+ key_name = "train_loss" if self.weight_strategy == "loss" else "wer"
+ weights = None
+
+ # Define ratio merge
+ if self.weight_strategy == "num":
+ weights_results = [
+ (parameters_to_ndarrays(fit_res.parameters), fit_res.num_examples)
+ for _, fit_res in results
+ ]
+ weights = aggregate(weights_results)
+ else:
+ weights_results = [
+ (
+ parameters_to_ndarrays(fit_res.parameters),
+ int(fit_res.metrics[key_name]),
+ )
+ for _, fit_res in results
+ ]
+ weights = aggregate(weights_results)
+
+ # Free memory for next round
+ del results, weights_results
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+ gc.collect()
+ return ndarrays_to_parameters(weights), {}
diff --git a/baselines/fedwav2vec2/fedwav2vec2/utils.py b/baselines/fedwav2vec2/fedwav2vec2/utils.py
new file mode 100644
index 000000000000..9a831719d623
--- /dev/null
+++ b/baselines/fedwav2vec2/fedwav2vec2/utils.py
@@ -0,0 +1,6 @@
+"""Define any utility function.
+
+They are not directly relevant to the other (more FL specific) python modules. For
+example, you may define here things like: loading a model from a checkpoint, saving
+results, plotting.
+"""
diff --git a/baselines/fedwav2vec2/pyproject.toml b/baselines/fedwav2vec2/pyproject.toml
new file mode 100644
index 000000000000..1e7dbf55154b
--- /dev/null
+++ b/baselines/fedwav2vec2/pyproject.toml
@@ -0,0 +1,142 @@
+[build-system]
+requires = ["poetry-core>=1.4.0"]
+build-backend = "poetry.masonry.api"
+
+[tool.poetry]
+name = "fedwav2vec2" # <----- Ensure it matches the name of your baseline directory containing all the source code
+version = "1.0.0"
+description = "Federated Learning for ASR Based on wav2vec 2.0"
+license = "Apache-2.0"
+authors = ["The Flower Authors ", "Tuan Nguyen "]
+readme = "README.md"
+homepage = "https://flower.dev"
+repository = "https://github.com/adap/flower"
+documentation = "https://flower.dev"
+classifiers = [
+ "Development Status :: 3 - Alpha",
+ "Intended Audience :: Developers",
+ "Intended Audience :: Science/Research",
+ "License :: OSI Approved :: Apache Software License",
+ "Operating System :: MacOS :: MacOS X",
+ "Operating System :: POSIX :: Linux",
+ "Programming Language :: Python",
+ "Programming Language :: Python :: 3",
+ "Programming Language :: Python :: 3 :: Only",
+ "Programming Language :: Python :: 3.8",
+ "Programming Language :: Python :: 3.9",
+ "Programming Language :: Python :: 3.10",
+ "Programming Language :: Python :: 3.11",
+ "Programming Language :: Python :: Implementation :: CPython",
+ "Topic :: Scientific/Engineering",
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
+ "Topic :: Scientific/Engineering :: Mathematics",
+ "Topic :: Software Development",
+ "Topic :: Software Development :: Libraries",
+ "Topic :: Software Development :: Libraries :: Python Modules",
+ "Typing :: Typed",
+]
+
+[tool.poetry.dependencies]
+python = ">=3.10.0, <3.12.0" # don't change this
+flwr = { extras = ["simulation"], version = "1.5.0" }
+hydra-core = "1.3.2" # don't change this
+speechbrain = "0.5.15"
+pandas = "2.1.1"
+torch = { url = "https://download.pytorch.org/whl/cu116/torch-1.13.1%2Bcu116-cp310-cp310-linux_x86_64.whl"}
+torchaudio = { url = "https://download.pytorch.org/whl/cu116/torchaudio-0.13.1%2Bcu116-cp310-cp310-linux_x86_64.whl"}
+transformers = "4.33.2"
+
+[tool.poetry.dev-dependencies]
+isort = "==5.11.5"
+black = "==23.1.0"
+docformatter = "==1.5.1"
+mypy = "==1.4.1"
+pylint = "==2.8.2"
+flake8 = "==3.9.2"
+pytest = "==6.2.4"
+pytest-watch = "==4.2.0"
+ruff = "==0.0.272"
+types-requests = "==2.27.7"
+
+[tool.isort]
+line_length = 88
+indent = " "
+multi_line_output = 3
+include_trailing_comma = true
+force_grid_wrap = 0
+use_parentheses = true
+
+[tool.black]
+line-length = 88
+target-version = ["py38", "py39", "py310", "py311"]
+
+[tool.pytest.ini_options]
+minversion = "6.2"
+addopts = "-qq"
+testpaths = [
+ "flwr_baselines",
+]
+
+[tool.mypy]
+ignore_missing_imports = true
+strict = false
+plugins = "numpy.typing.mypy_plugin"
+
+[tool.pylint."MESSAGES CONTROL"]
+disable = "bad-continuation,duplicate-code,too-few-public-methods,useless-import-alias"
+good-names = "i,j,k,_,x,y,X,Y"
+signature-mutators="hydra.main.main"
+
+[tool.pylint.typecheck]
+generated-members="numpy.*, torch.*, tensorflow.*"
+
+[[tool.mypy.overrides]]
+module = [
+ "importlib.metadata.*",
+ "importlib_metadata.*",
+]
+follow_imports = "skip"
+follow_imports_for_stubs = true
+disallow_untyped_calls = false
+
+[[tool.mypy.overrides]]
+module = "torch.*"
+follow_imports = "skip"
+follow_imports_for_stubs = true
+
+[tool.docformatter]
+wrap-summaries = 88
+wrap-descriptions = 88
+
+[tool.ruff]
+target-version = "py38"
+line-length = 88
+select = ["D", "E", "F", "W", "B", "ISC", "C4"]
+fixable = ["D", "E", "F", "W", "B", "ISC", "C4"]
+ignore = ["B024", "B027"]
+exclude = [
+ ".bzr",
+ ".direnv",
+ ".eggs",
+ ".git",
+ ".hg",
+ ".mypy_cache",
+ ".nox",
+ ".pants.d",
+ ".pytype",
+ ".ruff_cache",
+ ".svn",
+ ".tox",
+ ".venv",
+ "__pypackages__",
+ "_build",
+ "buck-out",
+ "build",
+ "dist",
+ "node_modules",
+ "venv",
+ "proto",
+]
+
+[tool.ruff.pydocstyle]
+convention = "numpy"
diff --git a/baselines/fjord/.gitignore b/baselines/fjord/.gitignore
new file mode 100644
index 000000000000..8199f9d1a17f
--- /dev/null
+++ b/baselines/fjord/.gitignore
@@ -0,0 +1,3 @@
+data/
+runs/
+exp_logs/
diff --git a/baselines/fjord/LICENSE b/baselines/fjord/LICENSE
new file mode 100644
index 000000000000..d64569567334
--- /dev/null
+++ b/baselines/fjord/LICENSE
@@ -0,0 +1,202 @@
+
+ Apache License
+ Version 2.0, January 2004
+ http://www.apache.org/licenses/
+
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
+
+ 1. Definitions.
+
+ "License" shall mean the terms and conditions for use, reproduction,
+ and distribution as defined by Sections 1 through 9 of this document.
+
+ "Licensor" shall mean the copyright owner or entity authorized by
+ the copyright owner that is granting the License.
+
+ "Legal Entity" shall mean the union of the acting entity and all
+ other entities that control, are controlled by, or are under common
+ control with that entity. For the purposes of this definition,
+ "control" means (i) the power, direct or indirect, to cause the
+ direction or management of such entity, whether by contract or
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
+ outstanding shares, or (iii) beneficial ownership of such entity.
+
+ "You" (or "Your") shall mean an individual or Legal Entity
+ exercising permissions granted by this License.
+
+ "Source" form shall mean the preferred form for making modifications,
+ including but not limited to software source code, documentation
+ source, and configuration files.
+
+ "Object" form shall mean any form resulting from mechanical
+ transformation or translation of a Source form, including but
+ not limited to compiled object code, generated documentation,
+ and conversions to other media types.
+
+ "Work" shall mean the work of authorship, whether in Source or
+ Object form, made available under the License, as indicated by a
+ copyright notice that is included in or attached to the work
+ (an example is provided in the Appendix below).
+
+ "Derivative Works" shall mean any work, whether in Source or Object
+ form, that is based on (or derived from) the Work and for which the
+ editorial revisions, annotations, elaborations, or other modifications
+ represent, as a whole, an original work of authorship. For the purposes
+ of this License, Derivative Works shall not include works that remain
+ separable from, or merely link (or bind by name) to the interfaces of,
+ the Work and Derivative Works thereof.
+
+ "Contribution" shall mean any work of authorship, including
+ the original version of the Work and any modifications or additions
+ to that Work or Derivative Works thereof, that is intentionally
+ submitted to Licensor for inclusion in the Work by the copyright owner
+ or by an individual or Legal Entity authorized to submit on behalf of
+ the copyright owner. For the purposes of this definition, "submitted"
+ means any form of electronic, verbal, or written communication sent
+ to the Licensor or its representatives, including but not limited to
+ communication on electronic mailing lists, source code control systems,
+ and issue tracking systems that are managed by, or on behalf of, the
+ Licensor for the purpose of discussing and improving the Work, but
+ excluding communication that is conspicuously marked or otherwise
+ designated in writing by the copyright owner as "Not a Contribution."
+
+ "Contributor" shall mean Licensor and any individual or Legal Entity
+ on behalf of whom a Contribution has been received by Licensor and
+ subsequently incorporated within the Work.
+
+ 2. Grant of Copyright License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ copyright license to reproduce, prepare Derivative Works of,
+ publicly display, publicly perform, sublicense, and distribute the
+ Work and such Derivative Works in Source or Object form.
+
+ 3. Grant of Patent License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ (except as stated in this section) patent license to make, have made,
+ use, offer to sell, sell, import, and otherwise transfer the Work,
+ where such license applies only to those patent claims licensable
+ by such Contributor that are necessarily infringed by their
+ Contribution(s) alone or by combination of their Contribution(s)
+ with the Work to which such Contribution(s) was submitted. If You
+ institute patent litigation against any entity (including a
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
+ or a Contribution incorporated within the Work constitutes direct
+ or contributory patent infringement, then any patent licenses
+ granted to You under this License for that Work shall terminate
+ as of the date such litigation is filed.
+
+ 4. Redistribution. You may reproduce and distribute copies of the
+ Work or Derivative Works thereof in any medium, with or without
+ modifications, and in Source or Object form, provided that You
+ meet the following conditions:
+
+ (a) You must give any other recipients of the Work or
+ Derivative Works a copy of this License; and
+
+ (b) You must cause any modified files to carry prominent notices
+ stating that You changed the files; and
+
+ (c) You must retain, in the Source form of any Derivative Works
+ that You distribute, all copyright, patent, trademark, and
+ attribution notices from the Source form of the Work,
+ excluding those notices that do not pertain to any part of
+ the Derivative Works; and
+
+ (d) If the Work includes a "NOTICE" text file as part of its
+ distribution, then any Derivative Works that You distribute must
+ include a readable copy of the attribution notices contained
+ within such NOTICE file, excluding those notices that do not
+ pertain to any part of the Derivative Works, in at least one
+ of the following places: within a NOTICE text file distributed
+ as part of the Derivative Works; within the Source form or
+ documentation, if provided along with the Derivative Works; or,
+ within a display generated by the Derivative Works, if and
+ wherever such third-party notices normally appear. The contents
+ of the NOTICE file are for informational purposes only and
+ do not modify the License. You may add Your own attribution
+ notices within Derivative Works that You distribute, alongside
+ or as an addendum to the NOTICE text from the Work, provided
+ that such additional attribution notices cannot be construed
+ as modifying the License.
+
+ You may add Your own copyright statement to Your modifications and
+ may provide additional or different license terms and conditions
+ for use, reproduction, or distribution of Your modifications, or
+ for any such Derivative Works as a whole, provided Your use,
+ reproduction, and distribution of the Work otherwise complies with
+ the conditions stated in this License.
+
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
+ any Contribution intentionally submitted for inclusion in the Work
+ by You to the Licensor shall be under the terms and conditions of
+ this License, without any additional terms or conditions.
+ Notwithstanding the above, nothing herein shall supersede or modify
+ the terms of any separate license agreement you may have executed
+ with Licensor regarding such Contributions.
+
+ 6. Trademarks. This License does not grant permission to use the trade
+ names, trademarks, service marks, or product names of the Licensor,
+ except as required for reasonable and customary use in describing the
+ origin of the Work and reproducing the content of the NOTICE file.
+
+ 7. Disclaimer of Warranty. Unless required by applicable law or
+ agreed to in writing, Licensor provides the Work (and each
+ Contributor provides its Contributions) on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
+ implied, including, without limitation, any warranties or conditions
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
+ PARTICULAR PURPOSE. You are solely responsible for determining the
+ appropriateness of using or redistributing the Work and assume any
+ risks associated with Your exercise of permissions under this License.
+
+ 8. Limitation of Liability. In no event and under no legal theory,
+ whether in tort (including negligence), contract, or otherwise,
+ unless required by applicable law (such as deliberate and grossly
+ negligent acts) or agreed to in writing, shall any Contributor be
+ liable to You for damages, including any direct, indirect, special,
+ incidental, or consequential damages of any character arising as a
+ result of this License or out of the use or inability to use the
+ Work (including but not limited to damages for loss of goodwill,
+ work stoppage, computer failure or malfunction, or any and all
+ other commercial damages or losses), even if such Contributor
+ has been advised of the possibility of such damages.
+
+ 9. Accepting Warranty or Additional Liability. While redistributing
+ the Work or Derivative Works thereof, You may choose to offer,
+ and charge a fee for, acceptance of support, warranty, indemnity,
+ or other liability obligations and/or rights consistent with this
+ License. However, in accepting such obligations, You may act only
+ on Your own behalf and on Your sole responsibility, not on behalf
+ of any other Contributor, and only if You agree to indemnify,
+ defend, and hold each Contributor harmless for any liability
+ incurred by, or claims asserted against, such Contributor by reason
+ of your accepting any such warranty or additional liability.
+
+ END OF TERMS AND CONDITIONS
+
+ APPENDIX: How to apply the Apache License to your work.
+
+ To apply the Apache License to your work, attach the following
+ boilerplate notice, with the fields enclosed by brackets "[]"
+ replaced with your own identifying information. (Don't include
+ the brackets!) The text should be enclosed in the appropriate
+ comment syntax for the file format. We also recommend that a
+ file or class name and description of purpose be included on the
+ same "printed page" as the copyright notice for easier
+ identification within third-party archives.
+
+ Copyright [yyyy] [name of copyright owner]
+
+ Licensed under the Apache License, Version 2.0 (the "License");
+ you may not use this file except in compliance with the License.
+ You may obtain a copy of the License at
+
+ http://www.apache.org/licenses/LICENSE-2.0
+
+ Unless required by applicable law or agreed to in writing, software
+ distributed under the License is distributed on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ See the License for the specific language governing permissions and
+ limitations under the License.
diff --git a/baselines/fjord/README.md b/baselines/fjord/README.md
new file mode 100644
index 000000000000..563f583082a8
--- /dev/null
+++ b/baselines/fjord/README.md
@@ -0,0 +1,118 @@
+---
+title: "FjORD: Fair and Accurate Federated Learning under heterogeneous targets with Ordered Dropout"
+url: "https://openreview.net/forum?id=4fLr7H5D_eT"
+labels: ["Federated Learning", "Heterogeneity", "Efficient DNNs", "Distributed Systems"]
+dataset: ["CIFAR-10"]
+---
+
+# FjORD: Fair and Accurate Federated Learning under heterogeneous targets with Ordered Dropout
+
+**Paper:** [openreview.net/forum?id=4fLr7H5D_eT](https://openreview.net/forum?id=4fLr7H5D_eT)
+
+**Authors:** Samuel Horváth\*, Stefanos Laskaridis\*, Mario Almeida\*, Ilias Leontiadis, Stylianos Venieris, Nicholas Donald Lane
+
+
+**Abstract:** Federated Learning (FL) has been gaining significant traction across different ML tasks, ranging from vision to keyboard predictions. In large-scale deployments, client heterogeneity is a fact and constitutes a primary problem for fairness, training performance and accuracy. Although significant efforts have been made into tackling statistical data heterogeneity, the diversity in the processing capabilities and network bandwidth of clients, termed system heterogeneity, has remained largely unexplored. Current solutions either disregard a large portion of available devices or set a uniform limit on the model's capacity, restricted by the least capable participants.
+
+In this work, we introduce Ordered Dropout, a mechanism that achieves an ordered, nested representation of knowledge in Neural Networks and enables the extraction of lower footprint submodels without the need for retraining. We further show that for linear maps our Ordered Dropout is equivalent to SVD. We employ this technique, along with a self-distillation methodology, in the realm of FL in a framework called FjORD. FjORD alleviates the problem of client system heterogeneity by tailoring the model width to the client's capabilities.
+Extensive evaluation on both CNNs and RNNs across diverse modalities shows that FjORD consistently leads to significant performance gains over state-of-the-art baselines while maintaining its nested structure.
+
+
+## About this baseline
+
+**What’s implemented:** The code in this directory implements the two variants of FjORD, with and without knowledge distillation.
+
+**Datasets:** CIFAR-10
+
+**Hardware Setup:** We trained the baseline on an Nvidia RTX 4090.
+
+**Contributors:** @stevelaskaridis ([Brave Software](https://brave.com/)), @SamuelHorvath ([MBZUAI](https://mbzuai.ac.ae/))
+
+
+## Experimental Setup
+
+**Task:** Image Classification
+
+**Model:** ResNet-18
+
+**Dataset:**
+
+| **Feature** | **Value** |
+| -------------------------- | ---------------------------- |
+| **Dataset** | CIFAR-10 |
+| **Partition** | Randomised Sequential Split |
+| **Number of Partitions** | 100 clients |
+| **Data points per client** | 500 samples |
+
+**Training Hyperparameters:**
+
+| **Hyperparameter** | **Value** |
+| ----------------------- | ------------------------- |
+| batch size | 32 |
+| learning rate | 0.1 |
+| learning rate scheduler | static |
+| optimiser | sgd |
+| momentum | 0 |
+| nesterov | False |
+| weight decay | 1e-4 |
+| sample per round | 10 |
+| local epochs | 1 |
+| p-values | [0.2, 0.4, 0.6, 0.8, 1.0] |
+| client tier allocation | uniform |
+
+
+## Environment Setup
+
+### Through regular pip
+
+```bash
+pip install -r requirements.txt
+python setup.py install
+```
+
+### Through poetry
+
+```bash
+# Set python version
+pyenv install 3.10.6
+pyenv local 3.10.6
+
+# Tell poetry to use python 3.10
+poetry env use 3.10.6
+
+# install the base Poetry environment
+poetry install
+
+# activate the environment
+poetry shell
+```
+
+## Running the Experiments
+
+### Through your environment
+
+
+```bash
+python -m fjord.main # without knowledge distillation
+# or
+python -m fjord.main +train_mode=fjord_kd # with knowledge distillation
+```
+
+### Through poetry
+
+```bash
+poetry run python -m fjord.main # without knowledge distillation
+# or
+poetry run python -m fjord.main +train_mode=fjord_kd # with knowledge distillation
+```
+
+## Expected Results
+
+```bash
+cd scripts/
+./run.sh
+```
+
+Plots and the associated code reside in `fjord/notebooks/visualise.ipynb`.
+
+![resnet18_cifar10_500_global_rounds_acc_pvalues](./_static/resnet18_cifar10_500_global_rounds_acc_pvalues.png)
\ No newline at end of file
diff --git a/baselines/fjord/_static/resnet18_cifar10_500_global_rounds_acc_pvalues.png b/baselines/fjord/_static/resnet18_cifar10_500_global_rounds_acc_pvalues.png
new file mode 100644
index 000000000000..de3ad61a5d55
Binary files /dev/null and b/baselines/fjord/_static/resnet18_cifar10_500_global_rounds_acc_pvalues.png differ
diff --git a/baselines/fjord/_static/resnet18_cifar10_fjord_convergence.png b/baselines/fjord/_static/resnet18_cifar10_fjord_convergence.png
new file mode 100644
index 000000000000..12b137e3d196
Binary files /dev/null and b/baselines/fjord/_static/resnet18_cifar10_fjord_convergence.png differ
diff --git a/baselines/fjord/_static/resnet18_cifar10_fjord_kd_convergence.png b/baselines/fjord/_static/resnet18_cifar10_fjord_kd_convergence.png
new file mode 100644
index 000000000000..358a5d19a281
Binary files /dev/null and b/baselines/fjord/_static/resnet18_cifar10_fjord_kd_convergence.png differ
diff --git a/baselines/fjord/fjord/__init__.py b/baselines/fjord/fjord/__init__.py
new file mode 100644
index 000000000000..7aa11d2a7b9f
--- /dev/null
+++ b/baselines/fjord/fjord/__init__.py
@@ -0,0 +1 @@
+"""FjORD package."""
diff --git a/baselines/fjord/fjord/client.py b/baselines/fjord/fjord/client.py
new file mode 100644
index 000000000000..2b18d9547086
--- /dev/null
+++ b/baselines/fjord/fjord/client.py
@@ -0,0 +1,240 @@
+"""Flower client implementing FjORD."""
+from collections import OrderedDict
+from copy import deepcopy
+from types import SimpleNamespace
+from typing import Any, Dict, List, Tuple, Union
+
+import flwr as fl
+import numpy as np
+import torch
+from torch import Tensor
+from torch.nn import Module
+from torch.utils.data import DataLoader
+
+from .dataset import load_data
+from .models import get_net, test, train
+from .od.layers import ODBatchNorm2d, ODConv2d, ODLinear
+from .od.samplers import ODSampler
+from .utils.logger import Logger
+from .utils.utils import save_model
+
+FJORD_CONFIG_TYPE = Dict[
+ Union[str, float],
+ List[Any],
+]
+
+
+def get_layer_from_state_dict(model: Module, state_dict_key: str) -> Module:
+ """Get the layer corresponding to the given state dict key.
+
+ :param model: The model.
+ :param state_dict_key: The state dict key.
+ :return: The module corresponding to the given state dict key.
+ """
+ keys = state_dict_key.split(".")
+ module = model
+ # The last keyc orresponds to the parameter name
+ # (e.g., weight or bias)
+ for key in keys[:-1]:
+ module = getattr(module, key)
+ return module
+
+
+def net_to_state_dict_layers(net: Module) -> List[Module]:
+ """Get the state_dict of the model.
+
+ :param net: The model.
+ :return: The state_dict of the model.
+ """
+ layers = []
+ for key, _ in net.state_dict().items():
+ layer = get_layer_from_state_dict(net, key)
+ layers.append(layer)
+ return layers
+
+
+def get_agg_config(
+ net: Module, trainloader: DataLoader, p_s: List[float]
+) -> FJORD_CONFIG_TYPE:
+ """Get the aggregation configuration of the model.
+
+ :param net: The model.
+ :param trainloader: The training set.
+ :param p_s: The p values used
+ :return: The aggregation configuration of the model.
+ """
+ Logger.get().info("Constructing OD model configuration for aggregation.")
+ device = next(net.parameters()).device
+ images, _ = next(iter(trainloader))
+ images = images.to(device)
+ layers = net_to_state_dict_layers(net)
+ # init min dims in networks
+ config: FJORD_CONFIG_TYPE = {p: [{} for _ in layers] for p in p_s}
+ config["layer"] = []
+ config["layer_p"] = []
+ with torch.no_grad():
+ for p in p_s:
+ max_sampler = ODSampler(
+ p_s=[p],
+ max_p=p,
+ model=net,
+ )
+ net(images, sampler=max_sampler)
+ for i, layer in enumerate(layers):
+ if isinstance(layer, (ODConv2d, ODLinear)):
+ config[p][i]["in_dim"] = layer.last_input_dim
+ config[p][i]["out_dim"] = layer.last_output_dim
+ elif isinstance(layer, ODBatchNorm2d):
+ config[p][i]["in_dim"] = None
+ config[p][i]["out_dim"] = layer.p_to_num_features[p]
+ elif isinstance(layer, torch.nn.BatchNorm2d):
+ pass
+ else:
+ raise ValueError(f"Unsupported layer {layer.__class__.__name__}")
+ for layer in layers:
+ config["layer"].append(layer.__class__.__name__)
+ if hasattr(layer, "p"):
+ config["layer_p"].append(layer.p)
+ else:
+ config["layer_p"].append(None)
+ return config
+
+
+# Define Flower client
+class FjORDClient(
+ fl.client.NumPyClient
+): # pylint: disable=too-many-instance-attributes
+ """Flower client training on CIFAR-10."""
+
+ def __init__( # pylint: disable=too-many-arguments
+ self,
+ cid: int,
+ model_name: str,
+ model_path: str,
+ data_path: str,
+ know_distill: bool,
+ max_p: float,
+ p_s: List[float],
+ train_config: SimpleNamespace,
+ fjord_config: FJORD_CONFIG_TYPE,
+ log_config: Dict[str, str],
+ device: torch.device,
+ seed: int,
+ ) -> None:
+ """Initialise the client.
+
+ :param cid: The client ID.
+ :param model_name: The model name.
+ :param model_path: The path to save the model.
+ :param data_path: The path to the dataset.
+ :param know_distill: Whether the model uses knowledge distillation.
+ :param max_p: The maximum p value.
+ :param p_s: The p values to use for training.
+ :param train_config: The training configuration.
+ :param fjord_config: The configuration for Fjord.
+ :param log_config: The logging configuration.
+ :param device: The device to use.
+ :param seed: The seed to use for the random number generator.
+ """
+ Logger.setup_logging(**log_config)
+ self.cid = cid
+ self.p_s = p_s
+ self.net = get_net(model_name, p_s, device)
+ self.trainloader, self.valloader = load_data(
+ data_path, int(cid), train_config.batch_size, seed
+ )
+
+ self.know_distill = know_distill
+ self.max_p = max_p
+ self.fjord_config = fjord_config
+ self.train_config = train_config
+ self.model_path = model_path
+
+ def get_parameters(self, config: Dict[str, fl.common.Scalar]) -> List[np.ndarray]:
+ """Get the parameters of the model to return to the server.
+
+ :param config: The configuration.
+ :return: The parameters of the model.
+ """
+ Logger.get().info(f"Getting parameters from client {self.cid}")
+ return [val.cpu().numpy() for _, val in self.net.state_dict().items()]
+
+ def net_to_state_dict_layers(self) -> List[Module]:
+ """Model to state dict layers."""
+ return net_to_state_dict_layers(self.net)
+
+ def set_parameters(self, parameters: List[np.ndarray]) -> None:
+ """Set the parameters of the model.
+
+ :param parameters: The parameters of the model.
+ """
+ params_dict = zip(self.net.state_dict().keys(), parameters)
+ state_dict = OrderedDict({k: torch.tensor(v) for k, v in params_dict})
+ self.net.load_state_dict(state_dict, strict=True)
+
+ def fit(
+ self, parameters: List[Tensor], config: Dict[str, fl.common.Scalar]
+ ) -> Tuple[List[np.ndarray], int, Dict[str, Any]]:
+ """Train the model on the training set.
+
+ :param parameters: The parameters of the model.
+ :param config: The train configuration.
+ :return: The parameters of the model, the number of samples used for training,
+ and the training metrics
+ """
+ Logger.get().info(
+ f"Training on client {self.cid} for round "
+ f"{config['current_round']!r}/{config['total_rounds']!r}"
+ )
+
+ original_parameters = deepcopy(parameters)
+
+ self.set_parameters(parameters)
+ self.train_config.lr = config["lr"]
+
+ loss = train(
+ self.net,
+ self.trainloader,
+ self.know_distill,
+ self.max_p,
+ p_s=self.p_s,
+ epochs=self.train_config.local_epochs,
+ current_round=int(config["current_round"]),
+ total_rounds=int(config["total_rounds"]),
+ train_config=self.train_config,
+ )
+
+ final_parameters = self.get_parameters(config={})
+
+ return (
+ final_parameters,
+ len(self.trainloader.dataset),
+ {
+ "max_p": self.max_p,
+ "p_s": self.p_s,
+ "fjord_config": self.fjord_config,
+ "original_parameters": original_parameters,
+ "loss": loss,
+ },
+ )
+
+ def evaluate(
+ self, parameters: List[np.ndarray], config: Dict[str, fl.common.Scalar]
+ ) -> Tuple[float, int, Dict[str, Union[bool, bytes, float, int, str]]]:
+ """Validate the model on the test set.
+
+ :param parameters: The parameters of the model.
+ :param config: The eval configuration.
+ :return: The loss on the test set, the number of samples used for evaluation,
+ and the evaluation metrics.
+ """
+ Logger.get().info(
+ f"Evaluating on client {self.cid} for round "
+ f"{config['current_round']!r}/{config['total_rounds']!r}"
+ )
+
+ self.set_parameters(parameters)
+ loss, accuracy = test(self.net, self.valloader, [self.max_p])
+ save_model(self.net, self.model_path, cid=self.cid)
+
+ return loss[0], len(self.valloader.dataset), {"accuracy": accuracy[0]}
diff --git a/baselines/fjord/fjord/conf/__init__.py b/baselines/fjord/fjord/conf/__init__.py
new file mode 100644
index 000000000000..39fdacc8e90b
--- /dev/null
+++ b/baselines/fjord/fjord/conf/__init__.py
@@ -0,0 +1 @@
+"""Fjord configuration."""
diff --git a/baselines/fjord/fjord/conf/common.yaml b/baselines/fjord/fjord/conf/common.yaml
new file mode 100644
index 000000000000..d0f392faf4ff
--- /dev/null
+++ b/baselines/fjord/fjord/conf/common.yaml
@@ -0,0 +1,39 @@
+# @package _global_
+---
+loglevel: info
+logfile: run.log
+
+manual_seed: 123
+model: resnet18
+dataset: cifar10
+num_clients: 100
+data_path: "./data"
+num_workers: 4
+evaluate_every: 10
+
+cuda: true
+batch_size: 32
+lr: 0.1
+lr_scheduler: static
+optimiser: sgd
+momentum: 0
+nesterov: false
+weight_decay: 1e-4
+
+sampled_clients: 10
+min_fit_clients: 2
+client_selection: random # or balanced
+num_rounds: 500
+local_epochs: 1
+strategy: fjord_fedavg
+client_resources:
+ num_cpus: 1
+ num_gpus: 0.2
+knowledge_distillation: ???
+p_s:
+ - 0.2
+ - 0.4
+ - 0.6
+ - 0.8
+ - 1.0
+client_tier_allocation: uniform
diff --git a/baselines/fjord/fjord/conf/config.yaml b/baselines/fjord/fjord/conf/config.yaml
new file mode 100644
index 000000000000..a1cdc87c63ce
--- /dev/null
+++ b/baselines/fjord/fjord/conf/config.yaml
@@ -0,0 +1,8 @@
+---
+hydra:
+ run:
+ dir: ./runs/${now:%Y-%m-%d}:${now:%H-%M-%S}
+
+defaults:
+ - train_mode/fjord
+ - override hydra/job_logging: disabled
diff --git a/baselines/fjord/fjord/conf/train_mode/fjord.yaml b/baselines/fjord/fjord/conf/train_mode/fjord.yaml
new file mode 100644
index 000000000000..33b0e17957e0
--- /dev/null
+++ b/baselines/fjord/fjord/conf/train_mode/fjord.yaml
@@ -0,0 +1,6 @@
+# @package _global_
+---
+defaults:
+ - ../common@
+
+knowledge_distillation: false
\ No newline at end of file
diff --git a/baselines/fjord/fjord/conf/train_mode/fjord_kd.yaml b/baselines/fjord/fjord/conf/train_mode/fjord_kd.yaml
new file mode 100644
index 000000000000..d344d95314e9
--- /dev/null
+++ b/baselines/fjord/fjord/conf/train_mode/fjord_kd.yaml
@@ -0,0 +1,6 @@
+# @package _global_
+---
+defaults:
+ - ../common@
+
+knowledge_distillation: true
\ No newline at end of file
diff --git a/baselines/fjord/fjord/dataset.py b/baselines/fjord/fjord/dataset.py
new file mode 100644
index 000000000000..478826c2cf64
--- /dev/null
+++ b/baselines/fjord/fjord/dataset.py
@@ -0,0 +1,174 @@
+"""Dataset for CIFAR10."""
+import random
+from typing import Optional, Tuple
+
+import numpy as np
+import torch
+from PIL import Image
+from torch.nn import Module
+from torch.utils.data import DataLoader, Dataset
+from torchvision import transforms
+from torchvision.datasets import CIFAR10
+
+CIFAR_NORMALIZATION = ((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))
+
+
+class FLCifar10Client(Dataset):
+ """Class implementing the partitioned CIFAR10 dataset."""
+
+ def __init__(self, fl_dataset: Dataset, client_id: Optional[int] = None) -> None:
+ """Ctor.
+
+ Args:
+ :param fl_dataset: The CIFAR10 dataset.
+ :param client_id: The client id to be used.
+ """
+ self.fl_dataset = fl_dataset
+ self.set_client(client_id)
+
+ def set_client(self, index: Optional[int] = None) -> None:
+ """Set the client to the given index. If index is None, use the whole dataset.
+
+ Args:
+ :param index: Index of the client to be used.
+ """
+ fl = self.fl_dataset
+ if index is None:
+ self.client_id = None
+ self.length = len(fl.data)
+ self.data = fl.data
+ else:
+ if index < 0 or index >= fl.num_clients:
+ raise ValueError("Number of clients is out of bounds.")
+ self.client_id = index
+ indices = fl.partition[self.client_id]
+ self.length = len(indices)
+ self.data = fl.data[indices]
+ self.targets = [fl.targets[i] for i in indices]
+
+ def __getitem__(self, index: int):
+ """Return the item at the given index.
+
+ :param index: Index of the item to be returned.
+ :return: The item at the given index.
+ """
+ fl = self.fl_dataset
+ img, target = self.data[index], self.targets[index]
+
+ # doing this so that it is consistent with all other fl_datasets
+ # to return a PIL Image
+ img = Image.fromarray(img)
+
+ if fl.transform is not None:
+ img = fl.transform(img)
+
+ if fl.target_transform is not None:
+ target = fl.target_transform(target)
+
+ return img, target
+
+ def __len__(self):
+ """Return the length of the dataset."""
+ return self.length
+
+
+class FLCifar10(CIFAR10):
+ """CIFAR10 Federated Dataset."""
+
+ def __init__( # pylint: disable=too-many-arguments
+ self,
+ root: str,
+ train: Optional[bool] = True,
+ transform: Optional[Module] = None,
+ target_transform: Optional[Module] = None,
+ download: Optional[bool] = False,
+ ) -> None:
+ """Ctor.
+
+ :param root: Root directory of dataset
+ :param train: If True, creates dataset from training set
+ :param transform: A function/transform that takes in an PIL image and returns a
+ transformed version.
+ :param target_transform: A function/transform that takes in the target and
+ transforms it.
+ :param download: If true, downloads the dataset from the internet.
+ """
+ super().__init__(
+ root,
+ train=train,
+ transform=transform,
+ target_transform=target_transform,
+ download=download,
+ )
+
+ # Uniform shuffle
+ shuffle = np.arange(len(self.data))
+ rng = np.random.default_rng(12345)
+ rng.shuffle(shuffle)
+ self.partition = shuffle.reshape([100, -1])
+ self.num_clients = len(self.partition)
+
+
+def get_transforms() -> Tuple[transforms.Compose, transforms.Compose]:
+ """Get the transforms for the CIFAR10 dataset.
+
+ :return: The transforms for the CIFAR10 dataset.
+ """
+ transform_train = transforms.Compose(
+ [
+ transforms.RandomCrop(32, padding=4),
+ transforms.RandomHorizontalFlip(),
+ transforms.ToTensor(),
+ transforms.Normalize(*CIFAR_NORMALIZATION),
+ ]
+ )
+
+ transform_test = transforms.Compose(
+ [
+ transforms.ToTensor(),
+ transforms.Normalize(*CIFAR_NORMALIZATION),
+ ]
+ )
+
+ return transform_train, transform_test
+
+
+def load_data(
+ path: str, cid: int, train_bs: int, seed: int, eval_bs: int = 1024
+) -> Tuple[DataLoader, DataLoader]:
+ """Load the CIFAR10 dataset.
+
+ :param path: The path to the dataset.
+ :param cid: The client ID.
+ :param train_bs: The batch size for training.
+ :param seed: The seed to use for the random number generator.
+ :param eval_bs: The batch size for evaluation.
+ :return: The training and test sets.
+ """
+
+ def seed_worker(worker_id): # pylint: disable=unused-argument
+ worker_seed = torch.initial_seed() % 2**32
+ np.random.seed(worker_seed)
+ random.seed(worker_seed)
+
+ g = torch.Generator()
+ g.manual_seed(seed)
+ transform_train, transform_test = get_transforms()
+
+ fl_dataset = FLCifar10(
+ root=path, train=True, download=True, transform=transform_train
+ )
+
+ trainset = FLCifar10Client(fl_dataset, client_id=cid)
+ testset = CIFAR10(root=path, train=False, download=True, transform=transform_test)
+
+ train_loader = DataLoader(
+ trainset,
+ batch_size=train_bs,
+ shuffle=True,
+ worker_init_fn=seed_worker,
+ generator=g,
+ )
+ test_loader = DataLoader(testset, batch_size=eval_bs)
+
+ return train_loader, test_loader
diff --git a/baselines/fjord/fjord/dataset_preparation.py b/baselines/fjord/fjord/dataset_preparation.py
new file mode 100644
index 000000000000..fe70679d3351
--- /dev/null
+++ b/baselines/fjord/fjord/dataset_preparation.py
@@ -0,0 +1 @@
+"""All dataset-related logic happens in dataset.py."""
diff --git a/baselines/fjord/fjord/main.py b/baselines/fjord/fjord/main.py
new file mode 100644
index 000000000000..f85fb9ccf158
--- /dev/null
+++ b/baselines/fjord/fjord/main.py
@@ -0,0 +1,278 @@
+"""Main script for FjORD."""
+import math
+import os
+import random
+from types import SimpleNamespace
+from typing import Any, Callable, Dict, List, Optional, Union
+
+import flwr as fl
+import hydra
+import numpy as np
+import torch
+from flwr.client import Client, NumPyClient
+from omegaconf import OmegaConf, open_dict
+
+from .client import FJORD_CONFIG_TYPE, FjORDClient, get_agg_config
+from .dataset import load_data
+from .models import get_net
+from .server import get_eval_fn
+from .strategy import FjORDFedAVG
+from .utils.logger import Logger
+from .utils.utils import get_parameters
+
+
+def get_fit_config_fn(
+ total_rounds: int, lr: float
+) -> Callable[[int], Dict[str, fl.common.Scalar]]:
+ """Get fit config function.
+
+ :param total_rounds: Total number of rounds
+ :param lr: Learning rate
+ :return: Fit config function
+ """
+
+ def fit_config(rnd: int) -> Dict[str, fl.common.Scalar]:
+ config: Dict[str, fl.common.Scalar] = {
+ "current_round": rnd,
+ "total_rounds": total_rounds,
+ "lr": lr,
+ }
+ return config
+
+ return fit_config
+
+
+def get_client_fn( # pylint: disable=too-many-arguments
+ args: Any,
+ model_path: str,
+ cid_to_max_p: Dict[int, float],
+ config: FJORD_CONFIG_TYPE,
+ train_config: SimpleNamespace,
+ device: torch.device,
+) -> Callable[[str], Union[Client, NumPyClient]]:
+ """Get client function that creates Flower client.
+
+ :param args: CLI/Config Arguments
+ :param model_path: Path to save the model
+ :param cid_to_max_p: Dictionary mapping client id to max p-value
+ :param config: Aggregation config
+ :param train_config: Training config
+ :param device: Device to be used
+ :return: Client function that returns Flower client
+ """
+
+ def client_fn(cid) -> FjORDClient:
+ max_p = cid_to_max_p[int(cid)]
+ log_config = {
+ "loglevel": args.loglevel,
+ "logfile": args.logfile,
+ }
+ return FjORDClient(
+ cid=cid,
+ model_name=args.model,
+ data_path=args.data_path,
+ model_path=model_path,
+ know_distill=args.knowledge_distillation,
+ max_p=max_p,
+ p_s=args.p_s,
+ fjord_config=config,
+ train_config=train_config,
+ log_config=log_config,
+ seed=args.manual_seed,
+ device=device,
+ )
+
+ return client_fn
+
+
+class FjORDBalancedClientManager(fl.server.SimpleClientManager):
+ """Balanced client manager for FjORD.
+
+ This class samples equal number of clients per p-value and the rest in RR.
+ """
+
+ def __init__(self, cid_to_max_p: Dict[int, float]) -> None:
+ """Ctor.
+
+ Args:
+ :param cid_to_max_p: Dictionary mapping client id to max p-value
+ """
+ super().__init__()
+ self.cid_to_max_p = cid_to_max_p
+ self.p_s = sorted(set(self.cid_to_max_p.values()))
+
+ def sample(
+ self,
+ num_clients: int,
+ min_num_clients: Optional[int] = None,
+ criterion: Optional[fl.server.criterion.Criterion] = None,
+ ) -> List[fl.server.client_proxy.ClientProxy]:
+ """Sample clients in a balanced way (equal per tier, remainder in Round-Robin).
+
+ Args:
+ :param num_clients: Number of clients to sample
+ :param min_num_clients: Minimum number of clients to sample
+ :param criterion: Client selection criterion
+ :return: List of sampled clients
+ """
+ if min_num_clients is None:
+ min_num_clients = num_clients
+ self.wait_for(min_num_clients)
+ available_cids = list(self.clients)
+ if criterion is not None:
+ available_cids = [
+ cid for cid in available_cids if criterion.select(self.clients[cid])
+ ]
+ if num_clients > len(available_cids):
+ Logger.get().info(
+ "Sampling failed: number of available clients"
+ " (%s) is less than number of requested clients (%s).",
+ len(available_cids),
+ num_clients,
+ )
+ return []
+
+ # construct p to available cids
+ max_p_to_cids: Dict[float, List[int]] = {p: [] for p in self.p_s}
+ random.shuffle(available_cids)
+ for cid_s in available_cids:
+ client_id = int(cid_s)
+ client_p = self.cid_to_max_p[client_id]
+ max_p_to_cids[client_p].append(client_id)
+
+ cl_per_tier = math.floor(num_clients / len(self.p_s))
+ remainder = num_clients - cl_per_tier * len(self.p_s)
+
+ selected_cids = set()
+ for p in self.p_s:
+ for cid in random.sample(max_p_to_cids[p], cl_per_tier):
+ selected_cids.add(cid)
+
+ for p in self.p_s:
+ if remainder == 0:
+ break
+ cid = random.choice(max_p_to_cids[p])
+ while cid not in selected_cids:
+ cid = random.choice(max_p_to_cids[p])
+ selected_cids.add(cid)
+ remainder -= 1
+
+ Logger.get().debug(f"Sampled {selected_cids}")
+ return [self.clients[str(cid)] for cid in selected_cids]
+
+
+def main(args: Any) -> None:
+ """Enter main functionality.
+
+ Args:
+ :param args: CLI/Config Arguments
+ """
+ torch.manual_seed(args.manual_seed)
+ torch.use_deterministic_algorithms(True)
+ np.random.seed(args.manual_seed)
+ random.seed(args.manual_seed)
+
+ path = args.data_path
+ device = torch.device("cuda") if args.cuda else torch.device("cpu")
+ model_path = hydra.core.hydra_config.HydraConfig.get().runtime.output_dir
+
+ Logger.get().info(
+ f"Training on {device} using PyTorch "
+ f"{torch.__version__} and Flower {fl.__version__}"
+ )
+
+ trainloader, testloader = load_data(
+ path, cid=0, seed=args.manual_seed, train_bs=args.batch_size
+ )
+ NUM_CLIENTS = args.num_clients
+ if args.client_tier_allocation == "uniform":
+ cid_to_max_p = {cid: (cid // 20) * 0.2 + 0.2 for cid in range(100)}
+ else:
+ raise ValueError(
+ f"Client to tier allocation strategy "
+ f"{args.client_tier_allocation} not currently"
+ "supported"
+ )
+
+ model = get_net(args.model, args.p_s, device=device)
+ config = get_agg_config(model, trainloader, args.p_s)
+ train_config = SimpleNamespace(
+ **{
+ "batch_size": args.batch_size,
+ "lr": args.lr,
+ "optimiser": args.optimiser,
+ "momentum": args.momentum,
+ "nesterov": args.nesterov,
+ "lr_scheduler": args.lr_scheduler,
+ "weight_decay": args.weight_decay,
+ "local_epochs": args.local_epochs,
+ }
+ )
+
+ if args.strategy == "fjord_fedavg":
+ strategy = FjORDFedAVG(
+ fraction_fit=args.sampled_clients / args.num_clients,
+ fraction_evaluate=0.0,
+ min_fit_clients=args.min_fit_clients,
+ min_evaluate_clients=1,
+ min_available_clients=NUM_CLIENTS,
+ evaluate_fn=get_eval_fn(args, model_path, testloader, device),
+ on_fit_config_fn=get_fit_config_fn(args.num_rounds, args.lr),
+ initial_parameters=fl.common.ndarrays_to_parameters(
+ get_parameters(get_net(args.model, args.p_s, device=device))
+ ),
+ )
+ else:
+ raise ValueError(f"Strategy {args.strategy} is not currently supported")
+
+ client_resources = args.client_resources
+ if device.type != "cuda":
+ client_resources = {
+ "num_cpus": args.client_resources["num_cpus"],
+ "num_gpus": 0,
+ }
+
+ if args.client_selection == "balanced":
+ cl_manager = FjORDBalancedClientManager(cid_to_max_p)
+ elif args.client_selection == "random":
+ cl_manager = None
+ else:
+ raise ValueError(
+ f"Client selection {args.client_selection} is not currently supported"
+ )
+
+ Logger.get().info("Starting simulated run.")
+ # Start simulation
+ fl.simulation.start_simulation(
+ client_fn=get_client_fn(
+ args, model_path, cid_to_max_p, config, train_config, device
+ ),
+ num_clients=NUM_CLIENTS,
+ config=fl.server.ServerConfig(num_rounds=args.num_rounds),
+ strategy=strategy,
+ client_resources=client_resources,
+ client_manager=cl_manager,
+ ray_init_args={"include_dashboard": False},
+ )
+
+
+@hydra.main(version_base=None, config_path="conf", config_name="config")
+def run_app(cfg):
+ """Run the application.
+
+ Args:
+ :param cfg: Hydra configuration
+ """
+ OmegaConf.resolve(cfg)
+ logfile = os.path.join(
+ hydra.core.hydra_config.HydraConfig.get()["runtime"]["output_dir"], cfg.logfile
+ )
+ with open_dict(cfg):
+ cfg.logfile = logfile
+ Logger.setup_logging(loglevel=cfg.loglevel, logfile=logfile)
+ Logger.get().info(f"Hydra configuration: {OmegaConf.to_yaml(cfg)}")
+ main(cfg)
+
+
+if __name__ == "__main__":
+ run_app()
diff --git a/baselines/fjord/fjord/models.py b/baselines/fjord/fjord/models.py
new file mode 100644
index 000000000000..0f3fc276decf
--- /dev/null
+++ b/baselines/fjord/fjord/models.py
@@ -0,0 +1,319 @@
+"""ResNet model for Fjord."""
+from types import SimpleNamespace
+from typing import List, Optional, Tuple
+
+import torch
+import torch.nn.functional as F
+from torch import nn
+from torch.nn import Module
+from torch.optim import Optimizer
+from torch.optim.lr_scheduler import MultiStepLR
+from torch.utils.data import DataLoader
+from tqdm import tqdm
+
+from .od.models.utils import (
+ SequentialWithSampler,
+ create_bn_layer,
+ create_conv_layer,
+ create_linear_layer,
+)
+from .od.samplers import BaseSampler, ODSampler
+
+
+class BasicBlock(nn.Module):
+ """Basic Block for resnet."""
+
+ expansion = 1
+
+ def __init__(
+ self, od, p_s, in_planes, planes, stride=1
+ ): # pylint: disable=too-many-arguments
+ super().__init__()
+ self.od = od
+ self.conv1 = create_conv_layer(
+ od,
+ True,
+ in_planes,
+ planes,
+ kernel_size=3,
+ stride=stride,
+ padding=1,
+ bias=False,
+ )
+ self.bn1 = create_bn_layer(od=od, p_s=p_s, num_features=planes)
+ self.conv2 = create_conv_layer(
+ od, True, planes, planes, kernel_size=3, stride=1, padding=1, bias=False
+ )
+ self.bn2 = create_bn_layer(od=od, p_s=p_s, num_features=planes)
+
+ self.shortcut = SequentialWithSampler()
+ if stride != 1 or in_planes != self.expansion * planes:
+ self.shortcut = SequentialWithSampler(
+ create_conv_layer(
+ od,
+ True,
+ in_planes,
+ self.expansion * planes,
+ kernel_size=1,
+ stride=stride,
+ bias=False,
+ ),
+ create_bn_layer(od=od, p_s=p_s, num_features=self.expansion * planes),
+ )
+
+ def forward(self, x, sampler):
+ """Forward method for basic block.
+
+ Args:
+ :param x: input
+ :param sampler: sampler
+ :return: Output of forward pass
+ """
+ if sampler is None:
+ out = F.relu(self.bn1(self.conv1(x)))
+ out = self.bn2(self.conv2(out))
+ out += self.shortcut(x)
+ out = F.relu(out)
+ else:
+ out = F.relu(self.bn1(self.conv1(x, p=sampler())))
+ out = self.bn2(self.conv2(out, p=sampler()))
+ shortcut = self.shortcut(x, sampler=sampler)
+ assert (
+ shortcut.shape == out.shape
+ ), f"Shortcut shape: {shortcut.shape} out.shape: {out.shape}"
+ out += shortcut
+ # out += self.shortcut(x, sampler=sampler)
+ out = F.relu(out)
+ return out
+
+
+# Adapted from:
+# https://github.com/kuangliu/pytorch-cifar/blob/master/models/resnet.py
+class ResNet(nn.Module): # pylint: disable=too-many-instance-attributes
+ """ResNet in PyTorch.
+
+ Reference:
+ [1] Kaiming He, Xiangyu Zhang, Shaoqing Ren, Jian Sun
+ Deep Residual Learning for Image Recognition. arXiv:1512.03385
+ """
+
+ def __init__(
+ self, od, p_s, block, num_blocks, num_classes=10
+ ): # pylint: disable=too-many-arguments
+ super().__init__()
+ self.od = od
+ self.in_planes = 64
+
+ self.conv1 = create_conv_layer(
+ od, True, 3, 64, kernel_size=3, stride=1, padding=1, bias=False
+ )
+ self.bn1 = create_bn_layer(od=od, p_s=p_s, num_features=64)
+ self.layer1 = self._make_layer(od, p_s, block, 64, num_blocks[0], stride=1)
+ self.layer2 = self._make_layer(od, p_s, block, 128, num_blocks[1], stride=2)
+ self.layer3 = self._make_layer(od, p_s, block, 256, num_blocks[2], stride=2)
+ self.layer4 = self._make_layer(od, p_s, block, 512, num_blocks[3], stride=2)
+ self.linear = create_linear_layer(od, False, 512 * block.expansion, num_classes)
+
+ def _make_layer(
+ self, od, p_s, block, planes, num_blocks, stride
+ ): # pylint: disable=too-many-arguments
+ strides = [stride] + [1] * (num_blocks - 1)
+ layers = []
+ for strd in strides:
+ layers.append(block(od, p_s, self.in_planes, planes, strd))
+ self.in_planes = planes * block.expansion
+ return SequentialWithSampler(*layers)
+
+ def forward(self, x, sampler=None):
+ """Forward method for ResNet.
+
+ Args:
+ :param x: input
+ :param sampler: sampler
+ :return: Output of forward pass
+ """
+ if self.od:
+ if sampler is None:
+ sampler = BaseSampler(self)
+ out = F.relu(self.bn1(self.conv1(x, p=sampler())))
+ out = self.layer1(out, sampler=sampler)
+ out = self.layer2(out, sampler=sampler)
+ out = self.layer3(out, sampler=sampler)
+ out = self.layer4(out, sampler=sampler)
+ out = F.avg_pool2d(out, 4) # pylint: disable=not-callable
+ out = out.view(out.size(0), -1)
+ out = self.linear(out)
+ else:
+ out = F.relu(self.bn1(self.conv1(x)))
+ out = self.layer1(out)
+ out = self.layer2(out)
+ out = self.layer3(out)
+ out = self.layer4(out)
+ out = F.avg_pool2d(out, 4) # pylint: disable=not-callable
+ out = out.view(out.size(0), -1)
+ out = self.linear(out)
+ return out
+
+
+def ResNet18(od=False, p_s=(1.0,)):
+ """Construct a ResNet-18 model.
+
+ Args:
+ :param od: whether to create OD (Ordered Dropout) layer
+ :param p_s: list of p-values
+ """
+ return ResNet(od, p_s, BasicBlock, [2, 2, 2, 2])
+
+
+def get_net(
+ model_name: str,
+ p_s: List[float],
+ device: torch.device,
+) -> torch.nn.Module:
+ """Initialise model.
+
+ :param model_name: name of the model
+ :param p_s: list of p-values
+ :param device: device to be used
+ :return: initialised model
+ """
+ if model_name == "resnet18":
+ net = ResNet18(od=True, p_s=p_s).to(device)
+ else:
+ raise ValueError(f"Model {model_name} is not supported")
+
+ return net
+
+
+def train( # pylint: disable=too-many-locals, too-many-arguments
+ net: Module,
+ trainloader: DataLoader,
+ know_distill: bool,
+ max_p: float,
+ current_round: int,
+ total_rounds: int,
+ p_s: List[float],
+ epochs: int,
+ train_config: SimpleNamespace,
+) -> float:
+ """Train the model on the training set.
+
+ :param net: The model to train.
+ :param trainloader: The training set.
+ :param know_distill: Whether the model being trained uses knowledge distillation.
+ :param max_p: The maximum p value.
+ :param current_round: The current round of training.
+ :param total_rounds: The total number of rounds of training.
+ :param p_s: The p values to use for training.
+ :param epochs: The number of epochs to train for.
+ :param train_config: The training configuration.
+ :return: The loss on the training set.
+ """
+ device = next(net.parameters()).device
+ criterion = torch.nn.CrossEntropyLoss()
+ net.train()
+ if train_config.optimiser == "sgd":
+ optimizer = torch.optim.SGD(
+ net.parameters(),
+ lr=train_config.lr,
+ momentum=train_config.momentum,
+ nesterov=train_config.nesterov,
+ weight_decay=train_config.weight_decay,
+ )
+ else:
+ raise ValueError(f"Optimiser {train_config.optimiser} not supported")
+ lr_scheduler = get_lr_scheduler(
+ optimizer, total_rounds, method=train_config.lr_scheduler
+ )
+ for _ in range(current_round):
+ lr_scheduler.step()
+
+ sampler = ODSampler(
+ p_s=p_s,
+ max_p=max_p,
+ model=net,
+ )
+ max_sampler = ODSampler(
+ p_s=[max_p],
+ max_p=max_p,
+ model=net,
+ )
+
+ loss = 0.0
+ samples = 0
+ for _ in range(epochs):
+ for images, labels in trainloader:
+ optimizer.zero_grad()
+ target = labels.to(device)
+ images = images.to(device)
+ batch_size = images.shape[0]
+ if know_distill:
+ full_output = net(images.to(device), sampler=max_sampler)
+ full_loss = criterion(full_output, target)
+ full_loss.backward()
+ target = full_output.detach().softmax(dim=1)
+ partial_loss = criterion(net(images, sampler=sampler), target)
+ partial_loss.backward()
+ optimizer.step()
+ loss += partial_loss.item() * batch_size
+ samples += batch_size
+
+ return loss / samples
+
+
+def test(
+ net: Module, testloader: DataLoader, p_s: List[float]
+) -> Tuple[List[float], List[float]]:
+ """Validate the model on the test set.
+
+ :param net: The model to validate.
+ :param testloader: The test set.
+ :param p_s: The p values to use for validation.
+ :return: The loss and accuracy on the test set.
+ """
+ device = next(net.parameters()).device
+ criterion = torch.nn.CrossEntropyLoss()
+ losses = []
+ accuracies = []
+ net.eval()
+
+ for p in p_s:
+ correct, loss = 0, 0.0
+ p_sampler = ODSampler(
+ p_s=[p],
+ max_p=p,
+ model=net,
+ )
+
+ with torch.no_grad():
+ for images, labels in tqdm(testloader):
+ outputs = net(images.to(device), sampler=p_sampler)
+ labels = labels.to(device)
+ loss += criterion(outputs, labels).item() * images.shape[0]
+ correct += (torch.max(outputs.data, 1)[1] == labels).sum().item()
+ accuracy = correct / len(testloader.dataset)
+ losses.append(loss / len(testloader.dataset))
+ accuracies.append(accuracy)
+
+ return losses, accuracies
+
+
+def get_lr_scheduler(
+ optimiser: Optimizer,
+ total_epochs: int,
+ method: Optional[str] = "static",
+) -> torch.optim.lr_scheduler.LRScheduler:
+ """Get the learning rate scheduler.
+
+ :param optimiser: The optimiser for which to get the scheduler.
+ :param total_epochs: The total number of epochs.
+ :param method: The method to use for the scheduler. Supports static and cifar10.
+ :return: The learning rate scheduler.
+ """
+ if method == "static":
+ return MultiStepLR(optimiser, [total_epochs + 1])
+ if method == "cifar10":
+ return MultiStepLR(
+ optimiser, [int(0.5 * total_epochs), int(0.75 * total_epochs)], gamma=0.1
+ )
+ raise ValueError(f"{method} scheduler not currently supported.")
diff --git a/baselines/fjord/fjord/od/__init__.py b/baselines/fjord/fjord/od/__init__.py
new file mode 100644
index 000000000000..f2b055c479f2
--- /dev/null
+++ b/baselines/fjord/fjord/od/__init__.py
@@ -0,0 +1 @@
+"""Ordered dropout package."""
diff --git a/baselines/fjord/fjord/od/layers/__init__.py b/baselines/fjord/fjord/od/layers/__init__.py
new file mode 100644
index 000000000000..a87c70401d4c
--- /dev/null
+++ b/baselines/fjord/fjord/od/layers/__init__.py
@@ -0,0 +1,6 @@
+"""Ordered Dropout layers."""
+from .batch_norm import ODBatchNorm2d
+from .conv import ODConv2d
+from .linear import ODLinear
+
+__all__ = ["ODBatchNorm2d", "ODConv2d", "ODLinear"]
diff --git a/baselines/fjord/fjord/od/layers/batch_norm.py b/baselines/fjord/fjord/od/layers/batch_norm.py
new file mode 100644
index 000000000000..5fce4dff0910
--- /dev/null
+++ b/baselines/fjord/fjord/od/layers/batch_norm.py
@@ -0,0 +1,75 @@
+"""BatchNorm using Ordered Dropout."""
+from typing import List, Optional
+
+import numpy as np
+import torch
+from torch import Tensor, nn
+
+__all__ = ["ODBatchNorm2d"]
+
+
+class ODBatchNorm2d(nn.Module): # pylint: disable=too-many-instance-attributes
+ """Ordered Dropout BatchNorm2d."""
+
+ def __init__(
+ self,
+ *args,
+ p_s: List[float],
+ num_features: int,
+ affine: Optional[bool] = True,
+ **kwargs,
+ ) -> None:
+ super().__init__()
+ self.p_s = p_s
+ self.is_od = False # no sampling is happening here
+ self.num_features = num_features
+ self.num_features_s = [int(np.ceil(num_features * p)) for p in p_s]
+ self.p_to_num_features = dict(zip(p_s, self.num_features_s))
+ self.width = np.max(self.num_features_s)
+ self.last_input_dim = None
+
+ self.bn = nn.ModuleDict(
+ {
+ str(num_features): nn.BatchNorm2d(
+ num_features, *args, **kwargs, affine=False
+ )
+ for num_features in self.num_features_s
+ }
+ )
+
+ # single track_running_stats
+ if affine:
+ self.affine = True
+ self.weight = nn.Parameter(torch.Tensor(self.width, 1, 1))
+ self.bias = nn.Parameter(torch.Tensor(self.width, 1, 1))
+
+ self.reset_parameters()
+
+ # get p into the layer
+ for m, p in zip(self.bn, self.p_s):
+ self.bn[m].p = p
+ self.bn[m].num_batches_tracked = torch.tensor(1, dtype=torch.long)
+
+ def reset_parameters(self):
+ """Reset parameters."""
+ if self.affine:
+ nn.init.ones_(self.weight)
+ nn.init.zeros_(self.bias)
+ for m in self.bn:
+ self.bn[m].reset_parameters()
+
+ def forward(self, x: Tensor) -> Tensor:
+ """Forward pass.
+
+ Args:
+ :param x: Input tensor.
+ :return: Output of forward pass.
+ """
+ in_dim = x.size(1) # second dimension is input dimension
+ assert (
+ in_dim in self.num_features_s
+ ), "input dimension not in selected num_features_s"
+ out = self.bn[str(in_dim)](x)
+ if self.affine:
+ out = out * self.weight[:in_dim] + self.bias[:in_dim]
+ return out
diff --git a/baselines/fjord/fjord/od/layers/conv.py b/baselines/fjord/fjord/od/layers/conv.py
new file mode 100644
index 000000000000..544f3a578418
--- /dev/null
+++ b/baselines/fjord/fjord/od/layers/conv.py
@@ -0,0 +1,140 @@
+"""Convolutional layer using Ordered Dropout."""
+from typing import Optional, Tuple, Union
+
+import numpy as np
+from torch import Tensor, nn
+from torch.nn import Module
+
+from .utils import check_layer
+
+__all__ = ["ODConv1d", "ODConv2d", "ODConv3d"]
+
+
+def od_conv_forward(
+ layer: Module, x: Tensor, p: Optional[Union[Tuple[Module, float], float]] = None
+) -> Tensor:
+ """Ordered dropout forward pass for convolution networks.
+
+ Args:
+ :param layer: The layer being forwarded.
+ :param x: Input tensor.
+ :param p: Tuple of layer and p or p.
+ :return: Output of forward pass.
+ """
+ p = check_layer(layer, p)
+ if not layer.is_od and p is not None:
+ raise ValueError("p must be None if is_od is False")
+ in_dim = x.size(1) # second dimension is input dimension
+ layer.last_input_dim = in_dim
+ if not p: # i.e., don't apply OD
+ out_dim = layer.width
+ else:
+ out_dim = int(np.ceil(layer.width * p))
+ layer.last_output_dim = out_dim
+ # subsampled weights and bias
+ weights_red = layer.weight[:out_dim, :in_dim]
+ bias_red = layer.bias[:out_dim] if layer.bias is not None else None
+ return layer._conv_forward( # pylint: disable=protected-access
+ x, weights_red, bias_red
+ )
+
+
+def get_slice(layer: Module, in_dim: int, out_dim: int) -> Tuple[Tensor, Tensor]:
+ """Get slice of weights and bias.
+
+ Args:
+ :param layer: The layer.
+ :param in_dim: The input dimension.
+ :param out_dim: The output dimension.
+ :return: The slice of weights and bias.
+ """
+ weight_slice = layer.weight[:in_dim, :out_dim]
+ bias_slice = layer.bias[:out_dim] if layer.bias is not None else None
+ return weight_slice, bias_slice
+
+
+class ODConv1d(nn.Conv1d):
+ """Ordered Dropout Conv1d."""
+
+ def __init__(self, *args, is_od: bool = True, **kwargs) -> None:
+ self.is_od = is_od
+ super().__init__(*args, **kwargs)
+ self.width = self.out_channels
+ self.last_input_dim = None
+ self.last_output_dim = None
+
+ def forward( # pylint: disable=arguments-differ
+ self,
+ input: Tensor, # pylint: disable=redefined-builtin
+ p: Optional[Union[Tuple[Module, float], float]] = None,
+ ) -> Tensor:
+ """Forward pass.
+
+ Args:
+ :param input: Input tensor.
+ :param p: Tuple of layer and p or p.
+ :return: Output of forward pass.
+ """
+ return od_conv_forward(self, input, p)
+
+ def get_slice(self, *args, **kwargs) -> Tuple[Tensor, Tensor]:
+ """Get slice of weights and bias."""
+ return get_slice(self, *args, **kwargs)
+
+
+class ODConv2d(nn.Conv2d):
+ """Ordered Dropout Conv2d."""
+
+ def __init__(self, *args, is_od: bool = True, **kwargs) -> None:
+ self.is_od = is_od
+ super().__init__(*args, **kwargs)
+ self.width = self.out_channels
+ self.last_input_dim = None
+ self.last_output_dim = None
+
+ def forward( # pylint: disable=arguments-differ
+ self,
+ input: Tensor, # pylint: disable=redefined-builtin
+ p: Optional[Union[Tuple[Module, float], float]] = None,
+ ) -> Tensor:
+ """Forward pass.
+
+ Args:
+ :param input: Input tensor.
+ :param p: Tuple of layer and p or p.
+ :return: Output of forward pass.
+ """
+ return od_conv_forward(self, input, p)
+
+ def get_slice(self, *args, **kwargs) -> Tuple[Tensor, Tensor]:
+ """Get slice of weights and bias."""
+ return get_slice(self, *args, **kwargs)
+
+
+class ODConv3d(nn.Conv3d):
+ """Ordered Dropout Conv3d."""
+
+ def __init__(self, *args, is_od: bool = True, **kwargs) -> None:
+ self.is_od = is_od
+ super().__init__(*args, **kwargs)
+ self.width = self.out_channels
+ self.last_input_dim = None
+ self.last_output_dim = None
+
+ def forward( # pylint: disable=arguments-differ
+ self,
+ input: Tensor, # pylint: disable=redefined-builtin
+ p: Optional[Union[Tuple[Module, float], float]] = None,
+ ) -> Tensor:
+ """Forward pass.
+
+ Args:
+ :param input: Input tensor.
+ :param p: Tuple of layer and p or p.
+ :return: Output of forward pass.
+ """
+ return od_conv_forward(self, input, p)
+
+ def get_slice(self, *args, **kwargs) -> Tuple[Tensor, Tensor]:
+ """Get slice of weights and bias."""
+ return get_slice(self, *args, **kwargs)
diff --git a/baselines/fjord/fjord/od/layers/linear.py b/baselines/fjord/fjord/od/layers/linear.py
new file mode 100644
index 000000000000..927ae4c8d516
--- /dev/null
+++ b/baselines/fjord/fjord/od/layers/linear.py
@@ -0,0 +1,62 @@
+"""Liner layer using Ordered Dropout."""
+from typing import Optional, Tuple, Union
+
+import numpy as np
+import torch.nn.functional as F
+from torch import Tensor, nn
+from torch.nn import Module
+
+from .utils import check_layer
+
+__all__ = ["ODLinear"]
+
+
+class ODLinear(nn.Linear):
+ """Ordered Dropout Linear."""
+
+ def __init__(self, *args, is_od: bool = True, **kwargs) -> None:
+ super().__init__(*args, **kwargs)
+ self.is_od = is_od
+ self.width = self.out_features
+ self.last_input_dim = None
+ self.last_output_dim = None
+
+ def forward( # pylint: disable=arguments-differ
+ self,
+ input: Tensor, # pylint: disable=redefined-builtin
+ p: Optional[Union[Tuple[Module, float], float]] = None,
+ ) -> Tensor:
+ """Forward pass.
+
+ Args:
+ :param input: Input tensor.
+ :param p: Tuple of layer and p or p.
+ :return: Output of forward pass.
+ """
+ if not self.is_od and p is not None:
+ raise ValueError("p must be None if is_od is False")
+ p = check_layer(self, p)
+ in_dim = input.size(1) # second dimension is input dimension
+ self.last_input_dim = in_dim
+ if not p: # i.e., don't apply OD
+ out_dim = self.width
+ else:
+ out_dim = int(np.ceil(self.width * p))
+ self.last_output_dim = out_dim
+ # subsampled weights and bias
+ weights_red = self.weight[:out_dim, :in_dim]
+ bias_red = self.bias[:out_dim] if self.bias is not None else None
+ return F.linear(input, weights_red, bias_red) # pylint: disable=not-callable
+
+ def get_slice(self, in_dim: int, out_dim: int) -> Tuple[Tensor, Tensor]:
+ """Get slice of weights and bias.
+
+ Args:
+ :param layer: The layer.
+ :param in_dim: The input dimension.
+ :param out_dim: The output dimension.
+ :return: The slice of weights and bias.
+ """
+ weight_slice = self.weight[:in_dim, :out_dim]
+ bias_slice = self.bias[:out_dim] if self.bias is not None else None
+ return weight_slice, bias_slice
diff --git a/baselines/fjord/fjord/od/layers/utils.py b/baselines/fjord/fjord/od/layers/utils.py
new file mode 100644
index 000000000000..46649a51de96
--- /dev/null
+++ b/baselines/fjord/fjord/od/layers/utils.py
@@ -0,0 +1,23 @@
+"""Utils function for Ordered Dropout layers."""
+from typing import Optional, Tuple, Union
+
+from torch.nn import Module
+
+
+def check_layer(
+ layer: Module, p: Union[Tuple[Module, Optional[float]], Optional[float]]
+) -> Optional[float]:
+ """Check if layer is valid and return p.
+
+ Args:
+ layer: PyTorch layer
+ p: Ordered dropout p
+ """
+ # if p is tuple, check layer validity
+ if isinstance(p, tuple):
+ p_, sampled_layer = p
+ assert layer == sampled_layer, "Layer mismatch"
+ else:
+ p_ = p
+
+ return p_
diff --git a/baselines/fjord/fjord/od/models/__init__.py b/baselines/fjord/fjord/od/models/__init__.py
new file mode 100644
index 000000000000..b0e5ede4f93b
--- /dev/null
+++ b/baselines/fjord/fjord/od/models/__init__.py
@@ -0,0 +1 @@
+"""Functions for creatingin OD models."""
diff --git a/baselines/fjord/fjord/od/models/utils.py b/baselines/fjord/fjord/od/models/utils.py
new file mode 100644
index 000000000000..4a1707587ef4
--- /dev/null
+++ b/baselines/fjord/fjord/od/models/utils.py
@@ -0,0 +1,77 @@
+"""Utility functions for models."""
+from torch import nn
+
+from ..layers import ODBatchNorm2d, ODConv2d, ODLinear
+
+
+def create_linear_layer(od, is_od, *args, **kwargs):
+ """Create linear layer.
+
+ :param od: whether to create OD layer
+ :param is_od: whether to create OD layer
+ :param args: arguments for nn.Linear
+ :param kwargs: keyword arguments for nn.Linear
+ :return: nn.Linear or ODLinear
+ """
+ if od:
+ return ODLinear(*args, is_od=is_od, **kwargs)
+
+ return nn.Linear(*args, **kwargs)
+
+
+def create_conv_layer(od, is_od, *args, **kwargs):
+ """Create conv layer.
+
+ :param od: whether to create OD layer
+ :param is_od: whether to create OD layer
+ :param args: arguments for nn.Conv2d
+ :param kwargs: keyword arguments for nn.Conv2d
+ :return: nn.Conv2d or ODConv2d
+ """
+ if od:
+ return ODConv2d(*args, is_od=is_od, **kwargs)
+
+ return nn.Conv2d(*args, **kwargs)
+
+
+def create_bn_layer(od, p_s, *args, **kwargs):
+ """Create batch norm layer.
+
+ :param od: whether to create OD layer
+ :param p_s: list of p-values
+ :param args: arguments for nn.BatchNorm2d
+ :param kwargs: keyword arguments for nn.BatchNorm2d
+ :return: nn.BatchNorm2d or ODBatchNorm2d
+ """
+ if od:
+ num_features = kwargs["num_features"]
+ del kwargs["num_features"]
+ return ODBatchNorm2d(*args, p_s=p_s, num_features=num_features, **kwargs)
+
+ return nn.BatchNorm2d(*args, **kwargs)
+
+
+class SequentialWithSampler(nn.Sequential):
+ """Implements sequential model with sampler."""
+
+ def forward(
+ self, input, sampler=None
+ ): # pylint: disable=redefined-builtin, arguments-differ
+ """Forward method for custom Sequential.
+
+ :param input: input
+ :param sampler: the sampler to use.
+ :return: Output of sequential
+ """
+ if sampler is None:
+ for module in self:
+ input = module(input)
+ else:
+ for module in self:
+ if hasattr(module, "od") and module.od:
+ input = module(input, sampler=sampler)
+ elif hasattr(module, "is_od") and module.is_od:
+ input = module(input, p=sampler())
+ else:
+ input = module(input)
+ return input
diff --git a/baselines/fjord/fjord/od/samplers/__init__.py b/baselines/fjord/fjord/od/samplers/__init__.py
new file mode 100644
index 000000000000..dad08b4236c4
--- /dev/null
+++ b/baselines/fjord/fjord/od/samplers/__init__.py
@@ -0,0 +1,5 @@
+"""OD samplers."""
+from .base_sampler import BaseSampler
+from .fixed_od import ODSampler
+
+__all__ = ["BaseSampler", "ODSampler"]
diff --git a/baselines/fjord/fjord/od/samplers/base_sampler.py b/baselines/fjord/fjord/od/samplers/base_sampler.py
new file mode 100644
index 000000000000..28eac929df81
--- /dev/null
+++ b/baselines/fjord/fjord/od/samplers/base_sampler.py
@@ -0,0 +1,49 @@
+"""Base sampler class."""
+from collections.abc import Generator
+
+from torch.nn import Module
+
+
+class BaseSampler:
+ """Base class implementing p-value sampling per layer."""
+
+ def __init__(self, model: Module, with_layer: bool = False) -> None:
+ """Initialise sampler.
+
+ :param model: OD model
+ :param with_layer: whether to return layer upon call.
+ """
+ self.model = model
+ self.with_layer = with_layer
+ self.prepare_sampler()
+ self.width_samples = self.width_sampler()
+ self.layer_samples = self.layer_sampler()
+
+ def prepare_sampler(self) -> None:
+ """Prepare sampler."""
+ self.num_od_layers = 0
+ self.widths = []
+ self.od_layers = []
+ for m in self.model.modules():
+ if hasattr(m, "is_od") and m.is_od:
+ self.num_od_layers += 1
+ self.widths.append(m.width)
+ self.od_layers.append(m)
+
+ def width_sampler(self) -> Generator: # pylint: disable=no-self-use
+ """Sample width."""
+ while True:
+ yield None
+
+ def layer_sampler(self) -> Module:
+ """Sample layer."""
+ while True:
+ for m in self.od_layers:
+ yield m
+
+ def __call__(self):
+ """Call sampler."""
+ if self.with_layer:
+ return next(self.width_samples), next(self.layer_samples)
+
+ return next(self.width_samples)
diff --git a/baselines/fjord/fjord/od/samplers/fixed_od.py b/baselines/fjord/fjord/od/samplers/fixed_od.py
new file mode 100644
index 000000000000..b90912a7b5c2
--- /dev/null
+++ b/baselines/fjord/fjord/od/samplers/fixed_od.py
@@ -0,0 +1,27 @@
+"""Ordered Dropout stochastic sampler."""
+from collections.abc import Generator
+from typing import List
+
+import numpy as np
+
+from .base_sampler import BaseSampler
+
+
+class ODSampler(BaseSampler):
+ """Implements OD sampling per layer up to p-max value.
+
+ :param p_s: list of p-values
+ :param max_p: maximum p-value
+ """
+
+ def __init__(self, p_s: List[float], max_p: float, *args, **kwargs) -> None:
+ super().__init__(*args, **kwargs)
+ self.p_s = np.array([p for p in p_s if p <= max_p])
+ self.max_p = max_p
+
+ def width_sampler(self) -> Generator:
+ """Sample width."""
+ while True:
+ p = np.random.choice(self.p_s)
+ for _ in range(self.num_od_layers):
+ yield p
diff --git a/baselines/fjord/fjord/server.py b/baselines/fjord/fjord/server.py
new file mode 100644
index 000000000000..d25e8f17156a
--- /dev/null
+++ b/baselines/fjord/fjord/server.py
@@ -0,0 +1,50 @@
+"""Global evaluation function."""
+from typing import Any, Dict, Optional, Tuple
+
+import flwr as fl
+import torch
+from torch.utils.data import DataLoader
+
+from .models import get_net, test
+from .utils.logger import Logger
+from .utils.utils import save_model, set_parameters
+
+
+def get_eval_fn(
+ args: Any, model_path: str, testloader: DataLoader, device: torch.device
+):
+ """Get evaluation function.
+
+ :param args: Arguments
+ :param model_path: Path to save the model
+ :param testloader: Test data loader
+ :param device: Device to be used
+ :return: Evaluation function
+ """
+
+ def evaluate(
+ server_round: int,
+ parameters: fl.common.NDArrays,
+ config: Dict[str, fl.common.Scalar], # pylint: disable=unused-argument
+ ) -> Optional[Tuple[float, Dict[str, fl.common.Scalar]]]:
+ if server_round and (server_round % args.evaluate_every == 0):
+ net = get_net(args.model, args.p_s, device)
+ set_parameters(net, parameters)
+ # Update model with the latest parameters
+ losses, accuracies = test(net, testloader, args.p_s)
+ avg_loss = sum(losses) / len(losses)
+ for p, loss, accuracy in zip(args.p_s, losses, accuracies):
+ Logger.get().info(
+ f"Server-side evaluation (global round={server_round})"
+ f" {p=}: {loss=} / {accuracy=}"
+ )
+ save_model(net, model_path)
+
+ return avg_loss, {
+ f"Accuracy[{p}]": acc for p, acc in zip(args.p_s, accuracies)
+ }
+
+ Logger.get().debug(f"Evaluation skipped for global round={server_round}.")
+ return float("inf"), {"accuracy": "None"}
+
+ return evaluate
diff --git a/baselines/fjord/fjord/strategy.py b/baselines/fjord/fjord/strategy.py
new file mode 100644
index 000000000000..d3ec99a419bd
--- /dev/null
+++ b/baselines/fjord/fjord/strategy.py
@@ -0,0 +1,235 @@
+"""FjORD strategy."""
+from copy import deepcopy
+from functools import reduce
+from typing import Dict, List, Optional, Tuple, Union
+
+import numpy as np
+from flwr.common import (
+ FitRes,
+ Metrics,
+ NDArrays,
+ Parameters,
+ Scalar,
+ ndarrays_to_parameters,
+ parameters_to_ndarrays,
+)
+from flwr.server.client_proxy import ClientProxy
+from flwr.server.strategy import FedAvg
+
+from .client import FJORD_CONFIG_TYPE
+from .utils.logger import Logger
+
+
+# Define metric aggregation function
+def weighted_average(metrics: List[Tuple[int, Metrics]]) -> Metrics:
+ """Aggregate using weighted average based on number of samples.
+
+ :param metrics: List of tuples (num_examples, metrics)
+ :return: Aggregated metrics
+ """
+ # Multiply accuracy of each client by number of examples used
+ accuracies = np.array([num_examples * m["accuracy"] for num_examples, m in metrics])
+ examples = np.array([num_examples for num_examples, _ in metrics])
+
+ # Aggregate and return custom metric (weighted average)
+ return {"accuracy": accuracies.sum() / examples.sum()}
+
+
+def get_p_layer_updates(
+ p: float,
+ layer_updates: List[np.ndarray],
+ num_examples: List[int],
+ p_max_s: List[float],
+) -> Tuple[List[np.ndarray], int]:
+ """Get layer updates for given p width.
+
+ :param p: p-value
+ :param layer_updates: list of layer updates from clients
+ :param num_examples: list of number of examples from clients
+ :param p_max_s: list of p_max values from clients
+ """
+ # get layers that were updated for given p
+ # i.e., for the clients with p_max >= p
+ layer_updates_p = [
+ layer_update
+ for p_max, layer_update in zip(p_max_s, layer_updates)
+ if p_max >= p
+ ]
+ num_examples_p = sum(n for p_max, n in zip(p_max_s, num_examples) if p_max >= p)
+ return layer_updates_p, num_examples_p
+
+
+def fjord_average( # pylint: disable=too-many-arguments
+ i: int,
+ layer_updates: List[np.ndarray],
+ num_examples: List[int],
+ p_max_s: List[float],
+ p_s: List[float],
+ fjord_config: FJORD_CONFIG_TYPE,
+ original_parameters: List[np.ndarray],
+) -> np.ndarray:
+ """Compute average per layer for given updates.
+
+ :param i: index of the layer
+ :param layer_updates: list of layer updates from clients
+ :param num_examples: list of number of examples from clients
+ :param p_max_s: list of p_max values from clients
+ :param p_s: list of p values
+ :param fjord_config: fjord config
+ :param original_parameters: original model parameters
+ :return: average of layer
+ """
+ # if no client updated the given part of the model,
+ # reuse previous parameters
+ update = deepcopy(original_parameters[i])
+
+ # BatchNorm2d layers, only average over the p_max_s
+ # that are greater than corresponding p of the layer
+ # i.e., only update the layers that were updated
+ if fjord_config["layer_p"][i] is not None:
+ p = fjord_config["layer_p"][i]
+ layer_updates_p, num_examples_p = get_p_layer_updates(
+ p, layer_updates, num_examples, p_max_s
+ )
+ if len(layer_updates_p) == 0:
+ return update
+
+ assert num_examples_p > 0
+ return reduce(np.add, layer_updates_p) / num_examples_p
+ if fjord_config["layer"][i] in ["ODLinear", "ODConv2d", "ODBatchNorm2d"]:
+ # perform nested updates
+ for p in p_s[::-1]:
+ layer_updates_p, num_examples_p = get_p_layer_updates(
+ p, layer_updates, num_examples, p_max_s
+ )
+ if len(layer_updates_p) == 0:
+ continue
+ in_dim = (
+ int(fjord_config[p][i]["in_dim"])
+ if fjord_config[p][i]["in_dim"]
+ else None
+ )
+ out_dim = (
+ int(fjord_config[p][i]["out_dim"])
+ if fjord_config[p][i]["out_dim"]
+ else None
+ )
+ assert num_examples_p > 0
+ # check whether the parameter to update is bias or weight
+ if len(update.shape) == 1:
+ # bias or ODBatchNorm2d
+ layer_updates_p = [
+ layer_update[:out_dim] for layer_update in layer_updates_p
+ ]
+ update[:out_dim] = reduce(np.add, layer_updates_p) / num_examples_p
+ else:
+ # weight
+ layer_updates_p = [
+ layer_update[:out_dim, :in_dim] for layer_update in layer_updates_p
+ ]
+ update[:out_dim, :in_dim] = (
+ reduce(np.add, layer_updates_p) / num_examples_p
+ )
+ return update
+
+ raise ValueError(f"Unsupported layer {fjord_config['layer'][i]}")
+
+
+def aggregate(
+ results: List[Tuple[NDArrays, int, float, List[float], FJORD_CONFIG_TYPE]],
+ original_parameters,
+) -> NDArrays:
+ """Compute weighted average.
+
+ :param results: list of tuples (layer_updates, num_examples, p_max, p_s)
+ :param original_parameters: original model parameters
+ :return: weighted average of layer updates
+ """
+ # Create a list of weights, each multiplied
+ # by the related number of examples
+ weights = [
+ [param * num_examples for param in params]
+ for params, num_examples, _, _, _ in results
+ ]
+ p_max_s = [p_max for _, _, p_max, _, _ in results]
+
+ # Calculate the total number of examples used during training
+ num_examples = [num_examples for _, num_examples, _, _, _ in results]
+ p_s = results[0][3]
+ fjord_config = results[0][4]
+
+ weights_prime: NDArrays = [
+ fjord_average(
+ i,
+ layer_updates,
+ num_examples,
+ p_max_s,
+ p_s,
+ fjord_config,
+ original_parameters,
+ )
+ for i, layer_updates in enumerate(zip(*weights))
+ ]
+ return weights_prime
+
+
+class FjORDFedAVG(FedAvg):
+ """FedAvg strategy with FjORD aggregation."""
+
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+
+ def aggregate_fit(
+ self,
+ server_round: int,
+ results: List[Tuple[ClientProxy, FitRes]],
+ failures: List[Union[Tuple[ClientProxy, FitRes], BaseException]],
+ ) -> Tuple[Optional[Parameters], Dict[str, Scalar]]:
+ """Aggregate fit results using weighted average."""
+ if not results:
+ return None, {}
+ # Do not aggregate if there are failures and failures are not accepted
+ if not self.accept_failures and failures:
+ return None, {}
+
+ Logger.get().info(f"Aggregating for global round {server_round}")
+ # Convert results
+ weights_results: List[
+ Tuple[NDArrays, int, float, List[float], FJORD_CONFIG_TYPE]
+ ] = [
+ ( # type: ignore
+ parameters_to_ndarrays(fit_res.parameters),
+ fit_res.num_examples,
+ fit_res.metrics["max_p"],
+ fit_res.metrics["p_s"],
+ fit_res.metrics["fjord_config"],
+ )
+ for _, fit_res in results
+ ]
+
+ p_max_values_str = ", ".join([str(val[2]) for val in weights_results])
+ Logger.get().info(f"\t - p_max values: {p_max_values_str}")
+
+ # all clients start with the same model
+ for _, fit_res in results:
+ original_parameters = fit_res.metrics["original_parameters"]
+ break
+
+ training_losses_str = ", ".join(
+ [str(fit_res.metrics["loss"]) for _, fit_res in results]
+ )
+ Logger.get().info(f"\t - train losses: {training_losses_str}")
+
+ agg = aggregate(weights_results, original_parameters)
+
+ parameters_aggregated = ndarrays_to_parameters(agg)
+
+ # Aggregate custom metrics if aggregation fn was provided
+ metrics_aggregated = {}
+ if self.fit_metrics_aggregation_fn:
+ fit_metrics = [(res.num_examples, res.metrics) for _, res in results]
+ metrics_aggregated = self.fit_metrics_aggregation_fn(fit_metrics)
+ elif server_round == 1: # Only log this warning once
+ Logger.get().warn("No fit_metrics_aggregation_fn provided")
+
+ return parameters_aggregated, metrics_aggregated
diff --git a/baselines/fjord/fjord/utils.py b/baselines/fjord/fjord/utils.py
new file mode 100644
index 000000000000..77b28f3d68ad
--- /dev/null
+++ b/baselines/fjord/fjord/utils.py
@@ -0,0 +1 @@
+"""Find the utils in the utils/ directory."""
diff --git a/baselines/fjord/fjord/utils/__init__.py b/baselines/fjord/fjord/utils/__init__.py
new file mode 100644
index 000000000000..46856dadddd5
--- /dev/null
+++ b/baselines/fjord/fjord/utils/__init__.py
@@ -0,0 +1 @@
+"""Utility functions for Fjord."""
diff --git a/baselines/fjord/fjord/utils/logger.py b/baselines/fjord/fjord/utils/logger.py
new file mode 100644
index 000000000000..b0eb2194bfef
--- /dev/null
+++ b/baselines/fjord/fjord/utils/logger.py
@@ -0,0 +1,129 @@
+"""Logger functionality."""
+import logging
+
+import coloredlogs
+
+
+class Logger:
+ """Logger class to be used by all modules in the project."""
+
+ log_format = (
+ "[%(asctime)s] (%(process)s) {%(filename)s:%(lineno)d}"
+ " %(levelname)s - %(message)s"
+ )
+ log_level = None
+
+ @classmethod
+ def setup_logging(cls, loglevel="INFO", logfile=""):
+ """Stateful setup of the logging infrastructure.
+
+ :param loglevel: log level to be used
+ :param logfile: file to log to
+ """
+ cls.registered_loggers = {}
+ cls.log_level = loglevel
+ numeric_level = getattr(logging, loglevel.upper(), None)
+
+ if not isinstance(numeric_level, int):
+ raise ValueError(f"Invalid log level: {loglevel}")
+ if logfile:
+ logging.basicConfig(
+ handlers=[logging.FileHandler(logfile), logging.StreamHandler()],
+ level=numeric_level,
+ format=cls.log_format,
+ datefmt="%Y-%m-%d %H:%M:%S",
+ )
+ else:
+ logging.basicConfig(
+ level=numeric_level,
+ format=cls.log_format,
+ datefmt="%Y-%m-%d %H:%M:%S",
+ )
+
+ @classmethod
+ def get(cls, logger_name="default"):
+ """Get logger instance.
+
+ :param logger_name: name of the logger
+ :return: logger instance
+ """
+ if logger_name in cls.registered_loggers:
+ return cls.registered_loggers[logger_name]
+
+ return cls(logger_name)
+
+ def __init__(self, logger_name="default"):
+ """Initialise logger not previously registered.
+
+ :param logger_name: name of the logger
+ """
+ if logger_name in self.registered_loggers:
+ raise ValueError(
+ f"Logger {logger_name} already exists. "
+ f'Call with Logger.get("{logger_name}")'
+ )
+
+ self.name = logger_name
+ self.logger = logging.getLogger(self.name)
+ self.registered_loggers[self.name] = self.logger
+ coloredlogs.install(
+ level=self.log_level,
+ logger=self.logger,
+ fmt=self.log_format,
+ datefmt="%Y-%m-%d %H:%M:%S",
+ )
+
+ self.warn = self.warning
+
+ def log(self, loglevel, msg):
+ """Log message.
+
+ :param loglevel: log level to be used
+ :param msg: message to be logged
+ """
+ loglevel = loglevel.upper()
+ if loglevel == "DEBUG":
+ self.logger.debug(msg)
+ elif loglevel == "INFO":
+ self.logger.info(msg)
+ elif loglevel == "WARNING":
+ self.logger.warning(msg)
+ elif loglevel == "ERROR":
+ self.logger.error(msg)
+ elif loglevel == "CRITICAL":
+ self.logger.critical(msg)
+
+ def debug(self, msg):
+ """Log debug message.
+
+ :param msg: message to be logged
+ """
+ self.log("debug", msg)
+
+ def info(self, msg):
+ """Log info message.
+
+ :param msg: message to be logged
+ """
+ self.log("info", msg)
+
+ def warning(self, msg):
+ """Log warning message.
+
+ :param msg: message to be logged
+ """
+ self.log("warning", msg)
+
+ def error(self, msg):
+ """Log error message.
+
+ :param msg: message to be logged
+ """
+ self.log("error", msg)
+
+ def critical(self, msg):
+ """Log critical message.
+
+ :param msg: message to be logged
+ """
+ self.log("critical", msg)
diff --git a/baselines/fjord/fjord/utils/utils.py b/baselines/fjord/fjord/utils/utils.py
new file mode 100644
index 000000000000..3a1a327dd555
--- /dev/null
+++ b/baselines/fjord/fjord/utils/utils.py
@@ -0,0 +1,52 @@
+"""Utility functions for fjord."""
+import os
+from typing import List, Optional, OrderedDict
+
+import numpy as np
+import torch
+from torch.nn import Module
+
+from .logger import Logger
+
+
+def get_parameters(net: Module) -> List[np.ndarray]:
+ """Get statedict parameters as a list of numpy arrays.
+
+ :param net: PyTorch model
+ :return: List of numpy arrays
+ """
+ return [val.cpu().numpy() for _, val in net.state_dict().items()]
+
+
+def set_parameters(net: Module, parameters: List[np.ndarray]) -> None:
+ """Load parameters into PyTorch model.
+
+ :param net: PyTorch model
+ :param parameters: List of numpy arrays
+ """
+ params_dict = zip(net.state_dict().keys(), parameters)
+ state_dict = OrderedDict({k: torch.Tensor(v) for k, v in params_dict})
+ net.load_state_dict(state_dict, strict=True)
+
+
+def save_model(
+ model: torch.nn.Module,
+ model_path: str,
+ is_best: bool = False,
+ cid: Optional[int] = None,
+) -> None:
+ """Checkpoint model.
+
+ :param model: model to be saved
+ :param model_path: path to save the model
+ :param is_best: whether this is the best model
+ :param cid: client id
+ """
+ suffix = "best" if is_best else "last"
+ if cid:
+ suffix += f"_{cid}"
+ filename = os.path.join(model_path, f"model_{suffix}.checkpoint")
+ Logger.get().info(f"Persisting model in {filename}")
+ if not os.path.isdir(model_path):
+ os.makedirs(model_path)
+ torch.save(model.state_dict(), filename)
diff --git a/baselines/fjord/notebooks/visualise.ipynb b/baselines/fjord/notebooks/visualise.ipynb
new file mode 100644
index 000000000000..04f9a7f768ec
--- /dev/null
+++ b/baselines/fjord/notebooks/visualise.ipynb
@@ -0,0 +1,277 @@
+{
+ "cells": [
+ {
+ "cell_type": "code",
+ "execution_count": 1,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "import re\n",
+ "import os\n",
+ "import glob\n",
+ "\n",
+ "import pandas as pd\n",
+ "import numpy as np\n",
+ "import matplotlib.pyplot as plt\n",
+ "\n",
+ "%matplotlib inline"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": 2,
+ "metadata": {},
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "['2023-09-23:12-53-16', '2023-09-23:11-25-43', '2023-09-23:13-32-56', '2023-09-23:14-20-22', '2023-09-23:12-06-14', '2023-09-23:15-00-26']\n"
+ ]
+ }
+ ],
+ "source": [
+ "log_root = \"../runs/best_config\"\n",
+ "\n",
+ "filenames = [os.path.basename(f) for f in glob.glob(os.path.join(log_root, \"*\"))]\n",
+ "print(filenames)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": 3,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "setups = {}\n",
+ "for f in filenames:\n",
+ " fq = os.path.join(log_root, f, \"run.log\")\n",
+ " with open(fq, \"r\") as fr:\n",
+ " s = fr.readlines()\n",
+ " # get CLI params\n",
+ " args_str = \"\\n\".join(s[:100])\n",
+ " manual_seed = re.search(r\"manual_seed: (\\d+)\", args_str).group(1)\n",
+ " knowledge_distillation = re.search(\n",
+ " r\"knowledge_distillation: (\\w+)\", args_str\n",
+ " ).group(1)\n",
+ " knowledge_distillation = \"kd\" if knowledge_distillation == \"true\" else \"nokd\"\n",
+ " client_selection = re.search(r\"client_selection: (\\w+)\", args_str).group(1)\n",
+ "\n",
+ " # get evaluation results\n",
+ " eval_timeline = []\n",
+ " eval_regex = r\".*Server-side evaluation \\(global round=(\\d+)\\) p=(\\d\\.\\d+): loss=(\\d+\\.\\d+) / accuracy=(\\d+\\.\\d+)\"\n",
+ " for line in s:\n",
+ " if re.match(eval_regex, line):\n",
+ " global_round, p, loss, accuracy = re.match(eval_regex, line).groups()\n",
+ " global_round, p, loss, accuracy = (\n",
+ " int(global_round),\n",
+ " float(p),\n",
+ " float(loss),\n",
+ " float(accuracy),\n",
+ " )\n",
+ " eval_timeline.append(\n",
+ " {\n",
+ " \"global_round\": global_round,\n",
+ " \"p\": p,\n",
+ " \"loss\": loss,\n",
+ " \"accuracy\": accuracy,\n",
+ " }\n",
+ " )\n",
+ "\n",
+ " setups[\n",
+ " f\"{client_selection}_{knowledge_distillation}_{manual_seed}\"\n",
+ " ] = eval_timeline"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": 4,
+ "metadata": {},
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ " global_round p loss accuracy kd seed client_selection\n",
+ "0 10 0.2 1.879755 0.2835 False 124 random\n",
+ "1 10 0.4 1.863002 0.3122 False 124 random\n",
+ "2 10 0.6 1.828429 0.3165 False 124 random\n",
+ "3 10 0.8 1.885398 0.2739 False 124 random\n",
+ "4 10 1.0 1.943324 0.2384 False 124 random\n"
+ ]
+ }
+ ],
+ "source": [
+ "dfs = []\n",
+ "for k, v in setups.items():\n",
+ " df = pd.DataFrame(v)\n",
+ " client_selection, kd, seed = k.split(\"_\")\n",
+ " df[\"kd\"] = False if kd == \"nokd\" else True\n",
+ " df[\"seed\"] = seed\n",
+ " df[\"client_selection\"] = client_selection\n",
+ " dfs.append(df)\n",
+ "df = pd.concat(dfs)\n",
+ "print(df.head())"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": 5,
+ "metadata": {},
+ "outputs": [
+ {
+ "data": {
+ "image/png": "",
+ "text/plain": [
+ ""
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ }
+ ],
+ "source": [
+ "grouped_df = df.groupby([\"kd\", \"global_round\", \"p\"])\n",
+ "df_mean = grouped_df[[\"loss\", \"accuracy\"]].mean()\n",
+ "df_std = grouped_df[[\"loss\", \"accuracy\"]].std()\n",
+ "\n",
+ "df_plot = df_mean.merge(\n",
+ " df_std, left_index=True, right_index=True, suffixes=(\"_mean\", \"_std\")\n",
+ ")\n",
+ "df_plot = df_plot.loc[:, 500, :]\n",
+ "grouped_df = df_plot.reset_index().groupby(\"kd\")\n",
+ "\n",
+ "plt.figure(figsize=(10, 4))\n",
+ "for i, (group_name, group_data) in enumerate(grouped_df):\n",
+ " label = \"FjORD w/ KD\" if group_name else \"FjORD\"\n",
+ " plt.plot(group_data.p, group_data.accuracy_mean * 100, label=label, marker=\"x\")\n",
+ " plt.fill_between(\n",
+ " group_data.p,\n",
+ " (group_data.accuracy_mean - group_data.accuracy_std) * 100,\n",
+ " (group_data.accuracy_mean + group_data.accuracy_std) * 100,\n",
+ " alpha=0.2,\n",
+ " )\n",
+ "\n",
+ "plt.legend()\n",
+ "plt.grid()\n",
+ "plt.title(\"ResNet18 - CIFAR10 - 500 global rounds\")\n",
+ "plt.xlabel(\"Submodel (p-value)\")\n",
+ "plt.ylabel(\"Accuracy (%)\")\n",
+ "plt.xticks(np.linspace(0.2, 1, 5))\n",
+ "\n",
+ "plt.savefig(\n",
+ " \"../_static/resnet18_cifar10_500_global_rounds_acc_pvalues.png\",\n",
+ " dpi=300,\n",
+ " bbox_inches=\"tight\",\n",
+ ")"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": 8,
+ "metadata": {},
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "False\n",
+ "True\n"
+ ]
+ },
+ {
+ "data": {
+ "image/png": "",
+ "text/plain": [
+ ""
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "data": {
+ "image/png": "",
+ "text/plain": [
+ ""
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ }
+ ],
+ "source": [
+ "grouped_df = df.groupby([\"kd\", \"global_round\", \"p\"])\n",
+ "df_mean = grouped_df[[\"loss\", \"accuracy\"]].mean()\n",
+ "df_std = grouped_df[[\"loss\", \"accuracy\"]].std()\n",
+ "\n",
+ "df_plot = df_mean.merge(\n",
+ " df_std, left_index=True, right_index=True, suffixes=(\"_mean\", \"_std\")\n",
+ ")\n",
+ "grouped_df = df_plot.reset_index().groupby([\"kd\"])\n",
+ "\n",
+ "for i, (group_name, group_data) in enumerate(grouped_df):\n",
+ " gd = group_data.groupby(\"p\")\n",
+ " plt.figure(figsize=(10, 4))\n",
+ " title = \"FjORD w/ KD\" if bool(group_name) else \"FjORD\"\n",
+ " filename_suffix = \"fjord_kd\" if bool(group_name) else \"fjord\"\n",
+ " plt.title(f\"ResNet18 - CIFAR10 - {title}\")\n",
+ " colors = plt.cm.viridis(np.linspace(0, 1, len(gd)))\n",
+ " for j, (p, p_data) in enumerate(gd):\n",
+ " plt.plot(\n",
+ " p_data[\"global_round\"],\n",
+ " p_data[\"loss_mean\"],\n",
+ " color=colors[j],\n",
+ " alpha=0.8,\n",
+ " label=f\"p={p}\",\n",
+ " )\n",
+ " plt.fill_between(\n",
+ " p_data[\"global_round\"],\n",
+ " p_data[\"loss_mean\"] - p_data[\"loss_std\"],\n",
+ " p_data[\"loss_mean\"] + p_data[\"loss_std\"],\n",
+ " alpha=0.1,\n",
+ " color=colors[j],\n",
+ " )\n",
+ " plt.xlabel(\"Global round\")\n",
+ " plt.ylabel(\"Loss\")\n",
+ " plt.legend()\n",
+ " plt.grid()\n",
+ "\n",
+ " plt.savefig(\n",
+ " f\"../_static/resnet18_cifar10_{filename_suffix}_convergence.png\",\n",
+ " dpi=300,\n",
+ " bbox_inches=\"tight\",\n",
+ " )"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": []
+ }
+ ],
+ "metadata": {
+ "kernelspec": {
+ "display_name": "fjord",
+ "language": "python",
+ "name": "python3"
+ },
+ "language_info": {
+ "codemirror_mode": {
+ "name": "ipython",
+ "version": 3
+ },
+ "file_extension": ".py",
+ "mimetype": "text/x-python",
+ "name": "python",
+ "nbconvert_exporter": "python",
+ "pygments_lexer": "ipython3",
+ "version": "3.9.12"
+ },
+ "orig_nbformat": 4
+ },
+ "nbformat": 4,
+ "nbformat_minor": 2
+}
diff --git a/baselines/fjord/pyproject.toml b/baselines/fjord/pyproject.toml
new file mode 100644
index 000000000000..d8a9ae307d7c
--- /dev/null
+++ b/baselines/fjord/pyproject.toml
@@ -0,0 +1,146 @@
+[build-system]
+requires = ["poetry-core>=1.4.0"]
+build-backend = "poetry.masonry.api"
+
+[tool.poetry]
+name = "fjord"
+version = "1.0.0"
+description = "FjORD implementation of Federated Ordered Dropout in Flower"
+license = "Apache-2.0"
+authors = ["Steve Laskaridis ", "Samuel Horvath "]
+readme = "README.md"
+homepage = "https://flower.dev"
+repository = "https://github.com/adap/flower"
+documentation = "https://flower.dev"
+classifiers = [
+ "Development Status :: 3 - Alpha",
+ "Intended Audience :: Developers",
+ "Intended Audience :: Science/Research",
+ "License :: OSI Approved :: Apache Software License",
+ "Operating System :: MacOS :: MacOS X",
+ "Operating System :: POSIX :: Linux",
+ "Programming Language :: Python",
+ "Programming Language :: Python :: 3",
+ "Programming Language :: Python :: 3 :: Only",
+ "Programming Language :: Python :: 3.8",
+ "Programming Language :: Python :: 3.9",
+ "Programming Language :: Python :: 3.10",
+ "Programming Language :: Python :: 3.11",
+ "Programming Language :: Python :: Implementation :: CPython",
+ "Topic :: Scientific/Engineering",
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
+ "Topic :: Scientific/Engineering :: Mathematics",
+ "Topic :: Software Development",
+ "Topic :: Software Development :: Libraries",
+ "Topic :: Software Development :: Libraries :: Python Modules",
+ "Typing :: Typed",
+]
+
+[tool.poetry.dependencies]
+python = ">=3.10.0, <3.11.0"
+flwr = { extras = ["simulation"], version = "1.5.0" }
+hydra-core = "1.3.2"
+matplotlib = "3.7.1"
+coloredlogs = "15.0.1"
+omegaconf = "2.3.0"
+tqdm = "4.65.0"
+torch = { url = "https://download.pytorch.org/whl/cu117/torch-2.0.1%2Bcu117-cp310-cp310-linux_x86_64.whl"}
+torchvision = { url = "https://download.pytorch.org/whl/cu117/torchvision-0.15.2%2Bcu117-cp310-cp310-linux_x86_64.whl"}
+
+
+
+[tool.poetry.dev-dependencies]
+isort = "==5.11.5"
+black = "==23.1.0"
+docformatter = "==1.5.1"
+mypy = "==1.4.1"
+pylint = "==2.8.2"
+flake8 = "==3.9.2"
+pytest = "==6.2.4"
+pytest-watch = "==4.2.0"
+ruff = "==0.0.272"
+types-requests = "==2.27.7"
+
+[tool.isort]
+line_length = 88
+indent = " "
+multi_line_output = 3
+include_trailing_comma = true
+force_grid_wrap = 0
+use_parentheses = true
+
+[tool.black]
+line-length = 88
+target-version = ["py38", "py39", "py310", "py311"]
+
+[tool.pytest.ini_options]
+minversion = "6.2"
+addopts = "-qq"
+testpaths = [
+ "flwr_baselines",
+]
+
+[tool.mypy]
+ignore_missing_imports = true
+strict = false
+plugins = "numpy.typing.mypy_plugin"
+
+[tool.pylint."MESSAGES CONTROL"]
+disable = "duplicate-code,too-few-public-methods,useless-import-alias"
+good-names = "i,j,k,_,x,y,X,Y,fl,lr,p,p_,bn,NUM_CLIENTS,od,m,g,ResNet18,FJORD_CONFIG_TYPE"
+signature-mutators="hydra.main.main"
+
+[tool.pylint.typecheck]
+generated-members="numpy.*, torch.*, tensorflow.*"
+
+
+[[tool.mypy.overrides]]
+module = [
+ "importlib.metadata.*",
+ "importlib_metadata.*",
+]
+follow_imports = "skip"
+follow_imports_for_stubs = true
+disallow_untyped_calls = false
+
+[[tool.mypy.overrides]]
+module = "torch.*"
+follow_imports = "skip"
+follow_imports_for_stubs = true
+
+[tool.docformatter]
+wrap-summaries = 88
+wrap-descriptions = 88
+
+[tool.ruff]
+target-version = "py38"
+line-length = 88
+select = ["D", "E", "F", "W", "B", "ISC", "C4"]
+fixable = ["D", "E", "F", "W", "B", "ISC", "C4"]
+ignore = ["B024", "B027"]
+exclude = [
+ ".bzr",
+ ".direnv",
+ ".eggs",
+ ".git",
+ ".hg",
+ ".mypy_cache",
+ ".nox",
+ ".pants.d",
+ ".pytype",
+ ".ruff_cache",
+ ".svn",
+ ".tox",
+ ".venv",
+ "__pypackages__",
+ "_build",
+ "buck-out",
+ "build",
+ "dist",
+ "node_modules",
+ "venv",
+ "proto",
+]
+
+[tool.ruff.pydocstyle]
+convention = "numpy"
diff --git a/baselines/fjord/requirements.txt b/baselines/fjord/requirements.txt
new file mode 100644
index 000000000000..35583b1a45c4
--- /dev/null
+++ b/baselines/fjord/requirements.txt
@@ -0,0 +1,8 @@
+coloredlogs==15.0.1
+hydra-core==1.3.2
+flwr==1.5.0
+omegaconf==2.3.0
+ray==2.6.3
+torch==2.0.1
+torchvision==0.15.2
+tqdm==4.65.0
diff --git a/baselines/fjord/scripts/run.sh b/baselines/fjord/scripts/run.sh
new file mode 100755
index 000000000000..ab4571724e2f
--- /dev/null
+++ b/baselines/fjord/scripts/run.sh
@@ -0,0 +1,16 @@
+#!/bin/bash
+
+RUN_LOG_DIR=${RUN_LOG_DIR:-"exp_logs"}
+
+pushd ../
+mkdir -p $RUN_LOG_DIR
+for seed in 123 124 125; do
+ echo "Running seed $seed"
+
+ echo "Running without KD ..."
+ poetry run python -m fjord.main ++manual_seed=$seed |& tee $RUN_LOG_DIR/wout_kd_$seed.log
+
+ echo "Running with KD ..."
+ poetry run python -m fjord.main +train_mode=fjord_kd ++manual_seed=$seed |& tee $RUN_LOG_DIR/w_kd_$seed.log
+done
+popd
diff --git a/baselines/fjord/setup.py b/baselines/fjord/setup.py
new file mode 100644
index 000000000000..aa09948c34fc
--- /dev/null
+++ b/baselines/fjord/setup.py
@@ -0,0 +1,14 @@
+"""Setup fjord package."""
+from setuptools import find_packages, setup
+
+VERSION = "0.1.0"
+DESCRIPTION = "FjORD Flwr package"
+LONG_DESCRIPTION = "Implementation of FjORD as a flwr baseline"
+
+setup(
+ name="fjord",
+ version=VERSION,
+ description=DESCRIPTION,
+ long_description=LONG_DESCRIPTION,
+ packages=find_packages(),
+)
diff --git a/baselines/flwr_baselines/dev/bootstrap.sh b/baselines/flwr_baselines/dev/bootstrap.sh
index eaa3a0bb046b..0bc322edc0de 100755
--- a/baselines/flwr_baselines/dev/bootstrap.sh
+++ b/baselines/flwr_baselines/dev/bootstrap.sh
@@ -6,8 +6,8 @@ cd "$( cd "$( dirname "${BASH_SOURCE[0]}" )" >/dev/null 2>&1 && pwd )"/../
./dev/rm-caches.sh
# Upgrade/install spcific versions of `pip`, `setuptools`, and `poetry`
-python -m pip install -U pip==23.1.2
-python -m pip install -U setuptools==68.0.0
+python -m pip install -U pip==23.3.1
+python -m pip install -U setuptools==68.2.2
python -m pip install -U poetry==1.5.1
# Use `poetry` to install project dependencies
diff --git a/baselines/moon/LICENSE b/baselines/moon/LICENSE
new file mode 100644
index 000000000000..d64569567334
--- /dev/null
+++ b/baselines/moon/LICENSE
@@ -0,0 +1,202 @@
+
+ Apache License
+ Version 2.0, January 2004
+ http://www.apache.org/licenses/
+
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
+
+ 1. Definitions.
+
+ "License" shall mean the terms and conditions for use, reproduction,
+ and distribution as defined by Sections 1 through 9 of this document.
+
+ "Licensor" shall mean the copyright owner or entity authorized by
+ the copyright owner that is granting the License.
+
+ "Legal Entity" shall mean the union of the acting entity and all
+ other entities that control, are controlled by, or are under common
+ control with that entity. For the purposes of this definition,
+ "control" means (i) the power, direct or indirect, to cause the
+ direction or management of such entity, whether by contract or
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
+ outstanding shares, or (iii) beneficial ownership of such entity.
+
+ "You" (or "Your") shall mean an individual or Legal Entity
+ exercising permissions granted by this License.
+
+ "Source" form shall mean the preferred form for making modifications,
+ including but not limited to software source code, documentation
+ source, and configuration files.
+
+ "Object" form shall mean any form resulting from mechanical
+ transformation or translation of a Source form, including but
+ not limited to compiled object code, generated documentation,
+ and conversions to other media types.
+
+ "Work" shall mean the work of authorship, whether in Source or
+ Object form, made available under the License, as indicated by a
+ copyright notice that is included in or attached to the work
+ (an example is provided in the Appendix below).
+
+ "Derivative Works" shall mean any work, whether in Source or Object
+ form, that is based on (or derived from) the Work and for which the
+ editorial revisions, annotations, elaborations, or other modifications
+ represent, as a whole, an original work of authorship. For the purposes
+ of this License, Derivative Works shall not include works that remain
+ separable from, or merely link (or bind by name) to the interfaces of,
+ the Work and Derivative Works thereof.
+
+ "Contribution" shall mean any work of authorship, including
+ the original version of the Work and any modifications or additions
+ to that Work or Derivative Works thereof, that is intentionally
+ submitted to Licensor for inclusion in the Work by the copyright owner
+ or by an individual or Legal Entity authorized to submit on behalf of
+ the copyright owner. For the purposes of this definition, "submitted"
+ means any form of electronic, verbal, or written communication sent
+ to the Licensor or its representatives, including but not limited to
+ communication on electronic mailing lists, source code control systems,
+ and issue tracking systems that are managed by, or on behalf of, the
+ Licensor for the purpose of discussing and improving the Work, but
+ excluding communication that is conspicuously marked or otherwise
+ designated in writing by the copyright owner as "Not a Contribution."
+
+ "Contributor" shall mean Licensor and any individual or Legal Entity
+ on behalf of whom a Contribution has been received by Licensor and
+ subsequently incorporated within the Work.
+
+ 2. Grant of Copyright License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ copyright license to reproduce, prepare Derivative Works of,
+ publicly display, publicly perform, sublicense, and distribute the
+ Work and such Derivative Works in Source or Object form.
+
+ 3. Grant of Patent License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ (except as stated in this section) patent license to make, have made,
+ use, offer to sell, sell, import, and otherwise transfer the Work,
+ where such license applies only to those patent claims licensable
+ by such Contributor that are necessarily infringed by their
+ Contribution(s) alone or by combination of their Contribution(s)
+ with the Work to which such Contribution(s) was submitted. If You
+ institute patent litigation against any entity (including a
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
+ or a Contribution incorporated within the Work constitutes direct
+ or contributory patent infringement, then any patent licenses
+ granted to You under this License for that Work shall terminate
+ as of the date such litigation is filed.
+
+ 4. Redistribution. You may reproduce and distribute copies of the
+ Work or Derivative Works thereof in any medium, with or without
+ modifications, and in Source or Object form, provided that You
+ meet the following conditions:
+
+ (a) You must give any other recipients of the Work or
+ Derivative Works a copy of this License; and
+
+ (b) You must cause any modified files to carry prominent notices
+ stating that You changed the files; and
+
+ (c) You must retain, in the Source form of any Derivative Works
+ that You distribute, all copyright, patent, trademark, and
+ attribution notices from the Source form of the Work,
+ excluding those notices that do not pertain to any part of
+ the Derivative Works; and
+
+ (d) If the Work includes a "NOTICE" text file as part of its
+ distribution, then any Derivative Works that You distribute must
+ include a readable copy of the attribution notices contained
+ within such NOTICE file, excluding those notices that do not
+ pertain to any part of the Derivative Works, in at least one
+ of the following places: within a NOTICE text file distributed
+ as part of the Derivative Works; within the Source form or
+ documentation, if provided along with the Derivative Works; or,
+ within a display generated by the Derivative Works, if and
+ wherever such third-party notices normally appear. The contents
+ of the NOTICE file are for informational purposes only and
+ do not modify the License. You may add Your own attribution
+ notices within Derivative Works that You distribute, alongside
+ or as an addendum to the NOTICE text from the Work, provided
+ that such additional attribution notices cannot be construed
+ as modifying the License.
+
+ You may add Your own copyright statement to Your modifications and
+ may provide additional or different license terms and conditions
+ for use, reproduction, or distribution of Your modifications, or
+ for any such Derivative Works as a whole, provided Your use,
+ reproduction, and distribution of the Work otherwise complies with
+ the conditions stated in this License.
+
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
+ any Contribution intentionally submitted for inclusion in the Work
+ by You to the Licensor shall be under the terms and conditions of
+ this License, without any additional terms or conditions.
+ Notwithstanding the above, nothing herein shall supersede or modify
+ the terms of any separate license agreement you may have executed
+ with Licensor regarding such Contributions.
+
+ 6. Trademarks. This License does not grant permission to use the trade
+ names, trademarks, service marks, or product names of the Licensor,
+ except as required for reasonable and customary use in describing the
+ origin of the Work and reproducing the content of the NOTICE file.
+
+ 7. Disclaimer of Warranty. Unless required by applicable law or
+ agreed to in writing, Licensor provides the Work (and each
+ Contributor provides its Contributions) on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
+ implied, including, without limitation, any warranties or conditions
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
+ PARTICULAR PURPOSE. You are solely responsible for determining the
+ appropriateness of using or redistributing the Work and assume any
+ risks associated with Your exercise of permissions under this License.
+
+ 8. Limitation of Liability. In no event and under no legal theory,
+ whether in tort (including negligence), contract, or otherwise,
+ unless required by applicable law (such as deliberate and grossly
+ negligent acts) or agreed to in writing, shall any Contributor be
+ liable to You for damages, including any direct, indirect, special,
+ incidental, or consequential damages of any character arising as a
+ result of this License or out of the use or inability to use the
+ Work (including but not limited to damages for loss of goodwill,
+ work stoppage, computer failure or malfunction, or any and all
+ other commercial damages or losses), even if such Contributor
+ has been advised of the possibility of such damages.
+
+ 9. Accepting Warranty or Additional Liability. While redistributing
+ the Work or Derivative Works thereof, You may choose to offer,
+ and charge a fee for, acceptance of support, warranty, indemnity,
+ or other liability obligations and/or rights consistent with this
+ License. However, in accepting such obligations, You may act only
+ on Your own behalf and on Your sole responsibility, not on behalf
+ of any other Contributor, and only if You agree to indemnify,
+ defend, and hold each Contributor harmless for any liability
+ incurred by, or claims asserted against, such Contributor by reason
+ of your accepting any such warranty or additional liability.
+
+ END OF TERMS AND CONDITIONS
+
+ APPENDIX: How to apply the Apache License to your work.
+
+ To apply the Apache License to your work, attach the following
+ boilerplate notice, with the fields enclosed by brackets "[]"
+ replaced with your own identifying information. (Don't include
+ the brackets!) The text should be enclosed in the appropriate
+ comment syntax for the file format. We also recommend that a
+ file or class name and description of purpose be included on the
+ same "printed page" as the copyright notice for easier
+ identification within third-party archives.
+
+ Copyright [yyyy] [name of copyright owner]
+
+ Licensed under the Apache License, Version 2.0 (the "License");
+ you may not use this file except in compliance with the License.
+ You may obtain a copy of the License at
+
+ http://www.apache.org/licenses/LICENSE-2.0
+
+ Unless required by applicable law or agreed to in writing, software
+ distributed under the License is distributed on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ See the License for the specific language governing permissions and
+ limitations under the License.
diff --git a/baselines/moon/README.md b/baselines/moon/README.md
new file mode 100644
index 000000000000..05ab4ef68469
--- /dev/null
+++ b/baselines/moon/README.md
@@ -0,0 +1,146 @@
+---
+title: Model-Contrastive Federated Learning
+url: https://arxiv.org/abs/2103.16257
+labels: [data heterogeneity, image classification, cross-silo, constrastive-learning]
+dataset: [CIFAR-10, CIFAR-100]
+---
+
+# Model-Contrastive Federated Learning
+
+> Note: If you use this baseline in your work, please remember to cite the original authors of the paper as well as the Flower paper.
+
+
+**Paper:** [arxiv.org/abs/2103.16257](https://arxiv.org/abs/2103.16257)
+
+**Authors:** Qinbin Li, Bingsheng He, Dawn Song
+
+**Abstract:** Federated learning enables multiple parties to collaboratively train a machine learning model without communicating their local data. A key challenge in federated learning is to handle the heterogeneity of local data distribution across parties. Although many studies have been proposed to address this challenge, we find that they fail to achieve high performance in image datasets with deep learning models. In this paper, we propose MOON: modelcontrastive federated learning. MOON is a simple and effective federated learning framework. The key idea of MOON is to utilize the similarity between model representations to correct the local training of individual parties, i.e., conducting contrastive learning in model-level. Our extensive experiments show that MOON significantly outperforms the other state-of-the-art federated learning algorithms on various image classification tasks.
+
+
+
+## About this baseline
+
+**What’s implemented:** The code in this directory replicates the experiments in *Model-Contrastive Federated Learning* (Li et al., 2021), which proposed the MOON algorithm. Concretely ,it replicates the results of MOON for CIFAR-10 and CIFAR-100 in Table 1.
+
+**Datasets:** CIFAR-10 and CIFAR-100
+
+**Hardware Setup:** The experiments are run on a server with 4x Intel Xeon Gold 6226R and 8x Nvidia GeForce RTX 3090. A machine with at least 1x 16GB GPU should be able to run the experiments in a reasonable time.
+
+****Contributors:**** [Qinbin Li](https://qinbinli.com)
+
+**Description:** MOON requires to compute the model-contrastive loss in local training, which requires access to the local model of the previous round (Lines 14-17 of Algorithm 1 of the paper). Since currently `FlowerClient` does not preserve the states when starting a new round, we store the local models into the specified `model.dir` in local training indexed by the client ID, which will be loaded to the corresponding client in the next round.
+
+## Experimental Setup
+
+**Task:** Image classification.
+
+**Model:** This directory implements two models as same as the paper:
+* A simple-CNN with a projection head for CIFAR-10
+* A ResNet-50 with a projection head for CIFAR-100.
+
+**Dataset:** This directory includes CIFAR-10 and CIFAR-100. They are partitioned in the same way as the paper. The settings are as follow:
+
+| Dataset | partitioning method |
+| :------ | :---: |
+| CIFAR-10 | Dirichlet with beta 0.5 |
+| CIFAR-100 | Dirichlet with beta 0.5 |
+
+
+**Training Hyperparameters:**
+
+| Description | Default Value |
+| ----------- | ----- |
+| number of clients | 10 |
+| number of local epochs | 10 |
+| fraction fit | 1.0 |
+| batch size | 64 |
+| learning rate | 0.01 |
+| mu | 1 |
+| temperature | 0.5 |
+| alg | moon |
+| seed | 0 |
+| service_device | cpu |
+| number of rounds | 100 |
+| client resources | {'num_cpus': 2.0, 'num_gpus': 0.0 }|
+
+## Environment Setup
+
+To construct the Python environment follow these steps:
+
+```bash
+# Set local python version via pyenv
+pyenv local 3.10.6
+# Then fix that for poetry
+poetry env use 3.10.6
+# Then install poetry env
+poetry install
+
+# Activate the environment
+poetry shell
+```
+
+
+## Running the Experiments
+
+First ensure you have activated your Poetry environment (execute `poetry shell` from this directory). To run MOON on CIFAR-10 (Table 1 of the paper), you should run:
+```bash
+python -m moon.main --config-name cifar10
+```
+
+To run MOON on CIFAR-100 (Table 1 of the paper), you should run:
+```bash
+python -m moon.main --config-name cifar100
+```
+
+
+You can also run FedProx on CIFAR-10:
+```bash
+python -m moon.main --config-name cifar10_fedprox
+```
+
+To run FedProx on CIFAR-100:
+```bash
+python -m moon.main --config-name cifar100_fedprox
+```
+
+## Expected Results
+
+You can find the output logs of a single run in this [link](https://drive.google.com/drive/folders/1YZEU2NcHWEHVyuJMlc1QvBSAvNMjH-aR?usp=share_link). After running the above commands, you can see the accuracy list at the end of the ouput, which is the test accuracy of the global model. For example, in one running, for CIFAR-10 with MOON, the accuracy after running 100 rounds is 0.7071.
+
+For CIFAR-10 with FedProx, the accuracy after running 100 rounds is 0.6852. For CIFAR100 with MOON, the accuracy after running 100 rounds is 0.6636. For CIFAR100 with FedProx, the accuracy after running 100 rounds is 0.6494. The results are summarized below:
+
+
+| | CIFAR-10 | CIFAR-100 |
+| ----------- | ----- | ----- |
+| MOON | 0.7071 | 0.6636 |
+| FedProx| 0.6852 | 0.6494 |
+
+### Figure 6
+You can find the curve comparing MOON and FedProx on CIFAR-10 and CIFAR-100 below.
+
+
+
+
+
+You can tune the hyperparameter `mu` for both MOON and FedProx by changing the configuration file in `conf`.
+
+### Figure 8(a)
+You can run the experiments in Figure 8 of the paper. To run MOON (`mu=10`) on CIFAR-100 with 50 clients (Figure 8(a) of the paper):
+```bash
+python -m moon.main --config-name cifar100_50clients
+```
+
+To run FedProx on CIFAR-100 with 50 clients (Figure 8(a) of the paper):
+```bash
+python -m moon.main --config-name cifar100_50clients_fedprox
+```
+
+
+You can find the curve presenting MOON and FedProx below.
+
+
+
+You may also run MOON on CIFAR-100 with 100 clients (Figure 8(b) of the paper):
+```bash
+python -m moon.main --config-name cifar100_100clients
+```
\ No newline at end of file
diff --git a/baselines/moon/_static/cifar100_50clients_moon_fedprox.png b/baselines/moon/_static/cifar100_50clients_moon_fedprox.png
new file mode 100644
index 000000000000..ecc1c99de230
Binary files /dev/null and b/baselines/moon/_static/cifar100_50clients_moon_fedprox.png differ
diff --git a/baselines/moon/_static/cifar100_moon_fedprox.png b/baselines/moon/_static/cifar100_moon_fedprox.png
new file mode 100644
index 000000000000..798d778cd1cc
Binary files /dev/null and b/baselines/moon/_static/cifar100_moon_fedprox.png differ
diff --git a/baselines/moon/_static/cifar10_moon_fedprox.png b/baselines/moon/_static/cifar10_moon_fedprox.png
new file mode 100644
index 000000000000..f5f18f1c08e9
Binary files /dev/null and b/baselines/moon/_static/cifar10_moon_fedprox.png differ
diff --git a/baselines/moon/moon/__init__.py b/baselines/moon/moon/__init__.py
new file mode 100644
index 000000000000..a5e567b59135
--- /dev/null
+++ b/baselines/moon/moon/__init__.py
@@ -0,0 +1 @@
+"""Template baseline package."""
diff --git a/baselines/moon/moon/client.py b/baselines/moon/moon/client.py
new file mode 100644
index 000000000000..4903140009b5
--- /dev/null
+++ b/baselines/moon/moon/client.py
@@ -0,0 +1,155 @@
+"""Define your client class and a function to construct such clients.
+
+Please overwrite `flwr.client.NumPyClient` or `flwr.client.Client` and create a function
+to instantiate your client.
+"""
+
+import copy
+import os
+from collections import OrderedDict
+from typing import Callable, Dict, List, Tuple
+
+import flwr as fl
+import torch
+from flwr.common.typing import NDArrays, Scalar
+from omegaconf import DictConfig
+from torch.utils.data import DataLoader
+
+from moon.models import init_net, train_fedprox, train_moon
+
+
+# pylint: disable=too-many-instance-attributes
+class FlowerClient(fl.client.NumPyClient):
+ """Standard Flower client for CNN training."""
+
+ def __init__(
+ self,
+ # net: torch.nn.Module,
+ net_id: int,
+ dataset: str,
+ model: str,
+ output_dim: int,
+ trainloader: DataLoader,
+ valloader: DataLoader,
+ device: torch.device,
+ num_epochs: int,
+ learning_rate: float,
+ mu: float,
+ temperature: float,
+ model_dir: str,
+ alg: str,
+ ): # pylint: disable=too-many-arguments
+ self.net = init_net(dataset, model, output_dim)
+ self.net_id = net_id
+ self.dataset = dataset
+ self.model = model
+ self.output_dim = output_dim
+ self.trainloader = trainloader
+ self.valloader = valloader
+ self.device = device
+ self.num_epochs = num_epochs
+ self.learning_rate = learning_rate
+ self.mu = mu # pylint: disable=invalid-name
+ self.temperature = temperature
+ self.model_dir = model_dir
+ self.alg = alg
+
+ def get_parameters(self, config: Dict[str, Scalar]) -> NDArrays:
+ """Return the parameters of the current net."""
+ return [val.cpu().numpy() for _, val in self.net.state_dict().items()]
+
+ def set_parameters(self, parameters: NDArrays) -> None:
+ """Change the parameters of the model using the given ones."""
+ params_dict = zip(self.net.state_dict().keys(), parameters)
+ state_dict = OrderedDict({k: torch.from_numpy(v) for k, v in params_dict})
+ self.net.load_state_dict(state_dict, strict=True)
+
+ def fit(
+ self, parameters: NDArrays, config: Dict[str, Scalar]
+ ) -> Tuple[NDArrays, int, Dict]:
+ """Implement distributed fit function for a given client."""
+ self.set_parameters(parameters)
+ prev_net = init_net(self.dataset, self.model, self.output_dim)
+ if not os.path.exists(os.path.join(self.model_dir, str(self.net_id))):
+ prev_net = copy.deepcopy(self.net)
+ else:
+ # load previous model from model_dir
+ prev_net.load_state_dict(
+ torch.load(
+ os.path.join(self.model_dir, str(self.net_id), "prev_net.pt")
+ )
+ )
+ global_net = init_net(self.dataset, self.model, self.output_dim)
+ global_net.load_state_dict(self.net.state_dict())
+ if self.alg == "moon":
+ train_moon(
+ self.net,
+ global_net,
+ prev_net,
+ self.trainloader,
+ self.num_epochs,
+ self.learning_rate,
+ self.mu,
+ self.temperature,
+ self.device,
+ )
+ elif self.alg == "fedprox":
+ train_fedprox(
+ self.net,
+ global_net,
+ self.trainloader,
+ self.num_epochs,
+ self.learning_rate,
+ self.mu,
+ self.device,
+ )
+ if not os.path.exists(os.path.join(self.model_dir, str(self.net_id))):
+ os.makedirs(os.path.join(self.model_dir, str(self.net_id)))
+ torch.save(
+ self.net.state_dict(),
+ os.path.join(self.model_dir, str(self.net_id), "prev_net.pt"),
+ )
+ return self.get_parameters({}), len(self.trainloader), {"is_straggler": False}
+
+ def evaluate(
+ self, parameters: NDArrays, config: Dict[str, Scalar]
+ ) -> Tuple[float, int, Dict]:
+ """Implement distributed evaluation for a given client."""
+ self.set_parameters(parameters)
+ # skip evaluation in the client-side
+ loss = 0.0
+ accuracy = 0.0
+ return float(loss), len(self.valloader), {"accuracy": float(accuracy)}
+
+
+def gen_client_fn(
+ trainloaders: List[DataLoader],
+ testloaders: List[DataLoader],
+ cfg: DictConfig,
+) -> Callable[[str], FlowerClient]:
+ """Generate the client function that creates the Flower Clients."""
+
+ def client_fn(cid: str) -> FlowerClient:
+ """Create a Flower client representing a single organization."""
+ device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
+
+ trainloader = trainloaders[int(cid)]
+ testloader = testloaders[int(cid)]
+
+ return FlowerClient(
+ int(cid),
+ cfg.dataset.name,
+ cfg.model.name,
+ cfg.model.output_dim,
+ trainloader,
+ testloader,
+ device,
+ cfg.num_epochs,
+ cfg.learning_rate,
+ cfg.mu,
+ cfg.temperature,
+ cfg.model.dir,
+ cfg.alg,
+ )
+
+ return client_fn
diff --git a/baselines/moon/moon/conf/base.yaml b/baselines/moon/moon/conf/base.yaml
new file mode 100644
index 000000000000..a2d3ddfb7bde
--- /dev/null
+++ b/baselines/moon/moon/conf/base.yaml
@@ -0,0 +1,33 @@
+---
+# this is the config that will be loaded as default by main.py
+# Please follow the provided structure (this will ensuring all baseline follow
+# a similar configuration structure and hence be easy to customise)
+
+num_clients: 10
+num_epochs: 10
+fraction_fit: 1.0
+batch_size: 64
+learning_rate: 0.01
+mu: 1
+temperature: 0.5
+alg: moon
+seed: 0
+server_device: cpu
+num_rounds: 100
+
+client_resources:
+ num_cpus: 2
+ num_gpus: 1
+
+dataset:
+ # dataset config
+ name: cifar10
+ dir: ./data/moon/
+ partition: noniid
+ beta: 0.5
+
+model:
+ # model config
+ name: simple-cnn
+ output_dim: 256
+ dir: ./models/moon/
\ No newline at end of file
diff --git a/baselines/moon/moon/conf/cifar10.yaml b/baselines/moon/moon/conf/cifar10.yaml
new file mode 100644
index 000000000000..672427495dfe
--- /dev/null
+++ b/baselines/moon/moon/conf/cifar10.yaml
@@ -0,0 +1,33 @@
+---
+# this is the config that will be loaded as default by main.py
+# Please follow the provided structure (this will ensuring all baseline follow
+# a similar configuration structure and hence be easy to customise)
+
+num_clients: 10
+num_epochs: 10
+fraction_fit: 1.0
+batch_size: 64
+learning_rate: 0.01
+mu: 5
+temperature: 0.5
+alg: moon
+seed: 0
+server_device: cpu
+num_rounds: 100
+
+client_resources:
+ num_cpus: 4
+ num_gpus: 0.2
+
+dataset:
+ # dataset config
+ name: cifar10
+ dir: ./data/moon/
+ partition: noniid
+ beta: 0.5
+
+model:
+ # model config
+ name: simple-cnn
+ output_dim: 256
+ dir: ./client_states/moon/cifar10/
\ No newline at end of file
diff --git a/baselines/moon/moon/conf/cifar100.yaml b/baselines/moon/moon/conf/cifar100.yaml
new file mode 100644
index 000000000000..33dc6d289456
--- /dev/null
+++ b/baselines/moon/moon/conf/cifar100.yaml
@@ -0,0 +1,33 @@
+---
+# this is the config that will be loaded as default by main.py
+# Please follow the provided structure (this will ensuring all baseline follow
+# a similar configuration structure and hence be easy to customise)
+
+num_clients: 10
+num_epochs: 10
+fraction_fit: 1.0
+batch_size: 64
+learning_rate: 0.01
+mu: 1
+temperature: 0.5
+alg: moon
+seed: 0
+server_device: cpu
+num_rounds: 100
+
+client_resources:
+ num_cpus: 4
+ num_gpus: 0.5
+
+dataset:
+ # dataset config
+ name: cifar100
+ dir: ./data/moon/
+ partition: noniid
+ beta: 0.5
+
+model:
+ # model config
+ name: resnet50
+ output_dim: 256
+ dir: ./client_states/moon/cifar100/
\ No newline at end of file
diff --git a/baselines/moon/moon/conf/cifar100_100clients.yaml b/baselines/moon/moon/conf/cifar100_100clients.yaml
new file mode 100644
index 000000000000..b314497b9411
--- /dev/null
+++ b/baselines/moon/moon/conf/cifar100_100clients.yaml
@@ -0,0 +1,33 @@
+---
+# this is the config that will be loaded as default by main.py
+# Please follow the provided structure (this will ensuring all baseline follow
+# a similar configuration structure and hence be easy to customise)
+
+num_clients: 100
+num_epochs: 10
+fraction_fit: 0.2
+batch_size: 64
+learning_rate: 0.01
+mu: 10
+temperature: 0.5
+alg: moon
+seed: 0
+server_device: cpu
+num_rounds: 500
+
+client_resources:
+ num_cpus: 8
+ num_gpus: 0.5
+
+dataset:
+ # dataset config
+ name: cifar100
+ dir: ./data/moon/
+ partition: noniid
+ beta: 0.5
+
+model:
+ # model config
+ name: resnet50
+ output_dim: 256
+ dir: ./client_states/moon/cifar100_100c/
\ No newline at end of file
diff --git a/baselines/moon/moon/conf/cifar100_50clients.yaml b/baselines/moon/moon/conf/cifar100_50clients.yaml
new file mode 100644
index 000000000000..d8c5877a1dcc
--- /dev/null
+++ b/baselines/moon/moon/conf/cifar100_50clients.yaml
@@ -0,0 +1,33 @@
+---
+# this is the config that will be loaded as default by main.py
+# Please follow the provided structure (this will ensuring all baseline follow
+# a similar configuration structure and hence be easy to customise)
+
+num_clients: 50
+num_epochs: 10
+fraction_fit: 1.0
+batch_size: 64
+learning_rate: 0.01
+mu: 10
+temperature: 0.5
+alg: moon
+seed: 0
+server_device: cpu
+num_rounds: 200
+
+client_resources:
+ num_cpus: 4
+ num_gpus: 0.5
+
+dataset:
+ # dataset config
+ name: cifar100
+ dir: ./data/moon/
+ partition: noniid
+ beta: 0.5
+
+model:
+ # model config
+ name: resnet50
+ output_dim: 256
+ dir: ./client_states/moon/cifar100_50clients/
\ No newline at end of file
diff --git a/baselines/moon/moon/conf/cifar100_50clients_fedprox.yaml b/baselines/moon/moon/conf/cifar100_50clients_fedprox.yaml
new file mode 100644
index 000000000000..69691021438a
--- /dev/null
+++ b/baselines/moon/moon/conf/cifar100_50clients_fedprox.yaml
@@ -0,0 +1,33 @@
+---
+# this is the config that will be loaded as default by main.py
+# Please follow the provided structure (this will ensuring all baseline follow
+# a similar configuration structure and hence be easy to customise)
+
+num_clients: 50
+num_epochs: 10
+fraction_fit: 1.0
+batch_size: 64
+learning_rate: 0.01
+mu: 0.001
+temperature: 0.5
+alg: fedprox
+seed: 0
+server_device: cpu
+num_rounds: 200
+
+client_resources:
+ num_cpus: 4
+ num_gpus: 0.5
+
+dataset:
+ # dataset config
+ name: cifar100
+ dir: ./data/moon/
+ partition: noniid
+ beta: 0.5
+
+model:
+ # model config
+ name: resnet50
+ output_dim: 256
+ dir: ./client_states/fedprox/cifar100_50clients/
\ No newline at end of file
diff --git a/baselines/moon/moon/conf/cifar100_fedprox.yaml b/baselines/moon/moon/conf/cifar100_fedprox.yaml
new file mode 100644
index 000000000000..1544f8e3a348
--- /dev/null
+++ b/baselines/moon/moon/conf/cifar100_fedprox.yaml
@@ -0,0 +1,33 @@
+---
+# this is the config that will be loaded as default by main.py
+# Please follow the provided structure (this will ensuring all baseline follow
+# a similar configuration structure and hence be easy to customise)
+
+num_clients: 10
+num_epochs: 10
+fraction_fit: 1.0
+batch_size: 64
+learning_rate: 0.01
+mu: 0.001
+temperature: 0.5
+alg: moon
+seed: 0
+server_device: cpu
+num_rounds: 100
+
+client_resources:
+ num_cpus: 4
+ num_gpus: 0.5
+
+dataset:
+ # dataset config
+ name: cifar100
+ dir: ./data/moon/
+ partition: noniid
+ beta: 0.5
+
+model:
+ # model config
+ name: resnet50
+ output_dim: 256
+ dir: ./client_states/moon/cifar100_fedprox/
\ No newline at end of file
diff --git a/baselines/moon/moon/conf/cifar10_fedprox.yaml b/baselines/moon/moon/conf/cifar10_fedprox.yaml
new file mode 100644
index 000000000000..d0f9c5e8e163
--- /dev/null
+++ b/baselines/moon/moon/conf/cifar10_fedprox.yaml
@@ -0,0 +1,33 @@
+---
+# this is the config that will be loaded as default by main.py
+# Please follow the provided structure (this will ensuring all baseline follow
+# a similar configuration structure and hence be easy to customise)
+
+num_clients: 10
+num_epochs: 10
+fraction_fit: 1.0
+batch_size: 64
+learning_rate: 0.01
+mu: 0.001
+temperature: 0.5
+alg: fedprox
+seed: 0
+server_device: cpu
+num_rounds: 100
+
+client_resources:
+ num_cpus: 4
+ num_gpus: 0.2
+
+dataset:
+ # dataset config
+ name: cifar10
+ dir: ./data/moon/
+ partition: noniid
+ beta: 0.5
+
+model:
+ # model config
+ name: simple-cnn
+ output_dim: 256
+ dir: ./client_states/moon/cifar10_fedprox/
\ No newline at end of file
diff --git a/baselines/moon/moon/dataset.py b/baselines/moon/moon/dataset.py
new file mode 100644
index 000000000000..0ec5c6ae9e27
--- /dev/null
+++ b/baselines/moon/moon/dataset.py
@@ -0,0 +1,271 @@
+"""Handle basic dataset creation.
+
+In case of PyTorch it should return dataloaders for your dataset (for both the clients
+and the server). If you are using a custom dataset class, this module is the place to
+define it. If your dataset requires to be downloaded (and this is not done
+automatically -- e.g. as it is the case for many dataset in TorchVision) and
+partitioned, please include all those functions and logic in the
+`dataset_preparation.py` module. You can use all those functions from functions/methods
+defined here of course.
+"""
+
+# https://github.com/QinbinLi/MOON/blob/main/datasets.py
+
+import logging
+import os
+
+import numpy as np
+import torch.nn.functional as F
+import torch.utils.data as data
+import torchvision
+import torchvision.transforms as transforms
+from PIL import Image
+from torch.autograd import Variable
+from torchvision.datasets import CIFAR10, CIFAR100
+
+logging.basicConfig()
+logger = logging.getLogger()
+logger.setLevel(logging.INFO)
+
+IMG_EXTENSIONS = (
+ ".jpg",
+ ".jpeg",
+ ".png",
+ ".ppm",
+ ".bmp",
+ ".pgm",
+ ".tif",
+ ".tiff",
+ ".webp",
+)
+
+
+class CIFAR10Sub(data.Dataset):
+ """CIFAR-10 dataset with idxs."""
+
+ def __init__(
+ self,
+ root,
+ dataidxs=None,
+ train=True,
+ transform=None,
+ target_transform=None,
+ download=False,
+ ):
+ self.root = root
+ self.dataidxs = dataidxs
+ self.train = train
+ self.transform = transform
+ self.target_transform = target_transform
+ self.download = download
+
+ self.data, self.target = self.__build_sub_dataset__()
+
+ def __build_sub_dataset__(self):
+ """Build sub dataset given idxs."""
+ cifar_dataobj = CIFAR10(
+ self.root, self.train, self.transform, self.target_transform, self.download
+ )
+
+ if torchvision.__version__ == "0.2.1":
+ if self.train:
+ # pylint: disable=redefined-outer-name
+ data, target = cifar_dataobj.train_data, np.array(
+ cifar_dataobj.train_labels
+ )
+ else:
+ # pylint: disable=redefined-outer-name
+ data, target = cifar_dataobj.test_data, np.array(
+ cifar_dataobj.test_labels
+ )
+ else:
+ data = cifar_dataobj.data
+ target = np.array(cifar_dataobj.targets)
+
+ if self.dataidxs is not None:
+ data = data[self.dataidxs]
+ target = target[self.dataidxs]
+
+ return data, target
+
+ def __getitem__(self, index):
+ """Get item by index.
+
+ Args:
+ index (int): Index.
+
+ Returns
+ -------
+ tuple: (image, target) where target is index of the target class.
+ """
+ img, target = self.data[index], self.target[index]
+
+ if self.transform is not None:
+ img = self.transform(img)
+
+ if self.target_transform is not None:
+ target = self.target_transform(target)
+
+ return img, target
+
+ def __len__(self):
+ """Length.
+
+ Returns
+ -------
+ int: length of data
+ """
+ return len(self.data)
+
+
+class CIFAR100Sub(data.Dataset):
+ """CIFAR-100 dataset with idxs."""
+
+ def __init__(
+ self,
+ root,
+ dataidxs=None,
+ train=True,
+ transform=None,
+ target_transform=None,
+ download=False,
+ ):
+ self.root = root
+ self.dataidxs = dataidxs
+ self.train = train
+ self.transform = transform
+ self.target_transform = target_transform
+ self.download = download
+
+ self.data, self.target = self.__build_sub_dataset__()
+
+ def __build_sub_dataset__(self):
+ """Build sub dataset given idxs."""
+ cifar_dataobj = CIFAR100(
+ self.root, self.train, self.transform, self.target_transform, self.download
+ )
+
+ if torchvision.__version__ == "0.2.1":
+ if self.train:
+ # pylint: disable=redefined-outer-name
+ data, target = cifar_dataobj.train_data, np.array(
+ cifar_dataobj.train_labels
+ )
+ else:
+ data, target = cifar_dataobj.test_data, np.array(
+ cifar_dataobj.test_labels
+ ) # pylint: disable=redefined-outer-name
+ else:
+ data = cifar_dataobj.data
+ target = np.array(cifar_dataobj.targets)
+
+ if self.dataidxs is not None:
+ data = data[self.dataidxs]
+ target = target[self.dataidxs]
+
+ return data, target
+
+ def __getitem__(self, index):
+ """Get item by index.
+
+ Args:
+ index (int): Index.
+
+ Returns
+ -------
+ tuple: (image, target) where target is index of the target class.
+ """
+ img, target = self.data[index], self.target[index]
+ img = Image.fromarray(img)
+
+ if self.transform is not None:
+ img = self.transform(img)
+
+ if self.target_transform is not None:
+ target = self.target_transform(target)
+
+ return img, target
+
+ def __len__(self):
+ """Length.
+
+ Returns
+ -------
+ int: length of data
+ """
+ return len(self.data)
+
+
+def get_dataloader(dataset, datadir, train_bs, test_bs, dataidxs=None, noise_level=0):
+ """Get dataloader for a given dataset."""
+ if dataset == "cifar10":
+ dl_obj = CIFAR10Sub
+ normalize = transforms.Normalize(
+ mean=[x / 255.0 for x in [125.3, 123.0, 113.9]],
+ std=[x / 255.0 for x in [63.0, 62.1, 66.7]],
+ )
+ transform_train = transforms.Compose(
+ [
+ transforms.ToTensor(),
+ transforms.Lambda(
+ lambda x: F.pad(
+ Variable(x.unsqueeze(0), requires_grad=False),
+ (4, 4, 4, 4),
+ mode="reflect",
+ ).data.squeeze()
+ ),
+ transforms.ToPILImage(),
+ transforms.ColorJitter(brightness=noise_level),
+ transforms.RandomCrop(32),
+ transforms.RandomHorizontalFlip(),
+ transforms.ToTensor(),
+ normalize,
+ ]
+ )
+ # data prep for test set
+ transform_test = transforms.Compose([transforms.ToTensor(), normalize])
+
+ elif dataset == "cifar100":
+ dl_obj = CIFAR100Sub
+
+ normalize = transforms.Normalize(
+ mean=[0.5070751592371323, 0.48654887331495095, 0.4409178433670343],
+ std=[0.2673342858792401, 0.2564384629170883, 0.27615047132568404],
+ )
+
+ transform_train = transforms.Compose(
+ [
+ transforms.RandomCrop(32, padding=4),
+ transforms.RandomHorizontalFlip(),
+ transforms.RandomRotation(15),
+ transforms.ToTensor(),
+ normalize,
+ ]
+ )
+ # data prep for test set
+ transform_test = transforms.Compose([transforms.ToTensor(), normalize])
+ if dataset == "cifar10" and os.path.isdir(
+ os.path.join(datadir, "cifar-10-batches-py")
+ ):
+ download = False
+ elif dataset == "cifar100" and os.path.isdir(
+ os.path.join(datadir, "cifar-100-python")
+ ):
+ download = False
+ else:
+ download = True
+ train_ds = dl_obj(
+ datadir,
+ dataidxs=dataidxs,
+ train=True,
+ transform=transform_train,
+ download=download,
+ )
+ test_ds = dl_obj(datadir, train=False, transform=transform_test, download=download)
+
+ train_dl = data.DataLoader(
+ dataset=train_ds, batch_size=train_bs, drop_last=True, shuffle=True
+ )
+ test_dl = data.DataLoader(dataset=test_ds, batch_size=test_bs, shuffle=False)
+
+ return train_dl, test_dl, train_ds, test_ds
diff --git a/baselines/moon/moon/dataset_preparation.py b/baselines/moon/moon/dataset_preparation.py
new file mode 100644
index 000000000000..11103d37763b
--- /dev/null
+++ b/baselines/moon/moon/dataset_preparation.py
@@ -0,0 +1,100 @@
+"""Handle the dataset partitioning and (optionally) complex downloads.
+
+Please add here all the necessary logic to either download, uncompress, pre/post-process
+your dataset (or all of the above). If the desired way of running your baseline is to
+first download the dataset and partition it and then run the experiments, please
+uncomment the lines below and tell us in the README.md (see the "Running the Experiment"
+block) that this file should be executed first.
+"""
+
+import numpy as np
+import torchvision.transforms as transforms
+
+from moon.dataset import CIFAR10Sub, CIFAR100Sub
+
+
+def load_cifar10_data(datadir):
+ """Load CIFAR10 dataset."""
+ transform = transforms.Compose([transforms.ToTensor()])
+
+ cifar10_train_ds = CIFAR10Sub(
+ datadir, train=True, download=True, transform=transform
+ )
+ cifar10_test_ds = CIFAR10Sub(
+ datadir, train=False, download=True, transform=transform
+ )
+
+ X_train, y_train = cifar10_train_ds.data, cifar10_train_ds.target
+ X_test, y_test = cifar10_test_ds.data, cifar10_test_ds.target
+
+ return (X_train, y_train, X_test, y_test)
+
+
+def load_cifar100_data(datadir):
+ """Load CIFAR100 dataset."""
+ transform = transforms.Compose([transforms.ToTensor()])
+
+ cifar100_train_ds = CIFAR100Sub(
+ datadir, train=True, download=True, transform=transform
+ )
+ cifar100_test_ds = CIFAR100Sub(
+ datadir, train=False, download=True, transform=transform
+ )
+
+ X_train, y_train = cifar100_train_ds.data, cifar100_train_ds.target
+ X_test, y_test = cifar100_test_ds.data, cifar100_test_ds.target
+
+ return (X_train, y_train, X_test, y_test)
+
+
+# pylint: disable=too-many-locals
+def partition_data(dataset, datadir, partition, num_clients, beta):
+ """Partition data into train and test sets for IID and non-IID experiments."""
+ if dataset == "cifar10":
+ X_train, y_train, X_test, y_test = load_cifar10_data(datadir)
+ elif dataset == "cifar100":
+ X_train, y_train, X_test, y_test = load_cifar100_data(datadir)
+
+ n_train = y_train.shape[0]
+
+ if partition in ("homo", "iid"):
+ idxs = np.random.permutation(n_train)
+ batch_idxs = np.array_split(idxs, num_clients)
+ net_dataidx_map = {i: batch_idxs[i] for i in range(num_clients)}
+
+ elif partition in ("noniid-labeldir", "noniid"):
+ min_size = 0
+ min_require_size = 10
+ K = 10
+ if dataset == "cifar100":
+ K = 100
+ elif dataset == "tinyimagenet":
+ K = 200
+
+ N = y_train.shape[0]
+ net_dataidx_map = {}
+
+ while min_size < min_require_size:
+ idx_batch = [[] for _ in range(num_clients)]
+ for k in range(K):
+ idx_k = np.where(y_train == k)[0]
+ np.random.shuffle(idx_k)
+ proportions = np.random.dirichlet(np.repeat(beta, num_clients))
+ proportions = np.array(
+ [
+ p * (len(idx_j) < N / num_clients)
+ for p, idx_j in zip(proportions, idx_batch)
+ ]
+ )
+ proportions = proportions / proportions.sum()
+ proportions = (np.cumsum(proportions) * len(idx_k)).astype(int)[:-1]
+ idx_batch = [
+ idx_j + idx.tolist()
+ for idx_j, idx in zip(idx_batch, np.split(idx_k, proportions))
+ ]
+ min_size = min([len(idx_j) for idx_j in idx_batch])
+ for j in range(num_clients):
+ np.random.shuffle(idx_batch[j])
+ net_dataidx_map[j] = idx_batch[j]
+
+ return (X_train, y_train, X_test, y_test, net_dataidx_map)
diff --git a/baselines/moon/moon/main.py b/baselines/moon/moon/main.py
new file mode 100644
index 000000000000..902ccfa8395c
--- /dev/null
+++ b/baselines/moon/moon/main.py
@@ -0,0 +1,150 @@
+"""Create and connect the building blocks for your experiments; start the simulation.
+
+It includes processioning the dataset, instantiate strategy, specify how the global
+model is going to be evaluated, etc. At the end, this script saves the results.
+"""
+import os
+import random
+import shutil
+from pathlib import Path
+
+# these are the basic packages you'll need here
+# feel free to remove some if aren't needed
+import flwr as fl
+import hydra
+import numpy as np
+import torch
+from hydra.core.hydra_config import HydraConfig
+from omegaconf import DictConfig, OmegaConf
+
+from moon import client, server
+from moon.dataset import get_dataloader
+from moon.dataset_preparation import partition_data
+from moon.utils import plot_metric_from_history
+
+
+@hydra.main(config_path="conf", config_name="base", version_base=None)
+def main(cfg: DictConfig) -> None:
+ """Run the baseline.
+
+ Parameters
+ ----------
+ cfg : DictConfig
+ An omegaconf object that stores the hydra config.
+ """
+ # Clean the model directory to save models for MOON
+ if cfg.alg == "moon":
+ if os.path.exists(cfg.model.dir):
+ shutil.rmtree(cfg.model.dir)
+ # 1. Print parsed config
+ print(OmegaConf.to_yaml(cfg))
+
+ # 2. Prepare your dataset
+ np.random.seed(cfg.seed)
+ torch.manual_seed(cfg.seed)
+ if torch.cuda.is_available():
+ torch.cuda.manual_seed(cfg.seed)
+ random.seed(cfg.seed)
+ (
+ _,
+ _,
+ _,
+ _,
+ net_dataidx_map,
+ ) = partition_data(
+ dataset=cfg.dataset.name,
+ datadir=cfg.dataset.dir,
+ partition=cfg.dataset.partition,
+ num_clients=cfg.num_clients,
+ beta=cfg.dataset.beta,
+ )
+
+ _, test_global_dl, _, _ = get_dataloader(
+ dataset=cfg.dataset.name,
+ datadir=cfg.dataset.dir,
+ train_bs=cfg.batch_size,
+ test_bs=32,
+ )
+
+ trainloaders = []
+ testloaders = []
+ for idx in range(cfg.num_clients):
+ train_dl, test_dl, _, _ = get_dataloader(
+ cfg.dataset.name, cfg.dataset.dir, cfg.batch_size, 32, net_dataidx_map[idx]
+ )
+
+ trainloaders.append(train_dl)
+ testloaders.append(test_dl)
+ # 3. Define your clients
+ # Define a function that returns another function that will be used during
+ # simulation to instantiate each individual client
+ client_fn = client.gen_client_fn(
+ trainloaders=trainloaders,
+ testloaders=testloaders,
+ cfg=cfg,
+ )
+
+ # get function that will executed by the strategy's evaluate() method
+ # Set server's device
+ device = (
+ torch.device("cuda:0")
+ if torch.cuda.is_available() and cfg.server_device == "cuda"
+ else "cpu"
+ )
+ evaluate_fn = server.gen_evaluate_fn(test_global_dl, device=device, cfg=cfg)
+
+ # 4. Define your strategy
+ strategy = fl.server.strategy.FedAvg(
+ # Clients in MOON do not perform federated evaluation
+ # (see the client's evaluate())
+ fraction_fit=cfg.fraction_fit,
+ fraction_evaluate=0.0,
+ evaluate_fn=evaluate_fn,
+ )
+ # 5. Start Simulation
+ # history = fl.simulation.start_simulation()
+ history = fl.simulation.start_simulation(
+ client_fn=client_fn,
+ num_clients=cfg.num_clients,
+ config=fl.server.ServerConfig(num_rounds=cfg.num_rounds),
+ client_resources={
+ "num_cpus": cfg.client_resources.num_cpus,
+ "num_gpus": cfg.client_resources.num_gpus,
+ },
+ strategy=strategy,
+ )
+ # remove saved models
+ if cfg.alg == "moon":
+ shutil.rmtree(cfg.model.dir)
+
+ # 6. Save your results
+ # Experiment completed. Now we save the results and
+ # generate plots using the `history`
+ print("................")
+ print(history)
+
+ # Hydra automatically creates an output directory
+ # Let's retrieve it and save some results there
+ save_path = HydraConfig.get().runtime.output_dir
+
+ # plot results and include them in the readme
+ strategy_name = strategy.__class__.__name__
+ file_suffix: str = (
+ f"_{strategy_name}"
+ f"{'_dataset' if cfg.dataset.name else ''}"
+ f"_C={cfg.num_clients}"
+ f"_B={cfg.batch_size}"
+ f"_E={cfg.num_epochs}"
+ f"_R={cfg.num_rounds}"
+ f"_mu={cfg.mu}"
+ )
+
+ plot_metric_from_history(
+ history,
+ Path(save_path),
+ (file_suffix),
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/baselines/moon/moon/models.py b/baselines/moon/moon/models.py
new file mode 100644
index 000000000000..a323a8e74727
--- /dev/null
+++ b/baselines/moon/moon/models.py
@@ -0,0 +1,528 @@
+"""Define our models, and training and eval functions.
+
+If your model is 100% off-the-shelf (e.g. directly from torchvision without requiring
+modifications) you might be better off instantiating your model directly from the Hydra
+config. In this way, swapping your model for another one can be done without changing
+the python code at all
+"""
+
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torch.optim as optim
+
+from moon.utils import compute_accuracy
+
+
+def conv3x3(in_planes, out_planes, stride=1, groups=1, dilation=1):
+ """3x3 convolution with padding."""
+ return nn.Conv2d(
+ in_planes,
+ out_planes,
+ kernel_size=3,
+ stride=stride,
+ padding=dilation,
+ groups=groups,
+ bias=False,
+ dilation=dilation,
+ )
+
+
+def conv1x1(in_planes, out_planes, stride=1):
+ """1x1 convolution."""
+ return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)
+
+
+class BasicBlock(nn.Module):
+ """Basic Block for resnet."""
+
+ expansion = 1
+
+ def __init__(
+ self,
+ inplanes,
+ planes,
+ stride=1,
+ downsample=None,
+ groups=1,
+ base_width=64,
+ dilation=1,
+ norm_layer=None,
+ ):
+ super().__init__()
+ if norm_layer is None:
+ norm_layer = nn.BatchNorm2d
+ if groups != 1 or base_width != 64:
+ raise ValueError("BasicBlock only supports groups=1 and base_width=64")
+ if dilation > 1:
+ raise NotImplementedError("Dilation > 1 not supported in BasicBlock")
+ self.conv1 = conv3x3(inplanes, planes, stride)
+ self.bn1 = norm_layer(planes)
+ self.relu = nn.ReLU(inplace=True)
+ self.conv2 = conv3x3(planes, planes)
+ self.bn2 = norm_layer(planes)
+ self.downsample = downsample
+ self.stride = stride
+
+ def forward(self, x):
+ """Forward."""
+ identity = x
+
+ out = self.conv1(x)
+ out = self.bn1(out)
+ out = self.relu(out)
+
+ out = self.conv2(out)
+ out = self.bn2(out)
+
+ if self.downsample is not None:
+ identity = self.downsample(x)
+
+ out += identity
+ out = self.relu(out)
+
+ return out
+
+
+class Bottleneck(nn.Module):
+ """Bottleneck in torchvision places the stride."""
+
+ expansion = 4
+
+ def __init__(
+ self,
+ inplanes,
+ planes,
+ stride=1,
+ downsample=None,
+ groups=1,
+ base_width=64,
+ dilation=1,
+ norm_layer=None,
+ ):
+ super().__init__()
+ if norm_layer is None:
+ norm_layer = nn.BatchNorm2d
+ width = int(planes * (base_width / 64.0)) * groups
+ self.conv1 = conv1x1(inplanes, width)
+ self.bn1 = norm_layer(width)
+ self.conv2 = conv3x3(width, width, stride, groups, dilation)
+ self.bn2 = norm_layer(width)
+ self.conv3 = conv1x1(width, planes * self.expansion)
+ self.bn3 = norm_layer(planes * self.expansion)
+ self.relu = nn.ReLU(inplace=True)
+ self.downsample = downsample
+ self.stride = stride
+
+ def forward(self, x):
+ """Forward."""
+ identity = x
+
+ out = self.conv1(x)
+ out = self.bn1(out)
+ out = self.relu(out)
+
+ out = self.conv2(out)
+ out = self.bn2(out)
+ out = self.relu(out)
+
+ out = self.conv3(out)
+ out = self.bn3(out)
+
+ if self.downsample is not None:
+ identity = self.downsample(x)
+
+ out += identity
+ out = self.relu(out)
+
+ return out
+
+
+class ResNetCifar10(nn.Module):
+ """ResNet model."""
+
+ def __init__(
+ self,
+ block,
+ layers,
+ num_classes=1000,
+ zero_init_residual=False,
+ groups=1,
+ width_per_group=64,
+ replace_stride_with_dilation=None,
+ norm_layer=None,
+ ):
+ super().__init__()
+ if norm_layer is None:
+ norm_layer = nn.BatchNorm2d
+ self._norm_layer = norm_layer
+
+ self.inplanes = 64
+ self.dilation = 1
+ if replace_stride_with_dilation is None:
+ # each element in the tuple indicates if we should replace
+ # the 2x2 stride with a dilated convolution instead
+ replace_stride_with_dilation = [False, False, False]
+ if len(replace_stride_with_dilation) != 3:
+ raise ValueError(
+ "replace_stride_with_dilation should be None "
+ "or a 3-element tuple, got {}".format(replace_stride_with_dilation)
+ )
+ self.groups = groups
+ self.base_width = width_per_group
+ self.conv1 = nn.Conv2d(
+ 3, self.inplanes, kernel_size=3, stride=1, padding=1, bias=False
+ )
+ self.bn1 = norm_layer(self.inplanes)
+ self.relu = nn.ReLU(inplace=True)
+ self.layer1 = self._make_layer(block, 64, layers[0])
+ self.layer2 = self._make_layer(
+ block, 128, layers[1], stride=2, dilate=replace_stride_with_dilation[0]
+ )
+ self.layer3 = self._make_layer(
+ block, 256, layers[2], stride=2, dilate=replace_stride_with_dilation[1]
+ )
+ self.layer4 = self._make_layer(
+ block, 512, layers[3], stride=2, dilate=replace_stride_with_dilation[2]
+ )
+ self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
+ self.fc = nn.Linear(512 * block.expansion, num_classes)
+
+ for module in self.modules():
+ if isinstance(module, nn.Conv2d):
+ nn.init.kaiming_normal_(
+ module.weight, mode="fan_out", nonlinearity="relu"
+ )
+ elif isinstance(module, (nn.BatchNorm2d, nn.GroupNorm)):
+ nn.init.constant_(module.weight, 1)
+ nn.init.constant_(module.bias, 0)
+
+ if zero_init_residual:
+ for module in self.modules():
+ if isinstance(module, Bottleneck):
+ nn.init.constant_(module.bn3.weight, 0)
+ elif isinstance(module, BasicBlock):
+ nn.init.constant_(module.bn2.weight, 0)
+
+ def _make_layer(self, block, planes, blocks, stride=1, dilate=False):
+ norm_layer = self._norm_layer
+ downsample = None
+ previous_dilation = self.dilation
+ if dilate:
+ self.dilation *= stride
+ stride = 1
+ if stride != 1 or self.inplanes != planes * block.expansion:
+ downsample = nn.Sequential(
+ conv1x1(self.inplanes, planes * block.expansion, stride),
+ norm_layer(planes * block.expansion),
+ )
+
+ layers = []
+ layers.append(
+ block(
+ self.inplanes,
+ planes,
+ stride,
+ downsample,
+ self.groups,
+ self.base_width,
+ previous_dilation,
+ norm_layer,
+ )
+ )
+ self.inplanes = planes * block.expansion
+ for _ in range(1, blocks):
+ layers.append(
+ block(
+ self.inplanes,
+ planes,
+ groups=self.groups,
+ base_width=self.base_width,
+ dilation=self.dilation,
+ norm_layer=norm_layer,
+ )
+ )
+
+ return nn.Sequential(*layers)
+
+ def _forward_impl(self, x):
+ # See note [TorchScript super()]
+ x = self.conv1(x)
+ x = self.bn1(x)
+ x = self.relu(x)
+
+ x = self.layer1(x)
+ x = self.layer2(x)
+ x = self.layer3(x)
+ x = self.layer4(x)
+
+ x = self.avgpool(x)
+ x = torch.flatten(x, 1)
+ x = self.fc(x)
+
+ return x
+
+ def forward(self, x):
+ """Forward."""
+ return self._forward_impl(x)
+
+
+def resnet50_cifar10(**kwargs):
+ r"""ResNet-50 model from `"Deep Residual Learning for Image Recognition".
+
+ `_
+
+ Args:
+ pretrained (bool): If True, returns a model pre-trained on ImageNet
+ progress (bool): If True, displays a progress bar of the download to stderr
+ """
+ return ResNetCifar10(Bottleneck, [3, 4, 6, 3], **kwargs)
+
+
+class SimpleCNNHeader(nn.Module):
+ """Simple CNN model."""
+
+ def __init__(self, input_dim, hidden_dims):
+ super().__init__()
+ self.conv1 = nn.Conv2d(3, 6, 5)
+ self.relu = nn.ReLU()
+ self.pool = nn.MaxPool2d(2, 2)
+ self.conv2 = nn.Conv2d(6, 16, 5)
+
+ self.fc1 = nn.Linear(input_dim, hidden_dims[0])
+ self.fc2 = nn.Linear(hidden_dims[0], hidden_dims[1])
+
+ def forward(self, x):
+ """Forward."""
+ x = self.pool(self.relu(self.conv1(x)))
+ x = self.pool(self.relu(self.conv2(x)))
+ x = x.view(-1, 16 * 5 * 5)
+
+ x = self.relu(self.fc1(x))
+ x = self.relu(self.fc2(x))
+ # x = self.fc3(x)
+ return x
+
+
+class ModelMOON(nn.Module):
+ """Model for MOON."""
+
+ def __init__(self, base_model, out_dim, n_classes):
+ super().__init__()
+
+ if base_model in (
+ "resnet50-cifar10",
+ "resnet50-cifar100",
+ "resnet50-smallkernel",
+ "resnet50",
+ ):
+ basemodel = resnet50_cifar10()
+ self.features = nn.Sequential(*list(basemodel.children())[:-1])
+ num_ftrs = basemodel.fc.in_features
+ elif base_model == "simple-cnn":
+ self.features = SimpleCNNHeader(
+ input_dim=(16 * 5 * 5), hidden_dims=[120, 84]
+ )
+ num_ftrs = 84
+
+ # projection MLP
+ self.l1 = nn.Linear(num_ftrs, num_ftrs)
+ self.l2 = nn.Linear(num_ftrs, out_dim)
+
+ # last layer
+ self.l3 = nn.Linear(out_dim, n_classes)
+
+ def _get_basemodel(self, model_name):
+ try:
+ model = self.model_dict[model_name]
+ return model
+ except KeyError as err:
+ raise ValueError("Invalid model name.") from err
+
+ def forward(self, x):
+ """Forward."""
+ h = self.features(x)
+ h = h.squeeze()
+ x = self.l1(h)
+ x = F.relu(x)
+ x = self.l2(x)
+
+ y = self.l3(x)
+ return h, x, y
+
+
+def init_net(dataset, model, output_dim, device="cpu"):
+ """Initialize model."""
+ if dataset == "cifar10":
+ n_classes = 10
+ elif dataset == "cifar100":
+ n_classes = 100
+
+ net = ModelMOON(model, output_dim, n_classes)
+ if device == "cpu":
+ net.to(device)
+ else:
+ net = net.cuda()
+
+ return net
+
+
+def train_moon(
+ net,
+ global_net,
+ previous_net,
+ train_dataloader,
+ epochs,
+ lr,
+ mu,
+ temperature,
+ device="cpu",
+):
+ """Training function for MOON."""
+ net.to(device)
+ global_net.to(device)
+ previous_net.to(device)
+ train_acc, _ = compute_accuracy(net, train_dataloader, device=device)
+ optimizer = optim.SGD(
+ filter(lambda p: p.requires_grad, net.parameters()),
+ lr=lr,
+ momentum=0.9,
+ weight_decay=1e-5,
+ )
+
+ criterion = nn.CrossEntropyLoss().cuda()
+
+ previous_net.eval()
+ for param in previous_net.parameters():
+ param.requires_grad = False
+ previous_net.cuda()
+
+ cnt = 0
+ cos = torch.nn.CosineSimilarity(dim=-1)
+
+ for epoch in range(epochs):
+ epoch_loss_collector = []
+ epoch_loss1_collector = []
+ epoch_loss2_collector = []
+ for _, (x, target) in enumerate(train_dataloader):
+ x, target = x.to(device), target.to(device)
+
+ optimizer.zero_grad()
+ x.requires_grad = False
+ target.requires_grad = False
+ target = target.long()
+
+ # pro1 is the representation by the current model (Line 14 of Algorithm 1)
+ _, pro1, out = net(x)
+ # pro2 is the representation by the global model (Line 15 of Algorithm 1)
+ _, pro2, _ = global_net(x)
+ # posi is the positive pair
+ posi = cos(pro1, pro2)
+ logits = posi.reshape(-1, 1)
+
+ previous_net.to(device)
+ # pro 3 is the representation by the previous model (Line 16 of Algorithm 1)
+ _, pro3, _ = previous_net(x)
+ # nega is the negative pair
+ nega = cos(pro1, pro3)
+ logits = torch.cat((logits, nega.reshape(-1, 1)), dim=1)
+
+ previous_net.to("cpu")
+ logits /= temperature
+ labels = torch.zeros(x.size(0)).cuda().long()
+ # compute the model-contrastive loss (Line 17 of Algorithm 1)
+ loss2 = mu * criterion(logits, labels)
+ # compute the cross-entropy loss (Line 13 of Algorithm 1)
+ loss1 = criterion(out, target)
+ # compute the loss (Line 18 of Algorithm 1)
+ loss = loss1 + loss2
+
+ loss.backward()
+ optimizer.step()
+
+ cnt += 1
+ epoch_loss_collector.append(loss.item())
+ epoch_loss1_collector.append(loss1.item())
+ epoch_loss2_collector.append(loss2.item())
+
+ epoch_loss = sum(epoch_loss_collector) / len(epoch_loss_collector)
+ epoch_loss1 = sum(epoch_loss1_collector) / len(epoch_loss1_collector)
+ epoch_loss2 = sum(epoch_loss2_collector) / len(epoch_loss2_collector)
+ print(
+ "Epoch: %d Loss: %f Loss1: %f Loss2: %f"
+ % (epoch, epoch_loss, epoch_loss1, epoch_loss2)
+ )
+
+ previous_net.to("cpu")
+ train_acc, _ = compute_accuracy(net, train_dataloader, device=device)
+
+ print(">> Training accuracy: %f" % train_acc)
+ net.to("cpu")
+ global_net.to("cpu")
+ print(" ** Training complete **")
+ return net
+
+
+def train_fedprox(net, global_net, train_dataloader, epochs, lr, mu, device="cpu"):
+ """Training function for FedProx."""
+ net = nn.DataParallel(net)
+ net.cuda()
+
+ train_acc, _ = compute_accuracy(net, train_dataloader, device=device)
+
+ print(">> Pre-Training Training accuracy: {}".format(train_acc))
+
+ optimizer = optim.SGD(
+ filter(lambda p: p.requires_grad, net.parameters()),
+ lr=lr,
+ momentum=0.9,
+ weight_decay=1e-5,
+ )
+
+ criterion = nn.CrossEntropyLoss().cuda()
+
+ cnt = 0
+ global_weight_collector = list(global_net.cuda().parameters())
+
+ for _epoch in range(epochs):
+ epoch_loss_collector = []
+ for _, (x, target) in enumerate(train_dataloader):
+ x, target = x.cuda(), target.cuda()
+
+ optimizer.zero_grad()
+ x.requires_grad = False
+ target.requires_grad = False
+ target = target.long()
+
+ _, _, out = net(x)
+ loss = criterion(out, target)
+
+ fed_prox_reg = 0.0
+ for param_index, param in enumerate(net.parameters()):
+ fed_prox_reg += (mu / 2) * torch.norm(
+ (param - global_weight_collector[param_index])
+ ) ** 2
+ loss += fed_prox_reg
+
+ loss.backward()
+ optimizer.step()
+
+ cnt += 1
+ epoch_loss_collector.append(loss.item())
+
+ train_acc, _ = compute_accuracy(net, train_dataloader, device=device)
+
+ print(">> Training accuracy: %f" % train_acc)
+ net.to("cpu")
+ print(" ** Training complete **")
+ return net
+
+
+def test(net, test_dataloader, device="cpu"):
+ """Test function."""
+ net.to(device)
+ test_acc, loss = compute_accuracy(net, test_dataloader, device=device)
+ print(">> Test accuracy: %f" % test_acc)
+ net.to("cpu")
+ return test_acc, loss
diff --git a/baselines/moon/moon/server.py b/baselines/moon/moon/server.py
new file mode 100644
index 000000000000..0cf812b88666
--- /dev/null
+++ b/baselines/moon/moon/server.py
@@ -0,0 +1,40 @@
+"""Create global evaluation function.
+
+Optionally, also define a new Server class (please note this is not needed in most
+settings).
+"""
+
+from collections import OrderedDict
+from typing import Callable, Dict, Optional, Tuple
+
+import torch
+from flwr.common.typing import NDArrays, Scalar
+from omegaconf import DictConfig
+from torch.utils.data import DataLoader
+
+from moon.models import init_net, test
+
+
+def gen_evaluate_fn(
+ testloader: DataLoader,
+ device: torch.device,
+ cfg: DictConfig,
+) -> Callable[
+ [int, NDArrays, Dict[str, Scalar]], Optional[Tuple[float, Dict[str, Scalar]]]
+]:
+ """Generate the function for centralized evaluation."""
+
+ def evaluate(
+ server_round: int, parameters_ndarrays: NDArrays, config: Dict[str, Scalar]
+ ) -> Optional[Tuple[float, Dict[str, Scalar]]]:
+ # pylint: disable=unused-argument
+ net = init_net(cfg.dataset.name, cfg.model.name, cfg.model.output_dim)
+ params_dict = zip(net.state_dict().keys(), parameters_ndarrays)
+ state_dict = OrderedDict({k: torch.from_numpy(v) for k, v in params_dict})
+ net.load_state_dict(state_dict, strict=True)
+ net.to(device)
+
+ accuracy, loss = test(net, testloader, device=device)
+ return loss, {"accuracy": accuracy}
+
+ return evaluate
diff --git a/baselines/moon/moon/strategy.py b/baselines/moon/moon/strategy.py
new file mode 100644
index 000000000000..17436c401c30
--- /dev/null
+++ b/baselines/moon/moon/strategy.py
@@ -0,0 +1,5 @@
+"""Optionally define a custom strategy.
+
+Needed only when the strategy is not yet implemented in Flower or because you want to
+extend or modify the functionality of an existing strategy.
+"""
diff --git a/baselines/moon/moon/utils.py b/baselines/moon/moon/utils.py
new file mode 100644
index 000000000000..4b99a480f77b
--- /dev/null
+++ b/baselines/moon/moon/utils.py
@@ -0,0 +1,127 @@
+"""Define any utility function.
+
+They are not directly relevant to the other (more FL specific) python modules. For
+example, you may define here things like: loading a model from a checkpoint, saving
+results, plotting.
+"""
+from pathlib import Path
+from typing import Optional
+
+import matplotlib.pyplot as plt
+import numpy as np
+import torch
+import torch.nn as nn
+from flwr.server.history import History
+
+
+def compute_accuracy(model, dataloader, device="cpu", multiloader=False):
+ """Compute accuracy."""
+ was_training = False
+ if model.training:
+ model.eval()
+ was_training = True
+
+ true_labels_list, pred_labels_list = np.array([]), np.array([])
+
+ correct, total = 0, 0
+ if device == "cpu":
+ criterion = nn.CrossEntropyLoss()
+ elif "cuda" in device.type:
+ criterion = nn.CrossEntropyLoss().cuda()
+ loss_collector = []
+ if multiloader:
+ for loader in dataloader:
+ with torch.no_grad():
+ for _, (x, target) in enumerate(loader):
+ if device != "cpu":
+ x, target = x.cuda(), target.to(dtype=torch.int64).cuda()
+ _, _, out = model(x)
+ if len(target) == 1:
+ loss = criterion(out, target)
+ else:
+ loss = criterion(out, target)
+ _, pred_label = torch.max(out.data, 1)
+ loss_collector.append(loss.item())
+ total += x.data.size()[0]
+ correct += (pred_label == target.data).sum().item()
+
+ if device == "cpu":
+ pred_labels_list = np.append(
+ pred_labels_list, pred_label.numpy()
+ )
+ true_labels_list = np.append(
+ true_labels_list, target.data.numpy()
+ )
+ else:
+ pred_labels_list = np.append(
+ pred_labels_list, pred_label.cpu().numpy()
+ )
+ true_labels_list = np.append(
+ true_labels_list, target.data.cpu().numpy()
+ )
+ avg_loss = sum(loss_collector) / len(loss_collector)
+ else:
+ with torch.no_grad():
+ for _, (x, target) in enumerate(dataloader):
+ # print("x:",x)
+ if device != "cpu":
+ x, target = x.cuda(), target.to(dtype=torch.int64).cuda()
+ _, _, out = model(x)
+ loss = criterion(out, target)
+ _, pred_label = torch.max(out.data, 1)
+ loss_collector.append(loss.item())
+ total += x.data.size()[0]
+ correct += (pred_label == target.data).sum().item()
+
+ if device == "cpu":
+ pred_labels_list = np.append(pred_labels_list, pred_label.numpy())
+ true_labels_list = np.append(true_labels_list, target.data.numpy())
+ else:
+ pred_labels_list = np.append(
+ pred_labels_list, pred_label.cpu().numpy()
+ )
+ true_labels_list = np.append(
+ true_labels_list, target.data.cpu().numpy()
+ )
+ avg_loss = sum(loss_collector) / len(loss_collector)
+
+ if was_training:
+ model.train()
+
+ return correct / float(total), avg_loss
+
+
+def plot_metric_from_history(
+ hist: History,
+ save_plot_path: Path,
+ suffix: Optional[str] = "",
+) -> None:
+ """Plot data from Flower server History.
+
+ Parameters
+ ----------
+ hist : History
+ Object containing evaluation for all rounds.
+ save_plot_path : Path
+ Folder to save the plot to.
+ suffix: Optional[str]
+ Optional string to add at the end of the filename for the plot.
+ """
+ metric_type = "centralized"
+ metric_dict = (
+ hist.metrics_centralized
+ if metric_type == "centralized"
+ else hist.metrics_distributed
+ )
+ rounds, values = zip(*metric_dict["accuracy"])
+
+ # Plot the curve
+ plt.figure(figsize=(10, 6))
+ plt.plot(rounds, values)
+ plt.xlabel("#round")
+ plt.ylabel("Test accuracy")
+ plt.legend()
+ plt.show()
+
+ plt.savefig(Path(save_plot_path) / Path(f"{metric_type}_metrics{suffix}.png"))
+ plt.close()
diff --git a/baselines/moon/pyproject.toml b/baselines/moon/pyproject.toml
new file mode 100644
index 000000000000..e9f826abb2ea
--- /dev/null
+++ b/baselines/moon/pyproject.toml
@@ -0,0 +1,146 @@
+[build-system]
+requires = ["poetry-core>=1.4.0"]
+build-backend = "poetry.masonry.api"
+
+[tool.poetry]
+name = "moon" # <----- Ensure it matches the name of your baseline directory containing all the source code
+version = "1.0.0"
+description = "Model-Contrastive Federated Learning"
+license = "Apache-2.0"
+authors = ["The Flower Authors ", "Qinbin Li "]
+readme = "README.md"
+homepage = "https://flower.dev"
+repository = "https://github.com/adap/flower"
+documentation = "https://flower.dev"
+classifiers = [
+ "Development Status :: 3 - Alpha",
+ "Intended Audience :: Developers",
+ "Intended Audience :: Science/Research",
+ "License :: OSI Approved :: Apache Software License",
+ "Operating System :: MacOS :: MacOS X",
+ "Operating System :: POSIX :: Linux",
+ "Programming Language :: Python",
+ "Programming Language :: Python :: 3",
+ "Programming Language :: Python :: 3 :: Only",
+ "Programming Language :: Python :: 3.8",
+ "Programming Language :: Python :: 3.9",
+ "Programming Language :: Python :: 3.10",
+ "Programming Language :: Python :: 3.11",
+ "Programming Language :: Python :: Implementation :: CPython",
+ "Topic :: Scientific/Engineering",
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
+ "Topic :: Scientific/Engineering :: Mathematics",
+ "Topic :: Software Development",
+ "Topic :: Software Development :: Libraries",
+ "Topic :: Software Development :: Libraries :: Python Modules",
+ "Typing :: Typed",
+]
+
+[tool.poetry.dependencies]
+python = ">=3.10.0, <3.12.0" # don't change this```
+flwr = { extras = ["simulation"], version = "1.5.0" }
+hydra-core = "1.3.2" # don't change this
+scikit-learn = "1.3.0"
+matplotlib = "3.8.0"
+torch = { url = "https://download.pytorch.org/whl/cu116/torch-1.12.0%2Bcu116-cp310-cp310-linux_x86_64.whl"}
+torchvision = { url = "https://download.pytorch.org/whl/cu116/torchvision-0.13.0%2Bcu116-cp310-cp310-linux_x86_64.whl"}
+
+[tool.poetry.dev-dependencies]
+isort = "==5.11.5"
+black = "==23.1.0"
+docformatter = "==1.5.1"
+mypy = "==1.4.1"
+pylint = "==2.8.2"
+flake8 = "==3.9.2"
+pytest = "==6.2.4"
+pytest-watch = "==4.2.0"
+ruff = "==0.0.272"
+types-requests = "==2.27.7"
+
+[tool.isort]
+line_length = 88
+indent = " "
+multi_line_output = 3
+include_trailing_comma = true
+force_grid_wrap = 0
+use_parentheses = true
+
+[tool.black]
+line-length = 88
+target-version = ["py38", "py39", "py310", "py311"]
+
+[tool.pytest.ini_options]
+minversion = "6.2"
+addopts = "-qq"
+testpaths = [
+ "flwr_baselines",
+]
+
+[tool.mypy]
+ignore_missing_imports = true
+strict = false
+plugins = "numpy.typing.mypy_plugin"
+
+[tool.pylint."MESSAGES CONTROL"]
+disable = "bad-continuation,duplicate-code,too-few-public-methods,useless-import-alias"
+good-names = "i,j,k,_,x,y,X,Y,K,N,X_train,X_test,fc,l1,l2,l3,h,lr,mu"
+max-args = 10
+max-attributes = 15
+max-locals = 36
+max-branches = 20
+max-statements = 55
+signature-mutators="hydra.main.main"
+
+[tool.pylint.typecheck]
+generated-members="numpy.*, torch.*, tensorflow.*"
+
+[[tool.mypy.overrides]]
+module = [
+ "importlib.metadata.*",
+ "importlib_metadata.*",
+]
+follow_imports = "skip"
+follow_imports_for_stubs = true
+disallow_untyped_calls = false
+
+[[tool.mypy.overrides]]
+module = "torch.*"
+follow_imports = "skip"
+follow_imports_for_stubs = true
+
+[tool.docformatter]
+wrap-summaries = 88
+wrap-descriptions = 88
+
+[tool.ruff]
+target-version = "py38"
+line-length = 88
+select = ["D", "E", "F", "W", "B", "ISC", "C4"]
+fixable = ["D", "E", "F", "W", "B", "ISC", "C4"]
+ignore = ["B024", "B027"]
+exclude = [
+ ".bzr",
+ ".direnv",
+ ".eggs",
+ ".git",
+ ".hg",
+ ".mypy_cache",
+ ".nox",
+ ".pants.d",
+ ".pytype",
+ ".ruff_cache",
+ ".svn",
+ ".tox",
+ ".venv",
+ "__pypackages__",
+ "_build",
+ "buck-out",
+ "build",
+ "dist",
+ "node_modules",
+ "venv",
+ "proto",
+]
+
+[tool.ruff.pydocstyle]
+convention = "numpy"
diff --git a/dev/aws-ami-bootstrap-tf.sh b/dev/aws-ami-bootstrap-tf.sh
index bece7d21f1a0..8799a254cbcc 100755
--- a/dev/aws-ami-bootstrap-tf.sh
+++ b/dev/aws-ami-bootstrap-tf.sh
@@ -27,7 +27,7 @@ sudo apt-get install -y make build-essential libssl-dev zlib1g-dev libbz2-dev li
sudo apt install -y python3.7 python3-pip
# Install project dependencies
-python3.7 -m pip install -U pip==23.1.2 setuptools==68.0.0
+python3.7 -m pip install -U pip==23.3.1 setuptools==68.2.2
python3.7 -m pip install -U numpy==1.18.1 grpcio==1.27.2 google==2.0.3 protobuf==3.12.1 \
boto3==1.12.36 boto3_type_annotations==0.3.1 paramiko==2.7.1 docker==4.2.0 matplotlib==3.2.1 \
tensorflow-cpu==2.6.2
diff --git a/dev/aws-ami-bootstrap-torch.sh b/dev/aws-ami-bootstrap-torch.sh
index 1c44cb09673d..835a3994c28a 100755
--- a/dev/aws-ami-bootstrap-torch.sh
+++ b/dev/aws-ami-bootstrap-torch.sh
@@ -27,7 +27,7 @@ sudo apt-get install -y make build-essential libssl-dev zlib1g-dev libbz2-dev li
sudo apt install -y python3.7 python3-pip
# Install project dependencies
-python3.7 -m pip install -U pip==23.1.2 setuptools==68.0.0
+python3.7 -m pip install -U pip==23.3.1 setuptools==68.2.2
python3.7 -m pip install -U numpy==1.18.1 grpcio==1.27.2 google==2.0.3 protobuf==3.12.1 \
boto3==1.12.36 boto3_type_annotations==0.3.1 paramiko==2.7.1 docker==4.2.0 matplotlib==3.2.1 \
tqdm==4.48.2 torch==1.6.0 torchvision==0.7.0
diff --git a/dev/bootstrap.sh b/dev/bootstrap.sh
index 4451115cc151..1700c3774767 100755
--- a/dev/bootstrap.sh
+++ b/dev/bootstrap.sh
@@ -9,8 +9,8 @@ cd "$( cd "$( dirname "${BASH_SOURCE[0]}" )" >/dev/null 2>&1 && pwd )"/../
./dev/rm-caches.sh
# Upgrade/install spcific versions of `pip`, `setuptools`, and `poetry`
-python -m pip install -U pip==23.1.2
-python -m pip install -U setuptools==68.0.0
+python -m pip install -U pip==23.3.1
+python -m pip install -U setuptools==68.2.2
python -m pip install -U poetry==1.5.1
# Use `poetry` to install project dependencies
diff --git a/doc/source/ref-changelog.md b/doc/source/ref-changelog.md
index d2978bac0213..06e77fefedf0 100644
--- a/doc/source/ref-changelog.md
+++ b/doc/source/ref-changelog.md
@@ -2,11 +2,17 @@
## Unreleased
+### What's new?
+
+- **Add experimental support for Python 3.12** ([#2565](https://github.com/adap/flower/pull/2565))
+
- **Support custom** `ClientManager` **in** `start_driver()` ([#2292](https://github.com/adap/flower/pull/2292))
- **Update REST API to support create and delete nodes** ([#2283](https://github.com/adap/flower/pull/2283))
-### What's new?
+- **Update the C++ SDK** ([#2537](https://github/com/adap/flower/pull/2537), [#2528](https://github/com/adap/flower/pull/2528), [#2523](https://github.com/adap/flower/pull/2523), [#2522](https://github.com/adap/flower/pull/2522))
+
+ Add gRPC request-response capability to the C++ SDK.
- **Fix the incorrect return types of Strategy** ([#2432](https://github.com/adap/flower/pull/2432/files))
@@ -28,13 +34,23 @@
- FedMeta [#2438](https://github.com/adap/flower/pull/2438)
+ - FjORD [#2431](https://github.com/adap/flower/pull/2431)
+
+ - MOON [#2421](https://github.com/adap/flower/pull/2421)
+
+ - DepthFL [#2295](https://github.com/adap/flower/pull/2295)
+
+ - FedPer [#2266](https://github.com/adap/flower/pull/2266)
+
+ - FedWav2vec [#2551](https://github.com/adap/flower/pull/2551)
+
- **Update Flower Examples** ([#2384](https://github.com/adap/flower/pull/2384),[#2425](https://github.com/adap/flower/pull/2425), [#2526](https://github.com/adap/flower/pull/2526))
- **General updates to baselines** ([#2301](https://github.com/adap/flower/pull/2301), [#2305](https://github.com/adap/flower/pull/2305), [#2307](https://github.com/adap/flower/pull/2307), [#2327](https://github.com/adap/flower/pull/2327), [#2435](https://github.com/adap/flower/pull/2435))
- **General updates to the simulation engine** ([#2331](https://github.com/adap/flower/pull/2331), [#2447](https://github.com/adap/flower/pull/2447), [#2448](https://github.com/adap/flower/pull/2448))
-- **General improvements** ([#2309](https://github.com/adap/flower/pull/2309), [#2310](https://github.com/adap/flower/pull/2310), [2313](https://github.com/adap/flower/pull/2313), [#2316](https://github.com/adap/flower/pull/2316), [2317](https://github.com/adap/flower/pull/2317),[#2349](https://github.com/adap/flower/pull/2349), [#2360](https://github.com/adap/flower/pull/2360), [#2402](https://github.com/adap/flower/pull/2402), [#2446](https://github.com/adap/flower/pull/2446))
+- **General improvements** ([#2309](https://github.com/adap/flower/pull/2309), [#2310](https://github.com/adap/flower/pull/2310), [2313](https://github.com/adap/flower/pull/2313), [#2316](https://github.com/adap/flower/pull/2316), [2317](https://github.com/adap/flower/pull/2317),[#2349](https://github.com/adap/flower/pull/2349), [#2360](https://github.com/adap/flower/pull/2360), [#2402](https://github.com/adap/flower/pull/2402), [#2446](https://github.com/adap/flower/pull/2446) [#2561](https://github.com/adap/flower/pull/2561))
Flower received many improvements under the hood, too many to list here.
diff --git a/examples/quickstart-cpp/CMakeLists.txt b/examples/quickstart-cpp/CMakeLists.txt
index 79af6a0ef17e..552132b079c9 100644
--- a/examples/quickstart-cpp/CMakeLists.txt
+++ b/examples/quickstart-cpp/CMakeLists.txt
@@ -3,7 +3,6 @@ project(SimpleCppFlowerClient VERSION 0.10
DESCRIPTION "Creates a Simple C++ Flower client that trains a linear model on synthetic data."
LANGUAGES CXX)
set(CMAKE_CXX_STANDARD 17)
-set(ABSL_PROPAGATE_CXX_STD ON)
######################
### Download gRPC
@@ -27,62 +26,27 @@ else()
set(_GRPC_CPP_PLUGIN_EXECUTABLE $)
endif()
-
######################
-### FLWR_GRPC_PROTO
-
-get_filename_component(FLWR_PROTO "../../src/proto/flwr/proto/transport.proto" ABSOLUTE)
-get_filename_component(FLWR_PROTO_PATH "${FLWR_PROTO}" PATH)
-
-set(FLWR_PROTO_SRCS "${CMAKE_CURRENT_BINARY_DIR}/transport.pb.cc")
-set(FLWR_PROTO_HDRS "${CMAKE_CURRENT_BINARY_DIR}/transport.pb.h")
-set(FLWR_GRPC_SRCS "${CMAKE_CURRENT_BINARY_DIR}/transport.grpc.pb.cc")
-set(FLAR_GRPC_HDRS "${CMAKE_CURRENT_BINARY_DIR}/transport.grpc.pb.h")
+### FLWR_LIB
-# External building command to generate gRPC source files.
-add_custom_command(
- OUTPUT "${FLWR_PROTO_SRCS}" "${FLWR_PROTO_HDRS}" "${FLWR_GRPC_SRCS}" "${FLWR_GRPC_HDRS}"
- COMMAND ${_PROTOBUF_PROTOC}
- ARGS --grpc_out "${CMAKE_CURRENT_BINARY_DIR}"
- --cpp_out "${CMAKE_CURRENT_BINARY_DIR}"
- -I "${FLWR_PROTO_PATH}"
- --plugin=protoc-gen-grpc="${_GRPC_CPP_PLUGIN_EXECUTABLE}"
- "${FLWR_PROTO}"
- DEPENDS "${FLWR_PROTO}"
-)
+set(FLWR_SDK_PATH "../../src/cc/flwr")
-add_library(flwr_grpc_proto
- ${FLWR_GRPC_SRCS}
- ${FLWR_GRPC_HDRS}
- ${FLWR_PROTO_SRCS}
- ${FLWR_PROTO_HDRS}
-)
+file(GLOB FLWR_SRCS "${FLWR_SDK_PATH}/src/*.cc")
+file(GLOB FLWR_PROTO_SRCS "${FLWR_SDK_PATH}/include/flwr/proto/*.cc")
+set(FLWR_INCLUDE_DIR "${FLWR_SDK_PATH}/include")
-target_include_directories(flwr_grpc_proto PUBLIC ${CMAKE_CURRENT_BINARY_DIR})
+add_library(flwr ${FLWR_SRCS} ${FLWR_PROTO_SRCS})
-target_link_libraries(flwr_grpc_proto
+target_link_libraries(flwr
${_REFLECTION}
${_GRPC_GRPCPP}
${_PROTOBUF_LIBPROTOBUF}
)
-######################
-### FLWR_LIB
-
-file(GLOB FLWR_SRCS "../../src/cc/flwr/src/*.cc")
-set(FLWR_INCLUDE_DIR "../../src/cc/flwr/include")
-
-add_library(flwr ${FLWR_SRCS})
-
target_include_directories(flwr PUBLIC
- ${CMAKE_CURRENT_BINARY_DIR}
${FLWR_INCLUDE_DIR}
)
-target_link_libraries(flwr
- flwr_grpc_proto
-)
-
######################
### FLWR_CLIENT
file(GLOB FLWR_CLIENT_SRCS src/*.cc)
diff --git a/examples/quickstart-cpp/driver.py b/examples/quickstart-cpp/driver.py
new file mode 100644
index 000000000000..037623ee77cf
--- /dev/null
+++ b/examples/quickstart-cpp/driver.py
@@ -0,0 +1,10 @@
+import flwr as fl
+from fedavg_cpp import FedAvgCpp
+
+# Start Flower server for three rounds of federated learning
+if __name__ == "__main__":
+ fl.driver.start_driver(
+ server_address="0.0.0.0:9091",
+ config=fl.server.ServerConfig(num_rounds=3),
+ strategy=FedAvgCpp(),
+ )
diff --git a/examples/quickstart-cpp/include/simple_client.h b/examples/quickstart-cpp/include/simple_client.h
index ce598365f29c..894ecb267387 100644
--- a/examples/quickstart-cpp/include/simple_client.h
+++ b/examples/quickstart-cpp/include/simple_client.h
@@ -1,6 +1,6 @@
/***********************************************************************************************************
*
- * @file libtorch_client.h
+ * @file simple_client.h
*
* @brief Define an example flower client, train and test method
*
diff --git a/examples/quickstart-cpp/src/main.cc b/examples/quickstart-cpp/src/main.cc
index fb3c533a3841..f294f9d69473 100644
--- a/examples/quickstart-cpp/src/main.cc
+++ b/examples/quickstart-cpp/src/main.cc
@@ -2,44 +2,58 @@
#include "start.h"
int main(int argc, char **argv) {
- if (argc != 3) {
- std::cout << "Client takes three arguments as follows: " << std::endl;
- std::cout << "./client CLIENT_ID SERVER_URL" << std::endl;
- std::cout << "Example: ./flwr_client 0 '127.0.0.1:8080'" << std::endl;
- return 0;
- }
-
- // Parsing arguments
- const std::string CLIENT_ID = argv[1];
- const std::string SERVER_URL = argv[2];
-
- // Populate local datasets
- std::vector ms{3.5, 9.3}; // b + m_0*x0 + m_1*x1
- double b = 1.7;
- std::cout <<"Training set:" << std::endl;
- SyntheticDataset local_training_data = SyntheticDataset(ms, b, 1000);
- std::cout << std::endl;
-
- std::cout <<"Validation set:" << std::endl;
- SyntheticDataset local_validation_data = SyntheticDataset(ms, b, 100);
- std::cout << std::endl;
-
- std::cout <<"Test set:" << std::endl;
- SyntheticDataset local_test_data = SyntheticDataset(ms, b, 500);
- std::cout << std::endl;
-
- // Define a model
- LineFitModel model = LineFitModel(500, 0.01, ms.size());
-
- // Initialize TorchClient
- SimpleFlwrClient client(CLIENT_ID, model, local_training_data, local_validation_data, local_test_data);
-
- // Define a server address
- std::string server_add = SERVER_URL;
-
- // Start client
+ if (argc != 3 && argc != 4) {
+ std::cout << "Client takes three mandatory arguments and one optional as "
+ "follows: "
+ << std::endl;
+ std::cout << "./client CLIENT_ID SERVER_URL [GRPC_MODE]" << std::endl;
+ std::cout
+ << "GRPC_MODE is optional and can be either 'bidi' (default) or 'rere'."
+ << std::endl;
+ std::cout << "Example: ./flwr_client 0 '127.0.0.1:8080' bidi" << std::endl;
+ std::cout << "This is the same as: ./flwr_client 0 '127.0.0.1:8080'"
+ << std::endl;
+ return 0;
+ }
+
+ // Parsing arguments
+ const std::string CLIENT_ID = argv[1];
+ const std::string SERVER_URL = argv[2];
+
+ // Populate local datasets
+ std::vector ms{3.5, 9.3}; // b + m_0*x0 + m_1*x1
+ double b = 1.7;
+ std::cout << "Training set:" << std::endl;
+ SyntheticDataset local_training_data = SyntheticDataset(ms, b, 1000);
+ std::cout << std::endl;
+
+ std::cout << "Validation set:" << std::endl;
+ SyntheticDataset local_validation_data = SyntheticDataset(ms, b, 100);
+ std::cout << std::endl;
+
+ std::cout << "Test set:" << std::endl;
+ SyntheticDataset local_test_data = SyntheticDataset(ms, b, 500);
+ std::cout << std::endl;
+
+ // Define a model
+ LineFitModel model = LineFitModel(500, 0.01, ms.size());
+
+ // Initialize TorchClient
+ SimpleFlwrClient client(CLIENT_ID, model, local_training_data,
+ local_validation_data, local_test_data);
+
+ // Define a server address
+ std::string server_add = SERVER_URL;
+
+ if (argc == 4 && std::string(argv[3]) == "rere") {
+ std::cout << "Starting rere client" << std::endl;
+ // Start rere client
+ start::start_rere_client(server_add, &client);
+ } else {
+ std::cout << "Starting bidi client" << std::endl;
+ // Start bidi client
start::start_client(server_add, &client);
+ }
- return 0;
+ return 0;
}
-
diff --git a/pyproject.toml b/pyproject.toml
index b948c8d8b64d..261eacbf0c94 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -34,6 +34,7 @@ classifiers = [
"Programming Language :: Python :: 3.9",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
+ "Programming Language :: Python :: 3.12",
"Programming Language :: Python :: Implementation :: CPython",
"Topic :: Scientific/Engineering",
"Topic :: Scientific/Engineering :: Artificial Intelligence",
@@ -81,14 +82,14 @@ rest = ["requests", "starlette", "uvicorn"]
types-dataclasses = "==0.6.6"
types-protobuf = "==3.19.18"
types-requests = "==2.31.0.2"
-types-setuptools = "==68.0.0.3"
+types-setuptools = "==68.2.0.0"
clang-format = "==16.0.6"
isort = "==5.11.5"
black = { version = "==23.3.0", extras = ["jupyter"] }
docformatter = "==1.7.5"
mypy = "==1.5.1"
pylint = "==2.13.9"
-flake8 = "==3.9.2"
+flake8 = "==5.0.4"
pytest = "==7.4.0"
pytest-cov = "==3.0.0"
pytest-watch = "==4.2.0"
diff --git a/src/cc/flwr/CMakeLists.txt b/src/cc/flwr/CMakeLists.txt
index 8ab7dc4c2964..c242f52b237b 100644
--- a/src/cc/flwr/CMakeLists.txt
+++ b/src/cc/flwr/CMakeLists.txt
@@ -4,7 +4,6 @@ project(flwr VERSION 1.0
LANGUAGES CXX)
set(CMAKE_CXX_STANDARD 17)
-set(ABSL_PROPAGATE_CXX_STD ON)
# Assume gRPC and other dependencies are necessary
include(FetchContent)
@@ -26,34 +25,56 @@ else()
set(_GRPC_CPP_PLUGIN_EXECUTABLE $)
endif()
-# FLWR_GRPC_PROTO
-
-get_filename_component(FLWR_PROTO "../../proto/flwr/proto/transport.proto" ABSOLUTE)
-get_filename_component(FLWR_PROTO_PATH "${FLWR_PROTO}" PATH)
-
-set(FLWR_PROTO_SRCS "${CMAKE_CURRENT_BINARY_DIR}/transport.pb.cc")
-set(FLWR_PROTO_HDRS "${CMAKE_CURRENT_BINARY_DIR}/transport.pb.h")
-set(FLWR_GRPC_SRCS "${CMAKE_CURRENT_BINARY_DIR}/transport.grpc.pb.cc")
-set(FLAR_GRPC_HDRS "${CMAKE_CURRENT_BINARY_DIR}/transport.grpc.pb.h")
-
-# External building command to generate gRPC source files.
-add_custom_command(
- OUTPUT "${FLWR_PROTO_SRCS}" "${FLWR_PROTO_HDRS}" "${FLWR_GRPC_SRCS}" "${FLWR_GRPC_HDRS}"
- COMMAND ${_PROTOBUF_PROTOC}
- ARGS --grpc_out "${CMAKE_CURRENT_BINARY_DIR}"
- --cpp_out "${CMAKE_CURRENT_BINARY_DIR}"
- -I "${FLWR_PROTO_PATH}"
- --plugin=protoc-gen-grpc="${_GRPC_CPP_PLUGIN_EXECUTABLE}"
- "${FLWR_PROTO}"
- DEPENDS "${FLWR_PROTO}"
-)
-
-add_library(flwr_grpc_proto STATIC
- ${FLWR_GRPC_SRCS}
- ${FLWR_GRPC_HDRS}
- ${FLWR_PROTO_SRCS}
- ${FLWR_PROTO_HDRS}
-)
+# Paths and output directories
+get_filename_component(FLWR_PROTO_BASE_PATH "../../proto/" ABSOLUTE)
+set(INCLUDE_FLWR_PROTO_DIR "${CMAKE_CURRENT_SOURCE_DIR}/include/flwr/proto")
+
+# Generate source files and copy them
+macro(GENERATE_AND_COPY PROTO_NAME)
+ set(OUT_PROTO_SRCS "${CMAKE_CURRENT_BINARY_DIR}/flwr/proto/${PROTO_NAME}.pb.cc")
+ set(OUT_PROTO_HDRS "${CMAKE_CURRENT_BINARY_DIR}/flwr/proto/${PROTO_NAME}.pb.h")
+ set(OUT_GRPC_SRCS "${CMAKE_CURRENT_BINARY_DIR}/flwr/proto/${PROTO_NAME}.grpc.pb.cc")
+ set(OUT_GRPC_HDRS "${CMAKE_CURRENT_BINARY_DIR}/flwr/proto/${PROTO_NAME}.grpc.pb.h")
+ set(SOURCE_PROTO "${FLWR_PROTO_BASE_PATH}/flwr/proto/${PROTO_NAME}.proto")
+
+ add_custom_command(
+ OUTPUT "${OUT_PROTO_SRCS}" "${OUT_PROTO_HDRS}" "${OUT_GRPC_SRCS}" "${OUT_GRPC_HDRS}"
+ COMMAND ${_PROTOBUF_PROTOC}
+ ARGS --grpc_out "${CMAKE_CURRENT_BINARY_DIR}"
+ --cpp_out "${CMAKE_CURRENT_BINARY_DIR}"
+ -I "${FLWR_PROTO_BASE_PATH}"
+ --plugin=protoc-gen-grpc="${_GRPC_CPP_PLUGIN_EXECUTABLE}"
+ "${SOURCE_PROTO}"
+ )
+
+ add_custom_command(
+ OUTPUT "${INCLUDE_FLWR_PROTO_DIR}/${PROTO_NAME}.pb.cc"
+ "${INCLUDE_FLWR_PROTO_DIR}/${PROTO_NAME}.pb.h"
+ "${INCLUDE_FLWR_PROTO_DIR}/${PROTO_NAME}.grpc.pb.cc"
+ "${INCLUDE_FLWR_PROTO_DIR}/${PROTO_NAME}.grpc.pb.h"
+ COMMAND ${CMAKE_COMMAND} -E copy_if_different
+ "${OUT_PROTO_SRCS}" "${OUT_PROTO_HDRS}" "${OUT_GRPC_SRCS}" "${OUT_GRPC_HDRS}"
+ "${INCLUDE_FLWR_PROTO_DIR}"
+ DEPENDS "${OUT_PROTO_SRCS}" "${OUT_PROTO_HDRS}" "${OUT_GRPC_SRCS}" "${OUT_GRPC_HDRS}"
+ )
+
+ set(ALL_PROTO_FILES
+ ${ALL_PROTO_FILES}
+ "${INCLUDE_FLWR_PROTO_DIR}/${PROTO_NAME}.pb.cc"
+ "${INCLUDE_FLWR_PROTO_DIR}/${PROTO_NAME}.pb.h"
+ "${INCLUDE_FLWR_PROTO_DIR}/${PROTO_NAME}.grpc.pb.cc"
+ "${INCLUDE_FLWR_PROTO_DIR}/${PROTO_NAME}.grpc.pb.h"
+ CACHE INTERNAL "All generated proto files"
+ )
+endmacro()
+
+# Using the above macro for all proto files
+GENERATE_AND_COPY(transport)
+GENERATE_AND_COPY(node)
+GENERATE_AND_COPY(task)
+GENERATE_AND_COPY(fleet)
+
+add_library(flwr_grpc_proto STATIC ${ALL_PROTO_FILES})
target_include_directories(flwr_grpc_proto
PUBLIC
@@ -67,56 +88,14 @@ target_link_libraries(flwr_grpc_proto
${_GRPC_GRPCPP}
${_PROTOBUF_LIBPROTOBUF}
)
+
# For the internal use of flwr
file(GLOB FLWR_SRCS "src/*.cc")
-
add_library(flwr ${FLWR_SRCS})
target_include_directories(flwr PUBLIC
$
- $
)
# Link gRPC and other dependencies
target_link_libraries(flwr PRIVATE flwr_grpc_proto)
-
-# Merge the two libraries
-add_library(flwr_merged STATIC $ $)
-
-target_include_directories(flwr_merged PUBLIC
- $
- $
-)
-
-# This will create a 'flwrConfig.cmake' for users to find
-install(TARGETS flwr_merged EXPORT flwrTargets
- LIBRARY DESTINATION lib
- ARCHIVE DESTINATION lib
- RUNTIME DESTINATION bin
- PUBLIC_HEADER DESTINATION include
-)
-install(
- FILES
- ${CMAKE_CURRENT_BINARY_DIR}/transport.grpc.pb.h
- ${CMAKE_CURRENT_BINARY_DIR}/transport.pb.h
- DESTINATION include
-)
-install(DIRECTORY include/ DESTINATION include)
-
-install(EXPORT flwrTargets
- FILE flwrConfig.cmake
- NAMESPACE flwr::
- DESTINATION lib/cmake/flwr
-)
-
-# Optional: Generate and install package version file
-include(CMakePackageConfigHelpers)
-write_basic_package_version_file(
- "${CMAKE_CURRENT_BINARY_DIR}/flwrConfigVersion.cmake"
- VERSION ${PROJECT_VERSION}
- COMPATIBILITY AnyNewerVersion
-)
-install(FILES "${CMAKE_CURRENT_BINARY_DIR}/flwrConfigVersion.cmake"
- DESTINATION lib/cmake/flwr
-)
-
diff --git a/src/cc/flwr/include/flwr/proto/fleet.grpc.pb.cc b/src/cc/flwr/include/flwr/proto/fleet.grpc.pb.cc
new file mode 100644
index 000000000000..c71a6a3e1c45
--- /dev/null
+++ b/src/cc/flwr/include/flwr/proto/fleet.grpc.pb.cc
@@ -0,0 +1,214 @@
+// Generated by the gRPC C++ plugin.
+// If you make any local change, they will be lost.
+// source: flwr/proto/fleet.proto
+
+#include "flwr/proto/fleet.pb.h"
+#include "flwr/proto/fleet.grpc.pb.h"
+
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+namespace flwr {
+namespace proto {
+
+static const char* Fleet_method_names[] = {
+ "/flwr.proto.Fleet/CreateNode",
+ "/flwr.proto.Fleet/DeleteNode",
+ "/flwr.proto.Fleet/PullTaskIns",
+ "/flwr.proto.Fleet/PushTaskRes",
+};
+
+std::unique_ptr< Fleet::Stub> Fleet::NewStub(const std::shared_ptr< ::grpc::ChannelInterface>& channel, const ::grpc::StubOptions& options) {
+ (void)options;
+ std::unique_ptr< Fleet::Stub> stub(new Fleet::Stub(channel, options));
+ return stub;
+}
+
+Fleet::Stub::Stub(const std::shared_ptr< ::grpc::ChannelInterface>& channel, const ::grpc::StubOptions& options)
+ : channel_(channel), rpcmethod_CreateNode_(Fleet_method_names[0], options.suffix_for_stats(),::grpc::internal::RpcMethod::NORMAL_RPC, channel)
+ , rpcmethod_DeleteNode_(Fleet_method_names[1], options.suffix_for_stats(),::grpc::internal::RpcMethod::NORMAL_RPC, channel)
+ , rpcmethod_PullTaskIns_(Fleet_method_names[2], options.suffix_for_stats(),::grpc::internal::RpcMethod::NORMAL_RPC, channel)
+ , rpcmethod_PushTaskRes_(Fleet_method_names[3], options.suffix_for_stats(),::grpc::internal::RpcMethod::NORMAL_RPC, channel)
+ {}
+
+::grpc::Status Fleet::Stub::CreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::flwr::proto::CreateNodeResponse* response) {
+ return ::grpc::internal::BlockingUnaryCall< ::flwr::proto::CreateNodeRequest, ::flwr::proto::CreateNodeResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(channel_.get(), rpcmethod_CreateNode_, context, request, response);
+}
+
+void Fleet::Stub::async::CreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest* request, ::flwr::proto::CreateNodeResponse* response, std::function f) {
+ ::grpc::internal::CallbackUnaryCall< ::flwr::proto::CreateNodeRequest, ::flwr::proto::CreateNodeResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(stub_->channel_.get(), stub_->rpcmethod_CreateNode_, context, request, response, std::move(f));
+}
+
+void Fleet::Stub::async::CreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest* request, ::flwr::proto::CreateNodeResponse* response, ::grpc::ClientUnaryReactor* reactor) {
+ ::grpc::internal::ClientCallbackUnaryFactory::Create< ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(stub_->channel_.get(), stub_->rpcmethod_CreateNode_, context, request, response, reactor);
+}
+
+::grpc::ClientAsyncResponseReader< ::flwr::proto::CreateNodeResponse>* Fleet::Stub::PrepareAsyncCreateNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ return ::grpc::internal::ClientAsyncResponseReaderHelper::Create< ::flwr::proto::CreateNodeResponse, ::flwr::proto::CreateNodeRequest, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(channel_.get(), cq, rpcmethod_CreateNode_, context, request);
+}
+
+::grpc::ClientAsyncResponseReader< ::flwr::proto::CreateNodeResponse>* Fleet::Stub::AsyncCreateNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ auto* result =
+ this->PrepareAsyncCreateNodeRaw(context, request, cq);
+ result->StartCall();
+ return result;
+}
+
+::grpc::Status Fleet::Stub::DeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::flwr::proto::DeleteNodeResponse* response) {
+ return ::grpc::internal::BlockingUnaryCall< ::flwr::proto::DeleteNodeRequest, ::flwr::proto::DeleteNodeResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(channel_.get(), rpcmethod_DeleteNode_, context, request, response);
+}
+
+void Fleet::Stub::async::DeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest* request, ::flwr::proto::DeleteNodeResponse* response, std::function f) {
+ ::grpc::internal::CallbackUnaryCall< ::flwr::proto::DeleteNodeRequest, ::flwr::proto::DeleteNodeResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(stub_->channel_.get(), stub_->rpcmethod_DeleteNode_, context, request, response, std::move(f));
+}
+
+void Fleet::Stub::async::DeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest* request, ::flwr::proto::DeleteNodeResponse* response, ::grpc::ClientUnaryReactor* reactor) {
+ ::grpc::internal::ClientCallbackUnaryFactory::Create< ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(stub_->channel_.get(), stub_->rpcmethod_DeleteNode_, context, request, response, reactor);
+}
+
+::grpc::ClientAsyncResponseReader< ::flwr::proto::DeleteNodeResponse>* Fleet::Stub::PrepareAsyncDeleteNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ return ::grpc::internal::ClientAsyncResponseReaderHelper::Create< ::flwr::proto::DeleteNodeResponse, ::flwr::proto::DeleteNodeRequest, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(channel_.get(), cq, rpcmethod_DeleteNode_, context, request);
+}
+
+::grpc::ClientAsyncResponseReader< ::flwr::proto::DeleteNodeResponse>* Fleet::Stub::AsyncDeleteNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ auto* result =
+ this->PrepareAsyncDeleteNodeRaw(context, request, cq);
+ result->StartCall();
+ return result;
+}
+
+::grpc::Status Fleet::Stub::PullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::flwr::proto::PullTaskInsResponse* response) {
+ return ::grpc::internal::BlockingUnaryCall< ::flwr::proto::PullTaskInsRequest, ::flwr::proto::PullTaskInsResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(channel_.get(), rpcmethod_PullTaskIns_, context, request, response);
+}
+
+void Fleet::Stub::async::PullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest* request, ::flwr::proto::PullTaskInsResponse* response, std::function f) {
+ ::grpc::internal::CallbackUnaryCall< ::flwr::proto::PullTaskInsRequest, ::flwr::proto::PullTaskInsResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(stub_->channel_.get(), stub_->rpcmethod_PullTaskIns_, context, request, response, std::move(f));
+}
+
+void Fleet::Stub::async::PullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest* request, ::flwr::proto::PullTaskInsResponse* response, ::grpc::ClientUnaryReactor* reactor) {
+ ::grpc::internal::ClientCallbackUnaryFactory::Create< ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(stub_->channel_.get(), stub_->rpcmethod_PullTaskIns_, context, request, response, reactor);
+}
+
+::grpc::ClientAsyncResponseReader< ::flwr::proto::PullTaskInsResponse>* Fleet::Stub::PrepareAsyncPullTaskInsRaw(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::grpc::CompletionQueue* cq) {
+ return ::grpc::internal::ClientAsyncResponseReaderHelper::Create< ::flwr::proto::PullTaskInsResponse, ::flwr::proto::PullTaskInsRequest, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(channel_.get(), cq, rpcmethod_PullTaskIns_, context, request);
+}
+
+::grpc::ClientAsyncResponseReader< ::flwr::proto::PullTaskInsResponse>* Fleet::Stub::AsyncPullTaskInsRaw(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::grpc::CompletionQueue* cq) {
+ auto* result =
+ this->PrepareAsyncPullTaskInsRaw(context, request, cq);
+ result->StartCall();
+ return result;
+}
+
+::grpc::Status Fleet::Stub::PushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::flwr::proto::PushTaskResResponse* response) {
+ return ::grpc::internal::BlockingUnaryCall< ::flwr::proto::PushTaskResRequest, ::flwr::proto::PushTaskResResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(channel_.get(), rpcmethod_PushTaskRes_, context, request, response);
+}
+
+void Fleet::Stub::async::PushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest* request, ::flwr::proto::PushTaskResResponse* response, std::function f) {
+ ::grpc::internal::CallbackUnaryCall< ::flwr::proto::PushTaskResRequest, ::flwr::proto::PushTaskResResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(stub_->channel_.get(), stub_->rpcmethod_PushTaskRes_, context, request, response, std::move(f));
+}
+
+void Fleet::Stub::async::PushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest* request, ::flwr::proto::PushTaskResResponse* response, ::grpc::ClientUnaryReactor* reactor) {
+ ::grpc::internal::ClientCallbackUnaryFactory::Create< ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(stub_->channel_.get(), stub_->rpcmethod_PushTaskRes_, context, request, response, reactor);
+}
+
+::grpc::ClientAsyncResponseReader< ::flwr::proto::PushTaskResResponse>* Fleet::Stub::PrepareAsyncPushTaskResRaw(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::grpc::CompletionQueue* cq) {
+ return ::grpc::internal::ClientAsyncResponseReaderHelper::Create< ::flwr::proto::PushTaskResResponse, ::flwr::proto::PushTaskResRequest, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(channel_.get(), cq, rpcmethod_PushTaskRes_, context, request);
+}
+
+::grpc::ClientAsyncResponseReader< ::flwr::proto::PushTaskResResponse>* Fleet::Stub::AsyncPushTaskResRaw(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::grpc::CompletionQueue* cq) {
+ auto* result =
+ this->PrepareAsyncPushTaskResRaw(context, request, cq);
+ result->StartCall();
+ return result;
+}
+
+Fleet::Service::Service() {
+ AddMethod(new ::grpc::internal::RpcServiceMethod(
+ Fleet_method_names[0],
+ ::grpc::internal::RpcMethod::NORMAL_RPC,
+ new ::grpc::internal::RpcMethodHandler< Fleet::Service, ::flwr::proto::CreateNodeRequest, ::flwr::proto::CreateNodeResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(
+ [](Fleet::Service* service,
+ ::grpc::ServerContext* ctx,
+ const ::flwr::proto::CreateNodeRequest* req,
+ ::flwr::proto::CreateNodeResponse* resp) {
+ return service->CreateNode(ctx, req, resp);
+ }, this)));
+ AddMethod(new ::grpc::internal::RpcServiceMethod(
+ Fleet_method_names[1],
+ ::grpc::internal::RpcMethod::NORMAL_RPC,
+ new ::grpc::internal::RpcMethodHandler< Fleet::Service, ::flwr::proto::DeleteNodeRequest, ::flwr::proto::DeleteNodeResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(
+ [](Fleet::Service* service,
+ ::grpc::ServerContext* ctx,
+ const ::flwr::proto::DeleteNodeRequest* req,
+ ::flwr::proto::DeleteNodeResponse* resp) {
+ return service->DeleteNode(ctx, req, resp);
+ }, this)));
+ AddMethod(new ::grpc::internal::RpcServiceMethod(
+ Fleet_method_names[2],
+ ::grpc::internal::RpcMethod::NORMAL_RPC,
+ new ::grpc::internal::RpcMethodHandler< Fleet::Service, ::flwr::proto::PullTaskInsRequest, ::flwr::proto::PullTaskInsResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(
+ [](Fleet::Service* service,
+ ::grpc::ServerContext* ctx,
+ const ::flwr::proto::PullTaskInsRequest* req,
+ ::flwr::proto::PullTaskInsResponse* resp) {
+ return service->PullTaskIns(ctx, req, resp);
+ }, this)));
+ AddMethod(new ::grpc::internal::RpcServiceMethod(
+ Fleet_method_names[3],
+ ::grpc::internal::RpcMethod::NORMAL_RPC,
+ new ::grpc::internal::RpcMethodHandler< Fleet::Service, ::flwr::proto::PushTaskResRequest, ::flwr::proto::PushTaskResResponse, ::grpc::protobuf::MessageLite, ::grpc::protobuf::MessageLite>(
+ [](Fleet::Service* service,
+ ::grpc::ServerContext* ctx,
+ const ::flwr::proto::PushTaskResRequest* req,
+ ::flwr::proto::PushTaskResResponse* resp) {
+ return service->PushTaskRes(ctx, req, resp);
+ }, this)));
+}
+
+Fleet::Service::~Service() {
+}
+
+::grpc::Status Fleet::Service::CreateNode(::grpc::ServerContext* context, const ::flwr::proto::CreateNodeRequest* request, ::flwr::proto::CreateNodeResponse* response) {
+ (void) context;
+ (void) request;
+ (void) response;
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+}
+
+::grpc::Status Fleet::Service::DeleteNode(::grpc::ServerContext* context, const ::flwr::proto::DeleteNodeRequest* request, ::flwr::proto::DeleteNodeResponse* response) {
+ (void) context;
+ (void) request;
+ (void) response;
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+}
+
+::grpc::Status Fleet::Service::PullTaskIns(::grpc::ServerContext* context, const ::flwr::proto::PullTaskInsRequest* request, ::flwr::proto::PullTaskInsResponse* response) {
+ (void) context;
+ (void) request;
+ (void) response;
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+}
+
+::grpc::Status Fleet::Service::PushTaskRes(::grpc::ServerContext* context, const ::flwr::proto::PushTaskResRequest* request, ::flwr::proto::PushTaskResResponse* response) {
+ (void) context;
+ (void) request;
+ (void) response;
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+}
+
+
+} // namespace flwr
+} // namespace proto
+
diff --git a/src/cc/flwr/include/flwr/proto/fleet.grpc.pb.h b/src/cc/flwr/include/flwr/proto/fleet.grpc.pb.h
new file mode 100644
index 000000000000..03d445142c37
--- /dev/null
+++ b/src/cc/flwr/include/flwr/proto/fleet.grpc.pb.h
@@ -0,0 +1,747 @@
+// Generated by the gRPC C++ plugin.
+// If you make any local change, they will be lost.
+// source: flwr/proto/fleet.proto
+// Original file comments:
+// Copyright 2022 Flower Labs GmbH. All Rights Reserved.
+//
+// Licensed under the Apache License, Version 2.0 (the "License");
+// you may not use this file except in compliance with the License.
+// You may obtain a copy of the License at
+//
+// http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing, software
+// distributed under the License is distributed on an "AS IS" BASIS,
+// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+// See the License for the specific language governing permissions and
+// limitations under the License.
+// ==============================================================================
+//
+#ifndef GRPC_flwr_2fproto_2ffleet_2eproto__INCLUDED
+#define GRPC_flwr_2fproto_2ffleet_2eproto__INCLUDED
+
+#include "flwr/proto/fleet.pb.h"
+
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+
+namespace flwr {
+namespace proto {
+
+class Fleet final {
+ public:
+ static constexpr char const* service_full_name() {
+ return "flwr.proto.Fleet";
+ }
+ class StubInterface {
+ public:
+ virtual ~StubInterface() {}
+ virtual ::grpc::Status CreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::flwr::proto::CreateNodeResponse* response) = 0;
+ std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::CreateNodeResponse>> AsyncCreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::CreateNodeResponse>>(AsyncCreateNodeRaw(context, request, cq));
+ }
+ std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::CreateNodeResponse>> PrepareAsyncCreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::CreateNodeResponse>>(PrepareAsyncCreateNodeRaw(context, request, cq));
+ }
+ virtual ::grpc::Status DeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::flwr::proto::DeleteNodeResponse* response) = 0;
+ std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::DeleteNodeResponse>> AsyncDeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::DeleteNodeResponse>>(AsyncDeleteNodeRaw(context, request, cq));
+ }
+ std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::DeleteNodeResponse>> PrepareAsyncDeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::DeleteNodeResponse>>(PrepareAsyncDeleteNodeRaw(context, request, cq));
+ }
+ // Retrieve one or more tasks, if possible
+ //
+ // HTTP API path: /api/v1/fleet/pull-task-ins
+ virtual ::grpc::Status PullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::flwr::proto::PullTaskInsResponse* response) = 0;
+ std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PullTaskInsResponse>> AsyncPullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PullTaskInsResponse>>(AsyncPullTaskInsRaw(context, request, cq));
+ }
+ std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PullTaskInsResponse>> PrepareAsyncPullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PullTaskInsResponse>>(PrepareAsyncPullTaskInsRaw(context, request, cq));
+ }
+ // Complete one or more tasks, if possible
+ //
+ // HTTP API path: /api/v1/fleet/push-task-res
+ virtual ::grpc::Status PushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::flwr::proto::PushTaskResResponse* response) = 0;
+ std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PushTaskResResponse>> AsyncPushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PushTaskResResponse>>(AsyncPushTaskResRaw(context, request, cq));
+ }
+ std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PushTaskResResponse>> PrepareAsyncPushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PushTaskResResponse>>(PrepareAsyncPushTaskResRaw(context, request, cq));
+ }
+ class async_interface {
+ public:
+ virtual ~async_interface() {}
+ virtual void CreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest* request, ::flwr::proto::CreateNodeResponse* response, std::function) = 0;
+ virtual void CreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest* request, ::flwr::proto::CreateNodeResponse* response, ::grpc::ClientUnaryReactor* reactor) = 0;
+ virtual void DeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest* request, ::flwr::proto::DeleteNodeResponse* response, std::function) = 0;
+ virtual void DeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest* request, ::flwr::proto::DeleteNodeResponse* response, ::grpc::ClientUnaryReactor* reactor) = 0;
+ // Retrieve one or more tasks, if possible
+ //
+ // HTTP API path: /api/v1/fleet/pull-task-ins
+ virtual void PullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest* request, ::flwr::proto::PullTaskInsResponse* response, std::function) = 0;
+ virtual void PullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest* request, ::flwr::proto::PullTaskInsResponse* response, ::grpc::ClientUnaryReactor* reactor) = 0;
+ // Complete one or more tasks, if possible
+ //
+ // HTTP API path: /api/v1/fleet/push-task-res
+ virtual void PushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest* request, ::flwr::proto::PushTaskResResponse* response, std::function) = 0;
+ virtual void PushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest* request, ::flwr::proto::PushTaskResResponse* response, ::grpc::ClientUnaryReactor* reactor) = 0;
+ };
+ typedef class async_interface experimental_async_interface;
+ virtual class async_interface* async() { return nullptr; }
+ class async_interface* experimental_async() { return async(); }
+ private:
+ virtual ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::CreateNodeResponse>* AsyncCreateNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::grpc::CompletionQueue* cq) = 0;
+ virtual ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::CreateNodeResponse>* PrepareAsyncCreateNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::grpc::CompletionQueue* cq) = 0;
+ virtual ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::DeleteNodeResponse>* AsyncDeleteNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::grpc::CompletionQueue* cq) = 0;
+ virtual ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::DeleteNodeResponse>* PrepareAsyncDeleteNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::grpc::CompletionQueue* cq) = 0;
+ virtual ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PullTaskInsResponse>* AsyncPullTaskInsRaw(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::grpc::CompletionQueue* cq) = 0;
+ virtual ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PullTaskInsResponse>* PrepareAsyncPullTaskInsRaw(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::grpc::CompletionQueue* cq) = 0;
+ virtual ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PushTaskResResponse>* AsyncPushTaskResRaw(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::grpc::CompletionQueue* cq) = 0;
+ virtual ::grpc::ClientAsyncResponseReaderInterface< ::flwr::proto::PushTaskResResponse>* PrepareAsyncPushTaskResRaw(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::grpc::CompletionQueue* cq) = 0;
+ };
+ class Stub final : public StubInterface {
+ public:
+ Stub(const std::shared_ptr< ::grpc::ChannelInterface>& channel, const ::grpc::StubOptions& options = ::grpc::StubOptions());
+ ::grpc::Status CreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::flwr::proto::CreateNodeResponse* response) override;
+ std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::CreateNodeResponse>> AsyncCreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::CreateNodeResponse>>(AsyncCreateNodeRaw(context, request, cq));
+ }
+ std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::CreateNodeResponse>> PrepareAsyncCreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::CreateNodeResponse>>(PrepareAsyncCreateNodeRaw(context, request, cq));
+ }
+ ::grpc::Status DeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::flwr::proto::DeleteNodeResponse* response) override;
+ std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::DeleteNodeResponse>> AsyncDeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::DeleteNodeResponse>>(AsyncDeleteNodeRaw(context, request, cq));
+ }
+ std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::DeleteNodeResponse>> PrepareAsyncDeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::DeleteNodeResponse>>(PrepareAsyncDeleteNodeRaw(context, request, cq));
+ }
+ ::grpc::Status PullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::flwr::proto::PullTaskInsResponse* response) override;
+ std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::PullTaskInsResponse>> AsyncPullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::PullTaskInsResponse>>(AsyncPullTaskInsRaw(context, request, cq));
+ }
+ std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::PullTaskInsResponse>> PrepareAsyncPullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::PullTaskInsResponse>>(PrepareAsyncPullTaskInsRaw(context, request, cq));
+ }
+ ::grpc::Status PushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::flwr::proto::PushTaskResResponse* response) override;
+ std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::PushTaskResResponse>> AsyncPushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::PushTaskResResponse>>(AsyncPushTaskResRaw(context, request, cq));
+ }
+ std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::PushTaskResResponse>> PrepareAsyncPushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::grpc::CompletionQueue* cq) {
+ return std::unique_ptr< ::grpc::ClientAsyncResponseReader< ::flwr::proto::PushTaskResResponse>>(PrepareAsyncPushTaskResRaw(context, request, cq));
+ }
+ class async final :
+ public StubInterface::async_interface {
+ public:
+ void CreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest* request, ::flwr::proto::CreateNodeResponse* response, std::function) override;
+ void CreateNode(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest* request, ::flwr::proto::CreateNodeResponse* response, ::grpc::ClientUnaryReactor* reactor) override;
+ void DeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest* request, ::flwr::proto::DeleteNodeResponse* response, std::function) override;
+ void DeleteNode(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest* request, ::flwr::proto::DeleteNodeResponse* response, ::grpc::ClientUnaryReactor* reactor) override;
+ void PullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest* request, ::flwr::proto::PullTaskInsResponse* response, std::function) override;
+ void PullTaskIns(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest* request, ::flwr::proto::PullTaskInsResponse* response, ::grpc::ClientUnaryReactor* reactor) override;
+ void PushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest* request, ::flwr::proto::PushTaskResResponse* response, std::function) override;
+ void PushTaskRes(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest* request, ::flwr::proto::PushTaskResResponse* response, ::grpc::ClientUnaryReactor* reactor) override;
+ private:
+ friend class Stub;
+ explicit async(Stub* stub): stub_(stub) { }
+ Stub* stub() { return stub_; }
+ Stub* stub_;
+ };
+ class async* async() override { return &async_stub_; }
+
+ private:
+ std::shared_ptr< ::grpc::ChannelInterface> channel_;
+ class async async_stub_{this};
+ ::grpc::ClientAsyncResponseReader< ::flwr::proto::CreateNodeResponse>* AsyncCreateNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::grpc::CompletionQueue* cq) override;
+ ::grpc::ClientAsyncResponseReader< ::flwr::proto::CreateNodeResponse>* PrepareAsyncCreateNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::CreateNodeRequest& request, ::grpc::CompletionQueue* cq) override;
+ ::grpc::ClientAsyncResponseReader< ::flwr::proto::DeleteNodeResponse>* AsyncDeleteNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::grpc::CompletionQueue* cq) override;
+ ::grpc::ClientAsyncResponseReader< ::flwr::proto::DeleteNodeResponse>* PrepareAsyncDeleteNodeRaw(::grpc::ClientContext* context, const ::flwr::proto::DeleteNodeRequest& request, ::grpc::CompletionQueue* cq) override;
+ ::grpc::ClientAsyncResponseReader< ::flwr::proto::PullTaskInsResponse>* AsyncPullTaskInsRaw(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::grpc::CompletionQueue* cq) override;
+ ::grpc::ClientAsyncResponseReader< ::flwr::proto::PullTaskInsResponse>* PrepareAsyncPullTaskInsRaw(::grpc::ClientContext* context, const ::flwr::proto::PullTaskInsRequest& request, ::grpc::CompletionQueue* cq) override;
+ ::grpc::ClientAsyncResponseReader< ::flwr::proto::PushTaskResResponse>* AsyncPushTaskResRaw(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::grpc::CompletionQueue* cq) override;
+ ::grpc::ClientAsyncResponseReader< ::flwr::proto::PushTaskResResponse>* PrepareAsyncPushTaskResRaw(::grpc::ClientContext* context, const ::flwr::proto::PushTaskResRequest& request, ::grpc::CompletionQueue* cq) override;
+ const ::grpc::internal::RpcMethod rpcmethod_CreateNode_;
+ const ::grpc::internal::RpcMethod rpcmethod_DeleteNode_;
+ const ::grpc::internal::RpcMethod rpcmethod_PullTaskIns_;
+ const ::grpc::internal::RpcMethod rpcmethod_PushTaskRes_;
+ };
+ static std::unique_ptr NewStub(const std::shared_ptr< ::grpc::ChannelInterface>& channel, const ::grpc::StubOptions& options = ::grpc::StubOptions());
+
+ class Service : public ::grpc::Service {
+ public:
+ Service();
+ virtual ~Service();
+ virtual ::grpc::Status CreateNode(::grpc::ServerContext* context, const ::flwr::proto::CreateNodeRequest* request, ::flwr::proto::CreateNodeResponse* response);
+ virtual ::grpc::Status DeleteNode(::grpc::ServerContext* context, const ::flwr::proto::DeleteNodeRequest* request, ::flwr::proto::DeleteNodeResponse* response);
+ // Retrieve one or more tasks, if possible
+ //
+ // HTTP API path: /api/v1/fleet/pull-task-ins
+ virtual ::grpc::Status PullTaskIns(::grpc::ServerContext* context, const ::flwr::proto::PullTaskInsRequest* request, ::flwr::proto::PullTaskInsResponse* response);
+ // Complete one or more tasks, if possible
+ //
+ // HTTP API path: /api/v1/fleet/push-task-res
+ virtual ::grpc::Status PushTaskRes(::grpc::ServerContext* context, const ::flwr::proto::PushTaskResRequest* request, ::flwr::proto::PushTaskResResponse* response);
+ };
+ template
+ class WithAsyncMethod_CreateNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithAsyncMethod_CreateNode() {
+ ::grpc::Service::MarkMethodAsync(0);
+ }
+ ~WithAsyncMethod_CreateNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status CreateNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::CreateNodeRequest* /*request*/, ::flwr::proto::CreateNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ void RequestCreateNode(::grpc::ServerContext* context, ::flwr::proto::CreateNodeRequest* request, ::grpc::ServerAsyncResponseWriter< ::flwr::proto::CreateNodeResponse>* response, ::grpc::CompletionQueue* new_call_cq, ::grpc::ServerCompletionQueue* notification_cq, void *tag) {
+ ::grpc::Service::RequestAsyncUnary(0, context, request, response, new_call_cq, notification_cq, tag);
+ }
+ };
+ template
+ class WithAsyncMethod_DeleteNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithAsyncMethod_DeleteNode() {
+ ::grpc::Service::MarkMethodAsync(1);
+ }
+ ~WithAsyncMethod_DeleteNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status DeleteNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::DeleteNodeRequest* /*request*/, ::flwr::proto::DeleteNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ void RequestDeleteNode(::grpc::ServerContext* context, ::flwr::proto::DeleteNodeRequest* request, ::grpc::ServerAsyncResponseWriter< ::flwr::proto::DeleteNodeResponse>* response, ::grpc::CompletionQueue* new_call_cq, ::grpc::ServerCompletionQueue* notification_cq, void *tag) {
+ ::grpc::Service::RequestAsyncUnary(1, context, request, response, new_call_cq, notification_cq, tag);
+ }
+ };
+ template
+ class WithAsyncMethod_PullTaskIns : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithAsyncMethod_PullTaskIns() {
+ ::grpc::Service::MarkMethodAsync(2);
+ }
+ ~WithAsyncMethod_PullTaskIns() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status PullTaskIns(::grpc::ServerContext* /*context*/, const ::flwr::proto::PullTaskInsRequest* /*request*/, ::flwr::proto::PullTaskInsResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ void RequestPullTaskIns(::grpc::ServerContext* context, ::flwr::proto::PullTaskInsRequest* request, ::grpc::ServerAsyncResponseWriter< ::flwr::proto::PullTaskInsResponse>* response, ::grpc::CompletionQueue* new_call_cq, ::grpc::ServerCompletionQueue* notification_cq, void *tag) {
+ ::grpc::Service::RequestAsyncUnary(2, context, request, response, new_call_cq, notification_cq, tag);
+ }
+ };
+ template
+ class WithAsyncMethod_PushTaskRes : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithAsyncMethod_PushTaskRes() {
+ ::grpc::Service::MarkMethodAsync(3);
+ }
+ ~WithAsyncMethod_PushTaskRes() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status PushTaskRes(::grpc::ServerContext* /*context*/, const ::flwr::proto::PushTaskResRequest* /*request*/, ::flwr::proto::PushTaskResResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ void RequestPushTaskRes(::grpc::ServerContext* context, ::flwr::proto::PushTaskResRequest* request, ::grpc::ServerAsyncResponseWriter< ::flwr::proto::PushTaskResResponse>* response, ::grpc::CompletionQueue* new_call_cq, ::grpc::ServerCompletionQueue* notification_cq, void *tag) {
+ ::grpc::Service::RequestAsyncUnary(3, context, request, response, new_call_cq, notification_cq, tag);
+ }
+ };
+ typedef WithAsyncMethod_CreateNode > > > AsyncService;
+ template
+ class WithCallbackMethod_CreateNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithCallbackMethod_CreateNode() {
+ ::grpc::Service::MarkMethodCallback(0,
+ new ::grpc::internal::CallbackUnaryHandler< ::flwr::proto::CreateNodeRequest, ::flwr::proto::CreateNodeResponse>(
+ [this](
+ ::grpc::CallbackServerContext* context, const ::flwr::proto::CreateNodeRequest* request, ::flwr::proto::CreateNodeResponse* response) { return this->CreateNode(context, request, response); }));}
+ void SetMessageAllocatorFor_CreateNode(
+ ::grpc::MessageAllocator< ::flwr::proto::CreateNodeRequest, ::flwr::proto::CreateNodeResponse>* allocator) {
+ ::grpc::internal::MethodHandler* const handler = ::grpc::Service::GetHandler(0);
+ static_cast<::grpc::internal::CallbackUnaryHandler< ::flwr::proto::CreateNodeRequest, ::flwr::proto::CreateNodeResponse>*>(handler)
+ ->SetMessageAllocator(allocator);
+ }
+ ~WithCallbackMethod_CreateNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status CreateNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::CreateNodeRequest* /*request*/, ::flwr::proto::CreateNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ virtual ::grpc::ServerUnaryReactor* CreateNode(
+ ::grpc::CallbackServerContext* /*context*/, const ::flwr::proto::CreateNodeRequest* /*request*/, ::flwr::proto::CreateNodeResponse* /*response*/) { return nullptr; }
+ };
+ template
+ class WithCallbackMethod_DeleteNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithCallbackMethod_DeleteNode() {
+ ::grpc::Service::MarkMethodCallback(1,
+ new ::grpc::internal::CallbackUnaryHandler< ::flwr::proto::DeleteNodeRequest, ::flwr::proto::DeleteNodeResponse>(
+ [this](
+ ::grpc::CallbackServerContext* context, const ::flwr::proto::DeleteNodeRequest* request, ::flwr::proto::DeleteNodeResponse* response) { return this->DeleteNode(context, request, response); }));}
+ void SetMessageAllocatorFor_DeleteNode(
+ ::grpc::MessageAllocator< ::flwr::proto::DeleteNodeRequest, ::flwr::proto::DeleteNodeResponse>* allocator) {
+ ::grpc::internal::MethodHandler* const handler = ::grpc::Service::GetHandler(1);
+ static_cast<::grpc::internal::CallbackUnaryHandler< ::flwr::proto::DeleteNodeRequest, ::flwr::proto::DeleteNodeResponse>*>(handler)
+ ->SetMessageAllocator(allocator);
+ }
+ ~WithCallbackMethod_DeleteNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status DeleteNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::DeleteNodeRequest* /*request*/, ::flwr::proto::DeleteNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ virtual ::grpc::ServerUnaryReactor* DeleteNode(
+ ::grpc::CallbackServerContext* /*context*/, const ::flwr::proto::DeleteNodeRequest* /*request*/, ::flwr::proto::DeleteNodeResponse* /*response*/) { return nullptr; }
+ };
+ template
+ class WithCallbackMethod_PullTaskIns : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithCallbackMethod_PullTaskIns() {
+ ::grpc::Service::MarkMethodCallback(2,
+ new ::grpc::internal::CallbackUnaryHandler< ::flwr::proto::PullTaskInsRequest, ::flwr::proto::PullTaskInsResponse>(
+ [this](
+ ::grpc::CallbackServerContext* context, const ::flwr::proto::PullTaskInsRequest* request, ::flwr::proto::PullTaskInsResponse* response) { return this->PullTaskIns(context, request, response); }));}
+ void SetMessageAllocatorFor_PullTaskIns(
+ ::grpc::MessageAllocator< ::flwr::proto::PullTaskInsRequest, ::flwr::proto::PullTaskInsResponse>* allocator) {
+ ::grpc::internal::MethodHandler* const handler = ::grpc::Service::GetHandler(2);
+ static_cast<::grpc::internal::CallbackUnaryHandler< ::flwr::proto::PullTaskInsRequest, ::flwr::proto::PullTaskInsResponse>*>(handler)
+ ->SetMessageAllocator(allocator);
+ }
+ ~WithCallbackMethod_PullTaskIns() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status PullTaskIns(::grpc::ServerContext* /*context*/, const ::flwr::proto::PullTaskInsRequest* /*request*/, ::flwr::proto::PullTaskInsResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ virtual ::grpc::ServerUnaryReactor* PullTaskIns(
+ ::grpc::CallbackServerContext* /*context*/, const ::flwr::proto::PullTaskInsRequest* /*request*/, ::flwr::proto::PullTaskInsResponse* /*response*/) { return nullptr; }
+ };
+ template
+ class WithCallbackMethod_PushTaskRes : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithCallbackMethod_PushTaskRes() {
+ ::grpc::Service::MarkMethodCallback(3,
+ new ::grpc::internal::CallbackUnaryHandler< ::flwr::proto::PushTaskResRequest, ::flwr::proto::PushTaskResResponse>(
+ [this](
+ ::grpc::CallbackServerContext* context, const ::flwr::proto::PushTaskResRequest* request, ::flwr::proto::PushTaskResResponse* response) { return this->PushTaskRes(context, request, response); }));}
+ void SetMessageAllocatorFor_PushTaskRes(
+ ::grpc::MessageAllocator< ::flwr::proto::PushTaskResRequest, ::flwr::proto::PushTaskResResponse>* allocator) {
+ ::grpc::internal::MethodHandler* const handler = ::grpc::Service::GetHandler(3);
+ static_cast<::grpc::internal::CallbackUnaryHandler< ::flwr::proto::PushTaskResRequest, ::flwr::proto::PushTaskResResponse>*>(handler)
+ ->SetMessageAllocator(allocator);
+ }
+ ~WithCallbackMethod_PushTaskRes() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status PushTaskRes(::grpc::ServerContext* /*context*/, const ::flwr::proto::PushTaskResRequest* /*request*/, ::flwr::proto::PushTaskResResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ virtual ::grpc::ServerUnaryReactor* PushTaskRes(
+ ::grpc::CallbackServerContext* /*context*/, const ::flwr::proto::PushTaskResRequest* /*request*/, ::flwr::proto::PushTaskResResponse* /*response*/) { return nullptr; }
+ };
+ typedef WithCallbackMethod_CreateNode > > > CallbackService;
+ typedef CallbackService ExperimentalCallbackService;
+ template
+ class WithGenericMethod_CreateNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithGenericMethod_CreateNode() {
+ ::grpc::Service::MarkMethodGeneric(0);
+ }
+ ~WithGenericMethod_CreateNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status CreateNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::CreateNodeRequest* /*request*/, ::flwr::proto::CreateNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ };
+ template
+ class WithGenericMethod_DeleteNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithGenericMethod_DeleteNode() {
+ ::grpc::Service::MarkMethodGeneric(1);
+ }
+ ~WithGenericMethod_DeleteNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status DeleteNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::DeleteNodeRequest* /*request*/, ::flwr::proto::DeleteNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ };
+ template
+ class WithGenericMethod_PullTaskIns : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithGenericMethod_PullTaskIns() {
+ ::grpc::Service::MarkMethodGeneric(2);
+ }
+ ~WithGenericMethod_PullTaskIns() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status PullTaskIns(::grpc::ServerContext* /*context*/, const ::flwr::proto::PullTaskInsRequest* /*request*/, ::flwr::proto::PullTaskInsResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ };
+ template
+ class WithGenericMethod_PushTaskRes : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithGenericMethod_PushTaskRes() {
+ ::grpc::Service::MarkMethodGeneric(3);
+ }
+ ~WithGenericMethod_PushTaskRes() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status PushTaskRes(::grpc::ServerContext* /*context*/, const ::flwr::proto::PushTaskResRequest* /*request*/, ::flwr::proto::PushTaskResResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ };
+ template
+ class WithRawMethod_CreateNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithRawMethod_CreateNode() {
+ ::grpc::Service::MarkMethodRaw(0);
+ }
+ ~WithRawMethod_CreateNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status CreateNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::CreateNodeRequest* /*request*/, ::flwr::proto::CreateNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ void RequestCreateNode(::grpc::ServerContext* context, ::grpc::ByteBuffer* request, ::grpc::ServerAsyncResponseWriter< ::grpc::ByteBuffer>* response, ::grpc::CompletionQueue* new_call_cq, ::grpc::ServerCompletionQueue* notification_cq, void *tag) {
+ ::grpc::Service::RequestAsyncUnary(0, context, request, response, new_call_cq, notification_cq, tag);
+ }
+ };
+ template
+ class WithRawMethod_DeleteNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithRawMethod_DeleteNode() {
+ ::grpc::Service::MarkMethodRaw(1);
+ }
+ ~WithRawMethod_DeleteNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status DeleteNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::DeleteNodeRequest* /*request*/, ::flwr::proto::DeleteNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ void RequestDeleteNode(::grpc::ServerContext* context, ::grpc::ByteBuffer* request, ::grpc::ServerAsyncResponseWriter< ::grpc::ByteBuffer>* response, ::grpc::CompletionQueue* new_call_cq, ::grpc::ServerCompletionQueue* notification_cq, void *tag) {
+ ::grpc::Service::RequestAsyncUnary(1, context, request, response, new_call_cq, notification_cq, tag);
+ }
+ };
+ template
+ class WithRawMethod_PullTaskIns : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithRawMethod_PullTaskIns() {
+ ::grpc::Service::MarkMethodRaw(2);
+ }
+ ~WithRawMethod_PullTaskIns() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status PullTaskIns(::grpc::ServerContext* /*context*/, const ::flwr::proto::PullTaskInsRequest* /*request*/, ::flwr::proto::PullTaskInsResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ void RequestPullTaskIns(::grpc::ServerContext* context, ::grpc::ByteBuffer* request, ::grpc::ServerAsyncResponseWriter< ::grpc::ByteBuffer>* response, ::grpc::CompletionQueue* new_call_cq, ::grpc::ServerCompletionQueue* notification_cq, void *tag) {
+ ::grpc::Service::RequestAsyncUnary(2, context, request, response, new_call_cq, notification_cq, tag);
+ }
+ };
+ template
+ class WithRawMethod_PushTaskRes : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithRawMethod_PushTaskRes() {
+ ::grpc::Service::MarkMethodRaw(3);
+ }
+ ~WithRawMethod_PushTaskRes() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status PushTaskRes(::grpc::ServerContext* /*context*/, const ::flwr::proto::PushTaskResRequest* /*request*/, ::flwr::proto::PushTaskResResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ void RequestPushTaskRes(::grpc::ServerContext* context, ::grpc::ByteBuffer* request, ::grpc::ServerAsyncResponseWriter< ::grpc::ByteBuffer>* response, ::grpc::CompletionQueue* new_call_cq, ::grpc::ServerCompletionQueue* notification_cq, void *tag) {
+ ::grpc::Service::RequestAsyncUnary(3, context, request, response, new_call_cq, notification_cq, tag);
+ }
+ };
+ template
+ class WithRawCallbackMethod_CreateNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithRawCallbackMethod_CreateNode() {
+ ::grpc::Service::MarkMethodRawCallback(0,
+ new ::grpc::internal::CallbackUnaryHandler< ::grpc::ByteBuffer, ::grpc::ByteBuffer>(
+ [this](
+ ::grpc::CallbackServerContext* context, const ::grpc::ByteBuffer* request, ::grpc::ByteBuffer* response) { return this->CreateNode(context, request, response); }));
+ }
+ ~WithRawCallbackMethod_CreateNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status CreateNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::CreateNodeRequest* /*request*/, ::flwr::proto::CreateNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ virtual ::grpc::ServerUnaryReactor* CreateNode(
+ ::grpc::CallbackServerContext* /*context*/, const ::grpc::ByteBuffer* /*request*/, ::grpc::ByteBuffer* /*response*/) { return nullptr; }
+ };
+ template
+ class WithRawCallbackMethod_DeleteNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithRawCallbackMethod_DeleteNode() {
+ ::grpc::Service::MarkMethodRawCallback(1,
+ new ::grpc::internal::CallbackUnaryHandler< ::grpc::ByteBuffer, ::grpc::ByteBuffer>(
+ [this](
+ ::grpc::CallbackServerContext* context, const ::grpc::ByteBuffer* request, ::grpc::ByteBuffer* response) { return this->DeleteNode(context, request, response); }));
+ }
+ ~WithRawCallbackMethod_DeleteNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status DeleteNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::DeleteNodeRequest* /*request*/, ::flwr::proto::DeleteNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ virtual ::grpc::ServerUnaryReactor* DeleteNode(
+ ::grpc::CallbackServerContext* /*context*/, const ::grpc::ByteBuffer* /*request*/, ::grpc::ByteBuffer* /*response*/) { return nullptr; }
+ };
+ template
+ class WithRawCallbackMethod_PullTaskIns : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithRawCallbackMethod_PullTaskIns() {
+ ::grpc::Service::MarkMethodRawCallback(2,
+ new ::grpc::internal::CallbackUnaryHandler< ::grpc::ByteBuffer, ::grpc::ByteBuffer>(
+ [this](
+ ::grpc::CallbackServerContext* context, const ::grpc::ByteBuffer* request, ::grpc::ByteBuffer* response) { return this->PullTaskIns(context, request, response); }));
+ }
+ ~WithRawCallbackMethod_PullTaskIns() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status PullTaskIns(::grpc::ServerContext* /*context*/, const ::flwr::proto::PullTaskInsRequest* /*request*/, ::flwr::proto::PullTaskInsResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ virtual ::grpc::ServerUnaryReactor* PullTaskIns(
+ ::grpc::CallbackServerContext* /*context*/, const ::grpc::ByteBuffer* /*request*/, ::grpc::ByteBuffer* /*response*/) { return nullptr; }
+ };
+ template
+ class WithRawCallbackMethod_PushTaskRes : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithRawCallbackMethod_PushTaskRes() {
+ ::grpc::Service::MarkMethodRawCallback(3,
+ new ::grpc::internal::CallbackUnaryHandler< ::grpc::ByteBuffer, ::grpc::ByteBuffer>(
+ [this](
+ ::grpc::CallbackServerContext* context, const ::grpc::ByteBuffer* request, ::grpc::ByteBuffer* response) { return this->PushTaskRes(context, request, response); }));
+ }
+ ~WithRawCallbackMethod_PushTaskRes() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable synchronous version of this method
+ ::grpc::Status PushTaskRes(::grpc::ServerContext* /*context*/, const ::flwr::proto::PushTaskResRequest* /*request*/, ::flwr::proto::PushTaskResResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ virtual ::grpc::ServerUnaryReactor* PushTaskRes(
+ ::grpc::CallbackServerContext* /*context*/, const ::grpc::ByteBuffer* /*request*/, ::grpc::ByteBuffer* /*response*/) { return nullptr; }
+ };
+ template
+ class WithStreamedUnaryMethod_CreateNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithStreamedUnaryMethod_CreateNode() {
+ ::grpc::Service::MarkMethodStreamed(0,
+ new ::grpc::internal::StreamedUnaryHandler<
+ ::flwr::proto::CreateNodeRequest, ::flwr::proto::CreateNodeResponse>(
+ [this](::grpc::ServerContext* context,
+ ::grpc::ServerUnaryStreamer<
+ ::flwr::proto::CreateNodeRequest, ::flwr::proto::CreateNodeResponse>* streamer) {
+ return this->StreamedCreateNode(context,
+ streamer);
+ }));
+ }
+ ~WithStreamedUnaryMethod_CreateNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable regular version of this method
+ ::grpc::Status CreateNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::CreateNodeRequest* /*request*/, ::flwr::proto::CreateNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ // replace default version of method with streamed unary
+ virtual ::grpc::Status StreamedCreateNode(::grpc::ServerContext* context, ::grpc::ServerUnaryStreamer< ::flwr::proto::CreateNodeRequest,::flwr::proto::CreateNodeResponse>* server_unary_streamer) = 0;
+ };
+ template
+ class WithStreamedUnaryMethod_DeleteNode : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithStreamedUnaryMethod_DeleteNode() {
+ ::grpc::Service::MarkMethodStreamed(1,
+ new ::grpc::internal::StreamedUnaryHandler<
+ ::flwr::proto::DeleteNodeRequest, ::flwr::proto::DeleteNodeResponse>(
+ [this](::grpc::ServerContext* context,
+ ::grpc::ServerUnaryStreamer<
+ ::flwr::proto::DeleteNodeRequest, ::flwr::proto::DeleteNodeResponse>* streamer) {
+ return this->StreamedDeleteNode(context,
+ streamer);
+ }));
+ }
+ ~WithStreamedUnaryMethod_DeleteNode() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable regular version of this method
+ ::grpc::Status DeleteNode(::grpc::ServerContext* /*context*/, const ::flwr::proto::DeleteNodeRequest* /*request*/, ::flwr::proto::DeleteNodeResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ // replace default version of method with streamed unary
+ virtual ::grpc::Status StreamedDeleteNode(::grpc::ServerContext* context, ::grpc::ServerUnaryStreamer< ::flwr::proto::DeleteNodeRequest,::flwr::proto::DeleteNodeResponse>* server_unary_streamer) = 0;
+ };
+ template
+ class WithStreamedUnaryMethod_PullTaskIns : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithStreamedUnaryMethod_PullTaskIns() {
+ ::grpc::Service::MarkMethodStreamed(2,
+ new ::grpc::internal::StreamedUnaryHandler<
+ ::flwr::proto::PullTaskInsRequest, ::flwr::proto::PullTaskInsResponse>(
+ [this](::grpc::ServerContext* context,
+ ::grpc::ServerUnaryStreamer<
+ ::flwr::proto::PullTaskInsRequest, ::flwr::proto::PullTaskInsResponse>* streamer) {
+ return this->StreamedPullTaskIns(context,
+ streamer);
+ }));
+ }
+ ~WithStreamedUnaryMethod_PullTaskIns() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable regular version of this method
+ ::grpc::Status PullTaskIns(::grpc::ServerContext* /*context*/, const ::flwr::proto::PullTaskInsRequest* /*request*/, ::flwr::proto::PullTaskInsResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ // replace default version of method with streamed unary
+ virtual ::grpc::Status StreamedPullTaskIns(::grpc::ServerContext* context, ::grpc::ServerUnaryStreamer< ::flwr::proto::PullTaskInsRequest,::flwr::proto::PullTaskInsResponse>* server_unary_streamer) = 0;
+ };
+ template
+ class WithStreamedUnaryMethod_PushTaskRes : public BaseClass {
+ private:
+ void BaseClassMustBeDerivedFromService(const Service* /*service*/) {}
+ public:
+ WithStreamedUnaryMethod_PushTaskRes() {
+ ::grpc::Service::MarkMethodStreamed(3,
+ new ::grpc::internal::StreamedUnaryHandler<
+ ::flwr::proto::PushTaskResRequest, ::flwr::proto::PushTaskResResponse>(
+ [this](::grpc::ServerContext* context,
+ ::grpc::ServerUnaryStreamer<
+ ::flwr::proto::PushTaskResRequest, ::flwr::proto::PushTaskResResponse>* streamer) {
+ return this->StreamedPushTaskRes(context,
+ streamer);
+ }));
+ }
+ ~WithStreamedUnaryMethod_PushTaskRes() override {
+ BaseClassMustBeDerivedFromService(this);
+ }
+ // disable regular version of this method
+ ::grpc::Status PushTaskRes(::grpc::ServerContext* /*context*/, const ::flwr::proto::PushTaskResRequest* /*request*/, ::flwr::proto::PushTaskResResponse* /*response*/) override {
+ abort();
+ return ::grpc::Status(::grpc::StatusCode::UNIMPLEMENTED, "");
+ }
+ // replace default version of method with streamed unary
+ virtual ::grpc::Status StreamedPushTaskRes(::grpc::ServerContext* context, ::grpc::ServerUnaryStreamer< ::flwr::proto::PushTaskResRequest,::flwr::proto::PushTaskResResponse>* server_unary_streamer) = 0;
+ };
+ typedef WithStreamedUnaryMethod_CreateNode > > > StreamedUnaryService;
+ typedef Service SplitStreamedService;
+ typedef WithStreamedUnaryMethod_CreateNode > > > StreamedService;
+};
+
+} // namespace proto
+} // namespace flwr
+
+
+#endif // GRPC_flwr_2fproto_2ffleet_2eproto__INCLUDED
diff --git a/src/cc/flwr/include/flwr/proto/fleet.pb.cc b/src/cc/flwr/include/flwr/proto/fleet.pb.cc
new file mode 100644
index 000000000000..302331374db1
--- /dev/null
+++ b/src/cc/flwr/include/flwr/proto/fleet.pb.cc
@@ -0,0 +1,1932 @@
+// Generated by the protocol buffer compiler. DO NOT EDIT!
+// source: flwr/proto/fleet.proto
+
+#include "flwr/proto/fleet.pb.h"
+
+#include
+
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+// @@protoc_insertion_point(includes)
+#include
+
+PROTOBUF_PRAGMA_INIT_SEG
+namespace flwr {
+namespace proto {
+constexpr CreateNodeRequest::CreateNodeRequest(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized){}
+struct CreateNodeRequestDefaultTypeInternal {
+ constexpr CreateNodeRequestDefaultTypeInternal()
+ : _instance(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{}) {}
+ ~CreateNodeRequestDefaultTypeInternal() {}
+ union {
+ CreateNodeRequest _instance;
+ };
+};
+PROTOBUF_ATTRIBUTE_NO_DESTROY PROTOBUF_CONSTINIT CreateNodeRequestDefaultTypeInternal _CreateNodeRequest_default_instance_;
+constexpr CreateNodeResponse::CreateNodeResponse(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized)
+ : node_(nullptr){}
+struct CreateNodeResponseDefaultTypeInternal {
+ constexpr CreateNodeResponseDefaultTypeInternal()
+ : _instance(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{}) {}
+ ~CreateNodeResponseDefaultTypeInternal() {}
+ union {
+ CreateNodeResponse _instance;
+ };
+};
+PROTOBUF_ATTRIBUTE_NO_DESTROY PROTOBUF_CONSTINIT CreateNodeResponseDefaultTypeInternal _CreateNodeResponse_default_instance_;
+constexpr DeleteNodeRequest::DeleteNodeRequest(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized)
+ : node_(nullptr){}
+struct DeleteNodeRequestDefaultTypeInternal {
+ constexpr DeleteNodeRequestDefaultTypeInternal()
+ : _instance(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{}) {}
+ ~DeleteNodeRequestDefaultTypeInternal() {}
+ union {
+ DeleteNodeRequest _instance;
+ };
+};
+PROTOBUF_ATTRIBUTE_NO_DESTROY PROTOBUF_CONSTINIT DeleteNodeRequestDefaultTypeInternal _DeleteNodeRequest_default_instance_;
+constexpr DeleteNodeResponse::DeleteNodeResponse(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized){}
+struct DeleteNodeResponseDefaultTypeInternal {
+ constexpr DeleteNodeResponseDefaultTypeInternal()
+ : _instance(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{}) {}
+ ~DeleteNodeResponseDefaultTypeInternal() {}
+ union {
+ DeleteNodeResponse _instance;
+ };
+};
+PROTOBUF_ATTRIBUTE_NO_DESTROY PROTOBUF_CONSTINIT DeleteNodeResponseDefaultTypeInternal _DeleteNodeResponse_default_instance_;
+constexpr PullTaskInsRequest::PullTaskInsRequest(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized)
+ : task_ids_()
+ , node_(nullptr){}
+struct PullTaskInsRequestDefaultTypeInternal {
+ constexpr PullTaskInsRequestDefaultTypeInternal()
+ : _instance(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{}) {}
+ ~PullTaskInsRequestDefaultTypeInternal() {}
+ union {
+ PullTaskInsRequest _instance;
+ };
+};
+PROTOBUF_ATTRIBUTE_NO_DESTROY PROTOBUF_CONSTINIT PullTaskInsRequestDefaultTypeInternal _PullTaskInsRequest_default_instance_;
+constexpr PullTaskInsResponse::PullTaskInsResponse(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized)
+ : task_ins_list_()
+ , reconnect_(nullptr){}
+struct PullTaskInsResponseDefaultTypeInternal {
+ constexpr PullTaskInsResponseDefaultTypeInternal()
+ : _instance(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{}) {}
+ ~PullTaskInsResponseDefaultTypeInternal() {}
+ union {
+ PullTaskInsResponse _instance;
+ };
+};
+PROTOBUF_ATTRIBUTE_NO_DESTROY PROTOBUF_CONSTINIT PullTaskInsResponseDefaultTypeInternal _PullTaskInsResponse_default_instance_;
+constexpr PushTaskResRequest::PushTaskResRequest(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized)
+ : task_res_list_(){}
+struct PushTaskResRequestDefaultTypeInternal {
+ constexpr PushTaskResRequestDefaultTypeInternal()
+ : _instance(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{}) {}
+ ~PushTaskResRequestDefaultTypeInternal() {}
+ union {
+ PushTaskResRequest _instance;
+ };
+};
+PROTOBUF_ATTRIBUTE_NO_DESTROY PROTOBUF_CONSTINIT PushTaskResRequestDefaultTypeInternal _PushTaskResRequest_default_instance_;
+constexpr PushTaskResResponse_ResultsEntry_DoNotUse::PushTaskResResponse_ResultsEntry_DoNotUse(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized){}
+struct PushTaskResResponse_ResultsEntry_DoNotUseDefaultTypeInternal {
+ constexpr PushTaskResResponse_ResultsEntry_DoNotUseDefaultTypeInternal()
+ : _instance(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{}) {}
+ ~PushTaskResResponse_ResultsEntry_DoNotUseDefaultTypeInternal() {}
+ union {
+ PushTaskResResponse_ResultsEntry_DoNotUse _instance;
+ };
+};
+PROTOBUF_ATTRIBUTE_NO_DESTROY PROTOBUF_CONSTINIT PushTaskResResponse_ResultsEntry_DoNotUseDefaultTypeInternal _PushTaskResResponse_ResultsEntry_DoNotUse_default_instance_;
+constexpr PushTaskResResponse::PushTaskResResponse(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized)
+ : results_(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{})
+ , reconnect_(nullptr){}
+struct PushTaskResResponseDefaultTypeInternal {
+ constexpr PushTaskResResponseDefaultTypeInternal()
+ : _instance(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{}) {}
+ ~PushTaskResResponseDefaultTypeInternal() {}
+ union {
+ PushTaskResResponse _instance;
+ };
+};
+PROTOBUF_ATTRIBUTE_NO_DESTROY PROTOBUF_CONSTINIT PushTaskResResponseDefaultTypeInternal _PushTaskResResponse_default_instance_;
+constexpr Reconnect::Reconnect(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized)
+ : reconnect_(uint64_t{0u}){}
+struct ReconnectDefaultTypeInternal {
+ constexpr ReconnectDefaultTypeInternal()
+ : _instance(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized{}) {}
+ ~ReconnectDefaultTypeInternal() {}
+ union {
+ Reconnect _instance;
+ };
+};
+PROTOBUF_ATTRIBUTE_NO_DESTROY PROTOBUF_CONSTINIT ReconnectDefaultTypeInternal _Reconnect_default_instance_;
+} // namespace proto
+} // namespace flwr
+static ::PROTOBUF_NAMESPACE_ID::Metadata file_level_metadata_flwr_2fproto_2ffleet_2eproto[10];
+static constexpr ::PROTOBUF_NAMESPACE_ID::EnumDescriptor const** file_level_enum_descriptors_flwr_2fproto_2ffleet_2eproto = nullptr;
+static constexpr ::PROTOBUF_NAMESPACE_ID::ServiceDescriptor const** file_level_service_descriptors_flwr_2fproto_2ffleet_2eproto = nullptr;
+
+const ::PROTOBUF_NAMESPACE_ID::uint32 TableStruct_flwr_2fproto_2ffleet_2eproto::offsets[] PROTOBUF_SECTION_VARIABLE(protodesc_cold) = {
+ ~0u, // no _has_bits_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::CreateNodeRequest, _internal_metadata_),
+ ~0u, // no _extensions_
+ ~0u, // no _oneof_case_
+ ~0u, // no _weak_field_map_
+ ~0u, // no _inlined_string_donated_
+ ~0u, // no _has_bits_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::CreateNodeResponse, _internal_metadata_),
+ ~0u, // no _extensions_
+ ~0u, // no _oneof_case_
+ ~0u, // no _weak_field_map_
+ ~0u, // no _inlined_string_donated_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::CreateNodeResponse, node_),
+ ~0u, // no _has_bits_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::DeleteNodeRequest, _internal_metadata_),
+ ~0u, // no _extensions_
+ ~0u, // no _oneof_case_
+ ~0u, // no _weak_field_map_
+ ~0u, // no _inlined_string_donated_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::DeleteNodeRequest, node_),
+ ~0u, // no _has_bits_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::DeleteNodeResponse, _internal_metadata_),
+ ~0u, // no _extensions_
+ ~0u, // no _oneof_case_
+ ~0u, // no _weak_field_map_
+ ~0u, // no _inlined_string_donated_
+ ~0u, // no _has_bits_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PullTaskInsRequest, _internal_metadata_),
+ ~0u, // no _extensions_
+ ~0u, // no _oneof_case_
+ ~0u, // no _weak_field_map_
+ ~0u, // no _inlined_string_donated_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PullTaskInsRequest, node_),
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PullTaskInsRequest, task_ids_),
+ ~0u, // no _has_bits_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PullTaskInsResponse, _internal_metadata_),
+ ~0u, // no _extensions_
+ ~0u, // no _oneof_case_
+ ~0u, // no _weak_field_map_
+ ~0u, // no _inlined_string_donated_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PullTaskInsResponse, reconnect_),
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PullTaskInsResponse, task_ins_list_),
+ ~0u, // no _has_bits_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PushTaskResRequest, _internal_metadata_),
+ ~0u, // no _extensions_
+ ~0u, // no _oneof_case_
+ ~0u, // no _weak_field_map_
+ ~0u, // no _inlined_string_donated_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PushTaskResRequest, task_res_list_),
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PushTaskResResponse_ResultsEntry_DoNotUse, _has_bits_),
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PushTaskResResponse_ResultsEntry_DoNotUse, _internal_metadata_),
+ ~0u, // no _extensions_
+ ~0u, // no _oneof_case_
+ ~0u, // no _weak_field_map_
+ ~0u, // no _inlined_string_donated_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PushTaskResResponse_ResultsEntry_DoNotUse, key_),
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PushTaskResResponse_ResultsEntry_DoNotUse, value_),
+ 0,
+ 1,
+ ~0u, // no _has_bits_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PushTaskResResponse, _internal_metadata_),
+ ~0u, // no _extensions_
+ ~0u, // no _oneof_case_
+ ~0u, // no _weak_field_map_
+ ~0u, // no _inlined_string_donated_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PushTaskResResponse, reconnect_),
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::PushTaskResResponse, results_),
+ ~0u, // no _has_bits_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::Reconnect, _internal_metadata_),
+ ~0u, // no _extensions_
+ ~0u, // no _oneof_case_
+ ~0u, // no _weak_field_map_
+ ~0u, // no _inlined_string_donated_
+ PROTOBUF_FIELD_OFFSET(::flwr::proto::Reconnect, reconnect_),
+};
+static const ::PROTOBUF_NAMESPACE_ID::internal::MigrationSchema schemas[] PROTOBUF_SECTION_VARIABLE(protodesc_cold) = {
+ { 0, -1, -1, sizeof(::flwr::proto::CreateNodeRequest)},
+ { 6, -1, -1, sizeof(::flwr::proto::CreateNodeResponse)},
+ { 13, -1, -1, sizeof(::flwr::proto::DeleteNodeRequest)},
+ { 20, -1, -1, sizeof(::flwr::proto::DeleteNodeResponse)},
+ { 26, -1, -1, sizeof(::flwr::proto::PullTaskInsRequest)},
+ { 34, -1, -1, sizeof(::flwr::proto::PullTaskInsResponse)},
+ { 42, -1, -1, sizeof(::flwr::proto::PushTaskResRequest)},
+ { 49, 57, -1, sizeof(::flwr::proto::PushTaskResResponse_ResultsEntry_DoNotUse)},
+ { 59, -1, -1, sizeof(::flwr::proto::PushTaskResResponse)},
+ { 67, -1, -1, sizeof(::flwr::proto::Reconnect)},
+};
+
+static ::PROTOBUF_NAMESPACE_ID::Message const * const file_default_instances[] = {
+ reinterpret_cast(&::flwr::proto::_CreateNodeRequest_default_instance_),
+ reinterpret_cast(&::flwr::proto::_CreateNodeResponse_default_instance_),
+ reinterpret_cast(&::flwr::proto::_DeleteNodeRequest_default_instance_),
+ reinterpret_cast(&::flwr::proto::_DeleteNodeResponse_default_instance_),
+ reinterpret_cast(&::flwr::proto::_PullTaskInsRequest_default_instance_),
+ reinterpret_cast(&::flwr::proto::_PullTaskInsResponse_default_instance_),
+ reinterpret_cast(&::flwr::proto::_PushTaskResRequest_default_instance_),
+ reinterpret_cast(&::flwr::proto::_PushTaskResResponse_ResultsEntry_DoNotUse_default_instance_),
+ reinterpret_cast(&::flwr::proto::_PushTaskResResponse_default_instance_),
+ reinterpret_cast(&::flwr::proto::_Reconnect_default_instance_),
+};
+
+const char descriptor_table_protodef_flwr_2fproto_2ffleet_2eproto[] PROTOBUF_SECTION_VARIABLE(protodesc_cold) =
+ "\n\026flwr/proto/fleet.proto\022\nflwr.proto\032\025fl"
+ "wr/proto/node.proto\032\025flwr/proto/task.pro"
+ "to\"\023\n\021CreateNodeRequest\"4\n\022CreateNodeRes"
+ "ponse\022\036\n\004node\030\001 \001(\0132\020.flwr.proto.Node\"3\n"
+ "\021DeleteNodeRequest\022\036\n\004node\030\001 \001(\0132\020.flwr."
+ "proto.Node\"\024\n\022DeleteNodeResponse\"F\n\022Pull"
+ "TaskInsRequest\022\036\n\004node\030\001 \001(\0132\020.flwr.prot"
+ "o.Node\022\020\n\010task_ids\030\002 \003(\t\"k\n\023PullTaskInsR"
+ "esponse\022(\n\treconnect\030\001 \001(\0132\025.flwr.proto."
+ "Reconnect\022*\n\rtask_ins_list\030\002 \003(\0132\023.flwr."
+ "proto.TaskIns\"@\n\022PushTaskResRequest\022*\n\rt"
+ "ask_res_list\030\001 \003(\0132\023.flwr.proto.TaskRes\""
+ "\256\001\n\023PushTaskResResponse\022(\n\treconnect\030\001 \001"
+ "(\0132\025.flwr.proto.Reconnect\022=\n\007results\030\002 \003"
+ "(\0132,.flwr.proto.PushTaskResResponse.Resu"
+ "ltsEntry\032.\n\014ResultsEntry\022\013\n\003key\030\001 \001(\t\022\r\n"
+ "\005value\030\002 \001(\r:\0028\001\"\036\n\tReconnect\022\021\n\treconne"
+ "ct\030\001 \001(\0042\311\002\n\005Fleet\022M\n\nCreateNode\022\035.flwr."
+ "proto.CreateNodeRequest\032\036.flwr.proto.Cre"
+ "ateNodeResponse\"\000\022M\n\nDeleteNode\022\035.flwr.p"
+ "roto.DeleteNodeRequest\032\036.flwr.proto.Dele"
+ "teNodeResponse\"\000\022P\n\013PullTaskIns\022\036.flwr.p"
+ "roto.PullTaskInsRequest\032\037.flwr.proto.Pul"
+ "lTaskInsResponse\"\000\022P\n\013PushTaskRes\022\036.flwr"
+ ".proto.PushTaskResRequest\032\037.flwr.proto.P"
+ "ushTaskResResponse\"\000b\006proto3"
+ ;
+static const ::PROTOBUF_NAMESPACE_ID::internal::DescriptorTable*const descriptor_table_flwr_2fproto_2ffleet_2eproto_deps[2] = {
+ &::descriptor_table_flwr_2fproto_2fnode_2eproto,
+ &::descriptor_table_flwr_2fproto_2ftask_2eproto,
+};
+static ::PROTOBUF_NAMESPACE_ID::internal::once_flag descriptor_table_flwr_2fproto_2ffleet_2eproto_once;
+const ::PROTOBUF_NAMESPACE_ID::internal::DescriptorTable descriptor_table_flwr_2fproto_2ffleet_2eproto = {
+ false, false, 1028, descriptor_table_protodef_flwr_2fproto_2ffleet_2eproto, "flwr/proto/fleet.proto",
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_once, descriptor_table_flwr_2fproto_2ffleet_2eproto_deps, 2, 10,
+ schemas, file_default_instances, TableStruct_flwr_2fproto_2ffleet_2eproto::offsets,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto, file_level_enum_descriptors_flwr_2fproto_2ffleet_2eproto, file_level_service_descriptors_flwr_2fproto_2ffleet_2eproto,
+};
+PROTOBUF_ATTRIBUTE_WEAK const ::PROTOBUF_NAMESPACE_ID::internal::DescriptorTable* descriptor_table_flwr_2fproto_2ffleet_2eproto_getter() {
+ return &descriptor_table_flwr_2fproto_2ffleet_2eproto;
+}
+
+// Force running AddDescriptors() at dynamic initialization time.
+PROTOBUF_ATTRIBUTE_INIT_PRIORITY static ::PROTOBUF_NAMESPACE_ID::internal::AddDescriptorsRunner dynamic_init_dummy_flwr_2fproto_2ffleet_2eproto(&descriptor_table_flwr_2fproto_2ffleet_2eproto);
+namespace flwr {
+namespace proto {
+
+// ===================================================================
+
+class CreateNodeRequest::_Internal {
+ public:
+};
+
+CreateNodeRequest::CreateNodeRequest(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned)
+ : ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase(arena, is_message_owned) {
+ // @@protoc_insertion_point(arena_constructor:flwr.proto.CreateNodeRequest)
+}
+CreateNodeRequest::CreateNodeRequest(const CreateNodeRequest& from)
+ : ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase() {
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+ // @@protoc_insertion_point(copy_constructor:flwr.proto.CreateNodeRequest)
+}
+
+
+
+
+
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData CreateNodeRequest::_class_data_ = {
+ ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::CopyImpl,
+ ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::MergeImpl,
+};
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*CreateNodeRequest::GetClassData() const { return &_class_data_; }
+
+
+
+
+
+
+
+::PROTOBUF_NAMESPACE_ID::Metadata CreateNodeRequest::GetMetadata() const {
+ return ::PROTOBUF_NAMESPACE_ID::internal::AssignDescriptors(
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_getter, &descriptor_table_flwr_2fproto_2ffleet_2eproto_once,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto[0]);
+}
+
+// ===================================================================
+
+class CreateNodeResponse::_Internal {
+ public:
+ static const ::flwr::proto::Node& node(const CreateNodeResponse* msg);
+};
+
+const ::flwr::proto::Node&
+CreateNodeResponse::_Internal::node(const CreateNodeResponse* msg) {
+ return *msg->node_;
+}
+void CreateNodeResponse::clear_node() {
+ if (GetArenaForAllocation() == nullptr && node_ != nullptr) {
+ delete node_;
+ }
+ node_ = nullptr;
+}
+CreateNodeResponse::CreateNodeResponse(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned)
+ : ::PROTOBUF_NAMESPACE_ID::Message(arena, is_message_owned) {
+ SharedCtor();
+ if (!is_message_owned) {
+ RegisterArenaDtor(arena);
+ }
+ // @@protoc_insertion_point(arena_constructor:flwr.proto.CreateNodeResponse)
+}
+CreateNodeResponse::CreateNodeResponse(const CreateNodeResponse& from)
+ : ::PROTOBUF_NAMESPACE_ID::Message() {
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+ if (from._internal_has_node()) {
+ node_ = new ::flwr::proto::Node(*from.node_);
+ } else {
+ node_ = nullptr;
+ }
+ // @@protoc_insertion_point(copy_constructor:flwr.proto.CreateNodeResponse)
+}
+
+void CreateNodeResponse::SharedCtor() {
+node_ = nullptr;
+}
+
+CreateNodeResponse::~CreateNodeResponse() {
+ // @@protoc_insertion_point(destructor:flwr.proto.CreateNodeResponse)
+ if (GetArenaForAllocation() != nullptr) return;
+ SharedDtor();
+ _internal_metadata_.Delete<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+inline void CreateNodeResponse::SharedDtor() {
+ GOOGLE_DCHECK(GetArenaForAllocation() == nullptr);
+ if (this != internal_default_instance()) delete node_;
+}
+
+void CreateNodeResponse::ArenaDtor(void* object) {
+ CreateNodeResponse* _this = reinterpret_cast< CreateNodeResponse* >(object);
+ (void)_this;
+}
+void CreateNodeResponse::RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena*) {
+}
+void CreateNodeResponse::SetCachedSize(int size) const {
+ _cached_size_.Set(size);
+}
+
+void CreateNodeResponse::Clear() {
+// @@protoc_insertion_point(message_clear_start:flwr.proto.CreateNodeResponse)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ if (GetArenaForAllocation() == nullptr && node_ != nullptr) {
+ delete node_;
+ }
+ node_ = nullptr;
+ _internal_metadata_.Clear<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+const char* CreateNodeResponse::_InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) {
+#define CHK_(x) if (PROTOBUF_PREDICT_FALSE(!(x))) goto failure
+ while (!ctx->Done(&ptr)) {
+ ::PROTOBUF_NAMESPACE_ID::uint32 tag;
+ ptr = ::PROTOBUF_NAMESPACE_ID::internal::ReadTag(ptr, &tag);
+ switch (tag >> 3) {
+ // .flwr.proto.Node node = 1;
+ case 1:
+ if (PROTOBUF_PREDICT_TRUE(static_cast<::PROTOBUF_NAMESPACE_ID::uint8>(tag) == 10)) {
+ ptr = ctx->ParseMessage(_internal_mutable_node(), ptr);
+ CHK_(ptr);
+ } else
+ goto handle_unusual;
+ continue;
+ default:
+ goto handle_unusual;
+ } // switch
+ handle_unusual:
+ if ((tag == 0) || ((tag & 7) == 4)) {
+ CHK_(ptr);
+ ctx->SetLastTag(tag);
+ goto message_done;
+ }
+ ptr = UnknownFieldParse(
+ tag,
+ _internal_metadata_.mutable_unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(),
+ ptr, ctx);
+ CHK_(ptr != nullptr);
+ } // while
+message_done:
+ return ptr;
+failure:
+ ptr = nullptr;
+ goto message_done;
+#undef CHK_
+}
+
+::PROTOBUF_NAMESPACE_ID::uint8* CreateNodeResponse::_InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const {
+ // @@protoc_insertion_point(serialize_to_array_start:flwr.proto.CreateNodeResponse)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ // .flwr.proto.Node node = 1;
+ if (this->_internal_has_node()) {
+ target = stream->EnsureSpace(target);
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::
+ InternalWriteMessage(
+ 1, _Internal::node(this), target, stream);
+ }
+
+ if (PROTOBUF_PREDICT_FALSE(_internal_metadata_.have_unknown_fields())) {
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormat::InternalSerializeUnknownFieldsToArray(
+ _internal_metadata_.unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(::PROTOBUF_NAMESPACE_ID::UnknownFieldSet::default_instance), target, stream);
+ }
+ // @@protoc_insertion_point(serialize_to_array_end:flwr.proto.CreateNodeResponse)
+ return target;
+}
+
+size_t CreateNodeResponse::ByteSizeLong() const {
+// @@protoc_insertion_point(message_byte_size_start:flwr.proto.CreateNodeResponse)
+ size_t total_size = 0;
+
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ // .flwr.proto.Node node = 1;
+ if (this->_internal_has_node()) {
+ total_size += 1 +
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::MessageSize(
+ *node_);
+ }
+
+ return MaybeComputeUnknownFieldsSize(total_size, &_cached_size_);
+}
+
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData CreateNodeResponse::_class_data_ = {
+ ::PROTOBUF_NAMESPACE_ID::Message::CopyWithSizeCheck,
+ CreateNodeResponse::MergeImpl
+};
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*CreateNodeResponse::GetClassData() const { return &_class_data_; }
+
+void CreateNodeResponse::MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to,
+ const ::PROTOBUF_NAMESPACE_ID::Message& from) {
+ static_cast(to)->MergeFrom(
+ static_cast(from));
+}
+
+
+void CreateNodeResponse::MergeFrom(const CreateNodeResponse& from) {
+// @@protoc_insertion_point(class_specific_merge_from_start:flwr.proto.CreateNodeResponse)
+ GOOGLE_DCHECK_NE(&from, this);
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ if (from._internal_has_node()) {
+ _internal_mutable_node()->::flwr::proto::Node::MergeFrom(from._internal_node());
+ }
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+}
+
+void CreateNodeResponse::CopyFrom(const CreateNodeResponse& from) {
+// @@protoc_insertion_point(class_specific_copy_from_start:flwr.proto.CreateNodeResponse)
+ if (&from == this) return;
+ Clear();
+ MergeFrom(from);
+}
+
+bool CreateNodeResponse::IsInitialized() const {
+ return true;
+}
+
+void CreateNodeResponse::InternalSwap(CreateNodeResponse* other) {
+ using std::swap;
+ _internal_metadata_.InternalSwap(&other->_internal_metadata_);
+ swap(node_, other->node_);
+}
+
+::PROTOBUF_NAMESPACE_ID::Metadata CreateNodeResponse::GetMetadata() const {
+ return ::PROTOBUF_NAMESPACE_ID::internal::AssignDescriptors(
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_getter, &descriptor_table_flwr_2fproto_2ffleet_2eproto_once,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto[1]);
+}
+
+// ===================================================================
+
+class DeleteNodeRequest::_Internal {
+ public:
+ static const ::flwr::proto::Node& node(const DeleteNodeRequest* msg);
+};
+
+const ::flwr::proto::Node&
+DeleteNodeRequest::_Internal::node(const DeleteNodeRequest* msg) {
+ return *msg->node_;
+}
+void DeleteNodeRequest::clear_node() {
+ if (GetArenaForAllocation() == nullptr && node_ != nullptr) {
+ delete node_;
+ }
+ node_ = nullptr;
+}
+DeleteNodeRequest::DeleteNodeRequest(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned)
+ : ::PROTOBUF_NAMESPACE_ID::Message(arena, is_message_owned) {
+ SharedCtor();
+ if (!is_message_owned) {
+ RegisterArenaDtor(arena);
+ }
+ // @@protoc_insertion_point(arena_constructor:flwr.proto.DeleteNodeRequest)
+}
+DeleteNodeRequest::DeleteNodeRequest(const DeleteNodeRequest& from)
+ : ::PROTOBUF_NAMESPACE_ID::Message() {
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+ if (from._internal_has_node()) {
+ node_ = new ::flwr::proto::Node(*from.node_);
+ } else {
+ node_ = nullptr;
+ }
+ // @@protoc_insertion_point(copy_constructor:flwr.proto.DeleteNodeRequest)
+}
+
+void DeleteNodeRequest::SharedCtor() {
+node_ = nullptr;
+}
+
+DeleteNodeRequest::~DeleteNodeRequest() {
+ // @@protoc_insertion_point(destructor:flwr.proto.DeleteNodeRequest)
+ if (GetArenaForAllocation() != nullptr) return;
+ SharedDtor();
+ _internal_metadata_.Delete<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+inline void DeleteNodeRequest::SharedDtor() {
+ GOOGLE_DCHECK(GetArenaForAllocation() == nullptr);
+ if (this != internal_default_instance()) delete node_;
+}
+
+void DeleteNodeRequest::ArenaDtor(void* object) {
+ DeleteNodeRequest* _this = reinterpret_cast< DeleteNodeRequest* >(object);
+ (void)_this;
+}
+void DeleteNodeRequest::RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena*) {
+}
+void DeleteNodeRequest::SetCachedSize(int size) const {
+ _cached_size_.Set(size);
+}
+
+void DeleteNodeRequest::Clear() {
+// @@protoc_insertion_point(message_clear_start:flwr.proto.DeleteNodeRequest)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ if (GetArenaForAllocation() == nullptr && node_ != nullptr) {
+ delete node_;
+ }
+ node_ = nullptr;
+ _internal_metadata_.Clear<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+const char* DeleteNodeRequest::_InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) {
+#define CHK_(x) if (PROTOBUF_PREDICT_FALSE(!(x))) goto failure
+ while (!ctx->Done(&ptr)) {
+ ::PROTOBUF_NAMESPACE_ID::uint32 tag;
+ ptr = ::PROTOBUF_NAMESPACE_ID::internal::ReadTag(ptr, &tag);
+ switch (tag >> 3) {
+ // .flwr.proto.Node node = 1;
+ case 1:
+ if (PROTOBUF_PREDICT_TRUE(static_cast<::PROTOBUF_NAMESPACE_ID::uint8>(tag) == 10)) {
+ ptr = ctx->ParseMessage(_internal_mutable_node(), ptr);
+ CHK_(ptr);
+ } else
+ goto handle_unusual;
+ continue;
+ default:
+ goto handle_unusual;
+ } // switch
+ handle_unusual:
+ if ((tag == 0) || ((tag & 7) == 4)) {
+ CHK_(ptr);
+ ctx->SetLastTag(tag);
+ goto message_done;
+ }
+ ptr = UnknownFieldParse(
+ tag,
+ _internal_metadata_.mutable_unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(),
+ ptr, ctx);
+ CHK_(ptr != nullptr);
+ } // while
+message_done:
+ return ptr;
+failure:
+ ptr = nullptr;
+ goto message_done;
+#undef CHK_
+}
+
+::PROTOBUF_NAMESPACE_ID::uint8* DeleteNodeRequest::_InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const {
+ // @@protoc_insertion_point(serialize_to_array_start:flwr.proto.DeleteNodeRequest)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ // .flwr.proto.Node node = 1;
+ if (this->_internal_has_node()) {
+ target = stream->EnsureSpace(target);
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::
+ InternalWriteMessage(
+ 1, _Internal::node(this), target, stream);
+ }
+
+ if (PROTOBUF_PREDICT_FALSE(_internal_metadata_.have_unknown_fields())) {
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormat::InternalSerializeUnknownFieldsToArray(
+ _internal_metadata_.unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(::PROTOBUF_NAMESPACE_ID::UnknownFieldSet::default_instance), target, stream);
+ }
+ // @@protoc_insertion_point(serialize_to_array_end:flwr.proto.DeleteNodeRequest)
+ return target;
+}
+
+size_t DeleteNodeRequest::ByteSizeLong() const {
+// @@protoc_insertion_point(message_byte_size_start:flwr.proto.DeleteNodeRequest)
+ size_t total_size = 0;
+
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ // .flwr.proto.Node node = 1;
+ if (this->_internal_has_node()) {
+ total_size += 1 +
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::MessageSize(
+ *node_);
+ }
+
+ return MaybeComputeUnknownFieldsSize(total_size, &_cached_size_);
+}
+
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData DeleteNodeRequest::_class_data_ = {
+ ::PROTOBUF_NAMESPACE_ID::Message::CopyWithSizeCheck,
+ DeleteNodeRequest::MergeImpl
+};
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*DeleteNodeRequest::GetClassData() const { return &_class_data_; }
+
+void DeleteNodeRequest::MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to,
+ const ::PROTOBUF_NAMESPACE_ID::Message& from) {
+ static_cast(to)->MergeFrom(
+ static_cast(from));
+}
+
+
+void DeleteNodeRequest::MergeFrom(const DeleteNodeRequest& from) {
+// @@protoc_insertion_point(class_specific_merge_from_start:flwr.proto.DeleteNodeRequest)
+ GOOGLE_DCHECK_NE(&from, this);
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ if (from._internal_has_node()) {
+ _internal_mutable_node()->::flwr::proto::Node::MergeFrom(from._internal_node());
+ }
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+}
+
+void DeleteNodeRequest::CopyFrom(const DeleteNodeRequest& from) {
+// @@protoc_insertion_point(class_specific_copy_from_start:flwr.proto.DeleteNodeRequest)
+ if (&from == this) return;
+ Clear();
+ MergeFrom(from);
+}
+
+bool DeleteNodeRequest::IsInitialized() const {
+ return true;
+}
+
+void DeleteNodeRequest::InternalSwap(DeleteNodeRequest* other) {
+ using std::swap;
+ _internal_metadata_.InternalSwap(&other->_internal_metadata_);
+ swap(node_, other->node_);
+}
+
+::PROTOBUF_NAMESPACE_ID::Metadata DeleteNodeRequest::GetMetadata() const {
+ return ::PROTOBUF_NAMESPACE_ID::internal::AssignDescriptors(
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_getter, &descriptor_table_flwr_2fproto_2ffleet_2eproto_once,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto[2]);
+}
+
+// ===================================================================
+
+class DeleteNodeResponse::_Internal {
+ public:
+};
+
+DeleteNodeResponse::DeleteNodeResponse(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned)
+ : ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase(arena, is_message_owned) {
+ // @@protoc_insertion_point(arena_constructor:flwr.proto.DeleteNodeResponse)
+}
+DeleteNodeResponse::DeleteNodeResponse(const DeleteNodeResponse& from)
+ : ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase() {
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+ // @@protoc_insertion_point(copy_constructor:flwr.proto.DeleteNodeResponse)
+}
+
+
+
+
+
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData DeleteNodeResponse::_class_data_ = {
+ ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::CopyImpl,
+ ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::MergeImpl,
+};
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*DeleteNodeResponse::GetClassData() const { return &_class_data_; }
+
+
+
+
+
+
+
+::PROTOBUF_NAMESPACE_ID::Metadata DeleteNodeResponse::GetMetadata() const {
+ return ::PROTOBUF_NAMESPACE_ID::internal::AssignDescriptors(
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_getter, &descriptor_table_flwr_2fproto_2ffleet_2eproto_once,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto[3]);
+}
+
+// ===================================================================
+
+class PullTaskInsRequest::_Internal {
+ public:
+ static const ::flwr::proto::Node& node(const PullTaskInsRequest* msg);
+};
+
+const ::flwr::proto::Node&
+PullTaskInsRequest::_Internal::node(const PullTaskInsRequest* msg) {
+ return *msg->node_;
+}
+void PullTaskInsRequest::clear_node() {
+ if (GetArenaForAllocation() == nullptr && node_ != nullptr) {
+ delete node_;
+ }
+ node_ = nullptr;
+}
+PullTaskInsRequest::PullTaskInsRequest(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned)
+ : ::PROTOBUF_NAMESPACE_ID::Message(arena, is_message_owned),
+ task_ids_(arena) {
+ SharedCtor();
+ if (!is_message_owned) {
+ RegisterArenaDtor(arena);
+ }
+ // @@protoc_insertion_point(arena_constructor:flwr.proto.PullTaskInsRequest)
+}
+PullTaskInsRequest::PullTaskInsRequest(const PullTaskInsRequest& from)
+ : ::PROTOBUF_NAMESPACE_ID::Message(),
+ task_ids_(from.task_ids_) {
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+ if (from._internal_has_node()) {
+ node_ = new ::flwr::proto::Node(*from.node_);
+ } else {
+ node_ = nullptr;
+ }
+ // @@protoc_insertion_point(copy_constructor:flwr.proto.PullTaskInsRequest)
+}
+
+void PullTaskInsRequest::SharedCtor() {
+node_ = nullptr;
+}
+
+PullTaskInsRequest::~PullTaskInsRequest() {
+ // @@protoc_insertion_point(destructor:flwr.proto.PullTaskInsRequest)
+ if (GetArenaForAllocation() != nullptr) return;
+ SharedDtor();
+ _internal_metadata_.Delete<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+inline void PullTaskInsRequest::SharedDtor() {
+ GOOGLE_DCHECK(GetArenaForAllocation() == nullptr);
+ if (this != internal_default_instance()) delete node_;
+}
+
+void PullTaskInsRequest::ArenaDtor(void* object) {
+ PullTaskInsRequest* _this = reinterpret_cast< PullTaskInsRequest* >(object);
+ (void)_this;
+}
+void PullTaskInsRequest::RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena*) {
+}
+void PullTaskInsRequest::SetCachedSize(int size) const {
+ _cached_size_.Set(size);
+}
+
+void PullTaskInsRequest::Clear() {
+// @@protoc_insertion_point(message_clear_start:flwr.proto.PullTaskInsRequest)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ task_ids_.Clear();
+ if (GetArenaForAllocation() == nullptr && node_ != nullptr) {
+ delete node_;
+ }
+ node_ = nullptr;
+ _internal_metadata_.Clear<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+const char* PullTaskInsRequest::_InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) {
+#define CHK_(x) if (PROTOBUF_PREDICT_FALSE(!(x))) goto failure
+ while (!ctx->Done(&ptr)) {
+ ::PROTOBUF_NAMESPACE_ID::uint32 tag;
+ ptr = ::PROTOBUF_NAMESPACE_ID::internal::ReadTag(ptr, &tag);
+ switch (tag >> 3) {
+ // .flwr.proto.Node node = 1;
+ case 1:
+ if (PROTOBUF_PREDICT_TRUE(static_cast<::PROTOBUF_NAMESPACE_ID::uint8>(tag) == 10)) {
+ ptr = ctx->ParseMessage(_internal_mutable_node(), ptr);
+ CHK_(ptr);
+ } else
+ goto handle_unusual;
+ continue;
+ // repeated string task_ids = 2;
+ case 2:
+ if (PROTOBUF_PREDICT_TRUE(static_cast<::PROTOBUF_NAMESPACE_ID::uint8>(tag) == 18)) {
+ ptr -= 1;
+ do {
+ ptr += 1;
+ auto str = _internal_add_task_ids();
+ ptr = ::PROTOBUF_NAMESPACE_ID::internal::InlineGreedyStringParser(str, ptr, ctx);
+ CHK_(::PROTOBUF_NAMESPACE_ID::internal::VerifyUTF8(str, "flwr.proto.PullTaskInsRequest.task_ids"));
+ CHK_(ptr);
+ if (!ctx->DataAvailable(ptr)) break;
+ } while (::PROTOBUF_NAMESPACE_ID::internal::ExpectTag<18>(ptr));
+ } else
+ goto handle_unusual;
+ continue;
+ default:
+ goto handle_unusual;
+ } // switch
+ handle_unusual:
+ if ((tag == 0) || ((tag & 7) == 4)) {
+ CHK_(ptr);
+ ctx->SetLastTag(tag);
+ goto message_done;
+ }
+ ptr = UnknownFieldParse(
+ tag,
+ _internal_metadata_.mutable_unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(),
+ ptr, ctx);
+ CHK_(ptr != nullptr);
+ } // while
+message_done:
+ return ptr;
+failure:
+ ptr = nullptr;
+ goto message_done;
+#undef CHK_
+}
+
+::PROTOBUF_NAMESPACE_ID::uint8* PullTaskInsRequest::_InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const {
+ // @@protoc_insertion_point(serialize_to_array_start:flwr.proto.PullTaskInsRequest)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ // .flwr.proto.Node node = 1;
+ if (this->_internal_has_node()) {
+ target = stream->EnsureSpace(target);
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::
+ InternalWriteMessage(
+ 1, _Internal::node(this), target, stream);
+ }
+
+ // repeated string task_ids = 2;
+ for (int i = 0, n = this->_internal_task_ids_size(); i < n; i++) {
+ const auto& s = this->_internal_task_ids(i);
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::VerifyUtf8String(
+ s.data(), static_cast(s.length()),
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::SERIALIZE,
+ "flwr.proto.PullTaskInsRequest.task_ids");
+ target = stream->WriteString(2, s, target);
+ }
+
+ if (PROTOBUF_PREDICT_FALSE(_internal_metadata_.have_unknown_fields())) {
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormat::InternalSerializeUnknownFieldsToArray(
+ _internal_metadata_.unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(::PROTOBUF_NAMESPACE_ID::UnknownFieldSet::default_instance), target, stream);
+ }
+ // @@protoc_insertion_point(serialize_to_array_end:flwr.proto.PullTaskInsRequest)
+ return target;
+}
+
+size_t PullTaskInsRequest::ByteSizeLong() const {
+// @@protoc_insertion_point(message_byte_size_start:flwr.proto.PullTaskInsRequest)
+ size_t total_size = 0;
+
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ // repeated string task_ids = 2;
+ total_size += 1 *
+ ::PROTOBUF_NAMESPACE_ID::internal::FromIntSize(task_ids_.size());
+ for (int i = 0, n = task_ids_.size(); i < n; i++) {
+ total_size += ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::StringSize(
+ task_ids_.Get(i));
+ }
+
+ // .flwr.proto.Node node = 1;
+ if (this->_internal_has_node()) {
+ total_size += 1 +
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::MessageSize(
+ *node_);
+ }
+
+ return MaybeComputeUnknownFieldsSize(total_size, &_cached_size_);
+}
+
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData PullTaskInsRequest::_class_data_ = {
+ ::PROTOBUF_NAMESPACE_ID::Message::CopyWithSizeCheck,
+ PullTaskInsRequest::MergeImpl
+};
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*PullTaskInsRequest::GetClassData() const { return &_class_data_; }
+
+void PullTaskInsRequest::MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to,
+ const ::PROTOBUF_NAMESPACE_ID::Message& from) {
+ static_cast(to)->MergeFrom(
+ static_cast(from));
+}
+
+
+void PullTaskInsRequest::MergeFrom(const PullTaskInsRequest& from) {
+// @@protoc_insertion_point(class_specific_merge_from_start:flwr.proto.PullTaskInsRequest)
+ GOOGLE_DCHECK_NE(&from, this);
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ task_ids_.MergeFrom(from.task_ids_);
+ if (from._internal_has_node()) {
+ _internal_mutable_node()->::flwr::proto::Node::MergeFrom(from._internal_node());
+ }
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+}
+
+void PullTaskInsRequest::CopyFrom(const PullTaskInsRequest& from) {
+// @@protoc_insertion_point(class_specific_copy_from_start:flwr.proto.PullTaskInsRequest)
+ if (&from == this) return;
+ Clear();
+ MergeFrom(from);
+}
+
+bool PullTaskInsRequest::IsInitialized() const {
+ return true;
+}
+
+void PullTaskInsRequest::InternalSwap(PullTaskInsRequest* other) {
+ using std::swap;
+ _internal_metadata_.InternalSwap(&other->_internal_metadata_);
+ task_ids_.InternalSwap(&other->task_ids_);
+ swap(node_, other->node_);
+}
+
+::PROTOBUF_NAMESPACE_ID::Metadata PullTaskInsRequest::GetMetadata() const {
+ return ::PROTOBUF_NAMESPACE_ID::internal::AssignDescriptors(
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_getter, &descriptor_table_flwr_2fproto_2ffleet_2eproto_once,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto[4]);
+}
+
+// ===================================================================
+
+class PullTaskInsResponse::_Internal {
+ public:
+ static const ::flwr::proto::Reconnect& reconnect(const PullTaskInsResponse* msg);
+};
+
+const ::flwr::proto::Reconnect&
+PullTaskInsResponse::_Internal::reconnect(const PullTaskInsResponse* msg) {
+ return *msg->reconnect_;
+}
+void PullTaskInsResponse::clear_task_ins_list() {
+ task_ins_list_.Clear();
+}
+PullTaskInsResponse::PullTaskInsResponse(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned)
+ : ::PROTOBUF_NAMESPACE_ID::Message(arena, is_message_owned),
+ task_ins_list_(arena) {
+ SharedCtor();
+ if (!is_message_owned) {
+ RegisterArenaDtor(arena);
+ }
+ // @@protoc_insertion_point(arena_constructor:flwr.proto.PullTaskInsResponse)
+}
+PullTaskInsResponse::PullTaskInsResponse(const PullTaskInsResponse& from)
+ : ::PROTOBUF_NAMESPACE_ID::Message(),
+ task_ins_list_(from.task_ins_list_) {
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+ if (from._internal_has_reconnect()) {
+ reconnect_ = new ::flwr::proto::Reconnect(*from.reconnect_);
+ } else {
+ reconnect_ = nullptr;
+ }
+ // @@protoc_insertion_point(copy_constructor:flwr.proto.PullTaskInsResponse)
+}
+
+void PullTaskInsResponse::SharedCtor() {
+reconnect_ = nullptr;
+}
+
+PullTaskInsResponse::~PullTaskInsResponse() {
+ // @@protoc_insertion_point(destructor:flwr.proto.PullTaskInsResponse)
+ if (GetArenaForAllocation() != nullptr) return;
+ SharedDtor();
+ _internal_metadata_.Delete<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+inline void PullTaskInsResponse::SharedDtor() {
+ GOOGLE_DCHECK(GetArenaForAllocation() == nullptr);
+ if (this != internal_default_instance()) delete reconnect_;
+}
+
+void PullTaskInsResponse::ArenaDtor(void* object) {
+ PullTaskInsResponse* _this = reinterpret_cast< PullTaskInsResponse* >(object);
+ (void)_this;
+}
+void PullTaskInsResponse::RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena*) {
+}
+void PullTaskInsResponse::SetCachedSize(int size) const {
+ _cached_size_.Set(size);
+}
+
+void PullTaskInsResponse::Clear() {
+// @@protoc_insertion_point(message_clear_start:flwr.proto.PullTaskInsResponse)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ task_ins_list_.Clear();
+ if (GetArenaForAllocation() == nullptr && reconnect_ != nullptr) {
+ delete reconnect_;
+ }
+ reconnect_ = nullptr;
+ _internal_metadata_.Clear<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+const char* PullTaskInsResponse::_InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) {
+#define CHK_(x) if (PROTOBUF_PREDICT_FALSE(!(x))) goto failure
+ while (!ctx->Done(&ptr)) {
+ ::PROTOBUF_NAMESPACE_ID::uint32 tag;
+ ptr = ::PROTOBUF_NAMESPACE_ID::internal::ReadTag(ptr, &tag);
+ switch (tag >> 3) {
+ // .flwr.proto.Reconnect reconnect = 1;
+ case 1:
+ if (PROTOBUF_PREDICT_TRUE(static_cast<::PROTOBUF_NAMESPACE_ID::uint8>(tag) == 10)) {
+ ptr = ctx->ParseMessage(_internal_mutable_reconnect(), ptr);
+ CHK_(ptr);
+ } else
+ goto handle_unusual;
+ continue;
+ // repeated .flwr.proto.TaskIns task_ins_list = 2;
+ case 2:
+ if (PROTOBUF_PREDICT_TRUE(static_cast<::PROTOBUF_NAMESPACE_ID::uint8>(tag) == 18)) {
+ ptr -= 1;
+ do {
+ ptr += 1;
+ ptr = ctx->ParseMessage(_internal_add_task_ins_list(), ptr);
+ CHK_(ptr);
+ if (!ctx->DataAvailable(ptr)) break;
+ } while (::PROTOBUF_NAMESPACE_ID::internal::ExpectTag<18>(ptr));
+ } else
+ goto handle_unusual;
+ continue;
+ default:
+ goto handle_unusual;
+ } // switch
+ handle_unusual:
+ if ((tag == 0) || ((tag & 7) == 4)) {
+ CHK_(ptr);
+ ctx->SetLastTag(tag);
+ goto message_done;
+ }
+ ptr = UnknownFieldParse(
+ tag,
+ _internal_metadata_.mutable_unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(),
+ ptr, ctx);
+ CHK_(ptr != nullptr);
+ } // while
+message_done:
+ return ptr;
+failure:
+ ptr = nullptr;
+ goto message_done;
+#undef CHK_
+}
+
+::PROTOBUF_NAMESPACE_ID::uint8* PullTaskInsResponse::_InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const {
+ // @@protoc_insertion_point(serialize_to_array_start:flwr.proto.PullTaskInsResponse)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ // .flwr.proto.Reconnect reconnect = 1;
+ if (this->_internal_has_reconnect()) {
+ target = stream->EnsureSpace(target);
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::
+ InternalWriteMessage(
+ 1, _Internal::reconnect(this), target, stream);
+ }
+
+ // repeated .flwr.proto.TaskIns task_ins_list = 2;
+ for (unsigned int i = 0,
+ n = static_cast(this->_internal_task_ins_list_size()); i < n; i++) {
+ target = stream->EnsureSpace(target);
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::
+ InternalWriteMessage(2, this->_internal_task_ins_list(i), target, stream);
+ }
+
+ if (PROTOBUF_PREDICT_FALSE(_internal_metadata_.have_unknown_fields())) {
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormat::InternalSerializeUnknownFieldsToArray(
+ _internal_metadata_.unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(::PROTOBUF_NAMESPACE_ID::UnknownFieldSet::default_instance), target, stream);
+ }
+ // @@protoc_insertion_point(serialize_to_array_end:flwr.proto.PullTaskInsResponse)
+ return target;
+}
+
+size_t PullTaskInsResponse::ByteSizeLong() const {
+// @@protoc_insertion_point(message_byte_size_start:flwr.proto.PullTaskInsResponse)
+ size_t total_size = 0;
+
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ // repeated .flwr.proto.TaskIns task_ins_list = 2;
+ total_size += 1UL * this->_internal_task_ins_list_size();
+ for (const auto& msg : this->task_ins_list_) {
+ total_size +=
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::MessageSize(msg);
+ }
+
+ // .flwr.proto.Reconnect reconnect = 1;
+ if (this->_internal_has_reconnect()) {
+ total_size += 1 +
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::MessageSize(
+ *reconnect_);
+ }
+
+ return MaybeComputeUnknownFieldsSize(total_size, &_cached_size_);
+}
+
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData PullTaskInsResponse::_class_data_ = {
+ ::PROTOBUF_NAMESPACE_ID::Message::CopyWithSizeCheck,
+ PullTaskInsResponse::MergeImpl
+};
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*PullTaskInsResponse::GetClassData() const { return &_class_data_; }
+
+void PullTaskInsResponse::MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to,
+ const ::PROTOBUF_NAMESPACE_ID::Message& from) {
+ static_cast(to)->MergeFrom(
+ static_cast(from));
+}
+
+
+void PullTaskInsResponse::MergeFrom(const PullTaskInsResponse& from) {
+// @@protoc_insertion_point(class_specific_merge_from_start:flwr.proto.PullTaskInsResponse)
+ GOOGLE_DCHECK_NE(&from, this);
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ task_ins_list_.MergeFrom(from.task_ins_list_);
+ if (from._internal_has_reconnect()) {
+ _internal_mutable_reconnect()->::flwr::proto::Reconnect::MergeFrom(from._internal_reconnect());
+ }
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+}
+
+void PullTaskInsResponse::CopyFrom(const PullTaskInsResponse& from) {
+// @@protoc_insertion_point(class_specific_copy_from_start:flwr.proto.PullTaskInsResponse)
+ if (&from == this) return;
+ Clear();
+ MergeFrom(from);
+}
+
+bool PullTaskInsResponse::IsInitialized() const {
+ return true;
+}
+
+void PullTaskInsResponse::InternalSwap(PullTaskInsResponse* other) {
+ using std::swap;
+ _internal_metadata_.InternalSwap(&other->_internal_metadata_);
+ task_ins_list_.InternalSwap(&other->task_ins_list_);
+ swap(reconnect_, other->reconnect_);
+}
+
+::PROTOBUF_NAMESPACE_ID::Metadata PullTaskInsResponse::GetMetadata() const {
+ return ::PROTOBUF_NAMESPACE_ID::internal::AssignDescriptors(
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_getter, &descriptor_table_flwr_2fproto_2ffleet_2eproto_once,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto[5]);
+}
+
+// ===================================================================
+
+class PushTaskResRequest::_Internal {
+ public:
+};
+
+void PushTaskResRequest::clear_task_res_list() {
+ task_res_list_.Clear();
+}
+PushTaskResRequest::PushTaskResRequest(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned)
+ : ::PROTOBUF_NAMESPACE_ID::Message(arena, is_message_owned),
+ task_res_list_(arena) {
+ SharedCtor();
+ if (!is_message_owned) {
+ RegisterArenaDtor(arena);
+ }
+ // @@protoc_insertion_point(arena_constructor:flwr.proto.PushTaskResRequest)
+}
+PushTaskResRequest::PushTaskResRequest(const PushTaskResRequest& from)
+ : ::PROTOBUF_NAMESPACE_ID::Message(),
+ task_res_list_(from.task_res_list_) {
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+ // @@protoc_insertion_point(copy_constructor:flwr.proto.PushTaskResRequest)
+}
+
+void PushTaskResRequest::SharedCtor() {
+}
+
+PushTaskResRequest::~PushTaskResRequest() {
+ // @@protoc_insertion_point(destructor:flwr.proto.PushTaskResRequest)
+ if (GetArenaForAllocation() != nullptr) return;
+ SharedDtor();
+ _internal_metadata_.Delete<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+inline void PushTaskResRequest::SharedDtor() {
+ GOOGLE_DCHECK(GetArenaForAllocation() == nullptr);
+}
+
+void PushTaskResRequest::ArenaDtor(void* object) {
+ PushTaskResRequest* _this = reinterpret_cast< PushTaskResRequest* >(object);
+ (void)_this;
+}
+void PushTaskResRequest::RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena*) {
+}
+void PushTaskResRequest::SetCachedSize(int size) const {
+ _cached_size_.Set(size);
+}
+
+void PushTaskResRequest::Clear() {
+// @@protoc_insertion_point(message_clear_start:flwr.proto.PushTaskResRequest)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ task_res_list_.Clear();
+ _internal_metadata_.Clear<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+const char* PushTaskResRequest::_InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) {
+#define CHK_(x) if (PROTOBUF_PREDICT_FALSE(!(x))) goto failure
+ while (!ctx->Done(&ptr)) {
+ ::PROTOBUF_NAMESPACE_ID::uint32 tag;
+ ptr = ::PROTOBUF_NAMESPACE_ID::internal::ReadTag(ptr, &tag);
+ switch (tag >> 3) {
+ // repeated .flwr.proto.TaskRes task_res_list = 1;
+ case 1:
+ if (PROTOBUF_PREDICT_TRUE(static_cast<::PROTOBUF_NAMESPACE_ID::uint8>(tag) == 10)) {
+ ptr -= 1;
+ do {
+ ptr += 1;
+ ptr = ctx->ParseMessage(_internal_add_task_res_list(), ptr);
+ CHK_(ptr);
+ if (!ctx->DataAvailable(ptr)) break;
+ } while (::PROTOBUF_NAMESPACE_ID::internal::ExpectTag<10>(ptr));
+ } else
+ goto handle_unusual;
+ continue;
+ default:
+ goto handle_unusual;
+ } // switch
+ handle_unusual:
+ if ((tag == 0) || ((tag & 7) == 4)) {
+ CHK_(ptr);
+ ctx->SetLastTag(tag);
+ goto message_done;
+ }
+ ptr = UnknownFieldParse(
+ tag,
+ _internal_metadata_.mutable_unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(),
+ ptr, ctx);
+ CHK_(ptr != nullptr);
+ } // while
+message_done:
+ return ptr;
+failure:
+ ptr = nullptr;
+ goto message_done;
+#undef CHK_
+}
+
+::PROTOBUF_NAMESPACE_ID::uint8* PushTaskResRequest::_InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const {
+ // @@protoc_insertion_point(serialize_to_array_start:flwr.proto.PushTaskResRequest)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ // repeated .flwr.proto.TaskRes task_res_list = 1;
+ for (unsigned int i = 0,
+ n = static_cast(this->_internal_task_res_list_size()); i < n; i++) {
+ target = stream->EnsureSpace(target);
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::
+ InternalWriteMessage(1, this->_internal_task_res_list(i), target, stream);
+ }
+
+ if (PROTOBUF_PREDICT_FALSE(_internal_metadata_.have_unknown_fields())) {
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormat::InternalSerializeUnknownFieldsToArray(
+ _internal_metadata_.unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(::PROTOBUF_NAMESPACE_ID::UnknownFieldSet::default_instance), target, stream);
+ }
+ // @@protoc_insertion_point(serialize_to_array_end:flwr.proto.PushTaskResRequest)
+ return target;
+}
+
+size_t PushTaskResRequest::ByteSizeLong() const {
+// @@protoc_insertion_point(message_byte_size_start:flwr.proto.PushTaskResRequest)
+ size_t total_size = 0;
+
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ // repeated .flwr.proto.TaskRes task_res_list = 1;
+ total_size += 1UL * this->_internal_task_res_list_size();
+ for (const auto& msg : this->task_res_list_) {
+ total_size +=
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::MessageSize(msg);
+ }
+
+ return MaybeComputeUnknownFieldsSize(total_size, &_cached_size_);
+}
+
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData PushTaskResRequest::_class_data_ = {
+ ::PROTOBUF_NAMESPACE_ID::Message::CopyWithSizeCheck,
+ PushTaskResRequest::MergeImpl
+};
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*PushTaskResRequest::GetClassData() const { return &_class_data_; }
+
+void PushTaskResRequest::MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to,
+ const ::PROTOBUF_NAMESPACE_ID::Message& from) {
+ static_cast(to)->MergeFrom(
+ static_cast(from));
+}
+
+
+void PushTaskResRequest::MergeFrom(const PushTaskResRequest& from) {
+// @@protoc_insertion_point(class_specific_merge_from_start:flwr.proto.PushTaskResRequest)
+ GOOGLE_DCHECK_NE(&from, this);
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ task_res_list_.MergeFrom(from.task_res_list_);
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+}
+
+void PushTaskResRequest::CopyFrom(const PushTaskResRequest& from) {
+// @@protoc_insertion_point(class_specific_copy_from_start:flwr.proto.PushTaskResRequest)
+ if (&from == this) return;
+ Clear();
+ MergeFrom(from);
+}
+
+bool PushTaskResRequest::IsInitialized() const {
+ return true;
+}
+
+void PushTaskResRequest::InternalSwap(PushTaskResRequest* other) {
+ using std::swap;
+ _internal_metadata_.InternalSwap(&other->_internal_metadata_);
+ task_res_list_.InternalSwap(&other->task_res_list_);
+}
+
+::PROTOBUF_NAMESPACE_ID::Metadata PushTaskResRequest::GetMetadata() const {
+ return ::PROTOBUF_NAMESPACE_ID::internal::AssignDescriptors(
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_getter, &descriptor_table_flwr_2fproto_2ffleet_2eproto_once,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto[6]);
+}
+
+// ===================================================================
+
+PushTaskResResponse_ResultsEntry_DoNotUse::PushTaskResResponse_ResultsEntry_DoNotUse() {}
+PushTaskResResponse_ResultsEntry_DoNotUse::PushTaskResResponse_ResultsEntry_DoNotUse(::PROTOBUF_NAMESPACE_ID::Arena* arena)
+ : SuperType(arena) {}
+void PushTaskResResponse_ResultsEntry_DoNotUse::MergeFrom(const PushTaskResResponse_ResultsEntry_DoNotUse& other) {
+ MergeFromInternal(other);
+}
+::PROTOBUF_NAMESPACE_ID::Metadata PushTaskResResponse_ResultsEntry_DoNotUse::GetMetadata() const {
+ return ::PROTOBUF_NAMESPACE_ID::internal::AssignDescriptors(
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_getter, &descriptor_table_flwr_2fproto_2ffleet_2eproto_once,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto[7]);
+}
+
+// ===================================================================
+
+class PushTaskResResponse::_Internal {
+ public:
+ static const ::flwr::proto::Reconnect& reconnect(const PushTaskResResponse* msg);
+};
+
+const ::flwr::proto::Reconnect&
+PushTaskResResponse::_Internal::reconnect(const PushTaskResResponse* msg) {
+ return *msg->reconnect_;
+}
+PushTaskResResponse::PushTaskResResponse(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned)
+ : ::PROTOBUF_NAMESPACE_ID::Message(arena, is_message_owned),
+ results_(arena) {
+ SharedCtor();
+ if (!is_message_owned) {
+ RegisterArenaDtor(arena);
+ }
+ // @@protoc_insertion_point(arena_constructor:flwr.proto.PushTaskResResponse)
+}
+PushTaskResResponse::PushTaskResResponse(const PushTaskResResponse& from)
+ : ::PROTOBUF_NAMESPACE_ID::Message() {
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+ results_.MergeFrom(from.results_);
+ if (from._internal_has_reconnect()) {
+ reconnect_ = new ::flwr::proto::Reconnect(*from.reconnect_);
+ } else {
+ reconnect_ = nullptr;
+ }
+ // @@protoc_insertion_point(copy_constructor:flwr.proto.PushTaskResResponse)
+}
+
+void PushTaskResResponse::SharedCtor() {
+reconnect_ = nullptr;
+}
+
+PushTaskResResponse::~PushTaskResResponse() {
+ // @@protoc_insertion_point(destructor:flwr.proto.PushTaskResResponse)
+ if (GetArenaForAllocation() != nullptr) return;
+ SharedDtor();
+ _internal_metadata_.Delete<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+inline void PushTaskResResponse::SharedDtor() {
+ GOOGLE_DCHECK(GetArenaForAllocation() == nullptr);
+ if (this != internal_default_instance()) delete reconnect_;
+}
+
+void PushTaskResResponse::ArenaDtor(void* object) {
+ PushTaskResResponse* _this = reinterpret_cast< PushTaskResResponse* >(object);
+ (void)_this;
+ _this->results_. ~MapField();
+}
+inline void PushTaskResResponse::RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena* arena) {
+ if (arena != nullptr) {
+ arena->OwnCustomDestructor(this, &PushTaskResResponse::ArenaDtor);
+ }
+}
+void PushTaskResResponse::SetCachedSize(int size) const {
+ _cached_size_.Set(size);
+}
+
+void PushTaskResResponse::Clear() {
+// @@protoc_insertion_point(message_clear_start:flwr.proto.PushTaskResResponse)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ results_.Clear();
+ if (GetArenaForAllocation() == nullptr && reconnect_ != nullptr) {
+ delete reconnect_;
+ }
+ reconnect_ = nullptr;
+ _internal_metadata_.Clear<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+const char* PushTaskResResponse::_InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) {
+#define CHK_(x) if (PROTOBUF_PREDICT_FALSE(!(x))) goto failure
+ while (!ctx->Done(&ptr)) {
+ ::PROTOBUF_NAMESPACE_ID::uint32 tag;
+ ptr = ::PROTOBUF_NAMESPACE_ID::internal::ReadTag(ptr, &tag);
+ switch (tag >> 3) {
+ // .flwr.proto.Reconnect reconnect = 1;
+ case 1:
+ if (PROTOBUF_PREDICT_TRUE(static_cast<::PROTOBUF_NAMESPACE_ID::uint8>(tag) == 10)) {
+ ptr = ctx->ParseMessage(_internal_mutable_reconnect(), ptr);
+ CHK_(ptr);
+ } else
+ goto handle_unusual;
+ continue;
+ // map results = 2;
+ case 2:
+ if (PROTOBUF_PREDICT_TRUE(static_cast<::PROTOBUF_NAMESPACE_ID::uint8>(tag) == 18)) {
+ ptr -= 1;
+ do {
+ ptr += 1;
+ ptr = ctx->ParseMessage(&results_, ptr);
+ CHK_(ptr);
+ if (!ctx->DataAvailable(ptr)) break;
+ } while (::PROTOBUF_NAMESPACE_ID::internal::ExpectTag<18>(ptr));
+ } else
+ goto handle_unusual;
+ continue;
+ default:
+ goto handle_unusual;
+ } // switch
+ handle_unusual:
+ if ((tag == 0) || ((tag & 7) == 4)) {
+ CHK_(ptr);
+ ctx->SetLastTag(tag);
+ goto message_done;
+ }
+ ptr = UnknownFieldParse(
+ tag,
+ _internal_metadata_.mutable_unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(),
+ ptr, ctx);
+ CHK_(ptr != nullptr);
+ } // while
+message_done:
+ return ptr;
+failure:
+ ptr = nullptr;
+ goto message_done;
+#undef CHK_
+}
+
+::PROTOBUF_NAMESPACE_ID::uint8* PushTaskResResponse::_InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const {
+ // @@protoc_insertion_point(serialize_to_array_start:flwr.proto.PushTaskResResponse)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ // .flwr.proto.Reconnect reconnect = 1;
+ if (this->_internal_has_reconnect()) {
+ target = stream->EnsureSpace(target);
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::
+ InternalWriteMessage(
+ 1, _Internal::reconnect(this), target, stream);
+ }
+
+ // map results = 2;
+ if (!this->_internal_results().empty()) {
+ typedef ::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >::const_pointer
+ ConstPtr;
+ typedef ConstPtr SortItem;
+ typedef ::PROTOBUF_NAMESPACE_ID::internal::CompareByDerefFirst Less;
+ struct Utf8Check {
+ static void Check(ConstPtr p) {
+ (void)p;
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::VerifyUtf8String(
+ p->first.data(), static_cast(p->first.length()),
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::SERIALIZE,
+ "flwr.proto.PushTaskResResponse.ResultsEntry.key");
+ }
+ };
+
+ if (stream->IsSerializationDeterministic() &&
+ this->_internal_results().size() > 1) {
+ ::std::unique_ptr items(
+ new SortItem[this->_internal_results().size()]);
+ typedef ::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >::size_type size_type;
+ size_type n = 0;
+ for (::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >::const_iterator
+ it = this->_internal_results().begin();
+ it != this->_internal_results().end(); ++it, ++n) {
+ items[static_cast(n)] = SortItem(&*it);
+ }
+ ::std::sort(&items[0], &items[static_cast(n)], Less());
+ for (size_type i = 0; i < n; i++) {
+ target = PushTaskResResponse_ResultsEntry_DoNotUse::Funcs::InternalSerialize(2, items[static_cast(i)]->first, items[static_cast(i)]->second, target, stream);
+ Utf8Check::Check(&(*items[static_cast(i)]));
+ }
+ } else {
+ for (::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >::const_iterator
+ it = this->_internal_results().begin();
+ it != this->_internal_results().end(); ++it) {
+ target = PushTaskResResponse_ResultsEntry_DoNotUse::Funcs::InternalSerialize(2, it->first, it->second, target, stream);
+ Utf8Check::Check(&(*it));
+ }
+ }
+ }
+
+ if (PROTOBUF_PREDICT_FALSE(_internal_metadata_.have_unknown_fields())) {
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormat::InternalSerializeUnknownFieldsToArray(
+ _internal_metadata_.unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(::PROTOBUF_NAMESPACE_ID::UnknownFieldSet::default_instance), target, stream);
+ }
+ // @@protoc_insertion_point(serialize_to_array_end:flwr.proto.PushTaskResResponse)
+ return target;
+}
+
+size_t PushTaskResResponse::ByteSizeLong() const {
+// @@protoc_insertion_point(message_byte_size_start:flwr.proto.PushTaskResResponse)
+ size_t total_size = 0;
+
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ // map results = 2;
+ total_size += 1 *
+ ::PROTOBUF_NAMESPACE_ID::internal::FromIntSize(this->_internal_results_size());
+ for (::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >::const_iterator
+ it = this->_internal_results().begin();
+ it != this->_internal_results().end(); ++it) {
+ total_size += PushTaskResResponse_ResultsEntry_DoNotUse::Funcs::ByteSizeLong(it->first, it->second);
+ }
+
+ // .flwr.proto.Reconnect reconnect = 1;
+ if (this->_internal_has_reconnect()) {
+ total_size += 1 +
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::MessageSize(
+ *reconnect_);
+ }
+
+ return MaybeComputeUnknownFieldsSize(total_size, &_cached_size_);
+}
+
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData PushTaskResResponse::_class_data_ = {
+ ::PROTOBUF_NAMESPACE_ID::Message::CopyWithSizeCheck,
+ PushTaskResResponse::MergeImpl
+};
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*PushTaskResResponse::GetClassData() const { return &_class_data_; }
+
+void PushTaskResResponse::MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to,
+ const ::PROTOBUF_NAMESPACE_ID::Message& from) {
+ static_cast(to)->MergeFrom(
+ static_cast(from));
+}
+
+
+void PushTaskResResponse::MergeFrom(const PushTaskResResponse& from) {
+// @@protoc_insertion_point(class_specific_merge_from_start:flwr.proto.PushTaskResResponse)
+ GOOGLE_DCHECK_NE(&from, this);
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ results_.MergeFrom(from.results_);
+ if (from._internal_has_reconnect()) {
+ _internal_mutable_reconnect()->::flwr::proto::Reconnect::MergeFrom(from._internal_reconnect());
+ }
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+}
+
+void PushTaskResResponse::CopyFrom(const PushTaskResResponse& from) {
+// @@protoc_insertion_point(class_specific_copy_from_start:flwr.proto.PushTaskResResponse)
+ if (&from == this) return;
+ Clear();
+ MergeFrom(from);
+}
+
+bool PushTaskResResponse::IsInitialized() const {
+ return true;
+}
+
+void PushTaskResResponse::InternalSwap(PushTaskResResponse* other) {
+ using std::swap;
+ _internal_metadata_.InternalSwap(&other->_internal_metadata_);
+ results_.InternalSwap(&other->results_);
+ swap(reconnect_, other->reconnect_);
+}
+
+::PROTOBUF_NAMESPACE_ID::Metadata PushTaskResResponse::GetMetadata() const {
+ return ::PROTOBUF_NAMESPACE_ID::internal::AssignDescriptors(
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_getter, &descriptor_table_flwr_2fproto_2ffleet_2eproto_once,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto[8]);
+}
+
+// ===================================================================
+
+class Reconnect::_Internal {
+ public:
+};
+
+Reconnect::Reconnect(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned)
+ : ::PROTOBUF_NAMESPACE_ID::Message(arena, is_message_owned) {
+ SharedCtor();
+ if (!is_message_owned) {
+ RegisterArenaDtor(arena);
+ }
+ // @@protoc_insertion_point(arena_constructor:flwr.proto.Reconnect)
+}
+Reconnect::Reconnect(const Reconnect& from)
+ : ::PROTOBUF_NAMESPACE_ID::Message() {
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+ reconnect_ = from.reconnect_;
+ // @@protoc_insertion_point(copy_constructor:flwr.proto.Reconnect)
+}
+
+void Reconnect::SharedCtor() {
+reconnect_ = uint64_t{0u};
+}
+
+Reconnect::~Reconnect() {
+ // @@protoc_insertion_point(destructor:flwr.proto.Reconnect)
+ if (GetArenaForAllocation() != nullptr) return;
+ SharedDtor();
+ _internal_metadata_.Delete<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+inline void Reconnect::SharedDtor() {
+ GOOGLE_DCHECK(GetArenaForAllocation() == nullptr);
+}
+
+void Reconnect::ArenaDtor(void* object) {
+ Reconnect* _this = reinterpret_cast< Reconnect* >(object);
+ (void)_this;
+}
+void Reconnect::RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena*) {
+}
+void Reconnect::SetCachedSize(int size) const {
+ _cached_size_.Set(size);
+}
+
+void Reconnect::Clear() {
+// @@protoc_insertion_point(message_clear_start:flwr.proto.Reconnect)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ reconnect_ = uint64_t{0u};
+ _internal_metadata_.Clear<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>();
+}
+
+const char* Reconnect::_InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) {
+#define CHK_(x) if (PROTOBUF_PREDICT_FALSE(!(x))) goto failure
+ while (!ctx->Done(&ptr)) {
+ ::PROTOBUF_NAMESPACE_ID::uint32 tag;
+ ptr = ::PROTOBUF_NAMESPACE_ID::internal::ReadTag(ptr, &tag);
+ switch (tag >> 3) {
+ // uint64 reconnect = 1;
+ case 1:
+ if (PROTOBUF_PREDICT_TRUE(static_cast<::PROTOBUF_NAMESPACE_ID::uint8>(tag) == 8)) {
+ reconnect_ = ::PROTOBUF_NAMESPACE_ID::internal::ReadVarint64(&ptr);
+ CHK_(ptr);
+ } else
+ goto handle_unusual;
+ continue;
+ default:
+ goto handle_unusual;
+ } // switch
+ handle_unusual:
+ if ((tag == 0) || ((tag & 7) == 4)) {
+ CHK_(ptr);
+ ctx->SetLastTag(tag);
+ goto message_done;
+ }
+ ptr = UnknownFieldParse(
+ tag,
+ _internal_metadata_.mutable_unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(),
+ ptr, ctx);
+ CHK_(ptr != nullptr);
+ } // while
+message_done:
+ return ptr;
+failure:
+ ptr = nullptr;
+ goto message_done;
+#undef CHK_
+}
+
+::PROTOBUF_NAMESPACE_ID::uint8* Reconnect::_InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const {
+ // @@protoc_insertion_point(serialize_to_array_start:flwr.proto.Reconnect)
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ // uint64 reconnect = 1;
+ if (this->_internal_reconnect() != 0) {
+ target = stream->EnsureSpace(target);
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::WriteUInt64ToArray(1, this->_internal_reconnect(), target);
+ }
+
+ if (PROTOBUF_PREDICT_FALSE(_internal_metadata_.have_unknown_fields())) {
+ target = ::PROTOBUF_NAMESPACE_ID::internal::WireFormat::InternalSerializeUnknownFieldsToArray(
+ _internal_metadata_.unknown_fields<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(::PROTOBUF_NAMESPACE_ID::UnknownFieldSet::default_instance), target, stream);
+ }
+ // @@protoc_insertion_point(serialize_to_array_end:flwr.proto.Reconnect)
+ return target;
+}
+
+size_t Reconnect::ByteSizeLong() const {
+// @@protoc_insertion_point(message_byte_size_start:flwr.proto.Reconnect)
+ size_t total_size = 0;
+
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ // Prevent compiler warnings about cached_has_bits being unused
+ (void) cached_has_bits;
+
+ // uint64 reconnect = 1;
+ if (this->_internal_reconnect() != 0) {
+ total_size += ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::UInt64SizePlusOne(this->_internal_reconnect());
+ }
+
+ return MaybeComputeUnknownFieldsSize(total_size, &_cached_size_);
+}
+
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData Reconnect::_class_data_ = {
+ ::PROTOBUF_NAMESPACE_ID::Message::CopyWithSizeCheck,
+ Reconnect::MergeImpl
+};
+const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*Reconnect::GetClassData() const { return &_class_data_; }
+
+void Reconnect::MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to,
+ const ::PROTOBUF_NAMESPACE_ID::Message& from) {
+ static_cast(to)->MergeFrom(
+ static_cast(from));
+}
+
+
+void Reconnect::MergeFrom(const Reconnect& from) {
+// @@protoc_insertion_point(class_specific_merge_from_start:flwr.proto.Reconnect)
+ GOOGLE_DCHECK_NE(&from, this);
+ ::PROTOBUF_NAMESPACE_ID::uint32 cached_has_bits = 0;
+ (void) cached_has_bits;
+
+ if (from._internal_reconnect() != 0) {
+ _internal_set_reconnect(from._internal_reconnect());
+ }
+ _internal_metadata_.MergeFrom<::PROTOBUF_NAMESPACE_ID::UnknownFieldSet>(from._internal_metadata_);
+}
+
+void Reconnect::CopyFrom(const Reconnect& from) {
+// @@protoc_insertion_point(class_specific_copy_from_start:flwr.proto.Reconnect)
+ if (&from == this) return;
+ Clear();
+ MergeFrom(from);
+}
+
+bool Reconnect::IsInitialized() const {
+ return true;
+}
+
+void Reconnect::InternalSwap(Reconnect* other) {
+ using std::swap;
+ _internal_metadata_.InternalSwap(&other->_internal_metadata_);
+ swap(reconnect_, other->reconnect_);
+}
+
+::PROTOBUF_NAMESPACE_ID::Metadata Reconnect::GetMetadata() const {
+ return ::PROTOBUF_NAMESPACE_ID::internal::AssignDescriptors(
+ &descriptor_table_flwr_2fproto_2ffleet_2eproto_getter, &descriptor_table_flwr_2fproto_2ffleet_2eproto_once,
+ file_level_metadata_flwr_2fproto_2ffleet_2eproto[9]);
+}
+
+// @@protoc_insertion_point(namespace_scope)
+} // namespace proto
+} // namespace flwr
+PROTOBUF_NAMESPACE_OPEN
+template<> PROTOBUF_NOINLINE ::flwr::proto::CreateNodeRequest* Arena::CreateMaybeMessage< ::flwr::proto::CreateNodeRequest >(Arena* arena) {
+ return Arena::CreateMessageInternal< ::flwr::proto::CreateNodeRequest >(arena);
+}
+template<> PROTOBUF_NOINLINE ::flwr::proto::CreateNodeResponse* Arena::CreateMaybeMessage< ::flwr::proto::CreateNodeResponse >(Arena* arena) {
+ return Arena::CreateMessageInternal< ::flwr::proto::CreateNodeResponse >(arena);
+}
+template<> PROTOBUF_NOINLINE ::flwr::proto::DeleteNodeRequest* Arena::CreateMaybeMessage< ::flwr::proto::DeleteNodeRequest >(Arena* arena) {
+ return Arena::CreateMessageInternal< ::flwr::proto::DeleteNodeRequest >(arena);
+}
+template<> PROTOBUF_NOINLINE ::flwr::proto::DeleteNodeResponse* Arena::CreateMaybeMessage< ::flwr::proto::DeleteNodeResponse >(Arena* arena) {
+ return Arena::CreateMessageInternal< ::flwr::proto::DeleteNodeResponse >(arena);
+}
+template<> PROTOBUF_NOINLINE ::flwr::proto::PullTaskInsRequest* Arena::CreateMaybeMessage< ::flwr::proto::PullTaskInsRequest >(Arena* arena) {
+ return Arena::CreateMessageInternal< ::flwr::proto::PullTaskInsRequest >(arena);
+}
+template<> PROTOBUF_NOINLINE ::flwr::proto::PullTaskInsResponse* Arena::CreateMaybeMessage< ::flwr::proto::PullTaskInsResponse >(Arena* arena) {
+ return Arena::CreateMessageInternal< ::flwr::proto::PullTaskInsResponse >(arena);
+}
+template<> PROTOBUF_NOINLINE ::flwr::proto::PushTaskResRequest* Arena::CreateMaybeMessage< ::flwr::proto::PushTaskResRequest >(Arena* arena) {
+ return Arena::CreateMessageInternal< ::flwr::proto::PushTaskResRequest >(arena);
+}
+template<> PROTOBUF_NOINLINE ::flwr::proto::PushTaskResResponse_ResultsEntry_DoNotUse* Arena::CreateMaybeMessage< ::flwr::proto::PushTaskResResponse_ResultsEntry_DoNotUse >(Arena* arena) {
+ return Arena::CreateMessageInternal< ::flwr::proto::PushTaskResResponse_ResultsEntry_DoNotUse >(arena);
+}
+template<> PROTOBUF_NOINLINE ::flwr::proto::PushTaskResResponse* Arena::CreateMaybeMessage< ::flwr::proto::PushTaskResResponse >(Arena* arena) {
+ return Arena::CreateMessageInternal< ::flwr::proto::PushTaskResResponse >(arena);
+}
+template<> PROTOBUF_NOINLINE ::flwr::proto::Reconnect* Arena::CreateMaybeMessage< ::flwr::proto::Reconnect >(Arena* arena) {
+ return Arena::CreateMessageInternal< ::flwr::proto::Reconnect >(arena);
+}
+PROTOBUF_NAMESPACE_CLOSE
+
+// @@protoc_insertion_point(global_scope)
+#include
diff --git a/src/cc/flwr/include/flwr/proto/fleet.pb.h b/src/cc/flwr/include/flwr/proto/fleet.pb.h
new file mode 100644
index 000000000000..842e800f5b1c
--- /dev/null
+++ b/src/cc/flwr/include/flwr/proto/fleet.pb.h
@@ -0,0 +1,2202 @@
+// Generated by the protocol buffer compiler. DO NOT EDIT!
+// source: flwr/proto/fleet.proto
+
+#ifndef GOOGLE_PROTOBUF_INCLUDED_flwr_2fproto_2ffleet_2eproto
+#define GOOGLE_PROTOBUF_INCLUDED_flwr_2fproto_2ffleet_2eproto
+
+#include
+#include
+
+#include
+#if PROTOBUF_VERSION < 3018000
+#error This file was generated by a newer version of protoc which is
+#error incompatible with your Protocol Buffer headers. Please update
+#error your headers.
+#endif
+#if 3018001 < PROTOBUF_MIN_PROTOC_VERSION
+#error This file was generated by an older version of protoc which is
+#error incompatible with your Protocol Buffer headers. Please
+#error regenerate this file with a newer version of protoc.
+#endif
+
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include
+#include // IWYU pragma: export
+#include // IWYU pragma: export
+#include // IWYU pragma: export
+#include
+#include
+#include
+#include "flwr/proto/node.pb.h"
+#include "flwr/proto/task.pb.h"
+// @@protoc_insertion_point(includes)
+#include
+#define PROTOBUF_INTERNAL_EXPORT_flwr_2fproto_2ffleet_2eproto
+PROTOBUF_NAMESPACE_OPEN
+namespace internal {
+class AnyMetadata;
+} // namespace internal
+PROTOBUF_NAMESPACE_CLOSE
+
+// Internal implementation detail -- do not use these members.
+struct TableStruct_flwr_2fproto_2ffleet_2eproto {
+ static const ::PROTOBUF_NAMESPACE_ID::internal::ParseTableField entries[]
+ PROTOBUF_SECTION_VARIABLE(protodesc_cold);
+ static const ::PROTOBUF_NAMESPACE_ID::internal::AuxiliaryParseTableField aux[]
+ PROTOBUF_SECTION_VARIABLE(protodesc_cold);
+ static const ::PROTOBUF_NAMESPACE_ID::internal::ParseTable schema[10]
+ PROTOBUF_SECTION_VARIABLE(protodesc_cold);
+ static const ::PROTOBUF_NAMESPACE_ID::internal::FieldMetadata field_metadata[];
+ static const ::PROTOBUF_NAMESPACE_ID::internal::SerializationTable serialization_table[];
+ static const ::PROTOBUF_NAMESPACE_ID::uint32 offsets[];
+};
+extern const ::PROTOBUF_NAMESPACE_ID::internal::DescriptorTable descriptor_table_flwr_2fproto_2ffleet_2eproto;
+namespace flwr {
+namespace proto {
+class CreateNodeRequest;
+struct CreateNodeRequestDefaultTypeInternal;
+extern CreateNodeRequestDefaultTypeInternal _CreateNodeRequest_default_instance_;
+class CreateNodeResponse;
+struct CreateNodeResponseDefaultTypeInternal;
+extern CreateNodeResponseDefaultTypeInternal _CreateNodeResponse_default_instance_;
+class DeleteNodeRequest;
+struct DeleteNodeRequestDefaultTypeInternal;
+extern DeleteNodeRequestDefaultTypeInternal _DeleteNodeRequest_default_instance_;
+class DeleteNodeResponse;
+struct DeleteNodeResponseDefaultTypeInternal;
+extern DeleteNodeResponseDefaultTypeInternal _DeleteNodeResponse_default_instance_;
+class PullTaskInsRequest;
+struct PullTaskInsRequestDefaultTypeInternal;
+extern PullTaskInsRequestDefaultTypeInternal _PullTaskInsRequest_default_instance_;
+class PullTaskInsResponse;
+struct PullTaskInsResponseDefaultTypeInternal;
+extern PullTaskInsResponseDefaultTypeInternal _PullTaskInsResponse_default_instance_;
+class PushTaskResRequest;
+struct PushTaskResRequestDefaultTypeInternal;
+extern PushTaskResRequestDefaultTypeInternal _PushTaskResRequest_default_instance_;
+class PushTaskResResponse;
+struct PushTaskResResponseDefaultTypeInternal;
+extern PushTaskResResponseDefaultTypeInternal _PushTaskResResponse_default_instance_;
+class PushTaskResResponse_ResultsEntry_DoNotUse;
+struct PushTaskResResponse_ResultsEntry_DoNotUseDefaultTypeInternal;
+extern PushTaskResResponse_ResultsEntry_DoNotUseDefaultTypeInternal _PushTaskResResponse_ResultsEntry_DoNotUse_default_instance_;
+class Reconnect;
+struct ReconnectDefaultTypeInternal;
+extern ReconnectDefaultTypeInternal _Reconnect_default_instance_;
+} // namespace proto
+} // namespace flwr
+PROTOBUF_NAMESPACE_OPEN
+template<> ::flwr::proto::CreateNodeRequest* Arena::CreateMaybeMessage<::flwr::proto::CreateNodeRequest>(Arena*);
+template<> ::flwr::proto::CreateNodeResponse* Arena::CreateMaybeMessage<::flwr::proto::CreateNodeResponse>(Arena*);
+template<> ::flwr::proto::DeleteNodeRequest* Arena::CreateMaybeMessage<::flwr::proto::DeleteNodeRequest>(Arena*);
+template<> ::flwr::proto::DeleteNodeResponse* Arena::CreateMaybeMessage<::flwr::proto::DeleteNodeResponse>(Arena*);
+template<> ::flwr::proto::PullTaskInsRequest* Arena::CreateMaybeMessage<::flwr::proto::PullTaskInsRequest>(Arena*);
+template<> ::flwr::proto::PullTaskInsResponse* Arena::CreateMaybeMessage<::flwr::proto::PullTaskInsResponse>(Arena*);
+template<> ::flwr::proto::PushTaskResRequest* Arena::CreateMaybeMessage<::flwr::proto::PushTaskResRequest>(Arena*);
+template<> ::flwr::proto::PushTaskResResponse* Arena::CreateMaybeMessage<::flwr::proto::PushTaskResResponse>(Arena*);
+template<> ::flwr::proto::PushTaskResResponse_ResultsEntry_DoNotUse* Arena::CreateMaybeMessage<::flwr::proto::PushTaskResResponse_ResultsEntry_DoNotUse>(Arena*);
+template<> ::flwr::proto::Reconnect* Arena::CreateMaybeMessage<::flwr::proto::Reconnect>(Arena*);
+PROTOBUF_NAMESPACE_CLOSE
+namespace flwr {
+namespace proto {
+
+// ===================================================================
+
+class CreateNodeRequest final :
+ public ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase /* @@protoc_insertion_point(class_definition:flwr.proto.CreateNodeRequest) */ {
+ public:
+ inline CreateNodeRequest() : CreateNodeRequest(nullptr) {}
+ explicit constexpr CreateNodeRequest(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized);
+
+ CreateNodeRequest(const CreateNodeRequest& from);
+ CreateNodeRequest(CreateNodeRequest&& from) noexcept
+ : CreateNodeRequest() {
+ *this = ::std::move(from);
+ }
+
+ inline CreateNodeRequest& operator=(const CreateNodeRequest& from) {
+ CopyFrom(from);
+ return *this;
+ }
+ inline CreateNodeRequest& operator=(CreateNodeRequest&& from) noexcept {
+ if (this == &from) return *this;
+ if (GetOwningArena() == from.GetOwningArena()
+ #ifdef PROTOBUF_FORCE_COPY_IN_MOVE
+ && GetOwningArena() != nullptr
+ #endif // !PROTOBUF_FORCE_COPY_IN_MOVE
+ ) {
+ InternalSwap(&from);
+ } else {
+ CopyFrom(from);
+ }
+ return *this;
+ }
+
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* descriptor() {
+ return GetDescriptor();
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* GetDescriptor() {
+ return default_instance().GetMetadata().descriptor;
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Reflection* GetReflection() {
+ return default_instance().GetMetadata().reflection;
+ }
+ static const CreateNodeRequest& default_instance() {
+ return *internal_default_instance();
+ }
+ static inline const CreateNodeRequest* internal_default_instance() {
+ return reinterpret_cast(
+ &_CreateNodeRequest_default_instance_);
+ }
+ static constexpr int kIndexInFileMessages =
+ 0;
+
+ friend void swap(CreateNodeRequest& a, CreateNodeRequest& b) {
+ a.Swap(&b);
+ }
+ inline void Swap(CreateNodeRequest* other) {
+ if (other == this) return;
+ if (GetOwningArena() == other->GetOwningArena()) {
+ InternalSwap(other);
+ } else {
+ ::PROTOBUF_NAMESPACE_ID::internal::GenericSwap(this, other);
+ }
+ }
+ void UnsafeArenaSwap(CreateNodeRequest* other) {
+ if (other == this) return;
+ GOOGLE_DCHECK(GetOwningArena() == other->GetOwningArena());
+ InternalSwap(other);
+ }
+
+ // implements Message ----------------------------------------------
+
+ inline CreateNodeRequest* New() const final {
+ return new CreateNodeRequest();
+ }
+
+ CreateNodeRequest* New(::PROTOBUF_NAMESPACE_ID::Arena* arena) const final {
+ return CreateMaybeMessage(arena);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::CopyFrom;
+ inline void CopyFrom(const CreateNodeRequest& from) {
+ ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::CopyImpl(this, from);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::MergeFrom;
+ void MergeFrom(const CreateNodeRequest& from) {
+ ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::MergeImpl(this, from);
+ }
+ public:
+ friend class ::PROTOBUF_NAMESPACE_ID::internal::AnyMetadata;
+ static ::PROTOBUF_NAMESPACE_ID::StringPiece FullMessageName() {
+ return "flwr.proto.CreateNodeRequest";
+ }
+ protected:
+ explicit CreateNodeRequest(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned = false);
+ private:
+ public:
+
+ static const ClassData _class_data_;
+ const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*GetClassData() const final;
+
+ ::PROTOBUF_NAMESPACE_ID::Metadata GetMetadata() const final;
+
+ // nested types ----------------------------------------------------
+
+ // accessors -------------------------------------------------------
+
+ // @@protoc_insertion_point(class_scope:flwr.proto.CreateNodeRequest)
+ private:
+ class _Internal;
+
+ template friend class ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper;
+ typedef void InternalArenaConstructable_;
+ typedef void DestructorSkippable_;
+ mutable ::PROTOBUF_NAMESPACE_ID::internal::CachedSize _cached_size_;
+ friend struct ::TableStruct_flwr_2fproto_2ffleet_2eproto;
+};
+// -------------------------------------------------------------------
+
+class CreateNodeResponse final :
+ public ::PROTOBUF_NAMESPACE_ID::Message /* @@protoc_insertion_point(class_definition:flwr.proto.CreateNodeResponse) */ {
+ public:
+ inline CreateNodeResponse() : CreateNodeResponse(nullptr) {}
+ ~CreateNodeResponse() override;
+ explicit constexpr CreateNodeResponse(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized);
+
+ CreateNodeResponse(const CreateNodeResponse& from);
+ CreateNodeResponse(CreateNodeResponse&& from) noexcept
+ : CreateNodeResponse() {
+ *this = ::std::move(from);
+ }
+
+ inline CreateNodeResponse& operator=(const CreateNodeResponse& from) {
+ CopyFrom(from);
+ return *this;
+ }
+ inline CreateNodeResponse& operator=(CreateNodeResponse&& from) noexcept {
+ if (this == &from) return *this;
+ if (GetOwningArena() == from.GetOwningArena()
+ #ifdef PROTOBUF_FORCE_COPY_IN_MOVE
+ && GetOwningArena() != nullptr
+ #endif // !PROTOBUF_FORCE_COPY_IN_MOVE
+ ) {
+ InternalSwap(&from);
+ } else {
+ CopyFrom(from);
+ }
+ return *this;
+ }
+
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* descriptor() {
+ return GetDescriptor();
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* GetDescriptor() {
+ return default_instance().GetMetadata().descriptor;
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Reflection* GetReflection() {
+ return default_instance().GetMetadata().reflection;
+ }
+ static const CreateNodeResponse& default_instance() {
+ return *internal_default_instance();
+ }
+ static inline const CreateNodeResponse* internal_default_instance() {
+ return reinterpret_cast(
+ &_CreateNodeResponse_default_instance_);
+ }
+ static constexpr int kIndexInFileMessages =
+ 1;
+
+ friend void swap(CreateNodeResponse& a, CreateNodeResponse& b) {
+ a.Swap(&b);
+ }
+ inline void Swap(CreateNodeResponse* other) {
+ if (other == this) return;
+ if (GetOwningArena() == other->GetOwningArena()) {
+ InternalSwap(other);
+ } else {
+ ::PROTOBUF_NAMESPACE_ID::internal::GenericSwap(this, other);
+ }
+ }
+ void UnsafeArenaSwap(CreateNodeResponse* other) {
+ if (other == this) return;
+ GOOGLE_DCHECK(GetOwningArena() == other->GetOwningArena());
+ InternalSwap(other);
+ }
+
+ // implements Message ----------------------------------------------
+
+ inline CreateNodeResponse* New() const final {
+ return new CreateNodeResponse();
+ }
+
+ CreateNodeResponse* New(::PROTOBUF_NAMESPACE_ID::Arena* arena) const final {
+ return CreateMaybeMessage(arena);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::Message::CopyFrom;
+ void CopyFrom(const CreateNodeResponse& from);
+ using ::PROTOBUF_NAMESPACE_ID::Message::MergeFrom;
+ void MergeFrom(const CreateNodeResponse& from);
+ private:
+ static void MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to, const ::PROTOBUF_NAMESPACE_ID::Message& from);
+ public:
+ PROTOBUF_ATTRIBUTE_REINITIALIZES void Clear() final;
+ bool IsInitialized() const final;
+
+ size_t ByteSizeLong() const final;
+ const char* _InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) final;
+ ::PROTOBUF_NAMESPACE_ID::uint8* _InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const final;
+ int GetCachedSize() const final { return _cached_size_.Get(); }
+
+ private:
+ void SharedCtor();
+ void SharedDtor();
+ void SetCachedSize(int size) const final;
+ void InternalSwap(CreateNodeResponse* other);
+ friend class ::PROTOBUF_NAMESPACE_ID::internal::AnyMetadata;
+ static ::PROTOBUF_NAMESPACE_ID::StringPiece FullMessageName() {
+ return "flwr.proto.CreateNodeResponse";
+ }
+ protected:
+ explicit CreateNodeResponse(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned = false);
+ private:
+ static void ArenaDtor(void* object);
+ inline void RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena* arena);
+ public:
+
+ static const ClassData _class_data_;
+ const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*GetClassData() const final;
+
+ ::PROTOBUF_NAMESPACE_ID::Metadata GetMetadata() const final;
+
+ // nested types ----------------------------------------------------
+
+ // accessors -------------------------------------------------------
+
+ enum : int {
+ kNodeFieldNumber = 1,
+ };
+ // .flwr.proto.Node node = 1;
+ bool has_node() const;
+ private:
+ bool _internal_has_node() const;
+ public:
+ void clear_node();
+ const ::flwr::proto::Node& node() const;
+ PROTOBUF_MUST_USE_RESULT ::flwr::proto::Node* release_node();
+ ::flwr::proto::Node* mutable_node();
+ void set_allocated_node(::flwr::proto::Node* node);
+ private:
+ const ::flwr::proto::Node& _internal_node() const;
+ ::flwr::proto::Node* _internal_mutable_node();
+ public:
+ void unsafe_arena_set_allocated_node(
+ ::flwr::proto::Node* node);
+ ::flwr::proto::Node* unsafe_arena_release_node();
+
+ // @@protoc_insertion_point(class_scope:flwr.proto.CreateNodeResponse)
+ private:
+ class _Internal;
+
+ template friend class ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper;
+ typedef void InternalArenaConstructable_;
+ typedef void DestructorSkippable_;
+ ::flwr::proto::Node* node_;
+ mutable ::PROTOBUF_NAMESPACE_ID::internal::CachedSize _cached_size_;
+ friend struct ::TableStruct_flwr_2fproto_2ffleet_2eproto;
+};
+// -------------------------------------------------------------------
+
+class DeleteNodeRequest final :
+ public ::PROTOBUF_NAMESPACE_ID::Message /* @@protoc_insertion_point(class_definition:flwr.proto.DeleteNodeRequest) */ {
+ public:
+ inline DeleteNodeRequest() : DeleteNodeRequest(nullptr) {}
+ ~DeleteNodeRequest() override;
+ explicit constexpr DeleteNodeRequest(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized);
+
+ DeleteNodeRequest(const DeleteNodeRequest& from);
+ DeleteNodeRequest(DeleteNodeRequest&& from) noexcept
+ : DeleteNodeRequest() {
+ *this = ::std::move(from);
+ }
+
+ inline DeleteNodeRequest& operator=(const DeleteNodeRequest& from) {
+ CopyFrom(from);
+ return *this;
+ }
+ inline DeleteNodeRequest& operator=(DeleteNodeRequest&& from) noexcept {
+ if (this == &from) return *this;
+ if (GetOwningArena() == from.GetOwningArena()
+ #ifdef PROTOBUF_FORCE_COPY_IN_MOVE
+ && GetOwningArena() != nullptr
+ #endif // !PROTOBUF_FORCE_COPY_IN_MOVE
+ ) {
+ InternalSwap(&from);
+ } else {
+ CopyFrom(from);
+ }
+ return *this;
+ }
+
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* descriptor() {
+ return GetDescriptor();
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* GetDescriptor() {
+ return default_instance().GetMetadata().descriptor;
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Reflection* GetReflection() {
+ return default_instance().GetMetadata().reflection;
+ }
+ static const DeleteNodeRequest& default_instance() {
+ return *internal_default_instance();
+ }
+ static inline const DeleteNodeRequest* internal_default_instance() {
+ return reinterpret_cast(
+ &_DeleteNodeRequest_default_instance_);
+ }
+ static constexpr int kIndexInFileMessages =
+ 2;
+
+ friend void swap(DeleteNodeRequest& a, DeleteNodeRequest& b) {
+ a.Swap(&b);
+ }
+ inline void Swap(DeleteNodeRequest* other) {
+ if (other == this) return;
+ if (GetOwningArena() == other->GetOwningArena()) {
+ InternalSwap(other);
+ } else {
+ ::PROTOBUF_NAMESPACE_ID::internal::GenericSwap(this, other);
+ }
+ }
+ void UnsafeArenaSwap(DeleteNodeRequest* other) {
+ if (other == this) return;
+ GOOGLE_DCHECK(GetOwningArena() == other->GetOwningArena());
+ InternalSwap(other);
+ }
+
+ // implements Message ----------------------------------------------
+
+ inline DeleteNodeRequest* New() const final {
+ return new DeleteNodeRequest();
+ }
+
+ DeleteNodeRequest* New(::PROTOBUF_NAMESPACE_ID::Arena* arena) const final {
+ return CreateMaybeMessage(arena);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::Message::CopyFrom;
+ void CopyFrom(const DeleteNodeRequest& from);
+ using ::PROTOBUF_NAMESPACE_ID::Message::MergeFrom;
+ void MergeFrom(const DeleteNodeRequest& from);
+ private:
+ static void MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to, const ::PROTOBUF_NAMESPACE_ID::Message& from);
+ public:
+ PROTOBUF_ATTRIBUTE_REINITIALIZES void Clear() final;
+ bool IsInitialized() const final;
+
+ size_t ByteSizeLong() const final;
+ const char* _InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) final;
+ ::PROTOBUF_NAMESPACE_ID::uint8* _InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const final;
+ int GetCachedSize() const final { return _cached_size_.Get(); }
+
+ private:
+ void SharedCtor();
+ void SharedDtor();
+ void SetCachedSize(int size) const final;
+ void InternalSwap(DeleteNodeRequest* other);
+ friend class ::PROTOBUF_NAMESPACE_ID::internal::AnyMetadata;
+ static ::PROTOBUF_NAMESPACE_ID::StringPiece FullMessageName() {
+ return "flwr.proto.DeleteNodeRequest";
+ }
+ protected:
+ explicit DeleteNodeRequest(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned = false);
+ private:
+ static void ArenaDtor(void* object);
+ inline void RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena* arena);
+ public:
+
+ static const ClassData _class_data_;
+ const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*GetClassData() const final;
+
+ ::PROTOBUF_NAMESPACE_ID::Metadata GetMetadata() const final;
+
+ // nested types ----------------------------------------------------
+
+ // accessors -------------------------------------------------------
+
+ enum : int {
+ kNodeFieldNumber = 1,
+ };
+ // .flwr.proto.Node node = 1;
+ bool has_node() const;
+ private:
+ bool _internal_has_node() const;
+ public:
+ void clear_node();
+ const ::flwr::proto::Node& node() const;
+ PROTOBUF_MUST_USE_RESULT ::flwr::proto::Node* release_node();
+ ::flwr::proto::Node* mutable_node();
+ void set_allocated_node(::flwr::proto::Node* node);
+ private:
+ const ::flwr::proto::Node& _internal_node() const;
+ ::flwr::proto::Node* _internal_mutable_node();
+ public:
+ void unsafe_arena_set_allocated_node(
+ ::flwr::proto::Node* node);
+ ::flwr::proto::Node* unsafe_arena_release_node();
+
+ // @@protoc_insertion_point(class_scope:flwr.proto.DeleteNodeRequest)
+ private:
+ class _Internal;
+
+ template friend class ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper;
+ typedef void InternalArenaConstructable_;
+ typedef void DestructorSkippable_;
+ ::flwr::proto::Node* node_;
+ mutable ::PROTOBUF_NAMESPACE_ID::internal::CachedSize _cached_size_;
+ friend struct ::TableStruct_flwr_2fproto_2ffleet_2eproto;
+};
+// -------------------------------------------------------------------
+
+class DeleteNodeResponse final :
+ public ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase /* @@protoc_insertion_point(class_definition:flwr.proto.DeleteNodeResponse) */ {
+ public:
+ inline DeleteNodeResponse() : DeleteNodeResponse(nullptr) {}
+ explicit constexpr DeleteNodeResponse(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized);
+
+ DeleteNodeResponse(const DeleteNodeResponse& from);
+ DeleteNodeResponse(DeleteNodeResponse&& from) noexcept
+ : DeleteNodeResponse() {
+ *this = ::std::move(from);
+ }
+
+ inline DeleteNodeResponse& operator=(const DeleteNodeResponse& from) {
+ CopyFrom(from);
+ return *this;
+ }
+ inline DeleteNodeResponse& operator=(DeleteNodeResponse&& from) noexcept {
+ if (this == &from) return *this;
+ if (GetOwningArena() == from.GetOwningArena()
+ #ifdef PROTOBUF_FORCE_COPY_IN_MOVE
+ && GetOwningArena() != nullptr
+ #endif // !PROTOBUF_FORCE_COPY_IN_MOVE
+ ) {
+ InternalSwap(&from);
+ } else {
+ CopyFrom(from);
+ }
+ return *this;
+ }
+
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* descriptor() {
+ return GetDescriptor();
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* GetDescriptor() {
+ return default_instance().GetMetadata().descriptor;
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Reflection* GetReflection() {
+ return default_instance().GetMetadata().reflection;
+ }
+ static const DeleteNodeResponse& default_instance() {
+ return *internal_default_instance();
+ }
+ static inline const DeleteNodeResponse* internal_default_instance() {
+ return reinterpret_cast(
+ &_DeleteNodeResponse_default_instance_);
+ }
+ static constexpr int kIndexInFileMessages =
+ 3;
+
+ friend void swap(DeleteNodeResponse& a, DeleteNodeResponse& b) {
+ a.Swap(&b);
+ }
+ inline void Swap(DeleteNodeResponse* other) {
+ if (other == this) return;
+ if (GetOwningArena() == other->GetOwningArena()) {
+ InternalSwap(other);
+ } else {
+ ::PROTOBUF_NAMESPACE_ID::internal::GenericSwap(this, other);
+ }
+ }
+ void UnsafeArenaSwap(DeleteNodeResponse* other) {
+ if (other == this) return;
+ GOOGLE_DCHECK(GetOwningArena() == other->GetOwningArena());
+ InternalSwap(other);
+ }
+
+ // implements Message ----------------------------------------------
+
+ inline DeleteNodeResponse* New() const final {
+ return new DeleteNodeResponse();
+ }
+
+ DeleteNodeResponse* New(::PROTOBUF_NAMESPACE_ID::Arena* arena) const final {
+ return CreateMaybeMessage(arena);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::CopyFrom;
+ inline void CopyFrom(const DeleteNodeResponse& from) {
+ ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::CopyImpl(this, from);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::MergeFrom;
+ void MergeFrom(const DeleteNodeResponse& from) {
+ ::PROTOBUF_NAMESPACE_ID::internal::ZeroFieldsBase::MergeImpl(this, from);
+ }
+ public:
+ friend class ::PROTOBUF_NAMESPACE_ID::internal::AnyMetadata;
+ static ::PROTOBUF_NAMESPACE_ID::StringPiece FullMessageName() {
+ return "flwr.proto.DeleteNodeResponse";
+ }
+ protected:
+ explicit DeleteNodeResponse(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned = false);
+ private:
+ public:
+
+ static const ClassData _class_data_;
+ const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*GetClassData() const final;
+
+ ::PROTOBUF_NAMESPACE_ID::Metadata GetMetadata() const final;
+
+ // nested types ----------------------------------------------------
+
+ // accessors -------------------------------------------------------
+
+ // @@protoc_insertion_point(class_scope:flwr.proto.DeleteNodeResponse)
+ private:
+ class _Internal;
+
+ template friend class ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper;
+ typedef void InternalArenaConstructable_;
+ typedef void DestructorSkippable_;
+ mutable ::PROTOBUF_NAMESPACE_ID::internal::CachedSize _cached_size_;
+ friend struct ::TableStruct_flwr_2fproto_2ffleet_2eproto;
+};
+// -------------------------------------------------------------------
+
+class PullTaskInsRequest final :
+ public ::PROTOBUF_NAMESPACE_ID::Message /* @@protoc_insertion_point(class_definition:flwr.proto.PullTaskInsRequest) */ {
+ public:
+ inline PullTaskInsRequest() : PullTaskInsRequest(nullptr) {}
+ ~PullTaskInsRequest() override;
+ explicit constexpr PullTaskInsRequest(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized);
+
+ PullTaskInsRequest(const PullTaskInsRequest& from);
+ PullTaskInsRequest(PullTaskInsRequest&& from) noexcept
+ : PullTaskInsRequest() {
+ *this = ::std::move(from);
+ }
+
+ inline PullTaskInsRequest& operator=(const PullTaskInsRequest& from) {
+ CopyFrom(from);
+ return *this;
+ }
+ inline PullTaskInsRequest& operator=(PullTaskInsRequest&& from) noexcept {
+ if (this == &from) return *this;
+ if (GetOwningArena() == from.GetOwningArena()
+ #ifdef PROTOBUF_FORCE_COPY_IN_MOVE
+ && GetOwningArena() != nullptr
+ #endif // !PROTOBUF_FORCE_COPY_IN_MOVE
+ ) {
+ InternalSwap(&from);
+ } else {
+ CopyFrom(from);
+ }
+ return *this;
+ }
+
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* descriptor() {
+ return GetDescriptor();
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* GetDescriptor() {
+ return default_instance().GetMetadata().descriptor;
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Reflection* GetReflection() {
+ return default_instance().GetMetadata().reflection;
+ }
+ static const PullTaskInsRequest& default_instance() {
+ return *internal_default_instance();
+ }
+ static inline const PullTaskInsRequest* internal_default_instance() {
+ return reinterpret_cast(
+ &_PullTaskInsRequest_default_instance_);
+ }
+ static constexpr int kIndexInFileMessages =
+ 4;
+
+ friend void swap(PullTaskInsRequest& a, PullTaskInsRequest& b) {
+ a.Swap(&b);
+ }
+ inline void Swap(PullTaskInsRequest* other) {
+ if (other == this) return;
+ if (GetOwningArena() == other->GetOwningArena()) {
+ InternalSwap(other);
+ } else {
+ ::PROTOBUF_NAMESPACE_ID::internal::GenericSwap(this, other);
+ }
+ }
+ void UnsafeArenaSwap(PullTaskInsRequest* other) {
+ if (other == this) return;
+ GOOGLE_DCHECK(GetOwningArena() == other->GetOwningArena());
+ InternalSwap(other);
+ }
+
+ // implements Message ----------------------------------------------
+
+ inline PullTaskInsRequest* New() const final {
+ return new PullTaskInsRequest();
+ }
+
+ PullTaskInsRequest* New(::PROTOBUF_NAMESPACE_ID::Arena* arena) const final {
+ return CreateMaybeMessage(arena);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::Message::CopyFrom;
+ void CopyFrom(const PullTaskInsRequest& from);
+ using ::PROTOBUF_NAMESPACE_ID::Message::MergeFrom;
+ void MergeFrom(const PullTaskInsRequest& from);
+ private:
+ static void MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to, const ::PROTOBUF_NAMESPACE_ID::Message& from);
+ public:
+ PROTOBUF_ATTRIBUTE_REINITIALIZES void Clear() final;
+ bool IsInitialized() const final;
+
+ size_t ByteSizeLong() const final;
+ const char* _InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) final;
+ ::PROTOBUF_NAMESPACE_ID::uint8* _InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const final;
+ int GetCachedSize() const final { return _cached_size_.Get(); }
+
+ private:
+ void SharedCtor();
+ void SharedDtor();
+ void SetCachedSize(int size) const final;
+ void InternalSwap(PullTaskInsRequest* other);
+ friend class ::PROTOBUF_NAMESPACE_ID::internal::AnyMetadata;
+ static ::PROTOBUF_NAMESPACE_ID::StringPiece FullMessageName() {
+ return "flwr.proto.PullTaskInsRequest";
+ }
+ protected:
+ explicit PullTaskInsRequest(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned = false);
+ private:
+ static void ArenaDtor(void* object);
+ inline void RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena* arena);
+ public:
+
+ static const ClassData _class_data_;
+ const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*GetClassData() const final;
+
+ ::PROTOBUF_NAMESPACE_ID::Metadata GetMetadata() const final;
+
+ // nested types ----------------------------------------------------
+
+ // accessors -------------------------------------------------------
+
+ enum : int {
+ kTaskIdsFieldNumber = 2,
+ kNodeFieldNumber = 1,
+ };
+ // repeated string task_ids = 2;
+ int task_ids_size() const;
+ private:
+ int _internal_task_ids_size() const;
+ public:
+ void clear_task_ids();
+ const std::string& task_ids(int index) const;
+ std::string* mutable_task_ids(int index);
+ void set_task_ids(int index, const std::string& value);
+ void set_task_ids(int index, std::string&& value);
+ void set_task_ids(int index, const char* value);
+ void set_task_ids(int index, const char* value, size_t size);
+ std::string* add_task_ids();
+ void add_task_ids(const std::string& value);
+ void add_task_ids(std::string&& value);
+ void add_task_ids(const char* value);
+ void add_task_ids(const char* value, size_t size);
+ const ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField& task_ids() const;
+ ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField* mutable_task_ids();
+ private:
+ const std::string& _internal_task_ids(int index) const;
+ std::string* _internal_add_task_ids();
+ public:
+
+ // .flwr.proto.Node node = 1;
+ bool has_node() const;
+ private:
+ bool _internal_has_node() const;
+ public:
+ void clear_node();
+ const ::flwr::proto::Node& node() const;
+ PROTOBUF_MUST_USE_RESULT ::flwr::proto::Node* release_node();
+ ::flwr::proto::Node* mutable_node();
+ void set_allocated_node(::flwr::proto::Node* node);
+ private:
+ const ::flwr::proto::Node& _internal_node() const;
+ ::flwr::proto::Node* _internal_mutable_node();
+ public:
+ void unsafe_arena_set_allocated_node(
+ ::flwr::proto::Node* node);
+ ::flwr::proto::Node* unsafe_arena_release_node();
+
+ // @@protoc_insertion_point(class_scope:flwr.proto.PullTaskInsRequest)
+ private:
+ class _Internal;
+
+ template friend class ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper;
+ typedef void InternalArenaConstructable_;
+ typedef void DestructorSkippable_;
+ ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField task_ids_;
+ ::flwr::proto::Node* node_;
+ mutable ::PROTOBUF_NAMESPACE_ID::internal::CachedSize _cached_size_;
+ friend struct ::TableStruct_flwr_2fproto_2ffleet_2eproto;
+};
+// -------------------------------------------------------------------
+
+class PullTaskInsResponse final :
+ public ::PROTOBUF_NAMESPACE_ID::Message /* @@protoc_insertion_point(class_definition:flwr.proto.PullTaskInsResponse) */ {
+ public:
+ inline PullTaskInsResponse() : PullTaskInsResponse(nullptr) {}
+ ~PullTaskInsResponse() override;
+ explicit constexpr PullTaskInsResponse(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized);
+
+ PullTaskInsResponse(const PullTaskInsResponse& from);
+ PullTaskInsResponse(PullTaskInsResponse&& from) noexcept
+ : PullTaskInsResponse() {
+ *this = ::std::move(from);
+ }
+
+ inline PullTaskInsResponse& operator=(const PullTaskInsResponse& from) {
+ CopyFrom(from);
+ return *this;
+ }
+ inline PullTaskInsResponse& operator=(PullTaskInsResponse&& from) noexcept {
+ if (this == &from) return *this;
+ if (GetOwningArena() == from.GetOwningArena()
+ #ifdef PROTOBUF_FORCE_COPY_IN_MOVE
+ && GetOwningArena() != nullptr
+ #endif // !PROTOBUF_FORCE_COPY_IN_MOVE
+ ) {
+ InternalSwap(&from);
+ } else {
+ CopyFrom(from);
+ }
+ return *this;
+ }
+
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* descriptor() {
+ return GetDescriptor();
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* GetDescriptor() {
+ return default_instance().GetMetadata().descriptor;
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Reflection* GetReflection() {
+ return default_instance().GetMetadata().reflection;
+ }
+ static const PullTaskInsResponse& default_instance() {
+ return *internal_default_instance();
+ }
+ static inline const PullTaskInsResponse* internal_default_instance() {
+ return reinterpret_cast(
+ &_PullTaskInsResponse_default_instance_);
+ }
+ static constexpr int kIndexInFileMessages =
+ 5;
+
+ friend void swap(PullTaskInsResponse& a, PullTaskInsResponse& b) {
+ a.Swap(&b);
+ }
+ inline void Swap(PullTaskInsResponse* other) {
+ if (other == this) return;
+ if (GetOwningArena() == other->GetOwningArena()) {
+ InternalSwap(other);
+ } else {
+ ::PROTOBUF_NAMESPACE_ID::internal::GenericSwap(this, other);
+ }
+ }
+ void UnsafeArenaSwap(PullTaskInsResponse* other) {
+ if (other == this) return;
+ GOOGLE_DCHECK(GetOwningArena() == other->GetOwningArena());
+ InternalSwap(other);
+ }
+
+ // implements Message ----------------------------------------------
+
+ inline PullTaskInsResponse* New() const final {
+ return new PullTaskInsResponse();
+ }
+
+ PullTaskInsResponse* New(::PROTOBUF_NAMESPACE_ID::Arena* arena) const final {
+ return CreateMaybeMessage(arena);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::Message::CopyFrom;
+ void CopyFrom(const PullTaskInsResponse& from);
+ using ::PROTOBUF_NAMESPACE_ID::Message::MergeFrom;
+ void MergeFrom(const PullTaskInsResponse& from);
+ private:
+ static void MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to, const ::PROTOBUF_NAMESPACE_ID::Message& from);
+ public:
+ PROTOBUF_ATTRIBUTE_REINITIALIZES void Clear() final;
+ bool IsInitialized() const final;
+
+ size_t ByteSizeLong() const final;
+ const char* _InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) final;
+ ::PROTOBUF_NAMESPACE_ID::uint8* _InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const final;
+ int GetCachedSize() const final { return _cached_size_.Get(); }
+
+ private:
+ void SharedCtor();
+ void SharedDtor();
+ void SetCachedSize(int size) const final;
+ void InternalSwap(PullTaskInsResponse* other);
+ friend class ::PROTOBUF_NAMESPACE_ID::internal::AnyMetadata;
+ static ::PROTOBUF_NAMESPACE_ID::StringPiece FullMessageName() {
+ return "flwr.proto.PullTaskInsResponse";
+ }
+ protected:
+ explicit PullTaskInsResponse(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned = false);
+ private:
+ static void ArenaDtor(void* object);
+ inline void RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena* arena);
+ public:
+
+ static const ClassData _class_data_;
+ const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*GetClassData() const final;
+
+ ::PROTOBUF_NAMESPACE_ID::Metadata GetMetadata() const final;
+
+ // nested types ----------------------------------------------------
+
+ // accessors -------------------------------------------------------
+
+ enum : int {
+ kTaskInsListFieldNumber = 2,
+ kReconnectFieldNumber = 1,
+ };
+ // repeated .flwr.proto.TaskIns task_ins_list = 2;
+ int task_ins_list_size() const;
+ private:
+ int _internal_task_ins_list_size() const;
+ public:
+ void clear_task_ins_list();
+ ::flwr::proto::TaskIns* mutable_task_ins_list(int index);
+ ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField< ::flwr::proto::TaskIns >*
+ mutable_task_ins_list();
+ private:
+ const ::flwr::proto::TaskIns& _internal_task_ins_list(int index) const;
+ ::flwr::proto::TaskIns* _internal_add_task_ins_list();
+ public:
+ const ::flwr::proto::TaskIns& task_ins_list(int index) const;
+ ::flwr::proto::TaskIns* add_task_ins_list();
+ const ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField< ::flwr::proto::TaskIns >&
+ task_ins_list() const;
+
+ // .flwr.proto.Reconnect reconnect = 1;
+ bool has_reconnect() const;
+ private:
+ bool _internal_has_reconnect() const;
+ public:
+ void clear_reconnect();
+ const ::flwr::proto::Reconnect& reconnect() const;
+ PROTOBUF_MUST_USE_RESULT ::flwr::proto::Reconnect* release_reconnect();
+ ::flwr::proto::Reconnect* mutable_reconnect();
+ void set_allocated_reconnect(::flwr::proto::Reconnect* reconnect);
+ private:
+ const ::flwr::proto::Reconnect& _internal_reconnect() const;
+ ::flwr::proto::Reconnect* _internal_mutable_reconnect();
+ public:
+ void unsafe_arena_set_allocated_reconnect(
+ ::flwr::proto::Reconnect* reconnect);
+ ::flwr::proto::Reconnect* unsafe_arena_release_reconnect();
+
+ // @@protoc_insertion_point(class_scope:flwr.proto.PullTaskInsResponse)
+ private:
+ class _Internal;
+
+ template friend class ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper;
+ typedef void InternalArenaConstructable_;
+ typedef void DestructorSkippable_;
+ ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField< ::flwr::proto::TaskIns > task_ins_list_;
+ ::flwr::proto::Reconnect* reconnect_;
+ mutable ::PROTOBUF_NAMESPACE_ID::internal::CachedSize _cached_size_;
+ friend struct ::TableStruct_flwr_2fproto_2ffleet_2eproto;
+};
+// -------------------------------------------------------------------
+
+class PushTaskResRequest final :
+ public ::PROTOBUF_NAMESPACE_ID::Message /* @@protoc_insertion_point(class_definition:flwr.proto.PushTaskResRequest) */ {
+ public:
+ inline PushTaskResRequest() : PushTaskResRequest(nullptr) {}
+ ~PushTaskResRequest() override;
+ explicit constexpr PushTaskResRequest(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized);
+
+ PushTaskResRequest(const PushTaskResRequest& from);
+ PushTaskResRequest(PushTaskResRequest&& from) noexcept
+ : PushTaskResRequest() {
+ *this = ::std::move(from);
+ }
+
+ inline PushTaskResRequest& operator=(const PushTaskResRequest& from) {
+ CopyFrom(from);
+ return *this;
+ }
+ inline PushTaskResRequest& operator=(PushTaskResRequest&& from) noexcept {
+ if (this == &from) return *this;
+ if (GetOwningArena() == from.GetOwningArena()
+ #ifdef PROTOBUF_FORCE_COPY_IN_MOVE
+ && GetOwningArena() != nullptr
+ #endif // !PROTOBUF_FORCE_COPY_IN_MOVE
+ ) {
+ InternalSwap(&from);
+ } else {
+ CopyFrom(from);
+ }
+ return *this;
+ }
+
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* descriptor() {
+ return GetDescriptor();
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* GetDescriptor() {
+ return default_instance().GetMetadata().descriptor;
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Reflection* GetReflection() {
+ return default_instance().GetMetadata().reflection;
+ }
+ static const PushTaskResRequest& default_instance() {
+ return *internal_default_instance();
+ }
+ static inline const PushTaskResRequest* internal_default_instance() {
+ return reinterpret_cast(
+ &_PushTaskResRequest_default_instance_);
+ }
+ static constexpr int kIndexInFileMessages =
+ 6;
+
+ friend void swap(PushTaskResRequest& a, PushTaskResRequest& b) {
+ a.Swap(&b);
+ }
+ inline void Swap(PushTaskResRequest* other) {
+ if (other == this) return;
+ if (GetOwningArena() == other->GetOwningArena()) {
+ InternalSwap(other);
+ } else {
+ ::PROTOBUF_NAMESPACE_ID::internal::GenericSwap(this, other);
+ }
+ }
+ void UnsafeArenaSwap(PushTaskResRequest* other) {
+ if (other == this) return;
+ GOOGLE_DCHECK(GetOwningArena() == other->GetOwningArena());
+ InternalSwap(other);
+ }
+
+ // implements Message ----------------------------------------------
+
+ inline PushTaskResRequest* New() const final {
+ return new PushTaskResRequest();
+ }
+
+ PushTaskResRequest* New(::PROTOBUF_NAMESPACE_ID::Arena* arena) const final {
+ return CreateMaybeMessage(arena);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::Message::CopyFrom;
+ void CopyFrom(const PushTaskResRequest& from);
+ using ::PROTOBUF_NAMESPACE_ID::Message::MergeFrom;
+ void MergeFrom(const PushTaskResRequest& from);
+ private:
+ static void MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to, const ::PROTOBUF_NAMESPACE_ID::Message& from);
+ public:
+ PROTOBUF_ATTRIBUTE_REINITIALIZES void Clear() final;
+ bool IsInitialized() const final;
+
+ size_t ByteSizeLong() const final;
+ const char* _InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) final;
+ ::PROTOBUF_NAMESPACE_ID::uint8* _InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const final;
+ int GetCachedSize() const final { return _cached_size_.Get(); }
+
+ private:
+ void SharedCtor();
+ void SharedDtor();
+ void SetCachedSize(int size) const final;
+ void InternalSwap(PushTaskResRequest* other);
+ friend class ::PROTOBUF_NAMESPACE_ID::internal::AnyMetadata;
+ static ::PROTOBUF_NAMESPACE_ID::StringPiece FullMessageName() {
+ return "flwr.proto.PushTaskResRequest";
+ }
+ protected:
+ explicit PushTaskResRequest(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned = false);
+ private:
+ static void ArenaDtor(void* object);
+ inline void RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena* arena);
+ public:
+
+ static const ClassData _class_data_;
+ const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*GetClassData() const final;
+
+ ::PROTOBUF_NAMESPACE_ID::Metadata GetMetadata() const final;
+
+ // nested types ----------------------------------------------------
+
+ // accessors -------------------------------------------------------
+
+ enum : int {
+ kTaskResListFieldNumber = 1,
+ };
+ // repeated .flwr.proto.TaskRes task_res_list = 1;
+ int task_res_list_size() const;
+ private:
+ int _internal_task_res_list_size() const;
+ public:
+ void clear_task_res_list();
+ ::flwr::proto::TaskRes* mutable_task_res_list(int index);
+ ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField< ::flwr::proto::TaskRes >*
+ mutable_task_res_list();
+ private:
+ const ::flwr::proto::TaskRes& _internal_task_res_list(int index) const;
+ ::flwr::proto::TaskRes* _internal_add_task_res_list();
+ public:
+ const ::flwr::proto::TaskRes& task_res_list(int index) const;
+ ::flwr::proto::TaskRes* add_task_res_list();
+ const ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField< ::flwr::proto::TaskRes >&
+ task_res_list() const;
+
+ // @@protoc_insertion_point(class_scope:flwr.proto.PushTaskResRequest)
+ private:
+ class _Internal;
+
+ template friend class ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper;
+ typedef void InternalArenaConstructable_;
+ typedef void DestructorSkippable_;
+ ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField< ::flwr::proto::TaskRes > task_res_list_;
+ mutable ::PROTOBUF_NAMESPACE_ID::internal::CachedSize _cached_size_;
+ friend struct ::TableStruct_flwr_2fproto_2ffleet_2eproto;
+};
+// -------------------------------------------------------------------
+
+class PushTaskResResponse_ResultsEntry_DoNotUse : public ::PROTOBUF_NAMESPACE_ID::internal::MapEntry {
+public:
+ typedef ::PROTOBUF_NAMESPACE_ID::internal::MapEntry SuperType;
+ PushTaskResResponse_ResultsEntry_DoNotUse();
+ explicit constexpr PushTaskResResponse_ResultsEntry_DoNotUse(
+ ::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized);
+ explicit PushTaskResResponse_ResultsEntry_DoNotUse(::PROTOBUF_NAMESPACE_ID::Arena* arena);
+ void MergeFrom(const PushTaskResResponse_ResultsEntry_DoNotUse& other);
+ static const PushTaskResResponse_ResultsEntry_DoNotUse* internal_default_instance() { return reinterpret_cast(&_PushTaskResResponse_ResultsEntry_DoNotUse_default_instance_); }
+ static bool ValidateKey(std::string* s) {
+ return ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::VerifyUtf8String(s->data(), static_cast(s->size()), ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::PARSE, "flwr.proto.PushTaskResResponse.ResultsEntry.key");
+ }
+ static bool ValidateValue(void*) { return true; }
+ using ::PROTOBUF_NAMESPACE_ID::Message::MergeFrom;
+ ::PROTOBUF_NAMESPACE_ID::Metadata GetMetadata() const final;
+};
+
+// -------------------------------------------------------------------
+
+class PushTaskResResponse final :
+ public ::PROTOBUF_NAMESPACE_ID::Message /* @@protoc_insertion_point(class_definition:flwr.proto.PushTaskResResponse) */ {
+ public:
+ inline PushTaskResResponse() : PushTaskResResponse(nullptr) {}
+ ~PushTaskResResponse() override;
+ explicit constexpr PushTaskResResponse(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized);
+
+ PushTaskResResponse(const PushTaskResResponse& from);
+ PushTaskResResponse(PushTaskResResponse&& from) noexcept
+ : PushTaskResResponse() {
+ *this = ::std::move(from);
+ }
+
+ inline PushTaskResResponse& operator=(const PushTaskResResponse& from) {
+ CopyFrom(from);
+ return *this;
+ }
+ inline PushTaskResResponse& operator=(PushTaskResResponse&& from) noexcept {
+ if (this == &from) return *this;
+ if (GetOwningArena() == from.GetOwningArena()
+ #ifdef PROTOBUF_FORCE_COPY_IN_MOVE
+ && GetOwningArena() != nullptr
+ #endif // !PROTOBUF_FORCE_COPY_IN_MOVE
+ ) {
+ InternalSwap(&from);
+ } else {
+ CopyFrom(from);
+ }
+ return *this;
+ }
+
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* descriptor() {
+ return GetDescriptor();
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* GetDescriptor() {
+ return default_instance().GetMetadata().descriptor;
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Reflection* GetReflection() {
+ return default_instance().GetMetadata().reflection;
+ }
+ static const PushTaskResResponse& default_instance() {
+ return *internal_default_instance();
+ }
+ static inline const PushTaskResResponse* internal_default_instance() {
+ return reinterpret_cast(
+ &_PushTaskResResponse_default_instance_);
+ }
+ static constexpr int kIndexInFileMessages =
+ 8;
+
+ friend void swap(PushTaskResResponse& a, PushTaskResResponse& b) {
+ a.Swap(&b);
+ }
+ inline void Swap(PushTaskResResponse* other) {
+ if (other == this) return;
+ if (GetOwningArena() == other->GetOwningArena()) {
+ InternalSwap(other);
+ } else {
+ ::PROTOBUF_NAMESPACE_ID::internal::GenericSwap(this, other);
+ }
+ }
+ void UnsafeArenaSwap(PushTaskResResponse* other) {
+ if (other == this) return;
+ GOOGLE_DCHECK(GetOwningArena() == other->GetOwningArena());
+ InternalSwap(other);
+ }
+
+ // implements Message ----------------------------------------------
+
+ inline PushTaskResResponse* New() const final {
+ return new PushTaskResResponse();
+ }
+
+ PushTaskResResponse* New(::PROTOBUF_NAMESPACE_ID::Arena* arena) const final {
+ return CreateMaybeMessage(arena);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::Message::CopyFrom;
+ void CopyFrom(const PushTaskResResponse& from);
+ using ::PROTOBUF_NAMESPACE_ID::Message::MergeFrom;
+ void MergeFrom(const PushTaskResResponse& from);
+ private:
+ static void MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to, const ::PROTOBUF_NAMESPACE_ID::Message& from);
+ public:
+ PROTOBUF_ATTRIBUTE_REINITIALIZES void Clear() final;
+ bool IsInitialized() const final;
+
+ size_t ByteSizeLong() const final;
+ const char* _InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) final;
+ ::PROTOBUF_NAMESPACE_ID::uint8* _InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const final;
+ int GetCachedSize() const final { return _cached_size_.Get(); }
+
+ private:
+ void SharedCtor();
+ void SharedDtor();
+ void SetCachedSize(int size) const final;
+ void InternalSwap(PushTaskResResponse* other);
+ friend class ::PROTOBUF_NAMESPACE_ID::internal::AnyMetadata;
+ static ::PROTOBUF_NAMESPACE_ID::StringPiece FullMessageName() {
+ return "flwr.proto.PushTaskResResponse";
+ }
+ protected:
+ explicit PushTaskResResponse(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned = false);
+ private:
+ static void ArenaDtor(void* object);
+ inline void RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena* arena);
+ public:
+
+ static const ClassData _class_data_;
+ const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*GetClassData() const final;
+
+ ::PROTOBUF_NAMESPACE_ID::Metadata GetMetadata() const final;
+
+ // nested types ----------------------------------------------------
+
+
+ // accessors -------------------------------------------------------
+
+ enum : int {
+ kResultsFieldNumber = 2,
+ kReconnectFieldNumber = 1,
+ };
+ // map results = 2;
+ int results_size() const;
+ private:
+ int _internal_results_size() const;
+ public:
+ void clear_results();
+ private:
+ const ::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >&
+ _internal_results() const;
+ ::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >*
+ _internal_mutable_results();
+ public:
+ const ::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >&
+ results() const;
+ ::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >*
+ mutable_results();
+
+ // .flwr.proto.Reconnect reconnect = 1;
+ bool has_reconnect() const;
+ private:
+ bool _internal_has_reconnect() const;
+ public:
+ void clear_reconnect();
+ const ::flwr::proto::Reconnect& reconnect() const;
+ PROTOBUF_MUST_USE_RESULT ::flwr::proto::Reconnect* release_reconnect();
+ ::flwr::proto::Reconnect* mutable_reconnect();
+ void set_allocated_reconnect(::flwr::proto::Reconnect* reconnect);
+ private:
+ const ::flwr::proto::Reconnect& _internal_reconnect() const;
+ ::flwr::proto::Reconnect* _internal_mutable_reconnect();
+ public:
+ void unsafe_arena_set_allocated_reconnect(
+ ::flwr::proto::Reconnect* reconnect);
+ ::flwr::proto::Reconnect* unsafe_arena_release_reconnect();
+
+ // @@protoc_insertion_point(class_scope:flwr.proto.PushTaskResResponse)
+ private:
+ class _Internal;
+
+ template friend class ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper;
+ typedef void InternalArenaConstructable_;
+ typedef void DestructorSkippable_;
+ ::PROTOBUF_NAMESPACE_ID::internal::MapField<
+ PushTaskResResponse_ResultsEntry_DoNotUse,
+ std::string, ::PROTOBUF_NAMESPACE_ID::uint32,
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::TYPE_STRING,
+ ::PROTOBUF_NAMESPACE_ID::internal::WireFormatLite::TYPE_UINT32> results_;
+ ::flwr::proto::Reconnect* reconnect_;
+ mutable ::PROTOBUF_NAMESPACE_ID::internal::CachedSize _cached_size_;
+ friend struct ::TableStruct_flwr_2fproto_2ffleet_2eproto;
+};
+// -------------------------------------------------------------------
+
+class Reconnect final :
+ public ::PROTOBUF_NAMESPACE_ID::Message /* @@protoc_insertion_point(class_definition:flwr.proto.Reconnect) */ {
+ public:
+ inline Reconnect() : Reconnect(nullptr) {}
+ ~Reconnect() override;
+ explicit constexpr Reconnect(::PROTOBUF_NAMESPACE_ID::internal::ConstantInitialized);
+
+ Reconnect(const Reconnect& from);
+ Reconnect(Reconnect&& from) noexcept
+ : Reconnect() {
+ *this = ::std::move(from);
+ }
+
+ inline Reconnect& operator=(const Reconnect& from) {
+ CopyFrom(from);
+ return *this;
+ }
+ inline Reconnect& operator=(Reconnect&& from) noexcept {
+ if (this == &from) return *this;
+ if (GetOwningArena() == from.GetOwningArena()
+ #ifdef PROTOBUF_FORCE_COPY_IN_MOVE
+ && GetOwningArena() != nullptr
+ #endif // !PROTOBUF_FORCE_COPY_IN_MOVE
+ ) {
+ InternalSwap(&from);
+ } else {
+ CopyFrom(from);
+ }
+ return *this;
+ }
+
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* descriptor() {
+ return GetDescriptor();
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Descriptor* GetDescriptor() {
+ return default_instance().GetMetadata().descriptor;
+ }
+ static const ::PROTOBUF_NAMESPACE_ID::Reflection* GetReflection() {
+ return default_instance().GetMetadata().reflection;
+ }
+ static const Reconnect& default_instance() {
+ return *internal_default_instance();
+ }
+ static inline const Reconnect* internal_default_instance() {
+ return reinterpret_cast(
+ &_Reconnect_default_instance_);
+ }
+ static constexpr int kIndexInFileMessages =
+ 9;
+
+ friend void swap(Reconnect& a, Reconnect& b) {
+ a.Swap(&b);
+ }
+ inline void Swap(Reconnect* other) {
+ if (other == this) return;
+ if (GetOwningArena() == other->GetOwningArena()) {
+ InternalSwap(other);
+ } else {
+ ::PROTOBUF_NAMESPACE_ID::internal::GenericSwap(this, other);
+ }
+ }
+ void UnsafeArenaSwap(Reconnect* other) {
+ if (other == this) return;
+ GOOGLE_DCHECK(GetOwningArena() == other->GetOwningArena());
+ InternalSwap(other);
+ }
+
+ // implements Message ----------------------------------------------
+
+ inline Reconnect* New() const final {
+ return new Reconnect();
+ }
+
+ Reconnect* New(::PROTOBUF_NAMESPACE_ID::Arena* arena) const final {
+ return CreateMaybeMessage(arena);
+ }
+ using ::PROTOBUF_NAMESPACE_ID::Message::CopyFrom;
+ void CopyFrom(const Reconnect& from);
+ using ::PROTOBUF_NAMESPACE_ID::Message::MergeFrom;
+ void MergeFrom(const Reconnect& from);
+ private:
+ static void MergeImpl(::PROTOBUF_NAMESPACE_ID::Message* to, const ::PROTOBUF_NAMESPACE_ID::Message& from);
+ public:
+ PROTOBUF_ATTRIBUTE_REINITIALIZES void Clear() final;
+ bool IsInitialized() const final;
+
+ size_t ByteSizeLong() const final;
+ const char* _InternalParse(const char* ptr, ::PROTOBUF_NAMESPACE_ID::internal::ParseContext* ctx) final;
+ ::PROTOBUF_NAMESPACE_ID::uint8* _InternalSerialize(
+ ::PROTOBUF_NAMESPACE_ID::uint8* target, ::PROTOBUF_NAMESPACE_ID::io::EpsCopyOutputStream* stream) const final;
+ int GetCachedSize() const final { return _cached_size_.Get(); }
+
+ private:
+ void SharedCtor();
+ void SharedDtor();
+ void SetCachedSize(int size) const final;
+ void InternalSwap(Reconnect* other);
+ friend class ::PROTOBUF_NAMESPACE_ID::internal::AnyMetadata;
+ static ::PROTOBUF_NAMESPACE_ID::StringPiece FullMessageName() {
+ return "flwr.proto.Reconnect";
+ }
+ protected:
+ explicit Reconnect(::PROTOBUF_NAMESPACE_ID::Arena* arena,
+ bool is_message_owned = false);
+ private:
+ static void ArenaDtor(void* object);
+ inline void RegisterArenaDtor(::PROTOBUF_NAMESPACE_ID::Arena* arena);
+ public:
+
+ static const ClassData _class_data_;
+ const ::PROTOBUF_NAMESPACE_ID::Message::ClassData*GetClassData() const final;
+
+ ::PROTOBUF_NAMESPACE_ID::Metadata GetMetadata() const final;
+
+ // nested types ----------------------------------------------------
+
+ // accessors -------------------------------------------------------
+
+ enum : int {
+ kReconnectFieldNumber = 1,
+ };
+ // uint64 reconnect = 1;
+ void clear_reconnect();
+ ::PROTOBUF_NAMESPACE_ID::uint64 reconnect() const;
+ void set_reconnect(::PROTOBUF_NAMESPACE_ID::uint64 value);
+ private:
+ ::PROTOBUF_NAMESPACE_ID::uint64 _internal_reconnect() const;
+ void _internal_set_reconnect(::PROTOBUF_NAMESPACE_ID::uint64 value);
+ public:
+
+ // @@protoc_insertion_point(class_scope:flwr.proto.Reconnect)
+ private:
+ class _Internal;
+
+ template friend class ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper;
+ typedef void InternalArenaConstructable_;
+ typedef void DestructorSkippable_;
+ ::PROTOBUF_NAMESPACE_ID::uint64 reconnect_;
+ mutable ::PROTOBUF_NAMESPACE_ID::internal::CachedSize _cached_size_;
+ friend struct ::TableStruct_flwr_2fproto_2ffleet_2eproto;
+};
+// ===================================================================
+
+
+// ===================================================================
+
+#ifdef __GNUC__
+ #pragma GCC diagnostic push
+ #pragma GCC diagnostic ignored "-Wstrict-aliasing"
+#endif // __GNUC__
+// CreateNodeRequest
+
+// -------------------------------------------------------------------
+
+// CreateNodeResponse
+
+// .flwr.proto.Node node = 1;
+inline bool CreateNodeResponse::_internal_has_node() const {
+ return this != internal_default_instance() && node_ != nullptr;
+}
+inline bool CreateNodeResponse::has_node() const {
+ return _internal_has_node();
+}
+inline const ::flwr::proto::Node& CreateNodeResponse::_internal_node() const {
+ const ::flwr::proto::Node* p = node_;
+ return p != nullptr ? *p : reinterpret_cast(
+ ::flwr::proto::_Node_default_instance_);
+}
+inline const ::flwr::proto::Node& CreateNodeResponse::node() const {
+ // @@protoc_insertion_point(field_get:flwr.proto.CreateNodeResponse.node)
+ return _internal_node();
+}
+inline void CreateNodeResponse::unsafe_arena_set_allocated_node(
+ ::flwr::proto::Node* node) {
+ if (GetArenaForAllocation() == nullptr) {
+ delete reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(node_);
+ }
+ node_ = node;
+ if (node) {
+
+ } else {
+
+ }
+ // @@protoc_insertion_point(field_unsafe_arena_set_allocated:flwr.proto.CreateNodeResponse.node)
+}
+inline ::flwr::proto::Node* CreateNodeResponse::release_node() {
+
+ ::flwr::proto::Node* temp = node_;
+ node_ = nullptr;
+#ifdef PROTOBUF_FORCE_COPY_IN_RELEASE
+ auto* old = reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(temp);
+ temp = ::PROTOBUF_NAMESPACE_ID::internal::DuplicateIfNonNull(temp);
+ if (GetArenaForAllocation() == nullptr) { delete old; }
+#else // PROTOBUF_FORCE_COPY_IN_RELEASE
+ if (GetArenaForAllocation() != nullptr) {
+ temp = ::PROTOBUF_NAMESPACE_ID::internal::DuplicateIfNonNull(temp);
+ }
+#endif // !PROTOBUF_FORCE_COPY_IN_RELEASE
+ return temp;
+}
+inline ::flwr::proto::Node* CreateNodeResponse::unsafe_arena_release_node() {
+ // @@protoc_insertion_point(field_release:flwr.proto.CreateNodeResponse.node)
+
+ ::flwr::proto::Node* temp = node_;
+ node_ = nullptr;
+ return temp;
+}
+inline ::flwr::proto::Node* CreateNodeResponse::_internal_mutable_node() {
+
+ if (node_ == nullptr) {
+ auto* p = CreateMaybeMessage<::flwr::proto::Node>(GetArenaForAllocation());
+ node_ = p;
+ }
+ return node_;
+}
+inline ::flwr::proto::Node* CreateNodeResponse::mutable_node() {
+ ::flwr::proto::Node* _msg = _internal_mutable_node();
+ // @@protoc_insertion_point(field_mutable:flwr.proto.CreateNodeResponse.node)
+ return _msg;
+}
+inline void CreateNodeResponse::set_allocated_node(::flwr::proto::Node* node) {
+ ::PROTOBUF_NAMESPACE_ID::Arena* message_arena = GetArenaForAllocation();
+ if (message_arena == nullptr) {
+ delete reinterpret_cast< ::PROTOBUF_NAMESPACE_ID::MessageLite*>(node_);
+ }
+ if (node) {
+ ::PROTOBUF_NAMESPACE_ID::Arena* submessage_arena =
+ ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper<
+ ::PROTOBUF_NAMESPACE_ID::MessageLite>::GetOwningArena(
+ reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(node));
+ if (message_arena != submessage_arena) {
+ node = ::PROTOBUF_NAMESPACE_ID::internal::GetOwnedMessage(
+ message_arena, node, submessage_arena);
+ }
+
+ } else {
+
+ }
+ node_ = node;
+ // @@protoc_insertion_point(field_set_allocated:flwr.proto.CreateNodeResponse.node)
+}
+
+// -------------------------------------------------------------------
+
+// DeleteNodeRequest
+
+// .flwr.proto.Node node = 1;
+inline bool DeleteNodeRequest::_internal_has_node() const {
+ return this != internal_default_instance() && node_ != nullptr;
+}
+inline bool DeleteNodeRequest::has_node() const {
+ return _internal_has_node();
+}
+inline const ::flwr::proto::Node& DeleteNodeRequest::_internal_node() const {
+ const ::flwr::proto::Node* p = node_;
+ return p != nullptr ? *p : reinterpret_cast(
+ ::flwr::proto::_Node_default_instance_);
+}
+inline const ::flwr::proto::Node& DeleteNodeRequest::node() const {
+ // @@protoc_insertion_point(field_get:flwr.proto.DeleteNodeRequest.node)
+ return _internal_node();
+}
+inline void DeleteNodeRequest::unsafe_arena_set_allocated_node(
+ ::flwr::proto::Node* node) {
+ if (GetArenaForAllocation() == nullptr) {
+ delete reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(node_);
+ }
+ node_ = node;
+ if (node) {
+
+ } else {
+
+ }
+ // @@protoc_insertion_point(field_unsafe_arena_set_allocated:flwr.proto.DeleteNodeRequest.node)
+}
+inline ::flwr::proto::Node* DeleteNodeRequest::release_node() {
+
+ ::flwr::proto::Node* temp = node_;
+ node_ = nullptr;
+#ifdef PROTOBUF_FORCE_COPY_IN_RELEASE
+ auto* old = reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(temp);
+ temp = ::PROTOBUF_NAMESPACE_ID::internal::DuplicateIfNonNull(temp);
+ if (GetArenaForAllocation() == nullptr) { delete old; }
+#else // PROTOBUF_FORCE_COPY_IN_RELEASE
+ if (GetArenaForAllocation() != nullptr) {
+ temp = ::PROTOBUF_NAMESPACE_ID::internal::DuplicateIfNonNull(temp);
+ }
+#endif // !PROTOBUF_FORCE_COPY_IN_RELEASE
+ return temp;
+}
+inline ::flwr::proto::Node* DeleteNodeRequest::unsafe_arena_release_node() {
+ // @@protoc_insertion_point(field_release:flwr.proto.DeleteNodeRequest.node)
+
+ ::flwr::proto::Node* temp = node_;
+ node_ = nullptr;
+ return temp;
+}
+inline ::flwr::proto::Node* DeleteNodeRequest::_internal_mutable_node() {
+
+ if (node_ == nullptr) {
+ auto* p = CreateMaybeMessage<::flwr::proto::Node>(GetArenaForAllocation());
+ node_ = p;
+ }
+ return node_;
+}
+inline ::flwr::proto::Node* DeleteNodeRequest::mutable_node() {
+ ::flwr::proto::Node* _msg = _internal_mutable_node();
+ // @@protoc_insertion_point(field_mutable:flwr.proto.DeleteNodeRequest.node)
+ return _msg;
+}
+inline void DeleteNodeRequest::set_allocated_node(::flwr::proto::Node* node) {
+ ::PROTOBUF_NAMESPACE_ID::Arena* message_arena = GetArenaForAllocation();
+ if (message_arena == nullptr) {
+ delete reinterpret_cast< ::PROTOBUF_NAMESPACE_ID::MessageLite*>(node_);
+ }
+ if (node) {
+ ::PROTOBUF_NAMESPACE_ID::Arena* submessage_arena =
+ ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper<
+ ::PROTOBUF_NAMESPACE_ID::MessageLite>::GetOwningArena(
+ reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(node));
+ if (message_arena != submessage_arena) {
+ node = ::PROTOBUF_NAMESPACE_ID::internal::GetOwnedMessage(
+ message_arena, node, submessage_arena);
+ }
+
+ } else {
+
+ }
+ node_ = node;
+ // @@protoc_insertion_point(field_set_allocated:flwr.proto.DeleteNodeRequest.node)
+}
+
+// -------------------------------------------------------------------
+
+// DeleteNodeResponse
+
+// -------------------------------------------------------------------
+
+// PullTaskInsRequest
+
+// .flwr.proto.Node node = 1;
+inline bool PullTaskInsRequest::_internal_has_node() const {
+ return this != internal_default_instance() && node_ != nullptr;
+}
+inline bool PullTaskInsRequest::has_node() const {
+ return _internal_has_node();
+}
+inline const ::flwr::proto::Node& PullTaskInsRequest::_internal_node() const {
+ const ::flwr::proto::Node* p = node_;
+ return p != nullptr ? *p : reinterpret_cast(
+ ::flwr::proto::_Node_default_instance_);
+}
+inline const ::flwr::proto::Node& PullTaskInsRequest::node() const {
+ // @@protoc_insertion_point(field_get:flwr.proto.PullTaskInsRequest.node)
+ return _internal_node();
+}
+inline void PullTaskInsRequest::unsafe_arena_set_allocated_node(
+ ::flwr::proto::Node* node) {
+ if (GetArenaForAllocation() == nullptr) {
+ delete reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(node_);
+ }
+ node_ = node;
+ if (node) {
+
+ } else {
+
+ }
+ // @@protoc_insertion_point(field_unsafe_arena_set_allocated:flwr.proto.PullTaskInsRequest.node)
+}
+inline ::flwr::proto::Node* PullTaskInsRequest::release_node() {
+
+ ::flwr::proto::Node* temp = node_;
+ node_ = nullptr;
+#ifdef PROTOBUF_FORCE_COPY_IN_RELEASE
+ auto* old = reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(temp);
+ temp = ::PROTOBUF_NAMESPACE_ID::internal::DuplicateIfNonNull(temp);
+ if (GetArenaForAllocation() == nullptr) { delete old; }
+#else // PROTOBUF_FORCE_COPY_IN_RELEASE
+ if (GetArenaForAllocation() != nullptr) {
+ temp = ::PROTOBUF_NAMESPACE_ID::internal::DuplicateIfNonNull(temp);
+ }
+#endif // !PROTOBUF_FORCE_COPY_IN_RELEASE
+ return temp;
+}
+inline ::flwr::proto::Node* PullTaskInsRequest::unsafe_arena_release_node() {
+ // @@protoc_insertion_point(field_release:flwr.proto.PullTaskInsRequest.node)
+
+ ::flwr::proto::Node* temp = node_;
+ node_ = nullptr;
+ return temp;
+}
+inline ::flwr::proto::Node* PullTaskInsRequest::_internal_mutable_node() {
+
+ if (node_ == nullptr) {
+ auto* p = CreateMaybeMessage<::flwr::proto::Node>(GetArenaForAllocation());
+ node_ = p;
+ }
+ return node_;
+}
+inline ::flwr::proto::Node* PullTaskInsRequest::mutable_node() {
+ ::flwr::proto::Node* _msg = _internal_mutable_node();
+ // @@protoc_insertion_point(field_mutable:flwr.proto.PullTaskInsRequest.node)
+ return _msg;
+}
+inline void PullTaskInsRequest::set_allocated_node(::flwr::proto::Node* node) {
+ ::PROTOBUF_NAMESPACE_ID::Arena* message_arena = GetArenaForAllocation();
+ if (message_arena == nullptr) {
+ delete reinterpret_cast< ::PROTOBUF_NAMESPACE_ID::MessageLite*>(node_);
+ }
+ if (node) {
+ ::PROTOBUF_NAMESPACE_ID::Arena* submessage_arena =
+ ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper<
+ ::PROTOBUF_NAMESPACE_ID::MessageLite>::GetOwningArena(
+ reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(node));
+ if (message_arena != submessage_arena) {
+ node = ::PROTOBUF_NAMESPACE_ID::internal::GetOwnedMessage(
+ message_arena, node, submessage_arena);
+ }
+
+ } else {
+
+ }
+ node_ = node;
+ // @@protoc_insertion_point(field_set_allocated:flwr.proto.PullTaskInsRequest.node)
+}
+
+// repeated string task_ids = 2;
+inline int PullTaskInsRequest::_internal_task_ids_size() const {
+ return task_ids_.size();
+}
+inline int PullTaskInsRequest::task_ids_size() const {
+ return _internal_task_ids_size();
+}
+inline void PullTaskInsRequest::clear_task_ids() {
+ task_ids_.Clear();
+}
+inline std::string* PullTaskInsRequest::add_task_ids() {
+ std::string* _s = _internal_add_task_ids();
+ // @@protoc_insertion_point(field_add_mutable:flwr.proto.PullTaskInsRequest.task_ids)
+ return _s;
+}
+inline const std::string& PullTaskInsRequest::_internal_task_ids(int index) const {
+ return task_ids_.Get(index);
+}
+inline const std::string& PullTaskInsRequest::task_ids(int index) const {
+ // @@protoc_insertion_point(field_get:flwr.proto.PullTaskInsRequest.task_ids)
+ return _internal_task_ids(index);
+}
+inline std::string* PullTaskInsRequest::mutable_task_ids(int index) {
+ // @@protoc_insertion_point(field_mutable:flwr.proto.PullTaskInsRequest.task_ids)
+ return task_ids_.Mutable(index);
+}
+inline void PullTaskInsRequest::set_task_ids(int index, const std::string& value) {
+ task_ids_.Mutable(index)->assign(value);
+ // @@protoc_insertion_point(field_set:flwr.proto.PullTaskInsRequest.task_ids)
+}
+inline void PullTaskInsRequest::set_task_ids(int index, std::string&& value) {
+ task_ids_.Mutable(index)->assign(std::move(value));
+ // @@protoc_insertion_point(field_set:flwr.proto.PullTaskInsRequest.task_ids)
+}
+inline void PullTaskInsRequest::set_task_ids(int index, const char* value) {
+ GOOGLE_DCHECK(value != nullptr);
+ task_ids_.Mutable(index)->assign(value);
+ // @@protoc_insertion_point(field_set_char:flwr.proto.PullTaskInsRequest.task_ids)
+}
+inline void PullTaskInsRequest::set_task_ids(int index, const char* value, size_t size) {
+ task_ids_.Mutable(index)->assign(
+ reinterpret_cast(value), size);
+ // @@protoc_insertion_point(field_set_pointer:flwr.proto.PullTaskInsRequest.task_ids)
+}
+inline std::string* PullTaskInsRequest::_internal_add_task_ids() {
+ return task_ids_.Add();
+}
+inline void PullTaskInsRequest::add_task_ids(const std::string& value) {
+ task_ids_.Add()->assign(value);
+ // @@protoc_insertion_point(field_add:flwr.proto.PullTaskInsRequest.task_ids)
+}
+inline void PullTaskInsRequest::add_task_ids(std::string&& value) {
+ task_ids_.Add(std::move(value));
+ // @@protoc_insertion_point(field_add:flwr.proto.PullTaskInsRequest.task_ids)
+}
+inline void PullTaskInsRequest::add_task_ids(const char* value) {
+ GOOGLE_DCHECK(value != nullptr);
+ task_ids_.Add()->assign(value);
+ // @@protoc_insertion_point(field_add_char:flwr.proto.PullTaskInsRequest.task_ids)
+}
+inline void PullTaskInsRequest::add_task_ids(const char* value, size_t size) {
+ task_ids_.Add()->assign(reinterpret_cast(value), size);
+ // @@protoc_insertion_point(field_add_pointer:flwr.proto.PullTaskInsRequest.task_ids)
+}
+inline const ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField&
+PullTaskInsRequest::task_ids() const {
+ // @@protoc_insertion_point(field_list:flwr.proto.PullTaskInsRequest.task_ids)
+ return task_ids_;
+}
+inline ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField*
+PullTaskInsRequest::mutable_task_ids() {
+ // @@protoc_insertion_point(field_mutable_list:flwr.proto.PullTaskInsRequest.task_ids)
+ return &task_ids_;
+}
+
+// -------------------------------------------------------------------
+
+// PullTaskInsResponse
+
+// .flwr.proto.Reconnect reconnect = 1;
+inline bool PullTaskInsResponse::_internal_has_reconnect() const {
+ return this != internal_default_instance() && reconnect_ != nullptr;
+}
+inline bool PullTaskInsResponse::has_reconnect() const {
+ return _internal_has_reconnect();
+}
+inline void PullTaskInsResponse::clear_reconnect() {
+ if (GetArenaForAllocation() == nullptr && reconnect_ != nullptr) {
+ delete reconnect_;
+ }
+ reconnect_ = nullptr;
+}
+inline const ::flwr::proto::Reconnect& PullTaskInsResponse::_internal_reconnect() const {
+ const ::flwr::proto::Reconnect* p = reconnect_;
+ return p != nullptr ? *p : reinterpret_cast(
+ ::flwr::proto::_Reconnect_default_instance_);
+}
+inline const ::flwr::proto::Reconnect& PullTaskInsResponse::reconnect() const {
+ // @@protoc_insertion_point(field_get:flwr.proto.PullTaskInsResponse.reconnect)
+ return _internal_reconnect();
+}
+inline void PullTaskInsResponse::unsafe_arena_set_allocated_reconnect(
+ ::flwr::proto::Reconnect* reconnect) {
+ if (GetArenaForAllocation() == nullptr) {
+ delete reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(reconnect_);
+ }
+ reconnect_ = reconnect;
+ if (reconnect) {
+
+ } else {
+
+ }
+ // @@protoc_insertion_point(field_unsafe_arena_set_allocated:flwr.proto.PullTaskInsResponse.reconnect)
+}
+inline ::flwr::proto::Reconnect* PullTaskInsResponse::release_reconnect() {
+
+ ::flwr::proto::Reconnect* temp = reconnect_;
+ reconnect_ = nullptr;
+#ifdef PROTOBUF_FORCE_COPY_IN_RELEASE
+ auto* old = reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(temp);
+ temp = ::PROTOBUF_NAMESPACE_ID::internal::DuplicateIfNonNull(temp);
+ if (GetArenaForAllocation() == nullptr) { delete old; }
+#else // PROTOBUF_FORCE_COPY_IN_RELEASE
+ if (GetArenaForAllocation() != nullptr) {
+ temp = ::PROTOBUF_NAMESPACE_ID::internal::DuplicateIfNonNull(temp);
+ }
+#endif // !PROTOBUF_FORCE_COPY_IN_RELEASE
+ return temp;
+}
+inline ::flwr::proto::Reconnect* PullTaskInsResponse::unsafe_arena_release_reconnect() {
+ // @@protoc_insertion_point(field_release:flwr.proto.PullTaskInsResponse.reconnect)
+
+ ::flwr::proto::Reconnect* temp = reconnect_;
+ reconnect_ = nullptr;
+ return temp;
+}
+inline ::flwr::proto::Reconnect* PullTaskInsResponse::_internal_mutable_reconnect() {
+
+ if (reconnect_ == nullptr) {
+ auto* p = CreateMaybeMessage<::flwr::proto::Reconnect>(GetArenaForAllocation());
+ reconnect_ = p;
+ }
+ return reconnect_;
+}
+inline ::flwr::proto::Reconnect* PullTaskInsResponse::mutable_reconnect() {
+ ::flwr::proto::Reconnect* _msg = _internal_mutable_reconnect();
+ // @@protoc_insertion_point(field_mutable:flwr.proto.PullTaskInsResponse.reconnect)
+ return _msg;
+}
+inline void PullTaskInsResponse::set_allocated_reconnect(::flwr::proto::Reconnect* reconnect) {
+ ::PROTOBUF_NAMESPACE_ID::Arena* message_arena = GetArenaForAllocation();
+ if (message_arena == nullptr) {
+ delete reconnect_;
+ }
+ if (reconnect) {
+ ::PROTOBUF_NAMESPACE_ID::Arena* submessage_arena =
+ ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper<::flwr::proto::Reconnect>::GetOwningArena(reconnect);
+ if (message_arena != submessage_arena) {
+ reconnect = ::PROTOBUF_NAMESPACE_ID::internal::GetOwnedMessage(
+ message_arena, reconnect, submessage_arena);
+ }
+
+ } else {
+
+ }
+ reconnect_ = reconnect;
+ // @@protoc_insertion_point(field_set_allocated:flwr.proto.PullTaskInsResponse.reconnect)
+}
+
+// repeated .flwr.proto.TaskIns task_ins_list = 2;
+inline int PullTaskInsResponse::_internal_task_ins_list_size() const {
+ return task_ins_list_.size();
+}
+inline int PullTaskInsResponse::task_ins_list_size() const {
+ return _internal_task_ins_list_size();
+}
+inline ::flwr::proto::TaskIns* PullTaskInsResponse::mutable_task_ins_list(int index) {
+ // @@protoc_insertion_point(field_mutable:flwr.proto.PullTaskInsResponse.task_ins_list)
+ return task_ins_list_.Mutable(index);
+}
+inline ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField< ::flwr::proto::TaskIns >*
+PullTaskInsResponse::mutable_task_ins_list() {
+ // @@protoc_insertion_point(field_mutable_list:flwr.proto.PullTaskInsResponse.task_ins_list)
+ return &task_ins_list_;
+}
+inline const ::flwr::proto::TaskIns& PullTaskInsResponse::_internal_task_ins_list(int index) const {
+ return task_ins_list_.Get(index);
+}
+inline const ::flwr::proto::TaskIns& PullTaskInsResponse::task_ins_list(int index) const {
+ // @@protoc_insertion_point(field_get:flwr.proto.PullTaskInsResponse.task_ins_list)
+ return _internal_task_ins_list(index);
+}
+inline ::flwr::proto::TaskIns* PullTaskInsResponse::_internal_add_task_ins_list() {
+ return task_ins_list_.Add();
+}
+inline ::flwr::proto::TaskIns* PullTaskInsResponse::add_task_ins_list() {
+ ::flwr::proto::TaskIns* _add = _internal_add_task_ins_list();
+ // @@protoc_insertion_point(field_add:flwr.proto.PullTaskInsResponse.task_ins_list)
+ return _add;
+}
+inline const ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField< ::flwr::proto::TaskIns >&
+PullTaskInsResponse::task_ins_list() const {
+ // @@protoc_insertion_point(field_list:flwr.proto.PullTaskInsResponse.task_ins_list)
+ return task_ins_list_;
+}
+
+// -------------------------------------------------------------------
+
+// PushTaskResRequest
+
+// repeated .flwr.proto.TaskRes task_res_list = 1;
+inline int PushTaskResRequest::_internal_task_res_list_size() const {
+ return task_res_list_.size();
+}
+inline int PushTaskResRequest::task_res_list_size() const {
+ return _internal_task_res_list_size();
+}
+inline ::flwr::proto::TaskRes* PushTaskResRequest::mutable_task_res_list(int index) {
+ // @@protoc_insertion_point(field_mutable:flwr.proto.PushTaskResRequest.task_res_list)
+ return task_res_list_.Mutable(index);
+}
+inline ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField< ::flwr::proto::TaskRes >*
+PushTaskResRequest::mutable_task_res_list() {
+ // @@protoc_insertion_point(field_mutable_list:flwr.proto.PushTaskResRequest.task_res_list)
+ return &task_res_list_;
+}
+inline const ::flwr::proto::TaskRes& PushTaskResRequest::_internal_task_res_list(int index) const {
+ return task_res_list_.Get(index);
+}
+inline const ::flwr::proto::TaskRes& PushTaskResRequest::task_res_list(int index) const {
+ // @@protoc_insertion_point(field_get:flwr.proto.PushTaskResRequest.task_res_list)
+ return _internal_task_res_list(index);
+}
+inline ::flwr::proto::TaskRes* PushTaskResRequest::_internal_add_task_res_list() {
+ return task_res_list_.Add();
+}
+inline ::flwr::proto::TaskRes* PushTaskResRequest::add_task_res_list() {
+ ::flwr::proto::TaskRes* _add = _internal_add_task_res_list();
+ // @@protoc_insertion_point(field_add:flwr.proto.PushTaskResRequest.task_res_list)
+ return _add;
+}
+inline const ::PROTOBUF_NAMESPACE_ID::RepeatedPtrField< ::flwr::proto::TaskRes >&
+PushTaskResRequest::task_res_list() const {
+ // @@protoc_insertion_point(field_list:flwr.proto.PushTaskResRequest.task_res_list)
+ return task_res_list_;
+}
+
+// -------------------------------------------------------------------
+
+// -------------------------------------------------------------------
+
+// PushTaskResResponse
+
+// .flwr.proto.Reconnect reconnect = 1;
+inline bool PushTaskResResponse::_internal_has_reconnect() const {
+ return this != internal_default_instance() && reconnect_ != nullptr;
+}
+inline bool PushTaskResResponse::has_reconnect() const {
+ return _internal_has_reconnect();
+}
+inline void PushTaskResResponse::clear_reconnect() {
+ if (GetArenaForAllocation() == nullptr && reconnect_ != nullptr) {
+ delete reconnect_;
+ }
+ reconnect_ = nullptr;
+}
+inline const ::flwr::proto::Reconnect& PushTaskResResponse::_internal_reconnect() const {
+ const ::flwr::proto::Reconnect* p = reconnect_;
+ return p != nullptr ? *p : reinterpret_cast(
+ ::flwr::proto::_Reconnect_default_instance_);
+}
+inline const ::flwr::proto::Reconnect& PushTaskResResponse::reconnect() const {
+ // @@protoc_insertion_point(field_get:flwr.proto.PushTaskResResponse.reconnect)
+ return _internal_reconnect();
+}
+inline void PushTaskResResponse::unsafe_arena_set_allocated_reconnect(
+ ::flwr::proto::Reconnect* reconnect) {
+ if (GetArenaForAllocation() == nullptr) {
+ delete reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(reconnect_);
+ }
+ reconnect_ = reconnect;
+ if (reconnect) {
+
+ } else {
+
+ }
+ // @@protoc_insertion_point(field_unsafe_arena_set_allocated:flwr.proto.PushTaskResResponse.reconnect)
+}
+inline ::flwr::proto::Reconnect* PushTaskResResponse::release_reconnect() {
+
+ ::flwr::proto::Reconnect* temp = reconnect_;
+ reconnect_ = nullptr;
+#ifdef PROTOBUF_FORCE_COPY_IN_RELEASE
+ auto* old = reinterpret_cast<::PROTOBUF_NAMESPACE_ID::MessageLite*>(temp);
+ temp = ::PROTOBUF_NAMESPACE_ID::internal::DuplicateIfNonNull(temp);
+ if (GetArenaForAllocation() == nullptr) { delete old; }
+#else // PROTOBUF_FORCE_COPY_IN_RELEASE
+ if (GetArenaForAllocation() != nullptr) {
+ temp = ::PROTOBUF_NAMESPACE_ID::internal::DuplicateIfNonNull(temp);
+ }
+#endif // !PROTOBUF_FORCE_COPY_IN_RELEASE
+ return temp;
+}
+inline ::flwr::proto::Reconnect* PushTaskResResponse::unsafe_arena_release_reconnect() {
+ // @@protoc_insertion_point(field_release:flwr.proto.PushTaskResResponse.reconnect)
+
+ ::flwr::proto::Reconnect* temp = reconnect_;
+ reconnect_ = nullptr;
+ return temp;
+}
+inline ::flwr::proto::Reconnect* PushTaskResResponse::_internal_mutable_reconnect() {
+
+ if (reconnect_ == nullptr) {
+ auto* p = CreateMaybeMessage<::flwr::proto::Reconnect>(GetArenaForAllocation());
+ reconnect_ = p;
+ }
+ return reconnect_;
+}
+inline ::flwr::proto::Reconnect* PushTaskResResponse::mutable_reconnect() {
+ ::flwr::proto::Reconnect* _msg = _internal_mutable_reconnect();
+ // @@protoc_insertion_point(field_mutable:flwr.proto.PushTaskResResponse.reconnect)
+ return _msg;
+}
+inline void PushTaskResResponse::set_allocated_reconnect(::flwr::proto::Reconnect* reconnect) {
+ ::PROTOBUF_NAMESPACE_ID::Arena* message_arena = GetArenaForAllocation();
+ if (message_arena == nullptr) {
+ delete reconnect_;
+ }
+ if (reconnect) {
+ ::PROTOBUF_NAMESPACE_ID::Arena* submessage_arena =
+ ::PROTOBUF_NAMESPACE_ID::Arena::InternalHelper<::flwr::proto::Reconnect>::GetOwningArena(reconnect);
+ if (message_arena != submessage_arena) {
+ reconnect = ::PROTOBUF_NAMESPACE_ID::internal::GetOwnedMessage(
+ message_arena, reconnect, submessage_arena);
+ }
+
+ } else {
+
+ }
+ reconnect_ = reconnect;
+ // @@protoc_insertion_point(field_set_allocated:flwr.proto.PushTaskResResponse.reconnect)
+}
+
+// map results = 2;
+inline int PushTaskResResponse::_internal_results_size() const {
+ return results_.size();
+}
+inline int PushTaskResResponse::results_size() const {
+ return _internal_results_size();
+}
+inline void PushTaskResResponse::clear_results() {
+ results_.Clear();
+}
+inline const ::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >&
+PushTaskResResponse::_internal_results() const {
+ return results_.GetMap();
+}
+inline const ::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >&
+PushTaskResResponse::results() const {
+ // @@protoc_insertion_point(field_map:flwr.proto.PushTaskResResponse.results)
+ return _internal_results();
+}
+inline ::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >*
+PushTaskResResponse::_internal_mutable_results() {
+ return results_.MutableMap();
+}
+inline ::PROTOBUF_NAMESPACE_ID::Map< std::string, ::PROTOBUF_NAMESPACE_ID::uint32 >*
+PushTaskResResponse::mutable_results() {
+ // @@protoc_insertion_point(field_mutable_map:flwr.proto.PushTaskResResponse.results)
+ return _internal_mutable_results();
+}
+
+// -------------------------------------------------------------------
+
+// Reconnect
+
+// uint64 reconnect = 1;
+inline void Reconnect::clear_reconnect() {
+ reconnect_ = uint64_t{0u};
+}
+inline ::PROTOBUF_NAMESPACE_ID::uint64 Reconnect::_internal_reconnect() const {
+ return reconnect_;
+}
+inline ::PROTOBUF_NAMESPACE_ID::uint64 Reconnect::reconnect() const {
+ // @@protoc_insertion_point(field_get:flwr.proto.Reconnect.reconnect)
+ return _internal_reconnect();
+}
+inline void Reconnect::_internal_set_reconnect(::PROTOBUF_NAMESPACE_ID::uint64 value) {
+
+ reconnect_ = value;
+}
+inline void Reconnect::set_reconnect(::PROTOBUF_NAMESPACE_ID::uint64 value) {
+ _internal_set_reconnect(value);
+ // @@protoc_insertion_point(field_set:flwr.proto.Reconnect.reconnect)
+}
+
+#ifdef __GNUC__
+ #pragma GCC diagnostic pop
+#endif // __GNUC__
+// -------------------------------------------------------------------
+
+// -------------------------------------------------------------------
+
+// -------------------------------------------------------------------
+
+// -------------------------------------------------------------------
+
+// -------------------------------------------------------------------
+
+// -------------------------------------------------------------------
+
+// -------------------------------------------------------------------
+
+// -------------------------------------------------------------------
+
+// -------------------------------------------------------------------
+
+
+// @@protoc_insertion_point(namespace_scope)
+
+} // namespace proto
+} // namespace flwr
+
+// @@protoc_insertion_point(global_scope)
+
+#include
+#endif // GOOGLE_PROTOBUF_INCLUDED_GOOGLE_PROTOBUF_INCLUDED_flwr_2fproto_2ffleet_2eproto
diff --git a/src/cc/flwr/include/flwr/proto/node.grpc.pb.cc b/src/cc/flwr/include/flwr/proto/node.grpc.pb.cc
new file mode 100644
index 000000000000..9bb46c7e16ca
--- /dev/null
+++ b/src/cc/flwr/include/flwr/proto/node.grpc.pb.cc
@@ -0,0 +1,27 @@
+// Generated by the gRPC C++ plugin.
+// If you make any local change, they will be lost.
+// source: flwr/proto/node.proto
+
+#include "flwr/proto/node.pb.h"
+#include "flwr/proto/node.grpc.pb.h"
+
+#include
+#include
+#include