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

Class LLaMAModel

SwissArmyTransformer/sat/model/official/llama_model.py:91–112  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

89from sat.ops.layernorm import RMSNorm
90
91class 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

Callers 1

transform_param.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected