MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / TopQuerySelfAttention

Class TopQuerySelfAttention

codegeex/torch/codegeex_model.py:227–398  ·  view source on GitHub ↗

Top query self-attention layer abstract class. Self-attention layer takes input with size [b, s, h] and returns output of the same size.

Source from the content-addressed store, hash-verified

225
226
227class TopQuerySelfAttention(torch.nn.Module):
228 """Top query self-attention layer abstract class.
229
230 Self-attention layer takes input with size [b, s, h]
231 and returns output of the same size.
232 """
233
234 def __init__(
235 self,
236 hidden_size,
237 num_attention_heads,
238 layer_number,
239 fp16=True,
240 attention_softmax_in_fp32=True,
241 ):
242 super(TopQuerySelfAttention, self).__init__()
243 self.hidden_size = hidden_size
244 self.num_attention_heads = num_attention_heads
245 self.fp16 = fp16
246 self.attention_softmax_in_fp32 = attention_softmax_in_fp32
247 self.layer_number = max(1, layer_number)
248
249 assert self.hidden_size % self.num_attention_heads == 0
250 self.hidden_size_per_attention_head = int(self.hidden_size // self.num_attention_heads)
251
252 self.query = torch.nn.Linear(self.hidden_size, self.hidden_size)
253 self.key = torch.nn.Linear(self.hidden_size, self.hidden_size)
254 self.value = torch.nn.Linear(self.hidden_size, self.hidden_size)
255
256 self.norm_factor = math.sqrt(self.hidden_size_per_attention_head)
257 self.softmax = torch.nn.Softmax(dim=-1)
258
259 self.dense = torch.nn.Linear(self.hidden_size, self.hidden_size)
260
261 def forward(
262 self,
263 hidden_states,
264 query_hidden_state,
265 attention_mask,
266 layer_past=None,
267 get_key_value=False,
268 prompt_length=None,
269 context_length=None,
270 ):
271
272 # hidden_states: [sq, b, h]
273 query_layer = self.query(query_hidden_state)
274 key_layer = self.key(hidden_states)
275 value_layer = self.value(hidden_states)
276
277 new_query_layer_shape = query_layer.size()[:-1] + \
278 (self.num_attention_heads,
279 self.hidden_size_per_attention_head)
280 query_layer = query_layer.view(*new_query_layer_shape)
281
282 new_query_layer_shape = key_layer.size()[:-1] + \
283 (self.num_attention_heads,
284 self.hidden_size_per_attention_head)

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected