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
5253PrefixMode = Literal ["exempt" , "compete" ]
5354VSAImpl = 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+
468476def _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
555561def _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 )
0 commit comments