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

Class CaiT

SwissArmyTransformer/sat/model/official/cait_model.py:174–196  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

172import argparse
173
174class CaiT(EncoderDecoderModel):
175 def __init__(self, args, transformer=None, layernorm_epsilon=1e-6):
176 encoder = CaiTEncoder(args, transformer=transformer, layernorm_epsilon=layernorm_epsilon)
177 dec_args = argparse.Namespace(**vars(args))
178 # dec_args.enc_hidden_size = dec_args.hidden_size # used for cross attn
179 override_attrs = ['num_layers', 'hidden_size', 'num_attention_heads', 'layernorm_order',
180 'max_sequence_length', 'inner_hidden_size', 'hidden_size_per_attention_head']
181 for name in override_attrs:
182 dec_attr = getattr(dec_args, 'dec_' + name, None)
183 if dec_attr is not None: # else use encoder-config
184 setattr(dec_args, name, dec_attr)
185 decoder = CaiTDecoder(dec_args, transformer=transformer, layernorm_epsilon=layernorm_epsilon)
186 super().__init__(args, encoder=encoder, decoder=decoder)
187
188 def forward(self, input_ids, enc_position_ids, dec_position_ids, *, enc_attention_mask=None, dec_attention_mask=None, cross_attention_mask=None, **kw_args):
189 # Please use self.decoder for auto-regressive generation.
190 if enc_attention_mask is None:
191 enc_attention_mask = torch.ones(1, 1, dtype=self.encoder.transformer.word_embeddings.weight.dtype, device=input_ids.device)
192 if cross_attention_mask is None:
193 cross_attention_mask = enc_attention_mask
194 encoder_outputs = self.encode(input_ids, enc_position_ids, enc_attention_mask, **kw_args)
195 decoder_outputs, *mems = self.decode(input_ids, dec_position_ids, dec_attention_mask, encoder_outputs=encoder_outputs, cross_attention_mask=cross_attention_mask, **kw_args)
196 return (encoder_outputs, decoder_outputs, *mems)

Callers 1

transform_param.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected