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

Class MAE

SwissArmyTransformer/sat/model/official/mae_model.py:135–175  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

133import argparse
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)
152
153 def decode(self, input_ids, position_ids, attention_mask, encoder_outputs, ids_restore, **kw_args):
154 return self.decoder(input_ids, position_ids, attention_mask, encoder_outputs=encoder_outputs, ids_restore=ids_restore, **kw_args)
155
156 def forward(self, input_ids, enc_position_ids, dec_position_ids, *, enc_attention_mask=None, dec_attention_mask=None, **kw_args):
157 if enc_attention_mask is None:
158 enc_attention_mask = torch.ones(1, 1, dtype=self.encoder.transformer.word_embeddings.weight.dtype, device=input_ids.device)
159 encoder_outputs, *encoder_mems = self.encode(input_ids, enc_position_ids, enc_attention_mask, **kw_args)
160 decoder_outputs, *decoder_mems = self.decode(input_ids, dec_position_ids, dec_attention_mask, encoder_outputs=encoder_outputs, ids_restore=encoder_mems[0]["ids_restore"], **kw_args)
161 return encoder_outputs, decoder_outputs, encoder_mems, decoder_mems
162
163 def unpatchify(self, x):
164 """
165 x: (N, L, patch_size**2 *3)
166 imgs: (N, 3, H, W)
167 """
168 p = self.encoder.property.patch_size
169 h = w = int(x.shape[1]**.5)
170 assert h * w == x.shape[1]
171
172 x = x.reshape(shape=(x.shape[0], h, w, p, p, 3))
173 x = torch.einsum('nhwpqc->nchpwq', x)
174 imgs = x.reshape(shape=(x.shape[0], 3, h * p, h * p))
175 return imgs

Callers 1

transform_param.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected