|
17 | 17 | """Custom Pallas flash attention kernel for TPU.""" |
18 | 18 |
|
19 | 19 | import functools |
20 | | -import math |
21 | 20 |
|
22 | 21 | import jax |
23 | 22 | import jax.numpy as jnp |
24 | 23 | import numpy as np |
25 | 24 | from jax import lax |
26 | 25 | from jax.experimental import pallas as pl |
27 | 26 | 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 |
30 | 27 |
|
31 | 28 | DEFAULT_MASK_VALUE = -0.7 * float(np.finfo(np.dtype("float32")).max) |
32 | 29 | NUM_LANES = 128 |
33 | 30 | NUM_SUBLANES = 8 |
34 | 31 | NT_DIM_NUMBERS = (((1,), (1,)), ((), ())) |
35 | 32 |
|
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 | 33 | # Kept at 256 to protect VPU registers (V1 Optimization) |
42 | 34 | DEFAULT_BKVCOMPUTEINSIZE = 256 |
43 | 35 |
|
@@ -569,227 +561,3 @@ def _splash_attention(q, k, v): |
569 | 561 | return _splash_attention |
570 | 562 |
|
571 | 563 |
|
572 | | -# --------------------------------------------------------------------------- |
573 | | -# High-level attention function with shard_map |
574 | | -# TODO: Move `tpu_custom_attention` and `make_custom_splash_sdpa` |
575 | | -# out of this file and into `attention_flax.py` where other orchestration |
576 | | -# functions live. This file should be strictly Pallas kernel definitions. |
577 | | -# --------------------------------------------------------------------------- |
578 | | - |
579 | | -def tpu_custom_attention( |
580 | | - query, |
581 | | - key, |
582 | | - value, |
583 | | - mesh, |
584 | | - *, |
585 | | - scale=None, |
586 | | - block_q=None, |
587 | | - block_kv=None, |
588 | | - block_kv_compute=None, |
589 | | - block_kv_compute_in=None, |
590 | | - heads_per_tile=None, |
591 | | - use_base2_exp=True, |
592 | | - use_experimental_scheduler=False, |
593 | | - vmem_limit_bytes=None, |
594 | | - flash_block_sizes=None, |
595 | | -): |
596 | | - _LOG2_E = 1.44269504 |
597 | | - num_heads = query.shape[1] |
598 | | - |
599 | | - if flash_block_sizes is not None: |
600 | | - block_q = flash_block_sizes.get("block_q", block_q) |
601 | | - block_kv = flash_block_sizes.get("block_kv", block_kv) |
602 | | - block_kv_compute = flash_block_sizes.get("block_kv_compute", block_kv_compute) |
603 | | - block_kv_compute_in = flash_block_sizes.get("block_kv_compute_in", block_kv_compute_in) |
604 | | - heads_per_tile = flash_block_sizes.get("heads_per_tile", heads_per_tile) |
605 | | - vmem_limit_bytes = flash_block_sizes.get("vmem_limit_bytes", vmem_limit_bytes) |
606 | | - |
607 | | - block_q = block_q if block_q is not None else DEFAULT_BQSIZE |
608 | | - block_kv = block_kv if block_kv is not None else DEFAULT_BKVSIZE |
609 | | - block_kv_compute = block_kv_compute if block_kv_compute is not None else DEFAULT_BKVCOMPUTESIZE |
610 | | - block_kv_compute_in = block_kv_compute_in if block_kv_compute_in is not None else DEFAULT_BKVCOMPUTEINSIZE |
611 | | - heads_per_tile = heads_per_tile if heads_per_tile is not None else 1 |
612 | | - |
613 | | - def _attention_on_slices(q, k, v): |
614 | | - scale_factor = 1.0 / math.sqrt(q.shape[-1]) if scale is None else scale |
615 | | - if use_base2_exp: |
616 | | - q = q * scale_factor * _LOG2_E |
617 | | - else: |
618 | | - q = q * scale_factor |
619 | | - |
620 | | - def _pad_to_multiple(x, multiple, axis): |
621 | | - seq_len = x.shape[axis] |
622 | | - pad_len = (multiple - seq_len % multiple) % multiple |
623 | | - if pad_len == 0: |
624 | | - return x, seq_len |
625 | | - pad_width = [(0, 0)] * x.ndim |
626 | | - pad_width[axis] = (0, pad_len) |
627 | | - return jnp.pad(x, pad_width), seq_len |
628 | | - |
629 | | - def _kernel_3d(q_3d, k_3d, v_3d): |
630 | | - q_orig_len = q_3d.shape[1] |
631 | | - kv_orig_len = k_3d.shape[1] |
632 | | - |
633 | | - q_3d_padded, _ = _pad_to_multiple(q_3d, block_q, axis=1) |
634 | | - k_3d_padded, _ = _pad_to_multiple(k_3d, block_kv, axis=1) |
635 | | - v_3d_padded, _ = _pad_to_multiple(v_3d, block_kv, axis=1) |
636 | | - |
637 | | - padded_q_seq_len = q_3d_padded.shape[1] |
638 | | - padded_kv_seq_len = k_3d_padded.shape[1] |
639 | | - |
640 | | - bsizes = _BlockSizes( |
641 | | - block_q=min(block_q, padded_q_seq_len), |
642 | | - block_kv=min(block_kv, padded_kv_seq_len), |
643 | | - block_kv_compute=min(block_kv_compute, padded_kv_seq_len), |
644 | | - ) |
645 | | - splash_kernel = make_splash_mha( |
646 | | - block_sizes=bsizes, |
647 | | - bkv_compute_in=block_kv_compute_in, |
648 | | - orig_q_seq_len=q_orig_len, |
649 | | - orig_kv_seq_len=kv_orig_len, |
650 | | - heads_per_tile=heads_per_tile, |
651 | | - use_base2_exp=use_base2_exp, |
652 | | - use_experimental_scheduler=use_experimental_scheduler, |
653 | | - vmem_limit_bytes=vmem_limit_bytes, |
654 | | - ) |
655 | | - out = splash_kernel( |
656 | | - q_3d_padded.astype(jnp.bfloat16), |
657 | | - k_3d_padded, |
658 | | - v_3d_padded, |
659 | | - ) |
660 | | - out = jnp.swapaxes(out, 1, 2) |
661 | | - return out[:, :q_orig_len, ...] |
662 | | - |
663 | | - return jax.vmap(_kernel_3d, in_axes=(0, 0, 0), out_axes=0)(q, k, v) |
664 | | - |
665 | | - batch_size = query.shape[0] |
666 | | - if num_heads < mesh.size: |
667 | | - q_partition_spec = P() |
668 | | - kv_partition_spec = P() |
669 | | - out_constraint = P() |
670 | | - else: |
671 | | - # Name-based axis resolution for the standard 4D mesh ("data", "fsdp", "context", "tensor"). |
672 | | - # Axes are grouped by semantic role: |
673 | | - # - Batch-parallel: "data", "fsdp" (shard the batch dimension) |
674 | | - # - Head-parallel: "context", "tensor" (shard the head dimension) |
675 | | - # Only axes with size > 1 participate in sharding to avoid shard_map errors. |
676 | | - # Q and KV are sharded identically (symmetric head-parallel), matching the |
677 | | - # strategy that the legacy 2D (dp, tp) path used. |
678 | | - batch_axes = tuple(a for a in ("data", "fsdp") if a in mesh.shape and mesh.shape[a] > 1) |
679 | | - head_axes = tuple(a for a in ("context", "tensor") if a in mesh.shape and mesh.shape[a] > 1) |
680 | | - |
681 | | - batch_parallel_size = 1 |
682 | | - for a in batch_axes: |
683 | | - batch_parallel_size *= mesh.shape[a] |
684 | | - head_parallel_size = 1 |
685 | | - for a in head_axes: |
686 | | - head_parallel_size *= mesh.shape[a] |
687 | | - |
688 | | - # Flatten singleton tuples for clean PartitionSpecs. |
689 | | - batch_spec = batch_axes[0] if len(batch_axes) == 1 else (batch_axes if batch_axes else None) |
690 | | - head_spec = head_axes[0] if len(head_axes) == 1 else (head_axes if head_axes else None) |
691 | | - |
692 | | - if ( |
693 | | - batch_parallel_size > 1 |
694 | | - and batch_size >= batch_parallel_size |
695 | | - and num_heads % head_parallel_size == 0 |
696 | | - ): |
697 | | - # Head-parallel: batch on data/fsdp, heads on context/tensor. |
698 | | - q_partition_spec = P(batch_spec, head_spec, None, None) |
699 | | - kv_partition_spec = P(batch_spec, head_spec, None, None) |
700 | | - out_constraint = P(batch_spec, None, head_spec, None) |
701 | | - elif head_parallel_size > 1 or batch_parallel_size > 1: |
702 | | - # Batch too small for batch-parallel; shard all parallelism onto heads. |
703 | | - all_axes = batch_axes + head_axes |
704 | | - all_parallel_size = batch_parallel_size * head_parallel_size |
705 | | - all_spec = all_axes[0] if len(all_axes) == 1 else (all_axes if all_axes else None) |
706 | | - if num_heads % all_parallel_size == 0: |
707 | | - q_partition_spec = P(None, all_spec, None, None) |
708 | | - kv_partition_spec = P(None, all_spec, None, None) |
709 | | - out_constraint = P(None, None, all_spec, None) |
710 | | - else: |
711 | | - # Heads not divisible by available parallelism; fall back to replicated. |
712 | | - q_partition_spec = P() |
713 | | - kv_partition_spec = P() |
714 | | - out_constraint = P() |
715 | | - else: |
716 | | - # All axes are size-1; fully replicated (e.g. single-device). |
717 | | - q_partition_spec = P() |
718 | | - kv_partition_spec = P() |
719 | | - out_constraint = P() |
720 | | - |
721 | | - sharded_fn = shard_map( |
722 | | - _attention_on_slices, |
723 | | - mesh=mesh, |
724 | | - in_specs=(q_partition_spec, kv_partition_spec, kv_partition_spec), |
725 | | - out_specs=q_partition_spec, |
726 | | - check_rep=False, |
727 | | - ) |
728 | | - out = sharded_fn(query, key, value) |
729 | | - out = jax.lax.with_sharding_constraint(out, out_constraint) |
730 | | - return out |
731 | | - |
732 | | - |
733 | | -# --------------------------------------------------------------------------- |
734 | | -# TorchAX SDPA wrapper |
735 | | -# --------------------------------------------------------------------------- |
736 | | - |
737 | | - |
738 | | -def make_custom_splash_sdpa(mesh, env, **kwargs): |
739 | | - flash_block_sizes = kwargs.get("flash_block_sizes", None) |
740 | | - bq = kwargs.get("block_q", DEFAULT_BQSIZE) |
741 | | - bkv = kwargs.get("block_kv", DEFAULT_BKVSIZE) |
742 | | - bkv_compute = kwargs.get("block_kv_compute", DEFAULT_BKVCOMPUTESIZE) |
743 | | - bkv_compute_in = kwargs.get("block_kv_compute_in", DEFAULT_BKVCOMPUTEINSIZE) |
744 | | - hpt = kwargs.get("heads_per_tile", 1) |
745 | | - use_k_smooth = kwargs.get("use_k_smooth", True) |
746 | | - use_base2_exp = kwargs.get("use_base2_exp", True) |
747 | | - use_experimental_scheduler = kwargs.get("use_experimental_scheduler", False) |
748 | | - vmem_limit_bytes = kwargs.get("vmem_limit_bytes", None) |
749 | | - |
750 | | - def _simple_attention(q, k, v, scale=None): |
751 | | - s = scale if scale is not None else 1.0 / math.sqrt(q.shape[-1]) |
752 | | - attn = jnp.einsum("bhsd,bhtd->bhst", q * s, k) |
753 | | - attn = jax.nn.softmax(attn.astype(jnp.float32), axis=-1).astype(q.dtype) |
754 | | - return jnp.einsum("bhst,bhtd->bhsd", attn, v) |
755 | | - |
756 | | - def _sdpa( |
757 | | - query, |
758 | | - key, |
759 | | - value, |
760 | | - attn_mask=None, |
761 | | - dropout_p=0.0, |
762 | | - is_causal=False, |
763 | | - scale=None, |
764 | | - enable_gqa=False, |
765 | | - ): |
766 | | - jquery, jkey, jvalue = env.t2j_iso((query, key, value)) |
767 | | - num_heads = jquery.shape[1] |
768 | | - |
769 | | - if num_heads <= 8: |
770 | | - result = _simple_attention(jquery, jkey, jvalue, scale=scale) |
771 | | - return env.j2t_iso(result) |
772 | | - |
773 | | - if use_k_smooth: |
774 | | - key_mean = jnp.mean(jkey, axis=2, keepdims=True) |
775 | | - jkey = jkey - key_mean |
776 | | - |
777 | | - result = tpu_custom_attention( |
778 | | - jquery, |
779 | | - jkey, |
780 | | - jvalue, |
781 | | - mesh, |
782 | | - scale=scale, |
783 | | - block_q=bq, |
784 | | - block_kv=bkv, |
785 | | - block_kv_compute=bkv_compute, |
786 | | - block_kv_compute_in=bkv_compute_in, |
787 | | - heads_per_tile=hpt, |
788 | | - use_base2_exp=use_base2_exp, |
789 | | - use_experimental_scheduler=use_experimental_scheduler, |
790 | | - vmem_limit_bytes=vmem_limit_bytes, |
791 | | - flash_block_sizes=flash_block_sizes, |
792 | | - ) |
793 | | - return env.j2t_iso(result) |
794 | | - |
795 | | - return _sdpa |
0 commit comments