| 89 | from sat.ops.layernorm import RMSNorm |
| 90 | |
| 91 | class LLaMAModel(BaseModel): |
| 92 | def __init__(self, args, transformer=None, layernorm=RMSNorm, activation_func=nn.functional.silu, **kwargs): |
| 93 | super().__init__(args, transformer=transformer, layernorm=layernorm, activation_func=activation_func, init_method_std=0.01, **kwargs) |
| 94 | if 'inner_hidden_size' not in args: |
| 95 | args.inner_hidden_size = None |
| 96 | if not (hasattr(args, 'is_rotary_emb') and args.is_rotary_emb): |
| 97 | del self.transformer.position_embeddings |
| 98 | self.add_mixin("rotary", RotaryMixin(args.hidden_size, args.num_attention_heads)) |
| 99 | self.add_mixin("lm", LMMixin(args.vocab_size, args.hidden_size)) |
| 100 | if not (hasattr(args, 'is_gated_mlp') and args.is_gated_mlp): |
| 101 | self.add_mixin("mlp", LLaMAMlpMixin(args.num_layers, args.hidden_size, args.inner_hidden_size)) |
| 102 | |
| 103 | def position_embedding_forward(self, *args, **kwargs): |
| 104 | return None |
| 105 | |
| 106 | @classmethod |
| 107 | def add_model_specific_args(cls, parser): |
| 108 | group = parser.add_argument_group('LLaMA', 'LLaMA Configurations') |
| 109 | group.add_argument('--bos-token-id', type=int, default=0) |
| 110 | group.add_argument('--eos-token-id', type=int, default=1) |
| 111 | group.add_argument('--pad-token-id', type=int, default=-1) |
| 112 | return parser |
| 113 | |