Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 6 additions & 7 deletions tests/jax/test_praxis_layers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from transformer_engine.jax.flax import RelativePositionBiases as flax_RelativePositionBiases
from transformer_engine.jax.flax import TransformerLayer as flax_TransformerLayer
from transformer_engine.jax.flax.module import Softmax
from transformer_engine.jax.flax.transformer import AttentionType
from transformer_engine.jax.fp8 import FP8Helper, is_fp8_available
from transformer_engine.jax.praxis import LayerNorm
from transformer_engine.jax.praxis import FusedSoftmax, LayerNorm
Expand DownExpand Up@@ -666,32 +665,32 @@ class MultiHeadAttnAttr:
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}]


Expand Down
105 changes: 82 additions & 23 deletions transformer_engine/jax/flax/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -197,17 +197,16 @@ def core_attention(query: Array,
dynamic_vector_slice_in_dim = vmap(lax.dynamic_slice_in_dim, in_axes=(None, 0, None, None))


class AttentionType(Enum):
"""TransformerLayerType."""
PADDING = AttnMaskType.PADDING_MASK
CAUSAL = AttnMaskType.CAUSAL_MASK


class MultiHeadAttention(nn.Module):
r"""
Multi-head Attention (MHA), including Query,
Key, Value and Output projection.

.. warning::

Argument :attr:`attn_type` is deprecated and superseded by :attr:`attn_mask_type`.
:attr:`attn_type` is ignored in version 0.10 and will be fully removed in version 0.11.

Parameters
----------
head_dim : int
Expand DownExpand Up@@ -245,8 +244,11 @@ class MultiHeadAttention(nn.Module):
Indicate if apply residual connection with the output of layer normalization.
output_layernorm : bool, default = False
Indicate if apply a layer normalization at the end of MHA.
attn_type: AttentionType, defult = AttentionType.PADDING
Indicate the format of the attention mask in the core attention.
attn_type: Any, defult = None
*Deprecated*, will be ignored in v0.10 and be fully removed in v0.11.
Please use `attn_mask_type` to config the attention mask.
attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.

Optimization parameters
-----------------------
Expand DownExpand Up@@ -282,7 +284,9 @@ class MultiHeadAttention(nn.Module):
bias_init: Initializer = nn.initializers.zeros
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
dtype: DType = jnp.float32
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
Expand All@@ -293,6 +297,14 @@ class MultiHeadAttention(nn.Module):
def __post_init__(self):
if self.kernel_init is None:
self.kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in', 'normal')
# TODO(rewang): remove attn_type after v0.11
if self.attn_type is not None:
warnings.warn(
"The 'attn_type' argument in the 'MultiHeadAttention' is"
" deprecated in version 0.10 and will be removed in version 0.11."
" Passing value in attn_type will be ignored, please use `attn_mask_type`"
" to config the attention mask type.",
category=DeprecationWarning)
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -570,9 +582,23 @@ def kv_init(key, shape, dtype):
if use_fused_attn:
assert mask is not None and mask.ndim == 4 # (b, 1, s_q, s_kv)
assert not self.transpose_batch_sequence

# TODO(rewang): make it configurable for pre_scale_bias
attn_bias_type = AttnBiasType.NO_BIAS if bias is None else AttnBiasType.POST_SCALE_BIAS

def canonicalize_attn_mask_type(attn_mask_type):
"""
Convert the string to AttnMaskType
"""
if attn_mask_type == 'causal':
return AttnMaskType.CAUSAL_MASK
if attn_mask_type == 'padding':
return AttnMaskType.PADDING_MASK
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

attn_mask_type = canonicalize_attn_mask_type(self.attn_mask_type)

if inputs_q is inputs_kv:
qkv_proj = qkv_proj.reshape((*qkv_proj.shape[:-1], self.num_heads, self.head_dim))
qkv_sharding_constraint = ('batch', 'length', 'qkv_dim', 'heads', 'kv')
Expand All@@ -583,7 +609,7 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
Expand All@@ -602,18 +628,27 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
sharding_type=first_sharding_type)
else:
softmax_type = SoftmaxType.SCALED
if self.attn_type is AttentionType.PADDING:
if mask is not None:
softmax_type = SoftmaxType.SCALED_MASKED
else:
softmax_type = SoftmaxType.SCALED_UPPER_TRIANG_MASKED

def convert_to_softmax_type(attn_mask_type, mask):
"""
Convert the string to SoftmaxType
"""
if attn_mask_type == 'causal':
return SoftmaxType.SCALED_UPPER_TRIANG_MASKED
if attn_mask_type == 'padding':
if mask is not None:
return SoftmaxType.SCALED_MASKED
return SoftmaxType.SCALED
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

softmax_type = convert_to_softmax_type(self.attn_mask_type, mask)

x = core_attention(query,
key,
Expand DownExpand Up@@ -765,6 +800,18 @@ class TransformerLayer(nn.Module):
an attention block and a feedforward network (MLP).
This standard layer is based on the paper “Attention Is All You Need”.

.. warning::

Argument :attr:`self_attn_mask_type` is introduced in version 0.10.
Starting from version 0.11, the default value will be `"causal"`.
However, to ensure compatibility with earlier versions, before 0.11,
the default value will be `"padding"` for the encoder and `"causal"` for the decoder.

.. note::

Argument :attr:`attention_mask` will be ignored when
:attr:`self_attn_mask_type` is set to `"causal"`.

Parameters
----------
hidden_size: int, default = 512
Expand DownExpand Up@@ -825,6 +872,8 @@ class TransformerLayer(nn.Module):
If set to TransformerLayerType.DECODER, an additional cross-attention block
is added after self-attention.this can be used for structures like `T5`
Transformer in conjunction with the TransformerLayerType.ENCODER option.
self_attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.
enable_relative_embedding: bool, default = True
Whether to enable relative embedding as shifting of attention logits.
relative_embedding: flax.linen.Module, default = None
Expand DownExpand Up@@ -878,6 +927,7 @@ class TransformerLayer(nn.Module):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: nn.Module = None
dtype: DType = jnp.float32
Expand All@@ -893,6 +943,19 @@ def __post_init__(self):
if self.mlp_kernel_init is None:
self.mlp_kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in',
'truncated_normal')
# TODO(rewang): default to 'causal' in 0.11 (also updated the doc after 0.11)
if self.self_attn_mask_type is None:
warnings.warn(
"The 'self_attn_mask_type' argument in the 'TransformerLayer' is"
" introduced in version 0.10. Starting from version 0.11, the default"
" value will be 'causal'. However, to ensure compatibility with earlier"
" versions, before 0.11, the default value will be 'padding' for the"
" encoder and 'causal' for the decoder.",
category=FutureWarning)
if self.layer_type == TransformerLayerType.ENCODER:
self.self_attn_mask_type = 'padding'
else:
self.self_attn_mask_type = 'causal'
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -975,16 +1038,12 @@ def __call__(self,

assert inputs.ndim == 3

self_attn_type = None
# Make name be the exactly same as T5X, since names would affect
# RNGKey during init and apply. Myabe no need in the feature.
if self.layer_type == TransformerLayerType.ENCODER:
mha_name = 'attention'
self_attn_type = AttentionType.PADDING
else:
mha_name = 'self_attention'
self_attn_type = AttentionType.CAUSAL
assert self_attn_type is not None

# [batch, length, emb_dim] -> [batch, length, emb_dim]
x, residual = MultiHeadAttention(
Expand All@@ -1002,7 +1061,7 @@ def __call__(self,
zero_centered_gamma=self.zero_centered_gamma,
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self_attn_type,
attn_mask_type=self.self_attn_mask_type,
fuse_qkv=self.fuse_qkv_params,
kernel_init=self.mha_kernel_init,
use_bias=self.use_bias,
Expand DownExpand Up@@ -1049,7 +1108,7 @@ def hidden_dropout(x, deterministic):
apply_residual_connection_post_layernorm=self.
apply_residual_connection_post_layernorm,
output_layernorm=False, # Must do LayerNorm before MHA.
attn_type=AttentionType.PADDING,
attn_mask_type='padding',
float32_logits=self.float32_attention_logits,
scale_attn_logits=self.scale_attn_logits,
scaled_query_init=self.scaled_query_init,
Expand Down
12 changes: 8 additions & 4 deletions transformer_engine/jax/praxis/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -5,14 +5,14 @@
Praxis Modules related Transformer
"""
from functools import partial
from typing import Optional, Sequence, Tuple
from typing import Any, Optional, Sequence, Tuple

from praxis import pax_fiddle
from praxis.base_layer import WeightInit
from praxis.pytypes import JTensor

from .module import TransformerEngineBaseLayer
from ..flax.transformer import AttentionType, TransformerLayerType
from ..flax.transformer import TransformerLayerType
from ..flax.transformer import MultiHeadAttention as flax_MultiHeadAttention
from ..flax.transformer import RelativePositionBiases as flax_RelativePositionBiases
from ..flax.transformer import TransformerLayer as flax_TransformerLayer
Expand DownExpand Up@@ -73,7 +73,9 @@ class MultiHeadAttention(TransformerEngineBaseLayer):
bias_init: WeightInit = WeightInit.Constant(0.0)
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
scale_attn_logits: bool = False
Expand All@@ -99,7 +101,7 @@ def setup(self) -> None:
bias_init=TransformerEngineBaseLayer.generate_params_init("bias", self.bias_init),
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self.attn_type,
attn_mask_type=self.attn_mask_type,
fuse_qkv=self.fuse_qkv,
transpose_batch_sequence=self.transpose_batch_sequence,
scale_attn_logits=self.scale_attn_logits,
Expand DownExpand Up@@ -145,6 +147,7 @@ class TransformerLayer(TransformerEngineBaseLayer):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: pax_fiddle.Config[RelativePositionBiases] = pax_fiddle.template_field(None)
drop_path: float = 0.0
Expand DownExpand Up@@ -201,6 +204,7 @@ def setup(self) -> None:
output_layernorm=self.output_layernorm,
float32_attention_logits=self.float32_attention_logits,
layer_type=self.layer_type,
self_attn_mask_type=self.self_attn_mask_type,
enable_relative_embedding=self.enable_relative_embedding,
relative_embedding=relative_embedding_flax_module,
drop_path=self.drop_path,
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
[JAX] Add self_attn_mask_type and replace attn_type by zlsh80826 · Pull Request #273 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 6 additions & 7 deletions tests/jax/test_praxis_layers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from transformer_engine.jax.flax import RelativePositionBiases as flax_RelativePositionBiases
from transformer_engine.jax.flax import TransformerLayer as flax_TransformerLayer
from transformer_engine.jax.flax.module import Softmax
from transformer_engine.jax.flax.transformer import AttentionType
from transformer_engine.jax.fp8 import FP8Helper, is_fp8_available
from transformer_engine.jax.praxis import LayerNorm
from transformer_engine.jax.praxis import FusedSoftmax, LayerNorm
Expand DownExpand Up@@ -666,32 +665,32 @@ class MultiHeadAttnAttr:
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}]


