Skip to content

Commit fcefd14

Browse files
committed
Cleaning up custom_splash_attention
1 parent 4bb8821 commit fcefd14

5 files changed

Lines changed: 21 additions & 242 deletions

File tree

src/maxdiffusion/kernels/custom_splash_attention.py

Lines changed: 6 additions & 226 deletions
Original file line numberDiff line numberDiff line change
@@ -17,38 +17,28 @@
1717
"""Custom Pallas flash attention kernel for TPU."""
1818

1919
import functools
20-
import math
2120

2221
import jax
2322
import jax.numpy as jnp
2423
import numpy as np
2524
from jax import lax
2625
from jax.experimental import pallas as pl
2726
from 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

3128
DEFAULT_MASK_VALUE = -0.7 * float(np.finfo(np.dtype("float32")).max)
3229
NUM_LANES = 128
3330
NUM_SUBLANES = 8
3431
NT_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

4534
class _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

712692
def 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

Comments
 (0)