fla / build /torch-cuda /ops /utils /__init__.py
kernels-bot's picture
Uploaded using `kernel-builder`.
e19323e verified
Raw
History Blame
1.8 kB
# 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",
]