Expand Down
105 changes: 82 additions & 23 deletions transformer_engine/jax/flax/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -197,17 +197,16 @@ def core_attention(query: Array,
dynamic_vector_slice_in_dim = vmap(lax.dynamic_slice_in_dim, in_axes=(None, 0, None, None))


class AttentionType(Enum):
"""TransformerLayerType."""
PADDING = AttnMaskType.PADDING_MASK
CAUSAL = AttnMaskType.CAUSAL_MASK


class MultiHeadAttention(nn.Module):
r"""
Multi-head Attention (MHA), including Query,
Key, Value and Output projection.

.. warning::

Argument :attr:`attn_type` is deprecated and superseded by :attr:`attn_mask_type`.
:attr:`attn_type` is ignored in version 0.10 and will be fully removed in version 0.11.

Parameters
----------
head_dim : int
Expand DownExpand Up@@ -245,8 +244,11 @@ class MultiHeadAttention(nn.Module):
Indicate if apply residual connection with the output of layer normalization.
output_layernorm : bool, default = False
Indicate if apply a layer normalization at the end of MHA.
attn_type: AttentionType, defult = AttentionType.PADDING
Indicate the format of the attention mask in the core attention.
attn_type: Any, defult = None
*Deprecated*, will be ignored in v0.10 and be fully removed in v0.11.
Please use `attn_mask_type` to config the attention mask.
attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.

Optimization parameters
-----------------------
Expand DownExpand Up@@ -282,7 +284,9 @@ class MultiHeadAttention(nn.Module):
bias_init: Initializer = nn.initializers.zeros
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
dtype: DType = jnp.float32
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
Expand All@@ -293,6 +297,14 @@ class MultiHeadAttention(nn.Module):
def __post_init__(self):
if self.kernel_init is None:
self.kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in', 'normal')
# TODO(rewang): remove attn_type after v0.11
if self.attn_type is not None:
warnings.warn(
"The 'attn_type' argument in the 'MultiHeadAttention' is"
" deprecated in version 0.10 and will be removed in version 0.11."
" Passing value in attn_type will be ignored, please use `attn_mask_type`"
" to config the attention mask type.",
category=DeprecationWarning)
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -570,9 +582,23 @@ def kv_init(key, shape, dtype):
if use_fused_attn:
assert mask is not None and mask.ndim == 4 # (b, 1, s_q, s_kv)
assert not self.transpose_batch_sequence

# TODO(rewang): make it configurable for pre_scale_bias
attn_bias_type = AttnBiasType.NO_BIAS if bias is None else AttnBiasType.POST_SCALE_BIAS

def canonicalize_attn_mask_type(attn_mask_type):
"""
Convert the string to AttnMaskType
"""
if attn_mask_type == 'causal':
return AttnMaskType.CAUSAL_MASK
if attn_mask_type == 'padding':
return AttnMaskType.PADDING_MASK
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

attn_mask_type = canonicalize_attn_mask_type(self.attn_mask_type)

if inputs_q is inputs_kv:
qkv_proj = qkv_proj.reshape((*qkv_proj.shape[:-1], self.num_heads, self.head_dim))
qkv_sharding_constraint = ('batch', 'length', 'qkv_dim', 'heads', 'kv')
Expand All@@ -583,7 +609,7 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
Expand All@@ -602,18 +628,27 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
sharding_type=first_sharding_type)
else:
softmax_type = SoftmaxType.SCALED
if self.attn_type is AttentionType.PADDING:
if mask is not None:
softmax_type = SoftmaxType.SCALED_MASKED
else:
softmax_type = SoftmaxType.SCALED_UPPER_TRIANG_MASKED

def convert_to_softmax_type(attn_mask_type, mask):
"""
Convert the string to SoftmaxType
"""
if attn_mask_type == 'causal':
return SoftmaxType.SCALED_UPPER_TRIANG_MASKED
if attn_mask_type == 'padding':
if mask is not None:
return SoftmaxType.SCALED_MASKED
return SoftmaxType.SCALED
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

softmax_type = convert_to_softmax_type(self.attn_mask_type, mask)

x = core_attention(query,
key,
Expand DownExpand Up@@ -765,6 +800,18 @@ class TransformerLayer(nn.Module):
an attention block and a feedforward network (MLP).
This standard layer is based on the paper “Attention Is All You Need”.

.. warning::

Argument :attr:`self_attn_mask_type` is introduced in version 0.10.
Starting from version 0.11, the default value will be `"causal"`.
However, to ensure compatibility with earlier versions, before 0.11,
the default value will be `"padding"` for the encoder and `"causal"` for the decoder.

.. note::

Argument :attr:`attention_mask` will be ignored when
:attr:`self_attn_mask_type` is set to `"causal"`.

Parameters
----------
hidden_size: int, default = 512
Expand DownExpand Up@@ -825,6 +872,8 @@ class TransformerLayer(nn.Module):
If set to TransformerLayerType.DECODER, an additional cross-attention block
is added after self-attention.this can be used for structures like `T5`
Transformer in conjunction with the TransformerLayerType.ENCODER option.
self_attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.
enable_relative_embedding: bool, default = True
Whether to enable relative embedding as shifting of attention logits.
relative_embedding: flax.linen.Module, default = None
Expand DownExpand Up@@ -878,6 +927,7 @@ class TransformerLayer(nn.Module):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: nn.Module = None
dtype: DType = jnp.float32
Expand All@@ -893,6 +943,19 @@ def __post_init__(self):
if self.mlp_kernel_init is None:
self.mlp_kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in',
'truncated_normal')
# TODO(rewang): default to 'causal' in 0.11 (also updated the doc after 0.11)
if self.self_attn_mask_type is None:
warnings.warn(
"The 'self_attn_mask_type' argument in the 'TransformerLayer' is"
" introduced in version 0.10. Starting from version 0.11, the default"
" value will be 'causal'. However, to ensure compatibility with earlier"
" versions, before 0.11, the default value will be 'padding' for the"
" encoder and 'causal' for the decoder.",
category=FutureWarning)
if self.layer_type == TransformerLayerType.ENCODER:
self.self_attn_mask_type = 'padding'
else:
self.self_attn_mask_type = 'causal'
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -975,16 +1038,12 @@ def __call__(self,

assert inputs.ndim == 3

self_attn_type = None
# Make name be the exactly same as T5X, since names would affect
# RNGKey during init and apply. Myabe no need in the feature.
if self.layer_type == TransformerLayerType.ENCODER:
mha_name = 'attention'
self_attn_type = AttentionType.PADDING
else:
mha_name = 'self_attention'
self_attn_type = AttentionType.CAUSAL
assert self_attn_type is not None

# [batch, length, emb_dim] -> [batch, length, emb_dim]
x, residual = MultiHeadAttention(
Expand All@@ -1002,7 +1061,7 @@ def __call__(self,
zero_centered_gamma=self.zero_centered_gamma,
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self_attn_type,
attn_mask_type=self.self_attn_mask_type,
fuse_qkv=self.fuse_qkv_params,
kernel_init=self.mha_kernel_init,
use_bias=self.use_bias,
Expand DownExpand Up@@ -1049,7 +1108,7 @@ def hidden_dropout(x, deterministic):
apply_residual_connection_post_layernorm=self.
apply_residual_connection_post_layernorm,
output_layernorm=False, # Must do LayerNorm before MHA.
attn_type=AttentionType.PADDING,
attn_mask_type='padding',
float32_logits=self.float32_attention_logits,
scale_attn_logits=self.scale_attn_logits,
scaled_query_init=self.scaled_query_init,
Expand Down
12 changes: 8 additions & 4 deletions transformer_engine/jax/praxis/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -5,14 +5,14 @@
Praxis Modules related Transformer
"""
from functools import partial
from typing import Optional, Sequence, Tuple
from typing import Any, Optional, Sequence, Tuple

from praxis import pax_fiddle
from praxis.base_layer import WeightInit
from praxis.pytypes import JTensor

from .module import TransformerEngineBaseLayer
from ..flax.transformer import AttentionType, TransformerLayerType
from ..flax.transformer import TransformerLayerType
from ..flax.transformer import MultiHeadAttention as flax_MultiHeadAttention
from ..flax.transformer import RelativePositionBiases as flax_RelativePositionBiases
from ..flax.transformer import TransformerLayer as flax_TransformerLayer
Expand DownExpand Up@@ -73,7 +73,9 @@ class MultiHeadAttention(TransformerEngineBaseLayer):
bias_init: WeightInit = WeightInit.Constant(0.0)
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
scale_attn_logits: bool = False
Expand All@@ -99,7 +101,7 @@ def setup(self) -> None:
bias_init=TransformerEngineBaseLayer.generate_params_init("bias", self.bias_init),
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self.attn_type,
attn_mask_type=self.attn_mask_type,
fuse_qkv=self.fuse_qkv,
transpose_batch_sequence=self.transpose_batch_sequence,
scale_attn_logits=self.scale_attn_logits,
Expand DownExpand Up@@ -145,6 +147,7 @@ class TransformerLayer(TransformerEngineBaseLayer):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: pax_fiddle.Config[RelativePositionBiases] = pax_fiddle.template_field(None)
drop_path: float = 0.0
Expand DownExpand Up@@ -201,6 +204,7 @@ def setup(self) -> None:
output_layernorm=self.output_layernorm,
float32_attention_logits=self.float32_attention_logits,
layer_type=self.layer_type,
self_attn_mask_type=self.self_attn_mask_type,
enable_relative_embedding=self.enable_relative_embedding,
relative_embedding=relative_embedding_flax_module,
drop_path=self.drop_path,
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' [JAX] Add self_attn_mask_type and replace attn_type by zlsh80826 · Pull Request #273 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 6 additions & 7 deletions tests/jax/test_praxis_layers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from transformer_engine.jax.flax import RelativePositionBiases as flax_RelativePositionBiases
from transformer_engine.jax.flax import TransformerLayer as flax_TransformerLayer
from transformer_engine.jax.flax.module import Softmax
from transformer_engine.jax.flax.transformer import AttentionType
from transformer_engine.jax.fp8 import FP8Helper, is_fp8_available
from transformer_engine.jax.praxis import LayerNorm
from transformer_engine.jax.praxis import FusedSoftmax, LayerNorm
Expand DownExpand Up@@ -666,32 +665,32 @@ class MultiHeadAttnAttr:
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}]


