From fa3f8ea45243d74a48e7b5c173a776a42e2aff98 Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Tue, 18 Aug 2026 22:50:21 +0800 Subject: [PATCH] feat(linked): add MetaX FlashAttention providers --- src/linked/torch/metax/flash_attn.yaml | 2 + .../ops/flash_attn_varlen_func/flash_attn.cc | 72 ++++++++++++ .../ops/flash_attn_varlen_func/flash_attn.h | 35 ++++++ .../flash_attn_varlen_func/flash_attn.yaml | 4 + .../ops/flash_attn_with_kvcache/flash_attn.cc | 110 ++++++++++++++++++ .../ops/flash_attn_with_kvcache/flash_attn.h | 53 +++++++++ .../flash_attn_with_kvcache/flash_attn.yaml | 4 + 7 files changed, 280 insertions(+) create mode 100644 src/linked/torch/metax/flash_attn.yaml create mode 100644 src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.cc create mode 100644 src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.h create mode 100644 src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.yaml create mode 100644 src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.cc create mode 100644 src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.h create mode 100644 src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.yaml diff --git a/src/linked/torch/metax/flash_attn.yaml b/src/linked/torch/metax/flash_attn.yaml new file mode 100644 index 000000000..c54d82dbd --- /dev/null +++ b/src/linked/torch/metax/flash_attn.yaml @@ -0,0 +1,2 @@ +python_distribution_package: flash-attn +library_glob: flash_attn_2_cuda*.so diff --git a/src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.cc b/src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.cc new file mode 100644 index 000000000..27b62f01c --- /dev/null +++ b/src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.cc @@ -0,0 +1,72 @@ +#include "linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.h" + +#include + +#include "linked/torch/ops/flash_attn_varlen_func.h" +#include "torch/metax/c10.h" + +std::vector mha_varlen_fwd( + at::Tensor& q, const at::Tensor& k, const at::Tensor& v, + std::optional& out, const at::Tensor& cu_seqlens_q, + const at::Tensor& cu_seqlens_k, std::optional& seqused_k, + std::optional& leftpad_k, + std::optional& block_table, + std::optional& alibi_slopes, int max_seqlen_q, int max_seqlen_k, + float dropout_p, float softmax_scale, bool zero_tensors, bool causal, + int window_size_left, int window_size_right, float softcap, + bool return_softmax, std::optional generator, + std::optional& flash_attn_mars_ext); + +namespace infini::ops::linked::torch::metax { + +struct FlashAttnVarlen : C10 { + static std::vector Call( + at::Tensor& q, const at::Tensor& k, const at::Tensor& v, + std::optional& out, const at::Tensor& cu_seqlens_q, + const at::Tensor& cu_seqlens_k, std::optional& seqused_k, + std::optional& leftpad_k, + std::optional& block_table, + std::optional& alibi_slopes, int max_seqlen_q, + int max_seqlen_k, float dropout_p, float softmax_scale, bool zero_tensors, + bool causal, int window_size_left, int window_size_right, float softcap, + bool return_softmax, std::optional generator) { + std::optional flash_attn_mars_ext; + return ::mha_varlen_fwd(q, k, v, out, cu_seqlens_q, cu_seqlens_k, seqused_k, + leftpad_k, block_table, alibi_slopes, max_seqlen_q, + max_seqlen_k, dropout_p, softmax_scale, + zero_tensors, causal, window_size_left, + window_size_right, softcap, return_softmax, + generator, flash_attn_mars_ext); + } +}; + +} // namespace infini::ops::linked::torch::metax + +namespace infini::ops { + +void Operator::operator()( + const Tensor q, const Tensor k, const Tensor v, const Tensor cu_seqlens_q, + const Tensor cu_seqlens_k, const std::optional alibi_slopes, + const std::optional block_table, const int64_t max_seqlen_q, + const int64_t max_seqlen_k, const double dropout_p, + const std::optional softmax_scale, const bool causal, + const std::vector window_size, const double softcap, + const bool deterministic, const bool return_attn_probs, Tensor out, + std::optional softmax_lse, std::optional s_dmask) const { + using Delegate = linked::torch::TorchFlashAttnVarlenFunc< + linked::torch::metax::FlashAttnVarlen>; + if (!delegate_) { + delegate_ = std::make_unique( + q, k, v, cu_seqlens_q, cu_seqlens_k, alibi_slopes, block_table, + max_seqlen_q, max_seqlen_k, dropout_p, softmax_scale, causal, + window_size, softcap, deterministic, return_attn_probs, out, + softmax_lse, s_dmask); + } + delegate_->set_stream(stream_); + (*delegate_)(q, k, v, cu_seqlens_q, cu_seqlens_k, alibi_slopes, block_table, + max_seqlen_q, max_seqlen_k, dropout_p, softmax_scale, causal, + window_size, softcap, deterministic, return_attn_probs, out, + softmax_lse, s_dmask); +} + +} // namespace infini::ops diff --git a/src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.h b/src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.h new file mode 100644 index 000000000..c0d9289d8 --- /dev/null +++ b/src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.h @@ -0,0 +1,35 @@ +#ifndef INFINI_OPS_LINKED_TORCH_METAX_OPS_FLASH_ATTN_VARLEN_FUNC_FLASH_ATTN_H_ +#define INFINI_OPS_LINKED_TORCH_METAX_OPS_FLASH_ATTN_VARLEN_FUNC_FLASH_ATTN_H_ + +#include + +#include "base/flash_attn_varlen_func.h" + +namespace infini::ops { + +template <> +class Operator + : public FlashAttnVarlenFunc { + public: + using FlashAttnVarlenFunc::FlashAttnVarlenFunc; + using FlashAttnVarlenFunc::operator(); + + void operator()(const Tensor q, const Tensor k, const Tensor v, + const Tensor cu_seqlens_q, const Tensor cu_seqlens_k, + const std::optional alibi_slopes, + const std::optional block_table, + const int64_t max_seqlen_q, const int64_t max_seqlen_k, + const double dropout_p, + const std::optional softmax_scale, const bool causal, + const std::vector window_size, const double softcap, + const bool deterministic, const bool return_attn_probs, + Tensor out, std::optional softmax_lse, + std::optional s_dmask) const override; + + private: + mutable std::unique_ptr delegate_; +}; + +} // namespace infini::ops + +#endif // INFINI_OPS_LINKED_TORCH_METAX_OPS_FLASH_ATTN_VARLEN_FUNC_FLASH_ATTN_H_ diff --git a/src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.yaml b/src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.yaml new file mode 100644 index 000000000..c692e1abd --- /dev/null +++ b/src/linked/torch/metax/ops/flash_attn_varlen_func/flash_attn.yaml @@ -0,0 +1,4 @@ +library: flash_attn +required_symbols: + - >- + mha_varlen_fwd(at::Tensor&, at::Tensor const&, at::Tensor const&, std::optional&, at::Tensor const&, at::Tensor const&, std::optional&, std::optional&, std::optional&, std::optional&, int, int, float, float, bool, bool, int, int, float, bool, std::optional, std::optional&) diff --git a/src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.cc b/src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.cc new file mode 100644 index 000000000..5405ad575 --- /dev/null +++ b/src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.cc @@ -0,0 +1,110 @@ +#include "linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.h" + +#include "linked/torch/ops/flash_attn_with_kvcache.h" +#include "torch/metax/c10.h" + +std::vector mha_fwd_kvcache( + at::Tensor& q, const at::Tensor& k_cache, const at::Tensor& v_cache, + std::optional& k, std::optional& v, + std::optional& cache_seqlens, + std::optional& rotary_cos, + std::optional& rotary_sin, + std::optional& cache_batch_idx, + std::optional& cache_leftpad, + std::optional& block_table, + std::optional& alibi_slopes, std::optional& out, + float softmax_scale, bool causal, int window_size_left, + int window_size_right, float softcap, bool rotary_interleaved, + int num_splits, std::optional& flash_attn_mars_ext); + +namespace infini::ops::linked::torch::metax { + +struct FlashAttnKvcache : C10 { + static std::vector Call( + at::Tensor& q, const at::Tensor& k_cache, const at::Tensor& v_cache, + std::optional& k, std::optional& v, + std::optional& cache_seqlens, + std::optional& rotary_cos, + std::optional& rotary_sin, + std::optional& cache_batch_idx, + std::optional& cache_leftpad, + std::optional& block_table, + std::optional& alibi_slopes, std::optional& out, + float softmax_scale, bool causal, int window_size_left, + int window_size_right, float softcap, bool rotary_interleaved, + int num_splits) { + std::optional flash_attn_mars_ext; + return ::mha_fwd_kvcache( + q, k_cache, v_cache, k, v, cache_seqlens, rotary_cos, rotary_sin, + cache_batch_idx, cache_leftpad, block_table, alibi_slopes, out, + softmax_scale, causal, window_size_left, window_size_right, softcap, + rotary_interleaved, num_splits, flash_attn_mars_ext); + } +}; + +} // namespace infini::ops::linked::torch::metax + +namespace infini::ops { + +void Operator::operator()( + const Tensor q, Tensor k_cache, Tensor v_cache, + const std::optional k, const std::optional v, + const std::optional rotary_cos, + const std::optional rotary_sin, const int64_t cache_seqlens, + const std::optional cache_batch_idx, + const std::optional cache_leftpad, + const std::optional block_table, + const std::optional alibi_slopes, + const std::optional softmax_scale, const bool causal, + const std::vector window_size, const double softcap, + const bool rotary_interleaved, const int64_t num_splits, + const bool return_softmax_lse, Tensor out, + std::optional softmax_lse) const { + using Delegate = linked::torch::TorchFlashAttnWithKvcache< + linked::torch::metax::FlashAttnKvcache>; + if (!delegate_) { + delegate_ = std::make_unique( + q, k_cache, v_cache, k, v, rotary_cos, rotary_sin, cache_seqlens, + cache_batch_idx, cache_leftpad, block_table, alibi_slopes, + softmax_scale, causal, window_size, softcap, rotary_interleaved, + num_splits, return_softmax_lse, out, softmax_lse); + } + delegate_->set_stream(stream_); + (*delegate_)(q, k_cache, v_cache, k, v, rotary_cos, rotary_sin, cache_seqlens, + cache_batch_idx, cache_leftpad, block_table, alibi_slopes, + softmax_scale, causal, window_size, softcap, rotary_interleaved, + num_splits, return_softmax_lse, out, softmax_lse); +} + +void Operator::operator()( + const Tensor q, Tensor k_cache, Tensor v_cache, + const std::optional k, const std::optional v, + const std::optional rotary_cos, + const std::optional rotary_sin, + const std::optional cache_seqlens, + const std::optional cache_batch_idx, + const std::optional cache_leftpad, + const std::optional block_table, + const std::optional alibi_slopes, + const std::optional softmax_scale, const bool causal, + const std::vector window_size, const double softcap, + const bool rotary_interleaved, const int64_t num_splits, + const bool return_softmax_lse, Tensor out, + std::optional softmax_lse) const { + using Delegate = linked::torch::TorchFlashAttnWithKvcache< + linked::torch::metax::FlashAttnKvcache>; + if (!delegate_) { + delegate_ = std::make_unique( + q, k_cache, v_cache, k, v, rotary_cos, rotary_sin, cache_seqlens, + cache_batch_idx, cache_leftpad, block_table, alibi_slopes, + softmax_scale, causal, window_size, softcap, rotary_interleaved, + num_splits, return_softmax_lse, out, softmax_lse); + } + delegate_->set_stream(stream_); + (*delegate_)(q, k_cache, v_cache, k, v, rotary_cos, rotary_sin, cache_seqlens, + cache_batch_idx, cache_leftpad, block_table, alibi_slopes, + softmax_scale, causal, window_size, softcap, rotary_interleaved, + num_splits, return_softmax_lse, out, softmax_lse); +} + +} // namespace infini::ops diff --git a/src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.h b/src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.h new file mode 100644 index 000000000..c7f12005c --- /dev/null +++ b/src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.h @@ -0,0 +1,53 @@ +#ifndef INFINI_OPS_LINKED_TORCH_METAX_OPS_FLASH_ATTN_WITH_KVCACHE_FLASH_ATTN_H_ +#define INFINI_OPS_LINKED_TORCH_METAX_OPS_FLASH_ATTN_WITH_KVCACHE_FLASH_ATTN_H_ + +#include + +#include "base/flash_attn_with_kvcache.h" + +namespace infini::ops { + +template <> +class Operator + : public FlashAttnWithKvcache { + public: + using FlashAttnWithKvcache::FlashAttnWithKvcache; + using FlashAttnWithKvcache::operator(); + + void operator()(const Tensor q, Tensor k_cache, Tensor v_cache, + const std::optional k, const std::optional v, + const std::optional rotary_cos, + const std::optional rotary_sin, + const int64_t cache_seqlens, + const std::optional cache_batch_idx, + const std::optional cache_leftpad, + const std::optional block_table, + const std::optional alibi_slopes, + const std::optional softmax_scale, const bool causal, + const std::vector window_size, const double softcap, + const bool rotary_interleaved, const int64_t num_splits, + const bool return_softmax_lse, Tensor out, + std::optional softmax_lse) const override; + + void operator()(const Tensor q, Tensor k_cache, Tensor v_cache, + const std::optional k, const std::optional v, + const std::optional rotary_cos, + const std::optional rotary_sin, + const std::optional cache_seqlens, + const std::optional cache_batch_idx, + const std::optional cache_leftpad, + const std::optional block_table, + const std::optional alibi_slopes, + const std::optional softmax_scale, const bool causal, + const std::vector window_size, const double softcap, + const bool rotary_interleaved, const int64_t num_splits, + const bool return_softmax_lse, Tensor out, + std::optional softmax_lse) const override; + + private: + mutable std::unique_ptr delegate_; +}; + +} // namespace infini::ops + +#endif // INFINI_OPS_LINKED_TORCH_METAX_OPS_FLASH_ATTN_WITH_KVCACHE_FLASH_ATTN_H_ diff --git a/src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.yaml b/src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.yaml new file mode 100644 index 000000000..9e68fe967 --- /dev/null +++ b/src/linked/torch/metax/ops/flash_attn_with_kvcache/flash_attn.yaml @@ -0,0 +1,4 @@ +library: flash_attn +required_symbols: + - >- + mha_fwd_kvcache(at::Tensor&, at::Tensor const&, at::Tensor const&, std::optional&, std::optional&, std::optional&, std::optional&, std::optional&, std::optional&, std::optional&, std::optional&, std::optional&, std::optional&, float, bool, int, int, float, bool, int, std::optional&)