Skip to content

Keep the iOS18 fused scaled_dot_product_attention op when scale is given - #2864

Open
Yigtwxx wants to merge 1 commit into
apple:mainfrom
Yigtwxx:sdpa-scale-fused
Open

Yigtwxx wants to merge 1 commit into
apple:mainfrom
Yigtwxx:sdpa-scale-fused

Conversation

@Yigtwxx

@Yigtwxx Yigtwxx commented Sep 17, 2026

Copy link
Copy Markdown
Contributor

Summary

Since iOS 18 the torch scaled_dot_product_attention converter emits the fused Core ML scaled_dot_product_attention op, but only when torch's scale argument is left as None. Any explicit scale sends the op down the matmul / softmax / matmul decomposition instead, even when the value is just the default 1 / sqrt(head_dim).

In practice that means the fused op is almost never produced for transformer models coming out of Hugging Face transformers: sdpa_attention_forward (integrations/sdpa_attention.py) always calls SDPA with scale=scaling, where scaling is head_dim**-0.5 for Llama / Mistral / Qwen and query_pre_attn_scalar**-0.5 for Gemma 2 / Gemma 3. Converting a randomly initialised Llama attention layer with torch.export and minimum_deployment_target=ct.target.iOS18 gives, on main vs. this branch:

model fused scaled_dot_product_attention matmul softmax
Llama attention (scale=head_dim**-0.5), main 0 2 1
Llama attention, this PR 1 0 0
Gemma 2 attention (scale=query_pre_attn_scalar**-0.5), main 0 2 1
Gemma 2 attention, this PR 1 0 0

Minimal reproduction on main:

import torch, coremltools as ct
from torch.nn.functional import scaled_dot_product_attention as sdpa

class M(torch.nn.Module):
    def forward(self, q, k, v):
        return sdpa(q, k, v, scale=0.25)  # any explicit value, even 1 / sqrt(16)

q = torch.rand(1, 4, 8, 16)
prog = ct.convert(torch.jit.trace(M().eval(), (q, q, q)),
                  inputs=[ct.TensorType(name=n, shape=q.shape) for n in "qkv"],
                  convert_to="milinternal", minimum_deployment_target=ct.target.iOS18)
print([op.op_type for op in prog.functions["main"].operations if op.op_type != "const"])
# main:    ['mul', 'matmul', 'softmax', 'matmul']
# this PR: ['scaled_dot_product_attention']

Implementation

The fused op always applies 1 / sqrt(embed_size), so a custom scale is folded into the query beforehand:

sdpa(q, k, v, scale) == fused_sdpa(q * (scale * sqrt(embed_size)), k, v)
  • A new helper _fold_scale_into_query computes the factor when scale is a compile-time constant and the embedding size is static, emits a single mul on the query, and skips the mul entirely when the factor is (numerically) 1, which is the Hugging Face default. The factor is created in the query's dtype, mirroring what _decompose_scaled_dot_product_attention already does for the default scale in fp16.
  • When the fold is not possible (symbolic scale or symbolic embedding size) the converter falls back to the existing decomposition, so behaviour there is unchanged. Targets below iOS 18 are untouched.
  • The fold happens before the mask translation, so bool masks, is_causal, rank-2 inputs and enable_gqa combine with an explicit scale the same way they do without one.

Tests

test_scale is now parametrised over minimum_deployment_target in (None, iOS18) and scale in (1.5, 1 / sqrt(embedding_dim)). For iOS 18 it additionally checks the emitted program: exactly one fused scaled_dot_product_attention op, one mul for the custom scale and none for the default one. The existing pre-iOS18 cases still take the decomposition path.

Locally the emitted MIL programs were also evaluated op by op in NumPy and compared against torch eager (fp32) for embedding sizes 7 / 16 / 64, scales 0.25 / 1.5 / 1.0 / default, both TorchScript and torch.export, plus a symbolic batch and sequence length, causal, bool mask, float mask, rank-2 and grouped-query cases; the largest difference was 6e-6. TestScaledDotProductAttention shows no new failures compared to main.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant