Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / core / fully_shard / hsdp_scheduler.py: 88%
240 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-25 04:27 +0800
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-25 04:27 +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, Mapping, 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.utils import CommFusionPolicy, SourceShardMetaInfo
22from hyper_parallel.core.fully_shard.hsdp_utils import (
23 FSDPSchedulerState,
24 get_managed_modules_parameters,
25)
26from hyper_parallel.tools.logging import get_logger
28logger = get_logger("FSDP")
30platform = get_platform()
31ModuleClass = platform.Module
34class ParamGroupCommCtx:
35 """Track the in-flight parameter-group communication for one module tree."""
37 def __init__(self) -> None:
38 self.pre_param_group = None
39 self.all_reduce_param_group = None
40 # MindSpore keeps the reduce-scatter handle outside the parameter group.
41 self.comm_handle = None
44class HSDPSchedulerContext:
45 """Share scheduler and backward-pipeline state within one HSDP module tree."""
47 def __init__(self) -> None:
48 self.is_last_backward: bool = True
49 self.root_module = None
50 # Compile tracing may enter the root more than once; initialize shared
51 # parameter state only on the first real forward.
52 self.lazy_init_done: bool = False
53 self.root_bp_state = False
54 # all_hsdp_schedulers (for one module tree structure), deduplicated at
55 # registration time because ``fully_shard`` may be given a module list
56 # whose modules share one scheduler.
57 self.all_hsdp_schedulers = []
58 # Parameter FQNs are initialized once after all schedulers share this context.
59 self._param_fqn_initialized = False
60 # Backward pipeline queues shared only by schedulers in this module tree.
61 self.pre_reduce_scatter_params = []
62 self.pre_all_reduce_params = []
63 self.pre_direct_all_reduce_grads = []
64 self.pre_all_reduce_groups = []
65 self.pending_all_reduce_groups = []
66 self.param_group_comm_ctx = ParamGroupCommCtx()
69class HSDPSchedulerV2:
70 """HSDPScheduler is used to scheduler hsdp"""
72 def __init__(
73 self,
74 cell: Union[ModuleClass, Tuple[ModuleClass, ...]],
75 mesh,
76 reshard_after_forward,
77 shard_placement_fn,
78 mp_policy,
79 offload_policy,
80 ignored_params,
81 replicate_params,
82 device,
83 comm_fusion,
84 comm_fusion_zero_copy=False,
85 source_shard_infos: Optional[Mapping[platform.Parameter, SourceShardMetaInfo]] = None,
86 ):
87 """init hsdp scheduler.
89 Args:
90 cell: A single platform.Module or tuple of platform.Module to manage as one FSDP unit.
91 """
92 self.modules = (cell,) if isinstance(cell, platform.Module) else tuple(cell)
93 self.cell = self.modules[0]
94 self.mesh: DeviceMesh = mesh
95 self.shard_placement_fn = shard_placement_fn
96 self.mp_policy = mp_policy
97 self.offload_policy = offload_policy
98 self.comm_fusion_policy = CommFusionPolicy(comm_fusion, comm_fusion_zero_copy)
99 self.ignored_params = ignored_params
100 self.replicate_params = replicate_params
101 self.device = device
102 self.reshard_after_forward = reshard_after_forward
103 self.source_shard_infos = source_shard_infos
104 self.scheduler_state = None
105 self.forward_prefetch_cells = []
106 self.backward_prefetch_cells = []
107 self._backup_forward_fetch = None
108 # Flag to identify root module.
109 self._is_root = True
110 # module and its all sub-modules share one same 'HSDPSchedulerContext'
111 self.scheduler_ctx = HSDPSchedulerContext()
112 # When ``fully_shard`` is given multiple root modules, forward pre/post hooks coordinate
113 # so unshard / PostBackward / reshard run once per forward (aligned with PyTorch FSDP2).
114 self._fsdp_group_post_pending: Optional[set] = set() if len(self.modules) > 1 else None
115 self._init_platform()
116 self._new_cell_state()
117 self._register_hooks()
119 def _init_platform(self):
120 """Initialize the platform."""
121 raise NotImplementedError("HSDPScheduler subclasses must implement _init_platform")
123 def _new_cell_state(self):
124 """Create a new cell state."""
125 raise NotImplementedError("HSDPScheduler subclasses must implement _new_cell_state")
127 def _register_hooks(self):
128 """Register hooks."""
129 raise NotImplementedError("HSDPScheduler subclasses must implement _register_hooks.")
131 def _register_forward_backward_hooks(self):
132 """Register module forward and backward hook."""
133 raise NotImplementedError("HSDPScheduler subclasses must implement _register_forward_backward_hooks.")
135 def _get_managed_params(self):
136 """Return deduplicated parameters from all managed modules."""
137 return get_managed_modules_parameters(self.modules, self.ignored_params)
139 def set_reshard_after_forward(self, reshard_after_forward: bool) -> None:
140 """Set reshard_after_forward flag.
142 Args:
143 reshard_after_forward: Whether to reshard parameters after forward.
144 """
145 if not isinstance(reshard_after_forward, bool):
146 raise ValueError(f"reshard_after_forward should be a bool, got {type(reshard_after_forward)}")
147 self.reshard_after_forward = reshard_after_forward
149 def set_reshard_after_backward(self, reshard_after_backward: bool) -> None:
150 """Set reshard_after_backward flag.
152 Args:
153 reshard_after_backward: Whether to reshard after backward completes.
154 """
155 if not isinstance(reshard_after_backward, bool):
156 raise ValueError(f"reshard_after_backward should be a bool, got {type(reshard_after_backward)}")
157 if self.hsdp_state is not None:
158 self.hsdp_state.reshard_after_backward = reshard_after_backward
160 def set_requires_all_reduce(self, requires_all_reduce: bool) -> None:
161 """Set requires_all_reduce flag.
163 Args:
164 requires_all_reduce: Whether this unit participates in all-reduce.
165 """
166 if not isinstance(requires_all_reduce, bool):
167 raise ValueError(f"requires_all_reduce should be a bool, got {type(requires_all_reduce)}")
168 if self.hsdp_state is not None:
169 self.hsdp_state.set_requires_all_reduce(requires_all_reduce)
171 def reset_iter_state(self) -> None:
172 """Reset scheduler bookkeeping after a completed iteration."""
173 self.scheduler_ctx.root_bp_state = False
174 self.scheduler_ctx.pre_reduce_scatter_params.clear()
175 self.scheduler_ctx.pre_all_reduce_params.clear()
176 self.scheduler_ctx.pre_direct_all_reduce_grads.clear()
177 self.scheduler_ctx.pre_all_reduce_groups.clear()
178 self.scheduler_ctx.pending_all_reduce_groups.clear()
179 self.scheduler_ctx.param_group_comm_ctx.pre_param_group = None
180 self.scheduler_ctx.param_group_comm_ctx.all_reduce_param_group = None
181 self.scheduler_ctx.param_group_comm_ctx.comm_handle = None
182 self.scheduler_state = None
183 if self._fsdp_group_post_pending is not None:
184 self._fsdp_group_post_pending.clear()
185 self._restore_forward_prefetch_after_recompute()
187 def set_requires_grad_sync(self, requires_grad_sync: bool) -> None:
188 """Set flag controlling whether gradients are synchronized.
190 Args:
191 requires_grad_sync: When True, enable grad sync for this scheduler.
192 """
193 if not isinstance(requires_grad_sync, bool):
194 raise ValueError(f"requires_grad_sync should be a bool, got {type(requires_grad_sync)}")
195 self.hsdp_state.set_requires_grad_sync(requires_grad_sync)
197 # pylint: disable=W0613
198 def _hsdp_forward_pre_hook(self, cell, args, kwargs):
199 """Forward pre hook to unsharded parameter for forward process."""
200 logger.debug("hook=forward_pre enter module=%s", self.hsdp_state)
201 if self.scheduler_state == FSDPSchedulerState.PRE_BACKWARD:
202 logger.debug("hook=forward_pre skip module=%s reason=pre_backward", self.hsdp_state)
203 return args, kwargs
204 if self.scheduler_ctx.root_bp_state:
205 self._disable_forward_prefetch_for_recompute()
206 if self.scheduler_ctx.root_module is None:
207 tree_ctx = self.scheduler_ctx
208 tree_ctx.root_module = self.cell
209 registered_schedulers = set()
210 for module_name, module in platform.get_cells_and_names(tree_ctx.root_module):
211 from hyper_parallel.core.fully_shard.api import HSDPModule # pylint: disable=C0415
212 if isinstance(module, HSDPModule):
213 submod_scheduler = module.hsdp_scheduler
214 if submod_scheduler is None or id(submod_scheduler) in registered_schedulers:
215 continue
216 registered_schedulers.add(id(submod_scheduler))
217 if submod_scheduler.scheduler_ctx is not tree_ctx:
218 if submod_scheduler.scheduler_ctx.root_module is not None:
219 raise ValueError(
220 "HSDP scheduler already belongs to another initialized module tree"
221 )
222 submod_scheduler.scheduler_ctx = tree_ctx
223 submod_scheduler.hsdp_state.scheduler_ctx = tree_ctx
224 if submod_scheduler.hsdp_state.param_group is not None:
225 submod_scheduler.hsdp_state.param_group.comm_ctx = tree_ctx.param_group_comm_ctx
226 submod_scheduler._is_root = submod_scheduler is self # pylint: disable=protected-access
227 submod_scheduler.hsdp_state.module_name = module_name
228 tree_ctx.all_hsdp_schedulers.append(submod_scheduler)
230 self.scheduler_state = FSDPSchedulerState.PRE_FORWARD
231 if self._is_root and not self.scheduler_ctx.lazy_init_done:
232 self._init_params_fqn()
233 self._lazy_init_all_states()
234 self.scheduler_ctx.lazy_init_done = True
235 if self.mp_policy.cast_forward_inputs and self.mp_policy.param_dtype:
236 cast_fn = functools.partial(self.platform.cast_fp_tensor, self.mp_policy.param_dtype)
237 args = self.platform.apply_to_tensors(cast_fn, args)
238 kwargs = self.platform.apply_to_tensors(cast_fn, kwargs)
239 with self.platform.profiler_record(f"pre_forward unshard:{self.hsdp_state.module_name}"):
240 logger.debug("hook=forward_pre action=unshard module=%s", self.hsdp_state)
241 self.hsdp_state.unshard()
242 for prefetch_cell in self.forward_prefetch_cells:
243 prefetch_state = prefetch_cell.hsdp_scheduler.hsdp_state
244 with self.platform.profiler_record(f"pre_forward prefetch:"
245 f"{prefetch_state.module_name}"):
246 logger.debug(
247 "hook=forward_pre action=prefetch module=%s target=%s",
248 self.hsdp_state,
249 prefetch_state,
250 )
251 prefetch_state.prefetch()
252 return args, kwargs
254 def _lazy_init_all_states(self):
255 if self._is_root and self.scheduler_ctx.root_module is not None:
256 for submod_scheduler in self.scheduler_ctx.all_hsdp_schedulers:
257 hsdp_state = submod_scheduler.hsdp_state
258 if hsdp_state:
259 hsdp_state.lazy_init()
261 def _init_params_fqn(self): # pylint: disable=W0212
262 if not self._is_root or self.scheduler_ctx.root_module is None:
263 return
264 if self.scheduler_ctx._param_fqn_initialized: # pylint: disable=protected-access
265 return
266 # Build a map from original (sharded) parameter tensor to its HSDPParam wrapper.
267 param_to_hsdp_param = {}
268 for submod_scheduler in self.scheduler_ctx.all_hsdp_schedulers:
269 hsdp_state = submod_scheduler.hsdp_state
270 if hsdp_state is None:
271 continue
272 for hsdp_param in hsdp_state.hsdp_params:
273 orig_param = hsdp_param.sharded_param
274 # Shared parameters: keep only the first mapping to preserve the
275 # first-seen FQN (consistent with the deduplication in _init_hsdp_params).
276 if orig_param not in param_to_hsdp_param:
277 param_to_hsdp_param[orig_param] = hsdp_param
279 # Walk the full parameter tree and assign FQNs; skip params already seen
280 # (shared-parameter deduplication: first name wins).
281 visited_params = set()
282 for param_name, parameter in platform.parameters_dict(self.scheduler_ctx.root_module):
283 if parameter in visited_params:
284 continue
285 visited_params.add(parameter)
286 hsdp_param = param_to_hsdp_param.get(parameter)
287 if hsdp_param is not None:
288 hsdp_param._param_fqn = param_name # pylint: disable=W0212
289 self.scheduler_ctx._param_fqn_initialized = True # pylint: disable=protected-access
291 # pylint: disable=W0613, R1710
292 def _hsdp_forward_hook(self, cell, inputs, outputs):
293 """Forward hook to shard parameter for saving memory."""
294 logger.debug("hook=forward enter module=%s", self.hsdp_state)
295 if self.scheduler_state == FSDPSchedulerState.PRE_BACKWARD:
296 logger.debug("hook=forward skip module=%s reason=pre_backward", self.hsdp_state)
297 return
298 self.scheduler_state = FSDPSchedulerState.FORWARD
299 if self.reshard_after_forward:
300 with self.platform.profiler_record(f"forward reshard:{self.hsdp_state.module_name}"):
301 logger.debug("hook=forward action=reshard module=%s", self.hsdp_state)
302 self.hsdp_state.shard()
303 if self.mp_policy.output_dtype is not None:
304 outputs = self.platform.apply_to_tensors(
305 functools.partial(self.platform.cast_fp_tensor, self.mp_policy.output_dtype),
306 outputs,
307 )
308 return outputs
310 # pylint: disable=W0613
311 def _hsdp_backward_pre_hook(self, cell, grad_outputs):
312 """Backward pre hook to unsharded parameter for backward process."""
313 logger.debug("hook=backward_pre enter module=%s", self.hsdp_state)
314 self.scheduler_state = FSDPSchedulerState.PRE_BACKWARD
315 if self.reshard_after_forward:
316 with self.platform.profiler_record(f"pre_backward unshard:{self.hsdp_state.module_name}"):
317 logger.debug("hook=backward_pre action=unshard module=%s", self.hsdp_state)
318 self.hsdp_state.unshard()
319 for prefetch_cell in self.backward_prefetch_cells:
320 prefetch_state = prefetch_cell.hsdp_scheduler.hsdp_state
321 with self.platform.profiler_record(f"pre_backward prefetch:"
322 f"{prefetch_state.module_name}"):
323 logger.debug(
324 "hook=backward_pre action=prefetch module=%s target=%s",
325 self.hsdp_state,
326 prefetch_state,
327 )
328 prefetch_state.prefetch()
330 # pylint: disable=W0613
331 def _hsdp_backward_hook(self, cell, grad_inputs, grad_outputs):
332 """Backward hook to shard parameter for optimizer process or saving memory."""
333 logger.debug("hook=backward_hook enter module=%s", self.hsdp_state)
334 self.scheduler_state = FSDPSchedulerState.BACKWARD
335 with self.platform.profiler_record(f"post_backward:{self.hsdp_state.module_name}"):
336 logger.debug("hook=backward_hook action=post_backward module=%s", self.hsdp_state)
337 self.hsdp_state.post_backward()
338 if self._fsdp_group_post_pending is not None:
339 self._fsdp_group_post_pending.clear()
341 # pylint: disable=W0613
342 @staticmethod
343 def _grouped_forward_pre_hook_skip(cell, args, kwargs):
344 """Return value when grouped pre-forward should not run (first module already did).
346 Default matches MindSpore Cell forward pre-hooks (explicit ``(args, kwargs)``).
347 ``TorchHSDPSchedulerV2`` overrides this to return ``None`` (``nn.Module`` idiom).
348 """
349 return args, kwargs
351 @staticmethod
352 def _grouped_forward_post_hook_skip(outputs):
353 """Return value when grouped post-forward is deferred to a later module in the group.
355 Default returns ``outputs`` (MindSpore). ``TorchHSDPSchedulerV2`` overrides to ``None``.
356 """
357 return outputs
359 def _grouped_forward_pre_hook(self, cell, args, kwargs):
360 """Run FSDP pre-forward only for the first module in the group (PyTorch FSDP2-aligned)."""
361 pending = self._fsdp_group_post_pending
362 if pending is None:
363 return self._forward_pre_hook(cell, args, kwargs)
364 if len(pending) == 0:
365 pending.update(self.modules)
366 return self._forward_pre_hook(cell, args, kwargs)
367 return self._grouped_forward_pre_hook_skip(cell, args, kwargs)
369 def _make_grouped_forward_post_hook(self, mod):
370 """Build post-forward hook: last module in the group runs reshard + output backward hooks."""
372 def grouped_post_hook(cell, inputs, outputs):
373 pending = self._fsdp_group_post_pending
374 if pending is None:
375 return self._forward_hook(cell, inputs, outputs)
376 pending.discard(mod)
377 if len(pending) == 0:
378 return self._forward_hook(cell, inputs, outputs)
379 return self._grouped_forward_post_hook_skip(outputs)
381 return grouped_post_hook
383 def set_forward_prefetch_cells(self, hsdp_cell_list: List[Any]) -> None:
384 """Set cells prefetched during forward.
386 Args:
387 hsdp_cell_list: HSDP cells to prefetch ahead of forward.
388 """
389 self.forward_prefetch_cells = hsdp_cell_list
391 def set_backward_prefetch_cells(self, hsdp_cell_list: List[Any]) -> None:
392 """Set cells prefetched during backward.
394 Args:
395 hsdp_cell_list: HSDP cells to prefetch ahead of backward.
396 """
397 self.backward_prefetch_cells = hsdp_cell_list
399 def _disable_forward_prefetch_for_recompute(self) -> None:
400 """Temporarily disable forward prefetch during activation recompute."""
401 self._backup_forward_fetch = self.forward_prefetch_cells
402 self.forward_prefetch_cells = []
404 def _restore_forward_prefetch_after_recompute(self) -> bool:
405 """Restore forward prefetch list after a recompute forward hook finishes."""
406 if self._backup_forward_fetch is None:
407 return False
408 self.forward_prefetch_cells = self._backup_forward_fetch
409 self._backup_forward_fetch = None
410 return True