Skip to content

[Bug] FusedAttention hangs with long causal KV under CP all_gather on sm103 / cuDNN 9.21 #3285

Description

@Baibaifan

Describe the bug

transformer_engine FusedAttention hangs for a BF16 long-KV causal attention shape when using the cuDNN fused attention backend NVTE_F16_arbitrary_seqlen.

The hang happens after the fused attention forward is launched. The Python process blocks at torch.cuda.synchronize(), while GPU utilization stays at 100%.

Failing shape:

  • dtype: bf16
  • q shape: [1, 8192, 16, 128]
  • k shape: [1, 131072, 8, 128]
  • v shape: [1, 131072, 8, 128]
  • layout: bshd_bshd_bshd
  • mask: causal_bottom_right
  • dropout: 0.0
  • backend: NVTE_F16_arbitrary_seqlen

A smaller control shape with num_heads=8, num_gqa_groups=4, head_dim=64 completes successfully.

Steps/Code to reproduce bug

Run the attached single-GPU reproducer:

CUDA_VISIBLE_DEVICES=0 python repro_single_fused_attention.py \
  --te-source /path/to/TransformerEngine/transformerengine \
  --q-len 8192 \
  --kv-len 131072 \
  --num-heads 16 \
  --num-gqa-groups 8 \
  --head-dim 128 \
  --te-debug
#!/usr/bin/env python3
"""Single-GPU Transformer Engine FusedAttention reproducer.
"""

from __future__ import annotations

import argparse
import math
import os
import sys
import time
from pathlib import Path
from typing import Optional, Tuple


def parse_window_size(value: str) -> Tuple[int, int]:
    parts = value.split(",")
    if len(parts) != 2:
        raise argparse.ArgumentTypeError("window size must be 'left,right'")
    return int(parts[0]), int(parts[1])


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Run a single-GPU TE fused_attn_fwd call with long KV."
    )
    parser.add_argument("--te-source", type=str, default=None, help="Optional TE source root.")
    parser.add_argument("--device", type=int, default=0)
    parser.add_argument("--mode", choices=["single", "cp-two-step"], default="single")
    parser.add_argument("--q-len", type=int, default=8192)
    parser.add_argument("--kv-len", type=int, default=131072)
    parser.add_argument("--short-kv-len", type=int, default=8192)
    parser.add_argument("--num-heads", type=int, default=16)
    parser.add_argument("--num-gqa-groups", type=int, default=8)
    parser.add_argument("--head-dim", type=int, default=128)
    parser.add_argument("--dtype", choices=["bf16", "fp16"], default="bf16")
    parser.add_argument(
        "--attn-mask-type",
        choices=["causal", "causal_bottom_right", "no_mask"],
        default="causal_bottom_right",
    )
    parser.add_argument("--window-size", type=parse_window_size, default=(-1, 0))
    parser.add_argument("--iters", type=int, default=1)
    parser.add_argument("--training", action="store_true")
    parser.add_argument("--init", choices=["empty", "zeros", "normal"], default="zeros")
    parser.add_argument("--seed", type=int, default=1234)
    parser.add_argument("--te-debug", action="store_true")
    return parser.parse_args()


def normalize_te_source(path: Optional[str]) -> Optional[str]:
    if not path:
        return None
    src = Path(path).expanduser().resolve()
    candidates = [
        src,
        src / "transformerengine",
        src / "transformerengine" / "transformer_engine",
    ]
    for candidate in candidates:
        if (candidate / "transformer_engine" / "__init__.py").exists():
            return str(candidate)
        if candidate.name == "transformer_engine" and (candidate / "__init__.py").exists():
            return str(candidate.parent)
    return str(src)


def configure_env(args: argparse.Namespace) -> None:
    os.environ.setdefault("NVTE_FLASH_ATTN", "0")
    os.environ.setdefault("NVTE_FUSED_ATTN", "1")
    os.environ.setdefault("NVTE_UNFUSED_ATTN", "0")
    os.environ.setdefault("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")
    if args.te_debug:
        os.environ["NVTE_DEBUG"] = "1"
        os.environ["NVTE_DEBUG_LEVEL"] = "2"

    te_source = normalize_te_source(args.te_source)
    if te_source:
        sys.path.insert(0, te_source)


def make_tensor(
    shape: Tuple[int, ...],
    dtype,
    device,
    init: str,
    requires_grad: bool,
) -> "torch.Tensor":
    import torch

    tensor = torch.empty(shape, dtype=dtype, device=device)
    if init == "zeros":
        tensor.zero_()
    elif init == "normal":
        tensor.normal_(mean=0.0, std=0.02)
    if requires_grad:
        tensor.requires_grad_(True)
    return tensor


