MCPcopy Create free account
hub / github.com/amazon-science/mm-cot / __init__

Method __init__

model.py:25–48  ·  view source on GitHub ↗
(self, config, embed_tokens=None, patch_size=None)

Source from the content-addressed store, hash-verified

23
24class JointEncoder(T5Stack):
25 def __init__(self, config, embed_tokens=None, patch_size=None):
26 super().__init__(config)
27
28 self.embed_tokens = embed_tokens
29 self.is_decoder = config.is_decoder
30
31 self.patch_num, self.patch_dim = patch_size
32 self.image_dense = nn.Linear(self.patch_dim, config.d_model)
33 self.mha_layer = torch.nn.MultiheadAttention(embed_dim=config.hidden_size, kdim=config.hidden_size, vdim=config.hidden_size, num_heads=1, batch_first=True)
34 self.gate_dense = nn.Linear(2*config.hidden_size, config.hidden_size)
35 self.sigmoid = nn.Sigmoid()
36
37 self.block = nn.ModuleList(
38 [T5Block(config, has_relative_attention_bias=bool(i == 0)) for i in range(config.num_layers)]
39 )
40 self.final_layer_norm = T5LayerNorm(config.d_model, eps=config.layer_norm_epsilon)
41 self.dropout = nn.Dropout(config.dropout_rate)
42
43 # Initialize weights and apply final processing
44 self.post_init()
45 # Model parallel
46 self.model_parallel = False
47 self.device_map = None
48 self.gradient_checkpointing = False
49
50 def parallelize(self, device_map=None):
51 warnings.warn(

Callers 1

__init__Method · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected