| 85 | from sat.ops.layernorm import RMSNorm |
| 86 | |
| 87 | class 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 | |