Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions src/mcore_bridge/model/gpts/qwen4_exp.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,8 +180,9 @@ def _qsa_select(self, hidden_states, attn_kwargs, position_ids=None):
f'falling back to full attention ({"packing/thd" if is_thd else f"CP={cp_size}"}).')
return None, False
raise RuntimeError(f'QSA needs the sparse kernel here ({"packing/thd" if is_thd else f"CP={cp_size}"}), '
'but QSASparseCoreAttention was not installed -- triton is missing or '
f'kv_channels={getattr(self.config, "kv_channels", None)} is not a power of two. '
'but QSASparseCoreAttention was not installed -- triton is missing, '
f'kv_channels={getattr(self.config, "kv_channels", None)} is not a power of two, '
'or this device is not CUDA (the kernel is not compiled for NPU). '
'Use --padding_free false with context_parallel_size 1 to take the bool-mask path, '
f'or set {QSA_SPARSE_KERNEL_ENV}=0 to fall back to full attention.')
if cp_size > 1 and getattr(self.config, 'cp_comm_type', None) != 'all_gather':
Expand Down
13 changes: 9 additions & 4 deletions src/mcore_bridge/model/modules/kernels/qsa_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,10 +57,15 @@ def qsa_sparse_supported(head_dim: int) -> bool:
if not HAVE_TRITON or head_dim <= 0 or (head_dim & (head_dim - 1)):
return False
if not torch.cuda.is_available():
logger.warning_once('The QSA sparse kernel is only tested on CUDA GPUs and may fail to compile on '
'this device. If you hit triton compile errors, set '
f'{QSA_SPARSE_KERNEL_ENV}=0 to disable it (QSA then falls back to full '
'attention under packing/CP, and to the bool-mask path otherwise).')
# Triton-Ascend compiles this kernel with the CUDA tile sizes and fails
# (UB/Cc overflow). Do not install it there. Packing and CP then take the
# explicit QSA_SPARSE_KERNEL=0 full-attention fallback instead of dying
# inside the compiler. The bool-mask path (CP == 1, not packed) is unchanged.
logger.warning_once('The QSA sparse kernel is CUDA-only and is disabled on this device. '
f'Set {QSA_SPARSE_KERNEL_ENV}=0 to fall back to full attention under '
'packing or context parallelism. With CP == 1 and packing off, QSA '
'keeps the bool-mask path.')
return False
return True


Expand Down
52 changes: 52 additions & 0 deletions tests/test_qsa_sparse_supported.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
"""The QSA sparse kernel stays off unless CUDA can compile it."""
import ast
from pathlib import Path


def _load(have_triton, cuda):
path = Path(__file__).resolve().parents[1] / 'src/mcore_bridge/model/modules/kernels/qsa_kernels.py'
tree = ast.parse(path.read_text())
fn = next(node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == 'qsa_sparse_supported')
module = ast.Module(body=[fn], type_ignores=[])
ast.fix_missing_locations(module)
warnings = []

class _Cuda:

@staticmethod
def is_available():
return cuda

class _Logger:

@staticmethod
def warning_once(message):
warnings.append(message)

namespace = {
'HAVE_TRITON': have_triton,
'QSA_SPARSE_KERNEL_ENV': 'QSA_SPARSE_KERNEL',
'logger': _Logger(),
'torch': type('Torch', (), {'cuda': _Cuda})(),
'use_qsa_sparse_kernel': lambda: True,
}
exec(compile(module, str(path), 'exec'), namespace)
return namespace['qsa_sparse_supported'], warnings


def test_non_cuda_does_not_enable_the_kernel():
supported, warnings = _load(have_triton=True, cuda=False)
assert supported(64) is False
assert warnings and 'CUDA-only' in warnings[0]


def test_cuda_power_of_two_head_stays_enabled():
supported, warnings = _load(have_triton=True, cuda=True)
assert supported(128) is True
assert warnings == []


def test_non_power_of_two_head_stays_disabled():
supported, _ = _load(have_triton=True, cuda=True)
assert supported(96) is False