MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / MixtralModel

Class MixtralModel

SwissArmyTransformer/sat/model/official/mixtral_model.py:87–106  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

85from sat.ops.layernorm import RMSNorm
86
87class MixtralModel(BaseModel):
88 def __init__(self, args, transformer=None, layernorm=RMSNorm, activation_func=nn.functional.silu, **kwargs):
89 super().__init__(args, transformer=transformer, layernorm=layernorm, activation_func=activation_func, init_method_std=0.01, **kwargs)
90 del self.transformer.position_embeddings
91 if 'inner_hidden_size' not in args:
92 args.inner_hidden_size = None
93 self.add_mixin("rotary", RotaryMixin(args.hidden_size, args.num_attention_heads))
94 self.add_mixin("lm", LMMixin(args.vocab_size, args.hidden_size))
95 self.add_mixin("mlp", MixtralMlpMixin(args.num_layers, args.hidden_size, args.num_experts, args.num_experts_per_tok))
96
97 def position_embedding_forward(self, *args, **kwargs):
98 return None
99
100 @classmethod
101 def add_model_specific_args(cls, parser):
102 group = parser.add_argument_group('Mixtral-8x7b', 'Mixtral-8x7b Configurations')
103 group.add_argument('--bos-token-id', type=int, default=1)
104 group.add_argument('--eos-token-id', type=int, default=2)
105 group.add_argument('--num-experts-per-tok', type=int, default=2)
106 return parser
107

Callers 1

transform_param.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected