danieldk HF Staff commited on
Commit
65b4aab
·
1 Parent(s): 85033ba

Remove stable ABI 2.10 builds, since we have 2.9 now.

Browse files
build/torch-stable-abi210-cu128-x86_64-linux/__init__.py DELETED
@@ -1,17 +0,0 @@
1
- from .flash_attn_interface import (
2
- flash_attn_combine,
3
- flash_attn_func,
4
- flash_attn_qkvpacked_func,
5
- flash_attn_varlen_func,
6
- flash_attn_with_kvcache,
7
- get_scheduler_metadata,
8
- )
9
-
10
- __all__ = [
11
- "flash_attn_combine",
12
- "flash_attn_func",
13
- "flash_attn_qkvpacked_func",
14
- "flash_attn_varlen_func",
15
- "flash_attn_with_kvcache",
16
- "get_scheduler_metadata",
17
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu128-x86_64-linux/_flash_attn3_cuda_477ab85.abi3.so DELETED
@@ -1,3 +0,0 @@
1
- version https://git-lfs.github.com/spec/v1
2
- oid sha256:dc32f163c5e1fb6362e1a2ec165a3d2b6d622929c4a463fa23bc6e561e7e4680
3
- size 802035904
 
 
 
 
build/torch-stable-abi210-cu128-x86_64-linux/_ops.py DELETED
@@ -1,9 +0,0 @@
1
- import torch
2
- from . import _flash_attn3_cuda_477ab85
3
- ops = torch.ops._flash_attn3_cuda_477ab85
4
-
5
- def add_op_namespace_prefix(op_name: str):
6
- """
7
- Prefix op by namespace.
8
- """
9
- return f"_flash_attn3_cuda_477ab85::{op_name}"
 
 
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu128-x86_64-linux/flash_attn3/__init__.py DELETED
@@ -1,26 +0,0 @@
1
- import ctypes
2
- import importlib.util
3
- import sys
4
- from pathlib import Path
5
- from types import ModuleType
6
-
7
-
8
- def _import_from_path(file_path: Path) -> ModuleType:
9
- # We cannot use the module name as-is, after adding it to `sys.modules`,
10
- # it would also be used for other imports. So, we make a module name that
11
- # depends on the path for it to be unique using the hex-encoded hash of
12
- # the path.
13
- path_hash = "{:x}".format(ctypes.c_size_t(hash(file_path.absolute())).value)
14
- module_name = path_hash
15
- spec = importlib.util.spec_from_file_location(module_name, file_path)
16
- if spec is None:
17
- raise ImportError(f"Cannot load spec for {module_name} from {file_path}")
18
- module = importlib.util.module_from_spec(spec)
19
- if module is None:
20
- raise ImportError(f"Cannot load module {module_name} from spec")
21
- sys.modules[module_name] = module
22
- spec.loader.exec_module(module) # type: ignore
23
- return module
24
-
25
-
26
- globals().update(vars(_import_from_path(Path(__file__).parent.parent / "__init__.py")))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu128-x86_64-linux/flash_attn_config.py DELETED
@@ -1,7 +0,0 @@
1
- # Auto-generated by flash attention 3 setup.py
2
- CONFIG = {'build_flags': {'FLASHATTENTION_DISABLE_BACKWARD': False, 'FLASHATTENTION_DISABLE_SPLIT': False, 'FLASHATTENTION_DISABLE_PAGEDKV': False, 'FLASHATTENTION_DISABLE_APPENDKV': False, 'FLASHATTENTION_DISABLE_LOCAL': False, 'FLASHATTENTION_DISABLE_SOFTCAP': False, 'FLASHATTENTION_DISABLE_PACKGQA': False, 'FLASHATTENTION_DISABLE_FP16': False, 'FLASHATTENTION_DISABLE_FP8': False, 'FLASHATTENTION_DISABLE_VARLEN': False, 'FLASHATTENTION_DISABLE_CLUSTER': False, 'FLASHATTENTION_DISABLE_HDIM64': False, 'FLASHATTENTION_DISABLE_HDIM96': False, 'FLASHATTENTION_DISABLE_HDIM128': False, 'FLASHATTENTION_DISABLE_HDIM192': False, 'FLASHATTENTION_DISABLE_HDIM256': False, 'FLASHATTENTION_DISABLE_SM8x': False, 'FLASHATTENTION_ENABLE_VCOLMAJOR': False, 'FLASH_ATTENTION_DISABLE_HDIMDIFF64': False, 'FLASH_ATTENTION_DISABLE_HDIMDIFF192': False}}
3
-
4
- def show():
5
- from pprint import pprint
6
- pprint(CONFIG)
7
-
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu128-x86_64-linux/flash_attn_interface.py DELETED
@@ -1,1127 +0,0 @@
1
- # Copyright (c) 2023, Tri Dao.
2
-
3
- from typing import Optional, Union, List, Tuple
4
-
5
- import torch
6
- import torch.nn as nn
7
-
8
- from ._ops import ops as flash_attn_3_cuda
9
- from ._ops import add_op_namespace_prefix
10
-
11
- def maybe_contiguous(x):
12
- return x.contiguous() if x is not None and x.stride(-1) != 1 else x
13
-
14
-
15
- def round_multiple(x, m):
16
- return (x + m - 1) // m * m
17
-
18
-
19
- def round_up_headdim(head_size: int) -> int:
20
- from .flash_attn_config import CONFIG
21
-
22
- if not CONFIG["build_flags"]["FLASHATTENTION_DISABLE_HDIM64"]:
23
- if head_size <= 64:
24
- return 64
25
- if not CONFIG["build_flags"]["FLASHATTENTION_DISABLE_HDIM96"]:
26
- if head_size <= 96:
27
- return 96
28
- if not CONFIG["build_flags"]["FLASHATTENTION_DISABLE_HDIM128"]:
29
- if head_size <= 128:
30
- return 128
31
- if not CONFIG["build_flags"]["FLASHATTENTION_DISABLE_HDIM192"]:
32
- if head_size <= 192:
33
- return 192
34
- if not CONFIG["build_flags"]["FLASHATTENTION_DISABLE_HDIM256"]:
35
- if head_size <= 256:
36
- return 256
37
- return 256
38
-
39
-
40
- @torch.library.custom_op(add_op_namespace_prefix("_flash_attn_forward"), mutates_args=(), device_types="cuda")
41
- def _flash_attn_forward(
42
- q: torch.Tensor,
43
- k: torch.Tensor,
44
- v: torch.Tensor,
45
- k_new: Optional[torch.Tensor] = None,
46
- v_new: Optional[torch.Tensor] = None,
47
- qv: Optional[torch.Tensor] = None,
48
- out_: Optional[torch.Tensor] = None,
49
- cu_seqlens_q: Optional[torch.Tensor] = None,
50
- cu_seqlens_k: Optional[torch.Tensor] = None,
51
- cu_seqlens_k_new: Optional[torch.Tensor] = None,
52
- seqused_q: Optional[torch.Tensor] = None,
53
- seqused_k: Optional[torch.Tensor] = None,
54
- max_seqlen_q: Optional[int] = None,
55
- max_seqlen_k: Optional[int] = None,
56
- page_table: Optional[torch.Tensor] = None,
57
- kv_batch_idx: Optional[torch.Tensor] = None,
58
- leftpad_k: Optional[torch.Tensor] = None,
59
- rotary_cos: Optional[torch.Tensor] = None,
60
- rotary_sin: Optional[torch.Tensor] = None,
61
- seqlens_rotary: Optional[torch.Tensor] = None,
62
- q_descale: Optional[torch.Tensor] = None,
63
- k_descale: Optional[torch.Tensor] = None,
64
- v_descale: Optional[torch.Tensor] = None,
65
- softmax_scale: Optional[float] = None,
66
- causal: bool = False,
67
- window_size_left: int = -1,
68
- window_size_right: int = -1,
69
- attention_chunk: int = 0,
70
- softcap: float = 0.0,
71
- rotary_interleaved: bool = True,
72
- scheduler_metadata: Optional[torch.Tensor] = None,
73
- num_splits: int = 1,
74
- pack_gqa: Optional[bool] = None,
75
- sm_margin: int = 0,
76
- ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
77
- q, k, k_new, v_new = [maybe_contiguous(x) for x in (q, k, k_new, v_new)]
78
- v = v.contiguous() if v.stride(-1) != 1 and v.stride(-3) != 1 else v
79
- cu_seqlens_q, cu_seqlens_k, cu_seqlens_k_new = [
80
- maybe_contiguous(x) for x in (cu_seqlens_q, cu_seqlens_k, cu_seqlens_k_new)
81
- ]
82
- seqused_q, seqused_k = [maybe_contiguous(x) for x in (seqused_q, seqused_k)]
83
- page_table, kv_batch_idx, leftpad_k = [
84
- maybe_contiguous(x) for x in (page_table, kv_batch_idx, leftpad_k)
85
- ]
86
- rotary_cos, rotary_sin = [maybe_contiguous(x) for x in (rotary_cos, rotary_sin)]
87
- seqlens_rotary = maybe_contiguous(seqlens_rotary)
88
- out, softmax_lse, out_accum, softmax_lse_accum = flash_attn_3_cuda.fwd(
89
- q,
90
- k,
91
- v,
92
- k_new,
93
- v_new,
94
- qv,
95
- out_,
96
- cu_seqlens_q,
97
- cu_seqlens_k,
98
- cu_seqlens_k_new,
99
- seqused_q,
100
- seqused_k,
101
- max_seqlen_q,
102
- max_seqlen_k,
103
- page_table,
104
- kv_batch_idx,
105
- leftpad_k,
106
- rotary_cos,
107
- rotary_sin,
108
- seqlens_rotary,
109
- q_descale,
110
- k_descale,
111
- v_descale,
112
- softmax_scale,
113
- causal,
114
- window_size_left,
115
- window_size_right,
116
- attention_chunk,
117
- softcap,
118
- rotary_interleaved,
119
- scheduler_metadata,
120
- num_splits,
121
- pack_gqa,
122
- sm_margin,
123
- )
124
-
125
- if out_accum is None:
126
- out_accum = torch.tensor([], device=out.device)
127
-
128
- if softmax_lse_accum is None:
129
- softmax_lse_accum = torch.tensor([], device=out.device)
130
-
131
- return out, softmax_lse, out_accum, softmax_lse_accum
132
-
133
-
134
- @torch.library.register_fake(add_op_namespace_prefix("_flash_attn_forward"))
135
- def _flash_attn_forward_fake(
136
- q: torch.Tensor,
137
- k: torch.Tensor,
138
- v: torch.Tensor,
139
- k_new: Optional[torch.Tensor] = None,
140
- v_new: Optional[torch.Tensor] = None,
141
- qv: Optional[torch.Tensor] = None,
142
- out_: Optional[torch.Tensor] = None,
143
- cu_seqlens_q: Optional[torch.Tensor] = None,
144
- cu_seqlens_k: Optional[torch.Tensor] = None,
145
- cu_seqlens_k_new: Optional[torch.Tensor] = None,
146
- seqused_q: Optional[torch.Tensor] = None,
147
- seqused_k: Optional[torch.Tensor] = None,
148
- max_seqlen_q: Optional[int] = None,
149
- max_seqlen_k: Optional[int] = None,
150
- page_table: Optional[torch.Tensor] = None,
151
- kv_batch_idx: Optional[torch.Tensor] = None,
152
- leftpad_k: Optional[torch.Tensor] = None,
153
- rotary_cos: Optional[torch.Tensor] = None,
154
- rotary_sin: Optional[torch.Tensor] = None,
155
- seqlens_rotary: Optional[torch.Tensor] = None,
156
- q_descale: Optional[torch.Tensor] = None,
157
- k_descale: Optional[torch.Tensor] = None,
158
- v_descale: Optional[torch.Tensor] = None,
159
- softmax_scale: Optional[float] = None,
160
- causal: bool = False,
161
- window_size_left: int = -1,
162
- window_size_right: int = -1,
163
- attention_chunk: int = 0,
164
- softcap: float = 0.0,
165
- rotary_interleaved: bool = True,
166
- scheduler_metadata: Optional[torch.Tensor] = None,
167
- num_splits: int = 1,
168
- pack_gqa: Optional[bool] = None,
169
- sm_margin: int = 0,
170
- ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
171
- """
172
- Symbolic fake implementation of flash attention forward.
173
- Returns tensors with the correct shapes and dtypes without actual computation.
174
- """
175
-
176
- # Determine if we're in varlen mode
177
- is_varlen_q = cu_seqlens_q is not None
178
-
179
- # Get dimensions from query tensor
180
- if is_varlen_q:
181
- # varlen mode: q is (total_q, num_heads, head_size)
182
- total_q, num_heads, head_size = q.shape
183
- batch_size = cu_seqlens_q.shape[0] - 1
184
-
185
- if max_seqlen_q is None:
186
- raise ValueError("max_seqlen_q must be provided if cu_seqlens_q is provided")
187
- seqlen_q = max_seqlen_q
188
- else:
189
- # batch mode: q is (batch_size, seqlen_q, num_heads, head_size)
190
- batch_size, seqlen_q, num_heads, head_size = q.shape
191
- total_q = batch_size * q.shape[1]
192
- # Get value head dimension
193
- head_size_v = v.shape[-1]
194
-
195
- # Determine output dtype (FP8 inputs produce BF16 outputs)
196
- q_type = q.dtype
197
- if q_type == torch.float8_e4m3fn:
198
- out_dtype = torch.bfloat16
199
- else:
200
- out_dtype = q_type
201
-
202
- # Create output tensor
203
- if out_ is not None:
204
- # If out_ is provided, _flash_attn_forward becomes non-functional
205
- raise TypeError("Tracing (torch.compile/torch.export) with pre-allocated output tensor is not supported.")
206
-
207
- if is_varlen_q:
208
- out = torch.empty((total_q, num_heads, head_size_v), dtype=out_dtype, device=q.device)
209
- else:
210
- out = torch.empty((batch_size, seqlen_q, num_heads, head_size_v), dtype=out_dtype, device=q.device)
211
-
212
- # Create softmax_lse tensor
213
- if is_varlen_q:
214
- softmax_lse = torch.empty((num_heads, total_q), dtype=torch.float32, device=q.device)
215
- else:
216
- softmax_lse = torch.empty((batch_size, num_heads, seqlen_q), dtype=torch.float32, device=q.device)
217
-
218
- # TODO(guilhermeleobas): Implement "get_num_splits"
219
- # There's an heuristic to compute num_splits when "num_splits <= 0"
220
- # assert that num_splits is > 0 for now
221
- if num_splits <= 0:
222
- raise ValueError(f"tracing (torch.compile/torch.export) with num_splits <= 0 not supported. Got {num_splits=}")
223
-
224
- if num_splits > 1:
225
- if is_varlen_q:
226
- out_accum = torch.empty((num_splits, num_heads, total_q, head_size_v), dtype=torch.float32, device=q.device)
227
- softmax_lse_accum = torch.empty((num_splits, num_heads, total_q), dtype=torch.float32, device=q.device)
228
- else:
229
- out_accum = torch.empty((num_splits, batch_size, num_heads, seqlen_q, head_size_v), dtype=torch.float32, device=q.device)
230
- softmax_lse_accum = torch.empty((num_splits, batch_size, num_heads, seqlen_q), dtype=torch.float32, device=q.device)
231
- else:
232
- # Tensors are not set when num_splits < 1
233
- out_accum = torch.tensor([], device=out.device)
234
- softmax_lse_accum = torch.tensor([], device=out.device)
235
-
236
- return out, softmax_lse, out_accum, softmax_lse_accum
237
-
238
-
239
- @torch.library.custom_op(add_op_namespace_prefix("_flash_attn_backward"), mutates_args=("dq", "dk", "dv"), device_types="cuda")
240
- def _flash_attn_backward(
241
- dout: torch.Tensor,
242
- q: torch.Tensor,
243
- k: torch.Tensor,
244
- v: torch.Tensor,
245
- out: torch.Tensor,
246
- softmax_lse: torch.Tensor,
247
- cu_seqlens_q: Optional[torch.Tensor] = None,
248
- cu_seqlens_k: Optional[torch.Tensor] = None,
249
- sequed_q: Optional[torch.Tensor] = None,
250
- sequed_k: Optional[torch.Tensor] = None,
251
- max_seqlen_q: Optional[int] = None,
252
- max_seqlen_k: Optional[int] = None,
253
- dq: Optional[torch.Tensor] = None,
254
- dk: Optional[torch.Tensor] = None,
255
- dv: Optional[torch.Tensor] = None,
256
- softmax_scale: Optional[float] = None,
257
- is_causal: bool = False,
258
- window_size_left: int = -1,
259
- window_size_right: int = -1,
260
- softcap: float = 0.0,
261
- deterministic: bool = False,
262
- sm_margin: int = 0,
263
- ) -> torch.Tensor:
264
- # dq, dk, dv are allocated by us so they should already be contiguous
265
- dout, q, k, v, out = [maybe_contiguous(x) for x in (dout, q, k, v, out)]
266
- softmax_d, *rest = flash_attn_3_cuda.bwd(
267
- dout,
268
- q,
269
- k,
270
- v,
271
- out,
272
- softmax_lse,
273
- dq,
274
- dk,
275
- dv,
276
- cu_seqlens_q,
277
- cu_seqlens_k,
278
- sequed_q,
279
- sequed_k,
280
- max_seqlen_q,
281
- max_seqlen_k,
282
- softmax_scale,
283
- is_causal,
284
- window_size_left,
285
- window_size_right,
286
- softcap,
287
- deterministic,
288
- sm_margin,
289
- )
290
- return softmax_d
291
-
292
-
293
- @torch.library.register_fake(add_op_namespace_prefix("_flash_attn_backward"))
294
- def _flash_attn_backward_fake(
295
- dout: torch.Tensor,
296
- q: torch.Tensor,
297
- k: torch.Tensor,
298
- v: torch.Tensor,
299
- out: torch.Tensor,
300
- softmax_lse: torch.Tensor,
301
- cu_seqlens_q: Optional[torch.Tensor] = None,
302
- cu_seqlens_k: Optional[torch.Tensor] = None,
303
- sequed_q: Optional[torch.Tensor] = None,
304
- sequed_k: Optional[torch.Tensor] = None,
305
- max_seqlen_q: Optional[int] = None,
306
- max_seqlen_k: Optional[int] = None,
307
- dq: Optional[torch.Tensor] = None,
308
- dk: Optional[torch.Tensor] = None,
309
- dv: Optional[torch.Tensor] = None,
310
- softmax_scale: Optional[float] = None,
311
- is_causal: bool = False,
312
- window_size_left: int = -1,
313
- window_size_right: int = -1,
314
- softcap: float = 0.0,
315
- deterministic: bool = False,
316
- sm_margin: int = 0,
317
- ) -> torch.Tensor:
318
-
319
- is_varlen_q = cu_seqlens_q is not None
320
- is_varlen_k = cu_seqlens_q is not None
321
- is_varlen = is_varlen_q or is_varlen_k or sequed_q is not None or sequed_k is not None
322
-
323
- if not is_varlen_q:
324
- batch_size = q.size(0)
325
- seqlen_q = q.size(1)
326
- seqlen_k = k.size(1)
327
- total_q = batch_size * q.size(1)
328
- else:
329
- batch_size = cu_seqlens_q.size(0) - 1
330
- total_q = q.size(0)
331
- seqlen_q = max_seqlen_q
332
- seqlen_k = max_seqlen_k
333
-
334
- if window_size_left >= seqlen_k - 1:
335
- window_size_left = -1
336
-
337
- if window_size_right >= seqlen_q - 1:
338
- window_size_right = -1
339
-
340
- if is_causal:
341
- window_size_right = 0
342
-
343
- is_causal = window_size_left < 0 and window_size_right == 0
344
-
345
- head_size = q.size(-1)
346
- head_size_v = v.size(-1)
347
- head_size_rounded = round_up_headdim(max(head_size, head_size_v))
348
-
349
- # Hopper gpus uses cuda compute capabilities 9.0
350
- cap = torch.cuda.get_device_capability(q.device)
351
- arch = cap[0] * 10 + cap[1]
352
-
353
- is_local = (window_size_left >= 0 or window_size_right >= 0) and not is_causal
354
-
355
- if head_size_rounded <= 64:
356
- kBlockM_sm90 = 96 if (is_causal and softcap > 0.0) else 128
357
- elif head_size_rounded <= 96:
358
- kBlockM_sm90 = 64
359
- elif head_size_rounded <= 128:
360
- kBlockM_sm90 = 64 if (is_causal or is_local or softcap > 0.0) else 80
361
- else:
362
- kBlockM_sm90 = 64
363
-
364
- kBlockM_sm80 = 128 if head_size_rounded <= 64 else 64
365
- kBlockM_sm86 = 64 if head_size_rounded <= 192 else 32
366
-
367
- if arch >= 90:
368
- kBlockM = kBlockM_sm90
369
- elif arch == 86 or arch == 89:
370
- kBlockM = kBlockM_sm86
371
- else:
372
- kBlockM = kBlockM_sm80
373
-
374
- num_heads = q.shape[-2]
375
- seqlen_q_rounded = round_multiple(seqlen_q, kBlockM)
376
-
377
- total_q_padded_rounded = round_multiple(total_q + batch_size * kBlockM, kBlockM)
378
-
379
- dq = torch.empty_like(q) if dq is None else dq
380
- dk = torch.empty_like(k) if dk is None else dk
381
- dv = torch.empty_like(v) if dv is None else dv
382
-
383
- if not is_varlen:
384
- softmax_d = torch.empty((batch_size, num_heads, seqlen_q_rounded), dtype=torch.float32, device=q.device)
385
- else:
386
- softmax_d = torch.empty((num_heads, total_q_padded_rounded), dtype=torch.float32, device=q.device)
387
-
388
- return softmax_d
389
-
390
-
391
- def setup_context(ctx, inputs, output):
392
- q, k, v = inputs[:3]
393
- out, softmax_lse, _, _ = output
394
- ctx.save_for_backward(q, k, v, out, softmax_lse)
395
- ctx.softmax_scale = inputs[-11]
396
- ctx.causal = inputs[-10]
397
- ctx.window_size = [inputs[-9], inputs[-8]]
398
- ctx.attention_chunk = inputs[-7]
399
- ctx.softcap = inputs[-6]
400
- ctx.sm_margin = inputs[-1]
401
-
402
-
403
- def _backward(ctx, dout, *grads):
404
- q, k, v, out, softmax_lse = ctx.saved_tensors
405
- dq, dk, dv = torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
406
- _flash_attn_backward(
407
- dout,
408
- q,
409
- k,
410
- v,
411
- out,
412
- softmax_lse,
413
- None, None, # cu_seqlens_q, cu_seqlens_k,
414
- None, None, # sequed_q, sequed_k,
415
- None, None, # max_seqlen_q, max_seqlen_k,
416
- dq,
417
- dk,
418
- dv,
419
- ctx.softmax_scale,
420
- ctx.causal,
421
- ctx.window_size[0],
422
- ctx.window_size[1],
423
- ctx.softcap,
424
- False, # deterministic
425
- ctx.sm_margin,
426
- )
427
- return dq, dk, dv, *((None,) * 21)
428
-
429
-
430
- _flash_attn_forward.register_autograd(_backward, setup_context=setup_context)
431
-
432
-
433
-
434
- class FlashAttnQKVPackedFunc(torch.autograd.Function):
435
- @staticmethod
436
- def forward(
437
- ctx,
438
- qkv,
439
- softmax_scale,
440
- causal,
441
- q_descale=None, k_descale=None, v_descale=None,
442
- window_size=(-1, -1),
443
- attention_chunk=0,
444
- softcap=0.0,
445
- deterministic=False,
446
- num_heads_q=None,
447
- sm_margin=0,
448
- return_softmax=False,
449
- ):
450
- if softmax_scale is None:
451
- softmax_scale = qkv.shape[-1] ** (-0.5)
452
- if qkv.dim() == 5:
453
- assert qkv.shape[-3] == 3
454
- q, k, v = qkv.unbind(dim=-3)
455
- else:
456
- assert qkv.dim() == 4
457
- assert num_heads_q is not None
458
- num_heads_k = (qkv.shape[2] - num_heads_q) // 2
459
- assert num_heads_k * 2 + num_heads_q == qkv.shape[2]
460
- q, k, v = qkv.split([num_heads_q, num_heads_k, num_heads_k], dim=-2)
461
- out, softmax_lse, *rest = _flash_attn_forward(
462
- q,
463
- k,
464
- v,
465
- None, None, # k_new, v_new
466
- None, # qv
467
- None, # out
468
- None, None, None, # cu_seqlens_q/k/k_new
469
- None, None, # seqused_q/k
470
- None, None, # max_seqlen_q/k
471
- None, None, None, # page_table, kv_batch_idx, leftpad_k,
472
- None, None, None, # rotary_cos/sin, seqlens_rotary
473
- q_descale, k_descale, v_descale,
474
- softmax_scale,
475
- causal=causal,
476
- window_size_left=window_size[0],
477
- window_size_right=window_size[1],
478
- attention_chunk=attention_chunk,
479
- softcap=softcap,
480
- sm_margin=sm_margin,
481
- )
482
- # ctx.save_for_backward(q, k, v, out_padded, softmax_lse)
483
- ctx.save_for_backward(q, k, v, out, softmax_lse)
484
- ctx.softmax_scale = softmax_scale
485
- ctx.causal = causal
486
- ctx.window_size = window_size
487
- ctx.attention_chunk = attention_chunk
488
- ctx.softcap = softcap
489
- ctx.deterministic = deterministic
490
- ctx.ndim = qkv.dim()
491
- ctx.sm_margin = sm_margin
492
- return (out, softmax_lse) if return_softmax else out
493
-
494
- @staticmethod
495
- def backward(ctx, dout, *args):
496
- q, k, v, out, softmax_lse = ctx.saved_tensors
497
- assert ctx.attention_chunk == 0, "FA3 backward does not support attention_chunk"
498
- if ctx.ndim == 5:
499
- qkv_shape = q.shape[:-2] + (3, *q.shape[-2:])
500
- dqkv = torch.empty(qkv_shape, dtype=q.dtype, device=q.device)
501
- dq, dk, dv = dqkv.unbind(dim=-3)
502
- else:
503
- num_heads_q = q.shape[2]
504
- num_heads_k = k.shape[2]
505
- qkv_shape = q.shape[:-2] + (num_heads_q + num_heads_k * 2, *q.shape[-1:])
506
- dqkv = torch.empty(qkv_shape, dtype=q.dtype, device=q.device)
507
- dq, dk, dv = dqkv.split([num_heads_q, num_heads_k, num_heads_k], dim=-2)
508
- _flash_attn_backward(
509
- dout,
510
- q,
511
- k,
512
- v,
513
- out,
514
- softmax_lse,
515
- None, None, # cu_seqlens_q, cu_seqlens_k,
516
- None, None, # sequed_q, sequed_k,
517
- None, None, # max_seqlen_q, max_seqlen_k,
518
- dq,
519
- dk,
520
- dv,
521
- ctx.softmax_scale,
522
- ctx.causal,
523
- ctx.window_size[0],
524
- ctx.window_size[1],
525
- ctx.softcap,
526
- ctx.deterministic,
527
- ctx.sm_margin,
528
- )
529
- dqkv = dqkv[..., : dout.shape[-1]] # We could have padded the head dimension
530
- return dqkv, None, None, None, None, None, None, None, None, None, None, None, None
531
-
532
-
533
- class FlashAttnFunc(torch.autograd.Function):
534
-
535
- @staticmethod
536
- def forward(
537
- ctx,
538
- q,
539
- k,
540
- v,
541
- softmax_scale,
542
- causal,
543
- qv=None,
544
- q_descale=None, k_descale=None, v_descale=None,
545
- window_size=(-1, -1),
546
- attention_chunk=0,
547
- softcap=0.0,
548
- num_splits=1,
549
- pack_gqa=None,
550
- deterministic=False,
551
- sm_margin=0,
552
- return_softmax=False,
553
- ):
554
- if softmax_scale is None:
555
- softmax_scale = (q.shape[-1] + (qv.shape[-1] if qv is not None else 0)) ** (-0.5)
556
- # out, q, k, v, out_padded, softmax_lse = _flash_attn_forward(
557
- out, softmax_lse, *rest = _flash_attn_forward(
558
- q,
559
- k,
560
- v,
561
- None, None, # k_new, v_new
562
- qv, # qv
563
- None, # out
564
- None, None, None, # cu_seqlens_q/k/k_new
565
- None, None, # seqused_q/k
566
- None, None, # max_seqlen_q/k
567
- None, None, None, # page_table, kv_batch_idx, leftpad_k,
568
- None, None, None, # rotary_cos/sin, seqlens_rotary
569
- q_descale, k_descale, v_descale,
570
- softmax_scale,
571
- causal=causal,
572
- window_size_left=window_size[0],
573
- window_size_right=window_size[1],
574
- attention_chunk=attention_chunk,
575
- softcap=softcap,
576
- num_splits=num_splits,
577
- pack_gqa=pack_gqa,
578
- sm_margin=sm_margin,
579
- )
580
- # ctx.save_for_backward(q, k, v, out_padded, softmax_lse)
581
- ctx.save_for_backward(q, k, v, out, softmax_lse)
582
- ctx.softmax_scale = softmax_scale
583
- ctx.causal = causal
584
- ctx.window_size = window_size
585
- ctx.attention_chunk = attention_chunk
586
- ctx.softcap = softcap
587
- ctx.deterministic = deterministic
588
- ctx.sm_margin = sm_margin
589
- return (out, softmax_lse) if return_softmax else out
590
-
591
- @staticmethod
592
- def backward(ctx, dout, *args):
593
- q, k, v, out, softmax_lse = ctx.saved_tensors
594
- assert ctx.attention_chunk == 0, "FA3 backward does not support attention_chunk"
595
- dq, dk, dv = torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
596
- _flash_attn_backward(
597
- dout,
598
- q,
599
- k,
600
- v,
601
- out,
602
- softmax_lse,
603
- None, None, # cu_seqlens_q, cu_seqlens_k,
604
- None, None, # sequed_q, sequed_k,
605
- None, None, # max_seqlen_q, max_seqlen_k,
606
- dq,
607
- dk,
608
- dv,
609
- ctx.softmax_scale,
610
- ctx.causal,
611
- ctx.window_size[0],
612
- ctx.window_size[1],
613
- ctx.softcap,
614
- ctx.deterministic,
615
- ctx.sm_margin,
616
- )
617
- dq = dq[..., : q.shape[-1]] # We could have padded the head dimension
618
- dk = dk[..., : k.shape[-1]]
619
- dv = dv[..., : v.shape[-1]]
620
- return dq, dk, dv, None, None, None, None, None, None, None, None, None, None, None, None, None, None
621
-
622
-
623
- class FlashAttnVarlenFunc(torch.autograd.Function):
624
-
625
- @staticmethod
626
- def forward(
627
- ctx,
628
- q,
629
- k,
630
- v,
631
- cu_seqlens_q,
632
- cu_seqlens_k,
633
- seqused_q,
634
- seqused_k,
635
- max_seqlen_q,
636
- max_seqlen_k,
637
- softmax_scale,
638
- causal,
639
- qv=None,
640
- q_descale=None, k_descale=None, v_descale=None,
641
- window_size=(-1, -1),
642
- attention_chunk=0,
643
- softcap=0.0,
644
- num_splits=1,
645
- pack_gqa=None,
646
- deterministic=False,
647
- sm_margin=0,
648
- return_softmax=False,
649
- ):
650
- if softmax_scale is None:
651
- softmax_scale = (q.shape[-1] + (qv.shape[-1] if qv is not None else 0)) ** (-0.5)
652
- # out, q, k, v, out_padded, softmax_lse = _flash_attn_varlen_forward(
653
- out, softmax_lse, *rest = _flash_attn_forward(
654
- q,
655
- k,
656
- v,
657
- None, None, # k_new, v_new
658
- qv, # qv
659
- None, # out
660
- cu_seqlens_q,
661
- cu_seqlens_k,
662
- None, # cu_seqlens_k_new
663
- seqused_q,
664
- seqused_k,
665
- max_seqlen_q,
666
- max_seqlen_k,
667
- None, None, None, # page_table, kv_batch_idx, leftpad_k,
668
- None, None, None, # rotary_cos/sin, seqlens_rotary
669
- q_descale, k_descale, v_descale,
670
- softmax_scale,
671
- causal=causal,
672
- window_size_left=window_size[0],
673
- window_size_right=window_size[1],
674
- attention_chunk=attention_chunk,
675
- softcap=softcap,
676
- num_splits=num_splits,
677
- pack_gqa=pack_gqa,
678
- sm_margin=sm_margin,
679
- )
680
- # ctx.save_for_backward(q, k, v, out_padded, softmax_lse, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k)
681
- ctx.save_for_backward(q, k, v, out, softmax_lse, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k)
682
- ctx.max_seqlen_q = max_seqlen_q
683
- ctx.max_seqlen_k = max_seqlen_k
684
- ctx.softmax_scale = softmax_scale
685
- ctx.causal = causal
686
- ctx.window_size = window_size
687
- ctx.attention_chunk = attention_chunk
688
- ctx.softcap = softcap
689
- ctx.deterministic = deterministic
690
- ctx.sm_margin = sm_margin
691
- return (out, softmax_lse) if return_softmax else out
692
-
693
- @staticmethod
694
- def backward(ctx, dout, *args):
695
- q, k, v, out, softmax_lse, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k = ctx.saved_tensors
696
- assert ctx.attention_chunk == 0, "FA3 backward does not support attention_chunk"
697
- dq, dk, dv = torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
698
- _flash_attn_backward(
699
- dout,
700
- q,
701
- k,
702
- v,
703
- out,
704
- softmax_lse,
705
- cu_seqlens_q,
706
- cu_seqlens_k,
707
- seqused_q,
708
- seqused_k,
709
- ctx.max_seqlen_q,
710
- ctx.max_seqlen_k,
711
- dq,
712
- dk,
713
- dv,
714
- ctx.softmax_scale,
715
- ctx.causal,
716
- ctx.window_size[0],
717
- ctx.window_size[1],
718
- ctx.softcap,
719
- ctx.deterministic,
720
- ctx.sm_margin,
721
- )
722
- dq = dq[..., : q.shape[-1]] # We could have padded the head dimension
723
- dk = dk[..., : k.shape[-1]]
724
- dv = dv[..., : v.shape[-1]]
725
- return dq, dk, dv, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None
726
-
727
-
728
- def flash_attn_qkvpacked_func(
729
- qkv,
730
- softmax_scale=None,
731
- causal=False,
732
- q_descale=None, k_descale=None, v_descale=None,
733
- window_size=(-1, -1),
734
- attention_chunk=0,
735
- softcap=0.0,
736
- deterministic=False,
737
- num_heads_q=None,
738
- sm_margin=0,
739
- return_attn_probs=False,
740
- ):
741
- """dropout_p should be set to 0.0 during evaluation
742
- If Q, K, V are already stacked into 1 tensor, this function will be faster than
743
- calling flash_attn_func on Q, K, V since the backward pass avoids explicit concatenation
744
- of the gradients of Q, K, V.
745
- For multi-query and grouped-query attention (MQA/GQA), please see
746
- flash_attn_kvpacked_func and flash_attn_func.
747
-
748
- If window_size != (-1, -1), implements sliding window local attention. Query at position i
749
- will only attend to keys between [i - window_size[0], i + window_size[1]] inclusive.
750
-
751
- Arguments:
752
- qkv: (batch_size, seqlen, 3, nheads, headdim)
753
- dropout_p: float. Dropout probability.
754
- softmax_scale: float. The scaling of QK^T before applying softmax.
755
- Default to 1 / sqrt(headdim).
756
- causal: bool. Whether to apply causal attention mask (e.g., for auto-regressive modeling).
757
- window_size: (left, right). If not (-1, -1), implements sliding window local attention.
758
- softcap: float. Anything > 0 activates softcapping attention.
759
- alibi_slopes: (nheads,) or (batch_size, nheads), fp32. A bias of (-alibi_slope * |i - j|) is added to
760
- the attention score of query i and key j.
761
- deterministic: bool. Whether to use the deterministic implementation of the backward pass,
762
- which is slightly slower and uses more memory. The forward pass is always deterministic.
763
- return_attn_probs: bool. Whether to return the attention probabilities. This option is for
764
- testing only. The returned probabilities are not guaranteed to be correct
765
- (they might not have the right scaling).
766
- Return:
767
- out: (batch_size, seqlen, nheads, headdim).
768
- softmax_lse [optional, if return_attn_probs=True]: (batch_size, nheads, seqlen). The
769
- logsumexp of each row of the matrix QK^T * scaling (e.g., log of the softmax
770
- normalization factor).
771
- S_dmask [optional, if return_attn_probs=True]: (batch_size, nheads, seqlen, seqlen).
772
- The output of softmax (possibly with different scaling). It also encodes the dropout
773
- pattern (negative means that location was dropped, nonnegative means it was kept).
774
- """
775
- return FlashAttnQKVPackedFunc.apply(
776
- qkv,
777
- softmax_scale,
778
- causal,
779
- q_descale, k_descale, v_descale,
780
- window_size,
781
- attention_chunk,
782
- softcap,
783
- deterministic,
784
- num_heads_q,
785
- sm_margin,
786
- return_attn_probs,
787
- )
788
-
789
-
790
- def flash_attn_func(
791
- q,
792
- k,
793
- v,
794
- softmax_scale=None,
795
- causal=False,
796
- qv=None,
797
- q_descale=None, k_descale=None, v_descale=None,
798
- window_size=(-1, -1),
799
- attention_chunk=0,
800
- softcap=0.0,
801
- num_splits=1,
802
- pack_gqa=None,
803
- deterministic=False,
804
- sm_margin=0,
805
- return_attn_probs=False,
806
- ):
807
- """dropout_p should be set to 0.0 during evaluation
808
- Supports multi-query and grouped-query attention (MQA/GQA) by passing in KV with fewer heads
809
- than Q. Note that the number of heads in Q must be divisible by the number of heads in KV.
810
- For example, if Q has 6 heads and K, V have 2 heads, head 0, 1, 2 of Q will attention to head
811
- 0 of K, V, and head 3, 4, 5 of Q will attention to head 1 of K, V.
812
-
813
- If causal=True, the causal mask is aligned to the bottom right corner of the attention matrix.
814
- For example, if seqlen_q = 2 and seqlen_k = 5, the causal mask (1 = keep, 0 = masked out) is:
815
- 1 1 1 1 0
816
- 1 1 1 1 1
817
- If seqlen_q = 5 and seqlen_k = 2, the causal mask is:
818
- 0 0
819
- 0 0
820
- 0 0
821
- 1 0
822
- 1 1
823
- If the row of the mask is all zero, the output will be zero.
824
-
825
- If window_size != (-1, -1), implements sliding window local attention. Query at position i
826
- will only attend to keys between
827
- [i + seqlen_k - seqlen_q - window_size[0], i + seqlen_k - seqlen_q + window_size[1]] inclusive.
828
-
829
- Arguments:
830
- q: (batch_size, seqlen, nheads, headdim)
831
- k: (batch_size, seqlen, nheads_k, headdim)
832
- v: (batch_size, seqlen, nheads_k, headdim)
833
- dropout_p: float. Dropout probability.
834
- softmax_scale: float. The scaling of QK^T before applying softmax.
835
- Default to 1 / sqrt(headdim).
836
- causal: bool. Whether to apply causal attention mask (e.g., for auto-regressive modeling).
837
- window_size: (left, right). If not (-1, -1), implements sliding window local attention.
838
- alibi_slopes: (nheads,) or (batch_size, nheads), fp32. A bias of
839
- (-alibi_slope * |i + seqlen_k - seqlen_q - j|)
840
- is added to the attention score of query i and key j.
841
- deterministic: bool. Whether to use the deterministic implementation of the backward pass,
842
- which is slightly slower and uses more memory. The forward pass is always deterministic.
843
- return_attn_probs: bool. Whether to return the attention probabilities. This option is for
844
- testing only. The returned probabilities are not guaranteed to be correct
845
- (they might not have the right scaling).
846
- Return:
847
- out: (batch_size, seqlen, nheads, headdim).
848
- softmax_lse [optional, if return_attn_probs=True]: (batch_size, nheads, seqlen). The
849
- logsumexp of each row of the matrix QK^T * scaling (e.g., log of the softmax
850
- normalization factor).
851
- """
852
- return FlashAttnFunc.apply(
853
- q,
854
- k,
855
- v,
856
- softmax_scale,
857
- causal,
858
- qv,
859
- q_descale, k_descale, v_descale,
860
- window_size,
861
- attention_chunk,
862
- softcap,
863
- num_splits,
864
- pack_gqa,
865
- deterministic,
866
- sm_margin,
867
- return_attn_probs,
868
- )
869
-
870
-
871
- def flash_attn_varlen_func(
872
- q,
873
- k,
874
- v,
875
- cu_seqlens_q,
876
- cu_seqlens_k,
877
- max_seqlen_q,
878
- max_seqlen_k,
879
- seqused_q=None,
880
- seqused_k=None,
881
- softmax_scale=None,
882
- causal=False,
883
- qv=None,
884
- q_descale=None, k_descale=None, v_descale=None,
885
- window_size=(-1, -1),
886
- attention_chunk=0,
887
- softcap=0.0,
888
- num_splits=1,
889
- pack_gqa=None,
890
- deterministic=False,
891
- sm_margin=0,
892
- return_attn_probs=False,
893
- ):
894
- return FlashAttnVarlenFunc.apply(
895
- q,
896
- k,
897
- v,
898
- cu_seqlens_q,
899
- cu_seqlens_k,
900
- seqused_q,
901
- seqused_k,
902
- max_seqlen_q,
903
- max_seqlen_k,
904
- softmax_scale,
905
- causal,
906
- qv,
907
- q_descale, k_descale, v_descale,
908
- window_size,
909
- attention_chunk,
910
- softcap,
911
- num_splits,
912
- pack_gqa,
913
- deterministic,
914
- sm_margin,
915
- return_attn_probs,
916
- )
917
-
918
-
919
- def flash_attn_combine(out_partial, lse_partial, out=None, out_dtype=None):
920
- return flash_attn_3_cuda.fwd_combine(out_partial, lse_partial, out, out_dtype)
921
-
922
-
923
- def flash_attn_with_kvcache(
924
- q,
925
- k_cache,
926
- v_cache,
927
- k=None,
928
- v=None,
929
- qv=None,
930
- rotary_cos=None,
931
- rotary_sin=None,
932
- cache_seqlens: Optional[Union[(int, torch.Tensor)]] = None,
933
- cache_batch_idx: Optional[torch.Tensor] = None,
934
- cache_leftpad: Optional[torch.Tensor] = None,
935
- page_table: Optional[torch.Tensor] = None,
936
- cu_seqlens_q: Optional[torch.Tensor] = None,
937
- cu_seqlens_k_new: Optional[torch.Tensor] = None,
938
- max_seqlen_q: Optional[int] = None,
939
- rotary_seqlens: Optional[torch.Tensor] = None,
940
- q_descale: Optional[torch.Tensor] = None,
941
- k_descale: Optional[torch.Tensor] = None,
942
- v_descale: Optional[torch.Tensor] = None,
943
- softmax_scale=None,
944
- causal=False,
945
- window_size=(-1, -1), # -1 means infinite context window
946
- attention_chunk=0,
947
- softcap=0.0, # 0.0 means deactivated
948
- rotary_interleaved=True,
949
- scheduler_metadata=None,
950
- num_splits=0, # Can be tuned for speed
951
- pack_gqa=None, # Can be tuned for speed
952
- sm_margin=0, # Can be tuned if some SMs are used for communication
953
- return_softmax_lse=False,
954
- ):
955
- """
956
- If k and v are not None, k_cache and v_cache will be updated *inplace* with the new values from
957
- k and v. This is useful for incremental decoding: you can pass in the cached keys/values from
958
- the previous step, and update them with the new keys/values from the current step, and do
959
- attention with the updated cache, all in 1 kernel.
960
-
961
- If you pass in k / v, you must make sure that the cache is large enough to hold the new values.
962
- For example, the KV cache could be pre-allocated with the max sequence length, and you can use
963
- cache_seqlens to keep track of the current sequence lengths of each sequence in the batch.
964
-
965
- Also apply rotary embedding if rotary_cos and rotary_sin are passed in. The key @k will be
966
- rotated by rotary_cos and rotary_sin at indices cache_seqlens, cache_seqlens + 1, etc.
967
- If causal or local (i.e., window_size != (-1, -1)), the query @q will be rotated by rotary_cos
968
- and rotary_sin at indices cache_seqlens, cache_seqlens + 1, etc.
969
- If not causal and not local, the query @q will be rotated by rotary_cos and rotary_sin at
970
- indices cache_seqlens only (i.e. we consider all tokens in @q to be at position cache_seqlens).
971
-
972
- See tests/test_flash_attn.py::test_flash_attn_kvcache for examples of how to use this function.
973
-
974
- Supports multi-query and grouped-query attention (MQA/GQA) by passing in KV with fewer heads
975
- than Q. Note that the number of heads in Q must be divisible by the number of heads in KV.
976
- For example, if Q has 6 heads and K, V have 2 heads, head 0, 1, 2 of Q will attention to head
977
- 0 of K, V, and head 3, 4, 5 of Q will attention to head 1 of K, V.
978
-
979
- If causal=True, the causal mask is aligned to the bottom right corner of the attention matrix.
980
- For example, if seqlen_q = 2 and seqlen_k = 5, the causal mask (1 = keep, 0 = masked out) is:
981
- 1 1 1 1 0
982
- 1 1 1 1 1
983
- If seqlen_q = 5 and seqlen_k = 2, the causal mask is:
984
- 0 0
985
- 0 0
986
- 0 0
987
- 1 0
988
- 1 1
989
- If the row of the mask is all zero, the output will be zero.
990
-
991
- If window_size != (-1, -1), implements sliding window local attention. Query at position i
992
- will only attend to keys between
993
- [i + seqlen_k - seqlen_q - window_size[0], i + seqlen_k - seqlen_q + window_size[1]] inclusive.
994
-
995
- Note: Does not support backward pass.
996
-
997
- Arguments:
998
- q: (batch_size, seqlen, nheads, headdim)
999
- k_cache: (batch_size_cache, seqlen_cache, nheads_k, headdim) if there's no page_table,
1000
- or (num_blocks, page_block_size, nheads_k, headdim) if there's a page_table (i.e. paged KV cache)
1001
- page_block_size can be arbitrary (e.g, 1, 2, 3, 64, etc.).
1002
- v_cache: (batch_size_cache, seqlen_cache, nheads_k, headdim_v) if there's no page_table,
1003
- or (num_blocks, page_block_size, nheads_k, headdim_v) if there's a page_table (i.e. paged KV cache)
1004
- k [optional]: (batch_size, seqlen_new, nheads_k, headdim). If not None, we concatenate
1005
- k with k_cache, starting at the indices specified by cache_seqlens.
1006
- v [optional]: (batch_size, seqlen_new, nheads_k, headdim_v). Similar to k.
1007
- qv [optional]: (batch_size, seqlen, nheads, headdim_v)
1008
- rotary_cos [optional]: (seqlen_ro, rotary_dim / 2). If not None, we apply rotary embedding
1009
- to k and q. Only applicable if k and v are passed in. rotary_dim must be divisible by 16.
1010
- rotary_sin [optional]: (seqlen_ro, rotary_dim / 2). Similar to rotary_cos.
1011
- cache_seqlens: int, or (batch_size,), dtype torch.int32. The sequence lengths of the
1012
- KV cache.
1013
- cache_batch_idx: (batch_size,), dtype torch.int32. The indices used to index into the KV cache.
1014
- If None, we assume that the batch indices are [0, 1, 2, ..., batch_size - 1].
1015
- If the indices are not distinct, and k and v are provided, the values updated in the cache
1016
- might come from any of the duplicate indices.
1017
- cache_leftpad: (batch_size,), dtype torch.int32. The index that the KV cache starts. If None, assume 0.
1018
- page_table [optional]: (batch_size, max_num_blocks_per_seq), dtype torch.int32.
1019
- softmax_scale: float. The scaling of QK^T before applying softmax.
1020
- Default to 1 / sqrt(headdim).
1021
- causal: bool. Whether to apply causal attention mask (e.g., for auto-regressive modeling).
1022
- window_size: (left, right). If not (-1, -1), implements sliding window local attention.
1023
- softcap: float. Anything > 0 activates softcapping attention.
1024
- rotary_interleaved: bool. Only applicable if rotary_cos and rotary_sin are passed in.
1025
- If True, rotary embedding will combine dimensions 0 & 1, 2 & 3, etc. If False,
1026
- rotary embedding will combine dimensions 0 & rotary_dim / 2, 1 & rotary_dim / 2 + 1
1027
- (i.e. GPT-NeoX style).
1028
- num_splits: int. If > 1, split the key/value into this many chunks along the sequence.
1029
- If num_splits == 1, we don't split the key/value. If num_splits == 0, we use a heuristic
1030
- to automatically determine the number of splits.
1031
- Don't change this unless you know what you are doing.
1032
- return_softmax_lse: bool. Whether to return the logsumexp of the attention scores.
1033
-
1034
- Return:
1035
- out: (batch_size, seqlen, nheads, headdim).
1036
- softmax_lse [optional, if return_softmax_lse=True]: (batch_size, nheads, seqlen). The
1037
- logsumexp of each row of the matrix QK^T * scaling (e.g., log of the softmax
1038
- normalization factor).
1039
- """
1040
- assert k_cache.stride(-1) == 1, "k_cache must have contiguous last dimension"
1041
- assert v_cache.stride(-1) == 1, "v_cache must have contiguous last dimension"
1042
- if softmax_scale is None:
1043
- softmax_scale = (q.shape[-1] + (qv.shape[-1] if qv is not None else 0)) ** (-0.5)
1044
- if cache_seqlens is not None and isinstance(cache_seqlens, int):
1045
- cache_seqlens = torch.full(
1046
- (q.shape[0],), cache_seqlens, dtype=torch.int32, device=k_cache.device
1047
- )
1048
- cache_seqlens = maybe_contiguous(cache_seqlens)
1049
- out, softmax_lse, *rest = _flash_attn_forward(
1050
- q,
1051
- k_cache,
1052
- v_cache,
1053
- k,
1054
- v,
1055
- qv,
1056
- None, # out
1057
- cu_seqlens_q,
1058
- None, # cu_seqlens_k
1059
- cu_seqlens_k_new,
1060
- None, # seqused_q
1061
- cache_seqlens,
1062
- max_seqlen_q,
1063
- None, # max_seqlen_k
1064
- page_table,
1065
- cache_batch_idx,
1066
- cache_leftpad,
1067
- rotary_cos,
1068
- rotary_sin,
1069
- rotary_seqlens,
1070
- q_descale, k_descale, v_descale,
1071
- softmax_scale,
1072
- causal=causal,
1073
- window_size_left=window_size[0],
1074
- window_size_right=window_size[1],
1075
- attention_chunk=attention_chunk,
1076
- softcap=softcap,
1077
- rotary_interleaved=rotary_interleaved,
1078
- scheduler_metadata=scheduler_metadata,
1079
- num_splits=num_splits,
1080
- pack_gqa=pack_gqa,
1081
- sm_margin=sm_margin,
1082
- )
1083
- # return (out, softmax_lse) if return_softmax_lse else out
1084
- return (out, softmax_lse, *rest) if return_softmax_lse else out
1085
-
1086
-
1087
- def get_scheduler_metadata(
1088
- batch_size, max_seqlen_q, max_seqlen_k, num_heads_q, num_heads_kv, headdim,
1089
- cache_seqlens: torch.Tensor,
1090
- qkv_dtype=torch.bfloat16,
1091
- headdim_v=None,
1092
- cu_seqlens_q: Optional[torch.Tensor] = None,
1093
- cu_seqlens_k_new: Optional[torch.Tensor] = None,
1094
- cache_leftpad: Optional[torch.Tensor] = None,
1095
- page_size: Optional[int] = None,
1096
- max_seqlen_k_new=0,
1097
- causal=False,
1098
- window_size=(-1, -1), # -1 means infinite context window
1099
- attention_chunk=0,
1100
- has_softcap=False,
1101
- num_splits=0, # Can be tuned for speed
1102
- pack_gqa=None, # Can be tuned for speed
1103
- sm_margin=0, # Can be tuned if some SMs are used for communication
1104
- ):
1105
- cache_seqlens = maybe_contiguous(cache_seqlens)
1106
- if headdim_v is None:
1107
- headdim_v = headdim
1108
- scheduler_metadata = flash_attn_3_cuda.get_scheduler_metadata(
1109
- batch_size, max_seqlen_q, max_seqlen_k, num_heads_q, num_heads_kv, headdim, headdim_v,
1110
- qkv_dtype,
1111
- cache_seqlens,
1112
- cu_seqlens_q,
1113
- None, # cu_seqlens_k
1114
- cu_seqlens_k_new,
1115
- None, # seqused_q
1116
- cache_leftpad,
1117
- page_size,
1118
- max_seqlen_k_new,
1119
- causal,
1120
- window_size[0], window_size[1],
1121
- attention_chunk,
1122
- has_softcap,
1123
- num_splits,
1124
- pack_gqa,
1125
- sm_margin,
1126
- )
1127
- return scheduler_metadata
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu128-x86_64-linux/metadata.json DELETED
@@ -1,25 +0,0 @@
1
- {
2
- "name": "flash-attn3",
3
- "id": "_flash_attn3_cuda_477ab85",
4
- "version": 1,
5
- "license": "BSD-3-Clause",
6
- "python-depends": [],
7
- "backend": {
8
- "type": "cuda",
9
- "archs": [
10
- "8.0",
11
- "9.0a"
12
- ]
13
- },
14
- "digest": {
15
- "algorithm": "sha256",
16
- "files": {
17
- "__init__.py": "KXVmQJM+KhWc2UqWJesupCaP+mqAvTx7yGueFATglHE=",
18
- "_flash_attn3_cuda_477ab85.abi3.so": "3DLxY8Xh+2Ni4aLsFlo9K21iKSnEpGP6I7xuVh5+RoA=",
19
- "_ops.py": "hccJrV95SPARE3ynEBHpwny/YFXZqpZKdGdbTTnKTsc=",
20
- "flash_attn3/__init__.py": "DFYPlrhXwYjEqCl/8n0SmWGZV8NFml5DPhMjKfv98GY=",
21
- "flash_attn_config.py": "uxo+eyMcDit//8YDaS4Rtk5BTMKzTSK9tOKTjdWVe/w=",
22
- "flash_attn_interface.py": "y0bD2JYGMFgX6iDre3zJN2sBPuq/12N6Mo55Z+HQCNQ="
23
- }
24
- }
25
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu128-x86_64-linux/metadata.json.sigstore DELETED
@@ -1 +0,0 @@
1
- {"mediaType":"application/vnd.dev.sigstore.bundle.v0.3+json","verificationMaterial":{"certificate":{"rawBytes":"MIIHczCCBvmgAwIBAgIUDeWZNo2mRx6vtQAXx/Jikqf5DDYwCgYIKoZIzj0EAwMwNzEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MR4wHAYDVQQDExVzaWdzdG9yZS1pbnRlcm1lZGlhdGUwHhcNMjYwNjE4MDY1NTAwWhcNMjYwNjE4MDcwNTAwWjAAMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAEcAfjiZigLxvzHRjpXS+0iv7g8I2MseiPK2dHGroSp5aRE4gbd8YScSCL2hhyZxvU/K8C1tRNRJBjnih5nsOVhKOCBhgwggYUMA4GA1UdDwEB/wQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcDAzAdBgNVHQ4EFgQU2L/xTq3sG39X3TK03NOwEcwx3akwHwYDVR0jBBgwFoAU39Ppz1YkEZb5qNjpKFWixi4YZD8wdQYDVR0RAQH/BGswaYZnaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL3NpZ24tb2xkLWJ1aWxkcy55YW1sQHJlZnMvaGVhZHMvbWFpbjA5BgorBgEEAYO/MAEBBCtodHRwczovL3Rva2VuLmFjdGlvbnMuZ2l0aHVidXNlcmNvbnRlbnQuY29tMB8GCisGAQQBg78wAQIEEXdvcmtmbG93X2Rpc3BhdGNoMDYGCisGAQQBg78wAQMEKDhhNmJlN2JjNzc1NjVhZDY1YzhhZTJkZWI0NTY0ODJmZjZhYTUwZWQwHQYKKwYBBAGDvzABBAQPU2lnbiBvbGQgYnVpbGRzMCsGCisGAQQBg78wAQUEHWh1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MB0GCisGAQQBg78wAQYED3JlZnMvaGVhZHMvbWFpbjA7BgorBgEEAYO/MAEIBC0MK2h0dHBzOi8vdG9rZW4uYWN0aW9ucy5naXRodWJ1c2VyY29udGVudC5jb20wdwYKKwYBBAGDvzABCQRpDGdodHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHkvLmdpdGh1Yi93b3JrZmxvd3Mvc2lnbi1vbGQtYnVpbGRzLnlhbWxAcmVmcy9oZWFkcy9tYWluMDgGCisGAQQBg78wAQoEKgwoOGE2YmU3YmM3NzU2NWFkNjVjOGFlMmRlYjQ1NjQ4MmZmNmFhNTBlZDAbBgorBgEEAYO/MAELBA0MC3NlbGYtaG9zdGVkMEAGCisGAQQBg78wAQwEMgwwaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MDgGCisGAQQBg78wAQ0EKgwoOGE2YmU3YmM3NzU2NWFkNjVjOGFlMmRlYjQ1NjQ4MmZmNmFhNTBlZDAfBgorBgEEAYO/MAEOBBEMD3JlZnMvaGVhZHMvbWFpbjAaBgorBgEEAYO/MAEPBAwMCjEwNzE0NzU1MjkwLgYKKwYBBAGDvzABEAQgDB5odHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2UwGAYKKwYBBAGDvzABEQQKDAgyNTcyMDc0MzB3BgorBgEEAYO/MAESBGkMZ2h0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS8uZ2l0aHViL3dvcmtmbG93cy9zaWduLW9sZC1idWlsZHMueWFtbEByZWZzL2hlYWRzL21haW4wOAYKKwYBBAGDvzABEwQqDCg4YTZiZTdiYzc3NTY1YWQ2NWM4YWUyZGViNDU2NDgyZmY2YWE1MGVkMCEGCisGAQQBg78wARQEEwwRd29ya2Zsb3dfZGlzcGF0Y2gwZAYKKwYBBAGDvzABFQRWDFRodHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHkvYWN0aW9ucy9ydW5zLzI3NzQxNDY0ODQyL2F0dGVtcHRzLzEwFgYKKwYBBAGDvzABFgQIDAZwdWJsaWMwRgYKKwYBBAGDvzABGAQ4DDZyZXBvOmh1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5OnJlZjpyZWZzL2hlYWRzL21haW4wgYoGCisGAQQB1nkCBAIEfAR6AHgAdgDdPTBqxscRMmMZHhyZZzcCokpeuN48rf+HinKALynujgAAAZ7Zgvr/AAAEAwBHMEUCIDCysZoUBPjr8yrMndkRrH2yDxtb/WwkeonqZCuRvbdlAiEA7cGm/XJwc1LJtaM/Zep3cMa/2HTfg5wBJxCvrJWQphIwCgYIKoZIzj0EAwMDaAAwZQIwTZjYYD631m2LhXheH6yxXf1zZcGalZ55tBF5s9IGP1z0S9DWlIYDbedVcQcOgmZUAjEA5olo+1U3qINl2utRh/jYmmgNLuqWLYXinDelUQGbwQOplBv/h4TxXV3yf5gBwNB4"},"tlogEntries":[{"logIndex":"1857834973","logId":{"keyId":"wNI9atQGlz+VWfO6LRygH4QUfY/8W4RFwiT5i5WRgB0="},"kindVersion":{"kind":"hashedrekord","version":"0.0.1"},"integratedTime":"1781765700","inclusionPromise":{"signedEntryTimestamp":"MEYCIQCENrXlhp/ctD6wAP017I3C8WxfqxolWuW8eBzhvBxnXwIhAOcy6VbnXJT6bASHQw807FNmAcEo3UMYGOxRtcV2Ymqo"},"inclusionProof":{"logIndex":"1735930711","rootHash":"AeHYpsUox9yf3aBLMT5/w1OD/krCEgBNXtnvDdQ/AkY=","treeSize":"1735930715","hashes":["rnwEiVpQlWs/dE3ZIzpmfOxZpmlpQvnHYmr2ky4vPgE=","9pnrOR6oxCpBtO1XCi9QRw1kjhau9d0mgEYO2br17Wo=","R/7Wau4ucA+1bY0rFjGMPN6DlNhfnDUGzfl/k1kk2Gg=","VZz4wXXyfQk4DgDyUahbsLH3FjbAb52+MRMMfuE32W4=","lC3TbeUspaMecWNudNe6ZIxg4sgrm792r1DG33W/k4M=","gP/ZYVz9iLEfjKC7esTiDzwHeStGfT8SzNgpwEYK0Ic=","4lyHa8i8t5GIm7VEJ+TbO30Wa4dBA9X6oIrqgy/9z/g=","M9Vb0zzLaw2h2vZMtZalUKLUlQkj3euyLByPs5Zb8A8=","uMjlW8ZYxFHFISKF1gqKEB4T/pfjDh7vzUC1E7W0/xU=","Ore1A2Ceavap3m5U3ZAHLkUinE8F4IgttCeJDxxaK98=","95goX8Iw4H4P6V4eHvRqTdVdViLbhEhsF6i7ICnz/+A=","sWDh7SjDHbJf3HWKGRxiURh6iIYrOzn4Zs37yIij6OA=","x7kXd4VRJvVmTHoSla2KyPQdKvGKDX26/7ST9OONR78=","U8wMiVDzFlmyqT7Nw1RSZYU9+fftsSkRhjpbyXnXUk8=","mqM7J+i75IpuD09QejiUBjeH85AZa+fSm4RXXFmj3lI=","72FC5FYLhxB4a4iiC956o0B/fT54ip4R41vsw2QBKtU=","lYGQ9ibwC8+smMkPQ6TchJm3H9Nc/aTYLfdRacGFChw=","daxmZaajRpZV+JxHiOYZhJBiSKN5ucqjh2WnGbHhirw=","DOCeoSMovIvLExkhIvisow9AuNXgeWs4ECkyR6EcqYU="],"checkpoint":{"envelope":"rekor.sigstore.dev - 1193050959916656506\n1735930715\nAeHYpsUox9yf3aBLMT5/w1OD/krCEgBNXtnvDdQ/AkY=\n\n— rekor.sigstore.dev wNI9ajBEAiAEuZZGsmgewSHMnp4w8DQPS/mECQm+CveKExGyALolcgIgQpnKH8nq6GScB3k0mw5hCXjZWeh/ey9aqa8f/lf3SSo=\n"}},"canonicalizedBody":"eyJhcGlWZXJzaW9uIjoiMC4wLjEiLCJraW5kIjoiaGFzaGVkcmVrb3JkIiwic3BlYyI6eyJkYXRhIjp7Imhhc2giOnsiYWxnb3JpdGhtIjoic2hhMjU2IiwidmFsdWUiOiJkOWI4MGUyYjhmYzJhMzZlMzUxOWFmZGNlYTc4N2RkMTI5MWY3Y2I2OWYwZDQxNWU2ZWY0Yjk0ZTg0M2U3ZmI5In19LCJzaWduYXR1cmUiOnsiY29udGVudCI6Ik1FWUNJUURBcGJ4eXArcWVIb0dBNkRBd0kxUEI0blpBTzErNzA5bVhCUUYwbkxIZ2tRSWhBUGZYemp2cXpLZGVsdjhDOC8rYm8wWWVNZE9GblQ1aVZVVWlZYWsyTXBzeCIsInB1YmxpY0tleSI6eyJjb250ZW50IjoiTFMwdExTMUNSVWRKVGlCRFJWSlVTVVpKUTBGVVJTMHRMUzB0Q2sxSlNVaGpla05EUW5adFowRjNTVUpCWjBsVlJHVlhXazV2TW0xU2VEWjJkRkZCV0hndlNtbHJjV1kxUkVSWmQwTm5XVWxMYjFwSmVtb3dSVUYzVFhjS1RucEZWazFDVFVkQk1WVkZRMmhOVFdNeWJHNWpNMUoyWTIxVmRWcEhWakpOVWpSM1NFRlpSRlpSVVVSRmVGWjZZVmRrZW1SSE9YbGFVekZ3WW01U2JBcGpiVEZzV2tkc2FHUkhWWGRJYUdOT1RXcFpkMDVxUlRSTlJGa3hUbFJCZDFkb1kwNU5hbGwzVG1wRk5FMUVZM2RPVkVGM1YycEJRVTFHYTNkRmQxbElDa3R2V2tsNmFqQkRRVkZaU1V0dldrbDZhakJFUVZGalJGRm5RVVZqUVdacWFWcHBaMHg0ZG5wSVVtcHdXRk1yTUdsMk4yYzRTVEpOYzJWcFVFc3laRWdLUjNKdlUzQTFZVkpGTkdkaVpEaFpVMk5UUTB3eWFHaDVXbmgyVlM5TE9FTXhkRkpPVWtwQ2FtNXBhRFZ1YzA5V2FFdFBRMEpvWjNkbloxbFZUVUUwUndwQk1WVmtSSGRGUWk5M1VVVkJkMGxJWjBSQlZFSm5UbFpJVTFWRlJFUkJTMEpuWjNKQ1owVkdRbEZqUkVGNlFXUkNaMDVXU0ZFMFJVWm5VVlV5VEM5NENsUnhNM05ITXpsWU0xUkxNRE5PVDNkRlkzZDRNMkZyZDBoM1dVUldVakJxUWtKbmQwWnZRVlV6T1ZCd2VqRlphMFZhWWpWeFRtcHdTMFpYYVhocE5Ga0tXa1E0ZDJSUldVUldVakJTUVZGSUwwSkhjM2RoV1ZwdVlVaFNNR05JVFRaTWVUbHVZVmhTYjJSWFNYVlpNamwwVERKb01Wb3laSEJpYldSdFdWZE9iQXBNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk1lVFZ1WVZoU2IyUlhTWFprTWpsNVlUSmFjMkl6WkhwTU0wNXdXakkwZEdJeWVHdE1WMG94Q21GWGVHdGplVFUxV1ZjeGMxRklTbXhhYmsxMllVZFdhRnBJVFhaaVYwWndZbXBCTlVKbmIzSkNaMFZGUVZsUEwwMUJSVUpDUTNSdlpFaFNkMk42YjNZS1RETlNkbUV5Vm5WTWJVWnFaRWRzZG1KdVRYVmFNbXd3WVVoV2FXUllUbXhqYlU1MlltNVNiR0p1VVhWWk1qbDBUVUk0UjBOcGMwZEJVVkZDWnpjNGR3cEJVVWxGUlZoa2RtTnRkRzFpUnpreldESlNjR016UW1oa1IwNXZUVVJaUjBOcGMwZEJVVkZDWnpjNGQwRlJUVVZMUkdob1RtMUtiRTR5U21wT2VtTXhDazVxVm1oYVJGa3hXWHBvYUZwVVNtdGFWMGt3VGxSWk1FOUVTbTFhYWxwb1dWUlZkMXBYVVhkSVVWbExTM2RaUWtKQlIwUjJla0ZDUWtGUlVGVXliRzRLWW1sQ2RtSkhVV2RaYmxad1lrZFNlazFEYzBkRGFYTkhRVkZSUW1jM09IZEJVVlZGU0Zkb01Wb3laSEJpYldSdFdWZE9iRXd5ZEd4amJUVnNZa2hOZEFwWk1qbDBZbGhXZFdGWVVqVk5RakJIUTJselIwRlJVVUpuTnpoM1FWRlpSVVF6U214YWJrMTJZVWRXYUZwSVRYWmlWMFp3WW1wQk4wSm5iM0pDWjBWRkNrRlpUeTlOUVVWSlFrTXdUVXN5YURCa1NFSjZUMms0ZG1SSE9YSmFWelIxV1ZkT01HRlhPWFZqZVRWdVlWaFNiMlJYU2pGak1sWjVXVEk1ZFdSSFZuVUtaRU0xYW1JeU1IZGtkMWxMUzNkWlFrSkJSMFIyZWtGQ1ExRlNjRVJIWkc5a1NGSjNZM3B2ZGt3eVpIQmtSMmd4V1drMWFtSXlNSFpoU0ZadVdqSnNkUXBhTWxwb1dUSlZkbUV5Vm5saWJWWnpZM2t4YW1JeU1YUmtWelZ3WkVocmRreHRaSEJrUjJneFdXazVNMkl6U25KYWJYaDJaRE5OZG1NeWJHNWlhVEYyQ21KSFVYUlpibFp3WWtkU2VreHViR2hpVjNoQlkyMVdiV041T1c5YVYwWnJZM2s1ZEZsWGJIVk5SR2RIUTJselIwRlJVVUpuTnpoM1FWRnZSVXRuZDI4S1QwZEZNbGx0VlROWmJVMHpUbnBWTWs1WFJtdE9hbFpxVDBkR2JFMXRVbXhaYWxFeFRtcFJORTF0V20xT2JVWm9UbFJDYkZwRVFXSkNaMjl5UW1kRlJRcEJXVTh2VFVGRlRFSkJNRTFETTA1c1lrZFpkR0ZIT1hwa1IxWnJUVVZCUjBOcGMwZEJVVkZDWnpjNGQwRlJkMFZOWjNkM1lVaFNNR05JVFRaTWVUbHVDbUZZVW05a1YwbDFXVEk1ZEV3eWFERmFNbVJ3WW0xa2JWbFhUbXhNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk5SR2RIUTJselIwRlJVVUlLWnpjNGQwRlJNRVZMWjNkdlQwZEZNbGx0VlROWmJVMHpUbnBWTWs1WFJtdE9hbFpxVDBkR2JFMXRVbXhaYWxFeFRtcFJORTF0V20xT2JVWm9UbFJDYkFwYVJFRm1RbWR2Y2tKblJVVkJXVTh2VFVGRlQwSkNSVTFFTTBwc1dtNU5kbUZIVm1oYVNFMTJZbGRHY0dKcVFXRkNaMjl5UW1kRlJVRlpUeTlOUVVWUUNrSkJkMDFEYWtWM1RucEZNRTU2VlRGTmFtdDNUR2RaUzB0M1dVSkNRVWRFZG5wQlFrVkJVV2RFUWpWdlpFaFNkMk42YjNaTU1tUndaRWRvTVZscE5Xb0tZakl3ZG1GSVZtNWFNbXgxV2pKYWFGa3lWWGRIUVZsTFMzZFpRa0pCUjBSMmVrRkNSVkZSUzBSQlozbE9WR041VFVSak1FMTZRak5DWjI5eVFtZEZSUXBCV1U4dlRVRkZVMEpIYTAxYU1tZ3daRWhDZWs5cE9IWmFNbXd3WVVoV2FVeHRUblppVXpsdlpGZGtibUZYTlc1YWJVWnFXbE01Y2xwWVNuVmFWM2g2Q2t4WFRuWmlWekV4WW0xc01HVlRPSFZhTW13d1lVaFdhVXd6WkhaamJYUnRZa2M1TTJONU9YcGhWMlIxVEZjNWMxcERNV2xrVjJ4eldraE5kV1ZYUm5RS1lrVkNlVnBYV25wTU1taHNXVmRTZWt3eU1XaGhWelIzVDBGWlMwdDNXVUpDUVVkRWRucEJRa1YzVVhGRVEyYzBXVlJhYVZwVVpHbFplbU16VGxSWk1RcFpWMUV5VGxkTk5GbFhWWGxhUjFacFRrUlZNazVFWjNsYWJWa3lXVmRGTVUxSFZtdE5RMFZIUTJselIwRlJVVUpuTnpoM1FWSlJSVVYzZDFKa01qbDVDbUV5V25OaU0yUm1Xa2RzZW1OSFJqQlpNbWQzV2tGWlMwdDNXVUpDUVVkRWRucEJRa1pSVWxkRVJsSnZaRWhTZDJONmIzWk1NbVJ3WkVkb01WbHBOV29LWWpJd2RtRklWbTVhTW14MVdqSmFhRmt5VlhaaE1sWjVZbTFXYzJONU1XcGlNakYwWkZjMWNHUklhM1paVjA0d1lWYzVkV041T1hsa1Z6VjZUSHBKTXdwT2VsRjRUa1JaTUU5RVVYbE1Na1l3WkVkV2RHTklVbnBNZWtWM1JtZFpTMHQzV1VKQ1FVZEVkbnBCUWtablVVbEVRVnAzWkZkS2MyRlhUWGRTWjFsTENrdDNXVUpDUVVkRWRucEJRa2RCVVRSRVJGcDVXbGhDZGs5dGFERmFNbVJ3WW0xa2JWbFhUbXhNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVUtUMjVLYkZwcWNIbGFWMXA2VERKb2JGbFhVbnBNTWpGb1lWYzBkMmRaYjBkRGFYTkhRVkZSUWpGdWEwTkNRVWxGWmtGU05rRklaMEZrWjBSa1VGUkNjUXA0YzJOU1RXMU5Xa2hvZVZwYWVtTkRiMnR3WlhWT05EaHlaaXRJYVc1TFFVeDViblZxWjBGQlFWbzNXbWQyY2k5QlFVRkZRWGRDU0UxRlZVTkpSRU41Q25OYWIxVkNVR3B5T0hseVRXNWthMUp5U0RKNVJIaDBZaTlYZDJ0bGIyNXhXa04xVW5aaVpHeEJhVVZCTjJOSGJTOVlTbmRqTVV4S2RHRk5MMXBsY0RNS1kwMWhMekpJVkdabk5YZENTbmhEZG5KS1YxRndhRWwzUTJkWlNVdHZXa2w2YWpCRlFYZE5SR0ZCUVhkYVVVbDNWRnBxV1ZsRU5qTXhiVEpNYUZob1pRcElObmw0V0dZeGVscGpSMkZzV2pVMWRFSkdOWE01U1VkUU1Yb3dVemxFVjJ4SldVUmlaV1JXWTFGalQyZHRXbFZCYWtWQk5XOXNieXN4VlROeFNVNXNDakoxZEZKb0wycFpiVzFuVGt4MWNWZE1XVmhwYmtSbGJGVlJSMkozVVU5d2JFSjJMMmcwVkhoWVZqTjVaalZuUW5kT1FqUUtMUzB0TFMxRlRrUWdRMFZTVkVsR1NVTkJWRVV0TFMwdExRbz0ifX19fQ=="}],"timestampVerificationData":{"rfc3161Timestamps":[{"signedTimestamp":"MIICyzADAgEAMIICwgYJKoZIhvcNAQcCoIICszCCAq8CAQMxDTALBglghkgBZQMEAgEwgbgGCyqGSIb3DQEJEAEEoIGoBIGlMIGiAgEBBgkrBgEEAYO/MAIwMTANBglghkgBZQMEAgEFAAQgA88BUi/BpeOzOFlAGk0vdxxvc0jT7Oipfk8SVtNe4YICFQC+324dc0UezY70e9CfXt5OH7Y84hgPMjAyNjA2MTgwNjU1MDBaMAMCAQGgMqQwMC4xFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEVMBMGA1UEAxMMc2lnc3RvcmUtdHNhoAAxggHcMIIB2AIBATBRMDkxFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEgMB4GA1UEAxMXc2lnc3RvcmUtdHNhLXNlbGZzaWduZWQCFDoTVC8MkGHuvMFDL8uKjosqI4sMMAsGCWCGSAFlAwQCAaCB/DAaBgkqhkiG9w0BCQMxDQYLKoZIhvcNAQkQAQQwHAYJKoZIhvcNAQkFMQ8XDTI2MDYxODA2NTUwMFowLwYJKoZIhvcNAQkEMSIEIBRIxBixSu8CFkS56i4YkrF12st0ePLAZ26gsPKnbcAkMIGOBgsqhkiG9w0BCRACLzF/MH0wezB5BCCF+Se8B6tiysO0Q1bBDvyBssaIP9p6uebYcNnROs0FtzBVMD2kOzA5MRUwEwYDVQQKEwxzaWdzdG9yZS5kZXYxIDAeBgNVBAMTF3NpZ3N0b3JlLXRzYS1zZWxmc2lnbmVkAhQ6E1QvDJBh7rzBQy/Lio6LKiOLDDAKBggqhkjOPQQDAgRoMGYCMQCg8v0tyTIRKiRIo0MKD/0l+4yy7ClS7K3Up7KYKEo6a0PBUZB0QTTkJjaOeTtvEPgCMQC7tQW4JWB31uszMh33RoesJbL2zJvGfotzGK0bOKKb14nmEE4rGQJIhTuH906oXro="}]}},"messageSignature":{"messageDigest":{"algorithm":"SHA2_256","digest":"2bgOK4/Co241Ga/c6nh90SkffLafDUFebvS5ToQ+f7k="},"signature":"MEYCIQDApbxyp+qeHoGA6DAwI1PB4nZAO1+709mXBQF0nLHgkQIhAPfXzjvqzKdelv8C8/+bo0YeMdOFnT5iVUUiYak2Mpsx"}}
 
 
build/torch-stable-abi210-cu130-x86_64-linux/__init__.py DELETED
@@ -1,17 +0,0 @@
1
- from .flash_attn_interface import (
2
- flash_attn_combine,
3
- flash_attn_func,
4
- flash_attn_qkvpacked_func,
5
- flash_attn_varlen_func,
6
- flash_attn_with_kvcache,
7
- get_scheduler_metadata,
8
- )
9
-
10
- __all__ = [
11
- "flash_attn_combine",
12
- "flash_attn_func",
13
- "flash_attn_qkvpacked_func",
14
- "flash_attn_varlen_func",
15
- "flash_attn_with_kvcache",
16
- "get_scheduler_metadata",
17
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu130-x86_64-linux/_flash_attn3_cuda_477ab85.abi3.so DELETED
@@ -1,3 +0,0 @@
1
- version https://git-lfs.github.com/spec/v1
2
- oid sha256:a79c38ca3cbf5f6630e96eb1ab6022b21be7d93e2aa2461d0da7341a0f243f30
3
- size 821408616
 
 
 
 
build/torch-stable-abi210-cu130-x86_64-linux/_ops.py DELETED
@@ -1,9 +0,0 @@
1
- import torch
2
- from . import _flash_attn3_cuda_477ab85
3
- ops = torch.ops._flash_attn3_cuda_477ab85
4
-
5
- def add_op_namespace_prefix(op_name: str):
6
- """
7
- Prefix op by namespace.
8
- """
9
- return f"_flash_attn3_cuda_477ab85::{op_name}"
 
 
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu130-x86_64-linux/flash_attn3/__init__.py DELETED
@@ -1,26 +0,0 @@
1
- import ctypes
2
- import importlib.util
3
- import sys
4
- from pathlib import Path
5
- from types import ModuleType
6
-
7
-
8
- def _import_from_path(file_path: Path) -> ModuleType:
9
- # We cannot use the module name as-is, after adding it to `sys.modules`,
10
- # it would also be used for other imports. So, we make a module name that
11
- # depends on the path for it to be unique using the hex-encoded hash of
12
- # the path.
13
- path_hash = "{:x}".format(ctypes.c_size_t(hash(file_path.absolute())).value)
14
- module_name = path_hash
15
- spec = importlib.util.spec_from_file_location(module_name, file_path)
16
- if spec is None:
17
- raise ImportError(f"Cannot load spec for {module_name} from {file_path}")
18
- module = importlib.util.module_from_spec(spec)
19
- if module is None:
20
- raise ImportError(f"Cannot load module {module_name} from spec")
21
- sys.modules[module_name] = module
22
- spec.loader.exec_module(module) # type: ignore
23
- return module
24
-
25
-
26
- globals().update(vars(_import_from_path(Path(__file__).parent.parent / "__init__.py")))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu130-x86_64-linux/flash_attn_config.py DELETED
@@ -1,7 +0,0 @@
1
- # Auto-generated by flash attention 3 setup.py
2
- CONFIG = {'build_flags': {'FLASHATTENTION_DISABLE_BACKWARD': False, 'FLASHATTENTION_DISABLE_SPLIT': False, 'FLASHATTENTION_DISABLE_PAGEDKV': False, 'FLASHATTENTION_DISABLE_APPENDKV': False, 'FLASHATTENTION_DISABLE_LOCAL': False, 'FLASHATTENTION_DISABLE_SOFTCAP': False, 'FLASHATTENTION_DISABLE_PACKGQA': False, 'FLASHATTENTION_DISABLE_FP16': False, 'FLASHATTENTION_DISABLE_FP8': False, 'FLASHATTENTION_DISABLE_VARLEN': False, 'FLASHATTENTION_DISABLE_CLUSTER': False, 'FLASHATTENTION_DISABLE_HDIM64': False, 'FLASHATTENTION_DISABLE_HDIM96': False, 'FLASHATTENTION_DISABLE_HDIM128': False, 'FLASHATTENTION_DISABLE_HDIM192': False, 'FLASHATTENTION_DISABLE_HDIM256': False, 'FLASHATTENTION_DISABLE_SM8x': False, 'FLASHATTENTION_ENABLE_VCOLMAJOR': False, 'FLASH_ATTENTION_DISABLE_HDIMDIFF64': False, 'FLASH_ATTENTION_DISABLE_HDIMDIFF192': False}}
3
-
4
- def show():
5
- from pprint import pprint
6
- pprint(CONFIG)
7
-
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu130-x86_64-linux/flash_attn_interface.py DELETED
@@ -1,1127 +0,0 @@
1
- # Copyright (c) 2023, Tri Dao.
2
-
3
- from typing import Optional, Union, List, Tuple
4
-
5
- import torch
6
- import torch.nn as nn
7
-
8
- from ._ops import ops as flash_attn_3_cuda
9
- from ._ops import add_op_namespace_prefix
10
-
11
- def maybe_contiguous(x):
12
- return x.contiguous() if x is not None and x.stride(-1) != 1 else x
13
-
14
-
15
- def round_multiple(x, m):
16
- return (x + m - 1) // m * m
17
-
18
-
19
- def round_up_headdim(head_size: int) -> int:
20
- from .flash_attn_config import CONFIG
21
-
22
- if not CONFIG["build_flags"]["FLASHATTENTION_DISABLE_HDIM64"]:
23
- if head_size <= 64:
24
- return 64
25
- if not CONFIG["build_flags"]["FLASHATTENTION_DISABLE_HDIM96"]:
26
- if head_size <= 96:
27
- return 96
28
- if not CONFIG["build_flags"]["FLASHATTENTION_DISABLE_HDIM128"]:
29
- if head_size <= 128:
30
- return 128
31
- if not CONFIG["build_flags"]["FLASHATTENTION_DISABLE_HDIM192"]:
32
- if head_size <= 192:
33
- return 192
34
- if not CONFIG["build_flags"]["FLASHATTENTION_DISABLE_HDIM256"]:
35
- if head_size <= 256:
36
- return 256
37
- return 256
38
-
39
-
40
- @torch.library.custom_op(add_op_namespace_prefix("_flash_attn_forward"), mutates_args=(), device_types="cuda")
41
- def _flash_attn_forward(
42
- q: torch.Tensor,
43
- k: torch.Tensor,
44
- v: torch.Tensor,
45
- k_new: Optional[torch.Tensor] = None,
46
- v_new: Optional[torch.Tensor] = None,
47
- qv: Optional[torch.Tensor] = None,
48
- out_: Optional[torch.Tensor] = None,
49
- cu_seqlens_q: Optional[torch.Tensor] = None,
50
- cu_seqlens_k: Optional[torch.Tensor] = None,
51
- cu_seqlens_k_new: Optional[torch.Tensor] = None,
52
- seqused_q: Optional[torch.Tensor] = None,
53
- seqused_k: Optional[torch.Tensor] = None,
54
- max_seqlen_q: Optional[int] = None,
55
- max_seqlen_k: Optional[int] = None,
56
- page_table: Optional[torch.Tensor] = None,
57
- kv_batch_idx: Optional[torch.Tensor] = None,
58
- leftpad_k: Optional[torch.Tensor] = None,
59
- rotary_cos: Optional[torch.Tensor] = None,
60
- rotary_sin: Optional[torch.Tensor] = None,
61
- seqlens_rotary: Optional[torch.Tensor] = None,
62
- q_descale: Optional[torch.Tensor] = None,
63
- k_descale: Optional[torch.Tensor] = None,
64
- v_descale: Optional[torch.Tensor] = None,
65
- softmax_scale: Optional[float] = None,
66
- causal: bool = False,
67
- window_size_left: int = -1,
68
- window_size_right: int = -1,
69
- attention_chunk: int = 0,
70
- softcap: float = 0.0,
71
- rotary_interleaved: bool = True,
72
- scheduler_metadata: Optional[torch.Tensor] = None,
73
- num_splits: int = 1,
74
- pack_gqa: Optional[bool] = None,
75
- sm_margin: int = 0,
76
- ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
77
- q, k, k_new, v_new = [maybe_contiguous(x) for x in (q, k, k_new, v_new)]
78
- v = v.contiguous() if v.stride(-1) != 1 and v.stride(-3) != 1 else v
79
- cu_seqlens_q, cu_seqlens_k, cu_seqlens_k_new = [
80
- maybe_contiguous(x) for x in (cu_seqlens_q, cu_seqlens_k, cu_seqlens_k_new)
81
- ]
82
- seqused_q, seqused_k = [maybe_contiguous(x) for x in (seqused_q, seqused_k)]
83
- page_table, kv_batch_idx, leftpad_k = [
84
- maybe_contiguous(x) for x in (page_table, kv_batch_idx, leftpad_k)
85
- ]
86
- rotary_cos, rotary_sin = [maybe_contiguous(x) for x in (rotary_cos, rotary_sin)]
87
- seqlens_rotary = maybe_contiguous(seqlens_rotary)
88
- out, softmax_lse, out_accum, softmax_lse_accum = flash_attn_3_cuda.fwd(
89
- q,
90
- k,
91
- v,
92
- k_new,
93
- v_new,
94
- qv,
95
- out_,
96
- cu_seqlens_q,
97
- cu_seqlens_k,
98
- cu_seqlens_k_new,
99
- seqused_q,
100
- seqused_k,
101
- max_seqlen_q,
102
- max_seqlen_k,
103
- page_table,
104
- kv_batch_idx,
105
- leftpad_k,
106
- rotary_cos,
107
- rotary_sin,
108
- seqlens_rotary,
109
- q_descale,
110
- k_descale,
111
- v_descale,
112
- softmax_scale,
113
- causal,
114
- window_size_left,
115
- window_size_right,
116
- attention_chunk,
117
- softcap,
118
- rotary_interleaved,
119
- scheduler_metadata,
120
- num_splits,
121
- pack_gqa,
122
- sm_margin,
123
- )
124
-
125
- if out_accum is None:
126
- out_accum = torch.tensor([], device=out.device)
127
-
128
- if softmax_lse_accum is None:
129
- softmax_lse_accum = torch.tensor([], device=out.device)
130
-
131
- return out, softmax_lse, out_accum, softmax_lse_accum
132
-
133
-
134
- @torch.library.register_fake(add_op_namespace_prefix("_flash_attn_forward"))
135
- def _flash_attn_forward_fake(
136
- q: torch.Tensor,
137
- k: torch.Tensor,
138
- v: torch.Tensor,
139
- k_new: Optional[torch.Tensor] = None,
140
- v_new: Optional[torch.Tensor] = None,
141
- qv: Optional[torch.Tensor] = None,
142
- out_: Optional[torch.Tensor] = None,
143
- cu_seqlens_q: Optional[torch.Tensor] = None,
144
- cu_seqlens_k: Optional[torch.Tensor] = None,
145
- cu_seqlens_k_new: Optional[torch.Tensor] = None,
146
- seqused_q: Optional[torch.Tensor] = None,
147
- seqused_k: Optional[torch.Tensor] = None,
148
- max_seqlen_q: Optional[int] = None,
149
- max_seqlen_k: Optional[int] = None,
150
- page_table: Optional[torch.Tensor] = None,
151
- kv_batch_idx: Optional[torch.Tensor] = None,
152
- leftpad_k: Optional[torch.Tensor] = None,
153
- rotary_cos: Optional[torch.Tensor] = None,
154
- rotary_sin: Optional[torch.Tensor] = None,
155
- seqlens_rotary: Optional[torch.Tensor] = None,
156
- q_descale: Optional[torch.Tensor] = None,
157
- k_descale: Optional[torch.Tensor] = None,
158
- v_descale: Optional[torch.Tensor] = None,
159
- softmax_scale: Optional[float] = None,
160
- causal: bool = False,
161
- window_size_left: int = -1,
162
- window_size_right: int = -1,
163
- attention_chunk: int = 0,
164
- softcap: float = 0.0,
165
- rotary_interleaved: bool = True,
166
- scheduler_metadata: Optional[torch.Tensor] = None,
167
- num_splits: int = 1,
168
- pack_gqa: Optional[bool] = None,
169
- sm_margin: int = 0,
170
- ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
171
- """
172
- Symbolic fake implementation of flash attention forward.
173
- Returns tensors with the correct shapes and dtypes without actual computation.
174
- """
175
-
176
- # Determine if we're in varlen mode
177
- is_varlen_q = cu_seqlens_q is not None
178
-
179
- # Get dimensions from query tensor
180
- if is_varlen_q:
181
- # varlen mode: q is (total_q, num_heads, head_size)
182
- total_q, num_heads, head_size = q.shape
183
- batch_size = cu_seqlens_q.shape[0] - 1
184
-
185
- if max_seqlen_q is None:
186
- raise ValueError("max_seqlen_q must be provided if cu_seqlens_q is provided")
187
- seqlen_q = max_seqlen_q
188
- else:
189
- # batch mode: q is (batch_size, seqlen_q, num_heads, head_size)
190
- batch_size, seqlen_q, num_heads, head_size = q.shape
191
- total_q = batch_size * q.shape[1]
192
- # Get value head dimension
193
- head_size_v = v.shape[-1]
194
-
195
- # Determine output dtype (FP8 inputs produce BF16 outputs)
196
- q_type = q.dtype
197
- if q_type == torch.float8_e4m3fn:
198
- out_dtype = torch.bfloat16
199
- else:
200
- out_dtype = q_type
201
-
202
- # Create output tensor
203
- if out_ is not None:
204
- # If out_ is provided, _flash_attn_forward becomes non-functional
205
- raise TypeError("Tracing (torch.compile/torch.export) with pre-allocated output tensor is not supported.")
206
-
207
- if is_varlen_q:
208
- out = torch.empty((total_q, num_heads, head_size_v), dtype=out_dtype, device=q.device)
209
- else:
210
- out = torch.empty((batch_size, seqlen_q, num_heads, head_size_v), dtype=out_dtype, device=q.device)
211
-
212
- # Create softmax_lse tensor
213
- if is_varlen_q:
214
- softmax_lse = torch.empty((num_heads, total_q), dtype=torch.float32, device=q.device)
215
- else:
216
- softmax_lse = torch.empty((batch_size, num_heads, seqlen_q), dtype=torch.float32, device=q.device)
217
-
218
- # TODO(guilhermeleobas): Implement "get_num_splits"
219
- # There's an heuristic to compute num_splits when "num_splits <= 0"
220
- # assert that num_splits is > 0 for now
221
- if num_splits <= 0:
222
- raise ValueError(f"tracing (torch.compile/torch.export) with num_splits <= 0 not supported. Got {num_splits=}")
223
-
224
- if num_splits > 1:
225
- if is_varlen_q:
226
- out_accum = torch.empty((num_splits, num_heads, total_q, head_size_v), dtype=torch.float32, device=q.device)
227
- softmax_lse_accum = torch.empty((num_splits, num_heads, total_q), dtype=torch.float32, device=q.device)
228
- else:
229
- out_accum = torch.empty((num_splits, batch_size, num_heads, seqlen_q, head_size_v), dtype=torch.float32, device=q.device)
230
- softmax_lse_accum = torch.empty((num_splits, batch_size, num_heads, seqlen_q), dtype=torch.float32, device=q.device)
231
- else:
232
- # Tensors are not set when num_splits < 1
233
- out_accum = torch.tensor([], device=out.device)
234
- softmax_lse_accum = torch.tensor([], device=out.device)
235
-
236
- return out, softmax_lse, out_accum, softmax_lse_accum
237
-
238
-
239
- @torch.library.custom_op(add_op_namespace_prefix("_flash_attn_backward"), mutates_args=("dq", "dk", "dv"), device_types="cuda")
240
- def _flash_attn_backward(
241
- dout: torch.Tensor,
242
- q: torch.Tensor,
243
- k: torch.Tensor,
244
- v: torch.Tensor,
245
- out: torch.Tensor,
246
- softmax_lse: torch.Tensor,
247
- cu_seqlens_q: Optional[torch.Tensor] = None,
248
- cu_seqlens_k: Optional[torch.Tensor] = None,
249
- sequed_q: Optional[torch.Tensor] = None,
250
- sequed_k: Optional[torch.Tensor] = None,
251
- max_seqlen_q: Optional[int] = None,
252
- max_seqlen_k: Optional[int] = None,
253
- dq: Optional[torch.Tensor] = None,
254
- dk: Optional[torch.Tensor] = None,
255
- dv: Optional[torch.Tensor] = None,
256
- softmax_scale: Optional[float] = None,
257
- is_causal: bool = False,
258
- window_size_left: int = -1,
259
- window_size_right: int = -1,
260
- softcap: float = 0.0,
261
- deterministic: bool = False,
262
- sm_margin: int = 0,
263
- ) -> torch.Tensor:
264
- # dq, dk, dv are allocated by us so they should already be contiguous
265
- dout, q, k, v, out = [maybe_contiguous(x) for x in (dout, q, k, v, out)]
266
- softmax_d, *rest = flash_attn_3_cuda.bwd(
267
- dout,
268
- q,
269
- k,
270
- v,
271
- out,
272
- softmax_lse,
273
- dq,
274
- dk,
275
- dv,
276
- cu_seqlens_q,
277
- cu_seqlens_k,
278
- sequed_q,
279
- sequed_k,
280
- max_seqlen_q,
281
- max_seqlen_k,
282
- softmax_scale,
283
- is_causal,
284
- window_size_left,
285
- window_size_right,
286
- softcap,
287
- deterministic,
288
- sm_margin,
289
- )
290
- return softmax_d
291
-
292
-
293
- @torch.library.register_fake(add_op_namespace_prefix("_flash_attn_backward"))
294
- def _flash_attn_backward_fake(
295
- dout: torch.Tensor,
296
- q: torch.Tensor,
297
- k: torch.Tensor,
298
- v: torch.Tensor,
299
- out: torch.Tensor,
300
- softmax_lse: torch.Tensor,
301
- cu_seqlens_q: Optional[torch.Tensor] = None,
302
- cu_seqlens_k: Optional[torch.Tensor] = None,
303
- sequed_q: Optional[torch.Tensor] = None,
304
- sequed_k: Optional[torch.Tensor] = None,
305
- max_seqlen_q: Optional[int] = None,
306
- max_seqlen_k: Optional[int] = None,
307
- dq: Optional[torch.Tensor] = None,
308
- dk: Optional[torch.Tensor] = None,
309
- dv: Optional[torch.Tensor] = None,
310
- softmax_scale: Optional[float] = None,
311
- is_causal: bool = False,
312
- window_size_left: int = -1,
313
- window_size_right: int = -1,
314
- softcap: float = 0.0,
315
- deterministic: bool = False,
316
- sm_margin: int = 0,
317
- ) -> torch.Tensor:
318
-
319
- is_varlen_q = cu_seqlens_q is not None
320
- is_varlen_k = cu_seqlens_q is not None
321
- is_varlen = is_varlen_q or is_varlen_k or sequed_q is not None or sequed_k is not None
322
-
323
- if not is_varlen_q:
324
- batch_size = q.size(0)
325
- seqlen_q = q.size(1)
326
- seqlen_k = k.size(1)
327
- total_q = batch_size * q.size(1)
328
- else:
329
- batch_size = cu_seqlens_q.size(0) - 1
330
- total_q = q.size(0)
331
- seqlen_q = max_seqlen_q
332
- seqlen_k = max_seqlen_k
333
-
334
- if window_size_left >= seqlen_k - 1:
335
- window_size_left = -1
336
-
337
- if window_size_right >= seqlen_q - 1:
338
- window_size_right = -1
339
-
340
- if is_causal:
341
- window_size_right = 0
342
-
343
- is_causal = window_size_left < 0 and window_size_right == 0
344
-
345
- head_size = q.size(-1)
346
- head_size_v = v.size(-1)
347
- head_size_rounded = round_up_headdim(max(head_size, head_size_v))
348
-
349
- # Hopper gpus uses cuda compute capabilities 9.0
350
- cap = torch.cuda.get_device_capability(q.device)
351
- arch = cap[0] * 10 + cap[1]
352
-
353
- is_local = (window_size_left >= 0 or window_size_right >= 0) and not is_causal
354
-
355
- if head_size_rounded <= 64:
356
- kBlockM_sm90 = 96 if (is_causal and softcap > 0.0) else 128
357
- elif head_size_rounded <= 96:
358
- kBlockM_sm90 = 64
359
- elif head_size_rounded <= 128:
360
- kBlockM_sm90 = 64 if (is_causal or is_local or softcap > 0.0) else 80
361
- else:
362
- kBlockM_sm90 = 64
363
-
364
- kBlockM_sm80 = 128 if head_size_rounded <= 64 else 64
365
- kBlockM_sm86 = 64 if head_size_rounded <= 192 else 32
366
-
367
- if arch >= 90:
368
- kBlockM = kBlockM_sm90
369
- elif arch == 86 or arch == 89:
370
- kBlockM = kBlockM_sm86
371
- else:
372
- kBlockM = kBlockM_sm80
373
-
374
- num_heads = q.shape[-2]
375
- seqlen_q_rounded = round_multiple(seqlen_q, kBlockM)
376
-
377
- total_q_padded_rounded = round_multiple(total_q + batch_size * kBlockM, kBlockM)
378
-
379
- dq = torch.empty_like(q) if dq is None else dq
380
- dk = torch.empty_like(k) if dk is None else dk
381
- dv = torch.empty_like(v) if dv is None else dv
382
-
383
- if not is_varlen:
384
- softmax_d = torch.empty((batch_size, num_heads, seqlen_q_rounded), dtype=torch.float32, device=q.device)
385
- else:
386
- softmax_d = torch.empty((num_heads, total_q_padded_rounded), dtype=torch.float32, device=q.device)
387
-
388
- return softmax_d
389
-
390
-
391
- def setup_context(ctx, inputs, output):
392
- q, k, v = inputs[:3]
393
- out, softmax_lse, _, _ = output
394
- ctx.save_for_backward(q, k, v, out, softmax_lse)
395
- ctx.softmax_scale = inputs[-11]
396
- ctx.causal = inputs[-10]
397
- ctx.window_size = [inputs[-9], inputs[-8]]
398
- ctx.attention_chunk = inputs[-7]
399
- ctx.softcap = inputs[-6]
400
- ctx.sm_margin = inputs[-1]
401
-
402
-
403
- def _backward(ctx, dout, *grads):
404
- q, k, v, out, softmax_lse = ctx.saved_tensors
405
- dq, dk, dv = torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
406
- _flash_attn_backward(
407
- dout,
408
- q,
409
- k,
410
- v,
411
- out,
412
- softmax_lse,
413
- None, None, # cu_seqlens_q, cu_seqlens_k,
414
- None, None, # sequed_q, sequed_k,
415
- None, None, # max_seqlen_q, max_seqlen_k,
416
- dq,
417
- dk,
418
- dv,
419
- ctx.softmax_scale,
420
- ctx.causal,
421
- ctx.window_size[0],
422
- ctx.window_size[1],
423
- ctx.softcap,
424
- False, # deterministic
425
- ctx.sm_margin,
426
- )
427
- return dq, dk, dv, *((None,) * 21)
428
-
429
-
430
- _flash_attn_forward.register_autograd(_backward, setup_context=setup_context)
431
-
432
-
433
-
434
- class FlashAttnQKVPackedFunc(torch.autograd.Function):
435
- @staticmethod
436
- def forward(
437
- ctx,
438
- qkv,
439
- softmax_scale,
440
- causal,
441
- q_descale=None, k_descale=None, v_descale=None,
442
- window_size=(-1, -1),
443
- attention_chunk=0,
444
- softcap=0.0,
445
- deterministic=False,
446
- num_heads_q=None,
447
- sm_margin=0,
448
- return_softmax=False,
449
- ):
450
- if softmax_scale is None:
451
- softmax_scale = qkv.shape[-1] ** (-0.5)
452
- if qkv.dim() == 5:
453
- assert qkv.shape[-3] == 3
454
- q, k, v = qkv.unbind(dim=-3)
455
- else:
456
- assert qkv.dim() == 4
457
- assert num_heads_q is not None
458
- num_heads_k = (qkv.shape[2] - num_heads_q) // 2
459
- assert num_heads_k * 2 + num_heads_q == qkv.shape[2]
460
- q, k, v = qkv.split([num_heads_q, num_heads_k, num_heads_k], dim=-2)
461
- out, softmax_lse, *rest = _flash_attn_forward(
462
- q,
463
- k,
464
- v,
465
- None, None, # k_new, v_new
466
- None, # qv
467
- None, # out
468
- None, None, None, # cu_seqlens_q/k/k_new
469
- None, None, # seqused_q/k
470
- None, None, # max_seqlen_q/k
471
- None, None, None, # page_table, kv_batch_idx, leftpad_k,
472
- None, None, None, # rotary_cos/sin, seqlens_rotary
473
- q_descale, k_descale, v_descale,
474
- softmax_scale,
475
- causal=causal,
476
- window_size_left=window_size[0],
477
- window_size_right=window_size[1],
478
- attention_chunk=attention_chunk,
479
- softcap=softcap,
480
- sm_margin=sm_margin,
481
- )
482
- # ctx.save_for_backward(q, k, v, out_padded, softmax_lse)
483
- ctx.save_for_backward(q, k, v, out, softmax_lse)
484
- ctx.softmax_scale = softmax_scale
485
- ctx.causal = causal
486
- ctx.window_size = window_size
487
- ctx.attention_chunk = attention_chunk
488
- ctx.softcap = softcap
489
- ctx.deterministic = deterministic
490
- ctx.ndim = qkv.dim()
491
- ctx.sm_margin = sm_margin
492
- return (out, softmax_lse) if return_softmax else out
493
-
494
- @staticmethod
495
- def backward(ctx, dout, *args):
496
- q, k, v, out, softmax_lse = ctx.saved_tensors
497
- assert ctx.attention_chunk == 0, "FA3 backward does not support attention_chunk"
498
- if ctx.ndim == 5:
499
- qkv_shape = q.shape[:-2] + (3, *q.shape[-2:])
500
- dqkv = torch.empty(qkv_shape, dtype=q.dtype, device=q.device)
501
- dq, dk, dv = dqkv.unbind(dim=-3)
502
- else:
503
- num_heads_q = q.shape[2]
504
- num_heads_k = k.shape[2]
505
- qkv_shape = q.shape[:-2] + (num_heads_q + num_heads_k * 2, *q.shape[-1:])
506
- dqkv = torch.empty(qkv_shape, dtype=q.dtype, device=q.device)
507
- dq, dk, dv = dqkv.split([num_heads_q, num_heads_k, num_heads_k], dim=-2)
508
- _flash_attn_backward(
509
- dout,
510
- q,
511
- k,
512
- v,
513
- out,
514
- softmax_lse,
515
- None, None, # cu_seqlens_q, cu_seqlens_k,
516
- None, None, # sequed_q, sequed_k,
517
- None, None, # max_seqlen_q, max_seqlen_k,
518
- dq,
519
- dk,
520
- dv,
521
- ctx.softmax_scale,
522
- ctx.causal,
523
- ctx.window_size[0],
524
- ctx.window_size[1],
525
- ctx.softcap,
526
- ctx.deterministic,
527
- ctx.sm_margin,
528
- )
529
- dqkv = dqkv[..., : dout.shape[-1]] # We could have padded the head dimension
530
- return dqkv, None, None, None, None, None, None, None, None, None, None, None, None
531
-
532
-
533
- class FlashAttnFunc(torch.autograd.Function):
534
-
535
- @staticmethod
536
- def forward(
537
- ctx,
538
- q,
539
- k,
540
- v,
541
- softmax_scale,
542
- causal,
543
- qv=None,
544
- q_descale=None, k_descale=None, v_descale=None,
545
- window_size=(-1, -1),
546
- attention_chunk=0,
547
- softcap=0.0,
548
- num_splits=1,
549
- pack_gqa=None,
550
- deterministic=False,
551
- sm_margin=0,
552
- return_softmax=False,
553
- ):
554
- if softmax_scale is None:
555
- softmax_scale = (q.shape[-1] + (qv.shape[-1] if qv is not None else 0)) ** (-0.5)
556
- # out, q, k, v, out_padded, softmax_lse = _flash_attn_forward(
557
- out, softmax_lse, *rest = _flash_attn_forward(
558
- q,
559
- k,
560
- v,
561
- None, None, # k_new, v_new
562
- qv, # qv
563
- None, # out
564
- None, None, None, # cu_seqlens_q/k/k_new
565
- None, None, # seqused_q/k
566
- None, None, # max_seqlen_q/k
567
- None, None, None, # page_table, kv_batch_idx, leftpad_k,
568
- None, None, None, # rotary_cos/sin, seqlens_rotary
569
- q_descale, k_descale, v_descale,
570
- softmax_scale,
571
- causal=causal,
572
- window_size_left=window_size[0],
573
- window_size_right=window_size[1],
574
- attention_chunk=attention_chunk,
575
- softcap=softcap,
576
- num_splits=num_splits,
577
- pack_gqa=pack_gqa,
578
- sm_margin=sm_margin,
579
- )
580
- # ctx.save_for_backward(q, k, v, out_padded, softmax_lse)
581
- ctx.save_for_backward(q, k, v, out, softmax_lse)
582
- ctx.softmax_scale = softmax_scale
583
- ctx.causal = causal
584
- ctx.window_size = window_size
585
- ctx.attention_chunk = attention_chunk
586
- ctx.softcap = softcap
587
- ctx.deterministic = deterministic
588
- ctx.sm_margin = sm_margin
589
- return (out, softmax_lse) if return_softmax else out
590
-
591
- @staticmethod
592
- def backward(ctx, dout, *args):
593
- q, k, v, out, softmax_lse = ctx.saved_tensors
594
- assert ctx.attention_chunk == 0, "FA3 backward does not support attention_chunk"
595
- dq, dk, dv = torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
596
- _flash_attn_backward(
597
- dout,
598
- q,
599
- k,
600
- v,
601
- out,
602
- softmax_lse,
603
- None, None, # cu_seqlens_q, cu_seqlens_k,
604
- None, None, # sequed_q, sequed_k,
605
- None, None, # max_seqlen_q, max_seqlen_k,
606
- dq,
607
- dk,
608
- dv,
609
- ctx.softmax_scale,
610
- ctx.causal,
611
- ctx.window_size[0],
612
- ctx.window_size[1],
613
- ctx.softcap,
614
- ctx.deterministic,
615
- ctx.sm_margin,
616
- )
617
- dq = dq[..., : q.shape[-1]] # We could have padded the head dimension
618
- dk = dk[..., : k.shape[-1]]
619
- dv = dv[..., : v.shape[-1]]
620
- return dq, dk, dv, None, None, None, None, None, None, None, None, None, None, None, None, None, None
621
-
622
-
623
- class FlashAttnVarlenFunc(torch.autograd.Function):
624
-
625
- @staticmethod
626
- def forward(
627
- ctx,
628
- q,
629
- k,
630
- v,
631
- cu_seqlens_q,
632
- cu_seqlens_k,
633
- seqused_q,
634
- seqused_k,
635
- max_seqlen_q,
636
- max_seqlen_k,
637
- softmax_scale,
638
- causal,
639
- qv=None,
640
- q_descale=None, k_descale=None, v_descale=None,
641
- window_size=(-1, -1),
642
- attention_chunk=0,
643
- softcap=0.0,
644
- num_splits=1,
645
- pack_gqa=None,
646
- deterministic=False,
647
- sm_margin=0,
648
- return_softmax=False,
649
- ):
650
- if softmax_scale is None:
651
- softmax_scale = (q.shape[-1] + (qv.shape[-1] if qv is not None else 0)) ** (-0.5)
652
- # out, q, k, v, out_padded, softmax_lse = _flash_attn_varlen_forward(
653
- out, softmax_lse, *rest = _flash_attn_forward(
654
- q,
655
- k,
656
- v,
657
- None, None, # k_new, v_new
658
- qv, # qv
659
- None, # out
660
- cu_seqlens_q,
661
- cu_seqlens_k,
662
- None, # cu_seqlens_k_new
663
- seqused_q,
664
- seqused_k,
665
- max_seqlen_q,
666
- max_seqlen_k,
667
- None, None, None, # page_table, kv_batch_idx, leftpad_k,
668
- None, None, None, # rotary_cos/sin, seqlens_rotary
669
- q_descale, k_descale, v_descale,
670
- softmax_scale,
671
- causal=causal,
672
- window_size_left=window_size[0],
673
- window_size_right=window_size[1],
674
- attention_chunk=attention_chunk,
675
- softcap=softcap,
676
- num_splits=num_splits,
677
- pack_gqa=pack_gqa,
678
- sm_margin=sm_margin,
679
- )
680
- # ctx.save_for_backward(q, k, v, out_padded, softmax_lse, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k)
681
- ctx.save_for_backward(q, k, v, out, softmax_lse, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k)
682
- ctx.max_seqlen_q = max_seqlen_q
683
- ctx.max_seqlen_k = max_seqlen_k
684
- ctx.softmax_scale = softmax_scale
685
- ctx.causal = causal
686
- ctx.window_size = window_size
687
- ctx.attention_chunk = attention_chunk
688
- ctx.softcap = softcap
689
- ctx.deterministic = deterministic
690
- ctx.sm_margin = sm_margin
691
- return (out, softmax_lse) if return_softmax else out
692
-
693
- @staticmethod
694
- def backward(ctx, dout, *args):
695
- q, k, v, out, softmax_lse, cu_seqlens_q, cu_seqlens_k, seqused_q, seqused_k = ctx.saved_tensors
696
- assert ctx.attention_chunk == 0, "FA3 backward does not support attention_chunk"
697
- dq, dk, dv = torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
698
- _flash_attn_backward(
699
- dout,
700
- q,
701
- k,
702
- v,
703
- out,
704
- softmax_lse,
705
- cu_seqlens_q,
706
- cu_seqlens_k,
707
- seqused_q,
708
- seqused_k,
709
- ctx.max_seqlen_q,
710
- ctx.max_seqlen_k,
711
- dq,
712
- dk,
713
- dv,
714
- ctx.softmax_scale,
715
- ctx.causal,
716
- ctx.window_size[0],
717
- ctx.window_size[1],
718
- ctx.softcap,
719
- ctx.deterministic,
720
- ctx.sm_margin,
721
- )
722
- dq = dq[..., : q.shape[-1]] # We could have padded the head dimension
723
- dk = dk[..., : k.shape[-1]]
724
- dv = dv[..., : v.shape[-1]]
725
- return dq, dk, dv, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None
726
-
727
-
728
- def flash_attn_qkvpacked_func(
729
- qkv,
730
- softmax_scale=None,
731
- causal=False,
732
- q_descale=None, k_descale=None, v_descale=None,
733
- window_size=(-1, -1),
734
- attention_chunk=0,
735
- softcap=0.0,
736
- deterministic=False,
737
- num_heads_q=None,
738
- sm_margin=0,
739
- return_attn_probs=False,
740
- ):
741
- """dropout_p should be set to 0.0 during evaluation
742
- If Q, K, V are already stacked into 1 tensor, this function will be faster than
743
- calling flash_attn_func on Q, K, V since the backward pass avoids explicit concatenation
744
- of the gradients of Q, K, V.
745
- For multi-query and grouped-query attention (MQA/GQA), please see
746
- flash_attn_kvpacked_func and flash_attn_func.
747
-
748
- If window_size != (-1, -1), implements sliding window local attention. Query at position i
749
- will only attend to keys between [i - window_size[0], i + window_size[1]] inclusive.
750
-
751
- Arguments:
752
- qkv: (batch_size, seqlen, 3, nheads, headdim)
753
- dropout_p: float. Dropout probability.
754
- softmax_scale: float. The scaling of QK^T before applying softmax.
755
- Default to 1 / sqrt(headdim).
756
- causal: bool. Whether to apply causal attention mask (e.g., for auto-regressive modeling).
757
- window_size: (left, right). If not (-1, -1), implements sliding window local attention.
758
- softcap: float. Anything > 0 activates softcapping attention.
759
- alibi_slopes: (nheads,) or (batch_size, nheads), fp32. A bias of (-alibi_slope * |i - j|) is added to
760
- the attention score of query i and key j.
761
- deterministic: bool. Whether to use the deterministic implementation of the backward pass,
762
- which is slightly slower and uses more memory. The forward pass is always deterministic.
763
- return_attn_probs: bool. Whether to return the attention probabilities. This option is for
764
- testing only. The returned probabilities are not guaranteed to be correct
765
- (they might not have the right scaling).
766
- Return:
767
- out: (batch_size, seqlen, nheads, headdim).
768
- softmax_lse [optional, if return_attn_probs=True]: (batch_size, nheads, seqlen). The
769
- logsumexp of each row of the matrix QK^T * scaling (e.g., log of the softmax
770
- normalization factor).
771
- S_dmask [optional, if return_attn_probs=True]: (batch_size, nheads, seqlen, seqlen).
772
- The output of softmax (possibly with different scaling). It also encodes the dropout
773
- pattern (negative means that location was dropped, nonnegative means it was kept).
774
- """
775
- return FlashAttnQKVPackedFunc.apply(
776
- qkv,
777
- softmax_scale,
778
- causal,
779
- q_descale, k_descale, v_descale,
780
- window_size,
781
- attention_chunk,
782
- softcap,
783
- deterministic,
784
- num_heads_q,
785
- sm_margin,
786
- return_attn_probs,
787
- )
788
-
789
-
790
- def flash_attn_func(
791
- q,
792
- k,
793
- v,
794
- softmax_scale=None,
795
- causal=False,
796
- qv=None,
797
- q_descale=None, k_descale=None, v_descale=None,
798
- window_size=(-1, -1),
799
- attention_chunk=0,
800
- softcap=0.0,
801
- num_splits=1,
802
- pack_gqa=None,
803
- deterministic=False,
804
- sm_margin=0,
805
- return_attn_probs=False,
806
- ):
807
- """dropout_p should be set to 0.0 during evaluation
808
- Supports multi-query and grouped-query attention (MQA/GQA) by passing in KV with fewer heads
809
- than Q. Note that the number of heads in Q must be divisible by the number of heads in KV.
810
- For example, if Q has 6 heads and K, V have 2 heads, head 0, 1, 2 of Q will attention to head
811
- 0 of K, V, and head 3, 4, 5 of Q will attention to head 1 of K, V.
812
-
813
- If causal=True, the causal mask is aligned to the bottom right corner of the attention matrix.
814
- For example, if seqlen_q = 2 and seqlen_k = 5, the causal mask (1 = keep, 0 = masked out) is:
815
- 1 1 1 1 0
816
- 1 1 1 1 1
817
- If seqlen_q = 5 and seqlen_k = 2, the causal mask is:
818
- 0 0
819
- 0 0
820
- 0 0
821
- 1 0
822
- 1 1
823
- If the row of the mask is all zero, the output will be zero.
824
-
825
- If window_size != (-1, -1), implements sliding window local attention. Query at position i
826
- will only attend to keys between
827
- [i + seqlen_k - seqlen_q - window_size[0], i + seqlen_k - seqlen_q + window_size[1]] inclusive.
828
-
829
- Arguments:
830
- q: (batch_size, seqlen, nheads, headdim)
831
- k: (batch_size, seqlen, nheads_k, headdim)
832
- v: (batch_size, seqlen, nheads_k, headdim)
833
- dropout_p: float. Dropout probability.
834
- softmax_scale: float. The scaling of QK^T before applying softmax.
835
- Default to 1 / sqrt(headdim).
836
- causal: bool. Whether to apply causal attention mask (e.g., for auto-regressive modeling).
837
- window_size: (left, right). If not (-1, -1), implements sliding window local attention.
838
- alibi_slopes: (nheads,) or (batch_size, nheads), fp32. A bias of
839
- (-alibi_slope * |i + seqlen_k - seqlen_q - j|)
840
- is added to the attention score of query i and key j.
841
- deterministic: bool. Whether to use the deterministic implementation of the backward pass,
842
- which is slightly slower and uses more memory. The forward pass is always deterministic.
843
- return_attn_probs: bool. Whether to return the attention probabilities. This option is for
844
- testing only. The returned probabilities are not guaranteed to be correct
845
- (they might not have the right scaling).
846
- Return:
847
- out: (batch_size, seqlen, nheads, headdim).
848
- softmax_lse [optional, if return_attn_probs=True]: (batch_size, nheads, seqlen). The
849
- logsumexp of each row of the matrix QK^T * scaling (e.g., log of the softmax
850
- normalization factor).
851
- """
852
- return FlashAttnFunc.apply(
853
- q,
854
- k,
855
- v,
856
- softmax_scale,
857
- causal,
858
- qv,
859
- q_descale, k_descale, v_descale,
860
- window_size,
861
- attention_chunk,
862
- softcap,
863
- num_splits,
864
- pack_gqa,
865
- deterministic,
866
- sm_margin,
867
- return_attn_probs,
868
- )
869
-
870
-
871
- def flash_attn_varlen_func(
872
- q,
873
- k,
874
- v,
875
- cu_seqlens_q,
876
- cu_seqlens_k,
877
- max_seqlen_q,
878
- max_seqlen_k,
879
- seqused_q=None,
880
- seqused_k=None,
881
- softmax_scale=None,
882
- causal=False,
883
- qv=None,
884
- q_descale=None, k_descale=None, v_descale=None,
885
- window_size=(-1, -1),
886
- attention_chunk=0,
887
- softcap=0.0,
888
- num_splits=1,
889
- pack_gqa=None,
890
- deterministic=False,
891
- sm_margin=0,
892
- return_attn_probs=False,
893
- ):
894
- return FlashAttnVarlenFunc.apply(
895
- q,
896
- k,
897
- v,
898
- cu_seqlens_q,
899
- cu_seqlens_k,
900
- seqused_q,
901
- seqused_k,
902
- max_seqlen_q,
903
- max_seqlen_k,
904
- softmax_scale,
905
- causal,
906
- qv,
907
- q_descale, k_descale, v_descale,
908
- window_size,
909
- attention_chunk,
910
- softcap,
911
- num_splits,
912
- pack_gqa,
913
- deterministic,
914
- sm_margin,
915
- return_attn_probs,
916
- )
917
-
918
-
919
- def flash_attn_combine(out_partial, lse_partial, out=None, out_dtype=None):
920
- return flash_attn_3_cuda.fwd_combine(out_partial, lse_partial, out, out_dtype)
921
-
922
-
923
- def flash_attn_with_kvcache(
924
- q,
925
- k_cache,
926
- v_cache,
927
- k=None,
928
- v=None,
929
- qv=None,
930
- rotary_cos=None,
931
- rotary_sin=None,
932
- cache_seqlens: Optional[Union[(int, torch.Tensor)]] = None,
933
- cache_batch_idx: Optional[torch.Tensor] = None,
934
- cache_leftpad: Optional[torch.Tensor] = None,
935
- page_table: Optional[torch.Tensor] = None,
936
- cu_seqlens_q: Optional[torch.Tensor] = None,
937
- cu_seqlens_k_new: Optional[torch.Tensor] = None,
938
- max_seqlen_q: Optional[int] = None,
939
- rotary_seqlens: Optional[torch.Tensor] = None,
940
- q_descale: Optional[torch.Tensor] = None,
941
- k_descale: Optional[torch.Tensor] = None,
942
- v_descale: Optional[torch.Tensor] = None,
943
- softmax_scale=None,
944
- causal=False,
945
- window_size=(-1, -1), # -1 means infinite context window
946
- attention_chunk=0,
947
- softcap=0.0, # 0.0 means deactivated
948
- rotary_interleaved=True,
949
- scheduler_metadata=None,
950
- num_splits=0, # Can be tuned for speed
951
- pack_gqa=None, # Can be tuned for speed
952
- sm_margin=0, # Can be tuned if some SMs are used for communication
953
- return_softmax_lse=False,
954
- ):
955
- """
956
- If k and v are not None, k_cache and v_cache will be updated *inplace* with the new values from
957
- k and v. This is useful for incremental decoding: you can pass in the cached keys/values from
958
- the previous step, and update them with the new keys/values from the current step, and do
959
- attention with the updated cache, all in 1 kernel.
960
-
961
- If you pass in k / v, you must make sure that the cache is large enough to hold the new values.
962
- For example, the KV cache could be pre-allocated with the max sequence length, and you can use
963
- cache_seqlens to keep track of the current sequence lengths of each sequence in the batch.
964
-
965
- Also apply rotary embedding if rotary_cos and rotary_sin are passed in. The key @k will be
966
- rotated by rotary_cos and rotary_sin at indices cache_seqlens, cache_seqlens + 1, etc.
967
- If causal or local (i.e., window_size != (-1, -1)), the query @q will be rotated by rotary_cos
968
- and rotary_sin at indices cache_seqlens, cache_seqlens + 1, etc.
969
- If not causal and not local, the query @q will be rotated by rotary_cos and rotary_sin at
970
- indices cache_seqlens only (i.e. we consider all tokens in @q to be at position cache_seqlens).
971
-
972
- See tests/test_flash_attn.py::test_flash_attn_kvcache for examples of how to use this function.
973
-
974
- Supports multi-query and grouped-query attention (MQA/GQA) by passing in KV with fewer heads
975
- than Q. Note that the number of heads in Q must be divisible by the number of heads in KV.
976
- For example, if Q has 6 heads and K, V have 2 heads, head 0, 1, 2 of Q will attention to head
977
- 0 of K, V, and head 3, 4, 5 of Q will attention to head 1 of K, V.
978
-
979
- If causal=True, the causal mask is aligned to the bottom right corner of the attention matrix.
980
- For example, if seqlen_q = 2 and seqlen_k = 5, the causal mask (1 = keep, 0 = masked out) is:
981
- 1 1 1 1 0
982
- 1 1 1 1 1
983
- If seqlen_q = 5 and seqlen_k = 2, the causal mask is:
984
- 0 0
985
- 0 0
986
- 0 0
987
- 1 0
988
- 1 1
989
- If the row of the mask is all zero, the output will be zero.
990
-
991
- If window_size != (-1, -1), implements sliding window local attention. Query at position i
992
- will only attend to keys between
993
- [i + seqlen_k - seqlen_q - window_size[0], i + seqlen_k - seqlen_q + window_size[1]] inclusive.
994
-
995
- Note: Does not support backward pass.
996
-
997
- Arguments:
998
- q: (batch_size, seqlen, nheads, headdim)
999
- k_cache: (batch_size_cache, seqlen_cache, nheads_k, headdim) if there's no page_table,
1000
- or (num_blocks, page_block_size, nheads_k, headdim) if there's a page_table (i.e. paged KV cache)
1001
- page_block_size can be arbitrary (e.g, 1, 2, 3, 64, etc.).
1002
- v_cache: (batch_size_cache, seqlen_cache, nheads_k, headdim_v) if there's no page_table,
1003
- or (num_blocks, page_block_size, nheads_k, headdim_v) if there's a page_table (i.e. paged KV cache)
1004
- k [optional]: (batch_size, seqlen_new, nheads_k, headdim). If not None, we concatenate
1005
- k with k_cache, starting at the indices specified by cache_seqlens.
1006
- v [optional]: (batch_size, seqlen_new, nheads_k, headdim_v). Similar to k.
1007
- qv [optional]: (batch_size, seqlen, nheads, headdim_v)
1008
- rotary_cos [optional]: (seqlen_ro, rotary_dim / 2). If not None, we apply rotary embedding
1009
- to k and q. Only applicable if k and v are passed in. rotary_dim must be divisible by 16.
1010
- rotary_sin [optional]: (seqlen_ro, rotary_dim / 2). Similar to rotary_cos.
1011
- cache_seqlens: int, or (batch_size,), dtype torch.int32. The sequence lengths of the
1012
- KV cache.
1013
- cache_batch_idx: (batch_size,), dtype torch.int32. The indices used to index into the KV cache.
1014
- If None, we assume that the batch indices are [0, 1, 2, ..., batch_size - 1].
1015
- If the indices are not distinct, and k and v are provided, the values updated in the cache
1016
- might come from any of the duplicate indices.
1017
- cache_leftpad: (batch_size,), dtype torch.int32. The index that the KV cache starts. If None, assume 0.
1018
- page_table [optional]: (batch_size, max_num_blocks_per_seq), dtype torch.int32.
1019
- softmax_scale: float. The scaling of QK^T before applying softmax.
1020
- Default to 1 / sqrt(headdim).
1021
- causal: bool. Whether to apply causal attention mask (e.g., for auto-regressive modeling).
1022
- window_size: (left, right). If not (-1, -1), implements sliding window local attention.
1023
- softcap: float. Anything > 0 activates softcapping attention.
1024
- rotary_interleaved: bool. Only applicable if rotary_cos and rotary_sin are passed in.
1025
- If True, rotary embedding will combine dimensions 0 & 1, 2 & 3, etc. If False,
1026
- rotary embedding will combine dimensions 0 & rotary_dim / 2, 1 & rotary_dim / 2 + 1
1027
- (i.e. GPT-NeoX style).
1028
- num_splits: int. If > 1, split the key/value into this many chunks along the sequence.
1029
- If num_splits == 1, we don't split the key/value. If num_splits == 0, we use a heuristic
1030
- to automatically determine the number of splits.
1031
- Don't change this unless you know what you are doing.
1032
- return_softmax_lse: bool. Whether to return the logsumexp of the attention scores.
1033
-
1034
- Return:
1035
- out: (batch_size, seqlen, nheads, headdim).
1036
- softmax_lse [optional, if return_softmax_lse=True]: (batch_size, nheads, seqlen). The
1037
- logsumexp of each row of the matrix QK^T * scaling (e.g., log of the softmax
1038
- normalization factor).
1039
- """
1040
- assert k_cache.stride(-1) == 1, "k_cache must have contiguous last dimension"
1041
- assert v_cache.stride(-1) == 1, "v_cache must have contiguous last dimension"
1042
- if softmax_scale is None:
1043
- softmax_scale = (q.shape[-1] + (qv.shape[-1] if qv is not None else 0)) ** (-0.5)
1044
- if cache_seqlens is not None and isinstance(cache_seqlens, int):
1045
- cache_seqlens = torch.full(
1046
- (q.shape[0],), cache_seqlens, dtype=torch.int32, device=k_cache.device
1047
- )
1048
- cache_seqlens = maybe_contiguous(cache_seqlens)
1049
- out, softmax_lse, *rest = _flash_attn_forward(
1050
- q,
1051
- k_cache,
1052
- v_cache,
1053
- k,
1054
- v,
1055
- qv,
1056
- None, # out
1057
- cu_seqlens_q,
1058
- None, # cu_seqlens_k
1059
- cu_seqlens_k_new,
1060
- None, # seqused_q
1061
- cache_seqlens,
1062
- max_seqlen_q,
1063
- None, # max_seqlen_k
1064
- page_table,
1065
- cache_batch_idx,
1066
- cache_leftpad,
1067
- rotary_cos,
1068
- rotary_sin,
1069
- rotary_seqlens,
1070
- q_descale, k_descale, v_descale,
1071
- softmax_scale,
1072
- causal=causal,
1073
- window_size_left=window_size[0],
1074
- window_size_right=window_size[1],
1075
- attention_chunk=attention_chunk,
1076
- softcap=softcap,
1077
- rotary_interleaved=rotary_interleaved,
1078
- scheduler_metadata=scheduler_metadata,
1079
- num_splits=num_splits,
1080
- pack_gqa=pack_gqa,
1081
- sm_margin=sm_margin,
1082
- )
1083
- # return (out, softmax_lse) if return_softmax_lse else out
1084
- return (out, softmax_lse, *rest) if return_softmax_lse else out
1085
-
1086
-
1087
- def get_scheduler_metadata(
1088
- batch_size, max_seqlen_q, max_seqlen_k, num_heads_q, num_heads_kv, headdim,
1089
- cache_seqlens: torch.Tensor,
1090
- qkv_dtype=torch.bfloat16,
1091
- headdim_v=None,
1092
- cu_seqlens_q: Optional[torch.Tensor] = None,
1093
- cu_seqlens_k_new: Optional[torch.Tensor] = None,
1094
- cache_leftpad: Optional[torch.Tensor] = None,
1095
- page_size: Optional[int] = None,
1096
- max_seqlen_k_new=0,
1097
- causal=False,
1098
- window_size=(-1, -1), # -1 means infinite context window
1099
- attention_chunk=0,
1100
- has_softcap=False,
1101
- num_splits=0, # Can be tuned for speed
1102
- pack_gqa=None, # Can be tuned for speed
1103
- sm_margin=0, # Can be tuned if some SMs are used for communication
1104
- ):
1105
- cache_seqlens = maybe_contiguous(cache_seqlens)
1106
- if headdim_v is None:
1107
- headdim_v = headdim
1108
- scheduler_metadata = flash_attn_3_cuda.get_scheduler_metadata(
1109
- batch_size, max_seqlen_q, max_seqlen_k, num_heads_q, num_heads_kv, headdim, headdim_v,
1110
- qkv_dtype,
1111
- cache_seqlens,
1112
- cu_seqlens_q,
1113
- None, # cu_seqlens_k
1114
- cu_seqlens_k_new,
1115
- None, # seqused_q
1116
- cache_leftpad,
1117
- page_size,
1118
- max_seqlen_k_new,
1119
- causal,
1120
- window_size[0], window_size[1],
1121
- attention_chunk,
1122
- has_softcap,
1123
- num_splits,
1124
- pack_gqa,
1125
- sm_margin,
1126
- )
1127
- return scheduler_metadata
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu130-x86_64-linux/metadata.json DELETED
@@ -1,25 +0,0 @@
1
- {
2
- "name": "flash-attn3",
3
- "id": "_flash_attn3_cuda_477ab85",
4
- "version": 1,
5
- "license": "BSD-3-Clause",
6
- "python-depends": [],
7
- "backend": {
8
- "type": "cuda",
9
- "archs": [
10
- "8.0",
11
- "9.0a"
12
- ]
13
- },
14
- "digest": {
15
- "algorithm": "sha256",
16
- "files": {
17
- "__init__.py": "KXVmQJM+KhWc2UqWJesupCaP+mqAvTx7yGueFATglHE=",
18
- "_flash_attn3_cuda_477ab85.abi3.so": "p5w4yjy/X2Yw6W6xq2Aishvn2T4qokYdDac0Gg8kPzA=",
19
- "_ops.py": "hccJrV95SPARE3ynEBHpwny/YFXZqpZKdGdbTTnKTsc=",
20
- "flash_attn3/__init__.py": "DFYPlrhXwYjEqCl/8n0SmWGZV8NFml5DPhMjKfv98GY=",
21
- "flash_attn_config.py": "uxo+eyMcDit//8YDaS4Rtk5BTMKzTSK9tOKTjdWVe/w=",
22
- "flash_attn_interface.py": "y0bD2JYGMFgX6iDre3zJN2sBPuq/12N6Mo55Z+HQCNQ="
23
- }
24
- }
25
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
build/torch-stable-abi210-cu130-x86_64-linux/metadata.json.sigstore DELETED
@@ -1 +0,0 @@
1
- {"mediaType":"application/vnd.dev.sigstore.bundle.v0.3+json","verificationMaterial":{"certificate":{"rawBytes":"MIIHczCCBvmgAwIBAgIUJfQ+Lafjg9MCHV3TGRl395a32pUwCgYIKoZIzj0EAwMwNzEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MR4wHAYDVQQDExVzaWdzdG9yZS1pbnRlcm1lZGlhdGUwHhcNMjYwNjE4MDY1NTAxWhcNMjYwNjE4MDcwNTAxWjAAMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAENpvmnb5gaUtD1jcrytKnsYrnzaVQOD/2vticuK0XRM4jxJYYXF7keAmNJsyhS8mEShuBz1pKOAyef4lwZpfLEKOCBhgwggYUMA4GA1UdDwEB/wQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcDAzAdBgNVHQ4EFgQUUCF7G3dnHT6jBpsyZHYvlNf43aQwHwYDVR0jBBgwFoAU39Ppz1YkEZb5qNjpKFWixi4YZD8wdQYDVR0RAQH/BGswaYZnaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL3NpZ24tb2xkLWJ1aWxkcy55YW1sQHJlZnMvaGVhZHMvbWFpbjA5BgorBgEEAYO/MAEBBCtodHRwczovL3Rva2VuLmFjdGlvbnMuZ2l0aHVidXNlcmNvbnRlbnQuY29tMB8GCisGAQQBg78wAQIEEXdvcmtmbG93X2Rpc3BhdGNoMDYGCisGAQQBg78wAQMEKDhhNmJlN2JjNzc1NjVhZDY1YzhhZTJkZWI0NTY0ODJmZjZhYTUwZWQwHQYKKwYBBAGDvzABBAQPU2lnbiBvbGQgYnVpbGRzMCsGCisGAQQBg78wAQUEHWh1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MB0GCisGAQQBg78wAQYED3JlZnMvaGVhZHMvbWFpbjA7BgorBgEEAYO/MAEIBC0MK2h0dHBzOi8vdG9rZW4uYWN0aW9ucy5naXRodWJ1c2VyY29udGVudC5jb20wdwYKKwYBBAGDvzABCQRpDGdodHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHkvLmdpdGh1Yi93b3JrZmxvd3Mvc2lnbi1vbGQtYnVpbGRzLnlhbWxAcmVmcy9oZWFkcy9tYWluMDgGCisGAQQBg78wAQoEKgwoOGE2YmU3YmM3NzU2NWFkNjVjOGFlMmRlYjQ1NjQ4MmZmNmFhNTBlZDAbBgorBgEEAYO/MAELBA0MC3NlbGYtaG9zdGVkMEAGCisGAQQBg78wAQwEMgwwaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MDgGCisGAQQBg78wAQ0EKgwoOGE2YmU3YmM3NzU2NWFkNjVjOGFlMmRlYjQ1NjQ4MmZmNmFhNTBlZDAfBgorBgEEAYO/MAEOBBEMD3JlZnMvaGVhZHMvbWFpbjAaBgorBgEEAYO/MAEPBAwMCjEwNzE0NzU1MjkwLgYKKwYBBAGDvzABEAQgDB5odHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2UwGAYKKwYBBAGDvzABEQQKDAgyNTcyMDc0MzB3BgorBgEEAYO/MAESBGkMZ2h0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS8uZ2l0aHViL3dvcmtmbG93cy9zaWduLW9sZC1idWlsZHMueWFtbEByZWZzL2hlYWRzL21haW4wOAYKKwYBBAGDvzABEwQqDCg4YTZiZTdiYzc3NTY1YWQ2NWM4YWUyZGViNDU2NDgyZmY2YWE1MGVkMCEGCisGAQQBg78wARQEEwwRd29ya2Zsb3dfZGlzcGF0Y2gwZAYKKwYBBAGDvzABFQRWDFRodHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHkvYWN0aW9ucy9ydW5zLzI3NzQxNDY0ODQyL2F0dGVtcHRzLzEwFgYKKwYBBAGDvzABFgQIDAZwdWJsaWMwRgYKKwYBBAGDvzABGAQ4DDZyZXBvOmh1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5OnJlZjpyZWZzL2hlYWRzL21haW4wgYoGCisGAQQB1nkCBAIEfAR6AHgAdgDdPTBqxscRMmMZHhyZZzcCokpeuN48rf+HinKALynujgAAAZ7Zgv3+AAAEAwBHMEUCIE6dtVPTFDubQnuGhiiA44Z8K+lriLwqt46E5+jbFUnUAiEAlAhRFEBqfL+6w/oFp8HiRhfsQVaNEm9BkfCI0hd5p1UwCgYIKoZIzj0EAwMDaAAwZQIxAMmnhYQt49xjuMCWBVkqE9IsTm/zbWeQ4sPNlzaw5d2YwFFM3dVf3E1SR+GxM9E7SQIwLh8Ww9UDMsy7MBBTj2WILk2pOMgV9LKcquxjvZKtq6kXPSGIWOfzHwDu+sWcaktn"},"tlogEntries":[{"logIndex":"1857834987","logId":{"keyId":"wNI9atQGlz+VWfO6LRygH4QUfY/8W4RFwiT5i5WRgB0="},"kindVersion":{"kind":"hashedrekord","version":"0.0.1"},"integratedTime":"1781765701","inclusionPromise":{"signedEntryTimestamp":"MEQCIEAD3i4bKuxCq5swVTQ0FYYIi4BiIpRO7XUGKG7rPoO7AiAUXzxEMIq2+TNxISsCH0AcrIzbKqcvDF59HQ2vw9FtTA=="},"inclusionProof":{"logIndex":"1735930725","rootHash":"vDb9AH9FSARY5f3U5CNoY04d5KHKh45hqYgFtyLw620=","treeSize":"1735930730","hashes":["j0e/NdB0h5fx2ucsjiiOeM4T9zrebbjFjcjXMs2FAzU=","/Cy7iEAHBbxosFzFaZOeD5CYcb/vnZ2FcYiGujSqE3I=","a6PuAqpoQ1EFzAz+T/DG/OTN57hi6wSHVScV566Hsyk=","H2cZWajzm/8Z0uxjztzJLAE8LunGyjo3M8OJ9G/B6bo=","AKkFKs+fy2ovKwhkx7RasX6W34gOJbPuo/YgbATa+2A=","gP/ZYVz9iLEfjKC7esTiDzwHeStGfT8SzNgpwEYK0Ic=","4lyHa8i8t5GIm7VEJ+TbO30Wa4dBA9X6oIrqgy/9z/g=","M9Vb0zzLaw2h2vZMtZalUKLUlQkj3euyLByPs5Zb8A8=","uMjlW8ZYxFHFISKF1gqKEB4T/pfjDh7vzUC1E7W0/xU=","Ore1A2Ceavap3m5U3ZAHLkUinE8F4IgttCeJDxxaK98=","95goX8Iw4H4P6V4eHvRqTdVdViLbhEhsF6i7ICnz/+A=","sWDh7SjDHbJf3HWKGRxiURh6iIYrOzn4Zs37yIij6OA=","x7kXd4VRJvVmTHoSla2KyPQdKvGKDX26/7ST9OONR78=","U8wMiVDzFlmyqT7Nw1RSZYU9+fftsSkRhjpbyXnXUk8=","mqM7J+i75IpuD09QejiUBjeH85AZa+fSm4RXXFmj3lI=","72FC5FYLhxB4a4iiC956o0B/fT54ip4R41vsw2QBKtU=","lYGQ9ibwC8+smMkPQ6TchJm3H9Nc/aTYLfdRacGFChw=","daxmZaajRpZV+JxHiOYZhJBiSKN5ucqjh2WnGbHhirw=","DOCeoSMovIvLExkhIvisow9AuNXgeWs4ECkyR6EcqYU="],"checkpoint":{"envelope":"rekor.sigstore.dev - 1193050959916656506\n1735930730\nvDb9AH9FSARY5f3U5CNoY04d5KHKh45hqYgFtyLw620=\n\n— rekor.sigstore.dev wNI9ajBEAiAvyRAF+JA7P4JSV+GEhj82lITA58lzVnPBnJuwwY6iqAIgTgr36k/FODTBtOivs6yxAhGxWFNP2pWgxGzMjg+LLP0=\n"}},"canonicalizedBody":"eyJhcGlWZXJzaW9uIjoiMC4wLjEiLCJraW5kIjoiaGFzaGVkcmVrb3JkIiwic3BlYyI6eyJkYXRhIjp7Imhhc2giOnsiYWxnb3JpdGhtIjoic2hhMjU2IiwidmFsdWUiOiJhYTM4OGRmN2E4NDJiZWU4N2FjOTNjMTE3Y2IwNDE2ZDMzM2QzOTA2YzFhZDZkYTJlNDIzYmEyN2U3ZjUyMmI2In19LCJzaWduYXR1cmUiOnsiY29udGVudCI6Ik1FVUNJSEN0cEdUTEUwZnROZ1B3NnMyQ2wySDJLd0V0MFJrbGt2VERVcTBNWmwzN0FpRUE5RlRwTm5vQTZIQ0taajl5Yk9HQ3BhREhiT0JKUFF6RjRFTDZ5a25HbU1nPSIsInB1YmxpY0tleSI6eyJjb250ZW50IjoiTFMwdExTMUNSVWRKVGlCRFJWSlVTVVpKUTBGVVJTMHRMUzB0Q2sxSlNVaGpla05EUW5adFowRjNTVUpCWjBsVlNtWlJLMHhoWm1wbk9VMURTRll6VkVkU2JETTVOV0V6TW5CVmQwTm5XVWxMYjFwSmVtb3dSVUYzVFhjS1RucEZWazFDVFVkQk1WVkZRMmhOVFdNeWJHNWpNMUoyWTIxVmRWcEhWakpOVWpSM1NFRlpSRlpSVVVSRmVGWjZZVmRrZW1SSE9YbGFVekZ3WW01U2JBcGpiVEZzV2tkc2FHUkhWWGRJYUdOT1RXcFpkMDVxUlRSTlJGa3hUbFJCZUZkb1kwNU5hbGwzVG1wRk5FMUVZM2RPVkVGNFYycEJRVTFHYTNkRmQxbElDa3R2V2tsNmFqQkRRVkZaU1V0dldrbDZhakJFUVZGalJGRm5RVVZPY0hadGJtSTFaMkZWZEVReGFtTnllWFJMYm5OWmNtNTZZVlpSVDBRdk1uWjBhV01LZFVzd1dGSk5OR3A0U2xsWldFWTNhMlZCYlU1S2MzbG9Vemh0UlZOb2RVSjZNWEJMVDBGNVpXWTBiSGRhY0daTVJVdFBRMEpvWjNkbloxbFZUVUUwUndwQk1WVmtSSGRGUWk5M1VVVkJkMGxJWjBSQlZFSm5UbFpJVTFWRlJFUkJTMEpuWjNKQ1owVkdRbEZqUkVGNlFXUkNaMDVXU0ZFMFJVWm5VVlZWUTBZM0NrY3paRzVJVkRacVFuQnplVnBJV1hac1RtWTBNMkZSZDBoM1dVUldVakJxUWtKbmQwWnZRVlV6T1ZCd2VqRlphMFZhWWpWeFRtcHdTMFpYYVhocE5Ga0tXa1E0ZDJSUldVUldVakJTUVZGSUwwSkhjM2RoV1ZwdVlVaFNNR05JVFRaTWVUbHVZVmhTYjJSWFNYVlpNamwwVERKb01Wb3laSEJpYldSdFdWZE9iQXBNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk1lVFZ1WVZoU2IyUlhTWFprTWpsNVlUSmFjMkl6WkhwTU0wNXdXakkwZEdJeWVHdE1WMG94Q21GWGVHdGplVFUxV1ZjeGMxRklTbXhhYmsxMllVZFdhRnBJVFhaaVYwWndZbXBCTlVKbmIzSkNaMFZGUVZsUEwwMUJSVUpDUTNSdlpFaFNkMk42YjNZS1RETlNkbUV5Vm5WTWJVWnFaRWRzZG1KdVRYVmFNbXd3WVVoV2FXUllUbXhqYlU1MlltNVNiR0p1VVhWWk1qbDBUVUk0UjBOcGMwZEJVVkZDWnpjNGR3cEJVVWxGUlZoa2RtTnRkRzFpUnpreldESlNjR016UW1oa1IwNXZUVVJaUjBOcGMwZEJVVkZDWnpjNGQwRlJUVVZMUkdob1RtMUtiRTR5U21wT2VtTXhDazVxVm1oYVJGa3hXWHBvYUZwVVNtdGFWMGt3VGxSWk1FOUVTbTFhYWxwb1dWUlZkMXBYVVhkSVVWbExTM2RaUWtKQlIwUjJla0ZDUWtGUlVGVXliRzRLWW1sQ2RtSkhVV2RaYmxad1lrZFNlazFEYzBkRGFYTkhRVkZSUW1jM09IZEJVVlZGU0Zkb01Wb3laSEJpYldSdFdWZE9iRXd5ZEd4amJUVnNZa2hOZEFwWk1qbDBZbGhXZFdGWVVqVk5RakJIUTJselIwRlJVVUpuTnpoM1FWRlpSVVF6U214YWJrMTJZVWRXYUZwSVRYWmlWMFp3WW1wQk4wSm5iM0pDWjBWRkNrRlpUeTlOUVVWSlFrTXdUVXN5YURCa1NFSjZUMms0ZG1SSE9YSmFWelIxV1ZkT01HRlhPWFZqZVRWdVlWaFNiMlJYU2pGak1sWjVXVEk1ZFdSSFZuVUtaRU0xYW1JeU1IZGtkMWxMUzNkWlFrSkJSMFIyZWtGQ1ExRlNjRVJIWkc5a1NGSjNZM3B2ZGt3eVpIQmtSMmd4V1drMWFtSXlNSFpoU0ZadVdqSnNkUXBhTWxwb1dUSlZkbUV5Vm5saWJWWnpZM2t4YW1JeU1YUmtWelZ3WkVocmRreHRaSEJrUjJneFdXazVNMkl6U25KYWJYaDJaRE5OZG1NeWJHNWlhVEYyQ21KSFVYUlpibFp3WWtkU2VreHViR2hpVjNoQlkyMVdiV041T1c5YVYwWnJZM2s1ZEZsWGJIVk5SR2RIUTJselIwRlJVVUpuTnpoM1FWRnZSVXRuZDI4S1QwZEZNbGx0VlROWmJVMHpUbnBWTWs1WFJtdE9hbFpxVDBkR2JFMXRVbXhaYWxFeFRtcFJORTF0V20xT2JVWm9UbFJDYkZwRVFXSkNaMjl5UW1kRlJRcEJXVTh2VFVGRlRFSkJNRTFETTA1c1lrZFpkR0ZIT1hwa1IxWnJUVVZCUjBOcGMwZEJVVkZDWnpjNGQwRlJkMFZOWjNkM1lVaFNNR05JVFRaTWVUbHVDbUZZVW05a1YwbDFXVEk1ZEV3eWFERmFNbVJ3WW0xa2JWbFhUbXhNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk5SR2RIUTJselIwRlJVVUlLWnpjNGQwRlJNRVZMWjNkdlQwZEZNbGx0VlROWmJVMHpUbnBWTWs1WFJtdE9hbFpxVDBkR2JFMXRVbXhaYWxFeFRtcFJORTF0V20xT2JVWm9UbFJDYkFwYVJFRm1RbWR2Y2tKblJVVkJXVTh2VFVGRlQwSkNSVTFFTTBwc1dtNU5kbUZIVm1oYVNFMTJZbGRHY0dKcVFXRkNaMjl5UW1kRlJVRlpUeTlOUVVWUUNrSkJkMDFEYWtWM1RucEZNRTU2VlRGTmFtdDNUR2RaUzB0M1dVSkNRVWRFZG5wQlFrVkJVV2RFUWpWdlpFaFNkMk42YjNaTU1tUndaRWRvTVZscE5Xb0tZakl3ZG1GSVZtNWFNbXgxV2pKYWFGa3lWWGRIUVZsTFMzZFpRa0pCUjBSMmVrRkNSVkZSUzBSQlozbE9WR041VFVSak1FMTZRak5DWjI5eVFtZEZSUXBCV1U4dlRVRkZVMEpIYTAxYU1tZ3daRWhDZWs5cE9IWmFNbXd3WVVoV2FVeHRUblppVXpsdlpGZGtibUZYTlc1YWJVWnFXbE01Y2xwWVNuVmFWM2g2Q2t4WFRuWmlWekV4WW0xc01HVlRPSFZhTW13d1lVaFdhVXd6WkhaamJYUnRZa2M1TTJONU9YcGhWMlIxVEZjNWMxcERNV2xrVjJ4eldraE5kV1ZYUm5RS1lrVkNlVnBYV25wTU1taHNXVmRTZWt3eU1XaGhWelIzVDBGWlMwdDNXVUpDUVVkRWRucEJRa1YzVVhGRVEyYzBXVlJhYVZwVVpHbFplbU16VGxSWk1RcFpWMUV5VGxkTk5GbFhWWGxhUjFacFRrUlZNazVFWjNsYWJWa3lXVmRGTVUxSFZtdE5RMFZIUTJselIwRlJVVUpuTnpoM1FWSlJSVVYzZDFKa01qbDVDbUV5V25OaU0yUm1Xa2RzZW1OSFJqQlpNbWQzV2tGWlMwdDNXVUpDUVVkRWRucEJRa1pSVWxkRVJsSnZaRWhTZDJONmIzWk1NbVJ3WkVkb01WbHBOV29LWWpJd2RtRklWbTVhTW14MVdqSmFhRmt5VlhaaE1sWjVZbTFXYzJONU1XcGlNakYwWkZjMWNHUklhM1paVjA0d1lWYzVkV041T1hsa1Z6VjZUSHBKTXdwT2VsRjRUa1JaTUU5RVVYbE1Na1l3WkVkV2RHTklVbnBNZWtWM1JtZFpTMHQzV1VKQ1FVZEVkbnBCUWtablVVbEVRVnAzWkZkS2MyRlhUWGRTWjFsTENrdDNXVUpDUVVkRWRucEJRa2RCVVRSRVJGcDVXbGhDZGs5dGFERmFNbVJ3WW0xa2JWbFhUbXhNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVUtUMjVLYkZwcWNIbGFWMXA2VERKb2JGbFhVbnBNTWpGb1lWYzBkMmRaYjBkRGFYTkhRVkZSUWpGdWEwTkNRVWxGWmtGU05rRklaMEZrWjBSa1VGUkNjUXA0YzJOU1RXMU5Xa2hvZVZwYWVtTkRiMnR3WlhWT05EaHlaaXRJYVc1TFFVeDViblZxWjBGQlFWbzNXbWQyTXl0QlFVRkZRWGRDU0UxRlZVTkpSVFprQ25SV1VGUkdSSFZpVVc1MVIyaHBhVUUwTkZvNFN5dHNjbWxNZDNGME5EWkZOU3RxWWtaVmJsVkJhVVZCYkVGb1VrWkZRbkZtVENzMmR5OXZSbkE0U0drS1VtaG1jMUZXWVU1RmJUbENhMlpEU1RCb1pEVndNVlYzUTJkWlNVdHZXa2w2YWpCRlFYZE5SR0ZCUVhkYVVVbDRRVTF0Ym1oWlVYUTBPWGhxZFUxRFZ3cENWbXR4UlRsSmMxUnRMM3BpVjJWUk5ITlFUbXg2WVhjMVpESlpkMFpHVFROa1ZtWXpSVEZUVWl0SGVFMDVSVGRUVVVsM1RHZzRWM2M1VlVSTmMzazNDazFDUWxScU1sZEpUR3N5Y0U5TloxWTVURXRqY1hWNGFuWmFTM1J4Tm10WVVGTkhTVmRQWm5wSWQwUjFLM05YWTJGcmRHNEtMUzB0TFMxRlRrUWdRMFZTVkVsR1NVTkJWRVV0TFMwdExRbz0ifX19fQ=="}],"timestampVerificationData":{"rfc3161Timestamps":[{"signedTimestamp":"MIICyTADAgEAMIICwAYJKoZIhvcNAQcCoIICsTCCAq0CAQMxDTALBglghkgBZQMEAgEwgbcGCyqGSIb3DQEJEAEEoIGnBIGkMIGhAgEBBgkrBgEEAYO/MAIwMTANBglghkgBZQMEAgEFAAQgM3snYkK+QerqMTYeiLeFAp1K7IhYpebX4FpEImlgKpkCFF0+tylnK5nEQDYFLTQO33WbjKipGA8yMDI2MDYxODA2NTUwMVowAwIBAaAypDAwLjEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MRUwEwYDVQQDEwxzaWdzdG9yZS10c2GgADGCAdswggHXAgEBMFEwOTEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MSAwHgYDVQQDExdzaWdzdG9yZS10c2Etc2VsZnNpZ25lZAIUOhNULwyQYe68wUMvy4qOiyojiwwwCwYJYIZIAWUDBAIBoIH8MBoGCSqGSIb3DQEJAzENBgsqhkiG9w0BCRABBDAcBgkqhkiG9w0BCQUxDxcNMjYwNjE4MDY1NTAxWjAvBgkqhkiG9w0BCQQxIgQgi5mZIMFbzwb41aJ5IY7pDWxYdilqKxefDwc7L0ymHOwwgY4GCyqGSIb3DQEJEAIvMX8wfTB7MHkEIIX5J7wHq2LKw7RDVsEO/IGyxog/2nq55thw2dE6zQW3MFUwPaQ7MDkxFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEgMB4GA1UEAxMXc2lnc3RvcmUtdHNhLXNlbGZzaWduZWQCFDoTVC8MkGHuvMFDL8uKjosqI4sMMAoGCCqGSM49BAMCBGcwZQIxAIQfb2j3XFhXe40MKIL+mCBDSiBeMPyBBeFs1kNPXxnw7n7ti0s+p076USHwAD/xfgIwcp8Vj4Meen+vDFZGHlzsK1GUJzKjz6doIVYwX5cwH2kGBUiwJ0mm5xZsDwEzxUII"}]}},"messageSignature":{"messageDigest":{"algorithm":"SHA2_256","digest":"qjiN96hCvuh6yTwRfLBBbTM9OQbBrW2i5CO6J+f1IrY="},"signature":"MEUCIHCtpGTLE0ftNgPw6s2Cl2H2KwEt0RklkvTDUq0MZl37AiEA9FTpNnoA6HCKZj9ybOGCpaDHbOBJPQzF4EL6yknGmMg="}}