ai.onnx.preview - FlexAttention¶
FlexAttention - 1 (ai.onnx.preview)¶
Version¶
name: FlexAttention (GitHub)
domain:
ai.onnx.previewsince_version:
1function:
Truesupport_level:
SupportType.EXPERIMENTALshape inference:
True
No versioning maintained for experimental ops.
Summary¶
Computes scaled dot-product attention over rank-4 (batched, multi-head) inputs, with optional user-provided customization subgraphs at two stages:
score_mod: Modify the attention score tensor after Q·K^T
prob_mod: Modify the probability tensor after Softmax
This operator mirrors the capabilities of PyTorch’s flex_attention: https://docs.pytorch.org/docs/stable/nn.attention.flex_attention.html
Input Shapes (MUST be rank-4 tensors):
Q:
(batch_size, q_num_heads, q_sequence_length, head_size)K:
(batch_size, kv_num_heads, kv_sequence_length, head_size)V:
(batch_size, kv_num_heads, kv_sequence_length, v_head_size)
Output Shape:
Y:
(batch_size, q_num_heads, q_sequence_length, v_head_size)
FlexAttention Computation:
Scores = (Q @ K^T) * scale
Scores = score_mod(Scores) # if 'score_mod' is provided
Probs = Softmax(Scores, axis=-1)
Probs = prob_mod(Probs) # if 'prob_mod' is provided
Y = Probs @ V
Grouped Query Attention (GQA):
When q_num_heads != kv_num_heads, each K/V head is shared by a contiguous
group of query heads in head-index order. Let
group_size = q_num_heads / kv_num_heads; then query head h uses K/V head
floor(h / group_size). q_num_heads must be a multiple of
kv_num_heads.
Modifier Subgraphs (score_mod, prob_mod): Each modifier subgraph takes exactly one rank-4 tensor input and must produce exactly one rank-4 tensor output of the same shape and element type.
score_mod input/output shape:
(batch_size, q_num_heads, q_sequence_length, kv_sequence_length)prob_mod input/output shape:
(batch_size, q_num_heads, q_sequence_length, kv_sequence_length)The element type is determined by softmax_precision (defaults to float32 for non-double inputs, otherwise double).
Masking can be expressed in score_mod by writing masked positions as -inf (or a large negative value appropriate for the target precision).
Attributes¶
prob_mod - GRAPH :
Optional probability modifier subgraph with 1 rank-4 tensor input and 1 rank-4 tensor output of the same shape and element type: (probs) -> probs_out. probs has softmax_precision element type and shape (B, Hq, L, S). The output must preserve the input shape.
scale - FLOAT :
Scaling factor for Q*K^T. Defaults to 1/sqrt(head_size).
score_mod - GRAPH :
Optional score modifier subgraph with 1 rank-4 tensor input and 1 rank-4 tensor output of the same shape and element type: (scores) -> scores_out. scores has softmax_precision element type and shape (B, Hq, L, S). The output must preserve the input shape.
softmax_precision - INT :
Floating-point precision for softmax computation. Defaults to float32 for non-double inputs, otherwise uses double. Must be explicitly specified for non-float types.
Inputs¶
Q (heterogeneous) - T1:
Query tensor with shape
(batch_size, q_num_heads, q_seq_len, head_size).K (heterogeneous) - T1:
Key tensor with shape
(batch_size, kv_num_heads, kv_seq_len, head_size).V (heterogeneous) - T1:
Value tensor with shape
(batch_size, kv_num_heads, kv_seq_len, v_head_size).
Outputs¶
Y (heterogeneous) - T1:
Output tensor with shape
(batch_size, q_num_heads, q_seq_len, v_head_size).
Type Constraints¶
T1 in (
tensor(bfloat16),tensor(double),tensor(float),tensor(float16)):Constrain Q, K, V to float tensors.
Examples¶
_flexattention¶
import numpy as np
import onnx
"""Basic FlexAttention test with default settings."""
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
domain=AI_ONNX_PREVIEW_DOMAIN,
)
B, Hq, L, E = 2, 4, 8, 16
S, Ev = 6, 16
Q = np.random.rand(B, Hq, L, E).astype(np.float32)
K = np.random.rand(B, Hq, S, E).astype(np.float32)
V = np.random.rand(B, Hq, S, Ev).astype(np.float32)
(Y,) = _compute_flex_attention(Q, K, V)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)
_flexattention_scaled¶
import numpy as np
import onnx
"""FlexAttention with explicit scale attribute."""
scale = 0.1
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
scale=scale,
domain=AI_ONNX_PREVIEW_DOMAIN,
)
B, Hq, L, E = 2, 4, 8, 16
S, Ev = 6, 16
Q = np.random.rand(B, Hq, L, E).astype(np.float32)
K = np.random.rand(B, Hq, S, E).astype(np.float32)
V = np.random.rand(B, Hq, S, Ev).astype(np.float32)
(Y,) = _compute_flex_attention(Q, K, V, scale=scale)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention_scaled",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)
_flexattention_gqa¶
import numpy as np
import onnx
"""FlexAttention with Grouped Query Attention (GQA)."""
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
domain=AI_ONNX_PREVIEW_DOMAIN,
)
B, Hq, Hkv, L, S, E, Ev = 2, 8, 2, 4, 6, 16, 16
Q = np.random.rand(B, Hq, L, E).astype(np.float32)
K = np.random.rand(B, Hkv, S, E).astype(np.float32)
V = np.random.rand(B, Hkv, S, Ev).astype(np.float32)
(Y,) = _compute_flex_attention(Q, K, V)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention_gqa",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)
_flexattention_diff_head_sizes¶
import numpy as np
import onnx
"""FlexAttention with different head sizes for Q/K vs V."""
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
domain=AI_ONNX_PREVIEW_DOMAIN,
)
B, Hq, L, E = 2, 4, 8, 16
S, Ev = 6, 32 # V has different head size
Q = np.random.rand(B, Hq, L, E).astype(np.float32)
K = np.random.rand(B, Hq, S, E).astype(np.float32)
V = np.random.rand(B, Hq, S, Ev).astype(np.float32)
(Y,) = _compute_flex_attention(Q, K, V)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention_diff_head_sizes",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)
_flexattention_score_mod¶
import numpy as np
import onnx
"""FlexAttention with score_mod subgraph (adds bias to scores)."""
bias_value = 0.5
score_mod_graph = _make_score_mod_bias_graph(bias_value, TensorProto.FLOAT)
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
domain=AI_ONNX_PREVIEW_DOMAIN,
)
# Add score_mod as a graph attribute
score_mod_attr = helper.make_attribute("score_mod", score_mod_graph)
node.attribute.append(score_mod_attr)
B, Hq, L, E = 1, 2, 3, 4
S, Ev = 3, 4
Q = np.random.rand(B, Hq, L, E).astype(np.float32)
K = np.random.rand(B, Hq, S, E).astype(np.float32)
V = np.random.rand(B, Hq, S, Ev).astype(np.float32)
scale = 1.0 / np.sqrt(E)
scores = np.einsum("bhle,bhse->bhls", Q, K) * scale
scores = scores + bias_value # score_mod: add bias
probs = np.exp(scores - scores.max(axis=-1, keepdims=True))
probs = probs / probs.sum(axis=-1, keepdims=True)
Y = np.einsum("bhls,bhsv->bhlv", probs, V).astype(np.float32)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention_score_mod",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)
_flexattention_prob_mod¶
import numpy as np
import onnx
"""FlexAttention with prob_mod subgraph (scales probabilities)."""
scale_value = 0.5
prob_mod_graph = _make_prob_mod_scale_graph(scale_value, TensorProto.FLOAT)
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
domain=AI_ONNX_PREVIEW_DOMAIN,
)
prob_mod_attr = helper.make_attribute("prob_mod", prob_mod_graph)
node.attribute.append(prob_mod_attr)
B, Hq, L, E = 1, 2, 3, 4
S, Ev = 3, 4
Q = np.random.rand(B, Hq, L, E).astype(np.float32)
K = np.random.rand(B, Hq, S, E).astype(np.float32)
V = np.random.rand(B, Hq, S, Ev).astype(np.float32)
scale = 1.0 / np.sqrt(E)
scores = np.einsum("bhle,bhse->bhls", Q, K) * scale
probs = np.exp(scores - scores.max(axis=-1, keepdims=True))
probs = probs / probs.sum(axis=-1, keepdims=True)
probs = probs * scale_value
Y = np.einsum("bhls,bhsv->bhlv", probs, V).astype(np.float32)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention_prob_mod",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)
_flexattention_fp16¶
import numpy as np
import onnx
"""FlexAttention with float16 inputs."""
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
domain=AI_ONNX_PREVIEW_DOMAIN,
)
B, Hq, L, E = 2, 4, 8, 16
S, Ev = 6, 16
Q = np.random.rand(B, Hq, L, E).astype(np.float16)
K = np.random.rand(B, Hq, S, E).astype(np.float16)
V = np.random.rand(B, Hq, S, Ev).astype(np.float16)
(Y,) = _compute_flex_attention(Q, K, V)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention_fp16",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)
_flexattention_double¶
import numpy as np
import onnx
"""FlexAttention with double precision inputs."""
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
domain=AI_ONNX_PREVIEW_DOMAIN,
)
B, Hq, L, E = 2, 4, 8, 16
S, Ev = 6, 16
Q = np.random.rand(B, Hq, L, E).astype(np.float64)
K = np.random.rand(B, Hq, S, E).astype(np.float64)
V = np.random.rand(B, Hq, S, Ev).astype(np.float64)
(Y,) = _compute_flex_attention(Q, K, V)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention_double",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)
_flexattention_causal_mask¶
import numpy as np
import onnx
"""FlexAttention with causal masking score_mod (Qwen-3, Gemma-3, Llama-3 pattern)."""
score_mod_graph = _make_score_mod_causal_mask_graph(TensorProto.FLOAT)
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
domain=AI_ONNX_PREVIEW_DOMAIN,
)
score_mod_attr = helper.make_attribute("score_mod", score_mod_graph)
node.attribute.append(score_mod_attr)
B, Hq, L, E = 1, 2, 4, 8
S, Ev = 4, 8
Q = np.random.rand(B, Hq, L, E).astype(np.float32)
K = np.random.rand(B, Hq, S, E).astype(np.float32)
V = np.random.rand(B, Hq, S, Ev).astype(np.float32)
# Manually compute expected output with causal masking
scale = 1.0 / np.sqrt(E)
scores = np.einsum("bhle,bhse->bhls", Q, K) * scale
# Apply causal mask: set future positions to -inf
q_idx = np.arange(L).reshape(1, 1, L, 1)
k_idx = np.arange(S).reshape(1, 1, 1, S)
mask = q_idx >= k_idx
scores = np.where(mask, scores, -np.inf)
probs = np.exp(scores - scores.max(axis=-1, keepdims=True))
probs = probs / probs.sum(axis=-1, keepdims=True)
Y = np.einsum("bhls,bhsv->bhlv", probs, V).astype(np.float32)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention_causal_mask",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)
_flexattention_soft_cap¶
import numpy as np
import onnx
"""FlexAttention with soft capping score_mod (Gemma-2 pattern)."""
cap_value = 20.0
score_mod_graph = _make_score_mod_soft_cap_graph(cap_value, TensorProto.FLOAT)
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
domain=AI_ONNX_PREVIEW_DOMAIN,
)
score_mod_attr = helper.make_attribute("score_mod", score_mod_graph)
node.attribute.append(score_mod_attr)
B, Hq, L, E = 1, 2, 4, 8
S, Ev = 4, 8
Q = np.random.rand(B, Hq, L, E).astype(np.float32)
K = np.random.rand(B, Hq, S, E).astype(np.float32)
V = np.random.rand(B, Hq, S, Ev).astype(np.float32)
# Manually compute expected output with soft capping
scale = 1.0 / np.sqrt(E)
scores = np.einsum("bhle,bhse->bhls", Q, K) * scale
scores = np.tanh(scores / cap_value) * cap_value
probs = np.exp(scores - scores.max(axis=-1, keepdims=True))
probs = probs / probs.sum(axis=-1, keepdims=True)
Y = np.einsum("bhls,bhsv->bhlv", probs, V).astype(np.float32)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention_soft_cap",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)
_flexattention_relative_positional¶
import numpy as np
import onnx
"""FlexAttention with relative positional bias score_mod."""
score_mod_graph = _make_score_mod_relative_positional_graph(TensorProto.FLOAT)
node = helper.make_node(
"FlexAttention",
inputs=["Q", "K", "V"],
outputs=["Y"],
domain=AI_ONNX_PREVIEW_DOMAIN,
)
score_mod_attr = helper.make_attribute("score_mod", score_mod_graph)
node.attribute.append(score_mod_attr)
B, Hq, L, E = 1, 2, 4, 8
S, Ev = 4, 8
Q = np.random.rand(B, Hq, L, E).astype(np.float32)
K = np.random.rand(B, Hq, S, E).astype(np.float32)
V = np.random.rand(B, Hq, S, Ev).astype(np.float32)
# Manually compute expected output with relative positional bias
scale = 1.0 / np.sqrt(E)
scores = np.einsum("bhle,bhse->bhls", Q, K) * scale
q_idx = np.arange(L).reshape(-1, 1)
k_idx = np.arange(S).reshape(1, -1)
rel_pos = (q_idx - k_idx).astype(np.float32)
scores = scores + rel_pos
probs = np.exp(scores - scores.max(axis=-1, keepdims=True))
probs = probs / probs.sum(axis=-1, keepdims=True)
Y = np.einsum("bhls,bhsv->bhlv", probs, V).astype(np.float32)
expect(
node,
inputs=[Q, K, V],
outputs=[Y],
name="test_flexattention_relative_positional",
opset_imports=[
helper.make_opsetid("", 26),
helper.make_opsetid(AI_ONNX_PREVIEW_DOMAIN, 1),
],
)