Expand Down
105 changes: 82 additions & 23 deletions transformer_engine/jax/flax/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -197,17 +197,16 @@ def core_attention(query: Array,
dynamic_vector_slice_in_dim = vmap(lax.dynamic_slice_in_dim, in_axes=(None, 0, None, None))


class AttentionType(Enum):
"""TransformerLayerType."""
PADDING = AttnMaskType.PADDING_MASK
CAUSAL = AttnMaskType.CAUSAL_MASK


class MultiHeadAttention(nn.Module):
r"""
Multi-head Attention (MHA), including Query,
Key, Value and Output projection.

.. warning::

Argument :attr:`attn_type` is deprecated and superseded by :attr:`attn_mask_type`.
:attr:`attn_type` is ignored in version 0.10 and will be fully removed in version 0.11.

Parameters
----------
head_dim : int
Expand DownExpand Up@@ -245,8 +244,11 @@ class MultiHeadAttention(nn.Module):
Indicate if apply residual connection with the output of layer normalization.
output_layernorm : bool, default = False
Indicate if apply a layer normalization at the end of MHA.
attn_type: AttentionType, defult = AttentionType.PADDING
Indicate the format of the attention mask in the core attention.
attn_type: Any, defult = None
*Deprecated*, will be ignored in v0.10 and be fully removed in v0.11.
Please use `attn_mask_type` to config the attention mask.
attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.

Optimization parameters
-----------------------
Expand DownExpand Up@@ -282,7 +284,9 @@ class MultiHeadAttention(nn.Module):
bias_init: Initializer = nn.initializers.zeros
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
dtype: DType = jnp.float32
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
Expand All@@ -293,6 +297,14 @@ class MultiHeadAttention(nn.Module):
def __post_init__(self):
if self.kernel_init is None:
self.kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in', 'normal')
# TODO(rewang): remove attn_type after v0.11
if self.attn_type is not None:
warnings.warn(
"The 'attn_type' argument in the 'MultiHeadAttention' is"
" deprecated in version 0.10 and will be removed in version 0.11."
" Passing value in attn_type will be ignored, please use `attn_mask_type`"
" to config the attention mask type.",
category=DeprecationWarning)
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -570,9 +582,23 @@ def kv_init(key, shape, dtype):
if use_fused_attn:
assert mask is not None and mask.ndim == 4 # (b, 1, s_q, s_kv)
assert not self.transpose_batch_sequence

# TODO(rewang): make it configurable for pre_scale_bias
attn_bias_type = AttnBiasType.NO_BIAS if bias is None else AttnBiasType.POST_SCALE_BIAS

def canonicalize_attn_mask_type(attn_mask_type):
"""
Convert the string to AttnMaskType
"""
if attn_mask_type == 'causal':
return AttnMaskType.CAUSAL_MASK
if attn_mask_type == 'padding':
return AttnMaskType.PADDING_MASK
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

attn_mask_type = canonicalize_attn_mask_type(self.attn_mask_type)

if inputs_q is inputs_kv:
qkv_proj = qkv_proj.reshape((*qkv_proj.shape[:-1], self.num_heads, self.head_dim))
qkv_sharding_constraint = ('batch', 'length', 'qkv_dim', 'heads', 'kv')
Expand All@@ -583,7 +609,7 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
Expand All@@ -602,18 +628,27 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
sharding_type=first_sharding_type)
else:
softmax_type = SoftmaxType.SCALED
if self.attn_type is AttentionType.PADDING:
if mask is not None:
softmax_type = SoftmaxType.SCALED_MASKED
else:
softmax_type = SoftmaxType.SCALED_UPPER_TRIANG_MASKED

def convert_to_softmax_type(attn_mask_type, mask):
"""
Convert the string to SoftmaxType
"""
if attn_mask_type == 'causal':
return SoftmaxType.SCALED_UPPER_TRIANG_MASKED
if attn_mask_type == 'padding':
if mask is not None:
return SoftmaxType.SCALED_MASKED
return SoftmaxType.SCALED
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

softmax_type = convert_to_softmax_type(self.attn_mask_type, mask)

x = core_attention(query,
key,
Expand DownExpand Up@@ -765,6 +800,18 @@ class TransformerLayer(nn.Module):
an attention block and a feedforward network (MLP).
This standard layer is based on the paper “Attention Is All You Need”.

.. warning::

Argument :attr:`self_attn_mask_type` is introduced in version 0.10.
Starting from version 0.11, the default value will be `"causal"`.
However, to ensure compatibility with earlier versions, before 0.11,
the default value will be `"padding"` for the encoder and `"causal"` for the decoder.

.. note::

Argument :attr:`attention_mask` will be ignored when
:attr:`self_attn_mask_type` is set to `"causal"`.

Parameters
----------
hidden_size: int, default = 512
Expand DownExpand Up@@ -825,6 +872,8 @@ class TransformerLayer(nn.Module):
If set to TransformerLayerType.DECODER, an additional cross-attention block
is added after self-attention.this can be used for structures like `T5`
Transformer in conjunction with the TransformerLayerType.ENCODER option.
self_attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.
enable_relative_embedding: bool, default = True
Whether to enable relative embedding as shifting of attention logits.
relative_embedding: flax.linen.Module, default = None
Expand DownExpand Up@@ -878,6 +927,7 @@ class TransformerLayer(nn.Module):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: nn.Module = None
dtype: DType = jnp.float32
Expand All@@ -893,6 +943,19 @@ def __post_init__(self):
if self.mlp_kernel_init is None:
self.mlp_kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in',
'truncated_normal')
# TODO(rewang): default to 'causal' in 0.11 (also updated the doc after 0.11)
if self.self_attn_mask_type is None:
warnings.warn(
"The 'self_attn_mask_type' argument in the 'TransformerLayer' is"
" introduced in version 0.10. Starting from version 0.11, the default"
" value will be 'causal'. However, to ensure compatibility with earlier"
" versions, before 0.11, the default value will be 'padding' for the"
" encoder and 'causal' for the decoder.",
category=FutureWarning)
if self.layer_type == TransformerLayerType.ENCODER:
self.self_attn_mask_type = 'padding'
else:
self.self_attn_mask_type = 'causal'
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -975,16 +1038,12 @@ def __call__(self,

assert inputs.ndim == 3

self_attn_type = None
# Make name be the exactly same as T5X, since names would affect
# RNGKey during init and apply. Myabe no need in the feature.
if self.layer_type == TransformerLayerType.ENCODER:
mha_name = 'attention'
self_attn_type = AttentionType.PADDING
else:
mha_name = 'self_attention'
self_attn_type = AttentionType.CAUSAL
assert self_attn_type is not None

# [batch, length, emb_dim] -> [batch, length, emb_dim]
x, residual = MultiHeadAttention(
Expand All@@ -1002,7 +1061,7 @@ def __call__(self,
zero_centered_gamma=self.zero_centered_gamma,
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self_attn_type,
attn_mask_type=self.self_attn_mask_type,
fuse_qkv=self.fuse_qkv_params,
kernel_init=self.mha_kernel_init,
use_bias=self.use_bias,
Expand DownExpand Up@@ -1049,7 +1108,7 @@ def hidden_dropout(x, deterministic):
apply_residual_connection_post_layernorm=self.
apply_residual_connection_post_layernorm,
output_layernorm=False, # Must do LayerNorm before MHA.
attn_type=AttentionType.PADDING,
attn_mask_type='padding',
float32_logits=self.float32_attention_logits,
scale_attn_logits=self.scale_attn_logits,
scaled_query_init=self.scaled_query_init,
Expand Down
12 changes: 8 additions & 4 deletions transformer_engine/jax/praxis/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -5,14 +5,14 @@
Praxis Modules related Transformer
"""
from functools import partial
from typing import Optional, Sequence, Tuple
from typing import Any, Optional, Sequence, Tuple

from praxis import pax_fiddle
from praxis.base_layer import WeightInit
from praxis.pytypes import JTensor

from .module import TransformerEngineBaseLayer
from ..flax.transformer import AttentionType, TransformerLayerType
from ..flax.transformer import TransformerLayerType
from ..flax.transformer import MultiHeadAttention as flax_MultiHeadAttention
from ..flax.transformer import RelativePositionBiases as flax_RelativePositionBiases
from ..flax.transformer import TransformerLayer as flax_TransformerLayer
Expand DownExpand Up@@ -73,7 +73,9 @@ class MultiHeadAttention(TransformerEngineBaseLayer):
bias_init: WeightInit = WeightInit.Constant(0.0)
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
scale_attn_logits: bool = False
Expand All@@ -99,7 +101,7 @@ def setup(self) -> None:
bias_init=TransformerEngineBaseLayer.generate_params_init("bias", self.bias_init),
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self.attn_type,
attn_mask_type=self.attn_mask_type,
fuse_qkv=self.fuse_qkv,
transpose_batch_sequence=self.transpose_batch_sequence,
scale_attn_logits=self.scale_attn_logits,
Expand DownExpand Up@@ -145,6 +147,7 @@ class TransformerLayer(TransformerEngineBaseLayer):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: pax_fiddle.Config[RelativePositionBiases] = pax_fiddle.template_field(None)
drop_path: float = 0.0
Expand DownExpand Up@@ -201,6 +204,7 @@ def setup(self) -> None:
output_layernorm=self.output_layernorm,
float32_attention_logits=self.float32_attention_logits,
layer_type=self.layer_type,
self_attn_mask_type=self.self_attn_mask_type,
enable_relative_embedding=self.enable_relative_embedding,
relative_embedding=relative_embedding_flax_module,
drop_path=self.drop_path,
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' [JAX] Add self_attn_mask_type and replace attn_type by zlsh80826 · Pull Request #273 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 6 additions & 7 deletions tests/jax/test_praxis_layers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from transformer_engine.jax.flax import RelativePositionBiases as flax_RelativePositionBiases
from transformer_engine.jax.flax import TransformerLayer as flax_TransformerLayer
from transformer_engine.jax.flax.module import Softmax
from transformer_engine.jax.flax.transformer import AttentionType
from transformer_engine.jax.fp8 import FP8Helper, is_fp8_available
from transformer_engine.jax.praxis import LayerNorm
from transformer_engine.jax.praxis import FusedSoftmax, LayerNorm
Expand DownExpand Up@@ -666,32 +665,32 @@ class MultiHeadAttnAttr:
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}]


Expand Down
105 changes: 82 additions & 23 deletions transformer_engine/jax/flax/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -197,17 +197,16 @@ def core_attention(query: Array,
dynamic_vector_slice_in_dim = vmap(lax.dynamic_slice_in_dim, in_axes=(None, 0, None, None))


class AttentionType(Enum):
"""TransformerLayerType."""
PADDING = AttnMaskType.PADDING_MASK
CAUSAL = AttnMaskType.CAUSAL_MASK


class MultiHeadAttention(nn.Module):
r"""
Multi-head Attention (MHA), including Query,
Key, Value and Output projection.

.. warning::

Argument :attr:`attn_type` is deprecated and superseded by :attr:`attn_mask_type`.
:attr:`attn_type` is ignored in version 0.10 and will be fully removed in version 0.11.

Parameters
----------
head_dim : int
Expand DownExpand Up@@ -245,8 +244,11 @@ class MultiHeadAttention(nn.Module):
Indicate if apply residual connection with the output of layer normalization.
output_layernorm : bool, default = False
Indicate if apply a layer normalization at the end of MHA.
attn_type: AttentionType, defult = AttentionType.PADDING
Indicate the format of the attention mask in the core attention.
attn_type: Any, defult = None
*Deprecated*, will be ignored in v0.10 and be fully removed in v0.11.
Please use `attn_mask_type` to config the attention mask.
attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.

Optimization parameters
-----------------------
Expand DownExpand Up@@ -282,7 +284,9 @@ class MultiHeadAttention(nn.Module):
bias_init: Initializer = nn.initializers.zeros
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
dtype: DType = jnp.float32
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
Expand All@@ -293,6 +297,14 @@ class MultiHeadAttention(nn.Module):
def __post_init__(self):
if self.kernel_init is None:
self.kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in', 'normal')
# TODO(rewang): remove attn_type after v0.11
if self.attn_type is not None:
warnings.warn(
"The 'attn_type' argument in the 'MultiHeadAttention' is"
" deprecated in version 0.10 and will be removed in version 0.11."
" Passing value in attn_type will be ignored, please use `attn_mask_type`"
" to config the attention mask type.",
category=DeprecationWarning)
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -570,9 +582,23 @@ def kv_init(key, shape, dtype):
if use_fused_attn:
assert mask is not None and mask.ndim == 4 # (b, 1, s_q, s_kv)
assert not self.transpose_batch_sequence

# TODO(rewang): make it configurable for pre_scale_bias
attn_bias_type = AttnBiasType.NO_BIAS if bias is None else AttnBiasType.POST_SCALE_BIAS

def canonicalize_attn_mask_type(attn_mask_type):
"""
Convert the string to AttnMaskType
"""
if attn_mask_type == 'causal':
return AttnMaskType.CAUSAL_MASK
if attn_mask_type == 'padding':
return AttnMaskType.PADDING_MASK
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

attn_mask_type = canonicalize_attn_mask_type(self.attn_mask_type)

if inputs_q is inputs_kv:
qkv_proj = qkv_proj.reshape((*qkv_proj.shape[:-1], self.num_heads, self.head_dim))
qkv_sharding_constraint = ('batch', 'length', 'qkv_dim', 'heads', 'kv')
Expand All@@ -583,7 +609,7 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
Expand All@@ -602,18 +628,27 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
sharding_type=first_sharding_type)
else:
softmax_type = SoftmaxType.SCALED
if self.attn_type is AttentionType.PADDING:
if mask is not None:
softmax_type = SoftmaxType.SCALED_MASKED
else:
softmax_type = SoftmaxType.SCALED_UPPER_TRIANG_MASKED

def convert_to_softmax_type(attn_mask_type, mask):
"""
Convert the string to SoftmaxType
"""
if attn_mask_type == 'causal':
return SoftmaxType.SCALED_UPPER_TRIANG_MASKED
if attn_mask_type == 'padding':
if mask is not None:
return SoftmaxType.SCALED_MASKED
return SoftmaxType.SCALED
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

softmax_type = convert_to_softmax_type(self.attn_mask_type, mask)

x = core_attention(query,
key,
Expand DownExpand Up@@ -765,6 +800,18 @@ class TransformerLayer(nn.Module):
an attention block and a feedforward network (MLP).
This standard layer is based on the paper “Attention Is All You Need”.

.. warning::

Argument :attr:`self_attn_mask_type` is introduced in version 0.10.
Starting from version 0.11, the default value will be `"causal"`.
However, to ensure compatibility with earlier versions, before 0.11,
the default value will be `"padding"` for the encoder and `"causal"` for the decoder.

.. note::

Argument :attr:`attention_mask` will be ignored when
:attr:`self_attn_mask_type` is set to `"causal"`.

Parameters
----------
hidden_size: int, default = 512
Expand DownExpand Up@@ -825,6 +872,8 @@ class TransformerLayer(nn.Module):
If set to TransformerLayerType.DECODER, an additional cross-attention block
is added after self-attention.this can be used for structures like `T5`
Transformer in conjunction with the TransformerLayerType.ENCODER option.
self_attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.
enable_relative_embedding: bool, default = True
Whether to enable relative embedding as shifting of attention logits.
relative_embedding: flax.linen.Module, default = None
Expand DownExpand Up@@ -878,6 +927,7 @@ class TransformerLayer(nn.Module):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: nn.Module = None
dtype: DType = jnp.float32
Expand All@@ -893,6 +943,19 @@ def __post_init__(self):
if self.mlp_kernel_init is None:
self.mlp_kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in',
'truncated_normal')
# TODO(rewang): default to 'causal' in 0.11 (also updated the doc after 0.11)
if self.self_attn_mask_type is None:
warnings.warn(
"The 'self_attn_mask_type' argument in the 'TransformerLayer' is"
" introduced in version 0.10. Starting from version 0.11, the default"
" value will be 'causal'. However, to ensure compatibility with earlier"
" versions, before 0.11, the default value will be 'padding' for the"
" encoder and 'causal' for the decoder.",
category=FutureWarning)
if self.layer_type == TransformerLayerType.ENCODER:
self.self_attn_mask_type = 'padding'
else:
self.self_attn_mask_type = 'causal'
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -975,16 +1038,12 @@ def __call__(self,

assert inputs.ndim == 3

self_attn_type = None
# Make name be the exactly same as T5X, since names would affect
# RNGKey during init and apply. Myabe no need in the feature.
if self.layer_type == TransformerLayerType.ENCODER:
mha_name = 'attention'
self_attn_type = AttentionType.PADDING
else:
mha_name = 'self_attention'
self_attn_type = AttentionType.CAUSAL
assert self_attn_type is not None

# [batch, length, emb_dim] -> [batch, length, emb_dim]
x, residual = MultiHeadAttention(
Expand All@@ -1002,7 +1061,7 @@ def __call__(self,
zero_centered_gamma=self.zero_centered_gamma,
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self_attn_type,
attn_mask_type=self.self_attn_mask_type,
fuse_qkv=self.fuse_qkv_params,
kernel_init=self.mha_kernel_init,
use_bias=self.use_bias,
Expand DownExpand Up@@ -1049,7 +1108,7 @@ def hidden_dropout(x, deterministic):
apply_residual_connection_post_layernorm=self.
apply_residual_connection_post_layernorm,
output_layernorm=False, # Must do LayerNorm before MHA.
attn_type=AttentionType.PADDING,
attn_mask_type='padding',
float32_logits=self.float32_attention_logits,
scale_attn_logits=self.scale_attn_logits,
scaled_query_init=self.scaled_query_init,
Expand Down
12 changes: 8 additions & 4 deletions transformer_engine/jax/praxis/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -5,14 +5,14 @@
Praxis Modules related Transformer
"""
from functools import partial
from typing import Optional, Sequence, Tuple
from typing import Any, Optional, Sequence, Tuple

from praxis import pax_fiddle
from praxis.base_layer import WeightInit
from praxis.pytypes import JTensor

from .module import TransformerEngineBaseLayer
from ..flax.transformer import AttentionType, TransformerLayerType
from ..flax.transformer import TransformerLayerType
from ..flax.transformer import MultiHeadAttention as flax_MultiHeadAttention
from ..flax.transformer import RelativePositionBiases as flax_RelativePositionBiases
from ..flax.transformer import TransformerLayer as flax_TransformerLayer
Expand DownExpand Up@@ -73,7 +73,9 @@ class MultiHeadAttention(TransformerEngineBaseLayer):
bias_init: WeightInit = WeightInit.Constant(0.0)
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
scale_attn_logits: bool = False
Expand All@@ -99,7 +101,7 @@ def setup(self) -> None:
bias_init=TransformerEngineBaseLayer.generate_params_init("bias", self.bias_init),
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self.attn_type,
attn_mask_type=self.attn_mask_type,
fuse_qkv=self.fuse_qkv,
transpose_batch_sequence=self.transpose_batch_sequence,
scale_attn_logits=self.scale_attn_logits,
Expand DownExpand Up@@ -145,6 +147,7 @@ class TransformerLayer(TransformerEngineBaseLayer):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: pax_fiddle.Config[RelativePositionBiases] = pax_fiddle.template_field(None)
drop_path: float = 0.0
Expand DownExpand Up@@ -201,6 +204,7 @@ def setup(self) -> None:
output_layernorm=self.output_layernorm,
float32_attention_logits=self.float32_attention_logits,
layer_type=self.layer_type,
self_attn_mask_type=self.self_attn_mask_type,
enable_relative_embedding=self.enable_relative_embedding,
relative_embedding=relative_embedding_flax_module,
drop_path=self.drop_path,
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + ' [JAX] Add self_attn_mask_type and replace attn_type by zlsh80826 · Pull Request #273 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 6 additions & 7 deletions tests/jax/test_praxis_layers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from transformer_engine.jax.flax import RelativePositionBiases as flax_RelativePositionBiases
from transformer_engine.jax.flax import TransformerLayer as flax_TransformerLayer
from transformer_engine.jax.flax.module import Softmax
from transformer_engine.jax.flax.transformer import AttentionType
from transformer_engine.jax.fp8 import FP8Helper, is_fp8_available
from transformer_engine.jax.praxis import LayerNorm
from transformer_engine.jax.praxis import FusedSoftmax, LayerNorm
Expand DownExpand Up@@ -666,32 +665,32 @@ class MultiHeadAttnAttr:
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}]


