Merged
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
146 changes: 146 additions & 0 deletions backends/cuda/tests/test_triton_sdpa_splitk.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,6 +12,7 @@

import importlib
import itertools
import math
import unittest
from unittest import mock

Expand DownExpand Up@@ -42,6 +43,12 @@ def _import_sdpa_module():
return importlib.import_module("executorch.backends.cuda.triton.kernels.sdpa")


def _import_splitk_config():
from executorch.backends.cuda.triton.kernels.sdpa import _decode_splitk_config

return _decode_splitk_config


def _reference_sdpa(q, k, v, attn_mask=None, scale=None):
"""Compute reference SDPA in float32 with expanded KV heads for GQA."""
H_q = q.shape[1]
Expand DownExpand Up@@ -71,6 +78,9 @@ def _max_abs_error(out, ref):
# bf16 kernel vs fp32 reference tolerance.
# Matches benchmark cross-validation and test_triton_sdpa.py.
MAX_ABS_TOL = 1e-2
LEGACY_SPLITK_PHI = 5.0
FLOAT32_LOG_MAX = math.log(torch.finfo(torch.float32).max)
FLOAT32_LOG_MIN_SUBNORMAL = math.log(2**-149)


HEAD_DIMS_POW2 = [64, 128, 256]
Expand All@@ -95,6 +105,7 @@ def setUpClass(cls):
cls.small_query_splitk = _import_small_query_splitk()
cls.sdpa_module = _import_sdpa_module()
cls.sdpa = cls.sdpa_module.sdpa
cls.splitk_config = staticmethod(_import_splitk_config())

