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

Method __init__

SwissArmyTransformer/sat/model/official/mae_model.py:136–148  ·  view source on GitHub ↗
(self, args, transformer=None, layernorm_epsilon=1e-6)

Source from the content-addressed store, hash-verified

134
135class MAE(EncoderDecoderModel):
136 def __init__(self, args, transformer=None, layernorm_epsilon=1e-6):
137 encoder = MAEEncoder(args, transformer=transformer, layernorm_epsilon=layernorm_epsilon)
138 dec_args = argparse.Namespace(**vars(args))
139 # dec_args.enc_hidden_size = dec_args.hidden_size # used for cross attn
140 override_attrs = ['num_layers', 'hidden_size', 'num_attention_heads',
141 'max_sequence_length', 'inner_hidden_size', 'hidden_size_per_attention_head']
142 for name in override_attrs:
143 dec_attr = getattr(dec_args, 'dec_' + name, None)
144 if dec_attr is not None: # else use encoder-config
145 setattr(dec_args, name, dec_attr)
146 setattr(dec_args, 'enc_hidden_size', args.hidden_size)
147 decoder = MAEDecoder(dec_args, transformer=transformer, layernorm_epsilon=layernorm_epsilon)
148 super().__init__(args, encoder=encoder, decoder=decoder, tie_word_embeddings=False)
149
150 def encode(self, input_ids, position_ids, attention_mask=None, **kw_args):
151 return self.encoder(input_ids, position_ids, attention_mask, **kw_args)

Callers 4

__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45

Calls 2

MAEEncoderClass · 0.85
MAEDecoderClass · 0.85

Tested by

no test coverage detected