File size: 1,629 Bytes
16533aa | 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 | from transformers import PretrainedConfig, AutoConfig
class ParallelMLPConfig(PretrainedConfig):
model_type = "parallel_mlp_causal_lm"
def __init__(
self,
base_model_name_or_path: str = "meta-llama/Meta-Llama-3.1-8B-Instruct",
mlp_positions: list[int] | None = None,
mlp_intermediate_size: int | None = None,
**kwargs,
):
self.base_model_name_or_path = base_model_name_or_path
self.mlp_positions = mlp_positions if mlp_positions is not None else []
try:
base_cfg = AutoConfig.from_pretrained(base_model_name_or_path, trust_remote_code=True)
self.hidden_size = getattr(base_cfg, "hidden_size", 4096)
self.rms_norm_eps = getattr(base_cfg, "rms_norm_eps", 1e-5)
self.num_hidden_layers = getattr(base_cfg, "num_hidden_layers", 32)
except Exception:
self.hidden_size = 4096
self.rms_norm_eps = 1e-5
self.num_hidden_layers = 32
# Intermediate size: ~8 × hidden_size when matched to CALYREX late8th budget
if mlp_intermediate_size is not None:
self.mlp_intermediate_size = mlp_intermediate_size
else:
self.mlp_intermediate_size = 8 * self.hidden_size
for pos in self.mlp_positions:
if pos < 1 or pos > self.num_hidden_layers:
raise ValueError(
f"mlp_position {pos} is out of bounds (1 to {self.num_hidden_layers})"
)
super().__init__(**kwargs)
AutoConfig.register("parallel_mlp_causal_lm", ParallelMLPConfig)
|