File size: 1,629 Bytes
b72f63a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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)