Kernels:
Trusted publisher
File size: 1,800 Bytes
e19323e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 | # Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
from .csr import prepare_block_csr
from .cumsum import (
chunk_global_cumsum,
chunk_global_cumsum_scalar,
chunk_global_cumsum_vector,
chunk_local_cumsum,
chunk_local_cumsum_scalar,
chunk_local_cumsum_vector,
)
from .index import (
get_max_num_splits,
prepare_chunk_indices,
prepare_chunk_offsets,
prepare_cu_seqlens_from_lens,
prepare_cu_seqlens_from_mask,
prepare_lens,
prepare_lens_from_mask,
prepare_position_ids,
prepare_sequence_ids,
prepare_token_indices,
)
from .logsumexp import logsumexp_fwd
from .matmul import addmm, matmul
from .pack import pack_sequence, unpack_sequence
from .pooling import mean_pooling
from .softmax import softmax_bwd, softmax_fwd
from .softplus import softplus
from .solve_tril import solve_tril
__all__ = [
"addmm",
"chunk_global_cumsum",
"chunk_global_cumsum_scalar",
"chunk_global_cumsum_vector",
"chunk_local_cumsum",
"chunk_local_cumsum_scalar",
"chunk_local_cumsum_vector",
"get_max_num_splits",
"logsumexp_fwd",
"matmul",
"mean_pooling",
"pack_sequence",
"prepare_block_csr",
"prepare_chunk_indices",
"prepare_chunk_offsets",
"prepare_cu_seqlens_from_lens",
"prepare_cu_seqlens_from_mask",
"prepare_lens",
"prepare_lens_from_mask",
"prepare_position_ids",
"prepare_sequence_ids",
"prepare_token_indices",
"softmax_bwd",
"softmax_fwd",
"softplus",
"solve_tril",
"unpack_sequence",
]
|