Coverage for  / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_nd / memory_estimation / _ppb.py: 96%

174 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-08-21 04:29 +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"""PPB input module""" 

16from __future__ import annotations 

17 

18from typing import Callable, TYPE_CHECKING 

19from types import SimpleNamespace 

20 

21from hyper_parallel.auto_parallel.sapp_nd.nd.common.config import Config 

22from hyper_parallel.auto_parallel.sapp_nd.nd.common.layer_type import LayerType 

23from hyper_parallel.auto_parallel.sapp_nd.memory_estimation._context import Context 

24from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.utils import EvalUtils 

25 

26if TYPE_CHECKING: 

27 from hyper_parallel.auto_parallel.sapp_nd.nd.common.cost_model_preprocess import CostModelConfig 

28 

29 

30class _PPB: 

31 """Pipeline balance payload builder.""" 

32 

33 def __init__(self, eval_cfg: Config, inner_dyn_fun: Callable) -> None: 

34 """Initialize _PPB with evaluation config and dynamic memory function. 

35 

36 Args: 

37 eval_cfg: Evaluation configuration object. 

38 inner_dyn_fun: Function to compute inner dynamic memory. 

39 """ 

40 self.eval_cfg = eval_cfg 

41 self._inner_dynamic_mem = inner_dyn_fun 

42 self.mb = EvalUtils.mb 

43 

44 def add_to_ppb_list(self, ppb_lay_desc: list, desc: dict) -> None: 

45 """layer description list preparation""" 

46 if desc: 

47 already_comp = False 

48 body_idx = 0 

49 for d in ppb_lay_desc: 

50 if all(v == d[k] for k, v in desc.items()): 

51 # already exist desc 

52 d["nb_layer"] += 1 

53 already_comp = True 

54 if d["type"] == "BODY": 

55 body_idx += 1 

56 if desc and not already_comp: 

57 desc["nb_layer"] = 1 

58 if desc["type"] == "BODY": 

59 desc["name"] = f"BODY_{body_idx}" 

60 else: 

61 desc["name"] = desc["type"] 

62 ppb_lay_desc += [desc] 

63 

64 def lay_ppb(self, ccfg: CostModelConfig, ctx: Context, res_stat: float) -> dict: 

65 """layer description preparation""" 

66 original_enable_node_log = ctx.enable_node_log 

67 ctx.enable_node_log = False 

68 try: 

69 desc = {} 

70 desc["model_name"] = ccfg.model_name 

71 if ctx.current_node == ctx.head_node: 

72 d_emb = self.mb(sum(self._inner_dynamic_mem(ppb=True))) 

73 desc["type"] = "HEAD" 

74 desc["memory_parameter"] = self.mb(res_stat) + d_emb 

75 desc["time"] = 1 

76 elif ctx.current_node == ctx.tail_node: 

77 d_out = self.mb(sum(self._inner_dynamic_mem(ppb=True))) 

78 desc["type"] = "TAIL" 

79 desc["memory_parameter"] = self.mb(res_stat) + d_out 

80 desc["time"] = 1 

81 else: 

82 ctx.current_node = LayerType.NOT_REC_LAYER 

83 dyn_nrec = self._inner_dynamic_mem(ppb=True) 

84 ctx.current_node = LayerType.SEL_REC_LAYER 

85 dyn_srec = self._inner_dynamic_mem(ppb=True) 

86 ctx.current_node = LayerType.FULL_REC_LAYER 

87 dyn_frec = self._inner_dynamic_mem(ppb=True) 

88 c = max(dyn_nrec[1], dyn_srec[1], dyn_frec[1]) 

89 desc["type"] = "BODY" 

90 desc["memory_parameter"] = self.mb(res_stat) 

91 desc["memory_parameter"] += self.mb(c) 

92 desc["memory_activation"] = self.mb(dyn_nrec[0]) 

93 desc["memory_select_rec"] = self.mb(dyn_srec[0]) 

94 desc["memory_recompute"] = self.mb(dyn_frec[0]) 

95 desc["time"] = 1 

96 finally: 

97 ctx.enable_node_log = original_enable_node_log 

98 return desc 

99 

100 def ppb_combine_bodies(self, ppb_lay_desc: list) -> None: 

101 """combine descriptions into a new body""" 

102 if not self.eval_cfg.ppb_combined: 

103 return 

104 for new_body in self.eval_cfg.ppb_combined: 

