Skip to content

Commit 4f62d31

Browse files
committed
[bugfix]: bound H3 VSA reference attention memory
1 parent 2dca66d commit 4f62d31

2 files changed

Lines changed: 30 additions & 11 deletions

File tree

fastvideo/mlx_runtime/minimax_h3_vsa.py

Lines changed: 17 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,7 @@
4848
# SIMD is opt-in until its dynamically routed output is accepted as the default.
4949
_AUTO_PREFERS_SIMD = False
5050
_REFERENCE_FULL_MASK_TILE_LIMIT = 24
51+
_REFERENCE_FULL_MASK_MAX_ELEMENTS = 256 * 1024**2
5152

5253
PrefixMode = Literal["exempt", "compete"]
5354
VSAImpl = Literal["auto", "reference", "simd"]
@@ -465,6 +466,13 @@ def _reference_token_sdpa(q_tiled, k_tiled, v_tiled, block_mask: np.ndarray, geo
465466
)[0].transpose(1, 0, 2)
466467

467468

469+
def _reference_full_mask_fits(geometry: MiniMaxH3VSAGeometry, heads: int) -> bool:
470+
"""Bound the dense token mask by both tile count and materialized size."""
471+
mask_elements = heads * geometry.padded_length**2
472+
return (geometry.num_tiles <= _REFERENCE_FULL_MASK_TILE_LIMIT
473+
and mask_elements <= _REFERENCE_FULL_MASK_MAX_ELEMENTS)
474+
475+
468476
def _gather_selected_kv(k_tiles, v_tiles, block_idx, tile_elems: int):
469477
"""Gather K/V tiles for one query-tile batch.
470478
@@ -511,7 +519,7 @@ def _reference_gather_sdpa(
511519
geometry: MiniMaxH3VSAGeometry,
512520
scale: float,
513521
):
514-
"""Grouped gather + batched SDPA over video query tiles; prefix stays dense-equivalent.
522+
"""Grouped gather + batched SDPA over video query tiles.
515523
516524
Full-sequence gather at 720p materializes tens of GiB of selected K/V, so
517525
query tiles are processed in memory-bounded chunks. This is the correctness
@@ -547,9 +555,7 @@ def _reference_gather_sdpa(
547555
mx.eval(out)
548556
chunks.append(out)
549557
out = mx.concatenate(chunks, axis=1) if len(chunks) > 1 else chunks[0]
550-
video_flat = out.transpose(1, 2, 0, 3).reshape(n_q * tile_elems, heads, dim)
551-
prefix = q_tiled[:geometry.num_prefix_tiles * tile_elems]
552-
return mx.concatenate([prefix, video_flat], axis=0)
558+
return out.transpose(1, 2, 0, 3).reshape(n_q * tile_elems, heads, dim)
553559

554560

555561
def _gate_compress_output(scores, v_tiled, gate_tiled, geometry: MiniMaxH3VSAGeometry):
@@ -642,19 +648,20 @@ def h3_vsa_attention(
642648
)
643649
chosen = resolve_impl(impl, dim, geometry.tile_elems)
644650
prefix_out = _dense_sdpa(query[:geometry.prefix_length], key, value, scale)
651+
n_prefix_pad = geometry.num_prefix_tiles * geometry.tile_elems
645652
if chosen == "simd":
646653
from fastvideo.mlx_runtime.minimax_h3_vsa_simd import simd_block_sparse
647654

648655
try:
649656
video_tiled = simd_block_sparse(q_tiled, k_tiled, v_tiled, block_idx, block_num, geometry, scale)
657+
video_tiled = video_tiled[n_prefix_pad:]
650658
except Exception as error: # noqa: BLE001 - keep generation alive on kernel failure
651659
logger.warning("SIMD VSA kernel failed (%s); falling back to reference gather+SDPA", error)
652660
chosen = "reference"
653661
if stats is not None:
654662
stats.dense_fallback_reason = f"simd kernel failed: {error}"
655663
video_tiled = _reference_gather_sdpa(q_tiled, k_tiled, v_tiled, block_idx, geometry, scale)
656-
out_tiled = video_tiled
657-
elif geometry.num_tiles <= _REFERENCE_FULL_MASK_TILE_LIMIT:
664+
elif _reference_full_mask_fits(geometry, heads):
658665
scores_np = np.array(scores, dtype=np.float32)
659666
mask = build_block_mask(
660667
scores_np,
@@ -663,10 +670,10 @@ def h3_vsa_attention(
663670
sparsity,
664671
exempt,
665672
)
666-
out_tiled = _reference_token_sdpa(q_tiled, k_tiled, v_tiled, mask, geometry, scale)
673+
video_tiled = _reference_token_sdpa(q_tiled, k_tiled, v_tiled, mask, geometry, scale)[n_prefix_pad:]
667674
else:
668-
out_tiled = _reference_gather_sdpa(q_tiled, k_tiled, v_tiled, block_idx, geometry, scale)
669-
out_tiled = out_tiled.astype(query.dtype)
675+
video_tiled = _reference_gather_sdpa(q_tiled, k_tiled, v_tiled, block_idx, geometry, scale)
676+
video_tiled = video_tiled.astype(query.dtype)
670677
prefix_tiled = _tile_hidden(
671678
mx.concatenate([
672679
prefix_out,
@@ -676,8 +683,7 @@ def h3_vsa_attention(
676683
geometry,
677684
)
678685
# Prefix query tiles are dense; keep fused-SDPA prefix rows and sparse video tiles.
679-
n_prefix_pad = geometry.num_prefix_tiles * geometry.tile_elems
680-
out_tiled = mx.concatenate([prefix_tiled[:n_prefix_pad], out_tiled[n_prefix_pad:]], axis=0)
686+
out_tiled = mx.concatenate([prefix_tiled[:n_prefix_pad], video_tiled], axis=0)
681687

682688
if gate_compress is not None:
683689
gate_tiled = _tile_hidden(gate_compress, geometry)

fastvideo/tests/mlx/test_mlx_minimax_h3_vsa.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -411,6 +411,19 @@ def test_reference_gather_path_matches_token_mask() -> None:
411411
assert mx.allclose(gather, masked, atol=2e-4, rtol=2e-4).item()
412412

413413

414+
def test_reference_full_mask_is_bounded_by_materialized_size() -> None:
415+
from fastvideo.mlx_runtime import minimax_h3_vsa as vsa_mod
416+
417+
small = build_h3_tile_geometry((8, 0, 8), (4, 8, 8), tile_size=64)
418+
assert vsa_mod._reference_full_mask_fits(small, heads=40)
419+
420+
# Twenty-four 256-token tiles pass the tile-count gate, but a 40-head
421+
# token mask would contain more than 1.5 billion elements.
422+
large = build_h3_tile_geometry((256, ), (4, 8, 184), tile_size=256)
423+
assert large.num_tiles == 24
424+
assert not vsa_mod._reference_full_mask_fits(large, heads=40)
425+
426+
414427
def test_prefix_segments_from_t2va_layout_drop_empty_condition() -> None:
415428
layout = MiniMaxH3PackedLayout(
416429
sequence_length=20,

0 commit comments

Comments
 (0)