Spark-H3 / tests /test_dynamic_fused_cache.py
Aazeus's picture
Publish Spark-H3 code and model card (part 3)
b78342b verified
Raw History Blame Contribute Delete
8.39 kB
"""Fused SM80/SM120 callables reuse dynamic text extents and Top-K ratios."""
import pytest
import torch
from h3_sparse_attention.sol_numerator_virtual_q import (
_dynamic_tensor_cache_signature,
)
def test_dynamic_signature_ignores_non_singleton_extents_and_compact_strides():
first = torch.empty((1, 73560, 56, 128), dtype=torch.bfloat16)
second = torch.empty((1, 73583, 56, 128), dtype=torch.bfloat16)
assert _dynamic_tensor_cache_signature(first) == _dynamic_tensor_cache_signature(second)
def test_dynamic_signature_preserves_static_abi_properties():
base = torch.empty((1, 128, 4, 128), dtype=torch.bfloat16)
different_dtype = torch.empty((1, 128, 4, 128), dtype=torch.float32)
different_order = torch.empty((1, 4, 128, 128), dtype=torch.bfloat16).permute(0, 2, 1, 3)
different_singletons = torch.empty((1, 128, 1, 128), dtype=torch.bfloat16)
signature = _dynamic_tensor_cache_signature(base)
assert signature != _dynamic_tensor_cache_signature(different_dtype)
assert signature != _dynamic_tensor_cache_signature(different_order)
assert signature != _dynamic_tensor_cache_signature(different_singletons)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
@pytest.mark.parametrize(
"execution",
("threshold", "fused", "packed_external", "packed_external_no_route_qk"),
)
def test_sm120_compiled_callable_reuses_different_sink_lengths(execution):
if torch.cuda.get_device_capability() != (12, 0):
pytest.skip("SM120 required")
import h3_sparse_attention.sol_numerator_virtual_q as fused
torch.manual_seed(813)
video_tokens = 8192
def inputs(total_tokens):
heads = 2
blocks = (total_tokens + 63) // 64
q, k, v = [
torch.randn(
1,
total_tokens,
heads,
128,
device="cuda",
dtype=torch.bfloat16,
)
for _ in range(3)
]
ranges = torch.tensor(
[[0, video_tokens], [video_tokens, total_tokens]],
device="cuda",
dtype=torch.int64,
)
mapping = torch.tensor(
[0] * (video_tokens // 64)
+ [1] * (blocks - video_tokens // 64),
device="cuda",
dtype=torch.int64,
)
anchors = fused.build_virtual_anchors(q, ranges)
centroids = fused.reduce_virtual_key_centroids(k)
threshold = (
torch.zeros((1, blocks, heads), device="cuda", dtype=torch.float32)
if execution == "threshold"
else None
)
route = None
if execution.startswith("packed_external"):
route = torch.full(
(1, video_tokens // 64, heads, (blocks + 31) // 32),
-1,
device="cuda",
dtype=torch.int32,
)
return q, k, v, anchors, ranges, mapping, centroids, threshold, route
options = {
"force_local_blocks": False,
"_query_tokens": video_tokens,
"fused_topk_ratio": 0.1 if execution == "fused" else 0.0,
"skip_external_route_qk": execution == "packed_external_no_route_qk",
}
first = inputs(video_tokens + 64)
second = inputs(video_tokens + 81)
fused._FUSED_COMPILED.clear()
fused._FUSED_COMPILE_CALLS = 0
fused._FUSED_COMPILE_SECONDS = 0.0
fused._fused_virtual(
*first,
video_tokens,
64,
**options,
)
reused = fused._fused_virtual(
*second,
video_tokens,
81,
**options,
)[:, :video_tokens].clone()
reused_stats = fused.fused_compile_cache_stats()
assert reused_stats["entries"] == 1
assert reused_stats["compile_calls"] == 1
assert torch.isfinite(reused).all()
fused._FUSED_COMPILED.clear()
fresh = fused._fused_virtual(
*second,
video_tokens,
81,
**options,
)[:, :video_tokens]
torch.cuda.synchronize()
assert torch.equal(reused, fresh)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_sm120_compiled_callable_reuses_different_fused_topk_ratios(monkeypatch):
if torch.cuda.get_device_capability() != (12, 0):
pytest.skip("SM120 required")
import h3_sparse_attention.sol_numerator_virtual_q as fused
monkeypatch.setenv("H3_SM120_RUNTIME_TOPK_RATIO", "1")
torch.manual_seed(814)
video_tokens = 8192
sink_tokens = 64
total_tokens = video_tokens + sink_tokens
heads = 2
blocks = (total_tokens + 63) // 64
q, k, v = [
torch.randn(
1, total_tokens, heads, 128, device="cuda", dtype=torch.bfloat16
)
for _ in range(3)
]
ranges = torch.tensor(
[[0, video_tokens], [video_tokens, total_tokens]],
device="cuda",
dtype=torch.int64,
)
mapping = torch.tensor(
[0] * (video_tokens // 64) + [1] * (blocks - video_tokens // 64),
device="cuda",
dtype=torch.int64,
)
anchors = fused.build_virtual_anchors(q, ranges)
centroids = fused.reduce_virtual_key_centroids(k)
def run(ratio):
return fused._fused_virtual(
q, k, v, anchors, ranges, mapping, centroids, None, None,
video_tokens, sink_tokens,
force_local_blocks=False,
_query_tokens=video_tokens,
fused_topk_ratio=ratio,
)[:, :video_tokens].clone()
fused._FUSED_COMPILED.clear()
fused._FUSED_COMPILE_CALLS = 0
fused._FUSED_COMPILE_SECONDS = 0.0
first = run(0.1)
reused = run(0.2)
reused_stats = fused.fused_compile_cache_stats()
assert reused_stats["entries"] == 1
assert reused_stats["compile_calls"] == 1
assert not torch.equal(first, reused)
fused._FUSED_COMPILED.clear()
fresh = run(0.2)
torch.cuda.synchronize()
assert torch.equal(reused, fresh)
monkeypatch.setenv("H3_SM120_RUNTIME_TOPK_RATIO", "0")
fused._FUSED_COMPILED.clear()
static = run(0.2)
torch.cuda.synchronize()
assert torch.equal(reused, static)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_sm80_compiled_callable_reuses_text_lengths_and_topk_ratios(monkeypatch):
if torch.cuda.get_device_capability() != (8, 0):
pytest.skip("SM80 required")
import h3_sparse_attention.sol_numerator_virtual_q as fused
monkeypatch.setenv("H3_SM80_RUNTIME_TOPK_RATIO", "1")
torch.manual_seed(915)
video_tokens = 512
def inputs(sink_tokens):
total_tokens = video_tokens + sink_tokens
blocks = (total_tokens + 63) // 64
q, k, v = [
torch.randn(1, total_tokens, 2, 128, device="cuda", dtype=torch.bfloat16)
for _ in range(3)
]
ranges = torch.tensor(
[[0, video_tokens], [video_tokens, total_tokens]],
device="cuda", dtype=torch.int64,
)
mapping = torch.tensor(
[0] * (video_tokens // 64) + [1] * (blocks - video_tokens // 64),
device="cuda", dtype=torch.int64,
)
return (
q, k, v, fused.build_virtual_anchors(q, ranges), ranges, mapping,
fused.reduce_virtual_key_centroids(k),
)
def run(data, sink_tokens, ratio):
return fused._fused_virtual(
*data, None, None, video_tokens, sink_tokens,
force_local_blocks=False, _query_tokens=video_tokens,
fused_topk_ratio=ratio,
)[:, :video_tokens].clone()
shorter = inputs(64)
longer = inputs(81)
fused._FUSED_COMPILED.clear()
fused._FUSED_COMPILE_CALLS = 0
run(shorter, 64, 0.1)
first_ratio = run(longer, 81, 0.1)
second_ratio = run(longer, 81, 0.2)
assert fused.fused_compile_cache_stats()["entries"] == 1
assert fused.fused_compile_cache_stats()["compile_calls"] == 1
assert not torch.equal(first_ratio, second_ratio)
with pytest.raises(ValueError, match="fused_topk_ratio"):
run(longer, 81, 1.01)
fused._FUSED_COMPILED.clear()
fresh = run(longer, 81, 0.2)
monkeypatch.setenv("H3_SM80_RUNTIME_TOPK_RATIO", "0")
fused._FUSED_COMPILED.clear()
static = run(longer, 81, 0.2)
torch.cuda.synchronize()
assert torch.equal(second_ratio, fresh)
assert torch.equal(second_ratio, static)