def run_fused_call(
    *,
    label: str,
    q_len: int,
    kv_len: int,
    args: argparse.Namespace,
    dtype,
    device,
    fused_attn_fwd,
    tex,
):
    import torch

    q = make_tensor(
        (1, q_len, args.num_heads, args.head_dim),
        dtype,
        device,
        args.init,
        args.training,
    )
    k = make_tensor(
        (1, kv_len, args.num_gqa_groups, args.head_dim),
        dtype,
        device,
        args.init,
        args.training,
    )
    v = make_tensor(
        (1, kv_len, args.num_gqa_groups, args.head_dim),
        dtype,
        device,
        args.init,
        args.training,
    )
    cu_seqlens_q = torch.tensor([0, q_len], dtype=torch.int32, device=device)
    cu_seqlens_kv = torch.tensor([0, kv_len], dtype=torch.int32, device=device)

    print(
        f"{label}: before fused_attn_fwd "
        f"q={tuple(q.shape)} k={tuple(k.shape)} v={tuple(v.shape)} "
        f"mask={args.attn_mask_type} window={args.window_size}",
        flush=True,
    )
    start = time.perf_counter()
    out, aux = fused_attn_fwd(
        args.training,
        q_len,
        kv_len,
        cu_seqlens_q,
        cu_seqlens_kv,
        q,
        k,
        v,
        dtype,
        tex.NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen,
        attn_scale=1.0 / math.sqrt(args.head_dim),
        dropout=0.0,
        qkv_layout="bshd_bshd_bshd",
        o_format="bshd",
        attn_mask_type=args.attn_mask_type,
        attn_bias_type="no_bias",
        attn_bias=None,
        window_size=args.window_size,
        return_max_logit=False,
        cuda_graph=False,
    )
    print(
        f"{label}: host returned after {time.perf_counter() - start:.3f}s, "
        f"out={tuple(out.shape)}, aux={len(aux)}; before cuda sync",
        flush=True,
    )
    return out, aux


def main() -> int:
    args = parse_args()
    configure_env(args)

    import torch
    import transformer_engine
    import transformer_engine_torch as tex
    from transformer_engine.pytorch.cpp_extensions.fused_attn import fused_attn_fwd

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required for this reproducer.")
    torch.cuda.set_device(args.device)
    device = torch.device("cuda", args.device)
    dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16
    torch.manual_seed(args.seed)

    print(
        "env: "
        f"CUDA_VISIBLE_DEVICES={os.environ.get('CUDA_VISIBLE_DEVICES')} "
        f"NVTE_FLASH_ATTN={os.environ.get('NVTE_FLASH_ATTN')} "
        f"NVTE_FUSED_ATTN={os.environ.get('NVTE_FUSED_ATTN')} "
        f"NVTE_UNFUSED_ATTN={os.environ.get('NVTE_UNFUSED_ATTN')}",
        flush=True,
    )
    print(
        f"imports: transformer_engine={transformer_engine.__file__} "
        f"transformer_engine_torch={getattr(tex, '__file__', '<builtin>')}",
        flush=True,
    )
    print(
        f"config: mode={args.mode} q_len={args.q_len} kv_len={args.kv_len} "
        f"heads={args.num_heads} gqa={args.num_gqa_groups} head_dim={args.head_dim} "
        f"dtype={dtype} training={args.training}",
        flush=True,
    )

    torch.cuda.synchronize()

    for iteration in range(args.iters):
        print(f"iter {iteration}: start", flush=True)
        if args.mode == "single":
            run_fused_call(
                label=f"iter {iteration} long",
                q_len=args.q_len,
                kv_len=args.kv_len,
                args=args,
                dtype=dtype,
                device=device,
                fused_attn_fwd=fused_attn_fwd,
                tex=tex,
            )
        else:
            aux_stream = torch.cuda.Stream(device=device)
            with torch.cuda.stream(torch.cuda.current_stream()):
                run_fused_call(
                    label=f"iter {iteration} short/default-stream",
                    q_len=args.q_len,
                    kv_len=args.short_kv_len,
                    args=args,
                    dtype=dtype,
                    device=device,
                    fused_attn_fwd=fused_attn_fwd,
                    tex=tex,
                )
            with torch.cuda.stream(aux_stream):
                run_fused_call(
                    label=f"iter {iteration} long/aux-stream",
                    q_len=args.q_len,
                    kv_len=args.kv_len,
                    args=args,
                    dtype=dtype,
                    device=device,
                    fused_attn_fwd=fused_attn_fwd,
                    tex=tex,
                )

        print(f"iter {iteration}: synchronizing", flush=True)
        torch.cuda.synchronize()
        print(f"iter {iteration}: synchronized", flush=True)

    print(f"done, max_memory_allocated={torch.cuda.max_memory_allocated(device) / 1024**3:.2f} GiB")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())

Expected behavior

FusedAttention should either complete successfully or reject/fallback for unsupported shapes. It should not launch a kernel that hangs indefinitely.

Environment overview

  • NGC version: 26.04
  • Transformer Engine version: 2.17
  • PyTorch version: 2.12.0a0+0291f960b6.nv26.4.48445190
  • Python version: 3.12
  • CUDA version: 13.2
  • CUDNN version: 9.21.0

Device details

  • GPU model: NVIDIA B300 SXM6 AC
  • Compute capability reported by TE: sm103

Metadata

Metadata

Labels

bugSomething isn't working

Type

No type

Projects

No projects

Milestone

No milestone

Relationships

None yet

Development

No branches or pull requests

Issue actions