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

260 statements  

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

1# Copyright 2025 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"""hook manager module""" 

16from __future__ import annotations 

17from typing import TYPE_CHECKING 

18import ast 

19import textwrap 

20import inspect 

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.nd.common.cost_model_preprocess import CostModelConfig 

24from hyper_parallel.auto_parallel.sapp_nd.nd.common.arch_hooks import check_and_apply_custom_hook 

25from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.logger import logger 

26from hyper_parallel.auto_parallel.sapp_nd.memory_estimation._context import ( 

27 NodeEval, 

28 NodeStatEval, 

29 NodeDynEval, 

30 NodeCommEval, 

31 NodeComputeEval, 

32) 

33from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.compute import EvalExpertCompute 

34from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.head import EvalHead 

35from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.tail import EvalTail 

36from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.body import EvalBody 

37from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.layer_block import EvalAttn, EvalFFn 

38from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.layer_block import EvalNorm 

39from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.comm import EvalLayerComm 

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

41from hyper_parallel.auto_parallel.sapp_nd.memory_estimation._backbone import _Backbone 

42from hyper_parallel.auto_parallel.sapp_nd.memory_estimation._func_tracer import _FuncTracer 

43from hyper_parallel.auto_parallel.sapp_nd.memory_estimation._context import MemType 

44 

45if TYPE_CHECKING: 

46 from typing import Callable, Any 

47 

48 

49class _HookManager(_Backbone): 

50 """Hook manager class.""" 

51 

52 def __init__(self, *args, **kwargs): 

53 super().__init__(*args, **kwargs) 

54 self.func_tracer = _FuncTracer() 

55 self.toggle_func_trace = kwargs.get("trace_fun", False) 

56 self.import_eval_yaml() 

57 self.fetch_hook_if_unimodal() 

58 

59 def fetch_hook_if_unimodal(self): 

60 """Process custom model config""" 

61 if not self._ccfg.multimodal: 

62 if not self._ccfg.hooks_dict: 

63 logger.info( 

64 "'hook_cls' not specified," 

65 "search in predefined arch_hooks" 

66 ) 

67 check_and_apply_custom_hook(self) 

68 else: 

69 hook = list(self._ccfg.hooks_dict.values())[0] 

70 hook(self) 

71 

72 def __is_valid_eval_func(self, fun: Any) -> bool: 

73 """check hook definition, return type""" 

74 if fun is None: 

75 return False 

76 if isinstance(fun, (str, int, float)): 

77 return True 

78 

79 source = inspect.getsource(fun) 

80 tree = ast.parse(textwrap.dedent(source)) 

81 for instruction in ast.walk(tree): 

82 if ( 

83 isinstance(instruction, ast.Return) 

84 and instruction.value is None 

85 ): 

86 return False 

87 return True 

88 

89 # Evaluation context setters 

90 

91 def set_passes( 

92 self, 

93 vpp_less_mem: bool = None, 

94 swap_os: bool = None, 

95 dropless_tok_factor: float = None, 

96 ) -> None: 

97 """toggle features""" 

98 if isinstance(vpp_less_mem, bool): 

99 self._ctx.vpp_less_mem = vpp_less_mem 

100 if isinstance(swap_os, bool): 

101 self._ctx.swap_os = swap_os 

102 if isinstance(dropless_tok_factor, (int, float)): 

103 self._ctx.dropless_tok_factor = dropless_tok_factor 

104 

105 def __set_node_eval_fun(self, cls_obj, target_node, *args, **kwargs): 

106 """overwrite given node's formulas""" 

107 c_stat, c_dyn = None, None 

108 if (args and args[-1] == 0) or kwargs.get("stat", None) == 0: 

109 c_stat = 0 

110 if (args and args[-1] == 0) or kwargs.get("dyn", None) == 0: 

111 c_dyn = 0 

112 num_p = kwargs.get("num_p", c_stat) 

113 stat_p = kwargs.get("stat_p", c_stat) 

114 stat_os = kwargs.get("stat_os", c_stat) 

115 stat_grad = kwargs.get("stat_grad", c_dyn) 

116 dyn_activ = kwargs.get("dyn_activ", c_dyn) 

