From 88589c1be55b8c8c3fda2d99ebbe7965cf61a6c8 Mon Sep 17 00:00:00 2001 From: vedanuj Date: Tue, 3 Oct 2023 08:16:22 -0700 Subject: [PATCH 1/5] changes for main_grad before fwd --- .../fully_sharded_data_parallel.py | 68 +++++++++++++------ fairscale/nn/misc/flatten_params_wrapper.py | 10 ++- 2 files changed, 53 insertions(+), 25 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index dee596e9a..f3cad4538 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1680,7 +1680,7 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # then subsequent hook callbacks will see POST state. self.assert_state([TrainingState.BACKWARD_PRE, TrainingState.BACKWARD_POST]) self.training_state = TrainingState.BACKWARD_POST - if param.grad is None: + if param.grad is None and param.main_grad is None: return if hasattr(param, "_linked_param"): @@ -1693,9 +1693,9 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: if hasattr(param._linked_param, "_is_shared") and param._linked_param._is_shared: param = param._linked_param - assert param.grad is not None, param.shape - if param.grad.requires_grad: - raise RuntimeError("FSDP only works with gradients that don't require gradients") + # assert param.grad is not None, param.shape + # 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 @@ -1714,6 +1714,14 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # Switch to FP32 shard after backward. self._use_fp32_param_shard([param]) + if self.fp32_reduce_scatter: + if param.grad is not None: + if param.main_grad is not None: + param.main_grad.add_(param.grad.data.float()) + else: + param.main_grad = param.grad.data.float() + param.grad = None + if not self._require_backward_grad_sync: return @@ -1721,15 +1729,24 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # reductions in post_backward stream. self._streams["post_backward"].wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(self._streams["post_backward"]): - orig_grad_data = param.grad.data + # orig_grad_data = param.main_grad.data if self.fp32_reduce_scatter: - # Cast grad to FP32. - param.grad.data = param.grad.data.float() + # Cast grad to FP32. with .main_grad params are already in FP32. + if param.main_grad is not None: + orig_grad_data = param.main_grad.data + else: + orig_grad_data = param.grad.data.to(torch.float32) + else: + orig_grad_data = param.grad.data if self.gradient_predivide_factor > 1: # Average grad by world_size for consistency with PyTorch DDP. - param.grad.data.div_(self.gradient_predivide_factor) + # param.grad.data.div_(self.gradient_predivide_factor) + if param.main_grad is not None: + param.main_grad.data.div_(self.gradient_predivide_factor) + else: + param.grad.data.div_(self.gradient_predivide_factor) if param._is_sharded: assert self._reducer is not None @@ -1737,19 +1754,23 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # param._saved_grad_shard. If this FSDP module was called multiple times it's possible that multiple # gradient reductions will happen in an undefined order. But addition commutes, so this order doesn't # matter, neglecting rounding. - grad = param.grad.data - # Clear grad on the tensor, so any repeated gradient computations do not interfere with this reduction. - # - # The effect on memory consumption is not usually significant. No extra memory is allocated if this - # module is called only once, reduction happens quickly, or the tensor is bucketed. If the module is - # called multiple times, and the backwards pass runs far enough ahead of the `post_backward` stream, - # then we can end up with multiple unsharded gradients allocated and queued for reduction. - # - # We could guard against this by using CUDA events (see record_event, wait_event in torch.cuda.Stream). - # This ensures the `default` stream will wait for the `post_backward` stream to complete the last - # reduction for this module, before scheduling additional reduction work. Then at most there are two - # unsharded gradients allocated; one for a pending reduction, and one for gradient computation. - param.grad = None + if param.main_grad is not None: + grad = param.main_grad.data + param.main_grad = None + else: + grad = param.grad.data + # Clear grad on the tensor, so any repeated gradient computations do not interfere with this reduction. + # + # The effect on memory consumption is not usually significant. No extra memory is allocated if this + # module is called only once, reduction happens quickly, or the tensor is bucketed. If the module is + # called multiple times, and the backwards pass runs far enough ahead of the `post_backward` stream, + # then we can end up with multiple unsharded gradients allocated and queued for reduction. + # + # We could guard against this by using CUDA events (see record_event, wait_event in torch.cuda.Stream). + # This ensures the `default` stream will wait for the `post_backward` stream to complete the last + # reduction for this module, before scheduling additional reduction work. Then at most there are two + # unsharded gradients allocated; one for a pending reduction, and one for gradient computation. + param.grad = None callback_fn = functools.partial(self._post_reduction_hook, param) self._reducer.reduce_scatter_async( grad, group=self.process_group_reduce_scatter, callback_fn=callback_fn @@ -1759,7 +1780,10 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # world_size == 1. This could be relaxed in the future, in which # case grads should be all-reduced here. assert self.world_size == 1 - self._post_reduction_hook(param, param.grad.data) + if param.main_grad is not None: + self._post_reduction_hook(param, param.main_grad.data) + else: + self._post_reduction_hook(param, param.grad.data) # After _post_backward_hook returns, orig_grad_data will eventually # go out of scope, at which point it could otherwise be freed for diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 38265dd2b..80fd5d5b5 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -227,6 +227,7 @@ def __init__( flat_param.set_file_params(fname, 0) else: flat_param = FlatParameter(params, params[0].requires_grad) + flat_param.main_grad = torch.zeros_like(flat_param, dtype=torch.float32) flat_param._param_infos = param_infos flat_param._shared_param_infos = shared_param_infos self.flat_params.append(flat_param) @@ -369,10 +370,11 @@ def _unflatten_params_as_views(self) -> None: self.flat_param unchanged. """ assert self.is_flattened - ps = self.get_param_views() + ps, ps_main_grad = self.get_param_views() param_views = [] - for (_, m, n), p in zip(self._param_infos, ps): + for (_, m, n), p, p_main_grad in zip(self._param_infos, ps, ps_main_grad): setattr(p, '_fsdp_weight', True) + p.main_grad = p_main_grad setattr(m, n, p) # This will set as plain attr param_views.append(p) @@ -499,10 +501,12 @@ def get_param_views(self, external_data_list: Optional[List[Optional[Tensor]]] = ), f"Incorrect external data list: {len(external_data_list)} vs. {len(params)}" gens = [] + gens_main_grad = [] for p, data in zip(params, external_data_list): gens.append(p.get_param_views(data)) + gens_main_grad.append(p.get_param_views(p.main_grad)) - return chain(*gens) + return chain(*gens), chain(*gens_main_grad) def metadata(self, flat_param_idx: int) -> Tuple[List[str], Sequence[torch.Size], List[int]]: """Return metadata for a flat param given its index in the flat_params list.""" From f9083cf3431aeea9b958b25f00c9909bc428cb38 Mon Sep 17 00:00:00 2001 From: vedanuj Date: Wed, 4 Oct 2023 01:01:45 -0700 Subject: [PATCH 2/5] move changes after orig_grad_data --- .../fully_sharded_data_parallel.py | 38 ++++++++----------- 1 file changed, 16 insertions(+), 22 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index f3cad4538..73b924cfd 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -1714,14 +1714,6 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # Switch to FP32 shard after backward. self._use_fp32_param_shard([param]) - if self.fp32_reduce_scatter: - if param.grad is not None: - if param.main_grad is not None: - param.main_grad.add_(param.grad.data.float()) - else: - param.main_grad = param.grad.data.float() - param.grad = None - if not self._require_backward_grad_sync: return @@ -1729,24 +1721,26 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # reductions in post_backward stream. self._streams["post_backward"].wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(self._streams["post_backward"]): - # orig_grad_data = param.main_grad.data + if param.main_grad is not None: + orig_grad_data = param.main_grad + else: + orig_grad_data = param.grad if self.fp32_reduce_scatter: - # Cast grad to FP32. with .main_grad params are already in FP32. - if param.main_grad is not None: - orig_grad_data = param.main_grad.data - else: - orig_grad_data = param.grad.data.to(torch.float32) - else: - orig_grad_data = param.grad.data + if param.grad is not None: + if param.main_grad is not None: + param.main_grad.copy_(param.grad.float()) + else: + param.main_grad = param.grad.float() + param.grad = None if self.gradient_predivide_factor > 1: # Average grad by world_size for consistency with PyTorch DDP. # param.grad.data.div_(self.gradient_predivide_factor) if param.main_grad is not None: - param.main_grad.data.div_(self.gradient_predivide_factor) + param.main_grad.div_(self.gradient_predivide_factor) else: - param.grad.data.div_(self.gradient_predivide_factor) + param.grad.div_(self.gradient_predivide_factor) if param._is_sharded: assert self._reducer is not None @@ -1755,10 +1749,10 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # gradient reductions will happen in an undefined order. But addition commutes, so this order doesn't # matter, neglecting rounding. if param.main_grad is not None: - grad = param.main_grad.data + grad = param.main_grad param.main_grad = None else: - grad = param.grad.data + grad = param.grad # Clear grad on the tensor, so any repeated gradient computations do not interfere with this reduction. # # The effect on memory consumption is not usually significant. No extra memory is allocated if this @@ -1781,9 +1775,9 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # case grads should be all-reduced here. assert self.world_size == 1 if param.main_grad is not None: - self._post_reduction_hook(param, param.main_grad.data) + self._post_reduction_hook(param, param.main_grad) else: - self._post_reduction_hook(param, param.grad.data) + self._post_reduction_hook(param, param.grad) # After _post_backward_hook returns, orig_grad_data will eventually # go out of scope, at which point it could otherwise be freed for From c1169272e2d4c056fe25f52c0412a169f0e48136 Mon Sep 17 00:00:00 2001 From: vedanuj Date: Thu, 5 Oct 2023 15:53:31 -0700 Subject: [PATCH 3/5] address comments --- .../fully_sharded_data_parallel.py | 37 +++++++++++-------- 1 file changed, 21 insertions(+), 16 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index 73b924cfd..c20c4e9b3 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -43,6 +43,7 @@ from torch.nn.parameter import Parameter from fairscale.nn.misc import FlattenParamsWrapper +from fairscale.nn.misc.flatten_params_wrapper import FlatParameter from fairscale.nn.wrap import auto_wrap, config_auto_wrap_policy, enable_wrap from fairscale.utils.containers import apply_to_tensors from fairscale.utils.parallel import ( @@ -1680,7 +1681,7 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # then subsequent hook callbacks will see POST state. self.assert_state([TrainingState.BACKWARD_PRE, TrainingState.BACKWARD_POST]) self.training_state = TrainingState.BACKWARD_POST - if param.grad is None and param.main_grad is None: + if param.grad is None and getattr(param, "main_grad", None) is None: return if hasattr(param, "_linked_param"): @@ -1721,26 +1722,24 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # reductions in post_backward stream. self._streams["post_backward"].wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(self._streams["post_backward"]): - if param.main_grad is not None: + if param.main_grad is not None and not param.main_grad.eq(0).all(): orig_grad_data = param.main_grad + param.grad = None else: orig_grad_data = param.grad if self.fp32_reduce_scatter: if param.grad is not None: - if param.main_grad is not None: - param.main_grad.copy_(param.grad.float()) - else: - param.main_grad = param.grad.float() + param.main_grad.copy_(param.grad) param.grad = None if self.gradient_predivide_factor > 1: # Average grad by world_size for consistency with PyTorch DDP. # param.grad.data.div_(self.gradient_predivide_factor) - if param.main_grad is not None: - param.main_grad.div_(self.gradient_predivide_factor) - else: + if param.grad is not None: param.grad.div_(self.gradient_predivide_factor) + else: + param.main_grad.div_(self.gradient_predivide_factor) if param._is_sharded: assert self._reducer is not None @@ -1748,10 +1747,7 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # param._saved_grad_shard. If this FSDP module was called multiple times it's possible that multiple # gradient reductions will happen in an undefined order. But addition commutes, so this order doesn't # matter, neglecting rounding. - if param.main_grad is not None: - grad = param.main_grad - param.main_grad = None - else: + if param.grad is not None: grad = param.grad # Clear grad on the tensor, so any repeated gradient computations do not interfere with this reduction. # @@ -1765,6 +1761,9 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # reduction for this module, before scheduling additional reduction work. Then at most there are two # unsharded gradients allocated; one for a pending reduction, and one for gradient computation. param.grad = None + else: + grad = param.main_grad + param.main_grad = None callback_fn = functools.partial(self._post_reduction_hook, param) self._reducer.reduce_scatter_async( grad, group=self.process_group_reduce_scatter, callback_fn=callback_fn @@ -1774,10 +1773,10 @@ def _post_backward_hook(self, param: Parameter, *unused: Any) -> None: # world_size == 1. This could be relaxed in the future, in which # case grads should be all-reduced here. assert self.world_size == 1 - if param.main_grad is not None: - self._post_reduction_hook(param, param.main_grad) - else: + if param.grad is not None: self._post_reduction_hook(param, param.grad) + else: + self._post_reduction_hook(param, param.main_grad) # After _post_backward_hook returns, orig_grad_data will eventually # go out of scope, at which point it could otherwise be freed for @@ -2154,6 +2153,12 @@ def _prep_grads_for_backward(self) -> None: right shape, device, accumulated values, etc. """ for p in self.params: + if isinstance(p, FlatParameter): + if getattr(p, "main_grad", None) is None: + p.main_grad = torch.zeros_like(p.data, dtype=torch.float32) + main_grad_views = p.get_param_views(p.main_grad) + for (_, m, n), main_grad in zip(p._param_infos, main_grad_views): + getattr(m, n).main_grad = main_grad if p.grad is not None: if p.grad.device != p.data.device: p.grad = None From e8f54b6c153735033181e0848c3817cc06469747 Mon Sep 17 00:00:00 2001 From: vedanuj Date: Fri, 6 Oct 2023 09:10:55 -0700 Subject: [PATCH 4/5] ensure grads are not downcasted to bf16 --- .../nn/data_parallel/fully_sharded_data_parallel.py | 12 +++--------- fairscale/nn/misc/flatten_params_wrapper.py | 4 +++- 2 files changed, 6 insertions(+), 10 deletions(-) diff --git a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py index c20c4e9b3..5324d08f4 100644 --- a/fairscale/nn/data_parallel/fully_sharded_data_parallel.py +++ b/fairscale/nn/data_parallel/fully_sharded_data_parallel.py @@ -688,7 +688,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] + return [p for p in self.parameters() if p.grad is not None or getattr(p, "main_grad", None) is not None] @torch.no_grad() def clip_grad_norm_( @@ -1802,7 +1802,7 @@ def _post_reduction_hook(self, param: Parameter, reduced_grad: torch.Tensor) -> # non-blocking. The downside is a bit more D2H transfer in that case. if self.fp32_reduce_scatter: orig_param_grad_data = reduced_grad.data - reduced_grad.data = reduced_grad.data.to(dtype=param.data.dtype) + # reduced_grad.data = reduced_grad.data.to(dtype=param.data.dtype) # Don't let this memory get reused until after the transfer. orig_param_grad_data.record_stream(torch.cuda.current_stream()) @@ -1904,7 +1904,7 @@ def _finalize_parameters(fsdp_module: FullyShardedDataParallel) -> None: if p.shape != p._saved_grad_shard.shape: self._use_fp32_param_shard([p]) if p._saved_grad_shard.dtype != p.dtype: - p.grad = p._saved_grad_shard.to(p.dtype) + p.main_grad = p._saved_grad_shard else: p.grad = p._saved_grad_shard @@ -2153,12 +2153,6 @@ def _prep_grads_for_backward(self) -> None: right shape, device, accumulated values, etc. """ for p in self.params: - if isinstance(p, FlatParameter): - if getattr(p, "main_grad", None) is None: - p.main_grad = torch.zeros_like(p.data, dtype=torch.float32) - main_grad_views = p.get_param_views(p.main_grad) - for (_, m, n), main_grad in zip(p._param_infos, main_grad_views): - getattr(m, n).main_grad = main_grad if p.grad is not None: if p.grad.device != p.data.device: p.grad = None diff --git a/fairscale/nn/misc/flatten_params_wrapper.py b/fairscale/nn/misc/flatten_params_wrapper.py index 80fd5d5b5..91c4f59b9 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -227,7 +227,6 @@ def __init__( flat_param.set_file_params(fname, 0) else: flat_param = FlatParameter(params, params[0].requires_grad) - flat_param.main_grad = torch.zeros_like(flat_param, dtype=torch.float32) flat_param._param_infos = param_infos flat_param._shared_param_infos = shared_param_infos self.flat_params.append(flat_param) @@ -370,6 +369,9 @@ def _unflatten_params_as_views(self) -> None: self.flat_param unchanged. """ assert self.is_flattened + for p in self.flat_params: + if not hasattr(p, 'main_grad') or p.main_grad.shape != p.shape: + p.main_grad = torch.zeros_like(p, dtype=torch.float32) ps, ps_main_grad = self.get_param_views() param_views = [] for (_, m, n), p, p_main_grad in zip(self._param_infos, ps, ps_main_grad): From 8cf28fa758e491d3e52e75a01d24ec2fae4c6a79 Mon Sep 17 00:00:00 2001 From: vedanuj Date: Sun, 8 Oct 2023 12:40:15 -0700 Subject: [PATCH 5/5] guard main_grad None --- 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 91c4f59b9..8f0c8f341 100644 --- a/fairscale/nn/misc/flatten_params_wrapper.py +++ b/fairscale/nn/misc/flatten_params_wrapper.py @@ -370,7 +370,7 @@ def _unflatten_params_as_views(self) -> None: """ assert self.is_flattened for p in self.flat_params: - if not hasattr(p, 'main_grad') or p.main_grad.shape != p.shape: + if getattr(p, 'main_grad', None) is None or p.main_grad.shape != p.shape: p.main_grad = torch.zeros_like(p, dtype=torch.float32) ps, ps_main_grad = self.get_param_views() param_views = []