1717"""Custom Pallas flash attention kernel for TPU."""
1818
1919import functools
20- import math
2120
2221import jax
2322import jax .numpy as jnp
2423import numpy as np
2524from jax import lax
2625from jax .experimental import pallas as pl
2726from jax .experimental .pallas import tpu as pltpu
28- from jax .experimental .shard_map import shard_map
29- from jax .sharding import PartitionSpec as P
3027
3128DEFAULT_MASK_VALUE = - 0.7 * float (np .finfo (np .dtype ("float32" )).max )
3229NUM_LANES = 128
3330NUM_SUBLANES = 8
3431NT_DIM_NUMBERS = (((1 ,), (1 ,)), ((), ()))
3532
36- # Default block sizes (tuned for 720p Wan2.1 on v6e/v7x)
37- DEFAULT_BQSIZE = 3328
38- DEFAULT_BKVSIZE = 2816
39- # Cranked up to 1024 for massive MXU throughput
40- DEFAULT_BKVCOMPUTESIZE = 1024
41- # Kept at 256 to protect VPU registers (V1 Optimization)
42- DEFAULT_BKVCOMPUTEINSIZE = 256
43-
4433
4534class _BlockSizes :
46- __slots__ = ("block_q" , "block_kv" , "block_kv_compute" )
35+ __slots__ = ("block_q" , "block_kv" , "block_kv_compute" , "block_kv_compute_in" )
4736
48- def __init__ (self , block_q : int , block_kv : int , block_kv_compute : int | None = None ):
37+ def __init__ (self , block_q : int , block_kv : int , block_kv_compute : int | None = None , block_kv_compute_in : int = 256 ):
4938 self .block_q = block_q
5039 self .block_kv = block_kv
5140 self .block_kv_compute = block_kv_compute if block_kv_compute is not None else block_kv
41+ self .block_kv_compute_in = block_kv_compute_in
5242
5343
5444# Fixed-m softmax-bound constants. Instead of tracking the online-softmax
@@ -78,12 +68,10 @@ def _flash_attention_kernel(
7868 * ,
7969 mask_value : float ,
8070 grid_width : int ,
81- bq : int ,
8271 bkv : int ,
8372 bkv_compute : int ,
8473 bkv_compute_in : int ,
8574 head_dim_v : int ,
86- q_seq_len : int ,
8775 kv_seq_len : int ,
8876 use_base2_exp : bool = True ,
8977 fuse_reciprocal : bool = True ,
@@ -260,12 +248,10 @@ def _flash_attention_kernel_mhpt(
260248 * ,
261249 mask_value : float ,
262250 grid_width : int ,
263- bq : int ,
264251 bkv : int ,
265252 bkv_compute : int ,
266253 bkv_compute_in : int ,
267254 head_dim_v : int ,
268- q_seq_len : int ,
269255 kv_seq_len : int ,
270256 heads_per_tile : int ,
271257 use_base2_exp : bool = True ,
@@ -407,7 +393,6 @@ def _splash_attention_forward(
407393 k : jax .Array ,
408394 v : jax .Array ,
409395 block_sizes : _BlockSizes ,
410- bkv_compute_in : int ,
411396 q_seq_len : int | None = None ,
412397 kv_seq_len : int | None = None ,
413398 use_base2_exp : bool = True ,
@@ -425,6 +410,7 @@ def _splash_attention_forward(
425410 mk = jnp .zeros ((2 , num_q_heads ), jnp .float32 )
426411 bq , bkv = block_sizes .block_q , block_sizes .block_kv
427412 bkv_compute = block_sizes .block_kv_compute
413+ bkv_compute_in = block_sizes .block_kv_compute_in
428414 num_kv_heads = k .shape [0 ]
429415 padded_kv_seq_len = k .shape [1 ]
430416
@@ -470,12 +456,10 @@ def v_index_map(h, i, j, *_):
470456 _flash_attention_kernel ,
471457 mask_value = DEFAULT_MASK_VALUE ,
472458 grid_width = grid_width ,
473- bq = bq ,
474459 bkv = bkv ,
475460 bkv_compute = bkv_compute ,
476461 bkv_compute_in = bkv_compute_in ,
477462 head_dim_v = head_dim_v ,
478- q_seq_len = actual_q_seq_len ,
479463 kv_seq_len = actual_kv_seq_len ,
480464 use_base2_exp = use_base2_exp ,
481465 use_fixed_m = use_fixed_m ,
@@ -503,7 +487,6 @@ def _splash_attention_forward_ring(
503487 k : jax .Array ,
504488 v : jax .Array ,
505489 block_sizes : _BlockSizes ,
506- bkv_compute_in : int ,
507490 q_seq_len : int | None = None ,
508491 kv_seq_len : int | None = None ,
509492 use_base2_exp : bool = True ,
@@ -530,6 +513,7 @@ def _splash_attention_forward_ring(
530513 head_dim_v = v .shape [- 1 ]
531514 bq , bkv = block_sizes .block_q , block_sizes .block_kv
532515 bkv_compute = block_sizes .block_kv_compute
516+ bkv_compute_in = block_sizes .block_kv_compute_in
533517 num_kv_heads = k .shape [0 ]
534518 padded_kv_seq_len = k .shape [1 ]
535519
@@ -586,12 +570,10 @@ def v_index_map(h, i, j, *_):
586570 _flash_attention_kernel ,
587571 mask_value = DEFAULT_MASK_VALUE ,
588572 grid_width = grid_width ,
589- bq = bq ,
590573 bkv = bkv ,
591574 bkv_compute = bkv_compute ,
592575 bkv_compute_in = bkv_compute_in ,
593576 head_dim_v = head_dim_v ,
594- q_seq_len = actual_q_seq_len ,
595577 kv_seq_len = actual_kv_seq_len ,
596578 use_base2_exp = use_base2_exp ,
597579 fuse_reciprocal = False ,
@@ -623,7 +605,6 @@ def _splash_attention_forward_mhpt(
623605 k : jax .Array ,
624606 v : jax .Array ,
625607 block_sizes : _BlockSizes ,
626- bkv_compute_in : int ,
627608 heads_per_tile : int ,
628609 q_seq_len : int | None = None ,
629610 kv_seq_len : int | None = None ,
@@ -635,6 +616,7 @@ def _splash_attention_forward_mhpt(
635616 head_dim_v = v .shape [- 1 ]
636617 bq , bkv = block_sizes .block_q , block_sizes .block_kv
637618 bkv_compute = block_sizes .block_kv_compute
619+ bkv_compute_in = block_sizes .block_kv_compute_in
638620 num_kv_heads = k .shape [0 ]
639621 actual_q_seq_len = q_seq_len if q_seq_len is not None else padded_q_seq_len
640622 actual_kv_seq_len = kv_seq_len if kv_seq_len is not None else k .shape [1 ]
@@ -681,12 +663,10 @@ def out_index_map(h, i, j, *_):
681663 _flash_attention_kernel_mhpt ,
682664 mask_value = DEFAULT_MASK_VALUE ,
683665 grid_width = grid_width ,
684- bq = bq ,
685666 bkv = bkv ,
686667 bkv_compute = bkv_compute ,
687668 bkv_compute_in = bkv_compute_in ,
688669 head_dim_v = head_dim_v ,
689- q_seq_len = actual_q_seq_len ,
690670 kv_seq_len = actual_kv_seq_len ,
691671 heads_per_tile = hpt ,
692672 use_base2_exp = use_base2_exp ,
@@ -711,7 +691,6 @@ def out_index_map(h, i, j, *_):
711691
712692def make_splash_mha (
713693 block_sizes : _BlockSizes ,
714- bkv_compute_in : int = DEFAULT_BKVCOMPUTEINSIZE ,
715694 orig_q_seq_len : int | None = None ,
716695 orig_kv_seq_len : int | None = None ,
717696 heads_per_tile : int = 1 ,
@@ -729,7 +708,6 @@ def _splash_attention(q, k, v, mk=None):
729708 k ,
730709 v ,
731710 block_sizes ,
732- bkv_compute_in ,
733711 heads_per_tile ,
734712 q_seq_len = orig_q_seq_len ,
735713 kv_seq_len = orig_kv_seq_len ,
@@ -742,7 +720,6 @@ def _splash_attention(q, k, v, mk=None):
742720 k ,
743721 v ,
744722 block_sizes ,
745- bkv_compute_in ,
746723 q_seq_len = orig_q_seq_len ,
747724 kv_seq_len = orig_kv_seq_len ,
748725 use_base2_exp = use_base2_exp ,
@@ -753,200 +730,3 @@ def _splash_attention(q, k, v, mk=None):
753730 )
754731
755732 return _splash_attention
756-
757-
758- # ---------------------------------------------------------------------------
759- # High-level attention function with shard_map
760- # ---------------------------------------------------------------------------
761-
762-
763- def tpu_custom_attention (
764- query ,
765- key ,
766- value ,
767- mesh ,
768- * ,
769- scale = None ,
770- block_q = None ,
771- block_kv = None ,
772- block_kv_compute = None ,
773- block_kv_compute_in = None ,
774- heads_per_tile = None ,
775- use_base2_exp = True ,
776- use_experimental_scheduler = False ,
777- vmem_limit_bytes = None ,
778- flash_block_sizes = None ,
779- ):
780- _LOG2_E = 1.44269504
781- num_heads = query .shape [1 ]
782-
783- if flash_block_sizes is not None :
784- block_q = flash_block_sizes .get ("block_q" , block_q )
785- block_kv = flash_block_sizes .get ("block_kv" , block_kv )
786- block_kv_compute = flash_block_sizes .get ("block_kv_compute" , block_kv_compute )
787- block_kv_compute_in = flash_block_sizes .get ("block_kv_compute_in" , block_kv_compute_in )
788- heads_per_tile = flash_block_sizes .get ("heads_per_tile" , heads_per_tile )
789- vmem_limit_bytes = flash_block_sizes .get ("vmem_limit_bytes" , vmem_limit_bytes )
790-
791- block_q = block_q if block_q is not None else DEFAULT_BQSIZE
792- block_kv = block_kv if block_kv is not None else DEFAULT_BKVSIZE
793- block_kv_compute = block_kv_compute if block_kv_compute is not None else DEFAULT_BKVCOMPUTESIZE
794- block_kv_compute_in = block_kv_compute_in if block_kv_compute_in is not None else DEFAULT_BKVCOMPUTEINSIZE
795- heads_per_tile = heads_per_tile if heads_per_tile is not None else 1
796-
797- def _attention_on_slices (q , k , v ):
798- scale_factor = 1.0 / math .sqrt (q .shape [- 1 ]) if scale is None else scale
799- if use_base2_exp :
800- q = q * scale_factor * _LOG2_E
801- else :
802- q = q * scale_factor
803-
804- def _pad_to_multiple (x , multiple , axis ):
805- seq_len = x .shape [axis ]
806- pad_len = (multiple - seq_len % multiple ) % multiple
807- if pad_len == 0 :
808- return x , seq_len
809- pad_width = [(0 , 0 )] * x .ndim
810- pad_width [axis ] = (0 , pad_len )
811- return jnp .pad (x , pad_width ), seq_len
812-
813- def _kernel_3d (q_3d , k_3d , v_3d ):
814- q_orig_len = q_3d .shape [1 ]
815- kv_orig_len = k_3d .shape [1 ]
816-
817- q_3d_padded , _ = _pad_to_multiple (q_3d , block_q , axis = 1 )
818- k_3d_padded , _ = _pad_to_multiple (k_3d , block_kv , axis = 1 )
819- v_3d_padded , _ = _pad_to_multiple (v_3d , block_kv , axis = 1 )
820-
821- padded_q_seq_len = q_3d_padded .shape [1 ]
822- padded_kv_seq_len = k_3d_padded .shape [1 ]
823-
824- bsizes = _BlockSizes (
825- block_q = min (block_q , padded_q_seq_len ),
826- block_kv = min (block_kv , padded_kv_seq_len ),
827- block_kv_compute = min (block_kv_compute , padded_kv_seq_len ),
828- )
829- splash_kernel = make_splash_mha (
830- block_sizes = bsizes ,
831- bkv_compute_in = block_kv_compute_in ,
832- orig_q_seq_len = q_orig_len ,
833- orig_kv_seq_len = kv_orig_len ,
834- heads_per_tile = heads_per_tile ,
835- use_base2_exp = use_base2_exp ,
836- use_experimental_scheduler = use_experimental_scheduler ,
837- vmem_limit_bytes = vmem_limit_bytes ,
838- )
839- out = splash_kernel (
840- q_3d_padded .astype (jnp .bfloat16 ),
841- k_3d_padded ,
842- v_3d_padded ,
843- )
844- out = jnp .swapaxes (out , 1 , 2 )
845- return out [:, :q_orig_len , ...]
846-
847- return jax .vmap (_kernel_3d , in_axes = (0 , 0 , 0 ), out_axes = 0 )(q , k , v )
848-
849- batch_size = query .shape [0 ]
850- if num_heads < mesh .size :
851- q_partition_spec = P ()
852- kv_partition_spec = P ()
853- out_constraint = P ()
854- else :
855- axis_names = mesh .axis_names
856- if len (axis_names ) == 1 :
857- tp_axis = axis_names [0 ]
858- q_partition_spec = P (None , tp_axis , None , None )
859- kv_partition_spec = P (None , tp_axis , None , None )
860- out_constraint = P (None , None , tp_axis , None )
861- elif len (axis_names ) == 2 :
862- dp_axis , tp_axis = axis_names [0 ], axis_names [1 ]
863- dp_size = mesh .shape [dp_axis ]
864- if batch_size >= dp_size :
865- q_partition_spec = P (dp_axis , tp_axis , None , None )
866- kv_partition_spec = P (dp_axis , tp_axis , None , None )
867- out_constraint = P (dp_axis , None , tp_axis , None )
868- else :
869- all_axes = tuple (axis_names )
870- q_partition_spec = P (None , all_axes , None , None )
871- kv_partition_spec = P (None , all_axes , None , None )
872- out_constraint = P (None , None , all_axes , None )
873- else :
874- q_partition_spec = P (axis_names [0 ], axis_names [1 ], axis_names [2 ], None )
875- kv_partition_spec = P (axis_names [0 ], axis_names [1 ], None , None )
876- out_constraint = P (axis_names [0 ], None , (axis_names [1 ], axis_names [2 ]), None )
877-
878- sharded_fn = shard_map (
879- _attention_on_slices ,
880- mesh = mesh ,
881- in_specs = (q_partition_spec , kv_partition_spec , kv_partition_spec ),
882- out_specs = q_partition_spec ,
883- check_rep = False ,
884- )
885- out = sharded_fn (query , key , value )
886- out = jax .lax .with_sharding_constraint (out , out_constraint )
887- return out
888-
889-
890- # ---------------------------------------------------------------------------
891- # TorchAX SDPA wrapper
892- # ---------------------------------------------------------------------------
893-
894-
895- def make_custom_splash_sdpa (mesh , env , ** kwargs ):
896- flash_block_sizes = kwargs .get ("flash_block_sizes" , None )
897- bq = kwargs .get ("block_q" , DEFAULT_BQSIZE )
898- bkv = kwargs .get ("block_kv" , DEFAULT_BKVSIZE )
899- bkv_compute = kwargs .get ("block_kv_compute" , DEFAULT_BKVCOMPUTESIZE )
900- bkv_compute_in = kwargs .get ("block_kv_compute_in" , DEFAULT_BKVCOMPUTEINSIZE )
901- hpt = kwargs .get ("heads_per_tile" , 1 )
902- use_k_smooth = kwargs .get ("use_k_smooth" , True )
903- use_base2_exp = kwargs .get ("use_base2_exp" , True )
904- use_experimental_scheduler = kwargs .get ("use_experimental_scheduler" , False )
905- vmem_limit_bytes = kwargs .get ("vmem_limit_bytes" , None )
906-
907- def _simple_attention (q , k , v , scale = None ):
908- s = scale if scale is not None else 1.0 / math .sqrt (q .shape [- 1 ])
909- attn = jnp .einsum ("bhsd,bhtd->bhst" , q * s , k )
910- attn = jax .nn .softmax (attn .astype (jnp .float32 ), axis = - 1 ).astype (q .dtype )
911- return jnp .einsum ("bhst,bhtd->bhsd" , attn , v )
912-
913- def _sdpa (
914- query ,
915- key ,
916- value ,
917- attn_mask = None ,
918- dropout_p = 0.0 ,
919- is_causal = False ,
920- scale = None ,
921- enable_gqa = False ,
922- ):
923- jquery , jkey , jvalue = env .t2j_iso ((query , key , value ))
924- num_heads = jquery .shape [1 ]
925-
926- if num_heads <= 8 :
927- result = _simple_attention (jquery , jkey , jvalue , scale = scale )
928- return env .j2t_iso (result )
929-
930- if use_k_smooth :
931- key_mean = jnp .mean (jkey , axis = 2 , keepdims = True )
932- jkey = jkey - key_mean
933-
934- result = tpu_custom_attention (
935- jquery ,
936- jkey ,
937- jvalue ,
938- mesh ,
939- scale = scale ,
940- block_q = bq ,
941- block_kv = bkv ,
942- block_kv_compute = bkv_compute ,
943- block_kv_compute_in = bkv_compute_in ,
944- heads_per_tile = hpt ,
945- use_base2_exp = use_base2_exp ,
946- use_experimental_scheduler = use_experimental_scheduler ,
947- vmem_limit_bytes = vmem_limit_bytes ,
948- flash_block_sizes = flash_block_sizes ,
949- )
950- return env .j2t_iso (result )
951-
952- return _sdpa
0 commit comments