117 if not self.__is_valid_eval_func(num_p): 

118 num_p = self._ctx.node_eval[target_node].num_p 

119 if not self.__is_valid_eval_func(stat_p): 

120 stat_p = self._ctx.node_eval[target_node].stat.p 

121 if not self.__is_valid_eval_func(stat_os): 

122 stat_os = self._ctx.node_eval[target_node].stat.os 

123 if not self.__is_valid_eval_func(stat_grad): 

124 stat_grad = self._ctx.node_eval[target_node].stat.grad 

125 if not self.__is_valid_eval_func(dyn_activ): 

126 dyn_activ = self._ctx.node_eval[target_node].dyn.activation 

127 # Compute formulas (FLOPs, not memory — bypass mem_counter) 

128 compute_eval = self.__set_node_eval_compute_fun(target_node, **kwargs) 

129 self._ctx.node_eval[target_node] = NodeEval( 

130 self.__custom_getattr(cls_obj, num_p), 

131 NodeStatEval( 

132 self.__custom_getattr(cls_obj, stat_p, MemType.MODEL_PARAM), 

133 self.__custom_getattr(cls_obj, stat_os, MemType.OPTIM_STATE), 

134 self.__custom_getattr(cls_obj, stat_grad, MemType.ACCU_GRAD), 

135 ), 

136 NodeDynEval( 

137 self.__custom_getattr(cls_obj, dyn_activ), 

138 self.__set_node_eval_comm_fun( 

139 cls_obj, target_node, *args, **kwargs 

140 ), 

141 compute=compute_eval, 

142 ), 

143 ) 

144 

145 def __set_node_eval_comm_fun(self, cls_obj, target_node, *args, **kwargs): 

146 """overwrite given node's comm formulas""" 

147 c_comm = None 

148 last_arg_is_zero = args and args[-1] == 0 

149 dyn_comm_is_zero = kwargs.get("dyn_comm") == 0 

150 dyn_is_zero = kwargs.get("dyn") == 0 

151 

152 if last_arg_is_zero or dyn_comm_is_zero or dyn_is_zero: 

153 c_comm = 0 

154 dyn_dp_comm = kwargs.get("dyn_dp_comm", c_comm) 

155 dyn_tp_comm = kwargs.get("dyn_tp_comm", c_comm) 

156 dyn_cp_comm = kwargs.get("dyn_cp_comm", c_comm) 

157 dyn_ep_comm = kwargs.get("dyn_ep_comm", c_comm) 

158 dyn_ep_comm_balanced = kwargs.get("dyn_ep_comm_balanced", None) 

159 dyn_ep_comm_imbalanced = kwargs.get("dyn_ep_comm_imbalanced", None) 

160 if not self.__is_valid_eval_func(dyn_dp_comm): 

161 dyn_dp_comm = self._ctx.node_eval[target_node].dyn.comm.dp 

162 if not self.__is_valid_eval_func(dyn_tp_comm): 

163 dyn_tp_comm = self._ctx.node_eval[target_node].dyn.comm.tp 

164 if not self.__is_valid_eval_func(dyn_cp_comm): 

165 dyn_cp_comm = self._ctx.node_eval[target_node].dyn.comm.cp 

166 if not self.__is_valid_eval_func(dyn_ep_comm): 

167 dyn_ep_comm = self._ctx.node_eval[target_node].dyn.comm.ep 

168 comm_cls_obj = cls_obj 

169 if self.is_regular_layer(target_node): 

170 comm_cls_obj = EvalLayerComm 

171 ep_balanced = ( 

172 self.__custom_getattr(comm_cls_obj, dyn_ep_comm_balanced) 

173 if self.__is_valid_eval_func(dyn_ep_comm_balanced) 

174 else None 

175 ) 

176 ep_imbalanced = ( 

177 self.__custom_getattr(comm_cls_obj, dyn_ep_comm_imbalanced) 

178 if self.__is_valid_eval_func(dyn_ep_comm_imbalanced) 

179 else None 

180 ) 

181 return NodeCommEval( 

182 self.__custom_getattr(comm_cls_obj, dyn_dp_comm), 

183 self.__custom_getattr(comm_cls_obj, dyn_tp_comm), 

184 self.__custom_getattr(comm_cls_obj, dyn_cp_comm), 

185 self.__custom_getattr(comm_cls_obj, dyn_ep_comm), 

186 ep_balanced=ep_balanced, 

187 ep_imbalanced=ep_imbalanced, 

188 ) 

189 

190 def __resolve_compute_fun(self, name): 

191 """Resolve compute formula name — bypasses __wrap_mem_counter. 

192 

193 Compute FLOPs are not memory bytes and must not be accumulated 

194 into accu_mem_type. This method directly retrieves the function 

195 from EvalExpertCompute without the mem_counter wrapper. 

196 Falls back to a zero-returning function if the name is not found, 

197 so YAML typos produce zero FLOPs instead of crashing. 

198 """ 

199 fun = getattr(EvalExpertCompute, name, None) 

200 if fun is not None: 

201 return fun 

202 logger.warning("Compute formula '%s' not found in EvalExpertCompute, " 

203 "falling back to zero", name) 

204 

205 def zero(*_): 

206 return 0 

207 

208 zero.__qualname__ = f"0({name})" 

209 return zero 

210 

211 def __set_node_eval_compute_fun(self, target_node, **kwargs): 

212 """Build NodeComputeEval from kwargs, falling back to existing values.""" 

213 compute_kwargs = kwargs.get("compute", None) 

214 if compute_kwargs is None: 

215 # No compute config provided — preserve existing or return None 

216 existing = self._ctx.node_eval.get(target_node) 

217 if existing and existing.dyn.compute is not None: 

218 return existing.dyn.compute 

219 return None 

220 if not isinstance(compute_kwargs, dict): 

221 return None 

222 router_name = compute_kwargs.get("router", None) 

223 expert_balanced_name = compute_kwargs.get("expert_balanced", None) 

224 expert_imbalanced_name = compute_kwargs.get("expert_imbalanced", None) 

225 shared_expert_name = compute_kwargs.get("shared_expert", None) 

226 # Resolve from EvalExpertCompute (bypass mem_counter) 

227 router = self.__resolve_compute_fun(router_name) if router_name else None 

228 expert_balanced = ( 

229 self.__resolve_compute_fun(expert_balanced_name) 

230 if expert_balanced_name else None 

231 ) 

232 expert_imbalanced = ( 

233 self.__resolve_compute_fun(expert_imbalanced_name) 

234 if expert_imbalanced_name else None 

235 ) 

236 shared_expert = ( 

237 self.__resolve_compute_fun(shared_expert_name) 

238 if shared_expert_name else None 

239 ) 

240 if not any([router, expert_balanced, expert_imbalanced, shared_expert]): 

241 return None 

242 return NodeComputeEval( 

243 router=router, 

244 expert_balanced=expert_balanced, 

245 expert_imbalanced=expert_imbalanced, 

246 shared_expert=shared_expert, 

247 ) 

248 

249 def set_head_eval_fun(self, *arg, **kwarg): 

250 """overwrite head formulas""" 

251 self.__set_node_eval_fun(EvalHead, self._ctx.head_node, *arg, **kwarg) 

252 

253 def set_tail_eval_fun(self, *arg, **kwarg): 

254 """overwrite tail formulas""" 

255 self.__set_node_eval_fun(EvalTail, self._ctx.tail_node, *arg, **kwarg) 

256 

257 def set_body_eval_fun(self, *args, **kwargs): 

258 """overwrite body formulas""" 

259 lay_type = kwargs.get("lay_type", args[0] if args else None) 

260 if not lay_type: 

261 lt = [ 

262 b_obj 

263 for b_obj in list(LayerType) 

264 if self.is_regular_layer(b_obj) 

265 ] 

266 else: 

267 if not isinstance(lay_type, LayerType): 

268 b_obj = self.__custom_getattr(LayerType, lay_type) 

269 lt = [b_obj] 

270 else: 

271 lt = [lay_type] 

272 for b_obj in lt: 

273 self.__set_node_eval_fun(EvalBody, b_obj, *args, **kwargs) 

274 

275 def set_attn_eval_fun( 

276 self, 

277 num_p: Any = None, 

278 qkv: Any = None, 

279 score: Any = None, 

280 proj: Any = None, 

281 ) -> None: 

