Skip to content
16 changes: 13 additions & 3 deletions fairscale/nn/data_parallel/fully_sharded_data_parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -494,7 +494,10 @@ def __init__(
del param_names

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
)
del module # free original module in case it helps garbage collection

Expand Down Expand Up @@ -1099,12 +1102,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:
Expand Down Expand Up @@ -1851,7 +1856,10 @@ 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]):

# TODO(chriscai): find a better way to handle the state transition check
# when we need to skip FSDP flatten parameter backward()
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)
Expand Down Expand Up @@ -1923,12 +1931,14 @@ def _finalize_parameters(fsdp_module: FullyShardedDataParallel) -> None:
_finalize_parameters(m)
m._free_ssd_offload()
m._pre_backward_hook_has_run = False
# TODO(chriscai): find a better way to handle the state transition check
# when we need to skip FSDP flatten parameter backward()
if any(p.requires_grad for p in m.parameters()):
# Check if the module has params and if any of them has
# 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)
Expand Down
94 changes: 92 additions & 2 deletions fairscale/nn/misc/flatten_params_wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -162,6 +164,7 @@ def __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:
param_list = list(module.parameters())
Expand Down Expand Up @@ -243,6 +246,8 @@ def __init__(
# explicitly requests a flat state dict via flat_state_dict().
self._auto_unflatten_state_dict = True

self._new_params = []

@property
def module(self) -> Any:
"""Support fpw.module in case we are immitating DDP, which has .module
Expand Down Expand Up @@ -368,13 +373,98 @@ def _unflatten_params_as_views(self) -> None:
"""Unlike ``_unflatten_params``, this function unflatten into views and keep
self.flat_param unchanged.
"""
logger.info("CHRISLOG: _unflatten_params_as_views() called")
assert self.is_flattened
ps = self.get_param_views()
# ps = self.get_param_views()

"""Return a generator of views that map to the original parameters."""

"""Used to get a generator over all views from a list of external data list."""
params = self.flat_params
external_data_list = [None] * len(params)
assert len(external_data_list) == len(
params
), f"Incorrect external data list: {len(external_data_list)} vs. {len(params)}"

# Post accumulation hook so we can call backward() on original leaf params at last microbatch
import functools

def _post_accumulation_hook(new_param_stop_grad, param_tuples, new_param_index):
if self._require_backward_grad_sync:
param_tuples[new_param_index][1] = new_param_stop_grad.grad
if any([t[1] is None for t in param_tuples]):
logger.info(
f"CHRISLOG: _post_accumulation_hook() not the last parameter in current FSDP module, param {new_param_index=} {len(param_tuples)=}"
)
else:
logger.info(
f"CHRISLOG: _post_accumulation_hook() all grads are generated, param {new_param_index=} {len(param_tuples)=}"
)
torch.autograd.backward([t[0] for t in param_tuples], grad_tensors=[t[1] for t in param_tuples])
logger.info(
f"CHRISLOG: _post_accumulation_hook() torch.autograd.backward() called with {len(param_tuples)=}"
)
else:
logger.info(
f"CHRISLOG: _post_accumulation_hook() {self._require_backward_grad_sync=} skipping calling backward() on param with {new_param_index=}, {len(param_tuples)=}"
)


_new_params_and_new_params_stop_grad_tuples = []

gens = []
logger.info(
f"CHRISLOG: {len(params)=}"
)
for p, data in zip(params, external_data_list):
# Sanity check
assert p.data.numel() <= sum(
p._param_numels
), f"Incorrect internal state {p.data.numel()} vs. {sum(p._param_numels)}"
data = data if data is not None else p
if data.numel() != sum(p._param_numels):
raise ValueError(
f"Incorrect numel of supplied data: got {data.numel()} but expected {sum(p._param_numels)}"
)

# Split the data into views of each parameter.
param_views_stop_grad = []
for t, s in zip(data.split(p._param_numels), p._param_shapes):
# Create unflattened view for the param.
new_param = t.view(s)
_new_params_and_new_params_stop_grad_tuples.append([new_param, None])
# Create a new_param_stop_grad param via detaching original leaf params after .view()
# as the new leaf nodes so that autograd.backward() won't call
# grad_fn of view() (which will be cat())
# TODO: need to figure out how to remove the clone()
new_param_stop_grad = new_param.detach().requires_grad_(True)
# Register post-accumulation hook to the new_param_stop_grad parameters so that
# we can still manually call backward() function
# to propogate gradients to the original leaf params, e.g. after last_microbatch
# backward()
new_param_stop_grad.register_post_accumulate_grad_hook(
functools.partial(_post_accumulation_hook,
param_tuples=_new_params_and_new_params_stop_grad_tuples,
new_param_index=len(_new_params_and_new_params_stop_grad_tuples) - 1)
)
param_views_stop_grad.append(new_param_stop_grad)

logger.info(
f"CHRISLOG: appending {len(param_views_stop_grad)=}"
)
gens.append(param_views_stop_grad)
ps = chain(*gens)

param_index = 0
# Set the param with unflattened view as the new attribute
# under original param name
param_views = []
for (_, m, n), p in zip(self._param_infos, ps):
setattr(p, '_fsdp_weight', True)
setattr(p, "_fsdp_weight", True)
setattr(m, n, p) # This will set as plain attr
logger.info(f"CHRISLOG: {m.__class__.__module__=} {n=} {p.is_leaf=} {p.size()=} {param_index=}")
param_views.append(p)
param_index += 1

# Save param views for easy access if anyone still wants to access
# parameters of the module.
Expand Down