2222import jax
2323import jax .numpy as jnp
2424from maxtext .kernels .megablox import backend
25+ from maxtext .kernels .megablox import pallas_mosaic_tpu_v2_gmm_kernel as gmm_v2
26+ from maxtext .kernels .megablox import pallas_mosaic_tpu_v2_tgmm_kernel as tgmm_v2
2527from maxtext .layers import quantizations
2628import qwix
2729import qwix .pallas as qpl
2830import tokamax
2931
3032
33+ DLHS_RAGGED_DOT_DIM_NUMS = jax .lax .RaggedDotDimensionNumbers (
34+ dot_dimension_numbers = (([1 ], [2 ]), ([], [])),
35+ lhs_ragged_dimensions = [0 ],
36+ rhs_group_dimensions = [0 ],
37+ )
38+
3139DRHS_RAGGED_DOT_DIM_NUMS = jax .lax .RaggedDotDimensionNumbers (
3240 dot_dimension_numbers = (([0 ], [0 ]), ([], [])),
3341 lhs_ragged_dimensions = [0 ],
@@ -65,6 +73,7 @@ def gmm(
6573 # TODO(amandaliang): get rid of the qwix_rule in favor of Qwix's interception feature
6674 qwix_rule : qwix .QtRule | None = None ,
6775 use_manual_quantization : bool = False , # used in batchsplit
76+ use_gmm_v2 : bool = False ,
6877):
6978 """Grouped matrix multiplication operation."""
7079 quantization_rule = None
@@ -84,7 +93,7 @@ def gmm(
8493 )
8594
8695 gmm_fwd_bwd = lambda * args : _gmm_fwd (* args )[0 ] # pylint: disable=C3001
87- gmm_fwd_bwd = jax .custom_vjp (gmm_fwd_bwd , nondiff_argnums = (3 , 4 , 7 , 8 , 9 , 10 , 11 , 12 , 13 , 14 ))
96+ gmm_fwd_bwd = jax .custom_vjp (gmm_fwd_bwd , nondiff_argnums = (3 , 4 , 7 , 8 , 9 , 10 , 11 , 12 , 13 , 14 , 15 ))
8897 gmm_fwd_bwd .defvjp (_gmm_fwd , functools .partial (_gmm_bwd , lhs .dtype , rhs .dtype ))
8998 return gmm_fwd_bwd (
9099 lhs ,
@@ -102,6 +111,7 @@ def gmm(
102111 use_manual_quantization ,
103112 lhs_vma_axes ,
104113 rhs_vma_axes ,
114+ use_gmm_v2 ,
105115 )
106116
107117
@@ -131,6 +141,7 @@ def _gmm_fwd(
131141 use_manual_quantization : bool = False ,
132142 lhs_vma_axes : tuple = tuple (),
133143 rhs_vma_axes : tuple = tuple (),
144+ use_gmm_v2 : bool = False ,
134145) -> tuple [
135146 jnp .ndarray ,
136147 tuple [
@@ -178,6 +189,19 @@ def _gmm_fwd(
178189 if transpose_rhs :
179190 rhs = rhs .swapaxes (1 , 2 )
180191
192+ if use_gmm_v2 :
193+ out = gmm_v2 .gmm_v2 (
194+ lhs = lhs ,
195+ rhs = rhs ,
196+ group_sizes = group_sizes ,
197+ tile_info = gmm_v2 .TileSizes (
198+ tile_m = tiling [0 ],
199+ tile_k = tiling [1 ],
200+ tile_n = tiling [2 ],
201+ ),
202+ preferred_element_type = preferred_element_type ,
203+ )
204+ elif use_tokamax_backend :
181205 # manual_axis_type is for gmm with shard_map check_vma=True, needs tokamax > 0.0.12
182206 out_kwargs = {}
183207 if use_manual_quantization :
@@ -226,6 +250,7 @@ def _gmm_bwd(
226250 use_manual_quantization : bool ,
227251 lhs_vma_axes : tuple ,
228252 rhs_vma_axes : tuple ,
253+ use_gmm_v2 : bool ,
229254 residual : tuple [
230255 jnp .ndarray | qpl .QArray ,
231256 jnp .ndarray | qpl .QArray ,
@@ -274,10 +299,10 @@ def _gmm_bwd(
274299 channelwise_axes = [] if quantization_rule .disable_channelwise_axes else [1 ],
275300 calibration_method = quantization_rule .bwd_calibration_method ,
276301 )
277- if use_tokamax_backend :
302+ if use_tokamax_backend or use_gmm_v2 :
278303 # Handle transpose_rhs manually
279304 dlhs_rhs = rhs
280- if not transpose_rhs :
305+ if transpose_rhs :
281306 dlhs_rhs = dlhs_rhs .swapaxes (1 , 2 )
282307
283308 # manual_axis_type is for gmm with shard_map check_vma=True, needs tokamax > 0.0.12
@@ -290,29 +315,63 @@ def _gmm_bwd(
290315 varying = frozenset (["expert" ]), unreduced = frozenset (["data" , "fsdp" ])
291316 )
292317
293- dlhs = tokamax .ragged_dot (
294- lhs = dlhs_dout ,
295- rhs = dlhs_rhs ,
296- group_sizes = group_sizes ,
297- precision = jax .lax .Precision .DEFAULT ,
298- preferred_element_type = lhs_dtype ,
299- # `group_offset` is not yet supported
300- group_offset = None ,
301- implementation = "mosaic" ,
302- ** dlhs_kwargs ,
303- )
304- drhs = tokamax .ragged_dot_general (
305- lhs = lhs ,
306- rhs = drhs_dout ,
307- group_sizes = group_sizes ,
308- ragged_dot_dimension_numbers = DRHS_RAGGED_DOT_DIM_NUMS ,
309- precision = jax .lax .Precision .DEFAULT ,
310- preferred_element_type = rhs_dtype ,
311- # `group_offset` is not yet supported
312- group_offset = None ,
313- implementation = "mosaic" ,
314- ** drhs_kwargs ,
315- )
318+ if use_gmm_v2 :
319+ dlhs = gmm_v2 .gmm_v2 (
320+ lhs = dlhs_dout ,
321+ rhs = dlhs_rhs .swapaxes (1 , 2 ), # requires rhs to be [g, n, k]
322+ group_sizes = group_sizes ,
323+ tile_info = gmm_v2 .TileSizes (
324+ tile_m = tiling [3 ],
325+ tile_k = tiling [4 ],
326+ tile_n = tiling [5 ],
327+ ),
328+ preferred_element_type = lhs_dtype ,
329+ )
330+
331+ # tgmm_v2_op requires lhs and rhs to have the same dtype.
332+ if lhs .dtype != drhs_dout .dtype :
333+ drhs_dout = drhs_dout .astype (lhs .dtype )
334+
335+ drhs = tgmm_v2 .tgmm_v2 (
336+ lhs = lhs ,
337+ rhs = drhs_dout ,
338+ group_sizes = group_sizes ,
339+ num_actual_groups = num_actual_groups ,
340+ precision = jax .lax .Precision .DEFAULT ,
341+ preferred_element_type = rhs_dtype ,
342+ group_offset = group_offset ,
343+ tile_info = gmm_v2 .TileSizes (
344+ tile_m = tiling [6 ],
345+ tile_k = tiling [7 ],
346+ tile_n = tiling [8 ],
347+ ),
348+ )
349+ else :
350+ dlhs = tokamax .ragged_dot_general (
351+ lhs = dlhs_dout ,
352+ rhs = dlhs_rhs ,
353+ group_sizes = group_sizes ,
354+ ragged_dot_dimension_numbers = DLHS_RAGGED_DOT_DIM_NUMS ,
355+ precision = jax .lax .Precision .DEFAULT ,
356+ preferred_element_type = lhs_dtype ,
357+ # `group_offset` is not yet supported
358+ group_offset = None ,
359+ implementation = "mosaic" ,
360+ ** dlhs_kwargs ,
361+ )
362+
363+ drhs = tokamax .ragged_dot_general (
364+ lhs = lhs ,
365+ rhs = drhs_dout ,
366+ group_sizes = group_sizes ,
367+ ragged_dot_dimension_numbers = DRHS_RAGGED_DOT_DIM_NUMS ,
368+ precision = jax .lax .Precision .DEFAULT ,
369+ preferred_element_type = rhs_dtype ,
370+ # `group_offset` is not yet supported
371+ group_offset = None ,
372+ implementation = "mosaic" ,
373+ ** drhs_kwargs ,
374+ )
316375 if quantization_rule and quantization_rule .bwd_qtype and weight_gather_axes :
317376 # Scatter back in reverse order of gather
318377 for axis_name , axis_idx in reversed (weight_gather_axes ):
0 commit comments