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

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 

17 

18from collections import defaultdict 

19from typing import List, Mapping, Optional 

20 

21import torch 

22 

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 

38 

39logger = get_logger("FSDP") 

40 

41 

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. 

46 

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 

54 

55 

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. 

75 

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() 

100 

101 def _init_param_group(self): 

102 """Initialize fused parameter group for communication fusion. 

103 

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 ) 

119 

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) 

133 

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) 

151 

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) 

167 

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 ) 

182 

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 ) 

201 

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) 

206 

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 ) 

223 

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() 

233 

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 ) 

247 

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 ) 

270 

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() 

294 

295 # Step 2: wait previous reduce-scatter outputs that skip replicate all-reduce 

296 self._wait_prev_reduce_scatter_without_all_reduce() 

297 

298 # Step 3: issue current reduce_scatter 

299 self._issue_reduce_scatter_for_current_module() 

300 

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() 

305 

306 def _issue_reduce_scatter_for_current_module(self): 

307 """Issue reduce_scatter for current module's parameters with fused all-reduce support. 

308 

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) 

328 

329 if not params_to_reduce: 

330 return 

331 

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) 

342 

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) 

355 

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 

360 

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 ) 

367 

368 # Allocate fused buffer with 512-byte alignment 

369 group.allocate_fused_buffer(self.device) 

370 

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 ) 

383 

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) 

386 

387 def _wait_prev_reduce_scatter(self) -> List[AllReduceParamGroup]: 

388 """Step 1: wait prev reduce_scatter. 

389 

390 This enables overlapping: 

391 - Layer N-1's reduce_scatter wait with Layer N's backward compute 

392 

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 [] 

414 

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. 

417 

418 The all-reduce work is collected in ``pending_all_reduce_groups`` 

419 and is waited in the root backward hook. 

420 

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) 

433 

434 def _wait_prev_reduce_scatter_without_all_reduce(self) -> None: 

435 """Wait previous RS outputs that do not enter a replicate all-reduce. 

436 

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 

459 

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 

464 

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() 

475 

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 

496 

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 

500 

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 

515 

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 )