| 22 | from torch.utils.checkpoint import checkpoint |
| 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( |
| 52 | "`T5Stack.parallelize` is deprecated and will be removed in v5 of Transformers, you should load your model" |
| 53 | " with `device_map='balanced'` in the call to `from_pretrained`. You can also provide your own" |
| 54 | " `device_map` but it needs to be a dictionary module_name to device, so for instance {'block.0': 0," |
| 55 | " 'block.1': 1, ...}", |
| 56 | FutureWarning, |
| 57 | ) |
| 58 | # Check validity of device_map |
| 59 | self.device_map = ( |
| 60 | get_device_map(len(self.block), range(torch.cuda.device_count())) if device_map is None else device_map |
| 61 | ) |
| 62 | assert_device_map(self.device_map, len(self.block)) |
| 63 | self.model_parallel = True |
| 64 | self.first_device = "cpu" if "cpu" in self.device_map.keys() else "cuda:" + str(min(self.device_map.keys())) |
| 65 | self.last_device = "cuda:" + str(max(self.device_map.keys())) |
| 66 | # Load onto devices |
| 67 | for k, v in self.device_map.items(): |
| 68 | for layer in v: |
| 69 | cuda_device = "cuda:" + str(k) |
| 70 | self.block[layer] = self.block[layer].to(cuda_device) |
| 71 | |
| 72 | # Set embed_tokens to first layer |
| 73 | self.embed_tokens = self.embed_tokens.to(self.first_device) |
| 74 | # Set final layer norm to last device |
| 75 | self.final_layer_norm = self.final_layer_norm.to(self.last_device) |
| 76 | |
| 77 | def deparallelize(self): |
| 78 | warnings.warn( |
| 79 | "Like `parallelize`, `deparallelize` is deprecated and will be removed in v5 of Transformers.", |
| 80 | FutureWarning, |
| 81 | ) |