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
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. ๐Ÿ™‹ Ask for provider support