From b6544c9e82366315f264894a95bdecf395e603a3 Mon Sep 17 00:00:00 2001 From: "Shekhawat, Nisha" Date: Sun, 18 May 2025 19:29:48 -0700 Subject: [PATCH 1/5] fix_torch_mu_value Signed-off-by: Shekhawat, Nisha --- ...Prox_PyTorch_MNIST_Workflow_Tutorial.ipynb | 733 +++++++++++------- openfl/utilities/optimizers/torch/fedprox.py | 23 +- 2 files changed, 471 insertions(+), 285 deletions(-) diff --git a/openfl-tutorials/experimental/workflow/403_Federated_FedProx_PyTorch_MNIST_Workflow_Tutorial.ipynb b/openfl-tutorials/experimental/workflow/403_Federated_FedProx_PyTorch_MNIST_Workflow_Tutorial.ipynb index 349832dc89..9511e63b06 100644 --- a/openfl-tutorials/experimental/workflow/403_Federated_FedProx_PyTorch_MNIST_Workflow_Tutorial.ipynb +++ b/openfl-tutorials/experimental/workflow/403_Federated_FedProx_PyTorch_MNIST_Workflow_Tutorial.ipynb @@ -47,7 +47,7 @@ }, { "cell_type": "code", - "execution_count": 21, + "execution_count": null, "metadata": {}, "outputs": [], "source": [ @@ -80,7 +80,7 @@ }, { "cell_type": "code", - "execution_count": 22, + "execution_count": null, "metadata": {}, "outputs": [], "source": [ @@ -126,7 +126,7 @@ }, { "cell_type": "code", - "execution_count": 23, + "execution_count": null, "metadata": {}, "outputs": [], "source": [ @@ -177,7 +177,7 @@ }, { "cell_type": "code", - "execution_count": 24, + "execution_count": null, "metadata": {}, "outputs": [], "source": [ @@ -210,22 +210,9 @@ }, { "cell_type": "code", - "execution_count": 25, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Aggregator step \"start\" registered\n", - "Collaborator step \"aggregated_model_validation\" registered\n", - "Collaborator step \"train\" registered\n", - "Collaborator step \"local_model_validation\" registered\n", - "Aggregator step \"join\" registered\n", - "Aggregator step \"end\" registered\n" - ] - } - ], + "outputs": [], "source": [ "class FederatedFlow(FLSpec):\n", " def __init__(self, model=None, optimizer=None, rounds=10, **kwargs):\n", @@ -273,13 +260,18 @@ "\n", " self.model.train()\n", " self.optimizer = get_optimizer(self.model)\n", + " \n", + " # Set old weights ONCE at the beginning of training\n", + " # This sets the reference weights to the global model weights \n", + " # received from the aggregator, implementing FedProx correctly\n", + " self.optimizer.set_old_weights([p.clone().detach() for p in self.model.parameters()])\n", + " \n", " for batch_idx, (data, target) in enumerate(self.train_loader):\n", " self.optimizer.zero_grad()\n", " output = self.model(data)\n", " loss = F.cross_entropy(output, target)\n", " loss.backward()\n", " \n", - " self.optimizer.set_old_weights([p.clone().detach() for p in self.model.parameters()])\n", " self.optimizer.step()\n", "\n", " if (len(data) * batch_idx) / len(self.train_loader.dataset) >= log_threshold:\n", @@ -335,277 +327,456 @@ }, { "cell_type": "code", - "execution_count": 26, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "\n", - "Calling start\n", - "\u001b[94mPerforming initialization for model\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator0, model: 140162497619616\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 4.6833, Accuracy: 171/2500 (7%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 1.889274\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 1.279191\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.994200\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator0\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.7548, Accuracy: 1929/2500 (77%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator0, Accuracy: 0.7716000080108643\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator1, model: 140158910463952\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 4.7259, Accuracy: 173/2500 (7%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 1.675623\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 1.068585\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.687561\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator1\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.6366, Accuracy: 2004/2500 (80%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator1, Accuracy: 0.8015999794006348\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator2, model: 140162497661872\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 4.6549, Accuracy: 215/2500 (9%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 1.879489\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 1.325507\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.968176\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator2\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.7462, Accuracy: 1901/2500 (76%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator2, Accuracy: 0.7603999972343445\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator3, model: 140162498346528\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 4.7129, Accuracy: 193/2500 (8%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 1.720635\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 1.061211\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.762026\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator3\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.6378, Accuracy: 1992/2500 (80%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator3, Accuracy: 0.7968000173568726\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling join\n", - "\u001b[94mAverage aggregated model accuracy = 0.07520000264048576\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mAverage training loss = 0.8529909627063148\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mAverage local model validation values = 0.782600000500679\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator0, model: 140158910552480\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.6740, Accuracy: 1996/2500 (80%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 0.974921\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 0.633429\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.591566\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator0\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.3951, Accuracy: 2214/2500 (89%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator0, Accuracy: 0.8855999708175659\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator1, model: 140162497608672\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.6877, Accuracy: 1981/2500 (79%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 0.824028\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 0.515538\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.410188\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator1\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.4668, Accuracy: 2166/2500 (87%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator1, Accuracy: 0.8664000034332275\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator2, model: 140162498107328\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.6919, Accuracy: 1981/2500 (79%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 1.025180\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 0.616896\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.483282\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator2\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.4406, Accuracy: 2163/2500 (87%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator2, Accuracy: 0.8651999831199646\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator3, model: 140162498345664\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.6698, Accuracy: 2000/2500 (80%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 0.725868\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 0.450241\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.388554\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator3\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.4106, Accuracy: 2211/2500 (88%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator3, Accuracy: 0.8844000101089478\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling join\n", - "\u001b[94mAverage aggregated model accuracy = 0.7958000004291534\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mAverage training loss = 0.4683974838455107\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mAverage local model validation values = 0.8753999918699265\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator0, model: 140162648091376\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.3590, Accuracy: 2230/2500 (89%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 0.406638\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 0.313662\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.326520\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator0\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.2096, Accuracy: 2338/2500 (94%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator0, Accuracy: 0.9351999759674072\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator1, model: 140162646717344\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.3773, Accuracy: 2228/2500 (89%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 0.392126\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 0.228912\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.200197\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator1\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.2601, Accuracy: 2317/2500 (93%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator1, Accuracy: 0.926800012588501\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator2, model: 140162498503728\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.3683, Accuracy: 2240/2500 (90%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 0.583415\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 0.407979\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.299050\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator2\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.2664, Accuracy: 2305/2500 (92%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator2, Accuracy: 0.921999990940094\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling aggregated_model_validation\n", - "\u001b[94mPerforming aggregated model validation for collaborator collaborator3, model: 140162497621488\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.3626, Accuracy: 2226/2500 (89%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling train\n", - "\u001b[94mTrain Epoch: [4096/15000 (27%)]\tLoss: 0.371595\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [8192/15000 (53%)]\tLoss: 0.234668\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mTrain Epoch: [11264/15000 (73%)]\tLoss: 0.177007\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling local_model_validation\n", - "\u001b[94mPerforming local model validation for collaborator collaborator3\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94m\n", - "Test set: Avg. loss: 0.2615, Accuracy: 2305/2500 (92%)\n", - "\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mDone with local model validation for collaborator collaborator3, Accuracy: 0.921999990940094\u001b[0m\u001b[94m\n", - "\u001b[0mShould transfer from local_model_validation to join\n", - "\n", - "Calling join\n", - "\u001b[94mAverage aggregated model accuracy = 0.8924000114202499\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mAverage training loss = 0.250693485351252\u001b[0m\u001b[94m\n", - "\u001b[0m\u001b[94mAverage local model validation values = 0.926499992609024\u001b[0m\u001b[94m\n", - "\u001b[0m\n", - "Calling end\n", - "\u001b[94mFlow ended\u001b[0m\u001b[94m\n", - "\u001b[0m" - ] - } - ], + "outputs": [], "source": [ "model = Net()\n", "flflow = FederatedFlow(model, get_optimizer(model), rounds=3, checkpoint=False)\n", "flflow.runtime = local_runtime\n", "flflow.run()" ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Model Comparison with Different Mu Values\n", + "Now let's check if trained models with the same seeding but different mu values produce the same results. We'll train multiple models with different mu values and then compare their weights." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "import copy\n", + "import random\n", + "import numpy as np\n", + "\n", + "# Function to set seed for reproducibility\n", + "def set_seed(seed):\n", + " \"\"\"Set seed for reproducibility\"\"\"\n", + " torch.manual_seed(seed)\n", + " torch.cuda.manual_seed_all(seed)\n", + " np.random.seed(seed)\n", + " random.seed(seed)\n", + " torch.backends.cudnn.deterministic = True\n", + " torch.backends.cudnn.benchmark = False\n", + "\n", + "# Function to get optimizer with different mu values\n", + "def get_optimizer_with_mu(model, mu_value):\n", + " \"\"\"Get FedProxAdam optimizer with specific mu value\"\"\"\n", + " return FedProxAdam(model.parameters(), lr=1e-3, mu=mu_value)\n", + "\n", + "# Function to compare model weights\n", + "def compare_models(models, model_names):\n", + " \"\"\"Compare weights between different models\"\"\"\n", + " print(\"Comparing model weights:\")\n", + " \n", + " # Compare weights between each pair of models\n", + " for i in range(len(models)):\n", + " for j in range(i+1, len(models)):\n", + " model1 = models[i]\n", + " model2 = models[j]\n", + " name1 = model_names[i]\n", + " name2 = model_names[j]\n", + " \n", + " print(f\"\\nComparing {name1} and {name2}:\")\n", + " \n", + " # Compare each layer's weights\n", + " all_equal = True\n", + " max_diff = 0.0\n", + " \n", + " for (p1, p2) in zip(model1.parameters(), model2.parameters()):\n", + " # Check if parameters are equal\n", + " if not torch.allclose(p1, p2, atol=1e-5):\n", + " all_equal = False\n", + " # Calculate maximum difference\n", + " diff = torch.max(torch.abs(p1 - p2)).item()\n", + " max_diff = max(max_diff, diff)\n", + " \n", + " if all_equal:\n", + " print(f\"Models {name1} and {name2} have identical weights\")\n", + " else:\n", + " print(f\"Models {name1} and {name2} have different weights\")\n", + " print(f\"Maximum difference in weights: {max_diff:.6f}\")\n", + "\n", + "# Function to train a model with specific mu value\n", + "def train_model(model, mu_value, seed, epochs=1):\n", + " \"\"\"Train a model with specific mu value and seed\"\"\"\n", + " set_seed(seed) # Set seed for reproducibility\n", + " \n", + " # Create a copy of the model\n", + " model_copy = copy.deepcopy(model)\n", + " \n", + " # Get optimizer with specific mu value\n", + " optimizer = get_optimizer_with_mu(model_copy, mu_value)\n", + " \n", + " # Training loop\n", + " model_copy.train()\n", + " \n", + " # Get a small dataset for quick testing\n", + " train_images, train_labels = mnist_train.train_data[:1000], np.array(mnist_train.train_labels[:1000])\n", + " train_images = torch.from_numpy(np.expand_dims(train_images, axis=1)).float()\n", + " train_labels = one_hot(train_labels, 10)\n", + " \n", + " train_dataset = CustomDataset(train_images, train_labels)\n", + " train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=32, shuffle=True)\n", + " \n", + " # CRITICAL FIX: Set the old weights ONCE at the beginning, before training\n", + " # This properly implements FedProx by setting the reference weights\n", + " # to the initial model weights (simulating the global model)\n", + " optimizer.set_old_weights([p.clone().detach() for p in model_copy.parameters()])\n", + " \n", + " for epoch in range(epochs):\n", + " running_loss = 0.0\n", + " for batch_idx, (data, target) in enumerate(train_loader):\n", + " optimizer.zero_grad()\n", + " output = model_copy(data)\n", + " loss = F.cross_entropy(output, target)\n", + " loss.backward()\n", + " \n", + " # REMOVED: Don't call set_old_weights here - this was causing the issue\n", + " optimizer.step()\n", + " \n", + " running_loss += loss.item()\n", + " \n", + " print(f\"Epoch {epoch+1}, Mu={mu_value}, Loss: {running_loss/len(train_loader):.6f}\")\n", + " \n", + " return model_copy" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# Define mu values to test\n", + "mu_values = [0.0, 0.01, 0.1, 0.5]\n", + "seed = 42 # Fixed seed for reproducibility\n", + "models = []\n", + "model_names = []\n", + "\n", + "print(\"Training models with different mu values but same seed...\")\n", + "\n", + "# Train models with different mu values\n", + "for mu in mu_values:\n", + " model_name = f\"Model_mu{mu}\"\n", + " print(f\"\\nTraining {model_name} with mu={mu}...\")\n", + " \n", + " # Initialize a new model\n", + " base_model = Net()\n", + " \n", + " # Train the model with current mu value\n", + " trained_model = train_model(base_model, mu, seed, epochs=2)\n", + " \n", + " # Save the model and name\n", + " models.append(trained_model)\n", + " model_names.append(model_name)\n", + "\n", + "# Compare trained models\n", + "compare_models(models, model_names)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# Visualize a few sample predictions to qualitatively check differences\n", + "def visualize_predictions(models, model_names):\n", + " \"\"\"Visualize predictions from different models on the same test samples\"\"\"\n", + " # Get a few test samples\n", + " test_images, test_labels = mnist_test.test_data[:5], mnist_test.test_labels[:5]\n", + " test_images = torch.from_numpy(np.expand_dims(test_images, axis=1)).float()\n", + " \n", + " print(\"Predictions from different models:\")\n", + " for i, (image, label) in enumerate(zip(test_images, test_labels)):\n", + " print(f\"\\nSample {i+1}, True label: {label}\")\n", + " \n", + " # Make predictions with each model\n", + " for model, name in zip(models, model_names):\n", + " model.eval()\n", + " with torch.no_grad():\n", + " output = model(image.unsqueeze(0))\n", + " pred = output.argmax(dim=1, keepdim=True).item()\n", + " confidence = torch.nn.functional.softmax(output, dim=1).max().item()\n", + " print(f\" {name}: Predicted {pred} with confidence {confidence:.4f}\")\n", + "\n", + "# Visualize predictions\n", + "visualize_predictions(models, model_names)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Analysis of the Results\n", + "\n", + "The results above show how different mu values in the FedProx algorithm affect model training. The mu parameter controls the strength of the proximal term, which penalizes the local model for deviating too much from the global model.\n", + "\n", + "If the trained models with different mu values have identical weights, it would suggest that the mu parameter isn't having any effect on the training process with the current configuration. However, if they have different weights, it confirms that the mu parameter is working as expected - different values lead to different optimization paths.\n", + "\n", + "The weight differences and prediction differences provide insights into how much the mu parameter affects the training process and final model behavior." + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Comparing Models in the Federated Setting with Different Mu Values\n", + "\n", + "Let's now run a more comprehensive experiment to compare models trained in the federated setting with different mu values but the same random seed. This will help us understand how the mu parameter affects federated training." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# Function to run federated training with specific mu value\n", + "def run_federated_with_mu(mu_value, rounds=3, seed=42):\n", + " \"\"\"Run federated training with specific mu value and seed\"\"\"\n", + " print(f\"\\n\\n--- Starting federated training with mu={mu_value}, seed={seed} ---\\n\")\n", + " \n", + " # Set global seed\n", + " set_seed(seed)\n", + " \n", + " # Define optimizer getter function with specific mu\n", + " def get_optimizer_with_mu_value(model):\n", + " return FedProxAdam(model.parameters(), lr=1e-3, mu=mu_value)\n", + " \n", + " # Initialize model\n", + " model = Net()\n", + " \n", + " # Create and run federated flow with the specified mu value\n", + " flflow = FederatedFlow(model, get_optimizer_with_mu_value(model), rounds=rounds, checkpoint=False)\n", + " flflow.runtime = local_runtime\n", + " flflow.run()\n", + " \n", + " return model, flflow\n", + "\n", + "# Let's save the original get_optimizer function to restore later\n", + "original_get_optimizer = get_optimizer\n", + "\n", + "# Run experiment with different mu values\n", + "mu_values_federated = [0.0, 0.01, 0.1]\n", + "seed_value = 42\n", + "federated_models = []\n", + "federated_flows = []\n", + "federated_model_names = []\n", + "\n", + "for mu in mu_values_federated:\n", + " # Modify the get_optimizer function temporarily\n", + " def get_optimizer_with_current_mu(model):\n", + " return FedProxAdam(model.parameters(), lr=1e-3, mu=mu)\n", + " \n", + " # Replace the global function\n", + " globals()['get_optimizer'] = get_optimizer_with_current_mu\n", + " \n", + " # Run federated training\n", + " model_name = f\"Federated_Model_mu{mu}\"\n", + " model, flow = run_federated_with_mu(mu, rounds=2, seed=seed_value)\n", + " \n", + " # Save results\n", + " federated_models.append(model)\n", + " federated_flows.append(flow)\n", + " federated_model_names.append(model_name)\n", + "\n", + "# Restore original get_optimizer function\n", + "globals()['get_optimizer'] = original_get_optimizer\n", + "\n", + "# Compare the federated models\n", + "compare_models(federated_models, federated_model_names)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Comparing Convergence and Performance Metrics\n", + "\n", + "Now let's analyze how different mu values affect the convergence and final performance of the federated models." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# Extract metrics from each flow for comparison\n", + "def extract_metrics_from_flow(flow):\n", + " \"\"\"Extract metrics from a completed flow for comparison\"\"\"\n", + " rounds = flow.current_round\n", + " metrics = {\n", + " 'aggregated_accuracy': [],\n", + " 'local_accuracy': [],\n", + " 'loss': []\n", + " }\n", + " \n", + " # This is a simplification - in practice, we would extract metrics from the flow's history\n", + " # Here we're just using the final metrics as a proxy\n", + " metrics['aggregated_accuracy'].append(flow.aggregated_model_accuracy)\n", + " metrics['local_accuracy'].append(flow.local_model_accuracy)\n", + " metrics['average_loss'] = flow.average_loss\n", + " \n", + " return metrics\n", + "\n", + "# Collect metrics from all flows\n", + "all_metrics = []\n", + "for i, flow in enumerate(federated_flows):\n", + " mu = mu_values_federated[i]\n", + " metrics = extract_metrics_from_flow(flow)\n", + " metrics['mu'] = mu\n", + " all_metrics.append(metrics)\n", + "\n", + "# Print comparison of metrics\n", + "print(\"\\nComparison of metrics across different mu values:\")\n", + "print(\"-\" * 60)\n", + "print(f\"{'Mu Value':<10} | {'Final Aggregated Accuracy':<25} | {'Final Loss':<15}\")\n", + "print(\"-\" * 60)\n", + "\n", + "for metrics in all_metrics:\n", + " mu = metrics['mu']\n", + " agg_acc = metrics['aggregated_accuracy'][-1] if metrics['aggregated_accuracy'] else 'N/A'\n", + " loss = metrics['average_loss']\n", + " print(f\"{mu:<10.2f} | {agg_acc:<25.4f} | {loss:<15.6f}\")\n", + "\n", + "print(\"-\" * 60)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Testing Weight Divergence Directly \n", + "\n", + "The mu parameter in FedProx is designed to limit the divergence of local models from the global model. Let's directly measure this divergence for different mu values." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "def measure_weight_divergence(initial_model, trained_models, mu_values):\n", + " \"\"\"\n", + " Measure how much each trained model has diverged from the initial model\n", + " for different mu values\n", + " \"\"\"\n", + " print(\"\\nMeasuring weight divergence from initial model:\")\n", + " \n", + " # Extract the initial model weights\n", + " initial_weights = [p.clone().detach() for p in initial_model.parameters()]\n", + " \n", + " divergences = []\n", + " for i, (model, mu) in enumerate(zip(trained_models, mu_values)):\n", + " # Calculate divergence as the Euclidean distance between weight vectors\n", + " total_divergence = 0\n", + " for p_trained, p_initial in zip(model.parameters(), initial_weights):\n", + " # Calculate squared Frobenius norm of the difference\n", + " diff = p_trained - p_initial\n", + " divergence = torch.norm(diff.flatten(), p=2).item()\n", + " total_divergence += divergence\n", + " \n", + " divergences.append(total_divergence)\n", + " print(f\"Mu = {mu}: Total weight divergence = {total_divergence:.6f}\")\n", + " \n", + " return divergences\n", + "\n", + "# Create a fresh model to serve as the initial reference point\n", + "initial_model = Net()\n", + "\n", + "# Measure divergence\n", + "divergences = measure_weight_divergence(initial_model, federated_models, mu_values_federated)\n", + "\n", + "# We would expect models with higher mu to show less divergence from the initial point\n", + "# as the proximal term penalizes moving away from the global model" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Conclusion\n", + "\n", + "These experiments help us understand how the `mu` parameter in FedProx affects model training and convergence. The `mu` parameter controls the strength of the proximal term, which penalizes the local model for deviating too much from the global model.\n", + "\n", + "Key observations:\n", + "\n", + "1. Different `mu` values should lead to different model weights if the proximal term is working as expected.\n", + "2. Higher `mu` values should result in less divergence between local models and the global model.\n", + "3. The effect of `mu` on accuracy and convergence speed can help determine the optimal value for a specific federated learning task.\n", + "\n", + "This analysis provides insights into how to properly configure FedProx for different federated learning scenarios." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Important Note on FedProx Implementation\n", + "\n", + "This notebook includes a critical fix for how FedProx is applied in OpenFL:\n", + "\n", + "### The FedProx Proximal Term\n", + "\n", + "FedProx adds a proximal term to the objective function:\n", + "\n", + "L(w) = F_i(w) + (μ/2) ||w - w^t||^2\n", + "\n", + "Where:\n", + "- F_i(w) is the original loss function\n", + "- w^t is the global model from the previous round\n", + "- μ is the regularization parameter controlling how far local models can deviate\n", + "\n", + "### Correct Implementation\n", + "\n", + "The key insight is that `set_old_weights` should be called **once at the beginning** of each local training round, not before each optimization step:\n", + "\n", + "```python\n", + "# At the beginning of local training:\n", + "self.optimizer = get_optimizer(self.model)\n", + "self.optimizer.set_old_weights([p.clone().detach() for p in self.model.parameters()])\n", + "\n", + "# Then during training loop:\n", + "for batch_idx, (data, target) in enumerate(self.train_loader):\n", + " self.optimizer.zero_grad()\n", + " output = self.model(data)\n", + " loss = F.cross_entropy(output, target)\n", + " loss.backward()\n", + " # DO NOT call set_old_weights here\n", + " self.optimizer.step()\n", + "```\n", + "\n", + "This ensures that the proximal term properly penalizes deviation from the global model, allowing different mu values to have the expected effect on model convergence." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [] } ], "metadata": { "kernelspec": { - "display_name": "venv", + "display_name": "Python 3 (ipykernel)", "language": "python", "name": "python3" }, @@ -619,7 +790,7 @@ "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", - "version": "3.8.10" + "version": "3.11.12" } }, "nbformat": 4, diff --git a/openfl/utilities/optimizers/torch/fedprox.py b/openfl/utilities/optimizers/torch/fedprox.py index 2475260e32..ec61961346 100644 --- a/openfl/utilities/optimizers/torch/fedprox.py +++ b/openfl/utilities/optimizers/torch/fedprox.py @@ -75,6 +75,7 @@ def __init__( "mu": mu, "nesterov": nesterov, "weight_decay": weight_decay, + "w_old": None, # Initialize w_old as None } if nesterov and (momentum <= 0 or dampening != 0): @@ -115,7 +116,11 @@ def step(self, closure=None): nesterov = group["nesterov"] mu = group["mu"] w_old = group["w_old"] - for p, w_old_p in zip(group["params"], w_old): + + # Skip FedProx regularization if w_old is not set or mu is 0 + apply_proximal = w_old is not None and mu > 0 + + for i, p in enumerate(group["params"]): if p.grad is None: continue d_p = p.grad @@ -132,7 +137,9 @@ def step(self, closure=None): d_p = d_p.add(buf, alpha=momentum) else: d_p = buf - if w_old is not None: + if apply_proximal: + # Apply proximal term: mu * (p - w_old_p) + w_old_p = w_old[i] d_p.add_(p - w_old_p, alpha=mu) p.add_(d_p, alpha=-group["lr"]) @@ -212,6 +219,7 @@ def __init__( "weight_decay": weight_decay, "amsgrad": amsgrad, "mu": mu, + "w_old": None, # Initialize w_old as None } super().__init__(params, defaults) @@ -348,10 +356,17 @@ def adam( mu (float): Proximal term coefficient. w_old: The old weights. """ + # Skip FedProx regularization if w_old is not set or mu is 0 + apply_proximal = w_old is not None and mu > 0 + for i, param in enumerate(params): - w_old_p = w_old[i] grad = grads[i] - grad.add_(param - w_old_p, alpha=mu) + + # Apply proximal term only if we have valid old weights and mu > 0 + if apply_proximal: + w_old_p = w_old[i] + grad.add_(param - w_old_p, alpha=mu) + exp_avg = exp_avgs[i] exp_avg_sq = exp_avg_sqs[i] step = state_steps[i] From 6fce86e97694b819dfa1b3463f15de0f07b57dfc Mon Sep 17 00:00:00 2001 From: "Shekhawat, Nisha" Date: Sun, 18 May 2025 21:55:50 -0700 Subject: [PATCH 2/5] change Signed-off-by: Shekhawat, Nisha --- ...Prox_PyTorch_MNIST_Workflow_Tutorial.ipynb | 438 +----------------- 1 file changed, 1 insertion(+), 437 deletions(-) diff --git a/openfl-tutorials/experimental/workflow/403_Federated_FedProx_PyTorch_MNIST_Workflow_Tutorial.ipynb b/openfl-tutorials/experimental/workflow/403_Federated_FedProx_PyTorch_MNIST_Workflow_Tutorial.ipynb index 9511e63b06..45410f8719 100644 --- a/openfl-tutorials/experimental/workflow/403_Federated_FedProx_PyTorch_MNIST_Workflow_Tutorial.ipynb +++ b/openfl-tutorials/experimental/workflow/403_Federated_FedProx_PyTorch_MNIST_Workflow_Tutorial.ipynb @@ -336,447 +336,11 @@ "flflow.runtime = local_runtime\n", "flflow.run()" ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Model Comparison with Different Mu Values\n", - "Now let's check if trained models with the same seeding but different mu values produce the same results. We'll train multiple models with different mu values and then compare their weights." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "import copy\n", - "import random\n", - "import numpy as np\n", - "\n", - "# Function to set seed for reproducibility\n", - "def set_seed(seed):\n", - " \"\"\"Set seed for reproducibility\"\"\"\n", - " torch.manual_seed(seed)\n", - " torch.cuda.manual_seed_all(seed)\n", - " np.random.seed(seed)\n", - " random.seed(seed)\n", - " torch.backends.cudnn.deterministic = True\n", - " torch.backends.cudnn.benchmark = False\n", - "\n", - "# Function to get optimizer with different mu values\n", - "def get_optimizer_with_mu(model, mu_value):\n", - " \"\"\"Get FedProxAdam optimizer with specific mu value\"\"\"\n", - " return FedProxAdam(model.parameters(), lr=1e-3, mu=mu_value)\n", - "\n", - "# Function to compare model weights\n", - "def compare_models(models, model_names):\n", - " \"\"\"Compare weights between different models\"\"\"\n", - " print(\"Comparing model weights:\")\n", - " \n", - " # Compare weights between each pair of models\n", - " for i in range(len(models)):\n", - " for j in range(i+1, len(models)):\n", - " model1 = models[i]\n", - " model2 = models[j]\n", - " name1 = model_names[i]\n", - " name2 = model_names[j]\n", - " \n", - " print(f\"\\nComparing {name1} and {name2}:\")\n", - " \n", - " # Compare each layer's weights\n", - " all_equal = True\n", - " max_diff = 0.0\n", - " \n", - " for (p1, p2) in zip(model1.parameters(), model2.parameters()):\n", - " # Check if parameters are equal\n", - " if not torch.allclose(p1, p2, atol=1e-5):\n", - " all_equal = False\n", - " # Calculate maximum difference\n", - " diff = torch.max(torch.abs(p1 - p2)).item()\n", - " max_diff = max(max_diff, diff)\n", - " \n", - " if all_equal:\n", - " print(f\"Models {name1} and {name2} have identical weights\")\n", - " else:\n", - " print(f\"Models {name1} and {name2} have different weights\")\n", - " print(f\"Maximum difference in weights: {max_diff:.6f}\")\n", - "\n", - "# Function to train a model with specific mu value\n", - "def train_model(model, mu_value, seed, epochs=1):\n", - " \"\"\"Train a model with specific mu value and seed\"\"\"\n", - " set_seed(seed) # Set seed for reproducibility\n", - " \n", - " # Create a copy of the model\n", - " model_copy = copy.deepcopy(model)\n", - " \n", - " # Get optimizer with specific mu value\n", - " optimizer = get_optimizer_with_mu(model_copy, mu_value)\n", - " \n", - " # Training loop\n", - " model_copy.train()\n", - " \n", - " # Get a small dataset for quick testing\n", - " train_images, train_labels = mnist_train.train_data[:1000], np.array(mnist_train.train_labels[:1000])\n", - " train_images = torch.from_numpy(np.expand_dims(train_images, axis=1)).float()\n", - " train_labels = one_hot(train_labels, 10)\n", - " \n", - " train_dataset = CustomDataset(train_images, train_labels)\n", - " train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=32, shuffle=True)\n", - " \n", - " # CRITICAL FIX: Set the old weights ONCE at the beginning, before training\n", - " # This properly implements FedProx by setting the reference weights\n", - " # to the initial model weights (simulating the global model)\n", - " optimizer.set_old_weights([p.clone().detach() for p in model_copy.parameters()])\n", - " \n", - " for epoch in range(epochs):\n", - " running_loss = 0.0\n", - " for batch_idx, (data, target) in enumerate(train_loader):\n", - " optimizer.zero_grad()\n", - " output = model_copy(data)\n", - " loss = F.cross_entropy(output, target)\n", - " loss.backward()\n", - " \n", - " # REMOVED: Don't call set_old_weights here - this was causing the issue\n", - " optimizer.step()\n", - " \n", - " running_loss += loss.item()\n", - " \n", - " print(f\"Epoch {epoch+1}, Mu={mu_value}, Loss: {running_loss/len(train_loader):.6f}\")\n", - " \n", - " return model_copy" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "# Define mu values to test\n", - "mu_values = [0.0, 0.01, 0.1, 0.5]\n", - "seed = 42 # Fixed seed for reproducibility\n", - "models = []\n", - "model_names = []\n", - "\n", - "print(\"Training models with different mu values but same seed...\")\n", - "\n", - "# Train models with different mu values\n", - "for mu in mu_values:\n", - " model_name = f\"Model_mu{mu}\"\n", - " print(f\"\\nTraining {model_name} with mu={mu}...\")\n", - " \n", - " # Initialize a new model\n", - " base_model = Net()\n", - " \n", - " # Train the model with current mu value\n", - " trained_model = train_model(base_model, mu, seed, epochs=2)\n", - " \n", - " # Save the model and name\n", - " models.append(trained_model)\n", - " model_names.append(model_name)\n", - "\n", - "# Compare trained models\n", - "compare_models(models, model_names)" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "# Visualize a few sample predictions to qualitatively check differences\n", - "def visualize_predictions(models, model_names):\n", - " \"\"\"Visualize predictions from different models on the same test samples\"\"\"\n", - " # Get a few test samples\n", - " test_images, test_labels = mnist_test.test_data[:5], mnist_test.test_labels[:5]\n", - " test_images = torch.from_numpy(np.expand_dims(test_images, axis=1)).float()\n", - " \n", - " print(\"Predictions from different models:\")\n", - " for i, (image, label) in enumerate(zip(test_images, test_labels)):\n", - " print(f\"\\nSample {i+1}, True label: {label}\")\n", - " \n", - " # Make predictions with each model\n", - " for model, name in zip(models, model_names):\n", - " model.eval()\n", - " with torch.no_grad():\n", - " output = model(image.unsqueeze(0))\n", - " pred = output.argmax(dim=1, keepdim=True).item()\n", - " confidence = torch.nn.functional.softmax(output, dim=1).max().item()\n", - " print(f\" {name}: Predicted {pred} with confidence {confidence:.4f}\")\n", - "\n", - "# Visualize predictions\n", - "visualize_predictions(models, model_names)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Analysis of the Results\n", - "\n", - "The results above show how different mu values in the FedProx algorithm affect model training. The mu parameter controls the strength of the proximal term, which penalizes the local model for deviating too much from the global model.\n", - "\n", - "If the trained models with different mu values have identical weights, it would suggest that the mu parameter isn't having any effect on the training process with the current configuration. However, if they have different weights, it confirms that the mu parameter is working as expected - different values lead to different optimization paths.\n", - "\n", - "The weight differences and prediction differences provide insights into how much the mu parameter affects the training process and final model behavior." - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Comparing Models in the Federated Setting with Different Mu Values\n", - "\n", - "Let's now run a more comprehensive experiment to compare models trained in the federated setting with different mu values but the same random seed. This will help us understand how the mu parameter affects federated training." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "# Function to run federated training with specific mu value\n", - "def run_federated_with_mu(mu_value, rounds=3, seed=42):\n", - " \"\"\"Run federated training with specific mu value and seed\"\"\"\n", - " print(f\"\\n\\n--- Starting federated training with mu={mu_value}, seed={seed} ---\\n\")\n", - " \n", - " # Set global seed\n", - " set_seed(seed)\n", - " \n", - " # Define optimizer getter function with specific mu\n", - " def get_optimizer_with_mu_value(model):\n", - " return FedProxAdam(model.parameters(), lr=1e-3, mu=mu_value)\n", - " \n", - " # Initialize model\n", - " model = Net()\n", - " \n", - " # Create and run federated flow with the specified mu value\n", - " flflow = FederatedFlow(model, get_optimizer_with_mu_value(model), rounds=rounds, checkpoint=False)\n", - " flflow.runtime = local_runtime\n", - " flflow.run()\n", - " \n", - " return model, flflow\n", - "\n", - "# Let's save the original get_optimizer function to restore later\n", - "original_get_optimizer = get_optimizer\n", - "\n", - "# Run experiment with different mu values\n", - "mu_values_federated = [0.0, 0.01, 0.1]\n", - "seed_value = 42\n", - "federated_models = []\n", - "federated_flows = []\n", - "federated_model_names = []\n", - "\n", - "for mu in mu_values_federated:\n", - " # Modify the get_optimizer function temporarily\n", - " def get_optimizer_with_current_mu(model):\n", - " return FedProxAdam(model.parameters(), lr=1e-3, mu=mu)\n", - " \n", - " # Replace the global function\n", - " globals()['get_optimizer'] = get_optimizer_with_current_mu\n", - " \n", - " # Run federated training\n", - " model_name = f\"Federated_Model_mu{mu}\"\n", - " model, flow = run_federated_with_mu(mu, rounds=2, seed=seed_value)\n", - " \n", - " # Save results\n", - " federated_models.append(model)\n", - " federated_flows.append(flow)\n", - " federated_model_names.append(model_name)\n", - "\n", - "# Restore original get_optimizer function\n", - "globals()['get_optimizer'] = original_get_optimizer\n", - "\n", - "# Compare the federated models\n", - "compare_models(federated_models, federated_model_names)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Comparing Convergence and Performance Metrics\n", - "\n", - "Now let's analyze how different mu values affect the convergence and final performance of the federated models." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "# Extract metrics from each flow for comparison\n", - "def extract_metrics_from_flow(flow):\n", - " \"\"\"Extract metrics from a completed flow for comparison\"\"\"\n", - " rounds = flow.current_round\n", - " metrics = {\n", - " 'aggregated_accuracy': [],\n", - " 'local_accuracy': [],\n", - " 'loss': []\n", - " }\n", - " \n", - " # This is a simplification - in practice, we would extract metrics from the flow's history\n", - " # Here we're just using the final metrics as a proxy\n", - " metrics['aggregated_accuracy'].append(flow.aggregated_model_accuracy)\n", - " metrics['local_accuracy'].append(flow.local_model_accuracy)\n", - " metrics['average_loss'] = flow.average_loss\n", - " \n", - " return metrics\n", - "\n", - "# Collect metrics from all flows\n", - "all_metrics = []\n", - "for i, flow in enumerate(federated_flows):\n", - " mu = mu_values_federated[i]\n", - " metrics = extract_metrics_from_flow(flow)\n", - " metrics['mu'] = mu\n", - " all_metrics.append(metrics)\n", - "\n", - "# Print comparison of metrics\n", - "print(\"\\nComparison of metrics across different mu values:\")\n", - "print(\"-\" * 60)\n", - "print(f\"{'Mu Value':<10} | {'Final Aggregated Accuracy':<25} | {'Final Loss':<15}\")\n", - "print(\"-\" * 60)\n", - "\n", - "for metrics in all_metrics:\n", - " mu = metrics['mu']\n", - " agg_acc = metrics['aggregated_accuracy'][-1] if metrics['aggregated_accuracy'] else 'N/A'\n", - " loss = metrics['average_loss']\n", - " print(f\"{mu:<10.2f} | {agg_acc:<25.4f} | {loss:<15.6f}\")\n", - "\n", - "print(\"-\" * 60)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Testing Weight Divergence Directly \n", - "\n", - "The mu parameter in FedProx is designed to limit the divergence of local models from the global model. Let's directly measure this divergence for different mu values." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "def measure_weight_divergence(initial_model, trained_models, mu_values):\n", - " \"\"\"\n", - " Measure how much each trained model has diverged from the initial model\n", - " for different mu values\n", - " \"\"\"\n", - " print(\"\\nMeasuring weight divergence from initial model:\")\n", - " \n", - " # Extract the initial model weights\n", - " initial_weights = [p.clone().detach() for p in initial_model.parameters()]\n", - " \n", - " divergences = []\n", - " for i, (model, mu) in enumerate(zip(trained_models, mu_values)):\n", - " # Calculate divergence as the Euclidean distance between weight vectors\n", - " total_divergence = 0\n", - " for p_trained, p_initial in zip(model.parameters(), initial_weights):\n", - " # Calculate squared Frobenius norm of the difference\n", - " diff = p_trained - p_initial\n", - " divergence = torch.norm(diff.flatten(), p=2).item()\n", - " total_divergence += divergence\n", - " \n", - " divergences.append(total_divergence)\n", - " print(f\"Mu = {mu}: Total weight divergence = {total_divergence:.6f}\")\n", - " \n", - " return divergences\n", - "\n", - "# Create a fresh model to serve as the initial reference point\n", - "initial_model = Net()\n", - "\n", - "# Measure divergence\n", - "divergences = measure_weight_divergence(initial_model, federated_models, mu_values_federated)\n", - "\n", - "# We would expect models with higher mu to show less divergence from the initial point\n", - "# as the proximal term penalizes moving away from the global model" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Conclusion\n", - "\n", - "These experiments help us understand how the `mu` parameter in FedProx affects model training and convergence. The `mu` parameter controls the strength of the proximal term, which penalizes the local model for deviating too much from the global model.\n", - "\n", - "Key observations:\n", - "\n", - "1. Different `mu` values should lead to different model weights if the proximal term is working as expected.\n", - "2. Higher `mu` values should result in less divergence between local models and the global model.\n", - "3. The effect of `mu` on accuracy and convergence speed can help determine the optimal value for a specific federated learning task.\n", - "\n", - "This analysis provides insights into how to properly configure FedProx for different federated learning scenarios." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Important Note on FedProx Implementation\n", - "\n", - "This notebook includes a critical fix for how FedProx is applied in OpenFL:\n", - "\n", - "### The FedProx Proximal Term\n", - "\n", - "FedProx adds a proximal term to the objective function:\n", - "\n", - "L(w) = F_i(w) + (μ/2) ||w - w^t||^2\n", - "\n", - "Where:\n", - "- F_i(w) is the original loss function\n", - "- w^t is the global model from the previous round\n", - "- μ is the regularization parameter controlling how far local models can deviate\n", - "\n", - "### Correct Implementation\n", - "\n", - "The key insight is that `set_old_weights` should be called **once at the beginning** of each local training round, not before each optimization step:\n", - "\n", - "```python\n", - "# At the beginning of local training:\n", - "self.optimizer = get_optimizer(self.model)\n", - "self.optimizer.set_old_weights([p.clone().detach() for p in self.model.parameters()])\n", - "\n", - "# Then during training loop:\n", - "for batch_idx, (data, target) in enumerate(self.train_loader):\n", - " self.optimizer.zero_grad()\n", - " output = self.model(data)\n", - " loss = F.cross_entropy(output, target)\n", - " loss.backward()\n", - " # DO NOT call set_old_weights here\n", - " self.optimizer.step()\n", - "```\n", - "\n", - "This ensures that the proximal term properly penalizes deviation from the global model, allowing different mu values to have the expected effect on model convergence." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [] } ], "metadata": { "kernelspec": { - "display_name": "Python 3 (ipykernel)", + "display_name": "env_name", "language": "python", "name": "python3" }, From 3387aaacfbb534c04f355504d2975ac2f3a0844b Mon Sep 17 00:00:00 2001 From: "Shekhawat, Nisha" Date: Mon, 19 May 2025 23:07:32 -0700 Subject: [PATCH 3/5] change_for_mu_values_comment Signed-off-by: Shekhawat, Nisha --- openfl/utilities/optimizers/torch/fedprox.py | 78 ++++++++++++++++++-- 1 file changed, 70 insertions(+), 8 deletions(-) diff --git a/openfl/utilities/optimizers/torch/fedprox.py b/openfl/utilities/optimizers/torch/fedprox.py index ec61961346..4b84fd29e0 100644 --- a/openfl/utilities/optimizers/torch/fedprox.py +++ b/openfl/utilities/optimizers/torch/fedprox.py @@ -20,6 +20,14 @@ class FedProxOptimizer(Optimizer): It introduces a proximal term to the federated averaging algorithm to reduce the impact of devices with outlying updates. + IMPORTANT: This optimizer requires a reference to the original (global) model parameters + to calculate the proximal term. These must be set explicitly using the set_old_weights() + method before training begins. The old weights (w_old) must match the order and structure + of the model's parameters. Typically, w_old should be set to the initial global model + parameters received from the aggregator at the beginning of each round. + + If mu > 0 and w_old is not set, the optimizer will raise a ValueError. + Paper: https://arxiv.org/pdf/1812.06127.pdf Attributes: @@ -67,7 +75,12 @@ def __init__( if weight_decay < 0.0: raise ValueError(f"Invalid weight_decay value: {weight_decay}") if mu < 0.0: - raise ValueError(f"Invalid mu value: {mu}") + import warnings + warnings.warn( + f"Negative mu value ({mu}) will cause the proximal term to reward " + f"deviations from global weights, which may be counterintuitive.", + UserWarning, + ) defaults = { "dampening": dampening, "lr": lr, @@ -117,8 +130,15 @@ def step(self, closure=None): mu = group["mu"] w_old = group["w_old"] - # Skip FedProx regularization if w_old is not set or mu is 0 - apply_proximal = w_old is not None and mu > 0 + # Check if FedProx regularization should be applied (mu > 0) + if mu > 0 and w_old is None: + raise ValueError( + "FedProx requires old weights to be set when mu > 0. " + "Please call set_old_weights() before optimization step." + ) + + # Apply proximal term when mu != 0 + apply_proximal = w_old is not None and mu != 0 for i, p in enumerate(group["params"]): if p.grad is None: @@ -147,9 +167,20 @@ def step(self, closure=None): def set_old_weights(self, old_weights): """Set the global weights parameter to `old_weights` value. + + This method must be called before training begins to set the reference point for + calculating the proximal term in FedProx. Typically, this should be set to the + initial global model parameters received from the aggregator at the beginning + of each federated learning round. + + If mu > 0 and this method is not called, the optimizer will raise a ValueError + during the optimization step. Args: - old_weights: The old weights to be set. + old_weights: List of parameter tensors representing the global model weights. + Must match the order and structure of the model's parameters + being optimized (typically obtained by calling + [p.clone().detach() for p in model.parameters()]). """ for param_group in self.param_groups: param_group["w_old"] = old_weights @@ -160,6 +191,14 @@ class FedProxAdam(Optimizer): Implements the FedProx optimization algorithm with Adam optimizer. + IMPORTANT: This optimizer requires a reference to the original (global) model parameters + to calculate the proximal term. These must be set explicitly using the set_old_weights() + method before training begins. The old weights (w_old) must match the order and structure + of the model's parameters. Typically, w_old should be set to the initial global model + parameters received from the aggregator at the beginning of each round. + + If mu > 0 and w_old is not set, the optimizer will raise a ValueError. + Attributes: params: Parameters to be stored for optimization. mu: Proximal term coefficient. @@ -211,7 +250,12 @@ def __init__( if not 0.0 <= weight_decay: raise ValueError(f"Invalid weight_decay value: {weight_decay}") if mu < 0.0: - raise ValueError(f"Invalid mu value: {mu}") + import warnings + warnings.warn( + f"Negative mu value ({mu}) will cause the proximal term to reward " + f"deviations from global weights, which may be counterintuitive.", + UserWarning, + ) defaults = { "lr": lr, "betas": betas, @@ -231,9 +275,20 @@ def __setstate__(self, state): def set_old_weights(self, old_weights): """Set the global weights parameter to `old_weights` value. + + This method must be called before training begins to set the reference point for + calculating the proximal term in FedProx. Typically, this should be set to the + initial global model parameters received from the aggregator at the beginning + of each federated learning round. + + If mu > 0 and this method is not called, the optimizer will raise a ValueError + during the optimization step. Args: - old_weights: The old weights to be set. + old_weights: List of parameter tensors representing the global model weights. + Must match the order and structure of the model's parameters + being optimized (typically obtained by calling + [p.clone().detach() for p in model.parameters()]). """ for param_group in self.param_groups: param_group["w_old"] = old_weights @@ -356,8 +411,15 @@ def adam( mu (float): Proximal term coefficient. w_old: The old weights. """ - # Skip FedProx regularization if w_old is not set or mu is 0 - apply_proximal = w_old is not None and mu > 0 + # Check if FedProx regularization should be applied (mu > 0) + if mu > 0 and w_old is None: + raise ValueError( + "FedProx requires old weights to be set when mu > 0. " + "Please call set_old_weights() before optimization step." + ) + + # Apply proximal term when mu != 0 + apply_proximal = w_old is not None and mu != 0 for i, param in enumerate(params): grad = grads[i] From 04a7967d07470387715723b7ca293d92299ddcab Mon Sep 17 00:00:00 2001 From: "sys_svc_tpe_perf@intel.com" Date: Mon, 19 May 2025 23:36:45 -0700 Subject: [PATCH 4/5] fix_pre_commit_flake Signed-off-by: Shekhawat, Nisha --- openfl/utilities/optimizers/torch/fedprox.py | 252 +++++++++++++------ 1 file changed, 180 insertions(+), 72 deletions(-) diff --git a/openfl/utilities/optimizers/torch/fedprox.py b/openfl/utilities/optimizers/torch/fedprox.py index 4b84fd29e0..50eb9093ec 100644 --- a/openfl/utilities/optimizers/torch/fedprox.py +++ b/openfl/utilities/optimizers/torch/fedprox.py @@ -20,12 +20,12 @@ class FedProxOptimizer(Optimizer): It introduces a proximal term to the federated averaging algorithm to reduce the impact of devices with outlying updates. - IMPORTANT: This optimizer requires a reference to the original (global) model parameters - to calculate the proximal term. These must be set explicitly using the set_old_weights() - method before training begins. The old weights (w_old) must match the order and structure - of the model's parameters. Typically, w_old should be set to the initial global model + IMPORTANT: This optimizer requires a reference to the original (global) model parameters + to calculate the proximal term. These must be set explicitly using the set_old_weights() + method before training begins. The old weights (w_old) must match the order and structure + of the model's parameters. Typically, w_old should be set to the initial global model parameters received from the aggregator at the beginning of each round. - + If mu > 0 and w_old is not set, the optimizer will raise a ValueError. Paper: https://arxiv.org/pdf/1812.06127.pdf @@ -76,10 +76,12 @@ def __init__( raise ValueError(f"Invalid weight_decay value: {weight_decay}") if mu < 0.0: import warnings + warnings.warn( f"Negative mu value ({mu}) will cause the proximal term to reward " f"deviations from global weights, which may be counterintuitive.", UserWarning, + stacklevel=2, ) defaults = { "dampening": dampening, @@ -107,6 +109,47 @@ def __setstate__(self, state): for group in self.param_groups: group.setdefault("nesterov", False) + def _validate_old_weights(self, mu, w_old): + """Validate old weights for FedProx regularization. + + Args: + mu: Proximal term coefficient + w_old: Old weights reference + + Raises: + ValueError: If mu > 0 and w_old is None + """ + if mu > 0 and w_old is None: + raise ValueError( + "FedProx requires old weights to be set when mu > 0. " + "Please call set_old_weights() before optimization step." + ) + + def _apply_momentum(self, p, d_p, momentum, dampening, nesterov): + """Apply momentum to gradient. + + Args: + p: Parameter + d_p: Gradient + momentum: Momentum factor + dampening: Dampening factor + nesterov: Whether to use Nesterov momentum + + Returns: + Modified gradient + """ + param_state = self.state[p] + if "momentum_buffer" not in param_state: + buf = param_state["momentum_buffer"] = torch.clone(d_p).detach() + else: + buf = param_state["momentum_buffer"] + buf.mul_(momentum).add_(d_p, alpha=1 - dampening) + if nesterov: + d_p = d_p.add(buf, alpha=momentum) + else: + d_p = buf + return d_p + @torch.no_grad() def step(self, closure=None): """Perform a single optimization step. @@ -129,50 +172,45 @@ def step(self, closure=None): nesterov = group["nesterov"] mu = group["mu"] w_old = group["w_old"] - - # Check if FedProx regularization should be applied (mu > 0) - if mu > 0 and w_old is None: - raise ValueError( - "FedProx requires old weights to be set when mu > 0. " - "Please call set_old_weights() before optimization step." - ) - + + # Validate old weights for FedProx + self._validate_old_weights(mu, w_old) + # Apply proximal term when mu != 0 apply_proximal = w_old is not None and mu != 0 - + for i, p in enumerate(group["params"]): if p.grad is None: continue + d_p = p.grad + + # Apply weight decay if weight_decay != 0: d_p = d_p.add(p, alpha=weight_decay) + + # Apply momentum if momentum != 0: - param_state = self.state[p] - if "momentum_buffer" not in param_state: - buf = param_state["momentum_buffer"] = torch.clone(d_p).detach() - else: - buf = param_state["momentum_buffer"] - buf.mul_(momentum).add_(d_p, alpha=1 - dampening) - if nesterov: - d_p = d_p.add(buf, alpha=momentum) - else: - d_p = buf + d_p = self._apply_momentum(p, d_p, momentum, dampening, nesterov) + + # Apply proximal term if apply_proximal: - # Apply proximal term: mu * (p - w_old_p) w_old_p = w_old[i] d_p.add_(p - w_old_p, alpha=mu) + + # Apply gradient step p.add_(d_p, alpha=-group["lr"]) return loss def set_old_weights(self, old_weights): """Set the global weights parameter to `old_weights` value. - + This method must be called before training begins to set the reference point for calculating the proximal term in FedProx. Typically, this should be set to the initial global model parameters received from the aggregator at the beginning of each federated learning round. - + If mu > 0 and this method is not called, the optimizer will raise a ValueError during the optimization step. @@ -191,20 +229,19 @@ class FedProxAdam(Optimizer): Implements the FedProx optimization algorithm with Adam optimizer. - IMPORTANT: This optimizer requires a reference to the original (global) model parameters - to calculate the proximal term. These must be set explicitly using the set_old_weights() - method before training begins. The old weights (w_old) must match the order and structure - of the model's parameters. Typically, w_old should be set to the initial global model + IMPORTANT: This optimizer requires a reference to the original (global) model parameters + to calculate the proximal term. These must be set explicitly using the set_old_weights() + method before training begins. The old weights (w_old) must match the order and structure + of the model's parameters. Typically, w_old should be set to the initial global model parameters received from the aggregator at the beginning of each round. - + If mu > 0 and w_old is not set, the optimizer will raise a ValueError. Attributes: params: Parameters to be stored for optimization. mu: Proximal term coefficient. lr: Learning rate. - betas: Coefficients used for computing running averages of gradient - and its square. + betas: Coefficients used for computing running averages of gradient and its square. eps: Value for computational stability. weight_decay: Weight decay (L2 penalty). amsgrad: Whether to use the AMSGrad variant of this algorithm. @@ -251,10 +288,12 @@ def __init__( raise ValueError(f"Invalid weight_decay value: {weight_decay}") if mu < 0.0: import warnings + warnings.warn( f"Negative mu value ({mu}) will cause the proximal term to reward " f"deviations from global weights, which may be counterintuitive.", UserWarning, + stacklevel=2, ) defaults = { "lr": lr, @@ -275,12 +314,12 @@ def __setstate__(self, state): def set_old_weights(self, old_weights): """Set the global weights parameter to `old_weights` value. - + This method must be called before training begins to set the reference point for calculating the proximal term in FedProx. Typically, this should be set to the initial global model parameters received from the aggregator at the beginning of each federated learning round. - + If mu > 0 and this method is not called, the optimizer will raise a ValueError during the optimization step. @@ -373,6 +412,88 @@ def step(self, closure=None): ) return loss + def _validate_old_weights(self, mu, w_old): + """Validate old weights for FedProx regularization. + + Args: + mu: Proximal term coefficient + w_old: Old weights reference + + Raises: + ValueError: If mu > 0 and w_old is None + """ + if mu > 0 and w_old is None: + raise ValueError( + "FedProx requires old weights to be set when mu > 0. " + "Please call set_old_weights() before optimization step.", + ) + + def _apply_proximal_term(self, grad, param, w_old_p, mu): + """Apply proximal term to gradient. + + Args: + grad: Gradient + param: Parameter + w_old_p: Old weight parameter + mu: Proximal term coefficient + + Returns: + Modified gradient + """ + return grad.add(param - w_old_p, alpha=mu) + + def _compute_adam_step( + self, + param, + grad, + exp_avg, + exp_avg_sq, + max_exp_avg_sq, + step, + amsgrad, + beta1, + beta2, + lr, + weight_decay, + eps, + ): + """Compute Adam optimization step. + + Args: + param: Parameter + grad: Gradient + exp_avg: Exponential moving average + exp_avg_sq: Exponential moving average squared + max_exp_avg_sq: Maximum exponential moving average squared + step: Step count + amsgrad: Whether to use AMSGrad + beta1: Beta1 coefficient + beta2: Beta2 coefficient + lr: Learning rate + weight_decay: Weight decay + eps: Epsilon value + """ + bias_correction1 = 1 - beta1**step + bias_correction2 = 1 - beta2**step + + if weight_decay != 0: + grad = grad.add(param, alpha=weight_decay) + + # Decay the first and second moment running average coefficient + exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1) + exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2) + + if amsgrad: + # Maintains the maximum of all 2nd moment running avg. till now + torch.maximum(max_exp_avg_sq, exp_avg_sq, out=max_exp_avg_sq) + # Use the max. for normalizing running avg. of gradient + denom = (max_exp_avg_sq.sqrt() / math.sqrt(bias_correction2)).add_(eps) + else: + denom = (exp_avg_sq.sqrt() / math.sqrt(bias_correction2)).add_(eps) + + step_size = lr / bias_correction1 + param.addcdiv_(exp_avg, denom, value=-step_size) + def adam( self, params, @@ -411,45 +532,32 @@ def adam( mu (float): Proximal term coefficient. w_old: The old weights. """ - # Check if FedProx regularization should be applied (mu > 0) - if mu > 0 and w_old is None: - raise ValueError( - "FedProx requires old weights to be set when mu > 0. " - "Please call set_old_weights() before optimization step." - ) - + # Validate old weights for FedProx + self._validate_old_weights(mu, w_old) + # Apply proximal term when mu != 0 apply_proximal = w_old is not None and mu != 0 - + for i, param in enumerate(params): grad = grads[i] - - # Apply proximal term only if we have valid old weights and mu > 0 + + # Apply proximal term if needed if apply_proximal: w_old_p = w_old[i] - grad.add_(param - w_old_p, alpha=mu) - - exp_avg = exp_avgs[i] - exp_avg_sq = exp_avg_sqs[i] - step = state_steps[i] - - bias_correction1 = 1 - beta1**step - bias_correction2 = 1 - beta2**step - - if weight_decay != 0: - grad = grad.add(param, alpha=weight_decay) - - # Decay the first and second moment running average coefficient - exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1) - exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2) - if amsgrad: - # Maintains the maximum of all 2nd moment running avg. till now - torch.maximum(max_exp_avg_sqs[i], exp_avg_sq, out=max_exp_avg_sqs[i]) - # Use the max. for normalizing running avg. of gradient - denom = (max_exp_avg_sqs[i].sqrt() / math.sqrt(bias_correction2)).add_(eps) - else: - denom = (exp_avg_sq.sqrt() / math.sqrt(bias_correction2)).add_(eps) - - step_size = lr / bias_correction1 - - param.addcdiv_(exp_avg, denom, value=-step_size) + grad = self._apply_proximal_term(grad, param, w_old_p, mu) + + # Apply Adam optimization steps + self._compute_adam_step( + param, + grad, + exp_avgs[i], + exp_avg_sqs[i], + max_exp_avg_sqs[i] if amsgrad else None, + state_steps[i], + amsgrad, + beta1, + beta2, + lr, + weight_decay, + eps, + ) From d432356fc1c44f67c77a30e74fd5373f6a2fbd3c Mon Sep 17 00:00:00 2001 From: "Shekhawat, Nisha" Date: Tue, 20 May 2025 21:46:53 -0700 Subject: [PATCH 5/5] remove_mu_is_0 Signed-off-by: Shekhawat, Nisha --- openfl/utilities/optimizers/torch/fedprox.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) mode change 100644 => 100755 openfl/utilities/optimizers/torch/fedprox.py diff --git a/openfl/utilities/optimizers/torch/fedprox.py b/openfl/utilities/optimizers/torch/fedprox.py old mode 100644 new mode 100755 index 50eb9093ec..d86ffec453 --- a/openfl/utilities/optimizers/torch/fedprox.py +++ b/openfl/utilities/optimizers/torch/fedprox.py @@ -177,7 +177,7 @@ def step(self, closure=None): self._validate_old_weights(mu, w_old) # Apply proximal term when mu != 0 - apply_proximal = w_old is not None and mu != 0 + apply_proximal = w_old is not None for i, p in enumerate(group["params"]): if p.grad is None: @@ -536,7 +536,7 @@ def adam( self._validate_old_weights(mu, w_old) # Apply proximal term when mu != 0 - apply_proximal = w_old is not None and mu != 0 + apply_proximal = w_old is not None for i, param in enumerate(params): grad = grads[i]