Skip to content

Honor the dtype keyword of zeros_like for torch.export (inverted MultiheadAttention masks) - #2865

Open
cdeil wants to merge 1 commit into
apple:mainfrom
cdeil:fix-zeros-like-dtype-kwarg
Open

cdeil wants to merge 1 commit into
apple:mainfrom
cdeil:fix-zeros-like-dtype-kwarg

Conversation

@cdeil

@cdeil cdeil commented Sep 23, 2026

Copy link
Copy Markdown
Contributor

Summary

torch.export and ExecuTorch pass zeros_like's dtype as a keyword argument. The torch frontend reads it only positionally (the TorchScript form), so torch.zeros_like(x, dtype=...) keeps the dtype of x.

This matters most in nn.MultiheadAttention, which turns a bool attn_mask into a float mask with torch.zeros_like(mask, dtype=q.dtype).masked_fill_(mask, float("-inf")) (F._canonical_mask). With bool zeros:

  1. masked_fill yields a bool tensor, because it casts the fill value to the tensor's dtype.
  2. scaled_dot_product_attention then treats that bool tensor with SDPA's bool semantics (True = attend), the opposite of the float mask it replaced.

For a causal mask the converted model attends to exactly the wrong positions. For a mask that blocks nothing, every logit gets a −3e4 bias. That cancels in exact arithmetic but costs precision, and float16 softmax on the CPU loses most of it.

import numpy as np, torch, coremltools as ct

torch.manual_seed(0)


class M(torch.nn.Module):
    def __init__(self, mask):
        super().__init__()
        self.attn = torch.nn.MultiheadAttention(16, 2, batch_first=True)
        self.register_buffer("mask", mask)

    def forward(self, x):
        return self.attn(x, x, x, attn_mask=self.mask, need_weights=False)[0]


x = torch.randn(1, 6, 16)
for name, mask in [
    ("none blocked", torch.zeros(6, 6, dtype=torch.bool)),
    ("causal", torch.triu(torch.ones(6, 6, dtype=torch.bool), 1)),
]:
    m = M(mask).eval()
    ep = torch.export.export(m, (x,)).run_decompositions({})
    ml = ct.convert(
        ep,
        minimum_deployment_target=ct.target.iOS16,
        compute_precision=ct.precision.FLOAT32,
        compute_units=ct.ComputeUnit.CPU_ONLY,
    )
    out = list(
        ml.predict({ml.get_spec().description.input[0].name: x.numpy()}).values()
    )[0]
    print(name, np.abs(out - m(x).detach().numpy()).max())

On main (and 9.0, 9.1.dev1) this prints about 2.6e-04 and 1.18. With this change it prints 1.0e-07 and 1.2e-07. In float16, a detection transformer whose keypoint head uses such an all-False mask lost all its detections on the CPU, because the softmax inputs sat near −30000, where float16 values are 16 apart. It recovers with this change.

zeros_like now reads the keyword the same way ones_like already does.

Testing

pytest coremltools/converters/mil/frontend/torch/test/test_torch_ops.py \
  -k "test_zeros_like or test_multihead_attention_bool_attn_mask or TestOnesLike or TestMaskedFill or TestScaledDotProductAttention or TestTransformer"
  • TestZeros::test_zeros_like_dtype_from_bool uses the result in arithmetic. The existing test_zeros_like_types already exercises the keyword under torch.export, but zeros compare equal across dtypes, so it passes without the fix.
  • TestTransformer::test_multihead_attention_bool_attn_mask covers a causal mask and one that blocks nothing, for the torch.export frontends; torch.jit.trace fails its own sanity check on nn.MultiheadAttention.

Both tests fail before the change (value mismatches of 6.0 and 0.63) and pass after it.

torch.export and ExecuTorch pass zeros_like's dtype as a keyword argument,
but the converter only read it positionally (the TorchScript form), so the
result kept the input's dtype. nn.MultiheadAttention builds its float
attention mask as zeros_like(bool_mask, dtype=q.dtype).masked_fill(bool_mask,
-inf) (F._canonical_mask). With bool zeros the masked_fill result stays bool,
and scaled_dot_product_attention then reads it with the opposite bool
semantics: a causal mask is inverted (max abs error 1.18 against PyTorch in
float32), and a mask that blocks nothing becomes a -3e4 bias on every logit,
which costs float16 softmax most of its precision.

Read the keyword the same way ones_like already does. The existing
test_zeros_like_types passed despite the bug because zeros compare equal
across dtypes; the new tests use the result in arithmetic.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@TobyRoseman

Copy link
Copy Markdown
Collaborator

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

2 participants