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

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

16 

17from collections import defaultdict 

18from typing import List 

19 

20import mindspore as ms 

21 

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 

38 

39logger = get_logger("FSDP") 

40 

41 

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 

47 

48 

49class MindSporeHSDPStateV2(HSDPState): 

50 """Own MindSpore fully_shard parameters and gradient communication state.""" 

51 

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

80 

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 ) 

95 

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) 

110 

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 ) 

122 

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) 

137 

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 ) 

152 

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 ) 

172 

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) 

177 

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 ) 

194 

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

204 

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 ) 

218 

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

223 

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 ) 

241 

242 def _resolve_reduce_op(self) -> str: 

243 """Return the active mint reduction operation.""" 

244 return self.reduce_op_type 

245 

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 

268 

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) 

273 

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 

290 

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) 

299 

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) 

308 

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) 

327 

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 

343 

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) 

355 

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

385 

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

396 

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

416 

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 

420 

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 

429 

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