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

Method __init__

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

Source from the content-addressed store, hash-verified

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.

Callers 5

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

Calls 2

CaiTEncoderClass · 0.85
CaiTDecoderClass · 0.85

Tested by

no test coverage detected