File size: 18,069 Bytes
c4cbbbc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
from .NSA_select_attn_fwd_hmma import HopperSelectAttentionFwd
from cudnn.datatypes import _convert_to_cutlass_data_type
from cudnn.api_base import APIBase, TupleDict

import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_fake_stream
from cuda.bindings import driver as cuda
import torch
from typing import Tuple, Optional
import math


class SelectionAttention(APIBase):
    def __init__(
        self,
        sample_q: torch.Tensor,
        sample_k: torch.Tensor,
        sample_v: torch.Tensor,
        sample_o: torch.Tensor,
        sample_l: torch.Tensor,
        sample_m: torch.Tensor,
        sample_block_indices: torch.Tensor,
        sample_block_counts: torch.Tensor,
        sample_cum_seqlen_q: Optional[torch.Tensor] = None,
        sample_cum_seqlen_k: Optional[torch.Tensor] = None,
        max_s_q: Optional[int] = 1024,
        max_s_k: Optional[int] = 1024,
        acc_dtype: torch.dtype = torch.float32,
        block_size: int = 64,
        scale_softmax: Optional[float] = None,
    ):
        super().__init__()
        self._kernel = HopperSelectAttentionFwd

        self._logger.warning("SelectionAttention is an experimental API")
        self._logger.debug("Entering __init__")

        self.q_desc = self._make_tensor_desc(sample_q, name="sample_q")
        self.k_desc = self._make_tensor_desc(sample_k, name="sample_k")
        self.v_desc = self._make_tensor_desc(sample_v, name="sample_v")
        self.o_desc = self._make_tensor_desc(sample_o, name="sample_o")
        self.l_desc = self._make_tensor_desc(sample_l, name="sample_l")
        self.m_desc = self._make_tensor_desc(sample_m, name="sample_m")
        self.block_indices_desc = self._make_tensor_desc(sample_block_indices, name="sample_block_indices")
        self.block_counts_desc = self._make_tensor_desc(sample_block_counts, name="sample_block_counts")
        self.cum_seqlen_q_desc = self._make_tensor_desc(sample_cum_seqlen_q, name="sample_cum_seqlen_q")
        self.cum_seqlen_k_desc = self._make_tensor_desc(sample_cum_seqlen_k, name="sample_cum_seqlen_k")
        if sample_cum_seqlen_q is not None and sample_cum_seqlen_k is not None:
            if not torch.equal(sample_cum_seqlen_q, sample_cum_seqlen_k):
                raise NotImplementedError("sample_cum_seqlen_k is not yet supported. Must be None or identical to sample_cum_seqlen_q")
        self.max_s_q = max_s_q
        self.max_s_k = max_s_k

        # Types and kernel configuration
        self.acc_dtype = acc_dtype
        self.block_size = block_size

        # Derived attributes (populated in check_support)
        self.input_layout = None
        self.dtype = None
        self.h_q = None
        self.h_kv = None
        self.gqa_group_size = None
        self.head_dim = None
        self.value_dim = None

        self.scale_softmax = scale_softmax

        self._logger.debug(
            f"__init__ completed with args: sample_q {self.q_desc.shape}, sample_k {self.k_desc.shape}, sample_v {self.v_desc.shape}, sample_o {self.o_desc.shape}, sample_l {self.l_desc.shape}, sample_m {self.m_desc.shape}, sample_block_indices {self.block_indices_desc.shape}, sample_block_counts {self.block_counts_desc.shape}, sample_cum_seqlen_q {self.cum_seqlen_q_desc.shape if self.cum_seqlen_q_desc is not None else 'None'}, sample_cum_seqlen_k {self.cum_seqlen_k_desc.shape if self.cum_seqlen_k_desc is not None else 'None'}, acc_dtype {acc_dtype}, max_s_q {max_s_q}, max_s_k {max_s_k}, block_size {block_size}, scale_softmax {scale_softmax}"
        )

    def check_support(self) -> bool:
        self._logger.debug("Entering check_support")

        # Shape normalization and validation
        self._logger.debug("Checking shape normalization and validation")
        if self.q_desc.ndim == 4:
            # B, H_q, S, D  format
            self.input_layout = "B,H,S,D"

            raise NotImplementedError("B, H_q, S, D format not implemented")
        elif self.q_desc.ndim == 3:
            # T, H_q, D  format
            self.input_layout = "T,H,D"

            t, h_q, d_qk = self.q_desc.shape
            t, h_kv, d_qk = self.k_desc.shape
            t, h_kv, d_v = self.v_desc.shape
            t, h_q, d_v = self.o_desc.shape

            self._check_tensor_shape(self.q_desc, (t, h_q, d_qk), name="Q")
            self._check_tensor_shape(self.k_desc, (t, h_kv, d_qk), name="K")
            self._check_tensor_shape(self.v_desc, (t, h_kv, d_v), name="V")
            self._check_tensor_shape(self.o_desc, (t, h_q, d_v), name="O")
            self.l_desc = self._unpad_tensor_to_ndim(self.l_desc, 2, "sample_l")
            self._check_tensor_shape(self.l_desc, (t, h_q), name="L")
            self.m_desc = self._unpad_tensor_to_ndim(self.m_desc, 2, "sample_m")
            self._check_tensor_shape(self.m_desc, (t, h_q), name="M")
            if self.cum_seqlen_q_desc is None:
                raise ValueError(f"cum_seqlen_q must be provided for T,H,D format, got {self.cum_seqlen_q_desc}")
            if self.max_s_q is None:
                raise ValueError(f"max_s_q must be provided for T,H,D format, got {self.max_s_q}")
            if self.max_s_k is not None and self.max_s_q != self.max_s_k:
                raise NotImplementedError(f"SelectionAttention requires max_s_q and max_s_k to be identical, but got {self.max_s_q} and {self.max_s_k}")

            self.batch_size = self.cum_seqlen_q_desc.shape[0] - 1
            if self.batch_size <= 0:
                raise ValueError(f"batch_size (len(cum_seqlen_q) - 1) must be greater than 0, got {self.batch_size}")
            if self.cum_seqlen_q_desc.dtype not in (torch.int32, torch.int64):
                raise ValueError(f"cum_seqlen_q must be int32 or int64, got {self.cum_seqlen_q_desc.dtype}")

            if self.block_indices_desc.shape[:2] != (t, h_kv) and self.block_indices_desc.ndim != 3:
                raise ValueError(f"block_indices shape mismatch: expected {(t, h_kv, 'K')}, got {tuple(self.block_indices_desc.shape)}")
            if self.block_counts_desc.shape != (t, h_kv):
                raise ValueError(f"block_counts shape mismatch: expected {(t, h_kv)}, got {tuple(self.block_counts_desc.shape)}")
            if self.block_indices_desc.dtype != torch.int32 or self.block_counts_desc.dtype != torch.int32:
                raise ValueError(f"block_indices and block_counts must be int32, got {self.block_indices_desc.dtype} and {self.block_counts_desc.dtype}")
        else:
            raise ValueError(f"q must be rank-3 (T,H,D) or rank-4 (B,H,S,D), got {self.q_desc.ndim}")

        # Shared derived attributes
        if h_q % h_kv != 0:
            raise ValueError("H_q must be a multiple of H_kv (GQA/MQA constraint)")
        self.h_q = h_q
        self.h_kv = h_kv
        self.gqa_group_size = h_q // h_kv
        self.head_dim = d_qk
        self.value_dim = d_v

        # Validate dtypes and config
        self._logger.debug("Checking dtypes and config")
        self.dtype = self._check_dtype(self.q_desc, dtype=[torch.float16, torch.bfloat16], name="Q")
        _ = self._check_dtype(self.k_desc, dtype=self.dtype, name="K", extra_error_msg="K must have the same dtype as Q")
        _ = self._check_dtype(self.v_desc, dtype=self.dtype, name="V", extra_error_msg="V must have the same dtype as Q")
        _ = self._check_dtype(self.o_desc, dtype=self.dtype, name="O", extra_error_msg="O must have the same dtype as Q")
        _ = self._check_dtype(self.acc_dtype, dtype=torch.float32, name="Acc", extra_error_msg="acc_dtype must be Float32")

        if self.block_size not in {16, 32, 64}:
            raise ValueError("block_size must be 16, 32, or 64")

        # Compute default scale_softmax if needed
        if self.scale_softmax is None:
            self.scale_softmax = 1.0 / math.sqrt(self.head_dim)

        if not torch.cuda.is_available():
            self._logger.error("CUDA is not available")
            raise RuntimeError("CUDA is not available")

        self._logger.debug("Checking environment")
        device = torch.cuda.current_device()
        major, minor = torch.cuda.get_device_capability(device)
        compute_capability = major * 10 + minor
        if compute_capability < 90:
            self._logger.error(f"Requires SM90+ compute capability, but found SM{compute_capability} on device {device}")
            raise RuntimeError(f"Requires SM90+ compute capability, but found SM{compute_capability} on device {device}")
        if compute_capability == 103:
            raise RuntimeError("cuteDSL SelectionAttention is not supported on SM103")

        self._is_supported = True
        self._logger.debug("check_support completed successfully")
        return True

    def compile(self) -> None:
        self._logger.debug("Entering compile")
        self._ensure_support_checked()
        if self._compiled_kernel is not None:
            self._logger.debug("Kernel already compiled; skipping recompilation")
            return

        selection_attention = self._kernel(
            head_dim=self.head_dim,
            value_dim=self.value_dim,
            GQA_group_size=self.gqa_group_size,
            block_size=self.block_size,
            dtype=_convert_to_cutlass_data_type(self.dtype),
            acc_dtype=_convert_to_cutlass_data_type(self.acc_dtype),
        )

        if self.input_layout == "T,H,D":
            _q_desc = self.q_desc.unsqueeze(0)
            _k_desc = self.k_desc.unsqueeze(0)
            _v_desc = self.v_desc.unsqueeze(0)
            _o_desc = self.o_desc.unsqueeze(0)
            _l_desc = self.l_desc.unsqueeze(0)
            _m_desc = self.m_desc.unsqueeze(0)
        else:
            raise NotImplementedError(f"Invalid input layout: {self.input_layout}")

        fake_stream = make_fake_stream(use_tvm_ffi_env_stream=False)

        self._logger.debug("Compiling selection_attention")
        _compiled_kernel = cute.compile(
            selection_attention,
            Q=self._make_fake_cute_tensor_from_desc(_q_desc, assumed_align=128),
            K=self._make_fake_cute_tensor_from_desc(_k_desc, assumed_align=128),
            V=self._make_fake_cute_tensor_from_desc(_v_desc, assumed_align=128),
            O=self._make_fake_cute_tensor_from_desc(_o_desc, assumed_align=128),
            L=self._make_fake_cute_tensor_from_desc(_l_desc),
            M=self._make_fake_cute_tensor_from_desc(_m_desc),
            block_indices=self._make_fake_cute_tensor_from_desc(self.block_indices_desc),
            block_counts=self._make_fake_cute_tensor_from_desc(self.block_counts_desc),
            max_length=self.max_s_q,
            seq_offsets=self._make_fake_cute_tensor_from_desc(self.cum_seqlen_q_desc),
            softmax_scale=self.scale_softmax,
            stream=fake_stream,
            options="--enable-tvm-ffi",
        )

        def tensor_api(
            q_tensor,
            k_tensor,
            v_tensor,
            o_tensor,
            l_tensor,
            m_tensor,
            block_indices_tensor,
            block_counts_tensor,
            cum_seqlen_q_tensor,
            softmax_scale,
            stream,
        ):
            # assumed T,H,D format
            q_tensor = q_tensor.unsqueeze(0)
            k_tensor = k_tensor.unsqueeze(0)
            v_tensor = v_tensor.unsqueeze(0)
            o_tensor = o_tensor.unsqueeze(0)
            l_tensor = self._unpad_tensor_to_ndim(l_tensor, 2, "l_tensor").unsqueeze(0)
            m_tensor = self._unpad_tensor_to_ndim(m_tensor, 2, "m_tensor").unsqueeze(0)

            return _compiled_kernel(
                q_tensor,
                k_tensor,
                v_tensor,
                o_tensor,
                l_tensor,
                m_tensor,
                block_indices_tensor,
                block_counts_tensor,
                self.max_s_q,
                cum_seqlen_q_tensor,
                softmax_scale,
                stream,
            )

        self._compiled_kernel = tensor_api

        self._logger.debug("Kernel compiled successfully")

    def execute(
        self,
        q_tensor: torch.Tensor,
        k_tensor: torch.Tensor,
        v_tensor: torch.Tensor,
        o_tensor: torch.Tensor,
        l_tensor: torch.Tensor,
        m_tensor: torch.Tensor,
        block_indices_tensor: torch.Tensor,
        block_counts_tensor: torch.Tensor,
        cum_seqlen_q_tensor: Optional[torch.Tensor] = None,
        cum_seqlen_k_tensor: Optional[torch.Tensor] = None,
        scale_softmax: Optional[float] = None,
        current_stream: Optional[cuda.CUstream] = None,
    ):
        self._logger.debug("Entering execute")
        current_stream = self._get_default_stream(current_stream)

        scale_softmax = self.scale_softmax if scale_softmax is None else scale_softmax

        if self._compiled_kernel is None:
            raise RuntimeError("SelectionAttention kernel not compiled")
        self._logger.debug("Executing with compiled kernel")
        self._compiled_kernel(
            q_tensor=q_tensor,
            k_tensor=k_tensor,
            v_tensor=v_tensor,
            o_tensor=o_tensor,
            l_tensor=l_tensor,
            m_tensor=m_tensor,
            block_indices_tensor=block_indices_tensor,
            block_counts_tensor=block_counts_tensor,
            cum_seqlen_q_tensor=cum_seqlen_q_tensor,
            softmax_scale=scale_softmax,
            stream=current_stream,
        )
        self._logger.debug("Executed with compiled kernel successfully")


