Skip to content

ONNX Attention CUDA: Performance Optimization Opportunities vs Contrib GQA #28352

Description

@titaiwangms

ONNX Attention Op — Performance Gap vs Contrib Op GQA

Parent issue: #27516
Related: #28351 (spec compliance gaps)

After #27992 merged, ONNX Attention CUDA uses the same Flash/MEA kernels as Contrib GQA. However, decode latency is ~34-111μs/token higher due to spec-level constraints and missing optimizations.

Baseline benchmark recipe needed. The ~34-111μs/decode figure comes from informal profiling during #27992 development and lacks reproducible context (GPU model, dtype, batch×heads×head_size×seq_len, KV cache size, single-stream vs concurrent). PR #27992 description does not preserve the raw numbers. TODO: capture a reproducible benchmark recipe (script + machine spec + model config) and link it here before relying on these numbers for prioritization. See "Baseline Benchmark Recipe" checklist below.

Performance Gap Summary

# Area Impact Root Cause Code Location Mitigation Success Criteria
1 KV cache concat (decode w/ past) ~34-111μs/decode ONNX spec: immutable past tensors → full copy each step Flash decode concat: core/providers/cuda/llm/attention.cc:398 (LaunchConcatNewToPastKV) · MEA decode concat: attention.cc:680 · Source: contrib_ops/cuda/bert/attention_kv_cache.cu (LaunchConcatNewToPastKV impl) TensorScatter mode (external cache) bypasses this TensorScatter decode latency within ±5% of Contrib GQA decode at matched config
2 No XQA / DecoderMaskedMHA kernel 2-5x slower decode Requires shared-buffer KV cache (blocked by spec for past/present path) Dispatch site (TensorScatter Flash path): attention.cc:297-310 (Path 1: nonpad_kv_seqlenmha_fwd_kvcache) · XQA reference impl: contrib_ops/cuda/bert/xqa/xqa_loader.h and xqa/mha.h · GQA dispatch precedent: contrib_ops/cuda/bert/group_query_attention.cc:318-335 (XQA eligibility predicate) Add XQA dispatch tier for TensorScatter mode (S_q==1, sm≥80, no softcap, no smooth-softmax, no local window — same predicates as GQA) TensorScatter+XQA decode within ±10% of GQA+XQA decode at matched config
3 Extra concat kernel launch ~5-10μs Explicit concat before Flash call (separate from prep) Same call sites as item 1 (attention.cc:398, 680) Fuse into prep kernel (see item 4) Eliminate 1+ kernel launch per decode step (verify via Nsight Systems)
4 No fused prep kernel 2-3 extra launches Prep is QKV reshape + (optional) RoPE + KV concat as separate launches TODO marker: core/providers/cuda/llm/attention.cc:345-348 · GQA fused-prep precedent (LaunchUnpackRoPEAppend): contrib_ops/cuda/bert/group_query_attention_impl.cu:126, 589, 818, 903 · Description in group_query_attention_qdq.cuh:200 Fuse QKV reshape + KV concat (and RoPE if added later) mirroring GQA's UnpackRoPEAppend Reduce decode prep from 4-6 kernel launches → ~2 (per the existing TODO comment)

Reference implementations to study

  • GQA contrib op (perf baseline): onnxruntime/contrib_ops/cuda/bert/group_query_attention.cc (op skeleton, XQA eligibility) and onnxruntime/contrib_ops/cuda/bert/group_query_attention_impl.cu (LaunchUnpackRoPEAppend fused prep at lines 126, 589, 818, 903 — direct precedent for items 3 and 4).
  • XQA kernel: onnxruntime/contrib_ops/cuda/bert/xqa/xqa_loader.h is the public entry, mha.h declares the launch API. Loader specializations (e.g., xqa_loader_fp16.cu, xqa_loader_bf16_int8_256.cu) cover the supported dtype/quantization matrix.
  • DecoderMaskedMHA (legacy decode kernel): onnxruntime/contrib_ops/cuda/bert/fastertransformer_decoder_attention/decoder_masked_multihead_attention_impl.cu — earlier optimized decode kernel; useful for comparing assumptions XQA can/cannot relax.

When ONNX Attention already matches GQA:

  • TensorScatter mode (external cache via nonpad_kv_seqlen) — no concat, near-parity. Dispatch: attention.cc:297-310 (Flash) and attention.cc:540-560 comment block (MEA).
  • Prefill (large S_q) — Flash kernel dominates, overhead amortized.

Proposed Optimizations (priority order)

  1. XQA kernel for TensorScatter decode — The external-cache path already has shared-buffer semantics. Integrating XQA via the same eligibility predicate used by group_query_attention.cc:318-335 would close the 2-5x gap for the most performance-critical path. Reuse the GQA-side scratch allocation pattern (group_query_attention.cc:306-307).

  2. Fused KV cache preparation — Combine past_key/past_value concat + (future) RoPE into a single kernel modeled on LaunchUnpackRoPEAppend. Replaces the explicit LaunchConcatNewToPastKV calls at attention.cc:398, 680.

  3. In-place KV buffer reuse — If the graph optimizer can prove past_key/past_value have no other consumers downstream, reuse the buffer for present_key/present_value (avoid O(seq_len) copy). Requires a graph-transform pass; not a kernel change.

  4. Fused prep kernel — Consolidate QKV reshape, type conversion, and layout transforms into one kernel (TODO already noted at attention.cc:345-348). Subsumes item 2 once layout transforms are folded in.

Baseline Benchmark Recipe (TODO before relying on the table's μs numbers)

  • Pick a representative decode config (e.g., Llama-3 8B: B=1, H=32, H_kv=8, head_size=128, S_kv ramps 128→4096)
  • Capture GPU model, driver, CUDA version, ORT build flags
  • Script: standalone harness invoking ONNX Attention vs Contrib GQA at matched config, single-stream
  • Profile with Nsight Systems for per-kernel breakdown (attribute the μs to KV concat vs Flash kernel vs prep)
  • Repeat at fp16 and bf16
  • Commit script under onnxruntime/test/python/transformers/benchmark_*.py and link from this issue

Notes

  • The immutable-past constraint is fundamental to the ONNX spec design (functional semantics). Items 1-2 work within this constraint via the TensorScatter mode escape hatch. Item 3 requires runtime graph analysis.
  • For users who need maximum decode performance TODAY, recommend using TensorScatter mode (external cache management) — this is the only path that can ever reach Contrib GQA parity without spec changes.

Metadata

Metadata

Assignees

Labels

ep:CUDAissues related to the CUDA execution providermodel:transformerissues related to a transformer model: BERT, GPT2, Hugging Face, Longformer, T5, etc.

Type

No type

Fields

No fields configured for issues without a type.

Projects

No projects

Milestone

No milestone

Relationships

None yet

Development

No branches or pull requests

Issue actions