Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / platform / torch / fully_shard / state.py: 79%
232 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"""Torch HSDP cell state"""
16# pylint: disable=protected-access
18from collections import defaultdict
19from typing import List, Mapping, Optional
21import torch
23from hyper_parallel.tools.logging import get_logger
24from hyper_parallel.core.dtensor.dtensor import DTensor
25from hyper_parallel.core.fully_shard.hsdp_state import HSDPState
26from hyper_parallel.core.fully_shard.hsdp_utils import (
27 _get_param_module_infos,
28)
29from hyper_parallel.core.fully_shard.utils import (
30 CPUOffloadPolicy,
31 DDPMeshInfo,
32 FSDPMeshInfo,
33 HSDPMeshInfo,
34 SourceShardMetaInfo,
35)
36from hyper_parallel.platform.torch.fully_shard.param import TorchHSDPParamV2
37from hyper_parallel.platform.torch.fully_shard.param_group import HSDPParamGroup, AllReduceParamGroup
39logger = get_logger("FSDP")
42def _to_dtype_if_needed(
43 tensor: torch.Tensor, dtype: Optional[torch.dtype]
44) -> torch.Tensor:
45 """Cast tensor to the given dtype if it differs from current dtype.
47 Args:
48 tensor: The input tensor to potentially cast.
49 dtype: Target dtype. If None or same as tensor dtype, no-op.
50 """
51 if dtype is not None and tensor.dtype != dtype:
52 return tensor.to(dtype)
53 return tensor
56class TorchHSDPStateV2(HSDPState):
57 """Torch HSDP cell state"""
58 def __init__(
59 self,
60 cell,
61 mesh,
62 shard_placement_fn,
63 comm_fusion_policy,
64 mp_policy,
65 offload_policy,
66 raw_ignored_params,
67 raw_replicate_params,
68 platform,
69 scheduler_ctx,
70 device,
71 source_shard_infos: Optional[Mapping[torch.nn.Parameter, SourceShardMetaInfo]] = None,
72 ):
73 """
74 Initialize TorchHSDPStateV2.
76 Args:
77 cell (nn.Module): The module whose parameters are managed by this state.
78 mesh: Mesh topology for shard/replicate dimensions.
79 comm_fusion:
80 mp_policy:
81 offload_policy:
82 platform (TorchPlatform): Torch platform abstraction.
83 device (torch.device): Target device.
84 """
85 self.source_shard_infos = source_shard_infos
86 super().__init__(
87 cell,
88 mesh,
89 shard_placement_fn,
90 comm_fusion_policy,
91 mp_policy,
92 offload_policy,
93 raw_ignored_params,
94 raw_replicate_params,
95 platform,
96 scheduler_ctx,
97 device,
98 )
99 self._init_param_group()
101 def _init_param_group(self):
102 """Initialize fused parameter group for communication fusion.
104 All managed parameters enter one ``HSDPParamGroup``. Parameters without
105 an FSDP shard dimension use the group's local all-gather and
106 reduce-scatter paths before entering replicate all-reduce buckets.
107 """
108 self.param_group = None
109 if not self.comm_fusion_policy.enable_comm_fusion:
110 return
111 if self.hsdp_params:
112 # pylint: disable=E1128
113 self.param_group = HSDPParamGroup(
114 self.hsdp_params,
115 self.device,
116 self.comm_fusion_policy.comm_fusion_zero_copy,
117 comm_ctx=self.scheduler_ctx.param_group_comm_ctx,
118 )
120 def _move_states_to_device(self):
121 """move states to device"""
122 for mod in self.modules:
123 for param in mod.parameters():
124 if hasattr(param, "_hsdp_param_initialized") and param._hsdp_param_initialized:
125 continue
126 if param.device == self.device or param.device.type == "meta":
127 continue
128 param.data = param.to(self.device)
129 for buffer in mod.buffers():
130 if buffer.device == self.device or buffer.device.type == "meta":
131 continue
132 buffer.data = buffer.to(self.device)
134 def _build_param_source_shard_info(
135 self, param: torch.nn.Parameter
136 ) -> Optional[SourceShardMetaInfo]:
137 """Build normalized source-layout metadata for one managed parameter."""
138 if isinstance(param, DTensor):
139 if self.source_shard_infos is not None:
140 raise ValueError(
141 "source_shard_infos cannot be provided when fully_shard manages a native DTensor parameter"
142 )
143 return SourceShardMetaInfo(
144 mesh=param.device_mesh,
145 placements=tuple(param.placements),
146 origin_is_dtensor=True,
147 )
148 if self.source_shard_infos is None:
149 return None
150 return self.source_shard_infos.get(param)
152 def _init_hsdp_params(self):
153 """Initialize all fully_shard-managed parameters for the module."""
154 # all parameters in the module tree(s), deduplicated
155 visited_params = set()
156 filtered_params = []
157 for mod in self.modules:
158 for _, param in mod.named_parameters():
159 if param in self.raw_ignored_params:
160 continue
161 if hasattr(param, "_hsdp_param_initialized") and param._hsdp_param_initialized:
162 continue
163 if param in visited_params:
164 continue
165 visited_params.add(param)
166 filtered_params.append(param)
168 module_infos = _get_param_module_infos(filtered_params, tuple(self.modules))
169 for param, module_info in zip(filtered_params, module_infos):
170 self.hsdp_params.append(
171 TorchHSDPParamV2(
172 param,
173 module_info,
174 self._build_param_mesh_info(param),
175 shard_placement_fn=self.shard_placement_fn,
176 mp_policy=self.mp_policy,
177 offload_policy=self.offload_policy,
178 device=self.device,
179 source_shard_info=self._build_param_source_shard_info(param),
180 )
181 )
183 def _build_param_mesh_info(self, parameter):
184 if self.mesh.ndim not in (1, 2):
185 raise ValueError(
186 "fully_shard only supports explicit 1D DP/FSDP meshes or 2D HSDP meshes. "
187 f"Got mesh.ndim={self.mesh.ndim}."
188 )
189 if parameter in self.raw_replicate_params:
190 return DDPMeshInfo(
191 mesh=self.mesh if self.mesh.ndim == 1 else self.mesh.flatten(),
192 replicate_mesh_dim=0,
193 )
194 if self.mesh.ndim == 1:
195 return FSDPMeshInfo(mesh=self.mesh, shard_mesh_dim=0)
196 return HSDPMeshInfo(
197 mesh=self.mesh,
198 shard_mesh_dim=1,
199 replicate_mesh_dim=0,
200 )
202 def _init_mp_dtypes(self):
203 """Initialize mixed-precision dtypes for all managed parameters."""
204 for hsdp_param in self.hsdp_params:
205 hsdp_param.init_dtype_attrs(self.mp_policy)
207 def _validate_cpu_offload_params(self):
208 """Validate that all parameters are on CPU when CPU offload policy is enabled."""
209 if not isinstance(self.offload_policy, CPUOffloadPolicy):
210 return
211 hsdp_params_not_on_cpu = [
212 hsdp_param
213 for hsdp_param in self.hsdp_params
214 if hsdp_param.sharded_param.device.type != "cpu"
215 ]
216 if hsdp_params_not_on_cpu:
217 raise RuntimeError(
218 "HSDP parameters should be materialized on CPU when enabling CPU offloading. "
219 'For example, load a CPU state dict or call module.to_empty(device="cpu"). '
220 "Found following parameters on non-CPU device: "
221 f"{[(p._param_fqn, p.sharded_param.device) for p in hsdp_params_not_on_cpu]}\n"
222 )
224 def lazy_init(self):
225 """Deferred initialization: reset sharded params, validate devices, and set mixed-precision dtypes."""
226 if self.is_shard and not self._reset_sharded_params:
227 for hsdp_param in self.hsdp_params:
228 hsdp_param.reset_sharded_param()
229 self._reset_sharded_params = True
230 self._validate_no_meta_params()
231 self._validate_cpu_offload_params()
232 self._init_mp_dtypes()
234 def _validate_no_meta_params(self):
235 param_names_on_meta = [
236 hsdp_param._param_fqn
237 for hsdp_param in self.hsdp_params
238 if hsdp_param.sharded_param.device.type == "meta"
239 ]
240 if param_names_on_meta:
241 raise RuntimeError(
242 "HSDP parameters should be materialized from meta device before training, "
243 f"but the following were still on meta device: {param_names_on_meta}\n"
244 "For example, call module.to_empty(device) to materialize to device and "
245 "call module.reset_parameters() on each module to initialize values."
246 )
248 def post_backward_for_comm_fusion(self):
249 """post_backward_for_comm_fusion."""
250 logger.debug("post_backward module=%s mode=comm_fusion enter", self)
251 # Fused gradient reduction path: first apply any pending async reduction
252 # from the previous module's backward (pipelined overlap), then issue
253 # this module's fused reduce-scatter (+ all-reduce for HSDP).
254 comm_ctx = self.scheduler_ctx.param_group_comm_ctx
255 # Phase 2: save gradients for the param group whose all-reduce is done.
256 if comm_ctx.all_reduce_param_group is not None:
257 logger.debug("post_backward module=%s wait=comm_fusion_all_reduce", self)
258 comm_ctx.all_reduce_param_group.wait_all_reduce_and_save_grad()
259 comm_ctx.all_reduce_param_group = None
260 # Phase 1: wait reduce_scatter, issue async all_reduce for previous layer
261 if comm_ctx.pre_param_group is not None:
262 logger.debug("post_backward module=%s wait=comm_fusion_reduce_scatter", self)
263 comm_ctx.pre_param_group.wait_reduce_scatter_and_issue_all_reduce()
264 comm_ctx.pre_param_group = None
265 if self.param_group is not None:
266 logger.debug("post_backward module=%s launch=comm_fusion_reduce_scatter", self)
267 self.param_group.foreach_reducescatter(
268 reduce_scatter_reduce_op=self.reduce_op_type,
269 )
271 def post_backward(self, *unused): # pylint: disable=unused-argument
272 """Reduce gradients and reshard parameters after backward."""
273 logger.debug(
274 "post_backward module=%s enter reduce_grads=%s comm_fusion=%s reshard_after_backward=%s",
275 self,
276 self.reduce_grads,
277 self.comm_fusion_policy.enable_comm_fusion,
278 self.reshard_after_backward,
279 )
280 for hsdp_param in self.hsdp_params:
281 hsdp_param.accumulate_unsharded_grad_if_needed()
282 if not self.reduce_grads:
283 if self.reshard_after_backward:
284 self.shard()
285 for hsdp_param in self.hsdp_params:
286 hsdp_param.to_accumulated_grad_if_needed()
287 return
288 if self.reshard_after_backward:
289 # Reshard before gradient communication to reduce backward memory peak.
290 self.shard()
291 if not self.comm_fusion_policy.enable_comm_fusion:
292 # Step 1: wait previous reduce-scatter (for params needing all-reduce)
293 prev_group = self._wait_prev_reduce_scatter()
295 # Step 2: wait previous reduce-scatter outputs that skip replicate all-reduce
296 self._wait_prev_reduce_scatter_without_all_reduce()
298 # Step 3: issue current reduce_scatter
299 self._issue_reduce_scatter_for_current_module()
301 # Step 4: issue previous fused all-reduce asynchronously
302 self._issue_prev_fused_all_reduce(prev_group)
303 else:
304 self.post_backward_for_comm_fusion()
306 def _issue_reduce_scatter_for_current_module(self):
307 """Issue reduce_scatter for current module's parameters with fused all-reduce support.
309 This method groups parameters by their replicate_process_group and:
310 1. For params without all_reduce needs: issue reduce_scatter directly
311 2. For params with all_reduce needs: allocate fused buffer and issue reduce_scatter
312 into aligned views, enabling zero-copy fused all_reduce later.
313 """
314 # Collect parameters that need gradient reduction
315 params_to_reduce = []
316 for hsdp_param in self.hsdp_params:
317 skip_param = (
318 not hsdp_param.unsharded_param_buffers
319 or not hsdp_param.sharded_param.requires_grad
320 or (
321 hsdp_param.unsharded_param.grad is None
322 and hsdp_param.unsharded_accumulated_grad_data is None
323 )
324 )
325 if skip_param:
326 continue
327 params_to_reduce.append(hsdp_param)
329 if not params_to_reduce:
330 return
332 # Group by replicate process group and reduction dtype so every fused
333 # all-reduce buffer has one communication group and one element type.
334 groups_by_comm = defaultdict(list)
335 for hsdp_param in params_to_reduce:
336 if self.requires_all_reduce and hsdp_param.replicate_world_size > 1:
337 replicate_process_group = hsdp_param.mesh_info.replicate_process_group
338 group_key = (id(replicate_process_group), hsdp_param.reduce_comm_dtype())
339 groups_by_comm[group_key].append(hsdp_param)
340 else:
341 groups_by_comm[None].append(hsdp_param)
343 # Handle params that don't need all_reduce (FSDP or single replica)
344 if None in groups_by_comm:
345 for hsdp_param in groups_by_comm[None]:
346 logger.debug(
347 "post_backward module=%s launch=reduce_scatter param=%s all_reduce=False",
348 self,
349 hsdp_param,
350 )
351 hsdp_param.reduce_scatter_grad(
352 reduce_op=self.reduce_op_type,
353 )
354 self.scheduler_ctx.pre_reduce_scatter_params.append(hsdp_param)
356 # Handle params that need all_reduce (HSDP with multiple replicas)
357 for group_key, hsdp_params in groups_by_comm.items():
358 if group_key is None:
359 continue
361 # Create AllReduceParamGroup for fused all-reduce
362 group = AllReduceParamGroup(
363 replicate_group=hsdp_params[0].mesh_info.replicate_process_group,
364 hsdp_params=hsdp_params,
365 reduce_op=self.reduce_op_type,
366 )
368 # Allocate fused buffer with 512-byte alignment
369 group.allocate_fused_buffer(self.device)
371 # Issue reduce_scatter with output directly into fused buffer views
372 logger.debug(
373 "post_backward module=%s launch=fused_reduce_scatter group_params=%s",
374 self,
375 hsdp_params,
376 )
377 for idx, hsdp_param in enumerate(hsdp_params):
378 buffer_view = group.get_param_buffer_view(idx)
379 hsdp_param.reduce_scatter_grad(
380 reduce_op=self.reduce_op_type,
381 output_buffer=buffer_view,
382 )
384 # Save the group so the next module hook can wait RS and launch AR.
385 self.scheduler_ctx.pre_all_reduce_groups.append(group)
387 def _wait_prev_reduce_scatter(self) -> List[AllReduceParamGroup]:
388 """Step 1: wait prev reduce_scatter.
390 This enables overlapping:
391 - Layer N-1's reduce_scatter wait with Layer N's backward compute
393 Returns:
394 List of previous AllReduceParamGroups (one per communication group).
395 """
396 if self.scheduler_ctx.pre_all_reduce_groups:
397 prev_groups = list(self.scheduler_ctx.pre_all_reduce_groups)
398 self.scheduler_ctx.pre_all_reduce_groups.clear()
399 for prev_group in prev_groups:
400 logger.debug(
401 "post_backward module=%s wait=fused_reduce_scatter group_params=%s",
402 self,
403 prev_group.hsdp_params,
404 )
405 for hsdp_param in prev_group.hsdp_params:
406 hsdp_param.reduce_scatter_output()
407 hsdp_param.clear_reduce_scatter_output()
408 if hsdp_param.unsharded_accumulated_grad_data is not None:
409 hsdp_param.unsharded_accumulated_grad = None
410 elif hsdp_param.unsharded_param.grad is not None:
411 hsdp_param.unsharded_param.grad = None
412 return prev_groups
413 return []
415 def _issue_prev_fused_all_reduce(self, prev_groups: List[AllReduceParamGroup]) -> None:
416 """Step 4: issue the previous module's fused all-reduce asynchronously.
418 The all-reduce work is collected in ``pending_all_reduce_groups``
419 and is waited in the root backward hook.
421 Args:
422 prev_groups: Previous parameter groups whose all-reduce should be issued.
423 """
424 for prev_group in prev_groups:
425 prev_group.accumulate_reduce_partial_outputs()
426 logger.debug(
427 "post_backward module=%s launch=fused_all_reduce group_params=%s",
428 self,
429 prev_group.hsdp_params,
430 )
431 prev_group.issue_async_allreduce()
432 self.scheduler_ctx.pending_all_reduce_groups.append(prev_group)
434 def _wait_prev_reduce_scatter_without_all_reduce(self) -> None:
435 """Wait previous RS outputs that do not enter a replicate all-reduce.
437 When the current micro-step disables all-reduce, outputs accumulate in
438 ``reduce_partial_output`` without being cast or applied to the parameter.
439 On the final synchronized micro-step, the partial result is merged into
440 the current RS output and retained for root-hook finalization.
441 """
442 while self.scheduler_ctx.pre_reduce_scatter_params:
443 pre_hsdp_param = self.scheduler_ctx.pre_reduce_scatter_params.pop(0)
444 logger.debug(
445 "post_backward module=%s wait=reduce_scatter param=%s",
446 self,
447 pre_hsdp_param,
448 )
449 reduced_grad = pre_hsdp_param.reduce_scatter_output()
450 if not self.requires_all_reduce:
451 if pre_hsdp_param.reduce_partial_output is None:
452 pre_hsdp_param.reduce_partial_output = reduced_grad
453 else:
454 pre_hsdp_param.reduce_partial_output.add_(reduced_grad)
455 pre_hsdp_param.clear_reduce_scatter_output()
456 elif pre_hsdp_param.reduce_partial_output is not None:
457 reduced_grad.add_(pre_hsdp_param.reduce_partial_output)
458 pre_hsdp_param.reduce_partial_output = None
460 if pre_hsdp_param.unsharded_accumulated_grad_data is not None:
461 pre_hsdp_param.unsharded_accumulated_grad = None
462 elif pre_hsdp_param.unsharded_param.grad is not None:
463 pre_hsdp_param.unsharded_param.grad = None
465 def wait_and_split_all_reduce_work_groups(self) -> None:
466 """Wait fused all-reduce work and expose each parameter result."""
467 for group in self.scheduler_ctx.pending_all_reduce_groups:
468 logger.debug(
469 "post_backward module=%s wait=fused_all_reduce group_params=%s",
470 self,
471 group.hsdp_params,
472 )
473 group.wait_and_split_grads()
474 self.scheduler_ctx.pending_all_reduce_groups.clear()
476 def reset_iter_state(self) -> None:
477 """Clear Torch communication bookkeeping without clearing optimizer gradients."""
478 self.scheduler_ctx.pre_reduce_scatter_params.clear()
479 self.scheduler_ctx.pre_all_reduce_params.clear()
480 self.scheduler_ctx.pre_all_reduce_groups.clear()
481 self.scheduler_ctx.pending_all_reduce_groups.clear()
482 if self.param_group is not None:
483 self.param_group.reset_iter_state()
484 for hsdp_param in self.hsdp_params:
485 hsdp_param.allgather_comm_ctx.allgather_handle = None
486 hsdp_param.allgather_comm_ctx.allgather_output = None
487 hsdp_param.reduce_scatter_comm_ctx.reduce_scatter_handle = None
488 hsdp_param.reduce_scatter_comm_ctx.reduce_scatter_output = None
489 hsdp_param.all_reduce_comm_ctx.all_reduce_handle = None
490 hsdp_param.all_reduce_comm_ctx.all_reduce_output = None
491 hsdp_param.reduce_partial_output = None
492 hsdp_param.unsharded_accumulated_grad = None
493 hsdp_param._grad = None
494 if hsdp_param.unsharded_param_buffers:
495 hsdp_param.unsharded_param.grad = None
497 def set_requires_grad_sync(self, requires_grad_sync: bool) -> None:
498 """set requires grad sync flag to control gradient sync."""
499 self.reduce_grads = requires_grad_sync
501 def set_reduce_op_type(self, reduce_op_type: str) -> None:
502 """set reduce op type for gradient reduction."""
503 fsdp_support_reduce_op = {
504 "sum": torch.distributed.ReduceOp.SUM,
505 "avg": torch.distributed.ReduceOp.AVG,
506 }
507 reduce_op = reduce_op_type.lower().strip() if isinstance(reduce_op_type, str) else reduce_op_type
508 reduce_op_value = fsdp_support_reduce_op.get(reduce_op)
509 if reduce_op_value is None:
510 raise ValueError(
511 f"Unsupported reduce op type {reduce_op_type}, "
512 f"supported types are {list(fsdp_support_reduce_op.keys())}"
513 )
514 self.reduce_op_type = reduce_op_value
516 def _sync_current_stream_if_needed(self, need_synchronize):
517 if need_synchronize:
518 if self.device.type == "npu":
519 torch.npu.current_stream().synchronize()
520 elif self.device.type == "cuda":
521 torch.cuda.current_stream().synchronize()
522 else:
523 raise NotImplementedError(
524 f"Unsupported device type {self.device.type} for synchronization after CPU offload."
525 )