import logging

_logger = logging.getLogger(__name__)
_cache_of_SelectionAttentionObjects = {}


def selection_attention_wrapper(
    q_tensor: torch.Tensor,
    k_tensor: torch.Tensor,
    v_tensor: torch.Tensor,
    block_indices_tensor: torch.Tensor,
    block_counts_tensor: torch.Tensor,
    cum_seqlen_q_tensor: Optional[torch.Tensor] = None,
    cum_seqlen_k_tensor: Optional[torch.Tensor] = None,
    block_size: int = 64,
    scale_softmax: Optional[float] = None,
    o_dtype: Optional[torch.dtype] = None,
    acc_dtype: torch.dtype = torch.float32,
    max_s_q: Optional[int] = None,
    max_s_k: Optional[int] = None,
    stream: Optional[cuda.CUstream] = None,
) -> TupleDict:
    """
    Selection Attention Wrapper that returns output tensors.

    Returns:
        TupleDict: (o_tensor, l_tensor, m_tensor) - Output, logsumexp, and max tensors
    """
    _logger.debug("selection_attention_wrapper: Creating empty output tensors o, l, and m")

    max_s_q = max(cum_seqlen_q_tensor[1:] - cum_seqlen_q_tensor[:-1]).item() if max_s_q is None else max_s_q
    max_s_k = max(cum_seqlen_k_tensor[1:] - cum_seqlen_k_tensor[:-1]).item() if max_s_k is None else max_s_k

    t, h_q, d = q_tensor.shape
    _, h_kv, d_v = v_tensor.shape

    o_dtype = o_dtype if o_dtype is not None else q_tensor.dtype
    o_tensor = torch.empty((t, h_q, d_v), dtype=o_dtype, device=q_tensor.device)
    l_tensor = torch.empty((t, h_q, 1), dtype=torch.float32, device=q_tensor.device)
    m_tensor = torch.empty((t, h_q, 1), dtype=torch.float32, device=q_tensor.device)

    cache_key = (
        q_tensor.shape,
        k_tensor.shape,
        v_tensor.shape,
        block_indices_tensor.shape,
        block_counts_tensor.shape,
        cum_seqlen_q_tensor.shape,
        cum_seqlen_k_tensor.shape,
        q_tensor.dtype,
        k_tensor.dtype,
        v_tensor.dtype,
        q_tensor.stride(),
        k_tensor.stride(),
        v_tensor.stride(),
        block_indices_tensor.stride(),
        block_counts_tensor.stride(),
        cum_seqlen_q_tensor.stride(),
        cum_seqlen_k_tensor.stride(),
        block_size,
        scale_softmax,
        acc_dtype,
        max_s_q,
        max_s_k,
    )
    if cache_key in _cache_of_SelectionAttentionObjects:
        _logger.debug("selection_attention_wrapper: Using previously cached SelectionAttention object")
        selection_attention = _cache_of_SelectionAttentionObjects[cache_key]
        selection_attention.execute(
            q_tensor=q_tensor,
            k_tensor=k_tensor,
            v_tensor=v_tensor,
            o_tensor=o_tensor,
            l_tensor=l_tensor,
            m_tensor=m_tensor,
            block_indices_tensor=block_indices_tensor,
            block_counts_tensor=block_counts_tensor,
            cum_seqlen_q_tensor=cum_seqlen_q_tensor,
            cum_seqlen_k_tensor=cum_seqlen_k_tensor,
            scale_softmax=scale_softmax,
            current_stream=stream,
        )
    else:
        _logger.debug("selection_attention_wrapper: No previously cached SelectionAttention object found, creating new SelectionAttention object")
        selection_attention = SelectionAttention(
            sample_q=q_tensor,
            sample_k=k_tensor,
            sample_v=v_tensor,
            sample_o=o_tensor,
            sample_l=l_tensor,
            sample_m=m_tensor,
            sample_block_indices=block_indices_tensor,
            sample_block_counts=block_counts_tensor,
            sample_cum_seqlen_q=cum_seqlen_q_tensor,
            sample_cum_seqlen_k=cum_seqlen_k_tensor,
            acc_dtype=acc_dtype,
            max_s_q=max_s_q,
            max_s_k=max_s_k,
            block_size=block_size,
            scale_softmax=scale_softmax,
        )
        assert selection_attention.check_support()
        selection_attention.compile()
        selection_attention.execute(
            q_tensor=q_tensor,
            k_tensor=k_tensor,
            v_tensor=v_tensor,
            o_tensor=o_tensor,
            l_tensor=l_tensor,
            m_tensor=m_tensor,
            block_indices_tensor=block_indices_tensor,
            block_counts_tensor=block_counts_tensor,
            cum_seqlen_q_tensor=cum_seqlen_q_tensor,
            cum_seqlen_k_tensor=cum_seqlen_k_tensor,
            scale_softmax=scale_softmax,
            current_stream=stream,
        )
        _cache_of_SelectionAttentionObjects[cache_key] = selection_attention

    return TupleDict(
        o_tensor=o_tensor,
        l_tensor=l_tensor,
        m_tensor=m_tensor,
    )