File size: 3,129 Bytes
d91766b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 | from __future__ import annotations
from typing import Any
MISSING = object()
def _read_external_attr(config: Any, names: tuple[str, ...], default: Any = MISSING) -> Any:
for name in names:
value = getattr(config, name, MISSING)
if value is not MISSING and value is not None:
return value
if default is MISSING:
joined = ", ".join(names)
raise AttributeError(f"Config does not define any of: {joined}.")
return default
def _read_external_int(config: Any, names: tuple[str, ...], default: int | object = MISSING) -> int:
return int(_read_external_attr(config, names, default))
def _read_external_float(config: Any, names: tuple[str, ...], default: float | object = MISSING) -> float:
return float(_read_external_attr(config, names, default))
def _read_external_bool(config: Any, names: tuple[str, ...], default: bool | object = MISSING) -> bool:
return bool(_read_external_attr(config, names, default))
def _read_external_int_tuple(config: Any, names: tuple[str, ...], default: tuple[int, ...] = ()) -> tuple[int, ...]:
value = _read_external_attr(config, names, default)
return tuple(int(layer_idx) for layer_idx in value)
def get_num_experts(config: Any) -> int:
value = _read_external_attr(
config,
("num_experts", "n_routed_experts", "num_local_experts"),
0,
)
return int(value or 0)
def get_num_experts_per_tok(config: Any) -> int:
value = _read_external_attr(
config,
("num_experts_per_tok", "moe_top_k", "top_k"),
2,
)
return int(value)
def get_moe_intermediate_size(config: Any) -> int:
value = _read_external_attr(
config,
("moe_intermediate_size", "expert_intermediate_size", "intermediate_size"),
)
return int(value)
def get_norm_topk_prob(config: Any) -> bool:
return _read_external_bool(config, ("norm_topk_prob",), True)
def get_mlp_only_layers(config: Any) -> tuple[int, ...]:
return _read_external_int_tuple(config, ("mlp_only_layers",), ())
def get_first_k_dense_replace(config: Any) -> int | None:
value = _read_external_attr(config, ("first_k_dense_replace",), None)
if value is None:
return None
return int(value)
def get_decoder_sparse_step(config: Any) -> int:
return _read_external_int(config, ("decoder_sparse_step",), 1)
def is_moe_layer(config: Any, layer_idx: int) -> bool:
if get_num_experts(config) <= 0:
return False
if layer_idx in get_mlp_only_layers(config):
return False
first_k_dense_replace = get_first_k_dense_replace(config)
if first_k_dense_replace is not None:
return layer_idx >= first_k_dense_replace
decoder_sparse_step = get_decoder_sparse_step(config)
if decoder_sparse_step <= 0:
return False
return (layer_idx + 1) % decoder_sparse_step == 0
__all__ = [
"get_mlp_only_layers",
"get_moe_intermediate_size",
"get_num_experts",
"get_num_experts_per_tok",
"get_norm_topk_prob",
"get_first_k_dense_replace",
"get_decoder_sparse_step",
"is_moe_layer",
]
|