Text-to-Video
Diffusers
Safetensors
MiniMax H3
MiniMaxH3ModularPipeline
image-to-video
audio-video-generation
sparse-attention
block-sparse
inference-acceleration
triton
comfyui
Instructions to use Aazeus/Spark-H3 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use Aazeus/Spark-H3 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("Aazeus/Spark-H3", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download tests/test_dynamic_fused_cache.py from Aazeus/Spark-H3: direct link, hf CLI and curl.
- Browser
- Download file 8.39 kB
-
https://huggingface.co/Aazeus/Spark-H3/resolve/main/tests/test_dynamic_fused_cache.py
- Command line
-
hf download hf://Aazeus/Spark-H3/tests/test_dynamic_fused_cache.py
-
curl -L -o test_dynamic_fused_cache.py https://huggingface.co/Aazeus/Spark-H3/resolve/main/tests/test_dynamic_fused_cache.py
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) | |
| 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) | |
| 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) | |
| 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) | |