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

Method forward

codegeex/torch/codegeex_model.py:438–476  ·  view source on GitHub ↗
(
        self,
        hidden_states,
        attention_mask,
        layer_past=None,
        get_key_value=False,
        prompt_length=None,
        context_length=None,
    )

Source from the content-addressed store, hash-verified

436 self.mlp = MLP(self.hidden_size)
437
438 def forward(
439 self,
440 hidden_states,
441 attention_mask,
442 layer_past=None,
443 get_key_value=False,
444 prompt_length=None,
445 context_length=None,
446 ):
447 # hidden_states: [b, s, h]
448 # Use FP32 for Layernorm
449 # layernorm_output = self.input_layernorm(hidden_states.float()).half()
450 layernorm_output = self.input_layernorm(hidden_states)
451
452 # Self attention.
453 attention_output = self.attention(layernorm_output,
454 attention_mask,
455 layer_past=layer_past,
456 get_key_value=get_key_value,
457 prompt_length=prompt_length,
458 context_length=context_length)
459
460 if get_key_value:
461 attention_output, presents = attention_output
462
463 # Residual connection.
464 residual = hidden_states
465 layernorm_input = attention_output + residual
466
467 # Use FP32 for Layernorm
468 # layernorm_output = self.post_attention_layernorm(layernorm_input.float()).half()
469 layernorm_output = self.post_attention_layernorm(layernorm_input)
470 mlp_output = self.mlp(layernorm_output)
471 output = mlp_output + layernorm_input
472
473 if get_key_value:
474 output = [output, presents]
475
476 return output
477
478
479class TopQueryLayer(torch.nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected