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

Class TransformerLayer

codegeex/paddle/codegeex_model.py:400–475  ·  view source on GitHub ↗

A single transformer layer. Transformore layer takes input with size [b, s, h] and returns an output of the same size.

Source from the content-addressed store, hash-verified

398
399
400class TransformerLayer(paddle.nn.Layer):
401 """A single transformer layer.
402
403 Transformore layer takes input with size [b, s, h] and returns an
404 output of the same size.
405 """
406
407 def __init__(
408 self,
409 hidden_size,
410 num_attention_heads,
411 layer_number,
412 layernorm_epsilon=1e-5,
413 fp16=True,
414 attention_softmax_in_fp32=True,
415 ):
416 super(TransformerLayer, self).__init__()
417 self.hidden_size = hidden_size
418 self.layernorm_epsilon = layernorm_epsilon
419 self.layer_number = layer_number
420
421 # Layernorm on the input data.
422 self.input_layernorm = paddle.nn.LayerNorm(hidden_size,
423 epsilon=self.layernorm_epsilon)
424
425 # Self attention.
426 self.attention = SelfAttention(hidden_size,
427 num_attention_heads,
428 layer_number,
429 fp16,
430 attention_softmax_in_fp32)
431
432 # Layernorm on the input data.
433 self.post_attention_layernorm = paddle.nn.LayerNorm(self.hidden_size,
434 epsilon=self.layernorm_epsilon)
435 self.mlp = MLP(self.hidden_size)
436
437 def forward(
438 self,
439 hidden_states,
440 attention_mask,
441 layer_past=None,
442 get_key_value=False,
443 prompt_length=None,
444 context_length=None,
445 ):
446 # hidden_states: [b, s, h]
447 # Use FP32 for Layernorm
448 # layernorm_output = self.input_layernorm(hidden_states.cast("float32")).cast("float16")
449 layernorm_output = self.input_layernorm(hidden_states)
450
451 # Self attention.
452 attention_output = self.attention(layernorm_output,
453 attention_mask,
454 layer_past=layer_past,
455 get_key_value=get_key_value,
456 prompt_length=prompt_length,
457 context_length=context_length)

Callers 1

build_layerMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected