@@ -43,6 +43,36 @@ void GetOutputKVSignal(const paddle::Tensor &x,
4343 int64_t rank_id,
4444 bool wait_flag);
4545
46+ std::vector<paddle::Tensor> BlockAttn (
47+ const paddle::Tensor& qkv,
48+ const paddle::Tensor& key_cache,
49+ const paddle::Tensor& value_cache,
50+ const paddle::Tensor& cum_offsets,
51+ const paddle::Tensor& rotary_embs,
52+ const paddle::Tensor& block_tables,
53+ const paddle::Tensor& prefix_block_tables,
54+ const paddle::Tensor& len_info_cpu,
55+ const paddle::Tensor& encoder_seq_lod_cpu,
56+ const paddle::Tensor& decoder_seq_lod_cpu,
57+ const paddle::Tensor& encoder_kv_lod_cpu,
58+ const paddle::Tensor& encoder_batch_map_cpu,
59+ const paddle::Tensor& decoder_context_len_cpu,
60+ const paddle::Tensor& decoder_context_len_cache_cpu,
61+ const paddle::Tensor& decoder_batch_map_cpu,
62+ const paddle::Tensor& prefix_len_cpu,
63+ const paddle::optional<paddle::Tensor>& k_scales,
64+ const paddle::optional<paddle::Tensor>& v_scales,
65+ const paddle::optional<paddle::Tensor>& k_scales_inv,
66+ const paddle::optional<paddle::Tensor>& v_scales_inv,
67+ const paddle::optional<paddle::Tensor>& k_zeros,
68+ const paddle::optional<paddle::Tensor>& v_zeros,
69+ const paddle::optional<paddle::Tensor>& shift,
70+ const paddle::optional<paddle::Tensor>& smooth,
71+ const paddle::optional<paddle::Tensor>& kv_signal_data_cpu,
72+ const paddle::optional<paddle::Tensor>& cachekv_signal_thread_cpu,
73+ const std::string &pos_emb_type=" NORMAL" ,
74+ bool rope_3d=false );
75+
4676std::vector<paddle::Tensor> MoERedundantTopKSelect (
4777 const paddle::Tensor& gating_logits,
4878 const paddle::Tensor& expert_id_to_ep_rank_array,
@@ -327,6 +357,37 @@ std::vector<paddle::Tensor> SpeculateGetSeqLensOutput(
327357 const paddle::Tensor& seq_lens_decoder);
328358
329359PYBIND11_MODULE (fastdeploy_ops, m) {
360+ m.def (" block_attn" ,
361+ &BlockAttn,
362+ py::arg (" qkv" ),
363+ py::arg (" key_cache" ),
364+ py::arg (" value_cache" ),
365+ py::arg (" cum_offsets" ),
366+ py::arg (" rotary_embs" ),
367+ py::arg (" block_tables" ),
368+ py::arg (" prefix_block_tables" ),
369+ py::arg (" len_info_cpu" ),
370+ py::arg (" encoder_seq_lod_cpu" ),
371+ py::arg (" decoder_seq_lod_cpu" ),
372+ py::arg (" encoder_kv_lod_cpu" ),
373+ py::arg (" encoder_batch_map_cpu" ),
374+ py::arg (" decoder_context_len_cpu" ),
375+ py::arg (" decoder_context_len_cache_cpu" ),
376+ py::arg (" decoder_batch_map_cpu" ),
377+ py::arg (" prefix_len_cpu" ),
378+ py::arg (" k_scales" ),
379+ py::arg (" v_scales" ),
380+ py::arg (" k_scales_inv" ),
381+ py::arg (" v_scales_inv" ),
382+ py::arg (" k_zeros" ),
383+ py::arg (" v_zeros" ),
384+ py::arg (" shift" ),
385+ py::arg (" smooth" ),
386+ py::arg (" kv_signal_data_cpu" ),
387+ py::arg (" cachekv_signal_thread_cpu" ),
388+ py::arg (" pos_emb_type" ) = " NORMAL" ,
389+ py::arg (" rope_3d" ) = false ,
390+ " block attention in XPU" );
330391 m.def (" cuda_host_alloc" ,
331392 &custom_xpu_host_alloc,
332393 " Allocate pinned memory" ,
0 commit comments