@@ -951,6 +951,8 @@ def __init__(self, config):
951951
952952 self ._recipe = TransformerEngineQuantization ._get_recipe (config .quantization )
953953
954+ self ._perform_collective_gemm = config .use_te_comm_gemm_overlap
955+
954956 def __hash__ (self ):
955957 return hash ((self .quant_mode , self ._recipe ))
956958
@@ -1001,11 +1003,13 @@ def _wrap(self, f, name=None):
10011003 2. Wraps the given function in a Flax linen module. This module does not store any Flax
10021004 parameters but can store Flax variables for quantizers if required by the recipe.
10031005
1004- 3. When the wrapper is called, it provides an additional argument to the given function `f`,
1005- 'generate_quantizer_set' as the first argument. 'generate_quantizer_set' is a function that
1006- can be called to generate a TransformerEngine/JAX quantizer set object used in
1007- TransformerEngine/JAX APIs. 'generate_quantizer_set' will generate quantizers based on the
1008- recipe of this TransformerEngineQuantizer object.
1006+ 3. When the wrapper is called, it provides two additional arguments to the given function `f`,
1007+ 'generate_quantizer_set' as the first argument and 'generate_collective_op_set' as the second argument.
1008+ 'generate_quantizer_set' is a function that can be called to generate a TransformerEngine/JAX quantizer
1009+ set object used in TransformerEngine/JAX APIs based on the recipe of this TransformerEngineQuantizer
1010+ object. Similarly, 'generate_collective_op_set' is a function that can be called to generate a
1011+ TransformerEngine/JAX collective operation set object used in TransformerEngine/JAX APIs based
1012+ the kernel's mesh axes.
10091013
10101014 Args:
10111015 f: The function to wrap. The first argument must be 'generate_quantizer_set'.
@@ -1016,6 +1020,7 @@ def _wrap(self, f, name=None):
10161020 """
10171021
10181022 import transformer_engine .jax # pylint: disable=import-outside-toplevel # pytype: disable=import-error
1023+ import transformer_engine .jax .cpp_extensions as tex # pylint: disable=import-outside-toplevel # pytype: disable=import-error
10191024 from transformer_engine .common import recipe # pylint: disable=import-outside-toplevel # pytype: disable=import-error
10201025
10211026 default_recipe = self ._recipe
@@ -1041,9 +1046,20 @@ def generate_quantizer_set(
10411046 n_groups = n_groups ,
10421047 )
10431048
1049+ def generate_collective_op_set (self , mesh_axes : Tuple [str , ...] = ()):
1050+ """Inspect the kernel's mesh axes to determine the type of collective operation to use for collective GEMM."""
1051+
1052+ if len (mesh_axes ) >= 1 :
1053+ if mesh_axes [0 ] == "embed" and mesh_axes [- 1 ] == "mlp" :
1054+ return tex .CollectiveOpSet .create (tex .CollectiveOp .ALL_GATHER )
1055+ elif mesh_axes [0 ] == "mlp" and mesh_axes [- 1 ] == "embed" :
1056+ return tex .CollectiveOpSet .create (tex .CollectiveOp .REDUCE_SCATTER )
1057+
1058+ return tex .noop_collective_op_set
1059+
10441060 @nn .compact
10451061 def __call__ (self , * args , ** kwargs ):
1046- return f (self .generate_quantizer_set , * args , ** kwargs )
1062+ return f (self .generate_quantizer_set , self . generate_collective_op_set , * args , ** kwargs )
10471063
10481064 TEWrapper .__name__ = f"TEWrapper_{ name if name else f .__name__ } "
10491065
@@ -1052,17 +1068,22 @@ def __call__(self, *args, **kwargs):
10521068 def dot_general_cls (self , mesh_axes : Tuple [str , ...] = ()):
10531069 """Placeholder for dot_general implementation in subclasses."""
10541070 import transformer_engine .jax # pylint: disable=import-outside-toplevel # pytype: disable=import-error
1071+ import transformer_engine .jax .cpp_extensions as tex # pylint: disable=import-outside-toplevel # pytype: disable=import-error
10551072
1056- def te_dot_general (generate_quantizer_set , x , kernel , dims , ** kwargs ):
1073+ def te_dot_general (generate_quantizer_set , generate_collective_op_set , x , kernel , dims , ** kwargs ):
10571074 contracting_dims , batch_dims = dims
10581075 assert batch_dims == ((), ()), "Batch dimensions must be empty for TransformerEngine dot."
10591076
10601077 quantizer_set = generate_quantizer_set ()
1078+ collective_op_set = (
1079+ generate_collective_op_set (mesh_axes ) if self ._perform_collective_gemm else tex .noop_collective_op_set
1080+ )
10611081 return transformer_engine .jax .dense .dense (
10621082 x ,
10631083 kernel ,
10641084 contracting_dims = contracting_dims ,
10651085 quantizer_set = quantizer_set ,
1086+ collective_op_set = collective_op_set ,
10661087 )
10671088
10681089 return self ._wrap (te_dot_general , "dot_general" )
0 commit comments