# ------------------------------------------------------------------
# Correctness
Expand DownExpand Up@@ -300,6 +311,105 @@ def test_qwen35_config(self):
self.assertFalse(torch.isnan(out).any())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_positive_logits_stable(self):
"""Large logits stay finite in both split-K kernel families."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
num_splits, chunk_size = self.splitk_config(Lk)
self.assertGreater(num_splits, 1)
self.assertNotEqual(0 // chunk_size, 300 // chunk_size)

k = torch.zeros(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

# Put nearly equal high scores in separate 256-token splits. The old
# fixed-phi path overflowed both partial softmaxes, and their proximity
# makes the result sensitive to correct cross-split rescaling.
k[:, :, 0, :] = 3.0
k[:, :, 300, :] = 2.96875
torch.manual_seed(42)
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# Guard against weakening the inputs below the old kernel's
# overflow boundary after subtracting its fixed phi.
high_scores = torch.stack(
[
(q[0, 0, 0].float() * k[0, 0, pos].float()).sum() / D**0.5
for pos in (0, 300)
]
)
self.assertTrue(
((high_scores - LEGACY_SPLITK_PHI) > FLOAT32_LOG_MAX).all()
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
# The tolerance includes BF16 output rounding for this
# concentrated two-key distribution.
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_negative_logits_do_not_underflow_to_zero(self):
"""Large negative logits retain their normalized weighted values."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
k = torch.full((B, H_kv, Lk, D), -3.0, dtype=torch.bfloat16, device="cuda")
v = torch.ones(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# The old fixed-phi exponent was below float32's smallest
# subnormal, so every weight became zero.
max_score = (q[0, 0, 0].float() * k[0, 0, 0].float()).sum() / D**0.5
self.assertLess(
max_score.item() - LEGACY_SPLITK_PHI,
FLOAT32_LOG_MIN_SUBNORMAL,
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_non_power_of_two_split_count(self):
"""Reduction masks lanes beyond the runtime split count."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 768, 128
num_splits, _ = self.splitk_config(Lk)
self.assertEqual(num_splits, 3)

torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_custom_scale(self):
"""Non-default attention scale."""
B, H_q, H_kv, Lq, Lk, D = 1, 8, 2, 1, 256, 128
Expand DownExpand Up@@ -354,6 +464,42 @@ def test_all_masked(self):
self.assertFalse(torch.isnan(out).any(), "All-masked should not NaN")
self.assertFalse(torch.isinf(out).any(), "All-masked should not Inf")

def test_kv_len_overwrites_poisoned_partial_buffers(self):
"""Every valid partial slot is written, including for empty splits."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 512, 128
valid_kv_len = 200
torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
kv_len = torch.tensor([valid_kv_len], dtype=torch.int32, device="cuda")

real_empty = torch.empty

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
poisoned_allocations = 0

def poisoned_empty(*args, **kwargs):
nonlocal poisoned_allocations
result = real_empty(*args, **kwargs)
if result.device.type == "cuda" and result.dtype == torch.float32:
result.fill_(float("nan"))
poisoned_allocations += 1
return result

with mock.patch.object(torch, "empty", side_effect=poisoned_empty):
out = splitk(q, k, v, kv_len=kv_len)

ref = _reference_sdpa(q, k[:, :, :valid_kv_len], v[:, :, :valid_kv_len])

self.assertGreaterEqual(poisoned_allocations, 3)
self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_lk_1(self):
"""Degenerate single KV position (num_splits=1)."""
B, H_q, H_kv, Lq, Lk, D = 1, 4, 2, 1, 1, 64
Expand Down
Loading
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all \u003cpre\u003e\u003ccode\u003e blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks"); } } catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); } })(); (function(){ try { var __m = "github.com"; var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Merged
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
146 changes: 146 additions & 0 deletions backends/cuda/tests/test_triton_sdpa_splitk.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,6 +12,7 @@

import importlib
import itertools
import math
import unittest
from unittest import mock

Expand DownExpand Up@@ -42,6 +43,12 @@ def _import_sdpa_module():
return importlib.import_module("executorch.backends.cuda.triton.kernels.sdpa")


def _import_splitk_config():
from executorch.backends.cuda.triton.kernels.sdpa import _decode_splitk_config

return _decode_splitk_config


def _reference_sdpa(q, k, v, attn_mask=None, scale=None):
"""Compute reference SDPA in float32 with expanded KV heads for GQA."""
H_q = q.shape[1]
Expand DownExpand Up@@ -71,6 +78,9 @@ def _max_abs_error(out, ref):
# bf16 kernel vs fp32 reference tolerance.
# Matches benchmark cross-validation and test_triton_sdpa.py.
MAX_ABS_TOL = 1e-2
LEGACY_SPLITK_PHI = 5.0
FLOAT32_LOG_MAX = math.log(torch.finfo(torch.float32).max)
FLOAT32_LOG_MIN_SUBNORMAL = math.log(2**-149)


HEAD_DIMS_POW2 = [64, 128, 256]
Expand All@@ -95,6 +105,7 @@ def setUpClass(cls):
cls.small_query_splitk = _import_small_query_splitk()
cls.sdpa_module = _import_sdpa_module()
cls.sdpa = cls.sdpa_module.sdpa
cls.splitk_config = staticmethod(_import_splitk_config())

# ------------------------------------------------------------------
# Correctness
Expand DownExpand Up@@ -300,6 +311,105 @@ def test_qwen35_config(self):
self.assertFalse(torch.isnan(out).any())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_positive_logits_stable(self):
"""Large logits stay finite in both split-K kernel families."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
num_splits, chunk_size = self.splitk_config(Lk)
self.assertGreater(num_splits, 1)
self.assertNotEqual(0 // chunk_size, 300 // chunk_size)

k = torch.zeros(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

# Put nearly equal high scores in separate 256-token splits. The old
# fixed-phi path overflowed both partial softmaxes, and their proximity
# makes the result sensitive to correct cross-split rescaling.
k[:, :, 0, :] = 3.0
k[:, :, 300, :] = 2.96875
torch.manual_seed(42)
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# Guard against weakening the inputs below the old kernel's
# overflow boundary after subtracting its fixed phi.
high_scores = torch.stack(
[
(q[0, 0, 0].float() * k[0, 0, pos].float()).sum() / D**0.5
for pos in (0, 300)
]
)
self.assertTrue(
((high_scores - LEGACY_SPLITK_PHI) > FLOAT32_LOG_MAX).all()
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
# The tolerance includes BF16 output rounding for this
# concentrated two-key distribution.
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_negative_logits_do_not_underflow_to_zero(self):
"""Large negative logits retain their normalized weighted values."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
k = torch.full((B, H_kv, Lk, D), -3.0, dtype=torch.bfloat16, device="cuda")
v = torch.ones(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# The old fixed-phi exponent was below float32's smallest
# subnormal, so every weight became zero.
max_score = (q[0, 0, 0].float() * k[0, 0, 0].float()).sum() / D**0.5
self.assertLess(
max_score.item() - LEGACY_SPLITK_PHI,
FLOAT32_LOG_MIN_SUBNORMAL,
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_non_power_of_two_split_count(self):
"""Reduction masks lanes beyond the runtime split count."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 768, 128
num_splits, _ = self.splitk_config(Lk)
self.assertEqual(num_splits, 3)

torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_custom_scale(self):
"""Non-default attention scale."""
B, H_q, H_kv, Lq, Lk, D = 1, 8, 2, 1, 256, 128
Expand DownExpand Up@@ -354,6 +464,42 @@ def test_all_masked(self):
self.assertFalse(torch.isnan(out).any(), "All-masked should not NaN")
self.assertFalse(torch.isinf(out).any(), "All-masked should not Inf")

def test_kv_len_overwrites_poisoned_partial_buffers(self):
"""Every valid partial slot is written, including for empty splits."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 512, 128
valid_kv_len = 200
torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
kv_len = torch.tensor([valid_kv_len], dtype=torch.int32, device="cuda")

real_empty = torch.empty

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
poisoned_allocations = 0

def poisoned_empty(*args, **kwargs):
nonlocal poisoned_allocations
result = real_empty(*args, **kwargs)
if result.device.type == "cuda" and result.dtype == torch.float32:
result.fill_(float("nan"))
poisoned_allocations += 1
return result

with mock.patch.object(torch, "empty", side_effect=poisoned_empty):
out = splitk(q, k, v, kv_len=kv_len)

ref = _reference_sdpa(q, k[:, :, :valid_kv_len], v[:, :, :valid_kv_len])

self.assertGreaterEqual(poisoned_allocations, 3)
self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_lk_1(self):
"""Degenerate single KV position (num_splits=1)."""
B, H_q, H_kv, Lq, Lk, D = 1, 4, 2, 1, 1, 64
Expand Down
Loading
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
146 changes: 146 additions & 0 deletions backends/cuda/tests/test_triton_sdpa_splitk.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,6 +12,7 @@

import importlib
import itertools
import math
import unittest
from unittest import mock

Expand DownExpand Up@@ -42,6 +43,12 @@ def _import_sdpa_module():
return importlib.import_module("executorch.backends.cuda.triton.kernels.sdpa")


def _import_splitk_config():
from executorch.backends.cuda.triton.kernels.sdpa import _decode_splitk_config

return _decode_splitk_config


def _reference_sdpa(q, k, v, attn_mask=None, scale=None):
"""Compute reference SDPA in float32 with expanded KV heads for GQA."""
H_q = q.shape[1]
Expand DownExpand Up@@ -71,6 +78,9 @@ def _max_abs_error(out, ref):
# bf16 kernel vs fp32 reference tolerance.
# Matches benchmark cross-validation and test_triton_sdpa.py.
MAX_ABS_TOL = 1e-2
LEGACY_SPLITK_PHI = 5.0
FLOAT32_LOG_MAX = math.log(torch.finfo(torch.float32).max)
FLOAT32_LOG_MIN_SUBNORMAL = math.log(2**-149)


HEAD_DIMS_POW2 = [64, 128, 256]
Expand All@@ -95,6 +105,7 @@ def setUpClass(cls):
cls.small_query_splitk = _import_small_query_splitk()
cls.sdpa_module = _import_sdpa_module()
cls.sdpa = cls.sdpa_module.sdpa
cls.splitk_config = staticmethod(_import_splitk_config())

# ------------------------------------------------------------------
# Correctness
Expand DownExpand Up@@ -300,6 +311,105 @@ def test_qwen35_config(self):
self.assertFalse(torch.isnan(out).any())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_positive_logits_stable(self):
"""Large logits stay finite in both split-K kernel families."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
num_splits, chunk_size = self.splitk_config(Lk)
self.assertGreater(num_splits, 1)
self.assertNotEqual(0 // chunk_size, 300 // chunk_size)

k = torch.zeros(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

# Put nearly equal high scores in separate 256-token splits. The old
# fixed-phi path overflowed both partial softmaxes, and their proximity
# makes the result sensitive to correct cross-split rescaling.
k[:, :, 0, :] = 3.0
k[:, :, 300, :] = 2.96875
torch.manual_seed(42)
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# Guard against weakening the inputs below the old kernel's
# overflow boundary after subtracting its fixed phi.
high_scores = torch.stack(
[
(q[0, 0, 0].float() * k[0, 0, pos].float()).sum() / D**0.5
for pos in (0, 300)
]
)
self.assertTrue(
((high_scores - LEGACY_SPLITK_PHI) > FLOAT32_LOG_MAX).all()
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
# The tolerance includes BF16 output rounding for this
# concentrated two-key distribution.
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_negative_logits_do_not_underflow_to_zero(self):
"""Large negative logits retain their normalized weighted values."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
k = torch.full((B, H_kv, Lk, D), -3.0, dtype=torch.bfloat16, device="cuda")
v = torch.ones(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# The old fixed-phi exponent was below float32's smallest
# subnormal, so every weight became zero.
max_score = (q[0, 0, 0].float() * k[0, 0, 0].float()).sum() / D**0.5
self.assertLess(
max_score.item() - LEGACY_SPLITK_PHI,
FLOAT32_LOG_MIN_SUBNORMAL,
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_non_power_of_two_split_count(self):
"""Reduction masks lanes beyond the runtime split count."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 768, 128
num_splits, _ = self.splitk_config(Lk)
self.assertEqual(num_splits, 3)

torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_custom_scale(self):
"""Non-default attention scale."""
B, H_q, H_kv, Lq, Lk, D = 1, 8, 2, 1, 256, 128
Expand DownExpand Up@@ -354,6 +464,42 @@ def test_all_masked(self):
self.assertFalse(torch.isnan(out).any(), "All-masked should not NaN")
self.assertFalse(torch.isinf(out).any(), "All-masked should not Inf")

def test_kv_len_overwrites_poisoned_partial_buffers(self):
"""Every valid partial slot is written, including for empty splits."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 512, 128
valid_kv_len = 200
torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
kv_len = torch.tensor([valid_kv_len], dtype=torch.int32, device="cuda")

real_empty = torch.empty

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
poisoned_allocations = 0

def poisoned_empty(*args, **kwargs):
nonlocal poisoned_allocations
result = real_empty(*args, **kwargs)
if result.device.type == "cuda" and result.dtype == torch.float32:
result.fill_(float("nan"))
poisoned_allocations += 1
return result

with mock.patch.object(torch, "empty", side_effect=poisoned_empty):
out = splitk(q, k, v, kv_len=kv_len)

ref = _reference_sdpa(q, k[:, :, :valid_kv_len], v[:, :, :valid_kv_len])

self.assertGreaterEqual(poisoned_allocations, 3)
self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_lk_1(self):
"""Degenerate single KV position (num_splits=1)."""
B, H_q, H_kv, Lq, Lk, D = 1, 4, 2, 1, 1, 64
Expand Down
Loading
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length \u003e 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
146 changes: 146 additions & 0 deletions backends/cuda/tests/test_triton_sdpa_splitk.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,6 +12,7 @@

import importlib
import itertools
import math
import unittest
from unittest import mock

Expand DownExpand Up@@ -42,6 +43,12 @@ def _import_sdpa_module():
return importlib.import_module("executorch.backends.cuda.triton.kernels.sdpa")


def _import_splitk_config():
from executorch.backends.cuda.triton.kernels.sdpa import _decode_splitk_config

return _decode_splitk_config


def _reference_sdpa(q, k, v, attn_mask=None, scale=None):
"""Compute reference SDPA in float32 with expanded KV heads for GQA."""
H_q = q.shape[1]
Expand DownExpand Up@@ -71,6 +78,9 @@ def _max_abs_error(out, ref):
# bf16 kernel vs fp32 reference tolerance.
# Matches benchmark cross-validation and test_triton_sdpa.py.
MAX_ABS_TOL = 1e-2
LEGACY_SPLITK_PHI = 5.0
FLOAT32_LOG_MAX = math.log(torch.finfo(torch.float32).max)
FLOAT32_LOG_MIN_SUBNORMAL = math.log(2**-149)


HEAD_DIMS_POW2 = [64, 128, 256]
Expand All@@ -95,6 +105,7 @@ def setUpClass(cls):
cls.small_query_splitk = _import_small_query_splitk()
cls.sdpa_module = _import_sdpa_module()
cls.sdpa = cls.sdpa_module.sdpa
cls.splitk_config = staticmethod(_import_splitk_config())

# ------------------------------------------------------------------
# Correctness
Expand DownExpand Up@@ -300,6 +311,105 @@ def test_qwen35_config(self):
self.assertFalse(torch.isnan(out).any())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_positive_logits_stable(self):
"""Large logits stay finite in both split-K kernel families."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
num_splits, chunk_size = self.splitk_config(Lk)
self.assertGreater(num_splits, 1)
self.assertNotEqual(0 // chunk_size, 300 // chunk_size)

k = torch.zeros(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

# Put nearly equal high scores in separate 256-token splits. The old
# fixed-phi path overflowed both partial softmaxes, and their proximity
# makes the result sensitive to correct cross-split rescaling.
k[:, :, 0, :] = 3.0
k[:, :, 300, :] = 2.96875
torch.manual_seed(42)
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# Guard against weakening the inputs below the old kernel's
# overflow boundary after subtracting its fixed phi.
high_scores = torch.stack(
[
(q[0, 0, 0].float() * k[0, 0, pos].float()).sum() / D**0.5
for pos in (0, 300)
]
)
self.assertTrue(
((high_scores - LEGACY_SPLITK_PHI) > FLOAT32_LOG_MAX).all()
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
# The tolerance includes BF16 output rounding for this
# concentrated two-key distribution.
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_negative_logits_do_not_underflow_to_zero(self):
"""Large negative logits retain their normalized weighted values."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
k = torch.full((B, H_kv, Lk, D), -3.0, dtype=torch.bfloat16, device="cuda")
v = torch.ones(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# The old fixed-phi exponent was below float32's smallest
# subnormal, so every weight became zero.
max_score = (q[0, 0, 0].float() * k[0, 0, 0].float()).sum() / D**0.5
self.assertLess(
max_score.item() - LEGACY_SPLITK_PHI,
FLOAT32_LOG_MIN_SUBNORMAL,
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_non_power_of_two_split_count(self):
"""Reduction masks lanes beyond the runtime split count."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 768, 128
num_splits, _ = self.splitk_config(Lk)
self.assertEqual(num_splits, 3)

torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_custom_scale(self):
"""Non-default attention scale."""
B, H_q, H_kv, Lq, Lk, D = 1, 8, 2, 1, 256, 128
Expand DownExpand Up@@ -354,6 +464,42 @@ def test_all_masked(self):
self.assertFalse(torch.isnan(out).any(), "All-masked should not NaN")
self.assertFalse(torch.isinf(out).any(), "All-masked should not Inf")

def test_kv_len_overwrites_poisoned_partial_buffers(self):
"""Every valid partial slot is written, including for empty splits."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 512, 128
valid_kv_len = 200
torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
kv_len = torch.tensor([valid_kv_len], dtype=torch.int32, device="cuda")

real_empty = torch.empty

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
poisoned_allocations = 0

def poisoned_empty(*args, **kwargs):
nonlocal poisoned_allocations
result = real_empty(*args, **kwargs)
if result.device.type == "cuda" and result.dtype == torch.float32:
result.fill_(float("nan"))
poisoned_allocations += 1
return result

with mock.patch.object(torch, "empty", side_effect=poisoned_empty):
out = splitk(q, k, v, kv_len=kv_len)

ref = _reference_sdpa(q, k[:, :, :valid_kv_len], v[:, :, :valid_kv_len])

self.assertGreaterEqual(poisoned_allocations, 3)
self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_lk_1(self):
"""Degenerate single KV position (num_splits=1)."""
B, H_q, H_kv, Lq, Lk, D = 1, 4, 2, 1, 1, 64
Expand Down
Loading
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Merged
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
146 changes: 146 additions & 0 deletions backends/cuda/tests/test_triton_sdpa_splitk.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,6 +12,7 @@

import importlib
import itertools
import math
import unittest
from unittest import mock

Expand DownExpand Up@@ -42,6 +43,12 @@ def _import_sdpa_module():
return importlib.import_module("executorch.backends.cuda.triton.kernels.sdpa")


def _import_splitk_config():
from executorch.backends.cuda.triton.kernels.sdpa import _decode_splitk_config

return _decode_splitk_config


def _reference_sdpa(q, k, v, attn_mask=None, scale=None):
"""Compute reference SDPA in float32 with expanded KV heads for GQA."""
H_q = q.shape[1]
Expand DownExpand Up@@ -71,6 +78,9 @@ def _max_abs_error(out, ref):
# bf16 kernel vs fp32 reference tolerance.
# Matches benchmark cross-validation and test_triton_sdpa.py.
MAX_ABS_TOL = 1e-2
LEGACY_SPLITK_PHI = 5.0
FLOAT32_LOG_MAX = math.log(torch.finfo(torch.float32).max)
FLOAT32_LOG_MIN_SUBNORMAL = math.log(2**-149)


HEAD_DIMS_POW2 = [64, 128, 256]
Expand All@@ -95,6 +105,7 @@ def setUpClass(cls):
cls.small_query_splitk = _import_small_query_splitk()
cls.sdpa_module = _import_sdpa_module()
cls.sdpa = cls.sdpa_module.sdpa
cls.splitk_config = staticmethod(_import_splitk_config())

# ------------------------------------------------------------------
# Correctness
Expand DownExpand Up@@ -300,6 +311,105 @@ def test_qwen35_config(self):
self.assertFalse(torch.isnan(out).any())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_positive_logits_stable(self):
"""Large logits stay finite in both split-K kernel families."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
num_splits, chunk_size = self.splitk_config(Lk)
self.assertGreater(num_splits, 1)
self.assertNotEqual(0 // chunk_size, 300 // chunk_size)

k = torch.zeros(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

# Put nearly equal high scores in separate 256-token splits. The old
# fixed-phi path overflowed both partial softmaxes, and their proximity
# makes the result sensitive to correct cross-split rescaling.
k[:, :, 0, :] = 3.0
k[:, :, 300, :] = 2.96875
torch.manual_seed(42)
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# Guard against weakening the inputs below the old kernel's
# overflow boundary after subtracting its fixed phi.
high_scores = torch.stack(
[
(q[0, 0, 0].float() * k[0, 0, pos].float()).sum() / D**0.5
for pos in (0, 300)
]
)
self.assertTrue(
((high_scores - LEGACY_SPLITK_PHI) > FLOAT32_LOG_MAX).all()
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
# The tolerance includes BF16 output rounding for this
# concentrated two-key distribution.
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_negative_logits_do_not_underflow_to_zero(self):
"""Large negative logits retain their normalized weighted values."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
k = torch.full((B, H_kv, Lk, D), -3.0, dtype=torch.bfloat16, device="cuda")
v = torch.ones(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# The old fixed-phi exponent was below float32's smallest
# subnormal, so every weight became zero.
max_score = (q[0, 0, 0].float() * k[0, 0, 0].float()).sum() / D**0.5
self.assertLess(
max_score.item() - LEGACY_SPLITK_PHI,
FLOAT32_LOG_MIN_SUBNORMAL,
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_non_power_of_two_split_count(self):
"""Reduction masks lanes beyond the runtime split count."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 768, 128
num_splits, _ = self.splitk_config(Lk)
self.assertEqual(num_splits, 3)

torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_custom_scale(self):
"""Non-default attention scale."""
B, H_q, H_kv, Lq, Lk, D = 1, 8, 2, 1, 256, 128
Expand DownExpand Up@@ -354,6 +464,42 @@ def test_all_masked(self):
self.assertFalse(torch.isnan(out).any(), "All-masked should not NaN")
self.assertFalse(torch.isinf(out).any(), "All-masked should not Inf")

def test_kv_len_overwrites_poisoned_partial_buffers(self):
"""Every valid partial slot is written, including for empty splits."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 512, 128
valid_kv_len = 200
torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
kv_len = torch.tensor([valid_kv_len], dtype=torch.int32, device="cuda")

real_empty = torch.empty

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
poisoned_allocations = 0

def poisoned_empty(*args, **kwargs):
nonlocal poisoned_allocations
result = real_empty(*args, **kwargs)
if result.device.type == "cuda" and result.dtype == torch.float32:
result.fill_(float("nan"))
poisoned_allocations += 1
return result

with mock.patch.object(torch, "empty", side_effect=poisoned_empty):
out = splitk(q, k, v, kv_len=kv_len)

ref = _reference_sdpa(q, k[:, :, :valid_kv_len], v[:, :, :valid_kv_len])

self.assertGreaterEqual(poisoned_allocations, 3)
self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_lk_1(self):
"""Degenerate single KV position (num_splits=1)."""
B, H_q, H_kv, Lq, Lk, D = 1, 4, 2, 1, 1, 64
Expand Down
Loading
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
146 changes: 146 additions & 0 deletions backends/cuda/tests/test_triton_sdpa_splitk.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,6 +12,7 @@

import importlib
import itertools
import math
import unittest
from unittest import mock

Expand DownExpand Up@@ -42,6 +43,12 @@ def _import_sdpa_module():
return importlib.import_module("executorch.backends.cuda.triton.kernels.sdpa")


def _import_splitk_config():
from executorch.backends.cuda.triton.kernels.sdpa import _decode_splitk_config

return _decode_splitk_config


def _reference_sdpa(q, k, v, attn_mask=None, scale=None):
"""Compute reference SDPA in float32 with expanded KV heads for GQA."""
H_q = q.shape[1]
Expand DownExpand Up@@ -71,6 +78,9 @@ def _max_abs_error(out, ref):
# bf16 kernel vs fp32 reference tolerance.
# Matches benchmark cross-validation and test_triton_sdpa.py.
MAX_ABS_TOL = 1e-2
LEGACY_SPLITK_PHI = 5.0
FLOAT32_LOG_MAX = math.log(torch.finfo(torch.float32).max)
FLOAT32_LOG_MIN_SUBNORMAL = math.log(2**-149)


HEAD_DIMS_POW2 = [64, 128, 256]
Expand All@@ -95,6 +105,7 @@ def setUpClass(cls):
cls.small_query_splitk = _import_small_query_splitk()
cls.sdpa_module = _import_sdpa_module()
cls.sdpa = cls.sdpa_module.sdpa
cls.splitk_config = staticmethod(_import_splitk_config())

# ------------------------------------------------------------------
# Correctness
Expand DownExpand Up@@ -300,6 +311,105 @@ def test_qwen35_config(self):
self.assertFalse(torch.isnan(out).any())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_positive_logits_stable(self):
"""Large logits stay finite in both split-K kernel families."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
num_splits, chunk_size = self.splitk_config(Lk)
self.assertGreater(num_splits, 1)
self.assertNotEqual(0 // chunk_size, 300 // chunk_size)

k = torch.zeros(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

# Put nearly equal high scores in separate 256-token splits. The old
# fixed-phi path overflowed both partial softmaxes, and their proximity
# makes the result sensitive to correct cross-split rescaling.
k[:, :, 0, :] = 3.0
k[:, :, 300, :] = 2.96875
torch.manual_seed(42)
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# Guard against weakening the inputs below the old kernel's
# overflow boundary after subtracting its fixed phi.
high_scores = torch.stack(
[
(q[0, 0, 0].float() * k[0, 0, pos].float()).sum() / D**0.5
for pos in (0, 300)
]
)
self.assertTrue(
((high_scores - LEGACY_SPLITK_PHI) > FLOAT32_LOG_MAX).all()
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
# The tolerance includes BF16 output rounding for this
# concentrated two-key distribution.
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_negative_logits_do_not_underflow_to_zero(self):
"""Large negative logits retain their normalized weighted values."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
k = torch.full((B, H_kv, Lk, D), -3.0, dtype=torch.bfloat16, device="cuda")
v = torch.ones(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# The old fixed-phi exponent was below float32's smallest
# subnormal, so every weight became zero.
max_score = (q[0, 0, 0].float() * k[0, 0, 0].float()).sum() / D**0.5
self.assertLess(
max_score.item() - LEGACY_SPLITK_PHI,
FLOAT32_LOG_MIN_SUBNORMAL,
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_non_power_of_two_split_count(self):
"""Reduction masks lanes beyond the runtime split count."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 768, 128
num_splits, _ = self.splitk_config(Lk)
self.assertEqual(num_splits, 3)

torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_custom_scale(self):
"""Non-default attention scale."""
B, H_q, H_kv, Lq, Lk, D = 1, 8, 2, 1, 256, 128
Expand DownExpand Up@@ -354,6 +464,42 @@ def test_all_masked(self):
self.assertFalse(torch.isnan(out).any(), "All-masked should not NaN")
self.assertFalse(torch.isinf(out).any(), "All-masked should not Inf")

def test_kv_len_overwrites_poisoned_partial_buffers(self):
"""Every valid partial slot is written, including for empty splits."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 512, 128
valid_kv_len = 200
torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
kv_len = torch.tensor([valid_kv_len], dtype=torch.int32, device="cuda")

real_empty = torch.empty

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
poisoned_allocations = 0

def poisoned_empty(*args, **kwargs):
nonlocal poisoned_allocations
result = real_empty(*args, **kwargs)
if result.device.type == "cuda" and result.dtype == torch.float32:
result.fill_(float("nan"))
poisoned_allocations += 1
return result

with mock.patch.object(torch, "empty", side_effect=poisoned_empty):
out = splitk(q, k, v, kv_len=kv_len)

ref = _reference_sdpa(q, k[:, :, :valid_kv_len], v[:, :, :valid_kv_len])

self.assertGreaterEqual(poisoned_allocations, 3)
self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_lk_1(self):
"""Degenerate single KV position (num_splits=1)."""
B, H_q, H_kv, Lq, Lk, D = 1, 4, 2, 1, 1, 64
Expand Down
Loading
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
146 changes: 146 additions & 0 deletions backends/cuda/tests/test_triton_sdpa_splitk.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,6 +12,7 @@

import importlib
import itertools
import math
import unittest
from unittest import mock

Expand DownExpand Up@@ -42,6 +43,12 @@ def _import_sdpa_module():
return importlib.import_module("executorch.backends.cuda.triton.kernels.sdpa")


def _import_splitk_config():
from executorch.backends.cuda.triton.kernels.sdpa import _decode_splitk_config

return _decode_splitk_config


def _reference_sdpa(q, k, v, attn_mask=None, scale=None):
"""Compute reference SDPA in float32 with expanded KV heads for GQA."""
H_q = q.shape[1]
Expand DownExpand Up@@ -71,6 +78,9 @@ def _max_abs_error(out, ref):
# bf16 kernel vs fp32 reference tolerance.
# Matches benchmark cross-validation and test_triton_sdpa.py.
MAX_ABS_TOL = 1e-2
LEGACY_SPLITK_PHI = 5.0
FLOAT32_LOG_MAX = math.log(torch.finfo(torch.float32).max)
FLOAT32_LOG_MIN_SUBNORMAL = math.log(2**-149)


HEAD_DIMS_POW2 = [64, 128, 256]
Expand All@@ -95,6 +105,7 @@ def setUpClass(cls):
cls.small_query_splitk = _import_small_query_splitk()
cls.sdpa_module = _import_sdpa_module()
cls.sdpa = cls.sdpa_module.sdpa
cls.splitk_config = staticmethod(_import_splitk_config())

# ------------------------------------------------------------------
# Correctness
Expand DownExpand Up@@ -300,6 +311,105 @@ def test_qwen35_config(self):
self.assertFalse(torch.isnan(out).any())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_positive_logits_stable(self):
"""Large logits stay finite in both split-K kernel families."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
num_splits, chunk_size = self.splitk_config(Lk)
self.assertGreater(num_splits, 1)
self.assertNotEqual(0 // chunk_size, 300 // chunk_size)

k = torch.zeros(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

# Put nearly equal high scores in separate 256-token splits. The old
# fixed-phi path overflowed both partial softmaxes, and their proximity
# makes the result sensitive to correct cross-split rescaling.
k[:, :, 0, :] = 3.0
k[:, :, 300, :] = 2.96875
torch.manual_seed(42)
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# Guard against weakening the inputs below the old kernel's
# overflow boundary after subtracting its fixed phi.
high_scores = torch.stack(
[
(q[0, 0, 0].float() * k[0, 0, pos].float()).sum() / D**0.5
for pos in (0, 300)
]
)
self.assertTrue(
((high_scores - LEGACY_SPLITK_PHI) > FLOAT32_LOG_MAX).all()
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
# The tolerance includes BF16 output rounding for this
# concentrated two-key distribution.
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_negative_logits_do_not_underflow_to_zero(self):
"""Large negative logits retain their normalized weighted values."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
k = torch.full((B, H_kv, Lk, D), -3.0, dtype=torch.bfloat16, device="cuda")
v = torch.ones(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# The old fixed-phi exponent was below float32's smallest
# subnormal, so every weight became zero.
max_score = (q[0, 0, 0].float() * k[0, 0, 0].float()).sum() / D**0.5
self.assertLess(
max_score.item() - LEGACY_SPLITK_PHI,
FLOAT32_LOG_MIN_SUBNORMAL,
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_non_power_of_two_split_count(self):
"""Reduction masks lanes beyond the runtime split count."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 768, 128
num_splits, _ = self.splitk_config(Lk)
self.assertEqual(num_splits, 3)

torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_custom_scale(self):
"""Non-default attention scale."""
B, H_q, H_kv, Lq, Lk, D = 1, 8, 2, 1, 256, 128
Expand DownExpand Up@@ -354,6 +464,42 @@ def test_all_masked(self):
self.assertFalse(torch.isnan(out).any(), "All-masked should not NaN")
self.assertFalse(torch.isinf(out).any(), "All-masked should not Inf")

def test_kv_len_overwrites_poisoned_partial_buffers(self):
"""Every valid partial slot is written, including for empty splits."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 512, 128
valid_kv_len = 200
torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
kv_len = torch.tensor([valid_kv_len], dtype=torch.int32, device="cuda")

real_empty = torch.empty

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
poisoned_allocations = 0

def poisoned_empty(*args, **kwargs):
nonlocal poisoned_allocations
result = real_empty(*args, **kwargs)
if result.device.type == "cuda" and result.dtype == torch.float32:
result.fill_(float("nan"))
poisoned_allocations += 1
return result

with mock.patch.object(torch, "empty", side_effect=poisoned_empty):
out = splitk(q, k, v, kv_len=kv_len)

ref = _reference_sdpa(q, k[:, :, :valid_kv_len], v[:, :, :valid_kv_len])

self.assertGreaterEqual(poisoned_allocations, 3)
self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_lk_1(self):
"""Degenerate single KV position (num_splits=1)."""
B, H_q, H_kv, Lq, Lk, D = 1, 4, 2, 1, 1, 64
Expand Down
Loading
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Merged
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
146 changes: 146 additions & 0 deletions backends/cuda/tests/test_triton_sdpa_splitk.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,6 +12,7 @@

import importlib
import itertools
import math
import unittest
from unittest import mock

Expand DownExpand Up@@ -42,6 +43,12 @@ def _import_sdpa_module():
return importlib.import_module("executorch.backends.cuda.triton.kernels.sdpa")


def _import_splitk_config():
from executorch.backends.cuda.triton.kernels.sdpa import _decode_splitk_config

return _decode_splitk_config


def _reference_sdpa(q, k, v, attn_mask=None, scale=None):
"""Compute reference SDPA in float32 with expanded KV heads for GQA."""
H_q = q.shape[1]
Expand DownExpand Up@@ -71,6 +78,9 @@ def _max_abs_error(out, ref):
# bf16 kernel vs fp32 reference tolerance.
# Matches benchmark cross-validation and test_triton_sdpa.py.
MAX_ABS_TOL = 1e-2
LEGACY_SPLITK_PHI = 5.0
FLOAT32_LOG_MAX = math.log(torch.finfo(torch.float32).max)
FLOAT32_LOG_MIN_SUBNORMAL = math.log(2**-149)


HEAD_DIMS_POW2 = [64, 128, 256]
Expand All@@ -95,6 +105,7 @@ def setUpClass(cls):
cls.small_query_splitk = _import_small_query_splitk()
cls.sdpa_module = _import_sdpa_module()
cls.sdpa = cls.sdpa_module.sdpa
cls.splitk_config = staticmethod(_import_splitk_config())

# ------------------------------------------------------------------
# Correctness
Expand DownExpand Up@@ -300,6 +311,105 @@ def test_qwen35_config(self):
self.assertFalse(torch.isnan(out).any())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_positive_logits_stable(self):
"""Large logits stay finite in both split-K kernel families."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
num_splits, chunk_size = self.splitk_config(Lk)
self.assertGreater(num_splits, 1)
self.assertNotEqual(0 // chunk_size, 300 // chunk_size)

k = torch.zeros(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

# Put nearly equal high scores in separate 256-token splits. The old
# fixed-phi path overflowed both partial softmaxes, and their proximity
# makes the result sensitive to correct cross-split rescaling.
k[:, :, 0, :] = 3.0
k[:, :, 300, :] = 2.96875
torch.manual_seed(42)
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# Guard against weakening the inputs below the old kernel's
# overflow boundary after subtracting its fixed phi.
high_scores = torch.stack(
[
(q[0, 0, 0].float() * k[0, 0, pos].float()).sum() / D**0.5
for pos in (0, 300)
]
)
self.assertTrue(
((high_scores - LEGACY_SPLITK_PHI) > FLOAT32_LOG_MAX).all()
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
# The tolerance includes BF16 output rounding for this
# concentrated two-key distribution.
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_large_negative_logits_do_not_underflow_to_zero(self):
"""Large negative logits retain their normalized weighted values."""
B, H_q, H_kv, Lk, D = 1, 32, 8, 512, 128
k = torch.full((B, H_kv, Lk, D), -3.0, dtype=torch.bfloat16, device="cuda")
v = torch.ones(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(2, self.small_query_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.full(
(B, H_q, Lq, D), 3.0, dtype=torch.bfloat16, device="cuda"
)

# The old fixed-phi exponent was below float32's smallest
# subnormal, so every weight became zero.
max_score = (q[0, 0, 0].float() * k[0, 0, 0].float()).sum() / D**0.5
self.assertLess(
max_score.item() - LEGACY_SPLITK_PHI,
FLOAT32_LOG_MIN_SUBNORMAL,
)

out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_non_power_of_two_split_count(self):
"""Reduction masks lanes beyond the runtime split count."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 768, 128
num_splits, _ = self.splitk_config(Lk)
self.assertEqual(num_splits, 3)

torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
out = splitk(q, k, v)
ref = _reference_sdpa(q, k, v)

self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_custom_scale(self):
"""Non-default attention scale."""
B, H_q, H_kv, Lq, Lk, D = 1, 8, 2, 1, 256, 128
Expand DownExpand Up@@ -354,6 +464,42 @@ def test_all_masked(self):
self.assertFalse(torch.isnan(out).any(), "All-masked should not NaN")
self.assertFalse(torch.isinf(out).any(), "All-masked should not Inf")

def test_kv_len_overwrites_poisoned_partial_buffers(self):
"""Every valid partial slot is written, including for empty splits."""
B, H_q, H_kv, Lk, D = 1, 8, 2, 512, 128
valid_kv_len = 200
torch.manual_seed(42)
k = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, H_kv, Lk, D, dtype=torch.bfloat16, device="cuda")
kv_len = torch.tensor([valid_kv_len], dtype=torch.int32, device="cuda")

real_empty = torch.empty

for Lq, splitk in [
(1, self.legacy_splitk),
(4, self.small_query_splitk),
]:
with self.subTest(Lq=Lq):
q = torch.randn(B, H_q, Lq, D, dtype=torch.bfloat16, device="cuda")
poisoned_allocations = 0

def poisoned_empty(*args, **kwargs):
nonlocal poisoned_allocations
result = real_empty(*args, **kwargs)
if result.device.type == "cuda" and result.dtype == torch.float32:
result.fill_(float("nan"))
poisoned_allocations += 1
return result

with mock.patch.object(torch, "empty", side_effect=poisoned_empty):
out = splitk(q, k, v, kv_len=kv_len)

ref = _reference_sdpa(q, k[:, :, :valid_kv_len], v[:, :, :valid_kv_len])

self.assertGreaterEqual(poisoned_allocations, 3)
self.assertTrue(torch.isfinite(out).all())
self.assertLess(_max_abs_error(out, ref), MAX_ABS_TOL)

def test_lk_1(self):
"""Degenerate single KV position (num_splits=1)."""
B, H_q, H_kv, Lq, Lk, D = 1, 4, 2, 1, 1, 64
Expand Down
Loading
Loading