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
« 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
18from typing import Callable, TYPE_CHECKING
19from types import SimpleNamespace
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
26if TYPE_CHECKING:
27 from hyper_parallel.auto_parallel.sapp_nd.nd.common.cost_model_preprocess import CostModelConfig
30class _PPB:
31 """Pipeline balance payload builder."""
33 def __init__(self, eval_cfg: Config, inner_dyn_fun: Callable) -> None:
34 """Initialize _PPB with evaluation config and dynamic memory function.
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
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]
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
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)
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)
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)
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)
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)
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
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
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)