| 23 | |
| 24 | class 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( |