| 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 |
|
|
| |
| 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) |
|
|