Expand Down
105 changes: 82 additions & 23 deletions transformer_engine/jax/flax/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -197,17 +197,16 @@ def core_attention(query: Array,
dynamic_vector_slice_in_dim = vmap(lax.dynamic_slice_in_dim, in_axes=(None, 0, None, None))


class AttentionType(Enum):
"""TransformerLayerType."""
PADDING = AttnMaskType.PADDING_MASK
CAUSAL = AttnMaskType.CAUSAL_MASK


class MultiHeadAttention(nn.Module):
r"""
Multi-head Attention (MHA), including Query,
Key, Value and Output projection.

.. warning::

Argument :attr:`attn_type` is deprecated and superseded by :attr:`attn_mask_type`.
:attr:`attn_type` is ignored in version 0.10 and will be fully removed in version 0.11.

Parameters
----------
head_dim : int
Expand DownExpand Up@@ -245,8 +244,11 @@ class MultiHeadAttention(nn.Module):
Indicate if apply residual connection with the output of layer normalization.
output_layernorm : bool, default = False
Indicate if apply a layer normalization at the end of MHA.
attn_type: AttentionType, defult = AttentionType.PADDING
Indicate the format of the attention mask in the core attention.
attn_type: Any, defult = None
*Deprecated*, will be ignored in v0.10 and be fully removed in v0.11.
Please use `attn_mask_type` to config the attention mask.
attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.

Optimization parameters
-----------------------
Expand DownExpand Up@@ -282,7 +284,9 @@ class MultiHeadAttention(nn.Module):
bias_init: Initializer = nn.initializers.zeros
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
dtype: DType = jnp.float32
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
Expand All@@ -293,6 +297,14 @@ class MultiHeadAttention(nn.Module):
def __post_init__(self):
if self.kernel_init is None:
self.kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in', 'normal')
# TODO(rewang): remove attn_type after v0.11
if self.attn_type is not None:
warnings.warn(
"The 'attn_type' argument in the 'MultiHeadAttention' is"
" deprecated in version 0.10 and will be removed in version 0.11."
" Passing value in attn_type will be ignored, please use `attn_mask_type`"
" to config the attention mask type.",
category=DeprecationWarning)
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -570,9 +582,23 @@ def kv_init(key, shape, dtype):
if use_fused_attn:
assert mask is not None and mask.ndim == 4 # (b, 1, s_q, s_kv)
assert not self.transpose_batch_sequence

# TODO(rewang): make it configurable for pre_scale_bias
attn_bias_type = AttnBiasType.NO_BIAS if bias is None else AttnBiasType.POST_SCALE_BIAS

def canonicalize_attn_mask_type(attn_mask_type):
"""
Convert the string to AttnMaskType
"""
if attn_mask_type == 'causal':
return AttnMaskType.CAUSAL_MASK
if attn_mask_type == 'padding':
return AttnMaskType.PADDING_MASK
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

attn_mask_type = canonicalize_attn_mask_type(self.attn_mask_type)

if inputs_q is inputs_kv:
qkv_proj = qkv_proj.reshape((*qkv_proj.shape[:-1], self.num_heads, self.head_dim))
qkv_sharding_constraint = ('batch', 'length', 'qkv_dim', 'heads', 'kv')
Expand All@@ -583,7 +609,7 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
Expand All@@ -602,18 +628,27 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
sharding_type=first_sharding_type)
else:
softmax_type = SoftmaxType.SCALED
if self.attn_type is AttentionType.PADDING:
if mask is not None:
softmax_type = SoftmaxType.SCALED_MASKED
else:
softmax_type = SoftmaxType.SCALED_UPPER_TRIANG_MASKED

def convert_to_softmax_type(attn_mask_type, mask):
"""
Convert the string to SoftmaxType
"""
if attn_mask_type == 'causal':
return SoftmaxType.SCALED_UPPER_TRIANG_MASKED
if attn_mask_type == 'padding':
if mask is not None:
return SoftmaxType.SCALED_MASKED
return SoftmaxType.SCALED
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

softmax_type = convert_to_softmax_type(self.attn_mask_type, mask)

x = core_attention(query,
key,
Expand DownExpand Up@@ -765,6 +800,18 @@ class TransformerLayer(nn.Module):
an attention block and a feedforward network (MLP).
This standard layer is based on the paper “Attention Is All You Need”.

.. warning::

Argument :attr:`self_attn_mask_type` is introduced in version 0.10.
Starting from version 0.11, the default value will be `"causal"`.
However, to ensure compatibility with earlier versions, before 0.11,
the default value will be `"padding"` for the encoder and `"causal"` for the decoder.

.. note::

Argument :attr:`attention_mask` will be ignored when
:attr:`self_attn_mask_type` is set to `"causal"`.

Parameters
----------
hidden_size: int, default = 512
Expand DownExpand Up@@ -825,6 +872,8 @@ class TransformerLayer(nn.Module):
If set to TransformerLayerType.DECODER, an additional cross-attention block
is added after self-attention.this can be used for structures like `T5`
Transformer in conjunction with the TransformerLayerType.ENCODER option.
self_attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.
enable_relative_embedding: bool, default = True
Whether to enable relative embedding as shifting of attention logits.
relative_embedding: flax.linen.Module, default = None
Expand DownExpand Up@@ -878,6 +927,7 @@ class TransformerLayer(nn.Module):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: nn.Module = None
dtype: DType = jnp.float32
Expand All@@ -893,6 +943,19 @@ def __post_init__(self):
if self.mlp_kernel_init is None:
self.mlp_kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in',
'truncated_normal')
# TODO(rewang): default to 'causal' in 0.11 (also updated the doc after 0.11)
if self.self_attn_mask_type is None:
warnings.warn(
"The 'self_attn_mask_type' argument in the 'TransformerLayer' is"
" introduced in version 0.10. Starting from version 0.11, the default"
" value will be 'causal'. However, to ensure compatibility with earlier"
" versions, before 0.11, the default value will be 'padding' for the"
" encoder and 'causal' for the decoder.",
category=FutureWarning)
if self.layer_type == TransformerLayerType.ENCODER:
self.self_attn_mask_type = 'padding'
else:
self.self_attn_mask_type = 'causal'
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -975,16 +1038,12 @@ def __call__(self,

assert inputs.ndim == 3

self_attn_type = None
# Make name be the exactly same as T5X, since names would affect
# RNGKey during init and apply. Myabe no need in the feature.
if self.layer_type == TransformerLayerType.ENCODER:
mha_name = 'attention'
self_attn_type = AttentionType.PADDING
else:
mha_name = 'self_attention'
self_attn_type = AttentionType.CAUSAL
assert self_attn_type is not None

# [batch, length, emb_dim] -> [batch, length, emb_dim]
x, residual = MultiHeadAttention(
Expand All@@ -1002,7 +1061,7 @@ def __call__(self,
zero_centered_gamma=self.zero_centered_gamma,
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self_attn_type,
attn_mask_type=self.self_attn_mask_type,
fuse_qkv=self.fuse_qkv_params,
kernel_init=self.mha_kernel_init,
use_bias=self.use_bias,
Expand DownExpand Up@@ -1049,7 +1108,7 @@ def hidden_dropout(x, deterministic):
apply_residual_connection_post_layernorm=self.
apply_residual_connection_post_layernorm,
output_layernorm=False, # Must do LayerNorm before MHA.
attn_type=AttentionType.PADDING,
attn_mask_type='padding',
float32_logits=self.float32_attention_logits,
scale_attn_logits=self.scale_attn_logits,
scaled_query_init=self.scaled_query_init,
Expand Down
12 changes: 8 additions & 4 deletions transformer_engine/jax/praxis/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -5,14 +5,14 @@
Praxis Modules related Transformer
"""
from functools import partial
from typing import Optional, Sequence, Tuple
from typing import Any, Optional, Sequence, Tuple

from praxis import pax_fiddle
from praxis.base_layer import WeightInit
from praxis.pytypes import JTensor

from .module import TransformerEngineBaseLayer
from ..flax.transformer import AttentionType, TransformerLayerType
from ..flax.transformer import TransformerLayerType
from ..flax.transformer import MultiHeadAttention as flax_MultiHeadAttention
from ..flax.transformer import RelativePositionBiases as flax_RelativePositionBiases
from ..flax.transformer import TransformerLayer as flax_TransformerLayer
Expand DownExpand Up@@ -73,7 +73,9 @@ class MultiHeadAttention(TransformerEngineBaseLayer):
bias_init: WeightInit = WeightInit.Constant(0.0)
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
scale_attn_logits: bool = False
Expand All@@ -99,7 +101,7 @@ def setup(self) -> None:
bias_init=TransformerEngineBaseLayer.generate_params_init("bias", self.bias_init),
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self.attn_type,
attn_mask_type=self.attn_mask_type,
fuse_qkv=self.fuse_qkv,
transpose_batch_sequence=self.transpose_batch_sequence,
scale_attn_logits=self.scale_attn_logits,
Expand DownExpand Up@@ -145,6 +147,7 @@ class TransformerLayer(TransformerEngineBaseLayer):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: pax_fiddle.Config[RelativePositionBiases] = pax_fiddle.template_field(None)
drop_path: float = 0.0
Expand DownExpand Up@@ -201,6 +204,7 @@ def setup(self) -> None:
output_layernorm=self.output_layernorm,
float32_attention_logits=self.float32_attention_logits,
layer_type=self.layer_type,
self_attn_mask_type=self.self_attn_mask_type,
enable_relative_embedding=self.enable_relative_embedding,
relative_embedding=relative_embedding_flax_module,
drop_path=self.drop_path,
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' [JAX] Add self_attn_mask_type and replace attn_type by zlsh80826 · Pull Request #273 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 6 additions & 7 deletions tests/jax/test_praxis_layers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from transformer_engine.jax.flax import RelativePositionBiases as flax_RelativePositionBiases
from transformer_engine.jax.flax import TransformerLayer as flax_TransformerLayer
from transformer_engine.jax.flax.module import Softmax
from transformer_engine.jax.flax.transformer import AttentionType
from transformer_engine.jax.fp8 import FP8Helper, is_fp8_available
from transformer_engine.jax.praxis import LayerNorm
from transformer_engine.jax.praxis import FusedSoftmax, LayerNorm
Expand DownExpand Up@@ -666,32 +665,32 @@ class MultiHeadAttnAttr:
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}]


