Coverage for  / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_nd / nd / global_config.py: 91%

153 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-08-04 05:18 +0800

1# Copyright 2024 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"""One configuration interface for parallelization""" 

16 

17import copy 

18from math import gcd 

19 

20from hyper_parallel.auto_parallel.sapp_nd.nd.common.arch_hooks import CWrap, check_and_apply_custom_hook 

21from hyper_parallel.auto_parallel.sapp_nd.nd.logger import logger 

22import hyper_parallel.auto_parallel.sapp_nd.nd.dimensions as Dim 

23import hyper_parallel.auto_parallel.sapp_nd.nd.common.hardware as Hard 

24import hyper_parallel.auto_parallel.sapp_nd.nd.balancing_adapter as BA 

25 

26 

27class GlobalConfig: 

28 """Union of cost model & parallel config""" 

29 

30 def __init__(self, config, dimensions=None, mppb=False): 

31 

32 self.wrap = CWrap(config) 

33 self.ccfg = self.wrap.ccfg 

34 

35 if dimensions is not None: 

36 logger.debug("dimensions = %s", str(dimensions)) 

37 self.dimensions = dimensions 

38 else: 

39 logger.debug("dimensions = %s", str(Dim.ALL_DIMS)) 

40 self.dimensions = Dim.ALL_DIMS.copy() 

41 logger.debug("self.dimensions = %s", str(self.dimensions)) 

42 logger.debug("layer_num_for_offset = %d", self.layer_num_for_offset()) 

43 logger.debug("total layer num = %d", self.total_layer_num()) 

44 self.balancing = BA.BalancingAdapter( 

45 self.layer_num_for_offset(), 

46 copy.deepcopy(self.ccfg.offset), 

47 copy.deepcopy(self.ccfg.full_rec), 

48 mppb, 

49 ) 

50 

51 def dim_val(self, dim, parallel_config): 

52 """Get the value of a parallel dimension""" 

53 if parallel_config.has_dim(dim): 

54 return parallel_config.val(dim) 

55 return dim.from_config(self.ccfg) 

56 

57 def global_batch_size(self, parallel_config): 

58 """Compute global batch size from hyperparameters""" 

59 dp = self.dim_val(Dim.DP, parallel_config) 

60 pp = self.dim_val(Dim.PP, parallel_config) 

61 mb = self.dim_val(Dim.MBN, parallel_config) 

62 bs = self.dim_val(Dim.MBS, parallel_config) 

63 if pp > 1: 

64 logger.info("GBS = %dDP * %dMB * %dBS", dp, mb, bs) 

65 return dp * mb * bs 

66 logger.info("GBS = %dDP * %dBS", dp, bs) 

67 return dp * bs 

68 

69 def layer_num_for_offset(self): 

70 """Compute layer number including MTP when necessary for offset""" 

71 layer_num = self.ccfg.n_lay 

72 if self.ccfg.emb_out_in_offset: 

73 layer_num += 2 

74 if self.ccfg.is_mtp_in_offset: 

75 layer_num += self.ccfg.n_mtp 

76 return layer_num 

77 

78 def total_layer_num(self): 

79 """Compute total layer number, always including MTP""" 

80 layer_num = self.ccfg.n_lay + self.ccfg.n_mtp 

81 return layer_num 

82 

83 def adapt_config_balancing(self, new_pp, new_vpp): 

84 """Adapt the layer-to-stage assignment to different PP""" 

85 logger.debug("new_pp=%d, new_vpp=%d", new_pp, new_vpp) 

86 

87 new_recompute_config = self.balancing.treat_recompute(new_pp, new_vpp) 

88 logger.debug("adapted recompute config: %s", str(new_recompute_config)) 

89 new_offset = self.balancing.treat_offset(new_pp, new_vpp) 

90 logger.debug("adapted offset: %s", str(new_offset)) 

91 ok = self.balancing.offset_checker(new_pp, new_vpp, new_offset) 

92 if not ok: 

93 logger.error("Offset {%s} NOT VALID", str(new_offset)) 

94 return new_offset, new_recompute_config 

95 

96 def adapt_config(self, pp, vpp): 

97 """Adapt configuration to different parallel config""" 

98 return self.adapt_config_balancing(pp, vpp) 

99 

100 def write(self, folder, parallel_config): 

101 """Dump config into a yaml file""" 

102 if folder: 

103 file_name = parallel_config.unique_name() 

104 self.ccfg.config.dump(file_name, folder) 

105 

106 def moe_valid(self, parallel_config): 

107 """Check whether the model is MoE""" 

108 expert_num = self.ccfg.n_exp 

109 if expert_num > 1: 

110 ep = self.dim_val(Dim.EP, parallel_config) 

111 dp = self.dim_val(Dim.DP, parallel_config) 

112 mp = self.dim_val(Dim.TP, parallel_config) 

113 logger.debug( 

114 "moe valid ? EP %d <= E %d & EP %d <= DP %d * MP %d", 

115 ep, 

116 expert_num, 

117 ep, 

118 dp, 

119 mp, 

120 ) 

121 return ep <= min(expert_num, dp * mp) 

122 return True 

123 

124 def ep_constraints_valid(self, parallel_config): 

125 """Check EP-specific divisibility constraints (C1, C2). 

126 

127 Runs only for MoE models (n_exp > 1). C1 ensures experts can be 

128 evenly partitioned across EP ranks; C2 ensures the expert FFN hidden 

129 dim can be evenly sharded by the expert TP degree. Both checks use 

130 architecture constants from ``self.ccfg`` and the candidate values 

131 from ``parallel_config``. 

132 

133 C3 (device count) is intentionally skipped here because the search 

134 loop borrows EP from the dp*tp budget, so dp*tp*pp*cp already 

135 equals total_devices by construction. 

136 

137 Args: 

138 parallel_config: candidate ``Dim.Dimensions``. 

139 

140 Returns: 

141 bool: True if all applicable EP constraints pass (or the model 

142 is dense), False otherwise. 

143 """ 

144 if self.ccfg.n_exp <= 1: 

145 return True 

146 # pylint: disable=C0415 

147 from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.validators.ep_constraints import EpConstraints 

148 ep = self.dim_val(Dim.EP, parallel_config) 

149 r1 = EpConstraints.check_ep_divisibility(self.ccfg.n_exp, ep) 

150 if not r1: 

151 logger.warning("EP constraint C1 failed: %s", r1.message) 

152 return False 

153 tp = self.dim_val(Dim.TP, parallel_config) 

154 etp = max(getattr(self.ccfg, "etp", 0), 0) 

155 t_exp = max(etp, 1) if etp > 1 else max(tp, 1) 

156 hff_exp = max(getattr(self.ccfg, "hff_exp", 0), 0) 

157 r2 = EpConstraints.check_expert_hidden_divisibility(hff_exp, t_exp) 

158 if not r2: 

159 logger.warning("EP constraint C2 failed: %s", r2.message) 

160 return False 

161 return True 

162 

163 def make_parallel_config_args(self, **kwargs): 

164 """Create a parallel config from parallel values""" 

165 logger.debug("dimensions considered: %s", str(self.dimensions)) 

166 

167 dims = [] 

168 # dims.append((Dim.DP, dp)) 

169 for dim in self.dimensions: 

170 dims.append((dim, kwargs.get(dim.lname()))) 

171 

172 has_mbn_not_in = Dim.MBN not in self.dimensions 

173 has_pp_in = Dim.PP in self.dimensions 

174 has_dp_or_mbs_in = Dim.DP in self.dimensions or Dim.MBS in self.dimensions 

175 if has_mbn_not_in and has_pp_in and has_dp_or_mbs_in: 

176 dims.append((Dim.MBN, kwargs.get(Dim.MBN.lname()))) 

177 self.dimensions.append(Dim.MBN) 

178 return Dim.Dimensions(dims, all_dims=self.dimensions) 

179 

180 def make_parallel_config(self, dtpc_p, mbsn, evos_p): 

181 """Create a parallel config from parallel values""" 

182 logger.debug("dimensions considered: %s", str(self.dimensions)) 

183 (dp, mp, pp, cp) = dtpc_p 

184 (mbs, mbn) = mbsn 

185 (ep, vpp, op, sp) = evos_p 

186 return self.make_parallel_config_args( 

187 dp=dp, 

188 mp=mp, 

189 pp=pp, 

190 cp=cp, 

191 mbs=mbs, 

192 mb=mbn, 

193 ep=ep, 

194 vpp=vpp, 

195 op=op, 

196 sp=sp, 

197 ) 

198 

199 def set_parallel_config(self, parallel_config): 

200 """Set a given parallel configuration in the config""" 

201 kwargs = {} 

202 ok = True 

203 new_pp = self.dim_val(Dim.PP, parallel_config) 

204 new_vp = self.dim_val(Dim.VPP, parallel_config) 

205 new_offset, new_recompute = self.adapt_config(new_pp, new_vp) 

206 kwargs["offset"] = new_offset 

207 kwargs["full_rec"] = new_recompute 

208 # kwargs["sel_rec"] = sel_rec 

209 for dim, value in parallel_config.dims_val.items(): 

210 kwargs[dim.name.lower()] = value 

211 

212 self.ccfg.set_strategy(**kwargs) 

213 if not self.ccfg.multimodal: 

214 if not self.ccfg.hooks_dict: 

215 logger.info( 

216 "'hook_cls' not specified," 

217 "search in predefined arch_hooks" 

218 ) 

219 check_and_apply_custom_hook(self.ccfg) 

220 else: 

221 logger.info("Apply hooks") 

222 hook = list(self.ccfg.hooks_dict.values())[0] 

223 hook(self.wrap) 

224 

225 return ok 

226 

227 def space(self, dim, divide, reverse=False): 

228 """Generate the space for a given dimension""" 

229 if dim in self.dimensions: 

230 if dim.get_bound() is not None: 

231 logger.debug( 

232 "Space of bounded dim %s is %s", 

233 str(dim), 

234 str( 

235 Hard.all_divisors( 

236 divide, reverse=reverse, max_bound=dim.get_bound() 

237 ) 

238 ), 

239 ) 

240 return Hard.all_divisors( 

241 divide, reverse=reverse, max_bound=dim.get_bound() 

242 ) 

243 logger.debug( 

244 "Space of dim %s is %s", 

245 str(dim), 

246 str(Hard.all_divisors(divide, reverse=reverse)), 

247 ) 

248 return Hard.all_divisors(divide, reverse=reverse) 

249 logger.debug( 

250 "Space of original dim %s is [%s]", 

251 str(dim), 

252 str(dim.from_config(self.ccfg)), 

253 ) 

254 return [dim.from_config(self.ccfg)] 

255 

256 def range_space(self, dim, bound): 

257 """Generate the space for a given dimension""" 

258 if dim in self.dimensions: 

259 return range(1, bound + 1) 

260 return [dim.from_config(self.ccfg)] 

261 

262 def bool_space(self, dim): 

263 """Generate the space for a given boolean dimension""" 

264 if dim in self.dimensions: 

265 return [False, True] 

266 return [dim.from_config(self.ccfg)] 

267 

268 def max_op(self, dp, tp, ep): 

269 """Compute bound for dimension OP""" 

270 if ( 

271 isinstance(self.ccfg.optimizer, str) 

272 and "muon" not in self.ccfg.optimizer.lower() 

273 ): 

274 return dp 

275 if self.ccfg.n_exp and self.ccfg.n_exp > 1: 

276 exp_gcd = gcd(dp * tp // max(tp, ep), self.ccfg.n_exp) 

277 else: 

278 exp_gcd = dp 

279 

280 dc_kv_valid = self.ccfg.dc_kv and self.ccfg.dc_kv > 1 

281 dhr_valid = self.ccfg.dhr and self.ccfg.dhr > 1 

282 if dc_kv_valid and dhr_valid: 

283 att_gcd = gcd(self.ccfg.h, self.ccfg.dc_kv + self.ccfg.dhr) 

284 else: 

285 att_gcd = self.ccfg.h 

286 return gcd(exp_gcd, att_gcd)