From d1102ce5c0045c0e8776f0bdd23522fd680bb893 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Sun, 28 Apr 2024 21:26:37 -0700 Subject: [PATCH 01/27] use torch.no_grad() to avoid calling cat() during FSDP backward except for last microbatch --- .../nn/data_parallel/fully_sharded_data_parallel.py | 7 ++++--- fairscale/nn/misc/flatten_params_wrapper.py | 13 ++++++++++++- 2 files changed, 16 insertions(+), 4 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index cdd6e6e8c..48c1bd0ff 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1099,12 +1099,14 @@ def no_sync(self) -> Generator: if isinstance(m, FullyShardedDataParallel): old_flags.append((m, m._require_backward_grad_sync)) m._require_backward_grad_sync = False + m._fsdp_wrapped_module._require_backward_grad_sync = False try: yield finally: for m, old_flag in old_flags: assert m._require_backward_grad_sync is False m._require_backward_grad_sync = old_flag + m._fsdp_wrapped_module._require_backward_grad_sync = old_flag @contextlib.contextmanager def summon_full_params(self, recurse: bool = True, volatile: bool = False) -> Generator: @@ -1458,7 +1460,6 @@ def forward(self, *args: Any, **kwargs: Any) -> torch.Tensor: # Register backward hooks to reshard params and reduce-scatter grads. # These need to be re-registered every forward pass. self._register_post_backward_hooks() - outputs = self.module(*args, **kwargs) if self.reshard_after_forward: @@ -1851,7 +1852,7 @@ def _wait_for_post_backward(self) -> None: # the `requires_grad` field set. If `requires_grad=False` for # all the params, the post_backward hook will not fire and the # state will remain in `TrainingState.BACKWARD_PRE`. - if any([p.requires_grad for p in self.params]): + if any([p.requires_grad for p in self.params]) and self._fsdp_wrapped_module._require_backward_grad_sync: self.assert_state(TrainingState.BACKWARD_POST) else: self.assert_state(TrainingState.BACKWARD_PRE) @@ -1928,7 +1929,7 @@ def _finalize_parameters(fsdp_module: FullyShardedDataParallel) -> None: # the `requires_grad` field set. If `requires_grad=False` for # all the params, the post_backward hook will not fire and the # state will remain in `TrainingState.BACKWARD_PRE`. - if any([p.requires_grad for p in m.params]): + if any([p.requires_grad for p in m.params]) and self._fsdp_wrapped_module._require_backward_grad_sync: m.assert_state(TrainingState.BACKWARD_POST) else: m.assert_state(TrainingState.BACKWARD_PRE) diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 38265dd2b..455442e69 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -37,6 +37,8 @@ if TYPE_CHECKING: from collections import OrderedDict # noqa: F401 +from logging import getLogger +logger = getLogger() class FlatParameter(nn.Parameter): """A parameter that is initialized from a list of parameters and can be @@ -161,6 +163,7 @@ def __init__( super().__init__() self._fpw_module = module self.is_flattened = False + self._require_backward_grad_sync = True # Handle param_list being None. if param_list is None: @@ -369,7 +372,15 @@ def _unflatten_params_as_views(self) -> None: self.flat_param unchanged. """ assert self.is_flattened - ps = self.get_param_views() + logger.info(f"CHRISLOG: {self._require_backward_grad_sync=}") + if self._require_backward_grad_sync: + logger.info("CHRISLOG: calling self.get_param_views() without torch.no_grad()") + ps = self.get_param_views() + else: + with torch.no_grad(): + logger.info("CHRISLOG: calling self.get_param_views() with torch.no_grad()") + ps = self.get_param_views() + param_views = [] for (_, m, n), p in zip(self._param_infos, ps): setattr(p, '_fsdp_weight', True) From 9a2262893b76a04c11a5744e0eae549059f578e7 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Sun, 28 Apr 2024 22:52:12 -0700 Subject: [PATCH 02/27] remove logging --- fairscale/nn/misc/flatten_params_wrapper.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 455442e69..7298a051e 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -372,13 +372,13 @@ def _unflatten_params_as_views(self) -> None: self.flat_param unchanged. """ assert self.is_flattened - logger.info(f"CHRISLOG: {self._require_backward_grad_sync=}") + #logger.info(f"CHRISLOG: {self._require_backward_grad_sync=}") if self._require_backward_grad_sync: - logger.info("CHRISLOG: calling self.get_param_views() without torch.no_grad()") + #logger.info("CHRISLOG: calling self.get_param_views() without torch.no_grad()") ps = self.get_param_views() else: with torch.no_grad(): - logger.info("CHRISLOG: calling self.get_param_views() with torch.no_grad()") + #logger.info("CHRISLOG: calling self.get_param_views() with torch.no_grad()") ps = self.get_param_views() param_views = [] From f787532e3a7f6533125b0d94ebb1fbead9fcbffd Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Mon, 29 Apr 2024 21:52:54 -0700 Subject: [PATCH 03/27] logging --- fairscale/nn/misc/flatten_params_wrapper.py | 1 + 1 file changed, 1 insertion(+) diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 7298a051e..bc1209bbb 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -385,6 +385,7 @@ def _unflatten_params_as_views(self) -> None: for (_, m, n), p in zip(self._param_infos, ps): setattr(p, '_fsdp_weight', True) setattr(m, n, p) # This will set as plain attr + #logger.info(f"CHRISLOG: {n=}, {p.requires_grad=}") param_views.append(p) # Save param views for easy access if anyone still wants to access From 3429f33b58513185fafb6c01d1416817c6577a0b Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Tue, 30 Apr 2024 21:29:21 -0700 Subject: [PATCH 04/27] logging --- fairscale/nn/misc/flatten_params_wrapper.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index bc1209bbb..30b88360d 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -372,7 +372,7 @@ def _unflatten_params_as_views(self) -> None: self.flat_param unchanged. """ assert self.is_flattened - #logger.info(f"CHRISLOG: {self._require_backward_grad_sync=}") + # logger.info(f"CHRISLOG: {self._require_backward_grad_sync=}") if self._require_backward_grad_sync: #logger.info("CHRISLOG: calling self.get_param_views() without torch.no_grad()") ps = self.get_param_views() @@ -385,7 +385,7 @@ def _unflatten_params_as_views(self) -> None: for (_, m, n), p in zip(self._param_infos, ps): setattr(p, '_fsdp_weight', True) setattr(m, n, p) # This will set as plain attr - #logger.info(f"CHRISLOG: {n=}, {p.requires_grad=}") + # logger.info(f"CHRISLOG: {n=}, {p.requires_grad=}, {p.grad_fn=}") param_views.append(p) # Save param views for easy access if anyone still wants to access From 4b5abe2541be244f1efe069c41ce902606073ea5 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Wed, 1 May 2024 18:17:48 -0700 Subject: [PATCH 05/27] use new field to accumulate per-parameter grads in fp32 and copy into flatten_parameter.unsharded_main_grad in last microbatch backward() --- .../fully_sharded_data_parallel.py | 17 +++++++---- fairscale/nn/misc/flatten_params_wrapper.py | 29 ++++++++++++++++++- 2 files changed, 40 insertions(+), 6 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 48c1bd0ff..27c39f66a 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1717,11 +1717,18 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: self._use_fp32_param_shard([param]) if self.fp32_reduce_scatter: - if getattr(param, "unsharded_main_grad", None) is None: - param.unsharded_main_grad = param.grad.to(torch.float32) - else: - param.unsharded_main_grad.add_(param.grad.data) - + #logger.info(f"CHRISLOG:{param.unsharded_main_grad.size()=}") + # logger.info(f"CHRISLOG:{len(self._fsdp_wrapped_module.fp32_grads)=}") + # grad_sizes = [grad.size() for grad in self._fsdp_wrapped_module.fp32_grads] + # logger.info(f"CHRISLOG:{grad_sizes=}") + + new_unsharded_main_grad_in_fp32 = torch.cat([grad.flatten() for grad in self._fsdp_wrapped_module.fp32_grads]) + logger.info(f"CHRISLOG: assigning new unsharded_main_grad with size {new_unsharded_main_grad_in_fp32.size()}, type:{new_unsharded_main_grad_in_fp32.dtype}, original grad size {param.grad.size()}") + # if getattr(param, "unsharded_main_grad", None) is None: + # param.unsharded_main_grad = param.grad.to(torch.float32) + # else: + # param.unsharded_main_grad.add_(param.grad.data) + param.unsharded_main_grad = new_unsharded_main_grad_in_fp32 param.grad = None if not self._require_backward_grad_sync: diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 30b88360d..613869364 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -92,6 +92,9 @@ def get_param_views(self, external_data: Optional[Tensor] = None) -> Iterator[Te raise ValueError( f"Incorrect numel of supplied data: got {data.numel()} but expected {sum(self._param_numels)}" ) + logger.info(f"CHRISLOG: {data.numel()=}") + logger.info(f"CHRISLOG: {self._param_numels=}") + logger.info(f"CHRISLOG: {self._param_shapes=}") return (t.view(s) for (t, s) in zip(data.split(self._param_numels), self._param_shapes)) def metadata(self) -> Tuple[List[str], List[torch.Size], List[int]]: @@ -164,6 +167,7 @@ def __init__( self._fpw_module = module self.is_flattened = False self._require_backward_grad_sync = True + self.fp32_grads = [] # Handle param_list being None. if param_list is None: @@ -367,6 +371,19 @@ def _unflatten_params(self, external_data: Optional[List[Optional[Tensor]]] = No delattr(self, n) self.flat_params = [] + + def _hook( + self, + grad, + param_index, + ): + logger.info(f"CHRISLOG: before post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") + if self.fp32_grads[param_index] is None: + self.fp32_grads[param_index] = grad.to(torch.float32) + else: + self.fp32_grads[param_index].add_(grad.data) + logger.info(f"CHRISLOG: after post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") + def _unflatten_params_as_views(self) -> None: """Unlike ``_unflatten_params``, this function unflatten into views and keep self.flat_param unchanged. @@ -385,8 +402,18 @@ def _unflatten_params_as_views(self) -> None: for (_, m, n), p in zip(self._param_infos, ps): setattr(p, '_fsdp_weight', True) setattr(m, n, p) # This will set as plain attr - # logger.info(f"CHRISLOG: {n=}, {p.requires_grad=}, {p.grad_fn=}") + #logger.info(f"CHRISLOG: {n=}, {p.requires_grad=}, {p.grad_fn=}, {p.grad=}") + + import functools + p.register_hook( + functools.partial( + self._hook, + param_index=len(param_views) - 1 + ) + ) param_views.append(p) + if len(self.fp32_grads) == 0: + self.fp32_grads = [None] * len(param_views) # Save param views for easy access if anyone still wants to access # parameters of the module. From c97bfd91a927f1474442494aabd190f1410ef1e1 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Wed, 1 May 2024 18:40:10 -0700 Subject: [PATCH 06/27] clean up accumulated fp32 grads between data batches --- fairscale/nn/data_parallel/fully_sharded_data_parallel.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 27c39f66a..48eac72ee 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1729,6 +1729,8 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # else: # param.unsharded_main_grad.add_(param.grad.data) param.unsharded_main_grad = new_unsharded_main_grad_in_fp32 + # Clean up accumulated grads between data batches + self._fsdp_wrapped_module.fp32_grads = [] param.grad = None if not self._require_backward_grad_sync: From d2a88b730958a2fa85dfaaf7f53c90416746edc5 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Thu, 2 May 2024 15:21:24 -0700 Subject: [PATCH 07/27] logging --- .../nn/data_parallel/fully_sharded_data_parallel.py | 2 +- fairscale/nn/misc/flatten_params_wrapper.py | 10 +++++----- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 48eac72ee..58f59b1f3 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1723,7 +1723,7 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # logger.info(f"CHRISLOG:{grad_sizes=}") new_unsharded_main_grad_in_fp32 = torch.cat([grad.flatten() for grad in self._fsdp_wrapped_module.fp32_grads]) - logger.info(f"CHRISLOG: assigning new unsharded_main_grad with size {new_unsharded_main_grad_in_fp32.size()}, type:{new_unsharded_main_grad_in_fp32.dtype}, original grad size {param.grad.size()}") + # logger.info(f"CHRISLOG: assigning new unsharded_main_grad with size {new_unsharded_main_grad_in_fp32.size()}, type:{new_unsharded_main_grad_in_fp32.dtype}, original grad size {param.grad.size()}") # if getattr(param, "unsharded_main_grad", None) is None: # param.unsharded_main_grad = param.grad.to(torch.float32) # else: diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 613869364..a7b3a09e1 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -92,9 +92,9 @@ def get_param_views(self, external_data: Optional[Tensor] = None) -> Iterator[Te raise ValueError( f"Incorrect numel of supplied data: got {data.numel()} but expected {sum(self._param_numels)}" ) - logger.info(f"CHRISLOG: {data.numel()=}") - logger.info(f"CHRISLOG: {self._param_numels=}") - logger.info(f"CHRISLOG: {self._param_shapes=}") + # logger.info(f"CHRISLOG: {data.numel()=}") + # logger.info(f"CHRISLOG: {self._param_numels=}") + # logger.info(f"CHRISLOG: {self._param_shapes=}") return (t.view(s) for (t, s) in zip(data.split(self._param_numels), self._param_shapes)) def metadata(self) -> Tuple[List[str], List[torch.Size], List[int]]: @@ -377,12 +377,12 @@ def _hook( grad, param_index, ): - logger.info(f"CHRISLOG: before post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") + #logger.info(f"CHRISLOG: before post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") if self.fp32_grads[param_index] is None: self.fp32_grads[param_index] = grad.to(torch.float32) else: self.fp32_grads[param_index].add_(grad.data) - logger.info(f"CHRISLOG: after post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") + #logger.info(f"CHRISLOG: after post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") def _unflatten_params_as_views(self) -> None: """Unlike ``_unflatten_params``, this function unflatten into views and keep From 901fb86d2c6557de7858de4f714162970196bf20 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Sun, 5 May 2024 22:15:45 -0700 Subject: [PATCH 08/27] logging --- fairscale/nn/misc/flatten_params_wrapper.py | 43 +++++++++++++++++---- 1 file changed, 36 insertions(+), 7 deletions(-) diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index a7b3a09e1..0e87bd8be 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -7,7 +7,9 @@ # Licensed under the MIT License. from contextlib import contextmanager +import functools from itertools import chain +from re import split import tempfile import typing from typing import ( @@ -81,7 +83,7 @@ def __init__(self, params: Sequence[nn.Parameter], requires_grad: bool = True): self._param_infos: List[Tuple[str, nn.Module, str]] = [] self._shared_param_infos: List[Tuple[str, str, nn.Module, str, nn.Module, str]] = [] - def get_param_views(self, external_data: Optional[Tensor] = None) -> Iterator[Tensor]: + def get_param_views(self, require_backward_grad_sync, external_data: Optional[Tensor] = None) -> Iterator[Tensor]: """Return a generator of views that map to the original parameters.""" # Note, self.data could be sharded, so its numel is <= to the sum. assert self.data.numel() <= sum( @@ -95,7 +97,34 @@ def get_param_views(self, external_data: Optional[Tensor] = None) -> Iterator[Te # logger.info(f"CHRISLOG: {data.numel()=}") # logger.info(f"CHRISLOG: {self._param_numels=}") # logger.info(f"CHRISLOG: {self._param_shapes=}") - return (t.view(s) for (t, s) in zip(data.split(self._param_numels), self._param_shapes)) + + # logger.info(f"CHRISLOG: {data.is_leaf=}, {data.grad_fn=}") + + # def post_accumulate_grad_hook( + # param + # ): + # logger.info(f"CHRISLOG: cleaning up {param.grad=}") + # param.grad = None + + # data.register_post_accumulate_grad_hook( + # functools.partial( + # post_accumulate_grad_hook + # ) + # ) + # logger.info("CHRISLOG: registered post_accumulate_grad_hook for bf16 grad cleanup on data") + + split_outputs = data.split(self._param_numels) + # for split_output in split_outputs: + # logger.info(f"CHRISLOG: {require_backward_grad_sync=} {split_output.is_leaf=}, {split_output.grad_fn=}, {split_output.grad=}") # + # if not require_backward_grad_sync: + # split_output.register_hook( + # functools.partial( + # post_accumulate_grad_hook + # ) + # ) + # logger.info("CHRISLOG: registered post_accumulate_grad_hook for bf16 grad cleanup on split_output") + + return (t.view(s) for (t, s) in zip(split_outputs, self._param_shapes)) def metadata(self) -> Tuple[List[str], List[torch.Size], List[int]]: """Return tuple of (names, shapes, numels) metadata for this flat parameter.""" @@ -382,6 +411,7 @@ def _hook( self.fp32_grads[param_index] = grad.to(torch.float32) else: self.fp32_grads[param_index].add_(grad.data) + #logger.info(f"CHRISLOG: after post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") def _unflatten_params_as_views(self) -> None: @@ -392,18 +422,17 @@ def _unflatten_params_as_views(self) -> None: # logger.info(f"CHRISLOG: {self._require_backward_grad_sync=}") if self._require_backward_grad_sync: #logger.info("CHRISLOG: calling self.get_param_views() without torch.no_grad()") - ps = self.get_param_views() + ps = self.get_param_views(require_backward_grad_sync=self._require_backward_grad_sync) else: with torch.no_grad(): #logger.info("CHRISLOG: calling self.get_param_views() with torch.no_grad()") - ps = self.get_param_views() + ps = self.get_param_views(require_backward_grad_sync=self._require_backward_grad_sync) param_views = [] for (_, m, n), p in zip(self._param_infos, ps): setattr(p, '_fsdp_weight', True) setattr(m, n, p) # This will set as plain attr #logger.info(f"CHRISLOG: {n=}, {p.requires_grad=}, {p.grad_fn=}, {p.grad=}") - import functools p.register_hook( functools.partial( @@ -528,7 +557,7 @@ def forward(self, *inputs: Any, **kwinputs: Any) -> Any: self._unflatten_params_as_views() return self.module(*inputs, **kwinputs) - def get_param_views(self, external_data_list: Optional[List[Optional[Tensor]]] = None) -> Iterator[Tensor]: + def get_param_views(self, require_backward_grad_sync, external_data_list: Optional[List[Optional[Tensor]]] = None) -> Iterator[Tensor]: """Used to get a generator over all views from a list of external data list.""" params = self.flat_params if external_data_list is None: @@ -539,7 +568,7 @@ def get_param_views(self, external_data_list: Optional[List[Optional[Tensor]]] = gens = [] for p, data in zip(params, external_data_list): - gens.append(p.get_param_views(data)) + gens.append(p.get_param_views(require_backward_grad_sync, data)) return chain(*gens) From ad40f246e641b62955aa5ae4a5ce36836fcb84ba Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Tue, 7 May 2024 20:48:10 -0700 Subject: [PATCH 09/27] return grad in post_backward_hook() --- fairscale/nn/misc/flatten_params_wrapper.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 0e87bd8be..d6af249c6 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -411,8 +411,8 @@ def _hook( self.fp32_grads[param_index] = grad.to(torch.float32) else: self.fp32_grads[param_index].add_(grad.data) - #logger.info(f"CHRISLOG: after post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") + return grad def _unflatten_params_as_views(self) -> None: """Unlike ``_unflatten_params``, this function unflatten into views and keep From 14499fedd29f14d8d9fc4c5576b736b1d524345a Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Thu, 9 May 2024 00:25:01 -0700 Subject: [PATCH 10/27] correct param_index --- fairscale/nn/misc/flatten_params_wrapper.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index d6af249c6..fb75ef739 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -406,12 +406,12 @@ def _hook( grad, param_index, ): - #logger.info(f"CHRISLOG: before post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") + logger.info(f"CHRISLOG: {param_index=} before post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") if self.fp32_grads[param_index] is None: self.fp32_grads[param_index] = grad.to(torch.float32) else: self.fp32_grads[param_index].add_(grad.data) - #logger.info(f"CHRISLOG: after post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") + logger.info(f"CHRISLOG: {param_index=} after post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") return grad def _unflatten_params_as_views(self) -> None: @@ -434,10 +434,12 @@ def _unflatten_params_as_views(self) -> None: setattr(m, n, p) # This will set as plain attr #logger.info(f"CHRISLOG: {n=}, {p.requires_grad=}, {p.grad_fn=}, {p.grad=}") import functools + param_index = len(param_views) + logger.info(f"CHRISLOG: {param_index=}") p.register_hook( functools.partial( self._hook, - param_index=len(param_views) - 1 + param_index=param_index ) ) param_views.append(p) From ad7aa1fe6f3b96e7967a86cae92dd1540207ab4d Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Thu, 9 May 2024 00:34:04 -0700 Subject: [PATCH 11/27] logging --- fairscale/nn/data_parallel/fully_sharded_data_parallel.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 58f59b1f3..7fe9d569d 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1723,6 +1723,14 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # logger.info(f"CHRISLOG:{grad_sizes=}") new_unsharded_main_grad_in_fp32 = torch.cat([grad.flatten() for grad in self._fsdp_wrapped_module.fp32_grads]) + baseline_grad = param.grad.to(torch.float32) + + + logger.info(f"CHRISLOG: baseline grad {baseline_grad=}, {baseline_grad.size()=}") + logger.info(f"CHRISLOG: new grad {new_unsharded_main_grad_in_fp32=}, {new_unsharded_main_grad_in_fp32.size()=}") + torch.allclose(baseline_grad, new_unsharded_main_grad_in_fp32, atol=0, rtol=0) + logger.info(f"CHRISLOG: baseline grad and new grad passed allclose check") + # logger.info(f"CHRISLOG: assigning new unsharded_main_grad with size {new_unsharded_main_grad_in_fp32.size()}, type:{new_unsharded_main_grad_in_fp32.dtype}, original grad size {param.grad.size()}") # if getattr(param, "unsharded_main_grad", None) is None: # param.unsharded_main_grad = param.grad.to(torch.float32) From b835770df0ca5d52bde2619a5981df368f45ec6a Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Thu, 9 May 2024 00:45:31 -0700 Subject: [PATCH 12/27] add torch.testing.assert_allclose() to compare baseline and new grads --- .../nn/data_parallel/fully_sharded_data_parallel.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 7fe9d569d..ef701c5e5 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1726,9 +1726,9 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: baseline_grad = param.grad.to(torch.float32) - logger.info(f"CHRISLOG: baseline grad {baseline_grad=}, {baseline_grad.size()=}") - logger.info(f"CHRISLOG: new grad {new_unsharded_main_grad_in_fp32=}, {new_unsharded_main_grad_in_fp32.size()=}") - torch.allclose(baseline_grad, new_unsharded_main_grad_in_fp32, atol=0, rtol=0) + logger.info(f"CHRISLOG: baseline grad {baseline_grad=}, {baseline_grad.size()=}, {baseline_grad.dtype=}") + logger.info(f"CHRISLOG: new grad {new_unsharded_main_grad_in_fp32=}, {new_unsharded_main_grad_in_fp32.size()=}, {new_unsharded_main_grad_in_fp32.dtype=}") + torch.testing.assert_allclose(baseline_grad, new_unsharded_main_grad_in_fp32, atol=0, rtol=0) logger.info(f"CHRISLOG: baseline grad and new grad passed allclose check") # logger.info(f"CHRISLOG: assigning new unsharded_main_grad with size {new_unsharded_main_grad_in_fp32.size()}, type:{new_unsharded_main_grad_in_fp32.dtype}, original grad size {param.grad.size()}") @@ -1737,6 +1737,9 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # else: # param.unsharded_main_grad.add_(param.grad.data) param.unsharded_main_grad = new_unsharded_main_grad_in_fp32 + logger.info(f"CHRISLOG: {param.unsharded_main_grad.dtype=}") + + # Clean up accumulated grads between data batches self._fsdp_wrapped_module.fp32_grads = [] param.grad = None From d689f38b67ec146a35c75f674c1b50de610d920e Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Thu, 9 May 2024 12:23:41 -0700 Subject: [PATCH 13/27] logging --- .../nn/data_parallel/fully_sharded_data_parallel.py | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index ef701c5e5..c2457cc8a 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1723,13 +1723,12 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # logger.info(f"CHRISLOG:{grad_sizes=}") new_unsharded_main_grad_in_fp32 = torch.cat([grad.flatten() for grad in self._fsdp_wrapped_module.fp32_grads]) - baseline_grad = param.grad.to(torch.float32) - - logger.info(f"CHRISLOG: baseline grad {baseline_grad=}, {baseline_grad.size()=}, {baseline_grad.dtype=}") - logger.info(f"CHRISLOG: new grad {new_unsharded_main_grad_in_fp32=}, {new_unsharded_main_grad_in_fp32.size()=}, {new_unsharded_main_grad_in_fp32.dtype=}") - torch.testing.assert_allclose(baseline_grad, new_unsharded_main_grad_in_fp32, atol=0, rtol=0) - logger.info(f"CHRISLOG: baseline grad and new grad passed allclose check") + # baseline_grad = param.grad.to(torch.float32) + # logger.info(f"CHRISLOG: baseline grad {baseline_grad=}, {baseline_grad.size()=}, {baseline_grad.dtype=}") + # logger.info(f"CHRISLOG: new grad {new_unsharded_main_grad_in_fp32=}, {new_unsharded_main_grad_in_fp32.size()=}, {new_unsharded_main_grad_in_fp32.dtype=}") + # torch.testing.assert_allclose(baseline_grad, new_unsharded_main_grad_in_fp32, atol=0, rtol=0) + # logger.info(f"CHRISLOG: baseline grad and new grad passed allclose check") # logger.info(f"CHRISLOG: assigning new unsharded_main_grad with size {new_unsharded_main_grad_in_fp32.size()}, type:{new_unsharded_main_grad_in_fp32.dtype}, original grad size {param.grad.size()}") # if getattr(param, "unsharded_main_grad", None) is None: From e8df583cfe03ba4c411a928c22b5af14d467ab90 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Sun, 12 May 2024 22:57:22 -0700 Subject: [PATCH 14/27] logging --- .../data_parallel/fully_sharded_data_parallel.py | 14 -------------- fairscale/nn/misc/flatten_params_wrapper.py | 2 +- 2 files changed, 1 insertion(+), 15 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index c2457cc8a..6f1398b6c 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1717,27 +1717,13 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: self._use_fp32_param_shard([param]) if self.fp32_reduce_scatter: - #logger.info(f"CHRISLOG:{param.unsharded_main_grad.size()=}") - # logger.info(f"CHRISLOG:{len(self._fsdp_wrapped_module.fp32_grads)=}") - # grad_sizes = [grad.size() for grad in self._fsdp_wrapped_module.fp32_grads] - # logger.info(f"CHRISLOG:{grad_sizes=}") - new_unsharded_main_grad_in_fp32 = torch.cat([grad.flatten() for grad in self._fsdp_wrapped_module.fp32_grads]) - - # baseline_grad = param.grad.to(torch.float32) - # logger.info(f"CHRISLOG: baseline grad {baseline_grad=}, {baseline_grad.size()=}, {baseline_grad.dtype=}") - # logger.info(f"CHRISLOG: new grad {new_unsharded_main_grad_in_fp32=}, {new_unsharded_main_grad_in_fp32.size()=}, {new_unsharded_main_grad_in_fp32.dtype=}") - # torch.testing.assert_allclose(baseline_grad, new_unsharded_main_grad_in_fp32, atol=0, rtol=0) - # logger.info(f"CHRISLOG: baseline grad and new grad passed allclose check") - # logger.info(f"CHRISLOG: assigning new unsharded_main_grad with size {new_unsharded_main_grad_in_fp32.size()}, type:{new_unsharded_main_grad_in_fp32.dtype}, original grad size {param.grad.size()}") # if getattr(param, "unsharded_main_grad", None) is None: # param.unsharded_main_grad = param.grad.to(torch.float32) # else: # param.unsharded_main_grad.add_(param.grad.data) param.unsharded_main_grad = new_unsharded_main_grad_in_fp32 - logger.info(f"CHRISLOG: {param.unsharded_main_grad.dtype=}") - # Clean up accumulated grads between data batches self._fsdp_wrapped_module.fp32_grads = [] diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index fb75ef739..f6a261551 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -435,7 +435,7 @@ def _unflatten_params_as_views(self) -> None: #logger.info(f"CHRISLOG: {n=}, {p.requires_grad=}, {p.grad_fn=}, {p.grad=}") import functools param_index = len(param_views) - logger.info(f"CHRISLOG: {param_index=}") + #logger.info(f"CHRISLOG: {param_index=}") p.register_hook( functools.partial( self._hook, From 5926a79b15b7cb66729f3b247284b78cfbf8df67 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Tue, 14 May 2024 22:48:56 -0700 Subject: [PATCH 15/27] honor optimize_backward_concat flag --- .../fully_sharded_data_parallel.py | 49 +++++++++++++------ fairscale/nn/misc/flatten_params_wrapper.py | 43 ++++++++-------- 2 files changed, 56 insertions(+), 36 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 6f1398b6c..927fe39e7 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -369,6 +369,7 @@ def __init__( limit_all_gather_events: bool = False, limit_reduce_scatter_events: bool = False, should_validate_process_group: bool = True, + optimize_backward_concat: bool = False, ): try: import torch._C @@ -493,8 +494,12 @@ def __init__( param_name_groups = [param_names] del param_names + self.optimize_backward_concat = optimize_backward_concat + if self.optimize_backward_concat: + assert self.fp32_reduce_scatter, f"{optimize_backward_concat=} requires self.fp32_reduce_scatter=True" + self._fsdp_wrapped_module: nn.Module = FlattenParamsWrapper( - module, param_list=to_be_flatten_params, ssd_offload=self.ssd_offload, ssd_directory=self.ssd_directory + module, param_list=to_be_flatten_params, ssd_offload=self.ssd_offload, ssd_directory=self.ssd_directory, optimize_backward_concat=self.optimize_backward_concat, ) del module # free original module in case it helps garbage collection @@ -851,6 +856,7 @@ def extra_repr(self) -> str: f"bucket_cap_mb={self.bucket_cap_mb}, " f"clear_autocast_cache={self.clear_autocast_cache}" f"force_input_to_fp32={self.force_input_to_fp32}" + f"optimize_backward_concat={self.optimize_backward_concat}" ) return repr @@ -1717,16 +1723,17 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: self._use_fp32_param_shard([param]) if self.fp32_reduce_scatter: - new_unsharded_main_grad_in_fp32 = torch.cat([grad.flatten() for grad in self._fsdp_wrapped_module.fp32_grads]) - # logger.info(f"CHRISLOG: assigning new unsharded_main_grad with size {new_unsharded_main_grad_in_fp32.size()}, type:{new_unsharded_main_grad_in_fp32.dtype}, original grad size {param.grad.size()}") - # if getattr(param, "unsharded_main_grad", None) is None: - # param.unsharded_main_grad = param.grad.to(torch.float32) - # else: - # param.unsharded_main_grad.add_(param.grad.data) - param.unsharded_main_grad = new_unsharded_main_grad_in_fp32 - - # Clean up accumulated grads between data batches - self._fsdp_wrapped_module.fp32_grads = [] + + if self.optimize_backward_concat: + param.unsharded_main_grad = torch.cat([grad.flatten() for grad in self._fsdp_wrapped_module.fp32_grads]) + # Clean up accumulated grads between data batches + self._fsdp_wrapped_module.fp32_grads = [] + else: + if getattr(param, "unsharded_main_grad", None) is None: + param.unsharded_main_grad = param.grad.to(torch.float32) + else: + param.unsharded_main_grad.add_(param.grad.data) + param.grad = None if not self._require_backward_grad_sync: @@ -1857,8 +1864,14 @@ def _wait_for_post_backward(self) -> None: # the `requires_grad` field set. If `requires_grad=False` for # all the params, the post_backward hook will not fire and the # state will remain in `TrainingState.BACKWARD_PRE`. - if any([p.requires_grad for p in self.params]) and self._fsdp_wrapped_module._require_backward_grad_sync: - self.assert_state(TrainingState.BACKWARD_POST) + if any([p.requires_grad for p in self.params]): + if self.optimize_backward_concat: + if self._fsdp_wrapped_module._require_backward_grad_sync: + self.assert_state(TrainingState.BACKWARD_POST) + else: + self.assert_state(TrainingState.BACKWARD_PRE) + else: + self.assert_state(TrainingState.BACKWARD_POST) else: self.assert_state(TrainingState.BACKWARD_PRE) @@ -1934,8 +1947,14 @@ def _finalize_parameters(fsdp_module: FullyShardedDataParallel) -> None: # the `requires_grad` field set. If `requires_grad=False` for # all the params, the post_backward hook will not fire and the # state will remain in `TrainingState.BACKWARD_PRE`. - if any([p.requires_grad for p in m.params]) and self._fsdp_wrapped_module._require_backward_grad_sync: - m.assert_state(TrainingState.BACKWARD_POST) + if any([p.requires_grad for p in m.params]): + if self.optimize_backward_concat: + if self._fsdp_wrapped_module._require_backward_grad_sync: + m.assert_state(TrainingState.BACKWARD_POST) + else: + m.assert_state(TrainingState.BACKWARD_PRE) + else: + m.assert_state(TrainingState.BACKWARD_POST) else: m.assert_state(TrainingState.BACKWARD_PRE) else: diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index f6a261551..84b67f977 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -35,6 +35,7 @@ from fairscale.experimental.nn.ssd_offload import SsdFlatParameter from fairscale.utils.state_dict import replace_by_prefix_ +import functools if TYPE_CHECKING: from collections import OrderedDict # noqa: F401 @@ -191,10 +192,13 @@ def __init__( flat_param_names: Optional[List[str]] = None, ssd_offload: bool = False, ssd_directory: str = "", + optimize_backward_concat: bool = False, ): super().__init__() self._fpw_module = module self.is_flattened = False + + self.optimize_backward_concat = optimize_backward_concat self._require_backward_grad_sync = True self.fp32_grads = [] @@ -401,17 +405,15 @@ def _unflatten_params(self, external_data: Optional[List[Optional[Tensor]]] = No self.flat_params = [] - def _hook( + def _grad_accumulation_hook( self, grad, param_index, ): - logger.info(f"CHRISLOG: {param_index=} before post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") if self.fp32_grads[param_index] is None: self.fp32_grads[param_index] = grad.to(torch.float32) else: self.fp32_grads[param_index].add_(grad.data) - logger.info(f"CHRISLOG: {param_index=} after post-backward hook, self.fp32_grads[param_index] is None: {self.fp32_grads[param_index] is None}") return grad def _unflatten_params_as_views(self) -> None: @@ -419,31 +421,30 @@ def _unflatten_params_as_views(self) -> None: self.flat_param unchanged. """ assert self.is_flattened - # logger.info(f"CHRISLOG: {self._require_backward_grad_sync=}") - if self._require_backward_grad_sync: - #logger.info("CHRISLOG: calling self.get_param_views() without torch.no_grad()") - ps = self.get_param_views(require_backward_grad_sync=self._require_backward_grad_sync) + if self.optimize_backward_concat: + if self._require_backward_grad_sync: + ps = self.get_param_views() + else: + with torch.no_grad(): + ps = self.get_param_views() else: - with torch.no_grad(): - #logger.info("CHRISLOG: calling self.get_param_views() with torch.no_grad()") - ps = self.get_param_views(require_backward_grad_sync=self._require_backward_grad_sync) + ps = self.get_param_views() param_views = [] for (_, m, n), p in zip(self._param_infos, ps): setattr(p, '_fsdp_weight', True) setattr(m, n, p) # This will set as plain attr - #logger.info(f"CHRISLOG: {n=}, {p.requires_grad=}, {p.grad_fn=}, {p.grad=}") - import functools param_index = len(param_views) - #logger.info(f"CHRISLOG: {param_index=}") - p.register_hook( - functools.partial( - self._hook, - param_index=param_index + if self.optimize_backward_concat: + p.register_hook( + functools.partial( + self._grad_accumulation_hook, + param_index=param_index + ) ) - ) param_views.append(p) - if len(self.fp32_grads) == 0: + + if self.optimize_backward_concat and len(self.fp32_grads) == 0: self.fp32_grads = [None] * len(param_views) # Save param views for easy access if anyone still wants to access @@ -559,7 +560,7 @@ def forward(self, *inputs: Any, **kwinputs: Any) -> Any: self._unflatten_params_as_views() return self.module(*inputs, **kwinputs) - def get_param_views(self, require_backward_grad_sync, external_data_list: Optional[List[Optional[Tensor]]] = None) -> Iterator[Tensor]: + def get_param_views(self, external_data_list: Optional[List[Optional[Tensor]]] = None) -> Iterator[Tensor]: """Used to get a generator over all views from a list of external data list.""" params = self.flat_params if external_data_list is None: @@ -570,7 +571,7 @@ def get_param_views(self, require_backward_grad_sync, external_data_list: Option gens = [] for p, data in zip(params, external_data_list): - gens.append(p.get_param_views(require_backward_grad_sync, data)) + gens.append(p.get_param_views(data)) return chain(*gens) From 5d08aa3f0b1c7f1641759b7e061f705a77d222d3 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Tue, 14 May 2024 23:55:43 -0700 Subject: [PATCH 16/27] documentation --- .../fully_sharded_data_parallel.py | 19 +++++-- fairscale/nn/misc/flatten_params_wrapper.py | 50 ++++++------------- 2 files changed, 31 insertions(+), 38 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 927fe39e7..35b88c67c 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1105,14 +1105,20 @@ def no_sync(self) -> Generator: if isinstance(m, FullyShardedDataParallel): old_flags.append((m, m._require_backward_grad_sync)) m._require_backward_grad_sync = False - m._fsdp_wrapped_module._require_backward_grad_sync = False + if self.optimize_backward_concat: + # Set the flag on the wrapped FlattenParamsWrapper module as well, + # so that FlattenParamsWrapper could accumulate grads at corresponding + # leaf nodes without triggering concat operations when gradient + # synchronization is not needed. + m._fsdp_wrapped_module._require_backward_grad_sync = False try: yield finally: for m, old_flag in old_flags: assert m._require_backward_grad_sync is False m._require_backward_grad_sync = old_flag - m._fsdp_wrapped_module._require_backward_grad_sync = old_flag + if self.optimize_backward_concat: + m._fsdp_wrapped_module._require_backward_grad_sync = old_flag @contextlib.contextmanager def summon_full_params(self, recurse: bool = True, volatile: bool = False) -> Generator: @@ -1723,8 +1729,9 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: self._use_fp32_param_shard([param]) if self.fp32_reduce_scatter: - if self.optimize_backward_concat: + # Flatten and concat the accumulated fp32 grads + # and assign them to param.unsharded_main_grad param.unsharded_main_grad = torch.cat([grad.flatten() for grad in self._fsdp_wrapped_module.fp32_grads]) # Clean up accumulated grads between data batches self._fsdp_wrapped_module.fp32_grads = [] @@ -1866,6 +1873,9 @@ def _wait_for_post_backward(self) -> None: # state will remain in `TrainingState.BACKWARD_PRE`. if any([p.requires_grad for p in self.params]): if self.optimize_backward_concat: + # If self.optimize_backward_concat==True, FSDP backward should + # only be triggered (which will invoke concat()) + # when self._fsdp_wrapped_module._require_backward_grad_sync = True if self._fsdp_wrapped_module._require_backward_grad_sync: self.assert_state(TrainingState.BACKWARD_POST) else: @@ -1949,6 +1959,9 @@ def _finalize_parameters(fsdp_module: FullyShardedDataParallel) -> None: # state will remain in `TrainingState.BACKWARD_PRE`. if any([p.requires_grad for p in m.params]): if self.optimize_backward_concat: + # If self.optimize_backward_concat==True, FSDP backward should + # only be triggered (which will invoke concat()) + # when self._fsdp_wrapped_module._require_backward_grad_sync = True if self._fsdp_wrapped_module._require_backward_grad_sync: m.assert_state(TrainingState.BACKWARD_POST) else: diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 84b67f977..d55ee6704 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -9,7 +9,6 @@ from contextlib import contextmanager import functools from itertools import chain -from re import split import tempfile import typing from typing import ( @@ -35,14 +34,10 @@ from fairscale.experimental.nn.ssd_offload import SsdFlatParameter from fairscale.utils.state_dict import replace_by_prefix_ -import functools if TYPE_CHECKING: from collections import OrderedDict # noqa: F401 -from logging import getLogger -logger = getLogger() - class FlatParameter(nn.Parameter): """A parameter that is initialized from a list of parameters and can be turned into a list of views as needed. @@ -95,36 +90,8 @@ def get_param_views(self, require_backward_grad_sync, external_data: Optional[Te raise ValueError( f"Incorrect numel of supplied data: got {data.numel()} but expected {sum(self._param_numels)}" ) - # logger.info(f"CHRISLOG: {data.numel()=}") - # logger.info(f"CHRISLOG: {self._param_numels=}") - # logger.info(f"CHRISLOG: {self._param_shapes=}") - - # logger.info(f"CHRISLOG: {data.is_leaf=}, {data.grad_fn=}") - - # def post_accumulate_grad_hook( - # param - # ): - # logger.info(f"CHRISLOG: cleaning up {param.grad=}") - # param.grad = None - - # data.register_post_accumulate_grad_hook( - # functools.partial( - # post_accumulate_grad_hook - # ) - # ) - # logger.info("CHRISLOG: registered post_accumulate_grad_hook for bf16 grad cleanup on data") split_outputs = data.split(self._param_numels) - # for split_output in split_outputs: - # logger.info(f"CHRISLOG: {require_backward_grad_sync=} {split_output.is_leaf=}, {split_output.grad_fn=}, {split_output.grad=}") # - # if not require_backward_grad_sync: - # split_output.register_hook( - # functools.partial( - # post_accumulate_grad_hook - # ) - # ) - # logger.info("CHRISLOG: registered post_accumulate_grad_hook for bf16 grad cleanup on split_output") - return (t.view(s) for (t, s) in zip(split_outputs, self._param_shapes)) def metadata(self) -> Tuple[List[str], List[torch.Size], List[int]]: @@ -183,6 +150,11 @@ class FlattenParamsWrapper(nn.Module): flat_param_names (Optional[List[str]]): originally, give each flat_param a unique name. Note a "flat_param_" prefix will be added to those names. + optimize_backward_concat (bool): + If True, only trigger the self.flat_params backward(), which will + invoke the parent FSDP module's _post_backward_hook() and concat() op, + when self._require_backward_grad_sync is True (e.g. last microbatch) + NOTE: this likely will incur more GPU memory usage """ def __init__( @@ -197,9 +169,12 @@ def __init__( super().__init__() self._fpw_module = module self.is_flattened = False - self.optimize_backward_concat = optimize_backward_concat + # If self.optimize_backward_concat == True, used to propagate the + # parent FSDP modules's _require_backward_grad_sync flag self._require_backward_grad_sync = True + # If self.optimize_backward_concat == True, used to accumulate the + # fp32 gradients for the flattened parameters self.fp32_grads = [] # Handle param_list being None. @@ -404,7 +379,7 @@ def _unflatten_params(self, external_data: Optional[List[Optional[Tensor]]] = No delattr(self, n) self.flat_params = [] - + # The post backward hook used to accumulate fp32 gradients def _grad_accumulation_hook( self, grad, @@ -434,8 +409,12 @@ def _unflatten_params_as_views(self) -> None: for (_, m, n), p in zip(self._param_infos, ps): setattr(p, '_fsdp_weight', True) setattr(m, n, p) # This will set as plain attr + # The param_index of p used to accumulate the correspnding + # gradients in self.fp32_grads param_index = len(param_views) if self.optimize_backward_concat: + # Register post backward hook to accumulate the gradients + # in self.fp32_grads p.register_hook( functools.partial( self._grad_accumulation_hook, @@ -445,6 +424,7 @@ def _unflatten_params_as_views(self) -> None: param_views.append(p) if self.optimize_backward_concat and len(self.fp32_grads) == 0: + # Allocate self.fp32_grads at the beginning of each data batch's forward() self.fp32_grads = [None] * len(param_views) # Save param views for easy access if anyone still wants to access From c91cb721a8faf5aba1b3b0d42ac906f3bb8f6e2f Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Wed, 15 May 2024 00:38:34 -0700 Subject: [PATCH 17/27] update documentation --- fairscale/nn/data_parallel/fully_sharded_data_parallel.py | 5 +++++ fairscale/nn/misc/flatten_params_wrapper.py | 6 ++++++ 2 files changed, 11 insertions(+) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 35b88c67c..4126a6142 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -338,6 +338,11 @@ class FullyShardedDataParallel(nn.Module): rank 0 and return empty dict non-rank 0, which allow FullyShardedDataParallel to skip the GPU -> CPU copy on non-rank 0 altogether and prevent OOM. Default: False + optimize_backward_concat (bool): + If True, only trigger the self._fsdp_wrapped_module.flat_params backward(), which will + invoke the _post_backward_hook() and concat() op, + when self._require_backward_grad_sync is True (e.g. last microbatch) + NOTE: this likely will incur more GPU memory usage """ def __init__( diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index d55ee6704..c5ad43bca 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -397,6 +397,12 @@ def _unflatten_params_as_views(self) -> None: """ assert self.is_flattened if self.optimize_backward_concat: + # If self._require_backward_grad_sync == True (e.g. last microbatch), + # we use the original flat_params as autograd leaf nodes and backward + # pass should propagate all the way back to FSDP module and thus invoke + # FSDP post_backward() hook and concat() op + # Otherwise we stop the backward propagation before FSDP module to avoid + # invoking concat() and store the accumulated fp32 grads if self._require_backward_grad_sync: ps = self.get_param_views() else: From fd3f3fc73b981b0dd65b624ee86e3dadeefd1b54 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Wed, 15 May 2024 00:44:29 -0700 Subject: [PATCH 18/27] update documentation --- .../nn/data_parallel/fully_sharded_data_parallel.py | 6 +++--- fairscale/nn/misc/flatten_params_wrapper.py | 12 ++++++------ 2 files changed, 9 insertions(+), 9 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 4126a6142..52f059872 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -339,9 +339,9 @@ class FullyShardedDataParallel(nn.Module): skip the GPU -> CPU copy on non-rank 0 altogether and prevent OOM. Default: False optimize_backward_concat (bool): - If True, only trigger the self._fsdp_wrapped_module.flat_params backward(), which will - invoke the _post_backward_hook() and concat() op, - when self._require_backward_grad_sync is True (e.g. last microbatch) + If True, only let backward pass propagate to self.params, which will + invoke the _post_backward_hook() and concat() op, when self._require_backward_grad_sync + is True (e.g. last microbatch) NOTE: this likely will incur more GPU memory usage """ diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index c5ad43bca..717f1b047 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -151,9 +151,9 @@ class FlattenParamsWrapper(nn.Module): originally, give each flat_param a unique name. Note a "flat_param_" prefix will be added to those names. optimize_backward_concat (bool): - If True, only trigger the self.flat_params backward(), which will - invoke the parent FSDP module's _post_backward_hook() and concat() op, - when self._require_backward_grad_sync is True (e.g. last microbatch) + If True, only let backward pass propagate to the corresponding FSDP.params, which will + invoke the FSDP._post_backward_hook() and concat() op, when _require_backward_grad_sync + is True (e.g. last microbatch) NOTE: this likely will incur more GPU memory usage """ @@ -170,10 +170,10 @@ def __init__( self._fpw_module = module self.is_flattened = False self.optimize_backward_concat = optimize_backward_concat - # If self.optimize_backward_concat == True, used to propagate the - # parent FSDP modules's _require_backward_grad_sync flag + # If optimize_backward_concat == True, used to propagate the + # corresponding FSDP modules's _require_backward_grad_sync flag self._require_backward_grad_sync = True - # If self.optimize_backward_concat == True, used to accumulate the + # If optimize_backward_concat == True, used to accumulate the # fp32 gradients for the flattened parameters self.fp32_grads = [] From 76785034442bcbbfe9f8864caae31af45daf5816 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Wed, 15 May 2024 15:02:05 -0700 Subject: [PATCH 19/27] use grad instead of grad.data --- fairscale/nn/data_parallel/fully_sharded_data_parallel.py | 3 +-- fairscale/nn/misc/flatten_params_wrapper.py | 2 +- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 52f059872..07f5eaacc 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -373,7 +373,6 @@ def __init__( gradient_predivide_factor: Optional[float] = None, limit_all_gather_events: bool = False, limit_reduce_scatter_events: bool = False, - should_validate_process_group: bool = True, optimize_backward_concat: bool = False, ): try: @@ -458,7 +457,7 @@ def __init__( raise ValueError(f"offload type: '{offload_config.offload_type}' requires flatten_parameters=True") # skip validation if the process group was created above - if process_group and should_validate_process_group: + if process_group: validate_process_group(self.compute_device, self.process_group) # enable pytorch sync_bn just in case model contains sync_bn layers. diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 717f1b047..45870a026 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -388,7 +388,7 @@ def _grad_accumulation_hook( if self.fp32_grads[param_index] is None: self.fp32_grads[param_index] = grad.to(torch.float32) else: - self.fp32_grads[param_index].add_(grad.data) + self.fp32_grads[param_index].add_(grad) return grad def _unflatten_params_as_views(self) -> None: From c55a0d16d2f4504cd4464e059338be92f3fcc720 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Wed, 15 May 2024 15:06:24 -0700 Subject: [PATCH 20/27] clean up --- fairscale/nn/misc/flatten_params_wrapper.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 45870a026..dfbdaf60f 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -79,7 +79,7 @@ def __init__(self, params: Sequence[nn.Parameter], requires_grad: bool = True): self._param_infos: List[Tuple[str, nn.Module, str]] = [] self._shared_param_infos: List[Tuple[str, str, nn.Module, str, nn.Module, str]] = [] - def get_param_views(self, require_backward_grad_sync, external_data: Optional[Tensor] = None) -> Iterator[Tensor]: + def get_param_views(self, external_data: Optional[Tensor] = None) -> Iterator[Tensor]: """Return a generator of views that map to the original parameters.""" # Note, self.data could be sharded, so its numel is <= to the sum. assert self.data.numel() <= sum( From 688b90268442c3b3d01d4dceef59460ed763fbf8 Mon Sep 17 00:00:00 2001 From: Andrew Gu Date: Fri, 12 Jan 2024 11:19:17 -0800 Subject: [PATCH 21/27] Added reshard hook for frozen params in backward --- .../fully_sharded_data_parallel.py | 75 ++++++++++++--- .../test_fsdp_freezing_weights.py | 96 +++++++++++++++++++ 2 files changed, 160 insertions(+), 11 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 07f5eaacc..e9419e2a8 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -9,6 +9,7 @@ from dataclasses import dataclass from enum import Enum, auto import functools +import itertools import logging from math import inf import os @@ -47,7 +48,6 @@ from fairscale.utils.containers import apply_to_tensors from fairscale.utils.parallel import ( ProcessGroupName, - chunk_and_pad, enable_pytorch_sync_bn, get_process_group_cached, validate_process_group, @@ -1476,6 +1476,8 @@ def forward(self, *args: Any, **kwargs: Any) -> torch.Tensor: # Register backward hooks to reshard params and reduce-scatter grads. # These need to be re-registered every forward pass. self._register_post_backward_hooks() + self._register_post_backward_reshard_hooks(args, kwargs) + outputs = self.module(*args, **kwargs) if self.reshard_after_forward: @@ -1673,6 +1675,37 @@ def _register_post_backward_hooks(self) -> None: p._shard_bwd_hooks.append((grad_acc, handle)) # p._shard_bwd_hook = (grad_acc, handle) + def _register_post_backward_reshard_hooks( + self, args: Tuple[Any, ...], kwargs: Dict[str, Any] + ) -> None: + if not hasattr(torch.autograd.graph, "register_multi_grad_hook"): + return # unsupported + if not torch.is_grad_enabled(): + return + from torch.utils._pytree import tree_flatten + from torch.autograd.graph import register_multi_grad_hook + # Construct `inp_tensors` lazily to avoid CPU overhead in typical case + # where each parameter requires gradient + inp_tensors: Optional[List[torch.Tensor]] = None + for param in self.params: + # Only register for parameters that do not require gradient + if param.requires_grad: + continue + if inp_tensors is None: + args_list, _ = tree_flatten(args) + kwargs_list, _ = tree_flatten(kwargs) + inp_tensors = [ + obj + for obj in itertools.chain(args_list, kwargs_list) + if torch.is_tensor(obj) and obj.requires_grad + ] + hook_handle = register_multi_grad_hook( + inp_tensors, functools.partial(self._post_backward_reshard_hook, param) + ) + if not hasattr(param, "_shard_bwd_hooks"): + param._shard_bwd_hooks = [] + param._shard_bwd_hooks.append((hook_handle,)) + @torch.no_grad() def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: """ @@ -1715,12 +1748,8 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: if param.grad.requires_grad: raise RuntimeError("FSDP only works with gradients that don't require gradients") - if self._require_backward_grad_sync or self.reshard_after_forward: - # Free full params. As a special case, we don't free the full params - # when in a ``no_sync`` context (as inversely indicated by - # ``self._require_backward_grad_sync``), since the params will not - # get updated before the next forward. This saves networking - # bandwidth but uses more GPU memory. + if self._should_free_in_backward(): + # Free full params. self._free_full_params([param]) if self.mixed_precision: @@ -1854,6 +1883,22 @@ def _post_reduction_hook(self, param: Parameter, reduced_grad: torch.Tensor) -> # Don't let this memory get reused until after the transfer. reduced_grad.data.record_stream(torch.cuda.current_stream()) + @torch.no_grad() + def _post_backward_reshard_hook(self, param: Parameter, *unused: Any) -> None: + if self._should_free_in_backward(): + self._free_full_params([param]) + if self.mixed_precision: + self._free_fp16_param_shard([param]) + self._use_fp32_param_shard([param]) + + def _should_free_in_backward(self): + # As a special case, we don't free the full params + # when in a ``no_sync`` context (as inversely indicated by + # ``self._require_backward_grad_sync``), since the params will not + # get updated before the next forward. This saves networking + # bandwidth but uses more GPU memory. + return self._require_backward_grad_sync or self.reshard_after_forward + def _queue_wait_for_post_backward(self) -> None: """Try to queue a `wait_for_post_backward` callback. @@ -1912,16 +1957,24 @@ def _wait_for_post_backward(self) -> None: def _finalize_parameters(fsdp_module: FullyShardedDataParallel) -> None: """Helper used below on all fsdp modules.""" for p in fsdp_module.params: - if not p.requires_grad: - continue if hasattr(p, "_shard_bwd_hook"): p_assert(len(p._shard_bwd_hook) == 2, f"WFPB: incorrect hook num: {len(p._shard_bwd_hook)}") # p._shard_bwd_hook[1].remove() # delattr(p, "_shard_bwd_hook") if hasattr(p, "_shard_bwd_hooks") and self._require_backward_grad_sync: - for _, handle in p._shard_bwd_hooks: - handle.remove() + for hook_state in p._shard_bwd_hooks: + if len(hook_state) == 1: + hook_state[0].remove() + elif len(hook_state) == 2: + hook_state[1].remove() p._shard_bwd_hooks.clear() + if not p.requires_grad: + # For the 1st layer, if the forward inputs did not require + # gradient, then we cannot run a reshard hook for it, and + # we instead free here. + if p._full_param_padded.untyped_storage().size() > 0: + fsdp_module._post_backward_reshard_hook(p) + continue # Leave the gradient accumulation state as-is if not synchronizing this pass. This ensures p.grad # remains the unsharded gradient accumulated from prior no-sync passes, and p._saved_grad_shard diff --git a/tests/nn/data_parallel/test_fsdp_freezing_weights.py b/tests/nn/data_parallel/test_fsdp_freezing_weights.py index c6ad364f7..7baadc5d9 100644 --- a/tests/nn/data_parallel/test_fsdp_freezing_weights.py +++ b/tests/nn/data_parallel/test_fsdp_freezing_weights.py @@ -12,6 +12,8 @@ from enum import Enum from itertools import product +from unittest import mock +import copy import tempfile import pytest @@ -275,3 +277,97 @@ def test_freezing_weights(temp_files, nested_trunk): nprocs=world_size, ) temp_file_idx += 3 + + +@skip_if_single_gpu +def test_reshard_frozen_weights(): + world_size = 2 + for flatten_parameters, reshard_after_forward, inp_requires_grad in product( + [False, True], [False, True], [False, True] + ): + print( + "Testing FSDP reshard frozen weights with " + f"flatten_parameters={flatten_parameters}, " + f"reshard_after_forward={reshard_after_forward}, " + f"inp_requires_grad={inp_requires_grad}" + ) + mp.spawn( + _distributed_worker_reshard, + (world_size, flatten_parameters, reshard_after_forward, inp_requires_grad), + nprocs=world_size, + ) + + +def _distributed_worker_reshard( + rank: int, + world_size: int, + flatten_parameters: bool, + reshard_after_forward: bool, + inp_requires_grad: bool, +): + import os + os.environ["MASTER_ADDR"] = "localhost" + os.environ["MASTER_PORT"] = "12355" + torch.cuda.set_device(rank) + torch.distributed.init_process_group(backend="nccl", rank=rank, world_size=world_size) + + torch.manual_seed(0) + + num_linears = 6 + modules = [] + for _ in range(num_linears): + modules += [nn.Linear(5, 5, device="cuda"), nn.ReLU()] + model = nn.Sequential(*modules) + # Freeze every other linear + for i in range(num_linears): + if i % 2 == 0: + for param in model[i * 2].parameters(recurse=False): + param.requires_grad = False + num_frozen_linears = num_linears // 2 + + ref_model = DistributedDataParallel(copy.deepcopy(model), device_ids=[rank]) + ref_optim = torch.optim.AdamW(ref_model.parameters(), lr=1e-2) + + for i, module in enumerate(model): + if isinstance(module, nn.Linear): + model[i] = FSDP( + module, + flatten_parameters=flatten_parameters, + reshard_after_forward=reshard_after_forward, + ) + fsdp_model = FSDP( + model, + flatten_parameters=flatten_parameters, + reshard_after_forward=reshard_after_forward, + ) + fsdp_optim = torch.optim.AdamW(fsdp_model.parameters(), lr=1e-2) + + orig_post_backward_reshard_hook = FSDP._post_backward_reshard_hook + reshard_hook_count = 0 + + def post_backward_reshard_hook_with_count(*args, **kwargs): + nonlocal reshard_hook_count + reshard_hook_count += 1 + return orig_post_backward_reshard_hook(*args, **kwargs) + + with mock.patch( + "fairscale.nn.data_parallel.FullyShardedDataParallel._post_backward_reshard_hook", + post_backward_reshard_hook_with_count, + ): + inp = torch.randn((8, 5), device="cuda", requires_grad=inp_requires_grad) + for i in range(6): + losses = [] + for model, optim in ((fsdp_model, fsdp_optim), (ref_model, ref_optim)): + optim.zero_grad() + loss = model(inp).sum() + losses.append(loss) + loss.backward() + optim.step() + expected_reshard_hook_count = num_frozen_linears + if not flatten_parameters: + expected_reshard_hook_count *= 2 # weight and bias per linear + assert ( + reshard_hook_count == expected_reshard_hook_count + ), f"Expected {expected_reshard_hook_count} but got {reshard_hook_count}" + assert losses[0].eq(losses[1]).all().item(), f"Expected {losses[1]} but got {losses[0]}" + reshard_hook_count = 0 From a3ff5c4369e7ca7de0a01fcf867e1215e9b40a44 Mon Sep 17 00:00:00 2001 From: Jiecao Yu Date: Wed, 21 Feb 2024 03:38:59 -0800 Subject: [PATCH 22/27] Avoid calling _free_fp16_param_shard() too early with PR 1159 --- fairscale/nn/data_parallel/fully_sharded_data_parallel.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index e9419e2a8..a057f8495 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1752,7 +1752,7 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # Free full params. self._free_full_params([param]) - if self.mixed_precision: + if self.mixed_precision and (self._require_backward_grad_sync or self.reshard_after_forward): # This is a no-op if reshard_after_forward is True, since we already # free the param shard when rebuilding the full params in the # pre_backward_hook. @@ -1887,7 +1887,7 @@ def _post_reduction_hook(self, param: Parameter, reduced_grad: torch.Tensor) -> def _post_backward_reshard_hook(self, param: Parameter, *unused: Any) -> None: if self._should_free_in_backward(): self._free_full_params([param]) - if self.mixed_precision: + if self.mixed_precision and (self._require_backward_grad_sync or self.reshard_after_forward): self._free_fp16_param_shard([param]) self._use_fp32_param_shard([param]) @@ -1972,7 +1972,7 @@ def _finalize_parameters(fsdp_module: FullyShardedDataParallel) -> None: # For the 1st layer, if the forward inputs did not require # gradient, then we cannot run a reshard hook for it, and # we instead free here. - if p._full_param_padded.untyped_storage().size() > 0: + if p._is_sharded and p._full_param_padded.untyped_storage().size() > 0: fsdp_module._post_backward_reshard_hook(p) continue From 9d0e41e56542b5cc10297dd6410223ab8836b4ab Mon Sep 17 00:00:00 2001 From: Jie Wang Date: Mon, 25 Mar 2024 11:56:10 -0700 Subject: [PATCH 23/27] Added requires_grad check for params_with_grad method (#1171) Co-authored-by: Jie Wang --- fairscale/nn/data_parallel/fully_sharded_data_parallel.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index a057f8495..e681870ec 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -697,7 +697,7 @@ def _cast_buffers( @property def params_with_grad(self) -> List[Parameter]: """[p for p in self.parameters() if p.grad is not None]""" - return [p for p in self.parameters() if (p.grad is not None or p.main_grad is not None)] + return [p for p in self.parameters() if (p.requires_grad and (p.grad is not None or p.main_grad is not None))] @torch.no_grad() def clip_grad_norm_( From e43a22fd1f3fbda98b247f77ab99a711b32913a1 Mon Sep 17 00:00:00 2001 From: Andrew Gu <31054793+awgu@users.noreply.github.com> Date: Mon, 1 Apr 2024 14:08:57 -0400 Subject: [PATCH 24/27] Changed to only run reshard hook if all gradients computed (#1166) * Changed to only run reshard hook if all gradients computed * Fix decreasing it/s with multi-grad hook --- .../fully_sharded_data_parallel.py | 76 ++++++++++++++++++- 1 file changed, 73 insertions(+), 3 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index e681870ec..fb4e2c5c5 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -28,6 +28,7 @@ Mapping, NamedTuple, Optional, + Sequence, Set, Tuple, Union, @@ -42,6 +43,7 @@ import torch.nn as nn import torch.nn.functional as F from torch.nn.parameter import Parameter +from torch.utils.hooks import RemovableHandle from fairscale.nn.misc import FlattenParamsWrapper from fairscale.nn.wrap import auto_wrap, config_auto_wrap_policy, enable_wrap @@ -1678,12 +1680,9 @@ def _register_post_backward_hooks(self) -> None: def _register_post_backward_reshard_hooks( self, args: Tuple[Any, ...], kwargs: Dict[str, Any] ) -> None: - if not hasattr(torch.autograd.graph, "register_multi_grad_hook"): - return # unsupported if not torch.is_grad_enabled(): return from torch.utils._pytree import tree_flatten - from torch.autograd.graph import register_multi_grad_hook # Construct `inp_tensors` lazily to avoid CPU overhead in typical case # where each parameter requires gradient inp_tensors: Optional[List[torch.Tensor]] = None @@ -2867,3 +2866,74 @@ def auto_wrap_bn( enable_wrap(config_auto_wrap_policy, wrapper_cls=FullyShardedDataParallel) if wrap_it else contextlib.suppress() ): return auto_wrap(module) + + +class Handle(RemovableHandle): + handles: Tuple[RemovableHandle, ...] + + def __init__(self, handles: Tuple[RemovableHandle, ...]): + self.handles = handles + + def remove(self): + for handle in self.handles: + handle.remove() + + def __getstate__(self): + return self.handles + + def __setstate__(self, state): + self.handles = state + + +def register_multi_grad_hook( + tensors: Sequence[torch.Tensor], + fn: Callable[[Sequence[Optional[torch.Tensor]]], None] +): + count: Dict[int, int] = dict() + nb_calls = None + buffer: Dict[int, List[Optional[torch.Tensor]]] = dict() + + grad_fns = list(map(_get_grad_fn_or_grad_acc, tensors)) + len_tensors = len(tensors) + + def get_inner_hook(idx): + def inner_hook(grad: torch.Tensor): + nonlocal count, nb_calls, buffer, fn + id = torch._C._current_graph_task_id() + assert ( + id != -1 + ), "expected this hook to be called inside a backward call" + count[id] = count.get(id, 0) + buffer[id] = buffer.get(id, [None] * len_tensors) + + if count[id] == 0: + # On the first call, compute the actual nb_calls and buffer + # nb_calls = sum(torch._C._will_engine_execute_node(g) for g in grad_fns) # type: ignore[attr-defined] + + # NOTE: To avoid resharding too early when microbatches share + # some same module inputs, let us require all gradients to be + # computed in this backward for the hook to run. + nb_calls = len(grad_fns) + + buffer[id][idx] = grad + count[id] += 1 + + if count[id] == nb_calls: + fn = cast(Callable[[Sequence[Optional[torch.Tensor]]], None], fn) + fn(buffer[id]) + del count[id] + del buffer[id] + + return inner_hook + + handles: Tuple[RemovableHandle, ...] = tuple( + t.register_hook(get_inner_hook(i)) for i, t in enumerate(tensors) + ) + return Handle(handles) + + +def _get_grad_fn_or_grad_acc(t): + if t.requires_grad and t.grad_fn is None: + return t.view_as(t).grad_fn.next_functions[0][0] + else: + return t.grad_fn From f039a3ae07c2ef100a5fb98bab3544b2dacb1047 Mon Sep 17 00:00:00 2001 From: Jie Wang Date: Fri, 5 Apr 2024 12:45:35 -0700 Subject: [PATCH 25/27] Add cast input argument (#1175) Co-authored-by: Jie Wang --- fairscale/nn/data_parallel/fully_sharded_data_parallel.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index fb4e2c5c5..13f1e5577 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -375,6 +375,7 @@ def __init__( gradient_predivide_factor: Optional[float] = None, limit_all_gather_events: bool = False, limit_reduce_scatter_events: bool = False, + cast_input: bool = True, optimize_backward_concat: bool = False, ): try: @@ -426,6 +427,7 @@ def __init__( self.reshard_after_forward = self._orig_reshard_after_forward = reshard_after_forward self.disable_reshard_on_root = disable_reshard_on_root self.mixed_precision = mixed_precision + self.cast_input = cast_input self.fp32_reduce_scatter = fp32_reduce_scatter self.flatten_parameters = flatten_parameters self.move_params_to_cpu = move_params_to_cpu or cpu_offload @@ -1450,7 +1452,7 @@ def forward(self, *args: Any, **kwargs: Any) -> torch.Tensor: # For root and mixed precision, we convert the input to FP16 (no_grad is needed for # the conversion). is_bf16 = self.compute_dtype == torch.bfloat16 - if self._is_root and self.mixed_precision: + if self._is_root and self.mixed_precision and self.cast_input: args, kwargs = cast_floats_to_right_precision(True, True, is_bf16, *args, **kwargs) if self not in self._fsdp_forward_ordering: From 529998259e43d0a0b43568e0753fbde714d01d8c Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Tue, 14 May 2024 22:48:56 -0700 Subject: [PATCH 26/27] honor optimize_backward_concat flag --- fairscale/nn/misc/flatten_params_wrapper.py | 20 ++++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index dfbdaf60f..35268f850 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -34,6 +34,7 @@ from fairscale.experimental.nn.ssd_offload import SsdFlatParameter from fairscale.utils.state_dict import replace_by_prefix_ +import functools if TYPE_CHECKING: from collections import OrderedDict # noqa: F401 @@ -379,7 +380,11 @@ def _unflatten_params(self, external_data: Optional[List[Optional[Tensor]]] = No delattr(self, n) self.flat_params = [] +<<<<<<< HEAD # The post backward hook used to accumulate fp32 gradients +======= + +>>>>>>> cbc3b89 (honor optimize_backward_concat flag) def _grad_accumulation_hook( self, grad, @@ -388,7 +393,11 @@ def _grad_accumulation_hook( if self.fp32_grads[param_index] is None: self.fp32_grads[param_index] = grad.to(torch.float32) else: +<<<<<<< HEAD self.fp32_grads[param_index].add_(grad) +======= + self.fp32_grads[param_index].add_(grad.data) +>>>>>>> cbc3b89 (honor optimize_backward_concat flag) return grad def _unflatten_params_as_views(self) -> None: @@ -397,12 +406,15 @@ def _unflatten_params_as_views(self) -> None: """ assert self.is_flattened if self.optimize_backward_concat: +<<<<<<< HEAD # If self._require_backward_grad_sync == True (e.g. last microbatch), # we use the original flat_params as autograd leaf nodes and backward # pass should propagate all the way back to FSDP module and thus invoke # FSDP post_backward() hook and concat() op # Otherwise we stop the backward propagation before FSDP module to avoid # invoking concat() and store the accumulated fp32 grads +======= +>>>>>>> cbc3b89 (honor optimize_backward_concat flag) if self._require_backward_grad_sync: ps = self.get_param_views() else: @@ -415,12 +427,17 @@ def _unflatten_params_as_views(self) -> None: for (_, m, n), p in zip(self._param_infos, ps): setattr(p, '_fsdp_weight', True) setattr(m, n, p) # This will set as plain attr +<<<<<<< HEAD # The param_index of p used to accumulate the correspnding # gradients in self.fp32_grads param_index = len(param_views) if self.optimize_backward_concat: # Register post backward hook to accumulate the gradients # in self.fp32_grads +======= + param_index = len(param_views) + if self.optimize_backward_concat: +>>>>>>> cbc3b89 (honor optimize_backward_concat flag) p.register_hook( functools.partial( self._grad_accumulation_hook, @@ -430,7 +447,10 @@ def _unflatten_params_as_views(self) -> None: param_views.append(p) if self.optimize_backward_concat and len(self.fp32_grads) == 0: +<<<<<<< HEAD # Allocate self.fp32_grads at the beginning of each data batch's forward() +======= +>>>>>>> cbc3b89 (honor optimize_backward_concat flag) self.fp32_grads = [None] * len(param_views) # Save param views for easy access if anyone still wants to access From b5e138fa2ae64fe89f367f99e9439fe30ef7f020 Mon Sep 17 00:00:00 2001 From: Chris Cai Date: Wed, 15 May 2024 15:02:05 -0700 Subject: [PATCH 27/27] use grad instead of grad.data --- fairscale/nn/misc/flatten_params_wrapper.py | 19 ------------------- 1 file changed, 19 deletions(-) diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 35268f850..7ef9ea1a9 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -380,11 +380,7 @@ def _unflatten_params(self, external_data: Optional[List[Optional[Tensor]]] = No delattr(self, n) self.flat_params = [] -<<<<<<< HEAD # The post backward hook used to accumulate fp32 gradients -======= - ->>>>>>> cbc3b89 (honor optimize_backward_concat flag) def _grad_accumulation_hook( self, grad, @@ -393,11 +389,7 @@ def _grad_accumulation_hook( if self.fp32_grads[param_index] is None: self.fp32_grads[param_index] = grad.to(torch.float32) else: -<<<<<<< HEAD self.fp32_grads[param_index].add_(grad) -======= - self.fp32_grads[param_index].add_(grad.data) ->>>>>>> cbc3b89 (honor optimize_backward_concat flag) return grad def _unflatten_params_as_views(self) -> None: @@ -406,15 +398,12 @@ def _unflatten_params_as_views(self) -> None: """ assert self.is_flattened if self.optimize_backward_concat: -<<<<<<< HEAD # If self._require_backward_grad_sync == True (e.g. last microbatch), # we use the original flat_params as autograd leaf nodes and backward # pass should propagate all the way back to FSDP module and thus invoke # FSDP post_backward() hook and concat() op # Otherwise we stop the backward propagation before FSDP module to avoid # invoking concat() and store the accumulated fp32 grads -======= ->>>>>>> cbc3b89 (honor optimize_backward_concat flag) if self._require_backward_grad_sync: ps = self.get_param_views() else: @@ -427,17 +416,12 @@ def _unflatten_params_as_views(self) -> None: for (_, m, n), p in zip(self._param_infos, ps): setattr(p, '_fsdp_weight', True) setattr(m, n, p) # This will set as plain attr -<<<<<<< HEAD # The param_index of p used to accumulate the correspnding # gradients in self.fp32_grads param_index = len(param_views) if self.optimize_backward_concat: # Register post backward hook to accumulate the gradients # in self.fp32_grads -======= - param_index = len(param_views) - if self.optimize_backward_concat: ->>>>>>> cbc3b89 (honor optimize_backward_concat flag) p.register_hook( functools.partial( self._grad_accumulation_hook, @@ -447,10 +431,7 @@ def _unflatten_params_as_views(self) -> None: param_views.append(p) if self.optimize_backward_concat and len(self.fp32_grads) == 0: -<<<<<<< HEAD # Allocate self.fp32_grads at the beginning of each data batch's forward() -======= ->>>>>>> cbc3b89 (honor optimize_backward_concat flag) self.fp32_grads = [None] * len(param_views) # Save param views for easy access if anyone still wants to access