282 """overwrite attention formulas""" 

283 if self.__is_valid_eval_func(num_p): 

284 self._ctx.attn_num_p = self.__custom_getattr(EvalAttn, num_p) 

285 if self.__is_valid_eval_func(qkv): 

286 self._ctx.attn_qkv_activ = self.__custom_getattr( 

287 EvalAttn, qkv, MemType.ATTN_ACTIV 

288 ) 

289 if self.__is_valid_eval_func(score): 

290 self._ctx.attn_score_activ = self.__custom_getattr( 

291 EvalAttn, score, MemType.ATTN_ACTIV 

292 ) 

293 if self.__is_valid_eval_func(proj): 

294 self._ctx.attn_proj_activ = self.__custom_getattr( 

295 EvalAttn, proj, MemType.ATTN_ACTIV 

296 ) 

297 

298 def set_ffn_eval_fun(self, num_p: Any = None, activation=None, moe_activ=None): 

299 """overwrite feedforward formulas""" 

300 if self.__is_valid_eval_func(num_p): 

301 self._ctx.ffn_num_p = self.__custom_getattr(EvalFFn, num_p) 

302 if self.__is_valid_eval_func(activation): 

303 self._ctx.ffn_activ = self.__custom_getattr( 

304 EvalFFn, activation, MemType.FFN_ACTIV 

305 ) 

306 if self.__is_valid_eval_func(moe_activ): 

307 self._ctx.ffn_moe_activ = self.__custom_getattr( 

308 EvalFFn, moe_activ, MemType.FFN_ACTIV 

309 ) 

310 

311 def set_expert_param_eval_fun( 

312 self, routed_num_p: Any = None, shared_num_p: Any = None 

313 ) -> None: 

314 """overwrite expert param count formulas for routed/shared breakdown""" 

315 if self.__is_valid_eval_func(routed_num_p): 

316 self._ctx.ffn_routed_num_p = self.__custom_getattr( 

317 EvalFFn, routed_num_p 

318 ) 

319 if self.__is_valid_eval_func(shared_num_p): 

320 self._ctx.ffn_shared_num_p = self.__custom_getattr( 

321 EvalFFn, shared_num_p 

322 ) 

323 

324 def set_norm_eval_fun(self, num_p: Any = None, activation=None): 

325 """overwrite norm formulas""" 

326 if self.__is_valid_eval_func(num_p): 

327 self._ctx.norm_num_p = self.__custom_getattr(EvalNorm, num_p) 

328 if self.__is_valid_eval_func(activation): 

329 self._ctx.norm_activ = self.__custom_getattr( 

330 EvalNorm, activation, MemType.NORM_ACTIV 

331 ) 

332 

333 def set_pp_micro_factor_eval_fun(self, sched_name, fun): 

334 """overwrite PP microfactor formulas""" 

335 if sched_name and self.__is_valid_eval_func(fun): 

336 self._ctx.pp_micro_eval[sched_name] = self.__custom_getattr( 

337 EvalUtils, fun 

338 ) 

339 

340 # Cost Model Config setter 

341 

342 def set_strategy(self, **kwargs): 

343 """overwrite parallelism""" 

344 self._ccfg.set_strategy(**kwargs) 

345 self.fetch_hook_if_unimodal() 

346 

347 def set_ccfg(self, hook): 

348 """overwrite cost model variable (except strategy)""" 

349 if hook and callable(hook): 

350 

351 def custom_setter(self, name, value): 

352 strat_vars = ["d", "t", "ep", "p", "vp", "cp", "os_max_shard"] 

353 if name in strat_vars: 

354 raise AttributeError( 

355 f"Cannot directly modify {name}, use set_strategy()" 

356 ) 

357 self.__dict__[name] = value 

358 

359 CostModelConfig.__setattr__ = custom_setter 

360 hook(self._ccfg) 

361 CostModelConfig.__setattr__ = object.__setattr__ 

362 

363 def __wrap_mem_counter(self, mem_type: MemType, fun: Callable) -> None: 

364 """Wrap formula calls to accumulate memory by type.""" 

365 if mem_type and not hasattr(fun, "wrapped_with_counter"): 

366 

367 def wrap(*args, **kwargs): 

368 res = fun(*args, **kwargs) 

369 self._ctx.accu_mem_type[mem_type] += res 

370 self._ctx.save2log(mem_type, res) 

371 return res 

372 

373 wrap.__qualname__ = fun.__qualname__ 

374 wrap.wrapped_with_counter = True 

375 return wrap 

376 return fun 

377 

378 def __custom_getattr( 

379 self, eval_class: Any, field: Any, mem_type: MemType = None 

380 ) -> Callable: 

381 """formula retrieve/wrap""" 

382 # Definition priority order : 

383 # 1. Callable (user defined in code) 

384 # OR Numeric value 

385 # 2. Overriding list from cost_model_preprocess (ccfg attribute) 

386 # 3. Function name (config_eval yaml) 

387 res = None 

388 if callable(field): 

389 res = field 

390 if self.toggle_func_trace == field.__name__: 

391 res = self.func_tracer.wrap(field) 

392 if isinstance(field, (int, float)): 

393 

394 def constant(*_): 

395 return field 

396 

397 constant.__qualname__ = str(field) 

398 res = constant 

399 if field in self._ccfg.overwrite_eval_functions: 

400 res = self._ccfg.overwrite_eval_functions[field] 

401 if isinstance(field, str): 

402 

403 def zero(*_): 

404 return 0 

405 

406 zero.__qualname__ = "0" 

407 res = getattr(eval_class, field, zero) 

408 if self.toggle_func_trace == field: 

409 res = self.func_tracer.wrap(getattr(eval_class, field)) 

410 res = self.__wrap_mem_counter(mem_type, res) 

411 if not res: 

412 raise TypeError(f"In eval config yaml, non valid field: {field}") 

413 return res 

414 

415 # Import Eval Config 

416 

417 def import_eval_yaml(self) -> None: 

418 """import evaluator config file, init ctx (inner call only)""" 

419 if not self.toggle_func_trace: 

420 self.toggle_func_trace = self.eval_cfg.trace_fun 

421 # head 

422 h_obj = getattr(LayerType, self.eval_cfg.nodes_mem_comp.head.name) 

423 self._ctx.head_node = h_obj 

424 self.set_head_eval_fun( 

425 num_p=self.eval_cfg.nodes_mem_comp.head.num_param_fun, 

426 stat_p=self.eval_cfg.nodes_mem_comp.head.stat_fun.p, 

427 stat_os=self.eval_cfg.nodes_mem_comp.head.stat_fun.os, 

428 stat_grad=self.eval_cfg.nodes_mem_comp.head.stat_fun.grad, 

429 dyn_activ=self.eval_cfg.nodes_mem_comp.head.dyn_fun.activation, 

430 dyn_dp_comm=self.eval_cfg.nodes_mem_comp.head.dyn_fun.comm.dp, 

431 dyn_tp_comm=self.eval_cfg.nodes_mem_comp.head.dyn_fun.comm.tp, 

432 dyn_cp_comm=self.eval_cfg.nodes_mem_comp.head.dyn_fun.comm.cp, 

433 dyn_ep_comm=self.eval_cfg.nodes_mem_comp.head.dyn_fun.comm.ep, 

434 ) 

435 

436 # tail 

437 t_obj = getattr(LayerType, self.eval_cfg.nodes_mem_comp.tail.name) 

438 self._ctx.tail_node = t_obj 

439 self.set_tail_eval_fun( 

440 num_p=self.eval_cfg.nodes_mem_comp.tail.num_param_fun, 

441 stat_p=self.eval_cfg.nodes_mem_comp.tail.stat_fun.p, 

442 stat_os=self.eval_cfg.nodes_mem_comp.tail.stat_fun.os, 

443 stat_grad=self.eval_cfg.nodes_mem_comp.tail.stat_fun.grad, 

444 dyn_activ=self.eval_cfg.nodes_mem_comp.tail.dyn_fun.activation, 

445 dyn_dp_comm=self.eval_cfg.nodes_mem_comp.tail.dyn_fun.comm.dp, 

446 dyn_tp_comm=self.eval_cfg.nodes_mem_comp.tail.dyn_fun.comm.tp, 

447 dyn_cp_comm=self.eval_cfg.nodes_mem_comp.tail.dyn_fun.comm.cp, 

448 dyn_ep_comm=self.eval_cfg.nodes_mem_comp.tail.dyn_fun.comm.ep, 

449 ) 

450 

451 # body 

452 for b in self.eval_cfg.nodes_mem_comp.body: 

453 b_cfg = Config(b) 

454 comm_cfg = b_cfg.dyn_fun.comm 

455 comm_kwargs = { 

456 "dyn_dp_comm": comm_cfg.dp, 

457 "dyn_tp_comm": comm_cfg.tp, 

458 "dyn_cp_comm": comm_cfg.cp, 

459 "dyn_ep_comm": comm_cfg.ep, 

460 } 

461 if hasattr(comm_cfg, "ep_balanced"): 

462 comm_kwargs["dyn_ep_comm_balanced"] = comm_cfg.ep_balanced 

463 if hasattr(comm_cfg, "ep_imbalanced"): 

464 comm_kwargs["dyn_ep_comm_imbalanced"] = comm_cfg.ep_imbalanced 

465 # Compute config (FLOPs, not memory) 

466 compute_kwargs = {} 

467 compute_cfg = getattr(b_cfg.dyn_fun, "compute", None) 

468 if compute_cfg and hasattr(compute_cfg, "__dict__"): 

469 compute_kwargs = {"compute": dict(vars(compute_cfg))} 

470 self.set_body_eval_fun( 

471 lay_type=b_cfg.name, 

472 num_p=b_cfg.num_param_fun, 

473 stat_p=b_cfg.stat_fun.p, 

474 stat_os=b_cfg.stat_fun.os, 

475 stat_grad=b_cfg.stat_fun.grad, 

476 dyn_activ=b_cfg.dyn_fun.activation, 

477 **comm_kwargs, 

478 **compute_kwargs, 

479 ) 

480 

481 # pp micro factor 

482 for sc in self.eval_cfg.pp_sched: 

483 self.set_pp_micro_factor_eval_fun(sc["name"], sc["fun"]) 

484 if not self._ccfg.pp_sched: 

485 self._ccfg.pp_sched = self.eval_cfg.default_pp_sched 

486 

487 # layerblock 

488 self.set_attn_eval_fun( 

489 self.eval_cfg.base_arch_mem_comp.attention.num_param_fun, 

490 self.eval_cfg.base_arch_mem_comp.attention.qkv, 

491 self.eval_cfg.base_arch_mem_comp.attention.score, 

492 self.eval_cfg.base_arch_mem_comp.attention.proj, 

493 ) 

494 self.set_ffn_eval_fun( 

495 self.eval_cfg.base_arch_mem_comp.feedforward.num_param_fun, 

496 self.eval_cfg.base_arch_mem_comp.feedforward.activation, 

497 self.eval_cfg.base_arch_mem_comp.feedforward.moe_activ, 

498 ) 

499 self.set_expert_param_eval_fun( 

500 routed_num_p=self.eval_cfg.base_arch_mem_comp.feedforward.routed_num_fun, 

501 shared_num_p=self.eval_cfg.base_arch_mem_comp.feedforward.shared_num_fun, 

502 ) 

503 self.set_norm_eval_fun( 

504 self.eval_cfg.base_arch_mem_comp.norm.num_param_fun, 

505 self.eval_cfg.base_arch_mem_comp.norm.activation, 

506 ) 

507 

508 # passes 

509 self.set_passes( 

510 vpp_less_mem=self.eval_cfg.passes.vpp_less_memory, 

511 swap_os=self.eval_cfg.passes.swap_optimizer, 

512 dropless_tok_factor=self.eval_cfg.passes.dropless_tok_factor, 

513 ) 

514 

515 self._ctx.comm_expr = self.eval_cfg.comm_expr 

516 

517 def is_regular_layer(self, lay): 

518 """check if layer is not head/tail""" 

519 if isinstance(lay, str): 

520 return lay[0] not in [ 

521 self._ctx.head_node.name[0], 

522 self._ctx.tail_node.name[0], 

523 ] 

524 return lay not in [self._ctx.head_node, self._ctx.tail_node]