From f24c8d03d851b4a5595a9d1b65ae5f5d158e8694 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Thu, 11 Dec 2025 16:11:39 -0600 Subject: [PATCH] memory improvement --- cezo_fl/client.py | 25 ++++++++++++++++----- cezo_fl/gradient_estimators/adam_forward.py | 2 +- 2 files changed, 21 insertions(+), 6 deletions(-) diff --git a/cezo_fl/client.py b/cezo_fl/client.py index 1bd737c..ce16c85 100644 --- a/cezo_fl/client.py +++ b/cezo_fl/client.py @@ -77,11 +77,16 @@ def __init__( criterion: CriterionType, accuracy_func, device: torch.device, + offload_to_cpu: bool = True, ): self.model = model self.model_inference = model_inference self.dataloader = dataloader + if offload_to_cpu and not isinstance(optimizer, torch.optim.SGD): + raise Exception("offload to cpu only works with SGD at this point of time") + self.offload_to_cpu = offload_to_cpu + self.device = device self.grad_estimator = grad_estimator @@ -166,14 +171,24 @@ def local_update(self, seeds: Sequence[int]) -> LocalUpdateResult: def reset_model(self) -> None: """Reset the mode to the state before the local_update.""" assert self.last_pull_state_dict is not None - self.model.load_state_dict(self.last_pull_state_dict["model"]) - self.optimizer.load_state_dict(self.last_pull_state_dict["optimizer"]) + if self.offload_to_cpu: + self.model.cpu() + self.model.load_state_dict(self.last_pull_state_dict["model"]) + self.model.to(self.device) + self.optimizer.load_state_dict(self.last_pull_state_dict["optimizer"]) + else: + self.model.load_state_dict(self.last_pull_state_dict["model"]) + self.optimizer.load_state_dict(self.last_pull_state_dict["optimizer"]) def screenshot(self) -> dict: # deepcopy current model.state_dict and optimizer.state_dict - return deepcopy( - {"model": self.model.state_dict(), "optimizer": self.optimizer.state_dict()} - ) + if self.offload_to_cpu: + model_state_dict = {k: deepcopy(v).cpu() for k, v in self.model.state_dict().items()} + return {"model": model_state_dict, "optimizer": deepcopy(self.optimizer.state_dict())} + else: + return deepcopy( + {"model": self.model.state_dict(), "optimizer": self.optimizer.state_dict()} + ) def pull_model( self, diff --git a/cezo_fl/gradient_estimators/adam_forward.py b/cezo_fl/gradient_estimators/adam_forward.py index 8d29ab3..9e23d2c 100644 --- a/cezo_fl/gradient_estimators/adam_forward.py +++ b/cezo_fl/gradient_estimators/adam_forward.py @@ -184,7 +184,7 @@ def generate_perturbation_norm_paramwise( param_k = self.K_param_list[param_index] return torch.randn( *param.shape, device=self.device, dtype=self.torch_dtype, generator=rng - ) / torch.sqrt(param_k) + ).div_(torch.sqrt(param_k)) def compute_grad( self,