Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / platform / mindspore / fully_shard / state.py: 62%
219 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 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"""MindSpore HSDP state aligned with the Torch fully_shard lifecycle."""
17from collections import defaultdict
18from typing import List
20import mindspore as ms
22from hyper_parallel.core.fully_shard.hsdp_state import HSDPState
23from hyper_parallel.core.fully_shard.hsdp_utils import _get_param_module_infos, unwrap_dtensor_param
24from hyper_parallel.core.fully_shard.utils import (
25 CPUOffloadPolicy,
26 DDPMeshInfo,
27 FSDPMeshInfo,
28 HSDPMeshInfo,
29 SourceShardMetaInfo,
30)
31from hyper_parallel.platform.mindspore.fully_shard.param import MindSporeHSDPParamV2
32from hyper_parallel.platform.mindspore.fully_shard.param_group import (
33 AllReduceParamGroup,
34 HSDPParamGroup,
35)
36from hyper_parallel.platform.mindspore.utils import normalize_runtime_device
37from hyper_parallel.tools.logging import get_logger
39logger = get_logger("FSDP")
42def _to_dtype_if_needed(tensor: ms.Tensor, dtype: ms.Type | None) -> ms.Tensor:
43 """Cast ``tensor`` only when a different MindSpore dtype is requested."""
44 if isinstance(dtype, ms.Type) and tensor.dtype != dtype:
45 return tensor.to(dtype)
46 return tensor
49class MindSporeHSDPStateV2(HSDPState):
50 """Own MindSpore fully_shard parameters and gradient communication state."""
52 def __init__(
53 self,
54 cell,
55 mesh,
56 shard_placement_fn,
57 comm_fusion_policy,
58 mp_policy,
59 offload_policy,
60 raw_ignored_params,
61 raw_replicate_params,
62 platform,
63 scheduler_ctx,
64 device=None,
65 ):
66 super().__init__(
67 cell,
68 mesh,
69 shard_placement_fn,
70 comm_fusion_policy,
71 mp_policy,
72 offload_policy,
73 raw_ignored_params,
74 raw_replicate_params,
75 platform,
76 scheduler_ctx,
77 device,
78 )
79 self._init_param_group()
81 def _init_param_group(self) -> None:
82 """Initialize fused communication for the single managed parameter list."""
83 self.param_group = None
84 if not self.comm_fusion_policy.enable_comm_fusion or not self.hsdp_params:
85 return
86 self.param_group = HSDPParamGroup(
87 self.hsdp_params,
88 self.device,
89 self.mp_policy,
90 # MindSpore does not provide the view + in-place semantics needed by
91 # the Torch parameter-group zero-copy path. Keep the safe copy path.
92 False,
93 comm_ctx=self.scheduler_ctx.param_group_comm_ctx,
94 )
96 def _move_states_to_device(self) -> None:
97 """Move parameters and buffers to the configured runtime device."""
98 for module in self.modules:
99 for param in module.get_parameters():
100 if getattr(param, "_hsdp_param_initialized", False):
101 continue
102 param_device = normalize_runtime_device(param.device)
103 if param_device in (self.device, "meta"):
104 continue
105 param.data = param.to(self.device)
106 for buffer in module.buffers():
107 if normalize_runtime_device(buffer.device) in (self.device, "meta"):
108 continue
109 buffer.data = buffer.to(self.device)
111 @staticmethod
112 def _build_param_source_shard_info(param):
113 """Build normalized source-layout metadata for a native DTensor parameter."""
114 dtensor_payload = unwrap_dtensor_param(param)
115 if dtensor_payload is None:
116 return None
117 return SourceShardMetaInfo(
118 mesh=dtensor_payload.device_mesh,
119 placements=tuple(dtensor_payload.placements),
120 origin_is_dtensor=True,
121 )
123 def _init_hsdp_params(self) -> None:
124 """Initialize all fully_shard-managed parameters for this module unit."""
125 visited_params = set()
126 filtered_params = []
127 for module in self.modules:
128 for _, param in module.parameters_and_names():
129 if param in self.raw_ignored_params:
130 continue
131 if getattr(param, "_hsdp_param_initialized", False):
132 continue
133 if param in visited_params:
134 continue
135 visited_params.add(param)
136 filtered_params.append(param)
138 module_infos = _get_param_module_infos(filtered_params, tuple(self.modules))
139 for param, module_info in zip(filtered_params, module_infos):
140 self.hsdp_params.append(
141 MindSporeHSDPParamV2(
142 param,
143 module_info,
144 self._build_param_mesh_info(param),
145 shard_placement_fn=self.shard_placement_fn,
146 mp_policy=self.mp_policy,
147 offload_policy=self.offload_policy,
148 device=self.device,
149 source_shard_info=self._build_param_source_shard_info(param),
150 )
151 )
153 def _build_param_mesh_info(self, parameter):
154 """Return the parameter-specific data-parallel route."""
155 if self.mesh.ndim not in (1, 2):
156 raise ValueError(
157 "fully_shard only supports explicit 1D DP/FSDP meshes or 2D HSDP meshes. "
158 f"Got mesh.ndim={self.mesh.ndim}."
159 )
160 if parameter in self.raw_replicate_params:
161 return DDPMeshInfo(
162 mesh=self.mesh if self.mesh.ndim == 1 else self.mesh.flatten(),
163 replicate_mesh_dim=0,
164 )
165 if self.mesh.ndim == 1:
166 return FSDPMeshInfo(mesh=self.mesh, shard_mesh_dim=0)
167 return HSDPMeshInfo(
168 mesh=self.mesh,
169 shard_mesh_dim=1,
170 replicate_mesh_dim=0,
171 )
173 def _init_mp_dtypes(self) -> None:
174 """Initialize mixed-precision metadata for all managed parameters."""
175 for hsdp_param in self.hsdp_params:
176 hsdp_param.init_dtype_attrs(self.mp_policy)
178 def _validate_cpu_offload_params(self) -> None:
179 """Validate CPU placement when CPU offload is configured."""
180 if not isinstance(self.offload_policy, CPUOffloadPolicy):
181 return
182 params_not_on_cpu = [
183 hsdp_param
184 for hsdp_param in self.hsdp_params
185 if not str(hsdp_param.sharded_param.device).lower().startswith("cpu")
186 ]
187 if params_not_on_cpu:
188 raise RuntimeError(
189 "HSDP parameters should be materialized on CPU when enabling CPU offloading. "
190 "Found following parameters on non-CPU device: "
191 f"{[(p._param_fqn, p.sharded_param.device) for p in params_not_on_cpu]}\n"
192 "MindSpore backend will support this feature in future version."
193 )
195 def lazy_init(self) -> None:
196 """Refresh parameter views and validate runtime state before execution."""
197 if self.is_shard and not self._reset_sharded_params:
198 for hsdp_param in self.hsdp_params:
199 hsdp_param.reset_sharded_param()
200 self._reset_sharded_params = True
201 self._validate_no_meta_params()
202 self._validate_cpu_offload_params()
203 self._init_mp_dtypes()
205 def _validate_no_meta_params(self) -> None:
206 """Validate that managed parameters have been materialized."""
207 param_names_on_meta = [
208 hsdp_param._param_fqn
209 for hsdp_param in self.hsdp_params
210 if normalize_runtime_device(hsdp_param.sharded_param.device) == "meta"
211 ]
212 if param_names_on_meta:
213 raise RuntimeError(
214 "HSDP parameters should be materialized from meta device before training, "
215 f"but the following were still on meta device: {param_names_on_meta}\n"
216 "For example, initialize the module weights on a real device before running training."
217 )
219 def zero_grad(self) -> None:
220 """Clear gradients for all managed parameters."""
221 for hsdp_param in self.hsdp_params:
222 hsdp_param.zero_grad()
224 def post_backward_for_comm_fusion(self) -> None:
225 """Pipeline fused reduce-scatter and all-reduce communication."""
226 logger.debug("post_backward module=%s mode=comm_fusion enter", self)
227 comm_ctx = self.scheduler_ctx.param_group_comm_ctx
228 if comm_ctx.all_reduce_param_group is not None:
229 logger.debug("post_backward module=%s wait=comm_fusion_all_reduce", self)
230 comm_ctx.all_reduce_param_group.wait_all_reduce_and_save_grad()
231 comm_ctx.all_reduce_param_group = None
232 if comm_ctx.pre_param_group is not None:
233 logger.debug("post_backward module=%s wait=comm_fusion_reduce_scatter", self)
234 comm_ctx.pre_param_group.wait_reduce_scatter_and_issue_all_reduce()
235 comm_ctx.pre_param_group = None
236 if self.param_group is not None:
237 logger.debug("post_backward module=%s launch=comm_fusion_reduce_scatter", self)
238 self.param_group.foreach_reducescatter(
239 reduce_scatter_reduce_op=self.reduce_op_type,
240 )
242 def _resolve_reduce_op(self) -> str:
243 """Return the active mint reduction operation."""
244 return self.reduce_op_type
246 def post_backward(self, *unused) -> None: # pylint: disable=unused-argument
247 """Accumulate gradients, reshard parameters, and launch reductions."""
248 logger.debug(
249 "post_backward module=%s enter reduce_grads=%s comm_fusion=%s reshard_after_backward=%s",
250 self,
251 self.reduce_grads,
252 self.comm_fusion_policy.enable_comm_fusion,
253 self.reshard_after_backward,
254 )
255 for hsdp_param in self.hsdp_params:
256 hsdp_param.accumulate_unsharded_grad_if_needed()
257 if not self.reduce_grads:
258 if self.reshard_after_backward:
259 self.shard()
260 for hsdp_param in self.hsdp_params:
261 hsdp_param.to_accumulated_grad_if_needed()
262 return
263 if self.reshard_after_backward:
264 self.shard()
265 if self.comm_fusion_policy.enable_comm_fusion:
266 self.post_backward_for_comm_fusion()
267 return
269 previous_groups = self._wait_prev_reduce_scatter()
270 self._wait_prev_reduce_scatter_without_all_reduce()
271 self._issue_reduce_scatter_for_current_module()
272 self._issue_prev_fused_all_reduce(previous_groups)
274 def _issue_reduce_scatter_for_current_module(self) -> None:
275 """Issue per-parameter reduce-scatter and fuse compatible HSDP all-reduces."""
276 params_to_reduce = []
277 for hsdp_param in self.hsdp_params:
278 skip_param = (
279 not hsdp_param.unsharded_param_buffers
280 or not hsdp_param.sharded_param.requires_grad
281 or (
282 hsdp_param.unsharded_param.grad is None
283 and hsdp_param.unsharded_accumulated_grad_data is None
284 )
285 )
286 if not skip_param:
287 params_to_reduce.append(hsdp_param)
288 if not params_to_reduce:
289 return
291 groups_by_comm = defaultdict(list)
292 for hsdp_param in params_to_reduce:
293 if self.requires_all_reduce and hsdp_param.replicate_world_size > 1:
294 replicate_group = hsdp_param.mesh_info.replicate_process_group
295 group_key = (replicate_group, hsdp_param.reduce_comm_dtype())
296 groups_by_comm[group_key].append(hsdp_param)
297 else:
298 groups_by_comm[None].append(hsdp_param)
300 for hsdp_param in groups_by_comm.get(None, ()):
301 logger.debug(
302 "post_backward module=%s launch=reduce_scatter param=%s all_reduce=False",
303 self,
304 hsdp_param,
305 )
306 hsdp_param.reduce_scatter_grad(reduce_op=self.reduce_op_type)
307 self.scheduler_ctx.pre_reduce_scatter_params.append(hsdp_param)
309 for group_key, hsdp_params in groups_by_comm.items():
310 if group_key is None:
311 continue
312 group = AllReduceParamGroup(
313 replicate_group=hsdp_params[0].mesh_info.replicate_process_group,
314 hsdp_params=hsdp_params,
315 reduce_op=self.reduce_op_type,
316 )
317 logger.debug(
318 "post_backward module=%s launch=fused_reduce_scatter group_params=%s",
319 self,
320 hsdp_params,
321 )
322 for hsdp_param in hsdp_params:
323 hsdp_param.reduce_scatter_grad(
324 reduce_op=self.reduce_op_type,
325 )
326 self.scheduler_ctx.pre_all_reduce_groups.append(group)
328 def _wait_prev_reduce_scatter(self) -> List[AllReduceParamGroup]:
329 """Wait previous fused reduce-scatter groups before all-reduce."""
330 if not self.scheduler_ctx.pre_all_reduce_groups:
331 return []
332 previous_groups = list(self.scheduler_ctx.pre_all_reduce_groups)
333 self.scheduler_ctx.pre_all_reduce_groups.clear()
334 for previous_group in previous_groups:
335 logger.debug(
336 "post_backward module=%s wait=fused_reduce_scatter group_params=%s",
337 self,
338 previous_group.hsdp_params,
339 )
340 for hsdp_param in previous_group.hsdp_params:
341 hsdp_param.reduce_scatter_output()
342 return previous_groups
344 def _issue_prev_fused_all_reduce(self, previous_groups: List[AllReduceParamGroup]) -> None:
345 """Launch the previous module's fused all-reduce asynchronously."""
346 for previous_group in previous_groups:
347 previous_group.accumulate_reduce_partial_outputs()
348 logger.debug(
349 "post_backward module=%s launch=fused_all_reduce group_params=%s",
350 self,
351 previous_group.hsdp_params,
352 )
353 previous_group.issue_async_allreduce()
354 self.scheduler_ctx.pending_all_reduce_groups.append(previous_group)
356 def _wait_prev_reduce_scatter_without_all_reduce(self) -> None:
357 """Wait reduce-scatter outputs that do not enter a DP all-reduce."""
358 while self.scheduler_ctx.pre_reduce_scatter_params:
359 hsdp_param = self.scheduler_ctx.pre_reduce_scatter_params.pop(0)
360 logger.debug(
361 "post_backward module=%s wait=reduce_scatter param=%s",
362 self,
363 hsdp_param,
364 )
365 reduced_grad = hsdp_param.reduce_scatter_output()
366 if not self.requires_all_reduce:
367 if hsdp_param.reduce_partial_output is None:
368 hsdp_param.reduce_partial_output = reduced_grad
369 else:
370 hsdp_param.reduce_partial_output = ms.mint.add(
371 hsdp_param.reduce_partial_output,
372 reduced_grad,
373 )
374 hsdp_param.clear_reduce_scatter_output()
375 elif hsdp_param.reduce_partial_output is not None:
376 reduced_grad = ms.mint.add(
377 reduced_grad,
378 hsdp_param.reduce_partial_output,
379 )
380 hsdp_param.reduce_scatter_comm_ctx.reduce_scatter_output = reduced_grad
381 hsdp_param.reduce_partial_output = None
382 else:
383 hsdp_param.reduce_scatter_comm_ctx.reduce_scatter_output = reduced_grad
384 hsdp_param.clear_unsharded_source_grad()
386 def wait_and_split_all_reduce_work_groups(self) -> None:
387 """Wait fused all-reduce work and expose each parameter result."""
388 for group in self.scheduler_ctx.pending_all_reduce_groups:
389 logger.debug(
390 "post_backward module=%s wait=fused_all_reduce group_params=%s",
391 self,
392 group.hsdp_params,
393 )
394 group.wait_and_split_grads()
395 self.scheduler_ctx.pending_all_reduce_groups.clear()
397 def reset_iter_state(self) -> None:
398 """Clear communication bookkeeping without clearing optimizer gradients."""
399 self.scheduler_ctx.pre_reduce_scatter_params.clear()
400 self.scheduler_ctx.pre_all_reduce_params.clear()
401 self.scheduler_ctx.pre_direct_all_reduce_grads.clear()
402 self.scheduler_ctx.pre_all_reduce_groups.clear()
403 self.scheduler_ctx.pending_all_reduce_groups.clear()
404 if self.param_group is not None:
405 self.param_group.reset_iter_state()
406 for hsdp_param in self.hsdp_params:
407 hsdp_param.allgather_comm_ctx.allgather_input = None
408 hsdp_param.allgather_comm_ctx.allgather_output = None
409 hsdp_param.allgather_comm_ctx.allgather_handle = None
410 hsdp_param.reduce_scatter_comm_ctx.reduce_scatter_handle = None
411 hsdp_param.clear_reduce_scatter_output()
412 hsdp_param.all_reduce_comm_ctx.all_reduce_handle = None
413 hsdp_param.clear_all_reduce_output()
414 hsdp_param.reduce_partial_output = None
415 hsdp_param.clear_unsharded_source_grad()
417 def set_requires_grad_sync(self, requires_grad_sync: bool) -> None:
418 """Set whether this state synchronizes gradients in backward."""
419 self.reduce_grads = requires_grad_sync
421 def set_reduce_op_type(self, reduce_op_type: str) -> None:
422 """Set the reduction operation accepted by ``mindspore.mint.distributed``."""
423 reduce_op = reduce_op_type.lower().strip()
424 if reduce_op not in ("sum", "avg"):
425 raise ValueError(
426 f"Unsupported reduce op type {reduce_op_type}, supported types are ['sum', 'avg']"
427 )
428 self.reduce_op_type = reduce_op
430 @staticmethod
431 def _sync_current_stream_if_needed(need_synchronize: bool) -> None:
432 """Synchronize after a non-blocking CPU-offload copy when required."""
433 if need_synchronize:
434 ms.runtime.current_stream().synchronize()