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

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 

18 

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 

27 

28logger = get_logger("FSDP") 

29 

30platform = get_platform() 

31ModuleClass = platform.Module 

32 

33 

34class ParamGroupCommCtx: 

35 """Track the in-flight parameter-group communication for one module tree.""" 

36 

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 

42 

43 

44class HSDPSchedulerContext: 

45 """Share scheduler and backward-pipeline state within one HSDP module tree.""" 

46 

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

67 

68 

69class HSDPSchedulerV2: 

70 """HSDPScheduler is used to scheduler hsdp""" 

71 

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. 

88 

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

118 

119 def _init_platform(self): 

120 """Initialize the platform.""" 

121 raise NotImplementedError("HSDPScheduler subclasses must implement _init_platform") 

122 

123 def _new_cell_state(self): 

124 """Create a new cell state.""" 

125 raise NotImplementedError("HSDPScheduler subclasses must implement _new_cell_state") 

126 

127 def _register_hooks(self): 

128 """Register hooks.""" 

129 raise NotImplementedError("HSDPScheduler subclasses must implement _register_hooks.") 

130 

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.") 

134 

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) 

138 

139 def set_reshard_after_forward(self, reshard_after_forward: bool) -> None: 

140 """Set reshard_after_forward flag. 

141 

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 

148 

149 def set_reshard_after_backward(self, reshard_after_backward: bool) -> None: 

150 """Set reshard_after_backward flag. 

151 

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 

159 

160 def set_requires_all_reduce(self, requires_all_reduce: bool) -> None: 

161 """Set requires_all_reduce flag. 

162 

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) 

170 

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

186 

187 def set_requires_grad_sync(self, requires_grad_sync: bool) -> None: 

188 """Set flag controlling whether gradients are synchronized. 

189 

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) 

196 

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) 

229 

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 

253 

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

260 

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 

278 

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 

290 

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 

309 

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

329 

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

340 

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

345 

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 

350 

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. 

354 

355 Default returns ``outputs`` (MindSpore). ``TorchHSDPSchedulerV2`` overrides to ``None``. 

356 """ 

357 return outputs 

358 

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) 

368 

369 def _make_grouped_forward_post_hook(self, mod): 

370 """Build post-forward hook: last module in the group runs reshard + output backward hooks.""" 

371 

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) 

380 

381 return grouped_post_hook 

382 

383 def set_forward_prefetch_cells(self, hsdp_cell_list: List[Any]) -> None: 

384 """Set cells prefetched during forward. 

385 

386 Args: 

387 hsdp_cell_list: HSDP cells to prefetch ahead of forward. 

388 """ 

389 self.forward_prefetch_cells = hsdp_cell_list 

390 

391 def set_backward_prefetch_cells(self, hsdp_cell_list: List[Any]) -> None: 

392 """Set cells prefetched during backward. 

393 

394 Args: 

395 hsdp_cell_list: HSDP cells to prefetch ahead of backward. 

396 """ 

397 self.backward_prefetch_cells = hsdp_cell_list 

398 

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

403 

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