Shared path prediction improves federated learning efficiency

CTP-FL: Common-Trajectory Gradient Prediction for Federated Learning

Machine LearningArtificial Intelligence

Summary

Federated learning helps train AI models across many devices without sharing raw data. Usually, devices train their models independently then share updates, which can be hard to combine if devices have different kinds of data. The authors propose a new method where all devices predict a common path for updating the model and send gradient information based on that path. This makes the combined updates more reliable and efficient without extra communication. Their math shows this approach balances the benefits and drawbacks of looking ahead along this shared path.

What this means in practice

  • For mobile app developers: Improve efficiency and accuracy of training AI models on user devices with diverse local data while minimizing communication costs.
  • For edge device engineers: Design systems that better coordinate updates from multiple edge devices by synchronizing gradient computations along a common update trajectory.

Authors

Junkang Liu

Abstract

Communication-efficient federated optimization commonly spends several gradient evaluations between server updates. Existing local-update methods use this computation to advance an independent model on each client. Under heterogeneous data, however, these models evaluate gradients at different locations, making the aggregated update difficult to interpret as a gradient of the global objective. We study an alternative use of the same computation budget: \emph{evaluate the global objective along a shared, predicted path}. We propose Common-Trajectory Predictive Federated Learning (\texttt{CTP-FL}). At each round, all clients construct the same sequence of query points from the current global model and the previous aggregated direction, evaluate $K$ stochastic gradients along this sequence, and upload their average. The server then performs a single global update. Thus, \texttt{CTP-FL} uses $K$ mini-batch gradients per client and one model-sized vector in each communication direction, matching the per-round computation and communication of full-participation FedAvg-M. Shared query points make the aggregated direction an unbiased estimator of the average \emph{global} gradient along the predicted path. The remaining discrepancy from the gradient at the current model is controlled by the path length, without assuming bounded client-gradient dissimilarity or bounded gradients. For smooth non-convex objectives, we establish an $\mathcal{O}\!\left( \sqrt{LΔσ^2/(NKR)}+LΔ/R \right)$ average-stationarity bound under full participation. The analysis isolates a testable trade-off: extending the prediction path provides more forward-looking gradient information but increases its displacement bias.