fix gelu | gelu_new

This commit is contained in:
Patrick von Platen
2020-12-12 13:29:41 +00:00
parent 2787f5251b
commit a5bc339c4d
3 changed files with 83 additions and 62 deletions
+3 -15
View File
@@ -34,20 +34,8 @@ logger = logging.get_logger(__name__)
@jax.jit
def gelu(x):
r"""
Gaussian error linear unit activation function.
Computes the element-wise function:
.. math::
\mathrm{gelu}(x) = \frac{x}{2} \left(1 + \mathrm{tanh} \left(
\sqrt{\frac{2}{\pi}} \left(x + 0.044715 x^3 \right) \right) \right)
We explicitly use the approximation rather than the exact formulation for speed. For more information, see
`Gaussian Error Linear Units (GELUs) <https://arxiv.org/abs/1606.08415>`_, section 2.
"""
return x * 0.5 * (1.0 + jax.lax.erf(x / jnp.sqrt(2.0)))
def gelu_new(x):
return nn.gelu(x, approximate=True)
ACT2FN = {
@@ -55,7 +43,7 @@ ACT2FN = {
"relu": nn.relu,
"silu": nn.swish,
"swish": nn.swish,
"gelu_new": gelu,
"gelu_new": gelu_new,
}
@@ -23,7 +23,7 @@ import jax.numpy as jnp
from jax.random import PRNGKey
from ...file_utils import add_start_docstrings, add_start_docstrings_to_model_forward
from ...modeling_flax_utils import FlaxPreTrainedModel, gelu
from ...modeling_flax_utils import ACT2FN, FlaxPreTrainedModel
from ...utils import logging
from .configuration_bert import BertConfig
@@ -205,7 +205,7 @@ class FlaxBertAttention(nn.Module):
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@nn.compact
def __call__(self, hidden_state, attention_mask, deterministic: bool = True):
def __call__(self, hidden_states, attention_mask, deterministic: bool = True):
# Attention mask comes in as attention_mask.shape == (*batch_sizes, kv_length)
# FLAX expects: attention_mask.shape == (*batch_sizes, 1, 1, kv_length) such that it is broadcastable
# with attn_weights.shape == (*batch_sizes, num_heads, q_length, kv_length)
@@ -219,27 +219,28 @@ class FlaxBertAttention(nn.Module):
bias_init=jax.nn.initializers.zeros,
name="self",
dtype=self.dtype,
)(hidden_state, attention_mask)
)(hidden_states, attention_mask)
layer_norm = FlaxBertLayerNorm(name="layer_norm", dtype=self.dtype)(self_att + hidden_state)
layer_norm = FlaxBertLayerNorm(name="layer_norm", dtype=self.dtype)(self_att + hidden_states)
return layer_norm
class FlaxBertIntermediate(nn.Module):
output_size: int
hidden_act: str
kernel_init_scale: float = 0.2
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@nn.compact
def __call__(self, hidden_state):
# TODO: Add ACT2FN reference to change activation function
dense = nn.Dense(
def __call__(self, hidden_states):
hidden_states = nn.Dense(
features=self.output_size,
kernel_init=jax.nn.initializers.normal(self.kernel_init_scale, self.dtype),
name="dense",
dtype=self.dtype,
)(hidden_state)
return gelu(dense)
)(hidden_states)
hidden_states = ACT2FN[self.hidden_act](hidden_states)
return hidden_states
class FlaxBertOutput(nn.Module):
@@ -249,27 +250,28 @@ class FlaxBertOutput(nn.Module):
@nn.compact
def __call__(self, intermediate_output, attention_output, deterministic: bool = True):
hidden_state = nn.Dense(
hidden_states = nn.Dense(
attention_output.shape[-1],
kernel_init=jax.nn.initializers.normal(self.kernel_init_scale, self.dtype),
name="dense",
dtype=self.dtype,
)(intermediate_output)
hidden_state = nn.Dropout(rate=self.dropout_rate)(hidden_state, deterministic=deterministic)
hidden_state = FlaxBertLayerNorm(name="layer_norm", dtype=self.dtype)(hidden_state + attention_output)
return hidden_state
hidden_states = nn.Dropout(rate=self.dropout_rate)(hidden_states, deterministic=deterministic)
hidden_states = FlaxBertLayerNorm(name="layer_norm", dtype=self.dtype)(hidden_states + attention_output)
return hidden_states
class FlaxBertLayer(nn.Module):
num_heads: int
head_size: int
intermediate_size: int
hidden_act: str
dropout_rate: float = 0.0
kernel_init_scale: float = 0.2
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@nn.compact
def __call__(self, hidden_state, attention_mask, deterministic: bool = True):
def __call__(self, hidden_states, attention_mask, deterministic: bool = True):
attention = FlaxBertAttention(
self.num_heads,
self.head_size,
@@ -277,9 +279,13 @@ class FlaxBertLayer(nn.Module):
dropout_rate=self.dropout_rate,
name="attention",
dtype=self.dtype,
)(hidden_state, attention_mask, deterministic=deterministic)
)(hidden_states, attention_mask, deterministic=deterministic)
intermediate = FlaxBertIntermediate(
self.intermediate_size, kernel_init_scale=self.kernel_init_scale, name="intermediate", dtype=self.dtype
self.intermediate_size,
kernel_init_scale=self.kernel_init_scale,
hidden_act=self.hidden_act,
name="intermediate",
dtype=self.dtype,
)(attention)
output = FlaxBertOutput(
kernel_init_scale=self.kernel_init_scale, dropout_rate=self.dropout_rate, name="output", dtype=self.dtype
@@ -297,6 +303,7 @@ class FlaxBertLayerCollection(nn.Module):
num_heads: int
head_size: int
intermediate_size: int
hidden_act: str
dropout_rate: float = 0.0
kernel_init_scale: float = 0.2
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@@ -316,6 +323,7 @@ class FlaxBertLayerCollection(nn.Module):
self.intermediate_size,
kernel_init_scale=self.kernel_init_scale,
dropout_rate=self.dropout_rate,
hidden_act=self.hidden_act,
name=f"{i}",
dtype=self.dtype,
)
@@ -328,22 +336,24 @@ class FlaxBertEncoder(nn.Module):
num_heads: int
head_size: int
intermediate_size: int
hidden_act: str
dropout_rate: float = 0.0
kernel_init_scale: float = 0.2
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@nn.compact
def __call__(self, hidden_state, attention_mask, deterministic: bool = True):
def __call__(self, hidden_states, attention_mask, deterministic: bool = True):
layer = FlaxBertLayerCollection(
self.num_layers,
self.num_heads,
self.head_size,
self.intermediate_size,
hidden_act=self.hidden_act,
kernel_init_scale=self.kernel_init_scale,
dropout_rate=self.dropout_rate,
name="layer",
dtype=self.dtype,
)(hidden_state, attention_mask, deterministic=deterministic)
)(hidden_states, attention_mask, deterministic=deterministic)
return layer
@@ -352,10 +362,10 @@ class FlaxBertPooler(nn.Module):
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@nn.compact
def __call__(self, hidden_state):
cls_token = hidden_state[:, 0]
def __call__(self, hidden_states):
cls_token = hidden_states[:, 0]
out = nn.Dense(
hidden_state.shape[-1],
hidden_states.shape[-1],
kernel_init=jax.nn.initializers.normal(self.kernel_init_scale, self.dtype),
name="dense",
dtype=self.dtype,
@@ -364,17 +374,19 @@ class FlaxBertPooler(nn.Module):
class FlaxBertPredictionHeadTransform(nn.Module):
hidden_act: str
dtype: jnp.dtype = jnp.float32
@nn.compact
def __call__(self, hidden_states):
hidden_states = nn.Dense(hidden_states.shape[-1], name="dense", dtype=self.dtype)(hidden_states)
hidden_states = gelu(hidden_states) # TODO: ACT2FN[config.hidden_act]
hidden_states = ACT2FN[self.hidden_act](hidden_states)
return FlaxBertLayerNorm(name="layer_norm", dtype=self.dtype)(hidden_states)
class FlaxBertLMPredictionHead(nn.Module):
vocab_size: int
hidden_act: str
dtype: jnp.dtype = jnp.float32
@nn.compact
@@ -384,20 +396,23 @@ class FlaxBertLMPredictionHead(nn.Module):
# Need a link between the two variables so that the bias is correctly
# resized with `resize_token_embeddings`
hidden_states = FlaxBertPredictionHeadTransform(name="transform", dtype=self.dtype)(hidden_states)
hidden_states = FlaxBertPredictionHeadTransform(
name="transform", hidden_act=self.hidden_act, dtype=self.dtype
)(hidden_states)
hidden_states = nn.Dense(self.vocab_size, name="decoder", dtype=self.dtype)(hidden_states)
return hidden_states
class FlaxBertOnlyMLMHead(nn.Module):
vocab_size: int
hidden_act: str
dtype: jnp.dtype = jnp.float32
@nn.compact
def __call__(self, hidden_states):
hidden_states = FlaxBertLMPredictionHead(vocab_size=self.vocab_size, name="predictions", dtype=self.dtype)(
hidden_states
)
hidden_states = FlaxBertLMPredictionHead(
vocab_size=self.vocab_size, hidden_act=self.hidden_act, name="predictions", dtype=self.dtype
)(hidden_states)
return hidden_states
@@ -530,6 +545,7 @@ class FlaxBertModel(FlaxBertPreTrainedModel):
head_size=config.hidden_size,
intermediate_size=config.intermediate_size,
dropout_rate=config.hidden_dropout_prob,
hidden_act=config.hidden_act,
dtype=dtype,
)
@@ -575,6 +591,7 @@ class FlaxBertModule(nn.Module):
num_heads: int
head_size: int
intermediate_size: int
hidden_act: str = "gelu"
dropout_rate: float = 0.0
kernel_init_scale: float = 0.2
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@@ -603,6 +620,7 @@ class FlaxBertModule(nn.Module):
self.intermediate_size,
kernel_init_scale=self.kernel_init_scale,
dropout_rate=self.dropout_rate,
hidden_act=self.hidden_act,
name="encoder",
dtype=self.dtype,
)(embeddings, attention_mask, deterministic=deterministic)
@@ -627,6 +645,7 @@ class FlaxBertForMaskedLM(FlaxBertPreTrainedModel):
num_heads=config.num_attention_heads,
num_encoder_layers=config.num_hidden_layers,
max_length=config.max_length,
hidden_act=config.hidden_act,
**kwargs,
)
@@ -673,6 +692,7 @@ class FlaxBertForMaskedLMModule(nn.Module):
num_encoder_layers: int
type_vocab_size: int
max_length: int
hidden_act: str
dropout_rate: float = 0.0
dtype: jnp.dtype = jnp.float32
@@ -691,6 +711,7 @@ class FlaxBertForMaskedLMModule(nn.Module):
num_encoder_layers=self.num_encoder_layers,
max_length=self.max_length,
dropout_rate=self.dropout_rate,
hidden_act=self.hidden_act,
dtype=self.dtype,
add_pooling_layer=False,
name="bert",
@@ -698,6 +719,8 @@ class FlaxBertForMaskedLMModule(nn.Module):
# Compute the prediction scores
encoder = nn.Dropout(rate=self.dropout_rate)(encoder, deterministic=deterministic)
logits = FlaxBertOnlyMLMHead(vocab_size=self.vocab_size, name="cls", dtype=self.dtype)(encoder)
logits = FlaxBertOnlyMLMHead(
vocab_size=self.vocab_size, hidden_act=self.hidden_act, name="cls", dtype=self.dtype
)(encoder)
return logits
@@ -22,7 +22,7 @@ import jax.numpy as jnp
from jax.random import PRNGKey
from ...file_utils import add_start_docstrings, add_start_docstrings_to_model_forward
from ...modeling_flax_utils import FlaxPreTrainedModel, gelu
from ...modeling_flax_utils import ACT2FN, FlaxPreTrainedModel
from ...utils import logging
from .configuration_roberta import RobertaConfig
@@ -225,7 +225,7 @@ class FlaxRobertaAttention(nn.Module):
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@nn.compact
def __call__(self, hidden_state, attention_mask, deterministic: bool = True):
def __call__(self, hidden_states, attention_mask, deterministic: bool = True):
# Attention mask comes in as attention_mask.shape == (*batch_sizes, kv_length)
# FLAX expects: attention_mask.shape == (*batch_sizes, 1, 1, kv_length) such that it is broadcastable
# with attn_weights.shape == (*batch_sizes, num_heads, q_length, kv_length)
@@ -239,28 +239,29 @@ class FlaxRobertaAttention(nn.Module):
bias_init=jax.nn.initializers.zeros,
name="self",
dtype=self.dtype,
)(hidden_state, attention_mask)
)(hidden_states, attention_mask)
layer_norm = FlaxRobertaLayerNorm(name="layer_norm", dtype=self.dtype)(self_att + hidden_state)
layer_norm = FlaxRobertaLayerNorm(name="layer_norm", dtype=self.dtype)(self_att + hidden_states)
return layer_norm
# Copied from transformers.models.bert.modeling_flax_bert.FlaxBertIntermediate with Bert->Roberta
class FlaxRobertaIntermediate(nn.Module):
output_size: int
hidden_act: str
kernel_init_scale: float = 0.2
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@nn.compact
def __call__(self, hidden_state):
# TODO: Add ACT2FN reference to change activation function
dense = nn.Dense(
def __call__(self, hidden_states):
hidden_states = nn.Dense(
features=self.output_size,
kernel_init=jax.nn.initializers.normal(self.kernel_init_scale, self.dtype),
name="dense",
dtype=self.dtype,
)(hidden_state)
return gelu(dense)
)(hidden_states)
hidden_states = ACT2FN[self.hidden_act](hidden_states)
return hidden_states
# Copied from transformers.models.bert.modeling_flax_bert.FlaxBertOutput with Bert->Roberta
@@ -271,27 +272,28 @@ class FlaxRobertaOutput(nn.Module):
@nn.compact
def __call__(self, intermediate_output, attention_output, deterministic: bool = True):
hidden_state = nn.Dense(
hidden_states = nn.Dense(
attention_output.shape[-1],
kernel_init=jax.nn.initializers.normal(self.kernel_init_scale, self.dtype),
name="dense",
dtype=self.dtype,
)(intermediate_output)
hidden_state = nn.Dropout(rate=self.dropout_rate)(hidden_state, deterministic=deterministic)
hidden_state = FlaxRobertaLayerNorm(name="layer_norm", dtype=self.dtype)(hidden_state + attention_output)
return hidden_state
hidden_states = nn.Dropout(rate=self.dropout_rate)(hidden_states, deterministic=deterministic)
hidden_states = FlaxRobertaLayerNorm(name="layer_norm", dtype=self.dtype)(hidden_states + attention_output)
return hidden_states
class FlaxRobertaLayer(nn.Module):
num_heads: int
head_size: int
intermediate_size: int
hidden_act: str
dropout_rate: float = 0.0
kernel_init_scale: float = 0.2
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@nn.compact
def __call__(self, hidden_state, attention_mask, deterministic: bool = True):
def __call__(self, hidden_states, attention_mask, deterministic: bool = True):
attention = FlaxRobertaAttention(
self.num_heads,
self.head_size,
@@ -299,10 +301,11 @@ class FlaxRobertaLayer(nn.Module):
dropout_rate=self.dropout_rate,
name="attention",
dtype=self.dtype,
)(hidden_state, attention_mask, deterministic=deterministic)
)(hidden_states, attention_mask, deterministic=deterministic)
intermediate = FlaxRobertaIntermediate(
self.intermediate_size,
kernel_init_scale=self.kernel_init_scale,
hidden_act=self.hidden_act,
name="intermediate",
dtype=self.dtype,
)(attention)
@@ -323,6 +326,7 @@ class FlaxRobertaLayerCollection(nn.Module):
num_heads: int
head_size: int
intermediate_size: int
hidden_act: str
dropout_rate: float = 0.0
kernel_init_scale: float = 0.2
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@@ -342,6 +346,7 @@ class FlaxRobertaLayerCollection(nn.Module):
self.intermediate_size,
kernel_init_scale=self.kernel_init_scale,
dropout_rate=self.dropout_rate,
hidden_act=self.hidden_act,
name=f"{i}",
dtype=self.dtype,
)
@@ -355,22 +360,24 @@ class FlaxRobertaEncoder(nn.Module):
num_heads: int
head_size: int
intermediate_size: int
hidden_act: str
dropout_rate: float = 0.0
kernel_init_scale: float = 0.2
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@nn.compact
def __call__(self, hidden_state, attention_mask, deterministic: bool = True):
def __call__(self, hidden_states, attention_mask, deterministic: bool = True):
layer = FlaxRobertaLayerCollection(
self.num_layers,
self.num_heads,
self.head_size,
self.intermediate_size,
hidden_act=self.hidden_act,
kernel_init_scale=self.kernel_init_scale,
dropout_rate=self.dropout_rate,
name="layer",
dtype=self.dtype,
)(hidden_state, attention_mask, deterministic=deterministic)
)(hidden_states, attention_mask, deterministic=deterministic)
return layer
@@ -380,10 +387,10 @@ class FlaxRobertaPooler(nn.Module):
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@nn.compact
def __call__(self, hidden_state):
cls_token = hidden_state[:, 0]
def __call__(self, hidden_states):
cls_token = hidden_states[:, 0]
out = nn.Dense(
hidden_state.shape[-1],
hidden_states.shape[-1],
kernel_init=jax.nn.initializers.normal(self.kernel_init_scale, self.dtype),
name="dense",
dtype=self.dtype,
@@ -513,6 +520,7 @@ class FlaxRobertaModel(FlaxRobertaPreTrainedModel):
num_encoder_layers=config.num_hidden_layers,
num_heads=config.num_attention_heads,
head_size=config.hidden_size,
hidden_act=config.hidden_act,
intermediate_size=config.intermediate_size,
dropout_rate=config.hidden_dropout_prob,
dtype=dtype,
@@ -561,6 +569,7 @@ class FlaxRobertaModule(nn.Module):
num_heads: int
head_size: int
intermediate_size: int
hidden_act: str = "gelu"
dropout_rate: float = 0.0
kernel_init_scale: float = 0.2
dtype: jnp.dtype = jnp.float32 # the dtype of the computation
@@ -589,6 +598,7 @@ class FlaxRobertaModule(nn.Module):
self.intermediate_size,
kernel_init_scale=self.kernel_init_scale,
dropout_rate=self.dropout_rate,
hidden_act=self.hidden_act,
name="encoder",
dtype=self.dtype,
)(embeddings, attention_mask, deterministic=deterministic)