File size: 6,605 Bytes
d401c91
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
46467b5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d401c91
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9b05071
4278535
d401c91
4278535
d401c91
 
7dfb205
4278535
7dfb205
4278535
7dfb205
4278535
d401c91
 
 
 
 
 
cf145d5
 
 
883dc36
cf145d5
 
 
 
 
 
 
d401c91
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
"""AOTI (Ahead-Of-Time Inductor) support for the VSA-H3 sparse attention soup.

Phase A of the AOTI experiment (torch 2.11.0+cu130 on ZeroGPU, where AOTI is officially supported):

1. The two raw Triton kernels the sparse path depends on — `map_to_index` and
   `triton_block_sparse_attn_forward` — are wrapped as `torch.library` custom ops with fake
   implementations, which is what lets `torch.export` trace through them.
2. `sparse_attention_functional` reimplements `vsa_h3.sparse_attention` as a *pure tensor function*:
   op-for-op the same math (fresh tile buffers, fp32 pooled scores, prefix-exempt top-k mask,
   sparse kernel, untile, compression branch with the trained per-row gate), but the cached
   `VSAGeometry` tensors arrive as plain arguments. No weights, no globals — small artifacts.
3. `/aoti_diag` on the Space (a) checks the functional reimplementation against the eager
   reference on the real bench geometry, and (b) when `H3_AOTI=1`, exports + AOTI-compiles the
   function for this GPU and checks the compiled artifact against eager, with timings.

The compiled function is only valid for one static layout (the frozen benchmark spec's). Any other
canvas falls back to the eager path.
"""

from __future__ import annotations

import math

import torch

_LIB = None  # registered-once torch.library fragment


def _ns() -> torch.library.Library:
    """Register the custom ops once and return the library fragment."""
    global _LIB
    if _LIB is not None:
        return _LIB

    from vsa_kernel import map_to_index as _map_to_index
    from vsa_kernel import triton_block_sparse_attn_forward as _fwd

    lib = torch.library.Library("fasth3", "FRAGMENT")

    def map_fake(block_map: torch.Tensor):
        return (
            torch.empty_like(block_map, dtype=torch.int32),
            torch.empty(block_map.shape[:-1], dtype=torch.int32, device=block_map.device),
        )

    def fwd_fake(q, k, v, q2k_index, q2k_num, vbs):
        return (
            torch.empty_like(q),
            torch.empty(q.shape[:3], dtype=torch.float32, device=q.device),
        )

    lib.define("map_to_index(Tensor block_map) -> (Tensor, Tensor)")
    lib.impl("map_to_index", _map_to_index, "CUDA")
    lib._register_fake("map_to_index", map_fake)

    lib.define(
        "block_sparse_fwd(Tensor q, Tensor k, Tensor v, Tensor q2k_index, Tensor q2k_num, "
        "Tensor vbs) -> (Tensor, Tensor)"
    )
    lib.impl("block_sparse_fwd", _fwd, "CUDA")
    lib._register_fake("block_sparse_fwd", fwd_fake)

    _LIB = lib
    return lib


def sparse_module(topk: int, num_prefix_tiles: int) -> torch.nn.Module:
    """`torch.export.export` requires an `nn.Module`; wrap the functional with baked scalars."""

    class _SparseModule(torch.nn.Module):
        def __init__(self):
            super().__init__()
            self.topk = topk
            self.num_prefix_tiles = num_prefix_tiles

        def forward(self, q, k, v, gate, untile_index, variable_block_sizes, tile_divisor):
            return sparse_attention_functional(
                q, k, v, gate, untile_index, variable_block_sizes, tile_divisor,
                self.topk, self.num_prefix_tiles,
            )

    return _SparseModule()


def sparse_attention_functional(
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    gate: torch.Tensor | None,
    untile_index: torch.Tensor,
    variable_block_sizes: torch.Tensor,
    tile_divisor: torch.Tensor,
    topk: int,
    num_prefix_tiles: int,
) -> torch.Tensor:
    """`vsa_h3.sparse_attention` with every cached buffer inlined as an argument.

    Op-for-op the same math as the eager path: fp32 pooled tile scores, top-k video tiles with the
    prefix exempt (prefix key tiles always selected, prefix query tiles dense), the sparse kernel,
    untiling, then the compression branch scaled by the trained per-row gate. `gate` is
    `[B, H, S_real, D]` (the eager path's `gate.transpose(1, 2)`), applied after untililing.
    """
    _ns()  # ensure the fasth3 custom ops exist before any torch.ops.fasth3.* call
    batch, heads, seq_len, dim = query.shape  # inputs are `[B, H, S_real, D]` (BHSD)
    n_tiles = variable_block_sizes.numel()
    padded_len = n_tiles * 64

    # Tile: scatter the packed rows into the kernels' `[B, H, S_pad, D]` layout, pad slots zero.
    query_tiled = torch.zeros((batch, heads, padded_len, dim), dtype=query.dtype, device=query.device)
    query_tiled.index_copy_(2, untile_index, query)
    key_tiled = torch.zeros_like(query_tiled)
    key_tiled.index_copy_(2, untile_index, key)
    value_tiled = torch.zeros_like(query_tiled)
    value_tiled.index_copy_(2, untile_index, value)

    # Per-head fp32 pooled tile scores.
    q_pool = query_tiled.view(batch, heads, n_tiles, 64, dim).sum(dim=3, dtype=torch.float32) / tile_divisor
    k_pool = key_tiled.view(batch, heads, n_tiles, 64, dim).sum(dim=3, dtype=torch.float32) / tile_divisor
    scores = torch.matmul(q_pool, k_pool.transpose(-2, -1)) / math.sqrt(dim)

    # Prefix-exempt top-k mask, assembled from fresh tensors (functional, export-traceable):
    # video query rows keep the top-k VIDEO kv tiles, prefix kv tiles are always selected,
    # prefix query rows are dense over everything.
    num_q_video = n_tiles - num_prefix_tiles
    video_kv_scores = scores[..., num_prefix_tiles:, num_prefix_tiles:]          # [B, H, nq_video, kv_video]
    indices = video_kv_scores.topk(topk, dim=-1, sorted=False).indices
    mask_video_kv = torch.zeros_like(video_kv_scores, dtype=torch.bool).scatter(-1, indices, True)
    prefix_cols = torch.ones(
        (batch, heads, num_q_video, num_prefix_tiles), dtype=torch.bool, device=query.device
    )
    mask_video = torch.cat([prefix_cols, mask_video_kv], dim=-1)                 # [B, H, nq_video, n_tiles]
    mask = torch.cat(
        [
            torch.ones((batch, heads, num_prefix_tiles, n_tiles), dtype=torch.bool, device=query.device),
            mask_video,
        ],
        dim=2,
    )

    q2k_index, q2k_num = torch.ops.fasth3.map_to_index(mask)
    out_tiled, _ = torch.ops.fasth3.block_sparse_fwd(
        query_tiled, key_tiled, value_tiled, q2k_index, q2k_num, variable_block_sizes
    )
    out = out_tiled.index_select(2, untile_index)

    if gate is not None:
        v_pool = value_tiled.view(batch, heads, n_tiles, 64, dim).sum(dim=3, dtype=torch.float32) / tile_divisor
        pooled = torch.matmul(torch.softmax(scores, dim=-1), v_pool).to(out.dtype)
        out = out + pooled.index_select(2, untile_index // 64) * gate

    return out.transpose(1, 2)