Expand Down
105 changes: 82 additions & 23 deletions transformer_engine/jax/flax/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -197,17 +197,16 @@ def core_attention(query: Array,
dynamic_vector_slice_in_dim = vmap(lax.dynamic_slice_in_dim, in_axes=(None, 0, None, None))


class AttentionType(Enum):
"""TransformerLayerType."""
PADDING = AttnMaskType.PADDING_MASK
CAUSAL = AttnMaskType.CAUSAL_MASK


class MultiHeadAttention(nn.Module):
r"""
Multi-head Attention (MHA), including Query,
Key, Value and Output projection.

.. warning::

Argument :attr:`attn_type` is deprecated and superseded by :attr:`attn_mask_type`.
:attr:`attn_type` is ignored in version 0.10 and will be fully removed in version 0.11.

Parameters
----------
head_dim : int
Expand DownExpand Up@@ -245,8 +244,11 @@ class MultiHeadAttention(nn.Module):
Indicate if apply residual connection with the output of layer normalization.
output_layernorm : bool, default = False
Indicate if apply a layer normalization at the end of MHA.
attn_type: AttentionType, defult = AttentionType.PADDING
Indicate the format of the attention mask in the core attention.
attn_type: Any, defult = None
*Deprecated*, will be ignored in v0.10 and be fully removed in v0.11.
Please use `attn_mask_type` to config the attention mask.
attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.

Optimization parameters
-----------------------
Expand DownExpand Up@@ -282,7 +284,9 @@ class MultiHeadAttention(nn.Module):
bias_init: Initializer = nn.initializers.zeros
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
dtype: DType = jnp.float32
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
Expand All@@ -293,6 +297,14 @@ class MultiHeadAttention(nn.Module):
def __post_init__(self):
if self.kernel_init is None:
self.kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in', 'normal')
# TODO(rewang): remove attn_type after v0.11
if self.attn_type is not None:
warnings.warn(
"The 'attn_type' argument in the 'MultiHeadAttention' is"
" deprecated in version 0.10 and will be removed in version 0.11."
" Passing value in attn_type will be ignored, please use `attn_mask_type`"
" to config the attention mask type.",
category=DeprecationWarning)
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -570,9 +582,23 @@ def kv_init(key, shape, dtype):
if use_fused_attn:
assert mask is not None and mask.ndim == 4 # (b, 1, s_q, s_kv)
assert not self.transpose_batch_sequence

# TODO(rewang): make it configurable for pre_scale_bias
attn_bias_type = AttnBiasType.NO_BIAS if bias is None else AttnBiasType.POST_SCALE_BIAS

def canonicalize_attn_mask_type(attn_mask_type):
"""
Convert the string to AttnMaskType
"""
if attn_mask_type == 'causal':
return AttnMaskType.CAUSAL_MASK
if attn_mask_type == 'padding':
return AttnMaskType.PADDING_MASK
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

attn_mask_type = canonicalize_attn_mask_type(self.attn_mask_type)

if inputs_q is inputs_kv:
qkv_proj = qkv_proj.reshape((*qkv_proj.shape[:-1], self.num_heads, self.head_dim))
qkv_sharding_constraint = ('batch', 'length', 'qkv_dim', 'heads', 'kv')
Expand All@@ -583,7 +609,7 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
Expand All@@ -602,18 +628,27 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
sharding_type=first_sharding_type)
else:
softmax_type = SoftmaxType.SCALED
if self.attn_type is AttentionType.PADDING:
if mask is not None:
softmax_type = SoftmaxType.SCALED_MASKED
else:
softmax_type = SoftmaxType.SCALED_UPPER_TRIANG_MASKED

def convert_to_softmax_type(attn_mask_type, mask):
"""
Convert the string to SoftmaxType
"""
if attn_mask_type == 'causal':
return SoftmaxType.SCALED_UPPER_TRIANG_MASKED
if attn_mask_type == 'padding':
if mask is not None:
return SoftmaxType.SCALED_MASKED
return SoftmaxType.SCALED
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

softmax_type = convert_to_softmax_type(self.attn_mask_type, mask)

x = core_attention(query,
key,
Expand DownExpand Up@@ -765,6 +800,18 @@ class TransformerLayer(nn.Module):
an attention block and a feedforward network (MLP).
This standard layer is based on the paper “Attention Is All You Need”.

.. warning::

Argument :attr:`self_attn_mask_type` is introduced in version 0.10.
Starting from version 0.11, the default value will be `"causal"`.
However, to ensure compatibility with earlier versions, before 0.11,
the default value will be `"padding"` for the encoder and `"causal"` for the decoder.

.. note::

Argument :attr:`attention_mask` will be ignored when
:attr:`self_attn_mask_type` is set to `"causal"`.

Parameters
----------
hidden_size: int, default = 512
Expand DownExpand Up@@ -825,6 +872,8 @@ class TransformerLayer(nn.Module):
If set to TransformerLayerType.DECODER, an additional cross-attention block
is added after self-attention.this can be used for structures like `T5`
Transformer in conjunction with the TransformerLayerType.ENCODER option.
self_attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.
enable_relative_embedding: bool, default = True
Whether to enable relative embedding as shifting of attention logits.
relative_embedding: flax.linen.Module, default = None
Expand DownExpand Up@@ -878,6 +927,7 @@ class TransformerLayer(nn.Module):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: nn.Module = None
dtype: DType = jnp.float32
Expand All@@ -893,6 +943,19 @@ def __post_init__(self):
if self.mlp_kernel_init is None:
self.mlp_kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in',
'truncated_normal')
# TODO(rewang): default to 'causal' in 0.11 (also updated the doc after 0.11)
if self.self_attn_mask_type is None:
warnings.warn(
"The 'self_attn_mask_type' argument in the 'TransformerLayer' is"
" introduced in version 0.10. Starting from version 0.11, the default"
" value will be 'causal'. However, to ensure compatibility with earlier"
" versions, before 0.11, the default value will be 'padding' for the"
" encoder and 'causal' for the decoder.",
category=FutureWarning)
if self.layer_type == TransformerLayerType.ENCODER:
self.self_attn_mask_type = 'padding'
else:
self.self_attn_mask_type = 'causal'
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -975,16 +1038,12 @@ def __call__(self,

assert inputs.ndim == 3

self_attn_type = None
# Make name be the exactly same as T5X, since names would affect
# RNGKey during init and apply. Myabe no need in the feature.
if self.layer_type == TransformerLayerType.ENCODER:
mha_name = 'attention'
self_attn_type = AttentionType.PADDING
else:
mha_name = 'self_attention'
self_attn_type = AttentionType.CAUSAL
assert self_attn_type is not None

# [batch, length, emb_dim] -> [batch, length, emb_dim]
x, residual = MultiHeadAttention(
Expand All@@ -1002,7 +1061,7 @@ def __call__(self,
zero_centered_gamma=self.zero_centered_gamma,
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self_attn_type,
attn_mask_type=self.self_attn_mask_type,
fuse_qkv=self.fuse_qkv_params,
kernel_init=self.mha_kernel_init,
use_bias=self.use_bias,
Expand DownExpand Up@@ -1049,7 +1108,7 @@ def hidden_dropout(x, deterministic):
apply_residual_connection_post_layernorm=self.
apply_residual_connection_post_layernorm,
output_layernorm=False, # Must do LayerNorm before MHA.
attn_type=AttentionType.PADDING,
attn_mask_type='padding',
float32_logits=self.float32_attention_logits,
scale_attn_logits=self.scale_attn_logits,
scaled_query_init=self.scaled_query_init,
Expand Down
12 changes: 8 additions & 4 deletions transformer_engine/jax/praxis/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -5,14 +5,14 @@
Praxis Modules related Transformer
"""
from functools import partial
from typing import Optional, Sequence, Tuple
from typing import Any, Optional, Sequence, Tuple

from praxis import pax_fiddle
from praxis.base_layer import WeightInit
from praxis.pytypes import JTensor

from .module import TransformerEngineBaseLayer
from ..flax.transformer import AttentionType, TransformerLayerType
from ..flax.transformer import TransformerLayerType
from ..flax.transformer import MultiHeadAttention as flax_MultiHeadAttention
from ..flax.transformer import RelativePositionBiases as flax_RelativePositionBiases
from ..flax.transformer import TransformerLayer as flax_TransformerLayer
Expand DownExpand Up@@ -73,7 +73,9 @@ class MultiHeadAttention(TransformerEngineBaseLayer):
bias_init: WeightInit = WeightInit.Constant(0.0)
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
scale_attn_logits: bool = False
Expand All@@ -99,7 +101,7 @@ def setup(self) -> None:
bias_init=TransformerEngineBaseLayer.generate_params_init("bias", self.bias_init),
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self.attn_type,
attn_mask_type=self.attn_mask_type,
fuse_qkv=self.fuse_qkv,
transpose_batch_sequence=self.transpose_batch_sequence,
scale_attn_logits=self.scale_attn_logits,
Expand DownExpand Up@@ -145,6 +147,7 @@ class TransformerLayer(TransformerEngineBaseLayer):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: pax_fiddle.Config[RelativePositionBiases] = pax_fiddle.template_field(None)
drop_path: float = 0.0
Expand DownExpand Up@@ -201,6 +204,7 @@ def setup(self) -> None:
output_layernorm=self.output_layernorm,
float32_attention_logits=self.float32_attention_logits,
layer_type=self.layer_type,
self_attn_mask_type=self.self_attn_mask_type,
enable_relative_embedding=self.enable_relative_embedding,
relative_embedding=relative_embedding_flax_module,
drop_path=self.drop_path,
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' [JAX] Add self_attn_mask_type and replace attn_type by zlsh80826 · Pull Request #273 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 6 additions & 7 deletions tests/jax/test_praxis_layers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from transformer_engine.jax.flax import RelativePositionBiases as flax_RelativePositionBiases
from transformer_engine.jax.flax import TransformerLayer as flax_TransformerLayer
from transformer_engine.jax.flax.module import Softmax
from transformer_engine.jax.flax.transformer import AttentionType
from transformer_engine.jax.fp8 import FP8Helper, is_fp8_available
from transformer_engine.jax.praxis import LayerNorm
from transformer_engine.jax.praxis import FusedSoftmax, LayerNorm
Expand DownExpand Up@@ -666,32 +665,32 @@ class MultiHeadAttnAttr:
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}]