105 desc = { 

106 "model_name": "combined", 

107 "type": "BODY", 

108 "memory_parameter": 0, 

109 "memory_activation": 0, 

110 "memory_select_rec": 0, 

111 "memory_recompute": 0, 

112 "time": 1, 

113 "nb_layer": 1, 

114 "name": "COMBINED", 

115 } 

116 idx = -1 

117 for mod, t in new_body: 

118 target = next( 

119 ( 

120 d 

121 for d in ppb_lay_desc 

122 if d["model_name"] == mod and d["type"] == t.upper() 

123 ), 

124 None, 

125 ) 

126 if target: 

127 desc["model_name"] += "_" + mod 

128 desc["name"] += "_" + target["name"] 

129 for m in desc: 

130 if m.startswith("memory") and m in target: 

131 desc[m] += target[m] 

132 target_idx = ppb_lay_desc.index(target) 

133 idx = target_idx if idx < 0 else min(idx, target_idx) 

134 del ppb_lay_desc[target_idx] 

135 idx = max(idx, 0) 

136 ppb_lay_desc.insert(idx, desc) 

137 

138 def lay_ppb_new(self, ccfg: CostModelConfig, ctx: Context, res_stat: float) -> dict: 

139 """layer description preparation""" 

140 original_enable_node_log = ctx.enable_node_log 

141 ctx.enable_node_log = False 

142 try: 

143 desc = {} 

144 desc["model_name"] = ccfg.model_name 

145 if ctx.current_node == ctx.head_node: 

146 desc["memory_activation"] = {"NONE": 0, "FULL": 0} 

147 d_emb = self.mb(sum(self._inner_dynamic_mem(ppb=True))) 

148 desc["memory_parameter"] = self.mb(res_stat) + d_emb 

149 desc["type"] = "HEAD" 

150 desc["options"] = ["NONE", "FULL"] 

151 desc["forward_time"] = {"NONE": 1, "FULL": 1} 

152 desc["backward_time"] = {"NONE": 1, "FULL": 1} 

153 elif ctx.current_node == ctx.tail_node: 

154 desc["memory_activation"] = {"NONE": 0, "FULL": 0} 

155 d_out = self.mb(sum(self._inner_dynamic_mem(ppb=True))) 

156 desc["memory_parameter"] = self.mb(res_stat) + d_out 

157 desc["type"] = "TAIL" 

158 desc["options"] = ["NONE", "FULL"] 

159 desc["forward_time"] = {"NONE": 1, "FULL": 1} 

160 desc["backward_time"] = {"NONE": 1, "FULL": 1} 

161 else: 

162 desc["memory_activation"] = {"NONE": 0, "COMM": 0, "SLCT": 0, "BOTH": 0, "FULL": 0} 

163 original_current_node = ctx.current_node 

164 synthetic_rec_op = False 

165 if not hasattr(ccfg, 'rec_op'): 

166 ccfg.rec_op = SimpleNamespace( 

167 attBMM=1, headCast=1, dropout=1, softmax=1, normOp=1, gather=1, ffAct=1 

168 ) 

169 synthetic_rec_op = True 

170 original_rec_op = {} 

171 rec_op_keys = ['attBMM', 'headCast', 'dropout', 'softmax', 'normOp', 'gather', 'ffAct'] 

172 for key in rec_op_keys: 

173 original_rec_op[key] = getattr(ccfg.rec_op, key, 1) 

174 try: 

175 # NOT_REC_LAYER: No recompute (save all activations) 

176 ctx.current_node = LayerType.NOT_REC_LAYER 

177 for key in rec_op_keys: 

178 setattr(ccfg.rec_op, key, 1) 

179 dyn_nrec = self._inner_dynamic_mem(ppb=True) 

180 

181 # SLCT recompute: Recompute operators only (saves ~4% memory) 

182 # rec_op=0 means recompute (saves memory), rec_op=1 means don't recompute (uses memory) 

183 ctx.current_node = LayerType.SEL_REC_LAYER 

184 for key in ['attBMM', 'headCast', 'dropout', 'softmax', 'normOp', 'ffAct']: 

185 setattr(ccfg.rec_op, key, 0) 

186 setattr(ccfg.rec_op, 'gather', 1) 

187 dyn_srec = self._inner_dynamic_mem(ppb=True) 

188 

189 # COMM recompute: Recompute communication only (saves ~12.5% memory) 

190 ctx.current_node = LayerType.SEL_REC_LAYER 

191 for key in ['attBMM', 'headCast', 'dropout', 'softmax', 'normOp', 'ffAct']: 

192 setattr(ccfg.rec_op, key, 1) 

193 setattr(ccfg.rec_op, 'gather', 0) 

194 dyn_comm = self._inner_dynamic_mem(ppb=True) 

195 

196 # BOTH recompute: Recompute both operators and communication 

197 ctx.current_node = LayerType.SEL_REC_LAYER 

198 for key in rec_op_keys: 

199 setattr(ccfg.rec_op, key, 0) 

200 dyn_both = self._inner_dynamic_mem(ppb=True) 

201 

202 # FULL_REC_LAYER: Full recompute 

203 ctx.current_node = LayerType.FULL_REC_LAYER 

204 dyn_frec = self._inner_dynamic_mem(ppb=True) 

205 finally: 

206 for key, val in original_rec_op.items(): 

207 setattr(ccfg.rec_op, key, val) 

208 if synthetic_rec_op: 

209 delattr(ccfg, 'rec_op') 

210 ctx.current_node = original_current_node 

211 

212 c = max(dyn_nrec[1], dyn_srec[1], dyn_comm[1], dyn_both[1], dyn_frec[1]) 

213 desc["memory_parameter"] = self.mb(res_stat) 

214 desc["memory_parameter"] += self.mb(c) 

215 desc["memory_activation"]["NONE"] = self.mb(dyn_nrec[0]) 

216 desc["memory_activation"]["COMM"] = self.mb(dyn_comm[0]) 

217 desc["memory_activation"]["SLCT"] = self.mb(dyn_srec[0]) 

218 desc["memory_activation"]["BOTH"] = self.mb(dyn_both[0]) 

219 desc["memory_activation"]["FULL"] = self.mb(dyn_frec[0]) 

220 desc["type"] = "BODY" 

221 desc["options"] = ["NONE", "COMM", "SLCT", "BOTH", "FULL"] 

222 desc["forward_time"] = {"NONE": 1, "COMM": 1, "SLCT": 1, "BOTH": 1, "FULL": 1} 

223 desc["backward_time"] = {"NONE": 1, "COMM": 1, "SLCT": 1, "BOTH": 1, "FULL": 1} 

224 desc["time"] = 1 

225 finally: 

226 ctx.enable_node_log = original_enable_node_log 

227 return desc 

228 

229 def ppb_combine_bodies_new(self, ppb_lay_desc: list) -> None: 

230 """combine descriptions into a new body""" 

231 if not self.eval_cfg.ppb_combined: 

232 return 

233 for new_body in self.eval_cfg.ppb_combined: 

234 desc = { 

235 "model_name": "combined", 

236 "type": "BODY", 

237 "memory_parameter": 0, 

238 "memory_activation": {"NONE": 0, "COMM": 0, "SLCT": 0, "BOTH": 0, "FULL": 0}, 

239 "options": ["NONE", "COMM", "SLCT", "BOTH", "FULL"], 

240 "forward_time": {"NONE": 1, "COMM": 1, "SLCT": 1, "BOTH": 1, "FULL": 1}, 

241 "backward_time": {"NONE": 1, "COMM": 1, "SLCT": 1, "BOTH": 1, "FULL": 1}, 

242 "time": 1, 

243 "nb_layer": 1, 

244 "name": "COMBINED", 

245 } 

246 idx = -1 

247 for mod, t in new_body: 

248 target = next( 

249 ( 

250 d 

251 for d in ppb_lay_desc 

252 if d["model_name"] == mod and d["type"] == t.upper() 

253 ), 

254 None, 

255 ) 

256 if target: 

257 desc["model_name"] += "_" + mod 

258 desc["name"] += "_" + target["name"] 

259 desc["memory_parameter"] += target["memory_parameter"] 

260 desc["memory_activation"]["NONE"] += target["memory_activation"]["NONE"] 

261 desc["memory_activation"]["COMM"] += target["memory_activation"].get("COMM", 0) 

262 desc["memory_activation"]["SLCT"] += target["memory_activation"].get("SLCT", 0) 

263 desc["memory_activation"]["BOTH"] += target["memory_activation"].get("BOTH", 0) 

264 desc["memory_activation"]["FULL"] += target["memory_activation"]["FULL"] 

265 target_idx = ppb_lay_desc.index(target) 

266 idx = target_idx if idx < 0 else min(idx, target_idx) 

267 del ppb_lay_desc[target_idx] 

268 idx = max(idx, 0) 

269 ppb_lay_desc.insert(idx, desc)