Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / core / fully_shard / hsdp_scheduler.py: 83%
201 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-22 04:23 +0800
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-22 04:23 +0800
1# Copyright 2025-2026 Huawei Technologies Co., Ltd
2#
3# Licensed under the Apache License, Version 2.0 (the "License");
4# you may not use this file except in compliance with the License.
5# You may obtain a copy of the License at
6#
7# http://www.apache.org/licenses/LICENSE-2.0
8#
9# Unless required by applicable law or agreed to in writing, software
10# distributed under the License is distributed on an "AS IS" BASIS,
11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12# See the License for the specific language governing permissions and
13# limitations under the License.
14# ============================================================================
15"""HSDP scheduler"""
16import functools
17from typing import Any, List, Optional, Tuple, Union
19from hyper_parallel.platform import get_platform
20from hyper_parallel.core.dtensor.device_mesh import DeviceMesh
21from hyper_parallel.core.fully_shard.hsdp_utils import (
22 FSDPSchedulerState,
23 HSDPConfigV2,
24 get_managed_modules_parameters,
25 get_hsdp_state
26)
27from hyper_parallel.tools.logging import get_logger
29logger = get_logger("FSDP")
31platform = get_platform()
34class HSDPSchedulerContext:
35 """HSDPSchedulerContext"""
37 def __init__(self) -> None:
38 # Currently only record is_last_backward flag for scheduler context.
39 self.is_last_backward: bool = True
40 # flag to identify "root_module"
41 self.root_module = None
42 # Compile tracing may enter the root more than once; initialize shared
43 # parameter state only on the first real forward.
44 self.lazy_init_done: bool = False
47class HSDPSchedulerV2:
48 """HSDPScheduler is used to scheduler hsdp"""
49 root_bp_state = False
52 def __init__(self, cell: Union[platform.Module, Tuple[platform.Module, ...]], mesh,
53 reshard_after_forward, shard_placement_fn,
54 mp_policy, offload_policy, ignored_params, replicate_params, device, comm_fusion,
55 comm_fusion_zero_copy=False):
56 """init hsdp scheduler.
58 Args:
59 cell: A single platform.Module or tuple of platform.Module to manage as one FSDP unit.
60 """
61 self.modules = (cell,) if isinstance(cell, platform.Module) else tuple(cell)
62 self.cell = self.modules[0]
63 self.mesh: DeviceMesh = mesh
64 self.reshard_after_forward = reshard_after_forward
65 self.shard_placement_fn = shard_placement_fn
66 self.mp_policy = mp_policy
67 self.offload_policy = offload_policy
68 self.ignored_params = ignored_params
69 self.replicate_params = replicate_params
70 self.device = device
71 self.scheduler_state = None
72 self.forward_prefetch_cells = []
73 self.backward_prefetch_cells = []
74 self._backup_forward_fetch = None
75 # Flag to identify root module.
76 self._is_root = False
77 # module and its all sub-modules share one same 'HSDPSchedulerContext'
78 self.scheduler_ctx = HSDPSchedulerContext()
79 # When ``fully_shard`` is given multiple root modules, forward pre/post hooks coordinate
80 # so unshard / PostBackward / reshard run once per forward (aligned with PyTorch FSDP2).
81 self._fsdp_group_post_pending: Optional[set] = set() if len(self.modules) > 1 else None
82 self.config = HSDPConfigV2(
83 mesh,
84 reshard_after_forward,
85 shard_placement_fn,
86 mp_policy,
87 offload_policy,
88 ignored_params,
89 replicate_params,
90 comm_fusion=comm_fusion,
91 comm_fusion_zero_copy=comm_fusion_zero_copy,
92 )
93 self._init_platform()
94 self._new_cell_state()
95 self._register_hooks()
97 def _init_platform(self):
98 """Initialize the platform."""
99 raise NotImplementedError("HSDPScheduler subclasses must implement _init_platform")
101 def _new_cell_state(self):
102 """Create a new cell state."""
103 raise NotImplementedError("HSDPScheduler subclasses must implement _new_cell_state")
105 def _register_hooks(self):
106 """Register hooks."""
107 raise NotImplementedError("HSDPScheduler subclasses must implement _register_hooks.")
109 def _register_forward_backward_hooks(self):
110 """Register module forward and backward hook."""
111 raise NotImplementedError("HSDPScheduler subclasses must implement _register_forward_backward_hooks.")
113 def _get_managed_params(self):
114 """Return deduplicated parameters from all managed modules."""
115 return get_managed_modules_parameters(self.modules, self.ignored_params)
117 def set_reshard_after_forward(self, reshard_after_forward: bool) -> None:
118 """Set reshard_after_forward flag.
120 Args:
121 reshard_after_forward: Whether to reshard parameters after forward.
122 """
123 if not isinstance(reshard_after_forward, bool):
124 raise ValueError(f"reshard_after_forward should be a bool, got {type(reshard_after_forward)}")
125 self.reshard_after_forward = reshard_after_forward
126 self.config.reshard_after_forward = reshard_after_forward
128 def set_reshard_after_backward(self, reshard_after_backward: bool) -> None:
129 """Set reshard_after_backward flag.
131 Args:
132 reshard_after_backward: Whether to reshard after backward completes.
133 """
134 if not isinstance(reshard_after_backward, bool):
135 raise ValueError(f"reshard_after_backward should be a bool, got {type(reshard_after_backward)}")
136 if self.hsdp_state is not None:
137 self.hsdp_state.reshard_after_backward = reshard_after_backward
139 def set_requires_all_reduce(self, requires_all_reduce: bool) -> None:
140 """Set requires_all_reduce flag.
142 Args:
143 requires_all_reduce: Whether this unit participates in all-reduce.
144 """
145 if not isinstance(requires_all_reduce, bool):
146 raise ValueError(f"requires_all_reduce should be a bool, got {type(requires_all_reduce)}")
147 if self.hsdp_state is not None:
148 self.hsdp_state.requires_all_reduce = requires_all_reduce
150 def set_requires_grad_sync(self, requires_grad_sync: bool) -> None:
151 """Set flag controlling whether gradients are synchronized.
153 Args:
154 requires_grad_sync: When True, enable grad sync for this scheduler.
155 """
156 if not isinstance(requires_grad_sync, bool):
157 raise ValueError(f"requires_grad_sync should be a bool, got {type(requires_grad_sync)}")
158 self.hsdp_state.set_requires_grad_sync(requires_grad_sync)
160 # pylint: disable=W0613
161 def _hsdp_forward_pre_hook(self, cell, args, kwargs):
162 """Forward pre hook to unsharded parameter for forward process."""
163 logger.debug("hook=forward_pre enter module=%s", self.hsdp_state)
164 if self.scheduler_state == FSDPSchedulerState.PRE_BACKWARD:
165 logger.debug("hook=forward_pre skip module=%s reason=pre_backward", self.hsdp_state)
166 return args, kwargs
167 if HSDPSchedulerV2.root_bp_state:
168 self._disable_forward_prefetch_for_recompute()
169 if self.scheduler_ctx.root_module is None:
170 self.scheduler_ctx.root_module = self.cell
171 self._is_root = True
172 for _, module in platform.get_cells_and_names(self.scheduler_ctx.root_module):
173 from hyper_parallel.core.fully_shard.api import HSDPModule # pylint: disable=C0415
174 if isinstance(module, HSDPModule):
175 submod_scheduler = getattr(module, "hsdp_scheduler", None)
176 if submod_scheduler and submod_scheduler.scheduler_ctx is not self.scheduler_ctx:
177 submod_scheduler.scheduler_ctx = self.scheduler_ctx
179 if not self._is_root and not self.hsdp_state.module_name:
180 for module_name, module in platform.get_cells_and_names(self.scheduler_ctx.root_module):
181 if module == self.cell:
182 self.hsdp_state.module_name = module_name
183 break
184 self.scheduler_state = FSDPSchedulerState.PRE_FORWARD
185 if self._is_root and not self.scheduler_ctx.lazy_init_done:
186 self._init_params_fqn()
187 self._lazy_init_all_states()
188 self.scheduler_ctx.lazy_init_done = True
189 if self.mp_policy.cast_forward_inputs and self.mp_policy.param_dtype:
190 cast_fn = functools.partial(self.platform.cast_fp_tensor, self.mp_policy.param_dtype)
191 args = self.platform.apply_to_tensors(cast_fn, args)
192 kwargs = self.platform.apply_to_tensors(cast_fn, kwargs)
193 with self.platform.profiler_record(f"pre_forward unshard:{self.hsdp_state.module_name}"):
194 logger.debug("hook=forward_pre action=unshard module=%s", self.hsdp_state)
195 self.hsdp_state.unshard()
196 for prefetch_cell in self.forward_prefetch_cells:
197 prefetch_state = prefetch_cell.hsdp_scheduler.hsdp_state
198 with self.platform.profiler_record(f"pre_forward prefetch:"
199 f"{prefetch_state.module_name}"):
200 logger.debug(
201 "hook=forward_pre action=prefetch module=%s target=%s",
202 self.hsdp_state,
203 prefetch_state,
204 )
205 prefetch_state.prefetch()
206 return args, kwargs
208 def _lazy_init_all_states(self):
209 if self._is_root and self.scheduler_ctx.root_module is not None:
210 for _, module in platform.get_cells_and_names(self.scheduler_ctx.root_module):
211 hsdp_state = get_hsdp_state(module)
212 if hsdp_state:
213 hsdp_state.lazy_init()
215 def _init_params_fqn(self): # pylint: disable=W0212
216 if not self._is_root or self.scheduler_ctx.root_module is None:
217 return
218 # Build a map from original (sharded) parameter tensor → hsdp_param wrapper,
219 # covering both sharded hsdp_params and replicate_params.
220 param_to_hsdp_param = {}
221 for _, module in platform.get_cells_and_names(self.scheduler_ctx.root_module):
222 hsdp_state = get_hsdp_state(module)
223 if hsdp_state is None:
224 continue
225 for hsdp_param in hsdp_state._iter_managed_params(): # pylint: disable=W0212
226 orig_param = hsdp_param.sharded_param
227 # Shared parameters: keep only the first mapping to preserve the
228 # first-seen FQN (consistent with the deduplication in _init_hsdp_params).
229 if orig_param not in param_to_hsdp_param:
230 param_to_hsdp_param[orig_param] = hsdp_param
232 # Walk the full parameter tree and assign FQNs; skip params already seen
233 # (shared-parameter deduplication: first name wins).
234 visited_params = set()
235 for param_name, parameter in platform.parameters_dict(self.scheduler_ctx.root_module):
236 if parameter in visited_params:
237 continue
238 visited_params.add(parameter)
239 hsdp_param = param_to_hsdp_param.get(parameter)
240 if hsdp_param is not None:
241 hsdp_param._param_fqn = param_name # pylint: disable=W0212
243 # pylint: disable=W0613, R1710
244 def _hsdp_forward_hook(self, cell, inputs, outputs):
245 """Forward hook to shard parameter for saving memory."""
246 logger.debug("hook=forward enter module=%s", self.hsdp_state)
247 if self.scheduler_state == FSDPSchedulerState.PRE_BACKWARD:
248 logger.debug("hook=forward skip module=%s reason=pre_backward", self.hsdp_state)
249 return
250 self.scheduler_state = FSDPSchedulerState.FORWARD
251 if self.reshard_after_forward:
252 with self.platform.profiler_record(f"forward reshard:{self.hsdp_state.module_name}"):
253 logger.debug("hook=forward action=reshard module=%s", self.hsdp_state)
254 self.hsdp_state.shard(shard_replicate=False)
255 if self.mp_policy.output_dtype is not None:
256 outputs = self.platform.apply_to_tensors(
257 functools.partial(self.platform.cast_fp_tensor, self.mp_policy.output_dtype),
258 outputs,
259 )
260 return outputs
262 # pylint: disable=W0613
263 def _hsdp_backward_pre_hook(self, cell, grad_outputs):
264 """Backward pre hook to unsharded parameter for backward process."""
265 logger.debug("hook=backward_pre enter module=%s", self.hsdp_state)
266 self.scheduler_state = FSDPSchedulerState.PRE_BACKWARD
267 if self.reshard_after_forward:
268 with self.platform.profiler_record(f"pre_backward unshard:{self.hsdp_state.module_name}"):
269 logger.debug("hook=backward_pre action=unshard module=%s", self.hsdp_state)
270 self.hsdp_state.unshard(unshard_replicate=False)
271 for prefetch_cell in self.backward_prefetch_cells:
272 prefetch_state = prefetch_cell.hsdp_scheduler.hsdp_state
273 with self.platform.profiler_record(f"pre_backward prefetch:"
274 f"{prefetch_state.module_name}"):
275 logger.debug(
276 "hook=backward_pre action=prefetch module=%s target=%s",
277 self.hsdp_state,
278 prefetch_state,
279 )
280 prefetch_state.prefetch(unshard_replicate=False)
282 # pylint: disable=W0613
283 def _hsdp_backward_hook(self, cell, grad_inputs, grad_outputs):
284 """Backward hook to shard parameter for optimizer process or saving memory."""
285 logger.debug("hook=backward_hook enter module=%s", self.hsdp_state)
286 self.scheduler_state = FSDPSchedulerState.BACKWARD
287 with self.platform.profiler_record(f"post_backward:{self.hsdp_state.module_name}"):
288 logger.debug("hook=backward_hook action=post_backward module=%s", self.hsdp_state)
289 self.hsdp_state.post_backward()
290 if self._fsdp_group_post_pending is not None:
291 self._fsdp_group_post_pending.clear()
293 # pylint: disable=W0613
294 @staticmethod
295 def _grouped_forward_pre_hook_skip(cell, args, kwargs):
296 """Return value when grouped pre-forward should not run (first module already did).
298 Default matches MindSpore Cell forward pre-hooks (explicit ``(args, kwargs)``).
299 ``TorchHSDPSchedulerV2`` overrides this to return ``None`` (``nn.Module`` idiom).
300 """
301 return args, kwargs
303 @staticmethod
304 def _grouped_forward_post_hook_skip(outputs):
305 """Return value when grouped post-forward is deferred to a later module in the group.
307 Default returns ``outputs`` (MindSpore). ``TorchHSDPSchedulerV2`` overrides to ``None``.
308 """
309 return outputs
311 def _grouped_forward_pre_hook(self, cell, args, kwargs):
312 """Run FSDP pre-forward only for the first module in the group (PyTorch FSDP2-aligned)."""
313 pending = self._fsdp_group_post_pending
314 if pending is None:
315 return self._forward_pre_hook(cell, args, kwargs)
316 if len(pending) == 0:
317 pending.update(self.modules)
318 return self._forward_pre_hook(cell, args, kwargs)
319 return self._grouped_forward_pre_hook_skip(cell, args, kwargs)
321 def _make_grouped_forward_post_hook(self, mod):
322 """Build post-forward hook: last module in the group runs reshard + output backward hooks."""
324 def grouped_post_hook(cell, inputs, outputs):
325 pending = self._fsdp_group_post_pending
326 if pending is None:
327 return self._forward_hook(cell, inputs, outputs)
328 pending.discard(mod)
329 if len(pending) == 0:
330 return self._forward_hook(cell, inputs, outputs)
331 return self._grouped_forward_post_hook_skip(outputs)
333 return grouped_post_hook
335 def set_forward_prefetch_cells(self, hsdp_cell_list: List[Any]) -> None:
336 """Set cells prefetched during forward.
338 Args:
339 hsdp_cell_list: HSDP cells to prefetch ahead of forward.
340 """
341 self.forward_prefetch_cells = hsdp_cell_list
343 def set_backward_prefetch_cells(self, hsdp_cell_list: List[Any]) -> None:
344 """Set cells prefetched during backward.
346 Args:
347 hsdp_cell_list: HSDP cells to prefetch ahead of backward.
348 """
349 self.backward_prefetch_cells = hsdp_cell_list
351 def _disable_forward_prefetch_for_recompute(self) -> None:
352 """Temporarily disable forward prefetch during activation recompute."""
353 self._backup_forward_fetch = self.forward_prefetch_cells
354 self.forward_prefetch_cells = []
356 def _restore_forward_prefetch_after_recompute(self) -> bool:
357 """Restore forward prefetch list after a recompute forward hook finishes."""
358 if self._backup_forward_fetch is None:
359 return False
360 self.forward_prefetch_cells = self._backup_forward_fetch
361 self._backup_forward_fetch = None
362 return True