YAML Metadata Warning:empty or missing yaml metadata in repo card
Check out the documentation for more information.
pytorch-mps-flash-sdpa
Pre-built PyTorch wheel with MPS FlashAttention wiring (pytorch/pytorch PR #198564).
Fixes NotImplementedError when sdpa_kernel(FLASH_ATTENTION) or _fused_sdp_choice is called on MPS tensors. Routes SDPBackend::flash_attention through the existing sdpa_prefill_mps Metal kernel on Apple Silicon.
Requirements
- Apple Silicon Mac (M1 or later)
- macOS 14.0+
- Python 3.14
Install
pip install "https://huggingface.co/EvanOLeary/pytorch-mps-flash-sdpa/resolve/main/torch-2.15.0a0-cp314-cp314-macosx_14_0_arm64.whl"
Quick test
import torch
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
# ESMC-300M config
q = torch.randn(1, 15, 512, 64, device="mps", dtype=torch.float16)
k, v = torch.randn_like(q), torch.randn_like(q)
# This used to raise NotImplementedError on MPS โ now works
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
out = F.scaled_dot_product_attention(q, k, v)
print(out.shape) # torch.Size([1, 15, 512, 64])
# Check dispatch
choice = torch.ops.aten._fused_sdp_choice(q, k, v)
print(SDPBackend(choice).name) # FLASH_ATTENTION
What's fixed
| Before | After | |
|---|---|---|
_fused_sdp_choice on MPS |
NotImplementedError |
Returns FLASH_ATTENTION for supported dims |
sdpa_kernel(FLASH_ATTENTION) on MPS |
NotImplementedError |
Works |
| Supported head dims | โ | 32, 64, 72, 80, 96, 128, 256 |
| ESMC-300M (D=64) | broken | โ |
| ESMC-600M (D=72) | broken | โ |
Build info
- Base:
pytorch/pytorch@199627e - Commits:
9e09d92(wiring) +27181af(use_mpp fix + tests) - Hardware: Apple M3 8GB, macOS 14.4.1
- Python: 3.14.7
Inference Providers NEW
This model isn't deployed by any Inference Provider. ๐ Ask for provider support