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

Class JointEncoder

model.py:24–315  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

22from torch.utils.checkpoint import checkpoint
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(
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 )

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected