MCPcopy Create free account
hub / github.com/THUDM/GLM / BertPredictionHeadTransform

Class BertPredictionHeadTransform

model/modeling_bert.py:590–608  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

588
589
590class BertPredictionHeadTransform(nn.Module):
591 def __init__(self, config):
592 super(BertPredictionHeadTransform, self).__init__()
593 self.dense = nn.Linear(config.hidden_size, config.hidden_size)
594 self.transform_act_fn = ACT2FN[config.hidden_act] \
595 if isinstance(config.hidden_act, str) else config.hidden_act
596 self.LayerNorm = BertLayerNorm(config.hidden_size, eps=config.layernorm_epsilon)
597 self.fp32_layernorm = config.fp32_layernorm
598
599 def forward(self, hidden_states):
600 hidden_states = self.dense(hidden_states)
601 hidden_states = self.transform_act_fn(hidden_states)
602 previous_type = hidden_states.type()
603 if self.fp32_layernorm:
604 hidden_states = hidden_states.float()
605 hidden_states = self.LayerNorm(hidden_states)
606 if self.fp32_layernorm:
607 hidden_states = hidden_states.type(previous_type)
608 return hidden_states
609
610
611class BertLMPredictionHead(nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected