Coverage for  / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / core / fully_shard / hsdp_scheduler.py: 83%

201 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-08-22 04:23 +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, 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.hsdp_utils import ( 

22 FSDPSchedulerState, 

23 HSDPConfigV2, 

24 get_managed_modules_parameters, 

25 get_hsdp_state 

26) 

27from hyper_parallel.tools.logging import get_logger 

28 

29logger = get_logger("FSDP") 

30 

31platform = get_platform() 

32 

33 

34class HSDPSchedulerContext: 

35 """HSDPSchedulerContext""" 

36 

37 def __init__(self) -> None: 

38 # Currently only record is_last_backward flag for scheduler context. 

39 self.is_last_backward: bool = True 

40 # flag to identify "root_module" 

41 self.root_module = None 

42 # Compile tracing may enter the root more than once; initialize shared 

43 # parameter state only on the first real forward. 

44 self.lazy_init_done: bool = False 

45 

46 

47class HSDPSchedulerV2: 

48 """HSDPScheduler is used to scheduler hsdp""" 

49 root_bp_state = False 

50 

51 

52 def __init__(self, cell: Union[platform.Module, Tuple[platform.Module, ...]], mesh, 

53 reshard_after_forward, shard_placement_fn, 

54 mp_policy, offload_policy, ignored_params, replicate_params, device, comm_fusion, 

55 comm_fusion_zero_copy=False): 

56 """init hsdp scheduler. 

57 

58 Args: 

59 cell: A single platform.Module or tuple of platform.Module to manage as one FSDP unit. 

60 """ 

61 self.modules = (cell,) if isinstance(cell, platform.Module) else tuple(cell) 

62 self.cell = self.modules[0] 

63 self.mesh: DeviceMesh = mesh 

64 self.reshard_after_forward = reshard_after_forward 

65 self.shard_placement_fn = shard_placement_fn 

66 self.mp_policy = mp_policy 

67 self.offload_policy = offload_policy 

68 self.ignored_params = ignored_params 

69 self.replicate_params = replicate_params 

70 self.device = device 

71 self.scheduler_state = None 

72 self.forward_prefetch_cells = [] 

73 self.backward_prefetch_cells = [] 

74 self._backup_forward_fetch = None 

75 # Flag to identify root module. 

76 self._is_root = False 

77 # module and its all sub-modules share one same 'HSDPSchedulerContext' 

78 self.scheduler_ctx = HSDPSchedulerContext() 

79 # When ``fully_shard`` is given multiple root modules, forward pre/post hooks coordinate 

80 # so unshard / PostBackward / reshard run once per forward (aligned with PyTorch FSDP2). 

81 self._fsdp_group_post_pending: Optional[set] = set() if len(self.modules) > 1 else None 

82 self.config = HSDPConfigV2( 

83 mesh, 

84 reshard_after_forward, 

85 shard_placement_fn, 

86 mp_policy, 

87 offload_policy, 

88 ignored_params, 

89 replicate_params, 

90 comm_fusion=comm_fusion, 

91 comm_fusion_zero_copy=comm_fusion_zero_copy, 

92 ) 

93 self._init_platform() 

94 self._new_cell_state() 

95 self._register_hooks() 

96 

97 def _init_platform(self): 

98 """Initialize the platform.""" 

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

100 

101 def _new_cell_state(self): 

102 """Create a new cell state.""" 

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

104 

105 def _register_hooks(self): 

106 """Register hooks.""" 

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

108 

109 def _register_forward_backward_hooks(self): 

110 """Register module forward and backward hook.""" 

111 raise NotImplementedError("HSDPScheduler subclasses must implement _register_forward_backward_hooks.") 

112 

113 def _get_managed_params(self): 

114 """Return deduplicated parameters from all managed modules.""" 

115 return get_managed_modules_parameters(self.modules, self.ignored_params) 

116 

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

118 """Set reshard_after_forward flag. 

119 

120 Args: 

121 reshard_after_forward: Whether to reshard parameters after forward. 

122 """ 

123 if not isinstance(reshard_after_forward, bool): 

124 raise ValueError(f"reshard_after_forward should be a bool, got {type(reshard_after_forward)}") 

125 self.reshard_after_forward = reshard_after_forward 

126 self.config.reshard_after_forward = reshard_after_forward 

127 

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

129 """Set reshard_after_backward flag. 

130 

131 Args: 

132 reshard_after_backward: Whether to reshard after backward completes. 

133 """ 

134 if not isinstance(reshard_after_backward, bool): 

135 raise ValueError(f"reshard_after_backward should be a bool, got {type(reshard_after_backward)}") 

136 if self.hsdp_state is not None: 

137 self.hsdp_state.reshard_after_backward = reshard_after_backward 

138 

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

140 """Set requires_all_reduce flag. 

141 

142 Args: 

143 requires_all_reduce: Whether this unit participates in all-reduce. 

144 """ 

145 if not isinstance(requires_all_reduce, bool): 

146 raise ValueError(f"requires_all_reduce should be a bool, got {type(requires_all_reduce)}") 

147 if self.hsdp_state is not None: 

148 self.hsdp_state.requires_all_reduce = requires_all_reduce 

149 

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

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

152 

153 Args: 

154 requires_grad_sync: When True, enable grad sync for this scheduler. 

155 """ 

156 if not isinstance(requires_grad_sync, bool): 

157 raise ValueError(f"requires_grad_sync should be a bool, got {type(requires_grad_sync)}") 

158 self.hsdp_state.set_requires_grad_sync(requires_grad_sync) 

159 

160 # pylint: disable=W0613 

161 def _hsdp_forward_pre_hook(self, cell, args, kwargs): 

162 """Forward pre hook to unsharded parameter for forward process.""" 

163 logger.debug("hook=forward_pre enter module=%s", self.hsdp_state) 

164 if self.scheduler_state == FSDPSchedulerState.PRE_BACKWARD: 

165 logger.debug("hook=forward_pre skip module=%s reason=pre_backward", self.hsdp_state) 

166 return args, kwargs 

167 if HSDPSchedulerV2.root_bp_state: 

168 self._disable_forward_prefetch_for_recompute() 

169 if self.scheduler_ctx.root_module is None: 

170 self.scheduler_ctx.root_module = self.cell 

171 self._is_root = True 

172 for _, module in platform.get_cells_and_names(self.scheduler_ctx.root_module): 

173 from hyper_parallel.core.fully_shard.api import HSDPModule # pylint: disable=C0415 

174 if isinstance(module, HSDPModule): 

175 submod_scheduler = getattr(module, "hsdp_scheduler", None) 

176 if submod_scheduler and submod_scheduler.scheduler_ctx is not self.scheduler_ctx: 

177 submod_scheduler.scheduler_ctx = self.scheduler_ctx 

178 

179 if not self._is_root and not self.hsdp_state.module_name: 

180 for module_name, module in platform.get_cells_and_names(self.scheduler_ctx.root_module): 

181 if module == self.cell: 

182 self.hsdp_state.module_name = module_name 

183 break 

184 self.scheduler_state = FSDPSchedulerState.PRE_FORWARD 

185 if self._is_root and not self.scheduler_ctx.lazy_init_done: 

186 self._init_params_fqn() 

187 self._lazy_init_all_states() 

188 self.scheduler_ctx.lazy_init_done = True 

189 if self.mp_policy.cast_forward_inputs and self.mp_policy.param_dtype: 

190 cast_fn = functools.partial(self.platform.cast_fp_tensor, self.mp_policy.param_dtype) 

191 args = self.platform.apply_to_tensors(cast_fn, args) 

192 kwargs = self.platform.apply_to_tensors(cast_fn, kwargs) 

193 with self.platform.profiler_record(f"pre_forward unshard:{self.hsdp_state.module_name}"): 

194 logger.debug("hook=forward_pre action=unshard module=%s", self.hsdp_state) 

195 self.hsdp_state.unshard() 

196 for prefetch_cell in self.forward_prefetch_cells: 

197 prefetch_state = prefetch_cell.hsdp_scheduler.hsdp_state 

198 with self.platform.profiler_record(f"pre_forward prefetch:" 

199 f"{prefetch_state.module_name}"): 

200 logger.debug( 

201 "hook=forward_pre action=prefetch module=%s target=%s", 

202 self.hsdp_state, 

203 prefetch_state, 

204 ) 

205 prefetch_state.prefetch() 

206 return args, kwargs 

207 

208 def _lazy_init_all_states(self): 

209 if self._is_root and self.scheduler_ctx.root_module is not None: 

210 for _, module in platform.get_cells_and_names(self.scheduler_ctx.root_module): 

211 hsdp_state = get_hsdp_state(module) 

212 if hsdp_state: 

213 hsdp_state.lazy_init() 

214 

215 def _init_params_fqn(self): # pylint: disable=W0212 

216 if not self._is_root or self.scheduler_ctx.root_module is None: 

217 return 

218 # Build a map from original (sharded) parameter tensor → hsdp_param wrapper, 

219 # covering both sharded hsdp_params and replicate_params. 

220 param_to_hsdp_param = {} 

221 for _, module in platform.get_cells_and_names(self.scheduler_ctx.root_module): 

222 hsdp_state = get_hsdp_state(module) 

223 if hsdp_state is None: 

224 continue 

225 for hsdp_param in hsdp_state._iter_managed_params(): # pylint: disable=W0212 

226 orig_param = hsdp_param.sharded_param 

227 # Shared parameters: keep only the first mapping to preserve the 

228 # first-seen FQN (consistent with the deduplication in _init_hsdp_params). 

229 if orig_param not in param_to_hsdp_param: 

230 param_to_hsdp_param[orig_param] = hsdp_param 

231 

232 # Walk the full parameter tree and assign FQNs; skip params already seen 

233 # (shared-parameter deduplication: first name wins). 

234 visited_params = set() 

235 for param_name, parameter in platform.parameters_dict(self.scheduler_ctx.root_module): 

236 if parameter in visited_params: 

237 continue 

238 visited_params.add(parameter) 

239 hsdp_param = param_to_hsdp_param.get(parameter) 

240 if hsdp_param is not None: 

241 hsdp_param._param_fqn = param_name # pylint: disable=W0212 

242 

243 # pylint: disable=W0613, R1710 

244 def _hsdp_forward_hook(self, cell, inputs, outputs): 

245 """Forward hook to shard parameter for saving memory.""" 

246 logger.debug("hook=forward enter module=%s", self.hsdp_state) 

247 if self.scheduler_state == FSDPSchedulerState.PRE_BACKWARD: 

248 logger.debug("hook=forward skip module=%s reason=pre_backward", self.hsdp_state) 

249 return 

250 self.scheduler_state = FSDPSchedulerState.FORWARD 

251 if self.reshard_after_forward: 

252 with self.platform.profiler_record(f"forward reshard:{self.hsdp_state.module_name}"): 

253 logger.debug("hook=forward action=reshard module=%s", self.hsdp_state) 

254 self.hsdp_state.shard(shard_replicate=False) 

255 if self.mp_policy.output_dtype is not None: 

256 outputs = self.platform.apply_to_tensors( 

257 functools.partial(self.platform.cast_fp_tensor, self.mp_policy.output_dtype), 

258 outputs, 

259 ) 

260 return outputs 

261 

262 # pylint: disable=W0613 

263 def _hsdp_backward_pre_hook(self, cell, grad_outputs): 

264 """Backward pre hook to unsharded parameter for backward process.""" 

265 logger.debug("hook=backward_pre enter module=%s", self.hsdp_state) 

266 self.scheduler_state = FSDPSchedulerState.PRE_BACKWARD 

267 if self.reshard_after_forward: 

268 with self.platform.profiler_record(f"pre_backward unshard:{self.hsdp_state.module_name}"): 

269 logger.debug("hook=backward_pre action=unshard module=%s", self.hsdp_state) 

270 self.hsdp_state.unshard(unshard_replicate=False) 

271 for prefetch_cell in self.backward_prefetch_cells: 

272 prefetch_state = prefetch_cell.hsdp_scheduler.hsdp_state 

273 with self.platform.profiler_record(f"pre_backward prefetch:" 

274 f"{prefetch_state.module_name}"): 

275 logger.debug( 

276 "hook=backward_pre action=prefetch module=%s target=%s", 

277 self.hsdp_state, 

278 prefetch_state, 

279 ) 

280 prefetch_state.prefetch(unshard_replicate=False) 

281 

282 # pylint: disable=W0613 

283 def _hsdp_backward_hook(self, cell, grad_inputs, grad_outputs): 

284 """Backward hook to shard parameter for optimizer process or saving memory.""" 

285 logger.debug("hook=backward_hook enter module=%s", self.hsdp_state) 

286 self.scheduler_state = FSDPSchedulerState.BACKWARD 

287 with self.platform.profiler_record(f"post_backward:{self.hsdp_state.module_name}"): 

288 logger.debug("hook=backward_hook action=post_backward module=%s", self.hsdp_state) 

289 self.hsdp_state.post_backward() 

290 if self._fsdp_group_post_pending is not None: 

291 self._fsdp_group_post_pending.clear() 

292 

293 # pylint: disable=W0613 

294 @staticmethod 

295 def _grouped_forward_pre_hook_skip(cell, args, kwargs): 

296 """Return value when grouped pre-forward should not run (first module already did). 

297 

298 Default matches MindSpore Cell forward pre-hooks (explicit ``(args, kwargs)``). 

299 ``TorchHSDPSchedulerV2`` overrides this to return ``None`` (``nn.Module`` idiom). 

300 """ 

301 return args, kwargs 

302 

303 @staticmethod 

304 def _grouped_forward_post_hook_skip(outputs): 

305 """Return value when grouped post-forward is deferred to a later module in the group. 

306 

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

308 """ 

309 return outputs 

310 

311 def _grouped_forward_pre_hook(self, cell, args, kwargs): 

312 """Run FSDP pre-forward only for the first module in the group (PyTorch FSDP2-aligned).""" 

313 pending = self._fsdp_group_post_pending 

314 if pending is None: 

315 return self._forward_pre_hook(cell, args, kwargs) 

316 if len(pending) == 0: 

317 pending.update(self.modules) 

318 return self._forward_pre_hook(cell, args, kwargs) 

319 return self._grouped_forward_pre_hook_skip(cell, args, kwargs) 

320 

321 def _make_grouped_forward_post_hook(self, mod): 

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

323 

324 def grouped_post_hook(cell, inputs, outputs): 

325 pending = self._fsdp_group_post_pending 

326 if pending is None: 

327 return self._forward_hook(cell, inputs, outputs) 

328 pending.discard(mod) 

329 if len(pending) == 0: 

330 return self._forward_hook(cell, inputs, outputs) 

331 return self._grouped_forward_post_hook_skip(outputs) 

332 

333 return grouped_post_hook 

334 

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

336 """Set cells prefetched during forward. 

337 

338 Args: 

339 hsdp_cell_list: HSDP cells to prefetch ahead of forward. 

340 """ 

341 self.forward_prefetch_cells = hsdp_cell_list 

342 

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

344 """Set cells prefetched during backward. 

345 

346 Args: 

347 hsdp_cell_list: HSDP cells to prefetch ahead of backward. 

348 """ 

349 self.backward_prefetch_cells = hsdp_cell_list 

350 

351 def _disable_forward_prefetch_for_recompute(self) -> None: 

352 """Temporarily disable forward prefetch during activation recompute.""" 

353 self._backup_forward_fetch = self.forward_prefetch_cells 

354 self.forward_prefetch_cells = [] 

355 

356 def _restore_forward_prefetch_after_recompute(self) -> bool: 

357 """Restore forward prefetch list after a recompute forward hook finishes.""" 

358 if self._backup_forward_fetch is None: 

359 return False 

360 self.forward_prefetch_cells = self._backup_forward_fetch 

361 self._backup_forward_fetch = None 

362 return True