Expand Down
105 changes: 82 additions & 23 deletions transformer_engine/jax/flax/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -197,17 +197,16 @@ def core_attention(query: Array,
dynamic_vector_slice_in_dim = vmap(lax.dynamic_slice_in_dim, in_axes=(None, 0, None, None))


class AttentionType(Enum):
"""TransformerLayerType."""
PADDING = AttnMaskType.PADDING_MASK
CAUSAL = AttnMaskType.CAUSAL_MASK


class MultiHeadAttention(nn.Module):
r"""
Multi-head Attention (MHA), including Query,
Key, Value and Output projection.

.. warning::

Argument :attr:`attn_type` is deprecated and superseded by :attr:`attn_mask_type`.
:attr:`attn_type` is ignored in version 0.10 and will be fully removed in version 0.11.

Parameters
----------
head_dim : int
Expand DownExpand Up@@ -245,8 +244,11 @@ class MultiHeadAttention(nn.Module):
Indicate if apply residual connection with the output of layer normalization.
output_layernorm : bool, default = False
Indicate if apply a layer normalization at the end of MHA.
attn_type: AttentionType, defult = AttentionType.PADDING
Indicate the format of the attention mask in the core attention.
attn_type: Any, defult = None
*Deprecated*, will be ignored in v0.10 and be fully removed in v0.11.
Please use `attn_mask_type` to config the attention mask.
attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.

Optimization parameters
-----------------------
Expand DownExpand Up@@ -282,7 +284,9 @@ class MultiHeadAttention(nn.Module):
bias_init: Initializer = nn.initializers.zeros
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
dtype: DType = jnp.float32
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
Expand All@@ -293,6 +297,14 @@ class MultiHeadAttention(nn.Module):
def __post_init__(self):
if self.kernel_init is None:
self.kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in', 'normal')
# TODO(rewang): remove attn_type after v0.11
if self.attn_type is not None:
warnings.warn(
"The 'attn_type' argument in the 'MultiHeadAttention' is"
" deprecated in version 0.10 and will be removed in version 0.11."
" Passing value in attn_type will be ignored, please use `attn_mask_type`"
" to config the attention mask type.",
category=DeprecationWarning)
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -570,9 +582,23 @@ def kv_init(key, shape, dtype):
if use_fused_attn:
assert mask is not None and mask.ndim == 4 # (b, 1, s_q, s_kv)
assert not self.transpose_batch_sequence

# TODO(rewang): make it configurable for pre_scale_bias
attn_bias_type = AttnBiasType.NO_BIAS if bias is None else AttnBiasType.POST_SCALE_BIAS

def canonicalize_attn_mask_type(attn_mask_type):
"""
Convert the string to AttnMaskType
"""
if attn_mask_type == 'causal':
return AttnMaskType.CAUSAL_MASK
if attn_mask_type == 'padding':
return AttnMaskType.PADDING_MASK
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

attn_mask_type = canonicalize_attn_mask_type(self.attn_mask_type)

if inputs_q is inputs_kv:
qkv_proj = qkv_proj.reshape((*qkv_proj.shape[:-1], self.num_heads, self.head_dim))
qkv_sharding_constraint = ('batch', 'length', 'qkv_dim', 'heads', 'kv')
Expand All@@ -583,7 +609,7 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
Expand All@@ -602,18 +628,27 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
sharding_type=first_sharding_type)
else:
softmax_type = SoftmaxType.SCALED
if self.attn_type is AttentionType.PADDING:
if mask is not None:
softmax_type = SoftmaxType.SCALED_MASKED
else:
softmax_type = SoftmaxType.SCALED_UPPER_TRIANG_MASKED

def convert_to_softmax_type(attn_mask_type, mask):
"""
Convert the string to SoftmaxType
"""
if attn_mask_type == 'causal':
return SoftmaxType.SCALED_UPPER_TRIANG_MASKED
if attn_mask_type == 'padding':
if mask is not None:
return SoftmaxType.SCALED_MASKED
return SoftmaxType.SCALED
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

softmax_type = convert_to_softmax_type(self.attn_mask_type, mask)

x = core_attention(query,
key,
Expand DownExpand Up@@ -765,6 +800,18 @@ class TransformerLayer(nn.Module):
an attention block and a feedforward network (MLP).
This standard layer is based on the paper “Attention Is All You Need”.

.. warning::

Argument :attr:`self_attn_mask_type` is introduced in version 0.10.
Starting from version 0.11, the default value will be `"causal"`.
However, to ensure compatibility with earlier versions, before 0.11,
the default value will be `"padding"` for the encoder and `"causal"` for the decoder.

.. note::

Argument :attr:`attention_mask` will be ignored when
:attr:`self_attn_mask_type` is set to `"causal"`.

Parameters
----------
hidden_size: int, default = 512
Expand DownExpand Up@@ -825,6 +872,8 @@ class TransformerLayer(nn.Module):
If set to TransformerLayerType.DECODER, an additional cross-attention block
is added after self-attention.this can be used for structures like `T5`
Transformer in conjunction with the TransformerLayerType.ENCODER option.
self_attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.
enable_relative_embedding: bool, default = True
Whether to enable relative embedding as shifting of attention logits.
relative_embedding: flax.linen.Module, default = None
Expand DownExpand Up@@ -878,6 +927,7 @@ class TransformerLayer(nn.Module):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: nn.Module = None
dtype: DType = jnp.float32
Expand All@@ -893,6 +943,19 @@ def __post_init__(self):
if self.mlp_kernel_init is None:
self.mlp_kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in',
'truncated_normal')
# TODO(rewang): default to 'causal' in 0.11 (also updated the doc after 0.11)
if self.self_attn_mask_type is None:
warnings.warn(
"The 'self_attn_mask_type' argument in the 'TransformerLayer' is"
" introduced in version 0.10. Starting from version 0.11, the default"
" value will be 'causal'. However, to ensure compatibility with earlier"
" versions, before 0.11, the default value will be 'padding' for the"
" encoder and 'causal' for the decoder.",
category=FutureWarning)
if self.layer_type == TransformerLayerType.ENCODER:
self.self_attn_mask_type = 'padding'
else:
self.self_attn_mask_type = 'causal'
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -975,16 +1038,12 @@ def __call__(self,

assert inputs.ndim == 3

self_attn_type = None
# Make name be the exactly same as T5X, since names would affect
# RNGKey during init and apply. Myabe no need in the feature.
if self.layer_type == TransformerLayerType.ENCODER:
mha_name = 'attention'
self_attn_type = AttentionType.PADDING
else:
mha_name = 'self_attention'
self_attn_type = AttentionType.CAUSAL
assert self_attn_type is not None

# [batch, length, emb_dim] -> [batch, length, emb_dim]
x, residual = MultiHeadAttention(
Expand All@@ -1002,7 +1061,7 @@ def __call__(self,
zero_centered_gamma=self.zero_centered_gamma,
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self_attn_type,
attn_mask_type=self.self_attn_mask_type,
fuse_qkv=self.fuse_qkv_params,
kernel_init=self.mha_kernel_init,
use_bias=self.use_bias,
Expand DownExpand Up@@ -1049,7 +1108,7 @@ def hidden_dropout(x, deterministic):
apply_residual_connection_post_layernorm=self.
apply_residual_connection_post_layernorm,
output_layernorm=False, # Must do LayerNorm before MHA.
attn_type=AttentionType.PADDING,
attn_mask_type='padding',
float32_logits=self.float32_attention_logits,
scale_attn_logits=self.scale_attn_logits,
scaled_query_init=self.scaled_query_init,
Expand Down
12 changes: 8 additions & 4 deletions transformer_engine/jax/praxis/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -5,14 +5,14 @@
Praxis Modules related Transformer
"""
from functools import partial
from typing import Optional, Sequence, Tuple
from typing import Any, Optional, Sequence, Tuple

from praxis import pax_fiddle
from praxis.base_layer import WeightInit
from praxis.pytypes import JTensor

from .module import TransformerEngineBaseLayer
from ..flax.transformer import AttentionType, TransformerLayerType
from ..flax.transformer import TransformerLayerType
from ..flax.transformer import MultiHeadAttention as flax_MultiHeadAttention
from ..flax.transformer import RelativePositionBiases as flax_RelativePositionBiases
from ..flax.transformer import TransformerLayer as flax_TransformerLayer
Expand DownExpand Up@@ -73,7 +73,9 @@ class MultiHeadAttention(TransformerEngineBaseLayer):
bias_init: WeightInit = WeightInit.Constant(0.0)
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
scale_attn_logits: bool = False
Expand All@@ -99,7 +101,7 @@ def setup(self) -> None:
bias_init=TransformerEngineBaseLayer.generate_params_init("bias", self.bias_init),
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self.attn_type,
attn_mask_type=self.attn_mask_type,
fuse_qkv=self.fuse_qkv,
transpose_batch_sequence=self.transpose_batch_sequence,
scale_attn_logits=self.scale_attn_logits,
Expand DownExpand Up@@ -145,6 +147,7 @@ class TransformerLayer(TransformerEngineBaseLayer):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: pax_fiddle.Config[RelativePositionBiases] = pax_fiddle.template_field(None)
drop_path: float = 0.0
Expand DownExpand Up@@ -201,6 +204,7 @@ def setup(self) -> None:
output_layernorm=self.output_layernorm,
float32_attention_logits=self.float32_attention_logits,
layer_type=self.layer_type,
self_attn_mask_type=self.self_attn_mask_type,
enable_relative_embedding=self.enable_relative_embedding,
relative_embedding=relative_embedding_flax_module,
drop_path=self.drop_path,
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Universal Dark Mode - works on any site (function() { var enabled = true; function applyDarkMode() { if (!enabled) return; // Create style element if it doesn't exist var style = document.getElementById('universal-dark-mode-style'); if (!style) { style = document.createElement('style'); style.id = 'universal-dark-mode-style'; document.head.appendChild(style); } // Dark mode CSS - inverts colors but preserves images/video style.textContent = ' /* Invert everything except media */ html { filter: invert(1) hue-rotate(180deg) !important; background: #1a1a2e !important; } /* Restore images, videos, iframes, canvas */ img, video, iframe, canvas, svg, picture, [style*="background-image"] { filter: invert(1) hue-rotate(180deg) !important; } /* Preserve specific elements that should not be inverted */ .no-dark-mode, .no-dark-mode *, [data-theme="light"], [data-theme="light"], .ace_editor, .ace_editor *, .CodeMirror, .CodeMirror *, .monaco-editor, .monaco-editor *, .markdown-body pre, .markdown-body pre *, .highlight, .highlight *, pre code, pre code * { filter: none !important; } /* Fix common UI elements */ .modal, .popup, .dropdown-menu, .tooltip, .popover { filter: invert(1) hue-rotate(180deg) !important; background: #2d2d44 !important; border-color: #444 !important; } /* Scrollbars */ ::-webkit-scrollbar { background: #1a1a2e !important; } ::-webkit-scrollbar-thumb { background: #444 !important; } ::-webkit-scrollbar-thumb:hover { background: #555 !important; } /* Selection */ ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; } ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; } '; } function removeDarkMode() { var style = document.getElementById('universal-dark-mode-style'); if (style) style.remove(); } // Toggle with Alt+Shift+D document.addEventListener('keydown', function(e) { if (e.altKey && e.shiftKey && e.key === 'D') { e.preventDefault(); enabled = !enabled; if (enabled) { applyDarkMode(); console.log('[Universal Dark Mode] Enabled'); } else { removeDarkMode(); console.log('[Universal Dark Mode] Disabled'); } } }); // Apply on load applyDarkMode(); // Re-apply on dynamic content var observer = new MutationObserver(function(mutations) { if (enabled && !document.getElementById('universal-dark-mode-style')) { applyDarkMode(); } }); observer.observe(document.head, { childList: true }); console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle'); })(); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })(); [JAX] Add self_attn_mask_type and replace attn_type by zlsh80826 · Pull Request #273 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 6 additions & 7 deletions tests/jax/test_praxis_layers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from transformer_engine.jax.flax import RelativePositionBiases as flax_RelativePositionBiases
from transformer_engine.jax.flax import TransformerLayer as flax_TransformerLayer
from transformer_engine.jax.flax.module import Softmax
from transformer_engine.jax.flax.transformer import AttentionType
from transformer_engine.jax.fp8 import FP8Helper, is_fp8_available
from transformer_engine.jax.praxis import LayerNorm
from transformer_engine.jax.praxis import FusedSoftmax, LayerNorm
Expand DownExpand Up@@ -666,32 +665,32 @@ class MultiHeadAttnAttr:
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.PADDING
ATTN_TYPE: 'padding'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'layernorm',
ZERO_CEN: True,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}, {
USE_BIAS: True,
LN_TYPE: 'rmsnorm',
ZERO_CEN: False,
ATTN_TYPE: AttentionType.CAUSAL
ATTN_TYPE: 'causal'
}]


Expand Down
105 changes: 82 additions & 23 deletions transformer_engine/jax/flax/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -197,17 +197,16 @@ def core_attention(query: Array,
dynamic_vector_slice_in_dim = vmap(lax.dynamic_slice_in_dim, in_axes=(None, 0, None, None))


class AttentionType(Enum):
"""TransformerLayerType."""
PADDING = AttnMaskType.PADDING_MASK
CAUSAL = AttnMaskType.CAUSAL_MASK


class MultiHeadAttention(nn.Module):
r"""
Multi-head Attention (MHA), including Query,
Key, Value and Output projection.

.. warning::

Argument :attr:`attn_type` is deprecated and superseded by :attr:`attn_mask_type`.
:attr:`attn_type` is ignored in version 0.10 and will be fully removed in version 0.11.

Parameters
----------
head_dim : int
Expand DownExpand Up@@ -245,8 +244,11 @@ class MultiHeadAttention(nn.Module):
Indicate if apply residual connection with the output of layer normalization.
output_layernorm : bool, default = False
Indicate if apply a layer normalization at the end of MHA.
attn_type: AttentionType, defult = AttentionType.PADDING
Indicate the format of the attention mask in the core attention.
attn_type: Any, defult = None
*Deprecated*, will be ignored in v0.10 and be fully removed in v0.11.
Please use `attn_mask_type` to config the attention mask.
attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.

Optimization parameters
-----------------------
Expand DownExpand Up@@ -282,7 +284,9 @@ class MultiHeadAttention(nn.Module):
bias_init: Initializer = nn.initializers.zeros
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
dtype: DType = jnp.float32
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
Expand All@@ -293,6 +297,14 @@ class MultiHeadAttention(nn.Module):
def __post_init__(self):
if self.kernel_init is None:
self.kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in', 'normal')
# TODO(rewang): remove attn_type after v0.11
if self.attn_type is not None:
warnings.warn(
"The 'attn_type' argument in the 'MultiHeadAttention' is"
" deprecated in version 0.10 and will be removed in version 0.11."
" Passing value in attn_type will be ignored, please use `attn_mask_type`"
" to config the attention mask type.",
category=DeprecationWarning)
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -570,9 +582,23 @@ def kv_init(key, shape, dtype):
if use_fused_attn:
assert mask is not None and mask.ndim == 4 # (b, 1, s_q, s_kv)
assert not self.transpose_batch_sequence

# TODO(rewang): make it configurable for pre_scale_bias
attn_bias_type = AttnBiasType.NO_BIAS if bias is None else AttnBiasType.POST_SCALE_BIAS

def canonicalize_attn_mask_type(attn_mask_type):
"""
Convert the string to AttnMaskType
"""
if attn_mask_type == 'causal':
return AttnMaskType.CAUSAL_MASK
if attn_mask_type == 'padding':
return AttnMaskType.PADDING_MASK
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

attn_mask_type = canonicalize_attn_mask_type(self.attn_mask_type)

if inputs_q is inputs_kv:
qkv_proj = qkv_proj.reshape((*qkv_proj.shape[:-1], self.num_heads, self.head_dim))
qkv_sharding_constraint = ('batch', 'length', 'qkv_dim', 'heads', 'kv')
Expand All@@ -583,7 +609,7 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
Expand All@@ -602,18 +628,27 @@ def kv_init(key, shape, dtype):
mask,
dropout_rng,
attn_bias_type=attn_bias_type,
attn_mask_type=self.attn_type.value,
attn_mask_type=attn_mask_type,
scaling_factor=scale_factor,
dropout_probability=self.dropout_rate,
is_training=not deterministic,
sharding_type=first_sharding_type)
else:
softmax_type = SoftmaxType.SCALED
if self.attn_type is AttentionType.PADDING:
if mask is not None:
softmax_type = SoftmaxType.SCALED_MASKED
else:
softmax_type = SoftmaxType.SCALED_UPPER_TRIANG_MASKED

def convert_to_softmax_type(attn_mask_type, mask):
"""
Convert the string to SoftmaxType
"""
if attn_mask_type == 'causal':
return SoftmaxType.SCALED_UPPER_TRIANG_MASKED
if attn_mask_type == 'padding':
if mask is not None:
return SoftmaxType.SCALED_MASKED
return SoftmaxType.SCALED
raise ValueError(f"Unsupported {attn_mask_type=}, "
"supported attn_mask_type = {'causal', 'padding'}")

softmax_type = convert_to_softmax_type(self.attn_mask_type, mask)

x = core_attention(query,
key,
Expand DownExpand Up@@ -765,6 +800,18 @@ class TransformerLayer(nn.Module):
an attention block and a feedforward network (MLP).
This standard layer is based on the paper “Attention Is All You Need”.

.. warning::

Argument :attr:`self_attn_mask_type` is introduced in version 0.10.
Starting from version 0.11, the default value will be `"causal"`.
However, to ensure compatibility with earlier versions, before 0.11,
the default value will be `"padding"` for the encoder and `"causal"` for the decoder.

.. note::

Argument :attr:`attention_mask` will be ignored when
:attr:`self_attn_mask_type` is set to `"causal"`.

Parameters
----------
hidden_size: int, default = 512
Expand DownExpand Up@@ -825,6 +872,8 @@ class TransformerLayer(nn.Module):
If set to TransformerLayerType.DECODER, an additional cross-attention block
is added after self-attention.this can be used for structures like `T5`
Transformer in conjunction with the TransformerLayerType.ENCODER option.
self_attn_mask_type: {'causal', 'padding'}, default = 'causal'
Type of attention mask passed into softmax operation.
enable_relative_embedding: bool, default = True
Whether to enable relative embedding as shifting of attention logits.
relative_embedding: flax.linen.Module, default = None
Expand DownExpand Up@@ -878,6 +927,7 @@ class TransformerLayer(nn.Module):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: nn.Module = None
dtype: DType = jnp.float32
Expand All@@ -893,6 +943,19 @@ def __post_init__(self):
if self.mlp_kernel_init is None:
self.mlp_kernel_init = nn.initializers.variance_scaling(1.0, 'fan_in',
'truncated_normal')
# TODO(rewang): default to 'causal' in 0.11 (also updated the doc after 0.11)
if self.self_attn_mask_type is None:
warnings.warn(
"The 'self_attn_mask_type' argument in the 'TransformerLayer' is"
" introduced in version 0.10. Starting from version 0.11, the default"
" value will be 'causal'. However, to ensure compatibility with earlier"
" versions, before 0.11, the default value will be 'padding' for the"
" encoder and 'causal' for the decoder.",
category=FutureWarning)
if self.layer_type == TransformerLayerType.ENCODER:
self.self_attn_mask_type = 'padding'
else:
self.self_attn_mask_type = 'causal'
super().__post_init__()

@nn.compact
Expand DownExpand Up@@ -975,16 +1038,12 @@ def __call__(self,

assert inputs.ndim == 3

self_attn_type = None
# Make name be the exactly same as T5X, since names would affect
# RNGKey during init and apply. Myabe no need in the feature.
if self.layer_type == TransformerLayerType.ENCODER:
mha_name = 'attention'
self_attn_type = AttentionType.PADDING
else:
mha_name = 'self_attention'
self_attn_type = AttentionType.CAUSAL
assert self_attn_type is not None

# [batch, length, emb_dim] -> [batch, length, emb_dim]
x, residual = MultiHeadAttention(
Expand All@@ -1002,7 +1061,7 @@ def __call__(self,
zero_centered_gamma=self.zero_centered_gamma,
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self_attn_type,
attn_mask_type=self.self_attn_mask_type,
fuse_qkv=self.fuse_qkv_params,
kernel_init=self.mha_kernel_init,
use_bias=self.use_bias,
Expand DownExpand Up@@ -1049,7 +1108,7 @@ def hidden_dropout(x, deterministic):
apply_residual_connection_post_layernorm=self.
apply_residual_connection_post_layernorm,
output_layernorm=False, # Must do LayerNorm before MHA.
attn_type=AttentionType.PADDING,
attn_mask_type='padding',
float32_logits=self.float32_attention_logits,
scale_attn_logits=self.scale_attn_logits,
scaled_query_init=self.scaled_query_init,
Expand Down
12 changes: 8 additions & 4 deletions transformer_engine/jax/praxis/transformer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -5,14 +5,14 @@
Praxis Modules related Transformer
"""
from functools import partial
from typing import Optional, Sequence, Tuple
from typing import Any, Optional, Sequence, Tuple

from praxis import pax_fiddle
from praxis.base_layer import WeightInit
from praxis.pytypes import JTensor

from .module import TransformerEngineBaseLayer
from ..flax.transformer import AttentionType, TransformerLayerType
from ..flax.transformer import TransformerLayerType
from ..flax.transformer import MultiHeadAttention as flax_MultiHeadAttention
from ..flax.transformer import RelativePositionBiases as flax_RelativePositionBiases
from ..flax.transformer import TransformerLayer as flax_TransformerLayer
Expand DownExpand Up@@ -73,7 +73,9 @@ class MultiHeadAttention(TransformerEngineBaseLayer):
bias_init: WeightInit = WeightInit.Constant(0.0)
apply_residual_connection_post_layernorm: bool = False
output_layernorm: bool = False
attn_type: AttentionType = AttentionType.PADDING
# TODO(rewang): remove attn_type and the related doc after v0.11
attn_type: Any = None
attn_mask_type: str = 'causal'
fuse_qkv: bool = True
transpose_batch_sequence: bool = True
scale_attn_logits: bool = False
Expand All@@ -99,7 +101,7 @@ def setup(self) -> None:
bias_init=TransformerEngineBaseLayer.generate_params_init("bias", self.bias_init),
apply_residual_connection_post_layernorm=self.apply_residual_connection_post_layernorm,
output_layernorm=self.output_layernorm,
attn_type=self.attn_type,
attn_mask_type=self.attn_mask_type,
fuse_qkv=self.fuse_qkv,
transpose_batch_sequence=self.transpose_batch_sequence,
scale_attn_logits=self.scale_attn_logits,
Expand DownExpand Up@@ -145,6 +147,7 @@ class TransformerLayer(TransformerEngineBaseLayer):
output_layernorm: bool = False
float32_attention_logits: bool = False
layer_type: TransformerLayerType = TransformerLayerType.ENCODER
self_attn_mask_type: str = None # TODO(rewang): default to 'causal' after 0.11
enable_relative_embedding: bool = True
relative_embedding: pax_fiddle.Config[RelativePositionBiases] = pax_fiddle.template_field(None)
drop_path: float = 0.0
Expand DownExpand Up@@ -201,6 +204,7 @@ def setup(self) -> None:
output_layernorm=self.output_layernorm,
float32_attention_logits=self.float32_attention_logits,
layer_type=self.layer_type,
self_attn_mask_type=self.self_attn_mask_type,
enable_relative_embedding=self.enable_relative_embedding,
relative_embedding=relative_embedding_flax_module,
drop_path=self.drop_path,
Expand Down